perf: reuse rollout behavior logprobs

Feed sampler-aligned behavior log-probabilities directly into online GRPO instead of allocating, synchronizing, and forwarding a duplicate old-policy model. Keep the old-model path as an offline compatibility fallback and validate supplied rollout tensors before loss computation.
This commit is contained in:
0z5a
2026-09-02 19:29:53 +08:00
committed by ViperEkura
parent e58a728b80
commit 4019ddac31
7 changed files with 115 additions and 44 deletions
+14 -19
View File
@@ -67,12 +67,11 @@ class _RecordingRunner:
def _make_grpo(device, executor=None):
model, _ = make_model(device)
old_model = make_frozen(model, device)
ref_model = make_frozen(model, device)
return GRPOStrategy(
model=model,
device=device,
old_model=old_model,
old_model=None,
ref_model=ref_model,
clip_eps=0.2,
kl_coef=0.01,
@@ -141,6 +140,7 @@ def test_grpo_prepare_from_rollout_mapping(device):
assert batch["responses"] is r.responses
assert batch["masks"] is r.response_mask
assert batch["rewards"] is r.rewards
assert batch["logprobs_old"] is r.logprobs_old
def test_dpo_prepare_from_rollout_conditions_responses_on_prompt(device):
@@ -197,13 +197,14 @@ def test_dpo_prepare_from_rollout_same_response_keeps_distinct_prompts():
assert not batch["rejected_mask"][:, :3].any()
def test_call_without_runner_falls_back_to_compute_loss_grpo(device):
def test_call_without_runner_accepts_behavior_logprobs_grpo(device):
strat = _make_grpo(device)
batch = {
"prompts": torch.randint(3, 200, (2, 4), device=device),
"responses": torch.randint(3, 200, (2, 4, 6), device=device),
"masks": torch.ones(2, 4, 6, device=device),
"rewards": torch.randn(2, 4, device=device),
"logprobs_old": torch.zeros(2, 4, 6, device=device),
}
loss = strat(batch)["loss"]
assert torch.isfinite(loss).item()
@@ -232,25 +233,19 @@ def test_call_invokes_runner_each_time(device):
assert runner.calls == 2
def test_grpo_syncs_old_model_on_first_rollout(device):
def test_grpo_reuses_rollout_logprobs_without_old_model(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
result = _make_rollout_result(device=device)
result.logprobs_old.normal_().requires_grad_()
runner = _RecordingRunner(result)
strat.set_rollout_runner(runner)
with torch.no_grad():
for p in strat.model.parameters():
p.add_(0.1)
old_before = {k: v.clone() for k, v in strat.old_model.state_dict().items()}
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
old_after = strat.old_model.state_dict()
synced = any(
not torch.allclose(old_before[k], old_after[k])
for k in old_before
if k in old_after
)
assert synced
assert strat.old_model is None
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})["loss"]
loss.backward()
assert result.logprobs_old.grad is None
def test_grpo_no_resync_when_same_cached_result(device):
def test_grpo_reuses_same_cached_result(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
@@ -262,7 +257,7 @@ def test_grpo_no_resync_when_same_cached_result(device):
assert runner.step_calls == 2
def test_grpo_resync_when_new_rollout_result(device):
def test_grpo_accepts_new_rollout_result(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)