diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index ad7f59c0d..ab4cf841f 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -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]