diff --git a/benchmark/kernels/bench_fused_gate_sigmoid_mul_add.py b/benchmark/kernels/bench_fused_gate_sigmoid_mul_add.py index 14ffe9465..be42c6834 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.srt.layers.elementwise import fused_gate_sigmoid_mul_add +from sglang.kernels.ops.layernorm.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 1f3cd6e67..b1d0c14ab 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.srt.layers.elementwise import fused_sigmoid_mul +from sglang.kernels.ops.layernorm.elementwise import fused_sigmoid_mul NUM_HEADS = 32 HEAD_DIM = 256 diff --git a/python/sglang/jit_kernel/dsv4/elementwise.py b/python/sglang/jit_kernel/dsv4/elementwise.py index 445d72eb3..b1395824d 100644 --- a/python/sglang/jit_kernel/dsv4/elementwise.py +++ b/python/sglang/jit_kernel/dsv4/elementwise.py @@ -123,7 +123,9 @@ def fused_rope_inplace( inverse: if True, apply inverse rotation (conjugate freqs) """ if _is_hip or _is_xpu: - from sglang.srt.layers.deepseek_v4_rope import apply_rotary_emb_triton + from sglang.kernels.ops.attention.deepseek_v4_rope import ( + apply_rotary_emb_triton, + ) apply_rotary_emb_triton(q, freqs_cis, positions=positions, inverse=inverse) if k is not None: diff --git a/python/sglang/kernels/ops/attention/__init__.py b/python/sglang/kernels/ops/attention/__init__.py index 1410c3070..7e46e7803 100644 --- a/python/sglang/kernels/ops/attention/__init__.py +++ b/python/sglang/kernels/ops/attention/__init__.py @@ -41,3 +41,24 @@ for _mod, _fn in _TRITON_KERNELS: del _mod, _fn __all__ = [] + + +# RoPE / QK-norm fusion kernels migrated from srt/layers top-level strays +# (RFC #29630, Phase 2.5); registered for inventory. +for _mod, _fn in [ + ("deepseek_v4_rope", "precompute_freqs_cis"), + ("fused_qk_norm_rope_store", "fused_qk_norm_rope_swa_store"), + ("fused_qk_rmsnorm_rope_gate", "fused_qk_gemma_rmsnorm_rope_gate"), + ("fused_qk_norm", "fused_qk_norm"), + ("rotary_triton", "triton_mrope_fused"), + ("rotary_triton", "triton_ernie45_rope_fused_inplace"), + ("mrope", "apply_interleaved_rope_triton"), +]: + register_kernel( + KernelSpec( + op=f"attention.{_fn}", + backend=KernelBackend.TRITON, + target=f"sglang.kernels.ops.attention.{_mod}:{_fn}", + ) + ) +del _mod, _fn diff --git a/python/sglang/srt/layers/deepseek_v4_rope.py b/python/sglang/kernels/ops/attention/deepseek_v4_rope.py similarity index 100% rename from python/sglang/srt/layers/deepseek_v4_rope.py rename to python/sglang/kernels/ops/attention/deepseek_v4_rope.py diff --git a/python/sglang/srt/layers/fused_qk_norm.py b/python/sglang/kernels/ops/attention/fused_qk_norm.py similarity index 100% rename from python/sglang/srt/layers/fused_qk_norm.py rename to python/sglang/kernels/ops/attention/fused_qk_norm.py diff --git a/python/sglang/srt/layers/fused_qk_norm_rope_store.py b/python/sglang/kernels/ops/attention/fused_qk_norm_rope_store.py similarity index 100% rename from python/sglang/srt/layers/fused_qk_norm_rope_store.py rename to python/sglang/kernels/ops/attention/fused_qk_norm_rope_store.py diff --git a/python/sglang/srt/layers/fused_qk_rmsnorm_rope_gate.py b/python/sglang/kernels/ops/attention/fused_qk_rmsnorm_rope_gate.py similarity index 100% rename from python/sglang/srt/layers/fused_qk_rmsnorm_rope_gate.py rename to python/sglang/kernels/ops/attention/fused_qk_rmsnorm_rope_gate.py diff --git a/python/sglang/kernels/ops/attention/mrope.py b/python/sglang/kernels/ops/attention/mrope.py new file mode 100644 index 000000000..71e48155b --- /dev/null +++ b/python/sglang/kernels/ops/attention/mrope.py @@ -0,0 +1,89 @@ +"""Interleaved M-RoPE Triton kernel, migrated from +``sglang.srt.layers.rotary_embedding.mrope`` (RFC #29630, Phase 2.5). +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def apply_interleaved_rope_kernel( + x_ptr, + out_ptr, + S: tl.constexpr, + D: tl.constexpr, + stride_x_m, + stride_x_s, + stride_out_s, + section_1_end, + section_2_end, + BLOCK_S: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + start_s = tl.program_id(0) * BLOCK_S + s_offsets = start_s + tl.arange(0, BLOCK_S) + + dim_offset = tl.program_id(1) * BLOCK_SIZE + dim_indices = dim_offset + tl.arange(0, BLOCK_SIZE) + + mask_s = s_offsets < S + mask_d = dim_indices < D + mask = mask_s[:, None] & mask_d[None, :] + + val_ptr = ( + x_ptr + 0 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :] + ) + val = tl.load(val_ptr, mask=mask, other=0.0) + + cond_a = (dim_indices[None, :] % 3 == 1) & ( + dim_indices[None, :] < section_1_end * 3 + ) + val_a_ptr = ( + x_ptr + 1 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :] + ) + val_a = tl.load(val_a_ptr, mask=mask & cond_a, other=0.0) + + cond_b = (dim_indices[None, :] % 3 == 2) & ( + dim_indices[None, :] < section_2_end * 3 + ) + val_b_ptr = ( + x_ptr + 2 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :] + ) + val_b = tl.load(val_b_ptr, mask=mask & cond_b, other=0.0) + + val = tl.where(cond_a, val_a, val) + val = tl.where(cond_b, val_b, val) + + out_ptr = out_ptr + s_offsets[:, None] * stride_out_s + dim_indices[None, :] + tl.store(out_ptr, val, mask=mask) + + +def apply_interleaved_rope_triton(x: torch.Tensor, mrope_section: list) -> torch.Tensor: + x = x.contiguous() + M, S, D = x.shape + + out = torch.empty((S, D), dtype=x.dtype, device=x.device) + + BLOCK_S = 64 + BLOCK_SIZE = 128 + + grid = (triton.cdiv(S, BLOCK_S), triton.cdiv(D, BLOCK_SIZE)) + + section_1_end = mrope_section[1] + section_2_end = mrope_section[2] + + apply_interleaved_rope_kernel[grid]( + x, + out, + S, + D, + x.stride(0), + x.stride(1), + out.stride(0), + section_1_end, + section_2_end, + BLOCK_S=BLOCK_S, + BLOCK_SIZE=BLOCK_SIZE, + ) + return out diff --git a/python/sglang/srt/layers/rotary_embedding/triton_kernels.py b/python/sglang/kernels/ops/attention/rotary_triton.py similarity index 100% rename from python/sglang/srt/layers/rotary_embedding/triton_kernels.py rename to python/sglang/kernels/ops/attention/rotary_triton.py diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index ae22499ea..a6e0fab43 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -336,3 +336,26 @@ __all__ = [ "gemma_rmsnorm", "gemma_fused_add_rmsnorm", ] + + +from sglang.kernels.registry import register_kernel +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"), +] +for _mod, _fn, _bk in _PHASE25_KERNELS: + register_kernel( + KernelSpec( + op=f"layernorm.{_fn}", + backend=KernelBackend(_bk), + target=f"sglang.kernels.ops.layernorm.{_mod}:{_fn}", + ) + ) +del _mod, _fn, _bk diff --git a/python/sglang/srt/layers/elementwise.py b/python/sglang/kernels/ops/layernorm/elementwise.py similarity index 100% rename from python/sglang/srt/layers/elementwise.py rename to python/sglang/kernels/ops/layernorm/elementwise.py diff --git a/python/sglang/srt/layers/gemma4_fused_ops.py b/python/sglang/kernels/ops/layernorm/gemma4_fused_ops.py similarity index 100% rename from python/sglang/srt/layers/gemma4_fused_ops.py rename to python/sglang/kernels/ops/layernorm/gemma4_fused_ops.py diff --git a/python/sglang/srt/layers/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py similarity index 100% rename from python/sglang/srt/layers/mhc.py rename to python/sglang/kernels/ops/layernorm/mhc.py diff --git a/python/sglang/srt/layers/mhc_head.py b/python/sglang/kernels/ops/layernorm/mhc_head.py similarity index 100% rename from python/sglang/srt/layers/mhc_head.py rename to python/sglang/kernels/ops/layernorm/mhc_head.py diff --git a/python/sglang/kernels/ops/memory/__init__.py b/python/sglang/kernels/ops/memory/__init__.py index 37f2967c6..1e9ae9d1e 100644 --- a/python/sglang/kernels/ops/memory/__init__.py +++ b/python/sglang/kernels/ops/memory/__init__.py @@ -28,3 +28,20 @@ for _mod, _fn in _TRITON_KERNELS: del _mod, _fn __all__ = [] + + +# Migrated from srt/layers (RFC #29630, Phase 2.5). +register_kernel( + KernelSpec( + op="memory.gpu_tensor_hash", + backend=KernelBackend.TRITON, + target="sglang.kernels.ops.memory.gpu_tensor_hash:gpu_tensor_hash", + ) +) +register_kernel( + KernelSpec( + op="memory.memcpy_triton", + backend=KernelBackend.TRITON, + target="sglang.kernels.ops.memory.memcpy_triton:memcpy_triton", + ) +) diff --git a/python/sglang/srt/layers/multimodal.py b/python/sglang/kernels/ops/memory/gpu_tensor_hash.py similarity index 100% rename from python/sglang/srt/layers/multimodal.py rename to python/sglang/kernels/ops/memory/gpu_tensor_hash.py diff --git a/python/sglang/kernels/ops/memory/memcpy_triton.py b/python/sglang/kernels/ops/memory/memcpy_triton.py new file mode 100644 index 000000000..5756f165c --- /dev/null +++ b/python/sglang/kernels/ops/memory/memcpy_triton.py @@ -0,0 +1,49 @@ +"""Offset/size-driven device memcpy kernel, migrated from +``sglang.srt.layers.dp_attention`` (RFC #29630, Phase 2.5). +""" + +import functools + +import triton +import triton.language as tl + + +@triton.jit +def memcpy_triton_kernel( + dst_ptr, + src_ptr, + offset_ptr, + sz_ptr, + offset_src: tl.constexpr, + chunk_size, # multiplied for offset and sz + BLOCK_SIZE: tl.constexpr, +): + pid = tl.program_id(axis=0).to(tl.int64) + offset = tl.load(offset_ptr).to(tl.int64) * chunk_size + sz = tl.load(sz_ptr).to(tl.int64) * chunk_size + + start_index = pid * BLOCK_SIZE + offs = tl.arange(0, BLOCK_SIZE) + mask = start_index + offs < sz + + if offset_src: + data = tl.load(src_ptr + offset + start_index + offs, mask=mask) + tl.store(dst_ptr + start_index + offs, data, mask=mask) + else: + data = tl.load(src_ptr + start_index + offs, mask=mask) + tl.store(dst_ptr + offset + start_index + offs, data, mask=mask) + + +def prod(x): + return functools.reduce(lambda a, b: a * b, x, 1) + + +def memcpy_triton(dst, src, dim, offset, sz, offset_src): + max_size = min(src.numel(), dst.numel()) + assert dim == 0, "dim != 0 unsupported" + assert src.shape[1:] == dst.shape[1:], "src and dst must have same shape" + chunk_size = prod(src.shape[1:]) + BLOCK_SIZE = 8192 + grid = (triton.cdiv(max_size, BLOCK_SIZE),) + + memcpy_triton_kernel[grid](dst, src, offset, sz, offset_src, chunk_size, BLOCK_SIZE) diff --git a/python/sglang/kernels/ops/sampling/__init__.py b/python/sglang/kernels/ops/sampling/__init__.py index 8f9587e58..00694f554 100644 --- a/python/sglang/kernels/ops/sampling/__init__.py +++ b/python/sglang/kernels/ops/sampling/__init__.py @@ -54,3 +54,13 @@ def top_p_renorm_probs( __all__ = ["top_k_renorm_probs", "top_p_renorm_probs"] + + +# Migrated from srt/layers/utils/hash.py (RFC #29630, Phase 2.5). +register_kernel( + KernelSpec( + op="sampling.murmur_hash32", + backend=KernelBackend.TRITON, + target="sglang.kernels.ops.sampling.murmur_hash:murmur_hash32", + ) +) diff --git a/python/sglang/srt/layers/utils/hash.py b/python/sglang/kernels/ops/sampling/murmur_hash.py similarity index 100% rename from python/sglang/srt/layers/utils/hash.py rename to python/sglang/kernels/ops/sampling/murmur_hash.py diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py index 90df5446c..d45ec4fda 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py @@ -312,7 +312,7 @@ class CompressorAscendBackendMixin(CompressorBackendMixin): ): return self._forward_compress_native(compressor, x, forward_batch) - from sglang.srt.layers.deepseek_v4_rope import ( + from sglang.kernels.ops.attention.deepseek_v4_rope import ( get_fused_compressor_rope_cos_sin, ) @@ -649,7 +649,7 @@ class CompressorAscendBackendMixin(CompressorBackendMixin): # Use the same contig cache as the outer rope path; .real/.imag on a # complex tensor are strided views and aclnnIndex over them triggers # StridedSlice (see _get_contig_freqs_real_imag in deepseek_v4_rope.py). - from sglang.srt.layers.deepseek_v4_rope import ( + from sglang.kernels.ops.attention.deepseek_v4_rope import ( _get_contig_freqs_real_imag, ) @@ -890,7 +890,7 @@ class C4IndexerAscendBackendMixin(C4IndexerBackendMixin): def _compute_q_npu( self, c4_indexer, q_lora: torch.Tensor, positions: torch.Tensor ) -> torch.Tensor: - from sglang.srt.layers.deepseek_v4_rope import v4_rope_inplace_npu + from sglang.kernels.ops.attention.deepseek_v4_rope import v4_rope_inplace_npu bs = q_lora.shape[0] q, _ = c4_indexer.wq_b(q_lora) diff --git a/python/sglang/srt/layers/attention/dsv4/compress_hip.py b/python/sglang/srt/layers/attention/dsv4/compress_hip.py index 004d49c5b..370eb8606 100644 --- a/python/sglang/srt/layers/attention/dsv4/compress_hip.py +++ b/python/sglang/srt/layers/attention/dsv4/compress_hip.py @@ -9,6 +9,11 @@ import torch.nn as nn import triton import triton.language as tl +from sglang.kernels.ops.attention.deepseek_v4_rope import ( + apply_rotary_emb_triton, + fused_norm_rope_inplace_triton, + fused_softmax_pool_triton, +) from sglang.srt.environ import envs from sglang.srt.layers.attention.dsa.dsa_indexer import rotate_activation from sglang.srt.layers.attention.dsv4.compressor import Compressor as _CompressorBase @@ -16,14 +21,9 @@ from sglang.srt.layers.attention.dsv4.fused_compress_triton import ( fused_ape_pool_norm_rope, ) from sglang.srt.layers.attention.nsa.nsa_indexer import rotate_activation -from sglang.srt.layers.deepseek_v4_rope import ( - apply_rotary_emb_triton, - fused_norm_rope_inplace_triton, - fused_softmax_pool_triton, -) try: - from sglang.srt.layers.deepseek_v4_rope import fused_softmax_pool_triton + from sglang.kernels.ops.attention.deepseek_v4_rope import fused_softmax_pool_triton except ImportError: fused_softmax_pool_triton = None from sglang.srt.mem_cache.deepseek_v4_compress_state import ( diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index 908545ff9..a783eac8a 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -562,12 +562,14 @@ class CompressorBackendMixin: layer_id: int, ) -> None: """HIP-specific forward path using PyTorch/Triton fallbacks.""" + from sglang.kernels.ops.attention.deepseek_v4_rope import ( + fused_norm_rope_inplace_triton, + ) from sglang.srt.layers.attention.dsv4.quant_k_cache import ( quant_to_nope_fp8_rope_bf16_pack_triton, ) from sglang.srt.layers.attention.nsa.nsa_indexer import rotate_activation from sglang.srt.layers.attention.nsa.triton_kernel import act_quant - from sglang.srt.layers.deepseek_v4_rope import fused_norm_rope_inplace_triton compress_ratio = compressor.ratio head_dim = compressor.head_dim diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 8432a62a9..99762aa78 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -1,14 +1,11 @@ from __future__ import annotations -import functools import logging from contextlib import contextmanager from enum import IntEnum, auto from typing import TYPE_CHECKING, List, Optional, Tuple import torch -import triton -import triton.language as tl from sglang.srt.distributed import ( GroupCoordinator, @@ -376,45 +373,7 @@ def get_dp_local_slice_cpu( return local_start_pos, local_num_tokens -@triton.jit -def memcpy_triton_kernel( - dst_ptr, - src_ptr, - offset_ptr, - sz_ptr, - offset_src: tl.constexpr, - chunk_size, # multiplied for offset and sz - BLOCK_SIZE: tl.constexpr, -): - pid = tl.program_id(axis=0).to(tl.int64) - offset = tl.load(offset_ptr).to(tl.int64) * chunk_size - sz = tl.load(sz_ptr).to(tl.int64) * chunk_size - - start_index = pid * BLOCK_SIZE - offs = tl.arange(0, BLOCK_SIZE) - mask = start_index + offs < sz - - if offset_src: - data = tl.load(src_ptr + offset + start_index + offs, mask=mask) - tl.store(dst_ptr + start_index + offs, data, mask=mask) - else: - data = tl.load(src_ptr + start_index + offs, mask=mask) - tl.store(dst_ptr + offset + start_index + offs, data, mask=mask) - - -def prod(x): - return functools.reduce(lambda a, b: a * b, x, 1) - - -def memcpy_triton(dst, src, dim, offset, sz, offset_src): - max_size = min(src.numel(), dst.numel()) - assert dim == 0, "dim != 0 unsupported" - assert src.shape[1:] == dst.shape[1:], "src and dst must have same shape" - chunk_size = prod(src.shape[1:]) - BLOCK_SIZE = 8192 - grid = (triton.cdiv(max_size, BLOCK_SIZE),) - - memcpy_triton_kernel[grid](dst, src, offset, sz, offset_src, chunk_size, BLOCK_SIZE) +from sglang.kernels.ops.memory.memcpy_triton import memcpy_triton def _dp_gather_via_all_reduce( diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index 609488e7d..575e34a12 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -7,11 +7,11 @@ from typing import List, Optional, Tuple import torch -from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding -from sglang.srt.layers.rotary_embedding.triton_kernels import ( +from sglang.kernels.ops.attention.rotary_triton import ( triton_ernie45_rope_fused_inplace, triton_mrope_fused, ) +from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.layers.rotary_embedding.yarn import ( yarn_find_correction_range, @@ -41,94 +41,10 @@ if _is_npu: if _is_xpu: from sgl_kernel import multimodal_rotary_embedding -import triton -import triton.language as tl - +from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton from sglang.srt.runtime_context import get_server_args -@triton.jit -def apply_interleaved_rope_kernel( - x_ptr, - out_ptr, - S: tl.constexpr, - D: tl.constexpr, - stride_x_m, - stride_x_s, - stride_out_s, - section_1_end, - section_2_end, - BLOCK_S: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - start_s = tl.program_id(0) * BLOCK_S - s_offsets = start_s + tl.arange(0, BLOCK_S) - - dim_offset = tl.program_id(1) * BLOCK_SIZE - dim_indices = dim_offset + tl.arange(0, BLOCK_SIZE) - - mask_s = s_offsets < S - mask_d = dim_indices < D - mask = mask_s[:, None] & mask_d[None, :] - - val_ptr = ( - x_ptr + 0 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :] - ) - val = tl.load(val_ptr, mask=mask, other=0.0) - - cond_a = (dim_indices[None, :] % 3 == 1) & ( - dim_indices[None, :] < section_1_end * 3 - ) - val_a_ptr = ( - x_ptr + 1 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :] - ) - val_a = tl.load(val_a_ptr, mask=mask & cond_a, other=0.0) - - cond_b = (dim_indices[None, :] % 3 == 2) & ( - dim_indices[None, :] < section_2_end * 3 - ) - val_b_ptr = ( - x_ptr + 2 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :] - ) - val_b = tl.load(val_b_ptr, mask=mask & cond_b, other=0.0) - - val = tl.where(cond_a, val_a, val) - val = tl.where(cond_b, val_b, val) - - out_ptr = out_ptr + s_offsets[:, None] * stride_out_s + dim_indices[None, :] - tl.store(out_ptr, val, mask=mask) - - -def apply_interleaved_rope_triton(x: torch.Tensor, mrope_section: list) -> torch.Tensor: - x = x.contiguous() - M, S, D = x.shape - - out = torch.empty((S, D), dtype=x.dtype, device=x.device) - - BLOCK_S = 64 - BLOCK_SIZE = 128 - - grid = (triton.cdiv(S, BLOCK_S), triton.cdiv(D, BLOCK_SIZE)) - - section_1_end = mrope_section[1] - section_2_end = mrope_section[2] - - apply_interleaved_rope_kernel[grid]( - x, - out, - S, - D, - x.stride(0), - x.stride(1), - out.stride(0), - section_1_end, - section_2_end, - BLOCK_S=BLOCK_S, - BLOCK_SIZE=BLOCK_SIZE, - ) - return out - - def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor: x_t = x[0].clone() x_t[..., 1 : mrope_section[1] * 3 : 3] = x[1, ..., 1 : mrope_section[1] * 3 : 3] diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index d28917eb5..dcbdc27ed 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -5,12 +5,12 @@ import torch import torch.distributed as dist from torch import nn +from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32 from sglang.srt.distributed import get_tp_group from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) from sglang.srt.layers.logits_processor import LogitsProcessorOutput -from sglang.srt.layers.utils.hash import murmur_hash32 from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 137266df4..545a70b56 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -16,8 +16,8 @@ import numpy as np import torch from torch import nn +from sglang.kernels.ops.memory.gpu_tensor_hash import gpu_tensor_hash from sglang.srt.environ import envs -from sglang.srt.layers.multimodal import gpu_tensor_hash from sglang.srt.managers.io_struct import ( BaseBatchReq, TokenizedEmbeddingReqInput, diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 5773395a0..00377253a 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -25,6 +25,9 @@ from sglang.jit_kernel.dsv4 import ( fused_rope_inplace, sglang_per_token_group_quant_fp8_dsv4_wo_a, ) +from sglang.kernels.ops.attention.deepseek_v4_rope import ( + v4_rope_inplace_npu, +) from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config from sglang.srt.distributed import ( @@ -47,9 +50,6 @@ from sglang.srt.layers.communicator_dsa_cp import ( dsa_cp_gather_hidden_states, dsa_cp_reduce_scatter_hidden_states, ) -from sglang.srt.layers.deepseek_v4_rope import ( - v4_rope_inplace_npu, -) from sglang.srt.layers.dp_attention import ( _tbo_event, attn_tp_all_gather, @@ -132,7 +132,11 @@ if not _is_hip: if _is_xpu: from sgl_kernel import hc_split_sinkhorn else: - from sglang.srt.layers.mhc import hc_split_sinkhorn, mhc_fused_post_pre, npu_hc_pre + from sglang.kernels.ops.layernorm.mhc import ( + hc_split_sinkhorn, + mhc_fused_post_pre, + npu_hc_pre, + ) from sglang.srt.utils import ( LazyValue, @@ -476,7 +480,7 @@ class MqaAttentionBase(nn.Module): tp_size=self.attn_tp_size, ) - from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis + from sglang.kernels.ops.attention.deepseek_v4_rope import precompute_freqs_cis rope_theta, rope_scaling = get_rope_config(config) self.rope_scaling = rope_scaling @@ -801,7 +805,7 @@ class MQALayer(MqaAttentionBase): else self.wkv(x_linear)[0] ) - from sglang.srt.layers.fused_qk_norm_rope_store import ( + from sglang.kernels.ops.attention.fused_qk_norm_rope_store import ( fused_qk_norm_rope_swa_store, ) @@ -915,7 +919,7 @@ class MQALayer(MqaAttentionBase): False, ) - from sglang.srt.layers.fused_qk_norm_rope_store import ( + from sglang.kernels.ops.attention.fused_qk_norm_rope_store import ( fused_qk_norm_rope_swa_store, ) @@ -1366,7 +1370,7 @@ class DeepseekV4DecoderLayer(nn.Module): return y, post, comb, False if envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get(): - from sglang.srt.layers.mhc import mhc_pre + from sglang.kernels.ops.layernorm.mhc import mhc_pre norm_kwargs = {} if norm is not None: @@ -1450,7 +1454,7 @@ class DeepseekV4DecoderLayer(nn.Module): return torch.ops.custom.npu_hc_post(x, residual, post, comb) if envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get(): - from sglang.srt.layers.mhc import mhc_post + from sglang.kernels.ops.layernorm.mhc import mhc_post return mhc_post(x, residual, post, comb) @@ -2048,7 +2052,7 @@ class DeepseekV4Model(nn.Module): hc_base: torch.Tensor, ): if x.numel() > 0: - from sglang.srt.layers.mhc_head import fused_hc_head + from sglang.kernels.ops.layernorm.mhc_head import fused_hc_head return fused_hc_head( x.contiguous(), @@ -2315,7 +2319,7 @@ class DeepseekV4ForCausalLM(nn.Module): # batched/contiguous-load rope kernels (faster on gfx95; . # Module-level toggles default OFF; flipped True here for DSV4 if _is_hip: - from sglang.srt.layers.deepseek_v4_rope import set_batched_rope + from sglang.kernels.ops.attention.deepseek_v4_rope import set_batched_rope from sglang.srt.layers.quantization.fp8_utils import set_force_ck_w8a8 set_force_ck_w8a8(True) @@ -2600,7 +2604,7 @@ class DeepseekV4ForCausalLM(nn.Module): if layer is None: return - from sglang.srt.layers.mhc import prewarm_mhc_pre + from sglang.kernels.ops.layernorm.mhc import prewarm_mhc_pre tic = time.perf_counter() prewarm_mhc_pre( diff --git a/python/sglang/srt/models/gemma4_causal.py b/python/sglang/srt/models/gemma4_causal.py index b8c800f5c..f2d9a645e 100644 --- a/python/sglang/srt/models/gemma4_causal.py +++ b/python/sglang/srt/models/gemma4_causal.py @@ -24,16 +24,16 @@ from transformers import ( PreTrainedModel, ) -from sglang.srt.distributed import ( - get_pp_group, -) -from sglang.srt.layers.gemma4_fused_ops import ( +from sglang.kernels.ops.layernorm.gemma4_fused_ops import ( gemma4_fused_routing, gemma_dual_rmsnorm_residual_scalar, gemma_qkv_rmsnorm, gemma_rmsnorm_residual_scalar, gemma_routing_post_topk, ) +from sglang.srt.distributed import ( + get_pp_group, +) from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm from sglang.srt.layers.linear import ( QKVParallelLinear, diff --git a/python/sglang/srt/models/grok.py b/python/sglang/srt/models/grok.py index 8d583afbc..f06912025 100644 --- a/python/sglang/srt/models/grok.py +++ b/python/sglang/srt/models/grok.py @@ -22,16 +22,16 @@ import torch.nn.functional as F from torch import nn from transformers import PretrainedConfig +from sglang.kernels.ops.layernorm.elementwise import ( + fused_dual_residual_rmsnorm, + fused_rmsnorm, + gelu_and_mul_triton, +) from sglang.kernels.ops.moe.router import fused_moe_router_shim from sglang.srt.distributed import ( tensor_model_parallel_all_reduce, ) from sglang.srt.layers.activation import GeluAndMul -from sglang.srt.layers.elementwise import ( - fused_dual_residual_rmsnorm, - fused_rmsnorm, - gelu_and_mul_triton, -) from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 0ce5b00ad..72e8ddee4 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -27,6 +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.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.distributed import ( get_pp_group, @@ -48,7 +49,6 @@ from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) -from sglang.srt.layers.elementwise import fused_gate_sigmoid_mul_add from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index bb1c90979..e5d386f90 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -25,6 +25,7 @@ import triton from sglang.jit_kernel.triton.gdn_fused_proj import ( fused_qkvzba_split_reshape_cat_contiguous, ) +from sglang.kernels.ops.layernorm.elementwise import fused_sigmoid_mul # Configs from sglang.srt.configs.qwen3_5 import ( @@ -45,7 +46,6 @@ from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) -from sglang.srt.layers.elementwise import fused_sigmoid_mul # Layers - Others from sglang.srt.layers.layernorm import GemmaRMSNorm @@ -142,7 +142,7 @@ def _disable_shared_experts_fusion() -> bool: if _is_cuda: - from sglang.srt.layers.fused_qk_rmsnorm_rope_gate import ( + from sglang.kernels.ops.attention.fused_qk_rmsnorm_rope_gate import ( fused_qk_gemma_rmsnorm_rope_gate, ) 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 bf46806d3..cc3fccc4a 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.srt.layers.elementwise import fused_gate_sigmoid_mul_add +from sglang.kernels.ops.layernorm.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 a585254a2..3d5cbf44d 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.srt.layers.elementwise import fused_sigmoid_mul +from sglang.kernels.ops.layernorm.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/jit/deepseek_v4/test_fp4_indexer.py b/test/registered/jit/deepseek_v4/test_fp4_indexer.py index 39de4ef0c..de332658c 100644 --- a/test/registered/jit/deepseek_v4/test_fp4_indexer.py +++ b/test/registered/jit/deepseek_v4/test_fp4_indexer.py @@ -11,14 +11,14 @@ from sglang.jit_kernel.dsv4 import ( fused_q_indexer_rope_hadamard_fp4_quant, ) from sglang.jit_kernel.hadamard import hadamard_transform +from sglang.kernels.ops.attention.deepseek_v4_rope import ( + apply_rotary_emb_triton, + precompute_freqs_cis, +) from sglang.srt.layers.attention.dsv4.fp4_indexer import ( quantize_fp4_indexer_tensor, store_fp4_index_k_cache, ) -from sglang.srt.layers.deepseek_v4_rope import ( - apply_rotary_emb_triton, - precompute_freqs_cis, -) from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/kernels/test_gemma4_fused_routing.py b/test/registered/kernels/test_gemma4_fused_routing.py index 487ab8200..bb94c4784 100644 --- a/test/registered/kernels/test_gemma4_fused_routing.py +++ b/test/registered/kernels/test_gemma4_fused_routing.py @@ -26,7 +26,7 @@ pytestmark = pytest.mark.skipif( @pytest.fixture(scope="module") def fused_routing(): - from sglang.srt.layers.gemma4_fused_ops import gemma4_fused_routing + from sglang.kernels.ops.layernorm.gemma4_fused_ops import gemma4_fused_routing return gemma4_fused_routing diff --git a/test/registered/kernels/test_mhc_kernels.py b/test/registered/kernels/test_mhc_kernels.py index 8f646b785..91c2b30dd 100644 --- a/test/registered/kernels/test_mhc_kernels.py +++ b/test/registered/kernels/test_mhc_kernels.py @@ -1,8 +1,8 @@ import pytest import torch -import sglang.srt.layers.mhc as mhc -from sglang.srt.layers.mhc import mhc_fused_post_pre, mhc_post, mhc_pre +import sglang.kernels.ops.layernorm.mhc as mhc +from sglang.kernels.ops.layernorm.mhc import mhc_fused_post_pre, mhc_post, mhc_pre from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-large")