Bidirectional Saturation-Aware Trust Region / experiment.py

Failed on benchmark

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