Reduce CP shared KV materialize and direct-write overhead
Shared KV now relies on page-aligned CP metadata and compute-owner page allocation so persistent MLA KV and NSA index shards can be written by the rank that computed them. The compatibility read path keeps the dense full-view contract for existing topk and attention kernels, but removes duplicated prev/next index materialize, adds optional tai materialize integration, and tightens tests/docs around the fallback boundaries. Constraint: Decode remains non-CP while prefill CP owns the shared-KV changes Constraint: Existing attention/topk kernels still expect dense full-view KV/index inputs Rejected: Change attention kernels to read owner-sharded KV directly | larger semantic change reserved for later phases Rejected: Merge index K/scale storage with MLA KV storage | would couple topk and attention cache lifecycles before materialize overhead is isolated Confidence: medium Scope-risk: broad Directive: Do not remove fallback logging or debug-gated assertions without reproducing long-context chunked/radix-hit paths Tested: git diff --check --cached Not-tested: Local pytest/runtime server verification not run in this commit step per current workflow constraints
This commit is contained in:
@@ -20,7 +20,7 @@ Page-aligned memory pool.
|
||||
"""
|
||||
|
||||
import abc
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
import torch
|
||||
import triton
|
||||
@@ -555,3 +555,130 @@ class CPSharedPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
|
||||
self.physical_size = physical_size
|
||||
self.cp_size = cp_size
|
||||
self.cp_rank = cp_rank
|
||||
|
||||
def compute_owner_lane_stats(
|
||||
self,
|
||||
page_compute_owners: List[int],
|
||||
) -> tuple[List[int], List[int], List[int]]:
|
||||
required = [0 for _ in range(self.cp_size)]
|
||||
for owner in page_compute_owners:
|
||||
if owner < 0 or owner >= self.cp_size:
|
||||
raise ValueError(
|
||||
f"compute owner must be in [0, {self.cp_size}), got {owner}"
|
||||
)
|
||||
required[owner] += 1
|
||||
|
||||
free_pages = self.free_pages
|
||||
if len(self.release_pages) > 0:
|
||||
free_pages = torch.cat((free_pages, self.release_pages))
|
||||
available = [
|
||||
int(
|
||||
(
|
||||
torch.remainder(free_pages - 1, self.cp_size) == owner
|
||||
).sum().item()
|
||||
)
|
||||
for owner in range(self.cp_size)
|
||||
]
|
||||
deficits = [
|
||||
max(0, required_count - available_count)
|
||||
for required_count, available_count in zip(required, available)
|
||||
]
|
||||
return required, available, deficits
|
||||
|
||||
def _select_compute_owner_pages(
|
||||
self,
|
||||
page_compute_owners: List[int],
|
||||
) -> Optional[torch.Tensor]:
|
||||
selected_pages = []
|
||||
lane_offsets = [0 for _ in range(self.cp_size)]
|
||||
lane_pages = [
|
||||
self.free_pages[
|
||||
torch.remainder(self.free_pages - 1, self.cp_size) == owner
|
||||
]
|
||||
for owner in range(self.cp_size)
|
||||
]
|
||||
|
||||
for owner in page_compute_owners:
|
||||
if owner < 0 or owner >= self.cp_size:
|
||||
raise ValueError(
|
||||
f"compute owner must be in [0, {self.cp_size}), got {owner}"
|
||||
)
|
||||
lane_offset = lane_offsets[owner]
|
||||
if lane_offset >= lane_pages[owner].numel():
|
||||
return None
|
||||
selected_pages.append(lane_pages[owner][lane_offset])
|
||||
lane_offsets[owner] = lane_offset + 1
|
||||
|
||||
if not selected_pages:
|
||||
return torch.empty((0,), dtype=torch.int64, device=self.device)
|
||||
return torch.stack(selected_pages).to(torch.int64)
|
||||
|
||||
def alloc_extend_compute_owner(
|
||||
self,
|
||||
prefix_lens: torch.Tensor,
|
||||
prefix_lens_cpu: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
last_loc: torch.Tensor,
|
||||
extend_num_tokens: int,
|
||||
page_compute_owners: List[int],
|
||||
):
|
||||
"""Allocate extend KV locs so logical page owner matches CP compute rank.
|
||||
|
||||
The returned logical `out_cache_loc` is still full-order and identical on
|
||||
every CP rank. Only the chosen logical page ids change: each newly
|
||||
allocated request page comes from the modulo-owner lane that will compute
|
||||
and directly persist that page.
|
||||
"""
|
||||
|
||||
if len(prefix_lens_cpu) != 1 or len(seq_lens_cpu) != 1:
|
||||
raise ValueError("compute-owner allocation supports batch size 1 only")
|
||||
|
||||
num_new_pages = get_num_new_pages(
|
||||
seq_lens=seq_lens_cpu,
|
||||
page_size=self.page_size,
|
||||
prefix_lens=prefix_lens_cpu,
|
||||
)
|
||||
if num_new_pages != len(page_compute_owners):
|
||||
raise ValueError(
|
||||
"compute-owner page count mismatch: "
|
||||
f"{num_new_pages=} page_compute_owners={len(page_compute_owners)}"
|
||||
)
|
||||
|
||||
if self.need_sort and num_new_pages > len(self.free_pages):
|
||||
self.merge_and_sort_free()
|
||||
|
||||
selected_pages = self._select_compute_owner_pages(page_compute_owners)
|
||||
if selected_pages is None and self.need_sort and len(self.release_pages) > 0:
|
||||
self.merge_and_sort_free()
|
||||
selected_pages = self._select_compute_owner_pages(page_compute_owners)
|
||||
if selected_pages is None:
|
||||
return None
|
||||
|
||||
out_indices = torch.empty(
|
||||
(extend_num_tokens,), dtype=torch.int64, device=self.device
|
||||
)
|
||||
alloc_extend_naive(
|
||||
prefix_lens_cpu,
|
||||
seq_lens_cpu,
|
||||
last_loc,
|
||||
selected_pages,
|
||||
out_indices,
|
||||
self.page_size,
|
||||
self.device,
|
||||
)
|
||||
|
||||
selected_mask = torch.isin(self.free_pages, selected_pages)
|
||||
self.free_pages = self.free_pages[~selected_mask]
|
||||
|
||||
if self.debug_mode:
|
||||
assert len(torch.unique(out_indices)) == len(out_indices)
|
||||
selected_owners = torch.remainder(selected_pages - 1, self.cp_size)
|
||||
expected_owners = torch.tensor(
|
||||
page_compute_owners,
|
||||
dtype=selected_owners.dtype,
|
||||
device=selected_owners.device,
|
||||
)
|
||||
assert torch.equal(selected_owners, expected_owners)
|
||||
|
||||
return out_indices
|
||||
|
||||
Reference in New Issue
Block a user