Spectrally safeguarded DFP / bench_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 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()