diff --git a/python/sglang/kernels/ops/attention/__init__.py b/python/sglang/kernels/ops/attention/__init__.py index 5d10a0662..dc9dcf962 100644 --- a/python/sglang/kernels/ops/attention/__init__.py +++ b/python/sglang/kernels/ops/attention/__init__.py @@ -43,6 +43,28 @@ del _mod, _fn __all__ = [] +# DeepSeek DSA / DSV4 kernels migrated in Phase 2.5 (RFC #29630); +# registered for inventory. Import them from their modules. +for _mod, _fn in [ + ("dsa.triton_sparse_mla", "triton_sparse_mla_fwd"), + ("dsa.transform_index", "transform_index_page_table_prefill"), + ("dsa.transform_index", "transform_index_page_table_decode"), + ("dsa.cp_split", "dsa_cp_round_robin_split_q_seqs_kernel"), + ("dsv4.fp4_indexer", "quantize_fp4_indexer_tensor"), + ("dsv4.fp4_indexer", "store_fp4_index_k_cache"), + ("dsv4.fused_scale", "fused_scale"), + ("dsv4.rms_normalize_hip", "rms_normalize_triton"), + ("dsv4.compress_c128_hip", "_compress_forward_c128_triton"), +]: + register_kernel( + KernelSpec( + op=f"attention.{_fn.lstrip('_')}", + backend=KernelBackend.TRITON, + target=f"sglang.kernels.ops.attention.{_mod}:{_fn}", + ) + ) +del _mod, _fn + # Generic attention kernels migrated in Phase 2.5 (RFC #29630). for _mod, _fn in [ ("utils", "mla_quantize_and_rope_for_fp8"), diff --git a/python/sglang/kernels/ops/attention/dsa/__init__.py b/python/sglang/kernels/ops/attention/dsa/__init__.py new file mode 100644 index 000000000..68cae4146 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsa/__init__.py @@ -0,0 +1 @@ +"""DeepSeek DSA kernels (RFC #29630, Phase 2.5).""" diff --git a/python/sglang/kernels/ops/attention/dsa/cp_split.py b/python/sglang/kernels/ops/attention/dsa/cp_split.py new file mode 100644 index 000000000..d20f3aca3 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsa/cp_split.py @@ -0,0 +1,29 @@ +"""Round-robin CP q-sequence split kernel for DSA prefill. + +Migrated from ``sglang.srt.layers.attention.dsa.utils`` (RFC #29630, Phase 2.5). +""" + +import triton +import triton.language as tl + + +@triton.jit +def dsa_cp_round_robin_split_q_seqs_kernel( + in_seqs_ptr, + out_seqs_ptr, + bs_idx_ptr, + tokens: tl.constexpr, + cp_size: tl.constexpr, + cp_rank: tl.constexpr, +): + extra_seq = 0 + bs_idx = 0 + for bs in range(tokens): + cur_len = tl.load(in_seqs_ptr + bs) + cur_len += extra_seq + cur_seq = cur_len // cp_size + (cur_len % cp_size > cp_rank) + if cur_seq > 0: + tl.store(bs_idx_ptr + bs_idx, bs) + tl.store(out_seqs_ptr + bs_idx, cur_seq) + bs_idx += 1 + extra_seq = cur_len - cur_seq * cp_size diff --git a/python/sglang/srt/layers/attention/dsa/dequant_k_cache.py b/python/sglang/kernels/ops/attention/dsa/dequant_k_cache.py similarity index 100% rename from python/sglang/srt/layers/attention/dsa/dequant_k_cache.py rename to python/sglang/kernels/ops/attention/dsa/dequant_k_cache.py diff --git a/python/sglang/srt/layers/attention/dsa/index_buf_accessor.py b/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py similarity index 100% rename from python/sglang/srt/layers/attention/dsa/index_buf_accessor.py rename to python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py diff --git a/python/sglang/srt/layers/attention/dsa/quant_k_cache.py b/python/sglang/kernels/ops/attention/dsa/quant_k_cache.py similarity index 100% rename from python/sglang/srt/layers/attention/dsa/quant_k_cache.py rename to python/sglang/kernels/ops/attention/dsa/quant_k_cache.py diff --git a/python/sglang/srt/layers/attention/dsa/tilelang_kernel.py b/python/sglang/kernels/ops/attention/dsa/tilelang_kernel.py similarity index 100% rename from python/sglang/srt/layers/attention/dsa/tilelang_kernel.py rename to python/sglang/kernels/ops/attention/dsa/tilelang_kernel.py diff --git a/python/sglang/srt/layers/attention/dsa/transform_index.py b/python/sglang/kernels/ops/attention/dsa/transform_index.py similarity index 100% rename from python/sglang/srt/layers/attention/dsa/transform_index.py rename to python/sglang/kernels/ops/attention/dsa/transform_index.py diff --git a/python/sglang/srt/layers/attention/dsa/triton_kernel.py b/python/sglang/kernels/ops/attention/dsa/triton_kernel.py similarity index 100% rename from python/sglang/srt/layers/attention/dsa/triton_kernel.py rename to python/sglang/kernels/ops/attention/dsa/triton_kernel.py diff --git a/python/sglang/srt/layers/attention/dsa/triton_sparse_mla.py b/python/sglang/kernels/ops/attention/dsa/triton_sparse_mla.py similarity index 100% rename from python/sglang/srt/layers/attention/dsa/triton_sparse_mla.py rename to python/sglang/kernels/ops/attention/dsa/triton_sparse_mla.py diff --git a/python/sglang/kernels/ops/attention/dsv4/__init__.py b/python/sglang/kernels/ops/attention/dsv4/__init__.py new file mode 100644 index 000000000..9ad2409ba --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/__init__.py @@ -0,0 +1 @@ +"""DeepSeek-V4 attention kernels (RFC #29630, Phase 2.5).""" diff --git a/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py b/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py new file mode 100644 index 000000000..6838ebe8b --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py @@ -0,0 +1,292 @@ +"""HIP c128 compression kernels for DSV4 (RFC #29630, Phase 2.5). + +Migrated from ``sglang.srt.layers.attention.dsv4.compressor_v2``; the +kernels are defined under an ``is_hip`` guard exactly as before. +""" + +from __future__ import annotations + +from typing import Union + +import torch + +from sglang.jit_kernel.dsv4 import ( + CompressorDecodePlan, + CompressorPrefillPlan, +) +from sglang.jit_kernel.utils import is_hip_runtime + +_is_hip = is_hip_runtime() + +if _is_hip: + import triton + import triton.language as tl + + @triton.jit + def _c128_compress_decode_kernel( + buf_ptr, + input_ptr, + ape_ptr, + out_ptr, + plan_ptr, + buf_stride_slot, + input_stride_b, + ape_stride_r, + out_stride_b, + bs, + HEAD_DIM: tl.constexpr, + BLOCK_D: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, + ): + """Fused C128 decode: write to state buffer + online softmax-pool. + + plan_ptr points to int32 view: [bs, 4] where each row is + {seq_len, write_loc, read_page_0, read_page_1}. + """ + bid = tl.program_id(0) + if bid >= bs: + return + + # Parse plan + plan_base = plan_ptr + bid * 4 + seq_len = tl.load(plan_base).to(tl.int32) + write_loc = tl.load(plan_base + 1).to(tl.int32) + read_page_0 = tl.load(plan_base + 2).to(tl.int32) + + d = tl.arange(0, BLOCK_D) + last_dim: tl.constexpr = HEAD_DIM * 2 + + # Step 1: Write kv_score_input to state buffer at write_loc + d_mask_full = d < last_dim + input_val = tl.load( + input_ptr + bid * input_stride_b + d, mask=d_mask_full, other=0.0 + ) + tl.store(buf_ptr + write_loc * buf_stride_slot + d, input_val, mask=d_mask_full) + + # Step 2: Check boundary condition + d_mask_hd = d < HEAD_DIM + if seq_len % COMPRESS_RATIO != 0: + tl.store( + out_ptr + bid * out_stride_b + d, + tl.zeros([BLOCK_D], tl.float32), + mask=d_mask_hd, + ) + return + + # Step 3: Online softmax-pool over 128 slots in the page + page_base = read_page_0 * COMPRESS_RATIO * buf_stride_slot + m_prev = tl.full([BLOCK_D], float("-inf"), tl.float32) + kv_acc = tl.zeros([BLOCK_D], tl.float32) + w_acc = tl.zeros([BLOCK_D], tl.float32) + + for k in tl.static_range(COMPRESS_RATIO): + slot_addr = page_base + k * buf_stride_slot + kv_val = tl.load(buf_ptr + slot_addr + d, mask=d_mask_hd, other=0.0).to( + tl.float32 + ) + sc_val = tl.load( + buf_ptr + slot_addr + HEAD_DIM + d, mask=d_mask_hd, other=0.0 + ).to(tl.float32) + ape_val = tl.load( + ape_ptr + k * ape_stride_r + d, mask=d_mask_hd, other=0.0 + ).to(tl.float32) + score_k = sc_val + ape_val + + m_new = tl.maximum(m_prev, score_k) + exp_old = tl.where(m_prev == float("-inf"), 0.0, tl.exp(m_prev - m_new)) + exp_cur = tl.where(score_k == float("-inf"), 0.0, tl.exp(score_k - m_new)) + kv_acc = kv_acc * exp_old + exp_cur * kv_val + w_acc = w_acc * exp_old + exp_cur + m_prev = m_new + + compressed = kv_acc / w_acc + tl.store(out_ptr + bid * out_stride_b + d, compressed, mask=d_mask_hd) + + @triton.jit + def _c128_compress_prefill_write_kernel( + buf_ptr, + input_ptr, + plan_w_ptr, + buf_stride_slot, + input_stride_b, + num_w, + BLOCK_D: tl.constexpr, + LAST_DIM: tl.constexpr, + ): + """Prefill write phase: scatter kv_score_input tokens into state buffer.""" + wid = tl.program_id(0) + if wid >= num_w: + return + + # WritePlan: {ragged_id(u32), write_loc(i32)} = 8 bytes = 2 int32s + plan_base = plan_w_ptr + wid * 2 + ragged_id = (tl.load(plan_base).to(tl.int32)) & 0xFFFF + write_loc = tl.load(plan_base + 1).to(tl.int32) + + d = tl.arange(0, BLOCK_D) + d_mask = d < LAST_DIM + + if write_loc >= 0: + input_val = tl.load( + input_ptr + ragged_id * input_stride_b + d, mask=d_mask, other=0.0 + ) + tl.store(buf_ptr + write_loc * buf_stride_slot + d, input_val, mask=d_mask) + + @triton.jit + def _c128_compress_prefill_compress_kernel( + buf_ptr, + ape_ptr, + out_ptr, + plan_c_ptr, + buf_stride_slot, + ape_stride_r, + out_stride_b, + num_c, + HEAD_DIM: tl.constexpr, + BLOCK_D: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, + ): + """Prefill compress phase: online softmax-pool for each compress plan entry.""" + cid = tl.program_id(0) + if cid >= num_c: + return + + # CompressPlan: {seq_len(u32), ragged_id(u16)|buffer_len(u16), read_page_0(i32), read_page_1(i32)} + plan_base = plan_c_ptr + cid * 4 + read_page_0 = tl.load(plan_base + 2).to(tl.int32) + + d = tl.arange(0, BLOCK_D) + d_mask_hd = d < HEAD_DIM + + if read_page_0 < 0: + tl.store( + out_ptr + cid * out_stride_b + d, + tl.zeros([BLOCK_D], tl.float32), + mask=d_mask_hd, + ) + return + + page_base = read_page_0 * COMPRESS_RATIO * buf_stride_slot + m_prev = tl.full([BLOCK_D], float("-inf"), tl.float32) + kv_acc = tl.zeros([BLOCK_D], tl.float32) + w_acc = tl.zeros([BLOCK_D], tl.float32) + + for k in tl.static_range(COMPRESS_RATIO): + slot_addr = page_base + k * buf_stride_slot + kv_val = tl.load(buf_ptr + slot_addr + d, mask=d_mask_hd, other=0.0).to( + tl.float32 + ) + sc_val = tl.load( + buf_ptr + slot_addr + HEAD_DIM + d, mask=d_mask_hd, other=0.0 + ).to(tl.float32) + ape_val = tl.load( + ape_ptr + k * ape_stride_r + d, mask=d_mask_hd, other=0.0 + ).to(tl.float32) + score_k = sc_val + ape_val + + m_new = tl.maximum(m_prev, score_k) + exp_old = tl.where(m_prev == float("-inf"), 0.0, tl.exp(m_prev - m_new)) + exp_cur = tl.where(score_k == float("-inf"), 0.0, tl.exp(score_k - m_new)) + kv_acc = kv_acc * exp_old + exp_cur * kv_val + w_acc = w_acc * exp_old + exp_cur + m_prev = m_new + + compressed = kv_acc / w_acc + tl.store(out_ptr + cid * out_stride_b + d, compressed, mask=d_mask_hd) + + +def _compress_forward_c128_triton( + kv_score_buffer: torch.Tensor, + kv_score_input: torch.Tensor, + ape: torch.Tensor, + plan: Union[CompressorDecodePlan, CompressorPrefillPlan], + head_dim: int, +) -> torch.Tensor: + """Triton C128 compress_forward for HIP (wave64). + + Fuses write + online-softmax-pool into Triton kernels. + CUDA graph compatible. + """ + num_total_slots = kv_score_buffer.shape[0] * kv_score_buffer.shape[1] + num_pages = kv_score_buffer.shape[0] + last_dim = kv_score_buffer.shape[-1] + compress_ratio = 128 + + buf_flat = kv_score_buffer.view(-1, last_dim) + buf_stride_slot = last_dim # elements per slot + + BLOCK_D = triton.next_power_of_2(last_dim) + + if plan.is_decode: + # Decode path: single kernel does write + compress + plan_raw = plan[1].view(torch.int32) # [bs, 4] + bs = plan_raw.shape[0] + out = torch.empty( + bs, head_dim, dtype=torch.float32, device=kv_score_input.device + ) + + if bs > 0 and num_total_slots > 0: + grid = (bs,) + _c128_compress_decode_kernel[grid]( + buf_flat, + kv_score_input, + ape, + out, + plan_raw, + buf_stride_slot, + kv_score_input.stride(0), + ape.stride(0), + out.stride(0), + bs, + HEAD_DIM=head_dim, + BLOCK_D=triton.next_power_of_2(head_dim), + COMPRESS_RATIO=compress_ratio, + num_warps=8, + ) + return out + else: + # Prefill path: separate write kernel + compress kernel + plan_c_raw = plan[1].view(torch.int32) # [num_c, 4] + plan_w = plan[2] # [num_w, 8] uint8 + plan_w_raw = plan_w.view(torch.int32) # [num_w, 2] + num_c = plan_c_raw.shape[0] + num_w = plan_w_raw.shape[0] + + out = torch.empty( + num_c, head_dim, dtype=torch.float32, device=kv_score_input.device + ) + + # Phase 1: Write + if num_w > 0 and num_total_slots > 0: + grid_w = (num_w,) + _c128_compress_prefill_write_kernel[grid_w]( + buf_flat, + kv_score_input, + plan_w_raw, + buf_stride_slot, + kv_score_input.stride(0), + num_w, + BLOCK_D=BLOCK_D, + LAST_DIM=last_dim, + num_warps=4, + ) + + # Phase 2: Compress + if num_c > 0 and num_pages > 0: + grid_c = (num_c,) + _c128_compress_prefill_compress_kernel[grid_c]( + buf_flat, + ape, + out, + plan_c_raw, + buf_stride_slot, + ape.stride(0), + out.stride(0), + num_c, + HEAD_DIM=head_dim, + BLOCK_D=triton.next_power_of_2(head_dim), + COMPRESS_RATIO=compress_ratio, + num_warps=8, + ) + + return out diff --git a/python/sglang/srt/layers/attention/dsv4/dequant_k_cache.py b/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/dequant_k_cache.py rename to python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py diff --git a/python/sglang/srt/layers/attention/dsv4/fp4_indexer.py b/python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/fp4_indexer.py rename to python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py diff --git a/python/sglang/srt/layers/attention/dsv4/fused_compress_triton.py b/python/sglang/kernels/ops/attention/dsv4/fused_compress_triton.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/fused_compress_triton.py rename to python/sglang/kernels/ops/attention/dsv4/fused_compress_triton.py diff --git a/python/sglang/kernels/ops/attention/dsv4/fused_scale.py b/python/sglang/kernels/ops/attention/dsv4/fused_scale.py new file mode 100644 index 000000000..74ba6144f --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/fused_scale.py @@ -0,0 +1,51 @@ +"""Fused scale kernel for the DSV4 indexer. + +Migrated from ``sglang.srt.layers.attention.dsv4.indexer`` (RFC #29630, Phase 2.5). +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _fused_scale_kernel( + weight_ptr, + q_scale_ptr, + out_ptr, + numel, + out_scale, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + mask = offs < numel + + w = tl.load(weight_ptr + offs, mask=mask) + qs = tl.load(q_scale_ptr + offs, mask=mask) + + acc = w.to(tl.float32) * out_scale * qs.to(tl.float32) + tl.store(out_ptr + offs, acc.to(out_ptr.dtype.element_ty), mask=mask) + + +def fused_scale( + weight: torch.Tensor, + out_scale: float, + q_scale: torch.Tensor, +) -> torch.Tensor: + assert weight.is_contiguous() and q_scale.is_contiguous() + B, H = weight.shape + numel = B * H + out_dtype = torch.promote_types(weight.dtype, q_scale.dtype) + out = torch.empty((B, H, 1), device=weight.device, dtype=out_dtype) + BLOCK = 1024 + grid = (triton.cdiv(numel, BLOCK),) + _fused_scale_kernel[grid]( + weight, + q_scale, + out, + numel, + out_scale, + BLOCK=BLOCK, + ) + return out diff --git a/python/sglang/srt/layers/attention/dsv4/index_buf_accessor.py b/python/sglang/kernels/ops/attention/dsv4/index_buf_accessor.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/index_buf_accessor.py rename to python/sglang/kernels/ops/attention/dsv4/index_buf_accessor.py diff --git a/python/sglang/srt/layers/attention/dsv4/metadata_kernel.py b/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/metadata_kernel.py rename to python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py diff --git a/python/sglang/srt/layers/attention/dsv4/quant_k_cache.py b/python/sglang/kernels/ops/attention/dsv4/quant_k_cache.py similarity index 97% rename from python/sglang/srt/layers/attention/dsv4/quant_k_cache.py rename to python/sglang/kernels/ops/attention/dsv4/quant_k_cache.py index 8f50df84b..21d44f33d 100644 --- a/python/sglang/srt/layers/attention/dsv4/quant_k_cache.py +++ b/python/sglang/kernels/ops/attention/dsv4/quant_k_cache.py @@ -2,8 +2,8 @@ import torch import triton import triton.language as tl +from sglang.kernels.ops.attention.dsv4.index_buf_accessor import NopeFp8RopeBf16Pack from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz -from sglang.srt.layers.attention.dsv4.index_buf_accessor import NopeFp8RopeBf16Pack fp8_dtype = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn diff --git a/python/sglang/kernels/ops/attention/dsv4/rms_normalize_hip.py b/python/sglang/kernels/ops/attention/dsv4/rms_normalize_hip.py new file mode 100644 index 000000000..bbd1e5476 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/rms_normalize_hip.py @@ -0,0 +1,52 @@ +"""RMS-normalize kernel used by the HIP DSV4 compressor. + +Migrated from ``sglang.srt.layers.attention.dsv4.compress_hip`` (RFC #29630, Phase 2.5). +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _rms_normalize_kernel( + x_ptr, + weight_ptr, + eps, + stride_row, + dim, + BLOCK_SIZE: tl.constexpr, + HAS_WEIGHT: tl.constexpr, +): + pid = tl.program_id(0) + offs = tl.arange(0, BLOCK_SIZE) + mask = offs < dim + base = pid * stride_row + x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32) + mean_sq = tl.sum(x * x, axis=0) / dim + rms_inv = tl.rsqrt(mean_sq + eps) + out = x * rms_inv + if HAS_WEIGHT: + weight = tl.load(weight_ptr + offs, mask=mask, other=0.0) + out = out * weight + tl.store(x_ptr + base + offs, out, mask=mask) + + +def rms_normalize_triton( + x: torch.Tensor, eps: float, weight: torch.Tensor = None +) -> torch.Tensor: + dim = x.shape[-1] + x_flat = x.view(-1, dim) + num_rows = x_flat.shape[0] + BLOCK_SIZE = triton.next_power_of_2(dim) + grid = (num_rows,) + _rms_normalize_kernel[grid]( + x_flat, + weight, + eps, + x_flat.stride(0), + dim, + BLOCK_SIZE=BLOCK_SIZE, + HAS_WEIGHT=(weight is not None), + ) + return x diff --git a/python/sglang/kernels/ops/attention/dsv4/sparse_prefill_kernels.py b/python/sglang/kernels/ops/attention/dsv4/sparse_prefill_kernels.py new file mode 100644 index 000000000..2c6d5887a --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/sparse_prefill_kernels.py @@ -0,0 +1,110 @@ +"""SWA token-id build and topk+SWA index combine kernels for DSV4 sparse prefill. + +Migrated from ``sglang.srt.layers.attention.dsv4.sparse_prefill_utils`` (RFC #29630, Phase 2.5). +""" + +import triton +import triton.language as tl + + +@triton.jit +def _build_swa_token_ids_kernel( + out_ptr, + swa_first_pos_ptr, + swa_gather_lens_ptr, + swa_offsets_ptr, + req_pool_indices_ptr, + req_to_token_ptr, + req_to_token_stride, + full_to_swa_ptr, +): + batch_idx = tl.program_id(0) + worker_id = tl.program_id(1) + num_workers = tl.num_programs(1) + + first_pos = tl.load(swa_first_pos_ptr + batch_idx) + gather_len = tl.load(swa_gather_lens_ptr + batch_idx) + out_off = tl.load(swa_offsets_ptr + batch_idx).to(tl.int64) + req_pool_idx = tl.load(req_pool_indices_ptr + batch_idx).to(tl.int64) + + for i in range(worker_id, gather_len, num_workers): + pos = first_pos + i + full_id = tl.load( + req_to_token_ptr + req_pool_idx * req_to_token_stride + pos + ).to(tl.int64) + swa_id = tl.load(full_to_swa_ptr + full_id).to(tl.int32) + tl.store(out_ptr + out_off + i, swa_id) + + +@triton.jit(do_not_specialize=["top_k"]) +def _combine_topk_swa_indices_kernel( + combined_indices_ptr, + combined_indices_stride, + combined_lens_ptr, + topk_indices_ptr, + topk_indices_stride, + query_start_loc_ptr, + seq_lens_ptr, + gather_lens_ptr, + compressed_base_ptr, + swa_base_ptr, + top_k, + COMPRESS_RATIO: tl.constexpr, + WINDOW_SIZE: tl.constexpr, + PADDED_TOP_K: tl.constexpr, +): + batch_idx = tl.program_id(0) + worker_id = tl.program_id(1) + num_workers = tl.num_programs(1) + + # query_start_loc may be a global tensor; rebase to chunk-local offsets + # by subtracting the chunk's starting value. + base = tl.load(query_start_loc_ptr) + query_start = tl.load(query_start_loc_ptr + batch_idx) - base + query_end = tl.load(query_start_loc_ptr + batch_idx + 1) - base + query_len = query_end - query_start + seq_len = tl.load(seq_lens_ptr + batch_idx) + gather_len = tl.load(gather_lens_ptr + batch_idx) + compressed_base = tl.load(compressed_base_ptr + batch_idx) + swa_base = tl.load(swa_base_ptr + batch_idx) + start_pos = seq_len - query_len + # SWA portion of the gathered buffer starts from position + # (seq_len - gather_len), not 0. The +pos-gather_start formula maps a + # query's window back into the workspace's SWA region. + gather_start = seq_len - gather_len + + for token_idx in range(query_start + worker_id, query_end, num_workers): + token_idx_in_query = token_idx - query_start + pos = start_pos + token_idx_in_query + # Both the C4 indexer and the C128 metadata builder emit + # min((pos+1)//compress_ratio, topk_tokens) valid entries. Caller + # passes top_k=0 for SWA-only layers to zero this out. + topk_len = tl.minimum((pos + 1) // COMPRESS_RATIO, top_k) + swa_len = tl.minimum(pos + 1, WINDOW_SIZE) + + combined_row = token_idx.to(tl.int64) * combined_indices_stride + topk_row = token_idx.to(tl.int64) * topk_indices_stride + + offset = tl.arange(0, PADDED_TOP_K) + mask = offset < topk_len + topk_vals = tl.load( + topk_indices_ptr + topk_row + offset, + mask=mask, + ) + tl.store( + combined_indices_ptr + combined_row + offset, + topk_vals + compressed_base, + mask=mask, + ) + + offset = tl.arange(0, WINDOW_SIZE) + # Workspace SWA index: swa_base[r] + (gather_offset_in_buffer). + # For positions [pos - swa_len + 1, pos], the buffer offsets are + # [pos - swa_len + 1 - gather_start, pos - gather_start]. + tl.store( + combined_indices_ptr + combined_row + topk_len + offset, + swa_base + offset + pos - swa_len + 1 - gather_start, + mask=offset < swa_len, + ) + + tl.store(combined_lens_ptr + token_idx, topk_len + swa_len) diff --git a/python/sglang/srt/layers/attention/dsv4/tilelang_kernel.py b/python/sglang/kernels/ops/attention/dsv4/tilelang_kernel.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/tilelang_kernel.py rename to python/sglang/kernels/ops/attention/dsv4/tilelang_kernel.py diff --git a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/env_gate.py b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/env_gate.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/env_gate.py rename to python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/env_gate.py diff --git a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/paged_decode.py b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/paged_decode.py rename to python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode.py diff --git a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/paged_decode_indices.py b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode_indices.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/paged_decode_indices.py rename to python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode_indices.py diff --git a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/paged_prefill.py b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_prefill.py similarity index 100% rename from python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/paged_prefill.py rename to python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_prefill.py diff --git a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/runtime.py b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/runtime.py similarity index 98% rename from python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/runtime.py rename to python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/runtime.py index 437fc3e8b..039af217d 100644 --- a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/runtime.py +++ b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/runtime.py @@ -31,13 +31,13 @@ import torch.nn.functional as F import triton import triton.language as tl -from sglang.srt.layers.attention.dsv4.unified_kv_kernels.paged_decode import ( +from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.paged_decode import ( sparse_attn_v4_paged_decode, ) -from sglang.srt.layers.attention.dsv4.unified_kv_kernels.paged_decode_indices import ( +from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.paged_decode_indices import ( write_v4_paged_decode_indices, ) -from sglang.srt.layers.attention.dsv4.unified_kv_kernels.paged_prefill import ( +from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.paged_prefill import ( sparse_attn_v4_paged_prefill, ) diff --git a/python/sglang/srt/arg_groups/hisparse_hook.py b/python/sglang/srt/arg_groups/hisparse_hook.py index d6c437bdb..c59784f6c 100644 --- a/python/sglang/srt/arg_groups/hisparse_hook.py +++ b/python/sglang/srt/arg_groups/hisparse_hook.py @@ -107,7 +107,7 @@ def validate_hisparse(server_args: ServerArgs) -> None: # In unified-KV mode c4_kv_pool is None, so DeepSeekV4HiSparseTokenToKVPoolAllocator # cannot attach and pool init dies with a cryptic AssertionError. Fail fast # at startup with a clear message instead. Remove once unified-KV HiSparse lands. - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index a15e41b0b..7dc34d158 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -19,6 +19,15 @@ import torch import torch.nn.functional as F from sglang.jit_kernel.dsv4.online_c128_mtp import OnlineC128MTPController +from sglang.kernels.ops.attention.dsv4.dequant_k_cache import ( + dequantize_k_cache_paged, +) +from sglang.kernels.ops.attention.dsv4.metadata_kernel import ( + init_compression_metadata as _init_compression_metadata_triton, +) +from sglang.kernels.ops.attention.dsv4.quant_k_cache import ( + quant_to_nope_fp8_rope_bf16_pack_triton, +) from sglang.srt.environ import envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.dsv4.attn_metadata_kernels import ( @@ -31,9 +40,6 @@ from sglang.srt.layers.attention.dsv4.compressor_v2 import ( FusedCompressMetadata, create_paged_compressor_data, ) -from sglang.srt.layers.attention.dsv4.dequant_k_cache import ( - dequantize_k_cache_paged, -) from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin from sglang.srt.layers.attention.dsv4.metadata import ( _LARGE_INDEXER_QUERY_THRESHOLD, @@ -41,12 +47,6 @@ from sglang.srt.layers.attention.dsv4.metadata import ( copy_metadata, maybe_copy_inplace, ) -from sglang.srt.layers.attention.dsv4.metadata_kernel import ( - init_compression_metadata as _init_compression_metadata_triton, -) -from sglang.srt.layers.attention.dsv4.quant_k_cache import ( - quant_to_nope_fp8_rope_bf16_pack_triton, -) from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import ( SparsePrefillChunkCache, SparsePrefillWorkspace, diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 28199de80..16b307a7e 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -18,6 +18,12 @@ from typing import ( import torch import torch.nn.functional as F +from sglang.kernels.ops.attention.dsv4.metadata_kernel import ( + init_compression_metadata as _init_compression_metadata_triton, +) +from sglang.kernels.ops.attention.dsv4.quant_k_cache import ( + quant_to_nope_fp8_rope_bf16_pack_triton, +) from sglang.srt.environ import envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.dsv4.compressor_v2 import ( @@ -31,12 +37,6 @@ from sglang.srt.layers.attention.dsv4.metadata import ( copy_metadata, maybe_copy_inplace, ) -from sglang.srt.layers.attention.dsv4.metadata_kernel import ( - init_compression_metadata as _init_compression_metadata_triton, -) -from sglang.srt.layers.attention.dsv4.quant_k_cache import ( - quant_to_nope_fp8_rope_bf16_pack_triton, -) from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.runtime_context import get_parallel @@ -1058,13 +1058,13 @@ class DeepseekV4HipRadixBackend( self, core: DSV4AttnMetadata, req_pool_indices: torch.Tensor ) -> None: """build the ragged decode index streams once per forward""" - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) if not is_unified_kv_triton(): return - from sglang.srt.layers.attention.dsv4.unified_kv_kernels import runtime + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels import runtime pool = self.token_to_kv_pool N = core.positions_casual.shape[0] @@ -1104,7 +1104,7 @@ class DeepseekV4HipRadixBackend( seq_lens: torch.Tensor, extend_seq_lens: torch.Tensor, ) -> None: - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) @@ -1139,7 +1139,7 @@ class DeepseekV4HipRadixBackend( save_kv_cache: bool = True, ) -> torch.Tensor: """unified_kv paged-attention path over the bf16 unified_kv""" - from sglang.srt.layers.attention.dsv4.unified_kv_kernels import runtime + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels import runtime pool = self.token_to_kv_pool layer_id = layer.layer_id @@ -1408,7 +1408,7 @@ class DeepseekV4HipRadixBackend( token_to_kv_pool = self.token_to_kv_pool assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 2702c860d..e6a1b8331 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -1555,7 +1555,7 @@ class Indexer(MultiPlatformOp): "piecewise/breakable CUDA graph" ) if not _is_npu: - from sglang.srt.layers.attention.dsa.tilelang_kernel import fp8_index + from sglang.kernels.ops.attention.dsa.tilelang_kernel import fp8_index page_size = get_token_to_kv_pool().page_size assert page_size == 64, "only support page size 64" @@ -1734,9 +1734,9 @@ class Indexer(MultiPlatformOp): return_indices: bool = True, ) -> Optional[torch.Tensor]: if _is_hip: - from sglang.srt.layers.attention.dsa.tilelang_kernel import act_quant + from sglang.kernels.ops.attention.dsa.tilelang_kernel import act_quant elif not _is_npu: - from sglang.srt.layers.attention.dsa.triton_kernel import act_quant + from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant if TYPE_CHECKING: assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool) @@ -2419,7 +2419,7 @@ def pcg_dsa_indexer_prefill_split( # captured graph reads it at a fixed address; eager code instead allocates # and returns a fresh, naturally-sized tensor each call. assert _is_cuda, "Internal error: DSA graph dispatch is only supported on CUDA" - from sglang.srt.layers.attention.dsa.triton_kernel import act_quant + from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant forward_context = get_tc_piecewise_forward_context() forward_batch = forward_context.forward_batch diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index e15c822db..62f47a57b 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -3,7 +3,6 @@ from typing import TYPE_CHECKING, List, Tuple, Union import torch import triton -import triton.language as tl from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import ( @@ -223,26 +222,9 @@ def can_dsa_cp_split(seq_len: int, cp_size: int, use_dsa: bool, forward_batch): return False -@triton.jit -def dsa_cp_round_robin_split_q_seqs_kernel( - in_seqs_ptr, - out_seqs_ptr, - bs_idx_ptr, - tokens: tl.constexpr, - cp_size: tl.constexpr, - cp_rank: tl.constexpr, -): - extra_seq = 0 - bs_idx = 0 - for bs in range(tokens): - cur_len = tl.load(in_seqs_ptr + bs) - cur_len += extra_seq - cur_seq = cur_len // cp_size + (cur_len % cp_size > cp_rank) - if cur_seq > 0: - tl.store(bs_idx_ptr + bs_idx, bs) - tl.store(out_seqs_ptr + bs_idx, cur_seq) - bs_idx += 1 - extra_seq = cur_len - cur_seq * cp_size +from sglang.kernels.ops.attention.dsa.cp_split import ( + dsa_cp_round_robin_split_q_seqs_kernel, +) def dsa_cp_round_robin_split_q_seqs_cpu(extend_seqs): diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index dae22035d..35a5487fe 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -18,6 +18,12 @@ from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa from sglang.srt.runtime_context import get_parallel logger = logging.getLogger(__name__) +from sglang.kernels.ops.attention.dsa.dequant_k_cache import dequantize_k_cache_paged +from sglang.kernels.ops.attention.dsa.quant_k_cache import quantize_k_cache +from sglang.kernels.ops.attention.dsa.transform_index import ( + transform_index_page_table_decode, + transform_index_page_table_prefill, +) from sglang.kernels.ops.attention.utils import ( concat_mla_absorb_q_general, mla_quantize_and_rope_for_fp8, @@ -25,7 +31,6 @@ from sglang.kernels.ops.attention.utils import ( ) from sglang.srt.environ import envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend -from sglang.srt.layers.attention.dsa.dequant_k_cache import dequantize_k_cache_paged from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import ( DeepseekSparseAttnBackendMTPPrecomputeMixin, PrecomputedMetadata, @@ -36,11 +41,6 @@ from sglang.srt.layers.attention.dsa.dsa_topk_backend import ( DSATopKBackend, TopkTransformMethod, ) -from sglang.srt.layers.attention.dsa.quant_k_cache import quantize_k_cache -from sglang.srt.layers.attention.dsa.transform_index import ( - transform_index_page_table_decode, - transform_index_page_table_prefill, -) from sglang.srt.layers.attention.dsa.utils import ( can_dsa_prefill_cp_round_robin_split, compute_dsa_seqlens, @@ -104,8 +104,8 @@ def _all_gather_dsa_trtllm_fp8_kv( _is_hip = is_hip() if _is_hip: + from sglang.kernels.ops.attention.dsa.triton_kernel import get_valid_kv_indices from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype - from sglang.srt.layers.attention.dsa.triton_kernel import get_valid_kv_indices try: from aiter import ( # noqa: F401 @@ -1961,7 +1961,7 @@ class DeepseekSparseAttnBackend( and page_table_1.shape[-1] == 2048 and q_nope.shape[0] >= 512 ): - from sglang.srt.layers.attention.dsa.triton_sparse_mla import ( + from sglang.kernels.ops.attention.dsa.triton_sparse_mla import ( triton_sparse_mla_fwd, ) @@ -2426,7 +2426,7 @@ class DeepseekSparseAttnBackend( page_table_1: torch.Tensor, sm_scale: float, ) -> torch.Tensor: - from sglang.srt.layers.attention.dsa.tilelang_kernel import tilelang_sparse_fwd + from sglang.kernels.ops.attention.dsa.tilelang_kernel import tilelang_sparse_fwd return tilelang_sparse_fwd( q=q_all, diff --git a/python/sglang/srt/layers/attention/dsv4/compress_hip.py b/python/sglang/srt/layers/attention/dsv4/compress_hip.py index 370eb8606..9151232c5 100644 --- a/python/sglang/srt/layers/attention/dsv4/compress_hip.py +++ b/python/sglang/srt/layers/attention/dsv4/compress_hip.py @@ -6,20 +6,18 @@ from typing import TYPE_CHECKING, Any import torch 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.kernels.ops.attention.dsv4.fused_compress_triton import ( + fused_ape_pool_norm_rope, +) 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 -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 try: @@ -39,49 +37,7 @@ if TYPE_CHECKING: ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch - -@triton.jit -def _rms_normalize_kernel( - x_ptr, - weight_ptr, - eps, - stride_row, - dim, - BLOCK_SIZE: tl.constexpr, - HAS_WEIGHT: tl.constexpr, -): - pid = tl.program_id(0) - offs = tl.arange(0, BLOCK_SIZE) - mask = offs < dim - base = pid * stride_row - x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32) - mean_sq = tl.sum(x * x, axis=0) / dim - rms_inv = tl.rsqrt(mean_sq + eps) - out = x * rms_inv - if HAS_WEIGHT: - weight = tl.load(weight_ptr + offs, mask=mask, other=0.0) - out = out * weight - tl.store(x_ptr + base + offs, out, mask=mask) - - -def rms_normalize_triton( - x: torch.Tensor, eps: float, weight: torch.Tensor = None -) -> torch.Tensor: - dim = x.shape[-1] - x_flat = x.view(-1, dim) - num_rows = x_flat.shape[0] - BLOCK_SIZE = triton.next_power_of_2(dim) - grid = (num_rows,) - _rms_normalize_kernel[grid]( - x_flat, - weight, - eps, - x_flat.stride(0), - dim, - BLOCK_SIZE=BLOCK_SIZE, - HAS_WEIGHT=(weight is not None), - ) - return x +from sglang.kernels.ops.attention.dsv4.rms_normalize_hip import rms_normalize_triton class DeepseekRefRMSNorm(nn.Module): diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index 0b3694b0e..d5ad03c51 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -12,13 +12,13 @@ from sglang.jit_kernel.dsv4.compress_old import ( compress_forward, compress_fused_norm_rope_inplace, ) -from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config -from sglang.srt.environ import envs -from sglang.srt.layers.attention.dsa.triton_kernel import act_quant -from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp -from sglang.srt.layers.attention.dsv4.quant_k_cache import ( +from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant +from sglang.kernels.ops.attention.dsv4.quant_k_cache import ( quant_to_nope_fp8_rope_bf16_pack_triton, ) +from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config +from sglang.srt.environ import envs +from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output @@ -99,7 +99,7 @@ class CompressorBackendMixin: if not is_paged: raise NotImplementedError("HIP fused compressor expects paged metadata") - from sglang.srt.layers.attention.dsv4.fused_compress_triton import ( + from sglang.kernels.ops.attention.dsv4.fused_compress_triton import ( hip_compress_forward, hip_compress_fused_norm_rope_hadamard_inplace, hip_compress_fused_norm_rope_inplace, diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index a783eac8a..12f5a8bf2 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -27,279 +27,6 @@ FusedCompressMetadata: TypeAlias = CompressMetadata _is_hip = is_hip_runtime() -if _is_hip: - import triton - import triton.language as tl - - @triton.jit - def _c128_compress_decode_kernel( - buf_ptr, - input_ptr, - ape_ptr, - out_ptr, - plan_ptr, - buf_stride_slot, - input_stride_b, - ape_stride_r, - out_stride_b, - bs, - HEAD_DIM: tl.constexpr, - BLOCK_D: tl.constexpr, - COMPRESS_RATIO: tl.constexpr, - ): - """Fused C128 decode: write to state buffer + online softmax-pool. - - plan_ptr points to int32 view: [bs, 4] where each row is - {seq_len, write_loc, read_page_0, read_page_1}. - """ - bid = tl.program_id(0) - if bid >= bs: - return - - # Parse plan - plan_base = plan_ptr + bid * 4 - seq_len = tl.load(plan_base).to(tl.int32) - write_loc = tl.load(plan_base + 1).to(tl.int32) - read_page_0 = tl.load(plan_base + 2).to(tl.int32) - - d = tl.arange(0, BLOCK_D) - last_dim: tl.constexpr = HEAD_DIM * 2 - - # Step 1: Write kv_score_input to state buffer at write_loc - d_mask_full = d < last_dim - input_val = tl.load( - input_ptr + bid * input_stride_b + d, mask=d_mask_full, other=0.0 - ) - tl.store(buf_ptr + write_loc * buf_stride_slot + d, input_val, mask=d_mask_full) - - # Step 2: Check boundary condition - d_mask_hd = d < HEAD_DIM - if seq_len % COMPRESS_RATIO != 0: - tl.store( - out_ptr + bid * out_stride_b + d, - tl.zeros([BLOCK_D], tl.float32), - mask=d_mask_hd, - ) - return - - # Step 3: Online softmax-pool over 128 slots in the page - page_base = read_page_0 * COMPRESS_RATIO * buf_stride_slot - m_prev = tl.full([BLOCK_D], float("-inf"), tl.float32) - kv_acc = tl.zeros([BLOCK_D], tl.float32) - w_acc = tl.zeros([BLOCK_D], tl.float32) - - for k in tl.static_range(COMPRESS_RATIO): - slot_addr = page_base + k * buf_stride_slot - kv_val = tl.load(buf_ptr + slot_addr + d, mask=d_mask_hd, other=0.0).to( - tl.float32 - ) - sc_val = tl.load( - buf_ptr + slot_addr + HEAD_DIM + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - ape_val = tl.load( - ape_ptr + k * ape_stride_r + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - score_k = sc_val + ape_val - - m_new = tl.maximum(m_prev, score_k) - exp_old = tl.where(m_prev == float("-inf"), 0.0, tl.exp(m_prev - m_new)) - exp_cur = tl.where(score_k == float("-inf"), 0.0, tl.exp(score_k - m_new)) - kv_acc = kv_acc * exp_old + exp_cur * kv_val - w_acc = w_acc * exp_old + exp_cur - m_prev = m_new - - compressed = kv_acc / w_acc - tl.store(out_ptr + bid * out_stride_b + d, compressed, mask=d_mask_hd) - - @triton.jit - def _c128_compress_prefill_write_kernel( - buf_ptr, - input_ptr, - plan_w_ptr, - buf_stride_slot, - input_stride_b, - num_w, - BLOCK_D: tl.constexpr, - LAST_DIM: tl.constexpr, - ): - """Prefill write phase: scatter kv_score_input tokens into state buffer.""" - wid = tl.program_id(0) - if wid >= num_w: - return - - # WritePlan: {ragged_id(u32), write_loc(i32)} = 8 bytes = 2 int32s - plan_base = plan_w_ptr + wid * 2 - ragged_id = (tl.load(plan_base).to(tl.int32)) & 0xFFFF - write_loc = tl.load(plan_base + 1).to(tl.int32) - - d = tl.arange(0, BLOCK_D) - d_mask = d < LAST_DIM - - if write_loc >= 0: - input_val = tl.load( - input_ptr + ragged_id * input_stride_b + d, mask=d_mask, other=0.0 - ) - tl.store(buf_ptr + write_loc * buf_stride_slot + d, input_val, mask=d_mask) - - @triton.jit - def _c128_compress_prefill_compress_kernel( - buf_ptr, - ape_ptr, - out_ptr, - plan_c_ptr, - buf_stride_slot, - ape_stride_r, - out_stride_b, - num_c, - HEAD_DIM: tl.constexpr, - BLOCK_D: tl.constexpr, - COMPRESS_RATIO: tl.constexpr, - ): - """Prefill compress phase: online softmax-pool for each compress plan entry.""" - cid = tl.program_id(0) - if cid >= num_c: - return - - # CompressPlan: {seq_len(u32), ragged_id(u16)|buffer_len(u16), read_page_0(i32), read_page_1(i32)} - plan_base = plan_c_ptr + cid * 4 - read_page_0 = tl.load(plan_base + 2).to(tl.int32) - - d = tl.arange(0, BLOCK_D) - d_mask_hd = d < HEAD_DIM - - if read_page_0 < 0: - tl.store( - out_ptr + cid * out_stride_b + d, - tl.zeros([BLOCK_D], tl.float32), - mask=d_mask_hd, - ) - return - - page_base = read_page_0 * COMPRESS_RATIO * buf_stride_slot - m_prev = tl.full([BLOCK_D], float("-inf"), tl.float32) - kv_acc = tl.zeros([BLOCK_D], tl.float32) - w_acc = tl.zeros([BLOCK_D], tl.float32) - - for k in tl.static_range(COMPRESS_RATIO): - slot_addr = page_base + k * buf_stride_slot - kv_val = tl.load(buf_ptr + slot_addr + d, mask=d_mask_hd, other=0.0).to( - tl.float32 - ) - sc_val = tl.load( - buf_ptr + slot_addr + HEAD_DIM + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - ape_val = tl.load( - ape_ptr + k * ape_stride_r + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - score_k = sc_val + ape_val - - m_new = tl.maximum(m_prev, score_k) - exp_old = tl.where(m_prev == float("-inf"), 0.0, tl.exp(m_prev - m_new)) - exp_cur = tl.where(score_k == float("-inf"), 0.0, tl.exp(score_k - m_new)) - kv_acc = kv_acc * exp_old + exp_cur * kv_val - w_acc = w_acc * exp_old + exp_cur - m_prev = m_new - - compressed = kv_acc / w_acc - tl.store(out_ptr + cid * out_stride_b + d, compressed, mask=d_mask_hd) - - -def _compress_forward_c128_triton( - kv_score_buffer: torch.Tensor, - kv_score_input: torch.Tensor, - ape: torch.Tensor, - plan: Union[CompressorDecodePlan, CompressorPrefillPlan], - head_dim: int, -) -> torch.Tensor: - """Triton C128 compress_forward for HIP (wave64). - - Fuses write + online-softmax-pool into Triton kernels. - CUDA graph compatible. - """ - num_total_slots = kv_score_buffer.shape[0] * kv_score_buffer.shape[1] - num_pages = kv_score_buffer.shape[0] - last_dim = kv_score_buffer.shape[-1] - compress_ratio = 128 - - buf_flat = kv_score_buffer.view(-1, last_dim) - buf_stride_slot = last_dim # elements per slot - - BLOCK_D = triton.next_power_of_2(last_dim) - - if plan.is_decode: - # Decode path: single kernel does write + compress - plan_raw = plan[1].view(torch.int32) # [bs, 4] - bs = plan_raw.shape[0] - out = torch.empty( - bs, head_dim, dtype=torch.float32, device=kv_score_input.device - ) - - if bs > 0 and num_total_slots > 0: - grid = (bs,) - _c128_compress_decode_kernel[grid]( - buf_flat, - kv_score_input, - ape, - out, - plan_raw, - buf_stride_slot, - kv_score_input.stride(0), - ape.stride(0), - out.stride(0), - bs, - HEAD_DIM=head_dim, - BLOCK_D=triton.next_power_of_2(head_dim), - COMPRESS_RATIO=compress_ratio, - num_warps=8, - ) - return out - else: - # Prefill path: separate write kernel + compress kernel - plan_c_raw = plan[1].view(torch.int32) # [num_c, 4] - plan_w = plan[2] # [num_w, 8] uint8 - plan_w_raw = plan_w.view(torch.int32) # [num_w, 2] - num_c = plan_c_raw.shape[0] - num_w = plan_w_raw.shape[0] - - out = torch.empty( - num_c, head_dim, dtype=torch.float32, device=kv_score_input.device - ) - - # Phase 1: Write - if num_w > 0 and num_total_slots > 0: - grid_w = (num_w,) - _c128_compress_prefill_write_kernel[grid_w]( - buf_flat, - kv_score_input, - plan_w_raw, - buf_stride_slot, - kv_score_input.stride(0), - num_w, - BLOCK_D=BLOCK_D, - LAST_DIM=last_dim, - num_warps=4, - ) - - # Phase 2: Compress - if num_c > 0 and num_pages > 0: - grid_c = (num_c,) - _c128_compress_prefill_compress_kernel[grid_c]( - buf_flat, - ape, - out, - plan_c_raw, - buf_stride_slot, - ape.stride(0), - out.stride(0), - num_c, - HEAD_DIM=head_dim, - BLOCK_D=triton.next_power_of_2(head_dim), - COMPRESS_RATIO=compress_ratio, - num_warps=8, - ) - - return out - def _use_online_compress(compress_ratio: int) -> bool: """Online state-pool path is c128-only.""" @@ -487,7 +214,7 @@ class CompressorBackendMixin: kv_score_input = compressor.compute_kv_score(x, forward_batch) state_pool = compressor.get_state_pool(self) - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) @@ -565,7 +292,7 @@ class CompressorBackendMixin: from sglang.kernels.ops.attention.deepseek_v4_rope import ( fused_norm_rope_inplace_triton, ) - from sglang.srt.layers.attention.dsv4.quant_k_cache import ( + from sglang.kernels.ops.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 diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 0f92d65d4..e11fbdeec 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -5,8 +5,6 @@ from typing import TYPE_CHECKING, Any, List, Optional, Tuple, TypeAlias, Union import torch import torch.nn as nn import torch.nn.functional as F -import triton -import triton.language as tl from sglang.jit_kernel.dsv4 import ( fused_q_indexer_rope_hadamard_fp4_quant, @@ -313,49 +311,6 @@ def topk_transform_512_pytorch_vectorized( out_raw_indices.copy_(raw_indices) -@triton.jit -def _fused_scale_kernel( - weight_ptr, - q_scale_ptr, - out_ptr, - numel, - out_scale, - BLOCK: tl.constexpr, -): - pid = tl.program_id(0) - offs = pid * BLOCK + tl.arange(0, BLOCK) - mask = offs < numel - - w = tl.load(weight_ptr + offs, mask=mask) - qs = tl.load(q_scale_ptr + offs, mask=mask) - - acc = w.to(tl.float32) * out_scale * qs.to(tl.float32) - tl.store(out_ptr + offs, acc.to(out_ptr.dtype.element_ty), mask=mask) - - -def fused_scale( - weight: torch.Tensor, - out_scale: float, - q_scale: torch.Tensor, -) -> torch.Tensor: - assert weight.is_contiguous() and q_scale.is_contiguous() - B, H = weight.shape - numel = B * H - out_dtype = torch.promote_types(weight.dtype, q_scale.dtype) - out = torch.empty((B, H, 1), device=weight.device, dtype=out_dtype) - BLOCK = 1024 - grid = (triton.cdiv(numel, BLOCK),) - _fused_scale_kernel[grid]( - weight, - q_scale, - out, - numel, - out_scale, - BLOCK=BLOCK, - ) - return out - - class C4IndexerBackendMixin: def __init__(self): super().__init__() @@ -658,7 +613,7 @@ class C4IndexerBackendMixin: raise RuntimeError("DeepSeek V4 FP4 indexer requires DeepGEMM indexer.") from deep_gemm import fp8_fp4_paged_mqa_logits as fn elif envs.SGLANG_OPT_USE_TILELANG_INDEXER.get(): - from sglang.srt.layers.attention.dsa.tilelang_kernel import ( + from sglang.kernels.ops.attention.dsa.tilelang_kernel import ( tilelang_fp8_paged_mqa_logits as fn, ) elif envs.SGLANG_OPT_USE_AITER_INDEXER.get(): diff --git a/python/sglang/srt/layers/attention/dsv4/sparse_prefill_utils.py b/python/sglang/srt/layers/attention/dsv4/sparse_prefill_utils.py index 351263cb0..931b791ee 100644 --- a/python/sglang/srt/layers/attention/dsv4/sparse_prefill_utils.py +++ b/python/sglang/srt/layers/attention/dsv4/sparse_prefill_utils.py @@ -37,9 +37,8 @@ from typing import Optional import torch import triton -import triton.language as tl -from sglang.srt.layers.attention.dsv4.dequant_k_cache import DIM_NOPE, DIM_ROPE +from sglang.kernels.ops.attention.dsv4.dequant_k_cache import DIM_NOPE, DIM_ROPE from sglang.srt.utils import ceil_align # FlashMLA sparse prefill asserts ``params.topk % B_TOPK == 0``. B_TOPK is 64 @@ -50,6 +49,12 @@ SPARSE_PREFILL_TOPK_ALIGNMENT = 128 WORKSPACE_DIM = DIM_NOPE + DIM_ROPE +from sglang.kernels.ops.attention.dsv4.sparse_prefill_kernels import ( + _build_swa_token_ids_kernel, + _combine_topk_swa_indices_kernel, +) + + class SparsePrefillWorkspace: """Backend-owned scratch storage for sparse prefill KV dequantization. @@ -253,109 +258,6 @@ def build_swa_token_ids( return swa_token_ids, swa_first_pos, swa_gather_lens, swa_offsets -@triton.jit -def _build_swa_token_ids_kernel( - out_ptr, - swa_first_pos_ptr, - swa_gather_lens_ptr, - swa_offsets_ptr, - req_pool_indices_ptr, - req_to_token_ptr, - req_to_token_stride, - full_to_swa_ptr, -): - batch_idx = tl.program_id(0) - worker_id = tl.program_id(1) - num_workers = tl.num_programs(1) - - first_pos = tl.load(swa_first_pos_ptr + batch_idx) - gather_len = tl.load(swa_gather_lens_ptr + batch_idx) - out_off = tl.load(swa_offsets_ptr + batch_idx).to(tl.int64) - req_pool_idx = tl.load(req_pool_indices_ptr + batch_idx).to(tl.int64) - - for i in range(worker_id, gather_len, num_workers): - pos = first_pos + i - full_id = tl.load( - req_to_token_ptr + req_pool_idx * req_to_token_stride + pos - ).to(tl.int64) - swa_id = tl.load(full_to_swa_ptr + full_id).to(tl.int32) - tl.store(out_ptr + out_off + i, swa_id) - - -@triton.jit(do_not_specialize=["top_k"]) -def _combine_topk_swa_indices_kernel( - combined_indices_ptr, - combined_indices_stride, - combined_lens_ptr, - topk_indices_ptr, - topk_indices_stride, - query_start_loc_ptr, - seq_lens_ptr, - gather_lens_ptr, - compressed_base_ptr, - swa_base_ptr, - top_k, - COMPRESS_RATIO: tl.constexpr, - WINDOW_SIZE: tl.constexpr, - PADDED_TOP_K: tl.constexpr, -): - batch_idx = tl.program_id(0) - worker_id = tl.program_id(1) - num_workers = tl.num_programs(1) - - # query_start_loc may be a global tensor; rebase to chunk-local offsets - # by subtracting the chunk's starting value. - base = tl.load(query_start_loc_ptr) - query_start = tl.load(query_start_loc_ptr + batch_idx) - base - query_end = tl.load(query_start_loc_ptr + batch_idx + 1) - base - query_len = query_end - query_start - seq_len = tl.load(seq_lens_ptr + batch_idx) - gather_len = tl.load(gather_lens_ptr + batch_idx) - compressed_base = tl.load(compressed_base_ptr + batch_idx) - swa_base = tl.load(swa_base_ptr + batch_idx) - start_pos = seq_len - query_len - # SWA portion of the gathered buffer starts from position - # (seq_len - gather_len), not 0. The +pos-gather_start formula maps a - # query's window back into the workspace's SWA region. - gather_start = seq_len - gather_len - - for token_idx in range(query_start + worker_id, query_end, num_workers): - token_idx_in_query = token_idx - query_start - pos = start_pos + token_idx_in_query - # Both the C4 indexer and the C128 metadata builder emit - # min((pos+1)//compress_ratio, topk_tokens) valid entries. Caller - # passes top_k=0 for SWA-only layers to zero this out. - topk_len = tl.minimum((pos + 1) // COMPRESS_RATIO, top_k) - swa_len = tl.minimum(pos + 1, WINDOW_SIZE) - - combined_row = token_idx.to(tl.int64) * combined_indices_stride - topk_row = token_idx.to(tl.int64) * topk_indices_stride - - offset = tl.arange(0, PADDED_TOP_K) - mask = offset < topk_len - topk_vals = tl.load( - topk_indices_ptr + topk_row + offset, - mask=mask, - ) - tl.store( - combined_indices_ptr + combined_row + offset, - topk_vals + compressed_base, - mask=mask, - ) - - offset = tl.arange(0, WINDOW_SIZE) - # Workspace SWA index: swa_base[r] + (gather_offset_in_buffer). - # For positions [pos - swa_len + 1, pos], the buffer offsets are - # [pos - swa_len + 1 - gather_start, pos - gather_start]. - tl.store( - combined_indices_ptr + combined_row + topk_len + offset, - swa_base + offset + pos - swa_len + 1 - gather_start, - mask=offset < swa_len, - ) - - tl.store(combined_lens_ptr + token_idx, topk_len + swa_len) - - @dataclass class SparsePrefillChunkCache: """Chunk-invariant scaffolding for ``_forward_prefill_sparse``. diff --git a/python/sglang/srt/layers/attention/hip_flash_mla.py b/python/sglang/srt/layers/attention/hip_flash_mla.py index 7643d190e..4576873a7 100644 --- a/python/sglang/srt/layers/attention/hip_flash_mla.py +++ b/python/sglang/srt/layers/attention/hip_flash_mla.py @@ -31,7 +31,7 @@ def flash_mla_with_kvcache_entrypoint(backend: str, **kwargs): return flash_mla_with_kvcache_torch(**kwargs) if backend == "tilelang": - from sglang.srt.layers.attention.dsa.tilelang_kernel import ( + from sglang.kernels.ops.attention.dsa.tilelang_kernel import ( dpsk_v4_fp8_attention_fwd, ) diff --git a/python/sglang/srt/layers/attention/nsa/dequant_k_cache.py b/python/sglang/srt/layers/attention/nsa/dequant_k_cache.py index 43c080a6c..c86091f15 100644 --- a/python/sglang/srt/layers/attention/nsa/dequant_k_cache.py +++ b/python/sglang/srt/layers/attention/nsa/dequant_k_cache.py @@ -3,8 +3,8 @@ import warnings warnings.warn( "sglang.srt.layers.attention.nsa.dequant_k_cache is deprecated; " - "use sglang.srt.layers.attention.dsa.dequant_k_cache instead.", + "use sglang.kernels.ops.attention.dsa.dequant_k_cache instead.", DeprecationWarning, stacklevel=2, ) -from sglang.srt.layers.attention.dsa.dequant_k_cache import * # noqa: F401, F403 +from sglang.kernels.ops.attention.dsa.dequant_k_cache import * # noqa: F401, F403 diff --git a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py index 7ee993b4f..72b92ec44 100644 --- a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py +++ b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py @@ -3,8 +3,8 @@ import warnings warnings.warn( "sglang.srt.layers.attention.nsa.index_buf_accessor is deprecated; " - "use sglang.srt.layers.attention.dsa.index_buf_accessor instead.", + "use sglang.kernels.ops.attention.dsa.index_buf_accessor instead.", DeprecationWarning, stacklevel=2, ) -from sglang.srt.layers.attention.dsa.index_buf_accessor import * # noqa: F401, F403 +from sglang.kernels.ops.attention.dsa.index_buf_accessor import * # noqa: F401, F403 diff --git a/python/sglang/srt/layers/attention/nsa/quant_k_cache.py b/python/sglang/srt/layers/attention/nsa/quant_k_cache.py index 65bca9e39..1b3a2274c 100644 --- a/python/sglang/srt/layers/attention/nsa/quant_k_cache.py +++ b/python/sglang/srt/layers/attention/nsa/quant_k_cache.py @@ -3,8 +3,8 @@ import warnings warnings.warn( "sglang.srt.layers.attention.nsa.quant_k_cache is deprecated; " - "use sglang.srt.layers.attention.dsa.quant_k_cache instead.", + "use sglang.kernels.ops.attention.dsa.quant_k_cache instead.", DeprecationWarning, stacklevel=2, ) -from sglang.srt.layers.attention.dsa.quant_k_cache import * # noqa: F401, F403 +from sglang.kernels.ops.attention.dsa.quant_k_cache import * # noqa: F401, F403 diff --git a/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py b/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py index 7e653f7ec..3fb456f0c 100644 --- a/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py +++ b/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py @@ -3,8 +3,8 @@ import warnings warnings.warn( "sglang.srt.layers.attention.nsa.tilelang_kernel is deprecated; " - "use sglang.srt.layers.attention.dsa.tilelang_kernel instead.", + "use sglang.kernels.ops.attention.dsa.tilelang_kernel instead.", DeprecationWarning, stacklevel=2, ) -from sglang.srt.layers.attention.dsa.tilelang_kernel import * # noqa: F401, F403 +from sglang.kernels.ops.attention.dsa.tilelang_kernel import * # noqa: F401, F403 diff --git a/python/sglang/srt/layers/attention/nsa/transform_index.py b/python/sglang/srt/layers/attention/nsa/transform_index.py index 67ccaa772..90dfc4924 100644 --- a/python/sglang/srt/layers/attention/nsa/transform_index.py +++ b/python/sglang/srt/layers/attention/nsa/transform_index.py @@ -3,8 +3,8 @@ import warnings warnings.warn( "sglang.srt.layers.attention.nsa.transform_index is deprecated; " - "use sglang.srt.layers.attention.dsa.transform_index instead.", + "use sglang.kernels.ops.attention.dsa.transform_index instead.", DeprecationWarning, stacklevel=2, ) -from sglang.srt.layers.attention.dsa.transform_index import * # noqa: F401, F403 +from sglang.kernels.ops.attention.dsa.transform_index import * # noqa: F401, F403 diff --git a/python/sglang/srt/layers/attention/nsa/triton_kernel.py b/python/sglang/srt/layers/attention/nsa/triton_kernel.py index efa6c354f..aa54fa950 100644 --- a/python/sglang/srt/layers/attention/nsa/triton_kernel.py +++ b/python/sglang/srt/layers/attention/nsa/triton_kernel.py @@ -3,8 +3,8 @@ import warnings warnings.warn( "sglang.srt.layers.attention.nsa.triton_kernel is deprecated; " - "use sglang.srt.layers.attention.dsa.triton_kernel instead.", + "use sglang.kernels.ops.attention.dsa.triton_kernel instead.", DeprecationWarning, stacklevel=2, ) -from sglang.srt.layers.attention.dsa.triton_kernel import * # noqa: F401, F403 +from sglang.kernels.ops.attention.dsa.triton_kernel import * # noqa: F401, F403 diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index 71ad833b1..c36ee39ea 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -11,13 +11,13 @@ from sglang.jit_kernel.dsv4 import ( fused_k_norm_rope_flashmla, fused_store_cache, ) -from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE -from sglang.srt.environ import envs -from sglang.srt.layers.attention.dsa import index_buf_accessor -from sglang.srt.layers.attention.dsv4 import ( +from sglang.kernels.ops.attention.dsa import index_buf_accessor +from sglang.kernels.ops.attention.dsv4 import ( index_buf_accessor as dsv4_index_buf_accessor, ) -from sglang.srt.layers.attention.dsv4.index_buf_accessor import NopeFp8RopeBf16Pack +from sglang.kernels.ops.attention.dsv4.index_buf_accessor import NopeFp8RopeBf16Pack +from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE +from sglang.srt.environ import envs from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.memory_pool import KVCache @@ -368,7 +368,7 @@ class DeepSeekV4IndexerPool(KVCache): loc: torch.Tensor, cache_k: torch.Tensor, ) -> None: - from sglang.srt.layers.attention.dsv4.fp4_indexer import ( + from sglang.kernels.ops.attention.dsv4.fp4_indexer import ( store_fp4_index_k_cache, ) @@ -565,7 +565,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): c4_page_size = page_size // 4 c128_page_size = page_size // 128 - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) diff --git a/python/sglang/srt/mem_cache/dsa_cache_layer_split.py b/python/sglang/srt/mem_cache/dsa_cache_layer_split.py index 9df0b2061..75f12030a 100644 --- a/python/sglang/srt/mem_cache/dsa_cache_layer_split.py +++ b/python/sglang/srt/mem_cache/dsa_cache_layer_split.py @@ -35,7 +35,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.srt.layers.attention.dsa import index_buf_accessor +from sglang.kernels.ops.attention.dsa import index_buf_accessor from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range from sglang.srt.mem_cache.memory_pool import ( GPU_MEMORY_TYPE_KV_CACHE, diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 4b50d2f4a..05d7c1241 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -36,6 +36,11 @@ import triton import triton.language as tl from sglang.jit_kernel.kvcache import can_use_store_cache, store_cache +from sglang.kernels.ops.attention.dsa import index_buf_accessor +from sglang.kernels.ops.attention.dsa.quant_k_cache import ( + quantize_k_cache, + quantize_k_cache_separate, +) from sglang.kernels.ops.kvcache.cache_move import ( copy_all_layer_kv_cache_func, set_kv_buffer_prefix_valid_tiled, @@ -45,11 +50,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz from sglang.srt.configs.mamba_utils import BaseLinearStateParams from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE from sglang.srt.environ import envs -from sglang.srt.layers.attention.dsa import index_buf_accessor -from sglang.srt.layers.attention.dsa.quant_k_cache import ( - quantize_k_cache, - quantize_k_cache_separate, -) from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator diff --git a/python/sglang/srt/mem_cache/swa_radix_cache.py b/python/sglang/srt/mem_cache/swa_radix_cache.py index 013a87001..516516734 100644 --- a/python/sglang/srt/mem_cache/swa_radix_cache.py +++ b/python/sglang/srt/mem_cache/swa_radix_cache.py @@ -383,7 +383,7 @@ class SWARadixCache(KVCacheEventMixin, BasePrefixCache): No-op (0) for the index-addressed SWA pool, whose slots are content-stable and safe to reuse. """ - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index 037e868fc..214244a1a 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -4,9 +4,9 @@ from typing import TYPE_CHECKING import torch +from sglang.kernels.ops.attention.dsa.dequant_k_cache import dequantize_k_cache_paged from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton from sglang.srt.environ import envs -from sglang.srt.layers.attention.dsa.dequant_k_cache import dequantize_k_cache_paged from sglang.srt.layers.attention.tbo_backend import TboAttnBackend from sglang.srt.layers.communicator import get_attn_tp_context from sglang.srt.layers.dcp import ( diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 3fc207069..29e285803 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -902,7 +902,7 @@ class MQALayer(MqaAttentionBase): use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch) kv: Optional[torch.Tensor] - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) @@ -1156,7 +1156,7 @@ class MQALayer(MqaAttentionBase): # (no DSA-CP), pass `q` as a sentinel for the `k is v` assert; the # attention path doesn't read it once `save_kv_cache=False`. attn_k = kv if kv is not None else q - from sglang.srt.layers.attention.dsv4.unified_kv_kernels.env_gate import ( + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py index 0c488a29f..f95670520 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py @@ -1145,7 +1145,7 @@ def dsa_impl_capability(impl: str) -> tuple[bool, str]: if impl == "tilelang": try: - from sglang.srt.layers.attention.dsa.tilelang_kernel import ( # noqa: F401 + from sglang.kernels.ops.attention.dsa.tilelang_kernel import ( # noqa: F401 tilelang_sparse_fwd, ) except ImportError as exc: diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py index 356fc4db4..0d41fef66 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py @@ -19,12 +19,12 @@ from typing import Any import torch from torch import nn +from sglang.kernels.ops.attention.dsv4.quant_k_cache import ( + quant_to_nope_fp8_rope_bf16_pack_triton, +) from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS -from sglang.srt.layers.attention.dsv4.quant_k_cache import ( - quant_to_nope_fp8_rope_bf16_pack_triton, -) from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool diff --git a/test/manual/layers/attention/dsa/test_act_quant_triton.py b/test/manual/layers/attention/dsa/test_act_quant_triton.py index f24c30285..e9a4d4f62 100644 --- a/test/manual/layers/attention/dsa/test_act_quant_triton.py +++ b/test/manual/layers/attention/dsa/test_act_quant_triton.py @@ -9,8 +9,8 @@ from typing import Tuple import pytest import torch -from sglang.srt.layers.attention.dsa.tilelang_kernel import act_quant -from sglang.srt.layers.attention.dsa.triton_kernel import act_quant as act_quant_triton +from sglang.kernels.ops.attention.dsa.tilelang_kernel import act_quant +from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant as act_quant_triton def benchmark_kernel( diff --git a/test/manual/layers/attention/dsa/test_get_k_scale_triton_kernel.py b/test/manual/layers/attention/dsa/test_get_k_scale_triton_kernel.py index a0c5956e1..2f86dff3d 100644 --- a/test/manual/layers/attention/dsa/test_get_k_scale_triton_kernel.py +++ b/test/manual/layers/attention/dsa/test_get_k_scale_triton_kernel.py @@ -1,6 +1,6 @@ import torch -from sglang.srt.layers.attention.dsa.index_buf_accessor import ( +from sglang.kernels.ops.attention.dsa.index_buf_accessor import ( _get_k_and_s_triton_kernel, ) diff --git a/test/manual/layers/attention/dsa/test_index_buf_accessor.py b/test/manual/layers/attention/dsa/test_index_buf_accessor.py index 3e9f9e374..6c9dcb3d0 100644 --- a/test/manual/layers/attention/dsa/test_index_buf_accessor.py +++ b/test/manual/layers/attention/dsa/test_index_buf_accessor.py @@ -13,7 +13,7 @@ Test coverage: import pytest import torch -from sglang.srt.layers.attention.dsa.index_buf_accessor import GetK, GetKAndS, GetS +from sglang.kernels.ops.attention.dsa.index_buf_accessor import GetK, GetKAndS, GetS class MockDSATokenToKVPool: diff --git a/test/registered/jit/deepseek_v4/test_fp4_indexer.py b/test/registered/jit/deepseek_v4/test_fp4_indexer.py index de332658c..8d4264a71 100644 --- a/test/registered/jit/deepseek_v4/test_fp4_indexer.py +++ b/test/registered/jit/deepseek_v4/test_fp4_indexer.py @@ -15,7 +15,7 @@ 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 ( +from sglang.kernels.ops.attention.dsv4.fp4_indexer import ( quantize_fp4_indexer_tensor, store_fp4_index_k_cache, ) diff --git a/test/registered/jit/test_fused_store_index_cache.py b/test/registered/jit/test_fused_store_index_cache.py index 7b4ceb0b4..1502b6822 100644 --- a/test/registered/jit/test_fused_store_index_cache.py +++ b/test/registered/jit/test_fused_store_index_cache.py @@ -188,7 +188,7 @@ def _reference_quantize_and_store( def _import_act_quant(): try: - from sglang.srt.layers.attention.dsa.triton_kernel import act_quant + from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant return act_quant except Exception: diff --git a/test/registered/kernels/test_dsa_indexer.py b/test/registered/kernels/test_dsa_indexer.py index ba841c19c..e06cf0f80 100644 --- a/test/registered/kernels/test_dsa_indexer.py +++ b/test/registered/kernels/test_dsa_indexer.py @@ -720,7 +720,7 @@ class TestDSAIndexer(CustomTestCase): self.assertEqual(indexer.layer_id, self.config["layer_id"]) @patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm") - @patch("sglang.srt.layers.attention.dsa.triton_kernel.act_quant") + @patch("sglang.kernels.ops.attention.dsa.triton_kernel.act_quant") def test_forward_extend_mode(self, mock_act_quant, mock_deep_gemm): """Test indexer forward pass in extend mode.""" if not self.supports_fp8: @@ -802,7 +802,7 @@ class TestDSAIndexer(CustomTestCase): ) @patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm") - @patch("sglang.srt.layers.attention.dsa.triton_kernel.act_quant") + @patch("sglang.kernels.ops.attention.dsa.triton_kernel.act_quant") def test_forward_decode_mode(self, mock_act_quant, mock_deep_gemm): """Test indexer forward pass in decode mode.""" if not self.supports_fp8: diff --git a/test/registered/kernels/test_dsa_transform_index.py b/test/registered/kernels/test_dsa_transform_index.py index d32219487..e4c32e359 100644 --- a/test/registered/kernels/test_dsa_transform_index.py +++ b/test/registered/kernels/test_dsa_transform_index.py @@ -3,8 +3,8 @@ from unittest.mock import patch import torch -import sglang.srt.layers.attention.dsa.transform_index as transform_index_module -from sglang.srt.layers.attention.dsa.transform_index import ( +import sglang.kernels.ops.attention.dsa.transform_index as transform_index_module +from sglang.kernels.ops.attention.dsa.transform_index import ( transform_index_page_table_decode_fast, transform_index_page_table_prefill_fast, )