def grad_sq_norm(model, X, Y):
logits = model(X)
loss = F.cross_entropy(logits.reshape(-1, logits.shape[-1]),
Y.reshape(-1))
grads = torch.autograd.grad(loss, list(model.parameters()))
return sum((g ** 2).sum() for g in grads)
def noise_scale(model, X, Y, b_small=16, b_big=2048, m=400):
def mean_sq_norm(b, m):
idx = torch.randint(0, len(Y), (m, b), device=X.device)
return torch.stack([grad_sq_norm(model, X[i], Y[i])
for i in idx]).mean()
n_small, n_big = mean_sq_norm(b_small, m), mean_sq_norm(b_big, m // 8)
tr_sigma = (n_small - n_big) / (1 / b_small - 1 / b_big)
sq_norm = (b_big * n_big - b_small * n_small) / (b_big - b_small)
return (tr_sigma / sq_norm).item()