summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
authorVoid Agent <void@jayrup.hermes>2026-08-02 14:22:13 +0100
committerVoid Agent <void@jayrup.hermes>2026-08-02 14:22:13 +0100
commita16d9b3177248ae9e27b195f117b5b615af6da8c (patch)
treee0e9e3c1d21ce2a0fa51f6aa1d31a477cd14e173 /src
parent2ba0e14c3559e5786c324a89f26f159363a230b5 (diff)
Add src/stats_decomp.py: reproducible W_U decomposition, p-values, Spearman, synthetic-pair CI (backs results.md numbers)
Diffstat (limited to 'src')
-rw-r--r--src/stats_decomp.py116
1 files changed, 116 insertions, 0 deletions
diff --git a/src/stats_decomp.py b/src/stats_decomp.py
new file mode 100644
index 0000000..f82e720
--- /dev/null
+++ b/src/stats_decomp.py
@@ -0,0 +1,116 @@
+"""Decomposition + robustness statistics for the committed results (results.md).
+
+Computes, from the saved jlens_v3 artifacts and the base checkpoint:
+ 1. W_U row-norm decomposition: r(||W_U[k]||, freq), r vs log-frequency,
+ Spearman, and per-layer partial correlation of faithful norm with
+ frequency after regressing out ||W_U[k]||.
+ 2. Pearson p-values (normal-approx two-sided) and Spearman per layer.
+ 3. Synthetic-pair ratio with across-seed mean/SD and bootstrap CI.
+
+Run inside the meru container where artifacts exist:
+ python3 src/stats_decomp.py
+"""
+import os
+import pickle
+import sys
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+import numpy as np
+import torch
+
+V = 65
+TRAIN_BIN = 'data/shakespeare_char/train.bin'
+CKPT = 'out-shakespeare-char/ckpt.pt'
+
+
+def rankdata(a):
+ """Average ranks (ties averaged) — minimal scipy-free rankdata."""
+ sorter = np.argsort(a, kind='mergesort')
+ ranks = np.empty(len(a), dtype=float)
+ ranks[sorter] = np.arange(1, len(a) + 1)
+ i = 0
+ while i < len(a):
+ j = i
+ while j + 1 < len(a) and a[sorter[j + 1]] == a[sorter[i]]:
+ j += 1
+ if j > i:
+ ranks[sorter[i:j + 1]] = (i + j) / 2.0 + 1
+ i = j + 1
+ return ranks
+
+
+def spearman(x, y):
+ return np.corrcoef(rankdata(x), rankdata(y))[0, 1]
+
+
+def pearson_p(r, n):
+ """Two-sided p for Pearson r via normal approx on Fisher's z (fine for n=65)."""
+ import math
+ from math import erf, sqrt
+ z = abs(r) * math.sqrt((n - 2) / max(1 - r * r, 1e-12))
+ return 2 * (1 - 0.5 * (1 + erf(z / sqrt(2))))
+
+
+def load_layer(l):
+ for p in (f'outputs/jlens_v3/layer{l}.pt', f'outputs/jlens_v3_layer{l}.pt'):
+ if os.path.exists(p):
+ return torch.load(p, map_location='cpu')
+ raise FileNotFoundError(f'no jlens artifact for layer {l}')
+
+
+def main():
+ train = np.memmap(TRAIN_BIN, dtype=np.uint16, mode='r')
+ counts = np.bincount(train, minlength=V).astype(float)
+ freq = counts / counts.sum() * 100
+ logf = np.log10(freq)
+
+ ckpt = torch.load(CKPT, map_location='cpu')
+ wu = ckpt['model']['lm_head.weight'].float().numpy()
+ wu_norms = np.linalg.norm(wu, axis=1)
+
+ print("=== W_U row-norm decomposition ===")
+ print(f"r(||W_U[k]||, freq) = {np.corrcoef(wu_norms, freq)[0, 1]:+.3f}")
+ print(f"r(||W_U[k]||, log10 freq) = {np.corrcoef(wu_norms, logf)[0, 1]:+.3f}")
+ print(f"Spearman(||W_U[k]||, freq) = {spearman(wu_norms, freq):+.3f}")
+
+ print("\n=== per-layer decomposition ===")
+ print(f"{'layer':<6}{'r(faith,freq)':>14}{'p':>10}{'r(faith,WU)':>13}"
+ f"{'Spear(faith)':>14}{'partial|WU':>13}")
+ for l in range(6):
+ d = load_layer(l)
+ fn = np.array([d['faithful_norms'][k] for k in range(V)])
+ rf = np.corrcoef(fn, freq)[0, 1]
+ rfw = np.corrcoef(fn, wu_norms)[0, 1]
+ sf = spearman(fn, freq)
+ A = np.vstack([wu_norms, np.ones(V)]).T
+ resid = fn - A @ np.linalg.lstsq(A, fn, rcond=None)[0]
+ r_part = np.corrcoef(resid, freq)[0, 1]
+ print(f"L{l:<5}{rf:>+14.3f}{pearson_p(rf, V):>10.1e}{rfw:>+13.3f}"
+ f"{sf:>+14.3f}{r_part:>+13.3f}")
+
+ print("\n=== synthetic pair ratio ===")
+ with open('data/synth_pair/meta.pkl', 'rb') as f:
+ meta = pickle.load(f)
+ iid_s = meta['stoi']['@']
+ iid_n = meta['stoi']['#']
+ ratios = []
+ for s in range(3):
+ vals_s, vals_n = [], []
+ for l in (2, 3, 4):
+ d = torch.load(f'outputs/synth_pair/seed{s}/layer{l}.pt',
+ map_location='cpu')
+ fn = d['faithful_norms']
+ vals_s.append(fn[iid_s])
+ vals_n.append(fn[iid_n])
+ ratios.append(np.mean(vals_s) / max(np.mean(vals_n), 1e-9))
+ r_arr = np.array(ratios)
+ rng = np.random.RandomState(0)
+ boot = [np.mean(rng.choice(r_arr, 3, replace=True)) for _ in range(10000)]
+ print(f"ratio per seed: {[f'{r:.3f}' for r in ratios]}")
+ print(f"mean +/- SD: {r_arr.mean():.3f} +/- {r_arr.std(ddof=1):.3f}")
+ print(f"bootstrap 95% CI: [{np.percentile(boot, 2.5):.3f}, "
+ f"{np.percentile(boot, 97.5):.3f}]")
+
+
+if __name__ == '__main__':
+ main()