flashinfer.fused_moe¶
This module provides fused Mixture-of-Experts (MoE) operations optimized for different backends and data types.
Types and Enums¶
|
An enumeration. |
|
An enumeration. |
Shared activation helpers live in flashinfer.tllm_enums and are used by
both the TRT-LLM and CuteDSL MoE paths.
|
Return whether the given activation type is a gated activation (e.g. SwiGLU family). |
Utility Functions¶
|
Reshape a 2-D tensor into a 3-D block layout. |
Reorder rows of a weight tensor for the TensorRT-LLM gated-activation GEMM layout. |
|
Interleave 4-bit packed MoE weights for the SM90 mixed-input GEMM. |
|
Fold weight scales for the SM90 mixed-input MoE GEMM. |
|
Prepare MXFP4 weights for the SM90 Humming-style FP8 activation path. |
|
|
Fused expert routing with top-k selection for DeepSeek-V3. |
|
Hash-based MoE expert routing for DeepSeek-V4. |
The E8M0 range-clamping, residual-scale factorization, and FP4 payload-rewrite
scheme used by preprocess_moe_weights_for_sm90_mixed_gemm_humming is adapted
from Humming.
Multi-LoRA MoE (BGMV)¶
Batched Gather-Matrix-Vector kernels for serving multiple LoRA adapters on top of a Mixture-of-Experts layer (shrink + expand).
|
High-level multi-LoRA MoE BGMV: shrink + expand in one call. |
|
MoE LoRA shrink operation: project input through LoRA-A matrices. |
|
MoE LoRA expand operation: project through LoRA-B matrices. |
|
FC1 (gate_up_proj) LoRA delta for a routed MoE, in the layout consumed by |
|
FC2 (down_proj) LoRA delta for a routed MoE, to be ADDED to the MoE output. |
CUTLASS Fused MoE¶
|
Compute a Mixture of Experts (MoE) layer using CUTLASS backend. |
TensorRT-LLM Fused MoE¶
|
BF16 MoE operation with autotuning support. |
|
Pre-routed BF16 MoE operation with autotuning support. |
|
FP4 block-scaled MoE operation. |
|
FP4 block scale MoE operation with pre-computed routing. |
|
FP8 block-scaled MoE operation. |
|
Pre-routed FP8 block-scaled MoE operation. |
|
FP8 per-tensor-scale MoE operation. |
Pre-routed FP8 per-tensor-scale MoE operation. |
|
|
MXINT4 block-scaled MoE operation. |
|
MxInt4 block-scale MoE with pre-computed routing. |
Standalone TRT-LLM Gen Routing¶
The routing stage the TRT-LLM Gen fused MoE launchers run before their GEMMs, exposed on its own so expert selection and the permutation/padding bookkeeping can be used (and tested) independently of quantization and GEMM configuration.
|
Standalone trtllm-gen MoE routing (expert selection + permutation). |
|
Outputs of the trtllm-gen MoE routing stage. |
CuteDSL Fused MoE¶
The CuteDSL backends are conditionally available when the
nvidia-cutlass-dsl package is installed.
|
Run a fused MoE forward pass using the CuTe-DSL NVFP4 kernels. |
|
Run fused MoE with MXFP8 activations and packed MXFP4 weights. |
|
Run fused MoE on SM120/SM121 using b12x CuTe-DSL kernels. |
- class flashinfer.fused_moe.CuteDslMoEWrapper(num_experts: int, top_k: int, hidden_size: int, intermediate_size: int, use_cuda_graph: bool = False, max_num_tokens: int | None = None, num_local_experts: int | None = None, local_expert_offset: int = 0, tile_size: int = 128, sf_vec_size: int = 16, output_dtype: dtype = torch.bfloat16, device: str = 'cuda', enable_pdl: bool = True, activation_type: int = 3, swiglu_alpha: float = 1.0, swiglu_beta: float = 0.0, swiglu_limit: float = 3.4028234663852886e+38, situ_beta: float | None = None, situ_linear_beta: float | None = None, use_fused_finalize: bool = True, quant_mode: str = 'w4a4')¶
Bases:
objectWrapper class for CuteDSL MoE with CUDA graph and auto-tuning support.
With use_cuda_graph=True, the wrapper creates persistent CUDA stream and event resources outside graph capture, enabling async-memset / GEMM1 overlap during capture and replay. Auto-tuning is supported via the tactic parameter or autotune() context.
Supported architectures: SM100, SM103.
- num_experts¶
Total number of experts.
- top_k¶
Number of experts per token.
Hidden dimension size.
- intermediate_size¶
Intermediate dimension size.
- use_cuda_graph¶
Whether the wrapper holds persistent stream/event resources for CUDA graph capture.
- use_fused_finalize¶
Use atomic fused finalize; otherwise use the deterministic two-stage finalize.
- quant_mode¶
Selected W4A4 or W4A16 compute mode.
- max_num_tokens¶
Deprecated; accepted for backwards compatibility but ignored.
- Example (CUDA Graph):
>>> moe = CuteDslMoEWrapper( ... num_experts=256, top_k=8, ... hidden_size=7168, intermediate_size=2048, ... use_cuda_graph=True, ... ) >>> # Warmup >>> for _ in range(3): ... output = moe.run(x, x_sf, topk_ids, topk_weights, w1, w1_sf, ...) >>> # Capture >>> g = torch.cuda.CUDAGraph() >>> with torch.cuda.graph(g): ... output = moe.run(x, x_sf, topk_ids, topk_weights, w1, w1_sf, ...) >>> # Replay >>> g.replay()
- Example (Auto-tuning):
>>> moe = CuteDslMoEWrapper(num_experts=256, top_k=8, ...) >>> # Run with auto-tuning >>> with autotune(True): ... output = moe.run(x, x_sf, topk_ids, topk_weights, w1, w1_sf, ...)
- __init__(num_experts: int, top_k: int, hidden_size: int, intermediate_size: int, use_cuda_graph: bool = False, max_num_tokens: int | None = None, num_local_experts: int | None = None, local_expert_offset: int = 0, tile_size: int = 128, sf_vec_size: int = 16, output_dtype: dtype = torch.bfloat16, device: str = 'cuda', enable_pdl: bool = True, activation_type: int = 3, swiglu_alpha: float = 1.0, swiglu_beta: float = 0.0, swiglu_limit: float = 3.4028234663852886e+38, situ_beta: float | None = None, situ_linear_beta: float | None = None, use_fused_finalize: bool = True, quant_mode: str = 'w4a4')¶
Configure the CuTe-DSL NVFP4 fused-MoE wrapper.
- Parameters:
num_experts (int) – Total number of experts.
top_k (int) – Number of experts routed to per token.
hidden_size (int) – Hidden dimension size.
intermediate_size (int) – Intermediate dimension size after the fused activation.
use_cuda_graph (bool) – Create persistent CUDA stream/events for W4A4 async-memset overlap. W4A16 is CUDA-graph safe without those resources. Defaults to
False.max_num_tokens (Optional[int]) – Deprecated; accepted for backwards compatibility but ignored.
num_local_experts (Optional[int]) – Local experts for expert parallelism. Defaults to
num_experts.local_expert_offset (int) – Offset of local experts in the global expert space. Defaults to
0.tile_size (int) – Tile size for
moe_sort. Defaults to128.sf_vec_size (int) – Scale-factor vector size. Defaults to
16.output_dtype (torch.dtype) – Output dtype. Defaults to
torch.bfloat16.device (str) – Device on which to allocate buffers. Defaults to
"cuda".enable_pdl (bool) – Enable Programmatic Dependent Launch. Defaults to
True.activation_type (int) – FC1 activation type. Use
ActivationType.Swiglufor gated SwiGLU/SiTU,ActivationType.GegluTanhfor tanh-approximate GeGLU, andActivationType.Relu2for non-gated ReLU^2. Settingsitu_betaselects SiTU.swiglu_alpha (float) – SwiGLU parameters.
swiglu_oaiis represented asActivationType.Swigluwith non-default values.swiglu_beta (float) – SwiGLU parameters.
swiglu_oaiis represented asActivationType.Swigluwith non-default values.swiglu_limit (float) – SwiGLU parameters.
swiglu_oaiis represented asActivationType.Swigluwith non-default values.situ_beta (Optional[float]) – When set with
ActivationType.Swiglu, use the SiTU gatebeta * tanh(gate / beta) * sigmoid(gate).situ_linear_beta (Optional[float]) – Optional SiTU tanh clamp for the up branch.
use_fused_finalize (bool) – Use atomic fused finalize; otherwise use the deterministic two-stage finalize. Defaults to
True.quant_mode (str) – Compute mode:
"w4a4"/"nvfp4"or"w4a16". Defaults to"w4a4".
- get_valid_tactics() list¶
Return list of valid tactics for this MoE configuration.
- run(x: Tensor, x_sf: Tensor | None, token_selected_experts: Tensor, token_final_scales: Tensor, w1_weight: Tensor, w1_weight_sf: Tensor, w1_alpha: Tensor, fc2_input_scale: Tensor | None, w2_weight: Tensor, w2_weight_sf: Tensor, w2_alpha: Tensor, tactic: Tuple | None = None, *, per_token_scale: Tensor | None = None) Tensor¶
Run the CuTe-DSL NVFP4 fused-MoE forward pass.
CUDA-graph safe when the wrapper was constructed with
use_cuda_graph=True. Supports auto-tuning via thetacticargument or the surroundingautotune()context manager.- Parameters:
x (torch.Tensor) – Packed NVFP4 input for
quant_mode="w4a4"or BF16 input forquant_mode="w4a16".x_sf (Optional[torch.Tensor]) – Scale factors for
quant_mode="w4a4"; must beNoneforquant_mode="w4a16".token_selected_experts (torch.Tensor) – Expert assignments of shape
[num_tokens, top_k].token_final_scales (torch.Tensor) – Routing weights of shape
[num_tokens, top_k].w1_weight (torch.Tensor) – GEMM1 weights (gate + up fused for gated activations, or a single projection for non-gated activations).
w1_weight_sf (torch.Tensor) – Scale factors for
w1_weight.w1_alpha (torch.Tensor) – Per-expert global scale for GEMM1.
fc2_input_scale (Optional[torch.Tensor]) – Global scale for W4A4 GEMM2 input quantization; must be
Nonefor W4A16 because GEMM1 output stays in BF16.w2_weight (torch.Tensor) – GEMM2 weights (down projection).
w2_weight_sf (torch.Tensor) – Scale factors for
w2_weight.w2_alpha (torch.Tensor) – Per-expert global scale for GEMM2.
tactic (Optional[Tuple]) – Tactic tuple, or
Nonefor auto-selection via the runtime tuner.per_token_scale (Optional[torch.Tensor]) – Optional W4A4 per-token input row scale for GEMM1.
- Returns:
Output tensor of shape
[num_tokens, hidden_size].- Return type:
torch.Tensor
- class flashinfer.fused_moe.CuteDslMxfp8Mxfp4MoEWrapper(num_experts: int, top_k: int, hidden_size: int, intermediate_size: int, max_num_tokens: int | None = None, num_local_experts: int | None = None, local_expert_offset: int = 0, use_cuda_graph: bool = False, device: str = 'cuda', enable_pdl: bool = True, activation_type: int = 3, swiglu_alpha: float = 1.0, swiglu_beta: float = 0.0, swiglu_limit: float = 3.4028234663852886e+38)¶
Bases:
objectProduction wrapper for the MXFP8 x MXFP4 fused-MoE pipeline.
With
use_cuda_graph=Truethe wrapper holds persistent CUDA stream and event resources, created outside graph capture so they can be reused inside it. Workspace itself is not pre-allocated: graph capture records allocations made during capture in its private pool, so pre-sizing buffers for a maximum batch buys nothing but memory.Because the stream and event resources are reused, one wrapper instance is not reentrant or safe for concurrent calls. The first
runbinds the instance to that call’s CUDA stream; create one wrapper per stream.- __init__(num_experts: int, top_k: int, hidden_size: int, intermediate_size: int, max_num_tokens: int | None = None, num_local_experts: int | None = None, local_expert_offset: int = 0, use_cuda_graph: bool = False, device: str = 'cuda', enable_pdl: bool = True, activation_type: int = 3, swiglu_alpha: float = 1.0, swiglu_beta: float = 0.0, swiglu_limit: float = 3.4028234663852886e+38) None¶
Initialize a reusable mixed-precision fused-MoE runner.
- Parameters:
num_experts (int) – Global expert count and experts selected per token.
top_k (int) – Global expert count and experts selected per token.
hidden_size (int) – Model hidden size and per-expert intermediate size.
intermediate_size (int) – Model hidden size and per-expert intermediate size.
max_num_tokens (int, optional) – Deprecated compatibility argument; accepted but ignored.
num_local_experts (int, optional) – Experts resident on this rank; defaults to
num_experts.local_expert_offset (int) – Global index of the first local expert.
use_cuda_graph (bool) – Create persistent stream and event resources for graph capture.
device (str) – CUDA device used for persistent resources.
enable_pdl (bool) – Enable programmatic dependent launch in the generated kernels.
activation_type (int) – Fused activation identifier.
swiglu_alpha (float) – SwiGLU activation parameters.
swiglu_beta (float) – SwiGLU activation parameters.
swiglu_limit (float) – SwiGLU activation parameters.
- run(x: Tensor, x_sf: Tensor, token_selected_experts: Tensor, token_final_scales: Tensor, w1_weight: Tensor, w1_weight_sf: Tensor, w1_alpha: Tensor, w2_weight: Tensor, w2_weight_sf: Tensor, w2_alpha: Tensor, tactic: Tuple[Any, ...] | None = None) Tensor¶
Execute one mixed-precision fused-MoE forward pass.
x_sfis the linear[M, H / 32]E8M0 byte layout. Weight scales must already be in the MMA layout returned byconvert_sf_to_mma_layout(..., sf_vec_size=32).The returned tensor is freshly allocated and owned by the caller.
- Parameters:
x (torch.Tensor) – MXFP8 activations and their linear block-32 E8M0 scales.
x_sf (torch.Tensor) – MXFP8 activations and their linear block-32 E8M0 scales.
token_selected_experts (torch.Tensor) – Per-token expert indices and routing scales.
token_final_scales (torch.Tensor) – Per-token expert indices and routing scales.
w1_weight (torch.Tensor) – Packed MXFP4 expert weights.
w2_weight (torch.Tensor) – Packed MXFP4 expert weights.
w1_weight_sf (torch.Tensor) – Block-32 weight scales in MMA layout.
w2_weight_sf (torch.Tensor) – Block-32 weight scales in MMA layout.
w1_alpha (torch.Tensor) – Per-expert dequantization multipliers.
w2_alpha (torch.Tensor) – Per-expert dequantization multipliers.
tactic (tuple, optional) – Explicit kernel tactic; the autotuner selects one when omitted.
- class flashinfer.fused_moe.B12xMoEWrapper(num_experts: int, top_k: int, hidden_size: int, intermediate_size: int, *, use_cuda_graph: bool = False, max_num_tokens: int = 4096, num_local_experts: int | None = None, output_dtype: dtype = torch.bfloat16, device: str = 'cuda', activation: str = 'silu', swiglu_alpha: float = 1.702, swiglu_beta: float = 1.0, swiglu_limit: float | None = None, activation_precision: str = 'fp4', quant_mode: str | None = None, source_format: str = 'modelopt')¶
Bases:
objectB12x fused MoE wrapper for SM120/SM121 with CUDA graph support.
Pre-allocates workspace buffers for CUDA graph compatibility. Automatically selects micro/static/dynamic backend per call.
- Parameters:
num_experts – Total number of experts.
top_k – Number of experts per token.
hidden_size – Hidden dimension size.
intermediate_size – Intermediate size.
use_cuda_graph – Pre-allocate buffers for CUDA graph compatibility.
max_num_tokens – Maximum tokens (only for use_cuda_graph=True).
num_local_experts – Local experts for EP. Default: num_experts.
output_dtype – Output data type. Only torch.bfloat16 is currently supported. Default: torch.bfloat16.
device – Device for buffer allocation. Default: “cuda”.
activation – Activation — “silu”, “gelu_tanh”, “swigluoai_uninterleave”, or “relu2”. Default: “silu”. swiglu_alpha/beta/limit apply to swigluoai.
activation_precision – Backward-compatible alias for quant_mode. “fp4” selects quant_mode=”nvfp4”; “bf16” selects quant_mode=”w4a16”.
quant_mode – Quantization mode, “nvfp4”/”w4a4”, “mxfp4”, or “w4a16”. When set, this selects the backend and internal workspace family.
source_format – Source weight format for quant_mode=”w4a16”. Supports “modelopt” and “compressed_tensors”. Default: “modelopt”.
Example
>>> moe = B12xMoEWrapper(num_experts=256, top_k=8, ...) >>> output = moe.run(x=hidden_states_bf16, ...)
- __init__(num_experts: int, top_k: int, hidden_size: int, intermediate_size: int, *, use_cuda_graph: bool = False, max_num_tokens: int = 4096, num_local_experts: int | None = None, output_dtype: dtype = torch.bfloat16, device: str = 'cuda', activation: str = 'silu', swiglu_alpha: float = 1.702, swiglu_beta: float = 1.0, swiglu_limit: float | None = None, activation_precision: str = 'fp4', quant_mode: str | None = None, source_format: str = 'modelopt')¶
Configure the b12x fused-MoE wrapper.
- Parameters:
num_experts (int) – Total number of experts.
top_k (int) – Number of experts routed to per token.
hidden_size (int) – Hidden dimension size.
intermediate_size (int) – Intermediate dimension size.
use_cuda_graph (bool) – If
True, pre-allocate workspace buffers sized formax_num_tokensso the wrapper can be captured into a CUDA graph. Defaults toFalse.max_num_tokens (int) – Maximum batch size, only used when
use_cuda_graph=True. Defaults to4096.num_local_experts (Optional[int]) – Number of local experts for expert parallelism. Defaults to
num_experts.output_dtype (torch.dtype) – Output dtype. Only
torch.bfloat16is currently supported.device (str) – Device on which to allocate workspace buffers. Defaults to
"cuda".activation (str) – Activation function —
"silu"(gated SwiGLU),"gelu_tanh"(gated GeGLU, tanh-approx GELU),"swigluoai_uninterleave"(gated SwiGLU-OAI) or"relu2"(non-gated). Defaults to"silu".swiglu_alpha (float) – SwiGLU-OAI parameters (only for
"swigluoai_uninterleave"):gate*sigmoid(alpha*gate)*(up+beta)with optional clamp toswiglu_limit(Nonedisables). Defaults 1.702 / 1.0 / None.swiglu_beta (float) – SwiGLU-OAI parameters (only for
"swigluoai_uninterleave"):gate*sigmoid(alpha*gate)*(up+beta)with optional clamp toswiglu_limit(Nonedisables). Defaults 1.702 / 1.0 / None.swiglu_limit (float) – SwiGLU-OAI parameters (only for
"swigluoai_uninterleave"):gate*sigmoid(alpha*gate)*(up+beta)with optional clamp toswiglu_limit(Nonedisables). Defaults 1.702 / 1.0 / None.activation_precision (str) – Backward-compatible alias for
quant_mode."fp4"selectsquant_mode="nvfp4";"bf16"selectsquant_mode="w4a16".quant_mode (Optional[str]) – Quantization mode,
"nvfp4"/"w4a4","mxfp4", or"w4a16".source_format (str) – Source weight format for
quant_mode="w4a16"—"modelopt"(default) or"compressed_tensors".
- run(x: Tensor, w1_weight: Tensor, w1_weight_sf: Tensor, w2_weight: Tensor, w2_weight_sf: Tensor, token_selected_experts: Tensor, token_final_scales: Tensor, *, w1_alpha: Tensor, w2_alpha: Tensor, fc2_input_scale: Tensor | None = None, input_global_scale: Tensor | None = None) Tensor¶
Run the b12x fused-MoE forward pass.
- Parameters:
x (torch.Tensor) – Input activations of shape
[num_tokens, hidden_size],bfloat16.w1_weight (torch.Tensor) – FC1 weights, FP4-packed.
w1_weight_sf (torch.Tensor) – Scale factors for
w1_weight.w2_weight (torch.Tensor) – FC2 weights, FP4-packed.
w2_weight_sf (torch.Tensor) – Scale factors for
w2_weight.token_selected_experts (torch.Tensor) – Expert assignments of shape
[num_tokens, top_k].token_final_scales (torch.Tensor) – Routing weights of shape
[num_tokens, top_k].w1_alpha (torch.Tensor) – Per-expert global scale for FC1.
w2_alpha (torch.Tensor) – Per-expert global scale for FC2.
fc2_input_scale (Optional[torch.Tensor]) – Global scale for FC2 input quantization. Required for
quant_mode="nvfp4"; accepted but ignored for"w4a16".input_global_scale (Optional[torch.Tensor]) – Global scale for FC1 input quantization, scalar or
[num_experts]. Defaults tow1_alpha; seeb12x_fused_moe(). Ignored for"w4a16".
- Returns:
Output tensor of shape
[num_tokens, hidden_size].- Return type:
torch.Tensor
MonoMoE (Single-Kernel Block-FP8, SM90a)¶
Single-kernel top-K Mixture-of-Experts implementation specialized for the
Qwen3.5-35B block-FP8 shape on Hopper (SM90a). The full pipeline — routing,
up-projection, SiLU, down-projection and reduction — runs inside one kernel
launch. Use has_monomoe() to check availability before calling.
Return True if the monomoe CUDA extension can be built and loaded. |
|
Return the global scratchpad size (bytes) required by the kernel. |
|
|
Allocate a zero-initialized scratchpad on |
|
Repack fp8 up-projection weights for the Pair_Layout WGMMA A-tile. |
|
Single-kernel block-FP8 top-K MoE (fixed E256/N512/K2048 shape, SM90a). |