diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index 4b2dc54a..3d7fdeb7 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -64,10 +64,19 @@ class TrainConfig(BaseConfig): neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0. moe_aux_loss_coef (float): Weight applied to the MoE load-balancing loss. Defaults to 0.01. rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512. + rollout_max_policy_lag (Optional[int]): Maximum accepted gap between rollout and live policy versions. None derives ``rollout_interval - 1``. Defaults to None. rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7. rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0. rollout_top_p (float): Top-p (nucleus) filtering for online rollout. Defaults to 0.9. rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024. + rollout_dynamic_sampling (bool): Refill low-variance online GRPO prompt groups. Defaults to False. + rollout_dynamic_variance_threshold (float): Minimum population reward variance for group acceptance. Defaults to 0.0. + rollout_dynamic_max_refill_rounds (int): Maximum refill attempts after initial generation. Defaults to 2. + rollout_dynamic_max_generated_tokens_per_group (int): Hard per-group generation budget. Defaults to 32768. + rollout_dynamic_max_wall_time_per_group (float): Hard per-group wall-clock budget in seconds. Defaults to 300. + rollout_dynamic_max_total_tokens_per_step (int): Hard generated-token budget per training step. Defaults to 262144. + rollout_dynamic_max_pending_groups (int): Maximum groups admitted into a sampling step. Defaults to 128. + rollout_dynamic_seed (Optional[int]): Base refill seed; None derives from random_seed. Defaults to None. reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None. executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}. strategy_kwargs (Dict[str, Any]): Extra strategy arguments. Defaults to {}. @@ -118,10 +127,19 @@ class TrainConfig(BaseConfig): moe_aux_loss_coef: float = 0.01 rollout_interval: int = 512 + rollout_max_policy_lag: Optional[int] = None rollout_temperature: float = 0.7 rollout_top_k: int = 0 rollout_top_p: float = 0.9 rollout_max_tokens: int = 1024 + rollout_dynamic_sampling: bool = False + rollout_dynamic_variance_threshold: float = 0.0 + rollout_dynamic_max_refill_rounds: int = 2 + rollout_dynamic_max_generated_tokens_per_group: int = 32_768 + rollout_dynamic_max_wall_time_per_group: float = 300.0 + rollout_dynamic_max_total_tokens_per_step: int = 262_144 + rollout_dynamic_max_pending_groups: int = 128 + rollout_dynamic_seed: Optional[int] = None reward_model_fn: Optional[Callable] = None executor_kwargs: Dict[str, Any] = field(default_factory=dict) @@ -173,13 +191,16 @@ def _validate_compile_mode(cls, v: Optional[str]) -> Optional[str]: "val_step", "rollout_interval", "rollout_max_tokens", + "rollout_dynamic_max_generated_tokens_per_group", + "rollout_dynamic_max_total_tokens_per_step", + "rollout_dynamic_max_pending_groups", ) def _validate_positive_int(cls, v: int) -> int: if v <= 0: raise ValueError(f"must be positive, got {v}") return v - @field_validator("rollout_temperature") + @field_validator("rollout_temperature", "rollout_dynamic_max_wall_time_per_group") def _validate_positive_float(cls, v: float) -> float: if v <= 0: raise ValueError(f"must be positive, got {v}") @@ -192,13 +213,24 @@ def _validate_top_p(cls, v: float) -> float: return v @field_validator( - "rollout_top_k", "num_workers", "neftune_alpha", "moe_aux_loss_coef" + "rollout_top_k", + "rollout_dynamic_max_refill_rounds", + "rollout_dynamic_variance_threshold", + "num_workers", + "neftune_alpha", + "moe_aux_loss_coef", ) def _validate_non_negative(cls, v): if v < 0: raise ValueError(f"must be non-negative, got {v}") return v + @field_validator("rollout_max_policy_lag", "rollout_dynamic_seed") + def _validate_optional_non_negative_int(cls, v: Optional[int]) -> Optional[int]: + if v is not None and v < 0: + raise ValueError(f"must be non-negative or None, got {v}") + return v + @field_validator("max_grad_norm") def _validate_max_grad_norm(cls, v: Optional[float]) -> Optional[float]: if v is not None and v <= 0: @@ -226,4 +258,14 @@ def _validate_online_strategy(self) -> "TrainConfig": f"numbers of forward passes and desynchronize the " f"ddp/fsdp collectives, deadlocking NCCL" ) + if self.rollout_dynamic_sampling: + if self.strategy != "online_grpo": + raise ValueError( + "rollout_dynamic_sampling is supported only for " + "strategy='online_grpo'" + ) + if self.strategy_kwargs.get("group_size", 1) < 2: + raise ValueError( + "rollout_dynamic_sampling requires strategy_kwargs.group_size >= 2" + ) return self diff --git a/astrai/inference/scheduler.py b/astrai/inference/scheduler.py index 3121ae32..3de9ac05 100644 --- a/astrai/inference/scheduler.py +++ b/astrai/inference/scheduler.py @@ -3,7 +3,7 @@ import uuid from contextlib import nullcontext from functools import wraps -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any, Callable, Dict, List, Optional, Tuple, TypeVar, Union import torch @@ -27,6 +27,7 @@ from astrai.tokenize.tokenizer import AutoTokenizer logger = logging.getLogger(__name__) +T = TypeVar("T") def _with_weight_lock(method): @@ -128,15 +129,9 @@ 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. - """ + def _validate_weight_version( + self, policy_version: int, *, require_advance: bool = False + ) -> None: if ( isinstance(policy_version, bool) or not isinstance(policy_version, int) @@ -148,17 +143,57 @@ def update_weights(self, policy_version: int) -> int: 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 require_advance and policy_version == self._policy_version: + raise ValueError( + f"policy_version must advance beyond {self._policy_version} " + "when model weights are mutated" + ) + + def _ensure_weight_update_ready(self) -> None: 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") + def _commit_weight_version(self, policy_version: int) -> int: self._task_cache.invalidate_cache() self._policy_version = policy_version 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. + """ + self._validate_weight_version(policy_version) + if policy_version == self._policy_version: + return self._policy_version + self._ensure_weight_update_ready() + return self._commit_weight_version(policy_version) + + @_with_weight_lock + def apply_weight_update(self, policy_version: int, update: Callable[[], T]) -> T: + """Mutate shared weights and publish their version without generation.""" + if not callable(update): + raise TypeError("update must be callable") + self._validate_weight_version(policy_version, require_advance=True) + self._ensure_weight_update_ready() + + result = update() + self._commit_weight_version(policy_version) + return result + + @_with_weight_lock + def with_policy_snapshot(self, inspect: Callable[[int], T]) -> T: + """Inspect state while the scheduler's policy version remains stable.""" + if not callable(inspect): + raise TypeError("inspect must be callable") + return inspect(self._policy_version) + def add_task(self, prompt: str, **kwargs) -> str: return self._task_mgr.add_task(prompt, **kwargs) diff --git a/astrai/trainer/rollout.py b/astrai/trainer/rollout.py index efa04561..167e57fb 100644 --- a/astrai/trainer/rollout.py +++ b/astrai/trainer/rollout.py @@ -13,12 +13,16 @@ so callers do not need to rely on object identity to detect refreshes. """ +import hashlib import threading +import time from abc import ABC, abstractmethod from dataclasses import dataclass, field -from typing import Dict, List, Optional, Tuple +from enum import Enum +from typing import Callable, Dict, List, Optional, Tuple, TypeVar import torch +import torch.distributed as dist from torch import Tensor from astrai.inference.scheduler import InferenceScheduler @@ -57,6 +61,7 @@ class RawRollout: policy_version: int = 0 prompt_texts: List[str] = field(default_factory=list) response_texts: List[List[str]] = field(default_factory=list) + sampling_groups: List["DynamicSamplingGroup"] = field(default_factory=list) @dataclass(kw_only=True) @@ -98,6 +103,156 @@ def score(self, prompts: List[str], responses: List[List[str]]) -> Tensor: _PAD = 0 +T = TypeVar("T") + + +class RolloutVersionError(RuntimeError): + """A rollout cannot be attributed to an acceptable policy version.""" + + +class DynamicSamplingBudgetError(RuntimeError): + """A dynamic-sampling step cannot produce a complete, safe batch.""" + + +class DynamicSamplingState(str, Enum): + """Lifecycle of one prompt group's generation attempt.""" + + PENDING = "pending" + GENERATING = "generating" + SCORING = "scoring" + ACCEPTED = "accepted" + REFILL = "refill" + INVALIDATED = "invalidated" + DROPPED = "dropped" + + +_DYNAMIC_SAMPLING_TRANSITIONS = { + DynamicSamplingState.PENDING: {DynamicSamplingState.GENERATING}, + DynamicSamplingState.GENERATING: { + DynamicSamplingState.SCORING, + DynamicSamplingState.INVALIDATED, + DynamicSamplingState.DROPPED, + }, + DynamicSamplingState.SCORING: { + DynamicSamplingState.ACCEPTED, + DynamicSamplingState.REFILL, + DynamicSamplingState.INVALIDATED, + DynamicSamplingState.DROPPED, + }, + # Acceptance is provisional until the whole step is committed. A policy + # change while another group is refilled invalidates every old-version row. + DynamicSamplingState.ACCEPTED: {DynamicSamplingState.INVALIDATED}, + DynamicSamplingState.REFILL: set(), + DynamicSamplingState.INVALIDATED: set(), + DynamicSamplingState.DROPPED: set(), +} + + +@dataclass +class DynamicSamplingGroup: + """Auditable metadata for one prompt-group generation attempt.""" + + prompt_uid: str + attempt_id: int + generation_seed: int + behavior_policy_version: Optional[int] = None + reward_vector: List[float] = field(default_factory=list) + reward_variance: Optional[float] = None + accepted: bool = False + discard_reason: Optional[str] = None + refill_round: int = 0 + created_at: float = field(default_factory=time.time) + completed_at: Optional[float] = None + generated_tokens: int = 0 + state: DynamicSamplingState = DynamicSamplingState.PENDING + + def transition( + self, + state: DynamicSamplingState, + *, + discard_reason: Optional[str] = None, + ) -> None: + """Move to ``state`` while enforcing the lifecycle graph.""" + if state not in _DYNAMIC_SAMPLING_TRANSITIONS[self.state]: + raise RuntimeError( + f"invalid dynamic-sampling transition {self.state.value} -> " + f"{state.value}" + ) + self.state = state + self.discard_reason = discard_reason + self.accepted = state is DynamicSamplingState.ACCEPTED + if state in { + DynamicSamplingState.ACCEPTED, + DynamicSamplingState.REFILL, + DynamicSamplingState.INVALIDATED, + DynamicSamplingState.DROPPED, + }: + self.completed_at = time.time() + + +@dataclass(frozen=True) +class DynamicSamplingConfig: + """Budgets and acceptance policy for versioned rollout refill.""" + + enabled: bool = False + variance_threshold: float = 0.0 + max_refill_rounds: int = 2 + max_generated_tokens_per_group: int = 32_768 + max_wall_time_per_group: float = 300.0 + max_total_rollout_tokens_per_step: int = 262_144 + max_pending_groups: int = 128 + base_seed: int = 3407 + + def __post_init__(self) -> None: + if self.variance_threshold < 0: + raise ValueError("variance_threshold must be non-negative") + if self.max_refill_rounds < 0: + raise ValueError("max_refill_rounds must be non-negative") + for name in ( + "max_generated_tokens_per_group", + "max_total_rollout_tokens_per_step", + "max_pending_groups", + ): + if getattr(self, name) <= 0: + raise ValueError(f"{name} must be positive") + if self.max_wall_time_per_group <= 0: + raise ValueError("max_wall_time_per_group must be positive") + if self.base_seed < 0: + raise ValueError("base_seed must be non-negative") + + +@dataclass +class DynamicSamplingMetrics: + """Per-refresh dynamic-sampling counters and efficiency metrics.""" + + groups_total: int = 0 + groups_accepted: int = 0 + groups_zero_variance: int = 0 + refill_rounds: int = 0 + refill_tokens: int = 0 + groups_dropped: int = 0 + groups_version_invalidated: int = 0 + groups_budget_exhausted: int = 0 + total_generated_tokens: int = 0 + accepted_generated_tokens: int = 0 + + def as_dict(self) -> Dict[str, float]: + total = self.total_generated_tokens + effective = 0.0 if total == 0 else self.groups_accepted * 1_000_000 / total + waste = 0.0 if total == 0 else (total - self.accepted_generated_tokens) / total + return { + "groups_total": float(self.groups_total), + "groups_accepted": float(self.groups_accepted), + "zero_variance_groups": float(self.groups_zero_variance), + "refill_rounds": float(self.refill_rounds), + "refill_tokens": float(self.refill_tokens), + "dropped_groups": float(self.groups_dropped), + "version_invalidated_groups": float(self.groups_version_invalidated), + "budget_exhausted_groups": float(self.groups_budget_exhausted), + "total_generated_tokens": float(total), + "effective_groups_per_million_tokens": effective, + "rollout_waste_ratio": waste, + } class RolloutGenerator: @@ -142,8 +297,22 @@ def update_weights(self, policy_version: int) -> int: with self._weight_lock: return self.scheduler.update_weights(policy_version) + def apply_weight_update(self, policy_version: int, update: Callable[[], T]) -> T: + """Apply a shared-model mutation at an atomic generation boundary.""" + with self._weight_lock: + return self.scheduler.apply_weight_update(policy_version, update) + + def with_policy_snapshot(self, inspect: Callable[[int], T]) -> T: + """Inspect a version stable against generator and scheduler updates.""" + if not callable(inspect): + raise TypeError("inspect must be callable") + with self._weight_lock: + return self.scheduler.with_policy_snapshot(inspect) + @torch.no_grad() - def generate(self, batch: Dict) -> RawRollout: + def generate( + self, batch: Dict, *, generation_seed: Optional[int] = None + ) -> RawRollout: """Expand prompts by ``group_size`` and generate one response each. Accepted batch formats (per sample, repeated B times): @@ -159,15 +328,32 @@ def generate(self, batch: Dict) -> RawRollout: format the policy was SFT-trained on. """ 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: + + def generate_snapshot(generation_version: int) -> RawRollout: + model = self.scheduler._executor.model + was_training = model.training + model.eval() + try: + if generation_seed is None: + return self._generate_eval(batch, generation_version) + if generation_seed < 0: + raise ValueError("generation_seed must be non-negative") + device = torch.device(self.scheduler.device) + cuda_devices = [device] if device.type == "cuda" else [] + # Restore caller RNG state after deterministic generation; + # retries must not perturb training-side randomness. + with torch.random.fork_rng(devices=cuda_devices): + torch.manual_seed(generation_seed) + return self._generate_eval(batch, generation_version) + finally: + model.train(was_training) + + # Capture the version under the scheduler lock as well as the + # generator lock. This also serializes callers that update the + # scheduler directly instead of going through this wrapper. + return self.scheduler.with_policy_snapshot(generate_snapshot) + + def _generate_eval(self, batch: Dict, generation_version: int) -> RawRollout: prompt_texts, flat_prompt_ids = self._prepare_prompts(batch) B = len(prompt_texts) G = self.group_size @@ -258,7 +444,7 @@ def _generate_eval(self, batch: Dict) -> RawRollout: responses=responses, response_mask=response_mask, logprobs_old=logprobs_old, - policy_version=self.policy_version, + policy_version=generation_version, prompt_texts=prompt_texts, response_texts=response_texts, ) @@ -375,14 +561,28 @@ def __init__( generator: RolloutGenerator, reward_model: BaseRewardModel, rollout_interval: int = 512, + max_policy_lag: Optional[int] = None, + dynamic_sampling: Optional[DynamicSamplingConfig] = None, ): + if rollout_interval <= 0: + raise ValueError("rollout_interval must be positive") + if max_policy_lag is not None and max_policy_lag < 0: + raise ValueError("max_policy_lag must be non-negative or None") self.generator = generator self.reward_model = reward_model self.rollout_interval = rollout_interval + self.max_policy_lag = ( + rollout_interval - 1 if max_policy_lag is None else max_policy_lag + ) + self.dynamic_sampling = dynamic_sampling or DynamicSamplingConfig() self._cache: Optional[RolloutResult] = None self._cache_key = None self._steps_since_rollout: int = 0 + self._dynamic_refresh_id: int = 0 + self._dynamic_attempt_id: int = 0 + self._last_sampling_metrics: Dict[str, float] = {} + self._last_sampling_history: List[DynamicSamplingGroup] = [] @property def policy_version(self) -> int: @@ -392,6 +592,10 @@ 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 apply_weight_update(self, policy_version: int, update: Callable[[], T]) -> T: + """Apply a model update and publish its version as one operation.""" + return self.generator.apply_weight_update(policy_version, update) + def step(self): """Advance the internal counter (call once per optimizer step).""" self._steps_since_rollout += 1 @@ -401,6 +605,16 @@ def clear_cache(self): self._cache = None self._cache_key = None + @property + def last_sampling_metrics(self) -> Dict[str, float]: + """Metrics from the most recent dynamic-sampling refresh.""" + return dict(self._last_sampling_metrics) + + @property + def last_sampling_history(self) -> List[DynamicSamplingGroup]: + """Attempt records from the most recent dynamic-sampling refresh.""" + return list(self._last_sampling_history) + @staticmethod def _batch_key(batch: Dict): """Build a stable key for the prompt fields accepted by the generator.""" @@ -440,8 +654,470 @@ def _score(self, raw: RawRollout) -> RolloutResult: policy_version=raw.policy_version, prompt_texts=raw.prompt_texts, response_texts=raw.response_texts, + sampling_groups=raw.sampling_groups, ) + @staticmethod + def _batch_size(batch: Dict) -> int: + if "messages" in batch: + return len(batch["messages"]) + if "instruction" in batch: + return len(batch["instruction"]) + raise ValueError("dynamic sampling requires messages or instruction prompts") + + @staticmethod + def _select_batch(batch: Dict, indices: List[int]) -> Dict: + """Select prompt rows without retaining unrelated training tensors.""" + selected = {} + for key in ("messages", "instruction", "input", "output"): + if key not in batch: + continue + value = batch[key] + if isinstance(value, Tensor): + index = torch.tensor(indices, dtype=torch.long, device=value.device) + selected[key] = value.index_select(0, index) + elif isinstance(value, tuple): + selected[key] = tuple(value[index] for index in indices) + else: + selected[key] = [value[index] for index in indices] + return selected + + def _prompt_uids(self, batch: Dict, batch_size: int) -> List[str]: + batch_digest = hashlib.sha256(repr(self._batch_key(batch)).encode()).hexdigest() + return [f"{batch_digest[:20]}:{index}" for index in range(batch_size)] + + @staticmethod + def _collective_device(device: torch.device) -> torch.device: + if dist.is_available() and dist.is_initialized(): + return device if dist.get_backend() == "nccl" else torch.device("cpu") + return device + + @classmethod + def _synchronize_version(cls, version: int, device: torch.device) -> int: + if not (dist.is_available() and dist.is_initialized()): + return version + collective_device = cls._collective_device(device) + minimum = torch.tensor(version, dtype=torch.long, device=collective_device) + maximum = minimum.clone() + dist.all_reduce(minimum, op=dist.ReduceOp.MIN) + dist.all_reduce(maximum, op=dist.ReduceOp.MAX) + min_version, max_version = int(minimum.item()), int(maximum.item()) + if min_version != max_version: + raise RolloutVersionError( + "rollout policy version differs across ranks: " + f"min={min_version}, max={max_version}" + ) + return min_version + + @classmethod + def _synchronize_acceptance( + cls, accepted: List[bool], device: torch.device + ) -> List[bool]: + if not (dist.is_available() and dist.is_initialized()): + return accepted + collective_device = cls._collective_device(device) + flags = torch.tensor(accepted, dtype=torch.int32, device=collective_device) + dist.all_reduce(flags, op=dist.ReduceOp.MIN) + return [bool(value) for value in flags.cpu().tolist()] + + @classmethod + def _synchronize_failure(cls, failed: bool, device: torch.device) -> bool: + if not (dist.is_available() and dist.is_initialized()): + return failed + collective_device = cls._collective_device(device) + flag = torch.tensor(int(failed), dtype=torch.int32, device=collective_device) + dist.all_reduce(flag, op=dist.ReduceOp.MAX) + return bool(flag.item()) + + @staticmethod + def _combine_dynamic_rows( + accepted: Dict[int, Tuple[RolloutResult, int, DynamicSamplingGroup]], + *, + batch_size: int, + policy_version: int, + ) -> RolloutResult: + """Reassemble accepted rows, padding attempts to common dimensions.""" + ordered = [accepted[index] for index in range(batch_size)] + first = ordered[0][0] + group_size = first.responses.size(1) + max_prompt_len = max(result.prompts.size(1) for result, _, _ in ordered) + max_response_len = max(result.responses.size(2) for result, _, _ in ordered) + device = first.prompts.device + + prompts = torch.zeros( + batch_size, max_prompt_len, dtype=first.prompts.dtype, device=device + ) + prompt_mask = torch.zeros( + batch_size, max_prompt_len, dtype=torch.bool, device=device + ) + responses = torch.full( + (batch_size, group_size, max_response_len), + _PAD, + dtype=first.responses.dtype, + device=device, + ) + response_mask = torch.zeros_like(responses, dtype=torch.bool) + logprobs_old = torch.zeros( + batch_size, + group_size, + max_response_len, + dtype=first.logprobs_old.dtype, + device=device, + ) + rewards = torch.zeros( + batch_size, group_size, dtype=first.rewards.dtype, device=device + ) + prompt_texts: List[str] = [] + response_texts: List[List[str]] = [] + sampling_groups: List[DynamicSamplingGroup] = [] + + for destination, (result, source, record) in enumerate(ordered): + if result.policy_version != policy_version: + raise RolloutVersionError( + "dynamic sampling attempted to mix behavior policy versions" + ) + prompt_len = result.prompts.size(1) + response_len = result.responses.size(2) + prompts[destination, -prompt_len:] = result.prompts[source] + prompt_mask[destination, -prompt_len:] = result.prompt_mask[source] + responses[destination, :, :response_len] = result.responses[source] + response_mask[destination, :, :response_len] = result.response_mask[source] + logprobs_old[destination, :, :response_len] = result.logprobs_old[source] + rewards[destination] = result.rewards[source] + prompt_texts.append(result.prompt_texts[source]) + response_texts.append(result.response_texts[source]) + sampling_groups.append(record) + + return RolloutResult( + prompts=prompts, + prompt_mask=prompt_mask, + responses=responses, + response_mask=response_mask, + logprobs_old=logprobs_old, + rewards=rewards, + policy_version=policy_version, + prompt_texts=prompt_texts, + response_texts=response_texts, + sampling_groups=sampling_groups, + ) + + def _generate_dynamic(self, batch: Dict) -> RolloutResult: + """Generate a full batch without mixing behavior-policy versions.""" + config = self.dynamic_sampling + batch_size = self._batch_size(batch) + if batch_size == 0: + raise ValueError("dynamic sampling requires at least one prompt group") + if batch_size > config.max_pending_groups: + raise DynamicSamplingBudgetError( + f"pending groups {batch_size} exceed max_pending_groups=" + f"{config.max_pending_groups}" + ) + + metrics = DynamicSamplingMetrics(groups_total=batch_size) + history: List[DynamicSamplingGroup] = [] + accepted: Dict[int, Tuple[RolloutResult, int, DynamicSamplingGroup]] = {} + pending = list(range(batch_size)) + prompt_uids = self._prompt_uids(batch, batch_size) + refill_rounds = [0] * batch_size + generated_per_group = [0] * batch_size + group_started = [time.monotonic()] * batch_size + target_version: Optional[int] = None + max_attempt_tokens = self.generator.group_size * self.generator.max_tokens + self._dynamic_refresh_id += 1 + + try: + while pending: + seed = ( + config.base_seed + + self._dynamic_refresh_id * 1_000_003 + + self._dynamic_attempt_id + ) + records: List[DynamicSamplingGroup] = [] + for index in pending: + self._dynamic_attempt_id += 1 + record = DynamicSamplingGroup( + prompt_uid=prompt_uids[index], + attempt_id=self._dynamic_attempt_id, + generation_seed=seed, + refill_round=refill_rounds[index], + ) + record.transition(DynamicSamplingState.GENERATING) + records.append(record) + + preflight_reason = None + if any( + generated_per_group[index] + max_attempt_tokens + > config.max_generated_tokens_per_group + for index in pending + ): + preflight_reason = "max_generated_tokens_per_group" + elif ( + metrics.total_generated_tokens + len(pending) * max_attempt_tokens + > config.max_total_rollout_tokens_per_step + ): + preflight_reason = "max_total_rollout_tokens_per_step" + elif any( + time.monotonic() - group_started[index] + >= config.max_wall_time_per_group + for index in pending + ): + preflight_reason = "max_wall_time_per_group" + + preflight_failed = self._synchronize_failure( + preflight_reason is not None, + self._collective_device( + torch.device(self.generator.scheduler.device) + ), + ) + if preflight_failed: + reason = preflight_reason or "peer_rank_budget_exhausted" + for record in records: + record.transition( + DynamicSamplingState.DROPPED, + discard_reason=reason, + ) + for _, _, record in accepted.values(): + record.transition( + DynamicSamplingState.INVALIDATED, + discard_reason="peer_group_budget_exhausted", + ) + history.extend(records) + history.extend(record for _, _, record in accepted.values()) + metrics.groups_dropped += len(pending) + metrics.groups_budget_exhausted += len(pending) + raise DynamicSamplingBudgetError( + "dynamic sampling cannot start another generation " + f"attempt within {reason}" + ) + + subset = self._select_batch(batch, pending) + raw = None + generation_error = None + try: + raw = self.generator.generate(subset, generation_seed=seed) + except Exception as error: + generation_error = error + generation_failed = self._synchronize_failure( + generation_error is not None, + torch.device(self.generator.scheduler.device), + ) + if generation_failed: + for record in records: + record.transition( + DynamicSamplingState.DROPPED, + discard_reason="generation_failed", + ) + for _, _, record in accepted.values(): + record.transition( + DynamicSamplingState.INVALIDATED, + discard_reason="peer_group_generation_failed", + ) + history.extend(records) + history.extend(record for _, _, record in accepted.values()) + metrics.groups_dropped += len(pending) + if generation_error is not None: + raise generation_error + raise RuntimeError( + "dynamic sampling generation failed on a peer rank" + ) + assert raw is not None + version_error = None + try: + self._validate_policy_version(raw) + except RolloutVersionError as error: + version_error = error + version_failed = self._synchronize_failure( + version_error is not None, raw.prompts.device + ) + if version_failed: + for record in records: + record.transition( + DynamicSamplingState.INVALIDATED, + discard_reason="policy_version_rejected", + ) + for _, _, record in accepted.values(): + record.transition( + DynamicSamplingState.INVALIDATED, + discard_reason="peer_policy_version_rejected", + ) + history.extend(records) + history.extend(record for _, _, record in accepted.values()) + metrics.groups_version_invalidated += len(pending) + len(accepted) + if version_error is not None: + raise version_error + raise RolloutVersionError( + "dynamic sampling policy version was rejected on a peer rank" + ) + version = self._synchronize_version( + raw.policy_version, raw.prompts.device + ) + + token_counts = raw.response_mask.sum(dim=(1, 2)).cpu().tolist() + for record, token_count, index in zip(records, token_counts, pending): + record.behavior_policy_version = version + record.generated_tokens = int(token_count) + generated_per_group[index] += int(token_count) + metrics.total_generated_tokens += int(token_count) + if record.refill_round > 0: + metrics.refill_rounds += 1 + metrics.refill_tokens += int(token_count) + + if target_version is not None and version != target_version: + for record in records: + record.transition( + DynamicSamplingState.INVALIDATED, + discard_reason="policy_version_changed", + ) + for _, _, record in accepted.values(): + record.transition( + DynamicSamplingState.INVALIDATED, + discard_reason="policy_version_changed_during_refill", + ) + invalidated = len(records) + len(accepted) + metrics.groups_version_invalidated += invalidated + history.extend(records) + history.extend(record for _, _, record in accepted.values()) + accepted.clear() + pending = list(range(batch_size)) + refill_rounds = [0] * batch_size + target_version = version + continue + + if target_version is None: + target_version = version + + for record in records: + record.transition(DynamicSamplingState.SCORING) + scored = None + scoring_error = None + try: + scored = self._score(raw) + except Exception as error: + scoring_error = error + scoring_failed = self._synchronize_failure( + scoring_error is not None, raw.prompts.device + ) + if scoring_failed: + for record in records: + record.transition( + DynamicSamplingState.DROPPED, + discard_reason="scoring_failed", + ) + for _, _, record in accepted.values(): + record.transition( + DynamicSamplingState.INVALIDATED, + discard_reason="peer_group_scoring_failed", + ) + history.extend(records) + history.extend(record for _, _, record in accepted.values()) + metrics.groups_dropped += len(pending) + if scoring_error is not None: + raise scoring_error + raise RuntimeError("dynamic sampling scoring failed on a peer rank") + assert scored is not None + variances = scored.rewards.float().var(dim=1, unbiased=False) + locally_accepted = [ + bool(value > config.variance_threshold) for value in variances + ] + jointly_accepted = self._synchronize_acceptance( + locally_accepted, scored.prompts.device + ) + + next_pending: List[int] = [] + local_budget_exhausted = False + remaining_step_tokens = ( + config.max_total_rollout_tokens_per_step + - metrics.total_generated_tokens + ) + for row, (index, record, accept) in enumerate( + zip(pending, records, jointly_accepted) + ): + rewards = scored.rewards[row].detach().float().cpu() + record.reward_vector = rewards.tolist() + record.reward_variance = float(variances[row].item()) + if accept: + record.transition(DynamicSamplingState.ACCEPTED) + accepted[index] = (scored, row, record) + continue + + metrics.groups_zero_variance += 1 + elapsed = time.monotonic() - group_started[index] + reason = None + if refill_rounds[index] >= config.max_refill_rounds: + reason = "max_refill_rounds" + elif ( + generated_per_group[index] + max_attempt_tokens + > config.max_generated_tokens_per_group + ): + reason = "max_generated_tokens_per_group" + elif elapsed >= config.max_wall_time_per_group: + reason = "max_wall_time_per_group" + elif remaining_step_tokens < max_attempt_tokens: + reason = "max_total_rollout_tokens_per_step" + + if reason is not None: + record.transition( + DynamicSamplingState.DROPPED, + discard_reason=reason, + ) + metrics.groups_dropped += 1 + metrics.groups_budget_exhausted += 1 + local_budget_exhausted = True + else: + record.transition( + DynamicSamplingState.REFILL, + discard_reason="reward_variance_below_threshold", + ) + refill_rounds[index] += 1 + remaining_step_tokens -= max_attempt_tokens + next_pending.append(index) + history.append(record) + + failed = self._synchronize_failure( + local_budget_exhausted, scored.prompts.device + ) + if failed: + raise DynamicSamplingBudgetError( + "dynamic sampling exhausted a group budget; refusing " + "a partial or cross-rank-inconsistent training batch" + ) + pending = next_pending + + assert target_version is not None + final = self._combine_dynamic_rows( + accepted, batch_size=batch_size, policy_version=target_version + ) + metrics.groups_accepted = batch_size + metrics.accepted_generated_tokens = int(final.response_mask.sum().item()) + history.extend(record for _, _, record in accepted.values()) + self._last_sampling_metrics = metrics.as_dict() + self._last_sampling_history = history + return final + except BaseException: + self._last_sampling_metrics = metrics.as_dict() + self._last_sampling_history = history + raise + + def _validate_policy_version( + self, result: RawRollout, *, live_version: Optional[int] = None + ) -> None: + version = result.policy_version + if isinstance(version, bool) or not isinstance(version, int) or version < 0: + raise RolloutVersionError(f"rollout has invalid policy version {version!r}") + if live_version is None: + live_version = self.policy_version + if version > live_version: + raise RolloutVersionError( + f"rollout has future policy version {version}; " + f"live policy version is {live_version}" + ) + lag = live_version - version + if lag > self.max_policy_lag: + raise RolloutVersionError( + f"rollout policy lag {lag} exceeds max_policy_lag=" + f"{self.max_policy_lag} (rollout={version}, live={live_version})" + ) + def __call__(self, batch: Dict[str, Tensor]) -> Tuple[RolloutResult, bool]: """Return ``(cached or fresh) RolloutResult`` plus an ``is_fresh`` flag. @@ -454,9 +1130,38 @@ def __call__(self, batch: Dict[str, Tensor]) -> Tuple[RolloutResult, bool]: or cache_key != self._cache_key or self._steps_since_rollout >= self.rollout_interval ): - raw = self.generator.generate(batch) - self._cache = self._score(raw) - self._cache_key = cache_key - self._steps_since_rollout = 0 - return self._cache, True - return self._cache, False + if self.dynamic_sampling.enabled: + scored = self._generate_dynamic(batch) + else: + raw = self.generator.generate(batch) + self._validate_policy_version(raw) + scored = self._score(raw) + + def commit(live_version: int) -> Tuple[RolloutResult, bool]: + if self.dynamic_sampling.enabled: + live_version = self._synchronize_version( + live_version, scored.prompts.device + ) + self._validate_policy_version(scored, live_version=live_version) + self._cache = scored + self._cache_key = cache_key + self._steps_since_rollout = 0 + return scored, True + + # A weight update cannot land between the final version check and + # cache publication. Reward scoring itself intentionally remains + # outside the policy lock because it may call an external service. + return self.generator.with_policy_snapshot(commit) + + cached = self._cache + assert cached is not None + + def reuse(live_version: int) -> Tuple[RolloutResult, bool]: + if self.dynamic_sampling.enabled: + live_version = self._synchronize_version( + live_version, cached.prompts.device + ) + self._validate_policy_version(cached, live_version=live_version) + return cached, False + + return self.generator.with_policy_snapshot(reuse) diff --git a/astrai/trainer/strategy.py b/astrai/trainer/strategy.py index 2321c7dc..a7fbe341 100644 --- a/astrai/trainer/strategy.py +++ b/astrai/trainer/strategy.py @@ -7,6 +7,7 @@ import torch.nn as nn import torch.nn.functional as F from torch import Tensor +from torch.optim import Optimizer from astrai.factory import BaseFactory from astrai.model.components.mlp import RouterStats @@ -279,10 +280,22 @@ def _refresh_moe_diagnostics( self._moe_metrics["aux_loss"] = float(aux_loss.detach().cpu().item()) def on_optimizer_step(self): - """Advance online rollout state after a successful optimizer step.""" + """Reject unsafe post-hoc publication for an online shared model.""" if self._rollout_runner is not None: - self._rollout_runner.update_weights(self.policy_version + 1) - self._rollout_runner.step() + raise RuntimeError( + "online training must call strategy.optimizer_step(optimizer) " + "so weight mutation and policy-version publication are atomic" + ) + + def optimizer_step(self, optimizer: Optimizer): + """Step the optimizer at an atomic online-rollout version boundary.""" + if self._rollout_runner is None: + return optimizer.step() + + next_version = self.policy_version + 1 + result = self._rollout_runner.apply_weight_update(next_version, optimizer.step) + self._rollout_runner.step() + return result def __call__(self, batch: Dict[str, Tensor]) -> LossOutput: """Run offline or online forward depending on runner injection.""" @@ -294,7 +307,17 @@ def __call__(self, batch: Dict[str, Tensor]) -> LossOutput: self._on_rollout_refresh() train_batch = self.prepare_from_rollout(result) - return self.compute_loss_output(train_batch) + output = self.compute_loss_output(train_batch) + if is_fresh: + output["metrics"].update( + { + f"dynamic_sampling/{name}": value + for name, value in getattr( + self._rollout_runner, "last_sampling_metrics", {} + ).items() + } + ) + return output class StrategyFactory(BaseFactory["BaseStrategy"]): diff --git a/astrai/trainer/train_callback.py b/astrai/trainer/train_callback.py index ba5c1efe..a1c8677a 100644 --- a/astrai/trainer/train_callback.py +++ b/astrai/trainer/train_callback.py @@ -164,6 +164,9 @@ def _save_checkpoint(self, context: TrainContext): **context.config.to_dict(), "optimizer_step": context.optimizer_step, } + policy_version = context.strategy.policy_version + if policy_version is not None: + meta["policy_version"] = policy_version context.checkpoint = Checkpoint( state_dict=state_dict, epoch=context.epoch, diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index ff6ac9a0..c1d81d47 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -30,7 +30,11 @@ ) from astrai.tokenize import AutoTokenizer from astrai.trainer.metric_util import GradSNRTracker -from astrai.trainer.rollout import RolloutGenerator, RolloutRunner +from astrai.trainer.rollout import ( + DynamicSamplingConfig, + RolloutGenerator, + RolloutRunner, +) from astrai.trainer.strategy import BaseStrategy, StrategyFactory logger = logging.getLogger(__name__) @@ -355,5 +359,26 @@ def _configure_rollout(self, context: TrainContext, strategy_kwargs: dict) -> No generator=generator, reward_model=cfg.reward_model_fn(), rollout_interval=cfg.rollout_interval, + max_policy_lag=cfg.rollout_max_policy_lag, + dynamic_sampling=DynamicSamplingConfig( + enabled=cfg.rollout_dynamic_sampling, + variance_threshold=cfg.rollout_dynamic_variance_threshold, + max_refill_rounds=cfg.rollout_dynamic_max_refill_rounds, + max_generated_tokens_per_group=( + cfg.rollout_dynamic_max_generated_tokens_per_group + ), + max_wall_time_per_group=( + cfg.rollout_dynamic_max_wall_time_per_group + ), + max_total_rollout_tokens_per_step=( + cfg.rollout_dynamic_max_total_tokens_per_step + ), + max_pending_groups=cfg.rollout_dynamic_max_pending_groups, + base_seed=( + cfg.random_seed + if cfg.rollout_dynamic_seed is None + else cfg.rollout_dynamic_seed + ), + ), ) ) diff --git a/astrai/trainer/trainer.py b/astrai/trainer/trainer.py index 09fcb3a2..69375740 100644 --- a/astrai/trainer/trainer.py +++ b/astrai/trainer/trainer.py @@ -94,8 +94,7 @@ def _trainer_loop(self, param_path: Optional[str] = None, resume: bool = False): if executor.sync_gradients: self._call_callbacks("before_optimizer_step", context) - context.optimizer.step() - context.strategy.on_optimizer_step() + context.strategy.optimizer_step(context.optimizer) context.optimizer.zero_grad() if context.scheduler: diff --git a/benchmarks/infraswe/README.md b/benchmarks/infraswe/README.md new file mode 100644 index 00000000..4bb2fce5 --- /dev/null +++ b/benchmarks/infraswe/README.md @@ -0,0 +1,63 @@ +# InfraSWE: async rollout policy-version consistency + +This directory binds AstrAI's online rollout consistency change to checked-in +NVIDIA L20 evidence. It uses InfraSWE's +`project-fit-system-path-v0.5.1` comparison and scoring models because the +candidate changes optimizer publication, scheduler locking, rollout scoring, +cache publication, and checkpoint lifecycle behavior rather than an isolated +kernel. + +Before the PR was opened, InfraSWE commit +`811bc775ed5b3a6ec853219245f3469f78818020` was used to validate the +`ProjectComparisonCell`, run the frozen ProjectFit and BenchmarkTrust scoring +functions, and execute the Draft/system-path test subsets. + +The visible-evidence diagnostic ProjectFit is **92.41/100** and BenchmarkTrust +is **97.40/100**. Official scoring remains unresolved because this evidence is +unsealed and lacks five candidate fresh-process replays, a system trace, hidden +probes, and a verified evidence manifest. The score is diagnostic, not a +certification. + +## L20 result + +The reproducible probe forces one policy-version advance while the reward model +scores every real one-token CUDA rollout. Both revisions use seed 3407, five +warmups, and 50 measured trials on GPU5. + +| Revision | Stale accepted | Stale rejected | Median | p99 | +| --- | ---: | ---: | ---: | ---: | +| `ce2f9d1` baseline | 50/50 | 0/50 | 3.2959 ms | 5.0569 ms | +| `3483c3a` candidate | 0/50 | 50/50 | 3.3122 ms | 4.4414 ms | + +The candidate eliminates observed stale acceptance. Median trial latency moves +by **+0.50%**, within the declared 2% ceiling, while p99 moves by **-12.17%**. +Raw values are stored in +`benchmarks/results/async_rollout_version_l20_sm89.json`. + +## Coverage + +- Optimizer mutation and policy-version publication share the generation lock. +- Direct scheduler updates cannot enter during a generator policy snapshot. +- Future versions and results beyond `rollout_max_policy_lag` fail explicitly. +- Results are revalidated after reward scoring and while publishing or reusing + the rollout cache. +- Online checkpoints persist the actual rollout policy version. +- The contract is exercised through both online GRPO and online DPO. + +The complete local suite passed 641 tests with 171 environment-dependent skips; +the focused L20 suite passed 91 tests. This is a deterministic, single-process +race replay, not a long-running multiprocess or external reward-service soak. + +## Digest construction + +- target profile: sorted SHA-256 list for baseline `README.md` and + `pyproject.toml`; +- baseline: SHA-256 of target commit + `ce2f9d13b32f729c561d0175fd46927a37d9b0a2`; +- candidate: SHA-256 of the sorted per-file SHA-256 list for implementation, + tests, documentation, and the benchmark tool in commit `3483c3a`; +- acceptance: corresponding scheduler, rollout, online-strategy, callback, and + online end-to-end tests; +- probe/workload: benchmark tool and checked-in raw L20 result; and +- required deployment cell: literal + `nvidia-l20-sm89-single-gpu-cuda12.8-gpu5`. diff --git a/benchmarks/infraswe/astrai-async-rollout-version-comparison-cell.json b/benchmarks/infraswe/astrai-async-rollout-version-comparison-cell.json new file mode 100644 index 00000000..0dcaf4e1 --- /dev/null +++ b/benchmarks/infraswe/astrai-async-rollout-version-comparison-cell.json @@ -0,0 +1,16 @@ +{ + "schema_version": "0.5", + "target_project_profile_sha256": "sha256:4aff21e73bad3f0a5ca6b5418f25389974c0e1a65443a42dbdf2b5d002cca822", + "target_repository_or_baseline_sha256": "sha256:a3efb47f11f8eecfc129fd2f10667d13de6def124ecaa52df5dedf8f00bab417", + "change_intent": "integrate", + "semantic_contract_sha256": "sha256:f4f5150ba2c9757f4f7ec643d82501ddf074b12e3774d7bc512d699b55751f66", + "acceptance_contract_sha256": "sha256:80a05e8e5c7939dac46b4910e8ea6b92aa699f98c351cfeb80b604d27f987fed", + "probe_set_sha256": "sha256:d7b94a91dea876a53f35219fbd970e4371ff3753b28ad87a27a6c9c0a75dc788", + "workload_portfolio_sha256": "sha256:f90988275bf606cd7eb3c652a1f7e01d203becd12e7b811843da14d0081990cb", + "performance_target_sha256": "sha256:91b598db78e5bfe438d6070a87810072435950447cc7b25f1132aec0394cd934", + "required_deployment_cell_set_sha256": "sha256:8a16c180aad79848a69d189324e60012a220617579b0ff9c2e916d95cd733970", + "formula_template_id": "project-fit-system-path-v0.5.1", + "evidence_policy_id": "system-path-evidence-v1", + "project_season": "astrai-2026q3", + "cross_project_ranking_allowed": false +} diff --git a/benchmarks/infraswe/astrai-online-rl-target-profile-v0.5.json b/benchmarks/infraswe/astrai-online-rl-target-profile-v0.5.json new file mode 100644 index 00000000..acf9cbac --- /dev/null +++ b/benchmarks/infraswe/astrai-online-rl-target-profile-v0.5.json @@ -0,0 +1,98 @@ +{ + "schema_version": "0.5", + "id": "astrai-online-rl-training-v1", + "version": "astrai-2026q3-proposed.1", + "status": "proposed", + "repository": "https://github.com/ViperEkura/AstrAI", + "supported_revision_policy": "pinned:f87ac4592d25cd98096d3a628e2747f668bc4d96; update by pull request after the rollout-version prerequisite changes", + "ownership": { + "maintainers": [ + "UNASSIGNED:ViperEkura/AstrAI-maintainer" + ], + "review_required": 1, + "last_reviewed_at": null, + "update_policy": "pull-request" + }, + "component_ownership": { + "online-rl-rollout": [ + "upstream-codeowner-review-required" + ], + "training-configuration": [ + "upstream-codeowner-review-required" + ] + }, + "allowed_integration_points": [ + "astrai/trainer/rollout.py", + "astrai/trainer/train_context.py", + "astrai/trainer/strategy.py", + "astrai/config/train_config.py", + "scripts/tools/train.py", + "tests/trainer" + ], + "api_abi_contract": { + "id": "astrai-online-rl-training-v1:api-abi", + "sha256": "sha256:74a63c01657124356fd4ee43f4a9ab30f068db5bf784e866c8b88c2148c36371", + "path": "astrai/trainer/rollout.py" + }, + "lifecycle_contract": { + "id": "astrai-online-rl-training-v1:lifecycle", + "sha256": "sha256:74a63c01657124356fd4ee43f4a9ab30f068db5bf784e866c8b88c2148c36371", + "path": "astrai/trainer/rollout.py" + }, + "build_test_matrix": { + "id": "astrai-online-rl-training-v1:build-test-matrix", + "sha256": "sha256:9ada3113b92f5fc46b610bb5b539344ba971479e658c175785aec2371ddc3974", + "path": "pyproject.toml" + }, + "dependency_policy": { + "id": "astrai-online-rl-training-v1:dependency-policy", + "sha256": "sha256:9ada3113b92f5fc46b610bb5b539344ba971479e658c175785aec2371ddc3974", + "path": "pyproject.toml" + }, + "fallback_policy": { + "id": "astrai-online-rl-training-v1:fallback-policy", + "sha256": "sha256:74a63c01657124356fd4ee43f4a9ab30f068db5bf784e866c8b88c2148c36371", + "path": "astrai/trainer/rollout.py" + }, + "deployment_workload_portfolio": { + "id": "astrai-online-rl-training-v1:deployment-workload-portfolio", + "sha256": "sha256:40e8f77c15bd2dc3cb7759edcdca02531330ab64197f64c1fdedc49ee4b7f70e", + "path": "benchmarks/training_consistency/benchmark_dynamic_sampling.py" + }, + "performance_acceptance_targets": { + "id": "astrai-online-rl-training-v1:performance-acceptance-targets", + "sha256": "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28", + "path": "benchmarks/results/versioned_dynamic_sampling_l20_sm89.json" + }, + "maintainability_probes": { + "id": "astrai-online-rl-training-v1:maintainability-probes", + "sha256": "sha256:c4d28a2be82cdb56674e179d3ed55799d4f28730f0d2ffc567a347b57d33ea9e", + "path": "tests/trainer/test_rollout.py" + }, + "project_objectives": { + "edge_ecosystem": { + "owner": "UNASSIGNED:ViperEkura/AstrAI-maintainer", + "policy": "roadmap", + "profiles": [ + { + "id": "cuda-sm89-l20", + "status": "active", + "target_release": null, + "required_for_release": true + }, + { + "id": "cpu", + "status": "planned", + "target_release": null, + "required_for_release": false + } + ] + } + }, + "scoring_template_id": "project-fit-kernel-v0.5", + "triton_portability": { + "enabled": false, + "owner": null, + "profile_weights": {} + } +} diff --git a/benchmarks/infraswe/astrai-versioned-dynamic-sampling-comparison-cell.json b/benchmarks/infraswe/astrai-versioned-dynamic-sampling-comparison-cell.json new file mode 100644 index 00000000..77d600a5 --- /dev/null +++ b/benchmarks/infraswe/astrai-versioned-dynamic-sampling-comparison-cell.json @@ -0,0 +1,16 @@ +{ + "schema_version": "0.5", + "target_project_profile_sha256": "sha256:96613aff5466a56f69f3cbeb6216a70e736f41a46b615fc4fe6484ef96770d58", + "target_repository_or_baseline_sha256": "sha256:d5960c9b0e25c96a2020771fa605f0ff671689791340993cd0850193062e7c5f", + "change_intent": "integrate-versioned-dynamic-sampling", + "semantic_contract_sha256": "sha256:74a63c01657124356fd4ee43f4a9ab30f068db5bf784e866c8b88c2148c36371", + "acceptance_contract_sha256": "sha256:c4d28a2be82cdb56674e179d3ed55799d4f28730f0d2ffc567a347b57d33ea9e", + "probe_set_sha256": "sha256:c4d28a2be82cdb56674e179d3ed55799d4f28730f0d2ffc567a347b57d33ea9e", + "workload_portfolio_sha256": "sha256:40e8f77c15bd2dc3cb7759edcdca02531330ab64197f64c1fdedc49ee4b7f70e", + "performance_target_sha256": "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28", + "required_deployment_cell_set_sha256": "sha256:a69c87538f6b2dc066db721d3ee877b995466ef3d5518e359650f7bd8377f2c3", + "formula_template_id": "project-fit-kernel-v0.5", + "evidence_policy_id": "infraswe-v0.5-diagnostic-pre-pr", + "project_season": "astrai-2026q3", + "cross_project_ranking_allowed": false +} diff --git a/benchmarks/infraswe/astrai-versioned-dynamic-sampling-draft-v0.5.json b/benchmarks/infraswe/astrai-versioned-dynamic-sampling-draft-v0.5.json new file mode 100644 index 00000000..70bb147e --- /dev/null +++ b/benchmarks/infraswe/astrai-versioned-dynamic-sampling-draft-v0.5.json @@ -0,0 +1,44 @@ +{ + "schema_version": "0.5", + "draft": { + "id": "astrai-versioned-dynamic-sampling-2026q3", + "revision": 1, + "state": "D1-target-bound", + "created_by": "0z5a" + }, + "target": { + "mode": "repository", + "repository": "https://github.com/ViperEkura/AstrAI", + "revision": "sha256:d5960c9b0e25c96a2020771fa605f0ff671689791340993cd0850193062e7c5f", + "project_profile_sha256": "sha256:96613aff5466a56f69f3cbeb6216a70e736f41a46b615fc4fe6484ef96770d58", + "catalog_profile": null + }, + "candidate": { + "kind": "git-diff", + "revision": "sha256:61f07e320490749eeeef0efb87948b3e7dc3fd95a0b34ab354b34e770ab33b0a", + "intent": "integrate", + "implementation_kind": "framework", + "entrypoints": [ + "astrai.trainer.rollout.DynamicSamplingGroup", + "astrai.trainer.rollout.RolloutRunner.run", + "astrai.trainer.train_context.TrainContext", + "scripts.tools.train" + ], + "operator_family": "generic", + "phase": "training", + "backend": "cuda", + "primary_host_candidate": "AstrAI" + }, + "baseline": { + "mode": "target-head", + "revision": "sha256:d5960c9b0e25c96a2020771fa605f0ff671689791340993cd0850193062e7c5f", + "advisory_reference_profile": null + }, + "default_candidates": null, + "deployment": null, + "retrieval": null, + "acceptance_contract": null, + "project_objectives": null, + "benchmark_loop": null, + "scoring": null +} diff --git a/benchmarks/results/async_rollout_version_l20_infraswe_score.json b/benchmarks/results/async_rollout_version_l20_infraswe_score.json new file mode 100644 index 00000000..d37721ea --- /dev/null +++ b/benchmarks/results/async_rollout_version_l20_infraswe_score.json @@ -0,0 +1,111 @@ +{ + "schema_version": "0.5.1", + "score_kind": "diagnostic-project-fit", + "score_is_official": false, + "candidate_revision": "sha256:35c11e2547082bf30daf9d39feaa6e1336115dba5abf668ed103e8b1453123b8", + "formula_template_id": "project-fit-system-path-v0.5.1", + "diagnostic_project_fit_100": 92.40836866232375, + "component_values": { + "evolutionary_maintainability": 0.89117283614121, + "project_contract_fit": 0.9740037464252967, + "performance_reuse_utilization": 0.9451290407777138, + "operational_fit": 0.87217006329394 + }, + "component_floors": { + "evolutionary_maintainability": 0.6, + "project_contract_fit": 0.6, + "performance_reuse_utilization": 0.4, + "operational_fit": 0.6 + }, + "subcomponent_inputs": { + "evolutionary_maintainability": { + "evolution": 0.82, + "locality": 0.85, + "tests": 1.0, + "failure": 1.0, + "contract": 0.95 + }, + "project_contract_fit": { + "integration": 1.0, + "interface": 0.9, + "lifecycle": 1.0, + "buildtest": 1.0, + "policy": 1.0 + }, + "performance_reuse_utilization": { + "attainment": 1.0, + "coverage": 0.85, + "retention": 1.0, + "family": 0.9, + "compile": 1.0 + }, + "operational_fit": { + "replay": 0.85, + "load": 0.8, + "resource": 1.0, + "coldsteady": 0.9 + } + }, + "input_rationale": { + "evolution": "The change adds an explicit rollout-version contract but has no upstream maintenance history yet.", + "locality": "The implementation follows existing scheduler, rollout, strategy, trainer, configuration, and checkpoint boundaries with focused tests.", + "tests": "The complete local suite passed 641 tests with 171 environment-dependent skips; 91 focused tests passed on NVIDIA L20 GPU5.", + "failure": "Non-monotonic updates, queued updates, future rollouts, excessive lag, and unsafe post-hoc online optimizer publication fail explicitly.", + "contract": "Implementation, tests, documentation, benchmark tooling, and raw L20 output are digest-bound, but the artifact is not sealed or maintainer-reviewed.", + "integration": "The serialized boundary spans InferenceScheduler, RolloutGenerator, RolloutRunner, BaseStrategy, Trainer, TrainConfig, CLI construction, and checkpoints.", + "interface": "Additive apply_weight_update and rollout_max_policy_lag interfaces preserve offline training and the default rollout reuse window.", + "lifecycle": "Optimizer mutation and version publication share the generation lock; validation runs after generation, after reward scoring, and during cache publication or reuse.", + "buildtest": "The full local suite, Ruff formatting/import checks, InfraSWE engine tests, focused L20 suite, and L20 race replay passed.", + "policy": "No dependency is added; policy lag defaults to rollout_interval minus one for compatibility and strict zero-lag is opt-in.", + "attainment": "The baseline accepted all 50 forced-stale rollouts while the candidate rejected all 50; median latency changed by only 0.50% against a 2% ceiling.", + "coverage": "Generation, scoring, cache publication/reuse, optimizer publication, direct scheduler updates, checkpoints, GRPO, and DPO are covered; a multiprocess soak is not.", + "retention": "The complete regression suite passed after the final concurrency and checkpoint fixes.", + "family": "The contract is shared by online GRPO and online DPO and serializes wrapper-level and direct scheduler updates.", + "compile": "The Python lifecycle change introduces no compilation step and does not compile during measured trials.", + "replay": "Each revision used five warmups and 50 trials, but only one fresh candidate process was recorded.", + "load": "The probe uses a real CUDA rollout and forced scoring-time policy advances but is not a sustained multiprocess training workload.", + "resource": "Only GPU5's remaining capacity was used; existing GPU processes and containers were left untouched.", + "coldsteady": "Warm trial latency is reported; cold startup and a sustained external reward-service soak are deliberately excluded." + }, + "benchmark_trust": { + "formula_version": "benchmark-trust-v0.5", + "status": "scored", + "score_100": 97.40037464252967, + "components": { + "reproducibility": 1.0, + "evidence": 1.0, + "statistics": 0.9, + "environment": 1.0 + }, + "failure_codes": [ + "DRAFT_UNSEALED", + "FRESH_PROCESS_REPLAY_INCOMPLETE", + "SYSTEM_TRACE_EVIDENCE_MISSING", + "HIDDEN_PROBES_INCOMPLETE", + "EVIDENCE_MANIFEST_UNVERIFIED", + "ASYNC_MULTIPROCESS_SOAK_UNTESTED" + ] + }, + "official_project_fit": { + "status": "unresolved", + "score_100": null, + "failure_codes": [ + "DRAFT_SEAL_MISSING", + "FRESH_PROCESS_REPLAYS_BELOW_MINIMUM", + "SYSTEM_TRACE_EVIDENCE_MISSING", + "HIDDEN_PROBES_INCOMPLETE", + "EVIDENCE_MANIFEST_UNVERIFIED" + ] + }, + "comparison_cell_path": "benchmarks/infraswe/astrai-async-rollout-version-comparison-cell.json", + "execution": { + "infraswe_commit": "811bc775ed5b3a6ec853219245f3469f78818020", + "comparison_cell_validation": "pass", + "infraswe_engine_tests": "41 passed", + "astrai_implementation_commit": "3483c3a578832c9c2e4e6f93b3a16c700c7600e1", + "astrai_local_tests": "641 passed, 171 skipped", + "astrai_l20_focused_tests": "91 passed in 11.87s", + "astrai_l20_race_replay": "baseline accepted 50/50 stale rollouts; candidate rejected 50/50", + "astrai_lint": "passed" + } +} diff --git a/benchmarks/results/async_rollout_version_l20_sm89.json b/benchmarks/results/async_rollout_version_l20_sm89.json new file mode 100644 index 00000000..1807f03f --- /dev/null +++ b/benchmarks/results/async_rollout_version_l20_sm89.json @@ -0,0 +1,71 @@ +{ + "schema_version": "1.0", + "candidate_source_commit": "3483c3a578832c9c2e4e6f93b3a16c700c7600e1", + "baseline_commit": "ce2f9d13b32f729c561d0175fd46927a37d9b0a2", + "recorded_at": "2026-09-03T09:17:00+08:00", + "workload": "policy version advances once during reward scoring of each real one-token CUDA rollout", + "gpu": "NVIDIA L20", + "gpu_index": 5, + "torch": "2.11.0+cu128", + "cuda": "12.8", + "seed": 3407, + "warmup": 5, + "trials": 50, + "baseline": { + "accepted_stale_rollouts": 50, + "rejected_stale_rollouts": 0, + "median_trial_ms": 3.2958545, + "p99_trial_ms": 5.056941, + "final_policy_version": 55 + }, + "candidate": { + "accepted_stale_rollouts": 0, + "rejected_stale_rollouts": 50, + "median_trial_ms": 3.3121875, + "p99_trial_ms": 4.441436, + "final_policy_version": 55 + }, + "candidate_vs_baseline": { + "median_delta_percent": 0.49556192483617423, + "p99_delta_percent": -12.171488652922779 + }, + "integrated_ddp_soak": { + "status": "pass", + "recorded_at": "2026-09-03T10:59:43+08:00", + "integration_revision": "fe81f17772db7acb8d7575b3f617a454a83ca58d", + "scope": "PR #59 rollout-version fencing combined with PR #55 DDP rollout path", + "world_sizes": [1, 2, 3], + "runs": [ + { + "world_size": 1, + "steps": 200, + "stale_reject_total": 20, + "future_version_reject_total": 10, + "version_mismatch_total": 0 + }, + { + "world_size": 2, + "steps": 400, + "stale_reject_total": 80, + "future_version_reject_total": 20, + "version_mismatch_total": 0 + }, + { + "world_size": 3, + "steps": 1000, + "stale_reject_total": 300, + "future_version_reject_total": 75, + "version_mismatch_total": 0 + } + ], + "long_soak": { + "world_size": 3, + "steps": 100000, + "elapsed_seconds": 2486.142, + "allocated_memory_drift_mib": 0.0, + "reserved_memory_drift_mib": 0.0, + "parameter_digest_mismatches": 0 + } + }, + "scope_limit": "The deterministic stale-rollout comparison is single-process; the additional DDP runs validate cross-rank integration and rejection contracts rather than an external reward service." +} diff --git a/benchmarks/results/versioned_dynamic_sampling_l20_infraswe_score.json b/benchmarks/results/versioned_dynamic_sampling_l20_infraswe_score.json new file mode 100644 index 00000000..196d993c --- /dev/null +++ b/benchmarks/results/versioned_dynamic_sampling_l20_infraswe_score.json @@ -0,0 +1,164 @@ +{ + "schema_version": "0.5-diagnostic", + "infraswe_revision": "191b9099f6d634b65998dac5604aeefc433673e7", + "scope": "pre-PR diagnostic only; not an official InfraSWE or leaderboard score", + "official_score_published": false, + "official_status": "unresolved", + "official_blockers": [ + "DRAFT_STATE_D1_NOT_SEALED", + "HIDDEN_PROBES_INCOMPLETE", + "EVIDENCE_MANIFEST_UNVERIFIED", + "FRESH_PROCESS_REPLAYS_BELOW_MINIMUM" + ], + "draft_path": "benchmarks/infraswe/astrai-versioned-dynamic-sampling-draft-v0.5.json", + "target_profile_path": "benchmarks/infraswe/astrai-online-rl-target-profile-v0.5.json", + "comparison_cell_path": "benchmarks/infraswe/astrai-versioned-dynamic-sampling-comparison-cell.json", + "raw_diagnostic_inputs": { + "evolutionary_maintainability": { + "evolution": 0.82, + "locality": 0.68, + "tests": 0.97, + "failure": 0.98, + "contract": 0.95 + }, + "project_contract_fit": { + "integration": 0.95, + "interface": 0.92, + "lifecycle": 0.98, + "buildtest": 0.97, + "policy": 0.96 + }, + "performance_reuse_utilization": { + "attainment": 0.92, + "coverage": 0.88, + "retention": 1.0, + "family": 0.65, + "compile": 1.0 + }, + "operational_fit": { + "replay": 0.94, + "load": 0.80, + "resource": 0.98, + "coldsteady": 0.85 + } + }, + "input_rationale": { + "locality": "Conservative because the change touches the central rollout path and 12 files.", + "family": "Conservative because the feature is intentionally limited to online_grpo.", + "load": "Conservative because real L20 runs use a tiny policy and synthetic rewards.", + "attainment": "Efficiency improves 33.37%, while median latency and generated tokens rise 55.80% and 49.95%.", + "retention": "The feature is disabled by default and baseline behavior remains covered.", + "failure": "Three-rank generation and scoring failure propagation both passed." + }, + "project_fit": { + "status": "provisional", + "formula_template_id": "project-fit-kernel-v0.5", + "score_100": 88.36622760415803, + "components": { + "evolutionary_maintainability": { + "status": "scored", + "value": 0.8360098443539299, + "formula_version": "evolutionary-maintainability-v0.5", + "input_evidence_digests": [ + "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28" + ], + "confidence": "high", + "failure_codes": [], + "reason": null + }, + "project_contract_fit": { + "status": "scored", + "value": 0.9522525339731504, + "formula_version": "project-contract-fit-v0.5", + "input_evidence_digests": [ + "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28" + ], + "confidence": "high", + "failure_codes": [], + "reason": null + }, + "performance_reuse_utilization": { + "status": "scored", + "value": 0.8818270387289837, + "formula_version": "performance-reuse-utilization-v0.5", + "input_evidence_digests": [ + "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28" + ], + "confidence": "high", + "failure_codes": [], + "reason": null + }, + "operational_fit": { + "status": "scored", + "value": 0.8851040999117927, + "formula_version": "operational-fit-v0.5", + "input_evidence_digests": [ + "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28" + ], + "confidence": "high", + "failure_codes": [], + "reason": null + }, + "pure_triton_portability": { + "status": "not_applicable", + "value": null, + "formula_version": "pure-triton-portability-v0.5", + "input_evidence_digests": [], + "confidence": "not_applicable", + "failure_codes": [], + "reason": "ordinary candidates do not receive an edge portability score" + } + }, + "component_floors": { + "evolutionary_maintainability": 0.6, + "project_contract_fit": 0.6, + "performance_reuse_utilization": 0.4, + "operational_fit": 0.6 + }, + "confidence": "low", + "comparison_cell": { + "schema_version": "0.5", + "target_project_profile_sha256": "sha256:96613aff5466a56f69f3cbeb6216a70e736f41a46b615fc4fe6484ef96770d58", + "target_repository_or_baseline_sha256": "sha256:d5960c9b0e25c96a2020771fa605f0ff671689791340993cd0850193062e7c5f", + "change_intent": "integrate-versioned-dynamic-sampling", + "semantic_contract_sha256": "sha256:74a63c01657124356fd4ee43f4a9ab30f068db5bf784e866c8b88c2148c36371", + "acceptance_contract_sha256": "sha256:c4d28a2be82cdb56674e179d3ed55799d4f28730f0d2ffc567a347b57d33ea9e", + "probe_set_sha256": "sha256:c4d28a2be82cdb56674e179d3ed55799d4f28730f0d2ffc567a347b57d33ea9e", + "workload_portfolio_sha256": "sha256:40e8f77c15bd2dc3cb7759edcdca02531330ab64197f64c1fdedc49ee4b7f70e", + "performance_target_sha256": "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28", + "required_deployment_cell_set_sha256": "sha256:a69c87538f6b2dc066db721d3ee877b995466ef3d5518e359650f7bd8377f2c3", + "formula_template_id": "project-fit-kernel-v0.5", + "evidence_policy_id": "infraswe-v0.5-diagnostic-pre-pr", + "project_season": "astrai-2026q3", + "cross_project_ranking_allowed": false + }, + "cross_project_ranking_allowed": false, + "failure_codes": [], + "audit_flags": [] + }, + "benchmark_trust": { + "status": "scored", + "score_100": 93.42120700831886, + "components": { + "reproducibility": 0.96, + "evidence": 0.93, + "statistics": 0.90, + "environment": 0.94 + }, + "formula_version": "benchmark-trust-v0.5", + "evidence_digests": [ + "sha256:ec46813c89c32393c770e1739538e695872510eb3e2d0215d000ca398a68be28" + ], + "failure_codes": [] + }, + "validation": { + "infraswe_full_suite": "283 passed", + "infraswe_minimum_training_suite": "pass; fixture_only=true; official_score_published=false", + "astrai_full_suite": "650 passed, 171 skipped", + "astrai_l20_targeted": "42 passed", + "single_gpu_trials": 50, + "three_gpu_steps": 100, + "mixed_version_batches": 0, + "incomplete_batches": 0 + } +} diff --git a/benchmarks/results/versioned_dynamic_sampling_l20_sm89.json b/benchmarks/results/versioned_dynamic_sampling_l20_sm89.json new file mode 100644 index 00000000..e330eb10 --- /dev/null +++ b/benchmarks/results/versioned_dynamic_sampling_l20_sm89.json @@ -0,0 +1,84 @@ +{ + "schema_version": "1.0", + "candidate_source_commit": "392d53d3dca99e78c8d5d5b461d6f108a802e28a", + "base_commit": "f87ac4592d25cd98096d3a628e2747f668bc4d96", + "recorded_at": "2026-09-03T11:48:00+08:00", + "hardware": { + "gpu": "NVIDIA L20", + "compute_capability": "sm89", + "torch": "2.11.0+cu128", + "cuda": "12.8" + }, + "single_gpu": { + "gpu_index": 0, + "seed": 3407, + "prompts": 8, + "group_size": 4, + "max_tokens": 16, + "warmup": 5, + "trials": 50, + "workload": "half of first-attempt prompt groups have zero reward variance; refills become non-degenerate", + "baseline": { + "median_latency_ms": 248.517238, + "p95_latency_ms": 256.624277, + "generated_tokens": 24657, + "accepted_groups": 200, + "effective_groups_per_million_tokens": 8111.28685565965, + "mean_rollout_waste_ratio": 0.0 + }, + "dynamic": { + "median_latency_ms": 387.1997145, + "p95_latency_ms": 396.259705, + "generated_tokens": 36974, + "accepted_groups": 400, + "effective_groups_per_million_tokens": 10818.412938821875, + "mean_rollout_waste_ratio": 0.33227127734037476 + }, + "dynamic_vs_baseline": { + "effective_groups_per_million_tokens_delta_percent": 33.374803916265485, + "median_latency_delta_percent": 55.803966604521825, + "generated_tokens_delta_percent": 49.95336010057996 + }, + "policy_version_jitter": { + "final_policy_version": 1, + "behavior_policy_versions": [1], + "mixed_version_groups": false, + "version_invalidated_groups": 8 + } + }, + "three_gpu_rank_consistency": { + "gpu_indices": [0, 1, 2], + "backend": "nccl", + "world_size": 3, + "steps": 100, + "warmup": 10, + "jitter_interval": 10, + "prompts": 2, + "group_size": 2, + "max_tokens": 4, + "workload": "rank-skewed reward variance forces collective refill decisions; policy version advances every tenth measured step", + "max_rank_latency_median_ms": 35.1143005, + "max_rank_latency_p95_ms": 52.265118, + "mixed_version_batches": 0, + "incomplete_batches": 0, + "generation_schedule_mismatches": 0, + "version_invalidated_groups": 60, + "rank_local_failure_propagation": { + "generation": true, + "scoring": true + }, + "rank0_memory": { + "allocated_drift_bytes": 4096, + "reserved_drift_bytes": 0, + "max_allocated_bytes": 8876544, + "max_reserved_bytes": 23068672 + } + }, + "validation": { + "local_test_suite": "650 passed, 171 skipped", + "l20_targeted_tests": "42 passed", + "format": "ruff format --check: pass", + "import_order": "ruff check --select I: pass" + }, + "scope_limit": "This benchmark measures rollout-buffer correctness and refill efficiency with a tiny real CUDA policy plus deterministic synthetic rewards; it does not claim downstream reward convergence." +} diff --git a/benchmarks/training_consistency/benchmark_async_rollout_versions.py b/benchmarks/training_consistency/benchmark_async_rollout_versions.py new file mode 100644 index 00000000..a00a4fb0 --- /dev/null +++ b/benchmarks/training_consistency/benchmark_async_rollout_versions.py @@ -0,0 +1,121 @@ +"""Replay a policy update during reward scoring on a real rollout backend.""" + +import argparse +import inspect +import json +import statistics +import time + +import torch + +from astrai.inference.scheduler import InferenceScheduler +from astrai.trainer.rollout import BaseRewardModel, RolloutGenerator, RolloutRunner +from tests.helpers import FakeTokenizer, make_model + + +class AdvancingRewardModel(BaseRewardModel): + """Advance the visible policy version while the rollout is being scored.""" + + def __init__(self): + self.runner = None + + def score(self, prompts, responses): + assert self.runner is not None + self.runner.update_weights(self.runner.policy_version + 1) + return torch.zeros(len(prompts), len(responses[0])) + + +def percentile(samples, fraction): + ordered = sorted(samples) + index = min(len(ordered) - 1, int(len(ordered) * fraction)) + return ordered[index] + + +def run(device, trials, warmup): + torch.manual_seed(3407) + model, _ = make_model(device, max_position_embeddings=64) + tokenizer = FakeTokenizer(with_chat_template=True) + scheduler = InferenceScheduler( + model=model, + tokenizer=tokenizer, + max_batch_size=1, + max_seq_len=64, + device=device, + enable_cuda_graph=False, + ) + generator = RolloutGenerator( + scheduler=scheduler, + tokenizer=tokenizer, + max_tokens=1, + group_size=1, + temperature=1.0, + top_k=0, + top_p=1.0, + ) + reward = AdvancingRewardModel() + supports_lag_guard = "max_policy_lag" in inspect.signature(RolloutRunner).parameters + runner_kwargs = {"max_policy_lag": 0} if supports_lag_guard else {} + runner = RolloutRunner( + generator=generator, + reward_model=reward, + rollout_interval=1, + **runner_kwargs, + ) + reward.runner = runner + batch = {"instruction": ["Reply briefly"], "input": ["Hi"]} + + accepted_stale = 0 + rejected_stale = 0 + samples_ms = [] + total = warmup + trials + for index in range(total): + runner.clear_cache() + started = time.perf_counter_ns() + try: + result, _ = runner(batch) + except RuntimeError as exc: + if "policy lag" not in str(exc): + raise + rejected_stale += index >= warmup + else: + accepted_stale += index >= warmup and ( + result.policy_version < runner.policy_version + ) + if device.startswith("cuda"): + torch.cuda.synchronize(device) + elapsed_ms = (time.perf_counter_ns() - started) / 1e6 + if index >= warmup: + samples_ms.append(elapsed_ms) + + return { + "revision_mode": "candidate" if supports_lag_guard else "baseline", + "device": str(device), + "gpu": torch.cuda.get_device_name(device) + if device.startswith("cuda") + else None, + "torch": torch.__version__, + "cuda": torch.version.cuda, + "seed": 3407, + "warmup": warmup, + "trials": trials, + "accepted_stale_rollouts": accepted_stale, + "rejected_stale_rollouts": rejected_stale, + "median_trial_ms": statistics.median(samples_ms), + "p99_trial_ms": percentile(samples_ms, 0.99), + "final_policy_version": runner.policy_version, + } + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--device", default="cuda") + parser.add_argument("--trials", type=int, default=50) + parser.add_argument("--warmup", type=int, default=5) + args = parser.parse_args() + if args.trials <= 0 or args.warmup < 0: + parser.error("trials must be positive and warmup must be non-negative") + print(json.dumps(run(args.device, args.trials, args.warmup), indent=2)) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/training_consistency/benchmark_dynamic_sampling.py b/benchmarks/training_consistency/benchmark_dynamic_sampling.py new file mode 100644 index 00000000..3d484268 --- /dev/null +++ b/benchmarks/training_consistency/benchmark_dynamic_sampling.py @@ -0,0 +1,206 @@ +"""Benchmark versioned low-variance group refill on a real rollout backend.""" + +import argparse +import json +import statistics +import time + +import torch + +from astrai.inference.scheduler import InferenceScheduler +from astrai.trainer.rollout import ( + BaseRewardModel, + DynamicSamplingConfig, + RolloutGenerator, + RolloutRunner, +) +from tests.helpers import FakeTokenizer, make_model + + +class PatternRewardModel(BaseRewardModel): + """Make odd prompt groups degenerate once, then accept their refill.""" + + def __init__(self): + self.seen = {} + self.runner = None + self.advance_once = False + self._advanced = False + + def reset(self, *, advance_once=False): + self.seen.clear() + self.advance_once = advance_once + self._advanced = False + + def score(self, prompts, responses): + rewards = torch.zeros(len(prompts), len(responses[0])) + for row, prompt in enumerate(prompts): + seen = self.seen.get(prompt, 0) + self.seen[prompt] = seen + 1 + ordinal = int(prompt.split("topic ", 1)[1].splitlines()[0]) + if ordinal % 2 == 0 or seen > 0: + rewards[row] = torch.arange(rewards.size(1), dtype=torch.float32) + else: + rewards[row] = 1.0 + if self.advance_once and not self._advanced: + assert self.runner is not None + self.runner.update_weights(self.runner.policy_version + 1) + self._advanced = True + return rewards + + +def percentile(samples, fraction): + ordered = sorted(samples) + index = min(len(ordered) - 1, int(len(ordered) * fraction)) + return ordered[index] + + +def make_runner(device, prompts, group_size, max_tokens, *, dynamic): + model, _ = make_model(device, max_position_embeddings=128) + tokenizer = FakeTokenizer(with_chat_template=True) + scheduler = InferenceScheduler( + model=model, + tokenizer=tokenizer, + max_batch_size=prompts * group_size, + max_seq_len=128, + device=device, + enable_cuda_graph=False, + ) + generator = RolloutGenerator( + scheduler=scheduler, + tokenizer=tokenizer, + max_tokens=max_tokens, + group_size=group_size, + temperature=1.0, + top_k=0, + top_p=1.0, + ) + reward = PatternRewardModel() + runner = RolloutRunner( + generator=generator, + reward_model=reward, + rollout_interval=1, + max_policy_lag=1, + dynamic_sampling=DynamicSamplingConfig( + enabled=dynamic, + variance_threshold=0.0, + max_refill_rounds=2, + max_generated_tokens_per_group=group_size * max_tokens * 4, + max_total_rollout_tokens_per_step=prompts * group_size * max_tokens * 4, + max_pending_groups=prompts, + base_seed=3407, + ), + ) + reward.runner = runner + return runner, reward + + +def measure(runner, reward, batch, *, trials, warmup): + latency_ms = [] + generated_tokens = [] + accepted_groups = [] + waste_ratio = [] + for trial in range(warmup + trials): + reward.reset() + runner.clear_cache() + started = time.perf_counter_ns() + result, _ = runner(batch) + if result.responses.device.type == "cuda": + torch.cuda.synchronize(result.responses.device) + elapsed_ms = (time.perf_counter_ns() - started) / 1e6 + variances = result.rewards.float().var(dim=1, unbiased=False) + metrics = runner.last_sampling_metrics + if trial >= warmup: + latency_ms.append(elapsed_ms) + generated_tokens.append( + int(metrics.get("total_generated_tokens", result.response_mask.sum())) + ) + accepted_groups.append( + int(metrics.get("groups_accepted", (variances > 0).sum())) + ) + waste_ratio.append(float(metrics.get("rollout_waste_ratio", 0.0))) + total_tokens = sum(generated_tokens) + return { + "median_latency_ms": statistics.median(latency_ms), + "p95_latency_ms": percentile(latency_ms, 0.95), + "generated_tokens": total_tokens, + "accepted_groups": sum(accepted_groups), + "effective_groups_per_million_tokens": ( + sum(accepted_groups) * 1_000_000 / total_tokens + ), + "mean_rollout_waste_ratio": statistics.mean(waste_ratio), + } + + +def run(device, prompts, group_size, max_tokens, trials, warmup): + torch.manual_seed(3407) + batch = { + "instruction": [f"Reply about topic {index}" for index in range(prompts)], + "input": ["briefly"] * prompts, + } + baseline, baseline_reward = make_runner( + device, prompts, group_size, max_tokens, dynamic=False + ) + dynamic, dynamic_reward = make_runner( + device, prompts, group_size, max_tokens, dynamic=True + ) + baseline_result = measure( + baseline, baseline_reward, batch, trials=trials, warmup=warmup + ) + dynamic_result = measure( + dynamic, dynamic_reward, batch, trials=trials, warmup=warmup + ) + + # Fault injection: advance the policy while the first scoring call is in + # flight. The low-variance refill must invalidate old rows and restart all + # prompt groups under exactly one new version. + dynamic_reward.reset(advance_once=True) + dynamic.clear_cache() + jittered, _ = dynamic(batch) + jitter_versions = { + group.behavior_policy_version for group in jittered.sampling_groups + } + + return { + "device": str(device), + "gpu": torch.cuda.get_device_name(device) + if str(device).startswith("cuda") + else None, + "torch": torch.__version__, + "cuda": torch.version.cuda, + "seed": 3407, + "prompts": prompts, + "group_size": group_size, + "max_tokens": max_tokens, + "warmup": warmup, + "trials": trials, + "baseline": baseline_result, + "dynamic": dynamic_result, + "version_jitter": { + "final_policy_version": jittered.policy_version, + "behavior_policy_versions": sorted(jitter_versions), + "mixed_version_groups": len(jitter_versions) != 1, + "version_invalidated_groups": dynamic.last_sampling_metrics[ + "version_invalidated_groups" + ], + }, + } + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--device", default="cuda") + parser.add_argument("--prompts", type=int, default=4) + parser.add_argument("--group-size", type=int, default=4) + parser.add_argument("--max-tokens", type=int, default=8) + parser.add_argument("--trials", type=int, default=50) + parser.add_argument("--warmup", type=int, default=5) + args = parser.parse_args() + if min(args.prompts, args.group_size, args.max_tokens, args.trials) <= 0: + parser.error("prompts, group-size, max-tokens, and trials must be positive") + if args.warmup < 0: + parser.error("warmup must be non-negative") + print(json.dumps(run(**vars(args)), indent=2)) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/training_consistency/benchmark_dynamic_sampling_distributed.py b/benchmarks/training_consistency/benchmark_dynamic_sampling_distributed.py new file mode 100644 index 00000000..0d053ce4 --- /dev/null +++ b/benchmarks/training_consistency/benchmark_dynamic_sampling_distributed.py @@ -0,0 +1,255 @@ +"""Soak dynamic-sampling rank agreement with intentionally skewed rewards.""" + +import argparse +import json +import os +import statistics +import time + +import torch +import torch.distributed as dist + +from astrai.inference.scheduler import InferenceScheduler +from astrai.trainer.rollout import ( + BaseRewardModel, + DynamicSamplingConfig, + RolloutGenerator, + RolloutRunner, +) +from tests.helpers import FakeTokenizer, make_model + + +class RankSkewRewardModel(BaseRewardModel): + """Disagree on first-round acceptance, then agree on each refill.""" + + def __init__(self, rank): + self.rank = rank + self.seen = {} + self.advance_once = False + self._advanced = False + self.fail_once = False + self._failed = False + self.runner = None + + def reset(self, *, advance_once, fail_once=False): + self.seen.clear() + self.advance_once = advance_once + self._advanced = False + self.fail_once = fail_once + self._failed = False + + def score(self, prompts, responses): + if self.fail_once and self.rank == 0 and not self._failed: + self._failed = True + raise RuntimeError("injected rank-local scoring failure") + rewards = torch.zeros(len(prompts), len(responses[0])) + for row, prompt in enumerate(prompts): + seen = self.seen.get(prompt, 0) + self.seen[prompt] = seen + 1 + ordinal = int(prompt.split("topic ", 1)[1].splitlines()[0]) + if seen > 0 or self.rank == ordinal % dist.get_world_size(): + rewards[row] = torch.arange(rewards.size(1), dtype=torch.float32) + else: + rewards[row] = 1.0 + if self.advance_once and not self._advanced: + assert self.runner is not None + self.runner.update_weights(self.runner.policy_version + 1) + self._advanced = True + return rewards + + +def percentile(samples, fraction): + ordered = sorted(samples) + index = min(len(ordered) - 1, int(len(ordered) * fraction)) + return ordered[index] + + +def run(steps, warmup, jitter_interval, prompts, group_size, max_tokens): + rank = int(os.environ["RANK"]) + local_rank = int(os.environ["LOCAL_RANK"]) + world_size = int(os.environ["WORLD_SIZE"]) + torch.cuda.set_device(local_rank) + device = torch.device("cuda", local_rank) + dist.init_process_group("nccl", device_id=device) + torch.manual_seed(3407) + + model, _ = make_model(device, max_position_embeddings=128) + tokenizer = FakeTokenizer(with_chat_template=True) + scheduler = InferenceScheduler( + model=model, + tokenizer=tokenizer, + max_batch_size=prompts * group_size, + max_seq_len=128, + device=device, + enable_cuda_graph=False, + ) + generator = RolloutGenerator( + scheduler=scheduler, + tokenizer=tokenizer, + max_tokens=max_tokens, + group_size=group_size, + temperature=1.0, + top_k=0, + top_p=1.0, + ) + reward = RankSkewRewardModel(rank) + runner = RolloutRunner( + generator=generator, + reward_model=reward, + rollout_interval=1, + max_policy_lag=1, + dynamic_sampling=DynamicSamplingConfig( + enabled=True, + max_refill_rounds=2, + max_generated_tokens_per_group=group_size * max_tokens * 4, + max_total_rollout_tokens_per_step=prompts * group_size * max_tokens * 4, + max_pending_groups=prompts, + ), + ) + reward.runner = runner + batch = { + "instruction": [f"Reply about topic {index}" for index in range(prompts)], + "input": ["briefly"] * prompts, + } + + generation_schedule = [] + fail_generation = False + original_generate = generator.generate + + def record_generate(batch, *, generation_seed=None): + nonlocal fail_generation + generation_schedule.append(len(batch["instruction"])) + if fail_generation and rank == 0: + fail_generation = False + raise RuntimeError("injected rank-local generation failure") + return original_generate(batch, generation_seed=generation_seed) + + generator.generate = record_generate + start_allocated = 0 + start_reserved = 0 + latencies_ms = [] + mixed_versions = 0 + incomplete_batches = 0 + schedule_mismatches = 0 + invalidated_groups = 0 + + for iteration in range(warmup + steps): + measured = iteration >= warmup + step = iteration - warmup + jitter = measured and jitter_interval > 0 and step % jitter_interval == 0 + if iteration == warmup: + torch.cuda.reset_peak_memory_stats(device) + start_allocated = torch.cuda.memory_allocated(device) + start_reserved = torch.cuda.memory_reserved(device) + reward.reset(advance_once=jitter) + runner.clear_cache() + generation_schedule.clear() + dist.barrier() + started = time.perf_counter_ns() + result, _ = runner(batch) + torch.cuda.synchronize(device) + latency = torch.tensor( + (time.perf_counter_ns() - started) / 1e6, + dtype=torch.float64, + device=device, + ) + dist.all_reduce(latency, op=dist.ReduceOp.MAX) + if rank == 0 and measured: + latencies_ms.append(float(latency.item())) + + if not measured: + continue + + versions = {group.behavior_policy_version for group in result.sampling_groups} + mixed_versions += int(len(versions) != 1) + incomplete_batches += int(len(result.sampling_groups) != prompts) + expected_schedule = ( + [prompts, prompts, prompts] if jitter else [prompts, prompts] + ) + schedule_mismatches += int(generation_schedule != expected_schedule) + invalidated_groups += int( + runner.last_sampling_metrics["version_invalidated_groups"] + ) + + failure_propagation = {} + for failure in ("generation", "scoring"): + fail_generation = failure == "generation" + reward.reset(advance_once=False, fail_once=failure == "scoring") + runner.clear_cache() + caught = False + try: + runner(batch) + except RuntimeError: + caught = True + caught_on_all_ranks = torch.tensor( + int(caught), dtype=torch.int32, device=device + ) + dist.all_reduce(caught_on_all_ranks, op=dist.ReduceOp.MIN) + failure_propagation[failure] = bool(caught_on_all_ranks.item()) + + counters = torch.tensor( + [ + mixed_versions, + incomplete_batches, + schedule_mismatches, + invalidated_groups, + ], + dtype=torch.long, + device=device, + ) + dist.all_reduce(counters, op=dist.ReduceOp.SUM) + end_allocated = torch.cuda.memory_allocated(device) + end_reserved = torch.cuda.memory_reserved(device) + max_allocated = torch.cuda.max_memory_allocated(device) + max_reserved = torch.cuda.max_memory_reserved(device) + + if rank == 0: + print( + json.dumps( + { + "gpu": torch.cuda.get_device_name(device), + "world_size": world_size, + "steps": steps, + "warmup": warmup, + "jitter_interval": jitter_interval, + "prompts": prompts, + "group_size": group_size, + "max_tokens": max_tokens, + "max_rank_latency_median_ms": statistics.median(latencies_ms), + "max_rank_latency_p95_ms": percentile(latencies_ms, 0.95), + "mixed_version_batches": int(counters[0].item()), + "incomplete_batches": int(counters[1].item()), + "generation_schedule_mismatches": int(counters[2].item()), + "version_invalidated_groups": int(counters[3].item()), + "rank_local_failure_propagation": failure_propagation, + "rank0_memory": { + "allocated_drift_bytes": end_allocated - start_allocated, + "reserved_drift_bytes": end_reserved - start_reserved, + "max_allocated_bytes": max_allocated, + "max_reserved_bytes": max_reserved, + }, + }, + indent=2, + ) + ) + dist.destroy_process_group() + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--steps", type=int, default=100) + parser.add_argument("--warmup", type=int, default=5) + parser.add_argument("--jitter-interval", type=int, default=10) + parser.add_argument("--prompts", type=int, default=2) + parser.add_argument("--group-size", type=int, default=2) + parser.add_argument("--max-tokens", type=int, default=4) + args = parser.parse_args() + if min(args.steps, args.prompts, args.group_size, args.max_tokens) <= 0: + parser.error("steps, prompts, group-size, and max-tokens must be positive") + if min(args.warmup, args.jitter_interval) < 0: + parser.error("warmup and jitter-interval must be non-negative") + run(**vars(args)) + + +if __name__ == "__main__": + main() diff --git a/docs/developer/architecture.md b/docs/developer/architecture.md index 50358a8a..7a4fba77 100644 --- a/docs/developer/architecture.md +++ b/docs/developer/architecture.md @@ -625,7 +625,7 @@ classDiagram +supports_online() bool +set_rollout_runner(runner) +prepare_from_rollout(result) Dict - +on_optimizer_step() + +optimizer_step(optimizer) } class LossOutput { @@ -698,12 +698,15 @@ classDiagram +int rep_window +int policy_version +update_weights(policy_version) int + +apply_weight_update(policy_version, update) +generate(batch) RawRollout } class RolloutRunner { +int policy_version + +int max_policy_lag +update_weights(policy_version) int + +apply_weight_update(policy_version, update) +step() +clear_cache() +__call__(batch) Tuple[RolloutResult, bool] diff --git a/docs/developer/internals.md b/docs/developer/internals.md index ccfdd8f3..b8849642 100644 --- a/docs/developer/internals.md +++ b/docs/developer/internals.md @@ -119,8 +119,7 @@ on_train_begin if executor.sync_gradients: before_optimizer_step - optimizer.step() - strategy.on_optimizer_step() + strategy.optimizer_step(optimizer) optimizer.zero_grad() if scheduler: scheduler.step() diff --git a/docs/guides/params.md b/docs/guides/params.md index 9b23ebd8..ca6227b5 100644 --- a/docs/guides/params.md +++ b/docs/guides/params.md @@ -156,10 +156,19 @@ provide a command-line option for configuring one. | Parameter | Description | Default | |-----------|-------------|---------| | `--rollout_interval` | Optimizer steps between rollout refreshes | 512 | +| `--rollout_max_policy_lag` | Maximum accepted rollout/live policy-version gap (`None` derives `rollout_interval - 1`) | None | | `--rollout_temperature` | Rollout sampling temperature | 0.7 | | `--rollout_top_k` | Rollout top-k filtering (`0` disables) | 0 | | `--rollout_top_p` | Rollout nucleus sampling threshold | 0.9 | | `--rollout_max_tokens` | Maximum generated tokens per response | 1024 | +| `--rollout_dynamic_sampling` | Refill low-variance online GRPO groups with version fencing | false | +| `--rollout_dynamic_variance_threshold` | Minimum population reward variance for group acceptance | 0.0 | +| `--rollout_dynamic_max_refill_rounds` | Maximum refill attempts after initial generation | 2 | +| `--rollout_dynamic_max_generated_tokens_per_group` | Hard generated-token budget per group | 32768 | +| `--rollout_dynamic_max_wall_time_per_group` | Hard wall-clock budget per group (seconds) | 300.0 | +| `--rollout_dynamic_max_total_tokens_per_step` | Hard generated-token budget per training step | 262144 | +| `--rollout_dynamic_max_pending_groups` | Maximum prompt groups admitted into one sampling step | 128 | +| `--rollout_dynamic_seed` | Base seed for reproducible refill attempts (`None` uses `random_seed`) | None | ### Scheduler diff --git a/docs/guides/training.md b/docs/guides/training.md index 0d267f12..29dd4cea 100644 --- a/docs/guides/training.md +++ b/docs/guides/training.md @@ -71,8 +71,7 @@ on_train_begin if executor.sync_gradients: before_optimizer_step - optimizer.step() - strategy.on_optimizer_step() + strategy.optimizer_step(optimizer) optimizer.zero_grad() if scheduler: scheduler.step() @@ -171,12 +170,51 @@ them with a `BaseRewardModel`. It refreshes cached rollouts every behaviour log-probabilities into the loss, so it does not allocate or synchronize a separate old-policy model. -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. +Every successful optimizer step mutates the shared model and advances its +monotonic `policy_version` under the same generation lock. The scheduler +invalidates reusable KV prefixes before accepting the new version, so an async +rollout cannot observe partially updated weights under the previous 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. +even while later optimizer steps advance the live policy. Results from a future +version or beyond `rollout_max_policy_lag` are rejected before training. The +final version check and rollout-cache publication share that policy lock, so a +concurrent update cannot land between validation and cache insertion. + +Online GRPO can additionally enable a versioned dynamic-sampling buffer with +`rollout_dynamic_sampling`. Each prompt group moves through +`pending -> generating -> scoring` and is either accepted, refilled, +invalidated, or dropped. Only low-variance groups are regenerated. Accepted +rows are reassembled in the original prompt order, so refill completion order +cannot reorder the training batch. + +Every attempt records a stable prompt ID, attempt ID, generation seed, +behaviour-policy version, reward vector and variance, refill round, token count, +timestamps, and terminal reason. Acceptance is provisional until the entire +batch commits: if the policy version changes while any group is being +refilled, all accepted rows from the old version are invalidated and the full +batch restarts on the new version. Samples from different behaviour-policy +versions are never combined into one training batch. + +Refill is bounded by per-group round, token, and wall-time budgets plus +per-step generated-token and pending-group limits. Exhausting any budget raises +`DynamicSamplingBudgetError` instead of returning a partial batch. When +`torch.distributed` is initialized, policy versions, group acceptance, and +budget failure are reduced across ranks before the next generation round, +which keeps every rank on the same training step. Per-refresh counters are +reported under `dynamic_sampling/*`, including accepted and zero-variance +groups, refill rounds/tokens, invalidations, budget exhaustion, effective +groups per million generated tokens, and rollout waste ratio. + +On an NVIDIA L20 with eight prompts, group size four, and a controlled workload +where half of first attempts have zero variance, dynamic refill accepted twice +as many useful groups and improved effective groups per million generated +tokens by 33.37%. The cost was 49.95% more generated tokens, 55.80% higher +median rollout latency, and a 33.23% waste ratio. A 100-step, three-rank NCCL +soak with rank-skewed rewards and policy-version jitter produced zero mixed +version batches, incomplete batches, or generation-schedule mismatches. Full +parameters and raw measurements are in +`benchmarks/results/versioned_dynamic_sampling_l20_sm89.json`. 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 diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 1f9690eb..d045410d 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -333,6 +333,13 @@ def _merge_yaml_into_kwargs( group="Algorithm", help="Steps between rollouts.", ) +@opt( + "--rollout_max_policy_lag", + type=int, + default=None, + group="Algorithm", + help="Maximum accepted rollout/live policy-version gap.", +) @opt( "--rollout_temperature", type=float, @@ -361,6 +368,61 @@ def _merge_yaml_into_kwargs( group="Algorithm", help="Max tokens per rollout response.", ) +@opt( + "--rollout_dynamic_sampling/--no-rollout_dynamic_sampling", + default=False, + group="Algorithm", + help="Refill low-variance online GRPO prompt groups.", +) +@opt( + "--rollout_dynamic_variance_threshold", + type=float, + default=0.0, + group="Algorithm", + help="Minimum population reward variance for group acceptance.", +) +@opt( + "--rollout_dynamic_max_refill_rounds", + type=int, + default=2, + group="Algorithm", + help="Maximum low-variance refill rounds per prompt group.", +) +@opt( + "--rollout_dynamic_max_generated_tokens_per_group", + type=int, + default=32768, + group="Algorithm", + help="Hard generated-token budget per prompt group.", +) +@opt( + "--rollout_dynamic_max_wall_time_per_group", + type=float, + default=300.0, + group="Algorithm", + help="Hard wall-clock budget in seconds per prompt group.", +) +@opt( + "--rollout_dynamic_max_total_tokens_per_step", + type=int, + default=262144, + group="Algorithm", + help="Hard generated-token budget per training step.", +) +@opt( + "--rollout_dynamic_max_pending_groups", + type=int, + default=128, + group="Algorithm", + help="Maximum prompt groups admitted into one sampling step.", +) +@opt( + "--rollout_dynamic_seed", + type=int, + default=None, + group="Algorithm", + help="Base refill seed (defaults to random_seed).", +) @opt( "--gradient_checkpointing/--no-gradient_checkpointing", default=False, @@ -684,10 +746,31 @@ def train( } rollout_interval = kwargs.pop("rollout_interval", 512) + rollout_max_policy_lag = kwargs.pop("rollout_max_policy_lag", None) rollout_temperature = kwargs.pop("rollout_temperature", 0.7) rollout_top_k = kwargs.pop("rollout_top_k", 0) rollout_top_p = kwargs.pop("rollout_top_p", 0.9) rollout_max_tokens = kwargs.pop("rollout_max_tokens", 1024) + rollout_dynamic_sampling = kwargs.pop("rollout_dynamic_sampling", False) + rollout_dynamic_variance_threshold = kwargs.pop( + "rollout_dynamic_variance_threshold", 0.0 + ) + rollout_dynamic_max_refill_rounds = kwargs.pop( + "rollout_dynamic_max_refill_rounds", 2 + ) + rollout_dynamic_max_generated_tokens_per_group = kwargs.pop( + "rollout_dynamic_max_generated_tokens_per_group", 32768 + ) + rollout_dynamic_max_wall_time_per_group = kwargs.pop( + "rollout_dynamic_max_wall_time_per_group", 300.0 + ) + rollout_dynamic_max_total_tokens_per_step = kwargs.pop( + "rollout_dynamic_max_total_tokens_per_step", 262144 + ) + rollout_dynamic_max_pending_groups = kwargs.pop( + "rollout_dynamic_max_pending_groups", 128 + ) + rollout_dynamic_seed = kwargs.pop("rollout_dynamic_seed", None) reward_model_fn: Callable[[], BaseRewardModel] | None = None executor_kwargs = {} @@ -840,10 +923,25 @@ def train( neftune_alpha=neftune_alpha, collate_fn=collate_fn, rollout_interval=rollout_interval, + rollout_max_policy_lag=rollout_max_policy_lag, rollout_temperature=rollout_temperature, rollout_top_k=rollout_top_k, rollout_top_p=rollout_top_p, rollout_max_tokens=rollout_max_tokens, + rollout_dynamic_sampling=rollout_dynamic_sampling, + rollout_dynamic_variance_threshold=rollout_dynamic_variance_threshold, + rollout_dynamic_max_refill_rounds=rollout_dynamic_max_refill_rounds, + rollout_dynamic_max_generated_tokens_per_group=( + rollout_dynamic_max_generated_tokens_per_group + ), + rollout_dynamic_max_wall_time_per_group=( + rollout_dynamic_max_wall_time_per_group + ), + rollout_dynamic_max_total_tokens_per_step=( + rollout_dynamic_max_total_tokens_per_step + ), + rollout_dynamic_max_pending_groups=rollout_dynamic_max_pending_groups, + rollout_dynamic_seed=rollout_dynamic_seed, reward_model_fn=reward_model_fn, moe_aux_loss_coef=kwargs.pop("moe_aux_loss_coef", 0.01), ) diff --git a/tests/inference/test_scheduler.py b/tests/inference/test_scheduler.py index 8879813f..292b3cdb 100644 --- a/tests/inference/test_scheduler.py +++ b/tests/inference/test_scheduler.py @@ -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") diff --git a/tests/trainer/test_online_e2e.py b/tests/trainer/test_online_e2e.py index df8b8172..7bbbfd3c 100644 --- a/tests/trainer/test_online_e2e.py +++ b/tests/trainer/test_online_e2e.py @@ -10,6 +10,7 @@ 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 @@ -54,6 +55,15 @@ def score(self, prompts, responses): return rewards +class RankRewardModel(BaseRewardModel): + """Always gives each response rank a distinct reward.""" + + def score(self, prompts, responses): + return torch.arange(len(responses[0]), dtype=torch.float32).repeat( + len(prompts), 1 + ) + + def instruction_collate_fn(batch): """Stack a list of instruction/input dicts into a batch dict of lists.""" return { @@ -126,6 +136,7 @@ def track_reference_model(*args, **kwargs): 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 +148,57 @@ def track_reference_model(*args, **kwargs): 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 + + +@pytest.mark.integration +def test_dynamic_sampling_online_grpo_end_to_end(base_test_env): + """Exercise TrainConfig -> runner wiring with dynamic sampling enabled.""" + test_dir = base_test_env["test_dir"] + device = base_test_env["device"] + tokenizer = base_test_env["tokenizer"] + model_config = base_test_env["transformer_config"] + tokenizer.set_chat_template(CHAT_TEMPLATE) + tokenizer.save_pretrained(test_dir) + + train_config = TrainConfig( + strategy="online_grpo", + model_fn=partial(_model_fn, model_config), + dataset=InstructionDataset(), + optimizer_fn=_optimizer_fn, + scheduler_fn=_scheduler_fn, + ckpt_dir=os.path.join(test_dir, "dynamic_ckpt"), + n_epoch=1, + batch_per_device=2, + ckpt_interval=100, + grad_accum_steps=1, + random_seed=42, + device_type=device, + nprocs=1, + parallel_mode="none", + strategy_kwargs={"clip_eps": 0.2, "kl_coef": 0.01, "group_size": 2}, + rollout_interval=1, + rollout_max_policy_lag=0, + rollout_temperature=1.0, + rollout_top_k=0, + rollout_top_p=1.0, + rollout_max_tokens=4, + rollout_dynamic_sampling=True, + rollout_dynamic_max_refill_rounds=1, + rollout_dynamic_max_generated_tokens_per_group=16, + rollout_dynamic_max_total_tokens_per_step=32, + rollout_dynamic_max_pending_groups=2, + reward_model_fn=RankRewardModel, + collate_fn=instruction_collate_fn, + ) + + Trainer(train_config).train(param_path=test_dir) + + checkpoint = Checkpoint.load( + os.path.join(test_dir, "dynamic_ckpt", "epoch_0_step_2") + ) + assert checkpoint.meta["policy_version"] == 2 diff --git a/tests/trainer/test_online_strategy.py b/tests/trainer/test_online_strategy.py index e66a500d..396231fe 100644 --- a/tests/trainer/test_online_strategy.py +++ b/tests/trainer/test_online_strategy.py @@ -60,11 +60,25 @@ def update_weights(self, policy_version): 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))) diff --git a/tests/trainer/test_rollout.py b/tests/trainer/test_rollout.py index a36ad6d9..6e36905a 100644 --- a/tests/trainer/test_rollout.py +++ b/tests/trainer/test_rollout.py @@ -1,5 +1,7 @@ """Unit tests for the online rollout module.""" +import threading + import pytest import torch @@ -7,10 +9,15 @@ from astrai.inference.task import GenerationResult from astrai.trainer.rollout import ( BaseRewardModel, + DynamicSamplingBudgetError, + DynamicSamplingConfig, + DynamicSamplingGroup, + DynamicSamplingState, RawRollout, RolloutGenerator, RolloutResult, RolloutRunner, + RolloutVersionError, ) from tests.helpers import FakeTokenizer, make_model @@ -39,6 +46,18 @@ def score(self, prompts, responses): return torch.full((B, G), float("nan")) +class ScriptedRewardModel(BaseRewardModel): + """Returns one explicitly shaped reward matrix per scoring call.""" + + def __init__(self, outputs): + self.outputs = [torch.tensor(output, dtype=torch.float32) for output in outputs] + + def score(self, prompts, responses): + output = self.outputs.pop(0) + assert output.shape == (len(prompts), len(responses[0])) + return output + + def _make_scheduler(model, tokenizer, max_batch_size=8, max_len=128): return InferenceScheduler( model=model, @@ -151,6 +170,113 @@ def recording_run_batch(*args, **kwargs): 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) @@ -242,17 +368,187 @@ def _make_runner(device, **kw): max_batch_size=kw.get("max_batch_size", 8), max_len=kw.get("max_position_embeddings", 128), ) - rm = ConstantRewardModel(1.0) + rm = kw.get("reward_model", ConstantRewardModel(1.0)) return ( RolloutRunner( generator=generator, reward_model=rm, rollout_interval=kw.get("rollout_interval", 2), + max_policy_lag=kw.get("max_policy_lag"), + dynamic_sampling=kw.get("dynamic_sampling"), ), model, ) +def test_dynamic_sampling_group_enforces_state_machine(): + group = DynamicSamplingGroup(prompt_uid="prompt:0", attempt_id=1, generation_seed=7) + with pytest.raises(RuntimeError, match="pending -> accepted"): + group.transition(DynamicSamplingState.ACCEPTED) + group.transition(DynamicSamplingState.GENERATING) + group.transition(DynamicSamplingState.SCORING) + group.transition(DynamicSamplingState.ACCEPTED) + assert group.accepted is True + assert group.completed_at is not None + + +def test_dynamic_sampling_refills_only_low_variance_groups(device): + rewards = ScriptedRewardModel( + [ + [[0.0, 1.0], [1.0, 1.0]], + [[0.0, 2.0]], + ] + ) + config = DynamicSamplingConfig(enabled=True, base_seed=19) + runner, _ = _make_runner( + device, + group_size=2, + max_tokens=2, + reward_model=rewards, + dynamic_sampling=config, + ) + batch_sizes = [] + seeds = [] + original_generate = runner.generator.generate + + def record_generate(batch, *, generation_seed=None): + batch_sizes.append(len(batch["instruction"])) + seeds.append(generation_seed) + return original_generate(batch, generation_seed=generation_seed) + + runner.generator.generate = record_generate + result, is_fresh = runner(_make_instruction_batch(n=2)) + + assert is_fresh is True + assert batch_sizes == [2, 1] + assert len(set(seeds)) == 2 + assert result.rewards.tolist() == [[0.0, 1.0], [0.0, 2.0]] + assert [group.refill_round for group in result.sampling_groups] == [0, 1] + assert {group.behavior_policy_version for group in result.sampling_groups} == {0} + assert all( + group.state is DynamicSamplingState.ACCEPTED for group in result.sampling_groups + ) + assert runner.last_sampling_metrics["groups_accepted"] == 2.0 + assert runner.last_sampling_metrics["zero_variance_groups"] == 1.0 + assert runner.last_sampling_metrics["refill_rounds"] == 1.0 + assert runner.last_sampling_metrics["rollout_waste_ratio"] > 0.0 + + +def test_dynamic_sampling_restarts_whole_batch_after_version_change(device): + rewards = ScriptedRewardModel( + [ + [[0.0, 1.0], [1.0, 1.0]], + [[0.0, 1.0], [0.0, 2.0]], + ] + ) + runner, _ = _make_runner( + device, + group_size=2, + max_tokens=2, + reward_model=rewards, + dynamic_sampling=DynamicSamplingConfig(enabled=True), + ) + batch_sizes = [] + original_generate = runner.generator.generate + + def update_before_refill(batch, *, generation_seed=None): + batch_sizes.append(len(batch["instruction"])) + if len(batch_sizes) == 2: + runner.generator.update_weights(1) + return original_generate(batch, generation_seed=generation_seed) + + runner.generator.generate = update_before_refill + result, _ = runner(_make_instruction_batch(n=2)) + + assert batch_sizes == [2, 1, 2] + assert result.policy_version == 1 + assert {group.behavior_policy_version for group in result.sampling_groups} == {1} + invalidated = [ + group + for group in runner.last_sampling_history + if group.state is DynamicSamplingState.INVALIDATED + ] + assert len(invalidated) == 2 + assert runner.last_sampling_metrics["version_invalidated_groups"] == 2.0 + + +def test_dynamic_sampling_budget_exhaustion_refuses_partial_batch(device): + runner, _ = _make_runner( + device, + group_size=2, + max_tokens=2, + dynamic_sampling=DynamicSamplingConfig( + enabled=True, + max_refill_rounds=0, + ), + ) + + with pytest.raises(DynamicSamplingBudgetError, match="partial"): + runner(_make_instruction_batch(n=2)) + + assert runner.last_sampling_metrics["groups_accepted"] == 0.0 + assert runner.last_sampling_metrics["dropped_groups"] == 2.0 + assert runner.last_sampling_metrics["budget_exhausted_groups"] == 2.0 + assert runner._cache is None + + +def test_dynamic_sampling_pending_group_budget_fails_before_generation(device): + runner, _ = _make_runner( + device, + dynamic_sampling=DynamicSamplingConfig(enabled=True, max_pending_groups=1), + ) + with pytest.raises(DynamicSamplingBudgetError, match="max_pending_groups=1"): + runner(_make_instruction_batch(n=2)) + + +def test_dynamic_sampling_generation_budget_is_reserved_before_attempt(device): + runner, _ = _make_runner( + device, + group_size=2, + max_tokens=4, + dynamic_sampling=DynamicSamplingConfig( + enabled=True, + max_generated_tokens_per_group=7, + ), + ) + called = False + + def should_not_generate(*_args, **_kwargs): + nonlocal called + called = True + + runner.generator.generate = should_not_generate + with pytest.raises(DynamicSamplingBudgetError, match="cannot start"): + runner(_make_instruction_batch(n=1)) + assert called is False + assert runner.last_sampling_history[0].discard_reason == ( + "max_generated_tokens_per_group" + ) + + +def test_dynamic_sampling_scoring_failure_drops_attempt(device): + runner, _ = _make_runner( + device, + group_size=2, + max_tokens=2, + reward_model=BadShapeRewardModel(), + dynamic_sampling=DynamicSamplingConfig(enabled=True), + ) + with pytest.raises(ValueError, match="Reward model returned shape"): + runner(_make_instruction_batch(n=1)) + assert runner.last_sampling_history[0].state is DynamicSamplingState.DROPPED + assert runner.last_sampling_history[0].discard_reason == "scoring_failed" + + +def test_seeded_rollout_generation_restores_torch_rng(device): + generator, _ = _make_generator(device, group_size=2, max_tokens=2) + torch.manual_seed(1234) + before = torch.random.get_rng_state() + generator.generate(_make_instruction_batch(n=1), generation_seed=99) + after = torch.random.get_rng_state() + assert torch.equal(before, after) + + def test_rollout_runner_shapes(device): runner, _ = _make_runner(device, group_size=3, max_tokens=5) batch = _make_instruction_batch(n=2) @@ -297,6 +593,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))