[CP] Fuse zigzag attention into a single call (#33137)
This commit is contained in:
@@ -75,6 +75,9 @@ class TRTLLMMHAMetadata:
|
|||||||
page_table: torch.Tensor = None
|
page_table: torch.Tensor = None
|
||||||
# Page table for SWA layers (translated from full pool indices to SWA pool indices)
|
# Page table for SWA layers (translated from full pool indices to SWA pool indices)
|
||||||
swa_page_table: torch.Tensor = None
|
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)
|
# full->SWA translated out_cache_loc (SWA KV-store write target)
|
||||||
swa_out_cache_loc: torch.Tensor = None
|
swa_out_cache_loc: torch.Tensor = None
|
||||||
is_ragged_verify: bool = False
|
is_ragged_verify: bool = False
|
||||||
@@ -339,6 +342,25 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
return swa_pt
|
return swa_pt
|
||||||
return self.forward_metadata.page_table
|
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
|
@staticmethod
|
||||||
def _get_scalar_scale(
|
def _get_scalar_scale(
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
@@ -962,6 +984,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
self._fill_page_table_device(
|
self._fill_page_table_device(
|
||||||
metadata, forward_batch.req_pool_indices, metadata.cache_seqlens_int32
|
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):
|
if self._needs_encoder_only_expand(forward_batch.forward_mode, metadata):
|
||||||
row_map = (
|
row_map = (
|
||||||
@@ -1276,13 +1299,23 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
cu_seqlens_q,
|
cu_seqlens_q,
|
||||||
cache_seqlens,
|
cache_seqlens,
|
||||||
max_seqlen_q,
|
max_seqlen_q,
|
||||||
|
*,
|
||||||
cu_seqlens_kv,
|
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(
|
return flashinfer.prefill.trtllm_batch_context_with_kv_cache(
|
||||||
query=q_chunk,
|
query=q_chunk,
|
||||||
kv_cache=kv_cache,
|
kv_cache=kv_cache,
|
||||||
workspace_buffer=self.workspace_buffer,
|
workspace_buffer=self.workspace_buffer,
|
||||||
block_tables=page_table,
|
block_tables=block_tables,
|
||||||
seq_lens=cache_seqlens,
|
seq_lens=cache_seqlens,
|
||||||
max_q_len=max_seqlen_q,
|
max_q_len=max_seqlen_q,
|
||||||
max_kv_len=self.max_context_len,
|
max_kv_len=self.max_context_len,
|
||||||
|
|||||||
@@ -79,11 +79,18 @@ class ZigzagContextParallelMetadata(BaseContextParallelMetadata):
|
|||||||
cu_seqlens_q_prev_tensor: Optional[Any] = None
|
cu_seqlens_q_prev_tensor: Optional[Any] = None
|
||||||
cu_seqlens_q_next_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.
|
# Scalars derived from the per-sequence lists above.
|
||||||
total_q_prev_tokens: int = 0
|
total_q_prev_tokens: int = 0
|
||||||
total_q_next_tokens: int = 0
|
total_q_next_tokens: int = 0
|
||||||
max_seqlen_q_prev: int = 0
|
max_seqlen_q_prev: int = 0
|
||||||
max_seqlen_q_next: int = 0
|
max_seqlen_q_next: int = 0
|
||||||
|
max_seqlen_q_combined: int = 0
|
||||||
|
|
||||||
# Per-sequence CPU lists, useful for indexers and diagnostics.
|
# Per-sequence CPU lists, useful for indexers and diagnostics.
|
||||||
kv_len_prev_list: Optional[List[int]] = None
|
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_next = [0] + list(accumulate(actual_seq_q_next_list))
|
||||||
cu_kv_prev = [0] + list(accumulate(kv_len_prev_list))
|
cu_kv_prev = [0] + list(accumulate(kv_len_prev_list))
|
||||||
cu_kv_next = [0] + list(accumulate(kv_len_next_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)
|
total_seq_lens = sum(extend_seqs_len)
|
||||||
assert len(split_list) == bs * cp_segment_num
|
assert len(split_list) == bs * cp_segment_num
|
||||||
@@ -256,6 +267,18 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
cu_seqlens_q_next_tensor=torch.tensor(
|
cu_seqlens_q_next_tensor=torch.tensor(
|
||||||
cu_next, device=device, dtype=torch.int32
|
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_prev_tokens=cu_prev[-1],
|
||||||
total_q_next_tokens=cu_next[-1],
|
total_q_next_tokens=cu_next[-1],
|
||||||
max_seqlen_q_prev=(
|
max_seqlen_q_prev=(
|
||||||
@@ -264,6 +287,9 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
max_seqlen_q_next=(
|
max_seqlen_q_next=(
|
||||||
max(actual_seq_q_next_list) if actual_seq_q_next_list else 0
|
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_prev_list=kv_len_prev_list,
|
||||||
kv_len_next_list=kv_len_next_list,
|
kv_len_next_list=kv_len_next_list,
|
||||||
actual_seq_q_prev_list=actual_seq_q_prev_list,
|
actual_seq_q_prev_list=actual_seq_q_prev_list,
|
||||||
@@ -334,9 +360,15 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
prev_kwargs = {}
|
prev_kwargs = {}
|
||||||
next_kwargs = {}
|
next_kwargs = {}
|
||||||
if attention_backend == CPAttentionBackendKind.TRTLLM_MHA:
|
if attention_backend == CPAttentionBackendKind.TRTLLM_MHA:
|
||||||
prev_kwargs["cu_seqlens_kv"] = meta.cu_seqlens_kv_prev_tensor
|
result = attn_fn(
|
||||||
next_kwargs["cu_seqlens_kv"] = meta.cu_seqlens_kv_next_tensor
|
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(
|
result_prev = attn_fn(
|
||||||
q_prev,
|
q_prev,
|
||||||
meta.cu_seqlens_q_prev_tensor,
|
meta.cu_seqlens_q_prev_tensor,
|
||||||
@@ -352,6 +384,7 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
**next_kwargs,
|
**next_kwargs,
|
||||||
)
|
)
|
||||||
result = torch.cat([result_prev, result_next], dim=0)
|
result = torch.cat([result_prev, result_next], dim=0)
|
||||||
|
|
||||||
pad_size = q.shape[0] - logical_tokens
|
pad_size = q.shape[0] - logical_tokens
|
||||||
assert pad_size >= 0
|
assert pad_size >= 0
|
||||||
if pad_size > 0:
|
if pad_size > 0:
|
||||||
|
|||||||
@@ -547,6 +547,108 @@ class TestCPZigzagStrategy(CustomTestCase):
|
|||||||
self.assertTrue(torch.equal(calls[1][0], q[2:]))
|
self.assertTrue(torch.equal(calls[1][0], q[2:]))
|
||||||
self.assertTrue(torch.equal(out, q + 100))
|
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):
|
class TestCPInterleaveStrategy(CustomTestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user