flashinfer.fused_moe.trtllm_fp8_per_tensor_scale_moe¶
- flashinfer.fused_moe.trtllm_fp8_per_tensor_scale_moe(routing_logits: Tensor, routing_bias: Tensor | None, hidden_states: Tensor, gemm1_weights: Tensor, output1_scales_scalar: Tensor, output1_scales_gate_scalar: Tensor, gemm2_weights: Tensor, output2_scales_scalar: 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, use_routing_scales_on_input: bool, routing_method_type: 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] 输入隐藏状态张量
gemm1_weights – [num_experts, 2*intermediate_size, hidden_size] 第一层权重张量
output1_scales_scalar – [local_num_experts] 第一层输出缩放
output1_scales_gate_scalar – [local_num_experts] 第一层门缩放
gemm2_weights – [num_experts, hidden_size, intermediate_size] 第二层权重张量
output2_scales_scalar – [local_num_experts] 第二层输出缩放
num_experts – 专家总数
top_k – 每个 token 路由到专家的数量
n_group – 专家组的数量
topk_group – 用于 top-k 路由要考虑的组数
intermediate_size – 中间层的大小
local_expert_offset – 在全局专家空间中本地专家的偏移量
local_num_experts – 此设备处理的专家数量
routed_scaling_factor – 路由缩放因子
use_routing_scales_on_input – 是否在输入上使用路由缩放
routing_method_type – 要使用的路由方法类型(默认值:0)
enable_pdl – 是否启用程序依赖启动 (PDL)。对于 >= sm90,自动启用。
tune_max_num_tokens (int) – 调优的最大 token 数。(默认:8192)
- 返回值:
输出张量形状为 [seq_len, hidden_size]
- 返回值类型:
torch.Tensor