Preserve request boundaries in CP shared-KV index top-k

W4-1 needs target index/top-k sync correctness before current/partial-current reuse can be restored. Batch-size>1 in-seq CP produces local q/weights in request-segment order, so top-k must consume req0_prev, req0_next, req1_prev, req1_next rather than treating the flattened batch as one scalar prev/next pair.

The implementation adds a batch dispatch for _get_topk_in_seq_cp_pair, reuses one synchronous shared-index materialization per layer, and calls _get_topk_ragged_with_cp per request segment with an explicit batch_idx. The scalar bs=1 path remains unchanged.

Constraint: This is W4-1 target index/top-k sync correctness; original W4 current/partial-current reuse remains a separate follow-up.

Constraint: Phase W4-1 must not enable bs>1 index prefetch, current reuse, partial-current reuse, or the cp_index multi-batch branch.

Rejected: Use cp_index branch for multi-batch | source marks that path as having accuracy issues.

Rejected: Pad batch requests to max length | wastes compute and violates packed/ragged batch contract.

Confidence: high

Scope-risk: moderate

Directive: Keep bs>1 target top-k ordered by request segment unless a later fused descriptor proves identical ordering and correctness.

Tested: Remote g0034 py_compile for nsa_indexer.py

Tested: Remote g0034 PYTHONPATH=python pytest test/registered/unit/layers/test_nsa_cp_utils.py -> 45 passed

Tested: Remote g0034 PYTHONPATH=python pytest test_nsa_cp_utils.py test_cp_shared_kv_layout.py test_cp_shared_kv_runtime.py -> 172 passed, 2 subtests passed

Not-tested: Full ETE bs>1 serving run with live traffic
This commit is contained in:
laoyao0822
2026-06-03 02:36:31 +08:00
parent f8b4f1915e
commit a7472c415f
4 changed files with 329 additions and 41 deletions

View File

@@ -1099,6 +1099,7 @@ class Indexer(MultiPlatformOp):
shared_block_tables: Optional[torch.Tensor] = None,
actual_seq_q_tensor: Optional[torch.Tensor] = None,
actual_seq_q_cu_tensor: Optional[torch.Tensor] = None,
batch_idx: int = 0,
) -> torch.Tensor:
if TYPE_CHECKING:
assert isinstance(forward_batch.token_to_kv_pool, NSATokenToKVPool)
@@ -1221,14 +1222,16 @@ class Indexer(MultiPlatformOp):
batch_idx_list=batch_idx_list,
)
else:
seq_len = int(forward_batch.seq_lens_cpu[batch_idx].item())
extend_seq_len = int(forward_batch.extend_seq_lens_cpu[batch_idx])
cp_kv_end = (
forward_batch.seq_lens_cpu[0].item()
- forward_batch.extend_seq_lens_cpu[0]
seq_len
- extend_seq_len
+ kv_len
)
page_table_1 = metadata.get_page_table_1()
logical_kv_limit = min(
int(forward_batch.seq_lens_cpu[0].item()),
seq_len,
int(page_table_1.shape[1]),
)
valid_q_count = _compute_contiguous_valid_cp_query_count(
@@ -1253,7 +1256,7 @@ class Indexer(MultiPlatformOp):
assert index_buffer is not None
tai_prepared = try_tai_prepare_cp_mqa_index(
index_buffer=index_buffer,
page_indices=block_tables[0],
page_indices=block_tables[batch_idx],
kv_len=kv_len,
valid_q_count=valid_q_count,
ke_start=ke_start,
@@ -1283,13 +1286,13 @@ class Indexer(MultiPlatformOp):
forward_batch.token_to_kv_pool,
index_buffer,
seq_len=kv_len,
page_indices=block_tables[0],
page_indices=block_tables[batch_idx],
)
k_scale = index_buf_accessor.GetS.execute(
forward_batch.token_to_kv_pool,
index_buffer,
seq_len=kv_len,
page_indices=block_tables[0],
page_indices=block_tables[batch_idx],
)
k_fp8 = k_fp8.view(torch.float8_e4m3fn)
@@ -1378,6 +1381,16 @@ class Indexer(MultiPlatformOp):
current_index_kv: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
assert forward_batch.nsa_cp_metadata is not None
if int(getattr(forward_batch.nsa_cp_metadata, "batch_size", 1) or 1) > 1:
return self._get_topk_in_seq_cp_pair_batch(
forward_batch,
layer_id,
q_fp8,
weights,
metadata,
current_index_kv=current_index_kv,
)
kv_len_prev = forward_batch.nsa_cp_metadata.kv_len_prev
kv_len_next = forward_batch.nsa_cp_metadata.kv_len_next
actual_seq_q_prev = forward_batch.nsa_cp_metadata.actual_seq_q_prev
@@ -1456,6 +1469,131 @@ class Indexer(MultiPlatformOp):
)
return torch.cat([topk_result_prev, topk_result_next], dim=0)
def _get_topk_in_seq_cp_pair_batch(
self,
forward_batch: ForwardBatch,
layer_id: int,
q_fp8: torch.Tensor,
weights: torch.Tensor,
metadata: BaseIndexerMetadata,
current_index_kv: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
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 [])
request_actual_seq_q_prev = list(
getattr(cp_metadata, "request_actual_seq_q_prev", []) or []
)
request_actual_seq_q_next = list(
getattr(cp_metadata, "request_actual_seq_q_next", []) or []
)
if not (
len(request_kv_len_prev) == batch_size
and len(request_kv_len_next) == batch_size
and len(request_actual_seq_q_prev) == batch_size
and len(request_actual_seq_q_next) == batch_size
):
raise RuntimeError(
"[CP_SHARED_KV_FAIL_FAST][index_topk] "
"reason=batch_gt1_index_metadata_incomplete "
f"batch_size={batch_size} layer_id={layer_id} "
f"kv_prev={request_kv_len_prev} kv_next={request_kv_len_next} "
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,
)
)
outputs = []
cursor = 0
def call_segment(
*,
req_id: int,
segment_len: int,
kv_len: int,
) -> torch.Tensor:
nonlocal cursor
segment_len = int(segment_len)
kv_len = int(kv_len)
q_segment = q_fp8[cursor : cursor + segment_len]
weights_segment = weights[cursor : cursor + segment_len]
cursor += segment_len
if segment_len == 0:
return torch.empty(
(0, self.index_topk),
dtype=torch.int32,
device=q_fp8.device,
)
actual_seq_q_tensor = torch.tensor(
[segment_len],
dtype=torch.int32,
device=q_fp8.device,
)
actual_seq_q_cu_tensor = torch.tensor(
[0, segment_len],
dtype=torch.int32,
device=q_fp8.device,
)
return self._get_topk_ragged_with_cp(
forward_batch,
layer_id,
q_segment,
weights_segment,
metadata,
kv_len,
segment_len,
current_index_kv=None,
shared_index_buffer=shared_index_buffer,
shared_block_tables=shared_block_tables,
actual_seq_q_tensor=actual_seq_q_tensor,
actual_seq_q_cu_tensor=actual_seq_q_cu_tensor,
batch_idx=req_id,
)
for req_id in range(batch_size):
outputs.append(
call_segment(
req_id=req_id,
segment_len=request_actual_seq_q_prev[req_id],
kv_len=request_kv_len_prev[req_id],
)
)
outputs.append(
call_segment(
req_id=req_id,
segment_len=request_actual_seq_q_next[req_id],
kv_len=request_kv_len_next[req_id],
)
)
if cursor != int(q_fp8.shape[0]) or cursor != int(weights.shape[0]):
raise RuntimeError(
"[CP_SHARED_KV_FAIL_FAST][index_topk] "
"reason=batch_gt1_index_q_length_mismatch "
f"batch_size={batch_size} layer_id={layer_id} cursor={cursor} "
f"q_tokens={int(q_fp8.shape[0])} weights_tokens={int(weights.shape[0])}"
)
if not outputs:
return torch.empty((0, self.index_topk), dtype=torch.int32, device=q_fp8.device)
return torch.cat(outputs, dim=0)
def forward_indexer(
self,
q_fp8: torch.Tensor,