feat(mem-cache): enable CP shared-KV batch owner allocation

Route supported multi-request CP shared-KV extend allocation through request-local owner plans so downstream direct-write paths can trust flattened out_cache_loc owner lanes. Unsupported page plans now fail fast instead of falling back to legacy allocation.

Include the normative govctl RFC and completed work-item trail in the tracked repository; keep transient superpowers planning outside the sglang branch.

Constraint: RFC-0001 W2 requires no legacy multi_batch fallback for supported bs>1 owner-lane allocation
Rejected: concatenating requests before owner planning | CP ownership is request-relative
Rejected: committing only docs/superpowers plan | it omits the normative RFC and work-item trail
Confidence: high
Scope-risk: moderate
Directive: Preserve flattened request order when changing CP shared-KV allocation; run govctl from the sglang repository root
Tested: remote uv test_nsa_cp_utils.py 39 OK; remote uv test_cp_shared_kv_layout.py 34 OK; remote uv test_alloc_pages_with_owners.py 10 OK; govctl check; git diff --check
Not-tested: CUDA ETE/perf paths beyond focused W1/W2 unit suites
This commit is contained in:
wxiwnd
2026-06-03 03:26:08 +08:00
parent e4cf8d18b4
commit 520770da4c
23 changed files with 1431 additions and 19 deletions

View File

@@ -996,9 +996,6 @@ class CPSharedPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
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,

View File

@@ -14,6 +14,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
)
from sglang.srt.mem_cache.allocator import compute_owner_lane_free_room_deficits
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_batch_in_seq_page_compute_owners,
build_in_seq_page_compute_owners,
get_in_seq_page_compute_owner_unavailable_reason,
)
@@ -66,6 +67,27 @@ def _log_cp_shared_kv_alloc_fallback(
)
def _raise_cp_shared_kv_alloc_fail_fast(
*,
reason: str,
batch_size: int,
extend_num_tokens: int,
page_size: int,
) -> None:
message = (
"CP shared-KV compute-owner page assignment is unavailable; "
"rejecting extend allocation instead of falling back to legacy page "
f"allocation. batch_size={batch_size} extend_num_tokens={extend_num_tokens} "
f"page_size={page_size} reason={reason}"
)
logger.error(
"[CP_SHARED_KV_FAIL_FAST][compute_owner_alloc] reason=%s %s",
reason,
message,
)
raise RuntimeError(message)
def _compute_owner_lane_stats_for_eviction(
*,
tree_cache: BasePrefixCache,
@@ -453,7 +475,8 @@ def alloc_paged_token_slots_extend(
)
page_compute_owners = None
compute_owner_unavailable_reason = None
if alloc_extend_compute_owner is not None and len(prefix_lens_cpu) == 1:
compute_owner_planning_required = False
if alloc_extend_compute_owner is not None:
try:
server_args = get_global_server_args()
except ValueError:
@@ -464,28 +487,52 @@ def alloc_paged_token_slots_extend(
and server_args.enable_nsa_prefill_context_parallel
and server_args.nsa_prefill_cp_mode == "in-seq-split"
):
extend_len = int(seq_lens_cpu[0].item() - prefix_lens_cpu[0].item())
page_compute_owners = build_in_seq_page_compute_owners(
extend_len=extend_len,
extend_prefix_len=int(prefix_lens_cpu[0].item()),
page_size=int(allocator.page_size),
cp_size=int(allocator.cp_size),
)
compute_owner_planning_required = True
extend_lens = [
int(seq_lens_cpu[i].item() - prefix_lens_cpu[i].item())
for i in range(len(prefix_lens_cpu))
]
extend_prefix_lens = [
int(prefix_lens_cpu[i].item()) for i in range(len(prefix_lens_cpu))
]
if len(prefix_lens_cpu) == 1:
page_compute_owners = build_in_seq_page_compute_owners(
extend_len=extend_lens[0],
extend_prefix_len=extend_prefix_lens[0],
page_size=int(allocator.page_size),
cp_size=int(allocator.cp_size),
)
else:
page_compute_owners = build_batch_in_seq_page_compute_owners(
extend_lens=extend_lens,
extend_prefix_lens=extend_prefix_lens,
page_size=int(allocator.page_size),
cp_size=int(allocator.cp_size),
)
if page_compute_owners is None:
compute_owner_unavailable_reason = (
get_in_seq_page_compute_owner_unavailable_reason(
for extend_len, extend_prefix_len in zip(
extend_lens, extend_prefix_lens
):
reason = get_in_seq_page_compute_owner_unavailable_reason(
extend_len=extend_len,
extend_prefix_len=int(prefix_lens_cpu[0].item()),
extend_prefix_len=extend_prefix_len,
page_size=int(allocator.page_size),
cp_size=int(allocator.cp_size),
)
or "unknown"
)
if reason is not None:
compute_owner_unavailable_reason = reason
break
if compute_owner_unavailable_reason is None:
compute_owner_unavailable_reason = "unknown"
else:
compute_owner_unavailable_reason = "server_args_not_enabled"
elif alloc_extend_compute_owner is not None:
compute_owner_unavailable_reason = (
"multi_batch" if len(prefix_lens_cpu) != 1 else "unknown"
if page_compute_owners is None and compute_owner_planning_required:
_raise_cp_shared_kv_alloc_fail_fast(
reason=compute_owner_unavailable_reason or "compute_owner_not_available",
batch_size=len(prefix_lens_cpu),
extend_num_tokens=extend_num_tokens,
page_size=int(allocator.page_size),
)
if page_compute_owners is not None:

View File

@@ -75,3 +75,31 @@ def build_in_seq_page_compute_owners(
owners.extend([owner] * unit_count)
return owners
def build_batch_in_seq_page_compute_owners(
*,
extend_lens: list[int],
extend_prefix_lens: list[int],
page_size: int,
cp_size: int,
) -> Optional[List[int]]:
if len(extend_lens) != len(extend_prefix_lens):
raise ValueError(
"extend_lens and extend_prefix_lens must have the same length, "
f"got {len(extend_lens)} and {len(extend_prefix_lens)}"
)
batch_owners: List[int] = []
for extend_len, extend_prefix_len in zip(extend_lens, extend_prefix_lens):
owners = build_in_seq_page_compute_owners(
extend_len=int(extend_len),
extend_prefix_len=int(extend_prefix_len),
page_size=page_size,
cp_size=cp_size,
)
if owners is None:
return None
batch_owners.extend(owners)
return batch_owners