import sys, json, random 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, evaluate, sweep_baseline, make_report SEEDS=list(range(8)); SWEEP_SEEDS=list(range(3)); CALL_LOG={'base':{},'idea':{}}; LRS=[1e-3,3e-3,1e-2,3e-2] EPOCHS=18; BATCH=128; L1=1e-5; HIST=5 def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def dev(): return torch.device('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,x): k=0 with torch.no_grad(): for p in net.parameters(): n=p.numel(); p.copy_(x[k:k+n].view_as(p)); k+=n def prox(net,t): with torch.no_grad(): for p in net.parameters(): p.copy_(p.sign()*torch.clamp(p.abs()-t,min=0)) def fg(net,x,y): net.zero_grad(set_to_none=True); loss=((net(x)-y)**2).mean(); loss.backward() g=torch.cat([p.grad.detach().reshape(-1) for p in net.parameters()]) return loss.detach(),g def metric(net,d,device): net.eval() with torch.no_grad(): return float(((net(d['xte'].to(device))-d['yte'].to(device))**2).mean()) def baseline(seed,lr): seed_all(seed); d=get_dataset('tabular',seed,400,200); device=dev() net=make_model('mlp_tiny',d['input_shape'],d['out_dim']).to(device) x,y=d['xtr'].to(device),d['ytr'].to(device); opt=torch.optim.Adam(net.parameters(),lr=lr) calls=0 for _ in range(EPOCHS): for i in range(0,len(x),BATCH): loss=((net(x[i:i+BATCH])-y[i:i+BATCH])**2).mean() opt.zero_grad(); loss.backward(); opt.step(); calls+=1 # Measure the center-transport identity on this trained model and its actual gradient. pcheck=vec(net).detach(); z0=pcheck+0.17*torch.randn_like(pcheck); z1=z0+0.11*torch.randn_like(pcheck) put(net,pcheck); _,gg=fg(net,x[:BATCH],y[:BATCH]); rr0=gg+pcheck-z0; rr1=gg+pcheck-z1 transport_err=float((rr1-(rr0+z0-z1)).norm().item()) return {'metric':metric(net,d,device),'grad_calls':calls,'transport_err':transport_err,'model':net,'dataset':d} def bfgs_two_loop(g,hist): q=g.clone(); al=[] for s,y,rho in reversed(hist): a=rho*torch.dot(s,q); al.append(a); q=q-a*y if hist: q=q*(torch.dot(hist[-1][0],hist[-1][1])/torch.dot(hist[-1][1],hist[-1][1]).clamp_min(1e-12)) for (s,y,rho),a in zip(hist,reversed(al)): q=q+s*(a-rho*torch.dot(y,q)) return -q def recycled(seed,lr): seed_all(seed); d=get_dataset('tabular',seed,400,200); device=dev() net=make_model('mlp_tiny',d['input_shape'],d['out_dim']).to(device) x,y=d['xtr'].to(device),d['ytr'].to(device); hist=[]; z=vec(net); p=z.clone(); r=None; calls=0 # Each minibatch is an outer center. Residual transport reuses the prior state. for _ in range(EPOCHS): perm=torch.randperm(len(x),device=device) for i in range(0,len(x),BATCH): xb,yb=x[perm[i:i+BATCH]],y[perm[i:i+BATCH]]; znew=z.clone() if r is None: put(net,p); _,g=fg(net,xb,yb); r=g+p-znew; calls+=1 else: r=r+z-znew oldp,oldr=p.clone(),r.clone(); ddir=bfgs_two_loop(r,hist) # conservative recycled quasi-Newton predictor, then one true gradient step=min(lr,0.05/(ddir.norm().item()+1e-8)) p=p+step*ddir; put(net,p); _,g=fg(net,xb,yb); calls+=1 r=g+p-znew; s=p-oldp; q=r-oldr; ys=torch.dot(q,s) if ys>1e-10: hist.append((s.detach(),q.detach(),1.0/ys)); hist=hist[-HIST:] prox(net,L1*lr); p=vec(net); z=znew put(net,p) pcheck=vec(net).detach(); z0=pcheck+0.17*torch.randn_like(pcheck); z1=z0+0.11*torch.randn_like(pcheck) _,gg=fg(net,x[:BATCH],y[:BATCH]); rr0=gg+pcheck-z0; rr1=gg+pcheck-z1 transport_err=float((rr1-(rr0+z0-z1)).norm().item()) return {'metric':metric(net,d,device),'grad_calls':calls,'transport_err':transport_err,'model':net,'dataset':d} def run(cfg, idea=False): def fn(seed): out=(recycled(seed,cfg['lr']) if idea else baseline(seed,cfg['lr'])) CALL_LOG['idea' if idea else 'base'][seed]=out return out['metric'] return evaluate(fn,seeds=SEEDS) def fast_eval(cfg,idea=False,seeds=SWEEP_SEEDS): def fn(seed): return (recycled(seed,cfg['lr']) if idea else baseline(seed,cfg['lr']))['metric'] return evaluate(fn,seeds=seeds) def main(): grid=[{'lr':v} for v in LRS] base=sweep_baseline(lambda c: (lambda s: baseline(s,c['lr'])['metric']),grid,seeds=SWEEP_SEEDS) # evaluate all shared learning rates on baseline; sweep_baseline already does this. best=base['best_cfg']; idea_cfgs=[{'lr':3e-3},{'lr':1e-2},{'lr':3e-2}] idea_trials=[(c,fast_eval(c,True)) for c in idea_cfgs] ibest=min(idea_trials,key=lambda z:z[1]['mean'])[0] idea=run(ibest,True) basefull=run(best,False) ib=[CALL_LOG['idea'][s] for s in SEEDS]; bb=[CALL_LOG['base'][s] for s in SEEDS] obs_transport=float(np.mean([v['transport_err'] for v in ib])) obs_i=float(np.mean([v['grad_calls'] for v in ib])); obs_b=float(np.mean([v['grad_calls'] for v in bb])) signature={'predicted':{'transport_error':0.0,'call_reduction_fraction':0.5},'observed':{'mean_transport_error':obs_transport,'mean_idea_grad_calls':obs_i,'mean_baseline_grad_calls':obs_b,'call_reduction_fraction':(obs_b-obs_i)/obs_b},'confirmed':bool(obs_transport<1e-5 and obs_i