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