diff --git a/tests/v1/kv_connector/unit/test_moriio_connector.py b/tests/v1/kv_connector/unit/test_moriio_connector.py index cfac6fa5a36..a8da6cf36d1 100644 --- a/tests/v1/kv_connector/unit/test_moriio_connector.py +++ b/tests/v1/kv_connector/unit/test_moriio_connector.py @@ -36,13 +36,33 @@ from vllm.utils.network_utils import ( get_ip, make_zmq_path, ) -from vllm.v1.kv_cache_interface import KVCacheConfig +from vllm.v1.kv_cache_interface import ( + FullAttentionSpec, + KVCacheConfig, + KVCacheGroupSpec, + KVCacheTensor, +) from .utils import create_request, create_scheduler def _make_test_kv_cache_config() -> KVCacheConfig: - return KVCacheConfig(num_blocks=0, kv_cache_tensors=[], kv_cache_groups=[]) + layer_names = ["layer0", "layer1", "layer2"] + return KVCacheConfig( + num_blocks=2, + kv_cache_tensors=[KVCacheTensor(size=0, shared_by=layer_names)], + kv_cache_groups=[ + KVCacheGroupSpec( + layer_names=layer_names, + kv_cache_spec=FullAttentionSpec( + block_size=16, + num_kv_heads=4, + head_size=64, + dtype=torch.float16, + ), + ) + ], + ) aiter_available = importlib.util.find_spec("aiter") is not None @@ -175,9 +195,18 @@ class FakeMoRIIOConnectorWorker(MoRIIOConnectorWorker): REMOTE_ENGINE_ID = "remote_engine" def __init__( - self, *args, hand_shake_latency: float = 1.8, kv_cache_layout="HND", **kwargs + self, + vllm_config, + engine_id, + *args, + hand_shake_latency: float = 1.8, + kv_cache_layout="HND", + kv_cache_config=None, + **kwargs, ): - super().__init__(*args, **kwargs) + super().__init__( + vllm_config, engine_id, kv_cache_config or _make_test_kv_cache_config() + ) def create_vllm_config( diff --git a/tests/v1/kv_connector/unit/test_moriio_kv_layout.py b/tests/v1/kv_connector/unit/test_moriio_kv_layout.py new file mode 100644 index 00000000000..5b3219db867 --- /dev/null +++ b/tests/v1/kv_connector/unit/test_moriio_kv_layout.py @@ -0,0 +1,228 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import importlib.util +from types import SimpleNamespace + +import pytest +import torch + +from vllm.platforms import current_platform +from vllm.v1.kv_cache_interface import FullAttentionSpec, MLAAttentionSpec + +aiter_available = importlib.util.find_spec("aiter") is not None +mori_available = importlib.util.find_spec("mori") is not None + +if not (current_platform.is_rocm() and mori_available): + pytest.skip( + "MoRIIOs are only available on ROCm with mori package installed", + allow_module_level=True, + ) + +moriio_layout = importlib.import_module( + "vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_layout" +) + + +def _full_spec(block_size: int = 4) -> FullAttentionSpec: + return FullAttentionSpec( + block_size=block_size, + num_kv_heads=2, + head_size=3, + dtype=torch.bfloat16, + ) + + +def _mla_spec(block_size: int = 4) -> MLAAttentionSpec: + return MLAAttentionSpec( + block_size=block_size, + num_kv_heads=1, + head_size=3, + dtype=torch.bfloat16, + ) + + +def _worker( + kv_caches: dict[str, torch.Tensor], + layer_to_spec: dict[str, object], + num_blocks: int = 8, +) -> SimpleNamespace: + return SimpleNamespace( + kv_caches=kv_caches, + layer_to_spec=layer_to_spec, + num_blocks=num_blocks, + block_size=4, + ) + + +def _remote_meta(num_blocks: int = 16) -> SimpleNamespace: + return SimpleNamespace(num_blocks=num_blocks) + + +def test_separated_kv_layout_uses_kv_axis_zero_and_block_axis_one(): + cache = torch.empty((2, 8, 4, 2, 3), dtype=torch.bfloat16) + worker = _worker({"layer": cache}, {"layer": _full_spec()}) + + geometry = moriio_layout.get_layer_transfer_geometry( + "layer", cache, worker.layer_to_spec, remote_num_blocks=16 + ) + assert geometry.block_stride == 24 + assert geometry.local_kv_stride == 192 + assert geometry.remote_kv_stride == 384 + assert geometry.split_kv_regions + + assert moriio_layout.compute_block_transfer_offsets( + "layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks + ) == ([48, 144, 432, 528], [192, 240, 960, 1008], [48, 48, 48, 48]) + + +def test_interleaved_kv_layout_uses_block_axis_zero_and_kv_axis_one(): + cache = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16) + worker = _worker({"layer": cache}, {"layer": _full_spec()}) + + geometry = moriio_layout.get_layer_transfer_geometry( + "layer", cache, worker.layer_to_spec, remote_num_blocks=16 + ) + assert geometry.block_stride == 48 + assert geometry.local_kv_stride == 24 + assert geometry.remote_kv_stride == 24 + assert not geometry.split_kv_regions + + assert moriio_layout.compute_block_transfer_offsets( + "layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks + ) == ([96, 288], [384, 480], [96, 96]) + + +def test_mla_key_only_layout_transfers_one_slab_per_block(): + cache = torch.empty((8, 4, 3), dtype=torch.bfloat16) + worker = _worker({"layer": cache}, {"layer": _mla_spec()}) + + geometry = moriio_layout.get_layer_transfer_geometry( + "layer", cache, worker.layer_to_spec, remote_num_blocks=16 + ) + assert geometry.block_stride == 12 + assert geometry.local_kv_stride is None + assert geometry.remote_kv_stride is None + assert geometry.transfers_per_block == 1 + + assert moriio_layout.compute_block_transfer_offsets( + "layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks + ) == ([24, 72], [96, 120], [24, 24]) + + +def test_mixed_layers_compute_distinct_offsets_per_layer(): + kv_caches = { + "separated": torch.empty((2, 8, 4, 2, 3), dtype=torch.bfloat16), + "interleaved": torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16), + "indexer": torch.empty((8, 4, 3), dtype=torch.bfloat16), + } + worker = _worker( + kv_caches, + { + "separated": _full_spec(), + "interleaved": _full_spec(), + "indexer": _mla_spec(), + }, + ) + + separated = moriio_layout.compute_block_transfer_offsets( + "separated", + kv_caches["separated"], + worker.layer_to_spec, + [1, 3], + [4, 5], + _remote_meta().num_blocks, + ) + interleaved = moriio_layout.compute_block_transfer_offsets( + "interleaved", + kv_caches["interleaved"], + worker.layer_to_spec, + [1, 3], + [4, 5], + _remote_meta().num_blocks, + ) + indexer = moriio_layout.compute_block_transfer_offsets( + "indexer", + kv_caches["indexer"], + worker.layer_to_spec, + [1, 3], + [4, 5], + _remote_meta().num_blocks, + ) + + assert separated != interleaved + assert separated != indexer + assert interleaved != indexer + + +def test_block_id_length_mismatch_raises_value_error(): + cache = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16) + worker = _worker({"layer": cache}, {"layer": _full_spec()}) + + with pytest.raises(ValueError, match="must have the same length"): + moriio_layout.compute_block_transfer_offsets( + "layer", cache, worker.layer_to_spec, [1, 3], [4], _remote_meta().num_blocks + ) + + +def test_registration_regions_do_not_split_interleaved_or_mla_cache(): + separated = torch.empty((2, 8, 4, 2, 3), dtype=torch.bfloat16) + interleaved = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16) + indexer = torch.empty((8, 4, 3), dtype=torch.bfloat16) + worker = _worker( + { + "separated": separated, + "interleaved": interleaved, + "indexer": indexer, + }, + { + "separated": _full_spec(), + "interleaved": _full_spec(), + "indexer": _mla_spec(), + }, + ) + + separated_regions = moriio_layout.iter_layer_registration_regions( + "separated", separated, worker.layer_to_spec + ) + interleaved_regions = moriio_layout.iter_layer_registration_regions( + "interleaved", interleaved, worker.layer_to_spec + ) + indexer_regions = moriio_layout.iter_layer_registration_regions( + "indexer", indexer, worker.layer_to_spec + ) + + assert [region[0].data_ptr() for region in separated_regions] == [ + separated[0].data_ptr(), + separated[1].data_ptr(), + ] + assert separated_regions[0][1] == 8 * 48 + assert separated_regions[1][1] == 8 * 48 + + assert len(interleaved_regions) == 1 + assert interleaved_regions[0][0].data_ptr() == interleaved.data_ptr() + assert interleaved_regions[0][1] == 8 * 2 * 48 + + assert len(indexer_regions) == 1 + assert indexer_regions[0][0].data_ptr() == indexer.data_ptr() + assert indexer_regions[0][1] == 8 * 24 + + +def test_registration_regions_use_layer_num_blocks(): + cache = torch.empty((4, 2, 4, 2, 3), dtype=torch.bfloat16) + worker = _worker({"layer": cache}, {"layer": _full_spec()}, num_blocks=8) + + regions = moriio_layout.iter_layer_registration_regions( + "layer", cache, worker.layer_to_spec + ) + + assert len(regions) == 1 + assert regions[0][1] == 4 * 2 * 48 + + +def test_unsupported_shape_raises_value_error(): + cache = torch.empty((8, 4, 2, 3), dtype=torch.bfloat16) + worker = _worker({"layer": cache}, {"layer": _full_spec()}) + + with pytest.raises(ValueError, match="Unsupported MoRIIO K/V cache shape"): + moriio_layout.get_layer_transfer_geometry("layer", cache, worker.layer_to_spec) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py index b5552f72046..a41bb5789f0 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py @@ -47,6 +47,14 @@ from vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_engine import ( MoRIIOWrapper, MoRIIOWriter, ) +from vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_layout import ( + LayerTransferGeometry, + build_layer_to_spec, + compute_block_transfer_offsets, + get_layer_transfer_geometry, + is_mla_cache_layer, + iter_layer_registration_regions, +) from vllm.distributed.parallel_state import ( get_tensor_model_parallel_world_size, get_tp_group, @@ -71,6 +79,7 @@ if TYPE_CHECKING: logger = init_logger(__name__) + try: from mori.io import ( BackendType, @@ -117,7 +126,9 @@ class MoRIIOConnector(KVConnectorBase_V1): self.connector_worker: MoRIIOConnectorWorker | None = None elif role == KVConnectorRole.WORKER: self.connector_scheduler = None - self.connector_worker = MoRIIOConnectorWorker(vllm_config, self.engine_id) + self.connector_worker = MoRIIOConnectorWorker( + vllm_config, self.engine_id, kv_cache_config + ) logger.info( "Initialized MoRIIO Connector,engine_id:%s,role: %s", self.engine_id, @@ -683,7 +694,12 @@ class MoRIIOConnectorScheduler: class MoRIIOConnectorWorker: """Implementation of Worker side methods""" - def __init__(self, vllm_config: VllmConfig, engine_id: str): + def __init__( + self, + vllm_config: VllmConfig, + engine_id: str, + kv_cache_config: "KVCacheConfig", + ): if not is_moriio_available(): raise RuntimeError( "MoRIIO is not available. Please ensure the 'mori' package " @@ -707,6 +723,7 @@ class MoRIIOConnectorWorker: ) self.kv_transfer_config = vllm_config.kv_transfer_config self.is_producer = self.kv_transfer_config.is_kv_producer + self.layer_to_spec = build_layer_to_spec(kv_cache_config) if self.is_producer: set_role(ROLE.PRODUCER) @@ -809,6 +826,8 @@ class MoRIIOConnectorWorker: self.kv_cache_shape = None self.block_shape = None self.kv_element_size = 0 + self.kv_cache_shapes: dict[str, torch.Size] = {} + self.block_lens: dict[str, int] = {} # Map of engine_id -> {agent_name0, agent_name1..}. self._remote_agents: dict[EngineId, set[str]] = {} @@ -1218,51 +1237,86 @@ class MoRIIOConnectorWorker: all_done_future = self._handshake_initiation_executor.submit(wait_all_dp) all_done_future.add_done_callback(request_ready) + def _is_mla_cache_layer(self, layer_name: str) -> bool: + return is_mla_cache_layer(self.layer_to_spec, layer_name) + + def _get_layer_transfer_geometry( + self, layer_name: str, remote_num_blocks: int | None = None + ) -> LayerTransferGeometry: + return get_layer_transfer_geometry( + layer_name, + self.kv_caches[layer_name], + self.layer_to_spec, + remote_num_blocks, + ) + + def _iter_layer_registration_regions( + self, layer_name: str + ) -> list[tuple[torch.Tensor, int]]: + return iter_layer_registration_regions( + layer_name, + self.kv_caches[layer_name], + self.layer_to_spec, + ) + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): """Register the KV Cache data in moriio.""" - _, first_kv_cache = next(iter(kv_caches.items())) + self.kv_caches = kv_caches # layer name to kv cache + self.kv_cache_shapes = { + layer_name: kv_cache.shape for layer_name, kv_cache in kv_caches.items() + } + + first_layer_name, first_kv_cache = next( + ( + (layer_name, kv_cache) + for layer_name, kv_cache in kv_caches.items() + if ( + not self._is_mla_cache_layer(layer_name) + and len(kv_cache.shape) == 5 + and (kv_cache.shape[0] == 2 or kv_cache.shape[1] == 2) + ) + ), + next(iter(kv_caches.items())), + ) kv_elem_size = first_kv_cache.element_size() - use_mla = len(first_kv_cache.shape) == 3 - assert use_mla == self.use_mla + use_mla = self._is_mla_cache_layer(first_layer_name) + first_geometry = self._get_layer_transfer_geometry(first_layer_name) if use_mla: # MLA case. - self.num_blocks = first_kv_cache.shape[0] block_rank = 2 # [block_size, latent_dim] block_shape = first_kv_cache.shape[-block_rank:] - block_size, kv_latent_dim = block_shape - self.slot_size_bytes = kv_elem_size * kv_latent_dim else: - # [2 (k and v), num_blocks, ...] - self.num_blocks = first_kv_cache.shape[1] + # [2, num_blocks, ...] or [num_blocks, 2, ...] block_rank = 3 # [block_size, kv_heads, head_dim] block_shape = first_kv_cache.shape[-block_rank:] - block_size, n_kv_heads, head_dim = block_shape[-3:] - # head size in bytes. - self.slot_size_bytes = ( - kv_elem_size * n_kv_heads * head_dim - ) # 1 token 1 layer size , slot size - assert block_size == self.block_size + self.num_blocks = first_geometry.num_blocks + self.slot_size_bytes = first_geometry.slot_size_bytes + assert first_geometry.block_size == self.block_size # TODO(tms): self.block_len needs to be per-layer for sliding window, # hybrid attn, etc # block size in bytes - self.block_len = kv_elem_size * math.prod(block_shape) + self.block_len = first_geometry.block_len self.kv_cache_shape = first_kv_cache.shape self.block_shape = block_shape self.kv_element_size = kv_elem_size self.dst_num_blocks[self.engine_id] = self.num_blocks - self.kv_caches = kv_caches # layer name to kv cache kv_caches_base_addr = [] caches_data = [] - for cache_or_caches in kv_caches.values(): - cache_list = [cache_or_caches] if use_mla else cache_or_caches - for cache in cache_list: + for layer_name in kv_caches: + geometry = self._get_layer_transfer_geometry(layer_name) + if geometry.block_size != self.block_size: + raise ValueError( + "MoRIIO KV cache block size mismatch for layer " + f"{layer_name}: {geometry.block_size} != {self.block_size}" + ) + self.block_lens[layer_name] = geometry.block_len + for cache, region_len in self._iter_layer_registration_regions(layer_name): base_addr = cache.data_ptr() - region_len = self.num_blocks * self.block_len caches_data.append((base_addr, region_len, cache.device.index, "")) kv_caches_base_addr.append(base_addr) @@ -1275,7 +1329,9 @@ class MoRIIOConnectorWorker: moriio_mem_metadata ) - self.local_kv_cache_size.append(cache.nelement() * cache.element_size()) + self.local_kv_cache_size.append( + kv_cache.nelement() * kv_cache.element_size() + ) self.kv_caches_base_addr[self.engine_id] = kv_caches_base_addr self.num_regions = len(caches_data) @@ -1666,47 +1722,17 @@ class MoRIIOConnectorWorker: Returns: Tuple of (local_offsets, remote_offsets, transfer_sizes) """ - assert self.kv_cache_shape is not None, "KV caches shape not initialized" - is_mla = len(self.kv_cache_shape) == 3 - stride = self.kv_caches[layer_name].stride() - sz = self.kv_caches[layer_name].element_size() - if is_mla: - blknum, blksize, hs = self.kv_cache_shape - hn = 1 - block_stride = stride[0] - else: - _, blknum, blksize, hn, hs = self.kv_cache_shape - local_ktov_stride = stride[0] - block_stride = stride[1] - remote_ktov_stride = block_stride * remote_moriio_meta.num_blocks - - transfer_size_byte = blksize * hn * hs * sz - per_block = 1 if is_mla else 2 - total = len(local_block_ids) * per_block - offset_local = [0] * total - offset_remote = [0] * total - sizes = [transfer_size_byte] * total - - w = 0 - for i, lb in enumerate(local_block_ids): - rb = remote_block_ids[i] - # K - offset_local[w] = sz * (lb * block_stride) - offset_remote[w] = sz * (rb * block_stride) - w += 1 - if not is_mla: - # V - # Handle num_block variations originating from PD (different kv strides) - # TODO: address block_sz differences in heterogeneous TP scenarios - # In MLA, we don't need to consider these two cases. - offset_local[w] = sz * (1 * local_ktov_stride + lb * block_stride) - offset_remote[w] = sz * (1 * remote_ktov_stride + rb * block_stride) - w += 1 - - merged_l, merged_r, merged_s = self.merge_contiguous_blocks( - offset_local, offset_remote, sizes, assume_sorted=False + return compute_block_transfer_offsets( + layer_name=layer_name, + kv_cache=self.kv_caches[layer_name], + layer_to_spec=self.layer_to_spec, + local_block_ids=local_block_ids, + remote_block_ids=remote_block_ids, + remote_num_blocks=remote_moriio_meta.num_blocks, + merge_fn=lambda local, remote, sizes: self.merge_contiguous_blocks( + local, remote, sizes, assume_sorted=False + ), ) - return merged_l, merged_r, merged_s def _read_blocks( self, @@ -1724,15 +1750,13 @@ class MoRIIOConnectorWorker: dp0_engine_id = self.get_engine_name_with_dp(dst_engine_id, 0) sessions, remote_moriio_meta = self._get_built_session(dp0_engine_id) - first_layer = list(self.layer_name_to_local_kv_cache_metadata.keys())[0] - offs = self._compute_block_transfer_offsets( - first_layer, local_block_ids, remote_block_ids, remote_moriio_meta - ) - for layer_name in self.layer_name_to_local_kv_cache_metadata: sess_idx = list(self.layer_name_to_local_kv_cache_metadata.keys()).index( layer_name ) + offs = self._compute_block_transfer_offsets( + layer_name, local_block_ids, remote_block_ids, remote_moriio_meta + ) # TODO : apply multi-session batch-read when moriio support it transfer_status = self.moriio_wrapper.read_remote_data( offs[2], offs[0], offs[1], sessions[sess_idx] diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_layout.py b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_layout.py new file mode 100644 index 00000000000..8a6aced9daa --- /dev/null +++ b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_layout.py @@ -0,0 +1,213 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from collections.abc import Callable, Mapping +from typing import NamedTuple + +import torch + +from vllm.v1.kv_cache_interface import ( + KVCacheConfig, + KVCacheSpec, + MLAAttentionSpec, + SlidingWindowMLASpec, + UniformTypeKVCacheSpecs, +) + + +class LayerTransferGeometry(NamedTuple): + num_blocks: int + block_size: int + block_len: int + slot_size_bytes: int + block_stride: int + local_kv_stride: int | None + remote_kv_stride: int | None + transfers_per_block: int + regions_per_block: int + split_kv_regions: bool + + +def build_layer_to_spec(kv_cache_config: KVCacheConfig) -> dict[str, KVCacheSpec]: + layer_to_spec: dict[str, KVCacheSpec] = {} + for group in kv_cache_config.kv_cache_groups: + group_spec = group.kv_cache_spec + if isinstance(group_spec, UniformTypeKVCacheSpecs): + layer_to_spec.update( + { + layer_name: group_spec.kv_cache_specs[layer_name] + for layer_name in group.layer_names + } + ) + else: + layer_to_spec.update( + {layer_name: group_spec for layer_name in group.layer_names} + ) + return layer_to_spec + + +def is_mla_cache_layer( + layer_to_spec: Mapping[str, KVCacheSpec], layer_name: str +) -> bool: + try: + spec = layer_to_spec[layer_name] + except KeyError as e: + raise ValueError(f"Missing KV cache spec for layer {layer_name}") from e + return isinstance(spec, (MLAAttentionSpec, SlidingWindowMLASpec)) + + +def get_layer_transfer_geometry( + layer_name: str, + kv_cache: torch.Tensor, + layer_to_spec: Mapping[str, KVCacheSpec], + remote_num_blocks: int | None = None, +) -> LayerTransferGeometry: + shape = kv_cache.shape + stride = kv_cache.stride() + element_size = kv_cache.element_size() + is_mla_cache = is_mla_cache_layer(layer_to_spec, layer_name) + + if is_mla_cache and len(shape) == 3: + num_blocks, block_size, latent_dim = shape + slot_size_bytes = latent_dim * element_size + block_len = block_size * slot_size_bytes + return LayerTransferGeometry( + num_blocks=num_blocks, + block_size=block_size, + block_len=block_len, + slot_size_bytes=slot_size_bytes, + block_stride=stride[0], + local_kv_stride=None, + remote_kv_stride=None, + transfers_per_block=1, + regions_per_block=1, + split_kv_regions=False, + ) + + if not is_mla_cache and len(shape) == 5 and shape[0] == 2: + _, num_blocks, block_size, num_kv_heads, head_dim = shape + slot_size_bytes = num_kv_heads * head_dim * element_size + block_len = block_size * slot_size_bytes + remote_kv_stride = stride[1] * (remote_num_blocks or num_blocks) + return LayerTransferGeometry( + num_blocks=num_blocks, + block_size=block_size, + block_len=block_len, + slot_size_bytes=slot_size_bytes, + block_stride=stride[1], + local_kv_stride=stride[0], + remote_kv_stride=remote_kv_stride, + transfers_per_block=2, + regions_per_block=1, + split_kv_regions=True, + ) + + if not is_mla_cache and len(shape) == 5 and shape[1] == 2: + num_blocks, _, block_size, num_kv_heads, head_dim = shape + slot_size_bytes = num_kv_heads * head_dim * element_size + block_len = block_size * slot_size_bytes + return LayerTransferGeometry( + num_blocks=num_blocks, + block_size=block_size, + block_len=block_len, + slot_size_bytes=slot_size_bytes, + block_stride=stride[0], + local_kv_stride=stride[1], + remote_kv_stride=stride[1], + transfers_per_block=2, + regions_per_block=2, + split_kv_regions=False, + ) + + cache_kind = "MLA" if is_mla_cache else "K/V" + raise ValueError( + f"Unsupported MoRIIO {cache_kind} cache shape for layer " + f"{layer_name}: {tuple(shape)}" + ) + + +def iter_layer_registration_regions( + layer_name: str, + kv_cache: torch.Tensor, + layer_to_spec: Mapping[str, KVCacheSpec], +) -> list[tuple[torch.Tensor, int]]: + geometry = get_layer_transfer_geometry(layer_name, kv_cache, layer_to_spec) + region_len = geometry.num_blocks * geometry.regions_per_block * geometry.block_len + if geometry.split_kv_regions: + return [(cache, region_len) for cache in kv_cache] + return [(kv_cache, region_len)] + + +def merge_contiguous_offsets( + offsets_local: list[int], + offsets_remote: list[int], + sizes: list[int], +) -> tuple[list[int], list[int], list[int]]: + if not offsets_local: + return [], [], [] + if not (len(offsets_local) == len(offsets_remote) == len(sizes)): + raise ValueError("Input list lengths mismatch") + + rows = sorted(zip(offsets_local, offsets_remote, sizes), key=lambda row: row[0]) + merged: list[list[int]] = [] + for local, remote, size in rows: + if ( + merged + and local == merged[-1][0] + merged[-1][2] + and remote == merged[-1][1] + merged[-1][2] + ): + merged[-1][2] += size + else: + merged.append([local, remote, size]) + + return ( + [row[0] for row in merged], + [row[1] for row in merged], + [row[2] for row in merged], + ) + + +def compute_block_transfer_offsets( + layer_name: str, + kv_cache: torch.Tensor, + layer_to_spec: Mapping[str, KVCacheSpec], + local_block_ids: list[int], + remote_block_ids: list[int], + remote_num_blocks: int, + merge_fn: Callable[ + [list[int], list[int], list[int]], tuple[list[int], list[int], list[int]] + ] = merge_contiguous_offsets, +) -> tuple[list[int], list[int], list[int]]: + if len(local_block_ids) != len(remote_block_ids): + raise ValueError( + "local_block_ids and remote_block_ids must have the same length: " + f"{len(local_block_ids)} != {len(remote_block_ids)}" + ) + geometry = get_layer_transfer_geometry( + layer_name, kv_cache, layer_to_spec, remote_num_blocks + ) + element_size = kv_cache.element_size() + transfer_size_byte = geometry.block_len + per_block = geometry.transfers_per_block + total = len(local_block_ids) * per_block + offset_local = [0] * total + offset_remote = [0] * total + sizes = [transfer_size_byte] * total + + w = 0 + for lb, rb in zip(local_block_ids, remote_block_ids): + offset_local[w] = element_size * (lb * geometry.block_stride) + offset_remote[w] = element_size * (rb * geometry.block_stride) + w += 1 + if per_block == 2: + assert geometry.local_kv_stride is not None + assert geometry.remote_kv_stride is not None + offset_local[w] = element_size * ( + geometry.local_kv_stride + lb * geometry.block_stride + ) + offset_remote[w] = element_size * ( + geometry.remote_kv_stride + rb * geometry.block_stride + ) + w += 1 + + return merge_fn(offset_local, offset_remote, sizes)