import sys,json,random from pathlib import Path import numpy as np, torch torch.set_num_threads(1) import torch.nn as nn sys.path.insert(0,'/home/maxwelhelp/all/math2nn') from bench import get_dataset,make_model,sweep_baseline,make_report from bench.protocol import evaluate,DEFAULT_SEEDS OUT=Path('bench_report.json') def seed(s): random.seed(s);np.random.seed(s);torch.manual_seed(s) def ds(s): return get_dataset('tabular',s,100,50) def base_fn(c): def run(s): seed(s);d=ds(s);m=make_model('mlp_tiny',d['input_shape'],1);o=torch.optim.Adam(m.parameters(),lr=c['lr'],weight_decay=c['wd']);x,y=d['xtr'],d['ytr']; lf=nn.MSELoss() for _ in range(c['epochs']): q=torch.randperm(len(x)) for i in range(0,len(x),128): z=q[i:i+128];o.zero_grad();lf(m(x[z]),y[z]).backward();o.step() with torch.no_grad():return float(((m(d['xte'])-d['yte'])**2).mean()) return run def chunks(m,size=64): n=sum(p.numel() for p in m.parameters());return [list(range(i,min(i+size,n))) for i in range(0,n,size)] def flatgrad(m):return torch.cat([p.grad.reshape(-1) for p in m.parameters()]) def setgrad(m,g): k=0 for p in m.parameters(): n=p.numel();p.grad.copy_(g[k:k+n].reshape_as(p));k+=n def phi(e,t): if t<1e-12:return 0. return 1. if e==0 else 2*t/(e+np.sqrt(e*e+4*t*t)) def make_factors(m,x,y,groups,c): # Per-example gradient sketch, but only block covariance is materialized. rows=[];lf=nn.MSELoss() for j in range(min(c['sketch'],len(x))): m.zero_grad(set_to_none=True);lf(m(x[j:j+1]),y[j:j+1]).backward();rows.append(flatgrad(m).detach()) J=torch.stack(rows); fac=[]; cert=[] for i,g in enumerate(groups): a=J[:,g]; B=(a.T@a)/len(J)+c['damp']*torch.eye(len(g));fac.append(torch.linalg.inv(B)) if i: cross=J[:,g].T@J[:,groups[i-1]]/len(J); eps=float(torch.linalg.matrix_norm(cross,2)); cert.append(phi(0.,eps)*np.sqrt(2)*eps) return J,fac,cert def idea_run(c,s,info=False): seed(s);d=ds(s);m=make_model('mlp_tiny',d['input_shape'],1);x,y=d['xtr'],d['ytr'];g=chunks(m,c['block']);J,fac,cs=make_factors(m,x,y,g,c) # Conservative eta=0 is used when estimated spectral intervals overlap. active=[list(z) for z in g] # Merge adjacent blocks if certified distortion exceeds tau; rebuild merged factors. while len(active)>1: merged=False for i in range(len(active)-1): a,b=active[i],active[i+1];cross=J[:,b].T@J[:,a]/len(J);eps=float(torch.linalg.matrix_norm(cross,2));cert=np.sqrt(2)*eps if cert>c['tau']: active[i]=a+b;active.pop(i+1);merged=True;break if not merged:break fac=[] for z in active: A=J[:,z].T@J[:,z]/len(J)+c['damp']*torch.eye(len(z));fac.append(torch.linalg.inv(A)) lf=nn.MSELoss() for _ in range(c['epochs']): q=torch.randperm(len(x)) for i in range(0,len(x),128): z=q[i:i+128];m.zero_grad();lf(m(x[z]),y[z]).backward();v=flatgrad(m);out=torch.zeros_like(v) for b,F in zip(active,fac):out[b]=F@v[b] setgrad(m,out) with torch.no_grad(): for p in m.parameters():p.add_(p.grad,alpha=-c['lr']) with torch.no_grad(): val=float(((m(d['xte'])-d['yte'])**2).mean()) if not info:return val # trained-model signature: compare full small projected Fisher spectrum to block deletion. # Projection uses 12 fixed random directions, avoiding an unmaterialized full matrix. K=min(12,J.shape[1]);R=torch.randn(J.shape[1],K);JR=J@R;full=JR.T@JR/len(J);bd=torch.zeros_like(J) for b in chunks(m,c['block']):bd[:,b]=J[:,b] block= (bd@R).T@(bd@R)/len(J);obs=float(torch.linalg.matrix_norm(full-block,2));pred=float(max(cs or [0.])) return val,{'blocks':len(active),'predicted_certificate':pred,'observed_projected_shift':obs} def main(): lrs=[.001,.003,.01];grid=[{'lr':r,'epochs':1,'wd':w} for r in lrs for w in [0.,1e-4]] base=sweep_baseline(base_fn,grid,seeds=(0,1,2,3)) ig=[{'lr':r,'epochs':1,'block':64,'tau':t,'damp':.01,'sketch':3} for r in lrs for t in [.02]] tuned=[(evaluate(lambda s,c=c:idea_run(c,s),seeds=(0,1,2,3))['mean'],c) for c in ig];chosen=min(tuned,key=lambda z:z[0])[1] vals=[];sign=[] for s in DEFAULT_SEEDS: v,z=idea_run(chosen,s,True);vals.append(v);z['seed']=s;sign.append(z) idea={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':8} pred=[z['predicted_certificate'] for z in sign];obs=[z['observed_projected_shift'] for z in sign];cor=float(np.corrcoef(pred,obs)[0,1]) if np.std(pred)>0 and np.std(obs)>0 else 0. extra={'prediction':'larger cross-block certificate predicts larger observed projected Fisher spectral distortion','predicted_values':pred,'observed_values':obs,'correlation':cor,'confirmed':bool(cor>.5)} rep=make_report('tabular','mlp_tiny',base,idea,extra);rep['idea_sweep']=[{'cfg':c,'mean':float(v)} for v,c in tuned];rep['chosen_idea_cfg']=chosen;rep['protocol_notes']='Tabular matches an optimizer/preconditioning intervention; both systems use identical mlp_tiny and paired seeds. Fisher uses a small per-example gradient sketch and block-local factors.';OUT.write_text(json.dumps(rep,indent=2));print(json.dumps(rep,indent=2)) if __name__=='__main__':main()