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:
@@ -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).",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user