From 6c76c16480aefa407eb11fc0ce9c488b25d141f0 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sat, 1 Aug 2026 08:32:28 +0800 Subject: [PATCH] feat: group train CLI options in --help output - add GroupedOption/GroupedCommand (no third-party dep) that tags each option with a group label and renders help in labeled sections - add opt() shorthand wrapping click.option with cls=GroupedOption - tag all ~55 options into 10 groups aligned with params.md chapters --- scripts/tools/train.py | 396 +++++++++++++++++++++++++++++++++-------- 1 file changed, 323 insertions(+), 73 deletions(-) diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 249a739..5a4e05a 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -1,4 +1,5 @@ import os +from collections import OrderedDict from collections.abc import Callable from functools import partial @@ -17,6 +18,37 @@ from astrai.trainer import SchedulerFactory, Trainer from astrai.trainer.rollout import BaseRewardModel +class GroupedOption(click.Option): + """A ``click.Option`` that carries a ``group`` label for help output.""" + + def __init__(self, *args, group: str = "Options", **kwargs): + super().__init__(*args, **kwargs) + self.group = group + + +class GroupedCommand(click.Command): + """A ``click.Command`` that renders options grouped by their ``group``.""" + + def format_options(self, ctx, formatter): + groups: OrderedDict[str, list] = OrderedDict() + for param in self.get_params(ctx): + record = param.get_help_record(ctx) + if record is None: + continue + group = getattr(param, "group", "Options") + groups.setdefault(group, []).append(record) + for group_name, records in groups.items(): + with formatter.section(group_name): + formatter.write_dl(records) + + +def opt(*param_decls, group: str, **kwargs): + """Shorthand for ``click.option`` that tags the option with a group.""" + kwargs.setdefault("cls", GroupedOption) + kwargs["group"] = group + return click.option(*param_decls, **kwargs) + + def _merge_yaml_into_kwargs( config_path: str, passed_kwargs: dict, @@ -52,176 +84,394 @@ _START_METHODS = ["spawn", "fork", "forkserver"] @click.command( name="train", + cls=GroupedCommand, help="Start model training (pretrain / SFT / DPO / GRPO).", context_settings={"show_default": True}, ) -@click.option( +@opt( "--config", "-c", "config_path", type=click.Path(exists=True), + group="Paths & Setup", help="YAML config file. CLI flags override YAML values.", ) -@click.option( +@opt( "--train_type", type=click.Choice(_TRAIN_TYPE), required=False, + group="Paths & Setup", help="Training type.", ) -@click.option( +@opt( "--data_root_path", type=click.Path(exists=True), + group="Paths & Setup", help="Root directory of the dataset.", ) -@click.option( +@opt( "--param_path", type=click.Path(exists=True), + group="Paths & Setup", help="Path to model parameters or resume checkpoint.", ) -@click.option("--resume", is_flag=True, default=False, help="Resume from checkpoint.") -@click.option("--n_epoch", type=int, default=1, help="Number of epochs.") -@click.option("--batch_per_device", type=int, default=1, help="Batch size per GPU.") -@click.option( - "--grad_accum_steps", type=int, default=1, help="Gradient accumulation steps." +@opt( + "--resume", + is_flag=True, + default=False, + group="Paths & Setup", + help="Resume from checkpoint.", ) -@click.option( +@opt("--n_epoch", type=int, default=1, group="Training", help="Number of epochs.") +@opt( + "--batch_per_device", + type=int, + default=1, + group="Training", + help="Batch size per GPU.", +) +@opt( + "--grad_accum_steps", + type=int, + default=1, + group="Training", + help="Gradient accumulation steps.", +) +@opt( "--warmup_ratio", type=float, default=0.05, + group="LR Schedule", help="Fraction of total steps for LR warmup.", ) -@click.option("--max_lr", type=float, default=3e-4, help="Max learning rate.") -@click.option( +@opt( + "--max_lr", + type=float, + default=3e-4, + group="Optimizer", + help="Max learning rate.", +) +@opt( "--optimizer", type=click.Choice(_OPTIMIZERS), default="muon_adamw", + group="Optimizer", help="Built-in optimizer.", ) -@click.option( - "--max_grad_norm", type=float, default=1.0, help="Max gradient norm for clipping." +@opt( + "--max_grad_norm", + type=float, + default=1.0, + group="Training", + help="Max gradient norm for clipping.", ) -@click.option( +@opt( "--weight_decay", type=float, default=0.1, + group="Optimizer", help="Weight decay for eligible optimizer parameters.", ) -@click.option("--nora_lr", type=float, default=5e-3, help="Nora learning rate.") -@click.option("--nora_beta", type=float, default=0.95, help="Nora EMA factor.") -@click.option("--nora_momentum", type=float, default=0.95, help="Nora update momentum.") -@click.option("--nora_weight_decay", type=float, default=0.0, help="Nora weight decay.") -@click.option("--muon_momentum", type=float, default=0.95, help="Muon momentum factor.") -@click.option("--muon_nesterov/--no-muon_nesterov", default=True, help="Muon Nesterov.") -@click.option("--muon_ns_steps", type=int, default=5, help="Muon Newton-Schulz steps.") -@click.option( +@opt( + "--nora_lr", type=float, default=5e-3, group="Optimizer", help="Nora learning rate." +) +@opt( + "--nora_beta", type=float, default=0.95, group="Optimizer", help="Nora EMA factor." +) +@opt( + "--nora_momentum", + type=float, + default=0.95, + group="Optimizer", + help="Nora update momentum.", +) +@opt( + "--nora_weight_decay", + type=float, + default=0.0, + group="Optimizer", + help="Nora weight decay.", +) +@opt( + "--muon_momentum", + type=float, + default=0.95, + group="Optimizer", + help="Muon momentum factor.", +) +@opt( + "--muon_nesterov/--no-muon_nesterov", + default=True, + group="Optimizer", + help="Muon Nesterov.", +) +@opt( + "--muon_ns_steps", + type=int, + default=5, + group="Optimizer", + help="Muon Newton-Schulz steps.", +) +@opt( "--muon_adjust_lr", type=click.Choice(["original", "match_rms_adamw"]), default="match_rms_adamw", + group="Optimizer", help="Muon LR adjustment strategy.", ) -@click.option("--random_seed", type=int, default=3407, help="Random seed.") -@click.option("--num_workers", type=int, default=4, help="DataLoader workers.") -@click.option("--pin_memory/--no-pin_memory", default=True, help="Pin memory.") -@click.option( - "--window_size", type=int, default=None, help="Max input sequence length." +@opt( + "--random_seed", + type=int, + default=3407, + group="Data Loading", + help="Random seed.", ) -@click.option("--stride", type=int, default=None, help="Step size for sliding window.") -@click.option("--dpo_beta", type=float, default=0.1, help="DPO beta.") -@click.option("--group_size", type=int, default=4, help="GRPO group size.") -@click.option("--grpo_clip_eps", type=float, default=0.2, help="GRPO clip epsilon.") -@click.option( - "--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient." +@opt( + "--num_workers", + type=int, + default=4, + group="Data Loading", + help="DataLoader workers.", ) -@click.option("--label_smoothing", type=float, default=0.0, help="Label smoothing.") -@click.option( - "--rollout_interval", type=int, default=512, help="Steps between rollouts." +@opt( + "--pin_memory/--no-pin_memory", + default=True, + group="Data Loading", + help="Pin memory.", ) -@click.option( - "--rollout_temperature", type=float, default=0.7, help="Rollout temperature." +@opt( + "--window_size", + type=int, + default=None, + group="Data Loading", + help="Max input sequence length.", ) -@click.option("--rollout_top_k", type=int, default=0, help="Rollout top-k (0=disable).") -@click.option("--rollout_top_p", type=float, default=0.9, help="Rollout top-p.") -@click.option( +@opt( + "--stride", + type=int, + default=None, + group="Data Loading", + help="Step size for sliding window.", +) +@opt("--dpo_beta", type=float, default=0.1, group="Algorithm", help="DPO beta.") +@opt("--group_size", type=int, default=4, group="Algorithm", help="GRPO group size.") +@opt( + "--grpo_clip_eps", + type=float, + default=0.2, + group="Algorithm", + help="GRPO clip epsilon.", +) +@opt( + "--grpo_kl_coef", + type=float, + default=0.01, + group="Algorithm", + help="GRPO KL penalty coefficient.", +) +@opt( + "--label_smoothing", + type=float, + default=0.0, + group="Data Loading", + help="Label smoothing.", +) +@opt( + "--rollout_interval", + type=int, + default=512, + group="Algorithm", + help="Steps between rollouts.", +) +@opt( + "--rollout_temperature", + type=float, + default=0.7, + group="Algorithm", + help="Rollout temperature.", +) +@opt( + "--rollout_top_k", + type=int, + default=0, + group="Algorithm", + help="Rollout top-k (0=disable).", +) +@opt( + "--rollout_top_p", + type=float, + default=0.9, + group="Algorithm", + help="Rollout top-p.", +) +@opt( "--rollout_max_tokens", type=int, default=1024, + group="Algorithm", help="Max tokens per rollout response.", ) -@click.option( +@opt( "--gradient_checkpointing/--no-gradient_checkpointing", default=False, + group="Misc", help="Enable activation checkpointing.", ) -@click.option( +@opt( "--compile", "compile_mode", type=click.Choice(["default", "reduce-overhead", "max-autotune"]), default=None, + group="Misc", help="torch.compile mode. Omit to disable.", ) -@click.option( - "--ckpt_interval", type=int, default=5000, help="Steps between checkpoints." +@opt( + "--ckpt_interval", + type=int, + default=5000, + group="Checkpoint", + help="Steps between checkpoints.", ) -@click.option( - "--ckpt_dir", type=click.Path(), default="checkpoint", help="Checkpoint directory." +@opt( + "--ckpt_dir", + type=click.Path(), + default="checkpoint", + group="Checkpoint", + help="Checkpoint directory.", ) -@click.option("--val_split", type=float, default=None, help="Validation split ratio.") -@click.option( - "--val_step", type=int, default=1000, help="Steps between validation runs." +@opt( + "--val_split", + type=float, + default=None, + group="Validation", + help="Validation split ratio.", ) -@click.option( +@opt( + "--val_step", + type=int, + default=1000, + group="Validation", + help="Steps between validation runs.", +) +@opt( "--metrics", multiple=True, default=("loss", "lr", "grad_norm"), + group="Validation", help="Metrics to log (repeatable).", ) -@click.option("--start_epoch", type=int, default=0, help="Start epoch.") -@click.option("--start_samples", type=int, default=0, help="Start samples (per rank).") -@click.option( - "--master_addr", type=str, default="localhost", help="Master node address." +@opt("--start_epoch", type=int, default=0, group="Checkpoint", help="Start epoch.") +@opt( + "--start_samples", + type=int, + default=0, + group="Checkpoint", + help="Start samples (per rank).", ) -@click.option("--master_port", type=str, default="29500", help="Master node port.") -@click.option( +@opt( + "--master_addr", + type=str, + default="localhost", + group="Distributed", + help="Master node address.", +) +@opt( + "--master_port", + type=str, + default="29500", + group="Distributed", + help="Master node port.", +) +@opt( "--backend", type=click.Choice(_BACKENDS), default="nccl", + group="Distributed", help="Distributed backend.", ) -@click.option("--nprocs", type=int, default=1, help="Number of GPUs.") -@click.option( +@opt("--nprocs", type=int, default=1, group="Distributed", help="Number of GPUs.") +@opt( "--parallel_mode", type=click.Choice(_PARALLEL), default="fsdp", + group="Distributed", help="Parallel strategy.", ) -@click.option("--device_type", type=str, default="cuda", help="Device type.") -@click.option( +@opt( + "--device_type", + type=str, + default="cuda", + group="Distributed", + help="Device type.", +) +@opt( "--start_method", type=click.Choice(_START_METHODS), default="spawn", + group="Distributed", help="Multiprocessing start method.", ) -@click.option("--neftune_alpha", type=float, default=0.0, help="NEFTune noise alpha.") -@click.option( +@opt( + "--neftune_alpha", + type=float, + default=0.0, + group="Algorithm", + help="NEFTune noise alpha.", +) +@opt( "--schedule_type", type=click.Choice(_SCHEDULES), default="cosine", + group="LR Schedule", help="LR scheduler.", ) -@click.option( - "--min_rate", type=float, default=None, help="Minimum LR as fraction of base LR." +@opt( + "--min_rate", + type=float, + default=None, + group="LR Schedule", + help="Minimum LR as fraction of base LR.", ) -@click.option("--cycle_length", type=int, default=None, help="SGDR first cycle length.") -@click.option("--t_mult", type=int, default=2, help="SGDR cycle length multiplier.") -@click.option( - "--stable_steps", type=int, default=None, help="WSD stable plateau steps." +@opt( + "--cycle_length", + type=int, + default=None, + group="LR Schedule", + help="SGDR first cycle length.", ) -@click.option("--decay_steps", type=int, default=None, help="WSD decay steps.") -@click.option("--tp_size", type=int, default=None, help="Tensor parallelism (future).") -@click.option( +@opt( + "--t_mult", + type=int, + default=2, + group="LR Schedule", + help="SGDR cycle length multiplier.", +) +@opt( + "--stable_steps", + type=int, + default=None, + group="LR Schedule", + help="WSD stable plateau steps.", +) +@opt( + "--decay_steps", + type=int, + default=None, + group="LR Schedule", + help="WSD decay steps.", +) +@opt( + "--tp_size", + type=int, + default=None, + group="Distributed", + help="Tensor parallelism (future).", +) +@opt( "--dry-run", is_flag=True, default=False, + group="Misc", help="Validate config and print plan, do not train.", ) @click.pass_context