import sys, json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report # Spectrally safeguarded block DFP. Blocks make the intervention usable at NN scale. def flat(m, grad=False): a=[] for p in m.parameters(): z = p.grad if grad else p a.append((torch.zeros_like(p) if z is None and grad else z).detach().reshape(-1)) return torch.cat(a).cpu().numpy().astype(np.float64) def put(m, x): k=0 with torch.no_grad(): for p in m.parameters(): n=p.numel(); p.copy_(torch.as_tensor(x[k:k+n],dtype=p.dtype,device=p.device).reshape_as(p)); k+=n def dfp(H,s,y): sy=float(s@y); Hy=H@y; den=float(y@Hy) 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 def safe_update(H,s,y,P,eps,tau,kappa): raw=dfp(H,s,y); raw=(raw+raw.T)/2 w,V=np.linalg.eigh(raw); r=min(2,len(w)); Q=V[:,:r]@V[:,:r].T q=0. if P is None else float(np.linalg.norm(Q-P,2)) unsafe=bool(w[0]kappa or q>tau) if unsafe: # Blend toward the previous safe matrix, then always project to the floor. best=H for rho in (.5,.25,.1,0.): B=(1-rho)*H+rho*raw; ew=np.linalg.eigvalsh((B+B.T)/2) if ew[0]>=eps and ew[-1]<=kappa*eps: best=B; break raw=best w,V=np.linalg.eigh((raw+raw.T)/2); w=np.clip(w,eps,kappa*eps) return (V*w)@V.T,Q,q,unsafe def run_idea(ds, lr=0.08, epochs=8, eps=1e-4, tau=.75, kappa=1e4, block=256): seed=int(ds.get('_seed',0)); torch.manual_seed(seed+1000); np.random.seed(seed+1000) m=make_model('mlp_tiny',ds['input_shape'],ds['out_dim']); lossf=nn.MSELoss() n=sum(p.numel() for p in m.parameters()); H=[]; P=[] for a in range(0,n,block): z=min(block,n-a); H.append(np.eye(z)*1.0); P.append(None) xtr,ytr=ds['xtr'],ds['ytr']; lastg=None; lastx=None; safeguards=0; eigmin=1e99 for ep in range(epochs): # full-batch keeps gradient evaluations and secant pairs unambiguous m.zero_grad(); l=lossf(m(xtr),ytr); l.backward(); g=flat(m,True); x=flat(m) d=np.zeros(n) for j,a in enumerate(range(0,n,block)): d[a:a+len(H[j])]=-H[j]@g[a:a+len(H[j])] alpha=lr # conservative Armijo backtracking, same objective and model as baseline f0=float(l); gd=float(g@d) while alpha>1e-7: put(m,x+alpha*d); f=float(lossf(m(xtr),ytr)) if f<=f0+1e-4*alpha*gd: break alpha*=.5 xn=x+alpha*d; put(m,xn); m.zero_grad(); ln=lossf(m(xtr),ytr); ln.backward(); gn=flat(m,True) s=xn-x; y=gn-g for j,a in enumerate(range(0,n,block)): sb=s[a:a+len(H[j])]; yb=y[a:a+len(H[j])]; sy=float(sb@yb) if sy>1e-10*np.linalg.norm(sb)*np.linalg.norm(yb): H[j],P[j],q,tr=safe_update(H[j],sb,yb,P[j],eps,tau,kappa); safeguards+=int(tr) eigmin=min(eigmin,float(np.linalg.eigvalsh(H[j])[0])) m.eval() with torch.no_grad(): metric=float(((m(ds['xte'])-ds['yte'])**2).mean()) return metric, {'safeguards':safeguards,'min_eig':eigmin} def make_ds(seed): d=get_dataset('tabular',seed,n_train=120,n_test=60); d['_seed']=seed; return d def baseline_one(cfg, seed): d=make_ds(seed); torch.manual_seed(seed+1000); np.random.seed(seed+1000) m=make_model('mlp_tiny',d['input_shape'],d['out_dim']) _,metric,_=train_model(m,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128) return metric def main(): torch.set_num_threads(4); seeds=tuple(range(8)) grid=[{'lr':x,'epochs':2} for x in (.001,.003,.01)] base=sweep_baseline(lambda c: (lambda s: baseline_one(c,s)),grid,seeds=tuple(range(4))) cfg=base['best_cfg']; settings=[cfg,{'lr':cfg['lr']*.5,'epochs':cfg['epochs']},{'lr':cfg['lr']*2,'epochs':cfg['epochs']}] allsets=[] for c in settings: vals=[]; aux=[] for s in seeds: v,a=run_idea(make_ds(s),lr=c['lr'],epochs=c['epochs'],block=2048); vals.append(v); aux.append(a) allsets.append({'cfg':c,'per_seed':vals,'aux':aux,'mean':float(np.mean(vals))}) best=min(allsets,key=lambda z:z['mean']); idea={'per_seed':best['per_seed'],'best_cfg':best['cfg'],'sweep':allsets} 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}) Path('bench_report.json').write_text(json.dumps(rep,indent=2)) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()