Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion dspy/teleprompt/grpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
47 changes: 47 additions & 0 deletions tests/teleprompt/test_grpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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!")