"""Third-order Langevin optimizer MVP. State: parameter x, velocity v, acceleration a. Noise is injected only into a: a <- a + dt*(-grad - gamma*a) + sqrt(2*gamma*temperature*dt)*N(0,I) v <- v + dt*a x <- x + dt*v """ import math import torch from torch.optim import Optimizer def cubic_unstable_rate(gamma: float, kappa: float) -> float: """Positive root r of r^3 + gamma*r^2 - kappa = 0.""" if gamma <= 0: raise ValueError("gamma must be positive") if kappa <= 0: return 0.0 lo, hi = 0.0, kappa ** (1.0 / 3.0) + math.sqrt(kappa / gamma) + 1.0 for _ in range(80): mid = (lo + hi) * 0.5 if mid**3 + gamma * mid**2 < kappa: lo = mid else: hi = mid return (lo + hi) * 0.5 class ThirdOrderLangevin(Optimizer): """Noise-on-acceleration third-order Langevin optimizer. ``adapt_dt=True`` optionally enforces dt*r <= rate_limit using a supplied negative-curvature estimate via ``set_negative_curvature``. Curvature estimation is deliberately kept outside the optimizer so callers can use Hessian-vector products at a chosen cadence. """ def __init__(self, params, dt=1e-2, gamma=1.0, temperature=0.0, adapt_dt=False, rate_limit=0.2): if dt <= 0 or gamma <= 0 or temperature < 0 or rate_limit <= 0: raise ValueError("invalid dt, gamma, temperature, or rate_limit") defaults = dict(dt=float(dt), gamma=float(gamma), temperature=float(temperature), adapt_dt=adapt_dt, rate_limit=float(rate_limit)) super().__init__(params, defaults) self._kappa = 0.0 @torch.no_grad() def set_negative_curvature(self, kappa: float): self._kappa = max(0.0, float(kappa)) @torch.no_grad() def step(self, closure=None): loss = None if closure is not None: with torch.enable_grad(): loss = closure() # CUDA errors are allowed to propagate to the caller, which can rerun # the experiment on CPU as required by the experiment harness. for group in self.param_groups: dt = group['dt'] gamma = group['gamma'] if group['adapt_dt'] and self._kappa > 0: rate = cubic_unstable_rate(gamma, self._kappa) dt = min(dt, group['rate_limit'] / rate) noise_scale = math.sqrt(2.0 * gamma * group['temperature'] * dt) for p in group['params']: if p.grad is None: continue if p.grad.is_sparse: raise RuntimeError("ThirdOrderLangevin does not support sparse gradients") state = self.state[p] if not state: state['v'] = torch.zeros_like(p, memory_format=torch.preserve_format) state['a'] = torch.zeros_like(p, memory_format=torch.preserve_format) # Keep optimizer state fp32 where parameters are lower precision. v, a = state['v'], state['a'] g = p.grad if not torch.is_floating_point(g): g = g.float() a.add_(-dt * g).add_(-dt * gamma * a) if noise_scale: a.add_(torch.randn_like(a) * noise_scale) v.add_(a, alpha=dt) p.add_(v, alpha=dt) return loss