adamw4bit provides PyTorch implementations of ZIP-SR and ZE-EDEN for
4-bit AdamW moment storage. Both use NF4 for the first moment. ZIP-SR uses
zero-inclusive Dyn4 with stochastic rounding in preconditioner space for the
second moment; ZE-EDEN uses zero-exclusive Dyn4 with EDEN scale calibration.
The methods are described in Rounding in Preconditioner Space: Redesigning
4-bit AdamW Optimizer-State Quantization.
Requires Python 3.11 or 3.12 and PyTorch 2.13 or newer.
pip install . # the optimizer
pip install '.[all]' # plus the Hugging Face dependencies used by the examplesUse ZEEDENAdamW4Bit or ZIPSRAdamW4Bit like any PyTorch optimizer:
import torch
from adamw4bit import ZEEDENAdamW4Bit, ZIPSRAdamW4Bit
model = torch.nn.Linear(64, 64)
optimizer = ZEEDENAdamW4Bit(
model.parameters(),
lr=1e-3,
betas=(0.9, 0.95),
weight_decay=0.1,
)For ZIP-SR, use ZIPSRAdamW4Bit(model.parameters(), lr=1e-3) instead.
Both constructors select the method's first- and second-moment quantizers and
EDEN setting. Learning rate, betas, epsilon and weight decay remain configurable.
ZIPSRAdamW4Bit and ZEEDENAdamW4Bit are thin subclasses of QuantizedAdamW, the
shared implementation for 4-bit, 8-bit and FP32 moment storage. For ablations,
use QuantizedAdamW with explicit m1_quant_scheme, m2_quant_scheme and
use_eden_m2 settings. Its defaults use 8-bit linear quantization;
AdamW8bit remains a compatibility wrapper with its original defaults.
Four-bit blocks default to 128 values; tensors with fewer than 4,096 values or sizes not divisible by the block size remain in FP32. This quantizes optimizer moments, not model weights or gradients.
The paper's complete recipes also switch only the LM-head first moment to NF4 stochastic rounding during the final 10% of training. The training loop must identify the head and apply this schedule before each optimizer update:
boundary = round(0.9 * total_optimizer_steps)
head = model.get_output_embeddings()
scheme = "nf4_sr" if completed_optimizer_steps >= boundary else "nf4"
optimizer.set_m1_quant_scheme_for_parameters(head.parameters(), scheme)For 6,179 updates, SR begins at update 5,562. Tied embeddings share the same parameter, so use untied weights for a head-only switch. The examples below keep first-moment RTN throughout and demonstrate the quantization settings; they do not implement the full paper training protocol.
The reference implementation remains the default. Enable the bounded-workspace backend explicitly:
optimizer = ZEEDENAdamW4Bit(model.parameters(), optimized=True)
# ZIPSRAdamW4Bit and QuantizedAdamW accept the same flag.For supported contiguous tensors, the backend updates packed moments in place. New optimized instances use a default working chunk of 512 Ki elements (524,288), rounded down as needed to preserve quantization blocks and packed-code boundaries. Persistent moments, parameters, and gradients still scale with model size.
CPU uses the shared eager chunk update. For FP32 CUDA parameters, torch.compile
compiles that same update, including moment decoding, AdamW arithmetic, and
encoding. Stochastic rounding keeps the optimizer's existing generators. For
non-EDEN FP32 CUDA updates where only the second moment uses stochastic rounding,
its draws are prepared eagerly before the compiled chunk. This removes a graph
break. Other configurations retain their existing draw schedule to limit peak
memory. Non-EDEN reductions use deterministic kernel selection to limit startup
autotuning memory. Other supported parameter dtypes use the eager chunk update.
Compilation is lazy. The first call for a new configuration can take substantially longer and temporarily use more GPU memory than warmed updates. Measure cold and warmed peaks and timings separately; warmed optimizer measurements do not describe the complete training peak.
Eight-bit schemes, noncontiguous tensors, telemetry, recorded diagnostics,
optimizer-level update clipping, unsupported research read variants, overlapping
parameter/gradient buffers, and incompatible or overlapping moment storage use
the reference update. Those fallbacks can require full-tensor workspace.
Ordinary gradient clipping before optimizer.step() is unaffected.
This backend remains experimental and can be slower than the reference. Chunking and compiled arithmetic can change floating-point results and stochastic rounding samples. Validate training quality, peak memory, and speed on your workload; optimizer memory savings may not reduce a training peak dominated by activations.
Resume with the same optimized setting and quantization seed, and use matching
settings across distributed replicas. Optimized checkpoints retain their saved
chunk limit and fallback policy, including older 1,048,576-element limits. Loading
a checkpoint without a saved chunk limit with optimized=True adopts the current
512 Ki-element default. The constructor's optimized flag must still be selected
when recreating the optimizer. Reapply any per-parameter first-moment scheme
overrides with set_m1_quant_scheme_for_parameters; these overrides are not saved
in the optimizer checkpoint. With quant_rng_seed=None, also save and restore
the global PyTorch RNG state as part of the training checkpoint.
Save optimizer state alongside model state. New optimizer checkpoints load with PyTorch's restricted loader:
torch.save(optimizer.state_dict(), "optimizer.pt")
optimizer.load_state_dict(torch.load("optimizer.pt", weights_only=True))For a trusted older adamw4bit checkpoint containing QuantState objects, use
a scoped allowlist, then save again to use the current format:
from adamw4bit.quantization import QuantState
with torch.serialization.safe_globals([QuantState]):
state = torch.load("optimizer.pt", weights_only=True)
optimizer.load_state_dict(state)
torch.save(optimizer.state_dict(), "optimizer.pt")ZE-EDEN starts with an exactly zero second moment, matching the paper's
initialization. This corrects the historical implementation's approximately
3.25e-15 initial value on eligible tensors and can change training trajectories.
For the six smaller pretraining sizes, FP32/TorchAO use the final scheduled
validation while ZE-EDEN/ZIP-SR use terminal validation, so their evaluation
steps differ slightly.
The corrected Qwen3-8B SFT runs use separate LM-head/non-head optimizer groups
under ZeRO-2; comparisons to earlier one-group runs retain that topology difference.
Run from the repository root after pip install '.[all]'. Each script accepts
--recipe (ze-eden or zip-sr; the GPT-small script also accepts fp32).
-
examples/quickstart.pytrains a tiny GPT-2 on synthetic tokens and does not download anything:python examples/quickstart.py --recipe zip-sr
-
examples/train_gpt_small.pytrains GPT-small on FineWeb-Edu with the Hugging Face Trainer. Use--smoke-testfor a reduced run, or launch the full run on eight GPUs:python examples/train_gpt_small.py --recipe ze-eden --smoke-test torchrun --standalone --nproc-per-node=8 \ examples/train_gpt_small.py --recipe ze-eden --seed 42
-
examples/run_gpt_small_experiment.shruns the full paired comparison of all three recipes over three seeds. Seeexamples/README.md.
uv sync
uv run pytestSecurity fixes are provided for the
latest release, as described in SECURITY.md. Questions and
non-security bugs belong in GitHub issues. The project is maintained by
Nubank.
Licensed under the Apache License, Version 2.0. See LICENSE.