"""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")