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