flashinfer.fused_moe.prims_ts_fp4_block_scale_routed_moe

flashinfer.fused_moe.prims_ts_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, weight_layout: 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) List[Tensor]

FP4 block-scaled MoE with precomputed routing (Prims-TS / SM100).

Same arguments and return value as trtllm_fp4_block_scale_routed_moe().

Parameters:
  • topk_ids (torch.Tensor or Tuple[torch.Tensor, torch.Tensor]) – Packed (expert_id, weight) tensor or unpacked (topk_ids, topk_weights) pair.

  • routing_bias (Optional[torch.Tensor]) – Optional [num_experts] routing bias.

  • hidden_states (torch.Tensor) – Activations (NVFP4 packed uint8, MXFP8, or BF16 depending on mode).

  • hidden_states_scale (Optional[torch.Tensor]) – Block scales for hidden_states when required by the quant mode.

  • gemm1_weights (torch.Tensor) – Packed FP4 FC1 weights.

  • gemm1_weights_scale (torch.Tensor) – FC1 block scales.

  • gemm1_bias (Optional[torch.Tensor]) – Optional FC1 bias.

  • gemm1_alpha (Optional[torch.Tensor]) – Optional per-expert SwiGLU / SiTU alpha.

  • gemm1_beta (Optional[torch.Tensor]) – Optional per-expert SwiGLU / SiTU beta.

  • gemm1_clamp_limit (Optional[torch.Tensor]) – Optional per-expert clamp limit.

  • gemm2_weights (torch.Tensor) – Packed FP4 FC2 weights.

  • gemm2_weights_scale (torch.Tensor) – FC2 block scales.

  • gemm2_bias (Optional[torch.Tensor]) – Optional FC2 bias.

  • output1_scale_scalar (Optional[torch.Tensor]) – Per-expert FC1 output scale.

  • output1_scale_gate_scalar (Optional[torch.Tensor]) – Per-expert FC1 gate scale.

  • output2_scale_scalar (Optional[torch.Tensor]) – Per-expert FC2 output scale.

  • num_experts (int) – Total number of experts.

  • top_k (int) – Experts selected per token.

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

  • topk_group (Optional[int]) – Groups considered for top-k routing.

  • intermediate_size (int) – Intermediate (FFN) width.

  • local_expert_offset (int) – Global offset of the first local expert.

  • local_num_experts (int) – Number of experts resident on this device.

  • routed_scaling_factor (Optional[float]) – Optional routing scale.

  • routing_method_type (int) – Routing method selector (default 0).

  • weight_layout (int) – Weight layout enum value (default MajorK).

  • do_finalize (bool) – If True, return the finalized MoE output.

  • enable_pdl (Optional[bool]) – Enable Programmatic Dependent Launch when supported.

  • activation_type (int) – Activation enum value (default Swiglu).

  • per_token_scale (Optional[torch.Tensor]) – Optional per-token scales.

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

  • tune_max_num_tokens (int) – Autotune token-bucket upper bound (default 8192).

Returns:

Same return contract as trtllm_fp4_block_scale_routed_moe().

Return type:

List[torch.Tensor]