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: torch.Tensor = None
|
||||||
# Page table for Sliding Window Attention
|
# Page table for Sliding Window Attention
|
||||||
swa_page_table: torch.Tensor = None
|
swa_page_table: torch.Tensor = None
|
||||||
|
# Precomputed FA3 scheduler metadata (avoids per-layer prepare_varlen_num_blocks)
|
||||||
|
scheduler_metadata: torch.Tensor = None
|
||||||
|
|
||||||
# Encoder metadata
|
# Encoder metadata
|
||||||
# Cumulative sequence lengths for encoder key
|
# Cumulative sequence lengths for encoder key
|
||||||
@@ -167,18 +169,37 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
from sgl_kernel.flash_attn import (
|
from sgl_kernel.flash_attn import (
|
||||||
flash_attn_varlen_func,
|
flash_attn_varlen_func,
|
||||||
flash_attn_with_kvcache,
|
flash_attn_with_kvcache,
|
||||||
|
get_scheduler_metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self._get_scheduler_metadata = get_scheduler_metadata
|
||||||
elif self.fa_impl_ver == 4:
|
elif self.fa_impl_ver == 4:
|
||||||
from sglang.jit_kernel.flash_attention_v4 import (
|
from sglang.jit_kernel.flash_attention_v4 import (
|
||||||
flash_attn_varlen_func,
|
flash_attn_varlen_func,
|
||||||
flash_attn_with_kvcache,
|
flash_attn_with_kvcache,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self._get_scheduler_metadata = None
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid version: {self.fa_impl_ver=}")
|
raise ValueError(f"Invalid version: {self.fa_impl_ver=}")
|
||||||
|
|
||||||
self.flash_attn_varlen_func = flash_attn_varlen_func
|
self.flash_attn_varlen_func = flash_attn_varlen_func
|
||||||
self.flash_attn_with_kvcache = flash_attn_with_kvcache
|
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.
|
# 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.
|
# We set nums splits to 1 if deterministic inference is enabled.
|
||||||
# See https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ for more details.
|
# See https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ for more details.
|
||||||
@@ -193,6 +214,33 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else 0
|
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):
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
"""Initialize forward metadata hence all layers in the forward pass can reuse it."""
|
"""Initialize forward metadata hence all layers in the forward pass can reuse it."""
|
||||||
metadata = FlashAttentionMetadata()
|
metadata = FlashAttentionMetadata()
|
||||||
@@ -285,6 +333,14 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[
|
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[
|
||||||
forward_batch.req_pool_indices, : metadata.max_seq_len_k
|
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
|
# TODO: we need to test this part for llama 4 eagle case
|
||||||
self._maybe_init_local_attn_metadata(forward_batch, metadata, device)
|
self._maybe_init_local_attn_metadata(forward_batch, metadata, device)
|
||||||
elif forward_batch.forward_mode.is_target_verify():
|
elif forward_batch.forward_mode.is_target_verify():
|
||||||
@@ -1059,6 +1115,15 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Default: single-token self-attention
|
# 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(
|
result = flash_attn_with_kvcache(
|
||||||
q=q_reshaped,
|
q=q_reshaped,
|
||||||
k_cache=key_cache,
|
k_cache=key_cache,
|
||||||
@@ -1076,6 +1141,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
return_softmax_lse=use_cascade_attn,
|
return_softmax_lse=use_cascade_attn,
|
||||||
num_splits=self.num_splits,
|
num_splits=self.num_splits,
|
||||||
ver=self.fa_impl_ver,
|
ver=self.fa_impl_ver,
|
||||||
|
scheduler_metadata=sched_meta,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
if use_cascade_attn:
|
if use_cascade_attn:
|
||||||
@@ -1220,6 +1286,16 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
0, self.max_context_len, self.page_size, device=self.device
|
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
|
# Only allocate local attention buffers if local attention is enabled
|
||||||
# This prevents OOM errors when local attention is not being used
|
# This prevents OOM errors when local attention is not being used
|
||||||
if self.has_local_attention:
|
if self.has_local_attention:
|
||||||
@@ -1589,6 +1665,20 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
|
|
||||||
self._maybe_update_local_attn_metadata_for_capture(metadata, batch_size)
|
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():
|
elif forward_mode.is_target_verify():
|
||||||
if self.topk <= 1:
|
if self.topk <= 1:
|
||||||
metadata.cache_seqlens_int32 = self.target_verify_metadata[
|
metadata.cache_seqlens_int32 = self.target_verify_metadata[
|
||||||
@@ -1855,6 +1945,23 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata,
|
metadata,
|
||||||
bs,
|
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():
|
elif forward_mode.is_target_verify():
|
||||||
if self.topk <= 1:
|
if self.topk <= 1:
|
||||||
metadata = self.target_verify_metadata[bs]
|
metadata = self.target_verify_metadata[bs]
|
||||||
|
|||||||
Reference in New Issue
Block a user