The weight gradient is an artifact of how autograd is staged β not something learning needs. FORGE removes it.
π Paper Β· FORGE: Fused On-Register Gradient Elimination for Memory-Efficient LLM Training
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.
For each weight tile, the gradient is accumulated in registers and the AdamW update is applied immediately β then the tile is dropped.
- 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.
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.
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.
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.
pip install -e ".[test]" # core + tests
# pip install -e ".[bench]" # + transformers/accelerate for the benchmarksimport 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.
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.
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 β setCUTLASS_PATHor clone it into the repo root ascutlass/. AMD/Apple backends are not yet validated.
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
@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}
}Apache License 2.0 β see LICENSE.



