Automagic 3 rework. Stable in my testing.
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user