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