nemo_automodel.components.quantization.mxfp4

View as Markdown

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

NameDescription
MXFP4GroupedMMGrouped GEMM over mxfp4-packed frozen weights with dequantization on the fly.

Functions

NameDescription
dequantize_mxfp4Unpack fp4 e2m1 packed-int8 values and apply the per-32-column e8m0 scale.
quantize_mxfp4Quantize along the last dim to the packed mxfp4 layout used by dequantize_mxfp4.

Data

MXFP4_BLOCK_SIZE

_FP4_BYTE_TABLE

_FP4_E2M1_MIDPOINTS

_FP4_E2M1_TABLE

API

class nemo_automodel.components.quantization.mxfp4.MXFP4GroupedMM()

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.

nemo_automodel.components.quantization.mxfp4.MXFP4GroupedMM.backward(
ctx: torch.autograd.function.FunctionCtx,
grad_out: torch.Tensor
) -> tuple[torch.Tensor, None, None, None]
staticmethod

Backpropagate through activations while keeping the packed base frozen.

Parameters:

grad_out
torch.Tensor

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

nemo_automodel.components.quantization.mxfp4.MXFP4GroupedMM.forward(
ctx: torch.autograd.function.FunctionCtx,
x: torch.Tensor,
packed: torch.Tensor,
scales: torch.Tensor,
offs: torch.Tensor
) -> torch.Tensor
staticmethod

Multiply activations grouped by expert by packed frozen weights.

Parameters:

x
torch.Tensor

Tensor of shape [tokens, in_dim], grouped contiguously by local expert.

packed
torch.Tensor

Int8 tensor of shape [local_experts, out_dim, in_dim // 2].

scales
torch.Tensor

E8M0 tensor of shape [local_experts, out_dim, in_dim // 32].

offs
torch.Tensor

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.

nemo_automodel.components.quantization.mxfp4.dequantize_mxfp4(
packed: torch.Tensor,
scales: torch.Tensor,
dtype: torch.dtype
) -> torch.Tensor

Unpack fp4 e2m1 packed-int8 values and apply the per-32-column e8m0 scale.

Parameters:

packed
torch.Tensor

int8 tensor of shape [..., K // 2] holding two e2m1 values per byte.

scales
torch.Tensor

float8_e8m0fnu tensor of shape [..., K // 32].

dtype
torch.dtype

Output dtype.

Returns: torch.Tensor

Dequantized tensor of shape [..., K] in dtype.

nemo_automodel.components.quantization.mxfp4.quantize_mxfp4(
weight: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]

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:

weight
torch.Tensor

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]).

nemo_automodel.components.quantization.mxfp4.MXFP4_BLOCK_SIZE = 32
nemo_automodel.components.quantization.mxfp4._FP4_BYTE_TABLE = torch.stack([_FP4_E2M1_TABLE[torch.arange(256) & 15], _FP4_E2M1_TABLE[torch.aran...
nemo_automodel.components.quantization.mxfp4._FP4_E2M1_MIDPOINTS = torch.tensor([0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0], dtype=(torch.float32))
nemo_automodel.components.quantization.mxfp4._FP4_E2M1_TABLE = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, 0.0, -0.5, -1.0, -1.5, -2....