Another complete rework of automagic3. Added a decay to the LR spread to the mean to prevent LRs fighting with eachother

This commit is contained in:
Jaret Burkett
2026-06-09 12:08:45 -06:00
parent 9e99d3ce5d
commit 01b6a9806b

View File

@@ -10,35 +10,39 @@ class Automagic3(torch.optim.Optimizer):
"""
Automagic v3.
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 whether the per-element update direction *flipped* vs the previous
step (RProp-style edge-of-stability control).
A single learning rate is kept per parameter tensor (one lr per weight
matrix / layer). Each step the lr is nudged by whether the per-element update
direction *flipped* vs the previous step (RProp-style edge-of-stability
control), aggregated over the whole tensor.
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)``.
restoring force. Each element votes: agree with last step -> +``lr_bump_rate``;
flip -> -``lr_bump_rate`` (symmetric). The votes are reduced to ONE value for
the whole tensor and EMA-smoothed over ~``lr_smoothing_steps`` steps, then
applied multiplicatively: ``lr *= exp(nudge)``. Reducing over the whole tensor
(rather than per row/channel) is deliberate: per-channel lrs let coupled
channels fight -- one drives its lr up while a neighbour drives its down to
compensate -- and split to opposite extremes, which visibly wrecks the model;
one lr per tensor makes opposing channels cancel into a single vote.
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].
The flip equilibrium is flip fraction == 0.5, the only rate that is both the
pure-noise point and the edge of stability: a tensor descending cleanly flips
less than half the time -> its lr grows; once it overshoots it flips more than
half -> its lr shrinks; pure noise flips ~half -> the lr holds. Symmetry is
load-bearing (any asymmetry drags a noise-dominated tensor to zero). Elements
with an exactly-zero update (dead/masked grads, low-precision underflow) carry
no direction and abstain from the vote.
On top of the per-tensor flip control, every tensor's lr is gently pulled
toward the GLOBAL average lr each step (``lr_pull``, geometric / log-space:
``lr *= (avg_lr / lr) ** lr_pull``). This mean-reversion is the restoring
force that replaces hard min/max clamps: layers may settle at their own level,
but cannot drift apart to opposite extremes (one frozen near zero, one running
away) -- the failure mode that destroys full finetunes. The target is the
emergent average across layers, not a fixed number, so it stays fully
automatic; there are no min/max lr bounds. ``lr_bump_rate`` sets how fast the
lr moves, not where it lands.
With ``fused=True`` (default) the step is fused into the backward pass via
``register_post_accumulate_grad_hook``: each parameter is updated and its
@@ -59,20 +63,21 @@ class Automagic3(torch.optim.Optimizer):
Parameters
----------
lr : float
Starting learning rate for every row. The controller adapts away from
Starting learning rate for every layer. The controller adapts away from
this, so it is a launch point, not a tuned target -- a low value just
lets the lr ramp up on its own (there is no warmup). Values above 1e-3
are rejected and forced back to 1e-6.
min_lr, max_lr : float
Hard clamps on every per-row lr. ``max_lr`` doubles as the per-step
trust region: since the per-element update is capped at
``clip_threshold`` (~1), a weight moves at most ~``max_lr`` per step, so
keep it modest -- a high ceiling lets the hottest rows take destabilising
steps (and, when fused, the trainer's grad clip can't catch them).
are rejected and forced back to 1e-6. There are no min/max lr clamps;
the mean-reversion pull keeps the spread bounded instead.
lr_bump_rate : float
Fractional, log-space size of each lr nudge (~10% at 0.1). Sets how fast
the lr moves, NOT where it settles (the flip dynamics fix that); a full
up-nudge and a full down-nudge cancel exactly.
lr_pull : float
Per-step strength of the mean-reversion pulling each layer's lr toward
the global average (log space; 0 disables). This is the restoring force
that replaces min/max clamps and stops layers drifting to opposite
extremes. Small (default 0.05) lets layers keep their own level while
bounding the spread; larger forces all layers toward a common lr.
beta2 : float
EMA decay for the second moment, as in Adam/Adafactor.
eps : float
@@ -87,7 +92,7 @@ class Automagic3(torch.optim.Optimizer):
How many steps of the flip signal to EMA-average before nudging the lr
(>=1, default 3). Higher = smoother/slower lr, lower = twitchier/faster;
it does not change where the lr lands. Held as an EMA, so it costs O(1)
state per row regardless of the value.
state per layer regardless of the value.
fused : bool
If True (default), each param is updated inside the backward pass the
moment its grad is ready -- low peak VRAM, but it bypasses the trainer's
@@ -97,12 +102,15 @@ class Automagic3(torch.optim.Optimizer):
Improvements over v2
--------------------
1. Per-row learning rate (was a single scalar per parameter tensor).
v2 kept one lr for an entire weight matrix; v3 keeps one per output
channel (per element for 1D params). Plain English: different neurons in
the same layer can now learn at different speeds instead of being forced
to share one rate, so a layer where some rows have converged and others
have not is handled gracefully.
1. Per-layer lr with mean-reversion (v2 had one static lr per tensor and no
coupling between layers). v3 still keeps one lr per tensor, but it is
adaptive (driven by the flip controller) and every layer's lr is pulled
toward the global average (``lr_pull``). Plain English: each layer finds
its own learning rate automatically, but no layer can run away or freeze
relative to the others -- which is what used to split a full finetune into
over-cooked and dead layers and destroy it. (An earlier v3 used a separate
lr per output channel; coupled channels fought and split to opposite
extremes, so it was reduced back to one lr per tensor.)
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
@@ -117,8 +125,8 @@ class Automagic3(torch.optim.Optimizer):
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
min_lr and a negligible one near max_lr. v3 multiplies by
fixed absolute amount, so the same bump was a huge relative jump when the
lr was tiny and a negligible one when it was large. v3 multiplies by
``exp(signal * lr_bump_rate)`` -- a fixed *percentage* step. Plain
English: the lr moves at the same relative pace whether it is tiny or
large, traverses its whole range in a predictable number of steps, and a
@@ -145,9 +153,8 @@ class Automagic3(torch.optim.Optimizer):
self,
params,
lr: float = 1e-6,
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
lr_pull: float = 0.025, # per-step pull of each layer's lr toward the global avg (log space)
beta2: float = 0.999,
eps: float = 1e-30,
clip_threshold: float = 1.0,
@@ -162,21 +169,24 @@ class Automagic3(torch.optim.Optimizer):
lr_smoothing_steps = max(1, int(lr_smoothing_steps))
defaults = dict(
lr=lr,
min_lr=min_lr,
max_lr=max_lr,
lr_bump_rate=lr_bump_rate,
lr_pull=max(0.0, float(lr_pull)),
beta2=beta2,
eps=eps,
clip_threshold=clip_threshold,
weight_decay=weight_decay,
lr_smoothing_steps=lr_smoothing_steps,
# EMA decay for the per-row lr nudge, derived from the smoothing
# EMA decay for the per-layer 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)
self.fused = fused
# Global geometric-mean lr across all layers; the mean-reversion pull
# targets this. Seeded at the start lr (all layers equal) and refreshed
# each .step() from the current per-layer lrs.
self._avg_lr = float(lr)
self._hook_handles = []
for group in self.param_groups:
for p in group["params"]:
@@ -282,18 +292,20 @@ class Automagic3(torch.optim.Optimizer):
def _init_state(self, p: torch.Tensor, group: dict) -> None:
state = self.state[p]
state["step"] = 0
# Per-row lr: one entry per output channel for >=2D params, per element
# for 1D params, a scalar for 0D params.
lr_shape = (p.shape[0],) if p.dim() >= 2 else p.shape
state["lr"] = torch.full(
lr_shape, float(group["lr"]), dtype=torch.float32, device=p.device
# ONE lr per parameter tensor (scalar), not per row/channel. Per-channel
# lrs let coupled channels fight -- one drives up while another drives
# down to compensate -- so they split to the rails and wreck the model.
# A single lr per tensor averages the vote over all elements, so opposing
# channels cancel instead of diverging.
state["lr"] = torch.tensor(
float(group["lr"]), dtype=torch.float32, device=p.device
)
# 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)
# EMA of the (scalar) log lr-nudge, smoothing the flip signal over time.
state["dir_ema"] = torch.zeros((), 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
@@ -401,50 +413,45 @@ class Automagic3(torch.optim.Optimizer):
# negatives rather than bucketed with them by a bare ``> 0``.
cur_sign = update.sign().to(torch.int8)
prev_sign = state["prev_sign"]
# dims: the within-row axes to reduce the per-element vote over (so each
# output channel gets one nudge). lr_b: the per-row lr reshaped to
# broadcast across the full param for the weight update below. For 1D
# params there is no row to reduce over (one lr per element).
lr_t = state["lr"]
if p.dim() >= 2:
dims = tuple(range(1, p.dim()))
lr_b = lr_t.view(lr_t.shape[0], *([1] * (p.dim() - 1)))
else:
dims = None
lr_b = lr_t
lr_t = state["lr"] # scalar (one lr for the whole tensor)
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.
# isn't counted), no separate masking needed.
#
# 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.
# Reduced over the WHOLE tensor (not per row) this is
# bump*(1 - 2*flip_fraction): the lr grows while the layer mostly
# holds its direction, shrinks once it mostly flips, and holds at the
# flip_fraction == 0.5 point. Aggregating over all elements means
# channels pushing opposite ways cancel into one vote instead of
# splitting to opposite lr extremes.
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]).
# Reduce directly on the int8 product (sum -> int64, no full-size
# float cast; count_nonzero -> valid votes) -- two reductions instead
# of casting the whole tensor to float twice.
num = prod.sum()
den = prod.count_nonzero().clamp_(min=1)
log_dir = num.float().div_(den.float()).mul_(bump)
# EMA-smooth the nudge so a single noisy step doesn't swing the lr,
# then apply it multiplicatively (a fixed fractional move).
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"])
lr_t.mul_(torch.exp(ema))
# Mean-reversion: pull this layer's lr toward the global average lr
# (geometric, in log space) -- lr *= (avg/lr)**lr_pull. This is the
# restoring force that replaces the hard min/max rails: layers can
# still settle at their own level, but can't drift apart to opposite
# extremes (one frozen, one runaway) and wreck the model. The pull is
# toward an emergent average, not a fixed target, so it stays
# automatic. self._avg_lr is refreshed once per .step().
pull = group["lr_pull"]
if pull > 0.0 and self._avg_lr > 0.0:
lr_t.mul_(lr_t.reciprocal().mul_(self._avg_lr).pow_(pull))
state["prev_sign"] = cur_sign
state["step"] += 1
@@ -453,10 +460,10 @@ class Automagic3(torch.optim.Optimizer):
if p.dtype == torch.float32:
# Decoupled weight decay folded in (update += wd*p), then a single
# fused p -= lr_b * update.
# fused p -= lr * update (lr is a scalar, broadcasts).
if wd != 0.0:
update.add_(p, alpha=wd)
p.addcmul_(update, lr_b, value=-1.0)
p.addcmul_(update, lr_t, value=-1.0)
else:
# Low precision: apply the update in fp32 then stochastically round
# back, so tiny updates aren't lost to round-to-nearest. Single
@@ -464,7 +471,7 @@ class Automagic3(torch.optim.Optimizer):
new_p_fp32 = p.to(torch.float32)
if wd != 0.0:
update.add_(new_p_fp32, alpha=wd)
new_p_fp32.addcmul_(update, lr_b, value=-1.0)
new_p_fp32.addcmul_(update, lr_t, value=-1.0)
self._stochastic_copy_(p, new_p_fp32)
p.grad = None
@@ -493,25 +500,39 @@ class Automagic3(torch.optim.Optimizer):
if p.grad is None:
continue
self._update_param(p, group)
self._refresh_avg_lr()
return loss
def _all_lrs(self) -> list:
# The per-layer lr scalars (0-d tensors), gathered without any device
# sync. Callers stack these and reduce in one op so there is a single
# GPU->CPU sync instead of one per layer.
return [
st["lr"]
for group in self.param_groups
for p in group["params"]
if (st := self.state.get(p)) is not None and "lr" in st
]
def _refresh_avg_lr(self) -> None:
# Global geometric-mean lr across all layers, the mean-reversion target.
# Geometric (mean of log) because the lr is controlled/pulled in log
# space. One stacked reduction + one sync, refreshed once per step.
lrs = self._all_lrs()
if lrs:
self._avg_lr = float(torch.stack(lrs).log_().mean().exp_())
def get_learning_rates(self) -> List[float]:
# Reporting helper: one representative lr per param group. Uses the
# arithmetic mean of the per-row lrs because it is magnitude-weighted --
# the rows actually taking sizeable steps dominate, so the number reads
# like the scalar lr you'd set in another trainer. (A geometric mean or
# median instead gets dragged toward min_lr by frozen rows and reads
# misleadingly low.) A handful of rows riding at max_lr can lift this
# average even while the typical row is flat; that's expected and bounded
# by the max_lr clamp, not a runaway.
# Reporting helper: one representative lr per param group -- the mean of
# the per-layer (per-tensor) lrs. Stacked reduction -> one sync per group.
out = []
for group in self.param_groups:
lrs = [
float(self.state[p]["lr"].mean())
self.state[p]["lr"]
for p in group["params"]
if p in self.state and "lr" in self.state[p]
]
out.append(sum(lrs) / len(lrs) if lrs else float(group["lr"]))
out.append(float(torch.stack(lrs).mean()) if lrs else float(group["lr"]))
return out
def get_avg_learning_rate(self) -> float:
@@ -536,3 +557,5 @@ class Automagic3(torch.optim.Optimizer):
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)
# Rebuild the global average lr from the restored per-layer lrs.
self._refresh_avg_lr()