Universal Trust-Region Neural Optimizer / trust_region_experiment.py

Failed on benchmark

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