Reduce inactive NSA index-cache transfer safely

Centralize the IndexCache skip formula and thread the resulting active logical index layers into NSA KV pools. HiCache now skips only the indexer H2D/D2H payload for inactive target layers while preserving per-layer MLA KV transfer, keeping allocation shape unchanged for this phase.

Constraint: P0-P2 must not compact device or host allocation yet; prefill/decode state transfer still has no logical layer-id metadata.

Rejected: Recompute the skip formula separately in mem_cache | formula drift would corrupt cache or waste transfers when offset/pattern settings change.

Rejected: Skip whole-layer HiCache load/backup | MLA KV remains required for every attention layer.

Confidence: medium

Scope-risk: moderate

Directive: Before enabling compact state buffers or compact allocation, add layer-id metadata validation to PD transfer.

Tested: Local py_compile for touched files; remote pytest in g0034 container: test_nsa_index_layers.py and TestNSAIndexerPageIndices, 20 passed.

Not-tested: ETE replay/GSM8K with --nsa-index-topk-freq 4; PD state-transfer compaction remains unimplemented.
This commit is contained in:
laoyao0822
2026-06-10 04:28:26 +08:00
parent 6229c7da60
commit d21952b903
8 changed files with 567 additions and 90 deletions
@@ -0,0 +1,83 @@
from types import SimpleNamespace
import pytest
from sglang.srt.configs.nsa_index_layers import (
build_nsa_index_layer_plan,
nsa_index_skip_flags,
)
def test_default_freq_one_marks_all_target_layers_active():
cfg = SimpleNamespace(index_topk_freq=1)
plan = build_nsa_index_layer_plan(cfg, 0, 6)
assert plan.active_layer_ids == (0, 1, 2, 3, 4, 5)
assert [nsa_index_skip_flags(cfg, i)[0] for i in range(6)] == [False] * 6
def test_freq_four_without_offset_matches_current_model_formula():
cfg = SimpleNamespace(index_topk_freq=4)
plan = build_nsa_index_layer_plan(cfg, 0, 12)
assert plan.active_layer_ids == (0, 1, 5, 9)
assert [nsa_index_skip_flags(cfg, i)[0] for i in range(10)] == [
False,
False,
True,
True,
True,
False,
True,
True,
True,
False,
]
def test_freq_four_with_offset_one_uses_layers_zero_four_eight():
cfg = SimpleNamespace(index_topk_freq=4, index_skip_topk_offset=1)
plan = build_nsa_index_layer_plan(cfg, 0, 12)
assert plan.active_layer_ids == (0, 4, 8)
assert [nsa_index_skip_flags(cfg, i)[0] for i in range(9)] == [
False,
True,
True,
True,
False,
True,
True,
True,
False,
]
def test_pattern_marks_non_shared_layers_active():
cfg = SimpleNamespace(index_topk_freq=1, index_topk_pattern="CSSSCSS")
plan = build_nsa_index_layer_plan(cfg, 0, len(cfg.index_topk_pattern))
assert plan.active_layer_ids == (0, 4)
def test_nextn_keeps_all_draft_layers_active_for_state_safety():
cfg = SimpleNamespace(index_topk_freq=4, index_skip_topk_offset=1)
plan = build_nsa_index_layer_plan(cfg, 0, 1, is_nextn=True)
assert plan.active_layer_ids == (0,)
assert nsa_index_skip_flags(cfg, 0, is_nextn=True) == (True, True)
def test_nonzero_start_layer_preserves_logical_layer_ids():
cfg = SimpleNamespace(index_topk_freq=4, index_skip_topk_offset=1)
plan = build_nsa_index_layer_plan(cfg, 4, 13)
assert plan.active_layer_ids == (4, 8, 12)
assert plan.slot_for_layer(8) == 1
def test_inactive_slot_lookup_fails_fast():
cfg = SimpleNamespace(index_topk_freq=4, index_skip_topk_offset=1)
plan = build_nsa_index_layer_plan(cfg, 0, 8)
with pytest.raises(RuntimeError, match="inactive index layer requested"):
plan.slot_for_layer(1)
def test_invalid_offset_fails_before_layer_zero_can_skip_without_prior_topk():
cfg = SimpleNamespace(index_topk_freq=4, index_skip_topk_offset=0)
with pytest.raises(ValueError, match="index_skip_topk_offset"):
nsa_index_skip_flags(cfg, 0)
@@ -1684,6 +1684,41 @@ class TestNSAIndexerPageIndices(CustomTestCase):
host_pool.page_size = page_size
return host_pool
def test_nsa_device_pool_active_index_layers_use_full_allocation_slots(self):
pool = object.__new__(NSATokenToKVPool)
pool.start_layer = 4
pool.end_layer = 12
pool.layer_num = 8
pool._init_index_layer_metadata(
index_active_layer_ids=(4, 8),
compact_index_layers=False,
)
self.assertEqual(pool.index_active_layer_ids, (4, 8))
self.assertTrue(pool.is_index_layer_active(4))
self.assertTrue(pool.is_index_layer_active(8))
self.assertFalse(pool.is_index_layer_active(5))
self.assertEqual(pool.get_index_layer_slot(4), 0)
self.assertEqual(pool.get_index_layer_slot(8), 4)
with self.assertRaisesRegex(RuntimeError, "inactive index layer requested"):
pool.get_index_layer_slot(5)
def test_nsa_device_pool_default_active_layers_cover_local_range(self):
pool = object.__new__(NSATokenToKVPool)
pool.start_layer = 4
pool.end_layer = 12
pool.layer_num = 8
pool._init_index_layer_metadata(
index_active_layer_ids=None,
compact_index_layers=False,
)
self.assertEqual(pool.index_active_layer_ids, tuple(range(4, 12)))
self.assertTrue(pool.is_index_layer_active(11))
self.assertEqual(pool.get_index_layer_slot(11), 7)
def test_indexer_page_indices_accepts_valid_page_spans(self):
host_pool = self.make_host_pool_stub(page_size=4)
@@ -1773,6 +1808,118 @@ class TestNSAIndexerPageIndices(CustomTestCase):
self.assertEqual(call["dst_indices"].tolist(), [0, 1])
self.assertEqual(call["page_size"], 1)
def test_page_first_direct_all_layer_indexer_backup_skips_inactive_layers(self):
host_pool = self.make_host_pool_stub(page_size=4)
host_pool.layout = "page_first_direct"
host_pool.indexer_page_stride_size = 8
host_pool.layer_num = 4
host_pool.index_k_with_scale_buffer = "host-page-first-indexer"
class FakeDevicePool:
index_active_layer_ids = (0, 2)
index_k_with_scale_buffer = [
"device-layer-0",
"device-layer-1",
"device-layer-2",
"device-layer-3",
]
def is_index_layer_active(self, layer_id):
return layer_id in self.index_active_layer_ids
def get_index_layer_slot(self, layer_id):
return layer_id
calls = []
def fake_tai_transfer(**kwargs):
calls.append(kwargs)
with patch(
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_pf",
return_value=fake_tai_transfer,
):
host_pool._backup_indexer_from_device_all_layer(
FakeDevicePool(),
torch.tensor([0, 1, 2, 3], dtype=torch.int64),
torch.tensor([8, 9, 10, 11], dtype=torch.int64),
"direct",
)
self.assertEqual([call["layer_id"] for call in calls], [0, 2])
self.assertEqual(calls[0]["src_ptrs"], ["device-layer-0"])
self.assertEqual(calls[1]["src_ptrs"], ["device-layer-2"])
def test_per_layer_indexer_backup_and_load_skip_inactive_layers(self):
host_pool = self.make_host_pool_stub(page_size=4)
host_pool.layout = "page_first_direct"
host_pool.indexer_page_stride_size = 8
host_pool.index_k_with_scale_buffer = "host-page-first-indexer"
class FakeDevicePool:
index_active_layer_ids = (2,)
index_k_with_scale_buffer = [
"device-layer-0",
"device-layer-1",
"device-layer-2",
]
def is_index_layer_active(self, layer_id):
return layer_id in self.index_active_layer_ids
def get_index_layer_slot(self, layer_id):
return layer_id
calls = []
def fake_backup_transfer(**kwargs):
calls.append(("backup", kwargs))
def fake_load_transfer(**kwargs):
calls.append(("load", kwargs))
with (
patch(
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_pf",
return_value=fake_backup_transfer,
),
patch(
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_pf_lf",
return_value=fake_load_transfer,
),
):
host_pool._backup_indexer_from_device_per_layer(
FakeDevicePool(),
torch.tensor([0, 1, 2, 3], dtype=torch.int64),
torch.tensor([8, 9, 10, 11], dtype=torch.int64),
1,
"direct",
)
host_pool._load_indexer_to_device_per_layer(
FakeDevicePool(),
torch.tensor([0, 1, 2, 3], dtype=torch.int64),
torch.tensor([8, 9, 10, 11], dtype=torch.int64),
1,
"direct",
)
host_pool._backup_indexer_from_device_per_layer(
FakeDevicePool(),
torch.tensor([0, 1, 2, 3], dtype=torch.int64),
torch.tensor([8, 9, 10, 11], dtype=torch.int64),
2,
"direct",
)
host_pool._load_indexer_to_device_per_layer(
FakeDevicePool(),
torch.tensor([0, 1, 2, 3], dtype=torch.int64),
torch.tensor([8, 9, 10, 11], dtype=torch.int64),
2,
"direct",
)
self.assertEqual([kind for kind, _ in calls], ["backup", "load"])
self.assertEqual([kwargs["layer_id"] for _, kwargs in calls], [2, 2])
def test_mla_layer_page_first_all_layer_backup_uses_tai_per_layer_route(self):
host_pool = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost)
host_pool.layout = "layer_page_first"