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-08 17:39:24 +08:00
parent aa27a444f6
commit 3fc7a5c18c
7 changed files with 850 additions and 107 deletions

View File

@@ -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)

View File

@@ -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:

View File

@@ -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():

View File

@@ -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(

View File

@@ -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

View File

@@ -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: