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