Skip to content
Closed
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 24 additions & 1 deletion vllm/distributed/kv_transfer/kv_connector/v1/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The docstring states that this dataclass "ensures read-only access to scheduler state". While the dataclass itself is frozen=True (preventing reassignment of its fields), the kv_cache_manager object it contains is mutable. A connector could technically call mutating methods on the manager. This is more of a design guideline than a technical enforcement, but the docstring might be slightly misleading in its current phrasing.

expansion of the scheduler state in the future.
"""

kv_cache_manager: "KVCacheManager"


class KVConnectorWorkerMetadata(ABC):
"""
Abstract Metadata used to communicate back
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
KVConnectorBase_V1,
KVConnectorMetadata,
KVConnectorRole,
SchedulerState,
SupportsHMA,
)
from vllm.logger import init_logger
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down
14 changes: 7 additions & 7 deletions vllm/v1/core/sched/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
Loading