Files
wehub-resource-sync 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
chore: import upstream snapshot with attribution
2026-07-13 12:38:16 +08:00

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)