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 <omx@oh-my-codex.dev>
This commit is contained in:
laoyao0822
2026-06-10 05:37:06 +08:00
co-authored by OmX
parent d21952b903
commit 1ebde44e59
14 changed files with 611 additions and 16 deletions
@@ -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",
@@ -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],
)
@@ -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)
@@ -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