feat: version rollout weight updates

Track a monotonic policy version across optimizer steps, scheduler updates, and rollout results. Serialize synchronous generation with weight acknowledgements and invalidate reusable prefix KV entries so cached samples remain attributable to the behavior policy that generated them.
This commit is contained in:
0z5a
2026-09-02 19:01:41 +08:00
parent 1fad50d847
commit e58a728b80
13 changed files with 232 additions and 7 deletions
+8
View File
@@ -338,6 +338,14 @@ class TaskCacheManager:
if state is not None: if state is not None:
self._strategy.record_hashes(state, prompt_ids, start_logical_page) self._strategy.record_hashes(state, prompt_ids, start_logical_page)
def invalidate_cache(self) -> int:
"""Drop reusable KV entries once all task-owned entries are released."""
if self._states:
raise RuntimeError("Cannot invalidate KV cache while tasks are active")
self._bind_state = None
self._bind_was_steady = False
return self._strategy.invalidate_cache()
@staticmethod @staticmethod
def task_cacheable_ids(task_id: str, prompt_ids: List[int], output_ids: List[int]): def task_cacheable_ids(task_id: str, prompt_ids: List[int], output_ids: List[int]):
return list(prompt_ids) + list(output_ids[:-1]) return list(prompt_ids) + list(output_ids[:-1])
+22
View File
@@ -89,6 +89,19 @@ class Allocator:
if idx in self._lru: if idx in self._lru:
self._lru.move_to_end(idx) self._lru.move_to_end(idx)
def clear_cached(self) -> int:
"""Release every unreferenced LRU page back to the free pool."""
with self._lock:
cached = list(self._lru)
self._lru.clear()
for idx in cached:
if self._refs[idx] != 0:
raise RuntimeError("Cannot invalidate a referenced cache page")
if self.on_evict:
self.on_evict(idx)
self._free_mask |= 1 << idx
return len(cached)
class RadixNode: class RadixNode:
"""A page-aligned edge in the CPU-side prefix radix trie.""" """A page-aligned edge in the CPU-side prefix radix trie."""
@@ -200,6 +213,10 @@ class AllocationStrategy(ABC):
start: int, start: int,
) -> None: ... ) -> None: ...
def invalidate_cache(self) -> int:
"""Drop reusable KV entries after an inference weight update."""
return 0
class ContiguousStrategy(AllocationStrategy): class ContiguousStrategy(AllocationStrategy):
"""Static contiguous allocation: slots are pre-assigned at pool init. """Static contiguous allocation: slots are pre-assigned at pool init.
@@ -316,3 +333,8 @@ class PagedStrategy(AllocationStrategy):
full = len(prompt_ids) // self._page_size full = len(prompt_ids) // self._page_size
for i in range(start, min(full, len(state.pages))): for i in range(start, min(full, len(state.pages))):
self._prefix.record(state.pages[i], prompt_ids, i) self._prefix.record(state.pages[i], prompt_ids, i)
def invalidate_cache(self) -> int:
if self._prefix is None:
return 0
return self._alloc.clear_cached()
+57
View File
@@ -2,6 +2,7 @@ import logging
import threading import threading
import uuid import uuid
from contextlib import nullcontext from contextlib import nullcontext
from functools import wraps
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any, Dict, List, Optional, Tuple, Union
import torch import torch
@@ -28,6 +29,15 @@ from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def _with_weight_lock(method):
@wraps(method)
def synchronized(self, *args, **kwargs):
with self._weight_lock:
return method(self, *args, **kwargs)
return synchronized
class InferenceScheduler: class InferenceScheduler:
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups).""" """Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
@@ -42,7 +52,14 @@ class InferenceScheduler:
cache: Optional[PagePool] = None, cache: Optional[PagePool] = None,
enable_cuda_graph: bool = True, enable_cuda_graph: bool = True,
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None, backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
policy_version: int = 0,
): ):
if (
isinstance(policy_version, bool)
or not isinstance(policy_version, int)
or policy_version < 0
):
raise ValueError("policy_version must be a non-negative integer")
config = model.config config = model.config
if max_seq_len is not None: if max_seq_len is not None:
@@ -103,6 +120,44 @@ class InferenceScheduler:
self._stop_event = threading.Event() self._stop_event = threading.Event()
self._loop_thread: Optional[threading.Thread] = None self._loop_thread: Optional[threading.Thread] = None
self._weight_lock = threading.RLock()
self._policy_version = policy_version
@property
def policy_version(self) -> int:
"""Version of the model weights used for subsequent generations."""
return self._policy_version
@_with_weight_lock
def update_weights(self, policy_version: int) -> int:
"""Acknowledge an in-place weight update and invalidate stale KV state.
The scheduler owns the same model object as the in-process trainer, so
weights have already changed when this method is called. The explicit
version update makes that lifecycle visible and prevents prefix KV
entries produced by older weights from being reused.
"""
if (
isinstance(policy_version, bool)
or not isinstance(policy_version, int)
or policy_version < 0
):
raise ValueError("policy_version must be a non-negative integer")
if policy_version < self._policy_version:
raise ValueError(
f"policy_version cannot move backwards from "
f"{self._policy_version} to {policy_version}"
)
if policy_version == self._policy_version:
return self._policy_version
if self._loop_thread is not None and self._loop_thread.is_alive():
raise RuntimeError("Stop the scheduler before updating model weights")
if self._task_mgr.get_active_tasks() or self._task_mgr.get_waiting_tasks():
raise RuntimeError("Cannot update model weights while tasks are queued")
self._task_cache.invalidate_cache()
self._policy_version = policy_version
return self._policy_version
def add_task(self, prompt: str, **kwargs) -> str: def add_task(self, prompt: str, **kwargs) -> str:
return self._task_mgr.add_task(prompt, **kwargs) return self._task_mgr.add_task(prompt, **kwargs)
@@ -125,6 +180,7 @@ class InferenceScheduler:
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
stats = self._task_mgr.get_stats() stats = self._task_mgr.get_stats()
stats["kv_cache_tasks"] = self._task_cache.task_count stats["kv_cache_tasks"] = self._task_cache.task_count
stats["policy_version"] = self._policy_version
return stats return stats
@property @property
@@ -343,6 +399,7 @@ class InferenceScheduler:
) )
self._task_mgr.clear_queues() self._task_mgr.clear_queues()
@_with_weight_lock
def run_batch( def run_batch(
self, self,
prompt_ids_list: List[List[int]], prompt_ids_list: List[List[int]],
+30 -7
View File
@@ -13,6 +13,7 @@ Provides:
so callers do not need to rely on object identity to detect refreshes. so callers do not need to rely on object identity to detect refreshes.
""" """
import threading
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple from typing import Dict, List, Optional, Tuple
@@ -53,6 +54,7 @@ class RawRollout:
responses: Tensor responses: Tensor
response_mask: Tensor response_mask: Tensor
logprobs_old: Tensor logprobs_old: Tensor
policy_version: int = 0
prompt_texts: List[str] = field(default_factory=list) prompt_texts: List[str] = field(default_factory=list)
response_texts: List[List[str]] = field(default_factory=list) response_texts: List[List[str]] = field(default_factory=list)
@@ -129,6 +131,16 @@ class RolloutGenerator:
self.top_p = top_p self.top_p = top_p
self.frequency_penalty = frequency_penalty self.frequency_penalty = frequency_penalty
self.rep_window = rep_window self.rep_window = rep_window
self._weight_lock = threading.RLock()
@property
def policy_version(self) -> int:
return self.scheduler.policy_version
def update_weights(self, policy_version: int) -> int:
"""Acknowledge shared-model weights and invalidate older scheduler KV."""
with self._weight_lock:
return self.scheduler.update_weights(policy_version)
@torch.no_grad() @torch.no_grad()
def generate(self, batch: Dict) -> RawRollout: def generate(self, batch: Dict) -> RawRollout:
@@ -146,13 +158,14 @@ class RolloutGenerator:
``add_generation_prompt=True`` so rollout prompts match the ``add_generation_prompt=True`` so rollout prompts match the
format the policy was SFT-trained on. format the policy was SFT-trained on.
""" """
model = self.scheduler._executor.model with self._weight_lock:
was_training = model.training model = self.scheduler._executor.model
model.eval() was_training = model.training
try: model.eval()
return self._generate_eval(batch) try:
finally: return self._generate_eval(batch)
model.train(was_training) finally:
model.train(was_training)
def _generate_eval(self, batch: Dict) -> RawRollout: def _generate_eval(self, batch: Dict) -> RawRollout:
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch) prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
@@ -245,6 +258,7 @@ class RolloutGenerator:
responses=responses, responses=responses,
response_mask=response_mask, response_mask=response_mask,
logprobs_old=logprobs_old, logprobs_old=logprobs_old,
policy_version=self.policy_version,
prompt_texts=prompt_texts, prompt_texts=prompt_texts,
response_texts=response_texts, response_texts=response_texts,
) )
@@ -370,6 +384,14 @@ class RolloutRunner:
self._cache_key = None self._cache_key = None
self._steps_since_rollout: int = 0 self._steps_since_rollout: int = 0
@property
def policy_version(self) -> int:
return self.generator.policy_version
def update_weights(self, policy_version: int) -> int:
"""Publish the shared policy's new version to the rollout backend."""
return self.generator.update_weights(policy_version)
def step(self): def step(self):
"""Advance the internal counter (call once per optimizer step).""" """Advance the internal counter (call once per optimizer step)."""
self._steps_since_rollout += 1 self._steps_since_rollout += 1
@@ -415,6 +437,7 @@ class RolloutRunner:
response_mask=raw.response_mask, response_mask=raw.response_mask,
rewards=rewards.to(device=device), rewards=rewards.to(device=device),
logprobs_old=raw.logprobs_old, logprobs_old=raw.logprobs_old,
policy_version=raw.policy_version,
prompt_texts=raw.prompt_texts, prompt_texts=raw.prompt_texts,
response_texts=raw.response_texts, response_texts=raw.response_texts,
) )
+7
View File
@@ -239,6 +239,12 @@ class BaseStrategy(ABC):
"""Inject a :class:`RolloutRunner` to enable online rollout mode.""" """Inject a :class:`RolloutRunner` to enable online rollout mode."""
self._rollout_runner = runner self._rollout_runner = runner
@property
def policy_version(self) -> Optional[int]:
if self._rollout_runner is None:
return None
return self._rollout_runner.policy_version
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]: def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
"""Map a :class:`RolloutResult` to the batch layout expected by """Map a :class:`RolloutResult` to the batch layout expected by
:meth:`compute_loss`. :meth:`compute_loss`.
@@ -275,6 +281,7 @@ class BaseStrategy(ABC):
def on_optimizer_step(self): def on_optimizer_step(self):
"""Advance online rollout state after a successful optimizer step.""" """Advance online rollout state after a successful optimizer step."""
if self._rollout_runner is not None: if self._rollout_runner is not None:
self._rollout_runner.update_weights(self.policy_version + 1)
self._rollout_runner.step() self._rollout_runner.step()
def __call__(self, batch: Dict[str, Tensor]) -> LossOutput: def __call__(self, batch: Dict[str, Tensor]) -> LossOutput:
+5
View File
@@ -318,6 +318,11 @@ class TrainContextBuilder:
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=group_size * max(1, cfg.batch_per_device), max_batch_size=group_size * max(1, cfg.batch_per_device),
max_seq_len=getattr(context.model.config, "max_position_embeddings", None), max_seq_len=getattr(context.model.config, "max_position_embeddings", None),
policy_version=(
context.checkpoint.meta.get("policy_version", context.optimizer_step)
if context.checkpoint is not None
else context.optimizer_step
),
) )
generator = RolloutGenerator( generator = RolloutGenerator(
scheduler=scheduler, scheduler=scheduler,
+10
View File
@@ -673,6 +673,7 @@ classDiagram
+Tensor responses +Tensor responses
+Tensor response_mask +Tensor response_mask
+Tensor logprobs_old +Tensor logprobs_old
+int policy_version
+List[str] prompt_texts +List[str] prompt_texts
+List[List[str]] response_texts +List[List[str]] response_texts
} }
@@ -695,10 +696,14 @@ classDiagram
+float top_p +float top_p
+float frequency_penalty +float frequency_penalty
+int rep_window +int rep_window
+int policy_version
+update_weights(policy_version) int
+generate(batch) RawRollout +generate(batch) RawRollout
} }
class RolloutRunner { class RolloutRunner {
+int policy_version
+update_weights(policy_version) int
+step() +step()
+clear_cache() +clear_cache()
+__call__(batch) Tuple[RolloutResult, bool] +__call__(batch) Tuple[RolloutResult, bool]
@@ -854,11 +859,13 @@ classDiagram
+int max_seq_len +int max_seq_len
+str device +str device
+torch.dtype dtype +torch.dtype dtype
+int policy_version
+add_task(prompt, **kwargs) str +add_task(prompt, **kwargs) str
+remove_task(task_id) +remove_task(task_id)
+start() +start()
+stop() +stop()
+get_stats() Dict +get_stats() Dict
+update_weights(policy_version) int
+run_batch(prompt_ids_list, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window, return_logprobs) Union[List[List[int]], List[Tuple[List[int], List[float]]]] +run_batch(prompt_ids_list, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window, return_logprobs) Union[List[List[int]], List[Tuple[List[int], List[float]]]]
} }
@@ -871,6 +878,7 @@ classDiagram
+inc_ref(idx) +inc_ref(idx)
+touch(idx) +touch(idx)
+ref_count(idx) int +ref_count(idx) int
+clear_cached() int
} }
class RadixNode { class RadixNode {
@@ -897,6 +905,7 @@ classDiagram
+extend(state, pos) bool +extend(state, pos) bool
+write_indices(state, prompt_ids) +write_indices(state, prompt_ids)
+record_hashes(state, prompt_ids, start_logical_page) +record_hashes(state, prompt_ids, start_logical_page)
+invalidate_cache() int
} }
class ContiguousStrategy { class ContiguousStrategy {
@@ -960,6 +969,7 @@ classDiagram
+task_extend(task_id, pos) bool +task_extend(task_id, pos) bool
+task_cached(task_id) int +task_cached(task_id) int
+task_record_hashes(task_id, prompt_ids, start_logical_page) +task_record_hashes(task_id, prompt_ids, start_logical_page)
+invalidate_cache() int
+bind(task_ids, workspace) KVCache +bind(task_ids, workspace) KVCache
} }
+6
View File
@@ -139,6 +139,12 @@ attention backends share the same rotary dispatch — it is backend-agnostic.
4. Decode → Run single-token forward for each same-position group 4. Decode → Run single-token forward for each same-position group
``` ```
For in-process training rollout, `InferenceScheduler.update_weights(version)`
acknowledges that the shared model was updated in place. Versions are monotonic;
the scheduler rejects updates while requests are queued and invalidates reusable
prefix KV pages before exposing the new version. Synchronous `run_batch()` and
weight updates are serialized so a generation cannot straddle two versions.
## Sampling (Strategy Pattern) ## Sampling (Strategy Pattern)
``` ```
+7
View File
@@ -166,6 +166,13 @@ them with a `BaseRewardModel`. It refreshes cached rollouts every
`rollout_interval` optimizer steps. `online_grpo` synchronizes `old_model` when `rollout_interval` optimizer steps. `online_grpo` synchronizes `old_model` when
a fresh rollout is produced. a fresh rollout is produced.
Every successful optimizer step advances a monotonic `policy_version` and
acknowledges the shared-model weight update to the rollout scheduler. The
scheduler invalidates reusable KV prefixes before accepting the new version.
`RawRollout` and `RolloutResult` retain the version that actually generated
their behavior log-probabilities, so cached rollout samples remain attributable
even while later optimizer steps advance the live policy.
Online strategies require `TrainConfig.reward_model_fn`. `train.py` exposes the Online strategies require `TrainConfig.reward_model_fn`. `train.py` exposes the
rollout sampling parameters but does not yet offer a CLI argument for the reward rollout sampling parameters but does not yet offer a CLI argument for the reward
model factory. model factory.
+21
View File
@@ -1,5 +1,6 @@
"""Unit tests for inference cache components.""" """Unit tests for inference cache components."""
import pytest
import torch import torch
from astrai.inference.cache import ( from astrai.inference.cache import (
@@ -435,6 +436,26 @@ def test_page_pool_prefix_hit_populates_request_mapping():
) )
def test_task_cache_invalidation_drops_cross_version_prefix_hits():
pool = _make_paged_pool_ps64(page_size=2, max_seq_len=8, n_tokens=16)
task_cache = _make_task_cache(pool)
prompt = [11, 12, 13, 14]
assert task_cache.task_alloc("first", prompt)
task_cache.task_record_hashes("first", prompt)
task_cache.task_free("first")
assert task_cache.task_alloc("cached", prompt)
assert task_cache.task_cached("cached") == len(prompt)
with pytest.raises(RuntimeError, match="while tasks are active"):
task_cache.invalidate_cache()
task_cache.task_free("cached")
assert task_cache.invalidate_cache() == 2
assert task_cache.task_alloc("after_update", prompt)
assert task_cache.task_cached("after_update") == 0
def test_page_pool_paged_ps64_bind_roundtrip(): def test_page_pool_paged_ps64_bind_roundtrip():
pool = _make_paged_pool_ps64(n_layers=1, n_kv_heads=2, head_dim=4) pool = _make_paged_pool_ps64(n_layers=1, n_kv_heads=2, head_dim=4)
task_cache = _make_task_cache(pool) task_cache = _make_task_cache(pool)
+28
View File
@@ -545,6 +545,34 @@ def test_run_batch_empty_prompts(device):
scheduler.stop() scheduler.stop()
def test_scheduler_weight_versions_are_monotonic_and_acknowledged(device):
scheduler, _tok, _model = _make_real_scheduler(device)
try:
assert scheduler.policy_version == 0
assert scheduler.update_weights(1) == 1
assert scheduler.policy_version == 1
assert scheduler.get_stats()["policy_version"] == 1
assert scheduler.update_weights(1) == 1
with pytest.raises(ValueError, match="cannot move backwards"):
scheduler.update_weights(0)
with pytest.raises(ValueError, match="non-negative integer"):
scheduler.update_weights(True)
finally:
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")
try:
with pytest.raises(RuntimeError, match="while tasks are queued"):
scheduler.update_weights(1)
scheduler.remove_task(task_id)
assert scheduler.update_weights(1) == 1
finally:
scheduler.stop()
def test_run_batch_too_long_prompt_skipped(device): def test_run_batch_too_long_prompt_skipped(device):
"""A prompt longer than max_seq_len yields an empty result slot.""" """A prompt longer than max_seq_len yields an empty result slot."""
scheduler, _tok, _model = _make_real_scheduler(device) scheduler, _tok, _model = _make_real_scheduler(device)
+9
View File
@@ -43,6 +43,8 @@ class _RecordingRunner:
self.calls = 0 self.calls = 0
self.step_calls = 0 self.step_calls = 0
self._fresh = True self._fresh = True
self.policy_version = result.policy_version
self.weight_updates = []
def __call__(self, batch): def __call__(self, batch):
self.calls += 1 self.calls += 1
@@ -53,6 +55,11 @@ class _RecordingRunner:
def step(self): def step(self):
self.step_calls += 1 self.step_calls += 1
def update_weights(self, policy_version):
self.policy_version = policy_version
self.weight_updates.append(policy_version)
return policy_version
def swap_result(self, result): def swap_result(self, result):
self.result = result self.result = result
self._fresh = True self._fresh = True
@@ -302,6 +309,8 @@ def test_step_called_when_sync_gradients_true(device):
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step() strat.on_optimizer_step()
assert runner.step_calls == 1 assert runner.step_calls == 1
assert runner.weight_updates == [1]
assert strat.policy_version == 1
def test_loss_is_differentiable_dpo(device): def test_loss_is_differentiable_dpo(device):
+22
View File
@@ -65,6 +65,7 @@ def test_raw_rollout_fields():
) )
assert r.prompts.shape == (2, 4) assert r.prompts.shape == (2, 4)
assert r.responses.shape == (2, 3, 5) assert r.responses.shape == (2, 3, 5)
assert r.policy_version == 0
assert r.prompt_texts == [] assert r.prompt_texts == []
assert r.response_texts == [] assert r.response_texts == []
@@ -131,6 +132,7 @@ def test_rollout_generator_shapes(device):
assert len(r.prompt_texts) == 2 assert len(r.prompt_texts) == 2
assert len(r.response_texts) == 2 assert len(r.response_texts) == 2
assert len(r.response_texts[0]) == 3 assert len(r.response_texts[0]) == 3
assert r.policy_version == 0
def test_rollout_generator_uses_eval_and_restores_mode(device): def test_rollout_generator_uses_eval_and_restores_mode(device):
@@ -275,6 +277,26 @@ def test_rollout_runner_cache_returns_stale_flag(device):
assert fresh2 is False assert fresh2 is False
def test_rollout_runner_tags_generation_version_and_preserves_cached_behavior(device):
runner, _ = _make_runner(device, rollout_interval=100)
batch = _make_instruction_batch(n=1)
first, first_fresh = runner(batch)
assert first_fresh is True
assert first.policy_version == 0
assert runner.update_weights(1) == 1
cached, cached_fresh = runner(batch)
assert cached is first
assert cached_fresh is False
assert cached.policy_version == 0
runner.clear_cache()
refreshed, refreshed_fresh = runner(batch)
assert refreshed_fresh is True
assert refreshed.policy_version == 1
def test_rollout_runner_refreshes_for_different_batch(device): def test_rollout_runner_refreshes_for_different_batch(device):
runner, _ = _make_runner(device, rollout_interval=100) runner, _ = _make_runner(device, rollout_interval=100)
r1, fresh1 = runner(_make_instruction_batch(n=1)) r1, fresh1 = runner(_make_instruction_batch(n=1))