From 3fc7a5c18c217938e0aa6cb79c62b9a3795a1b4d Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Fri, 8 May 2026 17:39:24 +0800 Subject: [PATCH] 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. --- python/sglang/srt/environ.py | 1 + .../attention/nsa/cp_shared_kv_prefetch.py | 425 +++++++++++++++--- .../attention/nsa/cp_shared_kv_runtime.py | 134 +++++- .../srt/layers/attention/nsa/nsa_indexer.py | 45 +- .../srt/layers/attention/nsa_backend.py | 93 +++- .../attention_forward_methods/forward_mla.py | 38 ++ .../mem_cache/test_cp_shared_kv_runtime.py | 221 ++++++++- 7 files changed, 850 insertions(+), 107 deletions(-) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 929ddf3d5..f54c814b4 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -211,6 +211,7 @@ class Envs: SGLANG_CP_SHARED_KV_FUSED_INDEX_MQA_PREPARE = EnvBool(False) SGLANG_CP_SHARED_KV_ENABLE_MLA_PREFETCH = EnvBool(False) SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH = EnvBool(False) + SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX = EnvBool(False) SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES = EnvInt(-1) SGLANG_TEST_REQUEST_TIME_STATS = EnvBool(False) SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py index 613487d2d..3a8cf61f9 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py @@ -12,6 +12,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( cp_shared_kv_debug_enabled, cp_shared_kv_mla_prefetch_enabled, cp_shared_kv_mla_prefetch_log, + cp_shared_kv_mla_prefetch_log_enabled, cp_shared_kv_mla_prefetch_min_prefix_pages, cp_shared_kv_mla_prefetch_should_log_layer, filter_locs_mappable_to_physical_pool, @@ -53,18 +54,46 @@ def _is_cuda_stream_capturing() -> bool: return False +def _debug_owned_pages_count( + layout: CpSharedKVLayout, + logical_pages: torch.Tensor, +) -> int: + if not cp_shared_kv_mla_prefetch_log_enabled(): + return -1 + try: + if logical_pages.numel() == 0: + return 0 + return int(layout.owned_pages_mask(logical_pages).sum().item()) + except Exception: + logger.exception("Failed to count CP shared KV prefetch owned pages.") + return -1 + + +def _debug_handle_keys( + layer_id: int, + handles: dict[int, Any], +) -> tuple[int, ...]: + if not cp_shared_kv_mla_prefetch_should_log_layer(layer_id): + return () + if not cp_shared_kv_mla_prefetch_log_enabled(): + return () + return tuple(sorted(handles.keys())) + + @dataclass class CpSharedKVMlaPrefetchHandle: layer_id: int dense_kv_cache: torch.Tensor - event: torch.cuda.Event + prefix_rows: slice + event: Optional[torch.cuda.Event] = None @dataclass class CpSharedKVIndexPrefetchHandle: layer_id: int dense_page_buffer: torch.Tensor - event: torch.cuda.Event + prefix_rows: slice + event: Optional[torch.cuda.Event] = None class CpSharedKVMlaPrefetcher: @@ -84,6 +113,8 @@ class CpSharedKVMlaPrefetcher: slot_logical_pages: torch.Tensor, page_inverse: torch.Tensor, dense_num_pages: int, + owned_prefix_pages: int = -1, + owned_total_pages: int = -1, stream: Optional[torch.cuda.Stream] = None, ) -> None: self.layout = layout @@ -92,6 +123,8 @@ class CpSharedKVMlaPrefetcher: self.slot_logical_pages = slot_logical_pages self.page_inverse = page_inverse self.dense_num_pages = dense_num_pages + self.owned_prefix_pages = owned_prefix_pages + self.owned_total_pages = owned_total_pages self.total_slots = int(slot_logical_pages.numel()) self.stream = stream if stream is not None else torch.cuda.Stream() self.handles: dict[int, CpSharedKVMlaPrefetchHandle] = {} @@ -186,6 +219,14 @@ class CpSharedKVMlaPrefetcher: return None min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages(layout.cp_size) if prefix_pages < min_prefix_pages: + _prefetch_log( + "create_skip reason=prefix_below_min cp_rank=%s cp_size=%s " + "prefix_pages=%s min_prefix_pages=%s", + layout.cp_rank, + layout.cp_size, + prefix_pages, + min_prefix_pages, + ) return None cp_group = get_attention_cp_group() @@ -211,12 +252,20 @@ class CpSharedKVMlaPrefetcher: logger.exception("Failed to initialize CP shared KV MLA prefetcher.") return None + owned_prefix_pages = _debug_owned_pages_count( + layout, remap.slot_logical_pages[:prefix_pages] + ) + owned_total_pages = _debug_owned_pages_count(layout, remap.slot_logical_pages) + _prefetch_log( - "create cp_rank=%s cp_size=%s prefix_pages=%s total_slots=%s dense_pages=%s page_size=%s", + "create cp_rank=%s cp_size=%s prefix_pages=%s total_slots=%s " + "owned_prefix_pages=%s owned_total_pages=%s dense_pages=%s page_size=%s", layout.cp_rank, layout.cp_size, prefix_pages, int(remap.slot_logical_pages.numel()), + owned_prefix_pages, + owned_total_pages, remap.dense_num_pages, page_size, ) @@ -228,6 +277,8 @@ class CpSharedKVMlaPrefetcher: slot_logical_pages=remap.slot_logical_pages, page_inverse=remap.page_inverse, dense_num_pages=remap.dense_num_pages, + owned_prefix_pages=owned_prefix_pages, + owned_total_pages=owned_total_pages, stream=stream, ) @@ -253,10 +304,29 @@ class CpSharedKVMlaPrefetcher: ) return None - handle = self.handles.pop(layer_id, None) + self._log_layer( + layer_id, + "consume_enter layer=%s prefix_pages=%s total_slots=%s handles=%s", + layer_id, + self.prefix_pages, + self.total_slots, + _debug_handle_keys(layer_id, self.handles), + ) + handle = self.handles.get(layer_id) if handle is None: self._log_layer(layer_id, "consume_miss layer=%s", layer_id) return None + if handle.event is None: + self.launch_pending_reduce() + handle = self.handles.get(layer_id) + if handle is None or handle.event is None: + self._log_layer( + layer_id, + "consume_miss reason=prefix_reduce_not_ready layer=%s", + layer_id, + ) + return None + handle = self.handles.pop(layer_id) if self.pending_attention_handle is handle: self.pending_attention_handle = None if handle.layer_id != layer_id: @@ -274,6 +344,15 @@ class CpSharedKVMlaPrefetcher: suffix_slots = self.total_slots - self.prefix_pages if self.prefix_pages < self.total_slots: + self._log_layer( + layer_id, + "consume_suffix_begin layer=%s start_slot=%s end_slot=%s " + "suffix_slots=%s", + layer_id, + self.prefix_pages, + self.total_slots, + suffix_slots, + ) materialize_local_token_kv_page_slots_into( kv_cache=kv_cache, dense_kv_cache=dense_kv_cache, @@ -293,6 +372,16 @@ class CpSharedKVMlaPrefetcher: self.layout.cp_size, suffix_rows.start, suffix_rows.stop, + nvtx_source="mla.consume_suffix", + nvtx_layer_id=layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) + self._log_layer( + layer_id, + "consume_suffix_done layer=%s rows=%s:%s", + layer_id, + suffix_rows.start, + suffix_rows.stop, ) self._log_layer( @@ -322,6 +411,22 @@ class CpSharedKVMlaPrefetcher: next_layer_id: int, token_to_kv_pool: Any, ) -> None: + layer_in_pool = self._layer_in_pool(token_to_kv_pool, next_layer_id) + self._log_next_layer( + next_layer_id, + "start_enter next_layer=%s disabled=%s already_started=%s " + "layer_in_pool=%s prefix_pages=%s total_slots=%s " + "owned_prefix_pages=%s owned_total_pages=%s handles=%s", + next_layer_id, + self.disabled, + next_layer_id in self.handles, + layer_in_pool, + self.prefix_pages, + self.total_slots, + self.owned_prefix_pages, + self.owned_total_pages, + _debug_handle_keys(next_layer_id, self.handles), + ) if self.disabled: self._log_next_layer( next_layer_id, @@ -336,7 +441,7 @@ class CpSharedKVMlaPrefetcher: next_layer_id, ) return - if not self._layer_in_pool(token_to_kv_pool, next_layer_id): + if not layer_in_pool: self._log_next_layer( next_layer_id, "start_skip reason=layer_out_of_pool next_layer=%s", @@ -358,40 +463,32 @@ class CpSharedKVMlaPrefetcher: ) return - current_stream = torch.cuda.current_stream() - self.stream.wait_stream(current_stream) try: - with torch.cuda.stream(self.stream): - dense_kv_cache = kv_cache.new_zeros( - (self.dense_num_pages * self.page_size, *kv_cache.shape[1:]) - ) - materialize_local_token_kv_page_slots_into( - kv_cache=kv_cache, - dense_kv_cache=dense_kv_cache, - slot_logical_pages=self.slot_logical_pages, - layout=self.layout, - page_size=self.page_size, - start_slot=0, - end_slot=self.prefix_pages, - ) - prefix_rows = slot_range_to_token_slice( - self.page_size, - 0, - self.prefix_pages, - ) - event = _all_reduce_materialized_buffer_async( - dense_kv_cache[prefix_rows], - cp_size=self.layout.cp_size, - stream=self.stream, - ) - if event is None: - self.disabled = True - self._log_next_layer( - next_layer_id, - "start_disable reason=async_reduce_unavailable next_layer=%s", - next_layer_id, - ) - return + dense_kv_cache = kv_cache.new_zeros( + (self.dense_num_pages * self.page_size, *kv_cache.shape[1:]) + ) + self._log_next_layer( + next_layer_id, + "start_prefix_begin next_layer=%s start_slot=0 end_slot=%s " + "dense_rows=%s", + next_layer_id, + self.prefix_pages, + int(dense_kv_cache.shape[0]), + ) + materialize_local_token_kv_page_slots_into( + kv_cache=kv_cache, + dense_kv_cache=dense_kv_cache, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + page_size=self.page_size, + start_slot=0, + end_slot=self.prefix_pages, + ) + prefix_rows = slot_range_to_token_slice( + self.page_size, + 0, + self.prefix_pages, + ) except Exception: logger.exception("Failed to start CP shared KV MLA prefix prefetch.") self.disabled = True @@ -405,7 +502,7 @@ class CpSharedKVMlaPrefetcher: handle = CpSharedKVMlaPrefetchHandle( layer_id=next_layer_id, dense_kv_cache=dense_kv_cache, - event=event, + prefix_rows=prefix_rows, ) self.handles[next_layer_id] = handle self.pending_attention_handle = handle @@ -418,6 +515,64 @@ class CpSharedKVMlaPrefetcher: int(dense_kv_cache.shape[0]), ) + def launch_pending_reduce(self) -> None: + handle = self.pending_attention_handle + if handle is None: + return + if handle.event is not None: + self._log_next_layer( + handle.layer_id, + "start_prefix_reduce_skip reason=already_enqueued next_layer=%s", + handle.layer_id, + ) + return + + prefix_rows = handle.prefix_rows + try: + current_stream = torch.cuda.current_stream() + self.stream.wait_stream(current_stream) + with torch.cuda.stream(self.stream): + event = _all_reduce_materialized_buffer_async( + handle.dense_kv_cache[prefix_rows], + cp_size=self.layout.cp_size, + stream=self.stream, + nvtx_source="mla.prefetch_prefix", + nvtx_layer_id=handle.layer_id, + nvtx_cp_rank=self.layout.cp_rank, + nvtx_rows=(prefix_rows.start, prefix_rows.stop), + ) + if event is None: + self.disabled = True + self.handles.pop(handle.layer_id, None) + if self.pending_attention_handle is handle: + self.pending_attention_handle = None + self._log_next_layer( + handle.layer_id, + "start_disable reason=async_reduce_unavailable next_layer=%s", + handle.layer_id, + ) + return + handle.event = event + self._log_next_layer( + handle.layer_id, + "start_prefix_reduce_enqueued next_layer=%s rows=%s:%s", + handle.layer_id, + prefix_rows.start, + prefix_rows.stop, + ) + except Exception: + logger.exception("Failed to launch CP shared KV MLA prefix prefetch reduce.") + self.disabled = True + self.handles.pop(handle.layer_id, None) + if self.pending_attention_handle is handle: + self.pending_attention_handle = None + self._log_next_layer( + handle.layer_id, + "start_disable reason=reduce_launch_exception next_layer=%s", + handle.layer_id, + ) + return + def wait_attention_window(self) -> None: handle = self.pending_attention_handle if handle is None: @@ -462,6 +617,8 @@ class CpSharedKVIndexPrefetcher: slot_logical_pages: torch.Tensor, page_inverse: torch.Tensor, dense_num_pages: int, + owned_prefix_pages: int = -1, + owned_total_pages: int = -1, stream: Optional[torch.cuda.Stream] = None, ) -> None: self.layout = layout @@ -469,6 +626,8 @@ class CpSharedKVIndexPrefetcher: self.slot_logical_pages = slot_logical_pages self.page_inverse = page_inverse self.dense_num_pages = dense_num_pages + self.owned_prefix_pages = owned_prefix_pages + self.owned_total_pages = owned_total_pages self.total_slots = int(slot_logical_pages.numel()) self.stream = stream if stream is not None else torch.cuda.Stream() self.handles: dict[int, CpSharedKVIndexPrefetchHandle] = {} @@ -606,6 +765,14 @@ class CpSharedKVIndexPrefetcher: return None min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages(layout.cp_size) if prefix_pages < min_prefix_pages: + _prefetch_log( + "index_create_skip reason=prefix_below_min cp_rank=%s cp_size=%s " + "prefix_pages=%s min_prefix_pages=%s", + layout.cp_rank, + layout.cp_size, + prefix_pages, + min_prefix_pages, + ) return None cp_group = get_attention_cp_group() @@ -638,12 +805,20 @@ class CpSharedKVIndexPrefetcher: logger.exception("Failed to initialize CP shared KV index prefetcher.") return None + owned_prefix_pages = _debug_owned_pages_count( + layout, remap.slot_logical_pages[:prefix_pages] + ) + owned_total_pages = _debug_owned_pages_count(layout, remap.slot_logical_pages) + _prefetch_log( - "index_create cp_rank=%s cp_size=%s prefix_pages=%s total_slots=%s dense_pages=%s page_size=%s", + "index_create cp_rank=%s cp_size=%s prefix_pages=%s total_slots=%s " + "owned_prefix_pages=%s owned_total_pages=%s dense_pages=%s page_size=%s", layout.cp_rank, layout.cp_size, prefix_pages, int(remap.slot_logical_pages.numel()), + owned_prefix_pages, + owned_total_pages, remap.dense_num_pages, page_size, ) @@ -654,6 +829,8 @@ class CpSharedKVIndexPrefetcher: slot_logical_pages=remap.slot_logical_pages, page_inverse=remap.page_inverse, dense_num_pages=remap.dense_num_pages, + owned_prefix_pages=owned_prefix_pages, + owned_total_pages=owned_total_pages, stream=stream, ) @@ -679,10 +856,29 @@ class CpSharedKVIndexPrefetcher: ) return None - handle = self.handles.pop(layer_id, None) + self._log_layer( + layer_id, + "index_consume_enter layer=%s prefix_pages=%s total_slots=%s handles=%s", + layer_id, + self.prefix_pages, + self.total_slots, + _debug_handle_keys(layer_id, self.handles), + ) + handle = self.handles.get(layer_id) if handle is None: self._log_layer(layer_id, "index_consume_miss layer=%s", layer_id) return None + if handle.event is None: + self.launch_pending_reduce() + handle = self.handles.get(layer_id) + if handle is None or handle.event is None: + self._log_layer( + layer_id, + "index_consume_miss reason=prefix_reduce_not_ready layer=%s", + layer_id, + ) + return None + handle = self.handles.pop(layer_id) if self.pending_attention_handle is handle: self.pending_attention_handle = None if handle.layer_id != layer_id: @@ -700,6 +896,15 @@ class CpSharedKVIndexPrefetcher: suffix_slots = self.total_slots - self.prefix_pages if self.prefix_pages < self.total_slots: + self._log_layer( + layer_id, + "index_consume_suffix_begin layer=%s start_slot=%s end_slot=%s " + "suffix_slots=%s", + layer_id, + self.prefix_pages, + self.total_slots, + suffix_slots, + ) materialize_local_paged_buffer_page_slots_into( page_buffer=page_buffer, dense_page_buffer=dense_page_buffer, @@ -717,6 +922,16 @@ class CpSharedKVIndexPrefetcher: self.layout.cp_size, suffix_rows.start, suffix_rows.stop, + nvtx_source="index.consume_suffix", + nvtx_layer_id=layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) + self._log_layer( + layer_id, + "index_consume_suffix_done layer=%s rows=%s:%s", + layer_id, + suffix_rows.start, + suffix_rows.stop, ) self._log_layer( @@ -745,6 +960,22 @@ class CpSharedKVIndexPrefetcher: next_layer_id: int, token_to_kv_pool: Any, ) -> None: + layer_in_pool = self._layer_in_pool(token_to_kv_pool, next_layer_id) + self._log_next_layer( + next_layer_id, + "index_start_enter next_layer=%s disabled=%s already_started=%s " + "layer_in_pool=%s prefix_pages=%s total_slots=%s " + "owned_prefix_pages=%s owned_total_pages=%s handles=%s", + next_layer_id, + self.disabled, + next_layer_id in self.handles, + layer_in_pool, + self.prefix_pages, + self.total_slots, + self.owned_prefix_pages, + self.owned_total_pages, + _debug_handle_keys(next_layer_id, self.handles), + ) if self.disabled: self._log_next_layer( next_layer_id, @@ -759,7 +990,7 @@ class CpSharedKVIndexPrefetcher: next_layer_id, ) return - if not self._layer_in_pool(token_to_kv_pool, next_layer_id): + if not layer_in_pool: self._log_next_layer( next_layer_id, "index_start_skip reason=layer_out_of_pool next_layer=%s", @@ -783,35 +1014,27 @@ class CpSharedKVIndexPrefetcher: ) return - current_stream = torch.cuda.current_stream() - self.stream.wait_stream(current_stream) try: - with torch.cuda.stream(self.stream): - dense_page_buffer = page_buffer.new_zeros( - (self.dense_num_pages, *page_buffer.shape[1:]) - ) - materialize_local_paged_buffer_page_slots_into( - page_buffer=page_buffer, - dense_page_buffer=dense_page_buffer, - slot_logical_pages=self.slot_logical_pages, - layout=self.layout, - start_slot=0, - end_slot=self.prefix_pages, - ) - prefix_rows = slot_range_to_page_slice(0, self.prefix_pages) - event = _all_reduce_materialized_buffer_async( - dense_page_buffer[prefix_rows], - cp_size=self.layout.cp_size, - stream=self.stream, - ) - if event is None: - self.disabled = True - self._log_next_layer( - next_layer_id, - "index_start_disable reason=async_reduce_unavailable next_layer=%s", - next_layer_id, - ) - return + dense_page_buffer = page_buffer.new_zeros( + (self.dense_num_pages, *page_buffer.shape[1:]) + ) + self._log_next_layer( + next_layer_id, + "index_start_prefix_begin next_layer=%s start_slot=0 " + "end_slot=%s dense_pages=%s", + next_layer_id, + self.prefix_pages, + int(dense_page_buffer.shape[0]), + ) + materialize_local_paged_buffer_page_slots_into( + page_buffer=page_buffer, + dense_page_buffer=dense_page_buffer, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + start_slot=0, + end_slot=self.prefix_pages, + ) + prefix_rows = slot_range_to_page_slice(0, self.prefix_pages) except Exception: logger.exception("Failed to start CP shared KV index prefix prefetch.") self.disabled = True @@ -825,7 +1048,7 @@ class CpSharedKVIndexPrefetcher: handle = CpSharedKVIndexPrefetchHandle( layer_id=next_layer_id, dense_page_buffer=dense_page_buffer, - event=event, + prefix_rows=prefix_rows, ) self.handles[next_layer_id] = handle self.pending_attention_handle = handle @@ -837,6 +1060,66 @@ class CpSharedKVIndexPrefetcher: int(dense_page_buffer.shape[0]), ) + def launch_pending_reduce(self) -> None: + handle = self.pending_attention_handle + if handle is None: + return + if handle.event is not None: + self._log_next_layer( + handle.layer_id, + "index_start_prefix_reduce_skip reason=already_enqueued next_layer=%s", + handle.layer_id, + ) + return + + prefix_rows = handle.prefix_rows + try: + current_stream = torch.cuda.current_stream() + self.stream.wait_stream(current_stream) + with torch.cuda.stream(self.stream): + event = _all_reduce_materialized_buffer_async( + handle.dense_page_buffer[prefix_rows], + cp_size=self.layout.cp_size, + stream=self.stream, + nvtx_source="index.prefetch_prefix", + nvtx_layer_id=handle.layer_id, + nvtx_cp_rank=self.layout.cp_rank, + nvtx_rows=(prefix_rows.start, prefix_rows.stop), + ) + if event is None: + self.disabled = True + self.handles.pop(handle.layer_id, None) + if self.pending_attention_handle is handle: + self.pending_attention_handle = None + self._log_next_layer( + handle.layer_id, + "index_start_disable reason=async_reduce_unavailable next_layer=%s", + handle.layer_id, + ) + return + handle.event = event + self._log_next_layer( + handle.layer_id, + "index_start_prefix_reduce_enqueued next_layer=%s rows=%s:%s", + handle.layer_id, + prefix_rows.start, + prefix_rows.stop, + ) + except Exception: + logger.exception( + "Failed to launch CP shared KV index prefix prefetch reduce." + ) + self.disabled = True + self.handles.pop(handle.layer_id, None) + if self.pending_attention_handle is handle: + self.pending_attention_handle = None + self._log_next_layer( + handle.layer_id, + "index_start_disable reason=reduce_launch_exception next_layer=%s", + handle.layer_id, + ) + return + def wait_attention_window(self) -> None: handle = self.pending_attention_handle if handle is None: diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py index 602c93f95..61a4c1e88 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -1,6 +1,7 @@ from __future__ import annotations import logging +from contextlib import contextmanager from dataclasses import dataclass from functools import lru_cache from typing import Any @@ -19,6 +20,7 @@ _TAI_FUSED_MLA_STORE_FALLBACK_LOG_COUNTS: dict[str, int] = {} _SLOT_REMAP_CACHE_LOG_COUNTS: dict[str, int] = {} _MLA_PREFETCH_LOG_PROBE_LAYER = 2 _SORT_NVTX_ENABLED = envs.SGLANG_DEBUG_SORT_NVTX.get() +_MATERIALIZE_NVTX_ENABLED = envs.SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX.get() def cp_shared_kv_debug_enabled() -> bool: @@ -33,6 +35,10 @@ def cp_shared_kv_sort_nvtx_enabled() -> bool: return _SORT_NVTX_ENABLED +def cp_shared_kv_materialize_nvtx_enabled() -> bool: + return _MATERIALIZE_NVTX_ENABLED + + def cp_shared_kv_tai_materialize_enabled() -> bool: return envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get() @@ -647,7 +653,16 @@ def _try_tai_materialize_token_kv_page_slots_into( _log_tai_materialize_fallback( "token_range_failed", "CP shared KV tai token range materialize failed; falling back to " - "torch materialize. error=%s", + "torch materialize. cp_rank=%s cp_size=%s start_slot=%s end_slot=%s " + "slot_count=%s page_size=%s kv_shape=%s dense_shape=%s error=%s", + layout.cp_rank, + layout.cp_size, + start_slot, + end_slot, + int(slot_logical_pages_range.numel()), + page_size, + tuple(kv_cache.shape), + tuple(dense_kv_cache.shape), exc, ) return False @@ -1730,20 +1745,65 @@ def _comm_view(buffer: torch.Tensor) -> torch.Tensor: return buffer -def _all_reduce_materialized_buffer(buffer: torch.Tensor, cp_size: int) -> torch.Tensor: +@contextmanager +def _materialize_all_reduce_nvtx_range( + *, + source: str, + buffer: torch.Tensor, + cp_size: int, + layer_id: int | None = None, + cp_rank: int | None = None, + rows: tuple[int, int] | None = None, +): + if not cp_shared_kv_materialize_nvtx_enabled() or not torch.cuda.is_available(): + yield + return + + row_text = f"{rows[0]}:{rows[1]}" if rows is not None else "full" + layer_text = "none" if layer_id is None else str(layer_id) + rank_text = "none" if cp_rank is None else str(cp_rank) + torch.cuda.nvtx.range_push( + "CP_SHARED_KV:materialize_allreduce " + f"source={source} layer={layer_text} cp_rank={rank_text} " + f"cp_size={cp_size} rows={row_text} shape={tuple(buffer.shape)} " + f"dtype={buffer.dtype} numel={buffer.numel()}" + ) + try: + yield + finally: + torch.cuda.nvtx.range_pop() + + +def _all_reduce_materialized_buffer( + buffer: torch.Tensor, + cp_size: int, + *, + nvtx_source: str = "unknown", + nvtx_layer_id: int | None = None, + nvtx_cp_rank: int | None = None, + nvtx_rows: tuple[int, int] | None = None, +) -> torch.Tensor: if cp_size > 1: - comm_buffer = _comm_view(buffer) - cp_group = get_attention_cp_group() - if comm_buffer.dtype in (torch.float32, torch.float16, torch.bfloat16): - reduced = cp_group.all_reduce(comm_buffer) - if ( - reduced is not None - and reduced.untyped_storage().data_ptr() - != comm_buffer.untyped_storage().data_ptr() - ): - comm_buffer.copy_(reduced) - else: - torch.distributed.all_reduce(comm_buffer, group=cp_group.device_group) + with _materialize_all_reduce_nvtx_range( + source=nvtx_source, + buffer=buffer, + cp_size=cp_size, + layer_id=nvtx_layer_id, + cp_rank=nvtx_cp_rank, + rows=nvtx_rows, + ): + comm_buffer = _comm_view(buffer) + cp_group = get_attention_cp_group() + if comm_buffer.dtype in (torch.float32, torch.float16, torch.bfloat16): + reduced = cp_group.all_reduce(comm_buffer) + if ( + reduced is not None + and reduced.untyped_storage().data_ptr() + != comm_buffer.untyped_storage().data_ptr() + ): + comm_buffer.copy_(reduced) + else: + torch.distributed.all_reduce(comm_buffer, group=cp_group.device_group) return buffer @@ -1752,6 +1812,10 @@ def _all_reduce_materialized_buffer_range( cp_size: int, start_row: int, end_row: int, + *, + nvtx_source: str = "range", + nvtx_layer_id: int | None = None, + nvtx_cp_rank: int | None = None, ) -> torch.Tensor: if start_row < 0 or end_row < start_row or end_row > buffer.shape[0]: raise ValueError( @@ -1760,7 +1824,14 @@ def _all_reduce_materialized_buffer_range( ) if start_row == end_row: return buffer - _all_reduce_materialized_buffer(buffer[start_row:end_row], cp_size) + _all_reduce_materialized_buffer( + buffer[start_row:end_row], + cp_size, + nvtx_source=nvtx_source, + nvtx_layer_id=nvtx_layer_id, + nvtx_cp_rank=nvtx_cp_rank, + nvtx_rows=(start_row, end_row), + ) return buffer @@ -1768,6 +1839,11 @@ def _all_reduce_materialized_buffer_async( buffer: torch.Tensor, cp_size: int, stream: torch.cuda.Stream, + *, + nvtx_source: str = "async", + nvtx_layer_id: int | None = None, + nvtx_cp_rank: int | None = None, + nvtx_rows: tuple[int, int] | None = None, ) -> torch.cuda.Event | None: """Enqueue an in-place CP all-reduce on `stream`. @@ -1791,7 +1867,14 @@ def _all_reduce_materialized_buffer_async( comm_buffer = _comm_view(buffer) try: - with pynccl_comm.change_state(enable=True, stream=stream): + with _materialize_all_reduce_nvtx_range( + source=nvtx_source, + buffer=buffer, + cp_size=cp_size, + layer_id=nvtx_layer_id, + cp_rank=nvtx_cp_rank, + rows=nvtx_rows, + ), pynccl_comm.change_state(enable=True, stream=stream): pynccl_comm.all_reduce(comm_buffer, stream=stream) event.record(stream) except Exception as exc: @@ -1814,6 +1897,8 @@ def materialize_shared_token_kv_buffer( remap_logical_locs: torch.Tensor | None = None, remap_logical_pages: torch.Tensor | None = None, slot_remap: SharedTokenKVSlotRemap | None = None, + nvtx_source: str = "mla.full_materialize", + nvtx_layer_id: int | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: _debug_assert_no_tensor_values_below( logical_locs, @@ -1944,7 +2029,13 @@ def materialize_shared_token_kv_buffer( ), ) - dense_kv_cache = _all_reduce_materialized_buffer(dense_kv_cache, layout.cp_size) + dense_kv_cache = _all_reduce_materialized_buffer( + dense_kv_cache, + layout.cp_size, + nvtx_source=nvtx_source, + nvtx_layer_id=nvtx_layer_id, + nvtx_cp_rank=layout.cp_rank, + ) if cp_shared_kv_debug_enabled(): cp_shared_kv_debug_log( @@ -1964,6 +2055,9 @@ def materialize_shared_paged_buffer( logical_pages: torch.Tensor, layout: CpSharedKVLayout, slot_remap: SharedPagedBufferSlotRemap | None = None, + *, + nvtx_source: str = "index.full_materialize", + nvtx_layer_id: int | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: _debug_assert_no_negative_tensor_values( logical_pages, @@ -2026,7 +2120,11 @@ def materialize_shared_paged_buffer( ) dense_page_buffer = _all_reduce_materialized_buffer( - dense_page_buffer, layout.cp_size + dense_page_buffer, + layout.cp_size, + nvtx_source=nvtx_source, + nvtx_layer_id=nvtx_layer_id, + nvtx_cp_rank=layout.cp_rank, ) if cp_shared_kv_debug_enabled(): diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index b6fff7700..3a33fff0d 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -18,6 +18,9 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( cp_shared_kv_debug_enabled, cp_shared_kv_debug_log, cp_shared_kv_current_reuse_enabled, + cp_shared_kv_mla_prefetch_log, + cp_shared_kv_mla_prefetch_log_enabled, + cp_shared_kv_mla_prefetch_should_log_layer, filter_owned_logical_locs, get_or_build_shared_paged_buffer_slot_remap, is_current_only_extend_batch, @@ -358,6 +361,8 @@ class Indexer(MultiPlatformOp): logical_pages=logical_page_table, layout=layout, slot_remap=slot_remap, + nvtx_source="index.full_materialize", + nvtx_layer_id=layer_id, ) if cp_shared_kv_debug_enabled(): cp_shared_kv_debug_log( @@ -375,13 +380,32 @@ class Indexer(MultiPlatformOp): forward_batch: ForwardBatch, layer_id: int, ) -> None: + next_layer_id = layer_id + 1 index_prefetcher = getattr( forward_batch, "cp_shared_kv_index_prefetcher", None ) + if ( + cp_shared_kv_mla_prefetch_log_enabled() + and cp_shared_kv_mla_prefetch_should_log_layer(next_layer_id) + ): + layout = getattr(forward_batch, "cp_shared_kv_layout", None) + seq_lens_cpu = getattr(forward_batch, "seq_lens_cpu", None) + cp_shared_kv_mla_prefetch_log( + "index_start_request cp_rank=%s cp_size=%s layer=%s " + "next_layer=%s has_index=%s uses_cp_shared_kv=%s " + "seq_lens_cpu_len=%s", + getattr(layout, "cp_rank", None), + getattr(layout, "cp_size", None), + layer_id, + next_layer_id, + index_prefetcher is not None, + getattr(forward_batch, "uses_cp_shared_kv", None), + len(seq_lens_cpu) if seq_lens_cpu is not None else None, + ) if index_prefetcher is None: return index_prefetcher.start_next_layer_prefix( - next_layer_id=layer_id + 1, + next_layer_id=next_layer_id, token_to_kv_pool=forward_batch.token_to_kv_pool, ) @@ -1666,6 +1690,25 @@ class Indexer(MultiPlatformOp): if _is_cuda or _is_hip: assert forward_batch.seq_lens_cpu is not None if len(forward_batch.seq_lens_cpu) == 0: + if ( + cp_shared_kv_mla_prefetch_log_enabled() + and cp_shared_kv_mla_prefetch_should_log_layer(layer_id) + ): + layout = getattr(forward_batch, "cp_shared_kv_layout", None) + index_prefetcher = getattr( + forward_batch, "cp_shared_kv_index_prefetcher", None + ) + cp_shared_kv_mla_prefetch_log( + "index_forward_early_return_empty_seq_lens cp_rank=%s " + "cp_size=%s layer=%s has_index=%s uses_cp_shared_kv=%s " + "x_meta_shape=%s", + getattr(layout, "cp_rank", None), + getattr(layout, "cp_size", None), + layer_id, + index_prefetcher is not None, + getattr(forward_batch, "uses_cp_shared_kv", None), + tuple(x_meta.shape), + ) # this seems b/c max-pad, no worries? # if x.shape[0] != 0: # print( diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index 696253ae9..76be1bbda 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -53,7 +53,7 @@ from sglang.srt.layers.attention.utils import ( mla_quantize_and_rope_for_fp8, seqlens_expand_triton, ) -from sglang.srt.layers.dp_attention import get_attention_tp_size +from sglang.srt.layers.dp_attention import get_attention_cp_group, get_attention_tp_size from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.utils import is_cuda, is_hip @@ -905,6 +905,47 @@ class NativeSparseAttnBackend( stream=shared_prefetch_stream, ) ) + if cp_shared_kv_mla_prefetch_log_enabled(): + layout = getattr(forward_batch, "cp_shared_kv_layout", None) + prefix_lens_cpu = getattr(forward_batch, "extend_prefix_lens_cpu", None) + extend_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None) + seq_lens_cpu = getattr(forward_batch, "seq_lens_cpu", None) + try: + pynccl_available = ( + getattr(get_attention_cp_group(), "pynccl_comm", None) is not None + ) + except Exception: + pynccl_available = False + cp_shared_kv_mla_prefetch_log( + "create_result cp_rank=%s cp_size=%s uses_cp_shared_kv=%s " + "forward_mode=%s topk_transform=%s batch_size=%s has_mla=%s " + "has_index=%s prefix_lens=%s extend_lens=%s seq_lens=%s " + "real_pages=%s page_table_shape=%s pynccl=%s", + getattr(layout, "cp_rank", None), + getattr(layout, "cp_size", None), + getattr(forward_batch, "uses_cp_shared_kv", None), + getattr(forward_batch, "forward_mode", None), + getattr(topk_transform_method, "name", topk_transform_method), + getattr(forward_batch, "batch_size", None), + mla_prefetcher is not None, + forward_batch.cp_shared_kv_index_prefetcher is not None, + [int(x) for x in prefix_lens_cpu] + if prefix_lens_cpu is not None + else None, + [int(x) for x in extend_lens_cpu] + if extend_lens_cpu is not None + else None, + [int(x) for x in seq_lens_cpu.tolist()] + if seq_lens_cpu is not None + else None, + int(metadata.real_page_table.numel()) + if metadata.real_page_table is not None + else None, + tuple(metadata.page_table_1.shape) + if metadata.page_table_1 is not None + else None, + pynccl_available, + ) def _cal_indexer_k_start_end( self, @@ -1637,6 +1678,45 @@ class NativeSparseAttnBackend( ) ) + if ( + cp_shared_kv_mla_prefetch_log_enabled() + and cp_shared_kv_mla_prefetch_should_log_layer(layer.layer_id) + ): + page_table_for_log = locals().get("page_table_1", None) + layout = getattr(forward_batch, "cp_shared_kv_layout", None) + mla_prefetcher = getattr( + forward_batch, "cp_shared_kv_mla_prefetcher", None + ) + index_prefetcher = getattr( + forward_batch, "cp_shared_kv_index_prefetcher", None + ) + prefix_lens_cpu = getattr(forward_batch, "extend_prefix_lens_cpu", None) + extend_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None) + cp_shared_kv_mla_prefetch_log( + "forward_probe cp_rank=%s cp_size=%s layer=%s " + "uses_cp_shared_kv=%s topk_transform=%s has_mla=%s " + "has_index=%s prefix_lens=%s extend_lens=%s q_rows=%s " + "page_table_shape=%s topk_shape=%s", + getattr(layout, "cp_rank", None), + getattr(layout, "cp_size", None), + layer.layer_id, + getattr(forward_batch, "uses_cp_shared_kv", None), + getattr(topk_transform_method, "name", topk_transform_method), + mla_prefetcher is not None, + index_prefetcher is not None, + [int(x) for x in prefix_lens_cpu] + if prefix_lens_cpu is not None + else None, + [int(x) for x in extend_lens_cpu] + if extend_lens_cpu is not None + else None, + int(q_nope.shape[0]), + tuple(page_table_for_log.shape) + if page_table_for_log is not None + else None, + tuple(topk_indices.shape) if topk_indices is not None else None, + ) + if ( forward_batch.uses_cp_shared_kv and topk_transform_method == TopkTransformMethod.PAGED @@ -1755,6 +1835,8 @@ class NativeSparseAttnBackend( layout=forward_batch.cp_shared_kv_layout, page_size=forward_batch.token_to_kv_pool.page_size, slot_remap=slot_remap, + nvtx_source="mla.full_materialize", + nvtx_layer_id=layer.layer_id, ) if mla_prefetcher is not None: @@ -1765,6 +1847,12 @@ class NativeSparseAttnBackend( else: mla_prefetcher = None + index_prefetcher = getattr( + forward_batch, "cp_shared_kv_index_prefetcher", None + ) + if index_prefetcher is not None: + index_prefetcher.launch_pending_reduce() + try: if nsa_impl == "tilelang": if q_rope is not None: @@ -1794,6 +1882,8 @@ class NativeSparseAttnBackend( logical_locs=page_table_1_flattened, layout=forward_batch.cp_shared_kv_layout, page_size=forward_batch.token_to_kv_pool.page_size, + nvtx_source="mla.ragged_full_materialize", + nvtx_layer_id=layer.layer_id, ) ) kv_cache = dequantize_k_cache_paged( @@ -1853,6 +1943,7 @@ class NativeSparseAttnBackend( ) finally: if mla_prefetcher is not None: + mla_prefetcher.launch_pending_reduce() mla_prefetcher.wait_attention_window() index_prefetcher = getattr( forward_batch, "cp_shared_kv_index_prefetcher", None diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 2c06599c3..b5475f982 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -94,6 +94,42 @@ class DeepseekMLAForwardMixin: get_global_server_args().flashinfer_mla_disable_ragged ) + def _maybe_start_cp_shared_next_layer_prefetch( + self: DeepseekV2AttentionMLA, + forward_batch: ForwardBatch, + ) -> None: + """Launch next-layer CP shared-KV prefix prefetch before rank-skewed work. + + The original Phase8 hook launched prefetch from the indexer/backend after + current-layer MQA/materialize work. Those paths can be rank-skewed in + in-seq CP, so rank 0 may enqueue the next-layer collective early while + other ranks reach it too late to overlap. Starting from layer prepare + keeps the launch point before the current layer's per-rank MQA/topk and + materialize imbalance. Existing later hooks remain as fallbacks and + become no-ops via the prefetchers' already-started guard. + """ + + token_to_kv_pool = getattr(forward_batch, "token_to_kv_pool", None) + if token_to_kv_pool is None: + return + + next_layer_id = int(self.layer_id) + 1 + index_prefetcher = getattr( + forward_batch, "cp_shared_kv_index_prefetcher", None + ) + if index_prefetcher is not None: + index_prefetcher.start_next_layer_prefix( + next_layer_id=next_layer_id, + token_to_kv_pool=token_to_kv_pool, + ) + + mla_prefetcher = getattr(forward_batch, "cp_shared_kv_mla_prefetcher", None) + if mla_prefetcher is not None: + mla_prefetcher.start_next_layer_prefix( + next_layer_id=next_layer_id, + token_to_kv_pool=token_to_kv_pool, + ) + def forward_absorb_prepare( self: DeepseekV2AttentionMLA, positions: torch.Tensor, @@ -104,6 +140,8 @@ class DeepseekMLAForwardMixin: ): from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode + self._maybe_start_cp_shared_next_layer_prefetch(forward_batch) + q_lora = None topk_indices = None if self.q_lora_rank is not None: diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py index a6ec657d3..4c7fb66ec 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py @@ -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,