[Bug] Fix input arguments of flashinfer_trtllm_moe (#9317)

This commit is contained in:
Jiaqi Gu
2025-08-18 18:03:19 -07:00
committed by GitHub
parent a31ea44824
commit 3c2c9f6c9e
3 changed files with 30 additions and 21 deletions
+14 -14
View File
@@ -85,8 +85,8 @@ if _is_npu:
class TopKConfig:
top_k: int
use_grouped_topk: bool = False
topk_group: int = 0
num_expert_group: int = 0
topk_group: Optional[int] = None
num_expert_group: Optional[int] = None
renormalize: bool = True
num_fused_shared_experts: int = 0
custom_routing_function: Optional[Callable] = None
@@ -189,8 +189,8 @@ class TopK(CustomOp):
top_k: int,
*,
use_grouped_topk: bool = False,
topk_group: int = 0,
num_expert_group: int = 0,
topk_group: Optional[int] = None,
num_expert_group: Optional[int] = None,
renormalize: bool = True,
num_fused_shared_experts: int = 0,
custom_routing_function: Optional[Callable] = None,
@@ -427,8 +427,8 @@ def grouped_topk_gpu(
gating_output: torch.Tensor,
topk: int,
renormalize: bool,
num_expert_group: int = 0,
topk_group: int = 0,
num_expert_group: Optional[int] = None,
topk_group: Optional[int] = None,
num_fused_shared_experts: int = 0,
routed_scaling_factor: Optional[float] = None,
num_token_non_padded: Optional[torch.Tensor] = None,
@@ -492,8 +492,8 @@ def grouped_topk_cpu(
gating_output: torch.Tensor,
topk: int,
renormalize: bool,
num_expert_group: int = 0,
topk_group: int = 0,
num_expert_group: Optional[int] = None,
topk_group: Optional[int] = None,
num_fused_shared_experts: int = 0,
routed_scaling_factor: Optional[float] = None,
num_token_non_padded: Optional[torch.Tensor] = None,
@@ -522,8 +522,8 @@ def biased_grouped_topk_impl(
correction_bias: torch.Tensor,
topk: int,
renormalize: bool,
num_expert_group: int = 0,
topk_group: int = 0,
num_expert_group: Optional[int] = None,
topk_group: Optional[int] = None,
num_fused_shared_experts: int = 0,
routed_scaling_factor: Optional[float] = None,
num_token_non_padded: Optional[torch.Tensor] = None,
@@ -615,8 +615,8 @@ def biased_grouped_topk_gpu(
correction_bias: torch.Tensor,
topk: int,
renormalize: bool,
num_expert_group: int = 0,
topk_group: int = 0,
num_expert_group: Optional[int] = None,
topk_group: Optional[int] = None,
num_fused_shared_experts: int = 0,
routed_scaling_factor: Optional[float] = None,
num_token_non_padded: Optional[torch.Tensor] = None,
@@ -690,8 +690,8 @@ def biased_grouped_topk_cpu(
correction_bias: torch.Tensor,
topk: int,
renormalize: bool,
num_expert_group: int = 0,
topk_group: int = 0,
num_expert_group: Optional[int] = None,
topk_group: Optional[int] = None,
compiled: bool = True,
num_fused_shared_experts: int = 0,
routed_scaling_factor: Optional[float] = None,