Bidirectional Saturation-Aware Trust Region / experiment.py
Failed on benchmark
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED = 2993
8
9def seed_all(seed=SEED):
10 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
11 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
12
13class AdaptiveTR:
14 def __init__(self, r0, rmin, rmax, kp=.02, km=.01, ema_decay=.9, eps=1e-12):
15 self.r=float(r0); self.rmin=float(rmin); self.rmax=float(rmax)
16 self.kp=kp; self.km=km; self.decay=ema_decay; self.eps=eps
17 self.q_ema=0.; self.steps=0
18 def apply(self, params, lr):
19 # params have gradients; proposed update is ordinary SGD p=-lr*g.
20 sq=0.;
21 for p in params:
22 if p.grad is not None: sq += float((lr*p.grad).pow(2).sum().item())
23 norm=math.sqrt(sq)
24 q=min(1., norm/(self.r+self.eps))
25 scale=min(1., self.r/(norm+self.eps))
26 with torch.no_grad():
27 for p in params:
28 if p.grad is not None: p.add_(p.grad, alpha=-lr*scale)
29 self.q_ema=self.decay*self.q_ema+(1-self.decay)*q
30 self.r=float(np.clip(self.r*math.exp(self.kp*q-self.km*(1-q)),self.rmin,self.rmax))
31 self.steps += 1
32 return norm,q,scale,self.r
33
34def toy_check():
35 # Formula-level bidirectional check: saturated proposals expand; unsaturated ones contract.
36 r0=.1; tr=AdaptiveTR(r0,.01,10.,kp=.08,km=.04,ema_decay=0.)
37 rs=[]; qs=[]
38 for proposal in [.5]*8+[.01]*12:
39 # emulate a one-dimensional gradient whose proposed norm is proposal
40 p=torch.tensor([proposal], requires_grad=True)
41 p.grad=torch.tensor([-proposal]) # lr=1 gives proposed norm proposal
42 _,q,_,r=tr.apply([p],1.)
43 qs.append(q); rs.append(r)
44 expansion=rs[7]>r0 and all(rs[i+1]>rs[i] for i in range(7))
45 contraction=rs[-1]<rs[8] and all(rs[i+1]<rs[i] for i in range(8,len(rs)-1))
46 # Verify smooth quadratic descent inequality in unsaturated regime.
47 L=3.; eta=.4; x=torch.tensor([1.2,-.7]); g=L*x
48 old=.5*L*float((x*x).sum()); y=x-eta*g
49 new=.5*L*float((y*y).sum()); bound=old-(eta-L*eta*eta/2)*float((g*g).sum())
50 inequality_ok=new <= bound+1e-6
51 return {'radius_start':r0,'radius_after_saturated':rs[7],'radius_end':rs[-1],
52 'q_saturated':qs[0],'q_unsaturated':qs[-1],
53 'expands_when_saturated':expansion,'contracts_when_unsaturated':contraction,
54 'quadratic_descent_inequality_ok':inequality_ok,'quadratic_loss_old':old,'quadratic_loss_new':new,'quadratic_bound':bound}
55
56def make_data(n=2400):
57 g=np.random.default_rng(SEED)
58 x=g.normal(size=(n,2)).astype('float32')
59 y=((x[:,0]*x[:,1]>0).astype('int64'))
60 # modest label noise makes the task nontrivial
61 flip=g.random(n)<.04; y[flip]=1-y[flip]
62 return torch.tensor(x),torch.tensor(y)
63
64class MLP(nn.Module):
65 def __init__(self):
66 super().__init__(); self.net=nn.Sequential(nn.Linear(2,32),nn.Tanh(),nn.Linear(32,32),nn.Tanh(),nn.Linear(32,2))
67 def forward(self,x): return self.net(x)
68
69def run(mode, x, y, steps=500, batch=64):
70 seed_all(SEED+17) # identical initialization for both methods
71 model=MLP(); lossfn=nn.CrossEntropyLoss()
72 # r0 intentionally causes early saturation; fixed control uses same safety cap.
73 lr=.35; r0=.055
74 tr=AdaptiveTR(r0,.005,.55,kp=.04,km=.02,ema_decay=.9) if mode=='adaptive' else None
75 losses=[]; accs=[]; qs=[]; radii=[]; scales=[]
76 gen=torch.Generator().manual_seed(SEED+99)
77 for t in range(steps):
78 ix=torch.randint(0,len(x),(batch,),generator=gen)
79 model.zero_grad(set_to_none=True); out=model(x[ix]); loss=lossfn(out,y[ix]); loss.backward()
80 if mode=='adaptive': norm,q,scale,r=tr.apply(list(model.parameters()),lr); radii.append(r)
81 else:
82 sq=sum(float((lr*p.grad).pow(2).sum()) for p in model.parameters() if p.grad is not None)
83 norm=math.sqrt(sq); q=min(1.,norm/(r0+1e-12)); scale=min(1.,r0/(norm+1e-12))
84 with torch.no_grad():
85 for p in model.parameters():
86 if p.grad is not None: p.add_(p.grad,alpha=-lr*scale)
87 losses.append(float(loss)); qs.append(q); scales.append(scale)
88 if (t+1)%50==0:
89 with torch.no_grad(): accs.append(float((model(x).argmax(1)==y).float().mean()))
90 with torch.no_grad(): final_loss=float(lossfn(model(x),y)); final_acc=float((model(x).argmax(1)==y).float().mean())
91 return {'final_loss':final_loss,'final_accuracy':final_acc,'mean_q':float(np.mean(qs)),
92 'clip_fraction':float(np.mean(np.array(qs)>=1-1e-9)),'mean_scale':float(np.mean(scales)),
93 'initial_q':float(qs[0]),'last100_q':float(np.mean(qs[-100:])),
94 'radius_initial':r0,'radius_final':float(radii[-1]) if radii else r0,
95 'loss_at_100':losses[99],'loss_at_500':losses[-1]}
96
97def main():
98 seed_all(); toy=toy_check(); x,y=make_data()
99 results={'toy_check':toy,'adaptive':run('adaptive',x,y),'fixed_clip':run('fixed',x,y)}
100 Path('results.json').write_text(json.dumps(results,indent=2))
101 print(json.dumps(results,indent=2))
102if __name__=='__main__': main()