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