diff --git a/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md b/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md
index c15542dca..606b28818 100644
--- a/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md
+++ b/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md
@@ -50,7 +50,7 @@ in-flight row as shipped.
| AITER allreduce fusion | ROCm all-reduce plus RMSNorm still split | `python/sglang/srt/layers/layernorm.py::forward_with_allreduce_fusion`
`python/sglang/srt/distributed/communication_op.py::tensor_model_parallel_fused_allreduce_rmsnorm`
`python/sglang/srt/layers/communicator.py::apply_aiter_all_reduce_fusion` | ROCm-side fused TP all-reduce + RMSNorm with fallback to plain all-reduce plus norm | On AMD, rule out existing AITER fusion before proposing a new communication fusion. |
| Fused activation-and-mul (`SwiGLU` / `GeGLU`) | `silu_and_mul`
`gelu_and_mul`
`npu_swiglu` | `python/sglang/srt/layers/activation.py` | Single op covers activation plus elementwise multiply across CUDA / CPU / NPU / XPU backends | Treat separate activation + mul on packed MLP outputs as missing existing fusion. |
| Fused dual residual RMSNorm | residual add plus two RMSNorm-like kernels around Grok blocks | `python/sglang/srt/layers/elementwise.py::fused_dual_residual_rmsnorm`
`python/sglang/srt/models/grok.py` | One Triton kernel computes intermediate residual update and next RMSNorm output together | On Grok-like residual layouts, treat split residual + norm as missing existing fusion. |
-| In-place QK RMSNorm | split `q_norm` / `k_norm` kernels | `python/sglang/srt/models/utils.py::apply_qk_norm`
`python/sglang/kernels/ops/layernorm/_jit_norm.py::fused_inplace_qknorm` | In-place JIT QK norm plus optional `alt_stream` overlap for K | Check shape, dtype, deterministic mode, and in-place legality before proposing a new QK fuse. |
+| In-place QK RMSNorm | split `q_norm` / `k_norm` kernels | `python/sglang/srt/models/utils.py::apply_qk_norm`
`python/sglang/kernels/ops/layernorm/norm.py::fused_inplace_qknorm` | In-place JIT QK norm plus optional `alt_stream` overlap for K | Check shape, dtype, deterministic mode, and in-place legality before proposing a new QK fuse. |
| TorchInductor horizontal Q/K norm combo-kernels | `combo_kernels`
`benchmark_combo_kernel`
`q_norm`
`k_norm`
`split_with_sizes` | `torch._inductor.config.combo_kernels` | TorchInductor can horizontally fuse sibling Q-norm and K-norm kernels in compiled traces, often deleting `split_with_sizes` / `clone` ladders | Treat separate Q/K norm ladders in compile-heavy traces as an existing compiler-fusion family first. |
| MiniMax TP fused QK RMSNorm | `MiniMaxM2RMSNormTP`
`rms_sumsq_serial`
`rms_apply_serial`
`forward_qk` | `python/sglang/srt/models/minimax_m2.py` | Triton kernels compute Q / K sumsq together, TP all-reduces shared stats, then apply both RMSNorms together | On MiniMax traces, separate Q norm and K norm are usually a missed model-specific Triton fusion. |
| Fused QK RMSNorm + RoPE | `qknorm*` + `rope*` + `rotary*` as separate steps | `python/sglang/kernels/ops/attention/fused_qknorm_rope.py`
`python/sglang/srt/models/qwen3_moe.py` | One JIT kernel applies QK RMSNorm and RoPE in-place on packed QKV | For compatible LLMs, classify split QK norm + RoPE as a missing existing fusion. |
diff --git a/.claude/skills/llm-torch-profiler-analysis/scripts/triage_kernel_helpers.py b/.claude/skills/llm-torch-profiler-analysis/scripts/triage_kernel_helpers.py
index 11f3a06b7..cb410a1d8 100644
--- a/.claude/skills/llm-torch-profiler-analysis/scripts/triage_kernel_helpers.py
+++ b/.claude/skills/llm-torch-profiler-analysis/scripts/triage_kernel_helpers.py
@@ -453,7 +453,7 @@ FUSION_PATTERN_REGISTRY: Tuple[FusionPatternSpec, ...] = (
pattern="In-place QK RMSNorm",
candidate_path=(
"python/sglang/srt/models/utils.py"
- "
python/sglang/kernels/ops/layernorm/_jit_norm.py"
+ "
python/sglang/kernels/ops/layernorm/norm.py"
),
active_keywords=("fused_inplace_qknorm", "minimaxm2rmsnormtp"),
split_groups=(("apply_qk_norm", "q_norm", "k_norm", "qknorm"),),
diff --git a/benchmark/kernels/bench_fused_gate_sigmoid_mul_add.py b/benchmark/kernels/bench_fused_gate_sigmoid_mul_add.py
index be42c6834..83d103212 100644
--- a/benchmark/kernels/bench_fused_gate_sigmoid_mul_add.py
+++ b/benchmark/kernels/bench_fused_gate_sigmoid_mul_add.py
@@ -7,7 +7,7 @@ over the Qwen3.5 MoE target hidden size.
import torch
import triton
-from sglang.kernels.ops.layernorm.elementwise import fused_gate_sigmoid_mul_add
+from sglang.kernels.ops.elementwise.elementwise import fused_gate_sigmoid_mul_add
HIDDEN_DIMS = [4096]
diff --git a/benchmark/kernels/bench_fused_sigmoid_mul.py b/benchmark/kernels/bench_fused_sigmoid_mul.py
index b1d0c14ab..c8fec53bf 100644
--- a/benchmark/kernels/bench_fused_sigmoid_mul.py
+++ b/benchmark/kernels/bench_fused_sigmoid_mul.py
@@ -10,7 +10,7 @@ a fair comparison — the reshape/contiguous cost is included.
import torch
import triton
-from sglang.kernels.ops.layernorm.elementwise import fused_sigmoid_mul
+from sglang.kernels.ops.elementwise.elementwise import fused_sigmoid_mul
NUM_HEADS = 32
HEAD_DIM = 256
diff --git a/python/sglang/kernels/README.md b/python/sglang/kernels/README.md
index 7b69d06e7..598d1ddaa 100644
--- a/python/sglang/kernels/README.md
+++ b/python/sglang/kernels/README.md
@@ -25,9 +25,9 @@ sglang/kernels/
```
Operator groups (all populated): `activation`, `attention`, `communication`,
-`diffusion`, `embeddings`, `gemm`, `grammar`, `kv_canary`, `kvcache`,
-`layernorm`, `lplb`, `mamba`, `memory`, `model`, `moe`, `quantization`,
-`sampling`, `spatial`, `speculative`.
+`diffusion`, `elementwise`, `embeddings`, `gemm`, `grammar`, `kv_canary`,
+`kvcache`, `layernorm`, `lplb`, `mamba`, `memory`, `moe`, `quantization`,
+`sampling`, `speculative`.
As of the RFC #29630 finale (#32072) the legacy `sglang.jit_kernel` package has
been **removed**: its shared build/runtime infra moved to `sglang.kernels.jit`
diff --git a/python/sglang/kernels/ops/__init__.py b/python/sglang/kernels/ops/__init__.py
index 307b02402..9574dc847 100644
--- a/python/sglang/kernels/ops/__init__.py
+++ b/python/sglang/kernels/ops/__init__.py
@@ -19,6 +19,7 @@ _GROUPS = (
"attention",
"communication",
"diffusion",
+ "elementwise",
"embeddings",
"gemm",
"grammar",
@@ -29,11 +30,9 @@ _GROUPS = (
"moe",
"quantization",
"sampling",
- "spatial",
"speculative",
"lplb",
"kv_canary",
- "model",
)
for _group in _GROUPS:
diff --git a/python/sglang/kernels/ops/activation/__init__.py b/python/sglang/kernels/ops/activation/__init__.py
index fe5248c3c..39fb60ba9 100644
--- a/python/sglang/kernels/ops/activation/__init__.py
+++ b/python/sglang/kernels/ops/activation/__init__.py
@@ -31,7 +31,7 @@ _HIP = frozenset({CapabilityRequirement.HIP})
# — the canonical OR-semantics case that a device-baked backend name couldn't.
_CUDA_HIP = frozenset({CapabilityRequirement.CUDA, CapabilityRequirement.HIP})
# JIT before AOT to match the production path (srt/layers/activation.py imports
-# from sglang.kernels.ops.activation._jit_activation on CUDA); auto-selection must not invert it.
+# from sglang.kernels.ops.activation.activation on CUDA); auto-selection must not invert it.
_ACT_PRIORITY = (
KernelBackend.JIT,
KernelBackend.AOT,
@@ -82,7 +82,7 @@ class _GatedActivationOp(BaseFusedOp):
expert_ids: Optional[torch.Tensor] = None,
expert_step: int = 1,
) -> torch.Tensor:
- import sglang.kernels.ops.activation._jit_activation as jit_activation
+ import sglang.kernels.ops.activation.activation as jit_activation
return getattr(jit_activation, self.kernel_attr)(
input, out, expert_ids, expert_step
@@ -183,7 +183,7 @@ class GeluTanhAndMulOp(_GatedActivationOp):
class ReLU2Op(BaseFusedOp):
"""``out = relu(input) ** 2`` (single-input, not gated).
- The real kernel is the CUDA JIT path (``sglang.kernels.ops.activation._jit_activation.relu2``,
+ The real kernel is the CUDA JIT path (``sglang.kernels.ops.activation.activation.relu2``,
used in production on CUDA); elsewhere the torch reference runs.
"""
@@ -214,7 +214,7 @@ class ReLU2Op(BaseFusedOp):
def forward_jit(
self, input: torch.Tensor, out: Optional[torch.Tensor] = None
) -> torch.Tensor:
- from sglang.kernels.ops.activation._jit_activation import relu2
+ from sglang.kernels.ops.activation.activation import relu2
result = relu2(input)
if out is None:
diff --git a/python/sglang/kernels/ops/activation/_jit_activation.py b/python/sglang/kernels/ops/activation/activation.py
similarity index 95%
rename from python/sglang/kernels/ops/activation/_jit_activation.py
rename to python/sglang/kernels/ops/activation/activation.py
index 7f39a89b0..481f74a9c 100644
--- a/python/sglang/kernels/ops/activation/_jit_activation.py
+++ b/python/sglang/kernels/ops/activation/activation.py
@@ -29,7 +29,7 @@ def _fast_math_flags() -> list[str]:
@cache_once
-def _jit_activation_module(dtype: torch.dtype) -> Module:
+def activation_module(dtype: torch.dtype) -> Module:
args = make_cpp_args(dtype, is_arch_support_pdl())
return load_jit(
"activation",
@@ -59,7 +59,7 @@ def _run_activation_inplace(
op_name: str, input: torch.Tensor, out: torch.Tensor
) -> None:
hidden_size = input.shape[-1] // 2
- module = _jit_activation_module(input.dtype)
+ module = activation_module(input.dtype)
input_2d = input.view(-1, hidden_size * 2)
out_2d = out.view(-1, hidden_size)
module.run_activation(input_2d, out_2d, op_name)
@@ -74,7 +74,7 @@ def _run_activation_filtered_inplace(
expert_step: int,
) -> None:
hidden_size = input.shape[-1] // 2
- module = _jit_activation_module(input.dtype)
+ module = activation_module(input.dtype)
input_2d = input.view(-1, hidden_size * 2)
out_2d = out.view(-1, hidden_size)
module.run_activation_filtered(input_2d, out_2d, expert_ids, expert_step, op_name)
@@ -110,7 +110,7 @@ def _run_unary_activation_inplace(
op_name: str, input: torch.Tensor, out: torch.Tensor
) -> None:
last = input.shape[-1]
- module = _jit_activation_module(input.dtype)
+ module = activation_module(input.dtype)
module.run_unary_activation(input.view(-1, last), out.view(-1, last), op_name)
diff --git a/python/sglang/kernels/ops/attention/dsv4/compress_old.py b/python/sglang/kernels/ops/attention/dsv4/compress_old.py
index 515f900b3..7af950b52 100644
--- a/python/sglang/kernels/ops/attention/dsv4/compress_old.py
+++ b/python/sglang/kernels/ops/attention/dsv4/compress_old.py
@@ -57,7 +57,7 @@ def _jit_compress_128_online_module(head_dim: int) -> Module:
@cache_once
-def _jit_norm_rope_module(
+def norm_rope_module(
dtype: torch.dtype,
head_dim: int,
rope_dim: int,
@@ -276,7 +276,7 @@ def compress_fused_norm_rope_inplace(
plan: Union[CompressorDecodePlan, CompressorPrefillPlan],
) -> None:
freq_cis = torch.view_as_real(freq_cis).flatten(-2)
- module = _jit_norm_rope_module(kv.dtype, kv.shape[-1], freq_cis.shape[-1])
+ module = norm_rope_module(kv.dtype, kv.shape[-1], freq_cis.shape[-1])
module.forward(
kv,
weight,
@@ -296,7 +296,7 @@ def fused_norm_rope_inplace(
positions: torch.Tensor,
) -> None:
freq_cis = torch.view_as_real(freq_cis).flatten(-2)
- module = _jit_norm_rope_module(kv.dtype, kv.shape[-1], freq_cis.shape[-1])
+ module = norm_rope_module(kv.dtype, kv.shape[-1], freq_cis.shape[-1])
module.forward(
kv,
weight,
diff --git a/python/sglang/kernels/ops/model/inkling/inkling_attn_prologue.py b/python/sglang/kernels/ops/attention/inkling_attn_prologue.py
similarity index 100%
rename from python/sglang/kernels/ops/model/inkling/inkling_attn_prologue.py
rename to python/sglang/kernels/ops/attention/inkling_attn_prologue.py
diff --git a/python/sglang/kernels/ops/model/inkling/inkling_rel_proj.py b/python/sglang/kernels/ops/attention/inkling_rel_proj.py
similarity index 100%
rename from python/sglang/kernels/ops/model/inkling/inkling_rel_proj.py
rename to python/sglang/kernels/ops/attention/inkling_rel_proj.py
diff --git a/python/sglang/kernels/ops/model/inkling/inkling_row_scale.py b/python/sglang/kernels/ops/attention/inkling_row_scale.py
similarity index 100%
rename from python/sglang/kernels/ops/model/inkling/inkling_row_scale.py
rename to python/sglang/kernels/ops/attention/inkling_row_scale.py
diff --git a/python/sglang/kernels/ops/attention/log_scaling_tau.py b/python/sglang/kernels/ops/attention/log_scaling_tau.py
index f0d191da9..7502421da 100644
--- a/python/sglang/kernels/ops/attention/log_scaling_tau.py
+++ b/python/sglang/kernels/ops/attention/log_scaling_tau.py
@@ -50,7 +50,7 @@ def apply_log_scaling_tau(x: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
# Vectorized JIT kernel (16B loads, one row divide per vector) --
# bit-identical output (same fp32-mul + bf16-round), ~2-3x the
# scalar triton kernel below at every size.
- from sglang.kernels.ops.model.inkling.inkling_row_scale import row_scale_bf16
+ from sglang.kernels.ops.attention.inkling_row_scale import row_scale_bf16
x2d = torch.as_strided(x, (rows, inner), (x.stride(0), 1))
return row_scale_bf16(x2d, tau.reshape(rows).float()).view(x.shape)
diff --git a/python/sglang/kernels/ops/model/inkling/inkling_all_reduce.py b/python/sglang/kernels/ops/communication/inkling_all_reduce.py
similarity index 100%
rename from python/sglang/kernels/ops/model/inkling/inkling_all_reduce.py
rename to python/sglang/kernels/ops/communication/inkling_all_reduce.py
diff --git a/python/sglang/kernels/ops/model/inkling/inkling_ar_fused.py b/python/sglang/kernels/ops/communication/inkling_ar_fused.py
similarity index 100%
rename from python/sglang/kernels/ops/model/inkling/inkling_ar_fused.py
rename to python/sglang/kernels/ops/communication/inkling_ar_fused.py
diff --git a/python/sglang/kernels/ops/model/inkling/inkling_ar_scattered_sconv.py b/python/sglang/kernels/ops/communication/inkling_ar_scattered_sconv.py
similarity index 100%
rename from python/sglang/kernels/ops/model/inkling/inkling_ar_scattered_sconv.py
rename to python/sglang/kernels/ops/communication/inkling_ar_scattered_sconv.py
diff --git a/python/sglang/kernels/ops/diffusion/norm_scale_shift_native.py b/python/sglang/kernels/ops/diffusion/norm_scale_shift_native.py
index fc093e0a7..30ff50ced 100644
--- a/python/sglang/kernels/ops/diffusion/norm_scale_shift_native.py
+++ b/python/sglang/kernels/ops/diffusion/norm_scale_shift_native.py
@@ -59,7 +59,7 @@ def _row_bf16(t, device: torch.device):
@cache_once
-def _jit_norm_scale_shift_module() -> Module:
+def norm_scale_shift_module() -> Module:
return load_jit(
"qwen_image_norm_scale_shift_native",
cuda_files=["diffusion/norm_scale_shift.cuh"],
@@ -77,7 +77,7 @@ def _jit_norm_scale_shift_module() -> Module:
)
-_module = _jit_norm_scale_shift_module
+_module = norm_scale_shift_module
def try_fused_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps):
diff --git a/python/sglang/kernels/ops/elementwise/__init__.py b/python/sglang/kernels/ops/elementwise/__init__.py
new file mode 100644
index 000000000..d895402cc
--- /dev/null
+++ b/python/sglang/kernels/ops/elementwise/__init__.py
@@ -0,0 +1,11 @@
+"""Generic elementwise / fused-pointwise kernels.
+
+Home for cross-cutting pointwise kernels that do not belong to a single
+functional group: the fused-pointwise Triton collection (``elementwise``:
+softcap, sigmoid-mul, gated-activation and fused-rmsnorm variants shared
+across models) and the ``add_constant`` JIT reference kernel used by the
+developer guide. Individual functions register (or are imported) under the
+functional op id they logically belong to.
+"""
+
+__all__ = []
diff --git a/python/sglang/kernels/ops/attention/add_constant.py b/python/sglang/kernels/ops/elementwise/add_constant.py
similarity index 100%
rename from python/sglang/kernels/ops/attention/add_constant.py
rename to python/sglang/kernels/ops/elementwise/add_constant.py
diff --git a/python/sglang/kernels/ops/layernorm/elementwise.py b/python/sglang/kernels/ops/elementwise/elementwise.py
similarity index 100%
rename from python/sglang/kernels/ops/layernorm/elementwise.py
rename to python/sglang/kernels/ops/elementwise/elementwise.py
diff --git a/python/sglang/kernels/ops/gemm/__init__.py b/python/sglang/kernels/ops/gemm/__init__.py
index 5fafd7652..b6670d8bd 100644
--- a/python/sglang/kernels/ops/gemm/__init__.py
+++ b/python/sglang/kernels/ops/gemm/__init__.py
@@ -59,7 +59,7 @@ register_kernel(
KernelSpec(
op="gemm.dsv3_fused_a_gemm",
backend=KernelBackend.JIT,
- target="sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm:dsv3_fused_a_gemm",
+ target="sglang.kernels.ops.gemm.dsv3_fused_a_gemm:dsv3_fused_a_gemm",
capabilities=_CUDA,
format_signature=FormatSignature(
supported_dtypes=("bfloat16",),
@@ -72,7 +72,7 @@ register_kernel(
KernelSpec(
op="gemm.dsv3_router_gemm",
backend=KernelBackend.JIT,
- target="sglang.kernels.ops.gemm._jit_dsv3_router_gemm:dsv3_router_gemm",
+ target="sglang.kernels.ops.gemm.dsv3_router_gemm:dsv3_router_gemm",
capabilities=_CUDA,
format_signature=FormatSignature(
supported_dtypes=("bfloat16",),
diff --git a/python/sglang/kernels/ops/gemm/_jit_dsv3_fused_a_gemm.py b/python/sglang/kernels/ops/gemm/dsv3_fused_a_gemm.py
similarity index 95%
rename from python/sglang/kernels/ops/gemm/_jit_dsv3_fused_a_gemm.py
rename to python/sglang/kernels/ops/gemm/dsv3_fused_a_gemm.py
index 0cef3cbcf..7652d1b0d 100644
--- a/python/sglang/kernels/ops/gemm/_jit_dsv3_fused_a_gemm.py
+++ b/python/sglang/kernels/ops/gemm/dsv3_fused_a_gemm.py
@@ -25,7 +25,7 @@ if TYPE_CHECKING:
@cache_once
-def _jit_dsv3_fused_a_gemm_module(hd_in: int, hd_out: int, use_pdl: bool) -> Module:
+def dsv3_fused_a_gemm_module(hd_in: int, hd_out: int, use_pdl: bool) -> Module:
args = make_cpp_args(hd_in, hd_out, use_pdl)
return load_jit(
"dsv3_fused_a_gemm",
@@ -44,7 +44,7 @@ def _dsv3_fused_a_gemm_run(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Te
device=mat_a.device,
dtype=mat_a.dtype,
)
- module = _jit_dsv3_fused_a_gemm_module(
+ module = dsv3_fused_a_gemm_module(
mat_a.shape[1], mat_b.shape[1], is_arch_support_pdl()
)
module.dsv3_fused_a_gemm(mat_a, mat_b, output)
diff --git a/python/sglang/kernels/ops/gemm/_jit_dsv3_router_gemm.py b/python/sglang/kernels/ops/gemm/dsv3_router_gemm.py
similarity index 97%
rename from python/sglang/kernels/ops/gemm/_jit_dsv3_router_gemm.py
rename to python/sglang/kernels/ops/gemm/dsv3_router_gemm.py
index 7279f1474..05ed46d72 100644
--- a/python/sglang/kernels/ops/gemm/_jit_dsv3_router_gemm.py
+++ b/python/sglang/kernels/ops/gemm/dsv3_router_gemm.py
@@ -25,7 +25,7 @@ if TYPE_CHECKING:
@cache_once
-def _jit_dsv3_router_gemm_module(
+def dsv3_router_gemm_module(
num_experts: int,
hidden_dim: int,
use_pdl: bool,
@@ -54,7 +54,7 @@ def _dsv3_router_gemm_custom_op(
num_experts = router_weights.shape[0]
hidden_dim = hidden_states.shape[1]
out_float = output.dtype == torch.float32
- module = _jit_dsv3_router_gemm_module(
+ module = dsv3_router_gemm_module(
num_experts, hidden_dim, is_arch_support_pdl(), out_float
)
module.dsv3_router_gemm(hidden_states, router_weights, output)
diff --git a/python/sglang/kernels/ops/gemm/fused_a_gemm.py b/python/sglang/kernels/ops/gemm/fused_a_gemm.py
index 4671e48d9..917fc24e1 100644
--- a/python/sglang/kernels/ops/gemm/fused_a_gemm.py
+++ b/python/sglang/kernels/ops/gemm/fused_a_gemm.py
@@ -2,7 +2,7 @@
Dispatches to one of two interchangeable implementations via ``backend``:
-- ``"jit"``: runtime-compiled CUDA C++ (``sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm``).
+- ``"jit"``: runtime-compiled CUDA C++ (``sglang.kernels.ops.gemm.dsv3_fused_a_gemm``).
- ``"cutedsl"``: CuTe DSL (``sglang.kernels.ops.gemm.cutedsl_dsv3_fused_a_gemm``).
- ``"auto"``: CuTe DSL on SM120+, otherwise the JIT kernel.
@@ -69,9 +69,7 @@ def dsv3_fused_a_gemm(
backend = _AUTO_BACKEND
if backend == FusedAGemmBackend.JIT:
- from sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm import (
- dsv3_fused_a_gemm as impl,
- )
+ from sglang.kernels.ops.gemm.dsv3_fused_a_gemm import dsv3_fused_a_gemm as impl
else:
from sglang.kernels.ops.gemm.cutedsl_dsv3_fused_a_gemm import (
dsv3_fused_a_gemm as impl,
diff --git a/python/sglang/kernels/ops/kvcache/mla_buffer.py b/python/sglang/kernels/ops/kvcache/mla_buffer.py
index 7249e9d9b..ea179daa3 100644
--- a/python/sglang/kernels/ops/kvcache/mla_buffer.py
+++ b/python/sglang/kernels/ops/kvcache/mla_buffer.py
@@ -116,10 +116,10 @@ def set_mla_kv_buffer_triton(
Name retained for caller compatibility; the implementation is no longer
Triton-only.
"""
- from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import (
+ from sglang.kernels.ops.kvcache.set_mla_kv_buffer import (
can_use_set_mla_kv_buffer,
)
- from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import (
+ from sglang.kernels.ops.kvcache.set_mla_kv_buffer import (
set_mla_kv_buffer as jit_set_mla_kv_buffer,
)
diff --git a/python/sglang/kernels/ops/kvcache/_jit_set_mla_kv_buffer.py b/python/sglang/kernels/ops/kvcache/set_mla_kv_buffer.py
similarity index 93%
rename from python/sglang/kernels/ops/kvcache/_jit_set_mla_kv_buffer.py
rename to python/sglang/kernels/ops/kvcache/set_mla_kv_buffer.py
index ad8cc879c..da462a780 100644
--- a/python/sglang/kernels/ops/kvcache/_jit_set_mla_kv_buffer.py
+++ b/python/sglang/kernels/ops/kvcache/set_mla_kv_buffer.py
@@ -27,9 +27,7 @@ logger = logging.getLogger(__name__)
@cache_once
-def _jit_set_mla_kv_buffer_module(
- nope_bytes: int, rope_bytes: int, use_pdl: bool
-) -> Module:
+def set_mla_kv_buffer_module(nope_bytes: int, rope_bytes: int, use_pdl: bool) -> Module:
args = make_cpp_args(nope_bytes, rope_bytes, use_pdl)
return load_jit(
f"set_mla_kv_buffer_{nope_bytes}_{rope_bytes}",
@@ -66,7 +64,7 @@ def can_use_set_mla_kv_buffer(nope_bytes: int, rope_bytes: int) -> bool:
)
return False
try:
- _jit_set_mla_kv_buffer_module(nope_bytes, rope_bytes, is_arch_support_pdl())
+ set_mla_kv_buffer_module(nope_bytes, rope_bytes, is_arch_support_pdl())
return True
except Exception as e: # pragma: no cover - compile-time only
logger.warning(
@@ -115,7 +113,5 @@ def set_mla_kv_buffer(
if num_warps <= 0:
num_warps = _pick_num_warps(n_loc)
- module = _jit_set_mla_kv_buffer_module(
- nope_bytes, rope_bytes, is_arch_support_pdl()
- )
+ module = set_mla_kv_buffer_module(nope_bytes, rope_bytes, is_arch_support_pdl())
module.set_mla_kv_buffer(buf, loc, src_nope, src_rope, num_warps)
diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py
index 8a70c0bd1..2e04ccdbe 100644
--- a/python/sglang/kernels/ops/layernorm/__init__.py
+++ b/python/sglang/kernels/ops/layernorm/__init__.py
@@ -113,7 +113,7 @@ class RMSNormOp(BaseFusedOp):
) -> torch.Tensor:
import torch
- from sglang.kernels.ops.layernorm._jit_norm import rmsnorm as jit_rmsnorm
+ from sglang.kernels.ops.layernorm.norm import rmsnorm as jit_rmsnorm
if out is None:
out = torch.empty_like(input)
@@ -227,7 +227,7 @@ class FusedAddRMSNormOp(BaseFusedOp):
eps: float = 1e-6,
enable_pdl: Optional[bool] = None,
) -> None:
- from sglang.kernels.ops.layernorm._jit_norm import (
+ from sglang.kernels.ops.layernorm.norm import (
fused_add_rmsnorm as jit_fused_add_rmsnorm,
)
@@ -512,8 +512,6 @@ from sglang.kernels.spec import KernelSpec
# Triton / TileLang kernels migrated from srt/layers top-level strays
# (RFC #29630, Phase 2.5); registered for inventory.
_PHASE25_KERNELS = [
- ("elementwise", "fused_dual_residual_rmsnorm", "triton"),
- ("elementwise", "fused_rmsnorm", "triton"),
("gemma4_fused_ops", "gemma4_fused_routing", "triton"),
("gemma4_fused_ops", "gemma_qkv_rmsnorm", "triton"),
("mhc_head", "fused_hc_head", "triton"),
@@ -527,3 +525,15 @@ for _mod, _fn, _bk in _PHASE25_KERNELS:
)
)
del _mod, _fn, _bk
+
+# The fused-rmsnorm variants physically live in the shared fused-pointwise
+# collection (sglang.kernels.ops.elementwise.elementwise) but stay layernorm ops.
+for _fn in ("fused_dual_residual_rmsnorm", "fused_rmsnorm"):
+ register_kernel(
+ KernelSpec(
+ op=f"layernorm.{_fn}",
+ backend=KernelBackend.TRITON,
+ target=f"sglang.kernels.ops.elementwise.elementwise:{_fn}",
+ )
+ )
+del _fn
diff --git a/python/sglang/kernels/ops/layernorm/_jit_norm.py b/python/sglang/kernels/ops/layernorm/norm.py
similarity index 100%
rename from python/sglang/kernels/ops/layernorm/_jit_norm.py
rename to python/sglang/kernels/ops/layernorm/norm.py
diff --git a/python/sglang/kernels/ops/model/__init__.py b/python/sglang/kernels/ops/model/__init__.py
deleted file mode 100644
index 34afe6ec2..000000000
--- a/python/sglang/kernels/ops/model/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Model-specific JIT kernels (RFC #29630)."""
diff --git a/python/sglang/kernels/ops/model/inkling/__init__.py b/python/sglang/kernels/ops/model/inkling/__init__.py
deleted file mode 100644
index 9f4d42f8d..000000000
--- a/python/sglang/kernels/ops/model/inkling/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Inkling model-family JIT kernels."""
diff --git a/python/sglang/kernels/ops/model/inkling/inkling_gate_topk_renorm.py b/python/sglang/kernels/ops/moe/inkling_gate_topk_renorm.py
similarity index 100%
rename from python/sglang/kernels/ops/model/inkling/inkling_gate_topk_renorm.py
rename to python/sglang/kernels/ops/moe/inkling_gate_topk_renorm.py
diff --git a/python/sglang/kernels/ops/quantization/__init__.py b/python/sglang/kernels/ops/quantization/__init__.py
index 2e377121c..a6f9ede74 100644
--- a/python/sglang/kernels/ops/quantization/__init__.py
+++ b/python/sglang/kernels/ops/quantization/__init__.py
@@ -56,7 +56,7 @@ register_kernel(
KernelSpec(
op="quantization.per_token_group_quant",
backend=KernelBackend.JIT,
- target="sglang.kernels.ops.quantization._jit_per_token_group_quant:per_token_group_quant",
+ target="sglang.kernels.ops.quantization.per_token_group_quant:per_token_group_quant",
capabilities=_CUDA,
format_signature=FormatSignature(
supported_dtypes=("float8_e4m3fn", "int8"),
diff --git a/python/sglang/kernels/ops/quantization/fp8_kernel.py b/python/sglang/kernels/ops/quantization/fp8_kernel.py
index cf4c87d9d..ddf478487 100644
--- a/python/sglang/kernels/ops/quantization/fp8_kernel.py
+++ b/python/sglang/kernels/ops/quantization/fp8_kernel.py
@@ -59,7 +59,7 @@ if _is_cuda or _is_musa:
per_token_group_quant,
sgl_per_token_quant_fp8,
)
- from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+ from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8,
)
@@ -543,7 +543,7 @@ def _run_per_token_group_quant_8bit_kernel(
``sglang_per_token_quant_fp8``.
"""
if scale_ue8m0 and x_s.dtype == torch.float32 and not _is_musa:
- from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
+ from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
diff --git a/python/sglang/kernels/ops/attention/hadamard.py b/python/sglang/kernels/ops/quantization/hadamard.py
similarity index 100%
rename from python/sglang/kernels/ops/attention/hadamard.py
rename to python/sglang/kernels/ops/quantization/hadamard.py
diff --git a/python/sglang/kernels/ops/quantization/_jit_per_tensor_quant_fp8.py b/python/sglang/kernels/ops/quantization/per_tensor_quant_fp8.py
similarity index 93%
rename from python/sglang/kernels/ops/quantization/_jit_per_tensor_quant_fp8.py
rename to python/sglang/kernels/ops/quantization/per_tensor_quant_fp8.py
index 4a45eb61d..98172fb81 100644
--- a/python/sglang/kernels/ops/quantization/_jit_per_tensor_quant_fp8.py
+++ b/python/sglang/kernels/ops/quantization/per_tensor_quant_fp8.py
@@ -12,7 +12,7 @@ if TYPE_CHECKING:
@cache_once
-def _jit_per_tensor_quant_fp8_module(is_static: bool, dtype: torch.dtype) -> Module:
+def per_tensor_quant_fp8_module(is_static: bool, dtype: torch.dtype) -> Module:
args = make_cpp_args(is_static, dtype)
return load_jit(
"per_tensor_quant_fp8",
@@ -41,7 +41,7 @@ def per_tensor_quant_fp8(
output_s: Output scale tensor (float scalar or 1D tensor with 1 element)
is_static: If True, assumes scale is pre-computed and skips absmax computation
"""
- module = _jit_per_tensor_quant_fp8_module(is_static, input.dtype)
+ module = per_tensor_quant_fp8_module(is_static, input.dtype)
module.per_tensor_quant_fp8(input.view(-1), output_q.view(-1), output_s.view(-1))
diff --git a/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant.py b/python/sglang/kernels/ops/quantization/per_token_group_quant.py
similarity index 100%
rename from python/sglang/kernels/ops/quantization/_jit_per_token_group_quant.py
rename to python/sglang/kernels/ops/quantization/per_token_group_quant.py
diff --git a/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant_8bit_v2.py b/python/sglang/kernels/ops/quantization/per_token_group_quant_8bit_v2.py
similarity index 97%
rename from python/sglang/kernels/ops/quantization/_jit_per_token_group_quant_8bit_v2.py
rename to python/sglang/kernels/ops/quantization/per_token_group_quant_8bit_v2.py
index 54eb87739..70d92e96d 100644
--- a/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant_8bit_v2.py
+++ b/python/sglang/kernels/ops/quantization/per_token_group_quant_8bit_v2.py
@@ -1,4 +1,4 @@
-"""DEPRECATED: superseded by ``sglang.kernels.ops.quantization._jit_per_token_group_quant`` (the
+"""DEPRECATED: superseded by ``sglang.kernels.ops.quantization.per_token_group_quant`` (the
default CUDA path). No sglang runtime code may call this kernel; it is kept
only as the perf baseline for the per_token_group_quant benchmarks and its own
bit-parity tests, and will be deleted once those move to torch references.
diff --git a/python/sglang/kernels/ops/spatial/__init__.py b/python/sglang/kernels/ops/spatial/__init__.py
deleted file mode 100644
index 7c8836334..000000000
--- a/python/sglang/kernels/ops/spatial/__init__.py
+++ /dev/null
@@ -1,49 +0,0 @@
-"""Spatial / green-context stream helpers."""
-
-from __future__ import annotations
-
-from typing import Optional
-
-from sglang.kernels.registry import register_kernel
-from sglang.kernels.selector import get_kernel
-from sglang.kernels.spec import FormatSignature, KernelBackend, KernelSpec
-
-register_kernel(
- KernelSpec(
- op="spatial.get_sm_available",
- backend=KernelBackend.AOT,
- target="sgl_kernel.spatial:get_sm_available",
- format_signature=FormatSignature(
- description="number of SMs available on device"
- ),
- description="Query available SM count (sgl_kernel wheel).",
- )
-)
-register_kernel(
- KernelSpec(
- op="spatial.create_greenctx_stream_by_value",
- backend=KernelBackend.AOT,
- target="sgl_kernel.spatial:create_greenctx_stream_by_value",
- format_signature=FormatSignature(
- description="create two green-context streams partitioned by SM count"
- ),
- description="Green-context stream creation (sgl_kernel wheel).",
- )
-)
-
-
-def get_sm_available(device_id: Optional[int] = None) -> int:
- """Return the number of SMs available on ``device_id``."""
- return get_kernel("spatial.get_sm_available", KernelBackend.AOT)(device_id)
-
-
-def create_greenctx_stream_by_value(
- SM_a: int, SM_b: int, device_id: Optional[int] = None
-):
- """Create two green-context streams partitioned by ``SM_a`` / ``SM_b``."""
- return get_kernel("spatial.create_greenctx_stream_by_value", KernelBackend.AOT)(
- SM_a, SM_b, device_id
- )
-
-
-__all__ = ["get_sm_available", "create_greenctx_stream_by_value"]
diff --git a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md
index 1acb738d4..c559f4917 100644
--- a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md
+++ b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md
@@ -29,7 +29,7 @@ framework-specific optimization workflow.
- `test/registered/jit/benchmark/diffusion/bench_qwen_image_modulation.py`
- `test/registered/jit/benchmark/diffusion/bench_group_norm_silu.py`
- `test/registered/jit/benchmark/diffusion/bench_residual_gate_add.py`
-- `python/sglang/kernels/ops/layernorm/_jit_norm.py`
+- `python/sglang/kernels/ops/layernorm/norm.py`
- `python/sglang/multimodal_gen/runtime/platforms/cuda.py`
- `python/sglang/multimodal_gen/runtime/layers/attention/selector.py`
- `docs_new/docs/sglang-diffusion/attention_backends.mdx` (repo root)
@@ -137,7 +137,7 @@ framework-specific optimization workflow.
**QK Norm Optimization**
- Entry point: `apply_qk_norm` in `layernorm.py`.
-- Fast path: JIT fused inplace QK norm from `python/sglang/kernels/ops/layernorm/_jit_norm.py` via `fused_inplace_qknorm`.
+- Fast path: JIT fused inplace QK norm from `python/sglang/kernels/ops/layernorm/norm.py` via `fused_inplace_qknorm`.
- Preconditions for fused path:
- CUDA only.
- `allow_inplace=True` and `q_eps == k_eps`.
diff --git a/python/sglang/multimodal_gen/runtime/layers/activation.py b/python/sglang/multimodal_gen/runtime/layers/activation.py
index 1e116dc10..23d0d06b0 100644
--- a/python/sglang/multimodal_gen/runtime/layers/activation.py
+++ b/python/sglang/multimodal_gen/runtime/layers/activation.py
@@ -19,7 +19,7 @@ _is_npu = current_platform.is_npu()
_is_xpu = current_platform.is_xpu()
if _is_cuda:
- from sglang.kernels.ops.activation._jit_activation import silu_and_mul
+ from sglang.kernels.ops.activation.activation import silu_and_mul
elif _is_hip or _is_xpu:
from sgl_kernel import silu_and_mul
diff --git a/python/sglang/multimodal_gen/runtime/layers/layernorm.py b/python/sglang/multimodal_gen/runtime/layers/layernorm.py
index 57f05d68b..d57e86373 100755
--- a/python/sglang/multimodal_gen/runtime/layers/layernorm.py
+++ b/python/sglang/multimodal_gen/runtime/layers/layernorm.py
@@ -17,7 +17,7 @@ from sglang.kernels.ops.diffusion.qknorm_rope import (
)
from sglang.kernels.ops.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm
from sglang.kernels.ops.diffusion.triton.scale_shift import fuse_scale_shift_kernel
-from sglang.kernels.ops.layernorm._jit_norm import (
+from sglang.kernels.ops.layernorm.norm import (
can_use_fused_inplace_qknorm,
fused_inplace_qknorm,
)
diff --git a/python/sglang/srt/layers/activation.py b/python/sglang/srt/layers/activation.py
index 2c9a2a123..9bd538397 100644
--- a/python/sglang/srt/layers/activation.py
+++ b/python/sglang/srt/layers/activation.py
@@ -57,7 +57,7 @@ _is_xpu = is_xpu()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _is_cuda:
- from sglang.kernels.ops.activation._jit_activation import (
+ from sglang.kernels.ops.activation.activation import (
gelu_and_mul,
gelu_tanh_and_mul,
relu2,
diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py
index 526c61b3e..036db24d6 100644
--- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py
+++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py
@@ -343,7 +343,7 @@ def rotate_activation(x: torch.Tensor) -> torch.Tensor:
elif _is_xpu:
from sgl_kernel import hadamard_transform
else:
- from sglang.kernels.ops.attention.hadamard import hadamard_transform
+ from sglang.kernels.ops.quantization.hadamard import hadamard_transform
hidden_size = x.size(-1)
assert (
diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py
index 3722af7eb..8eeb0807c 100644
--- a/python/sglang/srt/layers/attention/vision.py
+++ b/python/sglang/srt/layers/attention/vision.py
@@ -12,7 +12,7 @@ import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
-from sglang.kernels.ops.layernorm._jit_norm import (
+from sglang.kernels.ops.layernorm.norm import (
can_use_fused_inplace_qknorm as can_use_jit_qk_norm,
)
from sglang.srt.environ import envs
diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py
index b23a7cce6..66f596593 100644
--- a/python/sglang/srt/layers/layernorm.py
+++ b/python/sglang/srt/layers/layernorm.py
@@ -147,10 +147,10 @@ if _is_cuda:
_jit_rmsnorm_hf = None
- from sglang.kernels.ops.layernorm._jit_norm import (
+ from sglang.kernels.ops.layernorm.norm import (
fused_add_rmsnorm as _jit_fused_add_rmsnorm,
)
- from sglang.kernels.ops.layernorm._jit_norm import (
+ from sglang.kernels.ops.layernorm.norm import (
is_supported_jit_fused_add_rmsnorm_hidden_size,
)
diff --git a/python/sglang/srt/layers/moe/cutlass_moe.py b/python/sglang/srt/layers/moe/cutlass_moe.py
index dc80ed9f5..afc636fb1 100755
--- a/python/sglang/srt/layers/moe/cutlass_moe.py
+++ b/python/sglang/srt/layers/moe/cutlass_moe.py
@@ -18,7 +18,7 @@ if _is_cuda:
shuffle_rows,
)
- from sglang.kernels.ops.activation._jit_activation import silu_and_mul
+ from sglang.kernels.ops.activation.activation import silu_and_mul
def cutlass_fused_experts_fp8(
diff --git a/python/sglang/srt/layers/moe/cutlass_w4a8_moe.py b/python/sglang/srt/layers/moe/cutlass_w4a8_moe.py
index 3e965580a..613d4a0eb 100644
--- a/python/sglang/srt/layers/moe/cutlass_w4a8_moe.py
+++ b/python/sglang/srt/layers/moe/cutlass_w4a8_moe.py
@@ -18,7 +18,7 @@ if _is_cuda_alike:
)
if _is_cuda:
- from sglang.kernels.ops.activation._jit_activation import silu_and_mul
+ from sglang.kernels.ops.activation.activation import silu_and_mul
else:
from sgl_kernel import silu_and_mul
@@ -35,7 +35,7 @@ from sglang.kernels.ops.moe.ep_moe_kernels import (
silu_mul_dynamic_tensorwise_quant_for_cutlass_moe,
silu_mul_static_tensorwise_quant_for_cutlass_moe,
)
-from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_absmax_fp8,
per_tensor_quant_fp8,
)
diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py
index c7779b61f..9ab1dd29c 100644
--- a/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py
+++ b/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py
@@ -11,7 +11,7 @@ _is_cuda = is_cuda()
if _is_cuda:
from sgl_kernel import moe_sum_reduce
- from sglang.kernels.ops.activation._jit_activation import silu_and_mul
+ from sglang.kernels.ops.activation.activation import silu_and_mul
from sglang.kernels.ops.moe.moe_wna16_marlin import moe_wna16_marlin_gemm
diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/triton_kernels_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/triton_kernels_moe.py
index f0e8ee6a1..df283a7a2 100644
--- a/python/sglang/srt/layers/moe/fused_moe_triton/triton_kernels_moe.py
+++ b/python/sglang/srt/layers/moe/fused_moe_triton/triton_kernels_moe.py
@@ -30,7 +30,7 @@ if is_sm120_supported():
update_opt_flags_constraints({"is_persistent": False})
if is_cuda():
- from sglang.kernels.ops.activation._jit_activation import gelu_and_mul, silu_and_mul
+ from sglang.kernels.ops.activation.activation import gelu_and_mul, silu_and_mul
else:
from sgl_kernel import gelu_and_mul, silu_and_mul
diff --git a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py
index 16a7697ab..5d160b282 100644
--- a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py
+++ b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py
@@ -56,7 +56,7 @@ _is_musa = is_musa()
# Imported only for the SGLANG_OPT_FIX_MEGA_MOE_MEMORY=False fallback path.
if not (_is_npu or _is_hip) and _is_cuda:
- from sglang.kernels.ops.activation._jit_activation import (
+ from sglang.kernels.ops.activation.activation import (
silu_and_mul as _legacy_silu_and_mul,
)
elif _is_musa:
diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py
index 90f940495..acc9bc00c 100644
--- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py
+++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py
@@ -60,7 +60,7 @@ _is_musa = is_musa()
if _is_cuda:
from sgl_kernel import moe_sum_reduce
- from sglang.kernels.ops.activation._jit_activation import gelu_and_mul, silu_and_mul
+ from sglang.kernels.ops.activation.activation import gelu_and_mul, silu_and_mul
elif _is_cpu and _is_cpu_amx_available:
pass
elif _is_hip:
diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py
index 635317361..e6a0a7cdf 100644
--- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py
+++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py
@@ -13,7 +13,7 @@ import triton
import triton.language as tl
from sglang.kernels.jit.utils import is_arch_support_pdl
-from sglang.kernels.ops.model.inkling.inkling_gate_topk_renorm import (
+from sglang.kernels.ops.moe.inkling_gate_topk_renorm import (
inkling_gate_topk_renorm_v2,
)
from sglang.srt.environ import envs
diff --git a/python/sglang/srt/layers/quantization/gguf.py b/python/sglang/srt/layers/quantization/gguf.py
index e4eda93fa..f326d6d50 100644
--- a/python/sglang/srt/layers/quantization/gguf.py
+++ b/python/sglang/srt/layers/quantization/gguf.py
@@ -51,7 +51,7 @@ if _is_cuda:
ggml_mul_mat_vec_a8,
)
- from sglang.kernels.ops.activation._jit_activation import gelu_and_mul, silu_and_mul
+ from sglang.kernels.ops.activation.activation import gelu_and_mul, silu_and_mul
elif _is_musa:
from sgl_kernel import gelu_and_mul, moe_align_block_size, moe_sum, silu_and_mul
from sgl_kernel.quantization import (
diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py
index 312be7fc0..fdb863a01 100644
--- a/python/sglang/srt/models/deepseek_v2.py
+++ b/python/sglang/srt/models/deepseek_v2.py
@@ -213,8 +213,8 @@ if _use_aiter:
pass
if _is_cuda:
- from sglang.kernels.ops.gemm._jit_dsv3_router_gemm import (
- dsv3_router_gemm as _jit_dsv3_router_gemm,
+ from sglang.kernels.ops.gemm.dsv3_router_gemm import (
+ dsv3_router_gemm as dsv3_router_gemm,
)
elif _is_npu:
from sglang.srt.hardware_backend.npu.modules.deepseek_v2_attention_mla_npu import (
@@ -512,7 +512,7 @@ class MoEGate(nn.Module):
and (self.weight.shape[0] == 256 or self.weight.shape[0] == 384)
and _device_sm >= 90
):
- logits = _jit_dsv3_router_gemm(
+ logits = dsv3_router_gemm(
hidden_states, self.weight, out_dtype=torch.float32
)
diff --git a/python/sglang/srt/models/grok.py b/python/sglang/srt/models/grok.py
index f06912025..34a43ef9d 100644
--- a/python/sglang/srt/models/grok.py
+++ b/python/sglang/srt/models/grok.py
@@ -22,7 +22,7 @@ import torch.nn.functional as F
from torch import nn
from transformers import PretrainedConfig
-from sglang.kernels.ops.layernorm.elementwise import (
+from sglang.kernels.ops.elementwise.elementwise import (
fused_dual_residual_rmsnorm,
fused_rmsnorm,
gelu_and_mul_triton,
diff --git a/python/sglang/srt/models/inkling.py b/python/sglang/srt/models/inkling.py
index c2e10f487..6d71fb7b1 100644
--- a/python/sglang/srt/models/inkling.py
+++ b/python/sglang/srt/models/inkling.py
@@ -625,7 +625,7 @@ class InklingCausalLLM(nn.Module):
and sconv0 is not None
and world in (4, 8) # symm-mem multimem worlds, power-of-two
):
- from sglang.kernels.ops.model.inkling.inkling_ar_fused import (
+ from sglang.kernels.ops.communication.inkling_ar_fused import (
compile_inkling_ar_sconv_norm,
)
@@ -646,7 +646,7 @@ class InklingCausalLLM(nn.Module):
# local/SWA layer (head_dim != 128) that never uses the prologue while
# later full-attention layers do.
if is_cuda() and envs.SGLANG_OPT_USE_INKLING_FUSED_ATTN_PROLOGUE.get():
- from sglang.kernels.ops.model.inkling.inkling_attn_prologue import (
+ from sglang.kernels.ops.attention.inkling_attn_prologue import (
compile_inkling_attn_prologue,
)
diff --git a/python/sglang/srt/models/inkling_common/attn.py b/python/sglang/srt/models/inkling_common/attn.py
index 14b12e576..6b6c7300b 100644
--- a/python/sglang/srt/models/inkling_common/attn.py
+++ b/python/sglang/srt/models/inkling_common/attn.py
@@ -6,14 +6,14 @@ from functools import cache
import torch
from torch import nn
+from sglang.kernels.ops.attention.inkling_rel_proj import rel_proj_small_t
+from sglang.kernels.ops.attention.inkling_row_scale import row_compact_bf16
from sglang.kernels.ops.attention.log_scaling_tau import (
apply_log_scaling_tau as _apply_log_scaling_tau,
)
from sglang.kernels.ops.attention.score_mod import (
relative_bias_score_mod as triton_relative_bias_score_mod,
)
-from sglang.kernels.ops.model.inkling.inkling_rel_proj import rel_proj_small_t
-from sglang.kernels.ops.model.inkling.inkling_row_scale import row_compact_bf16
from sglang.srt.environ import envs
from sglang.srt.layers.linear import MergedColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
@@ -364,7 +364,7 @@ class InklingAttention(nn.Module):
def _fused_attn_prologue_verify(self, q, k, v, forward_batch, log_scaling_tau=None):
"""Fused target-verify {k/v sconv + save_windows + qk-norm (+ KV store)}
- (kernels/ops/model/inkling/inkling_attn_prologue.py); returns
+ (kernels/ops/attention/inkling_attn_prologue.py); returns
``(q, k, v, did_store)``.
The fused kernel writes raw bf16 KV, so it only does the store when the
@@ -377,7 +377,7 @@ class InklingAttention(nn.Module):
qk-norm stay fused either way. For the FA4 MXFP8 pool, the prologue can
quantize Q and directly fill the fp8 K/V cache plus interleaved scale
buffers, returning Q's per-token scales as ``q_descale``/``sfq``."""
- from sglang.kernels.ops.model.inkling.inkling_attn_prologue import (
+ from sglang.kernels.ops.attention.inkling_attn_prologue import (
inkling_attn_prologue_verify,
)
from sglang.srt.model_executor.forward_context import (
@@ -481,7 +481,7 @@ class InklingAttention(nn.Module):
the backend store. Store gating (bf16 NHD / FA4 MXFP8 pools, SWA loc
translation) is identical to the verify prologue. Returns
(q, k, v, did_store, q_descale)."""
- from sglang.kernels.ops.model.inkling.inkling_attn_prologue import (
+ from sglang.kernels.ops.attention.inkling_attn_prologue import (
inkling_attn_prologue_extend,
)
from sglang.srt.model_executor.forward_context import (
@@ -615,7 +615,7 @@ class InklingAttention(nn.Module):
kernel. Decode is one token/seq so the conv taps come from the working
cache (no cross-token reads, no barrier). Returns
(q, k, v, did_store, q_descale)."""
- from sglang.kernels.ops.model.inkling.inkling_attn_prologue import (
+ from sglang.kernels.ops.attention.inkling_attn_prologue import (
inkling_attn_prologue_decode,
)
from sglang.srt.model_executor.forward_context import (
diff --git a/python/sglang/srt/models/inkling_common/kernels/comm.py b/python/sglang/srt/models/inkling_common/kernels/comm.py
index 02f979b81..5fc0b987d 100644
--- a/python/sglang/srt/models/inkling_common/kernels/comm.py
+++ b/python/sglang/srt/models/inkling_common/kernels/comm.py
@@ -94,7 +94,7 @@ def _ar_jit():
lazy so importing comm.py doesn't pull in the JIT machinery)."""
if not is_cuda():
return None
- from sglang.kernels.ops.model.inkling import inkling_all_reduce
+ from sglang.kernels.ops.communication import inkling_all_reduce
return inkling_all_reduce
@@ -103,7 +103,7 @@ def _ar_jit():
def _ar_fused_jit():
if not is_cuda():
return None
- from sglang.kernels.ops.model.inkling import inkling_ar_fused
+ from sglang.kernels.ops.communication import inkling_ar_fused
return inkling_ar_fused
@@ -241,7 +241,7 @@ def ar_sconv_norm_fusable(
(attn-side: wo_ud AR -> attn_sconv -> mlp_norm; MoE-side: MoE AR ->
mlp_sconv -> next attn_norm)
can run as the single fused kernel
- (kernels/ops/model/inkling/inkling_ar_fused.py). Must be
+ (kernels/ops/communication/inkling_ar_fused.py). Must be
evaluated identically by the producing layer (MoE ``reduce=False``) and the
consuming layer/tail -- it is a pure function of per-forward state."""
if not is_cuda():
@@ -675,7 +675,7 @@ def all_gather_hidden(input: torch.Tensor, group: GroupCoordinator) -> torch.Ten
def _ar_ssconv_jit():
if not is_cuda():
return None
- from sglang.kernels.ops.model.inkling import inkling_ar_scattered_sconv
+ from sglang.kernels.ops.communication import inkling_ar_scattered_sconv
return inkling_ar_scattered_sconv
@@ -689,7 +689,7 @@ def scattered_ar_sconv_fusable(
) -> bool:
"""True when an extend {reduce_scatter_hidden -> sconv(shard) ->
all_gather_hidden} chain can run as the single fused v3/v3b-style kernel
- (kernels/ops/model/inkling/inkling_ar_scattered_sconv.py). Pure function of
+ (kernels/ops/communication/inkling_ar_scattered_sconv.py). Pure function of
per-forward
state -- the producing layer (reduce=False) and the consuming site must
evaluate it identically."""
diff --git a/python/sglang/srt/models/inkling_common/moe.py b/python/sglang/srt/models/inkling_common/moe.py
index e6f7bf918..c4036069b 100644
--- a/python/sglang/srt/models/inkling_common/moe.py
+++ b/python/sglang/srt/models/inkling_common/moe.py
@@ -10,7 +10,7 @@ from torch import nn
from triton.language.extra import libdevice
from sglang.kernels.jit.utils import is_arch_support_pdl
-from sglang.kernels.ops.model.inkling.inkling_gate_topk_renorm import (
+from sglang.kernels.ops.moe.inkling_gate_topk_renorm import (
ensure_gate_gemv_fused_scratch,
inkling_gate_gemv,
inkling_gate_gemv_fused,
diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py
index 4c35be1f7..d2b744a68 100644
--- a/python/sglang/srt/models/qwen2_moe.py
+++ b/python/sglang/srt/models/qwen2_moe.py
@@ -27,7 +27,7 @@ import torch.nn.functional as F
from torch import nn
from transformers import PretrainedConfig
-from sglang.kernels.ops.layernorm.elementwise import fused_gate_sigmoid_mul_add
+from sglang.kernels.ops.elementwise.elementwise import fused_gate_sigmoid_mul_add
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
from sglang.srt.distributed import (
get_pp_group,
diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py
index ecbb4c239..fcdcb92e0 100644
--- a/python/sglang/srt/models/qwen3_5.py
+++ b/python/sglang/srt/models/qwen3_5.py
@@ -27,7 +27,7 @@ from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormG
from sglang.kernels.ops.attention.triton_gdn_fused_proj import (
fused_qkvzba_split_reshape_cat_contiguous,
)
-from sglang.kernels.ops.layernorm.elementwise import fused_sigmoid_mul
+from sglang.kernels.ops.elementwise.elementwise import fused_sigmoid_mul
# Configs
from sglang.srt.configs.qwen3_5 import (
diff --git a/python/sglang/srt/models/utils.py b/python/sglang/srt/models/utils.py
index 97e0f2325..de686c226 100644
--- a/python/sglang/srt/models/utils.py
+++ b/python/sglang/srt/models/utils.py
@@ -25,7 +25,7 @@ import triton
import triton.language as tl
from sglang.kernels.ops.attention.rope import FusedSetKVBufferArg
-from sglang.kernels.ops.layernorm._jit_norm import (
+from sglang.kernels.ops.layernorm.norm import (
can_use_fused_inplace_qknorm,
fused_inplace_qknorm,
)
diff --git a/sgl-kernel/benchmark/bench_fp8_gemm.py b/sgl-kernel/benchmark/bench_fp8_gemm.py
index 35e10cb2e..2205d8e39 100644
--- a/sgl-kernel/benchmark/bench_fp8_gemm.py
+++ b/sgl-kernel/benchmark/bench_fp8_gemm.py
@@ -8,7 +8,7 @@ import torch
import triton
from sgl_kernel import fp8_scaled_mm as sgl_scaled_mm
-from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.utils import is_in_ci
diff --git a/sgl-kernel/benchmark/bench_fp8_gemm_swap_ab.py b/sgl-kernel/benchmark/bench_fp8_gemm_swap_ab.py
index 2d2127dd2..230330f5b 100644
--- a/sgl-kernel/benchmark/bench_fp8_gemm_swap_ab.py
+++ b/sgl-kernel/benchmark/bench_fp8_gemm_swap_ab.py
@@ -18,7 +18,7 @@ import torch
import triton
from sgl_kernel import fp8_scaled_mm as sgl_scaled_mm
-from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.utils import is_in_ci
diff --git a/sgl-kernel/benchmark/bench_per_tensor_quant_fp8.py b/sgl-kernel/benchmark/bench_per_tensor_quant_fp8.py
index fee2e4158..be70dfb88 100644
--- a/sgl-kernel/benchmark/bench_per_tensor_quant_fp8.py
+++ b/sgl-kernel/benchmark/bench_per_tensor_quant_fp8.py
@@ -8,7 +8,7 @@ import torch
import triton
import triton.testing
-from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.utils import is_in_ci
diff --git a/sgl-kernel/tests/test_cutlass_w4a8_moe_mm.py b/sgl-kernel/tests/test_cutlass_w4a8_moe_mm.py
index 969775fdd..57a0b0cee 100644
--- a/sgl-kernel/tests/test_cutlass_w4a8_moe_mm.py
+++ b/sgl-kernel/tests/test_cutlass_w4a8_moe_mm.py
@@ -5,7 +5,7 @@ import torch
from sgl_kernel import cutlass_w4a8_moe_mm
from utils import is_hopper
-from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
diff --git a/test/manual/layers/test_fused_gate_sigmoid_mul_add.py b/test/manual/layers/test_fused_gate_sigmoid_mul_add.py
index cc3fccc4a..c081c187e 100644
--- a/test/manual/layers/test_fused_gate_sigmoid_mul_add.py
+++ b/test/manual/layers/test_fused_gate_sigmoid_mul_add.py
@@ -3,7 +3,7 @@ import itertools
import pytest
import torch
-from sglang.kernels.ops.layernorm.elementwise import fused_gate_sigmoid_mul_add
+from sglang.kernels.ops.elementwise.elementwise import fused_gate_sigmoid_mul_add
DTYPES = [torch.float16, torch.bfloat16]
TOKEN_COUNTS = [1, 2, 4, 8, 16, 64, 512, 1024, 2048, 4096, 8192]
diff --git a/test/manual/layers/test_fused_sigmoid_mul.py b/test/manual/layers/test_fused_sigmoid_mul.py
index 3d5cbf44d..56b3d78bd 100644
--- a/test/manual/layers/test_fused_sigmoid_mul.py
+++ b/test/manual/layers/test_fused_sigmoid_mul.py
@@ -3,7 +3,7 @@ import itertools
import pytest
import torch
-from sglang.kernels.ops.layernorm.elementwise import fused_sigmoid_mul
+from sglang.kernels.ops.elementwise.elementwise import fused_sigmoid_mul
DTYPES = [torch.float16, torch.bfloat16]
TOKEN_COUNTS = [1, 2, 4, 8, 16, 64, 512, 1024, 2048, 4096, 8192]
diff --git a/test/registered/kernels/benchmark/activation/bench_activation.py b/test/registered/kernels/benchmark/activation/bench_activation.py
index c80d8bab0..983ef269f 100644
--- a/test/registered/kernels/benchmark/activation/bench_activation.py
+++ b/test/registered/kernels/benchmark/activation/bench_activation.py
@@ -6,16 +6,12 @@ from sgl_kernel import silu_and_mul as silu_and_mul_aot
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
-from sglang.kernels.ops.activation._jit_activation import (
- gelu_and_mul as gelu_and_mul_jit,
-)
-from sglang.kernels.ops.activation._jit_activation import (
+from sglang.kernels.ops.activation.activation import gelu_and_mul as gelu_and_mul_jit
+from sglang.kernels.ops.activation.activation import (
gelu_tanh_and_mul as gelu_tanh_and_mul_jit,
)
-from sglang.kernels.ops.activation._jit_activation import relu2 as relu2_jit
-from sglang.kernels.ops.activation._jit_activation import (
- silu_and_mul as silu_and_mul_jit,
-)
+from sglang.kernels.ops.activation.activation import relu2 as relu2_jit
+from sglang.kernels.ops.activation.activation import silu_and_mul as silu_and_mul_jit
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
diff --git a/test/registered/kernels/benchmark/attention/bench_add_constant.py b/test/registered/kernels/benchmark/attention/bench_add_constant.py
index e0d20985d..a9996b673 100644
--- a/test/registered/kernels/benchmark/attention/bench_add_constant.py
+++ b/test/registered/kernels/benchmark/attention/bench_add_constant.py
@@ -7,7 +7,7 @@ from sglang.kernels.jit.benchmark.utils import (
get_benchmark_range,
run_benchmark_no_cudagraph,
)
-from sglang.kernels.ops.attention.add_constant import (
+from sglang.kernels.ops.elementwise.add_constant import (
_jit_add_constant_module,
add_constant,
)
diff --git a/test/registered/kernels/benchmark/attention/bench_hadamard.py b/test/registered/kernels/benchmark/attention/bench_hadamard.py
index 5a2da9f23..7332dfa36 100644
--- a/test/registered/kernels/benchmark/attention/bench_hadamard.py
+++ b/test/registered/kernels/benchmark/attention/bench_hadamard.py
@@ -13,7 +13,7 @@ from sglang.kernels.jit.benchmark.utils import (
get_benchmark_range,
run_benchmark,
)
-from sglang.kernels.ops.attention.hadamard import hadamard_transform
+from sglang.kernels.ops.quantization.hadamard import hadamard_transform
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
diff --git a/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py b/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py
index 3ff6af016..fff0ddee6 100644
--- a/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py
+++ b/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py
@@ -17,10 +17,8 @@ from sglang.kernels.jit.benchmark.utils import DEFAULT_DEVICE
from sglang.kernels.jit.utils import KERNEL_PATH
from sglang.kernels.ops.diffusion.triton.norm import norm_infer, rms_norm_fn
from sglang.kernels.ops.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm
-from sglang.kernels.ops.layernorm._jit_norm import (
- fused_add_rmsnorm as jit_fused_add_rmsnorm,
-)
-from sglang.kernels.ops.layernorm._jit_norm import rmsnorm as jit_rmsnorm
+from sglang.kernels.ops.layernorm.norm import fused_add_rmsnorm as jit_fused_add_rmsnorm
+from sglang.kernels.ops.layernorm.norm import rmsnorm as jit_rmsnorm
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.utils import is_in_ci
diff --git a/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py b/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py
index 6ed86c581..7189ebf0c 100644
--- a/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py
+++ b/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py
@@ -131,7 +131,7 @@ def clone_inputs(
def split_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None:
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
- from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
+ from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
q = inputs["q"]
k = inputs["k"]
diff --git a/test/registered/kernels/benchmark/gemm/bench_dsv3_fused_a_gemm.py b/test/registered/kernels/benchmark/gemm/bench_dsv3_fused_a_gemm.py
index 5de153d71..8f6e4acc8 100644
--- a/test/registered/kernels/benchmark/gemm/bench_dsv3_fused_a_gemm.py
+++ b/test/registered/kernels/benchmark/gemm/bench_dsv3_fused_a_gemm.py
@@ -10,10 +10,10 @@ import triton.testing
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.utils import get_jit_cuda_arch, is_hip_runtime
-from sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm import dsv3_fused_a_gemm
from sglang.kernels.ops.gemm.cutedsl_dsv3_fused_a_gemm import (
dsv3_fused_a_gemm as cutedsl_dsv3_fused_a_gemm,
)
+from sglang.kernels.ops.gemm.dsv3_fused_a_gemm import dsv3_fused_a_gemm
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
diff --git a/test/registered/kernels/benchmark/gemm/bench_dsv3_router_gemm.py b/test/registered/kernels/benchmark/gemm/bench_dsv3_router_gemm.py
index e85af0b8b..1a4442bab 100644
--- a/test/registered/kernels/benchmark/gemm/bench_dsv3_router_gemm.py
+++ b/test/registered/kernels/benchmark/gemm/bench_dsv3_router_gemm.py
@@ -10,7 +10,7 @@ import torch.nn.functional as F
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.jit.utils import get_jit_cuda_arch, is_hip_runtime
-from sglang.kernels.ops.gemm._jit_dsv3_router_gemm import dsv3_router_gemm
+from sglang.kernels.ops.gemm.dsv3_router_gemm import dsv3_router_gemm
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
diff --git a/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py b/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py
index a6d369c0e..90bbfa1a8 100644
--- a/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py
+++ b/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py
@@ -21,9 +21,7 @@ from sglang.kernels.jit.benchmark.utils import (
get_benchmark_range,
)
from sglang.kernels.jit.utils import is_arch_support_pdl
-from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import (
- set_mla_kv_buffer as jit_set,
-)
+from sglang.kernels.ops.kvcache.set_mla_kv_buffer import set_mla_kv_buffer as jit_set
from sglang.srt.mem_cache.utils import set_mla_kv_buffer_kernel as sglang_triton_kernel
from sglang.srt.mem_cache.utils import set_mla_kv_buffer_triton as sglang_wrapper
from sglang.test.ci.ci_register import register_cuda_ci
diff --git a/test/registered/kernels/benchmark/layernorm/bench_norm.py b/test/registered/kernels/benchmark/layernorm/bench_norm.py
index f4a52596d..be1605992 100644
--- a/test/registered/kernels/benchmark/layernorm/bench_norm.py
+++ b/test/registered/kernels/benchmark/layernorm/bench_norm.py
@@ -7,10 +7,8 @@ from flashinfer.norm import fused_add_rmsnorm as fi_fused_add_rmsnorm
from flashinfer.norm import rmsnorm as fi_rmsnorm
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
-from sglang.kernels.ops.layernorm._jit_norm import (
- fused_add_rmsnorm as jit_fused_add_rmsnorm,
-)
-from sglang.kernels.ops.layernorm._jit_norm import rmsnorm as jit_rmsnorm
+from sglang.kernels.ops.layernorm.norm import fused_add_rmsnorm as jit_fused_add_rmsnorm
+from sglang.kernels.ops.layernorm.norm import rmsnorm as jit_rmsnorm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
diff --git a/test/registered/kernels/benchmark/layernorm/bench_qknorm.py b/test/registered/kernels/benchmark/layernorm/bench_qknorm.py
index fdc738b1b..78f149a13 100644
--- a/test/registered/kernels/benchmark/layernorm/bench_qknorm.py
+++ b/test/registered/kernels/benchmark/layernorm/bench_qknorm.py
@@ -2,7 +2,7 @@ import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
-from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
+from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
from sglang.srt.utils import get_current_device_stream_fast
from sglang.test.ci.ci_register import register_cuda_ci
diff --git a/test/registered/kernels/benchmark/layernorm/bench_qknorm_across_heads.py b/test/registered/kernels/benchmark/layernorm/bench_qknorm_across_heads.py
index be6726f20..298e3961d 100644
--- a/test/registered/kernels/benchmark/layernorm/bench_qknorm_across_heads.py
+++ b/test/registered/kernels/benchmark/layernorm/bench_qknorm_across_heads.py
@@ -7,7 +7,7 @@ import triton.testing
from sgl_kernel import rmsnorm
from sglang.kernels.jit.benchmark.utils import run_benchmark
-from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm_across_heads
+from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm_across_heads
from sglang.srt.utils import get_current_device_stream_fast
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
diff --git a/test/registered/kernels/benchmark/quantization/bench_per_tensor_quant_fp8.py b/test/registered/kernels/benchmark/quantization/bench_per_tensor_quant_fp8.py
index f9c9cef0b..5f8c046cd 100644
--- a/test/registered/kernels/benchmark/quantization/bench_per_tensor_quant_fp8.py
+++ b/test/registered/kernels/benchmark/quantization/bench_per_tensor_quant_fp8.py
@@ -5,7 +5,7 @@ import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
-from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.test.ci.ci_register import register_cuda_ci
diff --git a/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant.py b/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant.py
index 8ba53f2cc..a22e6bfed 100644
--- a/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant.py
+++ b/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant.py
@@ -1,20 +1,20 @@
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_empty, create_random
-
-# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
-# kernel is kept only as the perf baseline for this benchmark.
-from sglang.kernels.ops.quantization._jit_per_token_group_quant import (
- per_token_group_quant,
-)
-from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
- per_token_group_quant_8bit_v2,
-)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
fp8_min,
)
+
+# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
+# kernel is kept only as the perf baseline for this benchmark.
+from sglang.kernels.ops.quantization.per_token_group_quant import (
+ per_token_group_quant,
+)
+from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
+ per_token_group_quant_8bit_v2,
+)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
diff --git a/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_8bit_v2.py b/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_8bit_v2.py
index 7b6b940f8..47b9b00b4 100644
--- a/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_8bit_v2.py
+++ b/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_8bit_v2.py
@@ -3,15 +3,15 @@ from sgl_kernel import sgl_per_token_group_quant_8bit
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
-from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
- per_token_group_quant_8bit_v2,
-)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
fp8_min,
)
+from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
+ per_token_group_quant_8bit_v2,
+)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
diff --git a/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_masked.py b/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_masked.py
index a34c578e8..c9ca5d824 100644
--- a/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_masked.py
+++ b/test/registered/kernels/benchmark/quantization/bench_per_token_group_quant_masked.py
@@ -4,21 +4,21 @@ import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_empty, create_random
-
-# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
-# kernel is kept only as the perf baseline for this benchmark.
-from sglang.kernels.ops.quantization._jit_per_token_group_quant import (
- per_token_group_quant,
-)
-from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
- per_token_group_quant_8bit_v2,
-)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
fp8_min,
)
+
+# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
+# kernel is kept only as the perf baseline for this benchmark.
+from sglang.kernels.ops.quantization.per_token_group_quant import (
+ per_token_group_quant,
+)
+from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
+ per_token_group_quant_8bit_v2,
+)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
diff --git a/test/registered/kernels/ops/activation/test_activation.py b/test/registered/kernels/ops/activation/test_activation.py
index e8244c5b1..cb4737fdd 100644
--- a/test/registered/kernels/ops/activation/test_activation.py
+++ b/test/registered/kernels/ops/activation/test_activation.py
@@ -5,7 +5,7 @@ import torch
import torch.nn.functional as F
from sglang.kernels.jit.utils import get_ci_test_range
-from sglang.kernels.ops.activation._jit_activation import (
+from sglang.kernels.ops.activation.activation import (
SUPPORTED_ACTIVATIONS,
relu2,
run_activation,
diff --git a/test/registered/kernels/ops/attention/test_add_constant.py b/test/registered/kernels/ops/attention/test_add_constant.py
index 2e97174ad..c5fc47fd8 100644
--- a/test/registered/kernels/ops/attention/test_add_constant.py
+++ b/test/registered/kernels/ops/attention/test_add_constant.py
@@ -3,7 +3,7 @@ import sys
import pytest
import torch
-from sglang.kernels.ops.attention.add_constant import add_constant
+from sglang.kernels.ops.elementwise.add_constant import add_constant
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="1-gpu-large")
diff --git a/test/registered/kernels/ops/attention/test_fp4_indexer.py b/test/registered/kernels/ops/attention/test_fp4_indexer.py
index a64be2b6e..45639a098 100644
--- a/test/registered/kernels/ops/attention/test_fp4_indexer.py
+++ b/test/registered/kernels/ops/attention/test_fp4_indexer.py
@@ -27,7 +27,7 @@ _is_xpu = is_xpu()
if _is_xpu:
from sgl_kernel import hadamard_transform
else:
- from sglang.kernels.ops.attention.hadamard import hadamard_transform
+ from sglang.kernels.ops.quantization.hadamard import hadamard_transform
HEAD_DIM = 128
FP4_DIM = HEAD_DIM // 2
diff --git a/test/registered/kernels/ops/attention/test_hadamard_jit.py b/test/registered/kernels/ops/attention/test_hadamard_jit.py
index dec9981fc..3c53bd29f 100644
--- a/test/registered/kernels/ops/attention/test_hadamard_jit.py
+++ b/test/registered/kernels/ops/attention/test_hadamard_jit.py
@@ -7,7 +7,7 @@ import torch
import torch.nn.functional as F
from scipy.linalg import hadamard
-from sglang.kernels.ops.attention.hadamard import (
+from sglang.kernels.ops.quantization.hadamard import (
hadamard_transform,
hadamard_transform_12n,
hadamard_transform_20n,
diff --git a/test/registered/kernels/ops/attention/test_inkling_attn_prologue_tau.py b/test/registered/kernels/ops/attention/test_inkling_attn_prologue_tau.py
index 5198bec0f..a5e4a7804 100644
--- a/test/registered/kernels/ops/attention/test_inkling_attn_prologue_tau.py
+++ b/test/registered/kernels/ops/attention/test_inkling_attn_prologue_tau.py
@@ -7,7 +7,7 @@ must be untouched by tau.
import pytest
import torch
-from sglang.kernels.ops.model.inkling.inkling_attn_prologue import (
+from sglang.kernels.ops.attention.inkling_attn_prologue import (
inkling_attn_prologue_decode,
)
from sglang.test.ci.ci_register import register_cuda_ci
diff --git a/test/registered/kernels/ops/attention/test_inkling_row_scale.py b/test/registered/kernels/ops/attention/test_inkling_row_scale.py
index 5bb1fd872..68b5febad 100644
--- a/test/registered/kernels/ops/attention/test_inkling_row_scale.py
+++ b/test/registered/kernels/ops/attention/test_inkling_row_scale.py
@@ -5,10 +5,10 @@ including on the row-strided qkvr-slice layouts."""
import pytest
import torch
+from sglang.kernels.ops.attention.inkling_row_scale import row_scale_bf16
from sglang.kernels.ops.attention.log_scaling_tau import (
_apply_log_scaling_tau_kernel,
)
-from sglang.kernels.ops.model.inkling.inkling_row_scale import row_scale_bf16
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -51,7 +51,7 @@ def test_row_compact_bitexact(rows, inner, strided):
"""The tau-less compaction flavor (kHasTau=false) must reproduce
.contiguous() exactly on the same strided layouts row_scale handles --
no other test exercises run_compact."""
- from sglang.kernels.ops.model.inkling.inkling_row_scale import row_compact_bf16
+ from sglang.kernels.ops.attention.inkling_row_scale import row_compact_bf16
torch.manual_seed(rows + inner)
if strided:
diff --git a/test/registered/kernels/ops/diffusion/test_qknorm_rope.py b/test/registered/kernels/ops/diffusion/test_qknorm_rope.py
index bd617adb2..5033351e8 100644
--- a/test/registered/kernels/ops/diffusion/test_qknorm_rope.py
+++ b/test/registered/kernels/ops/diffusion/test_qknorm_rope.py
@@ -48,7 +48,7 @@ def split_qknorm_rope(
) -> None:
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
- from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
+ from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
fused_inplace_qknorm(q, k, q_weight, k_weight)
apply_rope_with_cos_sin_cache_inplace(
diff --git a/test/registered/kernels/ops/gemm/test_dsv3_fused_a_gemm.py b/test/registered/kernels/ops/gemm/test_dsv3_fused_a_gemm.py
index 128e8734a..d53c446d4 100644
--- a/test/registered/kernels/ops/gemm/test_dsv3_fused_a_gemm.py
+++ b/test/registered/kernels/ops/gemm/test_dsv3_fused_a_gemm.py
@@ -11,7 +11,7 @@ from sglang.kernels.jit.utils import (
get_jit_cuda_arch,
is_hip_runtime,
)
-from sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm import dsv3_fused_a_gemm
+from sglang.kernels.ops.gemm.dsv3_fused_a_gemm import dsv3_fused_a_gemm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
diff --git a/test/registered/kernels/ops/gemm/test_dsv3_router_gemm.py b/test/registered/kernels/ops/gemm/test_dsv3_router_gemm.py
index 00204776f..d5d4d76c0 100644
--- a/test/registered/kernels/ops/gemm/test_dsv3_router_gemm.py
+++ b/test/registered/kernels/ops/gemm/test_dsv3_router_gemm.py
@@ -11,7 +11,7 @@ from sglang.kernels.jit.utils import (
get_jit_cuda_arch,
is_hip_runtime,
)
-from sglang.kernels.ops.gemm._jit_dsv3_router_gemm import dsv3_router_gemm
+from sglang.kernels.ops.gemm.dsv3_router_gemm import dsv3_router_gemm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=37, stage="base-b-kernel-unit", runner_config="1-gpu-large")
diff --git a/test/registered/kernels/ops/kvcache/test_set_mla_kv_buffer.py b/test/registered/kernels/ops/kvcache/test_set_mla_kv_buffer.py
index 6ff6adec4..ac2947258 100644
--- a/test/registered/kernels/ops/kvcache/test_set_mla_kv_buffer.py
+++ b/test/registered/kernels/ops/kvcache/test_set_mla_kv_buffer.py
@@ -4,7 +4,7 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
-from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import (
+from sglang.kernels.ops.kvcache.set_mla_kv_buffer import (
can_use_set_mla_kv_buffer,
set_mla_kv_buffer,
)
diff --git a/test/registered/kernels/ops/layernorm/test_fused_add_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_fused_add_rmsnorm.py
index 1bb41dd92..0726b3675 100644
--- a/test/registered/kernels/ops/layernorm/test_fused_add_rmsnorm.py
+++ b/test/registered/kernels/ops/layernorm/test_fused_add_rmsnorm.py
@@ -20,7 +20,7 @@ def sglang_jit_fused_add_rmsnorm(
*,
cast_x_before_out_mul: bool = False,
) -> None:
- from sglang.kernels.ops.layernorm._jit_norm import fused_add_rmsnorm
+ from sglang.kernels.ops.layernorm.norm import fused_add_rmsnorm
fused_add_rmsnorm(
input, residual, weight, eps, cast_x_before_out_mul=cast_x_before_out_mul
diff --git a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py
index 4e13e7230..5fb7a1f09 100644
--- a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py
+++ b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py
@@ -21,6 +21,7 @@ GROUPS = [
"attention",
"communication",
"diffusion",
+ "elementwise",
"embeddings",
"gemm",
"grammar",
@@ -31,7 +32,6 @@ GROUPS = [
"moe",
"quantization",
"sampling",
- "spatial",
"speculative",
]
diff --git a/test/registered/kernels/ops/layernorm/test_qknorm.py b/test/registered/kernels/ops/layernorm/test_qknorm.py
index 0669782ee..d36b008fb 100644
--- a/test/registered/kernels/ops/layernorm/test_qknorm.py
+++ b/test/registered/kernels/ops/layernorm/test_qknorm.py
@@ -34,7 +34,7 @@ def sglang_jit_qknorm(
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
- from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
+ from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
fused_inplace_qknorm(q, k, q_weight, k_weight)
diff --git a/test/registered/kernels/ops/layernorm/test_qknorm_across_heads.py b/test/registered/kernels/ops/layernorm/test_qknorm_across_heads.py
index e2efb28ab..c7dbf5007 100644
--- a/test/registered/kernels/ops/layernorm/test_qknorm_across_heads.py
+++ b/test/registered/kernels/ops/layernorm/test_qknorm_across_heads.py
@@ -19,7 +19,7 @@ def sglang_jit_qknorm_across_heads(
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
- from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm_across_heads
+ from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm_across_heads
fused_inplace_qknorm_across_heads(q, k, q_weight, k_weight)
diff --git a/test/registered/kernels/ops/layernorm/test_rmsnorm.py b/test/registered/kernels/ops/layernorm/test_rmsnorm.py
index e4d726408..148e40641 100644
--- a/test/registered/kernels/ops/layernorm/test_rmsnorm.py
+++ b/test/registered/kernels/ops/layernorm/test_rmsnorm.py
@@ -26,7 +26,7 @@ def sglang_jit_rmsnorm(
output: torch.Tensor | None = None,
eps: float = EPS,
) -> None:
- from sglang.kernels.ops.layernorm._jit_norm import rmsnorm
+ from sglang.kernels.ops.layernorm.norm import rmsnorm
rmsnorm(input, weight, out=output, eps=eps)
@@ -127,7 +127,7 @@ def test_rmsnorm(
@pytest.mark.parametrize("hidden_size", [64, 128, 256, 512, 8192, 8704, 16384])
def test_rmsnorm_hidden_size_support(hidden_size: int) -> None:
- from sglang.kernels.ops.layernorm._jit_norm import _is_supported_rmsnorm_hidden_size
+ from sglang.kernels.ops.layernorm.norm import _is_supported_rmsnorm_hidden_size
assert _is_supported_rmsnorm_hidden_size(hidden_size)
@@ -148,7 +148,7 @@ def test_rmsnorm_hidden_size_support(hidden_size: int) -> None:
],
)
def test_rmsnorm_kernel_dispatch(hidden_size: int, expected: str) -> None:
- from sglang.kernels.ops.layernorm._jit_norm import _rmsnorm_kernel_class
+ from sglang.kernels.ops.layernorm.norm import _rmsnorm_kernel_class
assert _rmsnorm_kernel_class(hidden_size) == expected
diff --git a/test/registered/kernels/ops/model/test_inkling_rel_proj.py b/test/registered/kernels/ops/model/test_inkling_rel_proj.py
index 2ed7ae5f1..3f90b250d 100644
--- a/test/registered/kernels/ops/model/test_inkling_rel_proj.py
+++ b/test/registered/kernels/ops/model/test_inkling_rel_proj.py
@@ -6,7 +6,7 @@ production strided-r layout and contiguous inputs."""
import pytest
import torch
-from sglang.kernels.ops.model.inkling.inkling_rel_proj import rel_proj_small_t
+from sglang.kernels.ops.attention.inkling_rel_proj import rel_proj_small_t
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
diff --git a/test/registered/kernels/ops/quantization/test_per_tensor_quant_fp8.py b/test/registered/kernels/ops/quantization/test_per_tensor_quant_fp8.py
index 0af83ef3d..92cbf546e 100644
--- a/test/registered/kernels/ops/quantization/test_per_tensor_quant_fp8.py
+++ b/test/registered/kernels/ops/quantization/test_per_tensor_quant_fp8.py
@@ -6,7 +6,7 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
-from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
+from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.test.ci.ci_register import register_cuda_ci
diff --git a/test/registered/kernels/ops/quantization/test_per_token_group_quant.py b/test/registered/kernels/ops/quantization/test_per_token_group_quant.py
index d5ae71e33..576052f6c 100644
--- a/test/registered/kernels/ops/quantization/test_per_token_group_quant.py
+++ b/test/registered/kernels/ops/quantization/test_per_token_group_quant.py
@@ -22,14 +22,14 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
-from sglang.kernels.ops.quantization._jit_per_token_group_quant import (
- per_token_group_quant,
-)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
)
+from sglang.kernels.ops.quantization.per_token_group_quant import (
+ per_token_group_quant,
+)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=90, stage="base-b-kernel-unit", runner_config="1-gpu-large")
diff --git a/test/registered/kernels/ops/quantization/test_per_token_group_quant_8bit_v2.py b/test/registered/kernels/ops/quantization/test_per_token_group_quant_8bit_v2.py
index 4d7403c13..be19861fc 100644
--- a/test/registered/kernels/ops/quantization/test_per_token_group_quant_8bit_v2.py
+++ b/test/registered/kernels/ops/quantization/test_per_token_group_quant_8bit_v2.py
@@ -4,7 +4,7 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
-from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
+from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci