[Fix] Harden FlashAttention CUDA graph metadata bounds (#35454)

Co-authored-by: Yuwei An <ayw.sirius19@gmail.com>
This commit is contained in:
Aurick Qiao
2026-08-23 21:11:46 -07:00
committed by GitHub
co-authored by Yuwei An
parent 5b5b29d4e2
commit 3e30649064
2 changed files with 82 additions and 2 deletions
@@ -616,11 +616,14 @@ class FlashAttentionBackend(AttentionBackend):
# upper bound) so captured graphs keep a valid address; each replay
# refills a [:num_tokens] view.
if self.use_sliding_window_kv_pool:
assert forward_batch.out_cache_loc is not None
m.swa_page_table = torch.zeros(
(bs, self.max_num_pages), dtype=torch.int32, device=device
)
self.full_cg_prefill_swa_out_cache_loc = torch.zeros(
(self.max_context_len,), dtype=torch.int64, device=device
(forward_batch.out_cache_loc.shape[0],),
dtype=torch.int64,
device=device,
)
self.full_cg_prefill_metadata = m
m = self.full_cg_prefill_metadata
@@ -657,6 +660,11 @@ class FlashAttentionBackend(AttentionBackend):
# SWA write targets for the new tokens (KVWriteLoc.swa_loc), refilled
# into the pointer-stable buffer and bound as a [:num_tokens] view.
num_out = forward_batch.out_cache_loc.shape[0]
assert_buffer_fits(
num_out,
self.full_cg_prefill_swa_out_cache_loc.shape[0],
"full-CG prefill SWA write-location buffer",
)
self.full_cg_prefill_swa_out_cache_loc[:num_out].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
@@ -686,7 +694,7 @@ class FlashAttentionBackend(AttentionBackend):
seq_lens_cpu = forward_batch.seq_lens_cpu
eager_max_k = (
seq_lens_cpu.max().item()
if seq_lens_cpu is not None
if seq_lens_cpu is not None and seq_lens_cpu.numel() > 0
else self.max_context_len
)
@@ -0,0 +1,72 @@
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionBackend
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=2, stage="base-b", runner_config="1-gpu-small")
class TestFlashAttentionGraphMetadata(CustomTestCase):
def test_full_prefill_swa_buffer_uses_capture_token_capacity(self):
backend = FlashAttentionBackend.__new__(FlashAttentionBackend)
backend.full_cg_prefill_metadata = None
backend.use_sliding_window_kv_pool = True
backend.max_num_pages = 4
backend.max_context_len = 4
backend.page_size = 1
backend.req_to_token = torch.zeros((1, 4), dtype=torch.int32)
backend.token_to_kv_pool = SimpleNamespace(
translate_loc_from_full_to_swa=lambda locations: locations
)
forward_batch = SimpleNamespace(
batch_size=1,
seq_lens=torch.zeros(1, dtype=torch.int64),
extend_seq_lens=torch.zeros(1, dtype=torch.int64),
seq_lens_cpu=torch.zeros(1, dtype=torch.int64),
req_pool_indices=torch.zeros(1, dtype=torch.int64),
out_cache_loc=torch.arange(8, dtype=torch.int64),
positions=torch.arange(8, dtype=torch.int64),
)
backend._init_full_cg_prefill_metadata(forward_batch, in_capture=True)
self.assertEqual(backend.full_cg_prefill_swa_out_cache_loc.shape, (8,))
forward_batch.out_cache_loc = torch.arange(9, dtype=torch.int64)
with self.assertRaisesRegex(AssertionError, "used 9 > capacity 8"):
backend._init_full_cg_prefill_metadata(forward_batch, in_capture=False)
def test_empty_cpu_sequence_lengths_use_static_bound(self):
backend = FlashAttentionBackend.__new__(FlashAttentionBackend)
backend.max_context_len = 16
backend.req_to_token_pool = SimpleNamespace(
req_to_token=torch.zeros((1, 16), dtype=torch.int32)
)
backend.is_prefill_aware_swa = False
backend.has_swa = False
backend.use_sliding_window_kv_pool = False
backend._unified_dense = False
backend.page_size = 1
backend._compute_scheduler_metadata = lambda *_: None
backend._maybe_init_local_attn_metadata = lambda *_: None
forward_batch = SimpleNamespace(
forward_mode=ForwardMode.DECODE,
seq_lens=torch.empty(0, dtype=torch.int64),
batch_size=0,
seq_lens_cpu=torch.empty(0, dtype=torch.int64),
spec_info=None,
req_pool_indices=torch.empty(0, dtype=torch.int64),
encoder_lens=None,
)
backend.init_forward_metadata(forward_batch)
self.assertEqual(backend.forward_metadata.max_seq_len_k, 16)
if __name__ == "__main__":
unittest.main()