Universal Trust-Region Neural Optimizer / trust_region_experiment.py
Failed on benchmark
1import json, math, random, time
2from pathlib import Path
3import numpy as np
4import torch
5
6SEED=17
7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
8
9def tr_step_quadratic(x, lam, b, delta, eta1=.1, eta2=.75):
10 g=lam*x
11 # exact solution for scalar quadratic model with curvature b
12 s=-g/b if b>0 else (-np.sign(g)*delta)
13 s=float(np.clip(s,-delta,delta))
14 pred=-(g*s+.5*b*s*s)
15 ared=.5*lam*x*x-.5*lam*(x+s)*(x+s)
16 rho=ared/pred if pred>1e-15 else -np.inf
17 accepted=rho>=eta1
18 newdelta=delta
19 if not accepted: newdelta*=.25
20 elif rho>.75 and abs(s)>=.99*delta: newdelta=min(2*delta,1e6)
21 return (x+s if accepted else x),newdelta,rho,abs(s)>=.99*delta
22
23def toy_checks():
24 # Prediction A: exact model gives rho=1, independent of curvature and radius.
25 exact=[]
26 for lam in [0.1,1,10,100]:
27 for d in [.01,.2,10]:
28 _,_,r,_=tr_step_quadratic(2.,lam,lam,d); exact.append(r)
29 # Prediction B: for an unconstrained mismatched model, rho=2-lambda/B.
30 observed=[]; predicted=[]
31 for lam in [.5,1,2,8]:
32 for ratio in [.5,1,2,4]:
33 b=ratio*lam
34 _,_,r,_=tr_step_quadratic(1.,lam,b,100)
35 observed.append(r); predicted.append(2-1/ratio)
36 # Prediction C: the near-boundary flag (|s| >= .99 Delta) transitions at
37 # Delta <= Delta*/.99, where Delta*=|g|/B=lambda/B for x=1.
38 transition=[]
39 for lam,b in [(1.,.5),(1.,2.),(3.,1.)]:
40 dstar=lam/b
41 for mult in [.5,.99,1.01,2.]:
42 _,_,r,bound=tr_step_quadratic(1.,lam,b,dstar*mult)
43 transition.append({'lambda':lam,'B':b,'delta_over_delta_star':mult,
44 'predicted_boundary':mult<=1/.99,
45 'observed_near_boundary':bool(bound),'rho':float(r)})
46 # Prediction D: poor unconstrained agreement is rejected and Delta contracts by 4x.
47 _,d_good,r_good,bound_good=tr_step_quadratic(1.,1,1,.1)
48 _,d_bad,r_bad,bound_bad=tr_step_quadratic(1.,1,.01,100.)
49 return {
50 'exact_model_rho_minmax':[float(min(exact)),float(max(exact))],
51 'mismatch_rho_max_abs_error':float(max(abs(np.array(observed)-np.array(predicted)))),
52 'mismatch_rho_pairs':[[float(p),float(o)] for p,o in zip(predicted,observed)],
53 'boundary_transition':transition,
54 'good_boundary':{'rho':float(r_good),'radius_ratio':float(d_good/.1),'boundary':bool(bound_good)},
55 'poor_agreement':{'rho':float(r_bad),'radius_ratio':float(d_bad/100.),'boundary':bool(bound_bad)},
56 'acceptance_threshold_prediction':'rho >= 0.1; for unconstrained mismatch this means B/lambda >= 1/1.9 = 0.5263'
57 }
58
59class TinyMLP(torch.nn.Module):
60 def __init__(self):
61 super().__init__(); self.net=torch.nn.Sequential(torch.nn.Linear(2,16),torch.nn.Tanh(),torch.nn.Linear(16,2))
62 def forward(self,x): return self.net(x)
63
64def make_data():
65 rng=np.random.RandomState(SEED)
66 x=rng.randn(96,2).astype('float32')
67 y=((x[:,0]*x[:,1]>0).astype('int64'))
68 return torch.tensor(x),torch.tensor(y)
69
70def train_adam(x,y,steps=120):
71 torch.manual_seed(SEED); m=TinyMLP(); opt=torch.optim.Adam(m.parameters(),lr=.03)
72 lossfn=torch.nn.CrossEntropyLoss(); losses=[]; spikes=0
73 for _ in range(steps):
74 opt.zero_grad(); loss=lossfn(m(x),y); loss.backward(); opt.step(); v=float(loss); losses.append(v)
75 if len(losses)>1 and v>2*losses[-2]: spikes+=1
76 with torch.no_grad(): acc=float((m(x).argmax(1)==y).float().mean())
77 return {'final_loss':losses[-1],'best_loss':min(losses),'accuracy':acc,'spikes':spikes,'losses':losses}
78
79def train_tr(x,y,steps=120,delta0=.3):
80 torch.manual_seed(SEED); m=TinyMLP(); lossfn=torch.nn.CrossEntropyLoss(); delta=delta0; losses=[]; rejects=0; radii=[]; rhos=[]
81 for _ in range(steps):
82 # B=I is a damped diagonal curvature model; Cauchy step is exact for this model.
83 m.zero_grad(); old=float(lossfn(m(x),y)); lossfn(m(x),y).backward()
84 params=list(m.parameters()); flatg=torch.cat([p.grad.reshape(-1) for p in params]); gn=float(flatg.norm())
85 normstep=min(delta,gn); step=(-normstep/(gn+1e-12))*flatg
86 pred=gn*normstep-.5*normstep*normstep
87 oldvals=[p.detach().clone() for p in params]
88 pos=0
89 with torch.no_grad():
90 for p in params: n=p.numel(); p.add_(step[pos:pos+n].view_as(p)); pos+=n
91 new=float(lossfn(m(x),y)); ared=old-new; rho=ared/pred if pred>1e-12 else -1
92 accepted=rho>=.1
93 if not accepted:
94 with torch.no_grad():
95 for p,v in zip(params,oldvals): p.copy_(v)
96 rejects+=1; delta*=.25
97 else:
98 if rho>.75 and normstep>=.99*delta: delta=min(2*delta,10.)
99 losses.append(old if not accepted else new); radii.append(delta); rhos.append(rho)
100 with torch.no_grad(): acc=float((m(x).argmax(1)==y).float().mean())
101 return {'final_loss':losses[-1],'best_loss':min(losses),'accuracy':acc,'rejections':rejects,'reject_fraction':rejects/steps,'final_radius':radii[-1],'median_rho':float(np.median(rhos)),'losses':losses}
102
103def main():
104 checks=toy_checks(); x,y=make_data(); a=train_adam(x,y); t=train_tr(x,y)
105 out={'seed':SEED,'toy_checks':checks,'mlp':{'adam':{k:v for k,v in a.items() if k!='losses'},'trust_region':{k:v for k,v in t.items() if k!='losses'}}}
106 Path('results.json').write_text(json.dumps(out,indent=2))
107 print(json.dumps(out,indent=2))
108if __name__=='__main__': main()