diff --git a/tests/v1/kv_connector/unit/offloading_connector/test_worker.py b/tests/v1/kv_connector/unit/offloading_connector/test_worker.py index 25bb664ff40..b4557bd7b62 100644 --- a/tests/v1/kv_connector/unit/offloading_connector/test_worker.py +++ b/tests/v1/kv_connector/unit/offloading_connector/test_worker.py @@ -125,8 +125,7 @@ def test_register_kv_caches(backend): own dedicated tensors. Uses the real GPUModelRunner.initialize_kv_cache_tensors to produce - kv_caches, which automatically applies - _update_hybrid_attention_mamba_layout for hybrid models. + the raw per-layer kv_caches registered by the connector. Verifies that the canonicalized CanonicalKVCaches has the correct block tensors, tensor_idx references, and page sizes across all groups. diff --git a/vllm/distributed/kv_transfer/kv_connector/utils.py b/vllm/distributed/kv_transfer/kv_connector/utils.py index 77f043ef527..ac8744d8c46 100644 --- a/vllm/distributed/kv_transfer/kv_connector/utils.py +++ b/vllm/distributed/kv_transfer/kv_connector/utils.py @@ -21,12 +21,10 @@ from vllm.logger import init_logger from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase from vllm.platforms import current_platform from vllm.v1.attention.backend import AttentionBackend -from vllm.v1.kv_cache_interface import MambaSpec from vllm.v1.outputs import KVConnectorOutput, ModelRunnerOutput if TYPE_CHECKING: from vllm.distributed.kv_transfer.kv_connector.base import KVConnectorBase - from vllm.v1.kv_cache_interface import KVCacheSpec logger = init_logger(__name__) @@ -594,32 +592,6 @@ class TransferTopology: abs_ratio = -tp_ratio return [self.tp_rank * abs_ratio + i for i in range(abs_ratio)] - def get_transfer_cache_regions( - self, cache: torch.Tensor, layer_spec: "KVCacheSpec" - ) -> list[torch.Tensor] | torch.Tensor: - """Return the cache tensor(s) to register as NIXL memory regions, - also accounting for hybrid SSM models specificities. - """ - if isinstance(layer_spec, MambaSpec): - # Register the whole kv cache shared tensor, including - # SSM/Conv. - conv, ssm = cache - return [conv] - - # Check may be hacky but it's matching - # `_update_hybrid_attention_mamba_layout`. - if self.is_mamba and cache.shape[0] == 2: - # When MAMBA is present, all backends are blocks first, so - # that blocks can be shared between attention layers and mamba - # layers. Runner already adjusted strides for FlashAttn-like - # backends so its num_blocks first. - # Swap [2<>num_blocks] dims for hybrid SSM layout. - cache = cache.transpose(0, 1) - - # K and V are packed into one tensor (content dim), so each layer - # registers as a single region. - return [cache] - def describe(self, remote_engine_id: EngineId, remote_pp_rank: int = 0) -> str: """One-line summary of transfer config for logging.""" info = self._engines[(remote_engine_id, remote_pp_rank)] diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py index 34f99b7cfac..dd94ffb2f9e 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py @@ -1678,9 +1678,9 @@ class MooncakeConnectorWorker: conv, _ = cache_or_caches cache_list = [conv] else: - cache_list = self.transfer_topo.get_transfer_cache_regions( - cache_or_caches, layer_spec - ) + # K and V are packed into one blocks-first tensor per layer, + # so each layer registers as a single region. + cache_list = [cache_or_caches] logger.debug( "registering layer %s with %d cache tensor(s)", diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py index 4cdc54f7e76..e9315f6776c 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py @@ -1092,7 +1092,7 @@ class NixlBaseConnectorWorker: # to better exploit the memory layout (ie num_blocks is the first dim). tensor_size_bytes = None - for layer_name, cache_or_caches in xfer_buffers.items(): + for layer_name, cache in xfer_buffers.items(): # NOTE (NickLucche) Hybrid SSM models assume a layout that is similar to # that of FI, with block laid out as in `get_backend_aware_kv_block_len`. # However, physical page_size may differ when kernel requires a specific @@ -1109,9 +1109,6 @@ class NixlBaseConnectorWorker: if isinstance(layer_spec, UniformTypeKVCacheSpecs): # MLA DSv32 Indexer case: UniformTypeKVCacheSpecs merges kv_cache_specs layer_spec = layer_spec.kv_cache_specs[layer_name] - cache_list = self.transfer_topo.get_transfer_cache_regions( - cache_or_caches, layer_spec - ) # `layer_spec.page_size_bytes` only accounts for logical page_size, that is # the page_size assuming constant `self._logical_num_blocks`. physical_page_size = ( @@ -1120,8 +1117,6 @@ class NixlBaseConnectorWorker: else layer_spec.page_size_bytes // self._physical_blocks_per_logical_kv_block ) - # For when registering multiple tensors eg K/V in separate regions. - physical_page_size = physical_page_size // len(cache_list) if self.transfer_topo._cross_layers_blocks: # When cross-layers blocks are used, multiply by number of layers physical_page_size = physical_page_size * len( @@ -1136,66 +1131,61 @@ class NixlBaseConnectorWorker: # [`num_blocks` * `page_size`] curr_tensor_size_bytes = num_blocks * physical_page_size - # TODO (NickLucche) we could eventually unify how we handle FA/FI regions, - # registering a single tensor for both K/V and splitting logically like FI. - for cache in cache_list: - base_addr = cache.data_ptr() - if base_addr in seen_base_addresses: - # NOTE (NickLucche) HMA employs memory pooling to share tensors - # across groups. This results in skipping all tensors but the ones - # pointed to by group0. Also, generally we will have more blocks - # per tensor but fewer regions. - logger.debug("Skipping %s because it's already seen", layer_name) - continue - logger.debug( - "Registering layer %s with cache shape: %s", layer_name, cache.shape + base_addr = cache.data_ptr() + if base_addr in seen_base_addresses: + # NOTE (NickLucche) HMA employs memory pooling to share tensors + # across groups. This results in skipping all tensors but the ones + # pointed to by group0. Also, generally we will have more blocks + # per tensor but fewer regions. + logger.debug("Skipping %s because it's already seen", layer_name) + continue + logger.debug( + "Registering layer %s with cache shape: %s", layer_name, cache.shape + ) + seen_base_addresses.append(base_addr) + # Only record non-Mamba page sizes. + if isinstance(layer_spec, MambaSpec): + self.block_len_per_layer.append( + physical_page_size // self._physical_blocks_per_logical_kv_block ) - seen_base_addresses.append(base_addr) - # Only record non-Mamba page sizes. - if isinstance(layer_spec, MambaSpec): - self.block_len_per_layer.append( - physical_page_size // self._physical_blocks_per_logical_kv_block - ) - else: - self.block_len_per_layer.append(physical_page_size) - is_mla_region = isinstance( - layer_spec, (MLAAttentionSpec, SlidingWindowMLASpec) + else: + self.block_len_per_layer.append(physical_page_size) + is_mla_region = isinstance( + layer_spec, (MLAAttentionSpec, SlidingWindowMLASpec) + ) + self._region_is_mla.append(is_mla_region) + + if not is_mla_region: + if tensor_size_bytes is None: + tensor_size_bytes = curr_tensor_size_bytes + assert tensor_size_bytes == curr_tensor_size_bytes, ( + "All non-MLA kv cache tensors must have the same size" ) - self._region_is_mla.append(is_mla_region) - if not is_mla_region: - if tensor_size_bytes is None: - tensor_size_bytes = curr_tensor_size_bytes - assert tensor_size_bytes == curr_tensor_size_bytes, ( - "All non-MLA kv cache tensors must have the same size" - ) - - # When there's a mismatch between kbs<>bs, we rely on HMA to ensure - # caches are either [NB, PS] or [NB*r, PS/r] where r is bs/kbs. - if ( - self._physical_blocks_per_logical_kv_block == 1 - and cache.shape[0] != num_blocks - ): - raise AssertionError( - "All kv cache tensors must have the same number of " - f"blocks; layer={layer_name}, " - f"expected_num_blocks={num_blocks}, " - f"cache_shape={tuple(cache.shape)}, " - f"cache_stride={tuple(cache.stride())}, " - f"layer_spec={type(layer_spec).__name__}, " - f"backend={self.backend_name}, " - "all_backends=" - f"{[backend.get_name() for backend in self.attn_backends]}, " - f"kv_cache_layout={self.kv_cache_layout}" - ) - - # Need to make sure the device ID is non-negative for NIXL, - # Torch uses -1 to indicate CPU tensors. - self.device_id = max(cache.get_device(), 0) - caches_data.append( - (base_addr, curr_tensor_size_bytes, self.device_id, "") + # When there's a mismatch between kbs<>bs, we rely on HMA to ensure + # caches are either [NB, PS] or [NB*r, PS/r] where r is bs/kbs. + if ( + self._physical_blocks_per_logical_kv_block == 1 + and cache.shape[0] != num_blocks + ): + raise AssertionError( + "All kv cache tensors must have the same number of " + f"blocks; layer={layer_name}, " + f"expected_num_blocks={num_blocks}, " + f"cache_shape={tuple(cache.shape)}, " + f"cache_stride={tuple(cache.stride())}, " + f"layer_spec={type(layer_spec).__name__}, " + f"backend={self.backend_name}, " + "all_backends=" + f"{[backend.get_name() for backend in self.attn_backends]}, " + f"kv_cache_layout={self.kv_cache_layout}" ) + # Need to make sure the device ID is non-negative for NIXL, + # Torch uses -1 to indicate CPU tensors. + self.device_id = max(cache.get_device(), 0) + caches_data.append((base_addr, curr_tensor_size_bytes, self.device_id, "")) + logger.debug( "Different block lengths collected: %s", set(self.block_len_per_layer) ) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/worker.py index 045e513c3f4..a61ec917543 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/worker.py @@ -56,9 +56,7 @@ class OffloadingConnectorWorker: def _init_worker(self, kv_caches: CanonicalKVCaches) -> None: self.worker = self.spec.get_worker(kv_caches) - def register_kv_caches( - self, kv_caches: dict[str, torch.Tensor | list[torch.Tensor]] - ): + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): kv_cache_config = self.kv_cache_config num_blocks = kv_cache_config.num_blocks @@ -120,24 +118,13 @@ class OffloadingConnectorWorker: ) elif isinstance(layer_kv_cache_spec, MambaSpec): - state_tensors = kv_caches[layer_name] - assert isinstance(state_tensors, list) - - # re-construct the raw (num_blocks, page_size) tensor - # from the first state tensor - assert len(state_tensors) > 0 - first_state_tensor = state_tensors[0] - assert first_state_tensor.storage_offset() == 0 - tensor = ( - torch.tensor( - [], - dtype=torch.int8, - device=first_state_tensor.device, - ) - .set_(first_state_tensor.untyped_storage()) - .view((num_blocks, layer_kv_cache_spec.page_size_bytes)) + layer_kv_cache = kv_caches[layer_name] + assert layer_kv_cache.dtype == torch.int8 + tensors_per_block[layer_name] = ( + layer_kv_cache.view( + num_blocks, layer_kv_cache_spec.page_size_bytes + ), ) - tensors_per_block[layer_name] = (tensor,) page_size_bytes[layer_name] = layer_kv_cache_spec.page_size_bytes unpadded_page_size_bytes[layer_name] = replace( diff --git a/vllm/model_executor/layers/attention_layer_base.py b/vllm/model_executor/layers/attention_layer_base.py index 53994114558..197213dc79f 100644 --- a/vllm/model_executor/layers/attention_layer_base.py +++ b/vllm/model_executor/layers/attention_layer_base.py @@ -4,6 +4,8 @@ from abc import ABC, abstractmethod +import torch + from vllm.config import VllmConfig from vllm.v1.attention.backend import AttentionBackend, AttentionImpl from vllm.v1.kv_cache_interface import KVCacheSpec @@ -21,6 +23,14 @@ class AttentionLayerBase(ABC): impl: "AttentionImpl" supports_dcp: bool = True + def bind_kv_cache(self, kv_cache: torch.Tensor) -> None: + """Bind the allocated KV cache tensor to this layer. + + The default stores the cache view as-is; subclasses (e.g. Mamba) + override this to unpack the raw buffer into per-state views. + """ + self.kv_cache = kv_cache + @abstractmethod def get_attn_backend(self) -> type[AttentionBackend]: """Get the attention backend class for this layer.""" diff --git a/vllm/model_executor/layers/mamba/abstract.py b/vllm/model_executor/layers/mamba/abstract.py index 62fd64bd0de..d06916f697c 100644 --- a/vllm/model_executor/layers/mamba/abstract.py +++ b/vllm/model_executor/layers/mamba/abstract.py @@ -2,11 +2,13 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from abc import abstractmethod from collections.abc import Iterable +from math import prod import torch from vllm.config import VllmConfig from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase +from vllm.utils.torch_utils import get_dtype_size from vllm.v1.attention.backend import AttentionBackend from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum from vllm.v1.attention.selector import get_mamba_attn_backend @@ -24,6 +26,22 @@ class MambaBase(AttentionLayerBase): kv_cache: tuple[torch.Tensor, ...] supports_dcp: bool = False + def bind_kv_cache(self, kv_cache: torch.Tensor) -> None: + """Unpack a raw ``[B, 1, 1, C]`` int8 page view into per-state views. + + Each block's ``C`` bytes hold the layer's states (e.g. conv, ssm) + packed contiguously; slice them out and reinterpret per dtype/shape. + """ + pages = kv_cache.squeeze(dim=(1, 2)) + states: list[torch.Tensor] = [] + offset = 0 + for shape, dtype in zip(self.get_state_shape(), self.get_state_dtype()): + nbytes = prod(shape) * get_dtype_size(dtype) + state = pages[:, offset : offset + nbytes].view(dtype) + states.append(state.view(-1, *shape)) + offset += nbytes + self.kv_cache = tuple(states) + @abstractmethod def get_state_shape(self) -> Iterable[tuple[int, ...]]: """ diff --git a/vllm/v1/worker/gpu/attn_utils.py b/vllm/v1/worker/gpu/attn_utils.py index a97a4e39597..b6bbe28e5f2 100644 --- a/vllm/v1/worker/gpu/attn_utils.py +++ b/vllm/v1/worker/gpu/attn_utils.py @@ -262,7 +262,7 @@ def _reshape_kv_cache( kv_cache_config: "KVCacheConfig | None" = None, ) -> dict[str, Any]: kv_caches: dict[str, Any] = {} - has_attn, has_mamba = False, False + has_attn = False layer_packing: dict[str, tuple[int, int]] = {} if kv_cache_config is not None: @@ -340,38 +340,21 @@ def _reshape_kv_cache( ) elif isinstance(kv_cache_spec, MambaSpec): - has_mamba = True - state_tensors = [] - storage_offset_bytes = 0 - for shape, dtype in zip(kv_cache_spec.shapes, kv_cache_spec.dtypes): - dtype_size = get_dtype_size(dtype) - num_element_per_page = kv_cache_spec.page_size_bytes // dtype_size - target_shape = (num_blocks, *shape) - stride = torch.empty(target_shape).stride() - target_stride = (num_element_per_page, *stride[1:]) - assert storage_offset_bytes % dtype_size == 0 - tensor = torch.as_strided( - kv_raw_tensor.view(dtype), - size=target_shape, - stride=target_stride, - storage_offset=storage_offset_bytes // dtype_size, - ) - state_tensors.append(tensor) - storage_offset_bytes += stride[0] * dtype_size - kv_caches[layer_name] = state_tensors + page_size_bytes = kv_cache_spec.page_size_bytes + # Hold a single contiguous [num_blocks, 1, 1, page_size_bytes] + # int8 page view per layer; the layer's bind_kv_cache unpacks + # each block's bytes into its conv/ssm state views. Keeping + # one tensor per layer lets the KV connector register it + # without special-casing Mamba. + kv_caches[layer_name] = kv_raw_tensor[ + : num_blocks * page_size_bytes + ].view(num_blocks, 1, 1, page_size_bytes) else: raise NotImplementedError( f"Unsupported KV cache spec type: {type(kv_cache_spec)}" ) - if has_attn and has_mamba: - _update_hybrid_attention_layout( - attn_groups=attn_groups, - kv_caches=kv_caches, - kernel_block_sizes=kernel_block_sizes, - cache_dtype=cache_dtype, - ) - elif has_attn and kv_cache_config is not None: + if has_attn and kv_cache_config is not None: _align_mixed_attention_kv_cache_views( attn_groups=attn_groups, kv_caches=kv_caches, @@ -458,65 +441,6 @@ def _restride_blocks_first_kv_cache_to_kv_first_storage( ) -def _update_hybrid_attention_layout( - attn_groups: Iterable[AttentionGroup], - kv_caches: dict[str, Any], - kernel_block_sizes: list[int], - cache_dtype: str, -) -> None: - for group in attn_groups: - if group.kv_cache_group_id >= len(kernel_block_sizes): - continue - - kv_cache_spec = group.kv_cache_spec - if not isinstance(kv_cache_spec, AttentionSpec): - continue - # Mirror the per-layer dtype selection used when building the shape - # above. The block-dim index is dtype-independent for current backends - # (quantization only changes the last dim), so this is a no-op today, - # but it keeps both call sites consistent for skip layers. - layer_cache_dtype = ( - "auto" - if kv_cache_spec.kv_quant_mode == KVQuantMode.NONE - and not isinstance(kv_cache_spec, TQFullAttentionSpec) - else cache_dtype - ) - block_dim = group.backend.get_kv_cache_block_dim( - kernel_block_sizes[group.kv_cache_group_id], - kv_cache_spec.num_kv_heads, - kv_cache_spec.head_size, - cache_dtype_str=layer_cache_dtype, - ) - # if the first dim of the kvcache's layout is already num_blocks, continue - if block_dim == 0: - continue - - assert block_dim == 1, ( - "Expected the dim `num_blocks` at the second dim when updating" - " the kvcache's layout of full attention layer" - ) - - for layer_name in group.layer_names: - if layer_name not in kv_caches: - # Shared layer — will be aliased to its target after this pass. - continue - - kv_cache = kv_caches[layer_name] - if kv_cache.shape[0] == 2: - assert kv_cache.shape[1] != 2, ( - f"Cannot determine layout for tensor of shape {kv_cache.shape}" - ) - hidden_size = kv_cache.shape[2:].numel() - kv_cache.as_strided_( - size=kv_cache.shape, - stride=( - hidden_size, - 2 * hidden_size, - *kv_cache.stride()[2:], - ), - ) - - def init_kv_cache( runner_kv_caches: list[torch.Tensor | list[torch.Tensor]], forward_context: dict[str, Any], diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 3d68cb08d4b..56a27f3b2a0 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -126,7 +126,6 @@ from vllm.utils.torch_utils import ( PIN_MEMORY, async_tensor_h2d, current_stream, - get_dtype_size, is_quantized_kv_cache, kv_cache_dtype_str_to_dtype, ) @@ -7381,27 +7380,15 @@ class GPUModelRunner( elif isinstance(kv_cache_spec, MambaSpec): has_mamba = True raw_tensor = kv_cache_raw_tensors[layer_name] - state_tensors = [] - storage_offset_bytes = 0 - for shape, dtype in zip(kv_cache_spec.shapes, kv_cache_spec.dtypes): - dtype_size = get_dtype_size(dtype) - num_element_per_page = ( - kv_cache_spec.page_size_bytes // dtype_size - ) - target_shape = (num_blocks, *shape) - stride = torch.empty(target_shape).stride() - target_stride = (num_element_per_page, *stride[1:]) - assert storage_offset_bytes % dtype_size == 0 - tensor = torch.as_strided( - raw_tensor.view(dtype), - size=target_shape, - stride=target_stride, - storage_offset=storage_offset_bytes // dtype_size, - ) - state_tensors.append(tensor) - storage_offset_bytes += stride[0] * dtype_size - - kv_caches[layer_name] = state_tensors + page_size_bytes = kv_cache_spec.page_size_bytes + # Hold a single contiguous [num_blocks, 1, 1, page_size_bytes] + # int8 page view per layer; the layer's bind_kv_cache unpacks + # each block's bytes into its conv/ssm state views. Keeping + # one tensor per layer lets the KV connector register it + # without special-casing Mamba. + kv_caches[layer_name] = raw_tensor[ + : num_blocks * page_size_bytes + ].view(num_blocks, 1, 1, page_size_bytes) else: raise NotImplementedError diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index 9e60d2ae7fa..fda37dcdabe 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -533,9 +533,12 @@ def bind_kv_cache( for layer_name in layer_names: runner_kv_caches.append(kv_caches[layer_name]) - # Bind kv_caches to forward context + # Bind kv_caches to forward context. Each layer's bind_kv_cache unpacks + # its raw allocation into the per-layer view(s) it needs (e.g. Mamba + # splits conv/ssm), so the kv_caches dict can hold a single tensor per + # layer for the KV connector to register. for layer_name, kv_cache in kv_caches.items(): - forward_context[layer_name].kv_cache = kv_cache + forward_context[layer_name].bind_kv_cache(kv_cache) def copy_kv_cache_blocks_inplace(