flashinfer.fused_moe.cutlass_fused_moe

flashinfer.fused_moe.cutlass_fused_moe(input: Tensor, token_selected_experts: Tensor, token_final_scales: Tensor, fc1_expert_weights: Tensor, fc2_expert_weights: Tensor, output_dtype: dtype, quant_scales: List[Tensor], fc1_expert_biases: Tensor | None = None, fc2_expert_biases: Tensor | None = None, input_sf: Tensor | None = None, swiglu_alpha: Tensor | None = None, swiglu_beta: Tensor | None = None, swiglu_limit: Tensor | None = None, tp_size: int = 1, tp_rank: int = 0, ep_size: int = 1, ep_rank: int = 0, cluster_size: int = 1, cluster_rank: int = 0, output: Tensor | None = None, enable_alltoall: bool = False, use_deepseek_fp8_block_scale: bool = False, use_w4_group_scaling: bool = False, use_mxfp8_act_scaling: bool = False, min_latency_mode: bool = False, use_packed_weights: bool = False, tune_max_num_tokens: int = 8192, enable_pdl: bool | None = None, activation_type: ActivationType = ActivationType.Swiglu) Tensor

使用 CUTLASS 后端计算混合专家 (MoE) 层。

此函数实现了一个融合的 MoE 层,它将专家选择、专家计算和输出组合到一个操作中。它使用 CUTLASS 进行高效的矩阵乘法,并支持各种数据类型和并行策略。

参数:
  • input (torch.Tensor) – 输入张量,形状为 [num_tokens, hidden_size]。支持 float、float16、bfloat16、float8_e4m3fn 和 nvfp4。对于 FP8,输入必须量化。对于 NVFP4,支持量化和非量化输入。

  • token_selected_experts (torch.Tensor) – 每个 token 选择的专家的索引。

  • token_final_scales (torch.Tensor) – 每个 token 专家输出的缩放因子。

  • fc1_expert_weights (torch.Tensor) – 每个专家的 GEMM1 权重。

  • fc2_expert_weights (torch.Tensor) – 每个专家的 GEMM2 权重。

  • output_dtype (torch.dtype) – 期望的输出数据类型。

  • quant_scales (List[torch.Tensor]) –

    操作的量化比例。

    NVFP4
    • gemm1 激活全局比例

    • gemm1 权重块比例

    • gemm1 反量化比例

    • gemm2 激活全局比例

    • gemm2 权重块比例

    • gemm2 反量化比例

    FP8
    • gemm1 反量化比例

    • gemm2 激活量化比例

    • gemm2 反量化比例

    • gemm1 输入反量化比例

  • fc1_expert_biases (Optional[torch.Tensor]) – 每个专家的 GEMM1 偏差。

  • fc2_expert_biases (Optional[torch.Tensor]) – 每个专家的 GEMM1 偏差。

  • input_sf (Optional[torch.Tensor]) – 输入缩放因子,用于量化。

  • swiglu_alpha (Optional[torch.Tensor]) – Swiglu 激活的 Swiglu alpha。

  • swiglu_beta (Optional[torch.Tensor]) – Swiglu 激活的 Swiglu beta。

  • swiglu_limit (Optional[torch.Tensor]) – Swiglu 激活的 Swiglu limit。

  • tp_size (int = 1) – 张量并行大小。默认为 1。

  • tp_rank (int = 0) – 张量并行等级。默认为 0。

  • ep_size (int = 1) – 专家并行大小。默认为 1。

  • ep_rank (int = 0) – 专家并行等级。默认为 0。

  • cluster_size (int = 1) – 集群大小。默认为 1。

  • cluster_rank (int = 0) – 集群等级。默认为 0。

  • output (Optional[torch.Tensor] = None) – 输出张量,如果未提供,将在内部分配。

  • enable_alltoall (bool = False) – 是否为专家输出启用 all-to-all 通信。默认为 False。

  • use_deepseek_fp8_block_scale (bool = False) – 是否使用 FP8 块缩放。默认为 False。

  • use_w4_group_scaling (bool = False) – 是否使用 W4A8 组缩放。默认为 False。

  • use_mxfp8_act_scaling (bool = False) – 是否使用 MXFP8 激活缩放。默认为 False。

  • min_latency_mode (bool = False) – 是否使用最小延迟模式。默认为 False。

  • use_packed_weights (bool = False) – 是否使用打包的 uint4x2 权重,以打包的 uint8 值传递。默认为 False。

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

  • activation_type (ActivationType = ActivationType.Swiglu) – GEMM1 上的激活函数,请注意,Relu2 表示非门控 GEMM1

返回值:

out – 输出张量,形状为 [seq_len, hidden_size]。

返回值类型:

torch.Tensor

抛出:

NotImplementedError: – 如果请求了以下任何功能但尚未实现:- 最小延迟模式

注意

  • 该函数支持各种数据类型,包括 FP32、FP16、BF16、FP8 和 NVFP4。

  • 它实现了张量并行和专家并行。

  • 目前,诸如 FP8 块缩放和最小延迟模式等一些高级功能

    尚未针对 Blackwell 架构实现。