@@ -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)
|
||||
@@ -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",
|
||||
},
|
||||
}
|
||||
@@ -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
|
||||
@@ -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.")
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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"]
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
324
vllm_ascend/distributed/kv_transfer/kv_pool/ucm_connector.py
Normal file
324
vllm_ascend/distributed/kv_transfer/kv_pool/ucm_connector.py
Normal 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,
|
||||
)
|
||||
Reference in New Issue
Block a user