1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
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")
|