From 90c8397e29ced72d6b209617c26995de6453b7be Mon Sep 17 00:00:00 2001 From: Daoyuan Li <94409450+DaoyuanLi2816@users.noreply.github.com> Date: Tue, 25 Aug 2026 20:30:54 -0700 Subject: [PATCH] Fix unpairing iterable dataset dictionaries --- tests/test_data_utils.py | 11 ++++++++++- trl/data_utils.py | 2 +- 2 files changed, 11 insertions(+), 2 deletions(-) diff --git a/tests/test_data_utils.py b/tests/test_data_utils.py index c57faa727da..36f20f2ac48 100644 --- a/tests/test_data_utils.py +++ b/tests/test_data_utils.py @@ -18,7 +18,7 @@ import pytest import transformers -from datasets import Dataset, DatasetDict +from datasets import Dataset, DatasetDict, IterableDatasetDict from packaging.version import Version from transformers import AutoProcessor, AutoTokenizer, is_vision_available @@ -1090,6 +1090,15 @@ def test_unpair_preference_dataset_dict(self): "The paired dataset should be converted to unpaired." ) + def test_unpair_preference_iterable_dataset_dict(self): + # Test that a paired iterable dataset dict is correctly converted to unpaired + paired_dataset_dict = IterableDatasetDict({"abc": self.paired_dataset.to_iterable_dataset()}) + unpaired_dataset_dict = unpair_preference_dataset(paired_dataset_dict) + assert list(unpaired_dataset_dict["abc"]) == [ + dict(zip(self.unpaired_dataset.column_names, vals, strict=False)) + for vals in zip(*self.unpaired_dataset.to_dict().values(), strict=False) + ] + def test_maybe_unpair_preference_dataset(self): # Test that a paired dataset is correctly converted to unpaired with maybe_unpair_preference_dataset unpaired_dataset = maybe_unpair_preference_dataset(self.paired_dataset) diff --git a/trl/data_utils.py b/trl/data_utils.py index 671d78f7d64..b05bdf70af5 100644 --- a/trl/data_utils.py +++ b/trl/data_utils.py @@ -494,7 +494,7 @@ def unpair_preference_dataset( {'prompt': 'The sky is', 'completion': ' blue.', 'label': True} ``` """ - if isinstance(dataset, DatasetDict): + if isinstance(dataset, (DatasetDict, IterableDatasetDict)): column_names = next(iter(dataset.values())).column_names elif isinstance(dataset, Dataset): column_names = dataset.column_names