Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
db34d8c
:building_construction: refactor(npu): delegate Gemma RMSNorm to kern…
TallMessiWu Jul 30, 2026
4e87d80
:recycle: refactor(npu): simplify Gemma RMSNorm provider
TallMessiWu Jul 30, 2026
0650cea
refactor(npu): dispatch Gemma backend by package target
TallMessiWu Jul 30, 2026
973a698
:recycle: refactor(norm): Delegate Gemma dispatch to kernel
TallMessiWu Jul 31, 2026
a0dda07
:art: style(norm): merge NPU Gemma imports
TallMessiWu Jul 31, 2026
fed0a01
:bug: fix(norm): restore explicit NPU Gemma aliases
TallMessiWu Jul 31, 2026
3b2840b
:bug: fix(norm): preserve Gemma NPU API names
TallMessiWu Jul 31, 2026
d4c310d
:bug: fix(norm): use legacy Gemma add API
TallMessiWu Jul 31, 2026
48e9fb8
:recycle: refactor(layernorm): route Gemma RMSNorm via sgl_kernel_npu
TallMessiWu Aug 3, 2026
62e320f
Merge remote-tracking branch 'upstream/main' into junlin_qwen3.5_dens…
TallMessiWu Aug 4, 2026
7e7b795
:twisted_rightwards_arrows: merge(npu): sync PR 32745 with upstream
TallMessiWu Aug 12, 2026
558ef51
:twisted_rightwards_arrows: merge(npu): resolve PR 32745 conflicts
TallMessiWu Aug 20, 2026
ef127a4
:twisted_rightwards_arrows: merge(npu): sync PR 32745 with upstream
TallMessiWu Sep 17, 2026
a88e917
:white_check_mark: test(kernels): update NPU Gemma backend expectation
TallMessiWu Sep 17, 2026
02bc394
:memo: docs(npu): point Ascend 950 users to the 950 kernel package
TallMessiWu Sep 17, 2026
74d99be
:bug: fix(npu): rebind captured graph inputs before replaying
TallMessiWu Sep 18, 2026
668222c
:truck: test(npu): move the Gemma RMSNorm registry test to unit/npu
TallMessiWu Sep 18, 2026
4f35f3e
:bug: fix(modelslim): preserve partial MXFP8 scales
TallMessiWu Jul 23, 2026
c34d054
:twisted_rightwards_arrows: merge(npu): sync PR 32745 with upstream
TallMessiWu Sep 20, 2026
a2d59e1
:twisted_rightwards_arrows: merge(npu): Sync PR 32745 with upstream
TallMessiWu Oct 8, 2026
54b0865
:bug: fix(ci): Add Gemma test script entry
TallMessiWu Oct 8, 2026
f3adade
:twisted_rightwards_arrows: merge(npu): Sync PR 32745 with main
TallMessiWu Oct 10, 2026
cf48b47
:twisted_rightwards_arrows: merge(npu): Sync final main and dependenc…
TallMessiWu Oct 10, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -252,6 +252,18 @@ 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).

<Note>
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"`.
</Note>

#### 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).
Expand Down
5 changes: 5 additions & 0 deletions python/sglang/kernels/fused_op.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,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_torch_npu",
}

Expand All @@ -130,6 +131,7 @@
KernelBackend.DEEPGEMM,
KernelBackend.CUTE_DSL,
KernelBackend.AITER,
KernelBackend.SGL_KERNEL_NPU,
KernelBackend.TORCH_NPU,
KernelBackend.TRITON,
KernelBackend.TORCH,
Expand Down Expand Up @@ -448,6 +450,9 @@ def forward_deepgemm(self, *args, **kwargs):
def forward_aiter(self, *args, **kwargs):
raise NotImplementedError(f"{self._op_label()}: no aiter backend")

def forward_sgl_kernel_npu(self, *args, **kwargs):
raise NotImplementedError(f"{self._op_label()}: no sgl_kernel_npu backend")

def forward_torch_npu(self, *args, **kwargs):
raise NotImplementedError(f"{self._op_label()}: no torch_npu backend")

Expand Down
69 changes: 57 additions & 12 deletions python/sglang/kernels/ops/layernorm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,14 +5,16 @@
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, ``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``.
"""

from __future__ import annotations

import importlib
from typing import TYPE_CHECKING, Optional

from sglang.kernels.fused_op import BaseFusedOp, register_fused_op
Expand All @@ -35,17 +37,36 @@
# 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,
KernelBackend.AITER,
KernelBackend.SGL_KERNEL_NPU,
KernelBackend.TORCH_NPU,
KernelBackend.TORCH,
)


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.

Expand Down Expand Up @@ -276,11 +297,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.TORCH_NPU: _NPU,
KernelBackend.SGL_KERNEL_NPU: _NPU,
}
format_signature = FormatSignature(
supported_dtypes=_NORM_DTYPES,
Expand All @@ -291,7 +312,9 @@ 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: "Gemma-style RMS normalization (pure-torch reference).",
}

Expand Down Expand Up @@ -344,17 +367,18 @@ def forward_jit(
out.copy_(result)
return out

def forward_torch_npu(
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:
import torch_npu

result = torch_npu.npu_gemma_rms_norm(input, weight, eps)[0]
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
out.copy_(result)
Expand All @@ -367,11 +391,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.
# 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,
KernelBackend.SGL_KERNEL_NPU: _NPU,
}
format_signature = FormatSignature(
supported_dtypes=_NORM_DTYPES,
Expand All @@ -384,6 +409,10 @@ 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: (
"Gemma-style fused residual-add + RMS normalization (pure-torch reference)."
),
Expand Down Expand Up @@ -438,6 +467,22 @@ 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:
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")
_FUSED_ADD_RMSNORM = register_fused_op(
Expand Down
3 changes: 2 additions & 1 deletion python/sglang/kernels/spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,8 +48,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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -204,10 +205,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
Expand Down
17 changes: 15 additions & 2 deletions python/sglang/srt/layers/layernorm.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,19 @@ def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool:
import torch_npu
from sgl_kernel_npu.norm.add_rmsnorm_bias import add_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

_NPU_GEMMA_RMS_NORM_TRITON_MAX_HIDDEN_SIZE = 5120


Expand Down Expand Up @@ -1248,7 +1261,7 @@ def forward_npu(
)
return norm_out, residual

x, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.variance_epsilon)
x, _ = npu_gemma_rms_norm(x, self.weight, self.variance_epsilon)
return x

def forward_xpu(
Expand Down Expand Up @@ -1370,7 +1383,7 @@ def forward_hip(self, x, residual: Optional[torch.Tensor] = None):
def forward_npu(self, x, residual: Optional[torch.Tensor] = None):
if residual is not None:
return self.forward_native(x, residual)
output, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.eps)
output, _ = npu_gemma_rms_norm(x, self.weight, self.eps)
return output

def extra_repr(self):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
2 changes: 1 addition & 1 deletion test/registered/unit/kernels/test_kernels_namespace.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,7 +91,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):
Expand Down
61 changes: 61 additions & 0 deletions test/registered/unit/layers/quantization/test_modelslim_mxfp8.py
Original file line number Diff line number Diff line change
@@ -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()
Loading
Loading