Recycled-curvature proximal optimizer / bench_experiment.py
Failed on benchmark
1import sys, json, random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
7
8SEEDS=list(range(8)); SWEEP_SEEDS=list(range(3)); CALL_LOG={'base':{},'idea':{}}; LRS=[1e-3,3e-3,1e-2,3e-2]
9EPOCHS=18; BATCH=128; L1=1e-5; HIST=5
10
11def seed_all(s):
12 random.seed(s); np.random.seed(s); torch.manual_seed(s)
13 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
14
15def dev(): return torch.device('cuda' if torch.cuda.is_available() else 'cpu')
16
17def vec(net): return torch.cat([p.detach().reshape(-1) for p in net.parameters()])
18def put(net,x):
19 k=0
20 with torch.no_grad():
21 for p in net.parameters():
22 n=p.numel(); p.copy_(x[k:k+n].view_as(p)); k+=n
23
24def prox(net,t):
25 with torch.no_grad():
26 for p in net.parameters(): p.copy_(p.sign()*torch.clamp(p.abs()-t,min=0))
27
28def fg(net,x,y):
29 net.zero_grad(set_to_none=True); loss=((net(x)-y)**2).mean(); loss.backward()
30 g=torch.cat([p.grad.detach().reshape(-1) for p in net.parameters()])
31 return loss.detach(),g
32
33def metric(net,d,device):
34 net.eval()
35 with torch.no_grad(): return float(((net(d['xte'].to(device))-d['yte'].to(device))**2).mean())
36
37def baseline(seed,lr):
38 seed_all(seed); d=get_dataset('tabular',seed,400,200); device=dev()
39 net=make_model('mlp_tiny',d['input_shape'],d['out_dim']).to(device)
40 x,y=d['xtr'].to(device),d['ytr'].to(device); opt=torch.optim.Adam(net.parameters(),lr=lr)
41 calls=0
42 for _ in range(EPOCHS):
43 for i in range(0,len(x),BATCH):
44 loss=((net(x[i:i+BATCH])-y[i:i+BATCH])**2).mean()
45 opt.zero_grad(); loss.backward(); opt.step(); calls+=1
46 # Measure the center-transport identity on this trained model and its actual gradient.
47 pcheck=vec(net).detach(); z0=pcheck+0.17*torch.randn_like(pcheck); z1=z0+0.11*torch.randn_like(pcheck)
48 put(net,pcheck); _,gg=fg(net,x[:BATCH],y[:BATCH]); rr0=gg+pcheck-z0; rr1=gg+pcheck-z1
49 transport_err=float((rr1-(rr0+z0-z1)).norm().item())
50 return {'metric':metric(net,d,device),'grad_calls':calls,'transport_err':transport_err,'model':net,'dataset':d}
51
52def bfgs_two_loop(g,hist):
53 q=g.clone(); al=[]
54 for s,y,rho in reversed(hist):
55 a=rho*torch.dot(s,q); al.append(a); q=q-a*y
56 if hist: q=q*(torch.dot(hist[-1][0],hist[-1][1])/torch.dot(hist[-1][1],hist[-1][1]).clamp_min(1e-12))
57 for (s,y,rho),a in zip(hist,reversed(al)):
58 q=q+s*(a-rho*torch.dot(y,q))
59 return -q
60
61def recycled(seed,lr):
62 seed_all(seed); d=get_dataset('tabular',seed,400,200); device=dev()
63 net=make_model('mlp_tiny',d['input_shape'],d['out_dim']).to(device)
64 x,y=d['xtr'].to(device),d['ytr'].to(device); hist=[]; z=vec(net); p=z.clone(); r=None; calls=0
65 # Each minibatch is an outer center. Residual transport reuses the prior state.
66 for _ in range(EPOCHS):
67 perm=torch.randperm(len(x),device=device)
68 for i in range(0,len(x),BATCH):
69 xb,yb=x[perm[i:i+BATCH]],y[perm[i:i+BATCH]]; znew=z.clone()
70 if r is None:
71 put(net,p); _,g=fg(net,xb,yb); r=g+p-znew; calls+=1
72 else: r=r+z-znew
73 oldp,oldr=p.clone(),r.clone(); ddir=bfgs_two_loop(r,hist)
74 # conservative recycled quasi-Newton predictor, then one true gradient
75 step=min(lr,0.05/(ddir.norm().item()+1e-8))
76 p=p+step*ddir; put(net,p); _,g=fg(net,xb,yb); calls+=1
77 r=g+p-znew; s=p-oldp; q=r-oldr; ys=torch.dot(q,s)
78 if ys>1e-10:
79 hist.append((s.detach(),q.detach(),1.0/ys)); hist=hist[-HIST:]
80 prox(net,L1*lr); p=vec(net); z=znew
81 put(net,p)
82 pcheck=vec(net).detach(); z0=pcheck+0.17*torch.randn_like(pcheck); z1=z0+0.11*torch.randn_like(pcheck)
83 _,gg=fg(net,x[:BATCH],y[:BATCH]); rr0=gg+pcheck-z0; rr1=gg+pcheck-z1
84 transport_err=float((rr1-(rr0+z0-z1)).norm().item())
85 return {'metric':metric(net,d,device),'grad_calls':calls,'transport_err':transport_err,'model':net,'dataset':d}
86
87def run(cfg, idea=False):
88 def fn(seed):
89 out=(recycled(seed,cfg['lr']) if idea else baseline(seed,cfg['lr']))
90 CALL_LOG['idea' if idea else 'base'][seed]=out
91 return out['metric']
92 return evaluate(fn,seeds=SEEDS)
93
94def fast_eval(cfg,idea=False,seeds=SWEEP_SEEDS):
95 def fn(seed): return (recycled(seed,cfg['lr']) if idea else baseline(seed,cfg['lr']))['metric']
96 return evaluate(fn,seeds=seeds)
97
98def main():
99 grid=[{'lr':v} for v in LRS]
100 base=sweep_baseline(lambda c: (lambda s: baseline(s,c['lr'])['metric']),grid,seeds=SWEEP_SEEDS)
101 # evaluate all shared learning rates on baseline; sweep_baseline already does this.
102 best=base['best_cfg']; idea_cfgs=[{'lr':3e-3},{'lr':1e-2},{'lr':3e-2}]
103 idea_trials=[(c,fast_eval(c,True)) for c in idea_cfgs]
104 ibest=min(idea_trials,key=lambda z:z[1]['mean'])[0]
105 idea=run(ibest,True)
106 basefull=run(best,False)
107 ib=[CALL_LOG['idea'][s] for s in SEEDS]; bb=[CALL_LOG['base'][s] for s in SEEDS]
108 obs_transport=float(np.mean([v['transport_err'] for v in ib]))
109 obs_i=float(np.mean([v['grad_calls'] for v in ib])); obs_b=float(np.mean([v['grad_calls'] for v in bb]))
110 signature={'predicted':{'transport_error':0.0,'call_reduction_fraction':0.5},'observed':{'mean_transport_error':obs_transport,'mean_idea_grad_calls':obs_i,'mean_baseline_grad_calls':obs_b,'call_reduction_fraction':(obs_b-obs_i)/obs_b},'confirmed':bool(obs_transport<1e-5 and obs_i<obs_b)}
111 report=make_report('tabular','mlp_tiny',{'best_cfg':best,'sweep':base['sweep'],'full':basefull},idea,signature)
112 report['idea']['chosen_cfg']=ibest; report['idea']['trials']=[{'cfg':c,'mean':r['mean']} for c,r in idea_trials]
113 with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
114 print(json.dumps(report,indent=2))
115if __name__=='__main__': main()