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 BF16 out tensor, then convert back to out once. It does not run expert selection. The temporary uses 4 * M * K bytes and is compatible with CUDA graph capture.

hidden_states, gemm1_weights, and gemm2_weights store two E2M1 values per uint8 byte, 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_e4m3fn and 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_ids and expert_ids follow the aligned MoE plan used by vLLM/SGLang. num_tokens_post_padded is 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 than expert_ids.numel() * block_m and 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 least K / 2 that 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] with N divisible 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 and routed_scaling_factor.

  • sorted_token_ids (torch.Tensor) – Contiguous int32 aligned-plan entries.

  • expert_ids (torch.Tensor) – Contiguous int32 expert id per block_m plan 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 K is derived as 2 * hidden_states.shape[1] and must be at least 256 and divisible by 256. This function mutates out and returns None. The FP32 bulk reduction is order-dependent and flushes subnormal inputs and results to signed zero, as specified by PTX cp.reduce.async.bulk.add.f32.