[Bugfix] Fix Llama 4 FA3 local attention with paged KV cache (#32902)

Co-authored-by: kpham-sgl <khoa.pham@radixark.ai>
This commit is contained in:
Po-Han Huang (NVIDIA)
2026-09-03 23:15:23 -07:00
committed by GitHub
co-authored by kpham-sgl
parent 4e756ecc4a
commit d122ca99b2
@@ -3183,6 +3183,9 @@ class FlashAttentionBackend(AttentionBackend):
if cu_seqlens_q is None or cache_seqlens_int32 is None or page_table is None:
metadata.local_attn_metadata = None
return
if self.page_size > 1:
# Convert the eager token table to physical page indices.
page_table = page_table[:, :: self.page_size] // self.page_size
cu_seqlens_q_np = cu_seqlens_q.cpu().numpy()
seq_lens_np = cache_seqlens_int32.cpu().numpy()
@@ -3197,6 +3200,7 @@ class FlashAttentionBackend(AttentionBackend):
seq_lens_np,
page_table,
self.page_size,
preserve_attn_chunk_size=True,
)
local_metadata = FlashAttentionMetadata.LocalAttentionMetadata(
@@ -3237,6 +3241,7 @@ class FlashAttentionBackend(AttentionBackend):
seqlens_np,
page_table_capture,
self.page_size,
preserve_attn_chunk_size=True,
)
# Get exact dimensions from the calculation
@@ -3321,6 +3326,7 @@ class FlashAttentionBackend(AttentionBackend):
seqlens_np,
sliced_page_table,
self.page_size,
preserve_attn_chunk_size=True,
)
# Convert back to tensors
@@ -3564,6 +3570,7 @@ def make_local_attention_virtual_batches(
seq_lens_np: np.ndarray,
block_table: torch.Tensor,
page_size: int = 0,
preserve_attn_chunk_size: bool = False,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, torch.Tensor]:
"""
Take in `query_start_loc_np` and `seq_lens_np` and break the sequences into
@@ -3576,6 +3583,7 @@ def make_local_attention_virtual_batches(
seq_lens_np: Sequence lengths (numpy array)
block_table: Block table for KV cache
page_size: Size of each page in the KV cache
preserve_attn_chunk_size: Skip sequence-length-based chunk normalization.
Returns:
seqlens_q_local: Query sequence lengths for local attention
@@ -3583,15 +3591,13 @@ def make_local_attention_virtual_batches(
seqlens_k_local: Key sequence lengths for local attention
block_table_local: Block table for local attention
"""
# Adjust attention_chunk_size based on the actual sequence length
# to avoid index out of bounds errors
max_seq_len = seq_lens_np.max()
effective_chunk_size = min(attn_chunk_size, max_seq_len)
# Make sure effective_chunk_size is divisible by page_size
effective_chunk_size = (effective_chunk_size // page_size) * page_size
if effective_chunk_size < page_size:
effective_chunk_size = page_size
attn_chunk_size = effective_chunk_size
if not preserve_attn_chunk_size:
max_seq_len = seq_lens_np.max()
effective_chunk_size = min(attn_chunk_size, max_seq_len)
effective_chunk_size = (effective_chunk_size // page_size) * page_size
if effective_chunk_size < page_size:
effective_chunk_size = page_size
attn_chunk_size = effective_chunk_size
q_seqlens = query_start_loc_np[1:] - query_start_loc_np[:-1]
actual_batch_size = seq_lens_np.shape[0]