fp4-train / README.md
phanerozoic's picture
Card: standardized form with hero
75066be verified
|
Raw
History Blame
5.7 kB
metadata
library_name: kernels
license: apache-2.0

fp4-train

fp4-train is an NVFP4 quantizer with stochastic rounding and error feedback, the primitives for stable 4-bit training, loadable through kernels. The reference baseline is an independent NVFP4 encoding implementation, matched to zero error, with fp32 training as the convergence reference.

Training with weights stored in 4-bit fails under ordinary rounding: once the per-step update is smaller than a quantization step, round-to-nearest discards every update and the model stops learning. This kernel quantizes to NVFP4 (e2m1 values with an e4m3 scale per 16-element block) with stochastic rounding, so quantization is unbiased and sub-step updates accumulate in expectation, and it carries a per-element error-feedback residual so the dropped fraction of every update is banked instead of lost. The packed output is the layout Blackwell FP4 tensor cores consume.

Four training curves: round-to-nearest flatlines while stochastic rounding and error feedback descend toward the fp32 reference, above an e2m1 grid where a value hops between neighbors

Weights stored in NVFP4 and updated in place, 500 steps of sub-step updates: round-to-nearest never leaves its starting loss (6.56); stochastic rounding reaches 1.80, error feedback 0.77, against the fp32 reference at 0.02. Below, quantizing 2.3 on the e2m1 grid hops between 2 and 3 with the running mean at 2.26.

Usage

import torch
from kernels import get_kernel

fp4 = get_kernel("phanerozoic/fp4-train", version=1, trust_remote_code=True)

# NVFP4-stored-weight training: the weight lives in 4-bit, updated in place
ef = fp4.ErrorFeedback(stochastic=False)                      # RTN + error feedback
W = W0.clone()
for step in range(num_steps):
    grad = compute_grad(W)
    W = ef.qdq("layer0.weight", W - lr * grad, step=step)

# quantization primitives
packed, scale = fp4.quantize(W, stochastic=True, seed=0, step=global_step)
Wq = fp4.dequantize(packed, scale, *W.shape)

# native NVFP4 tensor-core matmul on Blackwell (a @ b^T, K a multiple of 16)
y = fp4.fp4_mm(a, b)

version selects the release branch; trust_remote_code is required by kernels for publishers without the trusted-publisher mark.

API

Symbol Purpose
quantize(x, stochastic, residual, seed, step) NVFP4 quantize; error feedback if residual is a [M,K] f32 tensor (updated in place)
dequantize(packed, scale, M, K) NVFP4 -> bf16
qdq(x, stochastic, residual, seed, step) quantize then dequantize (the value the tensor cores use)
fp4_linear(x, weight, stochastic, seed, step) quantized-operand linear (straight-through gradient)
fp4_matmul(a, b, stochastic, seed, step) a @ b^T FP4-operand product, dequant + bf16 accumulate (portable, cc 8.0+)
fp4_mm(a, b, stochastic, seed, step) a @ b^T on the Blackwell NVFP4 tensor cores (native _scaled_mm, ~10x faster)
to_fp4_operands(x, stochastic, seed, step) packed fp4 + e4m3 scale as the torch FP4 dtypes for _scaled_mm
ErrorFeedback per-tensor residual state for weight quantization

Method

Each 16-element block gets a scale of block_absmax / 6 encoded to e4m3, so the largest element maps near the e2m1 maximum. Every element is divided by the decoded scale and encoded to the e2m1 grid {0, .5, 1, 1.5, 2, 3, 4, 6} with a sign bit. Rounding is stochastic: for a value between two grid points, the higher is chosen with probability equal to the fractional distance, using a counter-based (seed, step, global-index) RNG, so the expected quantized value equals the input. With error feedback, the residual x - dequant(quantize(x)) is carried and added into the next quantization, so the time-averaged value equals the true value even under round-to-nearest. Codes pack two per byte; block scales are e4m3, the layout NVFP4 tensor cores read.

Correctness

Measured on L4, H200, and RTX PRO 6000 through get_kernel:

  • Exact encoding: quantize/dequantize matches an independent reference to 0 (max abs diff = 0).
  • Stochastic rounding is unbiased: for 2.3 between grid points 2 and 3, the SR mean is 2.300 where round-to-nearest is biased to 2.000.
  • Error feedback cancels the drift: RTN plus error feedback time-averages to 2.297.
  • With NVFP4-stored weights and sub-step updates, RTN stalls at the starting loss while SR and EF recover near-full-precision loss: a linear model reaches MSE 1.0 (EF) and 13 (SR) against bf16 0.0 with RTN at 77; a small MLP reaches 0.159 (EF) and 0.171 (SR) against bf16 0.158 with RTN at 1.03.
  • On an RTX PRO 6000 the hardware fp4_mm path returns a bit-identical result to the dequantize path (relative error 0) and is about 10-14x faster, and 1.5-1.7x faster than a bf16 matmul.

Requirements and limits

  • NVIDIA GPU with compute capability 8.0+ (e4m3 scale via software conversion); fp4_mm requires Blackwell.
  • K a multiple of 16; float32 or bfloat16 input.
  • The rounding modes matter when the low precision is in the training state itself (weights stored and updated in NVFP4). With a full-precision master weight, the master carries the sub-ULP information and the choice of rounding is second order.
  • Blocks whose absmax falls below the e4m3 scale range quantize to zero; keep stored-weight scales inside the representable range.

References

NVFP4 and OCP microscaling (MXFP) formats; Gupta et al., "Deep Learning with Limited Numerical Precision" (stochastic rounding, 2015); Seide et al., "1-bit SGD" (error feedback, 2014).

License

Apache-2.0.