Merge CP shared-KV batch owner allocation

PR #10 provides the W2 owner-lane allocation side needed by the current bs>1 CP shared-KV work: supported multi-request extends now derive page owners per request, preserve flattened request order, and use owner-lane allocation instead of legacy page allocation.\n\nThe merge applies cleanly on top of the W3/W4 target current-reuse work because the touched runtime surface is allocator/mem_cache only. Focused remote tests cover the new owner-lane planner/allocation behavior and the existing W4 target reuse regressions.\n\nConstraint: CP shared-KV page ownership is request-relative and page-granular.\nConstraint: Supported bs>1 CP shared-KV allocation must not silently fall back to legacy allocation.\nRejected: Treat multiple requests as one concatenated extend | would assign owners relative to the wrong request boundary.\nRejected: Keep multi_batch fallback | breaks downstream direct-write owner assumptions.\nConfidence: medium\nScope-risk: moderate\nDirective: Preserve request-order flattening when changing owner allocation; do not reorder by lane for convenience.\nTested: Local py_compile for allocator.py, common.py, cp_shared_kv_compute_owner.py, and test_cp_shared_kv_layout.py.\nTested: Local git diff --check --cached.\nTested: Remote g0034 py_compile for touched mem_cache files and test_cp_shared_kv_layout.py.\nTested: Remote g0034 pytest test_cp_shared_kv_layout.py => 34 passed, 3 warnings.\nTested: Remote g0034 pytest test_nsa_cp_utils.py test_cp_shared_kv_runtime.py => 157 passed, 5 warnings, 2 subtests passed.\nNot-tested: Full ETE/perf run; production HiCache load/back up with bs>1 allocation under sustained traffic.
This commit is contained in:
laoyao0822
2026-06-03 07:22:04 +08:00
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