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
+38
View File
@@ -71,6 +71,44 @@ def test_grpo_loss_backward(grpo_strategy):
assert has_grad
def test_grpo_reuses_supplied_behavior_logprobs(grpo_strategy):
"""A rollout batch must not forward the old policy again."""
strategy, device = grpo_strategy
class _FailingOldPolicy(torch.nn.Module):
def forward(self, *args, **kwargs):
raise AssertionError("old policy forward should not run")
strategy.old_model = _FailingOldPolicy()
batch = _make_batch(device=device)
batch["logprobs_old"] = torch.zeros_like(batch["responses"], dtype=torch.float)
loss = strategy.compute_loss(batch)
assert torch.isfinite(loss).item()
def test_grpo_requires_behavior_source(grpo_strategy):
strategy, device = grpo_strategy
strategy.old_model = None
with pytest.raises(ValueError, match="must provide logprobs_old"):
strategy.compute_loss(_make_batch(device=device))
@pytest.mark.parametrize("invalid", ["shape", "nonfinite"])
def test_grpo_rejects_invalid_behavior_logprobs(grpo_strategy, invalid):
strategy, device = grpo_strategy
batch = _make_batch(device=device)
if invalid == "shape":
batch["logprobs_old"] = torch.zeros(1, device=device)
match = "shape must match responses"
else:
batch["logprobs_old"] = torch.zeros_like(batch["responses"], dtype=torch.float)
batch["logprobs_old"][0, 0, 0] = float("nan")
match = "only finite values"
with pytest.raises(ValueError, match=match):
strategy.compute_loss(batch)
@pytest.mark.parametrize("model_name", ["ref_model", "old_model"])
def test_grpo_frozen_models_not_updated(grpo_strategy, model_name):
"""Backward should not populate gradients on ref_model or old_model."""
+14 -1
View File
@@ -7,6 +7,7 @@ import pytest
import torch
from torch.utils.data import Dataset
import astrai.trainer.train_context as train_context
from astrai.config import TrainConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.trainer.rollout import BaseRewardModel
@@ -87,8 +88,19 @@ _ONLINE_STRATEGIES = [
@pytest.mark.integration
@pytest.mark.parametrize(("strategy", "strategy_kwargs"), _ONLINE_STRATEGIES)
def test_online_rollout_end_to_end(base_test_env, strategy, strategy_kwargs):
def test_online_rollout_end_to_end(
base_test_env, strategy, strategy_kwargs, monkeypatch
):
"""Run one epoch of online RL rollout with KV-cache-backed generation."""
created_reference_models = []
create_ref_model = train_context.create_ref_model
def track_reference_model(*args, **kwargs):
created_reference_models.append(strategy)
return create_ref_model(*args, **kwargs)
monkeypatch.setattr(train_context, "create_ref_model", track_reference_model)
test_dir = base_test_env["test_dir"]
device = base_test_env["device"]
tokenizer = base_test_env["tokenizer"]
@@ -126,3 +138,4 @@ def test_online_rollout_end_to_end(base_test_env, strategy, strategy_kwargs):
trainer.train(param_path=test_dir)
assert os.path.isdir(os.path.join(test_dir, "ckpt"))
assert len(created_reference_models) == 1
+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)