94057c3d3e
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
204 lines
6.4 KiB
Python
204 lines
6.4 KiB
Python
"""Triton kernels MoE runner backend skeleton."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from typing import TYPE_CHECKING, Any, Optional
|
|
|
|
import torch
|
|
|
|
from sglang.srt.layers.moe.moe_runner.base import (
|
|
MoeQuantInfo,
|
|
MoeRunnerConfig,
|
|
MoeRunnerCore,
|
|
RunnerInput,
|
|
RunnerOutput,
|
|
register_post_permute,
|
|
register_pre_permute,
|
|
)
|
|
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
|
|
|
if TYPE_CHECKING:
|
|
from triton_kernels.matmul_ogs import (
|
|
GatherIndx,
|
|
PrecisionConfig,
|
|
RoutingData,
|
|
ScatterIndx,
|
|
)
|
|
|
|
from sglang.srt.layers.moe.token_dispatcher.standard import (
|
|
StandardCombineInput,
|
|
StandardDispatchOutput,
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Runner IO dataclasses
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@dataclass
|
|
class TritonKernelsRunnerInput(RunnerInput):
|
|
"""Input bundle passed to the triton-kernels runner core."""
|
|
|
|
hidden_states: torch.Tensor
|
|
routing_data: RoutingData
|
|
gather_indx: GatherIndx
|
|
scatter_indx: ScatterIndx
|
|
|
|
@property
|
|
def runner_backend(self) -> MoeRunnerBackend:
|
|
return MoeRunnerBackend.TRITON_KERNELS
|
|
|
|
|
|
@dataclass
|
|
class TritonKernelsRunnerOutput(RunnerOutput):
|
|
"""Output bundle returned from the triton-kernels runner core."""
|
|
|
|
hidden_states: torch.Tensor
|
|
|
|
@property
|
|
def runner_backend(self) -> MoeRunnerBackend:
|
|
return MoeRunnerBackend.TRITON_KERNELS
|
|
|
|
|
|
@dataclass
|
|
class TritonKernelsQuantInfo(MoeQuantInfo):
|
|
"""Quantization payload consumed by the triton-kernels backend."""
|
|
|
|
w13_weight: torch.Tensor
|
|
w2_weight: torch.Tensor
|
|
w13_bias: Optional[torch.Tensor] = None
|
|
w2_bias: Optional[torch.Tensor] = None
|
|
w13_precision_config: Optional[PrecisionConfig] = None
|
|
w2_precision_config: Optional[PrecisionConfig] = None
|
|
global_num_experts: int = -1
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Runner core
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TritonKernelsRunnerCore(MoeRunnerCore):
|
|
"""Execute MoE experts via the external triton_kernels package."""
|
|
|
|
def run(
|
|
self,
|
|
runner_input: TritonKernelsRunnerInput,
|
|
quant_info: TritonKernelsQuantInfo,
|
|
running_state: dict,
|
|
hooks: Optional[Any] = None,
|
|
) -> TritonKernelsRunnerOutput:
|
|
from sglang.srt.layers.moe.fused_moe_triton.triton_kernels_moe import (
|
|
triton_kernel_fused_experts,
|
|
triton_kernel_fused_experts_with_bias,
|
|
)
|
|
|
|
assert (
|
|
self.config.is_gated
|
|
), "Only gated MoEs are supported for Triton Kernels runner"
|
|
|
|
hidden_states = runner_input.hidden_states
|
|
|
|
common_kwargs = dict(
|
|
routing_data=runner_input.routing_data,
|
|
gather_indx=runner_input.gather_indx,
|
|
scatter_indx=None if self.config.no_combine else runner_input.scatter_indx,
|
|
inplace=False,
|
|
activation=self.config.activation,
|
|
apply_router_weight_on_input=self.config.apply_router_weight_on_input,
|
|
global_num_experts=quant_info.global_num_experts,
|
|
)
|
|
|
|
has_bias = quant_info.w13_bias is not None or quant_info.w2_bias is not None
|
|
|
|
if has_bias:
|
|
assert (
|
|
quant_info.w13_bias is not None and quant_info.w2_bias is not None
|
|
), "Bias execution requires both w13_bias and w2_bias"
|
|
output = triton_kernel_fused_experts_with_bias(
|
|
hidden_states=hidden_states,
|
|
w1=quant_info.w13_weight,
|
|
w1_pcg=quant_info.w13_precision_config,
|
|
b1=quant_info.w13_bias,
|
|
w2=quant_info.w2_weight,
|
|
w2_pcg=quant_info.w2_precision_config,
|
|
b2=quant_info.w2_bias,
|
|
gemm1_alpha=self.config.gemm1_alpha,
|
|
gemm1_clamp_limit=self.config.gemm1_clamp_limit,
|
|
**common_kwargs,
|
|
)
|
|
else:
|
|
output = triton_kernel_fused_experts(
|
|
hidden_states=hidden_states,
|
|
w1=quant_info.w13_weight,
|
|
w2=quant_info.w2_weight,
|
|
**common_kwargs,
|
|
)
|
|
|
|
if self.config.no_combine:
|
|
tokens = runner_input.hidden_states.shape[0]
|
|
hidden = runner_input.hidden_states.shape[-1]
|
|
total_rows = output.shape[0]
|
|
top_k = total_rows // tokens
|
|
output = output.view(tokens, top_k, hidden)
|
|
|
|
return TritonKernelsRunnerOutput(hidden_states=output)
|
|
|
|
@property
|
|
def runner_backend(self) -> MoeRunnerBackend:
|
|
return MoeRunnerBackend.TRITON_KERNELS
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Permute / fused hooks
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@register_pre_permute("standard", "triton_kernel")
|
|
def pre_permute_standard_to_triton_kernels(
|
|
dispatch_output: StandardDispatchOutput,
|
|
quant_info: TritonKernelsQuantInfo,
|
|
runner_config: MoeRunnerConfig,
|
|
running_state: dict,
|
|
) -> TritonKernelsRunnerInput:
|
|
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
|
|
|
hidden_states = dispatch_output.hidden_states
|
|
topk_output = dispatch_output.topk_output
|
|
|
|
assert TopKOutputChecker.format_is_triton_kernels(
|
|
topk_output
|
|
), "Triton-kernel runner expects TritonKernelTopKOutput"
|
|
|
|
routing_data, gather_indx, scatter_indx = topk_output
|
|
|
|
return TritonKernelsRunnerInput(
|
|
hidden_states=hidden_states,
|
|
routing_data=routing_data,
|
|
gather_indx=gather_indx,
|
|
scatter_indx=scatter_indx,
|
|
)
|
|
|
|
|
|
@register_post_permute("triton_kernel", "standard")
|
|
def post_permute_triton_kernels_to_standard(
|
|
runner_output: TritonKernelsRunnerOutput,
|
|
quant_info: TritonKernelsQuantInfo,
|
|
runner_config: MoeRunnerConfig,
|
|
running_state: dict,
|
|
) -> StandardCombineInput:
|
|
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
|
|
|
hidden_states = runner_output.hidden_states
|
|
|
|
if (
|
|
runner_config.routed_scaling_factor is not None
|
|
and runner_config.routed_scaling_factor != 1.0
|
|
and not runner_config.no_combine
|
|
):
|
|
hidden_states.mul_(runner_config.routed_scaling_factor)
|
|
|
|
return StandardCombineInput(hidden_states=hidden_states)
|