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_size contract: flattened token * top_k + route positions grouped by expert, padded per expert to a multiple of block_m with the sentinel num_tokens * top_k. expert_ids carries one expert index per block_m-sized block, and num_tokens_post_padded is a device-side scalar naming the valid plan extent; blocks past it are skipped, so worst-case-sized plan buffers are fine. sorted_token_ids must still contain at least expert_ids.numel() * block_m entries.

Parameters:
  • hidden_states (torch.Tensor) – (num_tokens, hidden_size) float8_e4m3fn activations, quantized per token in groups of 128 along hidden_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_e4m3fn gate+up weights in the interleaved device layout produced by alphamoe_interleave_gated_weights(). 2 * intermediate_size must be a multiple of 256 and hidden_size a 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_e4m3fn down 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_size plan positions (see above).

  • expert_ids (torch.Tensor) – int32, one expert index per block_m plan 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.