Support FlashAttention3 page_size > 1 and topk > 1 case with paged attn and spec decode (#7725)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user