From db34d8cc89fa88a1ee229499ed394ccc32dd2b84 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Thu, 30 Jul 2026 12:46:43 +0800 Subject: [PATCH 01/15] :building_construction: refactor(npu): delegate Gemma RMSNorm to kernel provider --- python/sglang/kernels/fused_op.py | 5 + .../sglang/kernels/ops/layernorm/__init__.py | 113 ++++++- python/sglang/kernels/spec.py | 3 +- python/sglang/srt/layers/layernorm.py | 23 +- .../ops/layernorm/test_npu_gemma_rmsnorm.py | 316 ++++++++++++++++++ 5 files changed, 443 insertions(+), 17 deletions(-) create mode 100644 test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py diff --git a/python/sglang/kernels/fused_op.py b/python/sglang/kernels/fused_op.py index d39ed4d1940b..548f26ab2a78 100644 --- a/python/sglang/kernels/fused_op.py +++ b/python/sglang/kernels/fused_op.py @@ -70,6 +70,7 @@ KernelBackend.FLASHINFER: "forward_flashinfer", KernelBackend.DEEPGEMM: "forward_deepgemm", KernelBackend.AITER: "forward_aiter", + KernelBackend.SGL_KERNEL_NPU: "forward_sgl_kernel_npu", KernelBackend.TORCH_NPU: "forward_npu", } @@ -83,6 +84,7 @@ KernelBackend.DEEPGEMM, KernelBackend.CUTE_DSL, KernelBackend.AITER, + KernelBackend.SGL_KERNEL_NPU, KernelBackend.TORCH_NPU, KernelBackend.TRITON, KernelBackend.TORCH, @@ -271,6 +273,9 @@ def forward_deepgemm(self, *args, **kwargs): def forward_aiter(self, *args, **kwargs): raise NotImplementedError(f"{self.op}: no aiter backend") + def forward_sgl_kernel_npu(self, *args, **kwargs): + raise NotImplementedError(f"{self.op}: no sgl_kernel_npu backend") + def forward_npu(self, *args, **kwargs): raise NotImplementedError(f"{self.op}: no npu backend") diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index 2e04ccdbe1f4..819ff4355229 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -5,7 +5,8 @@ all behind one signature. The public module-level functions are thin wrappers over module-level instances; auto-selection follows the production default for the live device: AOT ``sgl_kernel`` on CUDA, ``aiter`` (or rocm-triton for -gemma) on ROCm, ``torch_npu`` on Ascend, native reference otherwise. +gemma) on ROCm, ``sgl_kernel_npu`` then ``torch_npu`` on Ascend, native +reference otherwise. Pick a specific backend with e.g. ``_RMSNORM.forward(x, w, backend=KernelBackend.JIT)`` or globally via ``SGLANG_FORCE_FUSED_OP_BACKEND``. @@ -13,7 +14,8 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Optional +from functools import lru_cache +from typing import TYPE_CHECKING, Callable, Optional, Tuple from sglang.kernels.fused_op import BaseFusedOp, register_fused_op from sglang.kernels.spec import ( @@ -29,6 +31,24 @@ _CUDA = frozenset({CapabilityRequirement.CUDA}) _HIP = frozenset({CapabilityRequirement.HIP}) _NPU = frozenset({CapabilityRequirement.NPU}) + + +@lru_cache(maxsize=1) +def _load_sgl_kernel_npu_gemma_ops() -> Optional[Tuple[Callable, Callable]]: + """Load the stable NPU Gemma API without imposing a package version floor.""" + + try: + from sgl_kernel_npu.norm.gemma_rmsnorm import ( + add_gemma_rms_norm, + gemma_rms_norm, + ) + except (ImportError, AttributeError, OSError): + return None + if not callable(gemma_rms_norm) or not callable(add_gemma_rms_norm): + return None + return gemma_rms_norm, add_gemma_rms_norm + + # Unlike the gated-activation ops, sgl_kernel does *not* build the rmsnorm ops # for ROCm (production: ``if _is_cuda or _is_xpu or _is_musa: from sgl_kernel # import rmsnorm`` — HIP is absent), so AOT here is CUDA-only. ROCm instead has @@ -41,6 +61,7 @@ KernelBackend.AOT, KernelBackend.JIT, KernelBackend.AITER, + KernelBackend.SGL_KERNEL_NPU, KernelBackend.TORCH_NPU, KernelBackend.TORCH, ) @@ -280,6 +301,7 @@ class GemmaRMSNormOp(BaseFusedOp): capabilities = { KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, + KernelBackend.SGL_KERNEL_NPU: _NPU, KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( @@ -291,7 +313,12 @@ class GemmaRMSNormOp(BaseFusedOp): KernelBackend.JIT: ( "Gemma-style RMS normalization (rocm-triton, sglang.kernels.jit)." ), - KernelBackend.TORCH_NPU: ("Gemma-style RMS normalization (torch_npu, Ascend)."), + KernelBackend.SGL_KERNEL_NPU: ( + "Gemma-style RMS normalization (sgl_kernel_npu, Ascend)." + ), + KernelBackend.TORCH_NPU: ( + "Gemma-style RMS normalization (torch_npu fallback, Ascend)." + ), KernelBackend.TORCH: "Gemma-style RMS normalization (pure-torch reference).", } @@ -354,12 +381,38 @@ def forward_npu( ) -> torch.Tensor: import torch_npu - result = torch_npu.npu_gemma_rms_norm(input, weight, eps)[0] + # Cross-SoC safe fallback: Gemma stores an offset from one as weight. + result = torch_npu.npu_rms_norm(input, 1.0 + weight, eps)[0] if out is None: return result out.copy_(result) return out + def forward_sgl_kernel_npu( + self, + input: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + out: Optional[torch.Tensor] = None, + enable_pdl: Optional[bool] = None, + ) -> torch.Tensor: + ops = _load_sgl_kernel_npu_gemma_ops() + if ops is None: + raise ImportError("sgl_kernel_npu does not provide norm.gemma_rmsnorm") + result = ops[0](input, weight, eps) + if out is None: + return result + out.copy_(result) + return out + + def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: + if not super().backend_eligible(backend, *args, **kwargs): + return False + return ( + backend is not KernelBackend.SGL_KERNEL_NPU + or _load_sgl_kernel_npu_gemma_ops() is not None + ) + class GemmaFusedAddRMSNormOp(BaseFusedOp): """In-place ``residual += input; input = GemmaRMSNorm(residual) * (weight + 1)``.""" @@ -367,11 +420,12 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): op = "layernorm.gemma_fused_add_rmsnorm" priority = _NORM_PRIORITY # AOT (sgl_kernel) on CUDA; JIT is the ROCm rocm-triton path on HIP. - # NPU here would use ``sgl_kernel_npu.add_gemma_rms_norm`` (a distinct AOT-npu - # wheel provenance, not torch_npu) — deferred until that provenance lands. + # Ascend prefers the stable sgl_kernel_npu API, with torch_npu as fallback. capabilities = { KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, + KernelBackend.SGL_KERNEL_NPU: _NPU, + KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( supported_dtypes=_NORM_DTYPES, @@ -384,6 +438,14 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): "Gemma-style fused residual-add + RMS normalization " "(rocm-triton, sglang.kernels.jit)." ), + KernelBackend.SGL_KERNEL_NPU: ( + "Gemma-style fused residual-add + RMS normalization " + "(sgl_kernel_npu, Ascend)." + ), + KernelBackend.TORCH_NPU: ( + "Gemma-style fused residual-add + RMS normalization " + "(torch_npu fallback, Ascend)." + ), KernelBackend.TORCH: ( "Gemma-style fused residual-add + RMS normalization " "(pure-torch reference)." @@ -439,6 +501,45 @@ def forward_jit( input.copy_(norm_out) residual.copy_(residual_out) + def forward_sgl_kernel_npu( + self, + input: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + enable_pdl: Optional[bool] = None, + ) -> None: + ops = _load_sgl_kernel_npu_gemma_ops() + if ops is None: + raise ImportError("sgl_kernel_npu does not provide norm.gemma_rmsnorm") + norm_output, residual_sum = ops[1](input, weight, residual, eps) + input.copy_(norm_output) + residual.copy_(residual_sum) + + def forward_npu( + self, + input: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + enable_pdl: Optional[bool] = None, + ) -> None: + import torch_npu + + norm_output, _, residual_sum = torch_npu.npu_add_rms_norm( + residual, input, 1.0 + weight, eps + ) + input.copy_(norm_output) + residual.copy_(residual_sum) + + def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: + if not super().backend_eligible(backend, *args, **kwargs): + return False + return ( + backend is not KernelBackend.SGL_KERNEL_NPU + or _load_sgl_kernel_npu_gemma_ops() is not None + ) + _RMSNORM = register_fused_op(RMSNormOp(), __name__, "_RMSNORM") _FUSED_ADD_RMSNORM = register_fused_op( diff --git a/python/sglang/kernels/spec.py b/python/sglang/kernels/spec.py index dba590ac1614..d7baae3a268a 100644 --- a/python/sglang/kernels/spec.py +++ b/python/sglang/kernels/spec.py @@ -46,8 +46,9 @@ class KernelBackend(str, Enum): FLASHINFER = "flashinfer" DEEPGEMM = "deepgemm" AITER = "aiter" # AMD aiter library (device=HIP) + SGL_KERNEL_NPU = "sgl_kernel_npu" # SGLang Ascend kernel package (device=NPU) TORCH_NPU = "torch_npu" # Ascend NPU vendor runtime (device=NPU) - # TODO(RFC #29630): more provenance as needed (cpu-avx, sgl_kernel_npu, ...) + # TODO(RFC #29630): more provenance as needed (cpu-avx, ...) class DeviceType(str, Enum): diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index f7643efe29d1..1846390ea603 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -159,7 +159,11 @@ def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool: if _is_npu: import torch_npu - from sgl_kernel_npu.norm.add_rmsnorm_bias import add_gemma_rms_norm + + from sglang.kernels.ops.layernorm import ( + gemma_fused_add_rmsnorm as npu_gemma_fused_add_rmsnorm, + ) + from sglang.kernels.ops.layernorm import gemma_rmsnorm as npu_gemma_rmsnorm @lru_cache(maxsize=1) @@ -995,13 +999,10 @@ def forward_npu( if residual is not None: if post_residual_addition is not None: residual = residual + post_residual_addition - norm_out, residual = add_gemma_rms_norm( - x, self.weight, residual, self.variance_epsilon - ) - return norm_out, residual + npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.variance_epsilon) + return x, residual - x, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.variance_epsilon) - return x + return npu_gemma_rmsnorm(x, self.weight, self.variance_epsilon) def forward_xpu( self, @@ -1093,10 +1094,12 @@ def forward_hip(self, x, residual: Optional[torch.Tensor] = None): return self.forward_native(x, residual) def forward_npu(self, x, residual: Optional[torch.Tensor] = None): - if residual is not None: + if envs.SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM.get(): return self.forward_native(x, residual) - output, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.eps) - return output + if residual is not None: + npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.eps) + return x, residual + return npu_gemma_rmsnorm(x, self.weight, self.eps) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.eps}" diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py new file mode 100644 index 000000000000..51fcf5c01bb7 --- /dev/null +++ b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py @@ -0,0 +1,316 @@ +import sys +from types import ModuleType, SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +import torch + +import sglang.kernels as kernels +import sglang.kernels.fused_op as fused_op_module +import sglang.kernels.ops.layernorm as unified_layernorm +from sglang.kernels.fused_op import BACKEND_METHODS +from sglang.kernels.ops.layernorm import ( + GemmaFusedAddRMSNormOp, + GemmaRMSNormOp, +) +from sglang.kernels.spec import CapabilityRequirement, KernelBackend, PlatformInfo +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +@pytest.fixture(autouse=True) +def _reset_fused_op_state(): + unified_layernorm._load_sgl_kernel_npu_gemma_ops.cache_clear() + yield + unified_layernorm._load_sgl_kernel_npu_gemma_ops.cache_clear() + kernels.set_fused_op_backend(None) + kernels.disable_fused_op_trace() + kernels.clear_fused_op_trace() + + +def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): + package = ModuleType("sgl_kernel_npu") + package.__path__ = [] + norm_package = ModuleType("sgl_kernel_npu.norm") + norm_package.__path__ = [] + module = ModuleType("sgl_kernel_npu.norm.gemma_rmsnorm") + module.gemma_rms_norm = gemma_kernel + module.add_gemma_rms_norm = add_gemma_kernel + return { + "sgl_kernel_npu": package, + "sgl_kernel_npu.norm": norm_package, + "sgl_kernel_npu.norm.gemma_rmsnorm": module, + } + + +def _fake_old_sgl_kernel_npu(): + package = ModuleType("sgl_kernel_npu") + package.__path__ = [] + norm_package = ModuleType("sgl_kernel_npu.norm") + norm_package.__path__ = [] + return { + "sgl_kernel_npu": package, + "sgl_kernel_npu.norm": norm_package, + } + + +def test_backend_enum_method_mapping_and_priority(): + assert KernelBackend.SGL_KERNEL_NPU.value == "sgl_kernel_npu" + assert BACKEND_METHODS[KernelBackend.SGL_KERNEL_NPU] == "forward_sgl_kernel_npu" + assert unified_layernorm._NORM_PRIORITY.index( + KernelBackend.SGL_KERNEL_NPU + ) < unified_layernorm._NORM_PRIORITY.index(KernelBackend.TORCH_NPU) + + +@pytest.mark.parametrize( + "op_name", + ["layernorm.gemma_rmsnorm", "layernorm.gemma_fused_add_rmsnorm"], +) +def test_registry_exposes_sgl_kernel_npu_with_npu_capability(op_name): + spec = kernels.registry.get_backend(op_name, KernelBackend.SGL_KERNEL_NPU) + + assert spec.capabilities == frozenset({CapabilityRequirement.NPU}) + assert spec.target.endswith(".forward_sgl_kernel_npu") + + +def test_lazy_loader_finds_stable_kernel_api(): + gemma_kernel = MagicMock() + add_gemma_kernel = MagicMock() + + with patch.dict(sys.modules, _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel)): + ops = unified_layernorm._load_sgl_kernel_npu_gemma_ops() + + assert ops == (gemma_kernel, add_gemma_kernel) + + +def test_lazy_loader_rejects_old_package_without_stable_api(): + with patch.dict(sys.modules, _fake_old_sgl_kernel_npu()): + assert unified_layernorm._load_sgl_kernel_npu_gemma_ops() is None + + +def test_new_kernel_api_is_preferred_on_npu(): + op = GemmaRMSNormOp() + x = torch.randn(2, 4) + weight = torch.randn(4) + ops = (MagicMock(), MagicMock()) + + with ( + patch.object( + unified_layernorm, "_load_sgl_kernel_npu_gemma_ops", return_value=ops + ), + patch.object( + fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") + ), + ): + backend = op._resolve_backend(x, weight) + + assert backend is KernelBackend.SGL_KERNEL_NPU + + +def test_old_kernel_package_falls_back_to_torch_npu(): + op = GemmaRMSNormOp() + x = torch.randn(2, 4) + weight = torch.randn(4) + + with ( + patch.object( + unified_layernorm, "_load_sgl_kernel_npu_gemma_ops", return_value=None + ), + patch.object( + fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") + ), + ): + backend = op._resolve_backend(x, weight) + + assert backend is KernelBackend.TORCH_NPU + + +def test_non_npu_selection_does_not_import_sgl_kernel_npu(): + op = GemmaRMSNormOp() + + with ( + patch.object( + unified_layernorm, + "_load_sgl_kernel_npu_gemma_ops", + side_effect=AssertionError("NPU package import is forbidden"), + ), + patch.object( + fused_op_module, + "_platform", + return_value=PlatformInfo(device_type="cpu"), + ), + ): + assert not op.backend_eligible(KernelBackend.SGL_KERNEL_NPU) + + +def test_sgl_kernel_npu_selection_does_not_query_soc(): + gemma_kernel = MagicMock() + add_gemma_kernel = MagicMock() + get_soc_version = MagicMock(side_effect=AssertionError("SoC query is forbidden")) + torch_npu = SimpleNamespace(npu=SimpleNamespace(get_soc_version=get_soc_version)) + op = GemmaRMSNormOp() + + with ( + patch.dict( + sys.modules, + { + **_fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel), + "torch_npu": torch_npu, + }, + ), + patch.object( + fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") + ), + ): + assert op._resolve_backend(torch.randn(2, 4), torch.randn(4)) is ( + KernelBackend.SGL_KERNEL_NPU + ) + + get_soc_version.assert_not_called() + + +def test_sgl_kernel_npu_normal_out_contract(): + x = torch.randn(2, 4) + weight = torch.randn(4) + expected = torch.randn_like(x) + out = torch.empty_like(x) + gemma_kernel = MagicMock(return_value=expected) + + with patch.object( + unified_layernorm, + "_load_sgl_kernel_npu_gemma_ops", + return_value=(gemma_kernel, MagicMock()), + ): + result = GemmaRMSNormOp().forward_sgl_kernel_npu(x, weight, 1e-5, out=out) + + assert result is out + torch.testing.assert_close(out, expected) + gemma_kernel.assert_called_once_with(x, weight, 1e-5) + + +def test_sgl_kernel_npu_fused_in_place_contract(): + x = torch.randn(2, 4) + residual = torch.randn(2, 4) + weight = torch.randn(4) + norm_output = torch.randn_like(x) + residual_sum = torch.randn_like(residual) + add_gemma_kernel = MagicMock(return_value=(norm_output, residual_sum)) + + with patch.object( + unified_layernorm, + "_load_sgl_kernel_npu_gemma_ops", + return_value=(MagicMock(), add_gemma_kernel), + ): + result = GemmaFusedAddRMSNormOp().forward_sgl_kernel_npu( + x, residual, weight, 1e-5 + ) + + assert result is None + torch.testing.assert_close(x, norm_output) + torch.testing.assert_close(residual, residual_sum) + add_gemma_kernel.assert_called_once() + args = add_gemma_kernel.call_args.args + assert args[0] is x + assert args[1] is weight + assert args[2] is residual + assert args[3] == 1e-5 + + +def test_torch_npu_normal_fallback_uses_offset_weight(): + x = torch.randn(2, 4) + weight = torch.randn(4) + fallback_kernel = MagicMock(return_value=(x, None)) + torch_npu = SimpleNamespace(npu_rms_norm=fallback_kernel) + + with patch.dict(sys.modules, {"torch_npu": torch_npu}): + result = GemmaRMSNormOp().forward_npu(x, weight) + + assert result is x + args = fallback_kernel.call_args.args + assert args[0] is x + torch.testing.assert_close(args[1], 1.0 + weight) + assert args[2] == 1e-6 + + +def test_torch_npu_fused_fallback_uses_offset_weight_and_writes_back(): + x = torch.randn(2, 4) + residual = torch.randn(2, 4) + weight = torch.randn(4) + norm_output = torch.randn_like(x) + residual_sum = torch.randn_like(residual) + fallback_kernel = MagicMock(return_value=(norm_output, None, residual_sum)) + torch_npu = SimpleNamespace(npu_add_rms_norm=fallback_kernel) + + with patch.dict(sys.modules, {"torch_npu": torch_npu}): + result = GemmaFusedAddRMSNormOp().forward_npu(x, residual, weight) + + assert result is None + torch.testing.assert_close(x, norm_output) + torch.testing.assert_close(residual, residual_sum) + args = fallback_kernel.call_args.args + assert args[0] is residual + assert args[1] is x + torch.testing.assert_close(args[2], 1.0 + weight) + assert args[3] == 1e-6 + + +def test_force_backend_and_trace_use_sgl_kernel_npu(): + x = torch.randn(2, 4) + weight = torch.randn(4) + gemma_kernel = MagicMock(return_value=x) + kernels.set_fused_op_backend(KernelBackend.SGL_KERNEL_NPU) + kernels.enable_fused_op_trace() + + with patch.object( + unified_layernorm, + "_load_sgl_kernel_npu_gemma_ops", + return_value=(gemma_kernel, MagicMock()), + ): + result = unified_layernorm.gemma_rmsnorm(x, weight) + + assert result is x + (record,) = kernels.get_fused_op_trace() + assert record.op == "layernorm.gemma_rmsnorm" + assert record.backend == "sgl_kernel_npu" + + +@pytest.mark.parametrize("layer_name", ["GemmaRMSNorm", "Gemma3RMSNorm"]) +def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): + from sglang.srt.layers import layernorm as layernorm_module + + layer_cls = getattr(layernorm_module, layer_name) + layer = layer_cls(4) + x = torch.randn(2, 4) + unified_op = MagicMock(return_value=x) + + with patch.object(layernorm_module, "npu_gemma_rmsnorm", unified_op, create=True): + result = layer.forward_npu(x) + + assert result is x + eps = layer.variance_epsilon if hasattr(layer, "variance_epsilon") else layer.eps + unified_op.assert_called_once_with(x, layer.weight, eps) + + +@pytest.mark.parametrize("layer_name", ["GemmaRMSNorm", "Gemma3RMSNorm"]) +def test_srt_gemma_layers_delegate_residual_npu_path(layer_name): + from sglang.srt.layers import layernorm as layernorm_module + + layer_cls = getattr(layernorm_module, layer_name) + layer = layer_cls(4) + x = torch.randn(2, 4) + residual = torch.randn(2, 4) + fused_op = MagicMock() + + with patch.object( + layernorm_module, + "npu_gemma_fused_add_rmsnorm", + fused_op, + create=True, + ): + result = layer.forward_npu(x, residual) + + assert result[0] is x + assert result[1] is residual + eps = layer.variance_epsilon if hasattr(layer, "variance_epsilon") else layer.eps + fused_op.assert_called_once_with(x, residual, layer.weight, eps) From 4e87d80156a32a282852f6ed078cca0f9a928ff0 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Thu, 30 Jul 2026 16:02:31 +0800 Subject: [PATCH 02/15] :recycle: refactor(npu): simplify Gemma RMSNorm provider --- .../sglang/kernels/ops/layernorm/__init__.py | 105 ++-------- .../ops/layernorm/test_npu_gemma_rmsnorm.py | 183 ++---------------- 2 files changed, 24 insertions(+), 264 deletions(-) diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index 819ff4355229..5f9841b6e4d5 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -5,8 +5,8 @@ all behind one signature. The public module-level functions are thin wrappers over module-level instances; auto-selection follows the production default for the live device: AOT ``sgl_kernel`` on CUDA, ``aiter`` (or rocm-triton for -gemma) on ROCm, ``sgl_kernel_npu`` then ``torch_npu`` on Ascend, native -reference otherwise. +gemma) on ROCm, ``sgl_kernel_npu`` or ``torch_npu`` on Ascend depending on the +operator, and the native reference otherwise. Pick a specific backend with e.g. ``_RMSNORM.forward(x, w, backend=KernelBackend.JIT)`` or globally via ``SGLANG_FORCE_FUSED_OP_BACKEND``. @@ -14,8 +14,7 @@ from __future__ import annotations -from functools import lru_cache -from typing import TYPE_CHECKING, Callable, Optional, Tuple +from typing import TYPE_CHECKING, Optional from sglang.kernels.fused_op import BaseFusedOp, register_fused_op from sglang.kernels.spec import ( @@ -32,31 +31,14 @@ _HIP = frozenset({CapabilityRequirement.HIP}) _NPU = frozenset({CapabilityRequirement.NPU}) - -@lru_cache(maxsize=1) -def _load_sgl_kernel_npu_gemma_ops() -> Optional[Tuple[Callable, Callable]]: - """Load the stable NPU Gemma API without imposing a package version floor.""" - - try: - from sgl_kernel_npu.norm.gemma_rmsnorm import ( - add_gemma_rms_norm, - gemma_rms_norm, - ) - except (ImportError, AttributeError, OSError): - return None - if not callable(gemma_rms_norm) or not callable(add_gemma_rms_norm): - return None - return gemma_rms_norm, add_gemma_rms_norm - - # Unlike the gated-activation ops, sgl_kernel does *not* build the rmsnorm ops # for ROCm (production: ``if _is_cuda or _is_xpu or _is_musa: from sgl_kernel # import rmsnorm`` — HIP is absent), so AOT here is CUDA-only. ROCm instead has # an ``aiter`` path, and Ascend a ``torch_npu`` path — a clean illustration that # the same ``AOT`` provenance covers different devices per op. # Priority (best -> fallback) is device-agnostic; per-op CapabilityRequirement -# decides eligibility, so on CUDA this resolves to AOT, on HIP to AITER, on NPU -# to TORCH_NPU, each matching the production default for that device. +# decides eligibility, so on CUDA this resolves to AOT, on HIP to AITER, and on +# NPU to the provider implemented by each operator. _NORM_PRIORITY = ( KernelBackend.AOT, KernelBackend.JIT, @@ -297,12 +279,11 @@ class GemmaRMSNormOp(BaseFusedOp): priority = _NORM_PRIORITY # AOT (sgl_kernel) on CUDA; JIT is the ROCm rocm-triton path # (sglang.kernels.ops.moe.minimax_m3_swiglu) — a JIT provenance pinned to HIP, distinct - # from the CUDA-only JIT on the plain rmsnorm ops; torch_npu on Ascend. + # from the CUDA-only JIT on the plain rmsnorm ops; sgl_kernel_npu on Ascend. capabilities = { KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, KernelBackend.SGL_KERNEL_NPU: _NPU, - KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( supported_dtypes=_NORM_DTYPES, @@ -316,9 +297,6 @@ class GemmaRMSNormOp(BaseFusedOp): KernelBackend.SGL_KERNEL_NPU: ( "Gemma-style RMS normalization (sgl_kernel_npu, Ascend)." ), - KernelBackend.TORCH_NPU: ( - "Gemma-style RMS normalization (torch_npu fallback, Ascend)." - ), KernelBackend.TORCH: "Gemma-style RMS normalization (pure-torch reference).", } @@ -371,23 +349,6 @@ def forward_jit( out.copy_(result) return out - def forward_npu( - self, - input: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6, - out: Optional[torch.Tensor] = None, - enable_pdl: Optional[bool] = None, - ) -> torch.Tensor: - import torch_npu - - # Cross-SoC safe fallback: Gemma stores an offset from one as weight. - result = torch_npu.npu_rms_norm(input, 1.0 + weight, eps)[0] - if out is None: - return result - out.copy_(result) - return out - def forward_sgl_kernel_npu( self, input: torch.Tensor, @@ -396,23 +357,14 @@ def forward_sgl_kernel_npu( out: Optional[torch.Tensor] = None, enable_pdl: Optional[bool] = None, ) -> torch.Tensor: - ops = _load_sgl_kernel_npu_gemma_ops() - if ops is None: - raise ImportError("sgl_kernel_npu does not provide norm.gemma_rmsnorm") - result = ops[0](input, weight, eps) + from sgl_kernel_npu.norm.gemma_rmsnorm import gemma_rms_norm + + result = gemma_rms_norm(input, weight, eps) if out is None: return result out.copy_(result) return out - def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: - if not super().backend_eligible(backend, *args, **kwargs): - return False - return ( - backend is not KernelBackend.SGL_KERNEL_NPU - or _load_sgl_kernel_npu_gemma_ops() is not None - ) - class GemmaFusedAddRMSNormOp(BaseFusedOp): """In-place ``residual += input; input = GemmaRMSNorm(residual) * (weight + 1)``.""" @@ -420,12 +372,11 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): op = "layernorm.gemma_fused_add_rmsnorm" priority = _NORM_PRIORITY # AOT (sgl_kernel) on CUDA; JIT is the ROCm rocm-triton path on HIP. - # Ascend prefers the stable sgl_kernel_npu API, with torch_npu as fallback. + # Ascend delegates provider selection to the stable sgl_kernel_npu API. capabilities = { KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, KernelBackend.SGL_KERNEL_NPU: _NPU, - KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( supported_dtypes=_NORM_DTYPES, @@ -442,13 +393,8 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): "Gemma-style fused residual-add + RMS normalization " "(sgl_kernel_npu, Ascend)." ), - KernelBackend.TORCH_NPU: ( - "Gemma-style fused residual-add + RMS normalization " - "(torch_npu fallback, Ascend)." - ), KernelBackend.TORCH: ( - "Gemma-style fused residual-add + RMS normalization " - "(pure-torch reference)." + "Gemma-style fused residual-add + RMS normalization (pure-torch reference)." ), } @@ -509,37 +455,12 @@ def forward_sgl_kernel_npu( eps: float = 1e-6, enable_pdl: Optional[bool] = None, ) -> None: - ops = _load_sgl_kernel_npu_gemma_ops() - if ops is None: - raise ImportError("sgl_kernel_npu does not provide norm.gemma_rmsnorm") - norm_output, residual_sum = ops[1](input, weight, residual, eps) - input.copy_(norm_output) - residual.copy_(residual_sum) + from sgl_kernel_npu.norm.gemma_rmsnorm import add_gemma_rms_norm - def forward_npu( - self, - input: torch.Tensor, - residual: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6, - enable_pdl: Optional[bool] = None, - ) -> None: - import torch_npu - - norm_output, _, residual_sum = torch_npu.npu_add_rms_norm( - residual, input, 1.0 + weight, eps - ) + norm_output, residual_sum = add_gemma_rms_norm(input, weight, residual, eps) input.copy_(norm_output) residual.copy_(residual_sum) - def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: - if not super().backend_eligible(backend, *args, **kwargs): - return False - return ( - backend is not KernelBackend.SGL_KERNEL_NPU - or _load_sgl_kernel_npu_gemma_ops() is not None - ) - _RMSNORM = register_fused_op(RMSNormOp(), __name__, "_RMSNORM") _FUSED_ADD_RMSNORM = register_fused_op( diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py index 51fcf5c01bb7..2e0bc397f6dd 100644 --- a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py +++ b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py @@ -3,12 +3,9 @@ from unittest.mock import MagicMock, patch import pytest -import torch - -import sglang.kernels as kernels import sglang.kernels.fused_op as fused_op_module -import sglang.kernels.ops.layernorm as unified_layernorm -from sglang.kernels.fused_op import BACKEND_METHODS +import torch +from sglang import kernels from sglang.kernels.ops.layernorm import ( GemmaFusedAddRMSNormOp, GemmaRMSNormOp, @@ -19,16 +16,6 @@ register_cpu_ci(est_time=2, suite="base-a-test-cpu") -@pytest.fixture(autouse=True) -def _reset_fused_op_state(): - unified_layernorm._load_sgl_kernel_npu_gemma_ops.cache_clear() - yield - unified_layernorm._load_sgl_kernel_npu_gemma_ops.cache_clear() - kernels.set_fused_op_backend(None) - kernels.disable_fused_op_trace() - kernels.clear_fused_op_trace() - - def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): package = ModuleType("sgl_kernel_npu") package.__path__ = [] @@ -44,121 +31,39 @@ def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): } -def _fake_old_sgl_kernel_npu(): - package = ModuleType("sgl_kernel_npu") - package.__path__ = [] - norm_package = ModuleType("sgl_kernel_npu.norm") - norm_package.__path__ = [] - return { - "sgl_kernel_npu": package, - "sgl_kernel_npu.norm": norm_package, - } - - -def test_backend_enum_method_mapping_and_priority(): - assert KernelBackend.SGL_KERNEL_NPU.value == "sgl_kernel_npu" - assert BACKEND_METHODS[KernelBackend.SGL_KERNEL_NPU] == "forward_sgl_kernel_npu" - assert unified_layernorm._NORM_PRIORITY.index( - KernelBackend.SGL_KERNEL_NPU - ) < unified_layernorm._NORM_PRIORITY.index(KernelBackend.TORCH_NPU) - - @pytest.mark.parametrize( "op_name", ["layernorm.gemma_rmsnorm", "layernorm.gemma_fused_add_rmsnorm"], ) -def test_registry_exposes_sgl_kernel_npu_with_npu_capability(op_name): +def test_registry_exposes_sgl_kernel_npu_as_only_npu_provider(op_name): + registered_backends = {spec.backend for spec in kernels.registry.get(op_name)} spec = kernels.registry.get_backend(op_name, KernelBackend.SGL_KERNEL_NPU) + assert KernelBackend.TORCH_NPU not in registered_backends assert spec.capabilities == frozenset({CapabilityRequirement.NPU}) assert spec.target.endswith(".forward_sgl_kernel_npu") -def test_lazy_loader_finds_stable_kernel_api(): - gemma_kernel = MagicMock() - add_gemma_kernel = MagicMock() - - with patch.dict(sys.modules, _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel)): - ops = unified_layernorm._load_sgl_kernel_npu_gemma_ops() - - assert ops == (gemma_kernel, add_gemma_kernel) - - -def test_lazy_loader_rejects_old_package_without_stable_api(): - with patch.dict(sys.modules, _fake_old_sgl_kernel_npu()): - assert unified_layernorm._load_sgl_kernel_npu_gemma_ops() is None - - -def test_new_kernel_api_is_preferred_on_npu(): +def test_sgl_kernel_npu_is_selected_on_npu(): op = GemmaRMSNormOp() x = torch.randn(2, 4) weight = torch.randn(4) - ops = (MagicMock(), MagicMock()) - with ( - patch.object( - unified_layernorm, "_load_sgl_kernel_npu_gemma_ops", return_value=ops - ), - patch.object( - fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") - ), + with patch.object( + fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") ): backend = op._resolve_backend(x, weight) assert backend is KernelBackend.SGL_KERNEL_NPU -def test_old_kernel_package_falls_back_to_torch_npu(): - op = GemmaRMSNormOp() - x = torch.randn(2, 4) - weight = torch.randn(4) - - with ( - patch.object( - unified_layernorm, "_load_sgl_kernel_npu_gemma_ops", return_value=None - ), - patch.object( - fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") - ), - ): - backend = op._resolve_backend(x, weight) - - assert backend is KernelBackend.TORCH_NPU - - -def test_non_npu_selection_does_not_import_sgl_kernel_npu(): - op = GemmaRMSNormOp() - - with ( - patch.object( - unified_layernorm, - "_load_sgl_kernel_npu_gemma_ops", - side_effect=AssertionError("NPU package import is forbidden"), - ), - patch.object( - fused_op_module, - "_platform", - return_value=PlatformInfo(device_type="cpu"), - ), - ): - assert not op.backend_eligible(KernelBackend.SGL_KERNEL_NPU) - - def test_sgl_kernel_npu_selection_does_not_query_soc(): - gemma_kernel = MagicMock() - add_gemma_kernel = MagicMock() get_soc_version = MagicMock(side_effect=AssertionError("SoC query is forbidden")) torch_npu = SimpleNamespace(npu=SimpleNamespace(get_soc_version=get_soc_version)) op = GemmaRMSNormOp() with ( - patch.dict( - sys.modules, - { - **_fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel), - "torch_npu": torch_npu, - }, - ), + patch.dict(sys.modules, {"torch_npu": torch_npu}), patch.object( fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") ), @@ -177,11 +82,7 @@ def test_sgl_kernel_npu_normal_out_contract(): out = torch.empty_like(x) gemma_kernel = MagicMock(return_value=expected) - with patch.object( - unified_layernorm, - "_load_sgl_kernel_npu_gemma_ops", - return_value=(gemma_kernel, MagicMock()), - ): + with patch.dict(sys.modules, _fake_sgl_kernel_npu(gemma_kernel, MagicMock())): result = GemmaRMSNormOp().forward_sgl_kernel_npu(x, weight, 1e-5, out=out) assert result is out @@ -197,11 +98,7 @@ def test_sgl_kernel_npu_fused_in_place_contract(): residual_sum = torch.randn_like(residual) add_gemma_kernel = MagicMock(return_value=(norm_output, residual_sum)) - with patch.object( - unified_layernorm, - "_load_sgl_kernel_npu_gemma_ops", - return_value=(MagicMock(), add_gemma_kernel), - ): + with patch.dict(sys.modules, _fake_sgl_kernel_npu(MagicMock(), add_gemma_kernel)): result = GemmaFusedAddRMSNormOp().forward_sgl_kernel_npu( x, residual, weight, 1e-5 ) @@ -217,64 +114,6 @@ def test_sgl_kernel_npu_fused_in_place_contract(): assert args[3] == 1e-5 -def test_torch_npu_normal_fallback_uses_offset_weight(): - x = torch.randn(2, 4) - weight = torch.randn(4) - fallback_kernel = MagicMock(return_value=(x, None)) - torch_npu = SimpleNamespace(npu_rms_norm=fallback_kernel) - - with patch.dict(sys.modules, {"torch_npu": torch_npu}): - result = GemmaRMSNormOp().forward_npu(x, weight) - - assert result is x - args = fallback_kernel.call_args.args - assert args[0] is x - torch.testing.assert_close(args[1], 1.0 + weight) - assert args[2] == 1e-6 - - -def test_torch_npu_fused_fallback_uses_offset_weight_and_writes_back(): - x = torch.randn(2, 4) - residual = torch.randn(2, 4) - weight = torch.randn(4) - norm_output = torch.randn_like(x) - residual_sum = torch.randn_like(residual) - fallback_kernel = MagicMock(return_value=(norm_output, None, residual_sum)) - torch_npu = SimpleNamespace(npu_add_rms_norm=fallback_kernel) - - with patch.dict(sys.modules, {"torch_npu": torch_npu}): - result = GemmaFusedAddRMSNormOp().forward_npu(x, residual, weight) - - assert result is None - torch.testing.assert_close(x, norm_output) - torch.testing.assert_close(residual, residual_sum) - args = fallback_kernel.call_args.args - assert args[0] is residual - assert args[1] is x - torch.testing.assert_close(args[2], 1.0 + weight) - assert args[3] == 1e-6 - - -def test_force_backend_and_trace_use_sgl_kernel_npu(): - x = torch.randn(2, 4) - weight = torch.randn(4) - gemma_kernel = MagicMock(return_value=x) - kernels.set_fused_op_backend(KernelBackend.SGL_KERNEL_NPU) - kernels.enable_fused_op_trace() - - with patch.object( - unified_layernorm, - "_load_sgl_kernel_npu_gemma_ops", - return_value=(gemma_kernel, MagicMock()), - ): - result = unified_layernorm.gemma_rmsnorm(x, weight) - - assert result is x - (record,) = kernels.get_fused_op_trace() - assert record.op == "layernorm.gemma_rmsnorm" - assert record.backend == "sgl_kernel_npu" - - @pytest.mark.parametrize("layer_name", ["GemmaRMSNorm", "Gemma3RMSNorm"]) def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): from sglang.srt.layers import layernorm as layernorm_module From 0650cea3a942ca9ab95729eb2b57bcf58bec8cd3 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Thu, 30 Jul 2026 17:37:57 +0800 Subject: [PATCH 03/15] refactor(npu): dispatch Gemma backend by package target --- .../sglang/kernels/ops/layernorm/__init__.py | 86 ++++++++++++- .../ops/layernorm/test_npu_gemma_rmsnorm.py | 117 ++++++++++++++++-- 2 files changed, 190 insertions(+), 13 deletions(-) diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index 5f9841b6e4d5..1dc21ef10cda 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -14,6 +14,7 @@ from __future__ import annotations +from functools import lru_cache from typing import TYPE_CHECKING, Optional from sglang.kernels.fused_op import BaseFusedOp, register_fused_op @@ -31,6 +32,22 @@ _HIP = frozenset({CapabilityRequirement.HIP}) _NPU = frozenset({CapabilityRequirement.NPU}) + +@lru_cache(maxsize=1) +def _load_sgl_kernel_npu_gemma_api(): + """Load the Gemma API only when it is present in the target-specific wheel.""" + try: + from sgl_kernel_npu.norm.gemma_rmsnorm import ( + add_gemma_rms_norm, + gemma_rms_norm, + ) + except ModuleNotFoundError as error: + if error.name == "sgl_kernel_npu.norm.gemma_rmsnorm": + return None + raise + return gemma_rms_norm, add_gemma_rms_norm + + # Unlike the gated-activation ops, sgl_kernel does *not* build the rmsnorm ops # for ROCm (production: ``if _is_cuda or _is_xpu or _is_musa: from sgl_kernel # import rmsnorm`` — HIP is absent), so AOT here is CUDA-only. ROCm instead has @@ -284,6 +301,7 @@ class GemmaRMSNormOp(BaseFusedOp): KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, KernelBackend.SGL_KERNEL_NPU: _NPU, + KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( supported_dtypes=_NORM_DTYPES, @@ -297,9 +315,19 @@ class GemmaRMSNormOp(BaseFusedOp): KernelBackend.SGL_KERNEL_NPU: ( "Gemma-style RMS normalization (sgl_kernel_npu, Ascend)." ), + KernelBackend.TORCH_NPU: ( + "Gemma-style RMS normalization (native torch_npu, Ascend)." + ), KernelBackend.TORCH: "Gemma-style RMS normalization (pure-torch reference).", } + def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: + if not super().backend_eligible(backend, *args, **kwargs): + return False + if backend is KernelBackend.SGL_KERNEL_NPU: + return _load_sgl_kernel_npu_gemma_api() is not None + return True + def forward_native( self, input: torch.Tensor, @@ -357,7 +385,12 @@ def forward_sgl_kernel_npu( out: Optional[torch.Tensor] = None, enable_pdl: Optional[bool] = None, ) -> torch.Tensor: - from sgl_kernel_npu.norm.gemma_rmsnorm import gemma_rms_norm + api = _load_sgl_kernel_npu_gemma_api() + if api is None: + raise RuntimeError( + "sgl_kernel_npu was not built with Gemma RMSNorm support" + ) + gemma_rms_norm, _ = api result = gemma_rms_norm(input, weight, eps) if out is None: @@ -365,6 +398,22 @@ def forward_sgl_kernel_npu( out.copy_(result) return out + def forward_npu( + self, + input: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + out: Optional[torch.Tensor] = None, + enable_pdl: Optional[bool] = None, + ) -> torch.Tensor: + import torch_npu + + result = torch_npu.npu_gemma_rms_norm(input, weight, eps)[0] + if out is None: + return result + out.copy_(result) + return out + class GemmaFusedAddRMSNormOp(BaseFusedOp): """In-place ``residual += input; input = GemmaRMSNorm(residual) * (weight + 1)``.""" @@ -377,6 +426,7 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, KernelBackend.SGL_KERNEL_NPU: _NPU, + KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( supported_dtypes=_NORM_DTYPES, @@ -393,11 +443,22 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): "Gemma-style fused residual-add + RMS normalization " "(sgl_kernel_npu, Ascend)." ), + KernelBackend.TORCH_NPU: ( + "Gemma-style fused residual-add + RMS normalization " + "(native torch_npu, Ascend)." + ), KernelBackend.TORCH: ( "Gemma-style fused residual-add + RMS normalization (pure-torch reference)." ), } + def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: + if not super().backend_eligible(backend, *args, **kwargs): + return False + if backend is KernelBackend.SGL_KERNEL_NPU: + return _load_sgl_kernel_npu_gemma_api() is not None + return True + def forward_native( self, input: torch.Tensor, @@ -455,12 +516,33 @@ def forward_sgl_kernel_npu( eps: float = 1e-6, enable_pdl: Optional[bool] = None, ) -> None: - from sgl_kernel_npu.norm.gemma_rmsnorm import add_gemma_rms_norm + api = _load_sgl_kernel_npu_gemma_api() + if api is None: + raise RuntimeError( + "sgl_kernel_npu was not built with Gemma RMSNorm support" + ) + _, add_gemma_rms_norm = api norm_output, residual_sum = add_gemma_rms_norm(input, weight, residual, eps) input.copy_(norm_output) residual.copy_(residual_sum) + def forward_npu( + self, + input: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + enable_pdl: Optional[bool] = None, + ) -> None: + import torch_npu + + norm_output, _, residual_sum = torch_npu.npu_add_rms_norm( + residual, input, 1.0 + weight, eps + ) + input.copy_(norm_output) + residual.copy_(residual_sum) + _RMSNORM = register_fused_op(RMSNormOp(), __name__, "_RMSNORM") _FUSED_ADD_RMSNORM = register_fused_op( diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py index 2e0bc397f6dd..71db25117d09 100644 --- a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py +++ b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py @@ -3,12 +3,14 @@ from unittest.mock import MagicMock, patch import pytest -import sglang.kernels.fused_op as fused_op_module import torch + +import sglang.kernels.fused_op as fused_op_module from sglang import kernels from sglang.kernels.ops.layernorm import ( GemmaFusedAddRMSNormOp, GemmaRMSNormOp, + _load_sgl_kernel_npu_gemma_api, ) from sglang.kernels.spec import CapabilityRequirement, KernelBackend, PlatformInfo from sglang.test.ci.ci_register import register_cpu_ci @@ -16,6 +18,13 @@ register_cpu_ci(est_time=2, suite="base-a-test-cpu") +@pytest.fixture(autouse=True) +def clear_gemma_api_cache(): + _load_sgl_kernel_npu_gemma_api.cache_clear() + yield + _load_sgl_kernel_npu_gemma_api.cache_clear() + + def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): package = ModuleType("sgl_kernel_npu") package.__path__ = [] @@ -35,28 +44,67 @@ def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): "op_name", ["layernorm.gemma_rmsnorm", "layernorm.gemma_fused_add_rmsnorm"], ) -def test_registry_exposes_sgl_kernel_npu_as_only_npu_provider(op_name): +def test_registry_exposes_kernel_and_native_npu_providers(op_name): registered_backends = {spec.backend for spec in kernels.registry.get(op_name)} spec = kernels.registry.get_backend(op_name, KernelBackend.SGL_KERNEL_NPU) + native_spec = kernels.registry.get_backend(op_name, KernelBackend.TORCH_NPU) - assert KernelBackend.TORCH_NPU not in registered_backends + assert KernelBackend.TORCH_NPU in registered_backends assert spec.capabilities == frozenset({CapabilityRequirement.NPU}) + assert native_spec.capabilities == frozenset({CapabilityRequirement.NPU}) assert spec.target.endswith(".forward_sgl_kernel_npu") + assert native_spec.target.endswith(".forward_npu") -def test_sgl_kernel_npu_is_selected_on_npu(): - op = GemmaRMSNormOp() - x = torch.randn(2, 4) - weight = torch.randn(4) - - with patch.object( - fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") +@pytest.mark.parametrize("op_cls", [GemmaRMSNormOp, GemmaFusedAddRMSNormOp]) +def test_sgl_kernel_npu_is_selected_on_npu(op_cls): + with ( + patch.object( + fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") + ), + patch( + "sglang.kernels.ops.layernorm._load_sgl_kernel_npu_gemma_api", + return_value=(MagicMock(), MagicMock()), + ), ): - backend = op._resolve_backend(x, weight) + backend = op_cls()._resolve_backend() assert backend is KernelBackend.SGL_KERNEL_NPU +@pytest.mark.parametrize("op_cls", [GemmaRMSNormOp, GemmaFusedAddRMSNormOp]) +def test_torch_npu_is_selected_when_target_wheel_has_no_gemma_api(op_cls): + with ( + patch.object( + fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") + ), + patch( + "sglang.kernels.ops.layernorm._load_sgl_kernel_npu_gemma_api", + return_value=None, + ), + ): + backend = op_cls()._resolve_backend() + + assert backend is KernelBackend.TORCH_NPU + + +def test_loader_reports_missing_target_specific_api(): + package = ModuleType("sgl_kernel_npu") + package.__path__ = [] + norm_package = ModuleType("sgl_kernel_npu.norm") + norm_package.__path__ = [] + + with patch.dict( + sys.modules, + { + "sgl_kernel_npu": package, + "sgl_kernel_npu.norm": norm_package, + }, + ): + sys.modules.pop("sgl_kernel_npu.norm.gemma_rmsnorm", None) + assert _load_sgl_kernel_npu_gemma_api() is None + + def test_sgl_kernel_npu_selection_does_not_query_soc(): get_soc_version = MagicMock(side_effect=AssertionError("SoC query is forbidden")) torch_npu = SimpleNamespace(npu=SimpleNamespace(get_soc_version=get_soc_version)) @@ -67,6 +115,10 @@ def test_sgl_kernel_npu_selection_does_not_query_soc(): patch.object( fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") ), + patch( + "sglang.kernels.ops.layernorm._load_sgl_kernel_npu_gemma_api", + return_value=(MagicMock(), MagicMock()), + ), ): assert op._resolve_backend(torch.randn(2, 4), torch.randn(4)) is ( KernelBackend.SGL_KERNEL_NPU @@ -114,6 +166,49 @@ def test_sgl_kernel_npu_fused_in_place_contract(): assert args[3] == 1e-5 +def test_torch_npu_normal_out_contract(): + x = torch.randn(2, 4) + weight = torch.randn(4) + expected = torch.randn_like(x) + out = torch.empty_like(x) + native_kernel = MagicMock(return_value=(expected, None)) + + with patch.dict( + sys.modules, + {"torch_npu": SimpleNamespace(npu_gemma_rms_norm=native_kernel)}, + ): + result = GemmaRMSNormOp().forward_npu(x, weight, 1e-5, out=out) + + assert result is out + torch.testing.assert_close(out, expected) + native_kernel.assert_called_once_with(x, weight, 1e-5) + + +def test_torch_npu_fused_in_place_contract(): + x = torch.randn(2, 4) + residual = torch.randn(2, 4) + weight = torch.randn(4) + norm_output = torch.randn_like(x) + residual_sum = torch.randn_like(residual) + native_kernel = MagicMock(return_value=(norm_output, None, residual_sum)) + + with patch.dict( + sys.modules, + {"torch_npu": SimpleNamespace(npu_add_rms_norm=native_kernel)}, + ): + result = GemmaFusedAddRMSNormOp().forward_npu(x, residual, weight, 1e-5) + + assert result is None + torch.testing.assert_close(x, norm_output) + torch.testing.assert_close(residual, residual_sum) + native_kernel.assert_called_once() + args = native_kernel.call_args.args + assert args[0] is residual + assert args[1] is x + torch.testing.assert_close(args[2], 1.0 + weight) + assert args[3] == 1e-5 + + @pytest.mark.parametrize("layer_name", ["GemmaRMSNorm", "Gemma3RMSNorm"]) def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): from sglang.srt.layers import layernorm as layernorm_module From 973a69899dd85846b5a7e7cba8e4dd0d26479cd5 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Fri, 31 Jul 2026 10:45:26 +0800 Subject: [PATCH 04/15] :recycle: refactor(norm): Delegate Gemma dispatch to kernel --- .../sglang/kernels/ops/layernorm/__init__.py | 98 +++----------- .../ops/layernorm/test_npu_gemma_rmsnorm.py | 124 +++--------------- 2 files changed, 36 insertions(+), 186 deletions(-) diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index 1dc21ef10cda..eb296786c682 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -6,7 +6,8 @@ over module-level instances; auto-selection follows the production default for the live device: AOT ``sgl_kernel`` on CUDA, ``aiter`` (or rocm-triton for gemma) on ROCm, ``sgl_kernel_npu`` or ``torch_npu`` on Ascend depending on the -operator, and the native reference otherwise. +operator, and the native reference otherwise. Gemma normalization on Ascend +always delegates its SoC-specific implementation to ``sgl_kernel_npu``. Pick a specific backend with e.g. ``_RMSNORM.forward(x, w, backend=KernelBackend.JIT)`` or globally via ``SGLANG_FORCE_FUSED_OP_BACKEND``. @@ -14,7 +15,6 @@ from __future__ import annotations -from functools import lru_cache from typing import TYPE_CHECKING, Optional from sglang.kernels.fused_op import BaseFusedOp, register_fused_op @@ -33,21 +33,6 @@ _NPU = frozenset({CapabilityRequirement.NPU}) -@lru_cache(maxsize=1) -def _load_sgl_kernel_npu_gemma_api(): - """Load the Gemma API only when it is present in the target-specific wheel.""" - try: - from sgl_kernel_npu.norm.gemma_rmsnorm import ( - add_gemma_rms_norm, - gemma_rms_norm, - ) - except ModuleNotFoundError as error: - if error.name == "sgl_kernel_npu.norm.gemma_rmsnorm": - return None - raise - return gemma_rms_norm, add_gemma_rms_norm - - # Unlike the gated-activation ops, sgl_kernel does *not* build the rmsnorm ops # for ROCm (production: ``if _is_cuda or _is_xpu or _is_musa: from sgl_kernel # import rmsnorm`` — HIP is absent), so AOT here is CUDA-only. ROCm instead has @@ -301,7 +286,6 @@ class GemmaRMSNormOp(BaseFusedOp): KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, KernelBackend.SGL_KERNEL_NPU: _NPU, - KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( supported_dtypes=_NORM_DTYPES, @@ -315,19 +299,9 @@ class GemmaRMSNormOp(BaseFusedOp): KernelBackend.SGL_KERNEL_NPU: ( "Gemma-style RMS normalization (sgl_kernel_npu, Ascend)." ), - KernelBackend.TORCH_NPU: ( - "Gemma-style RMS normalization (native torch_npu, Ascend)." - ), KernelBackend.TORCH: "Gemma-style RMS normalization (pure-torch reference).", } - def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: - if not super().backend_eligible(backend, *args, **kwargs): - return False - if backend is KernelBackend.SGL_KERNEL_NPU: - return _load_sgl_kernel_npu_gemma_api() is not None - return True - def forward_native( self, input: torch.Tensor, @@ -385,12 +359,14 @@ def forward_sgl_kernel_npu( out: Optional[torch.Tensor] = None, enable_pdl: Optional[bool] = None, ) -> torch.Tensor: - api = _load_sgl_kernel_npu_gemma_api() - if api is None: + try: + from sgl_kernel_npu.norm.gemma_rmsnorm import gemma_rms_norm + except ImportError as error: raise RuntimeError( - "sgl_kernel_npu was not built with Gemma RMSNorm support" - ) - gemma_rms_norm, _ = api + "Gemma RMSNorm on Ascend requires a target-specific " + "sgl-kernel-npu wheel that provides " + "sgl_kernel_npu.norm.gemma_rmsnorm" + ) from error result = gemma_rms_norm(input, weight, eps) if out is None: @@ -398,22 +374,6 @@ def forward_sgl_kernel_npu( out.copy_(result) return out - def forward_npu( - self, - input: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6, - out: Optional[torch.Tensor] = None, - enable_pdl: Optional[bool] = None, - ) -> torch.Tensor: - import torch_npu - - result = torch_npu.npu_gemma_rms_norm(input, weight, eps)[0] - if out is None: - return result - out.copy_(result) - return out - class GemmaFusedAddRMSNormOp(BaseFusedOp): """In-place ``residual += input; input = GemmaRMSNorm(residual) * (weight + 1)``.""" @@ -426,7 +386,6 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, KernelBackend.SGL_KERNEL_NPU: _NPU, - KernelBackend.TORCH_NPU: _NPU, } format_signature = FormatSignature( supported_dtypes=_NORM_DTYPES, @@ -443,22 +402,11 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): "Gemma-style fused residual-add + RMS normalization " "(sgl_kernel_npu, Ascend)." ), - KernelBackend.TORCH_NPU: ( - "Gemma-style fused residual-add + RMS normalization " - "(native torch_npu, Ascend)." - ), KernelBackend.TORCH: ( "Gemma-style fused residual-add + RMS normalization (pure-torch reference)." ), } - def backend_eligible(self, backend: KernelBackend, *args, **kwargs) -> bool: - if not super().backend_eligible(backend, *args, **kwargs): - return False - if backend is KernelBackend.SGL_KERNEL_NPU: - return _load_sgl_kernel_npu_gemma_api() is not None - return True - def forward_native( self, input: torch.Tensor, @@ -516,33 +464,19 @@ def forward_sgl_kernel_npu( eps: float = 1e-6, enable_pdl: Optional[bool] = None, ) -> None: - api = _load_sgl_kernel_npu_gemma_api() - if api is None: + try: + from sgl_kernel_npu.norm.gemma_rmsnorm import add_gemma_rms_norm + except ImportError as error: raise RuntimeError( - "sgl_kernel_npu was not built with Gemma RMSNorm support" - ) - _, add_gemma_rms_norm = api + "Gemma RMSNorm on Ascend requires a target-specific " + "sgl-kernel-npu wheel that provides " + "sgl_kernel_npu.norm.gemma_rmsnorm" + ) from error norm_output, residual_sum = add_gemma_rms_norm(input, weight, residual, eps) input.copy_(norm_output) residual.copy_(residual_sum) - def forward_npu( - self, - input: torch.Tensor, - residual: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6, - enable_pdl: Optional[bool] = None, - ) -> None: - import torch_npu - - norm_output, _, residual_sum = torch_npu.npu_add_rms_norm( - residual, input, 1.0 + weight, eps - ) - input.copy_(norm_output) - residual.copy_(residual_sum) - _RMSNORM = register_fused_op(RMSNormOp(), __name__, "_RMSNORM") _FUSED_ADD_RMSNORM = register_fused_op( diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py index 71db25117d09..4d9f8b69cb02 100644 --- a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py +++ b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py @@ -1,5 +1,5 @@ import sys -from types import ModuleType, SimpleNamespace +from types import ModuleType from unittest.mock import MagicMock, patch import pytest @@ -10,7 +10,6 @@ from sglang.kernels.ops.layernorm import ( GemmaFusedAddRMSNormOp, GemmaRMSNormOp, - _load_sgl_kernel_npu_gemma_api, ) from sglang.kernels.spec import CapabilityRequirement, KernelBackend, PlatformInfo from sglang.test.ci.ci_register import register_cpu_ci @@ -18,13 +17,6 @@ register_cpu_ci(est_time=2, suite="base-a-test-cpu") -@pytest.fixture(autouse=True) -def clear_gemma_api_cache(): - _load_sgl_kernel_npu_gemma_api.cache_clear() - yield - _load_sgl_kernel_npu_gemma_api.cache_clear() - - def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): package = ModuleType("sgl_kernel_npu") package.__path__ = [] @@ -44,51 +36,26 @@ def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): "op_name", ["layernorm.gemma_rmsnorm", "layernorm.gemma_fused_add_rmsnorm"], ) -def test_registry_exposes_kernel_and_native_npu_providers(op_name): +def test_registry_exposes_only_sgl_kernel_npu_provider(op_name): registered_backends = {spec.backend for spec in kernels.registry.get(op_name)} spec = kernels.registry.get_backend(op_name, KernelBackend.SGL_KERNEL_NPU) - native_spec = kernels.registry.get_backend(op_name, KernelBackend.TORCH_NPU) - assert KernelBackend.TORCH_NPU in registered_backends + assert KernelBackend.TORCH_NPU not in registered_backends assert spec.capabilities == frozenset({CapabilityRequirement.NPU}) - assert native_spec.capabilities == frozenset({CapabilityRequirement.NPU}) assert spec.target.endswith(".forward_sgl_kernel_npu") - assert native_spec.target.endswith(".forward_npu") @pytest.mark.parametrize("op_cls", [GemmaRMSNormOp, GemmaFusedAddRMSNormOp]) def test_sgl_kernel_npu_is_selected_on_npu(op_cls): - with ( - patch.object( - fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") - ), - patch( - "sglang.kernels.ops.layernorm._load_sgl_kernel_npu_gemma_api", - return_value=(MagicMock(), MagicMock()), - ), + with patch.object( + fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") ): backend = op_cls()._resolve_backend() assert backend is KernelBackend.SGL_KERNEL_NPU -@pytest.mark.parametrize("op_cls", [GemmaRMSNormOp, GemmaFusedAddRMSNormOp]) -def test_torch_npu_is_selected_when_target_wheel_has_no_gemma_api(op_cls): - with ( - patch.object( - fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") - ), - patch( - "sglang.kernels.ops.layernorm._load_sgl_kernel_npu_gemma_api", - return_value=None, - ), - ): - backend = op_cls()._resolve_backend() - - assert backend is KernelBackend.TORCH_NPU - - -def test_loader_reports_missing_target_specific_api(): +def test_missing_kernel_package_reports_actionable_error(): package = ModuleType("sgl_kernel_npu") package.__path__ = [] norm_package = ModuleType("sgl_kernel_npu.norm") @@ -102,29 +69,21 @@ def test_loader_reports_missing_target_specific_api(): }, ): sys.modules.pop("sgl_kernel_npu.norm.gemma_rmsnorm", None) - assert _load_sgl_kernel_npu_gemma_api() is None - - -def test_sgl_kernel_npu_selection_does_not_query_soc(): - get_soc_version = MagicMock(side_effect=AssertionError("SoC query is forbidden")) - torch_npu = SimpleNamespace(npu=SimpleNamespace(get_soc_version=get_soc_version)) - op = GemmaRMSNormOp() - - with ( - patch.dict(sys.modules, {"torch_npu": torch_npu}), - patch.object( - fused_op_module, "_platform", return_value=PlatformInfo(device_type="npu") - ), - patch( - "sglang.kernels.ops.layernorm._load_sgl_kernel_npu_gemma_api", - return_value=(MagicMock(), MagicMock()), - ), - ): - assert op._resolve_backend(torch.randn(2, 4), torch.randn(4)) is ( - KernelBackend.SGL_KERNEL_NPU - ) + with pytest.raises(RuntimeError, match="requires a target-specific"): + GemmaRMSNormOp().forward_sgl_kernel_npu( + torch.randn(2, 4), torch.randn(4), 1e-5 + ) - get_soc_version.assert_not_called() + +def test_incompatible_kernel_package_reports_actionable_error(): + modules = _fake_sgl_kernel_npu(MagicMock(), MagicMock()) + del modules["sgl_kernel_npu.norm.gemma_rmsnorm"].gemma_rms_norm + + with patch.dict(sys.modules, modules): + with pytest.raises(RuntimeError, match="requires a target-specific"): + GemmaRMSNormOp().forward_sgl_kernel_npu( + torch.randn(2, 4), torch.randn(4), 1e-5 + ) def test_sgl_kernel_npu_normal_out_contract(): @@ -166,49 +125,6 @@ def test_sgl_kernel_npu_fused_in_place_contract(): assert args[3] == 1e-5 -def test_torch_npu_normal_out_contract(): - x = torch.randn(2, 4) - weight = torch.randn(4) - expected = torch.randn_like(x) - out = torch.empty_like(x) - native_kernel = MagicMock(return_value=(expected, None)) - - with patch.dict( - sys.modules, - {"torch_npu": SimpleNamespace(npu_gemma_rms_norm=native_kernel)}, - ): - result = GemmaRMSNormOp().forward_npu(x, weight, 1e-5, out=out) - - assert result is out - torch.testing.assert_close(out, expected) - native_kernel.assert_called_once_with(x, weight, 1e-5) - - -def test_torch_npu_fused_in_place_contract(): - x = torch.randn(2, 4) - residual = torch.randn(2, 4) - weight = torch.randn(4) - norm_output = torch.randn_like(x) - residual_sum = torch.randn_like(residual) - native_kernel = MagicMock(return_value=(norm_output, None, residual_sum)) - - with patch.dict( - sys.modules, - {"torch_npu": SimpleNamespace(npu_add_rms_norm=native_kernel)}, - ): - result = GemmaFusedAddRMSNormOp().forward_npu(x, residual, weight, 1e-5) - - assert result is None - torch.testing.assert_close(x, norm_output) - torch.testing.assert_close(residual, residual_sum) - native_kernel.assert_called_once() - args = native_kernel.call_args.args - assert args[0] is residual - assert args[1] is x - torch.testing.assert_close(args[2], 1.0 + weight) - assert args[3] == 1e-5 - - @pytest.mark.parametrize("layer_name", ["GemmaRMSNorm", "Gemma3RMSNorm"]) def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): from sglang.srt.layers import layernorm as layernorm_module From a0dda0785dcb0212e3e871ef2cecb4d5d1401ebd Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Fri, 31 Jul 2026 11:51:27 +0800 Subject: [PATCH 05/15] :art: style(norm): merge NPU Gemma imports --- python/sglang/srt/layers/layernorm.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 1846390ea603..47a14c984b90 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -161,9 +161,9 @@ def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool: import torch_npu from sglang.kernels.ops.layernorm import ( - gemma_fused_add_rmsnorm as npu_gemma_fused_add_rmsnorm, + gemma_fused_add_rmsnorm, + gemma_rmsnorm, ) - from sglang.kernels.ops.layernorm import gemma_rmsnorm as npu_gemma_rmsnorm @lru_cache(maxsize=1) @@ -999,10 +999,10 @@ def forward_npu( if residual is not None: if post_residual_addition is not None: residual = residual + post_residual_addition - npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.variance_epsilon) + gemma_fused_add_rmsnorm(x, residual, self.weight, self.variance_epsilon) return x, residual - return npu_gemma_rmsnorm(x, self.weight, self.variance_epsilon) + return gemma_rmsnorm(x, self.weight, self.variance_epsilon) def forward_xpu( self, @@ -1097,9 +1097,9 @@ def forward_npu(self, x, residual: Optional[torch.Tensor] = None): if envs.SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM.get(): return self.forward_native(x, residual) if residual is not None: - npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.eps) + gemma_fused_add_rmsnorm(x, residual, self.weight, self.eps) return x, residual - return npu_gemma_rmsnorm(x, self.weight, self.eps) + return gemma_rmsnorm(x, self.weight, self.eps) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.eps}" From fed0a019f6cc9ecbc8a50fb93239046fade62919 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Fri, 31 Jul 2026 14:17:22 +0800 Subject: [PATCH 06/15] :bug: fix(norm): restore explicit NPU Gemma aliases --- python/sglang/srt/layers/layernorm.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 47a14c984b90..07b3da77efd2 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -161,9 +161,9 @@ def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool: import torch_npu from sglang.kernels.ops.layernorm import ( - gemma_fused_add_rmsnorm, - gemma_rmsnorm, + gemma_fused_add_rmsnorm as npu_gemma_fused_add_rmsnorm, ) + from sglang.kernels.ops.layernorm import gemma_rmsnorm as npu_gemma_rmsnorm @lru_cache(maxsize=1) @@ -999,10 +999,11 @@ def forward_npu( if residual is not None: if post_residual_addition is not None: residual = residual + post_residual_addition - gemma_fused_add_rmsnorm(x, residual, self.weight, self.variance_epsilon) + # The unified fused op updates both x and residual in place. + npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.variance_epsilon) return x, residual - return gemma_rmsnorm(x, self.weight, self.variance_epsilon) + return npu_gemma_rmsnorm(x, self.weight, self.variance_epsilon) def forward_xpu( self, @@ -1097,9 +1098,10 @@ def forward_npu(self, x, residual: Optional[torch.Tensor] = None): if envs.SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM.get(): return self.forward_native(x, residual) if residual is not None: - gemma_fused_add_rmsnorm(x, residual, self.weight, self.eps) + # The unified fused op updates both x and residual in place. + npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.eps) return x, residual - return gemma_rmsnorm(x, self.weight, self.eps) + return npu_gemma_rmsnorm(x, self.weight, self.eps) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.eps}" From 3b2840b23c79af08196561f9cf6767fc1371ca88 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Fri, 31 Jul 2026 14:25:46 +0800 Subject: [PATCH 07/15] :bug: fix(norm): preserve Gemma NPU API names --- .../sglang/kernels/ops/layernorm/__init__.py | 4 ++-- python/sglang/srt/layers/layernorm.py | 23 +++++++++--------- .../ops/layernorm/test_npu_gemma_rmsnorm.py | 24 +++++++++++-------- 3 files changed, 28 insertions(+), 23 deletions(-) diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index eb296786c682..8d7ed66b8c29 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -360,7 +360,7 @@ def forward_sgl_kernel_npu( enable_pdl: Optional[bool] = None, ) -> torch.Tensor: try: - from sgl_kernel_npu.norm.gemma_rmsnorm import gemma_rms_norm + from sgl_kernel_npu.norm.gemma_rmsnorm import npu_gemma_rms_norm except ImportError as error: raise RuntimeError( "Gemma RMSNorm on Ascend requires a target-specific " @@ -368,7 +368,7 @@ def forward_sgl_kernel_npu( "sgl_kernel_npu.norm.gemma_rmsnorm" ) from error - result = gemma_rms_norm(input, weight, eps) + result, _ = npu_gemma_rms_norm(input, weight, eps) if out is None: return result out.copy_(result) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 07b3da77efd2..396d2bb8fa84 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -160,10 +160,10 @@ def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool: if _is_npu: import torch_npu - from sglang.kernels.ops.layernorm import ( - gemma_fused_add_rmsnorm as npu_gemma_fused_add_rmsnorm, + from sgl_kernel_npu.norm.gemma_rmsnorm import ( + add_gemma_rms_norm, + npu_gemma_rms_norm, ) - from sglang.kernels.ops.layernorm import gemma_rmsnorm as npu_gemma_rmsnorm @lru_cache(maxsize=1) @@ -999,11 +999,13 @@ def forward_npu( if residual is not None: if post_residual_addition is not None: residual = residual + post_residual_addition - # The unified fused op updates both x and residual in place. - npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.variance_epsilon) - return x, residual + norm_out, residual = add_gemma_rms_norm( + x, self.weight, residual, self.variance_epsilon + ) + return norm_out, residual - return npu_gemma_rmsnorm(x, self.weight, self.variance_epsilon) + x, _ = npu_gemma_rms_norm(x, self.weight, self.variance_epsilon) + return x def forward_xpu( self, @@ -1098,10 +1100,9 @@ def forward_npu(self, x, residual: Optional[torch.Tensor] = None): if envs.SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM.get(): return self.forward_native(x, residual) if residual is not None: - # The unified fused op updates both x and residual in place. - npu_gemma_fused_add_rmsnorm(x, residual, self.weight, self.eps) - return x, residual - return npu_gemma_rmsnorm(x, self.weight, self.eps) + return add_gemma_rms_norm(x, self.weight, residual, self.eps) + output, _ = npu_gemma_rms_norm(x, self.weight, self.eps) + return output def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.eps}" diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py index 4d9f8b69cb02..a652a6a6917c 100644 --- a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py +++ b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py @@ -23,7 +23,7 @@ def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): norm_package = ModuleType("sgl_kernel_npu.norm") norm_package.__path__ = [] module = ModuleType("sgl_kernel_npu.norm.gemma_rmsnorm") - module.gemma_rms_norm = gemma_kernel + module.npu_gemma_rms_norm = gemma_kernel module.add_gemma_rms_norm = add_gemma_kernel return { "sgl_kernel_npu": package, @@ -77,7 +77,7 @@ def test_missing_kernel_package_reports_actionable_error(): def test_incompatible_kernel_package_reports_actionable_error(): modules = _fake_sgl_kernel_npu(MagicMock(), MagicMock()) - del modules["sgl_kernel_npu.norm.gemma_rmsnorm"].gemma_rms_norm + del modules["sgl_kernel_npu.norm.gemma_rmsnorm"].npu_gemma_rms_norm with patch.dict(sys.modules, modules): with pytest.raises(RuntimeError, match="requires a target-specific"): @@ -91,7 +91,7 @@ def test_sgl_kernel_npu_normal_out_contract(): weight = torch.randn(4) expected = torch.randn_like(x) out = torch.empty_like(x) - gemma_kernel = MagicMock(return_value=expected) + gemma_kernel = MagicMock(return_value=(expected, "rstd")) with patch.dict(sys.modules, _fake_sgl_kernel_npu(gemma_kernel, MagicMock())): result = GemmaRMSNormOp().forward_sgl_kernel_npu(x, weight, 1e-5, out=out) @@ -132,9 +132,11 @@ def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): layer_cls = getattr(layernorm_module, layer_name) layer = layer_cls(4) x = torch.randn(2, 4) - unified_op = MagicMock(return_value=x) + unified_op = MagicMock(return_value=(x, "rstd")) - with patch.object(layernorm_module, "npu_gemma_rmsnorm", unified_op, create=True): + with patch.object( + layernorm_module, "npu_gemma_rms_norm", unified_op, create=True + ): result = layer.forward_npu(x) assert result is x @@ -150,17 +152,19 @@ def test_srt_gemma_layers_delegate_residual_npu_path(layer_name): layer = layer_cls(4) x = torch.randn(2, 4) residual = torch.randn(2, 4) - fused_op = MagicMock() + norm_output = torch.randn_like(x) + residual_sum = torch.randn_like(residual) + fused_op = MagicMock(return_value=(norm_output, residual_sum)) with patch.object( layernorm_module, - "npu_gemma_fused_add_rmsnorm", + "add_gemma_rms_norm", fused_op, create=True, ): result = layer.forward_npu(x, residual) - assert result[0] is x - assert result[1] is residual + assert result[0] is norm_output + assert result[1] is residual_sum eps = layer.variance_epsilon if hasattr(layer, "variance_epsilon") else layer.eps - fused_op.assert_called_once_with(x, residual, layer.weight, eps) + fused_op.assert_called_once_with(x, layer.weight, residual, eps) From d4c310dd0f82dcf124887df46704aafb134d355b Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Fri, 31 Jul 2026 14:48:02 +0800 Subject: [PATCH 08/15] :bug: fix(norm): use legacy Gemma add API --- python/sglang/kernels/ops/layernorm/__init__.py | 2 +- python/sglang/srt/layers/layernorm.py | 7 ++----- .../ops/layernorm/test_npu_gemma_rmsnorm.py | 14 +++++++------- 3 files changed, 10 insertions(+), 13 deletions(-) diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index 8d7ed66b8c29..4ff967b33870 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -465,7 +465,7 @@ def forward_sgl_kernel_npu( enable_pdl: Optional[bool] = None, ) -> None: try: - from sgl_kernel_npu.norm.gemma_rmsnorm import add_gemma_rms_norm + from sgl_kernel_npu.norm.add_rmsnorm_bias import add_gemma_rms_norm except ImportError as error: raise RuntimeError( "Gemma RMSNorm on Ascend requires a target-specific " diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 396d2bb8fa84..3eaccf218df7 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -159,11 +159,8 @@ def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool: if _is_npu: import torch_npu - - from sgl_kernel_npu.norm.gemma_rmsnorm import ( - add_gemma_rms_norm, - npu_gemma_rms_norm, - ) + from sgl_kernel_npu.norm.add_rmsnorm_bias import add_gemma_rms_norm + from sgl_kernel_npu.norm.gemma_rmsnorm import npu_gemma_rms_norm @lru_cache(maxsize=1) diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py index a652a6a6917c..501b26e09422 100644 --- a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py +++ b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py @@ -22,13 +22,15 @@ def _fake_sgl_kernel_npu(gemma_kernel, add_gemma_kernel): package.__path__ = [] norm_package = ModuleType("sgl_kernel_npu.norm") norm_package.__path__ = [] - module = ModuleType("sgl_kernel_npu.norm.gemma_rmsnorm") - module.npu_gemma_rms_norm = gemma_kernel - module.add_gemma_rms_norm = add_gemma_kernel + gemma_module = ModuleType("sgl_kernel_npu.norm.gemma_rmsnorm") + gemma_module.npu_gemma_rms_norm = gemma_kernel + add_module = ModuleType("sgl_kernel_npu.norm.add_rmsnorm_bias") + add_module.add_gemma_rms_norm = add_gemma_kernel return { "sgl_kernel_npu": package, "sgl_kernel_npu.norm": norm_package, - "sgl_kernel_npu.norm.gemma_rmsnorm": module, + "sgl_kernel_npu.norm.gemma_rmsnorm": gemma_module, + "sgl_kernel_npu.norm.add_rmsnorm_bias": add_module, } @@ -134,9 +136,7 @@ def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): x = torch.randn(2, 4) unified_op = MagicMock(return_value=(x, "rstd")) - with patch.object( - layernorm_module, "npu_gemma_rms_norm", unified_op, create=True - ): + with patch.object(layernorm_module, "npu_gemma_rms_norm", unified_op, create=True): result = layer.forward_npu(x) assert result is x From 48e9fb8ce0ed23e3ad9a8912b330a236a52169f2 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Mon, 3 Aug 2026 10:47:40 +0800 Subject: [PATCH 09/15] :recycle: refactor(layernorm): route Gemma RMSNorm via sgl_kernel_npu 1. Guard the new provider import in srt/layers/layernorm.py and fall back to torch_npu, so NPU deployments on an sgl-kernel-npu wheel predating the staged provider keep importing; a hard import would break every model, not just Gemma, at module import time 2. Restore Gemma3RMSNorm.forward_npu's residual path to forward_native and drop the added env switch; both were outside the scope of replacing direct torch_npu calls and would shift numerics on all Ascend SoCs 3. Extract the duplicated lazy-import guard into _sgl_kernel_npu_gemma, which fixes the fused-add error message naming the wrong module 4. Keep GemmaFusedAddRMSNormOp on the SGL_KERNEL_NPU provenance alongside GemmaRMSNormOp 5. Pin each op's error to its own module, drop the reverted-behaviour case, and replace defensive hasattr with explicit parametrization 6. Tell Ascend 950 users in the Ascend install guide that released kernel wheels are 910-only --- .../getting-started/installation.mdx | 10 +++ .../sglang/kernels/ops/layernorm/__init__.py | 63 ++++++++-------- python/sglang/srt/layers/layernorm.py | 18 +++-- .../ops/layernorm/test_npu_gemma_rmsnorm.py | 71 +++++++++++-------- 4 files changed, 101 insertions(+), 61 deletions(-) diff --git a/docs_new/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx b/docs_new/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx index 944741cb58a8..be6fb6429add 100644 --- a/docs_new/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx +++ b/docs_new/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx @@ -168,6 +168,16 @@ For installation of Triton on Ascend nightly builds or from sources, follow [ins We provide SGL kernels for Ascend NPU, check [installation guide](https://github.com/sgl-project/sgl-kernel-npu/blob/main/python/sgl_kernel_npu/README.md). + +Released packages are built for the `910` target only. Gemma-family models +(including Qwen3.5) call `sgl_kernel_npu.norm.gemma_rmsnorm`, whose +implementation is chosen when the wheel is built — the native `torch_npu` +operator on Ascend 910B/910C, and ACLNN RMSNorm on Ascend 950, which does not +register that operator. **Ascend 950 users must therefore build the wheel +from source** with `bash build.sh -a kernels 950`; a `910` wheel fails on the +first Gemma forward. + + #### DeepEP-compatible Library We provide a DeepEP-compatible Library as a drop-in replacement of deepseek-ai's DeepEP library, check the [installation guide](https://github.com/sgl-project/sgl-kernel-npu/blob/main/python/deep_ep/README.md). diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index 4ff967b33870..ba767090e5e9 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -5,9 +5,8 @@ all behind one signature. The public module-level functions are thin wrappers over module-level instances; auto-selection follows the production default for the live device: AOT ``sgl_kernel`` on CUDA, ``aiter`` (or rocm-triton for -gemma) on ROCm, ``sgl_kernel_npu`` or ``torch_npu`` on Ascend depending on the -operator, and the native reference otherwise. Gemma normalization on Ascend -always delegates its SoC-specific implementation to ``sgl_kernel_npu``. +gemma) on ROCm, ``torch_npu`` on Ascend (``sgl_kernel_npu`` for gemma, whose +implementation is SoC-specific), native reference otherwise. Pick a specific backend with e.g. ``_RMSNORM.forward(x, w, backend=KernelBackend.JIT)`` or globally via ``SGLANG_FORCE_FUSED_OP_BACKEND``. @@ -15,6 +14,7 @@ from __future__ import annotations +import importlib from typing import TYPE_CHECKING, Optional from sglang.kernels.fused_op import BaseFusedOp, register_fused_op @@ -31,8 +31,6 @@ _CUDA = frozenset({CapabilityRequirement.CUDA}) _HIP = frozenset({CapabilityRequirement.HIP}) _NPU = frozenset({CapabilityRequirement.NPU}) - - # Unlike the gated-activation ops, sgl_kernel does *not* build the rmsnorm ops # for ROCm (production: ``if _is_cuda or _is_xpu or _is_musa: from sgl_kernel # import rmsnorm`` — HIP is absent), so AOT here is CUDA-only. ROCm instead has @@ -51,6 +49,24 @@ ) +def _sgl_kernel_npu_gemma(module: str, symbol: str): + """Resolve a Gemma kernel from ``sgl_kernel_npu``, or explain what is missing. + + The Gemma provider is picked when the sgl-kernel-npu wheel is built (native + ``torch_npu`` operator on Ascend 910, ACLNN on Ascend 950), so an older or + mismatched wheel shows up here as a plain ImportError. Unlike + ``srt/layers/layernorm.py``, this path does not fall back to ``torch_npu``: + a backend selected by name must not silently run a different provenance. + """ + try: + return getattr(importlib.import_module(module), symbol) + except (ImportError, AttributeError) as error: + raise RuntimeError( + "Gemma RMSNorm on Ascend requires a target-specific sgl-kernel-npu " + f"wheel that provides {module}.{symbol}" + ) from error + + class RMSNormOp(BaseFusedOp): """``out = (input / RMS(input)) * weight``; returns a tensor. @@ -359,15 +375,9 @@ def forward_sgl_kernel_npu( out: Optional[torch.Tensor] = None, enable_pdl: Optional[bool] = None, ) -> torch.Tensor: - try: - from sgl_kernel_npu.norm.gemma_rmsnorm import npu_gemma_rms_norm - except ImportError as error: - raise RuntimeError( - "Gemma RMSNorm on Ascend requires a target-specific " - "sgl-kernel-npu wheel that provides " - "sgl_kernel_npu.norm.gemma_rmsnorm" - ) from error - + npu_gemma_rms_norm = _sgl_kernel_npu_gemma( + "sgl_kernel_npu.norm.gemma_rmsnorm", "npu_gemma_rms_norm" + ) result, _ = npu_gemma_rms_norm(input, weight, eps) if out is None: return result @@ -381,7 +391,8 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): op = "layernorm.gemma_fused_add_rmsnorm" priority = _NORM_PRIORITY # AOT (sgl_kernel) on CUDA; JIT is the ROCm rocm-triton path on HIP. - # Ascend delegates provider selection to the stable sgl_kernel_npu API. + # NPU uses ``sgl_kernel_npu.norm.add_rmsnorm_bias.add_gemma_rms_norm`` — the + # SGL_KERNEL_NPU provenance, not torch_npu. capabilities = { KernelBackend.AOT: _CUDA, KernelBackend.JIT: _HIP, @@ -403,7 +414,8 @@ class GemmaFusedAddRMSNormOp(BaseFusedOp): "(sgl_kernel_npu, Ascend)." ), KernelBackend.TORCH: ( - "Gemma-style fused residual-add + RMS normalization (pure-torch reference)." + "Gemma-style fused residual-add + RMS normalization " + "(pure-torch reference)." ), } @@ -464,18 +476,13 @@ def forward_sgl_kernel_npu( eps: float = 1e-6, enable_pdl: Optional[bool] = None, ) -> None: - try: - from sgl_kernel_npu.norm.add_rmsnorm_bias import add_gemma_rms_norm - except ImportError as error: - raise RuntimeError( - "Gemma RMSNorm on Ascend requires a target-specific " - "sgl-kernel-npu wheel that provides " - "sgl_kernel_npu.norm.gemma_rmsnorm" - ) from error - - norm_output, residual_sum = add_gemma_rms_norm(input, weight, residual, eps) - input.copy_(norm_output) - residual.copy_(residual_sum) + add_gemma_rms_norm = _sgl_kernel_npu_gemma( + "sgl_kernel_npu.norm.add_rmsnorm_bias", "add_gemma_rms_norm" + ) + # sgl_kernel_npu returns (normed, new_residual); honor the in-place contract. + norm_out, residual_out = add_gemma_rms_norm(input, weight, residual, eps) + input.copy_(norm_out) + residual.copy_(residual_out) _RMSNORM = register_fused_op(RMSNormOp(), __name__, "_RMSNORM") diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 3eaccf218df7..d043c8528e2e 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -160,7 +160,19 @@ def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool: if _is_npu: import torch_npu from sgl_kernel_npu.norm.add_rmsnorm_bias import add_gemma_rms_norm - from sgl_kernel_npu.norm.gemma_rmsnorm import npu_gemma_rms_norm + + try: + from sgl_kernel_npu.norm.gemma_rmsnorm import npu_gemma_rms_norm + except ImportError: + # sgl-kernel-npu wheels built before the target-specific Gemma provider + # landed expose only the torch_npu operator, which is exactly what that + # provider binds to on A2/A3 — so those deployments keep working + # unchanged. On Ascend 950 the operator is unregistered and the first Gemma + # forward fails loudly, same as before this indirection existed. The + # kernels registry (sglang/kernels/ops/layernorm) deliberately does not + # fall back: a backend selected by name there must not silently run a + # different provenance. + npu_gemma_rms_norm = torch_npu.npu_gemma_rms_norm @lru_cache(maxsize=1) @@ -1094,10 +1106,8 @@ def forward_hip(self, x, residual: Optional[torch.Tensor] = None): return self.forward_native(x, residual) def forward_npu(self, x, residual: Optional[torch.Tensor] = None): - if envs.SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM.get(): - return self.forward_native(x, residual) if residual is not None: - return add_gemma_rms_norm(x, self.weight, residual, self.eps) + return self.forward_native(x, residual) output, _ = npu_gemma_rms_norm(x, self.weight, self.eps) return output diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py index 501b26e09422..fe49531503ca 100644 --- a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py +++ b/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py @@ -1,4 +1,6 @@ +import re import sys +from pathlib import Path from types import ModuleType from unittest.mock import MagicMock, patch @@ -78,14 +80,25 @@ def test_missing_kernel_package_reports_actionable_error(): def test_incompatible_kernel_package_reports_actionable_error(): + """A wheel that imports but lacks the symbol must name the module it lacks. + + Covers the AttributeError branch (the test above covers ImportError), and + pins each op to its own module so the message cannot drift into a + copy-pasted sibling name. + """ modules = _fake_sgl_kernel_npu(MagicMock(), MagicMock()) del modules["sgl_kernel_npu.norm.gemma_rmsnorm"].npu_gemma_rms_norm + del modules["sgl_kernel_npu.norm.add_rmsnorm_bias"].add_gemma_rms_norm with patch.dict(sys.modules, modules): - with pytest.raises(RuntimeError, match="requires a target-specific"): + with pytest.raises(RuntimeError, match=r"norm\.gemma_rmsnorm\."): GemmaRMSNormOp().forward_sgl_kernel_npu( torch.randn(2, 4), torch.randn(4), 1e-5 ) + with pytest.raises(RuntimeError, match=r"norm\.add_rmsnorm_bias\."): + GemmaFusedAddRMSNormOp().forward_sgl_kernel_npu( + torch.randn(2, 4), torch.randn(2, 4), torch.randn(4), 1e-5 + ) def test_sgl_kernel_npu_normal_out_contract(): @@ -127,12 +140,14 @@ def test_sgl_kernel_npu_fused_in_place_contract(): assert args[3] == 1e-5 -@pytest.mark.parametrize("layer_name", ["GemmaRMSNorm", "Gemma3RMSNorm"]) -def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): +@pytest.mark.parametrize( + ("layer_name", "eps_attr"), + [("GemmaRMSNorm", "variance_epsilon"), ("Gemma3RMSNorm", "eps")], +) +def test_srt_gemma_layers_delegate_plain_npu_path(layer_name, eps_attr): from sglang.srt.layers import layernorm as layernorm_module - layer_cls = getattr(layernorm_module, layer_name) - layer = layer_cls(4) + layer = getattr(layernorm_module, layer_name)(4) x = torch.randn(2, 4) unified_op = MagicMock(return_value=(x, "rstd")) @@ -140,31 +155,29 @@ def test_srt_gemma_layers_delegate_plain_npu_path(layer_name): result = layer.forward_npu(x) assert result is x - eps = layer.variance_epsilon if hasattr(layer, "variance_epsilon") else layer.eps - unified_op.assert_called_once_with(x, layer.weight, eps) - + unified_op.assert_called_once_with(x, layer.weight, getattr(layer, eps_attr)) -@pytest.mark.parametrize("layer_name", ["GemmaRMSNorm", "Gemma3RMSNorm"]) -def test_srt_gemma_layers_delegate_residual_npu_path(layer_name): - from sglang.srt.layers import layernorm as layernorm_module - layer_cls = getattr(layernorm_module, layer_name) - layer = layer_cls(4) - x = torch.randn(2, 4) - residual = torch.randn(2, 4) - norm_output = torch.randn_like(x) - residual_sum = torch.randn_like(residual) - fused_op = MagicMock(return_value=(norm_output, residual_sum)) +def test_srt_falls_back_to_torch_npu_on_wheels_without_the_provider(): + """The srt-layer provider import must stay guarded, with a torch_npu fallback. - with patch.object( - layernorm_module, - "add_gemma_rms_norm", - fused_op, - create=True, - ): - result = layer.forward_npu(x, residual) + ``srt/layers/layernorm.py`` imports the provider at module scope under + ``if _is_npu``, which CPU CI never executes -- hence a source-shape check. + Turning it back into a hard import would break every NPU deployment running + an sgl-kernel-npu wheel from before the staged Gemma provider, at + ``import sglang.srt.layers.layernorm`` time and for all models, not just + Gemma ones. + """ + from sglang.srt.layers import layernorm as layernorm_module - assert result[0] is norm_output - assert result[1] is residual_sum - eps = layer.variance_epsilon if hasattr(layer, "variance_epsilon") else layer.eps - fused_op.assert_called_once_with(x, layer.weight, residual, eps) + source = Path(layernorm_module.__file__).read_text(encoding="utf-8") + guarded_import = re.search( + r"try:\s*\n" + r"\s*from sgl_kernel_npu\.norm\.gemma_rmsnorm import npu_gemma_rms_norm\s*\n" + r"\s*except ImportError:\s*\n" + r"(?:\s*#.*\n)*" + r"\s*npu_gemma_rms_norm = torch_npu\.npu_gemma_rms_norm\s*\n", + source, + ) + + assert guarded_import is not None From a88e917b6489b45acaf8a703f908d5288c8bfa62 Mon Sep 17 00:00:00 2001 From: Junlin Wu Date: Thu, 17 Sep 2026 10:34:32 +0800 Subject: [PATCH 10/15] :white_check_mark: test(kernels): update NPU Gemma backend expectation --- test/registered/kernels/ops/layernorm/test_kernels_namespace.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py index ae05236314aa..474c04109ad8 100644 --- a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py +++ b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py @@ -89,7 +89,7 @@ def test_activation_default_backend(monkeypatch, device, expect): ("_RMSNORM", "npu", "torch_npu"), ("_GEMMA_RMSNORM", "cuda", "aot"), ("_GEMMA_RMSNORM", "hip", "jit"), # rocm-triton JIT pinned to HIP - ("_GEMMA_RMSNORM", "npu", "torch_npu"), + ("_GEMMA_RMSNORM", "npu", "sgl_kernel_npu"), ], ) def test_layernorm_default_backend(monkeypatch, op_attr, device, expect): From 02bc3940208d27134cd3856004fc8bef8c3deb94 Mon Sep 17 00:00:00 2001 From: Junlin Wu Date: Thu, 17 Sep 2026 11:24:49 +0800 Subject: [PATCH 11/15] :memo: docs(npu): point Ascend 950 users to the 950 kernel package --- .../ascend-npus/getting-started/installation.mdx | 16 +++++++++------- 1 file changed, 9 insertions(+), 7 deletions(-) diff --git a/docs/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx b/docs/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx index 16ca3946ceea..d805b4328eda 100644 --- a/docs/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx +++ b/docs/docs/hardware-platforms/ascend-npus/getting-started/installation.mdx @@ -253,13 +253,15 @@ For installation of Triton on Ascend nightly builds or from sources, follow [ins We provide SGL kernels for Ascend NPU, check [installation guide](https://github.com/sgl-project/sgl-kernel-npu/blob/main/python/sgl_kernel_npu/README.md). -Released packages are built for the `910` target only. Gemma-family models -(including Qwen3.5) call `sgl_kernel_npu.norm.gemma_rmsnorm`, whose -implementation is chosen when the wheel is built — the native `torch_npu` -operator on Ascend 910B/910C, and ACLNN RMSNorm on Ascend 950, which does not -register that operator. **Ascend 950 users must therefore build the wheel -from source** with `bash build.sh -a kernels 950`; a `910` wheel fails on the -first Gemma forward. +Gemma-family models (including Qwen3.5) call `sgl_kernel_npu.norm.gemma_rmsnorm`, +whose implementation is chosen when the wheel is built: the native `torch_npu` +operator for Ascend 910B/910C, and ACLNN RMSNorm for Ascend 950, which does not +register that operator. **Ascend 950 users must install the `950` release +package** (`950` appears in the asset name), or build from source with +`bash build.sh -a kernels Ascend950PR_9599`, the target that package is built +with. On Ascend 950, a `910b` or `a3` package fails on the first Gemma forward, +and so does a package from a release that predates this module; check with +`python -c "import sgl_kernel_npu.norm.gemma_rmsnorm"`. #### DeepEP-compatible Library From 74d99befdbebe9c77bb1cb357f887d87094bf8e5 Mon Sep 17 00:00:00 2001 From: Junlin Wu Date: Fri, 18 Sep 2026 10:20:10 +0800 Subject: [PATCH 12/15] :bug: fix(npu): rebind captured graph inputs before replaying replay_with_input_update rebound the captured NPU graph's seq_lens from a background thread while the main thread issued graph.replay(), joining the thread only after the replay had been sent. Nothing ordered the rebind before the execution it was meant to feed: the replay reaches the driver's BindSqCq while the rebind is still in flight, the driver refuses the inconsistent state ("Stream not inited or stream_mem not match" in plog), and rtModelExecute fails on a decode that succeeded on the previous run. Rebind on the calling thread, then replay. Reproduced on Ascend with Qwen3.5-27B: serving crashed intermittently on the first decode graph replay with rtModelExecute retCode 0x7020023, surfaced as Insufficient_Resources(EL0006) although the failure is not about memory. It hit both BF16 and ModelSlim W8A8 MXFP8, on different devices, with no other process on the card. Repeated runs with the rebind serialized no longer reproduce it, and restoring the concurrent rebind brings the failure straight back. The same stack has been reported elsewhere with retCode 0x7020004 / Invalid_Argument(EL0003). --- .../npu/graph_runner/npu_cudagraph_backend.py | 17 +++++++---------- 1 file changed, 7 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py index 105034c03c0f..1f83422d367b 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py @@ -12,7 +12,6 @@ from __future__ import annotations -import threading from contextlib import AbstractContextManager, contextmanager from functools import partial from typing import TYPE_CHECKING, Any, Callable, Dict, Optional @@ -150,8 +149,11 @@ def replay_with_input_update( attr_type: Any = None, cpu_update_input: list = None, ) -> Any: - """Rebind seq_lens on the recorded NPU graph in a background - thread, then replay. Used when the model is not deepseek-nsa. + """Rebind seq_lens on the recorded NPU graph, then replay it. + Used when the model is not deepseek-nsa. + + The rebind must land before the replay is issued: a rebind still + in flight races the driver's BindSqCq and rtModelExecute fails. Two calling conventions: 1. (legacy) seq_lens + attr_name + attr_type: @@ -166,14 +168,9 @@ def replay_with_input_update( graph = self._graphs[shape_key] - def _update(): - self._device_module.set_device(self._device_id) - graph.update(cpu_update_input=cpu_update_input) - - thread = threading.Thread(target=_update) - thread.start() + self._device_module.set_device(self._device_id) + graph.update(cpu_update_input=cpu_update_input) graph.replay() - thread.join() return self._outputs[shape_key] def cleanup(self) -> None: From 668222c7f090cb71d15d402b706a14c312fb1b78 Mon Sep 17 00:00:00 2001 From: Junlin Wu Date: Fri, 18 Sep 2026 11:06:29 +0800 Subject: [PATCH 13/15] :truck: test(npu): move the Gemma RMSNorm registry test to unit/npu The lint gate rejects new files under test/registered/kernels/ unless they register a *-kernel-* suite, and no such suite runs on CPU: the workflows define base-b-kernel-unit-test-* on GPU runners only. This test mocks sgl_kernel_npu and asserts registry dispatch and error messages, so it is a CPU unit test rather than a kernel test. The rejection failed lint, which gates pr-gate, so every test job in the matrix was skipped -- including base-a-test-cpu, where this test runs. Path and suite now agree: test/registered/unit/npu/, next to the other NPU CPU unit tests, keeping register_cpu_ci(suite="base-a-test-cpu"). All 11 tests still pass. --- .../{kernels/ops/layernorm => unit/npu}/test_npu_gemma_rmsnorm.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename test/registered/{kernels/ops/layernorm => unit/npu}/test_npu_gemma_rmsnorm.py (100%) diff --git a/test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py b/test/registered/unit/npu/test_npu_gemma_rmsnorm.py similarity index 100% rename from test/registered/kernels/ops/layernorm/test_npu_gemma_rmsnorm.py rename to test/registered/unit/npu/test_npu_gemma_rmsnorm.py From 4f35f3ef7e83d6bf2d10d10b5d3c4ddeea51e6c0 Mon Sep 17 00:00:00 2001 From: "(Messi) Junlin Wu" Date: Thu, 23 Jul 2026 11:05:27 +0800 Subject: [PATCH 14/15] :bug: fix(modelslim): preserve partial MXFP8 scales (cherry picked from commit fc9cd5bad6de985b62b992febb6362ab60196c59) (cherry picked from commit 2a248ca40381cef4ba6e5c1c4d0f21c5445ea72e) --- .../npu/quantization/linear_method_npu.py | 12 +++- .../modelslim/schemes/modelslim_mxfp8.py | 8 ++- .../quantization/test_modelslim_mxfp8.py | 61 +++++++++++++++++++ 3 files changed, 76 insertions(+), 5 deletions(-) create mode 100644 test/registered/unit/layers/quantization/test_modelslim_mxfp8.py diff --git a/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py b/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py index d7491e916a6c..bf27415a1833 100644 --- a/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py +++ b/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py @@ -2,6 +2,7 @@ from typing import TYPE_CHECKING, List, Optional, Tuple import torch +import torch.nn.functional as F from torch.nn.parameter import Parameter from sglang.srt.hardware_backend.npu.utils import NPUACLFormat, npu_format_cast @@ -200,10 +201,15 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: weight = layer.weight.data if weight.dtype == torch.float8_e4m3fn: # Offline (ModelSlim) path: weight is already MXFP8-quantised and - # layer.weight_scale holds the uint8 block scales [out, in/32]. Only - # re-layout to [in, out] / [in//64, out, 2] strided views below. + # layer.weight_scale holds one uint8 scale per 32 input elements. + # Pair the scales for the NPU kernel, padding an odd final count + # (e.g. K=4304 -> 135 scales -> 136 -> 68 pairs). n_dim, k_dim = layer.weight_scale.data.shape - scale = layer.weight_scale.data.reshape(n_dim, k_dim // 2, 2) + scale_data = layer.weight_scale.data + if k_dim % 2 != 0: + scale_data = F.pad(scale_data, (0, 1), mode="constant", value=0) + k_dim += 1 + scale = scale_data.reshape(n_dim, k_dim // 2, 2) layer.weight = Parameter(weight.transpose(0, 1), requires_grad=False) layer.weight_scale_inv = Parameter( scale.transpose(0, 1), requires_grad=False diff --git a/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8.py b/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8.py index 8c82ed93293f..d82f977bc398 100644 --- a/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8.py +++ b/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8.py @@ -60,11 +60,15 @@ def create_weights( ) layer.register_parameter("weight", weight) - # msmodelslim exports weight_scale as uint8, shape [out, in/32]. + # msmodelslim exports one uint8 scale per 32 input elements. Keep the + # final partial block: Qwen3-VL has visual projections such as K=4304, + # whose checkpoint scale shape is [out, ceil(4304/32)] = [out, 135]. # NOTE: Named "weight_scale" (not "weight_scale_inv") to match the # checkpoint key exported by msmodelslim; the kernel re-layouts it into # weight_scale_inv during process_weights_after_loading. - scale_dim = input_size_per_partition // MXFP8_BLOCK_SIZE + scale_dim = ( + input_size_per_partition + MXFP8_BLOCK_SIZE - 1 + ) // MXFP8_BLOCK_SIZE weight_scale = GroupQuantScaleParameter( data=torch.empty( (output_size_per_partition, scale_dim), diff --git a/test/registered/unit/layers/quantization/test_modelslim_mxfp8.py b/test/registered/unit/layers/quantization/test_modelslim_mxfp8.py new file mode 100644 index 000000000000..60e4ed1c734a --- /dev/null +++ b/test/registered/unit/layers/quantization/test_modelslim_mxfp8.py @@ -0,0 +1,61 @@ +"""CPU regression tests for ModelSlim MXFP8 weight-scale layouts.""" + +import unittest + +import torch + +from sglang.srt.layers.quantization.modelslim.schemes.modelslim_mxfp8 import ( + ModelSlimMXFP8Scheme, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +class TestModelSlimMXFP8ScaleLayout(CustomTestCase): + def setUp(self): + super().setUp() + self.scheme = ModelSlimMXFP8Scheme() + + def test_scale_placeholder_rounds_up_partial_block(self): + layer = torch.nn.Module() + + self.scheme.create_weights( + layer=layer, + input_size_per_partition=4304, + output_partition_sizes=[2], + input_size=4304, + output_size=2, + params_dtype=torch.bfloat16, + ) + + self.assertEqual(layer.weight_scale.shape, (2, 135)) + + def test_post_load_pads_odd_scale_count_for_pair_layout(self): + layer = torch.nn.Module() + layer.register_parameter( + "weight", + torch.nn.Parameter( + torch.empty((2, 4304), dtype=torch.float8_e4m3fn), + requires_grad=False, + ), + ) + layer.register_parameter( + "weight_scale", + torch.nn.Parameter( + torch.ones((2, 135), dtype=torch.uint8), requires_grad=False + ), + ) + layer.register_parameter("bias", None) + + self.scheme.process_weights_after_loading(layer) + + self.assertEqual(layer.weight.shape, (4304, 2)) + self.assertEqual(layer.weight_scale_inv.shape, (68, 2, 2)) + self.assertTrue(torch.all(layer.weight_scale_inv[-1, :, 1] == 0)) + self.assertFalse(hasattr(layer, "weight_scale")) + + +if __name__ == "__main__": + unittest.main() From 54b086566559f3ae65bbd8528cd9b925be7b834c Mon Sep 17 00:00:00 2001 From: Junlin Wu Date: Thu, 8 Oct 2026 14:53:22 +0800 Subject: [PATCH 15/15] :bug: fix(ci): Add Gemma test script entry --- test/registered/unit/npu/test_npu_gemma_rmsnorm.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/test/registered/unit/npu/test_npu_gemma_rmsnorm.py b/test/registered/unit/npu/test_npu_gemma_rmsnorm.py index b28873a42bc7..ae15b28a14e7 100644 --- a/test/registered/unit/npu/test_npu_gemma_rmsnorm.py +++ b/test/registered/unit/npu/test_npu_gemma_rmsnorm.py @@ -181,3 +181,7 @@ def test_srt_falls_back_to_torch_npu_on_wheels_without_the_provider(): ) assert guarded_import is not None + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"]))