From d7695b40e328576f1089ce62be1b0079a80a3d15 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 19 Jul 2026 00:08:18 +0800 Subject: [PATCH] feat: make max_grad_norm optional (None disables clipping) - TrainConfig.max_grad_norm defaults to None - executor.clip_grad_norm returns grad norm without clipping when None - train.py --max_grad_norm defaults to None --- astrai/config/train_config.py | 5 +++-- astrai/parallel/executor.py | 18 ++++++++++++++++-- scripts/tools/train.py | 4 ++-- 3 files changed, 21 insertions(+), 6 deletions(-) diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index fcefa12..f03a47c 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -37,8 +37,9 @@ class TrainConfig(BaseConfig): grad_accum_steps: int = field( default=1, metadata={"help": "Number of iterations between steps."} ) - max_grad_norm: float = field( - default=1.0, metadata={"help": "Maximum gradient norm."} + max_grad_norm: Optional[float] = field( + default=None, + metadata={"help": "Maximum gradient norm. None disables clipping."}, ) gradient_checkpointing_modules: List[str] = field( default_factory=list, diff --git a/astrai/parallel/executor.py b/astrai/parallel/executor.py index 2d43a73..84f3c4c 100644 --- a/astrai/parallel/executor.py +++ b/astrai/parallel/executor.py @@ -148,7 +148,14 @@ class BaseExecutor: def grad_accum_steps(self) -> int: return self.gradient_state.num_steps - def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float: + 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() total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) if isinstance(total_norm, torch.Tensor): return total_norm.item() @@ -289,7 +296,14 @@ class FSDPExecutor(BaseExecutor): return model.no_sync() return contextlib.nullcontext() - def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float: + 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() 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 e35c56e..0eb03f5 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -148,8 +148,8 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--max_grad_norm", type=float, - default=1.0, - help="Max gradient norm for clipping.", + default=None, + help="Max gradient norm for clipping. None disables clipping.", ) parser.add_argument( "--weight_decay",