diff --git a/tests/v1/core/test_kv_cache_utils.py b/tests/v1/core/test_kv_cache_utils.py index 3be24d7fb34..3f5b7a12433 100644 --- a/tests/v1/core/test_kv_cache_utils.py +++ b/tests/v1/core/test_kv_cache_utils.py @@ -117,6 +117,7 @@ def new_kv_cache_spec( page_size_padded=None, sliding_window=None, attention_chunk_size=None, + indexes_kv_by_block_stride=False, ): return FullAttentionSpec( block_size=block_size, @@ -126,6 +127,7 @@ def new_kv_cache_spec( page_size_padded=page_size_padded, sliding_window=sliding_window, attention_chunk_size=attention_chunk_size, + indexes_kv_by_block_stride=indexes_kv_by_block_stride, ) @@ -136,6 +138,7 @@ def new_sliding_window_spec( dtype=torch.float32, page_size_padded=None, sliding_window=1, + indexes_kv_by_block_stride=False, ): return SlidingWindowSpec( block_size=block_size, @@ -144,6 +147,7 @@ def new_sliding_window_spec( dtype=dtype, page_size_padded=page_size_padded, sliding_window=sliding_window, + indexes_kv_by_block_stride=indexes_kv_by_block_stride, ) @@ -1799,16 +1803,38 @@ def test_get_kv_cache_config_one_worker(): ], ) - # different hidden size that cannot be aligned by using different block size + # different hidden size that cannot be aligned by using different block size, + # but can be aligned by padding the smaller physical page. + swa_spec = new_sliding_window_spec(head_size=96, indexes_kv_by_block_stride=True) kv_cache_specs_hybrid = { - "layer_1": new_kv_cache_spec(head_size=64), - "layer_2": new_sliding_window_spec(head_size=96), + "layer_1": new_kv_cache_spec(head_size=64, indexes_kv_by_block_stride=True), + "layer_2": swa_spec, } - with pytest.raises(NotImplementedError): - get_kv_cache_configs( - vllm_config, [kv_cache_specs_hybrid], [mem_per_block_per_layer * 2 * 32] - )[0] + kv_cache_config_hybrid = get_kv_cache_configs( + vllm_config, [kv_cache_specs_hybrid], [mem_per_block_per_layer * 2 * 32] + )[0] + padded_page_size = swa_spec.page_size_bytes + assert kv_cache_config_hybrid == KVCacheConfig( + num_blocks=42, + kv_cache_tensors=[ + KVCacheTensor(size=padded_page_size * 42, shared_by=["layer_1", "layer_2"]), + ], + kv_cache_groups=[ + KVCacheGroupSpec( + ["layer_1"], + new_kv_cache_spec( + head_size=64, + page_size_padded=padded_page_size, + indexes_kv_by_block_stride=True, + ), + ), + KVCacheGroupSpec( + ["layer_2"], + new_sliding_window_spec(head_size=96, indexes_kv_by_block_stride=True), + ), + ], + ) # Test num_gpu_blocks_override vllm_config.cache_config.num_gpu_blocks_override = 16 @@ -2322,6 +2348,75 @@ def test_check_enough_kv_cache_memory_respects_num_gpu_blocks_override(): get_kv_cache_configs(vllm_config, [kv_cache_specs], [large_available_memory]) +def test_unify_kv_cache_page_size_uses_padding_for_non_divisible_sizes(): + """DFlash drafters can have a smaller head size than the target model. + + For example, MiMo uses 192-dim target KV heads while its DFlash draft uses + 128-dim KV heads. The resulting page sizes are 3:2 rather than an integer + block-size multiple, so the smaller page must be padded instead. + """ + # Both layers' backends opt into the padded-page strided view (e.g. + # FlashAttention / its DiffKV subclass), so padding is allowed. + target_spec = new_kv_cache_spec( + block_size=16, + num_kv_heads=1, + head_size=192, + dtype=torch.bfloat16, + indexes_kv_by_block_stride=True, + ) + draft_spec = new_sliding_window_spec( + block_size=16, + num_kv_heads=1, + head_size=128, + dtype=torch.bfloat16, + sliding_window=1024, + indexes_kv_by_block_stride=True, + ) + + unified_specs = kv_cache_utils.unify_kv_cache_spec_page_size( + { + "target_attn": target_spec, + "draft_attn": draft_spec, + } + ) + + assert unified_specs["target_attn"] == target_spec + unified_draft_spec = unified_specs["draft_attn"] + assert unified_draft_spec.block_size == draft_spec.block_size + assert unified_draft_spec.real_page_size_bytes == draft_spec.real_page_size_bytes + assert unified_draft_spec.page_size_padded == target_spec.page_size_bytes + assert unified_draft_spec.page_size_bytes == target_spec.page_size_bytes + + +def test_unify_kv_cache_page_size_padding_requires_backend_support(): + """Padding is gated on the backend declaring ``indexes_kv_by_block_stride``. + + A backend that does not support the strided padded-page view must raise + rather than silently padding (and misreading KV at runtime). + """ + target_spec = new_kv_cache_spec( + block_size=16, + num_kv_heads=1, + head_size=192, + dtype=torch.bfloat16, + indexes_kv_by_block_stride=True, + ) + # The non-divisible draft layer needs padding but its backend does not + # support the strided padded-page view -> must raise, not silently pad. + draft_spec = new_sliding_window_spec( + block_size=16, + num_kv_heads=1, + head_size=128, + dtype=torch.bfloat16, + sliding_window=1024, + indexes_kv_by_block_stride=False, + ) + specs = {"target_attn": target_spec, "draft_attn": draft_spec} + + with pytest.raises(NotImplementedError): + kv_cache_utils.unify_kv_cache_spec_page_size(specs) + + def test_unify_hybrid_kv_cache_specs(): # 1. has_full_attention and has_sliding_window before_spec_1 = new_kv_cache_spec() diff --git a/tests/v1/worker/test_attn_utils.py b/tests/v1/worker/test_attn_utils.py new file mode 100644 index 00000000000..7e65d650f7e --- /dev/null +++ b/tests/v1/worker/test_attn_utils.py @@ -0,0 +1,242 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import torch + +from vllm.v1.kv_cache_interface import FullAttentionSpec, KVQuantMode +from vllm.v1.worker.gpu.attn_utils import _reshape_kv_cache +from vllm.v1.worker.utils import AttentionGroup + + +class FakeFlashAttentionBackend: + @staticmethod + def get_kv_cache_shape( + num_blocks: int, + block_size: int, + num_kv_heads: int, + head_size: int, + cache_dtype_str: str = "auto", + ) -> tuple[int, ...]: + return (num_blocks, 2, block_size, num_kv_heads, head_size) + + @staticmethod + def get_kv_cache_stride_order( + include_num_layers_dimension: bool = False, + ) -> tuple[int, ...]: + assert not include_num_layers_dimension + return (0, 1, 2, 3, 4) + + +class FakeHNDFlashAttentionBackend(FakeFlashAttentionBackend): + @staticmethod + def get_kv_cache_stride_order( + include_num_layers_dimension: bool = False, + ) -> tuple[int, ...]: + assert not include_num_layers_dimension + return (0, 1, 3, 2, 4) + + +def test_reshape_padded_flash_attention_kv_cache_strides_by_page(): + num_blocks = 3 + spec = FullAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=2, + dtype=torch.float32, + page_size_padded=384, + ) + assert spec.real_page_size_bytes == 256 + + raw_tensors = { + "layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8) + } + attn_groups = [ + AttentionGroup( + backend=FakeFlashAttentionBackend, + layer_names=["layer"], + kv_cache_spec=spec, + kv_cache_group_id=0, + ) + ] + + kv_cache = _reshape_kv_cache( + attn_groups, + raw_tensors, + "auto", + [spec.block_size], + {}, + )["layer"] + + assert kv_cache.shape == (num_blocks, 2, 16, 1, 2) + assert kv_cache.stride(0) == spec.page_size_bytes // 4 + assert kv_cache.stride(1) == spec.real_page_size_bytes // 2 // 4 + assert kv_cache[1, 0].storage_offset() == spec.page_size_bytes // 4 + assert ( + kv_cache[1, 1].storage_offset() + == (spec.page_size_bytes + spec.real_page_size_bytes // 2) // 4 + ) + + +def test_reshape_padded_hnd_flash_attention_kv_cache_strides_by_page(): + num_blocks = 3 + spec = FullAttentionSpec( + block_size=16, + num_kv_heads=3, + head_size=2, + dtype=torch.float32, + page_size_padded=1024, + ) + assert spec.real_page_size_bytes == 768 + + raw_tensors = { + "layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8) + } + attn_groups = [ + AttentionGroup( + backend=FakeHNDFlashAttentionBackend, + layer_names=["layer"], + kv_cache_spec=spec, + kv_cache_group_id=0, + ) + ] + + kv_cache = _reshape_kv_cache( + attn_groups, + raw_tensors, + "auto", + [spec.block_size], + {}, + )["layer"] + + assert kv_cache.shape == (num_blocks, 2, 16, 3, 2) + assert kv_cache.stride(0) == spec.page_size_bytes // 4 + assert kv_cache.stride(1) == spec.real_page_size_bytes // 2 // 4 + assert kv_cache.stride(2) == 2 + assert kv_cache.stride(3) == spec.block_size * spec.head_size + assert kv_cache[1, 0].storage_offset() == spec.page_size_bytes // 4 + assert ( + kv_cache[1, 1].storage_offset() + == (spec.page_size_bytes + spec.real_page_size_bytes // 2) // 4 + ) + assert ( + kv_cache[1, 1, 3, 2].storage_offset() + == ( + spec.page_size_bytes + + spec.real_page_size_bytes // 2 + + 3 * spec.head_size * 4 + + 2 * spec.block_size * spec.head_size * 4 + ) + // 4 + ) + + +class FakeDiffKVBackend: + @staticmethod + def get_kv_cache_shape( + num_blocks: int, + block_size: int, + num_kv_heads: int, + head_size: int, + cache_dtype_str: str = "auto", + ) -> tuple[int, ...]: + return (num_blocks, block_size, num_kv_heads, head_size * 2) + + @staticmethod + def get_kv_cache_stride_order( + include_num_layers_dimension: bool = False, + ) -> tuple[int, ...]: + assert not include_num_layers_dimension + return (0, 1, 2, 3) + + +def test_reshape_padded_diff_kv_cache_does_not_infer_kv_dim(): + num_blocks = 3 + spec = FullAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=2, + dtype=torch.float32, + page_size_padded=384, + ) + + raw_tensors = { + "layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8) + } + attn_groups = [ + AttentionGroup( + backend=FakeDiffKVBackend, + layer_names=["layer"], + kv_cache_spec=spec, + kv_cache_group_id=0, + ) + ] + + kv_cache = _reshape_kv_cache( + attn_groups, + raw_tensors, + "auto", + [spec.block_size], + {}, + )["layer"] + + assert kv_cache.shape == (num_blocks, 16, 1, 4) + assert kv_cache.stride(0) == spec.page_size_bytes // 4 + assert kv_cache.stride(1) == 4 + + +class FakePerTokenScaleBackend: + @staticmethod + def get_kv_cache_shape( + num_blocks: int, + block_size: int, + num_kv_heads: int, + head_size: int, + cache_dtype_str: str = "auto", + ) -> tuple[int, ...]: + return (num_blocks, 2, block_size, num_kv_heads, head_size + 4) + + @staticmethod + def get_kv_cache_stride_order( + include_num_layers_dimension: bool = False, + ) -> tuple[int, ...]: + assert not include_num_layers_dimension + return (0, 1, 2, 3, 4) + + +def test_reshape_padded_quantized_kv_cache_preserves_scale_stride(): + num_blocks = 3 + spec = FullAttentionSpec( + block_size=16, + num_kv_heads=1, + head_size=4, + dtype=torch.int8, + kv_quant_mode=KVQuantMode.INT8_PER_TOKEN_HEAD, + page_size_padded=384, + ) + assert spec.real_page_size_bytes == 128 + assert spec.page_size_bytes == 384 + + raw_tensors = { + "layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8) + } + attn_groups = [ + AttentionGroup( + backend=FakePerTokenScaleBackend, + layer_names=["layer"], + kv_cache_spec=spec, + kv_cache_group_id=0, + ) + ] + + kv_cache = _reshape_kv_cache( + attn_groups, + raw_tensors, + "int8_per_token_head", + [spec.block_size], + {}, + )["layer"] + + assert kv_cache.shape == (num_blocks, 2, 16, 1, 8) + assert kv_cache.stride(0) == spec.page_size_bytes + assert kv_cache.stride(1) == 16 * 1 * 8 + assert kv_cache[1, 1].storage_offset() == spec.page_size_bytes + 16 * 1 * 8 diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index 03a203a1bcf..ebf607b65a7 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -201,6 +201,38 @@ class AttentionBackend(ABC): return min(s.base if isinstance(s, MultipleOf) else s for s in supported_sizes) + @classmethod + def indexes_kv_by_block_stride(cls) -> bool: + """Whether the backend reads KV pages by the runtime block stride. + + True when ``num_blocks`` is the outermost physical dimension of the KV + cache, so the backend tolerates a non-contiguous block dim. This gates + page size padding and cross-layer uniform KV layout. + + Returns: + True if the backend's physical KV layout is num-blocks-first. False + otherwise, including when the backend does not define a layered + stride order. + """ + try: + kv_cache_stride_order = cls.get_kv_cache_stride_order( + include_num_layers_dimension=False + ) + layered_kv_cache_stride_order = cls.get_kv_cache_stride_order( + include_num_layers_dimension=True + ) + except (AttributeError, NotImplementedError): + return False + + # Check that attention backend includes a layers dimension. + if len(layered_kv_cache_stride_order) != len(kv_cache_stride_order) + 1: + return False + + # stride_order[0] == 0 means num_layers stays first in physical + # layout (identity permutation), so indexing by block stride is + # not supported. + return layered_kv_cache_stride_order[0] != 0 + @classmethod def is_mla(cls) -> bool: return False diff --git a/vllm/v1/core/kv_cache_utils.py b/vllm/v1/core/kv_cache_utils.py index 4e1d28d7d5d..95b8fba4ccf 100644 --- a/vllm/v1/core/kv_cache_utils.py +++ b/vllm/v1/core/kv_cache_utils.py @@ -20,6 +20,7 @@ from vllm.utils.math_utils import cdiv, round_up from vllm.utils.mem_utils import format_gib from vllm.utils.torch_utils import get_dtype_size from vllm.v1.kv_cache_interface import ( + AttentionSpec, ChunkedLocalAttentionSpec, FullAttentionSpec, HiddenStateCacheSpec, @@ -1029,9 +1030,14 @@ def unify_kv_cache_spec_page_size( ) -> dict[str, KVCacheSpec]: """ Unify the page size of the given KVCacheSpec. If the page size of all layers - are the same, return the original KVCacheSpec. If not same, unify the page - size by increasing the block size of layers with smaller page size. Raise - NotImplementedError if failed to unify the page size. + are the same, return the original KVCacheSpec. If not same, first try to + unify page size by increasing the block size of layers with smaller page + size. If a smaller attention page does not evenly divide the maximum page + size, keep its logical block size and pad its physical page instead --- but + only for attention layers whose backend opts in via + ``AttentionSpec.indexes_kv_by_block_stride`` (the padded page is read through + a strided view, which not every backend handles). Raise NotImplementedError + if failed to unify the page size. Args: kv_cache_spec: The KVCacheSpec of each attention layer in the model @@ -1051,14 +1057,23 @@ def unify_kv_cache_spec_page_size( new_kv_cache_spec[layer_name] = layer_spec else: layer_page_size = layer_spec.page_size_bytes - if max_page_size % layer_page_size != 0: + if max_page_size % layer_page_size == 0: + ratio = max_page_size // layer_page_size + new_block_size = layer_spec.block_size * ratio + new_spec = replace(layer_spec, block_size=new_block_size) + elif ( + isinstance(layer_spec, AttentionSpec) + and layer_spec.indexes_kv_by_block_stride + ): + new_spec = replace(layer_spec, page_size_padded=max_page_size) + else: raise NotImplementedError( - "The page size of the layer is not divisible by the " - "maximum page size. Cannot unify by adjusting block_size." + f"Layer {layer_name}: page size is not divisible by the " + "maximum page size and cannot be padded. Padding is only " + "supported for attention layers whose backend indexes KV " + "pages by the block stride (indexes_kv_by_block_stride is " + "True)." ) - ratio = max_page_size // layer_page_size - new_block_size = layer_spec.block_size * ratio - new_spec = replace(layer_spec, block_size=new_block_size) assert new_spec.page_size_bytes == max_page_size new_kv_cache_spec[layer_name] = new_spec return new_kv_cache_spec diff --git a/vllm/v1/kv_cache_interface.py b/vllm/v1/kv_cache_interface.py index 2e779b2c2a4..5a2a5c5e298 100644 --- a/vllm/v1/kv_cache_interface.py +++ b/vllm/v1/kv_cache_interface.py @@ -163,6 +163,7 @@ class AttentionSpec(KVCacheSpec): dtype: torch.dtype kv_quant_mode: KVQuantMode = KVQuantMode.NONE page_size_padded: int | None = None + indexes_kv_by_block_stride: bool = False @property def page_size_bytes(self) -> int: @@ -283,6 +284,7 @@ class FullAttentionSpec(AttentionSpec): dtype=specs[0].dtype, kv_quant_mode=specs[0].kv_quant_mode, page_size_padded=specs[0].page_size_padded, + indexes_kv_by_block_stride=specs[0].indexes_kv_by_block_stride, sliding_window=cls.merge_window_sizes(sliding_window), attention_chunk_size=cls.merge_window_sizes(attention_chunk_size), # If any layer in the group is non-causal, treat the group as @@ -403,13 +405,16 @@ class MLAAttentionSpec(FullAttentionSpec): cache_dtype_str_set = set(spec.cache_dtype_str for spec in specs) compress_ratio_set = set(spec.compress_ratio for spec in specs) model_version_set = set(spec.model_version for spec in specs) + block_stride_set = set(spec.indexes_kv_by_block_stride for spec in specs) assert ( len(cache_dtype_str_set) == 1 and len(compress_ratio_set) == 1 and len(model_version_set) == 1 + and len(block_stride_set) == 1 ), ( "All attention layers in the same KV cache group must use the same " - "quantization method, compress ratio, and model version." + "quantization method, compress ratio, model version, and KV block " + "stride indexing." ) return cls( block_size=specs[0].block_size, @@ -418,6 +423,7 @@ class MLAAttentionSpec(FullAttentionSpec): dtype=specs[0].dtype, kv_quant_mode=specs[0].kv_quant_mode, page_size_padded=specs[0].page_size_padded, + indexes_kv_by_block_stride=block_stride_set.pop(), cache_dtype_str=cache_dtype_str_set.pop(), compress_ratio=compress_ratio_set.pop(), model_version=model_version_set.pop(), @@ -584,15 +590,17 @@ class SlidingWindowMLASpec(SlidingWindowSpec): compress_ratio_set = set(spec.compress_ratio for spec in specs) model_version_set = set(spec.model_version for spec in specs) sliding_window_set = set(spec.sliding_window for spec in specs) + block_stride_set = set(spec.indexes_kv_by_block_stride for spec in specs) assert ( len(cache_dtype_str_set) == 1 and len(compress_ratio_set) == 1 and len(model_version_set) == 1 and len(sliding_window_set) == 1 + and len(block_stride_set) == 1 ), ( "All attention layers in the same KV cache group must use the same " - "quantization method, compress ratio, model version and sliding " - "window size." + "quantization method, compress ratio, model version, sliding " + "window size, and KV block stride indexing." ) return cls( block_size=specs[0].block_size, @@ -600,6 +608,7 @@ class SlidingWindowMLASpec(SlidingWindowSpec): head_size=specs[0].head_size, dtype=specs[0].dtype, page_size_padded=specs[0].page_size_padded, + indexes_kv_by_block_stride=block_stride_set.pop(), sliding_window=sliding_window_set.pop(), cache_dtype_str=cache_dtype_str_set.pop(), compress_ratio=compress_ratio_set.pop(), @@ -711,6 +720,7 @@ class SinkFullAttentionSpec(FullAttentionSpec): dtype=specs[0].dtype, kv_quant_mode=specs[0].kv_quant_mode, page_size_padded=specs[0].page_size_padded, + indexes_kv_by_block_stride=specs[0].indexes_kv_by_block_stride, sliding_window=cls.merge_window_sizes(sliding_window), attention_chunk_size=cls.merge_window_sizes(attention_chunk_size), non_causal=any(spec.non_causal for spec in specs), diff --git a/vllm/v1/worker/gpu/attn_utils.py b/vllm/v1/worker/gpu/attn_utils.py index 7b85e6fa316..737feb7d277 100644 --- a/vllm/v1/worker/gpu/attn_utils.py +++ b/vllm/v1/worker/gpu/attn_utils.py @@ -1,13 +1,17 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from collections.abc import Iterable, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, replace from math import prod from typing import Any, cast import torch -from vllm.config import VllmConfig, get_layers_from_vllm_config +from vllm.config import ( + VllmConfig, + get_layers_from_vllm_config, + set_current_vllm_config, +) from vllm.model_executor.layers.attention import Attention from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase from vllm.utils.torch_utils import get_dtype_size @@ -47,6 +51,13 @@ def get_kv_cache_spec(vllm_config: VllmConfig) -> dict[str, KVCacheSpec]: continue # Skip modules that don't need KV cache (eg encoder-only attention) if spec := attn_module.get_kv_cache_spec(vllm_config): + if isinstance(spec, AttentionSpec): + backend = attn_module.get_attn_backend() + # indexes_kv_by_block_stride() -> get_kv_cache_stride_order() -> + # get_kv_cache_layout() needs the current vLLM config. + with set_current_vllm_config(vllm_config): + indexes = backend.indexes_kv_by_block_stride() + spec = replace(spec, indexes_kv_by_block_stride=indexes) kv_cache_spec[layer_name] = spec return kv_cache_spec @@ -180,6 +191,62 @@ def _allocate_kv_cache( return kv_cache_raw_tensors +def _reshape_attention_kv_cache( + kv_raw_tensor: torch.Tensor, + kv_cache_spec: AttentionSpec, + kv_cache_shape: tuple[int, ...], + kv_cache_stride_order: tuple[int, ...], + num_blocks: int, + packing: tuple[int, int] | None, +) -> torch.Tensor: + permuted_kv_cache_shape = tuple(kv_cache_shape[i] for i in kv_cache_stride_order) + inv_order = [ + kv_cache_stride_order.index(i) for i in range(len(kv_cache_stride_order)) + ] + dtype = kv_cache_spec.dtype + + if packing is not None: + offset, block_stride = packing + assert inv_order[0] == 0 + page_bytes = prod(kv_cache_shape[1:]) * get_dtype_size(dtype) + kv_cache = ( + kv_raw_tensor.view(-1, block_stride)[:, offset : offset + page_bytes] + .view(dtype) + .view(kv_cache_shape) + ) + elif kv_cache_spec.page_size_padded is not None: + # Use a strided view to skip the padding between physical pages. + # + # Only num-blocks-first layouts are supported (the block dimension is + # dim 0 of the unpermuted shape). kv-first layouts such as ROCm's + # ``(2, num_blocks, ...)`` are intentionally not supported here. For a + # num-blocks-first layout the only stride that must change is the block + # stride: every other (contiguous) stride already steps within the + # unpadded region of a page, so no further adjustment is needed. + assert kv_cache_shape[0] == num_blocks, ( + "Padded KV pages require a num-blocks-first KV cache layout (got " + f"shape {kv_cache_shape} with num_blocks={num_blocks}); " + "kv-first layouts are not supported." + ) + dtype_size = get_dtype_size(kv_cache_spec.dtype) + page_stride = kv_cache_spec.page_size_bytes // dtype_size + + num_blocks_dim = inv_order[0] + strides = list(torch.empty(permuted_kv_cache_shape).stride()) + strides[num_blocks_dim] = page_stride + + kv_cache = torch.as_strided( + kv_raw_tensor.view(dtype), + size=permuted_kv_cache_shape, + stride=tuple(strides), + ) + else: + # No padding — safe to use a contiguous view. + kv_cache = kv_raw_tensor.view(dtype).view(permuted_kv_cache_shape) + + return kv_cache.permute(*inv_order) + + def _reshape_kv_cache( attn_groups: Sequence[AttentionGroup], kv_cache_raw_tensors: dict[str, torch.Tensor], @@ -248,45 +315,14 @@ def _reshape_kv_cache( except (AttributeError, NotImplementedError): kv_cache_stride_order = tuple(range(len(kv_cache_shape))) - kv_cache_shape = tuple(kv_cache_shape[i] for i in kv_cache_stride_order) - inv_order = [ - kv_cache_stride_order.index(i) - for i in range(len(kv_cache_stride_order)) - ] - - dtype = kv_cache_spec.dtype - if packing is not None: - offset, block_stride = packing - assert inv_order[0] == 0 - page_bytes = prod(kv_cache_shape[1:]) * get_dtype_size(dtype) - kv_cache = ( - kv_raw_tensor.view(-1, block_stride)[ - :, offset : offset + page_bytes - ] - .view(dtype) - .view(kv_cache_shape) - ) - elif kv_cache_spec.page_size_padded is not None: - # Use strided view to handle page_size_bytes that - # include padding. This follows the same pattern as - # MambaSpec handling in gpu_model_runner.py. - # NOTE: This assumes kv_cache_shape[0] == num_blocks - # (i.e. the first physical dimension is the block - # index), which holds for all current backends - # (MLA, FlashAttention, TritonAttention, etc.). - dtype_size = get_dtype_size(dtype) - page_stride = kv_cache_spec.page_size_bytes // dtype_size - strides = list(torch.empty(kv_cache_shape).stride()) - strides[inv_order[0]] = page_stride - kv_cache = torch.as_strided( - kv_raw_tensor.view(dtype), - size=kv_cache_shape, - stride=tuple(strides), - ) - else: - # No padding — safe to use a contiguous view. - kv_cache = kv_raw_tensor.view(dtype).view(kv_cache_shape) - kv_caches[layer_name] = kv_cache.permute(*inv_order) + kv_caches[layer_name] = _reshape_attention_kv_cache( + kv_raw_tensor, + kv_cache_spec, + kv_cache_shape, + kv_cache_stride_order, + kernel_num_blocks, + packing, + ) elif isinstance(kv_cache_spec, MambaSpec): has_mamba = True diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index b554542e65d..0b72870fc4d 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -12,7 +12,6 @@ from contextlib import contextmanager from copy import copy, deepcopy from dataclasses import dataclass, replace from functools import reduce -from math import prod from typing import TYPE_CHECKING, Any, NamedTuple, TypeAlias, cast import numpy as np @@ -203,6 +202,7 @@ from vllm.v1.worker.cp_utils import ( ) from vllm.v1.worker.dp_utils import coordinate_batch_across_dp from vllm.v1.worker.ec_connector_model_runner_mixin import ECConnectorModelRunnerMixin +from vllm.v1.worker.gpu.attn_utils import _reshape_attention_kv_cache from vllm.v1.worker.gpu.pool.late_interaction_runner import LateInteractionRunner from vllm.v1.worker.gpu_input_batch import CachedRequestState, InputBatch from vllm.v1.worker.gpu_ubatch_wrapper import UBatchWrapper @@ -7125,62 +7125,20 @@ class GPUModelRunner( kv_cache_spec.head_size, cache_dtype_str=self.cache_config.cache_dtype, ) - dtype = kv_cache_spec.dtype try: kv_cache_stride_order = attn_backend.get_kv_cache_stride_order() assert len(kv_cache_stride_order) == len(kv_cache_shape) except (AttributeError, NotImplementedError): kv_cache_stride_order = tuple(range(len(kv_cache_shape))) - # The allocation respects the backend-defined stride order - # to ensure the semantic remains consistent for each - # backend. We first obtain the generic kv cache shape and - # then permute it according to the stride order which could - # result in a non-contiguous tensor. - kv_cache_shape = tuple( - kv_cache_shape[i] for i in kv_cache_stride_order + raw_tensor = kv_cache_raw_tensors[layer_name] + kv_caches[layer_name] = _reshape_attention_kv_cache( + raw_tensor, + kv_cache_spec, + kv_cache_shape, + kv_cache_stride_order, + kernel_num_blocks, + packing, ) - # Maintain original KV shape view. - inv_order = [ - kv_cache_stride_order.index(i) - for i in range(len(kv_cache_stride_order)) - ] - - if packing is not None: - offset, block_stride = packing - assert inv_order[0] == 0 - page_bytes = prod(kv_cache_shape[1:]) * get_dtype_size(dtype) - kv_cache = ( - kv_cache_raw_tensors[layer_name] - .view(-1, block_stride)[:, offset : offset + page_bytes] - .view(dtype) - .view(kv_cache_shape) - ) - elif kv_cache_spec.page_size_padded is not None: - # Use strided view to handle page_size_bytes that - # include padding. This follows - # the same pattern as MambaSpec handling below. - # NOTE: This assumes kv_cache_shape[0] == num_blocks - # (i.e. the first physical dimension is the block - # index), which holds for MLA backends but NOT for - # standard attention backends whose shape starts with - # a K/V dimension of size 2. - dtype_size = get_dtype_size(dtype) - page_stride = kv_cache_spec.page_size_bytes // dtype_size - strides = list(torch.empty(kv_cache_shape).stride()) - strides[inv_order[0]] = page_stride - kv_cache = torch.as_strided( - kv_cache_raw_tensors[layer_name].view(dtype), - size=kv_cache_shape, - stride=tuple(strides), - ) - else: - # No padding — safe to use a contiguous view. - kv_cache = ( - kv_cache_raw_tensors[layer_name] - .view(dtype) - .view(kv_cache_shape) - ) - kv_caches[layer_name] = kv_cache.permute(*inv_order) elif isinstance(kv_cache_spec, MambaSpec): has_mamba = True @@ -7265,7 +7223,7 @@ class GPUModelRunner( # Try creating KV caches optimized for kv-connector transfers cache_dtype = self.cache_config.cache_dtype - if self.use_uniform_kv_cache(self.attn_groups, cache_dtype): + if self.use_uniform_kv_cache(self.attn_groups): kv_caches, cross_layers_kv_cache, attn_backend = ( self.allocate_uniform_kv_caches( kv_cache_config, @@ -7515,6 +7473,13 @@ class GPUModelRunner( continue # Skip modules that don't need KV cache (eg encoder-only attention) if spec := attn_module.get_kv_cache_spec(self.vllm_config): + if isinstance(spec, AttentionSpec): + backend = attn_module.get_attn_backend() + # indexes_kv_by_block_stride() -> get_kv_cache_stride_order() + # -> get_kv_cache_layout() needs the current vLLM config. + with set_current_vllm_config(self.vllm_config): + indexes = backend.indexes_kv_by_block_stride() + spec = replace(spec, indexes_kv_by_block_stride=indexes) kv_cache_spec[layer_name] = spec return kv_cache_spec diff --git a/vllm/v1/worker/kv_connector_model_runner_mixin.py b/vllm/v1/worker/kv_connector_model_runner_mixin.py index 797e59c0290..c2c54e647df 100644 --- a/vllm/v1/worker/kv_connector_model_runner_mixin.py +++ b/vllm/v1/worker/kv_connector_model_runner_mixin.py @@ -114,7 +114,6 @@ class KVConnectorModelRunnerMixin: @staticmethod def use_uniform_kv_cache( attn_groups: list[list[AttentionGroup]], - cache_dtype: CacheDType, ) -> bool: """ Determines whether a uniform KV layout should be used. @@ -128,9 +127,9 @@ class KVConnectorModelRunnerMixin: have the same page size. 2. A KV connector is configured, and the KV connector instance prefers to use this layout (prefer_cross_layer_blocks() returns True) - 2. The flash attention backend supports this layout - (get_kv_cache_stride_order(True) includes a placement for a - num_layers dimension) + 3. The attention backend indexes KV by the block stride + (kv_cache_spec.indexes_kv_by_block_stride), i.e. num_blocks is the + outermost physical dim so per-block all-layers data is contiguous. Note that the actual placement of the num_layers dimensions in the unified layers tensors will be determined by the attention @@ -140,7 +139,6 @@ class KVConnectorModelRunnerMixin: Args: attn_groups: The list of attention groups for this model - cache_dtype: The KV cache dtype Returns: True if we should use a uniform KV cache layout. """ @@ -157,30 +155,7 @@ class KVConnectorModelRunnerMixin: kv_cache_spec = attn_group.kv_cache_spec if not isinstance(kv_cache_spec, AttentionSpec): return False - - attn_backend = attn_group.backend - kv_cache_shape = attn_backend.get_kv_cache_shape( - 1234, - kv_cache_spec.block_size, - kv_cache_spec.num_kv_heads, - kv_cache_spec.head_size, - cache_dtype_str=cache_dtype, - ) - - try: - kv_cache_stride_order = attn_backend.get_kv_cache_stride_order( - include_num_layers_dimension=True - ) - except (AttributeError, NotImplementedError): - return False - - # check that attention backend includes a layers dimension - if len(kv_cache_stride_order) != len(kv_cache_shape) + 1: - return False - - # stride_order[0] == 0 means num_layers stays first in physical - # layout (identity permutation), so cross-layer is unsupported. - return kv_cache_stride_order[0] != 0 + return kv_cache_spec.indexes_kv_by_block_stride @staticmethod def allocate_uniform_kv_caches(