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