Training to converge early, with as little shown as possible. - Optimizers for deep learning.
Find a file
Repository files (latest commit first)
Filename Latest commit message Latest commit date
2026-10-09 23:23:14 -04:00
examples modularize weight decay 2026-09-14 19:17:58 -04:00
patches Add Mousse optimizer with DDP support and official parity tests 2026-10-02 09:22:01 -04:00
src/prejac Support projected Hyperball updates in Mechanic optimizers 2026-10-09 23:16:46 -04:00
tests Support projected Hyperball updates in Mechanic optimizers 2026-10-09 23:16:46 -04:00
.gitignore Reset harder 2026-09-19 14:31:19 -04:00
LICENSE Reset harder 2026-09-19 14:31:19 -04:00
NOTICE.md Add Mousse optimizer with DDP support and official parity tests 2026-10-02 09:22:01 -04:00
pyproject.toml Add Mousse optimizer with DDP support and official parity tests 2026-10-02 09:22:01 -04:00
README.md Support projected Hyperball updates in Mechanic optimizers 2026-10-09 23:16:46 -04:00
uv.lock Refresh lock after lowering PyTorch minimum 2026-09-24 09:27:23 -04:00

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 call nn.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 for ScheduleFreePlusAdamW. It is unset by default (no upper bound).

  • Step skipping: All non-vanilla optimizers (i.e. all optimizers but AdamW, Lion, and Muon) 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_loss MUST be set for each step() on ScheduleFree+ optimizers. A closure passed to step(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 to latest_loss (or is disabled if it was never set).

  • On ScheduleFree+ optimizers, if gradients are uniformly scaled between backward() and step() (for example by global-norm clipping), set latest_gradient_scale to 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.

  • lr means different things for the two schedule-free families. For ScheduleFree*, it is the fixed peak learning rate. For ScheduleFreePlus*, it scales the online Polyak step size.

  • set_total_steps(n) is required when anneal_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_mesh can be supplied to verify that parameters use the intended mesh.

References

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