Add swa page table to FA3 backend forward metadata (#16267)

This commit is contained in:
Ke Bao
2026-01-02 20:00:43 +08:00
committed by GitHub
parent dd93e4451c
commit 698629d1f2
@@ -47,6 +47,8 @@ class FlashAttentionMetadata:
window_size: tuple = (-1, -1) window_size: tuple = (-1, -1)
# Page table, the index of KV Cache Tables/Blocks # Page table, the index of KV Cache Tables/Blocks
page_table: torch.Tensor = None page_table: torch.Tensor = None
# Page table for Sliding Window Attention
swa_page_table: torch.Tensor = None
# Encoder metadata # Encoder metadata
# Cumulative sequence lengths for encoder key # Cumulative sequence lengths for encoder key
@@ -337,9 +339,11 @@ class FlashAttentionBackend(AttentionBackend):
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
self.use_sliding_window_kv_pool = isinstance( self.use_sliding_window_kv_pool = (
model_runner.token_to_kv_pool, SWAKVPool isinstance(model_runner.token_to_kv_pool, SWAKVPool)
and model_runner.token_to_kv_pool.swa_layer_nums > 0
) )
if self.use_sliding_window_kv_pool: if self.use_sliding_window_kv_pool:
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
@@ -646,11 +650,24 @@ class FlashAttentionBackend(AttentionBackend):
), ),
] ]
if self.use_sliding_window_kv_pool:
metadata.swa_page_table = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
metadata.page_table
)
)
# Convert the page table to a strided format which is needed by FA3 API # Convert the page table to a strided format which is needed by FA3 API
if self.page_size > 1: if self.page_size > 1:
self.strided_indices = torch.arange( self.strided_indices = torch.arange(
0, metadata.page_table.shape[1], self.page_size, device=self.device 0, metadata.page_table.shape[1], self.page_size, device=self.device
) )
if self.use_sliding_window_kv_pool:
metadata.swa_page_table = (
metadata.swa_page_table[:, self.strided_indices] // self.page_size
)
metadata.page_table = ( metadata.page_table = (
metadata.page_table[:, self.strided_indices] // self.page_size metadata.page_table[:, self.strided_indices] // self.page_size
) )
@@ -800,8 +817,11 @@ class FlashAttentionBackend(AttentionBackend):
else: else:
page_table = metadata.page_table page_table = metadata.page_table
if is_swa_layer and self.use_sliding_window_kv_pool: if is_swa_layer and self.use_sliding_window_kv_pool:
if metadata.swa_page_table is not None:
page_table = metadata.swa_page_table
else:
page_table = self.token_to_kv_pool.translate_loc_from_full_to_swa( page_table = self.token_to_kv_pool.translate_loc_from_full_to_swa(
page_table metadata.page_table
) )
cu_seqlens_q = metadata.cu_seqlens_q cu_seqlens_q = metadata.cu_seqlens_q
cache_seqlens = metadata.cache_seqlens_int32 cache_seqlens = metadata.cache_seqlens_int32
@@ -1155,8 +1175,13 @@ class FlashAttentionBackend(AttentionBackend):
else: else:
page_table = metadata.page_table page_table = metadata.page_table
if is_swa_layer and self.use_sliding_window_kv_pool: if is_swa_layer and self.use_sliding_window_kv_pool:
page_table = self.token_to_kv_pool.translate_loc_from_full_to_swa( if metadata.swa_page_table is not None:
page_table page_table = metadata.swa_page_table
else:
page_table = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
metadata.page_table
)
) )
cache_seqlens = metadata.cache_seqlens_int32 cache_seqlens = metadata.cache_seqlens_int32
cu_seqlens_k = metadata.cu_seqlens_k cu_seqlens_k = metadata.cu_seqlens_k
@@ -1350,6 +1375,14 @@ class FlashAttentionBackend(AttentionBackend):
), ),
} }
if self.use_sliding_window_kv_pool:
self.decode_cuda_graph_metadata["swa_page_table"] = torch.zeros(
max_bs,
max_num_pages,
dtype=torch.int32,
device=self.device,
)
# This is used by draft decode's first half of metadata when topk > 1 # This is used by draft decode's first half of metadata when topk > 1
if self.topk > 1: if self.topk > 1:
self.draft_decode_metadata_topk_normal = { self.draft_decode_metadata_topk_normal = {
@@ -1649,6 +1682,10 @@ class FlashAttentionBackend(AttentionBackend):
metadata.page_table = self.decode_cuda_graph_metadata["page_table"][ metadata.page_table = self.decode_cuda_graph_metadata["page_table"][
:bs, : :bs, :
] ]
if self.use_sliding_window_kv_pool:
metadata.swa_page_table = self.decode_cuda_graph_metadata[
"swa_page_table"
][:bs, :]
# Precompute cumulative sequence lengths # Precompute cumulative sequence lengths
metadata.cu_seqlens_q = torch.arange( metadata.cu_seqlens_q = torch.arange(
0, batch_size + 1, dtype=torch.int32, device=device 0, batch_size + 1, dtype=torch.int32, device=device
@@ -1900,6 +1937,8 @@ class FlashAttentionBackend(AttentionBackend):
seq_lens, seq_lens,
0, 0,
self.page_size, self.page_size,
metadata.swa_page_table,
self.token_to_kv_pool if self.use_sliding_window_kv_pool else None,
) )
self._update_local_attn_metadata_for_replay( self._update_local_attn_metadata_for_replay(
@@ -2534,6 +2573,8 @@ def normal_decode_set_metadata(
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
seq_len_delta: int, seq_len_delta: int,
page_size: int, page_size: int,
swa_page_table: Optional[torch.Tensor] = None,
token_to_kv_pool: Optional[SWAKVPool] = None,
): ):
cache_seqlens_int32.copy_(seq_lens + seq_len_delta) cache_seqlens_int32.copy_(seq_lens + seq_len_delta)
cu_seqlens_k[1:].copy_(torch.cumsum(cache_seqlens_int32, dim=0, dtype=torch.int32)) cu_seqlens_k[1:].copy_(torch.cumsum(cache_seqlens_int32, dim=0, dtype=torch.int32))
@@ -2543,6 +2584,11 @@ def normal_decode_set_metadata(
] ]
page_table[:, :max_seq_pages].copy_(page_indices // page_size) page_table[:, :max_seq_pages].copy_(page_indices // page_size)
if swa_page_table is not None and token_to_kv_pool is not None:
assert isinstance(token_to_kv_pool, SWAKVPool)
swa_page_indices = token_to_kv_pool.translate_loc_from_full_to_swa(page_indices)
swa_page_table[:, :max_seq_pages].copy_(swa_page_indices // page_size)
@torch.compile(dynamic=True, backend=get_compiler_backend()) @torch.compile(dynamic=True, backend=get_compiler_backend())
def draft_decode_set_expand_metadata( def draft_decode_set_expand_metadata(