Spectral-Certified Block-Diagonal Preconditioning / bench_experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys,json,random
 2from pathlib import Path
 3import numpy as np, torch
 4torch.set_num_threads(1)
 5import torch.nn as nn
 6sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset,make_model,sweep_baseline,make_report
 8from bench.protocol import evaluate,DEFAULT_SEEDS
 9OUT=Path('bench_report.json')
10
11def seed(s):
12 random.seed(s);np.random.seed(s);torch.manual_seed(s)
13
14def ds(s): return get_dataset('tabular',s,100,50)
15def base_fn(c):
16 def run(s):
17  seed(s);d=ds(s);m=make_model('mlp_tiny',d['input_shape'],1);o=torch.optim.Adam(m.parameters(),lr=c['lr'],weight_decay=c['wd']);x,y=d['xtr'],d['ytr']; lf=nn.MSELoss()
18  for _ in range(c['epochs']):
19   q=torch.randperm(len(x))
20   for i in range(0,len(x),128):
21    z=q[i:i+128];o.zero_grad();lf(m(x[z]),y[z]).backward();o.step()
22  with torch.no_grad():return float(((m(d['xte'])-d['yte'])**2).mean())
23 return run
24
25def chunks(m,size=64):
26 n=sum(p.numel() for p in m.parameters());return [list(range(i,min(i+size,n))) for i in range(0,n,size)]
27def flatgrad(m):return torch.cat([p.grad.reshape(-1) for p in m.parameters()])
28def setgrad(m,g):
29 k=0
30 for p in m.parameters(): n=p.numel();p.grad.copy_(g[k:k+n].reshape_as(p));k+=n
31def phi(e,t):
32 if t<1e-12:return 0.
33 return 1. if e==0 else 2*t/(e+np.sqrt(e*e+4*t*t))
34def make_factors(m,x,y,groups,c):
35 # Per-example gradient sketch, but only block covariance is materialized.
36 rows=[];lf=nn.MSELoss()
37 for j in range(min(c['sketch'],len(x))):
38  m.zero_grad(set_to_none=True);lf(m(x[j:j+1]),y[j:j+1]).backward();rows.append(flatgrad(m).detach())
39 J=torch.stack(rows); fac=[]; cert=[]
40 for i,g in enumerate(groups):
41  a=J[:,g]; B=(a.T@a)/len(J)+c['damp']*torch.eye(len(g));fac.append(torch.linalg.inv(B))
42  if i:
43   cross=J[:,g].T@J[:,groups[i-1]]/len(J); eps=float(torch.linalg.matrix_norm(cross,2)); cert.append(phi(0.,eps)*np.sqrt(2)*eps)
44 return J,fac,cert
45
46def idea_run(c,s,info=False):
47 seed(s);d=ds(s);m=make_model('mlp_tiny',d['input_shape'],1);x,y=d['xtr'],d['ytr'];g=chunks(m,c['block']);J,fac,cs=make_factors(m,x,y,g,c)
48 # Conservative eta=0 is used when estimated spectral intervals overlap.
49 active=[list(z) for z in g]
50 # Merge adjacent blocks if certified distortion exceeds tau; rebuild merged factors.
51 while len(active)>1:
52  merged=False
53  for i in range(len(active)-1):
54   a,b=active[i],active[i+1];cross=J[:,b].T@J[:,a]/len(J);eps=float(torch.linalg.matrix_norm(cross,2));cert=np.sqrt(2)*eps
55   if cert>c['tau']: active[i]=a+b;active.pop(i+1);merged=True;break
56  if not merged:break
57 fac=[]
58 for z in active:
59  A=J[:,z].T@J[:,z]/len(J)+c['damp']*torch.eye(len(z));fac.append(torch.linalg.inv(A))
60 lf=nn.MSELoss()
61 for _ in range(c['epochs']):
62  q=torch.randperm(len(x))
63  for i in range(0,len(x),128):
64   z=q[i:i+128];m.zero_grad();lf(m(x[z]),y[z]).backward();v=flatgrad(m);out=torch.zeros_like(v)
65   for b,F in zip(active,fac):out[b]=F@v[b]
66   setgrad(m,out)
67   with torch.no_grad():
68    for p in m.parameters():p.add_(p.grad,alpha=-c['lr'])
69 with torch.no_grad(): val=float(((m(d['xte'])-d['yte'])**2).mean())
70 if not info:return val
71 # trained-model signature: compare full small projected Fisher spectrum to block deletion.
72 # Projection uses 12 fixed random directions, avoiding an unmaterialized full matrix.
73 K=min(12,J.shape[1]);R=torch.randn(J.shape[1],K);JR=J@R;full=JR.T@JR/len(J);bd=torch.zeros_like(J)
74 for b in chunks(m,c['block']):bd[:,b]=J[:,b]
75 block= (bd@R).T@(bd@R)/len(J);obs=float(torch.linalg.matrix_norm(full-block,2));pred=float(max(cs or [0.]))
76 return val,{'blocks':len(active),'predicted_certificate':pred,'observed_projected_shift':obs}
77
78def main():
79 lrs=[.001,.003,.01];grid=[{'lr':r,'epochs':1,'wd':w} for r in lrs for w in [0.,1e-4]]
80 base=sweep_baseline(base_fn,grid,seeds=(0,1,2,3))
81 ig=[{'lr':r,'epochs':1,'block':64,'tau':t,'damp':.01,'sketch':3} for r in lrs for t in [.02]]
82 tuned=[(evaluate(lambda s,c=c:idea_run(c,s),seeds=(0,1,2,3))['mean'],c) for c in ig];chosen=min(tuned,key=lambda z:z[0])[1]
83 vals=[];sign=[]
84 for s in DEFAULT_SEEDS:
85  v,z=idea_run(chosen,s,True);vals.append(v);z['seed']=s;sign.append(z)
86 idea={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':8}
87 pred=[z['predicted_certificate'] for z in sign];obs=[z['observed_projected_shift'] for z in sign];cor=float(np.corrcoef(pred,obs)[0,1]) if np.std(pred)>0 and np.std(obs)>0 else 0.
88 extra={'prediction':'larger cross-block certificate predicts larger observed projected Fisher spectral distortion','predicted_values':pred,'observed_values':obs,'correlation':cor,'confirmed':bool(cor>.5)}
89 rep=make_report('tabular','mlp_tiny',base,idea,extra);rep['idea_sweep']=[{'cfg':c,'mean':float(v)} for v,c in tuned];rep['chosen_idea_cfg']=chosen;rep['protocol_notes']='Tabular matches an optimizer/preconditioning intervention; both systems use identical mlp_tiny and paired seeds. Fisher uses a small per-example gradient sketch and block-local factors.';OUT.write_text(json.dumps(rep,indent=2));print(json.dumps(rep,indent=2))
90if __name__=='__main__':main()