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:
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user