Skip to content

Latest commit

Β 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 

Repository files navigation

πŸ”₯ FORGE

Fused On-Register Gradient Elimination for Memory-Efficient LLM Training

The weight gradient is an artifact of how autograd is staged β€” not something learning needs. FORGE removes it.

arXiv PDF License: Apache 2.0

πŸ“„ Paper Β· FORGE: Fused On-Register Gradient Elimination for Memory-Efficient LLM Training

Python 3.10+ PyTorch 2.1+ Triton 3.4+ PRs Welcome

Quick start Β· Results Β· Convergence Β· Hardware Β· Distributed Β· How it works Β· Cite


Standard training computes grad_W = grad_output.T @ input as a full tensor in HBM, then runs optimizer.step() to read it back. For an 8B model the live gradients alone cost β‰ˆ15 GB β€” and at the seam between backward and the optimizer step, every layer's gradient is live at once, setting the memory ceiling of training.

FORGE fuses the optimizer step into the backward pass and applies it one tile at a time, entirely in GPU registers. Each weight-gradient tile is produced, consumed by the optimizer, and discarded before the next tile is computed. The full grad_W tensor never exists in HBM.

FORGE tile-by-tile fused backward + optimizer
For each weight tile, the gradient is accumulated in registers and the AdamW update is applied immediately β€” then the tile is dropped.

✨ Highlights

  • Deletes the gradient pool β€” on Llama-3.1-8B, peak memory falls from 62.0 GB under vanilla AdamW to 48.4 GB at matched state precision, and to 35.3 GB with int8 moments.
  • Faster, not just smaller β€” the update is folded into the weight-gradient GEMM, so the separate optimizer.step() pass disappears: 110.2 ms/step vs. 134.3 (fused AdamW) and 167.1 (vanilla) β€” 1.52Γ— faster than vanilla AdamW, and 2.2–2.6Γ— faster than bitsandbytes at matched int8 state bytes.
  • Provably exact β€” for any optimizer that updates a weight from its own gradient alone (AdamW, SGDΒ±momentum, Lion, RMSprop, Adagrad, NAdam, RAdam, …), the fused step produces exactly the standard result: bit-identical to a reference that accumulates the token axis in the same order.
  • Architecture-agnostic β€” one kernel, no architecture-specific code, trains GPT-2, Llama-3.1-8B, five Qwen3 sizes, vision transformers to 25B, Mamba-2 to 20B, and MLP-Mixers to 4.9B; thirteen optimizer families run end to end.
  • Converts to capability β€” with fp8 moments FORGE trains Qwen3-32B on a single H200 (134.4 GB) where fused AdamW runs out of memory; under FSDP on an 8-GPU node it reaches the lowest per-rank memory of any method that trains the model.

FORGE ships as the importable package fused_grad_optimizer.

πŸ“Š Key results (Llama-3.1-8B, single H200)

Peak memory collapse and memory-vs-speed on Llama-3.1-8B (H200)
The weight gradient (red) collapses under FORGE; its two arms differ only in moment precision. FORGE is 1.52Γ— faster than vanilla AdamW and 2.04Γ— faster than bitsandbytes 8-bit, at lower memory.

Single-GPU comparison on H200 (141 GB), batch 1, sequence 512, BF16 everywhere; step time is the median of 20 steps.

Method Peak (GB) Step (ms) TF/s
vanilla AdamW 62.04 167.1 Β± 18.0 149 Β± 17
fused AdamW 60.08 134.3 Β± 14.1 185 Β± 21
bitsandbytes 8-bit 45.36 316.3 Β± 20.6 78 Β± 5
FORGE 48.36 110.2 Β± 8.7 226 Β± 15
FORGE (int8) 35.32 155.0 Β± 4.4 159 Β± 5

FORGE is the only method that improves on fused AdamW on all three axes at once β€” the full comparison against FlashOptim, GaLore, APOLLO, optimi, and AdaLomo is in Table 1 of the paper. Standalone, the fused update reaches 74% of the measured 4,252 GB/s HBM ceiling, against 61% for fused AdamW, 24% for vanilla AdamW, and 8% for bitsandbytes.

Operating regime. What governs the saving is the token count BT = batch Γ— sequence: FORGE deletes a fixed β‰ˆ15 GB (the gradient pool), so at matched bf16 states the reduction fades from 22% at BT = 512 to nothing at BT β‰₯ 4096, where activations set the peak instead. FORGE is a small-BT method β€” the regime that dominates fine-tuning and continued pretraining. Model scale works the other way: the ratio improves with parameter count (Qwen3-14B fits in 87.8 GB vs. 110.3 for fused AdamW), up to the 32B-on-one-H200 point above.

πŸ“‰ Convergence parity

FORGE matches baseline convergence
1-epoch continued pretraining of Llama-3.1-8B (52k steps, identical hyperparameters): FORGE tracks PyTorch AdamW exactly, in bf16 and int8 states, while bitsandbytes 8-bit converges worse.
  • From scratch: GPT-2 124M on FineWeb-Edu tracks fused AdamW for 125k iterations and ends fractionally below it β€” 3.20 vs. 3.22 nats.
  • Continued pretraining: across Llama-3.1-8B and five Qwen3 sizes (20,000 steps, β‰₯ 3 seeds each), losses stay within 0.001 nats on average, 0.003 at worst.
  • Exactness, not approximation: the fused step is bit-identical to a reference that accumulates the token axis in the same order; against cuBLAS the only discrepancy is the summation order intrinsic to any GEMM. All thirteen implemented optimizer families train end to end.
GPT-2 124M pretrained from scratch: FORGE tracks fused AdamW throughout 20k-step continued pretraining on OpenMathInstruct-2: FORGE tracks fused AdamW on Qwen3-1.7B
Left: GPT-2 124M pretrained from random initialization on FineWeb-Edu β€” FORGE tracks fused AdamW throughout (3.20 vs. 3.22 nats). Right: continued pretraining on OpenMathInstruct-2 (20,000 steps, Qwen3-1.7B) β€” FORGE tracks fused AdamW, while bitsandbytes 8-bit drifts.
Qwen3-0.6B CPT parity Qwen3-4B CPT parity
Qwen3-8B CPT parity Qwen3-14B CPT parity
The same parity holds across model scale: Qwen3 0.6B, 4B, 8B, and 14B (20,000 steps each).

πŸš€ Quick start

pip install -e ".[test]"      # core + tests
# pip install -e ".[bench]"   # + transformers/accelerate for the benchmarks
import torch
from fused_grad_optimizer import FusedLinear, FusedOptimizerManager

model = YourModel().cuda()

# 1. Swap nn.Linear layers for FusedLinear
for name, module in model.named_modules():
    for child_name, child in list(module.named_children()):
        if isinstance(child, torch.nn.Linear):
            setattr(module, child_name,
                    FusedLinear.from_linear(child, optimizer_type="adamw"))

# 2. A manager coordinates the fused layers; a standard optimizer handles the rest
manager = FusedOptimizerManager(model)
optimizer = torch.optim.AdamW(manager.get_non_fused_params(), lr=1e-4, fused=True)

# 3. Train β€” fused layers update their weights DURING backward
for step, batch in enumerate(dataloader):
    manager.pre_step(lr=get_lr(step))
    loss = model(**batch).loss
    loss.backward()      # FORGE applies the optimizer here, tile-by-tile
    optimizer.step()     # only norms / embeddings (~0.1% of params)
    optimizer.zero_grad()

See examples/quickstart.py for a runnable toy example.

🧠 How it works

For each weight tile, FORGE accumulates grad_output.T @ input in fp32 registers via a loop over the token dimension, then applies the optimizer immediately β€” so the full grad_W is never written to HBM. A standard bf16 step streams sixteen bytes per parameter through HBM; FORGE moves twelve, and moves them closer to peak bandwidth. The trade-off is read amplification: activations are re-read once per weight tile. Autotuned tile sizes, a zero-cost virtual transpose, native bf16 tensor cores, and grouped tile ordering for L2 reuse keep that cost small β€” and it buys the elimination of the entire optimizer step.

The update is applied after the input gradient Ξ”X = Ξ”YΒ·W is read, so the chain rule is preserved. Weights with more than one gradient consumer in a step (tied embeddings) are left on the standard optimizer.

πŸ–₯️ Hardware support & results per GPU

Validated on NVIDIA datacenter / workstation GPUs via Triton, across the Qwen3 family and Llama-3.1-8B at sequence 512–4096:

GPU Arch Measured on this card
H200 141 GB SM90 Headline single-GPU results; Hopper TMA path (kernel.py); 8Γ—H200 NVLink
H100 SXM 80 GB SM90 Qwen3 family sweeps; the budget where baselines start to OOM
B200 180 GB SM100 Llama + Qwen3 sweeps; CUDA 12.8, Triton 3.6, FlashAttention-4
RTX PRO 6000 Blackwell 96 GB SM120 Thirteen-optimizer sweep; 8-GPU PCIe distributed node (below)
A100 40/80 GB, B300 SM80 / SM100 Cross-platform capability study

Peak memory is shape-deterministic and reproduces across cards to within rounding; step time is per-platform and is only ever compared within one card and one recipe. The full per-card grids β€” including the RTX PRO 6000 optimizer sweep, where peak falls 27–54% and step time 28–71% across every family β€” are in the paper's supplementary appendices (G: extended single-GPU grids, I: optimizer families, O: cross-platform capability).

Requires CUDA + Triton β‰₯ 3.4. The default path (kernel.py) is pure Triton and needs no extra setup. The arch-specific research kernels (hopper_* / cutlass_*) additionally JIT-compile against NVIDIA CUTLASS β€” set CUTLASS_PATH or clone it into the repo root as cutlass/. AMD/Apple backends are not yet validated.

πŸ—‚οΈ Repository layout

src/fused_grad_optimizer/   # the library
  kernel.py                 # core fused grad+optimizer Triton kernels (autotuned)
  autograd.py               # custom autograd.Function fusing backward + optimizer
  module.py                 # FusedLinear (nn.Module) + FusedOptimizerManager
  state.py                  # OptimizerConfig + lazy m/v state
  hopper_*/cutlass_*        # arch-specific kernels (H200 TMA, B200 EVT)
tests/                      # correctness: SGD/AdamW, bf16, int8, manager
examples/                   # runnable quickstart
assets/                     # figures

πŸ“ Citation

@article{kukreja2026forge,
  title   = {FORGE: Fused On-Register Gradient Elimination for Memory-Efficient LLM Training},
  author  = {Kukreja, Dikshant and Prasad, Kritarth and Anand, Avinash and Wang, Zhengkui
             and Cambria, Erik and Liu, Timothy and Ng, Aik Beng and See, Simon and Chatterjee, Bapi},
  journal = {arXiv preprint arXiv:2606.22932},
  year    = {2026}
}

πŸ“„ License

Apache License 2.0 β€” see LICENSE.

About

No description, website, or topics provided.

Resources

Contributing

Stars

12 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages