Reworked automagic v3 again. Seems more stable. Still testing.
This commit is contained in:
parent
5d6887fd98
commit
a1ac6e8b01
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue