KL-TopK Activation Bottleneck / stage2_kl_topk_bench.py

Failed on benchmark

Raw ⬇ ZIP
 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()