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
136 changes: 136 additions & 0 deletions tests/kernels/attention/test_flashinfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -705,6 +705,142 @@ def test_flashinfer_decode_with_paged_fp8_kv(
)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="Requires CUDA")
def test_packed_flashinfer_indices_refresh_on_graph_replay():
from vllm.v1.attention.backends.flashinfer import _copy_page_indices_kernel
from vllm.v1.attention.ops.flashinfer_repage import (
PackedKVPageGeometry,
repage_block_table,
)

cache = torch.empty((4, 2, 2, 128, 128), device="cuda", dtype=torch.bfloat16)[:, 1]
geometry = PackedKVPageGeometry.from_cache(cache, 32)
storage = torch.full((3, 10), 123, device="cuda", dtype=torch.int32)
source = storage[1:, 1:9]
source[0] = torch.tensor([0, 1, 2, 3, 8, 9, -1, 11], device="cuda")
lengths = torch.tensor([256, 0], device="cuda", dtype=torch.int32)
destination = torch.empty((2, 9), device="cuda", dtype=torch.int32)
indptr = torch.tensor([0, 8, 8], device="cuda", dtype=torch.int32)
flat = torch.empty(8, device="cuda", dtype=torch.int32)

def run():
repage_block_table(source, lengths, destination, geometry)
_copy_page_indices_kernel[(2,)](
flat,
source,
source.stride(0),
indptr,
BLOCK_SIZE=32,
PAGES_PER_BLOCK=geometry.pages_per_block,
BLOCK_STRIDE_PAGES=geometry.block_stride_pages,
)

run()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
run()
source[0, 4] = 4
original = source.clone()
destination.fill_(-999)
flat.fill_(-999)
graph.replay()
expected = torch.tensor(
[0, 1, 2, 3, 16, 33, -1, 35], device="cuda", dtype=torch.int32
)
torch.testing.assert_close(flat, expected)
torch.testing.assert_close(destination[0, :8], flat)
assert torch.all(destination[1, :8] == 0)
torch.testing.assert_close(source, original)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="Requires CUDA")
@pytest.mark.parametrize(
"backend",
[
"fa2",
pytest.param(
"trtllm-gen",
marks=pytest.mark.skipif(
not current_platform.is_device_capability_family(100),
reason="Requires Blackwell",
),
),
],
)
@torch.inference_mode()
def test_packed_flashinfer_matches_dense_pages(backend):
from vllm.v1.attention.ops.flashinfer_repage import (
PackedKVPageGeometry,
repage_block_table,
)

set_random_seed(0)
manager_size, page_size, heads, dim = 592, 16, 2, 64
cache = torch.randn(
(3, 2, heads, manager_size, 2 * dim), device="cuda", dtype=torch.bfloat16
)[:, 1]
geometry = PackedKVPageGeometry.from_cache(cache, page_size)
packed = geometry.read_view(cache).split(dim, dim=-1)
ratio = manager_size // page_size
dense = (
cache.reshape(3, heads, ratio, page_size, 2 * dim)
.permute(0, 2, 1, 3, 4)
.reshape(3 * ratio, heads, page_size, 2 * dim)
.contiguous()
.split(dim, dim=-1)
)
lengths = [manager_size + 5, manager_size - 1]
indptr, _, last, source = _make_paged_kv_metadata(lengths, page_size, 3 * ratio)
source[:, 0] = 3 * ratio - 1
seq_lens = torch.tensor(lengths, device="cuda", dtype=torch.int32)
mapped = repage_block_table(source, seq_lens, torch.empty_like(source), geometry)
query = torch.randn((6, 4 * heads, dim), device="cuda", dtype=torch.bfloat16)
workspace = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda")

def attend(kv, table):
if backend == "trtllm-gen":
return flashinfer.decode.trtllm_batch_decode_with_kv_cache(
query,
kv,
workspace,
table,
seq_lens,
max(lengths),
bmm1_scale=dim**-0.5,
bmm2_scale=1.0,
backend=backend,
q_len_per_req=3,
)
wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(
workspace, "HND", backend=backend
)
indices = torch.cat(
[
table[i, : (length + page_size - 1) // page_size]
for i, length in enumerate(lengths)
]
)
wrapper.plan(
torch.tensor([0, 3, 6], device="cpu", dtype=torch.int32),
indptr,
indices,
last,
4 * heads,
heads,
dim,
page_size,
causal=False,
sm_scale=dim**-0.5,
q_data_type=query.dtype,
kv_data_type=query.dtype,
)
return wrapper.run(query, kv, return_lse=True)

torch.testing.assert_close(
attend(packed, mapped), attend(dense, source), atol=1e-2, rtol=1e-2
)


@pytest.mark.skipif(
not any(current_platform.is_device_capability_family(f) for f in (80, 90, 120)),
reason="NVFP4 fa2 path",
Expand Down
22 changes: 22 additions & 0 deletions tests/v1/core/test_contiguous_kv_packing.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
MambaSpec,
MLAAttentionSpec,
SlidingWindowMLASpec,
SlidingWindowSpec,
UniformTypeKVCacheSpecs,
iter_layer_specs,
)
Expand Down Expand Up @@ -427,6 +428,27 @@ def test_mamba_split_measures_the_block_the_other_groups_already_force(self, wid


class TestSlidingWindowBucketCap:
@pytest.mark.parametrize("layout", ["BLHNC", "BLNHC"])
def test_sliding_window_keeps_manager_size_when_kernels_use_small_pages(
self, layout
):
config = _mock_vllm_config(layout)
config.speculative_config = None
specs = {
"target": replace(_full(), block_size=640, head_size=256, head_size_v=256),
"draft.sliding": SlidingWindowSpec(
block_size=64,
num_kv_heads=8,
head_size=128,
dtype=torch.float16,
sliding_window=4096,
),
}
groups = _get_packed_kv_cache_groups(config, specs)
sliding = next(g for g in groups if "draft.sliding" in g.layer_names)
assert sliding.kv_cache_spec.block_size == 640
assert specs["draft.sliding"].block_size == 64

def test_sliding_window_bucket_is_capped_at_the_main_page(self):
"""DeepSeek-V4.1 shape: 43 SlidingWindowMLASpec SWA caches beside an
unbalanced 8-layer paged MLA bucket. Left whole, the SWA bucket would
Expand Down
58 changes: 58 additions & 0 deletions tests/v1/worker/test_attn_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,10 @@
never addressed by the logical view.
"""

from dataclasses import replace
from types import SimpleNamespace
from typing import Any
from weakref import ref

import numpy as np
import pytest
Expand Down Expand Up @@ -41,6 +43,7 @@
from vllm.v1.worker.utils import (
AttentionGroup,
allocate_kv_cache,
clear_layer_kv_caches,
copy_kv_cache_blocks_inplace,
)

Expand Down Expand Up @@ -568,6 +571,61 @@ def test_copy_kv_cache_blocks_shared_storage(layout: KVCacheLayout):
torch.testing.assert_close(cache[1], expected[layer_idx][1])


@pytest.mark.parametrize("padded", [False, True])
def test_packed_allocation_retains_manager_blocks(padded):
class Backend(AttentionBackend):
@classmethod
def get_kv_cache_view_block_size(cls, spec, kernel_block_size, vllm_config):
return spec.block_size

spec = FullAttentionSpec(
block_size=128, num_kv_heads=2, head_size=64, dtype=torch.bfloat16
)
block_stride = 2 * spec.page_size_bytes
layers = ["layer.0"] if padded else ["layer.0", "layer.1"]
if padded:
spec = replace(spec, page_size_padded=block_stride)
config = KVCacheConfig(
num_blocks=3,
kv_cache_tensors=[
KVCacheTensor(
size=3 * block_stride,
layers=layers,
layer_stride=spec.page_size_bytes,
block_stride=block_stride,
)
],
kv_cache_groups=[KVCacheGroupSpec(layers, spec)],
)
context = {
name: SimpleNamespace(get_attn_backend=lambda: Backend) for name in layers
}
caches = allocate_kv_cache(
config,
torch.device("cpu"),
KVCacheLayout.BLHNC,
[32],
vllm_config=SimpleNamespace(
compilation_config=SimpleNamespace(static_forward_context=context)
),
)
for cache in caches.values():
assert cache.shape == (3, 2, 128, 128)
assert cache.stride(0) * cache.element_size() == block_stride


def test_clear_layer_kv_caches_releases_flashinfer_read_views():
cache = torch.empty((2, 1, 128, 128), dtype=torch.bfloat16, device="cpu")
cache_ref = ref(cache)
layer = SimpleNamespace(
kv_cache=cache,
impl=SimpleNamespace(_repage_source=cache, _repage_view=cache[:, :, :32]),
)
del cache
clear_layer_kv_caches([layer])
assert cache_ref() is None


def test_fixed_block_stride_propagates_outward_in_lhbnc():
num_blocks = 3
num_layers = 2
Expand Down
14 changes: 14 additions & 0 deletions vllm/v1/attention/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -373,6 +373,20 @@ def supported_kv_cache_layouts(cls) -> tuple[KVCacheLayout, ...] | None:
None when the kernels consume any layout and express no preference."""
return None

@classmethod
def get_kv_cache_view_block_size(
cls,
spec: "KVCacheSpec",
kernel_block_size: int,
vllm_config: "VllmConfig",
) -> int:
"""Physical view size when a non-dense manager block needs splitting.

Backends retaining manager blocks must adapt their read indices and
views themselves. The common block table still uses kernel blocks.
"""
return kernel_block_size

@classmethod
def is_ssm(cls) -> bool:
return False
Expand Down
Loading
Loading