diff options
| author | yurenh <blackhao0426@gmail.com> | 2026-08-31 18:14:09 -0500 |
|---|---|---|
| committer | yurenh <blackhao0426@gmail.com> | 2026-08-31 18:14:09 -0500 |
| commit | 6a544fabfc2af22e4d5823410dd2387b5af89ea9 (patch) | |
| tree | 0abd67bdda420deed27428b621fb59db8be07f41 /src/zbp_scaling/zbp/metrics.py | |
scaffold: model (OLMo2-ish + ZBP partition), trainer (DDP/config), data shards, bench
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe
Diffstat (limited to 'src/zbp_scaling/zbp/metrics.py')
| -rw-r--r-- | src/zbp_scaling/zbp/metrics.py | 61 |
1 files changed, 61 insertions, 0 deletions
diff --git a/src/zbp_scaling/zbp/metrics.py b/src/zbp_scaling/zbp/metrics.py new file mode 100644 index 0000000..99c3841 --- /dev/null +++ b/src/zbp_scaling/zbp/metrics.py @@ -0,0 +1,61 @@ +"""Bias / variance diagnostics for gradient estimators.""" +import math +import torch + + +def cosine(a, b): + a = a.reshape(a.shape[0], -1) if a.dim() > 1 else a.reshape(1, -1) + b = b.reshape(b.shape[0], -1) if b.dim() > 1 else b.reshape(1, -1) + return torch.nn.functional.cosine_similarity(a, b, dim=1) + + +def relerr(a, b): + a = a.reshape(a.shape[0], -1) + b = b.reshape(b.shape[0], -1) + return (a - b).norm(dim=1) / b.norm(dim=1).clamp_min(1e-30) + + +def audit(samples, g): + """samples: [M, *shape] independent estimates; g: exact [*shape]. + + Returns dict with relative bias, variance ratio, a chi-square test of zero bias, + and the distribution of cosines. + """ + M = samples.shape[0] + S = samples.reshape(M, -1).double() + g = g.reshape(-1).double() + mean = S.mean(0) + bias = mean - g + var = S.var(0, unbiased=True) # per-coordinate variance + gn2 = g.dot(g).clamp_min(1e-300) + rel_bias = (bias.norm() / gn2.sqrt()).item() + var_ratio = (var.sum() / gn2).item() # E||g_hat - mean||^2 / ||g||^2 + # chi-square statistic of H0: bias = 0, per coordinate z_j = bias_j / (sd_j / sqrt(M)), + # restricted to coordinates that are actually random (exact coordinates have zero variance) + active = var > 1e-12 * var.max().clamp_min(1e-300) + se = (var[active] / M).sqrt().clamp_min(1e-300) + z = bias[active] / se + d = max(int(z.numel()), 1) + chi2 = (z * z).sum().item() + # under H0 each z_j is Student-t with nu = M-1 dof: E[t^2] = nu/(nu-2), Var[t^2] = 2 nu^2 (nu-1) / ((nu-2)^2 (nu-4)) + nu = M - 1 + if nu > 4: + mu2 = nu / (nu - 2) + v2 = 2 * nu ** 2 * (nu - 1) / ((nu - 2) ** 2 * (nu - 4)) + else: + mu2, v2 = 1.0, 2.0 + zscore = (chi2 - d * mu2) / math.sqrt(d * v2) # approx N(0,1) under H0 for large d + # relative bias expected under H0 (pure noise): sqrt(sum var / M) / ||g|| + rel_bias_null = ((var.sum() / M).sqrt() / gn2.sqrt()).item() + cos = torch.nn.functional.cosine_similarity(S, g[None], dim=1) + return { + "M": M, "d": d, "d_total": g.numel(), + "rel_bias": rel_bias, + "rel_bias_null": rel_bias_null, + "bias_ratio": rel_bias / max(rel_bias_null, 1e-300), + "var_ratio": var_ratio, + "chi2": chi2, "chi2_z": zscore, + "cos_mean": cos.mean().item(), "cos_std": cos.std().item(), + "cos_q10": cos.quantile(0.1).item(), "cos_q90": cos.quantile(0.9).item(), + "cos_of_mean": torch.nn.functional.cosine_similarity(mean[None], g[None], dim=1).item(), + } |
