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
+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