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