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
This commit is contained in:
2026-08-01 08:54:44 +08:00
parent 6db276f37a
commit d6bfb09863
4 changed files with 59 additions and 1 deletions
+1 -1
View File
@@ -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).",
)