[CP] Fuse zigzag attention into a single call (#33137)

This commit is contained in:
Baizhou Zhang
2026-08-02 20:46:42 -07:00
committed by GitHub
parent 741e33db81
commit e824b24250
3 changed files with 186 additions and 18 deletions
@@ -75,6 +75,9 @@ class TRTLLMMHAMetadata:
page_table: torch.Tensor = None
# Page table for SWA layers (translated from full pool indices to SWA pool indices)
swa_page_table: torch.Tensor = None
# CP-v2 zigzag treats prev/next halves as a synthetic 2 * batch_size batch.
zigzag_page_table: torch.Tensor = None
zigzag_swa_page_table: torch.Tensor = None
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: torch.Tensor = None
is_ragged_verify: bool = False
@@ -339,6 +342,25 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
return swa_pt
return self.forward_metadata.page_table
def _maybe_build_cp_zigzag_page_tables(
self,
metadata: TRTLLMMHAMetadata,
forward_batch: ForwardBatch,
) -> None:
"""Duplicate request rows once for the combined prev-then-next CP launch."""
if not is_cp_v2_active(forward_batch):
return
# TODO: Avoid materializing duplicated page tables to reduce zigzag CP
# page-table memory usage.
metadata.zigzag_page_table = torch.cat(
(metadata.page_table, metadata.page_table), dim=0
)
if metadata.swa_page_table is not None:
metadata.zigzag_swa_page_table = torch.cat(
(metadata.swa_page_table, metadata.swa_page_table), dim=0
)
@staticmethod
def _get_scalar_scale(
layer: RadixAttention,
@@ -962,6 +984,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
self._fill_page_table_device(
metadata, forward_batch.req_pool_indices, metadata.cache_seqlens_int32
)
self._maybe_build_cp_zigzag_page_tables(metadata, forward_batch)
if self._needs_encoder_only_expand(forward_batch.forward_mode, metadata):
row_map = (
@@ -1276,13 +1299,23 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
cu_seqlens_q,
cache_seqlens,
max_seqlen_q,
*,
cu_seqlens_kv,
use_zigzag_page_table=False,
):
block_tables = page_table
if use_zigzag_page_table:
block_tables = self.forward_metadata.zigzag_page_table
zigzag_swa_pt = self.forward_metadata.zigzag_swa_page_table
if zigzag_swa_pt is not None:
_, is_swa = self._swa_kv_pool.layers_mapping[layer.layer_id]
if is_swa:
block_tables = zigzag_swa_pt
return flashinfer.prefill.trtllm_batch_context_with_kv_cache(
query=q_chunk,
kv_cache=kv_cache,
workspace_buffer=self.workspace_buffer,
block_tables=page_table,
block_tables=block_tables,
seq_lens=cache_seqlens,
max_q_len=max_seqlen_q,
max_kv_len=self.max_context_len,
+50 -17
View File
@@ -79,11 +79,18 @@ class ZigzagContextParallelMetadata(BaseContextParallelMetadata):
cu_seqlens_q_prev_tensor: Optional[Any] = None
cu_seqlens_q_next_tensor: Optional[Any] = None
# Combined prev-then-next TRT-LLM geometry (shape [2 * bs] or [2 * bs + 1]).
actual_seq_q_combined_tensor: Optional[Any] = None
kv_len_combined_tensor: Optional[Any] = None
cu_seqlens_q_combined_tensor: Optional[Any] = None
cu_seqlens_kv_combined_tensor: Optional[Any] = None
# Scalars derived from the per-sequence lists above.
total_q_prev_tokens: int = 0
total_q_next_tokens: int = 0
max_seqlen_q_prev: int = 0
max_seqlen_q_next: int = 0
max_seqlen_q_combined: int = 0
# Per-sequence CPU lists, useful for indexers and diagnostics.
kv_len_prev_list: Optional[List[int]] = None
@@ -216,6 +223,10 @@ class ZigzagCPStrategy(ContextParallelStrategy):
cu_next = [0] + list(accumulate(actual_seq_q_next_list))
cu_kv_prev = [0] + list(accumulate(kv_len_prev_list))
cu_kv_next = [0] + list(accumulate(kv_len_next_list))
actual_seq_q_combined_list = actual_seq_q_prev_list + actual_seq_q_next_list
kv_len_combined_list = kv_len_prev_list + kv_len_next_list
cu_q_combined = [0] + list(accumulate(actual_seq_q_combined_list))
cu_kv_combined = [0] + list(accumulate(kv_len_combined_list))
total_seq_lens = sum(extend_seqs_len)
assert len(split_list) == bs * cp_segment_num
@@ -256,6 +267,18 @@ class ZigzagCPStrategy(ContextParallelStrategy):
cu_seqlens_q_next_tensor=torch.tensor(
cu_next, device=device, dtype=torch.int32
),
actual_seq_q_combined_tensor=torch.tensor(
actual_seq_q_combined_list, device=device, dtype=torch.int32
),
kv_len_combined_tensor=torch.tensor(
kv_len_combined_list, device=device, dtype=torch.int32
),
cu_seqlens_q_combined_tensor=torch.tensor(
cu_q_combined, device=device, dtype=torch.int32
),
cu_seqlens_kv_combined_tensor=torch.tensor(
cu_kv_combined, device=device, dtype=torch.int32
),
total_q_prev_tokens=cu_prev[-1],
total_q_next_tokens=cu_next[-1],
max_seqlen_q_prev=(
@@ -264,6 +287,9 @@ class ZigzagCPStrategy(ContextParallelStrategy):
max_seqlen_q_next=(
max(actual_seq_q_next_list) if actual_seq_q_next_list else 0
),
max_seqlen_q_combined=(
max(actual_seq_q_combined_list) if actual_seq_q_combined_list else 0
),
kv_len_prev_list=kv_len_prev_list,
kv_len_next_list=kv_len_next_list,
actual_seq_q_prev_list=actual_seq_q_prev_list,
@@ -334,24 +360,31 @@ class ZigzagCPStrategy(ContextParallelStrategy):
prev_kwargs = {}
next_kwargs = {}
if attention_backend == CPAttentionBackendKind.TRTLLM_MHA:
prev_kwargs["cu_seqlens_kv"] = meta.cu_seqlens_kv_prev_tensor
next_kwargs["cu_seqlens_kv"] = meta.cu_seqlens_kv_next_tensor
result = attn_fn(
q[:logical_tokens],
meta.cu_seqlens_q_combined_tensor,
meta.kv_len_combined_tensor,
meta.max_seqlen_q_combined,
cu_seqlens_kv=meta.cu_seqlens_kv_combined_tensor,
use_zigzag_page_table=True,
)
else:
result_prev = attn_fn(
q_prev,
meta.cu_seqlens_q_prev_tensor,
meta.kv_len_prev_tensor,
meta.max_seqlen_q_prev,
**prev_kwargs,
)
result_next = attn_fn(
q_next,
meta.cu_seqlens_q_next_tensor,
meta.kv_len_next_tensor,
meta.max_seqlen_q_next,
**next_kwargs,
)
result = torch.cat([result_prev, result_next], dim=0)
result_prev = attn_fn(
q_prev,
meta.cu_seqlens_q_prev_tensor,
meta.kv_len_prev_tensor,
meta.max_seqlen_q_prev,
**prev_kwargs,
)
result_next = attn_fn(
q_next,
meta.cu_seqlens_q_next_tensor,
meta.kv_len_next_tensor,
meta.max_seqlen_q_next,
**next_kwargs,
)
result = torch.cat([result_prev, result_next], dim=0)
pad_size = q.shape[0] - logical_tokens
assert pad_size >= 0
if pad_size > 0: