diff --git a/tests/v1/attention/test_chunked_local_attention.py b/tests/v1/attention/test_chunked_local_attention.py index 0826b7073080..2b0fd1c337dd 100644 --- a/tests/v1/attention/test_chunked_local_attention.py +++ b/tests/v1/attention/test_chunked_local_attention.py @@ -8,7 +8,11 @@ from tests.v1.attention.utils import BatchSpec, create_common_attn_metadata from vllm.platforms import current_platform -from vllm.v1.attention.backends.utils import make_local_attention_virtual_batches +from vllm.utils.math_utils import cdiv +from vllm.v1.attention.backends.utils import ( + make_local_attention_virtual_batches, + max_local_attention_virtual_batches, +) @dataclass @@ -203,3 +207,57 @@ def test_local_attention_virtual_batches(test_data: LocalAttentionTestData): print(f"Actual block table:\n{result.block_table_tensor}") torch.testing.assert_close(result.block_table_tensor, expected_block_table_tensor) + + +@pytest.mark.parametrize( + "query_lens,seq_lens,attn_chunk_size,max_num_seqs,expected_num_reqs", + [ + # A single prefill longer than the chunk size, with max_num_seqs=1: + # the condition that overflowed FlashInfer's per-request buffers in + # https://github.com/vllm-project/vllm/issues/49980. + ([1000], [1000], 256, 1, 4), + # Partially computed context, so the first local block is partial. + ([1000], [1255], 256, 1, 5), + # Several requests, each spanning several chunks. + ([300, 700, 90], [300, 1400, 90], 128, 4, 10), + # Chunk size larger than every sequence: no extra virtual batches. + ([64, 64], [64, 64], 256, 2, 2), + # Decodes: one virtual batch each, so the token-count cap binds exactly. + ([1, 1, 1, 1], [16, 32, 48, 64], 16, 4, 4), + ], +) +def test_max_local_attention_virtual_batches_bounds_num_reqs( + query_lens: list[int], + seq_lens: list[int], + attn_chunk_size: int, + max_num_seqs: int, + expected_num_reqs: int, +): + """The virtual batch count must never exceed the advertised upper bound. + + Attention backends preallocate per-request buffers from this bound, so an + underestimate overflows or silently truncates them. `expected_num_reqs` + pins the split itself, so a bound that collapsed back to `max_num_seqs` + fails here rather than passing vacuously. + """ + block_size = 16 + common_attn_metadata = create_common_attn_metadata( + BatchSpec(query_lens=query_lens, seq_lens=seq_lens), + block_size, + torch.device("cpu"), + ) + result, _ = make_local_attention_virtual_batches( + attn_chunk_size, common_attn_metadata, block_size + ) + assert result.num_reqs == expected_num_reqs + + bound = max_local_attention_virtual_batches( + attn_chunk_size, max_num_seqs, sum(query_lens) + ) + assert result.num_reqs <= bound + + # Total pages must also fit `bound * pages_per_virtual_batch`, which is how + # the FlashInfer builder sizes `paged_kv_indices`. + pages_per_virtual_batch = cdiv(attn_chunk_size, block_size) + num_pages = sum(cdiv(int(k), block_size) for k in result.seq_lens) + assert num_pages <= bound * pages_per_virtual_batch diff --git a/tests/v1/attention/test_flashinfer_chunked_local_attention.py b/tests/v1/attention/test_flashinfer_chunked_local_attention.py new file mode 100644 index 000000000000..e3f158b6e7e4 --- /dev/null +++ b/tests/v1/attention/test_flashinfer_chunked_local_attention.py @@ -0,0 +1,231 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""FlashInfer paged-KV buffer sizing under chunked local attention.""" + +import unittest.mock + +import numpy as np +import pytest +import torch + +from vllm.platforms import current_platform + +if not current_platform.is_cuda(): + pytest.skip("FlashInfer backend requires a CUDA platform.", allow_module_level=True) + +from tests.v1.attention.utils import ( # noqa: E402 + BatchSpec, + create_common_attn_metadata, + create_vllm_config, +) +from vllm.config import set_current_vllm_config # noqa: E402 +from vllm.model_executor.layers.attention.chunked_local_attention import ( # noqa: E402 + create_chunked_local_attention_backend, +) +from vllm.v1.attention.backends.flashinfer import FlashInferBackend # noqa: E402 +from vllm.v1.attention.backends.utils import ( # noqa: E402 + PerLayerParameters, + make_local_attention_virtual_batches, + split_decodes_and_prefills, +) +from vllm.v1.kv_cache_interface import ( # noqa: E402 + ChunkedLocalAttentionSpec, + FullAttentionSpec, +) + +ATTN_CHUNK_SIZE = 256 +BLOCK_SIZE = 16 +# Longer than ATTN_CHUNK_SIZE, so the request is split into several virtual +# batches; with max_num_seqs=1 that is what overflowed the buffers. +QUERY_LEN = 1000 +MAX_NUM_SEQS = 1 + + +def _mock_get_per_layer_parameters(vllm_config, layer_names, impl_cls): + head_size = vllm_config.model_config.get_head_size() + return { + name: PerLayerParameters( + window_left=-1, + logits_soft_cap=0.0, + sm_scale=1.0 / (head_size**0.5), + ) + for name in layer_names + } + + +def _build_builder(vllm_config, kv_cache_spec): + backend = create_chunked_local_attention_backend(FlashInferBackend, ATTN_CHUNK_SIZE) + with ( + set_current_vllm_config(vllm_config), + unittest.mock.patch( + "vllm.v1.attention.backends.flashinfer.get_per_layer_parameters", + _mock_get_per_layer_parameters, + ), + ): + # Buffer sizing happens in __init__ and is device-independent, so the + # test stays on CPU and needs no particular GPU. + return backend.get_builder_cls()( + kv_cache_spec, ["layer.0"], vllm_config, torch.device("cpu") + ) + + +def _make_specs(vllm_config): + """Both specs a chunked-local layer can reach a builder with. + + With the hybrid KV cache manager enabled (the default on CUDA) the builder + sees a `ChunkedLocalAttentionSpec`. When it is disabled, the spec is promoted + to a `FullAttentionSpec` that keeps `attention_chunk_size` set. + """ + common = dict( + block_size=BLOCK_SIZE, + num_kv_heads=vllm_config.model_config.get_num_kv_heads( + vllm_config.parallel_config + ), + head_size=vllm_config.model_config.get_head_size(), + dtype=vllm_config.model_config.dtype, + ) + return { + "chunked_local": ChunkedLocalAttentionSpec( + attention_chunk_size=ATTN_CHUNK_SIZE, **common + ), + "promoted_full": FullAttentionSpec( + attention_chunk_size=ATTN_CHUNK_SIZE, **common + ), + } + + +@pytest.mark.parametrize("spec_name", ["chunked_local", "promoted_full"]) +def test_paged_kv_buffers_fit_local_attention_virtual_batches(spec_name: str): + """Regression test for https://github.com/vllm-project/vllm/issues/49980. + + `make_local_attention_virtual_batches` reports a `num_reqs` equal to the + virtual batch count, which is decoupled from `max_num_seqs`. Sizing the + paged-KV buffers from `max_num_seqs` made the cumsum in + `_compute_flashinfer_kv_metadata` raise "provided out is the wrong size for + the accumulation" whenever a prefill exceeded `attention_chunk_size`. + """ + vllm_config = create_vllm_config( + max_model_len=2048, + block_size=BLOCK_SIZE, + max_num_seqs=MAX_NUM_SEQS, + max_num_batched_tokens=2048, + ) + builder = _build_builder(vllm_config, _make_specs(vllm_config)[spec_name]) + + common_attn_metadata = create_common_attn_metadata( + BatchSpec(query_lens=[QUERY_LEN], seq_lens=[QUERY_LEN]), + BLOCK_SIZE, + torch.device("cpu"), + ) + local_metadata, _ = make_local_attention_virtual_batches( + ATTN_CHUNK_SIZE, common_attn_metadata, BLOCK_SIZE + ) + num_reqs = local_metadata.num_reqs + # Guard against the test passing vacuously if the split ever stops + # inflating the request count. + assert num_reqs > MAX_NUM_SEQS + + # The exact operation that raised before the fix. + seq_lens_np = local_metadata.seq_lens.numpy() + num_blocks_np = (seq_lens_np + BLOCK_SIZE - 1) // BLOCK_SIZE + np.cumsum( + num_blocks_np, + dtype=np.int32, + out=builder.paged_kv_indptr.np[1 : num_reqs + 1], + ) + + assert builder.paged_kv_last_page_len.np.shape[0] >= num_reqs + num_actual_pages = int(builder.paged_kv_indptr.np[num_reqs]) + assert builder.paged_kv_indices.shape[0] >= num_actual_pages + + +def test_xqa_decode_mask_covers_local_attention_virtual_decodes(): + """The uniform XQA draft mask must have a row for every virtual decode.""" + vllm_config = create_vllm_config( + max_model_len=2048, + block_size=BLOCK_SIZE, + max_num_seqs=MAX_NUM_SEQS, + max_num_batched_tokens=2048, + ) + builder = _build_builder(vllm_config, _make_specs(vllm_config)["promoted_full"]) + + # A 4-token speculative verify window split evenly by a chunk boundary. + common_attn_metadata = create_common_attn_metadata( + BatchSpec(query_lens=[4], seq_lens=[ATTN_CHUNK_SIZE + 2]), + BLOCK_SIZE, + torch.device("cpu"), + ) + local_metadata, _ = make_local_attention_virtual_batches( + ATTN_CHUNK_SIZE, common_attn_metadata, BLOCK_SIZE + ) + assert local_metadata.query_start_loc_cpu.diff().tolist() == [2, 2] + num_decodes, _, _, _ = split_decodes_and_prefills( + local_metadata, decode_threshold=4 + ) + assert num_decodes > MAX_NUM_SEQS + + mask = builder._get_decode_mask(2, None, num_decodes, causal=True) + assert mask.shape[0] == num_decodes + + +def test_chunked_local_sizing_never_shrinks_full_attention_capacity(): + """`attention_chunk_size` on a merged spec must not shrink the allocation. + + When the hybrid KV cache manager is disabled, every Llama-4 layer becomes a + `FullAttentionSpec` and they merge into one KV cache group whose spec keeps + `attention_chunk_size`. The global attention layers form their own attention + group but share that spec, and they attend over the whole sequence, so + `paged_kv_indices` must still cover `max_num_seqs * max_num_pages_per_req`. + """ + max_num_seqs = 4 + max_model_len = 8192 + vllm_config = create_vllm_config( + max_model_len=max_model_len, + block_size=BLOCK_SIZE, + max_num_seqs=max_num_seqs, + max_num_batched_tokens=max_model_len, + ) + kv_cache_spec = FullAttentionSpec( + block_size=BLOCK_SIZE, + num_kv_heads=vllm_config.model_config.get_num_kv_heads( + vllm_config.parallel_config + ), + head_size=vllm_config.model_config.get_head_size(), + dtype=vllm_config.model_config.dtype, + attention_chunk_size=ATTN_CHUNK_SIZE, + ) + builder = _build_builder(vllm_config, kv_cache_spec) + + full_attention_pages = max_num_seqs * -(-max_model_len // BLOCK_SIZE) + assert builder.paged_kv_indices.shape[0] >= full_attention_pages + + +def test_full_attention_buffer_sizing_is_unchanged(): + """A spec without `attention_chunk_size` must keep the original sizing.""" + vllm_config = create_vllm_config( + max_model_len=2048, block_size=BLOCK_SIZE, max_num_seqs=8 + ) + kv_cache_spec = FullAttentionSpec( + block_size=BLOCK_SIZE, + num_kv_heads=vllm_config.model_config.get_num_kv_heads( + vllm_config.parallel_config + ), + head_size=vllm_config.model_config.get_head_size(), + dtype=vllm_config.model_config.dtype, + ) + backend = create_chunked_local_attention_backend(FlashInferBackend, ATTN_CHUNK_SIZE) + with ( + set_current_vllm_config(vllm_config), + unittest.mock.patch( + "vllm.v1.attention.backends.flashinfer.get_per_layer_parameters", + _mock_get_per_layer_parameters, + ), + ): + builder = backend.get_builder_cls()( + kv_cache_spec, ["layer.0"], vllm_config, torch.device("cpu") + ) + + max_num_pages_per_req = -(-2048 // BLOCK_SIZE) + assert builder.paged_kv_indptr.np.shape[0] == 8 + 1 + assert builder.paged_kv_last_page_len.np.shape[0] == 8 + assert builder.paged_kv_indices.shape[0] == 8 * max_num_pages_per_req diff --git a/vllm/v1/attention/backends/flashinfer.py b/vllm/v1/attention/backends/flashinfer.py index 3113c261dd98..9dbc2a2d0603 100755 --- a/vllm/v1/attention/backends/flashinfer.py +++ b/vllm/v1/attention/backends/flashinfer.py @@ -75,6 +75,7 @@ get_per_layer_parameters, infer_global_hyperparameters, log2_lse_to_ln, + max_local_attention_virtual_batches, split_decodes_and_prefills, ) from vllm.v1.attention.ops.dcp import ( @@ -721,8 +722,26 @@ def __init__( self.model_config.max_model_len, self.kv_cache_spec.block_size ) max_num_reqs = vllm_config.scheduler_config.max_num_seqs - self.max_num_reqs = max_num_reqs max_num_pages = max_num_reqs * max_num_pages_per_req + # Chunked local attention emits one virtual batch per local attention + # block, so `build()` sees a `num_reqs` far above `max_num_seqs`. These + # buffers back fixed-address CUDA graph buffers, so size for that worst + # case rather than reallocating on overflow. A merged spec can carry + # `attention_chunk_size` for a group whose layers use full attention, + # so only ever grow the allocation. + attn_chunk_size = getattr(self.kv_cache_spec, "attention_chunk_size", None) + if attn_chunk_size is None: + self.max_buffer_reqs = max_num_reqs + else: + self.max_buffer_reqs = max_local_attention_virtual_batches( + attn_chunk_size, max_num_reqs, self.max_num_batched_tokens + ) + # Each virtual batch attends to at most `attn_chunk_size` KV tokens. + max_num_pages = max( + max_num_pages, + self.max_buffer_reqs + * cdiv(attn_chunk_size, self.kv_cache_spec.block_size), + ) # Persistent uniform masks keep stable addresses for CUDA graphs. self._decode_mask_cache: dict[tuple[int, bool], torch.Tensor] = {} speculative_config = vllm_config.speculative_config @@ -917,13 +936,19 @@ def __init__( ) # Preparing persistent buffers self.paged_kv_indptr = CpuGpuBuffer( - max_num_reqs + 1, dtype=torch.int32, device=self.device, pin_memory=False + self.max_buffer_reqs + 1, + dtype=torch.int32, + device=self.device, + pin_memory=False, ) self.paged_kv_indices = torch.zeros( max_num_pages, dtype=torch.int32, device=self.device ) self.paged_kv_last_page_len = CpuGpuBuffer( - max_num_reqs, dtype=torch.int32, device=self.device, pin_memory=False + self.max_buffer_reqs, + dtype=torch.int32, + device=self.device, + pin_memory=False, ) @property @@ -1111,7 +1136,9 @@ def _get_decode_mask( if buf is None: per_req = _make_xqa_draft_block_mask(q_len_per_req, causal, self.device) buf = ( - per_req.unsqueeze(0).expand(self.max_num_reqs, -1, -1).contiguous() + per_req.unsqueeze(0) + .expand(self.max_buffer_reqs, -1, -1) + .contiguous() ) self._decode_mask_cache[key] = buf return buf[:num_decodes] @@ -1245,6 +1272,37 @@ def _get_cascade_wrapper(self): ) return self._cascade_wrapper + def _ensure_paged_kv_capacity(self, num_reqs: int, num_pages: int = 0) -> None: + """Grow the paged-KV buffers if a batch outgrows the `__init__` sizing. + + CUDA-graph decode wrappers alias these buffers and `fast_decode_plan` does + not copy into them, so rebinding them would leave those wrappers stale. + """ + grow_reqs = num_reqs > self.max_buffer_reqs + grow_pages = num_pages > self.paged_kv_indices.shape[0] + if not grow_reqs and not grow_pages: + return + if self.enable_cuda_graph: + raise ValueError( + f"FlashInfer paged-KV buffers hold {self.max_buffer_reqs} " + f"requests / {self.paged_kv_indices.shape[0]} pages but the " + f"batch needs {num_reqs} / {num_pages}, and they cannot be " + "grown while CUDA-graph decode wrappers alias them." + ) + if grow_reqs: + self.max_buffer_reqs = num_reqs + self.paged_kv_indptr = CpuGpuBuffer( + num_reqs + 1, dtype=torch.int32, device=self.device, pin_memory=False + ) + self.paged_kv_last_page_len = CpuGpuBuffer( + num_reqs, dtype=torch.int32, device=self.device, pin_memory=False + ) + self._decode_mask_cache.clear() + if grow_pages: + self.paged_kv_indices = torch.zeros( + num_pages, dtype=torch.int32, device=self.device + ) + def _compute_flashinfer_kv_metadata( self, num_blocks_np: np.ndarray, @@ -1268,6 +1326,7 @@ def _compute_flashinfer_kv_metadata( # write self.paged_kv_indices inplace num_actual_pages = self.paged_kv_indptr.np[num_reqs] + self._ensure_paged_kv_capacity(num_reqs, int(num_actual_pages)) paged_kv_indices = self.paged_kv_indices[:num_actual_pages] _copy_page_indices_kernel[(num_reqs,)]( paged_kv_indices, @@ -1294,6 +1353,7 @@ def build( fast_build: bool = False, ) -> FlashInferMetadata: num_reqs = common_attn_metadata.num_reqs + self._ensure_paged_kv_capacity(num_reqs) num_actual_tokens = common_attn_metadata.num_actual_tokens causal = common_attn_metadata.causal route_decode = causal or self.use_xqa diff --git a/vllm/v1/attention/backends/utils.py b/vllm/v1/attention/backends/utils.py index 3413c5ad00ca..c0e69a9a4bb4 100644 --- a/vllm/v1/attention/backends/utils.py +++ b/vllm/v1/attention/backends/utils.py @@ -427,6 +427,39 @@ def infer_global_hyperparameters( return global_params +def max_local_attention_virtual_batches( + attn_chunk_size: int, + max_num_reqs: int, + max_num_batched_tokens: int, +) -> int: + """Upper bound on the ``num_reqs`` that local attention can produce. + + `make_local_attention_virtual_batches` emits one virtual batch per local + attention block, so the ``num_reqs`` it reports is decoupled from + `max_num_seqs`. Backends that preallocate per-request buffers must size them + from this bound instead. + + Each request contributes ``1 + cdiv(q_i - f_i, c)`` blocks with + ``f_i = q_tokens_in_first_block >= 1``, summing to at most + ``2 * max_num_reqs + cdiv(max_num_batched_tokens, attn_chunk_size)``; the + ``2 *`` term is required, not slack. Every virtual batch also owns at least + one query token, capping the count at ``max_num_batched_tokens``. + + Args: + attn_chunk_size: Local attention chunk size. + max_num_reqs: Maximum number of real requests in a batch. + max_num_batched_tokens: Maximum number of query tokens in a batch. + + Returns: + Maximum number of virtual batches. + + """ + return min( + 2 * max_num_reqs + cdiv(max_num_batched_tokens, attn_chunk_size), + max_num_batched_tokens, + ) + + # # Take in `query_start_loc_np` and `seq_lens_np` and break the sequences into # local attention blocks, where each block is passed to the attention kernel