diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/base.py b/vllm/distributed/kv_transfer/kv_connector/v1/base.py index ef143cba7fb5..686f3ad8ebac 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/base.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/base.py @@ -43,6 +43,7 @@ import enum from abc import ABC, abstractmethod from collections.abc import Callable, Iterable +from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Literal import torch @@ -62,7 +63,7 @@ PromMetricT, ) from vllm.forward_context import ForwardContext - from vllm.v1.core.kv_cache_manager import KVCacheBlocks + from vllm.v1.core.kv_cache_manager import KVCacheBlocks, KVCacheManager from vllm.v1.kv_cache_interface import KVCacheConfig from vllm.v1.request import Request @@ -146,6 +147,17 @@ class KVConnectorMetadata(ABC): # noqa: B024 pass +@dataclass(frozen=True) +class SchedulerState: + """ + State of the scheduler that the connector can access, scheduler-side. + This dataclass ensures read-only access to scheduler state, while enabling + expansion of the scheduler state in the future. + """ + + kv_cache_manager: "KVCacheManager" + + class KVConnectorWorkerMetadata(ABC): """ Abstract Metadata used to communicate back @@ -446,6 +458,17 @@ def build_connector_worker_meta(self) -> KVConnectorWorkerMetadata | None: # Scheduler-side methods # ============================== + def bind_scheduler_state(self, scheduler_state: SchedulerState): + """ + Bind the scheduler state to the connector. + This function is called by the scheduler after initialization + and before the first model execution. + + Args: + scheduler_state (SchedulerState): the scheduler state. + """ + return + @abstractmethod def get_num_new_matched_tokens( self, diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py index 4ef8f0ac9c90..3757188304be 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py @@ -18,6 +18,7 @@ KVConnectorMetadata, KVConnectorRole, KVConnectorWorkerMetadata, + SchedulerState, ) from vllm.distributed.kv_transfer.kv_connector.v1.metrics import ( KVConnectorPromMetrics, @@ -219,6 +220,10 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): for c in self._connectors: c.register_kv_caches(kv_caches) + def bind_scheduler_state(self, scheduler_state: SchedulerState): + for c in self._connectors: + c.bind_scheduler_state(scheduler_state) + # We must override the base class method here because we need to bind # the metadata to each connector in the order of the connectors in the # MultiKVConnectorMetadata. diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/simple_cpu_offload_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/simple_cpu_offload_connector.py index 6475b941ba59..87f156cfd620 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/simple_cpu_offload_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/simple_cpu_offload_connector.py @@ -13,6 +13,7 @@ KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole, + SchedulerState, SupportsHMA, ) from vllm.logger import init_logger @@ -31,7 +32,6 @@ 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 @@ -165,10 +165,11 @@ def build_connector_worker_meta(self): # --- Scheduler-side methods --- - # NOTE: New API only for SimpleCPUOffloadConnector. - def bind_gpu_block_pool(self, gpu_block_pool: "BlockPool") -> None: + def bind_scheduler_state(self, scheduler_state: SchedulerState) -> None: if self.scheduler_manager is not None: - self.scheduler_manager.bind_gpu_block_pool(gpu_block_pool) + self.scheduler_manager.bind_gpu_block_pool( + scheduler_state.kv_cache_manager.block_pool + ) def get_num_new_matched_tokens( self, diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index 395fa80bfe53..c76ca53d9a5e 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -24,7 +24,10 @@ KVConnectorRole, SupportsHMA, ) -from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorMetadata +from vllm.distributed.kv_transfer.kv_connector.v1.base import ( + KVConnectorMetadata, + SchedulerState, +) from vllm.distributed.kv_transfer.kv_connector.v1.metrics import KVConnectorStats from vllm.logger import init_logger from vllm.model_executor.layers.fused_moe.routed_experts_capturer import ( @@ -238,12 +241,6 @@ def __init__( hash_block_size=hash_block_size, metrics_collector=self.kv_metrics_collector, ) - # Bind GPU block pool to the KV connector. This must happen after - # kv_cache_manager is constructed so block_pool is available. - if self.connector is not None and hasattr( - self.connector, "bind_gpu_block_pool" - ): - self.connector.bind_gpu_block_pool(self.kv_cache_manager.block_pool) self.use_pp = self.parallel_config.pipeline_parallel_size > 1 self.use_v2_model_runner = envs.VLLM_USE_V2_MODEL_RUNNER @@ -298,6 +295,9 @@ def __init__( ) self._pause_state: PauseState = PauseState.UNPAUSED + if self.connector is not None: + state = SchedulerState(kv_cache_manager=self.kv_cache_manager) + self.connector.bind_scheduler_state(state) def _mamba_block_aligned_split( self,