Skip to content
Closed
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
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -260,8 +260,8 @@ areno train \
```

The tuner uses dummy-loaded model weights and synthetic token rows, respects
the configured sequence limits and tensor-parallel size, and enables
`--drop-rollout-state` for the tuned run. See `docs/cli/training.rst` for the
the configured sequence limits and tensor-parallel size, and uses the default
dropped rollout state for the tuned run. See `docs/cli/training.rst` for the
full tuning rules.

For Agentic RL, add `--agent-fn` to supply an agent function. The agent calls the local OpenAI-compatible endpoint, including `tools` and `tool_choice` when needed, and returns explicit `AgentTrajectoryTurn` objects. AReno converts those turns into trainable assistant outputs and masks tool results by default:
Expand Down
2 changes: 1 addition & 1 deletion areno/api/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ class MlxConfig:
prefill_step_size: int = 2048
max_kv_size: int | None = None
decode_progress_interval_s: float = 10.0
keep_rollout_state: bool = True
keep_rollout_state: bool = False
logits_chunk_size: int = 4096
compile_train_step: bool = True
gradient_checkpointing: bool = True
Expand Down
30 changes: 29 additions & 1 deletion areno/api/trainer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,12 @@ class TrainerConfig:
multimodal_projector_lr_decay_steps: int | None = None
multimodal_projector_lr_decay_style: str | None = None
activation_checkpointing: bool = True
keep_rollout_state: bool = True
fp8_checkpoint_activations: bool | None = None
fp8_checkpoint_group_size: int = 128
fp8_checkpoint_stochastic: bool = False
fp8_checkpoint_warmup_steps: int = 0
fp8_checkpoint_fallback_layers: tuple[int, ...] = ()
keep_rollout_state: bool = False
optimizer_state_offload: str | bool = "none"
optimizer_state_offload_dir: str | None = None
optimizer_state_offload_batch_size: int = 1
Expand All @@ -89,12 +94,25 @@ def __post_init__(self) -> None:
self.backend = default_backend_type().value.lower()
else:
self.backend = self.backend.lower()
if self.fp8_checkpoint_activations is None:
self.fp8_checkpoint_activations = self.backend == "cuda" and self.activation_checkpointing
self.fp8_checkpoint_fallback_layers = tuple(self.fp8_checkpoint_fallback_layers)
if self.backend not in {"cuda", "mlx"}:
raise ValueError("backend must be one of: cuda, mlx")
if self.adam_4bit and self.adam_8bit:
raise ValueError("adam_4bit and adam_8bit are mutually exclusive")
if self.adam_4bit and self.backend != "cuda":
raise ValueError("adam_4bit is only supported by the CUDA backend")
if self.fp8_checkpoint_activations and self.backend != "cuda":
raise ValueError("fp8_checkpoint_activations is only supported by the CUDA backend")
if self.fp8_checkpoint_activations and not self.activation_checkpointing:
raise ValueError("fp8_checkpoint_activations requires activation_checkpointing")
if self.fp8_checkpoint_group_size not in {0, 128, 256}:
raise ValueError("fp8_checkpoint_group_size must be one of: 0, 128, 256")
if self.fp8_checkpoint_warmup_steps < 0:
raise ValueError("fp8_checkpoint_warmup_steps must be non-negative")
if any(layer < 0 for layer in self.fp8_checkpoint_fallback_layers):
raise ValueError("fp8_checkpoint_fallback_layers must contain non-negative indices")
if self.attn_backend not in {"flash", "native"}:
raise ValueError("attn_backend must be one of: flash, native")
if self.model_hub not in {"hf", "modelscope"}:
Expand Down Expand Up @@ -216,6 +234,11 @@ def cuda_config(self):
optimizer=self.optimizer_config(),
runtime={
"activation_checkpointing": self.activation_checkpointing,
"fp8_checkpoint_activations": self.fp8_checkpoint_activations,
"fp8_checkpoint_group_size": self.fp8_checkpoint_group_size,
"fp8_checkpoint_stochastic": self.fp8_checkpoint_stochastic,
"fp8_checkpoint_warmup_steps": self.fp8_checkpoint_warmup_steps,
"fp8_checkpoint_fallback_layers": self.fp8_checkpoint_fallback_layers,
"keep_rollout_state": self.keep_rollout_state,
"optimizer_state_offload": self.optimizer_state_offload,
"optimizer_state_offload_dir": self.optimizer_state_offload_dir,
Expand Down Expand Up @@ -266,6 +289,11 @@ def cuda_config(self):
optimizer=self.optimizer_config(),
runtime={
"activation_checkpointing": self.activation_checkpointing,
"fp8_checkpoint_activations": self.fp8_checkpoint_activations,
"fp8_checkpoint_group_size": self.fp8_checkpoint_group_size,
"fp8_checkpoint_stochastic": self.fp8_checkpoint_stochastic,
"fp8_checkpoint_warmup_steps": self.fp8_checkpoint_warmup_steps,
"fp8_checkpoint_fallback_layers": self.fp8_checkpoint_fallback_layers,
"keep_rollout_state": self.keep_rollout_state,
"optimizer_state_offload": self.optimizer_state_offload,
"optimizer_state_offload_dir": self.optimizer_state_offload_dir,
Expand Down
20 changes: 18 additions & 2 deletions areno/cli/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,7 @@ def flash_attention_unsupported_model_reason(model_config):
"score_micro_bs",
"gradient_accumulation_steps",
"activation_checkpointing",
"fp8_checkpoint_activations",
"lora_rank",
"lora_alpha",
"lora_dropout",
Expand Down Expand Up @@ -240,6 +241,7 @@ def _trainer_config_from_options(**options) -> TrainerConfig:
args.rollout_devices = getattr(args, "rollout_devices", None)
args.policy_sync_bucket_mb = getattr(args, "policy_sync_bucket_mb", 64)
args.adam_4bit = getattr(args, "adam_4bit", False)
args.fp8_checkpoint_activations = getattr(args, "fp8_checkpoint_activations", None)
args.optimizer_state_offload = getattr(args, "optimizer_state_offload", "none")
args.optimizer_state_offload_dir = getattr(args, "optimizer_state_offload_dir", None)
args.optimizer_state_offload_batch_size = getattr(args, "optimizer_state_offload_batch_size", 1)
Expand Down Expand Up @@ -525,6 +527,7 @@ def _format_training_config_summary(
[
("max_steps", _format_optional(config.max_steps)),
("sequence_parallel", _sequence_parallel_for_summary(config, model_config)),
("fp8_ckpt_activations", _format_bool(config.fp8_checkpoint_activations)),
("mini_bs", str(config.mini_bs)),
("score_micro_bs", str(config.score_micro_bs)),
("gradient_accumulation_steps", _format_optional(config.gradient_accumulation_steps, default="auto")),
Expand Down Expand Up @@ -848,6 +851,7 @@ def _trainer_config_from_args(args) -> TrainerConfig:
args.rollout_devices = getattr(args, "rollout_devices", None)
args.policy_sync_bucket_mb = getattr(args, "policy_sync_bucket_mb", 64)
args.adam_4bit = getattr(args, "adam_4bit", False)
args.fp8_checkpoint_activations = getattr(args, "fp8_checkpoint_activations", None)
args.unfreeze_multimodal_tower = getattr(args, "unfreeze_multimodal_tower", False)
args.unfreeze_multimodal_projector = getattr(args, "unfreeze_multimodal_projector", False)
args.multimodal_tower_lr = getattr(args, "multimodal_tower_lr", None)
Expand Down Expand Up @@ -907,6 +911,7 @@ def _trainer_config_from_args(args) -> TrainerConfig:
multimodal_projector_lr_decay_steps=args.multimodal_projector_lr_decay_steps,
multimodal_projector_lr_decay_style=args.multimodal_projector_lr_decay_style,
activation_checkpointing=args.activation_checkpointing,
fp8_checkpoint_activations=args.fp8_checkpoint_activations,
keep_rollout_state=not args.drop_rollout_state,
optimizer_state_offload=args.optimizer_state_offload,
optimizer_state_offload_dir=args.optimizer_state_offload_dir,
Expand Down Expand Up @@ -967,6 +972,7 @@ def _trainer_config_from_args(args) -> TrainerConfig:
multimodal_projector_lr_decay_steps=args.multimodal_projector_lr_decay_steps,
multimodal_projector_lr_decay_style=args.multimodal_projector_lr_decay_style,
activation_checkpointing=args.activation_checkpointing,
fp8_checkpoint_activations=args.fp8_checkpoint_activations,
keep_rollout_state=not args.drop_rollout_state,
optimizer_state_offload=args.optimizer_state_offload,
optimizer_state_offload_dir=args.optimizer_state_offload_dir,
Expand Down Expand Up @@ -1035,6 +1041,7 @@ def _trainer_config_from_args(args) -> TrainerConfig:
multimodal_projector_lr_decay_steps=args.multimodal_projector_lr_decay_steps,
multimodal_projector_lr_decay_style=args.multimodal_projector_lr_decay_style,
activation_checkpointing=args.activation_checkpointing,
fp8_checkpoint_activations=args.fp8_checkpoint_activations,
keep_rollout_state=not args.drop_rollout_state,
optimizer_state_offload=args.optimizer_state_offload,
optimizer_state_offload_dir=args.optimizer_state_offload_dir,
Expand Down Expand Up @@ -1104,6 +1111,7 @@ def _trainer_config_from_args(args) -> TrainerConfig:
multimodal_projector_lr_decay_steps=args.multimodal_projector_lr_decay_steps,
multimodal_projector_lr_decay_style=args.multimodal_projector_lr_decay_style,
activation_checkpointing=args.activation_checkpointing,
fp8_checkpoint_activations=args.fp8_checkpoint_activations,
keep_rollout_state=not args.drop_rollout_state,
optimizer_state_offload=args.optimizer_state_offload,
optimizer_state_offload_dir=args.optimizer_state_offload_dir,
Expand Down Expand Up @@ -1232,6 +1240,7 @@ def section(title: str, names: list[str]) -> dict:
"attn_backend",
"eager_decode",
"activation_checkpointing",
"fp8_checkpoint_activations",
"keep_rollout_state",
"optimizer_state_offload",
"optimizer_state_offload_dir",
Expand Down Expand Up @@ -1753,8 +1762,15 @@ def _dataset_builder_for_suffix(suffix: str) -> str:
help="Enable decoder-layer activation recompute during training.",
)
@click.option(
"--drop-rollout-state",
is_flag=True,
"--fp8-ckpt-activations/--no-fp8-ckpt-activations",
"fp8_checkpoint_activations",
default=None,
help="Store activation-checkpoint boundary tensors in FP8 E4M3 (enabled by default on CUDA).",
)
@click.option(
"--drop-rollout-state/--keep-rollout-state",
default=True,
show_default=True,
help="Release completed rollout KV/cache state after each step.",
)
@click.option(
Expand Down
Loading
Loading