fix: report gradient snr in db
This commit is contained in:
@@ -1,3 +1,4 @@
|
|||||||
|
import math
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -27,6 +28,8 @@ class GradSNRTracker:
|
|||||||
|
|
||||||
SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2)
|
SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2)
|
||||||
|
|
||||||
|
The reported value is the power ratio in decibels: ``10 * log10(SNR)``.
|
||||||
|
|
||||||
The tracker accumulates per-parameter EMA moments across optimizer steps.
|
The tracker accumulates per-parameter EMA moments across optimizer steps.
|
||||||
Call ``update`` after backward (before ``optimizer.step``) and read
|
Call ``update`` after backward (before ``optimizer.step``) and read
|
||||||
``snr`` to get the aggregate SNR across all parameters.
|
``snr`` to get the aggregate SNR across all parameters.
|
||||||
@@ -64,7 +67,8 @@ class GradSNRTracker:
|
|||||||
noise = (v - m.pow(2)).clamp(min=0).sum().item()
|
noise = (v - m.pow(2)).clamp(min=0).sum().item()
|
||||||
total_signal += signal
|
total_signal += signal
|
||||||
total_noise += noise
|
total_noise += noise
|
||||||
return total_signal / (total_noise + self.eps)
|
snr = total_signal / (total_noise + self.eps)
|
||||||
|
return 10.0 * math.log10(max(snr, self.eps))
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_loss(ctx):
|
def ctx_get_loss(ctx):
|
||||||
|
|||||||
@@ -0,0 +1,17 @@
|
|||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.trainer.metric_util import GradSNRTracker
|
||||||
|
|
||||||
|
|
||||||
|
def test_grad_snr_is_reported_in_decibels():
|
||||||
|
model = torch.nn.Linear(1, 1, bias=False)
|
||||||
|
tracker = GradSNRTracker(beta=0.5, eps=1e-8)
|
||||||
|
|
||||||
|
model.weight.grad = torch.tensor([[1.0]])
|
||||||
|
tracker.update(model)
|
||||||
|
model.weight.grad = torch.tensor([[3.0]])
|
||||||
|
tracker.update(model)
|
||||||
|
|
||||||
|
# E[g]^2 / Var(g) = 4 / 1 = 4, which is 6.0206 dB.
|
||||||
|
assert tracker.snr == pytest.approx(10.0 * torch.log10(torch.tensor(4.0)).item())
|
||||||
Reference in New Issue
Block a user