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:
Minglei Zhu
2026-04-10 13:57:54 -07:00
committed by GitHub
co-authored by zminglei
parent 4ace144fae
commit 6af34b95b6
@@ -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]