albedo-sn97-workspace / code /tests /test_local_train_rl.py
dora7's picture
Albedo SN97 workspace v16: RECORD, merged v16, adapters v13/v15/v16-dpo, packs, harness, evals
2abcc30 verified
Raw
History Blame Contribute Delete
1.12 kB
import torch
from local_train.rl import gold_continuation, group_advantages
def test_gold_continuation_strips_duplicated_think():
prompt = "<|im_start|>assistant\n<think>\n"
completion = "<think>\nedit now\n</think>\n\n```bash\nsed -i 's/a/b/' foo.py\n```\n"
assert gold_continuation(prompt, completion).startswith("edit now")
assert gold_continuation("plain", completion) == completion
def test_group_advantages_prefers_higher_reward_in_pair():
rewards = torch.tensor([0.2, 0.8, 0.1, 0.9])
adv = group_advantages(rewards, num_generations=2)
assert adv[1] > adv[0]
assert adv[3] > adv[2]
def test_group_advantages_falls_back_when_groups_tie():
# Each pair is tied, but the two groups differ — batch baseline should fire.
rewards = torch.tensor([0.1, 0.1, 0.9, 0.9])
adv = group_advantages(rewards, num_generations=2)
assert adv[0] < 0
assert adv[2] > 0
def test_group_advantages_all_equal_is_zero():
rewards = torch.tensor([0.0, 0.0, 0.0, 0.0])
adv = group_advantages(rewards, num_generations=2)
assert torch.allclose(adv, torch.zeros_like(adv))