- Python 100%
| Filename | Latest commit message | Latest commit date |
|---|---|---|
|
|
||
| examples | ||
| patches | ||
| src/prejac | ||
| tests | ||
| .gitignore | ||
| LICENSE | ||
| NOTICE.md | ||
| pyproject.toml | ||
| README.md | ||
| uv.lock | ||
prejac
Training to converge early, with as little shown as possible.
Optimizers for deep learning.
Install
uv pip install prejac
# For faster* Gram Newton-Schulz kernels on Hopper/datacentre Blackwell, install the gns extra:
uv pip install prejac[gns]
Optimizers
Regular and Step-skipping Optimizers
| Optimizer | Description |
|---|---|
AdamW |
AdamW optimizer with selectable AdamW or AdamC weight decay. |
SkipStepAdamW |
AdamW with loss-spike step skipping. |
Muon |
Muon optimizer with Newton-Schulz iteration. Non-matrix params get AdamW. Includes decoupled weight decay (Moonlight). |
Mousse |
Muon with Shampoo-style curvature preconditioning and norm grafting. Supports ordinary tensors/DDP; rejects DTensors. |
SkipStepMuon |
Muon with loss-spike step skipping. |
Lion |
Lion optimizer. |
SkipStepLion |
Lion with loss-spike step skipping. |
NoOpOptimizer |
No-op optimizer for testing. |
ScheduleFree/ScheduleFree+ Optimizers
Refer to ScheduleFree+ (Defazio, 2026) for more details.
| Optimizer | Description |
|---|---|
ScheduleFreeAdamW |
SkipStepAdamW with schedule-free iterate averaging and a fixed learning rate. |
ScheduleFreePlusAdamW |
ScheduleFreeAdamW with a learning-rate-free Polyak step. |
Mechanic Optimizers
Mechanic (Cutkosky et al., 2023) learns a global scale for the cumulative updates produced by a base optimizer.
| Optimizer | Description |
|---|---|
MechanicAdamW |
Mechanic around prejac AdamW, including AdamW/AdamC decay options. |
MechanicMuon |
Mechanic around prejac Muon, including automatic AdamW routing for non-matrix parameters. |
MechanicOptimizer |
Wrap an existing first-order PyTorch optimizer instance. |
Miscellaneous Optimizers
| Optimizer | Description | Reference |
|---|---|---|
PoLoRA |
Preconditioned orthogonalized updates for explicitly paired LoRA A/B factors. | PoLoRA (Ghosh et al., 2026) |
Quick start
Mechanic
import prejac
import torch
optim = prejac.MechanicAdamW(model.parameters(), lr=1.0)
# Or: prejac.MechanicMuon(model.parameters(), lr=1.0)
# Or: prejac.MechanicOptimizer(torch.optim.SGD(model.parameters(), lr=1.0))
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optim, T_max=total_steps)
for batch in loader:
optim.zero_grad()
loss = model(batch).loss
loss.backward()
optim.step()
scheduler.step()
print(optim.scale) # learned scale; param_groups retain the base learning rates
Mechanic tunes the learning-rate scale; you can still use a learning-rate
schedule. Both named optimizers default to a base lr=1.0 and accept their
base optimizer's options and parameter groups. Groups can have different
learning rates, but share one tuner. Use the wrapper for training, schedulers,
and checkpoints; its base_optimizer exposes the underlying optimizer and
shares the same parameter-group dictionaries. Save and restore the model and
optim.state_dict() together as usual.
The keyword-only tuner options are s_init=1e-8, s_decay=0.01,
mechanic_betas=(0.9, 0.99, 0.999, 0.9999, 0.99999, 0.999999),
mechanic_eps=1e-8, and store_delta=True. The tuner's betas and epsilon
are independent of Adam's betas/eps and Muon's epsilon. s_decay
regularizes the tuner; the base optimizer's weight_decay retains its usual
meaning. An epsilon floor on the maximum inner-product magnitude bootstraps
the first scale to s_init; subsequent steps use Algorithm 1's clamped reward.
References and cumulative updates use FP32, or FP64 for FP64 parameters, and
retain that precision across checkpoint loading. This requires two additional
parameter-sized buffers. store_delta=False saves the cumulative-update
buffer by reconstructing it from model weights and the previous scale; that
approximation is less accurate at tiny scales and with BF16/FP16 models.
The base optimizer retains its own parameter-update and moment precision.
No loss report or train/eval parameter swapping is needed. Parameters without
gradients remain unchanged; their reference is captured on their first active
step. Gradient-free and skip-step-base rejected steps leave the tuner unchanged.
When wrapping a SkipStepOptimizer, report loss/norm on optim.base_optimizer;
a supplied closure automatically reports its loss there. Dense real tensors,
DDP, and DTensors are supported. Schedule-Free, LBFGS, and differentiable bases
are rejected because their update semantics do not compose with this wrapper.
Mechanic also supports weight_decay_mode="hyperball" on either named
optimizer or a Hyperball-capable base. For example:
optim = prejac.MechanicMuon([
{"params": transformer_matrix_params, "weight_decay_mode": "hyperball"},
{"params": other_params, "weight_decay_mode": "none"},
], lr=1.0)
This experimental composition retains Hyperball's normalized base updates,
then projects the final Mechanic query reference + scale * delta to the
matrix's saved radius. The scale gradient differentiates through that final
projection, so radial gradients cannot drive the tuner. The tuner's radial
regularization contributes only for unconstrained parameters. Base weight
decay is ignored in Hyperball groups, Muon's matrix shape adjustment is
bypassed, and non-matrices use the base optimizer's no-decay fallback with
ordinary Mechanic rescaling. Full local matrices/DDP are supported; Hyperball
matrix DTensors remain unsupported.
Both store_delta modes work. The memory-saving mode records two scalars per
constrained matrix (the unprojected query norm and its applied scale) to recover
the radial information removed by projection. It retains the usual rounding
and small-scale reconstruction limitations. Radii and projection metadata
survive checkpoint loading. This projected variant is an extension of
Mechanic's original unconstrained Algorithm 1.
Schedule-Free AdamW
import prejac
optim = prejac.ScheduleFreeAdamW(model.parameters(), lr=0.0025)
for step in range(total_steps):
optim.set_train_mode()
loss = model(inputs).loss
loss.backward()
optim.step() # no loss value is required
model.zero_grad()
optim.set_eval_mode()
ScheduleFree+ AdamW (learning-rate-free)
import prejac
import torch
model = MyModel()
optim = prejac.ScheduleFreePlusAdamW(model.parameters())
optim.set_total_steps(total_steps)
max_grad_norm = 1.0
for step in range(total_steps):
optim.set_train_mode() # set query point y into model params
loss = model(inputs).loss
loss.backward()
optim.latest_loss = loss # must be set before step()
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
optim.latest_gradient_scale = min(
1.0, max_grad_norm / (float(total_norm) + 1e-6)
)
optim.step()
model.zero_grad()
optim.set_eval_mode() # set averaged iterate x into model params
# model now holds the averaged parameters
set_train_mode()/set_eval_mode()swap model parameters between the schedule-free query point and the averaged iterate. They do not callnn.Module.train()/nn.Module.eval(). Call both if you need batch-norm, dropout, etc. to switch behaviour too.
Both Schedule-Free variants always keep their base iterate (z), averaged
iterate (x), and Adam moments in FP32. Gradients are promoted to FP32 for
optimizer calculations, and the query point is blended in FP32 before being
copied into the model's parameter dtype. BF16/FP16 model parameters remain
supported; forward and backward computation can still use mixed precision.
The dtype argument accepts only None or torch.float32.
Checkpoint loading preserves FP32 optimizer state even when the model uses BF16/FP16 parameters. Older low-precision optimizer state is promoted to FP32; updates already lost to rounding in an old checkpoint cannot be recovered.
Regular Muon (with learning rate)
import prejac
model = MyModel()
matrix_params = [p for p in model.parameters() if p.ndim == 2]
other_params = [p for p in model.parameters() if p.ndim != 2]
optim = prejac.Muon(
[
{"params": matrix_params, "algorithm": "muon"},
{"params": other_params, "algorithm": "adamw"},
],
lr=0.02,
weight_decay=0.1, # decoupled weight decay on all params (Moonlight)
weight_decay_mode="adamw", # or "adamc"
)
for step in range(total_steps):
loss = model(inputs).loss
loss.backward()
optim.step()
model.zero_grad()
PoLoRA
PoLoRA needs to know which LoRA
factors form each product. Register pairs as (A, B), with A shaped
(rank, in) and B shaped (out, rank). The paper assumes LoRA alpha=rank
so the adapter output uses B @ A without another scale factor.
import prejac
# For a PEFT model with one or more named adapters:
pairs = [
(module.lora_A[name].weight, module.lora_B[name].weight)
for module in model.modules()
if hasattr(module, "lora_A") and hasattr(module, "lora_B")
for name in module.lora_A
if name in module.lora_B and module.lora_A[name].weight.requires_grad
]
optim = prejac.PoLoRA(pairs, lr=2e-4)
for batch in loader:
loss = model(batch).loss
loss.backward()
optim.step()
optim.zero_grad()
If the LoRA adapter output is scaled by s (for example, s = alpha / rank),
pass that value as scale=s. PoLoRA divides both factor gradients by s and
uses lr / (abs(s) * (||A|| + ||B||)) for the factor-update magnitude, making
the step invariant to a uniform adapter-output scale. scale defaults to 1.
It can also be set per parameter group.
For different learning rates, pass groups such as
[{"pairs": [(A1, B1)], "lr": 1e-4}, {"pairs": [(A2, B2)], "lr": 2e-4}].
The pair order must be the same when loading an optimizer checkpoint. PoLoRA
only manages the registered LoRA factors; use a separate optimizer for any
other trainable parameters. A pair is skipped when both gradients are absent.
The factors must be full local matrices; sharded DTensor factors are unsupported.
With the default settings, compatible same-shaped pairs in a group are
automatically processed in batches to reduce kernel-launch and small-matrix
overhead.
SkipStepMuon (with spike detection)
import prejac
model = MyModel()
matrix_params = [p for p in model.parameters() if p.ndim == 2]
other_params = [p for p in model.parameters() if p.ndim != 2]
optim = prejac.SkipStepMuon(
[
{"params": matrix_params, "algorithm": "muon"},
{"params": other_params, "algorithm": "adamw"},
],
lr=0.02,
weight_decay=0.1,
weight_decay_mode="adamw", # or "adamc"
)
for step in range(total_steps):
loss = model(inputs).loss
loss.backward()
optim.latest_loss = loss # required for spike detection
optim.latest_grad_norm = ... # optional, tighter spike detection
optim.step()
model.zero_grad()
Muon
Muon is an optimizer designed primarily for neural network matrix parameters, using momentum and Newton–Schulz orthogonalization to produce its updates.
Our default Muon variant is largely based on the Moonlight implementation of Muon, scaling the update to AdamW's RMS magnitude and introducing weight decay to the update rule.
Weight decay
weight_decay_mode="hyperball" selects the norm-constrained wrapper from
Hyperball (arXiv:2606.16899v1), Equation (1).
It is available in AdamW, Lion, Muon, SOAP, Mousse, and their skip-step variants.
MechanicAdamW and MechanicMuon support the projected composition described
in the Mechanic quick start above.
For each selected matrix it stores the Frobenius norm at the first gradient-bearing
step as a fixed radius R, normalizes the base optimizer direction u, and applies
W <- R * normalize(W - lr * R * normalize(u)). Radii are saved in optimizer
checkpoints. This mode replaces decay entirely: weight_decay is ignored, even
when zero, and lr controls relative update length. Muon's shape-dependent LR
adjustment is bypassed because the direction is normalized.
Select attention and MLP matrices explicitly; embeddings and output heads can also be matrices and cannot be identified automatically. For example:
optimizer = Muon([
{"params": transformer_matrix_params, "weight_decay_mode": "hyperball"},
{"params": other_params, "algorithm": "adamw", "weight_decay_mode": "none"},
], lr=0.02)
Non-matrix parameters in Hyperball groups use the base optimizer without decay.
Zero directions, zero initial radii, and trial points at the origin leave weights
unchanged. Norms and projections use at least FP32 precision. Hyperball requires
full local matrices (ordinary tensors/DDP); DTensors are unsupported. Schedule-free
variants retain their existing rules. The functional step helpers accept
hyperball_radius (a float or scalar tensor, or a list for foreach) to hold the radius fixed
across calls; without it, stateless helpers use the current norm. hyperball_update
is also exported for direct use with a positive base update direction.
AdamW, SkipStepAdamW, Lion, and regular Muon variants accept weight_decay_mode="adamw"
(the default) or weight_decay_mode="adamc". The default applies AdamW-style
decoupled weight decay with one power of the learning rate. AdamC uses the
square of the parameter-update step size. For Muon matrices, that is the
shape-adjusted step lr * lr_adjust_ratio; Adam-routed Muon parameters use
the group learning rate. The default Muon rule uses the raw group learning
rate, following Moonlight.
The mode can also be overridden per parameter group:
optim = prejac.Muon(
[
{"params": matrix_params, "weight_decay_mode": "adamc"},
{"params": other_params, "weight_decay_mode": "adamw"},
],
lr=0.02,
weight_decay=0.1,
)
Because AdamC uses the squared step size, its weight_decay coefficient is
not numerically interchangeable with an AdamW coefficient and should be tuned
separately.
Schedule-Free variants default to weight_decay_mode="adamc", following the
ScheduleFree+ implementation. They also accept
"adamw", "none", and "hyperball", including per-group overrides. Conventional
decay updates z using the query point y: z -= alpha**2 * weight_decay * y
for AdamC, or z -= alpha * weight_decay * y for AdamW. Here alpha is the
warmup-adjusted fixed learning rate for SF, or the Polyak-sized step for SF+.
Schedule-Free with Hyperball
Both ScheduleFreeAdamW and ScheduleFreePlusAdamW support an experimental
Hyperball base optimizer. Select attention/MLP
weight matrices explicitly; leave embeddings, tied output weights, normalization
gains, and biases in conventional groups. Hyperball groups must contain only
2D matrices and use weight_decay=0, since projection replaces decay.
optim = prejac.ScheduleFreePlusAdamW(
[
{"params": matrix_params, "weight_decay_mode": "hyperball", "weight_decay": 0},
{"params": other_params, "weight_decay_mode": "adamc", "weight_decay": 5},
],
)
# ScheduleFreeAdamW accepts the same groups, with its fixed lr argument.
For each selected matrix, the radius R is its FP32 Frobenius norm when its
optimizer state is first initialized, and is saved as hyperball_radius in
its checkpoint state. Initialization must have a finite nonzero norm. With
Adam direction u, an accepted step is:
z_trial = z - alpha * R * u / ||u||_F
z_next = R * z_trial / ||z_trial||_F
x_next = (1 - c) * x + c * z_next
y_next = beta * x_next + (1 - beta) * z_next
Only z is projected. x remains the weighted average of the projected base
iterates, and y remains the usual SF query blend; their norms can be below
R. Projecting either blend separately would change the SF identities. When
sf_beta=0, the training query is the projected base iterate. A zero update
or zero step size leaves z unchanged; if the trial point is exactly zero,
the previous z is retained because its projection has no defined direction.
All norms include the full matrix across DTensor shards; state and calculations
remain FP32 even with BF16/FP16 model parameters.
SF+ scales each Hyperball matrix's contribution to its Polyak denominator by
R / ||u||_F:
sqrt(pi/2) * sum_i(scale_i * ||g_i||_1), with scale_i=1 for conventional
groups. It keeps the loss + beta * <g, z-x> numerator, denominator EMA,
warmup, gradient-clipping correction, and max_eta cap. Previewed Adam moments
are committed only for accepted steps. This is a normalized adaptation of the
existing L1 metric approximation, not a published Hyperball/SF+ algorithm or
convergence guarantee. For SF, lr sets the relative unprojected step length
on selected matrices; for SF+, that length is polyak_alpha.
Orthogonalization methods
Muon uses Newton-Schulz iteration to orthogonalize momentum. Two methods:
| Method | Description |
|---|---|
gram (default) |
Stabilized Gram Newton-Schulz. Iterates on the small n x n Gram matrix XX^T instead of the tall n x m matrix. Fewer FLOPs for rectangular matrices. Restarts at iteration 2 (0-indexed) for numerical stability. Dao AI Lab, 2026 |
standard |
Standard Newton-Schulz. The original algorithm. Iterates on the full n x m matrix. |
For square matrices both give similar results. For rectangular matrices (most LLM weights), gram is
theoretically faster because it operates on the smaller, square Gram matrix.
The bundled implementations use FP16 internally, restore the input dtype on
return, and accept one to five steps from the published Polar Express sequence.
When the optional gram-newton-schulz package is available on CUDA, batched
same-shaped matrices also use its optimized GNS kernels.
optim = prejac.Muon(params, ns_method="gram") # default
optim = prejac.Muon(params, ns_method="standard")
Safety/Stability features
-
max_eta: Optional nonnegative upper bound on the Polyak step size forScheduleFreePlusAdamW. It is unset by default (no upper bound). -
Step skipping: All non-vanilla optimizers (i.e. all optimizers but
AdamW,Lion, andMuon) support skipping steps when the grad norm or loss spike beyond the rolling mean of the grad norm/loss of the previous steps.
API footguns (Thanks LLMs -_-)
-
latest_lossMUST be set for eachstep()on ScheduleFree+ optimizers. A closure passed tostep(closure)supplies it automatically. Without a new loss,step()returns without updating parameters. Fixed-LR and regular optimizers (ScheduleFreeAdamW,SkipStepAdamW,SkipStepMuon, etc.) will step without it, but spike detection uses whatever value was last assigned tolatest_loss(or is disabled if it was never set). -
On ScheduleFree+ optimizers, if gradients are uniformly scaled between
backward()andstep()(for example by global-norm clipping), setlatest_gradient_scaleto that multiplier. It is consumed by the next step and resets to one. A Polyak ratio requires its loss and gradient statistics to have the same scale; omitting the clip coefficient makes clipping shrink only the denominator and can therefore increase the update instead of limiting it. -
lrmeans different things for the two schedule-free families. ForScheduleFree*, it is the fixed peak learning rate. ForScheduleFreePlus*, it scales the online Polyak step size. -
set_total_steps(n)is required whenanneal_steps=-1(anneal over the full training run). Call it before training begins or the optimizer will throw an exception. -
A rejected loss or gradient-norm spike is excluded from future rolling spike thresholds.
-
For Muon,
torch.DTensors are gathered before Newton-Schulz so the polar map is computed for the global matrix, then redistributed to the original placements. ScheduleFree+ Polyak statistics are reduced over the global tensors. Rank-local minibatch losses are averaged across the DTensor mesh before the global Polyak step and spike decision.distributed_meshcan be supplied to verify that parameters use the intended mesh.
References
- ScheduleFree+ (Defazio, 2026)
- Why Gradients Rapidly Increase Near the End of Training (Defazio, 2026) -- AdamC/corrected decoupled weight decay
- The Road Less Scheduled (Defazio et al., 2024)
- Moonlight / Muon is Scalable for LLM Training (Liu et al., 2025) -- weight decay and LR adjustment for Muon
- Polar Express (Amsel et al., 2025) -- Newton-Schulz coefficients used in our implementation
- Gram Newton-Schulz (Zhang et al., 2026) -- stabilized GNS algorithm
License
Apache-2.0
Mousse
optim = prejac.Mousse(
model.parameters(),
lr=0.02,
preconditioner_beta=0.95,
precondition_frequency=10,
preconditioner_epsilon=1e-5,
curvature_exponent=0.125,
precondition_side="both", # "left" or "right" uses only one factor
ns_method="reference", # authors' BF16 NS5; "gram"/"standard" are alternatives
)
Two-dimensional parameters use Mousse; all other ranks use AdamW. Put
embeddings and output heads in explicit algorithm="adamw" groups to exclude
them from curvature preconditioning, or manage them with a separate optimizer.
Mousse cannot identify these modules from tensor shape alone.
The implementation follows the authors' reference code
at revision d00c1bf17790fbe56424ee5567cce80d8e75f4b2: trace normalization,
spectral shifting after damping, scaling before and after NS, and norm grafting.
Factors refresh at updates 1, 11, 21, ... with the default frequency. Counters
are per parameter; missing gradients leave its state unchanged. Zero gradients
are handled without the reference implementation's division by zero.
Covariances and eigenbases use FP32 (FP64 for FP64 state). FP32 factor paths
round momentum directions and final updates to BF16, matching upstream.
dtype=torch.float32 optionally keeps momentum and AdamW fallback state in FP32
with low-precision model parameters. Checkpoint loading preserves state precision.
curvature_exponent=0 bypasses rotations; with ns_method="gram" it reduces to
this library's Muon direction. Defaults for LR and shape scaling follow prejac's
Muon; spectral damping 1e-5 follows the paper's experimental settings.
Mousse defaults to the fast approximate route: factor_method="qr", one
power/QR iteration per refresh, and preconditioner_precision="auto". Auto
precision uses BF16 products with FP32 outputs and covariance state on CUDA;
CPU/MPS and FP64 state retain full-precision products. QR uses full eigh at
initialization, then tracks the previous basis with Rayleigh eigenvalue estimates.
More qr_iterations trade refresh cost for approximation accuracy.
Compatible CUDA matrices are batched in groups of eight, with up to eight
concurrent factor workers. CPU/MPS use single-matrix execution. batch_size=1
and factor_workers=1 disable batching and concurrent solves respectively.
Persistent state shares batch allocations; AdamW fallback uses multi-tensor kernels.
For the reference numerical path, select explicitly:
optim = prejac.Mousse(
groups,
factor_method="eigh",
preconditioner_precision="float32",
)
QR and BF16 change numerical updates. In two 1,000-update WikiText runs on a 149M model, the fast route had 0.26–0.78% higher final perplexity than full-eigh Mousse. This is early convergence evidence, not a guarantee of equivalent full pretraining. A40 full-step overhead versus Muon fell from 96% at 4K tokens/update to 2.6% at 262K tokens/update; optimizer work still costs more at small batches. Existing checkpoints retain their saved settings; checkpoints predating these options load with full-eigh and full-precision products.
Ordinary PyTorch DDP synchronizes gradients before optim.step(); Mousse keeps
replicated optimizer state and requires no optimizer process group. It explicitly
rejects DTensors, including replicated DTensors, at construction, when adding
parameters, and before stepping. Distributed matrix scheduling, skip-step and
schedule-free Mousse variants are not implemented.
Official parity tests run the authors' actual optimizer in its working two-rank
DDP path, without rewriting its update equations. See
tests/test_mousse_official.py for the pinned
checkout and test command. The unused Triton import is stubbed for CPU tests,
and compilation decorators run eagerly. Unpatched upstream right-only mode raises
KeyError('L'); tests record that limitation explicitly. A
standalone upstream fix patch also repairs single-process
and zero-gradient behavior. Set MOUSSE_REFERENCE_PATCHED=1 to run parity for
all modes against the patched checkout, including padded DDP batches and
single-process updates.
For DDP integration and real DTensor rejection checks:
torchrun --standalone --nproc-per-node=2 tests/mousse_ddp_integration.py