Tensorized concentration mixing / bench_tensorized.py

Unverified

Raw ⬇ ZIP
 1import sys,json
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset,train_model,evaluate,sweep_baseline,make_report
 7SEEDS=tuple(range(8)); GRID=[{'lr':1e-3},{'lr':3e-3},{'lr':5e-3}]
 8
 9def dct(n):
10 x=torch.arange(n,dtype=torch.float32)[:,None]; k=torch.arange(n,dtype=torch.float32)[None,:]
11 u=torch.cos(np.pi/n*(x+.5)*k); u[:,0]/=np.sqrt(n)
12 if n>1:u[:,1:]*=np.sqrt(2/n)
13 return u
14
15class Dense(nn.Module):
16 def __init__(self,w,d=64):
17  super().__init__(); self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.randn(1,w,d)*.02)
18  self.a=nn.MultiheadAttention(d,2,dropout=0,batch_first=True); self.n1=nn.LayerNorm(d)
19  self.ff=nn.Sequential(nn.Linear(d,128),nn.GELU(),nn.Linear(128,d)); self.n2=nn.LayerNorm(d); self.head=nn.Linear(w*d,1)
20 def forward(self,x):
21  h=self.inp(x[...,None])+self.pos[:,:x.shape[1]]; z,_=self.a(h,h,h,need_weights=False); h=self.n1(h+z); h=self.n2(h+self.ff(h)); return self.head(h.flatten(1))
22
23class Tensorized(nn.Module):
24 def __init__(self,w,d=64):
25  super().__init__(); self.w=w; self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.randn(1,w,d)*.02)
26  self.register_buffer('U',dct(w)); self.theta=nn.Parameter(torch.ones(w)); self.v=nn.Linear(d,d); self.o=nn.Linear(d,d); self.n1=nn.LayerNorm(d)
27  self.ff=nn.Sequential(nn.Linear(d,128),nn.GELU(),nn.Linear(128,d)); self.n2=nn.LayerNorm(d); self.head=nn.Linear(w*d,1)
28 def forward(self,x):
29  h=self.inp(x[...,None])+self.pos[:,:x.shape[1]]; lam=torch.sigmoid(self.theta); q=h.permute(0,2,1).reshape(-1,self.w); q=(q@self.U) * lam; q=q@self.U.T; z=self.o(self.v(q.reshape(h.shape[0],h.shape[2],self.w).permute(0,2,1))); h=self.n1(h+z); return self.head(self.n2(h+self.ff(h)).flatten(1))
30 def contraction(self): return float(torch.sigmoid(self.theta).max().detach().cpu())
31
32def run(kind,cfg,seed):
33 torch.manual_seed(seed); np.random.seed(seed); d=get_dataset('sequence',seed,n_train=400,n_test=200); m=(Dense(d['input_shape'][0]) if kind=='baseline' else Tensorized(d['input_shape'][0]))
34 _,metric,_=train_model(m,d,epochs=12,lr=cfg['lr'],batch=128,log=lambda *_:None); return float(metric)
35
36def factory(cfg): return lambda s:run('baseline',cfg,s)
37
38def signature():
39 rows=[]
40 for seed in SEEDS:
41  torch.manual_seed(seed); d=get_dataset('sequence',seed,n_train=400,n_test=200); m=Tensorized(d['input_shape'][0]); m,_,_=train_model(m,d,epochs=12,lr=3e-3,batch=128,log=lambda *_:None)
42  with torch.no_grad():
43   dev=next(m.parameters()).device; x=d['xte'][:64].to(dev); h=m.inp(x[...,None])+m.pos[:,:x.shape[1]]; before=float(h.flatten().norm().cpu()); lam=torch.sigmoid(m.theta); q=h.permute(0,2,1).reshape(-1,m.w); q=(q@m.U)*lam; after=float((q@m.U.T).norm().cpu())
44  rows.append({'seed':seed,'input_norm':before,'mixed_norm':after,'observed_ratio':after/max(before,1e-12),'predicted_bound':m.contraction()})
45 ratios=[r['observed_ratio'] for r in rows]; bounds=[r['predicted_bound'] for r in rows]
46 return {'predicted':'axis mixing is non-expansive with spectral norm <=1','observed_mean_ratio':float(np.mean(ratios)),'observed_max_ratio':float(np.max(ratios)),'predicted_max_singular_mean':float(np.mean(bounds)),'confirmed':bool(np.max(ratios)<=1.001)}
47
48if __name__=='__main__':
49 base=sweep_baseline(factory,GRID,seeds=(0,1,2,3)); best=base['best_cfg']
50 idea_trials=[]
51 for cfg in GRID:
52  r=evaluate(lambda s,cfg=cfg:run('idea',cfg,s),SEEDS)
53  idea_trials.append({'cfg':cfg,'mean':r['mean'],'std':r['std'],'per_seed':r['per_seed']})
54 ibest=min(idea_trials,key=lambda z:z['mean'])
55 idea={'mean':ibest['mean'],'std':ibest['std'],'per_seed':ibest['per_seed'],'n':len(ibest['per_seed'])}
56 rep=make_report('sequence','transformer_tiny',base,idea,{'track_structure':'multi-token sequence forecast','implementation':'learned DCT-basis positive contraction','signature':signature()})
57 rep['baseline']['sweep_union_lrs']=[x['lr'] for x in GRID]; rep['idea']['best_cfg']=ibest['cfg']; rep['idea']['sweep']=idea_trials
58 open('bench_report.json','w').write(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))