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:
Vendored
+8
@@ -338,6 +338,14 @@ class TaskCacheManager:
|
||||
if state is not None:
|
||||
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
|
||||
def task_cacheable_ids(task_id: str, prompt_ids: List[int], output_ids: List[int]):
|
||||
return list(prompt_ids) + list(output_ids[:-1])
|
||||
|
||||
Vendored
+22
@@ -89,6 +89,19 @@ class Allocator:
|
||||
if idx in self._lru:
|
||||
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:
|
||||
"""A page-aligned edge in the CPU-side prefix radix trie."""
|
||||
@@ -200,6 +213,10 @@ class AllocationStrategy(ABC):
|
||||
start: int,
|
||||
) -> None: ...
|
||||
|
||||
def invalidate_cache(self) -> int:
|
||||
"""Drop reusable KV entries after an inference weight update."""
|
||||
return 0
|
||||
|
||||
|
||||
class ContiguousStrategy(AllocationStrategy):
|
||||
"""Static contiguous allocation: slots are pre-assigned at pool init.
|
||||
@@ -316,3 +333,8 @@ class PagedStrategy(AllocationStrategy):
|
||||
full = len(prompt_ids) // self._page_size
|
||||
for i in range(start, min(full, len(state.pages))):
|
||||
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()
|
||||
|
||||
@@ -2,6 +2,7 @@ import logging
|
||||
import threading
|
||||
import uuid
|
||||
from contextlib import nullcontext
|
||||
from functools import wraps
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
@@ -28,6 +29,15 @@ from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
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:
|
||||
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||
|
||||
@@ -42,7 +52,14 @@ class InferenceScheduler:
|
||||
cache: Optional[PagePool] = None,
|
||||
enable_cuda_graph: bool = True,
|
||||
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
|
||||
|
||||
if max_seq_len is not None:
|
||||
@@ -103,6 +120,44 @@ class InferenceScheduler:
|
||||
|
||||
self._stop_event = threading.Event()
|
||||
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:
|
||||
return self._task_mgr.add_task(prompt, **kwargs)
|
||||
@@ -125,6 +180,7 @@ class InferenceScheduler:
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
stats = self._task_mgr.get_stats()
|
||||
stats["kv_cache_tasks"] = self._task_cache.task_count
|
||||
stats["policy_version"] = self._policy_version
|
||||
return stats
|
||||
|
||||
@property
|
||||
@@ -343,6 +399,7 @@ class InferenceScheduler:
|
||||
)
|
||||
self._task_mgr.clear_queues()
|
||||
|
||||
@_with_weight_lock
|
||||
def run_batch(
|
||||
self,
|
||||
prompt_ids_list: List[List[int]],
|
||||
|
||||
Reference in New Issue
Block a user