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:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user