Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 44 additions & 2 deletions astrai/config/train_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 {}.
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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}")
Expand All @@ -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:
Expand Down Expand Up @@ -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
59 changes: 47 additions & 12 deletions astrai/inference/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -27,6 +27,7 @@
from astrai.tokenize.tokenizer import AutoTokenizer

logger = logging.getLogger(__name__)
T = TypeVar("T")


def _with_weight_lock(method):
Expand Down Expand Up @@ -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)
Expand All @@ -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)

Expand Down
Loading
Loading