fix: keep async rollouts version-consistent

- serialize shared-model optimizer updates with generation
- reject future or over-lagged rollout results after asynchronous scoring
- close cache publication races
- persist policy versions in online checkpoints
This commit is contained in:
0z5a
2026-09-03 12:01:24 +08:00
parent ce2f9d13b3
commit 587b0ee046
16 changed files with 531 additions and 49 deletions
+45 -7
View File
@@ -60,11 +60,25 @@ class _RecordingRunner:
self.weight_updates.append(policy_version)
return policy_version
def apply_weight_update(self, policy_version, update):
result = update()
self.update_weights(policy_version)
return result
def swap_result(self, result):
self.result = result
self._fresh = True
class _NoOpOptimizer:
def step(self):
return None
def _step(strat):
strat.optimizer_step(_NoOpOptimizer())
def _make_grpo(device, executor=None):
model, _ = make_model(device)
ref_model = make_frozen(model, device)
@@ -250,9 +264,9 @@ def test_grpo_reuses_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()
_step(strat)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
_step(strat)
assert runner.calls == 2
assert runner.step_calls == 2
@@ -262,10 +276,10 @@ def test_grpo_accepts_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()
_step(strat)
runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
_step(strat)
assert runner.calls == 2
assert runner.step_calls == 2
@@ -280,10 +294,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()
_step(strat)
runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
_step(strat)
assert runner.step_calls == 2
@@ -302,12 +316,36 @@ 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()
_step(strat)
assert runner.step_calls == 1
assert runner.weight_updates == [1]
assert strat.policy_version == 1
def test_post_hoc_online_optimizer_step_is_rejected(device):
strat = _make_grpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
with pytest.raises(RuntimeError, match="strategy.optimizer_step"):
strat.on_optimizer_step()
def test_optimizer_step_publishes_version_with_weight_update(device):
strat = _make_grpo(device)
runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner)
parameter = next(strat.model.parameters())
parameter.grad = torch.ones_like(parameter)
optimizer = torch.optim.SGD(strat.model.parameters(), lr=0.1)
before = parameter.detach().clone()
strat.optimizer_step(optimizer)
assert not torch.equal(parameter, before)
assert runner.weight_updates == [1]
assert runner.step_calls == 1
def test_loss_is_differentiable_dpo(device):
strat = _make_dpo(device)
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))