Enable batched CP shared-KV current reuse correctness
Batch-size > 1 cache-hit traffic must not lose the current/partial-current reuse fast path. This change extends the target MLA/index sync-correct path to validate batched current suffix rows, compose page-aligned prefix spans, and route batched index top-k through current-only or partial-current reuse without falling back to scalar guards.\n\nThe implementation keeps page as the minimum cache unit: prefix cache coverage is page-aligned, while current suffix rows are sliced by valid extend lengths so padded tail rows are not exposed to attention. The index top-k batch path now mirrors the single-request contract: current-only reuses current index KV directly, partial-current materializes once and shares the dense buffer across request segments.\n\nConstraint: CP shared-KV supports in-seq zigzag only; round-robin remains fail-fast for shared-KV.\nConstraint: No new collectives are introduced; this is a sync correctness path, not a new communication scheme.\nRejected: Keep bs>1 current_index_kv fail-fast | disables the cache-hit path that W4 is meant to restore.\nRejected: Pad requests to a common batch length | violates the page-granular contract and exposes padded tails to consumers.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not reintroduce batch_size_not_one or batch_gt1 current-reuse guards without proving an equivalent batched fast path exists.\nTested: Local py_compile for touched runtime/test files.\nTested: Remote g0034 pytest test_nsa_cp_utils.py test_cp_shared_kv_runtime.py => 157 passed, 5 warnings, 2 subtests passed.\nTested: Remote g0034 pytest test_cp_shared_kv_layout.py => 27 passed, 3 warnings.\nNot-tested: Full ETE with live traffic; CUDA kernel-backed batched descriptors; L1 prefetch, HiCache load/backup, and draft/EAGLE bs>1 current reuse.
This commit is contained in:
@@ -1783,11 +1783,8 @@ def _current_extend_kv_reuse_miss_reason(forward_batch) -> str | None:
|
||||
if not forward_mode.is_extend_without_speculative():
|
||||
return "not_extend_without_speculative"
|
||||
|
||||
batch_size = int(getattr(forward_batch, "batch_size", 0))
|
||||
if batch_size != 1:
|
||||
return "batch_size_not_one"
|
||||
|
||||
extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None)
|
||||
extend_prefix_lens_cpu = getattr(forward_batch, "extend_prefix_lens_cpu", None)
|
||||
seq_lens_cpu = getattr(forward_batch, "seq_lens_cpu", None)
|
||||
out_cache_loc = getattr(forward_batch, "out_cache_loc", None)
|
||||
if extend_seq_lens_cpu is None:
|
||||
@@ -1796,22 +1793,42 @@ def _current_extend_kv_reuse_miss_reason(forward_batch) -> str | None:
|
||||
return "missing_seq_lens_cpu"
|
||||
if out_cache_loc is None:
|
||||
return "missing_out_cache_loc"
|
||||
if len(extend_seq_lens_cpu) != 1:
|
||||
return "extend_batch_not_one"
|
||||
if int(seq_lens_cpu.numel()) != 1:
|
||||
return "seq_lens_batch_not_one"
|
||||
|
||||
extend_len = int(extend_seq_lens_cpu[0])
|
||||
seq_len = int(seq_lens_cpu[0].item())
|
||||
if extend_len <= 0:
|
||||
return "non_positive_extend_len"
|
||||
if seq_len < extend_len:
|
||||
return "seq_len_smaller_than_extend_len"
|
||||
batch_size = int(getattr(forward_batch, "batch_size", 0) or 0)
|
||||
extend_batch_size = len(extend_seq_lens_cpu)
|
||||
seq_lens_batch_size = int(seq_lens_cpu.numel())
|
||||
if batch_size <= 0:
|
||||
batch_size = extend_batch_size
|
||||
if extend_batch_size != batch_size:
|
||||
return "extend_batch_size_mismatch"
|
||||
if seq_lens_batch_size != batch_size:
|
||||
return "seq_lens_batch_size_mismatch"
|
||||
if extend_prefix_lens_cpu is not None and len(extend_prefix_lens_cpu) != batch_size:
|
||||
return "prefix_batch_size_mismatch"
|
||||
|
||||
valid_current_rows = 0
|
||||
for req_id, (extend_len_raw, seq_len_raw) in enumerate(
|
||||
zip(extend_seq_lens_cpu, seq_lens_cpu)
|
||||
):
|
||||
extend_len = int(extend_len_raw)
|
||||
seq_len = int(seq_len_raw.item())
|
||||
if extend_len <= 0:
|
||||
return f"non_positive_extend_len_req_{req_id}"
|
||||
if seq_len < extend_len:
|
||||
return f"seq_len_smaller_than_extend_len_req_{req_id}"
|
||||
if extend_prefix_lens_cpu is not None:
|
||||
prefix_len = int(extend_prefix_lens_cpu[req_id])
|
||||
if prefix_len < 0:
|
||||
return f"negative_prefix_len_req_{req_id}"
|
||||
if prefix_len + extend_len != seq_len:
|
||||
return f"prefix_extend_seq_len_mismatch_req_{req_id}"
|
||||
valid_current_rows += extend_len
|
||||
|
||||
# ForwardBatch pads tensors such as out_cache_loc at the tail for CUDA graph
|
||||
# and CP alignment. The first extend_len rows still cover the valid current
|
||||
# suffix, so padded batches remain eligible for current reuse as long as the
|
||||
# valid suffix is fully present.
|
||||
if int(out_cache_loc.numel()) < extend_len:
|
||||
if int(out_cache_loc.numel()) < valid_current_rows:
|
||||
return "out_cache_loc_shorter_than_extend"
|
||||
return None
|
||||
|
||||
@@ -1887,9 +1904,9 @@ def current_extend_kv_rows_for_reuse(
|
||||
return None
|
||||
|
||||
extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None)
|
||||
if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) != 1:
|
||||
if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) == 0:
|
||||
return None
|
||||
valid_current_rows = int(extend_seq_lens_cpu[0])
|
||||
valid_current_rows = sum(int(x) for x in extend_seq_lens_cpu)
|
||||
if valid_current_rows <= 0:
|
||||
return None
|
||||
|
||||
@@ -1903,6 +1920,84 @@ def current_extend_kv_rows_for_reuse(
|
||||
return valid_current_rows
|
||||
|
||||
|
||||
def build_batch_prefix_slot_span(
|
||||
*,
|
||||
logical_pages: torch.Tensor,
|
||||
prefix_lens_cpu,
|
||||
page_size: int,
|
||||
) -> tuple[int, int]:
|
||||
"""Return the flattened page-table slot span covering batched prefix pages.
|
||||
|
||||
CP shared KV page-table slot layout is row-major by request. For bs>1
|
||||
partial-current reuse, each request has its own prefix length, so a scalar
|
||||
``prefix_pages`` is not sufficient. This helper returns the smallest
|
||||
contiguous flattened slot span that covers all request prefix page ranges.
|
||||
|
||||
The span can include current/suffix slots between two request prefixes. Those
|
||||
slots are overwritten by current rows and masked before consumption; using a
|
||||
single span keeps the synchronous fallback to one range collective instead of
|
||||
one collective per request.
|
||||
"""
|
||||
|
||||
if page_size <= 0:
|
||||
raise ValueError(f"page_size must be positive, got {page_size}")
|
||||
if prefix_lens_cpu is None:
|
||||
raise ValueError("prefix_lens_cpu is required for batch prefix slot span")
|
||||
batch_size = len(prefix_lens_cpu)
|
||||
if batch_size == 0:
|
||||
return (0, 0)
|
||||
|
||||
if logical_pages.dim() == 1:
|
||||
if batch_size != 1:
|
||||
raise ValueError(
|
||||
"1D logical_pages can only describe one request for prefix slot span: "
|
||||
f"batch_size={batch_size} logical_pages_shape={tuple(logical_pages.shape)}"
|
||||
)
|
||||
pages_per_request = int(logical_pages.numel())
|
||||
else:
|
||||
if int(logical_pages.shape[0]) < batch_size:
|
||||
raise ValueError(
|
||||
"logical_pages has fewer rows than prefix_lens_cpu: "
|
||||
f"rows={int(logical_pages.shape[0])} batch_size={batch_size}"
|
||||
)
|
||||
pages_per_request = int(logical_pages.reshape(logical_pages.shape[0], -1).shape[1])
|
||||
|
||||
if pages_per_request < 0:
|
||||
raise ValueError(f"pages_per_request must be non-negative: {pages_per_request}")
|
||||
|
||||
start_slot: int | None = None
|
||||
end_slot: int | None = None
|
||||
for req_id, prefix_len_raw in enumerate(prefix_lens_cpu):
|
||||
prefix_len = int(prefix_len_raw)
|
||||
if prefix_len < 0:
|
||||
raise ValueError(
|
||||
f"prefix_lens_cpu contains negative prefix length: req_id={req_id} "
|
||||
f"prefix_len={prefix_len}"
|
||||
)
|
||||
if prefix_len % page_size != 0:
|
||||
raise ValueError(
|
||||
"CP shared KV batch prefix slot span requires page-aligned prefixes: "
|
||||
f"req_id={req_id} prefix_len={prefix_len} page_size={page_size}"
|
||||
)
|
||||
prefix_pages = prefix_len // page_size
|
||||
if prefix_pages == 0:
|
||||
continue
|
||||
if prefix_pages > pages_per_request:
|
||||
raise ValueError(
|
||||
"prefix pages exceed per-request logical page-table width: "
|
||||
f"req_id={req_id} prefix_pages={prefix_pages} "
|
||||
f"pages_per_request={pages_per_request}"
|
||||
)
|
||||
req_start = req_id * pages_per_request
|
||||
req_end = req_start + prefix_pages
|
||||
start_slot = req_start if start_slot is None else min(start_slot, req_start)
|
||||
end_slot = req_end if end_slot is None else max(end_slot, req_end)
|
||||
|
||||
if start_slot is None or end_slot is None:
|
||||
return (0, 0)
|
||||
return (start_slot, end_slot)
|
||||
|
||||
|
||||
def current_loc_remap_fast_path_args(
|
||||
forward_batch,
|
||||
) -> tuple[int | None, int | None]:
|
||||
@@ -2892,6 +2987,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
layout: CpSharedKVLayout,
|
||||
page_size: int,
|
||||
prefix_pages: int,
|
||||
prefix_slot_span: tuple[int, int] | None = None,
|
||||
layer_id: int | None = None,
|
||||
nvtx_source: str = "mla.partial_current_sync",
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
@@ -2905,11 +3001,28 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
"""
|
||||
|
||||
total_slots = int(slot_remap.slot_logical_pages.numel())
|
||||
if prefix_pages < 0 or prefix_pages > total_slots:
|
||||
raise ValueError(
|
||||
"Invalid CP shared KV partial-current prefix range: "
|
||||
f"prefix_pages={prefix_pages} total_slots={total_slots}"
|
||||
if prefix_slot_span is None:
|
||||
if prefix_pages < 0 or prefix_pages > total_slots:
|
||||
raise ValueError(
|
||||
"Invalid CP shared KV partial-current prefix range: "
|
||||
f"prefix_pages={prefix_pages} total_slots={total_slots}"
|
||||
)
|
||||
prefix_start_slot = 0
|
||||
prefix_end_slot = int(prefix_pages)
|
||||
else:
|
||||
prefix_start_slot, prefix_end_slot = (
|
||||
int(prefix_slot_span[0]),
|
||||
int(prefix_slot_span[1]),
|
||||
)
|
||||
if (
|
||||
prefix_start_slot < 0
|
||||
or prefix_end_slot < prefix_start_slot
|
||||
or prefix_end_slot > total_slots
|
||||
):
|
||||
raise ValueError(
|
||||
"Invalid CP shared KV partial-current prefix slot span: "
|
||||
f"prefix_slot_span={prefix_slot_span} total_slots={total_slots}"
|
||||
)
|
||||
|
||||
dense_kv_cache = kv_cache.new_zeros(
|
||||
(slot_remap.dense_num_pages * page_size, *kv_cache.shape[1:])
|
||||
@@ -2920,8 +3033,8 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
slot_logical_pages=slot_remap.slot_logical_pages,
|
||||
layout=layout,
|
||||
page_size=page_size,
|
||||
start_slot=0,
|
||||
end_slot=prefix_pages,
|
||||
start_slot=prefix_start_slot,
|
||||
end_slot=prefix_end_slot,
|
||||
)
|
||||
if not materialized_by_ipc:
|
||||
materialize_local_token_kv_page_slots_into(
|
||||
@@ -2930,11 +3043,15 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
slot_logical_pages=slot_remap.slot_logical_pages,
|
||||
layout=layout,
|
||||
page_size=page_size,
|
||||
start_slot=0,
|
||||
end_slot=prefix_pages,
|
||||
start_slot=prefix_start_slot,
|
||||
end_slot=prefix_end_slot,
|
||||
)
|
||||
|
||||
prefix_rows = slot_range_to_token_slice(page_size, 0, prefix_pages)
|
||||
prefix_rows = slot_range_to_token_slice(
|
||||
page_size,
|
||||
prefix_start_slot,
|
||||
prefix_end_slot,
|
||||
)
|
||||
_all_reduce_materialized_buffer_range(
|
||||
dense_kv_cache,
|
||||
layout.cp_size,
|
||||
@@ -2979,17 +3096,35 @@ def materialize_prefix_and_reuse_current_index_page_slots(
|
||||
page_size: int,
|
||||
index_head_dim: int,
|
||||
prefix_pages: int,
|
||||
prefix_slot_span: tuple[int, int] | None = None,
|
||||
layer_id: int | None = None,
|
||||
nvtx_source: str = "index.partial_current_sync",
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Synchronously compose prefix index materialization with current index rows."""
|
||||
|
||||
total_slots = int(slot_remap.slot_logical_pages.numel())
|
||||
if prefix_pages < 0 or prefix_pages > total_slots:
|
||||
raise ValueError(
|
||||
"Invalid CP shared KV index partial-current prefix range: "
|
||||
f"prefix_pages={prefix_pages} total_slots={total_slots}"
|
||||
if prefix_slot_span is None:
|
||||
if prefix_pages < 0 or prefix_pages > total_slots:
|
||||
raise ValueError(
|
||||
"Invalid CP shared KV index partial-current prefix range: "
|
||||
f"prefix_pages={prefix_pages} total_slots={total_slots}"
|
||||
)
|
||||
prefix_start_slot = 0
|
||||
prefix_end_slot = int(prefix_pages)
|
||||
else:
|
||||
prefix_start_slot, prefix_end_slot = (
|
||||
int(prefix_slot_span[0]),
|
||||
int(prefix_slot_span[1]),
|
||||
)
|
||||
if (
|
||||
prefix_start_slot < 0
|
||||
or prefix_end_slot < prefix_start_slot
|
||||
or prefix_end_slot > total_slots
|
||||
):
|
||||
raise ValueError(
|
||||
"Invalid CP shared KV index partial-current prefix slot span: "
|
||||
f"prefix_slot_span={prefix_slot_span} total_slots={total_slots}"
|
||||
)
|
||||
|
||||
dense_page_buffer = page_buffer.new_zeros(
|
||||
(slot_remap.dense_num_pages, *page_buffer.shape[1:])
|
||||
@@ -2999,8 +3134,8 @@ def materialize_prefix_and_reuse_current_index_page_slots(
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=slot_remap.slot_logical_pages,
|
||||
layout=layout,
|
||||
start_slot=0,
|
||||
end_slot=prefix_pages,
|
||||
start_slot=prefix_start_slot,
|
||||
end_slot=prefix_end_slot,
|
||||
)
|
||||
if not materialized_by_ipc:
|
||||
materialize_local_paged_buffer_page_slots_into(
|
||||
@@ -3008,10 +3143,10 @@ def materialize_prefix_and_reuse_current_index_page_slots(
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=slot_remap.slot_logical_pages,
|
||||
layout=layout,
|
||||
start_slot=0,
|
||||
end_slot=prefix_pages,
|
||||
start_slot=prefix_start_slot,
|
||||
end_slot=prefix_end_slot,
|
||||
)
|
||||
prefix_rows = slot_range_to_page_slice(0, prefix_pages)
|
||||
prefix_rows = slot_range_to_page_slice(prefix_start_slot, prefix_end_slot)
|
||||
_all_reduce_materialized_buffer_range(
|
||||
dense_page_buffer,
|
||||
layout.cp_size,
|
||||
|
||||
@@ -15,12 +15,14 @@ from sglang.jit_kernel.fused_store_index_cache import (
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.nsa import index_buf_accessor
|
||||
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
build_batch_prefix_slot_span,
|
||||
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,
|
||||
current_extend_kv_rows_for_reuse,
|
||||
filter_owned_logical_locs,
|
||||
get_or_build_shared_paged_buffer_slot_remap,
|
||||
is_current_only_extend_batch,
|
||||
@@ -338,36 +340,55 @@ class Indexer(MultiPlatformOp):
|
||||
if extend_lens_cpu is not None
|
||||
else None
|
||||
)
|
||||
if (
|
||||
prefix_lens_cpu is None
|
||||
or len(prefix_lens_cpu) != 1
|
||||
or int(prefix_lens_cpu[0]) <= 0
|
||||
or int(prefix_lens_cpu[0]) % page_size != 0
|
||||
):
|
||||
prefix_lens_valid = prefix_lens_cpu is not None and len(prefix_lens_cpu) > 0
|
||||
if prefix_lens_valid:
|
||||
prefix_lens_valid = all(
|
||||
int(prefix_len) >= 0 and int(prefix_len) % page_size == 0
|
||||
for prefix_len in prefix_lens_cpu
|
||||
) and any(int(prefix_len) > 0 for prefix_len in prefix_lens_cpu)
|
||||
if not prefix_lens_valid:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][index_partial_current_sync] "
|
||||
"CP shared KV index partial-current compose requires one "
|
||||
"positive page-aligned prefix. "
|
||||
"CP shared KV index partial-current compose requires "
|
||||
"positive page-aligned prefix pages. "
|
||||
f"cp_rank={layout.cp_rank} layer_id={layer_id} "
|
||||
f"prefix_lens={prefix_lens} extend_lens={extend_lens} "
|
||||
f"logical_page_table_shape={tuple(logical_page_table.shape)} "
|
||||
f"page_size={page_size}"
|
||||
)
|
||||
current_locs = forward_batch.out_cache_loc
|
||||
if extend_lens_cpu is not None and len(extend_lens_cpu) == 1:
|
||||
valid_current_rows = int(extend_lens_cpu[0])
|
||||
if (
|
||||
valid_current_rows > 0
|
||||
and valid_current_rows < int(current_locs.numel())
|
||||
and valid_current_rows <= int(current_index_kv[0].shape[0])
|
||||
and valid_current_rows <= int(current_index_kv[1].shape[0])
|
||||
):
|
||||
current_locs = current_locs[:valid_current_rows]
|
||||
current_index_kv = (
|
||||
current_index_kv[0][:valid_current_rows],
|
||||
current_index_kv[1][:valid_current_rows],
|
||||
)
|
||||
prefix_pages = int(prefix_lens_cpu[0]) // page_size
|
||||
valid_current_rows = current_extend_kv_rows_for_reuse(
|
||||
forward_batch,
|
||||
current_index_kv[0],
|
||||
current_index_kv[1],
|
||||
)
|
||||
if valid_current_rows is None:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][index_partial_current_sync] "
|
||||
"CP shared KV index partial-current compose received "
|
||||
"current_index_kv that does not satisfy current reuse "
|
||||
"metadata. "
|
||||
f"cp_rank={layout.cp_rank} layer_id={layer_id} "
|
||||
f"prefix_lens={prefix_lens} extend_lens={extend_lens} "
|
||||
f"current_k_shape={tuple(current_index_kv[0].shape)} "
|
||||
f"current_scale_shape={tuple(current_index_kv[1].shape)} "
|
||||
f"out_cache_loc_shape={tuple(current_locs.shape)}"
|
||||
)
|
||||
current_locs = current_locs[:valid_current_rows]
|
||||
current_index_kv = (
|
||||
current_index_kv[0][:valid_current_rows],
|
||||
current_index_kv[1][:valid_current_rows],
|
||||
)
|
||||
prefix_slot_span = None
|
||||
if len(prefix_lens_cpu) == 1:
|
||||
prefix_pages = int(prefix_lens_cpu[0]) // page_size
|
||||
else:
|
||||
prefix_pages = 0
|
||||
prefix_slot_span = build_batch_prefix_slot_span(
|
||||
logical_pages=logical_page_table,
|
||||
prefix_lens_cpu=prefix_lens_cpu,
|
||||
page_size=page_size,
|
||||
)
|
||||
if index_prefetcher is not None:
|
||||
prefetched = index_prefetcher.consume_prefix_with_current(
|
||||
layer_id=layer_id,
|
||||
@@ -422,6 +443,7 @@ class Indexer(MultiPlatformOp):
|
||||
page_size=page_size,
|
||||
index_head_dim=forward_batch.token_to_kv_pool.index_head_dim,
|
||||
prefix_pages=prefix_pages,
|
||||
prefix_slot_span=prefix_slot_span,
|
||||
layer_id=layer_id,
|
||||
)
|
||||
)
|
||||
@@ -432,12 +454,13 @@ class Indexer(MultiPlatformOp):
|
||||
cp_shared_kv_mla_prefetch_log(
|
||||
"index_partial_current_sync_compose cp_rank=%s layer=%s "
|
||||
"prefix_lens=%s extend_lens=%s prefix_pages=%s "
|
||||
"current_rows=%s dense_pages=%s",
|
||||
"prefix_slot_span=%s current_rows=%s dense_pages=%s",
|
||||
layout.cp_rank,
|
||||
layer_id,
|
||||
prefix_lens,
|
||||
extend_lens,
|
||||
prefix_pages,
|
||||
prefix_slot_span,
|
||||
int(current_index_kv[0].shape[0]),
|
||||
int(materialized.shape[0]),
|
||||
)
|
||||
@@ -1481,12 +1504,6 @@ class Indexer(MultiPlatformOp):
|
||||
cp_metadata = forward_batch.nsa_cp_metadata
|
||||
assert cp_metadata is not None
|
||||
batch_size = int(getattr(cp_metadata, "batch_size", 1) or 1)
|
||||
if current_index_kv is not None:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][index_topk] "
|
||||
"reason=batch_gt1_index_current_reuse_unsupported "
|
||||
f"batch_size={batch_size} layer_id={layer_id}"
|
||||
)
|
||||
|
||||
request_kv_len_prev = list(getattr(cp_metadata, "request_kv_len_prev", []) or [])
|
||||
request_kv_len_next = list(getattr(cp_metadata, "request_kv_len_next", []) or [])
|
||||
@@ -1510,14 +1527,31 @@ class Indexer(MultiPlatformOp):
|
||||
f"q_prev={request_actual_seq_q_prev} q_next={request_actual_seq_q_next}"
|
||||
)
|
||||
|
||||
shared_block_tables = metadata.get_page_table_64()
|
||||
shared_index_buffer, shared_block_tables = (
|
||||
self._maybe_materialize_shared_index_buffer(
|
||||
forward_batch,
|
||||
layer_id,
|
||||
shared_block_tables,
|
||||
shared_index_buffer = None
|
||||
shared_block_tables = None
|
||||
current_index_kv_for_topk = current_index_kv
|
||||
if current_index_kv is not None and not is_current_only_extend_batch(
|
||||
forward_batch
|
||||
):
|
||||
current_index_kv_for_topk = None
|
||||
shared_block_tables = metadata.get_page_table_64()
|
||||
shared_index_buffer, shared_block_tables = (
|
||||
self._maybe_materialize_shared_index_buffer(
|
||||
forward_batch,
|
||||
layer_id,
|
||||
shared_block_tables,
|
||||
current_index_kv=current_index_kv,
|
||||
)
|
||||
)
|
||||
elif current_index_kv is None:
|
||||
shared_block_tables = metadata.get_page_table_64()
|
||||
shared_index_buffer, shared_block_tables = (
|
||||
self._maybe_materialize_shared_index_buffer(
|
||||
forward_batch,
|
||||
layer_id,
|
||||
shared_block_tables,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
outputs = []
|
||||
cursor = 0
|
||||
@@ -1558,7 +1592,7 @@ class Indexer(MultiPlatformOp):
|
||||
metadata,
|
||||
kv_len,
|
||||
segment_len,
|
||||
current_index_kv=None,
|
||||
current_index_kv=current_index_kv_for_topk,
|
||||
shared_index_buffer=shared_index_buffer,
|
||||
shared_block_tables=shared_block_tables,
|
||||
actual_seq_q_tensor=actual_seq_q_tensor,
|
||||
@@ -1943,14 +1977,8 @@ class Indexer(MultiPlatformOp):
|
||||
|
||||
current_index_kv = None
|
||||
if self._can_reuse_current_index_kv(forward_batch):
|
||||
extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None)
|
||||
valid_current_rows = int(forward_batch.out_cache_loc.numel())
|
||||
if extend_seq_lens_cpu is not None and len(extend_seq_lens_cpu) == 1:
|
||||
valid_current_rows = min(
|
||||
int(extend_seq_lens_cpu[0]),
|
||||
valid_current_rows,
|
||||
)
|
||||
if key.shape[0] >= valid_current_rows:
|
||||
valid_current_rows = current_extend_kv_rows_for_reuse(forward_batch, key)
|
||||
if valid_current_rows is not None and key.shape[0] >= valid_current_rows:
|
||||
current_k_fp8, current_k_scale = act_quant(
|
||||
key[:valid_current_rows].contiguous(),
|
||||
self.block_size,
|
||||
|
||||
@@ -102,7 +102,38 @@ def is_nsa_prefill_cp_round_robin_split():
|
||||
)
|
||||
|
||||
|
||||
def _is_cp_shared_kv_forward_batch(forward_batch: "ForwardBatch") -> bool:
|
||||
return bool(getattr(forward_batch, "uses_cp_shared_kv", False))
|
||||
|
||||
|
||||
def _fail_if_cp_shared_kv_round_robin(
|
||||
forward_batch: "ForwardBatch",
|
||||
*,
|
||||
op: str,
|
||||
) -> None:
|
||||
if forward_batch is None or not _is_cp_shared_kv_forward_batch(forward_batch):
|
||||
return
|
||||
try:
|
||||
round_robin_split = is_nsa_prefill_cp_round_robin_split()
|
||||
except ValueError:
|
||||
round_robin_split = False
|
||||
if not round_robin_split:
|
||||
return
|
||||
|
||||
error_msg = (
|
||||
"[CP_SHARED_KV_FAIL_FAST][round_robin_unsupported] "
|
||||
"CP shared KV only supports in-seq zigzag split. "
|
||||
f"op={op} nsa_prefill_cp_mode=round-robin-split"
|
||||
)
|
||||
logger.error(error_msg)
|
||||
raise RuntimeError(error_msg)
|
||||
|
||||
|
||||
def can_nsa_prefill_cp_round_robin_split(forward_batch: "ForwardBatch"):
|
||||
_fail_if_cp_shared_kv_round_robin(
|
||||
forward_batch,
|
||||
op="can_nsa_prefill_cp_round_robin_split",
|
||||
)
|
||||
if not forward_batch.forward_mode.is_context_parallel_extend():
|
||||
return False
|
||||
cp_size = get_attention_cp_size()
|
||||
@@ -934,6 +965,7 @@ def should_skip_cp_shared_kv_cp_split_for_short_page_extent(
|
||||
|
||||
|
||||
def can_cp_split(seq_len: int, cp_size: int, use_nsa: bool, forward_batch):
|
||||
_fail_if_cp_shared_kv_round_robin(forward_batch, op="can_cp_split")
|
||||
if is_nsa_prefill_cp_round_robin_split():
|
||||
cur_cp_seq_len = seq_len // cp_size
|
||||
assert seq_len % cp_size == 0, (
|
||||
@@ -964,6 +996,7 @@ def can_cp_split(seq_len: int, cp_size: int, use_nsa: bool, forward_batch):
|
||||
|
||||
|
||||
def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor):
|
||||
_fail_if_cp_shared_kv_round_robin(forward_batch, op="cp_split_and_rebuild_data")
|
||||
if is_nsa_prefill_cp_round_robin_split():
|
||||
cp_size = get_attention_cp_size()
|
||||
assert input_.shape[0] % cp_size == 0, (
|
||||
@@ -985,6 +1018,7 @@ def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor):
|
||||
|
||||
|
||||
def cp_split_and_rebuild_1d(forward_batch, input_: torch.Tensor):
|
||||
_fail_if_cp_shared_kv_round_robin(forward_batch, op="cp_split_and_rebuild_1d")
|
||||
try:
|
||||
round_robin_split = is_nsa_prefill_cp_round_robin_split()
|
||||
except ValueError:
|
||||
@@ -1188,6 +1222,7 @@ def get_cp_shared_kv_local_physical_out_cache_loc(forward_batch: "ForwardBatch")
|
||||
|
||||
|
||||
def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor):
|
||||
_fail_if_cp_shared_kv_round_robin(forward_batch, op="cp_split_and_rebuild_position")
|
||||
if is_nsa_prefill_cp_round_robin_split():
|
||||
cp_size = get_attention_cp_size()
|
||||
assert positions.shape[0] % cp_size == 0, (
|
||||
@@ -1327,13 +1362,19 @@ def _cp_attn_tp_all_gather_padded_tensor(
|
||||
max_len = (total_len + attn_tp_size - 1) // attn_tp_size
|
||||
pad_size = max_len - input_.shape[0]
|
||||
if pad_size > 0:
|
||||
input_ = F.pad(input_, (0, 0, 0, pad_size), mode="constant", value=0)
|
||||
input_ = torch.cat(
|
||||
(
|
||||
input_,
|
||||
input_.new_zeros((pad_size, *input_.shape[1:])),
|
||||
),
|
||||
dim=0,
|
||||
)
|
||||
with use_symmetric_memory(
|
||||
get_attention_cp_group(), disabled=not is_allocation_symmetric()
|
||||
):
|
||||
input_tensor_all = torch.empty(
|
||||
max_len * attn_tp_size,
|
||||
input_.shape[1],
|
||||
*input_.shape[1:],
|
||||
device=input_.device,
|
||||
dtype=input_.dtype,
|
||||
)
|
||||
@@ -1465,6 +1506,140 @@ def _torch_in_seq_all_gather_rerange(
|
||||
return output_tensor.view(-1, hidden_size)
|
||||
|
||||
|
||||
def _raise_batch_rerange_error(reason: str, message: str, *args) -> None:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][batch_gt1_full_rerange_unsupported] "
|
||||
+ f"reason={reason} "
|
||||
+ (message % args if args else message)
|
||||
)
|
||||
|
||||
|
||||
def _torch_batch_in_seq_all_gather_rerange(
|
||||
input_tensor_all: torch.Tensor,
|
||||
forward_batch: "ForwardBatch",
|
||||
*,
|
||||
cp_size: int,
|
||||
) -> torch.Tensor:
|
||||
"""Restore rank-major batched in-seq CP all-gather rows to request order.
|
||||
|
||||
The row payload is intentionally opaque. MLA bf16 latent rows, packed fp8
|
||||
KV rows, and NSA index byte rows all share the same first-dimension token
|
||||
order contract, so this helper only slices/copies rows and never interprets
|
||||
the dtype or payload shape.
|
||||
"""
|
||||
|
||||
metadata = getattr(forward_batch, "nsa_cp_metadata", None)
|
||||
batch_size = int(getattr(metadata, "batch_size", 1) or 1)
|
||||
request_split_lists = getattr(metadata, "request_split_lists", None)
|
||||
max_rank_len = getattr(metadata, "max_rank_len", None)
|
||||
if metadata is None:
|
||||
_raise_batch_rerange_error("missing_metadata", "nsa_cp_metadata is missing")
|
||||
if batch_size <= 1:
|
||||
_raise_batch_rerange_error(
|
||||
"not_batch",
|
||||
"batch-aware rerange requires batch_size > 1, got %s",
|
||||
batch_size,
|
||||
)
|
||||
if request_split_lists is None or len(request_split_lists) != batch_size:
|
||||
_raise_batch_rerange_error(
|
||||
"missing_request_split_lists",
|
||||
"request_split_lists is missing or incomplete. batch_size=%s value=%s",
|
||||
batch_size,
|
||||
request_split_lists,
|
||||
)
|
||||
if max_rank_len is None or len(max_rank_len) < cp_size:
|
||||
_raise_batch_rerange_error(
|
||||
"missing_max_rank_len",
|
||||
"max_rank_len is missing or shorter than cp_size. cp_size=%s value=%s",
|
||||
cp_size,
|
||||
max_rank_len,
|
||||
)
|
||||
if cp_size <= 0:
|
||||
_raise_batch_rerange_error("bad_cp_size", "cp_size must be positive: %s", cp_size)
|
||||
|
||||
split_lists: List[List[int]] = []
|
||||
for req_id, split_list in enumerate(request_split_lists):
|
||||
if split_list is None or len(split_list) != cp_size * 2:
|
||||
_raise_batch_rerange_error(
|
||||
"bad_request_split",
|
||||
"request split must have 2 * cp_size entries. req_id=%s "
|
||||
"cp_size=%s split=%s",
|
||||
req_id,
|
||||
cp_size,
|
||||
split_list,
|
||||
)
|
||||
split_lists.append([int(x) for x in split_list])
|
||||
|
||||
max_rank_token = int(max_rank_len[0])
|
||||
if max_rank_token < 0:
|
||||
_raise_batch_rerange_error(
|
||||
"bad_max_rank_token",
|
||||
"max_rank_len[0] must be non-negative, got %s",
|
||||
max_rank_token,
|
||||
)
|
||||
required_rows = max_rank_token * cp_size
|
||||
if int(input_tensor_all.shape[0]) < required_rows:
|
||||
_raise_batch_rerange_error(
|
||||
"input_rows_short",
|
||||
"input rows=%s required=%s max_rank_token=%s cp_size=%s",
|
||||
int(input_tensor_all.shape[0]),
|
||||
required_rows,
|
||||
max_rank_token,
|
||||
cp_size,
|
||||
)
|
||||
|
||||
rank_request_offsets: List[List[int]] = []
|
||||
for source_rank in range(cp_size):
|
||||
mirror = cp_size * 2 - source_rank - 1
|
||||
offsets: List[int] = []
|
||||
cursor = 0
|
||||
for split_list in split_lists:
|
||||
offsets.append(cursor)
|
||||
cursor += split_list[source_rank] + split_list[mirror]
|
||||
if cursor > max_rank_token:
|
||||
_raise_batch_rerange_error(
|
||||
"rank_payload_exceeds_max",
|
||||
"rank local rows exceed max_rank_token. rank=%s rows=%s "
|
||||
"max_rank_token=%s",
|
||||
source_rank,
|
||||
cursor,
|
||||
max_rank_token,
|
||||
)
|
||||
rank_request_offsets.append(offsets)
|
||||
|
||||
total_tokens = sum(sum(split_list) for split_list in split_lists)
|
||||
output_tensor = input_tensor_all.new_empty(
|
||||
(total_tokens, *input_tensor_all.shape[1:])
|
||||
)
|
||||
if total_tokens == 0:
|
||||
return output_tensor
|
||||
|
||||
output_request_base = 0
|
||||
for req_id, split_list in enumerate(split_lists):
|
||||
segment_prefix = [0] + list(accumulate(split_list))[:-1]
|
||||
for segment_id, segment_len in enumerate(split_list):
|
||||
if segment_len <= 0:
|
||||
continue
|
||||
if segment_id < cp_size:
|
||||
source_rank = segment_id
|
||||
source_segment_offset = 0
|
||||
else:
|
||||
source_rank = cp_size * 2 - segment_id - 1
|
||||
source_segment_offset = split_list[source_rank]
|
||||
source_start = (
|
||||
source_rank * max_rank_token
|
||||
+ rank_request_offsets[source_rank][req_id]
|
||||
+ source_segment_offset
|
||||
)
|
||||
output_start = output_request_base + segment_prefix[segment_id]
|
||||
output_tensor[output_start : output_start + segment_len].copy_(
|
||||
input_tensor_all[source_start : source_start + segment_len]
|
||||
)
|
||||
output_request_base += sum(split_list)
|
||||
|
||||
return output_tensor
|
||||
|
||||
|
||||
def cp_all_gather_rerange_output(input_tensor, cp_size, forward_batch, stream):
|
||||
"""
|
||||
# for in-seq-split
|
||||
@@ -1492,14 +1667,47 @@ def cp_all_gather_rerange_output(input_tensor, cp_size, forward_batch, stream):
|
||||
| token0, token1, token2, token3, token4, token5, token6, token7, ...
|
||||
| +-------------------------+
|
||||
"""
|
||||
_fail_if_cp_shared_kv_round_robin(
|
||||
forward_batch,
|
||||
op="cp_all_gather_rerange_output",
|
||||
)
|
||||
metadata = getattr(forward_batch, "nsa_cp_metadata", None)
|
||||
if getattr(metadata, "batch_size", 1) > 1:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][batch_gt1_full_rerange_unsupported] "
|
||||
"CP shared-KV bs>1 must not use scalar full hidden/KV rerange. "
|
||||
"Use batch-aware narrow output collection or add batch-aware full "
|
||||
"rerange metadata/kernels for this consumer."
|
||||
batch_size = int(getattr(metadata, "batch_size", 1) or 1)
|
||||
if batch_size > 1:
|
||||
if input_tensor.dim() < 2:
|
||||
_raise_batch_rerange_error(
|
||||
"bad_input_dim",
|
||||
"input tensor must have token rows plus payload dims, got shape=%s",
|
||||
tuple(input_tensor.shape),
|
||||
)
|
||||
if getattr(metadata, "request_split_lists", None) is None:
|
||||
_raise_batch_rerange_error(
|
||||
"missing_request_split_lists",
|
||||
"request_split_lists is required for batch-aware current rerange. "
|
||||
"batch_size=%s",
|
||||
batch_size,
|
||||
)
|
||||
total_seq_lens = getattr(metadata, "total_seq_lens", None)
|
||||
if total_seq_lens is None:
|
||||
_raise_batch_rerange_error(
|
||||
"missing_total_seq_lens",
|
||||
"total_seq_lens is required for batch-aware current rerange. "
|
||||
"batch_size=%s",
|
||||
batch_size,
|
||||
)
|
||||
input_tensor_all = _cp_attn_tp_all_gather_padded_tensor(
|
||||
input_tensor,
|
||||
total_seq_lens,
|
||||
cp_size,
|
||||
forward_batch,
|
||||
stream,
|
||||
)
|
||||
return _torch_batch_in_seq_all_gather_rerange(
|
||||
input_tensor_all,
|
||||
forward_batch,
|
||||
cp_size=cp_size,
|
||||
)
|
||||
|
||||
if is_nsa_prefill_cp_round_robin_split():
|
||||
with use_symmetric_memory(
|
||||
get_attention_cp_group(), disabled=not is_allocation_symmetric()
|
||||
@@ -1599,6 +1807,10 @@ def prepare_input_dp_with_cp_dsa(
|
||||
forward_batch: "ForwardBatch" = None,
|
||||
page_size: int = None,
|
||||
):
|
||||
_fail_if_cp_shared_kv_round_robin(
|
||||
forward_batch,
|
||||
op="prepare_input_dp_with_cp_dsa",
|
||||
)
|
||||
if is_nsa_prefill_cp_round_robin_split():
|
||||
return True
|
||||
"""prepare_input_dp_with_cp_dsa-zigzag index
|
||||
@@ -1769,6 +1981,10 @@ def cp_collect_last_token_hidden(
|
||||
forward_batch: "ForwardBatch",
|
||||
cp_size: int,
|
||||
) -> torch.Tensor:
|
||||
_fail_if_cp_shared_kv_round_robin(
|
||||
forward_batch,
|
||||
op="cp_collect_last_token_hidden",
|
||||
)
|
||||
if is_nsa_prefill_cp_round_robin_split():
|
||||
return _round_robin_collect_last_token(hidden_states, forward_batch, cp_size)
|
||||
return _in_seq_collect_last_token(hidden_states, forward_batch, cp_size)
|
||||
|
||||
@@ -15,6 +15,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_prefetch import (
|
||||
CpSharedKVMlaPrefetcher,
|
||||
)
|
||||
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
|
||||
build_batch_prefix_slot_span,
|
||||
build_current_loc_remap,
|
||||
cp_shared_kv_debug_enabled,
|
||||
cp_shared_kv_debug_log,
|
||||
@@ -1987,16 +1988,18 @@ class NativeSparseAttnBackend(
|
||||
int(current_kv_cache.shape[0]),
|
||||
tuple(logical_page_table_1.shape),
|
||||
)
|
||||
if (
|
||||
prefix_lens_cpu is None
|
||||
or len(prefix_lens_cpu) != 1
|
||||
or int(prefix_lens_cpu[0]) <= 0
|
||||
or int(prefix_lens_cpu[0]) % page_size != 0
|
||||
):
|
||||
prefix_lens_valid = prefix_lens_cpu is not None and len(prefix_lens_cpu) > 0
|
||||
if prefix_lens_valid:
|
||||
prefix_lens_valid = all(
|
||||
int(prefix_len) >= 0
|
||||
and int(prefix_len) % page_size == 0
|
||||
for prefix_len in prefix_lens_cpu
|
||||
) and any(int(prefix_len) > 0 for prefix_len in prefix_lens_cpu)
|
||||
if not prefix_lens_valid:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][mla_partial_current_sync] "
|
||||
"CP shared KV MLA partial-current sync compose "
|
||||
"requires one positive page-aligned prefix. "
|
||||
"requires positive page-aligned prefix pages. "
|
||||
f"reason={reason} "
|
||||
f"cp_rank={forward_batch.cp_shared_kv_layout.cp_rank} "
|
||||
f"layer_id={layer.layer_id} "
|
||||
@@ -2007,7 +2010,16 @@ class NativeSparseAttnBackend(
|
||||
f"current_locs_shape={tuple(current_locs_for_reuse.shape)} "
|
||||
f"page_size={page_size}"
|
||||
)
|
||||
prefix_pages = int(prefix_lens_cpu[0]) // page_size
|
||||
prefix_slot_span = None
|
||||
if len(prefix_lens_cpu) == 1:
|
||||
prefix_pages = int(prefix_lens_cpu[0]) // page_size
|
||||
else:
|
||||
prefix_pages = 0
|
||||
prefix_slot_span = build_batch_prefix_slot_span(
|
||||
logical_pages=metadata.real_page_table,
|
||||
prefix_lens_cpu=prefix_lens_cpu,
|
||||
page_size=page_size,
|
||||
)
|
||||
slot_remap = get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
@@ -2025,6 +2037,7 @@ class NativeSparseAttnBackend(
|
||||
layout=forward_batch.cp_shared_kv_layout,
|
||||
page_size=page_size,
|
||||
prefix_pages=prefix_pages,
|
||||
prefix_slot_span=prefix_slot_span,
|
||||
layer_id=layer.layer_id,
|
||||
)
|
||||
)
|
||||
@@ -2037,14 +2050,15 @@ class NativeSparseAttnBackend(
|
||||
cp_shared_kv_mla_prefetch_log(
|
||||
"forward_partial_current_sync_compose cp_rank=%s "
|
||||
"layer=%s reason=%s prefix_lens=%s extend_lens=%s "
|
||||
"prefix_pages=%s current_rows=%s kv_rows=%s "
|
||||
"page_table_shape=%s",
|
||||
"prefix_pages=%s prefix_slot_span=%s "
|
||||
"current_rows=%s kv_rows=%s page_table_shape=%s",
|
||||
forward_batch.cp_shared_kv_layout.cp_rank,
|
||||
layer.layer_id,
|
||||
reason,
|
||||
prefix_lens,
|
||||
extend_lens,
|
||||
prefix_pages,
|
||||
prefix_slot_span,
|
||||
int(current_kv_cache.shape[0]),
|
||||
int(kv_cache.shape[0]),
|
||||
tuple(page_table_1.shape),
|
||||
|
||||
Reference in New Issue
Block a user