nemo_automodel.components.moe.quantized_experts

View as Markdown

MXFP4-resident expert storage and model application for frozen MoE experts.

GroupedExpertsMXFP4 keeps the frozen routed-expert base weights packed as fp4-e2m1 + e8m0 block scales (the DeepSeek V4 Flash checkpoint format) and dequantizes on the fly inside the grouped GEMM, instead of holding them in bf16. This is the storage win for LoRA / frozen-base training of large MoE models, where the routed experts dominate parameter memory.

MXFP4ExpertStorageMixin shares packed storage and base GEMMs between the frozen and LoRA experts, with Torch or DeepEP token dispatch.

Module Contents

Classes

NameDescription
GroupedExpertsDeepEPMXFP4Frozen routed experts with mxfp4-resident base weights under DeepEP dispatch.
GroupedExpertsMXFP4Frozen routed experts with mxfp4-resident base weights and no adapter.
MXFP4ExpertStorageMixinPacked-mxfp4 base-weight storage and grouped GEMM for routed experts.

Functions

NameDescription
_to_localReturn the local shard of a DTensor, or the tensor unchanged.
apply_mxfp4_to_moe_expertsApply MXFP4-resident storage to the model’s common routed-expert modules.

Data

logger

API

class nemo_automodel.components.moe.quantized_experts.GroupedExpertsDeepEPMXFP4(
)

Bases: MXFP4ExpertStorageMixin, GroupedExpertsDeepEP

Frozen routed experts with mxfp4-resident base weights under DeepEP dispatch.

Drop-in replacement for GroupedExpertsDeepEP when the experts are frozen (e.g. LoRA on attention only). The DeepEP fused all-to-all token dispatch is reused unchanged — mxfp4 only changes the two post-dispatch grouped GEMMs, which read the packed base weights via MXFP4GroupedMM instead of bf16 torch._grouped_mm.

The DeepEP parent always uses native grouped MM.

nemo_automodel.components.moe.quantized_experts.GroupedExpertsDeepEPMXFP4.forward(
x: torch.Tensor,
token_mask: torch.Tensor,
weights: torch.Tensor,
indices: torch.Tensor
) -> torch.Tensor

Forward over mxfp4 base weights with DeepEP dispatch.

Preserves the tensor and EP contract of GroupedExpertsDeepEP.forward, replacing the two base torch._grouped_mm calls with MXFP4GroupedMM over the packed weights.

class nemo_automodel.components.moe.quantized_experts.GroupedExpertsMXFP4(
)

Bases: MXFP4ExpertStorageMixin, GroupedExperts

Frozen routed experts with mxfp4-resident base weights and no adapter.

Drop-in replacement for GroupedExperts when the experts are frozen (e.g. LoRA training that targets only attention). Forward mirrors GroupedExperts._forward_grouped_mm but reads the packed base weights.

use_torch_mm
= orig_module.use_torch_mm
nemo_automodel.components.moe.quantized_experts.GroupedExpertsMXFP4._forward_grouped_mm_mxfp4(
x: torch.Tensor,
token_mask: torch.Tensor,
weights: torch.Tensor,
indices: torch.Tensor,
n_local_experts: int,
experts_start_idx: int
) -> torch.Tensor

Compute the frozen local experts’ contribution before the EP reduction.

Parameters:

x
torch.Tensor

Tensor of shape [tokens, hidden], gathered across the EP group.

token_mask
torch.Tensor

Boolean tensor of shape [tokens] selecting valid tokens.

weights
torch.Tensor

Tensor of shape [tokens, top_k] with differentiable routing probabilities.

indices
torch.Tensor

Integer tensor of shape [tokens, top_k] with global expert IDs.

n_local_experts
int

Number of experts on this rank.

experts_start_idx
int

First global expert ID on this rank.

Returns: torch.Tensor

FP32 tensor of shape [tokens, hidden], before the EP reduction.

nemo_automodel.components.moe.quantized_experts.GroupedExpertsMXFP4.forward(
x: torch.Tensor,
token_mask: torch.Tensor,
weights: torch.Tensor,
indices: torch.Tensor
) -> torch.Tensor

Compute packed experts with the tensor and EP contract of GroupedExperts.forward.

class nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin()

Packed-mxfp4 base-weight storage and grouped GEMM for routed experts.

Mixed into a GroupedExperts (or GroupedExpertsLoRA) subclass. The base projections gate_and_up_projs / down_projs are stored as packed fp4 (int8, two e2m1 nibbles per byte) plus float8_e8m0fnu block scales, in checkpoint orientation [n_experts, out_dim, in_dim] so the block scales run along the contraction dim. Floating-point base parameters are dropped once packed.

Meta weights become packed placeholders for direct checkpoint loading. Materialized weights are quantized immediately during construction.

_MXFP4_BASE_NAMES
tuple[str, ...] = ('gate_and_up_projs', 'down_projs')
nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin._init_mxfp4_storage() -> None

Replace floating-point bases with packed placeholders or quantized weights.

nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin._init_packed_placeholders() -> None

Register packed meta parameters in checkpoint orientation [experts, out, in].

nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin._mxfp4_base_mm(
x: torch.Tensor,
name: str,
offs: torch.Tensor
) -> torch.Tensor

Multiply routed activations by a frozen packed base projection.

Parameters:

x
torch.Tensor

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

name
str

Base projection parameter name.

offs
torch.Tensor

Int32 tensor of shape [local_experts], holding cumulative token counts.

Returns: torch.Tensor

Tensor of shape [tokens, out_dim] with the activation dtype and device.

nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin._mxfp4_dequant_expert0(
name: str,
dtype: torch.dtype
) -> torch.Tensor

Dequantize expert 0 of base weight name to compute layout [in, out].

nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin._pack_base_weights() -> None

Quantize materialized base projections and release their floating-point storage.

nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin._register_packed_base_weight(
name: str,
packed: torch.Tensor,
scales: torch.Tensor
) -> None

Replace one floating-point base projection with frozen packed parameters.

Parameters:

name
str

Base projection parameter name.

packed
torch.Tensor

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

scales
torch.Tensor

E8M0 tensor of shape [experts, out_dim, in_dim // 32]. DTensors must retain the base parameter’s mesh and expert-axis placement.

nemo_automodel.components.moe.quantized_experts.MXFP4ExpertStorageMixin._validate_mxfp4_config() -> None

Reject execution modes that the packed expert computation does not implement.

nemo_automodel.components.moe.quantized_experts._to_local(
t
)

Return the local shard of a DTensor, or the tensor unchanged.

nemo_automodel.components.moe.quantized_experts.apply_mxfp4_to_moe_experts(
model: torch.nn.Module
) -> torch.nn.Module

Apply MXFP4-resident storage to the model’s common routed-expert modules.

Call this after LoRA injection and before distributed sharding. LoRA-targeted experts that are already MXFP4-resident are preserved; remaining plain GroupedExperts and GroupedExpertsDeepEP modules are replaced in place.

Meta experts receive packed placeholders for direct checkpoint loading, and their model state-dict adapter must opt into the MXFP4 expert storage format. Already materialized expert weights are quantized without changing the adapter’s checkpoint loading mode. Model-specific adapters own checkpoint layouts.

nemo_automodel.components.moe.quantized_experts.logger = logging.getLogger(__name__)