[ROCm][P/D] Support MiniMax-M3 mixed KV layouts in MoRIIO READ mode (#46039)

Signed-off-by: Jun Kang Chow <junkangchow@gmail.com>
Signed-off-by: tjtanaa <tunjian.tan@embeddedllm.com>
Co-authored-by: Hongxia Yang <hongxia.yang@amd.com>
Co-authored-by: Tan Pin Siang <tanpinsiang@gmail.com>
Co-authored-by: vllmellm <vllm.ellm@embeddedllm.com>
Co-authored-by: Chun Fang <chun.fang@amd.com>
Co-authored-by: TianDi101 <ditian12@amd.com>
Co-authored-by: functionstackx <47992694+functionstackx@users.noreply.github.com>
Co-authored-by: tjtanaa <tunjian.tan@embeddedllm.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
junkang1991
2026-06-21 12:55:19 +00:00
committed by GitHub
co-authored by Hongxia Yang Tan Pin Siang vllmellm Chun Fang TianDi101 functionstackx tjtanaa mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent d3ad8e8bcd
commit b91b7726e0
4 changed files with 566 additions and 72 deletions
@@ -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(
@@ -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)
@@ -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]
@@ -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)