diff --git a/dspy/teleprompt/grpo.py b/dspy/teleprompt/grpo.py index 7998d2608d..9a8165df4a 100644 --- a/dspy/teleprompt/grpo.py +++ b/dspy/teleprompt/grpo.py @@ -201,7 +201,7 @@ def update_shuffled_trainset(self, original_trainset): for id in self.shuffled_trainset_ids: self.id_freqs[id] += 1 - num_to_pad = self.num_dspy_examples_per_grpo_step - (len(original_trainset) % self.num_dspy_examples_per_grpo_step) + num_to_pad = (self.num_dspy_examples_per_grpo_step - (len(original_trainset) % self.num_dspy_examples_per_grpo_step)) % self.num_dspy_examples_per_grpo_step if num_to_pad > 0: # Select ids based on least frequent ids for _ in range(num_to_pad): diff --git a/tests/teleprompt/test_grpo.py b/tests/teleprompt/test_grpo.py index 0c2e61b509..4af4615118 100644 --- a/tests/teleprompt/test_grpo.py +++ b/tests/teleprompt/test_grpo.py @@ -61,8 +61,55 @@ def test_grpo_dataset_shuffler_with_num_ex_per_step_greater_dataset(): assert counter[i] == 10 +def test_grpo_dataset_shuffler_no_padding_when_divisible(): + dataset = [1, 2, 3, 4, 5, 6] + grpo = GRPO( + num_dspy_examples_per_grpo_step=3, + exclude_demos=True, + ) + + grpo.select_training_sample_and_update_shuffled_trainset(dataset, 0) + + assert len(grpo.shuffled_trainset_ids) == len(dataset) + assert len(grpo.shuffled_trainset_ids) % grpo.num_dspy_examples_per_grpo_step == 0 + assert grpo.id_freqs.total() == len(dataset) + assert all(v == 1 for v in grpo.id_freqs.values()) + + +def test_grpo_dataset_shuffler_across_epoch_boundary_divisible(): + from collections import Counter + + dataset = [1, 2, 3, 4, 5, 6] + grpo = GRPO( + num_dspy_examples_per_grpo_step=3, + exclude_demos=True, + ) + + batches, epochs = [], [] + for i in range(6): + batch = grpo.select_training_sample_and_update_shuffled_trainset(dataset, i) + batches.append(batch) + epochs.append(grpo.epoch) + + assert epochs == [0, 0, 1, 1, 2, 2] + assert len(grpo.shuffled_trainset_ids) == len(dataset) + + epoch_to_ids = {} + for batch, epoch in zip(batches, epochs, strict=True): + epoch_to_ids.setdefault(epoch, []).extend(batch) + + for epoch, ids in epoch_to_ids.items(): + assert len(ids) == len(set(ids)), f"epoch {epoch} contains duplicate ids: {Counter(ids)}" + assert set(ids) == set(dataset), f"epoch {epoch} contains unexpected ids: {set(ids)}" + + assert Counter(epoch_to_ids[0]) == Counter(dataset) + assert Counter(epoch_to_ids[1]) == Counter(dataset) + + if __name__ == "__main__": test_grpo_dataset_shuffler() test_grpo_dataset_shuffler_with_num_ex_per_step_less_dataset() test_grpo_dataset_shuffler_with_num_ex_per_step_greater_dataset() + test_grpo_dataset_shuffler_no_padding_when_divisible() + test_grpo_dataset_shuffler_across_epoch_boundary_divisible() print("All tests passed!")