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:
@@ -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",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user