Move CP shared-KV prefetch reduce into attention overlap

Phase8 prefetch was starting local materialization and the CP all-reduce from the same late hook, after rank-skewed MQA/topk/materialize work had already separated CP ranks. This commit separates local prefix materialization from reduce launch, starts next-layer prefetch earlier in MLA prepare, and explicitly launches pending MLA/index reduces at attention boundaries so the collectives can overlap the intended window instead of drifting behind skewed ranks.

Constraint: Prefetch remains optional and gated by the existing CP shared-KV prefetch environment controls.

Constraint: Profiling required per-source materialize all-reduce ranges, so NVTX labels are gated behind SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX.

Rejected: Launch all-reduce immediately in start_next_layer_prefix | this reproduced the old skew-sensitive timing and could put prefetch behind current-layer work.

Rejected: Remove later hooks entirely | they are still needed as fallbacks when the early MLA prepare hook is bypassed.

Confidence: medium

Scope-risk: moderate

Directive: Preserve the consume-time fallback launch; otherwise missed launch paths silently lose correctness or overlap.

Tested: Remote g0034 docker py_compile for changed SGLang files; python -m pytest -q test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -q (50 passed).

Not-tested: Full GLM5 multi-node prefill/decode profile after this exact commit.
This commit is contained in:
laoyao0822
2026-05-12 20:31:19 +08:00
parent aa27a444f6
commit 3fc7a5c18c
7 changed files with 850 additions and 107 deletions
@@ -9,7 +9,192 @@ from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
def _identity_all_reduce(buffer, *args, **kwargs):
return buffer
class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def test_mla_prefetch_materializes_on_current_stream_and_defers_reduce_to_launch(
self,
):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
class FakeStream:
def __init__(self, name):
self.name = name
self.waited = []
def wait_stream(self, stream):
self.waited.append(stream.name)
class FakeStreamContext:
def __init__(self, active_stream, stream):
self.active_stream = active_stream
self.stream = stream
self.previous = None
def __enter__(self):
self.previous = self.active_stream[0]
self.active_stream[0] = self.stream.name
def __exit__(self, exc_type, exc, tb):
self.active_stream[0] = self.previous
class FakePool:
start_layer = 0
page_size = 4
kv_buffer = [object(), object(), object()]
def __init__(self):
self.kv_cache = torch.zeros((32, 1, 2), dtype=torch.float32)
def get_key_buffer(self, layer_id):
return self.kv_cache
active_stream = ["current"]
current_stream = FakeStream("current")
prefetch_stream = FakeStream("prefetch")
calls = []
def record_materialize(**kwargs):
calls.append(("materialize", active_stream[0]))
def record_reduce(buffer, cp_size, stream, **kwargs):
calls.append(("reduce", active_stream[0], stream.name))
return object()
prefetcher = prefetch.CpSharedKVMlaPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
slot_logical_pages=torch.tensor([0, 1, 2], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2], dtype=torch.int64),
dense_num_pages=4,
stream=prefetch_stream,
)
with patch.object(
prefetch.torch.cuda, "current_stream", return_value=current_stream
), patch.object(
prefetch.torch.cuda,
"stream",
side_effect=lambda stream: FakeStreamContext(active_stream, stream),
), patch.object(
prefetch,
"materialize_local_token_kv_page_slots_into",
side_effect=record_materialize,
), patch.object(
prefetch,
"_all_reduce_materialized_buffer_async",
side_effect=record_reduce,
):
prefetcher.start_next_layer_prefix(
next_layer_id=1,
token_to_kv_pool=FakePool(),
)
self.assertEqual(calls, [("materialize", "current")])
self.assertEqual(prefetch_stream.waited, [])
prefetcher.launch_pending_reduce()
self.assertEqual(
calls,
[("materialize", "current"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
def test_index_prefetch_materializes_on_current_stream_and_defers_reduce_to_launch(
self,
):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
class FakeStream:
def __init__(self, name):
self.name = name
self.waited = []
def wait_stream(self, stream):
self.waited.append(stream.name)
class FakeStreamContext:
def __init__(self, active_stream, stream):
self.active_stream = active_stream
self.stream = stream
self.previous = None
def __enter__(self):
self.previous = self.active_stream[0]
self.active_stream[0] = self.stream.name
def __exit__(self, exc_type, exc, tb):
self.active_stream[0] = self.previous
class FakePool:
start_layer = 0
page_size = 4
kv_buffer = [object(), object(), object()]
def __init__(self):
self.page_buffer = torch.zeros((16, 3), dtype=torch.uint8)
def get_index_k_with_scale_buffer(self, layer_id):
return self.page_buffer
active_stream = ["current"]
current_stream = FakeStream("current")
prefetch_stream = FakeStream("prefetch")
calls = []
def record_materialize(**kwargs):
calls.append(("materialize", active_stream[0]))
def record_reduce(buffer, cp_size, stream, **kwargs):
calls.append(("reduce", active_stream[0], stream.name))
return object()
prefetcher = prefetch.CpSharedKVIndexPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
prefix_pages=2,
slot_logical_pages=torch.tensor([0, 1, 2], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2], dtype=torch.int64),
dense_num_pages=4,
stream=prefetch_stream,
)
with patch.object(
prefetch.torch.cuda, "current_stream", return_value=current_stream
), patch.object(
prefetch.torch.cuda,
"stream",
side_effect=lambda stream: FakeStreamContext(active_stream, stream),
), patch.object(
prefetch,
"materialize_local_paged_buffer_page_slots_into",
side_effect=record_materialize,
), patch.object(
prefetch,
"_all_reduce_materialized_buffer_async",
side_effect=record_reduce,
):
prefetcher.start_next_layer_prefix(
next_layer_id=1,
token_to_kv_pool=FakePool(),
)
self.assertEqual(calls, [("materialize", "current")])
self.assertEqual(prefetch_stream.waited, [])
prefetcher.launch_pending_reduce()
self.assertEqual(
calls,
[("materialize", "current"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
def test_all_reduce_uses_group_fast_path_for_float_buffers(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
@@ -278,7 +463,9 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIs(remap_a, remap_b)
with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
with patch.object(
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
):
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
kv_cache=kv_cache,
logical_locs=torch.tensor([4, 20, -1], dtype=torch.int64),
@@ -391,7 +578,9 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIs(remap_a, remap_b)
with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
with patch.object(
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
):
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
page_buffer=page_buffer,
logical_pages=logical_pages,
@@ -828,7 +1017,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with patch.object(
runtime, "cp_shared_kv_debug_enabled", return_value=True
), patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
), patch.object(runtime, "_all_reduce_materialized_buffer", _identity_all_reduce):
with self.assertRaisesRegex(
RuntimeError,
"CP shared KV materialize got logical token locs outside the physical pool",
@@ -853,7 +1042,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with patch.object(
runtime, "cp_shared_kv_debug_enabled", return_value=False
), patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
), patch.object(runtime, "_all_reduce_materialized_buffer", _identity_all_reduce):
_, dense_locs = runtime.materialize_shared_token_kv_buffer(
kv_cache=kv_cache,
logical_locs=logical_locs,
@@ -872,7 +1061,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
logical_locs = torch.tensor([4, -1, 8], dtype=torch.int64)
with patch.object(runtime, "cp_shared_kv_debug_enabled", return_value=True):
with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
with patch.object(runtime, "_all_reduce_materialized_buffer", _identity_all_reduce):
_, dense_locs = runtime.materialize_shared_token_kv_buffer(
kv_cache=kv_cache,
logical_locs=logical_locs,
@@ -944,7 +1133,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
logical_locs = torch.tensor([8, 20, -1], dtype=torch.int64)
remap_source_locs = torch.tensor([4, 8, 12, 16, 20], dtype=torch.int64)
with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
with patch.object(runtime, "_all_reduce_materialized_buffer", _identity_all_reduce):
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
kv_cache=kv_cache,
logical_locs=logical_locs,
@@ -971,7 +1160,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with patch.object(
runtime, "cp_shared_kv_debug_enabled", return_value=False
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
), patch.object(
runtime.torch, "any", side_effect=AssertionError("torch.any sync")
), patch.object(
@@ -995,7 +1184,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
kv_cache = torch.arange(0, 32, dtype=torch.float32).view(32, 1, 1)
remap_source_locs = torch.tensor([4, 8, 12, 16, 20], dtype=torch.int64)
with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
with patch.object(runtime, "_all_reduce_materialized_buffer", _identity_all_reduce):
dense_kv_a, dense_locs_a = runtime.materialize_shared_token_kv_buffer(
kv_cache=kv_cache,
logical_locs=torch.tensor([4, 20], dtype=torch.int64),
@@ -1026,7 +1215,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with patch.object(
runtime, "cp_shared_kv_debug_enabled", return_value=True
), patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x):
), patch.object(runtime, "_all_reduce_materialized_buffer", _identity_all_reduce):
with self.assertRaisesRegex(
RuntimeError,
"CP shared KV materialize got logical pages outside the physical page buffer",
@@ -1050,7 +1239,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with patch.object(
runtime, "cp_shared_kv_debug_enabled", return_value=False
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
), patch.object(
runtime.torch, "any", side_effect=AssertionError("torch.any sync")
), patch.object(
@@ -1076,7 +1265,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with patch.object(
runtime, "cp_shared_kv_debug_enabled", return_value=False
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
), patch.object(
runtime.torch, "unique", side_effect=AssertionError("torch.unique sync")
), patch.object(
@@ -1111,7 +1300,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with patch.object(
runtime, "cp_shared_kv_debug_enabled", return_value=False
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
), patch.object(
runtime.torch, "unique", side_effect=AssertionError("torch.unique sync")
), patch.object(
@@ -1240,7 +1429,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
), patch.object(
runtime, "_load_tai_materialize_kernels", return_value=fake_tai
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
):
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
page_buffer=page_buffer,
@@ -1307,7 +1496,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
), patch.object(
runtime, "_load_tai_materialize_kernels", return_value=fake_tai
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
):
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
kv_cache=kv_cache,
@@ -1379,7 +1568,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
"build_slot_page_remap",
side_effect=AssertionError("tai token path must not run torch remap"),
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
):
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
kv_cache=kv_cache,
@@ -1410,7 +1599,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
"_load_tai_materialize_kernels",
side_effect=AssertionError("tai path must stay off in debug mode"),
), patch.object(
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
):
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
page_buffer=page_buffer,