diff --git a/benchmarks/bench_staging_restore.py b/benchmarks/bench_staging_restore.py index 990e118..3b46e6b 100644 --- a/benchmarks/bench_staging_restore.py +++ b/benchmarks/bench_staging_restore.py @@ -20,7 +20,7 @@ # First Party from benchmarks.utils.constants import BLOCK_TOKENS # noqa: E402 -from daser.connector.staging import copy_staging_to_kv_cache # noqa: E402 +from daser.connector.worker.staging import copy_staging_to_kv_cache # noqa: E402 from daser.ops.rope_apply import clear_rope_apply_cache # noqa: E402 diff --git a/benchmarks/utils/vllm_bench.py b/benchmarks/utils/vllm_bench.py index 73f4005..c292ad9 100644 --- a/benchmarks/utils/vllm_bench.py +++ b/benchmarks/utils/vllm_bench.py @@ -11,7 +11,7 @@ import asyncio from collections.abc import Callable -from dataclasses import asdict +from dataclasses import asdict, replace import json import math from pathlib import Path @@ -297,9 +297,19 @@ def run_load( asyncio.run(_wait_lmcache_quiescent(manifest, settle_seconds=0.0)) elif backend_run.backend == "daser": _drain_daser(manifest, print_kv=print_kv) - warm_metrics, warm_hit_rate = _run_phase( - args, manifest, warm_raw, run_command=run_command - ) + if backend_run.backend == "daser" and args.evict: + warm_metrics, warm_hit_rate = _run_daser_evict_warm_phase( + args, + manifest, + warm_raw, + run_command=run_command, + ) + else: + warm_metrics, warm_hit_rate = _run_phase( + args, manifest, warm_raw, run_command=run_command + ) + if backend_run.backend == "daser" and args.evict: + _require_daser_evict_tier_activity(warm_metrics) cold_summary = _normalise_result(cold_raw) warm_summary = _normalise_result(warm_raw) _apply_phase_metrics(cold_summary, cold_hit_rate) @@ -373,6 +383,33 @@ def _run_phase( return _collect_phase_metrics(manifest, before_metrics) +def _run_daser_evict_warm_phase( + args: RunBenchArgs, + manifest: BenchmarkManifest, + raw_path: Path, + *, + run_command: Callable[[list[str]], Any], +) -> tuple[dict[str, Any], float | None]: + """Perturb LRU order before measuring a complete evict warm phase. + + A same-order scan of a working set larger than L1 can produce 100% L2 + reads even when most entries were resident. Replaying the first 20% of the + deterministic workload first makes the subsequent complete phase exercise + both resident L1 entries and evicted L2 entries. Both commands are included + in the returned warm metric delta; correctness and latency use only the + complete phase result. + """ + before_metrics = asyncio.run(collect_phase_metrics(manifest)) + prime_args = replace( + args, + bench_num_prompts=max(1, math.ceil(args.bench_num_prompts * 0.2)), + ) + prime_path = raw_path.with_name("vllm_bench_warm_prime.json") + run_command(_bench_command(prime_args, manifest.endpoints["vllm"], prime_path)) + run_command(_bench_command(args, manifest.endpoints["vllm"], raw_path)) + return _collect_phase_metrics(manifest, before_metrics) + + def _collect_phase_metrics( manifest: BenchmarkManifest, before_metrics: dict[str, Any] | None, @@ -414,11 +451,21 @@ def _drain_daser( endpoint = manifest.endpoints.get("daser") if endpoint is None: return - try: - response = httpx.post(f"{endpoint.url}/drain", timeout=30.0) - response.raise_for_status() - except Exception as exc: # noqa: BLE001 - print_kv("daser_drain_status", f"unavailable ({exc})") + response = httpx.post(f"{endpoint.url}/drain", timeout=30.0) + response.raise_for_status() + print_kv("daser_drain_status", "ok") + + +def _require_daser_evict_tier_activity(metrics: dict[str, Any]) -> None: + """Require one evict warm phase to exercise both DaseR cache tiers.""" + counters = metrics.get("backend_prometheus", {}) + l1_hits = float(counters.get("daser_l1_hits_total", 0.0)) + l2_reads = float(counters.get("daser_l2_reads_total", 0.0)) + if l1_hits <= 0 or l2_reads <= 0: + raise RuntimeError( + "DaseR evict warm phase must exercise both tiers: " + f"l1_hits={l1_hits:g} l2_reads={l2_reads:g}" + ) def _bench_command( diff --git a/daser/connector/daser_connector.py b/daser/connector/daser_connector.py index 32e00dd..44ef88d 100644 --- a/daser/connector/daser_connector.py +++ b/daser/connector/daser_connector.py @@ -1,12 +1,9 @@ # SPDX-License-Identifier: Apache-2.0 # Standard -import asyncio -import threading from typing import TYPE_CHECKING, Any # Third Party -import torch from vllm.distributed import get_tensor_model_parallel_rank from vllm.distributed.kv_transfer.kv_connector.v1.base import ( KVConnectorBase_V1, @@ -18,29 +15,15 @@ from vllm.config import VllmConfig # First Party -from daser.connector.helpers import PendingStore -from daser.connector.ipc_client import IPCClientAsync, IPCClientSync +from daser.connector.ipc_client import IPCClientSync from daser.connector.metadata import DaserConnectorMeta, ReqLoadSpec, ReqStoreSpec -from daser.connector.reuse import CacheReuseStrategy, build_cache_reuse_strategy -from daser.connector.scheduler import ( - SchedulerConnectorMixin, - _block_ids_for_chunk, - _contiguous_prefix_tokens, - _trim_chunk_to_external_window, -) -from daser.connector.staging import ( +from daser.connector.scheduler.adapter import SchedulerConnectorMixin +from daser.connector.scheduler.lifecycle import RequestLifecycle +from daser.connector.worker.adapter import WorkerConnectorMixin +from daser.connector.worker.runtime import WorkerRuntime +from daser.connector.worker.staging import ( DEFAULT_ROPE_DELTA_SCALE, ) -from daser.connector.staging import ( - apply_rope_delta_to_key_block as _apply_rope_delta_to_key_block, -) -from daser.connector.staging import ( - build_load_read_plan as _build_load_read_plan, -) -from daser.connector.staging import ( - copy_staging_to_kv_cache as _copy_staging_to_kv_cache, -) -from daser.connector.worker import _LOAD_REQUEST_MAX_INFLIGHT, WorkerConnectorMixin from daser.logging import init_logger logger = init_logger(__name__) @@ -51,15 +34,33 @@ "DaserConnectorMeta", "ReqLoadSpec", "ReqStoreSpec", - "_apply_rope_delta_to_key_block", - "_build_load_read_plan", - "_block_ids_for_chunk", - "_contiguous_prefix_tokens", - "_copy_staging_to_kv_cache", - "_trim_chunk_to_external_window", ] +def _extract_rope_config(vllm_config: "VllmConfig") -> tuple[float, int, bool]: + """Extract worker RoPE geometry from the vLLM model config.""" + model_config = getattr(vllm_config, "model_config", None) + if model_config is None: + return 10000.0, 0, True + try: + head_size = int(model_config.get_head_size()) + except Exception: # noqa: BLE001 + logger.warning("[CONNECTOR] could not infer RoPE head size") + return 10000.0, 0, True + hf_text_config = getattr(model_config, "hf_text_config", None) + rope_parameters = getattr(hf_text_config, "rope_parameters", None) or {} + if not isinstance(rope_parameters, dict): + rope_parameters = {} + model_type = str(getattr(hf_text_config, "model_type", "")) + rope_base = ( + 1000000.0 + if "qwen" in model_type and "rope_theta" not in rope_parameters + else float(rope_parameters.get("rope_theta", 10000.0)) + ) + partial = float(rope_parameters.get("partial_rotary_factor", 1.0)) + return rope_base, int(head_size * partial), True + + class DaserConnector( SchedulerConnectorMixin, WorkerConnectorMixin, @@ -69,8 +70,8 @@ class DaserConnector( The entrypoint remains in this module for vLLM's ``kv_connector_module_path``. Scheduler-role behavior lives in - ``daser.connector.scheduler`` and worker-role behavior lives in - ``daser.connector.worker``. + ``daser.connector.scheduler.adapter`` and worker-role behavior lives in + ``daser.connector.worker.adapter``. Args: vllm_config: full VllmConfig from vLLM. @@ -93,87 +94,46 @@ def __init__( ): extra = vllm_config.kv_transfer_config.kv_connector_extra_config or {} - self._socket_path: str = extra.get("socket_path", "/tmp/daser.sock") - self._store_path: str = "" - self._slot_size: int = 0 - parallel_config = getattr(vllm_config, "parallel_config", None) - self._tp_size = int(getattr(parallel_config, "tensor_parallel_size", 1) or 1) - self._tp_rank = ( - get_tensor_model_parallel_rank() - if role == KVConnectorRole.WORKER and self._tp_size > 1 - else 0 - ) - self._server_tp_size = 1 - self._local_slot_size = 0 - self._rank_stride_bytes = 0 - self._block_tokens: int = 16 - self._model_id: str = "default" - self._skip_l2: bool = bool(extra.get("skip_l2", False)) - self._cache_reuse_strategy: CacheReuseStrategy - self._set_cache_reuse_strategy(str(extra.get("cache_reuse_mode", "chunk"))) - self._runtime_config_ready = False - self._rope_base: float = 10000.0 - self._rope_rotary_dim: int = 0 - self._rope_is_neox_style: bool = True - self._rope_delta_scale: float = float( - extra.get("rope_delta_scale", DEFAULT_ROPE_DELTA_SCALE) - ) - self._load_key_scale: float = float(extra.get("load_key_scale", 1.0)) - self._load_value_scale: float = float(extra.get("load_value_scale", 1.0)) - self._init_rope_config(vllm_config) - + socket_path = str(extra.get("socket_path", "/tmp/daser.sock")) if role == KVConnectorRole.SCHEDULER: - self._ipc_sync = IPCClientSync(self._socket_path) - self._refresh_runtime_config() - self._pending_loads: dict[str, dict[str, Any]] = {} - self._pending_stores: dict[str, dict[str, Any]] = {} - self._pending_alloc: dict[str, PendingStore] = {} - self._pending_async_saves: set[str] = set() - self._req_tokens: dict[str, list[int]] = {} + self._request_lifecycle = RequestLifecycle( + ipc_client=IPCClientSync(socket_path), + block_tokens=16, + slot_size=0, + model_id="default", + cache_reuse_mode=str(extra.get("cache_reuse_mode", "chunk")), + runtime_config_ready=False, + ) + self._request_lifecycle.refresh_runtime_config() else: - self._transfer_ready = False - self._transfer_mode = str(extra.get("transfer_mode", "iouring")) - self._ipc_load_async = IPCClientAsync(self._socket_path) - self._ipc_load_async_pool = [ - self._ipc_load_async, - *[ - IPCClientAsync(self._socket_path) - for _ in range(max(0, _LOAD_REQUEST_MAX_INFLIGHT - 1)) - ], - ] - self._ipc_store_async = IPCClientAsync(self._socket_path) - self._kv_caches: dict[str, torch.Tensor] = {} - self._layer_names: list[str] = [] - self._layer_idx_map: dict[str, int] = {} - self._meta: DaserConnectorMeta | None = None - self._save_futures: list = [] - self._pending_save_staging_bytes = 0 - self._store_staging_bytes = 0 - self._pending_store_staging_limit_bytes = 0 - self._store_staging_pool = None - self._pending_commits: set[str] = set() - self._pending_finished_saves: dict[str, Any] = {} - self._pending_loads: dict[str, Any] = {} - self._invalid_load_block_ids: set[int] = set() - self._load_request_queue = None - self._load_request_queue_lock = threading.Lock() - self._load_request_dispatcher_future = None - self._load_loop = asyncio.new_event_loop() - self._store_loop = asyncio.new_event_loop() - self._load_thread = threading.Thread( - target=self._run_load_loop, - daemon=True, - name="daser-load-io", + parallel_config = getattr(vllm_config, "parallel_config", None) + tp_size = int(getattr(parallel_config, "tensor_parallel_size", 1) or 1) + tp_rank = get_tensor_model_parallel_rank() if tp_size > 1 else 0 + rope_base, rope_rotary_dim, rope_is_neox_style = _extract_rope_config( + vllm_config ) - self._store_thread = threading.Thread( - target=self._run_store_loop, - daemon=True, - name="daser-store-io", + self._worker_runtime = WorkerRuntime( + socket_path=socket_path, + transfer_mode=str(extra.get("transfer_mode", "iouring")), + skip_l2=bool(extra.get("skip_l2", False)), + tp_size=tp_size, + tp_rank=tp_rank, + server_tp_size=1, + slot_size=0, + store_path="", + rank_stride_bytes=0, + rope_base=rope_base, + rope_rotary_dim=rope_rotary_dim, + rope_is_neox_style=rope_is_neox_style, + rope_delta_scale=float( + extra.get("rope_delta_scale", DEFAULT_ROPE_DELTA_SCALE) + ), + load_key_scale=float(extra.get("load_key_scale", 1.0)), + load_value_scale=float(extra.get("load_value_scale", 1.0)), + kv_cache_config=kv_cache_config, ) - self._load_thread.start() - self._store_thread.start() - logger.info("[CONNECTOR] role=%s socket=%s", role.name, self._socket_path) + logger.info("[CONNECTOR] role=%s socket=%s", role.name, socket_path) @property def prefer_cross_layer_blocks(self) -> bool: @@ -204,101 +164,3 @@ def get_required_kvcache_layout(cls, vllm_config: "VllmConfig") -> str | None: Class-level config helper with no mutable state. """ return "NHD" - - def _refresh_runtime_config(self) -> None: - """Refresh server-owned runtime config over IPC when available.""" - client = getattr(self, "_ipc_sync", None) - owns_client = client is None - if client is None: - client = IPCClientSync(self._socket_path) - try: - config = client.get_runtime_config() - except Exception as exc: # noqa: BLE001 - logger.info("[CONNECTOR] runtime config unavailable: %s", exc) - return - finally: - if owns_client: - client.close() - - self._store_path = str(config.get("store_path", self._store_path)) - self._slot_size = int(config.get("slot_size", self._slot_size)) - self._server_tp_size = int( - config.get("tensor_parallel_size", self._server_tp_size) - ) - self._rank_stride_bytes = int( - config.get("rank_stride_bytes", self._rank_stride_bytes) - ) - self._block_tokens = int(config.get("block_tokens", self._block_tokens)) - self._model_id = str(config.get("model_id", self._model_id)) - self._set_cache_reuse_strategy(str(config["cache_reuse_mode"])) - self._skip_l2 = bool(config.get("skip_l2", self._skip_l2)) - self._runtime_config_ready = bool( - self._slot_size and (self._store_path or self._skip_l2) - ) - self._transfer_mode = str( - config.get("transfer_mode", getattr(self, "_transfer_mode", "iouring")) - ) - logger.info( - "[CONNECTOR] runtime config store=%s slot_size=%d block_tokens=%d " - "model=%s transfer=%s skip_l2=%s", - self._store_path, - self._slot_size, - self._block_tokens, - self._model_id, - getattr(self, "_transfer_mode", "iouring"), - self._skip_l2, - ) - - def _set_cache_reuse_strategy(self, cache_reuse_mode: str) -> None: - """Set scheduler cache reuse strategy. - - Args: - cache_reuse_mode: either ``"chunk"`` or ``"prefix"``. - """ - self._cache_reuse_strategy = build_cache_reuse_strategy( - cache_reuse_mode, - self._block_tokens, - ) - - def _init_rope_config(self, vllm_config: "VllmConfig") -> None: - """Extract default RoPE settings from vLLM model config. - - Args: - vllm_config: vLLM runtime config passed to the connector. - """ - model_config = getattr(vllm_config, "model_config", None) - if model_config is None: - return - try: - head_size = int(model_config.get_head_size()) - except Exception: # noqa: BLE001 - logger.warning("[CONNECTOR] could not infer RoPE head size") - return - - hf_text_config = getattr(model_config, "hf_text_config", None) - rope_parameters = getattr(hf_text_config, "rope_parameters", None) or {} - if not isinstance(rope_parameters, dict): - rope_parameters = {} - model_type = str(getattr(hf_text_config, "model_type", "")) - if "qwen" in model_type and "rope_theta" not in rope_parameters: - rope_base = 1000000.0 - else: - rope_base = float(rope_parameters.get("rope_theta", 10000.0)) - partial = float(rope_parameters.get("partial_rotary_factor", 1.0)) - rotary_dim = int(head_size * partial) - - self._rope_base = rope_base - self._rope_rotary_dim = rotary_dim - self._rope_is_neox_style = True - logger.info( - "[CONNECTOR] rope base=%s rotary_dim=%d neox=%s", - self._rope_base, - self._rope_rotary_dim, - self._rope_is_neox_style, - ) - logger.info( - "[CONNECTOR] load tuning rope_delta_scale=%s key_scale=%s value_scale=%s", - self._rope_delta_scale, - self._load_key_scale, - self._load_value_scale, - ) diff --git a/daser/connector/ipc_client.py b/daser/connector/ipc_client.py index f024c10..4fedc06 100644 --- a/daser/connector/ipc_client.py +++ b/daser/connector/ipc_client.py @@ -498,7 +498,7 @@ async def transfer_store_cuda( allocation_base_ptr: int, allocation_offset: int, producer_pid: int, - spans: list[dict[str, int]], + spans: list[dict[str, Any]], ) -> list[str]: """Store from a CUDA IPC buffer through the server transfer layer. diff --git a/daser/connector/scheduler/__init__.py b/daser/connector/scheduler/__init__.py new file mode 100644 index 0000000..87904f9 --- /dev/null +++ b/daser/connector/scheduler/__init__.py @@ -0,0 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 + +from daser.connector.scheduler.adapter import SchedulerConnectorMixin +from daser.connector.scheduler.lifecycle import RequestLifecycle + +__all__ = ["RequestLifecycle", "SchedulerConnectorMixin"] diff --git a/daser/connector/scheduler/adapter.py b/daser/connector/scheduler/adapter.py new file mode 100644 index 0000000..b0f009f --- /dev/null +++ b/daser/connector/scheduler/adapter.py @@ -0,0 +1,59 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from vllm.v1.core.kv_cache_utils import KVCacheBlocks + from vllm.v1.core.scheduler import SchedulerOutput + from vllm.v1.request import Request + +from daser.connector.metadata import DaserConnectorMeta + + +class SchedulerConnectorMixin: + """Adapt vLLM scheduler hooks to the request lifecycle interface.""" + + def get_num_new_matched_tokens( + self, + request: "Request", + num_computed_tokens: int, + ) -> tuple[int | None, bool]: + """Return DaseR cache credit for one request.""" + return self._request_lifecycle.get_num_new_matched_tokens( + request, + num_computed_tokens, + ) + + def update_state_after_alloc( + self, + request: "Request", + blocks: "KVCacheBlocks", + num_external_tokens: int, + ) -> None: + """Bind allocated vLLM blocks to pending lifecycle work.""" + self._request_lifecycle.update_state_after_alloc( + request, + blocks, + num_external_tokens, + ) + + def build_connector_meta( + self, + scheduler_output: "SchedulerOutput", + ) -> DaserConnectorMeta: + """Build worker metadata from request lifecycle state.""" + return self._request_lifecycle.build_connector_meta(scheduler_output) + + def request_finished( + self, + request: "Request", + block_ids: list[int], + ) -> tuple[bool, dict[str, Any] | None]: + """Return whether worker save completion still holds request blocks.""" + return self._request_lifecycle.request_finished(request, block_ids) + + def update_connector_output(self, connector_output: Any) -> None: + """Apply worker transfer completions to request lifecycle state.""" + self._request_lifecycle.update_connector_output(connector_output) diff --git a/daser/connector/scheduler.py b/daser/connector/scheduler/lifecycle.py similarity index 70% rename from daser/connector/scheduler.py rename to daser/connector/scheduler/lifecycle.py index 6064340..ddcc69c 100644 --- a/daser/connector/scheduler.py +++ b/daser/connector/scheduler/lifecycle.py @@ -1,312 +1,69 @@ # SPDX-License-Identifier: Apache-2.0 -# Standard +from __future__ import annotations + import logging import math from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - # Third Party from vllm.v1.core.kv_cache_utils import KVCacheBlocks from vllm.v1.core.scheduler import SchedulerOutput from vllm.v1.request import Request -# First Party -from daser.connector.helpers import PendingStore, base_req_id +from daser.connector.helpers import PendingStore from daser.connector.metadata import DaserConnectorMeta, ReqLoadSpec, ReqStoreSpec -from daser.connector.reuse import build_cache_reuse_strategy +from daser.connector.scheduler.planning import ( + _base_req_id, + _computed_tokens_after_step, + _contiguous_prefix_tokens, + _get_kv_transfer_flag, + _load_spec_from_chunk, + _matches_request_or_store_id, + _merge_adjacent_load_specs, + _store_slot_index, + _trim_chunk_to_external_window, +) +from daser.connector.scheduler.reuse import build_cache_reuse_strategy from daser.logging import init_logger logger = init_logger(__name__) -def _base_req_id(req_id: str) -> str: - """Compatibility wrapper for tests importing the scheduler-private helper.""" - return base_req_id(req_id) - - -def _store_slot_index(req_id: str) -> int | None: - """Return the rolling-prefix store slot encoded in a synthetic request ID. - - Args: - req_id: vLLM request ID or ``:store:`` synthetic ID. - - Returns: - Slot index for synthetic store IDs, or None for regular IDs. - """ - if ":store:" not in req_id: - return None - try: - return int(req_id.rsplit(":store:", 1)[1]) - except ValueError: - return None - - -def _matches_request_or_store_id(req_id: str, base_req_id: str) -> bool: - """Return whether ``req_id`` belongs to a base request. - - Args: - req_id: vLLM request ID or synthetic connector work ID. - base_req_id: Base vLLM request ID to match. - - Returns: - True when ``req_id`` is the base request or one of its synthetic - store/load entries. - """ - return ( - req_id == base_req_id - or req_id.startswith(f"{base_req_id}:store:") - or req_id.startswith(f"{base_req_id}:load:") - ) - - -def _computed_tokens_after_step( - scheduler_output: "SchedulerOutput", -) -> dict[str, int]: - """Return per-request token counts that are valid after this step. - - Args: - scheduler_output: vLLM SchedulerOutput for this step. - - Returns: - Mapping from request ID to ``num_computed_tokens + scheduled_tokens``. - Falls back to the scheduled token count when older or test scheduler - outputs do not expose prior computed-token metadata. - """ - scheduled = dict(getattr(scheduler_output, "num_scheduled_tokens", {})) - computed_after = {req_id: int(tokens) for req_id, tokens in scheduled.items()} - - for req_data in getattr(scheduler_output, "scheduled_new_reqs", []) or []: - req_id = str(getattr(req_data, "req_id", "")) - if req_id in scheduled: - computed_after[req_id] = int( - getattr(req_data, "num_computed_tokens", 0) - ) + int(scheduled[req_id]) - - cached_reqs = getattr(scheduler_output, "scheduled_cached_reqs", None) - if cached_reqs is not None: - req_ids = getattr(cached_reqs, "req_ids", []) - prior_counts = getattr(cached_reqs, "num_computed_tokens", []) - for req_id, prior in zip(req_ids, prior_counts, strict=False): - req_id = str(req_id) - if req_id in scheduled: - computed_after[req_id] = int(prior) + int(scheduled[req_id]) - - return computed_after - - -def _get_kv_transfer_flag(request: "Request", key: str) -> Any: - """Return ``request.kv_transfer_params[key]`` if present, else ``None``. - - Args: - request: vLLM ``Request`` or compatible object. - key: connector-specific flag name to extract. - - Returns: - The value under ``key``, or ``None`` when absent. - """ - params = getattr(request, "kv_transfer_params", None) - if not isinstance(params, dict): - return None - return params.get(key) - - -def _block_ids_for_chunk( - block_ids: list[int], - target_token_start: int, - num_slots: int, - block_tokens: int, - max_tokens: int | None = None, -) -> list[int]: - """Return vLLM block IDs for a chunk's target prompt range. - - Args: - block_ids: all block IDs allocated to the request. - target_token_start: token offset where the chunk starts in the prompt. - num_slots: number of blocks/slots covered by the chunk. - block_tokens: tokens per vLLM block. - max_tokens: optional upper bound on accepted external tokens. - - Returns: - Slice of block_ids for the chunk, or an empty list when the range - is not block-aligned or exceeds the allocated blocks. - """ - if target_token_start % block_tokens != 0: - return [] - target_block_start = target_token_start // block_tokens - effective_slots = num_slots - if max_tokens is not None: - remaining_tokens = max_tokens - target_token_start - if remaining_tokens <= 0: - return [] - effective_slots = min(num_slots, math.ceil(remaining_tokens / block_tokens)) - target_block_end = target_block_start + effective_slots - if target_block_start < 0 or target_block_end > len(block_ids): - return [] - return block_ids[target_block_start:target_block_end] - - -def _trim_chunk_to_external_window( - chunk: dict[str, Any], - block_ids: list[int], - external_start: int, - num_external_tokens: int, - block_tokens: int, - slot_size: int, -) -> bool: - """Trim chunk metadata to the external token interval vLLM requested. - - Args: - chunk: Mutable chunk metadata returned by the server. - block_ids: Full vLLM block allocation for the request. - external_start: Token offset where external KV loading begins. - num_external_tokens: Number of tokens accepted from the connector. - block_tokens: Tokens per vLLM block. - slot_size: Bytes per DaseR slot. - - Returns: - True when the chunk still covers at least one whole KV block. - """ - if external_start % block_tokens != 0 or num_external_tokens <= 0: - return False - target_start = int(chunk.get("target_token_start", 0)) - target_end = target_start + int(chunk["token_count"]) - external_end = external_start + num_external_tokens - load_start = max(target_start, external_start) - load_end = min(target_end, external_end) - load_start = ((load_start + block_tokens - 1) // block_tokens) * block_tokens - load_end = ((load_end + block_tokens - 1) // block_tokens) * block_tokens - load_end = min(load_end, target_end) - if load_end <= load_start: - return False - - skip_slots = (load_start - target_start) // block_tokens - num_slots = (load_end - load_start) // block_tokens - if load_start < external_start: - return False - block_start = load_start // block_tokens - block_end = block_start + num_slots - if block_start < 0 or block_end > len(block_ids): - return False - - chunk["start_slot"] = int(chunk["start_slot"]) + skip_slots - chunk["file_offset"] = int(chunk["file_offset"]) + skip_slots * slot_size - chunk["num_slots"] = num_slots - chunk["token_count"] = num_slots * block_tokens - chunk["target_token_start"] = load_start - chunk["block_ids"] = block_ids[block_start:block_end] - return bool(chunk["block_ids"]) - - -def _contiguous_prefix_tokens( - chunks: list[dict[str, Any]], num_computed_tokens: int -) -> int: - """Return external tokens covered contiguously after computed tokens. - - Args: - chunks: server chunk payloads with target_token_start and token_count. - num_computed_tokens: tokens vLLM already has locally. - - Returns: - Number of additional contiguous prefix tokens covered by chunks. - """ - covered_until = num_computed_tokens - for chunk in sorted( - chunks, - key=lambda item: int(item.get("target_token_start", 0)), - ): - target_start = int(chunk.get("target_token_start", 0)) - token_count = int(chunk["token_count"]) - target_end = target_start + token_count - if target_end <= covered_until: - continue - if target_start > covered_until: - break - covered_until = target_end - return covered_until - num_computed_tokens - - -def _load_spec_from_chunk(chunk: dict[str, Any]) -> ReqLoadSpec: - """Build a worker load specification from scheduler chunk metadata. - - Args: - chunk: Chunk metadata returned by the server and annotated with vLLM - block IDs during allocation. - - Returns: - ReqLoadSpec consumed by the worker load path. - - Async/thread-safety: - Pure scheduler-thread helper; it does not mutate connector state. - """ - return ReqLoadSpec( - chunk_key=str(chunk["chunk_key"]), - start_slot=int(chunk["start_slot"]), - num_slots=int(chunk["num_slots"]), - block_ids=list(chunk["block_ids"]), - file_offset=int(chunk["file_offset"]), - token_count=int(chunk["token_count"]), - target_token_start=int(chunk.get("target_token_start", 0)), - pos_offset=int(chunk.get("pos_offset", 0)), - ) - - -def _merge_adjacent_load_specs( - specs: list[ReqLoadSpec], - slot_size: int, -) -> list[ReqLoadSpec]: - """Merge adjacent load specs that describe one continuous KV byte range. - - Args: - specs: Load specs for one request in prompt order. - slot_size: Bytes represented by one DaseR KV slot. - - Returns: - Coalesced load specs. Chunk keys from the first spec in a run are kept - only as diagnostics; the worker load path addresses data by byte range. - - Async/thread-safety: - Pure scheduler-thread helper; it does not mutate connector state. - """ - merged: list[ReqLoadSpec] = [] - for spec in specs: - if not spec.block_ids: - continue - if not merged: - merged.append(spec) - continue - prev = merged[-1] - prev_slots = len(prev.block_ids) - adjacent = ( - prev.pos_offset == spec.pos_offset - and prev.start_slot + prev_slots == spec.start_slot - and prev.file_offset + prev_slots * slot_size == spec.file_offset - and prev.target_token_start + prev.token_count == spec.target_token_start - ) - if not adjacent: - merged.append(spec) - continue - merged[-1] = ReqLoadSpec( - chunk_key=prev.chunk_key, - start_slot=prev.start_slot, - num_slots=prev_slots + len(spec.block_ids), - block_ids=[*prev.block_ids, *spec.block_ids], - file_offset=prev.file_offset, - token_count=prev.token_count + spec.token_count, - target_token_start=prev.target_token_start, - pos_offset=prev.pos_offset, - ) - return merged - - -class SchedulerConnectorMixin: - """Scheduler-role vLLM connector behavior. +class RequestLifecycle: + """Own scheduler request state and synchronous IPC orchestration. Async/thread-safety: These methods run on vLLM's scheduler thread and use the synchronous IPC client owned by the connector instance. """ + def __init__( + self, + *, + ipc_client: Any, + block_tokens: int, + slot_size: int, + model_id: str, + cache_reuse_mode: str, + runtime_config_ready: bool, + ) -> None: + self._ipc_sync = ipc_client + self._block_tokens = block_tokens + self._slot_size = slot_size + self._model_id = model_id + self._cache_reuse_mode = cache_reuse_mode + self._runtime_config_ready = runtime_config_ready + self._cache_reuse_strategy = build_cache_reuse_strategy( + cache_reuse_mode, + block_tokens, + ) + self._pending_loads: dict[str, dict[str, Any]] = {} + self._pending_stores: dict[str, dict[str, Any]] = {} + self._pending_alloc: dict[str, PendingStore] = {} + self._pending_async_saves: set[str] = set() + self._req_tokens: dict[str, list[int]] = {} + def get_num_new_matched_tokens( self, request: "Request", @@ -574,6 +331,18 @@ def build_connector_meta( ) return meta + def refresh_runtime_config(self) -> None: + """Refresh scheduler geometry and reuse policy from the DaseR server. + + Returns: + None. + + Async/thread-safety: + Called on the scheduler thread and performs synchronous control-plane + IPC during connector initialization or recovery, never worker IO. + """ + self._refresh_runtime_config() + def _drop_preempted_pending_state( self, scheduler_output: "SchedulerOutput", @@ -734,6 +503,31 @@ def _record_pending_store_blocks(self, req_id: str, block_ids: list[int]) -> Non ] self._maybe_allocate_pending_store(req_id, pending_store) + def _refresh_runtime_config(self) -> None: + """Refresh scheduler geometry and reuse policy over owned sync IPC.""" + try: + config = self._ipc_sync.get_runtime_config() + except Exception as exc: # noqa: BLE001 + logger.info("[CONNECTOR] runtime config unavailable: %s", exc) + return + self._slot_size = int(config.get("slot_size", self._slot_size)) + block_tokens = int(config.get("block_tokens", self._block_tokens)) + self._model_id = str(config.get("model_id", self._model_id)) + cache_reuse_mode = str(config.get("cache_reuse_mode", self._cache_reuse_mode)) + if ( + cache_reuse_mode != self._cache_reuse_mode + or block_tokens != self._block_tokens + ): + self._cache_reuse_mode = cache_reuse_mode + self._block_tokens = block_tokens + self._cache_reuse_strategy = build_cache_reuse_strategy( + cache_reuse_mode, + self._block_tokens, + ) + else: + self._block_tokens = block_tokens + self._runtime_config_ready = bool(self._slot_size) + def _init_reuse_strategy(self) -> None: """Initialize the scheduler cache reuse strategy from current config.""" self._cache_reuse_strategy = build_cache_reuse_strategy( @@ -890,7 +684,58 @@ def _maybe_allocate_pending_store( tokens = self._req_tokens.get(req_id, []) if len(tokens) < requested_tokens: return - strategy.allocate_store(self, req_id, pending_store, tokens) + plan = strategy.plan_store( + req_id, + pending_store, + tokens, + set(self._pending_stores), + ) + if plan.invalid: + self._pending_alloc.pop(req_id, None) + return + if plan.intents: + try: + if len(plan.intents) == 1 and plan.intents[0].req_id == req_id: + intent = plan.intents[0] + allocations = [ + self.allocate_store_chunk( + intent.chunk_key, + intent.token_count, + ) + ] + else: + allocations = self.allocate_store_chunks( + [ + { + "chunk_key": intent.chunk_key, + "token_count": intent.token_count, + } + for intent in plan.intents + ] + ) + except Exception as exc: # noqa: BLE001 + logger.warning("[CONNECTOR] store allocation failed: %s", exc) + return + if len(allocations) != len(plan.intents): + logger.warning( + "[CONNECTOR] allocation returned %d entries for %d intents", + len(allocations), + len(plan.intents), + ) + return + for intent, alloc in zip(plan.intents, allocations, strict=True): + if bool(alloc.get("skipped", False)): + continue + alloc["chunk_key"] = str(alloc.get("chunk_key", intent.chunk_key)) + alloc["token_count"] = intent.token_count + alloc["num_slots"] = len(intent.block_ids) + alloc["block_ids"] = intent.block_ids + self._pending_stores[intent.req_id] = alloc + pending_store.rolling_key = plan.next_key + pending_store.rolling_slot_index = plan.next_slot + if plan.complete: + pending_store.chunk_key = plan.next_key + self._pending_alloc.pop(req_id, None) def request_finished( self, diff --git a/daser/connector/scheduler/planning.py b/daser/connector/scheduler/planning.py new file mode 100644 index 0000000..986011f --- /dev/null +++ b/daser/connector/scheduler/planning.py @@ -0,0 +1,296 @@ +# SPDX-License-Identifier: Apache-2.0 + +# Standard +import math +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + # Third Party + from vllm.v1.core.scheduler import SchedulerOutput + from vllm.v1.request import Request + +# First Party +from daser.connector.helpers import base_req_id +from daser.connector.metadata import ReqLoadSpec +from daser.logging import init_logger + +logger = init_logger(__name__) + + +def _base_req_id(req_id: str) -> str: + """Compatibility wrapper for tests importing the scheduler-private helper.""" + return base_req_id(req_id) + + +def _store_slot_index(req_id: str) -> int | None: + """Return the rolling-prefix store slot encoded in a synthetic request ID. + + Args: + req_id: vLLM request ID or ``:store:`` synthetic ID. + + Returns: + Slot index for synthetic store IDs, or None for regular IDs. + """ + if ":store:" not in req_id: + return None + try: + return int(req_id.rsplit(":store:", 1)[1]) + except ValueError: + return None + + +def _matches_request_or_store_id(req_id: str, base_req_id: str) -> bool: + """Return whether ``req_id`` belongs to a base request. + + Args: + req_id: vLLM request ID or synthetic connector work ID. + base_req_id: Base vLLM request ID to match. + + Returns: + True when ``req_id`` is the base request or one of its synthetic + store/load entries. + """ + return ( + req_id == base_req_id + or req_id.startswith(f"{base_req_id}:store:") + or req_id.startswith(f"{base_req_id}:load:") + ) + + +def _computed_tokens_after_step( + scheduler_output: "SchedulerOutput", +) -> dict[str, int]: + """Return per-request token counts that are valid after this step. + + Args: + scheduler_output: vLLM SchedulerOutput for this step. + + Returns: + Mapping from request ID to ``num_computed_tokens + scheduled_tokens``. + Falls back to the scheduled token count when older or test scheduler + outputs do not expose prior computed-token metadata. + """ + scheduled = dict(getattr(scheduler_output, "num_scheduled_tokens", {})) + computed_after = {req_id: int(tokens) for req_id, tokens in scheduled.items()} + + for req_data in getattr(scheduler_output, "scheduled_new_reqs", []) or []: + req_id = str(getattr(req_data, "req_id", "")) + if req_id in scheduled: + computed_after[req_id] = int( + getattr(req_data, "num_computed_tokens", 0) + ) + int(scheduled[req_id]) + + cached_reqs = getattr(scheduler_output, "scheduled_cached_reqs", None) + if cached_reqs is not None: + req_ids = getattr(cached_reqs, "req_ids", []) + prior_counts = getattr(cached_reqs, "num_computed_tokens", []) + for req_id, prior in zip(req_ids, prior_counts, strict=False): + req_id = str(req_id) + if req_id in scheduled: + computed_after[req_id] = int(prior) + int(scheduled[req_id]) + + return computed_after + + +def _get_kv_transfer_flag(request: "Request", key: str) -> Any: + """Return ``request.kv_transfer_params[key]`` if present, else ``None``. + + Args: + request: vLLM ``Request`` or compatible object. + key: connector-specific flag name to extract. + + Returns: + The value under ``key``, or ``None`` when absent. + """ + params = getattr(request, "kv_transfer_params", None) + if not isinstance(params, dict): + return None + return params.get(key) + + +def _block_ids_for_chunk( + block_ids: list[int], + target_token_start: int, + num_slots: int, + block_tokens: int, + max_tokens: int | None = None, +) -> list[int]: + """Return vLLM block IDs for a chunk's target prompt range. + + Args: + block_ids: all block IDs allocated to the request. + target_token_start: token offset where the chunk starts in the prompt. + num_slots: number of blocks/slots covered by the chunk. + block_tokens: tokens per vLLM block. + max_tokens: optional upper bound on accepted external tokens. + + Returns: + Slice of block_ids for the chunk, or an empty list when the range + is not block-aligned or exceeds the allocated blocks. + """ + if target_token_start % block_tokens != 0: + return [] + target_block_start = target_token_start // block_tokens + effective_slots = num_slots + if max_tokens is not None: + remaining_tokens = max_tokens - target_token_start + if remaining_tokens <= 0: + return [] + effective_slots = min(num_slots, math.ceil(remaining_tokens / block_tokens)) + target_block_end = target_block_start + effective_slots + if target_block_start < 0 or target_block_end > len(block_ids): + return [] + return block_ids[target_block_start:target_block_end] + + +def _trim_chunk_to_external_window( + chunk: dict[str, Any], + block_ids: list[int], + external_start: int, + num_external_tokens: int, + block_tokens: int, + slot_size: int, +) -> bool: + """Trim chunk metadata to the external token interval vLLM requested. + + Args: + chunk: Mutable chunk metadata returned by the server. + block_ids: Full vLLM block allocation for the request. + external_start: Token offset where external KV loading begins. + num_external_tokens: Number of tokens accepted from the connector. + block_tokens: Tokens per vLLM block. + slot_size: Bytes per DaseR slot. + + Returns: + True when the chunk still covers at least one whole KV block. + """ + if external_start % block_tokens != 0 or num_external_tokens <= 0: + return False + target_start = int(chunk.get("target_token_start", 0)) + target_end = target_start + int(chunk["token_count"]) + external_end = external_start + num_external_tokens + load_start = max(target_start, external_start) + load_end = min(target_end, external_end) + load_start = ((load_start + block_tokens - 1) // block_tokens) * block_tokens + load_end = ((load_end + block_tokens - 1) // block_tokens) * block_tokens + load_end = min(load_end, target_end) + if load_end <= load_start: + return False + + skip_slots = (load_start - target_start) // block_tokens + num_slots = (load_end - load_start) // block_tokens + if load_start < external_start: + return False + block_start = load_start // block_tokens + block_end = block_start + num_slots + if block_start < 0 or block_end > len(block_ids): + return False + + chunk["start_slot"] = int(chunk["start_slot"]) + skip_slots + chunk["file_offset"] = int(chunk["file_offset"]) + skip_slots * slot_size + chunk["num_slots"] = num_slots + chunk["token_count"] = num_slots * block_tokens + chunk["target_token_start"] = load_start + chunk["block_ids"] = block_ids[block_start:block_end] + return bool(chunk["block_ids"]) + + +def _contiguous_prefix_tokens( + chunks: list[dict[str, Any]], num_computed_tokens: int +) -> int: + """Return external tokens covered contiguously after computed tokens. + + Args: + chunks: server chunk payloads with target_token_start and token_count. + num_computed_tokens: tokens vLLM already has locally. + + Returns: + Number of additional contiguous prefix tokens covered by chunks. + """ + covered_until = num_computed_tokens + for chunk in sorted( + chunks, + key=lambda item: int(item.get("target_token_start", 0)), + ): + target_start = int(chunk.get("target_token_start", 0)) + token_count = int(chunk["token_count"]) + target_end = target_start + token_count + if target_end <= covered_until: + continue + if target_start > covered_until: + break + covered_until = target_end + return covered_until - num_computed_tokens + + +def _load_spec_from_chunk(chunk: dict[str, Any]) -> ReqLoadSpec: + """Build a worker load specification from scheduler chunk metadata. + + Args: + chunk: Chunk metadata returned by the server and annotated with vLLM + block IDs during allocation. + + Returns: + ReqLoadSpec consumed by the worker load path. + + Async/thread-safety: + Pure scheduler-thread helper; it does not mutate connector state. + """ + return ReqLoadSpec( + chunk_key=str(chunk["chunk_key"]), + start_slot=int(chunk["start_slot"]), + num_slots=int(chunk["num_slots"]), + block_ids=list(chunk["block_ids"]), + file_offset=int(chunk["file_offset"]), + token_count=int(chunk["token_count"]), + target_token_start=int(chunk.get("target_token_start", 0)), + pos_offset=int(chunk.get("pos_offset", 0)), + ) + + +def _merge_adjacent_load_specs( + specs: list[ReqLoadSpec], + slot_size: int, +) -> list[ReqLoadSpec]: + """Merge adjacent load specs that describe one continuous KV byte range. + + Args: + specs: Load specs for one request in prompt order. + slot_size: Bytes represented by one DaseR KV slot. + + Returns: + Coalesced load specs. Chunk keys from the first spec in a run are kept + only as diagnostics; the worker load path addresses data by byte range. + + Async/thread-safety: + Pure scheduler-thread helper; it does not mutate connector state. + """ + merged: list[ReqLoadSpec] = [] + for spec in specs: + if not spec.block_ids: + continue + if not merged: + merged.append(spec) + continue + prev = merged[-1] + prev_slots = len(prev.block_ids) + adjacent = ( + prev.pos_offset == spec.pos_offset + and prev.start_slot + prev_slots == spec.start_slot + and prev.file_offset + prev_slots * slot_size == spec.file_offset + and prev.target_token_start + prev.token_count == spec.target_token_start + ) + if not adjacent: + merged.append(spec) + continue + merged[-1] = ReqLoadSpec( + chunk_key=prev.chunk_key, + start_slot=prev.start_slot, + num_slots=prev_slots + len(spec.block_ids), + block_ids=[*prev.block_ids, *spec.block_ids], + file_offset=prev.file_offset, + token_count=prev.token_count + spec.token_count, + target_token_start=prev.target_token_start, + pos_offset=prev.pos_offset, + ) + return merged diff --git a/daser/connector/reuse.py b/daser/connector/scheduler/reuse.py similarity index 72% rename from daser/connector/reuse.py rename to daser/connector/scheduler/reuse.py index 2920199..1fd03be 100644 --- a/daser/connector/reuse.py +++ b/daser/connector/scheduler/reuse.py @@ -3,6 +3,7 @@ # Standard from abc import ABC, abstractmethod +from dataclasses import dataclass import math from typing import Any @@ -19,6 +20,27 @@ logger = init_logger(__name__) +@dataclass(frozen=True) +class StoreIntent: + """Describe one server allocation needed for pending store work.""" + + req_id: str + chunk_key: str + token_count: int + block_ids: list[int] + + +@dataclass(frozen=True) +class StoreIntentPlan: + """Return store intents together with the next strategy cursor state.""" + + intents: tuple[StoreIntent, ...] + next_key: str + next_slot: int + complete: bool + invalid: bool = False + + class CacheReuseStrategy(ABC): """Compute store keys and allocate scheduler-side store work. @@ -65,21 +87,23 @@ def ready_to_allocate(self, pending_store: PendingStore) -> bool: return len(pending_store.block_ids) >= num_slots @abstractmethod - def allocate_store( + def plan_store( self, - owner: Any, req_id: str, pending_store: PendingStore, tokens: list[int], - ) -> None: - """Allocate server-side store metadata once block IDs are known. + pending_store_ids: set[str], + ) -> StoreIntentPlan: + """Build store allocation intents once block IDs are known. Args: - owner: scheduler connector object with ``_ipc_sync``, - ``_model_id``, ``_slot_size``, and ``_pending_*`` attributes. req_id: vLLM request ID. pending_store: pending store state for this request. tokens: full prompt token IDs. + pending_store_ids: synthetic or base IDs already allocated. + + Returns: + Immutable allocation intent plan for request lifecycle execution. """ @@ -122,56 +146,43 @@ def prepare_store( return None return PendingStore(chunk_key=chunk_key, token_count=aligned_tokens) - def allocate_store( + def plan_store( self, - owner: Any, req_id: str, pending_store: PendingStore, tokens: list[int], - ) -> None: - """Allocate one store covering the whole aligned prefix. + pending_store_ids: set[str], + ) -> StoreIntentPlan: + """Plan one store covering the whole aligned prefix. Args: - owner: scheduler connector object. req_id: vLLM request ID. pending_store: pending store state for this request. tokens: full prompt token IDs. + pending_store_ids: existing allocated work IDs. + + Returns: + One whole-prefix intent, or an invalid plan on key mismatch. """ + del pending_store_ids requested_tokens = pending_store.token_count num_slots = math.ceil(requested_tokens / self._block_tokens) chunk_key = pending_store.chunk_key if chunk_key != self.store_key(tokens, requested_tokens): logger.warning("[CONNECTOR] pending store key mismatch req=%s", req_id[:8]) - owner.drop_pending_alloc(req_id) - return - try: - alloc = owner.allocate_store_chunk( - chunk_key, - requested_tokens, - ) - except Exception as exc: # noqa: BLE001 - logger.warning("[CONNECTOR] alloc_chunk failed: %s", exc) - return - if bool(alloc.get("skipped", False)): - owner.drop_pending_alloc(req_id) - logger.debug( - "[CONNECTOR] skip duplicate store req=%s key=%s", - req_id[:8], - chunk_key[:8], - ) - return - alloc["chunk_key"] = chunk_key - alloc["token_count"] = requested_tokens - alloc["num_slots"] = num_slots - alloc["block_ids"] = pending_store.block_ids[:num_slots] - owner.set_pending_store(req_id, alloc) - owner.drop_pending_alloc(req_id) - logger.debug( - "[CONNECTOR] alloc store req=%s key=%s tokens=%d/%d", - req_id, - alloc["chunk_key"][:8], - requested_tokens, - requested_tokens, + return StoreIntentPlan((), chunk_key, num_slots, True, invalid=True) + return StoreIntentPlan( + intents=( + StoreIntent( + req_id=req_id, + chunk_key=chunk_key, + token_count=requested_tokens, + block_ids=pending_store.block_ids[:num_slots], + ), + ), + next_key=chunk_key, + next_slot=num_slots, + complete=True, ) @@ -227,20 +238,23 @@ def ready_to_allocate(self, pending_store: PendingStore) -> bool: """ return len(pending_store.block_ids) > pending_store.rolling_slot_index - def allocate_store( + def plan_store( self, - owner: Any, req_id: str, pending_store: PendingStore, tokens: list[int], - ) -> None: - """Allocate one store target for each missing rolling-prefix slot. + pending_store_ids: set[str], + ) -> StoreIntentPlan: + """Plan one store target for each missing rolling-prefix slot. Args: - owner: scheduler connector object. req_id: vLLM request ID. pending_store: pending store state for this request. tokens: full prompt token IDs. + pending_store_ids: existing allocated work IDs. + + Returns: + Missing slot intents and the next rolling-prefix cursor state. """ requested_tokens = pending_store.token_count num_slots = math.ceil(requested_tokens / self._block_tokens) @@ -260,60 +274,30 @@ def allocate_store( key = next_key if slot_i >= pending_store.start_slot_index: store_id = f"{req_id}:store:{slot_i}" - if not owner.has_pending_store(store_id): + if store_id not in pending_store_ids: run.append((slot_i, key)) slot_i += 1 - - if run: - try: - allocations = owner.allocate_store_chunks( - [ - {"chunk_key": chunk_key, "token_count": self._block_tokens} - for _slot_i, chunk_key in run - ] - ) - except Exception as exc: # noqa: BLE001 - logger.warning("[CONNECTOR] alloc_chunks failed: %s", exc) - return - if len(allocations) != len(run): - logger.warning( - "[CONNECTOR] alloc_chunks returned %d allocations for %d slots", - len(allocations), - len(run), - ) - return - for (store_slot_i, chunk_key), alloc in zip( - run, - allocations, - strict=True, - ): - if bool(alloc.get("skipped", False)): - continue - alloc["chunk_key"] = str(alloc.get("chunk_key", chunk_key)) - alloc["token_count"] = self._block_tokens - alloc["num_slots"] = 1 - alloc["block_ids"] = [pending_store.block_ids[store_slot_i]] - owner.set_pending_store(f"{req_id}:store:{store_slot_i}", alloc) - - pending_store.rolling_key = key - pending_store.rolling_slot_index = slot_i + intents = tuple( + StoreIntent( + req_id=f"{req_id}:store:{store_slot_i}", + chunk_key=chunk_key, + token_count=self._block_tokens, + block_ids=[pending_store.block_ids[store_slot_i]], + ) + for store_slot_i, chunk_key in run + ) if slot_i >= num_slots: if pending_store.chunk_key and pending_store.chunk_key != key: logger.warning( "[CONNECTOR] pending store key mismatch req=%s", req_id[:8] ) - owner.drop_pending_alloc(req_id) - return - pending_store.chunk_key = key - - if slot_i >= num_slots: - owner.drop_pending_alloc(req_id) - if run: - logger.debug( - "[CONNECTOR] alloc rolling-prefix stores req=%s slots=%d", - req_id[:8], - len(run), - ) + return StoreIntentPlan((), key, slot_i, True, invalid=True) + return StoreIntentPlan( + intents=intents, + next_key=key, + next_slot=slot_i, + complete=slot_i >= num_slots, + ) def build_cache_reuse_strategy( diff --git a/daser/connector/staging.py b/daser/connector/staging.py deleted file mode 100644 index 4e20cbd..0000000 --- a/daser/connector/staging.py +++ /dev/null @@ -1,1115 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 - -from __future__ import annotations - -# Standard -from collections.abc import Callable -from dataclasses import dataclass, replace -from typing import Any, Protocol - -# Third Party -import torch - -# First Party -from daser.connector.metadata import ReqLoadSpec, ReqStoreSpec, StoreWriteSpan -from daser.logging import init_logger -from daser.ops.rope_apply import ( - apply_rope_delta_to_key_block as _apply_rope_delta_to_key_block, -) -from daser.ops.rope_apply import ( - apply_rope_delta_to_kv_key_block as _apply_rope_delta_to_kv_key_block, -) -from daser.ops.rope_apply import ( - apply_rope_delta_to_kv_key_block_table, - restore_cross_layer_kv_cache_table, -) - -DEFAULT_ROPE_DELTA_SCALE = 1.0 -DEFAULT_STORE_STAGING_BYTES = 1536 << 20 -DEFAULT_PENDING_STORE_STAGING_BYTES = 3072 << 20 -MIN_STORE_STAGING_BYTES = 64 << 20 -CROSS_LAYER_KV_CACHE_KEY = "__cross_layers__" -FUSED_RESTORE_MIN_SLOTS = 32 - -logger = init_logger(__name__) -_rope_table_cache: dict[ - tuple[torch.device, int, float, int], - tuple[torch.Tensor, torch.Tensor], -] = {} - - -class _CudaStagingLeaseOwner(Protocol): - """Pool protocol required by ``CudaStagingLease``. - - Async/thread-safety: - Implementations define their own ownership model; release is invoked - by the worker thread after CUDA IPC users are finished with the lease. - """ - - def release(self, lease: "CudaStagingLease") -> None: - """Return a lease to its owning pool. - - Args: - lease: Lease previously returned by that pool. - """ - - -@dataclass -class CudaStagingLease: - """One logical staging allocation leased from a worker-side staging pool. - - Args: - pool: Owning pool that will receive the allocation on release. - tensor: Backing tensor, possibly larger than ``nbytes``. - nbytes: Logical byte count used by the current transfer. - - Async/thread-safety: - The lease is released on the vLLM worker thread after transfer - completion. It must not be reused while an async store future owns it. - """ - - pool: _CudaStagingLeaseOwner - tensor: torch.Tensor - nbytes: int - _released: bool = False - - @property - def view(self) -> torch.Tensor: - """Return the logical byte view for the active transfer. - - Returns: - A 1-D uint8 tensor slice with ``nbytes`` elements. - - Async/thread-safety: - The returned tensor remains valid until ``release`` is called. - """ - return self.tensor[: self.nbytes] - - def release(self) -> None: - """Return the lease to the owning pool. - - Async/thread-safety: - Call once no CUDA IPC transfer can access ``view``. - """ - if self._released: - return - self._released = True - self.pool.release(self) - - -class FixedCudaStagingPool: - """Fixed-size worker-side CUDA staging buffers. - - Args: - device: Device on which staging tensors are allocated. - buffer_bytes: Size of each fixed staging buffer. - depth: Number of fixed buffers to allocate. - - Async/thread-safety: - The pool is owned by one vLLM worker thread. It never allocates after - construction, so load and store staging cannot grow into unexpected - CUDA OOM. Callers that need blocking semantics can pass a release - callback to ``acquire`` when all buffers are currently leased. - """ - - def __init__( - self, - device: torch.device, - buffer_bytes: int, - depth: int, - ) -> None: - if buffer_bytes <= 0: - raise ValueError("buffer_bytes must be positive") - if depth <= 0: - raise ValueError("depth must be positive") - self._buffer_bytes = buffer_bytes - self._buffers: list[torch.Tensor] = [ - torch.empty(buffer_bytes, dtype=torch.uint8, device=device) - for _ in range(depth) - ] - self._free_indices: list[int] = list(range(depth)) - - @property - def buffer_bytes(self) -> int: - """Return the fixed capacity of each staging buffer.""" - return self._buffer_bytes - - @property - def available(self) -> int: - """Return the number of currently free staging buffers.""" - return len(self._free_indices) - - @property - def depth(self) -> int: - """Return the number of fixed staging buffers in this pool.""" - return len(self._buffers) - - def buffer(self, index: int) -> torch.Tensor: - """Return a fixed staging backing tensor by index. - - Args: - index: Fixed buffer index to inspect. - - Returns: - The full preallocated backing tensor for the requested index. - - Async/thread-safety: - Intended for one-time CUDA IPC registration before load traffic. - Callers must not mutate the returned tensor independently of the - pool lease protocol. - """ - if index < 0 or index >= len(self._buffers): - raise ValueError(f"fixed staging buffer index out of range: {index}") - return self._buffers[index] - - def acquire( - self, - nbytes: int, - wait_for_release: Callable[[], None] | None = None, - ) -> CudaStagingLease: - """Lease one preallocated staging buffer. - - Args: - nbytes: Logical transfer byte count. - wait_for_release: Optional callback invoked once when no buffer is - currently free. Store callers use this to wait for an older - asynchronous store to finish and release its lease. - - Returns: - Lease whose view is limited to ``nbytes``. - - Raises: - ValueError: If ``nbytes`` exceeds the fixed buffer size. - RuntimeError: If all fixed buffers are currently in use. - """ - if nbytes < 0: - raise ValueError("nbytes must be non-negative") - if nbytes > self._buffer_bytes: - raise ValueError( - f"staging request {nbytes} exceeds fixed staging buffer " - f"{self._buffer_bytes}" - ) - if not self._free_indices and wait_for_release is not None: - wait_for_release() - if not self._free_indices: - raise RuntimeError("no fixed staging buffers available") - return self.acquire_index(self._free_indices[0], nbytes) - - def acquire_index(self, index: int, nbytes: int) -> CudaStagingLease: - """Lease a specific preallocated staging buffer. - - Args: - index: Fixed buffer index to lease. - nbytes: Logical transfer byte count. - - Returns: - Lease whose view is limited to ``nbytes``. - - Raises: - ValueError: If ``index`` is invalid or ``nbytes`` exceeds the fixed - buffer size. - RuntimeError: If the requested buffer is currently in use. - """ - if index < 0 or index >= len(self._buffers): - raise ValueError(f"fixed staging buffer index out of range: {index}") - if nbytes < 0: - raise ValueError("nbytes must be non-negative") - if nbytes > self._buffer_bytes: - raise ValueError( - f"staging request {nbytes} exceeds fixed staging buffer " - f"{self._buffer_bytes}" - ) - if index not in self._free_indices: - raise RuntimeError(f"fixed staging buffer {index} is not available") - self._free_indices.remove(index) - return CudaStagingLease( - pool=self, - tensor=self._buffers[index], - nbytes=nbytes, - ) - - def release(self, lease: CudaStagingLease) -> None: - """Return a fixed staging lease to the free list. - - Args: - lease: Lease previously returned by ``acquire``. - """ - for index, tensor in enumerate(self._buffers): - if tensor is lease.tensor: - if index not in self._free_indices: - self._free_indices.append(index) - self._free_indices.sort() - return - raise ValueError("lease does not belong to this fixed staging pool") - - -StoreCudaStagingPool = FixedCudaStagingPool - - -@dataclass(frozen=True) -class StagedStoreBatch: - """Worker-owned CUDA staging batch ready for server-side transfer. - - Args: - buffer: Logical uint8 staging view exported through CUDA IPC. - ready_event: CUDA event recorded after KV -> staging copies. - spans: Server write spans targeting ``buffer``. - lease: Optional reusable staging lease backing ``buffer``. - - Async/thread-safety: - The batch crosses to the connector background loop; ``lease`` is - released only after the async server transfer future completes. - """ - - buffer: torch.Tensor - ready_event: torch.cuda.Event | None - spans: list[StoreWriteSpan] - lease: CudaStagingLease | None - - -@dataclass(frozen=True) -class LoadCopyRun: - """One contiguous staging range that can be copied with one layer loop.""" - - start: int - end: int - block_ids: list[int] - pos_offset: int - - -def _load_source_key(spec: ReqLoadSpec, nbytes: int) -> tuple[str, int, int, int, int]: - """Return the source identity for a load span.""" - return (spec.chunk_key, spec.start_slot, spec.num_slots, spec.file_offset, nbytes) - - -def _store_source_key(spec: ReqStoreSpec) -> tuple[str, int, int, int, int]: - """Return the destination identity for a full store spec.""" - return ( - spec.chunk_key, - spec.start_slot, - spec.num_slots, - spec.file_offset, - len(spec.block_ids), - ) - - -def synchronize_cuda_tensor(tensor: torch.Tensor) -> None: - """Synchronize pending CUDA work for a tensor before cross-process handoff. - - Args: - tensor: Tensor whose device stream must be visible across CUDA IPC. - - Async/thread-safety: - Synchronous barrier on the current worker thread. It is intentionally - conservative until CUDA IPC event handoff is added. - """ - if tensor.is_cuda: - torch.cuda.current_stream(tensor.device).synchronize() - - -def record_cuda_event(tensor: torch.Tensor) -> torch.cuda.Event | None: - """Record the tensor's current CUDA stream for deferred synchronization. - - Args: - tensor: Tensor whose producer stream should be observed. - - Returns: - A CUDA event recorded on the current stream, or ``None`` for CPU - tensors. - - Async/thread-safety: - Must be called on the producer thread before handing ``tensor`` to a - background task. - """ - if not tensor.is_cuda: - return None - event = torch.cuda.Event(blocking=False) - event.record(torch.cuda.current_stream(tensor.device)) - return event - - -def contiguous_block_range(block_ids: list[int]) -> tuple[int, int] | None: - """Return ``(start, stop)`` when block IDs are a contiguous range.""" - if not block_ids: - return None - start = block_ids[0] - for idx, block_id in enumerate(block_ids): - if block_id != start + idx: - return None - return start, start + len(block_ids) - - -def derive_store_staging_limits(device: torch.device) -> tuple[int, int]: - """Return bounded GPU staging caps for a CUDA device. - - Args: - device: Device that will own worker-side staging tensors. - - Returns: - ``(single_batch_bytes, pending_bytes)``. The cap is based on both total - and currently free VRAM after vLLM has allocated KV cache. Defaults are - intentionally modest because staging is an IPC transport buffer, not a - persistent cache tier. - - Async/thread-safety: - Reads CUDA device properties only; safe during worker initialization. - """ - if device.type != "cuda": - return DEFAULT_STORE_STAGING_BYTES, DEFAULT_PENDING_STORE_STAGING_BYTES - props = torch.cuda.get_device_properties(device) - total = int(props.total_memory) - try: - free, _ = torch.cuda.mem_get_info(device) - free = int(free) - except (RuntimeError, TypeError, ValueError): - free = total - batch = min( - DEFAULT_STORE_STAGING_BYTES, - max(MIN_STORE_STAGING_BYTES, min(total // 50, free // 10)), - ) - pending = min( - DEFAULT_PENDING_STORE_STAGING_BYTES, - max(batch, min(total // 25, free // 5)), - ) - return batch, pending - - -def apply_rope_delta_to_key_block( - key_block: torch.Tensor, - delta: int, - rope_base: float, - rotary_dim: int, - is_neox_style: bool, -) -> None: - """Rotate an already-RoPE'd K block by a relative position delta. - - Args: - key_block: K cache block with shape [..., block_tokens, heads, head_dim]. - delta: relative RoPE position delta to apply in place. - rope_base: RoPE theta/base. - rotary_dim: number of head dimensions covered by RoPE. - is_neox_style: True for split-half rotation, False for interleaved. - - Returns: - None. ``key_block`` is modified in place. - - Async/thread-safety: - Performs tensor work on the current PyTorch stream. - """ - _apply_rope_delta_to_key_block( - key_block, - delta=delta, - rope_base=rope_base, - rotary_dim=rotary_dim, - is_neox_style=is_neox_style, - ) - - -def apply_rope_delta_to_kv_key_block( - kv_block: torch.Tensor, - delta: int, - rope_base: float, - rotary_dim: int, - is_neox_style: bool, -) -> None: - """Rotate K entries inside a full KV staging block by a RoPE delta. - - Args: - kv_block: KV staging tensor with shape - ``[blocks, layers, 2, block_tokens, heads, head_dim]``. - delta: relative RoPE position delta to apply in place. - rope_base: RoPE theta/base. - rotary_dim: number of head dimensions covered by RoPE. - is_neox_style: True for split-half rotation, False for interleaved. - - Returns: - None. Only the key slice is modified in place. - - Async/thread-safety: - Performs tensor work on the current PyTorch stream. - """ - _apply_rope_delta_to_kv_key_block( - kv_block, - delta=delta, - rope_base=rope_base, - rotary_dim=rotary_dim, - is_neox_style=is_neox_style, - ) - - -def _transform_loaded_staging_batch( - staging_by_layer: torch.Tensor, - layer_sample: torch.Tensor, - load_key_scale: float, - load_value_scale: float, - pos_offset: int, - rope_delta_scale: float, - rope_base: float, - rope_rotary_dim: int, - rope_is_neox_style: bool, -) -> None: - """Apply load-time transforms once over all staging layers in a copy run.""" - if staging_by_layer.numel() == 0 or layer_sample.dim() < 4: - return - num_slots = int(staging_by_layer.shape[0]) - num_layers = int(staging_by_layer.shape[1]) - kv_batch = staging_by_layer.view(layer_sample.dtype).view( - num_slots, - num_layers, - *layer_sample.shape, - ) - if load_key_scale != 1.0: - kv_batch[:, :, 0].mul_(load_key_scale) - if load_value_scale != 1.0: - kv_batch[:, :, 1].mul_(load_value_scale) - if ( - not pos_offset - or rope_rotary_dim <= 0 - or layer_sample.shape[-1] < rope_rotary_dim - ): - return - if kv_batch.dim() == 6 and kv_batch.is_contiguous(): - apply_rope_delta_to_kv_key_block( - kv_batch, - delta=round(pos_offset * rope_delta_scale), - rope_base=rope_base, - rotary_dim=rope_rotary_dim, - is_neox_style=rope_is_neox_style, - ) - return - if layer_sample.dim() != 4: - return - apply_rope_delta_to_key_block( - kv_batch[:, :, 0], - delta=round(pos_offset * rope_delta_scale), - rope_base=rope_base, - rotary_dim=rope_rotary_dim, - is_neox_style=rope_is_neox_style, - ) - - -def _copy_staging_to_cross_layer_kv_cache( - staging_by_layer: torch.Tensor, - cross_layer_kv_cache: torch.Tensor, - block_ids: list[int], - load_key_scale: float, - load_value_scale: float, - pos_offset: int, - rope_delta_scale: float, - rope_base: float, - rope_rotary_dim: int, - rope_is_neox_style: bool, -) -> int: - """Copy staging bytes into a vLLM cross-layer KV cache in one bulk write.""" - num_slots = len(block_ids) - layer_sample = cross_layer_kv_cache[block_ids[0], 0] - src = staging_by_layer.view(cross_layer_kv_cache.dtype).view( - num_slots, - cross_layer_kv_cache.shape[1], - *layer_sample.shape, - ) - block_range = contiguous_block_range(block_ids) - dst_contiguous = False - start = 0 - stop = 0 - if block_range is not None: - start, stop = block_range - dst_contiguous = cross_layer_kv_cache[start:stop].is_contiguous() - can_rotate_target = ( - block_range is not None - and load_key_scale == 1.0 - and load_value_scale == 1.0 - and pos_offset - and rope_rotary_dim > 0 - and layer_sample.shape[-1] >= rope_rotary_dim - and src.is_contiguous() - and dst_contiguous - ) - if can_rotate_target: - dst = cross_layer_kv_cache[start:stop] - delta = round(pos_offset * rope_delta_scale) - if num_slots >= FUSED_RESTORE_MIN_SLOTS: - _restore_cross_layer_with_tables( - src, - dst, - delta=delta, - rope_base=rope_base, - rotary_dim=rope_rotary_dim, - is_neox_style=rope_is_neox_style, - ) - return 1 - dst.copy_(src) - _apply_rope_delta_with_tables( - dst, - delta=delta, - rope_base=rope_base, - rotary_dim=rope_rotary_dim, - is_neox_style=rope_is_neox_style, - ) - return 1 - _transform_loaded_staging_batch( - staging_by_layer, - layer_sample=layer_sample, - load_key_scale=load_key_scale, - load_value_scale=load_value_scale, - pos_offset=pos_offset, - rope_delta_scale=rope_delta_scale, - rope_base=rope_base, - rope_rotary_dim=rope_rotary_dim, - rope_is_neox_style=rope_is_neox_style, - ) - if block_range is None: - block_index = torch.tensor( - block_ids, - dtype=torch.long, - device=staging_by_layer.device, - ) - cross_layer_kv_cache.index_copy_(0, block_index, src) - else: - start, stop = block_range - cross_layer_kv_cache[start:stop].copy_(src) - return 1 - - -def _apply_rope_delta_with_tables( - kv_block: torch.Tensor, - delta: int, - rope_base: float, - rotary_dim: int, - is_neox_style: bool, -) -> None: - """Apply RoPE using cached trig tables.""" - if kv_block.device.type != "cuda" or kv_block.dtype not in ( - torch.bfloat16, - torch.float16, - torch.float32, - ): - raise ValueError("TileLang RoPE restore requires CUDA fp16/bf16/fp32 KV") - cos_table, sin_table = _get_rope_delta_tables( - kv_block.device, - delta=delta, - rope_base=rope_base, - rotary_dim=rotary_dim, - ) - apply_rope_delta_to_kv_key_block_table( - kv_block, - cos_table=cos_table, - sin_table=sin_table, - rotary_dim=rotary_dim, - is_neox_style=is_neox_style, - ) - - -def _restore_cross_layer_with_tables( - src_kv: torch.Tensor, - dst_kv: torch.Tensor, - delta: int, - rope_base: float, - rotary_dim: int, - is_neox_style: bool, -) -> None: - """Restore cross-layer KV using cached trig tables.""" - cos_table, sin_table = _get_rope_delta_tables( - src_kv.device, - delta=delta, - rope_base=rope_base, - rotary_dim=rotary_dim, - ) - restore_cross_layer_kv_cache_table( - src_kv, - dst_kv, - cos_table=cos_table, - sin_table=sin_table, - rotary_dim=rotary_dim, - is_neox_style=is_neox_style, - ) - - -def _get_rope_delta_tables( - device: torch.device, - delta: int, - rope_base: float, - rotary_dim: int, -) -> tuple[torch.Tensor, torch.Tensor]: - """Return cached fp32 RoPE delta cosine/sine tables.""" - key = (device, int(delta), float(rope_base), int(rotary_dim)) - cached = _rope_table_cache.get(key) - if cached is not None: - return cached - inv_freq = 1.0 / ( - rope_base - ** ( - torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=device) - / rotary_dim - ) - ) - freqs = int(delta) * inv_freq - tables = (freqs.cos().contiguous(), freqs.sin().contiguous()) - _rope_table_cache[key] = tables - return tables - - -def copy_staging_to_kv_cache( - staging: torch.Tensor, - kv_caches: dict[str, torch.Tensor], - layer_names: list[str], - block_ids: list[int], - slot_size: int, - load_key_scale: float = 1.0, - load_value_scale: float = 1.0, - pos_offset: int = 0, - rope_delta_scale: float = DEFAULT_ROPE_DELTA_SCALE, - rope_base: float = 10000.0, - rope_rotary_dim: int = 0, - rope_is_neox_style: bool = True, -) -> int: - """Copy slot-major staging bytes into vLLM KV cache blocks. - - Args: - staging: Contiguous uint8 tensor containing whole request KV bytes. - kv_caches: Per-layer vLLM KV cache tensors. - layer_names: Layer iteration order matching on-disk layout. - block_ids: vLLM KV block IDs corresponding to staging slots. - slot_size: Total bytes for all layers in one slot. - load_key_scale: Optional multiplier for loaded K tensors. - load_value_scale: Optional multiplier for loaded V tensors. - pos_offset: Position delta for loaded chunk reuse. - rope_delta_scale: Multiplier applied to pos_offset before RoPE update. - rope_base: RoPE theta/base. - rope_rotary_dim: Number of head dimensions covered by RoPE. - rope_is_neox_style: True for split-half rotation, False for interleaved. - - Returns: - Number of layer-level copy operations issued. - - Async/thread-safety: - Synchronous GPU tensor copies on the vLLM worker thread. - """ - if not block_ids or not layer_names: - return 0 - num_layers = len(layer_names) - layer_size = slot_size // num_layers - num_slots = len(block_ids) - staging_by_layer = staging.view(num_slots, num_layers, layer_size) - cross_layer_kv_cache = kv_caches.get(CROSS_LAYER_KV_CACHE_KEY) - if ( - cross_layer_kv_cache is not None - and cross_layer_kv_cache.dim() >= 6 - and cross_layer_kv_cache.shape[1] == num_layers - ): - return _copy_staging_to_cross_layer_kv_cache( - staging_by_layer=staging_by_layer, - cross_layer_kv_cache=cross_layer_kv_cache, - block_ids=block_ids, - load_key_scale=load_key_scale, - load_value_scale=load_value_scale, - pos_offset=pos_offset, - rope_delta_scale=rope_delta_scale, - rope_base=rope_base, - rope_rotary_dim=rope_rotary_dim, - rope_is_neox_style=rope_is_neox_style, - ) - first_kv = next( - (kv_caches[name] for name in layer_names if kv_caches.get(name) is not None), - None, - ) - if first_kv is not None: - layer_sample = ( - first_kv[:, block_ids[0]] if first_kv.dim() >= 2 else first_kv[block_ids[0]] - ) - _transform_loaded_staging_batch( - staging_by_layer, - layer_sample=layer_sample, - load_key_scale=load_key_scale, - load_value_scale=load_value_scale, - pos_offset=pos_offset, - rope_delta_scale=rope_delta_scale, - rope_base=rope_base, - rope_rotary_dim=rope_rotary_dim, - rope_is_neox_style=rope_is_neox_style, - ) - block_range = contiguous_block_range(block_ids) - block_index = ( - None - if block_range is not None - else torch.tensor(block_ids, dtype=torch.long, device=staging.device) - ) - - copies = 0 - for layer_idx, layer_name in enumerate(layer_names): - kv_tensor = kv_caches.get(layer_name) - if kv_tensor is None: - continue - # KV tensors are either block-major ([blocks, ...]) or kv-major - # ([2, blocks, ...]); block_dim points at the block axis. - block_dim = 1 if kv_tensor.dim() >= 2 else 0 - sample = ( - kv_tensor[:, block_ids[0]] if block_dim == 1 else kv_tensor[block_ids[0]] - ) - src = ( - staging_by_layer[:, layer_idx, :] - .view(kv_tensor.dtype) - .view(num_slots, *sample.shape) - ) - # staging is slot-major (slots first); align it to the block axis. - src = src.movedim(0, block_dim) - if block_range is None: - if block_index is None: - raise RuntimeError("block_index is required for non-contiguous IDs") - kv_tensor.index_copy_(block_dim, block_index, src) - else: - start, stop = block_range - if block_dim == 1: - kv_tensor[:, start:stop].copy_(src) - else: - kv_tensor[start:stop].copy_(src) - copies += 1 - return copies - - -def copy_kv_cache_to_staging( - staging: torch.Tensor, - kv_layer: torch.Tensor, - layer_idx: int, - block_ids: list[int], - num_layers: int, - slot_size: int, - block_index: torch.Tensor | None = None, -) -> None: - """Copy one vLLM KV layer for requested blocks into slot-major staging. - - Args: - staging: Contiguous uint8 tensor with slot-major DaseR layout. - kv_layer: vLLM KV cache tensor for one attention layer. - layer_idx: Index of ``kv_layer`` in the DaseR on-disk layer order. - block_ids: vLLM KV block IDs to persist. - num_layers: Total number of KV layers in the model. - slot_size: Total bytes for all layers in one slot. - block_index: Optional prebuilt CUDA/CPU tensor containing block IDs. - - Async/thread-safety: - Synchronous GPU tensor copies on the vLLM worker thread. - """ - if not block_ids: - return - layer_size = slot_size // num_layers - num_slots = len(block_ids) - staging_by_layer = staging.view(num_slots, num_layers, layer_size) - if block_index is None: - block_index = torch.tensor(block_ids, dtype=torch.long, device=kv_layer.device) - if kv_layer.dim() >= 2: - block_range = contiguous_block_range(block_ids) - if block_range is None: - src = kv_layer.index_select(1, block_index).movedim(1, 0) - else: - start, stop = block_range - src = kv_layer[:, start:stop].movedim(1, 0) - else: - block_range = contiguous_block_range(block_ids) - if block_range is None: - src = kv_layer.index_select(0, block_index) - else: - start, stop = block_range - src = kv_layer[start:stop] - dst = ( - staging_by_layer[:, layer_idx, :] - .view(kv_layer.dtype) - .view(num_slots, *src.shape[1:]) - ) - dst.copy_(src) - - -def copy_cross_layer_kv_cache_to_staging( - staging: torch.Tensor, - kv_cache: torch.Tensor, - block_ids: list[int], - num_layers: int, - slot_size: int, - block_index: torch.Tensor | None = None, -) -> None: - """Copy vLLM cross-layer KV blocks into slot-major staging bytes. - - Args: - staging: Contiguous uint8 tensor with slot-major DaseR layout. - kv_cache: vLLM cross-layer KV cache tensor with blocks as dim 0 and - layers as dim 1. - block_ids: vLLM KV block IDs to persist. - num_layers: Total number of KV layers in the model. - slot_size: Total bytes for all layers in one slot. - block_index: Optional prebuilt tensor containing block IDs. - - Async/thread-safety: - Synchronous GPU tensor copy on the vLLM worker thread. - """ - if not block_ids: - return - layer_size = slot_size // num_layers - num_slots = len(block_ids) - staging_by_layer = staging.view(num_slots, num_layers, layer_size) - block_range = contiguous_block_range(block_ids) - if block_range is None: - if block_index is None: - block_index = torch.tensor( - block_ids, - dtype=torch.long, - device=kv_cache.device, - ) - src = kv_cache.index_select(0, block_index) - else: - start, stop = block_range - src = kv_cache[start:stop] - dst = staging_by_layer.view(kv_cache.dtype).view( - num_slots, - num_layers, - *src.shape[2:], - ) - dst.copy_(src) - - -def build_load_read_plan( - reqs_to_load: dict[str, ReqLoadSpec], - slot_size: int, - include_req_ids: bool = False, -) -> tuple[int, list[dict[str, int]], list[Any]]: - """Build a combined transfer-load plan for one forward step. - - Args: - reqs_to_load: request ID to load spec from scheduler metadata. - slot_size: bytes per vLLM KV slot. - include_req_ids: when True, include request IDs in per-request ranges - for worker-side request completion tracking. - - Returns: - ``(total_bytes, spans, per_req_ranges)`` where spans target one - combined staging tensor and per-request ranges map slices back to - their original load specs. - """ - total_bytes = 0 - spans: list[dict[str, int]] = [] - per_req_ranges: list[Any] = [] - source_ranges: dict[tuple[str, int, int, int, int], tuple[int, int]] = {} - for req_id, spec in reqs_to_load.items(): - num_slots = len(spec.block_ids) - if num_slots == 0: - continue - nbytes = num_slots * slot_size - source_key = _load_source_key(spec, nbytes) - existing = source_ranges.get(source_key) - if existing is None: - start = total_bytes - end = start + nbytes - spans.append( - { - "target_offset": start, - "nbytes": nbytes, - "file_offset": spec.file_offset, - } - ) - source_ranges[source_key] = (start, end) - total_bytes = end - else: - start, end = existing - if include_req_ids: - per_req_ranges.append((start, end, req_id, spec)) - else: - per_req_ranges.append((start, end, spec)) - return total_bytes, spans, per_req_ranges - - -def build_load_read_batches( - reqs_to_load: dict[str, ReqLoadSpec], - slot_size: int, - max_batch_bytes: int, - include_req_ids: bool = False, -) -> list[tuple[int, list[dict[str, int]], list[Any]]]: - """Build bounded load staging plans for one forward step. - - Args: - reqs_to_load: request ID to load spec from scheduler metadata. - slot_size: bytes per vLLM KV slot. - max_batch_bytes: Maximum staging bytes for one transfer batch. - include_req_ids: when True, include request IDs in per-request ranges - for worker-side request completion tracking. - - Returns: - List of ``build_load_read_plan``-style tuples. Individual requests are - split at block boundaries when one request exceeds the staging cap. - - Async/thread-safety: - Pure CPU helper. It does not mutate connector state. - """ - if slot_size <= 0: - raise ValueError("slot_size must be positive") - if max_batch_bytes <= 0: - raise ValueError("max_batch_bytes must be positive") - max_slots = max(1, max_batch_bytes // slot_size) - batches: list[tuple[int, list[dict[str, int]], list[Any]]] = [] - current: dict[str, ReqLoadSpec] = {} - current_slots = 0 - synthetic_id = 0 - - def flush() -> None: - nonlocal current, current_slots - if current: - batches.append( - build_load_read_plan( - current, - slot_size, - include_req_ids=include_req_ids, - ) - ) - current = {} - current_slots = 0 - - for req_id, spec in reqs_to_load.items(): - cursor = 0 - while cursor < len(spec.block_ids): - if current_slots >= max_slots: - flush() - available = max_slots - current_slots - take = min(available, len(spec.block_ids) - cursor) - if take <= 0: - flush() - continue - part = spec.block_ids[cursor : cursor + take] - batch_spec = replace( - spec, - start_slot=spec.start_slot + cursor, - num_slots=take, - block_ids=part, - file_offset=spec.file_offset + cursor * slot_size, - ) - key = ( - req_id - if cursor == 0 and take == len(spec.block_ids) - else (f"{req_id}#{synthetic_id}") - ) - synthetic_id += 1 - current[key] = batch_spec - current_slots += take - cursor += take - flush() - return batches - - -def build_load_copy_runs( - per_req_ranges: list[tuple[int, int, ReqLoadSpec]], -) -> list[LoadCopyRun]: - """Merge adjacent load ranges that share the same KV transform. - - Args: - per_req_ranges: Per-request staging ranges from ``build_load_read_plan``. - - Returns: - Ordered copy runs. Each run covers a contiguous staging slice and the - matching flattened block ID list. - - Async/thread-safety: - Pure CPU helper. It does not mutate connector state. - """ - runs: list[LoadCopyRun] = [] - run_start = -1 - run_end = -1 - run_pos_offset = 0 - run_block_ids: list[int] = [] - - def flush() -> None: - nonlocal run_start, run_end, run_pos_offset, run_block_ids - if run_start >= 0 and run_block_ids: - runs.append( - LoadCopyRun( - start=run_start, - end=run_end, - block_ids=run_block_ids, - pos_offset=run_pos_offset, - ) - ) - run_start = -1 - run_end = -1 - run_pos_offset = 0 - run_block_ids = [] - - for item in per_req_ranges: - if len(item) == 3: - start, end, spec = item - else: - start, end, _req_id, spec = item - if not spec.block_ids: - continue - if run_start >= 0 and start == run_end and spec.pos_offset == run_pos_offset: - run_end = end - run_block_ids.extend(spec.block_ids) - continue - flush() - run_start = start - run_end = end - run_pos_offset = spec.pos_offset - run_block_ids = list(spec.block_ids) - flush() - return runs - - -def build_staging_store_batches( - reqs_to_store: dict[str, ReqStoreSpec], - slot_size: int, - max_batch_bytes: int = DEFAULT_STORE_STAGING_BYTES, -) -> list[tuple[list[int], list[StoreWriteSpan]]]: - """Split store requests into bounded slot-major staging batches. - - Args: - reqs_to_store: Store specs keyed by request ID. - slot_size: DaseR bytes per KV slot. - max_batch_bytes: Maximum GPU staging bytes per batch. - - Returns: - List of ``(block_ids, spans)`` batches. Span source offsets are relative - to that batch's staging tensor. - - Async/thread-safety: - Pure CPU helper. It does not mutate connector state. - """ - if slot_size <= 0: - raise ValueError("slot_size must be positive") - max_slots = max(1, max_batch_bytes // slot_size) - batches: list[tuple[list[int], list[StoreWriteSpan]]] = [] - batch_blocks: list[int] = [] - batch_spans: list[StoreWriteSpan] = [] - written_specs: set[tuple[str, int, int, int, int]] = set() - - def flush_batch() -> None: - nonlocal batch_blocks, batch_spans - if batch_blocks: - batches.append((batch_blocks, batch_spans)) - batch_blocks = [] - batch_spans = [] - - for spec in reqs_to_store.values(): - source_key = _store_source_key(spec) - if source_key in written_specs: - continue - written_specs.add(source_key) - cursor = 0 - while cursor < len(spec.block_ids): - if len(batch_blocks) >= max_slots: - flush_batch() - available = max_slots - len(batch_blocks) - take = min(available, len(spec.block_ids) - cursor) - if take <= 0: - flush_batch() - continue - source_slot = len(batch_blocks) - part = spec.block_ids[cursor : cursor + take] - batch_blocks.extend(part) - batch_spans.append( - StoreWriteSpan( - source_offset=source_slot * slot_size, - nbytes=take * slot_size, - file_offset=spec.file_offset + cursor * slot_size, - chunk_key=spec.chunk_key, - start_slot=spec.start_slot, - num_slots=spec.num_slots, - ) - ) - cursor += take - flush_batch() - return batches diff --git a/daser/connector/worker.py b/daser/connector/worker.py deleted file mode 100644 index 35b2375..0000000 --- a/daser/connector/worker.py +++ /dev/null @@ -1,2179 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 - -from __future__ import annotations - -# Standard -import asyncio -from collections import deque -from dataclasses import dataclass, replace -import os -import queue -import threading -import time -from typing import TYPE_CHECKING, Any - -# Third Party -import cupy -import torch -from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole - -if TYPE_CHECKING: - # Third Party - from vllm.attention import AttentionMetadata - from vllm.forward_context import ForwardContext - -# First Party -from daser.connector.helpers import base_req_id -from daser.connector.metadata import ( - DaserConnectorMeta, - ReqStoreSpec, - StoreWriteSpan, -) -from daser.connector.staging import ( - CROSS_LAYER_KV_CACHE_KEY, - DEFAULT_PENDING_STORE_STAGING_BYTES, - DEFAULT_STORE_STAGING_BYTES, - FUSED_RESTORE_MIN_SLOTS, - CudaStagingLease, - FixedCudaStagingPool, - StagedStoreBatch, - StoreCudaStagingPool, -) -from daser.connector.staging import ( - build_load_copy_runs as _build_load_copy_runs, -) -from daser.connector.staging import ( - build_load_read_batches as _build_load_read_batches, -) -from daser.connector.staging import ( - build_staging_store_batches as _build_staging_store_batches, -) -from daser.connector.staging import ( - copy_cross_layer_kv_cache_to_staging as _copy_cross_layer_kv_cache_to_staging, -) -from daser.connector.staging import ( - copy_kv_cache_to_staging as _copy_kv_cache_to_staging, -) -from daser.connector.staging import ( - copy_staging_to_kv_cache as _copy_staging_to_kv_cache, -) -from daser.connector.staging import ( - derive_store_staging_limits as _derive_store_staging_limits, -) -from daser.connector.staging import ( - record_cuda_event as _record_cuda_event, -) -from daser.connector.staging import ( - synchronize_cuda_tensor as _synchronize_cuda_tensor, -) -from daser.logging import init_logger -from daser.ops.rope_apply import ( - apply_rope_delta_to_key_block as _apply_rope_delta_to_key_block, -) -from daser.ops.rope_apply import ( - apply_rope_delta_to_kv_key_block_table, - restore_cross_layer_kv_cache_table, -) -from daser.transfer.cuda_ipc import ( - cuda_array_device_id, - cuda_array_pointer, - export_cuda_ipc_handle, -) - -logger = init_logger(__name__) - -_ROPE_WARMUP_BLOCKS = 1 -_LOAD_REQUEST_MAX_INFLIGHT = 8 -_LOAD_DISPATCH_WAIT_TIMEOUT_S = 0.001 -_LOAD_STAGING_RESERVE_BYTES = 1 << 30 -_MIN_STORE_STAGING_POOL_DEPTH = 1 -_LoadBatch = tuple[int, list[dict[str, int]], list[Any]] - - -def _rank_lane_offset( - start_slot: int, - local_slot_size: int, - rank_stride_bytes: int, - tp_rank: int, -) -> int: - """Return the physical offset for one rank-local logical slot range. - - Args: - start_slot: First server-owned logical slot. - local_slot_size: Bytes stored per logical slot by one TP rank. - rank_stride_bytes: Byte distance between adjacent rank lanes. - tp_rank: Current vLLM tensor-parallel rank. - - Returns: - Physical store offset for ``start_slot`` in ``tp_rank``'s lane. - - Async/thread-safety: - Pure arithmetic used on the worker thread before IPC submission. - """ - return tp_rank * rank_stride_bytes + start_slot * local_slot_size - - -def _local_slot_bytes(connector: Any) -> int: - """Return per-rank slot bytes, falling back for TP=1 test probes.""" - local_slot_size = int(getattr(connector, "_local_slot_size", 0)) - return local_slot_size or int(connector._slot_size) # noqa: SLF001 - - -def _validate_tp_layout( - local_slot_size: int, - storage_slot_size: int, - tp_size: int, - server_tp_size: int, - tp_rank: int, - rank_stride_bytes: int = 0, -) -> None: - """Validate worker KV geometry against the server-owned TP layout. - - Args: - local_slot_size: Slot bytes measured from the worker KV tensor. - storage_slot_size: Aggregate slot bytes reported by the server. - tp_size: vLLM worker tensor-parallel size. - server_tp_size: Tensor-parallel size reported by the server. - tp_rank: Current vLLM tensor-parallel rank. - rank_stride_bytes: Byte distance between server-owned rank lanes. - - Raises: - ValueError: if rank counts or slot geometry do not match. - - Async/thread-safety: - Pure startup validation called before request traffic. - """ - if tp_size <= 0 or not 0 <= tp_rank < tp_size: - raise ValueError(f"invalid TP rank {tp_rank} for size {tp_size}") - if not storage_slot_size: - return - if server_tp_size != tp_size: - raise ValueError( - f"vLLM TP size {tp_size} does not match DaseR TP size {server_tp_size}" - ) - if local_slot_size * tp_size != storage_slot_size: - raise ValueError( - "worker KV slot geometry does not match DaseR storage layout: " - f"local={local_slot_size} tp={tp_size} storage={storage_slot_size}" - ) - if tp_size > 1 and rank_stride_bytes <= 0: - raise ValueError("DaseR runtime config is missing TP rank stride") - - -def _store_staging_pool_depth(buffer_bytes: int, pending_limit_bytes: int) -> int: - """Return fixed store staging pool depth for the configured byte budget. - - Args: - buffer_bytes: Capacity of one fixed staging buffer. - pending_limit_bytes: Total pending store staging byte budget. - - Returns: - Number of fixed store staging buffers to preallocate. - - Async/thread-safety: - Pure helper used during worker-side pool initialization. - """ - if buffer_bytes <= 0: - raise ValueError("buffer_bytes must be positive") - if pending_limit_bytes <= 0: - return _MIN_STORE_STAGING_POOL_DEPTH - return max(_MIN_STORE_STAGING_POOL_DEPTH, pending_limit_bytes // buffer_bytes) - - -def _load_staging_pool_depth( - buffer_bytes: int, - pending_limit_bytes: int, - device: torch.device, -) -> int: - """Return fixed load staging depth under memory and inflight constraints. - - Args: - buffer_bytes: Capacity of one fixed staging buffer. - pending_limit_bytes: Existing staging byte budget from store limits. - device: CUDA device used for staging allocation. - - Returns: - Number of fixed load staging buffers to preallocate. - - Async/thread-safety: - Pure helper except for querying CUDA free memory. Called during worker - initialization before request traffic. - """ - depth = min( - _LOAD_REQUEST_MAX_INFLIGHT, - _store_staging_pool_depth(buffer_bytes, pending_limit_bytes), - ) - if device.type != "cuda": - return max(1, depth) - try: - free_bytes, _total_bytes = torch.cuda.mem_get_info(device) - except Exception: # noqa: BLE001 - return max(1, depth) - usable_bytes = max(0, int(free_bytes) - _LOAD_STAGING_RESERVE_BYTES) - memory_depth = max(1, usable_bytes // buffer_bytes) - return max(1, min(depth, memory_depth)) - - -@dataclass -class _DeferredFinishedSave: - """Store work held until vLLM reports a request as finished.""" - - commit_keys: set[str] - reqs_to_store: dict[str, ReqStoreSpec] - submitted: bool = False - future: Any | None = None - - -@dataclass -class _SaveFuture: - """One background save future and the staging it keeps alive. - - Attributes: - future: future returned by ``asyncio.run_coroutine_threadsafe``. - staging_bytes: GPU staging bytes held alive until completion. - lease: optional reusable staging lease released after completion. - """ - - future: Any - staging_bytes: int - lease: CudaStagingLease | None - - def release(self) -> None: - """Release the reusable staging lease, if any.""" - if self.lease is not None: - self.lease.release() - self.lease = None - - -@dataclass -class _PendingLoad: - """One background load future tracked until vLLM can resume the request. - - Attributes: - future: Future running the cross-layer load work. - block_ids: vLLM KV block IDs targeted by the load. - lease: optional staging lease held until the load completes. - """ - - future: Any - block_ids: list[int] - lease: CudaStagingLease | None - - def release(self) -> None: - """Release the reusable staging lease, if any.""" - if self.lease is not None: - self.lease.release() - self.lease = None - - -@dataclass -class _InflightLoadBatch: - """One submitted load batch and its fixed staging lease. - - Attributes: - buffer_index: Fixed load staging buffer index. - total_bytes: Logical bytes in this transfer batch. - per_req_ranges: Ranges used to restore staging bytes into vLLM KV cache. - staging_lease: Fixed staging lease retained until restore completes. - future: Future returned by the load event loop. - submitted_at: Wall-clock timestamp used for wait timing. - """ - - buffer_index: int - total_bytes: int - per_req_ranges: list[Any] - staging_lease: CudaStagingLease - future: Any - submitted_at: float - - -@dataclass -class _InflightRequestLoad: - """One request-level load submitted by the dispatcher. - - Attributes: - item: Original queued request item. - buffer_index: Fixed staging buffer owned by this request state. - batches: Bounded read batches for this request. - next_batch: Index of the next unsubmitted read batch. - remaining_batches: Number of read batches not yet consumed. - active: Current in-flight read batch for this request, if any. - completed: Consumed batch timing rows for logging. - """ - - item: "_QueuedLoadRequest" - buffer_index: int - batches: list[_LoadBatch] - next_batch: int - remaining_batches: int - active: _InflightLoadBatch | None - completed: list["_ConsumedLoadBatch"] - - -@dataclass -class _ConsumedLoadBatch: - """Timing and accounting data for one restored load batch.""" - - buffer_index: int - bytes: int - copies: int - copy_runs: int - ipc_ms: float - wait_ms: float - copy_ms: float - transfer_open_ms: float - transfer_load_ms: float - transfer_sync_ms: float - l1_hits: int - l1_misses: int - l2_reads: int - - -class _ImmediateLoadError: - """Completed future used when a load cannot be submitted. - - Args: - message: Error message raised when the future is collected. - - Async/thread-safety: - Immutable testable stand-in for a failed background future. It does not - spawn threads or perform IO. - """ - - def __init__(self, message: str) -> None: - self._message = message - - def done(self) -> bool: - """Return True because this failed future is already complete.""" - return True - - def result(self, timeout: float | None = None) -> None: - """Raise the seeded load-start failure. - - Args: - timeout: Ignored timeout for ``Future`` API compatibility. - """ - del timeout - raise RuntimeError(self._message) - - -class _RequestLoadFuture: - """Small future used to release request loads independently. - - Async/thread-safety: - The connector load executor marks completion and the vLLM worker thread - polls ``done``/``result`` from ``get_finished``. ``threading.Event`` - provides cross-thread visibility for the result state. - """ - - def __init__(self) -> None: - self._event = threading.Event() - self._error: BaseException | None = None - - def done(self) -> bool: - """Return whether this request load has completed.""" - return self._event.is_set() - - def result(self, timeout: float | None = None) -> None: - """Wait for completion and raise a stored load error, if any. - - Args: - timeout: Optional timeout in seconds. - """ - if not self._event.wait(timeout): - raise TimeoutError("request load did not complete before timeout") - if self._error is not None: - raise self._error - - def set_result(self) -> None: - """Mark this request load as successful.""" - self._event.set() - - def set_exception(self, error: BaseException) -> None: - """Mark this request load as failed. - - Args: - error: Exception to re-raise from ``result``. - """ - self._error = error - self._event.set() - - -@dataclass -class _QueuedLoadRequest: - """One request waiting in the worker-side load queue. - - Attributes: - req_id: Base vLLM request ID used for completion. - spec_id: Scheduler metadata ID for this load spec. - spec: Load spec to restore into the vLLM KV cache. - future: Per-request completion future observed by ``get_finished``. - """ - - req_id: str - spec_id: str - spec: Any - future: _RequestLoadFuture - - -class LoadRequestDispatcher: - """Bound request-level load concurrency by max in-flight and staging depth. - - Args: - max_inflight: Maximum request loads allowed in flight. - staging_depth: Number of fixed staging buffers available to request - loads. - - Async/thread-safety: - The dispatcher object is owned by the connector load asyncio loop. Test - helpers may call its pure synchronous scheduling helpers directly. - """ - - def __init__( - self, - max_inflight: int, - staging_depth: int, - ) -> None: - self._effective_inflight = max(1, min(max_inflight, staging_depth)) - self._free_buffers: deque[int] = deque(range(self._effective_inflight)) - - @property - def effective_inflight(self) -> int: - """Return the active request limit after applying staging depth.""" - return self._effective_inflight - - def submit_ready( - self, - connector: Any, - queued: list[_QueuedLoadRequest], - sample_tensor: torch.Tensor, - ) -> list[_InflightRequestLoad]: - """Submit queued requests while in-flight slots and buffers are free. - - Args: - connector: Worker connector that owns submit helpers. - queued: Mutable FIFO list of queued request work. - sample_tensor: Representative KV cache tensor. - - Returns: - Newly submitted request states. - """ - submitted: list[_InflightRequestLoad] = [] - while queued and self._free_buffers: - buffer_index = self._free_buffers.popleft() - item = queued.pop(0) - state = connector._submit_request_load_for_dispatcher( # noqa: SLF001 - item, - buffer_index, - sample_tensor, - ) - if state.active is None: - self._free_buffers.append(state.buffer_index) - continue - submitted.append(state) - return submitted - - def consume_ready( - self, - connector: Any, - active: list[_InflightRequestLoad], - sample_tensor: torch.Tensor, - ) -> list[_InflightRequestLoad]: - """Consume completed request loads and release their buffers. - - Args: - connector: Worker connector that owns consume helpers. - active: Mutable list of active request load states. - sample_tensor: Representative KV cache tensor. - - Returns: - States consumed during this call. - """ - consumed: list[_InflightRequestLoad] = [] - for state in list(active): - active_batch = state.active - if active_batch is None: - active.remove(state) - self._free_buffers.append(state.buffer_index) - consumed.append(state) - continue - if not active_batch.future.done(): - continue - reusable_buffer, request_done = connector._consume_dispatcher_load( # noqa: SLF001 - state, - sample_tensor, - ) - if request_done: - active.remove(state) - self._free_buffers.append(reusable_buffer) - consumed.append(state) - return consumed - - -def _cuda_allocation_base_and_offset(device_ptr: int) -> tuple[int, int]: - """Return CUDA allocation base pointer and byte offset for ``device_ptr``. - - Args: - device_ptr: CUDA device pointer exported through IPC. - - Returns: - Tuple of ``(allocation_base_ptr, byte_offset)``. - """ - try: - from cuda.bindings import driver as cuda_driver - - result, base_ptr, _allocation_size = cuda_driver.cuMemGetAddressRange( - device_ptr - ) - if result == cuda_driver.CUresult.CUDA_SUCCESS: - base = int(base_ptr) - return base, int(device_ptr) - base - except Exception as exc: # noqa: BLE001 - logger.debug("[CONNECTOR] cuMemGetAddressRange failed: %s", exc) - return int(device_ptr), 0 - - -def _warm_rope_apply_backends( - device: torch.device, - dtype: torch.dtype, - block_tokens: int, - heads: int, - head_dim: int, - rotary_dim: int, - rope_base: float, - is_neox_style: bool, -) -> None: - """Warm dynamic-shape RoPE apply operators. - - Args: - device: device that owns the worker KV cache. - dtype: KV cache dtype. - block_tokens: tokens per cache block. - heads: number of KV heads. - head_dim: per-head dimension. - rotary_dim: number of dimensions covered by RoPE. - rope_base: RoPE theta/base. - is_neox_style: True for split-half rotation, False for interleaved. - - Async/thread-safety: - Runs synchronously during worker KV cache registration, before request - traffic starts. It launches CUDA work on the current stream. TileLang - failures are surfaced to avoid silently entering a slow restore path. - """ - if device.type != "cuda" or rotary_dim <= 0 or head_dim < rotary_dim: - return - sample = torch.empty( - (_ROPE_WARMUP_BLOCKS, block_tokens, heads, head_dim), - dtype=dtype, - device=device, - ) - _apply_rope_delta_to_key_block( - sample, - delta=1, - rope_base=rope_base, - rotary_dim=rotary_dim, - is_neox_style=is_neox_style, - ) - torch.cuda.synchronize(device) - - -def _warm_cross_layer_restore_backends( - device: torch.device, - dtype: torch.dtype, - layers: int, - block_tokens: int, - heads: int, - head_dim: int, - rotary_dim: int, - rope_base: float, - is_neox_style: bool, -) -> None: - """Warm cross-layer staging restore TileLang kernels. - - Args: - device: device that owns the worker KV cache. - dtype: KV cache dtype. - layers: number of model KV layers. - block_tokens: tokens per cache block. - heads: number of KV heads. - head_dim: per-head dimension. - rotary_dim: number of dimensions covered by RoPE. - rope_base: RoPE theta/base. - is_neox_style: True for split-half rotation, False for interleaved. - - Async/thread-safety: - Runs synchronously during worker KV cache registration, before request - traffic starts. TileLang import/compile failures are surfaced to avoid - silently entering a slow restore path. - """ - if device.type != "cuda" or rotary_dim <= 0 or head_dim < rotary_dim: - return - inv_freq = 1.0 / ( - rope_base - ** ( - torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=device) - / rotary_dim - ) - ) - freqs = inv_freq - cos_table = freqs.cos().contiguous() - sin_table = freqs.sin().contiguous() - for blocks, use_fused_restore in ( - (_ROPE_WARMUP_BLOCKS, False), - (FUSED_RESTORE_MIN_SLOTS, True), - ): - sample = torch.empty( - blocks, - layers, - 2, - block_tokens, - heads, - head_dim, - dtype=dtype, - device=device, - ) - if use_fused_restore: - dst = torch.empty_like(sample) - restore_cross_layer_kv_cache_table( - sample, - dst, - cos_table=cos_table, - sin_table=sin_table, - rotary_dim=rotary_dim, - is_neox_style=is_neox_style, - ) - else: - apply_rope_delta_to_kv_key_block_table( - sample, - cos_table=cos_table, - sin_table=sin_table, - rotary_dim=rotary_dim, - is_neox_style=is_neox_style, - ) - torch.cuda.synchronize(device) - - -class WorkerConnectorMixin: - """Worker-role vLLM connector behavior. - - Async/thread-safety: - Public methods are called on vLLM worker threads. Blocking NVMe work is - submitted to the connector's background asyncio loop. - """ - - def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]) -> None: - """Register the per-layer KV cache tensors. - - Args: - kv_caches: dict mapping layer_name -> KV tensor. - """ - self._kv_caches = kv_caches - self._layer_names = list(kv_caches.keys()) - self._layer_idx_map = {name: idx for idx, name in enumerate(self._layer_names)} - sample = next(iter(kv_caches.values()), None) - if kv_caches: - ( - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - ) = _derive_store_staging_limits(sample.device) - logger.info( - "[CONNECTOR] register_kv_caches: %d layers, first shape=%s dtype=%s", - len(kv_caches), - sample.shape, - sample.dtype, - ) - logger.info( - "[CONNECTOR] transient store staging caps: batch=%d pending=%d", - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - ) - - if self._layer_names and sample is not None: - num_blocks = sample.shape[1] if sample.dim() >= 2 else 1 - layer_size = sample.nbytes // num_blocks - local_slot_size = layer_size * len(self._layer_names) - tp_size = int(getattr(self, "_tp_size", 1)) - _validate_tp_layout( - local_slot_size, - self._slot_size, - tp_size, - int(getattr(self, "_server_tp_size", tp_size)), - int(getattr(self, "_tp_rank", 0)), - int(getattr(self, "_rank_stride_bytes", 0)), - ) - self._local_slot_size = local_slot_size - if self._slot_size == 0: - self._slot_size = local_slot_size * tp_size - logger.info( - "[CONNECTOR] registered local_slot_size=%d from %d layers", - self._local_slot_size, - len(self._layer_names), - ) - - if sample is not None: - self._store_staging_bytes = max( - self._store_staging_bytes or DEFAULT_STORE_STAGING_BYTES, - self._local_slot_size, - ) - self._store_staging_pool = StoreCudaStagingPool( - device=sample.device, - buffer_bytes=self._store_staging_bytes, - depth=_store_staging_pool_depth( - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - ), - ) - self._load_staging_pool = FixedCudaStagingPool( - device=sample.device, - buffer_bytes=self._store_staging_bytes, - depth=_load_staging_pool_depth( - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - sample.device, - ), - ) - self._load_staging_registered = False - logger.info( - "[CONNECTOR] preallocated staging buffer=%d cap=%d pending=%d " - "load_request_max_inflight=%d load_staging_depth=%d", - self._store_staging_bytes, - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - _LOAD_REQUEST_MAX_INFLIGHT, - self._load_staging_pool.depth, - ) - if sample.dim() >= 5: - _warm_rope_apply_backends( - device=sample.device, - dtype=sample.dtype, - block_tokens=int(sample.shape[-3]), - heads=int(sample.shape[-2]), - head_dim=int(sample.shape[-1]), - rotary_dim=int(getattr(self, "_rope_rotary_dim", 0)), - rope_base=float(getattr(self, "_rope_base", 10000.0)), - is_neox_style=bool(getattr(self, "_rope_is_neox_style", True)), - ) - - self._init_server_transfer() - - def register_cross_layers_kv_cache( - self, - kv_cache: torch.Tensor, - attn_backend: type[Any], - ) -> None: - """Register vLLM's cross-layer KV cache tensor. - - Args: - kv_cache: vLLM tensor whose logical layout starts with - ``[blocks, layers, 2, block_tokens, heads, head_dim]`` for the - NHD layout DaseR requests. - attn_backend: Attention backend that created ``kv_cache``. - - Async/thread-safety: - Called once during worker initialization before request traffic. - """ - kv_cache_config = getattr(self, "_kv_cache_config", None) - layer_names: list[str] = [] - if kv_cache_config is not None: - for group in getattr(kv_cache_config, "kv_cache_groups", []): - layer_names.extend(list(getattr(group, "layer_names", []))) - if not layer_names: - layer_count = int(kv_cache.shape[1]) if kv_cache.dim() >= 2 else 0 - layer_names = [f"layer.{idx}" for idx in range(layer_count)] - self._kv_caches = {CROSS_LAYER_KV_CACHE_KEY: kv_cache} - self._layer_names = layer_names - self._layer_idx_map = {name: idx for idx, name in enumerate(self._layer_names)} - if kv_cache.dim() < 6: - logger.warning( - "[CONNECTOR] cross-layer KV cache has unsupported shape=%s", - tuple(kv_cache.shape), - ) - return - ( - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - ) = _derive_store_staging_limits(kv_cache.device) - layer_size = kv_cache[0, 0].nbytes - local_slot_size = layer_size * len(self._layer_names) - tp_size = int(getattr(self, "_tp_size", 1)) - _validate_tp_layout( - local_slot_size, - self._slot_size, - tp_size, - int(getattr(self, "_server_tp_size", tp_size)), - int(getattr(self, "_tp_rank", 0)), - int(getattr(self, "_rank_stride_bytes", 0)), - ) - self._local_slot_size = local_slot_size - if self._slot_size == 0: - self._slot_size = local_slot_size * tp_size - logger.info( - "[CONNECTOR] registered cross-layer local_slot_size=%d from %d layers", - self._local_slot_size, - len(self._layer_names), - ) - self._store_staging_bytes = max( - self._store_staging_bytes or DEFAULT_STORE_STAGING_BYTES, - self._local_slot_size, - ) - self._store_staging_pool = StoreCudaStagingPool( - device=kv_cache.device, - buffer_bytes=self._store_staging_bytes, - depth=_store_staging_pool_depth( - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - ), - ) - self._load_staging_pool = FixedCudaStagingPool( - device=kv_cache.device, - buffer_bytes=self._store_staging_bytes, - depth=_load_staging_pool_depth( - self._store_staging_bytes, - self._pending_store_staging_limit_bytes, - kv_cache.device, - ), - ) - self._load_staging_registered = False - logger.info( - "[CONNECTOR] register_cross_layers_kv_cache: layers=%d shape=%s " - "dtype=%s load_request_max_inflight=%d load_staging_depth=%d", - len(self._layer_names), - tuple(kv_cache.shape), - kv_cache.dtype, - _LOAD_REQUEST_MAX_INFLIGHT, - self._load_staging_pool.depth, - ) - _warm_rope_apply_backends( - device=kv_cache.device, - dtype=kv_cache.dtype, - block_tokens=int(kv_cache.shape[-3]), - heads=int(kv_cache.shape[-2]), - head_dim=int(kv_cache.shape[-1]), - rotary_dim=int(getattr(self, "_rope_rotary_dim", 0)), - rope_base=float(getattr(self, "_rope_base", 10000.0)), - is_neox_style=bool(getattr(self, "_rope_is_neox_style", True)), - ) - _warm_cross_layer_restore_backends( - device=kv_cache.device, - dtype=kv_cache.dtype, - layers=int(kv_cache.shape[1]), - block_tokens=int(kv_cache.shape[-3]), - heads=int(kv_cache.shape[-2]), - head_dim=int(kv_cache.shape[-1]), - rotary_dim=int(getattr(self, "_rope_rotary_dim", 0)), - rope_base=float(getattr(self, "_rope_base", 10000.0)), - is_neox_style=bool(getattr(self, "_rope_is_neox_style", True)), - ) - self._init_server_transfer() - - def bind_connector_metadata(self, connector_metadata: DaserConnectorMeta) -> None: - """Receive scheduler metadata before each forward pass. - - Args: - connector_metadata: DaserConnectorMeta from build_connector_meta. - """ - super().bind_connector_metadata(connector_metadata) - self._meta = connector_metadata - self._reap_save_futures(block=False) - self._pending_commits = set() - for spec in connector_metadata.reqs_to_store.values(): - if spec.block_ids: - self._pending_commits.add(spec.chunk_key) - - def clear_connector_metadata(self) -> None: - """Clear metadata after forward pass completes.""" - super().clear_connector_metadata() - self._meta = None - - def start_load_kv(self, forward_context: "ForwardContext", **kwargs: Any) -> None: - """Submit async KV cache loads for cache-hit requests. - - Args: - forward_context: vLLM ForwardContext for this forward pass. - """ - del forward_context, kwargs - if self._meta is None or not self._meta.reqs_to_load: - return - logger.debug( - "[CONNECTOR] start_load_kv: %d reqs to load", - len(self._meta.reqs_to_load), - ) - reqs_to_load = dict(self._meta.reqs_to_load) - if not self._ensure_transfer_ready(): - self._mark_load_start_failed( - reqs_to_load, - "server transfer config is not ready", - ) - return - - num_layers = len(self._layer_names) - if num_layers == 0: - self._mark_load_start_failed(reqs_to_load, "no registered KV cache layers") - return - - sample_tensor = next(iter(self._kv_caches.values()), None) - if sample_tensor is None: - self._mark_load_start_failed(reqs_to_load, "no registered KV cache tensor") - return - - pending_loads = getattr(self, "_pending_loads", None) - if pending_loads is None: - pending_loads = {} - self._pending_loads = pending_loads - load_queue = self._ensure_load_request_queue(sample_tensor) - for spec_id, spec in reqs_to_load.items(): - req_id = base_req_id(spec_id) - request_future = _RequestLoadFuture() - pending_loads[req_id] = _PendingLoad( - future=request_future, - block_ids=list(spec.block_ids), - lease=None, - ) - self._enqueue_load_request( - load_queue, - _QueuedLoadRequest( - req_id=req_id, - spec_id=spec_id, - spec=spec, - future=request_future, - ), - ) - self._ensure_load_request_dispatcher(sample_tensor) - - def _mark_load_start_failed( - self, - reqs_to_load: dict[str, Any], - reason: str, - ) -> None: - """Record failed load submission so vLLM can release waiting requests. - - Args: - reqs_to_load: Load metadata that could not be submitted. - reason: Human-readable failure reason used in diagnostics. - - Async/thread-safety: - Called on the vLLM worker thread before any background load is - started. Completion is later reported through ``get_finished``. - """ - if not reqs_to_load: - return - block_ids = [ - block_id for spec in reqs_to_load.values() for block_id in spec.block_ids - ] - failed_future = _ImmediateLoadError(reason) - pending_loads = getattr(self, "_pending_loads", None) - if pending_loads is None: - pending_loads = {} - self._pending_loads = pending_loads - for req_id in {base_req_id(req_id) for req_id in reqs_to_load}: - pending_loads[req_id] = _PendingLoad( - future=failed_future, - block_ids=list(block_ids), - lease=None, - ) - - def _ensure_load_request_queue( - self, sample_tensor: torch.Tensor | None = None - ) -> Any: - """Return the worker-side request load queue, creating it if needed. - - Returns: - Queue used to group request loads across scheduler windows. - - Async/thread-safety: - Called on the vLLM worker thread. Attribute installation is guarded - by a small lock because multiple model-runner calls may race during - startup. - """ - load_queue = getattr(self, "_load_request_queue", None) - if load_queue is not None: - return load_queue - lock = getattr(self, "_load_request_queue_lock", None) - if lock is None: - lock = threading.Lock() - self._load_request_queue_lock = lock - with lock: - load_queue = getattr(self, "_load_request_queue", None) - if load_queue is None: - if sample_tensor is not None and hasattr(self, "_load_loop"): - future = asyncio.run_coroutine_threadsafe( - self._create_load_request_queue(), - self._load_loop, - ) - load_queue = future.result(timeout=10.0) - else: - load_queue = queue.Queue() - self._load_request_queue = load_queue - return load_queue - - async def _create_load_request_queue(self) -> asyncio.Queue[Any]: - """Create an asyncio request queue on the load event loop.""" - return asyncio.Queue() - - def _enqueue_load_request( - self, - load_queue: Any, - request: _QueuedLoadRequest, - ) -> None: - """Enqueue one request from a worker thread into the load queue. - - Args: - load_queue: Queue created by ``_ensure_load_request_queue``. - request: Request-level load work item. - - Async/thread-safety: - Called on vLLM worker threads. Production queues are asyncio queues - owned by ``_load_loop``; test queues may be synchronous ``queue.Queue`` - instances. - """ - if isinstance(load_queue, asyncio.Queue): - self._load_loop.call_soon_threadsafe(load_queue.put_nowait, request) - else: - load_queue.put(request) - - def _ensure_load_request_dispatcher(self, sample_tensor: torch.Tensor) -> None: - """Start the persistent request load dispatcher if it is not running. - - Args: - sample_tensor: Representative KV cache tensor used by load workers. - - Async/thread-safety: - Called on the vLLM worker thread. The dispatcher itself runs on the - connector load asyncio loop and keeps request-level loads in flight - up to the fixed max-inflight/staging-pool limit. - """ - dispatcher_future = getattr(self, "_load_request_dispatcher_future", None) - if dispatcher_future is not None and not dispatcher_future.done(): - return - lock = getattr(self, "_load_request_queue_lock", None) - if lock is None: - lock = threading.Lock() - self._load_request_queue_lock = lock - with lock: - dispatcher_future = getattr(self, "_load_request_dispatcher_future", None) - if dispatcher_future is not None and not dispatcher_future.done(): - return - self._load_request_dispatcher_future = asyncio.run_coroutine_threadsafe( - self._run_load_request_dispatcher(sample_tensor), - self._load_loop, - ) - - async def _run_load_request_dispatcher(self, sample_tensor: torch.Tensor) -> None: - """Continuously dispatch request loads as in-flight slots become free. - - Args: - sample_tensor: Representative KV cache tensor used for load staging. - - Async/thread-safety: - Runs on the connector load asyncio loop. The dispatcher never waits - to coalesce requests; it submits each queued request immediately when - both an in-flight slot and a fixed staging buffer are available. - """ - if sample_tensor.device.type == "cuda": - torch.cuda.set_device(sample_tensor.device) - load_queue = self._ensure_load_request_queue(sample_tensor) - load_staging_pool = self._ensure_load_staging_pool(sample_tensor) - dispatcher = LoadRequestDispatcher( - max_inflight=_LOAD_REQUEST_MAX_INFLIGHT, - staging_depth=load_staging_pool.depth, - ) - queued: list[_QueuedLoadRequest] = [] - active: list[_InflightRequestLoad] = [] - try: - while True: - if not queued: - if active: - await self._drain_load_queue(load_queue, queued) - else: - item = await load_queue.get() - if item is None: - return - queued.append(item) - active.extend(dispatcher.submit_ready(self, queued, sample_tensor)) - consumed = dispatcher.consume_ready(self, active, sample_tensor) - if consumed: - continue - if not active: - continue - await self._wait_for_dispatcher_completion(active) - except BaseException as exc: - for state in active: - if not state.item.future.done(): - state.item.future.set_exception(exc) - active_batch = state.active - if active_batch is not None: - active_batch.staging_lease.release() - for item in queued: - if not item.future.done(): - item.future.set_exception(exc) - raise - - async def _drain_load_queue( - self, - load_queue: asyncio.Queue[Any], - queued: list[_QueuedLoadRequest], - ) -> None: - """Move immediately available queue items into a local FIFO list.""" - while True: - try: - item = load_queue.get_nowait() - except asyncio.QueueEmpty: - return - if item is None: - await load_queue.put(None) - return - queued.append(item) - - async def _wait_for_dispatcher_completion( - self, - active: list[_InflightRequestLoad], - ) -> None: - """Wait until any request-level active read future completes.""" - wrapped: dict[asyncio.Future[Any], _InflightRequestLoad] = {} - for state in active: - active_batch = state.active - if active_batch is None: - continue - wrapped[asyncio.wrap_future(active_batch.future)] = state - if not wrapped: - await asyncio.sleep(0) - return - done, _pending = await asyncio.wait( - wrapped.keys(), - timeout=_LOAD_DISPATCH_WAIT_TIMEOUT_S, - return_when=asyncio.FIRST_COMPLETED, - ) - if not done: - await asyncio.sleep(0) - - def _ensure_load_staging_pool( - self, - sample_tensor: torch.Tensor, - ) -> FixedCudaStagingPool: - """Return the fixed load staging pool, creating it before traffic if needed. - - Args: - sample_tensor: Representative KV cache tensor used for device - placement when the pool was not initialized during registration. - - Returns: - Fixed load staging pool with two preallocated buffers. - - Async/thread-safety: - Called from the worker load executor. Normal production flow creates - the pool during KV-cache registration; this fallback keeps tests and - deferred initialization paths explicit while still allocating only - once before batch submission. - """ - pool = getattr(self, "_load_staging_pool", None) - if pool is not None: - return pool - buffer_bytes = max( - self._store_staging_bytes or DEFAULT_STORE_STAGING_BYTES, - _local_slot_bytes(self), - ) - pool = FixedCudaStagingPool( - device=sample_tensor.device, - buffer_bytes=buffer_bytes, - depth=_load_staging_pool_depth( - buffer_bytes, - getattr(self, "_pending_store_staging_limit_bytes", 0), - sample_tensor.device, - ), - ) - self._load_staging_pool = pool - return pool - - def _submit_load_batch( - self, - batch: _LoadBatch, - buffer_index: int, - sample_tensor: torch.Tensor, - ) -> _InflightLoadBatch: - """Submit one load batch into a fixed staging buffer. - - Args: - batch: Tuple from ``build_load_read_batches``. - buffer_index: Fixed load staging buffer to use. - sample_tensor: Representative KV cache tensor for device context. - - Returns: - In-flight batch state consumed by ``_consume_loaded_batch``. - - Async/thread-safety: - Runs on the connector load executor. The returned state owns the - fixed staging lease until consumption finishes. - """ - del sample_tensor - total_bytes, spans, per_req_ranges = batch - pool = self._ensure_load_staging_pool(next(iter(self._kv_caches.values()))) - staging_lease = pool.acquire_index(buffer_index, total_bytes) - staging = staging_lease.view - - submitted_at = time.perf_counter() - if getattr(self, "_load_staging_registered", False): - transfer_coro = self._transfer_load_registered_cuda( - buffer_index=buffer_index, - producer_pid=os.getpid(), - nbytes=total_bytes, - spans=spans, - ) - else: - cp_staging = cupy.asarray(staging) - cuda_handle = export_cuda_ipc_handle(cp_staging) - device_id = cuda_array_device_id(cp_staging) - device_ptr = cuda_array_pointer(cp_staging) - ipc_base_ptr, ipc_offset = _cuda_allocation_base_and_offset(device_ptr) - transfer_coro = self._transfer_load_cuda( - buffer_index=buffer_index, - cuda_ipc_handle=cuda_handle, - nbytes=total_bytes, - device_id=device_id, - device_ptr=device_ptr, - allocation_base_ptr=ipc_base_ptr, - allocation_offset=ipc_offset, - producer_pid=os.getpid(), - spans=spans, - ) - future = self._submit_load_coroutine(transfer_coro) - return _InflightLoadBatch( - buffer_index=buffer_index, - total_bytes=total_bytes, - per_req_ranges=per_req_ranges, - staging_lease=staging_lease, - future=future, - submitted_at=submitted_at, - ) - - def _submit_request_load_for_dispatcher( - self, - item: _QueuedLoadRequest, - buffer_index: int, - sample_tensor: torch.Tensor, - ) -> _InflightRequestLoad: - """Submit the first read batch for one queued request load. - - Args: - item: Request-level queued load. - buffer_index: Fixed staging buffer index reserved by the dispatcher. - sample_tensor: Representative KV cache tensor. - - Returns: - Request state tracked by ``LoadRequestDispatcher``. - - Async/thread-safety: - Runs on the connector load asyncio loop. Each request owns at most one - active staging buffer at a time; multi-segment requests submit the - next segment only after the previous segment has been restored. - """ - load_staging_pool = self._ensure_load_staging_pool(sample_tensor) - spec = replace( - item.spec, - file_offset=_rank_lane_offset( - item.spec.start_slot, - _local_slot_bytes(self), - int(getattr(self, "_rank_stride_bytes", 0)), - int(getattr(self, "_tp_rank", 0)), - ), - ) - load_batches = _build_load_read_batches( - {item.spec_id: spec}, - _local_slot_bytes(self), - max_batch_bytes=load_staging_pool.buffer_bytes, - include_req_ids=True, - ) - if not load_batches: - item.future.set_result() - return _InflightRequestLoad( - item=item, - buffer_index=buffer_index, - batches=[], - next_batch=0, - remaining_batches=0, - active=None, - completed=[], - ) - active_batch = self._submit_load_batch( - load_batches[0], - buffer_index, - sample_tensor, - ) - return _InflightRequestLoad( - item=item, - buffer_index=buffer_index, - batches=load_batches, - next_batch=1, - remaining_batches=len(load_batches), - active=active_batch, - completed=[], - ) - - def _consume_dispatcher_load( - self, - state: _InflightRequestLoad, - sample_tensor: torch.Tensor, - ) -> tuple[int, bool]: - """Consume one completed request read and finish or advance the request. - - Args: - state: Request-level load state with a completed active batch. - sample_tensor: Representative KV cache tensor. - - Returns: - Tuple of released staging buffer index and whether the request is - fully complete. - - Async/thread-safety: - Runs on the connector load asyncio loop after the associated transfer - future is complete. The staging buffer is not reused until restore - kernels have synchronized and the lease has been released. - """ - active_batch = state.active - if active_batch is None: - return state.buffer_index, True - consumed = self._consume_loaded_batch(active_batch, sample_tensor) - state.completed.append(consumed) - state.remaining_batches = max(0, state.remaining_batches - 1) - reusable_buffer = consumed.buffer_index - state.buffer_index = reusable_buffer - if state.remaining_batches == 0: - if not state.item.future.done(): - state.item.future.set_result() - return reusable_buffer, True - state.active = self._submit_load_batch( - state.batches[state.next_batch], - reusable_buffer, - sample_tensor, - ) - state.next_batch += 1 - return reusable_buffer, False - - def _consume_loaded_batch( - self, - state: _InflightLoadBatch, - sample_tensor: torch.Tensor, - ) -> _ConsumedLoadBatch: - """Wait for one submitted load batch and restore it into vLLM KV cache. - - Args: - state: In-flight load batch returned by ``_submit_load_batch``. - sample_tensor: Representative KV cache tensor for synchronization. - - Returns: - Timing and accounting data for the restored batch. - - Async/thread-safety: - Runs on the connector load executor. It releases the fixed staging - lease only after restore kernels that read the staging view have - been synchronized. - """ - try: - wait_start = time.perf_counter() - load_response = state.future.result(timeout=120.0) - wait_ms = (time.perf_counter() - wait_start) * 1000 - ipc_ms = (time.perf_counter() - state.submitted_at) * 1000 - transfer_open_ms = float(load_response.get("transfer_open_ms", 0.0)) - transfer_load_ms = float(load_response.get("transfer_load_ms", 0.0)) - transfer_sync_ms = float(load_response.get("transfer_sync_ms", 0.0)) - stats = load_response.get("transfer_stats_delta", {}) - l1_hits = 0 - l1_misses = 0 - l2_reads = 0 - if isinstance(stats, dict): - l1_hits = int(stats.get("l1_hits", 0)) - l1_misses = int(stats.get("l1_misses", 0)) - l2_reads = int(stats.get("l2_reads", 0)) - - copy_runs = _build_load_copy_runs(state.per_req_ranges) - copy_start = time.perf_counter() - copies = 0 - staging = state.staging_lease.view - for run in copy_runs: - copies += _copy_staging_to_kv_cache( - staging=staging[run.start : run.end], - kv_caches=self._kv_caches, - layer_names=self._layer_names, - block_ids=run.block_ids, - slot_size=_local_slot_bytes(self), - load_key_scale=self._load_key_scale, - load_value_scale=self._load_value_scale, - pos_offset=run.pos_offset, - rope_delta_scale=self._rope_delta_scale, - rope_base=self._rope_base, - rope_rotary_dim=self._rope_rotary_dim, - rope_is_neox_style=self._rope_is_neox_style, - ) - copy_ms = (time.perf_counter() - copy_start) * 1000 - _synchronize_cuda_tensor(sample_tensor) - return _ConsumedLoadBatch( - buffer_index=state.buffer_index, - bytes=state.total_bytes, - copies=copies, - copy_runs=len(copy_runs), - ipc_ms=ipc_ms, - wait_ms=wait_ms, - copy_ms=copy_ms, - transfer_open_ms=transfer_open_ms, - transfer_load_ms=transfer_load_ms, - transfer_sync_ms=transfer_sync_ms, - l1_hits=l1_hits, - l1_misses=l1_misses, - l2_reads=l2_reads, - ) - finally: - state.staging_lease.release() - - def wait_for_layer_load(self, layer_name: str) -> None: - """No-op because async loads complete before vLLM resumes requests. - - Args: - layer_name: ignored. - """ - return - - def save_kv_layer( - self, - layer_name: str, - kv_layer: torch.Tensor, - attn_metadata: "AttentionMetadata", - **kwargs: Any, - ) -> None: - """Submit this layer's KV blocks for server-owned transfer. - - Args: - layer_name: name of the current attention layer. - kv_layer: full KV cache tensor for this layer. - attn_metadata: attention metadata (not directly used). - """ - if self._meta is None or not self._meta.reqs_to_store: - return - if not self._ensure_transfer_ready(): - return - - if layer_name not in self._layer_idx_map: - logger.warning( - "[CONNECTOR] save_kv_layer: unknown layer %s, skipping", layer_name - ) - - def wait_for_save(self) -> None: - """Queue stores until vLLM reports request completion.""" - if self._meta is None: - return - - commit_keys = list(self._pending_commits) - reqs_to_store = dict(self._meta.reqs_to_store) - if commit_keys and reqs_to_store: - pending_finished = getattr(self, "_pending_finished_saves", None) - if pending_finished is None: - pending_finished = {} - self._pending_finished_saves = pending_finished - for req_id, spec in reqs_to_store.items(): - base_id = base_req_id(req_id) - save = pending_finished.get(base_id) - if save is None: - save = _DeferredFinishedSave(commit_keys=set(), reqs_to_store={}) - pending_finished[base_id] = save - save.reqs_to_store[req_id] = spec - if spec.chunk_key in commit_keys: - save.commit_keys.add(spec.chunk_key) - self._pending_commits.clear() - - def get_finished( - self, finished_req_ids: set[str] - ) -> tuple[set[str] | None, set[str] | None]: - """Collect completed background transfers after a worker step. - - Args: - finished_req_ids: Request IDs that vLLM finished in this step. - - Returns: - Finished-saving request IDs and finished async-loading request IDs. - """ - self._reap_save_futures(block=False) - finished_recving = self._collect_finished_loads() - pending_finished = getattr(self, "_pending_finished_saves", {}) - if not pending_finished: - return None, finished_recving or None - - finished_sending: set[str] = set() - candidates = set(finished_req_ids) - candidates.update( - req_id for req_id, save in pending_finished.items() if save.submitted - ) - for req_id in list(candidates): - save = pending_finished.get(req_id) - if save is None: - continue - if not save.submitted: - save.future = self._submit_finished_save(save) - save.submitted = True - future = save.future - if future is None: - finished_sending.add(req_id) - del pending_finished[req_id] - elif future.done(): - future.result(timeout=120.0) - finished_sending.add(req_id) - del pending_finished[req_id] - return finished_sending or None, finished_recving or None - - def get_block_ids_with_load_errors(self) -> set[int]: - """Return and clear block IDs whose async load failed. - - Returns: - vLLM block IDs that should be treated as invalid. - """ - invalid_blocks = getattr(self, "_invalid_load_block_ids", None) - if invalid_blocks is None: - self._invalid_load_block_ids = set() - return set() - invalid = set(invalid_blocks) - invalid_blocks.clear() - return invalid - - def _collect_finished_loads(self) -> set[str]: - """Poll async load futures without blocking. - - Returns: - Base request IDs whose async load future completed in this poll. - """ - pending_loads = getattr(self, "_pending_loads", {}) - if not pending_loads: - return set() - - finished_recving: set[str] = set() - collected_futures: set[int] = set() - for req_id, load in list(pending_loads.items()): - if not load.future.done(): - continue - future_id = id(load.future) - try: - if future_id not in collected_futures: - load.future.result() - collected_futures.add(future_id) - except Exception as exc: # noqa: BLE001 - logger.warning( - "[CONNECTOR] async load failed req=%s blocks=%s: %s", - req_id, - load.block_ids, - exc, - ) - self._invalid_load_block_ids.update(load.block_ids) - collected_futures.add(future_id) - finally: - load.release() - del pending_loads[req_id] - finished_recving.add(req_id) - return finished_recving - - def shutdown(self) -> None: - """Stop the background IO loop.""" - if self._role != KVConnectorRole.WORKER: - return - pending_loads = getattr(self, "_pending_loads", {}) - for load in {id(load.future): load for load in pending_loads.values()}.values(): - if not load.future.done(): - try: - load.future.result(timeout=120.0) - except Exception: # noqa: BLE001 - pass - self._collect_finished_loads() - for req_id in list(getattr(self, "_pending_finished_saves", {})): - self.get_finished({req_id}) - self._reap_save_futures(block=True) - load_queue = getattr(self, "_load_request_queue", None) - if load_queue is not None: - if isinstance(load_queue, asyncio.Queue): - self._load_loop.call_soon_threadsafe(load_queue.put_nowait, None) - else: - load_queue.put(None) - dispatcher_future = getattr(self, "_load_request_dispatcher_future", None) - if dispatcher_future is not None: - try: - dispatcher_future.result(timeout=120.0) - except Exception: # noqa: BLE001 - pass - load_clients = list( - dict.fromkeys( - getattr(self, "_ipc_load_async_pool", []) - or [getattr(self, "_ipc_load_async", None)] - ) - ) - store_client = getattr(self, "_ipc_store_async", None) - for load_client in load_clients: - if load_client is not None: - self._submit_load_coroutine(load_client.close()).result(timeout=10.0) - if store_client is not None: - self._submit_store_coroutine(store_client.close()).result(timeout=10.0) - self._load_loop.call_soon_threadsafe(self._load_loop.stop) - self._store_loop.call_soon_threadsafe(self._store_loop.stop) - self._load_thread.join(timeout=5) - self._store_thread.join(timeout=5) - - def _run_load_loop(self) -> None: - """Run the foreground load asyncio IO loop.""" - asyncio.set_event_loop(self._load_loop) - self._load_loop.run_forever() - - def _run_store_loop(self) -> None: - """Run the background store asyncio IO loop.""" - asyncio.set_event_loop(self._store_loop) - self._store_loop.run_forever() - - def _submit_load_coroutine(self, coro: Any) -> Any: - """Submit foreground load work to the dedicated load event loop. - - Args: - coro: Coroutine object to schedule. - - Returns: - Future returned by ``asyncio.run_coroutine_threadsafe``. - - Async/thread-safety: - Called from vLLM worker threads. A load-only loop prevents cache-hit - reads from queueing behind background store coroutines. - """ - loop = self._load_loop - return asyncio.run_coroutine_threadsafe(coro, loop) - - def _submit_store_coroutine(self, coro: Any) -> Any: - """Submit background store and commit work to the store event loop. - - Args: - coro: Coroutine object to schedule. - - Returns: - Future returned by ``asyncio.run_coroutine_threadsafe``. - - Async/thread-safety: - Called from vLLM worker threads. Store work is serialized on the - store loop and does not occupy the foreground load loop. - """ - loop = self._store_loop - return asyncio.run_coroutine_threadsafe(coro, loop) - - def _ensure_transfer_ready(self) -> bool: - """Refresh server transfer config and mark worker data plane ready.""" - if getattr(self, "_transfer_ready", False): - return True - - self._refresh_runtime_config() - if not self._slot_size or ( - not getattr(self, "_skip_l2", False) and not self._store_path - ): - logger.warning( - "[CONNECTOR] server transfer config is not ready; start DaseR server " - "before sending requests", - ) - return False - - _validate_tp_layout( - _local_slot_bytes(self), - self._slot_size, - int(getattr(self, "_tp_size", 1)), - int(getattr(self, "_server_tp_size", 1)), - int(getattr(self, "_tp_rank", 0)), - int(getattr(self, "_rank_stride_bytes", 0)), - ) - - self._transfer_ready = True - logger.info("[CONNECTOR] server transfer mode=%s", self._transfer_mode) - return True - - def _init_server_transfer(self) -> None: - """Initialize the server-owned transfer layer on both IO loops. - - Async/thread-safety: - Called from a vLLM worker thread after KV-cache registration. Does - nothing until the server transfer config and both async IO loops - are ready. - """ - if not ( - self._ensure_transfer_ready() - and getattr(self, "_ipc_load_async", None) is not None - and getattr(self, "_ipc_store_async", None) is not None - and getattr(self, "_load_loop", None) is not None - and getattr(self, "_store_loop", None) is not None - ): - return - for load_client in self._load_ipc_clients(): - self._submit_load_coroutine(load_client.init_transfer()).result( - timeout=120.0 - ) - self._submit_store_coroutine(self._ipc_store_async.init_transfer()).result( - timeout=120.0 - ) - self._register_load_staging_buffers() - - def _load_ipc_clients(self) -> list[Any]: - """Return fixed load IPC clients used for parallel load RPCs. - - Returns: - Load IPC clients. The list falls back to the legacy single client - when the connector was constructed by tests or older harnesses. - - Async/thread-safety: - Called on the vLLM worker thread during initialization and from the - load event loop when selecting the client for a submitted batch. - """ - clients = getattr(self, "_ipc_load_async_pool", None) - if clients: - return list(clients) - client = getattr(self, "_ipc_load_async", None) - return [client] if client is not None else [] - - def _load_ipc_client_for_buffer(self, buffer_index: int | None = None) -> Any: - """Return the load IPC client assigned to a staging buffer. - - Args: - buffer_index: Fixed staging buffer index for the transfer. - - Returns: - Async IPC client for the selected load lane. - - Async/thread-safety: - Pure selection helper. Each returned client owns its own IPC socket, - so separate fixed staging buffers can have concurrent server RPCs - instead of serializing on one client lock. - """ - clients = self._load_ipc_clients() - if not clients: - raise RuntimeError("load IPC client is not initialized") - if buffer_index is None: - return clients[0] - return clients[int(buffer_index) % len(clients)] - - def _register_load_staging_buffers(self) -> None: - """Register fixed load staging buffers with the server. - - Async/thread-safety: - Called after server transfer initialization and before request - traffic. Registration failures are logged and leave the worker on - the compatible per-load CUDA IPC payload path. - """ - if getattr(self, "_load_staging_registered", False): - return - pool = getattr(self, "_load_staging_pool", None) - if pool is None or not self._load_ipc_clients(): - return - try: - for buffer_index in range(pool.depth): - tensor = pool.buffer(buffer_index) - cp_tensor = cupy.asarray(tensor) - device_ptr = cuda_array_pointer(cp_tensor) - ipc_base_ptr, ipc_offset = _cuda_allocation_base_and_offset(device_ptr) - load_client = self._load_ipc_client_for_buffer(buffer_index) - self._submit_load_coroutine( - load_client.register_load_staging_cuda( - buffer_index=buffer_index, - cuda_ipc_handle=export_cuda_ipc_handle(cp_tensor), - allocation_bytes=int(tensor.numel()), - device_id=cuda_array_device_id(cp_tensor), - device_ptr=device_ptr, - allocation_base_ptr=ipc_base_ptr, - allocation_offset=ipc_offset, - producer_pid=os.getpid(), - ) - ).result(timeout=120.0) - except Exception as exc: # noqa: BLE001 - logger.warning( - "[CONNECTOR] registered load staging unavailable; falling back " - "to per-load CUDA IPC payloads: %s", - exc, - ) - self._load_staging_registered = False - return - self._load_staging_registered = True - logger.info( - "[CONNECTOR] registered %d fixed load staging buffers", - pool.depth, - ) - - def _reap_save_futures(self, block: bool) -> None: - """Collect completed background save tasks. - - Args: - block: If True, wait for every pending save. If False, collect only - tasks that are already complete. - """ - remaining: list[_SaveFuture] = [] - pending_bytes = self._pending_save_staging_bytes - for record in self._save_futures: - if block or record.future.done(): - try: - record.future.result(timeout=120.0) - finally: - pending_bytes = max(0, pending_bytes - record.staging_bytes) - record.release() - else: - remaining.append(record) - self._save_futures = remaining - self._pending_save_staging_bytes = pending_bytes - - def _track_save_future( - self, - future: Any, - staging_bytes: int, - staging_lease: CudaStagingLease | None, - ) -> None: - """Track one background save future and its live staging bytes. - - Args: - future: Future returned by ``asyncio.run_coroutine_threadsafe``. - staging_bytes: GPU staging bytes kept alive by the future. - staging_lease: Optional reusable staging lease to release after - ``future`` completes. - - Async/thread-safety: - Called on the worker thread. Completion is collected by - ``_reap_save_futures``. - """ - self._pending_save_staging_bytes += staging_bytes - self._save_futures.append( - _SaveFuture(future=future, staging_bytes=staging_bytes, lease=staging_lease) - ) - - def _wait_for_save_staging_capacity(self, nbytes: int) -> None: - """Apply backpressure before allocating another store staging buffer. - - Args: - nbytes: Size of the next staging tensor. - - Async/thread-safety: - Called by vLLM's worker thread. It may wait for already-submitted - background stores when live staging would exceed the configured - cap. - """ - limit = max( - self._pending_store_staging_limit_bytes - or DEFAULT_PENDING_STORE_STAGING_BYTES, - nbytes, - ) - while self._pending_save_staging_bytes + nbytes > limit and self._save_futures: - record = self._save_futures.pop(0) - try: - record.future.result(timeout=120.0) - finally: - self._pending_save_staging_bytes = max( - 0, - self._pending_save_staging_bytes - record.staging_bytes, - ) - record.release() - self._reap_save_futures(block=False) - - def _wait_for_store_staging_release(self, nbytes: int) -> None: - """Wait until one store staging lease can return to the fixed pool. - - Args: - nbytes: Size of the staging lease requested by the caller. - - Async/thread-safety: - Called from the worker thread when the fixed store staging pool is - exhausted. It first applies byte-budget backpressure, then waits - for oldest store futures until a lease is released. - """ - pool = getattr(self, "_store_staging_pool", None) - self._wait_for_save_staging_capacity(nbytes) - while pool is not None and pool.available == 0 and self._save_futures: - record = self._save_futures.pop(0) - try: - record.future.result(timeout=120.0) - finally: - self._pending_save_staging_bytes = max( - 0, - self._pending_save_staging_bytes - record.staging_bytes, - ) - record.release() - self._reap_save_futures(block=False) - - def _acquire_staging( - self, - nbytes: int, - device: torch.device, - ) -> CudaStagingLease: - """Acquire a reusable staging buffer for a CUDA IPC transfer. - - Args: - nbytes: Logical byte count needed for the transfer. - device: Device used when the pool has not been initialized yet. - - Returns: - A staging lease whose ``view`` is safe to export through CUDA IPC. - - Async/thread-safety: - Called from the worker thread. Store-path callers must retain the - lease until the background server transfer completes. - """ - pool = getattr(self, "_store_staging_pool", None) - if pool is None: - max_bytes = max( - nbytes, - self._store_staging_bytes or DEFAULT_STORE_STAGING_BYTES, - ) - pool = StoreCudaStagingPool( - device=device, - buffer_bytes=max_bytes, - depth=_store_staging_pool_depth( - max_bytes, - self._pending_store_staging_limit_bytes - or DEFAULT_PENDING_STORE_STAGING_BYTES, - ), - ) - self._store_staging_pool = pool - return pool.acquire( - nbytes, - wait_for_release=lambda: self._wait_for_store_staging_release(nbytes), - ) - - def _stage_store_batch( - self, - block_ids: list[int], - spans: list[StoreWriteSpan], - ) -> StagedStoreBatch | None: - """Snapshot one bounded batch of KV blocks into CUDA staging. - - Args: - block_ids: vLLM KV block IDs to snapshot. - spans: Server store spans targeting this staging batch. - - Returns: - A staged batch ready for CUDA IPC transfer, or ``None`` when the - connector has no layer state. - - Async/thread-safety: - Runs on the vLLM worker thread so KV cache reads are launched before - vLLM can recycle the source blocks. The returned tensor is kept - alive by the background transfer future. - """ - num_layers = len(self._layer_names) - if num_layers == 0: - return None - sample_tensor = next(iter(self._kv_caches.values()), None) - if sample_tensor is None: - return None - if not block_ids or not spans: - return None - local_slot_size = _local_slot_bytes(self) - nbytes = len(block_ids) * local_slot_size - self._wait_for_save_staging_capacity(nbytes) - staging_lease = self._acquire_staging(nbytes, sample_tensor.device) - staging = staging_lease.view - block_index = torch.tensor( - block_ids, - dtype=torch.long, - device=sample_tensor.device, - ) - cross_layer_kv_cache = self._kv_caches.get(CROSS_LAYER_KV_CACHE_KEY) - if cross_layer_kv_cache is not None: - _copy_cross_layer_kv_cache_to_staging( - staging=staging, - kv_cache=cross_layer_kv_cache, - block_ids=block_ids, - num_layers=num_layers, - slot_size=local_slot_size, - block_index=block_index, - ) - else: - for layer_name in self._layer_names: - _copy_kv_cache_to_staging( - staging=staging, - kv_layer=self._kv_caches[layer_name], - layer_idx=self._layer_idx_map[layer_name], - block_ids=block_ids, - num_layers=num_layers, - slot_size=local_slot_size, - block_index=block_index, - ) - return StagedStoreBatch( - buffer=staging, - ready_event=_record_cuda_event(staging), - spans=spans, - lease=staging_lease, - ) - - def _submit_finished_save(self, save: _DeferredFinishedSave) -> Any | None: - """Submit one request's deferred KV store after request completion. - - Args: - save: Deferred store plan built during ``wait_for_save``. - - Returns: - Future that completes after all store batches are committed, or - ``None`` when no store batch could be staged. - - Async/thread-safety: - Called by vLLM's worker thread from ``get_finished`` while vLLM is - still holding the finished request's KV blocks. - """ - batch_futures = [] - reqs_to_store = { - req_id: replace( - spec, - file_offset=_rank_lane_offset( - spec.start_slot, - _local_slot_bytes(self), - int(getattr(self, "_rank_stride_bytes", 0)), - int(getattr(self, "_tp_rank", 0)), - ), - ) - for req_id, spec in save.reqs_to_store.items() - } - batches = _build_staging_store_batches( - reqs_to_store, - _local_slot_bytes(self), - max_batch_bytes=(self._store_staging_bytes or DEFAULT_STORE_STAGING_BYTES), - ) - for block_ids, spans in batches: - staged = self._stage_store_batch(block_ids, spans) - if staged is None: - continue - future = self._submit_store_coroutine( - self._write_cuda_buffer( - buffer=staged.buffer, - ready_event=staged.ready_event, - spans=staged.spans, - ) - ) - self._track_save_future(future, staged.buffer.nbytes, staged.lease) - batch_futures.append(future) - if not batch_futures: - return None - commit_future = self._submit_store_coroutine( - self._commit_after_store_futures(batch_futures, sorted(save.commit_keys)), - ) - self._track_save_future(commit_future, 0, None) - return commit_future - - async def _commit_after_store_futures( - self, - batch_futures: list[Any], - commit_keys: list[str], - ) -> None: - """Commit chunks after all staged transfer batches finish. - - Args: - batch_futures: Futures for each staged store batch. - commit_keys: Chunk keys to publish after stores complete. - - Async/thread-safety: - Runs on the connector background event loop and does not read vLLM - KV cache tensors. - """ - stored_keys: list[str] = [] - for future in batch_futures: - stored_keys.extend(await asyncio.wrap_future(future)) - await self._commit_stored_keys(stored_keys, commit_keys) - - async def _write_cuda_buffer( - self, - buffer: torch.Tensor, - ready_event: torch.cuda.Event | None, - spans: list[StoreWriteSpan], - ) -> list[str]: - """Write selected spans from one contiguous CUDA buffer. - - Args: - buffer: CUDA tensor exported over CUDA IPC. - ready_event: Producer-stream event for ``buffer``. - spans: Source/destination write spans. - - Returns: - Chunk keys accepted by the server for this buffer. - """ - if buffer.device.type == "cuda": - torch.cuda.set_device(buffer.device) - if ready_event is not None: - ready_event.synchronize() - else: - _synchronize_cuda_tensor(buffer) - cp_buffer = cupy.asarray(buffer) - cuda_ipc_handle = export_cuda_ipc_handle(cp_buffer) - device_id = cuda_array_device_id(cp_buffer) - device_ptr = cuda_array_pointer(cp_buffer) - ipc_base_ptr, ipc_offset = _cuda_allocation_base_and_offset(device_ptr) - stored_keys = await self._transfer_store_cuda( - cuda_ipc_handle=cuda_ipc_handle, - nbytes=buffer.nbytes, - device_id=device_id, - device_ptr=device_ptr, - allocation_base_ptr=ipc_base_ptr, - allocation_offset=ipc_offset, - producer_pid=os.getpid(), - spans=[ - { - "source_offset": span.source_offset, - "nbytes": span.nbytes, - "file_offset": span.file_offset, - "chunk_key": span.chunk_key, - "start_slot": span.start_slot, - "num_slots": span.num_slots, - } - for span in spans - ], - ) - return stored_keys - - async def _commit_stored_keys( - self, - stored_keys: list[str], - commit_keys: list[str], - ) -> None: - """Commit requested chunks whose store spans were accepted.""" - requested = set(commit_keys) - candidate_keys = [key for key in stored_keys if key in requested] - keys_to_commit = list(dict.fromkeys(candidate_keys)) - await self._ipc_store_async.commit_chunks( - keys_to_commit, - tp_rank=int(getattr(self, "_tp_rank", 0)), - tp_size=int(getattr(self, "_tp_size", 1)), - ) - - async def _transfer_load_cuda(self, **kwargs: Any) -> dict[str, Any]: - """Load through the dedicated worker load IPC client. - - Args: - **kwargs: forwarded CUDA transfer payload fields. - - Returns: - Server load response with timing counters. - - Async/thread-safety: - Runs on the worker load event loop. A dedicated client keeps - cache-hit loads from queueing behind store RPCs. - """ - buffer_index = kwargs.pop("buffer_index", None) - client = self._load_ipc_client_for_buffer(buffer_index) - return await client.transfer_load_cuda(**kwargs) - - async def _transfer_load_registered_cuda(self, **kwargs: Any) -> dict[str, Any]: - """Load through a pre-registered fixed CUDA staging buffer. - - Args: - **kwargs: forwarded registered-buffer transfer fields. - - Returns: - Server load response with timing counters. - - Async/thread-safety: - Runs on the worker load event loop. The server has already opened - the CUDA IPC mapping during initialization, so this hot-path call - only identifies the staging buffer index and logical byte range. - """ - buffer_index = int(kwargs.get("buffer_index", 0)) - client = self._load_ipc_client_for_buffer(buffer_index) - return await client.transfer_load_registered_cuda(**kwargs) - - async def _transfer_store_cuda(self, **kwargs: Any) -> list[str]: - """Store through the dedicated worker store IPC client. - - Args: - **kwargs: forwarded CUDA transfer payload fields. - - Returns: - Chunk keys accepted by the server. - - Async/thread-safety: - Runs on the worker store event loop and serializes only with other - store/commit traffic. - """ - return await self._ipc_store_async.transfer_store_cuda(**kwargs) diff --git a/daser/connector/worker/__init__.py b/daser/connector/worker/__init__.py new file mode 100644 index 0000000..c81b912 --- /dev/null +++ b/daser/connector/worker/__init__.py @@ -0,0 +1,13 @@ +# SPDX-License-Identifier: Apache-2.0 + +from daser.connector.worker.adapter import WorkerConnectorMixin +from daser.connector.worker.load import LoadPipeline +from daser.connector.worker.runtime import WorkerRuntime +from daser.connector.worker.store import StorePipeline + +__all__ = [ + "LoadPipeline", + "StorePipeline", + "WorkerConnectorMixin", + "WorkerRuntime", +] diff --git a/daser/connector/worker/adapter.py b/daser/connector/worker/adapter.py new file mode 100644 index 0000000..5d9d046 --- /dev/null +++ b/daser/connector/worker/adapter.py @@ -0,0 +1,83 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +import torch + +if TYPE_CHECKING: + from vllm.attention import AttentionMetadata + from vllm.forward_context import ForwardContext + +from daser.connector.metadata import DaserConnectorMeta + + +class WorkerConnectorMixin: + """Adapt vLLM worker hooks to the worker runtime interface.""" + + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]) -> None: + """Register per-layer KV tensors with the worker runtime.""" + self._worker_runtime.register_kv_caches(kv_caches) + + def register_cross_layers_kv_cache( + self, + kv_cache: torch.Tensor, + attn_backend: type[Any], + ) -> None: + """Register vLLM's cross-layer KV tensor with the worker runtime.""" + self._worker_runtime.register_cross_layers_kv_cache(kv_cache, attn_backend) + + def bind_connector_metadata(self, connector_metadata: DaserConnectorMeta) -> None: + """Bind one scheduler metadata step to the worker runtime.""" + super().bind_connector_metadata(connector_metadata) + self._worker_runtime.bind_connector_metadata(connector_metadata) + + def clear_connector_metadata(self) -> None: + """Clear the current metadata step from the worker runtime.""" + super().clear_connector_metadata() + self._worker_runtime.clear_connector_metadata() + + def start_load_kv(self, forward_context: "ForwardContext", **kwargs: Any) -> None: + """Submit scheduler-selected cache loads through the load pipeline.""" + self._worker_runtime.start_load_kv(forward_context, **kwargs) + + def wait_for_layer_load(self, layer_name: str) -> None: + """Observe the runtime's request-level load completion contract.""" + self._worker_runtime.wait_for_layer_load(layer_name) + + def save_kv_layer( + self, + layer_name: str, + kv_layer: torch.Tensor, + attn_metadata: "AttentionMetadata", + **kwargs: Any, + ) -> None: + """Forward one vLLM save hook to the worker runtime.""" + self._worker_runtime.save_kv_layer( + layer_name, + kv_layer, + attn_metadata, + **kwargs, + ) + + def wait_for_save(self) -> None: + """Defer step stores through the worker runtime.""" + self._worker_runtime.wait_for_save() + + def get_finished( + self, + finished_req_ids: set[str], + ) -> tuple[set[str] | None, set[str] | None]: + """Return completed store and load request IDs from the runtime.""" + return self._worker_runtime.get_finished(finished_req_ids) + + def get_block_ids_with_load_errors(self) -> set[int]: + """Return and clear runtime load-error block IDs.""" + return self._worker_runtime.get_block_ids_with_load_errors() + + def shutdown(self) -> None: + """Drain and stop the worker runtime when initialized.""" + runtime = getattr(self, "_worker_runtime", None) + if runtime is not None: + runtime.shutdown() diff --git a/daser/connector/worker/load.py b/daser/connector/worker/load.py new file mode 100644 index 0000000..2820286 --- /dev/null +++ b/daser/connector/worker/load.py @@ -0,0 +1,880 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +from collections import deque +from concurrent.futures import Future +from contextlib import nullcontext +from dataclasses import dataclass, replace +import os +import threading +import time +from typing import Any + +# Third Party +import cupy +import torch + +from daser.connector.helpers import base_req_id +from daser.connector.ipc_client import IPCClientAsync +from daser.connector.metadata import ReqLoadSpec +from daser.connector.worker.memory import ( + CudaStagingLease, + FixedCudaStagingPool, +) +from daser.connector.worker.staging import copy_staging_to_kv_cache +from daser.logging import init_logger +from daser.transfer.cuda_ipc import ( + cuda_allocation_base_and_offset, + cuda_array_device_id, + cuda_array_pointer, + export_cuda_ipc_handle, +) + +logger = init_logger(__name__) + +_LOAD_DISPATCH_WAIT_TIMEOUT_S = 0.001 +_LoadBatch = tuple[int, list[dict[str, int]], list[Any]] + + +@dataclass +class _LoadRequest: + """Own one base request's specs and completion future.""" + + req_id: str + specs: dict[str, ReqLoadSpec] + future: Future[None] + + @property + def block_ids(self) -> list[int]: + """Return all vLLM blocks affected by this request.""" + return [block_id for spec in self.specs.values() for block_id in spec.block_ids] + + +@dataclass +class _InflightLoadBatch: + """Hold one submitted load batch and its fixed staging lease.""" + + total_bytes: int + per_req_ranges: list[Any] + staging_lease: CudaStagingLease + future: Any + submitted_at: float + + +@dataclass(frozen=True) +class _LoadBatchTiming: + """Record transfer and restore accounting for one load batch.""" + + bytes: int + copies: int + copy_runs: int + ipc_ms: float + wait_ms: float + copy_ms: float + worker_sync_ms: float + transfer_open_ms: float + transfer_load_ms: float + transfer_sync_ms: float + l1_hits: int + l1_misses: int + l2_reads: int + + +@dataclass +class _InflightRequestLoad: + """Track active and remaining load batches for one request.""" + + request: _LoadRequest + buffer_index: int + batches: deque[_LoadBatch] + active: _InflightLoadBatch | None + completed: list[_LoadBatchTiming] + + +class LoadPipeline: + """Own the complete worker load state machine. + + Args: + socket_path: DaseR server Unix socket path. + client_count: Independent load IPC lanes and maximum inflight requests. + + Async/thread-safety: + Public methods are called on the vLLM worker thread. Queue dispatch, + IPC, and CUDA restore execute on the private load thread. + """ + + def __init__(self, socket_path: str, client_count: int) -> None: + self._clients = [IPCClientAsync(socket_path) for _ in range(client_count)] + self._loop = asyncio.new_event_loop() + self._thread = threading.Thread( + target=self._run_loop, + daemon=True, + name="daser-load-io", + ) + self._queue: asyncio.Queue[Any] | None = None + self._queue_lock = threading.Lock() + self._dispatcher_future: Any | None = None + self._pending: dict[str, _LoadRequest] = {} + self._invalid_block_ids: set[int] = set() + self._staging_pool: FixedCudaStagingPool | None = None + self._staging_registered = False + self._kv_caches: dict[str, torch.Tensor] = {} + self._layer_names: list[str] = [] + self._local_slot_size = 0 + self._rank_stride_bytes = 0 + self._tp_rank = 0 + self._load_key_scale = 1.0 + self._load_value_scale = 1.0 + self._rope_delta_scale = 1.0 + self._rope_base = 10000.0 + self._rope_rotary_dim = 0 + self._rope_is_neox_style = True + self._cuda_stream: torch.cuda.Stream | None = None + self._thread.start() + + def configure( + self, + *, + kv_caches: dict[str, torch.Tensor], + layer_names: list[str], + local_slot_size: int, + rank_stride_bytes: int, + tp_rank: int, + staging_pool: FixedCudaStagingPool, + load_key_scale: float, + load_value_scale: float, + rope_delta_scale: float, + rope_base: float, + rope_rotary_dim: int, + rope_is_neox_style: bool, + ) -> None: + """Configure immutable KV layout and transform state. + + Args: + kv_caches: Registered vLLM KV tensors. + layer_names: Stable storage layer order. + local_slot_size: Bytes stored per slot by this TP rank. + rank_stride_bytes: Byte distance between rank lanes. + tp_rank: Current tensor-parallel rank. + staging_pool: Fixed load staging buffers. + load_key_scale: Load-time key scaling factor. + load_value_scale: Load-time value scaling factor. + rope_delta_scale: Position-offset scaling factor. + rope_base: RoPE theta/base. + rope_rotary_dim: Number of dimensions covered by RoPE. + rope_is_neox_style: Whether RoPE uses split-half rotation. + + Async/thread-safety: + Called once on the worker thread before request traffic. + """ + self._kv_caches = kv_caches + self._layer_names = list(layer_names) + self._local_slot_size = local_slot_size + self._rank_stride_bytes = rank_stride_bytes + self._tp_rank = tp_rank + self._staging_pool = staging_pool + self._staging_registered = False + self._load_key_scale = load_key_scale + self._load_value_scale = load_value_scale + self._rope_delta_scale = rope_delta_scale + self._rope_base = rope_base + self._rope_rotary_dim = rope_rotary_dim + self._rope_is_neox_style = rope_is_neox_style + + def initialize_transfer(self) -> None: + """Initialize load IPC lanes and register staging buffers. + + Async/thread-safety: + Called on the worker thread during startup. IPC runs on the load + loop and is joined before this method returns. + """ + for client in self._clients: + self._submit(client.init_transfer()).result(timeout=120.0) + self._register_staging_buffers() + + def configure_rank_geometry(self, rank_stride_bytes: int, tp_rank: int) -> None: + """Apply server-finalized tensor-parallel lane geometry. + + Args: + rank_stride_bytes: Byte distance between server-owned rank lanes. + tp_rank: Current tensor-parallel rank. + + Async/thread-safety: + Called on the worker thread after runtime-config refresh and before + any load is submitted. + """ + self._rank_stride_bytes = rank_stride_bytes + self._tp_rank = tp_rank + + def start(self, reqs_to_load: dict[str, ReqLoadSpec]) -> None: + """Queue request loads for background transfer and restore. + + Args: + reqs_to_load: Scheduler load metadata keyed by request/spec ID. + + Async/thread-safety: + Called from the worker thread. Queue dispatch, IPC, and restore run + on the private load thread. + """ + if not reqs_to_load: + return + if not self._layer_names or not self._kv_caches: + self.mark_failed(reqs_to_load, "no registered KV cache layout") + return + queue = self._ensure_queue() + grouped: dict[str, dict[str, ReqLoadSpec]] = {} + for spec_id, spec in reqs_to_load.items(): + grouped.setdefault(base_req_id(spec_id), {})[spec_id] = spec + for req_id, specs in grouped.items(): + request = _LoadRequest(req_id, specs, Future()) + self._pending[req_id] = request + self._loop.call_soon_threadsafe( + queue.put_nowait, + request, + ) + self._ensure_dispatcher() + + def mark_failed( + self, + reqs_to_load: dict[str, ReqLoadSpec], + reason: str, + ) -> None: + """Record submission failures for completion polling. + + Args: + reqs_to_load: Load specs that could not be submitted. + reason: Diagnostic failure reason. + + Async/thread-safety: + Called on the worker thread before background submission. + """ + grouped: dict[str, dict[str, ReqLoadSpec]] = {} + for spec_id, spec in reqs_to_load.items(): + grouped.setdefault(base_req_id(spec_id), {})[spec_id] = spec + for req_id, specs in grouped.items(): + future: Future[None] = Future() + future.set_exception(RuntimeError(reason)) + self._pending[req_id] = _LoadRequest(req_id, specs, future) + + def collect_finished(self) -> set[str]: + """Collect completed loads without blocking the worker thread. + + Returns: + Base request IDs whose load lifecycle completed in this poll. + + Async/thread-safety: + Called on the worker thread; request futures provide cross-thread + visibility from the load thread. + """ + finished: set[str] = set() + collected_futures: set[int] = set() + for req_id, load in list(self._pending.items()): + if not load.future.done(): + continue + future_id = id(load.future) + try: + if future_id not in collected_futures: + load.future.result() + collected_futures.add(future_id) + except Exception as exc: # noqa: BLE001 + logger.warning( + "[CONNECTOR] async load failed req=%s blocks=%s: %s", + req_id, + load.block_ids, + exc, + ) + self._invalid_block_ids.update(load.block_ids) + collected_futures.add(future_id) + finally: + del self._pending[req_id] + finished.add(req_id) + return finished + + def take_invalid_block_ids(self) -> set[int]: + """Return and clear block IDs targeted by failed loads. + + Returns: + vLLM block IDs that must be invalidated. + + Async/thread-safety: + Called on the worker thread after ``collect_finished``. + """ + invalid = set(self._invalid_block_ids) + self._invalid_block_ids.clear() + return invalid + + def shutdown(self) -> None: + """Drain queue ownership, close IPC clients, and stop the load loop.""" + for load in {id(item.future): item for item in self._pending.values()}.values(): + if not load.future.done(): + try: + load.future.result(timeout=120.0) + except Exception: # noqa: BLE001 + pass + self.collect_finished() + if self._queue is not None: + self._loop.call_soon_threadsafe(self._queue.put_nowait, None) + if self._dispatcher_future is not None: + self._dispatcher_future.result(timeout=120.0) + for client in self._clients: + self._submit(client.close()).result(timeout=5.0) + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(timeout=5.0) + + def _submit(self, coro: Any) -> Any: + return asyncio.run_coroutine_threadsafe(coro, self._loop) + + def _client(self, buffer_index: int | None = None) -> IPCClientAsync: + index = 0 if buffer_index is None else int(buffer_index) + return self._clients[index % len(self._clients)] + + def _ensure_queue(self) -> asyncio.Queue[Any]: + with self._queue_lock: + if self._queue is None: + self._queue = self._submit(self._create_queue()).result(timeout=5.0) + return self._queue + + def _ensure_dispatcher(self) -> None: + with self._queue_lock: + if ( + self._dispatcher_future is not None + and not self._dispatcher_future.done() + ): + return + self._dispatcher_future = self._submit(self._run_dispatcher()) + + async def _create_queue(self) -> asyncio.Queue[Any]: + return asyncio.Queue() + + def _run_loop(self) -> None: + asyncio.set_event_loop(self._loop) + self._loop.run_forever() + + async def _run_dispatcher(self) -> None: + sample_tensor = next(iter(self._kv_caches.values())) + if sample_tensor.device.type == "cuda": + torch.cuda.set_device(sample_tensor.device) + if self._cuda_stream is None: + self._cuda_stream = torch.cuda.Stream(device=sample_tensor.device) + queue = self._ensure_queue() + if self._staging_pool is None: + raise RuntimeError("load staging pool is not configured") + free_buffers = deque( + range(max(1, min(len(self._clients), self._staging_pool.depth))) + ) + queued: deque[_LoadRequest] = deque() + active: list[_InflightRequestLoad] = [] + try: + while True: + if not queued: + if active: + await self._drain_queue(queue, queued) + else: + item = await queue.get() + if item is None: + return + queued.append(item) + while queued and free_buffers: + request = queued.popleft() + buffer_index = free_buffers.popleft() + try: + state = self._submit_request(request, buffer_index) + except BaseException as exc: + if not request.future.done(): + request.future.set_exception(exc) + free_buffers.append(buffer_index) + continue + if state.active is None: + free_buffers.append(state.buffer_index) + else: + active.append(state) + consumed = False + for state in list(active): + active_batch = state.active + if active_batch is not None and not active_batch.future.done(): + continue + try: + reusable_buffer, request_done = self._consume_request(state) + except BaseException as exc: + if not state.request.future.done(): + state.request.future.set_exception(exc) + active.remove(state) + free_buffers.append(state.buffer_index) + consumed = True + continue + if request_done: + active.remove(state) + free_buffers.append(reusable_buffer) + consumed = True + if consumed: + continue + if active: + await self._wait_for_completion(active) + except BaseException as exc: + for state in active: + if not state.request.future.done(): + state.request.future.set_exception(exc) + if state.active is not None: + state.active.staging_lease.release() + for item in queued: + if not item.future.done(): + item.future.set_exception(exc) + raise + + async def _drain_queue( + self, + queue: asyncio.Queue[Any], + queued: deque[_LoadRequest], + ) -> None: + while True: + try: + item = queue.get_nowait() + except asyncio.QueueEmpty: + return + if item is None: + await queue.put(None) + return + queued.append(item) + + async def _wait_for_completion( + self, + active: list[_InflightRequestLoad], + ) -> None: + wrapped = { + asyncio.wrap_future(state.active.future) + for state in active + if state.active is not None + } + if not wrapped: + await asyncio.sleep(0) + return + done, _pending = await asyncio.wait( + wrapped, + timeout=_LOAD_DISPATCH_WAIT_TIMEOUT_S, + return_when=asyncio.FIRST_COMPLETED, + ) + for future in done: + future.exception() + if not done: + await asyncio.sleep(0) + + def _submit_request( + self, + request: _LoadRequest, + buffer_index: int, + ) -> _InflightRequestLoad: + if self._staging_pool is None: + raise RuntimeError("load staging pool is not configured") + specs = { + spec_id: replace( + spec, + file_offset=( + self._tp_rank * self._rank_stride_bytes + + spec.start_slot * self._local_slot_size + ), + ) + for spec_id, spec in request.specs.items() + } + batches = deque( + build_load_read_batches( + specs, + self._local_slot_size, + max_batch_bytes=self._staging_pool.buffer_bytes, + include_req_ids=True, + ) + ) + if not batches: + request.future.set_result(None) + return _InflightRequestLoad(request, buffer_index, batches, None, []) + return _InflightRequestLoad( + request=request, + buffer_index=buffer_index, + batches=batches, + active=self._submit_batch(batches.popleft(), buffer_index), + completed=[], + ) + + def _submit_batch( + self, + batch: _LoadBatch, + buffer_index: int, + ) -> _InflightLoadBatch: + if self._staging_pool is None: + raise RuntimeError("load staging pool is not configured") + total_bytes, spans, per_req_ranges = batch + lease = self._staging_pool.acquire_index(buffer_index, total_bytes) + staging = lease.view + if self._staging_registered: + transfer = self._client(buffer_index).transfer_load_registered_cuda( + buffer_index=buffer_index, + producer_pid=os.getpid(), + nbytes=total_bytes, + spans=spans, + ) + else: + cp_staging = cupy.asarray(staging) + device_ptr = cuda_array_pointer(cp_staging) + allocation_base, allocation_offset = cuda_allocation_base_and_offset( + device_ptr + ) + transfer = self._client(buffer_index).transfer_load_cuda( + cuda_ipc_handle=export_cuda_ipc_handle(cp_staging), + nbytes=total_bytes, + device_id=cuda_array_device_id(cp_staging), + device_ptr=device_ptr, + allocation_base_ptr=allocation_base, + allocation_offset=allocation_offset, + producer_pid=os.getpid(), + spans=spans, + ) + submitted_at = time.perf_counter() + return _InflightLoadBatch( + total_bytes=total_bytes, + per_req_ranges=per_req_ranges, + staging_lease=lease, + future=self._submit(transfer), + submitted_at=submitted_at, + ) + + def _consume_request( + self, + state: _InflightRequestLoad, + ) -> tuple[int, bool]: + active = state.active + if active is None: + return state.buffer_index, True + state.completed.append(self._consume_batch(active)) + if not state.batches: + self._log_request_timing(state) + if not state.request.future.done(): + state.request.future.set_result(None) + return state.buffer_index, True + state.active = self._submit_batch( + state.batches.popleft(), + state.buffer_index, + ) + return state.buffer_index, False + + def _consume_batch(self, state: _InflightLoadBatch) -> _LoadBatchTiming: + try: + wait_start = time.perf_counter() + response = state.future.result(timeout=120.0) + wait_ms = (time.perf_counter() - wait_start) * 1000 + ipc_ms = (time.perf_counter() - state.submitted_at) * 1000 + copy_start = time.perf_counter() + copies, copy_runs = self._restore_batch(state) + copy_ms = (time.perf_counter() - copy_start) * 1000 + sync_ms = 0.0 + if self._cuda_stream is None: + pass + else: + sync_start = time.perf_counter() + self._cuda_stream.synchronize() + sync_ms = (time.perf_counter() - sync_start) * 1000 + payload = response if isinstance(response, dict) else {} + stats = payload.get("transfer_stats_delta", {}) + stats = stats if isinstance(stats, dict) else {} + return _LoadBatchTiming( + bytes=state.total_bytes, + copies=copies, + copy_runs=copy_runs, + ipc_ms=ipc_ms, + wait_ms=wait_ms, + copy_ms=copy_ms, + worker_sync_ms=sync_ms, + transfer_open_ms=float(payload.get("transfer_open_ms", 0.0)), + transfer_load_ms=float(payload.get("transfer_load_ms", 0.0)), + transfer_sync_ms=float(payload.get("transfer_sync_ms", 0.0)), + l1_hits=int(stats.get("l1_hits", 0)), + l1_misses=int(stats.get("l1_misses", 0)), + l2_reads=int(stats.get("l2_reads", 0)), + ) + finally: + state.staging_lease.release() + + def _restore_batch(self, state: _InflightLoadBatch) -> tuple[int, int]: + staging = state.staging_lease.view + runs = build_load_copy_runs(state.per_req_ranges) + copies = 0 + context = ( + torch.cuda.stream(self._cuda_stream) + if self._cuda_stream is not None + else nullcontext() + ) + with context: + for run in runs: + copies += copy_staging_to_kv_cache( + staging=staging[run.start : run.end], + kv_caches=self._kv_caches, + layer_names=self._layer_names, + block_ids=run.block_ids, + slot_size=self._local_slot_size, + load_key_scale=self._load_key_scale, + load_value_scale=self._load_value_scale, + pos_offset=run.pos_offset, + rope_delta_scale=self._rope_delta_scale, + rope_base=self._rope_base, + rope_rotary_dim=self._rope_rotary_dim, + rope_is_neox_style=self._rope_is_neox_style, + ) + return copies, len(runs) + + def _log_request_timing(self, state: _InflightRequestLoad) -> None: + rows = state.completed + logger.debug( + "[CONNECTOR] load timing req=%s batches=%d bytes=%d copy_runs=%d " + "gpu_copies=%d ipc_ms=%.3f dispatcher_wait_ms=%.3f copy_ms=%.3f " + "worker_sync_ms=%.3f transfer_open_ms=%.3f " + "transfer_load_ms=%.3f transfer_sync_ms=%.3f l1_hits=%d " + "l1_misses=%d l2_reads=%d", + state.request.req_id, + len(rows), + sum(row.bytes for row in rows), + sum(row.copy_runs for row in rows), + sum(row.copies for row in rows), + sum(row.ipc_ms for row in rows), + sum(row.wait_ms for row in rows), + sum(row.copy_ms for row in rows), + sum(row.worker_sync_ms for row in rows), + sum(row.transfer_open_ms for row in rows), + sum(row.transfer_load_ms for row in rows), + sum(row.transfer_sync_ms for row in rows), + sum(row.l1_hits for row in rows), + sum(row.l1_misses for row in rows), + sum(row.l2_reads for row in rows), + ) + + def _register_staging_buffers(self) -> None: + pool = self._staging_pool + if self._staging_registered or pool is None: + return + try: + for buffer_index in range(pool.depth): + tensor = pool.buffer(buffer_index) + cp_tensor = cupy.asarray(tensor) + device_ptr = cuda_array_pointer(cp_tensor) + allocation_base, allocation_offset = cuda_allocation_base_and_offset( + device_ptr + ) + self._submit( + self._client(buffer_index).register_load_staging_cuda( + buffer_index=buffer_index, + cuda_ipc_handle=export_cuda_ipc_handle(cp_tensor), + allocation_bytes=int(tensor.numel()), + device_id=cuda_array_device_id(cp_tensor), + device_ptr=device_ptr, + allocation_base_ptr=allocation_base, + allocation_offset=allocation_offset, + producer_pid=os.getpid(), + ) + ).result(timeout=120.0) + except Exception as exc: # noqa: BLE001 + logger.warning( + "[CONNECTOR] registered load staging unavailable; falling back " + "to per-load CUDA IPC payloads: %s", + exc, + ) + self._staging_registered = False + return + self._staging_registered = True + logger.info("[CONNECTOR] registered %d load staging buffers", pool.depth) + + +@dataclass(frozen=True) +class LoadCopyRun: + """Describe one contiguous staging range with a shared KV transform.""" + + start: int + end: int + block_ids: list[int] + pos_offset: int + + +def build_load_read_plan( + reqs_to_load: dict[str, ReqLoadSpec], + slot_size: int, + include_req_ids: bool = False, +) -> tuple[int, list[dict[str, int]], list[Any]]: + """Build one combined server read and staging restore plan. + + Args: + reqs_to_load: Request IDs mapped to scheduler load specifications. + slot_size: Bytes stored for one rank-local KV slot. + include_req_ids: Include request IDs in restore ranges when true. + + Returns: + Total bytes, server read spans, and ranges mapping staging back to + request specifications. + + Async/thread-safety: + Pure CPU planning; safe to call from worker or load-loop threads. + """ + total_bytes = 0 + spans: list[dict[str, int]] = [] + per_req_ranges: list[Any] = [] + source_ranges: dict[tuple[str, int, int, int, int], tuple[int, int]] = {} + for req_id, spec in reqs_to_load.items(): + num_slots = len(spec.block_ids) + if num_slots == 0: + continue + nbytes = num_slots * slot_size + source_key = ( + spec.chunk_key, + spec.start_slot, + spec.num_slots, + spec.file_offset, + nbytes, + ) + existing = source_ranges.get(source_key) + if existing is None: + start = total_bytes + end = start + nbytes + spans.append( + { + "target_offset": start, + "nbytes": nbytes, + "file_offset": spec.file_offset, + } + ) + source_ranges[source_key] = (start, end) + total_bytes = end + else: + start, end = existing + if include_req_ids: + per_req_ranges.append((start, end, req_id, spec)) + else: + per_req_ranges.append((start, end, spec)) + return total_bytes, spans, per_req_ranges + + +def build_load_read_batches( + reqs_to_load: dict[str, ReqLoadSpec], + slot_size: int, + max_batch_bytes: int, + include_req_ids: bool = False, +) -> list[tuple[int, list[dict[str, int]], list[Any]]]: + """Split load work into staging-capacity-bounded read plans. + + Args: + reqs_to_load: Request IDs mapped to scheduler load specifications. + slot_size: Bytes stored for one rank-local KV slot. + max_batch_bytes: Maximum staging bytes in one transfer. + include_req_ids: Include request IDs in restore ranges when true. + + Returns: + Ordered read plans; requests larger than the cap are split on slot + boundaries. + + Async/thread-safety: + Pure CPU planning; safe to call from worker or load-loop threads. + """ + if slot_size <= 0: + raise ValueError("slot_size must be positive") + if max_batch_bytes <= 0: + raise ValueError("max_batch_bytes must be positive") + max_slots = max(1, max_batch_bytes // slot_size) + batches: list[tuple[int, list[dict[str, int]], list[Any]]] = [] + current: dict[str, ReqLoadSpec] = {} + current_slots = 0 + synthetic_id = 0 + + def flush() -> None: + nonlocal current, current_slots + if current: + batches.append( + build_load_read_plan( + current, + slot_size, + include_req_ids=include_req_ids, + ) + ) + current = {} + current_slots = 0 + + for req_id, spec in reqs_to_load.items(): + cursor = 0 + while cursor < len(spec.block_ids): + if current_slots >= max_slots: + flush() + available = max_slots - current_slots + take = min(available, len(spec.block_ids) - cursor) + if take <= 0: + flush() + continue + part = spec.block_ids[cursor : cursor + take] + batch_spec = replace( + spec, + start_slot=spec.start_slot + cursor, + num_slots=take, + block_ids=part, + file_offset=spec.file_offset + cursor * slot_size, + ) + key = ( + req_id + if cursor == 0 and take == len(spec.block_ids) + else f"{req_id}#{synthetic_id}" + ) + synthetic_id += 1 + current[key] = batch_spec + current_slots += take + cursor += take + flush() + return batches + + +def build_load_copy_runs( + per_req_ranges: list[tuple[int, int, ReqLoadSpec]], +) -> list[LoadCopyRun]: + """Merge adjacent restore ranges with the same position transform. + + Args: + per_req_ranges: Per-request staging ranges from a read plan. + + Returns: + Ordered contiguous copy runs. + + Async/thread-safety: + Pure CPU planning; safe to call from worker or load-loop threads. + """ + runs: list[LoadCopyRun] = [] + run_start = -1 + run_end = -1 + run_pos_offset = 0 + run_block_ids: list[int] = [] + + def flush() -> None: + nonlocal run_start, run_end, run_pos_offset, run_block_ids + if run_start >= 0 and run_block_ids: + runs.append( + LoadCopyRun( + start=run_start, + end=run_end, + block_ids=run_block_ids, + pos_offset=run_pos_offset, + ) + ) + run_start = -1 + run_end = -1 + run_pos_offset = 0 + run_block_ids = [] + + for item in per_req_ranges: + if len(item) == 3: + start, end, spec = item + else: + start, end, _req_id, spec = item + if not spec.block_ids: + continue + if run_start >= 0 and start == run_end and spec.pos_offset == run_pos_offset: + run_end = end + run_block_ids.extend(spec.block_ids) + continue + flush() + run_start = start + run_end = end + run_pos_offset = spec.pos_offset + run_block_ids = list(spec.block_ids) + flush() + return runs diff --git a/daser/connector/worker/memory.py b/daser/connector/worker/memory.py new file mode 100644 index 0000000..7c84f4a --- /dev/null +++ b/daser/connector/worker/memory.py @@ -0,0 +1,268 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Bounded worker-side CUDA staging memory.""" + +from __future__ import annotations + +# Standard +from collections.abc import Callable +from dataclasses import dataclass +from typing import Protocol + +# Third Party +import torch + +DEFAULT_STORE_STAGING_BYTES = 1536 << 20 +DEFAULT_STAGING_BUDGET_BYTES = 6144 << 20 +MIN_STORE_STAGING_BYTES = 64 << 20 + + +class _CudaStagingLeaseOwner(Protocol): + """Pool interface required by ``CudaStagingLease``.""" + + def release(self, lease: "CudaStagingLease") -> None: + """Return a lease to its owning pool. + + Args: + lease: Lease previously returned by that pool. + """ + + +@dataclass +class CudaStagingLease: + """One logical staging allocation leased from a fixed CUDA pool. + + Args: + pool: Owning pool that receives the allocation on release. + tensor: Backing tensor, possibly larger than ``nbytes``. + nbytes: Logical byte count used by the current transfer. + + Async/thread-safety: + The worker thread releases a lease only after every CUDA IPC user has + finished with it. + """ + + pool: _CudaStagingLeaseOwner + tensor: torch.Tensor + nbytes: int + _released: bool = False + + @property + def view(self) -> torch.Tensor: + """Return the active one-dimensional uint8 tensor view.""" + return self.tensor[: self.nbytes] + + def release(self) -> None: + """Return this lease to its owning pool once.""" + if self._released: + return + self._released = True + self.pool.release(self) + + +class FixedCudaStagingPool: + """Preallocate and lease fixed-size worker CUDA staging buffers. + + Args: + device: Device on which staging tensors are allocated. + buffer_bytes: Size of each fixed staging buffer. + depth: Number of fixed buffers to allocate. + + Async/thread-safety: + One vLLM worker thread owns the pool. It never allocates after + construction; callers release leases after async transfer completes. + """ + + def __init__( + self, + device: torch.device, + buffer_bytes: int, + depth: int, + ) -> None: + if buffer_bytes <= 0: + raise ValueError("buffer_bytes must be positive") + if depth <= 0: + raise ValueError("depth must be positive") + self._buffer_bytes = buffer_bytes + self._buffers: list[torch.Tensor] = [ + torch.empty(buffer_bytes, dtype=torch.uint8, device=device) + for _ in range(depth) + ] + self._free_indices: list[int] = list(range(depth)) + + @property + def buffer_bytes(self) -> int: + """Return the fixed capacity of each staging buffer.""" + return self._buffer_bytes + + @property + def available(self) -> int: + """Return the number of currently free staging buffers.""" + return len(self._free_indices) + + @property + def depth(self) -> int: + """Return the number of fixed buffers in the pool.""" + return len(self._buffers) + + def buffer(self, index: int) -> torch.Tensor: + """Return a backing tensor for one-time CUDA IPC registration. + + Args: + index: Fixed buffer index. + + Returns: + The full preallocated backing tensor. + + Async/thread-safety: + The caller must not mutate the tensor outside the lease protocol. + """ + if index < 0 or index >= len(self._buffers): + raise ValueError(f"fixed staging buffer index out of range: {index}") + return self._buffers[index] + + def acquire( + self, + nbytes: int, + wait_for_release: Callable[[], None] | None = None, + ) -> CudaStagingLease: + """Lease one preallocated staging buffer. + + Args: + nbytes: Logical transfer byte count. + wait_for_release: Optional callback invoked once when every buffer + is leased. + + Returns: + Lease whose view is limited to ``nbytes``. + + Raises: + ValueError: If ``nbytes`` is invalid. + RuntimeError: If no buffer is available after the callback. + """ + if nbytes < 0: + raise ValueError("nbytes must be non-negative") + if nbytes > self._buffer_bytes: + raise ValueError( + f"staging request {nbytes} exceeds fixed staging buffer " + f"{self._buffer_bytes}" + ) + if not self._free_indices and wait_for_release is not None: + wait_for_release() + if not self._free_indices: + raise RuntimeError("no fixed staging buffers available") + return self.acquire_index(self._free_indices[0], nbytes) + + def acquire_index(self, index: int, nbytes: int) -> CudaStagingLease: + """Lease a specific preallocated staging buffer. + + Args: + index: Fixed buffer index. + nbytes: Logical transfer byte count. + + Returns: + Lease whose view is limited to ``nbytes``. + + Raises: + ValueError: If ``index`` or ``nbytes`` is invalid. + RuntimeError: If the requested buffer is already leased. + """ + if index < 0 or index >= len(self._buffers): + raise ValueError(f"fixed staging buffer index out of range: {index}") + if nbytes < 0: + raise ValueError("nbytes must be non-negative") + if nbytes > self._buffer_bytes: + raise ValueError( + f"staging request {nbytes} exceeds fixed staging buffer " + f"{self._buffer_bytes}" + ) + if index not in self._free_indices: + raise RuntimeError(f"fixed staging buffer {index} is not available") + self._free_indices.remove(index) + return CudaStagingLease( + pool=self, + tensor=self._buffers[index], + nbytes=nbytes, + ) + + def release(self, lease: CudaStagingLease) -> None: + """Return a lease to this pool. + + Args: + lease: Lease previously returned by this pool. + + Raises: + ValueError: If the lease belongs to another pool. + """ + for index, tensor in enumerate(self._buffers): + if tensor is lease.tensor: + if index not in self._free_indices: + self._free_indices.append(index) + self._free_indices.sort() + return + raise ValueError("lease does not belong to this fixed staging pool") + + +def derive_staging_layout( + device: torch.device, + local_slot_size: int, + max_load_inflight: int, + reserve_bytes: int, +) -> tuple[int, int, int, int]: + """Partition one CUDA staging budget between load and store pools. + + Args: + device: Device that owns worker-side staging tensors. + local_slot_size: Minimum buffer size required for one KV slot. + max_load_inflight: Maximum useful load pool depth. + reserve_bytes: Free CUDA memory kept outside staging pools. + + Returns: + Buffer bytes, load depth, store depth, and combined allocation bytes. + + Raises: + ValueError: If one buffer per direction cannot fit the budget. + + Async/thread-safety: + Reads CUDA device properties during worker initialization before + request traffic starts. + """ + if local_slot_size <= 0: + raise ValueError("local_slot_size must be positive") + if max_load_inflight <= 0: + raise ValueError("max_load_inflight must be positive") + if device.type != "cuda": + buffer_bytes = max(DEFAULT_STORE_STAGING_BYTES, local_slot_size) + budget_bytes = DEFAULT_STAGING_BUDGET_BYTES + else: + props = torch.cuda.get_device_properties(device) + total = int(props.total_memory) + try: + free, _ = torch.cuda.mem_get_info(device) + free = int(free) + except (RuntimeError, TypeError, ValueError): + free = total + usable = max(0, free - max(0, reserve_bytes)) + buffer_bytes = max( + local_slot_size, + min( + DEFAULT_STORE_STAGING_BYTES, + max(MIN_STORE_STAGING_BYTES, min(total // 50, free // 10)), + ), + ) + budget_bytes = min( + DEFAULT_STAGING_BUDGET_BYTES, + (2 * total) // 25, + usable, + ) + + minimum = 2 * buffer_bytes + if budget_bytes < minimum: + raise ValueError( + "CUDA staging budget cannot fit one load and one store buffer: " + f"required={minimum} available={budget_bytes}" + ) + total_depth = budget_bytes // buffer_bytes + store_depth = min(2, total_depth // 2) + load_depth = min(max_load_inflight, total_depth - store_depth) + allocated_bytes = buffer_bytes * (load_depth + store_depth) + return buffer_bytes, load_depth, store_depth, allocated_bytes diff --git a/daser/connector/worker/runtime.py b/daser/connector/worker/runtime.py new file mode 100644 index 0000000..f4d75cb --- /dev/null +++ b/daser/connector/worker/runtime.py @@ -0,0 +1,650 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +# Standard +from typing import TYPE_CHECKING, Any + +# Third Party +import torch +from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole + +if TYPE_CHECKING: + # Third Party + from vllm.attention import AttentionMetadata + from vllm.forward_context import ForwardContext + +# First Party +from daser.connector.ipc_client import IPCClientSync +from daser.connector.metadata import ( + DaserConnectorMeta, +) +from daser.connector.worker.load import LoadPipeline +from daser.connector.worker.memory import ( + FixedCudaStagingPool, + derive_staging_layout, +) +from daser.connector.worker.staging import ( + CROSS_LAYER_KV_CACHE_KEY, + FUSED_RESTORE_MIN_SLOTS, +) +from daser.connector.worker.store import StorePipeline +from daser.logging import init_logger +from daser.ops.rope_apply import ( + apply_rope_delta_to_key_block as _apply_rope_delta_to_key_block, +) +from daser.ops.rope_apply import ( + apply_rope_delta_to_kv_key_block_table, + restore_cross_layer_kv_cache_table, +) + +logger = init_logger(__name__) + +_ROPE_WARMUP_BLOCKS = 1 +_LOAD_REQUEST_MAX_INFLIGHT = 8 +_LOAD_STAGING_RESERVE_BYTES = 1 << 30 + + +def _local_slot_bytes(connector: Any) -> int: + """Return per-rank slot bytes, falling back for TP=1 test probes.""" + local_slot_size = int(getattr(connector, "_local_slot_size", 0)) + return local_slot_size or int(connector._slot_size) # noqa: SLF001 + + +def _validate_tp_layout( + local_slot_size: int, + storage_slot_size: int, + tp_size: int, + server_tp_size: int, + tp_rank: int, + rank_stride_bytes: int = 0, +) -> None: + """Validate worker KV geometry against the server-owned TP layout. + + Args: + local_slot_size: Slot bytes measured from the worker KV tensor. + storage_slot_size: Aggregate slot bytes reported by the server. + tp_size: vLLM worker tensor-parallel size. + server_tp_size: Tensor-parallel size reported by the server. + tp_rank: Current vLLM tensor-parallel rank. + rank_stride_bytes: Byte distance between server-owned rank lanes. + + Raises: + ValueError: if rank counts or slot geometry do not match. + + Async/thread-safety: + Pure startup validation called before request traffic. + """ + if tp_size <= 0 or not 0 <= tp_rank < tp_size: + raise ValueError(f"invalid TP rank {tp_rank} for size {tp_size}") + if not storage_slot_size: + return + if server_tp_size != tp_size: + raise ValueError( + f"vLLM TP size {tp_size} does not match DaseR TP size {server_tp_size}" + ) + if local_slot_size * tp_size != storage_slot_size: + raise ValueError( + "worker KV slot geometry does not match DaseR storage layout: " + f"local={local_slot_size} tp={tp_size} storage={storage_slot_size}" + ) + if tp_size > 1 and rank_stride_bytes <= 0: + raise ValueError("DaseR runtime config is missing TP rank stride") + + +def _warm_rope_apply_backends( + device: torch.device, + dtype: torch.dtype, + block_tokens: int, + heads: int, + head_dim: int, + rotary_dim: int, + rope_base: float, + is_neox_style: bool, +) -> None: + """Warm dynamic-shape RoPE apply operators. + + Args: + device: device that owns the worker KV cache. + dtype: KV cache dtype. + block_tokens: tokens per cache block. + heads: number of KV heads. + head_dim: per-head dimension. + rotary_dim: number of dimensions covered by RoPE. + rope_base: RoPE theta/base. + is_neox_style: True for split-half rotation, False for interleaved. + + Async/thread-safety: + Runs synchronously during worker KV cache registration, before request + traffic starts. It launches CUDA work on the current stream. TileLang + failures are surfaced to avoid silently entering a slow restore path. + """ + if device.type != "cuda" or rotary_dim <= 0 or head_dim < rotary_dim: + return + sample = torch.empty( + (_ROPE_WARMUP_BLOCKS, block_tokens, heads, head_dim), + dtype=dtype, + device=device, + ) + _apply_rope_delta_to_key_block( + sample, + delta=1, + rope_base=rope_base, + rotary_dim=rotary_dim, + is_neox_style=is_neox_style, + ) + torch.cuda.synchronize(device) + + +def _warm_cross_layer_restore_backends( + device: torch.device, + dtype: torch.dtype, + layers: int, + block_tokens: int, + heads: int, + head_dim: int, + rotary_dim: int, + rope_base: float, + is_neox_style: bool, +) -> None: + """Warm cross-layer staging restore TileLang kernels. + + Args: + device: device that owns the worker KV cache. + dtype: KV cache dtype. + layers: number of model KV layers. + block_tokens: tokens per cache block. + heads: number of KV heads. + head_dim: per-head dimension. + rotary_dim: number of dimensions covered by RoPE. + rope_base: RoPE theta/base. + is_neox_style: True for split-half rotation, False for interleaved. + + Async/thread-safety: + Runs synchronously during worker KV cache registration, before request + traffic starts. TileLang import/compile failures are surfaced to avoid + silently entering a slow restore path. + """ + if device.type != "cuda" or rotary_dim <= 0 or head_dim < rotary_dim: + return + inv_freq = 1.0 / ( + rope_base + ** ( + torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=device) + / rotary_dim + ) + ) + freqs = inv_freq + cos_table = freqs.cos().contiguous() + sin_table = freqs.sin().contiguous() + for blocks, use_fused_restore in ( + (_ROPE_WARMUP_BLOCKS, False), + (FUSED_RESTORE_MIN_SLOTS, True), + ): + sample = torch.empty( + blocks, + layers, + 2, + block_tokens, + heads, + head_dim, + dtype=dtype, + device=device, + ) + if use_fused_restore: + dst = torch.empty_like(sample) + restore_cross_layer_kv_cache_table( + sample, + dst, + cos_table=cos_table, + sin_table=sin_table, + rotary_dim=rotary_dim, + is_neox_style=is_neox_style, + ) + else: + apply_rope_delta_to_kv_key_block_table( + sample, + cos_table=cos_table, + sin_table=sin_table, + rotary_dim=rotary_dim, + is_neox_style=is_neox_style, + ) + torch.cuda.synchronize(device) + + +class WorkerRuntime: + """Own worker KV layout, step metadata, pipelines, and completion state. + + Async/thread-safety: + Public methods are called on vLLM worker threads. Blocking NVMe work is + submitted to the runtime's load or store pipeline loop. + """ + + def __init__( + self, + *, + socket_path: str, + transfer_mode: str, + skip_l2: bool, + tp_size: int, + tp_rank: int, + server_tp_size: int, + slot_size: int, + store_path: str, + rank_stride_bytes: int, + rope_base: float, + rope_rotary_dim: int, + rope_is_neox_style: bool, + rope_delta_scale: float, + load_key_scale: float, + load_value_scale: float, + kv_cache_config: Any, + ) -> None: + self._socket_path = socket_path + self._transfer_mode = transfer_mode + self._skip_l2 = skip_l2 + self._tp_size = tp_size + self._tp_rank = tp_rank + self._server_tp_size = server_tp_size + self._slot_size = slot_size + self._local_slot_size = 0 + self._store_path = store_path + self._rank_stride_bytes = rank_stride_bytes + self._rope_base = rope_base + self._rope_rotary_dim = rope_rotary_dim + self._rope_is_neox_style = rope_is_neox_style + self._rope_delta_scale = rope_delta_scale + self._load_key_scale = load_key_scale + self._load_value_scale = load_value_scale + self._kv_cache_config = kv_cache_config + self._role = KVConnectorRole.WORKER + self._transfer_ready = False + self._pipelines_initialized = False + self._load_pipeline = LoadPipeline(socket_path, _LOAD_REQUEST_MAX_INFLIGHT) + self._store_pipeline = StorePipeline(socket_path) + self._kv_caches: dict[str, torch.Tensor] = {} + self._layer_names: list[str] = [] + self._layer_idx_map: dict[str, int] = {} + self._meta: DaserConnectorMeta | None = None + + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]) -> None: + """Register the per-layer KV cache tensors. + + Args: + kv_caches: dict mapping layer_name -> KV tensor. + """ + self._kv_caches = kv_caches + self._layer_names = list(kv_caches.keys()) + self._layer_idx_map = {name: idx for idx, name in enumerate(self._layer_names)} + sample = next(iter(kv_caches.values()), None) + if sample is not None: + logger.info( + "[CONNECTOR] register_kv_caches: %d layers, first shape=%s dtype=%s", + len(kv_caches), + sample.shape, + sample.dtype, + ) + + if self._layer_names and sample is not None: + num_blocks = sample.shape[1] if sample.dim() >= 2 else 1 + layer_size = sample.nbytes // num_blocks + local_slot_size = layer_size * len(self._layer_names) + tp_size = self._tp_size + _validate_tp_layout( + local_slot_size, + self._slot_size, + tp_size, + self._server_tp_size, + self._tp_rank, + self._rank_stride_bytes, + ) + self._local_slot_size = local_slot_size + if self._slot_size == 0: + self._slot_size = local_slot_size * tp_size + logger.info( + "[CONNECTOR] registered local_slot_size=%d from %d layers", + self._local_slot_size, + len(self._layer_names), + ) + + if sample is not None: + self._configure_pipelines(sample) + if sample.dim() >= 5: + _warm_rope_apply_backends( + device=sample.device, + dtype=sample.dtype, + block_tokens=int(sample.shape[-3]), + heads=int(sample.shape[-2]), + head_dim=int(sample.shape[-1]), + rotary_dim=self._rope_rotary_dim, + rope_base=self._rope_base, + is_neox_style=self._rope_is_neox_style, + ) + + self._init_server_transfer() + + def register_cross_layers_kv_cache( + self, + kv_cache: torch.Tensor, + attn_backend: type[Any], + ) -> None: + """Register vLLM's cross-layer KV cache tensor. + + Args: + kv_cache: vLLM tensor whose logical layout starts with + ``[blocks, layers, 2, block_tokens, heads, head_dim]`` for the + NHD layout DaseR requests. + attn_backend: Attention backend that created ``kv_cache``. + + Async/thread-safety: + Called once during worker initialization before request traffic. + """ + kv_cache_config = self._kv_cache_config + layer_names: list[str] = [] + if kv_cache_config is not None: + for group in getattr(kv_cache_config, "kv_cache_groups", []): + layer_names.extend(list(getattr(group, "layer_names", []))) + if not layer_names: + layer_count = int(kv_cache.shape[1]) if kv_cache.dim() >= 2 else 0 + layer_names = [f"layer.{idx}" for idx in range(layer_count)] + self._kv_caches = {CROSS_LAYER_KV_CACHE_KEY: kv_cache} + self._layer_names = layer_names + self._layer_idx_map = {name: idx for idx, name in enumerate(self._layer_names)} + if kv_cache.dim() < 6: + logger.warning( + "[CONNECTOR] cross-layer KV cache has unsupported shape=%s", + tuple(kv_cache.shape), + ) + return + layer_size = kv_cache[0, 0].nbytes + local_slot_size = layer_size * len(self._layer_names) + tp_size = self._tp_size + _validate_tp_layout( + local_slot_size, + self._slot_size, + tp_size, + self._server_tp_size, + self._tp_rank, + self._rank_stride_bytes, + ) + self._local_slot_size = local_slot_size + if self._slot_size == 0: + self._slot_size = local_slot_size * tp_size + logger.info( + "[CONNECTOR] registered cross-layer local_slot_size=%d from %d layers", + self._local_slot_size, + len(self._layer_names), + ) + load_staging_depth = self._configure_pipelines(kv_cache) + logger.info( + "[CONNECTOR] register_cross_layers_kv_cache: layers=%d shape=%s " + "dtype=%s load_request_max_inflight=%d load_staging_depth=%d", + len(self._layer_names), + tuple(kv_cache.shape), + kv_cache.dtype, + _LOAD_REQUEST_MAX_INFLIGHT, + load_staging_depth, + ) + _warm_rope_apply_backends( + device=kv_cache.device, + dtype=kv_cache.dtype, + block_tokens=int(kv_cache.shape[-3]), + heads=int(kv_cache.shape[-2]), + head_dim=int(kv_cache.shape[-1]), + rotary_dim=self._rope_rotary_dim, + rope_base=self._rope_base, + is_neox_style=self._rope_is_neox_style, + ) + _warm_cross_layer_restore_backends( + device=kv_cache.device, + dtype=kv_cache.dtype, + layers=int(kv_cache.shape[1]), + block_tokens=int(kv_cache.shape[-3]), + heads=int(kv_cache.shape[-2]), + head_dim=int(kv_cache.shape[-1]), + rotary_dim=self._rope_rotary_dim, + rope_base=self._rope_base, + is_neox_style=self._rope_is_neox_style, + ) + self._init_server_transfer() + + def bind_connector_metadata(self, connector_metadata: DaserConnectorMeta) -> None: + """Receive scheduler metadata before each forward pass. + + Args: + connector_metadata: DaserConnectorMeta from build_connector_meta. + """ + self._meta = connector_metadata + + def clear_connector_metadata(self) -> None: + """Clear metadata after forward pass completes.""" + self._meta = None + + def start_load_kv(self, forward_context: "ForwardContext", **kwargs: Any) -> None: + """Submit cache-hit requests to the load pipeline. + + Args: + forward_context: vLLM forward context for this step. + **kwargs: Additional vLLM hook arguments, currently unused. + + Async/thread-safety: + Called on the vLLM worker thread. Transfer and restore execute on + the load pipeline thread. + """ + del forward_context, kwargs + if self._meta is None or not self._meta.reqs_to_load: + return + reqs_to_load = dict(self._meta.reqs_to_load) + if not self._ensure_transfer_ready(): + self._load_pipeline.mark_failed( + reqs_to_load, + "server transfer config is not ready", + ) + return + self._load_pipeline.start(reqs_to_load) + + def wait_for_layer_load(self, layer_name: str) -> None: + """Return after request-level load completion restored every layer. + + Args: + layer_name: Layer reported by vLLM; no per-layer wait is required. + + Async/thread-safety: + Called on the vLLM worker thread. + """ + del layer_name + + def save_kv_layer( + self, + layer_name: str, + kv_layer: torch.Tensor, + attn_metadata: "AttentionMetadata", + **kwargs: Any, + ) -> None: + """Submit this layer's KV blocks for server-owned transfer. + + Args: + layer_name: name of the current attention layer. + kv_layer: full KV cache tensor for this layer. + attn_metadata: attention metadata (not directly used). + """ + if self._meta is None or not self._meta.reqs_to_store: + return + if not self._ensure_transfer_ready(): + return + + if layer_name not in self._layer_idx_map: + logger.warning( + "[CONNECTOR] save_kv_layer: unknown layer %s, skipping", layer_name + ) + + def wait_for_save(self) -> None: + """Queue stores until vLLM reports request completion.""" + if self._meta is None: + return + reqs_to_store = dict(self._meta.reqs_to_store) + commit_keys = { + spec.chunk_key for spec in reqs_to_store.values() if spec.block_ids + } + if commit_keys: + self._store_pipeline.queue_finished(reqs_to_store, commit_keys) + + def get_finished( + self, finished_req_ids: set[str] + ) -> tuple[set[str] | None, set[str] | None]: + """Collect completed background transfers after a worker step. + + Args: + finished_req_ids: Request IDs that vLLM finished in this step. + + Returns: + Finished-saving request IDs and finished async-loading request IDs. + """ + finished_recving = self._load_pipeline.collect_finished() + finished_sending = self._store_pipeline.collect_finished(finished_req_ids) + return finished_sending or None, finished_recving or None + + def get_block_ids_with_load_errors(self) -> set[int]: + """Return and clear block IDs whose async load failed. + + Returns: + vLLM block IDs that should be treated as invalid. + """ + return self._load_pipeline.take_invalid_block_ids() + + def shutdown(self) -> None: + """Stop the background IO loop.""" + if self._role != KVConnectorRole.WORKER: + return + self._load_pipeline.shutdown() + self._store_pipeline.shutdown() + + def _ensure_transfer_ready(self) -> bool: + """Refresh config and initialize both pipeline transfer clients.""" + if not self._transfer_ready: + self._refresh_runtime_config() + if not self._slot_size or (not self._skip_l2 and not self._store_path): + logger.warning( + "[CONNECTOR] server transfer config is not ready; start DaseR " + "server before sending requests", + ) + return False + + _validate_tp_layout( + _local_slot_bytes(self), + self._slot_size, + self._tp_size, + self._server_tp_size, + self._tp_rank, + self._rank_stride_bytes, + ) + self._load_pipeline.configure_rank_geometry( + self._rank_stride_bytes, + self._tp_rank, + ) + self._store_pipeline.configure_rank_geometry( + self._rank_stride_bytes, + self._tp_rank, + self._tp_size, + ) + self._transfer_ready = True + logger.info("[CONNECTOR] server transfer mode=%s", self._transfer_mode) + + if not self._pipelines_initialized: + self._load_pipeline.initialize_transfer() + self._store_pipeline.initialize_transfer() + self._pipelines_initialized = True + return True + + def _refresh_runtime_config(self) -> None: + """Refresh worker-owned storage geometry directly over sync IPC.""" + client = IPCClientSync(self._socket_path) + try: + config = client.get_runtime_config() + except Exception as exc: # noqa: BLE001 + logger.info("[CONNECTOR] runtime config unavailable: %s", exc) + return + finally: + client.close() + self._store_path = str(config.get("store_path", self._store_path)) + self._slot_size = int(config.get("slot_size", self._slot_size)) + self._server_tp_size = int( + config.get("tensor_parallel_size", self._server_tp_size) + ) + self._rank_stride_bytes = int( + config.get("rank_stride_bytes", self._rank_stride_bytes) + ) + self._skip_l2 = bool(config.get("skip_l2", self._skip_l2)) + self._transfer_mode = str(config.get("transfer_mode", self._transfer_mode)) + + def _init_server_transfer(self) -> None: + """Initialize both pipeline-owned transfer clients. + + Async/thread-safety: + Called on the worker thread after KV-cache registration. Each + pipeline performs its initialization on its private event loop. + """ + self._ensure_transfer_ready() + + def _configure_pipelines(self, sample: torch.Tensor) -> int: + """Configure load and store pipelines from one finalized KV layout. + + Args: + sample: Representative registered KV-cache tensor. + + Returns: + Number of preallocated load staging buffers. + + Async/thread-safety: + Called once on the worker thread during KV-cache registration. + """ + staging_bytes, load_depth, store_depth, allocated_bytes = derive_staging_layout( + sample.device, + self._local_slot_size, + _LOAD_REQUEST_MAX_INFLIGHT, + _LOAD_STAGING_RESERVE_BYTES, + ) + store_pool = FixedCudaStagingPool( + device=sample.device, + buffer_bytes=staging_bytes, + depth=store_depth, + ) + load_pool = FixedCudaStagingPool( + device=sample.device, + buffer_bytes=staging_bytes, + depth=load_depth, + ) + self._store_pipeline.configure( + kv_caches=self._kv_caches, + layer_names=self._layer_names, + layer_idx_map=self._layer_idx_map, + local_slot_size=self._local_slot_size, + rank_stride_bytes=self._rank_stride_bytes, + tp_rank=self._tp_rank, + tp_size=self._tp_size, + staging_bytes=staging_bytes, + staging_pool=store_pool, + ) + self._load_pipeline.configure( + kv_caches=self._kv_caches, + layer_names=self._layer_names, + local_slot_size=self._local_slot_size, + rank_stride_bytes=self._rank_stride_bytes, + tp_rank=self._tp_rank, + staging_pool=load_pool, + load_key_scale=self._load_key_scale, + load_value_scale=self._load_value_scale, + rope_delta_scale=self._rope_delta_scale, + rope_base=self._rope_base, + rope_rotary_dim=self._rope_rotary_dim, + rope_is_neox_style=self._rope_is_neox_style, + ) + logger.info( + "[CONNECTOR] preallocated staging buffer_bytes=%d total_bytes=%d " + "load_depth=%d store_depth=%d", + staging_bytes, + allocated_bytes, + load_pool.depth, + store_pool.depth, + ) + return load_pool.depth diff --git a/daser/connector/worker/staging.py b/daser/connector/worker/staging.py new file mode 100644 index 0000000..d573b9d --- /dev/null +++ b/daser/connector/worker/staging.py @@ -0,0 +1,567 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +# Third Party +import torch + +# First Party +from daser.logging import init_logger +from daser.ops.rope_apply import ( + apply_rope_delta_to_key_block as _apply_rope_delta_to_key_block, +) +from daser.ops.rope_apply import ( + apply_rope_delta_to_kv_key_block as _apply_rope_delta_to_kv_key_block, +) +from daser.ops.rope_apply import ( + apply_rope_delta_to_kv_key_block_table, + restore_cross_layer_kv_cache_table, +) + +DEFAULT_ROPE_DELTA_SCALE = 1.0 +CROSS_LAYER_KV_CACHE_KEY = "__cross_layers__" +FUSED_RESTORE_MIN_SLOTS = 32 + +logger = init_logger(__name__) +_rope_table_cache: dict[ + tuple[torch.device, int, float, int], + tuple[torch.Tensor, torch.Tensor], +] = {} + + +def synchronize_cuda_tensor(tensor: torch.Tensor) -> None: + """Synchronize pending CUDA work for a tensor before cross-process handoff. + + Args: + tensor: Tensor whose device stream must be visible across CUDA IPC. + + Async/thread-safety: + Synchronous barrier on the current worker thread. It is intentionally + conservative until CUDA IPC event handoff is added. + """ + if tensor.is_cuda: + torch.cuda.current_stream(tensor.device).synchronize() + + +def record_cuda_event(tensor: torch.Tensor) -> torch.cuda.Event | None: + """Record the tensor's current CUDA stream for deferred synchronization. + + Args: + tensor: Tensor whose producer stream should be observed. + + Returns: + A CUDA event recorded on the current stream, or ``None`` for CPU + tensors. + + Async/thread-safety: + Must be called on the producer thread before handing ``tensor`` to a + background task. + """ + if not tensor.is_cuda: + return None + event = torch.cuda.Event(blocking=False) + event.record(torch.cuda.current_stream(tensor.device)) + return event + + +def contiguous_block_range(block_ids: list[int]) -> tuple[int, int] | None: + """Return ``(start, stop)`` when block IDs are a contiguous range.""" + if not block_ids: + return None + start = block_ids[0] + for idx, block_id in enumerate(block_ids): + if block_id != start + idx: + return None + return start, start + len(block_ids) + + +def apply_rope_delta_to_key_block( + key_block: torch.Tensor, + delta: int, + rope_base: float, + rotary_dim: int, + is_neox_style: bool, +) -> None: + """Rotate an already-RoPE'd K block by a relative position delta. + + Args: + key_block: K cache block with shape [..., block_tokens, heads, head_dim]. + delta: relative RoPE position delta to apply in place. + rope_base: RoPE theta/base. + rotary_dim: number of head dimensions covered by RoPE. + is_neox_style: True for split-half rotation, False for interleaved. + + Returns: + None. ``key_block`` is modified in place. + + Async/thread-safety: + Performs tensor work on the current PyTorch stream. + """ + _apply_rope_delta_to_key_block( + key_block, + delta=delta, + rope_base=rope_base, + rotary_dim=rotary_dim, + is_neox_style=is_neox_style, + ) + + +def apply_rope_delta_to_kv_key_block( + kv_block: torch.Tensor, + delta: int, + rope_base: float, + rotary_dim: int, + is_neox_style: bool, +) -> None: + """Rotate K entries inside a full KV staging block by a RoPE delta. + + Args: + kv_block: KV staging tensor with shape + ``[blocks, layers, 2, block_tokens, heads, head_dim]``. + delta: relative RoPE position delta to apply in place. + rope_base: RoPE theta/base. + rotary_dim: number of head dimensions covered by RoPE. + is_neox_style: True for split-half rotation, False for interleaved. + + Returns: + None. Only the key slice is modified in place. + + Async/thread-safety: + Performs tensor work on the current PyTorch stream. + """ + _apply_rope_delta_to_kv_key_block( + kv_block, + delta=delta, + rope_base=rope_base, + rotary_dim=rotary_dim, + is_neox_style=is_neox_style, + ) + + +def _transform_loaded_staging_batch( + staging_by_layer: torch.Tensor, + layer_sample: torch.Tensor, + load_key_scale: float, + load_value_scale: float, + pos_offset: int, + rope_delta_scale: float, + rope_base: float, + rope_rotary_dim: int, + rope_is_neox_style: bool, +) -> None: + """Apply load-time transforms once over all staging layers in a copy run.""" + if staging_by_layer.numel() == 0 or layer_sample.dim() < 4: + return + num_slots = int(staging_by_layer.shape[0]) + num_layers = int(staging_by_layer.shape[1]) + kv_batch = staging_by_layer.view(layer_sample.dtype).view( + num_slots, + num_layers, + *layer_sample.shape, + ) + if load_key_scale != 1.0: + kv_batch[:, :, 0].mul_(load_key_scale) + if load_value_scale != 1.0: + kv_batch[:, :, 1].mul_(load_value_scale) + if ( + not pos_offset + or rope_rotary_dim <= 0 + or layer_sample.shape[-1] < rope_rotary_dim + ): + return + if kv_batch.dim() == 6 and kv_batch.is_contiguous(): + apply_rope_delta_to_kv_key_block( + kv_batch, + delta=round(pos_offset * rope_delta_scale), + rope_base=rope_base, + rotary_dim=rope_rotary_dim, + is_neox_style=rope_is_neox_style, + ) + return + if layer_sample.dim() != 4: + return + apply_rope_delta_to_key_block( + kv_batch[:, :, 0], + delta=round(pos_offset * rope_delta_scale), + rope_base=rope_base, + rotary_dim=rope_rotary_dim, + is_neox_style=rope_is_neox_style, + ) + + +def _copy_staging_to_cross_layer_kv_cache( + staging_by_layer: torch.Tensor, + cross_layer_kv_cache: torch.Tensor, + block_ids: list[int], + load_key_scale: float, + load_value_scale: float, + pos_offset: int, + rope_delta_scale: float, + rope_base: float, + rope_rotary_dim: int, + rope_is_neox_style: bool, +) -> int: + """Copy staging bytes into a vLLM cross-layer KV cache in one bulk write.""" + num_slots = len(block_ids) + layer_sample = cross_layer_kv_cache[block_ids[0], 0] + src = staging_by_layer.view(cross_layer_kv_cache.dtype).view( + num_slots, + cross_layer_kv_cache.shape[1], + *layer_sample.shape, + ) + block_range = contiguous_block_range(block_ids) + dst_contiguous = False + start = 0 + stop = 0 + if block_range is not None: + start, stop = block_range + dst_contiguous = cross_layer_kv_cache[start:stop].is_contiguous() + can_rotate_target = ( + block_range is not None + and load_key_scale == 1.0 + and load_value_scale == 1.0 + and pos_offset + and rope_rotary_dim > 0 + and layer_sample.shape[-1] >= rope_rotary_dim + and src.is_contiguous() + and dst_contiguous + ) + if can_rotate_target: + dst = cross_layer_kv_cache[start:stop] + delta = round(pos_offset * rope_delta_scale) + if num_slots >= FUSED_RESTORE_MIN_SLOTS: + _restore_cross_layer_with_tables( + src, + dst, + delta=delta, + rope_base=rope_base, + rotary_dim=rope_rotary_dim, + is_neox_style=rope_is_neox_style, + ) + return 1 + dst.copy_(src) + _apply_rope_delta_with_tables( + dst, + delta=delta, + rope_base=rope_base, + rotary_dim=rope_rotary_dim, + is_neox_style=rope_is_neox_style, + ) + return 1 + _transform_loaded_staging_batch( + staging_by_layer, + layer_sample=layer_sample, + load_key_scale=load_key_scale, + load_value_scale=load_value_scale, + pos_offset=pos_offset, + rope_delta_scale=rope_delta_scale, + rope_base=rope_base, + rope_rotary_dim=rope_rotary_dim, + rope_is_neox_style=rope_is_neox_style, + ) + if block_range is None: + block_index = torch.tensor( + block_ids, + dtype=torch.long, + device=staging_by_layer.device, + ) + cross_layer_kv_cache.index_copy_(0, block_index, src) + else: + start, stop = block_range + cross_layer_kv_cache[start:stop].copy_(src) + return 1 + + +def _apply_rope_delta_with_tables( + kv_block: torch.Tensor, + delta: int, + rope_base: float, + rotary_dim: int, + is_neox_style: bool, +) -> None: + """Apply RoPE using cached trig tables.""" + if kv_block.device.type != "cuda" or kv_block.dtype not in ( + torch.bfloat16, + torch.float16, + torch.float32, + ): + raise ValueError("TileLang RoPE restore requires CUDA fp16/bf16/fp32 KV") + cos_table, sin_table = _get_rope_delta_tables( + kv_block.device, + delta=delta, + rope_base=rope_base, + rotary_dim=rotary_dim, + ) + apply_rope_delta_to_kv_key_block_table( + kv_block, + cos_table=cos_table, + sin_table=sin_table, + rotary_dim=rotary_dim, + is_neox_style=is_neox_style, + ) + + +def _restore_cross_layer_with_tables( + src_kv: torch.Tensor, + dst_kv: torch.Tensor, + delta: int, + rope_base: float, + rotary_dim: int, + is_neox_style: bool, +) -> None: + """Restore cross-layer KV using cached trig tables.""" + cos_table, sin_table = _get_rope_delta_tables( + src_kv.device, + delta=delta, + rope_base=rope_base, + rotary_dim=rotary_dim, + ) + restore_cross_layer_kv_cache_table( + src_kv, + dst_kv, + cos_table=cos_table, + sin_table=sin_table, + rotary_dim=rotary_dim, + is_neox_style=is_neox_style, + ) + + +def _get_rope_delta_tables( + device: torch.device, + delta: int, + rope_base: float, + rotary_dim: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Return cached fp32 RoPE delta cosine/sine tables.""" + key = (device, int(delta), float(rope_base), int(rotary_dim)) + cached = _rope_table_cache.get(key) + if cached is not None: + return cached + inv_freq = 1.0 / ( + rope_base + ** ( + torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=device) + / rotary_dim + ) + ) + freqs = int(delta) * inv_freq + tables = (freqs.cos().contiguous(), freqs.sin().contiguous()) + _rope_table_cache[key] = tables + return tables + + +def copy_staging_to_kv_cache( + staging: torch.Tensor, + kv_caches: dict[str, torch.Tensor], + layer_names: list[str], + block_ids: list[int], + slot_size: int, + load_key_scale: float = 1.0, + load_value_scale: float = 1.0, + pos_offset: int = 0, + rope_delta_scale: float = DEFAULT_ROPE_DELTA_SCALE, + rope_base: float = 10000.0, + rope_rotary_dim: int = 0, + rope_is_neox_style: bool = True, +) -> int: + """Copy slot-major staging bytes into vLLM KV cache blocks. + + Args: + staging: Contiguous uint8 tensor containing whole request KV bytes. + kv_caches: Per-layer vLLM KV cache tensors. + layer_names: Layer iteration order matching on-disk layout. + block_ids: vLLM KV block IDs corresponding to staging slots. + slot_size: Total bytes for all layers in one slot. + load_key_scale: Optional multiplier for loaded K tensors. + load_value_scale: Optional multiplier for loaded V tensors. + pos_offset: Position delta for loaded chunk reuse. + rope_delta_scale: Multiplier applied to pos_offset before RoPE update. + rope_base: RoPE theta/base. + rope_rotary_dim: Number of head dimensions covered by RoPE. + rope_is_neox_style: True for split-half rotation, False for interleaved. + + Returns: + Number of layer-level copy operations issued. + + Async/thread-safety: + Synchronous GPU tensor copies on the vLLM worker thread. + """ + if not block_ids or not layer_names: + return 0 + num_layers = len(layer_names) + layer_size = slot_size // num_layers + num_slots = len(block_ids) + staging_by_layer = staging.view(num_slots, num_layers, layer_size) + cross_layer_kv_cache = kv_caches.get(CROSS_LAYER_KV_CACHE_KEY) + if ( + cross_layer_kv_cache is not None + and cross_layer_kv_cache.dim() >= 6 + and cross_layer_kv_cache.shape[1] == num_layers + ): + return _copy_staging_to_cross_layer_kv_cache( + staging_by_layer=staging_by_layer, + cross_layer_kv_cache=cross_layer_kv_cache, + block_ids=block_ids, + load_key_scale=load_key_scale, + load_value_scale=load_value_scale, + pos_offset=pos_offset, + rope_delta_scale=rope_delta_scale, + rope_base=rope_base, + rope_rotary_dim=rope_rotary_dim, + rope_is_neox_style=rope_is_neox_style, + ) + first_kv = next( + (kv_caches[name] for name in layer_names if kv_caches.get(name) is not None), + None, + ) + if first_kv is not None: + layer_sample = ( + first_kv[:, block_ids[0]] if first_kv.dim() >= 2 else first_kv[block_ids[0]] + ) + _transform_loaded_staging_batch( + staging_by_layer, + layer_sample=layer_sample, + load_key_scale=load_key_scale, + load_value_scale=load_value_scale, + pos_offset=pos_offset, + rope_delta_scale=rope_delta_scale, + rope_base=rope_base, + rope_rotary_dim=rope_rotary_dim, + rope_is_neox_style=rope_is_neox_style, + ) + block_range = contiguous_block_range(block_ids) + block_index = ( + None + if block_range is not None + else torch.tensor(block_ids, dtype=torch.long, device=staging.device) + ) + + copies = 0 + for layer_idx, layer_name in enumerate(layer_names): + kv_tensor = kv_caches.get(layer_name) + if kv_tensor is None: + continue + # KV tensors are either block-major ([blocks, ...]) or kv-major + # ([2, blocks, ...]); block_dim points at the block axis. + block_dim = 1 if kv_tensor.dim() >= 2 else 0 + sample = ( + kv_tensor[:, block_ids[0]] if block_dim == 1 else kv_tensor[block_ids[0]] + ) + src = ( + staging_by_layer[:, layer_idx, :] + .view(kv_tensor.dtype) + .view(num_slots, *sample.shape) + ) + # staging is slot-major (slots first); align it to the block axis. + src = src.movedim(0, block_dim) + if block_range is None: + if block_index is None: + raise RuntimeError("block_index is required for non-contiguous IDs") + kv_tensor.index_copy_(block_dim, block_index, src) + else: + start, stop = block_range + if block_dim == 1: + kv_tensor[:, start:stop].copy_(src) + else: + kv_tensor[start:stop].copy_(src) + copies += 1 + return copies + + +def copy_kv_cache_to_staging( + staging: torch.Tensor, + kv_layer: torch.Tensor, + layer_idx: int, + block_ids: list[int], + num_layers: int, + slot_size: int, + block_index: torch.Tensor | None = None, +) -> None: + """Copy one vLLM KV layer for requested blocks into slot-major staging. + + Args: + staging: Contiguous uint8 tensor with slot-major DaseR layout. + kv_layer: vLLM KV cache tensor for one attention layer. + layer_idx: Index of ``kv_layer`` in the DaseR on-disk layer order. + block_ids: vLLM KV block IDs to persist. + num_layers: Total number of KV layers in the model. + slot_size: Total bytes for all layers in one slot. + block_index: Optional prebuilt CUDA/CPU tensor containing block IDs. + + Async/thread-safety: + Synchronous GPU tensor copies on the vLLM worker thread. + """ + if not block_ids: + return + layer_size = slot_size // num_layers + num_slots = len(block_ids) + staging_by_layer = staging.view(num_slots, num_layers, layer_size) + if block_index is None: + block_index = torch.tensor(block_ids, dtype=torch.long, device=kv_layer.device) + if kv_layer.dim() >= 2: + block_range = contiguous_block_range(block_ids) + if block_range is None: + src = kv_layer.index_select(1, block_index).movedim(1, 0) + else: + start, stop = block_range + src = kv_layer[:, start:stop].movedim(1, 0) + else: + block_range = contiguous_block_range(block_ids) + if block_range is None: + src = kv_layer.index_select(0, block_index) + else: + start, stop = block_range + src = kv_layer[start:stop] + dst = ( + staging_by_layer[:, layer_idx, :] + .view(kv_layer.dtype) + .view(num_slots, *src.shape[1:]) + ) + dst.copy_(src) + + +def copy_cross_layer_kv_cache_to_staging( + staging: torch.Tensor, + kv_cache: torch.Tensor, + block_ids: list[int], + num_layers: int, + slot_size: int, + block_index: torch.Tensor | None = None, +) -> None: + """Copy vLLM cross-layer KV blocks into slot-major staging bytes. + + Args: + staging: Contiguous uint8 tensor with slot-major DaseR layout. + kv_cache: vLLM cross-layer KV cache tensor with blocks as dim 0 and + layers as dim 1. + block_ids: vLLM KV block IDs to persist. + num_layers: Total number of KV layers in the model. + slot_size: Total bytes for all layers in one slot. + block_index: Optional prebuilt tensor containing block IDs. + + Async/thread-safety: + Synchronous GPU tensor copy on the vLLM worker thread. + """ + if not block_ids: + return + layer_size = slot_size // num_layers + num_slots = len(block_ids) + staging_by_layer = staging.view(num_slots, num_layers, layer_size) + block_range = contiguous_block_range(block_ids) + if block_range is None: + if block_index is None: + block_index = torch.tensor( + block_ids, + dtype=torch.long, + device=kv_cache.device, + ) + src = kv_cache.index_select(0, block_index) + else: + start, stop = block_range + src = kv_cache[start:stop] + dst = staging_by_layer.view(kv_cache.dtype).view( + num_slots, + num_layers, + *src.shape[2:], + ) + dst.copy_(src) diff --git a/daser/connector/worker/store.py b/daser/connector/worker/store.py new file mode 100644 index 0000000..40e2ec7 --- /dev/null +++ b/daser/connector/worker/store.py @@ -0,0 +1,476 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +from dataclasses import dataclass, replace +import os +import threading +from typing import Any + +import cupy +import torch + +from daser.connector.helpers import base_req_id +from daser.connector.ipc_client import IPCClientAsync +from daser.connector.metadata import ReqStoreSpec, StoreWriteSpan +from daser.connector.worker.memory import ( + DEFAULT_STORE_STAGING_BYTES, + CudaStagingLease, + FixedCudaStagingPool, +) +from daser.connector.worker.staging import ( + CROSS_LAYER_KV_CACHE_KEY, + copy_cross_layer_kv_cache_to_staging, + copy_kv_cache_to_staging, + record_cuda_event, +) +from daser.logging import init_logger +from daser.transfer.cuda_ipc import ( + cuda_allocation_base_and_offset, + cuda_array_device_id, + cuda_array_pointer, + export_cuda_ipc_handle, +) + +logger = init_logger(__name__) + + +@dataclass +class _DeferredFinishedSave: + """Hold request store work until vLLM reports it finished.""" + + commit_keys: set[str] + reqs_to_store: dict[str, ReqStoreSpec] + finished: bool = False + future: Any | None = None + + +class StorePipeline: + """Own the complete worker store state machine. + + Args: + socket_path: DaseR server Unix socket path. + + Async/thread-safety: + Public methods are called on the vLLM worker thread. Snapshot, IPC, and + commit execute on the private store thread and its fixed CUDA stream. + """ + + def __init__(self, socket_path: str) -> None: + self._client = IPCClientAsync(socket_path) + self._loop = asyncio.new_event_loop() + self._thread = threading.Thread( + target=self._run_loop, + daemon=True, + name="daser-store-io", + ) + self._staging_bytes = 0 + self._staging_pool: FixedCudaStagingPool | None = None + self._pending_finished_saves: dict[str, _DeferredFinishedSave] = {} + self._kv_caches: dict[str, torch.Tensor] = {} + self._layer_names: list[str] = [] + self._layer_idx_map: dict[str, int] = {} + self._local_slot_size = 0 + self._rank_stride_bytes = 0 + self._tp_rank = 0 + self._tp_size = 1 + self._cuda_stream: torch.cuda.Stream | None = None + self._thread.start() + + def configure( + self, + *, + kv_caches: dict[str, torch.Tensor], + layer_names: list[str], + layer_idx_map: dict[str, int], + local_slot_size: int, + rank_stride_bytes: int, + tp_rank: int, + tp_size: int, + staging_bytes: int, + staging_pool: FixedCudaStagingPool, + ) -> None: + """Configure immutable KV layout and staging state. + + Args: + kv_caches: Registered vLLM KV tensors. + layer_names: Stable storage layer order. + layer_idx_map: Layer names mapped to storage indices. + local_slot_size: Bytes stored per slot by this TP rank. + rank_stride_bytes: Byte distance between rank lanes. + tp_rank: Current tensor-parallel rank. + tp_size: Tensor-parallel world size. + staging_bytes: Maximum bytes per store batch. + staging_pool: Fixed store staging buffers. + + Async/thread-safety: + Called once on the worker thread before request traffic. + """ + self._kv_caches = kv_caches + self._layer_names = list(layer_names) + self._layer_idx_map = dict(layer_idx_map) + self._local_slot_size = local_slot_size + self._rank_stride_bytes = rank_stride_bytes + self._tp_rank = tp_rank + self._tp_size = tp_size + self._staging_bytes = staging_bytes + self._staging_pool = staging_pool + + def initialize_transfer(self) -> None: + """Initialize the store IPC transfer client on its event loop. + + Async/thread-safety: + Called on the worker thread during startup and waits only for the + store loop's initialization future. + """ + self._submit(self._client.init_transfer()).result(timeout=120.0) + + def configure_rank_geometry( + self, + rank_stride_bytes: int, + tp_rank: int, + tp_size: int, + ) -> None: + """Apply server-finalized tensor-parallel lane geometry. + + Args: + rank_stride_bytes: Byte distance between server-owned rank lanes. + tp_rank: Current tensor-parallel rank. + tp_size: Tensor-parallel world size used for commit coordination. + + Async/thread-safety: + Called on the worker thread after runtime-config refresh and before + any store is submitted. + """ + self._rank_stride_bytes = rank_stride_bytes + self._tp_rank = tp_rank + self._tp_size = tp_size + + def queue_finished( + self, + reqs_to_store: dict[str, ReqStoreSpec], + commit_keys: set[str], + ) -> None: + """Queue stores until vLLM reports their requests finished. + + Args: + reqs_to_store: Store metadata for the current worker step. + commit_keys: Chunk keys eligible for commit after transfer. + + Async/thread-safety: + Called on the worker thread to accumulate immutable store intent. + CUDA ordering is captured later, when vLLM reports completion. + """ + for req_id, spec in reqs_to_store.items(): + base_id = base_req_id(req_id) + save = self._pending_finished_saves.get(base_id) + if save is None: + save = _DeferredFinishedSave(set(), {}) + self._pending_finished_saves[base_id] = save + save.reqs_to_store[req_id] = spec + if spec.chunk_key in commit_keys: + save.commit_keys.add(spec.chunk_key) + + def collect_finished(self, finished_req_ids: set[str]) -> set[str]: + """Submit newly finished stores and collect completed requests. + + Args: + finished_req_ids: Requests vLLM finished in this step. + + Returns: + Request IDs whose store and commit lifecycle has completed. + + Async/thread-safety: + Called on the worker thread. Store work runs on the private loop. + """ + finished: set[str] = set() + for req_id in finished_req_ids: + save = self._pending_finished_saves.get(req_id) + if save is not None: + save.finished = True + for req_id, save in list(self._pending_finished_saves.items()): + if not save.finished and save.future is None: + continue + if save.future is not None and save.future.done(): + try: + save.future.result(timeout=120.0) + finished.add(req_id) + finally: + del self._pending_finished_saves[req_id] + + capacity = self._staging_pool.depth if self._staging_pool is not None else 1 + inflight = sum( + save.future is not None for save in self._pending_finished_saves.values() + ) + for save in self._pending_finished_saves.values(): + if inflight >= capacity: + break + if save.finished and save.future is None: + self._submit_save(save) + inflight += 1 + return finished + + def shutdown(self) -> None: + """Finish queued stores, close IPC, and stop the store loop. + + Async/thread-safety: + Called once on the worker thread after request traffic stops. + """ + first_error: BaseException | None = None + try: + submitted_ids = [ + req_id + for req_id, save in self._pending_finished_saves.items() + if save.future is not None + ] + for req_id in submitted_ids: + save = self._pending_finished_saves[req_id] + try: + if save.future is not None: + save.future.result(timeout=120.0) + except BaseException as exc: # preserve cleanup during shutdown + if first_error is None: + first_error = exc + finally: + del self._pending_finished_saves[req_id] + for req_id in list(self._pending_finished_saves): + save = self._pending_finished_saves[req_id] + try: + self._submit_save(save) + assert save.future is not None + save.future.result(timeout=120.0) + except BaseException as exc: + if first_error is None: + first_error = exc + finally: + del self._pending_finished_saves[req_id] + try: + self._submit(self._client.close()).result(timeout=5.0) + except BaseException as exc: + if first_error is None: + first_error = exc + finally: + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(timeout=5.0) + if first_error is not None: + raise first_error + + def _submit(self, coro: Any) -> Any: + return asyncio.run_coroutine_threadsafe(coro, self._loop) + + def _run_loop(self) -> None: + asyncio.set_event_loop(self._loop) + self._loop.run_forever() + + def _submit_save(self, save: _DeferredFinishedSave) -> None: + """Capture producer ordering and submit one save to the store thread.""" + sample = next(iter(self._kv_caches.values()), None) + producer_event = record_cuda_event(sample) if sample is not None else None + save.future = self._submit(self._store_finished_save(save, producer_event)) + + def _plan_finished_save( + self, + save: _DeferredFinishedSave, + ) -> list[tuple[list[int], list[StoreWriteSpan]]]: + if not self._kv_caches or self._staging_pool is None: + return [] + reqs_to_store = { + req_id: replace( + spec, + file_offset=( + self._tp_rank * self._rank_stride_bytes + + spec.start_slot * self._local_slot_size + ), + ) + for req_id, spec in save.reqs_to_store.items() + } + return build_staging_store_batches( + reqs_to_store, + self._local_slot_size, + max_batch_bytes=self._staging_bytes, + ) + + async def _store_finished_save( + self, + save: _DeferredFinishedSave, + producer_event: torch.cuda.Event | None, + ) -> None: + stored_keys: list[str] = [] + for block_ids, spans in self._plan_finished_save(save): + staged = self._stage_batch(block_ids, spans, producer_event) + try: + stored_keys.extend(await self._write_cuda_buffer(staged)) + finally: + staged.lease.release() + requested = save.commit_keys + keys_to_commit = list( + dict.fromkeys(key for key in stored_keys if key in requested) + ) + await self._client.commit_chunks( + keys_to_commit, + tp_rank=self._tp_rank, + tp_size=self._tp_size, + ) + + def _stage_batch( + self, + block_ids: list[int], + spans: list[StoreWriteSpan], + producer_event: torch.cuda.Event | None, + ) -> "StagedStoreBatch": + if self._staging_pool is None: + raise RuntimeError("store staging pool is not configured") + sample = next(iter(self._kv_caches.values())) + nbytes = len(block_ids) * self._local_slot_size + lease = self._staging_pool.acquire(nbytes) + stream = self._cuda_stream + if sample.device.type == "cuda" and stream is None: + torch.cuda.set_device(sample.device) + stream = torch.cuda.Stream(device=sample.device) + self._cuda_stream = stream + if stream is not None: + if producer_event is not None: + stream.wait_event(producer_event) + with torch.cuda.stream(stream): + self._copy_blocks(lease.view, block_ids, sample) + stream.synchronize() + else: + self._copy_blocks(lease.view, block_ids, sample) + return StagedStoreBatch(lease.view, spans, lease) + + def _copy_blocks( + self, + staging: torch.Tensor, + block_ids: list[int], + sample: torch.Tensor, + ) -> None: + block_index = torch.tensor(block_ids, dtype=torch.long, device=sample.device) + cross_layer = self._kv_caches.get(CROSS_LAYER_KV_CACHE_KEY) + if cross_layer is not None: + copy_cross_layer_kv_cache_to_staging( + staging, + cross_layer, + block_ids, + len(self._layer_names), + self._local_slot_size, + block_index, + ) + return + for layer_name in self._layer_names: + copy_kv_cache_to_staging( + staging, + self._kv_caches[layer_name], + self._layer_idx_map[layer_name], + block_ids, + len(self._layer_names), + self._local_slot_size, + block_index, + ) + + async def _write_cuda_buffer(self, staged: "StagedStoreBatch") -> list[str]: + if staged.buffer.device.type == "cuda": + torch.cuda.set_device(staged.buffer.device) + cp_buffer = cupy.asarray(staged.buffer) + device_ptr = cuda_array_pointer(cp_buffer) + allocation_base, allocation_offset = cuda_allocation_base_and_offset(device_ptr) + return await self._client.transfer_store_cuda( + cuda_ipc_handle=export_cuda_ipc_handle(cp_buffer), + nbytes=staged.buffer.nbytes, + device_id=cuda_array_device_id(cp_buffer), + device_ptr=device_ptr, + allocation_base_ptr=allocation_base, + allocation_offset=allocation_offset, + producer_pid=os.getpid(), + spans=[ + { + "source_offset": span.source_offset, + "nbytes": span.nbytes, + "file_offset": span.file_offset, + "chunk_key": span.chunk_key, + "start_slot": span.start_slot, + "num_slots": span.num_slots, + } + for span in staged.spans + ], + ) + + +@dataclass(frozen=True) +class StagedStoreBatch: + """Hold a worker CUDA snapshot until its async store completes.""" + + buffer: torch.Tensor + spans: list[StoreWriteSpan] + lease: CudaStagingLease + + +def build_staging_store_batches( + reqs_to_store: dict[str, ReqStoreSpec], + slot_size: int, + max_batch_bytes: int = DEFAULT_STORE_STAGING_BYTES, +) -> list[tuple[list[int], list[StoreWriteSpan]]]: + """Split store requests into bounded slot-major staging batches. + + Args: + reqs_to_store: Request IDs mapped to store specifications. + slot_size: Bytes stored for one rank-local KV slot. + max_batch_bytes: Maximum GPU staging bytes in one batch. + + Returns: + Ordered block ID and server write-span batches. + + Async/thread-safety: + Pure CPU planning; safe to call from worker or store-loop threads. + """ + if slot_size <= 0: + raise ValueError("slot_size must be positive") + max_slots = max(1, max_batch_bytes // slot_size) + batches: list[tuple[list[int], list[StoreWriteSpan]]] = [] + batch_blocks: list[int] = [] + batch_spans: list[StoreWriteSpan] = [] + written_specs: set[tuple[str, int, int, int, int]] = set() + + def flush_batch() -> None: + nonlocal batch_blocks, batch_spans + if batch_blocks: + batches.append((batch_blocks, batch_spans)) + batch_blocks = [] + batch_spans = [] + + for spec in reqs_to_store.values(): + source_key = ( + spec.chunk_key, + spec.start_slot, + spec.num_slots, + spec.file_offset, + len(spec.block_ids), + ) + if source_key in written_specs: + continue + written_specs.add(source_key) + cursor = 0 + while cursor < len(spec.block_ids): + if len(batch_blocks) >= max_slots: + flush_batch() + available = max_slots - len(batch_blocks) + take = min(available, len(spec.block_ids) - cursor) + if take <= 0: + flush_batch() + continue + source_slot = len(batch_blocks) + part = spec.block_ids[cursor : cursor + take] + batch_blocks.extend(part) + batch_spans.append( + StoreWriteSpan( + source_offset=source_slot * slot_size, + nbytes=take * slot_size, + file_offset=spec.file_offset + cursor * slot_size, + chunk_key=spec.chunk_key, + start_slot=spec.start_slot, + num_slots=spec.num_slots, + ) + ) + cursor += take + flush_batch() + return batches diff --git a/daser/server/ipc/server.py b/daser/server/ipc/server.py index 3571e62..07a6f0a 100644 --- a/daser/server/ipc/server.py +++ b/daser/server/ipc/server.py @@ -622,7 +622,7 @@ def _record_transfer_metrics( "Transfer size per operation in bytes.", buckets=(65536, 262144, 1048576, 4194304, 16777216, 67108864, 268435456), ).observe(nbytes, labels=labels) - self._record_l1_metrics() + self._record_tier_metrics() throughput_gbps = (nbytes / elapsed_s / 1_000_000_000) if elapsed_s > 0 else 0.0 logger.debug( "[IPC] transfer_%s summary backend=%s status=%s bytes=%d " @@ -635,18 +635,21 @@ def _record_transfer_metrics( throughput_gbps, ) - def _record_l1_metrics(self) -> None: - """Publish L1 cache hit/miss counters and usage gauges.""" + def _record_tier_metrics(self) -> None: + """Publish L1 cache metrics and the cumulative L2 read counter.""" transfer = self._transfer if transfer is None: return stats = transfer.stats current_hits = stats.l1_hits current_misses = stats.l1_misses + current_l2_reads = stats.l2_reads prev_hits = getattr(self, "_prev_l1_hits", 0) prev_misses = getattr(self, "_prev_l1_misses", 0) + prev_l2_reads = getattr(self, "_prev_l2_reads", 0) delta_hits = current_hits - prev_hits delta_misses = current_misses - prev_misses + delta_l2_reads = current_l2_reads - prev_l2_reads if delta_hits > 0: self._metrics.counter("daser_l1_hits_total", "L1 memory cache hits.").inc( delta_hits @@ -655,8 +658,13 @@ def _record_l1_metrics(self) -> None: self._metrics.counter( "daser_l1_misses_total", "L1 memory cache misses." ).inc(delta_misses) + if delta_l2_reads > 0: + self._metrics.counter( + "daser_l2_reads_total", "Reads served from the L2 storage tier." + ).inc(delta_l2_reads) self._prev_l1_hits = current_hits self._prev_l1_misses = current_misses + self._prev_l2_reads = current_l2_reads l1_used = transfer.l1_bytes_used l1_capacity = int(self._runtime_config.get("l1_size_bytes", 0)) self._metrics.gauge("daser_l1_bytes_used", "L1 memory cache bytes in use.").set( diff --git a/daser/transfer/cuda_ipc.py b/daser/transfer/cuda_ipc.py index 26868b0..0c5e8dd 100644 --- a/daser/transfer/cuda_ipc.py +++ b/daser/transfer/cuda_ipc.py @@ -110,3 +110,30 @@ def cuda_array_device_id(array: Any) -> int: CUDA device ordinal. """ return int(array.device.id) + + +def cuda_allocation_base_and_offset(device_ptr: int) -> tuple[int, int]: + """Return the CUDA allocation base and byte offset for a tensor pointer. + + Args: + device_ptr: CUDA device pointer exported through IPC. + + Returns: + Tuple of allocation base pointer and byte offset. When the CUDA driver + query is unavailable, the pointer itself is used as the base. + + Async/thread-safety: + Read-only CUDA driver query safe during worker transfer preparation. + """ + try: + from cuda.bindings import driver as cuda_driver + + result, base_ptr, _allocation_size = cuda_driver.cuMemGetAddressRange( + device_ptr + ) + if result == cuda_driver.CUresult.CUDA_SUCCESS: + base = int(base_ptr) + return base, int(device_ptr) - base + except Exception: # noqa: BLE001 + pass + return int(device_ptr), 0 diff --git a/docs/design/architecture.md b/docs/design/architecture.md index fce7760..a1f10ac 100644 --- a/docs/design/architecture.md +++ b/docs/design/architecture.md @@ -61,20 +61,29 @@ graph TB subgraph vllm["vLLM 进程"] VAPI["OpenAI-compatible HTTP API"] DC["DaserConnector
KVConnectorBase_V1"] - SCHED["scheduler.py
SCHEDULER role"] - WORKER["worker.py
WORKER role"] + SCHED["scheduler/adapter.py
SCHEDULER role"] + WORKER["worker/adapter.py
WORKER role"] + LIFE["RequestLifecycle
pending state + IPC orchestration"] + RUNTIME["WorkerRuntime
KV/meta/completion ownership"] + LOAD["LoadPipeline
queue + load loop"] + STORE["StorePipeline
deferred save + store loop"] VAPI --> DC DC --> SCHED DC --> WORKER + SCHED --> LIFE + WORKER --> RUNTIME + RUNTIME --> LOAD + RUNTIME --> STORE end NVMe[("NVMe
daser.store / daser.index")] User -- "HTTP" --> HTTP HTTP -- "prefill / completion HTTP" --> VAPI - SCHED -- "lookup / alloc / runtime config" --> IPC - WORKER -- "CUDA IPC handle + transfer ops" --> IPC + LIFE -- "lookup / alloc / runtime config" --> IPC + LOAD -- "CUDA IPC handle + load ops" --> IPC + STORE -- "CUDA IPC handle + store / commit ops" --> IPC GDS -- "GDS IO" --> NVMe L1 -- "L2 daser.store IO" --> NVMe CM -- "save / load metadata" --> NVMe @@ -233,11 +242,22 @@ batch,导出 CUDA IPC handle,请求 server 读回 spans,再按层批量拷 KV cache。load 和 store 使用同一套 worker-side staging 抽象。 `wait_for_layer_load` 是 no-op,以兼容 vLLM FULL CUDA graph 模式。 +`WorkerConnectorMixin` 只适配 vLLM hooks;`WorkerRuntime` 集中 KV layout、step +metadata、completion 和 shutdown,内部的 load/store pipeline 各自拥有 queue、 +event loop、IPC client、staging lease 和 future。固定 staging pool 由两个 pipeline +共享的 worker memory module 管理;slot-major tensor layout 和 RoPE restore 留在 +staging module。 + +Scheduler 侧同样把 hook adapter 与状态 implementation 分开:request lifecycle +集中持有 pending load/store/alloc/async-save 状态并编排同步 IPC,chunk/prefix reuse +strategy 只生成 store intent,不反向修改 connector 私有状态。 + ### 后台 asyncio IO loop -vLLM worker 线程不直接运行可重入 event loop。WORKER role 在初始化时创建 -`daser-io` 后台线程,所有 transfer IPC 和 async IPC commit 都通过 -`run_coroutine_threadsafe` 提交。 +vLLM worker 线程不直接运行可重入 event loop。Load/store pipeline 分别拥有 +`daser-load-io` 和 `daser-store-io` 后台线程,transfer IPC 和 async IPC commit +通过 `run_coroutine_threadsafe` 提交。独立 loop 避免 cache-hit load 排在后台 +store coroutine 后面。 ### Transfer backend 启动后不可切换 diff --git a/docs/design/components.md b/docs/design/components.md index 00a9824..73ae4e1 100644 --- a/docs/design/components.md +++ b/docs/design/components.md @@ -5,9 +5,13 @@ | 组件 | 进程 | 职责 | |------|------|------| | `DaserConnector` | vLLM | vLLM `KVConnectorBase_V1` 入口;保留在 `daser/connector/daser_connector.py` 供 `kv_connector_module_path` 加载 | -| `SchedulerConnectorMixin` | vLLM scheduler | `daser/connector/scheduler.py`;负责 lookup、pending load/store 跟踪、slot 分配和 connector metadata 构造 | -| `WorkerConnectorMixin` | vLLM worker | `daser/connector/worker.py`;负责 KV cache 注册、CUDA IPC handle 导出、后台 IPC loop | -| `CudaStagingPool` | vLLM worker | `daser/connector/staging.py`;负责 GDS 和 iouring 共享的 bounded slot-major GPU staging 复用 | +| `SchedulerConnectorMixin` | vLLM scheduler | `daser/connector/scheduler/adapter.py`;适配 vLLM scheduler hooks,不拥有请求 lifecycle 状态 | +| `RequestLifecycle` | vLLM scheduler | 集中 lookup、pending load/store/alloc/async-save、slot 分配、preemption、completion 和 connector metadata 构造 | +| `WorkerConnectorMixin` | vLLM worker | `daser/connector/worker/adapter.py`;适配 vLLM worker hooks 和启动期 kernel warmup,不拥有 pipeline 状态 | +| `WorkerRuntime` | vLLM worker | 集中 KV layout、step metadata、load/store completion 和 shutdown,组合两个独立 pipeline | +| `LoadPipeline` / `StorePipeline` | vLLM worker | 分别拥有 load/store event loop、IPC client、staging lease、future、backpressure 和 transfer plan | +| worker memory module | vLLM worker | 负责 load/store 独立 fixed pool、indexed lease,以及约 6 GiB combined budget 的单点推导 | +| staging tensor module | vLLM worker | 负责 slot-major KV copy、CUDA producer synchronization 和 RoPE restore,不拥有 request/IPC 状态 | | `IPCClientSync` | vLLM scheduler | 阻塞式 Unix socket 客户端,用于 `get_runtime_config`、`lookup`、`alloc_chunk` | | `IPCClientAsync` | vLLM worker | asyncio Unix socket 客户端,用于 `transfer_store`、`transfer_load`、`commit_chunks` | | `TransferLayer` | DaseR | `daser/transfer/base.py`;server-owned KV 数据传输抽象,由 `IPCServer` 按 runtime config 初始化 | diff --git a/docs/design/flows.md b/docs/design/flows.md index aaf047d..343a046 100644 --- a/docs/design/flows.md +++ b/docs/design/flows.md @@ -71,28 +71,30 @@ chunk 已在上传时缓存,task suffix 通常是一次性的,因此推理 ```mermaid sequenceDiagram - participant S as vLLM Scheduler + participant S as Scheduler hook + participant SL as RequestLifecycle participant IPC as IPC server participant C as ServerCore participant CM as ChunkManager - S->>S: get_num_new_matched_tokens(request) - S->>S: full_aligned = floor(len(tokens) / block_tokens) * block_tokens - S->>S: store_key = xxh3_128(tokens[:full_aligned]) - S->>IPC: match_and_alloc(prefix, "", model_id) + S->>SL: get_num_new_matched_tokens(request) + SL->>SL: full_aligned = floor(len(tokens) / block_tokens) * block_tokens + SL->>SL: reuse strategy builds store intent + SL->>IPC: match_and_alloc(prefix, "", model_id) IPC->>C: match_and_alloc(...) C-->>IPC: chunks=[] / alloc=null - IPC-->>S: miss - S->>S: track PendingStore(chunk_key, token_count) + IPC-->>SL: miss + SL->>SL: track PendingStore(chunk_key, token_count) - S->>S: update_state_after_alloc(block_ids) - S->>IPC: alloc_chunk(chunk_key, token_count, model_id) + S->>SL: update_state_after_alloc(block_ids) + SL->>IPC: alloc_chunk(chunk_key, token_count, model_id) IPC->>C: alloc_chunk(...) C->>CM: allocate slots, evict old chunks if needed CM-->>C: start_slot, num_slots C-->>IPC: file_offset, pos_offset - IPC-->>S: allocation - S->>S: build_connector_meta(reqs_to_store) + IPC-->>SL: allocation + S->>SL: build_connector_meta(scheduler_output) + SL-->>S: reqs_to_store ``` `alloc_chunk` 只预留 metadata,chunk 还不会进入 `RetrievalIndex`。 @@ -101,22 +103,23 @@ sequenceDiagram ```mermaid sequenceDiagram - participant W as vLLM Worker - participant BG as daser-io loop + participant W as Worker hook + participant WR as WorkerRuntime + participant BG as StorePipeline / daser-store-io participant IPC as IPC server participant TL as TransferLayer - W->>W: bind_connector_metadata(reqs_to_store) + W->>WR: bind_connector_metadata(reqs_to_store) loop each attention layer - W->>W: save_kv_layer(layer_name, kv_layer) + W->>WR: save_kv_layer(layer_name, kv_layer) end - W->>W: wait_for_save() - W->>W: split reqs into bounded staging batches + W->>WR: wait_for_save() + WR->>BG: defer/submit reqs_to_store + BG->>BG: split reqs into bounded staging batches loop each staging batch - W->>W: lease bounded GPU staging view - W->>W: copy selected block KV into slot-major staging - W->>W: record producer CUDA event - W->>BG: run_coroutine_threadsafe(_write_cuda_buffer) + BG->>BG: lease bounded GPU staging view + BG->>BG: copy selected block KV into slot-major staging + BG->>BG: record producer CUDA event BG->>BG: wait producer event BG->>IPC: transfer_store(cuda_ipc_handle, spans) IPC->>TL: store_bytes_grouped(staging slices, file_offset) @@ -127,7 +130,7 @@ sequenceDiagram ``` `wait_for_save` 在当前 worker step 内完成 KV -> staging snapshot,然后把 -server transfer 交给后台 `daser-io` loop。未完成 batch 的 staging lease 由 +server transfer 交给后台 `daser-store-io` loop。未完成 batch 的 staging lease 由 future 持有,完成后归还 `CudaStagingPool`;`shutdown` 会阻塞等待所有 pending store future 完成。 @@ -154,22 +157,24 @@ commit 完成后 chunk 才能被 lookup 命中。 ```mermaid sequenceDiagram - participant S as vLLM Scheduler + participant S as Scheduler hook + participant SL as RequestLifecycle participant IPC as IPC server participant C as ServerCore participant RI as RetrievalIndex - S->>S: get_num_new_matched_tokens(request) - S->>IPC: match_and_alloc(prefix, "", model_id) + S->>SL: get_num_new_matched_tokens(request) + SL->>IPC: match_and_alloc(prefix, "", model_id) IPC->>C: lookup(prefix, model_id) C->>RI: lookup(tokens, model_id) RI-->>C: matched chunks C-->>IPC: chunks - IPC-->>S: chunks - S->>S: compute extra_tokens and pending_loads - S->>S: update_state_after_alloc(block_ids) - S->>S: map chunk target ranges to vLLM block ids - S->>S: build_connector_meta(reqs_to_load) + IPC-->>SL: chunks + SL->>SL: compute extra_tokens and pending_loads + S->>SL: update_state_after_alloc(block_ids) + SL->>SL: map chunk target ranges to vLLM block ids + S->>SL: build_connector_meta(scheduler_output) + SL-->>S: reqs_to_load ``` `PrefixHashIndex` 返回 rolling-prefix 的连续 slot 命中; @@ -180,17 +185,18 @@ vLLM 的 external tokens 是连续可用的前缀范围。 ```mermaid sequenceDiagram - participant W as vLLM Worker - participant BG as daser-io loop + participant W as Worker hook + participant WR as WorkerRuntime + participant BG as LoadPipeline / daser-load-io participant IPC as IPC server participant TL as TransferLayer - W->>W: start_load_kv(forward_context) - W->>W: split spans into bounded staging batches + W->>WR: start_load_kv(forward_context) + WR->>BG: submit reqs_to_load + BG->>BG: split spans into bounded staging batches loop each load batch - W->>W: lease GPU uint8 staging view - W->>W: export CUDA IPC handle for staging - W->>BG: transfer_load(cuda_ipc_handle, spans).result(timeout=120s) + BG->>BG: lease GPU uint8 staging view + BG->>BG: export CUDA IPC handle for staging BG->>IPC: transfer_load(cuda_ipc_handle, spans) IPC->>TL: load_bytes_grouped(staging slices, file_offset) TL-->>IPC: bytes read diff --git a/docs/optimizations/3_server_managed_transfer.md b/docs/optimizations/3_server_managed_transfer.md index 3f5acac..8d52a38 100644 --- a/docs/optimizations/3_server_managed_transfer.md +++ b/docs/optimizations/3_server_managed_transfer.md @@ -103,7 +103,7 @@ A same-process pointer path is used only for local unit-test and benchmark cases where producer and consumer PIDs match. Worker-side staging is shared by both GDS and iouring modes through -`daser.connector.staging`. `register_kv_caches` creates a bounded +`daser.connector.worker.staging`. `register_kv_caches` creates a bounded `CudaStagingPool` and preallocates one reusable buffer so the hot path does not pay a fresh CUDA allocation for the common batch size. Store batches keep their lease alive until the background transfer future completes; load batches release diff --git a/docs/optimizations/4_rope_apply_compile.md b/docs/optimizations/4_rope_apply_compile.md index 86bbe53..f2174e1 100644 --- a/docs/optimizations/4_rope_apply_compile.md +++ b/docs/optimizations/4_rope_apply_compile.md @@ -30,7 +30,7 @@ delta implementation there: TileLang operator entrypoint. - Naive eager RoPE is no longer part of production code; it is kept only inside the temporary micro benchmark and tests as a correctness oracle. -- `daser.connector.staging.apply_rope_delta_to_key_block()` remains the +- `daser.connector.worker.staging.apply_rope_delta_to_key_block()` remains the connector-facing helper and delegates into `daser.ops`. The production path is TileLang-only. TileLang import/compile/runtime failures diff --git a/tests/connector/test_connector_layout.py b/tests/connector/test_connector_layout.py index ef18510..d60dd69 100644 --- a/tests/connector/test_connector_layout.py +++ b/tests/connector/test_connector_layout.py @@ -23,11 +23,11 @@ def test_connector_entrypoint_delegates_scheduler_and_worker_methods() -> None: ) assert ( inspect.getmodule(DaserConnector.get_num_new_matched_tokens).__name__ - == "daser.connector.scheduler" + == "daser.connector.scheduler.adapter" ) assert ( inspect.getmodule(DaserConnector.start_load_kv).__name__ - == "daser.connector.worker" + == "daser.connector.worker.adapter" ) diff --git a/tests/connector/test_daser_connector.py b/tests/connector/test_daser_connector.py index ea18b19..d33ea9d 100644 --- a/tests/connector/test_daser_connector.py +++ b/tests/connector/test_daser_connector.py @@ -1,7 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # Standard -import asyncio from types import SimpleNamespace # Third Party @@ -25,57 +24,51 @@ ReqStoreSpec, StoreWriteSpan, ) -from daser.connector.reuse import PrefixReuseStrategy -from daser.connector.scheduler import ( - SchedulerConnectorMixin, +from daser.connector.scheduler.lifecycle import RequestLifecycle +from daser.connector.scheduler.planning import ( _block_ids_for_chunk, _contiguous_prefix_tokens, _trim_chunk_to_external_window, ) -from daser.connector.staging import ( - DEFAULT_PENDING_STORE_STAGING_BYTES, - DEFAULT_ROPE_DELTA_SCALE, - DEFAULT_STORE_STAGING_BYTES, - MIN_STORE_STAGING_BYTES, - StoreCudaStagingPool, -) -from daser.connector.staging import ( - apply_rope_delta_to_key_block as _apply_rope_delta_to_key_block, -) -from daser.connector.staging import ( +from daser.connector.scheduler.reuse import PrefixReuseStrategy +from daser.connector.worker.load import ( build_load_copy_runs as _build_load_copy_runs, ) -from daser.connector.staging import ( +from daser.connector.worker.load import ( build_load_read_batches as _build_load_read_batches, ) -from daser.connector.staging import ( +from daser.connector.worker.load import ( build_load_read_plan as _build_load_read_plan, ) -from daser.connector.staging import ( - build_staging_store_batches as _build_staging_store_batches, +from daser.connector.worker.memory import ( + DEFAULT_STAGING_BUDGET_BYTES, + DEFAULT_STORE_STAGING_BYTES, + MIN_STORE_STAGING_BYTES, + FixedCudaStagingPool, + derive_staging_layout, ) -from daser.connector.staging import ( - copy_staging_to_kv_cache as _copy_staging_to_kv_cache, +from daser.connector.worker.runtime import WorkerRuntime +from daser.connector.worker.staging import ( + DEFAULT_ROPE_DELTA_SCALE, ) -from daser.connector.staging import ( - derive_store_staging_limits as _derive_store_staging_limits, +from daser.connector.worker.staging import ( + apply_rope_delta_to_key_block as _apply_rope_delta_to_key_block, +) +from daser.connector.worker.staging import ( + copy_staging_to_kv_cache as _copy_staging_to_kv_cache, ) -from daser.connector.staging import ( +from daser.connector.worker.staging import ( record_cuda_event as _record_cuda_event, ) -from daser.connector.staging import ( +from daser.connector.worker.staging import ( synchronize_cuda_tensor as _synchronize_cuda_tensor, ) -from daser.connector.worker import ( - LoadRequestDispatcher, - WorkerConnectorMixin, - _DeferredFinishedSave, - _InflightRequestLoad, - _load_staging_pool_depth, - _PendingLoad, - _rank_lane_offset, - _RequestLoadFuture, - _SaveFuture, +from daser.connector.worker.store import ( + StagedStoreBatch, + StorePipeline, +) +from daser.connector.worker.store import ( + build_staging_store_batches as _build_staging_store_batches, ) BLOCK_TOKENS = 4 @@ -105,73 +98,7 @@ def test_rolling_prefix_keys_match_single_step_helper() -> None: ) -class _RuntimeConfigProbe(DaserConnector): - """Test connector exposing runtime config state through public properties.""" - - @property - def runtime_state(self): - return ( - self._store_path, - self._slot_size, - self._block_tokens, - self._model_id, - ) - - -class _WorkerProbe(DaserConnector): - """Worker-side probe with minimal state for transfer readiness tests.""" - - def __init__(self, store_path: str) -> None: - self._meta = DaserConnectorMeta( - reqs_to_load={ - "req": ReqLoadSpec( - chunk_key="hit", - start_slot=0, - num_slots=1, - block_ids=[0], - file_offset=0, - token_count=BLOCK_TOKENS, - ) - } - ) - self._transfer_ready = False - self._store_path = store_path - self._slot_size = 1024 - self._block_tokens = 4 - self._layer_names = [] - self._transfer_mode = "gds" - self._skip_l2 = False - self._pending_loads = {} - self._invalid_load_block_ids = set() - self._pending_finished_saves = {} - self._save_futures = [] - self._pending_save_staging_bytes = 0 - - def _refresh_runtime_config(self) -> None: - return - - def _ensure_load_request_dispatcher(self, sample_tensor: torch.Tensor) -> None: - """Drain queued requests synchronously for worker probe tests.""" - dispatcher = LoadRequestDispatcher(max_inflight=8, staging_depth=1) - queued = [] - while not self._load_request_queue.empty(): - item = self._load_request_queue.get() - if item is None: - return - queued.append(item) - active = dispatcher.submit_ready(self, queued, sample_tensor) - while active: - consumed = dispatcher.consume_ready(self, active, sample_tensor) - if not consumed: - return - active.extend(dispatcher.submit_ready(self, queued, sample_tensor)) - - @property - def transfer_ready(self): - return self._transfer_ready - - -class _SchedulerProbe(DaserConnector): +class _SchedulerProbe(RequestLifecycle): """Scheduler-side probe that can emulate deferred runtime config.""" def __init__(self, ipc_client) -> None: @@ -196,7 +123,7 @@ def mark_runtime_ready(self, model_id: str) -> None: self._model_id = model_id -class _AllocatingSchedulerProbe(SchedulerConnectorMixin): +class _AllocatingSchedulerProbe(RequestLifecycle): """Minimal scheduler probe that records allocation RPCs.""" def __init__(self) -> None: @@ -342,306 +269,6 @@ def alloc_chunk( return self._owner.alloc_chunk(chunk_key, token_count, model_id) -class _CommitProbe(WorkerConnectorMixin): - """Minimal worker probe exposing commit filtering behavior.""" - - def __init__(self) -> None: - self.committed: list[list[str]] = [] - self.commit_metadata: list[tuple[int, int]] = [] - self._ipc_async = self - self._ipc_store_async = self - - async def commit_chunks( - self, chunk_keys: list[str], tp_rank: int = 0, tp_size: int = 1 - ) -> None: - """Record chunk keys submitted to the async IPC client.""" - self.committed.append(list(chunk_keys)) - self.commit_metadata.append((tp_rank, tp_size)) - - async def commit_stored_keys( - self, stored_keys: list[str], commit_keys: list[str] - ) -> None: - """Expose worker commit filtering through a public test helper.""" - await self._commit_stored_keys(stored_keys, commit_keys) - - -class _QueueProbe(WorkerConnectorMixin): - """Minimal worker probe exposing independent load/store IPC clients.""" - - def __init__(self) -> None: - self.load_calls: list[str] = [] - self.store_calls: list[str] = [] - self._ipc_load_async = self - self._ipc_store_async = self - self._ipc_async = self - - async def transfer_load_cuda(self, **_kwargs) -> dict[str, int]: - """Record load client usage.""" - self.load_calls.append("load") - return {} - - async def transfer_store_cuda(self, **_kwargs) -> list[str]: - """Record store client usage.""" - self.store_calls.append("store") - return [] - - async def load_via_public_helper(self) -> None: - """Issue one load through the worker helper under test.""" - await self._transfer_load_cuda() - - async def store_via_public_helper(self) -> None: - """Issue one store through the worker helper under test.""" - await self._transfer_store_cuda() - - -class _FinishedSaveProbe(WorkerConnectorMixin): - """Worker probe for finished-request save scheduling.""" - - def __init__(self) -> None: - self._meta = None - self._pending_commits = set() - self._pending_finished_saves = {} - self._save_futures = [] - self._pending_save_staging_bytes = 0 - self._slot_size = 32 - self._store_staging_bytes = 128 - self._pending_store_staging_limit_bytes = 128 - self._layer_names = ["layer.0"] - self._kv_caches = {"layer.0": torch.empty(1)} - self.staged_batches: list[tuple[list[int], list[StoreWriteSpan]]] = [] - self.submitted = 0 - self.tracked: list[tuple[int, object | None]] = [] - self.committed_after: list[tuple[int, list[str]]] = [] - - def _clear_save_state(self) -> None: - return - - def _reap_save_futures(self, block: bool) -> None: - if block: - for future, _bytes, lease in self._save_futures: - future.result(timeout=120.0) - if lease is not None: - lease.release() - self._save_futures = [] - - def _stage_store_batch( - self, - block_ids: list[int], - spans: list[StoreWriteSpan], - ): - self.staged_batches.append((list(block_ids), list(spans))) - - class _Staged: - buffer = torch.empty(len(block_ids) * 32, dtype=torch.uint8) - ready_event = None - lease = object() - - def __init__(self, spans): - self.spans = spans - - return _Staged(spans) - - def _submit_store_coroutine(self, coro): - self.submitted += 1 - if getattr(getattr(coro, "cr_code", None), "co_name", "") == ( - "_commit_after_store_futures" - ): - asyncio.get_event_loop().run_until_complete(coro) - else: - coro.close() - - class _Future: - def __init__(self, value): - self._value = value - - def done(self) -> bool: - return True - - def result(self, timeout: float): - del timeout - return self._value - - return _Future(["stored"]) - - def _track_save_future( - self, - future, - staging_bytes: int, - staging_lease, - ) -> None: - self.tracked.append((staging_bytes, staging_lease)) - self._save_futures.append((future, staging_bytes, staging_lease)) - - async def _commit_after_store_futures(self, batch_futures, commit_keys): - self.committed_after.append((len(batch_futures), list(commit_keys))) - - def set_pending_meta( - self, - reqs_to_store: dict[str, ReqStoreSpec], - commit_keys: set[str], - ) -> None: - """Seed worker metadata and commit keys for save tests.""" - self._meta = DaserConnectorMeta(reqs_to_store=reqs_to_store) - self._pending_commits = commit_keys - - def seed_finished_save( - self, - req_id: str, - reqs_to_store: dict[str, ReqStoreSpec], - commit_keys: list[str], - ) -> None: - """Seed a deferred finished save for worker completion tests.""" - self._pending_finished_saves[req_id] = _DeferredFinishedSave( - commit_keys=set(commit_keys), - reqs_to_store=reqs_to_store, - ) - - def pending_finished_save_ids(self) -> set[str]: - """Return request IDs with deferred save work.""" - return set(self._pending_finished_saves) - - def pending_commit_keys(self) -> set[str]: - """Return pending worker commit keys.""" - return set(self._pending_commits) - - def set_submit_store_coroutine(self, submitter) -> None: - """Replace the store coroutine submitter for worker tests.""" - self._submit_store_coroutine = submitter - - def disable_store_staging(self) -> None: - """Make staging fail for store lifecycle regression tests.""" - self._stage_store_batch = lambda _block_ids, _spans: None - - -class _AsyncLoadFuture: - """Controllable future for worker async-load completion tests.""" - - def __init__(self, *, done: bool, error: BaseException | None = None) -> None: - self._done = done - self._error = error - self.result_calls = 0 - - def done(self) -> bool: - """Return whether this fake load has completed.""" - return self._done - - def result(self, timeout: float | None = None) -> None: - """Record result collection and optionally raise the seeded error.""" - del timeout - self.result_calls += 1 - if self._error is not None: - raise self._error - - def mark_done(self) -> None: - """Mark this fake future as complete.""" - self._done = True - - -class _AsyncLoadProbe(WorkerConnectorMixin): - """Worker probe with pending async load futures.""" - - def __init__(self) -> None: - self._role = KVConnectorRole.WORKER - self._pending_loads = {} - self._pending_finished_saves = {} - self._save_futures = [] - self._invalid_load_block_ids = set() - self.released: list[object] = [] - self.load_loop_stopped = False - self.store_loop_stopped = False - self.load_thread_joined = False - self.store_thread_joined = False - - class _Loop: - def __init__(self, callback) -> None: - self._callback = callback - - def call_soon_threadsafe(self, callback) -> None: - del callback - self._callback() - - def stop(self) -> None: - return - - class _Thread: - def __init__(self, callback) -> None: - self._callback = callback - - def join(self, timeout: float | None = None) -> None: - del timeout - self._callback() - - self._load_loop = _Loop(lambda: setattr(self, "load_loop_stopped", True)) - self._store_loop = _Loop(lambda: setattr(self, "store_loop_stopped", True)) - self._load_thread = _Thread(lambda: setattr(self, "load_thread_joined", True)) - self._store_thread = _Thread(lambda: setattr(self, "store_thread_joined", True)) - - def _reap_save_futures(self, block: bool) -> None: - del block - - def get_finished(self, finished_req_ids: set[str]): - """Expose worker completion polling without shutdown test doubles.""" - return WorkerConnectorMixin.get_finished(self, finished_req_ids) - - def _submit_load_coroutine(self, coro): - coro.close() - - class _Done: - def result(self, timeout: float | None = None) -> None: - del timeout - - return _Done() - - def _submit_store_coroutine(self, coro): - coro.close() - - class _Done: - def result(self, timeout: float | None = None) -> None: - del timeout - - return _Done() - - def seed_load( - self, - req_id: str, - future: _AsyncLoadFuture, - *, - block_ids: list[int], - lease: object | None = None, - ) -> None: - """Seed one pending async load.""" - self._pending_loads[req_id] = _PendingLoad( - future=future, - block_ids=block_ids, - lease=lease, - ) - - -class _LoopProbe(WorkerConnectorMixin): - """Minimal worker probe exposing background loop selection.""" - - def __init__(self) -> None: - self._load_loop = object() - self._store_loop = object() - - @property - def loop_pair(self) -> tuple[object, object]: - """Return the load and store loops configured on this probe.""" - return self._load_loop, self._store_loop - - def submit_load(self) -> None: - """Submit a load coroutine through the worker helper.""" - self._submit_load_coroutine(self._noop()) - - def submit_store(self) -> None: - """Submit a store coroutine through the worker helper.""" - self._submit_store_coroutine(self._noop()) - - async def _noop(self) -> None: - """Return immediately for loop-selection tests.""" - return - - def test_dataclasses_instantiate(): """DaserConnectorMeta, ReqLoadSpec, ReqStoreSpec all instantiate cleanly.""" spec_load = ReqLoadSpec("k", 0, 1, [0], 0, 16, 0, 0) @@ -655,55 +282,6 @@ def test_dataclasses_instantiate(): assert spec_load.pos_offset == 0 -def test_connector_allows_runtime_config_from_ipc(monkeypatch, tmp_path): - """Worker startup can begin with socket_path only and fill config by IPC.""" - - class DummyIPCClient: - def __init__(self, socket_path): - self.socket_path = socket_path - - def get_runtime_config(self): - return { - "store_path": str(tmp_path / "daser.store"), - "slot_size": 1024, - "block_tokens": 4, - "model_id": "served-model", - "cache_reuse_mode": "prefix", - } - - class DummyBase: - def __init__(self, vllm_config, role, kv_cache_config=None): - self._role = role - - class DummyConfig: - kv_connector_extra_config = {"socket_path": "/tmp/daser.sock"} - - class DummyVLLMConfig: - kv_transfer_config = DummyConfig() - model_config = None - - monkeypatch.setattr( - "daser.connector.daser_connector.IPCClientSync", - DummyIPCClient, - ) - monkeypatch.setattr( - "daser.connector.daser_connector.KVConnectorBase_V1.__init__", - DummyBase.__init__, - ) - - connector = _RuntimeConfigProbe( - DummyVLLMConfig(), - role=KVConnectorRole.SCHEDULER, - ) - - assert connector.runtime_state == ( - str(tmp_path / "daser.store"), - 1024, - 4, - "served-model", - ) - - def test_connector_requests_cross_layer_nhd_layout() -> None: """DaseR asks vLLM for block-major cross-layer KV cache layout.""" assert DaserConnector.get_required_kvcache_layout(object()) == "NHD" @@ -751,529 +329,195 @@ class DummyVLLMConfig: role=KVConnectorRole.SCHEDULER, ) - assert isinstance(connector._cache_reuse_strategy, PrefixReuseStrategy) # noqa: SLF001 - - -def test_start_load_kv_initializes_gds_after_server_creates_store( - monkeypatch, tmp_path -): - """Worker load path marks server transfer ready after deferred startup.""" - store_path = tmp_path / "daser.store" - store_path.write_bytes(b"\0" * 4096) - - connector = _WorkerProbe(str(store_path)) - - connector.start_load_kv(forward_context=object()) - - assert connector.transfer_ready is True - - -def test_start_load_kv_does_not_emit_info_timing( - monkeypatch, - caplog: pytest.LogCaptureFixture, -) -> None: - """Worker load timing stays out of the hot INFO logging path.""" - connector = _WorkerProbe("") - connector._transfer_ready = True # noqa: SLF001 - connector._transfer_mode = "iouring" # noqa: SLF001 - connector._skip_l2 = True # noqa: SLF001 - connector._slot_size = 4 # noqa: SLF001 - connector._store_staging_bytes = 64 # noqa: SLF001 - connector._kv_caches = { # noqa: SLF001 - "layer.0": torch.empty(2, 1, 1, 4, dtype=torch.uint8) - } - connector._layer_names = ["layer.0"] # noqa: SLF001 - connector._load_key_scale = 1.0 # noqa: SLF001 - connector._load_value_scale = 1.0 # noqa: SLF001 - connector._rope_delta_scale = 1.0 # noqa: SLF001 - connector._rope_base = 10000.0 # noqa: SLF001 - connector._rope_rotary_dim = 0 # noqa: SLF001 - connector._rope_is_neox_style = True # noqa: SLF001 - connector._meta = DaserConnectorMeta( # noqa: SLF001 - reqs_to_load={ - "req": ReqLoadSpec( - chunk_key="hit", - start_slot=0, - num_slots=1, - block_ids=[0], - file_offset=0, - token_count=4, - ) - } + assert isinstance( # noqa: SLF001 + connector._request_lifecycle._reuse_strategy(), # noqa: SLF001 + PrefixReuseStrategy, ) - class _Lease: - view = torch.empty(4, dtype=torch.uint8) - def release(self) -> None: - return - - class _Future: - def done(self) -> bool: - return True - - def result(self, timeout: float): - del timeout - return { - "transfer_open_ms": 0.0, - "transfer_load_ms": 0.0, - "transfer_sync_ms": 0.0, - "transfer_stats_delta": {"l1_hits": 1, "l1_misses": 0, "l2_reads": 0}, - } - - def fake_submit_load_coroutine(coro): - coro.close() - return _Future() - - monkeypatch.setattr(connector, "_acquire_staging", lambda *args: _Lease()) - monkeypatch.setattr(connector, "_submit_load_coroutine", fake_submit_load_coroutine) - monkeypatch.setattr("daser.connector.worker.cupy.asarray", lambda tensor: tensor) - monkeypatch.setattr("daser.connector.worker.export_cuda_ipc_handle", lambda _: b"0") - monkeypatch.setattr("daser.connector.worker.cuda_array_device_id", lambda _: 0) - monkeypatch.setattr("daser.connector.worker.cuda_array_pointer", lambda _: 1) - monkeypatch.setattr( - "daser.connector.worker._cuda_allocation_base_and_offset", - lambda _: (1, 0), - ) +def test_staging_layout_respects_available_cuda_headroom(monkeypatch) -> None: + """Combined staging pools stay within available CUDA headroom.""" monkeypatch.setattr( - "daser.connector.worker._copy_staging_to_kv_cache", - lambda **kwargs: 1, + torch.cuda, + "get_device_properties", + lambda device: SimpleNamespace(total_memory=80 << 30), ) monkeypatch.setattr( - "daser.connector.worker._synchronize_cuda_tensor", - lambda tensor: None, + torch.cuda, + "mem_get_info", + lambda device=None: ((4 << 30), 80 << 30), ) - with caplog.at_level("INFO", logger="daser.connector.worker"): - connector.start_load_kv(forward_context=object()) - - assert "start_load_kv timing" not in caplog.text - - -def test_worker_load_batches_use_buffer_scoped_ipc_clients() -> None: - """Parallel load batches should not serialize on one async IPC client lock.""" - - class _Client: - def __init__(self, name: str) -> None: - self.name = name - self.calls: list[dict[str, object]] = [] - - async def transfer_load_registered_cuda(self, **kwargs): - self.calls.append(kwargs) - return {"client": self.name} - - class Probe(WorkerConnectorMixin): - def __init__(self) -> None: - self._ipc_load_async = _Client("fallback") - self._ipc_load_async_pool = [_Client("load0"), _Client("load1")] - - probe = Probe() - - coro0 = probe._transfer_load_registered_cuda( # noqa: SLF001 - buffer_index=0, - nbytes=4, - spans=[], - ) - coro1 = probe._transfer_load_registered_cuda( # noqa: SLF001 - buffer_index=1, - nbytes=4, - spans=[], + buffer_bytes, load_depth, store_depth, allocated = derive_staging_layout( + torch.device("cuda"), + local_slot_size=64 << 20, + max_load_inflight=8, + reserve_bytes=1 << 30, ) - assert asyncio.run(coro0) == {"client": "load0"} - assert asyncio.run(coro1) == {"client": "load1"} - assert probe._ipc_load_async.calls == [] # noqa: SLF001 - - -def test_get_finished_releases_each_request_load_independently() -> None: - """Completed request loads should not wait for unrelated pending loads.""" - connector = _AsyncLoadProbe() - done = _AsyncLoadFuture(done=True) - pending = _AsyncLoadFuture(done=False) - connector.seed_load("req-a", done, block_ids=[1]) - connector.seed_load("req-b", pending, block_ids=[2]) - - finished_sending, finished_recving = connector.get_finished(set()) + assert (buffer_bytes, load_depth, store_depth) == ((4 << 30) // 10, 5, 2) + assert allocated == buffer_bytes * 7 + assert allocated <= (4 << 30) - (1 << 30) - assert finished_sending is None - assert finished_recving == {"req-a"} - assert done.result_calls == 1 - assert pending.result_calls == 0 - assert set(connector._pending_loads) == {"req-b"} # noqa: SLF001 +def test_worker_transfer_ready_allows_skip_l2_without_store_path() -> None: + """L1-only mode has no store path but still has a valid transfer config.""" -def test_start_load_kv_releases_waiting_request_when_load_cannot_start() -> None: - """Worker should not leave async-hit requests waiting forever.""" - connector = _WorkerProbe("") - connector._transfer_ready = True # noqa: SLF001 + class Pipeline: + def __init__(self) -> None: + self.initialized = False + + def configure_rank_geometry(self, *args: int) -> None: + del args + + def initialize_transfer(self) -> None: + self.initialized = True + + connector = WorkerRuntime.__new__(WorkerRuntime) + connector._transfer_ready = False # noqa: SLF001 + connector._pipelines_initialized = False # noqa: SLF001 + connector._store_path = "" # noqa: SLF001 + connector._slot_size = 1024 # noqa: SLF001 + connector._local_slot_size = 1024 # noqa: SLF001 + connector._tp_size = 1 # noqa: SLF001 + connector._server_tp_size = 1 # noqa: SLF001 + connector._tp_rank = 0 # noqa: SLF001 + connector._rank_stride_bytes = 0 # noqa: SLF001 connector._transfer_mode = "iouring" # noqa: SLF001 connector._skip_l2 = True # noqa: SLF001 - connector._layer_names = [] # noqa: SLF001 - connector._meta = DaserConnectorMeta( # noqa: SLF001 - reqs_to_load={ - "req": ReqLoadSpec( - chunk_key="hit", - start_slot=0, - num_slots=1, - block_ids=[7], - file_offset=0, - token_count=4, - ) - } - ) + connector._refresh_runtime_config = lambda: None # noqa: SLF001 + connector._load_pipeline = Pipeline() # noqa: SLF001 + connector._store_pipeline = Pipeline() # noqa: SLF001 - connector.start_load_kv(forward_context=object()) - - assert connector.get_finished(set()) == (None, {"req"}) - assert connector.get_block_ids_with_load_errors() == {7} - - -def test_wait_for_save_defers_store_until_request_finished() -> None: - """Cold stores should be snapshotted after vLLM reports request finish.""" - connector = _FinishedSaveProbe() - connector.set_pending_meta( - { - "req": ReqStoreSpec( - chunk_key="stored", - start_slot=0, - num_slots=2, - block_ids=[4, 5], - file_offset=0, - token_count=8, - ) - }, - {"stored"}, - ) - - connector.wait_for_save() - - assert connector.staged_batches == [] - assert connector.submitted == 0 - assert connector.pending_finished_save_ids() == {"req"} - assert connector.pending_commit_keys() == set() - - finished_sending, finished_recving = connector.get_finished({"req"}) - - assert finished_recving is None - assert finished_sending == {"req"} - assert [block_ids for block_ids, _spans in connector.staged_batches] == [[4, 5]] - assert connector.submitted == 2 - assert connector.committed_after == [(1, ["stored"])] - - -def test_wait_for_save_groups_prefix_slot_stores_by_base_request() -> None: - """Synthetic prefix slot stores should finish with their base request.""" - connector = _FinishedSaveProbe() - connector.set_pending_meta( - { - "req:store:0": ReqStoreSpec( - chunk_key="stored-0", - start_slot=0, - num_slots=1, - block_ids=[4], - file_offset=0, - token_count=4, - ), - "req:store:1": ReqStoreSpec( - chunk_key="stored-1", - start_slot=1, - num_slots=1, - block_ids=[5], - file_offset=32, - token_count=4, - ), - }, - {"stored-0", "stored-1"}, - ) - - connector.wait_for_save() - - assert connector.pending_finished_save_ids() == {"req"} - - finished_sending, finished_recving = connector.get_finished({"req"}) - - assert finished_recving is None - assert finished_sending == {"req"} - assert [block_ids for block_ids, _spans in connector.staged_batches] == [[4, 5]] - assert connector.committed_after == [(1, ["stored-0", "stored-1"])] - - -def test_get_finished_holds_blocks_until_deferred_store_completes() -> None: - """Worker should not release finished request blocks before store is done.""" - connector = _FinishedSaveProbe() - connector.seed_finished_save( - "req", - { - "req": ReqStoreSpec( - chunk_key="stored", - start_slot=0, - num_slots=1, - block_ids=[4], - file_offset=0, - token_count=4, - ) - }, - ["stored"], - ) - - class _PendingFuture: - def done(self) -> bool: - return False - - def submit_pending(coro): - coro.close() - return _PendingFuture() - - connector.set_submit_store_coroutine(submit_pending) - - finished_sending, finished_recving = connector.get_finished({"req"}) - - assert finished_recving is None - assert finished_sending is None - assert connector.staged_batches - assert "req" in connector.pending_finished_save_ids() + assert connector._ensure_transfer_ready() is True # noqa: SLF001 + assert connector._transfer_ready is True # noqa: SLF001 + assert connector._pipelines_initialized is True # noqa: SLF001 + assert connector._load_pipeline.initialized is True # noqa: SLF001 + assert connector._store_pipeline.initialized is True # noqa: SLF001 -def test_get_finished_reports_completed_deferred_store_on_later_step() -> None: - """Completed saves should be reported even after the original finish step.""" - connector = _FinishedSaveProbe() - pending_future = None +def test_worker_transfer_ready_propagates_refreshed_tp_geometry() -> None: + """Delayed TP geometry must reach pipelines before transfer initialization.""" - class _PendingFuture: + class Pipeline: def __init__(self) -> None: - self.complete = False - - def done(self) -> bool: - return self.complete - - def result(self, timeout: float): - del timeout - return None - - def submit_pending(coro): - nonlocal pending_future - coro.close() - pending_future = _PendingFuture() - return pending_future - - connector.seed_finished_save( - "req", - { - "req": ReqStoreSpec( - chunk_key="stored", - start_slot=0, - num_slots=1, - block_ids=[4], - file_offset=0, - token_count=4, - ) - }, - ["stored"], - ) - connector.set_submit_store_coroutine(submit_pending) - - assert connector.get_finished({"req"}) == (None, None) - assert pending_future is not None - pending_future.complete = True - - assert connector.get_finished(set()) == ({"req"}, None) - - -def test_get_finished_releases_request_when_no_store_batch_can_be_staged() -> None: - """A skipped staging batch should not keep scheduler blocks forever.""" - connector = _FinishedSaveProbe() - connector.seed_finished_save( - "req", - { - "req": ReqStoreSpec( - chunk_key="stored", - start_slot=0, - num_slots=1, - block_ids=[4], - file_offset=0, - token_count=4, - ) - }, - ["stored"], - ) - connector.disable_store_staging() - - assert connector.get_finished({"req"}) == ({"req"}, None) - assert connector.pending_finished_save_ids() == set() - assert connector.submitted == 0 - - -def test_get_finished_does_not_block_on_pending_async_load() -> None: - """Worker polling should not block on incomplete async load futures.""" - connector = _AsyncLoadProbe() - future = _AsyncLoadFuture(done=False) - connector.seed_load("req", future, block_ids=[4, 5]) - - finished_sending, finished_recving = connector.get_finished(set()) - - assert finished_sending is None - assert finished_recving is None - assert future.result_calls == 0 - assert "req" in connector._pending_loads # noqa: SLF001 + self.geometry: tuple[int, ...] | None = None + self.initialized = False + + def configure_rank_geometry(self, *args: int) -> None: + self.geometry = args + + def initialize_transfer(self) -> None: + self.initialized = True + + connector = WorkerRuntime.__new__(WorkerRuntime) + connector._transfer_ready = False # noqa: SLF001 + connector._pipelines_initialized = False # noqa: SLF001 + connector._store_path = "" # noqa: SLF001 + connector._slot_size = 0 # noqa: SLF001 + connector._local_slot_size = 512 # noqa: SLF001 + connector._tp_size = 2 # noqa: SLF001 + connector._server_tp_size = 1 # noqa: SLF001 + connector._tp_rank = 1 # noqa: SLF001 + connector._rank_stride_bytes = 0 # noqa: SLF001 + connector._transfer_mode = "iouring" # noqa: SLF001 + connector._skip_l2 = True # noqa: SLF001 + def refresh() -> None: + connector._slot_size = 1024 # noqa: SLF001 + connector._server_tp_size = 2 # noqa: SLF001 + connector._rank_stride_bytes = 8192 # noqa: SLF001 -def test_get_finished_reports_completed_async_load() -> None: - """Completed async loads should release the waiting request.""" - connector = _AsyncLoadProbe() - future = _AsyncLoadFuture(done=True) - connector.seed_load("req", future, block_ids=[4, 5]) + connector._refresh_runtime_config = refresh # noqa: SLF001 + connector._load_pipeline = Pipeline() # noqa: SLF001 + connector._store_pipeline = Pipeline() # noqa: SLF001 - finished_sending, finished_recving = connector.get_finished(set()) + assert connector._ensure_transfer_ready() is True # noqa: SLF001 + assert connector._load_pipeline.geometry == (8192, 1) # noqa: SLF001 + assert connector._store_pipeline.geometry == (8192, 1, 2) # noqa: SLF001 + assert connector._load_pipeline.initialized is True # noqa: SLF001 + assert connector._store_pipeline.initialized is True # noqa: SLF001 - assert finished_sending is None - assert finished_recving == {"req"} - assert future.result_calls == 1 - assert connector._pending_loads == {} # noqa: SLF001 - assert connector._invalid_load_block_ids == set() # noqa: SLF001 +def test_worker_runtime_refreshes_l1_only_transfer_config(monkeypatch) -> None: + """Deferred worker config refresh must propagate L1-only transfer settings.""" -def test_load_request_dispatcher_limits_active_requests_by_staging_depth() -> None: - """Request dispatcher should cap active loads by max inflight and buffers.""" - dispatcher = LoadRequestDispatcher(max_inflight=8, staging_depth=3) + class DummyIPCClient: + def __init__(self, socket_path: str) -> None: + self.socket_path = socket_path - assert dispatcher.effective_inflight == 3 + def get_runtime_config(self) -> dict[str, object]: + return { + "store_path": "", + "slot_size": 1024, + "tensor_parallel_size": 2, + "rank_stride_bytes": 512, + "transfer_mode": "iouring", + "skip_l2": True, + } + def close(self) -> None: + return -def test_load_staging_pool_depth_respects_available_cuda_memory(monkeypatch) -> None: - """Load staging depth should not preallocate past available CUDA headroom.""" monkeypatch.setattr( - torch.cuda, - "mem_get_info", - lambda device=None: ((4 << 30), 80 << 30), - ) - - depth = _load_staging_pool_depth( - buffer_bytes=1536 << 20, - pending_limit_bytes=12 << 30, - device=torch.device("cuda"), + "daser.connector.worker.runtime.IPCClientSync", + DummyIPCClient, ) + connector = WorkerRuntime.__new__(WorkerRuntime) + connector._socket_path = "/unused/daser.sock" # noqa: SLF001 + connector._store_path = "" # noqa: SLF001 + connector._slot_size = 0 # noqa: SLF001 + connector._server_tp_size = 1 # noqa: SLF001 + connector._rank_stride_bytes = 0 # noqa: SLF001 + connector._transfer_mode = "gds" # noqa: SLF001 + connector._skip_l2 = False # noqa: SLF001 + + connector._refresh_runtime_config() # noqa: SLF001 + + assert connector._slot_size == 1024 # noqa: SLF001 + assert connector._server_tp_size == 2 # noqa: SLF001 + assert connector._rank_stride_bytes == 512 # noqa: SLF001 + assert connector._transfer_mode == "iouring" # noqa: SLF001 + assert connector._skip_l2 is True # noqa: SLF001 - assert depth == 2 - - -def test_load_request_dispatcher_submits_without_waiting_for_batch() -> None: - """Queued requests should be submitted immediately while slots are free.""" - - class Probe(_AsyncLoadProbe): - def __init__(self) -> None: - super().__init__() - self.submitted: list[str] = [] - self.completed: list[str] = [] - - def _submit_request_load_for_dispatcher( - self, item, buffer_index, sample_tensor - ): - del sample_tensor - self.submitted.append(item.req_id) - future = _AsyncLoadFuture(done=False) - active = SimpleNamespace(future=future) - return _InflightRequestLoad( - item=item, - buffer_index=buffer_index, - batches=[], - next_batch=0, - remaining_batches=1, - active=active, - completed=[], - ) - - def _consume_dispatcher_load(self, state, sample_tensor): - del sample_tensor - self.completed.append(state.item.req_id) - state.item.future.set_result() - return 0, True - - probe = Probe() - dispatcher = LoadRequestDispatcher(max_inflight=8, staging_depth=2) - sample = torch.empty(1) - first_future = _RequestLoadFuture() - second_future = _RequestLoadFuture() - third_future = _RequestLoadFuture() - queued = [ - SimpleNamespace( - req_id="req-a", - spec_id="req-a", - spec=ReqLoadSpec("hit-a", 0, 1, [1], 0, 4), - future=first_future, - ), - SimpleNamespace( - req_id="req-b", - spec_id="req-b", - spec=ReqLoadSpec("hit-b", 1, 1, [2], 4, 4), - future=second_future, - ), - SimpleNamespace( - req_id="req-c", - spec_id="req-c", - spec=ReqLoadSpec("hit-c", 2, 1, [3], 8, 4), - future=third_future, - ), - ] - - active = dispatcher.submit_ready(probe, queued, sample) - - assert probe.submitted == ["req-a", "req-b"] - assert [item.req_id for item in queued] == ["req-c"] - assert len(active) == 2 - - active_batch = active[0].active - assert active_batch is not None - active_batch.future.mark_done() - dispatcher.consume_ready(probe, active, sample) - dispatcher.submit_ready(probe, queued, sample) - - assert probe.completed == ["req-a"] - assert probe.submitted == ["req-a", "req-b", "req-c"] - assert first_future.done() - - -def test_get_finished_releases_request_after_async_load_failure() -> None: - """Failed async loads should report completion and mark invalid blocks.""" - connector = _AsyncLoadProbe() - future = _AsyncLoadFuture(done=True, error=RuntimeError("load failed")) - connector.seed_load("req", future, block_ids=[4, 5]) - - finished_sending, finished_recving = connector.get_finished(set()) - - assert finished_sending is None - assert finished_recving == {"req"} - assert connector._pending_loads == {} # noqa: SLF001 - assert connector._invalid_load_block_ids == {4, 5} # noqa: SLF001 - - -def test_shutdown_collects_failed_async_load_without_raising() -> None: - """Shutdown should still clean worker resources after load failures.""" - connector = _AsyncLoadProbe() - future = _AsyncLoadFuture(done=True, error=RuntimeError("load failed")) - connector.seed_load("req", future, block_ids=[4, 5]) - - connector.shutdown() - - assert connector._pending_loads == {} # noqa: SLF001 - assert connector._invalid_load_block_ids == {4, 5} # noqa: SLF001 - assert connector.load_loop_stopped - assert connector.store_loop_stopped - assert connector.load_thread_joined - assert connector.store_thread_joined +def test_request_lifecycle_rebuilds_prefix_keys_after_block_size_refresh() -> None: + """Deferred geometry refresh must rebuild same-mode rolling-prefix keys.""" -def test_worker_transfer_ready_allows_skip_l2_without_store_path() -> None: - """L1-only mode has no store path but still has a valid transfer config.""" - connector = _WorkerProbe("") - connector._transfer_mode = "iouring" # noqa: SLF001 - connector._skip_l2 = True # noqa: SLF001 + class DummyIPCClient: + def get_runtime_config(self) -> dict[str, object]: + return { + "slot_size": 1024, + "block_tokens": 128, + "model_id": "served-model", + "cache_reuse_mode": "prefix", + } - assert connector._ensure_transfer_ready() is True # noqa: SLF001 - assert connector.transfer_ready is True + lifecycle = RequestLifecycle( + ipc_client=DummyIPCClient(), + block_tokens=16, + slot_size=0, + model_id="default", + cache_reuse_mode="prefix", + runtime_config_ready=False, + ) + lifecycle._refresh_runtime_config() # noqa: SLF001 + tokens = list(range(256)) + pending = lifecycle._reuse_strategy().prepare_store(tokens, 256) # noqa: SLF001 + + assert pending is not None + pending.block_ids = [7, 8] + plan = lifecycle._reuse_strategy().plan_store( # noqa: SLF001 + "req", + pending, + tokens, + set(), + ) + assert [intent.chunk_key for intent in plan.intents] == rolling_keys(tokens, 128) -def test_runtime_config_ready_allows_skip_l2_without_store_path(monkeypatch): - """Scheduler runtime config should be ready when skip_l2 omits store_path.""" +def test_scheduler_runtime_config_is_owned_by_request_lifecycle(monkeypatch): + """Scheduler readiness is refreshed directly on its lifecycle owner.""" class DummyIPCClient: def __init__(self, socket_path): @@ -1315,8 +559,9 @@ class DummyVLLMConfig: role=KVConnectorRole.SCHEDULER, ) - assert connector._runtime_config_ready is True # noqa: SLF001 - assert connector._skip_l2 is True # noqa: SLF001 + lifecycle = connector._request_lifecycle # noqa: SLF001 + assert lifecycle._runtime_config_ready is True # noqa: SLF001 + assert lifecycle._model_id == "served-model" # noqa: SLF001 def test_scheduler_refreshes_runtime_config_before_lookup(monkeypatch): @@ -1542,7 +787,7 @@ def test_trim_chunk_to_external_window_skips_local_prefix_slots(): def test_update_state_after_alloc_single_hit_uses_external_window(): """Single-prefix hit maps the external suffix onto absolute request blocks.""" - class MockConnector(SchedulerConnectorMixin): + class MockConnector(RequestLifecycle): def __init__(self) -> None: self._block_tokens = BLOCK_TOKENS self._slot_size = 32 @@ -1575,7 +820,7 @@ class MockBlocks: connector = MockConnector() - DaserConnector.update_state_after_alloc( + RequestLifecycle.update_state_after_alloc( connector, MockRequest(), MockBlocks(), @@ -1623,7 +868,7 @@ def lookup( } ] - class MockConnector(SchedulerConnectorMixin): + class MockConnector(RequestLifecycle): def __init__(self) -> None: self._runtime_config_ready = True self._block_tokens = BLOCK_TOKENS @@ -1664,7 +909,7 @@ class MockRequest: def test_update_state_after_alloc_multi_hit_trims_each_chunk_to_external_window(): """Multi-chunk hits map onto absolute request block positions.""" - class MockConnector(SchedulerConnectorMixin): + class MockConnector(RequestLifecycle): def __init__(self) -> None: self._block_tokens = BLOCK_TOKENS self._slot_size = 32 @@ -1708,7 +953,7 @@ class MockBlocks: connector = MockConnector() - DaserConnector.update_state_after_alloc( + RequestLifecycle.update_state_after_alloc( connector, MockRequest(), MockBlocks(), @@ -2175,28 +1420,27 @@ def test_restore_cross_layer_kv_cache_table_tilelang_matches_reference_cuda(): def test_register_kv_caches_warms_dynamic_rope_apply_once(monkeypatch): - from daser.connector import worker + from daser.connector.worker import runtime as worker - class Probe(WorkerConnectorMixin): + class Probe(WorkerRuntime): def __init__(self) -> None: self._slot_size = 0 - self._store_staging_bytes = 0 - self._pending_store_staging_limit_bytes = 0 + self._tp_size = 1 + self._server_tp_size = 1 + self._tp_rank = 0 + self._rank_stride_bytes = 0 self._rope_rotary_dim = 8 self._rope_base = 10000.0 self._rope_is_neox_style = True - self._ipc_async = None - self._bg_loop = None def _ensure_transfer_ready(self) -> bool: return False + def _configure_pipelines(self, sample: torch.Tensor) -> int: + del sample + return 1 + calls = [] - monkeypatch.setattr( - worker, - "_derive_store_staging_limits", - lambda device: (4096, 8192), - ) monkeypatch.setattr( worker, "_warm_rope_apply_backends", @@ -2222,7 +1466,7 @@ def _ensure_transfer_ready(self) -> bool: def test_register_cross_layers_kv_cache_preserves_layer_order(monkeypatch): """Worker registration keeps vLLM layer names for slot-major staging.""" - from daser.connector import worker + from daser.connector.worker import runtime as worker class Group: layer_names = ["layer.0", "layer.1"] @@ -2230,12 +1474,14 @@ class Group: class Config: kv_cache_groups = [Group()] - class Probe(WorkerConnectorMixin): + class Probe(WorkerRuntime): def __init__(self) -> None: self._kv_cache_config = Config() self._slot_size = 0 - self._store_staging_bytes = 0 - self._pending_store_staging_limit_bytes = 0 + self._tp_size = 1 + self._server_tp_size = 1 + self._tp_rank = 0 + self._rank_stride_bytes = 0 self._rope_rotary_dim = 8 self._rope_base = 10000.0 self._rope_is_neox_style = True @@ -2244,6 +1490,13 @@ def __init__(self) -> None: def _ensure_transfer_ready(self) -> bool: return True + def _init_server_transfer(self) -> None: + return + + def _configure_pipelines(self, sample: torch.Tensor) -> int: + del sample + return 1 + @property def registration_state(self): return ( @@ -2253,11 +1506,6 @@ def registration_state(self): self._slot_size, ) - monkeypatch.setattr( - worker, - "_derive_store_staging_limits", - lambda device: (4096, 8192), - ) monkeypatch.setattr(worker, "_warm_rope_apply_backends", lambda **kwargs: None) kv_cache = torch.zeros((8, 2, 2, 4, 2, 8), dtype=torch.float32) @@ -2271,172 +1519,8 @@ def registration_state(self): assert slot_size == kv_cache[0].nbytes -def test_stage_store_batch_does_not_warm_dynamic_rope_again(monkeypatch): - from daser.connector import worker - - class Probe(WorkerConnectorMixin): - def __init__(self) -> None: - self.kv_cache = torch.zeros((2, 8, 4, 2, 8), dtype=torch.float32) - self._layer_names = ["layer.0"] - self._layer_idx_map = {"layer.0": 0} - self._kv_caches = {"layer.0": self.kv_cache} - self.slot_size = self.kv_cache[:, 0].nbytes - self._slot_size = self.slot_size - self._store_staging_bytes = 4096 - self._pending_store_staging_limit_bytes = 8192 - self._store_staging_pool = None - self._pending_save_staging_bytes = 0 - self._save_futures = [] - self._rope_rotary_dim = 8 - self._rope_base = 10000.0 - self._rope_is_neox_style = True - - def stage_store_batch(self, block_ids: list[int], spans: list[StoreWriteSpan]): - """Expose store staging through a public test helper.""" - return self._stage_store_batch(block_ids, spans) - - calls = [] - monkeypatch.setattr( - worker, - "_warm_rope_apply_backends", - lambda **kwargs: calls.append(kwargs), - ) - - probe = Probe() - staged = probe.stage_store_batch( - block_ids=[0, 1, 2], - spans=[StoreWriteSpan(0, probe.slot_size * 3, 0, "k0", 0, 3)], - ) - - assert staged is not None - assert calls == [] - - -def test_store_staging_wait_skips_future_without_lease() -> None: - """Store backpressure waits until a future returns a staging lease.""" - - class _Future: - def __init__(self, name: str) -> None: - self.name = name - - def done(self) -> bool: - return False - - def result(self, timeout: float) -> None: - assert timeout > 0 - completed.append(self.name) - - class Probe(WorkerConnectorMixin): - def __init__(self) -> None: - self._store_staging_pool = StoreCudaStagingPool( - device=torch.device("cpu"), - buffer_bytes=16, - depth=1, - ) - lease = self._store_staging_pool.acquire(8) - self._pending_store_staging_limit_bytes = 64 - self._pending_save_staging_bytes = 8 - self._save_futures = [ - _SaveFuture(_Future("commit"), 0, None), - _SaveFuture(_Future("store"), 8, lease), - ] - - def wait_for_staging(self) -> None: - self._wait_for_store_staging_release(8) - - completed: list[str] = [] - probe = Probe() - probe.wait_for_staging() - - assert completed == ["commit", "store"] - assert probe._store_staging_pool.available == 1 # noqa: SLF001 - - -def test_stage_store_batch_records_ready_event_without_synchronizing(monkeypatch): - from daser.connector import worker - - class Probe(WorkerConnectorMixin): - def __init__(self) -> None: - self.kv_cache = torch.zeros((2, 8, 4, 2, 8), dtype=torch.float32) - self._layer_names = ["layer.0"] - self._layer_idx_map = {"layer.0": 0} - self._kv_caches = {"layer.0": self.kv_cache} - self.slot_size = self.kv_cache[:, 0].nbytes - self._slot_size = self.slot_size - self._store_staging_bytes = 4096 - self._pending_store_staging_limit_bytes = 8192 - self._store_staging_pool = None - self._pending_save_staging_bytes = 0 - self._save_futures = [] - self._rope_rotary_dim = 8 - self._rope_base = 10000.0 - self._rope_is_neox_style = True - - def stage_store_batch(self, block_ids: list[int], spans: list[StoreWriteSpan]): - """Expose store staging through a public test helper.""" - return self._stage_store_batch(block_ids, spans) - - synced = [] - recorded = object() - monkeypatch.setattr( - worker, - "_synchronize_cuda_tensor", - lambda tensor: synced.append(tensor), - ) - monkeypatch.setattr(worker, "_record_cuda_event", lambda tensor: recorded) - - probe = Probe() - staged = probe.stage_store_batch( - block_ids=[0, 1, 2], - spans=[StoreWriteSpan(0, probe.slot_size * 3, 0, "k0", 0, 3)], - ) - - assert staged is not None - assert staged.ready_event is recorded - assert synced == [] - - -def test_stage_store_batch_keeps_dynamic_rope_warmup_out_of_store_path(monkeypatch): - from daser.connector import worker - - class Probe(WorkerConnectorMixin): - def __init__(self) -> None: - self.kv_cache = torch.zeros((2, 8, 4, 2, 8), dtype=torch.float32) - self._layer_names = ["layer.0"] - self._layer_idx_map = {"layer.0": 0} - self._kv_caches = {"layer.0": self.kv_cache} - self.slot_size = self.kv_cache[:, 0].nbytes - self._slot_size = self.slot_size - self._store_staging_bytes = 4096 - self._pending_store_staging_limit_bytes = 8192 - self._store_staging_pool = None - self._pending_save_staging_bytes = 0 - self._save_futures = [] - self._rope_rotary_dim = 8 - self._rope_base = 10000.0 - self._rope_is_neox_style = True - - def stage_store_batch(self, block_ids: list[int], spans: list[StoreWriteSpan]): - """Expose store staging through a public test helper.""" - return self._stage_store_batch(block_ids, spans) - - calls = [] - monkeypatch.setattr( - worker, - "_warm_rope_apply_backends", - lambda **kwargs: calls.append(kwargs), - ) - - probe = Probe() - spans = [StoreWriteSpan(0, probe.slot_size * 3, 0, "k0", 0, 3)] - assert probe.stage_store_batch([0, 1, 2], spans) is not None - assert probe.stage_store_batch([3, 4, 5], spans) is not None - - assert calls == [] - - def test_update_state_after_alloc_skips_chunks_beyond_external_prefix(): - class MockConnector(SchedulerConnectorMixin): + class MockConnector(RequestLifecycle): def __init__(self) -> None: self._block_tokens = BLOCK_TOKENS self._slot_size = 32 @@ -2486,7 +1570,7 @@ class MockBlocks: connector = MockConnector() - DaserConnector.update_state_after_alloc( + RequestLifecycle.update_state_after_alloc( connector, MockRequest(), MockBlocks(), num_external_tokens=8 ) @@ -2676,7 +1760,7 @@ def __init__(self) -> None: calls: list[tuple[list[int], int, str | None, int]] = [] monkeypatch.setattr( - "daser.connector.reuse.rolling_prefix_keys", + "daser.connector.scheduler.reuse.rolling_prefix_keys", lambda tokens, block_tokens, initial_key=None, start_slot=0, **kwargs: ( calls.append((list(tokens), block_tokens, initial_key, start_slot)) or rolling_prefix_keys( @@ -3096,7 +2180,7 @@ def lookup(self, tokens, model_id): } ] - class MockConnector(SchedulerConnectorMixin): + class MockConnector(RequestLifecycle): def __init__(self) -> None: self._runtime_config_ready = True self._block_tokens = BLOCK_TOKENS @@ -3429,7 +2513,7 @@ def fake_rope( kv_block[:, :, 0, ..., :rotary_dim].add_(10.0) monkeypatch.setattr( - "daser.connector.staging.apply_rope_delta_to_kv_key_block", + "daser.connector.worker.staging.apply_rope_delta_to_kv_key_block", fake_rope, ) @@ -3481,7 +2565,7 @@ def fail_rope(*args, **kwargs): raise AssertionError("RoPE should not run for pos_offset=0") monkeypatch.setattr( - "daser.connector.staging.apply_rope_delta_to_key_block", + "daser.connector.worker.staging.apply_rope_delta_to_key_block", fail_rope, ) @@ -3526,7 +2610,7 @@ def fake_rope(kv_block, delta, rope_base, rotary_dim, is_neox_style): kv_block[:, :, 0, ..., :rotary_dim].add_(10) monkeypatch.setattr( - "daser.connector.staging._apply_rope_delta_with_tables", + "daser.connector.worker.staging._apply_rope_delta_with_tables", fake_rope, ) @@ -3574,7 +2658,7 @@ def fake_rope( kv_block[:, :, 0, ..., :rotary_dim].add_(3.0) monkeypatch.setattr( - "daser.connector.staging._apply_rope_delta_with_tables", + "daser.connector.worker.staging._apply_rope_delta_with_tables", fake_rope, ) @@ -3634,14 +2718,14 @@ def fail_rope(*args, **kwargs): raise AssertionError("legacy RoPE path should not run") monkeypatch.setattr( - "daser.connector.staging.apply_rope_delta_to_kv_key_block_table", + "daser.connector.worker.staging.apply_rope_delta_to_kv_key_block_table", fake_table_rope, ) monkeypatch.setattr( - "daser.connector.staging._apply_rope_delta_to_kv_key_block", + "daser.connector.worker.staging._apply_rope_delta_to_kv_key_block", fail_rope, ) - monkeypatch.setattr("daser.connector.staging._rope_table_cache", {}) + monkeypatch.setattr("daser.connector.worker.staging._rope_table_cache", {}) copies = _copy_staging_to_kv_cache( staging=staging, @@ -3682,7 +2766,7 @@ def fake_table_restore( return True monkeypatch.setattr( - "daser.connector.staging._restore_cross_layer_with_tables", + "daser.connector.worker.staging._restore_cross_layer_with_tables", fake_table_restore, ) @@ -3721,7 +2805,7 @@ def failing_restore(*args, **kwargs): raise RuntimeError("backend unavailable") monkeypatch.setattr( - "daser.connector.staging._restore_cross_layer_with_tables", + "daser.connector.worker.staging._restore_cross_layer_with_tables", failing_restore, ) @@ -3789,65 +2873,57 @@ def test_build_staging_store_batches_uses_spec_file_offset(): @pytest.mark.asyncio -async def test_commit_empty_stored_keys_does_not_publish_requested_chunks(): - """Skipped stale stores must not commit requested chunks.""" - connector = _CommitProbe() +async def test_store_cuda_export_selects_staged_buffer_device(monkeypatch) -> None: + """Background CUDA IPC export must select the TP rank's staged device.""" + from daser.connector.worker import store as store_module - await connector.commit_stored_keys([], ["stale-key"]) + selected_devices: list[torch.device] = [] + transferred: list[dict] = [] - assert connector.committed == [[]] - assert connector.commit_metadata == [(0, 1)] + class Client: + async def transfer_store_cuda(self, **kwargs): + transferred.append(kwargs) + return [] + + pipeline = StorePipeline.__new__(StorePipeline) + pipeline._client = Client() # noqa: SLF001 + buffer = SimpleNamespace(device=torch.device("cuda:1"), nbytes=32) + staged = StagedStoreBatch(buffer=buffer, spans=[], lease=object()) + cupy_buffer = object() + + monkeypatch.setattr(torch.cuda, "set_device", selected_devices.append) + monkeypatch.setattr(store_module.cupy, "asarray", lambda tensor: cupy_buffer) + monkeypatch.setattr(store_module, "cuda_array_pointer", lambda array: 4096) + monkeypatch.setattr( + store_module, + "cuda_allocation_base_and_offset", + lambda pointer: (pointer, 0), + ) + monkeypatch.setattr(store_module, "export_cuda_ipc_handle", lambda array: b"ipc") + monkeypatch.setattr(store_module, "cuda_array_device_id", lambda array: 1) + + await pipeline._write_cuda_buffer(staged) # noqa: SLF001 + + assert selected_devices == [torch.device("cuda:1")] + assert transferred[0]["device_id"] == 1 def test_tensor_parallel_rank_lanes_are_contiguous_and_disjoint() -> None: """Each TP rank maps a logical slot run into its own contiguous lane.""" local_slot_size = 32 rank_stride = 10 * local_slot_size + start_slot = 3 - rank_0 = _rank_lane_offset(3, local_slot_size, rank_stride, tp_rank=0) - rank_1 = _rank_lane_offset(3, local_slot_size, rank_stride, tp_rank=1) + rank_0 = start_slot * local_slot_size + rank_1 = rank_stride + start_slot * local_slot_size assert rank_0 == 3 * local_slot_size assert rank_1 == rank_stride + 3 * local_slot_size assert rank_0 + 2 * local_slot_size <= rank_1 -@pytest.mark.asyncio -async def test_worker_load_and_store_use_separate_ipc_clients(): - """Worker load RPCs should not queue behind store RPCs on one IPC client.""" - connector = _QueueProbe() - - await connector.load_via_public_helper() - await connector.store_via_public_helper() - - assert connector.load_calls == ["load"] - assert connector.store_calls == ["store"] - - -def test_worker_load_and_store_use_separate_background_loops(monkeypatch): - """Foreground loads should not queue behind background store loop work.""" - connector = _LoopProbe() - submitted_loops = [] - - class Future: - def result(self, timeout: float | None = None) -> None: - return None - - def run_threadsafe(coro, loop): - submitted_loops.append(loop) - coro.close() - return Future() - - monkeypatch.setattr(asyncio, "run_coroutine_threadsafe", run_threadsafe) - - connector.submit_load() - connector.submit_store() - - assert submitted_loops == list(connector.loop_pair) - - -def test_derive_store_staging_limits_scale_with_vram(monkeypatch): - """GPU staging caps consider device size and currently free VRAM.""" +def test_derive_staging_layout_scales_with_vram(monkeypatch): + """One budget preserves balanced pools and the explicit 6 GiB ceiling.""" class Props: def __init__(self, total_memory: int) -> None: @@ -3863,15 +2939,15 @@ def __init__(self, total_memory: int) -> None: "mem_get_info", lambda device=None: (12 << 30, 24 << 30), ) - small_batch, small_pending = _derive_store_staging_limits(torch.device("cuda")) + small_batch, small_load, small_store, small_total = derive_staging_layout( + torch.device("cuda"), 64 << 20, 8, 1 << 30 + ) assert small_batch == max( MIN_STORE_STAGING_BYTES, min((24 << 30) // 50, (12 << 30) // 10), ) - assert small_pending == max( - small_batch, - min((24 << 30) // 25, (12 << 30) // 5), - ) + assert (small_load, small_store) == (2, 2) + assert small_total == small_batch * 4 monkeypatch.setattr( torch.cuda, @@ -3883,23 +2959,29 @@ def __init__(self, total_memory: int) -> None: "mem_get_info", lambda device=None: (64 << 30, 80 << 30), ) - large_batch, large_pending = _derive_store_staging_limits(torch.device("cuda")) + large_batch, large_load, large_store, large_total = derive_staging_layout( + torch.device("cuda"), 64 << 20, 8, 1 << 30 + ) assert large_batch == DEFAULT_STORE_STAGING_BYTES - assert large_pending == DEFAULT_PENDING_STORE_STAGING_BYTES + assert (large_load, large_store) == (2, 2) + assert large_total == DEFAULT_STAGING_BUDGET_BYTES monkeypatch.setattr( torch.cuda, "mem_get_info", lambda device=None: (8 << 30, 80 << 30), ) - tight_batch, tight_pending = _derive_store_staging_limits(torch.device("cuda")) + tight_batch, tight_load, tight_store, tight_total = derive_staging_layout( + torch.device("cuda"), 64 << 20, 8, 1 << 30 + ) assert tight_batch == (8 << 30) // 10 - assert tight_pending == (8 << 30) // 5 + assert (tight_load, tight_store) == (5, 2) + assert tight_total == tight_batch * 7 def test_store_cuda_staging_pool_reuses_preallocated_buffer(): """Store staging pool reuses its init-time allocation after release.""" - pool = StoreCudaStagingPool( + pool = FixedCudaStagingPool( device=torch.device("cpu"), buffer_bytes=128, depth=1, diff --git a/tests/connector/test_gds_transfer.py b/tests/connector/test_gds_transfer.py index 6627521..3394350 100644 --- a/tests/connector/test_gds_transfer.py +++ b/tests/connector/test_gds_transfer.py @@ -89,7 +89,7 @@ def test_missing_file_raises(tmp_path): def test_fixed_staging_pool_reuses_two_preallocated_buffers() -> None: """Fixed staging pool reuses bounded preallocated buffers.""" # First Party - from daser.connector.staging import FixedCudaStagingPool + from daser.connector.worker.memory import FixedCudaStagingPool pool = FixedCudaStagingPool( device=torch.device("cpu"), @@ -114,7 +114,7 @@ def test_fixed_staging_pool_reuses_two_preallocated_buffers() -> None: def test_fixed_staging_pool_can_block_until_buffer_release() -> None: """Fixed staging pool can wait for a callback to release capacity.""" # First Party - from daser.connector.staging import FixedCudaStagingPool + from daser.connector.worker.memory import FixedCudaStagingPool pool = FixedCudaStagingPool( device=torch.device("cpu"), @@ -139,7 +139,7 @@ def release_first() -> None: def test_fixed_staging_pool_rejects_oversized_request() -> None: """Fixed staging pool rejects requests larger than one buffer.""" # First Party - from daser.connector.staging import FixedCudaStagingPool + from daser.connector.worker.memory import FixedCudaStagingPool pool = FixedCudaStagingPool( device=torch.device("cpu"), diff --git a/tests/connector/test_worker_pipelines.py b/tests/connector/test_worker_pipelines.py new file mode 100644 index 0000000..686cbe7 --- /dev/null +++ b/tests/connector/test_worker_pipelines.py @@ -0,0 +1,219 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +from concurrent.futures import Future +import time +from types import SimpleNamespace +from typing import Any + +import pytest + +pytest.importorskip("torch") +pytest.importorskip("vllm") +pytest.importorskip("cupy") + +import torch + +from daser.connector.metadata import ReqLoadSpec, ReqStoreSpec +from daser.connector.worker.load import LoadPipeline +from daser.connector.worker.memory import FixedCudaStagingPool +from daser.connector.worker.store import StagedStoreBatch, StorePipeline + + +class _ManualFuture: + def __init__(self) -> None: + self.complete = False + + def done(self) -> bool: + return self.complete + + def result(self, timeout: float) -> None: + del timeout + + +def _store_spec(key: str, blocks: list[int]) -> ReqStoreSpec: + return ReqStoreSpec(key, 0, len(blocks), blocks, 0, len(blocks)) + + +def test_store_pipeline_defers_and_dispatches_up_to_pool_depth() -> None: + pipeline = StorePipeline.__new__(StorePipeline) + pipeline._pending_finished_saves = {} # noqa: SLF001 + pipeline._staging_pool = SimpleNamespace(depth=2) # noqa: SLF001 + submitted: list[_ManualFuture] = [] + + def submit(save: Any) -> None: + future = _ManualFuture() + save.future = future + submitted.append(future) + + pipeline._submit_save = submit # type: ignore[method-assign] # noqa: SLF001 + pipeline.queue_finished( + {req: _store_spec(req, [index]) for index, req in enumerate(("a", "b", "c"))}, + {"a", "b", "c"}, + ) + + assert pipeline.collect_finished(set()) == set() + assert submitted == [] + assert pipeline.collect_finished({"a", "b", "c"}) == set() + assert len(submitted) == 2 + + submitted[0].complete = True + assert pipeline.collect_finished(set()) == {"a"} + assert len(submitted) == 3 + submitted[1].complete = True + submitted[2].complete = True + assert pipeline.collect_finished(set()) == {"b", "c"} + + +def test_store_pipeline_streams_request_larger_than_pool_depth() -> None: + pipeline = StorePipeline.__new__(StorePipeline) + pipeline._pending_finished_saves = {} # noqa: SLF001 + pipeline._staging_pool = SimpleNamespace(depth=1) # noqa: SLF001 + pipeline._kv_caches = {"layer": torch.empty(1)} # noqa: SLF001 + pipeline._local_slot_size = 16 # noqa: SLF001 + pipeline._rank_stride_bytes = 0 # noqa: SLF001 + pipeline._tp_rank = 0 # noqa: SLF001 + pipeline._tp_size = 1 # noqa: SLF001 + pipeline._staging_bytes = 32 # noqa: SLF001 + released: list[int] = [] + writes: list[int] = [] + + class Lease: + view = torch.empty(1) + + def release(self) -> None: + released.append(1) + + class Client: + async def commit_chunks(self, keys: list[str], **kwargs: int) -> None: + assert keys == ["large"] + assert kwargs == {"tp_rank": 0, "tp_size": 1} + + def stage(block_ids: list[int], spans: list[Any], event: Any) -> StagedStoreBatch: + del event + return StagedStoreBatch(torch.empty(1), spans, Lease()) + + async def write(staged: StagedStoreBatch) -> list[str]: + writes.append(sum(span.nbytes for span in staged.spans) // 16) + return ["large"] + + def submit(coro: Any) -> Future[None]: + future: Future[None] = Future() + try: + asyncio.run(coro) + future.set_result(None) + except BaseException as exc: + future.set_exception(exc) + return future + + pipeline._client = Client() # type: ignore[assignment] # noqa: SLF001 + pipeline._stage_batch = stage # type: ignore[method-assign] # noqa: SLF001 + pipeline._write_cuda_buffer = write # type: ignore[method-assign] # noqa: SLF001 + pipeline._submit = submit # type: ignore[method-assign] # noqa: SLF001 + pipeline._submit_save = StorePipeline._submit_save.__get__(pipeline) # noqa: SLF001 + pipeline._kv_caches = {} # noqa: SLF001 + pipeline.queue_finished({"large": _store_spec("large", [0, 1, 2, 3, 4])}, {"large"}) + pipeline._kv_caches = {"layer": torch.empty(1)} # noqa: SLF001 + + assert pipeline.collect_finished({"large"}) == set() + assert pipeline.collect_finished(set()) == {"large"} + assert writes == [2, 2, 1] + assert len(released) == 3 + + +class _LoadClient: + def __init__(self, fail_offset: int | None = None) -> None: + self.calls: list[int] = [] + self.fail_offset = fail_offset + + async def transfer_load_registered_cuda(self, **kwargs: Any) -> dict[str, Any]: + offset = int(kwargs["spans"][0]["file_offset"]) + self.calls.append(offset) + if offset == self.fail_offset: + raise RuntimeError("load failed") + return { + "transfer_open_ms": 1.0, + "transfer_load_ms": 2.0, + "transfer_sync_ms": 3.0, + "transfer_stats_delta": {"l1_hits": 4, "l1_misses": 5, "l2_reads": 6}, + } + + async def close(self) -> None: + return None + + +def _load_spec(key: str, blocks: list[int], offset: int = 0) -> ReqLoadSpec: + return ReqLoadSpec(key, offset // 16, len(blocks), blocks, offset, len(blocks)) + + +def _load_pipeline( + monkeypatch: pytest.MonkeyPatch, client: _LoadClient +) -> LoadPipeline: + monkeypatch.setattr( + "daser.connector.worker.load.copy_staging_to_kv_cache", + lambda **kwargs: 1, + ) + pipeline = LoadPipeline("unused.sock", client_count=2) + pipeline._clients = [client, client] # type: ignore[assignment] # noqa: SLF001 + pipeline.configure( + kv_caches={"layer": torch.empty(1)}, + layer_names=["layer"], + local_slot_size=16, + rank_stride_bytes=0, + tp_rank=0, + staging_pool=FixedCudaStagingPool(torch.device("cpu"), 32, 2), + load_key_scale=1.0, + load_value_scale=1.0, + rope_delta_scale=1.0, + rope_base=10000.0, + rope_rotary_dim=0, + rope_is_neox_style=True, + ) + pipeline._staging_registered = True # noqa: SLF001 + return pipeline + + +def _wait_finished(pipeline: LoadPipeline, expected: set[str]) -> set[str]: + deadline = time.monotonic() + 2.0 + finished: set[str] = set() + while time.monotonic() < deadline and finished != expected: + finished.update(pipeline.collect_finished()) + time.sleep(0.005) + return finished + + +def test_load_pipeline_handles_empty_and_multibatch_requests( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = _LoadClient() + pipeline = _load_pipeline(monkeypatch, client) + try: + pipeline.start( + { + "empty": _load_spec("empty", []), + "large:load:0": _load_spec("large", [0, 1, 2, 3, 4]), + } + ) + assert _wait_finished(pipeline, {"empty", "large"}) == {"empty", "large"} + assert client.calls == [0, 32, 64] + finally: + pipeline.shutdown() + + +def test_load_failure_invalidates_only_failed_request( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = _LoadClient(fail_offset=0) + pipeline = _load_pipeline(monkeypatch, client) + try: + pipeline.start( + { + "bad": _load_spec("bad", [7], 0), + "good": _load_spec("good", [8], 16), + } + ) + assert _wait_finished(pipeline, {"bad", "good"}) == {"bad", "good"} + assert pipeline.take_invalid_block_ids() == {7} + assert sorted(client.calls) == [0, 16] + finally: + pipeline.shutdown() diff --git a/tests/server/test_ipc_server.py b/tests/server/test_ipc_server.py index f217b37..1994f79 100644 --- a/tests/server/test_ipc_server.py +++ b/tests/server/test_ipc_server.py @@ -3,6 +3,7 @@ # Standard import asyncio import os +from types import SimpleNamespace from typing import Any # Third Party @@ -19,7 +20,7 @@ from daser.server.doc_registry import DocRegistry from daser.server.ipc import IPCServer from daser.server.metadata_store import MetadataStore -from daser.transfer.base import TransferLayer +from daser.transfer.base import TransferLayer, TransferStats SLOT_SIZE = 4096 BLOCK_TOKENS = 4 @@ -227,6 +228,30 @@ async def test_ipc_server_records_operation_metrics(tmp_path) -> None: ) +def test_ipc_server_records_tier_counter_deltas(tmp_path) -> None: + """Tier metrics publish monotonic L1 and L2 deltas exactly once.""" + registry = MetricsRegistry() + server = IPCServer( + str(tmp_path / "test.sock"), + make_core(), + make_runtime_config(tmp_path), + metrics_registry=registry, + ) + transfer = SimpleNamespace( + stats=TransferStats(l1_hits=2, l1_misses=3, l2_reads=4), + l1_bytes_used=1024, + ) + server._transfer = transfer # type: ignore[assignment] # noqa: SLF001 + + server._record_tier_metrics() # noqa: SLF001 + server._record_tier_metrics() # noqa: SLF001 + + rendered = registry.render_prometheus() + assert "daser_l1_hits_total 2.0" in rendered + assert "daser_l1_misses_total 3.0" in rendered + assert "daser_l2_reads_total 4.0" in rendered + + @pytest.mark.asyncio async def test_ipc_server_records_external_prefix_cache_metrics(tmp_path) -> None: """IPC can publish vLLM-equivalent external prefix cache counters.""" diff --git a/tests/unit/test_benchmark_unified_utils.py b/tests/unit/test_benchmark_unified_utils.py index 4218bfc..564c955 100644 --- a/tests/unit/test_benchmark_unified_utils.py +++ b/tests/unit/test_benchmark_unified_utils.py @@ -8,6 +8,7 @@ from pathlib import Path import subprocess import sys +from types import SimpleNamespace from typing import Any import pytest @@ -2466,6 +2467,91 @@ async def fake_collect( assert hit_rate == 0.75 +def test_daser_evict_gate_requires_l1_and_l2_activity() -> None: + """One warm evict phase must include positive deltas from both tiers.""" + vllm_bench._require_daser_evict_tier_activity( # noqa: SLF001 + { + "backend_prometheus": { + "daser_l1_hits_total": 3.0, + "daser_l2_reads_total": 2.0, + } + } + ) + for counters in ( + {}, + {"daser_l1_hits_total": 1.0}, + {"daser_l1_hits_total": 1.0, "daser_l2_reads_total": 0.0}, + ): + with pytest.raises(RuntimeError, match="must exercise both tiers"): + vllm_bench._require_daser_evict_tier_activity( # noqa: SLF001 + {"backend_prometheus": counters} + ) + + +def test_daser_evict_warm_phase_primes_same_seed_prefix( + tmp_path: Path, + monkeypatch, +) -> None: + """Evict warm metrics include a short LRU perturbation and full phase.""" + commands: list[list[str]] = [] + + async def fake_collect( + manifest: Any, + before_metrics: dict[str, Any] | None = None, + ) -> dict[str, Any]: + del manifest + if before_metrics is None: + return {"backend_prometheus": {}} + return { + "backend_prometheus": { + "daser_l1_hits_total": 1.0, + "daser_l2_reads_total": 1.0, + }, + "hit_ratios": {}, + } + + monkeypatch.setattr(vllm_bench, "collect_phase_metrics", fake_collect) + args = RunBenchArgs(model="model", bench_num_prompts=10) + manifest = SimpleNamespace( + endpoints={"vllm": SimpleNamespace(url="http://127.0.0.1:8001")} + ) + + metrics, _hit_rate = vllm_bench._run_daser_evict_warm_phase( # noqa: SLF001 + args, + manifest, + tmp_path / "warm.json", + run_command=commands.append, + ) + + assert [command[command.index("--num-prompts") + 1] for command in commands] == [ + "2", + "10", + ] + assert [command[command.index("--seed") + 1] for command in commands] == [ + "42", + "42", + ] + assert metrics["backend_prometheus"]["daser_l1_hits_total"] == 1.0 + + +def test_daser_drain_failure_aborts_benchmark(monkeypatch) -> None: + """A failed cold-to-warm barrier is a benchmark failure.""" + monkeypatch.setattr( + vllm_bench.httpx, + "post", + lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("timeout")), + ) + manifest = SimpleNamespace( + endpoints={"daser": SimpleNamespace(url="http://127.0.0.1:9999")} + ) + + with pytest.raises(RuntimeError, match="timeout"): + vllm_bench._drain_daser( # noqa: SLF001 + manifest, + print_kv=lambda key, value: None, + ) + + def test_run_bench_vllm_bench_entrypoint_runs_openai_rows( tmp_path: Path, monkeypatch, diff --git a/tests/unit/test_block_aligned_bug.py b/tests/unit/test_block_aligned_bug.py index da03ed9..b812918 100644 --- a/tests/unit/test_block_aligned_bug.py +++ b/tests/unit/test_block_aligned_bug.py @@ -8,10 +8,11 @@ # Third Party import pytest -# First Party -from daser.connector.daser_connector import DaserConnector from daser.connector.helpers import PendingStore, hash_tokens +# First Party +from daser.connector.scheduler.lifecycle import RequestLifecycle + class TestTokeniseAndTruncateBug: def test_exact_block_aligned_length_not_multiple_of_block_tokens(self): @@ -73,7 +74,7 @@ def lookup(self, prefix, model_id): } ] - class MockDaserConnector(DaserConnector): + class MockDaserConnector(RequestLifecycle): def __init__(self): self._block_tokens = BLOCK_TOKENS self._socket_path = "/tmp/test.sock" @@ -121,7 +122,7 @@ def lookup(self, prefix, model_id): } ] - class MockDaserConnector(DaserConnector): + class MockDaserConnector(RequestLifecycle): def __init__(self): self._block_tokens = BLOCK_TOKENS self._socket_path = "/tmp/test.sock" @@ -169,7 +170,7 @@ def alloc_chunk(self, chunk_key, token_count, model_id): mock_ipc = MockIPCClientSync() - class MockDaserConnector(DaserConnector): + class MockDaserConnector(RequestLifecycle): def __init__(self): self._block_tokens = BLOCK_TOKENS self._socket_path = "/tmp/test.sock" @@ -230,7 +231,7 @@ def alloc_chunk(self, chunk_key, token_count, model_id): mock_ipc = MockIPCClientSync() - class MockDaserConnector(DaserConnector): + class MockDaserConnector(RequestLifecycle): def __init__(self): self._block_tokens = BLOCK_TOKENS self._socket_path = "/tmp/test.sock" @@ -299,7 +300,7 @@ def alloc_chunk(self, chunk_key, token_count, model_id): mock_ipc = MockIPCClientSync() tokens = list(range(630)) - class MockDaserConnector(DaserConnector): + class MockDaserConnector(RequestLifecycle): def __init__(self): self._block_tokens = BLOCK_TOKENS self._socket_path = "/tmp/test.sock" @@ -400,7 +401,7 @@ def lookup(self, prefix, model_id): } ] - class MockDaserConnector(DaserConnector): + class MockDaserConnector(RequestLifecycle): def __init__(self): self._block_tokens = BLOCK_TOKENS self._socket_path = "/tmp/test.sock" @@ -427,7 +428,9 @@ def __init__(self): ) def test_trim_external_window_keeps_full_block_for_lmcache_style_minus_one(self): - from daser.connector.scheduler import _trim_chunk_to_external_window + from daser.connector.scheduler.planning import ( + _trim_chunk_to_external_window, + ) chunk = { "chunk_key": "k", diff --git a/tests/unit/test_chunk_reuse_scheduler.py b/tests/unit/test_chunk_reuse_scheduler.py index 18ee782..bc6434d 100644 --- a/tests/unit/test_chunk_reuse_scheduler.py +++ b/tests/unit/test_chunk_reuse_scheduler.py @@ -4,12 +4,12 @@ from typing import Any # First Party -from daser.connector.scheduler import SchedulerConnectorMixin +from daser.connector.scheduler.lifecycle import RequestLifecycle BLOCK_TOKENS = 4 -class _SchedulerProbe(SchedulerConnectorMixin): +class _SchedulerProbe(RequestLifecycle): """Minimal scheduler-role connector for chunk reuse credit tests.""" def __init__(self, chunks: list[dict[str, Any]]) -> None: diff --git a/tests/unit/test_skip_save_flag.py b/tests/unit/test_skip_save_flag.py index 7d9977d..3468b1c 100644 --- a/tests/unit/test_skip_save_flag.py +++ b/tests/unit/test_skip_save_flag.py @@ -14,10 +14,11 @@ # Standard from typing import Any, Optional -# First Party -from daser.connector.daser_connector import DaserConnector from daser.connector.helpers import PendingStore, hash_tokens +# First Party +from daser.connector.scheduler.lifecycle import RequestLifecycle + BLOCK_TOKENS = 16 @@ -57,7 +58,7 @@ def record_external_prefix_cache(self, queries: int, hits: int) -> None: self.external_prefix_records.append((queries, hits)) -class _MockDaserConnector(DaserConnector): +class _MockDaserConnector(RequestLifecycle): """Test connector that bypasses vLLM init for scheduler-path testing. Mirrors the subclass pattern used by ``test_block_aligned_bug`` so