From b8339dfa9ab1df8a1a1839e699855c54f2087729 Mon Sep 17 00:00:00 2001 From: yunweili3 Date: Wed, 15 Jul 2026 15:10:52 -0700 Subject: [PATCH] Add paged-KV block_table bounds check in mha_fwd_kvcache The split-KV kernel (compute_attn_1rowblock_splitkv) indexes block_table[n_block * kBlockN / page_block_size], bounded only by actual_seqlen_k. In the kvcache path actual_seqlen_k is seqlens_k[b] + seqlen_knew, but block_table only has max_num_blocks_per_seq columns per sequence. If a caller passes a cache_seqlens (or appends new keys) exceeding max_num_blocks_per_seq * page_block_size, the kernel reads block_table out of bounds with no in-kernel check (see issue #2709). Validate the caller contract host-side and raise a clear error instead. The .max().item() sync is only paid on the paged-KV path. Add test_flash_attn_kvcache_paged_block_table_bounds covering both the cache-length overflow and the appended-new-keys overflow, plus a positive control exactly at capacity. --- csrc/flash_attn/flash_api.cpp | 16 +++++++++ tests/test_flash_attn.py | 63 +++++++++++++++++++++++++++++++++++ 2 files changed, 79 insertions(+) diff --git a/csrc/flash_attn/flash_api.cpp b/csrc/flash_attn/flash_api.cpp index ca974949740..cee5dc07450 100644 --- a/csrc/flash_attn/flash_api.cpp +++ b/csrc/flash_attn/flash_api.cpp @@ -1392,6 +1392,22 @@ mha_fwd_kvcache(at::Tensor &q, // batch_size x seqlen_q x num_he CHECK_DEVICE(seqlens_k); CHECK_CONTIGUOUS(seqlens_k); CHECK_SHAPE(seqlens_k, batch_size); + // Defense-in-depth for the paged KV cache. The split-KV kernel indexes block_table with + // block_table[n_block * kBlockN / page_block_size], bounded only by actual_seqlen_k, which + // in this path is seqlens_k[b] + seqlen_knew (leftpad_k is disallowed with paged KV below). + // block_table only has max_num_blocks_per_seq entries per sequence, so if any sequence length + // exceeds max_num_blocks_per_seq * page_block_size the kernel reads block_table out of bounds. + // The kernel itself does no such check, so validate the caller contract here. + // Note: .max().item() forces a device->host sync, so we only pay it for the paged KV case. + if (paged_KV) { + const int seqlen_knew = k_.has_value() ? k.size(1) : 0; + const int max_seqlen_k = seqlens_k.max().item() + seqlen_knew; + TORCH_CHECK(max_seqlen_k <= max_num_blocks_per_seq * page_block_size, + "Paged KV cache: max(seqlens_k)", seqlen_knew > 0 ? " + seqlen_knew" : "", " (= ", max_seqlen_k, + ") exceeds the capacity addressable by block_table (max_num_blocks_per_seq * page_block_size = ", + max_num_blocks_per_seq * page_block_size, "). Allocate more columns in block_table, otherwise the " + "kernel would index block_table out of bounds."); + } params.cu_seqlens_k = static_cast(seqlens_k.data_ptr()); } params.is_seqlens_k_cumulative = !(seqlens_k_.has_value()); diff --git a/tests/test_flash_attn.py b/tests/test_flash_attn.py index 0589d1b2cd9..b62f3c98f81 100644 --- a/tests/test_flash_attn.py +++ b/tests/test_flash_attn.py @@ -2581,3 +2581,66 @@ def test_flash_attn_varlen_paged_kv_num_splits(dtype): with pytest.raises(RuntimeError, match="num_splits > 1 is not supported"): _flash_attn_varlen_forward(q, k_cache, v_cache, **fwd_kwargs, num_splits=2) + + +@pytest.mark.parametrize("dtype", [torch.float16]) +@pytest.mark.parametrize("paged_kv_block_size", [256]) +@pytest.mark.parametrize("append_knew", [False, True]) +def test_flash_attn_kvcache_paged_block_table_bounds(append_knew, paged_kv_block_size, dtype): + # Regression test for the paged-KV out-of-bounds guard (issue #2709). + # block_table only has `max_num_blocks_per_seq` columns, so the split-KV kernel can + # only safely index it up to max_num_blocks_per_seq * page_block_size tokens. If any + # cache_seqlens[b] (+ appended new keys) exceeds that capacity, mha_fwd_kvcache must + # raise instead of letting the kernel read block_table out of bounds. + device = "cuda" + batch_size = 1 + nheads = 1 + d = 64 + max_num_blocks_per_seq = 1 + capacity = max_num_blocks_per_seq * paged_kv_block_size + + # A pool of pages large enough that the block_table indices are always valid; + # the guard must fire on the sequence length, not on missing pages. + num_blocks = 4 + k_cache_paged = torch.randn(num_blocks, paged_kv_block_size, nheads, d, device=device, dtype=dtype) + v_cache_paged = torch.randn(num_blocks, paged_kv_block_size, nheads, d, device=device, dtype=dtype) + block_table = torch.zeros(batch_size, max_num_blocks_per_seq, dtype=torch.int32, device=device) + + q = torch.randn(batch_size, 1, nheads, d, device=device, dtype=dtype) + + if append_knew: + # cache is full at capacity, appending even one new key overflows the block_table. + seqlen_knew = 1 + k_new = torch.randn(batch_size, seqlen_knew, nheads, d, device=device, dtype=dtype) + v_new = torch.randn(batch_size, seqlen_knew, nheads, d, device=device, dtype=dtype) + cache_seqlens = torch.full((batch_size,), capacity, dtype=torch.int32, device=device) + else: + seqlen_knew = 0 + k_new = None + v_new = None + cache_seqlens = torch.full((batch_size,), capacity + 1, dtype=torch.int32, device=device) + + with pytest.raises(RuntimeError, match="block_table"): + flash_attn_with_kvcache( + q, + k_cache_paged, + v_cache_paged, + k=k_new, + v=v_new, + cache_seqlens=cache_seqlens, + block_table=block_table, + causal=False, + ) + + # Positive control: exactly at capacity (and no appended keys) must NOT raise. + cache_seqlens_ok = torch.full((batch_size,), capacity, dtype=torch.int32, device=device) + out = flash_attn_with_kvcache( + q, + k_cache_paged, + v_cache_paged, + cache_seqlens=cache_seqlens_ok, + block_table=block_table, + causal=False, + ) + assert out.shape == (batch_size, 1, nheads, d) + assert not out.isnan().any()