fix(nsa): correct bs>1 CP shared-KV KV per-request page_inverse + CP/EAGLE compute-padding

Two correctness fixes in the bs>1 CP shared-KV prefill path, validated end-to-end on B300 (GLM-5.1-FP8, EAGLE, fp8 KV, CP8):

1) Cross-request KV contamination: the value-keyed 1-D page_inverse was last-writer-wins, so two requests sharing a physical page (radix shared-prefix / current-reuse) aliased each other's KV. Replace with a per-request flat page_inverse[batch_rows*capacity] indexed req_id*capacity+lp, threaded through build/remap/fill_current/prefetch and the TAI/Triton kernels. req_id sourced via slot//pages_per_request (build), repeat_interleave (logical_locs), and a companion tensor through the same CP valid-split (current_locs).

2) CP+EAGLE compute-padding cache-write: forward_absorb_prepare's rebuild_cp_kv_cache all-gathers current k back to global rows under current-reuse, but the compute-padding branch fed that full k straight to select_cp_local_valid_rows_for_cache_write (which requires per-rank compute rows) -> fail-fast. Add cp_localize_current_kv_to_compute_rows to re-localize via the batch-plan compute split first (mirrors the non-padding branch).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
(cherry picked from commit 28efef91b05ede1a7072ee451e2aea39ecc3a5bd)
This commit is contained in:
2026-06-09 00:01:04 +08:00
committed by laoyao0822
parent 81a0191331
commit 165b6a01f1
5 changed files with 487 additions and 104 deletions
@@ -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(
@@ -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
@@ -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,
@@ -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.
@@ -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",
)