flashinfer.fused_moe.alphamoe_nvfp4_aligned_moe¶
- flashinfer.fused_moe.alphamoe_nvfp4_aligned_moe(hidden_states: Tensor, hidden_states_scale: Tensor, gemm1_weights: Tensor, gemm1_weights_scale: Tensor, gemm2_weights: Tensor, gemm2_weights_scale: Tensor, output1_scale_gate_scalar: Tensor, output1_scale_scalar: Tensor, output2_scale_scalar: Tensor, sorted_token_ids: Tensor, expert_ids: Tensor, num_tokens_post_padded: Tensor, topk_weights: Tensor, out: Tensor, top_k: int, block_m: int = 8, routed_scaling_factor: float = 1.0, w1_scale_prepared: Tensor | None = None, w2_scale_prepared: Tensor | None = None) None¶
Run AlphaMoE NVFP4 gate/up, SwiGLU, requantization and down compute.
This SM100/SM103 operator consumes a pre-aligned routing plan. Contributions accumulate in a temporary FP32
[M, K]buffer initialized from the caller-owned BF16outtensor, then convert back tooutonce. It does not run expert selection. The temporary uses4 * M * Kbytes and is compatible with CUDA graph capture.hidden_states,gemm1_weights, andgemm2_weightsstore two E2M1 values peruint8byte, with the even logical value in the low nibble. Their E4M3 scales are linear, contiguous, and cover 16 logical values each:hidden_states:[M, K / 2]; scale[M, K / 16]gemm1_weights:[E, N, K / 2]in conventional[gate; up]row order; scale[E, N, K / 16]gemm2_weights:[E, K, N / 4]; scale[E, K, N / 32]
The scale tensors must use
torch.float8_e4m3fnand the linear per-16 layout above. FlashInfer’s 128x4-swizzled NVFP4 scale layout is a different contract and must not be passed to this kernel.sorted_token_idsandexpert_idsfollow the aligned MoE plan used by vLLM/SGLang.num_tokens_post_paddedis a one-element device tensor naming the valid plan extent; blocks in the capacity-sized launch grid that lie past this extent are skipped. The caller must keep that device value no larger thanexpert_ids.numel() * block_mand provide valid expert and token ids in the active plan.- Parameters:
hidden_states (torch.Tensor) – Packed E2M1 activations
[M, K / 2]. The innermost stride must be 1; row-strided views require a positive row stride of at leastK / 2that is divisible by 16, and a 16-byte-aligned data pointer.hidden_states_scale (torch.Tensor) – Linear E4M3 scales
[M, K / 16].gemm1_weights (torch.Tensor) – Packed gate/up weights
[E, N, K / 2]withNdivisible by 256.gemm1_weights_scale (torch.Tensor) – Linear E4M3 scales
[E, N, K / 16].gemm2_weights (torch.Tensor) – Packed down weights
[E, K, N / 4].gemm2_weights_scale (torch.Tensor) – Linear E4M3 scales
[E, K, N / 32].output1_scale_gate_scalar (torch.Tensor) – Contiguous FP32 per-expert gate dequantization scales
[E]. Each gate accumulator is multiplied by its routed expert’s value before applying SiLU.output1_scale_scalar (torch.Tensor) – Contiguous FP32 per-expert up-projection scales
[E]. For static ModelOpt FP4 this is the up-projection global scale divided by the second-activation quantization scale.output2_scale_scalar (torch.Tensor) – Contiguous FP32 per-expert down-projection scales
[E]. Each down accumulator is multiplied by its routed expert’s value before route weighting androuted_scaling_factor.sorted_token_ids (torch.Tensor) – Contiguous int32 aligned-plan entries.
expert_ids (torch.Tensor) – Contiguous int32 expert id per
block_mplan entries.num_tokens_post_padded (torch.Tensor) – One-element device int32 valid plan extent.
topk_weights (torch.Tensor) – FP32 route weights
[M, top_k].out (torch.Tensor) – Contiguous BF16 output
[M, K]. Contributions are added to its existing values in FP32 before the final BF16 conversion; its data pointer must be 16-byte aligned. Zero it before calling when a fresh result is wanted. It must not overlap any input tensor.top_k (int) – Routes per token.
block_m (int) – Routing plan block size, at least 8 and divisible by 8.
routed_scaling_factor (float) – Finite scalar applied to each routed contribution.
w1_scale_prepared (Optional[torch.Tensor]) – Immutable uint8 panels prepared once from the raw scale tensors with prepare_nvfp4_w1_scales/prepare_nvfp4_w2_scales before requests or graph capture. Raw tensors remain required.
w2_scale_prepared (Optional[torch.Tensor]) – Immutable uint8 panels prepared once from the raw scale tensors with prepare_nvfp4_w1_scales/prepare_nvfp4_w2_scales before requests or graph capture. Raw tensors remain required.
Notes
Logical
Kis derived as2 * hidden_states.shape[1]and must be at least 256 and divisible by 256. This function mutatesoutand returnsNone. The FP32 bulk reduction is order-dependent and flushes subnormal inputs and results to signed zero, as specified by PTXcp.reduce.async.bulk.add.f32.