diff --git a/gliclass/data_processing.py b/gliclass/data_processing.py index f9bfc12..ada835f 100644 --- a/gliclass/data_processing.py +++ b/gliclass/data_processing.py @@ -486,7 +486,9 @@ def __call__(self, batch): for key in keys: key_data = [item[key] for item in batch] if isinstance(key_data[0], torch.Tensor): - if key_data[0].dim() == 1: + if key_data[0].dim() == 0: + padded_batch[key] = torch.stack(key_data) + elif key_data[0].dim() == 1: padded_batch[key] = pad_sequence(key_data, batch_first=True) elif key_data[0].dim() == 2: padded_batch[key] = pad_2d_tensor(key_data) diff --git a/tests/test_data_processing.py b/tests/test_data_processing.py index d71d7b0..4e2640d 100644 --- a/tests/test_data_processing.py +++ b/tests/test_data_processing.py @@ -3,7 +3,31 @@ import pytest import torch -from gliclass.data_processing import pad_2d_tensor +from gliclass.data_processing import DataCollatorWithPadding, pad_2d_tensor + + +def test_collator_stacks_scalar_labels_for_single_label_classification(): + collator = DataCollatorWithPadding(device="cpu") + batch = [ + { + "input_ids": torch.tensor([1, 2]), + "attention_mask": torch.tensor([1, 1]), + "labels": torch.tensor(0), + "labels_text": ["first"], + }, + { + "input_ids": torch.tensor([3]), + "attention_mask": torch.tensor([1]), + "labels": torch.tensor(1), + "labels_text": ["second"], + }, + ] + + result = collator(batch) + + assert torch.equal(result["labels"], torch.tensor([0, 1])) + assert result["labels"].shape == (2,) + assert result["max_num_classes"] == 1 class TestPad2DTensor: