Sign-Reset PI Optimizer / bench_pi.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys,json,random
 2from pathlib import Path
 3import numpy as np, torch
 4sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 5from bench import get_dataset,make_model,sweep_baseline,evaluate,make_report
 6SEEDS=tuple(range(8)); LR=[.001,.003,.01]; MOM=[0.,.9]; EPOCHS=20; BATCH=128; NTR=800; NTE=300
 7
 8def seed(s):
 9 random.seed(s); np.random.seed(s); torch.manual_seed(s)
10 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
11def dev(): return torch.device('cuda' if torch.cuda.is_available() else 'cpu')
12def loss(m,x,y): return ((m(x)-y)**2).mean()
13def run_base(c,s):
14 seed(s); d=get_dataset('tabular',s,n_train=NTR,n_test=NTE); m=make_model('mlp_tiny',d['input_shape'],d['out_dim']).to(dev()); o=torch.optim.SGD(m.parameters(),lr=c['lr'],momentum=c['momentum']); x,y=d['xtr'].to(dev()),d['ytr'].to(dev()); g=torch.Generator().manual_seed(s+12345)
15 for _ in range(EPOCHS):
16  for ii in torch.randperm(len(x),generator=g).split(BATCH):
17   o.zero_grad(set_to_none=True); z=loss(m,x[ii],y[ii]); z.backward(); o.step()
18 m.eval()
19 with torch.no_grad(): return float(loss(m,d['xte'].to(dev()),d['yte'].to(dev())).cpu())
20def run_pi(c,s,diag=False):
21 seed(s); d=get_dataset('tabular',s,n_train=NTR,n_test=NTE); m=make_model('mlp_tiny',d['input_shape'],d['out_dim']).to(dev()); x,y=d['xtr'].to(dev()),d['ytr'].to(dev()); g=torch.Generator().manual_seed(s+12345); I=[torch.zeros_like(p) for p in m.parameters()]; prev=None; resets=0; norms=[]
22 for _ in range(EPOCHS):
23  for ii in torch.randperm(len(x),generator=g).split(BATCH):
24   m.zero_grad(set_to_none=True); z=loss(m,x[ii],y[ii]); z.backward(); gs=[p.grad.detach().clone() for p in m.parameters()]; dot=sum((a*b).sum() for a,b in zip(gs,prev)) if prev is not None else None
25   if dot is not None and float(dot)<0: I=[torch.zeros_like(p) for p in m.parameters()]; resets+=1
26   else: I=[a+b for a,b in zip(I,gs)]
27   with torch.no_grad():
28    for p,a,b in zip(m.parameters(),gs,I): p-=c['lr']*(c['kp']*a+c['ki']*b)
29   prev=gs; norms.append(float(torch.sqrt(sum((q*q).sum() for q in I)).cpu()))
30 m.eval()
31 with torch.no_grad(): metric=float(loss(m,d['xte'].to(dev()),d['yte'].to(dev())).cpu())
32 return {'metric':metric,'resets':resets,'mean_integral_norm':float(np.mean(norms)),'final_integral_norm':norms[-1]}
33def main():
34 baseline_grid=[{'lr':lr,'momentum':mo} for lr in sorted(set(LR+[x*.5 for x in LR]+[x*2 for x in LR]+[.004])) for mo in MOM]
35 b=sweep_baseline(lambda c:(lambda s:run_base(c,s)),baseline_grid,seeds=(0,1,2,3))
36 # Three nearby PI settings around the tuned baseline learning rate; all rates were in baseline_grid.
37 blr=b['best_cfg']['lr']; idea_grid=[{'lr':blr*.5,'kp':1.,'ki':.2},{'lr':blr,'kp':1.,'ki':.2},{'lr':blr*2,'kp':1.,'ki':.2}]
38 besti=min(idea_grid,key=lambda c:np.mean([run_pi(c,s)['metric'] for s in (0,1,2,3)]))
39 vals=[]; sig=[]
40 for s in SEEDS:
41  q=run_pi(besti,s,True); vals.append(q['metric']); sig.append(q)
42 idea={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals),'diagnostics':sig,'cfg':besti}
43 rep=make_report('tabular','mlp_tiny',b,idea,{'predicted_reset_effect':'global gradient sign reversal resets all integral tensors to zero','observed_reset_rate':float(np.mean([q['resets'] for q in sig])/ (EPOCHS*(NTR//BATCH+1))),'observed_mean_integral_norm':float(np.mean([q['mean_integral_norm'] for q in sig])),'confirmed':bool(np.mean([q['resets'] for q in sig])>0)})
44 Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
45if __name__=='__main__': main()