fix: correct online rollout lifecycle
This commit is contained in:
@@ -937,8 +937,9 @@ def test_grpo_collate_variable_lengths():
|
||||
assert result["masks"].shape == (2, 2, 4)
|
||||
assert result["rewards"].shape == (2, 2)
|
||||
|
||||
# Check padding: item 1 prompt is length 2, padded to 3
|
||||
assert result["prompts"][1, 2] == 0
|
||||
# Prompts are left-padded so each response follows its real prompt tokens.
|
||||
assert torch.equal(result["prompts"][1], torch.tensor([0, 10, 11]))
|
||||
assert torch.equal(result["prompt_mask"][1], torch.tensor([False, True, True]))
|
||||
|
||||
# Check response content: item 0, response 0 is [4,5] padded to 4
|
||||
assert result["responses"][0, 0, 0] == 4
|
||||
|
||||
@@ -63,6 +63,7 @@ def _make_frozen(model, device):
|
||||
def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"):
|
||||
return RolloutResult(
|
||||
prompts=torch.randint(3, 200, (B, P), device=device),
|
||||
prompt_mask=torch.ones(B, P, dtype=torch.bool, device=device),
|
||||
responses=torch.randint(3, 200, (B, G, R), device=device),
|
||||
response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device),
|
||||
rewards=torch.randn(B, G, device=device),
|
||||
@@ -177,6 +178,7 @@ def test_grpo_prepare_from_rollout_mapping(device):
|
||||
r = _make_rollout_result(device=device)
|
||||
batch = strat.prepare_from_rollout(r)
|
||||
assert batch["prompts"] is r.prompts
|
||||
assert batch["prompt_mask"] is r.prompt_mask
|
||||
assert batch["responses"] is r.responses
|
||||
assert batch["masks"] is r.response_mask
|
||||
assert batch["rewards"] is r.rewards
|
||||
@@ -255,9 +257,11 @@ def test_grpo_no_resync_when_same_cached_result(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.calls == 2
|
||||
assert runner.step_calls == 1
|
||||
assert runner.step_calls == 2
|
||||
|
||||
|
||||
def test_grpo_resync_when_new_rollout_result(device):
|
||||
@@ -265,8 +269,10 @@ def test_grpo_resync_when_new_rollout_result(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
runner.swap_result(_make_rollout_result(device=device))
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.calls == 2
|
||||
assert runner.step_calls == 2
|
||||
|
||||
@@ -281,8 +287,10 @@ def test_dpo_no_sync_hook_when_new_rollout_result(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
runner.swap_result(_make_rollout_result(device=device))
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.step_calls == 2
|
||||
|
||||
|
||||
@@ -301,6 +309,7 @@ def test_step_called_when_sync_gradients_true(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.step_calls == 1
|
||||
|
||||
|
||||
|
||||
@@ -71,6 +71,18 @@ class ConstantRewardModel(BaseRewardModel):
|
||||
return torch.full((B, G), float(self.value))
|
||||
|
||||
|
||||
class BadShapeRewardModel(BaseRewardModel):
|
||||
def score(self, prompts, responses):
|
||||
return torch.zeros(len(prompts))
|
||||
|
||||
|
||||
class NonFiniteRewardModel(BaseRewardModel):
|
||||
def score(self, prompts, responses):
|
||||
B = len(prompts)
|
||||
G = len(responses[0]) if B else 0
|
||||
return torch.full((B, G), float("nan"))
|
||||
|
||||
|
||||
def _make_config(vocab_size=200, max_position_embeddings=128):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
@@ -111,6 +123,7 @@ def _make_instruction_batch(n=2):
|
||||
def test_raw_rollout_fields():
|
||||
r = RawRollout(
|
||||
prompts=torch.zeros(2, 4, dtype=torch.long),
|
||||
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
|
||||
responses=torch.zeros(2, 3, 5, dtype=torch.long),
|
||||
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
|
||||
logprobs_old=torch.zeros(2, 3, 5),
|
||||
@@ -124,6 +137,7 @@ def test_raw_rollout_fields():
|
||||
def test_rollout_result_inherits_raw_rollout_fields():
|
||||
r = RolloutResult(
|
||||
prompts=torch.zeros(2, 4, dtype=torch.long),
|
||||
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
|
||||
responses=torch.zeros(2, 3, 5, dtype=torch.long),
|
||||
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
|
||||
logprobs_old=torch.zeros(2, 3, 5),
|
||||
@@ -132,6 +146,7 @@ def test_rollout_result_inherits_raw_rollout_fields():
|
||||
assert r.rewards.shape == (2, 3)
|
||||
assert r.prompts.shape == (2, 4)
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.prompt_mask.shape == (2, 4)
|
||||
|
||||
|
||||
def test_base_reward_model_is_abstract():
|
||||
@@ -179,11 +194,28 @@ def test_rollout_generator_shapes(device):
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.response_mask.shape == (2, 3, 5)
|
||||
assert r.logprobs_old.shape == (2, 3, 5)
|
||||
assert r.prompt_mask.shape == r.prompts.shape
|
||||
assert len(r.prompt_texts) == 2
|
||||
assert len(r.response_texts) == 2
|
||||
assert len(r.response_texts[0]) == 3
|
||||
|
||||
|
||||
def test_rollout_generator_uses_eval_and_restores_mode(device):
|
||||
gen, model = _make_generator(device, group_size=1, max_tokens=2)
|
||||
model.train()
|
||||
seen_training = []
|
||||
original = gen.scheduler.run_batch
|
||||
|
||||
def recording_run_batch(*args, **kwargs):
|
||||
seen_training.append(model.training)
|
||||
return original(*args, **kwargs)
|
||||
|
||||
gen.scheduler.run_batch = recording_run_batch
|
||||
gen.generate(_make_instruction_batch(n=1))
|
||||
assert seen_training == [False]
|
||||
assert model.training is True
|
||||
|
||||
|
||||
def test_rollout_generator_mask_matches_responses(device):
|
||||
"""Positions beyond a response's length are pad (mask False)."""
|
||||
gen, _ = _make_generator(device, group_size=2, max_tokens=6)
|
||||
@@ -291,6 +323,24 @@ def test_rollout_runner_cache_returns_stale_flag(device):
|
||||
assert fresh2 is False
|
||||
|
||||
|
||||
def test_rollout_runner_refreshes_for_different_batch(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=100)
|
||||
r1, fresh1 = runner(_make_instruction_batch(n=1))
|
||||
batch2 = {"instruction": ["Different prompt"], "input": [""]}
|
||||
r2, fresh2 = runner(batch2)
|
||||
assert fresh1 is True
|
||||
assert fresh2 is True
|
||||
assert r2 is not r1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("reward_model", [BadShapeRewardModel, NonFiniteRewardModel])
|
||||
def test_rollout_runner_rejects_invalid_rewards(device, reward_model):
|
||||
generator, _ = _make_generator(device, group_size=2, max_tokens=2)
|
||||
runner = RolloutRunner(generator, reward_model(), rollout_interval=1)
|
||||
with pytest.raises(ValueError):
|
||||
runner(_make_instruction_batch(n=1))
|
||||
|
||||
|
||||
def test_rollout_runner_step_triggers_new_rollout(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=2)
|
||||
batch = _make_instruction_batch()
|
||||
|
||||
Reference in New Issue
Block a user