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
+30 -7
View File
@@ -13,6 +13,7 @@ Provides:
so callers do not need to rely on object identity to detect refreshes.
"""
import threading
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple
@@ -53,6 +54,7 @@ class RawRollout:
responses: Tensor
response_mask: Tensor
logprobs_old: Tensor
policy_version: int = 0
prompt_texts: 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.frequency_penalty = frequency_penalty
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()
def generate(self, batch: Dict) -> RawRollout:
@@ -146,13 +158,14 @@ class RolloutGenerator:
``add_generation_prompt=True`` so rollout prompts match the
format the policy was SFT-trained on.
"""
model = self.scheduler._executor.model
was_training = model.training
model.eval()
try:
return self._generate_eval(batch)
finally:
model.train(was_training)
with self._weight_lock:
model = self.scheduler._executor.model
was_training = model.training
model.eval()
try:
return self._generate_eval(batch)
finally:
model.train(was_training)
def _generate_eval(self, batch: Dict) -> RawRollout:
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
@@ -245,6 +258,7 @@ class RolloutGenerator:
responses=responses,
response_mask=response_mask,
logprobs_old=logprobs_old,
policy_version=self.policy_version,
prompt_texts=prompt_texts,
response_texts=response_texts,
)
@@ -370,6 +384,14 @@ class RolloutRunner:
self._cache_key = None
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):
"""Advance the internal counter (call once per optimizer step)."""
self._steps_since_rollout += 1
@@ -415,6 +437,7 @@ class RolloutRunner:
response_mask=raw.response_mask,
rewards=rewards.to(device=device),
logprobs_old=raw.logprobs_old,
policy_version=raw.policy_version,
prompt_texts=raw.prompt_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."""
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]:
"""Map a :class:`RolloutResult` to the batch layout expected by
:meth:`compute_loss`.
@@ -275,6 +281,7 @@ class BaseStrategy(ABC):
def on_optimizer_step(self):
"""Advance online rollout state after a successful optimizer step."""
if self._rollout_runner is not None:
self._rollout_runner.update_weights(self.policy_version + 1)
self._rollout_runner.step()
def __call__(self, batch: Dict[str, Tensor]) -> LossOutput:
+5
View File
@@ -318,6 +318,11 @@ class TrainContextBuilder:
tokenizer=tokenizer,
max_batch_size=group_size * max(1, cfg.batch_per_device),
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(
scheduler=scheduler,