flashinfer.fused_moe.trtllm_fp4_block_scale_routed_moe

flashinfer.fused_moe.trtllm_fp4_block_scale_routed_moe(topk_ids: Tensor | Tuple[Tensor, Tensor], routing_bias: Tensor | None, hidden_states: Tensor, hidden_states_scale: Tensor | None, gemm1_weights: Tensor, gemm1_weights_scale: Tensor, gemm1_bias: Tensor | None, gemm1_alpha: Tensor | None, gemm1_beta: Tensor | None, gemm1_clamp_limit: Tensor | None, gemm2_weights: Tensor, gemm2_weights_scale: Tensor, gemm2_bias: Tensor | None, output1_scale_scalar: Tensor | None, output1_scale_gate_scalar: Tensor | None, output2_scale_scalar: Tensor | None, num_experts: int, top_k: int, n_group: int | None, topk_group: int | None, intermediate_size: int, local_expert_offset: int, local_num_experts: int, routed_scaling_factor: float | None, routing_method_type: int = 0, do_finalize: bool = True, enable_pdl: bool | None = None, activation_type: int = 3, per_token_scale: Tensor | None = None, output: Tensor | None = None, tune_max_num_tokens: int = 8192, gemm1_lora_delta: Tensor | None = None, valid_hidden_size: int | None = None, valid_intermediate_size: int | None = None, num_fused_shared_experts: int | None = None, hidden_states_scale_layout: SfLayout | None = None) → List[Tensor]

FP4 block scale MoE operation with pre-computed routing.

This function supports two pre-computed routing formats: 1. Packed format: topk_ids is a single int32 tensor with

(expert_id << 16) | weight entries (high 16 bits = int16 expert id, low 16 bits = float16/bfloat16 weight, matching PackedScoreIdx in include/flashinfer/trtllm/fused_moe/RoutingKernel.h).

  1. Unpacked format: topk_ids is a tuple (topk_ids, topk_weights).

Parameters:
  • topk_ids (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) – Pre-computed routing decision. Either a single int32 tensor of shape [seq_len, top_k] in packed format (expert_id << 16) | weight or a tuple (ids, weights) where ids is int32 of shape [seq_len, top_k] (plain expert indices) and weights is bfloat16 or float32 of the same shape (routing weights). The weights are consumed at their native dtype (no cast), so passing the float32 weights emitted by typical routers is copy-free.

  • routing_bias (Optional[torch.Tensor]) – [num_experts] routing bias, bfloat16 or float32. May be None.

  • hidden_states (torch.Tensor) – Hidden states of shape [seq_len, hidden_size // 2] (NVFP4) or [seq_len, hidden_size] (MXFP8 / bfloat16). Supports bfloat16, MXFP8 (float8_e4m3fn), and NVFP4 (packed into uint8).

  • hidden_states_scale (Optional[torch.Tensor]) – [seq_len, hidden_size // (32 if mxfp8 else 16)] block scales of the hidden states, float8. The equivalent flat [seq_len * hidden_size // (32 if mxfp8 else 16)] buffer returned by mxfp8_quantize() with is_sf_swizzled_layout=False (the linear layout) is also accepted. Declare which layout the buffer is in with hidden_states_scale_layout; only the linear layout is supported.

  • gemm1_weights (torch.Tensor) – [num_experts, 2 * intermediate_size, hidden_size // 2] packed FP4 FC1 weights, uint8.

  • gemm1_weights_scale (torch.Tensor) – [num_experts, 2 * intermediate_size, hidden_size // (32 if mxfp4 else 16)] FC1 weight block scales, float8.

  • gemm1_bias (Optional[torch.Tensor]) – [num_experts, 2 * intermediate_size] FC1 bias, float32.

  • gemm1_alpha (Optional[torch.Tensor]) – [num_experts] swiglu alpha, float32. For SiTU this is [local_num_experts], finite and positive; None materializes per-expert alpha=1.

  • gemm1_beta (Optional[torch.Tensor]) – [num_experts] swiglu beta, float32. For SiTU this is [local_num_experts], finite and positive; None materializes per-expert beta=1.

  • gemm1_clamp_limit (Optional[torch.Tensor]) – [num_experts] swiglu clamp limit, float32. For SiTU a provided limit is per-local-expert, finite, and positive; it clamps x0 to [-limit, limit] and x1 from above.

  • gemm2_weights (torch.Tensor) – [num_experts, hidden_size, intermediate_size] packed FP4 FC2 weights, uint8.

  • gemm2_weights_scale (torch.Tensor) – [num_experts, hidden_size, intermediate_size // (32 if mxfp4 else 16)] FC2 weight block scales, float8.

  • gemm2_bias (Optional[torch.Tensor]) – [num_experts, hidden_size] FC2 bias, float32.

  • output1_scale_scalar (Optional[torch.Tensor]) – [local_num_experts] scaling factors for the first-layer activation output.

  • output1_scale_gate_scalar (Optional[torch.Tensor]) – [local_num_experts] scaling factors for the first-layer gate output.

  • output2_scale_scalar (Optional[torch.Tensor]) – [local_num_experts] scaling factors for the second-layer output.

  • num_experts (int) – Total number of experts.

  • top_k (int) – Number of experts to route to per token.

  • n_group (Optional[int]) – Number of expert groups.

  • topk_group (Optional[int]) – Number of groups to consider for top-k routing.

  • intermediate_size (int) – Size of the intermediate layer.

  • local_expert_offset (int) – Offset of local experts in the global expert space.

  • local_num_experts (int) – Number of experts handled by this device.

  • routed_scaling_factor (Optional[float]) – Scaling factor for routing.

  • routing_method_type (int) –

    Routing method (default 0). Selects the routing-kernel pipeline; matches flashinfer.tllm_enums.RoutingMethodType.

    • 0 Default — Softmax → TopK.

    • 1 Renormalize — TopK → Softmax.

    • 2 DeepSeekV3 — Sigmoid → RoutingBiasAdd → Top-2 in group → Top-topk_group groups → Top-top_k experts from the selected groups.

    • 3 Llama4 — Top-1 → Sigmoid.

    • 4 RenormalizeNaive — Softmax → TopK → Renormalize (Qwen3 style).

    • 5 TopK — TopK only (no softmax/sigmoid).

    • 6 SigmoidRenorm — Sigmoid → TopK → Renormalize (divide by the sum of the top-K weights).

    • 7 MiniMax2 — Sigmoid + Bias → TopK → ScaledSumNormalize (routeScale = 1.0, epsilon = 1e-20).

    • 8 Sigmoid — Sigmoid → TopK (no renormalization).

    • 9 TopKSigmoid — TopK → Sigmoid (no renormalization).

    • 10 Unspecified — reserved.

  • do_finalize (bool) – Whether to finalize the output (default True).

  • enable_pdl (Optional[bool]) – Whether to enable Programmatic Dependent Launch.

  • activation_type (int) – Activation type (default 3 — Swiglu). 10 SiTU computes beta*tanh(x0/beta) * alpha*tanh(x1/alpha)*sigmoid(x1).

  • per_token_scale (Optional[torch.Tensor]) – [seq_len] per-token scaling factors, float32.

  • output (Optional[torch.Tensor]) – Optional in-place [seq_len, hidden_size] output tensor.

  • tune_max_num_tokens (int) – Maximum number of tokens for autotuning (default 8192).

  • gemm1_lora_delta (Optional[torch.Tensor]) – Optional MoE LoRA delta of shape [num_tokens, top_k, 2 * intermediate_size], bfloat16. When set it is added to FC1 before the fused gated activation and the post-activation FC1 output is appended to the return list.

  • hidden_states_scale_layout (Optional[flashinfer.tllm_enums.SfLayout]) –

    Layout of hidden_states_scale. See trtllm_fp4_block_scale_moe() — validation only, linear layout only, and omitting it is deprecated.

    Deprecated since version Omitting: this argument is deprecated. None (the default) infers layout_linear and emits a DeprecationWarning. Both the inference and this argument’s optionality are deprecated together.

  • valid_hidden_size (Optional[int]) – Valid (unpadded) hidden dimension. When provided, the hidden_size implied by the tensor shapes is treated as padded and only the valid region is contracted. Matching TensorRT-LLM, the returned tensor is exactly valid_hidden_size wide: GEMM2 computes roundUp(valid_hidden_size, 128) columns (that is what the FC2 weights, weight scales and bias are sized for) and finalize writes only the leading valid_hidden_size of them. A caller-supplied output must therefore have valid_hidden_size columns. Default None (use the full hidden_size).

  • valid_intermediate_size (Optional[int]) – Valid (unpadded) intermediate dimension. When provided, intermediate_size is treated as padded and only the valid region is computed. Default None (use the full intermediate_size).

  • num_fused_shared_experts (Optional[int]) – Number of shared experts to fuse into the MoE kernel (default None / 0). When > 0, every per-expert tensor must have num_experts + num_fused_shared_experts rows in the expert dimension — the shared-expert weights are appended after the routed ones. Every token is unconditionally routed to the shared experts with weight 1.0. Expert parallelism is not yet supported together with fused shared experts: require local_expert_offset == 0 and local_num_experts == num_experts. Only DeepSeekV3 routing is supported when this is > 0.

Returns:

Return shape depends on do_finalize and gemm1_lora_delta; see trtllm_bf16_routed_moe() for the table.

Return type:

torch.Tensor or List[torch.Tensor]