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
+4
View File
@@ -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