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