flashinfer.fused_moe.alphamoe_fused_router

flashinfer.fused_moe.alphamoe_fused_router(logits: Tensor, *, top_k: int, block_m: int = 8, has_shared_expert: bool = False, plan: AlphaMoERoutePlan | None = None) → AlphaMoERoutePlan

Route FP32 logits and build an AlphaMoE aligned plan (SM100/SM103).

One cooperative generated kernel performs top-k selection, selected-logit softmax, per-expert counting, block padding, and grouped route scatter. With has_shared_expert=True, the last expert is forced into the final route slot and the remaining top_k - 1 experts are selected from [0, num_experts - 1).

Parameters:
  • logits (torch.Tensor) – Contiguous FP32 CUDA tensor shaped (num_tokens, num_experts). num_tokens must be positive and num_experts must be in [1, 512]. Logits are expected to be finite.

  • top_k (int) – Routes per token, in [1, min(num_experts, 16)].

  • block_m (int) – Per-expert alignment in [1, 16]. Defaults to 8.

  • has_shared_expert (bool) – Force the last expert into every token’s last route slot. This mode requires top_k >= 2.

  • plan (Optional[AlphaMoERoutePlan]) – Reusable output/workspace allocation. When omitted, a worst-case-capacity plan is allocated. Supply a plan for steady-state or CUDA-graph use and warm up the JIT module once before capture.

Returns:

topk_weights and topk_ids have shape (num_tokens, top_k). sorted_token_ids stores flattened token * top_k + route indices grouped by expert and padded with sentinel num_tokens * top_k. expert_ids stores one expert per block_m entries, and num_tokens_post_padded is a device-side int32 scalar naming the valid extent. Atomic scatter order within an expert is deliberately unspecified.

Return type:

AlphaMoERoutePlan

Notes

When routed logits are exactly equal, the selected expert set and its order are unspecified; no lower-expert-ID tie break is guaranteed. The frozen CUDA device source is compiled with --use_fast_math; its SHA256 is ec5bc689e68264a11a56a17fb10f699bc3733a521dea916b71ecda51d4227801. The operation has no dependency on an AlphaMoE compute kernel and its plan can feed either the W8A8 or NVFP4 fused up/down path.