Support triton kernels v3.4.0 for fused_moe (#8258)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
Co-authored-by: Cheng Wan <cwan@x.ai>
Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
This commit is contained in:
Yuan Luo
2025-07-27 02:28:49 -07:00
committed by GitHub
co-authored by luoyuan.luo Cheng Wan Cheng Wan
parent 10ee89559e
commit b3eac168e7
3 changed files with 109 additions and 40 deletions
+84 -22
View File
@@ -15,7 +15,8 @@
from __future__ import annotations
import math
from typing import Callable, NamedTuple, Optional
from enum import Enum, auto
from typing import Callable, NamedTuple, Optional, Protocol, runtime_checkable
import torch
import torch.nn.functional as F
@@ -27,6 +28,7 @@ from sglang.srt.eplb.expert_location_dispatch import (
ExpertLocationDispatchInfo,
topk_ids_logical_to_physical,
)
from sglang.srt.managers.schedule_batch import global_server_args_dict
from sglang.srt.utils import (
cpu_has_amx_support,
get_bool_env_var,
@@ -37,6 +39,12 @@ from sglang.srt.utils import (
is_npu,
)
try:
from triton_kernels.routing import GatherIndx, RoutingData, ScatterIndx, routing
except ImportError:
pass
_is_cuda = is_cuda()
_is_hip = is_hip()
_is_cpu = is_cpu()
@@ -58,16 +66,59 @@ if _is_npu:
import torch_npu
class TopKOutput(NamedTuple):
# -------------------------------- TopKOutput ---------------------------------------
class TopKOutputFormat(Enum):
STANDARD = auto()
TRITON_KERNEL = auto()
def is_standard(self) -> bool:
return self == TopKOutputFormat.STANDARD
def is_triton_kernel(self) -> bool:
return self == TopKOutputFormat.TRITON_KERNEL
@runtime_checkable
class TopKOutput(Protocol):
"""Protocol for top-k outputs in different formats."""
@property
def format(self) -> TopKOutputFormat:
"""The format of the output."""
...
class StandardTopKOutput(NamedTuple):
"""Standard top-k output format."""
topk_weights: torch.Tensor
topk_ids: torch.Tensor
router_logits: torch.Tensor
@property
def format(self) -> TopKOutputFormat:
return TopKOutputFormat.STANDARD
class TritonKernelTopKOutput(NamedTuple):
"""Triton kernel top-k output format."""
routing_data: RoutingData
gather_indx: GatherIndx
scatter_indx: ScatterIndx
@property
def format(self) -> TopKOutputFormat:
return TopKOutputFormat.TRITON_KERNEL
# -------------------------------- TopK ---------------------------------------
class TopK(CustomOp):
# TODO(ch-wan): support triton_kernels
def __init__(
self,
top_k: int,
@@ -97,6 +148,8 @@ class TopK(CustomOp):
self.correction_bias = correction_bias
self.routed_scaling_factor = routed_scaling_factor
self.use_triton_kernels = global_server_args_dict["enable_triton_kernel_moe"]
def forward_native(
self,
hidden_states: torch.Tensor,
@@ -131,23 +184,29 @@ class TopK(CustomOp):
num_token_non_padded: Optional[torch.Tensor] = None,
expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None,
) -> TopKOutput:
torch_native = False
return select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=self.top_k,
use_grouped_topk=self.use_grouped_topk,
renormalize=self.renormalize,
topk_group=self.topk_group,
num_expert_group=self.num_expert_group,
num_fused_shared_experts=self.num_fused_shared_experts,
custom_routing_function=self.custom_routing_function,
correction_bias=self.correction_bias,
torch_native=torch_native,
routed_scaling_factor=self.routed_scaling_factor,
num_token_non_padded=num_token_non_padded,
expert_location_dispatch_info=expert_location_dispatch_info,
)
if self.use_triton_kernels:
routing_data, gather_idx, scatter_idx = routing(
router_logits, self.top_k, self.renormalize
)
return TritonKernelTopKOutput(routing_data, gather_idx, scatter_idx)
else:
torch_native = False
return select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=self.top_k,
use_grouped_topk=self.use_grouped_topk,
renormalize=self.renormalize,
topk_group=self.topk_group,
num_expert_group=self.num_expert_group,
num_fused_shared_experts=self.num_fused_shared_experts,
custom_routing_function=self.custom_routing_function,
correction_bias=self.correction_bias,
torch_native=torch_native,
routed_scaling_factor=self.routed_scaling_factor,
num_token_non_padded=num_token_non_padded,
expert_location_dispatch_info=expert_location_dispatch_info,
)
def forward_cpu(
self,
@@ -217,6 +276,9 @@ class TopK(CustomOp):
)
# ------------------------------- TopK implementation -------------------------------------
def fused_topk_torch_native(
hidden_states: torch.Tensor,
gating_output: torch.Tensor,
@@ -680,4 +742,4 @@ def select_experts(
get_global_expert_distribution_recorder().on_select_experts(topk_ids=topk_ids)
return TopKOutput(topk_weights, topk_ids, router_logits)
return StandardTopKOutput(topk_weights, topk_ids, router_logits)