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
42 changes: 38 additions & 4 deletions astrai/config/train_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,14 @@ class TrainConfig(BaseConfig):
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.
critic_model_fn (Optional[Callable]): Factory for the value (critic) model, required for online_ppo. Defaults to None.
critic_optimizer_fn (Optional[Callable]): Factory for the critic optimizer; None reuses optimizer_fn. Defaults to None.
Expand Down Expand Up @@ -133,6 +141,14 @@ class TrainConfig(BaseConfig):
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
critic_model_fn: Optional[Callable] = None
critic_optimizer_fn: Optional[Callable] = None
Expand Down Expand Up @@ -213,13 +229,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 @@ -232,17 +251,22 @@ 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")
@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"rollout_max_policy_lag must be non-negative, got {v}")
raise ValueError(f"must be non-negative or None, got {v}")
return v

@field_validator("max_grad_norm")
Expand Down Expand Up @@ -287,4 +311,14 @@ def _validate_online_strategy(self) -> "TrainConfig":
f"rollouts up to that lag, so a tighter bound guarantees "
f"a fatal RolloutVersionError mid-training"
)
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
Loading
Loading