> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# 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

| Name                                                                             | Description                                                                   |
| -------------------------------------------------------------------------------- | ----------------------------------------------------------------------------- |
| [`MXFP4GroupedMM`](#nemo_automodel-components-quantization-mxfp4-MXFP4GroupedMM) | Grouped GEMM over mxfp4-packed frozen weights with dequantization on the fly. |

### Functions

| Name                                                                                 | Description                                                                        |
| ------------------------------------------------------------------------------------ | ---------------------------------------------------------------------------------- |
| [`dequantize_mxfp4`](#nemo_automodel-components-quantization-mxfp4-dequantize_mxfp4) | Unpack fp4 e2m1 packed-int8 values and apply the per-32-column e8m0 scale.         |
| [`quantize_mxfp4`](#nemo_automodel-components-quantization-mxfp4-quantize_mxfp4)     | Quantize along the last dim to the packed mxfp4 layout used by `dequantize_mxfp4`. |

### Data

[`MXFP4_BLOCK_SIZE`](#nemo_automodel-components-quantization-mxfp4-MXFP4_BLOCK_SIZE)

[`_FP4_BYTE_TABLE`](#nemo_automodel-components-quantization-mxfp4-_FP4_BYTE_TABLE)

[`_FP4_E2M1_MIDPOINTS`](#nemo_automodel-components-quantization-mxfp4-_FP4_E2M1_MIDPOINTS)

[`_FP4_E2M1_TABLE`](#nemo_automodel-components-quantization-mxfp4-_FP4_E2M1_TABLE)

### API

```python
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.

```python
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

```python
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.

```python
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`.

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

```python
nemo_automodel.components.quantization.mxfp4.MXFP4_BLOCK_SIZE = 32
```

```python
nemo_automodel.components.quantization.mxfp4._FP4_BYTE_TABLE = torch.stack([_FP4_E2M1_TABLE[torch.arange(256) & 15], _FP4_E2M1_TABLE[torch.aran...
```

```python
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))
```

```python
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....
```