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
@@ -209,6 +209,57 @@ class TestCPSharedPagedAllocator(CustomTestCase):
[0, 0, 1, 1, 2, 2, 3, 3, 3, 3, 2, 2, 1, 1, 0, 0],
)
def test_batch_compute_owner_page_assignment_flattens_request_order(self):
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_batch_in_seq_page_compute_owners,
)
owners = build_batch_in_seq_page_compute_owners(
extend_lens=[64 * 3, 64 * 5],
extend_prefix_lens=[0, 64 * 2],
page_size=64,
cp_size=4,
)
self.assertEqual(owners, [0, 1, 2, 0, 1, 2, 3, 3])
def test_batch_compute_owner_page_assignment_matches_single_request(self):
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_batch_in_seq_page_compute_owners,
build_in_seq_page_compute_owners,
)
batch_owners = build_batch_in_seq_page_compute_owners(
extend_lens=[64 * 7],
extend_prefix_lens=[64 * 2],
page_size=64,
cp_size=4,
)
single_owners = build_in_seq_page_compute_owners(
extend_len=64 * 7,
extend_prefix_len=64 * 2,
page_size=64,
cp_size=4,
)
self.assertEqual(batch_owners, single_owners)
def test_batch_compute_owner_page_assignment_returns_none_for_misaligned_prefix(
self,
):
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_batch_in_seq_page_compute_owners,
)
owners = build_batch_in_seq_page_compute_owners(
extend_lens=[64, 64],
extend_prefix_lens=[0, 1],
page_size=64,
cp_size=4,
)
self.assertIsNone(owners)
def test_compute_owner_page_assignment_keeps_page_aligned_short_extend(self):
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,
@@ -355,6 +406,54 @@ class TestCPSharedPagedAllocator(CustomTestCase):
)
self.assertEqual(allocator.available_size(), page_size * 16)
def test_shared_allocator_compute_owner_alloc_supports_multi_request_flattened_order(
self,
):
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_batch_in_seq_page_compute_owners,
)
page_size = 64
cp_size = 4
owners = build_batch_in_seq_page_compute_owners(
extend_lens=[page_size * 3, page_size * 5],
extend_prefix_lens=[0, page_size * 2],
page_size=page_size,
cp_size=cp_size,
)
allocator = CPSharedPagedTokenToKVPoolAllocator(
logical_size=page_size * 32,
physical_size=page_size * 8,
page_size=page_size,
dtype=torch.bfloat16,
device="cpu",
kvcache=None,
need_sort=False,
cp_size=cp_size,
cp_rank=0,
)
locs = allocator.alloc_extend_compute_owner(
prefix_lens=torch.tensor([0, page_size * 2], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0, page_size * 2], dtype=torch.int64),
seq_lens=torch.tensor(
[page_size * 3, page_size * 7], dtype=torch.int64
),
seq_lens_cpu=torch.tensor(
[page_size * 3, page_size * 7], dtype=torch.int64
),
last_loc=torch.tensor([-1, page_size * 2 - 1], dtype=torch.int64),
extend_num_tokens=page_size * 8,
page_compute_owners=owners,
)
self.assertIsNotNone(locs)
req0_pages = locs[: page_size * 3].view(-1, page_size)[:, 0] // page_size
req1_pages = locs[page_size * 3 :].view(-1, page_size)[:, 0] // page_size
logical_pages = torch.cat([req0_pages, req1_pages])
self.assertEqual(((logical_pages - 1) % cp_size).tolist(), owners)
def test_compute_owner_alloc_does_not_use_torch_isin_for_page_removal(self):
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
@@ -555,6 +654,135 @@ class TestCPSharedPagedAllocator(CustomTestCase):
self.assertEqual(allocator.calls, 1)
self.assertEqual(out.numel(), page_size * 8)
def test_compute_owner_alloc_uses_batch_owner_plan_for_multi_request(self):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
page_size = 64
class FakeAllocator:
def __init__(self):
self.page_size = page_size
self.cp_size = 4
self.calls = []
def available_size(self):
return page_size * 100
def alloc_extend_compute_owner(
self,
_prefix_lens,
_prefix_lens_cpu,
_seq_lens,
_seq_lens_cpu,
_last_loc,
extend_num_tokens,
page_compute_owners,
):
self.calls.append(list(page_compute_owners))
return torch.arange(extend_num_tokens, dtype=torch.int64)
def alloc_extend(self, *_args, **_kwargs):
raise AssertionError("legacy allocation should not be used for bs>1")
class FakeTreeCache:
def __init__(self, allocator):
self.token_to_kv_pool_allocator = allocator
def is_chunk_cache(self):
return False
def evict(self, _params):
raise AssertionError("tree eviction should not be used on first success")
allocator = FakeAllocator()
server_args = SimpleNamespace(
enable_nsa_prefill_cp_shared_kv=True,
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
with patch.object(common, "get_global_server_args", return_value=server_args):
out = common.alloc_paged_token_slots_extend(
tree_cache=FakeTreeCache(allocator),
prefix_lens=torch.tensor([0, page_size * 2], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0, page_size * 2], dtype=torch.int64),
seq_lens=torch.tensor(
[page_size * 3, page_size * 7], dtype=torch.int64
),
seq_lens_cpu=torch.tensor(
[page_size * 3, page_size * 7], dtype=torch.int64
),
last_loc=torch.tensor([-1, page_size * 2 - 1], dtype=torch.int64),
extend_num_tokens=page_size * 8,
)
self.assertEqual(allocator.calls, [[0, 1, 2, 0, 1, 2, 3, 3]])
self.assertEqual(out.numel(), page_size * 8)
def test_compute_owner_alloc_fail_fast_for_multi_request_unsupported_owner_plan(
self,
):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
page_size = 64
class FakeAllocator:
def __init__(self):
self.page_size = page_size
self.cp_size = 4
self.owner_calls = 0
self.legacy_calls = 0
def available_size(self):
return page_size * 100
def alloc_extend_compute_owner(self, *_args, **_kwargs):
self.owner_calls += 1
raise AssertionError("owner allocation should not run without plan")
def alloc_extend(self, *_args, **_kwargs):
self.legacy_calls += 1
raise AssertionError("legacy allocation should not be used for bs>1")
class FakeTreeCache:
def __init__(self, allocator):
self.token_to_kv_pool_allocator = allocator
def is_chunk_cache(self):
return False
def evict(self, _params):
raise AssertionError("tree eviction should not be used")
allocator = FakeAllocator()
server_args = SimpleNamespace(
enable_nsa_prefill_cp_shared_kv=True,
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
with patch.object(common, "get_global_server_args", return_value=server_args):
with self.assertRaises(RuntimeError) as cm:
common.alloc_paged_token_slots_extend(
tree_cache=FakeTreeCache(allocator),
prefix_lens=torch.tensor([0, 1], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0, 1], dtype=torch.int64),
seq_lens=torch.tensor([page_size, page_size + 1], dtype=torch.int64),
seq_lens_cpu=torch.tensor(
[page_size, page_size + 1], dtype=torch.int64
),
last_loc=torch.tensor([-1, 0], dtype=torch.int64),
extend_num_tokens=page_size * 2,
)
self.assertIn("prefix_not_page_aligned", str(cm.exception))
self.assertEqual(allocator.owner_calls, 0)
self.assertEqual(allocator.legacy_calls, 0)
def test_compute_owner_alloc_skips_aggregate_evict_before_owner_attempt(self):
from types import SimpleNamespace
@@ -880,6 +1108,103 @@ class TestCPSharedPagedAllocator(CustomTestCase):
[params.owner_lane_deficits for params in tree_cache.evict_params],
)
def test_compute_owner_capacity_wait_reports_owner_lane_deficits_for_multi_request(
self,
):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
page_size = 64
class FakeAllocator:
def __init__(self):
self.page_size = page_size
self.cp_size = 4
self.owner_calls = []
def available_size(self):
return self.page_size * 4
def allocator_state_str(self):
return "allocator_state_for_test"
def alloc_extend_compute_owner(
self,
_prefix_lens,
_prefix_lens_cpu,
_seq_lens,
_seq_lens_cpu,
_last_loc,
_extend_num_tokens,
page_compute_owners,
):
self.owner_calls.append(list(page_compute_owners))
return None
def compute_owner_lane_stats(self, _page_compute_owners):
return [2, 2, 2, 2], [2, 2, 2, 0], [0, 0, 0, 2]
def alloc_extend(self, *_args, **_kwargs):
raise AssertionError("legacy allocation should not be used for bs>1")
class FakeTreeCache:
def __init__(self):
self.token_to_kv_pool_allocator = FakeAllocator()
self.evict_params = []
def is_chunk_cache(self):
return False
def evictable_size(self):
return page_size * 8
def evict(self, params):
self.evict_params.append(params)
return EvictResult(num_tokens_evicted=0)
def pretty_print(self):
raise AssertionError("recoverable capacity wait should not dump tree")
server_args = SimpleNamespace(
enable_nsa_prefill_cp_shared_kv=True,
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
tree_cache = FakeTreeCache()
with patch.object(common, "get_global_server_args", return_value=server_args):
with self.assertRaises(common.KVCapacityWaitError) as cm:
common.alloc_paged_token_slots_extend(
tree_cache=tree_cache,
prefix_lens=torch.tensor([0, page_size * 2], dtype=torch.int64),
prefix_lens_cpu=torch.tensor(
[0, page_size * 2], dtype=torch.int64
),
seq_lens=torch.tensor(
[page_size * 3, page_size * 7], dtype=torch.int64
),
seq_lens_cpu=torch.tensor(
[page_size * 3, page_size * 7], dtype=torch.int64
),
last_loc=torch.tensor([-1, page_size * 2 - 1], dtype=torch.int64),
extend_num_tokens=page_size * 8,
)
err = cm.exception
self.assertEqual(
tree_cache.token_to_kv_pool_allocator.owner_calls,
[[0, 1, 2, 0, 1, 2, 3, 3], [0, 1, 2, 0, 1, 2, 3, 3]],
)
self.assertEqual(err.required_by_owner, [2, 2, 2, 2])
self.assertEqual(err.available_by_owner, [2, 2, 2, 0])
self.assertEqual(err.deficit_by_owner, [0, 0, 0, 2])
self.assertIn(
[0, 0, 0, 2],
[params.owner_lane_deficits for params in tree_cache.evict_params],
)
def test_alloc_for_extend_releases_req_slots_on_recoverable_capacity_wait(self):
from types import SimpleNamespace