nemo_automodel.components.moe.quantized_experts
nemo_automodel.components.moe.quantized_experts
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
Functions
Data
API
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.
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.
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.
Compute the frozen local experts’ contribution before the EP reduction.
Parameters:
Tensor of shape [tokens, hidden], gathered across the EP group.
Boolean tensor of shape [tokens] selecting valid tokens.
Tensor of shape [tokens, top_k] with differentiable routing probabilities.
Integer tensor of shape [tokens, top_k] with global expert IDs.
Number of experts on this rank.
First global expert ID on this rank.
Returns: torch.Tensor
FP32 tensor of shape [tokens, hidden], before the EP reduction.
Compute packed experts with the tensor and EP contract of GroupedExperts.forward.
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.
Replace floating-point bases with packed placeholders or quantized weights.
Register packed meta parameters in checkpoint orientation [experts, out, in].
Multiply routed activations by a frozen packed base projection.
Parameters:
Tensor of shape [tokens, in_dim], grouped contiguously by local expert.
Base projection parameter name.
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.
Dequantize expert 0 of base weight name to compute layout [in, out].
Quantize materialized base projections and release their floating-point storage.
Replace one floating-point base projection with frozen packed parameters.
Parameters:
Base projection parameter name.
Int8 tensor of shape [experts, out_dim, in_dim // 2].
E8M0 tensor of shape [experts, out_dim, in_dim // 32]. DTensors must retain the base parameter’s mesh and expert-axis placement.
Reject execution modes that the packed expert computation does not implement.
Return the local shard of a DTensor, or the tensor unchanged.
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.