Inexact High-Order Moreau DC Optimizer / bench_experiment.py
Mechanism confirmed, baseline not beaten
1import sys, json, time
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
7
8MODEL='mlp_tiny'; EPOCHS=16; BATCH=128; LAMBDA=2e-4; CURV=0.02
9SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4))
10
11def dev(): return 'cuda' if torch.cuda.is_available() else 'cpu'
12def vec(net): return torch.cat([p.detach().reshape(-1) for p in net.parameters()])
13def put(net,z):
14 k=0
15 with torch.no_grad():
16 for p in net.parameters():
17 q=p.numel(); p.copy_(z[k:k+q].view_as(p)); k+=q
18
19def obj(net,x,y,kind):
20 z=nn.functional.mse_loss(net(x),y)
21 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())
22 elif kind=='h': z=0.5*CURV*sum((p*p).sum() for p in net.parameters())
23 return z
24
25def home_train(model,ds,epochs,lr,p,gamma,K=3,return_sig=False):
26 try:
27 d=dev(); model.to(d); x=ds['xtr'].to(d); y=ds['ytr'].to(d); n=len(x); rng=np.random.default_rng(991)
28 sig=[]
29 for ep in range(epochs):
30 order=rng.permutation(n)
31 gam=gamma*(0.97**ep)
32 for st in range(0,n,BATCH):
33 ix=torch.as_tensor(order[st:st+BATCH],device=d); xb=x[ix]; yb=y[ix]; s=vec(model)
34 us=s.clone(); vs=s.clone()
35 for _ in range(K):
36 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()])
37 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)
38 for _ in range(K):
39 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()])
40 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)
41 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
42 # record predicted envelope gradient attenuation against raw DC gradient
43 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()])
44 sig.append((float(G.norm()),float(raw.norm())))
45 gn=G.norm().clamp_min(1e-12); G=G*torch.clamp(5.0/gn,max=1.0); put(model,s-lr*G)
46 model.eval()
47 with torch.no_grad(): metric=float(nn.functional.mse_loss(model(xte:=ds['xte'].to(d)),ds['yte'].to(d)).cpu())
48 out={'metric':metric,'sig':sig}
49 return out if return_sig else metric
50 except Exception:
51 return {'metric':float('nan'),'sig':[]} if return_sig else float('nan')
52
53def baseline_fn(cfg):
54 def f(seed):
55 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'])
56 _,metric,_=train_model(m,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['wd'],log=lambda *a:None); return metric
57 return f
58
59def idea_fn(cfg, detailed=False):
60 def f(seed, want_detail=False):
61 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'])
62 return home_train(m,ds,EPOCHS,cfg['lr'],cfg['p'],cfg['gamma'],K=3,return_sig=want_detail)
63 return f
64
65def main():
66 # Union parity: every idea lr occurs in baseline grid; baseline's weight decay is swept too.
67 grid=[{'lr':lr,'wd':wd} for lr in (0.001,0.003,0.006) for wd in (0.0,1e-4)]
68 base=sweep_baseline(baseline_fn,grid,seeds=SWEEP_SEEDS)
69 best_lr=base['best_cfg']['lr']; idea_grid=[{'lr':best_lr,'p':2,'gamma':g} for g in (0.03,0.1,0.3)]
70 idea_grid += [{'lr':lr,'p':4,'gamma':0.1} for lr in (0.001,0.003,0.006)]
71 # Evaluate all idea settings on four seeds for selection, then best on all eight.
72 tried=[]
73 for c in idea_grid:
74 r=evaluate(lambda s: idea_fn(c)(s),seeds=SWEEP_SEEDS); tried.append({'cfg':c,'mean':r['mean'],'n':r['n']})
75 finite=[z for z in tried if np.isfinite(z['mean']) and z.get('n',0)==len(SWEEP_SEEDS)]
76 best=min(finite,key=lambda z:z['mean'])['cfg']; ir=evaluate(idea_fn(best),seeds=SEEDS)
77 # Signature is measured from trained models, not analytic identities.
78 pairs=[]
79 for s in SEEDS:
80 z=idea_fn(best)(s,True); pairs.extend(z['sig'][-10:])
81 ratios=[a/(b+1e-12) for a,b in pairs if np.isfinite(a+b)]
82 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)}
83 rep=make_report('tabular',MODEL,base,ir,{'mechanism_signature':sig,'idea_sweep':tried,'selection_cfg':best,'equal_budget':{'epochs':EPOCHS,'batch':BATCH}})
84 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
85 print(json.dumps(rep,indent=2))
86if __name__=='__main__': main()