Automagic 3 rework. Stable in my testing.

This commit is contained in:
Jaret Burkett
2026-06-12 07:52:44 -06:00
parent 53ebb93edb
commit 55ce6570f2

View File

@@ -10,39 +10,67 @@ class Automagic3(torch.optim.Optimizer):
"""
Automagic v3.
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 single learning rate is kept per param group (typically: one lr for
the whole run). The control principle: the lr RISES while elements hold
a decisively consistent update direction at the current step size, FALLS
while their signs decisively alternate (the overshoot signature: weights
hopping across a minimum flip sign step to step -- shrinking the step is
what makes a trajectory reappear at a finer scale), and HOLDS on
everything in between, which is treated as noise.
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 -> +``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.
Each element keeps a window of its last H (= ``polarity_history``,
default 4) update sign bits ("is the update positive", 1-bit packed) --
H/8 bytes per element (half a byte at the default), the only
per-element optimizer state. A short window suffices because verdicts
are pooled across the whole group: millions of voters make weak
common-mode evidence visible long before any single element is
decisive, and the window length is also the controller's reaction lag
and warmup.
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.
Vote rule (per element)
-----------------------
Only the two perfectly decisive window states vote; everything else is
noise:
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.
up all H signs agree +1 * |update| ("step too small")
down all H-1 transitions flip -1 * |update| ("step too large":
(perfect alternation) the overshoot signature)
else any imperfect window 0 (noise)
The two events are exact mirrors with IDENTICAL pure-noise probability
(2 of the 2^H possible windows each; ~0.8% per element at H=8), so equal
weights balance exactly -- no correction factors, no tiers. Per element
the events are rare, but the verdict is pooled over the whole group
(millions of elements -> tens of thousands of voters per step even
under pure noise, mean zero), so the pooled signal is smooth and a real
trend or real overshoot shifts it decisively. A majority being overshot
always outvotes a persistent minority, which is what anchors the lr's
absolute level without external rails. Weighting by |update| lets the
elements actually moving the weights dominate; an exact-zero update
records as the negative bit, but such dead/masked elements carry zero
weight anyway. A tensor abstains entirely until its window has filled
(the first H steps, and again after a history reset on resume).
ONE learning rate per param group -- not per tensor. Every element of
every tensor in the group votes into a single pool, and the group lr is
nudged once per step by the pooled result, applied multiplicatively
with NO gain knob: ``lr *= exp(vote)`` -- the lr moves at exactly the
rate the model votes for it. A fully unanimous pool (practically
unreachable) would move e ~= 2.7x per step; the silent majority dilutes
the pooled vote, so realistic moves are a few percent per step, and the
worst-case overshoot past the edge is bounded by the H-step window lag
before alternation votes answer. There is no
noise-floor estimation, no smoothing, no significance test: the polarity
windows are the only indicator. Pooling at group level (rather than per
tensor, and originally rather than per channel) is the load-bearing
choice: COUPLED tensors fight per-tensor lrs exactly like coupled
channels fight per-channel ones. A Q/K pair is the canonical case --
Q's weights scaling up while K's scale down preserves the attention
logits, so the gradients reward whichever asymmetry randomly seeded
first: Q votes "too slow" and climbs while K votes "too fast" and sinks,
self-reinforcing without bound. One shared lr makes those opposing votes
cancel in the pool instead of diverging, so only common-mode evidence
("the whole group's step is too small / too large") moves the 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
@@ -63,21 +91,10 @@ class Automagic3(torch.optim.Optimizer):
Parameters
----------
lr : float
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. 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.
Starting learning rate for every group. The controller adapts away
from this in whichever direction the pooled vote points, so it is a
launch point, not a tuned target. There are no min/max lr clamps
(only a numerical overflow guard far outside the usable range).
beta2 : float
EMA decay for the second moment, as in Adam/Adafactor.
eps : float
@@ -88,11 +105,14 @@ class Automagic3(torch.optim.Optimizer):
step.
weight_decay : float
Decoupled (AdamW-style) weight decay; 0 disables it.
lr_smoothing_steps : int
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 layer regardless of the value.
polarity_history : int
Sign-history window length H (2 to 64, default 4); H/8 bytes of
state per element. Longer windows make the two vote events rarer
and more decisive (probability 2^(1-H) each under noise -- a real
trend's excess grows ~(1+rho)^H), so detection sharpens, at the
cost of memory, an H-step reaction lag/warmup, and fewer voters
per step. Changing it on resume resets the histories cleanly (one
re-warmup of H steps).
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
@@ -102,36 +122,34 @@ class Automagic3(torch.optim.Optimizer):
Improvements over v2
--------------------
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.)
1. One adaptive lr per param group (v2 had one static lr per tensor).
Plain English: the group finds its learning rate automatically, and 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. (Earlier v3s used a separate lr per output channel, then
per tensor; each level let coupled units -- channels, then Q/K-style
tensor pairs -- fight and split to opposite extremes, so the lr was
pooled one level up each time until the fighting was structurally
impossible.)
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.
2. Direction-consistency lr control with a real equilibrium. v2 bumped
the lr from raw single-step 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
votes from each element's recent sign window (see the vote rule
above). Plain English: the lr speeds up while the model holds a
trajectory, backs off hard when it overshoots, and holds steady on
pure noise.
3. Multiplicative (geometric) lr bump (was additive). v2 added/subtracted a
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
``exp(vote)`` -- 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
full up bump is exactly cancelled by a full down bump (no drift). The
knob was renamed ``lr_bump`` -> ``lr_bump_rate`` to signal the change.
full up bump is exactly cancelled by a full down bump (no drift); the
gain knob was removed entirely once the vote became a pooled
fraction with natural log-units.
4. Stochastic rounding for fp16, not just bf16. v2 only rounded bf16
write-backs and let fp16 fall back to round-to-nearest, silently
@@ -142,51 +160,44 @@ class Automagic3(torch.optim.Optimizer):
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
and parameter update are fused into one ``addcmul_``; and the per-element
agree/flip vote is a single int8 sign-product (``cur_sign * prev_sign``)
instead of several boolean masks plus float casts. Plain English: each
step issues fewer GPU passes over the weights, so it runs faster (notably
in bf16/fp16) without changing the result.
and parameter update are fused into one ``addcmul_``; the per-element
direction and flip sums are recomputed from the 1-bit history planes
in a single batched unpack and two integer reductions, and scored
with three boolean compares and weighted sums. Plain English: each
step issues few GPU passes over the weights, so it runs fast
(notably in bf16/fp16).
"""
def __init__(
self,
params,
lr: float = 1e-6,
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,
weight_decay: float = 0.0,
lr_smoothing_steps: int = 3, # lr-nudge EMA smoothing horizon, in steps (min 1)
polarity_history: int = 8, # sign-history window length (2-64)
fused: bool = True,
):
if lr > 1e-3:
print(f"Warning! Start lr {lr} is very high; forcing to 1e-6.")
lr = 1e-6
# The lr nudge is EMA-smoothed over ~this many steps; at least 1.
lr_smoothing_steps = max(1, int(lr_smoothing_steps))
# No clamping: a too-high start just oscillates immediately and
# the controller drives it down.
print(
f"Note: start lr {lr} is high; the controller will correct it "
f"(the pooled vote will walk it down)."
)
defaults = dict(
lr=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-layer lr nudge, derived from the smoothing
# horizon (n steps -> beta = n/(n+1)).
dir_beta=lr_smoothing_steps / (lr_smoothing_steps + 1.0),
polarity_history=max(2, min(64, int(polarity_history))),
)
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._rebuild_group_index()
self._hook_handles = []
for group in self.param_groups:
for p in group["params"]:
@@ -260,6 +271,56 @@ class Automagic3(torch.optim.Optimizer):
noise = torch.rand_like(v).sub_(0.5).mul_(ulp)
return v.add_(noise).to(dtype)
# Per-device cached constants for pack/unpack (avoid re-allocating a tiny
# tensor on every call).
_PACK_CONSTS: dict = {}
@classmethod
def _pack_consts(cls, device):
consts = cls._PACK_CONSTS.get(device)
if consts is None:
consts = (
torch.tensor(
[1, 2, 4, 8, 16, 32, 64, 128], device=device, dtype=torch.uint8
),
torch.tensor(
[0, 1, 2, 3, 4, 5, 6, 7], device=device, dtype=torch.uint8
),
)
cls._PACK_CONSTS[device] = consts
return consts
@classmethod
def _pack_bits(cls, bits: torch.Tensor) -> torch.Tensor:
# Pack sign bits (bool / {0, 1}) 8 per byte (uint8), as a base-2 dot
# product per group of 8 (two kernels rather than per-slice shift/or
# chains).
weights, _ = cls._pack_consts(bits.device)
flat = bits.reshape(-1).to(torch.uint8)
pad = (-flat.numel()) % 8
if pad:
flat = torch.cat([flat, flat.new_zeros(pad)])
return (flat.view(-1, 8) * weights).sum(-1, dtype=torch.uint8)
@classmethod
def _unpack_bits(cls, packed: torch.Tensor, numel: int) -> torch.Tensor:
# Inverse of _pack_bits: uint8 -> flat uint8 {0, 1} of length numel.
_, shifts = cls._pack_consts(packed.device)
vals = (packed.unsqueeze(-1) >> shifts).bitwise_and_(1)
return vals.view(-1)[:numel]
def _rebuild_group_index(self) -> None:
# param -> index of its param group, plus per-group vote accumulators
# (weighted vote mass and total weight, gathered across every tensor
# in the group during the step and applied once in .step()). The map
# exists because the fused hooks cannot rely on group-dict identity:
# the parent's load_state_dict replaces the group dicts.
self._param_group_index = {
p: gi for gi, group in enumerate(self.param_groups) for p in group["params"]
}
self._group_num: List = [None] * len(self.param_groups)
self._group_den: List = [None] * len(self.param_groups)
@classmethod
def _stochastic_copy_(cls, dst: torch.Tensor, src_fp32: torch.Tensor) -> None:
# Stochastically round the fp32 ``src`` into the low-precision ``dst`` in
@@ -292,20 +353,27 @@ class Automagic3(torch.optim.Optimizer):
def _init_state(self, p: torch.Tensor, group: dict) -> None:
state = self.state[p]
state["step"] = 0
# 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.
# The group lr, mirrored per param (every param in a group receives
# identical multiplicative nudges, so these stay equal; storing per
# param rides the normal state_dict machinery and tolerates
# multi-device groups).
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 (scalar) log lr-nudge, smoothing the flip signal over time.
state["dir_ema"] = torch.zeros((), dtype=torch.float32, device=p.device)
# Ring buffer of per-element update sign bits, one 1-bit-packed
# plane per step (H/8 bytes per element). Sums are recomputed from
# the planes each step rather than stored -- the history is the
# ONLY per-element state.
H = group["polarity_history"]
width = (p.numel() + 7) // 8
state["sign_history"] = torch.zeros(
(H, width), dtype=torch.uint8, device=p.device
)
# Index of the OLDEST plane (the one overwritten next step).
state["hist_idx"] = 0
# Number of real sign planes stored so far; the controller is gated
# until the window is full (there is no per-element abstain state).
state["hist_fill"] = 0
if p.dim() >= 2:
state["exp_avg_sq_row"] = torch.zeros(
p.shape[:-1], dtype=p.dtype, device=p.device
@@ -338,15 +406,13 @@ class Automagic3(torch.optim.Optimizer):
if grad.dtype != torch.float32:
grad = grad.to(torch.float32)
# This step is fused into backward, so the trainer's grad clipping and
# nan/inf-skip run too late to protect us -- the weights are already
# updated here. A single non-finite gradient would poison the
# second-moment EMA (NaN*beta2 + ... stays NaN forever) and corrupt the
# weights, which surfaces as the model "randomly" blowing up. Neutralise
# non-finite grads in place (we own this fp32 grad) so those elements
# contribute nothing this step instead of destroying state. Large but
# finite grads are left alone -- the second-moment normalisation already
# bounds their effect.
# In fused mode this runs inside backward, so the trainer's grad
# clipping and nan/inf-skip come too late to protect us. A single
# non-finite gradient would poison the second-moment EMA (NaN stays
# NaN forever) and corrupt the weights, so neutralise non-finite
# grads in place (we own this fp32 copy); those elements contribute
# nothing this step. Large but finite grads are left alone -- the
# second-moment normalisation already bounds their effect.
grad.nan_to_num_(nan=0.0, posinf=0.0, neginf=0.0)
beta2 = group["beta2"]
@@ -400,60 +466,63 @@ 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"])
# 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"] # scalar (one lr for the whole tensor)
# Direction-consistency lr control (the vote rule -- see the class
# docstring). The second-moment scale, RMS clip and clamp are all
# positive, so the sign bit is exactly sign(grad); an exact-zero
# update records as the negative bit, harmless because its |update|
# vote weight is zero.
cur_bits = update.gt(0.0)
hist = state["sign_history"] # (H, numel/8) 1-bit packed uint8
idx = state["hist_idx"] # oldest plane (overwritten below)
H = hist.shape[0]
lr_t = state["lr"] # this param's mirror of the shared group lr
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.
#
# 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
# 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))
# 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))
# Slide the window first so the vote sees the freshest H signs.
hist[idx].copy_(self._pack_bits(cur_bits))
state["hist_idx"] = (idx + 1) % H
# The planes hold garbage until H real signs have been stored (fresh
# start or a history reset on resume): gate the controller, not the
# parameter update, until the window is full.
fill = min(H, state["hist_fill"] + 1)
state["hist_fill"] = fill
if fill == H:
# Extremes-only vote (see the class docstring): all H signs
# agreeing votes up, perfect alternation (all H-1 transitions
# flipping) votes down -- the two events have identical
# pure-noise probability (2 of the 2^H windows each), so equal
# +/-1 weights balance exactly. The planes are rolled into
# chronological order so adjacent rows are adjacent steps; XOR
# of neighbour rows marks per-bit flips. The weighted vote mass
# and total weight are ACCUMULATED into this tensor's group; the
# single group lr is nudged once per step in .step().
_, shifts = self._pack_consts(hist.device)
chron = torch.roll(hist, -state["hist_idx"], dims=0)
bits = (
(chron.unsqueeze(-1) >> shifts)
.bitwise_and_(1)
.view(H, -1)[:, : update.numel()]
)
s1 = bits.sum(0, dtype=torch.int16)
flips = (bits[1:] ^ bits[:-1]).sum(0, dtype=torch.int16)
up = s1.eq(H).logical_or_(s1.eq(0))
down = flips.eq(H - 1)
w = update.abs().view(-1)
num = (w * up).sum().sub_((w * down).sum())
den = w.sum()
gi = self._param_group_index.get(p)
if gi is not None:
if self._group_num[gi] is None:
self._group_num[gi] = num
self._group_den[gi] = den
else:
acc = self._group_num[gi]
if num.device != acc.device:
num = num.to(acc.device)
den = den.to(acc.device)
acc.add_(num)
self._group_den[gi].add_(den)
state["prev_sign"] = cur_sign
state["step"] += 1
wd = group["weight_decay"]
@@ -500,31 +569,37 @@ class Automagic3(torch.optim.Optimizer):
if p.grad is None:
continue
self._update_param(p, group)
self._refresh_avg_lr()
self._apply_group_votes()
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 _apply_group_votes(self) -> None:
# ONE lr nudge per group per step, from the pooled vote of every
# element of every tensor in the group (see the class docstring on
# why pooling at group level is load-bearing). Each param's lr tensor
# receives the same multiplicative factor, so they stay identical --
# effectively a single group lr, stored per param only so it rides
# the normal state_dict machinery. All tensor ops: no GPU sync.
for gi, group in enumerate(self.param_groups):
num = self._group_num[gi]
if num is None:
continue
den = self._group_den[gi]
signal = num.div_(den.clamp_(min=1e-30)).clamp_(-1.0, 1.0)
factor = torch.exp(signal)
for p in group["params"]:
st = self.state.get(p)
if st is None or "lr" not in st:
continue
lr_t = st["lr"]
f = factor if factor.device == lr_t.device else factor.to(lr_t.device)
# Numerical overflow guard only -- NOT a control rail
# (decades outside the usable range).
lr_t.mul_(f).clamp_(min=1e-30, max=1e3)
self._group_num[gi] = None
self._group_den[gi] = None
def get_learning_rates(self) -> List[float]:
# Reporting helper: one representative lr per param group -- the mean of
# the per-layer (per-tensor) lrs. Stacked reduction -> one sync per group.
# Reporting helper: the (shared) lr of each param group.
out = []
for group in self.param_groups:
lrs = [
@@ -543,19 +618,64 @@ class Automagic3(torch.optim.Optimizer):
# Parent casts every fp state tensor to param.dtype; force lr back to fp32
# so subsequent lr bumps aren't rounded away on bf16 weights.
super().load_state_dict(state_dict)
# Constructor args always win over whatever was saved in the checkpoint.
# Hyperparameters are NOT loaded from the checkpoint: constructor args
# always win, so any setting can be changed mid-run just by passing a
# different value when resuming. Only the adaptive state is restored
# -- the group lr and the sign history (when its geometry still
# matches the current config).
for group in self.param_groups:
for k, v in self.defaults.items():
group[k] = v
# One lr per group: unify the restored lrs to their geometric
# median (they are already identical for checkpoints from this
# version; older per-tensor checkpoints land on a sane middle).
lrs = [
st["lr"]
for p in group["params"]
if (st := self.state.get(p)) is not None
and isinstance(st.get("lr"), torch.Tensor)
]
med = None
if lrs:
dev = lrs[0].device
med = (
torch.stack([t.to(torch.float32).to(dev) for t in lrs])
.log_()
.median()
.exp_()
)
for p in group["params"]:
st = self.state.get(p)
if st is not None and isinstance(st.get("lr"), torch.Tensor):
if st is None:
continue
if isinstance(st.get("lr"), torch.Tensor):
st["lr"] = st["lr"].to(torch.float32)
# 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)
# Rebuild the global average lr from the restored per-layer lrs.
self._refresh_avg_lr()
if med is not None:
st["lr"].copy_(med.to(st["lr"].device))
# Sign history: keep it when its geometry matches the current
# config (the parent cast it to param dtype; recover by shape).
# On any mismatch (e.g. a checkpoint from an older window
# layout) -- start fresh.
numel = p.numel()
H = group["polarity_history"]
width = (numel + 7) // 8
sh = st.get("sign_history")
hist_ok = (
isinstance(sh, torch.Tensor)
and sh.shape == (H, width)
and isinstance(st.get("hist_idx"), int)
and 0 <= st["hist_idx"] < H
and isinstance(st.get("hist_fill"), int)
and 0 <= st["hist_fill"] <= H
)
if hist_ok:
st["sign_history"] = sh.to(torch.uint8)
else:
st["sign_history"] = torch.zeros(
(H, width), dtype=torch.uint8, device=p.device
)
st["hist_idx"] = 0
st["hist_fill"] = 0
# The parent rebuilt the group dicts; remap params to groups and
# reset the vote accumulators.
self._rebuild_group_index()