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