import sys, json, time 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 MODEL='mlp_tiny'; EPOCHS=16; BATCH=128; LAMBDA=2e-4; CURV=0.02 SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)) def dev(): return 'cuda' if torch.cuda.is_available() else 'cpu' def vec(net): return torch.cat([p.detach().reshape(-1) for p in net.parameters()]) def put(net,z): k=0 with torch.no_grad(): for p in net.parameters(): q=p.numel(); p.copy_(z[k:k+q].view_as(p)); k+=q def obj(net,x,y,kind): z=nn.functional.mse_loss(net(x),y) if kind=='g': z=z+LAMBDA*sum(torch.sqrt(p*p+1e-8).sum() for p in net.parameters())+0.5*CURV*sum((p*p).sum() for p in net.parameters()) elif kind=='h': z=0.5*CURV*sum((p*p).sum() for p in net.parameters()) return z def home_train(model,ds,epochs,lr,p,gamma,K=3,return_sig=False): try: d=dev(); model.to(d); x=ds['xtr'].to(d); y=ds['ytr'].to(d); n=len(x); rng=np.random.default_rng(991) sig=[] for ep in range(epochs): order=rng.permutation(n) gam=gamma*(0.97**ep) for st in range(0,n,BATCH): ix=torch.as_tensor(order[st:st+BATCH],device=d); xb=x[ix]; yb=y[ix]; s=vec(model) us=s.clone(); vs=s.clone() for _ in range(K): put(model,us); model.zero_grad(set_to_none=True); q=obj(model,xb,yb,'g'); q.backward(); gg=vec(model.grad_proxy) if False else torch.cat([(p_.grad if p_.grad is not None else torch.zeros_like(p_)).reshape(-1) for p_ in model.parameters()]) du=us-s; ng=du.norm(); pg=du*(ng.clamp_min(1e-12)**(p-2))/gam; us=us-(0.012 if p==2 else 0.003)*(gg+pg) for _ in range(K): put(model,vs); model.zero_grad(set_to_none=True); q=obj(model,xb,yb,'h'); q.backward(); gh=torch.cat([(p_.grad if p_.grad is not None else torch.zeros_like(p_)).reshape(-1) for p_ in model.parameters()]) dv=vs-s; nv=dv.norm(); ph=dv*(nv.clamp_min(1e-12)**(p-2))/gam; vs=vs-(0.012 if p==2 else 0.003)*(gh+ph) ag=(s-us)*(s-us).norm().clamp_min(1e-12)**(p-2)/gam; ah=(s-vs)*(s-vs).norm().clamp_min(1e-12)**(p-2)/gam; G=ag-ah # record predicted envelope gradient attenuation against raw DC gradient put(model,s); model.zero_grad(set_to_none=True); q=obj(model,xb,yb,'g')-obj(model,xb,yb,'h'); q.backward(); raw=torch.cat([(p_.grad if p_.grad is not None else torch.zeros_like(p_)).reshape(-1) for p_ in model.parameters()]) sig.append((float(G.norm()),float(raw.norm()))) gn=G.norm().clamp_min(1e-12); G=G*torch.clamp(5.0/gn,max=1.0); put(model,s-lr*G) model.eval() with torch.no_grad(): metric=float(nn.functional.mse_loss(model(xte:=ds['xte'].to(d)),ds['yte'].to(d)).cpu()) out={'metric':metric,'sig':sig} return out if return_sig else metric except Exception: return {'metric':float('nan'),'sig':[]} if return_sig else float('nan') def baseline_fn(cfg): def f(seed): torch.manual_seed(seed); np.random.seed(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=400); m=make_model(MODEL,ds['input_shape'],ds['out_dim']) _,metric,_=train_model(m,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['wd'],log=lambda *a:None); return metric return f def idea_fn(cfg, detailed=False): def f(seed, want_detail=False): torch.manual_seed(seed); np.random.seed(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=400); m=make_model(MODEL,ds['input_shape'],ds['out_dim']) return home_train(m,ds,EPOCHS,cfg['lr'],cfg['p'],cfg['gamma'],K=3,return_sig=want_detail) return f def main(): # Union parity: every idea lr occurs in baseline grid; baseline's weight decay is swept too. grid=[{'lr':lr,'wd':wd} for lr in (0.001,0.003,0.006) for wd in (0.0,1e-4)] base=sweep_baseline(baseline_fn,grid,seeds=SWEEP_SEEDS) best_lr=base['best_cfg']['lr']; idea_grid=[{'lr':best_lr,'p':2,'gamma':g} for g in (0.03,0.1,0.3)] idea_grid += [{'lr':lr,'p':4,'gamma':0.1} for lr in (0.001,0.003,0.006)] # Evaluate all idea settings on four seeds for selection, then best on all eight. tried=[] for c in idea_grid: r=evaluate(lambda s: idea_fn(c)(s),seeds=SWEEP_SEEDS); tried.append({'cfg':c,'mean':r['mean'],'n':r['n']}) finite=[z for z in tried if np.isfinite(z['mean']) and z.get('n',0)==len(SWEEP_SEEDS)] best=min(finite,key=lambda z:z['mean'])['cfg']; ir=evaluate(idea_fn(best),seeds=SEEDS) # Signature is measured from trained models, not analytic identities. pairs=[] for s in SEEDS: z=idea_fn(best)(s,True); pairs.extend(z['sig'][-10:]) ratios=[a/(b+1e-12) for a,b in pairs if np.isfinite(a+b)] sig={'quantity':'||G_HOME|| / ||raw DC gradient|| on minibatches','predicted':'smoothing should attenuate updates','observed_mean_ratio':float(np.mean(ratios)) if ratios else None,'observed_median_ratio':float(np.median(ratios)) if ratios else None,'n':len(ratios),'confirmed':bool(ratios and np.mean(ratios)<1.0)} rep=make_report('tabular',MODEL,base,ir,{'mechanism_signature':sig,'idea_sweep':tried,'selection_cfg':best,'equal_budget':{'epochs':EPOCHS,'batch':BATCH}}) with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()