"""Order-adaptive integral optimizer and deterministic toy verification.""" import json, math import numpy as np import torch from torch.optim import Optimizer class OrderAdaptiveIntegral(Optimizer): """Gradient descent plus nested integrated-gradient feedback. The order is global (not layerwise): p=0 is ordinary gradient descent. States are allocated lazily, and each newly enabled gain is ramped. """ def __init__(self, params, lr=1e-2, max_order=2, beta=.95, rho=.98, decision_interval=10, patience=2, ramp_steps=50, state_clip=100.0, a0=1.0, gains=(.05, .0005)): if lr <= 0 or max_order < 0: raise ValueError("bad lr/max_order") defaults = dict(lr=lr); super().__init__(params, defaults) self.max_order, self.beta, self.rho = max_order, beta, rho self.decision_interval, self.patience = decision_interval, patience self.ramp_steps, self.state_clip = ramp_steps, state_clip self.a0, self.gains = a0, tuple(gains) self.order, self.ema = 0, None self.step_count, self.bad_intervals, self.activations = 0, 0, [] if len(self.gains) < max_order: raise ValueError("need one gain per order") for group in self.param_groups: for p in group['params']: self.state[p]['active'] = 0 self.state[p]['gamma'] = {} self.state[p]['z'] = {} @torch.no_grad() def step(self, closure=None): loss = None if closure is not None: with torch.enable_grad(): loss = closure() grad_sq, n = 0.0, 0 for group in self.param_groups: for p in group['params']: if p.grad is not None: grad_sq += float(p.grad.detach().pow(2).sum()); n += p.numel() gnorm = math.sqrt(grad_sq) / max(math.sqrt(n), 1.0) self.ema = gnorm if self.ema is None else self.beta*self.ema+(1-self.beta)*gnorm self.step_count += 1 if self.order < self.max_order and self.step_count % self.decision_interval == 0: # Compare consecutive EMA samples at decision boundaries. old = getattr(self, '_last_decision_ema', None) if old is not None and self.ema > self.rho*old: self.bad_intervals += 1 else: self.bad_intervals = 0 self._last_decision_ema = self.ema if self.bad_intervals >= self.patience: self.order += 1; self.bad_intervals = 0 self.activations.append({'step': self.step_count, 'order': self.order, 'ema': self.ema}) for group in self.param_groups: for p in group['params']: self.state[p]['gamma'][self.order] = 0.0 self.state[p]['z'][self.order] = torch.zeros_like(p) for group in self.param_groups: lr = group['lr'] for p in group['params']: if p.grad is None: continue st = self.state[p]; g = p.grad update = self.a0*g # Update each chain state before computing the feedback. for r in range(1, self.order+1): z = st['z'][r] if r == 1: z.add_(g, alpha=lr) else: z.add_(st['z'][r-1], alpha=lr) norm = z.norm() if torch.isfinite(norm) and norm.item() > self.state_clip: z.mul_(self.state_clip / norm.item()) st['gamma'][r] = min(1.0, st['gamma'].get(r, 0.0) + 1.0/max(1,self.ramp_steps)) update = update + st['gamma'][r]*self.gains[r-1]*z p.add_(update, alpha=-lr) return loss def polynomial_roots(a0, gains, q): # P=lambda^(p+1)+a0*q lambda^p + a1*q lambda^(p-1)... p=len(gains); coeff=[1.0, a0*q] + [x*q for x in gains] return np.roots(coeff) def quadratic_run(curv, adaptive=True, steps=1200, lr=.01, seed=7): rng=np.random.default_rng(seed); theta=rng.normal(size=len(curv)) initial=theta.copy(); ema=None; order=0; bad=0; last=None; acts=[] z=[]; gamma=[]; beta=.95; rho=.98; K=2; interval=10; H=50 gains=[.05,.0005]; losses=[]; rms=[] for k in range(steps): g=curv*theta; gn=np.linalg.norm(g)/math.sqrt(len(g)) ema=gn if ema is None else beta*ema+(1-beta)*gn if adaptive and (k+1)%interval==0 and order<2: if last is not None and ema > rho*last: bad+=1 else: bad=0 last=ema if bad>=K: order+=1; bad=0; acts.append((k+1, order, float(ema))) z.append(np.zeros_like(theta)); gamma.append(0.) update=g.copy() for r in range(order): z[r] += lr*(g if r==0 else z[r-1]) gamma[r]=min(1.,gamma[r]+1/H) norm=np.linalg.norm(z[r]) if norm>100: z[r]*=100/norm update += gamma[r]*gains[r]*z[r] theta -= lr*update losses.append(.5*np.sum(curv*theta*theta)); rms.append(gn) return {'initial_loss':.5*np.sum(curv*initial*initial), 'final_loss':float(losses[-1]), 'loss_200':float(losses[199]), 'loss_600':float(losses[599]), 'rms_final':float(rms[-1]), 'activations':acts, 'max_loss':float(max(losses)), 'stable':bool(np.isfinite(losses).all() and max(losses)<1e12)} def verify(): # Directly test the stated coefficient condition and nested root stability. coeff_ok=(.05**2 > 4*1.0*.0005) roots={q: polynomial_roots(1., [.05,.0005], q).tolist() for q in [1.,100.]} easy_gd=quadratic_run(np.array([1.,2.,4.]), False) easy_ad=quadratic_run(np.array([1.,2.,4.]), True) ill_gd=quadratic_run(np.array([1.,100.]), False) ill_ad=quadratic_run(np.array([1.,100.]), True) return {'coefficient_condition':coeff_ok, 'roots':roots, 'root_stable':all(np.max(np.real(x))<0 for v in roots.values() for x in [np.array(v)]), 'easy_baseline':easy_gd, 'easy_adaptive':easy_ad, 'ill_baseline':ill_gd, 'ill_adaptive':ill_ad} if __name__ == '__main__': print(json.dumps(verify(), indent=2))