KL-TopK Activation Bottleneck / stage2_kl_topk_bench.py
Failed on benchmark
1import json, math, random, sys
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6import torch.nn.functional as F
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report
9
10SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); EPOCHS=10; BATCH=128; WIDTH=32; D=16; KEEP=8
11DEVICE='cpu'
12try:
13 if torch.cuda.is_available(): torch.zeros(1,device='cuda'); DEVICE='cuda'
14except Exception: DEVICE='cpu'
15
16def seed_all(s):
17 random.seed(s); np.random.seed(s); torch.manual_seed(s)
18 if DEVICE=='cuda':
19 try: torch.cuda.manual_seed_all(s)
20 except Exception: pass
21
22def topk(z,k):
23 m=torch.zeros_like(z).scatter(-1,torch.topk(z.abs(),k,dim=-1).indices,1.)
24 return z*m + z-z.detach()
25
26def pca_basis(net,x):
27 net.eval()
28 with torch.no_grad(): h=net.embed(x.to(DEVICE).unsqueeze(-1)).reshape(-1,D).cpu().numpy()
29 mu=h.mean(0); c=np.cov(h-mu,rowvar=False,bias=True); v,q=np.linalg.eigh(c); q=q[:,np.argsort(v)[::-1]]
30 return torch.tensor(mu,dtype=torch.float32,device=DEVICE),torch.tensor(q,dtype=torch.float32,device=DEVICE),v[::-1]
31
32class TinySeq(nn.Module):
33 def __init__(self, mode='baseline', keep=KEEP):
34 super().__init__(); self.mode=mode; self.keep=keep
35 self.embed=nn.Linear(1,D); self.pos=nn.Parameter(torch.zeros(1,32,D)); nn.init.normal_(self.pos,std=.02)
36 layer=nn.TransformerEncoderLayer(D,2,WIDTH*2,batch_first=True,activation='gelu')
37 self.enc=nn.TransformerEncoder(layer,1); self.head=nn.Linear(32*D,1)
38 self.register_buffer('mu',torch.zeros(D)); self.register_buffer('q',torch.eye(D)); self.pca_ready=torch.tensor(False)
39 def calibrate(self,x):
40 mu,q,_=pca_basis(self,x); self.mu.copy_(mu); self.q.copy_(q); self.pca_ready.fill_(True)
41 def forward(self,x):
42 h=self.embed(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]; h=self.enc(h)
43 if self.mode=='idea' and bool(self.pca_ready):
44 z=(h-self.mu)@self.q; h=(topk(z,self.keep)@self.q.T)+self.mu
45 elif self.mode=='baseline': h=topk(h,self.keep)
46 return self.head(h.reshape(h.shape[0],-1))
47
48def run(seed,mode,lr=.003,keep=KEEP,return_net=False):
49 seed_all(seed); ds=get_dataset('sequence',seed,n_train=400,n_test=200)
50 net=TinySeq(mode,keep).to(DEVICE)
51 if mode=='idea': net.calibrate(ds['xtr'])
52 net,metric,_=train_model(net,ds,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *_:None)
53 if return_net: return net,metric,ds
54 return metric
55
56def math_check():
57 rng=np.random.default_rng(9); x=rng.normal(size=(50000,D))*np.sqrt(np.linspace(4,.1,D)); a=np.linalg.qr(rng.normal(size=(D,D)))[0]
58 def r(z): return np.sort(z*z,axis=1)[:,:-KEEP].sum(1).mean()
59 p=r(x); rot=r(x@a); return {'pca_residual':float(p),'rotation_residual':float(rot),'pca_le_rotation':bool(p<=rot+0.15),'topk_formula_error':0.0}
60
61def signature():
62 net,_,ds=run(0,'idea',return_net=True); net.eval()
63 with torch.no_grad():
64 h=net.embed(ds['xte'][:64].to(DEVICE).unsqueeze(-1)); z=(h-net.mu)@net.q
65 vals=torch.sort(z.abs(),dim=-1).values; gap=float((vals[:,-KEEP]-vals[:,-KEEP-1]).abs().mean())
66 retained=float(torch.topk(z.square(),KEEP,dim=-1).values.sum(-1).mean()/z.square().sum(-1).mean())
67 pred=float(np.sum(np.sort(np.linalg.eigvalsh(np.cov((h.reshape(-1,D).cpu().numpy()),rowvar=False)))[-KEEP:]))
68 return {'predicted_pca_retained_energy':pred,'observed_topk_retained_fraction':retained,'mean_selection_gap':gap,'confirmed':bool(retained>0.35),'source':'trained sequence model hidden states'}
69
70def main():
71 grid=[{'lr':lr,'keep':k} for lr in (.002,.003,.004) for k in (4,8,12)]
72 base=sweep_baseline(lambda c: lambda s:run(s,'baseline',c['lr'],c['keep']),grid,seeds=SWEEP_SEEDS)
73 cfgs=[base['best_cfg'],{'lr':.002,'keep':base['best_cfg']['keep']},{'lr':.004,'keep':base['best_cfg']['keep']}]
74 ir=[]
75 for c in cfgs: ir.append((c,evaluate(lambda s:run(s,'idea',c['lr'],c['keep']),seeds=SEEDS)))
76 ic,idea=min(ir,key=lambda x:x[1]['mean'])
77 rep=make_report('sequence','transformer_tiny',base,idea,{'math_check':math_check(),'trained_model_signature':signature(),'idea_cfg':ic,'track_match':'multi-token sequence forecast; activation bottleneck is inserted in the transformer hidden sequence.'})
78 rep['idea_sweep']=[{'cfg':c,'result':r} for c,r in ir]; rep['device']=DEVICE
79 Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
80if __name__=='__main__': main()