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
+72
View File
@@ -561,6 +561,78 @@ def test_scheduler_weight_versions_are_monotonic_and_acknowledged(device):
scheduler.stop()
def test_scheduler_applies_weight_mutation_and_version_atomically(device):
scheduler, _tok, model = _make_real_scheduler(device)
before = next(model.parameters()).detach().clone()
def mutate():
with torch.no_grad():
next(model.parameters()).add_(1)
return "updated"
try:
assert scheduler.apply_weight_update(1, mutate) == "updated"
assert scheduler.policy_version == 1
assert not torch.equal(next(model.parameters()), before)
with pytest.raises(ValueError, match="must advance"):
scheduler.apply_weight_update(1, mutate)
def failed_mutation():
raise RuntimeError("optimizer failed")
with pytest.raises(RuntimeError, match="optimizer failed"):
scheduler.apply_weight_update(2, failed_mutation)
assert scheduler.policy_version == 1
finally:
scheduler.stop()
def test_scheduler_serializes_policy_snapshot_and_direct_update(device):
scheduler, _tok, _model = _make_real_scheduler(device)
snapshot_started = threading.Event()
release_snapshot = threading.Event()
update_finished = threading.Event()
errors = []
def inspect(version):
assert version == 0
snapshot_started.set()
assert release_snapshot.wait(timeout=5)
def take_snapshot():
try:
scheduler.with_policy_snapshot(inspect)
except BaseException as exc:
errors.append(exc)
def update():
try:
scheduler.update_weights(1)
update_finished.set()
except BaseException as exc:
errors.append(exc)
snapshot_thread = threading.Thread(target=take_snapshot)
update_thread = threading.Thread(target=update)
try:
snapshot_thread.start()
assert snapshot_started.wait(timeout=5)
update_thread.start()
assert not update_finished.wait(timeout=0.1)
release_snapshot.set()
snapshot_thread.join(timeout=5)
update_thread.join(timeout=5)
assert not snapshot_thread.is_alive()
assert not update_thread.is_alive()
assert errors == []
assert scheduler.policy_version == 1
finally:
release_snapshot.set()
snapshot_thread.join(timeout=5)
update_thread.join(timeout=5)
scheduler.stop()
def test_scheduler_rejects_weight_update_with_queued_tasks(device):
scheduler, _tok, _model = _make_real_scheduler(device)
task_id = scheduler.add_task("queued")
+6 -1
View File
@@ -10,6 +10,7 @@ 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.serialization import Checkpoint
from astrai.trainer.rollout import BaseRewardModel
from astrai.trainer.schedule import SchedulerFactory
from astrai.trainer.trainer import Trainer
@@ -126,6 +127,7 @@ def test_online_rollout_end_to_end(
parallel_mode="none",
strategy_kwargs=strategy_kwargs,
rollout_interval=1,
rollout_max_policy_lag=0,
rollout_temperature=1.0,
rollout_top_k=0,
rollout_top_p=1.0,
@@ -137,5 +139,8 @@ def test_online_rollout_end_to_end(
trainer = Trainer(train_config)
trainer.train(param_path=test_dir)
assert os.path.isdir(os.path.join(test_dir, "ckpt"))
checkpoint_dir = os.path.join(test_dir, "ckpt", "epoch_0_step_2")
assert os.path.isdir(checkpoint_dir)
checkpoint = Checkpoint.load(checkpoint_dir)
assert checkpoint.meta["policy_version"] == 2
assert len(created_reference_models) == 1
+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)))
+219
View File
@@ -1,5 +1,7 @@
"""Unit tests for the online rollout module."""
import threading
import pytest
import torch
@@ -11,6 +13,7 @@ from astrai.trainer.rollout import (
RolloutGenerator,
RolloutResult,
RolloutRunner,
RolloutVersionError,
)
from tests.helpers import FakeTokenizer, make_model
@@ -151,6 +154,113 @@ def test_rollout_generator_uses_eval_and_restores_mode(device):
assert model.training is True
def test_rollout_generator_serializes_generation_and_policy_update(device):
gen, _ = _make_generator(device, group_size=1, max_tokens=2)
generation_started = threading.Event()
allow_generation_to_finish = threading.Event()
update_finished = threading.Event()
thread_errors = []
original = gen._generate_eval
def blocking_generate(batch, generation_version):
generation_started.set()
assert allow_generation_to_finish.wait(timeout=5)
return original(batch, generation_version)
gen._generate_eval = blocking_generate
def generate():
try:
gen.generate(_make_instruction_batch(n=1))
except BaseException as exc:
thread_errors.append(exc)
def apply_update():
try:
gen.apply_weight_update(1, update_finished.set)
except BaseException as exc:
thread_errors.append(exc)
generation_thread = threading.Thread(target=generate)
update_thread = threading.Thread(target=apply_update)
generation_thread.start()
assert generation_started.wait(timeout=5)
update_thread.start()
assert not update_finished.wait(timeout=0.1)
allow_generation_to_finish.set()
generation_thread.join(timeout=5)
update_thread.join(timeout=5)
assert not generation_thread.is_alive()
assert not update_thread.is_alive()
assert thread_errors == []
assert update_finished.is_set()
assert gen.policy_version == 1
def test_rollout_generator_serializes_direct_scheduler_update(device):
gen, _ = _make_generator(device, group_size=1, max_tokens=2)
generation_started = threading.Event()
allow_generation_to_finish = threading.Event()
update_finished = threading.Event()
thread_errors = []
original = gen._generate_eval
def blocking_generate(batch, generation_version):
generation_started.set()
assert allow_generation_to_finish.wait(timeout=5)
return original(batch, generation_version)
gen._generate_eval = blocking_generate
rollout = []
def generate():
try:
rollout.append(gen.generate(_make_instruction_batch(n=1)))
except BaseException as exc:
thread_errors.append(exc)
def update_scheduler_directly():
try:
gen.scheduler.update_weights(1)
update_finished.set()
except BaseException as exc:
thread_errors.append(exc)
generation_thread = threading.Thread(target=generate)
update_thread = threading.Thread(target=update_scheduler_directly)
generation_thread.start()
assert generation_started.wait(timeout=5)
update_thread.start()
assert not update_finished.wait(timeout=0.1)
allow_generation_to_finish.set()
generation_thread.join(timeout=5)
update_thread.join(timeout=5)
assert not generation_thread.is_alive()
assert not update_thread.is_alive()
assert thread_errors == []
assert rollout[0].policy_version == 0
assert gen.policy_version == 1
def test_rollout_generator_keeps_generation_start_version(device):
gen, _ = _make_generator(device, group_size=1, max_tokens=2)
original_run_batch = gen.scheduler.run_batch
def update_after_generation(*args, **kwargs):
result = original_run_batch(*args, **kwargs)
gen.scheduler.update_weights(1)
return result
gen.scheduler.run_batch = update_after_generation
rollout = gen.generate(_make_instruction_batch(n=1))
assert rollout.policy_version == 0
assert gen.policy_version == 1
def test_rollout_generator_mask_matches_responses(device):
"""Positions beyond a response's length are pad (mask False)."""
gen, _ = _make_generator(device, group_size=2, max_tokens=6)
@@ -248,6 +358,7 @@ def _make_runner(device, **kw):
generator=generator,
reward_model=rm,
rollout_interval=kw.get("rollout_interval", 2),
max_policy_lag=kw.get("max_policy_lag"),
),
model,
)
@@ -297,6 +408,114 @@ def test_rollout_runner_tags_generation_version_and_preserves_cached_behavior(de
assert refreshed.policy_version == 1
def test_rollout_runner_rejects_future_generation_version(device):
runner, _ = _make_runner(device, rollout_interval=2)
raw = runner.generator.generate(_make_instruction_batch(n=1))
raw.policy_version = runner.policy_version + 1
runner.generator.generate = lambda _batch: raw
with pytest.raises(RolloutVersionError, match="future policy version"):
runner(_make_instruction_batch(n=1))
def test_rollout_runner_rejects_result_beyond_max_policy_lag(device):
runner, _ = _make_runner(device, rollout_interval=4, max_policy_lag=1)
batch = _make_instruction_batch(n=1)
result, _ = runner(batch)
assert result.policy_version == 0
runner.update_weights(2)
with pytest.raises(RolloutVersionError, match="exceeds max_policy_lag=1"):
runner(batch)
def test_rollout_runner_revalidates_version_after_async_scoring(device):
runner, _ = _make_runner(device, rollout_interval=4, max_policy_lag=0)
original_score = runner._score
def score_while_policy_advances(raw):
result = original_score(raw)
runner.update_weights(1)
return result
runner._score = score_while_policy_advances
with pytest.raises(RolloutVersionError, match="exceeds max_policy_lag=0"):
runner(_make_instruction_batch(n=1))
assert runner._cache is None
def test_rollout_runner_publishes_cache_before_concurrent_policy_update(device):
runner, _ = _make_runner(device, rollout_interval=4, max_policy_lag=1)
final_validation_started = threading.Event()
allow_final_validation_to_finish = threading.Event()
update_finished = threading.Event()
rollout_finished = threading.Event()
thread_errors = []
validation_calls = 0
original_validate = runner._validate_policy_version
def blocking_validate(result, *, live_version=None):
nonlocal validation_calls
validation_calls += 1
original_validate(result, live_version=live_version)
if validation_calls == 2:
final_validation_started.set()
assert allow_final_validation_to_finish.wait(timeout=5)
runner._validate_policy_version = blocking_validate
def produce_rollout():
try:
runner(_make_instruction_batch(n=1))
rollout_finished.set()
except BaseException as exc:
thread_errors.append(exc)
def apply_update():
try:
runner.apply_weight_update(1, update_finished.set)
except BaseException as exc:
thread_errors.append(exc)
rollout_thread = threading.Thread(target=produce_rollout)
update_thread = threading.Thread(target=apply_update)
rollout_thread.start()
assert final_validation_started.wait(timeout=5)
update_thread.start()
assert not update_finished.wait(timeout=0.1)
allow_final_validation_to_finish.set()
rollout_thread.join(timeout=5)
update_thread.join(timeout=5)
assert not rollout_thread.is_alive()
assert not update_thread.is_alive()
assert thread_errors == []
assert rollout_finished.is_set()
assert update_finished.is_set()
assert runner._cache is not None
assert runner._cache.policy_version == 0
assert runner.policy_version == 1
def test_rollout_runner_derives_default_policy_lag_from_interval(device):
runner, _ = _make_runner(device, rollout_interval=4)
assert runner.max_policy_lag == 3
@pytest.mark.parametrize(
("kwargs", "message"),
[
({"rollout_interval": 0}, "rollout_interval must be positive"),
({"max_policy_lag": -1}, "max_policy_lag must be non-negative"),
],
)
def test_rollout_runner_rejects_invalid_version_window(device, kwargs, message):
generator, _ = _make_generator(device)
with pytest.raises(ValueError, match=message):
RolloutRunner(generator, ConstantRewardModel(), **kwargs)
def test_rollout_runner_refreshes_for_different_batch(device):
runner, _ = _make_runner(device, rollout_interval=100)
r1, fresh1 = runner(_make_instruction_batch(n=1))