Skip to content

vllm.distributed.kv_transfer.kv_connector.v1.offloading_connector

Classes:

OffloadingConnector

Bases: KVConnectorBase_V1, SupportsHMA

Source code in vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py
class OffloadingConnector(KVConnectorBase_V1, SupportsHMA):
    @cached_property
    def _bounding_group_ids(self) -> tuple[int, ...]:
        """Prefix-cacheable groups this connector does not offload.

        Never offer tokens for a range some group cannot cover: past such a
        group's own cached prefix its KV would be left unwritten.
        """
        offloaded = set(get_offloading_group_ids(self._kv_cache_config))
        return tuple(
            group_id
            for group_id, group in enumerate(self._kv_cache_config.kv_cache_groups)
            if group.kv_cache_spec.prefix_cacheable and group_id not in offloaded
        )

    @property
    def scheduler(self) -> OffloadingConnectorScheduler:
        assert self.connector_scheduler is not None
        return self.connector_scheduler

    @property
    def requires_kv_delivery(self) -> bool:
        # Runs as kv_both, but is a best-effort cache: a dropped save is just a
        # future cache miss, so opt out of the producer-role default.
        return False

    def __init__(
        self,
        vllm_config: VllmConfig,
        role: KVConnectorRole,
        kv_cache_config: KVCacheConfig,
    ):
        super().__init__(vllm_config, role, kv_cache_config)

        offloading_config = build_offloading_config(vllm_config, kv_cache_config)
        self._canonical_layout = offloading_config.canonical_layout
        spec = OffloadingSpecFactory.create_spec(offloading_config)

        self.connector_scheduler: OffloadingConnectorScheduler | None = None
        self.connector_worker: OffloadingConnectorWorker | None = None
        if role == KVConnectorRole.SCHEDULER:
            self.connector_scheduler = OffloadingConnectorScheduler(
                spec, vllm_config, kv_cache_config
            )
        elif role == KVConnectorRole.WORKER:
            self.connector_worker = OffloadingConnectorWorker(
                spec, vllm_config, kv_cache_config
            )

    def shutdown(self) -> None:
        if self.connector_worker is not None:
            self.connector_worker.shutdown()
        if self.connector_scheduler is not None:
            self.connector_scheduler.shutdown()

    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 handle_preemptions(self, kv_connector_metadata: KVConnectorMetadata):
        assert self.connector_worker is not None
        assert isinstance(kv_connector_metadata, OffloadingConnectorMetadata)
        self.connector_worker.handle_preemptions(kv_connector_metadata)

    def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None:
        assert self.connector_worker is not None
        assert isinstance(self._connector_metadata, OffloadingConnectorMetadata)
        self.connector_worker.start_kv_transfers(self._connector_metadata)

    def wait_for_layer_load(self, layer_name: str) -> None:
        pass

    def save_kv_layer(
        self,
        layer_name: str,
        kv_layer: torch.Tensor,
        attn_metadata: "AttentionMetadata",
        **kwargs,
    ) -> None:
        pass

    def wait_for_save(self):
        assert self.connector_worker is not None
        assert isinstance(self._connector_metadata, OffloadingConnectorMetadata)
        # Defer store jobs to the next step's start_kv_transfers.
        self.connector_worker.prepare_store_kv(self._connector_metadata)

    def get_finished(self, finished_req_ids: set[str]) -> tuple[set[str], set[str]]:
        assert self.connector_worker is not None
        assert isinstance(self._connector_metadata, OffloadingConnectorMetadata)
        return self.connector_worker.get_finished(finished_req_ids)

    def build_connector_worker_meta(self) -> OffloadingWorkerMetadata | None:
        if self.connector_worker is not None:
            return self.connector_worker.build_connector_worker_meta()
        return None

    def on_new_request(self, request: "Request") -> None:
        assert self.connector_scheduler is not None
        self.connector_scheduler.on_new_request(request)

    def get_num_new_matched_tokens(
        self, request: "Request", num_computed_tokens: int
    ) -> tuple[int | None, bool]:
        assert self.connector_scheduler is not None
        return self.connector_scheduler.get_num_new_matched_tokens(
            request,
            num_computed_tokens,
            max_num_new_tokens=self._max_loadable_tokens(request, num_computed_tokens),
        )

    def _max_loadable_tokens(
        self, request: "Request", num_computed_tokens: int
    ) -> int | None:
        """How far past ``num_computed_tokens`` a load may reach, if bounded.

        Bounded by the deepest prefix the groups this connector does not
        offload already hold, since nothing refills them beyond it.
        """
        if not self._bounding_group_ids:
            return None
        assert self._kv_cache_manager is not None
        coordinator = self._kv_cache_manager.coordinator
        assert isinstance(coordinator, HybridKVCacheCoordinator)
        _, per_group_hits = coordinator.find_longest_cache_hit_per_group(
            request.block_hashes, request.num_tokens - 1
        )
        bound = min(per_group_hits[group_id] for group_id in self._bounding_group_ids)
        return max(0, bound - 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 has_pending_push_work(self) -> bool:
        assert self.connector_scheduler is not None
        return self.connector_scheduler.has_pending_push_work()

    def update_connector_output(self, connector_output: KVConnectorOutput):
        assert self.connector_scheduler is not None
        self.connector_scheduler.update_connector_output(connector_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)

    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(request)

    def take_events(self) -> Iterable[KVCacheEvent]:
        assert self.connector_scheduler is not None
        return self.connector_scheduler.take_events()

    @classmethod
    def get_required_kvcache_layout(cls, vllm_config: VllmConfig) -> str | None:
        if vllm_config.attention_config.hisparse_config is not None:
            return "BLHNC"
        return "LBHNC"

    def reset_cache(self) -> bool | None:
        assert self.connector_scheduler is not None
        self.connector_scheduler.reset_cache()
        return True

    def get_kv_connector_stats(self) -> KVConnectorStats | None:
        if self.connector_scheduler is not None:
            return self.connector_scheduler.get_stats()
        return None

    @classmethod
    def build_kv_connector_stats(
        cls, data: dict[str, Any] | None = None
    ) -> KVConnectorStats | None:
        return (
            OffloadingConnectorStats(data=data)
            if data is not None
            else OffloadingConnectorStats()
        )

    @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]],
    ) -> KVConnectorPromMetrics:
        return OffloadPromMetrics(
            vllm_config, metric_types, labelnames, per_engine_labelvalues
        )

_bounding_group_ids cached property

Prefix-cacheable groups this connector does not offload.

Never offer tokens for a range some group cannot cover: past such a group's own cached prefix its KV would be left unwritten.

_max_loadable_tokens(request, num_computed_tokens)

How far past num_computed_tokens a load may reach, if bounded.

Bounded by the deepest prefix the groups this connector does not offload already hold, since nothing refills them beyond it.

Source code in vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py
def _max_loadable_tokens(
    self, request: "Request", num_computed_tokens: int
) -> int | None:
    """How far past ``num_computed_tokens`` a load may reach, if bounded.

    Bounded by the deepest prefix the groups this connector does not
    offload already hold, since nothing refills them beyond it.
    """
    if not self._bounding_group_ids:
        return None
    assert self._kv_cache_manager is not None
    coordinator = self._kv_cache_manager.coordinator
    assert isinstance(coordinator, HybridKVCacheCoordinator)
    _, per_group_hits = coordinator.find_longest_cache_hit_per_group(
        request.block_hashes, request.num_tokens - 1
    )
    bound = min(per_group_hits[group_id] for group_id in self._bounding_group_ids)
    return max(0, bound - num_computed_tokens)