Support FlashAttention3 page_size > 1 and topk > 1 case with paged attn and spec decode (#7725)

This commit is contained in:
Yubo Wang
2025-11-26 11:44:41 +08:00
committed by GitHub
parent ca5c8b16f6
commit 18fb51583f
9 changed files with 706 additions and 86 deletions
@@ -14,6 +14,7 @@ from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import get_compiler_backend
if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
@@ -411,7 +412,6 @@ class FlashAttentionBackend(AttentionBackend):
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k
]
metadata_expand = FlashAttentionMetadata()
decode_length = self.speculative_step_id + 1
metadata_expand.cache_seqlens_int32 = torch.full(
@@ -645,6 +645,40 @@ class FlashAttentionBackend(AttentionBackend):
metadata.page_table[:, self.strided_indices] // self.page_size
)
if (
self.topk > 1
and forward_batch.forward_mode.is_decode_or_idle()
and forward_batch.spec_info is not None
):
# Modifies cache_seqlens_int32 and page_table(B, speculative_num_steps).
last_page_lens = forward_batch.seq_lens % self.page_size
# First attention handles prefix - last_page_len part.
metadata.cache_seqlens_int32 -= last_page_lens # Both (B, )
# Second attention handles last_page_len + decode part.
expanded_last_page_lens = last_page_lens.repeat_interleave(self.topk)
self.forward_metadata_spec_decode_expand.cache_seqlens_int32 += (
expanded_last_page_lens
)
decode_length = self.speculative_step_id + 1
expand_page_table = cache_loc[:, :decode_length].clone()
strided_indices_expand = torch.arange(
0,
decode_length,
self.page_size,
device=self.device,
)
last_page_lens_broadcast = expanded_last_page_lens.unsqueeze(-1).expand(
-1, expand_page_table.shape[1]
)
expand_page_table -= last_page_lens_broadcast
expand_page_table = (
expand_page_table[:, strided_indices_expand] // self.page_size
)
self.forward_metadata_spec_decode_expand.page_table = (
expand_page_table.to(torch.int32)
)
self.forward_metadata = metadata
def forward_extend(
@@ -798,8 +832,13 @@ class FlashAttentionBackend(AttentionBackend):
o, softmax_lse, *rest = result
o_expand, softmax_lse_expand, *rest_expand = flash_attn_with_kvcache(
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
k_cache=key_cache,
v_cache=value_cache,
# Here metadata_expand.page_table is not divided with page_size.
# This is because we loose the fine control of what token to attend,
# but has to attend to some block completely.
k_cache=key_cache.view(-1, 1, layer.tp_k_head_num, layer.head_dim),
v_cache=value_cache.view(
-1, 1, layer.tp_v_head_num, layer.head_dim
),
page_table=self.forward_metadata_spec_decode_expand.page_table,
cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32,
cu_seqlens_q=self.forward_metadata_spec_decode_expand.cu_seqlens_q,
@@ -1112,7 +1151,6 @@ class FlashAttentionBackend(AttentionBackend):
page_table=page_table,
cache_seqlens=cache_seqlens,
cu_seqlens_q=metadata.cu_seqlens_q,
cu_seqlens_k_new=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
softmax_scale=layer.scaling,
causal=False if use_cascade_attn else causal,
@@ -1344,6 +1382,17 @@ class FlashAttentionBackend(AttentionBackend):
),
}
if self.page_size > 1:
# Used for indicing expand page_table
self.draft_decode_metadata_topk_expand["strided_indices_expand"] = (
torch.arange(
0,
self.speculative_num_steps,
self.page_size,
device=self.device,
)
)
if (
self.speculative_num_draft_tokens is not None
and self.speculative_num_draft_tokens > 0
@@ -1778,30 +1827,58 @@ class FlashAttentionBackend(AttentionBackend):
# When top k > 1, we need two specific draft decode metadata, and then merge states
# 1. The first half of metadata for prefix tokens
metadata = self.draft_decode_metadata_topk_normal[bs]
if self.page_size > 1:
# First attention handles seq_lens - last_page_lens if page size > 1.
last_page_lens = seq_lens % self.page_size
seq_lens = seq_lens - last_page_lens
# last_page_lens_cpu = last_page_lens.max().item()
# seq_lens_cpu -= last_page_lens_cpu
metadata.cache_seqlens_int32.copy_(seq_lens)
# metadata.max_seq_len_q = self.topk, already set in capture
metadata.max_seq_len_k = seq_lens_cpu.max().item()
# metadata.cu_seqlens_q already set in capture
metadata.cu_seqlens_k[1:].copy_(
torch.cumsum(
metadata.cache_seqlens_int32, dim=0, dtype=torch.int32
)
# metadata.cu_seqlens_k is not needed
metadata.max_seq_len_k = seq_lens_cpu.max().item()
max_seq_pages = (
metadata.max_seq_len_k + self.page_size - 1
) // self.page_size
strided_indices = self.decode_cuda_graph_metadata["strided_indices"]
strided_indices = strided_indices[:max_seq_pages]
page_table = (
self.req_to_token[
req_pool_indices[:, None], # shape [bs, 1]
strided_indices[None, :], # shape [1, max_seq_pages]
]
// self.page_size
)
page_table = self.req_to_token[
req_pool_indices, : metadata.max_seq_len_k
]
metadata.page_table[:, : metadata.max_seq_len_k].copy_(page_table)
metadata.page_table[:, :max_seq_pages].copy_(page_table)
# 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk)
metadata_expand = self.draft_decode_metadata_topk_expand[bs]
decode_length = self.speculative_step_id + 1
# shape: [bs, num_steps, topk] -> [bs x topk, num_steps]
cache_loc = out_cache_loc.view(-1, self.speculative_num_steps)
metadata_expand.page_table[: cache_loc.shape[0]].copy_(
cache_loc[:, :decode_length]
)
if self.page_size > 1:
# Second attention handles last_page_len + decode part.
strided_indices_expand = (
self.draft_decode_metadata_topk_expand.get(
"strided_indices_expand"
)
)
update_draft_decode_set_expand_metadata_with_page_size(
metadata_expand.cache_seqlens_int32, # Modifies
metadata_expand.page_table, # Modifies
cache_loc,
last_page_lens,
strided_indices_expand,
decode_length,
bs,
self.topk,
self.page_size,
)
else:
metadata_expand.page_table[: cache_loc.shape[0]].copy_(
cache_loc[:, :decode_length]
)
# TODO: Handle local attention metadata for draft decode when llama4 eagle is supported
else:
# Normal Decode
@@ -1860,10 +1937,15 @@ class FlashAttentionBackend(AttentionBackend):
metadata.cu_seqlens_k[1:].copy_(
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
)
page_table = self.req_to_token[
req_pool_indices, : metadata.max_seq_len_k
max_seq_pages = (
metadata.max_seq_len_k + self.page_size - 1
) // self.page_size
page_indices = self.req_to_token[
req_pool_indices[:, None],
self.decode_cuda_graph_metadata["strided_indices"][:max_seq_pages],
]
metadata.page_table[:, : metadata.max_seq_len_k].copy_(page_table)
page_indices //= self.page_size
metadata.page_table[:, :max_seq_pages].copy_(page_indices)
# 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk)
metadata_expand = self.target_verify_metadata_topk_expand[bs]
@@ -1926,7 +2008,6 @@ class FlashAttentionBackend(AttentionBackend):
dtype=torch.int32,
)
)
if self.has_swa:
metadata_swa = self.target_verify_metadata_topk_swa[bs]
self._init_sliding_window_attn_spec_metadata(
@@ -2411,3 +2492,32 @@ def normal_decode_set_metadata(
strided_indices[:max_seq_pages][None, :],
]
page_table[:, :max_seq_pages].copy_(page_indices // page_size)
@torch.compile(dynamic=True, backend=get_compiler_backend())
def update_draft_decode_set_expand_metadata_with_page_size(
cache_seqlens_int32: torch.Tensor, # Modifies
page_table: torch.Tensor, # Modifies
cache_loc: torch.Tensor,
last_page_lens: torch.Tensor,
strided_indices_expand: torch.Tensor,
decode_length: int,
bs: int,
topk: int,
page_size: int,
):
expanded_last_page_lens = last_page_lens.repeat_interleave(topk)
cache_seqlens_int32.copy_(decode_length + expanded_last_page_lens)
expand_page_table = cache_loc[:, :decode_length].clone()
last_page_lens_broadcast = expanded_last_page_lens.unsqueeze(-1).expand(
-1, expand_page_table.shape[1]
)
expand_page_table -= last_page_lens_broadcast
expand_page_table = (
expand_page_table[
:, strided_indices_expand[: (decode_length + page_size - 1) // page_size]
]
// page_size
)
max_seq_pages_expand = (decode_length + page_size - 1) // page_size
page_table[:, :max_seq_pages_expand].copy_(expand_page_table)