init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,339 @@
import threading
from collections.abc import Iterable
from typing import TYPE_CHECKING, Any
import torch
import zmq
from vllm.config import VllmConfig
from vllm.distributed.kv_events import (
KVCacheEvent,
KVConnectorKVEvents,
KVEventAggregator,
)
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorBase_V1,
KVConnectorMetadata,
KVConnectorRole,
SupportsHMA,
)
from vllm.forward_context import ForwardContext
from vllm.logger import logger
from vllm.utils.network_utils import make_zmq_socket
from vllm.v1.attention.backend import AttentionMetadata # type: ignore
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.outputs import KVConnectorOutput
from vllm.v1.request import Request
from vllm.v1.serial_utils import MsgpackDecoder
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import AscendStoreKVConnectorWorkerMetadata
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.pool_scheduler import (
KVPoolScheduler,
get_zmq_rpc_path_lookup,
)
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.pool_worker import KVPoolWorker
if TYPE_CHECKING:
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorHandshakeMetadata
class AscendStoreKVEvents(KVConnectorKVEvents):
def __init__(self, num_workers: int) -> None:
self._aggregator = KVEventAggregator(num_workers)
def add_events(self, events: list[KVCacheEvent]) -> None:
self._aggregator.add_events(events)
def aggregate(self) -> "AscendStoreKVEvents":
"""
Aggregate KV events and retain only common events.
"""
common_events = self._aggregator.get_common_events()
self._aggregator.clear_events()
self._aggregator.add_events(common_events)
self._aggregator.reset_workers()
return self
def increment_workers(self, count: int = 1) -> None:
self._aggregator.increment_workers(count)
def get_all_events(self) -> list[KVCacheEvent]:
return self._aggregator.get_all_events()
def get_number_of_workers(self) -> int:
return self._aggregator.get_number_of_workers()
def clear_events(self) -> None:
self._aggregator.clear_events()
self._aggregator.reset_workers()
def __repr__(self) -> str:
return f"<AscendStoreKVEvents events={self.get_all_events()}>"
class AscendStoreConnector(KVConnectorBase_V1, SupportsHMA):
@classmethod
def requires_piecewise_for_cudagraph(cls, extra_config: dict[str, Any]) -> bool:
"""
AscendStore requires PIECEWISE CUDA graph mode when layerwise
operations are enabled.
"""
return extra_config.get("use_layerwise", False)
def __init__(self, vllm_config: VllmConfig, role: KVConnectorRole, kv_cache_config: KVCacheConfig | None = None):
super().__init__(vllm_config=vllm_config, role=role, kv_cache_config=kv_cache_config)
self.kv_role = vllm_config.kv_transfer_config.kv_role
self.use_layerwise = vllm_config.kv_transfer_config.kv_connector_extra_config.get("use_layerwise", False)
backend_name = vllm_config.kv_transfer_config.kv_connector_extra_config.get("backend", "mooncake")
self.backend_name = backend_name.lower()
self.use_gva_layerwise = self.use_layerwise and self.backend_name == "memcache"
self.consumer_is_to_put = vllm_config.kv_transfer_config.kv_connector_extra_config.get(
"consumer_is_to_put", False
)
connector_name = vllm_config.kv_transfer_config.kv_connector
if connector_name == "MooncakeConnectorStoreV1":
logger.warning(
"It is recommended to use the AscendStoreConnector, "
"as the MoonCakeStoreConnector will be removed in the future."
)
self.kv_caches: dict[str, torch.Tensor] = {}
self._kv_cache_events: AscendStoreKVEvents | None = None
self._current_step_has_real_forward = False
if role == KVConnectorRole.SCHEDULER:
assert kv_cache_config is not None
page_size_bytes = kv_cache_config.kv_cache_groups[0].kv_cache_spec.page_size_bytes
self.connector_scheduler = KVPoolScheduler(
vllm_config, self.use_layerwise, kv_cache_config, page_size_bytes=page_size_bytes
)
else:
self.connector_worker = KVPoolWorker(
vllm_config,
self.use_layerwise,
kv_cache_config,
)
assert self.connector_worker is not None
if not self.use_layerwise and vllm_config.parallel_config.rank == 0:
self.lookup_server = LookupKeyServer(self.connector_worker, vllm_config)
############################################################
# Scheduler Side Methods
############################################################
def set_xfer_handshake_metadata_pp_aware(
self,
metadata: dict[tuple[int, int], "KVConnectorHandshakeMetadata"],
) -> None:
"""Ignore P/D handshake metadata because AscendStore handles PP via pool keys."""
pass
def get_num_new_matched_tokens(self, request: "Request", num_computed_tokens: int) -> tuple[int, bool]:
assert self.connector_scheduler is not None
return self.connector_scheduler.get_num_new_matched_tokens(request, num_computed_tokens)
def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", num_external_tokens: int):
assert self.connector_scheduler is not None
return self.connector_scheduler.update_state_after_alloc(request, blocks, num_external_tokens)
def build_connector_meta(
self,
scheduler_output: SchedulerOutput,
) -> KVConnectorMetadata:
assert self.connector_scheduler is not None
return self.connector_scheduler.build_connector_meta(scheduler_output)
def request_finished(
self,
request: "Request",
block_ids: list[int],
) -> tuple[bool, dict[str, Any] | None]:
assert self.connector_scheduler is not None
return self.connector_scheduler.request_finished(request, block_ids)
def request_finished_all_groups(
self,
request: "Request",
block_ids: tuple[list[int], ...],
) -> tuple[bool, dict[str, Any] | None]:
assert self.connector_scheduler is not None
return self.connector_scheduler.request_finished_all_groups(request, block_ids)
def update_connector_output(self, connector_output: KVConnectorOutput):
"""
Update KVConnector state from worker-side connectors output.
Args:
connector_output (KVConnectorOutput): the worker-side connectors output.
"""
if self.connector_scheduler is not None:
self.connector_scheduler.update_connector_output(connector_output)
# Get the KV events
kv_cache_events = connector_output.kv_cache_events
if not kv_cache_events or not isinstance(kv_cache_events, AscendStoreKVEvents):
return
if self._kv_cache_events is None:
self._kv_cache_events = kv_cache_events
else:
self._kv_cache_events.add_events(kv_cache_events.get_all_events())
self._kv_cache_events.increment_workers(kv_cache_events.get_number_of_workers())
return
def take_events(self) -> Iterable["KVCacheEvent"]:
"""
Take the KV cache events from the connector.
Yields:
New KV cache events since the last call.
"""
if self._kv_cache_events is not None:
self._kv_cache_events.aggregate()
kv_cache_events = self._kv_cache_events.get_all_events()
yield from kv_cache_events
self._kv_cache_events.clear_events()
self._kv_cache_events = None
############################################################
# Worker Side Methods
############################################################
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
assert self.connector_worker is not None
self.connector_worker.register_kv_caches(kv_caches)
def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None:
assert self.connector_worker is not None
metadata = self._get_connector_metadata()
self._current_step_has_real_forward = forward_context is not None
logger.debug(
"KV pool connector start_load_kv metadata_requests=%d specs=%s",
len(metadata.requests),
[
(
request.req_id,
None if request.load_spec is None else request.load_spec.can_load,
None if request.load_spec is None else request.load_spec.vllm_cached_tokens,
None if request.load_spec is None else request.load_spec.kvpool_cached_tokens,
)
for request in metadata.requests
],
)
self.connector_worker.start_load_kv(metadata)
def wait_for_layer_load(self, layer_name: str) -> None:
if not self.use_layerwise:
return
self.connector_worker.wait_for_layer_load()
def save_kv_layer(
self, layer_name: str, kv_layer: torch.Tensor, attn_metadata: "AttentionMetadata", **kwargs
) -> None:
if not self.use_layerwise:
return
if self.kv_role == "kv_consumer":
# Don't do save if the role is kv_consumer
return
self.connector_worker.save_kv_layer(self._get_connector_metadata())
def wait_for_save(self):
if self.kv_role == "kv_consumer" and not self.consumer_is_to_put:
# Don't do save if the role is kv_consumer
return
if self.use_layerwise:
return
self.connector_worker.wait_for_save(self._get_connector_metadata())
def get_finished(self, finished_req_ids: set[str]) -> tuple[set[str], set[str]]:
"""Get the finished recving and sending requests."""
assert self.connector_worker is not None
metadata = self._get_connector_metadata()
if self._current_step_has_real_forward:
try:
self.connector_worker.ensure_store_initialized()
finally:
self._current_step_has_real_forward = False
done_sending, done_recving = self.connector_worker.get_finished(finished_req_ids, metadata)
return done_sending, done_recving
def get_block_ids_with_load_errors(self) -> set[int]:
"""Return KV block IDs that failed to load on the worker."""
assert self.connector_worker is not None
return self.connector_worker.get_block_ids_with_load_errors()
def get_kv_connector_kv_cache_events(self) -> AscendStoreKVEvents | None:
"""
Get the KV connector kv cache events collected during the last interval.
"""
events = self.connector_worker.get_kv_events()
if not events:
return None
ascend_store_kv_events = AscendStoreKVEvents(num_workers=1)
ascend_store_kv_events.add_events(events)
return ascend_store_kv_events
def bind_gpu_block_pool(self, gpu_block_pool: "BlockPool") -> None:
assert self.connector_scheduler is not None
self.connector_scheduler.bind_gpu_block_pool(gpu_block_pool)
def build_connector_worker_meta(self) -> AscendStoreKVConnectorWorkerMetadata | None:
assert self.connector_worker is not None
return self.connector_worker.build_connector_worker_meta()
class LookupKeyServer:
def __init__(
self,
pool_worker: KVPoolWorker,
vllm_config: "VllmConfig",
):
self.decoder = MsgpackDecoder()
self.ctx = zmq.Context() # type: ignore[attr-defined]
socket_path = get_zmq_rpc_path_lookup(vllm_config)
self.socket = make_zmq_socket(
self.ctx,
socket_path,
zmq.REP, # type: ignore[attr-defined]
bind=True,
)
self.pool_worker = pool_worker
self.running = True
def process_request():
while self.running:
all_frames = self.socket.recv_multipart(copy=False)
token_len = int.from_bytes(all_frames[0], byteorder="big")
kv_group_ids = self.decoder.decode([all_frames[1]])
hbm_hit_tokens = int.from_bytes(all_frames[2], byteorder="big")
hashes_str = self.decoder.decode(all_frames[3:])
result = self.pool_worker.lookup_scheduler(
token_len,
hashes_str,
kv_group_ids,
use_layerwise=False,
hbm_hit_tokens=hbm_hit_tokens,
)
logger.debug(
"KV pool lookup response token_len=%d groups=%s hit_tokens=%d",
token_len,
kv_group_ids,
result,
)
response = result.to_bytes(4, "big")
self.socket.send(response)
self.thread = threading.Thread(target=process_request, daemon=True)
self.thread.start()
def close(self):
self.socket.close(linger=0)

View File

@@ -0,0 +1,30 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
backend_map = {
"mooncake": {
"name": "MooncakeBackend",
"path": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend",
},
"memcache": {
"name": "MemcacheBackend",
"path": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend",
},
"yuanrong": {
"name": "YuanrongBackend",
"path": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend",
},
}

View File

@@ -0,0 +1,56 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Any
from vllm.config import ParallelConfig
class Backend(ABC):
store: Any | None = None
@abstractmethod
def __init__(self, parallel_config: ParallelConfig):
pass
@classmethod
def create_scheduler_client(cls, parallel_config: ParallelConfig):
return cls(parallel_config)
@abstractmethod
def set_device(self):
pass
@abstractmethod
def register_buffer(self, ptrs: list[int], lengths: list[int]):
pass
@abstractmethod
def exists(self, keys: list[str]) -> list[int]:
pass
def batch_is_exist(self, keys: list[str]) -> list[int]:
return self.exists(keys)
def batch_get_key_info(self, keys: list[str]):
raise NotImplementedError(f"{type(self).__name__} does not support batch_get_key_info")
def batch_alloc(self, keys: list[str], sizes: list[int]) -> list[int]:
raise NotImplementedError(f"{type(self).__name__} does not support batch_alloc")
def batch_add_lease(self, keys: list[str], lease_ttl_ms: int = 0) -> list[int]:
raise NotImplementedError(f"{type(self).__name__} does not support batch_add_lease")
def batch_remove_lease(self, keys: list[str]) -> int:
raise NotImplementedError(f"{type(self).__name__} does not support batch_remove_lease")
def batch_write_finish(self, keys: list[str], results: list[int]) -> list[int]:
raise NotImplementedError(f"{type(self).__name__} does not support batch_write_finish")
@abstractmethod
def put(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]):
pass
@abstractmethod
def get(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]):
pass

View File

@@ -0,0 +1,233 @@
# Standard
import os
import threading
import time
from enum import Enum
from typing import Any
import torch
from vllm.config import ParallelConfig
from vllm.distributed.parallel_state import get_world_group
from vllm.logger import logger
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.backend import Backend
MEMCACHE_THREAD_START_WAIT_S = 0.1
def _is_device_sdma() -> bool:
config_path = os.getenv("MMC_LOCAL_CONFIG_PATH")
if not config_path:
raise ValueError("The environment variable 'MMC_LOCAL_CONFIG_PATH' is not set.")
with open(config_path, encoding="utf-8") as config_file:
for line in config_file:
line = line.strip()
if not line or line.startswith(("#", ";")):
continue
key, separator, value = line.partition("=")
if separator and key.strip() == "ock.mmc.local_service.protocol":
return value.strip() == "device_sdma"
return False
class MmcDirect(Enum):
COPY_L2G = 0
COPY_G2L = 1
COPY_G2H = 2
COPY_H2G = 3
class MemcacheBackend(Backend):
def __init__(
self,
parallel_config: ParallelConfig,
local_rank: int | None = None,
init_bm: bool = True,
lazy_init: bool = False,
):
self.local_rank = local_rank if local_rank is not None else get_world_group().local_rank
self._init_bm = init_bm
self._lazy_init = lazy_init and _is_device_sdma()
self.store: Any | None = None
self._store_initialized = False
self._store_init_lock = threading.Lock()
self._pending_buffers: tuple[list[int], list[int]] | None = None
if not self._lazy_init:
self.store = self._setup_store()
self._store_initialized = True
def ensure_initialized(self):
if self._store_initialized:
return
with self._store_init_lock:
if self._store_initialized:
return
logger.info("Initializing Memcache store. local_rank=%d", self.local_rank)
self.store = self._setup_store()
self._store_initialized = True
self._register_buffers_if_needed()
def _setup_store(self):
try:
from memcache_hybrid import DistributedObjectStore # type: ignore
except ImportError as e:
raise ImportError(
"Please install memcache by following the instructions at "
"https://gitee.com/ascend/memfabric_hybrid " # noqa: E501
"to run vLLM with MemcacheConnector."
) from e
store = DistributedObjectStore()
try:
res = store.init(self.local_rank, init_bm=self._init_bm)
except ValueError as e:
logger.error("Configuration loading failed. error=%s. Check memcache config and environment.", e)
raise
except Exception as exc:
logger.error("Store initialization failed. error=%s. Check memcache setup and dependencies.", exc)
raise
assert res == 0
time.sleep(MEMCACHE_THREAD_START_WAIT_S)
return store
@classmethod
def create_scheduler_client(cls, parallel_config: ParallelConfig):
# The scheduler is a single metadata client. It is initialized before
# the world group exists and must not initialize memcache storage, so
# keep the old device_id=0/init_bm=False behavior here.
return cls(parallel_config, local_rank=0, init_bm=False)
def init_store(self, init_bm: bool = True):
if self.store is not None:
return
self._init_bm = init_bm
self.store = self._setup_store()
self._store_initialized = True
self._register_buffers_if_needed()
def set_device(self):
device = torch.device(f"npu:{self.local_rank}")
torch.npu.set_device(device)
def register_buffer(self, ptrs: list[int], sizes: list[int]):
self._pending_buffers = (list(ptrs), list(sizes))
self._register_buffers_if_needed()
def _register_buffers_if_needed(self):
if self._pending_buffers is None or not self._store_initialized:
return
assert self.store is not None
ptrs, sizes = self._pending_buffers
for ptr, size in zip(ptrs, sizes):
self.store.register_buffer(ptr, size)
self._pending_buffers = None
def exists(self, keys: list[str]) -> list[int]:
if self._lazy_init and not self._store_initialized:
logger.debug(
"MemcacheBackend.exists called before store initialization; treating %d keys as missing.",
len(keys),
)
return [0] * len(keys)
assert self.store is not None
return self.store.batch_is_exist(keys)
def batch_get_key_info(self, keys: list[str]) -> list[Any]:
if self._lazy_init and not self._store_initialized:
logger.debug(
"MemcacheBackend.batch_get_key_info called before store initialization; "
"returning empty list for %d keys.",
len(keys),
)
return []
assert self.store is not None
return self.store.batch_get_key_info(keys)
def batch_alloc(self, keys: list[str], sizes: list[int]) -> list[int]:
self.ensure_initialized()
assert self.store is not None
return self.store.batch_alloc(keys, sizes)
def batch_add_lease(self, keys: list[str], lease_ttl_ms: int = 0) -> list[int]:
assert self.store is not None
return self.store.batch_add_lease(keys, lease_ttl_ms)
def batch_remove_lease(self, keys: list[str]) -> int:
assert self.store is not None
return self.store.batch_remove_lease(keys)
def batch_write_finish(self, keys: list[str], results: list[int]) -> list[int]:
assert self.store is not None
return self.store.batch_write_finish(keys, results)
def get(self, key: list[str], addr: list[list[int]], size: list[list[int]]):
if self._lazy_init and not self._store_initialized:
logger.error(
"Failed to get %d keys out of %d. Store is not initialized; "
"call put() first to trigger initialization.",
len(key),
len(key),
)
logger.debug("Failed to get key details. keys=%s", key)
return
assert self.store is not None
try:
res = self.store.batch_get_into_layers(key, addr, size, MmcDirect.COPY_G2L.value)
failed_codes = [int(value) for value in res if value != 0]
failed_count = len(failed_codes)
if failed_count:
error_codes = sorted(set(failed_codes))
logger.error(
"Failed to get %d keys out of %d. error_codes=%s. Check key existence and memory state.",
failed_count,
len(key),
error_codes,
)
logger.debug("Failed to get key details. keys=%s, result=%s", key, res)
return res
except Exception as e:
logger.error(
"Failed to get %d keys out of %d. type=%s, error=%s. Check store state and network.",
len(key),
len(key),
type(e).__name__,
e,
)
logger.debug("Failed to get key details. keys=%s", key)
return None
def put(self, key: list[str], addr: list[list[int]], size: list[list[int]]):
self.ensure_initialized()
assert self.store is not None
try:
res = self.store.batch_put_from_layers(key, addr, size, MmcDirect.COPY_L2G.value)
failed_codes = [int(value) for value in res if value != 0]
failed_count = len(failed_codes)
if failed_count:
error_codes = sorted(set(failed_codes))
logger.error(
"Failed to put %d keys out of %d. error_codes=%s. Check memory and store capacity.",
failed_count,
len(key),
error_codes,
)
logger.debug("Failed to put key details. keys=%s, result=%s", key, res)
if self._lazy_init:
logger.warning("First DSV4(compress) request failure is expected. This is normal behavior.")
except Exception as e:
logger.error(
"Failed to put %d keys out of %d. type=%s, error=%s. Check store state and memory.",
len(key),
len(key),
type(e).__name__,
e,
)
logger.debug("Failed to put key details. keys=%s", key)
if self._lazy_init:
logger.warning("First DSV4(compress) request failure is expected. This is normal behavior.")

View File

@@ -0,0 +1,402 @@
# Standard
import functools
import json
import os
import threading
from dataclasses import dataclass
from typing import Any
import regex as re
import torch
# Third Party
from mooncake.store import ReplicateConfig # type: ignore
from vllm.config import ParallelConfig
from vllm.distributed.parallel_state import get_world_group
from vllm.logger import logger
from vllm.utils.network_utils import get_ip
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.backend import Backend
from vllm_ascend.distributed.kv_transfer.utils.mooncake_transfer_engine import global_te
from vllm_ascend.distributed.parallel_state import get_global_rank
DEFAULT_GLOBAL_SEGMENT_SIZE = 1073741824 # 1.0 GiB
DEFAULT_LOCAL_BUFFER_SIZE = 1073741824 # 1.0 GiB
@functools.lru_cache(maxsize=1)
def _mooncake_setup_supports_ssd_offload() -> bool:
"""True when installed Mooncake exposes SSD kwargs on setup() (v0.3.11+)."""
from mooncake.store import MooncakeDistributedStore # type: ignore
setup = MooncakeDistributedStore.setup
try:
import inspect
sig = inspect.signature(setup)
return "enable_ssd_offload" in sig.parameters
except (TypeError, ValueError):
# pybind11 overloaded bindings often reject inspect.signature
doc = setup.__doc__ or ""
return "enable_ssd_offload" in doc
def _ssd_setup_kwargs(config: "MooncakeStoreConfig") -> dict[str, object]:
"""Keyword args for store.setup(); empty on old Mooncake or when SSD is off."""
if not config.enable_ssd_offload:
return {}
if not _mooncake_setup_supports_ssd_offload():
raise RuntimeError(
"mooncake.json has enable_ssd_offload=true, but the installed "
"Mooncake does not support enable_ssd_offload/ssd_offload_path in "
"MooncakeDistributedStore.setup(). Upgrade Mooncake to v0.3.11 or "
"later (see Mooncake ssd-offload.md Step 3A), or set "
"enable_ssd_offload to false."
)
return {
"enable_ssd_offload": config.enable_ssd_offload,
"ssd_offload_path": config.ssd_offload_path,
}
class MooncakeBackend(Backend):
def __init__(self, parallel_config: ParallelConfig, lazy_init: bool = False, contribute_memory: bool = True):
self.parallel_config = parallel_config
self.config = MooncakeStoreConfig.load_from_env()
if self.config.protocol != "ascend":
raise NotImplementedError(f"MooncakeBackend does not support protocol {self.config.protocol!r}.")
self.store: Any | None = None
self.local_seg: str | None = None
self._use_fabric_mem = os.getenv("ASCEND_ENABLE_USE_FABRIC_MEM", "0") == "1"
self._lazy_init = lazy_init and self._use_fabric_mem
self._contribute_memory = contribute_memory
self._store_initialized = False
self._store_init_lock = threading.Lock()
if not self._lazy_init:
self.store = self._setup_store()
self._store_initialized = True
def ensure_initialized(self):
if self._store_initialized:
return
with self._store_init_lock:
if self._store_initialized:
return
logger.info("Initializing Mooncake store. metadata_server=%s", self.config.metadata_server)
self.store = self._setup_store()
self._store_initialized = True
def _setup_store(self):
try:
from mooncake.store import MooncakeDistributedStore # type: ignore
except ImportError as e:
raise ImportError(
"Please install mooncake by following the instructions at "
"https://github.com/kvcache-ai/Mooncake/blob/main/doc/en/build.md " # noqa: E501
"to run vLLM with MooncakeConnector."
) from e
store = MooncakeDistributedStore()
local_hostname = get_ip()
ssd_kwargs = _ssd_setup_kwargs(self.config)
# Scheduler-only clients (contribute_memory=False) do not contribute
# KV cache memory and therefore do not need SSD offload. Passing
# enable_ssd_offload=True for them would cause Mooncake to register
# an extra active client on the master, inflating both the client
# count and the reported SSD storage usage.
if ssd_kwargs and not self._contribute_memory:
ssd_kwargs = {}
# Each rank that contributes memory to the pool uses its own SSD
# directory to avoid bucket file collisions. Key by the globally unique
# rank so that DP/TP/PP/CP replicas never share a directory (dense and
# MoE alike); only ranks that contribute memory need an offload dir.
if ssd_kwargs and ssd_kwargs.get("ssd_offload_path"):
global_rank = get_global_rank(self.parallel_config)
rank_path = os.path.join(str(ssd_kwargs["ssd_offload_path"]), f"rank_{global_rank}")
try:
os.makedirs(rank_path, exist_ok=True)
except OSError as e:
raise RuntimeError(f"Failed to create per-rank SSD offload directory: {rank_path!r} ({e})")
ssd_kwargs["ssd_offload_path"] = rank_path
# ASCEND_ENABLE_USE_FABRIC_MEM: Enable unified memory address direct transmission scheme
# and only can be used for 800 I/T A3 series.
# Required supporting hardware versions are as follows:
if not self._use_fabric_mem:
transfer_engine = global_te.get_transfer_engine(local_hostname, device_name=None)
self.local_seg = local_hostname + ":" + str(transfer_engine.get_rpc_port())
ret = store.setup(
local_hostname=self.local_seg,
metadata_server=self.config.metadata_server,
global_segment_size=self.config.global_segment_size if self._contribute_memory else 0,
local_buffer_size=self.config.local_buffer_size if self._contribute_memory else 0,
protocol=self.config.protocol,
rdma_devices=self.config.device_name,
master_server_addr=self.config.master_server_address,
engine=transfer_engine.get_engine(),
**ssd_kwargs,
)
else:
self.local_seg = local_hostname
ret = store.setup(
local_hostname=self.local_seg,
metadata_server=self.config.metadata_server,
global_segment_size=self.config.global_segment_size if self._contribute_memory else 0,
local_buffer_size=0,
protocol=self.config.protocol,
rdma_devices=self.config.device_name,
master_server_addr=self.config.master_server_address,
**ssd_kwargs,
)
if ret != 0:
msg = "Initialize mooncake failed."
logger.error(
"Initialize mooncake failed. ret=%d, metadata_server=%s. Check mooncake config and network.",
ret,
self.config.metadata_server,
)
raise RuntimeError(msg)
if ssd_kwargs:
logger.info(
"Mooncake SSD offload enabled (Mode A): path=%s",
self.config.ssd_offload_path,
)
return store
@classmethod
def create_scheduler_client(cls, parallel_config: ParallelConfig):
torch.npu.set_device(0)
return cls(parallel_config, contribute_memory=False)
def set_device(self):
local_rank = get_world_group().local_rank
device = torch.device(f"npu:{local_rank}")
torch.npu.set_device(device)
def register_buffer(self, ptrs: list[int], lengths: list[int]):
if not self._use_fabric_mem:
local_hostname = get_ip()
global_te.get_transfer_engine(local_hostname, device_name=None)
global_te.register_buffer(ptrs, lengths)
def exists(self, keys: list[str]) -> list[int]:
if self._lazy_init and not self._store_initialized:
logger.debug(
"MooncakeBackend.exists called before store initialization; treating %d keys as missing.",
len(keys),
)
return [0] * len(keys)
assert self.store is not None
return self.store.batch_is_exist(keys)
def put(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]):
self.ensure_initialized()
assert self.store is not None
try:
config = ReplicateConfig()
if self.config.preferred_segment:
config.preferred_segment = self.local_seg
config.prefer_alloc_in_same_node = self.config.prefer_alloc_in_same_node
res = self.store.batch_put_from_multi_buffers(keys, addrs, sizes, config)
failed_codes = [int(value) for value in res if value < 0]
failed_count = len(failed_codes)
if failed_count:
error_codes = sorted(set(failed_codes))
logger.error(
"Failed to put %d keys out of %d. error_codes=%s. Check memory and store capacity.",
failed_count,
len(keys),
error_codes,
)
logger.debug("Failed to put key details. keys=%s, result=%s", keys, res)
if self._lazy_init:
logger.warning("First DSV4(compress) request failure is expected. This is normal behavior.")
except Exception as e:
logger.error(
"Failed to put %d keys out of %d. type=%s, error=%s. Check store state and memory.",
len(keys),
len(keys),
type(e).__name__,
e,
)
logger.debug("Failed to put key details. keys=%s", keys)
if self._lazy_init:
logger.warning("First DSV4(compress) request failure is expected. This is normal behavior.")
def get(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]):
if self._lazy_init and not self._store_initialized:
logger.error(
"Failed to get %d keys out of %d. Store is not initialized; "
"call put() first to trigger initialization.",
len(keys),
len(keys),
)
logger.debug("Failed to get key details. keys=%s", keys)
return
assert self.store is not None
logger.debug(
"MooncakeBackend.get enter keys=%d sample_keys=%s",
len(keys),
keys[:3],
)
try:
res = self.store.batch_get_into_multi_buffers(keys, addrs, sizes)
res_list = list(res)
failed_codes = [int(value) for value in res_list if value < 0]
failed_count = len(failed_codes)
error_codes = sorted(set(failed_codes))
if failed_count:
logger.error(
"Failed to get %d keys out of %d. error_codes=%s. Check key existence and memory state.",
failed_count,
len(keys),
error_codes,
)
logger.debug("Failed to get key details. keys=%s, result=%s", keys, res_list)
for i, value in enumerate(res_list):
if value > 0:
res_list[i] = 0
return res_list
except Exception as e:
logger.error(
"Failed to get %d keys out of %d. type=%s, error=%s. Check store state and network.",
len(keys),
len(keys),
type(e).__name__,
e,
)
logger.debug("Failed to get key details. keys=%s", keys)
return None
@dataclass
class MooncakeStoreConfig:
metadata_server: str
global_segment_size: int | str
local_buffer_size: int
protocol: str
device_name: str
master_server_address: str
preferred_segment: bool
prefer_alloc_in_same_node: bool
enable_ssd_offload: bool = False
ssd_offload_path: str = ""
def __post_init__(self) -> None:
if not self.enable_ssd_offload:
return
if not self.ssd_offload_path:
raise ValueError(
"enable_ssd_offload is true but ssd_offload_path is empty. Set ssd_offload_path in mooncake.json."
)
if not os.path.isabs(self.ssd_offload_path):
raise ValueError(f"ssd_offload_path must be an absolute path, got: {self.ssd_offload_path!r}")
@staticmethod
def from_file(file_path: str) -> "MooncakeStoreConfig":
with open(file_path) as file:
config = json.load(file)
master_server_address = os.getenv("MOONCAKE_MASTER", None)
global_segment_size_env = os.getenv("MOONCAKE_GLOBAL_SEGMENT_SIZE", None)
return MooncakeStoreConfig(
metadata_server=config.get("metadata_server"),
global_segment_size=_parse_global_segment_size(
global_segment_size_env
if global_segment_size_env is not None
else config.get("global_segment_size", DEFAULT_GLOBAL_SEGMENT_SIZE)
),
local_buffer_size=_parse_global_segment_size(config.get("local_buffer_size", DEFAULT_LOCAL_BUFFER_SIZE)),
protocol=config.get("protocol", "ascend"),
device_name=config.get("device_name", ""),
master_server_address=master_server_address
if master_server_address is not None
else config.get("master_server_address"),
preferred_segment=config.get("preferred_segment", False),
prefer_alloc_in_same_node=config.get("prefer_alloc_in_same_node", True),
enable_ssd_offload=bool(config.get("enable_ssd_offload", False)),
ssd_offload_path=config.get("ssd_offload_path", ""),
)
@staticmethod
def load_from_env() -> "MooncakeStoreConfig":
config_path = os.getenv("MOONCAKE_CONFIG_PATH")
if not config_path:
raise ValueError("The environment variable 'MOONCAKE_CONFIG_PATH' is not set.")
return MooncakeStoreConfig.from_file(config_path)
def _parse_global_segment_size(value) -> int:
"""
Parse storage size strings with support for units: GB, MB, KB, B
Args:
value: Input value (int, str, or other convertible types)
Returns:
int: Size in bytes
Raises:
ValueError: For invalid format, missing number, or negative values
TypeError: For unsupported input types
"""
if isinstance(value, int):
return value
elif not isinstance(value, str):
try:
return int(value)
except (TypeError, ValueError) as e:
raise TypeError(f"Unsupported type for global_segment_size: {type(value)}") from e
cleaned_input = value.strip().lower()
if not cleaned_input:
raise ValueError("global segment size cannot be empty.")
UNIT_MULTIPLIERS = {
"gb": 1024**3, # 1 GB = 1024^3 bytes
"mb": 1024**2, # 1 MB = 1024^2 bytes
"kb": 1024, # 1 KB = 1024 bytes
"b": 1, # 1 B = 1 byte
}
pattern = r"^\s*([\d.]+)\s*(gb|mb|kb|b)?\s*$"
match = re.match(pattern, cleaned_input)
if not match:
raise ValueError(f"Invalid format: '{value}'")
number_str = match.group(1)
unit = match.group(2) or "b"
multiplier = UNIT_MULTIPLIERS[unit]
return _convert_to_bytes(number_str, multiplier, value)
def _convert_to_bytes(number_str: str, multiplier: int, original_input: str) -> int:
"""
Convert numeric string to byte count
Args:
number_str: Numeric portion of input
multiplier: Unit conversion factor
original_input: Original input string (for error messages)
Returns:
int: Byte count
Raises:
ValueError: For invalid numbers or negative results
"""
try:
numeric_value = float(number_str)
except ValueError:
raise ValueError(f"Invalid numeric value '{number_str}' in: '{original_input}'")
# Calculate byte count
try:
byte_count = int(numeric_value * multiplier)
except OverflowError:
raise ValueError(f"Storage size too large: '{original_input}'")
return byte_count

View File

@@ -0,0 +1,238 @@
import hashlib
import os
from dataclasses import dataclass
from typing import Any
import regex as re
import torch
from vllm.config import ParallelConfig
from vllm.distributed.parallel_state import get_world_group
from vllm.logger import logger
from vllm.utils.network_utils import split_host_port
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.backend import Backend
from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type
def _iter_slices(total: int, batch_size: int):
for start in range(0, total, batch_size):
end = min(start + batch_size, total)
yield start, end
@dataclass
class YuanrongConfig:
worker_addr: str
enable_exclusive_connection: bool
enable_remote_h2d: bool
@staticmethod
def load_from_env() -> "YuanrongConfig":
worker_addr = os.getenv("DS_WORKER_ADDR")
if not worker_addr:
raise ValueError("Environment variable DS_WORKER_ADDR is required, expected format '<host>:<port>'.")
return YuanrongConfig(
worker_addr=worker_addr,
enable_exclusive_connection=bool(int(os.getenv("DS_ENABLE_EXCLUSIVE_CONNECTION", "0"))),
enable_remote_h2d=bool(int(os.getenv("DS_ENABLE_REMOTE_H2D", "0"))),
)
class YuanrongHelper:
_DS_KEY_MAX_LEN = 1024
_DS_KEY_ALLOWED_PATTERN = re.compile(r"^[a-zA-Z0-9\-_!@#%\^\*\(\)\+\=\:;]+$")
_DS_KEY_INVALID_CHAR_PATTERN = re.compile(r"[^a-zA-Z0-9\-_!@#%\^\*\(\)\+\=\:;]")
_DS_KEY_HASH_SUFFIX_LEN = 16
def __init__(self, blob_cls, blob_list_cls):
self._blob_cls = blob_cls
self._blob_list_cls = blob_list_cls
self._device_id: int | None = None
def normalize_keys(self, keys: list[str]) -> list[str]:
normalized: list[str] = []
for key in keys:
if len(key) <= self._DS_KEY_MAX_LEN and self._DS_KEY_ALLOWED_PATTERN.match(key):
normalized.append(key)
continue
sanitized = self._DS_KEY_INVALID_CHAR_PATTERN.sub("_", key)
hash_digest = hashlib.sha256(key.encode("utf-8")).hexdigest()
suffix = f"__{hash_digest[: self._DS_KEY_HASH_SUFFIX_LEN]}"
max_prefix_len = self._DS_KEY_MAX_LEN - len(suffix)
normalized.append(sanitized[:max_prefix_len] + suffix)
return normalized
def make_blob_lists(self, addrs_list: list[list[int]], sizes_list: list[list[int]]) -> list[Any]:
total = len(addrs_list)
if total != len(sizes_list):
raise ValueError("Address list and size list length mismatch.")
device_id = self._device_id
if device_id is None:
logger.error("Device id is not set. Check device initialization and configuration.")
raise RuntimeError("Yuanrong backend device id is not initialized.")
blob_lists: list[Any] = []
for addrs, sizes in zip(addrs_list, sizes_list):
if len(addrs) != len(sizes):
raise ValueError("Address list and size list length mismatch.")
blobs = [
self._blob_cls(addr, size) # type: ignore[misc]
for addr, size in zip(addrs, sizes)
]
blob_lists.append(
self._blob_list_cls(device_id, blobs) # type: ignore[misc]
)
return blob_lists
class YuanrongBackend(Backend):
_DS_MAX_BATCH_KEYS = 10000
def __init__(self, parallel_config: ParallelConfig):
try:
from yr.datasystem.hetero_client import Blob, DeviceBlobList, HeteroClient # type: ignore[import-not-found]
from yr.datasystem.kv_client import SetParam # type: ignore[import-not-found]
from yr.datasystem.object_client import WriteMode # type: ignore[import-not-found]
except ImportError as exc:
raise ImportError("Please install openyuanrong-datasystem to use the yuanrong backend.") from exc
self._helper = YuanrongHelper(Blob, DeviceBlobList)
self._ds_set_param = SetParam()
self._ds_set_param.write_mode = WriteMode.NONE_L2_CACHE_EVICT
self.config = YuanrongConfig.load_from_env()
try:
host, port = split_host_port(self.config.worker_addr)
except Exception as exc:
raise ValueError(f"Invalid DS_WORKER_ADDR '{self.config.worker_addr}', expected '<host>:<port>'.") from exc
self._hetero_client = HeteroClient(
host,
int(port),
enable_exclusive_connection=self.config.enable_exclusive_connection,
enable_remote_h2d=self.config.enable_remote_h2d,
)
self._hetero_client.init()
self._is_a2 = get_ascend_device_type() in {AscendDeviceType.A2}
self._registered_buffers: tuple[list[int], list[int]] | None = None
self._buffers_registered = False
def _ensure_device_ready(self):
if self._helper._device_id is None:
self.set_device()
def set_device(self):
local_rank = get_world_group().local_rank
device = torch.device(f"npu:{local_rank}")
torch.npu.set_device(device)
self._helper._device_id = int(torch.npu.current_device())
def register_buffer(self, ptrs: list[int], lengths: list[int]):
self._registered_buffers = (list(ptrs), list(lengths))
self._register_buffers_if_needed()
def _register_buffers_if_needed(self):
if self._is_a2:
return
if not self.config.enable_remote_h2d:
return
if self._registered_buffers is None or self._buffers_registered:
return
ptrs, lengths = self._registered_buffers
self._hetero_client.pre_register_device_memory(ptrs, lengths) # type: ignore[union-attr]
self._buffers_registered = True
def exists(self, keys: list[str]) -> list[int]:
if len(keys) == 0:
return []
try:
keys = self._helper.normalize_keys(keys)
if len(keys) <= self._DS_MAX_BATCH_KEYS:
exists = self._hetero_client.exist(keys) # type: ignore[union-attr]
return [1 if value else 0 for value in exists]
results: list[int] = []
for start, end in _iter_slices(len(keys), self._DS_MAX_BATCH_KEYS):
exists = self._hetero_client.exist(keys[start:end]) # type: ignore[union-attr]
results.extend(1 if value else 0 for value in exists)
return results
except Exception as exc:
logger.error(
"Failed to check keys. keys_count=%d, type=%s, error=%s. Check network and yuanrong service.",
len(keys),
type(exc).__name__,
exc,
)
return [0] * len(keys)
def get(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]) -> list[int] | None:
if len(keys) == 0:
return []
failed_keys_for_log = keys
try:
self._ensure_device_ready()
keys = self._helper.normalize_keys(keys)
failed_keys_for_log = keys
blob_lists = self._helper.make_blob_lists(addrs, sizes)
failed_keys: list[str] = []
if len(keys) <= self._DS_MAX_BATCH_KEYS:
failed_keys = self._hetero_client.mget_h2d( # type: ignore[union-attr]
keys, blob_lists, 0
)
else:
for start, end in _iter_slices(len(keys), self._DS_MAX_BATCH_KEYS):
failed_keys_for_log = keys[start:end]
failed_keys.extend(
self._hetero_client.mget_h2d( # type: ignore[union-attr]
keys[start:end], blob_lists[start:end], 0
)
)
if failed_keys:
logger.error(
"Failed to get %d keys out of %d. Check key existence and memory state.",
len(failed_keys),
len(keys),
)
logger.debug("Failed to get key details. failed_keys=%s", failed_keys)
failed_set = set(failed_keys)
return [1 if k in failed_set else 0 for k in keys]
except Exception as exc:
logger.error(
"Failed to get %d keys out of %d. type=%s, error=%s. Check network and yuanrong service.",
len(failed_keys_for_log),
len(keys),
type(exc).__name__,
exc,
)
logger.debug("Failed to get key details. keys=%s", failed_keys_for_log)
return None
def put(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]):
if len(keys) == 0:
return
failed_keys_for_log = keys
try:
self._ensure_device_ready()
keys = self._helper.normalize_keys(keys)
failed_keys_for_log = keys
blob_lists = self._helper.make_blob_lists(addrs, sizes)
if len(keys) <= self._DS_MAX_BATCH_KEYS:
self._hetero_client.mset_d2h( # type: ignore[union-attr]
keys, blob_lists, self._ds_set_param
)
else:
for start, end in _iter_slices(len(keys), self._DS_MAX_BATCH_KEYS):
failed_keys_for_log = keys[start:end]
self._hetero_client.mset_d2h( # type: ignore[union-attr]
keys[start:end], blob_lists[start:end], self._ds_set_param
)
except Exception as exc:
logger.error(
"Failed to put %d keys out of %d. type=%s, error=%s. Check network and yuanrong service.",
len(failed_keys_for_log),
len(keys),
type(exc).__name__,
exc,
)
logger.debug("Failed to put key details. keys=%s", failed_keys_for_log)

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,450 @@
from __future__ import annotations
from dataclasses import replace
from importlib import import_module
from typing import Any, cast
from vllm.logger import logger
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_utils import BlockHash, BlockHashList, KVCacheBlock
from vllm.v1.core.single_type_kv_cache_manager import SingleTypeKVCacheManager
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheGroupSpec,
KVCacheSpec,
UniformTypeKVCacheSpecs,
)
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
GroupedBlockHashCache,
block_hash_to_bytes,
get_block_hashes,
)
_CACHE_MISSING = object()
_MANAGER_CLASS_CACHE_ATTR = "_manager_class_cache"
class ExternalCachedBlockPool:
"""Duck-typed BlockPool backed by external AscendStore key existence."""
def __init__(self, exists: set[tuple[int, bytes]] | None = None) -> None:
# exists=None is used for load/store masks where hit length has already
# been decided and each manager only needs to apply its own reachability.
self._exists = exists
self.null_block = KVCacheBlock(block_id=0)
self._present_block = KVCacheBlock(block_id=1)
def get_cached_block(
self,
block_hash: BlockHash,
group_ids: list[int],
) -> list[KVCacheBlock] | None:
if self._exists is None:
return [self._present_block] * len(group_ids)
h = block_hash_to_bytes(block_hash)
if all((group_id, h) in self._exists for group_id in group_ids):
return [self._present_block] * len(group_ids)
return None
class AscendStoreCoordinator:
"""Hybrid cache-hit/mask coordinator for AscendStore external KV Pool.
This mirrors vLLM MooncakeStoreCoordinator but uses AscendStore's external
key granularity. For DSV4 compressed groups, keys are generated over the
raw-token span ``group_block_size * compress_ratio`` while transfer
addresses remain in cache-domain blocks.
"""
def __init__(
self,
kv_cache_groups: list[KVCacheGroupSpec],
scheduler_block_size: int,
hash_block_size: int,
group_block_sizes: list[int],
group_cache_families: list[str],
use_eagle: bool = False,
retention_interval: int | None = None,
) -> None:
assert len(kv_cache_groups) == len(group_block_sizes)
assert len(kv_cache_groups) == len(group_cache_families)
assert scheduler_block_size % hash_block_size == 0, (
f"scheduler_block_size ({scheduler_block_size}) must be a multiple of hash_block_size ({hash_block_size})"
)
self.kv_cache_groups = kv_cache_groups
self.hash_block_size = hash_block_size
self.lcm_block_size = scheduler_block_size
self.use_eagle = use_eagle
self.retention_interval = retention_interval
self.group_block_sizes = group_block_sizes
self.group_cache_families = group_cache_families
self.group_effective_block_sizes = [
_cache_family_granularity(block_size, family)
for block_size, family in zip(group_block_sizes, group_cache_families, strict=True)
]
for effective_block_size in self.group_effective_block_sizes:
assert effective_block_size % hash_block_size == 0, "block_size must be divisible by hash_block_size"
assert scheduler_block_size % effective_block_size == 0, (
"scheduler_block_size must be a multiple of each group's effective block_size"
)
self.eagle_group_ids = {i for i, group in enumerate(kv_cache_groups) if group.is_eagle_group}
if use_eagle and not self.eagle_group_ids:
self.eagle_group_ids = set(range(len(kv_cache_groups)))
self._verify_and_split_kv_cache_groups()
def _verify_and_split_kv_cache_groups(self) -> None:
attention_groups: list[tuple[KVCacheSpec, list[int], type[SingleTypeKVCacheManager]]] = []
self.group_effective_specs: list[KVCacheSpec] = []
for group_id, group in enumerate(self.kv_cache_groups):
spec = _unwrap_spec(group.kv_cache_spec)
effective_spec = _copy_spec_with_block_size(spec, self.group_effective_block_sizes[group_id])
if (
not _uses_reachable_mask(self.group_cache_families[group_id])
and getattr(effective_spec, "compress_ratio", 1) > 1
):
# The cache family already folds the compression ratio into
# the external key granularity. Avoid applying it again inside
# CompressAttentionManager.find_longest_cache_hit().
effective_spec = replace(effective_spec, compress_ratio=1)
self.group_effective_specs.append(effective_spec)
manager_cls = _get_manager_class(spec)
for existing_spec, group_ids, existing_cls in attention_groups:
if existing_spec == effective_spec:
assert manager_cls is existing_cls, "Expected same manager class for identical KV cache specs."
group_ids.append(group_id)
break
else:
attention_groups.append((effective_spec, [group_id], manager_cls))
self.attention_groups = sorted(
attention_groups,
key=lambda item: not isinstance(item[0], FullAttentionSpec),
)
self.eagle_attn_group_indices: set[int] = {
index
for index, (_, group_ids, _) in enumerate(self.attention_groups)
if any(group_id in self.eagle_group_ids for group_id in group_ids)
}
if self.use_eagle and not self.eagle_attn_group_indices:
self.eagle_attn_group_indices = set(range(len(self.attention_groups)))
self.eagle_reachable_group_ids: set[int] = {
group_id for index in self.eagle_attn_group_indices for group_id in self.attention_groups[index][1]
}
def find_longest_cache_hit(
self,
block_hashes: list[BlockHash],
max_length: int,
cached_block_pool: ExternalCachedBlockPool,
*,
apply_eagle: bool = True,
grouped_hash_cache: GroupedBlockHashCache | None = None,
) -> tuple[tuple[list[bool], ...], int]:
blocks_per_group, hit_length = self._find_hit_blocks(
block_hashes,
max_length,
cached_block_pool,
apply_eagle=apply_eagle,
grouped_hash_cache=grouped_hash_cache,
)
masks = tuple([block is not cached_block_pool.null_block for block in blocks] for blocks in blocks_per_group)
return masks, hit_length
def load_mask(
self,
block_hashes: list[BlockHash],
token_len: int,
grouped_hash_cache: GroupedBlockHashCache | None = None,
) -> tuple[list[bool], ...]:
masks, _ = self.find_longest_cache_hit(
block_hashes,
token_len,
ExternalCachedBlockPool(),
apply_eagle=False,
grouped_hash_cache=grouped_hash_cache,
)
return tuple(
[True] * _num_chunks(token_len, self.group_effective_block_sizes[group_id])
if not _uses_reachable_mask(self.group_cache_families[group_id])
else mask
for group_id, mask in enumerate(masks)
)
def _reachable_masks(
self,
aligned_token_len: int,
retention_interval: int | None,
num_prompt_tokens: int | None,
) -> list[tuple[int, list[bool] | None]]:
assert aligned_token_len % self.lcm_block_size == 0, (
f"aligned_token_len ({aligned_token_len}) must be a multiple of lcm_block_size ({self.lcm_block_size})"
)
masks: list[tuple[int, list[bool] | None]] = []
for group_id, spec in enumerate(self.group_effective_specs):
num_chunks = aligned_token_len // self.group_effective_block_sizes[group_id]
if not _uses_reachable_mask(self.group_cache_families[group_id]):
masks.append((num_chunks, None))
continue
manager_cls = _get_manager_class(_unwrap_spec(self.kv_cache_groups[group_id].kv_cache_spec))
mask = _reachable_block_mask(
manager_cls,
start_block=0,
end_block=num_chunks,
alignment_tokens=self.lcm_block_size,
kv_cache_spec=spec,
use_eagle=group_id in self.eagle_reachable_group_ids,
retention_interval=retention_interval,
num_prompt_tokens=num_prompt_tokens,
)
masks.append((num_chunks, mask))
return masks
def store_mask(
self,
aligned_token_len: int,
num_prompt_tokens: int | None = None,
) -> tuple[list[bool], ...]:
masks = self._reachable_masks(aligned_token_len, self.retention_interval, num_prompt_tokens)
return tuple([True] * num_chunks if mask is None else mask for num_chunks, mask in masks)
def lookup_mask(
self,
aligned_token_len: int,
) -> tuple[list[bool] | None, ...]:
masks = self._reachable_masks(aligned_token_len, None, None)
for num_chunks, mask in masks:
if mask is not None:
assert len(mask) == num_chunks
return tuple(None if mask is None or all(mask) else mask for _, mask in masks)
def block_hashes_for_spec(
self,
block_hashes: list[BlockHash],
spec: KVCacheSpec,
grouped_hash_cache: GroupedBlockHashCache | None = None,
) -> BlockHashList:
if spec.block_size == self.hash_block_size:
return block_hashes
return cast(
BlockHashList,
get_block_hashes(
block_hashes,
spec.block_size,
self.hash_block_size,
grouped_hash_cache=grouped_hash_cache,
),
)
def _find_hit_blocks(
self,
block_hashes: list[BlockHash],
max_length: int,
cached_block_pool: ExternalCachedBlockPool,
*,
apply_eagle: bool = True,
grouped_hash_cache: GroupedBlockHashCache | None = None,
) -> tuple[tuple[list[KVCacheBlock], ...], int]:
eagle_indices = self.eagle_attn_group_indices if apply_eagle else set()
if len(self.attention_groups) == 1:
spec, group_ids, manager_cls = self.attention_groups[0]
hashes = self.block_hashes_for_spec(block_hashes, spec, grouped_hash_cache)
hit_blocks = _find_longest_cache_hit(
manager_cls,
block_hashes=hashes,
max_length=max_length,
kv_cache_group_ids=group_ids,
block_pool=cast(BlockPool, cached_block_pool),
kv_cache_spec=spec,
drop_eagle_block=0 in eagle_indices,
alignment_tokens=spec.block_size,
)
blocks_by_group: list[list[KVCacheBlock]] = [[] for _ in range(len(self.kv_cache_groups))]
for group_id, blocks in zip(group_ids, hit_blocks, strict=True):
blocks_by_group[group_id] = blocks
return tuple(blocks_by_group), len(hit_blocks[0]) * spec.block_size
hit_length = max_length
hit_blocks_by_group: list[list[KVCacheBlock] | None] = [None] * len(self.kv_cache_groups)
is_simple_hybrid = len(self.attention_groups) == 2 and isinstance(
self.attention_groups[0][0], FullAttentionSpec
)
eagle_verified: set[int] = set()
while True:
curr_hit_length = hit_length
for index, (spec, group_ids, manager_cls) in enumerate(self.attention_groups):
cached = hit_blocks_by_group[group_ids[0]]
if isinstance(spec, FullAttentionSpec) and cached is not None:
curr_hit_length = curr_hit_length // spec.block_size * spec.block_size
continue
drop_eagle_block = index in eagle_indices and index not in eagle_verified
max_group_length = curr_hit_length
if drop_eagle_block:
max_group_length = min(curr_hit_length + spec.block_size, max_length)
hashes = self.block_hashes_for_spec(block_hashes, spec, grouped_hash_cache)
hit_blocks = _find_longest_cache_hit(
manager_cls,
block_hashes=hashes,
max_length=max_group_length,
kv_cache_group_ids=group_ids,
block_pool=cast(BlockPool, cached_block_pool),
kv_cache_spec=spec,
drop_eagle_block=drop_eagle_block,
alignment_tokens=self.lcm_block_size,
)
new_hit_length = len(hit_blocks[0]) * spec.block_size
if drop_eagle_block:
eagle_verified.add(index)
elif new_hit_length < curr_hit_length:
eagle_verified.clear()
curr_hit_length = new_hit_length
for group_id, blocks in zip(group_ids, hit_blocks, strict=True):
hit_blocks_by_group[group_id] = blocks
if curr_hit_length >= hit_length:
break
hit_length = curr_hit_length
if is_simple_hybrid:
break
spec0, group_ids0, _ = self.attention_groups[0]
if isinstance(spec0, FullAttentionSpec):
num_blocks = hit_length // spec0.block_size
for group_id in group_ids0:
full_blocks = hit_blocks_by_group[group_id]
assert full_blocks is not None
del full_blocks[num_blocks:]
return (
tuple(blocks if blocks is not None else [] for blocks in hit_blocks_by_group),
hit_length,
)
def _unwrap_spec(spec: KVCacheSpec) -> KVCacheSpec:
if isinstance(spec, UniformTypeKVCacheSpecs):
return next(iter(spec.kv_cache_specs.values()))
return spec
def _copy_spec_with_block_size(spec: KVCacheSpec, block_size: int) -> KVCacheSpec:
if spec.block_size == block_size:
return spec
copy_with_new_block_size = getattr(spec, "copy_with_new_block_size", None)
if copy_with_new_block_size is not None:
return copy_with_new_block_size(block_size)
return replace(spec, block_size=block_size)
def _get_manager_class_cache() -> dict[str, Any]:
cache = getattr(_get_manager_class, _MANAGER_CLASS_CACHE_ATTR, None)
if not isinstance(cache, dict):
cache = {}
setattr(_get_manager_class, _MANAGER_CLASS_CACHE_ATTR, cache)
return cast(dict[str, Any], cache)
def _get_manager_class(spec: KVCacheSpec) -> type[SingleTypeKVCacheManager]:
cache = _get_manager_class_cache()
compress_ratio = getattr(spec, "compress_ratio", None)
if compress_ratio is not None and compress_ratio > 1:
compress_manager = cache.get("compress_manager", _CACHE_MISSING)
if compress_manager is _CACHE_MISSING:
try:
from vllm_ascend.core.single_type_kv_cache_manager import CompressAttentionManager
except ImportError:
compress_manager = None
else:
compress_manager = CompressAttentionManager
cache["compress_manager"] = compress_manager
if compress_manager is not None:
return cast(type[SingleTypeKVCacheManager], compress_manager)
registry = cache.get("registry", _CACHE_MISSING)
if registry is _CACHE_MISSING:
try:
registry_module = import_module("vllm.v1.kv_cache_spec_registry")
registry = getattr(registry_module, "KVCacheSpecRegistry", None)
except ImportError:
registry = None
cache["registry"] = registry
if registry is not None:
manager_cls = registry.get_manager_class(spec)
if manager_cls is not None:
return manager_cls
spec_manager_map = cache.get("spec_manager_map", _CACHE_MISSING)
if spec_manager_map is _CACHE_MISSING:
try:
manager_module = import_module("vllm.v1.core.single_type_kv_cache_manager")
spec_manager_map = vars(manager_module)["spec_manager_map"]
except Exception as exc:
raise AssertionError(f"No manager registered for KVCacheSpec {type(spec)}") from exc
cache["spec_manager_map"] = spec_manager_map
try:
manager_cls = spec_manager_map[type(spec)]
except Exception as exc:
raise AssertionError(f"No manager registered for KVCacheSpec {type(spec)}") from exc
return manager_cls
def _find_longest_cache_hit(
manager_cls: type[SingleTypeKVCacheManager],
**kwargs: Any,
) -> tuple[list[KVCacheBlock], ...]:
try:
return manager_cls.find_longest_cache_hit(**kwargs)
except TypeError as exc:
if "drop_eagle_block" not in str(exc):
raise
kwargs["use_eagle"] = kwargs.pop("drop_eagle_block")
return manager_cls.find_longest_cache_hit(**kwargs)
def _reachable_block_mask(
manager_cls: type[SingleTypeKVCacheManager],
**kwargs: Any,
) -> list[bool] | None:
reachable_block_mask = getattr(manager_cls, "reachable_block_mask", None)
if reachable_block_mask is None:
return None
try:
return reachable_block_mask(**kwargs)
except TypeError as exc:
if "retention_interval" not in str(exc) and "num_prompt_tokens" not in str(exc):
logger.debug("KV cache manager does not support reachable_block_mask kwargs: %s", exc)
return reachable_block_mask(
start_block=kwargs["start_block"],
end_block=kwargs["end_block"],
alignment_tokens=kwargs["alignment_tokens"],
kv_cache_spec=kwargs["kv_cache_spec"],
use_eagle=kwargs["use_eagle"],
)
kwargs.pop("retention_interval", None)
kwargs.pop("num_prompt_tokens", None)
return reachable_block_mask(**kwargs)
def _cache_family_granularity(block_size: int, cache_family: str | None) -> int:
if not cache_family or not cache_family.startswith("c"):
return block_size
ratio = cache_family[1:]
return block_size * int(ratio) if ratio.isdigit() else block_size
def _uses_reachable_mask(cache_family: str | None) -> bool:
return cache_family in (None, "default", "c1")
def _num_chunks(token_len: int, block_size: int) -> int:
return (token_len + block_size - 1) // block_size

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,199 @@
import time
from collections import defaultdict
from vllm.logger import logger
from vllm.utils.hashing import sha256
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_utils import BlockHash, KVCacheBlock
from vllm.v1.kv_cache_interface import KVCacheSpec
from vllm.v1.metrics.stats import CachingMetrics, PrefixCacheStats
from vllm.v1.request import Request
from vllm_ascend.core.single_type_kv_cache_manager import get_manager_for_kv_cache_spec
class CPUCacheStats:
def __init__(self, enable_prefix_caching: bool, log_stats: bool = False):
self.enable_prefix_caching = enable_prefix_caching
self.log_stats = log_stats
self.prefix_cache_stats = PrefixCacheStats() if log_stats else None
self.cpu_prefix_cache_metrics = CachingMetrics()
self.time_sec = int(time.time())
def log(self):
current_time_sec = int(time.time())
# Log the prefix cache hit rate every 10 seconds.
if current_time_sec - self.time_sec >= 10:
self.time_sec = current_time_sec
logger.info("CPU Prefix cache hit rate: %.1f%%", self.cpu_prefix_cache_metrics.hit_rate * 100)
def make_prefix_cache_stats(self) -> PrefixCacheStats | None:
"""Get (and reset) the prefix cache stats.
Returns:
The current prefix caching stats, or None if logging is disabled.
"""
if not self.log_stats:
return None
stats = self.prefix_cache_stats
self.prefix_cache_stats = PrefixCacheStats()
return stats
def update(self, num_tokens, num_computed_tokens):
# Note the function is called by scheduler
if self.log_stats and self.enable_prefix_caching:
assert self.prefix_cache_stats is not None
self.prefix_cache_stats.requests += 1
self.prefix_cache_stats.queries += num_tokens
self.prefix_cache_stats.hits += num_computed_tokens
def set_cache_stats(self, num_tokens, num_computed_tokens):
assert self.prefix_cache_stats is not None
self.prefix_cache_stats.hits = num_computed_tokens
self.prefix_cache_stats.queries = num_tokens
self.prefix_cache_stats.requests = 1
class CPUKVCacheManager:
def __init__(
self,
kv_cache_spec: KVCacheSpec,
num_cpu_blocks: int,
caching_hash_algo: str = "builtin",
use_eagle: bool = False,
enable_kv_cache_events: bool = False,
) -> None:
self.block_size = kv_cache_spec.block_size
self.num_cpu_blocks = num_cpu_blocks
self.caching_hash_fn = sha256 if caching_hash_algo == "sha256" else hash
self.use_eagle = use_eagle
self.block_pool = BlockPool(self.num_cpu_blocks, True, self.block_size, enable_kv_cache_events)
max_model_len = self.num_cpu_blocks * self.block_size
manager_kwargs = dict(
kv_cache_spec=kv_cache_spec,
block_pool=self.block_pool,
enable_caching=True,
kv_cache_group_id=0,
max_num_batched_tokens=max_model_len,
max_model_len=max_model_len,
)
manager_kwargs["scheduler_block_size"] = kv_cache_spec.block_size
self.single_type_manager = get_manager_for_kv_cache_spec(**manager_kwargs)
# Record kv block hashes, avoid redundant computation.
self.req_to_block_hashes: defaultdict[str, list[BlockHash]] = defaultdict(list)
# Record blocks touched in get_matched_num_and_touch().
self.req_to_computed_blocks: defaultdict[str, list[KVCacheBlock]] = defaultdict(list)
# Record the request that failed to allocate.
self.req_failed_to_allocate: defaultdict[str, bool] = defaultdict(bool)
self.req_to_num_tokens: defaultdict[str, int] = defaultdict(int)
self.cpu_cache_stats = CPUCacheStats(enable_prefix_caching=True, log_stats=True)
# Record request that will be free after finish sending
self.req_to_free: defaultdict[str, Request] = defaultdict(Request)
def get_matched_num_and_touch(self, request: Request) -> tuple[int, bool]:
# When the request requires prompt logprobs, we skip prefix caching.
if request.sampling_params.prompt_logprobs is not None:
return 0, False
request_id = request.request_id
# The block hashes for the request may already be computed
# if the scheduler has tried to schedule the request before.
block_hashes = self.req_to_block_hashes[request_id]
if not block_hashes:
block_hashes = request.block_hashes
self.req_to_block_hashes[request_id] = block_hashes
max_cache_hit_length = request.num_tokens - 1
eagle_kwarg = {"drop_eagle_block": self.use_eagle}
computed_blocks = self.single_type_manager.find_longest_cache_hit(
block_hashes=block_hashes,
max_length=max_cache_hit_length,
kv_cache_group_ids=[0],
block_pool=self.block_pool,
kv_cache_spec=self.single_type_manager.kv_cache_spec,
**eagle_kwarg,
alignment_tokens=self.block_size,
)
num_computed_tokens = len(computed_blocks[0]) * self.block_size
self.req_to_computed_blocks[request_id] = computed_blocks[0]
# We should touch these blocks in the concurrent scenarios.
self.block_pool.touch(computed_blocks)
# cup prefix cache status set and log
assert self.cpu_cache_stats is not None and self.cpu_cache_stats.prefix_cache_stats is not None
self.cpu_cache_stats.set_cache_stats(request.num_tokens, num_computed_tokens)
self.cpu_cache_stats.cpu_prefix_cache_metrics.observe(self.cpu_cache_stats.prefix_cache_stats)
self.cpu_cache_stats.log()
return num_computed_tokens, False
def _release_ahead_touch(self, request_id: str):
computed_blocks = self.req_to_computed_blocks[request_id]
if computed_blocks:
self.single_type_manager.block_pool.free_blocks(reversed(computed_blocks))
self.req_to_computed_blocks.pop(request_id, None)
def allocate_slots(self, req_to_num_tokens: dict[str, int], unallocated_req_ids: set[str]) -> dict[str, list[int]]:
for request_id in unallocated_req_ids:
self._free_slots(request_id)
req_to_new_blocks = {}
for request_id, num_tokens in req_to_num_tokens.items():
if self.req_failed_to_allocate[request_id]:
continue
new_computed_blocks = self.req_to_computed_blocks[request_id]
num_local_computed_tokens = len(new_computed_blocks) * self.block_size
num_blocks_to_allocate = self.single_type_manager.get_num_blocks_to_allocate(
request_id=request_id,
num_tokens=num_tokens,
new_computed_blocks=new_computed_blocks,
total_computed_tokens=num_local_computed_tokens,
num_tokens_main_model=num_tokens,
)
if num_blocks_to_allocate > self.block_pool.get_num_free_blocks():
self._release_ahead_touch(request_id)
self.req_failed_to_allocate[request_id] = True
continue
# Append the new computed blocks to the request blocks until now to
# avoid the case where the new blocks cannot be allocated.
self.single_type_manager.allocate_new_computed_blocks(
request_id,
new_computed_blocks,
num_local_computed_tokens=num_local_computed_tokens,
num_external_computed_tokens=0,
)
# Allocate new blocks but do not cache now.
new_blocks = self.single_type_manager.allocate_new_blocks(
request_id,
num_tokens,
num_tokens,
)
self.req_to_num_tokens[request_id] = num_tokens
# No need to release ref_cnt because we use officially.
self.req_to_computed_blocks.pop(request_id, None)
req_to_new_blocks[request_id] = [block.block_id for block in new_computed_blocks + new_blocks]
return req_to_new_blocks
def record_request_cache_and_free_slots(self, request: Request):
logger.debug("record_request_cache_and_free_slots for request %s in cpu_kv_cache_manager", request.request_id)
self.req_to_free[request.request_id] = request
def cache_and_free_slots(self, request_id: str):
logger.debug("Cache and free slots for request %s in cpu_kv_cache_manager", request_id)
if request_id not in self.req_to_free:
logger.error("request %s not in req_to_free, maybe bug!", request_id)
return
request = self.req_to_free[request_id]
if not self.req_failed_to_allocate[request_id]:
self.single_type_manager.cache_blocks(
request,
self.req_to_num_tokens[request_id],
)
self._free_slots(request_id)
logger.debug("delete request %s in cpu_kv_cache_manager req_to_free", request_id)
del self.req_to_free[request_id]
def _free_slots(self, request_id: str):
# This function is designed to be reentrant.
self._release_ahead_touch(request_id)
self.single_type_manager.free(request_id)
self.req_to_block_hashes.pop(request_id, None)
self.req_to_computed_blocks.pop(request_id, None)
self.req_failed_to_allocate.pop(request_id, None)
self.req_to_num_tokens.pop(request_id, None)

View File

@@ -0,0 +1,448 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import copy
import queue
import threading
import time
from collections import defaultdict
from collections.abc import Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Optional
import torch
from vllm.config import VllmConfig, get_layers_from_vllm_config
from vllm.distributed.ec_transfer import get_ec_transfer, has_ec_transfer
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole
from vllm.distributed.parallel_state import get_pp_group, get_tp_group
from vllm.logger import logger
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.model_executor.layers.mamba.abstract import MambaBase
from vllm.utils.torch_utils import STR_DTYPE_TO_TORCH_DTYPE
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheSpec
from vllm_ascend.distributed.kv_transfer.kv_pool.cpu_offload.metadata import (
MetadataServer,
MetadataServerProc,
MLAConfig,
)
if TYPE_CHECKING:
from vllm.forward_context import ForwardContext
from vllm.v1.attention.backend import AttentionMetadata # type: ignore
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.request import Request
from vllm.model_executor.layers.attention import Attention, MLAAttention
@dataclass
class ReqMeta:
gpu_block_ids: list[int]
cpu_block_ids: list[int]
num_scheduled_tokens: int
num_computed_tokens: int
num_gpu_computed_tokens: int
num_cpu_computed_tokens: int
def update(self, other: "ReqMeta"):
self.gpu_block_ids.extend(other.gpu_block_ids)
self.cpu_block_ids.extend(other.cpu_block_ids)
self.num_scheduled_tokens = other.num_scheduled_tokens
self.num_computed_tokens = other.num_computed_tokens
self.num_gpu_computed_tokens = other.num_gpu_computed_tokens
self.num_cpu_computed_tokens = other.num_cpu_computed_tokens
@dataclass
class CPUOffloadingConnectorMetadata(KVConnectorMetadata):
requests: dict[str, ReqMeta]
finished_req_ids: set[str]
class CPUOffloadingConnector(KVConnectorBase_V1):
def __init__(
self, vllm_config: VllmConfig, role: KVConnectorRole, kv_cache_config: Optional["KVCacheConfig"] = None
):
self._connector_metadata = CPUOffloadingConnectorMetadata(requests={}, finished_req_ids=set())
if not vllm_config.cache_config.enable_prefix_caching:
self.connector_scheduler: CPUOffloadingConnectorScheduler | None = None
self.connector_worker: CPUOffloadingConnectorWorker | None = None
elif role == KVConnectorRole.SCHEDULER:
self.connector_scheduler = CPUOffloadingConnectorScheduler(vllm_config)
self.connector_worker = None
elif role == KVConnectorRole.WORKER:
self.connector_scheduler = None
self.connector_worker = CPUOffloadingConnectorWorker(vllm_config)
# ==============================
# Worker-side methods
# ==============================
def bind_connector_metadata(self, connector_metadata: KVConnectorMetadata) -> None:
if self.connector_worker is not None:
assert isinstance(connector_metadata, CPUOffloadingConnectorMetadata)
self.connector_worker.bind_connector_metadata(connector_metadata)
def clear_connector_metadata(self) -> None:
assert self.connector_worker is not None
self.connector_worker.clear_connector_metadata()
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
if self.connector_worker is not None:
self.connector_worker.register_kv_caches(kv_caches)
def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None:
if self.connector_worker is not None:
self.connector_worker.start_load_kv()
def wait_for_layer_load(self, layer_name: str) -> None:
if self.connector_worker is not None:
self.connector_worker.wait_for_layer_load()
def save_kv_layer(
self, layer_name: str, kv_layer: torch.Tensor, attn_metadata: "AttentionMetadata", **kwargs
) -> None:
pass
def wait_for_save(self):
pass
def get_finished(self, finished_req_ids: set[str]) -> tuple[set[str] | None, set[str] | None]:
assert self.connector_worker is not None
return self.connector_worker.get_finished(), None
# Scheduler-side methods
# ==============================
def get_num_new_matched_tokens(self, request: "Request", num_computed_tokens: int) -> tuple[int, bool]:
if self.connector_scheduler is not None:
return self.connector_scheduler.get_num_new_matched_tokens(request, num_computed_tokens)
return 0, False
def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", num_external_tokens: int):
if self.connector_scheduler is not None:
return self.connector_scheduler.update_state_after_alloc(request)
def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnectorMetadata:
if self.connector_scheduler is not None:
return self.connector_scheduler.build_connector_meta(scheduler_output)
return KVConnectorMetadata()
def request_finished(self, request: "Request", block_ids: list[int]) -> tuple[bool, dict[str, Any] | None]:
if self.connector_scheduler is not None:
self.connector_scheduler.request_finished(request)
return True, None
class CPUOffloadingConnectorScheduler:
def __init__(self, vllm_config: VllmConfig):
logger.info("init CPUOffloadingConnectorScheduler")
self.vllm_config = vllm_config
self.block_size = vllm_config.cache_config.block_size
self.use_mla = vllm_config.model_config.use_mla
self.num_gpu_computed_tokens: dict[str, int] = {}
self.num_cpu_computed_tokens: dict[str, int] = {}
self.allocated_req_ids: set[str] = set()
self.finished_req_ids: list[str] = []
self.zmq_rpc_client = MetadataServer.ZMQRPCClient()
self.zmq_rpc_client.call("post_init")
if vllm_config.kv_transfer_config is not None:
self.swap_in_threshold = vllm_config.kv_transfer_config.get_from_extra_config("swap_in_threshold", 0)
else:
self.swap_in_threshold = 0
logger.info("swap_in_threshold: %s", self.swap_in_threshold)
def get_num_new_matched_tokens(self, ori_request: "Request", num_computed_tokens: int) -> tuple[int, bool]:
request = copy.deepcopy(ori_request)
request.get_hash_new_full_blocks = None
num_cpu_computed_tokens, load_async = self.zmq_rpc_client.call("get_matched_num_and_touch", request)
self.num_gpu_computed_tokens[request.request_id] = num_computed_tokens
self.num_cpu_computed_tokens[request.request_id] = num_cpu_computed_tokens
if num_cpu_computed_tokens - num_computed_tokens >= self.swap_in_threshold:
return num_cpu_computed_tokens - num_computed_tokens, load_async
else:
return 0, load_async
def update_state_after_alloc(self, request: "Request"):
self.allocated_req_ids.add(request.request_id)
def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnectorMetadata:
num_tokens = {}
# process scheduled_new_reqs
for req in scheduler_output.scheduled_new_reqs:
req_id = req.req_id
num_tokens[req_id] = req.num_computed_tokens + scheduler_output.num_scheduled_tokens[req_id]
# process scheduled_cached_reqs
cached_reqs = scheduler_output.scheduled_cached_reqs
for idx, req_id in enumerate(cached_reqs.req_ids):
num_tokens[req_id] = cached_reqs.num_computed_tokens[idx] + scheduler_output.num_scheduled_tokens[req_id]
unallocated_req_ids = set(
self.num_gpu_computed_tokens.keys() - self.allocated_req_ids - scheduler_output.num_scheduled_tokens.keys()
)
new_cpu_block_ids = self.zmq_rpc_client.call("allocate_slots", num_tokens, unallocated_req_ids)
metadata = CPUOffloadingConnectorMetadata(
requests={},
finished_req_ids=set(self.finished_req_ids),
)
for req in scheduler_output.scheduled_new_reqs:
req_id = req.req_id
gpu_block_ids = req.block_ids[0]
metadata.requests[req_id] = ReqMeta(
gpu_block_ids=[] if gpu_block_ids is None else gpu_block_ids,
cpu_block_ids=new_cpu_block_ids.get(req_id, []),
num_scheduled_tokens=scheduler_output.num_scheduled_tokens[req_id],
num_computed_tokens=req.num_computed_tokens,
num_gpu_computed_tokens=self.num_gpu_computed_tokens[req_id],
num_cpu_computed_tokens=self.num_cpu_computed_tokens[req_id],
)
for idx, req_id in enumerate(cached_reqs.req_ids):
gpu_block_ids = cached_reqs.new_block_ids[idx]
metadata.requests[req_id] = ReqMeta(
gpu_block_ids=[] if gpu_block_ids is None else gpu_block_ids,
cpu_block_ids=new_cpu_block_ids.get(req_id, []),
num_scheduled_tokens=scheduler_output.num_scheduled_tokens[req_id],
num_computed_tokens=cached_reqs.num_computed_tokens[idx],
num_gpu_computed_tokens=cached_reqs.num_computed_tokens[idx],
num_cpu_computed_tokens=cached_reqs.num_computed_tokens[idx],
)
self.num_gpu_computed_tokens.clear()
self.num_cpu_computed_tokens.clear()
self.allocated_req_ids.clear()
self.finished_req_ids.clear()
return metadata
def request_finished(self, ori_request: "Request"):
request = copy.deepcopy(ori_request)
request.get_hash_new_full_blocks = None
self.finished_req_ids.append(request.request_id)
# inform metadata server to record request, and free it after finish sending
self.zmq_rpc_client.call("record_request_cache_and_free_slots", request)
class CPUOffloadingConnectorWorker:
def __init__(self, vllm_config: VllmConfig):
logger.info("init CPUOffloadingConnectorWorker")
self.vllm_config = vllm_config
self.block_size = vllm_config.cache_config.block_size
self.pp_rank = get_pp_group().rank_in_group
self.tp_group = get_tp_group()
self.tp_rank = self.tp_group.rank_in_group
self.tp_world_size = self.tp_group.world_size
self.use_mla = vllm_config.model_config.use_mla
self.requests: dict[str, ReqMeta] = {}
self.load_stream = torch.npu.Stream()
self.save_stream = torch.npu.Stream()
self.zmq_rpc_client = MetadataServer.ZMQRPCClient()
self.load_block_mapping: list[tuple[int, int]] = []
self.save_input_queue: queue.Queue[tuple[str, ReqMeta]] = queue.Queue()
self.save_output_queue: queue.Queue[str] = queue.Queue()
self.save_thread = threading.Thread(target=self._save_listener)
self.save_thread.start()
self.done_sending_count: defaultdict[str, int] = defaultdict(int)
# start metadata server to init cpu_kv_cache_manager and handle rpc requests
# all dp shared the same metadata server, only start the process on data_rank 0
if vllm_config.parallel_config.data_parallel_rank == 0 and self.tp_rank == 0 and self.pp_rank == 0:
config = VllmConfig()
config.cache_config = vllm_config.cache_config
config.parallel_config = vllm_config.parallel_config
config.kv_transfer_config = vllm_config.kv_transfer_config
self.init_metadata_server(config)
self._wait_for_metadata_process_start()
def init_metadata_server(self, vllm_config: VllmConfig):
self.metadata_thread = threading.Thread(
target=MetadataServerProc.run_metadata_server,
args=(vllm_config,),
)
self.metadata_thread.daemon = True
self.metadata_thread.start()
def _wait_for_metadata_process_start(self):
# TODO: wait for metadata server to start, add a rpc to check if ready
while True:
try:
if self.zmq_rpc_client.call("ready"):
break
except Exception as e:
logger.info("wait for metadata server to start, error: %s", e)
time.sleep(1)
def bind_connector_metadata(self, connector_metadata: CPUOffloadingConnectorMetadata) -> None:
for req_id, req in connector_metadata.requests.items():
if req_id in self.requests:
self.requests[req_id].update(req)
req = self.requests[req_id]
else:
self.requests[req_id] = req
for i in range(req.num_gpu_computed_tokens // self.block_size, req.num_computed_tokens // self.block_size):
self.load_block_mapping.append((req.cpu_block_ids[i], req.gpu_block_ids[i]))
for req_id in connector_metadata.finished_req_ids:
if req_id in self.requests:
self.save_input_queue.put((req_id, self.requests[req_id]))
def clear_connector_metadata(self) -> None:
self.load_block_mapping.clear()
def register_kv_caches(self, kv_caches: dict[str, Sequence[torch.Tensor]]):
self.gpu_kv_caches = kv_caches
model_config = self.vllm_config.model_config
mla_config: MLAConfig | None = None
if model_config.use_mla:
mla_config = MLAConfig(
model_config.hf_text_config.kv_lora_rank, model_config.hf_text_config.qk_rope_head_dim
)
self.cpu_kv_caches = list(
self.zmq_rpc_client.call(
"init_cpu_kv_caches",
self.pp_rank,
self.tp_rank,
get_kv_cache_spec(self.vllm_config),
mla_config,
).values()
)
def start_load_kv(self) -> None:
self.current_layer = 0
self.gpu_kv_caches_load_iter = iter(self.gpu_kv_caches.values())
self.load_kv_layer(0)
def wait_for_layer_load(self) -> None:
# TODO: Replace with `torch.npu.current_stream().wait_stream(self.load_stream)` after fixing the bug.
self.load_stream.synchronize()
self.current_layer += 1
self.load_kv_layer(self.current_layer)
def load_kv_layer(self, layer: int):
if layer == len(self.gpu_kv_caches):
return
gpu_kv_caches = next(self.gpu_kv_caches_load_iter)
cpu_kv_caches = self.cpu_kv_caches[layer]
with torch.npu.stream(self.load_stream):
for cpu_block_id, gpu_block_id in self.load_block_mapping:
for gpu_layer_part, cpu_layer_part in zip(gpu_kv_caches, cpu_kv_caches):
gpu_layer_part[gpu_block_id].copy_(cpu_layer_part[cpu_block_id], non_blocking=True)
def get_finished(self) -> set[str]:
done_sending: set[str] = set()
while True:
try:
id = self.save_output_queue.get_nowait()
except queue.Empty:
break
done_sending.add(id)
for id in done_sending:
del self.requests[id]
if self.tp_world_size == 1:
return done_sending
if self.tp_rank == 0:
for req_id in done_sending:
self.done_sending_count[req_id] += 1
other_ranks_finished_ids: list[str] = []
for i in range(1, self.tp_world_size):
other_ranks_finished_ids.extend(self.tp_group.recv_object(src=i))
for req_id in other_ranks_finished_ids:
self.done_sending_count[req_id] += 1
all_done_sending: set[str] = set()
for req_id in list(self.done_sending_count.keys()):
if self.done_sending_count[req_id] == self.tp_world_size:
del self.done_sending_count[req_id]
all_done_sending.add(req_id)
# release cpu_kv_cache after request sending finished
# to avoid rpc blocking, use thread to call rpc asynchronously
sending_finished_thread = threading.Thread(target=self._sending_finished, args=(all_done_sending,))
sending_finished_thread.daemon = True
sending_finished_thread.start()
return all_done_sending
else:
self.tp_group.send_object(done_sending, dst=0)
return done_sending
def _sending_finished(self, all_done_sending):
for req_id in all_done_sending:
logger.debug("call cache_and_free_slots for req_id: %s", req_id)
self.zmq_rpc_client.call("cache_and_free_slots", req_id)
def _save_listener(self):
save_block_mapping = []
while True:
req_id, req = self.save_input_queue.get()
for i in range(
req.num_cpu_computed_tokens // self.block_size,
min((req.num_computed_tokens + req.num_scheduled_tokens) // self.block_size, len(req.cpu_block_ids)),
):
save_block_mapping.append((req.gpu_block_ids[i], req.cpu_block_ids[i]))
with torch.npu.stream(self.save_stream):
# MLA: kv_layer is tuple[tensor, tensor] means (rope, nope).
# non-MLA: kv_layer is list[tensor], typically means [k, v].
if self.use_mla:
start, step = self.tp_rank, self.tp_world_size
else:
start, step = 0, 1
for i in range(start, len(save_block_mapping), step):
gpu_block_id, cpu_block_id = save_block_mapping[i]
for cpu_kv_caches, gpu_kv_caches in zip(self.cpu_kv_caches, self.gpu_kv_caches.values()):
for cpu_layer_part, gpu_layer_part in zip(cpu_kv_caches, gpu_kv_caches):
cpu_layer_part[cpu_block_id].copy_(gpu_layer_part[gpu_block_id], non_blocking=True)
self.save_stream.synchronize()
self.save_output_queue.put(req_id)
save_block_mapping.clear()
# copied and modified from vllm_ascend/worker/model_runner_v1.py
def get_kv_cache_spec(vllm_config: VllmConfig) -> dict[str, KVCacheSpec]:
"""
Generates the KVCacheSpec by parsing the kv cache format from each
Attention module in the static forward context.
Returns:
KVCacheSpec: A dictionary mapping layer names to their KV cache
format. Layers that do not need KV cache are not included.
"""
if has_ec_transfer() and get_ec_transfer().is_producer:
return {}
use_sparse = hasattr(vllm_config.model_config.hf_config, "index_topk")
if vllm_config.cache_config.cache_dtype == "auto":
kv_cache_dtype = vllm_config.model_config.dtype
else:
kv_cache_dtype = STR_DTYPE_TO_TORCH_DTYPE[vllm_config.cache_config.cache_dtype]
kv_cache_spec: dict[str, KVCacheSpec] = {}
attn_layers = get_layers_from_vllm_config(vllm_config, AttentionLayerBase)
# NOTE: Must process Attention/MLAAttention before MambaBase to maintain
# ordering expected by graph parameter update logic in attention backends.
mamba_layers: dict[str, MambaBase] = {}
for layer_name, attn_module in attn_layers.items():
if isinstance(attn_module, Attention):
if spec := attn_module.get_kv_cache_spec(vllm_config):
kv_cache_spec[layer_name] = spec
elif isinstance(attn_module, MLAAttention):
if use_sparse:
# TODO(cmq): This is a hack way to fix deepseek kvcache when
# using DSA. Fix the spec in vLLM is the final way.
block_size = vllm_config.cache_config.block_size
kv_cache_spec[layer_name] = FullAttentionSpec(
block_size=block_size, num_kv_heads=1, head_size=attn_module.head_size, dtype=kv_cache_dtype
)
elif spec := attn_module.get_kv_cache_spec(vllm_config):
kv_cache_spec[layer_name] = spec
elif isinstance(attn_module, MambaBase):
mamba_layers[layer_name] = attn_module
if len(mamba_layers) > 0:
if vllm_config.cache_config.enable_prefix_caching:
raise NotImplementedError("Prefix caching is not supported for Mamba yet.")
for layer_name, mamba_module in mamba_layers.items():
if spec := mamba_module.get_kv_cache_spec(vllm_config):
kv_cache_spec[layer_name] = spec
return kv_cache_spec

View File

@@ -0,0 +1,258 @@
import math
import os
import pickle
from collections.abc import Callable
from dataclasses import dataclass
from multiprocessing.shared_memory import SharedMemory
from typing import Any
import torch
import vllm.envs as envs
import zmq
from vllm.config import KVTransferConfig, VllmConfig
from vllm.logger import logger
from vllm.utils.network_utils import make_zmq_socket
from vllm.utils.torch_utils import get_dtype_size
from vllm.v1.kv_cache_interface import AttentionSpec
from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec
from vllm_ascend.distributed.kv_transfer.kv_pool.cpu_offload.cpu_kv_cache_manager import CPUKVCacheManager
@dataclass
class MLAConfig:
nope_dim: int
rope_dim: int
def get_cpu_offload_connector(vllm_config: VllmConfig) -> KVTransferConfig:
if vllm_config.kv_transfer_config is not None:
kv_transfer_config = vllm_config.kv_transfer_config
if kv_transfer_config.kv_connector == "CPUOffloadingConnector":
return kv_transfer_config
elif kv_transfer_config.kv_connector == "MultiConnector":
ktcs = kv_transfer_config.kv_connector_extra_config.get("connectors")
for ktc in ktcs:
kv_transfer_config = KVTransferConfig(**ktc)
if kv_transfer_config.kv_connector == "CPUOffloadingConnector":
return kv_transfer_config
return None
class MetadataServer:
METADATA_SERVER_ADDRESS = f"ipc://{envs.VLLM_RPC_BASE_PATH}/metadata.ipc"
DEFAULT_CPU_SWAP_SPACE_GB = 800
class ZMQRPCClient:
def __init__(self, identity=None):
if identity is None:
identity = f"worker-{os.getpid()}-{id(self)}"
logger.info("metadata client for worker %s started", identity)
self.ctx = zmq.Context() # type: ignore
self.socket = make_zmq_socket(
self.ctx,
MetadataServer.METADATA_SERVER_ADDRESS,
zmq.DEALER, # type: ignore
bind=False,
identity=identity.encode(),
linger=0,
)
def call(self, func_name: str, *args, **kwargs) -> Any:
request = (func_name, args, kwargs)
self.socket.send(b"", zmq.SNDMORE) # type: ignore
self.socket.send(pickle.dumps(request))
_ = self.socket.recv()
response = pickle.loads(self.socket.recv())
result, error = response
if error:
logger.exception("call metadata sever error: %s", error)
raise error
if func_name == "init_cpu_kv_caches":
(memory_dict, layer_size, layer_dtype, mla_config) = result
# shared_memory_dict is recorded in self to close
self.shared_memory_dict = memory_dict
result = {}
for key, shm in memory_dict.items():
tensor = torch.frombuffer(shm.buf, dtype=layer_dtype).reshape(layer_size)
if mla_config is not None:
tensor = tensor.split([mla_config.nope_dim, mla_config.rope_dim], dim=-1)
result[key] = tensor
return result
def __del__(self):
# will be finalized by outer process
self.socket.close()
self.ctx.term()
if hasattr(self, "shared_memory_dict"):
for shm in self.shared_memory_dict.values():
shm.close()
def __init__(self, vllm_config: VllmConfig):
self.world_size = vllm_config.parallel_config.world_size
self.pipeline_parallel_size = vllm_config.parallel_config.pipeline_parallel_size
kv_transfer_config = get_cpu_offload_connector(vllm_config)
assert kv_transfer_config is not None
available_memory_gb = kv_transfer_config.get_from_extra_config(
"cpu_swap_space_gb", MetadataServer.DEFAULT_CPU_SWAP_SPACE_GB
)
self.available_memory = available_memory_gb * 1024 * 1024 * 1024
logger.info("cpu swap space: %s bytes", self.available_memory)
self.ctx = zmq.Context() # type: ignore
self.socket = make_zmq_socket(
self.ctx,
MetadataServer.METADATA_SERVER_ADDRESS,
zmq.ROUTER, # type: ignore
bind=True,
linger=0,
)
self.functions: dict[str, Callable] = {
"init_cpu_kv_caches": self.init_cpu_kv_caches,
"post_init": self.post_init,
"ready": self.ready,
}
self.shared_memory = {} # type: ignore
self.num_cpu_blocks = -1
@staticmethod
def _safe_create_shared_memory(name: str, size: int) -> SharedMemory:
try:
existing_shm = SharedMemory(name=name, create=False)
existing_shm.close()
existing_shm.unlink()
except FileNotFoundError:
pass
return SharedMemory(name=name, create=True, size=size)
def ready(self):
return True
def init_cpu_kv_caches(
self,
pp_rank: int,
tp_rank: int,
kv_cache_specs: dict[str, AttentionSpec],
mla_config: MLAConfig,
) -> tuple[dict[str, SharedMemory], tuple[int, ...], torch.dtype, MLAConfig]:
logger.info("receive pp rank: %s, tp rank: %s", pp_rank, tp_rank)
# follow the assumption that each layer has the same spec
layer = next(iter(kv_cache_specs.values()))
assert all([layer.page_size_bytes == any.page_size_bytes for any in kv_cache_specs.values()])
use_mla = isinstance(layer, AscendMLAAttentionSpec)
# mla shares the same kv cache among different tp
if use_mla:
tp_rank = 0
if (pp_rank, tp_rank) in self.shared_memory:
return self.shared_memory[(pp_rank, tp_rank)]
available_memory = self.available_memory
shared_memory_dict = {}
if use_mla:
available_memory //= self.pipeline_parallel_size
available_memory //= len(kv_cache_specs)
num_blocks = available_memory // layer.page_size_bytes
layer_size = (num_blocks, layer.block_size, layer.num_kv_heads, layer.head_size) # type: ignore
else:
available_memory //= self.world_size
available_memory //= len(kv_cache_specs)
num_blocks = available_memory // layer.page_size_bytes
layer_size = (2, num_blocks, layer.block_size, layer.num_kv_heads, layer.head_size) # type: ignore
nbytes = math.prod(layer_size) * get_dtype_size(layer.dtype)
for layer_name in kv_cache_specs:
# only this format can share during ZeroMQ+pickle
shared_memory_dict[layer_name] = MetadataServer._safe_create_shared_memory(
f"cpu_kv_cache_{pp_rank}_{tp_rank}_{layer_name}", nbytes
)
if use_mla:
assert mla_config is not None
assert layer.head_size == mla_config.rope_dim + mla_config.nope_dim
self.shared_memory[(pp_rank, tp_rank)] = (shared_memory_dict, layer_size, layer.dtype, mla_config)
else:
self.shared_memory[(pp_rank, tp_rank)] = (shared_memory_dict, layer_size, layer.dtype, None)
if self.num_cpu_blocks == -1 or num_blocks < self.num_cpu_blocks:
self.num_cpu_blocks = num_blocks
self.layer = layer
return self.shared_memory[(pp_rank, tp_rank)]
def post_init(self):
# different processors in data parallel may call multiple times
if hasattr(self, "cpu_block_manager"):
return
# do shared_memory() at least once
logger.info("assign cpu num blocks: %s", self.num_cpu_blocks)
assert self.num_cpu_blocks >= 0
self.cpu_block_manager = CPUKVCacheManager(self.layer, self.num_cpu_blocks)
self.functions.update(
{
"get_matched_num_and_touch": self.cpu_block_manager.get_matched_num_and_touch,
"allocate_slots": self.cpu_block_manager.allocate_slots,
"record_request_cache_and_free_slots": self.cpu_block_manager.record_request_cache_and_free_slots,
"cache_and_free_slots": self.cpu_block_manager.cache_and_free_slots,
}
)
def serve_step(self):
client_id = self.socket.recv()
_ = self.socket.recv()
raw_msg = self.socket.recv()
try:
func_name, args, kwargs = pickle.loads(raw_msg)
except Exception as e:
response = (None, Exception(f"Invalid request: {str(e)}"))
else:
if func_name in self.functions:
try:
result = self.functions[func_name](*args, **kwargs)
response = (result, None) # type: ignore
except Exception as e:
logger.exception("metadata execute error: %s", e)
response = (None, e) # type: ignore
else:
response = (None, NameError(f"Function {func_name} not found"))
self.socket.send(client_id, zmq.SNDMORE) # type: ignore
self.socket.send(b"", zmq.SNDMORE) # type: ignore
self.socket.send(pickle.dumps(response))
def shutdown(self):
self.socket.close()
self.ctx.term()
socket_path = MetadataServer.METADATA_SERVER_ADDRESS.replace("ipc://", "")
if os.path.exists(socket_path):
os.remove(socket_path)
for cached in self.shared_memory.values():
for shm in cached[0].values():
shm.close()
shm.unlink()
class MetadataServerProc:
@staticmethod
def run_metadata_server(vllm_config: VllmConfig):
if not vllm_config.cache_config.enable_prefix_caching or get_cpu_offload_connector(vllm_config) is None:
return
shutdown_requested = False
def _signal_handler(signum, frame):
nonlocal shutdown_requested
if not shutdown_requested:
shutdown_requested = True
raise SystemExit()
# Either SIGTERM or SIGINT will terminate the worker
# signal.signal(signal.SIGTERM, _signal_handler)
# signal.signal(signal.SIGINT, _signal_handler)
metadata_server: MetadataServer | None = None
try:
metadata_server = MetadataServer(vllm_config)
logger.info("Metadata server started.")
while True:
metadata_server.serve_step()
except SystemExit:
logger.info("Metadata server exiting.")
raise
except Exception as e:
logger.exception("Metadata server error: %s.", e)
raise e
finally:
if metadata_server is not None:
metadata_server.shutdown()

View File

@@ -0,0 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import lmcache_ascend # noqa: F401
from vllm.distributed.kv_transfer.kv_connector.v1.lmcache_connector import LMCacheConnectorV1
__all__ = ["LMCacheConnectorV1"]

View File

@@ -0,0 +1,630 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Scheduler-side manager for recompute CPU offloading."""
import contextlib
from collections.abc import Iterable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from vllm.config import VllmConfig
from vllm.distributed.kv_events import KVCacheEvent
from vllm.logger import logger
from vllm.utils.math_utils import cdiv
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_coordinator import (
KVCacheCoordinator,
get_kv_cache_coordinator,
)
from vllm.v1.core.kv_cache_utils import resolve_kv_cache_block_sizes
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import SlidingWindowSpec, UniformTypeKVCacheSpecs
from vllm.v1.outputs import KVConnectorOutput
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.metadata import (
RecomputeCPUOffloadMetadata,
RecomputeCPUOffloadWorkerMetadata,
)
if TYPE_CHECKING:
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.core.kv_cache_utils import BlockHashWithGroupId, KVCacheBlock
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.request import Request
@dataclass
class TransferMeta:
gpu_block_ids: list[int]
cpu_block_ids: list[int]
@dataclass
class PreemptedRequestState:
req_id: str
cpu_block_ids: tuple[list[int], ...]
num_computed_tokens: int
store_transfer_meta: TransferMeta
store_event: int | None = None
load_event: int | None = None
load_transfer_meta: TransferMeta | None = None
load_start_tokens: int = 0
ready: bool = False
finished: bool = False
class RecomputeCPUOffloadScheduler:
"""Preserve preempted requests' KV blocks in CPU memory.
When offload prefix caching is enabled, full hashed blocks share CPU
blocks. Otherwise every offloaded block is private to its request.
"""
def __init__(
self,
vllm_config: VllmConfig,
kv_cache_config: "KVCacheConfig | None",
cpu_capacity_bytes: int,
enable_offload_prefix_caching: bool = True,
):
assert kv_cache_config is not None
self.vllm_config = vllm_config
self.enable_offload_prefix_caching = enable_offload_prefix_caching
self.cpu_kv_cache_config = self._derive_cpu_config(kv_cache_config, cpu_capacity_bytes)
self.num_cpu_blocks = self.cpu_kv_cache_config.num_blocks
self._group_is_sliding_window = self._get_group_is_sliding_window(kv_cache_config)
self.enable_kv_cache_events = (
vllm_config.kv_events_config is not None and vllm_config.kv_events_config.enable_kv_cache_events
)
logger.info(
"RecomputeCPUOffloadScheduler: allocating %d CPU blocks (%.2f GB) for recompute offload, prefix caching=%s",
self.num_cpu_blocks,
cpu_capacity_bytes / (1024**3),
self.enable_offload_prefix_caching,
)
dcp_world_size = vllm_config.parallel_config.decode_context_parallel_size
pcp_world_size = vllm_config.parallel_config.prefill_context_parallel_size
assert dcp_world_size == 1 and pcp_world_size == 1
scheduler_block_size, hash_block_size = resolve_kv_cache_block_sizes(kv_cache_config, vllm_config)
self.cpu_coordinator: KVCacheCoordinator = get_kv_cache_coordinator(
kv_cache_config=self.cpu_kv_cache_config,
max_model_len=vllm_config.model_config.max_model_len,
max_num_batched_tokens=(vllm_config.scheduler_config.max_num_batched_tokens),
use_eagle=False,
enable_caching=self.enable_offload_prefix_caching,
enable_kv_cache_events=self.enable_kv_cache_events,
dcp_world_size=dcp_world_size,
pcp_world_size=pcp_world_size,
scheduler_block_size=scheduler_block_size,
hash_block_size=hash_block_size,
)
self.cpu_block_pool: BlockPool = self.cpu_coordinator.block_pool
self._gpu_block_pool: BlockPool | None = None
self._preempted_req_states: dict[str, PreemptedRequestState] = {}
self._preempt_store_event_to_reqs: dict[int, list[str]] = {}
self._preempt_store_event_to_blocks: dict[int, TransferMeta] = {}
self._preempt_load_event_to_reqs: dict[int, list[str]] = {}
# Hash blocks created before build_connector_meta() are shared by all
# requests preempted in the same scheduling step.
self._pending_hash_blocks: dict[BlockHashWithGroupId, KVCacheBlock] = {}
self._load_event_counter = 0
self._store_event_counter = 0
self._expected_worker_count = vllm_config.parallel_config.world_size
self._store_event_pending_counts: dict[int, int] = {}
@staticmethod
def _get_group_is_sliding_window(kv_cache_config: "KVCacheConfig") -> list[bool]:
group_is_sliding_window: list[bool] = []
for group in kv_cache_config.kv_cache_groups:
if isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs):
group_is_sliding_window.append(
any(isinstance(spec, SlidingWindowSpec) for spec in group.kv_cache_spec.kv_cache_specs.values())
)
else:
group_is_sliding_window.append(isinstance(group.kv_cache_spec, SlidingWindowSpec))
return group_is_sliding_window
@staticmethod
def _derive_cpu_config(gpu_config: "KVCacheConfig", cpu_capacity_bytes: int) -> "KVCacheConfig":
from vllm.v1.kv_cache_interface import KVCacheConfig as KVCacheConfigCls
from vllm.v1.kv_cache_interface import KVCacheTensor
assert gpu_config.kv_cache_tensors
gpu_kv_cache_tensors = []
for t in gpu_config.kv_cache_tensors:
if t.shared_by:
gpu_kv_cache_tensors.append(t)
gpu_total_bytes = sum(t.size for t in gpu_kv_cache_tensors)
num_gpu_blocks = gpu_config.num_blocks
num_cpu_blocks = max(1, num_gpu_blocks * cpu_capacity_bytes // gpu_total_bytes)
cpu_tensors = [
KVCacheTensor(
size=t.size // num_gpu_blocks * num_cpu_blocks,
shared_by=list(t.shared_by),
)
for t in gpu_kv_cache_tensors
]
return KVCacheConfigCls(
num_blocks=num_cpu_blocks,
kv_cache_tensors=cpu_tensors,
kv_cache_groups=gpu_config.kv_cache_groups,
)
def _align_group_block_ids(
self,
group_idx: int,
group_block_ids: list[int],
logical_num_blocks: int,
) -> list[int]:
if logical_num_blocks <= 0:
return []
aligned_group_block_ids = list(group_block_ids)
if self._group_is_sliding_window[group_idx] and len(aligned_group_block_ids) < logical_num_blocks:
aligned_group_block_ids = [0] * (
logical_num_blocks - len(aligned_group_block_ids)
) + aligned_group_block_ids
return aligned_group_block_ids[:logical_num_blocks]
def bind_gpu_block_pool(self, gpu_block_pool: BlockPool) -> None:
self._gpu_block_pool = gpu_block_pool
def has_preempted_request(self, req_id: str) -> bool:
return req_id in self._preempted_req_states
def get_num_new_matched_tokens(self, request: "Request", num_computed_tokens: int) -> tuple[int | None, bool]:
state = self._preempted_req_states.get(request.request_id)
if state is None:
return 0, False
if not state.ready:
return None, False
restorable_tokens = min(state.num_computed_tokens, request.num_tokens)
hit_length = max(0, restorable_tokens - num_computed_tokens)
if hit_length <= 0:
self._cleanup_preempt_cache_request(request.request_id)
return 0, False
state.load_start_tokens = num_computed_tokens
logger.debug(
"Recompute offload cache hit for request %s: load_start=%d, load_tokens=%d, stored_tokens=%d.",
request.request_id,
num_computed_tokens,
hit_length,
state.num_computed_tokens,
)
return hit_length, True
def update_state_after_alloc(
self,
request: "Request",
blocks: "KVCacheBlocks",
num_external_tokens: int,
) -> None:
if num_external_tokens <= 0:
return
prepared = self._prepare_preempt_load_after_alloc(
request,
blocks.get_block_ids(),
num_external_tokens,
)
if not prepared:
raise RuntimeError(
"Failed to prepare recompute H2D load after KV block "
f"allocation: req_id={request.request_id}, "
f"num_external_tokens={num_external_tokens}"
)
def update_state_before_preempt(
self,
request: "Request",
block_ids: tuple[list[int], ...],
num_computed_tokens: int,
) -> bool:
if request.request_id in self._preempted_req_states:
return True
return self._create_preempt_state(
request.request_id,
block_ids,
num_computed_tokens,
)
def _create_preempt_state(
self,
req_id: str,
block_ids_by_group: tuple[list[int], ...],
num_computed_tokens: int,
) -> bool:
if num_computed_tokens <= 0 or self._gpu_block_pool is None:
return False
kv_cache_groups = self.cpu_kv_cache_config.kv_cache_groups
group_gpu_blocks: list[list[KVCacheBlock | None]] = []
group_gpu_hashes: list[list[BlockHashWithGroupId | None]] = []
missing_hashes: set[BlockHashWithGroupId] = set()
num_unhashed = 0
for g, group_gpu_ids in enumerate(block_ids_by_group):
group_block_size = kv_cache_groups[g].kv_cache_spec.block_size
logical_num_blocks = cdiv(num_computed_tokens, group_block_size)
aligned_group_gpu_ids = self._align_group_block_ids(g, group_gpu_ids, logical_num_blocks)
eviction_group_gpu_ids = self._align_group_block_ids(
g,
group_gpu_ids,
max(logical_num_blocks, len(group_gpu_ids)),
)
gpu_blocks: list[KVCacheBlock | None] = []
effective_hashes: list[BlockHashWithGroupId | None] = []
for block_idx, block_id in enumerate(eviction_group_gpu_ids):
if block_id <= 0:
continue
gpu_block = self._gpu_block_pool.blocks[block_id]
block_is_computed = (block_idx + 1) * group_block_size <= num_computed_tokens
if not block_is_computed and gpu_block.block_hash is not None:
# allocate_slots() may assign a hash using tokens planned
# for this scheduling step. If the request is then
# preempted before forward, that block does not contain the
# hashed KV and must not remain in the GPU prefix cache.
self._gpu_block_pool._maybe_evict_cached_block(gpu_block)
for block_idx, block_id in enumerate(aligned_group_gpu_ids):
if block_id <= 0:
gpu_blocks.append(None)
effective_hashes.append(None)
continue
gpu_block = self._gpu_block_pool.blocks[block_id]
block_is_computed = (block_idx + 1) * group_block_size <= num_computed_tokens
block_hash = gpu_block.block_hash if block_is_computed and self.enable_offload_prefix_caching else None
gpu_blocks.append(gpu_block)
effective_hashes.append(block_hash)
if block_hash is None:
num_unhashed += 1
elif (
self.cpu_block_pool.cached_block_hash_to_block.get_one_block(block_hash) is None
and block_hash not in self._pending_hash_blocks
):
missing_hashes.add(block_hash)
group_gpu_blocks.append(gpu_blocks)
group_gpu_hashes.append(effective_hashes)
num_needed = num_unhashed + len(missing_hashes)
if not any(any(gpu_block is not None for gpu_block in group) for group in group_gpu_blocks):
return False
if num_needed > self.cpu_block_pool.get_num_free_blocks():
logger.warning(
"Skip recompute offload for request %s: CPU cache has %d free blocks, but %d new blocks are required.",
req_id,
self.cpu_block_pool.get_num_free_blocks(),
num_needed,
)
return False
cpu_block_iter = iter(self.cpu_block_pool.get_new_blocks(num_needed))
cpu_block_ids_by_group: list[list[int]] = []
store_gpu_block_ids: list[int] = []
store_cpu_block_ids: list[int] = []
waiting_for_store = False
for gpu_blocks, effective_hashes in zip(group_gpu_blocks, group_gpu_hashes):
group_cpu_ids: list[int] = []
for gpu_block, block_hash in zip(gpu_blocks, effective_hashes):
if gpu_block is None:
group_cpu_ids.append(0)
continue
cpu_block = None
if block_hash is not None:
cpu_block = self.cpu_block_pool.cached_block_hash_to_block.get_one_block(block_hash)
if cpu_block is not None:
self.cpu_block_pool.touch([cpu_block])
else:
cpu_block = self._pending_hash_blocks.get(block_hash)
if cpu_block is not None:
self.cpu_block_pool.touch([cpu_block])
waiting_for_store = True
else:
cpu_block = next(cpu_block_iter)
cpu_block._block_hash = block_hash
self._pending_hash_blocks[block_hash] = cpu_block
store_gpu_block_ids.append(gpu_block.block_id)
store_cpu_block_ids.append(cpu_block.block_id)
waiting_for_store = True
else:
cpu_block = next(cpu_block_iter)
store_gpu_block_ids.append(gpu_block.block_id)
store_cpu_block_ids.append(cpu_block.block_id)
waiting_for_store = True
group_cpu_ids.append(cpu_block.block_id)
cpu_block_ids_by_group.append(group_cpu_ids)
store_transfer = TransferMeta(store_gpu_block_ids, store_cpu_block_ids)
self._preempted_req_states[req_id] = PreemptedRequestState(
req_id=req_id,
cpu_block_ids=tuple(cpu_block_ids_by_group),
num_computed_tokens=num_computed_tokens,
store_transfer_meta=store_transfer,
ready=not waiting_for_store,
)
logger.info(
"Created recompute offload state for request %s: "
"computed_tokens=%d, cpu_blocks=%d, store_blocks=%d, "
"ready=%s.",
req_id,
num_computed_tokens,
sum(len(ids) for ids in cpu_block_ids_by_group),
len(store_cpu_block_ids),
not waiting_for_store,
)
return True
def _prepare_preempt_store_specs(
self,
) -> tuple[list[int], list[int], list[str]]:
gpu_block_ids: list[int] = []
cpu_block_ids: list[int] = []
req_ids: list[str] = []
for req_id, state in self._preempted_req_states.items():
if state.store_event is not None or state.ready:
continue
gpu_block_ids.extend(state.store_transfer_meta.gpu_block_ids)
cpu_block_ids.extend(state.store_transfer_meta.cpu_block_ids)
req_ids.append(req_id)
return gpu_block_ids, cpu_block_ids, req_ids
def _prepare_preempt_load_after_alloc(
self,
request: "Request",
block_ids_by_group: tuple[list[int], ...],
num_external_tokens: int,
) -> bool:
state = self._preempted_req_states.get(request.request_id)
if state is None or not state.ready:
return False
load_start_tokens = state.load_start_tokens
load_end_tokens = min(
load_start_tokens + num_external_tokens,
state.num_computed_tokens,
)
if load_end_tokens <= load_start_tokens:
return False
if len(block_ids_by_group) != len(state.cpu_block_ids):
raise RuntimeError(
"Recompute H2D KV group count mismatch: "
f"req_id={request.request_id}, "
f"gpu_groups={len(block_ids_by_group)}, "
f"cpu_groups={len(state.cpu_block_ids)}"
)
gpu_block_ids: list[int] = []
cpu_block_ids: list[int] = []
for g, group_cpu_ids in enumerate(state.cpu_block_ids):
group_block_size = self.cpu_kv_cache_config.kv_cache_groups[g].kv_cache_spec.block_size
start_block = load_start_tokens // group_block_size
end_block = min(
len(group_cpu_ids),
len(
self._align_group_block_ids(
g,
block_ids_by_group[g],
max(
cdiv(load_end_tokens, group_block_size),
len(block_ids_by_group[g]),
),
)
),
cdiv(load_end_tokens, group_block_size),
)
if end_block == start_block:
continue
if end_block < start_block:
raise RuntimeError(
"Recompute H2D produced an empty block range: "
f"req_id={request.request_id}, group={g}, "
f"start_block={start_block}, end_block={end_block}, "
f"gpu_blocks={len(block_ids_by_group[g])}, "
f"cpu_blocks={len(group_cpu_ids)}"
)
aligned_group_gpu_ids = self._align_group_block_ids(
g,
block_ids_by_group[g],
end_block,
)
for block_idx in range(start_block, end_block):
cpu_block_id = group_cpu_ids[block_idx]
gpu_block_id = aligned_group_gpu_ids[block_idx]
if cpu_block_id <= 0 or gpu_block_id <= 0:
continue
cpu_block_ids.append(cpu_block_id)
gpu_block_ids.append(gpu_block_id)
if not cpu_block_ids or len(cpu_block_ids) != len(gpu_block_ids):
raise RuntimeError(
"Recompute H2D block mapping is incomplete: "
f"req_id={request.request_id}, "
f"gpu_blocks={len(gpu_block_ids)}, "
f"cpu_blocks={len(cpu_block_ids)}"
)
assert self._gpu_block_pool is not None
self._gpu_block_pool.touch([self._gpu_block_pool.blocks[block_id] for block_id in gpu_block_ids])
state.load_transfer_meta = TransferMeta(gpu_block_ids, cpu_block_ids)
logger.info(
"Prepared recompute offload H2D load for request %s: tokens=[%d, %d), blocks=%d.",
request.request_id,
load_start_tokens,
load_end_tokens,
len(gpu_block_ids),
)
return True
def build_connector_meta(
self,
scheduler_output: SchedulerOutput,
) -> RecomputeCPUOffloadMetadata:
store_event = -1
store_gpu, store_cpu, store_req_ids = self._prepare_preempt_store_specs()
if store_gpu:
store_event = self._store_event_counter
self._store_event_counter += 1
self._preempt_store_event_to_blocks[store_event] = TransferMeta(store_gpu, store_cpu)
self._preempt_store_event_to_reqs[store_event] = store_req_ids
for req_id in store_req_ids:
self._preempted_req_states[req_id].store_event = store_event
self._pending_hash_blocks.clear()
load_event = -1
load_gpu: list[int] = []
load_cpu: list[int] = []
load_req_ids: list[str] = []
for req_id, state in self._preempted_req_states.items():
if state.load_transfer_meta is None or state.load_event is not None:
continue
load_gpu.extend(state.load_transfer_meta.gpu_block_ids)
load_cpu.extend(state.load_transfer_meta.cpu_block_ids)
load_req_ids.append(req_id)
if load_req_ids:
load_event = self._load_event_counter
self._load_event_counter += 1
for req_id in load_req_ids:
self._preempted_req_states[req_id].load_event = load_event
self._preempt_load_event_to_reqs[load_event] = load_req_ids
return RecomputeCPUOffloadMetadata(
need_flush=bool(scheduler_output.preempted_req_ids),
preempt_store_event=store_event,
preempt_store_gpu_blocks=store_gpu,
preempt_store_cpu_blocks=store_cpu,
preempt_load_event=load_event,
preempt_load_gpu_blocks=load_gpu,
preempt_load_cpu_blocks=load_cpu,
preempt_load_event_to_reqs=self._preempt_load_event_to_reqs,
)
def update_connector_output(self, connector_output: KVConnectorOutput) -> None:
for req_id in list(connector_output.finished_recving or []):
if req_id in self._preempted_req_states:
self._cleanup_preempt_load_request(req_id)
meta = connector_output.kv_connector_worker_meta
if not isinstance(meta, RecomputeCPUOffloadWorkerMetadata):
return
for event_idx, count in meta.completed_store_events.items():
total = self._store_event_pending_counts.get(event_idx, 0) + count
if total >= self._expected_worker_count:
self._store_event_pending_counts.pop(event_idx, None)
self._process_preempt_store_event(event_idx)
else:
self._store_event_pending_counts[event_idx] = total
def _process_preempt_store_event(self, event_idx: int) -> None:
transfer = self._preempt_store_event_to_blocks.pop(event_idx)
req_ids = self._preempt_store_event_to_reqs.pop(event_idx, [])
for cpu_block_id in transfer.cpu_block_ids:
cpu_block = self.cpu_block_pool.blocks[cpu_block_id]
block_hash = cpu_block.block_hash
if block_hash is None:
continue
cached_block = self.cpu_block_pool.cached_block_hash_to_block.get_one_block(block_hash)
if cached_block is None:
self.cpu_block_pool.cached_block_hash_to_block.insert(block_hash, cpu_block)
elif cached_block.block_id != cpu_block.block_id:
cpu_block.reset_hash()
for req_id in req_ids:
state = self._preempted_req_states.get(req_id)
if state is not None:
state.ready = True
if state.finished:
self._cleanup_preempt_cache_request(req_id)
def has_pending_transfers(self) -> bool:
return bool(
self._store_event_pending_counts
or self._preempt_store_event_to_blocks
or any(
not state.ready or state.load_transfer_meta is not None for state in self._preempted_req_states.values()
)
)
def reset_cache(self) -> bool:
if self.has_pending_transfers():
logger.warning(
"Failed to reset recompute offload cache because transfers or request states are still pending."
)
return False
for req_id in list(self._preempted_req_states):
self._cleanup_preempt_cache_request(req_id)
self._preempt_store_event_to_reqs.clear()
self._preempt_store_event_to_blocks.clear()
self._preempt_load_event_to_reqs.clear()
self._pending_hash_blocks.clear()
return self.cpu_block_pool.reset_prefix_cache()
def request_finished(
self,
request: "Request",
block_ids: list[int],
) -> tuple[bool, dict[str, Any] | None]:
state = self._preempted_req_states.get(request.request_id)
if state is not None and state.load_event is None:
if state.ready:
self._cleanup_preempt_cache_request(request.request_id)
else:
state.finished = True
return False, None
def request_finished_all_groups(
self,
request: "Request",
block_ids: tuple[list[int], ...],
) -> tuple[bool, dict[str, Any] | None]:
return self.request_finished(request, block_ids=[])
def _cleanup_preempt_load_request(self, req_id: str) -> None:
state = self._preempted_req_states.get(req_id)
if state is None:
return
if state.load_event is not None:
reqs = self._preempt_load_event_to_reqs.get(state.load_event)
if reqs is not None:
with contextlib.suppress(ValueError):
reqs.remove(req_id)
if not reqs:
self._preempt_load_event_to_reqs.pop(state.load_event, None)
if state.load_transfer_meta is not None:
assert self._gpu_block_pool is not None
self._gpu_block_pool.free_blocks(
self._gpu_block_pool.blocks[block_id] for block_id in state.load_transfer_meta.gpu_block_ids
)
self._cleanup_preempt_cache_request(req_id)
def _cleanup_preempt_cache_request(self, req_id: str) -> None:
state = self._preempted_req_states.pop(req_id, None)
if state is None:
return
self.cpu_block_pool.free_blocks(
self.cpu_block_pool.blocks[block_id]
for group_cpu_ids in state.cpu_block_ids
for block_id in group_cpu_ids
if block_id > 0
)
def take_events(self) -> Iterable[KVCacheEvent]:
return self.cpu_block_pool.take_events()

View File

@@ -0,0 +1,52 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Metadata for RecomputeCPUOffloadConnector."""
from dataclasses import dataclass, field
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorMetadata,
KVConnectorWorkerMetadata,
)
INVALID_JOB_ID = -1
@dataclass
class RecomputeCPUOffloadMetadata(KVConnectorMetadata):
"""Recompute offload transfers passed from scheduler to worker."""
# Whether any requests were preempted this step and need flush pending transfers.
need_flush: bool = False
# Store blocks of newly preempted requests before their GPU blocks can
# be reused. The list may include a final partial block without a hash.
preempt_store_event: int = INVALID_JOB_ID
preempt_store_gpu_blocks: list[int] = field(default_factory=list)
preempt_store_cpu_blocks: list[int] = field(default_factory=list)
# Preemption load event. Used when a previously preempted request resumes.
preempt_load_event: int = INVALID_JOB_ID
preempt_load_gpu_blocks: list[int] = field(default_factory=list)
preempt_load_cpu_blocks: list[int] = field(default_factory=list)
preempt_load_event_to_reqs: dict[int, list[str]] = field(default_factory=dict)
@dataclass
class RecomputeCPUOffloadWorkerMetadata(KVConnectorWorkerMetadata):
"""Worker -> Scheduler metadata for completed store events.
Each worker reports {event_idx: 1} for newly completed stores.
``aggregate()`` sums counts across workers within a step.
The scheduler-side manager accumulates across steps and processes
a store completion only when count reaches ``world_size``.
"""
completed_store_events: dict[int, int]
def aggregate(self, other: "KVConnectorWorkerMetadata") -> "KVConnectorWorkerMetadata":
assert isinstance(other, RecomputeCPUOffloadWorkerMetadata)
merged = dict(self.completed_store_events)
for k, v in other.completed_store_events.items():
merged[k] = merged.get(k, 0) + v
return RecomputeCPUOffloadWorkerMetadata(completed_store_events=merged)

View File

@@ -0,0 +1,246 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""RecomputeCPUOffloadConnector: minimal CPU KV cache offloading."""
from collections.abc import Iterable
from typing import TYPE_CHECKING, Any
import torch
from vllm.config import VllmConfig
from vllm.distributed.kv_events import KVCacheEvent
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorBase_V1,
KVConnectorMetadata,
KVConnectorRole,
SupportsHMA,
)
from vllm.logger import logger
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.outputs import KVConnectorOutput
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.manager import (
RecomputeCPUOffloadScheduler,
)
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.metadata import (
RecomputeCPUOffloadMetadata,
)
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.worker import (
RecomputeCPUOffloadWorker,
)
if TYPE_CHECKING:
from vllm.forward_context import ForwardContext
from vllm.v1.attention.backend import AttentionMetadata
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.request import Request
# Default CPU capacity: 8 GB
DEFAULT_CPU_CAPACITY_BYTES = 8 * (1024**3)
class RecomputeCPUOffloadConnectorV1(KVConnectorBase_V1, SupportsHMA):
"""CPU KV cache preservation for recompute-preempted requests."""
def __init__(
self,
vllm_config: VllmConfig,
role: KVConnectorRole,
kv_cache_config: "KVCacheConfig | None" = None,
):
super().__init__(vllm_config, role, kv_cache_config)
extra_config = self._kv_transfer_config.kv_connector_extra_config or {}
cpu_capacity_bytes = int(extra_config.get("cpu_bytes_to_use", DEFAULT_CPU_CAPACITY_BYTES))
enable_offload_prefix_caching = extra_config.get("enable_offload_prefix_caching", False)
if not isinstance(enable_offload_prefix_caching, bool):
raise ValueError(f"enable_offload_prefix_caching must be a boolean, got {enable_offload_prefix_caching!r}")
world_size = vllm_config.parallel_config.world_size
cpu_capacity_per_rank = cpu_capacity_bytes // world_size
if "cpu_bytes_to_use_per_rank" in extra_config:
explicit = int(extra_config["cpu_bytes_to_use_per_rank"])
if explicit != cpu_capacity_per_rank:
logger.warning(
"cpu_bytes_to_use_per_rank (%.2f GB) != "
"cpu_bytes_to_use/world_size (%.2f GB). Using per-rank value.",
explicit / (1024**3),
cpu_capacity_per_rank / (1024**3),
)
cpu_capacity_per_rank = explicit
self.scheduler_manager: RecomputeCPUOffloadScheduler | None = None
self.worker_handler: RecomputeCPUOffloadWorker | None = None
logger.info(
"RecomputeCPUOffloadConnector: role=%s, per_rank=%.2f GB, world_size=%d, offload_prefix_caching=%s",
role.name,
cpu_capacity_per_rank / (1024**3),
world_size,
enable_offload_prefix_caching,
)
if role == KVConnectorRole.SCHEDULER:
self.scheduler_manager = RecomputeCPUOffloadScheduler(
vllm_config,
kv_cache_config,
cpu_capacity_per_rank,
enable_offload_prefix_caching,
)
elif role == KVConnectorRole.WORKER:
self.worker_handler = RecomputeCPUOffloadWorker(
vllm_config,
kv_cache_config,
cpu_capacity_per_rank,
)
# --- Worker-side methods ---
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]) -> None:
if self.worker_handler is not None:
self.worker_handler.register_kv_caches(kv_caches)
def bind_connector_metadata(
self,
connector_metadata: KVConnectorMetadata,
) -> None:
super().bind_connector_metadata(connector_metadata)
if self.worker_handler is not None:
assert isinstance(connector_metadata, RecomputeCPUOffloadMetadata)
self.worker_handler.bind_connector_metadata(connector_metadata)
def clear_connector_metadata(self) -> None:
super().clear_connector_metadata()
if self.worker_handler is not None:
self.worker_handler.clear_connector_metadata()
def handle_preemptions(self, kv_connector_metadata: KVConnectorMetadata) -> None:
if self.worker_handler is not None:
assert isinstance(kv_connector_metadata, RecomputeCPUOffloadMetadata)
self.worker_handler.handle_preemptions(kv_connector_metadata)
def start_load_kv(self, forward_context: "ForwardContext", **kwargs: Any) -> None:
if self.worker_handler is not None:
self.worker_handler.start_load_kv()
def wait_for_layer_load(self, layer_name: str) -> None:
if self.worker_handler is not None:
self.worker_handler.wait_for_layer_load()
def save_kv_layer(
self,
layer_name: str,
kv_layer: torch.Tensor,
attn_metadata: "AttentionMetadata",
**kwargs: Any,
) -> None:
pass
def wait_for_save(self) -> None:
pass
def get_finished(
self,
finished_req_ids: set[str],
) -> tuple[set[str] | None, set[str] | None]:
if self.worker_handler is not None:
return self.worker_handler.get_finished(finished_req_ids)
return None, None
def build_connector_worker_meta(self):
if self.worker_handler is not None:
return self.worker_handler.build_connector_worker_meta()
return None
# --- Scheduler-side methods ---
# NOTE: New API only for RecomputeCPUOffloadConnector.
def bind_gpu_block_pool(self, gpu_block_pool: "BlockPool") -> None:
if self.scheduler_manager is not None:
self.scheduler_manager.bind_gpu_block_pool(gpu_block_pool)
def get_num_new_matched_tokens(
self,
request: "Request",
num_computed_tokens: int,
) -> tuple[int | None, bool]:
if self.scheduler_manager is not None:
return self.scheduler_manager.get_num_new_matched_tokens(request, num_computed_tokens)
return 0, False
def update_state_after_alloc(
self,
request: "Request",
blocks: "KVCacheBlocks",
num_external_tokens: int,
) -> None:
if self.scheduler_manager is not None:
self.scheduler_manager.update_state_after_alloc(request, blocks, num_external_tokens)
def update_state_before_preempt(
self,
request: "Request",
block_ids: tuple[list[int], ...],
num_computed_tokens: int,
) -> bool:
if self.scheduler_manager is not None:
return self.scheduler_manager.update_state_before_preempt(
request,
block_ids,
num_computed_tokens,
)
return False
def build_connector_meta(
self,
scheduler_output: SchedulerOutput,
) -> KVConnectorMetadata:
if self.scheduler_manager is not None:
return self.scheduler_manager.build_connector_meta(scheduler_output)
return RecomputeCPUOffloadMetadata()
def update_connector_output(
self,
connector_output: KVConnectorOutput,
) -> None:
if self.scheduler_manager is not None:
self.scheduler_manager.update_connector_output(connector_output)
def request_finished(
self,
request: "Request",
block_ids: list[int],
) -> tuple[bool, dict[str, Any] | None]:
if self.scheduler_manager is not None:
return self.scheduler_manager.request_finished(request, block_ids)
return False, None
def request_finished_all_groups(
self,
request: "Request",
block_ids: tuple[list[int], ...],
) -> tuple[bool, dict[str, Any] | None]:
if self.scheduler_manager is not None:
return self.scheduler_manager.request_finished_all_groups(request, block_ids)
return False, None
# NOTE: New API only for RecomputeCPUOffloadConnector.
def has_pending_transfers(self) -> bool:
if self.scheduler_manager is not None:
return self.scheduler_manager.has_pending_transfers()
return False
def has_preempted_request(self, req_id: str) -> bool:
if self.scheduler_manager is not None:
return self.scheduler_manager.has_preempted_request(req_id)
return False
def take_events(self) -> Iterable[KVCacheEvent]:
if self.scheduler_manager is not None:
return self.scheduler_manager.take_events()
return []
def reset_cache(self) -> bool | None:
if self.scheduler_manager is not None:
return self.scheduler_manager.reset_cache()
return None

View File

@@ -0,0 +1,319 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Worker-side handler for Ascend RecomputeCPUOffloadConnector."""
from typing import TYPE_CHECKING
import torch
from vllm.config import VllmConfig
from vllm.logger import logger
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.metadata import (
RecomputeCPUOffloadMetadata,
RecomputeCPUOffloadWorkerMetadata,
)
if TYPE_CHECKING:
from vllm.v1.kv_cache_interface import KVCacheConfig
class RecomputeCPUOffloadWorker:
"""Worker-side handler for recompute CPU/NPU KV cache transfers."""
def __init__(
self,
vllm_config: VllmConfig,
kv_cache_config: "KVCacheConfig | None",
cpu_capacity_bytes: int,
):
self.vllm_config = vllm_config
self.kv_cache_config = kv_cache_config
self.cpu_capacity_bytes = cpu_capacity_bytes
self.gpu_kv_caches: dict[str, torch.Tensor] | None = None
self.cpu_kv_caches: dict[str, torch.Tensor] | None = None
self.device: torch.device | None = None
self.num_cpu_blocks: int = 0
self.load_stream: torch.npu.Stream | None = None
self.store_stream: torch.npu.Stream | None = None
self._load_events: list[tuple[int, torch.npu.Event]] = []
self._load_hwm: int = -1
self._connector_metadata: RecomputeCPUOffloadMetadata | None = None
self._pending_load_event_indices: set[int] = set()
self._submitted_load_event_indices: set[int] = set()
self._completed_store_events: dict[int, int] = {}
self._load_stream_waited = False
def register_kv_caches(
self,
kv_caches: dict[str, torch.Tensor],
) -> None:
"""Register KV caches and initialize CPU/NPU transfer resources."""
if not kv_caches:
logger.warning("No KV caches to offload.")
return
any_tensor = next(iter(kv_caches.values()))
if isinstance(any_tensor, (tuple, list)):
any_tensor = any_tensor[0]
self.device = any_tensor.device
assert self.kv_cache_config is not None
self.num_gpu_blocks = self.kv_cache_config.num_blocks
self.block_size_scale = {}
scheduler_gpu_kv_cache_tensors = []
for t in self.kv_cache_config.kv_cache_tensors:
if t.shared_by:
scheduler_gpu_kv_cache_tensors.append(t)
scheduler_gpu_total_bytes = sum(t.size for t in scheduler_gpu_kv_cache_tensors)
scheduler_num_cpu_blocks = max(1, self.num_gpu_blocks * self.cpu_capacity_bytes // scheduler_gpu_total_bytes)
unique_gpu_caches: dict[str, torch.Tensor] = {}
register_cache_ptrs = []
for layer_name, layer_tensor in kv_caches.items():
if isinstance(layer_tensor, (tuple, list)):
for idx, single_tensor in enumerate(layer_tensor):
if single_tensor.data_ptr() not in register_cache_ptrs:
unique_gpu_caches[f"{layer_name}.{idx}"] = single_tensor.view(single_tensor.shape[0], -1)
register_cache_ptrs.append(single_tensor.data_ptr())
self.block_size_scale[f"{layer_name}.{idx}"] = single_tensor.shape[0] // self.num_gpu_blocks
else:
if layer_tensor.data_ptr() not in register_cache_ptrs:
unique_gpu_caches[layer_name] = layer_tensor.view(layer_tensor.shape[0], -1)
register_cache_ptrs.append(layer_tensor.data_ptr())
self.block_size_scale[layer_name] = layer_tensor.shape[0] // self.num_gpu_blocks
per_tensor_bytes_per_block = [tensor.shape[-1] * tensor.element_size() for tensor in unique_gpu_caches.values()]
total_bytes_per_block = sum(per_tensor_bytes_per_block)
self.num_cpu_blocks = max(1, self.cpu_capacity_bytes // total_bytes_per_block)
if self.num_cpu_blocks != scheduler_num_cpu_blocks:
self.num_cpu_blocks = scheduler_num_cpu_blocks
logger.warning(
"RecomputeCPUOffloadScheduler has different num_blocks: %d,"
"worker-side num_block is set to %d to align with scheduler.",
scheduler_num_cpu_blocks,
scheduler_num_cpu_blocks,
)
self.gpu_kv_caches = unique_gpu_caches
self.cpu_kv_caches = {}
for name, gpu_tensor in unique_gpu_caches.items():
tensor_block_size_scale = self.block_size_scale[name]
cpu_shape = (self.num_cpu_blocks * tensor_block_size_scale,) + gpu_tensor.shape[1:]
self.cpu_kv_caches[name] = torch.zeros(
cpu_shape,
dtype=gpu_tensor.dtype,
pin_memory=True,
device="cpu",
)
self.load_stream = torch.npu.Stream()
self.store_stream = torch.npu.Stream()
logger.info(
"RecomputeCPUOffloadWorker scaffold registered %d unique KV tensors, allocating %d CPU blocks (%.2f GB).",
len(unique_gpu_caches),
self.num_cpu_blocks,
(self.num_cpu_blocks * total_bytes_per_block) / (1024**3),
)
def bind_connector_metadata(self, metadata: RecomputeCPUOffloadMetadata) -> None:
self._connector_metadata = metadata
self._load_stream_waited = False
if metadata.preempt_load_event >= 0:
self._pending_load_event_indices.add(metadata.preempt_load_event)
def clear_connector_metadata(self) -> None:
"""Clear metadata after the model runner finishes the current step."""
self._connector_metadata = None
def handle_preemptions(
self,
kv_connector_metadata: RecomputeCPUOffloadMetadata,
) -> None:
"""Save preempted blocks before input preparation can overwrite them."""
if kv_connector_metadata.need_flush:
self._flush_and_sync_all()
# The scheduler may immediately reuse preempted block IDs in this same
# step. This blocking D2H must therefore run before _update_states()
# processes new_block_ids_to_zero and before model forward writes KV.
self._submit_transfer(
kv_connector_metadata.preempt_store_gpu_blocks,
kv_connector_metadata.preempt_store_cpu_blocks,
kv_connector_metadata.preempt_store_event,
is_store=True,
sync=True,
)
def start_load_kv(self) -> None:
"""Submit pre-forward recompute H2D transfers."""
metadata = self._connector_metadata
if metadata is None:
return
self._submit_transfer(
metadata.preempt_load_cpu_blocks,
metadata.preempt_load_gpu_blocks,
metadata.preempt_load_event,
is_store=False,
sync=True,
)
def wait_for_layer_load(self) -> None:
"""Make the current forward stream wait for the recompute H2D copy."""
if self._load_stream_waited or self.load_stream is None:
return
metadata = self._connector_metadata
if metadata is None or metadata.preempt_load_event < 0:
return
torch.npu.current_stream().wait_stream(self.load_stream)
self._load_stream_waited = True
def _flush_and_sync_all(self) -> None:
"""Synchronize all in-flight transfer events."""
for event_idx, event in self._load_events:
event.synchronize()
self._load_hwm = event_idx
self._load_events.clear()
self._submitted_load_event_indices.clear()
def _poll_load_events(self) -> int:
"""Return the highest completed H2D event index."""
events = self._load_events
hwm = self._load_hwm
while events:
event_idx, event = events[0]
if not event.query():
break
hwm = event_idx
events.pop(0)
self._load_hwm = hwm
return hwm
def _submit_transfer(
self,
src_block_ids: list[int],
dst_block_ids: list[int],
event_idx: int,
is_store: bool,
sync: bool = False,
) -> None:
"""Submit a CPU<->NPU block copy and record a completion event."""
if event_idx < 0:
return
if not is_store and event_idx in self._submitted_load_event_indices:
return
if not is_store:
self._submitted_load_event_indices.add(event_idx)
if not src_block_ids:
if is_store:
self._completed_store_events[event_idx] = 1
else:
self._load_hwm = max(self._load_hwm, event_idx)
return
assert len(src_block_ids) == len(dst_block_ids)
assert self.gpu_kv_caches is not None
assert self.cpu_kv_caches is not None
stream = self.store_stream if is_store else self.load_stream
assert stream is not None
torch.npu.synchronize()
with torch.npu.stream(stream):
for src_block_id, dst_block_id in zip(src_block_ids, dst_block_ids):
for name, gpu_tensor in self.gpu_kv_caches.items():
cpu_tensor = self.cpu_kv_caches[name]
tensor_block_size_scale = self.block_size_scale[name]
if is_store:
# TODO: Replace this D2H torch copy with the NPU copy
# backend dedicated kernel.
if tensor_block_size_scale > 1:
cpu_tensor[
dst_block_id * tensor_block_size_scale : (dst_block_id + 1) * tensor_block_size_scale
].copy_(
gpu_tensor[
src_block_id * tensor_block_size_scale : (src_block_id + 1)
* tensor_block_size_scale
],
non_blocking=True,
)
else:
cpu_tensor[dst_block_id].copy_(
gpu_tensor[src_block_id],
non_blocking=True,
)
else:
# TODO: Replace this H2D torch copy with the NPU copy
# backend dedicated kernel.
if tensor_block_size_scale > 1:
gpu_tensor[
dst_block_id * tensor_block_size_scale : (dst_block_id + 1) * tensor_block_size_scale
].copy_(
cpu_tensor[
src_block_id * tensor_block_size_scale : (src_block_id + 1)
* tensor_block_size_scale
],
non_blocking=True,
)
else:
gpu_tensor[dst_block_id].copy_(
cpu_tensor[src_block_id],
non_blocking=True,
)
event = torch.npu.Event()
event.record(stream)
if sync:
event.synchronize()
if is_store:
self._completed_store_events[event_idx] = 1
else:
self._load_hwm = max(self._load_hwm, event_idx)
return
assert not is_store
self._load_events.append((event_idx, event))
def get_finished(
self,
finished_req_ids: set[str],
) -> tuple[set[str] | None, set[str] | None]:
"""Poll recompute transfers and report completed request restores."""
metadata = self._connector_metadata
if metadata is None:
return None, None
finished_recving: set[str] = set()
if self._pending_load_event_indices:
load_hwm = self._poll_load_events()
completed_loads = [event_idx for event_idx in self._pending_load_event_indices if event_idx <= load_hwm]
for event_idx in completed_loads:
self._pending_load_event_indices.discard(event_idx)
self._submitted_load_event_indices.discard(event_idx)
finished_recving.update(metadata.preempt_load_event_to_reqs.get(event_idx, []))
return None, finished_recving or None
def build_connector_worker_meta(self) -> RecomputeCPUOffloadWorkerMetadata | None:
"""Return completed store events since the previous call.
The scheduler aggregates this metadata across workers/ranks. A store
event becomes available to recompute requests only after all expected
workers have reported completion.
"""
if not self._completed_store_events:
return None
meta = RecomputeCPUOffloadWorkerMetadata(
completed_store_events=self._completed_store_events,
)
self._completed_store_events = {}
return meta

View File

@@ -0,0 +1,60 @@
"""Ascend NPU adaptation of vLLM's ``SimpleCPUOffloadConnector``.
The scheduler-side ``SimpleCPUOffloadScheduler`` is platform-agnostic
and reused as-is from upstream vLLM. The Ascend variant only swaps the
worker-side handler with an NPU-native implementation that uses
``aclrtMemcpyBatchAsync`` and ``torch.npu`` streams/events.
"""
from typing import TYPE_CHECKING
from vllm.config import VllmConfig
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
from vllm.distributed.kv_transfer.kv_connector.v1.simple_cpu_offload_connector import ( # noqa: E501
SimpleCPUOffloadConnector,
)
from vllm.logger import logger
from vllm_ascend.simple_kv_offload.worker import SimpleCPUOffloadNPUWorker
if TYPE_CHECKING:
from vllm.v1.kv_cache_interface import KVCacheConfig
class AscendSimpleCPUOffloadConnector(SimpleCPUOffloadConnector):
"""NPU-flavored ``SimpleCPUOffloadConnector``.
Inherits the full scheduler/worker plumbing from upstream and only
replaces the CUDA worker handler with the NPU one. All other public
APIs (``register_kv_caches``, ``bind_connector_metadata``,
``get_finished``, ``handle_preemptions``, every scheduler-side
method, etc.) are inherited verbatim — they all route through
``self.worker_handler`` / ``self.scheduler_manager``.
Why post-init swap (instead of skipping ``super().__init__``):
``SimpleCPUOffloadWorker.__init__`` and ``DmaCopyBackend.__init__``
only assign ``None``/empty-field defaults — no CUDA resource is
allocated until ``register_kv_caches`` runs. So letting the parent
construct a transient CUDA worker and then replacing it costs
nothing and keeps us free of duplicated configuration parsing.
"""
def __init__(
self,
vllm_config: VllmConfig,
role: KVConnectorRole,
kv_cache_config: "KVCacheConfig | None" = None,
) -> None:
super().__init__(vllm_config, role, kv_cache_config)
# If prefix caching is disabled, the parent leaves both handlers
# as None and the connector is a no-op — nothing to swap.
if role == KVConnectorRole.WORKER and self.worker_handler is not None:
cpu_capacity = self.worker_handler.cpu_capacity_bytes
self.worker_handler: SimpleCPUOffloadNPUWorker = SimpleCPUOffloadNPUWorker(
vllm_config, kv_cache_config, cpu_capacity
)
logger.info(
"AscendSimpleCPUOffloadConnector: swapped CUDA worker for NPU worker (per_rank=%.2f GB)",
cpu_capacity / (1024**3),
)

View File

@@ -0,0 +1,324 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Iterable
from typing import TYPE_CHECKING, Any, Optional
import torch
from ucm.integration.vllm.ucm_connector import UCMConnector
from vllm.config import VllmConfig
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
CopyBlocksOp,
KVConnectorBase_V1,
KVConnectorHandshakeMetadata,
KVConnectorMetadata,
KVConnectorRole,
KVConnectorWorkerMetadata,
SupportsHMA,
)
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.outputs import KVConnectorOutput
# isort: off
if TYPE_CHECKING:
from vllm.distributed.kv_events import KVCacheEvent, KVConnectorKVEvents
from vllm.distributed.kv_transfer.kv_connector.v1.metrics import (
KVConnectorPromMetrics,
KVConnectorStats,
PromMetric,
PromMetricT,
)
from vllm.forward_context import ForwardContext
from vllm.v1.attention.backend import AttentionMetadata
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.request import Request
# isort: on
class UCMConnectorV1(KVConnectorBase_V1, SupportsHMA):
def __init__(
self,
vllm_config: "VllmConfig",
role: KVConnectorRole,
kv_cache_config: "KVCacheConfig",
):
super().__init__(vllm_config=vllm_config, role=role, kv_cache_config=kv_cache_config)
assert vllm_config.kv_transfer_config is not None
ImplCls = UCMConnector
self._ucm_engine = ImplCls(vllm_config, role, kv_cache_config)
def _get_ucm_delegate_for(self, name: str) -> Any:
"""Return the UCM object that owns a reserved connector interface.
The imported UCMConnector is itself a thin dispatcher. Some reserved
KVConnectorBase_V1 hooks are not redeclared on that dispatcher yet, so
the base class no-op can otherwise hide the real inner connector hook.
Prefer an explicit method/property on the dispatcher; otherwise fall
through to the selected inner connector when present.
"""
if name not in type(self._ucm_engine).__dict__:
inner_connector = getattr(self._ucm_engine, "connector", None)
if inner_connector is not None:
return inner_connector
return self._ucm_engine
def _call_ucm_reserved_hook(self, name: str, *args: Any, **kwargs: Any) -> Any:
hook = getattr(self._get_ucm_delegate_for(name), name, None)
if callable(hook):
return hook(*args, **kwargs)
return None
# ==============================
# Worker-side methods
# ==============================
def shutdown(self) -> None:
self._call_ucm_reserved_hook("shutdown")
def has_connector_metadata(self) -> bool:
"""Check whether the connector metadata is currently set.
Returns:
bool: True if connector metadata exists, False otherwise.
"""
return self._ucm_engine.has_connector_metadata()
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]) -> None:
"""
Initialize with the KV caches. Useful for pre-registering the
KV Caches in the KVConnector (e.g. for NIXL).
Args:
kv_caches: A dictionary mapping layer names to KV cache tensors.
"""
self._ucm_engine.register_kv_caches(kv_caches)
def set_host_xfer_buffer_ops(self, copy_operation: CopyBlocksOp) -> None:
self._call_ucm_reserved_hook("set_host_xfer_buffer_ops", copy_operation)
def handle_preemptions(self, kv_connector_metadata: KVConnectorMetadata) -> None:
self._call_ucm_reserved_hook("handle_preemptions", kv_connector_metadata)
def start_load_kv(self, forward_context: "ForwardContext", **kwargs: Any) -> None:
"""
Start loading the KV cache from the connector to vLLM's paged
KV buffer. This is called from the forward context before the
forward pass to enable async loading during model execution.
Args:
forward_context (ForwardContext): the forward context.
**kwargs: additional arguments for the load operation
Note:
The number of elements in kv_caches and layer_names should be
the same.
"""
self._ucm_engine.start_load_kv(forward_context, **kwargs)
def wait_for_layer_load(self, layer_name: str) -> None:
"""
Block until the KV for a specific layer is loaded into vLLM's
paged buffer. This is called from within attention layer to ensure
async copying from start_load_kv is complete.
This interface will be useful for layer-by-layer pipelining.
Args:
layer_name: the name of that layer
"""
self._ucm_engine.wait_for_layer_load(layer_name)
def save_kv_layer(
self,
layer_name: str,
kv_layer: torch.Tensor,
attn_metadata: "AttentionMetadata",
**kwargs: Any,
) -> None:
"""
Start saving the a layer of KV cache from vLLM's paged buffer
to the connector. This is called from within attention layer to
enable async copying during execution.
Args:
layer_name (str): the name of the layer.
kv_layer (torch.Tensor): the paged KV buffer of the current
layer in vLLM.
attn_metadata (AttentionMetadata): the attention metadata.
**kwargs: additional arguments for the save operation.
"""
self._ucm_engine.save_kv_layer(layer_name, kv_layer, attn_metadata, **kwargs)
def wait_for_save(self) -> None:
"""
Block until all the save operations is done. This is called
as the forward context exits to ensure that the async saving
from save_kv_layer is complete before finishing the forward.
This prevents overwrites of paged KV buffer before saving done.
"""
self._ucm_engine.wait_for_save()
def clear_connector_metadata(self) -> None:
"""Clear the connector metadata.
This function should be called by the model runner every time
after the model execution.
"""
self._ucm_engine.clear_connector_metadata()
def bind_connector_metadata(self, connector_metadata: KVConnectorMetadata) -> None:
"""Set the connector metadata from the scheduler.
This function should be called by the model runner every time
before the model execution. The metadata will be used for runtime
KV cache loading and saving.
Args:
connector_metadata (dict): the connector metadata.
"""
self._ucm_engine.bind_connector_metadata(connector_metadata)
def get_block_ids_with_load_errors(self) -> set[int]:
"""
Get the set of block IDs that failed to load.
Returns:
Set of block IDs that encountered load errors.
Empty set if no load errors occurred.
"""
return self._ucm_engine.get_block_ids_with_load_errors()
# ==============================
# Scheduler-side methods
# ==============================
def get_num_new_matched_tokens(
self,
request: "Request",
num_computed_tokens: int,
) -> tuple[int | None, bool]:
"""
Get number of new tokens that can be loaded from the
external KV cache beyond the num_computed_tokens.
Args:
request (Request): the request object.
num_computed_tokens (int): the number of locally
computed tokens for this request
Returns:
the number of tokens that can be loaded from the
external KV cache beyond what is already computed.
"""
return self._ucm_engine.get_num_new_matched_tokens(request, num_computed_tokens)
def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", num_external_tokens: int) -> None:
"""
Update KVConnector state after block allocation.
"""
self._ucm_engine.update_state_after_alloc(request, blocks, num_external_tokens)
def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnectorMetadata:
"""
Build the connector metadata for this step.
This function should NOT modify fields in the scheduler_output.
Also, calling this function will reset the state of the connector.
Args:
scheduler_output (SchedulerOutput): the scheduler output object.
"""
return self._ucm_engine.build_connector_meta(scheduler_output)
def request_finished(
self,
request: "Request",
block_ids: list[int],
) -> tuple[bool, dict[str, Any] | None]:
"""
Called when a request has finished, before its blocks are freed.
Returns:
True if the request is being saved/sent asynchronously and blocks
should not be freed until the request_id is returned from
get_finished().
Optional KVTransferParams to be included in the request outputs
returned by the engine.
"""
return self._ucm_engine.request_finished(request, block_ids)
def request_finished_all_groups(
self,
request: "Request",
block_ids: tuple[list[int], ...],
) -> tuple[bool, dict[str, Any] | None]:
return self._ucm_engine.request_finished_all_groups(request, block_ids)
def get_finished(
self,
finished_req_ids: set[str],
) -> tuple[set[str] | None, set[str] | None]:
return self._ucm_engine.get_finished(finished_req_ids)
def build_connector_worker_meta(self) -> KVConnectorWorkerMetadata | None:
return self._call_ucm_reserved_hook("build_connector_worker_meta")
def update_connector_output(self, connector_output: KVConnectorOutput) -> None:
self._ucm_engine.update_connector_output(connector_output)
def take_events(self) -> Iterable["KVCacheEvent"]:
events = self._call_ucm_reserved_hook("take_events")
return () if events is None else events
def get_kv_connector_stats(self) -> Optional["KVConnectorStats"]:
return self._call_ucm_reserved_hook("get_kv_connector_stats")
def get_kv_connector_kv_cache_events(self) -> Optional["KVConnectorKVEvents"]:
return self._call_ucm_reserved_hook("get_kv_connector_kv_cache_events")
def get_handshake_metadata(self) -> KVConnectorHandshakeMetadata | None:
return self._call_ucm_reserved_hook("get_handshake_metadata")
def set_xfer_handshake_metadata(self, metadata: dict[int, KVConnectorHandshakeMetadata]) -> None:
self._call_ucm_reserved_hook("set_xfer_handshake_metadata", metadata)
def get_finished_count(self) -> int | None:
return self._call_ucm_reserved_hook("get_finished_count")
def reset_cache(self) -> bool | None:
return self._call_ucm_reserved_hook("reset_cache")
# ==============================
# Metrics & Stats
# ==============================
@classmethod
def build_kv_connector_stats(cls, data: dict[str, Any] | None = None) -> Optional["KVConnectorStats"]:
"""
KVConnectorStats resolution method. This method allows dynamically
registered connectors to return their own KVConnectorStats object,
which can implement custom aggregation logic on the data dict.
"""
return UCMConnector.build_kv_connector_stats(data)
@classmethod
def build_prom_metrics(
cls,
vllm_config: "VllmConfig",
metric_types: dict[type["PromMetric"], type["PromMetricT"]],
labelnames: list[str],
per_engine_labelvalues: dict[int, list[object]],
) -> Optional["KVConnectorPromMetrics"]:
"""
Create a KVConnectorPromMetrics subclass which should register
per-connector Prometheus metrics and implement observe() to
expose connector transfer stats via Prometheus.
This implementation forwards the call to the underlying
UCMConnector engine.
"""
return UCMConnector.build_prom_metrics(
vllm_config,
metric_types,
labelnames,
per_engine_labelvalues,
)