diff --git a/python/sglang/kernels/ops/kvcache/__init__.py b/python/sglang/kernels/ops/kvcache/__init__.py index e2ee65933873..94c5bd9e8170 100644 --- a/python/sglang/kernels/ops/kvcache/__init__.py +++ b/python/sglang/kernels/ops/kvcache/__init__.py @@ -62,6 +62,7 @@ def reshape_and_cache_flash( ("cache_ops", "concat_and_cast_mha_k_triton"), ("cache_ops", "launch_reshape_and_cache_flash"), ("pd_dcp_gather", "copy_mla_rows_into_pack"), + ("kv_read_table", "build_kv_read_table"), ("kv_indices", "create_flashinfer_kv_indices_triton"), ("kv_indices", "create_flashmla_kv_indices_triton"), ("kv_indices", "create_chunked_prefix_cache_kv_indices"), diff --git a/python/sglang/kernels/ops/kvcache/kv_indices.py b/python/sglang/kernels/ops/kvcache/kv_indices.py index 4f5f75f082f6..9f2de2e452b8 100644 --- a/python/sglang/kernels/ops/kvcache/kv_indices.py +++ b/python/sglang/kernels/ops/kvcache/kv_indices.py @@ -7,19 +7,33 @@ @triton.jit def create_flashinfer_kv_indices_triton( - req_to_token_ptr, # [max_batch, max_context_len] + req_to_token_ptr, # [max_batch, max_context_len] token table; at + # ENTRY_PAGE_SIZE > 1 a PAGE-granular table (the unified pool's read table) req_pool_indices_ptr, page_kernel_lens_ptr, kv_indptr, kv_start_idx, kv_indices_ptr, - req_to_token_ptr_stride: tl.constexpr, + # Runtime, not constexpr: the translator's eager table is allocated at the + # batch's live width, so a constexpr stride would JIT-specialize per width + # (a recompile every few decode steps at small page sizes). + req_to_token_ptr_stride, + ENTRY_PAGE_SIZE: tl.constexpr = 1, ): + """Gather per-request token ids into a flat CSR kv_indices stream. + + ``ENTRY_PAGE_SIZE == 1`` (default): the source table is token-granular and + entries are emitted verbatim -- byte-identical to the historical kernel. + ``ENTRY_PAGE_SIZE == ps``: the source is the translator's PAGE-granular + read table (entries already kernel-facing page ids); token ids are rebuilt + as ``token = entry * ps + pos % ps``, exact because converting an id keeps + its offset inside the page. + """ BLOCK_SIZE: tl.constexpr = 512 pid = tl.program_id(axis=0) # find the req pool idx, this is for batch to token - req_pool_index = tl.load(req_pool_indices_ptr + pid) + req_pool_index = tl.load(req_pool_indices_ptr + pid).to(tl.int64) kv_indices_offset = tl.load(kv_indptr + pid) kv_start = 0 @@ -34,13 +48,23 @@ def create_flashinfer_kv_indices_triton( # index into req_to_token_ptr needs to be int64 offset = tl.arange(0, BLOCK_SIZE).to(tl.int64) + i * BLOCK_SIZE mask = offset < kv_end - kv_start - data = tl.load( - req_to_token_ptr - + req_pool_index * req_to_token_ptr_stride - + kv_start - + offset, - mask=mask, - ) + if ENTRY_PAGE_SIZE == 1: + data = tl.load( + req_to_token_ptr + + req_pool_index * req_to_token_ptr_stride + + kv_start + + offset, + mask=mask, + ) + else: + pos = kv_start + offset + entry = tl.load( + req_to_token_ptr + + req_pool_index * req_to_token_ptr_stride + + pos // ENTRY_PAGE_SIZE, + mask=mask, + ) + data = entry.to(tl.int64) * ENTRY_PAGE_SIZE + pos % ENTRY_PAGE_SIZE tl.store(kv_indices_ptr + kv_indices_offset + offset, data, mask=mask) diff --git a/python/sglang/kernels/ops/kvcache/kv_read_table.py b/python/sglang/kernels/ops/kvcache/kv_read_table.py new file mode 100644 index 000000000000..6909c12873cf --- /dev/null +++ b/python/sglang/kernels/ops/kvcache/kv_read_table.py @@ -0,0 +1,140 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Builds the per-batch read table for the unified memory pool. + +One fused gather-and-translate. For each request row it reads the virtual ids +out of `req_to_token`, converts each to the id the kernels can use, and writes +the result into `out`: + + out[b, c] = clamp(v2p[req_to_token[req[b], c * ps] // ps] * multiplier, 0) + for c < ceil(seq_lens[b] / ps) -- the row's LIVE prefix + +`v2p` is the pool's virtual->physical page table and `multiplier` scales a +physical page into the id space the per-layer views use (1 when they are not +dense). Since only the page number is rewritten, a token-level consumer can +rebuild flat ids as `entry * ps + offset`. + +PREFIX-ONLY per row: columns past the live prefix are never written, so a +caller-owned buffer keeps what it had there -- which is what lets a captured +cuda-graph buffer be refreshed in place. Readers bound themselves by +`cache_seqlens` and never look past the prefix. + +A `-1` in `req_to_token` and a freed (`-1`) v2p row both clamp to entry 0, the +reserved padding slot, so a kernel dereferences padding, not a wild address. +""" + +from __future__ import annotations + +import torch +import triton +import triton.language as tl + +_BLOCK_COLS = 256 + + +@triton.jit +def build_kv_read_table_kernel( + req_to_token_ptr, # in: [max_reqs, max_context] -- VIRTUAL token ids + req_pool_indices_ptr, # in: [bs] -- row per batch lane + seq_lens_ptr, # in: [bs] + v2p_ptr, # in: [num_pages + 1] int64 -- virtual->physical page table + out_ptr, # out: [>=bs, >=max_pages] int32 -- the read table + req_stride, # runtime: req_to_token row stride (elements) + out_stride, # runtime: out row stride (elements) + mult, # runtime: kernel_page_multiplier of the target sub-pool + PAGE_SIZE: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + blk = tl.program_id(1) + req = tl.load(req_pool_indices_ptr + bid).to(tl.int64) + seqlen = tl.load(seq_lens_ptr + bid) + n_pages = (seqlen + PAGE_SIZE - 1) // PAGE_SIZE + + cols = blk * BLOCK + tl.arange(0, BLOCK) + mask = cols < n_pages + tok = tl.load( + req_to_token_ptr + req * req_stride + cols.to(tl.int64) * PAGE_SIZE, + mask=mask, + other=0, + ).to(tl.int64) + # Triton's `//` truncates toward zero, so `-1 // ps` is 0 for ps > 1 but + # -1 at ps == 1, which would read one element BEFORE `v2p`. + page = tl.where(tok < 0, 0, tok // PAGE_SIZE) + phys = tl.load(v2p_ptr + page, mask=mask, other=0) + entry = tl.maximum(phys * mult, 0).to(tl.int32) + tl.store(out_ptr + bid.to(tl.int64) * out_stride + cols, entry, mask=mask) + + +def build_kv_read_table( + *, + req_to_token: torch.Tensor, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + v2p: torch.Tensor, + multiplier: int, + page_size: int, + max_pages: int, + out: torch.Tensor, +) -> torch.Tensor: + """Fill ``out``'s live prefix with read-table entries. + + ``out`` is caller-owned (fresh zeros for the eager path, the module's + capture-stable buffer for replay) and only its ``[:bs, :max_pages]`` + region's live prefix is written -- never rebound, never tail-cleared. + """ + bs = int(req_pool_indices.numel()) + assert ( + out.dtype == torch.int32 + ), f"build_kv_read_table: out must be int32, got {out.dtype}" + assert out.dim() == 2 and out.shape[0] >= bs and out.shape[1] >= max_pages, ( + f"build_kv_read_table: out {tuple(out.shape)} cannot hold " + f"(bs={bs}, max_pages={max_pages})" + ) + assert out.stride(1) == 1, "build_kv_read_table: out rows must be packed" + assert (max_pages - 1) * page_size < req_to_token.shape[1], ( + f"build_kv_read_table: max_pages={max_pages} x ps={page_size} " + f"exceeds req_to_token width {req_to_token.shape[1]}" + ) + if bs == 0 or max_pages == 0: + return out + + if not req_to_token.is_cuda: + cols = torch.arange(max_pages, device=req_to_token.device) + live = cols[None, :] < ( + (seq_lens[:bs, None].to(torch.int64) + page_size - 1) // page_size + ) + tok = req_to_token[ + req_pool_indices[:bs, None].to(torch.int64), (cols * page_size)[None, :] + ].to(torch.int64) + pages = torch.where(tok < 0, 0, tok // page_size) + entry = (v2p[pages] * multiplier).clamp(min=0).to(torch.int32) + dst = out[:bs, :max_pages] + dst.copy_(torch.where(live, entry, dst)) + return out + + grid = (bs, triton.cdiv(max_pages, _BLOCK_COLS)) + build_kv_read_table_kernel[grid]( + req_to_token, + req_pool_indices, + seq_lens, + v2p, + out, + req_to_token.stride(0), + out.stride(0), + multiplier, + PAGE_SIZE=page_size, + BLOCK=_BLOCK_COLS, + ) + return out diff --git a/python/sglang/srt/arg_groups/kv_cache_hook.py b/python/sglang/srt/arg_groups/kv_cache_hook.py index 5a58b474e63d..1200ec1acf6c 100644 --- a/python/sglang/srt/arg_groups/kv_cache_hook.py +++ b/python/sglang/srt/arg_groups/kv_cache_hook.py @@ -262,11 +262,21 @@ def handle_unified_memory_pool(server_args: Any) -> None: # Only monolithic decode cuda-graph capture is wired; piecewise prefill # capture is not. Guard when the user opts into it. _cg_cfg = cfg.cuda_graph_config - if _cg_cfg is not None and _cg_cfg.prefill.backend == Backend.TC_PIECEWISE: - raise ValueError( - "--enable-unified-memory supports monolithic (decode) " - "cuda-graph capture only; disable piecewise prefill capture " - "(e.g. --cuda-graph-backend-prefill=disabled)." + if _cg_cfg is not None and _cg_cfg.prefill.backend != Backend.DISABLED: + if cfg.cuda_graph_backend_prefill is not None: + raise ValueError( + "--enable-unified-memory supports decode cuda-graph " + "capture only; prefill capture is not wired (the prefill " + "graph runner bypasses the unified virtual->physical loc " + "rebind). Got --cuda-graph-backend-prefill=" + f"{cfg.cuda_graph_backend_prefill!r}; pass " + "--cuda-graph-backend-prefill=disabled." + ) + _cg_cfg.prefill.backend = Backend.DISABLED + logger.warning( + "--enable-unified-memory: disabling prefill cuda-graph " + "capture (not wired for the unified pool's loc rebind); " + "decode capture is unaffected." ) diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 64dbe640553e..b9eb036c5126 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -29,6 +29,7 @@ get_dcp_lens, ) from sglang.srt.layers.radix_attention import AttentionType +from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.cuda_graph_config import ( @@ -208,14 +209,8 @@ def __init__( # Lets the Triton wrappers specialize on PAGE_SIZE; page_size=1 is # byte-identical to the slot-based envelope. self.page_size = getattr(model_runner, "page_size", 1) or 1 - # Unified pool v2p hook (None = no-op): req_to_token holds VIRTUAL ids but - # kernels need the kernel-facing id space — PHYSICAL for MHA, DENSE for the - # per-layer-view MLA pool (translate_kv_loc_for_kernel falls back to the physical - # translate when kernel_page_multiplier == 1, so preferring it is exact for - # both). Applied eagerly so the captured graph has no translate. - self._translate_kv_loc = getattr( - self.token_to_kv_pool_allocator, "translate_kv_loc_for_kernel", None - ) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None) + self.kv_index_translator = model_runner.kv_index_translator + self.kv_read_tables = None self.num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_steps = get_spec().speculative_num_steps self.topk = get_spec().speculative_eagle_topk or 0 @@ -466,19 +461,20 @@ def _fill_kv_indptr_and_indices( self, bs: int, seq_lens: torch.Tensor, - req_pool_indices: torch.Tensor, + index_table, kv_indices: torch.Tensor, ) -> torch.Tensor: kv_indptr = self.kv_indptr[: bs + 1] kv_indptr[1:] = torch.cumsum(seq_lens, dim=0) create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, + index_table.ids, + index_table.row_ids, seq_lens, kv_indptr, None, kv_indices, - self.req_to_token.stride(0), + index_table.row_stride, + ENTRY_PAGE_SIZE=index_table.entry_page_size, ) return kv_indptr @@ -487,6 +483,7 @@ def _update_decode_kv_buffers( bs: int, seq_lens: torch.Tensor, req_pool_indices: torch.Tensor, + index_table, ): """Fill KV (and SWA) cuda-graph buffers for decode/idle mode. @@ -495,6 +492,9 @@ def _update_decode_kv_buffers( ``num_kv_splits_lens`` is the per-request length used to size kv splits (per-DCP-rank length clamped to >=1 when DCP is enabled, full seq_lens otherwise). + + ``index_table`` is the captured read-index view: under the unified pool the + gathers below read the converted tables. """ seq_lens = seq_lens[:bs] req_pool_indices = req_pool_indices[:bs] @@ -512,27 +512,20 @@ def _update_decode_kv_buffers( num_kv_splits_lens = dcp_seq_lens.clamp_min(1) else: kv_indptr = self._fill_kv_indptr_and_indices( - bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices + bs, seq_lens, index_table, self.cuda_graph_kv_indices ) - # Unified pool: VIRTUAL ids written here are translated to PHYSICAL in - # init_forward_metadata_out_graph (replay-prep) so the captured graph - # carries zero translate nodes. num_kv_splits_lens = seq_lens window_kv_indptr = self.window_kv_indptr window_kv_lens = None if self.sliding_window_size is not None and self.sliding_window_size > 0: - # Unified pool: leave the window VIRTUAL too (translated alongside the - # full kv_indices later); baseline SWA keeps the eager window translate. window_kv_indptr, _, window_kv_lens, _ = update_sliding_window_buffer( self.window_kv_indptr, - self.req_to_token, + index_table, self.sliding_window_size, seq_lens, - req_pool_indices, bs, token_to_kv_pool=self.token_to_kv_pool, window_kv_indices=self.cuda_graph_window_kv_indices, - skip_full_to_swa_translation=(self._translate_kv_loc is not None), ) return kv_indptr, window_kv_indptr, window_kv_lens, num_kv_splits_lens @@ -540,8 +533,8 @@ def _update_target_verify_buffers( self, bs: int, seq_lens: torch.Tensor, - req_pool_indices: torch.Tensor, spec_info, + index_table, ): """Fill all cuda-graph buffers for target_verify mode.""" # Prefer the spec_info's per-request query length (DSpark draft propose @@ -561,7 +554,7 @@ def _update_target_verify_buffers( device=self.device, ) kv_indptr = self._fill_kv_indptr_and_indices( - bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices + bs, seq_lens, index_table, self.cuda_graph_kv_indices ) window_kv_indptr = self.window_kv_indptr window_kv_indices = None @@ -574,10 +567,9 @@ def _update_target_verify_buffers( window_kv_indptr, window_kv_indices, _, window_kv_offsets[:bs] = ( update_sliding_window_buffer( self.window_kv_indptr, - self.req_to_token, + index_table, self.sliding_window_size, seq_lens[:bs], - req_pool_indices, bs, token_to_kv_pool=self.token_to_kv_pool, window_kv_indices=window_kv_indices, @@ -611,9 +603,9 @@ def _update_draft_extend_buffers( self, bs: int, seq_lens: torch.Tensor, - req_pool_indices: torch.Tensor, forward_mode: ForwardMode, spec_info: Optional[SpecInput], + index_table, ): """Fill QO + KV cuda-graph buffers for draft_extend mode.""" seq_lens = seq_lens[:bs] @@ -644,7 +636,7 @@ def _update_draft_extend_buffers( extend_seq_lens = torch.zeros(bs, dtype=torch.int32, device=seq_lens.device) kv_lens = torch.clamp(seq_lens - extend_seq_lens, min=0).to(torch.int32) kv_indptr = self._fill_kv_indptr_and_indices( - bs, kv_lens, req_pool_indices, self.cuda_graph_kv_indices + bs, kv_lens, index_table, self.cuda_graph_kv_indices ) return qo_indptr, kv_indptr, num_tokens_per_req @@ -689,7 +681,7 @@ def init_forward_metadata_out_graph( forward_mode=forward_mode, spec_info=spec_info, ) - out_cache_loc_full_physical = self._translate_cuda_graph_shared_pool_locs( + out_cache_loc_full_physical = self._fill_cuda_graph_write_locs( forward_batch, bs ) swa_out_cache_loc = self._fill_cuda_graph_swa_out_cache_loc(forward_batch) @@ -709,15 +701,16 @@ def init_forward_metadata_out_graph( spec_info=spec_info, ) # Metadata view is reused from capture; just refill the buffers. - self._translate_cuda_graph_shared_pool_locs(forward_batch, bs) + self._fill_cuda_graph_write_locs(forward_batch, bs) self._fill_cuda_graph_swa_out_cache_loc(forward_batch) def _fill_cuda_graph_swa_out_cache_loc( self, forward_batch: ForwardBatch ) -> Optional[torch.Tensor]: - """Refill the SWA write-target buffer from live out_cache_loc, returning the - [:n] view (None for non-SWA / multi-step draft) so the captured store reads - fresh slots on replay.""" + """Refill the SWA write-target buffer from the batch's derived + sliding-window write loc, returning the [:n] view (None for non-SWA / + multi-step draft) so the captured store reads fresh slots on replay. + """ if not self.use_sliding_window_kv_pool: return None out_cache_loc = forward_batch.out_cache_loc @@ -726,68 +719,32 @@ def _fill_cuda_graph_swa_out_cache_loc( or out_cache_loc.shape[0] > self.cuda_graph_swa_out_cache_loc.shape[0] ): return None + swa_write_loc = self.kv_index_translator.sliding_window_write_loc_for( + out_cache_loc + ) n = out_cache_loc.shape[0] self.cuda_graph_swa_out_cache_loc[n:].zero_() - self.cuda_graph_swa_out_cache_loc[:n].copy_( - self.token_to_kv_pool.translate_loc_from_full_to_swa(out_cache_loc) - ) + self.cuda_graph_swa_out_cache_loc[:n].copy_(swa_write_loc) return self.cuda_graph_swa_out_cache_loc[:n] - def _translate_cuda_graph_shared_pool_locs( + def _fill_cuda_graph_write_locs( self, forward_batch: ForwardBatch, bs: int ) -> Optional[torch.Tensor]: - """Unified pool: eager v2p translate of the cuda-graph read+write LOC buffers, - run BEFORE graph.replay() reading the live post-compaction v2p, so the - captured graph carries zero translate nodes. No-op for non-unified pools. - - Read buffers (full kv_indices, SWA window) are translated IN PLACE; the - full-attn WRITE loc is RETURNED as the [:n] view of the backend-owned - out_cache_loc_full_physical buffer. Eager .item() bounds are fine here - (out-of-graph), so no in-graph translate variant is needed. + """Copy the cuda-graph WRITE loc into the capture-stable buffer and + return the ``[:n]`` view; no-op for non-unified pools. + + Runs BEFORE graph.replay() so it reads the live post-compaction v2p. + The capture batch is runner-built with zeros, which is safe because + slot 0 is the reserved sink in every id space. """ - if self._translate_kv_loc is None: + if not self.kv_index_translator.is_translating: return None - # seq_lens_sum is the reliable "mirror present" signal: it is - # None-preserving into the replay view, unlike seq_lens_cpu (always a - # non-None but stale slice for gpu_only batches). None -> fall back to a - # per-step D2H `.item()` on the indptr. - have_cpu_mirror = forward_batch.seq_lens_sum is not None - # Full-attention read path. kv_indptr[bs] == seq_lens_sum. - n_kv = ( - forward_batch.seq_lens_sum - if have_cpu_mirror - else int(self.kv_indptr[bs].item()) - ) - if n_kv > 0: - self.cuda_graph_kv_indices[:n_kv] = self._translate_kv_loc( - self.cuda_graph_kv_indices[:n_kv] - ) - # SWA window read path. window_kv_indptr[bs] == sum(min(seq_len, window)). - if self.sliding_window_size is not None and self.sliding_window_size > 0: - if have_cpu_mirror: - n_win = int( - forward_batch.seq_lens_cpu[:bs] - .clamp(max=self.sliding_window_size) - .sum() - ) - else: - n_win = int(self.window_kv_indptr[bs].item()) - if n_win > 0: - self.cuda_graph_window_kv_indices[:n_win] = ( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - self.cuda_graph_window_kv_indices[:n_win] - ) - ) - # Full-attention write path: translate out_cache_loc -> physical into the - # capture-stable buffer and RETURN the [:n] view. out_cache_loc = forward_batch.out_cache_loc n = out_cache_loc.shape[0] # Zero the padded tail first: a smaller replay batch leaves [n:] holding # stale ids that the captured store would write; send them to slot 0 (sink). self.cuda_graph_out_cache_loc_full_physical[n:].zero_() - self.cuda_graph_out_cache_loc_full_physical[:n].copy_( - self._translate_kv_loc(out_cache_loc) - ) + self.cuda_graph_out_cache_loc_full_physical[:n].copy_(out_cache_loc) return self.cuda_graph_out_cache_loc_full_physical[:n] def init_forward_metadata(self, forward_batch: ForwardBatch): @@ -806,6 +763,9 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): if forward_batch.forward_mode.is_decode_or_idle(): if spec_info is None or spec_info.kv_indptr is None: + index_table = self.kv_index_translator.index_table_for_batch( + forward_batch + ) # kv_indptr is None for draft-extend's idle batch; build from seq_lens. if self.dcp_size > 1: # DCP: per-rank sharded KV indices, else each rank reads the @@ -826,11 +786,9 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): kv_indptr = self._fill_kv_indptr_and_indices( bs, forward_batch.seq_lens, - forward_batch.req_pool_indices, + index_table, kv_indices, ) - if self._translate_kv_loc is not None: - kv_indices = self._translate_kv_loc(kv_indices) if ( self.sliding_window_size is not None and self.sliding_window_size > 0 @@ -838,10 +796,9 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): window_kv_indptr, window_kv_indices, window_kv_lens, _ = ( update_sliding_window_buffer( self.window_kv_indptr, - self.req_to_token, + index_table, self.sliding_window_size, forward_batch.seq_lens, - forward_batch.req_pool_indices, bs, self.device, self.token_to_kv_pool, @@ -931,10 +888,11 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): kv_indices = torch.empty( seq_lens_sum, dtype=torch.int64, device=self.device ) + index_table = self.kv_index_translator.index_table_for_batch(forward_batch) kv_indptr = self._fill_kv_indptr_and_indices( bs, forward_batch.seq_lens, - forward_batch.req_pool_indices, + index_table, kv_indices, ) @@ -947,10 +905,9 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): window_kv_offsets, ) = update_sliding_window_buffer( self.window_kv_indptr, - self.req_to_token, + index_table, self.sliding_window_size, forward_batch.seq_lens, - forward_batch.req_pool_indices, bs, self.device, self.token_to_kv_pool, @@ -969,6 +926,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): attn_lse = None else: + index_table = self.kv_index_translator.index_table_for_batch(forward_batch) if self.dcp_size > 1: kv_indptr, kv_indices, _ = self._dcp_kv_indices( forward_batch.req_pool_indices, @@ -989,11 +947,9 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): kv_indptr = self._fill_kv_indptr_and_indices( bs, forward_batch.extend_prefix_lens, - forward_batch.req_pool_indices, + index_table, kv_indices, ) - if self._translate_kv_loc is not None: - kv_indices = self._translate_kv_loc(kv_indices) if self.sliding_window_size is not None and self.sliding_window_size > 0: ( window_kv_indptr, @@ -1002,10 +958,9 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): window_kv_offsets, ) = update_sliding_window_buffer( self.window_kv_indptr, - self.req_to_token, + index_table, self.sliding_window_size, forward_batch.extend_prefix_lens, - forward_batch.req_pool_indices, bs, self.device, self.token_to_kv_pool, @@ -1027,18 +982,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): swa_out_cache_loc = None if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None: - swa_out_cache_loc = self.token_to_kv_pool.translate_loc_from_full_to_swa( - forward_batch.out_cache_loc - ) - - # Unified pool full-attention WRITE loc (virtual out_cache_loc -> physical), - # carried in the metadata (-> KVWriteLoc.full_loc). None for non-unified pools. - out_cache_loc_full_physical = None - if ( - self._translate_kv_loc is not None - and forward_batch.out_cache_loc is not None - ): - out_cache_loc_full_physical = self._translate_kv_loc( + swa_out_cache_loc = self.kv_index_translator.sliding_window_write_loc_for( forward_batch.out_cache_loc ) @@ -1058,7 +1002,11 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): window_kv_offsets, swa_attn_logits=swa_attn_logits, swa_out_cache_loc=swa_out_cache_loc, - out_cache_loc_full_physical=out_cache_loc_full_physical, + out_cache_loc_full_physical=( + forward_batch.out_cache_loc + if self.kv_index_translator.is_translating + else None + ), lean_Mp=lean_Mp, lean_Lp=lean_Lp, lean_Op=lean_Op, @@ -1179,7 +1127,7 @@ def init_cuda_graph_state( device=self.device, ) - if self._translate_kv_loc is not None: + if self.kv_index_translator.is_translating: # Unified pool full-attention write-target buffer, refilled at replay # (-> KVWriteLoc.full_loc). Capture-stable, mirrors cuda_graph_swa_out_cache_loc. self.cuda_graph_out_cache_loc_full_physical = torch.zeros( @@ -1187,6 +1135,9 @@ def init_cuda_graph_state( dtype=torch.int64, device=self.device, ) + self.kv_read_tables = self.kv_index_translator.make_capture_tables( + max_bs=max_bs, max_context_len=self.max_context_len + ) def _build_cuda_graph_forward_metadata( self, @@ -1300,16 +1251,21 @@ def _apply_cuda_graph_metadata( seq_lens: torch.Tensor, forward_mode: ForwardMode, spec_info: Optional[SpecInput], - ): + ) -> None: """Shared capture+replay body for the cuda-graph init path. Public entry: :py:meth:`init_forward_metadata_out_graph`. """ # NOTE: encoder_lens expected to be zeros or None + index_table = self.kv_index_translator.build_index_table( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + into=self.kv_read_tables, + ) if forward_mode.is_decode_or_idle(): assert spec_info is None, "Multi-step cuda graph init is not done here." _, _, window_kv_lens, num_kv_splits_lens = self._update_decode_kv_buffers( - bs, seq_lens, req_pool_indices + bs, seq_lens, req_pool_indices, index_table ) self.get_num_kv_splits( self.cuda_graph_num_kv_splits[:bs], num_kv_splits_lens[:bs] @@ -1320,12 +1276,10 @@ def _apply_cuda_graph_metadata( ) elif forward_mode.is_target_verify(): bs = len(req_pool_indices) - self._update_target_verify_buffers( - bs, seq_lens, req_pool_indices, spec_info - ) + self._update_target_verify_buffers(bs, seq_lens, spec_info, index_table) elif forward_mode.is_draft_extend_v2(): self._update_draft_extend_buffers( - bs, seq_lens, req_pool_indices, forward_mode, spec_info + bs, seq_lens, forward_mode, spec_info, index_table ) else: raise ValueError( @@ -1402,10 +1356,11 @@ def forward_extend( pool = self.token_to_kv_pool cache_loc = forward_batch.out_cache_loc if isinstance(pool, SWAKVPool) and pool.layers_mapping[layer.layer_id][1]: - cache_loc = pool.translate_loc_from_full_to_swa(cache_loc) - elif self._translate_kv_loc is not None: - # Unified pool: buffers are indexed in the kernel-facing id space. - cache_loc = self._translate_kv_loc(cache_loc) + assert self.forward_metadata.swa_out_cache_loc is not None, ( + "window-layer read-back before the metadata carried a " + "sliding-window write loc" + ) + cache_loc = self.forward_metadata.swa_out_cache_loc k_buffer, v_buffer = pool.get_kv_buffer(layer.layer_id) k = k_buffer[cache_loc] v = v_buffer[cache_loc] @@ -1776,15 +1731,12 @@ def _forward_extend_unified( and isinstance(pool, SWAKVPool) and pool.layers_mapping[layer.layer_id][1] ): - # Consumes VIRTUAL ids, so it must see out_cache_loc untranslated. - extend_kv_indices = pool.translate_loc_from_full_to_swa(extend_kv_indices) + extend_kv_indices = self.forward_metadata.swa_out_cache_loc + assert extend_kv_indices is not None, ( + "window-layer extend before the metadata carried a " + "sliding-window write loc" + ) elif self.forward_metadata.out_cache_loc_full_physical is not None: - # Unified pool: this kernel reads the extend half OUT OF THE POOL (the - # 2-stage path takes it from the k/v arguments), so it needs the same - # translated loc the KV write uses -- otherwise the prefix is read at - # physical ids and the extend tokens at virtual ones. Reuse the - # per-forward translation rather than re-translating: this runs once - # per layer. extend_kv_indices = self.forward_metadata.out_cache_loc_full_physical # Handle cases where extend_seq_lens or extend_start_loc might not be set @@ -2246,15 +2198,13 @@ def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: def update_sliding_window_buffer( window_kv_indptr, - req_to_token, + index_table, sliding_window_size, seq_lens, - req_pool_indices, bs, device=None, token_to_kv_pool=None, window_kv_indices=None, - skip_full_to_swa_translation=False, ): """Fill window KV buffers for sliding-window attention. @@ -2262,13 +2212,12 @@ def update_sliding_window_buffer( path); omit it (or pass ``None``) to allocate a fresh tensor (eager path, requires ``device``). - ``skip_full_to_swa_translation=True`` leaves ``window_kv_indices`` as VIRTUAL - full-token ids (no eager full->swa translate). The unified-memory-pool cuda-graph - builder passes this so the window translate is deferred to - ``TritonAttnBackend._translate_cuda_graph_shared_pool_locs`` (run in - ``init_forward_metadata_out_graph``, BEFORE ``graph.replay()``), which reads - the live v2p and rewrites the static window buffer to swa-physical in place; - baseline SWA leaves it False (eager). + ``index_table`` is the batch's read-index source view. Unified pool: the + gather reads the parallel SWA array (built directly from virtual ids + through the swa side's own v2p), so the window indices come out + already swa-side ids -- no translate here, eager or captured. Static SWA + pools gather full-token ids from req_to_token and keep the legacy + full->swa translate below. """ window_kv_lens = torch.minimum( seq_lens, @@ -2281,18 +2230,18 @@ def update_sliding_window_buffer( window_kv_indptr[-1], dtype=torch.int64, device=device ) window_kv_start_idx = seq_lens - window_kv_lens + source_ids = index_table.sliding_window_read_ids() create_flashinfer_kv_indices_triton[(bs,)]( - req_to_token, - req_pool_indices, + source_ids, + index_table.row_ids, window_kv_lens, window_kv_indptr, window_kv_start_idx, window_kv_indices, - req_to_token.stride(0), + source_ids.stride(0), + ENTRY_PAGE_SIZE=index_table.entry_page_size, ) - if not skip_full_to_swa_translation and hasattr( - token_to_kv_pool, "translate_loc_from_full_to_swa" - ): + if not index_table.is_translated and isinstance(token_to_kv_pool, BaseSWAKVPool): kv_last_index = window_kv_indptr[-1] window_kv_indices[:kv_last_index] = ( token_to_kv_pool.translate_loc_from_full_to_swa( diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index a43b28c7895b..24820d17ce50 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -294,11 +294,8 @@ def __init__( self._v2p_page_table = _hooks.v2p_page_table self._kernel_page_multiplier = _hooks.kernel_page_multiplier self._unified_mla = _hooks.enabled - # virtual token id -> DENSE kernel-facing id, for the KV write loc. - self._translate_kv_loc_dense = _hooks.translate_kv_loc_for_kernel - # Per-forward kernel-facing write loc ([:n] view of a capture-stable buffer), - # set by the cuda-graph out-graph hook; None on the eager path (where the - # write translates through the pool's _full_translate hook instead). + # Per-forward kernel-facing write loc ([:n] view of a capture-stable buffer); + # None on the eager path, which passes out_cache_loc straight through. self._decode_kernel_loc: Optional[torch.Tensor] = None self.cuda_graph_out_cache_loc_kernel: Optional[torch.Tensor] = None # Fused KV-scatter + q-concat on the decode dense-loc path (one launch @@ -636,7 +633,7 @@ def init_forward_metadata_out_graph( out_cache_loc = forward_batch.out_cache_loc n = out_cache_loc.shape[0] dst = self.cuda_graph_out_cache_loc_kernel[:n] - self._translate_kv_loc_dense(out_cache_loc, out=dst) + dst.copy_(out_cache_loc) # Replay-prep receives the RAW (unpadded) out_cache_loc # (build_replay_fb_view), but the captured write kernel consumes the # full captured tier of this buffer. Zero the tail so pad rows write @@ -651,8 +648,6 @@ def init_forward_metadata_out_graph( def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize the metadata for a forward pass.""" - # Eager path: no capture-stable kernel-facing write loc; the pool's _full_translate - # hook translates the write loc (safe out of a cuda graph). self._decode_kernel_loc = None # Delegate to parent for non-decode modes. if ( @@ -950,11 +945,9 @@ def _set_kv_and_concat_q_fused( q: torch.Tensor, q_rope: torch.Tensor, ) -> Optional[torch.Tensor]: - """Decode: scatter the KV row at ``loc`` (already physical — the - dense-loc buffer on the unified pool, or out_cache_loc on the static - pool where ``_full_translate`` is identity) and build the - [q_nope | q_rope] fmha query in one kernel launch (saves one launch - per MLA layer and keeps the PDL chain intact). + """Decode: scatter the KV row at ``loc`` (already kernel-facing) and + build the [q_nope | q_rope] fmha query in one kernel launch (saves one + launch per MLA layer and keeps the PDL chain intact). Returns the concatenated query, or None when the fused kernel does not cover the inputs (caller falls back to the two-kernel path). @@ -1100,8 +1093,6 @@ def forward_decode( k is not None and k_rope is not None ), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None." if self._decode_kernel_loc is not None: - # cuda-graph path: kernel-facing write loc precomputed out-of-graph, so - # the in-graph write captures no translate allocation. if merge_query and self._fused_set_kv_concat_q: # Fused: KV scatter + [q_nope | q_rope] concat in one # launch; None when the inputs are not covered. @@ -1115,21 +1106,16 @@ def forward_decode( ) if query is None: self.token_to_kv_pool.set_mla_kv_buffer( - layer, - self._decode_kernel_loc, - k, - k_rope, - loc_is_kernel_facing=True, + layer, self._decode_kernel_loc, k, k_rope ) else: - # eager (or static pool): the pool's _full_translate handles it. + # eager (or static pool): out_cache_loc is kernel-facing. if ( merge_query and self._fused_set_kv_concat_q and not self._unified_mla ): - # Static pool: _full_translate is identity, so - # out_cache_loc is already the physical write loc. + # Static pool only, conservatively. query = self._set_kv_and_concat_q_fused( layer=layer, loc=forward_batch.out_cache_loc, @@ -1265,7 +1251,7 @@ def forward_extend( ), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None." if self._decode_kernel_loc is not None: self.token_to_kv_pool.set_mla_kv_buffer( - layer, self._decode_kernel_loc, k, k_rope, loc_is_kernel_facing=True + layer, self._decode_kernel_loc, k, k_rope ) else: self.token_to_kv_pool.set_mla_kv_buffer( diff --git a/python/sglang/srt/mem_cache/kv_index_translator.py b/python/sglang/srt/mem_cache/kv_index_translator.py new file mode 100644 index 000000000000..0e2399cec308 --- /dev/null +++ b/python/sglang/srt/mem_cache/kv_index_translator.py @@ -0,0 +1,358 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Turns the KV ids stored in `req_to_token` into ids attention kernels can use. + +A KV slot can be named in three id spaces: + + * **virtual** - what `req_to_token` stores. Keeps naming the same logical + slot even after the pool moves data around. + * **physical** - where that slot sits in the pool right now. + * **kernel-facing** - what a kernel can index the per-layer K/V tensors + with. Same as physical on a plain pool; under the unified pool it is the + physical page scaled by the per-page block count. + +All three coincide on a plain pool, so nothing here does any work there. + +Backends get a `KVIndexTable`, which answers "what do I gather from, and +which row is mine?": + + ids[row_ids[b], pos] + + plain pool : ids = req_to_token, row_ids = req_pool_indices (those very + objects - no copy, no kernel) + unified : ids = a built array of kernel-facing ids, + row_ids = arange(batch_size) + +Backends call their own copy a *page table* (fa3) or a *block table* +(trtllm); here it is the **index table**. + +Converting only ever rewrites the page number and keeps the in-page offset, so +one page-granular table serves both kinds of consumer: a block-table backend +uses its rows as-is, and one that wants flat per-token ids rebuilds them as + + token_id = entry * entry_page_size + pos % entry_page_size + +WRITES, IN TWO PHASES. The full-side write loc is rebound to kernel-facing +ids at ForwardBatch construction - the earliest consumer can snapshot it +right after. The sliding-window write loc is derived at the same moment as +read table, into the same index table. +""" + +from __future__ import annotations + +import weakref +from typing import Optional, Tuple + +import msgspec +import torch + +from sglang.kernels.ops.kvcache.kv_read_table import build_kv_read_table +from sglang.srt.mem_cache.multi_ended_allocator import ( + UnifiedMambaTokenToKVPoolAllocator, + UnifiedSWATokenToKVPoolAllocator, +) +from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool + + +class KVReadTables(msgspec.Struct, frozen=True): + """One capture-stable destination, however many id spaces the pool has. + + A backend holds this and hands it back to `build_index_table(into=...)`; + it never has to know whether there is a sliding-window space behind it. + """ + + full: torch.Tensor + sliding_window: Optional[torch.Tensor] + + +class KVIndexTable(msgspec.Struct, frozen=True): + """Collection of what one batch gathers from.""" + + ids: torch.Tensor # 2-D array of KV ids to gather from + row_ids: torch.Tensor # which row belongs to batch lane b + row_stride: int # stride between rows of `ids`, in elements + entry_page_size: int # what one entry covers: 1 = a token, N = a page of N + is_translated: bool # entries are already kernel-facing ids + sliding_window_ids: Optional[torch.Tensor] # SWA models: the parallel swa array + + def sliding_window_read_ids(self) -> torch.Tensor: + """Which array a sliding-window gather reads: the parallel swa array + when translated, else the full-attention array, which the caller maps + through the pool's own full->swa map.""" + return self.sliding_window_ids if self.is_translated else self.ids + + +class KVIndexTranslator: + """Built once per ModelRunner.""" + + def __init__( + self, + *, + req_to_token: torch.Tensor, + token_to_kv_pool_allocator, + token_to_kv_pool, + page_size: int, + device: str, + ): + self.req_to_token = req_to_token + self.page_size = page_size + self.device = device + + self.is_translating = ( + isinstance( + token_to_kv_pool_allocator, + (UnifiedMambaTokenToKVPoolAllocator, UnifiedSWATokenToKVPoolAllocator), + ) + and token_to_kv_pool_allocator.get_kvcache() is token_to_kv_pool + ) + if self.is_translating: + alloc = token_to_kv_pool_allocator + self._full_v2p_table = alloc.full_v2p_page_table + self._full_p2v_table = alloc.full_p2v_page_table + self._full_page_multiplier = alloc.kernel_page_multiplier + self._translate_full = alloc.translate_kv_loc_for_kernel + if isinstance(alloc, UnifiedSWATokenToKVPoolAllocator): + self._swa_v2p_table = alloc.swa_v2p_page_table + self._swa_page_multiplier = alloc.swa_kernel_page_multiplier + self._swa_write_loc_from_full = self._swa_write_loc_unified + else: + self._swa_v2p_table = None + self._swa_page_multiplier = 1 + self._swa_write_loc_from_full = None + else: + self._full_v2p_table = None + self._full_p2v_table = None + self._full_page_multiplier = 1 + self._translate_full = None + self._swa_v2p_table = None + self._swa_page_multiplier = 1 + self._swa_write_loc_from_full = ( + token_to_kv_pool.translate_loc_from_full_to_swa + if isinstance(token_to_kv_pool, SWAKVPool) + else None + ) + + self._rows: Optional[torch.Tensor] = ( + torch.arange(req_to_token.shape[0], dtype=torch.int64, device=device) + if self.is_translating + else None + ) + self._index_table_memo: Optional[Tuple[weakref.ref, KVIndexTable]] = None + + def make_capture_tables( + self, *, max_bs: int, max_context_len: int + ) -> Optional[KVReadTables]: + """Capture-stable destinations for a backend to own, or None when this + pool needs no translation and the backend will never fill any. + + Zero-filled: entry 0 is the reserved padding slot in every id space, so + a captured graph replaying before its first refresh reads padding, not + garbage. + """ + if not self.is_translating: + return None + max_pages = -(-max_context_len // self.page_size) + + def _zeros(): + return torch.zeros( + (max_bs, max_pages), dtype=torch.int32, device=self.device + ) + + return KVReadTables( + full=_zeros(), + sliding_window=_zeros() if self._swa_v2p_table is not None else None, + ) + + # -- per-batch view -------------------------------------------------------- + + def build_index_table( + self, + *, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + max_pages: Optional[int] = None, + into: Optional[KVReadTables] = None, + ) -> KVIndexTable: + """The one per-batch entry point. + + Non-unified: the raw ``(req_to_token, req_pool_indices)`` passthrough, + no tensor ops and no copies. Unified: fills each row's live prefix and + returns the table WHOLE, so a caller needing a stable pointer (a + captured graph bakes it) passes its own tables in ``into``; + ``into=None`` allocates of width ``max_pages`` instead. + """ + if not self.is_translating: + return KVIndexTable( + ids=self.req_to_token, + row_ids=req_pool_indices, + row_stride=self.req_to_token.stride(0), + entry_page_size=1, + is_translated=False, + sliding_window_ids=None, + ) + + bs = int(req_pool_indices.numel()) + if into is not None: + out_full = into.full + out_swa = into.sliding_window + # A caller-owned table may be padded wider than req_to_token's span + # (trtllm_mla / flashmla pad to a page-count bound); the columns + # past it have no source to read, so stop there. + width = min( + out_full.shape[1] if max_pages is None else max_pages, + -(-self.req_to_token.shape[1] // self.page_size), + ) + else: + assert max_pages is not None, ( + "KVIndexTranslator.build_index_table: allocating needs max_pages " + "(from the batch's seq_lens_cpu max)" + ) + width = max_pages + out_full = torch.zeros((bs, width), dtype=torch.int32, device=self.device) + out_swa = ( + torch.zeros((bs, width), dtype=torch.int32, device=self.device) + if self._swa_v2p_table is not None + else None + ) + + build_kv_read_table( + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + v2p=self._full_v2p_table, + multiplier=self._full_page_multiplier, + page_size=self.page_size, + max_pages=width, + out=out_full, + ) + if out_swa is not None: + build_kv_read_table( + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + v2p=self._swa_v2p_table, + multiplier=self._swa_page_multiplier, + page_size=self.page_size, + max_pages=width, + out=out_swa, + ) + return KVIndexTable( + ids=out_full, + row_ids=self._rows[:bs], + row_stride=out_full.stride(0), + entry_page_size=self.page_size, + is_translated=True, + sliding_window_ids=out_swa, + ) + + def fill_read_table( + self, + *, + out: torch.Tensor, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + ) -> None: + """`build_index_table(into=...)` for a caller that owns a bare block + table rather than a KVReadTables: trtllm_mla / flashmla consume that + table directly, its rows already being the index table's rows. + """ + assert ( + self.is_translating + ), "KVIndexTranslator.fill_read_table on a pool that needs no translation" + self.build_index_table( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + into=KVReadTables(full=out, sliding_window=None), + ) + + def index_table_for_batch(self, forward_batch) -> KVIndexTable: + """Eager per-batch view, memoized in one slot keyed by batch identity + so multi-consumer metadata builds share a build. The next batch + replaces the slot; consumers only read during their own build. The + captured path does not memoize -- it refreshes its buffers per + replay.""" + memo = self._index_table_memo + if memo is not None and memo[0]() is forward_batch: + return memo[1] + max_pages = None + if self.is_translating: + # `seq_lens_cpu` is a non-None but STALE slice on a gpu_only + # batch; `seq_lens_sum` is the signal that it is live. A stale max + # under-sizes the table and the tail then reads as the sink. + slc = forward_batch.seq_lens_cpu + if ( + forward_batch.seq_lens_sum is not None + and slc is not None + and slc.numel() > 0 + ): + max_seq = int(slc.max()) + else: + max_seq = self.req_to_token.shape[1] + max_pages = max(-(-max_seq // self.page_size), 1) + view = self.build_index_table( + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + max_pages=max_pages, + ) + self._index_table_memo = (weakref.ref(forward_batch), view) + return view + + # -- write loc (phase 1; phase 2 lives in build_index_table) ---------------- + + def rebind_write_loc(self, forward_batch) -> None: + """Phase 1 of the WRITE contract: translate the batch's write loc to + FULL-side kernel-facing ids exactly once, at ForwardBatch + construction. No-op on non-unified pools. + + REBIND, never mutate: the translate returns a FRESH tensor, so the + ScheduleBatch's aliased tensor stays VIRTUAL for the radix / accept / + in-flight machinery that reads it. + """ + self._index_table_memo = None + if not self.is_translating or forward_batch.out_cache_loc is None: + return + forward_batch.out_cache_loc = self._translate_full(forward_batch.out_cache_loc) + + def sliding_window_write_loc_for( + self, out_cache_loc: Optional[torch.Tensor] + ) -> Optional[torch.Tensor]: + """This batch's sliding-window write loc, or None when there is no loc + this forward or the pool has no sliding-window id space.""" + if out_cache_loc is None or self._swa_write_loc_from_full is None: + return None + return self._swa_write_loc_from_full(out_cache_loc) + + def _swa_write_loc_unified(self, kernel_loc: torch.Tensor) -> torch.Tensor: + """Sliding-window write loc, derived pointwise from FULL-side + kernel-facing values (phase 2 of the write contract). + """ + full_stride = self.page_size * self._full_page_multiplier + offset = kernel_loc % full_stride # == virtual_token % page_size + # An unmapped physical page reads back as -1; clamp it rather than let + # the gather wrap onto the v2p table's last element. + virt_page = self._full_p2v_table[kernel_loc // full_stride].clamp_(min=0) + swa_stride = self.page_size * self._swa_page_multiplier + return (self._swa_v2p_table[virt_page] * swa_stride + offset).clamp_(min=0) + + # -- token-level translate surface (the mixin / local-attn consumers) ------ + + def translate_full_attn_ids( + self, kv_indices: torch.Tensor, *, out: Optional[torch.Tensor] = None + ) -> torch.Tensor: + """Virtual token ids -> kernel-facing full-attention ids (the identity + when no translation is needed, so callers never branch).""" + if not self.is_translating: + assert out is None, "passthrough translate takes no out=" + return kv_indices + return self._translate_full(kv_indices, out=out) diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index abc2ae2dda36..4e78d79155d8 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -81,7 +81,10 @@ is_npu, next_power_of_2, ) -from sglang.srt.utils.async_probe import maybe_detect_oob +from sglang.srt.utils.async_probe import ( + maybe_detect_kernel_facing_loc, + maybe_detect_oob, +) from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter if TYPE_CHECKING: @@ -1593,21 +1596,22 @@ class KVWriteLoc: """Write target(s) for ``KVCache.set_kv_buffer``. All location info lives here (in the attention metadata), NOT in the pool: - - ``loc``: the generic per-token write location (the allocated - ``out_cache_loc``). VIRTUAL under the unified memory pool (it indexes the - virtual slot space); already physical for a non-unified memory pool. - - ``swa_loc``: the pre-translated SWA-sub-pool PHYSICAL location for hybrid - SWA pools (``None`` otherwise). - - ``full_loc``: the pre-translated full-attention-sub-pool PHYSICAL location - for the unified memory pool (``None`` otherwise), computed once per forward in - attention metadata (``ForwardMetadata.out_cache_loc_full_physical``). The - shared full pool writes it directly; the pool never translates (replacing - the former per-layer v2p gather / ``set_full_loc`` pin). + - ``loc``: the generic per-token write location (``out_cache_loc``). + KERNEL-FACING on every pool: physical by allocation on non-unified + pools, rebound at ForwardBatch construction (``rebind_write_loc``) on + the unified pool. + - ``swa_loc``: the pre-resolved SWA-sub-pool location for hybrid SWA pools + (``None`` otherwise). + - ``full_loc``: the full-attention-sub-pool location for the unified + memory pool (``None`` otherwise), carried in attention metadata + (``ForwardMetadata.out_cache_loc_full_physical``). Since the + construction-time rebind it is the SAME id space as ``loc``; the shared + full pool writes it directly and never translates. ``swa_loc`` and ``full_loc`` are the parallel pair (each a pre-resolved - PHYSICAL loc into its sub-pool, mirroring ``swa_kv_pool`` / ``full_kv_pool``); - ``loc`` is the generic, possibly-virtual fallback. Bundling them lets a - backend issue one ``set_kv_buffer`` call regardless of pool type. + loc into its sub-pool, mirroring ``swa_kv_pool`` / ``full_kv_pool``); + ``loc`` is the generic fallback. Bundling them lets a backend issue one + ``set_kv_buffer`` call regardless of pool type. """ loc: torch.Tensor @@ -1688,6 +1692,10 @@ def __init__( ): self.size = size self.page_size = page_size + # Row-blocks one page holds in this pool's kernel-facing id space; >1 + # only where the per-layer views are dense (the unified pool), and then + # a write loc must have been translated into that space first. + self.kernel_page_blocks = 1 self.dtype = dtype self.device = device if dtype in (torch.float8_e5m2, torch.float8_e4m3fn, torch.float8_e4m3fnuz): @@ -2387,6 +2395,9 @@ def set_kv_buffer( # Catch stale slot ids here instead of as illegal-addr / silent KV # corruption in the store_kvcache write (gated on SGLANG_ENABLE_ASYNC_ASSERT). maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA)") + maybe_detect_kernel_facing_loc( + loc, self.page_size, self.kernel_page_blocks, "set_kv_buffer (MHA)" + ) layer_id = ( layer_id_override if layer_id_override is not None else layer.layer_id ) @@ -3652,9 +3663,10 @@ def __init__( # virtual->physical mamba-slot translate for the HiCache offload path; # identity for a static pool, the allocator's `translate` for the unified pool. self._mamba_translate = lambda ids: ids - # virtual->kernel-facing full-KV translate for the model-level MLA entry points - # (`set_mla_kv_buffer` / `get_mla_kv_buffer` receive VIRTUAL locs); - # identity for a static pool, `translate_kv_loc_for_kernel` for the unified pool. + # The MLA doors take DIFFERENT id spaces: `get_mla_kv_buffer` gets + # ForwardBatch-built read indices (prefix_chunk_kv_indices / + # fetch_mha_one_shot_kv_indices), still VIRTUAL, so it translates; + # `set_mla_kv_buffer` gets out_cache_loc, already kernel-facing. self._full_translate = lambda ids: ids self.use_mla = use_mla if full_kv_pool is not None: @@ -3941,17 +3953,8 @@ def set_mla_kv_buffer( loc: torch.Tensor, cache_k_nope: torch.Tensor, cache_k_rope: torch.Tensor, - loc_is_kernel_facing: bool = False, ): assert self.use_mla, "set_mla_kv_buffer called when use_mla is False" - # Model-level MLA entry point: `loc` is a VIRTUAL loc under the unified - # pool, so translate to the kernel-facing id space here. - # - # `loc_is_kernel_facing`: the caller already translated `loc` (the unified-pool - # cuda-graph decode precomputes it out-of-graph into a capture-stable - # buffer, so the in-graph write does not capture a translate allocation). - if not loc_is_kernel_facing: - loc = self._full_translate(loc) with self._transfer_id_context(layer): self.full_kv_pool.set_mla_kv_buffer(layer, loc, cache_k_nope, cache_k_rope) @@ -4090,6 +4093,9 @@ def set_kv_buffer( ): loc, _, _ = unwrap_write_loc(loc_info) maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)") + maybe_detect_kernel_facing_loc( + loc, self.page_size, self.kernel_page_blocks, "set_kv_buffer (MLA)" + ) layer_id = ( layer_id_override if layer_id_override is not None else layer.layer_id ) @@ -4174,6 +4180,9 @@ def set_mla_kv_buffer( (self.size + self.page_size) * get_parallel().attn_dcp_size, "set_mla_kv_buffer (MLA)", ) + maybe_detect_kernel_facing_loc( + loc, self.page_size, self.kernel_page_blocks, "set_mla_kv_buffer (MLA)" + ) layer_id = ( layer_id_override if layer_id_override is not None else layer.layer_id ) diff --git a/python/sglang/srt/mem_cache/multi_ended_allocator.py b/python/sglang/srt/mem_cache/multi_ended_allocator.py index a0674e7574b7..93a6c39e22ed 100644 --- a/python/sglang/srt/mem_cache/multi_ended_allocator.py +++ b/python/sglang/srt/mem_cache/multi_ended_allocator.py @@ -1906,6 +1906,11 @@ def full_v2p_page_table(self) -> torch.Tensor: """ return self.full_attn_allocator.virtual_to_physical + @property + def full_p2v_page_table(self) -> torch.Tensor: + """Page-level physical->virtual table of the full sub-pool.""" + return self.full_attn_allocator.physical_to_virtual + def translate_kv_loc_for_kernel( self, loc: torch.Tensor, @@ -2257,6 +2262,11 @@ def full_v2p_page_table(self) -> torch.Tensor: """Page-level virtual->physical table of the full sub-pool.""" return self.full_attn_allocator.virtual_to_physical + @property + def full_p2v_page_table(self) -> torch.Tensor: + """Page-level physical->virtual table of the full sub-pool.""" + return self.full_attn_allocator.physical_to_virtual + def translate_kv_loc_for_kernel( self, loc: torch.Tensor, diff --git a/python/sglang/srt/mem_cache/unified_memory_pool.py b/python/sglang/srt/mem_cache/unified_memory_pool.py index f974b1f96d70..72daecffce48 100644 --- a/python/sglang/srt/mem_cache/unified_memory_pool.py +++ b/python/sglang/srt/mem_cache/unified_memory_pool.py @@ -548,6 +548,7 @@ def __init__( enable_kv_cache_copy=False, kv_cache_layout="page_major", ) + self.kernel_page_blocks = spec.blocks_per_page() def _create_buffers(self): self.k_buffer = self._k_views @@ -641,7 +642,7 @@ def __init__( max_slots = unified_buffer.max_slots(sub_pool_name) self._num_pages = max_slots // page_size self._page_bytes = page_size * spec.entry_bytes() - self._view_rows = self._num_pages * spec.layer_num * page_size + self._view_rows = self._num_pages * spec.blocks_per_page() * page_size super().__init__( # OOB checks bound locs by `size + page_size`; kernel-facing ids run to @@ -655,6 +656,7 @@ def __init__( device=unified_buffer.device, enable_memory_saver=False, # buffer owned by UnifiedKVPool ) + self.kernel_page_blocks = spec.blocks_per_page() def _create_buffers(self): self.kv_buffer = self._kv_views @@ -1243,9 +1245,6 @@ def init_unified_mamba_pools( req_to_token_pool.mamba_allocator = mamba_slot_allocator token_to_kv_pool._mamba_translate = mamba_slot_allocator.translate if use_mla_backend: - # Model-level MLA entry points (`set_mla_kv_buffer` / `get_mla_kv_buffer`) - # receive VIRTUAL locs and translate to the kernel-facing space internally - # (eager-prefill-only paths; never captured in a cuda graph). token_to_kv_pool._full_translate = allocator.translate_kv_loc_for_kernel logger.info( diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 7c504656fdaf..3926aa58c3fd 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -827,6 +827,8 @@ def init_new( device = model_runner.device + model_runner.kv_index_translator.rebind_write_loc(ret) + if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get(): hashed = _hash_rids_to_tensor( rids=[req.rid for req in batch.reqs], diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index b878d07cffbe..064ff0dafb0c 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -89,6 +89,7 @@ from sglang.srt.mem_cache.kv_cache_configurator import ( KVCacheConfigurator, ) +from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.model_executor.cuda_graph_config import ( cuda_graph_fully_disabled, @@ -856,6 +857,18 @@ def account_preloaded_weights(self, preloaded_weights_bytes: int) -> None: return self.pre_model_load_memory += preloaded_weights_bytes / (1 << 30) + def init_kv_index_translator(self): + """The one object that converts KV ids for this runner: attention + backends build their read indices from the table it hands them instead + of probing the pool's id spaces themselves.""" + self.kv_index_translator = KVIndexTranslator( + req_to_token=self.req_to_token_pool.req_to_token, + token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, + token_to_kv_pool=self.token_to_kv_pool, + page_size=self.page_size or 1, + device=self.device, + ) + def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = None): """Allocate KV cache memory pools only (no backends or cuda graphs).""" if memory_pool_config is not None: @@ -882,6 +895,8 @@ def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = Non def _init_post_memory_pool_components(self): """Post-pool component wiring, split out of alloc_memory_pool so forks that build bespoke memory pools can reuse it after allocating them.""" + self.init_kv_index_translator() + # Must be called AFTER init_memory_pool so the pool object exists for # canary to monkey-patch, and BEFORE init_decode_cuda_graph so warmup # forwards captured into the graph see the patched pool methods. diff --git a/python/sglang/srt/utils/async_probe.py b/python/sglang/srt/utils/async_probe.py index abe25cd9a0e5..be514dc4e5fc 100644 --- a/python/sglang/srt/utils/async_probe.py +++ b/python/sglang/srt/utils/async_probe.py @@ -141,6 +141,27 @@ def maybe_detect_oob(indices: Optional[torch.Tensor], low: int, high: int, msg: ) +def maybe_detect_kernel_facing_loc( + indices: Optional[torch.Tensor], page_size: int, blocks_per_page: int, msg: str +): + """Async check that a write loc is in the pool's KERNEL-FACING id space. + + A kernel-facing id is `phys_page * (page_size * blocks_per_page) + offset` + with `offset < page_size`, so its remainder modulo the page stride is + below page_size; a VIRTUAL id satisfies that only in the first block. + Vacuous at blocks_per_page 1. Virtual ids are in range for the OOB probe, + so this is the only check that separates them. + """ + if blocks_per_page <= 1 or not envs.SGLANG_ENABLE_ASYNC_ASSERT.get(): + return + if indices is None or indices.numel() == 0: + return + torch._assert_async( + (indices % (page_size * blocks_per_page) < page_size).all(), + f"write loc outside the kernel-facing id space (virtual ids?): {msg}", + ) + + def maybe_detect_page_aligned( indices: Optional[torch.Tensor], page_size: int, msg: str ): diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py index 7ca676e61b2f..b352cd7c29ce 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py @@ -411,6 +411,7 @@ def __init__( page_size=case.page_size, get_kvcache=lambda: self.token_to_kv_pool, ) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py index 3d05a742343b..7cc4c340a540 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py @@ -404,6 +404,7 @@ def __init__( kv_cache_dim=pool_kv_cache_dim, ) self.token_to_kv_pool_allocator = SimpleNamespace(page_size=case.page_size) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py index 2e6831e72d97..3a62e1031c4f 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py @@ -384,6 +384,7 @@ def __init__( enable_alt_stream=False, ) self.token_to_kv_pool_allocator = SimpleNamespace(page_size=case.page_size) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py index 9ef502e9452e..556bc92931f2 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py @@ -318,6 +318,7 @@ def __init__( page_size=case.page_size, get_kvcache=lambda: self.token_to_kv_pool, ) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py index 899b2dbce927..00ddbc0a5321 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py @@ -317,6 +317,7 @@ def __init__( enable_alt_stream=False, ) self.token_to_kv_pool_allocator = SimpleNamespace(page_size=case.page_size) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py index 89b104f6747f..4c2cf225f7b1 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py @@ -326,6 +326,7 @@ def __init__( enable_alt_stream=False, ) self.token_to_kv_pool_allocator = SimpleNamespace(page_size=case.page_size) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index f46459a5048c..65b4ebe9b7ba 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -453,6 +453,7 @@ def __init__( enable_alt_stream=False, ) self.token_to_kv_pool_allocator = SimpleNamespace(page_size=case.page_size) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py index 21cfbdcd28a8..89cc7742f323 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py @@ -310,6 +310,7 @@ def __init__( enable_memory_saver=False, ) self.token_to_kv_pool_allocator = SimpleNamespace(page_size=case.page_size) + self.init_kv_index_translator() self.attn_cp_size = 1 self.attention_chunk_size = None self.hisparse_coordinator = None diff --git a/test/registered/attention/test_create_kvindices.py b/test/registered/attention/test_create_kvindices.py index 907a12f865a5..7e749301f0da 100644 --- a/test/registered/attention/test_create_kvindices.py +++ b/test/registered/attention/test_create_kvindices.py @@ -75,6 +75,60 @@ def test_create_kvindices(self): for batch in BATCH: self._run_test(batch, MAX_BATCH, MAX_CONTEXT_LEN) + def _run_page_table_test(self, batch, ps, with_window_start): + """ENTRY_PAGE_SIZE > 1: the source is a PAGE-granular table (the unified + pool's read table); the kernel must reconstruct token ids by the affine + rule token = entry * ps + pos % ps -- including the kv_start_idx + (sliding-window) offset path, whose pos is an absolute token position.""" + max_batch, max_pages = 64, 128 + page_table = torch.randint( + 0, 1 << 20, (max_batch, max_pages), dtype=torch.int32 + ) + req_pool_indices = torch.tensor( + np.random.choice(range(max_batch), size=batch, replace=False), + dtype=torch.int32, + ) + lens = torch.tensor( + np.random.randint(1, max_pages * ps, size=batch), dtype=torch.int32 + ) + if with_window_start: + start = torch.clamp( + lens - torch.randint(1, ps * 3, (batch,), dtype=torch.int32), min=0 + ) + gather_lens = lens - start + else: + start, gather_lens = None, lens + kv_indptr = torch.zeros((batch + 1,), dtype=torch.int32) + kv_indptr[1:] = torch.cumsum(gather_lens, dim=0) + + # ref: absolute positions [start, start+len) through the affine rule + refs = [] + for i in range(batch): + s = int(start[i]) if start is not None else 0 + pos = torch.arange(s, s + int(gather_lens[i]), dtype=torch.int64) + entry = page_table[int(req_pool_indices[i])][pos // ps].to(torch.int64) + refs.append(entry * ps + pos % ps) + ref = torch.cat(refs).contiguous() + + out = torch.empty(int(kv_indptr[-1]), dtype=torch.int64) + create_flashinfer_kv_indices_triton[(batch,)]( + page_table, + req_pool_indices, + gather_lens, + kv_indptr, + start, + out, + page_table.size(1), + ENTRY_PAGE_SIZE=ps, + ) + self.assertTrue(torch.equal(ref, out)) + + def test_page_table_source_reconstruction(self): + for batch in (1, 37): + for ps in (4, 64, 256): + self._run_page_table_test(batch, ps, with_window_start=False) + self._run_page_table_test(batch, ps, with_window_start=True) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_full_loc_fast_path.py b/test/registered/unit/mem_cache/test_full_loc_fast_path.py index 208c213ac55d..2a3a47d0ef05 100644 --- a/test/registered/unit/mem_cache/test_full_loc_fast_path.py +++ b/test/registered/unit/mem_cache/test_full_loc_fast_path.py @@ -19,7 +19,7 @@ always PHYSICAL. Two routing contracts are pinned here: 1. Full-attention. The full-physical loc is carried in `KVWriteLoc.full_loc` - (from `ForwardBatch.out_cache_loc_full_physical`) and written directly. + (from `ForwardMetadata.out_cache_loc_full_physical`) and written directly. `UnifiedSWAKVPool` asserts it's present (the unified memory pool always precomputes it); `HybridLinearKVPool` falls back to `loc` for a static (non-shared) pool, where `loc` is itself already physical. @@ -264,8 +264,10 @@ class TestHybridLinearMLARouting(unittest.TestCase): - `set_kv_buffer` (MLA branch) mirrors the MHA branch — write the pre-translated `KVWriteLoc.full_loc` when present (unified pool, where it carries the DENSE loc), else the raw `loc` (static pool, already physical). - - `set_mla_kv_buffer` / `get_mla_kv_buffer` receive VIRTUAL locs and apply - `_full_translate` exactly once (identity for a static pool).""" + - `set_mla_kv_buffer` forwards `loc` untouched (kernel-facing since the + ForwardBatch rebind); `get_mla_kv_buffer` applies `_full_translate` + exactly once (its indices are req_to_token-produced, virtual under the + unified pool).""" def _make_bare_pool(self, translate=None): from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool @@ -274,7 +276,7 @@ def _make_bare_pool(self, translate=None): pool.full_kv_pool = _RecordingMLAPool() pool.use_mla = True pool.full_attention_layer_id_mapping = {0: 0} - pool._full_translate = translate if translate is not None else (lambda x: x) + pool._full_translate = translate if translate is not None else (lambda ids: ids) return pool def test_mla_writes_full_loc_from_write_loc(self): @@ -311,28 +313,26 @@ def test_mla_falls_back_to_loc_when_absent(self): forwarded, _ = pool.full_kv_pool.calls[0] self.assertIs(forwarded, phys_loc) - def test_set_mla_kv_buffer_translates_exactly_once(self): - calls = [] - - def translate(ids): - calls.append(ids) - return ids + 100 - - pool = self._make_bare_pool(translate=translate) - virtual_loc = torch.tensor([7, 8, 9], dtype=torch.int64) + def test_set_mla_kv_buffer_door_never_translates(self): + """Physical-loc contract: the write door forwards `loc` UNTOUCHED. + The translate happens exactly once at ForwardBatch construction + (rebind_write_loc, kernel-facing-first); a door that translated + again would double-translate every unified MLA write. Deleting the + forward (or re-adding a door translate) turns this red.""" + pool = self._make_bare_pool() + loc = torch.tensor([107, 108, 109], dtype=torch.int64) layer = types.SimpleNamespace(layer_id=0) - pool.set_mla_kv_buffer( - layer, virtual_loc, torch.zeros(3, 1, 6), torch.zeros(3, 1, 2) - ) + pool.set_mla_kv_buffer(layer, loc, torch.zeros(3, 1, 6), torch.zeros(3, 1, 2)) - self.assertEqual(len(calls), 1) self.assertEqual(len(pool.full_kv_pool.mla_set_calls), 1) - self.assertTrue( - torch.all(pool.full_kv_pool.mla_set_calls[0] == virtual_loc + 100) - ) + self.assertIs(pool.full_kv_pool.mla_set_calls[0], loc) def test_get_mla_kv_buffer_translates_exactly_once(self): + """READ door: `loc` is produced from req_to_token (VIRTUAL under the + unified pool), so the get side still translates here — exactly once. + The WRITE door (case above) never translates: the split is the write + flip's contract.""" calls = [] def translate(ids): diff --git a/test/registered/unit/mem_cache/test_kv_index_translator.py b/test/registered/unit/mem_cache/test_kv_index_translator.py new file mode 100644 index 000000000000..a69a217e6228 --- /dev/null +++ b/test/registered/unit/mem_cache/test_kv_index_translator.py @@ -0,0 +1,684 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""KVIndexTranslator -- the read-path id translator. + +Covers, CPU-only (the builder's pure-torch reference path; GPU parity of the +Triton kernel is a later CUDA CI pin): + - strict passthrough: a non-unified source returns the SAME req_to_token / + req_pool_indices objects -- zero tensor ops, no copies (the property that + makes backend re-pointing byte-identical for every non-unified server); + - static SWA pools keep their legacy full->swa mapping on the view; + - the read table matches the hand formula + entry[b, c] = clamp(v2p[req_to_token[req[b], c*ps] // ps] * mult, 0) + over the REAL SWA composite's tables (full AND swa, ps in {1, 4}, + multiplier in {1, 2L}), with the swa table built from VIRTUAL ids; + - sink routing: dead lanes (seq_len 0), -1 req_to_token entries, and + tombstoned v2p pages all read entry 0; + - the capture contract: buffers are zero-filled and idempotent; a refresh + updates ONLY the live prefix (stale tails and rows beyond bs keep prior + contents); the returned table is the WHOLE buffer (pointer-stable); + - the eager-view memo: a single source-resident slot keyed by batch + identity (same batch shares one build; the next batch replaces it; a + dead batch never matches); + - the two-phase write contract: the rebind touches only the full side, and + the sliding-window write loc derives POINTWISE from the kernel-facing values + (pads, slices, and fresh copies included), for both pool families. + + python -m pytest test/registered/unit/mem_cache/test_kv_index_translator.py -v +""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=8, suite="base-a-test-cpu") + +import unittest +from types import SimpleNamespace + +import torch +from test_multi_ended_allocator import _FakeUnifiedSWAKVPool + +from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator +from sglang.srt.mem_cache.multi_ended_allocator import ( + UnifiedSWATokenToKVPoolAllocator, +) +from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool +from sglang.srt.mem_cache.unified_memory_pool import MHASubPoolSpec, UnifiedKVPool + +_DEV = "cpu" +_FULL_L = 2 +_SWA_L = 3 + + +def _build_composite(ps, collapse=False, n_full_pages=16, n_swa_pages=8): + full_spec = MHASubPoolSpec( + name="full", + layer_num=_FULL_L, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="up", + ) + swa_spec = MHASubPoolSpec( + name="swa", + layer_num=_SWA_L, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="down", + ) + n_full, n_swa = n_full_pages * ps, n_swa_pages * ps + total = n_full * full_spec.entry_bytes() + n_swa * swa_spec.entry_bytes() + pool = UnifiedKVPool( + total_bytes=total, + sub_pool_specs=[full_spec, swa_spec], + device=_DEV, + enable_memory_saver=False, + page_size=ps, + ) + kvcache = _FakeUnifiedSWAKVPool(pool) + allocator = UnifiedSWATokenToKVPoolAllocator( + unified_buffer=pool, + kvcache=kvcache, + device=_DEV, + full_max_total_num_tokens=n_full, + swa_max_total_num_tokens=n_swa, + page_size=ps, + need_sort=False, + forward_stream=None, + ) + if collapse: + # The multiplier-1 arm, where kernel-facing ids ARE the physical ones. + # No unified sub-pool reports 1 today, so pin the regime here. + allocator.full_attn_allocator.kernel_page_multiplier = 1 + allocator.swa_attn_allocator.kernel_page_multiplier = 1 + # The fake IS the runner's token_to_kv_pool, and the real UnifiedSWAKVPool + # carries the pool-level full->swa translate, so the fake must too. + kvcache.translate_loc_from_full_to_swa = allocator.translate_loc_from_full_to_swa + return allocator + + +def _make_source(allocator, req_to_token, ps): + """The owning runner's source: its token_to_kv_pool IS the allocator's own + kvcache. (A runner can share the allocator while owning a different pool -- + see TestPoolOwnership.)""" + return KVIndexTranslator( + req_to_token=req_to_token, + token_to_kv_pool_allocator=allocator, + token_to_kv_pool=allocator.get_kvcache(), + page_size=ps, + device=_DEV, + ) + + +def _reference_table(req_to_token, req_pool_indices, seq_lens, v2p, mult, ps, width): + """Independent python derivation of the read-table formula.""" + bs = req_pool_indices.numel() + out = torch.zeros((bs, width), dtype=torch.int32) + for b in range(bs): + req = int(req_pool_indices[b]) + n_pages = -(-int(seq_lens[b]) // ps) + for c in range(min(n_pages, width)): + tok = int(req_to_token[req, c * ps]) + page = 0 if tok < 0 else tok // ps + out[b, c] = max(int(v2p[page]) * mult, 0) + return out + + +class TestPassthrough(unittest.TestCase): + def test_non_unified_returns_same_objects(self): + """The strict-passthrough property: no copy, no branch, the exact + tensors backends read today. A regression here (any tensor op on the + non-unified path) breaks byte-identity for every static-pool server.""" + req_to_token = torch.arange(64, dtype=torch.int64).reshape(4, 16) + src = KVIndexTranslator( + req_to_token=req_to_token, + token_to_kv_pool_allocator=SimpleNamespace(), # not a composite + token_to_kv_pool=SimpleNamespace(), # not an SWAKVPool + page_size=1, + device=_DEV, + ) + self.assertFalse(src.is_translating) + rows = torch.tensor([2, 0]) + view = src.build_index_table( + req_pool_indices=rows, seq_lens=torch.tensor([5, 3]) + ) + self.assertIs(view.ids, req_to_token) + self.assertIs(view.row_ids, rows) + self.assertEqual(view.row_stride, req_to_token.stride(0)) + self.assertEqual(view.entry_page_size, 1) + self.assertFalse(view.is_translated) + self.assertIsNone(view.sliding_window_ids) + # And the translate surface is the identity, not a wrapped copy. + t = torch.tensor([1, 2, 3]) + self.assertIs(src.translate_full_attn_ids(t), t) + + +def _alloc_and_fill(allocator, ps, lens): + """Allocate per-request virtual runs and write them into a fake + req_to_token; returns (req_to_token, req_pool_indices, seq_lens).""" + width = 16 * ps + req_to_token = torch.full((len(lens), width), -1, dtype=torch.int64) + for r, n in enumerate(lens): + n_alloc = -(-n // ps) * ps # page-aligned virtual run + v = allocator.alloc(n_alloc) + assert v is not None + req_to_token[r, :n] = v[:n] + return ( + req_to_token, + torch.arange(len(lens), dtype=torch.int64), + torch.tensor(lens, dtype=torch.int64), + ) + + +class TestReadTableBuild(unittest.TestCase): + + def test_read_table_matches_reference_dense_and_strided(self): + """The load-bearing formula pin: full AND swa read tables equal + the independent per-element derivation, across page sizes and both + multiplier regimes (strided=1, dense=2L). The swa table agreeing with + a formula over VIRTUAL ids is also the never-chained-through- + full-physical proof.""" + for ps in (1, 4): + for collapse in (True, False): + allocator = _build_composite(ps, collapse=collapse) + full_mult = allocator.kernel_page_multiplier + swa_mult = allocator.swa_kernel_page_multiplier + req_to_token, rows, seq_lens = _alloc_and_fill( + allocator, ps, lens=[5 * ps, 2 * ps, 3 * ps - 1] + ) + src = _make_source(allocator, req_to_token, ps) + self.assertTrue(src.is_translating) + width = 6 + view = src.build_index_table( + req_pool_indices=rows, seq_lens=seq_lens, max_pages=width + ) + self.assertTrue(view.is_translated) + self.assertEqual(view.entry_page_size, ps) + self.assertTrue( + torch.equal(view.row_ids, torch.arange(3, dtype=torch.int64)) + ) + want_full = _reference_table( + req_to_token, + rows, + seq_lens, + allocator.full_v2p_page_table, + full_mult, + ps, + width, + ) + want_swa = _reference_table( + req_to_token, + rows, + seq_lens, + allocator.swa_v2p_page_table, + swa_mult, + ps, + width, + ) + self.assertTrue( + torch.equal(view.ids, want_full), + f"full read table off-formula (ps={ps}, mult={full_mult})", + ) + self.assertTrue( + torch.equal(view.sliding_window_ids, want_swa), + f"swa read table off-formula (ps={ps}, mult={swa_mult})", + ) + + def test_sink_routing(self): + """Dead lanes (seq_len 0), -1 slots inside the live prefix, and + tombstoned v2p pages must ALL read entry 0 -- one wild entry is a + captured-graph OOB read at replay.""" + ps = 4 + allocator = _build_composite(ps) + req_to_token, rows, seq_lens = _alloc_and_fill( + allocator, ps, lens=[3 * ps, 2 * ps, ps] + ) + seq_lens[1] = 0 # dead lane + req_to_token[0, ps] = -1 # unwritten slot inside the live prefix + # Tombstone row 2's first page on BOTH sides. + tomb_page = int(req_to_token[2, 0]) // ps + allocator.full_v2p_page_table[tomb_page] = -1 + allocator.swa_v2p_page_table[tomb_page] = -1 + src = _make_source(allocator, req_to_token, ps) + view = src.build_index_table( + req_pool_indices=rows, seq_lens=seq_lens, max_pages=4 + ) + for table in (view.ids, view.sliding_window_ids): + self.assertTrue(bool((table >= 0).all())) + self.assertTrue(bool((table[1] == 0).all()), "dead lane not sunk") + self.assertEqual(int(table[0, 1]), 0, "-1 slot not sunk") + self.assertEqual(int(table[2, 0]), 0, "tombstone not sunk") + + +class TestBuildInto(unittest.TestCase): + """fill_read_table fills a backend-owned padded block table's live prefix with + FULL-side read-table entries -- the trtllm_mla / flashmla consumption route + (their rows ARE the read table's rows).""" + + def test_prefix_filled_tail_sentinel_preserved_width_capped(self): + """Three contracts in one batch: entries equal the read-table formula, + lanes past each row's live pages keep the backend's -1 sentinel + (prefix-only -- a tail write scatters the trtllm sentinel contract), + and a table padded WIDER than the req_to_token page span (trtllm's + LCM alignment) is capped instead of tripping the builder's width + assert.""" + ps = 4 + allocator = _build_composite(ps) + full_mult = allocator.kernel_page_multiplier + lens = [5, 2 * ps + 1, 1] + req_to_token, rows, seq_lens = _alloc_and_fill(allocator, ps, lens=lens) + src = _make_source(allocator, req_to_token, ps) + self.assertTrue(src.is_translating) + + width_pages = req_to_token.shape[1] // ps + 3 # wider than the span + out = torch.full((len(lens), width_pages), -1, dtype=torch.int32) + src.fill_read_table(out=out, req_pool_indices=rows, seq_lens=seq_lens) + + want = _reference_table( + req_to_token, + rows, + seq_lens, + allocator.full_v2p_page_table, + full_mult, + ps, + width_pages, + ) + for b, n in enumerate(lens): + n_pages = -(-n // ps) + self.assertTrue( + torch.equal(out[b, :n_pages], want[b, :n_pages]), + f"row {b} live prefix off-formula", + ) + self.assertTrue( + bool((out[b, n_pages:] == -1).all()), + f"row {b} tail sentinel clobbered", + ) + + def test_passthrough_source_refuses(self): + """Callers dispatch on `enabled`; a passthrough source has no v2p to + build from and must fail loud, not fill garbage.""" + src = KVIndexTranslator( + req_to_token=torch.zeros((2, 4), dtype=torch.int64), + token_to_kv_pool_allocator=SimpleNamespace(), + token_to_kv_pool=SimpleNamespace(), + page_size=1, + device=_DEV, + ) + with self.assertRaises(AssertionError): + src.fill_read_table( + out=torch.zeros((1, 4), dtype=torch.int32), + req_pool_indices=torch.tensor([0]), + seq_lens=torch.tensor([1]), + ) + + +class TestPoolOwnership(unittest.TestCase): + """A runner only gets the kernel-facing id space when the pool IT reads and + writes is the one the allocator's ids address. + + Guarded shape: a runner handed a SHARED allocator (one slot index space, + one req_to_token) while owning a SEPARATE KV buffer sized to the + allocator's SLOT count. Probing the allocator alone reports "unified" for + that runner, so its indices would be mapped into the composite's + kernel-facing space (kernel-facing ids up to num_pages * multiplier) and then used + to address a buffer with only num_slots rows -- out of bounds on both the + read gather and the KV store. + """ + + def test_real_factory_bundle_satisfies_the_ownership_identity(self): + """The guard rests on `allocator.get_kvcache() is token_to_kv_pool` + holding for a REAL target bundle. If a factory ever returned a pool + the allocator does not hold, the guard would silently disable the + unified path for EVERY model -- so pin it against the real factory + rather than against this file's own construction.""" + from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools + + bundle = init_unified_swa_pools( + device="cpu", + kv_cache_dtype=torch.float16, + head_num=2, + head_dim=8, + v_head_dim=8, + swa_head_num=2, + swa_head_dim=8, + swa_v_head_dim=8, + page_size=1, + start_layer=0, + end_layer=4, + swa_attention_layer_ids=[1, 3], + full_attention_layer_ids=[0, 2], + full_max_total_num_tokens=64, + swa_max_total_num_tokens=32, + enable_memory_saver=False, + need_sort=False, + ) + self.assertIs( + bundle.token_to_kv_pool_allocator.get_kvcache(), + bundle.token_to_kv_pool, + ) + src = KVIndexTranslator( + req_to_token=torch.zeros((2, 8), dtype=torch.int32, device=_DEV), + token_to_kv_pool_allocator=bundle.token_to_kv_pool_allocator, + token_to_kv_pool=bundle.token_to_kv_pool, + page_size=1, + device=_DEV, + ) + self.assertTrue(src.is_translating) + + def test_runner_with_its_own_pool_is_disabled(self): + """Same allocator, different pool: must stay disabled.""" + alloc = _build_composite(ps=1) + req_to_token = torch.zeros((2, 8), dtype=torch.int32, device=_DEV) + own_pool = SimpleNamespace() # a separate buffer, not the composite's + src = KVIndexTranslator( + req_to_token=req_to_token, + token_to_kv_pool_allocator=alloc, + token_to_kv_pool=own_pool, + page_size=1, + device=_DEV, + ) + self.assertFalse(src.is_translating) + + def test_disabled_source_is_strict_passthrough(self): + """Consequence of the guard: such a runner must see RAW virtual ids on + the read table -- they index its own pool directly. A translate here is + the out-of-bounds bug the ownership identity exists to prevent.""" + alloc = _build_composite(ps=1) + req_to_token = torch.arange(16, dtype=torch.int32, device=_DEV).view(2, 8) + src = KVIndexTranslator( + req_to_token=req_to_token, + token_to_kv_pool_allocator=alloc, + token_to_kv_pool=SimpleNamespace(), + page_size=1, + device=_DEV, + ) + rows = torch.tensor([1, 0], dtype=torch.int32, device=_DEV) + view = src.build_index_table( + req_pool_indices=rows, + seq_lens=torch.tensor([3, 2], dtype=torch.int32, device=_DEV), + ) + # Read table: the EXACT objects a static-pool backend reads today. + self.assertIs(view.ids, req_to_token) + self.assertIs(view.row_ids, rows) + self.assertFalse(view.is_translated) + # And the token-level surface is the identity, same guard. + t = torch.tensor([5, 6], dtype=torch.int64, device=_DEV) + self.assertIs(src.translate_full_attn_ids(t), t) + + +class TestCaptureContract(unittest.TestCase): + def test_caller_owned_table_is_returned_whole_and_filled_prefix_only(self): + ps = 4 + allocator = _build_composite(ps) + req_to_token = torch.full((4, 16 * ps), -1, dtype=torch.int64) + v = allocator.alloc(2 * ps) + req_to_token[1, : 2 * ps] = v + src = _make_source(allocator, req_to_token, ps) + + tables = src.make_capture_tables(max_bs=4, max_context_len=8 * ps) + cap, cap_swa = tables.full, tables.sliding_window + self.assertTrue(bool((cap == 0).all()), "read tables must start zeroed") + self.assertIsNotNone(cap_swa, "the SWA composite has a second id space") + + # Poison everything, then refresh a 1-row batch: ONLY its live prefix + # may change -- stale tails and other rows are the fa3 contract. + cap.fill_(7) + cap_swa.fill_(7) + view = src.build_index_table( + req_pool_indices=torch.tensor([1]), + seq_lens=torch.tensor([2 * ps]), + into=tables, + ) + self.assertIs(view.ids, cap, "the caller's table comes back WHOLE") + want = allocator.full_v2p_page_table[req_to_token[1, ::ps][:2] // ps] * ( + 2 * _FULL_L + ) + self.assertTrue(torch.equal(cap[0, :2], want.to(torch.int32))) + self.assertTrue(bool((cap[0, 2:] == 7).all()), "stale tail was cleared") + self.assertTrue(bool((cap[1:] == 7).all()), "rows beyond bs were touched") + + def test_row_ids_not_reallocated_across_builds(self): + """`row_ids` is a constant arange sized once from the request pool, so + builds at different batch sizes hand back slices of ONE buffer. A + per-build `torch.arange` would be correct but would spend an allocation + and a launch on every replay prep.""" + allocator = _build_composite(1) + req_to_token, rows, seq_lens = _alloc_and_fill(allocator, 1, lens=[4, 2, 3]) + src = _make_source(allocator, req_to_token, 1) + self.assertTrue(src.is_translating) + first = src.build_index_table( + req_pool_indices=rows[:2], seq_lens=seq_lens[:2], max_pages=4 + ) + second = src.build_index_table( + req_pool_indices=rows, seq_lens=seq_lens, max_pages=4 + ) + self.assertEqual(first.row_ids.data_ptr(), second.row_ids.data_ptr()) + self.assertTrue(torch.equal(first.row_ids, torch.arange(2, device=_DEV))) + self.assertTrue(torch.equal(second.row_ids, torch.arange(3, device=_DEV))) + # Sized to bound any batch the request pool can hold. + self.assertGreaterEqual(src._rows.numel(), req_to_token.shape[0]) + + +class _FakeForwardBatch: + """Weakref-able stand-in (SimpleNamespace is not) carrying the fields + `index_table_for_batch` and `rebind_write_loc` read. `seq_lens_sum` + defaults to the real sum: it is the signal that the CPU mirror is live, + and a real ForwardBatch always carries it (None only when gpu_only).""" + + def __init__( + self, + *, + req_pool_indices=None, + seq_lens=None, + seq_lens_cpu=None, + out_cache_loc=None, + seq_lens_sum=-1, + ): + self.req_pool_indices = req_pool_indices + self.seq_lens = seq_lens + self.seq_lens_cpu = seq_lens_cpu + self.out_cache_loc = out_cache_loc + self.seq_lens_sum = ( + (None if seq_lens is None else int(seq_lens.sum())) + if seq_lens_sum == -1 + else seq_lens_sum + ) + + +class TestViewMemo(unittest.TestCase): + """The eager view is memoized ON THE SOURCE in a single slot keyed by + batch identity -- per-batch state stays out of the ForwardBatch (it does + not scale with the number of id spaces), and one metadata build's many + consumers still share one table build.""" + + def _fb(self, allocator, ps, lens): + req_to_token, rows, seq_lens = _alloc_and_fill(allocator, ps, lens=lens) + fb = _FakeForwardBatch( + req_pool_indices=rows, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens, + ) + return fb, req_to_token + + def test_same_batch_returns_the_memoized_view(self): + ps = 1 + allocator = _build_composite(ps) + fb, req_to_token = self._fb(allocator, ps, lens=[3, 2]) + src = _make_source(allocator, req_to_token, ps) + v1 = src.index_table_for_batch(fb) + v2 = src.index_table_for_batch(fb) + self.assertIs(v1, v2) + + def test_next_batch_replaces_the_single_slot(self): + ps = 1 + allocator = _build_composite(ps) + fb1, req_to_token = self._fb(allocator, ps, lens=[3, 2]) + src = _make_source(allocator, req_to_token, ps) + v1 = src.index_table_for_batch(fb1) + fb2 = _FakeForwardBatch( + req_pool_indices=fb1.req_pool_indices, + seq_lens=fb1.seq_lens, + seq_lens_cpu=fb1.seq_lens_cpu, + ) + v2 = src.index_table_for_batch(fb2) + self.assertIsNot(v1, v2) + # Single slot: fb1 no longer matches and rebuilds. + v1_again = src.index_table_for_batch(fb1) + self.assertIsNot(v1_again, v1) + + def test_dead_batch_never_matches(self): + """A garbage-collected batch's slot must not serve a later batch: the + weakref key goes dead and the build runs fresh.""" + import gc + + ps = 1 + allocator = _build_composite(ps) + fb1, req_to_token = self._fb(allocator, ps, lens=[3, 2]) + src = _make_source(allocator, req_to_token, ps) + v1 = src.index_table_for_batch(fb1) + del fb1 + gc.collect() + fb2, _ = self._fb(allocator, ps, lens=[2]) + v2 = src.index_table_for_batch(fb2) + self.assertIsNot(v2, v1) + self.assertEqual(v2.ids.shape[0], 1) + + +class TestWriteLoc(unittest.TestCase): + """The two-phase write contract: phase 1 (`rebind_write_loc`) rebinds the + full side once at ForwardBatch construction; phase 2 derives the + sliding-window write loc on demand, POINTWISE from the full-side + values. Value-based derivation is the property under test: pads, slices, + and fresh copies of the loc must all derive correctly with no handover + and no stored per-forward state.""" + + def _built(self, ps=1, n=4): + allocator = _build_composite(ps) + req_to_token, rows, seq_lens = _alloc_and_fill(allocator, ps, lens=[max(n, 1)]) + src = _make_source(allocator, req_to_token, ps) + virt = allocator.alloc(-(-n // ps) * ps)[:n] + want_full = allocator.translate_kv_loc_for_kernel(virt) + want_swa = allocator.translate_loc_from_full_to_swa(virt) + return src, allocator, rows, seq_lens, virt, want_full, want_swa + + def _field(self, src, rows, seq_lens, kernel_loc): + return src.sliding_window_write_loc_for(kernel_loc) + + def test_rebind_translates_full_side_only(self): + for ps in (1, 4): + src, _, _, _, virt, want_full, _ = self._built(ps=ps, n=3 * ps) + keep = virt.clone() + fb = _FakeForwardBatch(out_cache_loc=virt) + src.rebind_write_loc(fb) + # Full side: rebound to a FRESH kernel-facing tensor; the + # ScheduleBatch's aliased virtual tensor is untouched. + self.assertIsNot(fb.out_cache_loc, virt) + self.assertTrue(torch.equal(fb.out_cache_loc, want_full)) + self.assertTrue(torch.equal(virt, keep)) + + def test_swa_write_loc_round_trips_from_dense(self): + """The derived property behind phase 2: for any virtual run t, + deriving from the dense full-side values must equal the direct + virtual->swa translate — `field(full(t)) == swa(t)` across page sizes + and multipliers.""" + for ps in (1, 4, 64): + src, _, rows, seq_lens, _, want_full, want_swa = self._built( + ps=ps, n=3 * ps + ) + got = self._field(src, rows, seq_lens, want_full) + self.assertTrue(torch.equal(got, want_swa)) + + def test_pad_lanes_derive_to_sink(self): + """The DP pad appends zeros; dense 0 is the reserved padding slot in + every id space, so pad lanes must derive to swa slot 0 with no + `num_live` bookkeeping.""" + src, _, rows, seq_lens, _, want_full, want_swa = self._built(n=3) + padded = torch.cat([want_full, want_full.new_zeros(2)]) + got = self._field(src, rows, seq_lens, padded) + self.assertTrue(torch.equal(got[:3], want_swa)) + self.assertTrue(bool((got[3:] == 0).all()), "pad lanes must land on slot 0") + + def test_slice_and_copy_derive_pointwise_without_handover(self): + """REGRESSION (design): the retired identity-resolver refused any + tensor it had not been handed -- a TBO child's re-padded slice or a + registry's fresh copy raised. Value-based derivation must accept + both, pointwise, with no adopt/handover call.""" + src, _, rows, seq_lens, _, want_full, want_swa = self._built(n=4) + padded = torch.cat([want_full, want_full.new_zeros(2)]) + # TBO-child shape: a slice crossing the pad boundary. + got = self._field(src, rows, seq_lens, padded[2:6]) + self.assertTrue(torch.equal(got[:2], want_swa[2:4])) + self.assertTrue(bool((got[2:] == 0).all())) + # Registry shape: a fresh equal-value copy. + got2 = self._field(src, rows, seq_lens, want_full.clone()) + self.assertTrue(torch.equal(got2, want_swa)) + + def test_tombstoned_swa_page_clamps_to_sink(self): + src, allocator, rows, seq_lens, virt, want_full, _ = self._built(ps=1, n=2) + allocator.swa_v2p_page_table[int(virt[0])] = -1 + got = self._field(src, rows, seq_lens, want_full[:1]) + self.assertEqual(int(got[0]), 0) + + def test_static_swa_pool_derives_via_pool_translate(self): + """Static SWA pools: the field is the pool's own legacy full->swa + translate, computed at the same build; the rebind stays a no-op.""" + pool = SWAKVPool.__new__(SWAKVPool) + pool.full_to_swa_index_mapping = torch.arange(10, dtype=torch.int64) + pool.translate_loc_from_full_to_swa = lambda t: t + 100 + src = KVIndexTranslator( + req_to_token=torch.zeros((2, 4), dtype=torch.int64), + token_to_kv_pool_allocator=SimpleNamespace(), + token_to_kv_pool=pool, + page_size=1, + device=_DEV, + ) + loc = torch.tensor([5, 6], dtype=torch.int64) + fb = _FakeForwardBatch(out_cache_loc=loc) + src.rebind_write_loc(fb) + self.assertIs(fb.out_cache_loc, loc, "disabled rebind must be a no-op") + self.assertTrue(torch.equal(src.sliding_window_write_loc_for(loc), loc + 100)) + + def test_no_loc_or_no_swa_side_yields_none(self): + # Unified swa composite, but there is no write loc this forward. + src, _, rows, seq_lens, _, _, _ = self._built(n=2) + self.assertIsNone(src.sliding_window_write_loc_for(None)) + # Passthrough on a non-SWA pool: a loc is given, but there is no swa + # id space to derive into. + plain = KVIndexTranslator( + req_to_token=torch.zeros((2, 4), dtype=torch.int64), + token_to_kv_pool_allocator=SimpleNamespace(), + token_to_kv_pool=SimpleNamespace(), + page_size=1, + device=_DEV, + ) + self.assertIsNone( + plain.sliding_window_write_loc_for(torch.tensor([3], dtype=torch.int64)) + ) + + def test_rebind_retires_the_view_memo(self): + ps = 1 + allocator = _build_composite(ps) + req_to_token, rows, seq_lens = _alloc_and_fill(allocator, ps, lens=[3, 2]) + src = _make_source(allocator, req_to_token, ps) + fb = _FakeForwardBatch( + req_pool_indices=rows, seq_lens=seq_lens, seq_lens_cpu=seq_lens + ) + v1 = src.index_table_for_batch(fb) + src.rebind_write_loc(_FakeForwardBatch(out_cache_loc=None)) + v2 = src.index_table_for_batch(fb) + self.assertIsNot(v2, v1, "rebind starts the next forward: stale views die") + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_multi_ended_allocator.py b/test/registered/unit/mem_cache/test_multi_ended_allocator.py index 3329c0de9631..e1a02127b567 100644 --- a/test/registered/unit/mem_cache/test_multi_ended_allocator.py +++ b/test/registered/unit/mem_cache/test_multi_ended_allocator.py @@ -539,6 +539,31 @@ def test_translate_kv_loc_with_out_clamps_tombstoned_v2p(self): ) self.assertEqual(int(buf[1].item()), 0) + def test_slot_zero_sink_invariant_survives_churn(self): + """PINNED INVARIANT: virtual 0 <-> physical 0 (the padding sink), so + `translate_kv_loc(zeros) == zeros` -- after init AND after alloc/free/ + compaction churn. The cuda-graph capture path RELIES on this: the + physical-loc contract replaced capture-time translate with a plain + copy of the zero-filled static buffer, which is only equivalent while + v2p[0] == 0. If an allocator change breaks this, captured stores would + write pad lanes to a live slot.""" + _, full_alloc, _, full_kv, _ = self._build_pair() + zeros = torch.zeros(4, dtype=torch.int64) + + self.assertEqual(int(full_alloc.virtual_to_physical[0].item()), 0) + self.assertTrue(torch.equal(full_alloc.translate_kv_loc(zeros), zeros)) + + # Churn: allocate, free interior (forces compaction moves), re-allocate. + a = self._alloc(full_alloc, full_kv, 6) + b = self._alloc(full_alloc, full_kv, 6) + self._free(full_alloc, full_kv, a) + c = self._alloc(full_alloc, full_kv, 4) + self._free(full_alloc, full_kv, b) + self._free(full_alloc, full_kv, c) + + self.assertEqual(int(full_alloc.virtual_to_physical[0].item()), 0) + self.assertTrue(torch.equal(full_alloc.translate_kv_loc(zeros), zeros)) + # --------------------------------------------------------------------------- # Shared SWA composite — unit tests diff --git a/test/registered/unit/mem_cache/test_unified_mha_views.py b/test/registered/unit/mem_cache/test_unified_mha_views.py index a28fcbe0573c..3d8a173ee0d2 100644 --- a/test/registered/unit/mem_cache/test_unified_mha_views.py +++ b/test/registered/unit/mem_cache/test_unified_mha_views.py @@ -483,8 +483,9 @@ def test_hnd_env_cannot_hijack_layout(self): class TestFactoryDenseViews(unittest.TestCase): - """The real SWA factory builds both sub-pools and wires the matching - kernel-facing multipliers into the composite allocator.""" + """The real SWA factory builds dense sub-pools and wires the matching + kernel-facing multipliers into the composite allocator. End-to-end over + that factory, the rebind must emit BOTH kernel-facing write locs.""" # _swa_factory geometry: L_full = L_swa = 2, uniform 8/8 dims, ps = 1. FULL_MULT = 4 # 2 * L_full @@ -526,6 +527,42 @@ def test_factory_wires_matching_multipliers(self): self.assertEqual(b.token_to_kv_pool.swa_kv_pool.k_buffer[0].dim(), 3) self.assertGreater(pool.view_tail_pad_bytes, 0) + def test_rebind_emits_dense_full_and_build_derives_swa(self): + """End-to-end over the real factory: rebind_write_loc rebinds + out_cache_loc to FULL-kernel-facing ids (phase 1), and the per-batch build + derives the SWA-DENSE write loc pointwise from those kernel-facing values + (phase 2) — both checked against the formulas over the VIRTUAL + ids.""" + from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator + + b = self._bundle() + alloc = b.token_to_kv_pool_allocator + v = alloc.alloc(4) + self.assertIsNotNone(v) + expected_full = alloc.full_v2p_page_table[v] * self.FULL_MULT # ps=1 + expected_swa = alloc.swa_v2p_page_table[v] * self.SWA_MULT + + class _FB: + pass + + fb = _FB() + fb.out_cache_loc = v.clone() + source = KVIndexTranslator( + req_to_token=torch.zeros((2, 8), dtype=torch.int64), + token_to_kv_pool_allocator=alloc, + token_to_kv_pool=b.token_to_kv_pool, + page_size=1, + device="cpu", + ) + self.assertTrue(source.is_translating) + source.rebind_write_loc(fb) + self.assertTrue(torch.equal(fb.out_cache_loc, expected_full)) + self.assertTrue( + torch.equal( + source.sliding_window_write_loc_for(fb.out_cache_loc), expected_swa + ) + ) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/model_executor/test_unified_out_cache_loc_rebind.py b/test/registered/unit/model_executor/test_unified_out_cache_loc_rebind.py new file mode 100644 index 000000000000..e11706057d10 --- /dev/null +++ b/test/registered/unit/model_executor/test_unified_out_cache_loc_rebind.py @@ -0,0 +1,181 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""ForwardBatch construction wires the unified write-loc rebind. + +`init_new` must call `kv_index_translator.rebind_write_loc`: a construction +path that skips it ships VIRTUAL write ids to the kernels, a silent +wrong-slot store. Also runs the REAL `_pad_inputs_to_size` against a live +translator, since pad lanes are zeros and zeros must derive to the slot-0 +sink. Sliding-window semantics are pinned in test_kv_index_translator.py. + + python -m pytest test/registered/unit/model_executor/test_unified_out_cache_loc_rebind.py -v +""" + +import ast +import inspect +import textwrap +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator +from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +_DEV = "cpu" + + +def _make_fb(out_cache_loc, **kw): + """Minimal ForwardBatch with only the required core fields.""" + n = 0 if out_cache_loc is None else out_cache_loc.shape[0] + defaults = dict( + forward_mode=ForwardMode.DECODE, + batch_size=max(n, 1), + input_ids=torch.zeros(max(n, 1), dtype=torch.int64), + req_pool_indices=torch.zeros(max(n, 1), dtype=torch.int64), + seq_lens=torch.ones(max(n, 1), dtype=torch.int64), + out_cache_loc=out_cache_loc, + seq_lens_sum=max(n, 1), + ) + defaults.update(kw) + return ForwardBatch(**defaults) + + +def _armed_source(v2p, swa_map): + """A KVIndexTranslator hand-armed with fake translates: this file pins the + ForwardBatch-side wiring, not the composite's formulas (those are pinned + in test_kv_index_translator.py over the real allocator).""" + src = KVIndexTranslator( + req_to_token=torch.zeros((1, 4), dtype=torch.int64), + token_to_kv_pool_allocator=SimpleNamespace(), + token_to_kv_pool=SimpleNamespace(), + page_size=1, + device=_DEV, + ) + src.is_translating = True + src._translate_full = lambda t, out=None: v2p[t.to(torch.int64)] + # Phase 2 derives from DENSE values through p2v + the swa v2p; arm the + # inverse of the fake v2p (ps=1, both multipliers 1: dense == physical, + # and the expected swa loc for virtual t is swa_map[t]). + p2v = torch.zeros(int(v2p.max()) + 1, dtype=torch.int64) + p2v[v2p] = torch.arange(v2p.numel(), dtype=torch.int64) + src._full_p2v_table = p2v + src._swa_v2p_table = swa_map + src._full_page_multiplier = 1 + src._swa_page_multiplier = 1 + return src + + +def _call_names(func) -> list: + """Dotted call targets appearing in `func`'s body, e.g. + 'model_runner.kv_index_translator.rebind_write_loc'.""" + tree = ast.parse(textwrap.dedent(inspect.getsource(func))) + names = [] + for node in ast.walk(tree): + if isinstance(node, ast.Call): + parts = [] + cur = node.func + while isinstance(cur, ast.Attribute): + parts.append(cur.attr) + cur = cur.value + if isinstance(cur, ast.Name): + parts.append(cur.id) + names.append(".".join(reversed(parts))) + return names + + +class TestForwardBatchWiring(CustomTestCase): + """Critical-path bookkeeping: the construction-time call sites.""" + + def test_init_new_calls_the_rebind(self): + self.assertIn( + "model_runner.kv_index_translator.rebind_write_loc", + _call_names(ForwardBatch.init_new.__func__), + "init_new must rebind the write loc through the source; a batch " + "built without it ships virtual ids to the kernels", + ) + + +class TestPadComposesWithDerivation(CustomTestCase): + def _fake_runner_for_pad(self, src): + return SimpleNamespace( + attn_backend=SimpleNamespace(get_cuda_graph_seq_len_fill_value=lambda: 0), + kv_index_translator=src, + ) + + def test_pad_lanes_derive_to_sink_and_slices_stay_pointwise(self): + """The REAL `_pad_inputs_to_size` composes with phase 2: pad lanes are + zeros, zeros derive to the slot-0 sink, and any slice of the padded + tensor (the TBO-child shape) derives pointwise -- no handover call + exists for the pad to make.""" + n, padded = 3, 6 + v2p = torch.arange(64, dtype=torch.int64) * 3 + swa_map = torch.arange(64, dtype=torch.int64) * 5 + src = _armed_source(v2p, swa_map) + virt = torch.tensor([11, 12, 13], dtype=torch.int64) + fb = _make_fb(virt.clone()) + fb.positions = torch.arange(n, dtype=torch.int64) + fb.lora_ids = [None] * fb.batch_size + src.rebind_write_loc(fb) + self.assertTrue(torch.equal(fb.out_cache_loc, v2p[virt])) + + fb._pad_inputs_to_size(self._fake_runner_for_pad(src), padded, fb.batch_size) + + self.assertEqual(fb.out_cache_loc.shape[0], padded) + # Padded tail lanes go to slot 0 -- the reserved dummy-write sink. + self.assertTrue(bool((fb.out_cache_loc[n:] == 0).all())) + loc = src._swa_write_loc_unified(fb.out_cache_loc) + self.assertTrue(torch.equal(loc[:n], swa_map[virt])) + self.assertTrue(bool((loc[n:] == 0).all())) + self.assertEqual(loc.dtype, torch.int64) + # The TBO-child shape: a slice of the PADDED tensor derives pointwise. + sub = src._swa_write_loc_unified(fb.out_cache_loc[1:5]) + self.assertTrue(torch.equal(sub, loc[1:5])) + + def test_the_probe_separates_kernel_facing_from_virtual_ids(self): + """A skipped rebind is the failure mode this contract has no other + guard against: virtual ids stay inside the OOB probe's bounds (they are + `blocks_per_page` times SMALLER than a kernel-facing id), so the store lands on + the wrong slots and only the output is wrong. The kernel-facing probe + is what separates them -- the in-page offset of a kernel-facing id is always + below page_size, and a virtual id's is not unless it happens to fall in + the first block.""" + for page_size, blocks in ((1, 8), (4, 6)): + with self.subTest(page_size=page_size, blocks=blocks): + stride = page_size * blocks + virt = torch.arange(1, 2 * stride, dtype=torch.int64) + dense = (virt // page_size) * stride + virt % page_size + in_space = dense % stride < page_size + self.assertTrue(bool(in_space.all()), "kernel-facing ids must pass") + # Virtual ids pass only in the first block; that is why the + # probe needs a batch, not one id, to be conclusive. + caught = ~(virt % stride < page_size) + self.assertTrue(bool(caught.any()), "virtual ids must be caught") + + def test_empty_loc_rebinds_to_empty(self): + src = _armed_source( + torch.arange(8, dtype=torch.int64), torch.arange(8, dtype=torch.int64) + ) + fb = _make_fb(torch.empty(0, dtype=torch.int64)) + src.rebind_write_loc(fb) + self.assertEqual(fb.out_cache_loc.numel(), 0) + self.assertEqual(src._swa_write_loc_unified(fb.out_cache_loc).numel(), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py b/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py new file mode 100644 index 000000000000..3f950fe6ec70 --- /dev/null +++ b/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py @@ -0,0 +1,92 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""`--enable-unified-memory` disables PREFILL cuda-graph capture. + +BUG REGRESSION. Only decode capture is wired: the prefill graph runner builds +its ForwardBatch directly, so it never runs the unified pool's write-loc +rebind (rebind_write_loc) and the captured batch holds VIRTUAL ids -- the +captured store would silently write wrong slots. + +The old gate only rejected `TC_PIECEWISE`, but the generic prefill default is +`BREAKABLE` -- so the DEFAULT unified invocation was broken; it only ever +worked when `--disable-piecewise-cuda-graph` (a deprecated alias for +`--cuda-graph-backend-prefill=disabled`) happened to be passed. + +Pinned: the default is auto-disabled with a warning (unified boots out of the +box), an EXPLICIT prefill backend still raises (never silently override a +user's stated intent), and decode capture is untouched either way. + + python -m pytest test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py -v +""" + +import unittest +from types import SimpleNamespace + +from sglang.srt.arg_groups.kv_cache_hook import handle_unified_memory_pool +from sglang.srt.model_executor.cuda_graph_config import Backend +from sglang.srt.server_args import ServerArgs +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def _run_handler(*, prefill_backend, explicit): + """Run just `handle_unified_memory_pool` over a minimal stand-in.""" + sa = ServerArgs.__new__(ServerArgs) + cg = SimpleNamespace( + prefill=SimpleNamespace(backend=prefill_backend), + decode=SimpleNamespace(backend=Backend.FULL), + ) + for name, value in { + "enable_unified_memory": True, + "disaggregation_mode": "null", + "speculative_algorithm": None, + "speculative_eagle_topk": None, + "enable_hierarchical_cache": False, + "enable_lmcache": False, + "dcp_size": 1, + "cuda_graph_config": cg, + "cuda_graph_backend_prefill": prefill_backend if explicit else None, + }.items(): + object.__setattr__(sa, name, value) + handle_unified_memory_pool(sa) + return cg + + +class TestUnifiedPrefillCudaGraphGate(unittest.TestCase): + def test_default_prefill_capture_is_auto_disabled(self): + """The generic default (BREAKABLE) must be turned off, not crash the + server 30 seconds later inside graph capture.""" + for backend in (Backend.BREAKABLE, Backend.FULL, Backend.TC_PIECEWISE): + cg = _run_handler(prefill_backend=backend, explicit=False) + self.assertEqual(cg.prefill.backend, Backend.DISABLED) + # Decode capture is the wired path and must survive untouched. + self.assertEqual(cg.decode.backend, Backend.FULL) + + def test_explicit_prefill_backend_is_refused(self): + """A user who explicitly asked for prefill graphs gets a clear error, + not a silent override of their stated intent.""" + for backend in (Backend.BREAKABLE, Backend.FULL, Backend.TC_PIECEWISE): + with self.assertRaises(ValueError) as ctx: + _run_handler(prefill_backend=backend, explicit=True) + self.assertIn("prefill capture is not wired", str(ctx.exception)) + + def test_already_disabled_is_a_no_op(self): + cg = _run_handler(prefill_backend=Backend.DISABLED, explicit=True) + self.assertEqual(cg.prefill.backend, Backend.DISABLED) + self.assertEqual(cg.decode.backend, Backend.FULL) + + +if __name__ == "__main__": + unittest.main()