flashinfer.fused_moe.trtllm_fp4_block_scale_moe

flashinfer.fused_moe.trtllm_fp4_block_scale_moe(routing_logits: 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, gated_act_type: int = 0, output: Tensor | None = None, tune_max_num_tokens: int = 8192) List[Tensor]

FP4 块缩放 MoE 操作。

参数:
  • routing_logits (torch.Tensor) – shape [seq_len, num_experts] 路由 logits 的输入张量。支持 float32, bfloat16。

  • routing_bias (Optional[torch.Tensor]) – shape [num_experts] 路由偏差张量。对于某些路由方法可以是 None。必须与路由 logits 的类型相同。

  • hidden_states (torch.Tensor) – shape [seq_len, hidden_size // 2 if nvfp4 else hidden_size] 输入隐藏状态张量。支持 bfloat16, mxfp8 和 nvfp4(打包到 uint8 中)

  • hidden_states_scale (Optional[torch.Tensor]) – shape [seq_len, hidden_size // (32 if mxfp8, 16 if mxfp4)] mxfp8 / nvfp4 隐藏状态的缩放张量。Dtype 必须是 float8。

  • gemm1_weights (torch.Tensor) – shape [num_experts, 2 * intermediate_size, hidden_size // 2] FC1 权重张量。Dtype 必须是 uint8(打包 fp4)

  • gemm1_weights_scale (torch.Tensor) – shape [num_experts, 2 * intermediate_size, hidden_size // (32 if mxfp4 else 16)] FC1 权重的缩放张量。Dtype 必须是 float8。

  • gemm1_bias (Optional[torch.Tensor]) – shape [num_experts, 2 * intermediate_size] FC1 偏差张量。Dtype 是 float32。

  • gemm1_alpha (Optional[torch.Tensor]) – shape [num_experts] swiglu alpha 张量。Dtype 是 float32。

  • gemm1_beta (Optional[torch.Tensor]) – shape [num_experts] swiglu beta 张量。Dtype 是 float32。

  • gemm1_clamp_limit (Optional[torch.Tensor]) – shape [num_experts] swiglu clamp limit 张量。Dtype 是 float32。

  • gemm2_weights (torch.Tensor) – shape [num_experts, hidden_size, intermediate_size] FC2 权重张量。Dtype 必须是 uint8(打包 fp4)

  • gemm2_weights_scale (torch.Tensor) – shape [num_experts, hidden_size, intermediate_size // (32 if mxfp4 else 16)] FC2 权重的缩放张量。Dtype 必须是 float8。

  • gemm2_bias (Optional[torch.Tensor]) – shape [num_experts, hidden_size] FC2 偏差张量。Dtype 是 float32。

  • output1_scale_scalar (Optional[torch.Tensor]) – shape [local_num_experts] 第一层激活输出的缩放因子张量

  • output1_scale_gate_scalar (Optional[torch.Tensor]) – shape [local_num_experts] 第一层门输出的缩放因子张量

  • output2_scale_scalar (Optional[torch.Tensor]) – shape [local_num_experts] 第二层输出的缩放因子张量

  • num_experts (int) – 专家总数

  • top_k (int) – 每个 token 路由到专家的数量

  • n_group (Optional[int]) – 专家组的数量(对于某些路由方法可以是 None)

  • topk_group (Optional[int]) – 用于 top-k 路由要考虑的组的数量(对于某些路由方法可以是 None)

  • intermediate_size (int) – 中间层的大小

  • local_expert_offset (int) – 全局专家空间中本地专家的偏移量

  • local_num_experts (int) – 此设备处理的专家数量

  • routed_scaling_factor (Optional[float]) – 路由的缩放因子(对于某些路由方法可以是 None)

  • routing_method_type (int) – 要使用的路由方法类型(默认:0)- 0:默认(Softmax -> TopK)- 1:重归一化(TopK -> Softmax)- 2:DeepSeekV3(Sigmoid -> RoutingBiasAdd -> Top2 in group -> Top4 groups -> Top8 experts)- 3:Llama4(Top1 -> Sigmoid)- 4:RenormalizeNaive(Softmax -> TopK -> Renormalize)

  • do_finalize (bool) – 是否完成输出(默认:False)

  • enable_pdl (Optional[bool]) – 是否启用程序依赖启动 (PDL)。对于 >= sm90,自动启用。

  • gated_act_type (int) – 门控激活函数的类型(默认:0)- 0:SwiGlu - 1:GeGlu

  • tune_max_num_tokens (int) – 调优的最大 token 数。(默认:8192)

  • output (Optional[torch.Tensor]) – shape [seq_len, hidden_size] 可选的就地输出张量。

返回值:

输出张量列表。如果 do_finalize=True,则返回最终的 MoE 输出。

否则,返回需要进一步处理的中间结果(gemm2_output、expert_weights、expanded_idx_to_permuted_idx)。

返回值类型:

List[torch.Tensor]