Order-Adaptive Integral Optimizer / order_adaptive.py
Failed on benchmark
1"""Order-adaptive integral optimizer and deterministic toy verification."""
2import json, math
3import numpy as np
4import torch
5from torch.optim import Optimizer
6
7class OrderAdaptiveIntegral(Optimizer):
8 """Gradient descent plus nested integrated-gradient feedback.
9
10 The order is global (not layerwise): p=0 is ordinary gradient descent.
11 States are allocated lazily, and each newly enabled gain is ramped.
12 """
13 def __init__(self, params, lr=1e-2, max_order=2, beta=.95, rho=.98,
14 decision_interval=10, patience=2, ramp_steps=50,
15 state_clip=100.0, a0=1.0, gains=(.05, .0005)):
16 if lr <= 0 or max_order < 0: raise ValueError("bad lr/max_order")
17 defaults = dict(lr=lr); super().__init__(params, defaults)
18 self.max_order, self.beta, self.rho = max_order, beta, rho
19 self.decision_interval, self.patience = decision_interval, patience
20 self.ramp_steps, self.state_clip = ramp_steps, state_clip
21 self.a0, self.gains = a0, tuple(gains)
22 self.order, self.ema = 0, None
23 self.step_count, self.bad_intervals, self.activations = 0, 0, []
24 if len(self.gains) < max_order: raise ValueError("need one gain per order")
25 for group in self.param_groups:
26 for p in group['params']:
27 self.state[p]['active'] = 0
28 self.state[p]['gamma'] = {}
29 self.state[p]['z'] = {}
30
31 @torch.no_grad()
32 def step(self, closure=None):
33 loss = None
34 if closure is not None:
35 with torch.enable_grad(): loss = closure()
36 grad_sq, n = 0.0, 0
37 for group in self.param_groups:
38 for p in group['params']:
39 if p.grad is not None:
40 grad_sq += float(p.grad.detach().pow(2).sum()); n += p.numel()
41 gnorm = math.sqrt(grad_sq) / max(math.sqrt(n), 1.0)
42 self.ema = gnorm if self.ema is None else self.beta*self.ema+(1-self.beta)*gnorm
43 self.step_count += 1
44 if self.order < self.max_order and self.step_count % self.decision_interval == 0:
45 # Compare consecutive EMA samples at decision boundaries.
46 old = getattr(self, '_last_decision_ema', None)
47 if old is not None and self.ema > self.rho*old: self.bad_intervals += 1
48 else: self.bad_intervals = 0
49 self._last_decision_ema = self.ema
50 if self.bad_intervals >= self.patience:
51 self.order += 1; self.bad_intervals = 0
52 self.activations.append({'step': self.step_count, 'order': self.order,
53 'ema': self.ema})
54 for group in self.param_groups:
55 for p in group['params']:
56 self.state[p]['gamma'][self.order] = 0.0
57 self.state[p]['z'][self.order] = torch.zeros_like(p)
58 for group in self.param_groups:
59 lr = group['lr']
60 for p in group['params']:
61 if p.grad is None: continue
62 st = self.state[p]; g = p.grad
63 update = self.a0*g
64 # Update each chain state before computing the feedback.
65 for r in range(1, self.order+1):
66 z = st['z'][r]
67 if r == 1: z.add_(g, alpha=lr)
68 else: z.add_(st['z'][r-1], alpha=lr)
69 norm = z.norm()
70 if torch.isfinite(norm) and norm.item() > self.state_clip:
71 z.mul_(self.state_clip / norm.item())
72 st['gamma'][r] = min(1.0, st['gamma'].get(r, 0.0) + 1.0/max(1,self.ramp_steps))
73 update = update + st['gamma'][r]*self.gains[r-1]*z
74 p.add_(update, alpha=-lr)
75 return loss
76
77def polynomial_roots(a0, gains, q):
78 # P=lambda^(p+1)+a0*q lambda^p + a1*q lambda^(p-1)...
79 p=len(gains); coeff=[1.0, a0*q] + [x*q for x in gains]
80 return np.roots(coeff)
81
82def quadratic_run(curv, adaptive=True, steps=1200, lr=.01, seed=7):
83 rng=np.random.default_rng(seed); theta=rng.normal(size=len(curv))
84 initial=theta.copy(); ema=None; order=0; bad=0; last=None; acts=[]
85 z=[]; gamma=[]; beta=.95; rho=.98; K=2; interval=10; H=50
86 gains=[.05,.0005]; losses=[]; rms=[]
87 for k in range(steps):
88 g=curv*theta; gn=np.linalg.norm(g)/math.sqrt(len(g))
89 ema=gn if ema is None else beta*ema+(1-beta)*gn
90 if adaptive and (k+1)%interval==0 and order<2:
91 if last is not None and ema > rho*last: bad+=1
92 else: bad=0
93 last=ema
94 if bad>=K:
95 order+=1; bad=0; acts.append((k+1, order, float(ema)))
96 z.append(np.zeros_like(theta)); gamma.append(0.)
97 update=g.copy()
98 for r in range(order):
99 z[r] += lr*(g if r==0 else z[r-1])
100 gamma[r]=min(1.,gamma[r]+1/H)
101 norm=np.linalg.norm(z[r])
102 if norm>100: z[r]*=100/norm
103 update += gamma[r]*gains[r]*z[r]
104 theta -= lr*update
105 losses.append(.5*np.sum(curv*theta*theta)); rms.append(gn)
106 return {'initial_loss':.5*np.sum(curv*initial*initial), 'final_loss':float(losses[-1]),
107 'loss_200':float(losses[199]), 'loss_600':float(losses[599]),
108 'rms_final':float(rms[-1]), 'activations':acts, 'max_loss':float(max(losses)),
109 'stable':bool(np.isfinite(losses).all() and max(losses)<1e12)}
110
111def verify():
112 # Directly test the stated coefficient condition and nested root stability.
113 coeff_ok=(.05**2 > 4*1.0*.0005)
114 roots={q: polynomial_roots(1., [.05,.0005], q).tolist() for q in [1.,100.]}
115 easy_gd=quadratic_run(np.array([1.,2.,4.]), False)
116 easy_ad=quadratic_run(np.array([1.,2.,4.]), True)
117 ill_gd=quadratic_run(np.array([1.,100.]), False)
118 ill_ad=quadratic_run(np.array([1.,100.]), True)
119 return {'coefficient_condition':coeff_ok, 'roots':roots,
120 'root_stable':all(np.max(np.real(x))<0 for v in roots.values() for x in [np.array(v)]),
121 'easy_baseline':easy_gd, 'easy_adaptive':easy_ad,
122 'ill_baseline':ill_gd, 'ill_adaptive':ill_ad}
123
124if __name__ == '__main__':
125 print(json.dumps(verify(), indent=2))