Enable CP shared-KV compute padding without inflating cache state

Tiny extend requests can leave most CP lanes without query work, which has been tied to hangs and accept-length regressions. This change introduces a dual valid/compute metadata contract: forward paths may materialize compute-padded rows, while cache, current reuse, direct write, HiCache backup, and load remain valid/page based.

The implementation keeps radix/HiCache/device allocation on real page extents, filters dummy compute rows before MLA/index cache writes and current reuse, makes top-k/index consume compute rows while compacting valid rows, and opens tiny CP shared-KV in-seq split through compute padding. The accompanying plan document records the contract and P1-P7 evidence.

Constraint: CP shared KV and HiCache must stay page-granular; dummy compute rows must not allocate, write, backup, or load KV cache.

Constraint: Avoid silent fallback and avoid adding collectives on hot paths.

Rejected: Pad cache allocations to cp_size pages | would waste KV capacity and pollute radix/HiCache state.

Rejected: Keep tiny suffixes out of CP split | preserves the zero-lane behavior that compute padding is meant to remove.

Confidence: medium

Scope-risk: broad

Directive: Do not route compute-padded dummy rows into out_cache_loc, current reuse, HiCache reservation, or backup descriptors; keep valid/cache metadata explicit.

Tested: Remote g0034 container targeted P7 tests: 3 passed, 3 warnings.

Tested: Remote g0034 container full unit slice: PYTHONPATH=python python -m pytest -q test/registered/unit/layers/test_nsa_cp_utils.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py test/registered/unit/mem_cache/test_cp_shared_kv_layout.py => 214 passed, 5 warnings, 2 subtests passed.

Tested: Local py_compile for touched P7 test file.

Not-tested: Latest CUDA/ETE traffic validation for dummy top-k rows, accept len, output len, and detokenizer hang behavior.
This commit is contained in:
laoyao0822
2026-06-04 01:34:35 +08:00
parent b3913046b6
commit 3e3f1b776b
7 changed files with 2681 additions and 124 deletions
@@ -297,6 +297,173 @@ class TestCPSharedPagedAllocator(CustomTestCase):
self.assertEqual(owners, [0, 1])
def test_alloc_extend_compute_owner_uses_valid_pages_not_compute_padding_pages(
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 = 8
self.owner_calls = []
self.extend_num_tokens = []
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.extend_num_tokens.append(int(extend_num_tokens))
self.owner_calls.append(list(page_compute_owners))
return torch.arange(
1024, 1024 + int(extend_num_tokens), dtype=torch.int64
)
def alloc_extend(self, *_args, **_kwargs):
raise AssertionError("legacy allocation should not be used")
class FakeTreeCache:
def __init__(self):
self.token_to_kv_pool_allocator = FakeAllocator()
def is_chunk_cache(self):
return False
def evict(self, *_args, **_kwargs):
raise AssertionError("eviction should not be needed")
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):
out_cache_loc = common.alloc_paged_token_slots_extend(
tree_cache=tree_cache,
prefix_lens=torch.tensor([0], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
seq_lens=torch.tensor([65], dtype=torch.int64),
seq_lens_cpu=torch.tensor([65], dtype=torch.int64),
last_loc=torch.tensor([-1], dtype=torch.int64),
extend_num_tokens=65,
)
self.assertEqual(out_cache_loc.numel(), 65)
self.assertEqual(tree_cache.token_to_kv_pool_allocator.extend_num_tokens, [65])
self.assertEqual(tree_cache.token_to_kv_pool_allocator.owner_calls, [[0, 1]])
def test_cp_hicache_write_reservation_uses_page_tail_not_compute_padding_extent(
self,
):
from sglang.srt.managers.cache_controller import HiCacheController
page_size = 64
class FakeHostPool:
def __init__(self):
self.alloc_sizes = []
def alloc_contiguous_preferred(self, need_size):
self.alloc_sizes.append(int(need_size))
return torch.arange(1024, 1024 + int(need_size), dtype=torch.int64)
def alloc(self, need_size):
return self.alloc_contiguous_preferred(need_size)
def free(self, _indices):
raise AssertionError("reservation should not roll back")
controller = HiCacheController.__new__(HiCacheController)
controller.page_size = page_size
controller.cp_shared_kv_layout = CpSharedKVLayout(
page_size=page_size, cp_size=8, cp_rank=1
)
controller.mem_pool_host = FakeHostPool()
controller.draft_mem_pool_host = None
controller.draft_mem_pool_device = None
reservation = controller.reserve_write_cp(
torch.arange(page_size, page_size + 65, dtype=torch.int64),
node_id=123,
)
self.assertEqual(controller.mem_pool_host.alloc_sizes, [page_size])
self.assertEqual(reservation.metadata.logical_len, 65)
self.assertEqual(reservation.metadata.padded_len, page_size * 2)
self.assertEqual(reservation.metadata.page_owners.tolist(), [0, 1])
self.assertEqual(reservation.host_indices.numel(), page_size)
self.assertEqual(reservation.physical_device_indices.numel(), page_size)
self.assertEqual(
reservation.metadata.owned_positions.tolist(),
list(range(64, 128)),
)
def test_cp_hicache_load_returns_valid_visible_len_while_loading_owned_page_tail(
self,
):
from types import SimpleNamespace
from sglang.srt.managers.cache_controller import HiCacheController
from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata
page_size = 64
class FakeDeviceAllocator:
def __init__(self):
self.owner_calls = []
self.freed = []
def alloc_pages_with_owners(self, page_owners):
self.owner_calls.append(list(page_owners))
return torch.arange(page_size, page_size * 3, dtype=torch.int64)
def free(self, indices):
self.freed.append(indices.clone())
controller = HiCacheController.__new__(HiCacheController)
controller.page_size = page_size
controller.cp_shared_kv_layout = CpSharedKVLayout(
page_size=page_size, cp_size=8, cp_rank=1
)
controller.mem_pool_device_allocator = FakeDeviceAllocator()
controller.load_queue = []
controller.draft_load_queue = []
controller.draft_mem_pool_host = None
controller.draft_mem_pool_device = None
metadata = CpHiCacheNodeMetadata(
logical_len=65,
padded_len=page_size * 2,
owned_positions=torch.arange(page_size, page_size * 2, dtype=torch.int64),
host_indices=torch.arange(1024, 1024 + page_size, dtype=torch.int64),
page_owners=torch.tensor([0, 1], dtype=torch.int8),
page_size=page_size,
)
node = SimpleNamespace(cp_hicache=metadata, host_len=65, id=321)
visible_device_indices = controller.load_cp([node], node_id=321)
self.assertEqual(controller.mem_pool_device_allocator.owner_calls, [[0, 1]])
self.assertEqual(controller.mem_pool_device_allocator.freed, [])
self.assertEqual(visible_device_indices.numel(), 65)
self.assertEqual(visible_device_indices.tolist(), list(range(64, 129)))
self.assertEqual(len(controller.load_queue), 1)
load_op = controller.load_queue[0]
self.assertEqual(load_op.host_indices.tolist(), list(range(1024, 1024 + 64)))
self.assertEqual(load_op.device_indices.tolist(), list(range(64, 128)))
def test_compute_owner_page_assignment_allows_radix_hit_suffix_with_one_page_per_rank(
self,
):