flashinfer.fused_moe.trtllm_fp8_per_tensor_scale_routed_moe

flashinfer.fused_moe.trtllm_fp8_per_tensor_scale_routed_moe(topk_ids: Tensor, routing_bias: Tensor | None, hidden_states: Tensor, gemm1_weights: Tensor, output1_scales_scalar: Tensor, output1_scales_gate_scalar: Tensor, gemm2_weights: Tensor, output2_scales_scalar: Tensor, 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, use_routing_scales_on_input: bool, routing_method_type: int = 0, do_finalize: bool = True, enable_pdl: bool | None = None, tune_max_num_tokens: int = 8192, activation_type: int = 3, routing_replay_out: Tensor | None = None, output: Tensor | None = None) List[Tensor] | Tensor

Pre-routed FP8 per-tensor-scale MoE operation.

Like trtllm_fp8_per_tensor_scale_moe(), but consumes a pre-computed packed (expert_id, weight) tensor instead of routing logits. Use this entry point for distributed MoE where routing (top-k selection, including EPLB redundant-expert placement) happens in an external DP/EP dispatch, or for CUDA-graph capture (avoids the CPU-GPU sync from logits processing).

Parameters:
  • topk_ids (torch.Tensor) – [seq_len, top_k] int32 tensor of packed expert indices and weights with format (expert_id << 16) | (weight_bf16.view(int16)).

  • routing_bias (Optional[torch.Tensor]) – [num_experts] tensor of routing bias (may be None).

  • hidden_states (torch.Tensor) – [seq_len, hidden_size] tensor of input hidden states.

  • gemm1_weights (torch.Tensor) – [num_experts, 2 * intermediate_size, hidden_size] first-layer weights.

  • output1_scales_scalar (torch.Tensor) – [local_num_experts] first-layer output scales.

  • output1_scales_gate_scalar (torch.Tensor) – [local_num_experts] first-layer gate scales.

  • gemm2_weights (torch.Tensor) – [num_experts, hidden_size, intermediate_size] second-layer weights.

  • output2_scales_scalar (torch.Tensor) – [local_num_experts] second-layer output scales.

  • 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.

  • use_routing_scales_on_input (bool) – Whether to use routing scales on input (Llama4-style).

  • routing_method_type (int) – Routing method (default 0). Matches flashinfer.tllm_enums.RoutingMethodType; see trtllm_fp8_per_tensor_scale_moe() for the full list.

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

  • enable_pdl (Optional[bool]) – Whether to enable Programmatic Dependent Launch. None (default) lets the runtime auto-select on SM90+.

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

  • activation_type (int) – Activation type (default 3 — Swiglu).

  • routing_replay_out (Optional[torch.Tensor]) – Optional int16 tensor of shape (num_tokens_or_larger, top_k) used to capture the selected expert IDs during routing.

  • output (Optional[torch.Tensor]) – Optional in-place output tensor of shape [seq_len, hidden_size]. Allocated internally when None (default).

Returns:

Final MoE output when do_finalize is True, otherwise [gemm2_output, expert_weights, expanded_idx_to_permuted_idx].

Return type:

torch.Tensor or List[torch.Tensor]