flashinfer.fused_moe.alphamoe_fp8_block_scale_aligned_moe¶
- flashinfer.fused_moe.alphamoe_fp8_block_scale_aligned_moe(hidden_states: Tensor, hidden_states_scale: Tensor, gemm1_weights: Tensor, gemm1_weights_scale: Tensor, gemm2_weights: Tensor, gemm2_weights_scale: Tensor, sorted_token_ids: Tensor, expert_ids: Tensor, num_tokens_post_padded: Tensor, topk_weights: Tensor, *, top_k: int, block_m: int = 8, routed_scaling_factor: float = 1.0, out: Tensor | None = None) Tensor¶
Fused W8A8 block-scale MoE over a pre-aligned routing plan (SM100/SM103).
Runs the generated Alpha-MoE megakernel: routed gate/up projection, SwiGLU with per-token FP8 requantization of the intermediate, down projection, routing-weight/scaling application, and asynchronous BF16 reduce-add into
out— one kernel, no global intermediate tensor.The routing plan is the vLLM/SGLang
moe_align_block_sizecontract: flattenedtoken * top_k + routepositions grouped by expert, padded per expert to a multiple ofblock_mwith the sentinelnum_tokens * top_k.expert_idscarries one expert index perblock_m-sized block, andnum_tokens_post_paddedis a device-side scalar naming the valid plan extent; blocks past it are skipped, so worst-case-sized plan buffers are fine.sorted_token_idsmust still contain at leastexpert_ids.numel() * block_mentries.- Parameters:
hidden_states (torch.Tensor) –
(num_tokens, hidden_size)float8_e4m3fnactivations, quantized per token in groups of 128 alonghidden_size. The innermost stride must be 1; the row stride must be positive, non-overlapping, and divisible by 16 bytes; and the data pointer must be 16-byte aligned. Row-sliced activations satisfying these conditions are accepted.hidden_states_scale (torch.Tensor) –
(num_tokens, hidden_size // 128)float32 per-token-group scales.gemm1_weights (torch.Tensor) –
(num_experts, 2 * intermediate_size, hidden_size)float8_e4m3fngate+up weights in the interleaved device layout produced byalphamoe_interleave_gated_weights().2 * intermediate_sizemust be a multiple of 256 andhidden_sizea multiple of 128.gemm1_weights_scale (torch.Tensor) –
(num_experts, 2 * intermediate_size // 128, hidden_size // 128)float32 128x128 block scales, interleaved by the same helper.gemm2_weights (torch.Tensor) –
(num_experts, hidden_size, intermediate_size)float8_e4m3fndown weights (natural layout, no interleaving).gemm2_weights_scale (torch.Tensor) –
(num_experts, hidden_size // 128, intermediate_size // 128)float32 128x128 block scales.sorted_token_ids (torch.Tensor) – int32
moe_align_block_sizeplan positions (see above).expert_ids (torch.Tensor) – int32, one expert index per
block_mplan entries.num_tokens_post_padded (torch.Tensor) – int32 device tensor with one element: the valid plan extent.
topk_weights (torch.Tensor) –
(num_tokens, top_k)float32 routing weights.top_k (int) – Routes per token.
block_m (int) – Routing-plan block size; a positive multiple of 8. Defaults to 8.
routed_scaling_factor (float) – Extra scalar applied to every routed contribution. Defaults to 1.0.
out (Optional[torch.Tensor]) –
(num_tokens, hidden_size)bfloat16 accumulator. When omitted, a zeroed tensor is allocated and returned. When provided, the kernel accumulates into it (BF16 reduce-add); the caller is responsible for zeroing it (or seeding it with the intended initial values) before the call. Its data pointer must be 16-byte aligned.
- Returns:
out – The
(num_tokens, hidden_size)bfloat16 accumulator.- Return type:
torch.Tensor
Notes
Requires an SM100/SM103 (B200/B300 class) device; the kernel is a frozen generated schedule (see
csrc/alphamoe_sm100.cu), with a unit-scale guard for all-zero activation blocks.