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]