forked from Karylab-cklius/vllm
[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:
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
@@ -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)
|
||||
Reference in New Issue
Block a user