From c66a285c94bfeb224697560c88c936948f1fe923 Mon Sep 17 00:00:00 2001 From: Khoa Pham Date: Tue, 1 Sep 2026 22:17:09 -0700 Subject: [PATCH] [Kernel] GLM 5.3 Flash related kernels (ported from #36507) (#37477) Co-authored-by: Claude Opus 5 (1M context) --- .../jit/csrc/dsa/kpool_topk_transform.cuh | 412 ++++ .../sglang/kernels/ops/attention/__init__.py | 1 + .../ops/attention/dsa/transform_index.py | 52 + .../sglang/kernels/ops/attention/fla/kda.py | 5 + .../ops/attention/helion/kda_prefill.py | 3 + python/sglang/kernels/ops/attention/utils.py | 21 + .../sglang/kernels/ops/kvcache/mla_buffer.py | 99 +- python/sglang/kernels/ops/layernorm/mhc.py | 184 ++ .../kernels/ops/moe/kpool_topk_transform.py | 71 + .../layers/attention/dsa/kpool_fp8_index.py | 1701 +++++++++++++++++ .../srt/layers/attention/dsa/kpool_plan.py | 852 +++++++++ .../attention/linear/kernels/kda_cutedsl.py | 2 + .../attention/linear/kernels/kda_flashkda.py | 17 +- .../attention/linear/kernels/kda_triton.py | 4 +- .../kernels/ops/attention/test_kda_helion.py | 41 +- .../kernels/test_dsa_kpool_multi_pool.py | 213 +++ 16 files changed, 3665 insertions(+), 13 deletions(-) create mode 100644 python/sglang/kernels/jit/csrc/dsa/kpool_topk_transform.cuh create mode 100644 python/sglang/kernels/ops/moe/kpool_topk_transform.py create mode 100644 python/sglang/srt/layers/attention/dsa/kpool_fp8_index.py create mode 100644 python/sglang/srt/layers/attention/dsa/kpool_plan.py create mode 100644 test/registered/kernels/test_dsa_kpool_multi_pool.py diff --git a/python/sglang/kernels/jit/csrc/dsa/kpool_topk_transform.cuh b/python/sglang/kernels/jit/csrc/dsa/kpool_topk_transform.cuh new file mode 100644 index 000000000..e2dffd907 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/dsa/kpool_topk_transform.cuh @@ -0,0 +1,412 @@ +// Radix top-k core adapted from https://github.com/tile-ai/tilelang/blob/main/examples/deepseek_v32/topk_selector.py. +// This JIT variant adds pool expansion and page-table/ragged-offset transforms. +#include +#include + +#include + +#include +#include + +#include +#include +#include +#include + +namespace sglang { +namespace { + +#ifndef C10_LIKELY +#define C10_LIKELY(expr) (__builtin_expect(static_cast(expr), 1)) +#endif + +#ifndef SGL_GROUP_TOPK +#define SGL_GROUP_TOPK 256 +#endif + +inline constexpr int kGroupTopK = SGL_GROUP_TOPK; +inline constexpr int kThreadsPerBlock = 1024; + +inline constexpr std::size_t kSmem = 8 * 1024 * sizeof(uint32_t); // 32KB (bytes) + +struct FastTopKParams { + const float* __restrict__ input; // [B, input_stride] + const int32_t* __restrict__ row_starts; // [B] or nullptr + int32_t* __restrict__ indices; // unused here (kept for layout parity) + const int32_t* __restrict__ lengths; // [B] + int64_t input_stride; +}; + +__device__ __forceinline__ auto convert_to_uint8(float x) -> uint8_t { + __half h = __float2half_rn(x); + uint16_t bits = __half_as_ushort(h); + uint16_t key = (bits & 0x8000) ? static_cast(~bits) : static_cast(bits | 0x8000); + return static_cast(key >> 8); +} + +__device__ __forceinline__ auto convert_to_uint32(float x) -> uint32_t { + uint32_t bits = __float_as_uint(x); + return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); +} + +template +__device__ void +fast_topk_cuda_tl_impl(const float* __restrict__ input, int* __restrict__ index, int row_start, int length) { + // We assume length > K here, or it will crash + int topk = K; + constexpr auto BLOCK_SIZE = 1024; + constexpr auto RADIX = 256; + constexpr auto SMEM_INPUT_SIZE = kSmem / (2 * sizeof(int)); + + alignas(128) __shared__ int s_histogram_buf[2][RADIX + 128]; + alignas(128) __shared__ int s_counter; + alignas(128) __shared__ int s_threshold_bin_id; + alignas(128) __shared__ int s_num_input[2]; + + auto& s_histogram = s_histogram_buf[0]; + extern __shared__ int s_input_idx[][SMEM_INPUT_SIZE]; + + const int tx = threadIdx.x; + + if (tx < RADIX + 1) s_histogram[tx] = 0; + __syncthreads(); + + for (int idx = tx; idx < length; idx += BLOCK_SIZE) { + const auto bin = convert_to_uint8(input[idx + row_start]); + ::atomicAdd(&s_histogram[bin], 1); + } + __syncthreads(); + + const auto run_cumsum = [&] { +#pragma unroll 8 + for (int i = 0; i < 8; ++i) { + static_assert(1 << 8 == RADIX); + if (C10_LIKELY(tx < RADIX)) { + const auto j = 1 << i; + const auto k = i & 1; + auto value = s_histogram_buf[k][tx]; + if (tx < RADIX - j) { + value += s_histogram_buf[k][tx + j]; + } + s_histogram_buf[k ^ 1][tx] = value; + } + __syncthreads(); + } + }; + + run_cumsum(); + if (tx < RADIX && s_histogram[tx] > topk && s_histogram[tx + 1] <= topk) { + s_threshold_bin_id = tx; + s_num_input[0] = 0; + s_counter = 0; + } + __syncthreads(); + + const auto threshold_bin = s_threshold_bin_id; + topk -= s_histogram[threshold_bin + 1]; + + if (topk == 0) { + for (int idx = tx; idx < length; idx += BLOCK_SIZE) { + const auto bin = static_cast(convert_to_uint8(input[idx + row_start])); + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + index[pos] = idx; + } + } + __syncthreads(); + return; + } else { + __syncthreads(); + if (tx < RADIX + 1) { + s_histogram[tx] = 0; + } + __syncthreads(); + + for (int idx = tx; idx < length; idx += BLOCK_SIZE) { + const auto raw_input = input[idx + row_start]; + const auto bin = static_cast(convert_to_uint8(raw_input)); + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + index[pos] = idx; + } else if (bin == threshold_bin) { + const auto pos = ::atomicAdd(&s_num_input[0], 1); + if (C10_LIKELY(pos < SMEM_INPUT_SIZE)) { + s_input_idx[0][pos] = idx; + const auto bin = convert_to_uint32(raw_input); + const auto sub_bin = (bin >> 24) & 0xFF; + ::atomicAdd(&s_histogram[sub_bin], 1); + } + } + } + __syncthreads(); + } + +#pragma unroll 4 + for (int round = 0; round < 4; ++round) { + __shared__ int s_last_remain; + const auto r_idx = round % 2; + + const auto _raw_num_input = s_num_input[r_idx]; + const auto num_input = (_raw_num_input < int(SMEM_INPUT_SIZE)) ? _raw_num_input : int(SMEM_INPUT_SIZE); + + run_cumsum(); + if (tx < RADIX && s_histogram[tx] > topk && s_histogram[tx + 1] <= topk) { + s_threshold_bin_id = tx; + s_num_input[r_idx ^ 1] = 0; + s_last_remain = topk - s_histogram[tx + 1]; + } + __syncthreads(); + + const auto threshold_bin = s_threshold_bin_id; + topk -= s_histogram[threshold_bin + 1]; + + if (topk == 0) { + for (int i = tx; i < num_input; i += BLOCK_SIZE) { + const auto idx = s_input_idx[r_idx][i]; + const auto offset = 24 - round * 8; + const auto bin = (convert_to_uint32(input[idx + row_start]) >> offset) & 0xFF; + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + index[pos] = idx; + } + } + __syncthreads(); + break; + } else { + __syncthreads(); + if (tx < RADIX + 1) { + s_histogram[tx] = 0; + } + __syncthreads(); + for (int i = tx; i < num_input; i += BLOCK_SIZE) { + const auto idx = s_input_idx[r_idx][i]; + const auto raw_input = input[idx + row_start]; + const auto offset = 24 - round * 8; + const auto bin = (convert_to_uint32(raw_input) >> offset) & 0xFF; + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + index[pos] = idx; + } else if (bin == threshold_bin) { + if (round == 3) { + const auto pos = ::atomicAdd(&s_last_remain, -1); + if (pos > 0) { + index[K - pos] = idx; + } + } else { + const auto pos = ::atomicAdd(&s_num_input[r_idx ^ 1], 1); + if (C10_LIKELY(pos < SMEM_INPUT_SIZE)) { + s_input_idx[r_idx ^ 1][pos] = idx; + const auto bin = convert_to_uint32(raw_input); + const auto sub_bin = (bin >> (offset - 8)) & 0xFF; + ::atomicAdd(&s_histogram[sub_bin], 1); + } + } + } + } + __syncthreads(); + } + } +} + +__device__ __forceinline__ int32_t transform_kpool_token( + int32_t raw_token, + const int32_t* __restrict__ page_table_entry, + const int32_t* __restrict__ topk_indices_offset, + int32_t offset) { + if (page_table_entry != nullptr) { + return page_table_entry[raw_token]; + } + if (topk_indices_offset != nullptr) { + return raw_token + offset; + } + return raw_token; +} + +template +__global__ __launch_bounds__(kThreadsPerBlock) void kpool_topk_transform_kernel( + const __grid_constant__ FastTopKParams params, + int32_t* __restrict__ dst_token_indices, + const int64_t dst_stride, + const int32_t pool_size, + const int32_t token_topk, + const int32_t out_cols, + const int32_t* __restrict__ page_table, + const int64_t page_table_stride, + const int32_t* __restrict__ page_table_row_index, + const int32_t* __restrict__ topk_indices_offset, + const int32_t* __restrict__ seq_lens) { + const auto& [input, row_starts, _, lengths, input_stride] = params; + const auto bid = static_cast(blockIdx.x); + const auto tid = threadIdx.x; + const auto row_start = row_starts == nullptr ? 0 : row_starts[bid]; + const auto length = lengths[bid]; + const auto score = input + bid * input_stride; + const auto dst = dst_token_indices + bid * dst_stride; + const auto page_table_row = page_table_row_index == nullptr ? bid : static_cast(page_table_row_index[bid]); + const auto page_table_entry = page_table == nullptr ? nullptr : page_table + page_table_row * page_table_stride; + const auto offset = topk_indices_offset == nullptr ? 0 : topk_indices_offset[bid]; + const bool append_tail = seq_lens != nullptr; + const auto full_pool_token_len = length * pool_size; + const auto history_len = full_pool_token_len < token_topk ? full_pool_token_len : token_topk; + const auto tail_count = append_tail ? seq_lens[bid] % pool_size : 0; + + if (length <= K) { + for (int col = tid; col < out_cols; col += kThreadsPerBlock) { + if (col < history_len) { + const auto group_rank = col / pool_size; + const auto slot = col % pool_size; + const auto raw_token = group_rank * pool_size + slot; + dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); + } else if (append_tail && col < history_len + tail_count) { + const auto raw_token = length * pool_size + (col - history_len); + dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); + } else { + dst[col] = -1; + } + } + return; + } + + __shared__ int s_indices[K]; + fast_topk_cuda_tl_impl(score, s_indices, row_start, length); + for (int col = tid; col < out_cols; col += kThreadsPerBlock) { + if (col < history_len) { + const auto group_rank = col / pool_size; + const auto group_id = s_indices[group_rank]; + const auto slot = col % pool_size; + const auto raw_token = group_id * pool_size + slot; + dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); + } else if (append_tail && col < history_len + tail_count) { + const auto raw_token = length * pool_size + (col - history_len); + dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); + } else { + dst[col] = -1; + } + } +} + +template +void setup_kernel_smem_once(host::DebugInfo where = {}) { + [[maybe_unused]] + static const auto result = [] { + const auto fptr = std::bit_cast(f); + return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); + }(); + host::RuntimeDeviceCheck(result, where); +} + +template +const T* optional_data_ptr(const tvm::ffi::Optional& opt) { + if (!opt.has_value()) { + return nullptr; + } + return static_cast(opt.value().data_ptr()); +} + +struct KpoolTopKTransformKernel { + static constexpr auto kernel = kpool_topk_transform_kernel; + + static void transform( + const tvm::ffi::TensorView score, + const tvm::ffi::TensorView lengths, + const tvm::ffi::TensorView dst_token_indices, + const int64_t pool_size, + const tvm::ffi::Optional page_table_opt, + const tvm::ffi::Optional topk_indices_offset_opt, + const tvm::ffi::Optional row_starts_opt, + const tvm::ffi::Optional seq_lens_opt, + const tvm::ffi::Optional page_table_row_index_opt) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto S = SymbolicSize{"score_stride"}; + auto out_cols_sym = SymbolicSize{"out_cols"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({B, -1}).with_strides({S, 1}).with_dtype().with_device(device).verify(score); + TensorMatcher({B}).with_dtype().with_device(device).verify(lengths); + TensorMatcher({B, out_cols_sym}).with_dtype().with_device(device).verify(dst_token_indices); + + RuntimeCheck(pool_size > 1, "pool_size must be > 1, got ", pool_size); + RuntimeCheck( + !(page_table_opt.has_value() && topk_indices_offset_opt.has_value()), + "page_table and topk_indices_offset are mutually exclusive"); + RuntimeCheck( + !page_table_row_index_opt.has_value() || page_table_opt.has_value(), + "page_table_row_index requires page_table"); + + const auto out_cols = static_cast(out_cols_sym.unwrap()); + const auto tail_cols = seq_lens_opt.has_value() ? static_cast(pool_size) - 1 : 0; + RuntimeCheck(out_cols > tail_cols, "dst_token_indices columns ", out_cols, " must exceed tail ", tail_cols); + const auto token_topk = out_cols - tail_cols; + RuntimeCheck(token_topk % static_cast(pool_size) == 0, "token_topk must be a multiple of pool_size"); + RuntimeCheck( + token_topk / static_cast(pool_size) == kGroupTopK, + "this module is built for group_topk=", + kGroupTopK, + " but got ", + token_topk / static_cast(pool_size)); + + const auto batch_size = static_cast(B.unwrap()); + + int64_t page_table_stride = 0; + const int32_t* page_table_ptr = nullptr; + if (page_table_opt.has_value()) { + auto P = SymbolicSize{"page_table_stride"}; + if (page_table_row_index_opt.has_value()) { + auto page_table_rows = SymbolicSize{"page_table_rows"}; + TensorMatcher({page_table_rows, -1}) + .with_strides({P, 1}) + .with_dtype() + .with_device(device) + .verify(page_table_opt.value()); + } else { + TensorMatcher({B, -1}).with_strides({P, 1}).with_dtype().with_device(device).verify( + page_table_opt.value()); + } + page_table_ptr = static_cast(page_table_opt.value().data_ptr()); + page_table_stride = static_cast(P.unwrap()); + } + + if (topk_indices_offset_opt.has_value()) { + TensorMatcher({B}).with_dtype().with_device(device).verify(topk_indices_offset_opt.value()); + } + if (row_starts_opt.has_value()) { + TensorMatcher({B}).with_dtype().with_device(device).verify(row_starts_opt.value()); + } + if (seq_lens_opt.has_value()) { + TensorMatcher({B}).with_dtype().with_device(device).verify(seq_lens_opt.value()); + } + if (page_table_row_index_opt.has_value()) { + TensorMatcher({B}).with_dtype().with_device(device).verify(page_table_row_index_opt.value()); + } + + const auto params = FastTopKParams{ + .input = static_cast(score.data_ptr()), + .row_starts = optional_data_ptr(row_starts_opt), + .indices = nullptr, + .lengths = static_cast(lengths.data_ptr()), + .input_stride = static_cast(S.unwrap()), + }; + + setup_kernel_smem_once(); + LaunchKernel(batch_size, kThreadsPerBlock, device.unwrap(), kSmem)( + kernel, + params, + static_cast(dst_token_indices.data_ptr()), + static_cast(dst_token_indices.strides()[0]), + static_cast(pool_size), + token_topk, + out_cols, + page_table_ptr, + page_table_stride, + optional_data_ptr(page_table_row_index_opt), + optional_data_ptr(topk_indices_offset_opt), + optional_data_ptr(seq_lens_opt)); + } +}; + +} // namespace + +} // namespace sglang diff --git a/python/sglang/kernels/ops/attention/__init__.py b/python/sglang/kernels/ops/attention/__init__.py index 674cbd50d..b7516f92d 100644 --- a/python/sglang/kernels/ops/attention/__init__.py +++ b/python/sglang/kernels/ops/attention/__init__.py @@ -95,6 +95,7 @@ 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.transform_index", "prepare_trtllm_nope_sparse_metadata"), ("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"), diff --git a/python/sglang/kernels/ops/attention/dsa/transform_index.py b/python/sglang/kernels/ops/attention/dsa/transform_index.py index 2d7dd105e..8fbfa190f 100644 --- a/python/sglang/kernels/ops/attention/dsa/transform_index.py +++ b/python/sglang/kernels/ops/attention/dsa/transform_index.py @@ -14,6 +14,58 @@ def transform_index_page_table_decode(**kwargs): return transform_index_page_table_decode_fast(**kwargs) +@triton.jit +def prepare_trtllm_nope_sparse_metadata_kernel( + page_table_ptr: torch.Tensor, + topk_lens_ptr: torch.Tensor, + row_stride: tl.constexpr, + TOPK: tl.constexpr, + BLOCK_TOPK: tl.constexpr, +): + row = tl.program_id(0) + offsets = tl.arange(0, BLOCK_TOPK) + valid = offsets < TOPK + indices = tl.load( + page_table_ptr + row * row_stride + offsets, + mask=valid, + other=-1, + ) + topk_len = tl.sum((valid & (indices >= 0)).to(tl.int32), axis=0) + + # TRTLLM-GEN's native H512 dynamic sparse kernel produces NaNs for an + # empty row. CUDA-graph padding rows are never consumed, so point them at + # a valid dummy token and run a one-element attention instead. + is_empty = topk_len == 0 + tl.store(page_table_ptr + row * row_stride, 0, mask=is_empty) + tl.store(topk_lens_ptr + row, tl.maximum(topk_len, 1)) + + +def prepare_trtllm_nope_sparse_metadata( + page_table: torch.Tensor, +) -> torch.Tensor: + """Build per-query active top-k lengths for native H512 TRTLLM-GEN MLA. + + ``page_table`` must contain packed valid token locations followed by ``-1`` + padding. The tensor is only modified for fully empty CUDA-graph padding + rows, whose first entry is replaced with the valid dummy token location 0. + """ + assert page_table.ndim == 2 + assert page_table.dtype == torch.int32 + assert page_table.is_contiguous() + num_rows, topk = page_table.shape + topk_lens = torch.empty(num_rows, dtype=torch.int32, device=page_table.device) + block_topk = triton.next_power_of_2(topk) + prepare_trtllm_nope_sparse_metadata_kernel[(num_rows,)]( + page_table, + topk_lens, + page_table.stride(0), + TOPK=topk, + BLOCK_TOPK=block_topk, + num_warps=8, + ) + return topk_lens + + def _allocate_prefill_result( topk_indices: torch.Tensor, real_num_tokens: int, diff --git a/python/sglang/kernels/ops/attention/fla/kda.py b/python/sglang/kernels/ops/attention/fla/kda.py index c5d2defc1..ad9720c53 100644 --- a/python/sglang/kernels/ops/attention/fla/kda.py +++ b/python/sglang/kernels/ops/attention/fla/kda.py @@ -1155,6 +1155,7 @@ def chunk_kda_fwd( cu_seqlens=cu_seqlens, chunk_size=chunk_size, chunk_indices=chunk_indices, + safe_gate=lower_bound is not None, fuse_diagonal=_small_grid, fuse_recompute=_small_grid, ) @@ -1210,6 +1211,7 @@ def chunk_kda( dt_bias: Optional[torch.Tensor] = None, lower_bound: Optional[float] = None, output_intermediate_states: bool = False, + beta_is_raw: bool = False, **kwargs, ): if scale is None: @@ -1219,6 +1221,9 @@ def chunk_kda( q = l2norm_fwd(q.contiguous()) k = l2norm_fwd(k.contiguous()) + if beta_is_raw: + beta = beta.float().sigmoid() + # Returns o [B, T, H, V] when output_intermediate_states=False, or (o, h [B, NT, H, V, K]) when output_intermediate_states=True. return chunk_kda_fwd( q=q, diff --git a/python/sglang/kernels/ops/attention/helion/kda_prefill.py b/python/sglang/kernels/ops/attention/helion/kda_prefill.py index 25e0eb3c1..ad67fea72 100644 --- a/python/sglang/kernels/ops/attention/helion/kda_prefill.py +++ b/python/sglang/kernels/ops/attention/helion/kda_prefill.py @@ -1314,6 +1314,7 @@ def chunk_kda( dt_bias: torch.Tensor | None = None, lower_bound: float | None = None, output_intermediate_states: bool = False, + beta_is_raw: bool = False, **kwargs: object, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: """Match the public forward contract of SGLang's Triton ``chunk_kda``.""" @@ -1327,6 +1328,8 @@ def chunk_kda( raise ValueError("g and beta must cover every q token") g = g[:, :num_tokens] beta = beta[:, :num_tokens] + if beta_is_raw: + beta = beta.float().sigmoid() if num_tokens == 1: # Tracing constant-folds size-one dimensions, but the resulting kernel # can share a cache entry with longer inputs. Keep T=1 on Triton so a diff --git a/python/sglang/kernels/ops/attention/utils.py b/python/sglang/kernels/ops/attention/utils.py index d70fe65a1..baa7dd27a 100644 --- a/python/sglang/kernels/ops/attention/utils.py +++ b/python/sglang/kernels/ops/attention/utils.py @@ -180,6 +180,27 @@ def mla_quantize_and_rope_for_fp8( return q_out, k_nope_out, k_rope_out +def mla_quantize_for_fp8_no_rope( + q_nope: torch.Tensor, + q_rope: torch.Tensor, + k_nope: torch.Tensor, + k_rope: torch.Tensor, + kv_lora_rank: int, + qk_rope_head_dim: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + attn_dtype = torch.float8_e4m3fn + q_len, num_heads = q_rope.shape[:2] + q_out = q_rope.new_empty( + q_len, + num_heads, + kv_lora_rank + qk_rope_head_dim, + dtype=attn_dtype, + ) + q_out[..., :kv_lora_rank] = q_nope.to(attn_dtype) + q_out[..., kv_lora_rank:] = q_rope.to(attn_dtype) + return q_out, k_nope.to(attn_dtype), k_rope.to(attn_dtype) + + def mla_quantize_without_rope_for_fp8( q_nope: torch.Tensor, q_rope: torch.Tensor, diff --git a/python/sglang/kernels/ops/kvcache/mla_buffer.py b/python/sglang/kernels/ops/kvcache/mla_buffer.py index 84915d7eb..99ff8b8d2 100644 --- a/python/sglang/kernels/ops/kvcache/mla_buffer.py +++ b/python/sglang/kernels/ops/kvcache/mla_buffer.py @@ -1,5 +1,7 @@ from __future__ import annotations +from typing import Optional + import torch import triton import triton.language as tl @@ -81,6 +83,40 @@ def set_mla_kv_buffer_kernel( tl.extra.cuda.gdc_launch_dependents() +@triton.jit +def set_mla_kv_buffer_kernel_norope( + kv_buffer_ptr, + cache_k_nope_ptr, + loc_ptr, + buffer_stride: tl.constexpr, + nope_stride: tl.constexpr, + nope_dim: tl.constexpr, + BLOCK: tl.constexpr, + USE_GDC: tl.constexpr = False, +): + pid_loc = tl.program_id(0) + pid_blk = tl.program_id(1) + + base = pid_blk * BLOCK + offs = base + tl.arange(0, BLOCK) + mask = offs < nope_dim + + if USE_GDC: + tl.extra.cuda.gdc_wait() + + loc = tl.load(loc_ptr + pid_loc).to(tl.int64) + dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs + + src = tl.load( + cache_k_nope_ptr + pid_loc * nope_stride + offs, + mask=mask, + ) + tl.store(dst_ptr, src, mask=mask) + + if USE_GDC: + tl.extra.cuda.gdc_launch_dependents() + + # Above this loc count the TMA bulk-store path overtakes the single-CTA-per-loc # Triton kernel. Below it, Triton with BLOCK = next_pow2(total_dim) (one CTA # does the whole row in one tile, no boundary fan-out) is the winning fallback. @@ -92,7 +128,7 @@ def _set_mla_kv_buffer_impl( kv_buffer: torch.Tensor, loc: torch.Tensor, cache_k_nope: torch.Tensor, - cache_k_rope: torch.Tensor, + cache_k_rope: Optional[torch.Tensor] = None, *, reserved_skip_index: int, dcp_world_size: int, @@ -127,6 +163,28 @@ def _set_mla_kv_buffer_impl( Shared body of the two entry points below; the owner rule reaches it as ``1, 0`` (nothing to select) or as the live topology. """ + has_rope = cache_k_rope is not None and cache_k_rope.numel() > 0 + n_loc = loc.numel() + nope_dim = cache_k_nope.shape[-1] + + if not has_rope: + BLOCK = triton.next_power_of_2(nope_dim) + grid = (n_loc, 1) + pdl_kwargs = ( + {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {} + ) + set_mla_kv_buffer_kernel_norope[grid]( + kv_buffer, + cache_k_nope, + loc, + kv_buffer.stride(0), + cache_k_nope.stride(0), + nope_dim, + BLOCK=BLOCK, + **pdl_kwargs, + ) + return + from sglang.kernels.ops.kvcache.set_mla_kv_buffer import ( can_use_set_mla_kv_buffer, ) @@ -134,7 +192,6 @@ def _set_mla_kv_buffer_impl( set_mla_kv_buffer as jit_set_mla_kv_buffer, ) - n_loc = loc.numel() nope_bytes = cache_k_nope.shape[-1] * cache_k_nope.element_size() rope_bytes = cache_k_rope.shape[-1] * cache_k_rope.element_size() if ( @@ -157,7 +214,6 @@ def _set_mla_kv_buffer_impl( # ``set_mla_kv_buffer_kernel`` handles the over-allocation past total_dim # via the offs 0 + if not has_rope: + get_mla_kv_buffer_kernel_norope[grid]( + kv_buffer, + cache_k_nope, + loc, + kv_buffer.stride(0), + cache_k_nope.stride(0), + nope_dim, + ) + return + + rope_dim = cache_k_rope.shape[-1] # 64 get_mla_kv_buffer_kernel[grid]( kv_buffer, cache_k_nope, diff --git a/python/sglang/kernels/ops/layernorm/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py index 35a62ebe3..a6a56318e 100644 --- a/python/sglang/kernels/ops/layernorm/mhc.py +++ b/python/sglang/kernels/ops/layernorm/mhc.py @@ -1734,6 +1734,190 @@ def mhc_fused_post_pre( ) +def hc_expand(x: torch.Tensor, n: int) -> torch.Tensor: + return x.repeat(1, n) + + +def hc_contract(x: torch.Tensor, n: int) -> torch.Tensor: + return x.unflatten(-1, (n, -1)).mean(dim=-2) + + +def _mhc_pre_torch( + residual: torch.Tensor, + fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + rms_eps: float, + hc_pre_eps: float, + hc_sinkhorn_eps: float, + hc_post_mult_value: float, + sinkhorn_repeat: int, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + import torch.nn.functional as F + + s, n, h = residual.shape + dtype = residual.dtype + + x_flat = residual.view(s, n * h).float() + rsqrt = torch.rsqrt(x_flat.square().mean(-1, keepdim=True) + rms_eps) + mixes = F.linear(x_flat, fn) * rsqrt + + pre_raw = mixes[:, :n] + post_raw = mixes[:, n : 2 * n] + comb_raw = mixes[:, 2 * n :].view(s, n, n) + pre_base = hc_base[:n] + post_base = hc_base[n : 2 * n] + comb_base = hc_base[2 * n :].view(n, n) + + pre = torch.sigmoid(pre_raw * hc_scale[0] + pre_base) + hc_pre_eps + post = hc_post_mult_value * torch.sigmoid(post_raw * hc_scale[1] + post_base) + comb = comb_raw * hc_scale[2] + comb_base + + comb = comb.softmax(-1) + hc_sinkhorn_eps + comb = comb / (comb.sum(-2, keepdim=True) + hc_sinkhorn_eps) + for _ in range(sinkhorn_repeat - 1): + comb = comb / (comb.sum(-1, keepdim=True) + hc_sinkhorn_eps) + comb = comb / (comb.sum(-2, keepdim=True) + hc_sinkhorn_eps) + + layer_input = (pre.unsqueeze(-1) * residual.float()).sum(dim=1).to(dtype) + return post.unsqueeze(-1), comb, layer_input + + +def _mhc_post_torch( + x: torch.Tensor, + residual: torch.Tensor, + post_layer_mix: torch.Tensor, + comb_res_mix: torch.Tensor, +) -> torch.Tensor: + out = post_layer_mix * x.unsqueeze(1) + ( + comb_res_mix.unsqueeze(-1) * residual.unsqueeze(2) + ).sum(dim=1) + return out.type_as(x) + + +@torch._dynamo.disable +def _mhc_pre_dispatch( + residual: torch.Tensor, + fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + rms_eps: float, + hc_pre_eps: float, + hc_sinkhorn_eps: float, + hc_post_mult_value: float, + sinkhorn_repeat: int, + norm_weight: torch.Tensor | None = None, + norm_eps: float | None = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, bool]: + assert residual.dim() == 3, f"residual must be (s, n, h); got {residual.shape}" + if not envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get(): + post_mix, comb_mix, layer_input = _mhc_pre_torch( + residual=residual, + fn=fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=rms_eps, + hc_pre_eps=hc_pre_eps, + hc_sinkhorn_eps=hc_sinkhorn_eps, + hc_post_mult_value=hc_post_mult_value, + sinkhorn_repeat=sinkhorn_repeat, + ) + return post_mix, comb_mix, layer_input, False + + post_mix, comb_mix, layer_input = mhc_pre( + residual=residual, + fn=fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=rms_eps, + hc_pre_eps=hc_pre_eps, + hc_sinkhorn_eps=hc_sinkhorn_eps, + hc_post_mult_value=hc_post_mult_value, + sinkhorn_repeat=sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + return post_mix, comb_mix, layer_input, norm_weight is not None + + +@torch._dynamo.disable +def _mhc_post_dispatch( + x: torch.Tensor, + residual: torch.Tensor, + post_layer_mix: torch.Tensor, + comb_res_mix: torch.Tensor, +) -> torch.Tensor: + assert x.dim() == 2 and residual.dim() == 3 + assert post_layer_mix.dim() == 3 and comb_res_mix.dim() == 3 + if not envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get(): + return _mhc_post_torch(x, residual, post_layer_mix, comb_res_mix) + return mhc_post(x, residual, post_layer_mix, comb_res_mix) + + +def hc_pre( + x: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + hc_mult: int, + rms_eps: float, + hc_eps: float, + sinkhorn_iters: int, + post_mult_value: float = 2.0, + hc_norm_weight: torch.Tensor | None = None, + out_norm_weight: torch.Tensor | None = None, + out_norm_eps: float | None = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, bool]: + s, total = x.shape + hidden_size = total // hc_mult + if x.numel() == 0: + empty_layer_input = x.new_zeros((s, hidden_size)) + empty_h_res = torch.zeros( + (s, hc_mult * hc_mult), device=x.device, dtype=torch.float32 + ) + empty_h_post = torch.zeros((s, hc_mult), device=x.device, dtype=torch.float32) + return empty_layer_input, empty_h_res, empty_h_post, False + + fn = hc_fn if hc_norm_weight is None else hc_fn * hc_norm_weight + residual_3d = x.view(s, hc_mult, hidden_size) + post_mix, comb_mix, layer_input, norm_fused = _mhc_pre_dispatch( + residual=residual_3d, + fn=fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=rms_eps, + hc_pre_eps=hc_eps, + hc_sinkhorn_eps=hc_eps, + hc_post_mult_value=post_mult_value, + sinkhorn_repeat=sinkhorn_iters, + norm_weight=out_norm_weight, + norm_eps=out_norm_eps, + ) + return ( + layer_input, + comb_mix.reshape(s, hc_mult * hc_mult), + post_mix.reshape(s, hc_mult), + norm_fused, + ) + + +def hc_post( + x: torch.Tensor, + residual: torch.Tensor, + h_post: torch.Tensor, + h_res: torch.Tensor, + hc_mult: int, +) -> torch.Tensor: + s, hidden_size = x.shape + if s == 0: + return x.new_zeros((s, hc_mult * hidden_size)) + residual = residual.view(s, hc_mult, hidden_size) + h_post = h_post.view(s, hc_mult, 1) + h_res = h_res.view(s, hc_mult, hc_mult) + out = _mhc_post_dispatch(x, residual, h_post, h_res) + return out.view(s, -1) + + def npu_hc_pre( x: torch.Tensor, hc_fn: torch.Tensor, diff --git a/python/sglang/kernels/ops/moe/kpool_topk_transform.py b/python/sglang/kernels/ops/moe/kpool_topk_transform.py new file mode 100644 index 000000000..89f066f08 --- /dev/null +++ b/python/sglang/kernels/ops/moe/kpool_topk_transform.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernels.jit.utils import cache_once, load_jit + +if TYPE_CHECKING: + from tvm_ffi.module import Module + +SUPPORTED_GROUP_TOPK = (128, 160, 192, 224, 256, 512) + + +@cache_once +def _jit_kpool_topk_transform_module(group_topk: int) -> Module: + assert group_topk in SUPPORTED_GROUP_TOPK, ( + "fast_kpool_topk_transform supports pool-level topk " + f"{SUPPORTED_GROUP_TOPK}, got {group_topk}" + ) + return load_jit( + f"kpool_topk_transform_{group_topk}", + cuda_files=["dsa/kpool_topk_transform.cuh"], + cuda_wrappers=[("kpool_topk_transform", "KpoolTopKTransformKernel::transform")], + extra_cuda_cflags=[f"-DSGL_GROUP_TOPK={group_topk}"], + ) + + +def fast_kpool_topk_transform_fused( + score: torch.Tensor, + lengths: torch.Tensor, + pool_size: int, + topk: int, + page_table: Optional[torch.Tensor] = None, + topk_indices_offset: Optional[torch.Tensor] = None, + row_starts: Optional[torch.Tensor] = None, + seq_lens: Optional[torch.Tensor] = None, + page_table_row_index: Optional[torch.Tensor] = None, +) -> torch.Tensor: + assert topk % pool_size == 0 + group_topk = topk // pool_size + assert group_topk in SUPPORTED_GROUP_TOPK, ( + "fast_kpool_topk_transform supports pool-level topk " + f"{SUPPORTED_GROUP_TOPK}, got {group_topk}" + ) + assert score.dim() == 2 + assert page_table is None or topk_indices_offset is None + assert page_table_row_index is None or page_table is not None + if seq_lens is not None: + assert seq_lens.dim() == 1 + assert seq_lens.shape[0] == score.shape[0] + if page_table_row_index is not None: + assert page_table_row_index.dim() == 1 + assert page_table_row_index.shape[0] == score.shape[0] + + out_cols = topk + (pool_size - 1 if seq_lens is not None else 0) + dst_token_indices = score.new_empty((score.shape[0], out_cols), dtype=torch.int32) + + module = _jit_kpool_topk_transform_module(group_topk) + module.kpool_topk_transform( + score, + lengths, + dst_token_indices, + pool_size, + page_table, + topk_indices_offset, + row_starts, + seq_lens, + page_table_row_index, + ) + return dst_token_indices diff --git a/python/sglang/srt/layers/attention/dsa/kpool_fp8_index.py b/python/sglang/srt/layers/attention/dsa/kpool_fp8_index.py new file mode 100644 index 000000000..f9bf23137 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/kpool_fp8_index.py @@ -0,0 +1,1701 @@ +from typing import Optional, Tuple + +import torch +import triton +import triton.language as tl + +BLOCK_SIZE_K = 64 +INDEX_HEAD_DIM = 128 +KPOOL_SCORE_DTYPES = (torch.float16, torch.bfloat16, torch.float32) + + +def kpool_max_closed_pools(num_draft_tokens: int, pool_size: int) -> int: + return (num_draft_tokens + pool_size - 1) // pool_size + + +def build_pooled_page_table_64( + page_table_64: torch.Tensor, + pool_size: int, +) -> torch.Tensor: + # Advanced indexing is required: a (1, 1) strided slice can remain non-unit- + # stride even after contiguous(), which DeepGEMM rejects. + assert ( + BLOCK_SIZE_K % pool_size == 0 + ), f"pool_size ({pool_size}) must divide page_size ({BLOCK_SIZE_K})" + idx = torch.arange( + 0, page_table_64.shape[-1], pool_size, device=page_table_64.device + ) + return page_table_64[..., idx] + + +def gather_index_k_scale_prefix_into( + pool, + buf: torch.Tensor, + page_indices: torch.Tensor, + seq_len: int, + k_out: torch.Tensor, + scale_out: torch.Tensor, +) -> None: + assert buf.dtype == torch.uint8 + assert page_indices.dtype in (torch.int32, torch.int64) + assert k_out.dtype == torch.uint8 + assert scale_out.dtype == torch.float32 + assert pool.page_size == BLOCK_SIZE_K + assert k_out.shape[0] >= seq_len + assert k_out.shape[1] == INDEX_HEAD_DIM + assert scale_out.shape[0] >= seq_len + assert buf.is_contiguous() + assert page_indices.is_contiguous() + assert k_out.is_contiguous() + assert scale_out.is_contiguous() + if seq_len == 0: + return + + _gather_index_k_scale_prefix_into_kernel[(seq_len,)]( + buf, + buf.view(torch.float32), + page_indices, + k_out, + scale_out, + PAGE_SIZE=pool.page_size, + BUF_NUMEL_PER_PAGE=buf.shape[1], + HEAD_DIM=INDEX_HEAD_DIM, + S_OFFSET_NBYTES_IN_PAGE=pool.page_size * INDEX_HEAD_DIM, + BLOCK_D=triton.next_power_of_2(INDEX_HEAD_DIM), + ) + + +@triton.jit +def _gather_index_k_scale_prefix_into_kernel( + buf_u8_ptr, + buf_fp32_ptr, + page_indices_ptr, + k_out_ptr, + scale_out_ptr, + PAGE_SIZE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + HEAD_DIM: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + BLOCK_D: tl.constexpr, +): + token_id = tl.program_id(0) + page_idx = token_id // PAGE_SIZE + token_offset_in_page = token_id % PAGE_SIZE + page = tl.load(page_indices_ptr + page_idx) + + offs = tl.arange(0, BLOCK_D) + mask = offs < HEAD_DIM + src_k_offsets = page * BUF_NUMEL_PER_PAGE + token_offset_in_page * HEAD_DIM + offs + dst_k_offsets = token_id * HEAD_DIM + offs + k = tl.load(buf_u8_ptr + src_k_offsets, mask=mask) + tl.store(k_out_ptr + dst_k_offsets, k, mask=mask) + + src_s_offset = ( + page * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + token_offset_in_page + ) + scale = tl.load(buf_fp32_ptr + src_s_offset) + tl.store(scale_out_ptr + token_id, scale) + + +def kpool_build_ragged_layout( + full_page_table: torch.Tensor, + cu_pages_excl: torch.Tensor, + ragged_pool_pages: torch.Tensor, + cu_q_len_excl: torch.Tensor, + ragged_q_len: torch.Tensor, + pooled_seq_lens_expanded: torch.Tensor, + slots_per_page: int, + total_pool_pages: int, + total_q: int, + pool_size: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + # One packed-cache page represents 64 pools, so its source is every + # pool_size-th real token page. + device = full_page_table.device + n_rag = cu_pages_excl.shape[0] + concat_page_table = torch.empty( + (total_pool_pages,), dtype=full_page_table.dtype, device=device + ) + q_ks = torch.empty((total_q,), dtype=torch.int32, device=device) + q_ke = torch.empty((total_q,), dtype=torch.int32, device=device) + if n_rag == 0: + return concat_page_table, q_ks, q_ke + + max_pool_pages = full_page_table.shape[1] + _kpool_build_ragged_layout_kernel[(n_rag,)]( + full_page_table, + cu_pages_excl, + ragged_pool_pages, + cu_q_len_excl, + ragged_q_len, + pooled_seq_lens_expanded, + concat_page_table, + q_ks, + q_ke, + max_pool_pages, + slots_per_page, + pool_size, + BLOCK_PAGE=128, + BLOCK_Q=128, + ) + return concat_page_table, q_ks, q_ke + + +@triton.jit +def _kpool_build_ragged_layout_kernel( + full_page_table_ptr, + cu_pages_excl_ptr, + ragged_pool_pages_ptr, + cu_q_len_excl_ptr, + ragged_q_len_ptr, + pooled_seq_lens_ptr, + concat_page_table_ptr, + q_ks_ptr, + q_ke_ptr, + MAX_POOL_PAGES, + SLOTS_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + BLOCK_PAGE: tl.constexpr, + BLOCK_Q: tl.constexpr, +): + k = tl.program_id(0) + page_start = tl.load(cu_pages_excl_ptr + k) + n_pages = tl.load(ragged_pool_pages_ptr + k) + q_start = tl.load(cu_q_len_excl_ptr + k) + q_count = tl.load(ragged_q_len_ptr + k) + ks_val = page_start * SLOTS_PER_PAGE + + for p_off in tl.range(0, BLOCK_PAGE * tl.cdiv(n_pages, BLOCK_PAGE), BLOCK_PAGE): + p_offs = p_off + tl.arange(0, BLOCK_PAGE) + p_mask = p_offs < n_pages + source_cols = p_offs * POOL_SIZE + pages = tl.load( + full_page_table_ptr + k * MAX_POOL_PAGES + source_cols, + mask=p_mask, + other=0, + ) + tl.store(concat_page_table_ptr + page_start + p_offs, pages, mask=p_mask) + + for q_off in tl.range(0, BLOCK_Q * tl.cdiv(q_count, BLOCK_Q), BLOCK_Q): + q_offs = q_off + tl.arange(0, BLOCK_Q) + q_mask = q_offs < q_count + plen = tl.load(pooled_seq_lens_ptr + q_start + q_offs, mask=q_mask, other=0) + ke_val = tl.minimum( + ks_val + plen, + (page_start + n_pages) * SLOTS_PER_PAGE, + ) + tl.store( + q_ks_ptr + q_start + q_offs, + tl.full([BLOCK_Q], ks_val, tl.int32), + mask=q_mask, + ) + tl.store(q_ke_ptr + q_start + q_offs, ke_val, mask=q_mask) + + +def _prep_update_kpool_write_plan_launch( + write_start: torch.Tensor, + req_pool_indices: torch.Tensor, + real_page_table: torch.Tensor, + req_out: torch.Tensor, + write_start_out: torch.Tensor, + tail_logical_start_out: torch.Tensor, + write_loc_out: torch.Tensor, + pool_seqlens_per_q_out: Optional[torch.Tensor], + seqlens_per_q_out: Optional[torch.Tensor], + *, + pool_size: int, + num_draft_tokens: int, + slots_per_page: int, +): + # Eager and captured replay must share one launch spec, including + # N > pool_size where one write closes multiple pools. + max_closed_pools = kpool_max_closed_pools(num_draft_tokens, pool_size) + bs = write_start.shape[0] + + assert write_loc_out.shape == (bs, max_closed_pools), write_loc_out.shape + assert write_loc_out.stride(1) == 1, write_loc_out.stride() + + has_per_q_outputs = pool_seqlens_per_q_out is not None + assert has_per_q_outputs == ( + seqlens_per_q_out is not None + ), "pool_seqlens_per_q_out and seqlens_per_q_out must be both set or both None" + per_q_dummy = ( + pool_seqlens_per_q_out + if has_per_q_outputs + else torch.empty(1, dtype=torch.int32, device=write_start.device) + ) + + args = ( + write_start, + req_pool_indices, + real_page_table, + req_out, + write_start_out, + tail_logical_start_out, + write_loc_out, + pool_seqlens_per_q_out if has_per_q_outputs else per_q_dummy, + seqlens_per_q_out if has_per_q_outputs else per_q_dummy, + real_page_table.stride(0), + real_page_table.shape[1], + write_loc_out.stride(0), + ) + constexprs = dict( + POOL_SIZE=pool_size, + N=num_draft_tokens, + SLOTS_PER_PAGE=slots_per_page, + MAX_CLOSED_POOLS=max_closed_pools, + HAS_PER_Q=has_per_q_outputs, + ) + return bs, args, constexprs + + +def update_kpool_write_plan_cuda_graph( + write_start: torch.Tensor, + req_pool_indices: torch.Tensor, + real_page_table: torch.Tensor, + req_out: torch.Tensor, + write_start_out: torch.Tensor, + tail_logical_start_out: torch.Tensor, + write_loc_out: torch.Tensor, + pool_seqlens_per_q_out: Optional[torch.Tensor], + seqlens_per_q_out: Optional[torch.Tensor], + *, + pool_size: int, + num_draft_tokens: int, + slots_per_page: int, +) -> None: + if write_start.shape[0] == 0: + return + bs, args, constexprs = _prep_update_kpool_write_plan_launch( + write_start, + req_pool_indices, + real_page_table, + req_out, + write_start_out, + tail_logical_start_out, + write_loc_out, + pool_seqlens_per_q_out, + seqlens_per_q_out, + pool_size=pool_size, + num_draft_tokens=num_draft_tokens, + slots_per_page=slots_per_page, + ) + _update_kpool_write_plan_kernel[(bs,)](*args, **constexprs) + + +@triton.jit +def _update_kpool_write_plan_kernel( + write_start_ptr, + req_pool_indices_ptr, + real_page_table_ptr, + req_out_ptr, + write_start_out_ptr, + tail_logical_start_out_ptr, + write_loc_out_ptr, + pool_seqlens_per_q_out_ptr, + seqlens_per_q_out_ptr, + real_page_table_stride_0, + real_page_table_cols, + write_loc_out_stride_0, + POOL_SIZE: tl.constexpr, + N: tl.constexpr, + SLOTS_PER_PAGE: tl.constexpr, + MAX_CLOSED_POOLS: tl.constexpr, + HAS_PER_Q: tl.constexpr, +): + b = tl.program_id(0) + ws = tl.load(write_start_ptr + b).to(tl.int32) + req = tl.load(req_pool_indices_ptr + b) + base_pool = ws // POOL_SIZE + + if HAS_PER_Q: + for k in tl.static_range(0, N): + row = b * N + k + seqlen_per_q = ws + k + 1 + tl.store(seqlens_per_q_out_ptr + row, seqlen_per_q) + tl.store(pool_seqlens_per_q_out_ptr + row, seqlen_per_q // POOL_SIZE) + + tl.store(req_out_ptr + b, req) + tl.store(write_start_out_ptr + b, ws) + tl.store(tail_logical_start_out_ptr + b, (base_pool * POOL_SIZE).to(tl.int32)) + for p in tl.static_range(0, MAX_CLOSED_POOLS): + pool_id = base_pool + p + pool_page_group = pool_id // SLOTS_PER_PAGE + token_page_row = pool_page_group * POOL_SIZE + token_page_row = tl.minimum( + tl.maximum(token_page_row, 0), real_page_table_cols - 1 + ) + # Promote the row term before multiplication; 1M-context strides + # overflow int32 around row 2048. + packed_page = tl.load( + real_page_table_ptr + + (b * N).to(tl.int64) * real_page_table_stride_0 + + token_page_row + ).to(tl.int64) + write_loc = packed_page * SLOTS_PER_PAGE + (pool_id % SLOTS_PER_PAGE) + tl.store( + write_loc_out_ptr + b * write_loc_out_stride_0 + p, + write_loc.to(tl.int64), + ) + + +def compute_pooled_write_locs( + page_table_64: torch.Tensor, + pool_ids: torch.Tensor, + pool_size: int, +) -> torch.Tensor: + assert page_table_64.ndim == 1 + pool_ids = pool_ids.to(torch.int64) + pool_page_group = torch.div(pool_ids, BLOCK_SIZE_K, rounding_mode="floor") + token_page_row = pool_page_group * pool_size + packed_page = page_table_64.index_select(0, token_page_row.to(torch.int64)) + return packed_page.to(torch.int64) * BLOCK_SIZE_K + torch.remainder( + pool_ids, BLOCK_SIZE_K + ) + + +def history_group_budget_for_topk(topk: int, pool_size: int) -> int: + assert topk % pool_size == 0 + return topk // pool_size + + +def expand_pooled_groups_to_topk( + group_ids: torch.Tensor, + group_valid: torch.Tensor, + topk: int, + pool_size: int, + page_table: torch.Tensor | None = None, + topk_offsets: torch.Tensor | None = None, +) -> torch.Tensor: + assert group_ids.ndim == 2 + assert group_valid.shape == group_ids.shape + assert topk % pool_size == 0 + assert group_ids.shape[1] == history_group_budget_for_topk(topk, pool_size) + assert page_table is None or topk_offsets is None + + device = group_ids.device + offsets = torch.arange(pool_size, device=device, dtype=torch.int64) + token_ids = group_ids.to(torch.int64).unsqueeze(-1) * pool_size + offsets + token_ids = token_ids.reshape(group_ids.shape[0], topk) + valid = ( + group_valid.unsqueeze(-1) + .expand(-1, -1, pool_size) + .reshape(group_ids.shape[0], topk) + ) + + if page_table is not None: + assert page_table.ndim == 2 + assert page_table.shape[0] == group_ids.shape[0] + safe_ids = token_ids.clamp(min=0, max=page_table.shape[1] - 1) + output = torch.gather(page_table, dim=1, index=safe_ids).to(torch.int32) + elif topk_offsets is not None: + if topk_offsets.ndim == 2: + assert topk_offsets.shape[1] == 1 + topk_offsets = topk_offsets.squeeze(1) + assert topk_offsets.ndim == 1 + output = (token_ids + topk_offsets.to(torch.int64).unsqueeze(1)).to(torch.int32) + else: + output = token_ids.to(torch.int32) + + return torch.where(valid, output, torch.full_like(output, -1)) + + +def append_kpool_tail_to_topk( + topk_result: torch.Tensor, + seq_lens: torch.Tensor, + pool_lens: torch.Tensor, + pool_size: int, + page_table: torch.Tensor | None = None, + topk_offsets: torch.Tensor | None = None, +) -> torch.Tensor: + assert topk_result.dtype == torch.int32 + assert seq_lens.ndim == 1 + assert pool_lens.ndim == 1 + assert seq_lens.shape[0] == topk_result.shape[0] + assert pool_lens.shape[0] == topk_result.shape[0] + + tail_pool = pool_size - 1 + if tail_pool == 0: + return topk_result + + rows, n_cols = topk_result.shape + out_cols = n_cols + tail_pool + out = torch.empty( + (rows, out_cols), dtype=topk_result.dtype, device=topk_result.device + ) + + if page_table is None: + page_table = topk_result + has_page_table = False + page_table_cols = 1 + else: + assert page_table.ndim == 2 + has_page_table = True + page_table_cols = page_table.shape[1] + + if topk_offsets is None: + topk_offsets = seq_lens + has_topk_offsets = False + else: + if topk_offsets.ndim == 2: + assert topk_offsets.shape[1] == 1 + topk_offsets = topk_offsets.squeeze(1) + assert topk_offsets.ndim == 1 + has_topk_offsets = True + + block_cols = triton.next_power_of_2(out_cols) + _append_kpool_tail_to_topk_kernel[(rows,)]( + topk_result, + seq_lens, + pool_lens, + page_table, + topk_offsets, + out, + topk_result.stride(0), + topk_result.stride(1), + page_table.stride(0), + page_table.stride(1), + out.stride(0), + out.stride(1), + N_COLS=n_cols, + OUT_COLS=out_cols, + PAGE_TABLE_COLS=page_table_cols, + POOL_SIZE=pool_size, + HAS_PAGE_TABLE=has_page_table, + HAS_TOPK_OFFSETS=has_topk_offsets, + BLOCK_COLS=block_cols, + ) + return out + + +@triton.jit +def _append_kpool_tail_to_topk_kernel( + topk_ptr, + seq_lens_ptr, + pool_lens_ptr, + page_table_ptr, + topk_offsets_ptr, + out_ptr, + topk_stride_0, + topk_stride_1, + page_table_stride_0, + page_table_stride_1, + out_stride_0, + out_stride_1, + N_COLS: tl.constexpr, + OUT_COLS: tl.constexpr, + PAGE_TABLE_COLS: tl.constexpr, + POOL_SIZE: tl.constexpr, + HAS_PAGE_TABLE: tl.constexpr, + HAS_TOPK_OFFSETS: tl.constexpr, + BLOCK_COLS: tl.constexpr, +): + row = tl.program_id(0) + cols = tl.arange(0, BLOCK_COLS) + mask = cols < OUT_COLS + + seq_len = tl.load(seq_lens_ptr + row).to(tl.int32) + pool_len = tl.load(pool_lens_ptr + row).to(tl.int32) + tail_start = pool_len * POOL_SIZE + history_len = tl.minimum(tail_start, N_COLS) + tail_count = seq_len % POOL_SIZE + + is_history = cols < history_len + safe_history_cols = tl.minimum(cols, N_COLS - 1) + history_value = tl.load( + topk_ptr + row * topk_stride_0 + safe_history_cols * topk_stride_1, + mask=mask & is_history, + other=-1, + ) + + tail_offset = cols - history_len + is_tail = (tail_offset >= 0) & (tail_offset < tail_count) + tail_raw = tail_start + tail_offset + tail_value = tail_raw + if HAS_PAGE_TABLE: + safe_tail = tl.minimum(tl.maximum(tail_raw, 0), PAGE_TABLE_COLS - 1) + # Promote the row term before multiplication; 1M-context strides + # overflow int32 around row 2048. + tail_value = tl.load( + page_table_ptr + + row.to(tl.int64) * page_table_stride_0 + + safe_tail * page_table_stride_1, + mask=mask & is_tail, + other=-1, + ).to(tl.int32) + if HAS_TOPK_OFFSETS: + offset = tl.load(topk_offsets_ptr + row).to(tl.int32) + tail_value = tail_raw + offset + + value = tl.where(is_history, history_value, -1) + value = tl.where(is_tail, tail_value, value) + tl.store(out_ptr + row * out_stride_0 + cols * out_stride_1, value, mask=mask) + + +def topk_from_pooled_history_logits( + logits: torch.Tensor, + group_lengths: torch.Tensor, + pool_size: int, + topk: int, + page_table: torch.Tensor | None = None, + topk_offsets: torch.Tensor | None = None, + seq_lens: torch.Tensor | None = None, + row_starts: torch.Tensor | None = None, + out_rows: int | None = None, + page_table_row_index: torch.Tensor | None = None, +) -> torch.Tensor: + assert logits.ndim == 2 + assert group_lengths.ndim == 1 + assert logits.shape[0] == group_lengths.shape[0] + assert topk > 0 + assert topk % pool_size == 0 + assert out_rows is None or out_rows >= logits.shape[0] + assert page_table_row_index is None or page_table is not None + + _, cols = logits.shape + group_topk = history_group_budget_for_topk(topk, pool_size) + if topk_offsets is not None and topk_offsets.ndim == 2: + assert topk_offsets.shape[1] == 1 + topk_offsets = topk_offsets.squeeze(1) + + if group_topk not in (128, 160, 192, 224, 256, 512, 2048): + raise NotImplementedError( + "index_kpool topk only supports pooled group_topk in " + f"(128, 160, 192, 224, 256, 512, 2048), got {group_topk} " + f"(topk={topk}, pool_size={pool_size})." + ) + if not logits.is_cuda or logits.dtype != torch.float32: + raise NotImplementedError( + "index_kpool topk requires CUDA float32 logits; PyTorch topk fallback " + f"is disabled. Got device={logits.device}, dtype={logits.dtype}." + ) + + if group_topk in (128, 160, 192, 224, 256, 512): + from sglang.kernels.ops.moe.kpool_topk_transform import ( + fast_kpool_topk_transform_fused, + ) + + result = fast_kpool_topk_transform_fused( + score=logits, + lengths=group_lengths.to(torch.int32), + pool_size=pool_size, + topk=topk, + page_table=page_table, + topk_indices_offset=topk_offsets, + row_starts=row_starts, + seq_lens=seq_lens.to(torch.int32) if seq_lens is not None else None, + page_table_row_index=page_table_row_index, + ) + if out_rows is None or out_rows == result.shape[0]: + return result + padded = torch.full( + (out_rows, result.shape[1]), -1, dtype=result.dtype, device=result.device + ) + padded[: result.shape[0]] = result + return padded + + assert ( + page_table_row_index is None + ), "page_table_row_index requires the fused fast_kpool group_topk path" + + from sgl_kernel import fast_topk_v2 + + selected_groups = fast_topk_v2( + logits, + group_lengths.to(torch.int32), + group_topk, + row_starts=row_starts, + ) + + rank = torch.arange(group_topk, device=logits.device, dtype=torch.int32) + max_valid_groups = min(cols, group_topk) + valid_counts = torch.minimum( + group_lengths.to(torch.int32), + torch.full_like(group_lengths.to(torch.int32), max_valid_groups), + ) + group_valid = rank.unsqueeze(0) < valid_counts.unsqueeze(1) + expanded = expand_pooled_groups_to_topk( + selected_groups.contiguous(), + group_valid, + topk=topk, + pool_size=pool_size, + page_table=page_table, + topk_offsets=topk_offsets, + ) + if seq_lens is None: + result = expanded + else: + result = append_kpool_tail_to_topk( + expanded, + seq_lens=seq_lens, + pool_lens=group_lengths, + pool_size=pool_size, + page_table=page_table, + topk_offsets=topk_offsets, + ) + if out_rows is None or out_rows == result.shape[0]: + return result + padded = torch.full( + (out_rows, result.shape[1]), -1, dtype=result.dtype, device=result.device + ) + padded[: result.shape[0]] = result + return padded + + +def kpool_softmax_rotate_write_cache( + pool, + buf: torch.Tensor, + slot_k: torch.Tensor, + slot_score: torch.Tensor, + ape: torch.Tensor, + loc: torch.Tensor, + write_mask: torch.Tensor | None = None, + round_scale: bool = False, + return_compressed: bool = False, + write_cache: bool = True, +) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: + assert slot_k.ndim == 3 + assert slot_score.shape == slot_k.shape + assert ape.shape == slot_k.shape[1:] + assert slot_k.shape[2] == INDEX_HEAD_DIM + assert slot_k.dtype == torch.bfloat16 + assert slot_score.dtype in KPOOL_SCORE_DTYPES + assert ape.dtype == torch.float32 + assert buf.dtype == torch.uint8 + assert pool.page_size == BLOCK_SIZE_K + assert pool.index_head_dim == INDEX_HEAD_DIM + assert loc.dtype == torch.int64 + assert write_cache or return_compressed + + slot_k = slot_k.contiguous() + slot_score = slot_score.contiguous() + ape = ape.contiguous() + loc = loc.contiguous() + if write_mask is None: + write_mask = torch.empty((1,), dtype=torch.bool, device=slot_k.device) + has_write_mask = False + else: + assert write_mask.shape == (slot_k.shape[0],) + assert not return_compressed + write_mask = write_mask.contiguous() + has_write_mask = True + + if slot_k.shape[0] == 0: + if return_compressed: + return ( + torch.empty( + (0, slot_k.shape[2]), + dtype=torch.float8_e4m3fn, + device=slot_k.device, + ), + torch.empty((0,), dtype=torch.float32, device=slot_k.device), + ) + return None + + buf_fp8 = buf.view(torch.float8_e4m3fn) + buf_fp32 = buf.view(torch.float32) + if return_compressed: + compressed_k = torch.empty( + (slot_k.shape[0], slot_k.shape[2]), + dtype=torch.float8_e4m3fn, + device=slot_k.device, + ) + compressed_scale = torch.empty( + (slot_k.shape[0],), dtype=torch.float32, device=slot_k.device + ) + else: + compressed_k = buf_fp8 + compressed_scale = buf_fp32 + _kpool_softmax_rotate_write_cache_kernel[(slot_k.shape[0],)]( + buf_fp8, + buf_fp32, + slot_k, + slot_score, + ape, + loc, + write_mask, + compressed_k, + compressed_scale, + slot_k.stride(0), + slot_k.stride(1), + slot_score.stride(0), + slot_score.stride(1), + ape.stride(0), + PAGE_SIZE=pool.page_size, + BUF_NUMEL_PER_PAGE=buf.shape[1], + POOL_SIZE=slot_k.shape[1], + HEAD_DIM=slot_k.shape[2], + S_OFFSET_NBYTES_IN_PAGE=pool.page_size * pool.index_head_dim, + ROUND_SCALE=round_scale, + HAS_WRITE_MASK=has_write_mask, + RETURN_COMPRESSED=return_compressed, + WRITE_CACHE=write_cache, + BLOCK_D=triton.next_power_of_2(slot_k.shape[2]), + ) + if return_compressed: + return compressed_k, compressed_scale + return None + + +def kpool_decode_update_and_maybe_write_cache( + pool, + buf: torch.Tensor, + tail_k: torch.Tensor, + tail_score: torch.Tensor, + key: torch.Tensor, + slot_score: torch.Tensor, + ape: torch.Tensor, + block_tables: torch.Tensor, + req_pool_indices: torch.Tensor, + positions: torch.Tensor, + seq_lens: torch.Tensor, + out_cache_loc: torch.Tensor, + round_scale: bool = False, +) -> None: + assert tail_k.ndim == 3 + assert tail_score.shape == tail_k.shape + assert tail_k.shape[1] == pool.index_kpool + pool.tail_extra_slots + assert tail_k.shape[2] == INDEX_HEAD_DIM + assert key.ndim == 2 and key.shape[1] == INDEX_HEAD_DIM + assert slot_score.shape == key.shape + assert ape.shape == (pool.index_kpool, INDEX_HEAD_DIM) + assert tail_k.dtype == torch.bfloat16 + assert key.dtype == torch.bfloat16 + assert tail_score.dtype in KPOOL_SCORE_DTYPES + assert slot_score.dtype == tail_score.dtype + assert ape.dtype == torch.float32 + assert buf.dtype == torch.uint8 + assert pool.page_size == BLOCK_SIZE_K + assert pool.index_head_dim == INDEX_HEAD_DIM + assert tail_k.is_contiguous() + assert tail_score.is_contiguous() + + batch = key.shape[0] + if batch == 0: + return + + key = key.contiguous() + slot_score = slot_score.contiguous() + ape = ape.contiguous() + req_pool_indices = req_pool_indices.contiguous() + positions = positions.contiguous() + seq_lens = seq_lens.contiguous() + out_cache_loc = out_cache_loc.contiguous() + + assert req_pool_indices.shape[0] >= batch + assert positions.shape[0] >= batch + assert seq_lens.shape[0] >= batch + assert out_cache_loc.shape[0] >= batch + assert block_tables.ndim == 2 + assert block_tables.shape[0] >= batch + + buf_fp8 = buf.view(torch.float8_e4m3fn) + buf_fp32 = buf.view(torch.float32) + _kpool_decode_update_and_maybe_write_cache_kernel[(batch,)]( + buf_fp8, + buf_fp32, + tail_k, + tail_score, + key, + slot_score, + ape, + block_tables, + req_pool_indices, + positions, + seq_lens, + out_cache_loc, + tail_k.stride(0), + tail_k.stride(1), + tail_score.stride(0), + tail_score.stride(1), + key.stride(0), + slot_score.stride(0), + ape.stride(0), + block_tables.stride(0), + block_tables.stride(1), + REQ_POOL_SIZE=tail_k.shape[0], + PAGE_SIZE=pool.page_size, + BUF_NUMEL_PER_PAGE=buf.shape[1], + POOL_SIZE=pool.index_kpool, + TAIL_SIZE=tail_k.shape[1], + HEAD_DIM=tail_k.shape[2], + BLOCK_TABLE_COLS=block_tables.shape[1], + S_OFFSET_NBYTES_IN_PAGE=pool.slots_per_page * pool.index_head_dim, + ROUND_SCALE=round_scale, + BLOCK_D=triton.next_power_of_2(tail_k.shape[2]), + SLOTS_PER_PAGE=pool.slots_per_page, + ) + + +@triton.jit +def _hadamard128_stage(x, GROUPS: tl.constexpr, STRIDE: tl.constexpr): + x3 = tl.reshape(x, (GROUPS, 2, STRIDE)) + x3 = tl.trans(x3, 0, 2, 1) + a, b = tl.split(x3) + x3 = tl.join(a + b, a - b) + x3 = tl.trans(x3, 0, 2, 1) + return tl.reshape(x3, (128,)) + + +@triton.jit +def _hadamard128(x): + x = _hadamard128_stage(x, 64, 1) + x = _hadamard128_stage(x, 32, 2) + x = _hadamard128_stage(x, 16, 4) + x = _hadamard128_stage(x, 8, 8) + x = _hadamard128_stage(x, 4, 16) + x = _hadamard128_stage(x, 2, 32) + x = _hadamard128_stage(x, 1, 64) + return x * 0.08838834764831845 + + +@triton.jit +def _kpool_softmax_rotate_write_cache_kernel( + buf_fp8_ptr, + buf_fp32_ptr, + slot_k_ptr, + slot_score_ptr, + ape_ptr, + loc_ptr, + write_mask_ptr, + compressed_k_ptr, + compressed_scale_ptr, + slot_k_stride_0, + slot_k_stride_1, + slot_score_stride_0, + slot_score_stride_1, + ape_stride_0, + PAGE_SIZE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + HAS_WRITE_MASK: tl.constexpr, + RETURN_COMPRESSED: tl.constexpr, + WRITE_CACHE: tl.constexpr, + BLOCK_D: tl.constexpr, +): + row = tl.program_id(0) + do_write = True + if HAS_WRITE_MASK: + do_write = tl.load(write_mask_ptr + row) + + offs = tl.arange(0, BLOCK_D) + mask = (offs < HEAD_DIM) & do_write + + max_score = tl.full((BLOCK_D,), -float("inf"), tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + score = tl.load( + slot_score_ptr + + row * slot_score_stride_0 + + slot * slot_score_stride_1 + + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + score += tl.load(ape_ptr + slot * ape_stride_0 + offs, mask=mask, other=0.0).to( + tl.float32 + ) + max_score = tl.maximum(max_score, score) + + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + score = tl.load( + slot_score_ptr + + row * slot_score_stride_0 + + slot * slot_score_stride_1 + + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + score += tl.load(ape_ptr + slot * ape_stride_0 + offs, mask=mask, other=0.0).to( + tl.float32 + ) + prob = tl.exp(score - max_score) + denom += prob + k = tl.load( + slot_k_ptr + row * slot_k_stride_0 + slot * slot_k_stride_1 + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + acc += k * prob + + x = acc / denom + x = tl.where(do_write, x, 0.0).to(tl.bfloat16).to(tl.float32) + x = _hadamard128(x).to(tl.bfloat16).to(tl.float32) + + fp8_min = -448.0 + fp8_max = 448.0 + fp8_max_inv = 1.0 / fp8_max + absmax = tl.max(tl.abs(x), axis=0) + absmax = tl.maximum(absmax, 1e-4) + + if ROUND_SCALE: + log_val = tl.log2(absmax * fp8_max_inv) + scale = tl.exp2(tl.ceil(log_val)) + else: + scale = absmax * fp8_max_inv + + quantized = x / scale + quantized = tl.minimum(tl.maximum(quantized, fp8_min), fp8_max) + + if WRITE_CACHE: + loc = tl.load(loc_ptr + row, mask=do_write, other=0) + loc_page_index = loc // PAGE_SIZE + loc_token_offset_in_page = loc % PAGE_SIZE + out_k_offsets = ( + loc_page_index * BUF_NUMEL_PER_PAGE + + loc_token_offset_in_page * HEAD_DIM + + offs + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=mask) + tl.store(buf_fp32_ptr + out_s_offset, scale, mask=do_write) + if RETURN_COMPRESSED: + tl.store( + compressed_k_ptr + row * HEAD_DIM + offs, + quantized, + mask=offs < HEAD_DIM, + ) + tl.store(compressed_scale_ptr + row, scale) + + +@triton.jit +def _kpool_decode_update_and_maybe_write_cache_kernel( + buf_fp8_ptr, + buf_fp32_ptr, + tail_k_ptr, + tail_score_ptr, + key_ptr, + slot_score_ptr, + ape_ptr, + block_tables_ptr, + req_pool_indices_ptr, + positions_ptr, + seq_lens_ptr, + out_cache_loc_ptr, + tail_k_stride_0, + tail_k_stride_1, + tail_score_stride_0, + tail_score_stride_1, + key_stride_0, + slot_score_stride_0, + ape_stride_0, + block_tables_stride_0, + block_tables_stride_1, + REQ_POOL_SIZE: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + TAIL_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_TABLE_COLS: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + BLOCK_D: tl.constexpr, + SLOTS_PER_PAGE: tl.constexpr, +): + row = tl.program_id(0) + offs = tl.arange(0, BLOCK_D) + dim_mask = offs < HEAD_DIM + + req_raw = tl.load(req_pool_indices_ptr + row) + req_valid = (req_raw >= 0) & (req_raw < REQ_POOL_SIZE) + req = tl.minimum(tl.maximum(req_raw, 0), REQ_POOL_SIZE - 1) + + pos = tl.load(positions_ptr + row) + safe_pos = tl.maximum(pos, 0) + seq_len = tl.load(seq_lens_ptr + row) + cache_loc = tl.load(out_cache_loc_ptr + row) + pos_valid = req_valid & (cache_loc != 0) & (pos >= 0) & (pos < seq_len) + + slot = safe_pos % POOL_SIZE + phys_slot = safe_pos % TAIL_SIZE + + key = tl.load( + key_ptr + row * key_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score_current = tl.load( + slot_score_ptr + row * slot_score_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + + if pos_valid & (slot == POOL_SIZE - 1): + pool_logical_start = safe_pos - slot + max_score = tl.full((BLOCK_D,), -float("inf"), tl.float32) + for pool_slot in tl.static_range(0, POOL_SIZE): + is_current = pool_slot == slot + phys = (pool_logical_start + pool_slot) % TAIL_SIZE + score_buf = tl.load( + tail_score_ptr + + req * tail_score_stride_0 + + phys * tail_score_stride_1 + + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score = tl.where(is_current, score_current, score_buf) + score += tl.load( + ape_ptr + pool_slot * ape_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + max_score = tl.maximum(max_score, score) + + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for pool_slot in tl.static_range(0, POOL_SIZE): + is_current = pool_slot == slot + phys = (pool_logical_start + pool_slot) % TAIL_SIZE + score_buf = tl.load( + tail_score_ptr + + req * tail_score_stride_0 + + phys * tail_score_stride_1 + + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score = tl.where(is_current, score_current, score_buf) + score += tl.load( + ape_ptr + pool_slot * ape_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + prob = tl.exp(score - max_score) + denom += prob + k_buf = tl.load( + tail_k_ptr + req * tail_k_stride_0 + phys * tail_k_stride_1 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + k = tl.where(is_current, key, k_buf) + acc += k * prob + + x = (acc / denom).to(tl.bfloat16).to(tl.float32) + x = _hadamard128(x).to(tl.bfloat16).to(tl.float32) + + fp8_min = -448.0 + fp8_max = 448.0 + fp8_max_inv = 1.0 / fp8_max + absmax = tl.max(tl.abs(x), axis=0) + absmax = tl.maximum(absmax, 1e-4) + + if ROUND_SCALE: + log_val = tl.log2(absmax * fp8_max_inv) + scale = tl.exp2(tl.ceil(log_val)) + else: + scale = absmax * fp8_max_inv + + quantized = x / scale + quantized = tl.minimum(tl.maximum(quantized, fp8_min), fp8_max) + + pool_id = safe_pos // POOL_SIZE + pool_page_group = pool_id // SLOTS_PER_PAGE + token_page_row = pool_page_group * POOL_SIZE + token_page_row = tl.minimum(tl.maximum(token_page_row, 0), BLOCK_TABLE_COLS - 1) + packed_page = tl.load( + block_tables_ptr + + row * block_tables_stride_0 + + token_page_row * block_tables_stride_1, + ) + loc_page_index = packed_page.to(tl.int64) + loc_token_offset_in_page = pool_id % SLOTS_PER_PAGE + out_k_offsets = ( + loc_page_index * BUF_NUMEL_PER_PAGE + + loc_token_offset_in_page * HEAD_DIM + + offs + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=dim_mask) + tl.store(buf_fp32_ptr + out_s_offset, scale) + + tail_k_offset = req * tail_k_stride_0 + phys_slot * tail_k_stride_1 + offs + tail_score_offset = ( + req * tail_score_stride_0 + phys_slot * tail_score_stride_1 + offs + ) + update_mask = dim_mask & pos_valid + tl.store(tail_k_ptr + tail_k_offset, key, mask=update_mask) + tl.store(tail_score_ptr + tail_score_offset, score_current, mask=update_mask) + + +@triton.jit +def _hadamard_quantize_fp8(acc, denom, ROUND_SCALE: tl.constexpr): + x = (acc / denom).to(tl.bfloat16).to(tl.float32) + x = _hadamard128(x).to(tl.bfloat16).to(tl.float32) + + fp8_max_inv = 1.0 / 448.0 + absmax = tl.maximum(tl.max(tl.abs(x), axis=0), 1e-4) + if ROUND_SCALE: + scale = tl.exp2(tl.ceil(tl.log2(absmax * fp8_max_inv))) + else: + scale = absmax * fp8_max_inv + + quantized = tl.minimum(tl.maximum(x / scale, -448.0), 448.0) + return quantized, scale + + +@triton.jit +def _kpool_assemble_softmax_rotate_write_cache_kernel( + buf_fp8_ptr, + buf_fp32_ptr, + chunk_k_ptr, + chunk_score_ptr, + tail_k_ptr, + tail_score_ptr, + req_pool_idx_ptr, + n_from_tail_ptr, + chunk_src_start_ptr, + tail_logical_base_ptr, + ape_ptr, + loc_ptr, + write_mask_ptr, + chunk_stride_0, + tail_stride_0, + tail_stride_1, + ape_stride_0, + BUF_NUMEL_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + TAIL_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + HAS_WRITE_MASK: tl.constexpr, + BLOCK_D: tl.constexpr, + SLOTS_PER_PAGE: tl.constexpr, +): + row = tl.program_id(0) + if HAS_WRITE_MASK: + if not tl.load(write_mask_ptr + row): + return + + offs = tl.arange(0, BLOCK_D) + mask = offs < HEAD_DIM + + n_tail = tl.load(n_from_tail_ptr + row) + req = tl.load(req_pool_idx_ptr + row) + chunk_src = tl.load(chunk_src_start_ptr + row) + tail_base = tl.load(tail_logical_base_ptr + row) + + m = tl.full((BLOCK_D,), -float("inf"), tl.float32) + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + if slot < n_tail: + phys = (tail_base + slot) % TAIL_SIZE + off = req * tail_stride_0 + phys * tail_stride_1 + offs + score = tl.load(tail_score_ptr + off, mask=mask, other=0.0).to(tl.float32) + k = tl.load(tail_k_ptr + off, mask=mask, other=0.0).to(tl.float32) + else: + off = (chunk_src + (slot - n_tail)) * chunk_stride_0 + offs + score = tl.load(chunk_score_ptr + off, mask=mask, other=0.0).to(tl.float32) + k = tl.load(chunk_k_ptr + off, mask=mask, other=0.0).to(tl.float32) + + score += tl.load(ape_ptr + slot * ape_stride_0 + offs, mask=mask, other=0.0).to( + tl.float32 + ) + new_m = tl.maximum(m, score) + rescale = tl.exp(m - new_m) + prob = tl.exp(score - new_m) + denom = denom * rescale + prob + acc = acc * rescale + k * prob + m = new_m + + quantized, scale = _hadamard_quantize_fp8(acc, denom, ROUND_SCALE) + + loc = tl.load(loc_ptr + row) + loc_page_index = loc // SLOTS_PER_PAGE + loc_token_offset_in_page = loc % SLOTS_PER_PAGE + out_k_offsets = ( + loc_page_index * BUF_NUMEL_PER_PAGE + loc_token_offset_in_page * HEAD_DIM + offs + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=mask) + tl.store(buf_fp32_ptr + out_s_offset, scale) + + +def kpool_assemble_softmax_rotate_write_cache( + pool, + buf: torch.Tensor, + chunk_k: torch.Tensor, + chunk_score: torch.Tensor, + tail_k: torch.Tensor, + tail_score: torch.Tensor, + req_pool_idx: torch.Tensor, + n_from_tail: torch.Tensor, + chunk_src_start: torch.Tensor, + tail_logical_base: torch.Tensor, + ape: torch.Tensor, + loc: torch.Tensor, + write_mask: torch.Tensor | None = None, + round_scale: bool = False, +) -> None: + pool_size = pool.index_kpool + n_pools = req_pool_idx.shape[0] + if n_pools == 0: + return + + chunk_k = chunk_k.contiguous() + chunk_score = chunk_score.contiguous() + ape = ape.contiguous() + loc = loc.contiguous() + if write_mask is None: + write_mask = torch.empty((1,), dtype=torch.bool, device=chunk_k.device) + has_write_mask = False + else: + write_mask = write_mask.contiguous() + has_write_mask = True + + buf_fp8 = buf.view(torch.float8_e4m3fn) + buf_fp32 = buf.view(torch.float32) + slots_per_page = pool.slots_per_page + + _kpool_assemble_softmax_rotate_write_cache_kernel[(n_pools,)]( + buf_fp8, + buf_fp32, + chunk_k, + chunk_score, + tail_k, + tail_score, + req_pool_idx, + n_from_tail, + chunk_src_start, + tail_logical_base, + ape, + loc, + write_mask, + chunk_k.stride(0), + tail_k.stride(0), + tail_k.stride(1), + ape.stride(0), + BUF_NUMEL_PER_PAGE=buf.shape[1], + POOL_SIZE=pool_size, + TAIL_SIZE=tail_k.shape[1], + HEAD_DIM=INDEX_HEAD_DIM, + S_OFFSET_NBYTES_IN_PAGE=slots_per_page * INDEX_HEAD_DIM, + ROUND_SCALE=round_scale, + HAS_WRITE_MASK=has_write_mask, + BLOCK_D=triton.next_power_of_2(INDEX_HEAD_DIM), + SLOTS_PER_PAGE=slots_per_page, + ) + + +def scatter_kpool_tail_updates( + pool, + chunk_k: torch.Tensor, + chunk_score: torch.Tensor, + tail_k: torch.Tensor, + tail_score: torch.Tensor, + req_pool_idx: torch.Tensor, + dst_logical_start: torch.Tensor, + chunk_src_start: torch.Tensor, + n_write: torch.Tensor, +) -> None: + pool_size = pool.index_kpool + n_rows = req_pool_idx.shape[0] + if n_rows == 0: + return + + chunk_k = chunk_k.contiguous() + chunk_score = chunk_score.contiguous() + _scatter_kpool_tail_updates_kernel[(n_rows, pool_size)]( + chunk_k, + chunk_score, + tail_k, + tail_score, + req_pool_idx, + dst_logical_start, + chunk_src_start, + n_write, + chunk_k.stride(0), + tail_k.stride(0), + tail_k.stride(1), + POOL_SIZE=pool_size, + TAIL_SIZE=tail_k.shape[1], + HEAD_DIM=INDEX_HEAD_DIM, + BLOCK_D=triton.next_power_of_2(INDEX_HEAD_DIM), + ) + + +@triton.jit +def _scatter_kpool_tail_updates_kernel( + chunk_k_ptr, + chunk_score_ptr, + tail_k_ptr, + tail_score_ptr, + req_pool_idx_ptr, + dst_logical_start_ptr, + chunk_src_start_ptr, + n_write_ptr, + chunk_stride_0, + tail_stride_0, + tail_stride_1, + POOL_SIZE: tl.constexpr, + TAIL_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_D: tl.constexpr, +): + row = tl.program_id(0) + slot = tl.program_id(1) + + n_w = tl.load(n_write_ptr + row) + if slot >= n_w: + return + + req = tl.load(req_pool_idx_ptr + row) + dst_logical_start = tl.load(dst_logical_start_ptr + row) + src_off = tl.load(chunk_src_start_ptr + row) + slot + + offs = tl.arange(0, BLOCK_D) + mask = offs < HEAD_DIM + k = tl.load(chunk_k_ptr + src_off * chunk_stride_0 + offs, mask=mask) + s = tl.load(chunk_score_ptr + src_off * chunk_stride_0 + offs, mask=mask) + + dst = ( + req * tail_stride_0 + + ((dst_logical_start + slot) % TAIL_SIZE) * tail_stride_1 + + offs + ) + tl.store(tail_k_ptr + dst, k, mask=mask) + tl.store(tail_score_ptr + dst, s, mask=mask) + + +@triton.jit +def _pack_pool_slots_to_payload_kernel( + buf_ptr, + locs_ptr, + payload_ptr, + payload_bytes: tl.constexpr, + slots_per_page: tl.constexpr, + head_dim: tl.constexpr, + page_bytes: tl.constexpr, + scale_region_off: tl.constexpr, + BLOCK_D: tl.constexpr, +): + row = tl.program_id(0) + loc = tl.load(locs_ptr + row).to(tl.int64) + page = loc // slots_per_page + slot = loc % slots_per_page + page_base = page * page_bytes + + offs = tl.arange(0, BLOCK_D) + mask = offs < head_dim + src = page_base + slot * head_dim + offs + val = tl.load(buf_ptr + src, mask=mask, other=0).to(tl.uint8) + tl.store(payload_ptr + row * payload_bytes + offs, val, mask=mask) + + s_offs = tl.arange(0, 4) + s_src = page_base + scale_region_off + slot * 4 + s_offs + s_val = tl.load(buf_ptr + s_src).to(tl.uint8) + tl.store(payload_ptr + row * payload_bytes + head_dim + s_offs, s_val) + + +@triton.jit +def _select_and_scatter_pool_slots_kernel( + recv_ptr, + owner_ptr, + locs_ptr, + buf_ptr, + cp_rank: tl.constexpr, + payload_bytes: tl.constexpr, + slots_per_page: tl.constexpr, + head_dim: tl.constexpr, + page_bytes: tl.constexpr, + scale_region_off: tl.constexpr, + n_total: tl.constexpr, + BLOCK_D: tl.constexpr, +): + row = tl.program_id(0) + owner = tl.load(owner_ptr + row).to(tl.int64) + if owner == cp_rank: + return + + loc = tl.load(locs_ptr + row).to(tl.int64) + page = loc // slots_per_page + slot = loc % slots_per_page + page_base = page * page_bytes + recv_row_base = (owner * n_total + row) * payload_bytes + + offs = tl.arange(0, BLOCK_D) + mask = offs < head_dim + val = tl.load(recv_ptr + recv_row_base + offs, mask=mask, other=0).to(tl.uint8) + dst = page_base + slot * head_dim + offs + tl.store(buf_ptr + dst, val, mask=mask) + + s_offs = tl.arange(0, 4) + s_val = tl.load(recv_ptr + recv_row_base + head_dim + s_offs).to(tl.uint8) + s_dst = page_base + scale_region_off + slot * 4 + s_offs + tl.store(buf_ptr + s_dst, s_val) + + +def all_gather_and_scatter_pool_slots( + buf: torch.Tensor, + local_locs: torch.Tensor, + owner_rank: torch.Tensor, + cp_size: int, + cp_rank: int, + slots_per_page: int, +) -> None: + from sglang.srt.layers.dp_attention import attn_cp_all_gather_into_tensor + + assert buf.is_contiguous() + n_total = local_locs.shape[0] + if n_total == 0 or cp_size <= 1: + return + + head_dim = INDEX_HEAD_DIM + payload_bytes = head_dim + 4 + scale_region_off = slots_per_page * head_dim + page_bytes = buf.shape[1] + device = buf.device + + send_payload = torch.empty( + (n_total, payload_bytes), dtype=torch.uint8, device=device + ) + _pack_pool_slots_to_payload_kernel[(n_total,)]( + buf, + local_locs, + send_payload, + payload_bytes=payload_bytes, + slots_per_page=slots_per_page, + head_dim=head_dim, + page_bytes=page_bytes, + scale_region_off=scale_region_off, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + recv = torch.empty( + (cp_size, n_total, payload_bytes), dtype=torch.uint8, device=device + ) + attn_cp_all_gather_into_tensor( + recv.view(cp_size * n_total, payload_bytes), send_payload + ) + + _select_and_scatter_pool_slots_kernel[(n_total,)]( + recv, + owner_rank, + local_locs, + buf, + cp_rank=cp_rank, + payload_bytes=payload_bytes, + slots_per_page=slots_per_page, + head_dim=head_dim, + page_bytes=page_bytes, + scale_region_off=scale_region_off, + n_total=n_total, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + +@triton.jit +def _kpool_write_tail_and_maybe_compress_kernel( + key_ptr, + score_ptr, + tail_k_ptr, + tail_score_ptr, + ape_ptr, + req_pool_indices_ptr, + write_start_ptr, + tail_logical_start_ptr, + write_loc_ptr, + out_cache_loc_ptr, + effective_n_ptr, + buf_fp8_ptr, + buf_fp32_ptr, + key_stride_0, + score_stride_0, + tail_stride_0, + tail_stride_1, + ape_stride_0, + write_loc_stride_0, + N: tl.constexpr, + POOL_SIZE: tl.constexpr, + TAIL_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_D: tl.constexpr, + SLOTS_PER_PAGE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + HAS_EFFECTIVE_N: tl.constexpr, + MAX_CLOSED_POOLS: tl.constexpr, +): + b = tl.program_id(0) + cache_loc_0 = tl.load(out_cache_loc_ptr + b * N) + if cache_loc_0 == 0: + return + + req = tl.load(req_pool_indices_ptr + b) + write_start = tl.load(write_start_ptr + b) + offs = tl.arange(0, BLOCK_D) + dim_mask = offs < HEAD_DIM + + for i_n in tl.static_range(0, N): + row = b * N + i_n + k = tl.load(key_ptr + row * key_stride_0 + offs, mask=dim_mask) + s = tl.load(score_ptr + row * score_stride_0 + offs, mask=dim_mask) + phys = (write_start + i_n) % TAIL_SIZE + dst = req * tail_stride_0 + phys * tail_stride_1 + offs + tl.store(tail_k_ptr + dst, k, mask=dim_mask) + tl.store(tail_score_ptr + dst, s, mask=dim_mask) + + if HAS_EFFECTIVE_N: + gate_n = tl.load(effective_n_ptr + b).to(tl.int32) + else: + gate_n = N + base_pool = write_start // POOL_SIZE + n_pool = (write_start + gate_n) // POOL_SIZE - base_pool + if n_pool == 0: + return + + base0 = tl.load(tail_logical_start_ptr + b) + for p in tl.static_range(0, MAX_CLOSED_POOLS): + if p < n_pool: + base = base0 + p * POOL_SIZE + m = tl.full((BLOCK_D,), -float("inf"), tl.float32) + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + phys = (base + slot) % TAIL_SIZE + off = req * tail_stride_0 + phys * tail_stride_1 + offs + score = tl.load(tail_score_ptr + off, mask=dim_mask, other=0.0).to( + tl.float32 + ) + k_ld = tl.load(tail_k_ptr + off, mask=dim_mask, other=0.0).to( + tl.float32 + ) + score += tl.load( + ape_ptr + slot * ape_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + new_m = tl.maximum(m, score) + rescale = tl.exp(m - new_m) + prob = tl.exp(score - new_m) + denom = denom * rescale + prob + acc = acc * rescale + k_ld * prob + m = new_m + + quantized, scale = _hadamard_quantize_fp8(acc, denom, ROUND_SCALE) + loc = tl.load(write_loc_ptr + b * write_loc_stride_0 + p) + loc_page_index = loc // SLOTS_PER_PAGE + loc_token_offset_in_page = loc % SLOTS_PER_PAGE + out_k_offsets = ( + loc_page_index * BUF_NUMEL_PER_PAGE + + loc_token_offset_in_page * HEAD_DIM + + offs + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=dim_mask) + tl.store(buf_fp32_ptr + out_s_offset, scale) + + +def kpool_write_tail_and_maybe_compress( + pool, + buf: torch.Tensor, + key: torch.Tensor, + score: torch.Tensor, + tail_k: torch.Tensor, + tail_score: torch.Tensor, + ape: torch.Tensor, + req_pool_indices: torch.Tensor, + write_start: torch.Tensor, + tail_logical_start: torch.Tensor, + write_loc: torch.Tensor, + out_cache_loc: torch.Tensor, + num_draft_tokens: int, + round_scale: bool, + effective_n_per_batch: Optional[torch.Tensor] = None, +) -> None: + assert num_draft_tokens > 0 + assert key.dim() == 2 and key.shape[1] == INDEX_HEAD_DIM + assert score.shape == key.shape + assert tail_k.shape == tail_score.shape + assert tail_k.shape[1] == pool.index_kpool + pool.tail_extra_slots + assert tail_k.shape[2] == INDEX_HEAD_DIM + assert key.dtype == torch.bfloat16 + assert score.dtype in KPOOL_SCORE_DTYPES + assert tail_k.dtype == torch.bfloat16 + assert tail_score.dtype in KPOOL_SCORE_DTYPES + assert ape.dtype == torch.float32 + + bn = key.shape[0] + if bn == 0: + return + assert bn % num_draft_tokens == 0 + bs = bn // num_draft_tokens + max_closed_pools = kpool_max_closed_pools(num_draft_tokens, pool.index_kpool) + assert write_loc.shape == (bs, max_closed_pools), write_loc.shape + assert write_loc.stride(1) == 1, write_loc.stride() + + key = key.contiguous() + score = score.contiguous() + ape = ape.contiguous() + req_pool_indices = req_pool_indices.contiguous() + write_start = write_start.contiguous() + tail_logical_start = tail_logical_start.contiguous() + write_loc = write_loc.contiguous() + out_cache_loc = out_cache_loc.contiguous() + if effective_n_per_batch is not None: + effective_n_per_batch = effective_n_per_batch.contiguous() + + slots_per_page = pool.slots_per_page + buf_fp8 = buf.view(torch.float8_e4m3fn) + buf_fp32 = buf.view(torch.float32) + _kpool_write_tail_and_maybe_compress_kernel[(bs,)]( + key, + score, + tail_k, + tail_score, + ape, + req_pool_indices, + write_start, + tail_logical_start, + write_loc, + out_cache_loc, + effective_n_per_batch, + buf_fp8, + buf_fp32, + key.stride(0), + score.stride(0), + tail_k.stride(0), + tail_k.stride(1), + ape.stride(0), + write_loc.stride(0), + N=num_draft_tokens, + POOL_SIZE=pool.index_kpool, + TAIL_SIZE=tail_k.shape[1], + HEAD_DIM=INDEX_HEAD_DIM, + BLOCK_D=triton.next_power_of_2(INDEX_HEAD_DIM), + SLOTS_PER_PAGE=slots_per_page, + BUF_NUMEL_PER_PAGE=buf.shape[1], + S_OFFSET_NBYTES_IN_PAGE=slots_per_page * pool.index_head_dim, + ROUND_SCALE=round_scale, + HAS_EFFECTIVE_N=effective_n_per_batch is not None, + MAX_CLOSED_POOLS=max_closed_pools, + ) diff --git a/python/sglang/srt/layers/attention/dsa/kpool_plan.py b/python/sglang/srt/layers/attention/dsa/kpool_plan.py new file mode 100644 index 000000000..ea2d3a3fb --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/kpool_plan.py @@ -0,0 +1,852 @@ +from __future__ import annotations + +import dataclasses +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, List, NamedTuple, Optional + +import torch + +from sglang.srt.environ import envs +from sglang.srt.layers.attention.dsa.kpool_fp8_index import ( + INDEX_HEAD_DIM, + build_pooled_page_table_64, + kpool_build_ragged_layout, + kpool_max_closed_pools, + update_kpool_write_plan_cuda_graph, +) +from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp +from sglang.srt.model_executor.forward_context import get_req_to_token_pool +from sglang.srt.runtime_context import get_parallel +from sglang.srt.utils import is_cuda + +if TYPE_CHECKING: + from sglang.srt.layers.attention.dsa.dsa_topk_backend import TopkTransformMethod + from sglang.srt.layers.attention.dsa_backend import DSAMetadata + from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode + + +_RAGGED_SCRATCH_K_U8: Optional[torch.Tensor] = None +_RAGGED_SCRATCH_K_SCALE: Optional[torch.Tensor] = None + + +def _get_ragged_scratch( + total_k_rows: int, device: torch.device +) -> tuple[torch.Tensor, torch.Tensor]: + global _RAGGED_SCRATCH_K_U8, _RAGGED_SCRATCH_K_SCALE + cur = _RAGGED_SCRATCH_K_U8 + grow = ( + cur is None + or cur.device.type != device.type + or (device.index is not None and cur.device.index != device.index) + or cur.shape[0] < total_k_rows + ) + if grow: + _RAGGED_SCRATCH_K_U8 = torch.empty( + (total_k_rows, INDEX_HEAD_DIM), dtype=torch.uint8, device=device + ) + _RAGGED_SCRATCH_K_SCALE = torch.empty( + (total_k_rows,), dtype=torch.float32, device=device + ) + return _RAGGED_SCRATCH_K_U8[:total_k_rows], _RAGGED_SCRATCH_K_SCALE[:total_k_rows] + + +@dataclass(frozen=True) +class PoolWriteRows: + req: torch.Tensor + pool_id: torch.Tensor + n_from_tail: torch.Tensor + chunk_src: torch.Tensor + tail_logical_base: torch.Tensor + write_loc: torch.Tensor + + @property + def is_empty(self) -> bool: + return self.pool_id.shape[0] == 0 + + +@dataclass(frozen=True) +class TailWriteRows: + req: torch.Tensor + dst_logical_start: torch.Tensor + chunk_src: torch.Tensor + n_write: torch.Tensor + + @property + def is_empty(self) -> bool: + return self.req.shape[0] == 0 + + +@dataclass(frozen=True) +class KPoolCpInfo: + size: int + rank: int + owner_rank: torch.Tensor + local_write_mask: torch.Tensor + + +@dataclass(frozen=True) +class KPoolExtendPlan: + writes: PoolWriteRows + tails: TailWriteRows + pooled_seq_lens_expanded: torch.Tensor + seq_lens_expanded: torch.Tensor + ragged_concat_page_table: torch.Tensor + ragged_q_ks: torch.Tensor + ragged_q_ke: torch.Tensor + ragged_total_k_rows: int + ragged_k_u8: Optional[torch.Tensor] + ragged_k_scale: Optional[torch.Tensor] + ragged_paged_page_table: Optional[torch.Tensor] + ragged_paged_page_table_row_index: Optional[torch.Tensor] + cp: Optional[KPoolCpInfo] = None + + +@dataclass(frozen=True) +class KPoolWritePlan: + """``write_loc[b, p]`` is the compression destination for candidate closed + pool ``base_pool[b] + p``; the kernel decides which candidates closed.""" + + req: torch.Tensor + write_start: torch.Tensor + tail_logical_start: torch.Tensor + write_loc: torch.Tensor # int64 [B, max_closed_pools] + num_draft_tokens: int + pool_seqlens_per_q: Optional[torch.Tensor] = None + seqlens_per_q: Optional[torch.Tensor] = None + pool_schedule_metadata: Optional[torch.Tensor] = None + effective_n_per_batch: Optional[torch.Tensor] = None + + +def _is_kpool_layout_enabled(pool_size: int, real_page_size: int) -> bool: + return pool_size > 1 and real_page_size == 64 and real_page_size % pool_size == 0 + + +@dataclass +class _KPoolCpuPlan: + pool_batch_idx: List[int] = field(default_factory=list) + pool_req: List[int] = field(default_factory=list) + pool_pool_id: List[int] = field(default_factory=list) + pool_n_from_tail: List[int] = field(default_factory=list) + pool_chunk_src: List[int] = field(default_factory=list) + pool_tail_logical_base: List[int] = field(default_factory=list) + + tail_req: List[int] = field(default_factory=list) + tail_dst_logical_start: List[int] = field(default_factory=list) + tail_chunk_src: List[int] = field(default_factory=list) + tail_n_write: List[int] = field(default_factory=list) + + ragged_q_len: List[int] = field(default_factory=list) + ragged_pool_pages: List[int] = field(default_factory=list) + cu_pages_excl: List[int] = field(default_factory=list) + cu_q_len_excl: List[int] = field(default_factory=list) + total_pool_pages: int = 0 + + +class _KPoolDecompose(NamedTuple): + first_slot: int + base_pool: int + n_pool: int + tail_n_write: int + + +def _decompose_compress(start: int, length: int, pool_size: int) -> _KPoolDecompose: + first_slot = start % pool_size + base_pool = start // pool_size + n_pool = (start + length) // pool_size - base_pool + consumed = max(0, n_pool * pool_size - first_slot) + tail_n = length - consumed + return _KPoolDecompose( + first_slot=first_slot, + base_pool=base_pool, + n_pool=n_pool, + tail_n_write=tail_n, + ) + + +def _append_compress_rows( + plan: _KPoolCpuPlan, + pool_size: int, + batch_size: int, + extend_seq_lens_cpu: List[int], + seq_lens_cpu: List[int], + req_pool_indices_cpu: List[int], +) -> None: + q_offset = 0 + for i in range(batch_size): + q_len = extend_seq_lens_cpu[i] + assert q_len > 0, f"extend_seq_lens_cpu[{i}] = {q_len}; expected > 0" + + seq_len = seq_lens_cpu[i] + req = req_pool_indices_cpu[i] + d = _decompose_compress(seq_len - q_len, q_len, pool_size) + + if d.n_pool > 0: + plan.pool_batch_idx.extend([i] * d.n_pool) + plan.pool_req.extend([req] * d.n_pool) + plan.pool_pool_id.extend(range(d.base_pool, d.base_pool + d.n_pool)) + plan.pool_n_from_tail.append(d.first_slot) + plan.pool_n_from_tail.extend([0] * (d.n_pool - 1)) + bulk_start = q_offset + pool_size - d.first_slot + plan.pool_chunk_src.append(q_offset) + plan.pool_chunk_src.extend( + range(bulk_start, bulk_start + (d.n_pool - 1) * pool_size, pool_size) + ) + plan.pool_tail_logical_base.extend( + range( + d.base_pool * pool_size, + (d.base_pool + d.n_pool) * pool_size, + pool_size, + ) + ) + + if d.tail_n_write > 0: + consumed = q_len - d.tail_n_write + plan.tail_req.append(req) + plan.tail_dst_logical_start.append(seq_len - q_len + consumed) + plan.tail_chunk_src.append(q_offset + consumed) + plan.tail_n_write.append(d.tail_n_write) + + q_offset += q_len + + +def _append_local_rows( + plan: _KPoolCpuPlan, + pool_size: int, + slots_per_page: int, + local_extend_seq_lens_cpu: List[int], + local_seq_lens_cpu: List[int], +) -> None: + q_offset = 0 + for q_len, seq_len in zip( + local_extend_seq_lens_cpu, local_seq_lens_cpu, strict=True + ): + assert q_len > 0, f"local_extend_seq_lens_cpu has non-positive {q_len = }" + plan.ragged_q_len.append(q_len) + pool_seq_len = seq_len // pool_size + pool_pages_i = (pool_seq_len + slots_per_page - 1) // slots_per_page + plan.ragged_pool_pages.append(pool_pages_i) + plan.cu_pages_excl.append(plan.total_pool_pages) + plan.cu_q_len_excl.append(q_offset) + plan.total_pool_pages += pool_pages_i + q_offset += q_len + + +def _kpool_cpu_plan( + forward_batch: ForwardBatch, + pool_size: int, + slots_per_page: int, + *, + local_extend_seq_lens_cpu: Optional[List[int]] = None, + local_seq_lens_cpu: Optional[List[int]] = None, +) -> _KPoolCpuPlan: + plan = _KPoolCpuPlan() + + extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu + if isinstance(extend_seq_lens_cpu, torch.Tensor): + extend_seq_lens_cpu = extend_seq_lens_cpu.tolist() + seq_lens_cpu = forward_batch.seq_lens_cpu.tolist() + req_pool_indices_cpu = forward_batch.req_pool_indices.tolist() + + _append_compress_rows( + plan, + pool_size, + forward_batch.batch_size, + extend_seq_lens_cpu, + seq_lens_cpu, + req_pool_indices_cpu, + ) + + if local_extend_seq_lens_cpu is None: + local_extend_seq_lens_cpu = extend_seq_lens_cpu + local_seq_lens_cpu = seq_lens_cpu + + _append_local_rows( + plan, + pool_size, + slots_per_page, + local_extend_seq_lens_cpu, + local_seq_lens_cpu, + ) + return plan + + +def _kpool_plan_to_gpu( + cpu: _KPoolCpuPlan, + forward_batch: ForwardBatch, + full_real_page_table: torch.Tensor, + local_real_page_table: torch.Tensor, + local_seqlens_expanded: torch.Tensor, + local_req_pool_indices: torch.Tensor, + pool_size: int, + slots_per_page: int, + topk_transform_method: TopkTransformMethod, +) -> KPoolExtendPlan: + from sglang.srt.layers.attention.dsa.dsa_topk_backend import TopkTransformMethod + + device = forward_batch.seq_lens.device + n_pool = len(cpu.pool_pool_id) + n_tail = len(cpu.tail_req) + n_rag = len(cpu.ragged_q_len) + + total_pool_pages = cpu.total_pool_pages + ragged_total_k_rows = total_pool_pages * slots_per_page + + need_paged = ( + topk_transform_method == TopkTransformMethod.PAGED + and envs.SGLANG_DSA_FUSE_TOPK.get() + and n_rag > 0 + ) + + i64_total = 4 * n_pool + 2 * n_tail + if i64_total > 0: + i64_cpu = torch.tensor( + cpu.pool_req + + cpu.pool_pool_id + + cpu.pool_chunk_src + + cpu.pool_batch_idx + + cpu.tail_req + + cpu.tail_chunk_src, + dtype=torch.int64, + pin_memory=True, + ) + i64_gpu = i64_cpu.to(device, non_blocking=True) + c = 0 + pool_req_t = i64_gpu[c : c + n_pool] + c += n_pool + pool_pool_id_t = i64_gpu[c : c + n_pool] + c += n_pool + pool_chunk_src_t = i64_gpu[c : c + n_pool] + c += n_pool + pool_batch_idx_t = i64_gpu[c : c + n_pool] + c += n_pool + tail_req_t = i64_gpu[c : c + n_tail] + c += n_tail + tail_chunk_src_t = i64_gpu[c : c + n_tail] + else: + empty_i64 = torch.empty((0,), dtype=torch.int64, device=device) + pool_req_t = pool_pool_id_t = pool_chunk_src_t = pool_batch_idx_t = empty_i64 + tail_req_t = tail_chunk_src_t = empty_i64 + + i32_total = 2 * n_pool + 2 * n_tail + 4 * n_rag + if i32_total > 0: + i32_cpu = torch.tensor( + cpu.pool_n_from_tail + + cpu.pool_tail_logical_base + + cpu.tail_dst_logical_start + + cpu.tail_n_write + + cpu.ragged_pool_pages + + cpu.ragged_q_len + + cpu.cu_pages_excl + + cpu.cu_q_len_excl, + dtype=torch.int32, + pin_memory=True, + ) + i32_gpu = i32_cpu.to(device, non_blocking=True) + c = 0 + pool_n_from_tail_t = i32_gpu[c : c + n_pool] + c += n_pool + pool_tail_logical_base_t = i32_gpu[c : c + n_pool] + c += n_pool + tail_dst_logical_start_t = i32_gpu[c : c + n_tail] + c += n_tail + tail_n_write_t = i32_gpu[c : c + n_tail] + c += n_tail + ragged_pool_pages_t = i32_gpu[c : c + n_rag] + c += n_rag + ragged_q_len_t = i32_gpu[c : c + n_rag] + c += n_rag + cu_pages_excl_t = i32_gpu[c : c + n_rag] + c += n_rag + cu_q_len_excl_t = i32_gpu[c : c + n_rag] + else: + empty_i32 = torch.empty((0,), dtype=torch.int32, device=device) + pool_n_from_tail_t = pool_tail_logical_base_t = empty_i32 + tail_dst_logical_start_t = tail_n_write_t = empty_i32 + ragged_pool_pages_t = ragged_q_len_t = empty_i32 + cu_pages_excl_t = cu_q_len_excl_t = empty_i32 + + if n_pool > 0: + pool_page_group = torch.div( + pool_pool_id_t, slots_per_page, rounding_mode="floor" + ) + token_page_row = pool_page_group * pool_size + packed_page = full_real_page_table[pool_batch_idx_t, token_page_row].to( + torch.int64 + ) + pool_write_locs = packed_page * slots_per_page + torch.remainder( + pool_pool_id_t, slots_per_page + ) + else: + pool_write_locs = torch.empty((0,), dtype=torch.int64, device=device) + + pooled_seq_lens_expanded = torch.div( + local_seqlens_expanded, pool_size, rounding_mode="floor" + ).to(torch.int32) + + if n_rag > 0: + ( + ragged_concat_page_table, + ragged_q_ks, + ragged_q_ke, + ) = kpool_build_ragged_layout( + full_page_table=local_real_page_table, + cu_pages_excl=cu_pages_excl_t, + ragged_pool_pages=ragged_pool_pages_t, + cu_q_len_excl=cu_q_len_excl_t, + ragged_q_len=ragged_q_len_t, + pooled_seq_lens_expanded=pooled_seq_lens_expanded, + slots_per_page=slots_per_page, + total_pool_pages=total_pool_pages, + total_q=pooled_seq_lens_expanded.shape[0], + pool_size=pool_size, + ) + else: + empty_i32_dev = torch.empty((0,), dtype=torch.int32, device=device) + ragged_concat_page_table = empty_i32_dev + ragged_q_ks = empty_i32_dev + ragged_q_ke = empty_i32_dev + + ragged_paged_page_table = None + ragged_paged_page_table_row_index = None + if need_paged: + req_to_token = get_req_to_token_pool().req_to_token + ragged_paged_page_table_row_index = torch.repeat_interleave( + local_req_pool_indices.to(torch.int32), ragged_q_len_t + ) + ragged_paged_page_table = req_to_token + + if ragged_total_k_rows > 0: + ragged_k_u8, ragged_k_scale = _get_ragged_scratch(ragged_total_k_rows, device) + else: + ragged_k_u8 = None + ragged_k_scale = None + + return KPoolExtendPlan( + writes=PoolWriteRows( + req=pool_req_t, + pool_id=pool_pool_id_t, + n_from_tail=pool_n_from_tail_t, + chunk_src=pool_chunk_src_t, + tail_logical_base=pool_tail_logical_base_t, + write_loc=pool_write_locs, + ), + tails=TailWriteRows( + req=tail_req_t, + dst_logical_start=tail_dst_logical_start_t, + chunk_src=tail_chunk_src_t, + n_write=tail_n_write_t, + ), + pooled_seq_lens_expanded=pooled_seq_lens_expanded, + seq_lens_expanded=local_seqlens_expanded, + ragged_concat_page_table=ragged_concat_page_table, + ragged_q_ks=ragged_q_ks, + ragged_q_ke=ragged_q_ke, + ragged_total_k_rows=ragged_total_k_rows, + ragged_k_u8=ragged_k_u8, + ragged_k_scale=ragged_k_scale, + ragged_paged_page_table=ragged_paged_page_table, + ragged_paged_page_table_row_index=ragged_paged_page_table_row_index, + cp=_kpool_cp_owner_rank(forward_batch, n_pool, device), + ) + + +def _kpool_cp_owner_rank( + forward_batch: ForwardBatch, + n_pool: int, + device: torch.device, +) -> Optional[KPoolCpInfo]: + if not dsa_use_prefill_cp(forward_batch): + return None + + cp_size = get_parallel().attn_cp_size + if cp_size <= 1: + return None + + cp_rank = get_parallel().attn_cp_rank + if n_pool > 0: + owner = torch.arange(n_pool, dtype=torch.int32, device=device) % cp_size + local_write_mask = owner == cp_rank + else: + owner = torch.empty((0,), dtype=torch.int32, device=device) + local_write_mask = torch.empty((0,), dtype=torch.bool, device=device) + return KPoolCpInfo( + size=cp_size, + rank=cp_rank, + owner_rank=owner, + local_write_mask=local_write_mask, + ) + + +def init_kpool_extend_metadata( + metadata: DSAMetadata, + forward_batch: ForwardBatch, + *, + pool_size: int, + real_page_size: int, + slots_per_page: int, + topk_transform_method: TopkTransformMethod, + full_real_page_table: torch.Tensor, + full_seqlens_expanded: torch.Tensor, + local_real_page_table: Optional[torch.Tensor] = None, + local_seqlens_expanded: Optional[torch.Tensor] = None, + local_extend_seq_lens_cpu: Optional[List[int]] = None, + local_seq_lens_cpu: Optional[List[int]] = None, + local_req_pool_indices: Optional[torch.Tensor] = None, +) -> DSAMetadata: + mode = forward_batch.forward_mode + is_extend_like = mode.is_extend_without_speculative() or mode.is_draft_extend_v2() + if ( + not _is_kpool_layout_enabled(pool_size, real_page_size) + or not is_extend_like + or forward_batch.extend_seq_lens_cpu is None + or forward_batch.seq_lens_cpu is None + ): + return metadata + + if local_real_page_table is None: + local_real_page_table = full_real_page_table + if local_seqlens_expanded is None: + local_seqlens_expanded = full_seqlens_expanded + if local_req_pool_indices is None: + local_req_pool_indices = forward_batch.req_pool_indices + + cpu = _kpool_cpu_plan( + forward_batch, + pool_size, + slots_per_page, + local_extend_seq_lens_cpu=local_extend_seq_lens_cpu, + local_seq_lens_cpu=local_seq_lens_cpu, + ) + plan = _kpool_plan_to_gpu( + cpu, + forward_batch, + full_real_page_table, + local_real_page_table, + local_seqlens_expanded, + local_req_pool_indices, + pool_size, + slots_per_page, + topk_transform_method, + ) + return dataclasses.replace(metadata, kpool_extend_plan=plan) + + +_DEEP_GEMM_MODULE = None +_DEEP_GEMM_IMPORT_FAILED = False + + +def _get_deep_gemm(): + # Cache import failure too, because this helper runs on every graph-replay + # metadata refresh. + global _DEEP_GEMM_MODULE, _DEEP_GEMM_IMPORT_FAILED + if _DEEP_GEMM_MODULE is None and not _DEEP_GEMM_IMPORT_FAILED: + try: + import deep_gemm + except (ImportError, ModuleNotFoundError): + _DEEP_GEMM_IMPORT_FAILED = True + else: + _DEEP_GEMM_MODULE = deep_gemm + return _DEEP_GEMM_MODULE + + +def _compute_pool_schedule_metadata( + pool_seqlens: torch.Tensor, + *, + slots_per_page: int, +) -> Optional[torch.Tensor]: + if not is_cuda(): + return None + deep_gemm = _get_deep_gemm() + if deep_gemm is None: + return None + return deep_gemm.get_paged_mqa_logits_metadata( + pool_seqlens.contiguous().view(-1, 1).clamp(min=1), + slots_per_page, + deep_gemm.get_num_sms(), + ) + + +def init_pooled_paged_mqa_metadata( + metadata: DSAMetadata, + seqlens_32: torch.Tensor, + forward_mode: ForwardMode, + *, + pool_size: int, + real_page_size: int, + slots_per_page: int, + build_schedule_metadata: bool = True, +) -> DSAMetadata: + if ( + not _is_kpool_layout_enabled(pool_size, real_page_size) + or not is_cuda() + or not forward_mode.is_decode_or_idle() + ): + return metadata + + pool_seqlens = torch.div(seqlens_32, pool_size, rounding_mode="floor").to( + torch.int32 + ) + pooled_page_table = build_pooled_page_table_64( + metadata.real_page_table, pool_size + ).contiguous() + schedule = ( + _compute_pool_schedule_metadata( + pool_seqlens, + slots_per_page=slots_per_page, + ) + if build_schedule_metadata + else None + ) + return dataclasses.replace( + metadata, + pooled_index_kpool=pool_size, + pooled_cache_seqlens_int32=pool_seqlens, + pooled_real_page_table=pooled_page_table, + pooled_paged_mqa_schedule_metadata=schedule, + ) + + +def update_pooled_paged_mqa_metadata( + metadata: DSAMetadata, + seqlens_32: torch.Tensor, + forward_mode: ForwardMode, + *, + pool_size: int, + real_page_size: int, + slots_per_page: int, + build_schedule_metadata: bool = True, +) -> None: + if ( + not _is_kpool_layout_enabled(pool_size, real_page_size) + or not is_cuda() + or not forward_mode.is_decode_or_idle() + ): + return + + if ( + metadata.pooled_index_kpool != pool_size + or metadata.pooled_cache_seqlens_int32 is None + or metadata.pooled_real_page_table is None + ): + return + + pool_seqlens = torch.div(seqlens_32, pool_size, rounding_mode="floor").to( + torch.int32 + ) + metadata.pooled_cache_seqlens_int32[: pool_seqlens.shape[0]].copy_(pool_seqlens) + pooled_page_table = build_pooled_page_table_64( + metadata.real_page_table, pool_size + ).contiguous() + metadata.pooled_real_page_table[ + : pooled_page_table.shape[0], : pooled_page_table.shape[1] + ].copy_(pooled_page_table) + + if ( + build_schedule_metadata + and metadata.pooled_paged_mqa_schedule_metadata is not None + ): + new_schedule = _compute_pool_schedule_metadata( + metadata.pooled_cache_seqlens_int32, + slots_per_page=slots_per_page, + ) + if new_schedule is not None: + metadata.pooled_paged_mqa_schedule_metadata.copy_(new_schedule) + + +def _alloc_kpool_write_plan_buffers( + *, + max_bs: int, + num_draft_tokens: int, + pool_size: int, + device: torch.device, + is_verify: bool, + is_v2: bool = False, +) -> KPoolWritePlan: + max_closed_pools = kpool_max_closed_pools(num_draft_tokens, pool_size) + verify_extras = {} + if is_verify: + n_rows = max_bs * num_draft_tokens + verify_extras = dict( + pool_seqlens_per_q=torch.zeros(n_rows, dtype=torch.int32, device=device), + seqlens_per_q=torch.zeros(n_rows, dtype=torch.int32, device=device), + ) + if is_v2: + verify_extras["effective_n_per_batch"] = torch.zeros( + max_bs, dtype=torch.int32, device=device + ) + return KPoolWritePlan( + req=torch.zeros(max_bs, dtype=torch.int64, device=device), + write_start=torch.zeros(max_bs, dtype=torch.int32, device=device), + tail_logical_start=torch.zeros(max_bs, dtype=torch.int32, device=device), + write_loc=torch.zeros( + max_bs, max_closed_pools, dtype=torch.int64, device=device + ), + num_draft_tokens=num_draft_tokens, + **verify_extras, + ) + + +def init_kpool_write_plan_capture( + metadata: DSAMetadata, + *, + max_bs: int, + pool_size: int, + real_page_size: int, + num_draft_tokens: int, + device: torch.device, + is_verify: bool, + slots_per_page: int, + is_v2: bool = False, + build_schedule_metadata: bool = True, +) -> DSAMetadata: + if not _is_kpool_layout_enabled(pool_size, real_page_size) or num_draft_tokens == 0: + return metadata + + plan = _alloc_kpool_write_plan_buffers( + max_bs=max_bs, + num_draft_tokens=num_draft_tokens, + pool_size=pool_size, + device=device, + is_verify=is_verify, + is_v2=is_v2, + ) + if is_verify and build_schedule_metadata: + schedule = _compute_pool_schedule_metadata( + plan.pool_seqlens_per_q, + slots_per_page=slots_per_page, + ) + plan = dataclasses.replace(plan, pool_schedule_metadata=schedule) + return dataclasses.replace(metadata, kpool_write_plan=plan) + + +def update_kpool_write_plan( + metadata: DSAMetadata, + *, + write_start: torch.Tensor, + req_pool_indices: torch.Tensor, + real_page_table: torch.Tensor, + pool_size: int, + real_page_size: int, + num_draft_tokens: int, + forward_mode: ForwardMode, + slots_per_page: int, + effective_n_per_batch: Optional[torch.Tensor] = None, + include_deep_gemm_schedule: bool = True, +) -> None: + if not _is_kpool_layout_enabled(pool_size, real_page_size) or not is_cuda(): + return + is_verify = forward_mode.is_target_verify() + is_decode = forward_mode.is_decode_or_idle() + is_v2 = forward_mode.is_draft_extend_v2() + if not (is_verify or is_decode or is_v2): + return + + plan = metadata.kpool_write_plan + assert plan is not None, "kpool_write_plan must be allocated before update" + update_kpool_write_plan_cuda_graph( + write_start=write_start, + req_pool_indices=req_pool_indices, + real_page_table=real_page_table, + req_out=plan.req, + write_start_out=plan.write_start, + tail_logical_start_out=plan.tail_logical_start, + write_loc_out=plan.write_loc, + pool_seqlens_per_q_out=plan.pool_seqlens_per_q, + seqlens_per_q_out=plan.seqlens_per_q, + pool_size=pool_size, + num_draft_tokens=num_draft_tokens, + slots_per_page=slots_per_page, + ) + + if ( + is_v2 + and effective_n_per_batch is not None + and plan.effective_n_per_batch is not None + ): + plan.effective_n_per_batch[: effective_n_per_batch.shape[0]].copy_( + effective_n_per_batch.to(torch.int32) + ) + + # In-graph replay updates plan lengths too late for host schedule construction; + # the caller rebuilds the schedule from raw seq_lens out of graph. + if include_deep_gemm_schedule and plan.pool_schedule_metadata is not None: + new_schedule = _compute_pool_schedule_metadata( + plan.pool_seqlens_per_q, + slots_per_page=slots_per_page, + ) + if new_schedule is not None: + plan.pool_schedule_metadata.copy_(new_schedule) + + +def refresh_kpool_pool_schedule_from( + metadata: DSAMetadata, + pool_seqlens_per_q: torch.Tensor, + *, + slots_per_page: int, +) -> None: + """Use an explicit source because the captured plan buffer remains stale + until replay.""" + plan = metadata.kpool_write_plan + if plan is None or plan.pool_schedule_metadata is None: + return + new_schedule = _compute_pool_schedule_metadata( + pool_seqlens_per_q, + slots_per_page=slots_per_page, + ) + if new_schedule is not None: + plan.pool_schedule_metadata.copy_(new_schedule) + + +def init_kpool_write_plan( + metadata: DSAMetadata, + forward_batch: ForwardBatch, + *, + pool_size: int, + real_page_size: int, + real_page_table: torch.Tensor, + num_draft_tokens: int, + write_start: torch.Tensor, + slots_per_page: int, + effective_n_per_batch: Optional[torch.Tensor] = None, + build_schedule_metadata: bool = True, +) -> DSAMetadata: + forward_mode = forward_batch.forward_mode + is_verify = forward_mode.is_target_verify() + is_decode = forward_mode.is_decode_or_idle() + is_v2 = forward_mode.is_draft_extend_v2() + is_ring_write = is_verify or is_decode or is_v2 + if not _is_kpool_layout_enabled(pool_size, real_page_size) or not is_ring_write: + return metadata + + pool = getattr(forward_batch, "token_to_kv_pool", None) + if (is_verify or is_v2) and pool is not None: + assert pool.tail_extra_slots == num_draft_tokens, ( + f"tail_extra_slots mismatch: pool={pool.tail_extra_slots}, " + f"forward={num_draft_tokens}" + ) + + metadata = init_kpool_write_plan_capture( + metadata, + max_bs=forward_batch.seq_lens.shape[0], + pool_size=pool_size, + real_page_size=real_page_size, + num_draft_tokens=num_draft_tokens, + device=forward_batch.seq_lens.device, + is_verify=is_verify or is_v2, + slots_per_page=slots_per_page, + is_v2=is_v2, + build_schedule_metadata=build_schedule_metadata, + ) + update_kpool_write_plan( + metadata, + write_start=write_start, + req_pool_indices=forward_batch.req_pool_indices, + real_page_table=real_page_table, + pool_size=pool_size, + real_page_size=real_page_size, + num_draft_tokens=num_draft_tokens, + forward_mode=forward_mode, + slots_per_page=slots_per_page, + effective_n_per_batch=effective_n_per_batch, + ) + return metadata diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py b/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py index b06e0c5a2..b2ff3643d 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py @@ -138,6 +138,8 @@ class CuteDSLKDAKernel(LinearAttnKernelBase): num_tokens = q_n.shape[0] g_in = g[0][:num_tokens] # raw forget gate; activated inside chunk_kda_cutedsl beta_in = beta[0][:num_tokens].to(torch.float32) + if kwargs.get("beta_is_raw"): + beta_in = beta_in.sigmoid() cu_seqlens = query_start_loc.to(torch.int32) # Pool state I/O is fused into the h kernel's TMA load/store: pass the diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_flashkda.py b/python/sglang/srt/layers/attention/linear/kernels/kda_flashkda.py index 9872cf547..49c2c066e 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/kda_flashkda.py +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_flashkda.py @@ -39,6 +39,7 @@ def _triton_fallback( A_log=None, dt_bias=None, lower_bound=None, + beta_is_raw=False, return_intermediate_states=False, ): """Fall back to the Triton chunk_kda kernel (handles all preprocessing). @@ -64,6 +65,7 @@ def _triton_fallback( A_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound, + beta_is_raw=beta_is_raw, output_intermediate_states=return_intermediate_states, ) @@ -114,6 +116,7 @@ class FlashKDAKernel(LinearAttnKernelBase): lower_bound: Optional[float] = None, extend_seq_lens_cpu: Optional[list] = None, is_spec_decode: bool = False, + beta_is_raw: bool = False, return_intermediate_states: bool = False, **kwargs, ) -> torch.Tensor: @@ -136,6 +139,7 @@ class FlashKDAKernel(LinearAttnKernelBase): A_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound, + beta_is_raw=beta_is_raw, return_intermediate_states=return_intermediate_states, ) @@ -152,6 +156,7 @@ class FlashKDAKernel(LinearAttnKernelBase): A_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound, + beta_is_raw=beta_is_raw, ), None, ) @@ -206,6 +211,7 @@ class FlashKDAKernel(LinearAttnKernelBase): A_log: Optional[torch.Tensor] = None, dt_bias: Optional[torch.Tensor] = None, lower_bound: Optional[float] = None, + beta_is_raw: bool = False, ) -> torch.Tensor: flash_kda = _load_flash_kda() @@ -223,12 +229,11 @@ class FlashKDAKernel(LinearAttnKernelBase): v = v.contiguous() g = g.contiguous() - # KimiDeltaAttention.forward already applies sigmoid to beta on the - # prefill path, but flash_kda expects beta LOGITS (it sigmoids - # internally). Invert back so the kernel recovers the intended value: - # sigmoid(logit(p)) == p. (triton/cuLA consume the post-sigmoid beta.) - beta = torch.logit(beta.float().clamp_(1e-7, 1.0 - 1e-7)).to(torch.bfloat16) - beta = beta.contiguous() + # FlashKDA applies sigmoid internally; invert only the already-activated + # Kimi beta path. + if not beta_is_raw: + beta = torch.logit(beta.float().clamp_(1e-7, 1.0 - 1e-7)) + beta = beta.to(torch.bfloat16).contiguous() # flash_kda wants A_log [H] fp32 and dt_bias [H, K] fp32. The model # stores A_log as [1, 1, H, 1] and dt_bias as 1D [H*K], so reshape both. diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_triton.py b/python/sglang/srt/layers/attention/linear/kernels/kda_triton.py index 400a5303c..8ecbad158 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/kda_triton.py +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_triton.py @@ -171,7 +171,7 @@ class TritonKDAKernel(LinearAttnKernelBase): intermediate_states_buffer: torch.Tensor, intermediate_state_indices: torch.Tensor, cache_steps: int, - retrieve_parent_token: torch.Tensor, + retrieve_parent_token: Optional[torch.Tensor], lower_bound: Optional[float] = None, # fused ReplaySSM ring-write (dense verify only; off elsewhere). cache_ring: bool = False, @@ -229,6 +229,7 @@ class TritonKDAKernel(LinearAttnKernelBase): A_log: Optional[torch.Tensor] = None, dt_bias: Optional[torch.Tensor] = None, lower_bound: Optional[float] = None, + beta_is_raw: bool = False, return_intermediate_states: bool = False, **kwargs, ) -> tuple[torch.Tensor, torch.Tensor | None]: @@ -245,5 +246,6 @@ class TritonKDAKernel(LinearAttnKernelBase): A_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound, + beta_is_raw=beta_is_raw, output_intermediate_states=return_intermediate_states, ) diff --git a/test/registered/kernels/ops/attention/test_kda_helion.py b/test/registered/kernels/ops/attention/test_kda_helion.py index 36a8f34b3..f4f0c558a 100644 --- a/test/registered/kernels/ops/attention/test_kda_helion.py +++ b/test/registered/kernels/ops/attention/test_kda_helion.py @@ -706,6 +706,7 @@ def _compare_prefill( A_log: torch.Tensor | None = None, dt_bias: torch.Tensor | None = None, lower_bound: float | None = None, + beta_is_raw: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: batch, tokens, heads, key_dim = q.shape value_dim = v.size(-1) @@ -740,7 +741,8 @@ def _compare_prefill( k_rows = reference_k.view(batch * tokens, heads, key_dim) v_rows = v.view(batch * tokens, heads, value_dim).float() gate_rows = reference_gate.view(batch * tokens, heads, key_dim) - beta_rows = beta.view(batch * tokens, heads).float() + reference_beta = beta.float().sigmoid() if beta_is_raw else beta + beta_rows = reference_beta.view(batch * tokens, heads).float() out_rows = reference_out.view(batch * tokens, heads, value_dim) if cu_seqlens is None: @@ -811,6 +813,7 @@ def _compare_prefill( A_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound, + beta_is_raw=beta_is_raw, ) assert helion_out.data_ptr() == helion_v.data_ptr() @@ -852,6 +855,42 @@ def test_fixed_partial_prefill_and_state_pool_contract() -> None: assert torch.equal(helion_state[untouched], state[untouched]) +def test_raw_beta_prefill_contract() -> None: + torch.manual_seed(811) + batch, tokens, heads, key_dim, value_dim = 2, 17, 2, 32, 32 + q = torch.randn(batch, tokens, heads, key_dim, device="cuda", dtype=torch.bfloat16) + # Keep the unnormalized recurrence numerically contractive while still + # exercising the no-QK-L2-normalization path. Unit-scale random keys make + # (I - beta * k k^T) expansive and obscure the raw-beta contract with + # exponentially amplified BF16 round-off. + k = torch.randn_like(q) * 0.05 + v = torch.randn( + batch, tokens, heads, value_dim, device="cuda", dtype=torch.bfloat16 + ) + # Keep the recurrent decay contractive so the raw-beta check measures the + # sigmoid conversion instead of amplifying BF16 round-off exponentially. + gate = -torch.rand_like(q) * 0.2 + raw_beta = torch.linspace( + -2, + 2, + steps=batch * tokens * heads, + device="cuda", + ).reshape(batch, tokens, heads) + indices = torch.tensor([3, 1], device="cuda", dtype=torch.int32) + state = torch.randn(5, heads, value_dim, key_dim, device="cuda") * 0.01 + + _compare_prefill( + q, + k, + v, + gate, + raw_beta, + state, + indices, + beta_is_raw=True, + ) + + @pytest.mark.parametrize("is_varlen", [False, True], ids=["fixed", "varlen"]) def test_single_token_prefill_does_not_poison_later_shapes( is_varlen: bool, diff --git a/test/registered/kernels/test_dsa_kpool_multi_pool.py b/test/registered/kernels/test_dsa_kpool_multi_pool.py new file mode 100644 index 000000000..26ce44533 --- /dev/null +++ b/test/registered/kernels/test_dsa_kpool_multi_pool.py @@ -0,0 +1,213 @@ +"""CUDA regressions for DSA kpool speculative writes spanning multiple pools.""" + +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.layers.attention.dsa.kpool_fp8_index import ( + INDEX_HEAD_DIM, + kpool_assemble_softmax_rotate_write_cache, + kpool_max_closed_pools, + kpool_write_tail_and_maybe_compress, + update_kpool_write_plan_cuda_graph, +) +from sglang.srt.layers.attention.dsa.kpool_plan import ( + _alloc_kpool_write_plan_buffers, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +@unittest.skipUnless(torch.cuda.is_available(), "Test requires CUDA") +class TestDsaKpoolMultiPool(CustomTestCase): + POOL_SIZE = 4 + PAGE_SIZE = 64 + SLOTS_PER_PAGE = 64 + NUM_DRAFT_TOKENS = 6 + + def _pool(self) -> SimpleNamespace: + return SimpleNamespace( + page_size=self.PAGE_SIZE, + index_head_dim=INDEX_HEAD_DIM, + slots_per_page=self.SLOTS_PER_PAGE, + index_kpool=self.POOL_SIZE, + tail_extra_slots=self.NUM_DRAFT_TOKENS, + quant_block_size=128, + ) + + def _empty_cache(self) -> torch.Tensor: + page_nbytes = self.SLOTS_PER_PAGE * INDEX_HEAD_DIM + self.SLOTS_PER_PAGE * 4 + return torch.zeros((1, page_nbytes), dtype=torch.uint8, device="cuda") + + def test_write_plan_records_every_candidate_pool(self): + batch_size = 2 + num_draft_tokens = self.NUM_DRAFT_TOKENS + max_closed_pools = kpool_max_closed_pools(num_draft_tokens, self.POOL_SIZE) + self.assertEqual(max_closed_pools, 2) + + plan = _alloc_kpool_write_plan_buffers( + max_bs=batch_size, + num_draft_tokens=num_draft_tokens, + pool_size=self.POOL_SIZE, + device=torch.device("cuda"), + is_verify=True, + ) + self.assertEqual(plan.write_loc.shape, (batch_size, max_closed_pools)) + + write_start = torch.tensor([3, 255], dtype=torch.int32, device="cuda") + req_pool_indices = torch.tensor([7, 11], dtype=torch.int64, device="cuda") + real_page_table = torch.zeros( + (batch_size * num_draft_tokens, 8), + dtype=torch.int32, + device="cuda", + ) + real_page_table[:num_draft_tokens, 0] = 2 + real_page_table[:num_draft_tokens, 4] = 3 + real_page_table[num_draft_tokens:, 0] = 5 + real_page_table[num_draft_tokens:, 4] = 6 + + update_kpool_write_plan_cuda_graph( + write_start=write_start, + req_pool_indices=req_pool_indices, + real_page_table=real_page_table, + req_out=plan.req, + write_start_out=plan.write_start, + tail_logical_start_out=plan.tail_logical_start, + write_loc_out=plan.write_loc, + pool_seqlens_per_q_out=plan.pool_seqlens_per_q, + seqlens_per_q_out=plan.seqlens_per_q, + pool_size=self.POOL_SIZE, + num_draft_tokens=num_draft_tokens, + slots_per_page=self.SLOTS_PER_PAGE, + ) + + torch.testing.assert_close(plan.req, req_pool_indices) + torch.testing.assert_close(plan.write_start, write_start) + torch.testing.assert_close( + plan.tail_logical_start, + torch.tensor([0, 252], dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close( + plan.write_loc, + torch.tensor( + [ + [2 * self.SLOTS_PER_PAGE, 2 * self.SLOTS_PER_PAGE + 1], + [ + 5 * self.SLOTS_PER_PAGE + 63, + 6 * self.SLOTS_PER_PAGE, + ], + ], + dtype=torch.int64, + device="cuda", + ), + ) + + def _run_compress_case(self, effective_n: int, expected_closed_pools: int): + torch.manual_seed(42) + pool = self._pool() + num_draft_tokens = self.NUM_DRAFT_TOKENS + tail_size = self.POOL_SIZE + num_draft_tokens + write_start_value = 3 + + key = torch.randn( + num_draft_tokens, INDEX_HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + score = torch.randn_like(key) + ape = torch.randn( + self.POOL_SIZE, INDEX_HEAD_DIM, dtype=torch.float32, device="cuda" + ) + tail_k_initial = torch.randn( + 1, tail_size, INDEX_HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + tail_score_initial = torch.randn_like(tail_k_initial) + + tail_k_expected = tail_k_initial.clone() + tail_score_expected = tail_score_initial.clone() + for i in range(num_draft_tokens): + physical_slot = (write_start_value + i) % tail_size + tail_k_expected[0, physical_slot] = key[i] + tail_score_expected[0, physical_slot] = score[i] + + expected_cache = self._empty_cache() + dummy_chunk = torch.zeros( + 1, INDEX_HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + kpool_assemble_softmax_rotate_write_cache( + pool=pool, + buf=expected_cache, + chunk_k=dummy_chunk, + chunk_score=dummy_chunk, + tail_k=tail_k_expected, + tail_score=tail_score_expected, + req_pool_idx=torch.zeros( + expected_closed_pools, dtype=torch.int64, device="cuda" + ), + n_from_tail=torch.full( + (expected_closed_pools,), + self.POOL_SIZE, + dtype=torch.int32, + device="cuda", + ), + chunk_src_start=torch.zeros( + expected_closed_pools, dtype=torch.int64, device="cuda" + ), + tail_logical_base=torch.arange( + 0, + expected_closed_pools * self.POOL_SIZE, + self.POOL_SIZE, + dtype=torch.int32, + device="cuda", + ), + ape=ape, + loc=torch.arange(expected_closed_pools, dtype=torch.int64, device="cuda"), + round_scale=False, + ) + + actual_cache = self._empty_cache() + tail_k_actual = tail_k_initial.clone() + tail_score_actual = tail_score_initial.clone() + kpool_write_tail_and_maybe_compress( + pool=pool, + buf=actual_cache, + key=key, + score=score, + tail_k=tail_k_actual, + tail_score=tail_score_actual, + ape=ape, + req_pool_indices=torch.zeros(1, dtype=torch.int64, device="cuda"), + write_start=torch.tensor( + [write_start_value], dtype=torch.int32, device="cuda" + ), + tail_logical_start=torch.zeros(1, dtype=torch.int32, device="cuda"), + write_loc=torch.tensor([[0, 1]], dtype=torch.int64, device="cuda"), + out_cache_loc=torch.arange( + 1, num_draft_tokens + 1, dtype=torch.int64, device="cuda" + ), + num_draft_tokens=num_draft_tokens, + round_scale=False, + effective_n_per_batch=torch.tensor( + [effective_n], dtype=torch.int32, device="cuda" + ), + ) + + torch.testing.assert_close(tail_k_actual, tail_k_expected, atol=0, rtol=0) + torch.testing.assert_close( + tail_score_actual, tail_score_expected, atol=0, rtol=0 + ) + torch.testing.assert_close(actual_cache, expected_cache, atol=0, rtol=0) + + def test_compresses_two_pools_when_draft_window_closes_two(self): + self._run_compress_case( + effective_n=self.NUM_DRAFT_TOKENS, + expected_closed_pools=2, + ) + + def test_effective_n_only_compresses_accepted_pools(self): + self._run_compress_case(effective_n=2, expected_closed_pools=1) + + +if __name__ == "__main__": + unittest.main()