import json, math, random, time from pathlib import Path import numpy as np import torch SEED=17 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) def tr_step_quadratic(x, lam, b, delta, eta1=.1, eta2=.75): g=lam*x # exact solution for scalar quadratic model with curvature b s=-g/b if b>0 else (-np.sign(g)*delta) s=float(np.clip(s,-delta,delta)) pred=-(g*s+.5*b*s*s) ared=.5*lam*x*x-.5*lam*(x+s)*(x+s) rho=ared/pred if pred>1e-15 else -np.inf accepted=rho>=eta1 newdelta=delta if not accepted: newdelta*=.25 elif rho>.75 and abs(s)>=.99*delta: newdelta=min(2*delta,1e6) return (x+s if accepted else x),newdelta,rho,abs(s)>=.99*delta def toy_checks(): # Prediction A: exact model gives rho=1, independent of curvature and radius. exact=[] for lam in [0.1,1,10,100]: for d in [.01,.2,10]: _,_,r,_=tr_step_quadratic(2.,lam,lam,d); exact.append(r) # Prediction B: for an unconstrained mismatched model, rho=2-lambda/B. observed=[]; predicted=[] for lam in [.5,1,2,8]: for ratio in [.5,1,2,4]: b=ratio*lam _,_,r,_=tr_step_quadratic(1.,lam,b,100) observed.append(r); predicted.append(2-1/ratio) # Prediction C: the near-boundary flag (|s| >= .99 Delta) transitions at # Delta <= Delta*/.99, where Delta*=|g|/B=lambda/B for x=1. transition=[] for lam,b in [(1.,.5),(1.,2.),(3.,1.)]: dstar=lam/b for mult in [.5,.99,1.01,2.]: _,_,r,bound=tr_step_quadratic(1.,lam,b,dstar*mult) transition.append({'lambda':lam,'B':b,'delta_over_delta_star':mult, 'predicted_boundary':mult<=1/.99, 'observed_near_boundary':bool(bound),'rho':float(r)}) # Prediction D: poor unconstrained agreement is rejected and Delta contracts by 4x. _,d_good,r_good,bound_good=tr_step_quadratic(1.,1,1,.1) _,d_bad,r_bad,bound_bad=tr_step_quadratic(1.,1,.01,100.) return { 'exact_model_rho_minmax':[float(min(exact)),float(max(exact))], 'mismatch_rho_max_abs_error':float(max(abs(np.array(observed)-np.array(predicted)))), 'mismatch_rho_pairs':[[float(p),float(o)] for p,o in zip(predicted,observed)], 'boundary_transition':transition, 'good_boundary':{'rho':float(r_good),'radius_ratio':float(d_good/.1),'boundary':bool(bound_good)}, 'poor_agreement':{'rho':float(r_bad),'radius_ratio':float(d_bad/100.),'boundary':bool(bound_bad)}, 'acceptance_threshold_prediction':'rho >= 0.1; for unconstrained mismatch this means B/lambda >= 1/1.9 = 0.5263' } class TinyMLP(torch.nn.Module): def __init__(self): super().__init__(); self.net=torch.nn.Sequential(torch.nn.Linear(2,16),torch.nn.Tanh(),torch.nn.Linear(16,2)) def forward(self,x): return self.net(x) def make_data(): rng=np.random.RandomState(SEED) x=rng.randn(96,2).astype('float32') y=((x[:,0]*x[:,1]>0).astype('int64')) return torch.tensor(x),torch.tensor(y) def train_adam(x,y,steps=120): torch.manual_seed(SEED); m=TinyMLP(); opt=torch.optim.Adam(m.parameters(),lr=.03) lossfn=torch.nn.CrossEntropyLoss(); losses=[]; spikes=0 for _ in range(steps): opt.zero_grad(); loss=lossfn(m(x),y); loss.backward(); opt.step(); v=float(loss); losses.append(v) if len(losses)>1 and v>2*losses[-2]: spikes+=1 with torch.no_grad(): acc=float((m(x).argmax(1)==y).float().mean()) return {'final_loss':losses[-1],'best_loss':min(losses),'accuracy':acc,'spikes':spikes,'losses':losses} def train_tr(x,y,steps=120,delta0=.3): torch.manual_seed(SEED); m=TinyMLP(); lossfn=torch.nn.CrossEntropyLoss(); delta=delta0; losses=[]; rejects=0; radii=[]; rhos=[] for _ in range(steps): # B=I is a damped diagonal curvature model; Cauchy step is exact for this model. m.zero_grad(); old=float(lossfn(m(x),y)); lossfn(m(x),y).backward() params=list(m.parameters()); flatg=torch.cat([p.grad.reshape(-1) for p in params]); gn=float(flatg.norm()) normstep=min(delta,gn); step=(-normstep/(gn+1e-12))*flatg pred=gn*normstep-.5*normstep*normstep oldvals=[p.detach().clone() for p in params] pos=0 with torch.no_grad(): for p in params: n=p.numel(); p.add_(step[pos:pos+n].view_as(p)); pos+=n new=float(lossfn(m(x),y)); ared=old-new; rho=ared/pred if pred>1e-12 else -1 accepted=rho>=.1 if not accepted: with torch.no_grad(): for p,v in zip(params,oldvals): p.copy_(v) rejects+=1; delta*=.25 else: if rho>.75 and normstep>=.99*delta: delta=min(2*delta,10.) losses.append(old if not accepted else new); radii.append(delta); rhos.append(rho) with torch.no_grad(): acc=float((m(x).argmax(1)==y).float().mean()) 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} def main(): checks=toy_checks(); x,y=make_data(); a=train_adam(x,y); t=train_tr(x,y) 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'}}} Path('results.json').write_text(json.dumps(out,indent=2)) print(json.dumps(out,indent=2)) if __name__=='__main__': main()