From a1ac6e8b01336b6cf12999a70a34572ab941c92f Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Mon, 8 Jun 2026 22:03:54 -0600 Subject: [PATCH] Reworked automagic v3 again. Seems more stable. Still testing. --- toolkit/optimizers/automagic3.py | 214 +++++++++++++++---------------- 1 file changed, 107 insertions(+), 107 deletions(-) diff --git a/toolkit/optimizers/automagic3.py b/toolkit/optimizers/automagic3.py index 82bde608..b99f2eb2 100644 --- a/toolkit/optimizers/automagic3.py +++ b/toolkit/optimizers/automagic3.py @@ -1,7 +1,7 @@ """ NOTE: This is experimental and under active development; expect breaking changes and bugs. Feedback welcome. """ -import math + from typing import List import torch @@ -13,29 +13,33 @@ class Automagic3(torch.optim.Optimizer): A learning rate is kept per row of each parameter: one lr per output channel for >=2D weights (e.g. one lr per output neuron of a Linear layer) and one lr per element for 1D weights (biases, norms). Each step the lr is - nudged by how consistent the per-element update direction has been. + nudged by whether the per-element update direction *flipped* vs the previous + step (RProp-style edge-of-stability control). - Agreement is window-consistency: over the last ``polarity_history_count`` - update-sign snapshots plus the current sign, the per-element score is the - fraction sharing the dominant sign (``max(p, 1-p)``, in [0.5, 1]; direction- - agnostic, and the current step is just one vote so a single flip against a - long consistent history barely moves it). It is reduced to one value per row. - ``polarity_history_count`` is min 1 (default 2); a longer window makes the lr - react to a sustained trend rather than a single noisy step. + A sign flip means the step jumped past the local minimum (overshoot) -- the + one event whose frequency genuinely rises with the lr, so it provides a true + restoring force. Each element votes: agree with last step -> nudge its row's + lr up by ``lr_bump_rate``; flip -> nudge down by the same amount (symmetric). + The per-element log-nudges are averaged to one value per row and EMA-smoothed + over ~``lr_smoothing_steps`` steps (so the lr reacts to a sustained trend, + not a single noisy step), then applied multiplicatively: ``lr *= exp(nudge)``. - The agreement maps to a direction in [-1, 1] measured relative to the - window's noise floor ``b`` -- the consistency a pure-noise row shows purely - by chance, computed automatically from the window size (e.g. 0.75 for a - window of <=3 snapshots, ~0.69 for 4). Agreement == ``b`` -> 0 (lr steady), - 1.0 -> +1 (lr up), below ``b`` -> down. There is no manual target: holding at - the noise floor is self-balancing per row -- the lr grows while a row is more - consistent than chance and settles at whatever lr makes its consistency meet - ``b``, so noisy and clean layers each find their own operating point without - tuning, and it neither collapses to min_lr nor runs to max_lr. Measuring - relative to ``b`` also makes the signal window-size independent. The direction - scales a multiplicative (geometric) bump ``lr *= exp(direction * - lr_bump_rate)`` -- a fixed fractional move, uniform across the whole range. - lr is clamped to [min_lr, max_lr]. + This is self-balancing with no target and no noise floor, and the symmetry is + load-bearing: the equilibrium is flip fraction == 0.5, which is the only flip + rate that is simultaneously the pure-noise point and the edge of stability. + A row still descending cleanly flips less than half the time -> its lr grows; + once the lr is large enough to overshoot it flips more than half -> its lr + shrinks; and a row whose gradients are pure noise (a fresh LoRA's first + steps, or a converged layer) flips ~half the time -> its lr HOLDS. Any + up/down asymmetry moves the equilibrium off 0.5 and a noise-dominated row + then marches straight to min_lr, so the votes are kept symmetric. Elements + whose update is exactly zero (dead/masked grads, low-precision underflow) + carry no direction and abstain from the vote, so a pool of frozen elements + can't quietly bias a row's lr upward. Noisy and clean layers each find their + own operating point automatically; the lr + neither collapses to min_lr nor runs away to max_lr. ``lr_bump_rate`` only + sets how fast it gets there, not where it lands. lr is clamped to + [min_lr, max_lr]. With ``fused=True`` (default) the step is fused into the backward pass via ``register_post_accumulate_grad_hook``: each parameter is updated and its @@ -62,13 +66,17 @@ class Automagic3(torch.optim.Optimizer): to share one rate, so a layer where some rows have converged and others have not is handled gracefully. - 2. Proportional lr control (was a hard threshold flip). v2 bumped the lr up - or down by a fixed amount depending on whether agreement crossed a - threshold, which jitters when agreement hovers near the boundary. v3 - scales the bump by how far consistency is from its hold point relative to - the noise floor (direction in [-1, 1]). Plain English: the lr nudges gently when the - signal is weak and firmly when it is strong, and parks itself instead of - oscillating when gradients are basically noise. + 2. Overshoot-based (RProp-style) lr control with a real equilibrium. v2 + bumped the lr from raw direction agreement, which has no upper fixed point + -- a parameter that is simply still descending keeps agreeing at any lr, + so the lr ratchets up and eventually runs away on long runs. v3 drives the + lr from sign *flips* (overshoot) instead, nudging up on agree and down on + flip symmetrically; the equilibrium is a flip fraction of 0.5, which is + both the noise point and the edge of stability. Plain English: the lr + speeds up while a layer is making clean progress, backs off the moment it + starts overshooting, and simply holds when the gradient is pure noise -- + so it neither climbs without bound on long runs nor collapses to nothing + on a fresh, noisy LoRA. 3. Multiplicative (geometric) lr bump (was additive). v2 added/subtracted a fixed absolute amount, so the same bump was a huge relative jump near @@ -87,34 +95,31 @@ class Automagic3(torch.optim.Optimizer): weight updates, so it actually keeps learning instead of stalling. 5. Faster hot path, identical math. eps is folded into the small reduced - row/col vectors instead of the full gradient-square tensor; the lr scale, - weight decay and parameter update are fused into one ``addcmul_``; and the - sign-agreement is summed straight off the bool mask with no full-size - float cast. Plain English: each step issues fewer GPU passes over the - weights, so it runs faster (notably in bf16/fp16) without changing the - result. + row/col vectors instead of the full gradient-square tensor, and the lr + scale and parameter update are fused into one ``addcmul_``. Plain English: + each step issues fewer GPU passes over the weights, so it runs faster + (notably in bf16/fp16) without changing the result. """ def __init__( self, params, lr: float = 1e-6, - min_lr: float = 1e-8, - max_lr: float = 1e-2, + min_lr: float = 1e-7, + max_lr: float = 1e-3, lr_bump_rate: float = 0.1, # fractional/log step per bump (~10%); see step logic beta2: float = 0.999, eps: float = 1e-30, clip_threshold: float = 1.0, weight_decay: float = 0.0, - polarity_history_count: int = 3, # update-sign snapshots kept; current is compared vs all (min 1) + lr_smoothing_steps: int = 3, # lr-nudge EMA smoothing horizon, in steps (min 1) fused: bool = True, ): if lr > 1e-3: print(f"Warning! Start lr {lr} is very high; forcing to 1e-6.") lr = 1e-6 - # Agreement compares the current update sign against the stored history, - # so at least one snapshot must be kept (1 = compare to previous only). - polarity_history_count = max(1, int(polarity_history_count)) + # The lr nudge is EMA-smoothed over ~this many steps; at least 1. + lr_smoothing_steps = max(1, int(lr_smoothing_steps)) defaults = dict( lr=lr, min_lr=min_lr, @@ -124,10 +129,10 @@ class Automagic3(torch.optim.Optimizer): eps=eps, clip_threshold=clip_threshold, weight_decay=weight_decay, - polarity_history_count=polarity_history_count, - # Noise floor of the consistency measure for this window (history+1), - # subtracted off so the lr signal is window-size independent. - agreement_floor=self._noise_floor(polarity_history_count + 1), + lr_smoothing_steps=lr_smoothing_steps, + # EMA decay for the per-row lr nudge, derived from the smoothing + # horizon (n steps -> beta = n/(n+1)). + dir_beta=lr_smoothing_steps / (lr_smoothing_steps + 1.0), ) super().__init__(params, defaults) @@ -162,16 +167,6 @@ class Automagic3(torch.optim.Optimizer): def _rms(t: torch.Tensor) -> torch.Tensor: return t.norm(2) / (t.numel() ** 0.5) - @staticmethod - def _noise_floor(window: int) -> float: - # Expected window-consistency max(p, 1-p) of a pure-noise element over a - # window of `window` independent fair coin flips. This is > 0.5 and grows - # with smaller windows (0.75 for window<=3, ~0.688 for 4, ...), so it must - # be subtracted off for the lr signal to behave the same at any window. - total = sum(math.comb(window, k) * max(k, window - k) - for k in range(window + 1)) - return total / (window * (2 ** window)) - @staticmethod def _approx_sq_grad(row: torch.Tensor, col: torch.Tensor) -> torch.Tensor: r = (row / row.mean(dim=-1, keepdim=True)).rsqrt_().unsqueeze(-1) @@ -246,9 +241,12 @@ class Automagic3(torch.optim.Optimizer): state["lr"] = torch.full( lr_shape, float(group["lr"]), dtype=torch.float32, device=p.device ) - # Rolling history of the last polarity_history_count update-sign - # snapshots (bool); the current step is compared against all of them. - state["pol_hist"] = [] + # Previous update-sign snapshot (int8 {-1, 0, +1}, full param shape); the + # current sign is compared against it to detect per-element flips. Set on + # the first step. + state["prev_sign"] = None + # EMA of the per-row log lr-nudge, smoothing the flip signal over time. + state["dir_ema"] = torch.zeros(lr_shape, dtype=torch.float32, device=p.device) if p.dim() >= 2: state["exp_avg_sq_row"] = torch.zeros( p.shape[:-1], dtype=p.dtype, device=p.device @@ -336,18 +334,19 @@ class Automagic3(torch.optim.Optimizer): # max-norm trust region) so no single weight can take an outsized step. update.clamp_(-group["clip_threshold"], group["clip_threshold"]) - # Window-consistency agreement. We keep the last polarity_history_count - # update-sign snapshots; over the window of those plus the current sign, - # the per-element agreement is the fraction sharing the *dominant* sign, - # max(p, 1 - p) == 0.5 + |p - 0.5| where p is the fraction positive. This - # is direction-agnostic (a consistently-negative element scores as high - # as a consistently-positive one) and the current step is just one vote, - # so a single flip against a long consistent history barely moves it - # (e.g. 20-of-21 the same -> ~0.95). Values lie in [0.5, 1]: 0.5 is a - # 50/50 (chaotic) element, 1.0 is perfectly consistent. Reduced to one - # value per output channel for >=2D, per element for 1D. - cur_polarity = update > 0 - pol_hist = state["pol_hist"] + # RProp-style edge-of-stability lr control. The signal is whether each + # element's update direction *flipped* vs the previous step, not how + # steady it has been: a flip means we stepped past the local minimum + # (overshoot), the one event whose frequency actually rises with the lr, + # so it gives a true restoring force. Steadiness does not -- a parameter + # descending monotonically agrees with itself at any non-overshooting lr, + # which is why a consistency-vs-noise-floor signal has no upper + # equilibrium and runs away on long tunes. + # Trinary sign {-1, 0, +1}: zero updates (dead/masked grads, flat + # activation regions, low-precision underflow) are kept distinct from + # negatives rather than bucketed with them by a bare ``> 0``. + cur_sign = update.sign().to(torch.int8) + prev_sign = state["prev_sign"] lr_t = state["lr"] if p.dim() >= 2: @@ -357,40 +356,39 @@ class Automagic3(torch.optim.Optimizer): dims = None lr_b = lr_t - if pol_hist: - # per-element fraction positive over the window (current + history) - win = len(pol_hist) + 1 - frac = cur_polarity.to(torch.float32) - for h in pol_hist: - frac.add_(h) - frac.div_(win) - # dominant-sign fraction in [0.5, 1], then reduce to per-row - consist = frac.sub_(0.5).abs_().add_(0.5) - agreement = consist.mean(dim=dims) if dims is not None else consist + if prev_sign is not None: + # Per-element vote via the sign product. With signs in {-1, 0, +1}, + # cur_sign * prev_sign is +1 when the direction held (agree), -1 when + # it flipped (overshoot), and 0 whenever either step's update was zero + # -- so a zero update automatically ABSTAINS (contributes nothing and + # isn't counted), no separate masking needed. One int8 multiply + # replaces the agree/flip/valid masks and their float casts. + # + # Summed per row over the voting (nonzero) elements this is + # bump*(1 - 2*flip_fraction): the lr grows while a row mostly holds + # its direction, shrinks once it mostly flips, and holds at the + # flip_fraction == 0.5 noise/edge-of-stability point. Symmetric + # up/down is load-bearing -- any asymmetry drags a noisy row to + # min_lr (see class docstring) -- and abstaining (rather than counting + # frozen elements as agreement) keeps a pool of dead elements from + # quietly ratcheting the lr upward. + bump = group["lr_bump_rate"] + prod = cur_sign * prev_sign # int8 {-1, 0, +1} per element + if dims is not None: + num = prod.to(torch.float32).sum(dim=dims) + den = (prod != 0).to(torch.float32).sum(dim=dims).clamp_(min=1.0) + log_dir = num.div_(den).mul_(bump) + else: + log_dir = prod.to(torch.float32).mul_(bump) + # EMA-smooth the per-row nudge so a single noisy step doesn't swing + # the lr, then apply it multiplicatively (geometric move at every + # scale across [min_lr, max_lr]). + ema = state["dir_ema"] + beta = group["dir_beta"] + ema.mul_(beta).add_(log_dir, alpha=1.0 - beta) + lr_t.mul_(torch.exp(ema)).clamp_(min=group["min_lr"], max=group["max_lr"]) - # Map consistency to a direction in [-1, 1], measured relative to the - # window's noise floor b (the consistency a pure-noise row shows by - # chance, computed from the window size): agreement == b -> 0 (hold), - # 1.0 (perfect) -> +1 (lr up), below b -> down. There is no manual - # target -- holding at b is fully automatic and self-balancing per - # row: the lr grows while a row is more consistent than chance and - # settles at whatever lr makes its consistency meet b, so noisy and - # clean layers each find their own operating point. Measuring - # relative to b also makes the signal window-size independent. - b = group["agreement_floor"] - direction = agreement.sub_(b).div_(1.0 - b).clamp_(-1.0, 1.0) - # Multiplicative (geometric) bump: lr *= exp(direction * lr_bump_rate). - # lr_bump_rate is a fractional rate (~lr_bump_rate per step for small - # values), giving a uniform relative move at every scale across - # [min_lr, max_lr] instead of a fixed absolute amount. - lr_t.mul_(torch.exp(direction.mul_(group["lr_bump_rate"]))).clamp_( - min=group["min_lr"], max=group["max_lr"] - ) - - # Record this step's polarity and trim the window to the configured size. - pol_hist.append(cur_polarity) - if len(pol_hist) > group["polarity_history_count"]: - del pol_hist[0] + state["prev_sign"] = cur_sign state["step"] += 1 wd = group["weight_decay"] @@ -466,7 +464,9 @@ class Automagic3(torch.optim.Optimizer): st = self.state.get(p) if st is not None and isinstance(st.get("lr"), torch.Tensor): st["lr"] = st["lr"].to(torch.float32) - # Polarity history is transient; rebuild it after load rather - # than trying to persist/cast a list of bool tensors. - if st is not None and "pol_hist" in st: - st["pol_hist"] = [] + # prev_sign / dir_ema are transient; rebuild them after load + # rather than persisting a sign tensor and an fp32 EMA. + if st is not None and "prev_sign" in st: + st["prev_sign"] = None + if st is not None and isinstance(st.get("dir_ema"), torch.Tensor): + st["dir_ema"] = torch.zeros_like(st["dir_ema"], dtype=torch.float32)