flashinfer.fused_moe.trtllm_fp8_block_scale_moe

flashinfer.fused_moe.trtllm_fp8_block_scale_moe(routing_logits: Tensor, routing_bias: Tensor | None, hidden_states: Tensor, hidden_states_scale: Tensor, gemm1_weights: Tensor, gemm1_weights_scale: Tensor, gemm2_weights: Tensor, gemm2_weights_scale: 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, routing_method_type: int = 0, use_shuffled_weight: bool = False, weight_layout: int = 0, enable_pdl: bool | None = None, tune_max_num_tokens: int = 8192) Tensor

FP8 块缩放 MoE 操作。

参数:
  • routing_logits – [seq_len, num_experts] 路由 logits 张量

  • routing_bias – [num_experts] 路由偏差张量

  • hidden_states – [seq_len, hidden_size] 输入隐藏状态张量

  • hidden_states_scale – [hidden_size//128, seq_len] 隐藏状态块缩放张量

  • gemm1_weights – [num_experts, 2*intermediate_size, hidden_size] 第一层权重张量

  • gemm1_weights_scale – [num_experts, 2*intermediate_size//128, hidden_size//128] 第一层块缩放张量

  • gemm2_weights – [num_experts, hidden_size, intermediate_size] 第二层权重张量

  • gemm2_weights_scale – [num_experts, hidden_size//128, intermediate_size//128] 第二层块缩放张量

  • num_experts – 专家总数

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

  • n_group – 专家组的数量

  • topk_group – 用于 top-k 路由要考虑的组数

  • intermediate_size – 中间层大小

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

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

  • routed_scaling_factor – 路由缩放因子

  • routing_method_type – 要使用的路由方法类型 (默认: 0)

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

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

返回值:

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

返回值类型:

torch.Tensor