Keep CP shared-KV prefetch warnings actionable
Expected no-prefetch paths were polluting production logs: no cache prefix, tiny/first-layer windows, and FP8 RAGGED top-k were being reported as fallback warnings. The prefetch contract now treats zero-prefix and first-layer misses as normal skips, while preserving warnings for non-zero misaligned prefixes and real consume misses after the first layer. The same change keeps RAGGED cache-hit prefetch eligible and records the CE/IPM prefetch contract in the plan doc. Constraint: FP8 sparse prefill uses RAGGED top-k, but CP shared-KV prefix materialization is still page-slot based Constraint: Layer 0 has no previous attention-window hook that can have prefetched the layer Rejected: Warn whenever a prefetcher is absent | no-cache and too-short requests are expected synchronous paths and make logs unusable Confidence: high Scope-risk: moderate Directive: Keep CP_SHARED_KV_FALLBACK warnings for unexpected contract failures only; use debug logs for expected skip paths Tested: Local py_compile for cp_shared_kv_prefetch.py, nsa_indexer.py, nsa_backend.py Tested: Remote cjy-glm5-new targeted regression: 3 passed, 21 warnings Tested: Remote cjy-glm5-new full test_cp_shared_kv_runtime.py: 156 passed, 21 warnings, 2 subtests passed Not-tested: New ETE run after prefill restart to confirm log volume reduction in production traffic (cherry picked from commit e08e321e5929fdbb30102ec0b19c6ff0ecac7e7e)
This commit is contained in:
@@ -8,8 +8,6 @@ from typing import Any, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
_all_reduce_materialized_buffer_async,
|
||||
_all_reduce_materialized_buffer_range,
|
||||
build_batch_current_slot_spans,
|
||||
build_batch_prefix_slot_spans,
|
||||
cp_shared_kv_debug_enabled,
|
||||
@@ -34,13 +32,14 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
slot_range_to_token_slice,
|
||||
_raise_tai_ipc_materialize_required,
|
||||
_should_fail_fast_tai_ipc_materialize,
|
||||
_try_tai_ipc_ce_materialize_paged_buffer_page_slot_spans_into,
|
||||
_try_tai_ipc_ce_materialize_token_kv_page_slot_spans_into,
|
||||
_try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into,
|
||||
_try_tai_ipc_materialize_current_token_kv_page_slot_spans_into,
|
||||
_try_tai_ipc_materialize_paged_buffer_page_slot_spans_into,
|
||||
_try_tai_ipc_materialize_token_kv_page_slot_spans_into,
|
||||
)
|
||||
from sglang.srt.layers.attention.nsa.utils import is_nsa_prefill_cp_in_seq_split
|
||||
from sglang.srt.layers.dp_attention import get_attention_cp_group
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -215,32 +214,6 @@ def _materialize_local_paged_buffer_page_slot_spans_into(
|
||||
)
|
||||
|
||||
|
||||
def _all_reduce_materialized_buffer_ranges_async(
|
||||
*,
|
||||
dense_kv_cache: torch.Tensor,
|
||||
row_slices: list[slice],
|
||||
cp_size: int,
|
||||
stream: torch.cuda.Stream,
|
||||
nvtx_source: str,
|
||||
nvtx_layer_id: int,
|
||||
nvtx_cp_rank: int,
|
||||
) -> torch.cuda.Event | None:
|
||||
event: torch.cuda.Event | None = None
|
||||
for rows in row_slices:
|
||||
event = _all_reduce_materialized_buffer_async(
|
||||
dense_kv_cache[rows],
|
||||
cp_size=cp_size,
|
||||
stream=stream,
|
||||
nvtx_source=nvtx_source,
|
||||
nvtx_layer_id=nvtx_layer_id,
|
||||
nvtx_cp_rank=nvtx_cp_rank,
|
||||
nvtx_rows=(rows.start, rows.stop),
|
||||
)
|
||||
if event is None:
|
||||
return None
|
||||
return event
|
||||
|
||||
|
||||
def _record_event_on_stream(stream: torch.cuda.Stream) -> torch.cuda.Event:
|
||||
event = torch.cuda.Event()
|
||||
event.record(stream)
|
||||
@@ -277,7 +250,7 @@ class _PrefetchCpuTiming:
|
||||
start_max_ms: float = 0.0
|
||||
get_total_ms: float = 0.0
|
||||
materialize_total_ms: float = 0.0
|
||||
reduce_enqueue_total_ms: float = 0.0
|
||||
materialize_enqueue_total_ms: float = 0.0
|
||||
consume_count: int = 0
|
||||
consume_total_ms: float = 0.0
|
||||
consume_wait_total_ms: float = 0.0
|
||||
@@ -290,7 +263,7 @@ class _PrefetchCpuTiming:
|
||||
total_ms: float,
|
||||
get_ms: float,
|
||||
materialize_ms: float,
|
||||
reduce_enqueue_ms: float,
|
||||
materialize_enqueue_ms: float,
|
||||
) -> None:
|
||||
if total_ms < 0.0:
|
||||
return
|
||||
@@ -299,7 +272,7 @@ class _PrefetchCpuTiming:
|
||||
self.start_max_ms = max(self.start_max_ms, total_ms)
|
||||
self.get_total_ms += max(get_ms, 0.0)
|
||||
self.materialize_total_ms += max(materialize_ms, 0.0)
|
||||
self.reduce_enqueue_total_ms += max(reduce_enqueue_ms, 0.0)
|
||||
self.materialize_enqueue_total_ms += max(materialize_enqueue_ms, 0.0)
|
||||
|
||||
def record_consume(
|
||||
self,
|
||||
@@ -328,7 +301,7 @@ def _log_prefetch_cpu_start(
|
||||
total_ms: float,
|
||||
get_ms: float,
|
||||
materialize_ms: float,
|
||||
reduce_enqueue_ms: float,
|
||||
materialize_enqueue_ms: float,
|
||||
prefix_pages: int,
|
||||
total_slots: int,
|
||||
dense_units: int,
|
||||
@@ -338,14 +311,14 @@ def _log_prefetch_cpu_start(
|
||||
if cp_shared_kv_mla_prefetch_should_log_layer(layer_id):
|
||||
log_fn(
|
||||
"cpu_timing path=%s stage=start layer=%s total_ms=%.3f "
|
||||
"get_ms=%.3f materialize_ms=%.3f reduce_enqueue_ms=%.3f "
|
||||
"get_ms=%.3f materialize_ms=%.3f materialize_enqueue_ms=%.3f "
|
||||
"prefix_pages=%s total_slots=%s dense_units=%s",
|
||||
path,
|
||||
layer_id,
|
||||
total_ms,
|
||||
get_ms,
|
||||
materialize_ms,
|
||||
reduce_enqueue_ms,
|
||||
materialize_enqueue_ms,
|
||||
prefix_pages,
|
||||
total_slots,
|
||||
dense_units,
|
||||
@@ -354,14 +327,14 @@ def _log_prefetch_cpu_start(
|
||||
starts = timing.start_count
|
||||
log_fn(
|
||||
"cpu_summary path=%s starts=%s avg_start_ms=%.3f max_start_ms=%.3f "
|
||||
"avg_get_ms=%.3f avg_materialize_ms=%.3f avg_reduce_enqueue_ms=%.3f",
|
||||
"avg_get_ms=%.3f avg_materialize_ms=%.3f avg_materialize_enqueue_ms=%.3f",
|
||||
path,
|
||||
starts,
|
||||
timing.start_total_ms / starts,
|
||||
timing.start_max_ms,
|
||||
timing.get_total_ms / starts,
|
||||
timing.materialize_total_ms / starts,
|
||||
timing.reduce_enqueue_total_ms / starts,
|
||||
timing.materialize_enqueue_total_ms / starts,
|
||||
)
|
||||
|
||||
|
||||
@@ -453,6 +426,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None = None,
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None = None,
|
||||
dense_num_pages: int,
|
||||
slot_remap: Any | None = None,
|
||||
owned_prefix_pages: int = -1,
|
||||
owned_total_pages: int = -1,
|
||||
stream: Optional[torch.cuda.Stream] = None,
|
||||
@@ -465,6 +439,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
self.slot_sorted_logical_pages_by_row = slot_sorted_logical_pages_by_row
|
||||
self.slot_sorted_dense_pages_by_row = slot_sorted_dense_pages_by_row
|
||||
self.dense_num_pages = dense_num_pages
|
||||
self.slot_remap = slot_remap
|
||||
self.owned_prefix_pages = owned_prefix_pages
|
||||
self.owned_total_pages = owned_total_pages
|
||||
self.total_slots = int(slot_logical_pages.numel())
|
||||
@@ -519,8 +494,11 @@ class CpSharedKVMlaPrefetcher:
|
||||
_prefetch_log("create_skip reason=not_in_seq_split")
|
||||
return None
|
||||
if not topk_transform_is_paged:
|
||||
_prefetch_log("create_skip reason=not_paged_topk")
|
||||
return None
|
||||
# FP8 sparse prefill uses the RAGGED top-k transform by design.
|
||||
# MLA prefix prefetch materializes page-slot KV rows before the
|
||||
# final top-k offset semantics are consumed, so the RAGGED transform
|
||||
# is still a valid prefetch target.
|
||||
_prefetch_log("create_continue reason=ragged_topk")
|
||||
batch_size = int(getattr(forward_batch, "batch_size", 0))
|
||||
if batch_size <= 0:
|
||||
_prefetch_log(
|
||||
@@ -553,11 +531,17 @@ class CpSharedKVMlaPrefetcher:
|
||||
for prefix_len in prefix_lens
|
||||
if prefix_len < 0 or prefix_len % page_size != 0
|
||||
]
|
||||
if bad_prefix_lens or not any(prefix_len > 0 for prefix_len in prefix_lens):
|
||||
if not any(prefix_len > 0 for prefix_len in prefix_lens):
|
||||
_prefetch_log(
|
||||
"create_skip reason=no_cache_prefix prefix_lens=%s page_size=%s",
|
||||
prefix_lens,
|
||||
page_size,
|
||||
)
|
||||
return None
|
||||
if bad_prefix_lens:
|
||||
_mla_prefetch_fallback_log(
|
||||
"prefix_not_page_aligned",
|
||||
"prefix length is zero or not page-aligned. "
|
||||
"prefix_lens=%s page_size=%s",
|
||||
"non-zero prefix length is not page-aligned. prefix_lens=%s page_size=%s",
|
||||
prefix_lens,
|
||||
page_size,
|
||||
)
|
||||
@@ -654,15 +638,6 @@ class CpSharedKVMlaPrefetcher:
|
||||
)
|
||||
return None
|
||||
|
||||
cp_group = get_attention_cp_group()
|
||||
if getattr(cp_group, "pynccl_comm", None) is None and layout.cp_size > 1:
|
||||
_prefetch_log(
|
||||
"create_skip reason=missing_pynccl cp_rank=%s cp_size=%s",
|
||||
layout.cp_rank,
|
||||
layout.cp_size,
|
||||
)
|
||||
return None
|
||||
|
||||
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
|
||||
create_cpu = _cpu_timing_start()
|
||||
get_cpu = _cpu_timing_start()
|
||||
@@ -738,6 +713,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
slot_sorted_logical_pages_by_row=remap.slot_sorted_logical_pages_by_row,
|
||||
slot_sorted_dense_pages_by_row=remap.slot_sorted_dense_pages_by_row,
|
||||
dense_num_pages=remap.dense_num_pages,
|
||||
slot_remap=remap,
|
||||
owned_prefix_pages=owned_prefix_pages,
|
||||
owned_total_pages=owned_total_pages,
|
||||
stream=prefetch_stream,
|
||||
@@ -779,12 +755,12 @@ class CpSharedKVMlaPrefetcher:
|
||||
self._log_layer(layer_id, "consume_miss layer=%s", layer_id)
|
||||
return None
|
||||
if handle.event is None:
|
||||
self.launch_pending_reduce()
|
||||
self.finalize_pending_materialize()
|
||||
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",
|
||||
"consume_miss reason=prefix_materialize_not_ready layer=%s",
|
||||
layer_id,
|
||||
)
|
||||
return None
|
||||
@@ -830,7 +806,11 @@ class CpSharedKVMlaPrefetcher:
|
||||
spans=suffix_spans,
|
||||
)
|
||||
)
|
||||
if not materialized_suffix_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_kv_cache):
|
||||
if (
|
||||
not materialized_suffix_by_ipc
|
||||
and suffix_spans
|
||||
and self.layout.cp_size > 1
|
||||
):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"mla_prefetch_suffix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
@@ -847,19 +827,6 @@ class CpSharedKVMlaPrefetcher:
|
||||
page_size=self.page_size,
|
||||
spans=suffix_spans,
|
||||
)
|
||||
if not materialized_suffix_by_ipc:
|
||||
for suffix_rows in _slot_spans_to_token_slices(
|
||||
self.page_size, suffix_spans
|
||||
):
|
||||
_all_reduce_materialized_buffer_range(
|
||||
dense_kv_cache,
|
||||
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 suffix_spans=%s",
|
||||
@@ -929,7 +896,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
"""Consume the prefetched prefix and append current-layer KV rows.
|
||||
|
||||
This is the partial-current-reuse variant of :meth:`consume`: prefix
|
||||
pages are already materialized/reduced by the prefetch stream, while
|
||||
pages are already materialized by the prefetch stream, while
|
||||
current/suffix pages are not copied from the shared pool. Current locs
|
||||
in ``logical_locs`` are remapped to the appended current KV rows.
|
||||
"""
|
||||
@@ -954,12 +921,12 @@ class CpSharedKVMlaPrefetcher:
|
||||
self._log_layer(layer_id, "consume_prefix_current_miss layer=%s", layer_id)
|
||||
return None
|
||||
if handle.event is None:
|
||||
self.launch_pending_reduce()
|
||||
self.finalize_pending_materialize()
|
||||
handle = self.handles.get(layer_id)
|
||||
if handle is None or handle.event is None:
|
||||
self._log_layer(
|
||||
layer_id,
|
||||
"consume_prefix_current_miss reason=prefix_reduce_not_ready layer=%s",
|
||||
"consume_prefix_current_miss reason=prefix_materialize_not_ready layer=%s",
|
||||
layer_id,
|
||||
)
|
||||
return None
|
||||
@@ -1022,7 +989,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
if (
|
||||
not current_materialized_by_ipc
|
||||
and self.current_slot_spans
|
||||
and _should_fail_fast_tai_ipc_materialize(mixed_kv_cache)
|
||||
and self.layout.cp_size > 1
|
||||
):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"mla_prefetch_current_ipc_unavailable",
|
||||
@@ -1031,19 +998,6 @@ class CpSharedKVMlaPrefetcher:
|
||||
spans=self.current_slot_spans,
|
||||
dense_shape=tuple(mixed_kv_cache.shape),
|
||||
)
|
||||
if not current_materialized_by_ipc:
|
||||
for current_rows in _slot_spans_to_token_slices(
|
||||
self.page_size, self.current_slot_spans
|
||||
):
|
||||
_all_reduce_materialized_buffer_range(
|
||||
mixed_kv_cache,
|
||||
self.layout.cp_size,
|
||||
current_rows.start,
|
||||
current_rows.stop,
|
||||
nvtx_source="mla.prefetch_current",
|
||||
nvtx_layer_id=layer_id,
|
||||
nvtx_cp_rank=self.layout.cp_rank,
|
||||
)
|
||||
remap_ms = _cpu_timing_ms(remap_cpu)
|
||||
total_ms = _cpu_timing_ms(consume_cpu)
|
||||
self._log_layer(
|
||||
@@ -1106,25 +1060,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
start_cpu = _cpu_timing_start()
|
||||
current_stream = torch.cuda.current_stream()
|
||||
get_cpu = _cpu_timing_start()
|
||||
try:
|
||||
kv_cache = _prefetch_pool_get_key_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=next_layer_id,
|
||||
stream=current_stream,
|
||||
path="mla",
|
||||
)
|
||||
get_ms = _cpu_timing_ms(get_cpu)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to get next-layer KV cache for CP shared KV MLA prefetch."
|
||||
)
|
||||
self.disabled = True
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start_disable reason=get_kv_failed next_layer=%s",
|
||||
next_layer_id,
|
||||
)
|
||||
return
|
||||
self.stream.wait_stream(current_stream)
|
||||
|
||||
try:
|
||||
prefix_row_spans = _slot_spans_to_token_slices(
|
||||
@@ -1132,64 +1068,81 @@ class CpSharedKVMlaPrefetcher:
|
||||
)
|
||||
prefix_rows = prefix_row_spans[0] if prefix_row_spans else slice(0, 0)
|
||||
materialize_cpu = _cpu_timing_start()
|
||||
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 prefix_slot_spans=%s "
|
||||
"dense_rows=%s",
|
||||
next_layer_id,
|
||||
self.prefix_slot_spans,
|
||||
int(dense_kv_cache.shape[0]),
|
||||
)
|
||||
materialized_by_ipc = _try_tai_ipc_materialize_token_kv_page_slot_spans_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,
|
||||
spans=self.prefix_slot_spans,
|
||||
)
|
||||
if not materialized_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_kv_cache):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"mla_prefetch_prefix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
spans=self.prefix_slot_spans,
|
||||
dense_shape=tuple(dense_kv_cache.shape),
|
||||
with torch.cuda.stream(self.stream):
|
||||
kv_cache = _prefetch_pool_get_key_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=next_layer_id,
|
||||
stream=self.stream,
|
||||
path="mla",
|
||||
)
|
||||
if not materialized_by_ipc:
|
||||
_materialize_local_token_kv_page_slot_spans_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,
|
||||
spans=self.prefix_slot_spans,
|
||||
get_ms = _cpu_timing_ms(get_cpu)
|
||||
dense_kv_cache = kv_cache.new_zeros(
|
||||
(self.dense_num_pages * self.page_size, *kv_cache.shape[1:])
|
||||
)
|
||||
materialize_ms = _cpu_timing_ms(materialize_cpu)
|
||||
reduce_cpu = _cpu_timing_start()
|
||||
if materialized_by_ipc:
|
||||
event = _record_event_on_stream(current_stream)
|
||||
else:
|
||||
self.stream.wait_stream(current_stream)
|
||||
with torch.cuda.stream(self.stream):
|
||||
event = _all_reduce_materialized_buffer_ranges_async(
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start_prefix_begin next_layer=%s prefix_slot_spans=%s "
|
||||
"dense_rows=%s stream=prefetch",
|
||||
next_layer_id,
|
||||
self.prefix_slot_spans,
|
||||
int(dense_kv_cache.shape[0]),
|
||||
)
|
||||
materialized_by_ce = (
|
||||
_try_tai_ipc_ce_materialize_token_kv_page_slot_spans_into(
|
||||
kv_cache=kv_cache,
|
||||
dense_kv_cache=dense_kv_cache,
|
||||
row_slices=prefix_row_spans,
|
||||
cp_size=self.layout.cp_size,
|
||||
stream=self.stream,
|
||||
nvtx_source="mla.prefetch_prefix",
|
||||
nvtx_layer_id=next_layer_id,
|
||||
nvtx_cp_rank=self.layout.cp_rank,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
page_size=self.page_size,
|
||||
spans=self.prefix_slot_spans,
|
||||
slot_remap=self.slot_remap,
|
||||
)
|
||||
reduce_enqueue_ms = _cpu_timing_ms(reduce_cpu)
|
||||
)
|
||||
event = _record_event_on_stream(self.stream) if materialized_by_ce else None
|
||||
|
||||
materialized_by_ipc = materialized_by_ce
|
||||
if not materialized_by_ce:
|
||||
current_stream.wait_stream(self.stream)
|
||||
materialized_by_ipc = (
|
||||
_try_tai_ipc_materialize_token_kv_page_slot_spans_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,
|
||||
spans=self.prefix_slot_spans,
|
||||
slot_remap=self.slot_remap,
|
||||
)
|
||||
)
|
||||
if (
|
||||
not materialized_by_ipc
|
||||
and self.prefix_slot_spans
|
||||
and self.layout.cp_size > 1
|
||||
):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"mla_prefetch_prefix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
spans=self.prefix_slot_spans,
|
||||
dense_shape=tuple(dense_kv_cache.shape),
|
||||
)
|
||||
if not materialized_by_ipc:
|
||||
_materialize_local_token_kv_page_slot_spans_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,
|
||||
spans=self.prefix_slot_spans,
|
||||
)
|
||||
event = _record_event_on_stream(current_stream)
|
||||
materialize_ms = _cpu_timing_ms(materialize_cpu)
|
||||
materialize_enqueue_ms = 0.0
|
||||
if event is None:
|
||||
self.disabled = True
|
||||
logger.warning(
|
||||
"[CP_SHARED_KV_FALLBACK][mla_prefetch] "
|
||||
"reason=async_reduce_unavailable layer_id=%s cp_rank=%s "
|
||||
"reason=materialize_event_unavailable layer_id=%s cp_rank=%s "
|
||||
"cp_size=%s prefix_pages=%s total_slots=%s",
|
||||
next_layer_id,
|
||||
self.layout.cp_rank,
|
||||
@@ -1199,7 +1152,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start_disable reason=async_reduce_unavailable next_layer=%s",
|
||||
"start_disable reason=materialize_event_unavailable next_layer=%s",
|
||||
next_layer_id,
|
||||
)
|
||||
return
|
||||
@@ -1208,7 +1161,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
total_ms=total_ms,
|
||||
get_ms=get_ms,
|
||||
materialize_ms=materialize_ms,
|
||||
reduce_enqueue_ms=reduce_enqueue_ms,
|
||||
materialize_enqueue_ms=materialize_enqueue_ms,
|
||||
)
|
||||
_log_prefetch_cpu_start(
|
||||
log_fn=self._log,
|
||||
@@ -1219,7 +1172,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
total_ms=total_ms,
|
||||
get_ms=get_ms,
|
||||
materialize_ms=materialize_ms,
|
||||
reduce_enqueue_ms=reduce_enqueue_ms,
|
||||
materialize_enqueue_ms=materialize_enqueue_ms,
|
||||
prefix_pages=self.prefix_pages,
|
||||
total_slots=self.total_slots,
|
||||
dense_units=int(dense_kv_cache.shape[0]),
|
||||
@@ -1246,74 +1199,44 @@ class CpSharedKVMlaPrefetcher:
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start next_layer=%s prefix_pages=%s prefix_rows=%s dense_rows=%s "
|
||||
"reduce_enqueued=True",
|
||||
"materialize_enqueued=True",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
prefix_rows.stop - prefix_rows.start,
|
||||
int(dense_kv_cache.shape[0]),
|
||||
)
|
||||
|
||||
def launch_pending_reduce(self) -> None:
|
||||
def finalize_pending_materialize(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",
|
||||
"start_prefix_materialize_skip reason=already_enqueued next_layer=%s",
|
||||
handle.layer_id,
|
||||
)
|
||||
return
|
||||
|
||||
prefix_row_spans = list(handle.prefix_row_spans or (handle.prefix_rows,))
|
||||
try:
|
||||
if _should_fail_fast_tai_ipc_materialize(handle.dense_kv_cache):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"mla_prefetch_deferred_prefix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
dense_shape=tuple(handle.dense_kv_cache.shape),
|
||||
layer_id=handle.layer_id,
|
||||
)
|
||||
current_stream = torch.cuda.current_stream()
|
||||
self.stream.wait_stream(current_stream)
|
||||
with torch.cuda.stream(self.stream):
|
||||
event = _all_reduce_materialized_buffer_ranges_async(
|
||||
dense_kv_cache=handle.dense_kv_cache,
|
||||
row_slices=prefix_row_spans,
|
||||
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,
|
||||
)
|
||||
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 row_spans=%s",
|
||||
handle.layer_id,
|
||||
[(rows.start, rows.stop) for rows in prefix_row_spans],
|
||||
)
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"mla_prefetch_deferred_prefix_materialize_missing",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
dense_shape=tuple(handle.dense_kv_cache.shape),
|
||||
layer_id=handle.layer_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to launch CP shared KV MLA prefix prefetch reduce.")
|
||||
logger.exception(
|
||||
"CP shared KV MLA prefix prefetch has no IPC/CE completion event."
|
||||
)
|
||||
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",
|
||||
"start_disable reason=materialize_event_missing next_layer=%s",
|
||||
handle.layer_id,
|
||||
)
|
||||
return
|
||||
@@ -1366,6 +1289,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
slot_sorted_logical_pages_by_row: torch.Tensor | None = None,
|
||||
slot_sorted_dense_pages_by_row: torch.Tensor | None = None,
|
||||
dense_num_pages: int,
|
||||
slot_remap: Any | None = None,
|
||||
owned_prefix_pages: int = -1,
|
||||
owned_total_pages: int = -1,
|
||||
stream: Optional[torch.cuda.Stream] = None,
|
||||
@@ -1377,6 +1301,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
self.slot_sorted_logical_pages_by_row = slot_sorted_logical_pages_by_row
|
||||
self.slot_sorted_dense_pages_by_row = slot_sorted_dense_pages_by_row
|
||||
self.dense_num_pages = dense_num_pages
|
||||
self.slot_remap = slot_remap
|
||||
self.owned_prefix_pages = owned_prefix_pages
|
||||
self.owned_total_pages = owned_total_pages
|
||||
self.total_slots = int(slot_logical_pages.numel())
|
||||
@@ -1448,11 +1373,11 @@ class CpSharedKVIndexPrefetcher:
|
||||
)
|
||||
return None
|
||||
if not topk_transform_is_paged:
|
||||
_index_prefetch_fallback_log(
|
||||
"not_paged_topk",
|
||||
"topk transform is not PAGED.",
|
||||
)
|
||||
return None
|
||||
# FP8 prefill uses the RAGGED top-k transform by design. Index
|
||||
# prefix materialization is still page-slot based and remains a
|
||||
# valid prefetch target; the transform method only affects the
|
||||
# final top-k offset semantics after MQA logits are computed.
|
||||
_prefetch_log("index_create_continue reason=ragged_topk")
|
||||
batch_size = int(getattr(forward_batch, "batch_size", 0))
|
||||
if batch_size <= 0:
|
||||
_index_prefetch_fallback_log(
|
||||
@@ -1501,10 +1426,17 @@ class CpSharedKVIndexPrefetcher:
|
||||
for prefix_len in prefix_lens
|
||||
if prefix_len < 0 or prefix_len % page_size != 0
|
||||
]
|
||||
if bad_prefix_lens or not any(prefix_len > 0 for prefix_len in prefix_lens):
|
||||
if not any(prefix_len > 0 for prefix_len in prefix_lens):
|
||||
_prefetch_log(
|
||||
"index_create_skip reason=no_cache_prefix prefix_lens=%s page_size=%s",
|
||||
prefix_lens,
|
||||
page_size,
|
||||
)
|
||||
return None
|
||||
if bad_prefix_lens:
|
||||
_index_prefetch_fallback_log(
|
||||
"prefix_not_page_aligned",
|
||||
"prefix length is zero or not page-aligned. prefix_lens=%s page_size=%s",
|
||||
"non-zero prefix length is not page-aligned. prefix_lens=%s page_size=%s",
|
||||
prefix_lens,
|
||||
page_size,
|
||||
)
|
||||
@@ -1605,16 +1537,6 @@ class CpSharedKVIndexPrefetcher:
|
||||
)
|
||||
return None
|
||||
|
||||
cp_group = get_attention_cp_group()
|
||||
if getattr(cp_group, "pynccl_comm", None) is None and layout.cp_size > 1:
|
||||
_index_prefetch_fallback_log(
|
||||
"missing_pynccl",
|
||||
"pynccl communicator is unavailable. cp_rank=%s cp_size=%s",
|
||||
layout.cp_rank,
|
||||
layout.cp_size,
|
||||
)
|
||||
return None
|
||||
|
||||
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
|
||||
create_cpu = _cpu_timing_start()
|
||||
get_cpu = _cpu_timing_start()
|
||||
@@ -1690,6 +1612,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
slot_sorted_logical_pages_by_row=remap.slot_sorted_logical_pages_by_row,
|
||||
slot_sorted_dense_pages_by_row=remap.slot_sorted_dense_pages_by_row,
|
||||
dense_num_pages=remap.dense_num_pages,
|
||||
slot_remap=remap,
|
||||
owned_prefix_pages=owned_prefix_pages,
|
||||
owned_total_pages=owned_total_pages,
|
||||
stream=prefetch_stream,
|
||||
@@ -1730,12 +1653,12 @@ class CpSharedKVIndexPrefetcher:
|
||||
self._log_layer(layer_id, "index_consume_miss layer=%s", layer_id)
|
||||
return None
|
||||
if handle.event is None:
|
||||
self.launch_pending_reduce()
|
||||
self.finalize_pending_materialize()
|
||||
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",
|
||||
"index_consume_miss reason=prefix_materialize_not_ready layer=%s",
|
||||
layer_id,
|
||||
)
|
||||
return None
|
||||
@@ -1780,7 +1703,11 @@ class CpSharedKVIndexPrefetcher:
|
||||
spans=suffix_spans,
|
||||
)
|
||||
)
|
||||
if not materialized_suffix_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_page_buffer):
|
||||
if (
|
||||
not materialized_suffix_by_ipc
|
||||
and suffix_spans
|
||||
and self.layout.cp_size > 1
|
||||
):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"index_prefetch_suffix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
@@ -1796,17 +1723,6 @@ class CpSharedKVIndexPrefetcher:
|
||||
layout=self.layout,
|
||||
spans=suffix_spans,
|
||||
)
|
||||
if not materialized_suffix_by_ipc:
|
||||
for suffix_rows in _slot_spans_to_page_slices(suffix_spans):
|
||||
_all_reduce_materialized_buffer_range(
|
||||
dense_page_buffer,
|
||||
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 suffix_spans=%s",
|
||||
@@ -1895,13 +1811,13 @@ class CpSharedKVIndexPrefetcher:
|
||||
)
|
||||
return None
|
||||
if handle.event is None:
|
||||
self.launch_pending_reduce()
|
||||
self.finalize_pending_materialize()
|
||||
handle = self.handles.get(layer_id)
|
||||
if handle is None or handle.event is None:
|
||||
self._log_layer(
|
||||
layer_id,
|
||||
"index_consume_prefix_current_miss "
|
||||
"reason=prefix_reduce_not_ready layer=%s",
|
||||
"reason=prefix_materialize_not_ready layer=%s",
|
||||
layer_id,
|
||||
)
|
||||
return None
|
||||
@@ -1955,7 +1871,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
if (
|
||||
not current_materialized_by_ipc
|
||||
and self.current_slot_spans
|
||||
and _should_fail_fast_tai_ipc_materialize(dense_page_buffer)
|
||||
and self.layout.cp_size > 1
|
||||
):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"index_prefetch_current_ipc_unavailable",
|
||||
@@ -1964,17 +1880,6 @@ class CpSharedKVIndexPrefetcher:
|
||||
spans=self.current_slot_spans,
|
||||
dense_shape=tuple(dense_page_buffer.shape),
|
||||
)
|
||||
if not current_materialized_by_ipc:
|
||||
for current_pages in _slot_spans_to_page_slices(self.current_slot_spans):
|
||||
_all_reduce_materialized_buffer_range(
|
||||
dense_page_buffer,
|
||||
self.layout.cp_size,
|
||||
current_pages.start,
|
||||
current_pages.stop,
|
||||
nvtx_source="index.prefetch_current",
|
||||
nvtx_layer_id=layer_id,
|
||||
nvtx_cp_rank=self.layout.cp_rank,
|
||||
)
|
||||
remap_ms = _cpu_timing_ms(remap_cpu)
|
||||
total_ms = _cpu_timing_ms(consume_cpu)
|
||||
self._log_layer(
|
||||
@@ -2039,87 +1944,83 @@ class CpSharedKVIndexPrefetcher:
|
||||
start_cpu = _cpu_timing_start()
|
||||
current_stream = torch.cuda.current_stream()
|
||||
get_cpu = _cpu_timing_start()
|
||||
try:
|
||||
page_buffer = _prefetch_pool_get_index_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=next_layer_id,
|
||||
stream=current_stream,
|
||||
)
|
||||
get_ms = _cpu_timing_ms(get_cpu)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to get next-layer index buffer for CP shared KV index prefetch."
|
||||
)
|
||||
self.disabled = True
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start_disable reason=get_index_buffer_failed next_layer=%s",
|
||||
next_layer_id,
|
||||
)
|
||||
return
|
||||
self.stream.wait_stream(current_stream)
|
||||
|
||||
try:
|
||||
prefix_row_spans = _slot_spans_to_page_slices(self.prefix_slot_spans)
|
||||
prefix_rows = prefix_row_spans[0] if prefix_row_spans else slice(0, 0)
|
||||
materialize_cpu = _cpu_timing_start()
|
||||
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 prefix_slot_spans=%s "
|
||||
"dense_pages=%s",
|
||||
next_layer_id,
|
||||
self.prefix_slot_spans,
|
||||
int(dense_page_buffer.shape[0]),
|
||||
)
|
||||
materialized_by_ipc = (
|
||||
_try_tai_ipc_materialize_paged_buffer_page_slot_spans_into(
|
||||
page_buffer=page_buffer,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
spans=self.prefix_slot_spans,
|
||||
with torch.cuda.stream(self.stream):
|
||||
page_buffer = _prefetch_pool_get_index_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=next_layer_id,
|
||||
stream=self.stream,
|
||||
)
|
||||
)
|
||||
if not materialized_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_page_buffer):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"index_prefetch_prefix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
spans=self.prefix_slot_spans,
|
||||
dense_shape=tuple(dense_page_buffer.shape),
|
||||
get_ms = _cpu_timing_ms(get_cpu)
|
||||
dense_page_buffer = page_buffer.new_zeros(
|
||||
(self.dense_num_pages, *page_buffer.shape[1:])
|
||||
)
|
||||
if not materialized_by_ipc:
|
||||
_materialize_local_paged_buffer_page_slot_spans_into(
|
||||
page_buffer=page_buffer,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
spans=self.prefix_slot_spans,
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start_prefix_begin next_layer=%s prefix_slot_spans=%s "
|
||||
"dense_pages=%s stream=prefetch",
|
||||
next_layer_id,
|
||||
self.prefix_slot_spans,
|
||||
int(dense_page_buffer.shape[0]),
|
||||
)
|
||||
materialize_ms = _cpu_timing_ms(materialize_cpu)
|
||||
reduce_cpu = _cpu_timing_start()
|
||||
if materialized_by_ipc:
|
||||
event = _record_event_on_stream(current_stream)
|
||||
else:
|
||||
self.stream.wait_stream(current_stream)
|
||||
with torch.cuda.stream(self.stream):
|
||||
event = _all_reduce_materialized_buffer_ranges_async(
|
||||
dense_kv_cache=dense_page_buffer,
|
||||
row_slices=prefix_row_spans,
|
||||
cp_size=self.layout.cp_size,
|
||||
stream=self.stream,
|
||||
nvtx_source="index.prefetch_prefix",
|
||||
nvtx_layer_id=next_layer_id,
|
||||
nvtx_cp_rank=self.layout.cp_rank,
|
||||
materialized_by_ce = (
|
||||
_try_tai_ipc_ce_materialize_paged_buffer_page_slot_spans_into(
|
||||
page_buffer=page_buffer,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
spans=self.prefix_slot_spans,
|
||||
slot_remap=self.slot_remap,
|
||||
)
|
||||
reduce_enqueue_ms = _cpu_timing_ms(reduce_cpu)
|
||||
)
|
||||
event = _record_event_on_stream(self.stream) if materialized_by_ce else None
|
||||
|
||||
materialized_by_ipc = materialized_by_ce
|
||||
if not materialized_by_ce:
|
||||
current_stream.wait_stream(self.stream)
|
||||
materialized_by_ipc = (
|
||||
_try_tai_ipc_materialize_paged_buffer_page_slot_spans_into(
|
||||
page_buffer=page_buffer,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
spans=self.prefix_slot_spans,
|
||||
slot_remap=self.slot_remap,
|
||||
)
|
||||
)
|
||||
if (
|
||||
not materialized_by_ipc
|
||||
and self.prefix_slot_spans
|
||||
and self.layout.cp_size > 1
|
||||
):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"index_prefetch_prefix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
spans=self.prefix_slot_spans,
|
||||
dense_shape=tuple(dense_page_buffer.shape),
|
||||
)
|
||||
if not materialized_by_ipc:
|
||||
_materialize_local_paged_buffer_page_slot_spans_into(
|
||||
page_buffer=page_buffer,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
spans=self.prefix_slot_spans,
|
||||
)
|
||||
event = _record_event_on_stream(current_stream)
|
||||
materialize_ms = _cpu_timing_ms(materialize_cpu)
|
||||
materialize_enqueue_ms = 0.0
|
||||
if event is None:
|
||||
self.disabled = True
|
||||
_index_prefetch_fallback_log(
|
||||
"async_reduce_unavailable",
|
||||
"async reduce unavailable. layer_id=%s cp_rank=%s cp_size=%s "
|
||||
"materialize_event_unavailable",
|
||||
"materialize event unavailable. layer_id=%s cp_rank=%s cp_size=%s "
|
||||
"prefix_pages=%s total_slots=%s",
|
||||
next_layer_id,
|
||||
self.layout.cp_rank,
|
||||
@@ -2129,7 +2030,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start_disable reason=async_reduce_unavailable next_layer=%s",
|
||||
"index_start_disable reason=materialize_event_unavailable next_layer=%s",
|
||||
next_layer_id,
|
||||
)
|
||||
return
|
||||
@@ -2138,7 +2039,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
total_ms=total_ms,
|
||||
get_ms=get_ms,
|
||||
materialize_ms=materialize_ms,
|
||||
reduce_enqueue_ms=reduce_enqueue_ms,
|
||||
materialize_enqueue_ms=materialize_enqueue_ms,
|
||||
)
|
||||
_log_prefetch_cpu_start(
|
||||
log_fn=self._log,
|
||||
@@ -2149,7 +2050,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
total_ms=total_ms,
|
||||
get_ms=get_ms,
|
||||
materialize_ms=materialize_ms,
|
||||
reduce_enqueue_ms=reduce_enqueue_ms,
|
||||
materialize_enqueue_ms=materialize_enqueue_ms,
|
||||
prefix_pages=self.prefix_pages,
|
||||
total_slots=self.total_slots,
|
||||
dense_units=int(dense_page_buffer.shape[0]),
|
||||
@@ -2176,67 +2077,35 @@ class CpSharedKVIndexPrefetcher:
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start next_layer=%s prefix_pages=%s dense_pages=%s "
|
||||
"reduce_enqueued=True",
|
||||
"materialize_enqueued=True",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
int(dense_page_buffer.shape[0]),
|
||||
)
|
||||
|
||||
def launch_pending_reduce(self) -> None:
|
||||
def finalize_pending_materialize(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",
|
||||
"index_start_prefix_materialize_skip reason=already_enqueued next_layer=%s",
|
||||
handle.layer_id,
|
||||
)
|
||||
return
|
||||
|
||||
prefix_row_spans = list(handle.prefix_row_spans or (handle.prefix_rows,))
|
||||
try:
|
||||
if _should_fail_fast_tai_ipc_materialize(handle.dense_page_buffer):
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"index_prefetch_deferred_prefix_ipc_unavailable",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
dense_shape=tuple(handle.dense_page_buffer.shape),
|
||||
layer_id=handle.layer_id,
|
||||
)
|
||||
current_stream = torch.cuda.current_stream()
|
||||
self.stream.wait_stream(current_stream)
|
||||
with torch.cuda.stream(self.stream):
|
||||
event = _all_reduce_materialized_buffer_ranges_async(
|
||||
dense_kv_cache=handle.dense_page_buffer,
|
||||
row_slices=prefix_row_spans,
|
||||
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,
|
||||
)
|
||||
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 row_spans=%s",
|
||||
handle.layer_id,
|
||||
[(rows.start, rows.stop) for rows in prefix_row_spans],
|
||||
)
|
||||
_raise_tai_ipc_materialize_required(
|
||||
"index_prefetch_deferred_prefix_materialize_missing",
|
||||
cp_rank=self.layout.cp_rank,
|
||||
cp_size=self.layout.cp_size,
|
||||
dense_shape=tuple(handle.dense_page_buffer.shape),
|
||||
layer_id=handle.layer_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to launch CP shared KV index prefix prefetch reduce."
|
||||
"CP shared KV index prefix prefetch has no IPC/CE completion event."
|
||||
)
|
||||
self.disabled = True
|
||||
self.handles.pop(handle.layer_id, None)
|
||||
@@ -2244,7 +2113,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
self.pending_attention_handle = None
|
||||
self._log_next_layer(
|
||||
handle.layer_id,
|
||||
"index_start_disable reason=reduce_launch_exception next_layer=%s",
|
||||
"index_start_disable reason=materialize_event_missing next_layer=%s",
|
||||
handle.layer_id,
|
||||
)
|
||||
return
|
||||
|
||||
@@ -2679,6 +2679,253 @@ def _get_or_build_prefix_ipc_slot_descriptors(
|
||||
return descriptors
|
||||
|
||||
|
||||
def _build_prefix_ipc_slot_descriptors_uncached(
|
||||
*,
|
||||
slot_logical_pages: torch.Tensor,
|
||||
layout: CpSharedKVLayout,
|
||||
spans: list[tuple[int, int]],
|
||||
device: torch.device,
|
||||
physical_page_capacity: int | None,
|
||||
) -> _IpcPrefixSlotDescriptors:
|
||||
flat_slot_logical_pages = _contiguous_for_tai(
|
||||
slot_logical_pages.reshape(-1).to(device=device)
|
||||
)
|
||||
slot_indices = _slot_spans_to_cuda_slot_indices(
|
||||
spans,
|
||||
total_slots=int(flat_slot_logical_pages.numel()),
|
||||
device=device,
|
||||
)
|
||||
if slot_indices.numel() == 0:
|
||||
empty = torch.empty((0,), dtype=torch.long, device=device)
|
||||
return _IpcPrefixSlotDescriptors(empty, empty, empty, empty)
|
||||
slot_logical_pages_range = _contiguous_for_tai(
|
||||
flat_slot_logical_pages.index_select(0, slot_indices)
|
||||
)
|
||||
owner_ranks, src_page_indices = build_cp_shared_kv_ipc_page_descriptors(
|
||||
slot_logical_pages_range,
|
||||
layout,
|
||||
physical_page_capacity=physical_page_capacity,
|
||||
)
|
||||
return _IpcPrefixSlotDescriptors(
|
||||
slot_indices=slot_indices.contiguous(),
|
||||
owner_ranks=owner_ranks,
|
||||
src_page_indices=src_page_indices,
|
||||
dense_page_indices=(slot_indices + 1).to(torch.long).contiguous(),
|
||||
)
|
||||
|
||||
|
||||
def _get_prefix_ipc_slot_descriptors(
|
||||
*,
|
||||
slot_remap: SharedTokenKVSlotRemap | SharedPagedBufferSlotRemap | None,
|
||||
slot_logical_pages: torch.Tensor,
|
||||
layout: CpSharedKVLayout,
|
||||
spans: list[tuple[int, int]],
|
||||
device: torch.device,
|
||||
physical_page_capacity: int | None,
|
||||
cache_kind: str,
|
||||
) -> _IpcPrefixSlotDescriptors:
|
||||
if slot_remap is not None:
|
||||
return _get_or_build_prefix_ipc_slot_descriptors(
|
||||
slot_remap=slot_remap,
|
||||
layout=layout,
|
||||
spans=spans,
|
||||
device=device,
|
||||
physical_page_capacity=physical_page_capacity,
|
||||
cache_kind=cache_kind,
|
||||
)
|
||||
return _build_prefix_ipc_slot_descriptors_uncached(
|
||||
slot_logical_pages=slot_logical_pages,
|
||||
layout=layout,
|
||||
spans=spans,
|
||||
device=device,
|
||||
physical_page_capacity=physical_page_capacity,
|
||||
)
|
||||
|
||||
|
||||
def _filter_ipc_prefix_descriptors_for_ce(
|
||||
descriptors: _IpcPrefixSlotDescriptors,
|
||||
) -> _IpcPrefixSlotDescriptors:
|
||||
"""Drop invalid/sentinel slots before CE copy.
|
||||
|
||||
The SM IPC materialize kernel can zero-fill invalid slots. The CE path is
|
||||
intentionally just cudaMemcpyBatchAsync over valid spans; callers allocate
|
||||
the dense destination with zeros, so invalid slots should be skipped instead
|
||||
of submitted to the CE primitive.
|
||||
"""
|
||||
|
||||
if descriptors.owner_ranks.numel() == 0:
|
||||
return descriptors
|
||||
valid = (
|
||||
(descriptors.owner_ranks >= 0)
|
||||
& (descriptors.src_page_indices >= 0)
|
||||
& (descriptors.dense_page_indices >= 0)
|
||||
)
|
||||
if bool(valid.all().item()):
|
||||
return descriptors
|
||||
return _IpcPrefixSlotDescriptors(
|
||||
slot_indices=descriptors.slot_indices[valid].contiguous(),
|
||||
owner_ranks=descriptors.owner_ranks[valid].contiguous(),
|
||||
src_page_indices=descriptors.src_page_indices[valid].contiguous(),
|
||||
dense_page_indices=descriptors.dense_page_indices[valid].contiguous(),
|
||||
)
|
||||
|
||||
|
||||
def _try_tai_ipc_ce_materialize_token_kv_page_slot_spans_into(
|
||||
*,
|
||||
kv_cache: torch.Tensor,
|
||||
dense_kv_cache: torch.Tensor,
|
||||
slot_logical_pages: torch.Tensor,
|
||||
layout: CpSharedKVLayout,
|
||||
page_size: int,
|
||||
spans: list[tuple[int, int]],
|
||||
slot_remap: SharedTokenKVSlotRemap | None = None,
|
||||
) -> bool:
|
||||
if not spans:
|
||||
return True
|
||||
if not dense_kv_cache.is_cuda or not dense_kv_cache.is_contiguous():
|
||||
_log_tai_ipc_materialize_fallback(
|
||||
"ce_dense_token_tensor_unsupported",
|
||||
"CP shared KV tai IPC CE token span materialize requires a contiguous "
|
||||
"CUDA dense tensor; falling back to SM IPC or collective. "
|
||||
"device=%s contiguous=%s dense_shape=%s spans=%s",
|
||||
dense_kv_cache.device,
|
||||
dense_kv_cache.is_contiguous(),
|
||||
tuple(dense_kv_cache.shape),
|
||||
spans,
|
||||
limit=4,
|
||||
)
|
||||
return False
|
||||
|
||||
ipc_state = _get_or_open_tai_ipc_peer_ptrs(kv_cache, layout)
|
||||
if ipc_state is None:
|
||||
return False
|
||||
kernels, peer_ptrs = ipc_state
|
||||
ce_kernel = getattr(kernels, "materialize_cuda_ipc_peer_pages_slot_indices_ce", None)
|
||||
if ce_kernel is None:
|
||||
_log_tai_ipc_materialize_fallback(
|
||||
"ce_kernel_missing",
|
||||
"CP shared KV tai IPC CE token materialize is unavailable; "
|
||||
"falling back to SM IPC or collective. Upgrade tai-kernel.",
|
||||
limit=1,
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
descriptors = _get_prefix_ipc_slot_descriptors(
|
||||
slot_remap=slot_remap,
|
||||
slot_logical_pages=slot_logical_pages,
|
||||
layout=layout,
|
||||
spans=spans,
|
||||
device=torch.device("cpu"),
|
||||
physical_page_capacity=kv_cache.shape[0] // page_size,
|
||||
cache_kind="token-ce",
|
||||
)
|
||||
descriptors = _filter_ipc_prefix_descriptors_for_ce(descriptors)
|
||||
if descriptors.owner_ranks.numel() == 0:
|
||||
return True
|
||||
ce_kernel(
|
||||
peer_ptrs,
|
||||
dense_kv_cache,
|
||||
descriptors.owner_ranks,
|
||||
descriptors.src_page_indices,
|
||||
descriptors.dense_page_indices,
|
||||
page_nbytes=_token_kv_page_nbytes(kv_cache, page_size),
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
_log_tai_ipc_materialize_fallback(
|
||||
"token_span_ce_kernel_failed",
|
||||
"CP shared KV tai IPC CE token span materialize failed; falling back "
|
||||
"to SM IPC or collective materialize. cp_rank=%s cp_size=%s spans=%s "
|
||||
"page_size=%s kv_shape=%s dense_shape=%s error=%s",
|
||||
layout.cp_rank,
|
||||
layout.cp_size,
|
||||
spans,
|
||||
page_size,
|
||||
tuple(kv_cache.shape),
|
||||
tuple(dense_kv_cache.shape),
|
||||
exc,
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def _try_tai_ipc_ce_materialize_paged_buffer_page_slot_spans_into(
|
||||
*,
|
||||
page_buffer: torch.Tensor,
|
||||
dense_page_buffer: torch.Tensor,
|
||||
slot_logical_pages: torch.Tensor,
|
||||
layout: CpSharedKVLayout,
|
||||
spans: list[tuple[int, int]],
|
||||
slot_remap: SharedPagedBufferSlotRemap | None = None,
|
||||
) -> bool:
|
||||
if not spans:
|
||||
return True
|
||||
if not dense_page_buffer.is_cuda or not dense_page_buffer.is_contiguous():
|
||||
_log_tai_ipc_materialize_fallback(
|
||||
"ce_dense_paged_tensor_unsupported",
|
||||
"CP shared KV tai IPC CE paged span materialize requires a contiguous "
|
||||
"CUDA dense tensor; falling back to SM IPC or collective. "
|
||||
"device=%s contiguous=%s dense_shape=%s spans=%s",
|
||||
dense_page_buffer.device,
|
||||
dense_page_buffer.is_contiguous(),
|
||||
tuple(dense_page_buffer.shape),
|
||||
spans,
|
||||
limit=4,
|
||||
)
|
||||
return False
|
||||
|
||||
ipc_state = _get_or_open_tai_ipc_peer_ptrs(page_buffer, layout)
|
||||
if ipc_state is None:
|
||||
return False
|
||||
kernels, peer_ptrs = ipc_state
|
||||
ce_kernel = getattr(kernels, "materialize_cuda_ipc_peer_pages_slot_indices_ce", None)
|
||||
if ce_kernel is None:
|
||||
_log_tai_ipc_materialize_fallback(
|
||||
"ce_kernel_missing",
|
||||
"CP shared KV tai IPC CE paged materialize is unavailable; "
|
||||
"falling back to SM IPC or collective. Upgrade tai-kernel.",
|
||||
limit=1,
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
descriptors = _get_prefix_ipc_slot_descriptors(
|
||||
slot_remap=slot_remap,
|
||||
slot_logical_pages=slot_logical_pages,
|
||||
layout=layout,
|
||||
spans=spans,
|
||||
device=torch.device("cpu"),
|
||||
physical_page_capacity=page_buffer.shape[0],
|
||||
cache_kind="paged-ce",
|
||||
)
|
||||
descriptors = _filter_ipc_prefix_descriptors_for_ce(descriptors)
|
||||
if descriptors.owner_ranks.numel() == 0:
|
||||
return True
|
||||
ce_kernel(
|
||||
peer_ptrs,
|
||||
dense_page_buffer,
|
||||
descriptors.owner_ranks,
|
||||
descriptors.src_page_indices,
|
||||
descriptors.dense_page_indices,
|
||||
page_nbytes=_page_nbytes_from_page_tensor(page_buffer),
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
_log_tai_ipc_materialize_fallback(
|
||||
"paged_span_ce_kernel_failed",
|
||||
"CP shared KV tai IPC CE paged span materialize failed; falling back "
|
||||
"to SM IPC or collective materialize. cp_rank=%s cp_size=%s spans=%s "
|
||||
"page_shape=%s dense_shape=%s error=%s",
|
||||
layout.cp_rank,
|
||||
layout.cp_size,
|
||||
spans,
|
||||
tuple(page_buffer.shape),
|
||||
tuple(dense_page_buffer.shape),
|
||||
exc,
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def _get_or_build_current_ipc_slot_descriptors(
|
||||
*,
|
||||
slot_remap: SharedTokenKVSlotRemap | SharedPagedBufferSlotRemap,
|
||||
|
||||
@@ -20,7 +20,6 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
build_batch_prefix_slot_spans,
|
||||
cp_shared_kv_debug_enabled,
|
||||
cp_shared_kv_debug_log,
|
||||
cp_shared_kv_mla_prefetch_enabled,
|
||||
cp_shared_kv_mla_prefetch_log,
|
||||
cp_shared_kv_mla_prefetch_log_enabled,
|
||||
cp_shared_kv_mla_prefetch_should_log_layer,
|
||||
@@ -477,14 +476,6 @@ def _log_cp_shared_kv_index_prefetch_fallback(
|
||||
)
|
||||
|
||||
|
||||
def _should_log_missing_index_prefetcher() -> bool:
|
||||
# When MLA/index prefetch is disabled, a missing index prefetcher is the
|
||||
# expected sync-materialize path, not a fast-path fallback. Still log consume
|
||||
# misses when a prefetcher exists because those indicate an enabled prefetch
|
||||
# pipeline failed to provide the requested layer/current buffer.
|
||||
return cp_shared_kv_mla_prefetch_enabled()
|
||||
|
||||
|
||||
class BaseIndexerMetadata(ABC):
|
||||
@abstractmethod
|
||||
def get_seqlens_int32(self) -> torch.Tensor:
|
||||
@@ -798,25 +789,15 @@ class Indexer(MultiPlatformOp):
|
||||
)
|
||||
if prefetched is not None:
|
||||
return prefetched
|
||||
_log_cp_shared_kv_index_prefetch_fallback(
|
||||
"current_consume_miss",
|
||||
"prefetcher did not provide current index prefix buffer; "
|
||||
"falling back to sync partial-current index compose. "
|
||||
"layer=%s cp_rank=%s prefix_lens=%s extend_lens=%s "
|
||||
"logical_page_table_shape=%s",
|
||||
layer_id,
|
||||
layout.cp_rank,
|
||||
prefix_lens,
|
||||
extend_lens,
|
||||
tuple(logical_page_table.shape),
|
||||
)
|
||||
else:
|
||||
if _should_log_missing_index_prefetcher():
|
||||
if int(layer_id) > int(
|
||||
getattr(forward_batch.token_to_kv_pool, "start_layer", 0)
|
||||
):
|
||||
_log_cp_shared_kv_index_prefetch_fallback(
|
||||
"current_missing_prefetcher",
|
||||
"index prefetcher is unavailable; falling back to sync "
|
||||
"partial-current index compose. layer=%s cp_rank=%s "
|
||||
"prefix_lens=%s extend_lens=%s logical_page_table_shape=%s",
|
||||
"current_consume_miss",
|
||||
"prefetcher did not provide current index prefix buffer; "
|
||||
"falling back to sync partial-current index compose. "
|
||||
"layer=%s cp_rank=%s prefix_lens=%s extend_lens=%s "
|
||||
"logical_page_table_shape=%s",
|
||||
layer_id,
|
||||
layout.cp_rank,
|
||||
prefix_lens,
|
||||
@@ -901,16 +882,6 @@ class Indexer(MultiPlatformOp):
|
||||
layout.cp_rank,
|
||||
tuple(logical_page_table.shape),
|
||||
)
|
||||
else:
|
||||
if _should_log_missing_index_prefetcher():
|
||||
_log_cp_shared_kv_index_prefetch_fallback(
|
||||
"missing_prefetcher",
|
||||
"index prefetcher is unavailable; falling back to sync paged "
|
||||
"materialize. layer=%s cp_rank=%s logical_page_table_shape=%s",
|
||||
layer_id,
|
||||
layout.cp_rank,
|
||||
tuple(logical_page_table.shape),
|
||||
)
|
||||
|
||||
if cp_shared_kv_debug_enabled():
|
||||
cp_shared_kv_debug_log(
|
||||
|
||||
@@ -76,7 +76,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_cp_group, get_attention_tp_size
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.utils import is_cuda, is_hip
|
||||
|
||||
@@ -166,7 +166,7 @@ def _maybe_start_cp_shared_kv_attention_prefetch(
|
||||
This hook intentionally runs after current-layer index/MLA cache
|
||||
materialization has finished and immediately before the attention kernel.
|
||||
Earlier hooks can make the next-layer collective overlap current-layer KV
|
||||
materialization/reduce instead of attention, which can serialize or contend
|
||||
materialization instead of attention, which can serialize or contend
|
||||
with the current layer's required cache work.
|
||||
"""
|
||||
|
||||
@@ -1061,17 +1061,11 @@ class NativeSparseAttnBackend(
|
||||
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",
|
||||
"real_pages=%s page_table_shape=%s",
|
||||
getattr(layout, "cp_rank", None),
|
||||
getattr(layout, "cp_size", None),
|
||||
getattr(forward_batch, "uses_cp_shared_kv", None),
|
||||
@@ -1095,7 +1089,6 @@ class NativeSparseAttnBackend(
|
||||
tuple(metadata.page_table_1.shape)
|
||||
if metadata.page_table_1 is not None
|
||||
else None,
|
||||
pynccl_available,
|
||||
)
|
||||
|
||||
def _cal_indexer_k_start_end(
|
||||
@@ -2468,7 +2461,7 @@ class NativeSparseAttnBackend(
|
||||
forward_batch, "cp_shared_kv_index_prefetcher", None
|
||||
)
|
||||
if index_prefetcher is not None:
|
||||
index_prefetcher.launch_pending_reduce()
|
||||
index_prefetcher.finalize_pending_materialize()
|
||||
|
||||
try:
|
||||
if nsa_impl == "tilelang":
|
||||
@@ -2654,40 +2647,69 @@ class NativeSparseAttnBackend(
|
||||
ragged_current_req_id = ragged_current_req_id[
|
||||
: int(current_locs_for_reuse.shape[0])
|
||||
]
|
||||
slot_remap = get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
remap_logical_pages=metadata.real_page_table,
|
||||
layout=forward_batch.cp_shared_kv_layout,
|
||||
page_size=page_size,
|
||||
prefetched_kv = None
|
||||
mla_prefetcher = getattr(
|
||||
forward_batch, "cp_shared_kv_mla_prefetcher", None
|
||||
)
|
||||
kv_cache, page_table_1_flattened = (
|
||||
materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
if mla_prefetcher is not None:
|
||||
prefetched_kv = (
|
||||
mla_prefetcher.consume_prefix_with_current(
|
||||
layer_id=layer.layer_id,
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=page_table_1_flattened,
|
||||
current_kv_cache=current_kv_cache,
|
||||
current_locs=current_locs_for_reuse,
|
||||
loc_req_id=logical_locs_row_ids,
|
||||
current_req_id=ragged_current_req_id,
|
||||
)
|
||||
)
|
||||
if prefetched_kv is not None:
|
||||
kv_cache, page_table_1_flattened = prefetched_kv
|
||||
else:
|
||||
if (
|
||||
mla_prefetcher is not None
|
||||
and int(layer.layer_id) > int(getattr(forward_batch.token_to_kv_pool, "start_layer", 0))
|
||||
):
|
||||
_log_cp_shared_kv_mla_prefetch_fallback(
|
||||
"ragged_current_consume_miss",
|
||||
"MLA prefetch fast path did not provide "
|
||||
"RAGGED partial-current prefix; using "
|
||||
"synchronous compose. cp_rank=%s layer=%s "
|
||||
"prefix_lens=%s extend_lens=%s "
|
||||
"current_rows=%s page_table_shape=%s",
|
||||
forward_batch.cp_shared_kv_layout.cp_rank,
|
||||
layer.layer_id,
|
||||
[int(x) for x in prefix_lens_cpu],
|
||||
[int(x) for x in extend_lens_cpu],
|
||||
int(current_kv_cache.shape[0]),
|
||||
tuple(page_table_1_flattened.shape),
|
||||
)
|
||||
slot_remap = get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=page_table_1_flattened,
|
||||
current_kv_cache=current_kv_cache,
|
||||
current_locs=current_locs_for_reuse,
|
||||
slot_remap=slot_remap,
|
||||
remap_logical_pages=metadata.real_page_table,
|
||||
layout=forward_batch.cp_shared_kv_layout,
|
||||
page_size=page_size,
|
||||
prefix_pages=prefix_pages,
|
||||
prefix_slot_spans=prefix_slot_spans,
|
||||
current_slot_spans=current_slot_spans,
|
||||
loc_req_id=logical_locs_row_ids,
|
||||
current_req_id=ragged_current_req_id,
|
||||
layer_id=layer.layer_id,
|
||||
nvtx_source="mla.ragged_partial_current_sync",
|
||||
current_page_writer_ranks=(
|
||||
maybe_build_current_page_writer_ranks(
|
||||
forward_batch=forward_batch,
|
||||
prefix_lens_cpu=prefix_lens_cpu,
|
||||
extend_lens_cpu=extend_lens_cpu,
|
||||
page_size=page_size,
|
||||
layout=forward_batch.cp_shared_kv_layout,
|
||||
)
|
||||
),
|
||||
|
||||
)
|
||||
kv_cache, page_table_1_flattened = (
|
||||
materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=page_table_1_flattened,
|
||||
current_kv_cache=current_kv_cache,
|
||||
current_locs=current_locs_for_reuse,
|
||||
slot_remap=slot_remap,
|
||||
layout=forward_batch.cp_shared_kv_layout,
|
||||
page_size=page_size,
|
||||
prefix_pages=prefix_pages,
|
||||
prefix_slot_spans=prefix_slot_spans,
|
||||
current_slot_spans=current_slot_spans,
|
||||
loc_req_id=logical_locs_row_ids,
|
||||
current_req_id=ragged_current_req_id,
|
||||
layer_id=layer.layer_id,
|
||||
nvtx_source="mla.ragged_partial_current_sync",
|
||||
)
|
||||
)
|
||||
)
|
||||
if envs.SGLANG_NSA_DEQUANT_ONLY_TOPK.get():
|
||||
kv_cache = self._dequant_topk_inplace(
|
||||
kv_cache,
|
||||
@@ -2766,7 +2788,7 @@ class NativeSparseAttnBackend(
|
||||
)
|
||||
finally:
|
||||
if mla_prefetcher is not None:
|
||||
mla_prefetcher.launch_pending_reduce()
|
||||
mla_prefetcher.finalize_pending_materialize()
|
||||
mla_prefetcher.wait_attention_window()
|
||||
index_prefetcher = getattr(
|
||||
forward_batch, "cp_shared_kv_index_prefetcher", None
|
||||
|
||||
Reference in New Issue
Block a user