From e824b242507c7fd52e3a1c67aaf3590fc4021e0c Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sun, 2 Aug 2026 20:46:42 -0700 Subject: [PATCH] [CP] Fuse zigzag attention into a single call (#33137) --- .../layers/attention/trtllm_mha_backend.py | 35 +++++- python/sglang/srt/layers/cp/zigzag.py | 67 +++++++++--- test/registered/cp/test_cp_strategy_unit.py | 102 ++++++++++++++++++ 3 files changed, 186 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index 67db8a14c..b09f22664 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -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, diff --git a/python/sglang/srt/layers/cp/zigzag.py b/python/sglang/srt/layers/cp/zigzag.py index 6ad74e186..ce2da4e97 100644 --- a/python/sglang/srt/layers/cp/zigzag.py +++ b/python/sglang/srt/layers/cp/zigzag.py @@ -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: diff --git a/test/registered/cp/test_cp_strategy_unit.py b/test/registered/cp/test_cp_strategy_unit.py index d586416f4..392f7cde0 100644 --- a/test/registered/cp/test_cp_strategy_unit.py +++ b/test/registered/cp/test_cp_strategy_unit.py @@ -547,6 +547,108 @@ class TestCPZigzagStrategy(CustomTestCase): self.assertTrue(torch.equal(calls[1][0], q[2:])) self.assertTrue(torch.equal(out, q + 100)) + def test_zigzag_combined_attention_matches_two_half_reference(self): + def reference_attention( + q, + cu_seqlens_q, + cache_seqlens, + cu_seqlens_kv, + *, + sequence_offset, + ): + outputs = [] + for seq_id in range(cache_seqlens.numel()): + q_start = int(cu_seqlens_q[seq_id]) + q_end = int(cu_seqlens_q[seq_id + 1]) + q_seq = q[q_start:q_end] + q_len = q_end - q_start + kv_len = int(cache_seqlens[seq_id]) + self.assertEqual( + int(cu_seqlens_kv[seq_id + 1] - cu_seqlens_kv[seq_id]), + kv_len, + ) + + absolute_seq_id = sequence_offset + seq_id + 1 + positions = torch.arange(kv_len, dtype=q.dtype) + k_seq = torch.stack( + ( + positions / (kv_len + 1), + torch.sin(positions + absolute_seq_id), + torch.full_like(positions, absolute_seq_id / 10), + ), + dim=1, + ) + v_seq = torch.stack( + ( + torch.cos(positions + absolute_seq_id), + positions / (absolute_seq_id + 1), + torch.full_like(positions, absolute_seq_id), + ), + dim=1, + ) + q_positions = kv_len - q_len + torch.arange(q_len) + allowed = torch.arange(kv_len)[None, :] <= q_positions[:, None] + scores = q_seq @ k_seq.T / q.shape[-1] ** 0.5 + outputs.append( + torch.softmax(scores.masked_fill(~allowed, -torch.inf), dim=-1) + @ v_seq + ) + return torch.cat(outputs, dim=0) + + cp_size = 4 + seq_lens = [19, 27] + extend_seq_lens = [11, 13] + for rank in range(cp_size): + with self.subTest(rank=rank): + metadata = self._metadata_for_rank( + rank, + cp_size=cp_size, + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + ) + logical_tokens = ( + metadata.total_q_prev_tokens + metadata.total_q_next_tokens + ) + q = torch.linspace( + -0.75, + 0.75, + steps=logical_tokens * 3, + dtype=torch.float32, + ).view(logical_tokens, 3) + q_prev = q[: metadata.total_q_prev_tokens] + q_next = q[metadata.total_q_prev_tokens :] + + two_half_out = torch.cat( + ( + reference_attention( + q_prev, + metadata.cu_seqlens_q_prev_tensor, + metadata.kv_len_prev_tensor, + metadata.cu_seqlens_kv_prev_tensor, + sequence_offset=0, + ), + reference_attention( + q_next, + metadata.cu_seqlens_q_next_tensor, + metadata.kv_len_next_tensor, + metadata.cu_seqlens_kv_next_tensor, + sequence_offset=metadata.bs, + ), + ), + dim=0, + ) + combined_out = reference_attention( + q, + metadata.cu_seqlens_q_combined_tensor, + metadata.kv_len_combined_tensor, + metadata.cu_seqlens_kv_combined_tensor, + sequence_offset=0, + ) + + torch.testing.assert_close( + combined_out, two_half_out, atol=1e-5, rtol=1e-5 + ) + class TestCPInterleaveStrategy(CustomTestCase): def setUp(self):