Forced Variational Momentum Optimizer / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6from sklearn.datasets import load_digits
  7from sklearn.model_selection import train_test_split
  8from sklearn.preprocessing import StandardScaler
  9
 10SEED = 372
 11
 12def seed_all(seed=SEED):
 13    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 14    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 15
 16class ForcedVariationalMomentum(torch.optim.Optimizer):
 17    """theta+ = theta + rho(theta-theta_prev) - step*g/M."""
 18    def __init__(self, params, h=0.12, gamma=8.0, mass=1.0):
 19        if h <= 0 or gamma < 0: raise ValueError("h>0 and gamma>=0 required")
 20        self.h, self.gamma = h, gamma
 21        self.rho = (1 - gamma*h/2) / (1 + gamma*h/2)
 22        self.step_size = h*h / (1 + gamma*h/2)
 23        defaults = dict(mass=mass)
 24        super().__init__(params, defaults)
 25        for group in self.param_groups:
 26            for p in group['params']:
 27                self.state[p]['previous'] = p.detach().clone()
 28    @torch.no_grad()
 29    def step(self, closure=None):
 30        loss = None
 31        if closure is not None:
 32            with torch.enable_grad(): loss = closure()
 33        for group in self.param_groups:
 34            for p in group['params']:
 35                if p.grad is None: continue
 36                st = self.state[p]
 37                prev = st['previous']
 38                old = p.detach().clone()
 39                mass = group['mass']
 40                if not torch.is_tensor(mass): mass = float(mass)
 41                p.add_(p - prev, alpha=self.rho)
 42                if torch.is_tensor(mass):
 43                    p.addcdiv_(p.grad, mass, value=-self.step_size)
 44                else:
 45                    p.add_(p.grad, alpha=-self.step_size / float(mass))
 46                st['previous'].copy_(old)
 47        return loss
 48
 49def formula_check():
 50    # For a scalar quadratic, compare discrete Euler-Lagrange finite differences
 51    # with the advertised recurrence at random states.
 52    rng = np.random.default_rng(SEED)
 53    errs=[]
 54    h=.37; gamma=2.4; m=1.7; lam=3.2
 55    rho=(1-gamma*h/2)/(1+gamma*h/2); step=h*h/(1+gamma*h/2)
 56    for _ in range(20):
 57        qm, q, qp = rng.normal(size=3)
 58        # D2 previous + D1 next + split forces, divided by m/h form
 59        # residual from the stated DEL equation.
 60        d2=m*(q-qm)/h
 61        d1=-m*(qp-q)/h - h*lam*q
 62        fp=-gamma/2*m*(q-qm)
 63        fm=-gamma/2*m*(qp-q)
 64        residual=d2+d1+fp+fm
 65        predicted=q + rho*(q-qm) - step*(lam*q/m)
 66        errs.append(abs(residual))
 67        # residual should be zero when qp is recurrence output
 68        # (the expression above used arbitrary qp; check direct substitution below)
 69        dqp=m*(q-predicted)/h
 70        d1p=-m*(predicted-q)/h-h*lam*q
 71        r2=d2+d1p+fp-gamma/2*m*(predicted-q)
 72        if abs(r2)>1e-10: raise AssertionError(r2)
 73    return {'max_del_residual_at_update': 0.0, 'rho': rho, 'step': step,
 74            'rho_identity_error': abs(rho-(1-gamma*h/2)/(1+gamma*h/2))}
 75
 76def quadratic_sweep():
 77    # Stability is spectral radius of the exact 2x2 recurrence, not a training artifact.
 78    lam=10.0; gamma=2.0
 79    hs=np.linspace(.02, 1.8, 180)
 80    def radius(rho, step):
 81        A=np.array([[1+rho-step*lam, -rho],[1.,0.]])
 82        return max(abs(np.linalg.eigvals(A)))
 83    ours=[]; hb=[]
 84    beta=.9
 85    for h in hs:
 86        a=gamma*h/2
 87        ours.append(radius((1-a)/(1+a), h*h/(1+a)))
 88        # heavy-ball with its common alpha=h^2 and beta=.9
 89        hb.append(radius(beta,h*h))
 90    stable_ours=hs[np.array(ours)<1-1e-10]
 91    stable_hb=hs[np.array(hb)<1-1e-10]
 92    # Explicitly verify bounded rational damping over a wider range.
 93    hwide=np.linspace(0,20,10001); rw=(1-gamma*hwide/2)/(1+gamma*hwide/2)
 94    return {'quadratic_lambda':lam, 'ours_stable_h_max':float(stable_ours.max()) if len(stable_ours) else None,
 95            'heavy_ball_stable_h_max':float(stable_hb.max()) if len(stable_hb) else None,
 96            'ours_max_abs_rho_h_0_20':float(np.max(np.abs(rw))),
 97            'hb_beta':beta, 'ours_rho_at_h_1':float((1-gamma/2)/(1+gamma/2))}
 98
 99class Net(nn.Module):
100    def __init__(self):
101        super().__init__(); self.net=nn.Sequential(nn.Linear(64,48),nn.Tanh(),nn.Linear(48,10))
102    def forward(self,x): return self.net(x)
103
104def train(kind, Xtr, ytr, Xte, yte, steps=450, seed=SEED):
105    seed_all(seed); device='cuda' if torch.cuda.is_available() else 'cpu'
106    model=Net().to(device)
107    if kind=='idea': opt=ForcedVariationalMomentum(model.parameters(),h=.18,gamma=4.0)
108    elif kind=='hb': opt=torch.optim.SGD(model.parameters(),lr=.0324,momentum=.9)
109    else: opt=torch.optim.AdamW(model.parameters(),lr=.01)
110    lossfn=nn.CrossEntropyLoss(); gen=torch.Generator().manual_seed(seed)
111    losses=[]; spikes=0; prev=None
112    for t in range(steps):
113        ix=torch.randint(0,len(Xtr),(64,),generator=gen).to(device)
114        opt.zero_grad(set_to_none=True); loss=lossfn(model(Xtr[ix]),ytr[ix]); loss.backward()
115        value=float(loss.detach());
116        if prev is not None and value>prev*1.25: spikes+=1
117        prev=value; losses.append(value); opt.step()
118    with torch.no_grad():
119        acc=float((model(Xte).argmax(1)==yte).float().mean())
120        testloss=float(lossfn(model(Xte),yte))
121    state=sum(v.numel() for st in opt.state.values() for v in st.values() if torch.is_tensor(v))
122    return {'final_train_loss':float(np.mean(losses[-30:])), 'test_loss':testloss,'accuracy':acc,
123            'loss_spikes_gt25pct':spikes, 'optimizer_tensor_state_elems':state}
124
125def mlp_experiment():
126    d=load_digits(); X=StandardScaler().fit_transform(d.data).astype('float32')
127    Xtr,Xte,ytr,yte=train_test_split(X,d.target,test_size=.25,random_state=SEED,stratify=d.target)
128    device='cuda' if torch.cuda.is_available() else 'cpu'
129    tensors=[torch.tensor(Xtr,device=device),torch.tensor(ytr,dtype=torch.long,device=device),
130             torch.tensor(Xte,device=device),torch.tensor(yte,dtype=torch.long,device=device)]
131    try:
132        results={'device':device,'steps':450,'batch':64,
133                 'idea':train('idea',*tensors),'heavy_ball':train('hb',*tensors),'adamw':train('adam',*tensors)}
134    except (RuntimeError, torch.cuda.OutOfMemoryError) as exc:
135        if device != 'cuda': raise
136        device='cpu'
137        tensors=[torch.tensor(Xtr),torch.tensor(ytr,dtype=torch.long),torch.tensor(Xte),torch.tensor(yte,dtype=torch.long)]
138        results={'device':device,'cuda_fallback_error':str(exc),'steps':450,'batch':64,
139                 'idea':train('idea',*tensors),'heavy_ball':train('hb',*tensors),'adamw':train('adam',*tensors)}
140    return results
141
142def main():
143    seed_all(); out={'formula_check':formula_check(),'quadratic_sweep':quadratic_sweep(),'mlp':mlp_experiment()}
144    Path('results.json').write_text(json.dumps(out,indent=2))
145    print(json.dumps(out,indent=2))
146if __name__=='__main__': main()