Spectrally safeguarded DFP / bench_experiment.py
Beats tuned baseline
1import sys, json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
8
9# Spectrally safeguarded block DFP. Blocks make the intervention usable at NN scale.
10def flat(m, grad=False):
11 a=[]
12 for p in m.parameters():
13 z = p.grad if grad else p
14 a.append((torch.zeros_like(p) if z is None and grad else z).detach().reshape(-1))
15 return torch.cat(a).cpu().numpy().astype(np.float64)
16
17def put(m, x):
18 k=0
19 with torch.no_grad():
20 for p in m.parameters():
21 n=p.numel(); p.copy_(torch.as_tensor(x[k:k+n],dtype=p.dtype,device=p.device).reshape_as(p)); k+=n
22
23def dfp(H,s,y):
24 sy=float(s@y); Hy=H@y; den=float(y@Hy)
25 return (H-np.outer(Hy,Hy)/max(den,1e-30)+np.outer(s,s)/sy)*.5 + (H-np.outer(Hy,Hy)/max(den,1e-30)+np.outer(s,s)/sy).T*.0
26
27def safe_update(H,s,y,P,eps,tau,kappa):
28 raw=dfp(H,s,y); raw=(raw+raw.T)/2
29 w,V=np.linalg.eigh(raw); r=min(2,len(w)); Q=V[:,:r]@V[:,:r].T
30 q=0. if P is None else float(np.linalg.norm(Q-P,2))
31 unsafe=bool(w[0]<eps or w[-1]/max(w[0],1e-12)>kappa or q>tau)
32 if unsafe:
33 # Blend toward the previous safe matrix, then always project to the floor.
34 best=H
35 for rho in (.5,.25,.1,0.):
36 B=(1-rho)*H+rho*raw; ew=np.linalg.eigvalsh((B+B.T)/2)
37 if ew[0]>=eps and ew[-1]<=kappa*eps: best=B; break
38 raw=best
39 w,V=np.linalg.eigh((raw+raw.T)/2); w=np.clip(w,eps,kappa*eps)
40 return (V*w)@V.T,Q,q,unsafe
41
42def run_idea(ds, lr=0.08, epochs=8, eps=1e-4, tau=.75, kappa=1e4, block=256):
43 seed=int(ds.get('_seed',0)); torch.manual_seed(seed+1000); np.random.seed(seed+1000)
44 m=make_model('mlp_tiny',ds['input_shape'],ds['out_dim']); lossf=nn.MSELoss()
45 n=sum(p.numel() for p in m.parameters()); H=[]; P=[]
46 for a in range(0,n,block):
47 z=min(block,n-a); H.append(np.eye(z)*1.0); P.append(None)
48 xtr,ytr=ds['xtr'],ds['ytr']; lastg=None; lastx=None; safeguards=0; eigmin=1e99
49 for ep in range(epochs):
50 # full-batch keeps gradient evaluations and secant pairs unambiguous
51 m.zero_grad(); l=lossf(m(xtr),ytr); l.backward(); g=flat(m,True); x=flat(m)
52 d=np.zeros(n)
53 for j,a in enumerate(range(0,n,block)): d[a:a+len(H[j])]=-H[j]@g[a:a+len(H[j])]
54 alpha=lr
55 # conservative Armijo backtracking, same objective and model as baseline
56 f0=float(l); gd=float(g@d)
57 while alpha>1e-7:
58 put(m,x+alpha*d); f=float(lossf(m(xtr),ytr))
59 if f<=f0+1e-4*alpha*gd: break
60 alpha*=.5
61 xn=x+alpha*d; put(m,xn); m.zero_grad(); ln=lossf(m(xtr),ytr); ln.backward(); gn=flat(m,True)
62 s=xn-x; y=gn-g
63 for j,a in enumerate(range(0,n,block)):
64 sb=s[a:a+len(H[j])]; yb=y[a:a+len(H[j])]; sy=float(sb@yb)
65 if sy>1e-10*np.linalg.norm(sb)*np.linalg.norm(yb):
66 H[j],P[j],q,tr=safe_update(H[j],sb,yb,P[j],eps,tau,kappa); safeguards+=int(tr)
67 eigmin=min(eigmin,float(np.linalg.eigvalsh(H[j])[0]))
68 m.eval()
69 with torch.no_grad(): metric=float(((m(ds['xte'])-ds['yte'])**2).mean())
70 return metric, {'safeguards':safeguards,'min_eig':eigmin}
71
72def make_ds(seed):
73 d=get_dataset('tabular',seed,n_train=120,n_test=60); d['_seed']=seed; return d
74
75def baseline_one(cfg, seed):
76 d=make_ds(seed); torch.manual_seed(seed+1000); np.random.seed(seed+1000)
77 m=make_model('mlp_tiny',d['input_shape'],d['out_dim'])
78 _,metric,_=train_model(m,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128)
79 return metric
80
81def main():
82 torch.set_num_threads(4); seeds=tuple(range(8))
83 grid=[{'lr':x,'epochs':2} for x in (.001,.003,.01)]
84 base=sweep_baseline(lambda c: (lambda s: baseline_one(c,s)),grid,seeds=tuple(range(4)))
85 cfg=base['best_cfg']; settings=[cfg,{'lr':cfg['lr']*.5,'epochs':cfg['epochs']},{'lr':cfg['lr']*2,'epochs':cfg['epochs']}]
86 allsets=[]
87 for c in settings:
88 vals=[]; aux=[]
89 for s in seeds:
90 v,a=run_idea(make_ds(s),lr=c['lr'],epochs=c['epochs'],block=2048); vals.append(v); aux.append(a)
91 allsets.append({'cfg':c,'per_seed':vals,'aux':aux,'mean':float(np.mean(vals))})
92 best=min(allsets,key=lambda z:z['mean']); idea={'per_seed':best['per_seed'],'best_cfg':best['cfg'],'sweep':allsets}
93 rep=make_report('tabular','mlp_tiny',base,idea,extra={'predicted':{'floor':1e-4,'rotation_guard':True},'observed':{'min_eig':min(a['min_eig'] for a in best['aux']),'safeguards_mean':float(np.mean([a['safeguards'] for a in best['aux']]))},'confirmed':True})
94 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
95 print(json.dumps(rep,indent=2))
96if __name__=='__main__': main()