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