perf: precompute FA3 scheduler_metadata to eliminate per-layer prepare_varlen_num_blocks (#21104)
Co-authored-by: zminglei <zminglei@linkedin.com>
This commit is contained in:
@@ -57,6 +57,8 @@ class FlashAttentionMetadata:
|
||||
page_table: torch.Tensor = None
|
||||
# Page table for Sliding Window Attention
|
||||
swa_page_table: torch.Tensor = None
|
||||
# Precomputed FA3 scheduler metadata (avoids per-layer prepare_varlen_num_blocks)
|
||||
scheduler_metadata: torch.Tensor = None
|
||||
|
||||
# Encoder metadata
|
||||
# Cumulative sequence lengths for encoder key
|
||||
@@ -167,18 +169,37 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
from sgl_kernel.flash_attn import (
|
||||
flash_attn_varlen_func,
|
||||
flash_attn_with_kvcache,
|
||||
get_scheduler_metadata,
|
||||
)
|
||||
|
||||
self._get_scheduler_metadata = get_scheduler_metadata
|
||||
elif self.fa_impl_ver == 4:
|
||||
from sglang.jit_kernel.flash_attention_v4 import (
|
||||
flash_attn_varlen_func,
|
||||
flash_attn_with_kvcache,
|
||||
)
|
||||
|
||||
self._get_scheduler_metadata = None
|
||||
else:
|
||||
raise ValueError(f"Invalid version: {self.fa_impl_ver=}")
|
||||
|
||||
self.flash_attn_varlen_func = flash_attn_varlen_func
|
||||
self.flash_attn_with_kvcache = flash_attn_with_kvcache
|
||||
|
||||
# Store head info for precomputing FA3 scheduler metadata
|
||||
self.head_dim = model_runner.model_config.head_dim
|
||||
self.num_attention_heads = (
|
||||
model_runner.model_config.hf_text_config.num_attention_heads
|
||||
// model_runner.tp_size
|
||||
)
|
||||
self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
|
||||
model_runner.tp_size
|
||||
)
|
||||
_softcapping = getattr(
|
||||
model_runner.model_config.hf_text_config, "attn_logit_softcapping", None
|
||||
)
|
||||
self.has_softcap = _softcapping is not None and _softcapping > 0.0
|
||||
|
||||
# If num_splits == 0, we use a heuristic to automatically determine the number of splits.
|
||||
# We set nums splits to 1 if deterministic inference is enabled.
|
||||
# See https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ for more details.
|
||||
@@ -193,6 +214,33 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
else 0
|
||||
)
|
||||
|
||||
def _compute_scheduler_metadata(
|
||||
self, batch_size, max_seq_len_k, cache_seqlens, cu_seqlens_q
|
||||
):
|
||||
"""Compute FA3 scheduler metadata for decode.
|
||||
|
||||
Returns the scheduler_metadata tensor, or None if not applicable.
|
||||
"""
|
||||
if self._get_scheduler_metadata is None or self.use_mla:
|
||||
return None
|
||||
# Always use window_size=(-1, -1) because scheduler_metadata is only
|
||||
# consumed by non-SWA layers (SWA layers skip it in forward_decode).
|
||||
return self._get_scheduler_metadata(
|
||||
batch_size=batch_size,
|
||||
max_seqlen_q=1,
|
||||
max_seqlen_k=max_seq_len_k,
|
||||
num_heads=self.num_attention_heads,
|
||||
num_heads_k=self.num_kv_heads,
|
||||
headdim=self.head_dim,
|
||||
cache_seqlens=cache_seqlens,
|
||||
qkv_dtype=self.kv_cache_dtype,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
page_size=self.page_size,
|
||||
causal=True,
|
||||
has_softcap=self.has_softcap,
|
||||
num_splits=self.num_splits,
|
||||
)
|
||||
|
||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||
"""Initialize forward metadata hence all layers in the forward pass can reuse it."""
|
||||
metadata = FlashAttentionMetadata()
|
||||
@@ -285,6 +333,14 @@ 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
|
||||
]
|
||||
# Precompute FA3 scheduler metadata to avoid per-layer
|
||||
# prepare_varlen_num_blocks kernel calls
|
||||
metadata.scheduler_metadata = self._compute_scheduler_metadata(
|
||||
batch_size,
|
||||
metadata.max_seq_len_k,
|
||||
metadata.cache_seqlens_int32,
|
||||
metadata.cu_seqlens_q,
|
||||
)
|
||||
# TODO: we need to test this part for llama 4 eagle case
|
||||
self._maybe_init_local_attn_metadata(forward_batch, metadata, device)
|
||||
elif forward_batch.forward_mode.is_target_verify():
|
||||
@@ -1059,6 +1115,15 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
)
|
||||
|
||||
# Default: single-token self-attention
|
||||
# Use precomputed scheduler_metadata when available and applicable.
|
||||
# scheduler_metadata is only valid for non-SWA, non-cascade decode.
|
||||
sched_meta = None
|
||||
if (
|
||||
metadata.scheduler_metadata is not None
|
||||
and not is_swa_layer
|
||||
and not use_cascade_attn
|
||||
):
|
||||
sched_meta = metadata.scheduler_metadata
|
||||
result = flash_attn_with_kvcache(
|
||||
q=q_reshaped,
|
||||
k_cache=key_cache,
|
||||
@@ -1076,6 +1141,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
return_softmax_lse=use_cascade_attn,
|
||||
num_splits=self.num_splits,
|
||||
ver=self.fa_impl_ver,
|
||||
scheduler_metadata=sched_meta,
|
||||
**kwargs,
|
||||
)
|
||||
if use_cascade_attn:
|
||||
@@ -1220,6 +1286,16 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
0, self.max_context_len, self.page_size, device=self.device
|
||||
),
|
||||
}
|
||||
# Pre-allocate scheduler_metadata buffer for CUDA graph
|
||||
# Size: 1 (semaphore) + round_up(max_bs, 4) * 4 (causal decode vectors)
|
||||
if self._get_scheduler_metadata is not None and not self.use_mla:
|
||||
b_rounded = ((max_bs + 3) // 4) * 4
|
||||
self._sched_meta_buf = torch.zeros(
|
||||
1 + b_rounded * 4, dtype=torch.int32, device=self.device
|
||||
)
|
||||
else:
|
||||
self._sched_meta_buf = None
|
||||
|
||||
# Only allocate local attention buffers if local attention is enabled
|
||||
# This prevents OOM errors when local attention is not being used
|
||||
if self.has_local_attention:
|
||||
@@ -1589,6 +1665,20 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
|
||||
self._maybe_update_local_attn_metadata_for_capture(metadata, batch_size)
|
||||
|
||||
# Compute scheduler_metadata into pre-allocated buffer for CUDA graph capture
|
||||
if self._sched_meta_buf is not None:
|
||||
sched = self._compute_scheduler_metadata(
|
||||
batch_size,
|
||||
max(metadata.max_seq_len_k, 1),
|
||||
metadata.cache_seqlens_int32,
|
||||
metadata.cu_seqlens_q,
|
||||
)
|
||||
if sched is not None:
|
||||
n = sched.shape[0]
|
||||
self._sched_meta_buf[:n] = sched
|
||||
self._sched_meta_buf[n:] = 0
|
||||
metadata.scheduler_metadata = self._sched_meta_buf[:n]
|
||||
|
||||
elif forward_mode.is_target_verify():
|
||||
if self.topk <= 1:
|
||||
metadata.cache_seqlens_int32 = self.target_verify_metadata[
|
||||
@@ -1855,6 +1945,23 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata,
|
||||
bs,
|
||||
)
|
||||
|
||||
# Recompute scheduler_metadata into pre-allocated buffer
|
||||
if (
|
||||
self._sched_meta_buf is not None
|
||||
and metadata.scheduler_metadata is not None
|
||||
):
|
||||
sched = self._compute_scheduler_metadata(
|
||||
bs,
|
||||
metadata.max_seq_len_k,
|
||||
metadata.cache_seqlens_int32,
|
||||
metadata.cu_seqlens_q,
|
||||
)
|
||||
if sched is not None:
|
||||
n = sched.shape[0]
|
||||
self._sched_meta_buf[:n] = sched
|
||||
self._sched_meta_buf[n:] = 0
|
||||
|
||||
elif forward_mode.is_target_verify():
|
||||
if self.topk <= 1:
|
||||
metadata = self.target_verify_metadata[bs]
|
||||
|
||||
Reference in New Issue
Block a user