summaryrefslogtreecommitdiff
path: root/check_proxy.py
blob: 83d16021ba475d726803ffddcd9c084ed32dc54e (plain)
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")