summaryrefslogtreecommitdiff
path: root/check_proxy.py
diff options
context:
space:
mode:
authorVoid Agent <void@jayrup.hermes>2026-07-31 16:22:45 +0100
committerVoid Agent <void@jayrup.hermes>2026-07-31 16:22:45 +0100
commit3986b9f5e7e7efc0bd10143a860e19e1b889fe60 (patch)
tree2d20f55d23fcb83b1b45dffb135d838b65e2808c /check_proxy.py
parent84e5c0a215acccd6c8e6c52a9c4d08bc8c5c5b93 (diff)
Add faithful J-lens (jlens_v3): W_U-probed residual Jacobian per paper; both-ways comparison vs log-softmax proxy; 3-model adversarial reviews
Diffstat (limited to 'check_proxy.py')
-rw-r--r--check_proxy.py41
1 files changed, 41 insertions, 0 deletions
diff --git a/check_proxy.py b/check_proxy.py
new file mode 100644
index 0000000..83d1602
--- /dev/null
+++ b/check_proxy.py
@@ -0,0 +1,41 @@
+"""Verify batched proxy == per-token proxy on the same model (numerical agreement)."""
+import sys, os
+sys.path.insert(0, '.')
+sys.path.insert(0, 'src')
+import torch, numpy as np
+from model import GPT, GPTConfig
+import jlens_v3
+
+torch.manual_seed(0)
+cfg = GPTConfig(n_layer=3, n_head=4, n_embd=16, block_size=32, bias=False,
+ vocab_size=65, dropout=0.0)
+model = GPT(cfg).eval()
+data = np.random.randint(0, 65, 2000).astype(np.uint16)
+batches = jlens_v3.make_batches(data, 32, 4, 3, 'cpu')
+
+# batched (new) proxy
+batched = jlens_v3.compute_proxy_norms(model, 1, batches, 'cpu', 65)
+
+# per-token (old) proxy, inline
+d = model.config.n_embd
+accum = torch.zeros(65, d)
+counts = torch.zeros(65)
+for (x, y) in batches:
+ B, T = x.shape
+ rc = {}
+ def hook(m, i, o): rc['val'] = o
+ h = model.transformer.h[1].register_forward_hook(hook)
+ logits, loss = model(x, y)
+ h.remove()
+ lp = torch.nn.functional.log_softmax(logits, dim=-1)
+ for k in range(65):
+ g = torch.autograd.grad(lp[:, :, k].sum(), rc['val'],
+ retain_graph=(k < 64))[0]
+ accum[k] += g.detach().cpu().reshape(-1, d).sum(dim=0)
+ counts[k] += B * T
+per_token = {k: (accum[k] / counts[k]).norm().item() for k in range(65)}
+
+maxdiff = max(abs(batched[k] - per_token[k]) for k in range(65))
+print(f"max |batched - per_token| over 65 tokens: {maxdiff:.6e}")
+assert maxdiff < 1e-4, "proxy implementations disagree!"
+print("PROXY VARIANT AGREEMENT: PASSED")