Preserve request-slot identity for CP cache-hit reuse
Warm-cache bs>1 requests can carry duplicate logical page ids across request rows. A batch-global page inverse collapses those request slots and can alias current/prefix KV across requests. The runtime now carries compact row-scoped sorted page descriptors and remaps flattened logical locs with request row ids where needed, while retaining the legacy global inverse for no-row-context paths. Constraint: Avoid the rejected dense [batch, logical_page_capacity] inverse because bs up to 10 and long contexts make that memory cost unacceptable Rejected: Keep global page_inverse for bs>1 duplicate pages | it is lossy and matches the GSM8K warm-cache corruption shape Rejected: Allocate page_inverse_by_row | correctness-safe but too much GPU memory for production Confidence: medium Scope-risk: moderate Directive: Any future TAI materialize fast path for bs>1 duplicate pages must consume row-scoped descriptors or an equivalent request-row key Tested: Local py_compile for touched runtime files; remote py_compile; remote pytest test_cp_shared_kv_runtime.py, test_nsa_cp_utils.py, test_cp_shared_kv_layout.py => 263 passed, 5 warnings, 2 subtests passed Not-tested: Full GSM8K ETE after this cleanup pass
This commit is contained in:
@@ -25,8 +25,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_slot_dense_pages,
|
||||
remap_logical_locs_to_slot_dense_locs_optimized,
|
||||
remap_logical_pages_to_shared_paged_slot_dense_pages,
|
||||
remap_logical_locs_to_shared_token_slot_dense_locs,
|
||||
slot_range_to_page_slice,
|
||||
slot_range_to_token_slice,
|
||||
)
|
||||
@@ -295,6 +295,23 @@ class CpSharedKVIndexPrefetchHandle:
|
||||
event: Optional[torch.cuda.Event] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PrefetchTokenSlotRemapView:
|
||||
slot_logical_pages: torch.Tensor
|
||||
page_inverse: torch.Tensor
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PrefetchPagedSlotRemapView:
|
||||
slot_logical_pages: torch.Tensor
|
||||
page_inverse: torch.Tensor
|
||||
dense_pages: torch.Tensor | None
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None
|
||||
|
||||
|
||||
class CpSharedKVMlaPrefetcher:
|
||||
"""One-layer-ahead MLA prefix materialize prefetch for CP shared KV.
|
||||
|
||||
@@ -311,6 +328,8 @@ class CpSharedKVMlaPrefetcher:
|
||||
prefix_pages: int,
|
||||
slot_logical_pages: torch.Tensor,
|
||||
page_inverse: torch.Tensor,
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None = None,
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None = None,
|
||||
dense_num_pages: int,
|
||||
owned_prefix_pages: int = -1,
|
||||
owned_total_pages: int = -1,
|
||||
@@ -321,6 +340,8 @@ class CpSharedKVMlaPrefetcher:
|
||||
self.prefix_pages = prefix_pages
|
||||
self.slot_logical_pages = slot_logical_pages
|
||||
self.page_inverse = page_inverse
|
||||
self.slot_sorted_logical_pages_by_row = slot_sorted_logical_pages_by_row
|
||||
self.slot_sorted_dense_pages_by_row = slot_sorted_dense_pages_by_row
|
||||
self.dense_num_pages = dense_num_pages
|
||||
self.owned_prefix_pages = owned_prefix_pages
|
||||
self.owned_total_pages = owned_total_pages
|
||||
@@ -528,6 +549,8 @@ class CpSharedKVMlaPrefetcher:
|
||||
prefix_pages=prefix_pages,
|
||||
slot_logical_pages=remap.slot_logical_pages,
|
||||
page_inverse=remap.page_inverse,
|
||||
slot_sorted_logical_pages_by_row=remap.slot_sorted_logical_pages_by_row,
|
||||
slot_sorted_dense_pages_by_row=remap.slot_sorted_dense_pages_by_row,
|
||||
dense_num_pages=remap.dense_num_pages,
|
||||
owned_prefix_pages=owned_prefix_pages,
|
||||
owned_total_pages=owned_total_pages,
|
||||
@@ -657,9 +680,14 @@ class CpSharedKVMlaPrefetcher:
|
||||
layout=self.layout,
|
||||
physical_token_capacity=kv_cache.shape[0],
|
||||
)
|
||||
dense_locs = remap_logical_locs_to_slot_dense_locs_optimized(
|
||||
dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs(
|
||||
logical_locs,
|
||||
page_inverse=self.page_inverse,
|
||||
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_size=self.page_size,
|
||||
)
|
||||
remap_ms = _cpu_timing_ms(remap_cpu)
|
||||
@@ -757,9 +785,14 @@ class CpSharedKVMlaPrefetcher:
|
||||
layout=self.layout,
|
||||
physical_token_capacity=kv_cache.shape[0],
|
||||
)
|
||||
dense_locs = remap_logical_locs_to_slot_dense_locs_optimized(
|
||||
dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs(
|
||||
logical_locs,
|
||||
page_inverse=self.page_inverse,
|
||||
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_size=self.page_size,
|
||||
)
|
||||
mixed_kv_cache, mixed_locs, _ = fill_current_kv_page_slots_and_remap_locs(
|
||||
@@ -1078,6 +1111,8 @@ class CpSharedKVIndexPrefetcher:
|
||||
prefix_pages: int,
|
||||
slot_logical_pages: torch.Tensor,
|
||||
page_inverse: torch.Tensor,
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None = None,
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None = None,
|
||||
dense_num_pages: int,
|
||||
owned_prefix_pages: int = -1,
|
||||
owned_total_pages: int = -1,
|
||||
@@ -1087,6 +1122,8 @@ class CpSharedKVIndexPrefetcher:
|
||||
self.prefix_pages = prefix_pages
|
||||
self.slot_logical_pages = slot_logical_pages
|
||||
self.page_inverse = page_inverse
|
||||
self.slot_sorted_logical_pages_by_row = slot_sorted_logical_pages_by_row
|
||||
self.slot_sorted_dense_pages_by_row = slot_sorted_dense_pages_by_row
|
||||
self.dense_num_pages = dense_num_pages
|
||||
self.owned_prefix_pages = owned_prefix_pages
|
||||
self.owned_total_pages = owned_total_pages
|
||||
@@ -1332,6 +1369,8 @@ class CpSharedKVIndexPrefetcher:
|
||||
prefix_pages=prefix_pages,
|
||||
slot_logical_pages=remap.slot_logical_pages,
|
||||
page_inverse=remap.page_inverse,
|
||||
slot_sorted_logical_pages_by_row=remap.slot_sorted_logical_pages_by_row,
|
||||
slot_sorted_dense_pages_by_row=remap.slot_sorted_dense_pages_by_row,
|
||||
dense_num_pages=remap.dense_num_pages,
|
||||
owned_prefix_pages=owned_prefix_pages,
|
||||
owned_total_pages=owned_total_pages,
|
||||
@@ -1459,9 +1498,15 @@ class CpSharedKVIndexPrefetcher:
|
||||
layout=self.layout,
|
||||
physical_page_capacity=page_buffer.shape[0],
|
||||
)
|
||||
dense_pages = remap_logical_pages_to_slot_dense_pages(
|
||||
dense_pages = remap_logical_pages_to_shared_paged_slot_dense_pages(
|
||||
logical_pages,
|
||||
page_inverse=self.page_inverse,
|
||||
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,
|
||||
),
|
||||
)
|
||||
remap_ms = _cpu_timing_ms(remap_cpu)
|
||||
total_ms = _cpu_timing_ms(consume_cpu)
|
||||
@@ -1553,9 +1598,15 @@ class CpSharedKVIndexPrefetcher:
|
||||
dense_page_buffer = handle.dense_page_buffer
|
||||
|
||||
remap_cpu = _cpu_timing_start()
|
||||
dense_pages = remap_logical_pages_to_slot_dense_pages(
|
||||
dense_pages = remap_logical_pages_to_shared_paged_slot_dense_pages(
|
||||
logical_pages,
|
||||
page_inverse=self.page_inverse,
|
||||
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,
|
||||
),
|
||||
)
|
||||
dense_page_buffer = fill_current_index_page_slots(
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
|
||||
@@ -191,6 +191,9 @@ class SharedTokenKVSlotRemap:
|
||||
dense_locs: torch.Tensor | None
|
||||
logical_page_capacity: int
|
||||
dense_num_pages: int
|
||||
slot_dense_pages: torch.Tensor | None = None
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None = None
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -200,6 +203,8 @@ class SharedPagedBufferSlotRemap:
|
||||
dense_pages: torch.Tensor
|
||||
logical_page_capacity: int
|
||||
dense_num_pages: int
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None = None
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None = None
|
||||
|
||||
|
||||
def _tensor_identity_key(tensor: torch.Tensor) -> tuple[int, tuple[int, ...], str, str]:
|
||||
@@ -2974,6 +2979,63 @@ def build_slot_page_inverse(
|
||||
return page_inverse
|
||||
|
||||
|
||||
def build_slot_page_sorted_index_by_row(
|
||||
logical_pages: torch.Tensor,
|
||||
dense_pages: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor] | tuple[None, None]:
|
||||
"""Build compact row-scoped logical_page -> dense_page descriptors.
|
||||
|
||||
The global slot inverse is one-dimensional and cannot represent the same
|
||||
logical page appearing in multiple request rows. A dense
|
||||
[batch, logical_page_capacity] inverse fixes correctness but wastes GPU
|
||||
memory. Keep only the actual batch page table, sorted per request row, so
|
||||
row-aware remap can use searchsorted with O(actual_batch_pages) storage.
|
||||
"""
|
||||
|
||||
if logical_pages.dim() != 2 or dense_pages.dim() != 2:
|
||||
return None, None
|
||||
if logical_pages.shape != dense_pages.shape:
|
||||
raise ValueError(
|
||||
"CP shared KV row-scoped page index expects matching logical/dense "
|
||||
f"page shapes, got logical={tuple(logical_pages.shape)} "
|
||||
f"dense={tuple(dense_pages.shape)}"
|
||||
)
|
||||
|
||||
pages_long = logical_pages.to(torch.long)
|
||||
dense_long = dense_pages.to(torch.long)
|
||||
if pages_long.numel() == 0:
|
||||
return pages_long, dense_long
|
||||
|
||||
valid_pages = pages_long > 0
|
||||
invalid_sort_key = torch.full_like(
|
||||
pages_long,
|
||||
torch.iinfo(torch.long).max,
|
||||
)
|
||||
sort_keys = torch.where(valid_pages, pages_long, invalid_sort_key)
|
||||
order = torch.argsort(sort_keys, dim=1)
|
||||
sorted_pages = torch.gather(pages_long, 1, order)
|
||||
sorted_dense_pages = torch.gather(dense_long, 1, order)
|
||||
sorted_dense_pages = torch.where(
|
||||
sorted_pages > 0,
|
||||
sorted_dense_pages,
|
||||
torch.full_like(sorted_dense_pages, -1),
|
||||
)
|
||||
return sorted_pages, sorted_dense_pages
|
||||
|
||||
|
||||
def build_flattened_request_row_ids(
|
||||
seq_lens_cpu: Any,
|
||||
*,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
seq_lens = _to_cpu_int_list(seq_lens_cpu, name="seq_lens_cpu")
|
||||
if len(seq_lens) == 0:
|
||||
return torch.empty((0,), device=device, dtype=torch.long)
|
||||
lengths = torch.tensor(seq_lens, device=device, dtype=torch.long)
|
||||
rows = torch.arange(len(seq_lens), device=device, dtype=torch.long)
|
||||
return torch.repeat_interleave(rows, lengths)
|
||||
|
||||
|
||||
def remap_logical_locs_to_slot_dense_locs(
|
||||
logical_locs: torch.Tensor,
|
||||
page_inverse: torch.Tensor,
|
||||
@@ -3012,6 +3074,154 @@ def remap_logical_locs_to_slot_dense_locs(
|
||||
return torch.where(mapped, dense_values, dense_locs)
|
||||
|
||||
|
||||
def remap_logical_locs_to_slot_dense_locs_by_row(
|
||||
logical_locs: torch.Tensor,
|
||||
sorted_logical_pages_by_row: torch.Tensor | None,
|
||||
sorted_dense_pages_by_row: torch.Tensor | None,
|
||||
page_size: int,
|
||||
*,
|
||||
row_ids: torch.Tensor | None = None,
|
||||
) -> torch.Tensor | None:
|
||||
"""Map logical locs through a compact row-scoped slot descriptor."""
|
||||
|
||||
if (
|
||||
sorted_logical_pages_by_row is None
|
||||
or sorted_dense_pages_by_row is None
|
||||
or sorted_logical_pages_by_row.dim() != 2
|
||||
or sorted_dense_pages_by_row.dim() != 2
|
||||
):
|
||||
return None
|
||||
if sorted_logical_pages_by_row.shape != sorted_dense_pages_by_row.shape:
|
||||
raise ValueError(
|
||||
"CP shared KV row-scoped loc remap got mismatched sorted row maps: "
|
||||
f"pages={tuple(sorted_logical_pages_by_row.shape)} "
|
||||
f"dense={tuple(sorted_dense_pages_by_row.shape)}"
|
||||
)
|
||||
if logical_locs.numel() == 0:
|
||||
return torch.full_like(logical_locs, -1)
|
||||
|
||||
rows = int(sorted_logical_pages_by_row.shape[0])
|
||||
if row_ids is None:
|
||||
if logical_locs.dim() < 2:
|
||||
return None
|
||||
if int(logical_locs.shape[0]) != rows:
|
||||
return None
|
||||
dense_locs = torch.full_like(logical_locs, -1)
|
||||
for row in range(rows):
|
||||
dense_locs[row] = _remap_logical_locs_to_slot_dense_locs_one_row(
|
||||
logical_locs[row],
|
||||
sorted_logical_pages_by_row[row],
|
||||
sorted_dense_pages_by_row[row],
|
||||
page_size,
|
||||
)
|
||||
return dense_locs
|
||||
else:
|
||||
row_ids = row_ids.to(device=logical_locs.device, dtype=torch.long)
|
||||
if row_ids.shape != logical_locs.shape:
|
||||
if row_ids.numel() != logical_locs.numel():
|
||||
raise ValueError(
|
||||
"CP shared KV row-scoped loc remap got incompatible row_ids: "
|
||||
f"logical_locs_shape={tuple(logical_locs.shape)} "
|
||||
f"row_ids_shape={tuple(row_ids.shape)}"
|
||||
)
|
||||
row_ids = row_ids.reshape(logical_locs.shape)
|
||||
|
||||
dense_locs = torch.full_like(logical_locs, -1)
|
||||
logical_locs_flat = logical_locs.reshape(-1)
|
||||
row_ids_flat = row_ids.reshape(-1)
|
||||
dense_locs_flat = dense_locs.reshape(-1)
|
||||
for row in range(rows):
|
||||
row_mask = row_ids_flat == row
|
||||
dense_locs_flat[row_mask] = _remap_logical_locs_to_slot_dense_locs_one_row(
|
||||
logical_locs_flat[row_mask],
|
||||
sorted_logical_pages_by_row[row],
|
||||
sorted_dense_pages_by_row[row],
|
||||
page_size,
|
||||
)
|
||||
return dense_locs
|
||||
|
||||
|
||||
def _remap_logical_locs_to_slot_dense_locs_one_row(
|
||||
logical_locs: torch.Tensor,
|
||||
sorted_logical_pages: torch.Tensor,
|
||||
sorted_dense_pages: torch.Tensor,
|
||||
page_size: int,
|
||||
) -> torch.Tensor:
|
||||
dense_locs = torch.full_like(logical_locs, -1)
|
||||
if logical_locs.numel() == 0:
|
||||
return dense_locs
|
||||
|
||||
locs_long = logical_locs.to(torch.long)
|
||||
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)
|
||||
|
||||
dummy = valid_locs & (logical_pages == 0)
|
||||
dense_locs = torch.where(dummy, offsets.to(logical_locs.dtype), dense_locs)
|
||||
|
||||
positive = valid_locs & (logical_pages > 0)
|
||||
flat_positive = positive.reshape(-1)
|
||||
positive_indices = torch.nonzero(flat_positive, as_tuple=False).reshape(-1)
|
||||
if positive_indices.numel() == 0 or sorted_logical_pages.numel() == 0:
|
||||
return dense_locs
|
||||
|
||||
pages_flat = logical_pages.reshape(-1)
|
||||
offsets_flat = offsets.reshape(-1)
|
||||
out_flat = dense_locs.reshape(-1)
|
||||
pages_to_map = pages_flat[positive_indices]
|
||||
positions = torch.searchsorted(sorted_logical_pages, pages_to_map)
|
||||
in_range = positions < int(sorted_logical_pages.numel())
|
||||
safe_positions = torch.clamp(
|
||||
positions,
|
||||
min=0,
|
||||
max=max(int(sorted_logical_pages.numel()) - 1, 0),
|
||||
)
|
||||
found_pages = sorted_logical_pages[safe_positions]
|
||||
dense_pages = sorted_dense_pages[safe_positions]
|
||||
mapped = in_range & (found_pages == pages_to_map) & (dense_pages >= 0)
|
||||
mapped_indices = positive_indices[mapped]
|
||||
if mapped_indices.numel() > 0:
|
||||
dense_values = (
|
||||
dense_pages[mapped].to(logical_locs.dtype) * page_size
|
||||
+ offsets_flat[mapped_indices].to(logical_locs.dtype)
|
||||
)
|
||||
out_flat[mapped_indices] = dense_values
|
||||
|
||||
if cp_shared_kv_debug_enabled() and bool(torch.any(~mapped).item()):
|
||||
missing_pages = pages_to_map[~mapped]
|
||||
raise RuntimeError(
|
||||
"CP shared KV row slot remap got logical locs outside row map. "
|
||||
f"missing_page_min={int(missing_pages.min().item())} "
|
||||
f"missing_page_max={int(missing_pages.max().item())} "
|
||||
f"logical_locs={tensor_debug_summary(logical_locs)}"
|
||||
)
|
||||
return dense_locs
|
||||
|
||||
|
||||
def remap_logical_locs_to_shared_token_slot_dense_locs(
|
||||
logical_locs: torch.Tensor,
|
||||
slot_remap: SharedTokenKVSlotRemap,
|
||||
page_size: int,
|
||||
*,
|
||||
row_ids: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
row_scoped = remap_logical_locs_to_slot_dense_locs_by_row(
|
||||
logical_locs,
|
||||
slot_remap.slot_sorted_logical_pages_by_row,
|
||||
slot_remap.slot_sorted_dense_pages_by_row,
|
||||
page_size,
|
||||
row_ids=row_ids,
|
||||
)
|
||||
if row_scoped is not None:
|
||||
return row_scoped
|
||||
return remap_logical_locs_to_slot_dense_locs_optimized(
|
||||
logical_locs,
|
||||
page_inverse=slot_remap.page_inverse,
|
||||
page_size=page_size,
|
||||
)
|
||||
|
||||
|
||||
def remap_logical_pages_to_slot_dense_pages(
|
||||
logical_pages: torch.Tensor,
|
||||
page_inverse: torch.Tensor,
|
||||
@@ -3049,6 +3259,139 @@ def remap_logical_pages_to_slot_dense_pages(
|
||||
return torch.where(mapped, dense_pages.to(logical_pages.dtype), dense_pages_out)
|
||||
|
||||
|
||||
def remap_logical_pages_to_slot_dense_pages_by_row(
|
||||
logical_pages: torch.Tensor,
|
||||
sorted_logical_pages_by_row: torch.Tensor | None,
|
||||
sorted_dense_pages_by_row: torch.Tensor | None,
|
||||
*,
|
||||
row_ids: torch.Tensor | None = None,
|
||||
) -> torch.Tensor | None:
|
||||
if (
|
||||
sorted_logical_pages_by_row is None
|
||||
or sorted_dense_pages_by_row is None
|
||||
or sorted_logical_pages_by_row.dim() != 2
|
||||
or sorted_dense_pages_by_row.dim() != 2
|
||||
):
|
||||
return None
|
||||
if sorted_logical_pages_by_row.shape != sorted_dense_pages_by_row.shape:
|
||||
raise ValueError(
|
||||
"CP shared KV row-scoped page remap got mismatched sorted row maps: "
|
||||
f"pages={tuple(sorted_logical_pages_by_row.shape)} "
|
||||
f"dense={tuple(sorted_dense_pages_by_row.shape)}"
|
||||
)
|
||||
if logical_pages.numel() == 0:
|
||||
return torch.full_like(logical_pages, -1)
|
||||
|
||||
rows = int(sorted_logical_pages_by_row.shape[0])
|
||||
if row_ids is None:
|
||||
if logical_pages.dim() < 2:
|
||||
return None
|
||||
if int(logical_pages.shape[0]) != rows:
|
||||
return None
|
||||
dense_pages = torch.full_like(logical_pages, -1)
|
||||
for row in range(rows):
|
||||
dense_pages[row] = _remap_logical_pages_to_slot_dense_pages_one_row(
|
||||
logical_pages[row],
|
||||
sorted_logical_pages_by_row[row],
|
||||
sorted_dense_pages_by_row[row],
|
||||
)
|
||||
return dense_pages
|
||||
else:
|
||||
row_ids = row_ids.to(device=logical_pages.device, dtype=torch.long)
|
||||
if row_ids.shape != logical_pages.shape:
|
||||
if row_ids.numel() != logical_pages.numel():
|
||||
raise ValueError(
|
||||
"CP shared KV row-scoped page remap got incompatible row_ids: "
|
||||
f"logical_pages_shape={tuple(logical_pages.shape)} "
|
||||
f"row_ids_shape={tuple(row_ids.shape)}"
|
||||
)
|
||||
row_ids = row_ids.reshape(logical_pages.shape)
|
||||
|
||||
dense_pages = torch.full_like(logical_pages, -1)
|
||||
logical_pages_flat = logical_pages.reshape(-1)
|
||||
row_ids_flat = row_ids.reshape(-1)
|
||||
dense_pages_flat = dense_pages.reshape(-1)
|
||||
for row in range(rows):
|
||||
row_mask = row_ids_flat == row
|
||||
dense_pages_flat[row_mask] = _remap_logical_pages_to_slot_dense_pages_one_row(
|
||||
logical_pages_flat[row_mask],
|
||||
sorted_logical_pages_by_row[row],
|
||||
sorted_dense_pages_by_row[row],
|
||||
)
|
||||
return dense_pages
|
||||
|
||||
|
||||
def _remap_logical_pages_to_slot_dense_pages_one_row(
|
||||
logical_pages: torch.Tensor,
|
||||
sorted_logical_pages: torch.Tensor,
|
||||
sorted_dense_pages: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
dense_pages_out = torch.full_like(logical_pages, -1)
|
||||
if logical_pages.numel() == 0:
|
||||
return dense_pages_out
|
||||
|
||||
pages_long = logical_pages.to(torch.long)
|
||||
valid_pages = pages_long >= 0
|
||||
dummy = valid_pages & (pages_long == 0)
|
||||
dense_pages_out = torch.where(
|
||||
dummy,
|
||||
torch.zeros_like(dense_pages_out),
|
||||
dense_pages_out,
|
||||
)
|
||||
|
||||
positive = valid_pages & (pages_long > 0)
|
||||
positive_indices = torch.nonzero(positive.reshape(-1), as_tuple=False).reshape(-1)
|
||||
if positive_indices.numel() == 0 or sorted_logical_pages.numel() == 0:
|
||||
return dense_pages_out
|
||||
|
||||
pages_flat = pages_long.reshape(-1)
|
||||
out_flat = dense_pages_out.reshape(-1)
|
||||
pages_to_map = pages_flat[positive_indices]
|
||||
positions = torch.searchsorted(sorted_logical_pages, pages_to_map)
|
||||
in_range = positions < int(sorted_logical_pages.numel())
|
||||
safe_positions = torch.clamp(
|
||||
positions,
|
||||
min=0,
|
||||
max=max(int(sorted_logical_pages.numel()) - 1, 0),
|
||||
)
|
||||
found_pages = sorted_logical_pages[safe_positions]
|
||||
dense_pages = sorted_dense_pages[safe_positions]
|
||||
mapped = in_range & (found_pages == pages_to_map) & (dense_pages >= 0)
|
||||
mapped_indices = positive_indices[mapped]
|
||||
if mapped_indices.numel() > 0:
|
||||
out_flat[mapped_indices] = dense_pages[mapped].to(logical_pages.dtype)
|
||||
|
||||
if cp_shared_kv_debug_enabled() and bool(torch.any(~mapped).item()):
|
||||
missing_pages = pages_to_map[~mapped]
|
||||
raise RuntimeError(
|
||||
"CP shared KV row slot remap got logical pages outside row map. "
|
||||
f"missing_page_min={int(missing_pages.min().item())} "
|
||||
f"missing_page_max={int(missing_pages.max().item())} "
|
||||
f"logical_pages={tensor_debug_summary(logical_pages)}"
|
||||
)
|
||||
return dense_pages_out
|
||||
|
||||
|
||||
def remap_logical_pages_to_shared_paged_slot_dense_pages(
|
||||
logical_pages: torch.Tensor,
|
||||
slot_remap: SharedPagedBufferSlotRemap,
|
||||
*,
|
||||
row_ids: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
row_scoped = remap_logical_pages_to_slot_dense_pages_by_row(
|
||||
logical_pages,
|
||||
slot_remap.slot_sorted_logical_pages_by_row,
|
||||
slot_remap.slot_sorted_dense_pages_by_row,
|
||||
row_ids=row_ids,
|
||||
)
|
||||
if row_scoped is not None:
|
||||
return row_scoped
|
||||
return remap_logical_pages_to_slot_dense_pages(
|
||||
logical_pages,
|
||||
page_inverse=slot_remap.page_inverse,
|
||||
)
|
||||
|
||||
|
||||
def build_shared_token_kv_slot_remap(
|
||||
kv_cache: torch.Tensor,
|
||||
logical_locs: torch.Tensor | None,
|
||||
@@ -3079,26 +3422,42 @@ def build_shared_token_kv_slot_remap(
|
||||
kv_cache.shape[0] // page_size,
|
||||
layout,
|
||||
)
|
||||
slot_logical_pages, _ = build_slot_page_remap(remap_logical_pages)
|
||||
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,
|
||||
)
|
||||
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_optimized(
|
||||
remap_logical_locs_to_slot_dense_locs_by_row(
|
||||
logical_locs,
|
||||
page_inverse=page_inverse,
|
||||
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,
|
||||
)
|
||||
return SharedTokenKVSlotRemap(
|
||||
slot_logical_pages=slot_logical_pages,
|
||||
page_inverse=page_inverse,
|
||||
dense_locs=dense_locs,
|
||||
logical_page_capacity=logical_page_capacity,
|
||||
dense_num_pages=int(slot_logical_pages.numel()) + 1,
|
||||
slot_dense_pages=slot_dense_pages,
|
||||
slot_sorted_logical_pages_by_row=sorted_logical_pages_by_row,
|
||||
slot_sorted_dense_pages_by_row=sorted_dense_pages_by_row,
|
||||
)
|
||||
|
||||
|
||||
@@ -3125,12 +3484,20 @@ def build_shared_paged_buffer_slot_remap(
|
||||
slot_logical_pages,
|
||||
logical_page_capacity=logical_page_capacity,
|
||||
)
|
||||
sorted_logical_pages_by_row, sorted_dense_pages_by_row = (
|
||||
build_slot_page_sorted_index_by_row(
|
||||
logical_pages,
|
||||
dense_pages,
|
||||
)
|
||||
)
|
||||
return SharedPagedBufferSlotRemap(
|
||||
slot_logical_pages=slot_logical_pages,
|
||||
page_inverse=page_inverse,
|
||||
dense_pages=dense_pages,
|
||||
logical_page_capacity=logical_page_capacity,
|
||||
dense_num_pages=int(slot_logical_pages.numel()) + 1,
|
||||
slot_sorted_logical_pages_by_row=sorted_logical_pages_by_row,
|
||||
slot_sorted_dense_pages_by_row=sorted_dense_pages_by_row,
|
||||
)
|
||||
|
||||
|
||||
@@ -3594,6 +3961,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
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,
|
||||
logical_locs_row_ids: torch.Tensor | None = None,
|
||||
layer_id: int | None = None,
|
||||
nvtx_source: str = "mla.partial_current_sync",
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
@@ -3699,10 +4067,11 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
layout=layout,
|
||||
physical_token_capacity=kv_cache.shape[0],
|
||||
)
|
||||
dense_locs = remap_logical_locs_to_slot_dense_locs_optimized(
|
||||
dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs(
|
||||
logical_locs,
|
||||
page_inverse=slot_remap.page_inverse,
|
||||
page_size=page_size,
|
||||
slot_remap=slot_remap,
|
||||
row_ids=logical_locs_row_ids,
|
||||
)
|
||||
mixed_kv_cache, mixed_locs, _ = fill_current_kv_page_slots_and_remap_locs(
|
||||
dense_kv_cache=dense_kv_cache,
|
||||
@@ -4262,6 +4631,7 @@ def materialize_shared_token_kv_buffer(
|
||||
remap_logical_locs: torch.Tensor | None = None,
|
||||
remap_logical_pages: torch.Tensor | None = None,
|
||||
slot_remap: SharedTokenKVSlotRemap | None = None,
|
||||
logical_locs_row_ids: torch.Tensor | None = None,
|
||||
nvtx_source: str = "mla.full_materialize",
|
||||
nvtx_layer_id: int | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
@@ -4294,10 +4664,11 @@ def materialize_shared_token_kv_buffer(
|
||||
dense_kv_cache = None
|
||||
if slot_remap is not None:
|
||||
materialized_logical_pages = slot_remap.slot_logical_pages
|
||||
dense_locs = remap_logical_locs_to_slot_dense_locs_optimized(
|
||||
dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs(
|
||||
logical_locs,
|
||||
page_inverse=slot_remap.page_inverse,
|
||||
page_size=page_size,
|
||||
slot_remap=slot_remap,
|
||||
row_ids=logical_locs_row_ids,
|
||||
)
|
||||
use_slot_materialize = True
|
||||
elif remap_logical_pages is None:
|
||||
@@ -4325,7 +4696,11 @@ def materialize_shared_token_kv_buffer(
|
||||
layout,
|
||||
)
|
||||
tai_result = None
|
||||
if _tai_materialize_runtime_enabled():
|
||||
row_scoped_remap_required = (
|
||||
logical_locs_row_ids is not None
|
||||
or (remap_logical_pages.dim() == 2 and logical_locs.dim() >= 2)
|
||||
)
|
||||
if _tai_materialize_runtime_enabled() and not row_scoped_remap_required:
|
||||
materialized_logical_pages = remap_logical_pages.reshape(-1)
|
||||
tai_result = _try_tai_materialize_token_kv_pages_and_locs(
|
||||
kv_cache=kv_cache,
|
||||
@@ -4336,16 +4711,32 @@ def materialize_shared_token_kv_buffer(
|
||||
page_size=page_size,
|
||||
)
|
||||
if tai_result is None:
|
||||
materialized_logical_pages, _ = build_slot_page_remap(remap_logical_pages)
|
||||
materialized_logical_pages, slot_dense_pages = build_slot_page_remap(
|
||||
remap_logical_pages
|
||||
)
|
||||
page_inverse = build_slot_page_inverse_optimized(
|
||||
materialized_logical_pages,
|
||||
logical_page_capacity=logical_page_capacity,
|
||||
)
|
||||
dense_locs = remap_logical_locs_to_slot_dense_locs_optimized(
|
||||
logical_locs,
|
||||
page_inverse=page_inverse,
|
||||
page_size=page_size,
|
||||
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(
|
||||
logical_locs,
|
||||
sorted_logical_pages_by_row,
|
||||
sorted_dense_pages_by_row,
|
||||
page_size=page_size,
|
||||
row_ids=logical_locs_row_ids,
|
||||
)
|
||||
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
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_prefetch import (
|
||||
)
|
||||
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
build_batch_current_slot_spans,
|
||||
build_flattened_request_row_ids,
|
||||
build_batch_prefix_slot_spans,
|
||||
build_current_loc_remap,
|
||||
cp_shared_kv_debug_enabled,
|
||||
@@ -2522,6 +2523,20 @@ class NativeSparseAttnBackend(
|
||||
extend_lens_cpu=extend_lens_cpu,
|
||||
page_size=page_size,
|
||||
)
|
||||
logical_locs_row_ids = build_flattened_request_row_ids(
|
||||
metadata.indexer_seq_lens_cpu,
|
||||
device=page_table_1_flattened.device,
|
||||
)
|
||||
if int(logical_locs_row_ids.numel()) != int(
|
||||
page_table_1_flattened.numel()
|
||||
):
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][mla_ragged_current_sync] "
|
||||
"row-id metadata does not match flattened page table. "
|
||||
f"row_ids={int(logical_locs_row_ids.numel())} "
|
||||
f"page_table_1_flattened={int(page_table_1_flattened.numel())} "
|
||||
f"indexer_seq_lens_cpu={metadata.indexer_seq_lens_cpu}"
|
||||
)
|
||||
slot_remap = get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
@@ -2541,6 +2556,7 @@ 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,
|
||||
layer_id=layer.layer_id,
|
||||
nvtx_source="mla.ragged_partial_current_sync",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user