103 lines
4.3 KiB
Python
103 lines
4.3 KiB
Python
from typing import TYPE_CHECKING, Any, cast
|
|
|
|
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
|
|
KVConnectorRole,
|
|
SupportsHMA,
|
|
supports_hma,
|
|
)
|
|
from vllm.distributed.kv_transfer.kv_connector.v1.multi_connector import MultiConnector
|
|
|
|
from vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector import MooncakeLayerwiseConnector
|
|
|
|
if TYPE_CHECKING:
|
|
from vllm.config import VllmConfig
|
|
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
|
|
from vllm.v1.kv_cache_interface import KVCacheConfig
|
|
from vllm.v1.request import Request
|
|
|
|
|
|
class AscendMultiConnector(MultiConnector, 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,
|
|
)
|
|
|
|
self._all_support_hma = all(supports_hma(c) for c in self._connectors)
|
|
assert vllm_config.scheduler_config.disable_hybrid_kv_cache_manager or self._all_support_hma, (
|
|
"HMA should not be enabled unless all sub-connectors support it"
|
|
)
|
|
|
|
def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", num_external_tokens: int):
|
|
chosen_connector = self._requests_to_connector.get(request.request_id, -1)
|
|
empty_blocks = blocks.new_empty()
|
|
for i, c in enumerate(self._connectors):
|
|
if i == chosen_connector or isinstance(c, MooncakeLayerwiseConnector):
|
|
# Forward call to the chosen connector (if any).
|
|
c.update_state_after_alloc(request, blocks, num_external_tokens)
|
|
else:
|
|
# Call with empty blocks for other connectors.
|
|
c.update_state_after_alloc(request, empty_blocks, 0)
|
|
|
|
def get_num_new_matched_tokens(
|
|
self,
|
|
request: "Request",
|
|
num_computed_tokens: int,
|
|
) -> tuple[int | None, bool]:
|
|
# Recompute offload may contain an unhashed partial block that other
|
|
# prefix-cache connectors cannot restore. Give its request state
|
|
# priority regardless of connector ordering.
|
|
for i, connector in enumerate(self._connectors):
|
|
has_preempted_request = getattr(connector, "has_preempted_request", None)
|
|
if has_preempted_request is None or not has_preempted_request(request.request_id):
|
|
continue
|
|
tokens, load_async = connector.get_num_new_matched_tokens(request, num_computed_tokens)
|
|
if tokens is None:
|
|
return None, False
|
|
if tokens > 0:
|
|
self._requests_to_connector[request.request_id] = i
|
|
return tokens, load_async
|
|
break
|
|
|
|
return super().get_num_new_matched_tokens(request, num_computed_tokens)
|
|
|
|
def update_state_before_preempt(
|
|
self,
|
|
request: "Request",
|
|
block_ids: tuple[list[int], ...],
|
|
num_computed_tokens: int,
|
|
) -> bool:
|
|
offloaded = False
|
|
for c in self._connectors:
|
|
hook = getattr(c, "update_state_before_preempt", None)
|
|
if hook is not None:
|
|
offloaded = bool(hook(request, block_ids, num_computed_tokens)) or offloaded
|
|
return offloaded
|
|
|
|
def request_finished_all_groups(
|
|
self,
|
|
request: "Request",
|
|
block_ids: tuple[list[int], ...],
|
|
) -> tuple[bool, dict[str, Any] | None]:
|
|
if not self._all_support_hma:
|
|
assert len(block_ids) == 1, "HMA with multiple kv_cache_groups requires all sub-connectors to support HMA"
|
|
return super().request_finished(request, block_ids[0])
|
|
|
|
async_saves = 0
|
|
kv_txfer_params = None
|
|
for c in self._connectors:
|
|
async_save, txfer_params = cast(SupportsHMA, c).request_finished_all_groups(request, block_ids)
|
|
if async_save:
|
|
async_saves += 1
|
|
if txfer_params is not None:
|
|
if kv_txfer_params is not None:
|
|
raise RuntimeError("Only one connector can produce KV transfer params")
|
|
kv_txfer_params = txfer_params
|
|
if async_saves > 1:
|
|
self._extra_async_saves[request.request_id] = async_saves - 1
|
|
|
|
self._requests_to_connector.pop(request.request_id, None)
|
|
|
|
return async_saves > 0, kv_txfer_params
|