diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py index e6473c5b6..a1c221c23 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py @@ -17,6 +17,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( cp_shared_kv_mla_prefetch_min_async_extend_tokens, cp_shared_kv_mla_prefetch_min_prefix_pages, cp_shared_kv_mla_prefetch_should_log_layer, + build_page_table_row_req_id, filter_locs_mappable_to_physical_pool, filter_pages_mappable_to_physical_pool, fill_current_index_page_slots, @@ -25,8 +26,8 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( get_or_build_shared_token_kv_slot_remap, materialize_local_paged_buffer_page_slots_into, materialize_local_token_kv_page_slots_into, - remap_logical_pages_to_shared_paged_slot_dense_pages, - remap_logical_locs_to_shared_token_slot_dense_locs, + remap_logical_pages_to_slot_dense_pages, + remap_logical_locs_to_slot_dense_locs_optimized, slot_range_to_page_slice, slot_range_to_token_slice, ) @@ -570,6 +571,7 @@ class CpSharedKVMlaPrefetcher: layer_id: int, kv_cache: torch.Tensor, logical_locs: torch.Tensor, + loc_req_id: torch.Tensor | None = None, ) -> Optional[tuple[torch.Tensor, torch.Tensor]]: if self.disabled: self._log_layer( @@ -680,15 +682,13 @@ class CpSharedKVMlaPrefetcher: layout=self.layout, physical_token_capacity=kv_cache.shape[0], ) - dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs( + if loc_req_id is None: + loc_req_id = torch.zeros_like(logical_locs, dtype=torch.long) + dense_locs = remap_logical_locs_to_slot_dense_locs_optimized( logical_locs, - slot_remap=_PrefetchTokenSlotRemapView( - slot_logical_pages=self.slot_logical_pages, - page_inverse=self.page_inverse, - slot_sorted_logical_pages_by_row=self.slot_sorted_logical_pages_by_row, - slot_sorted_dense_pages_by_row=self.slot_sorted_dense_pages_by_row, - ), + page_inverse=self.page_inverse, page_size=self.page_size, + loc_req_id=loc_req_id, ) remap_ms = _cpu_timing_ms(remap_cpu) total_ms = _cpu_timing_ms(consume_cpu) @@ -720,6 +720,8 @@ class CpSharedKVMlaPrefetcher: logical_locs: torch.Tensor, current_kv_cache: torch.Tensor, current_locs: torch.Tensor, + loc_req_id: torch.Tensor | None = None, + current_req_id: torch.Tensor | None = None, current_remap_page_size: int | None = None, current_remap_logical_page_capacity: int | None = None, ) -> Optional[tuple[torch.Tensor, torch.Tensor]]: @@ -785,15 +787,15 @@ class CpSharedKVMlaPrefetcher: layout=self.layout, physical_token_capacity=kv_cache.shape[0], ) - dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs( + if loc_req_id is None: + loc_req_id = torch.zeros_like(logical_locs, dtype=torch.long) + if current_req_id is None: + current_req_id = torch.zeros_like(current_locs, dtype=torch.long) + dense_locs = remap_logical_locs_to_slot_dense_locs_optimized( logical_locs, - slot_remap=_PrefetchTokenSlotRemapView( - slot_logical_pages=self.slot_logical_pages, - page_inverse=self.page_inverse, - slot_sorted_logical_pages_by_row=self.slot_sorted_logical_pages_by_row, - slot_sorted_dense_pages_by_row=self.slot_sorted_dense_pages_by_row, - ), + page_inverse=self.page_inverse, page_size=self.page_size, + loc_req_id=loc_req_id, ) mixed_kv_cache, mixed_locs, _ = fill_current_kv_page_slots_and_remap_locs( dense_kv_cache=dense_kv_cache, @@ -803,6 +805,7 @@ class CpSharedKVMlaPrefetcher: current_locs=current_locs, page_inverse=self.page_inverse, page_size=self.page_size, + current_req_id=current_req_id, mask_non_current_in_current_pages=True, ) if self.layout.cp_size > 1 and self.prefix_pages < self.total_slots: @@ -1498,15 +1501,10 @@ class CpSharedKVIndexPrefetcher: layout=self.layout, physical_page_capacity=page_buffer.shape[0], ) - dense_pages = remap_logical_pages_to_shared_paged_slot_dense_pages( + dense_pages = remap_logical_pages_to_slot_dense_pages( logical_pages, - slot_remap=_PrefetchPagedSlotRemapView( - slot_logical_pages=self.slot_logical_pages, - page_inverse=self.page_inverse, - dense_pages=None, - slot_sorted_logical_pages_by_row=self.slot_sorted_logical_pages_by_row, - slot_sorted_dense_pages_by_row=self.slot_sorted_dense_pages_by_row, - ), + page_inverse=self.page_inverse, + page_req_id=build_page_table_row_req_id(logical_pages), ) remap_ms = _cpu_timing_ms(remap_cpu) total_ms = _cpu_timing_ms(consume_cpu) @@ -1540,6 +1538,7 @@ class CpSharedKVIndexPrefetcher: current_locs: torch.Tensor, page_size: int, index_head_dim: int, + current_req_id: torch.Tensor | None = None, ) -> Optional[tuple[torch.Tensor, torch.Tensor]]: if self.disabled: self._log_layer( @@ -1598,15 +1597,12 @@ class CpSharedKVIndexPrefetcher: dense_page_buffer = handle.dense_page_buffer remap_cpu = _cpu_timing_start() - dense_pages = remap_logical_pages_to_shared_paged_slot_dense_pages( + if current_req_id is None: + current_req_id = torch.zeros_like(current_locs, dtype=torch.long) + dense_pages = remap_logical_pages_to_slot_dense_pages( logical_pages, - slot_remap=_PrefetchPagedSlotRemapView( - slot_logical_pages=self.slot_logical_pages, - page_inverse=self.page_inverse, - dense_pages=None, - slot_sorted_logical_pages_by_row=self.slot_sorted_logical_pages_by_row, - slot_sorted_dense_pages_by_row=self.slot_sorted_dense_pages_by_row, - ), + page_inverse=self.page_inverse, + page_req_id=build_page_table_row_req_id(logical_pages), ) dense_page_buffer = fill_current_index_page_slots( dense_page_buffer=dense_page_buffer, @@ -1616,6 +1612,7 @@ class CpSharedKVIndexPrefetcher: page_inverse=self.page_inverse, page_size=page_size, index_head_dim=index_head_dim, + current_req_id=current_req_id, ) if self.layout.cp_size > 1 and self.prefix_pages < self.total_slots: current_pages = slot_range_to_page_slice( diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py index 0b925fe01..0c7fcde1a 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -519,11 +519,17 @@ def _tai_current_slot_fill_sparse_pages_self_test( device=device, dtype=torch.int64, ) - page_inverse = torch.full((32,), -1, device=device, dtype=torch.long) - page_inverse[0] = 0 - page_inverse[5] = 1 - page_inverse[10] = 2 - page_inverse[25] = 3 + # Per-request inverse [batch_rows, capacity]; this probe is single-request + # (batch_rows=1), so req_id is all-zeros and the result matches the old + # 1-D behaviour exactly. + page_inverse = torch.full((1, 32), -1, device=device, dtype=torch.long) + page_inverse[0, 0] = 0 + page_inverse[0, 5] = 1 + page_inverse[0, 10] = 2 + page_inverse[0, 25] = 3 + current_req_id = torch.zeros( + (int(current_locs.numel()),), device=device, dtype=torch.long + ) mixed_kv, mixed_locs, current_mask = fill_kernel( dense_kv, @@ -532,6 +538,7 @@ def _tai_current_slot_fill_sparse_pages_self_test( logical_locs, current_locs, page_inverse, + current_req_id, page_size=page_size, mask_non_current_in_current_pages=True, ) @@ -1325,6 +1332,8 @@ def _try_tai_materialize_token_kv_pages_and_locs( logical_page_capacity: int, layout: CpSharedKVLayout, page_size: int, + batch_rows: int, + loc_req_id: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor] | None: if not _tai_materialize_runtime_enabled(): _log_tai_materialize_runtime_disabled("token_pages_and_locs") @@ -1339,10 +1348,12 @@ def _try_tai_materialize_token_kv_pages_and_locs( page_inverse = kernels.build_slot_page_inverse( tai_slot_logical_pages, logical_page_capacity, + int(max(int(batch_rows), 1)), ) dense_locs = kernels.remap_logical_locs_to_slot_dense_locs( _contiguous_for_tai(logical_locs), page_inverse, + _row_req_id_for_tai(loc_req_id, logical_locs), page_size=page_size, ) dense_kv_cache = kernels.materialize_shared_token_kv_pages( @@ -1366,6 +1377,7 @@ def _try_tai_materialize_token_kv_pages_and_locs( def _try_tai_build_slot_page_inverse( slot_logical_pages: torch.Tensor, logical_page_capacity: int, + batch_rows: int, ) -> torch.Tensor | None: if not _tai_materialize_runtime_enabled(): _log_tai_materialize_runtime_disabled("slot_page_inverse") @@ -1379,6 +1391,7 @@ def _try_tai_build_slot_page_inverse( return kernels.build_slot_page_inverse( _contiguous_for_tai(slot_logical_pages.reshape(-1)), logical_page_capacity, + int(max(int(batch_rows), 1)), ) except Exception as exc: _log_tai_materialize_fallback( @@ -1390,19 +1403,134 @@ def _try_tai_build_slot_page_inverse( return None +def _row_req_id_for_tai(req_id: torch.Tensor, ref: torch.Tensor) -> torch.Tensor: + """Return a COMPACT int64 req-id with one entry per ROW of ``ref``. + + The TAI remap kernel infers ``inner = ref.numel() / num_rows`` and reads + ``req_id[loc_index / inner]``, and the fill kernels (1-D ``ref``) read + ``req_id[row]`` directly -- so a ``[num_rows]`` tensor suffices and we avoid + materializing a full ``[tokens, topk]`` broadcast (which for index_topk=2048 + would be a 100+ MB int64 copy per call). ``num_rows = ref.shape[0]``. + """ + + rows = int(ref.shape[0]) if ref.dim() >= 1 else 1 + r = req_id.reshape(-1) + if int(r.numel()) != rows: + if int(r.numel()) > rows: + r = r[:rows] + elif int(r.numel()) > 0: + pad = r[-1:].expand(rows - int(r.numel())) + r = torch.cat([r, pad], dim=0) + else: + r = req_id.new_zeros(rows) + return r.to(torch.long).contiguous() + + +def build_token_loc_req_id( + num_tokens: int, + extend_seq_lens: Any, + device: torch.device, +) -> torch.Tensor: + """Per-token owning-request id for token-major logical locs (page_table_1). + + ``page_table_1`` lays its first (token) dim out per request, contiguously, by + ``extend_seq_lens`` (see ``transform_index_page_table_prefill``). Returns a + 1-D ``[num_tokens]`` req-id that the materialize remap broadcasts over the + topk dim. Defaults to all-zeros (single request) when the lengths do not sum + to ``num_tokens`` (e.g. static-padded warmup) so it never misaligns a real + request onto the wrong page-inverse row. + """ + + if extend_seq_lens is None: + return torch.zeros(num_tokens, device=device, dtype=torch.long) + if isinstance(extend_seq_lens, torch.Tensor): + lens = [int(x) for x in extend_seq_lens.tolist()] + else: + lens = [int(x) for x in extend_seq_lens] + nreq = len(lens) + total = sum(lens) + if nreq <= 1 or total == 0: + return torch.zeros(num_tokens, device=device, dtype=torch.long) + req = torch.repeat_interleave( + torch.arange(nreq, device=device, dtype=torch.long), + torch.tensor(lens, device=device, dtype=torch.long), + ) + if int(req.numel()) == num_tokens: + return req + if int(req.numel()) > num_tokens: + return req[:num_tokens] + # Trailing static-padding rows (invalid locs, masked by loc>=0) -> last request. + pad = torch.full( + (num_tokens - int(req.numel()),), + nreq - 1, + device=device, + dtype=torch.long, + ) + return torch.cat([req, pad], dim=0) + + +def get_cp_shared_kv_token_loc_req_id( + forward_batch: Any, + logical_locs: torch.Tensor, + extend_seq_lens: Any, +) -> torch.Tensor: + """Per-batch-cached token-major ``loc_req_id`` for ``logical_locs``. + + ``loc_req_id`` depends only on the query-token count and ``extend_seq_lens`` + (both batch-invariant), so it is identical across all layers of a forward + pass. Cache it on the forward batch keyed by the token count to avoid + rebuilding it in the per-layer attention loop (mirrors + ``get_cp_shared_kv_local_current_req_id``'s batch caching). + """ + + num_tokens = int(logical_locs.shape[0]) + cached = getattr(forward_batch, "cp_token_loc_req_id", None) + cached_n = getattr(forward_batch, "cp_token_loc_req_id_ntok", None) + if ( + cached is not None + and cached_n == num_tokens + and cached.device == logical_locs.device + ): + return cached + req = build_token_loc_req_id(num_tokens, extend_seq_lens, logical_locs.device) + forward_batch.cp_token_loc_req_id = req + forward_batch.cp_token_loc_req_id_ntok = num_tokens + return req + + +def build_page_table_row_req_id(logical_pages: torch.Tensor) -> torch.Tensor: + """Owning-request id for a 2-D ``[batch_rows, pages_per_request]`` page table. + + The page table is row-major per request (row r == request r), so the owning + request of every page in row r is simply ``r``. Returns ``[batch_rows, 1]`` + (broadcast over the page columns by the remap) or all-zeros for a 1-D table. + """ + + if logical_pages.dim() >= 2: + return torch.arange( + int(logical_pages.shape[0]), + device=logical_pages.device, + dtype=torch.long, + ).unsqueeze(1) + return torch.zeros_like(logical_pages, dtype=torch.long) + + def build_slot_page_inverse_optimized( slot_logical_pages: torch.Tensor, logical_page_capacity: int, + batch_rows: int, ) -> torch.Tensor: tai_result = _try_tai_build_slot_page_inverse( slot_logical_pages, logical_page_capacity, + batch_rows, ) if tai_result is not None: return tai_result return build_slot_page_inverse( slot_logical_pages, logical_page_capacity=logical_page_capacity, + batch_rows=batch_rows, ) @@ -1410,6 +1538,7 @@ def remap_logical_locs_to_slot_dense_locs_optimized( logical_locs: torch.Tensor, page_inverse: torch.Tensor, page_size: int, + loc_req_id: torch.Tensor, ) -> torch.Tensor: if _tai_materialize_runtime_enabled(): kernels = _load_tai_materialize_kernels() @@ -1418,6 +1547,7 @@ def remap_logical_locs_to_slot_dense_locs_optimized( return kernels.remap_logical_locs_to_slot_dense_locs( _contiguous_for_tai(logical_locs), _contiguous_for_tai(page_inverse), + _row_req_id_for_tai(loc_req_id, logical_locs), page_size=page_size, ) except Exception as exc: @@ -1431,6 +1561,7 @@ def remap_logical_locs_to_slot_dense_locs_optimized( logical_locs, page_inverse=page_inverse, page_size=page_size, + loc_req_id=loc_req_id, ) @@ -1443,6 +1574,7 @@ def _try_tai_fill_current_kv_page_slots_and_remap_locs( current_locs: torch.Tensor, page_inverse: torch.Tensor, page_size: int, + current_req_id: torch.Tensor, mask_non_current_in_current_pages: bool, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None: if not _tai_materialize_runtime_enabled(): @@ -1494,6 +1626,7 @@ def _try_tai_fill_current_kv_page_slots_and_remap_locs( _contiguous_for_tai(logical_locs), _contiguous_for_tai(current_locs.reshape(-1)), _contiguous_for_tai(page_inverse), + _row_req_id_for_tai(current_req_id, current_locs.reshape(-1)), page_size=int(page_size), mask_non_current_in_current_pages=bool(mask_non_current_in_current_pages), ) @@ -1519,6 +1652,7 @@ def fill_current_kv_page_slots_and_remap_locs( current_locs: torch.Tensor, page_inverse: torch.Tensor, page_size: int, + current_req_id: torch.Tensor, mask_non_current_in_current_pages: bool = True, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Fill current suffix KV into preallocated dense page slots and remap locs. @@ -1526,6 +1660,9 @@ def fill_current_kv_page_slots_and_remap_locs( This is the CP shared-KV prefetch compose path. The dense buffer already has page-aligned slots for prefix and suffix pages; current rows should be copied into those suffix slots instead of appended to a second compact suffix. + ``current_req_id`` (aligned with ``current_locs``) routes each current row + through its own request's page-inverse row so divergent current writes never + land on another request's dense slot (bs>1 contamination). """ tai_result = _try_tai_fill_current_kv_page_slots_and_remap_locs( @@ -1536,6 +1673,7 @@ def fill_current_kv_page_slots_and_remap_locs( current_locs=current_locs, page_inverse=page_inverse, page_size=page_size, + current_req_id=current_req_id, mask_non_current_in_current_pages=mask_non_current_in_current_pages, ) if tai_result is not None: @@ -1557,6 +1695,7 @@ def fill_current_kv_page_slots_and_remap_locs( current_locs.reshape(-1), page_inverse=page_inverse, page_size=page_size, + loc_req_id=current_req_id.reshape(-1), ) valid_current_rows = current_dense_locs >= 0 if torch.any(valid_current_rows): @@ -1589,6 +1728,7 @@ def _try_tai_fill_current_index_page_slots( page_inverse: torch.Tensor, page_size: int, index_head_dim: int, + current_req_id: torch.Tensor, ) -> torch.Tensor | None: if not _tai_materialize_runtime_enabled(): _log_tai_materialize_runtime_disabled("fill_current_index_page_slots") @@ -1620,6 +1760,7 @@ def _try_tai_fill_current_index_page_slots( _contiguous_for_tai(current_index_scale), _contiguous_for_tai(current_locs.reshape(-1)), _contiguous_for_tai(page_inverse), + _row_req_id_for_tai(current_req_id, current_locs.reshape(-1)), page_size=int(page_size), index_head_dim=int(index_head_dim), ) @@ -1647,6 +1788,7 @@ def fill_current_index_page_slots( page_inverse: torch.Tensor, page_size: int, index_head_dim: int, + current_req_id: torch.Tensor, ) -> torch.Tensor: """Fill current index K/scale rows into a slot-dense page buffer. @@ -1654,9 +1796,12 @@ def fill_current_index_page_slots( the dummy page. Current rows are copied into the row/offset selected by their logical token loc; tail-page slack remains zero and is invisible because the index kernels still receive the valid sequence length separately. + ``current_req_id`` (aligned with ``current_locs``) routes each current row + through its own request's page-inverse row (bs>1 contamination fix). """ current_locs = current_locs.reshape(-1) + current_req_id = current_req_id.reshape(-1) current_rows = int(current_locs.numel()) if current_rows == 0: return dense_page_buffer @@ -1668,6 +1813,7 @@ def fill_current_index_page_slots( page_inverse=page_inverse, page_size=page_size, index_head_dim=index_head_dim, + current_req_id=current_req_id, ) if tai_result is not None: return tai_result @@ -1713,18 +1859,24 @@ def fill_current_index_page_slots( f"scale_bytes_per_token={int(scale_bytes.shape[1])} page_size={page_size}" ) + capacity = int(page_inverse.shape[-1]) + batch_rows = int(page_inverse.shape[0]) + req_long = current_req_id.to(torch.long) current_pages = torch.div(current_locs, page_size, rounding_mode="floor") valid_pages = ( (current_locs >= 0) & (current_pages >= 0) - & (current_pages < int(page_inverse.numel())) + & (current_pages < capacity) + & (req_long >= 0) + & (req_long < batch_rows) ) safe_pages = torch.clamp( current_pages, min=0, - max=max(int(page_inverse.numel()) - 1, 0), + max=max(capacity - 1, 0), ) - dense_pages = page_inverse[safe_pages.to(torch.long)].to(torch.long) + safe_req = torch.clamp(req_long, min=0, max=max(batch_rows - 1, 0)) + dense_pages = page_inverse[safe_req, safe_pages.to(torch.long)].to(torch.long) valid_rows = valid_pages & (dense_pages > 0) if not torch.any(valid_rows): return dense_page_buffer @@ -2948,11 +3100,28 @@ def _logical_page_capacity_from_physical_page_capacity( def build_slot_page_inverse( slot_logical_pages: torch.Tensor, logical_page_capacity: int, + batch_rows: int, ) -> torch.Tensor: - """Build logical_page -> slot_dense_page map without unique/searchsorted.""" + """Build a PER-REQUEST logical_page -> slot_dense_page map. + Returns a 2-D ``[batch_rows, logical_page_capacity]`` tensor. Each request row + is populated only from that request's own page-table slots, so duplicate + logical-page VALUES shared across requests (radix shared-prefix / cache reuse) + no longer collapse to a single slot. That collapse was the bs>1 cross-request + KV contamination bug: a value-keyed 1-D inverse is last-writer-wins, so two + requests sharing a physical page funnelled through one dense slot and a + divergent current-write on the surviving slot corrupted the other request. + With a per-request inverse, request ``i``'s locs remap through + ``page_inverse[i]`` and never alias request ``j``. + + Slots are row-major over ``[batch_rows, pages_per_request]`` (the same flatten + used by :func:`build_slot_page_remap` on ``real_page_table``), so the owning + request of flat slot ``s`` is ``s // pages_per_request``. + """ + + batch_rows = max(int(batch_rows), 1) page_inverse = torch.full( - (logical_page_capacity,), + (batch_rows, logical_page_capacity), -1, device=slot_logical_pages.device, dtype=torch.long, @@ -2960,22 +3129,36 @@ def build_slot_page_inverse( if logical_page_capacity == 0: return page_inverse - # Page 0 is the dummy/padding page and always maps to dense page 0. - page_inverse[0] = 0 + # Page 0 is the dummy/padding page and always maps to dense page 0 for every + # request row. + page_inverse[:, 0] = 0 if slot_logical_pages.numel() == 0: return page_inverse flat_pages = slot_logical_pages.reshape(-1).to(torch.long) + num_slots = int(flat_pages.numel()) + pages_per_request = num_slots // batch_rows slot_ids = torch.arange( 1, - flat_pages.numel() + 1, + num_slots + 1, device=flat_pages.device, dtype=torch.long, ) + if pages_per_request > 0: + req_ids = torch.div( + torch.arange(num_slots, device=flat_pages.device, dtype=torch.long), + pages_per_request, + rounding_mode="floor", + ).clamp_(max=batch_rows - 1) + else: + req_ids = torch.zeros(num_slots, device=flat_pages.device, dtype=torch.long) valid_pages = (flat_pages > 0) & (flat_pages < logical_page_capacity) - safe_pages = torch.where(valid_pages, flat_pages, torch.zeros_like(flat_pages)) + # Flatten (req, page) into the row-major inverse; invalid slots fold to dense + # page 0 at index 0 (the dummy), matching the 1-D reference's no-op behaviour. + flat_index = req_ids * logical_page_capacity + flat_pages + safe_index = torch.where(valid_pages, flat_index, torch.zeros_like(flat_index)) safe_slot_ids = torch.where(valid_pages, slot_ids, torch.zeros_like(slot_ids)) - page_inverse.scatter_(0, safe_pages, safe_slot_ids) + page_inverse.view(-1).scatter_(0, safe_index, safe_slot_ids) return page_inverse @@ -3040,27 +3223,42 @@ def remap_logical_locs_to_slot_dense_locs( logical_locs: torch.Tensor, page_inverse: torch.Tensor, page_size: int, + loc_req_id: torch.Tensor, ) -> torch.Tensor: - """Map logical token locs through a fixed-shape slot page inverse.""" + """Map logical token locs through a PER-REQUEST slot page inverse. + + ``page_inverse`` is ``[batch_rows, logical_page_capacity]`` and ``loc_req_id`` + carries the owning request of each loc (broadcastable to ``logical_locs``), so + each loc is mapped through ITS OWN request row. Two requests that share a + physical page value therefore no longer alias each other (bs>1 contamination). + """ dense_locs = torch.full_like(logical_locs, -1) if logical_locs.numel() == 0 or page_inverse.numel() == 0: return dense_locs + capacity = int(page_inverse.shape[-1]) + batch_rows = int(page_inverse.shape[0]) locs_long = logical_locs.to(torch.long) + req_long = loc_req_id.to(torch.long) + while req_long.dim() < locs_long.dim(): + req_long = req_long.unsqueeze(-1) valid_locs = locs_long >= 0 safe_locs = torch.where(valid_locs, locs_long, torch.zeros_like(locs_long)) logical_pages = torch.div(safe_locs, page_size, rounding_mode="floor") offsets = torch.remainder(safe_locs, page_size) - pages_in_range = logical_pages < page_inverse.numel() - safe_pages = torch.clamp(logical_pages, max=page_inverse.numel() - 1) - dense_pages = page_inverse[safe_pages] - mapped = valid_locs & pages_in_range & (dense_pages >= 0) + pages_in_range = logical_pages < capacity + req_in_range = (req_long >= 0) & (req_long < batch_rows) + safe_pages = torch.clamp(logical_pages, max=max(capacity - 1, 0)) + safe_req = torch.clamp(req_long, min=0, max=max(batch_rows - 1, 0)) + dense_pages = page_inverse[safe_req, safe_pages] + mapped = valid_locs & pages_in_range & req_in_range & (dense_pages >= 0) if cp_shared_kv_debug_enabled() and torch.any( - valid_locs & pages_in_range & (dense_pages < 0) + valid_locs & pages_in_range & req_in_range & (dense_pages < 0) ): - missing_pages = logical_pages[valid_locs & pages_in_range & (dense_pages < 0)] + missing_mask = valid_locs & pages_in_range & req_in_range & (dense_pages < 0) + missing_pages = logical_pages[missing_mask] raise RuntimeError( "CP shared KV slot remap got logical locs outside remap_logical_pages. " f"missing_page_min={int(missing_pages.min().item())} " @@ -3205,7 +3403,15 @@ def remap_logical_locs_to_shared_token_slot_dense_locs( page_size: int, *, row_ids: torch.Tensor | None = None, + loc_req_id: torch.Tensor | None = None, ) -> torch.Tensor: + if loc_req_id is not None: + return remap_logical_locs_to_slot_dense_locs_optimized( + logical_locs, + page_inverse=slot_remap.page_inverse, + page_size=page_size, + loc_req_id=loc_req_id, + ) row_scoped = remap_logical_locs_to_slot_dense_locs_by_row( logical_locs, slot_remap.slot_sorted_logical_pages_by_row, @@ -3215,40 +3421,54 @@ def remap_logical_locs_to_shared_token_slot_dense_locs( ) if row_scoped is not None: return row_scoped + fallback_req_id = torch.zeros_like(logical_locs, dtype=torch.long) return remap_logical_locs_to_slot_dense_locs_optimized( logical_locs, page_inverse=slot_remap.page_inverse, page_size=page_size, + loc_req_id=fallback_req_id, ) def remap_logical_pages_to_slot_dense_pages( logical_pages: torch.Tensor, page_inverse: torch.Tensor, + page_req_id: torch.Tensor, ) -> torch.Tensor: - """Map logical page ids through a fixed-shape slot page inverse. + """Map logical page ids through a PER-REQUEST slot page inverse. This is the paged-buffer equivalent of - `remap_logical_locs_to_slot_dense_locs`. Page 0 remains the dummy dense page - 0; negative sentinels and pages absent from `page_inverse` remain `-1`. + `remap_logical_locs_to_slot_dense_locs`. ``page_inverse`` is + ``[batch_rows, logical_page_capacity]`` and ``page_req_id`` carries the owning + request of each page (broadcastable to ``logical_pages``), so a page is mapped + through its own request row. Page 0 remains the dummy dense page 0; negative + sentinels and pages absent from the request's row remain `-1`. """ dense_pages_out = torch.full_like(logical_pages, -1) if logical_pages.numel() == 0 or page_inverse.numel() == 0: return dense_pages_out + capacity = int(page_inverse.shape[-1]) + batch_rows = int(page_inverse.shape[0]) pages_long = logical_pages.to(torch.long) + req_long = page_req_id.to(torch.long) + while req_long.dim() < pages_long.dim(): + req_long = req_long.unsqueeze(-1) valid_pages = pages_long >= 0 safe_pages = torch.where(valid_pages, pages_long, torch.zeros_like(pages_long)) - pages_in_range = safe_pages < page_inverse.numel() - clamped_pages = torch.clamp(safe_pages, max=page_inverse.numel() - 1) - dense_pages = page_inverse[clamped_pages] - mapped = valid_pages & pages_in_range & (dense_pages >= 0) + pages_in_range = safe_pages < capacity + req_in_range = (req_long >= 0) & (req_long < batch_rows) + clamped_pages = torch.clamp(safe_pages, max=max(capacity - 1, 0)) + clamped_req = torch.clamp(req_long, min=0, max=max(batch_rows - 1, 0)) + dense_pages = page_inverse[clamped_req, clamped_pages] + mapped = valid_pages & pages_in_range & req_in_range & (dense_pages >= 0) if cp_shared_kv_debug_enabled() and torch.any( - valid_pages & pages_in_range & (dense_pages < 0) + valid_pages & pages_in_range & req_in_range & (dense_pages < 0) ): - missing_pages = pages_long[valid_pages & pages_in_range & (dense_pages < 0)] + missing_mask = valid_pages & pages_in_range & req_in_range & (dense_pages < 0) + missing_pages = pages_long[missing_mask] raise RuntimeError( "CP shared KV slot remap got logical pages outside remap_logical_pages. " f"missing_page_min={int(missing_pages.min().item())} " @@ -3377,7 +3597,14 @@ def remap_logical_pages_to_shared_paged_slot_dense_pages( slot_remap: SharedPagedBufferSlotRemap, *, row_ids: torch.Tensor | None = None, + page_req_id: torch.Tensor | None = None, ) -> torch.Tensor: + if page_req_id is not None: + return remap_logical_pages_to_slot_dense_pages( + logical_pages, + page_inverse=slot_remap.page_inverse, + page_req_id=page_req_id, + ) row_scoped = remap_logical_pages_to_slot_dense_pages_by_row( logical_pages, slot_remap.slot_sorted_logical_pages_by_row, @@ -3389,6 +3616,7 @@ def remap_logical_pages_to_shared_paged_slot_dense_pages( return remap_logical_pages_to_slot_dense_pages( logical_pages, page_inverse=slot_remap.page_inverse, + page_req_id=build_page_table_row_req_id(logical_pages), ) @@ -3422,10 +3650,14 @@ def build_shared_token_kv_slot_remap( kv_cache.shape[0] // page_size, layout, ) + batch_rows = ( + int(remap_logical_pages.shape[0]) if remap_logical_pages.dim() >= 2 else 1 + ) slot_logical_pages, slot_dense_pages = build_slot_page_remap(remap_logical_pages) page_inverse = build_slot_page_inverse_optimized( slot_logical_pages, logical_page_capacity=logical_page_capacity, + batch_rows=batch_rows, ) sorted_logical_pages_by_row, sorted_dense_pages_by_row = ( build_slot_page_sorted_index_by_row( @@ -3433,22 +3665,16 @@ def build_shared_token_kv_slot_remap( slot_dense_pages, ) ) - dense_locs = ( - remap_logical_locs_to_slot_dense_locs_by_row( - logical_locs, - sorted_logical_pages_by_row, - sorted_dense_pages_by_row, - page_size=page_size, - ) - if logical_locs is not None - else None - ) - if dense_locs is None and logical_locs is not None: - dense_locs = remap_logical_locs_to_slot_dense_locs_optimized( - logical_locs, - page_inverse=page_inverse, - page_size=page_size, + # ``dense_locs`` is never consumed downstream (the orchestrators remap their + # own per-request locs against ``page_inverse``). ``get_or_build_*`` always + # passes ``logical_locs=None``; keep it None to avoid needing a per-loc req id + # in the cached, loc-agnostic remap. + if logical_locs is not None: + raise NotImplementedError( + "build_shared_token_kv_slot_remap no longer materializes dense_locs; " + "remap per-request locs at the materialize call site instead." ) + dense_locs = None return SharedTokenKVSlotRemap( slot_logical_pages=slot_logical_pages, page_inverse=page_inverse, @@ -3478,11 +3704,13 @@ def build_shared_paged_buffer_slot_remap( layout=layout, physical_page_capacity=page_buffer.shape[0], ) + batch_rows = int(logical_pages.shape[0]) if logical_pages.dim() >= 2 else 1 slot_logical_pages, dense_pages = build_slot_page_remap(logical_pages) logical_page_capacity = max(int(page_buffer.shape[0]) - 1, 0) * (layout.cp_size) + 1 page_inverse = build_slot_page_inverse_optimized( slot_logical_pages, logical_page_capacity=logical_page_capacity, + batch_rows=batch_rows, ) sorted_logical_pages_by_row, sorted_dense_pages_by_row = ( build_slot_page_sorted_index_by_row( @@ -3958,6 +4186,8 @@ def materialize_prefix_and_reuse_current_kv_page_slots( layout: CpSharedKVLayout, page_size: int, prefix_pages: int, + loc_req_id: torch.Tensor, + current_req_id: torch.Tensor, prefix_slot_span: tuple[int, int] | None = None, prefix_slot_spans: list[tuple[int, int]] | None = None, current_slot_spans: list[tuple[int, int]] | None = None, @@ -4069,9 +4299,9 @@ def materialize_prefix_and_reuse_current_kv_page_slots( ) dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs( logical_locs, - page_size=page_size, slot_remap=slot_remap, - row_ids=logical_locs_row_ids, + page_size=page_size, + loc_req_id=loc_req_id, ) mixed_kv_cache, mixed_locs, _ = fill_current_kv_page_slots_and_remap_locs( dense_kv_cache=dense_kv_cache, @@ -4081,6 +4311,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots( current_locs=current_locs, page_inverse=slot_remap.page_inverse, page_size=page_size, + current_req_id=current_req_id, mask_non_current_in_current_pages=True, ) if layout.cp_size > 1: @@ -4149,6 +4380,7 @@ def materialize_prefix_and_reuse_current_index_page_slots( page_size: int, index_head_dim: int, prefix_pages: int, + current_req_id: torch.Tensor, prefix_slot_span: tuple[int, int] | None = None, prefix_slot_spans: list[tuple[int, int]] | None = None, current_slot_spans: list[tuple[int, int]] | None = None, @@ -4245,6 +4477,7 @@ def materialize_prefix_and_reuse_current_index_page_slots( page_inverse=slot_remap.page_inverse, page_size=page_size, index_head_dim=index_head_dim, + current_req_id=current_req_id, ) if layout.cp_size > 1: if current_slot_spans is None: @@ -4628,6 +4861,7 @@ def materialize_shared_token_kv_buffer( logical_locs: torch.Tensor, layout: CpSharedKVLayout, page_size: int, + loc_req_id: torch.Tensor | None = None, remap_logical_locs: torch.Tensor | None = None, remap_logical_pages: torch.Tensor | None = None, slot_remap: SharedTokenKVSlotRemap | None = None, @@ -4635,6 +4869,13 @@ def materialize_shared_token_kv_buffer( nvtx_source: str = "mla.full_materialize", nvtx_layer_id: int | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: + """Materialize a dense shared token-KV buffer + remapped locs. + + The slot-remap and remap-logical-pages branches map locs through the + per-request ``page_inverse`` and therefore require ``loc_req_id`` (the owning + request of each loc). The ``remap_logical_pages is None`` branch uses the + unique-page dense remap (no per-request inverse) and ignores ``loc_req_id``. + """ _debug_assert_no_tensor_values_below( logical_locs, min_allowed=-1, @@ -4663,12 +4904,17 @@ def materialize_shared_token_kv_buffer( dense_kv_cache = None if slot_remap is not None: + if loc_req_id is None: + raise ValueError( + "materialize_shared_token_kv_buffer slot_remap path requires " + "loc_req_id for the per-request page-inverse remap." + ) materialized_logical_pages = slot_remap.slot_logical_pages dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs( logical_locs, - page_size=page_size, slot_remap=slot_remap, - row_ids=logical_locs_row_ids, + page_size=page_size, + loc_req_id=loc_req_id, ) use_slot_materialize = True elif remap_logical_pages is None: @@ -4681,6 +4927,11 @@ def materialize_shared_token_kv_buffer( ) use_slot_materialize = False else: + if loc_req_id is None: + raise ValueError( + "materialize_shared_token_kv_buffer remap_logical_pages path " + "requires loc_req_id for the per-request page-inverse remap." + ) _debug_assert_no_negative_tensor_values( remap_logical_pages, context="CP shared KV token materialize page remap", @@ -4695,12 +4946,11 @@ def materialize_shared_token_kv_buffer( kv_cache.shape[0] // page_size, layout, ) - tai_result = None - row_scoped_remap_required = ( - logical_locs_row_ids is not None - or (remap_logical_pages.dim() == 2 and logical_locs.dim() >= 2) + materialize_batch_rows = ( + int(remap_logical_pages.shape[0]) if remap_logical_pages.dim() >= 2 else 1 ) - if _tai_materialize_runtime_enabled() and not row_scoped_remap_required: + tai_result = None + if _tai_materialize_runtime_enabled(): materialized_logical_pages = remap_logical_pages.reshape(-1) tai_result = _try_tai_materialize_token_kv_pages_and_locs( kv_cache=kv_cache, @@ -4709,6 +4959,8 @@ def materialize_shared_token_kv_buffer( logical_page_capacity=logical_page_capacity, layout=layout, page_size=page_size, + batch_rows=materialize_batch_rows, + loc_req_id=loc_req_id, ) if tai_result is None: materialized_logical_pages, slot_dense_pages = build_slot_page_remap( @@ -4717,26 +4969,14 @@ def materialize_shared_token_kv_buffer( page_inverse = build_slot_page_inverse_optimized( materialized_logical_pages, logical_page_capacity=logical_page_capacity, + batch_rows=materialize_batch_rows, ) - sorted_logical_pages_by_row, sorted_dense_pages_by_row = ( - build_slot_page_sorted_index_by_row( - remap_logical_pages, - slot_dense_pages, - ) - ) - dense_locs = remap_logical_locs_to_slot_dense_locs_by_row( + dense_locs = remap_logical_locs_to_slot_dense_locs_optimized( logical_locs, - sorted_logical_pages_by_row, - sorted_dense_pages_by_row, + page_inverse=page_inverse, page_size=page_size, - row_ids=logical_locs_row_ids, + loc_req_id=loc_req_id, ) - if dense_locs is None: - dense_locs = remap_logical_locs_to_slot_dense_locs_optimized( - logical_locs, - page_inverse=page_inverse, - page_size=page_size, - ) use_slot_materialize = True else: dense_kv_cache, dense_locs = tai_result diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index f1f6603e1..300fdef5a 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -73,6 +73,7 @@ from sglang.srt.layers.attention.nsa.utils import ( cp_split_and_rebuild_data, cp_shared_kv_bs_gt1_debug_enabled, get_cp_shared_kv_batch_plan, + get_cp_shared_kv_local_current_req_id, get_cp_shared_kv_local_out_cache_loc, get_cp_shared_kv_local_physical_out_cache_loc, is_nsa_enable_prefill_cp, @@ -603,6 +604,13 @@ class Indexer(MultiPlatformOp): f"current_scale_shape={tuple(current_index_kv[1].shape)} " f"out_cache_loc_shape={tuple(current_locs.shape)}" ) + # Owning request id per current row, aligned with current_locs (the + # local out_cache_loc), for the bs>1 per-request page-inverse fill. + current_req_id = get_cp_shared_kv_local_current_req_id(forward_batch) + if current_req_id is None: + current_req_id = torch.zeros_like(current_locs, dtype=torch.long) + else: + current_req_id = current_req_id[: int(current_locs.shape[0])] prefix_slot_spans = None current_slot_spans = build_batch_current_slot_spans( logical_pages=logical_page_table, @@ -628,6 +636,7 @@ class Indexer(MultiPlatformOp): current_locs=current_locs, page_size=page_size, index_head_dim=forward_batch.token_to_kv_pool.index_head_dim, + current_req_id=current_req_id, ) if prefetched is not None: return prefetched @@ -673,6 +682,7 @@ class Indexer(MultiPlatformOp): page_size=page_size, index_head_dim=forward_batch.token_to_kv_pool.index_head_dim, prefix_pages=prefix_pages, + current_req_id=current_req_id, prefix_slot_spans=prefix_slot_spans, current_slot_spans=current_slot_spans, layer_id=layer_id, diff --git a/python/sglang/srt/layers/attention/nsa/utils.py b/python/sglang/srt/layers/attention/nsa/utils.py index af0fd36ef..735d5fdb5 100644 --- a/python/sglang/srt/layers/attention/nsa/utils.py +++ b/python/sglang/srt/layers/attention/nsa/utils.py @@ -1246,6 +1246,35 @@ def select_cp_current_valid_rows_for_reuse( ) +def cp_localize_current_kv_to_compute_rows(forward_batch, tensor: torch.Tensor): + """Re-localize a current-KV tensor to this CP rank's compute rows. + + Under CP shared-KV current reuse, ``forward_absorb_prepare`` calls + ``rebuild_cp_kv_cache`` which all-gathers the current KV back to the full + sequence (global token order) whenever ``should_reuse_current_extend_kv`` is + true. The compute-padding cache-write path + (``select_cp_local_valid_rows_for_cache_write``) requires this rank's per-rank + *compute* rows, so a global tensor must first be re-split through the batch + plan (``cp_split_and_rebuild_data`` uses ``split_kind="compute"`` -- exactly + the rows the valid-row selector expects before it drops padding to valid + rows). A tensor that is already CP-local (the rebuild did not fire) is + returned unchanged. + + This mirrors the re-split the non-compute-padding branch performs in + ``nsa_backend``; the only difference is that this targets the compute split + rather than the valid split. + """ + plan = get_cp_shared_kv_batch_plan(forward_batch) + if plan is None or not bool(getattr(plan, "compute_padding_enabled", False)): + return tensor + _, expected_compute_rows = _get_cp_local_valid_row_indices_cache( + forward_batch, plan, tensor.device + ) + if int(tensor.shape[0]) == int(expected_compute_rows): + return tensor + return cp_split_and_rebuild_data(forward_batch, tensor.contiguous()) + + def _pad_cp_request_tensor_for_split( tensor: torch.Tensor, *, @@ -2037,6 +2066,59 @@ def get_cp_shared_kv_local_out_cache_loc(forward_batch: "ForwardBatch"): return local_out_cache_loc +def get_cp_shared_kv_local_current_req_id(forward_batch: "ForwardBatch"): + """Per-row owning-request id aligned with ``get_cp_shared_kv_local_out_cache_loc``. + + ``current_locs`` (this rank's local out_cache_loc) is request-major but + CP-zigzag-interleaved within each request and only holds this rank's valid + rows, so a naive ``cumsum`` over ``extend_seq_lens`` would NOT line up. Build + the owning-request id by threading a per-token req-id companion through the + IDENTICAL ``split_tensor_by_cp_batch_plan(..., split_kind="valid")`` used to + build ``current_locs`` -- it is therefore aligned by construction. Used to + route each current row through its own request's per-request page-inverse row + (the bs>1 cross-request KV contamination fix). + """ + + cached = getattr(forward_batch, "cp_local_current_req_id", None) + if cached is not None: + return cached + + local_out_cache_loc = get_cp_shared_kv_local_out_cache_loc(forward_batch) + if local_out_cache_loc is None: + return None + + device = local_out_cache_loc.device + batch_plan = get_cp_shared_kv_batch_plan(forward_batch) + if batch_plan is None or local_out_cache_loc.numel() == 0: + # bs=1 / non-batched: every row belongs to request 0. + req_id = torch.zeros_like(local_out_cache_loc, dtype=torch.long) + forward_batch.cp_local_current_req_id = req_id + return req_id + + request_extend_lens = [ + int(x) for x in getattr(batch_plan, "request_extend_lens", []) or [] + ] + total = sum(request_extend_lens) + nreq = len(request_extend_lens) + if nreq <= 1 or total == 0: + req_id = torch.zeros_like(local_out_cache_loc, dtype=torch.long) + forward_batch.cp_local_current_req_id = req_id + return req_id + + global_req_id = torch.repeat_interleave( + torch.arange(nreq, device=device, dtype=torch.long), + torch.tensor(request_extend_lens, device=device, dtype=torch.long), + ) + local_req_id = split_tensor_by_cp_batch_plan( + global_req_id.contiguous(), + batch_plan, + mode="1d", + split_kind="valid", + ) + forward_batch.cp_local_current_req_id = local_req_id + return local_req_id + + def get_cp_shared_kv_local_physical_out_cache_loc(forward_batch: "ForwardBatch"): """Return cached physical rows for this CP rank's shared-KV direct writes. diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index a9b8f2993..b0c0af69c 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -29,6 +29,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( current_extend_kv_rows_for_reuse, current_loc_remap_fast_path_args, filter_owned_logical_locs, + get_cp_shared_kv_token_loc_req_id, get_or_build_shared_token_kv_slot_remap, is_current_only_extend_batch, is_packed_fp8_mla_kv_cache, @@ -53,10 +54,12 @@ from sglang.srt.layers.attention.nsa.transform_index import ( ) from sglang.srt.layers.attention.nsa.utils import ( can_nsa_prefill_cp_round_robin_split, + cp_localize_current_kv_to_compute_rows, cp_split_and_rebuild_data, compute_nsa_seqlens, cp_shared_kv_bs_gt1_debug_enabled, get_cp_shared_kv_batch_plan, + get_cp_shared_kv_local_current_req_id, get_cp_shared_kv_local_out_cache_loc, is_nsa_enable_prefill_cp, log_cp_shared_kv_bs_gt1_debug, @@ -2028,6 +2031,23 @@ class NativeSparseAttnBackend( current_remap_page_size, current_remap_logical_page_capacity = ( current_loc_remap_fast_path_args(forward_batch) ) + # Per-request ids for the bs>1 per-request page-inverse materialize. + # logical_page_table_1 is token-major, contiguous per request by + # nsa_extend_seq_lens_list (the lengths transform_index used); + # current_req_id_full tracks the local out_cache_loc rows and is + # sliced to whatever current_locs_for_reuse is passed below. + cp_loc_req_id = get_cp_shared_kv_token_loc_req_id( + forward_batch, + logical_page_table_1, + metadata.nsa_extend_seq_lens_list, + ) + cp_current_req_id_full = get_cp_shared_kv_local_current_req_id( + forward_batch + ) + if cp_current_req_id_full is None: + cp_current_req_id_full = torch.zeros_like( + current_locs_for_reuse, dtype=torch.long + ) if is_current_only_extend_batch(forward_batch): eagle_draft_mla_branch = "current_only" @@ -2085,6 +2105,10 @@ class NativeSparseAttnBackend( layout=forward_batch.cp_shared_kv_layout, page_size=page_size, prefix_pages=0, + loc_req_id=cp_loc_req_id, + current_req_id=cp_current_req_id_full[ + : int(current_locs_for_reuse.shape[0]) + ], current_slot_spans=current_slot_spans, layer_id=layer.layer_id, nvtx_source="mla.current_only_page_slots", @@ -2116,6 +2140,10 @@ class NativeSparseAttnBackend( logical_locs=logical_page_table_1, current_kv_cache=current_kv_cache, current_locs=current_locs_for_reuse, + loc_req_id=cp_loc_req_id, + current_req_id=cp_current_req_id_full[ + : int(current_locs_for_reuse.shape[0]) + ], current_remap_page_size=current_remap_page_size, current_remap_logical_page_capacity=current_remap_logical_page_capacity, ) @@ -2215,6 +2243,10 @@ class NativeSparseAttnBackend( layout=forward_batch.cp_shared_kv_layout, page_size=page_size, prefix_pages=prefix_pages, + loc_req_id=cp_loc_req_id, + current_req_id=cp_current_req_id_full[ + : int(current_locs_for_reuse.shape[0]) + ], prefix_slot_spans=prefix_slot_spans, current_slot_spans=current_slot_spans, layer_id=layer.layer_id, @@ -2288,6 +2320,11 @@ class NativeSparseAttnBackend( layer_id=layer.layer_id, kv_cache=kv_cache, logical_locs=page_table_1, + loc_req_id=get_cp_shared_kv_token_loc_req_id( + forward_batch, + page_table_1, + metadata.nsa_extend_seq_lens_list, + ), ) if prefetched_kv is not None: eagle_draft_mla_branch = "full_materialize_prefetch" @@ -2318,6 +2355,11 @@ class NativeSparseAttnBackend( kv_cache, page_table_1 = materialize_shared_token_kv_buffer( kv_cache=kv_cache, logical_locs=page_table_1, + loc_req_id=get_cp_shared_kv_token_loc_req_id( + forward_batch, + page_table_1, + metadata.nsa_extend_seq_lens_list, + ), remap_logical_locs=metadata.page_table_1, remap_logical_pages=metadata.real_page_table, layout=forward_batch.cp_shared_kv_layout, @@ -2563,6 +2605,17 @@ class NativeSparseAttnBackend( f"page_table_1_flattened={int(page_table_1_flattened.numel())} " f"indexer_seq_lens_cpu={metadata.indexer_seq_lens_cpu}" ) + ragged_current_req_id = ( + get_cp_shared_kv_local_current_req_id(forward_batch) + ) + if ragged_current_req_id is None: + ragged_current_req_id = torch.zeros_like( + current_locs_for_reuse, dtype=torch.long + ) + else: + ragged_current_req_id = ragged_current_req_id[ + : int(current_locs_for_reuse.shape[0]) + ] slot_remap = get_or_build_shared_token_kv_slot_remap( forward_batch, kv_cache=kv_cache, @@ -2582,7 +2635,8 @@ class NativeSparseAttnBackend( prefix_pages=prefix_pages, prefix_slot_spans=prefix_slot_spans, current_slot_spans=current_slot_spans, - logical_locs_row_ids=logical_locs_row_ids, + loc_req_id=logical_locs_row_ids, + current_req_id=ragged_current_req_id, layer_id=layer.layer_id, nvtx_source="mla.ragged_partial_current_sync", )