fix: fix fa3 cross-attention batched-decode for per-request varlen encoder (#25045)

This commit is contained in:
Shaojun Zhou
2026-05-26 14:26:15 +08:00
committed by GitHub
parent f0ba651d66
commit 7c0fbc8c2e
@@ -547,28 +547,37 @@ class FlashAttentionBackend(AttentionBackend):
if forward_batch.forward_mode == ForwardMode.EXTEND:
self._maybe_init_local_attn_metadata(forward_batch, metadata, device)
# Encoder metadata for cross attention
# Encoder metadata for cross attention. Supports per-request varlen
# encoder lengths (e.g. MossVL with different image sizes per request).
if forward_batch.encoder_lens is not None:
assert (
forward_batch.encoder_lens.numel() == 1
), "Only encoder size 1 is supported for now"
metadata.encoder_lens_int32 = forward_batch.encoder_lens.to(torch.int32)
metadata.encoder_cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(metadata.encoder_lens_int32, dim=0, dtype=torch.int32),
(1, 0),
)
metadata.encoder_max_seq_len_k = metadata.encoder_lens_int32.max().item()
# Cross-attn page_table: per-request rows. cache_seqlens
# (encoder_lens_int32) caps per-request reads so any garbage past
# encoder_lens[i] is never consumed.
metadata.encoder_page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.encoder_max_seq_len_k
]
# Currently only support forward_batch.encoder_lens.numel() == 1
# Self-attn (text) page_table: text starts at per-request offset
# encoder_lens[i], NOT at a single max. Use a fancy-index gather.
text_max = metadata.max_seq_len_k
arange_text = torch.arange(
text_max, device=forward_batch.req_pool_indices.device
)
text_col = forward_batch.encoder_lens.long().unsqueeze(
1
) + arange_text.unsqueeze(
0
) # (bs, max_seq_len_k)
text_row = forward_batch.req_pool_indices.unsqueeze(1).expand(-1, text_max)
metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices,
metadata.encoder_max_seq_len_k : (
metadata.encoder_max_seq_len_k + metadata.max_seq_len_k
),
text_row, text_col
]
if self.use_sliding_window_kv_pool:
@@ -2312,25 +2321,25 @@ class FlashAttentionBackend(AttentionBackend):
metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size)
if encoder_lens is not None:
# Only support encoder size 1 for now
metadata.encoder_max_seq_len_k = encoder_lens[0]
metadata.encoder_lens_int32.copy_(encoder_lens[:1])
metadata.encoder_cu_seqlens_k[1:].copy_(
torch.cumsum(metadata.encoder_lens_int32, dim=0, dtype=torch.int32)
# Per-request varlen encoder support (e.g. MossVL different images).
metadata.encoder_max_seq_len_k = int(encoder_lens.max().item())
metadata.encoder_lens_int32[:bs].copy_(encoder_lens[:bs].to(torch.int32))
metadata.encoder_cu_seqlens_k[1 : bs + 1].copy_(
torch.cumsum(metadata.encoder_lens_int32[:bs], dim=0, dtype=torch.int32)
)
metadata.encoder_page_table[:, : metadata.encoder_max_seq_len_k].copy_(
metadata.encoder_page_table[:bs, : metadata.encoder_max_seq_len_k].copy_(
self.req_to_token[req_pool_indices, : metadata.encoder_max_seq_len_k]
)
# Update the regular page table
page_table = self.req_to_token[
req_pool_indices,
metadata.encoder_max_seq_len_k : (
metadata.encoder_max_seq_len_k + metadata.max_seq_len_k
),
]
metadata.page_table[:, : metadata.max_seq_len_k].copy_(page_table)
# Self-attn (text) page_table: per-request offset = encoder_lens[i].
text_max = metadata.max_seq_len_k
arange_text = torch.arange(text_max, device=req_pool_indices.device)
text_col = encoder_lens[:bs].long().unsqueeze(1) + arange_text.unsqueeze(0)
text_row = req_pool_indices.unsqueeze(1).expand(-1, text_max)
metadata.page_table[:bs, :text_max].copy_(
self.req_to_token[text_row, text_col]
)
self.forward_metadata = metadata
self.forward_metadata_spec_decode_expand = metadata_expand