From 1ebde44e5937851464bb5c6ac3a9c458d973a63a Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Wed, 10 Jun 2026 05:37:06 +0800 Subject: [PATCH] Compact skipped NSA index-cache state safely Index skip reduces the number of target layers that own NSA index state, but PD transfer and HiCache still assumed dense full-layer state buffers. This change carries explicit state layer IDs through prefill/decode registration, compacts device and host index buffers to active layers, and maps logical layer IDs to compact slots on transfer paths. The PD side fails fast when prefill/decode disagree on NSA state layer identity instead of silently truncating or copying mismatched buffers. Host direct tests now use the same CPU-index descriptor contract required by the TAI cudaMemcpyBatchAsync path, and host registered memory is unregistered on tensor finalization to avoid stale cudaHostRegister state across CUDA tests. Constraint: CP shared-KV with index_topk skip must keep target/draft state identity explicit before compacting buffers Constraint: Direct HiCache TAI transfer rejects CUDA indices to avoid hidden D2H copies on the control path Rejected: Keep full-layer L1/L2 index buffers | wastes the memory/bandwidth that index skip is meant to save Rejected: Infer state buffer order by count only | can silently corrupt cache when active layer sets differ Confidence: high Scope-risk: moderate Directive: Do not compact or reorder NSA state buffers without carrying logical layer IDs through PD registration and validating both sides Tested: Remote container py_compile for touched runtime files Tested: Remote container pytest: test_nsa_pool_host_unit.py, test_model_runner_kv_cache_mixin.py, test_cp_shared_kv_transfer_mapping.py, test_pd_state_layer_ids.py, test_cp_per_layer_transfer.py, test_cp_shared_kv_runtime.py -> 200 passed, 2 subtests passed Not-tested: Full ETE GSM8K/replay after compacted P3-P6 changes Co-authored-by: OmX --- ..._cp_index_skip_cache_communication_plan.md | 35 ++++ python/sglang/srt/disaggregation/base/conn.py | 1 + python/sglang/srt/disaggregation/decode.py | 9 + .../srt/disaggregation/mooncake/conn.py | 35 ++++ python/sglang/srt/disaggregation/nixl/conn.py | 33 ++++ python/sglang/srt/disaggregation/prefill.py | 9 + python/sglang/srt/disaggregation/utils.py | 2 + python/sglang/srt/mem_cache/memory_pool.py | 20 +- .../sglang/srt/mem_cache/memory_pool_host.py | 79 +++++++- .../model_runner_kv_cache_mixin.py | 13 +- .../test_cp_shared_kv_transfer_mapping.py | 22 +++ .../disaggregation/test_pd_state_layer_ids.py | 182 ++++++++++++++++++ .../unit/mem_cache/test_nsa_pool_host_unit.py | 146 +++++++++++++- .../test_model_runner_kv_cache_mixin.py | 41 ++++ 14 files changed, 611 insertions(+), 16 deletions(-) create mode 100644 test/registered/unit/disaggregation/test_pd_state_layer_ids.py create mode 100644 test/registered/unit/model_executor/test_model_runner_kv_cache_mixin.py diff --git a/docs/advanced_features/nsa_prefill_cp_index_skip_cache_communication_plan.md b/docs/advanced_features/nsa_prefill_cp_index_skip_cache_communication_plan.md index 23b840976..48ff09622 100644 --- a/docs/advanced_features/nsa_prefill_cp_index_skip_cache_communication_plan.md +++ b/docs/advanced_features/nsa_prefill_cp_index_skip_cache_communication_plan.md @@ -1106,3 +1106,38 @@ Completion criteria: - Replay throughput does not regress relative to current bs>1 baseline. - Logs show active index state buffers are reduced when index_topk_freq > 1. ``` + +## 14. Implementation status ledger + +- P0-P2: implemented in `d21952b90 Reduce inactive NSA index-cache transfer safely`. + - Central helper now owns target/draft skip formula. + - `NSATokenToKVPool` exposes active-layer metadata without allocation compaction. + - HiCache indexer load/backup skips inactive target index layers while MLA KV still transfers every layer. +- P3: implemented locally after this plan revision. + - `KVArgs.state_layer_ids` is populated for NSA state buffers. + - CP draft NSA state is appended under negative layer IDs (`-1`, `-2`, ...). + - Mooncake and NIXL registration roundtrip optional signed `state_layer_ids`. + - Mooncake/NIXL layer-aware state transfer fails fast on layer-id mismatch; Mooncake also fails fast on layer-aware state-buffer count mismatch instead of truncating. + - Verification: remote container `test_cp_shared_kv_transfer_mapping.py`, `test_pd_state_layer_ids.py`, and `test_cp_per_layer_transfer.py` passed (`39 passed`). +- P4: implemented locally after this plan revision. + - `NSATokenToKVPool.get_state_buf_infos()` returns only active index-layer slots. + - `NSATokenToKVPool.get_state_layer_ids()` returns active logical layer IDs. + - Draft/EAGLE remains appended under P3 negative layer IDs. + - Verification: remote container `TestNSAIndexerPageIndices`, `test_cp_shared_kv_transfer_mapping.py`, `test_pd_state_layer_ids.py`, and `test_cp_per_layer_transfer.py` passed (`52 passed`). +- P5: implemented locally after this plan revision. + - Device/L1 NSA index cache allocation uses `len(index_active_layer_ids)` when `compact_index_layers=True`. + - Model runner enables compact device index allocation for NSA pools. + - Capacity estimation counts active index layers instead of all local layers. + - Verification: remote container `TestNSAIndexerPageIndices`, `test_model_runner_kv_cache_mixin.py`, and `test_cp_shared_kv_runtime.py` passed (`141 passed`, `2 subtests passed`). +- P6: implemented locally after this plan revision. + - `NSATokenToKVPoolHost` mirrors the device pool's active index-layer IDs and allocates host/L2 index buffers with active-layer depth only. + - Host index transfer maps logical layer IDs to compact host slots for `layer_first`, `page_first_direct`, and `layer_page_first`. + - Test-only/legacy host stubs without compact metadata keep dense-slot behavior; real compact host pools still fail fast if an inactive logical index layer is requested. + - Direct HiCache tests now follow the production CPU-index contract before calling TAI `cudaMemcpyBatchAsync` paths. + - Host-register allocation now checks `cudaHostRegister` errors and unregisters host tensors on object finalization; this prevents repeated CUDA unit tests from poisoning later CUDA ops with stale `cudaErrorHostMemoryAlreadyRegistered`. + - Verification: remote container P3-P6 suite passed: + `test_nsa_pool_host_unit.py`, `test_model_runner_kv_cache_mixin.py`, + `test_cp_shared_kv_transfer_mapping.py`, `test_pd_state_layer_ids.py`, + `test_cp_per_layer_transfer.py`, and `test_cp_shared_kv_runtime.py` + (`200 passed`, `2 subtests passed`). +- P7: not rerun after P3. Needs GSM8K/replay after P4+ compact changes. diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index c87b7ce41..8c08a3ef1 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -31,6 +31,7 @@ class KVArgs: state_data_ptrs: List[int] state_data_lens: List[int] state_item_lens: List[int] + state_layer_ids: List[int] state_type: str # "none", "mamba", "swa" # for mamba state different tp slice transfer state_dim_per_tensor: List[int] # dimension to slice for each state tensor diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index a426b93bd..07b2f4e0b 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -136,6 +136,13 @@ def _state_buf_infos(pool): return state_type, state_data_ptrs, state_data_lens, state_item_lens +def _state_layer_ids(pool): + get_state_layer_ids = getattr(pool, "get_state_layer_ids", None) + if get_state_layer_ids is None: + return [] + return list(get_state_layer_ids()) + + def _kv_locs_to_page_indices_cpu( kv_locs: torch.Tensor, page_size: int, @@ -429,6 +436,7 @@ class DecodePreallocQueue: kv_args.state_data_ptrs = state_data_ptrs kv_args.state_data_lens = state_data_lens kv_args.state_item_lens = state_item_lens + kv_args.state_layer_ids = _state_layer_ids(self.token_to_kv_pool) if isinstance(self.token_to_kv_pool, SWAKVPool): kv_args.state_type = "swa" @@ -447,6 +455,7 @@ class DecodePreallocQueue: kv_args.state_data_ptrs = [] kv_args.state_data_lens = [] kv_args.state_item_lens = [] + kv_args.state_layer_ids = [] kv_args.state_type = "none" draft_state_type = "none" diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index bc379c6ea..9bfda8089 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -175,6 +175,7 @@ class KVArgsRegisterInfo: # for mamba state different tp slice transfer dst_state_item_lens: list[int] dst_state_dim_per_tensor: list[int] + dst_state_layer_ids: list[int] @classmethod def from_zmq(cls, msg: List[bytes]): @@ -199,6 +200,11 @@ class KVArgsRegisterInfo: if len(msg) > 11 and len(msg[11]) > 0 else [] ), + dst_state_layer_ids=( + list(struct.unpack(f"{len(msg[12])//4}i", msg[12])) + if len(msg) > 12 and len(msg[12]) > 0 + else [] + ), ) @@ -957,6 +963,22 @@ class MooncakeKVManager(CommonKVManager): raise RuntimeError( f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." ) + src_state_layer_ids = list( + getattr(self.kv_args, "state_layer_ids", []) or [] + ) + dst_state_layer_ids = ( + list(getattr(target_rank_registration_info, "dst_state_layer_ids", []) or []) + if target_rank_registration_info is not None + else [] + ) + if src_state_layer_ids or dst_state_layer_ids: + if src_state_layer_ids != dst_state_layer_ids: + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][state_layer_ids] " + f"prefill={src_state_layer_ids} decode={dst_state_layer_ids} " + f"state_type={state_type} room={req.room} " + f"session={req.mooncake_session_id}" + ) effective_dst_state_indices = ( np.asarray(dst_state_indices, dtype=np.int32) if dst_state_indices is not None @@ -984,6 +1006,14 @@ class MooncakeKVManager(CommonKVManager): dst_state_ptrs = dst_state_data_ptrs state_item_lens = self.kv_args.state_item_lens if len(src_state_data_ptrs) != len(dst_state_ptrs): + if src_state_layer_ids or dst_state_layer_ids: + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][state_buffer_count] " + f"src={len(src_state_data_ptrs)} dst={len(dst_state_ptrs)} " + f"src_layers={src_state_layer_ids} " + f"dst_layers={dst_state_layer_ids} state_type={state_type} " + f"room={req.room} session={req.mooncake_session_id}" + ) transfer_buf_count = min(len(src_state_data_ptrs), len(dst_state_ptrs)) logger.warning( "State buffer count mismatch during PD transfer: src=%s dst=%s " @@ -1839,6 +1869,10 @@ class MooncakeKVReceiver(CommonKVReceiver): packed_state_dim_per_tensor = b"".join( struct.pack("I", dim) for dim in state_dim_per_tensor ) + packed_state_layer_ids = b"".join( + struct.pack("i", int(layer_id)) + for layer_id in getattr(self.kv_mgr.kv_args, "state_layer_ids", []) + ) # Note(shangming): No need to add pp rank here since decode pp size should be equal to prefill pp size or 1 tp_rank = self.kv_mgr.kv_args.engine_rank kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0] @@ -1880,6 +1914,7 @@ class MooncakeKVReceiver(CommonKVReceiver): dst_kv_item_len, packed_state_item_lens, packed_state_dim_per_tensor, + packed_state_layer_ids, ] ) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 764fd9e42..ec0fba292 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -84,6 +84,7 @@ class KVArgsRegisterInfo: decode_tp_size: int decode_tp_rank: int dst_kv_item_len: int + dst_state_layer_ids: list[int] @classmethod def from_zmq(cls, msg: List[bytes]): @@ -106,6 +107,11 @@ class KVArgsRegisterInfo: decode_tp_size=int(msg[9].decode("ascii")), decode_tp_rank=int(msg[10].decode("ascii")), dst_kv_item_len=int(msg[11].decode("ascii")), + dst_state_layer_ids=( + list(struct.unpack(f"{len(msg[12]) // 4}i", msg[12])) + if len(msg) > 12 and msg[12] != b"" + else [] + ), ) @@ -680,6 +686,7 @@ class NixlKVManager(CommonKVManager): dst_gpu_id: int, notif: str, decode_tp_size: int, + dst_state_layer_ids: Optional[List[int]] = None, ): """Send state or extra pool data with type-specific handling.""" state_type = getattr(self.kv_args, "state_type", "none") @@ -702,6 +709,26 @@ class NixlKVManager(CommonKVManager): raise RuntimeError( f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." ) + src_state_layer_ids = list( + getattr(self.kv_args, "state_layer_ids", []) or [] + ) + dst_state_layer_ids = list(dst_state_layer_ids or []) + if src_state_layer_ids or dst_state_layer_ids: + if src_state_layer_ids != dst_state_layer_ids: + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][state_layer_ids] " + f"prefill={src_state_layer_ids} decode={dst_state_layer_ids} " + f"state_type={state_type} peer={peer_name} notif={notif}" + ) + if len(self.kv_args.state_data_ptrs) != len(dst_state_data_ptrs): + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][state_buffer_count] " + f"src={len(self.kv_args.state_data_ptrs)} " + f"dst={len(dst_state_data_ptrs)} " + f"src_layers={src_state_layer_ids} " + f"dst_layers={dst_state_layer_ids} " + f"state_type={state_type} peer={peer_name} notif={notif}" + ) if len(prefill_state_indices) != len(dst_state_indices): raise RuntimeError( f"State index length mismatch: prefill={len(prefill_state_indices)}, " @@ -791,6 +818,7 @@ class NixlKVManager(CommonKVManager): dst_info.gpu_id, f"{req.room}_state_{self.kv_args.pp_rank}", decode_tp_size, + dst_info.dst_state_layer_ids, ) if state_xfer_handle is not None: handles.append(state_xfer_handle) @@ -1069,6 +1097,10 @@ class NixlKVReceiver(CommonKVReceiver): packed_state_data_ptrs = b"".join( struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs ) + packed_state_layer_ids = b"".join( + struct.pack("i", int(layer_id)) + for layer_id in getattr(self.kv_mgr.kv_args, "state_layer_ids", []) + ) with lock: sock.send_multipart( @@ -1086,6 +1118,7 @@ class NixlKVReceiver(CommonKVReceiver): str(self.kv_mgr.attn_tp_size).encode("ascii"), str(self.kv_mgr.kv_args.engine_rank).encode("ascii"), str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii"), + packed_state_layer_ids, ] ) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 15881b4f0..9298130f8 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -191,6 +191,13 @@ def _state_buf_infos(pool): return state_type, state_data_ptrs, state_data_lens, state_item_lens +def _state_layer_ids(pool): + get_state_layer_ids = getattr(pool, "get_state_layer_ids", None) + if get_state_layer_ids is None: + return [] + return list(get_state_layer_ids()) + + def _kv_locs_to_page_indices_cpu( kv_locs: torch.Tensor, page_size: int, @@ -376,6 +383,7 @@ class PrefillBootstrapQueue: kv_args.state_data_ptrs = state_data_ptrs kv_args.state_data_lens = state_data_lens kv_args.state_item_lens = state_item_lens + kv_args.state_layer_ids = _state_layer_ids(self.token_to_kv_pool) if isinstance(self.token_to_kv_pool, SWAKVPool): kv_args.state_type = "swa" @@ -394,6 +402,7 @@ class PrefillBootstrapQueue: kv_args.state_data_ptrs = [] kv_args.state_data_lens = [] kv_args.state_item_lens = [] + kv_args.state_layer_ids = [] kv_args.state_type = "none" draft_state_type = "none" diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 49a524609..e20f83b4a 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -259,6 +259,8 @@ def append_cp_draft_state_buffers( kv_args.state_data_ptrs += draft_state_data_ptrs kv_args.state_data_lens += draft_state_data_lens kv_args.state_item_lens += draft_state_item_lens + if hasattr(kv_args, "state_layer_ids"): + kv_args.state_layer_ids += [-(i + 1) for i in range(draft_state_bufs)] kv_args.draft_state_buffer_count = draft_state_bufs return True diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 2c05f9acc..41f5cc36f 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -1913,6 +1913,11 @@ class NSATokenToKVPool(MLATokenToKVPool): if self.custom_mem_pool else nullcontext() ): + index_buffer_layer_num = ( + len(self.index_active_layer_ids) + if self.index_compact_layers + else layer_num + ) self.index_k_with_scale_buffer = [ torch.zeros( # Layout: @@ -1931,7 +1936,7 @@ class NSATokenToKVPool(MLATokenToKVPool): dtype=self.index_k_with_scale_buffer_dtype, device=device, ) - for _ in range(layer_num) + for _ in range(index_buffer_layer_num) ] self._finalize_allocation_log(size) @@ -2054,17 +2059,24 @@ class NSATokenToKVPool(MLATokenToKVPool): ) def get_state_buf_infos(self): + slots = [ + self.get_index_layer_slot(layer_id) + for layer_id in self.index_active_layer_ids + ] data_ptrs = [ - self.index_k_with_scale_buffer[i].data_ptr() for i in range(self.layer_num) + self.index_k_with_scale_buffer[slot].data_ptr() for slot in slots ] data_lens = [ - self.index_k_with_scale_buffer[i].nbytes for i in range(self.layer_num) + self.index_k_with_scale_buffer[slot].nbytes for slot in slots ] item_lens = [ - self.index_k_with_scale_buffer[i][0].nbytes for i in range(self.layer_num) + self.index_k_with_scale_buffer[slot][0].nbytes for slot in slots ] return data_ptrs, data_lens, item_lens + def get_state_layer_ids(self): + return list(self.index_active_layer_ids) + def get_kv_size_bytes(self): kv_size_bytes = super().get_kv_size_bytes() for index_k_cache in self.index_k_with_scale_buffer: diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 4032d1b32..3957c05c4 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -3,6 +3,7 @@ import bisect import heapq import logging import threading +import weakref from collections import defaultdict from functools import lru_cache, wraps from typing import Optional @@ -196,12 +197,38 @@ def alloc_with_host_register( """ buffer = allocator.allocate(dims, dtype=dtype, device=device) if pin_memory: - torch.cuda.cudart().cudaHostRegister( - buffer.data_ptr(), buffer.numel() * buffer.element_size(), 0 + ptr = buffer.data_ptr() + _check_torch_cudart( + torch.cuda.cudart().cudaHostRegister( + ptr, buffer.numel() * buffer.element_size(), 0 + ), + "cudaHostRegister", ) + weakref.finalize(buffer, _cuda_host_unregister, ptr) return buffer +def _check_torch_cudart(err, op_name: str) -> None: + if int(err) != 0: + try: + err_str = torch.cuda.cudart().cudaGetErrorString(err) + except Exception: + err_str = repr(err) + raise RuntimeError(f"{op_name} failed: {err_str}") + + +def _cuda_host_unregister(ptr: int) -> None: + try: + _check_torch_cudart( + torch.cuda.cudart().cudaHostUnregister(ptr), "cudaHostUnregister" + ) + except Exception: + # Best-effort cleanup for host tensors. The owning process may already + # be tearing down CUDA state at exit; production host pools are + # long-lived, while unit tests rely on timely unregister when pools die. + pass + + def alloc_with_pin_memory( dims, dtype: torch.dtype, @@ -1834,6 +1861,18 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): self.index_head_dim = device_pool.index_head_dim self.indexer_quant_block_size = device_pool.quant_block_size self.indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype + self.index_active_layer_ids = tuple( + int(layer_id) + for layer_id in getattr( + device_pool, + "index_active_layer_ids", + range(device_pool.start_layer, device_pool.start_layer + device_pool.layer_num), + ) + ) + self.index_active_layer_num = len(self.index_active_layer_ids) + self.index_logical_to_slot = { + layer_id: slot for slot, layer_id in enumerate(self.index_active_layer_ids) + } self.indexer_size_per_token = ( self.index_head_dim + self.index_head_dim // self.indexer_quant_block_size * 4 @@ -1853,7 +1892,9 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): self.indexer_page_stride_size = ( self.indexer_size_per_token * self.page_size * self.indexer_dtype.itemsize ) - self.indexer_layout_dim = self.indexer_page_stride_size * self.layer_num + self.indexer_layout_dim = ( + self.indexer_page_stride_size * self.index_active_layer_num + ) self.indexer_page_num = (self.size + self.page_size + 1) // self.page_size self._init_indexer_buffers() logger.info( @@ -1864,7 +1905,9 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): base = super().get_size_per_token() return ( base - + self.indexer_size_per_token * self.layer_num * self.indexer_dtype.itemsize + + self.indexer_size_per_token + * self.index_active_layer_num + * self.indexer_dtype.itemsize ) def _init_indexer_buffers(self): @@ -1883,10 +1926,11 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): pin_memory=self.pin_memory, allocator=self.allocator, ) - for _ in range(self.layer_num) + for _ in range(self.index_active_layer_num) ] self.index_k_data_refs = [ - self.index_k_with_scale_buffer[i] for i in range(self.layer_num) + self.index_k_with_scale_buffer[i] + for i in range(self.index_active_layer_num) ] self.index_k_data_ptrs = torch.tensor( [x.data_ptr() for x in self.index_k_data_refs], @@ -1897,7 +1941,7 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): self.index_k_with_scale_buffer = alloc_func( ( self.indexer_page_num, - self.layer_num, + self.index_active_layer_num, 1, self.indexer_page_stride_size, ), @@ -1909,7 +1953,7 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): elif self.layout == "layer_page_first": self.index_k_with_scale_buffer = alloc_func( ( - self.layer_num, + self.index_active_layer_num, self.indexer_page_num, 1, self.indexer_page_stride_size, @@ -1920,7 +1964,8 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): allocator=self.allocator, ) self.index_k_data_refs = [ - self.index_k_with_scale_buffer[i] for i in range(self.layer_num) + self.index_k_with_scale_buffer[i] + for i in range(self.index_active_layer_num) ] self.index_k_data_ptrs = torch.tensor( [x.data_ptr() for x in self.index_k_data_refs], @@ -1958,7 +2003,21 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost): return int(layer_id - getattr(device_pool, "start_layer", 0)) def _host_index_layer_slot(self, layer_id: int) -> int: - return int(layer_id - getattr(self, "start_layer", 0)) + mapping = getattr(self, "index_logical_to_slot", None) + if mapping is None: + # Unit-test stubs and legacy in-memory host pools created before the + # compact-index-layer metadata existed still use dense logical slots. + # Real NSATokenToKVPoolHost instances always install the mapping in + # __init__, so production compact paths retain the fail-fast below. + return int(layer_id - getattr(self, "start_layer", 0)) + try: + return int(mapping[int(layer_id)]) + except KeyError as exc: + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][host_index_cache_layer] " + f"inactive host index layer requested: layer_id={layer_id} " + f"active_layer_ids={list(self.index_active_layer_ids)}" + ) from exc def _active_index_layer_ids_for_transfer(self, device_pool): active_layer_ids = getattr(device_pool, "index_active_layer_ids", None) diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index ab70da97b..7f09ddff4 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -107,7 +107,17 @@ class ModelRunnerKVCacheMixin: element_size = torch._utils._element_size( NSATokenToKVPool.index_k_with_scale_buffer_dtype ) - cell_size += indexer_size_per_token * num_layers * element_size + index_layer_plan = build_nsa_index_layer_plan( + self.model_config.hf_config, + self.start_layer, + self.end_layer, + is_nextn=self.is_draft_worker, + ) + cell_size += ( + indexer_size_per_token + * len(index_layer_plan.active_layer_ids) + * element_size + ) else: if self.model_config.is_hybrid_swa: full_layers_num = len(self.model_config.full_attention_layer_ids) @@ -515,6 +525,7 @@ class ModelRunnerKVCacheMixin: end_layer=self.end_layer, index_head_dim=get_nsa_index_head_dim(self.model_config.hf_config), index_active_layer_ids=index_layer_plan.active_layer_ids, + compact_index_layers=True, ) if self.enable_hisparse: from sglang.srt.mem_cache.sparsity import parse_hisparse_config diff --git a/test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py b/test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py index db5a3961f..36db9dc76 100644 --- a/test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py +++ b/test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py @@ -134,6 +134,28 @@ class TestCPSharedKVTransferMapping(unittest.TestCase): self.assertEqual(kv_args.state_data_lens, [12, 23, 24]) self.assertEqual(kv_args.state_item_lens, [13, 25, 26]) + def test_append_cp_draft_state_buffers_adds_negative_state_layer_ids(self): + kv_args = SimpleNamespace( + state_type="nsa", + state_data_ptrs=[11], + state_data_lens=[12], + state_item_lens=[13], + state_layer_ids=[0], + ) + + appended = append_cp_draft_state_buffers( + kv_args, + draft_state_type="nsa", + draft_state_data_ptrs=[21, 22], + draft_state_data_lens=[23, 24], + draft_state_item_lens=[25, 26], + role="prefill", + cp_rank=2, + ) + + self.assertTrue(appended) + self.assertEqual(kv_args.state_layer_ids, [0, -1, -2]) + def test_append_cp_draft_state_buffers_rejects_state_mismatch_under_shared_kv(self): kv_args = SimpleNamespace( state_type="nsa", diff --git a/test/registered/unit/disaggregation/test_pd_state_layer_ids.py b/test/registered/unit/disaggregation/test_pd_state_layer_ids.py new file mode 100644 index 000000000..6c8d1581e --- /dev/null +++ b/test/registered/unit/disaggregation/test_pd_state_layer_ids.py @@ -0,0 +1,182 @@ +import concurrent.futures +import struct +from types import SimpleNamespace + +import numpy as np +import pytest + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + +mooncake_conn = pytest.importorskip("sglang.srt.disaggregation.mooncake.conn") +KVArgsRegisterInfo = mooncake_conn.KVArgsRegisterInfo +MooncakeKVManager = mooncake_conn.MooncakeKVManager + + +def _pack_q(values): + return b"".join(struct.pack("Q", value) for value in values) + + +def _pack_i(values): + return b"".join(struct.pack("i", value) for value in values) + + +def _pack_u(values): + return b"".join(struct.pack("I", value) for value in values) + + +def _register_msg(*, state_layer_ids): + return [ + b"room-a", + b"127.0.0.1", + b"1234", + b"session-a", + _pack_q([101, 102]), + _pack_q([201]), + _pack_q([301, 302, 303]), + b"0", + b"8", + b"64", + _pack_u([4096, 4096, 4096]), + _pack_u([128, 128, 128]), + _pack_i(state_layer_ids), + ] + + +def test_mooncake_register_info_roundtrips_state_layer_ids(): + info = KVArgsRegisterInfo.from_zmq( + _register_msg(state_layer_ids=[0, 4, 8, -1]) + ) + + assert info.dst_state_layer_ids == [0, 4, 8, -1] + + +def test_mooncake_state_layer_id_mismatch_fails_fast(): + manager = MooncakeKVManager.__new__(MooncakeKVManager) + manager.kv_args = SimpleNamespace( + state_type="nsa", + state_data_ptrs=[101, 102], + state_item_lens=[64, 64], + state_layer_ids=[0, 4], + draft_state_type="none", + draft_state_buffer_count=0, + ) + manager.attn_cp_rank = 3 + manager.attn_tp_size = 8 + manager.is_mla_backend = True + manager._send_kvcache_generic = lambda **_: 0 + req = SimpleNamespace( + room="room-a", + mooncake_session_id="session-a", + dst_state_indices=np.array([11, 12], dtype=np.int32), + ) + target_info = SimpleNamespace( + dst_state_layer_ids=[0, 5], + dst_state_data_ptrs=[201, 202], + dst_attn_tp_size=8, + dst_state_item_lens=[64, 64], + dst_state_dim_per_tensor=[], + dst_tp_rank=0, + ) + + with pytest.raises(RuntimeError, match=r"\[CP_SHARED_KV_FAIL_FAST\]\[state_layer_ids\].*prefill=\[0, 4\].*decode=\[0, 5\]"): + MooncakeKVManager.maybe_send_extra( + manager, + req, + prefill_state_indices=[1, 2], + dst_state_data_ptrs=[201, 202], + executor=concurrent.futures.ThreadPoolExecutor(max_workers=1), + target_rank_registration_info=target_info, + ) + + +def test_mooncake_layer_aware_state_buffer_count_mismatch_fails_fast(): + manager = MooncakeKVManager.__new__(MooncakeKVManager) + manager.kv_args = SimpleNamespace( + state_type="nsa", + state_data_ptrs=[101, 102, 103], + state_item_lens=[64, 64, 64], + state_layer_ids=[0, 4, 8], + draft_state_type="none", + draft_state_buffer_count=0, + ) + manager.attn_cp_rank = 3 + manager.attn_tp_size = 8 + manager.is_mla_backend = True + manager._send_kvcache_generic = lambda **_: 0 + req = SimpleNamespace( + room="room-b", + mooncake_session_id="session-b", + dst_state_indices=np.array([11, 12, 13], dtype=np.int32), + ) + target_info = SimpleNamespace( + dst_state_layer_ids=[0, 4, 8], + dst_state_data_ptrs=[201, 202], + dst_attn_tp_size=8, + dst_state_item_lens=[64, 64], + dst_state_dim_per_tensor=[], + dst_tp_rank=0, + ) + + with pytest.raises(RuntimeError, match=r"\[CP_SHARED_KV_FAIL_FAST\]\[state_buffer_count\].*src=3.*dst=2"): + MooncakeKVManager.maybe_send_extra( + manager, + req, + prefill_state_indices=[1, 2, 3], + dst_state_data_ptrs=[201, 202], + executor=concurrent.futures.ThreadPoolExecutor(max_workers=1), + target_rank_registration_info=target_info, + ) + + +def test_nixl_register_info_roundtrips_state_layer_ids(): + nixl_conn = pytest.importorskip("sglang.srt.disaggregation.nixl.conn") + + msg = [ + b"room-a", + b"127.0.0.1", + b"1234", + b"agent-a", + b"metadata", + _pack_q([101, 102]), + _pack_q([201]), + _pack_q([301, 302]), + b"0", + b"8", + b"3", + b"64", + _pack_i([0, 4, -1]), + ] + + info = nixl_conn.KVArgsRegisterInfo.from_zmq(msg) + + assert info.dst_state_layer_ids == [0, 4, -1] + + +def test_nixl_state_layer_id_mismatch_fails_fast(): + nixl_conn = pytest.importorskip("sglang.srt.disaggregation.nixl.conn") + manager = nixl_conn.NixlKVManager.__new__(nixl_conn.NixlKVManager) + manager.kv_args = SimpleNamespace( + state_type="nsa", + state_data_ptrs=[101, 102], + state_item_lens=[64, 64], + state_layer_ids=[0, 4], + ) + manager.attn_cp_rank = 3 + manager.attn_tp_size = 8 + manager.is_mla_backend = True + manager._send_kvcache_generic = lambda **_: 0 + + with pytest.raises(RuntimeError, match=r"\[CP_SHARED_KV_FAIL_FAST\]\[state_layer_ids\].*prefill=\[0, 4\].*decode=\[0, 5\]"): + nixl_conn.NixlKVManager.maybe_send_extra( + manager, + peer_name="decode-a", + prefill_state_indices=[1, 2], + dst_state_data_ptrs=[201, 202], + dst_state_indices=[11, 12], + dst_gpu_id=0, + notif="notif-a", + decode_tp_size=8, + dst_state_layer_ids=[0, 5], + ) diff --git a/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py b/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py index da2db2f8c..b074b70f2 100644 --- a/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py +++ b/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py @@ -147,6 +147,85 @@ class TestLayerPageFirstDirectHostLayout(CustomTestCase): (host_pool.indexer_page_num, 1, host_pool.indexer_page_stride_size), ) + def test_nsa_page_first_direct_indexer_layout_compacts_active_layers(self): + indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype + device_pool = SimpleNamespace( + store_dtype=torch.float16, + size=8, + start_layer=0, + end_layer=4, + device="cpu", + kv_lora_rank=16, + qk_rope_head_dim=4, + kv_cache_dim=24, + layer_num=4, + index_head_dim=16, + quant_block_size=8, + index_active_layer_ids=(0, 2), + index_k_with_scale_buffer=[ + torch.empty((5, 96), dtype=indexer_dtype) for _ in range(2) + ], + ) + + host_pool = NSATokenToKVPoolHost( + device_pool=device_pool, + host_to_device_ratio=2.0, + host_size=0, + page_size=4, + layout="page_first_direct", + pin_memory=False, + device="cpu", + host_token_capacity=16, + ) + + self.assertEqual(host_pool.index_active_layer_ids, (0, 2)) + self.assertEqual(host_pool._host_index_layer_slot(0), 0) + self.assertEqual(host_pool._host_index_layer_slot(2), 1) + self.assertEqual( + tuple(host_pool.index_k_with_scale_buffer.shape), + (host_pool.indexer_page_num, 2, 1, host_pool.indexer_page_stride_size), + ) + + def test_nsa_layer_page_first_indexer_layout_compacts_active_layers(self): + indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype + device_pool = SimpleNamespace( + store_dtype=torch.float16, + size=8, + start_layer=0, + end_layer=4, + device="cpu", + kv_lora_rank=16, + qk_rope_head_dim=4, + kv_cache_dim=24, + layer_num=4, + index_head_dim=16, + quant_block_size=8, + index_active_layer_ids=(0, 2), + index_k_with_scale_buffer=[ + torch.empty((5, 96), dtype=indexer_dtype) for _ in range(2) + ], + ) + + host_pool = NSATokenToKVPoolHost( + device_pool=device_pool, + host_to_device_ratio=2.0, + host_size=0, + page_size=4, + layout="layer_page_first", + pin_memory=False, + device="cpu", + host_token_capacity=16, + ) + + self.assertEqual(host_pool.index_active_layer_ids, (0, 2)) + self.assertEqual(host_pool._host_index_layer_slot(0), 0) + self.assertEqual(host_pool._host_index_layer_slot(2), 1) + self.assertEqual( + tuple(host_pool.index_k_with_scale_buffer.shape), + (2, host_pool.indexer_page_num, 1, host_pool.indexer_page_stride_size), + ) + self.assertEqual(len(host_pool.index_k_data_refs), 2) + def test_mha_layer_page_first_page_buffer_meta_fails_fast_for_storage(self): device_pool = SimpleNamespace( store_dtype=torch.float16, @@ -285,8 +364,9 @@ class TestNSAHiCacheTransfer(CustomTestCase): device="cuda" if io_backend == "kernel" else "cpu", dtype=torch.int64, ) + index_device = "cuda" if io_backend == "kernel" else "cpu" device_indices = self._token_indices_for_pages( - device_pages, page_size, device="cuda" + device_pages, page_size, device=index_device ) host_indices = self._token_indices_for_pages( host_pages, @@ -1719,6 +1799,70 @@ class TestNSAIndexerPageIndices(CustomTestCase): self.assertTrue(pool.is_index_layer_active(11)) self.assertEqual(pool.get_index_layer_slot(11), 7) + def test_nsa_state_buf_infos_returns_active_index_layers_only(self): + pool = object.__new__(NSATokenToKVPool) + pool.start_layer = 4 + pool.end_layer = 12 + pool.layer_num = 8 + pool.index_k_with_scale_buffer = [ + torch.empty((3, 8), dtype=torch.uint8) for _ in range(pool.layer_num) + ] + pool._init_index_layer_metadata( + index_active_layer_ids=(4, 8), + compact_index_layers=False, + ) + + data_ptrs, data_lens, item_lens = pool.get_state_buf_infos() + + self.assertEqual( + data_ptrs, + [ + pool.index_k_with_scale_buffer[0].data_ptr(), + pool.index_k_with_scale_buffer[4].data_ptr(), + ], + ) + self.assertEqual( + data_lens, + [ + pool.index_k_with_scale_buffer[0].nbytes, + pool.index_k_with_scale_buffer[4].nbytes, + ], + ) + self.assertEqual( + item_lens, + [ + pool.index_k_with_scale_buffer[0][0].nbytes, + pool.index_k_with_scale_buffer[4][0].nbytes, + ], + ) + self.assertEqual(pool.get_state_layer_ids(), [4, 8]) + + def test_nsa_device_pool_compact_index_layers_allocates_active_slots_only(self): + if not torch.cuda.is_available(): + self.skipTest("CUDA is required for compact NSA pool allocation test.") + + pool = NSATokenToKVPool( + size=128, + page_size=64, + kv_lora_rank=128, + dtype=torch.bfloat16, + qk_rope_head_dim=32, + layer_num=8, + device="cuda", + enable_memory_saver=False, + kv_cache_dim=576, + index_head_dim=128, + start_layer=4, + end_layer=12, + index_active_layer_ids=(4, 8), + compact_index_layers=True, + ) + + self.assertEqual(len(pool.index_k_with_scale_buffer), 2) + self.assertEqual(pool.get_index_layer_slot(4), 0) + self.assertEqual(pool.get_index_layer_slot(8), 1) + self.assertEqual(len(pool.get_state_buf_infos()[0]), 2) + def test_indexer_page_indices_accepts_valid_page_spans(self): host_pool = self.make_host_pool_stub(page_size=4) diff --git a/test/registered/unit/model_executor/test_model_runner_kv_cache_mixin.py b/test/registered/unit/model_executor/test_model_runner_kv_cache_mixin.py new file mode 100644 index 000000000..0359632fb --- /dev/null +++ b/test/registered/unit/model_executor/test_model_runner_kv_cache_mixin.py @@ -0,0 +1,41 @@ +from types import SimpleNamespace + +import torch + +from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool +from sglang.srt.model_executor.model_runner_kv_cache_mixin import ( + ModelRunnerKVCacheMixin, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +def test_nsa_cell_size_uses_active_index_layer_count(): + hf_config = SimpleNamespace( + architectures=["DeepseekV3ForCausalLM"], + index_topk=2048, + index_topk_freq=4, + index_head_dim=128, + ) + runner = SimpleNamespace( + use_mla_backend=True, + kv_cache_dtype=torch.bfloat16, + model_config=SimpleNamespace( + hf_config=hf_config, + kv_lora_rank=128, + qk_rope_head_dim=32, + ), + start_layer=0, + end_layer=12, + is_draft_worker=False, + ) + + cell_size = ModelRunnerKVCacheMixin.get_cell_size_per_token(runner, num_layers=12) + + indexer_size_per_token = ( + hf_config.index_head_dim + + hf_config.index_head_dim // NSATokenToKVPool.quant_block_size * 4 + ) + expected = (128 + 32) * 12 * 2 + indexer_size_per_token * 4 + assert cell_size == expected