Reuse CP shared KV remaps across layer materialization
CP shared KV materialization repeatedly rebuilt the same logical-page slot remaps and page inverse metadata for each layer. Cache the token and paged remap metadata on the forward batch so MLA KV, index K/scale, and prefetch paths can reuse the layer-independent mapping while still materializing layer-specific data through the existing tai/torch runtime paths. Constraint: Only mapping metadata is batch-scoped; dense KV/index contents remain layer-specific and are not reused. Rejected: Cache fully materialized dense KV/index buffers | would add large per-layer memory residency and invalidation complexity. Confidence: medium Scope-risk: moderate Directive: Do not assume this removes materialize or CP all-reduce cost; profile tai fallback logs and Nsight kernels before attributing E2E gains or losses. Tested: git diff --check Tested: remote g0034 container PYTHONPATH=python python3 -m pytest test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -q (52 passed, 5 warnings) Not-tested: Full GLM-5 disaggregated E2E performance run
This commit is contained in:
@@ -252,6 +252,196 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
self.assertEqual(remap.dense_num_pages, 7)
|
||||
self.assertEqual(remap.dense_locs.tolist(), [[4, 8, -1], [16, 0, 20]])
|
||||
|
||||
def test_forward_batch_token_slot_remap_is_cached_across_layers(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0)
|
||||
forward_batch = SimpleNamespace()
|
||||
kv_cache = torch.arange(0, 32, dtype=torch.float32).view(32, 1, 1)
|
||||
remap_logical_pages = torch.tensor([[1, 2, 5]], dtype=torch.int64)
|
||||
|
||||
remap_a = runtime.get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
remap_b = runtime.get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
|
||||
self.assertIs(remap_a, remap_b)
|
||||
|
||||
with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
|
||||
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=torch.tensor([4, 20, -1], dtype=torch.int64),
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
slot_remap=remap_a,
|
||||
)
|
||||
|
||||
self.assertEqual(dense_locs.tolist(), [4, 12, -1])
|
||||
self.assertEqual(list(dense_kv.shape), [16, 1, 1])
|
||||
self.assertTrue(torch.equal(dense_kv[4:8], kv_cache[4:8]))
|
||||
self.assertTrue(torch.equal(dense_kv[12:16], kv_cache[20:24]))
|
||||
|
||||
def test_token_slot_remap_cache_miss_logs_after_warm_cache(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
runtime._SLOT_REMAP_CACHE_LOG_COUNTS.clear()
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0)
|
||||
forward_batch = SimpleNamespace(batch_size=1, forward_mode="extend")
|
||||
kv_cache = torch.zeros((32, 1, 1), dtype=torch.float32)
|
||||
logical_pages_a = torch.tensor([[1, 2, 5]], dtype=torch.int64)
|
||||
logical_pages_b = torch.tensor([[1, 2, 6]], dtype=torch.int64)
|
||||
|
||||
with patch.object(runtime.logger, "warning") as warning:
|
||||
remap_a = runtime.get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
remap_logical_pages=logical_pages_a,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
warning.assert_not_called()
|
||||
|
||||
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
|
||||
remap_b = runtime.get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
remap_logical_pages=logical_pages_b,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
|
||||
self.assertIsNot(remap_a, remap_b)
|
||||
self.assertTrue(
|
||||
any(
|
||||
"token slot remap cache not reused (key_mismatch)" in message
|
||||
for message in logs.output
|
||||
)
|
||||
)
|
||||
|
||||
def test_token_slot_remap_incomplete_cache_state_logs(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
runtime._SLOT_REMAP_CACHE_LOG_COUNTS.clear()
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0)
|
||||
forward_batch = SimpleNamespace(batch_size=1, forward_mode="extend")
|
||||
kv_cache = torch.zeros((32, 1, 1), dtype=torch.float32)
|
||||
logical_pages = torch.tensor([[1, 2, 5]], dtype=torch.int64)
|
||||
forward_batch.cp_shared_kv_token_slot_remap_key = runtime._slot_remap_cache_key(
|
||||
logical_pages=logical_pages,
|
||||
physical_page_capacity=kv_cache.shape[0] // 4,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
forward_batch.cp_shared_kv_token_slot_remap = None
|
||||
|
||||
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
|
||||
runtime.get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
remap_logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
|
||||
self.assertTrue(
|
||||
any(
|
||||
"token slot remap cache not reused (missing_cached_value)"
|
||||
in message
|
||||
for message in logs.output
|
||||
)
|
||||
)
|
||||
|
||||
def test_forward_batch_paged_slot_remap_is_cached_across_layers(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0)
|
||||
forward_batch = SimpleNamespace()
|
||||
page_buffer = torch.arange(0, 8 * 3, dtype=torch.uint8).view(8, 3)
|
||||
logical_pages = torch.tensor([[1, 2, 5]], dtype=torch.int64)
|
||||
|
||||
remap_a = runtime.get_or_build_shared_paged_buffer_slot_remap(
|
||||
forward_batch,
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
)
|
||||
remap_b = runtime.get_or_build_shared_paged_buffer_slot_remap(
|
||||
forward_batch,
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
)
|
||||
|
||||
self.assertIs(remap_a, remap_b)
|
||||
|
||||
with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
|
||||
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
slot_remap=remap_a,
|
||||
)
|
||||
|
||||
self.assertEqual(dense_pages.tolist(), [[1, 2, 3]])
|
||||
self.assertEqual(list(dense_page_buffer.shape), [4, 3])
|
||||
self.assertTrue(torch.equal(dense_page_buffer[1], page_buffer[1]))
|
||||
self.assertTrue(torch.equal(dense_page_buffer[2], page_buffer[2]))
|
||||
self.assertTrue(torch.equal(dense_page_buffer[3], page_buffer[5]))
|
||||
|
||||
def test_paged_slot_remap_cache_miss_logs_after_warm_cache(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
runtime._SLOT_REMAP_CACHE_LOG_COUNTS.clear()
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0)
|
||||
forward_batch = SimpleNamespace(batch_size=1, forward_mode="extend")
|
||||
page_buffer = torch.zeros((8, 3), dtype=torch.uint8)
|
||||
logical_pages_a = torch.tensor([[1, 2, 5]], dtype=torch.int64)
|
||||
logical_pages_b = torch.tensor([[1, 2, 6]], dtype=torch.int64)
|
||||
|
||||
with patch.object(runtime.logger, "warning") as warning:
|
||||
remap_a = runtime.get_or_build_shared_paged_buffer_slot_remap(
|
||||
forward_batch,
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages_a,
|
||||
layout=layout,
|
||||
)
|
||||
warning.assert_not_called()
|
||||
|
||||
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
|
||||
remap_b = runtime.get_or_build_shared_paged_buffer_slot_remap(
|
||||
forward_batch,
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages_b,
|
||||
layout=layout,
|
||||
)
|
||||
|
||||
self.assertIsNot(remap_a, remap_b)
|
||||
self.assertTrue(
|
||||
any(
|
||||
"paged slot remap cache not reused (key_mismatch)" in message
|
||||
for message in logs.output
|
||||
)
|
||||
)
|
||||
|
||||
def test_mla_prefetch_log_env_defaults_to_off_and_can_enable(self):
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
@@ -1350,6 +1540,8 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
|
||||
return self.index_buffer
|
||||
|
||||
class FakeLayout:
|
||||
page_size = 4
|
||||
cp_size = 1
|
||||
cp_rank = 3
|
||||
|
||||
class MissingPrefetcher:
|
||||
@@ -1403,6 +1595,8 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
|
||||
return self.index_buffer
|
||||
|
||||
class FakeLayout:
|
||||
page_size = 4
|
||||
cp_size = 1
|
||||
cp_rank = 0
|
||||
|
||||
class MissingPrefetcher:
|
||||
|
||||
Reference in New Issue
Block a user