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