Inexact High-Order Moreau DC Optimizer / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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()