Skip to content
nubankPublic

About

PyTorch implementations of ZIP-SR and ZE-EDEN for packed 4-bit AdamW optimizer states.

Resources

Contributing

Security policy

Stars

1 star

Watchers

1 watching

Forks

Repository files navigation

adamw4bit

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.

Installation

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 examples

Usage

Use 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.

Optional memory optimization

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.

Checkpoints

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")

Paper reproduction

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.

Examples

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.py trains a tiny GPT-2 on synthetic tokens and does not download anything:

    python examples/quickstart.py --recipe zip-sr
  • examples/train_gpt_small.py trains GPT-small on FineWeb-Edu with the Hugging Face Trainer. Use --smoke-test for 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.sh runs the full paired comparison of all three recipes over three seeds. See examples/README.md.

Development

uv sync
uv run pytest

Status

Security 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.

License

Licensed under the Apache License, Version 2.0. See LICENSE.

About

PyTorch implementations of ZIP-SR and ZE-EDEN for packed 4-bit AdamW optimizer states.

Resources

Contributing

Security policy

Stars

1 star

Watchers

1 watching

Forks

Releases

Packages

Used by

Contributors

Languages