From e5f9b1a3a92ce43e5efe6bd7f9e4aaf822e02cf1 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Mon, 20 Jul 2026 01:08:13 +0800 Subject: [PATCH] fix: default max_grad_norm to 1.0 and drop None branch --- assets/docs/params.md | 2 +- astrai/config/train_config.py | 2 +- astrai/parallel/executor.py | 13 ++----------- scripts/tools/train.py | 2 +- 4 files changed, 5 insertions(+), 14 deletions(-) diff --git a/assets/docs/params.md b/assets/docs/params.md index 53c547e..7f10942 100644 --- a/assets/docs/params.md +++ b/assets/docs/params.md @@ -26,7 +26,7 @@ |-----------|-------------|---------| | `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 | | `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 | -| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | None | +| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 | ### Optimizer (MuonMix) diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index f03a47c..65e9eb9 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -38,7 +38,7 @@ class TrainConfig(BaseConfig): default=1, metadata={"help": "Number of iterations between steps."} ) max_grad_norm: Optional[float] = field( - default=None, + default=1.0, metadata={"help": "Maximum gradient norm. None disables clipping."}, ) gradient_checkpointing_modules: List[str] = field( diff --git a/astrai/parallel/executor.py b/astrai/parallel/executor.py index af0a70e..7b4c112 100644 --- a/astrai/parallel/executor.py +++ b/astrai/parallel/executor.py @@ -148,14 +148,7 @@ class BaseExecutor: def grad_accum_steps(self) -> int: return self.gradient_state.num_steps - def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float: - if max_norm is None: - total_norm = torch.norm( - torch.stack( - [p.grad.norm(2) for p in model.parameters() if p.grad is not None] - ) - ) - return total_norm.item() + def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float: total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) if isinstance(total_norm, torch.Tensor): return total_norm.item() @@ -289,9 +282,7 @@ class FSDPExecutor(BaseExecutor): return model.no_sync() return contextlib.nullcontext() - def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float: - if max_norm is None: - return super().clip_grad_norm(model, max_norm) + def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float: if isinstance(model, FSDP) and self.use_distributed: total_norm = model.clip_grad_norm_(max_norm) if isinstance(total_norm, torch.Tensor): diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 0eb03f5..a3594e6 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -148,7 +148,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--max_grad_norm", type=float, - default=None, + default=1.0, help="Max gradient norm for clipping. None disables clipping.", ) parser.add_argument(