Order-Adaptive Integral Optimizer / order_adaptive.py

Failed on benchmark

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