nemo_automodel.components.quantization.mxfp4
nemo_automodel.components.quantization.mxfp4
MXFP4 (fp4 e2m1 + e8m0 block scales) pack/unpack utilities for MoE expert weights.
The packed layout matches the DeepSeek V4 Flash routed-expert checkpoint format:
two e2m1 values per int8 byte (low nibble at even column index, high nibble at the
following odd column) with one float8_e8m0fnu scale per 32 contiguous columns.
MXFP4GroupedMM provides a grouped GEMM over packed weights that re-dequantizes
in backward instead of saving the dequantized tensor, so frozen expert weights stay
packed at steady state during LoRA training.
Module Contents
Classes
Functions
Data
API
Bases: Function
Grouped GEMM over mxfp4-packed frozen weights with dequantization on the fly.
Saves only the packed weights for backward and re-dequantizes there, so the
bf16 weight tensor is a transient in both passes instead of being kept alive
by autograd. Weights are stored as [E, N, K] packed along K, which is
the natural dequantization output and the operand the backward GEMM needs
directly (grad_x = grad_out @ W). The forward needs [E, K, N], which
torch._grouped_mm consumes as a transposed view (cuBLAS transB) — no
contiguous copy required.
No weight gradient is produced — the base weights are frozen under LoRA.
Backpropagate through activations while keeping the packed base frozen.
Parameters:
Tensor of shape [tokens, out_dim], in the forward output dtype.
Returns: torch.Tensor
Activation gradient of shape [tokens, in_dim], followed by None for
Multiply activations grouped by expert by packed frozen weights.
Parameters:
Tensor of shape [tokens, in_dim], grouped contiguously by local expert.
Int8 tensor of shape [local_experts, out_dim, in_dim // 2].
E8M0 tensor of shape [local_experts, out_dim, in_dim // 32].
Int32 tensor of shape [local_experts], holding cumulative token counts. All tensors must be on the same device.
Returns: torch.Tensor
Tensor of shape [tokens, out_dim] with the activation dtype.
Unpack fp4 e2m1 packed-int8 values and apply the per-32-column e8m0 scale.
Parameters:
int8 tensor of shape [..., K // 2] holding two e2m1 values per byte.
float8_e8m0fnu tensor of shape [..., K // 32].
Output dtype.
Returns: torch.Tensor
Dequantized tensor of shape [..., K] in dtype.
Quantize along the last dim to the packed mxfp4 layout used by dequantize_mxfp4.
Block scales are computed as 2^(floor(log2(amax)) - 2) so that values that are
already exactly representable (e.g. a dequantized fp4 checkpoint) round-trip
value-exactly.
Parameters:
Floating-point tensor of shape [..., K] with K divisible by 32.
Returns: tuple[torch.Tensor, torch.Tensor]
Tuple of (int8 packed tensor [..., K // 2], float8_e8m0fnu scales [..., K // 32]).