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:
laoyao0822
2026-05-02 07:07:28 +08:00
parent 2317952a01
commit 5769b63082
17 changed files with 2524 additions and 96 deletions

View File

@@ -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