Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 59 additions & 1 deletion tests/v1/attention/test_chunked_local_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
231 changes: 231 additions & 0 deletions tests/v1/attention/test_flashinfer_chunked_local_attention.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading