Forced Variational Momentum Optimizer / experiment.py
Mechanism failed
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()