From d6bfb098636faffb62814a2a6b94c9f510c6c7e5 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sat, 1 Aug 2026 08:54:44 +0800 Subject: [PATCH] feat: add grad_snr metric with EMA-based gradient SNR tracking - add GradSNRTracker to metric_util.py computing SNR = E[g]^2 / Var(g) via per-parameter EMA moments - add grad_snr_tracker field to TrainContext (instantiated by default) - register grad_snr in MetricCallback, update tracker on each optimizer step before metrics are recorded - add grad_snr to default --metrics in train.py CLI --- astrai/trainer/metric_util.py | 52 ++++++++++++++++++++++++++++++++ astrai/trainer/train_callback.py | 4 +++ astrai/trainer/train_context.py | 2 ++ scripts/tools/train.py | 2 +- 4 files changed, 59 insertions(+), 1 deletion(-) diff --git a/astrai/trainer/metric_util.py b/astrai/trainer/metric_util.py index e81d9e9..8b89bbe 100644 --- a/astrai/trainer/metric_util.py +++ b/astrai/trainer/metric_util.py @@ -22,6 +22,51 @@ def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, fl return total_sq.sqrt().item() +class GradSNRTracker: + """Track gradient signal-to-noise ratio via EMA of first/second moments. + + SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2) + + The tracker accumulates per-parameter EMA moments across optimizer steps. + Call ``update`` after backward (before ``optimizer.step``) and read + ``snr`` to get the aggregate SNR across all parameters. + """ + + def __init__(self, beta: float = 0.999, eps: float = 1e-8): + self.beta = beta + self.eps = eps + self._first: Dict[int, torch.Tensor] = {} + self._second: Dict[int, torch.Tensor] = {} + + @torch.no_grad() + def update(self, model: nn.Module) -> None: + beta = self.beta + for param in model.parameters(): + if param.grad is None: + continue + pid = id(param) + g = param.grad.detach() + if pid not in self._first: + self._first[pid] = g.clone() + self._second[pid] = g.pow(2).clone() + else: + self._first[pid].mul_(beta).add_(g, alpha=1 - beta) + self._second[pid].mul_(beta).addcmul_(g, g, value=1 - beta) + + @property + def snr(self) -> float: + if not self._first: + return 0.0 + total_signal = 0.0 + total_noise = 0.0 + for m, v in zip(self._first.values(), self._second.values()): + signal = m.pow(2).sum().item() + noise = (v - m.pow(2)).clamp(min=0).sum().item() + total_signal += signal + total_noise += noise + return total_signal / (total_noise + self.eps) + + def ctx_get_loss(ctx): return ctx.loss @@ -36,3 +81,10 @@ def ctx_get_val_loss(ctx): def ctx_get_grad_norm(ctx): return ctx.grad_norm + + +def ctx_get_grad_snr(ctx): + tracker = getattr(ctx, "grad_snr_tracker", None) + if tracker is None: + return None + return tracker.snr diff --git a/astrai/trainer/train_callback.py b/astrai/trainer/train_callback.py index b39aec2..d477421 100644 --- a/astrai/trainer/train_callback.py +++ b/astrai/trainer/train_callback.py @@ -18,6 +18,7 @@ from astrai.parallel.setup import get_current_device from astrai.serialization import Checkpoint from astrai.trainer.metric_util import ( ctx_get_grad_norm, + ctx_get_grad_snr, ctx_get_loss, ctx_get_lr, ctx_get_val_loss, @@ -255,6 +256,7 @@ class MetricCallback(TrainCallback): "lr": ctx_get_lr, "val_loss": ctx_get_val_loss, "grad_norm": ctx_get_grad_norm, + "grad_snr": ctx_get_grad_snr, } def _metrics(self, context: TrainContext, names): @@ -312,6 +314,8 @@ class MetricCallback(TrainCallback): f.write(json.dumps(log) + "\n") def on_optimizer_step(self, context): + context.grad_snr_tracker.update(context.model) + if ( context.val_dataloader is not None and self.val_step > 0 diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index 0ff7bfd..0bdb984 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -17,6 +17,7 @@ from astrai.parallel.setup import get_current_device, get_rank, get_world_size from astrai.protocols import OptimizerProtocol, SchedulerProtocol from astrai.serialization import Checkpoint, load_json from astrai.tokenize import AutoTokenizer +from astrai.trainer.metric_util import GradSNRTracker from astrai.trainer.rollout import RolloutGenerator, RolloutRunner from astrai.trainer.strategy import BaseStrategy, StrategyFactory @@ -38,6 +39,7 @@ class TrainContext: consumed_samples: int = field(default=0) loss: float = field(default=0.0) grad_norm: Optional[float] = field(default=None) + grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker) val_dataloader: Optional[DataLoader] = field(default=None) val_loss: Optional[float] = field(default=None) diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 235c329..5c1dc54 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -369,7 +369,7 @@ _START_METHODS = ["spawn", "fork", "forkserver"] @opt( "--metrics", multiple=True, - default=("loss", "lr", "grad_norm"), + default=("loss", "lr", "grad_norm", "grad_snr"), group="Validation", help="Metrics to log (repeatable).", )