Skip to content

Commit 74dfedf

Browse files
committed
fix: unpair IterableDatasetDict inputs
1 parent 6d484ba commit 74dfedf

2 files changed

Lines changed: 34 additions & 3 deletions

File tree

tests/test_data_utils.py

Lines changed: 31 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818

1919
import pytest
2020
import transformers
21-
from datasets import Dataset, DatasetDict
21+
from datasets import Dataset, DatasetDict, IterableDataset, IterableDatasetDict
2222
from packaging.version import Version
2323
from transformers import AutoProcessor, AutoTokenizer, is_vision_available
2424

@@ -1082,6 +1082,36 @@ def test_unpair_preference_dataset_iterable_extra_columns(self):
10821082
dict(zip(expected.keys(), vals, strict=False)) for vals in zip(*expected.values(), strict=False)
10831083
]
10841084

1085+
def test_unpair_preference_dataset_iterable_dict(self):
1086+
# Test that an IterableDatasetDict is correctly unpaired
1087+
paired_dataset_dict = IterableDatasetDict({"abc": self.paired_dataset.to_iterable_dataset()})
1088+
unpaired_dataset_dict = unpair_preference_dataset(paired_dataset_dict)
1089+
assert list(unpaired_dataset_dict["abc"]) == [
1090+
dict(zip(self.unpaired_dataset.column_names, vals, strict=False))
1091+
for vals in zip(*self.unpaired_dataset.to_dict().values(), strict=False)
1092+
]
1093+
1094+
def test_unpair_preference_dataset_iterable_dict_without_column_names(self):
1095+
# Test that an IterableDatasetDict without schema metadata is correctly unpaired
1096+
def generate_examples():
1097+
for prompt, chosen, rejected in zip(
1098+
self.paired_dataset["prompt"],
1099+
self.paired_dataset["chosen"],
1100+
self.paired_dataset["rejected"],
1101+
strict=True,
1102+
):
1103+
yield {"prompt": prompt, "chosen": chosen, "rejected": rejected}
1104+
1105+
paired_dataset_dict = IterableDatasetDict({"abc": IterableDataset.from_generator(generate_examples)})
1106+
assert paired_dataset_dict["abc"].column_names is None
1107+
1108+
unpaired_dataset_dict = unpair_preference_dataset(paired_dataset_dict)
1109+
1110+
assert list(unpaired_dataset_dict["abc"]) == [
1111+
dict(zip(self.unpaired_dataset.column_names, vals, strict=False))
1112+
for vals in zip(*self.unpaired_dataset.to_dict().values(), strict=False)
1113+
]
1114+
10851115
def test_unpair_preference_dataset_dict(self):
10861116
# Test that a paired dataset dict is correctly converted to unpaired
10871117
paired_dataset_dict = DatasetDict({"abc": self.paired_dataset})

trl/data_utils.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -494,8 +494,9 @@ def unpair_preference_dataset(
494494
{'prompt': 'The sky is', 'completion': ' blue.', 'label': True}
495495
```
496496
"""
497-
if isinstance(dataset, DatasetDict):
498-
column_names = next(iter(dataset.values())).column_names
497+
if isinstance(dataset, (DatasetDict, IterableDatasetDict)):
498+
first_split = next(iter(dataset.values()))
499+
column_names = first_split.column_names or list(next(iter(first_split)).keys())
499500
elif isinstance(dataset, Dataset):
500501
column_names = dataset.column_names
501502
else: # IterableDataset

0 commit comments

Comments
 (0)