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:
laoyao0822
2026-06-03 07:06:21 +08:00
parent a7472c415f
commit 50d0008705
9 changed files with 1195 additions and 152 deletions

View File

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

View File

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

View File

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

View File

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