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