From a60339677c36744d48343513379f484aa33b338c Mon Sep 17 00:00:00 2001 From: Teng Ma Date: Mon, 4 May 2026 02:00:32 +0800 Subject: [PATCH 01/30] Add ECMooncakeConnector for encoder cache over Mooncake TransferEngine Wire factory and ec_transfer config; add two-process e2e test, EPD full-pipeline script, and README notes. Co-authored-by: Cursor Signed-off-by: Teng Ma --- tests/v1/ec_connector/integration/README.md | 14 +- .../run_epd_mooncake_ec_full_pipeline.sh | 224 ++++++++++ .../test_ec_mooncake_transfer_e2e.py | 203 +++++++++ vllm/config/ec_transfer.py | 5 + .../ec_transfer/ec_connector/factory.py | 6 + .../ec_connector/mooncake_ec_connector.py | 397 ++++++++++++++++++ 6 files changed, 848 insertions(+), 1 deletion(-) create mode 100755 tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh create mode 100644 tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py diff --git a/tests/v1/ec_connector/integration/README.md b/tests/v1/ec_connector/integration/README.md index a7dab5d5d9d1..376604dcb66e 100644 --- a/tests/v1/ec_connector/integration/README.md +++ b/tests/v1/ec_connector/integration/README.md @@ -13,7 +13,7 @@ The test ensures that disaggregated encoding produces **identical** outputs to t Note that currently PD disaggregation set up may give slightly different results from a single instance. Therefore, we need the result from 1P+1D as the baseline for 1E+1P+1D -Please refer to [Disaggregated Encoder Feature](../../../../docs/features/disagg_encoder.md) for the detailed explanation for the EPD features. +Please refer to [Disaggregated Encoder Feature](../../../../docs/features/disagg_encoder.md) for the detailed explanation for the EPD features. ## Files @@ -124,6 +124,18 @@ Quick sanity check: - Safe to run multiple times (idempotent) - We setup the PD disagg part with NixlConnector. Please read details about EPD in `examples/disaggregated/disaggregated_encoder/README.md` +## ECMooncakeConnector (TransferEngine) smoke test + +Two-process transfer over Mooncake (no full vLLM serve, no HF model download): + +```bash +cd vllm +PYTHONPATH=. python tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py +``` + +Requires: **2+ CUDA GPUs**, `mooncake-transfer-engine`, `pyzmq`, `httpx`, `fastapi`, `uvicorn`. +Optional: `MOONCAKE_EC_PROTOCOL=rdma` or `=tcp` (default in test mocks is `tcp`) to match your cluster. + ## Requirements - Multiple GPUs (3 for 1E+1P+1D, 2 for 1E+1PD, 1 for baseline) diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh new file mode 100755 index 000000000000..eb031c044bf8 --- /dev/null +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -0,0 +1,224 @@ +#!/usr/bin/env bash +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# +# Full-stack EPD validation with ECMooncakeConnector (Mooncake TransferEngine): +# 1) Single-GPU baseline (multimodal) -> saves baseline JSON +# 2) 1 Encoder + 1 PD with Mooncake EC -> compare outputs to baseline +# +# Usage (from repo root): +# ./tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +# +# Env: +# MODEL HF model id (default: Qwen/Qwen2.5-VL-3B-Instruct) +# GPU_SINGLE / GPU_E / GPU_PD GPU ids (defaults 0 / 1 / 2) +# ENDPOINT_PORT, ENCODE_PORT, PREFILL_DECODE_PORT +# EC_MOONCAKE_REGISTRY_PORT HTTP registry on encoder (default 19018) +# EC_REGISTRY_HOST Host PD uses to query registry (default 127.0.0.1) +# MOONCAKE_EC_PROTOCOL rdma | tcp (default rdma) +# USE_MM_PROMPTS 1 (default) or 0 for text-only quick sanity +# TIMEOUT_SECONDS wait_for_server timeout (default 1200) +# SKIP_BASELINE set to 1 to reuse existing BASELINE_FILE + +set -u + +GIT_ROOT=$(git rev-parse --show-toplevel) +cd "$GIT_ROOT" || exit 1 +export PYTHONPATH="${GIT_ROOT}:${PYTHONPATH:-}" + +MODEL="${MODEL:-Qwen/Qwen2.5-VL-3B-Instruct}" +USE_MM_PROMPTS="${USE_MM_PROMPTS:-1}" +MM_FLAG="" +if [[ "$USE_MM_PROMPTS" == "1" ]]; then + MM_FLAG="--use_mm_prompts" +fi + +GPU_SINGLE="${GPU_SINGLE:-0}" +GPU_E="${GPU_E:-1}" +GPU_PD="${GPU_PD:-2}" + +ENCODE_PORT="${ENCODE_PORT:-19534}" +PREFILL_DECODE_PORT="${PREFILL_DECODE_PORT:-19537}" +ENDPOINT_PORT="${ENDPOINT_PORT:-10002}" +BASELINE_PORT="${BASELINE_PORT:-10003}" + +EC_MOONCAKE_REGISTRY_PORT="${EC_MOONCAKE_REGISTRY_PORT:-19018}" +EC_REGISTRY_HOST="${EC_REGISTRY_HOST:-127.0.0.1}" +MOONCAKE_EC_PROTOCOL="${MOONCAKE_EC_PROTOCOL:-rdma}" +EC_REGISTRY_URL="http://${EC_REGISTRY_HOST}:${EC_MOONCAKE_REGISTRY_PORT}" +export EC_MOONCAKE_REGISTRY_PORT EC_REGISTRY_HOST MOONCAKE_EC_PROTOCOL EC_REGISTRY_URL + +LOG_PATH="${LOG_PATH:-/tmp}" +BASELINE_FILE="${BASELINE_FILE:-/tmp/vllm_epd_mooncake_baseline.txt}" +TIMEOUT_SECONDS="${TIMEOUT_SECONDS:-1200}" + +mkdir -p "$LOG_PATH" + +if command -v vllm &>/dev/null; then + VLLM_SERVE=(vllm serve) +else + VLLM_SERVE=(python -m vllm.entrypoints.cli.main serve) +fi + +ENC_EC_JSON=$(python3 </dev/null || true + pkill -f "vllm.entrypoints.cli.main serve" 2>/dev/null || true + pkill -f "disagg_epd_proxy.py" 2>/dev/null || true + sleep 2 +} + +trap 'cleanup_instances; kill $(jobs -pr) 2>/dev/null || true' EXIT INT TERM + +run_baseline() { + echo "================================" + echo "BASELINE (single vLLM, MM if enabled)" + echo "================================" + cleanup_instances + local PORT=$BASELINE_PORT + echo "Starting baseline on GPU $GPU_SINGLE port $PORT" + CUDA_VISIBLE_DEVICES="$GPU_SINGLE" "${VLLM_SERVE[@]}" "$MODEL" \ + --port "$PORT" \ + --enforce-eager \ + --gpu-memory-utilization 0.75 \ + --max-num-seqs 32 \ + --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ + >"${LOG_PATH}/mooncake_epd_baseline.log" 2>&1 & + local BASELINE_PID=$! + echo "Waiting for baseline..." + wait_for_server "$PORT" || { echo "Baseline failed to start; tail log:"; tail -80 "${LOG_PATH}/mooncake_epd_baseline.log"; return 1; } + curl -s "http://127.0.0.1:${PORT}/v1/models" | head -c 200 || true + echo "" + python "${GIT_ROOT}/tests/v1/ec_connector/integration/test_epd_correctness.py" \ + --service_url "http://localhost:$PORT" \ + --model_name "$MODEL" \ + --mode baseline \ + --baseline_file "$BASELINE_FILE" \ + $MM_FLAG + kill "$BASELINE_PID" 2>/dev/null || true + sleep 2 + cleanup_instances +} + +run_epd_mooncake() { + echo "================================" + echo "EPD 1E + 1PD with ECMooncakeConnector" + echo "Registry URL for consumer: $EC_REGISTRY_URL" + echo "Mooncake protocol: $MOONCAKE_EC_PROTOCOL" + echo "================================" + cleanup_instances + + declare -a PIDS=() + + echo "Starting ENCODER on GPU $GPU_E port $ENCODE_PORT" + CUDA_VISIBLE_DEVICES="$GPU_E" "${VLLM_SERVE[@]}" "$MODEL" \ + --port "$ENCODE_PORT" \ + --enforce-eager \ + --gpu-memory-utilization 0.35 \ + --enable-request-id-headers \ + --no-enable-prefix-caching \ + --max-num-batched-tokens 114688 \ + --max-num-seqs 32 \ + --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ + --ec-transfer-config "$ENC_EC_JSON" \ + >"${LOG_PATH}/mooncake_epd_encoder.log" 2>&1 & + PIDS+=($!) + + echo "Starting PD on GPU $GPU_PD port $PREFILL_DECODE_PORT" + CUDA_VISIBLE_DEVICES="$GPU_PD" "${VLLM_SERVE[@]}" "$MODEL" \ + --port "$PREFILL_DECODE_PORT" \ + --enforce-eager \ + --gpu-memory-utilization 0.75 \ + --enable-request-id-headers \ + --max-num-seqs 32 \ + --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ + --ec-transfer-config "$PD_EC_JSON" \ + >"${LOG_PATH}/mooncake_epd_pd.log" 2>&1 & + PIDS+=($!) + + echo "Waiting for encoder..." + wait_for_server "$ENCODE_PORT" || { echo "Encoder log:"; tail -100 "${LOG_PATH}/mooncake_epd_encoder.log"; return 1; } + echo "Waiting for PD..." + wait_for_server "$PREFILL_DECODE_PORT" || { echo "PD log:"; tail -100 "${LOG_PATH}/mooncake_epd_pd.log"; return 1; } + + echo "Starting EPD proxy on $ENDPOINT_PORT" + python "${GIT_ROOT}/examples/online_serving/disaggregated_encoder/disagg_epd_proxy.py" \ + --host "0.0.0.0" \ + --port "$ENDPOINT_PORT" \ + --encode-servers-urls "http://localhost:$ENCODE_PORT" \ + --prefill-servers-urls "disable" \ + --decode-servers-urls "http://localhost:$PREFILL_DECODE_PORT" \ + >"${LOG_PATH}/mooncake_epd_proxy.log" 2>&1 & + PIDS+=($!) + + echo "Waiting for proxy..." + wait_for_server "$ENDPOINT_PORT" || { echo "Proxy log:"; tail -80 "${LOG_PATH}/mooncake_epd_proxy.log"; return 1; } + curl -s "http://127.0.0.1:${ENDPOINT_PORT}/health" || true + echo "" + + python "${GIT_ROOT}/tests/v1/ec_connector/integration/test_epd_correctness.py" \ + --service_url "http://localhost:$ENDPOINT_PORT" \ + --model_name "$MODEL" \ + --mode disagg \ + --baseline_file "$BASELINE_FILE" \ + $MM_FLAG + + for pid in "${PIDS[@]}"; do + kill "$pid" 2>/dev/null || true + done + sleep 2 + cleanup_instances +} + +echo "================================" +echo "EPD + ECMooncake full pipeline" +echo "MODEL=$MODEL" +echo "================================" + +if [[ "${SKIP_BASELINE:-0}" != "1" ]]; then + run_baseline +else + echo "SKIP_BASELINE=1 -> using existing $BASELINE_FILE" + [[ -f "$BASELINE_FILE" ]] || { echo "Missing baseline file"; exit 1; } +fi + +run_epd_mooncake + +echo "================================" +echo "PASS: Mooncake EC EPD matches baseline" +echo "Logs: ${LOG_PATH}/mooncake_epd_*.log" +echo "================================" diff --git a/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py b/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py new file mode 100644 index 000000000000..6c4b97bcf449 --- /dev/null +++ b/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py @@ -0,0 +1,203 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +""" +End-to-end test: ECMooncakeConnector producer (GPU0) publishes EC metadata and +serves pull requests; consumer (GPU1) fetches layout over HTTP and pulls the +tensor via Mooncake TransferEngine + ZMQ. + +Requires: 2+ CUDA GPUs, mooncake-transfer-engine, pyzmq, httpx, fastapi, uvicorn. + +Protocol: ``mooncake_protocol`` defaults to ``tcp`` in mocks unless you set +``MOONCAKE_EC_PROTOCOL=rdma`` (matches ``ec_connector_extra_config.mooncake_protocol``). +Example RDMA run:: + + MOONCAKE_EC_PROTOCOL=rdma PYTHONPATH=. python tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py +""" + +from __future__ import annotations + +import multiprocessing as mp +import os +import time +from unittest.mock import Mock + +import torch + +try: + import zmq # noqa: F401 +except ImportError as e: + raise SystemExit("pyzmq is required: pip install pyzmq") from e +try: + import mooncake # noqa: F401 +except ImportError as e: + raise SystemExit("mooncake-transfer-engine is required") from e + +from vllm.config import VllmConfig +from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole +from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( + ECMooncakeConnector, + ECMooncakeConnectorMetadata, + ECMooncakeLoadSpec, +) + + +def _find_free_port() -> int: + import socket + + s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + s.bind(("127.0.0.1", 0)) + _, port = s.getsockname() + s.close() + return int(port) + + +def _mock_vllm_producer(registry_port: int) -> Mock: + cfg = Mock(spec=VllmConfig) + cfg.parallel_config = Mock() + cfg.parallel_config.tensor_parallel_size = 1 + cfg.parallel_config.pipeline_parallel_size = 1 + cfg.ec_transfer_config = Mock() + cfg.ec_transfer_config.is_ec_producer = True + cfg.ec_transfer_config.is_ec_consumer = False + cfg.ec_transfer_config.ec_buffer_device = "cuda" + cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), + "registry_http_port": registry_port, + } + return cfg + + +def _mock_vllm_consumer() -> Mock: + cfg = Mock(spec=VllmConfig) + cfg.parallel_config = Mock() + cfg.parallel_config.tensor_parallel_size = 1 + cfg.parallel_config.pipeline_parallel_size = 1 + cfg.ec_transfer_config = Mock() + cfg.ec_transfer_config.is_ec_producer = False + cfg.ec_transfer_config.is_ec_consumer = True + cfg.ec_transfer_config.ec_buffer_device = "cuda" + cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), + "remote_registry_url": "http://unused-on-worker", + } + return cfg + + +def _producer_entry( + mm_hash: str, + registry_port: int, + ready: mp.Queue, + done: mp.Event, + barrier: mp.Barrier, +) -> None: + os.environ["CUDA_VISIBLE_DEVICES"] = "0" + torch.cuda.init() + cfg = _mock_vllm_producer(registry_port) + conn = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) + torch.manual_seed(12345) + tensor = torch.randn(8, 64, device="cuda", dtype=torch.float32) + cache = {mm_hash: tensor} + conn.save_caches(cache, mm_hash) + ready.put("ok") + barrier.wait(timeout=120) + # Hold process until consumer finishes transfer + done.wait(timeout=180) + + +def _consumer_entry( + mm_hash: str, + registry_url: str, + barrier: mp.Barrier, + result_queue: mp.Queue, +) -> None: + os.environ["CUDA_VISIBLE_DEVICES"] = "1" + torch.cuda.init() + barrier.wait(timeout=120) + import httpx + + url = f"{registry_url.rstrip('/')}/ec/info/{mm_hash}" + for _ in range(60): + try: + r = httpx.get(url, timeout=2.0) + if r.status_code == 200: + break + except httpx.HTTPError: + pass + time.sleep(0.5) + else: + result_queue.put({"ok": False, "err": "registry never ready"}) + return + data = r.json() + spec = ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=1, + nbytes=int(data["nbytes"]), + shape=tuple(int(x) for x in data["shape"]), + dtype=str(data["dtype"]), + producer_zmq=str(data["producer_zmq"]), + ) + meta = ECMooncakeConnectorMetadata() + meta.add_load(spec) + cfg = _mock_vllm_consumer() + conn = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) + conn.bind_connector_metadata(meta) + enc: dict[str, torch.Tensor] = {} + try: + conn.start_load_caches(enc) + except Exception as e: + result_queue.put({"ok": False, "err": repr(e)}) + return + got = enc.get(mm_hash) + if got is None: + result_queue.put({"ok": False, "err": "missing tensor"}) + return + torch.manual_seed(12345) + expected = torch.randn(8, 64, device="cuda", dtype=torch.float32) + max_diff = (got.cpu() - expected.cpu()).abs().max().item() + result_queue.put({"ok": True, "max_diff": max_diff}) + + +def test_ec_mooncake_two_process_transfer(): + """Producer on cuda:0 and consumer on cuda:1 transfer one EC tensor.""" + mm_hash = "e2e_mm_test_hash" + registry_port = _find_free_port() + registry_url = f"http://127.0.0.1:{registry_port}" + + ctx = mp.get_context("spawn") + ready: mp.Queue = ctx.Queue() + result: mp.Queue = ctx.Queue() + done = ctx.Event() + barrier = ctx.Barrier(2) + prod = ctx.Process( + target=_producer_entry, + args=(mm_hash, registry_port, ready, done, barrier), + daemon=True, + ) + cons = ctx.Process( + target=_consumer_entry, + args=(mm_hash, registry_url, barrier, result), + daemon=True, + ) + prod.start() + assert ready.get(timeout=120) == "ok" + cons.start() + cons.join(timeout=180) + done.set() + prod.join(timeout=30) + + assert not cons.is_alive(), "consumer process hung" + assert cons.exitcode == 0, f"consumer exit {cons.exitcode}" + out = result.get(timeout=1) + assert out["ok"], out.get("err", out) + assert out["max_diff"] < 1e-4, f"tensor mismatch max_diff={out['max_diff']}" + + +def _main() -> None: + if torch.cuda.device_count() < 2: + raise SystemExit("Need at least 2 CUDA devices for this e2e test.") + test_ec_mooncake_two_process_transfer() + print("ECMooncake two-process transfer e2e: PASSED") + + +if __name__ == "__main__": + _main() diff --git a/vllm/config/ec_transfer.py b/vllm/config/ec_transfer.py index c64bf13ec0ca..3948f3205934 100644 --- a/vllm/config/ec_transfer.py +++ b/vllm/config/ec_transfer.py @@ -18,6 +18,11 @@ class ECTransferConfig: ec_connector: str | None = None """The EC connector for vLLM to transmit EC caches between vLLM instances. + + Built-in options include ``ECExampleConnector`` (shared filesystem via + safetensors) and ``ECMooncakeConnector`` (Mooncake TransferEngine RDMA; + requires ``mooncake-transfer-engine`` and matching producer/consumer + ``ec_connector_extra_config``; see ``mooncake_ec_connector`` module docstring). """ engine_id: str | None = None diff --git a/vllm/distributed/ec_transfer/ec_connector/factory.py b/vllm/distributed/ec_transfer/ec_connector/factory.py index 598b1b22e4af..1c33af551371 100644 --- a/vllm/distributed/ec_transfer/ec_connector/factory.py +++ b/vllm/distributed/ec_transfer/ec_connector/factory.py @@ -89,3 +89,9 @@ def get_connector_class( "vllm.distributed.ec_transfer.ec_connector.cpu.connector", "ECCPUConnector", ) + +ECConnectorFactory.register_connector( + "ECMooncakeConnector", + "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector", + "ECMooncakeConnector", +) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py new file mode 100644 index 000000000000..e65b7d086e82 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -0,0 +1,397 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +""" +Encoder-cache (EC) connector backed by Mooncake TransferEngine. + +Used in disaggregated setups where an encoder / prefill instance produces +multimodal encoder outputs and a decode instance loads them over RDMA-capable +Mooncake transport instead of shared filesystem. +""" + +from __future__ import annotations + +import json +import threading +import time +from dataclasses import dataclass, field +from typing import Any + +import httpx +import torch +import uvicorn +import zmq +from fastapi import FastAPI, HTTPException + +from vllm.distributed.ec_transfer.ec_connector.base import ( + ECConnectorBase, + ECConnectorMetadata, + ECConnectorRole, +) +from vllm.distributed.parallel_state import is_local_first_rank +from vllm.logger import init_logger +from vllm.utils.network_utils import get_ip +from vllm.config import VllmConfig +from vllm.v1.core.sched.output import SchedulerOutput + +logger = init_logger(__name__) + +try: + from mooncake.engine import TransferEngine +except ImportError as e: + TransferEngine = None # type: ignore[misc, assignment] + _MOONCAKE_IMPORT_ERROR = e +else: + _MOONCAKE_IMPORT_ERROR = None + + +@dataclass +class ECMooncakeLoadSpec: + """Per-item metadata shipped from scheduler to worker (pickle-friendly).""" + + mm_hash: str + num_token: int + nbytes: int + shape: tuple[int, ...] + dtype: str + producer_zmq: str + + +@dataclass +class ECMooncakeConnectorMetadata(ECConnectorMetadata): + """Worker-side metadata for one scheduler step.""" + + loads: list[ECMooncakeLoadSpec] = field(default_factory=list) + + def add_load(self, spec: ECMooncakeLoadSpec) -> None: + self.loads.append(spec) + + +class ECMooncakeRegistryServer: + """Lightweight HTTP registry on the producer for remote has_cache_item / info.""" + + def __init__(self, host: str, port: int): + self.host = host + self.port = port + self._entries: dict[str, dict[str, Any]] = {} + self._lock = threading.Lock() + self.app = FastAPI() + self._register_routes() + self.server_thread: threading.Thread | None = None + self.server: uvicorn.Server | None = None + + def _register_routes(self) -> None: + @self.app.get("/ec/info/{mm_hash}") + async def ec_info(mm_hash: str) -> dict[str, Any]: + with self._lock: + data = self._entries.get(mm_hash) + if data is None: + raise HTTPException(status_code=404, detail="unknown mm_hash") + return data + + def start(self) -> None: + if self.server_thread is not None: + return + config = uvicorn.Config(app=self.app, host=self.host, port=self.port) + self.server = uvicorn.Server(config=config) + self.server_thread = threading.Thread( + target=self.server.run, name="ec_mooncake_registry", daemon=True + ) + self.server_thread.start() + while self.server is not None and not self.server.started: + time.sleep(0.05) + logger.info( + "EC Mooncake registry listening on http://%s:%d", self.host, self.port + ) + + def shutdown(self) -> None: + if self.server is None or self.server_thread is None or not self.server.started: + return + self.server.should_exit = True + self.server_thread.join() + logger.info("EC Mooncake registry stopped.") + + def publish(self, mm_hash: str, payload: dict[str, Any]) -> None: + with self._lock: + self._entries[mm_hash] = payload + + def unpublish(self, mm_hash: str) -> None: + with self._lock: + self._entries.pop(mm_hash, None) + + +class ECMooncakeConnector(ECConnectorBase): + """ + EC connector using Mooncake TransferEngine for GPU tensor transport. + + Extra config (``ec_connector_extra_config``): + + - ``remote_registry_url`` (consumer, required): Base URL of the producer + registry, e.g. ``http://192.168.0.2:9018``. + - ``registry_http_port`` (producer, optional): Port for the in-process HTTP + registry (default ``9018``). + - ``mooncake_protocol`` (optional): Passed to ``TransferEngine.initialize`` + (default ``"rdma"``). + + Limitations: ``tensor_parallel_size`` and ``pipeline_parallel_size`` must + be ``1`` (same assumption as Mooncake KV connector for P2P handshake). + """ + + def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): + super().__init__(vllm_config=vllm_config, role=role) + if _MOONCAKE_IMPORT_ERROR is not None or TransferEngine is None: + raise ImportError( + "Install mooncake-transfer-engine (see " + "https://github.com/kvcache-ai/Mooncake ) to use ECMooncakeConnector." + ) from _MOONCAKE_IMPORT_ERROR + + if vllm_config.parallel_config.tensor_parallel_size > 1: + raise ValueError("ECMooncakeConnector requires tensor_parallel_size=1.") + if vllm_config.parallel_config.pipeline_parallel_size > 1: + raise ValueError( + "ECMooncakeConnector does not support pipeline parallelism yet." + ) + + self._role = role + self._ec_cfg = vllm_config.ec_transfer_config + assert self._ec_cfg is not None + self._extra = self._ec_cfg.ec_connector_extra_config + self._protocol: str = self._extra.get("mooncake_protocol", "rdma") + self._remote_registry_url: str | None = self._extra.get("remote_registry_url") + self._registry_http_port: int = int(self._extra.get("registry_http_port", 9018)) + + # Scheduler (consumer): mm_hash -> pending tensor layout from registry + self._pending_specs: dict[str, ECMooncakeLoadSpec] = {} + self._mm_datas_need_loads: dict[str, int] = {} + + # Worker producer + self._engine: TransferEngine | None = None + self._hostname = get_ip() + self._registry: ECMooncakeRegistryServer | None = None + self._zmq_listen_addr: str | None = None + self._zmq_thread: threading.Thread | None = None + self._zmq_ctx: zmq.Context | None = None + self._tensor_by_hash: dict[str, torch.Tensor] = {} + self._tensor_lock = threading.Lock() + self._producer_services_started = False + + if role == ECConnectorRole.SCHEDULER and self.is_consumer: + if not self._remote_registry_url: + raise ValueError( + "ec_consumer with ECMooncakeConnector requires " + "ec_connector_extra_config['remote_registry_url']." + ) + + def _ensure_engine(self) -> TransferEngine: + if self._engine is None: + eng = TransferEngine() + ret = eng.initialize(self._hostname, "P2PHANDSHAKE", self._protocol, "") + if ret != 0: + raise RuntimeError("Mooncake TransferEngine initialization failed.") + self._engine = eng + logger.info( + "ECMooncakeConnector TransferEngine ready at %s:%d", + self._hostname, + eng.get_rpc_port(), + ) + return self._engine + + def _start_producer_zmq_listener(self) -> None: + if self._zmq_thread is not None: + return + + def loop() -> None: + assert self._zmq_ctx is not None + sock = self._zmq_ctx.socket(zmq.REP) + port = sock.bind_to_random_port(f"tcp://{self._hostname}") + self._zmq_listen_addr = f"tcp://{self._hostname}:{port}" + logger.info("EC Mooncake pull listener at %s", self._zmq_listen_addr) + eng = self._ensure_engine() + while True: + try: + raw = sock.recv() + except zmq.ContextTerminated: + break + try: + req = json.loads(raw.decode("utf-8")) + if req.get("op") != "pull": + sock.send_json({"ok": False, "err": "unknown op"}) + continue + mm_hash = req["mm_hash"] + dst_session = req["dst_session"] + dst_ptr = int(req["dst_ptr"]) + nbytes = int(req["nbytes"]) + with self._tensor_lock: + tensor = self._tensor_by_hash.get(mm_hash) + if tensor is None: + sock.send_json({"ok": False, "err": "unknown mm_hash"}) + continue + src_ptr = tensor.data_ptr() + if tensor.nbytes != nbytes: + sock.send_json({"ok": False, "err": "size mismatch"}) + continue + ret = eng.batch_transfer_sync_write( + dst_session, [src_ptr], [dst_ptr], [nbytes] + ) + sock.send_json({"ok": ret == 0, "mooncake_ret": int(ret)}) + except Exception as e: + logger.exception("EC Mooncake pull handler error: %s", e) + try: + sock.send_json({"ok": False, "err": str(e)}) + except zmq.ZMQError: + break + + self._zmq_ctx = zmq.Context() + self._zmq_thread = threading.Thread(target=loop, name="ec-mooncake-zmq", daemon=True) + self._zmq_thread.start() + while self._zmq_listen_addr is None: + time.sleep(0.01) + + def _ensure_producer_services(self) -> None: + if self._producer_services_started: + return + if not self.is_producer or self._role != ECConnectorRole.WORKER: + return + self._ensure_engine() + self._start_producer_zmq_listener() + if is_local_first_rank(): + self._registry = ECMooncakeRegistryServer("0.0.0.0", self._registry_http_port) + self._registry.start() + self._producer_services_started = True + + def start_load_caches( + self, encoder_cache: dict[str, torch.Tensor], **kwargs: Any + ) -> None: + metadata = self._get_connector_metadata() + assert isinstance(metadata, ECMooncakeConnectorMetadata) + eng = self._ensure_engine() + raw_buf = self._ec_cfg.ec_buffer_device + buf = ( + raw_buf.lower() + if isinstance(raw_buf, str) and raw_buf + else "cuda" + ) + if buf == "cuda" and not torch.cuda.is_available(): + raise RuntimeError("ECMooncakeConnector requires CUDA for ec_buffer_device=cuda") + device = torch.device(buf) + + for spec in metadata.loads: + if spec.mm_hash in encoder_cache: + continue + torch_dtype = getattr(torch, spec.dtype, None) + if torch_dtype is None: + raise ValueError(f"Unsupported torch dtype string: {spec.dtype!r}") + t = torch.empty(spec.shape, dtype=torch_dtype, device=device) + ret = eng.batch_register_memory([t.data_ptr()], [t.nbytes]) + if ret != 0: + raise RuntimeError( + "Mooncake EC batch_register_memory failed on consumer." + ) + pull = { + "op": "pull", + "mm_hash": spec.mm_hash, + "dst_session": f"{self._hostname}:{eng.get_rpc_port()}", + "dst_ptr": t.data_ptr(), + "nbytes": t.nbytes, + } + ctx = zmq.Context() + sock = ctx.socket(zmq.REQ) + sock.setsockopt(zmq.RCVTIMEO, 120_000) + sock.connect(spec.producer_zmq) + try: + sock.send_json(pull) + resp = sock.recv_json() + finally: + sock.close(linger=0) + ctx.term() + if not resp.get("ok"): + raise RuntimeError(f"EC Mooncake pull failed: {resp}") + encoder_cache[spec.mm_hash] = t + logger.debug("Loaded EC tensor for mm_hash=%s via Mooncake", spec.mm_hash) + + def save_caches( + self, encoder_cache: dict[str, torch.Tensor], mm_hash: str, **kwargs: Any + ) -> None: + if not self.is_producer or self._role != ECConnectorRole.WORKER: + return + self._ensure_producer_services() + tensor = encoder_cache[mm_hash] + eng = self._ensure_engine() + ret = eng.batch_register_memory([tensor.data_ptr()], [tensor.nbytes]) + if ret != 0: + raise RuntimeError("Mooncake EC batch_register_memory failed on producer.") + with self._tensor_lock: + self._tensor_by_hash[mm_hash] = tensor + + dtype_str = str(tensor.dtype).split(".")[-1] + payload = { + "nbytes": tensor.nbytes, + "shape": list(tensor.shape), + "dtype": dtype_str, + "producer_zmq": self._zmq_listen_addr, + } + if self._registry is not None: + self._registry.publish(mm_hash, payload) + logger.debug("Published EC tensor mm_hash=%s to registry", mm_hash) + + def has_cache_item(self, identifier: str) -> bool: + if not self.is_consumer or self._role != ECConnectorRole.SCHEDULER: + return False + assert self._remote_registry_url is not None + url = self._remote_registry_url.rstrip("/") + f"/ec/info/{identifier}" + try: + r = httpx.get(url, timeout=5.0) + except httpx.HTTPError as e: + logger.warning("EC Mooncake registry query failed for %s: %s", identifier, e) + return False + if r.status_code != 200: + return False + data = r.json() + zmq_addr = data.get("producer_zmq") + if not zmq_addr: + return False + self._pending_specs[identifier] = ECMooncakeLoadSpec( + mm_hash=identifier, + num_token=0, + nbytes=int(data["nbytes"]), + shape=tuple(int(x) for x in data["shape"]), + dtype=str(data["dtype"]), + producer_zmq=str(zmq_addr), + ) + return True + + def update_state_after_alloc(self, request: Any, index: int) -> None: + mm_hash = request.mm_features[index].identifier + num_encoder_token = request.get_num_encoder_embeds(index) + self._mm_datas_need_loads[mm_hash] = num_encoder_token + + def build_connector_meta( + self, scheduler_output: SchedulerOutput + ) -> ECConnectorMetadata: + meta = ECMooncakeConnectorMetadata() + for mm_hash, num_token in self._mm_datas_need_loads.items(): + spec = self._pending_specs.get(mm_hash) + if spec is None: + logger.warning("Missing EC Mooncake spec for mm_hash=%s", mm_hash) + continue + meta.add_load( + ECMooncakeLoadSpec( + mm_hash=spec.mm_hash, + num_token=num_token, + nbytes=spec.nbytes, + shape=spec.shape, + dtype=spec.dtype, + producer_zmq=spec.producer_zmq, + ) + ) + self._pending_specs.pop(mm_hash, None) + self._mm_datas_need_loads.clear() + return meta + + def __del__(self) -> None: + try: + if self._registry is not None: + self._registry.shutdown() + if self._zmq_ctx is not None: + self._zmq_ctx.term() + except Exception: + pass From 023796417971f6c46cc0a98b21171b60d5d137a2 Mon Sep 17 00:00:00 2001 From: Teng Ma Date: Mon, 25 May 2026 14:17:49 +0800 Subject: [PATCH 02/30] Add Mooncake transfer reliability knobs and EC connector sync Introduce max_transfer_bytes splitting, sync_after_transfer, and optional integrity verification for Mooncake KV/EC transfers (vllm #42395). Document CUDA 13 wheel and tuning in mooncake_connector_usage.md. Add EC Mooncake unit tests and extend KV connector unit coverage. Co-authored-by: Claude Co-authored-by: Cursor Signed-off-by: Teng Ma --- docs/features/mooncake_connector_usage.md | 39 +++ .../test_ec_mooncake_transfer_e2e.py | 5 + .../unit/test_ec_mooncake_connector.py | 329 ++++++++++++++++++ .../unit/test_mooncake_connector.py | 53 +++ .../ec_connector/mooncake_ec_connector.py | 13 + .../v1/mooncake/mooncake_connector.py | 132 ++++++- vllm/envs.py | 19 + 7 files changed, 584 insertions(+), 6 deletions(-) create mode 100644 tests/v1/ec_connector/unit/test_ec_mooncake_connector.py diff --git a/docs/features/mooncake_connector_usage.md b/docs/features/mooncake_connector_usage.md index 6cce042fb6f8..b211d0a31891 100644 --- a/docs/features/mooncake_connector_usage.md +++ b/docs/features/mooncake_connector_usage.md @@ -14,6 +14,12 @@ Install mooncake through pip: `uv pip install mooncake-transfer-engine-cuda13`. vLLM defaults to CUDA 13. On a CUDA 12 environment install `mooncake-transfer-engine` instead — the two are the same release built against different CUDA majors, and the wrong one fails to import with `libcudart.so.: cannot open shared object file`. +If you observe PD transfer data mismatches (`dst != src`) with the CUDA 13 +package, enable the mitigations in +[Transfer reliability](#transfer-reliability) below (see +[vllm #42395](https://github.com/vllm-project/vllm/issues/42395), +[Mooncake #2086](https://github.com/kvcache-ai/Mooncake/issues/2086)). + Refer to [Mooncake official repository](https://github.com/kvcache-ai/Mooncake) for more installation instructions ## Usage @@ -56,6 +62,36 @@ Now you can send requests to the proxy server through port 8000. - Default: 480 - If a request is aborted and the decoder has not yet notified the prefiller, the prefill instance will release its KV-cache blocks after this timeout to avoid holding them indefinitely. +### Transfer reliability + +Under concurrent PD load, some Mooncake transfer-engine builds (notably +`mooncake-transfer-engine-cuda13==0.3.10.post2`) can produce destination bytes +that do not match the producer source for very large coalesced descriptors. +vLLM mitigations (env vars or `kv_connector_extra_config` keys): + +- `VLLM_MOONCAKE_MAX_TRANSFER_BYTES` / `max_transfer_bytes`: Split any single + transfer descriptor larger than this size into contiguous chunks (recommended + starting value: `262144` for multimodal PD workloads). +- `VLLM_MOONCAKE_SYNC_AFTER_TRANSFER` / `sync_after_transfer`: Call + `torch.cuda.synchronize()` after each Mooncake batch transfer on producer and + consumer (reduces visibility races at some throughput cost). +- `VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY` / `verify_transfer_integrity`: + Debug-only SHA-256 check that producer source memory is unchanged after + transfer (does not verify remote destination bytes). + +Example prefill/decode extra config: + +```json +{ + "kv_connector": "MooncakeConnector", + "kv_role": "kv_producer", + "kv_connector_extra_config": { + "max_transfer_bytes": 262144, + "sync_after_transfer": true + } +} +``` + ## KV Transfer Config ### KV Role Options @@ -69,6 +105,9 @@ Now you can send requests to the proxy server through port 8000. - **num_workers**: Size of thread pool for one prefiller worker to transfer KV caches by mooncake. (default 10) - **mooncake_protocol**: Mooncake connector protocol. (default "rdma") - **device_name**: Comma-separated whitelist of RDMA devices (e.g. `"mlx5_0,mlx5_1"`) to restrict topology discovery to. Empty discovers every device. Useful on hosts exposing a mix of InfiniBand and RoCE ports, where both peers must settle on the same link layer. +- **max_transfer_bytes**: Split descriptors larger than this many bytes (see [Transfer reliability](#transfer-reliability)) +- **sync_after_transfer**: Synchronize CUDA after each Mooncake batch transfer (default false) +- **verify_transfer_integrity**: Debug SHA-256 check of producer source after transfer (default false) ## Example Scripts/Code diff --git a/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py b/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py index 6c4b97bcf449..e849dfb5144e 100644 --- a/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py +++ b/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py @@ -21,6 +21,7 @@ import time from unittest.mock import Mock +import pytest import torch try: @@ -157,6 +158,10 @@ def _consumer_entry( result_queue.put({"ok": True, "max_diff": max_diff}) +@pytest.mark.skipif( + torch.cuda.device_count() < 2, + reason="Requires at least 2 CUDA devices", +) def test_ec_mooncake_two_process_transfer(): """Producer on cuda:0 and consumer on cuda:1 transfer one EC tensor.""" mm_hash = "e2e_mm_test_hash" diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py new file mode 100644 index 000000000000..9fb3f622d7c2 --- /dev/null +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -0,0 +1,329 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Unit tests for ECMooncakeConnector and its HTTP registry.""" + +from __future__ import annotations + +import ctypes +import socket +import time +from contextlib import contextmanager +from unittest.mock import Mock, patch + +import httpx +import pytest +import torch + +from vllm.config import VllmConfig +from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole +from vllm.distributed.ec_transfer.ec_connector.factory import ECConnectorFactory +from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( + ECMooncakeConnector, + ECMooncakeConnectorMetadata, + ECMooncakeLoadSpec, + ECMooncakeRegistryServer, +) +from vllm.v1.core.sched.output import SchedulerOutput + +from tests.v1.ec_connector.unit.test_ec_example_connector import ( + mock_request_with_3_mm, +) + + +class CopyingFakeTransferEngine: + def __init__(self, *args, **kwargs): + pass + + def initialize(self, local_hostname, metadata_server, protocol, device_name) -> int: + return 0 + + def get_rpc_port(self) -> int: + return 12345 + + def batch_transfer_sync_write( + self, target_hostname, buffers, peer_buffer_addresses, lengths + ) -> int: + for src, dst, nbytes in zip(buffers, peer_buffer_addresses, lengths): + ctypes.memmove(int(dst), int(src), int(nbytes)) + return 0 + + def batch_register_memory(self, buffer_addresses, capacities) -> int: + return 0 + + +def _find_free_port() -> int: + s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + s.bind(("127.0.0.1", 0)) + _, port = s.getsockname() + s.close() + return int(port) + + +@pytest.fixture +def mock_vllm_config_producer(): + config = Mock(spec=VllmConfig) + config.parallel_config = Mock() + config.parallel_config.tensor_parallel_size = 1 + config.parallel_config.pipeline_parallel_size = 1 + config.ec_transfer_config = Mock() + config.ec_transfer_config.is_ec_producer = True + config.ec_transfer_config.is_ec_consumer = False + config.ec_transfer_config.ec_buffer_device = "cuda" + config.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "registry_http_port": 19018, + } + return config + + +@pytest.fixture +def mock_vllm_config_consumer(): + config = Mock(spec=VllmConfig) + config.parallel_config = Mock() + config.parallel_config.tensor_parallel_size = 1 + config.parallel_config.pipeline_parallel_size = 1 + config.ec_transfer_config = Mock() + config.ec_transfer_config.is_ec_producer = False + config.ec_transfer_config.is_ec_consumer = True + config.ec_transfer_config.ec_buffer_device = "cuda" + config.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "remote_registry_url": "http://127.0.0.1:19018", + } + return config + + +@contextmanager +def patch_ec_mooncake_deps(): + with ( + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector.TransferEngine", + CopyingFakeTransferEngine, + ), + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector.get_ip", + return_value="127.0.0.1", + ), + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector.is_local_first_rank", + return_value=True, + ), + ): + yield + + +class TestECMooncakeRegistryServer: + def test_publish_and_lookup(self): + port = _find_free_port() + registry = ECMooncakeRegistryServer("127.0.0.1", port) + registry.start() + try: + payload = { + "nbytes": 128, + "shape": [4, 8], + "dtype": "float32", + "producer_zmq": "tcp://127.0.0.1:9999", + } + registry.publish("hash_a", payload) + r = httpx.get(f"http://127.0.0.1:{port}/ec/info/hash_a", timeout=2.0) + assert r.status_code == 200 + assert r.json() == payload + r404 = httpx.get( + f"http://127.0.0.1:{port}/ec/info/missing", timeout=2.0 + ) + assert r404.status_code == 404 + finally: + registry.shutdown() + + def test_unpublish_removes_entry(self): + port = _find_free_port() + registry = ECMooncakeRegistryServer("127.0.0.1", port) + registry.start() + try: + registry.publish("h", {"nbytes": 1, "shape": [], "dtype": "float32"}) + registry.unpublish("h") + r = httpx.get(f"http://127.0.0.1:{port}/ec/info/h", timeout=2.0) + assert r.status_code == 404 + finally: + registry.shutdown() + + +class TestECMooncakeFactory: + def test_factory_registers_connector(self): + cls = ECConnectorFactory.get_connector_class( + Mock(ec_connector="ECMooncakeConnector") + ) + assert cls.__name__ == "ECMooncakeConnector" + + +class TestECMooncakeConnectorValidation: + def test_consumer_scheduler_requires_remote_registry(self, mock_vllm_config_consumer): + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + } + with patch_ec_mooncake_deps(): + with pytest.raises(ValueError, match="remote_registry_url"): + ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + + def test_rejects_tensor_parallel_gt_one(self, mock_vllm_config_producer): + mock_vllm_config_producer.parallel_config.tensor_parallel_size = 2 + with patch_ec_mooncake_deps(): + with pytest.raises(ValueError, match="tensor_parallel_size"): + ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + +class TestECMooncakeSchedulerMetadata: + def test_has_cache_item_queries_registry( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + port = _find_free_port() + registry = ECMooncakeRegistryServer("127.0.0.1", port) + registry.start() + try: + mm_hash = mock_request_with_3_mm.mm_features[0].identifier + registry.publish( + mm_hash, + { + "nbytes": 64, + "shape": [2, 4], + "dtype": "float32", + "producer_zmq": "tcp://127.0.0.1:1", + }, + ) + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "remote_registry_url" + ] = f"http://127.0.0.1:{port}" + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + assert scheduler.has_cache_item(mm_hash) + assert mm_hash in scheduler._pending_specs + spec = scheduler._pending_specs[mm_hash] + assert spec.shape == (2, 4) + assert spec.dtype == "float32" + finally: + registry.shutdown() + + def test_has_cache_item_missing_returns_false( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + port = _find_free_port() + registry = ECMooncakeRegistryServer("127.0.0.1", port) + registry.start() + try: + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "remote_registry_url" + ] = f"http://127.0.0.1:{port}" + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + mm_hash = mock_request_with_3_mm.mm_features[0].identifier + assert not scheduler.has_cache_item(mm_hash) + finally: + registry.shutdown() + + def test_build_connector_meta_clears_pending( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + mm_hash = mock_request_with_3_mm.mm_features[0].identifier + scheduler._pending_specs[mm_hash] = ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=0, + nbytes=32, + shape=(2, 4), + dtype="float32", + producer_zmq="tcp://127.0.0.1:1", + ) + scheduler._mm_datas_need_loads[mm_hash] = 100 + meta = scheduler.build_connector_meta(Mock(spec=SchedulerOutput)) + assert isinstance(meta, ECMooncakeConnectorMetadata) + assert len(meta.loads) == 1 + assert meta.loads[0].mm_hash == mm_hash + assert meta.loads[0].num_token == 100 + assert scheduler._mm_datas_need_loads == {} + assert mm_hash not in scheduler._pending_specs + + +class TestECMooncakeWorkerTransfer: + def test_single_process_save_and_load(self, mock_vllm_config_producer): + """Host-memory pull path (fake engine uses memcpy; CUDA ptrs need e2e).""" + port = _find_free_port() + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "registry_http_port" + ] = port + mm_hash = "unit_test_hash" + torch.manual_seed(7) + source = torch.randn(4, 16, dtype=torch.float32) + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + producer.save_caches({mm_hash: source}, mm_hash) + for _ in range(100): + if producer._zmq_listen_addr is not None: + break + time.sleep(0.01) + assert producer._zmq_listen_addr is not None + + url = f"http://127.0.0.1:{port}/ec/info/{mm_hash}" + r = httpx.get(url, timeout=2.0) + assert r.status_code == 200 + data = r.json() + + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + } + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) + spec = ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=1, + nbytes=int(data["nbytes"]), + shape=tuple(int(x) for x in data["shape"]), + dtype=str(data["dtype"]), + producer_zmq=str(data["producer_zmq"]), + ) + meta = ECMooncakeConnectorMetadata() + meta.add_load(spec) + consumer.bind_connector_metadata(meta) + loaded: dict[str, torch.Tensor] = {} + consumer.start_load_caches(loaded) + assert mm_hash in loaded + assert torch.allclose(loaded[mm_hash].cpu(), source.cpu()) + + def test_producer_scheduler_has_cache_item_false( + self, mock_vllm_config_producer, mock_request_with_3_mm + ): + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + mm_hash = mock_request_with_3_mm.mm_features[0].identifier + assert not scheduler.has_cache_item(mm_hash) + + def test_consumer_worker_save_is_noop(self, mock_vllm_config_consumer): + with patch_ec_mooncake_deps(): + worker = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + mm_hash = "noop_hash" + tensor = torch.randn(2, 4) + worker.save_caches({mm_hash: tensor}, mm_hash) + assert mm_hash not in worker._tensor_by_hash diff --git a/tests/v1/kv_connector/unit/test_mooncake_connector.py b/tests/v1/kv_connector/unit/test_mooncake_connector.py index 4847b956b196..259f980ca185 100644 --- a/tests/v1/kv_connector/unit/test_mooncake_connector.py +++ b/tests/v1/kv_connector/unit/test_mooncake_connector.py @@ -27,6 +27,7 @@ _align_transfer_regions, get_mooncake_bootstrap_addr, should_launch_bootstrap_server, + split_transfer_descriptors, ) from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import ( MooncakeBootstrapServer, @@ -349,6 +350,58 @@ async def run_in_executor(self, executor, func, *args): prefill_worker.sender_loop = origin_sender_loop prefill_worker.shutdown() +def test_split_transfer_descriptors_disabled(): + src = [100, 200] + dst = [300, 400] + lengths = [16, 32] + assert split_transfer_descriptors(src, dst, lengths, 0) == (src, dst, lengths) + + +def test_split_transfer_descriptors_splits_large_descriptor(): + # Reproduces coalesced multimodal descriptor size from vllm #42395. + src_base, dst_base, total = 0x1000, 0x2000, 2_424_832 + max_chunk = 262_144 + out_src, out_dst, out_len = split_transfer_descriptors( + [src_base], [dst_base], [total], max_chunk + ) + assert sum(out_len) == total + assert len(out_len) == (total + max_chunk - 1) // max_chunk + offset = 0 + for src, dst, length in zip(out_src, out_dst, out_len): + assert src == src_base + offset + assert dst == dst_base + offset + assert 0 < length <= max_chunk + offset += length + + +def test_send_blocks_splits_when_max_transfer_bytes_set(): + """_send_blocks should chunk descriptors before calling TransferEngine.""" + from vllm.distributed.kv_transfer.kv_connector.v1.mooncake import ( + mooncake_connector as mc, + ) + + captured: list[list[int]] = [] + + class StubWorker: + _max_transfer_bytes = 8 + _sync_after_transfer = False + _verify_transfer_integrity = False + device_id = 0 + + engine = FakeMooncakeWrapper() + + def _send_blocks(self, remote_session, src_ptrs, dst_ptrs, lengths): + return mc.MooncakeConnectorWorker._send_blocks( + self, remote_session, src_ptrs, dst_ptrs, lengths + ) + + worker = StubWorker() + worker.engine.batch_transfer_sync_write = ( # type: ignore[method-assign] + lambda _session, _src, _dst, lengths: captured.append(lengths) or 0 + ) + ret = worker._send_blocks("host:1", [1000], [2000], [20]) + assert ret == 0 + assert captured == [[8, 8, 4]] def test_basic_interface(): diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index e65b7d086e82..5a9a996cccb9 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -22,6 +22,7 @@ import zmq from fastapi import FastAPI, HTTPException +from vllm import envs from vllm.distributed.ec_transfer.ec_connector.base import ( ECConnectorBase, ECConnectorMetadata, @@ -232,6 +233,13 @@ def loop() -> None: ret = eng.batch_transfer_sync_write( dst_session, [src_ptr], [dst_ptr], [nbytes] ) + if ( + ret == 0 + and envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER + and torch.cuda.is_available() + and tensor.is_cuda + ): + torch.cuda.synchronize(device=tensor.device) sock.send_json({"ok": ret == 0, "mooncake_ret": int(ret)}) except Exception as e: logger.exception("EC Mooncake pull handler error: %s", e) @@ -305,6 +313,11 @@ def start_load_caches( ctx.term() if not resp.get("ok"): raise RuntimeError(f"EC Mooncake pull failed: {resp}") + if ( + envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER + and device.type == "cuda" + ): + torch.cuda.synchronize(device=device) encoder_cache[spec.mm_hash] = t logger.debug("Loaded EC tensor for mm_hash=%s via Mooncake", spec.mm_hash) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py index dd94ffb2f9ed..478dedcb549e 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import asyncio +import hashlib import logging import threading import time @@ -232,6 +233,75 @@ def _compute_sender_transfer_plan( ) +def split_transfer_descriptors( + src_ptrs: list[int], + dst_ptrs: list[int], + lengths: list[int], + max_bytes: int, +) -> tuple[list[int], list[int], list[int]]: + """Split descriptors longer than ``max_bytes`` into contiguous chunks. + + Mitigates Mooncake transfer-engine issues with very large single + descriptors under concurrent PD load (see vllm #42395). + """ + if max_bytes <= 0: + return src_ptrs, dst_ptrs, lengths + out_src: list[int] = [] + out_dst: list[int] = [] + out_len: list[int] = [] + for src, dst, length in zip(src_ptrs, dst_ptrs, lengths): + offset = 0 + while offset < length: + chunk = min(max_bytes, length - offset) + out_src.append(src + offset) + out_dst.append(dst + offset) + out_len.append(chunk) + offset += chunk + return out_src, out_dst, out_len + + +def _digest_gpu_memory_region(ptr: int, length: int, device_id: int) -> bytes: + """Return SHA-256 digest of a CUDA memory region (debug / verify mode).""" + if length == 0: + return hashlib.sha256(b"").digest() + if not torch.cuda.is_available(): + raise RuntimeError("GPU memory digest requires CUDA.") + from torch.cuda import cudart + + host = torch.empty(length, dtype=torch.uint8, device="cpu", pin_memory=True) + with torch.cuda.device(device_id): + err = cudart.cudaMemcpy( + host.data_ptr(), + ptr, + length, + cudart.cudaMemcpyKind.cudaMemcpyDeviceToHost, + )[0] + if err != 0: + raise RuntimeError(f"cudaMemcpy D2H failed with error {err}") + return hashlib.sha256(host.numpy().tobytes()).digest() + + +def _resolve_mooncake_transfer_tuning( + extra_config: dict[str, Any], +) -> tuple[int | None, bool, bool]: + """Merge kv_connector_extra_config with Mooncake env overrides.""" + max_bytes: int | None = None + if envs.VLLM_MOONCAKE_MAX_TRANSFER_BYTES > 0: + max_bytes = envs.VLLM_MOONCAKE_MAX_TRANSFER_BYTES + cfg_max = extra_config.get("max_transfer_bytes") + if cfg_max is not None: + max_bytes = int(cfg_max) + + sync_after = envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER + if "sync_after_transfer" in extra_config: + sync_after = bool(extra_config["sync_after_transfer"]) + + verify = envs.VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY + if "verify_transfer_integrity" in extra_config: + verify = bool(extra_config["verify_transfer_integrity"]) + return max_bytes, sync_after, verify + + def _can_coalesce_block_transfers( local_region_block_len: int, remote_region_block_len: int, @@ -925,15 +995,26 @@ def __init__( # Tasks can await async events, so a surplus (2x is a robust heuristic) # prevents workers from idling. self.num_sender_tasks = self.num_sender_workers * 2 - protocol = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr] - "mooncake_protocol", "rdma" - ) - device_name = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr] - "device_name", "" - ) + extra_config = kv_transfer_config.kv_connector_extra_config # type: ignore[union-attr] + protocol = extra_config.get("mooncake_protocol", "rdma") + device_name = extra_config.get("device_name", "") + ( + self._max_transfer_bytes, + self._sync_after_transfer, + self._verify_transfer_integrity, + ) = _resolve_mooncake_transfer_tuning(extra_config) logger.info( "The Mooncake Transfer Engine is using %s as its protocol.", protocol ) + if self._max_transfer_bytes: + logger.info( + "MooncakeConnector will split transfer descriptors above %d bytes", + self._max_transfer_bytes, + ) + if self._sync_after_transfer: + logger.info("MooncakeConnector sync_after_transfer is enabled") + if self._verify_transfer_integrity: + logger.info("MooncakeConnector verify_transfer_integrity is enabled") ret_value = self.engine.initialize( self.hostname, "P2PHANDSHAKE", protocol, device_name ) @@ -1626,12 +1707,49 @@ def _send_blocks( dst_ptrs: list[int], lengths: list[int], ) -> int: + if self._max_transfer_bytes: + src_ptrs, dst_ptrs, lengths = split_transfer_descriptors( + src_ptrs, dst_ptrs, lengths, self._max_transfer_bytes + ) + + pre_hashes: list[bytes] | None = None + if self._verify_transfer_integrity: + try: + pre_hashes = [ + _digest_gpu_memory_region(src, length, self.device_id) + for src, length in zip(src_ptrs, lengths) + ] + except Exception as e: + logger.error("Mooncake transfer integrity pre-hash failed: %s", e) + return -1 + start_time = time.perf_counter() ret_value = self.engine.batch_transfer_sync_write( remote_session, src_ptrs, dst_ptrs, lengths ) duration = time.perf_counter() - start_time if ret_value == 0: + if self._sync_after_transfer and torch.cuda.is_available(): + torch.cuda.synchronize(device=self.device_id) + if pre_hashes is not None: + for idx, (src, length) in enumerate(zip(src_ptrs, lengths)): + try: + post_hash = _digest_gpu_memory_region( + src, length, self.device_id + ) + except Exception as e: + logger.error( + "Mooncake transfer integrity post-hash failed: %s", e + ) + return -1 + if post_hash != pre_hashes[idx]: + logger.error( + "Mooncake source memory changed after transfer " + "(descriptor_idx=%d length=%d)", + idx, + length, + ) + return -1 self.xfer_stats.record_transfer( duration_s=duration, total_bytes=sum(lengths), @@ -1893,6 +2011,8 @@ def process_pulling_result( self.finished_recving_reqs.add(pull_meta.d_req_id) if ok_reqs: + if self._sync_after_transfer and torch.cuda.is_available(): + torch.cuda.synchronize(device=self.device_id) logger.debug("pulling kv_caches for %s finished", ok_reqs) if response.err_reqs: diff --git a/vllm/envs.py b/vllm/envs.py index 819c6d30cd3e..7598b8462df6 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -245,6 +245,9 @@ VLLM_ROCM_QUICK_REDUCE_MIN_SIZE_BYTES_MB: int | None = None VLLM_ROCM_QUICK_REDUCE_QUANTIZATION_MIN_SIZE_KB: int | None = None VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT: int = 480 + VLLM_MOONCAKE_MAX_TRANSFER_BYTES: int = 0 + VLLM_MOONCAKE_SYNC_AFTER_TRANSFER: bool = False + VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY: bool = False VLLM_ENABLE_CUDAGRAPH_GC: bool = False VLLM_LOOPBACK_IP: str = "" VLLM_ALLOW_CHUNKED_LOCAL_ATTN_WITH_HYBRID_KV_CACHE: bool = True @@ -1773,6 +1776,22 @@ def _resolve_rust_cli_path() -> str | None: "VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT": lambda: int( os.getenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "480") ), + # Split Mooncake transfer descriptors larger than this many bytes. + # Mitigates data-integrity issues with very large coalesced copies under + # concurrent PD load (vllm #42395). 0 disables splitting. + "VLLM_MOONCAKE_MAX_TRANSFER_BYTES": lambda: int( + os.getenv("VLLM_MOONCAKE_MAX_TRANSFER_BYTES", "0") + ), + # CUDA synchronize after each Mooncake batch transfer (producer and + # consumer). May reduce dst/src mismatches at the cost of throughput. + "VLLM_MOONCAKE_SYNC_AFTER_TRANSFER": lambda: ( + os.getenv("VLLM_MOONCAKE_SYNC_AFTER_TRANSFER", "0").lower() in ("1", "true") + ), + # Hash GPU source regions before/after Mooncake writes (debug only). + "VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY": lambda: ( + os.getenv("VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY", "0").lower() + in ("1", "true") + ), # If set, it means we pre-downloaded cubin files and flashinfer will # read the cubin files directly. "VLLM_HAS_FLASHINFER_CUBIN": lambda: bool( From b6d1445330287e75ba8954e758852d4231e82603 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Wed, 12 Aug 2026 03:58:23 +0000 Subject: [PATCH 03/30] Fix issues with latest main Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 10 +- tests/v1/ec_connector/integration/README.md | 366 +++++++++--------- .../run_epd_mooncake_ec_full_pipeline.sh | 14 +- .../unit/test_ec_mooncake_connector.py | 79 +++- .../unit/test_mooncake_connector.py | 18 +- .../ec_connector/mooncake_ec_connector.py | 103 +++-- 6 files changed, 351 insertions(+), 239 deletions(-) diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index 9a1bd517215b..75b7a3c277ac 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -457,7 +457,15 @@ async def forward_non_stream( async with decode_session.post( f"{d_url}/v1/chat/completions", json=req_data, headers=headers ) as resp: - resp.raise_for_status() + if resp.status >= 400: + detail = await resp.text() + logger.error( + "[%s] Decode request returned status %s: %s", + req_id, + resp.status, + detail, + ) + raise HTTPException(status_code=resp.status, detail=detail) out = await resp.json() _t3 = time.perf_counter() logger.info( diff --git a/tests/v1/ec_connector/integration/README.md b/tests/v1/ec_connector/integration/README.md index 376604dcb66e..5feee2ae11ee 100644 --- a/tests/v1/ec_connector/integration/README.md +++ b/tests/v1/ec_connector/integration/README.md @@ -1,183 +1,183 @@ -# EPD Correctness Test - -This test verifies that EPD (Encoder-Prefill-Decode) disaggregation produces identical outputs to a baseline single instance. - -## What It Tests - -- **Baseline**: Single vLLM instance serving a multimodal model -- **EPD (1E+1PD)**: 1 Encoder + 1 Prefill-Decode instance -- **Baseline (1P+1D)**: 1 Prefill + 1 Decode instance -- **EPD (1E+1P+1D)**: 1 Encoder + 1 Prefill + 1 Decode instance - -The test ensures that disaggregated encoding produces **identical** outputs to the baseline. - -Note that currently PD disaggregation set up may give slightly different results from a single instance. Therefore, we need the result from 1P+1D as the baseline for 1E+1P+1D - -Please refer to [Disaggregated Encoder Feature](../../../../docs/features/disagg_encoder.md) for the detailed explanation for the EPD features. - -## Files - -- `run_epd_correctness_test.sh` - Main test script (starts all instances and runs tests) -- `test_epd_correctness.py` - Python test script (compares outputs) - -## Usage - -### Multimodal Prompts (Default) - -```bash -cd vllm -./tests/v1/ec_connector/integration/run_epd_correctness_test.sh -``` - -This runs the test with actual multimodal (image) prompts. - -### Text-Only Prompts - -```bash -cd vllm -USE_MM_PROMPTS=0 ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh -``` - -This runs a quick test with text-only prompts to verify the setup works. - -### Custom Configuration - -```bash -# Use specific GPUs -GPU_E=0 GPU_PD=1 GPU_P=1 GPU_D=2 bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh - -# Use specific ports -ENDPOINT_PORT=10001 bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh - -# Use specific model -MODEL="Qwen/Qwen2.5-VL-3B-Instruct" bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh - -# Use specific storage path -EC_SHARED_STORAGE_PATH="/tmp/my_ec_cache" bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh -``` - -## How It Works - -### Step 1: Baseline - -1. Start single vLLM instance on GPU -2. Run test prompts (multimodal or text-only) -3. Save outputs to `.vllm_epd_baseline.txt` -4. Shutdown instance - -### Step 2: EPD (1E + 1PD) - -1. Clear encoder cache storage -2. Start instances and proxy -3. Run same test prompts -4. Assert outputs match baseline exactly -5. Shutdown instances - -### Step 3: EPD (1E + 1P + 1D) - -1. Clear encoder cache storage -2. Start instances and proxy -3. Run same test prompts -4. Assert outputs match baseline exactly -5. Shutdown instances - -## Test Scenarios - -### Multimodal Prompts (--use_mm_prompts) - -Tests encoder cache transfer: - -- Single image query -- Multiple images in one request -- Mixed image and text -- Image with detailed questions - -### Text-Only Prompts (default) - -Quick sanity check: - -- Simple text queries -- Text-only explanations -- Verifies proxy routing works - -## Expected Behavior - -### ✅ Test Passes When - -- All disagg outputs match baseline outputs exactly -- No errors during instance startup -- Encoder cache is properly saved and loaded -- Proxy correctly routes requests - -### ❌ Test Fails When - -- Outputs differ between baseline and disagg -- Server startup fails -- Encoder cache not found (should fall back to local execution) -- Proxy routing errors - -## Notes - -- The test uses deterministic generation (`temperature=0.0`, `seed=42`) -- Encoder cache should enable exact output reproduction -- Test cleans up all instances and cache files after completion -- Safe to run multiple times (idempotent) -- We setup the PD disagg part with NixlConnector. Please read details about EPD in `examples/disaggregated/disaggregated_encoder/README.md` - -## ECMooncakeConnector (TransferEngine) smoke test - -Two-process transfer over Mooncake (no full vLLM serve, no HF model download): - -```bash -cd vllm -PYTHONPATH=. python tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py -``` - -Requires: **2+ CUDA GPUs**, `mooncake-transfer-engine`, `pyzmq`, `httpx`, `fastapi`, `uvicorn`. -Optional: `MOONCAKE_EC_PROTOCOL=rdma` or `=tcp` (default in test mocks is `tcp`) to match your cluster. - -## Requirements - -- Multiple GPUs (3 for 1E+1P+1D, 2 for 1E+1PD, 1 for baseline) - - 1E+1P+1D is runnable with 2 GPU by assign E and P on the same GPU now. -- Multimodal model (e.g., Qwen2.5-VL-3B-Instruct) -- Internet access (for accessing vllm test images) - -## Debugging - -### Check Logs - -Logs and baseline output are saved in `/tmp/` by default. -Can be customized by changing the environment variables. - -### Check Encoder Cache - -```bash -# Verify cache files are created -ls -la $EC_SHARED_STORAGE_PATH/ - -# Should see directories with mm_hash names -# Each containing encoder_cache.safetensors -``` - -### Manual Testing - -Run individual components: - -```bash -# Baseline only -python test_epd_correctness.py \ - --service_url http://localhost:8000 \ - --model_name Qwen/Qwen2.5-VL-3B-Instruct \ - --mode baseline \ - --baseline_file test_output.txt \ - --use_mm_prompts - -# Disagg only (requires baseline output file!) -python test_epd_correctness.py \ - --service_url http://localhost:8000 \ - --model_name Qwen/Qwen2.5-VL-3B-Instruct \ - --mode disagg \ - --baseline_file test_output.txt \ - --use_mm_prompts -``` +# EPD Correctness Test + +This test verifies that EPD (Encoder-Prefill-Decode) disaggregation produces identical outputs to a baseline single instance. + +## What It Tests + +- **Baseline**: Single vLLM instance serving a multimodal model +- **EPD (1E+1PD)**: 1 Encoder + 1 Prefill-Decode instance +- **Baseline (1P+1D)**: 1 Prefill + 1 Decode instance +- **EPD (1E+1P+1D)**: 1 Encoder + 1 Prefill + 1 Decode instance + +The test ensures that disaggregated encoding produces **identical** outputs to the baseline. + +Note that currently PD disaggregation set up may give slightly different results from a single instance. Therefore, we need the result from 1P+1D as the baseline for 1E+1P+1D + +Please refer to [Disaggregated Encoder Feature](../../../../docs/features/disagg_encoder.md) for the detailed explanation for the EPD features. + +## Files + +- `run_epd_correctness_test.sh` - Main test script (starts all instances and runs tests) +- `test_epd_correctness.py` - Python test script (compares outputs) + +## Usage + +### Multimodal Prompts (Default) + +```bash +cd vllm +./tests/v1/ec_connector/integration/run_epd_correctness_test.sh +``` + +This runs the test with actual multimodal (image) prompts. + +### Text-Only Prompts + +```bash +cd vllm +USE_MM_PROMPTS=0 ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh +``` + +This runs a quick test with text-only prompts to verify the setup works. + +### Custom Configuration + +```bash +# Use specific GPUs +GPU_E=0 GPU_PD=1 GPU_P=1 GPU_D=2 bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh + +# Use specific ports +ENDPOINT_PORT=10001 bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh + +# Use specific model +MODEL="Qwen/Qwen2.5-VL-3B-Instruct" bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh + +# Use specific storage path +EC_SHARED_STORAGE_PATH="/tmp/my_ec_cache" bash ./tests/v1/ec_connector/integration/run_epd_correctness_test.sh +``` + +## How It Works + +### Step 1: Baseline + +1. Start single vLLM instance on GPU +2. Run test prompts (multimodal or text-only) +3. Save outputs to `.vllm_epd_baseline.txt` +4. Shutdown instance + +### Step 2: EPD (1E + 1PD) + +1. Clear encoder cache storage +2. Start instances and proxy +3. Run same test prompts +4. Assert outputs match baseline exactly +5. Shutdown instances + +### Step 3: EPD (1E + 1P + 1D) + +1. Clear encoder cache storage +2. Start instances and proxy +3. Run same test prompts +4. Assert outputs match baseline exactly +5. Shutdown instances + +## Test Scenarios + +### Multimodal Prompts (--use_mm_prompts) + +Tests encoder cache transfer: + +- Single image query +- Multiple images in one request +- Mixed image and text +- Image with detailed questions + +### Text-Only Prompts (default) + +Quick sanity check: + +- Simple text queries +- Text-only explanations +- Verifies proxy routing works + +## Expected Behavior + +### ✅ Test Passes When + +- All disagg outputs match baseline outputs exactly +- No errors during instance startup +- Encoder cache is properly saved and loaded +- Proxy correctly routes requests + +### ❌ Test Fails When + +- Outputs differ between baseline and disagg +- Server startup fails +- Encoder cache not found (should fall back to local execution) +- Proxy routing errors + +## Notes + +- The test uses deterministic generation (`temperature=0.0`, `seed=42`) +- Encoder cache should enable exact output reproduction +- Test cleans up all instances and cache files after completion +- Safe to run multiple times (idempotent) +- We setup the PD disagg part with NixlConnector. Please read details about EPD in `examples/disaggregated/disaggregated_encoder/README.md` + +## ECMooncakeConnector (TransferEngine) smoke test + +Two-process transfer over Mooncake (no full vLLM serve, no HF model download): + +```bash +cd vllm +PYTHONPATH=. python tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py +``` + +Requires: **2+ CUDA GPUs**, `mooncake-transfer-engine`, `pyzmq`, `httpx`, `fastapi`, `uvicorn`. +Optional: `MOONCAKE_EC_PROTOCOL=rdma` or `=tcp` (default in test mocks is `tcp`) to match your cluster. + +## Requirements + +- Multiple GPUs (3 for 1E+1P+1D, 2 for 1E+1PD, 1 for baseline) + - 1E+1P+1D is runnable with 2 GPU by assign E and P on the same GPU now. +- Multimodal model (e.g., Qwen2.5-VL-3B-Instruct) +- Internet access (for accessing vllm test images) + +## Debugging + +### Check Logs + +Logs and baseline output are saved in `/tmp/` by default. +Can be customized by changing the environment variables. + +### Check Encoder Cache + +```bash +# Verify cache files are created +ls -la $EC_SHARED_STORAGE_PATH/ + +# Should see directories with mm_hash names +# Each containing encoder_cache.safetensors +``` + +### Manual Testing + +Run individual components: + +```bash +# Baseline only +python test_epd_correctness.py \ + --service_url http://localhost:8000 \ + --model_name Qwen/Qwen2.5-VL-3B-Instruct \ + --mode baseline \ + --baseline_file test_output.txt \ + --use_mm_prompts + +# Disagg only (requires baseline output file!) +python test_epd_correctness.py \ + --service_url http://localhost:8000 \ + --model_name Qwen/Qwen2.5-VL-3B-Instruct \ + --mode disagg \ + --baseline_file test_output.txt \ + --use_mm_prompts +``` diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh index eb031c044bf8..d68040baa098 100755 --- a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -20,7 +20,7 @@ # TIMEOUT_SECONDS wait_for_server timeout (default 1200) # SKIP_BASELINE set to 1 to reuse existing BASELINE_FILE -set -u +set -euo pipefail GIT_ROOT=$(git rev-parse --show-toplevel) cd "$GIT_ROOT" || exit 1 @@ -60,7 +60,7 @@ else VLLM_SERVE=(python -m vllm.entrypoints.cli.main serve) fi -ENC_EC_JSON=$(python3 < pending tensor layout from registry self._pending_specs: dict[str, ECMooncakeLoadSpec] = {} @@ -175,12 +179,15 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._tensor_lock = threading.Lock() self._producer_services_started = False - if role == ECConnectorRole.SCHEDULER and self.is_consumer: - if not self._remote_registry_url: - raise ValueError( - "ec_consumer with ECMooncakeConnector requires " - "ec_connector_extra_config['remote_registry_url']." - ) + if ( + role == ECConnectorRole.SCHEDULER + and self.is_consumer + and not self._remote_registry_url + ): + raise ValueError( + "ec_consumer with ECMooncakeConnector requires " + "ec_connector_extra_config['remote_registry_url']." + ) def _ensure_engine(self) -> TransferEngine: if self._engine is None: @@ -236,10 +243,10 @@ def loop() -> None: if ( ret == 0 and envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER - and torch.cuda.is_available() + and torch.accelerator.is_available() and tensor.is_cuda ): - torch.cuda.synchronize(device=tensor.device) + torch.accelerator.synchronize() sock.send_json({"ok": ret == 0, "mooncake_ret": int(ret)}) except Exception as e: logger.exception("EC Mooncake pull handler error: %s", e) @@ -249,7 +256,9 @@ def loop() -> None: break self._zmq_ctx = zmq.Context() - self._zmq_thread = threading.Thread(target=loop, name="ec-mooncake-zmq", daemon=True) + self._zmq_thread = threading.Thread( + target=loop, name="ec-mooncake-zmq", daemon=True + ) self._zmq_thread.start() while self._zmq_listen_addr is None: time.sleep(0.01) @@ -262,7 +271,9 @@ def _ensure_producer_services(self) -> None: self._ensure_engine() self._start_producer_zmq_listener() if is_local_first_rank(): - self._registry = ECMooncakeRegistryServer("0.0.0.0", self._registry_http_port) + self._registry = ECMooncakeRegistryServer( + "0.0.0.0", self._registry_http_port + ) self._registry.start() self._producer_services_started = True @@ -273,13 +284,11 @@ def start_load_caches( assert isinstance(metadata, ECMooncakeConnectorMetadata) eng = self._ensure_engine() raw_buf = self._ec_cfg.ec_buffer_device - buf = ( - raw_buf.lower() - if isinstance(raw_buf, str) and raw_buf - else "cuda" - ) - if buf == "cuda" and not torch.cuda.is_available(): - raise RuntimeError("ECMooncakeConnector requires CUDA for ec_buffer_device=cuda") + buf = raw_buf.lower() if isinstance(raw_buf, str) and raw_buf else "cuda" + if buf == "cuda" and not torch.accelerator.is_available(): + raise RuntimeError( + "ECMooncakeConnector requires CUDA for ec_buffer_device=cuda" + ) device = torch.device(buf) for spec in metadata.loads: @@ -313,11 +322,8 @@ def start_load_caches( ctx.term() if not resp.get("ok"): raise RuntimeError(f"EC Mooncake pull failed: {resp}") - if ( - envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER - and device.type == "cuda" - ): - torch.cuda.synchronize(device=device) + if envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER and device.type == "cuda": + torch.accelerator.synchronize() encoder_cache[spec.mm_hash] = t logger.debug("Loaded EC tensor for mm_hash=%s via Mooncake", spec.mm_hash) @@ -354,7 +360,9 @@ def has_cache_item(self, identifier: str) -> bool: try: r = httpx.get(url, timeout=5.0) except httpx.HTTPError as e: - logger.warning("EC Mooncake registry query failed for %s: %s", identifier, e) + logger.warning( + "EC Mooncake registry query failed for %s: %s", identifier, e + ) return False if r.status_code != 200: return False @@ -373,6 +381,8 @@ def has_cache_item(self, identifier: str) -> bool: return True def update_state_after_alloc(self, request: Any, index: int) -> None: + if not self.is_consumer: + return mm_hash = request.mm_features[index].identifier num_encoder_token = request.get_num_encoder_embeds(index) self._mm_datas_need_loads[mm_hash] = num_encoder_token @@ -400,6 +410,47 @@ def build_connector_meta( self._mm_datas_need_loads.clear() return meta + def _placeholder_metadata_fields(self, modality: str) -> set[str]: + if modality in self._metadata_fields_cache: + return self._metadata_fields_cache[modality] + + fields: set[str] = set() + try: + from vllm.multimodal import MULTIMODAL_REGISTRY + + info = MULTIMODAL_REGISTRY.create_processor(self._model_config).info + fields = info.data_parser.placeholder_metadata_fields(modality) + except Exception: + logger.warning( + "Could not determine the placeholder metadata fields for " + "modality %s; the consumer will preprocess the media itself.", + modality, + exc_info=True, + ) + + self._metadata_fields_cache[modality] = fields + return fields + + def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: + if not self.is_producer: + return False, None + + items = [] + for feature in request.mm_features: + metadata = {} + if feature.data is not None: + wanted = self._placeholder_metadata_fields(feature.modality) + metadata = { + key: value.tolist() + for key, value in feature.data.get_data().items() + if key in wanted and isinstance(value, torch.Tensor) + } + items.append({"mm_hash": feature.identifier, **metadata}) + + if not items: + return False, None + return False, {"ec_items": items} + def __del__(self) -> None: try: if self._registry is not None: From f942d70bde9eeda4a76ecd2ca0c00c011eabf568 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Wed, 12 Aug 2026 23:06:07 +0000 Subject: [PATCH 04/30] [EC] Manage Mooncake registration lifetime Signed-off-by: Tianyu Guo --- .../test_ec_mooncake_transfer_e2e.py | 25 +- .../unit/test_ec_mooncake_connector.py | 148 +++++++- .../ec_connector/mooncake_ec_connector.py | 318 +++++++++++++----- 3 files changed, 400 insertions(+), 91 deletions(-) diff --git a/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py b/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py index e849dfb5144e..b8ed8446c8ee 100644 --- a/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py +++ b/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py @@ -8,10 +8,11 @@ Requires: 2+ CUDA GPUs, mooncake-transfer-engine, pyzmq, httpx, fastapi, uvicorn. Protocol: ``mooncake_protocol`` defaults to ``tcp`` in mocks unless you set -``MOONCAKE_EC_PROTOCOL=rdma`` (matches ``ec_connector_extra_config.mooncake_protocol``). +``MOONCAKE_EC_PROTOCOL=rdma`` (matches the connector protocol configuration). Example RDMA run:: - MOONCAKE_EC_PROTOCOL=rdma PYTHONPATH=. python tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py + MOONCAKE_EC_PROTOCOL=rdma PYTHONPATH=. python \ + tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py """ from __future__ import annotations @@ -19,6 +20,7 @@ import multiprocessing as mp import os import time +from typing import Any from unittest.mock import Mock import pytest @@ -61,6 +63,7 @@ def _mock_vllm_producer(registry_port: int) -> Mock: cfg.ec_transfer_config.is_ec_producer = True cfg.ec_transfer_config.is_ec_consumer = False cfg.ec_transfer_config.ec_buffer_device = "cuda" + cfg.ec_transfer_config.ec_buffer_size = 1e9 cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), "registry_http_port": registry_port, @@ -77,6 +80,7 @@ def _mock_vllm_consumer() -> Mock: cfg.ec_transfer_config.is_ec_producer = False cfg.ec_transfer_config.is_ec_consumer = True cfg.ec_transfer_config.ec_buffer_device = "cuda" + cfg.ec_transfer_config.ec_buffer_size = 1e9 cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), "remote_registry_url": "http://unused-on-worker", @@ -87,12 +91,11 @@ def _mock_vllm_consumer() -> Mock: def _producer_entry( mm_hash: str, registry_port: int, - ready: mp.Queue, - done: mp.Event, - barrier: mp.Barrier, + ready: Any, + done: Any, + barrier: Any, ) -> None: os.environ["CUDA_VISIBLE_DEVICES"] = "0" - torch.cuda.init() cfg = _mock_vllm_producer(registry_port) conn = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) torch.manual_seed(12345) @@ -108,11 +111,10 @@ def _producer_entry( def _consumer_entry( mm_hash: str, registry_url: str, - barrier: mp.Barrier, - result_queue: mp.Queue, + barrier: Any, + result_queue: Any, ) -> None: os.environ["CUDA_VISIBLE_DEVICES"] = "1" - torch.cuda.init() barrier.wait(timeout=120) import httpx @@ -136,6 +138,7 @@ def _consumer_entry( shape=tuple(int(x) for x in data["shape"]), dtype=str(data["dtype"]), producer_zmq=str(data["producer_zmq"]), + lease_id=str(data["lease_id"]), ) meta = ECMooncakeConnectorMetadata() meta.add_load(spec) @@ -159,7 +162,7 @@ def _consumer_entry( @pytest.mark.skipif( - torch.cuda.device_count() < 2, + torch.accelerator.device_count() < 2, reason="Requires at least 2 CUDA devices", ) def test_ec_mooncake_two_process_transfer(): @@ -198,7 +201,7 @@ def test_ec_mooncake_two_process_transfer(): def _main() -> None: - if torch.cuda.device_count() < 2: + if torch.accelerator.device_count() < 2: raise SystemExit("Need at least 2 CUDA devices for this e2e test.") test_ec_mooncake_two_process_transfer() print("ECMooncake two-process transfer e2e: PASSED") diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index d9e79ec49b2f..cf4af7eb8a68 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -31,7 +31,10 @@ class CopyingFakeTransferEngine: def __init__(self, *args, **kwargs): - pass + self.registered: set[int] = set() + self.register_calls: list[list[int]] = [] + self.unregister_calls: list[int] = [] + self.batch_unregister_calls: list[list[int]] = [] def initialize(self, local_hostname, metadata_server, protocol, device_name) -> int: return 0 @@ -47,6 +50,21 @@ def batch_transfer_sync_write( return 0 def batch_register_memory(self, buffer_addresses, capacities) -> int: + addresses = [int(addr) for addr in buffer_addresses] + self.register_calls.append(addresses) + self.registered.update(addresses) + return 0 + + def unregister_memory(self, buffer_address) -> int: + address = int(buffer_address) + self.unregister_calls.append(address) + self.registered.discard(address) + return 0 + + def batch_unregister_memory(self, buffer_addresses) -> int: + addresses = [int(addr) for addr in buffer_addresses] + self.batch_unregister_calls.append(addresses) + self.registered.difference_update(addresses) return 0 @@ -68,6 +86,7 @@ def mock_vllm_config_producer(): config.ec_transfer_config.is_ec_producer = True config.ec_transfer_config.is_ec_consumer = False config.ec_transfer_config.ec_buffer_device = "cuda" + config.ec_transfer_config.ec_buffer_size = 1e9 config.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", "registry_http_port": 19018, @@ -85,6 +104,7 @@ def mock_vllm_config_consumer(): config.ec_transfer_config.is_ec_producer = False config.ec_transfer_config.is_ec_consumer = True config.ec_transfer_config.ec_buffer_device = "cuda" + config.ec_transfer_config.ec_buffer_size = 1e9 config.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", "remote_registry_url": "http://127.0.0.1:19018", @@ -130,7 +150,10 @@ def test_publish_and_lookup(self): registry.publish("hash_a", payload) r = httpx.get(f"http://127.0.0.1:{port}/ec/info/hash_a", timeout=2.0) assert r.status_code == 200 - assert r.json() == payload + data = r.json() + lease_id = data.pop("lease_id") + assert data == payload + assert registry.consume_lease("hash_a", lease_id) r404 = httpx.get(f"http://127.0.0.1:{port}/ec/info/missing", timeout=2.0) assert r404.status_code == 404 finally: @@ -246,6 +269,7 @@ def test_build_connector_meta_clears_pending( shape=(2, 4), dtype="float32", producer_zmq="tcp://127.0.0.1:1", + lease_id="lease", ) scheduler._mm_datas_need_loads[mm_hash] = 100 meta = scheduler.build_connector_meta(Mock(spec=SchedulerOutput)) @@ -333,6 +357,7 @@ def test_single_process_save_and_load(self, mock_vllm_config_producer): consumer_cfg.ec_transfer_config.is_ec_producer = False consumer_cfg.ec_transfer_config.is_ec_consumer = True consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 1e9 consumer_cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", } @@ -344,6 +369,7 @@ def test_single_process_save_and_load(self, mock_vllm_config_producer): shape=tuple(int(x) for x in data["shape"]), dtype=str(data["dtype"]), producer_zmq=str(data["producer_zmq"]), + lease_id=str(data["lease_id"]), ) meta = ECMooncakeConnectorMetadata() meta.add_load(spec) @@ -352,6 +378,124 @@ def test_single_process_save_and_load(self, mock_vllm_config_producer): consumer.start_load_caches(loaded) assert mm_hash in loaded assert torch.allclose(loaded[mm_hash].cpu(), source.cpu()) + consumer_engine = consumer._engine + assert isinstance(consumer_engine, CopyingFakeTransferEngine) + assert consumer_engine.registered == set() + assert consumer_engine.unregister_calls == [loaded[mm_hash].data_ptr()] + producer.shutdown() + consumer.shutdown() + + def test_producer_evicts_lru_registration_at_capacity( + self, mock_vllm_config_producer + ): + port = _find_free_port() + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 32 + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "registry_http_port" + ] = port + first = torch.randn(8) + second = torch.randn(8) + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + try: + producer.save_caches({"first": first}, "first") + producer.save_caches({"second": second}, "second") + + engine = producer._engine + assert isinstance(engine, CopyingFakeTransferEngine) + assert list(producer._tensor_by_hash) == ["second"] + assert producer._registered_bytes == second.nbytes + assert first.data_ptr() in engine.unregister_calls + assert second.data_ptr() in engine.registered + + base_url = f"http://127.0.0.1:{port}/ec/info" + assert httpx.get(f"{base_url}/first").status_code == 404 + assert httpx.get(f"{base_url}/second").status_code == 200 + finally: + producer.shutdown() + + def test_producer_does_not_evict_in_flight_registration( + self, mock_vllm_config_producer + ): + port = _find_free_port() + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 32 + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "registry_http_port" + ] = port + first = torch.randn(8) + second = torch.randn(8) + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + try: + producer.save_caches({"first": first}, "first") + producer._tensor_by_hash["first"].in_flight = 1 + with pytest.raises(RuntimeError, match="no evictable"): + producer.save_caches({"second": second}, "second") + assert list(producer._tensor_by_hash) == ["first"] + finally: + producer._tensor_by_hash["first"].in_flight = 0 + producer.shutdown() + + def test_producer_does_not_evict_leased_registration( + self, mock_vllm_config_producer + ): + port = _find_free_port() + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 32 + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "registry_http_port" + ] = port + first = torch.randn(8) + second = torch.randn(8) + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + try: + producer.save_caches({"first": first}, "first") + response = httpx.get(f"http://127.0.0.1:{port}/ec/info/first").json() + + with pytest.raises(RuntimeError, match="no evictable"): + producer.save_caches({"second": second}, "second") + + assert producer._registry is not None + assert producer._registry.consume_lease("first", response["lease_id"]) + producer.save_caches({"second": second}, "second") + assert list(producer._tensor_by_hash) == ["second"] + finally: + producer.shutdown() + + def test_shutdown_unregisters_all_producer_tensors(self, mock_vllm_config_producer): + port = _find_free_port() + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "registry_http_port" + ] = port + tensor = torch.randn(8) + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + producer.save_caches({"hash": tensor}, "hash") + engine = producer._engine + assert isinstance(engine, CopyingFakeTransferEngine) + + producer.shutdown() + + assert engine.batch_unregister_calls == [[tensor.data_ptr()]] + assert engine.registered == set() + assert producer._tensor_by_hash == {} + assert producer._registered_bytes == 0 def test_producer_scheduler_has_cache_item_false( self, mock_vllm_config_producer, mock_request_with_3_mm diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 94630889c5a5..5e5f13942b82 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -13,6 +13,9 @@ import json import threading import time +import uuid +from collections import OrderedDict +from contextlib import suppress from dataclasses import dataclass, field from typing import Any @@ -36,6 +39,8 @@ logger = init_logger(__name__) +_LEASE_TTL_SECONDS = 300 + _MOONCAKE_IMPORT_ERROR: ImportError | None try: from mooncake.engine import TransferEngine @@ -56,6 +61,7 @@ class ECMooncakeLoadSpec: shape: tuple[int, ...] dtype: str producer_zmq: str + lease_id: str @dataclass @@ -68,6 +74,12 @@ def add_load(self, spec: ECMooncakeLoadSpec) -> None: self.loads.append(spec) +@dataclass +class _RegisteredTensor: + tensor: torch.Tensor + in_flight: int = 0 + + class ECMooncakeRegistryServer: """Lightweight HTTP registry on the producer for remote has_cache_item / info.""" @@ -75,6 +87,7 @@ def __init__(self, host: str, port: int): self.host = host self.port = port self._entries: dict[str, dict[str, Any]] = {} + self._leases: dict[str, dict[str, float]] = {} self._lock = threading.Lock() self.app = FastAPI() self._register_routes() @@ -86,9 +99,13 @@ def _register_routes(self) -> None: async def ec_info(mm_hash: str) -> dict[str, Any]: with self._lock: data = self._entries.get(mm_hash) - if data is None: - raise HTTPException(status_code=404, detail="unknown mm_hash") - return data + if data is None: + raise HTTPException(status_code=404, detail="unknown mm_hash") + lease_id = uuid.uuid4().hex + self._leases.setdefault(mm_hash, {})[lease_id] = ( + time.monotonic() + _LEASE_TTL_SECONDS + ) + return {**data, "lease_id": lease_id} def start(self) -> None: if self.server_thread is not None: @@ -116,9 +133,31 @@ def publish(self, mm_hash: str, payload: dict[str, Any]) -> None: with self._lock: self._entries[mm_hash] = payload - def unpublish(self, mm_hash: str) -> None: + def unpublish(self, mm_hash: str) -> dict[str, Any] | None: + with self._lock: + self._leases.pop(mm_hash, None) + return self._entries.pop(mm_hash, None) + + def unpublish_if_unleased(self, mm_hash: str) -> tuple[bool, dict[str, Any] | None]: + with self._lock: + leases = self._leases.get(mm_hash, {}) + now = time.monotonic() + leases = {lease: expiry for lease, expiry in leases.items() if expiry > now} + if leases: + self._leases[mm_hash] = leases + return False, None + self._leases.pop(mm_hash, None) + return True, self._entries.pop(mm_hash, None) + + def consume_lease(self, mm_hash: str, lease_id: str) -> bool: with self._lock: - self._entries.pop(mm_hash, None) + leases = self._leases.get(mm_hash) + if leases is None: + return False + expiry = leases.pop(lease_id, 0) + if not leases: + self._leases.pop(mm_hash, None) + return expiry > time.monotonic() class ECMooncakeConnector(ECConnectorBase): @@ -161,6 +200,9 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._protocol: str = self._extra.get("mooncake_protocol", "rdma") self._remote_registry_url: str | None = self._extra.get("remote_registry_url") self._registry_http_port: int = int(self._extra.get("registry_http_port", 9018)) + self._registered_capacity = int(self._ec_cfg.ec_buffer_size) + if self._registered_capacity <= 0: + raise ValueError("ECMooncakeConnector requires ec_buffer_size > 0.") self._model_config = vllm_config.model_config self._metadata_fields_cache: dict[str, set[str]] = {} @@ -175,9 +217,13 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._zmq_listen_addr: str | None = None self._zmq_thread: threading.Thread | None = None self._zmq_ctx: zmq.Context | None = None - self._tensor_by_hash: dict[str, torch.Tensor] = {} + self._zmq_stop = threading.Event() + self._tensor_by_hash: OrderedDict[str, _RegisteredTensor] = OrderedDict() + self._registered_bytes = 0 + self._pending_unregister: dict[int, torch.Tensor] = {} self._tensor_lock = threading.Lock() self._producer_services_started = False + self._shutdown = False if ( role == ECConnectorRole.SCHEDULER @@ -203,6 +249,56 @@ def _ensure_engine(self) -> TransferEngine: ) return self._engine + def _unregister_memory(self, tensor: torch.Tensor) -> bool: + assert self._engine is not None + ret = self._engine.unregister_memory(tensor.data_ptr()) + if ret != 0: + logger.error( + "Mooncake EC memory unregistration failed for address %d: %d", + tensor.data_ptr(), + ret, + ) + self._pending_unregister[tensor.data_ptr()] = tensor + return False + self._pending_unregister.pop(tensor.data_ptr(), None) + return True + + def _release_tensor_locked(self, mm_hash: str) -> bool: + entry = self._tensor_by_hash.get(mm_hash) + if entry is None: + return True + if entry.in_flight: + return False + payload = None + if self._registry is not None: + released, payload = self._registry.unpublish_if_unleased(mm_hash) + if not released: + return False + if not self._unregister_memory(entry.tensor): + if self._registry is not None and payload is not None: + self._registry.publish(mm_hash, payload) + return False + self._tensor_by_hash.pop(mm_hash) + self._registered_bytes -= entry.tensor.nbytes + return True + + def _make_registration_space_locked(self, nbytes: int) -> None: + if nbytes > self._registered_capacity: + raise RuntimeError( + f"Encoder cache tensor ({nbytes} bytes) exceeds ec_buffer_size " + f"({self._registered_capacity} bytes)." + ) + while self._registered_bytes + nbytes > self._registered_capacity: + evicted = False + for mm_hash, entry in list(self._tensor_by_hash.items()): + if not entry.in_flight and self._release_tensor_locked(mm_hash): + evicted = True + break + if not evicted: + raise RuntimeError( + "ECMooncakeConnector has no evictable registered memory." + ) + def _start_producer_zmq_listener(self) -> None: if self._zmq_thread is not None: return @@ -210,50 +306,73 @@ def _start_producer_zmq_listener(self) -> None: def loop() -> None: assert self._zmq_ctx is not None sock = self._zmq_ctx.socket(zmq.REP) + sock.setsockopt(zmq.RCVTIMEO, 100) port = sock.bind_to_random_port(f"tcp://{self._hostname}") self._zmq_listen_addr = f"tcp://{self._hostname}:{port}" logger.info("EC Mooncake pull listener at %s", self._zmq_listen_addr) eng = self._ensure_engine() - while True: - try: - raw = sock.recv() - except zmq.ContextTerminated: - break - try: - req = json.loads(raw.decode("utf-8")) - if req.get("op") != "pull": - sock.send_json({"ok": False, "err": "unknown op"}) - continue - mm_hash = req["mm_hash"] - dst_session = req["dst_session"] - dst_ptr = int(req["dst_ptr"]) - nbytes = int(req["nbytes"]) - with self._tensor_lock: - tensor = self._tensor_by_hash.get(mm_hash) - if tensor is None: - sock.send_json({"ok": False, "err": "unknown mm_hash"}) - continue - src_ptr = tensor.data_ptr() - if tensor.nbytes != nbytes: - sock.send_json({"ok": False, "err": "size mismatch"}) - continue - ret = eng.batch_transfer_sync_write( - dst_session, [src_ptr], [dst_ptr], [nbytes] - ) - if ( - ret == 0 - and envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER - and torch.accelerator.is_available() - and tensor.is_cuda - ): - torch.accelerator.synchronize() - sock.send_json({"ok": ret == 0, "mooncake_ret": int(ret)}) - except Exception as e: - logger.exception("EC Mooncake pull handler error: %s", e) + try: + while True: try: - sock.send_json({"ok": False, "err": str(e)}) - except zmq.ZMQError: + raw = sock.recv() + except zmq.Again: + if self._zmq_stop.is_set(): + break + continue + except zmq.ContextTerminated: break + try: + req = json.loads(raw.decode("utf-8")) + if req.get("op") != "pull": + sock.send_json({"ok": False, "err": "unknown op"}) + continue + mm_hash = req["mm_hash"] + dst_session = req["dst_session"] + dst_ptr = int(req["dst_ptr"]) + nbytes = int(req["nbytes"]) + lease_id = str(req["lease_id"]) + with self._tensor_lock: + entry = self._tensor_by_hash.get(mm_hash) + if ( + entry is not None + and self._registry is not None + and self._registry.consume_lease(mm_hash, lease_id) + ): + entry.in_flight += 1 + self._tensor_by_hash.move_to_end(mm_hash) + else: + entry = None + if entry is None: + sock.send_json({"ok": False, "err": "invalid lease"}) + continue + try: + tensor = entry.tensor + src_ptr = tensor.data_ptr() + if tensor.nbytes != nbytes: + sock.send_json({"ok": False, "err": "size mismatch"}) + continue + ret = eng.batch_transfer_sync_write( + dst_session, [src_ptr], [dst_ptr], [nbytes] + ) + if ( + ret == 0 + and envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER + and torch.accelerator.is_available() + and tensor.is_cuda + ): + torch.accelerator.synchronize() + finally: + with self._tensor_lock: + entry.in_flight -= 1 + sock.send_json({"ok": ret == 0, "mooncake_ret": int(ret)}) + except Exception as e: + logger.exception("EC Mooncake pull handler error: %s", e) + try: + sock.send_json({"ok": False, "err": str(e)}) + except zmq.ZMQError: + break + finally: + sock.close(linger=0) self._zmq_ctx = zmq.Context() self._zmq_thread = threading.Thread( @@ -303,27 +422,34 @@ def start_load_caches( raise RuntimeError( "Mooncake EC batch_register_memory failed on consumer." ) - pull = { - "op": "pull", - "mm_hash": spec.mm_hash, - "dst_session": f"{self._hostname}:{eng.get_rpc_port()}", - "dst_ptr": t.data_ptr(), - "nbytes": t.nbytes, - } - ctx = zmq.Context() - sock = ctx.socket(zmq.REQ) - sock.setsockopt(zmq.RCVTIMEO, 120_000) - sock.connect(spec.producer_zmq) try: - sock.send_json(pull) - resp = sock.recv_json() + pull = { + "op": "pull", + "mm_hash": spec.mm_hash, + "dst_session": f"{self._hostname}:{eng.get_rpc_port()}", + "dst_ptr": t.data_ptr(), + "nbytes": t.nbytes, + "lease_id": spec.lease_id, + } + ctx = zmq.Context() + sock = ctx.socket(zmq.REQ) + sock.setsockopt(zmq.RCVTIMEO, 120_000) + sock.connect(spec.producer_zmq) + try: + sock.send_json(pull) + resp = sock.recv_json() + finally: + sock.close(linger=0) + ctx.term() + if not resp.get("ok"): + raise RuntimeError(f"EC Mooncake pull failed: {resp}") + if envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER and device.type == "cuda": + torch.accelerator.synchronize() finally: - sock.close(linger=0) - ctx.term() - if not resp.get("ok"): - raise RuntimeError(f"EC Mooncake pull failed: {resp}") - if envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER and device.type == "cuda": - torch.accelerator.synchronize() + if not self._unregister_memory(t): + logger.warning( + "Keeping EC tensor alive after Mooncake unregistration failure" + ) encoder_cache[spec.mm_hash] = t logger.debug("Loaded EC tensor for mm_hash=%s via Mooncake", spec.mm_hash) @@ -335,12 +461,6 @@ def save_caches( self._ensure_producer_services() tensor = encoder_cache[mm_hash] eng = self._ensure_engine() - ret = eng.batch_register_memory([tensor.data_ptr()], [tensor.nbytes]) - if ret != 0: - raise RuntimeError("Mooncake EC batch_register_memory failed on producer.") - with self._tensor_lock: - self._tensor_by_hash[mm_hash] = tensor - dtype_str = str(tensor.dtype).split(".")[-1] payload = { "nbytes": tensor.nbytes, @@ -348,8 +468,20 @@ def save_caches( "dtype": dtype_str, "producer_zmq": self._zmq_listen_addr, } - if self._registry is not None: - self._registry.publish(mm_hash, payload) + with self._tensor_lock: + if mm_hash in self._tensor_by_hash: + self._tensor_by_hash.move_to_end(mm_hash) + return + self._make_registration_space_locked(tensor.nbytes) + ret = eng.batch_register_memory([tensor.data_ptr()], [tensor.nbytes]) + if ret != 0: + raise RuntimeError( + "Mooncake EC batch_register_memory failed on producer." + ) + self._tensor_by_hash[mm_hash] = _RegisteredTensor(tensor) + self._registered_bytes += tensor.nbytes + if self._registry is not None: + self._registry.publish(mm_hash, payload) logger.debug("Published EC tensor mm_hash=%s to registry", mm_hash) def has_cache_item(self, identifier: str) -> bool: @@ -377,6 +509,7 @@ def has_cache_item(self, identifier: str) -> bool: shape=tuple(int(x) for x in data["shape"]), dtype=str(data["dtype"]), producer_zmq=str(zmq_addr), + lease_id=str(data["lease_id"]), ) return True @@ -404,6 +537,7 @@ def build_connector_meta( shape=spec.shape, dtype=spec.dtype, producer_zmq=spec.producer_zmq, + lease_id=spec.lease_id, ) ) self._pending_specs.pop(mm_hash, None) @@ -451,11 +585,39 @@ def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: return False, None return False, {"ec_items": items} + def shutdown(self) -> None: + if self._shutdown: + return + self._shutdown = True + if self._registry is not None: + self._registry.shutdown() + self._zmq_stop.set() + if self._zmq_thread is not None: + self._zmq_thread.join() + if self._zmq_ctx is not None: + self._zmq_ctx.term() + + if self._engine is not None: + with self._tensor_lock: + addresses = [ + entry.tensor.data_ptr() for entry in self._tensor_by_hash.values() + ] + addresses.extend(self._pending_unregister) + unregistered = True + if addresses: + ret = self._engine.batch_unregister_memory( + list(dict.fromkeys(addresses)) + ) + if ret != 0: + unregistered = False + logger.error( + "Mooncake EC batch memory unregistration failed: %d", ret + ) + if unregistered: + self._tensor_by_hash.clear() + self._pending_unregister.clear() + self._registered_bytes = 0 + def __del__(self) -> None: - try: - if self._registry is not None: - self._registry.shutdown() - if self._zmq_ctx is not None: - self._zmq_ctx.term() - except Exception: - pass + with suppress(Exception): + self.shutdown() From 320815b26441293ef4ec663534e1df1abaadcc96 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Thu, 13 Aug 2026 01:18:27 +0000 Subject: [PATCH 05/30] Batch Mooncake embedding transfers Signed-off-by: Tianyu Guo --- .../unit/test_ec_mooncake_connector.py | 78 +++++++- .../ec_connector/mooncake_ec_connector.py | 174 ++++++++++++------ 2 files changed, 195 insertions(+), 57 deletions(-) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index cf4af7eb8a68..44b9187518d6 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -35,6 +35,7 @@ def __init__(self, *args, **kwargs): self.register_calls: list[list[int]] = [] self.unregister_calls: list[int] = [] self.batch_unregister_calls: list[list[int]] = [] + self.transfer_calls: list[list[int]] = [] def initialize(self, local_hostname, metadata_server, protocol, device_name) -> int: return 0 @@ -45,6 +46,7 @@ def get_rpc_port(self) -> int: def batch_transfer_sync_write( self, target_hostname, buffers, peer_buffer_addresses, lengths ) -> int: + self.transfer_calls.append([int(length) for length in lengths]) for src, dst, nbytes in zip(buffers, peer_buffer_addresses, lengths): ctypes.memmove(int(dst), int(src), int(nbytes)) return 0 @@ -381,9 +383,83 @@ def test_single_process_save_and_load(self, mock_vllm_config_producer): consumer_engine = consumer._engine assert isinstance(consumer_engine, CopyingFakeTransferEngine) assert consumer_engine.registered == set() - assert consumer_engine.unregister_calls == [loaded[mm_hash].data_ptr()] + assert consumer_engine.batch_unregister_calls == [ + [loaded[mm_hash].data_ptr()] + ] + consumer.shutdown() producer.shutdown() + + def test_batches_multi_item_transfer_and_reuses_socket( + self, mock_vllm_config_producer + ): + port = _find_free_port() + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "registry_http_port" + ] = port + sources = {f"hash_{i}": torch.randn(4, 16) for i in range(3)} + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + for mm_hash, tensor in sources.items(): + producer.save_caches({mm_hash: tensor}, mm_hash) + + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 1e9 + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp" + } + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) + + def make_spec(mm_hash: str) -> ECMooncakeLoadSpec: + data = httpx.get(f"http://127.0.0.1:{port}/ec/info/{mm_hash}").json() + return ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=1, + nbytes=int(data["nbytes"]), + shape=tuple(data["shape"]), + dtype=str(data["dtype"]), + producer_zmq=str(data["producer_zmq"]), + lease_id=str(data["lease_id"]), + ) + + first_meta = ECMooncakeConnectorMetadata( + loads=[make_spec("hash_0"), make_spec("hash_1")] + ) + consumer.bind_connector_metadata(first_meta) + loaded: dict[str, torch.Tensor] = {} + consumer.start_load_caches(loaded) + socket = next(iter(consumer._client_sockets.values())) + + second_meta = ECMooncakeConnectorMetadata(loads=[make_spec("hash_2")]) + consumer.bind_connector_metadata(second_meta) + consumer.start_load_caches(loaded) + + producer_engine = producer._engine + consumer_engine = consumer._engine + assert isinstance(producer_engine, CopyingFakeTransferEngine) + assert isinstance(consumer_engine, CopyingFakeTransferEngine) + assert producer_engine.transfer_calls == [[256, 256], [256]] + assert len(consumer._client_sockets) == 1 + assert next(iter(consumer._client_sockets.values())) is socket + assert all( + torch.equal(loaded[key], value) for key, value in sources.items() + ) + register_sizes = [len(call) for call in consumer_engine.register_calls] + assert register_sizes == [2, 1] + assert [len(call) for call in consumer_engine.batch_unregister_calls] == [ + 2, + 1, + ] consumer.shutdown() + producer.shutdown() def test_producer_evicts_lru_registration_at_capacity( self, mock_vllm_config_producer diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 5e5f13942b82..ea75e5e774eb 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -149,15 +149,21 @@ def unpublish_if_unleased(self, mm_hash: str) -> tuple[bool, dict[str, Any] | No self._leases.pop(mm_hash, None) return True, self._entries.pop(mm_hash, None) - def consume_lease(self, mm_hash: str, lease_id: str) -> bool: + def consume_leases(self, leases: list[tuple[str, str]]) -> bool: with self._lock: - leases = self._leases.get(mm_hash) - if leases is None: - return False - expiry = leases.pop(lease_id, 0) - if not leases: - self._leases.pop(mm_hash, None) - return expiry > time.monotonic() + now = time.monotonic() + for mm_hash, lease_id in leases: + if self._leases.get(mm_hash, {}).get(lease_id, 0) <= now: + return False + for mm_hash, lease_id in leases: + item_leases = self._leases[mm_hash] + item_leases.pop(lease_id) + if not item_leases: + self._leases.pop(mm_hash) + return True + + def consume_lease(self, mm_hash: str, lease_id: str) -> bool: + return self.consume_leases([(mm_hash, lease_id)]) class ECMooncakeConnector(ECConnectorBase): @@ -218,6 +224,8 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._zmq_thread: threading.Thread | None = None self._zmq_ctx: zmq.Context | None = None self._zmq_stop = threading.Event() + self._client_zmq_ctx: zmq.Context | None = None + self._client_sockets: dict[str, zmq.Socket] = {} self._tensor_by_hash: OrderedDict[str, _RegisteredTensor] = OrderedDict() self._registered_bytes = 0 self._pending_unregister: dict[int, torch.Tensor] = {} @@ -326,44 +334,51 @@ def loop() -> None: if req.get("op") != "pull": sock.send_json({"ok": False, "err": "unknown op"}) continue - mm_hash = req["mm_hash"] dst_session = req["dst_session"] - dst_ptr = int(req["dst_ptr"]) - nbytes = int(req["nbytes"]) - lease_id = str(req["lease_id"]) + items = req["items"] with self._tensor_lock: - entry = self._tensor_by_hash.get(mm_hash) + entries: list[_RegisteredTensor] = [] + leases: list[tuple[str, str]] = [] + dst_ptrs: list[int] = [] + lengths: list[int] = [] + for item in items: + mm_hash = str(item["mm_hash"]) + entry = self._tensor_by_hash.get(mm_hash) + nbytes = int(item["nbytes"]) + if entry is None or entry.tensor.nbytes != nbytes: + raise ValueError( + "unknown EC tensor or size mismatch" + ) + entries.append(entry) + leases.append((mm_hash, str(item["lease_id"]))) + dst_ptrs.append(int(item["dst_ptr"])) + lengths.append(nbytes) if ( - entry is not None - and self._registry is not None - and self._registry.consume_lease(mm_hash, lease_id) + self._registry is None + or not self._registry.consume_leases(leases) ): + raise ValueError("invalid lease") + for (mm_hash, _), entry in zip(leases, entries): entry.in_flight += 1 self._tensor_by_hash.move_to_end(mm_hash) - else: - entry = None - if entry is None: - sock.send_json({"ok": False, "err": "invalid lease"}) - continue try: - tensor = entry.tensor - src_ptr = tensor.data_ptr() - if tensor.nbytes != nbytes: - sock.send_json({"ok": False, "err": "size mismatch"}) - continue ret = eng.batch_transfer_sync_write( - dst_session, [src_ptr], [dst_ptr], [nbytes] + dst_session, + [entry.tensor.data_ptr() for entry in entries], + dst_ptrs, + lengths, ) if ( ret == 0 and envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER and torch.accelerator.is_available() - and tensor.is_cuda + and any(entry.tensor.is_cuda for entry in entries) ): torch.accelerator.synchronize() finally: with self._tensor_lock: - entry.in_flight -= 1 + for entry in entries: + entry.in_flight -= 1 sock.send_json({"ok": ret == 0, "mooncake_ret": int(ret)}) except Exception as e: logger.exception("EC Mooncake pull handler error: %s", e) @@ -396,6 +411,43 @@ def _ensure_producer_services(self) -> None: self._registry.start() self._producer_services_started = True + def _get_pull_socket(self, producer_zmq: str) -> zmq.Socket: + sock = self._client_sockets.get(producer_zmq) + if sock is not None: + return sock + if self._client_zmq_ctx is None: + self._client_zmq_ctx = zmq.Context() + sock = self._client_zmq_ctx.socket(zmq.REQ) + sock.setsockopt(zmq.RCVTIMEO, 120_000) + sock.connect(producer_zmq) + self._client_sockets[producer_zmq] = sock + return sock + + def _send_pull(self, producer_zmq: str, pull: dict[str, Any]) -> dict[str, Any]: + sock = self._get_pull_socket(producer_zmq) + try: + sock.send_json(pull) + return sock.recv_json() + except zmq.ZMQError: + sock.close(linger=0) + self._client_sockets.pop(producer_zmq, None) + raise + + def _unregister_memories(self, tensors: list[torch.Tensor]) -> None: + assert self._engine is not None + addresses = [tensor.data_ptr() for tensor in tensors] + ret = self._engine.batch_unregister_memory(addresses) + if ret != 0: + for tensor in tensors: + self._pending_unregister[tensor.data_ptr()] = tensor + logger.warning( + "Keeping %d EC tensors alive after Mooncake unregistration failure", + len(tensors), + ) + return + for address in addresses: + self._pending_unregister.pop(address, None) + def start_load_caches( self, encoder_cache: dict[str, torch.Tensor], **kwargs: Any ) -> None: @@ -410,6 +462,7 @@ def start_load_caches( ) device = torch.device(buf) + pending: list[tuple[ECMooncakeLoadSpec, torch.Tensor]] = [] for spec in metadata.loads: if spec.mm_hash in encoder_cache: continue @@ -417,39 +470,43 @@ def start_load_caches( if torch_dtype is None: raise ValueError(f"Unsupported torch dtype string: {spec.dtype!r}") t = torch.empty(spec.shape, dtype=torch_dtype, device=device) - ret = eng.batch_register_memory([t.data_ptr()], [t.nbytes]) - if ret != 0: - raise RuntimeError( - "Mooncake EC batch_register_memory failed on consumer." - ) - try: + pending.append((spec, t)) + if not pending: + return + + tensors = [tensor for _, tensor in pending] + ret = eng.batch_register_memory( + [tensor.data_ptr() for tensor in tensors], + [tensor.nbytes for tensor in tensors], + ) + if ret != 0: + raise RuntimeError("Mooncake EC batch_register_memory failed on consumer.") + try: + batches: dict[str, list[tuple[ECMooncakeLoadSpec, torch.Tensor]]] = {} + for spec, tensor in pending: + batches.setdefault(spec.producer_zmq, []).append((spec, tensor)) + for producer_zmq, batch in batches.items(): pull = { "op": "pull", - "mm_hash": spec.mm_hash, "dst_session": f"{self._hostname}:{eng.get_rpc_port()}", - "dst_ptr": t.data_ptr(), - "nbytes": t.nbytes, - "lease_id": spec.lease_id, + "items": [ + { + "mm_hash": spec.mm_hash, + "dst_ptr": tensor.data_ptr(), + "nbytes": tensor.nbytes, + "lease_id": spec.lease_id, + } + for spec, tensor in batch + ], } - ctx = zmq.Context() - sock = ctx.socket(zmq.REQ) - sock.setsockopt(zmq.RCVTIMEO, 120_000) - sock.connect(spec.producer_zmq) - try: - sock.send_json(pull) - resp = sock.recv_json() - finally: - sock.close(linger=0) - ctx.term() + resp = self._send_pull(producer_zmq, pull) if not resp.get("ok"): raise RuntimeError(f"EC Mooncake pull failed: {resp}") - if envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER and device.type == "cuda": - torch.accelerator.synchronize() - finally: - if not self._unregister_memory(t): - logger.warning( - "Keeping EC tensor alive after Mooncake unregistration failure" - ) + if envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER and device.type == "cuda": + torch.accelerator.synchronize() + finally: + self._unregister_memories(tensors) + for spec, t in pending: encoder_cache[spec.mm_hash] = t logger.debug("Loaded EC tensor for mm_hash=%s via Mooncake", spec.mm_hash) @@ -596,6 +653,11 @@ def shutdown(self) -> None: self._zmq_thread.join() if self._zmq_ctx is not None: self._zmq_ctx.term() + for sock in self._client_sockets.values(): + sock.close(linger=0) + self._client_sockets.clear() + if self._client_zmq_ctx is not None: + self._client_zmq_ctx.term() if self._engine is not None: with self._tensor_lock: From 578126a936b7185def2e19f87ccd7ba2488d016c Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Thu, 13 Aug 2026 01:35:00 +0000 Subject: [PATCH 06/30] Clean up code Signed-off-by: Tianyu Guo --- docs/features/mooncake_connector_usage.md | 39 ------ .../unit/test_mooncake_connector.py | 63 --------- .../ec_connector/mooncake_ec_connector.py | 10 -- .../v1/mooncake/mooncake_connector.py | 132 +----------------- vllm/envs.py | 19 --- 5 files changed, 6 insertions(+), 257 deletions(-) diff --git a/docs/features/mooncake_connector_usage.md b/docs/features/mooncake_connector_usage.md index b211d0a31891..6cce042fb6f8 100644 --- a/docs/features/mooncake_connector_usage.md +++ b/docs/features/mooncake_connector_usage.md @@ -14,12 +14,6 @@ Install mooncake through pip: `uv pip install mooncake-transfer-engine-cuda13`. vLLM defaults to CUDA 13. On a CUDA 12 environment install `mooncake-transfer-engine` instead — the two are the same release built against different CUDA majors, and the wrong one fails to import with `libcudart.so.: cannot open shared object file`. -If you observe PD transfer data mismatches (`dst != src`) with the CUDA 13 -package, enable the mitigations in -[Transfer reliability](#transfer-reliability) below (see -[vllm #42395](https://github.com/vllm-project/vllm/issues/42395), -[Mooncake #2086](https://github.com/kvcache-ai/Mooncake/issues/2086)). - Refer to [Mooncake official repository](https://github.com/kvcache-ai/Mooncake) for more installation instructions ## Usage @@ -62,36 +56,6 @@ Now you can send requests to the proxy server through port 8000. - Default: 480 - If a request is aborted and the decoder has not yet notified the prefiller, the prefill instance will release its KV-cache blocks after this timeout to avoid holding them indefinitely. -### Transfer reliability - -Under concurrent PD load, some Mooncake transfer-engine builds (notably -`mooncake-transfer-engine-cuda13==0.3.10.post2`) can produce destination bytes -that do not match the producer source for very large coalesced descriptors. -vLLM mitigations (env vars or `kv_connector_extra_config` keys): - -- `VLLM_MOONCAKE_MAX_TRANSFER_BYTES` / `max_transfer_bytes`: Split any single - transfer descriptor larger than this size into contiguous chunks (recommended - starting value: `262144` for multimodal PD workloads). -- `VLLM_MOONCAKE_SYNC_AFTER_TRANSFER` / `sync_after_transfer`: Call - `torch.cuda.synchronize()` after each Mooncake batch transfer on producer and - consumer (reduces visibility races at some throughput cost). -- `VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY` / `verify_transfer_integrity`: - Debug-only SHA-256 check that producer source memory is unchanged after - transfer (does not verify remote destination bytes). - -Example prefill/decode extra config: - -```json -{ - "kv_connector": "MooncakeConnector", - "kv_role": "kv_producer", - "kv_connector_extra_config": { - "max_transfer_bytes": 262144, - "sync_after_transfer": true - } -} -``` - ## KV Transfer Config ### KV Role Options @@ -105,9 +69,6 @@ Example prefill/decode extra config: - **num_workers**: Size of thread pool for one prefiller worker to transfer KV caches by mooncake. (default 10) - **mooncake_protocol**: Mooncake connector protocol. (default "rdma") - **device_name**: Comma-separated whitelist of RDMA devices (e.g. `"mlx5_0,mlx5_1"`) to restrict topology discovery to. Empty discovers every device. Useful on hosts exposing a mix of InfiniBand and RoCE ports, where both peers must settle on the same link layer. -- **max_transfer_bytes**: Split descriptors larger than this many bytes (see [Transfer reliability](#transfer-reliability)) -- **sync_after_transfer**: Synchronize CUDA after each Mooncake batch transfer (default false) -- **verify_transfer_integrity**: Debug SHA-256 check of producer source after transfer (default false) ## Example Scripts/Code diff --git a/tests/v1/kv_connector/unit/test_mooncake_connector.py b/tests/v1/kv_connector/unit/test_mooncake_connector.py index 54057a38de65..4847b956b196 100644 --- a/tests/v1/kv_connector/unit/test_mooncake_connector.py +++ b/tests/v1/kv_connector/unit/test_mooncake_connector.py @@ -27,7 +27,6 @@ _align_transfer_regions, get_mooncake_bootstrap_addr, should_launch_bootstrap_server, - split_transfer_descriptors, ) from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import ( MooncakeBootstrapServer, @@ -352,68 +351,6 @@ async def run_in_executor(self, executor, func, *args): prefill_worker.shutdown() -def test_split_transfer_descriptors_disabled(): - src = [100, 200] - dst = [300, 400] - lengths = [16, 32] - assert split_transfer_descriptors(src, dst, lengths, 0) == (src, dst, lengths) - - -def test_split_transfer_descriptors_splits_large_descriptor(): - # Reproduces coalesced multimodal descriptor size from vllm #42395. - src_base, dst_base, total = 0x1000, 0x2000, 2_424_832 - max_chunk = 262_144 - out_src, out_dst, out_len = split_transfer_descriptors( - [src_base], [dst_base], [total], max_chunk - ) - assert sum(out_len) == total - assert len(out_len) == (total + max_chunk - 1) // max_chunk - offset = 0 - for src, dst, length in zip(out_src, out_dst, out_len): - assert src == src_base + offset - assert dst == dst_base + offset - assert 0 < length <= max_chunk - offset += length - - -def test_send_blocks_splits_when_max_transfer_bytes_set(): - """_send_blocks should chunk descriptors before calling TransferEngine.""" - from vllm.distributed.kv_transfer.kv_connector.v1.mooncake import ( - mooncake_connector as mc, - ) - - captured: list[list[int]] = [] - - class StubWorker: - _max_transfer_bytes = 8 - _sync_after_transfer = False - _verify_transfer_integrity = False - device_id = 0 - xfer_stats = MagicMock() - - engine = FakeMooncakeWrapper() - - def _send_blocks(self, remote_session, src_ptrs, dst_ptrs, lengths): - return mc.MooncakeConnectorWorker._send_blocks( - self, remote_session, src_ptrs, dst_ptrs, lengths - ) - - worker = StubWorker() - - def capture_transfer(_session, _src, _dst, lengths): - captured.append(lengths) - return 0 - - with patch.object( - worker.engine, - "batch_transfer_sync_write", - side_effect=capture_transfer, - ): - ret = worker._send_blocks("host:1", [1000], [2000], [20]) - assert ret == 0 - assert captured == [[8, 8, 4]] - - def test_basic_interface(): """Unit test for basic MooncakeConnector interface functionality.""" diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index ea75e5e774eb..a59638478190 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -25,7 +25,6 @@ import zmq from fastapi import FastAPI, HTTPException -from vllm import envs from vllm.config import VllmConfig from vllm.distributed.ec_transfer.ec_connector.base import ( ECConnectorBase, @@ -368,13 +367,6 @@ def loop() -> None: dst_ptrs, lengths, ) - if ( - ret == 0 - and envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER - and torch.accelerator.is_available() - and any(entry.tensor.is_cuda for entry in entries) - ): - torch.accelerator.synchronize() finally: with self._tensor_lock: for entry in entries: @@ -502,8 +494,6 @@ def start_load_caches( resp = self._send_pull(producer_zmq, pull) if not resp.get("ok"): raise RuntimeError(f"EC Mooncake pull failed: {resp}") - if envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER and device.type == "cuda": - torch.accelerator.synchronize() finally: self._unregister_memories(tensors) for spec, t in pending: diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py index 478dedcb549e..dd94ffb2f9ed 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py @@ -1,7 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import asyncio -import hashlib import logging import threading import time @@ -233,75 +232,6 @@ def _compute_sender_transfer_plan( ) -def split_transfer_descriptors( - src_ptrs: list[int], - dst_ptrs: list[int], - lengths: list[int], - max_bytes: int, -) -> tuple[list[int], list[int], list[int]]: - """Split descriptors longer than ``max_bytes`` into contiguous chunks. - - Mitigates Mooncake transfer-engine issues with very large single - descriptors under concurrent PD load (see vllm #42395). - """ - if max_bytes <= 0: - return src_ptrs, dst_ptrs, lengths - out_src: list[int] = [] - out_dst: list[int] = [] - out_len: list[int] = [] - for src, dst, length in zip(src_ptrs, dst_ptrs, lengths): - offset = 0 - while offset < length: - chunk = min(max_bytes, length - offset) - out_src.append(src + offset) - out_dst.append(dst + offset) - out_len.append(chunk) - offset += chunk - return out_src, out_dst, out_len - - -def _digest_gpu_memory_region(ptr: int, length: int, device_id: int) -> bytes: - """Return SHA-256 digest of a CUDA memory region (debug / verify mode).""" - if length == 0: - return hashlib.sha256(b"").digest() - if not torch.cuda.is_available(): - raise RuntimeError("GPU memory digest requires CUDA.") - from torch.cuda import cudart - - host = torch.empty(length, dtype=torch.uint8, device="cpu", pin_memory=True) - with torch.cuda.device(device_id): - err = cudart.cudaMemcpy( - host.data_ptr(), - ptr, - length, - cudart.cudaMemcpyKind.cudaMemcpyDeviceToHost, - )[0] - if err != 0: - raise RuntimeError(f"cudaMemcpy D2H failed with error {err}") - return hashlib.sha256(host.numpy().tobytes()).digest() - - -def _resolve_mooncake_transfer_tuning( - extra_config: dict[str, Any], -) -> tuple[int | None, bool, bool]: - """Merge kv_connector_extra_config with Mooncake env overrides.""" - max_bytes: int | None = None - if envs.VLLM_MOONCAKE_MAX_TRANSFER_BYTES > 0: - max_bytes = envs.VLLM_MOONCAKE_MAX_TRANSFER_BYTES - cfg_max = extra_config.get("max_transfer_bytes") - if cfg_max is not None: - max_bytes = int(cfg_max) - - sync_after = envs.VLLM_MOONCAKE_SYNC_AFTER_TRANSFER - if "sync_after_transfer" in extra_config: - sync_after = bool(extra_config["sync_after_transfer"]) - - verify = envs.VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY - if "verify_transfer_integrity" in extra_config: - verify = bool(extra_config["verify_transfer_integrity"]) - return max_bytes, sync_after, verify - - def _can_coalesce_block_transfers( local_region_block_len: int, remote_region_block_len: int, @@ -995,26 +925,15 @@ def __init__( # Tasks can await async events, so a surplus (2x is a robust heuristic) # prevents workers from idling. self.num_sender_tasks = self.num_sender_workers * 2 - extra_config = kv_transfer_config.kv_connector_extra_config # type: ignore[union-attr] - protocol = extra_config.get("mooncake_protocol", "rdma") - device_name = extra_config.get("device_name", "") - ( - self._max_transfer_bytes, - self._sync_after_transfer, - self._verify_transfer_integrity, - ) = _resolve_mooncake_transfer_tuning(extra_config) + protocol = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr] + "mooncake_protocol", "rdma" + ) + device_name = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr] + "device_name", "" + ) logger.info( "The Mooncake Transfer Engine is using %s as its protocol.", protocol ) - if self._max_transfer_bytes: - logger.info( - "MooncakeConnector will split transfer descriptors above %d bytes", - self._max_transfer_bytes, - ) - if self._sync_after_transfer: - logger.info("MooncakeConnector sync_after_transfer is enabled") - if self._verify_transfer_integrity: - logger.info("MooncakeConnector verify_transfer_integrity is enabled") ret_value = self.engine.initialize( self.hostname, "P2PHANDSHAKE", protocol, device_name ) @@ -1707,49 +1626,12 @@ def _send_blocks( dst_ptrs: list[int], lengths: list[int], ) -> int: - if self._max_transfer_bytes: - src_ptrs, dst_ptrs, lengths = split_transfer_descriptors( - src_ptrs, dst_ptrs, lengths, self._max_transfer_bytes - ) - - pre_hashes: list[bytes] | None = None - if self._verify_transfer_integrity: - try: - pre_hashes = [ - _digest_gpu_memory_region(src, length, self.device_id) - for src, length in zip(src_ptrs, lengths) - ] - except Exception as e: - logger.error("Mooncake transfer integrity pre-hash failed: %s", e) - return -1 - start_time = time.perf_counter() ret_value = self.engine.batch_transfer_sync_write( remote_session, src_ptrs, dst_ptrs, lengths ) duration = time.perf_counter() - start_time if ret_value == 0: - if self._sync_after_transfer and torch.cuda.is_available(): - torch.cuda.synchronize(device=self.device_id) - if pre_hashes is not None: - for idx, (src, length) in enumerate(zip(src_ptrs, lengths)): - try: - post_hash = _digest_gpu_memory_region( - src, length, self.device_id - ) - except Exception as e: - logger.error( - "Mooncake transfer integrity post-hash failed: %s", e - ) - return -1 - if post_hash != pre_hashes[idx]: - logger.error( - "Mooncake source memory changed after transfer " - "(descriptor_idx=%d length=%d)", - idx, - length, - ) - return -1 self.xfer_stats.record_transfer( duration_s=duration, total_bytes=sum(lengths), @@ -2011,8 +1893,6 @@ def process_pulling_result( self.finished_recving_reqs.add(pull_meta.d_req_id) if ok_reqs: - if self._sync_after_transfer and torch.cuda.is_available(): - torch.cuda.synchronize(device=self.device_id) logger.debug("pulling kv_caches for %s finished", ok_reqs) if response.err_reqs: diff --git a/vllm/envs.py b/vllm/envs.py index 7598b8462df6..819c6d30cd3e 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -245,9 +245,6 @@ VLLM_ROCM_QUICK_REDUCE_MIN_SIZE_BYTES_MB: int | None = None VLLM_ROCM_QUICK_REDUCE_QUANTIZATION_MIN_SIZE_KB: int | None = None VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT: int = 480 - VLLM_MOONCAKE_MAX_TRANSFER_BYTES: int = 0 - VLLM_MOONCAKE_SYNC_AFTER_TRANSFER: bool = False - VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY: bool = False VLLM_ENABLE_CUDAGRAPH_GC: bool = False VLLM_LOOPBACK_IP: str = "" VLLM_ALLOW_CHUNKED_LOCAL_ATTN_WITH_HYBRID_KV_CACHE: bool = True @@ -1776,22 +1773,6 @@ def _resolve_rust_cli_path() -> str | None: "VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT": lambda: int( os.getenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "480") ), - # Split Mooncake transfer descriptors larger than this many bytes. - # Mitigates data-integrity issues with very large coalesced copies under - # concurrent PD load (vllm #42395). 0 disables splitting. - "VLLM_MOONCAKE_MAX_TRANSFER_BYTES": lambda: int( - os.getenv("VLLM_MOONCAKE_MAX_TRANSFER_BYTES", "0") - ), - # CUDA synchronize after each Mooncake batch transfer (producer and - # consumer). May reduce dst/src mismatches at the cost of throughput. - "VLLM_MOONCAKE_SYNC_AFTER_TRANSFER": lambda: ( - os.getenv("VLLM_MOONCAKE_SYNC_AFTER_TRANSFER", "0").lower() in ("1", "true") - ), - # Hash GPU source regions before/after Mooncake writes (debug only). - "VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY": lambda: ( - os.getenv("VLLM_MOONCAKE_VERIFY_TRANSFER_INTEGRITY", "0").lower() - in ("1", "true") - ), # If set, it means we pre-downloaded cubin files and flashinfer will # read the cubin files directly. "VLLM_HAS_FLASHINFER_CUBIN": lambda: bool( From c58f20ded31c9711809a1a0e6216cc5514accf97 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Fri, 14 Aug 2026 01:46:11 +0000 Subject: [PATCH 07/30] Add consumer_buffer_pool Signed-off-by: Tianyu Guo --- .../unit/test_ec_mooncake_connector.py | 64 ++++++ .../ec_connector/mooncake_ec_connector.py | 210 ++++++++++++++++-- 2 files changed, 257 insertions(+), 17 deletions(-) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 44b9187518d6..2c21f41484fa 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -23,6 +23,7 @@ ECMooncakeConnectorMetadata, ECMooncakeLoadSpec, ECMooncakeRegistryServer, + _ContiguousAllocator, ) from vllm.v1.core.sched.output import SchedulerOutput @@ -182,6 +183,21 @@ def test_factory_registers_connector(self): assert cls.__name__ == "ECMooncakeConnector" +class TestContiguousAllocator: + def test_reuses_and_coalesces_contiguous_regions(self): + allocator = _ContiguousAllocator(1024, alignment=256) + + first = allocator.allocate(1) + second = allocator.allocate(300) + assert first == (0, 256) + assert second == (256, 512) + assert allocator.allocate(300) is None + + allocator.free(*first) + allocator.free(*second) + assert allocator.allocate(1024) == (0, 1024) + + class TestECMooncakeConnectorValidation: def test_consumer_scheduler_requires_remote_registry( self, mock_vllm_config_consumer @@ -326,6 +342,54 @@ def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer class TestECMooncakeWorkerTransfer: + @pytest.mark.skipif( + not torch.accelerator.is_available(), + reason="Requires an accelerator for registered pool", + ) + def test_consumer_reuses_registered_cuda_pool(self, mock_vllm_config_consumer): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + specs = [ + ECMooncakeLoadSpec( + mm_hash=f"hash_{index}", + num_token=1, + nbytes=256, + shape=(32, 2), + dtype="float32", + producer_zmq="tcp://127.0.0.1:1", + lease_id=f"lease_{index}", + ) + for index in range(2) + ] + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + consumer.bind_connector_metadata(ECMooncakeConnectorMetadata(loads=specs)) + cache: dict[str, torch.Tensor] = {} + with patch.object(consumer, "_send_pull", return_value={"ok": True}): + consumer.start_load_caches(cache) + + engine = consumer._engine + pool = consumer._consumer_pool + assert isinstance(engine, CopyingFakeTransferEngine) + assert pool is not None + assert engine.register_calls == [[pool.data_ptr()]] + assert engine.batch_unregister_calls == [] + assert cache["hash_0"].data_ptr() == pool.data_ptr() + assert cache["hash_1"].data_ptr() == pool.data_ptr() + 256 + + cache.clear() + consumer._release_stale_consumer_allocations(cache) + torch.accelerator.synchronize() + consumer._poll_consumer_pool_frees() + assert consumer._consumer_pool_allocator is not None + assert consumer._consumer_pool_allocator.allocate(4096) == (0, 4096) + consumer.shutdown() + def test_single_process_save_and_load(self, mock_vllm_config_producer): """Host-memory pull path (fake engine uses memcpy; CUDA ptrs need e2e).""" port = _find_free_port() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index a59638478190..b9b4350124df 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -11,6 +11,7 @@ from __future__ import annotations import json +import math import threading import time import uuid @@ -79,6 +80,44 @@ class _RegisteredTensor: in_flight: int = 0 +@dataclass +class _ConsumerPoolAllocation: + offset: int + size: int + tensor: torch.Tensor + + +class _ContiguousAllocator: + def __init__(self, capacity: int, alignment: int = 256): + self.capacity = capacity + self.alignment = alignment + self._free = [(0, capacity)] + + def allocate(self, nbytes: int) -> tuple[int, int] | None: + size = math.ceil(nbytes / self.alignment) * self.alignment + for index, (offset, available) in enumerate(self._free): + if size > available: + continue + if size == available: + self._free.pop(index) + else: + self._free[index] = (offset + size, available - size) + return offset, size + return None + + def free(self, offset: int, size: int) -> None: + self._free.append((offset, size)) + self._free.sort() + merged: list[tuple[int, int]] = [] + for free_offset, free_size in self._free: + if merged and sum(merged[-1]) == free_offset: + previous_offset, previous_size = merged[-1] + merged[-1] = (previous_offset, previous_size + free_size) + else: + merged.append((free_offset, free_size)) + self._free = merged + + class ECMooncakeRegistryServer: """Lightweight HTTP registry on the producer for remote has_cache_item / info.""" @@ -177,6 +216,8 @@ class ECMooncakeConnector(ECConnectorBase): registry (default ``9018``). - ``mooncake_protocol`` (optional): Passed to ``TransferEngine.initialize`` (default ``"rdma"``). + - ``consumer_buffer_pool_size`` (consumer, optional): Bytes reserved for a + long-lived registered CUDA receive arena (default ``ec_buffer_size``). Limitations: ``tensor_parallel_size`` and ``pipeline_parallel_size`` must be ``1`` (same assumption as Mooncake KV connector for P2P handshake). @@ -211,6 +252,18 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._model_config = vllm_config.model_config self._metadata_fields_cache: dict[str, set[str]] = {} + pool_size = self._extra.get( + "consumer_buffer_pool_size", self._registered_capacity + ) + self._consumer_pool_capacity = int(pool_size) + self._consumer_pool: torch.Tensor | None = None + self._consumer_pool_allocator: _ContiguousAllocator | None = None + self._consumer_allocations: dict[str, _ConsumerPoolAllocation] = {} + self._consumer_pending_frees: list[ + tuple[torch.Event, _ConsumerPoolAllocation] + ] = [] + self._consumer_pool_disabled = self._consumer_pool_capacity <= 0 + # Scheduler (consumer): mm_hash -> pending tensor layout from registry self._pending_specs: dict[str, ECMooncakeLoadSpec] = {} self._mm_datas_need_loads: dict[str, int] = {} @@ -440,6 +493,92 @@ def _unregister_memories(self, tensors: list[torch.Tensor]) -> None: for address in addresses: self._pending_unregister.pop(address, None) + def _ensure_consumer_pool(self, device: torch.device) -> None: + if ( + self._consumer_pool is not None + or self._consumer_pool_disabled + or device.type != "cuda" + ): + return + try: + pool = torch.empty( + self._consumer_pool_capacity, dtype=torch.uint8, device=device + ) + ret = self._ensure_engine().batch_register_memory( + [pool.data_ptr()], [pool.nbytes] + ) + if ret != 0: + raise RuntimeError(f"Mooncake returned {ret}") + except (RuntimeError, torch.OutOfMemoryError) as e: + self._consumer_pool_disabled = True + logger.warning( + "Could not initialize the EC consumer buffer pool; falling back " + "to per-tensor registration: %s", + e, + ) + return + self._consumer_pool = pool + self._consumer_pool_allocator = _ContiguousAllocator(pool.nbytes) + logger.info( + "Registered %d-byte CUDA receive pool for Mooncake EC", + pool.nbytes, + ) + + def _poll_consumer_pool_frees(self) -> None: + allocator = self._consumer_pool_allocator + if allocator is None: + return + pending = [] + for event, allocation in self._consumer_pending_frees: + if event.query(): + allocator.free(allocation.offset, allocation.size) + else: + pending.append((event, allocation)) + self._consumer_pending_frees = pending + + def _release_stale_consumer_allocations( + self, encoder_cache: dict[str, torch.Tensor] + ) -> None: + if self._consumer_pool is None: + return + for mm_hash, allocation in list(self._consumer_allocations.items()): + if encoder_cache.get(mm_hash) is allocation.tensor: + continue + self._consumer_allocations.pop(mm_hash) + event = torch.Event() + event.record(torch.accelerator.current_stream(self._consumer_pool.device)) + self._consumer_pending_frees.append((event, allocation)) + self._poll_consumer_pool_frees() + + def _allocate_consumer_tensor( + self, + spec: ECMooncakeLoadSpec, + dtype: torch.dtype, + device: torch.device, + ) -> tuple[torch.Tensor, _ConsumerPoolAllocation | None]: + expected_nbytes = ( + math.prod(spec.shape) * torch.empty((), dtype=dtype).element_size() + ) + if expected_nbytes != spec.nbytes: + raise ValueError( + f"EC tensor size mismatch for {spec.mm_hash}: metadata has " + f"{spec.nbytes} bytes, shape and dtype require {expected_nbytes}." + ) + + self._ensure_consumer_pool(device) + allocator = self._consumer_pool_allocator + pool = self._consumer_pool + if allocator is not None and pool is not None: + region = allocator.allocate(spec.nbytes) + if region is not None: + offset, size = region + tensor = ( + pool.narrow(0, offset, spec.nbytes).view(dtype).view(spec.shape) + ) + return tensor, _ConsumerPoolAllocation(offset, size, tensor) + + return torch.empty(spec.shape, dtype=dtype, device=device), None + def start_load_caches( self, encoder_cache: dict[str, torch.Tensor], **kwargs: Any ) -> None: @@ -454,29 +593,49 @@ def start_load_caches( ) device = torch.device(buf) - pending: list[tuple[ECMooncakeLoadSpec, torch.Tensor]] = [] + self._release_stale_consumer_allocations(encoder_cache) + + pending: list[ + tuple[ECMooncakeLoadSpec, torch.Tensor, _ConsumerPoolAllocation | None] + ] = [] for spec in metadata.loads: if spec.mm_hash in encoder_cache: continue torch_dtype = getattr(torch, spec.dtype, None) if torch_dtype is None: raise ValueError(f"Unsupported torch dtype string: {spec.dtype!r}") - t = torch.empty(spec.shape, dtype=torch_dtype, device=device) - pending.append((spec, t)) + tensor, allocation = self._allocate_consumer_tensor( + spec, torch_dtype, device + ) + pending.append((spec, tensor, allocation)) if not pending: return - tensors = [tensor for _, tensor in pending] - ret = eng.batch_register_memory( - [tensor.data_ptr() for tensor in tensors], - [tensor.nbytes for tensor in tensors], - ) - if ret != 0: - raise RuntimeError("Mooncake EC batch_register_memory failed on consumer.") + standalone = [tensor for _, tensor, allocation in pending if allocation is None] + standalone_registered = False try: - batches: dict[str, list[tuple[ECMooncakeLoadSpec, torch.Tensor]]] = {} - for spec, tensor in pending: - batches.setdefault(spec.producer_zmq, []).append((spec, tensor)) + if standalone: + ret = eng.batch_register_memory( + [tensor.data_ptr() for tensor in standalone], + [tensor.nbytes for tensor in standalone], + ) + if ret != 0: + raise RuntimeError( + "Mooncake EC batch_register_memory failed on consumer." + ) + standalone_registered = True + batches: dict[ + str, + list[ + tuple[ + ECMooncakeLoadSpec, + torch.Tensor, + _ConsumerPoolAllocation | None, + ] + ], + ] = {} + for item in pending: + batches.setdefault(item[0].producer_zmq, []).append(item) for producer_zmq, batch in batches.items(): pull = { "op": "pull", @@ -488,16 +647,26 @@ def start_load_caches( "nbytes": tensor.nbytes, "lease_id": spec.lease_id, } - for spec, tensor in batch + for spec, tensor, _ in batch ], } resp = self._send_pull(producer_zmq, pull) if not resp.get("ok"): raise RuntimeError(f"EC Mooncake pull failed: {resp}") + except Exception: + allocator = self._consumer_pool_allocator + if allocator is not None: + for _, _, allocation in pending: + if allocation is not None: + allocator.free(allocation.offset, allocation.size) + raise finally: - self._unregister_memories(tensors) - for spec, t in pending: - encoder_cache[spec.mm_hash] = t + if standalone_registered: + self._unregister_memories(standalone) + for spec, tensor, allocation in pending: + encoder_cache[spec.mm_hash] = tensor + if allocation is not None: + self._consumer_allocations[spec.mm_hash] = allocation logger.debug("Loaded EC tensor for mm_hash=%s via Mooncake", spec.mm_hash) def save_caches( @@ -650,6 +819,13 @@ def shutdown(self) -> None: self._client_zmq_ctx.term() if self._engine is not None: + if self._consumer_pool is not None and self._unregister_memory( + self._consumer_pool + ): + self._consumer_pool = None + self._consumer_pool_allocator = None + self._consumer_allocations.clear() + self._consumer_pending_frees.clear() with self._tensor_lock: addresses = [ entry.tensor.data_ptr() for entry in self._tensor_by_hash.values() From f50448cd516079b43060e9f3af8274e692a5aac9 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Sun, 16 Aug 2026 01:48:32 +0000 Subject: [PATCH 08/30] Fix stale encoder cache eviction notifications Signed-off-by: Tianyu Guo --- tests/v1/core/test_encoder_cache_manager.py | 19 +++++++++++++++++++ vllm/v1/core/encoder_cache_manager.py | 4 +++- 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/tests/v1/core/test_encoder_cache_manager.py b/tests/v1/core/test_encoder_cache_manager.py index e225666f8443..e56bcbf5c63a 100644 --- a/tests/v1/core/test_encoder_cache_manager.py +++ b/tests/v1/core/test_encoder_cache_manager.py @@ -167,6 +167,25 @@ def test_get_freed_mm_hashes_clears_freed_list(): assert manager.get_freed_mm_hashes() == [] +def test_reallocated_hash_is_not_reported_as_freed(): + manager = EncoderCacheManager(cache_size=8) + req_a = MockRequest("reqA", ["a"], [4]) + req_b = MockRequest("reqB", ["b"], [4]) + req_c = MockRequest("reqC", ["c"], [4]) + + manager.allocate(req_a, 0) + manager.allocate(req_b, 0) + manager.free(req_a) + manager.free(req_b) + + assert manager.can_allocate(req_c, 0, int(1e9), 0) + manager.allocate(req_c, 0) + assert manager.can_allocate(req_a, 0, int(1e9), 0) + manager.allocate(req_a, 0) + + assert manager.get_freed_mm_hashes() == ["b"] + + def test_schedule_request_multi_images_respect_space_limit(): manager = EncoderCacheManager(cache_size=10) req = MockRequest("reqA", ["a", "b"], [5, 6]) diff --git a/vllm/v1/core/encoder_cache_manager.py b/vllm/v1/core/encoder_cache_manager.py index 02133ff5b888..8d2a81a11b13 100644 --- a/vllm/v1/core/encoder_cache_manager.py +++ b/vllm/v1/core/encoder_cache_manager.py @@ -269,7 +269,9 @@ def get_freed_mm_hashes(self) -> list[str]: encoder outputs can be removed from their caches. The internal list is cleared after this call. """ - freed = self.freed + # An entry evicted early in the scheduling pass can be allocated again + # later in the same pass. Keep its worker-side tensor in that case. + freed = [mm_hash for mm_hash in self.freed if mm_hash not in self.cached] self.freed = [] return freed From 456b7dc4a99625798417ca64eca020700c5151c4 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Mon, 17 Aug 2026 23:15:33 +0000 Subject: [PATCH 09/30] Support async load and save for PUSH mode Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 289 +- tests/v1/core/test_encoder_cache_manager.py | 32 + tests/v1/core/test_scheduler.py | 93 +- tests/v1/ec_connector/integration/README.md | 12 - .../run_epd_mooncake_ec_full_pipeline.sh | 37 +- .../test_ec_mooncake_transfer_e2e.py | 211 -- .../unit/test_ec_mooncake_connector.py | 1509 +++++++--- .../unit/test_worker_ec_connector.py | 26 + .../ec_transfer/ec_connector/base.py | 38 +- .../ec_transfer/ec_connector/cpu/connector.py | 8 +- .../ec_connector/cpu/scheduler/__init__.py | 6 +- .../ec_connector/mooncake_ec_connector.py | 2564 ++++++++++++++--- vllm/v1/core/sched/scheduler.py | 29 +- .../worker/ec_connector_model_runner_mixin.py | 7 + vllm/v1/worker/gpu/ec_connector.py | 3 + vllm/v1/worker/gpu_model_runner.py | 25 +- vllm/v1/worker/gpu_worker.py | 5 + 17 files changed, 3769 insertions(+), 1125 deletions(-) delete mode 100644 tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index 75b7a3c277ac..fb12f3a0a572 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -61,6 +61,10 @@ # Diagnostic switch: forward the original request to the decoder so the # only difference from the rewrite path is the rewrite itself. NO_REWRITE = False +# Decode-side retries for a retryable internal error (`finish_reason="error"`, +# e.g. an encoder embedding the connector could not deliver). Re-issuing runs +# the encode again, which produces a fresh transfer. +DECODE_RETRIES = 1 # Grid metadata reported by the encoder instance, keyed by item index. @@ -102,6 +106,7 @@ def rewrite_for_decode(req_data: dict, item_meta: dict[int, dict]) -> dict: here -- a second derivation could disagree with the encoder's. """ rewritten = 0 + transfer_items = [] idx = 0 new_messages = [] for msg in req_data.get("messages", []): @@ -117,13 +122,19 @@ def rewrite_for_decode(req_data: dict, item_meta: dict[int, dict]) -> dict: meta = dict(item_meta.get(idx) or {}) idx += 1 item_uuid = meta.pop("mm_hash", None) + transfer_id = meta.pop("transfer_id", None) # Whatever keys the encoder reported are the metadata its model # declared as needed to size the placeholder range; the proxy does # not need to know their names. metadata = {k: _b64_tensor(v) for k, v in meta.items()} if not metadata or not item_uuid: - # The encoder reported no metadata (e.g. the item came from its - # processor cache); let the decoder process the media itself. + # Nothing to size the placeholder range with. A processor cache + # hit is not a cause on its own: with the default `lru` type the + # engine restores the item before the scheduler reports it. It + # goes missing when the encode request failed, or under + # `--mm-processor-cache-type shm`, where a hit replaces the item + # with its shared-memory address and only the worker restores + # it. Send the media so the decoder can derive the grid itself. new_content.append(item) continue new_content.append( @@ -133,13 +144,22 @@ def rewrite_for_decode(req_data: dict, item_meta: dict[int, dict]) -> dict: "uuid": item_uuid, } ) + if transfer_id is not None: + transfer_items.append( + {"mm_hash": item_uuid, "transfer_id": transfer_id} + ) rewritten += 1 new_messages.append({**msg, "content": new_content}) if not rewritten: return req_data logger.info("Rewrote %d image item(s) as metadata references", rewritten) - return {**req_data, "messages": new_messages} + rewritten_request = {**req_data, "messages": new_messages} + if transfer_items: + ec_transfer_params = dict(req_data.get("ec_transfer_params") or {}) + ec_transfer_params["ec_items"] = transfer_items + rewritten_request["ec_transfer_params"] = ec_transfer_params + return rewritten_request def extract_mm_items(request_data: dict) -> list[dict]: @@ -165,6 +185,7 @@ async def fanout_encoder_primer( orig_request: dict, e_urls: list[str], req_id: str, + consumer_zmq: str | None = None, ) -> dict[int, dict]: """ 1. Build one request *per MM item* with all text removed. @@ -187,6 +208,7 @@ async def fanout_encoder_primer( tasks = [] item_uuids: dict[int, str] = {} + item_transfer_ids: dict[int, str] = {} item_meta: dict[int, dict] = {} # Round-robin over encode servers to distribute load a bit @@ -204,6 +226,8 @@ async def fanout_encoder_primer( item_uuid = None if NO_REWRITE else content_uuid(item) if item_uuid is not None: item_uuids[idx] = item_uuid + transfer_id = uuid.uuid4().hex + item_transfer_ids[idx] = transfer_id encoder_req = { # You *may* need to keep additional fields @@ -220,6 +244,11 @@ async def fanout_encoder_primer( # once the prompt is encoded and its embeddings are published. "stream": False, } + if consumer_zmq is not None: + encoder_req["ec_transfer_params"] = { + "consumer_zmq": consumer_zmq, + "ec_items": [{"mm_hash": item_uuid, "transfer_id": transfer_id}], + } tasks.append( encode_session.post( f"{target_url}/v1/chat/completions", @@ -269,7 +298,11 @@ async def fanout_encoder_primer( reported = [] if reported and idx in item_uuids: # One item per encoder request, so the first entry is this item's. - item_meta[idx] = {**reported[0], "mm_hash": item_uuids[idx]} + item_meta[idx] = { + **reported[0], + "mm_hash": item_uuids[idx], + "transfer_id": item_transfer_ids[idx], + } logger.info( "[%s] All %d encoder requests completed successfully", req_id, len(mm_items) @@ -435,48 +468,85 @@ async def on_shutdown() -> None: ############################################################################### +async def prepare_for_decode( + req_data: dict, + req_id: str, + e_urls: list[str], + p_url: str, + consumer_zmq: str | None, +) -> tuple[dict, float, float]: + """Encode, rewrite and prefill, returning the body to send to decode. + + `req_data` is left untouched so a retry starts from the original media + rather than from a body whose images are already metadata references. + """ + _t0 = time.perf_counter() + item_meta = await fanout_encoder_primer(req_data, e_urls, req_id, consumer_zmq) + _t1 = time.perf_counter() + prepared = req_data if NO_REWRITE else rewrite_for_decode(req_data, item_meta) + _t2 = time.perf_counter() + prepared = await maybe_prefill(prepared, p_url, req_id) + return prepared, _t1 - _t0, _t2 - _t1 + + async def forward_non_stream( - req_data: dict, req_id: str, e_urls: list[str], p_url: str, d_url: str + req_data: dict, + req_id: str, + e_urls: list[str], + p_url: str, + d_url: str, + consumer_zmq: str | None, ) -> dict: try: - # Step 1: Process through Encoder instance (if has MM input) - _t0 = time.perf_counter() - item_meta = await fanout_encoder_primer(req_data, e_urls, req_id) - _t1 = time.perf_counter() - req_data = req_data if NO_REWRITE else rewrite_for_decode(req_data, item_meta) - _t2 = time.perf_counter() - - # Step 2: Process through Prefill instance - req_data = await maybe_prefill(req_data, p_url, req_id) - - # Step 3: Process through Decode instance - logger.info("[%s] Forwarding to decode: %s", req_id, d_url) - headers = {"x-request-id": req_id} - - # Non-streaming response - async with decode_session.post( - f"{d_url}/v1/chat/completions", json=req_data, headers=headers - ) as resp: - if resp.status >= 400: - detail = await resp.text() - logger.error( - "[%s] Decode request returned status %s: %s", - req_id, - resp.status, - detail, - ) - raise HTTPException(status_code=resp.status, detail=detail) - out = await resp.json() - _t3 = time.perf_counter() - logger.info( - "STAGE %s encode=%.1f rewrite=%.1f decode=%.1f total=%.1f", - "no-rewrite" if NO_REWRITE else "rewrite", - (_t1 - _t0) * 1e3, - (_t2 - _t1) * 1e3, - (_t3 - _t2) * 1e3, - (_t3 - _t0) * 1e3, + for attempt in range(DECODE_RETRIES + 1): + _t0 = time.perf_counter() + prepared, encode_s, rewrite_s = await prepare_for_decode( + req_data, req_id, e_urls, p_url, consumer_zmq ) - return out + _t2 = time.perf_counter() + + logger.info("[%s] Forwarding to decode: %s", req_id, d_url) + headers = {"x-request-id": req_id} + + async with decode_session.post( + f"{d_url}/v1/chat/completions", json=prepared, headers=headers + ) as resp: + if resp.status >= 400: + detail = await resp.text() + # 500 is the decoder's retryable internal error, which + # includes an encoder embedding it could not obtain. Redoing + # the encode publishes the item again. + if resp.status == 500 and attempt < DECODE_RETRIES: + logger.warning( + "[%s] Decode returned 500, re-encoding and retrying " + "(attempt %d/%d): %s", + req_id, + attempt + 1, + DECODE_RETRIES, + detail[:200], + ) + continue + logger.error( + "[%s] Decode request returned status %s: %s", + req_id, + resp.status, + detail, + ) + raise HTTPException(status_code=resp.status, detail=detail) + out = await resp.json() + _t3 = time.perf_counter() + logger.info( + "STAGE %s encode=%.1f rewrite=%.1f decode=%.1f total=%.1f " + "attempt=%d", + "no-rewrite" if NO_REWRITE else "rewrite", + encode_s * 1e3, + rewrite_s * 1e3, + (_t3 - _t2) * 1e3, + (_t3 - _t0) * 1e3, + attempt, + ) + return out + raise HTTPException(status_code=500, detail="Decode failed after re-encoding") except HTTPException: raise @@ -486,47 +556,63 @@ async def forward_non_stream( async def forward_stream( - req_data: dict, req_id: str, e_urls: list[str], p_url: str, d_url: str + req_data: dict, + req_id: str, + e_urls: list[str], + p_url: str, + d_url: str, + consumer_zmq: str | None, ) -> AsyncIterator[str]: try: - # Step 1: Process through Encoder instance (if has MM input) - _t0 = time.perf_counter() - item_meta = await fanout_encoder_primer(req_data, e_urls, req_id) - _t1 = time.perf_counter() - req_data = req_data if NO_REWRITE else rewrite_for_decode(req_data, item_meta) - _t2 = time.perf_counter() - - # Step 2: Process through Prefill instance - req_data = await maybe_prefill(req_data, p_url, req_id) - - # Step 3: Process through Decode instance - logger.info("[%s] Starting streaming from decode: %s", req_id, d_url) - headers = {"x-request-id": req_id} - - # Streaming response - _first = None - async with decode_session.post( - f"{d_url}/v1/chat/completions", - json=req_data, - headers=headers, - ) as resp: - resp.raise_for_status() - async for chunk in resp.content.iter_chunked(1024): - if chunk: - if _first is None: - _first = time.perf_counter() - yield chunk.decode("utf-8", errors="ignore") - _t3 = time.perf_counter() + for attempt in range(DECODE_RETRIES + 1): + _t0 = time.perf_counter() + prepared, encode_s, rewrite_s = await prepare_for_decode( + req_data, req_id, e_urls, p_url, consumer_zmq + ) + _t2 = time.perf_counter() - logger.info( - "STAGE %s encode=%.1f rewrite=%.2f decode_ttfb=%.1f decode_total=%.1f", - "no-rewrite" if NO_REWRITE else "rewrite", - (_t1 - _t0) * 1e3, - (_t2 - _t1) * 1e3, - ((_first or _t3) - _t2) * 1e3, - (_t3 - _t2) * 1e3, - ) - logger.info("[%s] Streaming completed", req_id) + logger.info("[%s] Starting streaming from decode: %s", req_id, d_url) + headers = {"x-request-id": req_id} + + _first = None + async with decode_session.post( + f"{d_url}/v1/chat/completions", + json=prepared, + headers=headers, + ) as resp: + # Retry only before the first chunk: once anything reached the + # client the response cannot be replaced. + if resp.status == 500 and attempt < DECODE_RETRIES: + detail = await resp.text() + logger.warning( + "[%s] Decode returned 500 before streaming, re-encoding " + "and retrying (attempt %d/%d): %s", + req_id, + attempt + 1, + DECODE_RETRIES, + detail[:200], + ) + continue + resp.raise_for_status() + async for chunk in resp.content.iter_chunked(1024): + if chunk: + if _first is None: + _first = time.perf_counter() + yield chunk.decode("utf-8", errors="ignore") + _t3 = time.perf_counter() + + logger.info( + "STAGE %s encode=%.1f rewrite=%.2f decode_ttfb=%.1f " + "decode_total=%.1f attempt=%d", + "no-rewrite" if NO_REWRITE else "rewrite", + encode_s * 1e3, + rewrite_s * 1e3, + ((_first or _t3) - _t2) * 1e3, + (_t3 - _t2) * 1e3, + attempt, + ) + logger.info("[%s] Streaming completed", req_id) + return except HTTPException: logger.exception("[%s] HTTPException in forward_stream", req_id) @@ -551,16 +637,22 @@ async def chat_completions(request: Request): e_urls = app.state.e_urls # we want the full list for fan-out p_url = random.choice(app.state.p_urls) if app.state.p_urls else None - d_url = random.choice(app.state.d_urls) + decode_index = random.randrange(len(app.state.d_urls)) + d_url = app.state.d_urls[decode_index] + consumer_zmq = ( + app.state.d_ec_urls[decode_index] if app.state.d_ec_urls else None + ) is_streaming = req_data.get("stream", False) if is_streaming: return StreamingResponse( - forward_stream(req_data, req_id, e_urls, p_url, d_url), + forward_stream(req_data, req_id, e_urls, p_url, d_url, consumer_zmq), media_type="text/event-stream", ) - result = await forward_non_stream(req_data, req_id, e_urls, p_url, d_url) + result = await forward_non_stream( + req_data, req_id, e_urls, p_url, d_url, consumer_zmq + ) return JSONResponse(content=result) except HTTPException: @@ -725,8 +817,8 @@ async def stop_profile(request: Request): "--prefill-servers-urls", required=True, help=( - 'Comma-separated prefill URLs ("http://p1:8003,http://p2:8004") ', - 'to enable E->P->D, set "disable" or "none" to enable E->PD', + 'Comma-separated prefill URLs ("http://p1:8003,http://p2:8004") ' + 'to enable E->P->D, set "disable" or "none" to enable E->PD' ), ) parser.add_argument( @@ -734,15 +826,42 @@ async def stop_profile(request: Request): required=True, help='Comma-separated decode URLs ("http://d1:8005,http://d2:8006")', ) + parser.add_argument( + "--decode-retries", + type=int, + default=1, + help=( + "Re-encode and re-send when decode returns 500, which is its " + "retryable internal error (an undeliverable encoder embedding " + "among them). 0 disables." + ), + ) + parser.add_argument( + "--decode-ec-transfer-zmq-addrs", + default="", + help=( + "Comma-separated Mooncake EC consumer ZMQ addresses, aligned " + "with --decode-servers-urls. Required when the decoders use the " + "Mooncake EC connector." + ), + ) args = parser.parse_args() NO_REWRITE = args.no_rewrite + DECODE_RETRIES = max(0, args.decode_retries) app.state.e_urls = [ u.strip() for u in args.encode_servers_urls.split(",") if u.strip() ] app.state.d_urls = [ u.strip() for u in args.decode_servers_urls.split(",") if u.strip() ] + app.state.d_ec_urls = [ + u.strip() for u in args.decode_ec_transfer_zmq_addrs.split(",") if u.strip() + ] + if app.state.d_ec_urls and len(app.state.d_ec_urls) != len(app.state.d_urls): + parser.error( + "--decode-ec-transfer-zmq-addrs must contain one address per decode server" + ) # handle prefill instances if args.prefill_servers_urls.lower() in ("disable", "none", ""): app.state.p_urls = [] diff --git a/tests/v1/core/test_encoder_cache_manager.py b/tests/v1/core/test_encoder_cache_manager.py index e56bcbf5c63a..07fbaeca8bc4 100644 --- a/tests/v1/core/test_encoder_cache_manager.py +++ b/tests/v1/core/test_encoder_cache_manager.py @@ -167,6 +167,38 @@ def test_get_freed_mm_hashes_clears_freed_list(): assert manager.get_freed_mm_hashes() == [] +def test_referencing_a_freeable_entry_protects_it_from_eviction(): + """A waiting request can hold an entry it is not scheduled for yet. + + Entries with no referent are evictable. A request that is deferred (e.g. + by an EC connector waiting on another item) still needs the items it + already has, and a connector that hands out one transfer per request + cannot always produce the same item a second time. + """ + manager = EncoderCacheManager(cache_size=10) + owner = MockRequest("owner", ["a"], [5]) + waiter = MockRequest("waiter", ["a"], [5]) + newcomer = MockRequest("newcomer", ["b"], [6]) + + manager.allocate(owner, 0) + manager.free_encoder_input(owner, 0) + assert "a" in manager.freeable + + # The waiting request references it; it is no longer reclaimable. + assert manager.check_and_update_cache(waiter, 0) + assert "a" not in manager.freeable + + # 'b' no longer fits, and 'a' must not be taken away to make room. + assert not manager.can_allocate(newcomer, 0, int(1e9), 0) + assert "a" in manager.cached + assert manager.get_freed_mm_hashes() == [] + + # Once the waiter is done with it, the entry is reclaimable again. + manager.free_encoder_input(waiter, 0) + assert manager.can_allocate(newcomer, 0, int(1e9), 0) + assert manager.get_freed_mm_hashes() == ["a"] + + def test_reallocated_hash_is_not_reported_as_freed(): manager = EncoderCacheManager(cache_size=8) req_a = MockRequest("reqA", ["a"], [4]) diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index c0932d975c18..da0460ac2f4f 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -5334,6 +5334,68 @@ def test_free_encoder_inputs_respects_unconfirmed_placeholders(): assert manager.get_cached_input_ids(request) == set() +def test_free_encoder_inputs_notifies_the_ec_connector(): + """The connector learns the item is consumed, not just that the request ended. + + Connectors hold per-item transfer state (remote buffers, reservations); + waiting for `request_finished` pins it for the whole generation. + """ + scheduler = create_scheduler(model="llava-hf/llava-1.5-7b-hf") + mm_start_pos, mm_length = 50, 100 + request = create_requests( + num_requests=1, + num_tokens=mm_start_pos + mm_length + 10, + mm_positions=[[PlaceholderRange(offset=mm_start_pos, length=mm_length)]], + )[0] + scheduler.encoder_cache_manager.allocate(request, 0) + scheduler.ec_connector = Mock() + + request.num_computed_tokens = mm_start_pos + mm_length - 1 + scheduler._free_encoder_inputs(request) + scheduler.ec_connector.update_state_after_free.assert_not_called() + + request.num_computed_tokens = mm_start_pos + mm_length + scheduler._free_encoder_inputs(request) + scheduler.ec_connector.update_state_after_free.assert_called_once_with(request, 0) + + +def test_unavailable_encoder_input_fails_the_request_as_retryable(): + """An encoder input the connector gave up on must end the request. + + `FinishReason.ERROR` is the retryable channel the KV connector already uses + for load failures, so the caller can re-issue; deferring instead parked the + request until the client timed out. + """ + scheduler = create_scheduler(model="llava-hf/llava-1.5-7b-hf") + request = create_requests( + num_requests=1, + num_tokens=160, + mm_positions=[[PlaceholderRange(offset=50, length=100)]], + )[0] + scheduler.add_request(request) + scheduler.ec_connector = Mock() + scheduler.ec_connector.take_unavailable_requests.return_value = {request.request_id} + # `request_finished` reports (delay_free, params) for the request teardown. + scheduler.ec_connector.request_finished.return_value = (False, None) + + scheduler_output = scheduler.schedule() + outputs = scheduler.update_from_output( + scheduler_output, + ModelRunnerOutput( + req_ids=[request.request_id], + req_id_to_index={request.request_id: 0}, + sampled_token_ids=[[]], + logprobs=None, + prompt_logprobs_dict={}, + pooler_output=[], + ), + ) + + assert request.status == RequestStatus.FINISHED_ERROR + engine_outputs = outputs[0].outputs + assert [o.finish_reason for o in engine_outputs] == [FinishReason.ERROR] + + def test_free_encoder_inputs_defers_for_eagle_lookahead(): """With EAGLE speculative decoding, the encoder input is retained one extra position so the drafter's +1 look-ahead mm-embedding gather (which reads one @@ -5531,9 +5593,9 @@ def test_ec_connector_ensure_cache_available_defers_request(use_kv_connector): # ensure_cache_available must have been called with (request, num_computed_tokens=0) # for a brand-new request that has no cached tokens yet. - scheduler.ec_connector.ensure_cache_available.assert_called_once_with( - request_deferred, 0 - ) + ensure_call = scheduler.ec_connector.ensure_cache_available.call_args + assert ensure_call.args[:2] == (request_deferred, 0) + assert not ensure_call.args[2] # Deferred request must NOT be scheduled assert request_deferred.request_id not in output.num_scheduled_tokens _assert_right_encoder_cache_allocated(scheduler, expected_total_allocated=0) @@ -5565,6 +5627,31 @@ def test_ec_connector_ensure_cache_available_defers_request(use_kv_connector): _assert_right_encoder_inputs(output, expected_total_reqs=0) +def test_ec_connector_defers_running_request_for_async_reload(): + scheduler = create_scheduler( + model="llava-hf/llava-1.5-7b-hf", + max_num_batched_tokens=32, + use_ec_connector=True, + ec_role="ec_consumer", + ) + request = create_requests( + num_requests=1, + num_tokens=128, + mm_positions=[[PlaceholderRange(offset=48, length=32)]], + req_ids=["request"], + )[0] + scheduler.ec_connector.ensure_cache_available = Mock(side_effect=[True, False]) + + scheduler.add_request(request) + first_output = scheduler.schedule() + assert first_output.num_scheduled_tokens[request.request_id] == 32 + + second_output = scheduler.schedule() + assert request.request_id not in second_output.num_scheduled_tokens + ensure_call = scheduler.ec_connector.ensure_cache_available.call_args + assert ensure_call.args[:2] == (request, 32) + + def test_ec_connector_pending_prefetch_only_checks_future_mm_features(): """Test that future mm feature filtering only yields features beyond the computed token frontier. diff --git a/tests/v1/ec_connector/integration/README.md b/tests/v1/ec_connector/integration/README.md index 5feee2ae11ee..1b3bac1dd3e2 100644 --- a/tests/v1/ec_connector/integration/README.md +++ b/tests/v1/ec_connector/integration/README.md @@ -124,18 +124,6 @@ Quick sanity check: - Safe to run multiple times (idempotent) - We setup the PD disagg part with NixlConnector. Please read details about EPD in `examples/disaggregated/disaggregated_encoder/README.md` -## ECMooncakeConnector (TransferEngine) smoke test - -Two-process transfer over Mooncake (no full vLLM serve, no HF model download): - -```bash -cd vllm -PYTHONPATH=. python tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py -``` - -Requires: **2+ CUDA GPUs**, `mooncake-transfer-engine`, `pyzmq`, `httpx`, `fastapi`, `uvicorn`. -Optional: `MOONCAKE_EC_PROTOCOL=rdma` or `=tcp` (default in test mocks is `tcp`) to match your cluster. - ## Requirements - Multiple GPUs (3 for 1E+1P+1D, 2 for 1E+1PD, 1 for baseline) diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh index d68040baa098..a30e942422f6 100755 --- a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -13,8 +13,6 @@ # MODEL HF model id (default: Qwen/Qwen2.5-VL-3B-Instruct) # GPU_SINGLE / GPU_E / GPU_PD GPU ids (defaults 0 / 1 / 2) # ENDPOINT_PORT, ENCODE_PORT, PREFILL_DECODE_PORT -# EC_MOONCAKE_REGISTRY_PORT HTTP registry on encoder (default 19018) -# EC_REGISTRY_HOST Host PD uses to query registry (default 127.0.0.1) # MOONCAKE_EC_PROTOCOL rdma | tcp (default rdma) # USE_MM_PROMPTS 1 (default) or 0 for text-only quick sanity # TIMEOUT_SECONDS wait_for_server timeout (default 1200) @@ -25,6 +23,7 @@ set -euo pipefail GIT_ROOT=$(git rev-parse --show-toplevel) cd "$GIT_ROOT" || exit 1 export PYTHONPATH="${GIT_ROOT}:${PYTHONPATH:-}" +PYTHON_BIN="${PYTHON_BIN:-${GIT_ROOT}/.venv/bin/python}" MODEL="${MODEL:-Qwen/Qwen2.5-VL-3B-Instruct}" USE_MM_PROMPTS="${USE_MM_PROMPTS:-1}" @@ -42,11 +41,10 @@ PREFILL_DECODE_PORT="${PREFILL_DECODE_PORT:-19537}" ENDPOINT_PORT="${ENDPOINT_PORT:-10002}" BASELINE_PORT="${BASELINE_PORT:-10003}" -EC_MOONCAKE_REGISTRY_PORT="${EC_MOONCAKE_REGISTRY_PORT:-19018}" -EC_REGISTRY_HOST="${EC_REGISTRY_HOST:-127.0.0.1}" +EC_MOONCAKE_RESERVATION_PORT="${EC_MOONCAKE_RESERVATION_PORT:-19019}" MOONCAKE_EC_PROTOCOL="${MOONCAKE_EC_PROTOCOL:-rdma}" -EC_REGISTRY_URL="http://${EC_REGISTRY_HOST}:${EC_MOONCAKE_REGISTRY_PORT}" -export EC_MOONCAKE_REGISTRY_PORT EC_REGISTRY_HOST MOONCAKE_EC_PROTOCOL EC_REGISTRY_URL +export EC_MOONCAKE_RESERVATION_PORT +export MOONCAKE_EC_PROTOCOL LOG_PATH="${LOG_PATH:-/tmp}" BASELINE_FILE="${BASELINE_FILE:-/tmp/vllm_epd_mooncake_baseline.txt}" @@ -54,33 +52,30 @@ TIMEOUT_SECONDS="${TIMEOUT_SECONDS:-1200}" mkdir -p "$LOG_PATH" -if command -v vllm &>/dev/null; then - VLLM_SERVE=(vllm serve) -else - VLLM_SERVE=(python -m vllm.entrypoints.cli.main serve) -fi +VLLM_SERVE=("$PYTHON_BIN" -m vllm.entrypoints.cli.main serve) -ENC_EC_JSON=$(python <"${LOG_PATH}/mooncake_epd_proxy.log" 2>&1 & PIDS+=($!) @@ -188,7 +187,7 @@ run_epd_mooncake() { curl -s "http://127.0.0.1:${ENDPOINT_PORT}/health" || true echo "" - python "${GIT_ROOT}/tests/v1/ec_connector/integration/test_epd_correctness.py" \ + "$PYTHON_BIN" "${GIT_ROOT}/tests/v1/ec_connector/integration/test_epd_correctness.py" \ --service_url "http://localhost:$ENDPOINT_PORT" \ --model_name "$MODEL" \ --mode disagg \ diff --git a/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py b/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py deleted file mode 100644 index b8ed8446c8ee..000000000000 --- a/tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py +++ /dev/null @@ -1,211 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -""" -End-to-end test: ECMooncakeConnector producer (GPU0) publishes EC metadata and -serves pull requests; consumer (GPU1) fetches layout over HTTP and pulls the -tensor via Mooncake TransferEngine + ZMQ. - -Requires: 2+ CUDA GPUs, mooncake-transfer-engine, pyzmq, httpx, fastapi, uvicorn. - -Protocol: ``mooncake_protocol`` defaults to ``tcp`` in mocks unless you set -``MOONCAKE_EC_PROTOCOL=rdma`` (matches the connector protocol configuration). -Example RDMA run:: - - MOONCAKE_EC_PROTOCOL=rdma PYTHONPATH=. python \ - tests/v1/ec_connector/integration/test_ec_mooncake_transfer_e2e.py -""" - -from __future__ import annotations - -import multiprocessing as mp -import os -import time -from typing import Any -from unittest.mock import Mock - -import pytest -import torch - -try: - import zmq # noqa: F401 -except ImportError as e: - raise SystemExit("pyzmq is required: pip install pyzmq") from e -try: - import mooncake # noqa: F401 -except ImportError as e: - raise SystemExit("mooncake-transfer-engine is required") from e - -from vllm.config import VllmConfig -from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole -from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( - ECMooncakeConnector, - ECMooncakeConnectorMetadata, - ECMooncakeLoadSpec, -) - - -def _find_free_port() -> int: - import socket - - s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - s.bind(("127.0.0.1", 0)) - _, port = s.getsockname() - s.close() - return int(port) - - -def _mock_vllm_producer(registry_port: int) -> Mock: - cfg = Mock(spec=VllmConfig) - cfg.parallel_config = Mock() - cfg.parallel_config.tensor_parallel_size = 1 - cfg.parallel_config.pipeline_parallel_size = 1 - cfg.ec_transfer_config = Mock() - cfg.ec_transfer_config.is_ec_producer = True - cfg.ec_transfer_config.is_ec_consumer = False - cfg.ec_transfer_config.ec_buffer_device = "cuda" - cfg.ec_transfer_config.ec_buffer_size = 1e9 - cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), - "registry_http_port": registry_port, - } - return cfg - - -def _mock_vllm_consumer() -> Mock: - cfg = Mock(spec=VllmConfig) - cfg.parallel_config = Mock() - cfg.parallel_config.tensor_parallel_size = 1 - cfg.parallel_config.pipeline_parallel_size = 1 - cfg.ec_transfer_config = Mock() - cfg.ec_transfer_config.is_ec_producer = False - cfg.ec_transfer_config.is_ec_consumer = True - cfg.ec_transfer_config.ec_buffer_device = "cuda" - cfg.ec_transfer_config.ec_buffer_size = 1e9 - cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), - "remote_registry_url": "http://unused-on-worker", - } - return cfg - - -def _producer_entry( - mm_hash: str, - registry_port: int, - ready: Any, - done: Any, - barrier: Any, -) -> None: - os.environ["CUDA_VISIBLE_DEVICES"] = "0" - cfg = _mock_vllm_producer(registry_port) - conn = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) - torch.manual_seed(12345) - tensor = torch.randn(8, 64, device="cuda", dtype=torch.float32) - cache = {mm_hash: tensor} - conn.save_caches(cache, mm_hash) - ready.put("ok") - barrier.wait(timeout=120) - # Hold process until consumer finishes transfer - done.wait(timeout=180) - - -def _consumer_entry( - mm_hash: str, - registry_url: str, - barrier: Any, - result_queue: Any, -) -> None: - os.environ["CUDA_VISIBLE_DEVICES"] = "1" - barrier.wait(timeout=120) - import httpx - - url = f"{registry_url.rstrip('/')}/ec/info/{mm_hash}" - for _ in range(60): - try: - r = httpx.get(url, timeout=2.0) - if r.status_code == 200: - break - except httpx.HTTPError: - pass - time.sleep(0.5) - else: - result_queue.put({"ok": False, "err": "registry never ready"}) - return - data = r.json() - spec = ECMooncakeLoadSpec( - mm_hash=mm_hash, - num_token=1, - nbytes=int(data["nbytes"]), - shape=tuple(int(x) for x in data["shape"]), - dtype=str(data["dtype"]), - producer_zmq=str(data["producer_zmq"]), - lease_id=str(data["lease_id"]), - ) - meta = ECMooncakeConnectorMetadata() - meta.add_load(spec) - cfg = _mock_vllm_consumer() - conn = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) - conn.bind_connector_metadata(meta) - enc: dict[str, torch.Tensor] = {} - try: - conn.start_load_caches(enc) - except Exception as e: - result_queue.put({"ok": False, "err": repr(e)}) - return - got = enc.get(mm_hash) - if got is None: - result_queue.put({"ok": False, "err": "missing tensor"}) - return - torch.manual_seed(12345) - expected = torch.randn(8, 64, device="cuda", dtype=torch.float32) - max_diff = (got.cpu() - expected.cpu()).abs().max().item() - result_queue.put({"ok": True, "max_diff": max_diff}) - - -@pytest.mark.skipif( - torch.accelerator.device_count() < 2, - reason="Requires at least 2 CUDA devices", -) -def test_ec_mooncake_two_process_transfer(): - """Producer on cuda:0 and consumer on cuda:1 transfer one EC tensor.""" - mm_hash = "e2e_mm_test_hash" - registry_port = _find_free_port() - registry_url = f"http://127.0.0.1:{registry_port}" - - ctx = mp.get_context("spawn") - ready: mp.Queue = ctx.Queue() - result: mp.Queue = ctx.Queue() - done = ctx.Event() - barrier = ctx.Barrier(2) - prod = ctx.Process( - target=_producer_entry, - args=(mm_hash, registry_port, ready, done, barrier), - daemon=True, - ) - cons = ctx.Process( - target=_consumer_entry, - args=(mm_hash, registry_url, barrier, result), - daemon=True, - ) - prod.start() - assert ready.get(timeout=120) == "ok" - cons.start() - cons.join(timeout=180) - done.set() - prod.join(timeout=30) - - assert not cons.is_alive(), "consumer process hung" - assert cons.exitcode == 0, f"consumer exit {cons.exitcode}" - out = result.get(timeout=1) - assert out["ok"], out.get("err", out) - assert out["max_diff"] < 1e-4, f"tensor mismatch max_diff={out['max_diff']}" - - -def _main() -> None: - if torch.accelerator.device_count() < 2: - raise SystemExit("Need at least 2 CUDA devices for this e2e test.") - test_ec_mooncake_two_process_transfer() - print("ECMooncake two-process transfer e2e: PASSED") - - -if __name__ == "__main__": - _main() diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 2c21f41484fa..250c268a769b 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -1,9 +1,10 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Unit tests for ECMooncakeConnector and its HTTP registry.""" +"""Unit tests for ECMooncakeConnector.""" from __future__ import annotations +import copy import ctypes import socket import time @@ -11,18 +12,21 @@ from types import SimpleNamespace from unittest.mock import Mock, patch -import httpx import pytest import torch +import zmq from vllm.config import VllmConfig from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole from vllm.distributed.ec_transfer.ec_connector.factory import ECConnectorFactory from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( + _LEASE_TTL_SECONDS, ECMooncakeConnector, ECMooncakeConnectorMetadata, ECMooncakeLoadSpec, - ECMooncakeRegistryServer, + ECMooncakePushSpec, + ECMooncakeWorkerMetadata, + _ConsumerPoolAllocation, _ContiguousAllocator, ) from vllm.v1.core.sched.output import SchedulerOutput @@ -33,6 +37,7 @@ class CopyingFakeTransferEngine: def __init__(self, *args, **kwargs): self.registered: set[int] = set() + self.regions: dict[int, int] = {} self.register_calls: list[list[int]] = [] self.unregister_calls: list[int] = [] self.batch_unregister_calls: list[list[int]] = [] @@ -54,20 +59,33 @@ def batch_transfer_sync_write( def batch_register_memory(self, buffer_addresses, capacities) -> int: addresses = [int(addr) for addr in buffer_addresses] + lengths = [int(length) for length in capacities] + # A real Transfer Engine refuses overlapping memory regions, so model + # that here: registering a range that intersects a live one fails. + regions = dict(self.regions) + for address, length in zip(addresses, lengths): + for other, other_length in regions.items(): + if address < other + other_length and other < address + length: + return 1 + regions[address] = length self.register_calls.append(addresses) self.registered.update(addresses) + self.regions = regions return 0 def unregister_memory(self, buffer_address) -> int: address = int(buffer_address) self.unregister_calls.append(address) self.registered.discard(address) + self.regions.pop(address, None) return 0 def batch_unregister_memory(self, buffer_addresses) -> int: addresses = [int(addr) for addr in buffer_addresses] self.batch_unregister_calls.append(addresses) self.registered.difference_update(addresses) + for address in addresses: + self.regions.pop(address, None) return 0 @@ -79,6 +97,19 @@ def _find_free_port() -> int: return int(port) +def _wait_for_worker_io( + connector: ECMooncakeConnector, timeout: float = 5.0 +) -> ECMooncakeWorkerMetadata: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + meta = connector.build_connector_worker_meta() + assert isinstance(meta, ECMooncakeWorkerMetadata) + if not meta.pending_loads and not meta.pending_saves: + return meta + time.sleep(0.01) + raise TimeoutError("EC Mooncake worker I/O did not finish") + + @pytest.fixture def mock_vllm_config_producer(): config = Mock(spec=VllmConfig) @@ -92,7 +123,6 @@ def mock_vllm_config_producer(): config.ec_transfer_config.ec_buffer_size = 1e9 config.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "registry_http_port": 19018, } return config @@ -110,7 +140,7 @@ def mock_vllm_config_consumer(): config.ec_transfer_config.ec_buffer_size = 1e9 config.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "remote_registry_url": "http://127.0.0.1:19018", + "reservation_zmq_port": 19019, } return config @@ -130,51 +160,10 @@ def patch_ec_mooncake_deps(): "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector.get_ip", return_value="127.0.0.1", ), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector.is_local_first_rank", - return_value=True, - ), ): yield -class TestECMooncakeRegistryServer: - def test_publish_and_lookup(self): - port = _find_free_port() - registry = ECMooncakeRegistryServer("127.0.0.1", port) - registry.start() - try: - payload = { - "nbytes": 128, - "shape": [4, 8], - "dtype": "float32", - "producer_zmq": "tcp://127.0.0.1:9999", - } - registry.publish("hash_a", payload) - r = httpx.get(f"http://127.0.0.1:{port}/ec/info/hash_a", timeout=2.0) - assert r.status_code == 200 - data = r.json() - lease_id = data.pop("lease_id") - assert data == payload - assert registry.consume_lease("hash_a", lease_id) - r404 = httpx.get(f"http://127.0.0.1:{port}/ec/info/missing", timeout=2.0) - assert r404.status_code == 404 - finally: - registry.shutdown() - - def test_unpublish_removes_entry(self): - port = _find_free_port() - registry = ECMooncakeRegistryServer("127.0.0.1", port) - registry.start() - try: - registry.publish("h", {"nbytes": 1, "shape": [], "dtype": "float32"}) - registry.unpublish("h") - r = httpx.get(f"http://127.0.0.1:{port}/ec/info/h", timeout=2.0) - assert r.status_code == 404 - finally: - registry.shutdown() - - class TestECMooncakeFactory: def test_factory_registers_connector(self): cls = ECConnectorFactory.get_connector_class( @@ -199,18 +188,6 @@ def test_reuses_and_coalesces_contiguous_regions(self): class TestECMooncakeConnectorValidation: - def test_consumer_scheduler_requires_remote_registry( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - with ( - patch_ec_mooncake_deps(), - pytest.raises(ValueError, match="remote_registry_url"), - ): - ECMooncakeConnector(mock_vllm_config_consumer, ECConnectorRole.SCHEDULER) - def test_rejects_tensor_parallel_gt_one(self, mock_vllm_config_producer): mock_vllm_config_producer.parallel_config.tensor_parallel_size = 2 with ( @@ -221,56 +198,416 @@ def test_rejects_tensor_parallel_gt_one(self, mock_vllm_config_producer): class TestECMooncakeSchedulerMetadata: - def test_has_cache_item_queries_registry( + def test_missing_push_event_is_tracked( self, mock_vllm_config_consumer, mock_request_with_3_mm ): - port = _find_free_port() - registry = ECMooncakeRegistryServer("127.0.0.1", port) - registry.start() - try: - mm_hash = mock_request_with_3_mm.mm_features[0].identifier - registry.publish( - mm_hash, - { - "nbytes": 64, - "shape": [2, 4], - "dtype": "float32", - "producer_zmq": "tcp://127.0.0.1:1", - }, - ) - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "remote_registry_url" - ] = f"http://127.0.0.1:{port}" - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19019, + } + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + with patch.object(scheduler, "_drain_push_notifications"): + assert not scheduler.ensure_cache_available(request, 0) + mm_hash = request.mm_features[0].identifier + assert scheduler._consumer_scheduler_metrics["missing_event"] == 1 + assert mm_hash in scheduler._consumer_missing_since + finally: + scheduler.shutdown() + + def test_item_with_no_transfer_in_flight_is_reported_as_stalled( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """A push that never arrives must not wait silently forever.""" + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19019, + "push_wait_timeout_s": 0, + } + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + mm_hash = request.mm_features[0].identifier + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + with patch.object(scheduler, "_drain_push_notifications"): + assert not scheduler.ensure_cache_available(request, 0) + assert not scheduler.ensure_cache_available(request, 0) + assert scheduler._consumer_scheduler_metrics["stalled"] == 1 + assert mm_hash in scheduler._stalled_hashes + # The stall is reported once, not once per scheduling pass. + with patch.object(scheduler, "_drain_push_notifications"): + assert not scheduler.ensure_cache_available(request, 0) + assert scheduler._consumer_scheduler_metrics["stalled"] == 1 + finally: + scheduler.shutdown() + + def test_pending_observation_ends_with_last_spec(self, mock_vllm_config_consumer): + specs = [ + ECMooncakeLoadSpec( + mm_hash="hash", + num_token=1, + nbytes=32, + shape=(8,), + dtype="float32", + pushed=True, + transfer_id=f"transfer-{index}", + ) + for index in range(2) + ] + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + for spec in specs: + scheduler._index_pending_spec(spec) + scheduler._pop_pending_spec("transfer-0") + assert "hash" in scheduler._consumer_pending_since + scheduler._pop_pending_spec("transfer-1") + assert "hash" not in scheduler._consumer_pending_since + finally: + scheduler.shutdown() + + def test_local_cache_hit_keeps_the_transfer( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """`local_cache_hashes` is a snapshot, so the transfer must survive it. + + Cancelling on a cache hit strands the request when the entry is + evicted before it is scheduled: the item is then unreachable. + """ + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + mm_hash = request.mm_features[0].identifier + request.ec_transfer_params = { + "ec_items": [ + {"mm_hash": mm_hash, "transfer_id": "request-transfer"} + ] + } + scheduler._index_pending_spec( + ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=0, + nbytes=16, + shape=(4,), + dtype="float32", + pushed=True, + transfer_id="request-transfer", + reservation_id="reservation", + ) ) - assert scheduler.has_cache_item(mm_hash) - assert mm_hash in scheduler._pending_specs - spec = scheduler._pending_specs[mm_hash] - assert spec.shape == (2, 4) - assert spec.dtype == "float32" - finally: - registry.shutdown() - - def test_has_cache_item_missing_returns_false( + with ( + patch.object(scheduler, "_drain_push_notifications"), + patch.object(scheduler, "_queue_cancel") as cancel, + ): + assert scheduler.ensure_cache_available(request, 0, {mm_hash}) + cancel.assert_not_called() + assert "request-transfer" in scheduler._pending_specs + + # Once the entry is evicted the request can still get it. + assert not scheduler.ensure_cache_available(request, 0, set()) + assert mm_hash in scheduler._loading_hashes + finally: + scheduler.shutdown() + + def test_consumed_item_releases_its_transfer_immediately( self, mock_vllm_config_consumer, mock_request_with_3_mm ): - port = _find_free_port() - registry = ECMooncakeRegistryServer("127.0.0.1", port) - registry.start() - try: - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "remote_registry_url" - ] = f"http://127.0.0.1:{port}" - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + """The buffer goes back as soon as the item is consumed. + + Holding it until `request_finished` would pin a pool slot for the + whole generation, long after the embedding was used. + """ + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + mm_hash = request.mm_features[0].identifier + request.ec_transfer_params = { + "ec_items": [{"mm_hash": mm_hash, "transfer_id": "consumed-transfer"}] + } + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._index_pending_spec( + ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=0, + nbytes=16, + shape=(4,), + dtype="float32", + pushed=True, + transfer_id="consumed-transfer", + reservation_id="reservation", + ) ) - mm_hash = mock_request_with_3_mm.mm_features[0].identifier - assert not scheduler.has_cache_item(mm_hash) - finally: - registry.shutdown() + with patch.object(scheduler, "_queue_cancel") as cancel: + scheduler.update_state_after_free(request, 0) + cancel.assert_called_once_with("consumed-transfer", "reservation") + assert "consumed-transfer" not in scheduler._pending_specs + finally: + scheduler.shutdown() + + def test_ready_hash_eviction_does_not_strand_a_later_transfer( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """The long-tail stall: an event arrives while the hash is still ready. + + Dropping it as redundant left the next request with no transfer at + all once the encoder cache entry was freed, and nothing could bring + the item back. + """ + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19019, + } + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + mm_hash = request.mm_features[0].identifier + request.ec_transfer_params = { + "ec_items": [{"mm_hash": mm_hash, "transfer_id": "later-transfer"}] + } + event = { + "mm_hash": mm_hash, + "transfer_id": "later-transfer", + "ready": True, + "reservation_id": "later", + "nbytes": 16, + "shape": [4], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._ready_hashes.add(mm_hash) + scheduler._event_zmq_socket = Mock() + scheduler._event_zmq_socket.recv_json.side_effect = [ + event, + zmq.Again(), + ] + with patch.object(scheduler, "_queue_cancel") as cancel: + scheduler._drain_push_notifications() + cancel.assert_not_called() + assert "later-transfer" in scheduler._pending_specs + + # The scheduler frees the encoder cache entry. + scheduler.build_connector_meta( + SimpleNamespace(free_encoder_mm_hashes=[mm_hash]) + ) + assert mm_hash not in scheduler._ready_hashes + + # The request that owns the transfer can still pick it up. + with patch.object(scheduler, "_drain_push_notifications"): + assert not scheduler.ensure_cache_available(request, 0, set()) + assert mm_hash in scheduler._loading_hashes + finally: + scheduler.shutdown() + + def test_item_that_never_arrives_fails_the_request( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """A push that never lands must end the request, not hold it forever. + + The failure is retryable: the caller re-issues, the encode runs again + and produces a fresh transfer. Deferring instead left the request + parked until the client timed out. + """ + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19019, + "push_wait_timeout_s": 0, + } + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + mm_hash = request.mm_features[0].identifier + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + with ( + patch.object(scheduler, "_drain_push_notifications"), + patch.object(scheduler, "_send_control", return_value=None), + ): + assert not scheduler.ensure_cache_available(request, 0, set()) + assert scheduler.take_unavailable_requests() == {request.request_id} + # Draining clears it: the scheduler acts on each id once. + assert scheduler.take_unavailable_requests() == set() + # A re-issued request gets a fresh window rather than the + # expired one, or it would fail before its push could land. + assert mm_hash not in scheduler._consumer_missing_since + finally: + scheduler.shutdown() + + def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """The stall at high concurrency: more needs than transfers. + + Requests sharing an image get one transfer each, but a load consumes + one spec and serves everyone at once. After an eviction the remaining + requests need the item again with no spec left. The receive pool still + holds it, so the reload must come from there. + """ + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19019, + "consumer_buffer_pool_size": 1 << 20, + } + first = mock_request_with_3_mm + first.mm_features = first.mm_features[:1] + mm_hash = first.mm_features[0].identifier + first.ec_transfer_params = { + "ec_items": [{"mm_hash": mm_hash, "transfer_id": "only-transfer"}] + } + second = copy.copy(first) + second.request_id = "second-request" + second.ec_transfer_params = None + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._event_zmq_socket = Mock() + scheduler._event_zmq_socket.recv_json.side_effect = [ + { + "mm_hash": mm_hash, + "transfer_id": "only-transfer", + "ready": True, + "reservation_id": "r0", + "nbytes": 16, + "shape": [4], + "dtype": "float32", + }, + zmq.Again(), + ] + scheduler._drain_push_notifications() + + assert not scheduler.ensure_cache_available(first, 0, set()) + meta = scheduler.build_connector_meta( + SimpleNamespace(free_encoder_mm_hashes=[]) + ) + assert [spec.transfer_id for spec in meta.loads] == ["only-transfer"] + scheduler.update_connector_output( + SimpleNamespace( + ec_connector_worker_meta=ECMooncakeWorkerMetadata( + loaded={mm_hash} + ) + ) + ) + assert not scheduler._pending_specs + + # The encoder cache evicts the entry. + scheduler.build_connector_meta( + SimpleNamespace(free_encoder_mm_hashes=[mm_hash]) + ) + + # The second request has no transfer of its own, and the only + # transfer is spent. It must still be served. + with patch.object(scheduler, "_drain_push_notifications"): + assert scheduler.has_cache_item(mm_hash) + assert not scheduler.ensure_cache_available(second, 0, set()) + assert mm_hash in scheduler._loading_hashes + reload = scheduler.build_connector_meta( + SimpleNamespace(free_encoder_mm_hashes=[]) + ) + assert [spec.local for spec in reload.loads] == [True] + finally: + scheduler.shutdown() + + def test_reclaimed_item_stops_being_offered_as_resident( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """Residency is a mirror of the worker's pool, not a promise.""" + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19019, + "consumer_buffer_pool_size": 1 << 20, + } + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + mm_hash = request.mm_features[0].identifier + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._note_resident( + ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=0, + nbytes=16, + shape=(4,), + dtype="float32", + ) + ) + with patch.object(scheduler, "_drain_push_notifications"): + assert scheduler.has_cache_item(mm_hash) + + scheduler.update_connector_output( + SimpleNamespace( + ec_connector_worker_meta=ECMooncakeWorkerMetadata( + reclaimed={mm_hash} + ) + ) + ) + with patch.object(scheduler, "_drain_push_notifications"): + assert not scheduler.has_cache_item(mm_hash) + assert scheduler._resident_bytes == 0 + finally: + scheduler.shutdown() + + def test_retains_new_completion_while_same_hash_is_loading( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19019, + } + event = { + "mm_hash": "hash", + "transfer_id": "next-transfer", + "ready": True, + "reservation_id": "next", + "nbytes": 64, + "shape": [2, 8], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + scheduler._event_zmq_socket = Mock() + scheduler._event_zmq_socket.recv_json.side_effect = [event, zmq.Again()] + scheduler._loading_hashes.add("hash") + + scheduler._drain_push_notifications() + + assert scheduler._pending_specs["next-transfer"].reservation_id == "next" def test_build_connector_meta_clears_pending( self, mock_vllm_config_consumer, mock_request_with_3_mm @@ -280,23 +617,26 @@ def test_build_connector_meta_clears_pending( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) mm_hash = mock_request_with_3_mm.mm_features[0].identifier - scheduler._pending_specs[mm_hash] = ECMooncakeLoadSpec( + load_spec = ECMooncakeLoadSpec( mm_hash=mm_hash, num_token=0, nbytes=32, shape=(2, 4), dtype="float32", - producer_zmq="tcp://127.0.0.1:1", - lease_id="lease", + transfer_id="transfer", ) + scheduler._index_pending_spec(load_spec) + scheduler._load_specs[mm_hash] = load_spec scheduler._mm_datas_need_loads[mm_hash] = 100 - meta = scheduler.build_connector_meta(Mock(spec=SchedulerOutput)) + meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) assert isinstance(meta, ECMooncakeConnectorMetadata) assert len(meta.loads) == 1 assert meta.loads[0].mm_hash == mm_hash assert meta.loads[0].num_token == 100 assert scheduler._mm_datas_need_loads == {} - assert mm_hash not in scheduler._pending_specs + assert "transfer" not in scheduler._pending_specs def test_producer_does_not_build_load_metadata( self, mock_vllm_config_producer, mock_request_with_3_mm @@ -306,11 +646,46 @@ def test_producer_does_not_build_load_metadata( mock_vllm_config_producer, ECConnectorRole.SCHEDULER ) scheduler.update_state_after_alloc(mock_request_with_3_mm, 0) - meta = scheduler.build_connector_meta(Mock(spec=SchedulerOutput)) + meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) assert isinstance(meta, ECMooncakeConnectorMetadata) assert meta.loads == [] + def test_producer_builds_push_metadata_after_preprocessing( + self, mock_vllm_config_producer, mock_request_with_3_mm + ): + request = mock_request_with_3_mm + request.ec_transfer_params = { + "consumer_zmq": "tcp://decode:19019", + "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], + } + mock_vllm_config_producer.model_config.dtype = torch.float32 + mock_vllm_config_producer.model_config.get_hidden_size.return_value = 16 + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + scheduler.update_state_after_alloc(request, 0) + meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) + + assert meta.loads == [] + assert meta.pushes == [ + ECMooncakePushSpec( + mm_hash="img_hash_1", + nbytes=100 * 16 * 4, + shape=(100, 16), + dtype="float32", + consumer_zmq="tcp://decode:19019", + transfer_id="transfer-1", + request_id="test_req_123", + ) + ] + def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer): feature = SimpleNamespace( identifier="image_uuid", @@ -342,300 +717,776 @@ def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer class TestECMooncakeWorkerTransfer: - @pytest.mark.skipif( - not torch.accelerator.is_available(), - reason="Requires an accelerator for registered pool", - ) - def test_consumer_reuses_registered_cuda_pool(self, mock_vllm_config_consumer): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 - specs = [ - ECMooncakeLoadSpec( - mm_hash=f"hash_{index}", - num_token=1, - nbytes=256, - shape=(32, 2), + def test_batches_pushes_from_one_model_step(self, mock_vllm_config_producer): + port = _find_free_port() + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config + consumer_cfg.model_config = Mock() + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": port, + "consumer_buffer_pool_size": 4096, + } + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + sources = { + "first": torch.randn(4, 16), + "second": torch.randn(8, 16), + } + pushes = [ + ECMooncakePushSpec( + mm_hash=mm_hash, + nbytes=tensor.nbytes, + shape=tuple(tensor.shape), dtype="float32", - producer_zmq="tcp://127.0.0.1:1", - lease_id=f"lease_{index}", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id=f"transfer-{mm_hash}", ) - for index in range(2) + for mm_hash, tensor in sources.items() ] with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER ) - consumer.bind_connector_metadata(ECMooncakeConnectorMetadata(loads=specs)) - cache: dict[str, torch.Tensor] = {} - with patch.object(consumer, "_send_pull", return_value={"ok": True}): - consumer.start_load_caches(cache) + consumer.start_worker_services() + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=pushes)) + try: + producer.start_save_caches(encoder_cache=sources) + _wait_for_worker_io(producer) + + engine = producer._engine + assert isinstance(engine, CopyingFakeTransferEngine) + assert len(engine.transfer_calls) == 1 + assert sorted(engine.transfer_calls[0]) == sorted( + tensor.nbytes for tensor in sources.values() + ) + assert all( + reservation.ready + for reservation in consumer._push_reservations.values() + ) + finally: + producer.shutdown() + consumer.shutdown() - engine = consumer._engine - pool = consumer._consumer_pool - assert isinstance(engine, CopyingFakeTransferEngine) - assert pool is not None - assert engine.register_calls == [[pool.data_ptr()]] - assert engine.batch_unregister_calls == [] - assert cache["hash_0"].data_ptr() == pool.data_ptr() - assert cache["hash_1"].data_ptr() == pool.data_ptr() + 256 - - cache.clear() - consumer._release_stale_consumer_allocations(cache) - torch.accelerator.synchronize() - consumer._poll_consumer_pool_frees() - assert consumer._consumer_pool_allocator is not None - assert consumer._consumer_pool_allocator.allocate(4096) == (0, 4096) - consumer.shutdown() - - def test_single_process_save_and_load(self, mock_vllm_config_producer): - """Host-memory pull path (fake engine uses memcpy; CUDA ptrs need e2e).""" + def test_push_reserves_before_encoder_output_is_saved( + self, mock_vllm_config_producer + ): port = _find_free_port() + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config + consumer_cfg.model_config = Mock() + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": port, + "consumer_buffer_pool_size": 4096, + } mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "registry_http_port" - ] = port - mm_hash = "unit_test_hash" - torch.manual_seed(7) - source = torch.randn(4, 16, dtype=torch.float32) + source = torch.randn(4, 16) + push = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id="transfer-1", + ) with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) + consumer.start_worker_services() + scheduler = ECMooncakeConnector(consumer_cfg, ECConnectorRole.SCHEDULER) producer = ECMooncakeConnector( mock_vllm_config_producer, ECConnectorRole.WORKER ) - producer.save_caches({mm_hash: source}, mm_hash) - for _ in range(100): - if producer._zmq_listen_addr is not None: - break - time.sleep(0.01) - assert producer._zmq_listen_addr is not None - - url = f"http://127.0.0.1:{port}/ec/info/{mm_hash}" - r = httpx.get(url, timeout=2.0) - assert r.status_code == 200 - data = r.json() - - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 1e9 - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - spec = ECMooncakeLoadSpec( - mm_hash=mm_hash, - num_token=1, - nbytes=int(data["nbytes"]), - shape=tuple(int(x) for x in data["shape"]), - dtype=str(data["dtype"]), - producer_zmq=str(data["producer_zmq"]), - lease_id=str(data["lease_id"]), - ) - meta = ECMooncakeConnectorMetadata() - meta.add_load(spec) - consumer.bind_connector_metadata(meta) - loaded: dict[str, torch.Tensor] = {} - consumer.start_load_caches(loaded) - assert mm_hash in loaded - assert torch.allclose(loaded[mm_hash].cpu(), source.cpu()) - consumer_engine = consumer._engine - assert isinstance(consumer_engine, CopyingFakeTransferEngine) - assert consumer_engine.registered == set() - assert consumer_engine.batch_unregister_calls == [ - [loaded[mm_hash].data_ptr()] - ] - consumer.shutdown() - producer.shutdown() - - def test_batches_multi_item_transfer_and_reuses_socket( + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) + try: + producer.start_save_caches(encoder_cache={}) + _, reservation = producer._pending_reservations["hash"][0] + reservation_data = reservation.result(timeout=2) + assert reservation_data["nbytes"] == source.nbytes + old_reservation_id = reservation_data["reservation_id"] + reservation_data["_received_at"] -= _LEASE_TTL_SECONDS + consumer._push_reservations["transfer-1"].expires_at = 0 + with patch.object( + scheduler, + "_send_control", + wraps=scheduler._send_control, + ) as send_control: + assert not scheduler.has_cache_item("hash") + assert not scheduler.has_cache_item("hash") + assert send_control.call_count == 1 + assert send_control.call_args.args[1] == {"op": "event_port"} + assert "transfer-1" in consumer._push_reservations + + producer.save_caches({"hash": source}, "hash") + _wait_for_worker_io(producer) + assert ( + consumer._push_reservations["transfer-1"].reservation_id + != old_reservation_id + ) + deadline = time.monotonic() + 2 + while not scheduler.has_cache_item("hash"): + assert time.monotonic() < deadline + time.sleep(0.01) + assert send_control.call_count == 1 + load = scheduler._pending_specs["transfer-1"] + consumer.bind_connector_metadata( + ECMooncakeConnectorMetadata(loads=[load]) + ) + loaded: dict[str, torch.Tensor] = {} + consumer.start_load_caches(loaded) + first_meta = consumer.build_connector_worker_meta() + assert first_meta.loaded == {"hash"} + assert torch.equal(loaded["hash"], source) + consumer_engine = consumer._engine + assert isinstance(consumer_engine, CopyingFakeTransferEngine) + assert consumer_engine.transfer_calls == [] + finally: + producer.shutdown() + scheduler.shutdown() + consumer.shutdown() + + def test_finished_request_cancels_unbound_reservation( self, mock_vllm_config_producer ): + """A pre-reservation without an encoder tensor must not outlive its request.""" port = _find_free_port() + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config + consumer_cfg.model_config = Mock() + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": port, + "consumer_buffer_pool_size": 4096, + } mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "registry_http_port" - ] = port - sources = {f"hash_{i}": torch.randn(4, 16) for i in range(3)} + source = torch.randn(4, 16) + push = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id="transfer-1", + request_id="request-1", + ) with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) producer = ECMooncakeConnector( mock_vllm_config_producer, ECConnectorRole.WORKER ) - for mm_hash, tensor in sources.items(): - producer.save_caches({mm_hash: tensor}, mm_hash) - - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 1e9 - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp" - } + consumer.start_worker_services() + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) + try: + producer.start_save_caches(encoder_cache={}) + _, reservation = producer._pending_reservations["hash"][0] + reservation.result(timeout=2) + assert "transfer-1" in consumer._push_reservations + + producer.get_finished({"request-1"}) + _wait_for_worker_io(producer) + assert "hash" not in producer._pending_reservations + assert "transfer-1" not in consumer._push_reservations + finally: + producer.shutdown() + consumer.shutdown() + + def test_duplicate_pushes_share_one_transfer_per_reservation( + self, mock_vllm_config_producer + ): + port = _find_free_port() + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config + consumer_cfg.model_config = Mock() + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": port, + "consumer_buffer_pool_size": 4096, + } + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + source = torch.randn(4, 16) + push = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id="transfer-1", + ) + + with patch_ec_mooncake_deps(): consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) + scheduler = ECMooncakeConnector(consumer_cfg, ECConnectorRole.SCHEDULER) + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + consumer.start_worker_services() + producer.bind_connector_metadata( + ECMooncakeConnectorMetadata(pushes=[push, push]) + ) + try: + producer.start_save_caches(encoder_cache={"hash": source}) + _wait_for_worker_io(producer) - def make_spec(mm_hash: str) -> ECMooncakeLoadSpec: - data = httpx.get(f"http://127.0.0.1:{port}/ec/info/{mm_hash}").json() - return ECMooncakeLoadSpec( - mm_hash=mm_hash, - num_token=1, - nbytes=int(data["nbytes"]), - shape=tuple(data["shape"]), - dtype=str(data["dtype"]), - producer_zmq=str(data["producer_zmq"]), - lease_id=str(data["lease_id"]), + engine = producer._engine + assert isinstance(engine, CopyingFakeTransferEngine) + assert engine.transfer_calls == [[source.nbytes]] + reservation = consumer._push_reservations["transfer-1"] + assert reservation.ready + assert consumer._consumer_worker_metrics["completions_accepted"] == 1 + assert consumer._consumer_worker_metrics["completions_repeated"] == 0 + + deadline = time.monotonic() + 2 + while not scheduler.has_cache_item("hash"): + assert time.monotonic() < deadline + time.sleep(0.01) + load = scheduler._pop_pending_spec("transfer-1") + assert load is not None + load.num_token = 4 + consumer.bind_connector_metadata( + ECMooncakeConnectorMetadata(loads=[load]) + ) + loaded: dict[str, torch.Tensor] = {} + consumer.start_load_caches(loaded) + + cached_push = ECMooncakePushSpec( + mm_hash=push.mm_hash, + nbytes=push.nbytes, + shape=push.shape, + dtype=push.dtype, + consumer_zmq=push.consumer_zmq, + transfer_id="transfer-2", ) + producer.bind_connector_metadata( + ECMooncakeConnectorMetadata(pushes=[cached_push]) + ) + producer.start_save_caches(encoder_cache={"hash": source}) + _wait_for_worker_io(producer) + assert engine.transfer_calls == [[source.nbytes]] + cached = consumer._push_reservations["transfer-2"] + assert cached.ready and not cached.owns_allocation + assert consumer._consumer_worker_metrics["reservations_cached"] == 1 + + deadline = time.monotonic() + 2 + while not scheduler.has_cache_item("hash"): + assert time.monotonic() < deadline + time.sleep(0.01) + cached_load = scheduler._pop_pending_spec("transfer-2") + assert cached_load is not None + cached_load.num_token = 4 + consumer.bind_connector_metadata( + ECMooncakeConnectorMetadata(loads=[cached_load]) + ) + consumer.start_load_caches(loaded) + cached_meta = consumer.build_connector_worker_meta() + assert cached_meta.loaded == {"hash"} + assert "transfer-2" not in consumer._push_reservations + assert torch.equal(loaded["hash"], source) + finally: + producer.shutdown() + scheduler.shutdown() + consumer.shutdown() - first_meta = ECMooncakeConnectorMetadata( - loads=[make_spec("hash_0"), make_spec("hash_1")] - ) - consumer.bind_connector_metadata(first_meta) - loaded: dict[str, torch.Tensor] = {} - consumer.start_load_caches(loaded) - socket = next(iter(consumer._client_sockets.values())) - - second_meta = ECMooncakeConnectorMetadata(loads=[make_spec("hash_2")]) - consumer.bind_connector_metadata(second_meta) - consumer.start_load_caches(loaded) - - producer_engine = producer._engine - consumer_engine = consumer._engine - assert isinstance(producer_engine, CopyingFakeTransferEngine) - assert isinstance(consumer_engine, CopyingFakeTransferEngine) - assert producer_engine.transfer_calls == [[256, 256], [256]] - assert len(consumer._client_sockets) == 1 - assert next(iter(consumer._client_sockets.values())) is socket - assert all( - torch.equal(loaded[key], value) for key, value in sources.items() - ) - register_sizes = [len(call) for call in consumer_engine.register_calls] - assert register_sizes == [2, 1] - assert [len(call) for call in consumer_engine.batch_unregister_calls] == [ - 2, - 1, - ] - consumer.shutdown() - producer.shutdown() - - def test_producer_evicts_lru_registration_at_capacity( + def test_retired_item_reserved_again_still_serves_a_local_load( self, mock_vllm_config_producer ): + """A push for a retired item makes it live, not gone. + + Reusing the allocation for a new reservation takes it out of the + reclaim order. Looking the load up there instead of in the residency + map failed it, and the request fell back to waiting for a transfer. + """ + port = _find_free_port() + cfg = Mock(spec=VllmConfig) + cfg.parallel_config = mock_vllm_config_producer.parallel_config + cfg.model_config = Mock() + cfg.ec_transfer_config = Mock() + cfg.ec_transfer_config.is_ec_producer = False + cfg.ec_transfer_config.is_ec_consumer = True + cfg.ec_transfer_config.ec_buffer_device = "cpu" + cfg.ec_transfer_config.ec_buffer_size = 4096 + cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": port, + "consumer_buffer_pool_size": 4096, + } + spec = ECMooncakeLoadSpec( + mm_hash="hash", + num_token=0, + nbytes=64, + shape=(4, 4), + dtype="float32", + local=True, + ) + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) + try: + consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) + pool = consumer._consumer_pool + allocator = consumer._consumer_pool_allocator + assert pool is not None and allocator is not None + offset, size = allocator.allocate(spec.nbytes) + tensor = ( + pool.narrow(0, offset, spec.nbytes).view(torch.float32).view(4, 4) + ) + allocation = _ConsumerPoolAllocation(offset, size, tensor) + consumer._consumer_residents.insert("hash", allocation, size) + consumer._consumer_residents.retire("hash") + + # A later push reserves the retired copy instead of transferring. + consumer._reserve_push_destination( + { + "transfer_id": "t1", + "mm_hash": "hash", + "nbytes": spec.nbytes, + "shape": list(spec.shape), + "dtype": spec.dtype, + } + ) + assert consumer._consumer_residents.num_evictable == 0 + + assert consumer._take_resident_tensor(spec) is tensor + finally: + consumer.shutdown() + + def test_pushes_stage_through_the_registered_pool(self, mock_vllm_config_producer): + """Repeated content must not register overlapping source storage.""" port = _find_free_port() + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config + consumer_cfg.model_config = Mock() + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": port, + "consumer_buffer_pool_size": 4096, + } mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 32 - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "registry_http_port" - ] = port - first = torch.randn(8) - second = torch.randn(8) + source = torch.randn(4, 16) + pushes = [ + ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id=f"transfer-{index}", + ) + for index in range(2) + ] with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) producer = ECMooncakeConnector( mock_vllm_config_producer, ECConnectorRole.WORKER ) + consumer.start_worker_services() + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=pushes)) try: - producer.save_caches({"first": first}, "first") - producer.save_caches({"second": second}, "second") + producer.start_save_caches(encoder_cache={"hash": source}) + _wait_for_worker_io(producer) engine = producer._engine assert isinstance(engine, CopyingFakeTransferEngine) - assert list(producer._tensor_by_hash) == ["second"] - assert producer._registered_bytes == second.nbytes - assert first.data_ptr() in engine.unregister_calls - assert second.data_ptr() in engine.registered - - base_url = f"http://127.0.0.1:{port}/ec/info" - assert httpx.get(f"{base_url}/first").status_code == 404 - assert httpx.get(f"{base_url}/second").status_code == 200 + # The staging pool is registered once; a transfer registers + # nothing of its own. + pool = producer._producer_pool + assert pool is not None + assert engine.register_calls == [[pool.data_ptr()]] + assert engine.batch_unregister_calls == [] + assert engine.transfer_calls == [[source.nbytes, source.nbytes]] + assert all( + reservation.ready + for reservation in consumer._push_reservations.values() + ) finally: producer.shutdown() + consumer.shutdown() - def test_producer_does_not_evict_in_flight_registration( + def test_push_falls_back_to_per_tensor_registration_without_a_pool( self, mock_vllm_config_producer ): + """A pool that cannot be created must not break pushes.""" port = _find_free_port() - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 32 + consumer_cfg = self._push_harness_config(mock_vllm_config_producer, port) mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "registry_http_port" - ] = port - first = torch.randn(8) - second = torch.randn(8) + "producer_buffer_pool_size" + ] = 0 + source = torch.randn(4, 16) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id="transfer", + ) with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) producer = ECMooncakeConnector( mock_vllm_config_producer, ECConnectorRole.WORKER ) + consumer.start_worker_services() + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) try: - producer.save_caches({"first": first}, "first") - producer._tensor_by_hash["first"].in_flight = 1 - with pytest.raises(RuntimeError, match="no evictable"): - producer.save_caches({"second": second}, "second") - assert list(producer._tensor_by_hash) == ["first"] + producer.start_save_caches(encoder_cache={"hash": source}) + _wait_for_worker_io(producer) + engine = producer._engine + assert isinstance(engine, CopyingFakeTransferEngine) + assert producer._producer_pool is None + assert engine.register_calls == [[source.data_ptr()]] + assert engine.batch_unregister_calls == [[source.data_ptr()]] + assert engine.transfer_calls == [[source.nbytes]] finally: - producer._tensor_by_hash["first"].in_flight = 0 producer.shutdown() + consumer.shutdown() - def test_producer_does_not_evict_leased_registration( + def test_concurrent_pushes_hold_source_registration_until_last_release( self, mock_vllm_config_producer ): - port = _find_free_port() + """Concurrent transfers share one MR until every user releases it.""" mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 32 - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "registry_http_port" - ] = port - first = torch.randn(8) - second = torch.randn(8) + source = torch.randn(4, 16) with patch_ec_mooncake_deps(): producer = ECMooncakeConnector( mock_vllm_config_producer, ECConnectorRole.WORKER ) try: - producer.save_caches({"first": first}, "first") - response = httpx.get(f"http://127.0.0.1:{port}/ec/info/first").json() + first = producer._acquire_push_source_registrations([source]) + second = producer._acquire_push_source_registrations([source]) + engine = producer._engine + assert isinstance(engine, CopyingFakeTransferEngine) + assert len(engine.register_calls) == 1 - with pytest.raises(RuntimeError, match="no evictable"): - producer.save_caches({"second": second}, "second") + producer._release_push_source_registrations(first) + assert engine.batch_unregister_calls == [] + producer._release_push_source_registrations(second) + assert engine.batch_unregister_calls == [first] + finally: + producer.shutdown() + + def _push_harness_config(self, producer_cfg, port: int): + consumer_cfg = Mock(spec=VllmConfig) + consumer_cfg.parallel_config = producer_cfg.parallel_config + consumer_cfg.model_config = Mock() + consumer_cfg.ec_transfer_config = Mock() + consumer_cfg.ec_transfer_config.is_ec_producer = False + consumer_cfg.ec_transfer_config.is_ec_consumer = True + consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": port, + "consumer_buffer_pool_size": 4096, + } + producer_cfg.ec_transfer_config.ec_buffer_device = "cpu" + return consumer_cfg - assert producer._registry is not None - assert producer._registry.consume_lease("first", response["lease_id"]) - producer.save_caches({"second": second}, "second") - assert list(producer._tensor_by_hash) == ["second"] + def test_batch_completion_sends_one_control_message( + self, mock_vllm_config_producer + ): + """Completion is per batch, not per item: k items used to cost k RTTs.""" + port = _find_free_port() + consumer_cfg = self._push_harness_config(mock_vllm_config_producer, port) + sources = {"a": torch.randn(4, 16), "b": torch.randn(4, 16)} + pushes = [ + ECMooncakePushSpec( + mm_hash=mm_hash, + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id=f"transfer-{mm_hash}", + ) + for mm_hash, source in sources.items() + ] + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + consumer.start_worker_services() + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=pushes)) + try: + with patch.object( + producer, "_send_control", wraps=producer._send_control + ) as send_control: + producer.start_save_caches(encoder_cache=sources) + _wait_for_worker_io(producer) + ops = [call.args[1]["op"] for call in send_control.call_args_list] + assert ops.count("complete_batch") == 1 + assert "complete" not in ops + assert all( + reservation.ready + for reservation in consumer._push_reservations.values() + ) finally: producer.shutdown() + consumer.shutdown() - def test_shutdown_unregisters_all_producer_tensors(self, mock_vllm_config_producer): + def test_failed_push_is_reported_not_raised(self, mock_vllm_config_producer): + """A transfer failure must not surface as a fatal engine error.""" port = _find_free_port() - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "registry_http_port" - ] = port - tensor = torch.randn(8) + consumer_cfg = self._push_harness_config(mock_vllm_config_producer, port) + source = torch.randn(4, 16) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=f"tcp://127.0.0.1:{port}", + transfer_id="transfer", + ) with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) producer = ECMooncakeConnector( mock_vllm_config_producer, ECConnectorRole.WORKER ) - producer.save_caches({"hash": tensor}, "hash") - engine = producer._engine - assert isinstance(engine, CopyingFakeTransferEngine) + consumer.start_worker_services() + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) + try: + engine = producer._engine or producer._ensure_engine() + with patch.object(engine, "batch_transfer_sync_write", return_value=1): + producer.start_save_caches(encoder_cache={"hash": source}) + # No raise: the batch reports itself and gives up the + # consumer-side reservation. + _wait_for_worker_io(producer) + assert "transfer" not in consumer._push_reservations + finally: + producer.shutdown() + consumer.shutdown() + + def test_complete_is_idempotent_without_republishing( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) + reservation = consumer._reserve_push_destination( + { + "mm_hash": "hash", + "transfer_id": "transfer-1", + "nbytes": 64, + "shape": [4, 4], + "dtype": "float32", + } + ) + reservation_id = reservation["reservation_id"] - producer.shutdown() + first = consumer._complete_push("transfer-1", reservation_id) + repeated = consumer._complete_push("transfer-1", reservation_id) - assert engine.batch_unregister_calls == [[tensor.data_ptr()]] - assert engine.registered == set() - assert producer._tensor_by_hash == {} - assert producer._registered_bytes == 0 + assert first.accepted and first.became_ready + assert repeated.accepted and not repeated.became_ready + finally: + consumer.shutdown() + + def test_same_hash_transfers_have_independent_lifecycles( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + + def payload(transfer_id: str) -> dict: + return { + "mm_hash": "shared-hash", + "transfer_id": transfer_id, + "nbytes": 64, + "shape": [4, 4], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) + first = consumer._reserve_push_destination(payload("first")) + second = consumer._reserve_push_destination(payload("second")) + + consumer._complete_push("first", first["reservation_id"]) + assert consumer._push_reservations["first"].ready + assert not consumer._push_reservations["second"].ready + + assert consumer._cancel_push("first", first["reservation_id"]) + assert "first" not in consumer._push_reservations + assert "second" in consumer._push_reservations + assert consumer._complete_push("second", second["reservation_id"]) + finally: + consumer.shutdown() + + def test_late_completion_cannot_complete_new_reservation( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + payload = { + "mm_hash": "hash", + "transfer_id": "transfer", + "nbytes": 64, + "shape": [4, 4], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) + old = consumer._reserve_push_destination(payload) + consumer._push_reservations["transfer"].expires_at = 0 + consumer._expire_push_reservations() + new = consumer._reserve_push_destination(payload) + + assert old["reservation_id"] != new["reservation_id"] + stale = consumer._complete_push("transfer", old["reservation_id"]) + assert not stale.accepted + assert not consumer._push_reservations["transfer"].ready + finally: + consumer.shutdown() + + def test_ready_reservation_has_a_terminal_expiry(self, mock_vllm_config_consumer): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) + reservation = consumer._reserve_push_destination( + { + "mm_hash": "hash", + "transfer_id": "transfer-1", + "nbytes": 64, + "shape": [4, 4], + "dtype": "float32", + } + ) + consumer._complete_push("transfer-1", reservation["reservation_id"]) + consumer._push_reservations["transfer-1"].expires_at = 0 + + assert consumer._expire_push_reservations() == 1 + assert "transfer-1" not in consumer._push_reservations + finally: + consumer.shutdown() + + def test_cancel_before_reserve_creates_bounded_tombstone( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + payload = { + "mm_hash": "hash", + "transfer_id": "cancelled-transfer", + "nbytes": 64, + "shape": [4, 4], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) + assert consumer._cancel_push("cancelled-transfer", "") + cancelled = consumer._reserve_push_destination(payload) + assert cancelled["cancelled"] and not cancelled["write"] + assert "cancelled-transfer" not in consumer._push_reservations + + consumer._cancelled_transfers["cancelled-transfer"] = 0 + consumer._expire_push_reservations() + replacement = consumer._reserve_push_destination(payload) + assert replacement["write"] + finally: + consumer.shutdown() + + def test_missing_push_reservation_reports_failed_load( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + spec = ECMooncakeLoadSpec( + mm_hash="hash", + num_token=1, + nbytes=32, + shape=(8,), + dtype="float32", + pushed=True, + transfer_id="missing-transfer", + reservation_id="missing-reservation", + ) + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + consumer.bind_connector_metadata( + ECMooncakeConnectorMetadata(loads=[spec]) + ) + cache: dict[str, torch.Tensor] = {} + consumer.start_load_caches(cache) + meta = consumer.build_connector_worker_meta() + assert meta.failed_loads == {"hash"} + assert cache == {} + finally: + consumer.shutdown() def test_producer_scheduler_has_cache_item_false( self, mock_vllm_config_producer, mock_request_with_3_mm @@ -646,13 +1497,3 @@ def test_producer_scheduler_has_cache_item_false( ) mm_hash = mock_request_with_3_mm.mm_features[0].identifier assert not scheduler.has_cache_item(mm_hash) - - def test_consumer_worker_save_is_noop(self, mock_vllm_config_consumer): - with patch_ec_mooncake_deps(): - worker = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - mm_hash = "noop_hash" - tensor = torch.randn(2, 4) - worker.save_caches({mm_hash: tensor}, mm_hash) - assert mm_hash not in worker._tensor_by_hash diff --git a/tests/v1/ec_connector/unit/test_worker_ec_connector.py b/tests/v1/ec_connector/unit/test_worker_ec_connector.py index 3dcad1e50ae5..01dad9d89ebd 100644 --- a/tests/v1/ec_connector/unit/test_worker_ec_connector.py +++ b/tests/v1/ec_connector/unit/test_worker_ec_connector.py @@ -12,6 +12,9 @@ ECConnectorMetadata, ) from vllm.v1.outputs import EMPTY_MODEL_RUNNER_OUTPUT +from vllm.v1.worker.ec_connector_model_runner_mixin import ( + ECConnectorModelRunnerMixin, +) from vllm.v1.worker.gpu.ec_connector import NO_OP_EC_CONNECTOR, ActiveECConnector pytestmark = pytest.mark.cpu_test @@ -52,6 +55,9 @@ def test_saves_newly_added_caches_for_every_producer(is_producer, is_consumer): saved = [call.kwargs["mm_hash"] for call in fake.save_caches.call_args_list] assert saved == (["mm_new"] if is_producer else []) + assert fake.start_save_caches.called == is_producer + if is_producer: + assert fake.start_save_caches.call_args.kwargs["encoder_cache"] is encoder_cache assert fake.start_load_caches.called == is_consumer @@ -66,6 +72,26 @@ def test_worker_meta_is_reported_on_context_exit(): assert fake.clear_connector_metadata.called +def test_v1_producer_receives_existing_encoder_cache(): + encoder_cache = {"mm_cached": None} + fake = MagicMock(spec=ECConnectorBase) + fake.is_producer = True + fake.is_consumer = False + fake.get_finished.return_value = (None, None) + + module = "vllm.v1.worker.ec_connector_model_runner_mixin" + with ( + patch(f"{module}.has_ec_transfer", return_value=True), + patch(f"{module}.get_ec_transfer", return_value=fake), + ECConnectorModelRunnerMixin.maybe_get_ec_connector_output( + _scheduler_output(), encoder_cache + ), + ): + pass + + assert fake.start_save_caches.call_args.kwargs["encoder_cache"] is encoder_cache + + def test_no_forward_reports_without_running_the_model(): connector, _ = _connector() diff --git a/vllm/distributed/ec_transfer/ec_connector/base.py b/vllm/distributed/ec_transfer/ec_connector/base.py index 9138203bf4dd..5a526def1487 100644 --- a/vllm/distributed/ec_transfer/ec_connector/base.py +++ b/vllm/distributed/ec_transfer/ec_connector/base.py @@ -26,6 +26,7 @@ import enum from abc import ABC, abstractmethod +from collections.abc import Collection from typing import TYPE_CHECKING, Any import torch @@ -159,6 +160,25 @@ def register_caches( # TODO: Implement this later for P2P feature return + def start_save_caches(self, **kwargs: Any) -> None: + """Start work that can overlap encoder execution.""" + return None + + def start_worker_services(self) -> None: + """Start services that require the worker device to be initialized.""" + return None + + def take_unavailable_requests(self) -> set[str]: + """Requests whose encoder inputs can no longer be obtained. + + A connector that cannot always deliver an item reports the affected + requests here instead of deferring them forever. The scheduler fails + them with a retryable error, leaving the caller to decide whether to + re-issue the request. Called once per scheduling pass; the returned + ids are cleared. + """ + return set() + @abstractmethod def start_load_caches( self, encoder_cache: dict[str, torch.Tensor], **kwargs @@ -245,7 +265,10 @@ def has_cache_item( pass def ensure_cache_available( - self, request: "Request", num_computed_tokens: int + self, + request: "Request", + num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, ) -> bool: """ Ensure encoder cache items are available for the given request. @@ -254,6 +277,7 @@ def ensure_cache_available( Args: request: the request whose multimodal features to check. num_computed_tokens: tokens already covered by cached KV blocks. + local_cache_hashes: encoder outputs already cached locally. Returns: True if all items are ready or no transfer is needed. @@ -271,6 +295,18 @@ def update_state_after_alloc(self, request: "Request", index: int): """ pass + def update_state_after_free(self, request: "Request", index: int): + """ + Called once the request has consumed the encoder input, well before it + finishes generating. Connectors that hold per-request transfer state + (buffers, reservations) should release this item's share here. + + Args: + request (Request): the request object. + index (int): the multimodal item index within the request. + """ + return + @abstractmethod def build_connector_meta( self, scheduler_output: SchedulerOutput diff --git a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py index 7de5c7fae3ff..c582c4faf08b 100644 --- a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py @@ -8,6 +8,7 @@ offloaded to CPU instead of recomputing them. """ +from collections.abc import Collection from typing import TYPE_CHECKING import torch @@ -89,11 +90,14 @@ def has_cache_item(self, identifier: str) -> bool: return self.connector_scheduler.has_cache_item(identifier) def ensure_cache_available( - self, request: "Request", num_computed_tokens: int + self, + request: "Request", + num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, ) -> bool: assert self.connector_scheduler is not None return self.connector_scheduler.ensure_cache_available( - request, num_computed_tokens + request, num_computed_tokens, local_cache_hashes ) def update_state_after_alloc(self, request: "Request", index: int) -> None: diff --git a/vllm/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py b/vllm/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py index 132efc1df0f7..fadeb1330726 100644 --- a/vllm/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py +++ b/vllm/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py @@ -7,6 +7,7 @@ for the ECCPUConnector. """ +from collections.abc import Collection from typing import TYPE_CHECKING from vllm.distributed.ec_transfer.ec_connector.cpu.common import ( @@ -60,7 +61,10 @@ def has_cache_item(self, identifier: str) -> bool: return entry is not None and entry.ready def ensure_cache_available( - self, request: "Request", num_computed_tokens: int + self, + request: "Request", + num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, ) -> bool: return True diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index b9b4350124df..187c9ca9f623 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -10,36 +10,41 @@ from __future__ import annotations -import json +import bisect import math import threading import time import uuid -from collections import OrderedDict +from collections import Counter, OrderedDict, deque +from collections.abc import Callable, Collection +from concurrent.futures import Future, ThreadPoolExecutor from contextlib import suppress from dataclasses import dataclass, field -from typing import Any +from typing import Any, Generic, TypeVar -import httpx import torch -import uvicorn import zmq -from fastapi import FastAPI, HTTPException from vllm.config import VllmConfig from vllm.distributed.ec_transfer.ec_connector.base import ( ECConnectorBase, ECConnectorMetadata, ECConnectorRole, + ECConnectorWorkerMetadata, ) -from vllm.distributed.parallel_state import is_local_first_rank from vllm.logger import init_logger from vllm.utils.network_utils import get_ip from vllm.v1.core.sched.output import SchedulerOutput +from vllm.v1.outputs import ECConnectorOutput logger = init_logger(__name__) +_T = TypeVar("_T") + _LEASE_TTL_SECONDS = 300 +_RESERVATION_REFRESH_SECONDS = _LEASE_TTL_SECONDS / 2 +_RESERVATION_REAP_INTERVAL_SECONDS = 1 +_DRAIN_MIN_INTERVAL = 0.005 _MOONCAKE_IMPORT_ERROR: ImportError | None try: @@ -60,8 +65,25 @@ class ECMooncakeLoadSpec: nbytes: int shape: tuple[int, ...] dtype: str - producer_zmq: str - lease_id: str + pushed: bool = False + transfer_id: str = "" + reservation_id: str = "" + # The consumer pool still holds this item, so the load is a local handoff: + # no transfer, no producer. + local: bool = False + + +@dataclass +class ECMooncakePushSpec: + """Destination reservation requested before an encoder tensor is ready.""" + + mm_hash: str + nbytes: int + shape: tuple[int, ...] + dtype: str + consumer_zmq: str + transfer_id: str + request_id: str = "" @dataclass @@ -69,15 +91,43 @@ class ECMooncakeConnectorMetadata(ECConnectorMetadata): """Worker-side metadata for one scheduler step.""" loads: list[ECMooncakeLoadSpec] = field(default_factory=list) + pushes: list[ECMooncakePushSpec] = field(default_factory=list) def add_load(self, spec: ECMooncakeLoadSpec) -> None: self.loads.append(spec) + def add_push(self, spec: ECMooncakePushSpec) -> None: + self.pushes.append(spec) + + +@dataclass +class ECMooncakeWorkerMetadata(ECConnectorWorkerMetadata): + """Completion state reported from workers to the scheduler.""" + + loaded: set[str] = field(default_factory=set) + failed_loads: set[str] = field(default_factory=set) + # Items the receive pool dropped under pressure. The scheduler assumes an + # evicted item stays resident until told otherwise. + reclaimed: set[str] = field(default_factory=set) + pending_loads: bool = False + pending_saves: bool = False + + def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadata: + assert isinstance(other, ECMooncakeWorkerMetadata) + return ECMooncakeWorkerMetadata( + loaded=self.loaded | other.loaded, + failed_loads=self.failed_loads | other.failed_loads, + reclaimed=self.reclaimed | other.reclaimed, + pending_loads=self.pending_loads or other.pending_loads, + pending_saves=self.pending_saves or other.pending_saves, + ) + @dataclass -class _RegisteredTensor: +class _PushSourceRegistration: tensor: torch.Tensor - in_flight: int = 0 + nbytes: int + users: int = 1 @dataclass @@ -87,6 +137,226 @@ class _ConsumerPoolAllocation: tensor: torch.Tensor +@dataclass +class _PushReservation: + mm_hash: str + reservation_id: str + allocation: _ConsumerPoolAllocation + shape: tuple[int, ...] + dtype: str + ready: bool = False + owns_allocation: bool = True + discard_on_complete: bool = False + created_at: float = field(default_factory=time.monotonic) + expires_at: float = 0 + + +@dataclass(frozen=True) +class _PushCompletion: + accepted: bool + became_ready: bool = False + + +@dataclass +class _PendingPush: + tensor: torch.Tensor + spec: ECMooncakePushSpec + reservation: Future[dict[str, Any]] + ready_event: torch.Event | None + enqueued_at: float + + +@dataclass +class _PushPerfWindow: + started_at: float = field(default_factory=time.monotonic) + batches: int = 0 + items: int = 0 + bytes: int = 0 + skipped_items: int = 0 + failures: int = 0 + stage_totals_ms: dict[str, float] = field(default_factory=dict) + stage_max_ms: dict[str, float] = field(default_factory=dict) + + +class _ControlChannel: + """Reusable REQ sockets for the ZMQ control plane. + + One context and one connection per message costs a thread spawn plus a + TCP handshake, and the push path sends one message per reserve, complete + and cancel. Sockets are cached per thread because a REQ socket is neither + thread-safe nor usable after a failed exchange. + """ + + def __init__(self, timeout_ms: int): + self._context = zmq.Context() + self._timeout_ms = timeout_ms + self._local = threading.local() + + def _sockets(self) -> dict[str, zmq.Socket]: + sockets = getattr(self._local, "sockets", None) + if sockets is None: + sockets = {} + self._local.sockets = sockets + return sockets + + def _discard(self, addr: str) -> None: + socket = self._sockets().pop(addr, None) + if socket is not None: + socket.close(linger=0) + + def send(self, addr: str, payload: dict[str, Any]) -> dict[str, Any]: + sockets = self._sockets() + socket = sockets.get(addr) + if socket is None: + socket = self._context.socket(zmq.REQ) + socket.setsockopt(zmq.RCVTIMEO, self._timeout_ms) + socket.setsockopt(zmq.SNDTIMEO, self._timeout_ms) + socket.setsockopt(zmq.LINGER, 0) + socket.connect(addr) + sockets[addr] = socket + try: + socket.send_json(payload) + response = socket.recv_json() + except Exception: + # A REQ socket cannot recover from a half-finished exchange. + self._discard(addr) + raise + assert isinstance(response, dict) + return response + + def request(self, addr: str, payload: dict[str, Any]) -> Any: + response = self.send(addr, payload) + if not response.get("ok"): + raise RuntimeError(response.get("error", "EC control request failed")) + return response.get("result") + + def close(self) -> None: + # Callers must have stopped every thread that used this channel. + self._context.destroy(linger=0) + + +class _ResidentPool(Generic[_T]): + """Content-addressed entries kept until their space is needed. + + Both sides of the connector hold the same thing under different names: a + map from mm_hash to a device resource, a count of who is using it, and an + eviction order over the rest. This is `BlockPool`'s accounting for + variable-sized entries: `acquire`/`release` mirror `touch`/`free_blocks`, + and `evict_lru` mirrors the reclaim inside `get_new_blocks`. + + An unreferenced entry stays resident. Eviction is driven by pressure, so + the entry serves whoever needs it next instead of being transferred again. + """ + + def __init__(self, capacity: int): + self.capacity = capacity + self.used = 0 + self._entries: dict[str, tuple[_T, int]] = {} + self._refs: Counter[str] = Counter() + # Unreferenced entries in eviction order, oldest first. + self._evictable: OrderedDict[str, None] = OrderedDict() + + def __len__(self) -> int: + return len(self._entries) + + def __contains__(self, key: str) -> bool: + return key in self._entries + + @property + def num_evictable(self) -> int: + return len(self._evictable) + + def referenced(self) -> list[str]: + """Keys that are in use. `_refs` only holds entries above zero.""" + return list(self._refs) + + def referenced_or_retired(self) -> list[str]: + """Every key held, in insertion order.""" + return list(self._entries) + + def get(self, key: str) -> _T | None: + entry = self._entries.get(key) + return entry[0] if entry is not None else None + + def insert(self, key: str, value: _T, nbytes: int) -> None: + """Add a referenced entry, replacing any previous one.""" + previous = self._entries.get(key) + if previous is not None: + self.used -= previous[1] + self._entries[key] = (value, nbytes) + self.used += nbytes + self.pin(key) + + def pin(self, key: str) -> _T | None: + """Mark an entry as in use without counting a new reference. + + For a holder whose references are discovered by scanning rather than + released in pairs, `pin`/`retire` are the matching operations. + """ + entry = self._entries.get(key) + if entry is None: + return None + self._evictable.pop(key, None) + self._refs[key] = max(1, self._refs[key]) + return entry[0] + + def retire(self, key: str) -> None: + """Drop every reference; the entry is evictable from now on.""" + if key not in self._entries: + return + self._refs.pop(key, None) + self._evictable[key] = None + + def refresh(self, key: str) -> None: + """Move an unreferenced entry to the back of the eviction order.""" + if key in self._evictable: + self._evictable.move_to_end(key) + + def acquire(self, key: str) -> _T | None: + """Take one reference so pressure cannot evict the entry.""" + entry = self._entries.get(key) + if entry is None: + return None + self._evictable.pop(key, None) + self._refs[key] += 1 + return entry[0] + + def release(self, key: str) -> None: + """Drop one reference; the entry becomes evictable at zero.""" + if key not in self._entries: + return + count = self._refs[key] - 1 + if count > 0: + self._refs[key] = count + return + self._refs.pop(key, None) + self._evictable[key] = None + + def evict_lru(self, evict: Callable[[str, _T], bool]) -> str | None: + """Drop the oldest entry `evict` accepts, and return its key. + + `evict` returns False for an entry that cannot go yet (a lease the + remote side still holds, a deregistration that failed). Those keep + their place in the order and the next candidate is tried. + """ + for key in list(self._evictable): + value, nbytes = self._entries[key] + if not evict(key, value): + continue + self._evictable.pop(key, None) + del self._entries[key] + self._refs.pop(key, None) + self.used -= nbytes + return key + return None + + def clear(self) -> None: + self._entries.clear() + self._refs.clear() + self._evictable.clear() + self.used = 0 + + class _ContiguousAllocator: def __init__(self, capacity: int, alignment: int = 256): self.capacity = capacity @@ -106,118 +376,235 @@ def allocate(self, nbytes: int) -> tuple[int, int] | None: return None def free(self, offset: int, size: int) -> None: - self._free.append((offset, size)) - self._free.sort() - merged: list[tuple[int, int]] = [] - for free_offset, free_size in self._free: - if merged and sum(merged[-1]) == free_offset: - previous_offset, previous_size = merged[-1] - merged[-1] = (previous_offset, previous_size + free_size) - else: - merged.append((free_offset, free_size)) - self._free = merged + index = bisect.bisect_left(self._free, (offset, size)) + self._free.insert(index, (offset, size)) + # Coalesce with the neighbours only; the rest of the list is already + # merged, so a full re-scan per free is wasted work. + if index + 1 < len(self._free): + next_offset, next_size = self._free[index + 1] + if offset + size == next_offset: + self._free[index] = (offset, size + next_size) + self._free.pop(index + 1) + if index > 0: + previous_offset, previous_size = self._free[index - 1] + current_offset, current_size = self._free[index] + if previous_offset + previous_size == current_offset: + self._free[index - 1] = ( + previous_offset, + previous_size + current_size, + ) + self._free.pop(index) -class ECMooncakeRegistryServer: - """Lightweight HTTP registry on the producer for remote has_cache_item / info.""" +class ECMooncakeControlServer: + """Expose consumer reservations over a lightweight ZMQ control channel.""" - def __init__(self, host: str, port: int): + def __init__( + self, + host: str, + port: int, + reserve: Callable[[dict[str, Any]], dict[str, Any]], + status: Callable[[str], dict[str, Any] | None], + complete: Callable[[str, str], _PushCompletion], + cancel: Callable[[str, str, bool], bool], + reap: Callable[[], int], + metrics_log_interval: float = 10, + ): self.host = host self.port = port - self._entries: dict[str, dict[str, Any]] = {} - self._leases: dict[str, dict[str, float]] = {} - self._lock = threading.Lock() - self.app = FastAPI() - self._register_routes() - self.server_thread: threading.Thread | None = None - self.server: uvicorn.Server | None = None - - def _register_routes(self) -> None: - @self.app.get("/ec/info/{mm_hash}") - async def ec_info(mm_hash: str) -> dict[str, Any]: - with self._lock: - data = self._entries.get(mm_hash) - if data is None: - raise HTTPException(status_code=404, detail="unknown mm_hash") - lease_id = uuid.uuid4().hex - self._leases.setdefault(mm_hash, {})[lease_id] = ( - time.monotonic() + _LEASE_TTL_SECONDS - ) - return {**data, "lease_id": lease_id} + self.event_port: int | None = None + self._reserve = reserve + self._status = status + self._complete = complete + self._cancel = cancel + self._reap = reap + self._metrics_log_interval = metrics_log_interval + self._stop = threading.Event() + self._started = threading.Event() + self._thread: threading.Thread | None = None + self._startup_error: Exception | None = None def start(self) -> None: - if self.server_thread is not None: - return - config = uvicorn.Config(app=self.app, host=self.host, port=self.port) - self.server = uvicorn.Server(config=config) - self.server_thread = threading.Thread( - target=self.server.run, name="ec_mooncake_registry", daemon=True + def loop() -> None: + context = zmq.Context() + socket = context.socket(zmq.REP) + event_socket = context.socket(zmq.PUSH) + pending_events: deque[dict[str, Any]] = deque() + metrics: Counter[str] = Counter() + metrics_started_at = time.monotonic() + last_reap_at = metrics_started_at + socket.setsockopt(zmq.RCVTIMEO, 100) + try: + socket.bind(f"tcp://{self.host}:{self.port}") + self.event_port = event_socket.bind_to_random_port(f"tcp://{self.host}") + except Exception as e: + self._startup_error = e + self._started.set() + socket.close(linger=0) + event_socket.close(linger=0) + context.term() + return + self._started.set() + try: + while not self._stop.is_set(): + while pending_events: + try: + event_socket.send_json( + pending_events[0], flags=zmq.DONTWAIT + ) + except zmq.Again: + break + pending_events.popleft() + metrics["events_sent"] += 1 + now = time.monotonic() + if now - last_reap_at >= _RESERVATION_REAP_INTERVAL_SECONDS: + metrics["reservations_reaped"] += self._reap() + last_reap_at = now + if ( + self._metrics_log_interval > 0 + and now - metrics_started_at >= self._metrics_log_interval + ): + logger.info( + "EC Mooncake consumer control: requests=%s, " + "events_queued=%d, events_sent=%d, event_backlog=%d, " + "reservations_reaped=%d", + { + key.removeprefix("request_"): value + for key, value in metrics.items() + if key.startswith("request_") + }, + metrics["events_queued"], + metrics["events_sent"], + len(pending_events), + metrics["reservations_reaped"], + ) + metrics.clear() + metrics_started_at = now + try: + request = socket.recv_json() + except zmq.Again: + continue + try: + op = request.get("op") + result: Any = None + metrics[f"request_{op}"] += 1 + if op == "reserve": + result = self._reserve(request) + if result.get("ready"): + transfer_id = str(request["transfer_id"]) + status = self._status(transfer_id) + if status is not None: + pending_events.append( + {"transfer_id": transfer_id, **status} + ) + metrics["events_queued"] += 1 + elif op == "status": + result = self._status(str(request["transfer_id"])) + elif op == "event_port": + result = self.event_port + elif op in ("complete", "complete_batch"): + items = ( + request["items"] + if op == "complete_batch" + else [request] + ) + completions = [] + for item in items: + transfer_id = str(item["transfer_id"]) + completion = self._complete( + transfer_id, + str(item["reservation_id"]), + ) + completions.append( + { + "completed": completion.accepted, + "became_ready": completion.became_ready, + } + ) + if not completion.became_ready: + continue + status = self._status(transfer_id) + if status is not None: + pending_events.append( + {"transfer_id": transfer_id, **status} + ) + metrics["events_queued"] += 1 + result = ( + {"items": completions} + if op == "complete_batch" + else completions[0] + ) + elif op == "cancel": + result = { + "cancelled": self._cancel( + str(request["transfer_id"]), + str(request.get("reservation_id", "")), + bool(request.get("abandon", False)), + ) + } + else: + raise ValueError(f"unknown control op: {op!r}") + socket.send_json({"ok": True, "result": result}) + except Exception as e: + socket.send_json({"ok": False, "error": str(e)}) + finally: + socket.close(linger=0) + event_socket.close(linger=0) + context.term() + + self._thread = threading.Thread( + target=loop, name="ec-mooncake-control", daemon=True ) - self.server_thread.start() - while self.server is not None and not self.server.started: - time.sleep(0.05) + self._thread.start() + if not self._started.wait(timeout=5): + raise RuntimeError("EC Mooncake control channel failed to start") + if self._startup_error is not None: + raise RuntimeError("EC Mooncake control channel failed to bind") from ( + self._startup_error + ) logger.info( - "EC Mooncake registry listening on http://%s:%d", self.host, self.port + "EC Mooncake control channel listening on tcp://%s:%d (events tcp://%s:%d)", + self.host, + self.port, + self.host, + self.event_port, ) def shutdown(self) -> None: - if self.server is None or self.server_thread is None or not self.server.started: - return - self.server.should_exit = True - self.server_thread.join() - logger.info("EC Mooncake registry stopped.") - - def publish(self, mm_hash: str, payload: dict[str, Any]) -> None: - with self._lock: - self._entries[mm_hash] = payload - - def unpublish(self, mm_hash: str) -> dict[str, Any] | None: - with self._lock: - self._leases.pop(mm_hash, None) - return self._entries.pop(mm_hash, None) - - def unpublish_if_unleased(self, mm_hash: str) -> tuple[bool, dict[str, Any] | None]: - with self._lock: - leases = self._leases.get(mm_hash, {}) - now = time.monotonic() - leases = {lease: expiry for lease, expiry in leases.items() if expiry > now} - if leases: - self._leases[mm_hash] = leases - return False, None - self._leases.pop(mm_hash, None) - return True, self._entries.pop(mm_hash, None) - - def consume_leases(self, leases: list[tuple[str, str]]) -> bool: - with self._lock: - now = time.monotonic() - for mm_hash, lease_id in leases: - if self._leases.get(mm_hash, {}).get(lease_id, 0) <= now: - return False - for mm_hash, lease_id in leases: - item_leases = self._leases[mm_hash] - item_leases.pop(lease_id) - if not item_leases: - self._leases.pop(mm_hash) - return True - - def consume_lease(self, mm_hash: str, lease_id: str) -> bool: - return self.consume_leases([(mm_hash, lease_id)]) + if self._thread is None: + return + self._stop.set() + self._thread.join() class ECMooncakeConnector(ECConnectorBase): """ EC connector using Mooncake TransferEngine for GPU tensor transport. + The producer pushes each encoder output into a receive buffer the consumer + reserved for it, so the transfer overlaps encoding instead of waiting for + the consumer to ask. An item the consumer's encoder cache evicted stays in + that pool and is handed back locally; when neither has it, the load fails + with a retryable error so the caller can re-issue the request. + Extra config (``ec_connector_extra_config``): - - ``remote_registry_url`` (consumer, required): Base URL of the producer - registry, e.g. ``http://192.168.0.2:9018``. - - ``registry_http_port`` (producer, optional): Port for the in-process HTTP - registry (default ``9018``). - ``mooncake_protocol`` (optional): Passed to ``TransferEngine.initialize`` (default ``"rdma"``). - ``consumer_buffer_pool_size`` (consumer, optional): Bytes reserved for a long-lived registered CUDA receive arena (default ``ec_buffer_size``). + - ``reservation_zmq_port`` (consumer worker, required): Exposes registered + receive addresses over ZMQ on this port so producers can push into them. + - ``reservation_zmq_addr`` (consumer scheduler, required): Address of the + consumer control channel. Defaults to ``tcp://127.0.0.1:``. + - ``transfer_max_workers`` (optional): Maximum concurrent Mooncake transfer + batches (default ``4``). + - ``control_max_workers`` (optional): Maximum concurrent reservation requests + issued by a producer (default ``8``). + - ``transfer_metrics_log_interval`` (optional): Seconds between aggregated + push-transfer performance logs (default ``10``; ``0`` disables them). + - ``consumer_metrics_log_interval`` (optional): Seconds between aggregated + consumer lifecycle logs (default ``10``; ``0`` disables them). Limitations: ``tensor_parallel_size`` and ``pipeline_parallel_size`` must be ``1`` (same assumption as Mooncake KV connector for P2P handshake). @@ -244,8 +631,16 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._ec_cfg = ec_cfg self._extra = self._ec_cfg.ec_connector_extra_config self._protocol: str = self._extra.get("mooncake_protocol", "rdma") - self._remote_registry_url: str | None = self._extra.get("remote_registry_url") - self._registry_http_port: int = int(self._extra.get("registry_http_port", 9018)) + reservation_port = self._extra.get("reservation_zmq_port") + self._reservation_zmq_port = ( + int(reservation_port) if reservation_port is not None else None + ) + self._reservation_zmq_addr: str | None = self._extra.get("reservation_zmq_addr") + if ( + self._reservation_zmq_addr is None + and self._reservation_zmq_port is not None + ): + self._reservation_zmq_addr = f"tcp://127.0.0.1:{self._reservation_zmq_port}" self._registered_capacity = int(self._ec_cfg.ec_buffer_size) if self._registered_capacity <= 0: raise ValueError("ECMooncakeConnector requires ec_buffer_size > 0.") @@ -258,45 +653,121 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._consumer_pool_capacity = int(pool_size) self._consumer_pool: torch.Tensor | None = None self._consumer_pool_allocator: _ContiguousAllocator | None = None - self._consumer_allocations: dict[str, _ConsumerPoolAllocation] = {} + # The receive pool is orders of magnitude larger than the encoder + # cache, so an item the encoder cache evicted stays resident here and + # a later request gets it for a dict lookup instead of a transfer. + self._consumer_residents: _ResidentPool[_ConsumerPoolAllocation] = ( + _ResidentPool(self._consumer_pool_capacity) + ) + self._consumer_retire_events: dict[str, torch.Event] = {} self._consumer_pending_frees: list[ tuple[torch.Event, _ConsumerPoolAllocation] ] = [] + self._consumer_reclaimed: set[str] = set() self._consumer_pool_disabled = self._consumer_pool_capacity <= 0 - - # Scheduler (consumer): mm_hash -> pending tensor layout from registry + self._consumer_lock = threading.Lock() + self._push_reservations: dict[str, _PushReservation] = {} + self._cancelled_transfers: dict[str, float] = {} + self._control_server: ECMooncakeControlServer | None = None + self._consumer_metrics_log_interval = float( + self._extra.get("consumer_metrics_log_interval", 10) + ) + self._consumer_metrics_started_at = time.monotonic() + self._consumer_worker_metrics: Counter[str] = Counter() + self._consumer_scheduler_metrics: Counter[str] = Counter() + self._consumer_missing_since: dict[str, float] = {} + self._stalled_hashes: set[str] = set() + self._unavailable_requests: set[str] = set() + self._active_push_sources: Counter[tuple[str, int]] = Counter() + self._active_push_sources_lock = threading.Lock() + self._push_wait_timeout = float(self._extra.get("push_wait_timeout_s", 60)) + self._drain_pending = True + self._drained_at = 0.0 + self._consumer_loading_since: dict[str, float] = {} + self._consumer_pending_since: dict[str, float] = {} + self._pending_spec_deadlines: dict[str, float] = {} + self._pending_cancels: dict[str, Future[Any]] = {} + self._cancelled_transfer_ids: set[str] = set() + + # Scheduler (consumer): transfer_id -> pending tensor layout. self._pending_specs: dict[str, ECMooncakeLoadSpec] = {} + self._pending_specs_by_hash: dict[str, deque[str]] = {} + self._load_specs: dict[str, ECMooncakeLoadSpec] = {} self._mm_datas_need_loads: dict[str, int] = {} + self._loading_hashes: set[str] = set() + self._ready_hashes: set[str] = set() + # Scheduler-side mirror of the worker's receive pool, oldest first. An + # item stays here after the encoder cache evicts it, so the next + # request that needs it is served locally instead of consuming another + # transfer. The worker reports what it reclaims under pressure; the + # byte budget only guards against drift. + self._resident_specs: OrderedDict[str, ECMooncakeLoadSpec] = OrderedDict() + self._resident_bytes = 0 + self._scheduler_pending_work = False + self._pushes_to_prepare: dict[str, ECMooncakePushSpec] = {} # Worker producer self._engine: TransferEngine | None = None + self._engine_lock = threading.Lock() self._hostname = get_ip() - self._registry: ECMooncakeRegistryServer | None = None - self._zmq_listen_addr: str | None = None - self._zmq_thread: threading.Thread | None = None - self._zmq_ctx: zmq.Context | None = None - self._zmq_stop = threading.Event() - self._client_zmq_ctx: zmq.Context | None = None - self._client_sockets: dict[str, zmq.Socket] = {} - self._tensor_by_hash: OrderedDict[str, _RegisteredTensor] = OrderedDict() - self._registered_bytes = 0 + # Published encoder outputs, referenced while a pull is reading them. self._pending_unregister: dict[int, torch.Tensor] = {} - self._tensor_lock = threading.Lock() - self._producer_services_started = False + self._push_source_registrations: dict[int, _PushSourceRegistration] = {} + self._push_source_registration_lock = threading.Lock() + producer_pool = self._extra.get( + "producer_buffer_pool_size", self._registered_capacity + ) + self._producer_pool_capacity = int(producer_pool) + self._producer_pool: torch.Tensor | None = None + self._producer_pool_allocator: _ContiguousAllocator | None = None + self._producer_pool_disabled = self._producer_pool_capacity <= 0 + self._producer_pool_lock = threading.Lock() + transfer_workers = int(self._extra.get("transfer_max_workers", 4)) + control_workers = int(self._extra.get("control_max_workers", 8)) + self._transfer_metrics_log_interval = float( + self._extra.get("transfer_metrics_log_interval", 10) + ) + self._control_channel = _ControlChannel( + int(float(self._extra.get("control_timeout_s", 30)) * 1000) + ) + self._producer_metrics: Counter[str] = Counter() + self._io_executor = ThreadPoolExecutor( + max_workers=transfer_workers, thread_name_prefix="ec-mooncake-transfer" + ) + self._control_executor = ThreadPoolExecutor( + max_workers=control_workers, thread_name_prefix="ec-mooncake-control" + ) + self._pending_saves: list[tuple[str, Future[None]]] = [] + self._pending_reservations: dict[ + str, deque[tuple[ECMooncakePushSpec, Future[dict[str, Any]]]] + ] = {} + self._pending_pushes: list[_PendingPush] = [] + self._push_perf_lock = threading.Lock() + self._push_perf = _PushPerfWindow() + self._active_transfer_batches = 0 + self._queued_transfer_batches = 0 + self._event_zmq_ctx: zmq.Context | None = None + self._event_zmq_socket: zmq.Socket | None = None + self._completed_loads: set[str] = set() + self._failed_loads: set[str] = set() self._shutdown = False if ( role == ECConnectorRole.SCHEDULER and self.is_consumer - and not self._remote_registry_url + and not self._reservation_zmq_addr ): raise ValueError( "ec_consumer with ECMooncakeConnector requires " - "ec_connector_extra_config['remote_registry_url']." + "reservation_zmq_port or reservation_zmq_addr." ) def _ensure_engine(self) -> TransferEngine: - if self._engine is None: + if self._engine is not None: + return self._engine + with self._engine_lock: + if self._engine is not None: + return self._engine eng = TransferEngine() ret = eng.initialize(self._hostname, "P2PHANDSHAKE", self._protocol, "") if ret != 0: @@ -309,6 +780,35 @@ def _ensure_engine(self) -> TransferEngine: ) return self._engine + def start_worker_services(self) -> None: + if ( + self._role != ECConnectorRole.WORKER + or not self.is_consumer + or self._reservation_zmq_port is None + or self._control_server is not None + ): + return + raw_device = self._ec_cfg.ec_buffer_device + device_name = ( + raw_device.lower() if isinstance(raw_device, str) and raw_device else "cuda" + ) + self._ensure_consumer_pool(torch.device(device_name), allow_host=True) + if self._consumer_pool is None: + raise RuntimeError( + "Mooncake push mode requires a registered consumer buffer pool." + ) + self._control_server = ECMooncakeControlServer( + "0.0.0.0", + self._reservation_zmq_port, + self._reserve_push_destination, + self._push_status, + self._complete_push, + self._cancel_push, + self._expire_push_reservations, + self._consumer_metrics_log_interval, + ) + self._control_server.start() + def _unregister_memory(self, tensor: torch.Tensor) -> bool: assert self._engine is not None ret = self._engine.unregister_memory(tensor.data_ptr()) @@ -323,161 +823,6 @@ def _unregister_memory(self, tensor: torch.Tensor) -> bool: self._pending_unregister.pop(tensor.data_ptr(), None) return True - def _release_tensor_locked(self, mm_hash: str) -> bool: - entry = self._tensor_by_hash.get(mm_hash) - if entry is None: - return True - if entry.in_flight: - return False - payload = None - if self._registry is not None: - released, payload = self._registry.unpublish_if_unleased(mm_hash) - if not released: - return False - if not self._unregister_memory(entry.tensor): - if self._registry is not None and payload is not None: - self._registry.publish(mm_hash, payload) - return False - self._tensor_by_hash.pop(mm_hash) - self._registered_bytes -= entry.tensor.nbytes - return True - - def _make_registration_space_locked(self, nbytes: int) -> None: - if nbytes > self._registered_capacity: - raise RuntimeError( - f"Encoder cache tensor ({nbytes} bytes) exceeds ec_buffer_size " - f"({self._registered_capacity} bytes)." - ) - while self._registered_bytes + nbytes > self._registered_capacity: - evicted = False - for mm_hash, entry in list(self._tensor_by_hash.items()): - if not entry.in_flight and self._release_tensor_locked(mm_hash): - evicted = True - break - if not evicted: - raise RuntimeError( - "ECMooncakeConnector has no evictable registered memory." - ) - - def _start_producer_zmq_listener(self) -> None: - if self._zmq_thread is not None: - return - - def loop() -> None: - assert self._zmq_ctx is not None - sock = self._zmq_ctx.socket(zmq.REP) - sock.setsockopt(zmq.RCVTIMEO, 100) - port = sock.bind_to_random_port(f"tcp://{self._hostname}") - self._zmq_listen_addr = f"tcp://{self._hostname}:{port}" - logger.info("EC Mooncake pull listener at %s", self._zmq_listen_addr) - eng = self._ensure_engine() - try: - while True: - try: - raw = sock.recv() - except zmq.Again: - if self._zmq_stop.is_set(): - break - continue - except zmq.ContextTerminated: - break - try: - req = json.loads(raw.decode("utf-8")) - if req.get("op") != "pull": - sock.send_json({"ok": False, "err": "unknown op"}) - continue - dst_session = req["dst_session"] - items = req["items"] - with self._tensor_lock: - entries: list[_RegisteredTensor] = [] - leases: list[tuple[str, str]] = [] - dst_ptrs: list[int] = [] - lengths: list[int] = [] - for item in items: - mm_hash = str(item["mm_hash"]) - entry = self._tensor_by_hash.get(mm_hash) - nbytes = int(item["nbytes"]) - if entry is None or entry.tensor.nbytes != nbytes: - raise ValueError( - "unknown EC tensor or size mismatch" - ) - entries.append(entry) - leases.append((mm_hash, str(item["lease_id"]))) - dst_ptrs.append(int(item["dst_ptr"])) - lengths.append(nbytes) - if ( - self._registry is None - or not self._registry.consume_leases(leases) - ): - raise ValueError("invalid lease") - for (mm_hash, _), entry in zip(leases, entries): - entry.in_flight += 1 - self._tensor_by_hash.move_to_end(mm_hash) - try: - ret = eng.batch_transfer_sync_write( - dst_session, - [entry.tensor.data_ptr() for entry in entries], - dst_ptrs, - lengths, - ) - finally: - with self._tensor_lock: - for entry in entries: - entry.in_flight -= 1 - sock.send_json({"ok": ret == 0, "mooncake_ret": int(ret)}) - except Exception as e: - logger.exception("EC Mooncake pull handler error: %s", e) - try: - sock.send_json({"ok": False, "err": str(e)}) - except zmq.ZMQError: - break - finally: - sock.close(linger=0) - - self._zmq_ctx = zmq.Context() - self._zmq_thread = threading.Thread( - target=loop, name="ec-mooncake-zmq", daemon=True - ) - self._zmq_thread.start() - while self._zmq_listen_addr is None: - time.sleep(0.01) - - def _ensure_producer_services(self) -> None: - if self._producer_services_started: - return - if not self.is_producer or self._role != ECConnectorRole.WORKER: - return - self._ensure_engine() - self._start_producer_zmq_listener() - if is_local_first_rank(): - self._registry = ECMooncakeRegistryServer( - "0.0.0.0", self._registry_http_port - ) - self._registry.start() - self._producer_services_started = True - - def _get_pull_socket(self, producer_zmq: str) -> zmq.Socket: - sock = self._client_sockets.get(producer_zmq) - if sock is not None: - return sock - if self._client_zmq_ctx is None: - self._client_zmq_ctx = zmq.Context() - sock = self._client_zmq_ctx.socket(zmq.REQ) - sock.setsockopt(zmq.RCVTIMEO, 120_000) - sock.connect(producer_zmq) - self._client_sockets[producer_zmq] = sock - return sock - - def _send_pull(self, producer_zmq: str, pull: dict[str, Any]) -> dict[str, Any]: - sock = self._get_pull_socket(producer_zmq) - try: - sock.send_json(pull) - return sock.recv_json() - except zmq.ZMQError: - sock.close(linger=0) - self._client_sockets.pop(producer_zmq, None) - raise - def _unregister_memories(self, tensors: list[torch.Tensor]) -> None: assert self._engine is not None addresses = [tensor.data_ptr() for tensor in tensors] @@ -493,11 +838,92 @@ def _unregister_memories(self, tensors: list[torch.Tensor]) -> None: for address in addresses: self._pending_unregister.pop(address, None) - def _ensure_consumer_pool(self, device: torch.device) -> None: + @staticmethod + def _push_source_range(tensor: torch.Tensor) -> tuple[int, int]: + # Register exactly the bytes that will be transferred. One encoder + # batch returns its items as views of a single storage (models split + # the batched embeddings, e.g. `image_embeds.split(sizes)`), so + # registering the whole storage would overlap the per-tensor + # registration a sibling item takes -- and Mooncake rejects + # overlapping memory regions. + return tensor.data_ptr(), tensor.nbytes + + def _acquire_push_source_registrations( + self, tensors: list[torch.Tensor] + ) -> list[int]: + ranges: dict[int, tuple[int, torch.Tensor]] = {} + for tensor in tensors: + address, nbytes = self._push_source_range(tensor) + ranges.setdefault(address, (nbytes, tensor)) + + eng = self._ensure_engine() + acquired: list[int] = [] + new_addresses: list[int] = [] + new_lengths: list[int] = [] + with self._push_source_registration_lock: + for address, (nbytes, tensor) in ranges.items(): + entry = self._push_source_registrations.get(address) + if entry is not None: + if entry.nbytes != nbytes: + raise RuntimeError( + "Mooncake EC source storage changed size while registered" + ) + entry.users += 1 + acquired.append(address) + continue + new_addresses.append(address) + new_lengths.append(nbytes) + self._push_source_registrations[address] = _PushSourceRegistration( + tensor=tensor, + nbytes=nbytes, + ) + acquired.append(address) + + if new_addresses: + ret = eng.batch_register_memory(new_addresses, new_lengths) + if ret != 0: + for address in acquired: + entry = self._push_source_registrations[address] + entry.users -= 1 + if entry.users == 0: + del self._push_source_registrations[address] + raise RuntimeError("Mooncake EC source registration failed") + return acquired + + def _release_push_source_registrations(self, addresses: list[int]) -> bool: + if not addresses: + return True + with self._push_source_registration_lock: + unused = [] + for address in addresses: + entry = self._push_source_registrations.get(address) + if entry is None: + continue + entry.users -= 1 + if entry.users == 0: + unused.append(address) + if not unused: + return True + ret = self._ensure_engine().batch_unregister_memory(unused) + if ret != 0: + logger.warning( + "Keeping %d EC source tensors registered after Mooncake " + "unregistration failure", + len(unused), + ) + return False + for address in unused: + del self._push_source_registrations[address] + self._pending_unregister.pop(address, None) + return True + + def _ensure_consumer_pool( + self, device: torch.device, *, allow_host: bool = False + ) -> None: if ( self._consumer_pool is not None or self._consumer_pool_disabled - or device.type != "cuda" + or (device.type != "cuda" and not allow_host) ): return try: @@ -524,242 +950,1469 @@ def _ensure_consumer_pool(self, device: torch.device) -> None: pool.nbytes, ) + def _ensure_producer_pool(self, device: torch.device) -> None: + """Register one staging slab so pushes never register per transfer. + + Registering the encoder output itself costs more than the transfer + (register+unregister dominated the push path); staging into a slab + that is registered once trades that for a device-to-device copy. + """ + if self._producer_pool is not None or self._producer_pool_disabled: + return + with self._producer_pool_lock: + if self._producer_pool is not None or self._producer_pool_disabled: + return + try: + pool = torch.empty( + self._producer_pool_capacity, dtype=torch.uint8, device=device + ) + ret = self._ensure_engine().batch_register_memory( + [pool.data_ptr()], [pool.nbytes] + ) + if ret != 0: + raise RuntimeError(f"Mooncake returned {ret}") + except (RuntimeError, torch.OutOfMemoryError) as e: + self._producer_pool_disabled = True + logger.warning( + "Could not initialize the EC producer staging pool; falling " + "back to per-transfer registration: %s", + e, + ) + return + self._producer_pool = pool + self._producer_pool_allocator = _ContiguousAllocator(pool.nbytes) + logger.info( + "Registered %d-byte staging pool for Mooncake EC pushes", + pool.nbytes, + ) + + def _stage_push_sources( + self, tensors: list[torch.Tensor] + ) -> tuple[list[torch.Tensor], list[tuple[int, int]]] | None: + """Copy the batch into the staging pool; None if it does not fit.""" + if not tensors: + return [], [] + self._ensure_producer_pool(tensors[0].device) + pool = self._producer_pool + allocator = self._producer_pool_allocator + if pool is None or allocator is None: + return None + staged: list[torch.Tensor] = [] + regions: list[tuple[int, int]] = [] + with self._producer_pool_lock: + for tensor in tensors: + region = allocator.allocate(tensor.nbytes) + if region is None: + for offset, size in regions: + allocator.free(offset, size) + return None + regions.append(region) + offset = region[0] + staged.append( + pool.narrow(0, offset, tensor.nbytes) + .view(tensor.dtype) + .view(tensor.shape) + ) + for destination, source in zip(staged, tensors): + destination.copy_(source, non_blocking=True) + return staged, regions + + def _release_push_staging(self, regions: list[tuple[int, int]]) -> None: + allocator = self._producer_pool_allocator + if allocator is None or not regions: + return + with self._producer_pool_lock: + for offset, size in regions: + allocator.free(offset, size) + def _poll_consumer_pool_frees(self) -> None: allocator = self._consumer_pool_allocator if allocator is None: return - pending = [] - for event, allocation in self._consumer_pending_frees: - if event.query(): + with self._consumer_lock: + pending = [] + for event, allocation in self._consumer_pending_frees: + if event.query(): + allocator.free(allocation.offset, allocation.size) + else: + pending.append((event, allocation)) + self._consumer_pending_frees = pending + + def _reclaim_residents_locked( + self, allocator: _ContiguousAllocator, nbytes: int + ) -> tuple[int, int] | None: + """Give up retired items, oldest first, until `nbytes` fits. + + Called only when the pool cannot satisfy an allocation, so a retired + item survives until its memory is genuinely needed. + """ + + def evict(mm_hash: str, allocation: _ConsumerPoolAllocation) -> bool: + event = self._consumer_retire_events.pop(mm_hash, None) + if event is None or event.query(): allocator.free(allocation.offset, allocation.size) else: - pending.append((event, allocation)) - self._consumer_pending_frees = pending + self._consumer_pending_frees.append((event, allocation)) + self._consumer_reclaimed.add(mm_hash) + self._consumer_worker_metrics["residents_reclaimed"] += 1 + return True + + while self._consumer_residents.evict_lru(evict) is not None: + region = allocator.allocate(nbytes) + if region is not None: + return region + return None + + def _take_resident_tensor(self, spec: ECMooncakeLoadSpec) -> torch.Tensor | None: + """Hand back a copy the pool still holds. + + Retired and in-use entries live in the same map, so an item a later + push reserved again still serves this load. + """ + with self._consumer_lock: + allocation = self._consumer_residents.get(spec.mm_hash) + if allocation is None: + self._consumer_worker_metrics["residents_missed"] += 1 + return None + tensor = allocation.tensor + if ( + tuple(tensor.shape) != tuple(spec.shape) + or str(tensor.dtype).split(".")[-1] != spec.dtype + ): + self._consumer_worker_metrics["residents_mismatched"] += 1 + return None + self._consumer_residents.pin(spec.mm_hash) + self._consumer_retire_events.pop(spec.mm_hash, None) + self._consumer_worker_metrics["residents_promoted"] += 1 + return tensor def _release_stale_consumer_allocations( self, encoder_cache: dict[str, torch.Tensor] ) -> None: if self._consumer_pool is None: return - for mm_hash, allocation in list(self._consumer_allocations.items()): - if encoder_cache.get(mm_hash) is allocation.tensor: - continue - self._consumer_allocations.pop(mm_hash) - event = torch.Event() - event.record(torch.accelerator.current_stream(self._consumer_pool.device)) - self._consumer_pending_frees.append((event, allocation)) + with self._consumer_lock: + reserved_allocations = { + id(reservation.allocation) + for reservation in self._push_reservations.values() + } + # Walk only the referenced entries: the retired set grows to + # thousands and none of it can change state here. + for mm_hash in self._consumer_residents.referenced(): + allocation = self._consumer_residents.get(mm_hash) + if allocation is None: + continue + if encoder_cache.get(mm_hash) is allocation.tensor: + continue + if id(allocation) in reserved_allocations: + continue + # Retire rather than free: the bytes stay valid and serve the + # next request that needs this item. The event orders the + # eventual reuse behind whatever still reads the tensor. + event = torch.Event() + event.record( + torch.accelerator.current_stream(self._consumer_pool.device) + ) + self._consumer_retire_events[mm_hash] = event + self._consumer_residents.retire(mm_hash) + self._consumer_worker_metrics["residents_retired"] += 1 self._poll_consumer_pool_frees() - def _allocate_consumer_tensor( + def _clear_item_timers(self, mm_hash: str) -> None: + self._consumer_missing_since.pop(mm_hash, None) + self._consumer_loading_since.pop(mm_hash, None) + self._consumer_pending_since.pop(mm_hash, None) + self._stalled_hashes.discard(mm_hash) + + def _note_awaiting_push( self, - spec: ECMooncakeLoadSpec, - dtype: torch.dtype, - device: torch.device, - ) -> tuple[torch.Tensor, _ConsumerPoolAllocation | None]: - expected_nbytes = ( - math.prod(spec.shape) * torch.empty((), dtype=dtype).element_size() + mm_hash: str, + transfer_id: str | None = None, + request_id: str | None = None, + ) -> bool: + """Wait for an item with nothing in flight, and give up on timeout. + + Nothing on this side can produce the item, so a push that never + arrives would defer the request forever. Past the timeout the request + is reported unavailable instead: the scheduler fails it with a + retryable error and the caller can re-issue it, which re-runs the + encode and produces a fresh transfer. + + Returns: + True once this request has been given up on. + """ + now = time.monotonic() + since = self._consumer_missing_since.setdefault(mm_hash, now) + self._consumer_scheduler_metrics["missing_event"] += 1 + elapsed = now - since + if elapsed < self._push_wait_timeout: + return False + stale = mm_hash in self._stalled_hashes + if request_id is not None: + self._unavailable_requests.add(request_id) + self._consumer_scheduler_metrics["given_up"] += 1 + # Start a fresh window: a re-issued request pushes this item again, + # and it must be allowed to wait for that push rather than inherit + # this one's deadline and be given up on immediately. Only the + # deadline resets -- `_stalled_hashes` keeps the warning to one per + # hash, while `given_up` counts every occurrence. + self._consumer_missing_since.pop(mm_hash, None) + if stale: + return request_id is not None + self._stalled_hashes.add(mm_hash) + self._consumer_scheduler_metrics["stalled"] += 1 + # Ask the worker what it knows about this transfer: whether the + # reservation exists at all separates "the producer never sent it" + # from "it arrived and the scheduler missed it". + reservation: Any = "unknown" + if transfer_id and self._reservation_zmq_addr is not None: + try: + reservation = self._send_control( + self._reservation_zmq_addr, + {"op": "status", "transfer_id": transfer_id}, + ) + except Exception as e: # noqa: BLE001 - diagnostic only + reservation = f"status failed: {e}" + logger.warning( + "EC Mooncake waited %.1fs for a push of mm_hash=%s " + "(transfer_id=%s) that never arrived; worker reservation=%s; " + "requests needing it fail with a retryable error.", + elapsed, + mm_hash, + transfer_id, + reservation, ) - if expected_nbytes != spec.nbytes: - raise ValueError( - f"EC tensor size mismatch for {spec.mm_hash}: metadata has " - f"{spec.nbytes} bytes, shape and dtype require {expected_nbytes}." + return request_id is not None + + def take_unavailable_requests(self) -> set[str]: + given_up = self._unavailable_requests + self._unavailable_requests = set() + return given_up + + @staticmethod + def _hash_samples(values: list[str], limit: int = 5) -> list[str]: + return [value[:16] for value in values[:limit]] + + def _maybe_log_consumer_worker_metrics(self) -> None: + now = time.monotonic() + if ( + self._consumer_metrics_log_interval <= 0 + or now - self._consumer_metrics_started_at + < self._consumer_metrics_log_interval + ): + return + with self._consumer_lock: + ready = [ + mm_hash + for mm_hash, reservation in self._push_reservations.items() + if reservation.ready + ] + pending = [ + mm_hash + for mm_hash, reservation in self._push_reservations.items() + if not reservation.ready + ] + metrics = dict(self._consumer_worker_metrics) + self._consumer_worker_metrics.clear() + residents = len(self._consumer_residents) + live = len(self._consumer_residents.referenced()) + retired = self._consumer_residents.num_evictable + pending_frees = len(self._consumer_pending_frees) + oldest_reservation_ms = max( + ( + (now - reservation.created_at) * 1000 + for reservation in self._push_reservations.values() + ), + default=0.0, ) + logger.info( + "EC Mooncake consumer worker: lifecycle=%s, reservations_ready=%d, " + "reservations_pending=%d, residents=%d, live=%d, retired=%d, " + "pending_frees=%d, " + "oldest_reservation_ms=%.1f, ready_hashes=%s, pending_hashes=%s", + metrics, + len(ready), + len(pending), + residents, + live, + retired, + pending_frees, + oldest_reservation_ms, + self._hash_samples(ready), + self._hash_samples(pending), + ) + self._consumer_metrics_started_at = now + + def _maybe_log_consumer_scheduler_metrics(self) -> None: + now = time.monotonic() + if ( + self._consumer_metrics_log_interval <= 0 + or now - self._consumer_metrics_started_at + < self._consumer_metrics_log_interval + ): + return + missing = sorted(self._consumer_missing_since.items(), key=lambda item: item[1]) + loading = sorted(self._consumer_loading_since.items(), key=lambda item: item[1]) + pending = sorted(self._consumer_pending_since.items(), key=lambda item: item[1]) + oldest_missing_ms = round((now - missing[0][1]) * 1000, 1) if missing else 0.0 + oldest_loading_ms = round((now - loading[0][1]) * 1000, 1) if loading else 0.0 + oldest_pending_ms = round((now - pending[0][1]) * 1000, 1) if pending else 0.0 + logger.info( + "EC Mooncake consumer scheduler: decisions=%s, ready=%d, loading=%d, " + "resident=%d, pending_specs=%d, needs_load=%d, missing=%d, " + "oldest_missing_ms=%.1f, oldest_loading_ms=%.1f, " + "oldest_pending_ms=%.1f, missing_hashes=%s, loading_hashes=%s, " + "pending_hashes=%s", + dict(self._consumer_scheduler_metrics), + len(self._ready_hashes), + len(self._loading_hashes), + len(self._resident_specs), + len(self._pending_specs), + len(self._mm_datas_need_loads), + len(missing), + oldest_missing_ms, + oldest_loading_ms, + oldest_pending_ms, + self._hash_samples([mm_hash for mm_hash, _ in missing]), + self._hash_samples([mm_hash for mm_hash, _ in loading]), + self._hash_samples([mm_hash for mm_hash, _ in pending]), + ) + self._consumer_scheduler_metrics.clear() + self._consumer_metrics_started_at = now - self._ensure_consumer_pool(device) + def _expire_push_reservations_locked(self) -> None: + now = time.monotonic() allocator = self._consumer_pool_allocator - pool = self._consumer_pool - if allocator is not None and pool is not None: - region = allocator.allocate(spec.nbytes) - if region is not None: - offset, size = region - tensor = ( - pool.narrow(0, offset, spec.nbytes).view(dtype).view(spec.shape) + assert allocator is not None + for transfer_id, reservation in list(self._push_reservations.items()): + if reservation.expires_at > now: + continue + if reservation.owns_allocation: + allocator.free( + reservation.allocation.offset, reservation.allocation.size + ) + self._push_reservations.pop(transfer_id) + self._consumer_worker_metrics["reservations_expired"] += 1 + for transfer_id, expires_at in list(self._cancelled_transfers.items()): + if expires_at <= now: + self._cancelled_transfers.pop(transfer_id) + + def _expire_push_reservations(self) -> int: + with self._consumer_lock: + before = len(self._push_reservations) + self._expire_push_reservations_locked() + return before - len(self._push_reservations) + + def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: + transfer_id = str(payload["transfer_id"]) + mm_hash = str(payload["mm_hash"]) + nbytes = int(payload["nbytes"]) + shape = tuple(int(value) for value in payload["shape"]) + dtype_name = str(payload["dtype"]) + dtype = getattr(torch, dtype_name, None) + if dtype is None: + raise ValueError(f"Unsupported torch dtype string: {dtype_name!r}") + expected_nbytes = math.prod(shape) * dtype.itemsize + if expected_nbytes != nbytes: + raise ValueError("shape and dtype do not match nbytes") + + with self._consumer_lock: + self._expire_push_reservations_locked() + if transfer_id in self._cancelled_transfers: + self._consumer_worker_metrics["reservations_cancelled_early"] += 1 + return { + "reservation_id": "", + "dst_session": "", + "dst_ptr": 0, + "nbytes": nbytes, + "write": False, + "ready": False, + "cancelled": True, + } + existing = self._push_reservations.get(transfer_id) + if existing is not None: + if ( + existing.mm_hash != mm_hash + or existing.shape != shape + or existing.dtype != dtype_name + ): + raise ValueError("conflicting reservation for transfer_id") + reservation = existing + should_write = False + key = ( + "reservations_reused_ready" + if existing.ready + else ("reservations_reused_pending") + ) + self._consumer_worker_metrics[key] += 1 + if not existing.ready: + existing.expires_at = time.monotonic() + _LEASE_TTL_SECONDS + else: + cached = self._consumer_residents.get(mm_hash) + if cached is not None: + if ( + tuple(cached.tensor.shape) != shape + or cached.tensor.dtype != dtype + ): + raise ValueError("conflicting cached tensor for mm_hash") + reservation = _PushReservation( + mm_hash=mm_hash, + reservation_id=uuid.uuid4().hex, + allocation=cached, + shape=shape, + dtype=dtype_name, + ready=True, + owns_allocation=False, + expires_at=time.monotonic() + _LEASE_TTL_SECONDS, + ) + should_write = False + # Live again: it must not be reclaimed under pressure. + self._consumer_residents.pin(mm_hash) + self._consumer_retire_events.pop(mm_hash, None) + self._consumer_worker_metrics["reservations_cached"] += 1 + else: + pool = self._consumer_pool + allocator = self._consumer_pool_allocator + assert pool is not None and allocator is not None + region = allocator.allocate(nbytes) + if region is None: + self._expire_push_reservations_locked() + region = allocator.allocate(nbytes) + if region is None: + region = self._reclaim_residents_locked(allocator, nbytes) + if region is None: + raise RuntimeError("EC consumer buffer pool is full") + offset, size = region + tensor = pool.narrow(0, offset, nbytes).view(dtype).view(shape) + allocation = _ConsumerPoolAllocation(offset, size, tensor) + reservation = _PushReservation( + mm_hash=mm_hash, + reservation_id=uuid.uuid4().hex, + allocation=allocation, + shape=shape, + dtype=dtype_name, + expires_at=time.monotonic() + _LEASE_TTL_SECONDS, + ) + should_write = True + self._consumer_worker_metrics["reservations_created"] += 1 + self._push_reservations[transfer_id] = reservation + + eng = self._ensure_engine() + return { + "reservation_id": reservation.reservation_id, + "dst_session": f"{self._hostname}:{eng.get_rpc_port()}", + "dst_ptr": reservation.allocation.tensor.data_ptr(), + "nbytes": reservation.allocation.tensor.nbytes, + "write": should_write, + "ready": reservation.ready, + "cached": not reservation.owns_allocation, + } + + def _push_status(self, transfer_id: str) -> dict[str, Any] | None: + with self._consumer_lock: + reservation = self._push_reservations.get(transfer_id) + if reservation is None: + return None + return { + "mm_hash": reservation.mm_hash, + "ready": reservation.ready, + "reservation_id": reservation.reservation_id, + "nbytes": reservation.allocation.tensor.nbytes, + "shape": list(reservation.shape), + "dtype": reservation.dtype, + } + + def _complete_push(self, transfer_id: str, reservation_id: str) -> _PushCompletion: + with self._consumer_lock: + reservation = self._push_reservations.get(transfer_id) + if reservation is None or reservation.reservation_id != reservation_id: + self._consumer_worker_metrics["completions_rejected"] += 1 + return _PushCompletion(False) + if reservation.ready: + self._consumer_worker_metrics["completions_repeated"] += 1 + return _PushCompletion(True) + self._consumer_worker_metrics["completions_accepted"] += 1 + if reservation.discard_on_complete: + allocator = self._consumer_pool_allocator + assert allocator is not None + self._push_reservations.pop(transfer_id) + if reservation.owns_allocation: + allocator.free( + reservation.allocation.offset, reservation.allocation.size + ) + self._consumer_worker_metrics["reservations_discarded"] += 1 + return _PushCompletion(True) + reservation.ready = True + reservation.expires_at = time.monotonic() + _LEASE_TTL_SECONDS + return _PushCompletion(True, became_ready=True) + + def _cancel_push( + self, transfer_id: str, reservation_id: str, abandon: bool = False + ) -> bool: + with self._consumer_lock: + reservation = self._push_reservations.get(transfer_id) + if ( + reservation is not None + and reservation_id + and reservation.reservation_id != reservation_id + ): + self._consumer_worker_metrics["cancellations_rejected"] += 1 + return False + self._cancelled_transfers[transfer_id] = ( + time.monotonic() + _LEASE_TTL_SECONDS + ) + if reservation is None: + self._consumer_worker_metrics["cancellations_pre_reserved"] += 1 + return True + allocator = self._consumer_pool_allocator + assert allocator is not None + if not reservation.ready and not abandon: + reservation.discard_on_complete = True + self._consumer_worker_metrics["cancellations_deferred"] += 1 + return True + self._push_reservations.pop(transfer_id) + if reservation.owns_allocation: + allocator.free( + reservation.allocation.offset, reservation.allocation.size + ) + self._consumer_worker_metrics["reservations_cancelled"] += 1 + return True + + def _take_pushed_tensor( + self, spec: ECMooncakeLoadSpec + ) -> tuple[torch.Tensor, _ConsumerPoolAllocation]: + with self._consumer_lock: + reservation = self._push_reservations.get(spec.transfer_id) + if ( + reservation is None + or not reservation.ready + or reservation.reservation_id != spec.reservation_id + ): + self._consumer_worker_metrics["takes_rejected"] += 1 + raise RuntimeError( + f"Pushed EC tensor is not ready for mm_hash={spec.mm_hash}" + ) + self._push_reservations.pop(spec.transfer_id) + self._consumer_residents.insert( + spec.mm_hash, reservation.allocation, reservation.allocation.size + ) + self._consumer_worker_metrics["reservations_taken"] += 1 + return reservation.allocation.tensor, reservation.allocation + + def _send_control(self, addr: str, request: dict[str, Any]) -> Any: + return self._control_channel.request(addr, request) + + def _reserve_remote(self, spec: ECMooncakePushSpec) -> dict[str, Any]: + result = self._send_control( + spec.consumer_zmq, + { + "op": "reserve", + "transfer_id": spec.transfer_id, + "mm_hash": spec.mm_hash, + "nbytes": spec.nbytes, + "shape": list(spec.shape), + "dtype": spec.dtype, + }, + ) + if not isinstance(result, dict): + raise RuntimeError("Invalid EC reservation response") + result["_received_at"] = time.monotonic() + return result + + def _cancel_remote( + self, consumer_zmq: str, transfer_id: str, reservation_id: str + ) -> bool: + result = self._send_control( + consumer_zmq, + { + "op": "cancel", + "transfer_id": transfer_id, + "reservation_id": reservation_id, + }, + ) + return isinstance(result, dict) and bool(result.get("cancelled")) + + def _poll_pending_cancels(self) -> None: + pending = {} + for transfer_id, future in self._pending_cancels.items(): + if not future.done(): + pending[transfer_id] = future + continue + try: + cancelled = future.result() + except Exception: + self._cancelled_transfer_ids.discard(transfer_id) + self._consumer_scheduler_metrics["cancellations_failed"] += 1 + logger.warning( + "EC Mooncake reservation cancellation failed", exc_info=True ) - return tensor, _ConsumerPoolAllocation(offset, size, tensor) + else: + key = "cancellations_completed" if cancelled else "cancellations_stale" + self._consumer_scheduler_metrics[key] += 1 + self._pending_cancels = pending - return torch.empty(spec.shape, dtype=dtype, device=device), None + def start_save_caches(self, **kwargs: Any) -> None: + metadata = self._get_connector_metadata() + assert isinstance(metadata, ECMooncakeConnectorMetadata) + for spec in metadata.pushes: + reservation = self._control_executor.submit(self._reserve_remote, spec) + self._pending_reservations.setdefault(spec.mm_hash, deque()).append( + (spec, reservation) + ) + encoder_cache = kwargs.get("encoder_cache") + if not isinstance(encoder_cache, dict): + return + for mm_hash in dict.fromkeys(spec.mm_hash for spec in metadata.pushes): + tensor = encoder_cache.get(mm_hash) + if tensor is not None: + self._submit_reserved_pushes(tensor, mm_hash) def start_load_caches( self, encoder_cache: dict[str, torch.Tensor], **kwargs: Any ) -> None: metadata = self._get_connector_metadata() assert isinstance(metadata, ECMooncakeConnectorMetadata) - eng = self._ensure_engine() + self._ensure_engine() raw_buf = self._ec_cfg.ec_buffer_device buf = raw_buf.lower() if isinstance(raw_buf, str) and raw_buf else "cuda" if buf == "cuda" and not torch.accelerator.is_available(): raise RuntimeError( "ECMooncakeConnector requires CUDA for ec_buffer_device=cuda" ) - device = torch.device(buf) - self._release_stale_consumer_allocations(encoder_cache) - pending: list[ - tuple[ECMooncakeLoadSpec, torch.Tensor, _ConsumerPoolAllocation | None] - ] = [] for spec in metadata.loads: if spec.mm_hash in encoder_cache: + if spec.pushed: + self._cancel_push(spec.transfer_id, spec.reservation_id) + self._completed_loads.add(spec.mm_hash) + continue + if spec.local: + resident = self._take_resident_tensor(spec) + if resident is None: + # Reclaimed before the scheduler heard about it; the load + # falls back to a transfer on a later step. + self._failed_loads.add(spec.mm_hash) + else: + encoder_cache[spec.mm_hash] = resident + self._completed_loads.add(spec.mm_hash) continue - torch_dtype = getattr(torch, spec.dtype, None) - if torch_dtype is None: - raise ValueError(f"Unsupported torch dtype string: {spec.dtype!r}") - tensor, allocation = self._allocate_consumer_tensor( - spec, torch_dtype, device + if spec.pushed: + try: + pushed_tensor, _ = self._take_pushed_tensor(spec) + except RuntimeError as e: + logger.warning("EC Mooncake pushed load failed: %s", e) + self._failed_loads.add(spec.mm_hash) + continue + encoder_cache[spec.mm_hash] = pushed_tensor + self._completed_loads.add(spec.mm_hash) + continue + logger.warning( + "EC Mooncake load for mm_hash=%s has no transfer to take", + spec.mm_hash, ) - pending.append((spec, tensor, allocation)) - if not pending: - return - - standalone = [tensor for _, tensor, allocation in pending if allocation is None] - standalone_registered = False + self._failed_loads.add(spec.mm_hash) + + def _push_batch(self, pushes: list[_PendingPush]) -> None: + started_at = time.monotonic() + with self._push_perf_lock: + self._queued_transfer_batches -= 1 + self._active_transfer_batches += 1 + + queue_waits_ms = [ + max(0, started_at - push.enqueued_at) * 1000 for push in pushes + ] + stage_ms = { + "queue": sum(queue_waits_ms), + "reserve": 0.0, + "cuda": 0.0, + "register": 0.0, + "rdma": 0.0, + "unregister": 0.0, + "complete": 0.0, + } + ready: list[tuple[_PendingPush, dict[str, Any]]] = [] + notifications: list[tuple[_PendingPush, dict[str, Any]]] = [] + failed = False try: - if standalone: - ret = eng.batch_register_memory( - [tensor.data_ptr() for tensor in standalone], - [tensor.nbytes for tensor in standalone], - ) - if ret != 0: + for push in pushes: + stage_started_at = time.monotonic() + reservation = push.reservation.result() + received_at = float(reservation.get("_received_at", started_at)) + if ( + not reservation.get("ready", False) + and time.monotonic() - received_at >= _RESERVATION_REFRESH_SECONDS + ): + reservation = self._reserve_remote(push.spec) + stage_ms["reserve"] += (time.monotonic() - stage_started_at) * 1000 + if reservation.get("cached", False) or reservation.get( + "cancelled", False + ): + continue + if not reservation.get("write", True): + continue + if push.ready_event is not None: + stage_started_at = time.monotonic() + push.ready_event.synchronize() + stage_ms["cuda"] += (time.monotonic() - stage_started_at) * 1000 + if int(reservation["nbytes"]) != push.tensor.nbytes: raise RuntimeError( - "Mooncake EC batch_register_memory failed on consumer." + "Reserved EC size does not match tensor for " + f"mm_hash={push.spec.mm_hash}" + ) + ready.append((push, reservation)) + notifications.append((push, reservation)) + if not ready and not notifications: + return + + if ready: + eng = self._ensure_engine() + tensors = [push.tensor for push, _ in ready] + lengths = [tensor.nbytes for tensor in tensors] + stage_started_at = time.monotonic() + staged = self._stage_push_sources(tensors) + registered_sources: list[int] = [] + staged_regions: list[tuple[int, int]] = [] + if staged is not None: + sources, staged_regions = staged + # The NIC reads outside the CUDA stream, so the staging + # copies have to have landed before the transfer starts. + if sources and sources[0].device.type == "cuda": + torch.accelerator.current_stream( + sources[0].device + ).synchronize() + else: + sources = tensors + registered_sources = self._acquire_push_source_registrations( + tensors ) - standalone_registered = True - batches: dict[ - str, - list[ - tuple[ - ECMooncakeLoadSpec, - torch.Tensor, - _ConsumerPoolAllocation | None, - ] - ], - ] = {} - for item in pending: - batches.setdefault(item[0].producer_zmq, []).append(item) - for producer_zmq, batch in batches.items(): - pull = { - "op": "pull", - "dst_session": f"{self._hostname}:{eng.get_rpc_port()}", + addresses = [tensor.data_ptr() for tensor in sources] + stage_ms["register"] = (time.monotonic() - stage_started_at) * 1000 + try: + by_session: dict[str, list[int]] = {} + for index, (_, reservation) in enumerate(ready): + by_session.setdefault( + str(reservation["dst_session"]), [] + ).append(index) + stage_started_at = time.monotonic() + for session, indices in by_session.items(): + ret = eng.batch_transfer_sync_write( + session, + [addresses[index] for index in indices], + [int(ready[index][1]["dst_ptr"]) for index in indices], + [lengths[index] for index in indices], + ) + if ret != 0: + raise RuntimeError( + f"Mooncake EC push failed with status {ret}" + ) + stage_ms["rdma"] = (time.monotonic() - stage_started_at) * 1000 + finally: + stage_started_at = time.monotonic() + self._release_push_staging(staged_regions) + self._release_push_source_registrations(registered_sources) + stage_ms["unregister"] = ( + time.monotonic() - stage_started_at + ) * 1000 + + stage_started_at = time.monotonic() + self._notify_completions(notifications) + stage_ms["complete"] = (time.monotonic() - stage_started_at) * 1000 + except Exception: + # A failed batch must not take the engine down with it: the + # consumer is told to drop its reservations and this item falls + # back to whatever the consumer can still do (pull, or a local + # re-encode). Raising here would surface in + # `build_connector_worker_meta` as a fatal EngineCore error. + failed = True + logger.exception( + "EC Mooncake push batch failed for mm_hashes=%s", + [push.spec.mm_hash for push in pushes], + ) + self._abandon_pushes(pushes) + finally: + with self._active_push_sources_lock: + for push in pushes: + key = (push.spec.mm_hash, id(push.tensor)) + self._active_push_sources[key] -= 1 + if self._active_push_sources[key] == 0: + del self._active_push_sources[key] + stage_ms["total"] = (time.monotonic() - started_at) * 1000 + self._record_push_perf( + stage_ms, + stage_max_ms={"queue": max(queue_waits_ms, default=0.0)}, + item_count=len(pushes), + byte_count=sum(push.tensor.nbytes for push, _ in ready), + skipped_items=len(pushes) - len(ready), + failed=failed, + ) + + def _notify_completions( + self, notifications: list[tuple[_PendingPush, dict[str, Any]]] + ) -> None: + """Tell the consumer, in one message per destination, what landed.""" + if not notifications: + return + by_destination: dict[str, list[tuple[_PendingPush, dict[str, Any]]]] = {} + for push, reservation in notifications: + by_destination.setdefault(push.spec.consumer_zmq, []).append( + (push, reservation) + ) + for consumer_zmq, items in by_destination.items(): + result = self._send_control( + consumer_zmq, + { + "op": "complete_batch", "items": [ { - "mm_hash": spec.mm_hash, - "dst_ptr": tensor.data_ptr(), - "nbytes": tensor.nbytes, - "lease_id": spec.lease_id, + "transfer_id": push.spec.transfer_id, + "reservation_id": reservation["reservation_id"], } - for spec, tensor, _ in batch + for push, reservation in items ], - } - resp = self._send_pull(producer_zmq, pull) - if not resp.get("ok"): - raise RuntimeError(f"EC Mooncake pull failed: {resp}") + }, + ) + completions = result.get("items", []) if isinstance(result, dict) else [] + if len(completions) != len(items): + raise RuntimeError("Malformed EC completion response") + for (push, _), completion in zip(items, completions): + if not completion.get("completed"): + raise RuntimeError( + f"Unknown EC reservation for mm_hash={push.spec.mm_hash}" + ) + + def _abandon_pushes(self, pushes: list[_PendingPush]) -> None: + """Release the consumer-side reservations of a batch that failed.""" + for push in pushes: + reservation_id = "" + if push.reservation.done() and not push.reservation.cancelled(): + with suppress(Exception): + result = push.reservation.result() + reservation_id = str(result.get("reservation_id", "")) + with suppress(Exception): + self._send_control( + push.spec.consumer_zmq, + { + "op": "cancel", + "transfer_id": push.spec.transfer_id, + "reservation_id": reservation_id, + "abandon": True, + }, + ) + + def _record_push_perf( + self, + stage_ms: dict[str, float], + *, + stage_max_ms: dict[str, float], + item_count: int, + byte_count: int, + skipped_items: int, + failed: bool, + ) -> None: + now = time.monotonic() + report: tuple[_PushPerfWindow, int, int] | None = None + with self._push_perf_lock: + self._active_transfer_batches -= 1 + perf = self._push_perf + perf.batches += 1 + perf.items += item_count + perf.bytes += byte_count + perf.skipped_items += skipped_items + perf.failures += int(failed) + for stage, elapsed_ms in stage_ms.items(): + perf.stage_totals_ms[stage] = ( + perf.stage_totals_ms.get(stage, 0.0) + elapsed_ms + ) + perf.stage_max_ms[stage] = max( + perf.stage_max_ms.get(stage, 0.0), + stage_max_ms.get(stage, elapsed_ms), + ) + if ( + self._transfer_metrics_log_interval > 0 + and now - perf.started_at >= self._transfer_metrics_log_interval + ): + report = ( + perf, + self._active_transfer_batches, + self._queued_transfer_batches, + ) + self._push_perf = _PushPerfWindow(started_at=now) + if report is None: + return + perf, active_batches, queued_batches = report + batches = max(perf.batches, 1) + items = max(perf.items, 1) + stage_parts = [] + for stage in ( + "queue", + "reserve", + "cuda", + "register", + "rdma", + "unregister", + "complete", + "total", + ): + divisor = items if stage == "queue" else batches + average = perf.stage_totals_ms.get(stage, 0.0) / divisor + maximum = perf.stage_max_ms.get(stage, 0.0) + stage_parts.append(f"{stage}_ms={average:.1f}/{maximum:.1f}") + stage_summary = " ".join(stage_parts) + producer_metrics = dict(self._producer_metrics) + self._producer_metrics.clear() + logger.info( + "EC Mooncake push perf: batches=%d items=%d bytes=%d " + "batch_items=%.1f skipped=%d failures=%d active=%d queued=%d " + "producer=%s queue_item_avg/max and stage_batch_avg/max: %s", + perf.batches, + perf.items, + perf.bytes, + perf.items / batches, + perf.skipped_items, + perf.failures, + active_batches, + queued_batches, + producer_metrics, + stage_summary, + ) + + def _flush_pending_pushes(self) -> None: + if not self._pending_pushes: + return + grouped: dict[str, list[_PendingPush]] = {} + for push in self._pending_pushes: + grouped.setdefault(push.spec.consumer_zmq, []).append(push) + self._pending_pushes = [] + for pushes in grouped.values(): + with self._push_perf_lock: + self._queued_transfer_batches += 1 + future = self._io_executor.submit(self._push_batch, pushes) + hashes = ",".join(push.spec.mm_hash for push in pushes) + self._pending_saves.append((hashes, future)) + + def _submit_push( + self, + tensor: torch.Tensor, + spec: ECMooncakePushSpec, + reservation: Future[dict[str, Any]], + ) -> None: + ready_event = None + if tensor.device.type == "cuda": + ready_event = torch.Event() + ready_event.record(torch.accelerator.current_stream(tensor.device)) + self._pending_pushes.append( + _PendingPush( + tensor=tensor, + spec=spec, + reservation=reservation, + ready_event=ready_event, + enqueued_at=time.monotonic(), + ) + ) + + def _submit_reserved_pushes(self, tensor: torch.Tensor, mm_hash: str) -> None: + reservations = self._pending_reservations.pop(mm_hash, deque()) + if reservations: + with self._active_push_sources_lock: + self._active_push_sources[(mm_hash, id(tensor))] += len(reservations) + for spec, reservation in reservations: + self._submit_push(tensor, spec, reservation) + + def _cancel_orphaned_reservation( + self, + spec: ECMooncakePushSpec, + reservation: Future[dict[str, Any]], + ) -> None: + try: + result = reservation.result() + if result.get("cached", False) or result.get("cancelled", False): + return + self._send_control( + spec.consumer_zmq, + { + "op": "cancel", + "transfer_id": spec.transfer_id, + "reservation_id": str(result.get("reservation_id", "")), + "abandon": True, + }, + ) except Exception: - allocator = self._consumer_pool_allocator - if allocator is not None: - for _, _, allocation in pending: - if allocation is not None: - allocator.free(allocation.offset, allocation.size) - raise - finally: - if standalone_registered: - self._unregister_memories(standalone) - for spec, tensor, allocation in pending: - encoder_cache[spec.mm_hash] = tensor - if allocation is not None: - self._consumer_allocations[spec.mm_hash] = allocation - logger.debug("Loaded EC tensor for mm_hash=%s via Mooncake", spec.mm_hash) + logger.exception( + "Failed to cancel orphaned EC reservation for transfer_id=%s", + spec.transfer_id, + ) + + def get_finished( + self, finished_req_ids: set[str] + ) -> tuple[set[str] | None, set[str] | None]: + if not self.is_producer or self._role != ECConnectorRole.WORKER: + return None, None + + orphaned: list[tuple[ECMooncakePushSpec, Future[dict[str, Any]]]] = [] + for mm_hash, reservations in list(self._pending_reservations.items()): + remaining: deque[tuple[ECMooncakePushSpec, Future[dict[str, Any]]]] = ( + deque() + ) + for spec, reservation in reservations: + if spec.request_id in finished_req_ids: + orphaned.append((spec, reservation)) + else: + remaining.append((spec, reservation)) + if remaining: + self._pending_reservations[mm_hash] = remaining + else: + self._pending_reservations.pop(mm_hash) + + for spec, reservation in orphaned: + future = self._io_executor.submit( + self._cancel_orphaned_reservation, spec, reservation + ) + self._pending_saves.append((f"cancel:{spec.transfer_id}", future)) + return None, None def save_caches( self, encoder_cache: dict[str, torch.Tensor], mm_hash: str, **kwargs: Any ) -> None: if not self.is_producer or self._role != ECConnectorRole.WORKER: return - self._ensure_producer_services() tensor = encoder_cache[mm_hash] - eng = self._ensure_engine() - dtype_str = str(tensor.dtype).split(".")[-1] - payload = { - "nbytes": tensor.nbytes, - "shape": list(tensor.shape), - "dtype": dtype_str, - "producer_zmq": self._zmq_listen_addr, - } - with self._tensor_lock: - if mm_hash in self._tensor_by_hash: - self._tensor_by_hash.move_to_end(mm_hash) + if mm_hash in self._pending_reservations: + self._submit_reserved_pushes(tensor, mm_hash) + + def _index_pending_spec(self, spec: ECMooncakeLoadSpec) -> None: + transfer_id = spec.transfer_id or spec.mm_hash + if transfer_id in self._pending_specs: + self._consumer_scheduler_metrics["events_duplicate"] += 1 + return + self._pending_specs[transfer_id] = spec + self._pending_specs_by_hash.setdefault(spec.mm_hash, deque()).append( + transfer_id + ) + self._pending_spec_deadlines[transfer_id] = ( + time.monotonic() + _LEASE_TTL_SECONDS + ) + self._consumer_missing_since.pop(spec.mm_hash, None) + self._consumer_pending_since.setdefault(spec.mm_hash, time.monotonic()) + + def _pop_pending_spec(self, transfer_id: str) -> ECMooncakeLoadSpec | None: + spec = self._pending_specs.pop(transfer_id, None) + self._pending_spec_deadlines.pop(transfer_id, None) + if spec is not None: + if not self._pending_specs_by_hash.get(spec.mm_hash): + self._consumer_pending_since.pop(spec.mm_hash, None) + transfer_ids = self._pending_specs_by_hash.get(spec.mm_hash) + if transfer_ids is not None: + with suppress(ValueError): + transfer_ids.remove(transfer_id) + if not transfer_ids: + self._pending_specs_by_hash.pop(spec.mm_hash, None) + self._consumer_pending_since.pop(spec.mm_hash, None) + else: + self._consumer_pending_since.pop(spec.mm_hash, None) + return spec + + def _first_pending_spec(self, mm_hash: str) -> ECMooncakeLoadSpec | None: + transfer_ids = self._pending_specs_by_hash.get(mm_hash) + if transfer_ids is None: + return None + while transfer_ids: + spec = self._pending_specs.get(transfer_ids[0]) + if spec is not None: + return spec + transfer_ids.popleft() + self._pending_specs_by_hash.pop(mm_hash, None) + self._consumer_pending_since.pop(mm_hash, None) + return None + + def _store_pushed_spec(self, data: dict[str, Any]) -> None: + transfer_id = str(data["transfer_id"]) + identifier = str(data["mm_hash"]) + reservation_id = str(data["reservation_id"]) + self._index_pending_spec( + ECMooncakeLoadSpec( + mm_hash=identifier, + num_token=0, + nbytes=int(data["nbytes"]), + shape=tuple(int(value) for value in data["shape"]), + dtype=str(data["dtype"]), + pushed=True, + transfer_id=transfer_id, + reservation_id=reservation_id, + ) + ) + + def _note_resident(self, spec: ECMooncakeLoadSpec) -> None: + """Record that the worker's receive pool now holds this item.""" + self._drop_resident(spec.mm_hash) + self._resident_specs[spec.mm_hash] = ECMooncakeLoadSpec( + mm_hash=spec.mm_hash, + num_token=0, + nbytes=spec.nbytes, + shape=spec.shape, + dtype=spec.dtype, + local=True, + ) + self._resident_bytes += spec.nbytes + while ( + self._resident_specs and self._resident_bytes > self._consumer_pool_capacity + ): + _, dropped = self._resident_specs.popitem(last=False) + self._resident_bytes -= dropped.nbytes + + def _drop_resident(self, mm_hash: str) -> None: + spec = self._resident_specs.pop(mm_hash, None) + if spec is not None: + self._resident_bytes -= spec.nbytes + + def _queue_cancel(self, transfer_id: str, reservation_id: str = "") -> None: + if ( + self._reservation_zmq_addr is None + or transfer_id in self._pending_cancels + or transfer_id in self._cancelled_transfer_ids + ): + return + self._cancelled_transfer_ids.add(transfer_id) + self._pending_cancels[transfer_id] = self._control_executor.submit( + self._cancel_remote, + self._reservation_zmq_addr, + transfer_id, + reservation_id, + ) + + def _expire_pending_specs(self) -> None: + now = time.monotonic() + for transfer_id, deadline in list(self._pending_spec_deadlines.items()): + if deadline > now: + continue + spec = self._pop_pending_spec(transfer_id) + if spec is not None: + self._consumer_pending_since.pop(spec.mm_hash, None) + self._consumer_scheduler_metrics["pending_specs_expired"] += 1 + self._queue_cancel(transfer_id, spec.reservation_id) + + def _ensure_event_channel(self) -> None: + if self._event_zmq_socket is not None: + return + assert self._reservation_zmq_addr is not None + event_port = self._send_control( + self._reservation_zmq_addr, + {"op": "event_port"}, + ) + address, _ = self._reservation_zmq_addr.rsplit(":", 1) + self._event_zmq_ctx = zmq.Context() + self._event_zmq_socket = self._event_zmq_ctx.socket(zmq.PULL) + self._event_zmq_socket.connect(f"{address}:{int(event_port)}") + + def _drain_push_notifications(self) -> None: + # `has_cache_item` and `ensure_cache_available` run once per request + # per multimodal item, so draining on every call rescans the cancel + # and deadline tables thousands of times per step. Once per step is + # enough: `build_connector_meta` re-arms this at the end of each one. + now = time.monotonic() + if not self._drain_pending and now - self._drained_at < _DRAIN_MIN_INTERVAL: + return + self._drain_pending = False + self._drained_at = now + self._poll_pending_cancels() + self._expire_pending_specs() + if self._reservation_zmq_addr is not None: + self._ensure_event_channel() + socket = self._event_zmq_socket + if socket is None: + return + while True: + try: + data = socket.recv_json(flags=zmq.DONTWAIT) + except zmq.Again: return - self._make_registration_space_locked(tensor.nbytes) - ret = eng.batch_register_memory([tensor.data_ptr()], [tensor.nbytes]) - if ret != 0: - raise RuntimeError( - "Mooncake EC batch_register_memory failed on producer." - ) - self._tensor_by_hash[mm_hash] = _RegisteredTensor(tensor) - self._registered_bytes += tensor.nbytes - if self._registry is not None: - self._registry.publish(mm_hash, payload) - logger.debug("Published EC tensor mm_hash=%s to registry", mm_hash) + identifier = str(data["mm_hash"]) + self._consumer_scheduler_metrics["events_received"] += 1 + if data.get("ready"): + self._consumer_scheduler_metrics["events_ready"] += 1 + if identifier in self._ready_hashes: + # Redundant only for as long as the hash stays ready; hold + # on to the spec so an eviction does not strand whoever + # this transfer belongs to. + self._consumer_scheduler_metrics["events_redundant"] += 1 + self._store_pushed_spec(data) + else: + self._consumer_scheduler_metrics["events_not_ready"] += 1 def has_cache_item(self, identifier: str) -> bool: if not self.is_consumer or self._role != ECConnectorRole.SCHEDULER: return False - assert self._remote_registry_url is not None - url = self._remote_registry_url.rstrip("/") + f"/ec/info/{identifier}" - try: - r = httpx.get(url, timeout=5.0) - except httpx.HTTPError as e: - logger.warning( - "EC Mooncake registry query failed for %s: %s", identifier, e - ) - return False - if r.status_code != 200: - return False - data = r.json() - zmq_addr = data.get("producer_zmq") - if not zmq_addr: + self._drain_push_notifications() + self._maybe_log_consumer_scheduler_metrics() + if identifier in self._ready_hashes: + self._consumer_scheduler_metrics["ready"] += 1 + self._clear_item_timers(identifier) + return True + if identifier in self._loading_hashes: + self._consumer_scheduler_metrics["loading"] += 1 return False - self._pending_specs[identifier] = ECMooncakeLoadSpec( - mm_hash=identifier, - num_token=0, - nbytes=int(data["nbytes"]), - shape=tuple(int(x) for x in data["shape"]), - dtype=str(data["dtype"]), - producer_zmq=str(zmq_addr), - lease_id=str(data["lease_id"]), + if identifier in self._resident_specs: + self._consumer_scheduler_metrics["resident"] += 1 + self._consumer_missing_since.pop(identifier, None) + return True + pending = self._first_pending_spec(identifier) + if pending is not None: + self._consumer_scheduler_metrics["pending_spec"] += 1 + self._consumer_missing_since.pop(identifier, None) + return True + self._consumer_scheduler_metrics["missing_event"] += 1 + self._consumer_missing_since.setdefault(identifier, time.monotonic()) + return False + + @staticmethod + def _request_transfer_id(request: Any, index: int) -> str | None: + params = getattr(request, "ec_transfer_params", None) or {} + items = params.get("ec_items") or [] + mm_hash = request.mm_features[index].identifier + if index < len(items): + item = items[index] + if item.get("mm_hash") in (None, mm_hash) and item.get("transfer_id"): + return str(item["transfer_id"]) + for item in items: + if item.get("mm_hash") == mm_hash and item.get("transfer_id"): + return str(item["transfer_id"]) + return None + + def ensure_cache_available( + self, + request: Any, + num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, + ) -> bool: + if self.is_producer: + for index, feature in enumerate(request.mm_features): + if ( + feature.mm_position.offset + feature.mm_position.length + > num_computed_tokens + ): + self._prepare_push_spec(request, index) + if not self.is_consumer or self._role != ECConnectorRole.SCHEDULER: + return True + + self._drain_push_notifications() + local_cache_hashes = local_cache_hashes or set() + all_ready = True + for index, feature in enumerate(request.mm_features): + if ( + feature.mm_position.offset + feature.mm_position.length + <= num_computed_tokens + ): + continue + mm_hash = feature.identifier + transfer_id = self._request_transfer_id(request, index) + if transfer_id is not None and transfer_id in self._pending_spec_deadlines: + # A live request still references this transfer, so keep it out + # of the orphan sweep in `_expire_pending_specs`. + self._pending_spec_deadlines[transfer_id] = ( + time.monotonic() + _LEASE_TTL_SECONDS + ) + if mm_hash in local_cache_hashes: + # Keep the transfer: `local_cache_hashes` is a snapshot, and + # the entry can be evicted before this request is scheduled. + # Cancelling here used to strand the request with no way to + # get the item back. `request_finished` releases it instead. + continue + if mm_hash in self._ready_hashes: + self._consumer_scheduler_metrics["ready"] += 1 + self._clear_item_timers(mm_hash) + continue + if mm_hash in self._loading_hashes: + self._consumer_scheduler_metrics["loading"] += 1 + all_ready = False + continue + # A resident copy is preferred over a transfer: it is already in + # this instance's memory, and using it leaves the transfer for + # whoever has no copy at all. + spec = self._resident_specs.get(mm_hash) + if spec is not None: + self._consumer_scheduler_metrics["resident_hit"] += 1 + else: + spec = ( + self._pending_specs.get(transfer_id) + if transfer_id is not None + else None + ) + if spec is None: + spec = self._first_pending_spec(mm_hash) + if spec is not None: + self._loading_hashes.add(mm_hash) + self._load_specs[mm_hash] = spec + self._consumer_loading_since.setdefault(mm_hash, time.monotonic()) + self._consumer_pending_since.pop(mm_hash, None) + self._mm_datas_need_loads[mm_hash] = request.get_num_encoder_embeds( + index + ) + self._scheduler_pending_work = True + all_ready = False + else: + # Keep waiting until the timeout, then let the request fail + # rather than hold a scheduler slot forever. + self._note_awaiting_push(mm_hash, transfer_id, request.request_id) + all_ready = False + return all_ready + + def _prepare_push_spec(self, request: Any, index: int) -> None: + params = getattr(request, "ec_transfer_params", None) or {} + consumer_zmq = params.get("consumer_zmq") + mm_hash = request.mm_features[index].identifier + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + transfer_id = f"{request.request_id}:{index}" + if not consumer_zmq or transfer_id in self._pushes_to_prepare: + return + num_tokens = request.get_num_encoder_embeds(index) + dtype = self._model_config.dtype + assert isinstance(dtype, torch.dtype) + dtype_name = str(dtype).split(".")[-1] + shape = (num_tokens, self._model_config.get_hidden_size()) + nbytes = math.prod(shape) * dtype.itemsize + self._pushes_to_prepare[transfer_id] = ECMooncakePushSpec( + mm_hash=mm_hash, + nbytes=nbytes, + shape=shape, + dtype=dtype_name, + consumer_zmq=str(consumer_zmq), + transfer_id=transfer_id, + request_id=request.request_id, ) - return True def update_state_after_alloc(self, request: Any, index: int) -> None: + mm_hash = request.mm_features[index].identifier + if self.is_producer: + self._prepare_push_spec(request, index) if not self.is_consumer: return - mm_hash = request.mm_features[index].identifier + if mm_hash in self._ready_hashes: + return + if mm_hash in self._loading_hashes: + return num_encoder_token = request.get_num_encoder_embeds(index) self._mm_datas_need_loads[mm_hash] = num_encoder_token + def update_state_after_free(self, request: Any, index: int) -> None: + """Release this request's transfer as soon as it consumed the item. + + Waiting for `request_finished` would keep a consumer buffer (and the + pool slot its reservation pins) alive for the whole generation. + """ + if not self.is_consumer or self._role != ECConnectorRole.SCHEDULER: + return + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + return + pending = self._pop_pending_spec(transfer_id) + self._queue_cancel( + transfer_id, + pending.reservation_id if pending is not None else "", + ) + def build_connector_meta( self, scheduler_output: SchedulerOutput ) -> ECConnectorMetadata: + for mm_hash in scheduler_output.free_encoder_mm_hashes: + self._ready_hashes.discard(mm_hash) + self._clear_item_timers(mm_hash) meta = ECMooncakeConnectorMetadata() + for push_spec in self._pushes_to_prepare.values(): + meta.add_push(push_spec) + self._pushes_to_prepare.clear() for mm_hash, num_token in self._mm_datas_need_loads.items(): - spec = self._pending_specs.get(mm_hash) - if spec is None: + load_spec = self._load_specs.pop(mm_hash, None) + if load_spec is None: logger.warning("Missing EC Mooncake spec for mm_hash=%s", mm_hash) continue meta.add_load( ECMooncakeLoadSpec( - mm_hash=spec.mm_hash, + mm_hash=load_spec.mm_hash, num_token=num_token, - nbytes=spec.nbytes, - shape=spec.shape, - dtype=spec.dtype, - producer_zmq=spec.producer_zmq, - lease_id=spec.lease_id, + nbytes=load_spec.nbytes, + shape=load_spec.shape, + dtype=load_spec.dtype, + pushed=load_spec.pushed, + transfer_id=load_spec.transfer_id, + reservation_id=load_spec.reservation_id, + local=load_spec.local, ) ) - self._pending_specs.pop(mm_hash, None) + # Either way the pool holds the item once this load lands, so it + # can serve the next request without another transfer. + self._note_resident(load_spec) + if not load_spec.local: + self._pop_pending_spec(load_spec.transfer_id or load_spec.mm_hash) self._mm_datas_need_loads.clear() + self._poll_pending_cancels() + self._maybe_log_consumer_scheduler_metrics() + self._drain_pending = True + return meta + + def build_connector_worker_meta(self) -> ECConnectorWorkerMetadata | None: + if self._role != ECConnectorRole.WORKER: + return None + + self._flush_pending_pushes() + saves = self._pending_saves + completed_saves = [] + self._pending_saves = [ + (mm_hash, future) for mm_hash, future in saves if not future.done() + ] + for mm_hash, future in saves: + if future.done(): + completed_saves.append((mm_hash, future)) + for mm_hash, future in completed_saves: + try: + future.result() + except Exception: + # Publishing is best-effort: a consumer that cannot fetch this + # item falls back to encoding it locally. Failing the step + # instead would take the whole engine down. + self._producer_metrics["saves_failed"] += 1 + logger.exception( + "EC Mooncake async save failed for mm_hash=%s", mm_hash + ) + with self._consumer_lock: + reclaimed = self._consumer_reclaimed + self._consumer_reclaimed = set() + meta = ECMooncakeWorkerMetadata( + loaded=self._completed_loads, + failed_loads=self._failed_loads, + reclaimed=reclaimed, + pending_loads=False, + pending_saves=bool(self._pending_saves), + ) + self._completed_loads = set() + self._failed_loads = set() + if self.is_consumer: + self._maybe_log_consumer_worker_metrics() return meta + def update_connector_output(self, connector_output: ECConnectorOutput) -> None: + meta = connector_output.ec_connector_worker_meta + if not isinstance(meta, ECMooncakeWorkerMetadata): + return + for mm_hash in meta.loaded: + self._loading_hashes.discard(mm_hash) + self._ready_hashes.add(mm_hash) + self._clear_item_timers(mm_hash) + self._consumer_scheduler_metrics["loads_completed"] += 1 + for mm_hash in meta.failed_loads: + self._loading_hashes.discard(mm_hash) + self._load_specs.pop(mm_hash, None) + self._drop_resident(mm_hash) + self._clear_item_timers(mm_hash) + self._consumer_scheduler_metrics["loads_failed"] += 1 + for mm_hash in meta.reclaimed: + self._drop_resident(mm_hash) + self._consumer_scheduler_metrics["resident_reclaimed"] += 1 + self._scheduler_pending_work = meta.pending_loads or meta.pending_saves + + def has_pending_push_work(self) -> bool: + return self._scheduler_pending_work + def _placeholder_metadata_fields(self, modality: str) -> set[str]: if modality in self._metadata_fields_cache: return self._metadata_fields_cache[modality] @@ -782,11 +2435,21 @@ def _placeholder_metadata_fields(self, modality: str) -> set[str]: return fields def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: + if self.is_consumer and self._role == ECConnectorRole.SCHEDULER: + for index in range(len(request.mm_features)): + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + continue + pending = self._pop_pending_spec(transfer_id) + self._queue_cancel( + transfer_id, + pending.reservation_id if pending is not None else "", + ) if not self.is_producer: return False, None items = [] - for feature in request.mm_features: + for index, feature in enumerate(request.mm_features): metadata = {} if feature.data is not None: wanted = self._placeholder_metadata_fields(feature.modality) @@ -795,7 +2458,11 @@ def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: for key, value in feature.data.get_data().items() if key in wanted and isinstance(value, torch.Tensor) } - items.append({"mm_hash": feature.identifier, **metadata}) + transfer_id = self._request_transfer_id(request, index) + item = {"mm_hash": feature.identifier, **metadata} + if transfer_id is not None: + item["transfer_id"] = transfer_id + items.append(item) if not items: return False, None @@ -805,18 +2472,17 @@ def shutdown(self) -> None: if self._shutdown: return self._shutdown = True - if self._registry is not None: - self._registry.shutdown() - self._zmq_stop.set() - if self._zmq_thread is not None: - self._zmq_thread.join() - if self._zmq_ctx is not None: - self._zmq_ctx.term() - for sock in self._client_sockets.values(): - sock.close(linger=0) - self._client_sockets.clear() - if self._client_zmq_ctx is not None: - self._client_zmq_ctx.term() + self._flush_pending_pushes() + self._io_executor.shutdown(wait=True, cancel_futures=True) + self._control_executor.shutdown(wait=True, cancel_futures=True) + # Every thread that could hold a control socket is stopped by now. + self._control_channel.close() + if self._control_server is not None: + self._control_server.shutdown() + if self._event_zmq_socket is not None: + self._event_zmq_socket.close(linger=0) + if self._event_zmq_ctx is not None: + self._event_zmq_ctx.term() if self._engine is not None: if self._consumer_pool is not None and self._unregister_memory( @@ -824,12 +2490,13 @@ def shutdown(self) -> None: ): self._consumer_pool = None self._consumer_pool_allocator = None - self._consumer_allocations.clear() + self._consumer_residents.clear() + self._consumer_retire_events.clear() self._consumer_pending_frees.clear() - with self._tensor_lock: - addresses = [ - entry.tensor.data_ptr() for entry in self._tensor_by_hash.values() - ] + # Published tensors and in-flight push sources share one refcounted + # registration table, so a single pass covers both. + with self._push_source_registration_lock: + addresses = list(self._push_source_registrations) addresses.extend(self._pending_unregister) unregistered = True if addresses: @@ -842,9 +2509,8 @@ def shutdown(self) -> None: "Mooncake EC batch memory unregistration failed: %d", ret ) if unregistered: - self._tensor_by_hash.clear() + self._push_source_registrations.clear() self._pending_unregister.clear() - self._registered_bytes = 0 def __del__(self) -> None: with suppress(Exception): diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index e4a21328660a..827ef460a48b 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -555,6 +555,18 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: req_index += 1 continue + if ( + self.ec_connector is not None + and request.mm_features + and not self.ec_connector.ensure_cache_available( + request, + request.num_computed_tokens - request.num_output_placeholders, + self.encoder_cache_manager.cached.keys(), + ) + ): + req_index += 1 + continue + num_new_tokens = ( request.num_tokens_with_spec + request.num_output_placeholders @@ -888,7 +900,9 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: self.ec_connector is not None and request.mm_features and not self.ec_connector.ensure_cache_available( - request, num_computed_tokens + request, + num_computed_tokens, + self.encoder_cache_manager.cached.keys(), ) ): request_queue.pop_request() @@ -2039,6 +2053,10 @@ def update_from_output( self.grammar_compile_error_reqs.clear() if failed_kv_load_req_ids and not self.recompute_kv_load_failures: error_req_ids.update(failed_kv_load_req_ids) + if self.ec_connector is not None: + # An encoder input the connector can no longer obtain. Failing is + # retryable: re-issuing the request re-runs the encode. + error_req_ids.update(self.ec_connector.take_unavailable_requests()) if error_req_ids: error_reqs = self.finish_requests( @@ -2221,7 +2239,7 @@ def _free_encoder_inputs(self, request: Request) -> None: # With Whisper, as soon as we've generated a single token, # we know we're done with the encoder input. Cross Attention # KVs have been calculated and cached already. - self.encoder_cache_manager.free_encoder_input(request, input_id) + self._free_encoder_input(request, input_id) elif ( start_pos + num_tokens + spec_lookahead <= request.num_computed_tokens - request.num_output_placeholders @@ -2229,7 +2247,12 @@ def _free_encoder_inputs(self, request: Request) -> None: # Processed, stored in the decoder KV cache, and far enough past # the placeholder range (plus the drafter's look-ahead) that no # rejection or drafter gather can reference it. - self.encoder_cache_manager.free_encoder_input(request, input_id) + self._free_encoder_input(request, input_id) + + def _free_encoder_input(self, request: Request, input_id: int) -> None: + self.encoder_cache_manager.free_encoder_input(request, input_id) + if self.ec_connector is not None: + self.ec_connector.update_state_after_free(request, input_id) def update_draft_token_ids(self, draft_token_ids: DraftTokenIds) -> None: for req_id, spec_token_ids in zip( diff --git a/vllm/v1/worker/ec_connector_model_runner_mixin.py b/vllm/v1/worker/ec_connector_model_runner_mixin.py index b3430a8d94da..be676666e84d 100644 --- a/vllm/v1/worker/ec_connector_model_runner_mixin.py +++ b/vllm/v1/worker/ec_connector_model_runner_mixin.py @@ -64,6 +64,12 @@ def _get_ec_connector_output( assert scheduler_output.ec_connector_metadata is not None ec_connector.bind_connector_metadata(scheduler_output.ec_connector_metadata) + if ec_connector.is_producer: + # Pass the cache explicitly: a producer needs it to publish an item + # that is already cached from an earlier step, which this step will + # not re-encode -- `save_caches` never fires for those. + ec_connector.start_save_caches(encoder_cache=encoder_cache, **kwargs) + # Load caches for consumer or both roles if ec_connector.is_consumer: ec_connector.start_load_caches(encoder_cache, **kwargs) @@ -74,5 +80,6 @@ def _get_ec_connector_output( output.finished_sending, output.finished_recving = ( ec_connector.get_finished(scheduler_output.finished_req_ids) ) + output.ec_connector_worker_meta = ec_connector.build_connector_worker_meta() ec_connector.clear_connector_metadata() diff --git a/vllm/v1/worker/gpu/ec_connector.py b/vllm/v1/worker/gpu/ec_connector.py index 5dc8d92359a9..1285d9fd14db 100644 --- a/vllm/v1/worker/gpu/ec_connector.py +++ b/vllm/v1/worker/gpu/ec_connector.py @@ -62,6 +62,9 @@ def maybe_get_output( assert scheduler_output.ec_connector_metadata is not None ec_connector.bind_connector_metadata(scheduler_output.ec_connector_metadata) + if ec_connector.is_producer: + ec_connector.start_save_caches(encoder_cache=self.encoder_cache) + if ec_connector.is_consumer: ec_connector.start_load_caches(self.encoder_cache) diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 9087f878546f..1b756d8b3c30 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -4333,7 +4333,10 @@ def execute_model( encoder_cache=self.encoder_cache, ) as ec_connector_output: self._execute_mm_encoder(scheduler_output) - return make_empty_encoder_model_runner_output(scheduler_output) + return ModelRunnerOutput.with_ec_conn_output( + make_empty_encoder_model_runner_output(scheduler_output), + ec_connector_output, + ) if not num_scheduled_tokens: if ( @@ -4348,10 +4351,22 @@ def execute_model( # dummy run to ensure coordinate_batch_across_dp # is called into to avoid out of sync issues. self._dummy_run(1) - if not has_kv_transfer_group(): - # Return empty ModelRunnerOutput if no work to do. - return EMPTY_MODEL_RUNNER_OUTPUT - return self.kv_connector_no_forward(scheduler_output, self.vllm_config) + if has_kv_transfer_group(): + empty_output = self.kv_connector_no_forward( + scheduler_output, self.vllm_config + ) + else: + empty_output = EMPTY_MODEL_RUNNER_OUTPUT + if has_ec_transfer(): + with self.maybe_get_ec_connector_output( + scheduler_output, + encoder_cache=self.encoder_cache, + ) as ec_connector_output: + pass + empty_output = ModelRunnerOutput.with_ec_conn_output( + empty_output, ec_connector_output + ) + return empty_output if self.cache_config.kv_sharing_fast_prefill: assert not self.num_prompt_logprobs, ( diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index a3b00aaad2a2..0b8967b16fdd 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -28,6 +28,8 @@ from vllm.distributed.ec_transfer import ( ensure_ec_transfer_initialized, ensure_ec_transfer_shutdown, + get_ec_transfer, + has_ec_transfer, ) from vllm.distributed.eplb.eplb_utils import override_envs_for_eplb from vllm.distributed.kv_transfer import ( @@ -456,6 +458,9 @@ def load_model(self, *, load_dummy_weights: bool = False) -> None: ): self.model_runner.load_model(load_dummy_weights=load_dummy_weights) + if has_ec_transfer(): + get_ec_transfer().start_worker_services() + if self.vllm_config.weight_transfer_config is not None: self.weight_transfer_engine = WeightTransferEngineFactory.create_engine( self.vllm_config.weight_transfer_config, From 10db3ae61ad0229715d743c9740fe51194b4aa9d Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Tue, 18 Aug 2026 07:25:54 +0000 Subject: [PATCH 10/30] Support consumer TP and PP Signed-off-by: Tianyu Guo --- .../run_epd_mooncake_ec_full_pipeline.sh | 5 + .../unit/test_ec_mooncake_connector.py | 116 ++++- .../ec_connector/mooncake_ec_connector.py | 434 +++++++++++++----- vllm/v1/worker/gpu_model_runner.py | 25 +- 4 files changed, 441 insertions(+), 139 deletions(-) diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh index a30e942422f6..04e28c70483d 100755 --- a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -45,6 +45,11 @@ EC_MOONCAKE_RESERVATION_PORT="${EC_MOONCAKE_RESERVATION_PORT:-19019}" MOONCAKE_EC_PROTOCOL="${MOONCAKE_EC_PROTOCOL:-rdma}" export EC_MOONCAKE_RESERVATION_PORT export MOONCAKE_EC_PROTOCOL +# Mooncake cannot register CUDA memory through the peer-memory path on hosts +# whose kernel lacks the OFED peer-memory API; every transfer then fails at +# setup with -202. Opting out selects the path that works there and is a no-op +# where GPUDirect is available. +export WITH_NVIDIA_PEERMEM="${WITH_NVIDIA_PEERMEM:-0}" LOG_PATH="${LOG_PATH:-/tmp}" BASELINE_FILE="${BASELINE_FILE:-/tmp/vllm_epd_mooncake_baseline.txt}" diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 250c268a769b..65f537f6c9ba 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -116,6 +116,7 @@ def mock_vllm_config_producer(): config.parallel_config = Mock() config.parallel_config.tensor_parallel_size = 1 config.parallel_config.pipeline_parallel_size = 1 + config.parallel_config.data_parallel_size = 1 config.ec_transfer_config = Mock() config.ec_transfer_config.is_ec_producer = True config.ec_transfer_config.is_ec_consumer = False @@ -133,6 +134,7 @@ def mock_vllm_config_consumer(): config.parallel_config = Mock() config.parallel_config.tensor_parallel_size = 1 config.parallel_config.pipeline_parallel_size = 1 + config.parallel_config.data_parallel_size = 1 config.ec_transfer_config = Mock() config.ec_transfer_config.is_ec_producer = False config.ec_transfer_config.is_ec_consumer = True @@ -188,7 +190,8 @@ def test_reuses_and_coalesces_contiguous_regions(self): class TestECMooncakeConnectorValidation: - def test_rejects_tensor_parallel_gt_one(self, mock_vllm_config_producer): + def test_rejects_sharded_producer(self, mock_vllm_config_producer): + """One copy of each encoder output, so sharding only duplicates it.""" mock_vllm_config_producer.parallel_config.tensor_parallel_size = 2 with ( patch_ec_mooncake_deps(), @@ -196,6 +199,48 @@ def test_rejects_tensor_parallel_gt_one(self, mock_vllm_config_producer): ): ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.WORKER) + def test_accepts_sharded_consumer(self, mock_vllm_config_consumer): + """Consumers shard: each rank gathers from its own encoder cache.""" + mock_vllm_config_consumer.parallel_config.tensor_parallel_size = 4 + mock_vllm_config_consumer.parallel_config.pipeline_parallel_size = 2 + with patch_ec_mooncake_deps(): + connector = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + connector.shutdown() + + def test_rejects_data_parallel(self, mock_vllm_config_consumer): + """One control channel per instance cannot address a replica.""" + mock_vllm_config_consumer.parallel_config.data_parallel_size = 2 + with ( + patch_ec_mooncake_deps(), + pytest.raises(ValueError, match="data parallelism"), + ): + ECMooncakeConnector(mock_vllm_config_consumer, ECConnectorRole.SCHEDULER) + + +class TestECMooncakeWorkerMetadataAggregation: + def test_an_item_one_rank_missed_is_not_loaded(self): + """Each rank gathers from its own cache, so all of them must have it. + + Reporting it as loaded because one rank succeeded left the scheduler + marking the hash ready while another rank raised on the cache miss. + """ + rank0 = ECMooncakeWorkerMetadata(loaded={"a", "b"}) + rank1 = ECMooncakeWorkerMetadata(loaded={"a"}, failed_loads={"b"}) + + merged = rank0.aggregate(rank1) + + assert merged.loaded == {"a"} + assert merged.failed_loads == {"b"} + + def test_a_reclaim_on_any_rank_invalidates_residency(self): + """The scheduler mirrors one pool, so the weakest rank decides.""" + merged = ECMooncakeWorkerMetadata(loaded={"a"}).aggregate( + ECMooncakeWorkerMetadata(loaded={"a"}, reclaimed={"c"}) + ) + assert merged.reclaimed == {"c"} + class TestECMooncakeSchedulerMetadata: def test_missing_push_event_is_tracked( @@ -360,7 +405,9 @@ def test_consumed_item_releases_its_transfer_immediately( ) with patch.object(scheduler, "_queue_cancel") as cancel: scheduler.update_state_after_free(request, 0) - cancel.assert_called_once_with("consumed-transfer", "reservation") + # Cancelled by transfer: a shard's reservation id means + # nothing to its peers, so it is not passed along. + cancel.assert_called_once_with("consumed-transfer") assert "consumed-transfer" not in scheduler._pending_specs finally: scheduler.shutdown() @@ -813,7 +860,10 @@ def test_push_reserves_before_encoder_output_is_saved( try: producer.start_save_caches(encoder_cache={}) _, reservation = producer._pending_reservations["hash"][0] - reservation_data = reservation.result(timeout=2) + shards = reservation.result(timeout=2) + # One reservation per consumer shard; this consumer is single. + assert len(shards) == 1 + reservation_data = shards[0] assert reservation_data["nbytes"] == source.nbytes old_reservation_id = reservation_data["reservation_id"] reservation_data["_received_at"] -= _LEASE_TTL_SECONDS @@ -1072,6 +1122,66 @@ def test_retired_item_reserved_again_still_serves_a_local_load( finally: consumer.shutdown() + def test_push_reaches_every_consumer_shard(self, mock_vllm_config_producer): + """A sharded consumer gets one copy per rank, from one source. + + Each rank gathers from its own encoder cache, so the push has to land + on all of them; the bytes are identical, so staging and registration + happen once however many destinations there are. + """ + shard_ports = [_find_free_port() for _ in range(3)] + base = f"tcp://127.0.0.1:{shard_ports[0]}" + source = torch.randn(4, 16) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq=base, + transfer_id="transfer-0", + ) + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + destinations = [torch.zeros_like(source) for _ in shard_ports] + + def fake_send(addr: str, request: dict): + if request["op"] == "peers": + return {"ports": shard_ports} + index = shard_ports.index(int(addr.rsplit(":", 1)[1])) + if request["op"] == "reserve": + return { + "reservation_id": f"r{index}", + "dst_session": f"session-{index}", + "dst_ptr": destinations[index].data_ptr(), + "nbytes": source.nbytes, + "write": True, + "ready": True, + "addr": addr, + } + if request["op"] == "complete_batch": + return {"items": [{"completed": True} for _ in request["items"]]} + return {} + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) + try: + with patch.object(producer, "_send_control", side_effect=fake_send): + producer.start_save_caches(encoder_cache={"hash": source}) + _wait_for_worker_io(producer) + + engine = producer._engine + assert isinstance(engine, CopyingFakeTransferEngine) + # One write per rank, and every rank got the same bytes. + assert len(engine.transfer_calls) == len(shard_ports) + for destination in destinations: + assert torch.equal(destination, source) + # The source is staged once, not once per destination. + assert len(engine.register_calls) == 1 + finally: + producer.shutdown() + def test_pushes_stage_through_the_registered_pool(self, mock_vllm_config_producer): """Repeated content must not register overlapping source storage.""" port = _find_free_port() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 187c9ca9f623..146d038ad769 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -115,7 +115,11 @@ class ECMooncakeWorkerMetadata(ECConnectorWorkerMetadata): def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadata: assert isinstance(other, ECMooncakeWorkerMetadata) return ECMooncakeWorkerMetadata( - loaded=self.loaded | other.loaded, + # Every tensor-parallel rank gathers the embedding from its own + # cache, so an item counts as loaded only where all of them have + # it; one rank falling short must fail the load rather than leave + # the scheduler believing it is ready. + loaded=self.loaded & other.loaded, failed_loads=self.failed_loads | other.failed_loads, reclaimed=self.reclaimed | other.reclaimed, pending_loads=self.pending_loads or other.pending_loads, @@ -161,7 +165,7 @@ class _PushCompletion: class _PendingPush: tensor: torch.Tensor spec: ECMooncakePushSpec - reservation: Future[dict[str, Any]] + reservation: Future[list[dict[str, Any]]] ready_event: torch.Event | None enqueued_at: float @@ -409,9 +413,13 @@ def __init__( cancel: Callable[[str, str, bool], bool], reap: Callable[[], int], metrics_log_interval: float = 10, + peer_ports: list[int] | None = None, + device: torch.device | None = None, ): self.host = host self.port = port + self.peer_ports = peer_ports or [port] + self._device = device self.event_port: int | None = None self._reserve = reserve self._status = status @@ -426,6 +434,15 @@ def __init__( def start(self) -> None: def loop() -> None: + if self._device is not None and self._device.type == "cuda": + # Reserving can retire an entry, and the event that orders its + # reuse is created on the recording thread's device rather + # than the stream's. A thread starts on device 0, which under + # a shard-local CUDA_VISIBLE_DEVICES is a peer's GPU, so + # without this every shard but the first strands a primary + # context there. The event orders correctly either way; what + # it costs is a few hundred MiB on someone else's card. + torch.accelerator.set_device_index(self._device.index or 0) context = zmq.Context() socket = context.socket(zmq.REP) event_socket = context.socket(zmq.PUSH) @@ -502,6 +519,10 @@ def loop() -> None: result = self._status(str(request["transfer_id"])) elif op == "event_port": result = self.event_port + elif op == "peers": + # Every consumer shard receives its own copy, so a + # producer holding one address needs the rest. + result = {"ports": self.peer_ports} elif op in ("complete", "complete_batch"): items = ( request["items"] @@ -594,7 +615,9 @@ class ECMooncakeConnector(ECConnectorBase): - ``consumer_buffer_pool_size`` (consumer, optional): Bytes reserved for a long-lived registered CUDA receive arena (default ``ec_buffer_size``). - ``reservation_zmq_port`` (consumer worker, required): Exposes registered - receive addresses over ZMQ on this port so producers can push into them. + receive addresses over ZMQ. Tensor-parallel rank ``r`` of the first + pipeline stage listens on ``port + r``, and rank 0 reports the whole set, + so a producer only needs the first address. - ``reservation_zmq_addr`` (consumer scheduler, required): Address of the consumer control channel. Defaults to ``tcp://127.0.0.1:``. - ``transfer_max_workers`` (optional): Maximum concurrent Mooncake transfer @@ -606,8 +629,13 @@ class ECMooncakeConnector(ECConnectorBase): - ``consumer_metrics_log_interval`` (optional): Seconds between aggregated consumer lifecycle logs (default ``10``; ``0`` disables them). - Limitations: ``tensor_parallel_size`` and ``pipeline_parallel_size`` must - be ``1`` (same assumption as Mooncake KV connector for P2P handshake). + Parallelism: consumers may use tensor and pipeline parallelism. Only the + first pipeline stage holds encoder outputs, and each tensor-parallel rank + there gathers from its own cache, so every rank exposes a control channel + and the producer writes into all of them concurrently from one registered + source. That costs bandwidth but not latency, and avoids the second hop a + receive-then-broadcast would add. Producers must be unsharded, and data + parallelism is unsupported on either side. """ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): @@ -618,12 +646,26 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): "https://github.com/kvcache-ai/Mooncake ) to use ECMooncakeConnector." ) from _MOONCAKE_IMPORT_ERROR - if vllm_config.parallel_config.tensor_parallel_size > 1: - raise ValueError("ECMooncakeConnector requires tensor_parallel_size=1.") - if vllm_config.parallel_config.pipeline_parallel_size > 1: + parallel_config = vllm_config.parallel_config + if parallel_config.data_parallel_size > 1: raise ValueError( - "ECMooncakeConnector does not support pipeline parallelism yet." + "ECMooncakeConnector does not support data parallelism yet: the " + "consumer exposes one control channel per instance, so a push " + "cannot be routed to the replica that will run the request." ) + ec_cfg_early = vllm_config.ec_transfer_config + assert ec_cfg_early is not None + if ec_cfg_early.is_ec_producer: + # The producer holds one copy of each encoder output and addresses + # consumers directly; sharding it would only duplicate the push. + if parallel_config.tensor_parallel_size > 1: + raise ValueError( + "ECMooncakeConnector producers require tensor_parallel_size=1." + ) + if parallel_config.pipeline_parallel_size > 1: + raise ValueError( + "ECMooncakeConnector producers do not support pipeline parallelism." + ) self._role = role ec_cfg = vllm_config.ec_transfer_config @@ -664,6 +706,10 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): tuple[torch.Event, _ConsumerPoolAllocation] ] = [] self._consumer_reclaimed: set[str] = set() + self._consumer_rank_resolved = False + self._is_receiving_rank = True + self._tp_rank = 0 + self._tp_size = 1 self._consumer_pool_disabled = self._consumer_pool_capacity <= 0 self._consumer_lock = threading.Lock() self._push_reservations: dict[str, _PushReservation] = {} @@ -737,9 +783,12 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._control_executor = ThreadPoolExecutor( max_workers=control_workers, thread_name_prefix="ec-mooncake-control" ) + self._consumer_shard_cache: dict[str, list[str]] = {} + self._shard_pool: ThreadPoolExecutor | None = None + self._shard_pool_lock = threading.Lock() self._pending_saves: list[tuple[str, Future[None]]] = [] self._pending_reservations: dict[ - str, deque[tuple[ECMooncakePushSpec, Future[dict[str, Any]]]] + str, deque[tuple[ECMooncakePushSpec, Future[list[dict[str, Any]]]]] ] = {} self._pending_pushes: list[_PendingPush] = [] self._push_perf_lock = threading.Lock() @@ -780,6 +829,32 @@ def _ensure_engine(self) -> TransferEngine: ) return self._engine + def _resolve_consumer_rank(self) -> None: + """Place this worker in the consumer's receive topology. + + Encoder outputs only exist on the first pipeline stage, and every + tensor-parallel rank there gathers from its own cache, so each of them + receives its own copy on its own control channel. Ports run + consecutively from the configured one so a producer holding the first + address can reach the rest. + """ + if self._consumer_rank_resolved: + return + self._consumer_rank_resolved = True + try: + from vllm.distributed.parallel_state import get_pp_group, get_tp_group + + tp_group = get_tp_group() + self._tp_rank = tp_group.rank_in_group + self._tp_size = tp_group.world_size + self._is_receiving_rank = get_pp_group().is_first_rank + except AssertionError: + # Groups are only absent outside a distributed run, where this + # worker is the whole consumer. + self._tp_rank = 0 + self._tp_size = 1 + self._is_receiving_rank = True + def start_worker_services(self) -> None: if ( self._role != ECConnectorRole.WORKER @@ -788,6 +863,11 @@ def start_worker_services(self) -> None: or self._control_server is not None ): return + self._resolve_consumer_rank() + if not self._is_receiving_rank: + # Later pipeline stages hold no encoder outputs, so they need + # neither a receive pool nor a control channel. + return raw_device = self._ec_cfg.ec_buffer_device device_name = ( raw_device.lower() if isinstance(raw_device, str) and raw_device else "cuda" @@ -799,13 +879,17 @@ def start_worker_services(self) -> None: ) self._control_server = ECMooncakeControlServer( "0.0.0.0", - self._reservation_zmq_port, + self._reservation_zmq_port + self._tp_rank, self._reserve_push_destination, self._push_status, self._complete_push, self._cancel_push, self._expire_push_reservations, self._consumer_metrics_log_interval, + peer_ports=[ + self._reservation_zmq_port + rank for rank in range(self._tp_size) + ], + device=self._consumer_pool.device, ) self._control_server.start() @@ -930,11 +1014,14 @@ def _ensure_consumer_pool( pool = torch.empty( self._consumer_pool_capacity, dtype=torch.uint8, device=device ) - ret = self._ensure_engine().batch_register_memory( - [pool.data_ptr()], [pool.nbytes] - ) - if ret != 0: - raise RuntimeError(f"Mooncake returned {ret}") + if self._is_receiving_rank: + # Producers write into this pool directly, so it needs a memory + # region. Later pipeline stages never receive and skip it. + ret = self._ensure_engine().batch_register_memory( + [pool.data_ptr()], [pool.nbytes] + ) + if ret != 0: + raise RuntimeError(f"Mooncake returned {ret}") except (RuntimeError, torch.OutOfMemoryError) as e: self._consumer_pool_disabled = True logger.warning( @@ -946,8 +1033,9 @@ def _ensure_consumer_pool( self._consumer_pool = pool self._consumer_pool_allocator = _ContiguousAllocator(pool.nbytes) logger.info( - "Registered %d-byte CUDA receive pool for Mooncake EC", + "Prepared %d-byte CUDA receive pool for Mooncake EC (registered=%s)", pool.nbytes, + self._is_receiving_rank, ) def _ensure_producer_pool(self, device: torch.device) -> None: @@ -1483,11 +1571,12 @@ def _take_pushed_tensor( ) -> tuple[torch.Tensor, _ConsumerPoolAllocation]: with self._consumer_lock: reservation = self._push_reservations.get(spec.transfer_id) - if ( - reservation is None - or not reservation.ready - or reservation.reservation_id != spec.reservation_id - ): + # Not compared against `spec.reservation_id`: each shard mints its + # own, while the spec carries the one from whichever shard's event + # the scheduler observed. `transfer_id` is assigned per request + # item and is already unique, and a stale reservation for a reused + # one is rejected by `_reserve_push_destination`. + if reservation is None or not reservation.ready: self._consumer_worker_metrics["takes_rejected"] += 1 raise RuntimeError( f"Pushed EC tensor is not ready for mm_hash={spec.mm_hash}" @@ -1502,9 +1591,58 @@ def _take_pushed_tensor( def _send_control(self, addr: str, request: dict[str, Any]) -> Any: return self._control_channel.request(addr, request) - def _reserve_remote(self, spec: ECMooncakePushSpec) -> dict[str, Any]: + def _shard_executor(self) -> ThreadPoolExecutor: + """Threads for the extra shards of a sharded consumer. + + Reserving and writing both fan out from a task that already holds a + worker of the control or transfer pool, so the extra shards need a + pool of their own: queueing them behind their own caller deadlocks as + soon as every worker there is waiting. Nothing submitted here fans out + again, so this pool cannot deadlock on itself. + """ + with self._shard_pool_lock: + if self._shard_pool is None: + self._shard_pool = ThreadPoolExecutor( + max_workers=32, thread_name_prefix="ec-mooncake-shard" + ) + return self._shard_pool + + def _consumer_shards(self, base_addr: str) -> list[str]: + """Every control channel of the consumer reachable at `base_addr`. + + A tensor-parallel consumer gathers from each rank's own cache, so each + rank receives its own copy. Asking the first one for the roster keeps + the address list out of the request and the proxy configuration. + """ + cached = self._consumer_shard_cache.get(base_addr) + if cached is not None: + return cached + shards = [base_addr] + try: + reply = self._send_control(base_addr, {"op": "peers"}) + ports = reply.get("ports") if isinstance(reply, dict) else None + if ports: + prefix = base_addr.rsplit(":", 1)[0] + shards = [f"{prefix}:{int(port)}" for port in ports] + except Exception: + # An older consumer does not answer this, and it can only be + # unsharded, so its single address is the whole roster. + logger.warning( + "EC Mooncake consumer at %s did not report its shards; " + "assuming it is unsharded.", + base_addr, + exc_info=True, + ) + self._consumer_shard_cache[base_addr] = shards + if len(shards) > 1: + logger.info( + "EC Mooncake consumer at %s has %d shards", base_addr, len(shards) + ) + return shards + + def _reserve_one(self, addr: str, spec: ECMooncakePushSpec) -> dict[str, Any]: result = self._send_control( - spec.consumer_zmq, + addr, { "op": "reserve", "transfer_id": spec.transfer_id, @@ -1517,20 +1655,43 @@ def _reserve_remote(self, spec: ECMooncakePushSpec) -> dict[str, Any]: if not isinstance(result, dict): raise RuntimeError("Invalid EC reservation response") result["_received_at"] = time.monotonic() + result["addr"] = addr return result + def _reserve_remote(self, spec: ECMooncakePushSpec) -> list[dict[str, Any]]: + """Reserve a destination on every shard of the consumer.""" + shards = self._consumer_shards(spec.consumer_zmq) + if len(shards) == 1: + return [self._reserve_one(shards[0], spec)] + # This already runs on the control pool, so the extra shards go to the + # fan-out pool: queueing them behind their own caller would deadlock + # once every control worker is holding a reservation. + extra = [ + self._shard_executor().submit(self._reserve_one, addr, spec) + for addr in shards[1:] + ] + return [self._reserve_one(shards[0], spec)] + [f.result() for f in extra] + def _cancel_remote( self, consumer_zmq: str, transfer_id: str, reservation_id: str ) -> bool: - result = self._send_control( - consumer_zmq, - { - "op": "cancel", - "transfer_id": transfer_id, - "reservation_id": reservation_id, - }, - ) - return isinstance(result, dict) and bool(result.get("cancelled")) + """Release this transfer on every shard that reserved for it. + + A sharded consumer holds one reservation per rank, so cancelling only + the first would leave the rest pinning pool slots until they expire. + """ + cancelled = False + for addr in self._consumer_shards(consumer_zmq): + result = self._send_control( + addr, + { + "op": "cancel", + "transfer_id": transfer_id, + "reservation_id": reservation_id, + }, + ) + cancelled |= isinstance(result, dict) and bool(result.get("cancelled")) + return cancelled def _poll_pending_cancels(self) -> None: pending = {} @@ -1570,6 +1731,12 @@ def start_save_caches(self, **kwargs: Any) -> None: def start_load_caches( self, encoder_cache: dict[str, torch.Tensor], **kwargs: Any ) -> None: + self._resolve_consumer_rank() + if not self._is_receiving_rank: + # Reached on steps with no work, from a stage that never gathers + # multimodal embeddings. Taking a transfer here would fail for + # want of a reservation and fail the load for everyone. + return metadata = self._get_connector_metadata() assert isinstance(metadata, ECMooncakeConnectorMetadata) self._ensure_engine() @@ -1584,7 +1751,8 @@ def start_load_caches( for spec in metadata.loads: if spec.mm_hash in encoder_cache: if spec.pushed: - self._cancel_push(spec.transfer_id, spec.reservation_id) + # The spec's id is one shard's; cancel by transfer. + self._cancel_push(spec.transfer_id, "") self._completed_loads.add(spec.mm_hash) continue if spec.local: @@ -1635,39 +1803,52 @@ def _push_batch(self, pushes: list[_PendingPush]) -> None: notifications: list[tuple[_PendingPush, dict[str, Any]]] = [] failed = False try: + synchronized: set[int] = set() for push in pushes: stage_started_at = time.monotonic() - reservation = push.reservation.result() - received_at = float(reservation.get("_received_at", started_at)) - if ( - not reservation.get("ready", False) - and time.monotonic() - received_at >= _RESERVATION_REFRESH_SECONDS - ): - reservation = self._reserve_remote(push.spec) + reservations = push.reservation.result() + stale = [ + index + for index, shard in enumerate(reservations) + if not shard.get("ready", False) + and time.monotonic() - float(shard.get("_received_at", started_at)) + >= _RESERVATION_REFRESH_SECONDS + ] + if stale: + reservations = self._reserve_remote(push.spec) stage_ms["reserve"] += (time.monotonic() - stage_started_at) * 1000 - if reservation.get("cached", False) or reservation.get( - "cancelled", False - ): - continue - if not reservation.get("write", True): - continue - if push.ready_event is not None: - stage_started_at = time.monotonic() - push.ready_event.synchronize() - stage_ms["cuda"] += (time.monotonic() - stage_started_at) * 1000 - if int(reservation["nbytes"]) != push.tensor.nbytes: - raise RuntimeError( - "Reserved EC size does not match tensor for " - f"mm_hash={push.spec.mm_hash}" - ) - ready.append((push, reservation)) - notifications.append((push, reservation)) + for shard in reservations: + if shard.get("cached", False) or shard.get("cancelled", False): + continue + if not shard.get("write", True): + continue + if push.ready_event is not None and id(push) not in synchronized: + stage_started_at = time.monotonic() + push.ready_event.synchronize() + stage_ms["cuda"] += (time.monotonic() - stage_started_at) * 1000 + synchronized.add(id(push)) + if int(shard["nbytes"]) != push.tensor.nbytes: + raise RuntimeError( + "Reserved EC size does not match tensor for " + f"mm_hash={push.spec.mm_hash}" + ) + ready.append((push, shard)) + notifications.append((push, shard)) if not ready and not notifications: return if ready: eng = self._ensure_engine() - tensors = [push.tensor for push, _ in ready] + # One source per push: a sharded consumer reads the same bytes + # into each of its ranks, so staging and registration happen + # once however many destinations there are. + unique: list[_PendingPush] = [] + source_index: dict[int, int] = {} + for push, _ in ready: + if id(push) not in source_index: + source_index[id(push)] = len(unique) + unique.append(push) + tensors = [push.tensor for push in unique] lengths = [tensor.nbytes for tensor in tensors] stage_started_at = time.monotonic() staged = self._stage_push_sources(tensors) @@ -1689,23 +1870,39 @@ def _push_batch(self, pushes: list[_PendingPush]) -> None: addresses = [tensor.data_ptr() for tensor in sources] stage_ms["register"] = (time.monotonic() - stage_started_at) * 1000 try: - by_session: dict[str, list[int]] = {} - for index, (_, reservation) in enumerate(ready): - by_session.setdefault( - str(reservation["dst_session"]), [] - ).append(index) + by_session: dict[str, list[tuple[int, int]]] = {} + for push, shard in ready: + by_session.setdefault(str(shard["dst_session"]), []).append( + (source_index[id(push)], int(shard["dst_ptr"])) + ) stage_started_at = time.monotonic() - for session, indices in by_session.items(): + + def write(session: str, items: list[tuple[int, int]]) -> None: ret = eng.batch_transfer_sync_write( session, - [addresses[index] for index in indices], - [int(ready[index][1]["dst_ptr"]) for index in indices], - [lengths[index] for index in indices], + [addresses[index] for index, _ in items], + [dst for _, dst in items], + [lengths[index] for index, _ in items], ) if ret != 0: raise RuntimeError( - f"Mooncake EC push failed with status {ret}" + f"Mooncake EC push to {session} failed with " + f"status {ret}" ) + + sessions = list(by_session.items()) + # Shards are written concurrently: serialising them would + # make the transfer cost the sum of the ranks instead of + # the slowest one. + extra = [ + self._shard_executor().submit(write, session, items) + for session, items in sessions[1:] + ] + try: + write(*sessions[0]) + finally: + for future in extra: + future.result() stage_ms["rdma"] = (time.monotonic() - stage_started_at) * 1000 finally: stage_started_at = time.monotonic() @@ -1742,8 +1939,12 @@ def _push_batch(self, pushes: list[_PendingPush]) -> None: stage_ms, stage_max_ms={"queue": max(queue_waits_ms, default=0.0)}, item_count=len(pushes), - byte_count=sum(push.tensor.nbytes for push, _ in ready), - skipped_items=len(pushes) - len(ready), + # `ready` holds one entry per destination shard, so count the + # distinct items rather than the writes. + byte_count=sum( + push.tensor.nbytes for push in {id(p): p for p, _ in ready}.values() + ), + skipped_items=len(pushes) - len({id(push) for push, _ in ready}), failed=failed, ) @@ -1755,9 +1956,9 @@ def _notify_completions( return by_destination: dict[str, list[tuple[_PendingPush, dict[str, Any]]]] = {} for push, reservation in notifications: - by_destination.setdefault(push.spec.consumer_zmq, []).append( - (push, reservation) - ) + by_destination.setdefault( + str(reservation.get("addr", push.spec.consumer_zmq)), [] + ).append((push, reservation)) for consumer_zmq, items in by_destination.items(): result = self._send_control( consumer_zmq, @@ -1784,21 +1985,23 @@ def _notify_completions( def _abandon_pushes(self, pushes: list[_PendingPush]) -> None: """Release the consumer-side reservations of a batch that failed.""" for push in pushes: - reservation_id = "" + shards: list[dict[str, Any]] = [] if push.reservation.done() and not push.reservation.cancelled(): with suppress(Exception): - result = push.reservation.result() - reservation_id = str(result.get("reservation_id", "")) - with suppress(Exception): - self._send_control( - push.spec.consumer_zmq, - { - "op": "cancel", - "transfer_id": push.spec.transfer_id, - "reservation_id": reservation_id, - "abandon": True, - }, - ) + shards = push.reservation.result() + if not shards: + shards = [{"addr": push.spec.consumer_zmq, "reservation_id": ""}] + for shard in shards: + with suppress(Exception): + self._send_control( + str(shard.get("addr", push.spec.consumer_zmq)), + { + "op": "cancel", + "transfer_id": push.spec.transfer_id, + "reservation_id": str(shard.get("reservation_id", "")), + "abandon": True, + }, + ) def _record_push_perf( self, @@ -1895,7 +2098,7 @@ def _submit_push( self, tensor: torch.Tensor, spec: ECMooncakePushSpec, - reservation: Future[dict[str, Any]], + reservation: Future[list[dict[str, Any]]], ) -> None: ready_event = None if tensor.device.type == "cuda": @@ -1922,21 +2125,21 @@ def _submit_reserved_pushes(self, tensor: torch.Tensor, mm_hash: str) -> None: def _cancel_orphaned_reservation( self, spec: ECMooncakePushSpec, - reservation: Future[dict[str, Any]], + reservation: Future[list[dict[str, Any]]], ) -> None: try: - result = reservation.result() - if result.get("cached", False) or result.get("cancelled", False): - return - self._send_control( - spec.consumer_zmq, - { - "op": "cancel", - "transfer_id": spec.transfer_id, - "reservation_id": str(result.get("reservation_id", "")), - "abandon": True, - }, - ) + for shard in reservation.result(): + if shard.get("cached", False) or shard.get("cancelled", False): + continue + self._send_control( + str(shard.get("addr", spec.consumer_zmq)), + { + "op": "cancel", + "transfer_id": spec.transfer_id, + "reservation_id": str(shard.get("reservation_id", "")), + "abandon": True, + }, + ) except Exception: logger.exception( "Failed to cancel orphaned EC reservation for transfer_id=%s", @@ -1949,11 +2152,10 @@ def get_finished( if not self.is_producer or self._role != ECConnectorRole.WORKER: return None, None - orphaned: list[tuple[ECMooncakePushSpec, Future[dict[str, Any]]]] = [] + Reserved = tuple[ECMooncakePushSpec, Future[list[dict[str, Any]]]] + orphaned: list[Reserved] = [] for mm_hash, reservations in list(self._pending_reservations.items()): - remaining: deque[tuple[ECMooncakePushSpec, Future[dict[str, Any]]]] = ( - deque() - ) + remaining: deque[Reserved] = deque() for spec, reservation in reservations: if spec.request_id in finished_req_ids: orphaned.append((spec, reservation)) @@ -2089,7 +2291,7 @@ def _expire_pending_specs(self) -> None: if spec is not None: self._consumer_pending_since.pop(spec.mm_hash, None) self._consumer_scheduler_metrics["pending_specs_expired"] += 1 - self._queue_cancel(transfer_id, spec.reservation_id) + self._queue_cancel(transfer_id) def _ensure_event_channel(self) -> None: if self._event_zmq_socket is not None: @@ -2305,11 +2507,8 @@ def update_state_after_free(self, request: Any, index: int) -> None: transfer_id = self._request_transfer_id(request, index) if transfer_id is None: return - pending = self._pop_pending_spec(transfer_id) - self._queue_cancel( - transfer_id, - pending.reservation_id if pending is not None else "", - ) + self._pop_pending_spec(transfer_id) + self._queue_cancel(transfer_id) def build_connector_meta( self, scheduler_output: SchedulerOutput @@ -2353,6 +2552,10 @@ def build_connector_meta( def build_connector_worker_meta(self) -> ECConnectorWorkerMetadata | None: if self._role != ECConnectorRole.WORKER: return None + if self.is_consumer and not self._is_receiving_rank: + # `loaded` is intersected across reporting ranks, so a stage that + # never loads must not report at all rather than report nothing. + return None self._flush_pending_pushes() saves = self._pending_saves @@ -2440,11 +2643,8 @@ def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: transfer_id = self._request_transfer_id(request, index) if transfer_id is None: continue - pending = self._pop_pending_spec(transfer_id) - self._queue_cancel( - transfer_id, - pending.reservation_id if pending is not None else "", - ) + self._pop_pending_spec(transfer_id) + self._queue_cancel(transfer_id) if not self.is_producer: return False, None @@ -2474,6 +2674,8 @@ def shutdown(self) -> None: self._shutdown = True self._flush_pending_pushes() self._io_executor.shutdown(wait=True, cancel_futures=True) + if self._shard_pool is not None: + self._shard_pool.shutdown(wait=True, cancel_futures=True) self._control_executor.shutdown(wait=True, cancel_futures=True) # Every thread that could hold a control socket is stopped by now. self._control_channel.close() diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 1b756d8b3c30..9087f878546f 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -4333,10 +4333,7 @@ def execute_model( encoder_cache=self.encoder_cache, ) as ec_connector_output: self._execute_mm_encoder(scheduler_output) - return ModelRunnerOutput.with_ec_conn_output( - make_empty_encoder_model_runner_output(scheduler_output), - ec_connector_output, - ) + return make_empty_encoder_model_runner_output(scheduler_output) if not num_scheduled_tokens: if ( @@ -4351,22 +4348,10 @@ def execute_model( # dummy run to ensure coordinate_batch_across_dp # is called into to avoid out of sync issues. self._dummy_run(1) - if has_kv_transfer_group(): - empty_output = self.kv_connector_no_forward( - scheduler_output, self.vllm_config - ) - else: - empty_output = EMPTY_MODEL_RUNNER_OUTPUT - if has_ec_transfer(): - with self.maybe_get_ec_connector_output( - scheduler_output, - encoder_cache=self.encoder_cache, - ) as ec_connector_output: - pass - empty_output = ModelRunnerOutput.with_ec_conn_output( - empty_output, ec_connector_output - ) - return empty_output + if not has_kv_transfer_group(): + # Return empty ModelRunnerOutput if no work to do. + return EMPTY_MODEL_RUNNER_OUTPUT + return self.kv_connector_no_forward(scheduler_output, self.vllm_config) if self.cache_config.kv_sharing_fast_prefill: assert not self.num_prompt_logprobs, ( From b429844d97bc14adf10c5824a457c79730068609 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Tue, 18 Aug 2026 12:25:34 +0000 Subject: [PATCH 11/30] Fix consumer TP>1 Signed-off-by: Tianyu Guo --- .../unit/test_ec_mooncake_connector.py | 61 +++++++++- .../ec_connector/mooncake_ec_connector.py | 105 +++++++++++++++--- 2 files changed, 145 insertions(+), 21 deletions(-) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 65f537f6c9ba..d2fabc233a66 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -507,6 +507,55 @@ def test_item_that_never_arrives_fails_the_request( finally: scheduler.shutdown() + def test_readiness_needs_every_consumer_shard(self, mock_vllm_config_consumer): + """A sharded consumer is only ready once every rank reports. + + Each rank runs its own control channel and pushes its own readiness + notifications. Subscribing to the first rank alone strands the other + ranks' queues and lets a load be scheduled that the last rank cannot + serve, which only `aggregate`'s `loaded` intersection then catches. + """ + ports = [19101, 19102, 19103] + event_ports = {port: 19201 + index for index, port in enumerate(ports)} + + def fake_send(addr: str, request: dict): + port = int(addr.rsplit(":", 1)[1]) + if request["op"] == "peers": + return {"ports": ports} + if request["op"] == "event_port": + return event_ports[port] + return {} + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._reservation_zmq_addr = f"tcp://127.0.0.1:{ports[0]}" + with patch.object( + scheduler, "_send_control", side_effect=fake_send + ) as send_control: + scheduler._ensure_event_channel() + + subscribed = [ + call.args[0] + for call in send_control.call_args_list + if call.args[1]["op"] == "event_port" + ] + assert len(subscribed) == len(ports) + assert scheduler._event_shard_count == len(ports) + + event = {"transfer_id": "transfer-0"} + assert not scheduler._note_shard_ready({**event, "shard": ports[0]}) + # The same rank reporting twice is not two ranks. + assert not scheduler._note_shard_ready({**event, "shard": ports[0]}) + assert not scheduler._note_shard_ready({**event, "shard": ports[1]}) + assert scheduler._note_shard_ready({**event, "shard": ports[2]}) + # Nothing is retained once the transfer is handed on. + assert "transfer-0" not in scheduler._event_ready_shards + finally: + scheduler.shutdown() + def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( self, mock_vllm_config_consumer, mock_request_with_3_mm ): @@ -875,8 +924,12 @@ def test_push_reserves_before_encoder_output_is_saved( ) as send_control: assert not scheduler.has_cache_item("hash") assert not scheduler.has_cache_item("hash") - assert send_control.call_count == 1 - assert send_control.call_args.args[1] == {"op": "event_port"} + # The channel is built once, not per call: the roster is + # fetched and every shard subscribed to on the first one. + assert [call.args[1] for call in send_control.call_args_list] == [ + {"op": "peers"}, + {"op": "event_port"}, + ] assert "transfer-1" in consumer._push_reservations producer.save_caches({"hash": source}, "hash") @@ -889,7 +942,9 @@ def test_push_reserves_before_encoder_output_is_saved( while not scheduler.has_cache_item("hash"): assert time.monotonic() < deadline time.sleep(0.01) - assert send_control.call_count == 1 + # Still just the two setup requests: polling for readiness + # must not re-open the channel. + assert send_control.call_count == 2 load = scheduler._pending_specs["transfer-1"] consumer.bind_connector_metadata( ECMooncakeConnectorMetadata(loads=[load]) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 146d038ad769..cd72b4ec1f85 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -45,6 +45,10 @@ _RESERVATION_REFRESH_SECONDS = _LEASE_TTL_SECONDS / 2 _RESERVATION_REAP_INTERVAL_SECONDS = 1 _DRAIN_MIN_INTERVAL = 0.005 +# Readiness notifications are advisory: the scheduler also learns from the +# reserve reply. Cap the queue so a shard nobody subscribed to cannot grow +# without bound. +_MAX_PENDING_EVENTS = 4096 _MOONCAKE_IMPORT_ERROR: ImportError | None try: @@ -448,6 +452,17 @@ def loop() -> None: event_socket = context.socket(zmq.PUSH) pending_events: deque[dict[str, Any]] = deque() metrics: Counter[str] = Counter() + + def queue_event(event: dict[str, Any]) -> None: + # The shard tag lets the scheduler tell each rank's readiness + # apart; a transfer is only loadable once every rank has it. + event["shard"] = self.port + if len(pending_events) >= _MAX_PENDING_EVENTS: + pending_events.popleft() + metrics["events_dropped"] += 1 + pending_events.append(event) + metrics["events_queued"] += 1 + metrics_started_at = time.monotonic() last_reap_at = metrics_started_at socket.setsockopt(zmq.RCVTIMEO, 100) @@ -483,8 +498,8 @@ def loop() -> None: ): logger.info( "EC Mooncake consumer control: requests=%s, " - "events_queued=%d, events_sent=%d, event_backlog=%d, " - "reservations_reaped=%d", + "events_queued=%d, events_sent=%d, events_dropped=%d, " + "event_backlog=%d, reservations_reaped=%d", { key.removeprefix("request_"): value for key, value in metrics.items() @@ -492,6 +507,7 @@ def loop() -> None: }, metrics["events_queued"], metrics["events_sent"], + metrics["events_dropped"], len(pending_events), metrics["reservations_reaped"], ) @@ -511,10 +527,7 @@ def loop() -> None: transfer_id = str(request["transfer_id"]) status = self._status(transfer_id) if status is not None: - pending_events.append( - {"transfer_id": transfer_id, **status} - ) - metrics["events_queued"] += 1 + queue_event({"transfer_id": transfer_id, **status}) elif op == "status": result = self._status(str(request["transfer_id"])) elif op == "event_port": @@ -546,10 +559,7 @@ def loop() -> None: continue status = self._status(transfer_id) if status is not None: - pending_events.append( - {"transfer_id": transfer_id, **status} - ) - metrics["events_queued"] += 1 + queue_event({"transfer_id": transfer_id, **status}) result = ( {"items": completions} if op == "complete_batch" @@ -797,6 +807,12 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._queued_transfer_batches = 0 self._event_zmq_ctx: zmq.Context | None = None self._event_zmq_socket: zmq.Socket | None = None + self._event_shard_count = 1 + # transfer_id -> shards that reported it ready, oldest first. A sharded + # consumer writes one copy per rank, so the item is only loadable once + # every rank has reported. Bounded: a transfer whose last rank never + # arrives is given up on by the push-wait timeout, not by this map. + self._event_ready_shards: OrderedDict[str, set[int]] = OrderedDict() self._completed_loads: set[str] = set() self._failed_loads: set[str] = set() self._shutdown = False @@ -2200,6 +2216,7 @@ def _index_pending_spec(self, spec: ECMooncakeLoadSpec) -> None: def _pop_pending_spec(self, transfer_id: str) -> ECMooncakeLoadSpec | None: spec = self._pending_specs.pop(transfer_id, None) self._pending_spec_deadlines.pop(transfer_id, None) + self._forget_shard_readiness(transfer_id) if spec is not None: if not self._pending_specs_by_hash.get(spec.mm_hash): self._consumer_pending_since.pop(spec.mm_hash, None) @@ -2227,6 +2244,36 @@ def _first_pending_spec(self, mm_hash: str) -> ECMooncakeLoadSpec | None: self._consumer_pending_since.pop(mm_hash, None) return None + def _note_shard_ready(self, data: dict[str, Any]) -> bool: + """Whether every consumer shard has now reported this transfer ready. + + Loading before the last rank has its copy makes that rank miss, which + `ECMooncakeWorkerMetadata.aggregate` catches by intersecting `loaded` + across ranks -- at the cost of rescheduling the whole load. + """ + if self._event_shard_count <= 1: + return True + transfer_id = str(data["transfer_id"]) + if transfer_id in self._pending_specs: + # Already indexed; later shards are just confirmations. + return False + shard = data.get("shard") + shards = self._event_ready_shards.setdefault(transfer_id, set()) + self._event_ready_shards.move_to_end(transfer_id) + shards.add(int(shard) if shard is not None else len(shards)) + if len(shards) < self._event_shard_count: + self._consumer_scheduler_metrics["events_awaiting_shards"] += 1 + while len(self._event_ready_shards) > _MAX_PENDING_EVENTS: + self._event_ready_shards.popitem(last=False) + self._consumer_scheduler_metrics["events_partial_dropped"] += 1 + return False + self._event_ready_shards.pop(transfer_id, None) + self._consumer_scheduler_metrics["events_all_shards_ready"] += 1 + return True + + def _forget_shard_readiness(self, transfer_id: str) -> None: + self._event_ready_shards.pop(transfer_id, None) + def _store_pushed_spec(self, data: dict[str, Any]) -> None: transfer_id = str(data["transfer_id"]) identifier = str(data["mm_hash"]) @@ -2297,14 +2344,34 @@ def _ensure_event_channel(self) -> None: if self._event_zmq_socket is not None: return assert self._reservation_zmq_addr is not None - event_port = self._send_control( - self._reservation_zmq_addr, - {"op": "event_port"}, - ) - address, _ = self._reservation_zmq_addr.rsplit(":", 1) - self._event_zmq_ctx = zmq.Context() - self._event_zmq_socket = self._event_zmq_ctx.socket(zmq.PULL) - self._event_zmq_socket.connect(f"{address}:{int(event_port)}") + shards = self._consumer_shards(self._reservation_zmq_addr) + ctx = zmq.Context() + socket = ctx.socket(zmq.PULL) + # One PULL fair-queues across every shard's PUSH. Subscribing to the + # first shard alone leaves the others' notifications queued on their + # side forever, and hides their readiness from the scheduler. + connected = 0 + for addr in shards: + try: + event_port = self._send_control(addr, {"op": "event_port"}) + address, _ = addr.rsplit(":", 1) + socket.connect(f"{address}:{int(event_port)}") + except Exception: + logger.warning( + "EC Mooncake could not subscribe to the event channel of " + "consumer shard %s; its readiness will only be seen " + "through reserve replies.", + addr, + ) + continue + connected += 1 + if not connected: + socket.close(linger=0) + ctx.term() + return + self._event_zmq_ctx = ctx + self._event_zmq_socket = socket + self._event_shard_count = connected def _drain_push_notifications(self) -> None: # `has_cache_item` and `ensure_cache_available` run once per request @@ -2337,6 +2404,8 @@ def _drain_push_notifications(self) -> None: # on to the spec so an eviction does not strand whoever # this transfer belongs to. self._consumer_scheduler_metrics["events_redundant"] += 1 + if not self._note_shard_ready(data): + continue self._store_pushed_spec(data) else: self._consumer_scheduler_metrics["events_not_ready"] += 1 From 72135d5eef6d1a2aa00b7aac343eef427fadee79 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Wed, 19 Aug 2026 00:46:50 +0000 Subject: [PATCH 12/30] Give each consumer replica its own block of control ports Data-parallel consumers each run their own scheduler and control channels, so sharing one port block collided at bind time and cross-subscribed their event channels. Offset the block by the replica's index so the addressing works, and narrow the data-parallel rejection to producers, which hold one copy of each encoder output and would only duplicate the push if replicated. The offset comes from `data_parallel_index` rather than `data_parallel_rank`: a non-MoE replica is reconfigured to look like DP=1, which resets the rank and the size, and reading the config rather than a process group keeps the scheduler, which has no groups, in agreement with its workers. Routing both halves of a request to the same replica remains the caller's job; the docstring says what the proxy has to do and how a mismatch surfaces. Signed-off-by: Tianyu Guo --- .../unit/test_ec_mooncake_connector.py | 44 ++++++++++-- .../ec_connector/mooncake_ec_connector.py | 67 +++++++++++++------ 2 files changed, 84 insertions(+), 27 deletions(-) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index d2fabc233a66..20ea692032bc 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -117,6 +117,7 @@ def mock_vllm_config_producer(): config.parallel_config.tensor_parallel_size = 1 config.parallel_config.pipeline_parallel_size = 1 config.parallel_config.data_parallel_size = 1 + config.parallel_config.data_parallel_index = 0 config.ec_transfer_config = Mock() config.ec_transfer_config.is_ec_producer = True config.ec_transfer_config.is_ec_consumer = False @@ -135,6 +136,7 @@ def mock_vllm_config_consumer(): config.parallel_config.tensor_parallel_size = 1 config.parallel_config.pipeline_parallel_size = 1 config.parallel_config.data_parallel_size = 1 + config.parallel_config.data_parallel_index = 0 config.ec_transfer_config = Mock() config.ec_transfer_config.is_ec_producer = False config.ec_transfer_config.is_ec_consumer = True @@ -209,14 +211,46 @@ def test_accepts_sharded_consumer(self, mock_vllm_config_consumer): ) connector.shutdown() - def test_rejects_data_parallel(self, mock_vllm_config_consumer): - """One control channel per instance cannot address a replica.""" - mock_vllm_config_consumer.parallel_config.data_parallel_size = 2 + def test_replicated_consumer_addresses_its_own_block( + self, mock_vllm_config_consumer + ): + """Each replica owns a distinct block of control ports. + + Replicas run their own schedulers and control channels, so sharing a + port would collide at bind time and cross-subscribe their event + channels. The block is derived from `data_parallel_index` because a + non-MoE replica is reconfigured to look like DP=1, which resets + `data_parallel_rank`. + """ + cfg = mock_vllm_config_consumer + cfg.parallel_config.tensor_parallel_size = 2 + cfg.parallel_config.data_parallel_size = 3 + cfg.parallel_config.data_parallel_index = 2 + # What a non-MoE replica actually looks like: reconfigured to DP=1, so + # `data_parallel_rank` no longer identifies it but the index still does. + cfg.parallel_config.data_parallel_rank = 0 + cfg.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "reservation_zmq_port": 19500, + } + with patch_ec_mooncake_deps(): + connector = ECMooncakeConnector(cfg, ECConnectorRole.SCHEDULER) + try: + # Replica 2 of a TP=2 consumer starts after two 2-port blocks. + assert connector._control_port_offset == 4 + assert connector._reservation_zmq_addr == "tcp://127.0.0.1:19504" + finally: + connector.shutdown() + + def test_rejects_replicated_producer(self, mock_vllm_config_producer): + """A producer holds one copy of each output, so replicating it only + duplicates the push.""" + mock_vllm_config_producer.parallel_config.data_parallel_size = 2 with ( patch_ec_mooncake_deps(), - pytest.raises(ValueError, match="data parallelism"), + pytest.raises(ValueError, match="data_parallel_size=1"), ): - ECMooncakeConnector(mock_vllm_config_consumer, ECConnectorRole.SCHEDULER) + ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.SCHEDULER) class TestECMooncakeWorkerMetadataAggregation: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index cd72b4ec1f85..b1103d1aa90c 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -625,9 +625,10 @@ class ECMooncakeConnector(ECConnectorBase): - ``consumer_buffer_pool_size`` (consumer, optional): Bytes reserved for a long-lived registered CUDA receive arena (default ``ec_buffer_size``). - ``reservation_zmq_port`` (consumer worker, required): Exposes registered - receive addresses over ZMQ. Tensor-parallel rank ``r`` of the first - pipeline stage listens on ``port + r``, and rank 0 reports the whole set, - so a producer only needs the first address. + receive addresses over ZMQ. Replica ``d`` of the first pipeline stage owns + the block starting at ``port + d * tensor_parallel_size``; tensor-parallel + rank ``r`` in that block listens on ``block + r``, and rank 0 reports the + whole block, so a producer only needs the block's first address. - ``reservation_zmq_addr`` (consumer scheduler, required): Address of the consumer control channel. Defaults to ``tcp://127.0.0.1:``. - ``transfer_max_workers`` (optional): Maximum concurrent Mooncake transfer @@ -639,13 +640,25 @@ class ECMooncakeConnector(ECConnectorBase): - ``consumer_metrics_log_interval`` (optional): Seconds between aggregated consumer lifecycle logs (default ``10``; ``0`` disables them). - Parallelism: consumers may use tensor and pipeline parallelism. Only the - first pipeline stage holds encoder outputs, and each tensor-parallel rank - there gathers from its own cache, so every rank exposes a control channel - and the producer writes into all of them concurrently from one registered - source. That costs bandwidth but not latency, and avoids the second hop a - receive-then-broadcast would add. Producers must be unsharded, and data - parallelism is unsupported on either side. + Parallelism: consumers may use tensor, pipeline and data parallelism. + Producers must be unsharded and unreplicated: one copy of each encoder + output is held and addressed directly, so splitting the producer would only + duplicate the push. + + Only the first pipeline stage holds encoder outputs, and each tensor-parallel + rank there gathers from its own cache, so every rank exposes a control + channel and the producer writes into all of them concurrently from one + registered source. That costs bandwidth but not latency, and avoids the + second hop a receive-then-broadcast would add. + + Data parallelism additionally requires the caller to route both halves of a + request to the same replica, because a push has to land where the request + will run. The proxy is the only component that knows which replica it picked: + it names the replica to the consumer (``X-data-parallel-rank``) and passes + that replica's control address to the producer. Getting this wrong is loud + rather than silent -- the replica that runs the request never sees its + embedding and gives up after ``push_wait_timeout_s`` -- but it is the + caller's responsibility, not something this connector can detect. """ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): @@ -657,17 +670,12 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): ) from _MOONCAKE_IMPORT_ERROR parallel_config = vllm_config.parallel_config - if parallel_config.data_parallel_size > 1: - raise ValueError( - "ECMooncakeConnector does not support data parallelism yet: the " - "consumer exposes one control channel per instance, so a push " - "cannot be routed to the replica that will run the request." - ) ec_cfg_early = vllm_config.ec_transfer_config assert ec_cfg_early is not None if ec_cfg_early.is_ec_producer: # The producer holds one copy of each encoder output and addresses - # consumers directly; sharding it would only duplicate the push. + # consumers directly; sharding or replicating it would only + # duplicate the push. if parallel_config.tensor_parallel_size > 1: raise ValueError( "ECMooncakeConnector producers require tensor_parallel_size=1." @@ -676,6 +684,21 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): raise ValueError( "ECMooncakeConnector producers do not support pipeline parallelism." ) + if parallel_config.data_parallel_size > 1: + raise ValueError( + "ECMooncakeConnector producers require data_parallel_size=1." + ) + + # Each data-parallel replica runs its own scheduler and its own control + # channels, so their ports must not overlap. `data_parallel_index` is the + # only field that identifies the replica in both cases: a non-MoE replica + # is reconfigured to look like DP=1, which resets `data_parallel_rank` + # and `data_parallel_size`. Deriving the offset from the config rather + # than from a process group keeps the scheduler, which has no groups, in + # agreement with its workers. + self._control_port_offset = ( + parallel_config.data_parallel_index * parallel_config.tensor_parallel_size + ) self._role = role ec_cfg = vllm_config.ec_transfer_config @@ -692,7 +715,8 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._reservation_zmq_addr is None and self._reservation_zmq_port is not None ): - self._reservation_zmq_addr = f"tcp://127.0.0.1:{self._reservation_zmq_port}" + base = self._reservation_zmq_port + self._control_port_offset + self._reservation_zmq_addr = f"tcp://127.0.0.1:{base}" self._registered_capacity = int(self._ec_cfg.ec_buffer_size) if self._registered_capacity <= 0: raise ValueError("ECMooncakeConnector requires ec_buffer_size > 0.") @@ -893,18 +917,17 @@ def start_worker_services(self) -> None: raise RuntimeError( "Mooncake push mode requires a registered consumer buffer pool." ) + base_port = self._reservation_zmq_port + self._control_port_offset self._control_server = ECMooncakeControlServer( "0.0.0.0", - self._reservation_zmq_port + self._tp_rank, + base_port + self._tp_rank, self._reserve_push_destination, self._push_status, self._complete_push, self._cancel_push, self._expire_push_reservations, self._consumer_metrics_log_interval, - peer_ports=[ - self._reservation_zmq_port + rank for rank in range(self._tp_size) - ], + peer_ports=[base_port + rank for rank in range(self._tp_size)], device=self._consumer_pool.device, ) self._control_server.start() From af3e6999eaeeef65c51205b50fc8a9444f20eab8 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Wed, 19 Aug 2026 01:27:39 +0000 Subject: [PATCH 13/30] Support DP Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 63 +++++++++++++++---- 1 file changed, 51 insertions(+), 12 deletions(-) diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index fb12f3a0a572..71f56ac51310 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -23,6 +23,7 @@ import asyncio import hashlib import io +import itertools import json import logging import os @@ -496,6 +497,7 @@ async def forward_non_stream( p_url: str, d_url: str, consumer_zmq: str | None, + dp_rank: int | None = None, ) -> dict: try: for attempt in range(DECODE_RETRIES + 1): @@ -507,6 +509,8 @@ async def forward_non_stream( logger.info("[%s] Forwarding to decode: %s", req_id, d_url) headers = {"x-request-id": req_id} + if dp_rank is not None: + headers["X-data-parallel-rank"] = str(dp_rank) async with decode_session.post( f"{d_url}/v1/chat/completions", json=prepared, headers=headers @@ -562,6 +566,7 @@ async def forward_stream( p_url: str, d_url: str, consumer_zmq: str | None, + dp_rank: int | None = None, ) -> AsyncIterator[str]: try: for attempt in range(DECODE_RETRIES + 1): @@ -573,6 +578,8 @@ async def forward_stream( logger.info("[%s] Starting streaming from decode: %s", req_id, d_url) headers = {"x-request-id": req_id} + if dp_rank is not None: + headers["X-data-parallel-rank"] = str(dp_rank) _first = None async with decode_session.post( @@ -639,19 +646,26 @@ async def chat_completions(request: Request): p_url = random.choice(app.state.p_urls) if app.state.p_urls else None decode_index = random.randrange(len(app.state.d_urls)) d_url = app.state.d_urls[decode_index] - consumer_zmq = ( - app.state.d_ec_urls[decode_index] if app.state.d_ec_urls else None - ) + dp_size = app.state.ec_consumer_dp_size + # Round-robin the replica, then name it to both halves: the decoder + # honours the rank header instead of its own balancer, and the encoder + # pushes to that replica's control channel. Choosing once here means a + # decode retry re-encodes to the same replica. + dp_rank = next(app.state.replica_counter) % dp_size if dp_size > 1 else None + ec_index = decode_index * dp_size + (dp_rank or 0) + consumer_zmq = app.state.d_ec_urls[ec_index] if app.state.d_ec_urls else None is_streaming = req_data.get("stream", False) if is_streaming: return StreamingResponse( - forward_stream(req_data, req_id, e_urls, p_url, d_url, consumer_zmq), + forward_stream( + req_data, req_id, e_urls, p_url, d_url, consumer_zmq, dp_rank + ), media_type="text/event-stream", ) result = await forward_non_stream( - req_data, req_id, e_urls, p_url, d_url, consumer_zmq + req_data, req_id, e_urls, p_url, d_url, consumer_zmq, dp_rank ) return JSONResponse(content=result) @@ -837,12 +851,23 @@ async def stop_profile(request: Request): ), ) parser.add_argument( - "--decode-ec-transfer-zmq-addrs", + "--ec-consumer-zmq-addrs", default="", help=( - "Comma-separated Mooncake EC consumer ZMQ addresses, aligned " - "with --decode-servers-urls. Required when the decoders use the " - "Mooncake EC connector." + "Comma-separated Mooncake EC consumer control addresses, aligned " + "with --decode-servers-urls. Required when the consumers use the " + "Mooncake EC connector. With --ec-consumer-dp-size > 1, list each " + "server's replicas consecutively: s0r0,s0r1,s1r0,s1r1." + ), + ) + parser.add_argument( + "--ec-consumer-dp-size", + type=int, + default=1, + help=( + "Data-parallel replicas per EC consumer. The proxy picks a replica " + "round-robin and names it to both halves of the request, because an " + "encoder push has to land where the request will run." ), ) @@ -856,11 +881,19 @@ async def stop_profile(request: Request): u.strip() for u in args.decode_servers_urls.split(",") if u.strip() ] app.state.d_ec_urls = [ - u.strip() for u in args.decode_ec_transfer_zmq_addrs.split(",") if u.strip() + u.strip() for u in args.ec_consumer_zmq_addrs.split(",") if u.strip() ] - if app.state.d_ec_urls and len(app.state.d_ec_urls) != len(app.state.d_urls): + if args.ec_consumer_dp_size < 1: + parser.error("--ec-consumer-dp-size must be at least 1") + app.state.ec_consumer_dp_size = args.ec_consumer_dp_size + app.state.replica_counter = itertools.count() + expected = len(app.state.d_urls) * args.ec_consumer_dp_size + if app.state.d_ec_urls and len(app.state.d_ec_urls) != expected: parser.error( - "--decode-ec-transfer-zmq-addrs must contain one address per decode server" + "--ec-consumer-zmq-addrs must contain one address per consumer " + f"replica: expected {expected} " + f"({len(app.state.d_urls)} servers x {args.ec_consumer_dp_size} replicas), " + f"got {len(app.state.d_ec_urls)}" ) # handle prefill instances if args.prefill_servers_urls.lower() in ("disable", "none", ""): @@ -878,6 +911,12 @@ async def stop_profile(request: Request): logger.info("Encode servers: %s", app.state.e_urls) logger.info("Prefill instances %s", app.state.p_urls) logger.info("Decode servers: %s", app.state.d_urls) + if app.state.ec_consumer_dp_size > 1: + logger.info( + "EC consumer replicas per server: %d (control addresses: %s)", + app.state.ec_consumer_dp_size, + app.state.d_ec_urls, + ) uvicorn.run( app, From d4ff1b44577f32e03dddce95af92fc5c224c47db Mon Sep 17 00:00:00 2001 From: Zhou ziheng Date: Mon, 24 Aug 2026 11:12:36 +0800 Subject: [PATCH 14/30] [Bugfix][EPD] Fix DeepStack encoder cache reservation sizing (#4) Assisted-by: OpenAI Codex Signed-off-by: Zhou ziheng --- .../unit/test_ec_mooncake_connector.py | 33 ++++++++++++++++++- .../ec_connector/mooncake_ec_connector.py | 5 ++- 2 files changed, 36 insertions(+), 2 deletions(-) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 20ea692032bc..7d3cef442fd6 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -792,7 +792,8 @@ def test_producer_builds_push_metadata_after_preprocessing( "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], } mock_vllm_config_producer.model_config.dtype = torch.float32 - mock_vllm_config_producer.model_config.get_hidden_size.return_value = 16 + mock_vllm_config_producer.model_config.hf_config = None + mock_vllm_config_producer.model_config.get_inputs_embeds_size.return_value = 16 with patch_ec_mooncake_deps(): scheduler = ECMooncakeConnector( @@ -816,6 +817,36 @@ def test_producer_builds_push_metadata_after_preprocessing( ) ] + def test_producer_uses_deepstack_encoder_cache_width( + self, mock_vllm_config_producer, mock_request_with_3_mm + ): + request = mock_request_with_3_mm + request.ec_transfer_params = { + "consumer_zmq": "tcp://decode:19019", + "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], + } + mock_vllm_config_producer.model_config.dtype = torch.bfloat16 + mock_vllm_config_producer.model_config.hf_config = SimpleNamespace( + vision_config=SimpleNamespace( + out_hidden_size=2560, + deepstack_visual_indexes=[5, 11, 17], + ) + ) + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + scheduler.update_state_after_alloc(request, 0) + meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) + + num_tokens = request.get_num_encoder_embeds(0) + spec = meta.pushes[0] + assert spec.shape == (num_tokens, 10240) + assert spec.nbytes == num_tokens * 10240 * torch.bfloat16.itemsize + def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer): feature = SimpleNamespace( identifier="image_uuid", diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index b1103d1aa90c..7f917727cdfb 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -32,6 +32,9 @@ ECConnectorRole, ECConnectorWorkerMetadata, ) +from vllm.distributed.ec_transfer.ec_connector.cpu.common import ( + _get_encoder_cache_hidden_dim, +) from vllm.logger import init_logger from vllm.utils.network_utils import get_ip from vllm.v1.core.sched.output import SchedulerOutput @@ -2563,7 +2566,7 @@ def _prepare_push_spec(self, request: Any, index: int) -> None: dtype = self._model_config.dtype assert isinstance(dtype, torch.dtype) dtype_name = str(dtype).split(".")[-1] - shape = (num_tokens, self._model_config.get_hidden_size()) + shape = (num_tokens, _get_encoder_cache_hidden_dim(self._vllm_config)) nbytes = math.prod(shape) * dtype.itemsize self._pushes_to_prepare[transfer_id] = ECMooncakePushSpec( mm_hash=mm_hash, From 720f98f89c2371e58441b398e3b352492fc018fa Mon Sep 17 00:00:00 2001 From: Zhou ziheng Date: Fri, 28 Aug 2026 07:58:13 +0800 Subject: [PATCH 15/30] [Bugfix] Prevent duplicate Mooncake push metadata across scheduler steps (#5) Assisted-by: OpenAI Codex Signed-off-by: Zhou ziheng --- .../unit/test_ec_mooncake_connector.py | 12 ++++++++++++ .../ec_connector/mooncake_ec_connector.py | 17 ++++++++++++++++- 2 files changed, 28 insertions(+), 1 deletion(-) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 7d3cef442fd6..89cabec2d0ac 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -787,6 +787,7 @@ def test_producer_builds_push_metadata_after_preprocessing( self, mock_vllm_config_producer, mock_request_with_3_mm ): request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] request.ec_transfer_params = { "consumer_zmq": "tcp://decode:19019", "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], @@ -804,6 +805,15 @@ def test_producer_builds_push_metadata_after_preprocessing( Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) ) + # The same request remains visible on a later scheduler step, but + # its worker push metadata must not be emitted a second time. + assert scheduler.ensure_cache_available(request, 0) + next_meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) + + scheduler.request_finished(request) + assert meta.loads == [] assert meta.pushes == [ ECMooncakePushSpec( @@ -816,6 +826,8 @@ def test_producer_builds_push_metadata_after_preprocessing( request_id="test_req_123", ) ] + assert next_meta.pushes == [] + assert "transfer-1" not in scheduler._prepared_push_transfer_ids def test_producer_uses_deepstack_encoder_cache_width( self, mock_vllm_config_producer, mock_request_with_3_mm diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 7f917727cdfb..3b9553c023c2 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -788,6 +788,10 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._resident_bytes = 0 self._scheduler_pending_work = False self._pushes_to_prepare: dict[str, ECMooncakePushSpec] = {} + # A producer request may be revisited across scheduler steps. Queue its + # initial push metadata once; the worker owns reservation refreshes + # after the scheduler emits it. + self._prepared_push_transfer_ids: set[str] = set() # Worker producer self._engine: TransferEngine | None = None @@ -2560,7 +2564,7 @@ def _prepare_push_spec(self, request: Any, index: int) -> None: transfer_id = self._request_transfer_id(request, index) if transfer_id is None: transfer_id = f"{request.request_id}:{index}" - if not consumer_zmq or transfer_id in self._pushes_to_prepare: + if not consumer_zmq or transfer_id in self._prepared_push_transfer_ids: return num_tokens = request.get_num_encoder_embeds(index) dtype = self._model_config.dtype @@ -2577,6 +2581,7 @@ def _prepare_push_spec(self, request: Any, index: int) -> None: transfer_id=transfer_id, request_id=request.request_id, ) + self._prepared_push_transfer_ids.add(transfer_id) def update_state_after_alloc(self, request: Any, index: int) -> None: mm_hash = request.mm_features[index].identifier @@ -2740,6 +2745,16 @@ def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: continue self._pop_pending_spec(transfer_id) self._queue_cancel(transfer_id) + if ( + self.is_producer + and self._role == ECConnectorRole.SCHEDULER + and self._prepared_push_transfer_ids + ): + for index in range(len(request.mm_features)): + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + transfer_id = f"{request.request_id}:{index}" + self._prepared_push_transfer_ids.discard(transfer_id) if not self.is_producer: return False, None From f67f3906b68fb01c2fa0687210e3569034445567 Mon Sep 17 00:00:00 2001 From: Zhou ziheng Date: Fri, 28 Aug 2026 08:38:23 +0800 Subject: [PATCH 16/30] [Bugfix] Ignore late Mooncake ready events for cancelled transfers (#6) Assisted-by: OpenAI Codex Signed-off-by: Zhou ziheng --- .../unit/test_ec_mooncake_connector.py | 35 +++++++++++++++++++ .../ec_connector/mooncake_ec_connector.py | 3 ++ 2 files changed, 38 insertions(+) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 89cabec2d0ac..d946043955ef 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -504,6 +504,41 @@ def test_ready_hash_eviction_does_not_strand_a_later_transfer( finally: scheduler.shutdown() + def test_cancelled_transfer_ignores_late_ready_events( + self, mock_vllm_config_consumer + ): + """Cancelled is terminal even when a ready event was already queued.""" + transfer_id = "cancelled-transfer" + ports = [19101, 19102, 19103, 19104] + event = { + "mm_hash": "hash", + "transfer_id": transfer_id, + "ready": True, + "reservation_id": "reservation", + "nbytes": 16, + "shape": [4], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._event_shard_count = len(ports) + scheduler._cancelled_transfer_ids.add(transfer_id) + scheduler._event_zmq_socket = Mock() + scheduler._event_zmq_socket.recv_json.side_effect = [ + {**event, "shard": port} for port in ports + ] + [zmq.Again()] + + scheduler._drain_push_notifications() + + assert transfer_id not in scheduler._pending_specs + assert transfer_id not in scheduler._event_ready_shards + finally: + scheduler.shutdown() + def test_item_that_never_arrives_fails_the_request( self, mock_vllm_config_consumer, mock_request_with_3_mm ): diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 3b9553c023c2..b3f8534da7b6 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -2429,6 +2429,9 @@ def _drain_push_notifications(self) -> None: self._consumer_scheduler_metrics["events_received"] += 1 if data.get("ready"): self._consumer_scheduler_metrics["events_ready"] += 1 + transfer_id = str(data["transfer_id"]) + if transfer_id in self._cancelled_transfer_ids: + continue if identifier in self._ready_hashes: # Redundant only for as long as the hash stays ready; hold # on to the spec so an eviction does not strand whoever From 4151de2e2983225ae04c488be847356aa14f62b5 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Fri, 28 Aug 2026 00:49:04 +0000 Subject: [PATCH 17/30] [Bugfix] Sweep Mooncake cancel records in deadline order Signed-off-by: Tianyu Guo --- .../unit/test_ec_mooncake_connector.py | 146 +++++++++++++++++- .../ec_connector/mooncake_ec_connector.py | 51 +++++- 2 files changed, 189 insertions(+), 8 deletions(-) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index d946043955ef..99cbee1d53f6 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -526,7 +526,9 @@ def test_cancelled_transfer_ignores_late_ready_events( ) try: scheduler._event_shard_count = len(ports) - scheduler._cancelled_transfer_ids.add(transfer_id) + scheduler._cancelled_transfer_ids[transfer_id] = ( + time.monotonic() + _LEASE_TTL_SECONDS + ) scheduler._event_zmq_socket = Mock() scheduler._event_zmq_socket.recv_json.side_effect = [ {**event, "shard": port} for port in ports @@ -536,6 +538,115 @@ def test_cancelled_transfer_ignores_late_ready_events( assert transfer_id not in scheduler._pending_specs assert transfer_id not in scheduler._event_ready_shards + assert scheduler._consumer_scheduler_metrics["events_cancelled"] == ( + len(ports) + ) + finally: + scheduler.shutdown() + + def test_cancel_between_shards_drops_the_partial_readiness( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """A cancel mid-aggregation leaves nothing for the late shards to finish. + + The early shards are already counted when the request releases the + item. Only clearing them keeps the remaining notifications from + completing the set and rebuilding a spec for a buffer the worker + freed as it cancelled. + """ + request = mock_request_with_3_mm + request.mm_features = request.mm_features[:1] + mm_hash = request.mm_features[0].identifier + transfer_id = "half-reported-transfer" + request.ec_transfer_params = { + "ec_items": [{"mm_hash": mm_hash, "transfer_id": transfer_id}] + } + event = { + "mm_hash": mm_hash, + "transfer_id": transfer_id, + "ready": True, + "reservation_id": "reservation", + "nbytes": 16, + "shape": [4], + "dtype": "float32", + } + + def deliver(scheduler, *shards): + scheduler._event_zmq_socket.recv_json.side_effect = [ + {**event, "shard": shard} for shard in shards + ] + [zmq.Again()] + scheduler._drain_pending = True + scheduler._drain_push_notifications() + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" + scheduler._event_shard_count = 4 + scheduler._event_zmq_socket = Mock() + + deliver(scheduler, 0, 1) + assert scheduler._event_ready_shards[transfer_id] == {0, 1} + + with patch.object(scheduler, "_cancel_remote", return_value=True): + scheduler.update_state_after_free(request, 0) + assert transfer_id in scheduler._cancelled_transfer_ids + assert transfer_id not in scheduler._event_ready_shards + + deliver(scheduler, 2, 3) + + assert transfer_id not in scheduler._pending_specs + assert transfer_id not in scheduler._event_ready_shards + assert scheduler._consumer_scheduler_metrics["events_cancelled"] == 2 + finally: + scheduler.shutdown() + + def test_cancelled_transfer_ids_stay_bounded(self, mock_vllm_config_consumer): + """The ignore list is swept, not accumulated. + + It is consulted for every readiness notification and grows by one + entry per multimodal item the instance serves, so retaining ids the + worker has itself forgotten leaks for the life of the process. + """ + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" + scheduler._event_zmq_socket = Mock() + scheduler._event_zmq_socket.recv_json.side_effect = zmq.Again() + with patch.object(scheduler, "_cancel_remote", return_value=True): + for name in ("first", "second", "third"): + scheduler._queue_cancel(name) + + now = time.monotonic() + assert scheduler._cancelled_transfer_ids["third"] > now + assert scheduler._cancelled_transfer_ids["third"] <= ( + now + _LEASE_TTL_SECONDS + ) + + # Ignored for exactly as long as the worker refuses to reserve + # the id again, and no longer. The drain is what sweeps. + scheduler._cancelled_transfer_ids["first"] = 0.0 + scheduler._drain_pending = True + scheduler._drain_push_notifications() + assert list(scheduler._cancelled_transfer_ids) == ["second", "third"] + + # The count is the backstop for a rate that outruns the TTL. + with patch( + "vllm.distributed.ec_transfer.ec_connector." + "mooncake_ec_connector._MAX_CANCELLED_TRANSFER_IDS", + 1, + ): + scheduler._drain_pending = True + scheduler._drain_push_notifications() + assert list(scheduler._cancelled_transfer_ids) == ["third"] + assert ( + scheduler._consumer_scheduler_metrics["cancel_records_dropped"] == 2 + ) finally: scheduler.shutdown() @@ -1734,6 +1845,39 @@ def test_cancel_before_reserve_creates_bounded_tombstone( finally: consumer.shutdown() + def test_repeated_cancel_does_not_strand_older_tombstones( + self, mock_vllm_config_consumer + ): + """Re-cancelling refreshes a tombstone without breaking the sweep order. + + The sweep stops at the first live record, so a refreshed one that kept + its original position would shield every older record behind it and + the table would grow for the life of the process. + """ + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) + assert consumer._cancel_push("refreshed-transfer", "") + assert consumer._cancel_push("stale-transfer", "") + consumer._cancelled_transfers["stale-transfer"] = 0.0 + assert consumer._cancel_push("refreshed-transfer", "") + + consumer._expire_push_reservations() + + assert list(consumer._cancelled_transfers) == ["refreshed-transfer"] + assert consumer._consumer_worker_metrics["cancel_records_dropped"] == 1 + finally: + consumer.shutdown() + def test_missing_push_reservation_reports_failed_load( self, mock_vllm_config_consumer ): diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index b3f8534da7b6..3d5a31a64358 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -52,6 +52,10 @@ # reserve reply. Cap the queue so a shard nobody subscribed to cannot grow # without bound. _MAX_PENDING_EVENTS = 4096 +# A cancelled transfer stays on the scheduler's ignore list for as long as the +# worker refuses to reserve it again. The count is a backstop for a rate that +# outruns that TTL; the race it guards is a single drain interval wide. +_MAX_CANCELLED_TRANSFER_IDS = 1 << 16 _MOONCAKE_IMPORT_ERROR: ImportError | None try: @@ -750,7 +754,7 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._consumer_pool_disabled = self._consumer_pool_capacity <= 0 self._consumer_lock = threading.Lock() self._push_reservations: dict[str, _PushReservation] = {} - self._cancelled_transfers: dict[str, float] = {} + self._cancelled_transfers: OrderedDict[str, float] = OrderedDict() self._control_server: ECMooncakeControlServer | None = None self._consumer_metrics_log_interval = float( self._extra.get("consumer_metrics_log_interval", 10) @@ -770,7 +774,9 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._consumer_pending_since: dict[str, float] = {} self._pending_spec_deadlines: dict[str, float] = {} self._pending_cancels: dict[str, Future[Any]] = {} - self._cancelled_transfer_ids: set[str] = set() + # Cancelled transfers, oldest deadline first, so the sweep can stop + # at the first live entry. + self._cancelled_transfer_ids: OrderedDict[str, float] = OrderedDict() # Scheduler (consumer): transfer_id -> pending tensor layout. self._pending_specs: dict[str, ECMooncakeLoadSpec] = {} @@ -1414,6 +1420,30 @@ def _maybe_log_consumer_scheduler_metrics(self) -> None: self._consumer_scheduler_metrics.clear() self._consumer_metrics_started_at = now + @staticmethod + def _expire_cancel_records(records: OrderedDict[str, float], now: float) -> int: + """Drop the cancels that can no longer be told apart from unknown ids. + + Both roles keep one record per multimodal item they handle, and both + consult it on a per-item hot path, so a full rescan costs the square + of the item rate: at 53 items/s the worker's 300 s window is 16k + entries and its sweep ran under `_consumer_lock` on every + reservation. Callers append in deadline order -- `move_to_end` when + refreshing one -- so the front is always the oldest and the sweep + stops at the first live entry. + + Returns: + How many records were dropped. + """ + dropped = 0 + while records: + expires_at = next(iter(records.values())) + if expires_at > now and len(records) <= _MAX_CANCELLED_TRANSFER_IDS: + break + records.popitem(last=False) + dropped += 1 + return dropped + def _expire_push_reservations_locked(self) -> None: now = time.monotonic() allocator = self._consumer_pool_allocator @@ -1427,9 +1457,9 @@ def _expire_push_reservations_locked(self) -> None: ) self._push_reservations.pop(transfer_id) self._consumer_worker_metrics["reservations_expired"] += 1 - for transfer_id, expires_at in list(self._cancelled_transfers.items()): - if expires_at <= now: - self._cancelled_transfers.pop(transfer_id) + self._consumer_worker_metrics["cancel_records_dropped"] += ( + self._expire_cancel_records(self._cancelled_transfers, now) + ) def _expire_push_reservations(self) -> int: with self._consumer_lock: @@ -1595,6 +1625,7 @@ def _cancel_push( self._cancelled_transfers[transfer_id] = ( time.monotonic() + _LEASE_TTL_SECONDS ) + self._cancelled_transfers.move_to_end(transfer_id) if reservation is None: self._consumer_worker_metrics["cancellations_pre_reserved"] += 1 return True @@ -1748,7 +1779,7 @@ def _poll_pending_cancels(self) -> None: try: cancelled = future.result() except Exception: - self._cancelled_transfer_ids.discard(transfer_id) + self._cancelled_transfer_ids.pop(transfer_id, None) self._consumer_scheduler_metrics["cancellations_failed"] += 1 logger.warning( "EC Mooncake reservation cancellation failed", exc_info=True @@ -2351,7 +2382,9 @@ def _queue_cancel(self, transfer_id: str, reservation_id: str = "") -> None: or transfer_id in self._cancelled_transfer_ids ): return - self._cancelled_transfer_ids.add(transfer_id) + self._cancelled_transfer_ids[transfer_id] = ( + time.monotonic() + _LEASE_TTL_SECONDS + ) self._pending_cancels[transfer_id] = self._control_executor.submit( self._cancel_remote, self._reservation_zmq_addr, @@ -2415,6 +2448,9 @@ def _drain_push_notifications(self) -> None: self._drained_at = now self._poll_pending_cancels() self._expire_pending_specs() + self._consumer_scheduler_metrics["cancel_records_dropped"] += ( + self._expire_cancel_records(self._cancelled_transfer_ids, now) + ) if self._reservation_zmq_addr is not None: self._ensure_event_channel() socket = self._event_zmq_socket @@ -2431,6 +2467,7 @@ def _drain_push_notifications(self) -> None: self._consumer_scheduler_metrics["events_ready"] += 1 transfer_id = str(data["transfer_id"]) if transfer_id in self._cancelled_transfer_ids: + self._consumer_scheduler_metrics["events_cancelled"] += 1 continue if identifier in self._ready_hashes: # Redundant only for as long as the hash stays ready; hold From e5aa2bc3ec348996a09e5182afaac3ab65036d31 Mon Sep 17 00:00:00 2001 From: jiangkuaixue123 Date: Sat, 29 Aug 2026 11:44:48 +0800 Subject: [PATCH 18/30] Fix Mooncake EPD proxy flag in integration script (#7) The proxy renamed its consumer control address option while the full-pipeline script retained the old name, causing argparse to reject the launch command. Assisted-by: OpenAI Codex Signed-off-by: jiangkuaixue123 --- .../integration/run_epd_mooncake_ec_full_pipeline.sh | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh index 04e28c70483d..c419784a7374 100755 --- a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -182,7 +182,7 @@ run_epd_mooncake() { --encode-servers-urls "http://localhost:$ENCODE_PORT" \ --prefill-servers-urls "disable" \ --decode-servers-urls "http://localhost:$PREFILL_DECODE_PORT" \ - --decode-ec-transfer-zmq-addrs \ + --ec-consumer-zmq-addrs \ "tcp://localhost:$EC_MOONCAKE_RESERVATION_PORT" \ >"${LOG_PATH}/mooncake_epd_proxy.log" 2>&1 & PIDS+=($!) From 2f4444608a6ecf9889dac0796f4307a9fca54c17 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Sat, 29 Aug 2026 04:04:24 +0000 Subject: [PATCH 19/30] [EPD] Make the proxy's per-request logging opt-in Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index 3f293db9c061..abbc12639f01 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -42,9 +42,7 @@ # FastAPI app & global state ############################################################################### -logging.basicConfig( - level=logging.DEBUG, format="%(asctime)s %(levelname)s: %(message)s" -) +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s") logger = logging.getLogger("proxy") app = FastAPI() @@ -422,7 +420,6 @@ async def process_prefill_stage( ############################################################################### -@app.middleware("http") async def log_requests(request: Request, call_next): """Middleware to log all incoming requests and responses""" req_id = request.headers.get("x-request-id", str(uuid.uuid4())) @@ -843,6 +840,14 @@ async def stop_profile(request: Request): parser = argparse.ArgumentParser() parser.add_argument("--host", default="0.0.0.0") parser.add_argument("--port", type=int, default=8000) + parser.add_argument( + "--log-requests", + action="store_true", + help=( + "Log every request in and out, and raise the log level to DEBUG. " + "Off by default: the proxy is on the request path." + ), + ) parser.add_argument( "--no-rewrite", action="store_true", @@ -898,6 +903,9 @@ async def stop_profile(request: Request): ) args = parser.parse_args() + if args.log_requests: + logging.getLogger().setLevel(logging.DEBUG) + app.middleware("http")(log_requests) NO_REWRITE = args.no_rewrite DECODE_RETRIES = max(0, args.decode_retries) app.state.e_urls = [ From 0b319fdcfe2ab13f1358b9de85bd0459fd93451d Mon Sep 17 00:00:00 2001 From: jiangkuaixue123 Date: Tue, 1 Sep 2026 11:52:24 +0800 Subject: [PATCH 20/30] [EPD] Refactor ECMooncakeConnector internals (#9) Signed-off-by: jiangkuaixue123 Co-authored-by: OpenAI Codex --- .../unit/test_ec_mooncake_connector.py | 4174 +++++++++++++++-- .../ec_connector/mooncake/__init__.py | 9 + .../ec_connector/mooncake/_availability.py | 28 + .../ec_connector/mooncake/config.py | 232 + .../ec_connector/mooncake/control.py | 527 +++ .../ec_connector/mooncake/memory.py | 616 +++ .../ec_connector/mooncake/metadata.py | 118 + .../ec_connector/mooncake/producer.py | 427 ++ .../ec_connector/mooncake/reservation.py | 441 ++ .../ec_connector/mooncake/scheduler.py | 636 +++ .../ec_connector/mooncake/state.py | 432 ++ .../ec_connector/mooncake/transfer.py | 216 + .../ec_connector/mooncake/worker.py | 1154 +++++ .../ec_connector/mooncake_ec_connector.py | 2879 +----------- 14 files changed, 8748 insertions(+), 3141 deletions(-) create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/__init__.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/config.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/control.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/state.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py create mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 99cbee1d53f6..580936a49054 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -1,33 +1,95 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Unit tests for ECMooncakeConnector.""" +"""Behavioral contract for the refactored Mooncake encoder-cache connector. + +The suite covers public compatibility, configuration, control and data planes, +memory ownership, the three role-specific lifecycle managers, Scheduler +metadata, Worker orchestration, and failure or cancellation races. +""" from __future__ import annotations import copy import ctypes +import gc +import importlib import socket +import sys +import threading import time +import weakref +from collections import Counter, OrderedDict +from concurrent.futures import Future, ThreadPoolExecutor from contextlib import contextmanager -from types import SimpleNamespace -from unittest.mock import Mock, patch +from dataclasses import FrozenInstanceError +from multiprocessing.reduction import ForkingPickler +from types import ModuleType, SimpleNamespace +from typing import Any +from unittest.mock import MagicMock, Mock, call, patch import pytest import torch import zmq -from vllm.config import VllmConfig -from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole +from vllm.config import ModelConfig, VllmConfig +from vllm.distributed.ec_transfer.ec_connector import mooncake_ec_connector +from vllm.distributed.ec_transfer.ec_connector.base import ( + ECConnectorMetadata, + ECConnectorRole, +) from vllm.distributed.ec_transfer.ec_connector.factory import ECConnectorFactory -from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( +from vllm.distributed.ec_transfer.ec_connector.mooncake import ( + control, + memory, + metadata, + producer, + state, + transfer, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.config import MooncakeECConfig +from vllm.distributed.ec_transfer.ec_connector.mooncake.control import ( + ConsumerControlServer, + ControlClient, + ControlCompletion, + EventInbox, + ShardTopology, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.memory import ( + ConsumerMemoryPool, + ContiguousAllocator, + ProducerMemoryPool, + ResidentPool, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.producer import ( + ProducerPushManager, + ProducerPushState, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.reservation import ( + CancellationOutcome, + ConsumerReservationManager, + ConsumerReservationState, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.scheduler import ( + ECMooncakeScheduler, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.state import ( + InvalidSchedulerTransferTransition, + SchedulerTransferState, + SchedulerTransferTable, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( + MooncakeTransfer, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.worker import ( _LEASE_TTL_SECONDS, + ECMooncakeWorker, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( ECMooncakeConnector, ECMooncakeConnectorMetadata, ECMooncakeLoadSpec, ECMooncakePushSpec, ECMooncakeWorkerMetadata, - _ConsumerPoolAllocation, - _ContiguousAllocator, ) from vllm.v1.core.sched.output import SchedulerOutput @@ -35,6 +97,19 @@ class CopyingFakeTransferEngine: + """Model Mooncake registration rules while copying bytes in-process. + + Attributes: + registered: Base addresses of currently registered ranges. + regions: Registered byte lengths keyed by base address. + register_calls: Address batches passed to memory registration. + unregister_calls: Addresses passed to single-range unregistration. + batch_unregister_calls: Address batches passed to unregistration. + transfer_calls: Byte lengths recorded for each transfer batch. + transfer_batches: Complete source and destination transfer arguments. + initialize_calls: Arguments used to initialize the fake engine. + """ + def __init__(self, *args, **kwargs): self.registered: set[int] = set() self.regions: dict[int, int] = {} @@ -42,8 +117,13 @@ def __init__(self, *args, **kwargs): self.unregister_calls: list[int] = [] self.batch_unregister_calls: list[list[int]] = [] self.transfer_calls: list[list[int]] = [] + self.transfer_batches: list[tuple[str, list[int], list[int], list[int]]] = [] + self.initialize_calls: list[tuple[str, str, str, str]] = [] def initialize(self, local_hostname, metadata_server, protocol, device_name) -> int: + self.initialize_calls.append( + (local_hostname, metadata_server, protocol, device_name) + ) return 0 def get_rpc_port(self) -> int: @@ -52,8 +132,14 @@ def get_rpc_port(self) -> int: def batch_transfer_sync_write( self, target_hostname, buffers, peer_buffer_addresses, lengths ) -> int: - self.transfer_calls.append([int(length) for length in lengths]) - for src, dst, nbytes in zip(buffers, peer_buffer_addresses, lengths): + sources = [int(address) for address in buffers] + destinations = [int(address) for address in peer_buffer_addresses] + sizes = [int(length) for length in lengths] + self.transfer_calls.append(sizes) + self.transfer_batches.append( + (str(target_hostname), sources, destinations, sizes) + ) + for src, dst, nbytes in zip(sources, destinations, sizes): ctypes.memmove(int(dst), int(src), int(nbytes)) return 0 @@ -110,9 +196,283 @@ def _wait_for_worker_io( raise TimeoutError("EC Mooncake worker I/O did not finish") +class TestECMooncakeControlPlane: + """Validate ZMQ client reuse, shard discovery, events, and server RPCs.""" + + def test_worker_get_ip_failure_does_not_construct_client( + self, mock_vllm_config_producer + ): + with ( + patch_ec_mooncake_deps(), + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake.worker.get_ip", + side_effect=RuntimeError("no address"), + ), + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "worker.ControlClient" + ) as client_cls, + pytest.raises(RuntimeError, match="no address"), + ): + ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.WORKER) + + client_cls.assert_not_called() + + def test_worker_constructor_failure_closes_client(self, mock_vllm_config_producer): + with ( + patch_ec_mooncake_deps(), + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "worker.ControlClient" + ) as client_cls, + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "worker.ThreadPoolExecutor", + side_effect=RuntimeError("executor failed"), + ), + pytest.raises(RuntimeError, match="executor failed"), + ): + ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.WORKER) + + client_cls.return_value.close.assert_called_once_with() + + def test_client_reuses_socket_and_discards_failed_exchange(self): + context = MagicMock() + socket = context.socket.return_value + socket.recv_json.side_effect = [ + {"ok": True, "result": {"ports": [19019]}}, + {"ok": True}, + {"ok": False, "error": "reservation rejected"}, + RuntimeError("timeout"), + ] + + with patch.object(control.zmq, "Context", return_value=context): + client = ControlClient(17) + assert client.request("tcp://consumer:19019", {"op": "peers"}) == { + "ports": [19019] + } + assert client.request("tcp://consumer:19019", {"op": "event_port"}) is None + with pytest.raises(RuntimeError, match="reservation rejected"): + client.request( + "tcp://consumer:19019", + {"op": "status", "transfer_id": "transfer"}, + ) + with pytest.raises(RuntimeError, match="timeout"): + client.request("tcp://consumer:19019", {"op": "event_port"}) + client.close() + client.close() + + context.socket.assert_called_once_with(zmq.REQ) + assert socket.setsockopt.call_args_list == [ + call(zmq.RCVTIMEO, 17), + call(zmq.SNDTIMEO, 17), + call(zmq.LINGER, 0), + ] + socket.connect.assert_called_once_with("tcp://consumer:19019") + socket.close.assert_called_once_with(linger=0) + context.destroy.assert_called_once_with(linger=0) + + def test_client_uses_one_socket_per_thread(self): + context = MagicMock() + sockets = [MagicMock(), MagicMock()] + for index, control_socket in enumerate(sockets): + control_socket.recv_json.return_value = {"ok": True, "result": index} + context.socket.side_effect = sockets + barrier = threading.Barrier(2) + results: list[int] = [] + + with patch.object(control.zmq, "Context", return_value=context): + client = ControlClient(20) + + def request() -> None: + barrier.wait() + results.append(client.request("tcp://consumer:19019", {"op": "peers"})) + + threads = [threading.Thread(target=request) for _ in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + client.close() + + assert sorted(results) == [0, 1] + assert context.socket.call_args_list == [call(zmq.REQ), call(zmq.REQ)] + for control_socket in sockets: + control_socket.connect.assert_called_once_with("tcp://consumer:19019") + + def test_topology_retries_transient_discovery_failures(self): + client = Mock(spec=ControlClient) + client.request.side_effect = [ + {"ports": [19019, 19020]}, + RuntimeError("old consumer"), + {"ports": [19029, 19030]}, + ] + topology = ShardTopology(client) + + assert topology.shards("tcp://consumer:19019") == [ + "tcp://consumer:19019", + "tcp://consumer:19020", + ] + assert topology.shards("tcp://consumer:19019") == [ + "tcp://consumer:19019", + "tcp://consumer:19020", + ] + assert topology.shards("tcp://legacy:19019") == ["tcp://legacy:19019"] + assert topology.shards("tcp://legacy:19019") == [ + "tcp://legacy:19029", + "tcp://legacy:19030", + ] + assert topology.shards("tcp://legacy:19019") == [ + "tcp://legacy:19029", + "tcp://legacy:19030", + ] + assert client.request.call_args_list == [ + call("tcp://consumer:19019", {"op": "peers"}), + call("tcp://legacy:19019", {"op": "peers"}), + call("tcp://legacy:19019", {"op": "peers"}), + ] + + def test_event_inbox_retries_until_every_shard_is_connected(self): + client = Mock(spec=ControlClient) + client.request.side_effect = [ + RuntimeError("peers not ready"), + {"ports": [19019, 19020]}, + 20001, + RuntimeError("event port not ready"), + 20001, + 20002, + ] + topology = ShardTopology(client) + context = MagicMock() + socket = context.socket.return_value + event = {"transfer_id": "transfer", "ready": True} + socket.recv_json.side_effect = [event, zmq.Again()] + + with patch.object(control.zmq, "Context", return_value=context) as create: + inbox = EventInbox(client, topology) + assert inbox.drain("tcp://consumer:19019") == [] + assert inbox.shard_count == 1 + create.assert_not_called() + assert inbox.drain("tcp://consumer:19019") == [] + assert inbox.shard_count == 1 + create.assert_not_called() + assert inbox.drain("tcp://consumer:19019") == [event] + assert inbox.shard_count == 2 + inbox.close() + inbox.close() + + assert client.request.call_args_list == [ + call("tcp://consumer:19019", {"op": "peers"}), + call("tcp://consumer:19019", {"op": "peers"}), + call("tcp://consumer:19019", {"op": "event_port"}), + call("tcp://consumer:19020", {"op": "event_port"}), + call("tcp://consumer:19019", {"op": "event_port"}), + call("tcp://consumer:19020", {"op": "event_port"}), + ] + create.assert_called_once_with() + assert socket.connect.call_args_list == [ + call("tcp://consumer:20001"), + call("tcp://consumer:20002"), + ] + assert socket.recv_json.call_args_list == [ + call(flags=zmq.DONTWAIT), + call(flags=zmq.DONTWAIT), + ] + socket.close.assert_called_once_with(linger=0) + context.term.assert_called_once_with() + + def test_server_preserves_wire_shapes_and_closes_twice(self): + port = _find_free_port() + completed: list[tuple[str, str]] = [] + cancelled: list[tuple[str, str, bool, bool]] = [] + + def status(transfer_id: str): + return {"transfer_id": transfer_id, "ready": False} + + def complete(transfer_id: str, reservation_id: str): + completed.append((transfer_id, reservation_id)) + return ControlCompletion(True, became_ready=True) + + def cancel( + transfer_id: str, + reservation_id: str, + abandon: bool, + refresh: bool, + ): + cancelled.append((transfer_id, reservation_id, abandon, refresh)) + return True + + server = ConsumerControlServer( + "127.0.0.1", + port, + reserve=lambda request: {"nbytes": request["nbytes"], "ready": False}, + status=status, + complete=complete, + cancel=cancel, + reap=lambda: 0, + peer_ports=[port, port + 1], + ) + client = ControlClient(1000) + server.start() + try: + addr = f"tcp://127.0.0.1:{port}" + assert client.request(addr, {"op": "peers"}) == {"ports": [port, port + 1]} + assert isinstance(client.request(addr, {"op": "event_port"}), int) + assert client.request( + addr, {"op": "status", "transfer_id": "transfer"} + ) == {"transfer_id": "transfer", "ready": False} + assert client.request( + addr, + { + "op": "reserve", + "transfer_id": "transfer", + "mm_hash": "hash", + "nbytes": 16, + "shape": [4], + "dtype": "float32", + }, + ) == {"nbytes": 16, "ready": False} + assert client.request( + addr, + { + "op": "complete", + "transfer_id": "transfer", + "reservation_id": "r0", + }, + ) == {"completed": True, "became_ready": True} + assert client.request( + addr, + { + "op": "complete_batch", + "items": [{"transfer_id": "transfer", "reservation_id": "r0"}], + }, + ) == {"items": [{"completed": True, "became_ready": True}]} + assert client.request( + addr, + { + "op": "cancel", + "transfer_id": "transfer", + "reservation_id": "r0", + "abandon": True, + }, + ) == {"cancelled": True} + finally: + client.close() + client.close() + server.close() + server.close() + + assert completed == [("transfer", "r0"), ("transfer", "r0")] + assert cancelled == [("transfer", "r0", True, False)] + + @pytest.fixture def mock_vllm_config_producer(): config = Mock(spec=VllmConfig) + config.model_config = Mock(spec=ModelConfig) + config.model_config.dtype = torch.float16 + config.model_config.hf_config = None + config.model_config.get_inputs_embeds_size.return_value = 16 config.parallel_config = Mock() config.parallel_config.tensor_parallel_size = 1 config.parallel_config.pipeline_parallel_size = 1 @@ -153,32 +513,186 @@ def mock_vllm_config_consumer(): def patch_ec_mooncake_deps(): with ( patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector.TransferEngine", + "vllm.distributed.ec_transfer.ec_connector.mooncake.transfer.TransferEngine", CopyingFakeTransferEngine, ), patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector._MOONCAKE_IMPORT_ERROR", + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "_availability._MOONCAKE_IMPORT_ERROR", None, ), patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector.get_ip", + "vllm.distributed.ec_transfer.ec_connector.mooncake.worker.get_ip", return_value="127.0.0.1", ), ): yield +class TestMooncakeTransfer: + """Validate lazy engine setup and source registration ownership.""" + + def test_initializes_engine_once_on_first_use(self): + engine = CopyingFakeTransferEngine() + with patch.object( + transfer, "TransferEngine", return_value=engine + ) as engine_cls: + data_plane = MooncakeTransfer("host", "tcp") + assert engine_cls.call_count == 0 + + assert data_plane.local_session() == "host:12345" + data_plane.ensure_ready() + + engine_cls.assert_called_once_with() + assert engine.initialize_calls == [("host", "P2PHANDSHAKE", "tcp", "")] + data_plane.close() + + def test_source_registration_is_refcounted_and_failed_release_keeps_owner(self): + engine = CopyingFakeTransferEngine() + data_plane = MooncakeTransfer("host", "tcp") + source = torch.randn(4, 4) + source_ref = weakref.ref(source) + with patch.object(transfer, "TransferEngine", return_value=engine): + first = data_plane.acquire_sources([source]) + second = data_plane.acquire_sources([source]) + assert engine.register_calls == [[source.data_ptr()]] + + assert data_plane.release_sources(first) + assert engine.batch_unregister_calls == [] + with patch.object( + engine, "batch_unregister_memory", return_value=1 + ) as unregister: + assert not data_plane.release_sources(second) + unregister.assert_called_once_with(second) + + del source + gc.collect() + assert source_ref() is not None + with patch.object( + engine, "batch_unregister_memory", return_value=0 + ) as unregister: + data_plane.close() + data_plane.close() + unregister.assert_called_once_with(second) + gc.collect() + assert source_ref() is None + + def test_failed_destination_unregister_is_retried_on_close(self): + engine = CopyingFakeTransferEngine() + data_plane = MooncakeTransfer("host", "tcp") + destination = torch.zeros(4, 4) + destination_ref = weakref.ref(destination) + with patch.object(transfer, "TransferEngine", return_value=engine): + assert data_plane.register_memory(destination) == 0 + with patch.object(engine, "unregister_memory", return_value=2): + assert not data_plane.unregister_memory(destination) + del destination + gc.collect() + assert destination_ref() is not None + + with patch.object( + engine, "batch_unregister_memory", return_value=0 + ) as unregister: + data_plane.close() + unregister.assert_called_once() + gc.collect() + assert destination_ref() is None + + def test_write_preserves_segments_and_reports_terminal_failure(self): + engine = CopyingFakeTransferEngine() + data_plane = MooncakeTransfer("host", "tcp") + sources = [torch.tensor([1, 2]), torch.tensor([3, 4])] + destinations = [torch.zeros_like(source) for source in sources] + source_addresses = [source.data_ptr() for source in sources] + destination_addresses = [tensor.data_ptr() for tensor in destinations] + lengths = [source.nbytes for source in sources] + with patch.object(transfer, "TransferEngine", return_value=engine): + data_plane.write("peer:1", source_addresses, destination_addresses, lengths) + assert engine.transfer_batches == [ + ("peer:1", source_addresses, destination_addresses, lengths) + ] + assert all( + torch.equal(source, destination) + for source, destination in zip(sources, destinations) + ) + + with ( + patch.object(engine, "batch_transfer_sync_write", return_value=9), + pytest.raises(RuntimeError, match="peer:2 failed with status 9"), + ): + data_plane.write( + "peer:2", source_addresses, destination_addresses, lengths + ) + data_plane.close() + + def test_write_returns_only_after_sync_engine_call_finishes(self): + engine = CopyingFakeTransferEngine() + data_plane = MooncakeTransfer("host", "tcp") + source = torch.ones(1, dtype=torch.uint8) + entered = threading.Event() + finish = threading.Event() + + def blocking_write(*args): + entered.set() + assert finish.wait(timeout=2) + return 0 + + with ( + patch.object(transfer, "TransferEngine", return_value=engine), + patch.object( + engine, "batch_transfer_sync_write", side_effect=blocking_write + ), + ): + addresses = data_plane.acquire_sources([source]) + completed = threading.Event() + + def write(): + try: + data_plane.write("peer:1", addresses, [2], [source.nbytes]) + finally: + data_plane.release_sources(addresses) + completed.set() + + thread = threading.Thread(target=write) + thread.start() + assert entered.wait(timeout=2) + assert not completed.is_set() + assert engine.batch_unregister_calls == [] + finish.set() + thread.join(timeout=2) + assert completed.is_set() + assert engine.batch_unregister_calls == [addresses] + data_plane.close() + + class TestECMooncakeFactory: + """Validate factory registration and compatibility exports.""" + def test_factory_registers_connector(self): cls = ECConnectorFactory.get_connector_class( Mock(ec_connector="ECMooncakeConnector") ) - assert cls.__name__ == "ECMooncakeConnector" + assert cls is ECMooncakeConnector + assert ( + cls.__module__ + == "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector" + ) + + def test_public_exports_are_compatible_and_narrow(self): + assert mooncake_ec_connector.__all__ == [ + "ECMooncakeConnector", + "ECMooncakeConnectorMetadata", + "ECMooncakeLoadSpec", + "ECMooncakePushSpec", + "ECMooncakeWorkerMetadata", + ] class TestContiguousAllocator: + """Validate aligned allocation, reuse, and range coalescing.""" + def test_reuses_and_coalesces_contiguous_regions(self): - allocator = _ContiguousAllocator(1024, alignment=256) + allocator = ContiguousAllocator(1024, alignment=256) first = allocator.allocate(1) second = allocator.allocate(300) @@ -190,8 +704,570 @@ def test_reuses_and_coalesces_contiguous_regions(self): allocator.free(*second) assert allocator.allocate(1024) == (0, 1024) + def test_splits_until_exhausted(self): + allocator = ContiguousAllocator(768, alignment=256) + + assert allocator.allocate(257) == (0, 512) + assert allocator.allocate(1) == (512, 256) + assert allocator.allocate(1) is None + + +class TestResidentPool: + """Validate resident pin, lease, replacement, and LRU semantics.""" + + def test_lru_skips_rejected_entry_and_replaces_without_losing_owner(self): + pool = ResidentPool[str]() + pool.insert("oldest", "first", 256) + pool.insert("next", "second", 256) + pool.retire("oldest") + pool.retire("next") + + evicted = pool.evict_lru(lambda key, _: key != "oldest") + + assert evicted == "next" + assert pool.get("oldest") == "first" + assert pool.insert("oldest", "replacement", 128) == "first" + assert pool.get("oldest") == "replacement" + assert pool.used == 128 + + def test_displaced_entry_waits_for_every_lease(self): + pool = ResidentPool[str]() + pool.insert("hash", "original", 256) + first = pool.acquire("hash") + second = pool.acquire("hash") + assert first is not None and second is not None + + assert pool.insert("hash", "replacement", 256) is None + assert pool.used == 512 + assert pool.release(first) is None + assert pool.release(second) == "original" + assert pool.used == 256 + assert pool.release(second) is None + + +class TestMooncakeMemoryPools: + """Validate Producer staging and Consumer residency ownership.""" + + class _Event: + """Minimal CUDA-event substitute controlling deferred frees.""" + + def __init__(self, complete: bool): + self.complete = complete + + def record(self, stream): + pass + + def query(self): + return self.complete + + def test_consumer_replacement_waits_for_cached_owner(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 0 + mooncake_transfer.unregister_memory.return_value = True + pool = ConsumerMemoryPool(768, mooncake_transfer) + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + first = pool.try_allocate(64, (16,), torch.float32) + replacement = pool.try_allocate(64, (16,), torch.float32) + assert first is not None and replacement is not None + pool.publish("hash", first) + held = pool.acquire_cached("hash", (16,), torch.float32) + assert held is not None + pool.publish("hash", replacement) + + third = pool.try_allocate(64, (16,), torch.float32) + assert third is not None + assert third.offset != first.offset + + pool.release_cached(held) + reused = pool.try_allocate(64, (16,), torch.float32) + assert reused is not None + assert reused.offset == first.offset + assert pool.take_resident("hash", (16,), "float32") is replacement.tensor + + def test_cached_consume_returns_newer_canonical_allocation(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 0 + pool = ConsumerMemoryPool(768, mooncake_transfer) + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + first = pool.try_allocate(64, (16,), torch.float32) + replacement = pool.try_allocate(64, (16,), torch.float32) + assert first is not None and replacement is not None + pool.publish("hash", first) + held = pool.acquire_cached("hash", (16,), torch.float32) + assert held is not None + pool.publish("hash", replacement) + + canonical = pool.publish("hash", held.value, held) + + assert canonical is replacement + reused = pool.try_allocate(64, (16,), torch.float32) + assert reused is not None + assert reused.offset == first.offset + + def test_consumer_defers_retired_reuse_until_event_completes(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 0 + pool = ConsumerMemoryPool(256, mooncake_transfer) + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + allocation = pool.try_allocate(64, (16,), torch.float32) + assert allocation is not None + pool.publish("hash", allocation) + event = self._Event(complete=False) + + with ( + patch.object(memory.torch, "Event", return_value=event), + patch.object(memory.torch.accelerator, "current_stream"), + ): + pool.retire_stale({}, set()) + assert pool.reclaim_and_allocate(64, (16,), torch.float32) is None + event.complete = True + reused = pool.try_allocate(64, (16,), torch.float32) + + assert reused is not None + assert reused.offset == allocation.offset + assert pool.drain_reclaimed() == {"hash"} + + def test_consumer_registration_failure_disables_pool(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 1 + pool = ConsumerMemoryPool(256, mooncake_transfer) + + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + + assert pool.tensor is None + mooncake_transfer.register_memory.assert_called_once() + + def test_nonreceiving_consumer_never_registers_or_unregisters_pool(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + pool = ConsumerMemoryPool(256, mooncake_transfer) + + pool.prepare(torch.device("cpu"), receiving_rank=False, allow_host=True) + pool.close() + pool.close() + + assert pool.tensor is None + mooncake_transfer.register_memory.assert_not_called() + mooncake_transfer.unregister_memory.assert_not_called() + + def test_consumer_close_unregisters_once_and_releases_parent(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 0 + mooncake_transfer.unregister_memory.return_value = True + pool = ConsumerMemoryPool(256, mooncake_transfer) + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + parent = pool.tensor + + pool.close() + pool.close() + + mooncake_transfer.unregister_memory.assert_called_once_with(parent) + assert pool.tensor is None + + def test_producer_reuses_staging_and_keeps_parent_for_later_close_phase(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 0 + pool = ProducerMemoryPool(256, mooncake_transfer) + source = torch.arange(16, dtype=torch.float32) + + first = pool.stage([source]) + assert first is not None + assert torch.equal(first.tensors[0], source) + pool.release(first) + second = pool.stage([source]) + assert second is not None + assert second.regions == first.regions + pool.release(second) + parent = pool.tensor + + pool.close() + pool.close() + + assert pool.tensor is parent + mooncake_transfer.unregister_memory.assert_not_called() + + def test_producer_falls_back_when_staging_pool_allocation_fails(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + pool = ProducerMemoryPool(256, mooncake_transfer) + + with patch.object(memory.torch, "empty", side_effect=torch.OutOfMemoryError): + assert pool.stage([torch.ones(16)]) is None + assert pool.stage([torch.ones(16)]) is None + + mooncake_transfer.register_memory.assert_not_called() + + +class TestMooncakeECConfig: + """Validate normalization, defaults, immutability, and bounds.""" + + def test_defaults_are_an_immutable_snapshot(self, mock_vllm_config_producer): + config = MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + + assert config == MooncakeECConfig( + is_producer=True, + is_consumer=False, + protocol="tcp", + buffer_device="cuda", + reservation_port=None, + reservation_addr=None, + control_timeout_s=30, + push_wait_timeout_s=60, + transfer_workers=4, + control_workers=8, + producer_pool_size=1_000_000_000, + consumer_pool_size=1_000_000_000, + transfer_metrics_log_interval=10, + consumer_metrics_log_interval=10, + ) + + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "mooncake_protocol" + ] = "rdma" + assert config.protocol == "tcp" + with pytest.raises(FrozenInstanceError): + config.protocol = "rdma" # type: ignore[misc] + + def test_custom_values_are_normalized(self, mock_vllm_config_consumer): + source = mock_vllm_config_consumer + source.parallel_config.tensor_parallel_size = 2 + source.parallel_config.data_parallel_size = 3 + source.parallel_config.data_parallel_index = 1 + source.ec_transfer_config.ec_buffer_device = " cpu " + source.ec_transfer_config.ec_buffer_size = 2048 + source.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": " tcp ", + "reservation_zmq_port": "5000", + "control_timeout_s": "1.5", + "push_wait_timeout_s": "2.5", + "transfer_max_workers": "3", + "control_max_workers": "5", + "producer_buffer_pool_size": "1024", + "consumer_buffer_pool_size": "1536", + "transfer_metrics_log_interval": "0", + "consumer_metrics_log_interval": "7", + } + + config = MooncakeECConfig.from_vllm_config(source, ECConnectorRole.WORKER) + + assert config.reservation_port == 5002 + assert config.reservation_addr == "tcp://127.0.0.1:5002" + assert config.control_timeout_s == 1.5 + assert config.push_wait_timeout_s == 2.5 + assert config.transfer_workers == 3 + assert config.control_workers == 5 + assert config.producer_pool_size == 1024 + assert config.consumer_pool_size == 1536 + assert config.transfer_metrics_log_interval == 0 + assert config.consumer_metrics_log_interval == 7 + assert config.buffer_device == "cpu" + assert config.protocol == "tcp" + + @pytest.mark.parametrize( + ("key", "value"), + [ + ("control_timeout_s", 0), + ("push_wait_timeout_s", -1), + ("transfer_max_workers", 0), + ("control_max_workers", -1), + ("producer_buffer_pool_size", 0), + ("consumer_buffer_pool_size", -1), + ], + ) + @pytest.mark.parametrize( + "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] + ) + def test_rejects_nonpositive_values( + self, mock_vllm_config_producer, key, value, role + ): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( + value + ) + + with pytest.raises(ValueError, match=key): + MooncakeECConfig.from_vllm_config(mock_vllm_config_producer, role) + + @pytest.mark.parametrize( + "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] + ) + def test_rejects_nonpositive_registered_buffer( + self, mock_vllm_config_producer, role + ): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 0 + + with pytest.raises(ValueError, match="ec_buffer_size > 0"): + MooncakeECConfig.from_vllm_config(mock_vllm_config_producer, role) + + @pytest.mark.parametrize("key", ["control_timeout_s", "push_wait_timeout_s"]) + @pytest.mark.parametrize("value", [float("nan"), float("inf"), float("-inf"), True]) + def test_rejects_invalid_timeouts(self, mock_vllm_config_producer, key, value): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( + value + ) + + with pytest.raises(ValueError, match=key): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + @pytest.mark.parametrize( + "key", + [ + "reservation_zmq_port", + "transfer_max_workers", + "control_max_workers", + "producer_buffer_pool_size", + "consumer_buffer_pool_size", + ], + ) + @pytest.mark.parametrize("value", [1.5, True]) + def test_rejects_noninteger_values(self, mock_vllm_config_producer, key, value): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( + value + ) + + with pytest.raises(ValueError, match=key): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + @pytest.mark.parametrize("value", [1.5, True]) + def test_rejects_noninteger_registered_buffer( + self, mock_vllm_config_producer, value + ): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = value + + with pytest.raises(ValueError, match="ec_buffer_size"): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + def test_accepts_integral_float_integer_values(self, mock_vllm_config_producer): + source = mock_vllm_config_producer + source.ec_transfer_config.ec_buffer_size = 8.0 + source.ec_transfer_config.ec_connector_extra_config.update( + { + "reservation_zmq_port": 5000.0, + "transfer_max_workers": 2.0, + "control_max_workers": 3.0, + "producer_buffer_pool_size": 4.0, + "consumer_buffer_pool_size": 5.0, + } + ) + + config = MooncakeECConfig.from_vllm_config(source, ECConnectorRole.WORKER) + + assert ( + config.reservation_port, + config.transfer_workers, + config.control_workers, + config.producer_pool_size, + config.consumer_pool_size, + ) == (5000, 2, 3, 4, 5) + + @pytest.mark.parametrize( + ("key", "value"), + [ + ("mooncake_protocol", ""), + ("mooncake_protocol", " "), + ("mooncake_protocol", None), + ("mooncake_protocol", 1), + ("reservation_zmq_addr", ""), + ("reservation_zmq_addr", " "), + ("reservation_zmq_addr", None), + ("reservation_zmq_addr", 1), + ], + ) + def test_rejects_invalid_strings(self, mock_vllm_config_producer, key, value): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( + value + ) + + with pytest.raises(ValueError, match=key): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + + @pytest.mark.parametrize("buffer_device", [None, "", " \t"]) + def test_normalizes_default_buffer_device( + self, mock_vllm_config_producer, buffer_device + ): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = buffer_device + + config = MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + assert config.buffer_device == "cuda" + + def test_strips_buffer_device(self, mock_vllm_config_producer): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = " cuda " + + config = MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + assert config.buffer_device == "cuda" + + @pytest.mark.parametrize("buffer_device", [1, True]) + def test_rejects_invalid_buffer_device( + self, mock_vllm_config_producer, buffer_device + ): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = buffer_device + + with pytest.raises(ValueError, match="ec_buffer_device"): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + @pytest.mark.parametrize( + "key", ["transfer_metrics_log_interval", "consumer_metrics_log_interval"] + ) + @pytest.mark.parametrize( + "value", [-1, float("nan"), float("inf"), float("-inf"), True, None] + ) + def test_rejects_invalid_metrics_intervals( + self, mock_vllm_config_producer, key, value + ): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( + value + ) + + with pytest.raises(ValueError, match=key): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + def test_zero_disables_metrics(self, mock_vllm_config_producer): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config.update( + { + "transfer_metrics_log_interval": 0, + "consumer_metrics_log_interval": 0, + } + ) + + config = MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + assert config.transfer_metrics_log_interval == 0 + assert config.consumer_metrics_log_interval == 0 + + def test_submillisecond_control_timeout_uses_one_millisecond( + self, mock_vllm_config_producer + ): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "control_timeout_s" + ] = 0.0001 + with ( + patch_ec_mooncake_deps(), + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "scheduler.ControlClient" + ) as scheduler_client, + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "worker.ControlClient" + ) as worker_client, + ): + scheduler = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + worker = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + scheduler.shutdown() + worker.shutdown() + + scheduler_client.assert_called_once_with(1) + worker_client.assert_called_once_with(1) + + @pytest.mark.parametrize("timeout", [1e308, sys.float_info.max]) + def test_rejects_control_timeout_too_large_for_zmq( + self, mock_vllm_config_producer, timeout + ): + mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ + "control_timeout_s" + ] = timeout + + with pytest.raises(ValueError, match="control_timeout_s"): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + + def test_rejects_pipeline_parallel_producer(self, mock_vllm_config_producer): + mock_vllm_config_producer.parallel_config.pipeline_parallel_size = 2 + + with pytest.raises(ValueError, match="pipeline parallelism"): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + + @pytest.mark.parametrize("port", [0, 65536]) + def test_rejects_out_of_range_base_port(self, mock_vllm_config_consumer, port): + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "reservation_zmq_port" + ] = port + + with pytest.raises(ValueError, match="1..65535"): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + + def test_rejects_topology_that_overflows_port_range( + self, mock_vllm_config_consumer + ): + source = mock_vllm_config_consumer + source.parallel_config.tensor_parallel_size = 4 + source.parallel_config.data_parallel_index = 1 + source.ec_transfer_config.ec_connector_extra_config["reservation_zmq_port"] = ( + 65530 + ) + + with pytest.raises(ValueError, match="ports must be in 1..65535"): + MooncakeECConfig.from_vllm_config(source, ECConnectorRole.SCHEDULER) + + def test_consumer_role_requirements_differ_only_by_process_role( + self, mock_vllm_config_consumer + ): + extra = {"reservation_zmq_addr": "tcp://consumer:19019"} + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = extra + + scheduler = MooncakeECConfig.from_vllm_config( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + assert scheduler.reservation_addr == "tcp://consumer:19019" + with pytest.raises(ValueError, match="workers require reservation_zmq_port"): + MooncakeECConfig.from_vllm_config( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + class TestECMooncakeConnectorValidation: + """Validate role construction, optional dependencies, and topology rules.""" + + @pytest.mark.parametrize( + "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] + ) + def test_requires_transfer_engine_symbol_for_each_role( + self, mock_vllm_config_producer, monkeypatch, role + ): + from vllm.distributed.ec_transfer.ec_connector.mooncake import _availability + + fake_package = ModuleType("mooncake") + fake_package.__path__ = [] + fake_engine = ModuleType("mooncake.engine") + try: + with monkeypatch.context() as context: + context.setitem(sys.modules, "mooncake", fake_package) + context.setitem(sys.modules, "mooncake.engine", fake_engine) + importlib.reload(_availability) + with pytest.raises(ImportError, match="mooncake-transfer-engine"): + ECMooncakeConnector(mock_vllm_config_producer, role) + finally: + importlib.reload(_availability) + def test_rejects_sharded_producer(self, mock_vllm_config_producer): """One copy of each encoder output, so sharding only duplicates it.""" mock_vllm_config_producer.parallel_config.tensor_parallel_size = 2 @@ -237,8 +1313,11 @@ def test_replicated_consumer_addresses_its_own_block( connector = ECMooncakeConnector(cfg, ECConnectorRole.SCHEDULER) try: # Replica 2 of a TP=2 consumer starts after two 2-port blocks. - assert connector._control_port_offset == 4 - assert connector._reservation_zmq_addr == "tcp://127.0.0.1:19504" + assert connector._scheduler is not None + assert ( + connector._scheduler._reservation_zmq_addr + == "tcp://127.0.0.1:19504" + ) finally: connector.shutdown() @@ -252,8 +1331,297 @@ def test_rejects_replicated_producer(self, mock_vllm_config_producer): ): ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.SCHEDULER) + def test_scheduler_hooks_route_exactly(self, mock_vllm_config_producer): + scheduler = Mock() + scheduler.take_unavailable_requests.return_value = {"unavailable"} + scheduler.has_cache_item.return_value = True + scheduler.ensure_cache_available.return_value = False + scheduler.build_connector_meta.return_value = "metadata" + scheduler.has_pending_push_work.return_value = True + scheduler.request_finished.return_value = (True, {"result": 1}) + with patch.object( + ECMooncakeScheduler, + "from_vllm_config", + return_value=scheduler, + ) as from_vllm_config: + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + request = Mock() + scheduler_output = Mock() + connector_output = Mock() + try: + assert connector.take_unavailable_requests() == {"unavailable"} + assert connector.has_cache_item("hash") is True + assert connector.ensure_cache_available(request, 7, {"local"}) is False + connector.update_state_after_alloc(request, 2) + connector.update_state_after_free(request, 3) + assert connector.build_connector_meta(scheduler_output) == "metadata" + connector.update_connector_output(connector_output) + assert connector.has_pending_push_work() is True + assert connector.request_finished(request) == (True, {"result": 1}) + finally: + connector.shutdown() + + from_vllm_config.assert_called_once_with(mock_vllm_config_producer) + scheduler.take_unavailable_requests.assert_called_once_with() + scheduler.has_cache_item.assert_called_once_with("hash") + scheduler.ensure_cache_available.assert_called_once_with(request, 7, {"local"}) + scheduler.update_state_after_alloc.assert_called_once_with(request, 2) + scheduler.update_state_after_free.assert_called_once_with(request, 3) + scheduler.build_connector_meta.assert_called_once_with(scheduler_output) + scheduler.update_connector_output.assert_called_once_with(connector_output) + scheduler.has_pending_push_work.assert_called_once_with() + scheduler.request_finished.assert_called_once_with(request) + scheduler.close.assert_called_once_with() + + def test_worker_hooks_route_exactly(self, mock_vllm_config_producer): + metadata = ECMooncakeConnectorMetadata() + worker = Mock() + worker.get_finished.return_value = ({"saved"}, {"loaded"}) + worker.build_connector_worker_meta.return_value = "worker-metadata" + with patch.object( + ECMooncakeWorker, + "from_vllm_config", + return_value=worker, + ) as from_vllm_config: + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + connector.bind_connector_metadata(metadata) + encoder_cache: dict[str, torch.Tensor] = {} + try: + connector.start_worker_services() + connector.start_save_caches(encoder_cache=encoder_cache, marker=1) + connector.start_load_caches(encoder_cache, marker=2) + connector.save_caches(encoder_cache, "hash", marker=3) + assert connector.get_finished({"finished"}) == ( + {"saved"}, + {"loaded"}, + ) + assert connector.build_connector_worker_meta() == "worker-metadata" + finally: + connector.shutdown() + + from_vllm_config.assert_called_once_with(mock_vllm_config_producer) + worker.start_services.assert_called_once_with() + worker.start_save_caches.assert_called_once_with( + metadata, encoder_cache=encoder_cache, marker=1 + ) + worker.start_load_caches.assert_called_once_with( + metadata, encoder_cache, marker=2 + ) + worker.save_caches.assert_called_once_with(encoder_cache, "hash", marker=3) + worker.get_finished.assert_called_once_with({"finished"}) + worker.build_connector_worker_meta.assert_called_once_with() + worker.close.assert_called_once_with() + + @pytest.mark.parametrize( + ("method", "args"), + [ + ("start_save_caches", ()), + ("start_load_caches", ({},)), + ], + ) + def test_worker_load_and_save_reject_wrong_metadata( + self, mock_vllm_config_producer, method, args + ): + class OtherMetadata(ECConnectorMetadata): + """Represent an incompatible connector metadata implementation.""" + + pass + + worker = Mock() + with patch.object( + ECMooncakeWorker, + "from_vllm_config", + return_value=worker, + ): + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + connector.bind_connector_metadata(OtherMetadata()) + try: + with pytest.raises(AssertionError): + getattr(connector, method)(*args) + finally: + connector.shutdown() + + getattr(worker, method).assert_not_called() + + @pytest.mark.parametrize( + ("method", "args", "kwargs"), + [ + ("start_worker_services", (), {}), + ("start_save_caches", (), {}), + ("start_load_caches", ({},), {}), + ("save_caches", ({}, "hash"), {}), + ("get_finished", (set(),), {}), + ("build_connector_worker_meta", (), {}), + ], + ) + def test_scheduler_rejects_worker_hooks( + self, mock_vllm_config_producer, method, args, kwargs + ): + with patch.object( + ECMooncakeScheduler, + "from_vllm_config", + return_value=Mock(), + ): + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + try: + with pytest.raises(AssertionError): + getattr(connector, method)(*args, **kwargs) + finally: + connector.shutdown() + + @pytest.mark.parametrize( + ("method", "args", "kwargs"), + [ + ("take_unavailable_requests", (), {}), + ("has_cache_item", ("hash",), {}), + ("ensure_cache_available", (Mock(), 0, set()), {}), + ("update_state_after_alloc", (Mock(), 0), {}), + ("update_state_after_free", (Mock(), 0), {}), + ("build_connector_meta", (Mock(),), {}), + ("update_connector_output", (Mock(),), {}), + ("has_pending_push_work", (), {}), + ("request_finished", (Mock(),), {}), + ], + ) + def test_worker_rejects_scheduler_hooks( + self, mock_vllm_config_producer, method, args, kwargs + ): + with patch.object( + ECMooncakeWorker, + "from_vllm_config", + return_value=Mock(), + ): + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + try: + with pytest.raises(AssertionError): + getattr(connector, method)(*args, **kwargs) + finally: + connector.shutdown() + + @pytest.mark.parametrize( + ("role", "active", "inactive"), + [ + (ECConnectorRole.SCHEDULER, "_scheduler", "_worker"), + (ECConnectorRole.WORKER, "_worker", "_scheduler"), + ], + ) + def test_exactly_one_role_and_idempotent_shutdown( + self, mock_vllm_config_producer, role, active, inactive + ): + scheduler = Mock() + worker = Mock() + with ( + patch.object( + ECMooncakeScheduler, + "from_vllm_config", + return_value=scheduler, + ), + patch.object( + ECMooncakeWorker, + "from_vllm_config", + return_value=worker, + ), + ): + connector = ECMooncakeConnector(mock_vllm_config_producer, role) + assert getattr(connector, active) is not None + assert getattr(connector, inactive) is None + assert set(connector.__dict__) - { + "_connector_metadata", + "_vllm_config", + "_role", + "_is_producer", + "_is_consumer", + } == {"_scheduler", "_worker", "_closed"} + connector.shutdown() + connector.shutdown() + + if role == ECConnectorRole.SCHEDULER: + scheduler.close.assert_called_once_with() + worker.close.assert_not_called() + else: + worker.close.assert_called_once_with() + scheduler.close.assert_not_called() + + def test_rejects_unknown_role(self, mock_vllm_config_producer): + invalid_role = Mock(name="invalid_role") + with pytest.raises(ValueError, match="Unknown EC connector role"): + ECMooncakeConnector(mock_vllm_config_producer, invalid_role) + + def test_del_is_best_effort(self): + connector = object.__new__(ECMooncakeConnector) + with patch.object( + ECMooncakeConnector, + "shutdown", + side_effect=RuntimeError("shutdown failed"), + ) as shutdown: + connector.__del__() + shutdown.assert_called_once_with() + + +class TestECMooncakeMetadata: + """Validate metadata compatibility, pickling, and aggregation inputs.""" + + def test_old_imports_reexport_packaged_metadata(self): + assert ECMooncakeLoadSpec is metadata.ECMooncakeLoadSpec + assert ECMooncakePushSpec is metadata.ECMooncakePushSpec + assert ECMooncakeConnectorMetadata is metadata.ECMooncakeConnectorMetadata + assert ECMooncakeWorkerMetadata is metadata.ECMooncakeWorkerMetadata + + @pytest.mark.parametrize( + "metadata", + [ + ECMooncakeConnectorMetadata( + loads=[ + ECMooncakeLoadSpec( + mm_hash="load", + num_token=2, + nbytes=8, + shape=(2, 4), + dtype="float16", + pushed=True, + transfer_id="transfer", + reservation_id="reservation", + local=True, + ) + ], + pushes=[ + ECMooncakePushSpec( + mm_hash="push", + nbytes=8, + shape=(2, 4), + dtype="float16", + consumer_zmq="tcp://127.0.0.1:1234", + transfer_id="transfer", + request_id="request", + ) + ], + ), + ECMooncakeWorkerMetadata( + loaded={"loaded"}, + failed_loads={"failed"}, + reclaimed={"reclaimed"}, + pending_loads=True, + pending_saves=True, + ), + ], + ) + def test_metadata_pickle_round_trip(self, metadata): + assert ForkingPickler.loads(ForkingPickler.dumps(metadata)) == metadata + class TestECMooncakeWorkerMetadataAggregation: + """Validate cross-rank success intersection and failure union rules.""" + def test_an_item_one_rank_missed_is_not_loaded(self): """Each rank gathers from its own cache, so all of them must have it. @@ -276,8 +1644,223 @@ def test_a_reclaim_on_any_rank_invalidates_residency(self): assert merged.reclaimed == {"c"} -class TestECMooncakeSchedulerMetadata: - def test_missing_push_event_is_tracked( +class TestSchedulerTransferTable: + """Validate Scheduler transfer transitions, indexes, and retention.""" + + @staticmethod + def pushed_spec(transfer_id: str, mm_hash: str = "hash") -> ECMooncakeLoadSpec: + return ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=0, + nbytes=16, + shape=(4,), + dtype="float32", + pushed=True, + transfer_id=transfer_id, + reservation_id=f"reservation-{transfer_id}", + ) + + def test_legal_load_and_resident_reload_use_authoritative_record(self): + assert SchedulerTransferTable is state.SchedulerTransferTable + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + record, accepted = table.observe_ready(self.pushed_spec("transfer"), 10) + + assert accepted and record.state is SchedulerTransferState.AVAILABLE + assert table.begin_load("hash", 7, "transfer", "request") is record + assert table.take_loads_to_dispatch() == [record] + assert record.spec is not None and record.spec.num_token == 7 + assert table.complete_load("hash") + table.release_ready("hash", 1) + assert record.state is SchedulerTransferState.RESIDENT + assert table.begin_load("hash", 9) is record + assert record.spec is not None and record.spec.local + + def test_illegal_transition_is_rejected(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + record, _ = table.observe_ready(self.pushed_spec("transfer"), 10) + + with pytest.raises(InvalidSchedulerTransferTransition): + table.mark_unavailable("transfer", "late", 1) + assert record.state is SchedulerTransferState.AVAILABLE + + def test_same_hash_index_preserves_transfer_order_and_identity(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + first, _ = table.observe_ready(self.pushed_spec("first"), 10) + second, _ = table.observe_ready(self.pushed_spec("second"), 10) + + assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [ + first, + second, + ] + assert table.begin_load("hash", 3) is first + assert ( + table.first_for_hash("hash", (SchedulerTransferState.AVAILABLE,)) is second + ) + with pytest.raises(ValueError): + table.observe_ready(self.pushed_spec("first", "other-hash"), 10) + + def test_unavailable_notification_drains_once_and_rejects_late_ready(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + record = table.wait_for_event("transfer", "request-r", "hash", 1) + table.mark_unavailable("transfer", "timed out", 2) + + assert table.take_unavailable_requests() == {"request-r"} + assert table.take_unavailable_requests() == set() + table.wait_for_event("transfer", "request-r", "hash", 3) + assert table.take_unavailable_requests() == set() + table.wait_for_event("transfer", "request-n", "hash", 3) + assert table.take_unavailable_requests() == {"request-n"} + assert table.take_unavailable_requests() == set() + table.wait_for_event("transfer", "request-n", "hash", 3) + assert table.take_unavailable_requests() == set() + same, accepted = table.observe_ready(self.pushed_spec("transfer"), 40) + assert same is record and not accepted + assert record.state is SchedulerTransferState.UNAVAILABLE + + def test_cancel_and_duplicate_completion_are_idempotent(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + cancelled = table.wait_for_event("cancelled", "request", "hash", 10) + assert table.cancel("cancelled", 1) + assert not table.cancel("cancelled", 2) + _, accepted = table.observe_ready(self.pushed_spec("cancelled"), 40) + assert not accepted and cancelled.state is SchedulerTransferState.CANCELLED + + record, _ = table.observe_ready(self.pushed_spec("completed", "other"), 10) + table.begin_load("other", 4, "completed") + assert table.complete_load("other") + assert table.complete_load("other") + assert record.state is SchedulerTransferState.READY + + def test_failed_record_expires_from_record_and_hash_index(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + record, _ = table.observe_ready(self.pushed_spec("failed"), 10) + table.begin_load("hash", 4, "failed") + + assert table.fail_load("hash", "copy failed", 20) + assert record.deadline == 50 + _, dropped = table.expire(51, terminal_limit=100) + assert dropped == 1 + assert table.get("failed") is None + assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [] + + def test_reclaimed_resident_tombstone_expires(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + record, _ = table.observe_ready(self.pushed_spec("reclaimed"), 10) + table.begin_load("hash", 4, "reclaimed") + table.complete_load("hash") + table.release_ready("hash", 20) + + table.reclaim("hash", 30) + assert record.state is SchedulerTransferState.EXPIRED + assert record.deadline == 60 + table.expire(61, terminal_limit=100) + assert table.get("reclaimed") is None + assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [] + + def test_capacity_eviction_tombstone_expires(self): + table = SchedulerTransferTable(resident_capacity=0, tombstone_ttl=30) + record, _ = table.observe_ready(self.pushed_spec("evicted"), 10) + table.begin_load("hash", 4, "evicted") + table.complete_load("hash") + + table.release_ready("hash", 20) + assert record.state is SchedulerTransferState.EXPIRED + assert record.deadline == 50 + table.expire(51, terminal_limit=100) + assert table.get("evicted") is None + assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [] + + def test_terminal_record_limit_prunes_oldest_records(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + for transfer_id in ("first", "second", "third"): + table.cancel(transfer_id, 1) + + _, dropped = table.expire(2, terminal_limit=1) + assert dropped == 2 + assert table.get("first") is None + assert table.get("second") is None + assert table.get("third") is not None + + def test_zero_terminal_limit_prunes_every_record(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + table.cancel("first", 1) + table.cancel("second", 1) + + _, dropped = table.expire(2, terminal_limit=0) + assert dropped == 2 + assert table.get("first") is None + assert table.get("second") is None + + def test_negative_terminal_limit_is_rejected_without_mutation(self): + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + record = table.wait_for_event("transfer", "request", "hash", 1) + + with pytest.raises(ValueError, match="terminal_limit"): + table.expire(2, terminal_limit=-1) + assert table.get("transfer") is record + assert record.state is SchedulerTransferState.WAITING_EVENT + + def test_same_hash_residency_uses_only_the_latest_completed_record(self): + table = SchedulerTransferTable(resident_capacity=32, tombstone_ttl=30) + first, _ = table.observe_ready(self.pushed_spec("first"), 10) + table.begin_load("hash", 4, "first") + table.complete_load("hash") + table.release_ready("hash", 20) + second, _ = table.observe_ready(self.pushed_spec("second"), 30) + table.begin_load("hash", 4, "second") + table.complete_load("hash") + table.release_ready("hash", 40) + third, _ = table.observe_ready(self.pushed_spec("third", "other"), 50) + table.begin_load("other", 4, "third") + table.complete_load("other") + table.release_ready("other", 60) + + assert first.state is SchedulerTransferState.EXPIRED + assert second.state is SchedulerTransferState.RESIDENT + assert third.state is SchedulerTransferState.RESIDENT + assert table.resident_bytes == 32 + + +class TestECMooncakeSchedulerMetadata: + """Validate Scheduler decisions and per-step Worker metadata.""" + + def test_cancel_confirms_topology_and_retries_only_failed_shards(self): + scheduler = object.__new__(ECMooncakeScheduler) + scheduler._topology = Mock(spec=ShardTopology) + scheduler._topology.discover.side_effect = [ + None, + ["shard-0", "shard-1", "shard-2"], + ] + scheduler._control_client = Mock(spec=ControlClient) + called = [] + + def request(addr, _payload): + called.append(addr) + if addr == "shard-0" and called.count(addr) == 1: + raise RuntimeError("cancel shard failed") + return {"cancelled": True} + + scheduler._control_client.request.side_effect = request + assert scheduler._cancel_remote("base", "transfer", "reservation") + assert scheduler._topology.discover.call_args_list == [ + call("base"), + call("base"), + ] + assert called == ["shard-0", "shard-1", "shard-2", "shard-0"] + + def test_cancel_rejects_unconfirmed_topology_without_sending(self): + scheduler = object.__new__(ECMooncakeScheduler) + scheduler._topology = Mock(spec=ShardTopology) + scheduler._topology.discover.return_value = None + scheduler._control_client = Mock(spec=ControlClient) + + with pytest.raises(RuntimeError, match="discover every EC consumer shard"): + scheduler._cancel_remote("base", "transfer", "reservation") + + assert scheduler._topology.discover.call_count == 2 + scheduler._control_client.request.assert_not_called() + + def test_missing_push_event_is_tracked( self, mock_vllm_config_consumer, mock_request_with_3_mm ): mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { @@ -292,11 +1875,15 @@ def test_missing_push_event_is_tracked( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - with patch.object(scheduler, "_drain_push_notifications"): + with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert not scheduler.ensure_cache_available(request, 0) - mm_hash = request.mm_features[0].identifier - assert scheduler._consumer_scheduler_metrics["missing_event"] == 1 - assert mm_hash in scheduler._consumer_missing_since + assert ( + scheduler._scheduler._consumer_scheduler_metrics["missing_event"] + == 1 + ) + record = scheduler._scheduler._transfers.get(f"{request.request_id}:0") + assert record is not None + assert record.state is SchedulerTransferState.WAITING_EVENT finally: scheduler.shutdown() @@ -307,26 +1894,36 @@ def test_item_with_no_transfer_in_flight_is_reported_as_stalled( mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", "reservation_zmq_port": 19019, - "push_wait_timeout_s": 0, + "push_wait_timeout_s": 0.001, } request = mock_request_with_3_mm request.mm_features = request.mm_features[:1] - mm_hash = request.mm_features[0].identifier - with patch_ec_mooncake_deps(): scheduler = ECMooncakeConnector( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - with patch.object(scheduler, "_drain_push_notifications"): + with ( + patch.object(scheduler._scheduler, "_drain_push_notifications"), + patch( + "vllm.distributed.ec_transfer.ec_connector." + "mooncake.scheduler.time.monotonic", + side_effect=[10, 10.002, 10.003], + ), + ): assert not scheduler.ensure_cache_available(request, 0) assert not scheduler.ensure_cache_available(request, 0) - assert scheduler._consumer_scheduler_metrics["stalled"] == 1 - assert mm_hash in scheduler._stalled_hashes - # The stall is reported once, not once per scheduling pass. - with patch.object(scheduler, "_drain_push_notifications"): + assert ( + scheduler._scheduler._consumer_scheduler_metrics["stalled"] == 1 + ) + record = scheduler._scheduler._transfers.get( + f"{request.request_id}:0" + ) + assert record is not None + assert record.state is SchedulerTransferState.UNAVAILABLE + # The stall is reported once, not once per scheduling pass. assert not scheduler.ensure_cache_available(request, 0) - assert scheduler._consumer_scheduler_metrics["stalled"] == 1 + assert scheduler._scheduler._consumer_scheduler_metrics["stalled"] == 1 finally: scheduler.shutdown() @@ -350,11 +1947,38 @@ def test_pending_observation_ends_with_last_spec(self, mock_vllm_config_consumer ) try: for spec in specs: - scheduler._index_pending_spec(spec) - scheduler._pop_pending_spec("transfer-0") - assert "hash" in scheduler._consumer_pending_since - scheduler._pop_pending_spec("transfer-1") - assert "hash" not in scheduler._consumer_pending_since + scheduler._scheduler._transfers.observe_ready(spec, 10) + scheduler._scheduler._transfers.cancel("transfer-0", 0) + available = scheduler._scheduler._transfers.records_for_hash( + "hash", (SchedulerTransferState.AVAILABLE,) + ) + assert [record.transfer_id for record in available] == ["transfer-1"] + scheduler._scheduler._transfers.cancel("transfer-1", 0) + assert ( + scheduler._scheduler._transfers.first_for_hash( + "hash", (SchedulerTransferState.AVAILABLE,) + ) + is None + ) + finally: + scheduler.shutdown() + + def test_available_expiry_is_cancelled_before_tombstone_cleanup( + self, mock_vllm_config_consumer + ): + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + record, _ = scheduler._scheduler._transfers.observe_ready( + TestSchedulerTransferTable.pushed_spec("expired"), 0 + ) + try: + with patch.object(scheduler._scheduler, "_queue_cancel") as cancel: + scheduler._scheduler._expire_transfers() + cancel.assert_called_once_with("expired") + assert record.state is SchedulerTransferState.EXPIRED + assert scheduler._scheduler._transfers.get("expired") is record finally: scheduler.shutdown() @@ -379,7 +2003,7 @@ def test_local_cache_hit_keeps_the_transfer( {"mm_hash": mm_hash, "transfer_id": "request-transfer"} ] } - scheduler._index_pending_spec( + scheduler._scheduler._transfers.observe_ready( ECMooncakeLoadSpec( mm_hash=mm_hash, num_token=0, @@ -389,19 +2013,27 @@ def test_local_cache_hit_keeps_the_transfer( pushed=True, transfer_id="request-transfer", reservation_id="reservation", - ) + ), + 10, ) with ( - patch.object(scheduler, "_drain_push_notifications"), - patch.object(scheduler, "_queue_cancel") as cancel, + patch.object(scheduler._scheduler, "_drain_push_notifications"), + patch.object(scheduler._scheduler, "_queue_cancel") as cancel, ): assert scheduler.ensure_cache_available(request, 0, {mm_hash}) cancel.assert_not_called() - assert "request-transfer" in scheduler._pending_specs + record = scheduler._scheduler._transfers.get("request-transfer") + assert record is not None + assert record.state is SchedulerTransferState.AVAILABLE # Once the entry is evicted the request can still get it. assert not scheduler.ensure_cache_available(request, 0, set()) - assert mm_hash in scheduler._loading_hashes + assert ( + scheduler._scheduler._transfers.first_for_hash( + mm_hash, (SchedulerTransferState.LOADING,) + ) + is not None + ) finally: scheduler.shutdown() @@ -425,7 +2057,7 @@ def test_consumed_item_releases_its_transfer_immediately( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._index_pending_spec( + scheduler._scheduler._transfers.observe_ready( ECMooncakeLoadSpec( mm_hash=mm_hash, num_token=0, @@ -435,14 +2067,13 @@ def test_consumed_item_releases_its_transfer_immediately( pushed=True, transfer_id="consumed-transfer", reservation_id="reservation", - ) + ), + 10, ) - with patch.object(scheduler, "_queue_cancel") as cancel: - scheduler.update_state_after_free(request, 0) - # Cancelled by transfer: a shard's reservation id means - # nothing to its peers, so it is not passed along. - cancel.assert_called_once_with("consumed-transfer") - assert "consumed-transfer" not in scheduler._pending_specs + scheduler.update_state_after_free(request, 0) + record = scheduler._scheduler._transfers.get("consumed-transfer") + assert record is not None + assert record.state is SchedulerTransferState.CANCELLED finally: scheduler.shutdown() @@ -480,27 +2111,47 @@ def test_ready_hash_eviction_does_not_strand_a_later_transfer( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._ready_hashes.add(mm_hash) - scheduler._event_zmq_socket = Mock() - scheduler._event_zmq_socket.recv_json.side_effect = [ - event, - zmq.Again(), - ] - with patch.object(scheduler, "_queue_cancel") as cancel: - scheduler._drain_push_notifications() + current = ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=0, + nbytes=16, + shape=(4,), + dtype="float32", + pushed=True, + transfer_id="current-transfer", + reservation_id="current", + ) + scheduler._scheduler._transfers.observe_ready( + current, time.monotonic() + _LEASE_TTL_SECONDS + ) + scheduler._scheduler._transfers.begin_load( + mm_hash, 4, "current-transfer" + ) + scheduler._scheduler._transfers.take_loads_to_dispatch() + scheduler._scheduler._transfers.complete_load(mm_hash) + scheduler._scheduler._event_inbox.drain = Mock(return_value=[event]) + with patch.object(scheduler._scheduler, "_queue_cancel") as cancel: + scheduler._scheduler._drain_push_notifications() cancel.assert_not_called() - assert "later-transfer" in scheduler._pending_specs + later = scheduler._scheduler._transfers.get("later-transfer") + assert later is not None + assert later.state is SchedulerTransferState.AVAILABLE # The scheduler frees the encoder cache entry. scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[mm_hash]) ) - assert mm_hash not in scheduler._ready_hashes + assert ( + scheduler._scheduler._transfers.first_for_hash( + mm_hash, (SchedulerTransferState.READY,) + ) + is None + ) # The request that owns the transfer can still pick it up. - with patch.object(scheduler, "_drain_push_notifications"): + with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert not scheduler.ensure_cache_available(request, 0, set()) - assert mm_hash in scheduler._loading_hashes + assert later.state is SchedulerTransferState.LOADING finally: scheduler.shutdown() @@ -525,22 +2176,46 @@ def test_cancelled_transfer_ignores_late_ready_events( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._event_shard_count = len(ports) - scheduler._cancelled_transfer_ids[transfer_id] = ( - time.monotonic() + _LEASE_TTL_SECONDS + scheduler._scheduler._event_inbox.shard_count = len(ports) + scheduler._scheduler._transfers.cancel( + transfer_id, time.monotonic(), mm_hash="hash" + ) + scheduler._scheduler._event_inbox.drain = Mock( + return_value=[{**event, "shard": port} for port in ports] ) - scheduler._event_zmq_socket = Mock() - scheduler._event_zmq_socket.recv_json.side_effect = [ - {**event, "shard": port} for port in ports - ] + [zmq.Again()] - scheduler._drain_push_notifications() + scheduler._scheduler._drain_push_notifications() - assert transfer_id not in scheduler._pending_specs - assert transfer_id not in scheduler._event_ready_shards - assert scheduler._consumer_scheduler_metrics["events_cancelled"] == ( - len(ports) - ) + record = scheduler._scheduler._transfers.get(transfer_id) + assert record is not None + assert record.state is SchedulerTransferState.CANCELLED + assert transfer_id not in scheduler._scheduler._event_ready_shards + assert scheduler._scheduler._consumer_scheduler_metrics[ + "events_cancelled" + ] == len(ports) + finally: + scheduler.shutdown() + + def test_cancel_rpc_failure_keeps_tombstone_and_rejects_late_ready( + self, mock_vllm_config_consumer + ): + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + spec = TestSchedulerTransferTable.pushed_spec("transfer") + record, _ = scheduler._scheduler._transfers.observe_ready(spec, 10) + scheduler._scheduler._transfers.cancel("transfer", 1) + failed = Mock() + failed.done.return_value = True + failed.result.side_effect = RuntimeError("unknown remote result") + scheduler._scheduler._pending_cancels["transfer"] = failed + try: + scheduler._scheduler._poll_pending_cancels() + assert record.state is SchedulerTransferState.CANCELLED + assert record.spec is spec + same, accepted = scheduler._scheduler._transfers.observe_ready(spec, 20) + assert same is record and not accepted finally: scheduler.shutdown() @@ -572,34 +2247,41 @@ def test_cancel_between_shards_drops_the_partial_readiness( } def deliver(scheduler, *shards): - scheduler._event_zmq_socket.recv_json.side_effect = [ + scheduler._scheduler._event_inbox.drain.return_value = [ {**event, "shard": shard} for shard in shards - ] + [zmq.Again()] - scheduler._drain_pending = True - scheduler._drain_push_notifications() + ] + scheduler._scheduler._drain_pending = True + scheduler._scheduler._drain_push_notifications() with patch_ec_mooncake_deps(): scheduler = ECMooncakeConnector( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" - scheduler._event_shard_count = 4 - scheduler._event_zmq_socket = Mock() + scheduler._scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" + scheduler._scheduler._event_inbox.shard_count = 4 + scheduler._scheduler._event_inbox.drain = Mock() deliver(scheduler, 0, 1) - assert scheduler._event_ready_shards[transfer_id] == {0, 1} + assert scheduler._scheduler._event_ready_shards[transfer_id] == {0, 1} - with patch.object(scheduler, "_cancel_remote", return_value=True): + with patch.object( + scheduler._scheduler, "_cancel_remote", return_value=True + ): scheduler.update_state_after_free(request, 0) - assert transfer_id in scheduler._cancelled_transfer_ids - assert transfer_id not in scheduler._event_ready_shards + record = scheduler._scheduler._transfers.get(transfer_id) + assert record is not None + assert record.state is SchedulerTransferState.CANCELLED + assert transfer_id not in scheduler._scheduler._event_ready_shards deliver(scheduler, 2, 3) - assert transfer_id not in scheduler._pending_specs - assert transfer_id not in scheduler._event_ready_shards - assert scheduler._consumer_scheduler_metrics["events_cancelled"] == 2 + assert record.state is SchedulerTransferState.CANCELLED + assert transfer_id not in scheduler._scheduler._event_ready_shards + assert ( + scheduler._scheduler._consumer_scheduler_metrics["events_cancelled"] + == 2 + ) finally: scheduler.shutdown() @@ -615,37 +2297,46 @@ def test_cancelled_transfer_ids_stay_bounded(self, mock_vllm_config_consumer): mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" - scheduler._event_zmq_socket = Mock() - scheduler._event_zmq_socket.recv_json.side_effect = zmq.Again() - with patch.object(scheduler, "_cancel_remote", return_value=True): + scheduler._scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" + scheduler._scheduler._event_inbox.drain = Mock(return_value=[]) + with patch.object( + scheduler._scheduler, "_cancel_remote", return_value=True + ): for name in ("first", "second", "third"): - scheduler._queue_cancel(name) + scheduler._scheduler._queue_cancel(name) now = time.monotonic() - assert scheduler._cancelled_transfer_ids["third"] > now - assert scheduler._cancelled_transfer_ids["third"] <= ( - now + _LEASE_TTL_SECONDS - ) + third = scheduler._scheduler._transfers.get("third") + assert third is not None and third.deadline is not None + assert third.deadline > now + assert third.deadline <= now + _LEASE_TTL_SECONDS # Ignored for exactly as long as the worker refuses to reserve # the id again, and no longer. The drain is what sweeps. - scheduler._cancelled_transfer_ids["first"] = 0.0 - scheduler._drain_pending = True - scheduler._drain_push_notifications() - assert list(scheduler._cancelled_transfer_ids) == ["second", "third"] + first = scheduler._scheduler._transfers.get("first") + assert first is not None + first.deadline = 0.0 + scheduler._scheduler._drain_pending = True + scheduler._scheduler._drain_push_notifications() + assert scheduler._scheduler._transfers.get("first") is None + assert scheduler._scheduler._transfers.get("second") is not None + assert scheduler._scheduler._transfers.get("third") is third # The count is the backstop for a rate that outruns the TTL. with patch( "vllm.distributed.ec_transfer.ec_connector." - "mooncake_ec_connector._MAX_CANCELLED_TRANSFER_IDS", + "mooncake.scheduler._MAX_TERMINAL_TRANSFER_RECORDS", 1, ): - scheduler._drain_pending = True - scheduler._drain_push_notifications() - assert list(scheduler._cancelled_transfer_ids) == ["third"] + scheduler._scheduler._drain_pending = True + scheduler._scheduler._drain_push_notifications() + assert scheduler._scheduler._transfers.get("second") is None + assert scheduler._scheduler._transfers.get("third") is third assert ( - scheduler._consumer_scheduler_metrics["cancel_records_dropped"] == 2 + scheduler._scheduler._consumer_scheduler_metrics[ + "cancel_records_dropped" + ] + == 2 ) finally: scheduler.shutdown() @@ -662,28 +2353,36 @@ def test_item_that_never_arrives_fails_the_request( mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", "reservation_zmq_port": 19019, - "push_wait_timeout_s": 0, + "push_wait_timeout_s": 0.001, } request = mock_request_with_3_mm request.mm_features = request.mm_features[:1] - mm_hash = request.mm_features[0].identifier - with patch_ec_mooncake_deps(): scheduler = ECMooncakeConnector( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: with ( - patch.object(scheduler, "_drain_push_notifications"), - patch.object(scheduler, "_send_control", return_value=None), + patch.object(scheduler._scheduler, "_drain_push_notifications"), + patch.object( + scheduler._scheduler._control_client, + "request", + return_value=None, + ), + patch( + "vllm.distributed.ec_transfer.ec_connector." + "mooncake.scheduler.time.monotonic", + side_effect=[10, 10.002], + ), ): assert not scheduler.ensure_cache_available(request, 0, set()) + assert not scheduler.ensure_cache_available(request, 0, set()) assert scheduler.take_unavailable_requests() == {request.request_id} # Draining clears it: the scheduler acts on each id once. assert scheduler.take_unavailable_requests() == set() - # A re-issued request gets a fresh window rather than the - # expired one, or it would fail before its push could land. - assert mm_hash not in scheduler._consumer_missing_since + record = scheduler._scheduler._transfers.get(f"{request.request_id}:0") + assert record is not None + assert record.state is SchedulerTransferState.UNAVAILABLE finally: scheduler.shutdown() @@ -711,11 +2410,15 @@ def fake_send(addr: str, request: dict): mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._reservation_zmq_addr = f"tcp://127.0.0.1:{ports[0]}" + scheduler._scheduler._reservation_zmq_addr = ( + f"tcp://127.0.0.1:{ports[0]}" + ) with patch.object( - scheduler, "_send_control", side_effect=fake_send + scheduler._scheduler._control_client, + "request", + side_effect=fake_send, ) as send_control: - scheduler._ensure_event_channel() + scheduler._scheduler._drain_push_notifications() subscribed = [ call.args[0] @@ -723,16 +2426,24 @@ def fake_send(addr: str, request: dict): if call.args[1]["op"] == "event_port" ] assert len(subscribed) == len(ports) - assert scheduler._event_shard_count == len(ports) + assert scheduler._scheduler._event_shard_count == len(ports) event = {"transfer_id": "transfer-0"} - assert not scheduler._note_shard_ready({**event, "shard": ports[0]}) + assert not scheduler._scheduler._note_shard_ready( + {**event, "shard": ports[0]} + ) # The same rank reporting twice is not two ranks. - assert not scheduler._note_shard_ready({**event, "shard": ports[0]}) - assert not scheduler._note_shard_ready({**event, "shard": ports[1]}) - assert scheduler._note_shard_ready({**event, "shard": ports[2]}) + assert not scheduler._scheduler._note_shard_ready( + {**event, "shard": ports[0]} + ) + assert not scheduler._scheduler._note_shard_ready( + {**event, "shard": ports[1]} + ) + assert scheduler._scheduler._note_shard_ready( + {**event, "shard": ports[2]} + ) # Nothing is retained once the transfer is handed on. - assert "transfer-0" not in scheduler._event_ready_shards + assert "transfer-0" not in scheduler._scheduler._event_ready_shards finally: scheduler.shutdown() @@ -766,20 +2477,20 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._event_zmq_socket = Mock() - scheduler._event_zmq_socket.recv_json.side_effect = [ - { - "mm_hash": mm_hash, - "transfer_id": "only-transfer", - "ready": True, - "reservation_id": "r0", - "nbytes": 16, - "shape": [4], - "dtype": "float32", - }, - zmq.Again(), - ] - scheduler._drain_push_notifications() + scheduler._scheduler._event_inbox.drain = Mock( + return_value=[ + { + "mm_hash": mm_hash, + "transfer_id": "only-transfer", + "ready": True, + "reservation_id": "r0", + "nbytes": 16, + "shape": [4], + "dtype": "float32", + } + ] + ) + scheduler._scheduler._drain_push_notifications() assert not scheduler.ensure_cache_available(first, 0, set()) meta = scheduler.build_connector_meta( @@ -793,7 +2504,9 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( ) ) ) - assert not scheduler._pending_specs + record = scheduler._scheduler._transfers.get("only-transfer") + assert record is not None + assert record.state is SchedulerTransferState.READY # The encoder cache evicts the entry. scheduler.build_connector_meta( @@ -802,10 +2515,10 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( # The second request has no transfer of its own, and the only # transfer is spent. It must still be served. - with patch.object(scheduler, "_drain_push_notifications"): + with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert scheduler.has_cache_item(mm_hash) assert not scheduler.ensure_cache_available(second, 0, set()) - assert mm_hash in scheduler._loading_hashes + assert record.state is SchedulerTransferState.LOADING reload = scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[]) ) @@ -831,16 +2544,21 @@ def test_reclaimed_item_stops_being_offered_as_resident( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._note_resident( - ECMooncakeLoadSpec( - mm_hash=mm_hash, - num_token=0, - nbytes=16, - shape=(4,), - dtype="float32", - ) + spec = ECMooncakeLoadSpec( + mm_hash=mm_hash, + num_token=0, + nbytes=16, + shape=(4,), + dtype="float32", + transfer_id="transfer", ) - with patch.object(scheduler, "_drain_push_notifications"): + table = scheduler._scheduler._transfers + table.observe_ready(spec, time.monotonic() + _LEASE_TTL_SECONDS) + table.begin_load(mm_hash, 4, "transfer") + table.take_loads_to_dispatch() + table.complete_load(mm_hash) + table.release_ready(mm_hash, time.monotonic()) + with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert scheduler.has_cache_item(mm_hash) scheduler.update_connector_output( @@ -850,9 +2568,48 @@ def test_reclaimed_item_stops_being_offered_as_resident( ) ) ) - with patch.object(scheduler, "_drain_push_notifications"): + with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert not scheduler.has_cache_item(mm_hash) - assert scheduler._resident_bytes == 0 + assert table.resident_bytes == 0 + finally: + scheduler.shutdown() + + def test_reclaim_keeps_ready_cache_visible_until_it_is_freed( + self, mock_vllm_config_consumer + ): + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + table = scheduler._scheduler._transfers + spec = ECMooncakeLoadSpec( + mm_hash="hash", + num_token=0, + nbytes=16, + shape=(4,), + dtype="float32", + transfer_id="transfer", + ) + try: + table.observe_ready(spec, time.monotonic() + _LEASE_TTL_SECONDS) + table.begin_load("hash", 4, "transfer") + table.take_loads_to_dispatch() + table.complete_load("hash") + + scheduler.update_connector_output( + SimpleNamespace( + ec_connector_worker_meta=ECMooncakeWorkerMetadata( + reclaimed={"hash"} + ) + ) + ) + with patch.object(scheduler._scheduler, "_drain_push_notifications"): + assert scheduler.has_cache_item("hash") + scheduler.build_connector_meta( + SimpleNamespace(free_encoder_mm_hashes=["hash"]) + ) + with patch.object(scheduler._scheduler, "_drain_push_notifications"): + assert not scheduler.has_cache_item("hash") finally: scheduler.shutdown() @@ -877,13 +2634,24 @@ def test_retains_new_completion_while_same_hash_is_loading( scheduler = ECMooncakeConnector( mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) - scheduler._event_zmq_socket = Mock() - scheduler._event_zmq_socket.recv_json.side_effect = [event, zmq.Again()] - scheduler._loading_hashes.add("hash") + scheduler._scheduler._event_inbox.drain = Mock(return_value=[event]) + current = ECMooncakeLoadSpec( + mm_hash="hash", + num_token=0, + nbytes=64, + shape=(2, 8), + dtype="float32", + transfer_id="current-transfer", + ) + scheduler._scheduler._transfers.observe_ready(current, time.monotonic() + 1) + scheduler._scheduler._transfers.begin_load("hash", 2, "current-transfer") - scheduler._drain_push_notifications() + scheduler._scheduler._drain_push_notifications() - assert scheduler._pending_specs["next-transfer"].reservation_id == "next" + pending = scheduler._scheduler._transfers.get("next-transfer") + assert pending is not None and pending.spec is not None + assert pending.state is SchedulerTransferState.AVAILABLE + assert pending.spec.reservation_id == "next" def test_build_connector_meta_clears_pending( self, mock_vllm_config_consumer, mock_request_with_3_mm @@ -901,9 +2669,10 @@ def test_build_connector_meta_clears_pending( dtype="float32", transfer_id="transfer", ) - scheduler._index_pending_spec(load_spec) - scheduler._load_specs[mm_hash] = load_spec - scheduler._mm_datas_need_loads[mm_hash] = 100 + scheduler._scheduler._transfers.observe_ready( + load_spec, time.monotonic() + _LEASE_TTL_SECONDS + ) + scheduler._scheduler._transfers.begin_load(mm_hash, 100, "transfer") meta = scheduler.build_connector_meta( Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) ) @@ -911,8 +2680,10 @@ def test_build_connector_meta_clears_pending( assert len(meta.loads) == 1 assert meta.loads[0].mm_hash == mm_hash assert meta.loads[0].num_token == 100 - assert scheduler._mm_datas_need_loads == {} - assert "transfer" not in scheduler._pending_specs + assert scheduler._scheduler._transfers.take_loads_to_dispatch() == [] + record = scheduler._scheduler._transfers.get("transfer") + assert record is not None + assert record.state is SchedulerTransferState.LOADING def test_producer_does_not_build_load_metadata( self, mock_vllm_config_producer, mock_request_with_3_mm @@ -943,99 +2714,1617 @@ def test_producer_builds_push_metadata_after_preprocessing( mock_vllm_config_producer.model_config.get_inputs_embeds_size.return_value = 16 with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - scheduler.update_state_after_alloc(request, 0) - meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + scheduler = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + scheduler.update_state_after_alloc(request, 0) + meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) + + # The same request remains visible on a later scheduler step, but + # its worker push metadata must not be emitted a second time. + assert scheduler.ensure_cache_available(request, 0) + next_meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) + + scheduler.request_finished(request) + + assert meta.loads == [] + assert meta.pushes == [ + ECMooncakePushSpec( + mm_hash="img_hash_1", + nbytes=100 * 16 * 4, + shape=(100, 16), + dtype="float32", + consumer_zmq="tcp://decode:19019", + transfer_id="transfer-1", + request_id="test_req_123", + ) + ] + assert next_meta.pushes == [] + assert "transfer-1" not in scheduler._scheduler._prepared_push_transfer_ids + + def test_producer_uses_deepstack_encoder_cache_width( + self, mock_vllm_config_producer, mock_request_with_3_mm + ): + request = mock_request_with_3_mm + request.ec_transfer_params = { + "consumer_zmq": "tcp://decode:19019", + "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], + } + mock_vllm_config_producer.model_config.dtype = torch.bfloat16 + mock_vllm_config_producer.model_config.hf_config = SimpleNamespace( + vision_config=SimpleNamespace( + out_hidden_size=2560, + deepstack_visual_indexes=[5, 11, 17], + ) + ) + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + scheduler.update_state_after_alloc(request, 0) + meta = scheduler.build_connector_meta( + Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + ) + + num_tokens = request.get_num_encoder_embeds(0) + spec = meta.pushes[0] + assert spec.shape == (num_tokens, 10240) + assert spec.nbytes == num_tokens * 10240 * torch.bfloat16.itemsize + + def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer): + feature = SimpleNamespace( + identifier="image_uuid", + modality="image", + data=SimpleNamespace( + get_data=lambda: { + "image_grid_thw": torch.tensor([1, 32, 48]), + "pixel_values": torch.ones(2), + } + ), + ) + request = SimpleNamespace(mm_features=[feature]) + + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.SCHEDULER + ) + with patch.object( + scheduler._scheduler, + "_placeholder_metadata_fields", + return_value={"image_grid_thw"}, + ): + delay_free, params = scheduler.request_finished(request) + + assert not delay_free + assert params == { + "ec_items": [{"mm_hash": "image_uuid", "image_grid_thw": [1, 32, 48]}] + } + + +class TestConsumerReservationManager: + """Validate Consumer destination ownership and cancellation races.""" + + @staticmethod + def manager(): + pool = Mock() + pool.lock = threading.RLock() + pool.acquire_cached.return_value = None + allocation = memory.MemoryAllocation(0, 64, torch.empty(16)) + pool.try_allocate.return_value = allocation + pool.reclaim_and_allocate.return_value = None + return ConsumerReservationManager(pool, 300, 16), pool, allocation + + @staticmethod + def reserve(manager: ConsumerReservationManager): + record, write, reused, _ = manager.reserve( + "transfer", "hash", 64, (16,), "float32", torch.float32 + ) + assert record is not None + return record, write, reused + + def test_writing_ready_and_repeated_completion_use_one_state_record(self): + manager, _, _ = self.manager() + record, write, _ = self.reserve(manager) + + assert write + assert record.state is ConsumerReservationState.WRITING + assert manager.status("transfer") is record + completed = manager.complete("transfer", record.reservation_id) + repeated = manager.complete("transfer", record.reservation_id) + + assert completed.accepted and completed.became_ready + assert repeated.accepted and repeated.repeated + assert record.state is ConsumerReservationState.READY + with pytest.raises(RuntimeError): + manager._transition(record, ConsumerReservationState.WRITING) + + def test_writing_cancel_defers_the_only_allocation_release(self): + manager, pool, allocation = self.manager() + record, _, _ = self.reserve(manager) + + assert manager.cancel("transfer", "wrong-id") == ( + CancellationOutcome.REJECTED, + 0, + ) + outcome, dropped = manager.cancel("transfer", record.reservation_id) + assert outcome is CancellationOutcome.DEFERRED + assert dropped == 0 + assert record.state is ConsumerReservationState.CANCEL_PENDING + pool.free.assert_not_called() + + completed = manager.complete("transfer", record.reservation_id) + repeated = manager.complete("transfer", record.reservation_id) + assert completed.accepted and completed.discarded + assert not repeated.accepted + assert record.state is ConsumerReservationState.CANCELLED + assert record.allocation is None + pool.free.assert_called_once_with(allocation) + + def test_ready_expiry_releases_once_and_keeps_a_tombstone(self): + manager, pool, allocation = self.manager() + record, _, _ = self.reserve(manager) + manager.complete("transfer", record.reservation_id) + record.expires_at = 0 + + first_expired, _, _ = manager.expire() + second_expired, _, _ = manager.expire() + + assert first_expired == 1 + assert second_expired == 0 + assert manager.status("transfer") is None + assert record.state is ConsumerReservationState.EXPIRED + pool.free.assert_called_once_with(allocation) + + def test_failed_allocation_returns_deferred_and_tombstone_counts(self): + manager, pool, _ = self.manager() + writing, _, _ = self.reserve(manager) + writing.expires_at = 1.5 + assert manager.cancel("stale", "") == ( + CancellationOutcome.PRE_RESERVED, + 0, + ) + manager.get("stale").expires_at = 0 + pool.try_allocate.side_effect = [None, None] + + monotonic = ( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "reservation.time.monotonic" + ) + with patch(monotonic, side_effect=[1.0, 2.0]): + record, write, reused, counts = manager.reserve( + "new", "new-hash", 64, (16,), "float32", torch.float32 + ) + + assert record is None and not write and not reused + assert counts == (0, 1, 1) + assert writing.state is ConsumerReservationState.EXPIRE_PENDING + assert manager.get("stale") is None + pool.free.assert_not_called() + + def test_expired_writer_refresh_precedes_re_reserve_and_old_completion(self): + manager, pool, old_allocation = self.manager() + new_allocation = memory.MemoryAllocation(256, 64, torch.ones(16)) + pool.try_allocate.side_effect = [old_allocation, new_allocation] + old, _, _ = self.reserve(manager) + old.expires_at = 0 + _, deferred, _ = manager.expire() + + with pytest.raises(RuntimeError, match="still has an active writer"): + self.reserve(manager) + assert old.allocation is old_allocation + pool.free.assert_not_called() + + (refreshed, dropped) = manager.cancel( + "transfer", old.reservation_id, abandon=True, refresh=True + ) + new, write, _ = self.reserve(manager) + late = manager.complete("transfer", old.reservation_id) + + assert deferred == 1 + assert refreshed is CancellationOutcome.CANCELLED + assert dropped == 0 + assert write and new.reservation_id != old.reservation_id + assert new.state is ConsumerReservationState.WRITING + assert new.allocation is new_allocation + assert not late.accepted + pool.free.assert_called_once_with(old_allocation) + + def test_expired_writer_single_slot_is_reused_only_after_refresh_abandon(self): + transfer_engine = Mock() + transfer_engine.register_memory.return_value = 0 + pool = ConsumerMemoryPool(256, transfer_engine) + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + manager = ConsumerReservationManager(pool, 300, 16) + + old, _, _, _ = manager.reserve( + "transfer", "hash", 64, (16,), "float32", torch.float32 + ) + assert old is not None + assert old.allocation is not None + old_offset = old.allocation.offset + old.expires_at = 0 + manager.expire() + + with pytest.raises(RuntimeError, match="still has an active writer"): + manager.reserve("transfer", "hash", 64, (16,), "float32", torch.float32) + assert pool.try_allocate(64, (16,), torch.float32) is None + + outcome, dropped = manager.cancel( + "transfer", old.reservation_id, abandon=True, refresh=True + ) + new, write, _, _ = manager.reserve( + "transfer", "hash", 64, (16,), "float32", torch.float32 + ) + assert new is not None + assert outcome is CancellationOutcome.CANCELLED + assert dropped == 0 + assert write and new.reservation_id != old.reservation_id + assert old.allocation is None + assert new.allocation is not None + assert new.allocation.offset == old_offset == 0 + + def test_expired_writer_completion_releases_before_re_reserve(self): + manager, pool, old_allocation = self.manager() + new_allocation = memory.MemoryAllocation(256, 64, torch.ones(16)) + pool.try_allocate.side_effect = [old_allocation, new_allocation] + old, _, _ = self.reserve(manager) + old.expires_at = 0 + manager.expire() + + completed = manager.complete("transfer", old.reservation_id) + new, write, _ = self.reserve(manager) + + assert completed.accepted and completed.discarded + assert old.state is ConsumerReservationState.EXPIRED + assert old.allocation is None + assert write and new.allocation is new_allocation + assert new.reservation_id != old.reservation_id + pool.free.assert_called_once_with(old_allocation) + + def test_expired_writer_cancel_stays_deferred_until_completion(self): + manager, pool, allocation = self.manager() + record, _, _ = self.reserve(manager) + record.expires_at = 0 + _, deferred, _ = manager.expire() + + cancelled, dropped = manager.cancel("transfer", record.reservation_id) + assert deferred == 1 + assert cancelled is CancellationOutcome.DEFERRED + assert dropped == 0 + assert record.state is ConsumerReservationState.EXPIRE_PENDING + assert record.allocation is allocation + pool.free.assert_not_called() + + completed = manager.complete("transfer", record.reservation_id) + assert completed.accepted and completed.discarded + assert record.state is ConsumerReservationState.EXPIRED + pool.free.assert_called_once_with(allocation) + + def test_expired_writer_abandon_releases_once(self): + manager, pool, allocation = self.manager() + record, _, _ = self.reserve(manager) + record.expires_at = 0 + manager.expire() + + abandoned, first_dropped = manager.cancel( + "transfer", record.reservation_id, abandon=True + ) + repeated, second_dropped = manager.cancel( + "transfer", record.reservation_id, abandon=True + ) + + assert abandoned is CancellationOutcome.CANCELLED + assert repeated is CancellationOutcome.PRE_RESERVED + assert first_dropped == second_dropped == 0 + assert record.state is ConsumerReservationState.CANCELLED + assert record.allocation is None + pool.free.assert_called_once_with(allocation) + + def test_cached_take_returns_the_memory_pool_canonical_allocation(self): + manager, pool, cached = self.manager() + lease = SimpleNamespace(value=cached) + canonical = memory.MemoryAllocation(256, 64, torch.ones(16)) + pool.acquire_cached.return_value = lease + pool.publish.return_value = canonical + + record, write, _ = self.reserve(manager) + assert not write and record.lease is lease + taken = manager.take("transfer", "hash") + + assert taken is canonical + assert record.state is ConsumerReservationState.RESIDENT + assert record.allocation is None and record.lease is None + pool.publish.assert_called_once_with("hash", cached, lease) + pool.free.assert_not_called() + pool.release_cached.assert_not_called() + + def test_tombstone_indexes_reap_by_prefix_without_scanning_records(self): + manager, _, _ = self.manager() + + class NoScanDict(dict): + """Fail if reservation code scans the complete record mapping.""" + + def __iter__(self): + raise AssertionError("record table must not be scanned") + + def items(self): + raise AssertionError("record table must not be scanned") + + def values(self): + raise AssertionError("record table must not be scanned") + + manager._tombstone_limit = 3 + manager._records = NoScanDict(manager._records) + for transfer_id in ("a", "b", "c"): + assert manager.cancel(transfer_id, "") == ( + CancellationOutcome.PRE_RESERVED, + 0, + ) + assert manager.cancel("a", "") == (CancellationOutcome.PRE_RESERVED, 0) + assert manager.cancel("d", "") == (CancellationOutcome.PRE_RESERVED, 1) + assert list(manager._tombstones) == ["c", "a", "d"] + assert set(manager._records.keys()) == {"c", "a", "d"} + + manager.get("c").expires_at = 0 + _, _, dropped = manager.expire() + assert dropped == 1 + assert list(manager._tombstones) == ["a", "d"] + assert set(manager._records.keys()) == {"a", "d"} + assert not manager._active_ids + + def test_active_index_tracks_reserve_complete_take_and_expiry(self): + manager, pool, allocation = self.manager() + pool.publish.return_value = allocation + record, _, _ = self.reserve(manager) + assert list(manager._active_ids) == ["transfer"] + assert not manager._tombstones + + manager.complete("transfer", record.reservation_id) + manager.take("transfer", "hash") + assert manager.get("transfer") is None + assert not manager._active_ids and not manager._tombstones + + replacement, _, _ = self.reserve(manager) + manager.complete("transfer", replacement.reservation_id) + replacement.expires_at = 0 + manager.expire() + assert not manager._active_ids + assert list(manager._tombstones) == ["transfer"] + assert manager.get("transfer").state is ConsumerReservationState.EXPIRED + + +class TestECMooncakeWorkerTransfer: + """Validate end-to-end Worker reservation, push, load, and cleanup flows.""" + + def test_allocation_retry_accounts_for_expiry_after_outer_sweep(self): + transfer_engine = Mock() + transfer_engine.register_memory.return_value = 0 + transfer_engine.local_session.return_value = "local-session" + pool = ConsumerMemoryPool(256, transfer_engine) + pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + manager = ConsumerReservationManager(pool, 300, 16) + old, _, _, counts = manager.reserve( + "old", "old-hash", 64, (16,), "float32", torch.float32 + ) + assert counts == (0, 0, 0) + assert old is not None + assert old.allocation is not None + old_offset = old.allocation.offset + manager.complete("old", old.reservation_id) + old.expires_at = 1.5 + + worker = object.__new__(ECMooncakeWorker) + worker._consumer_worker_metrics = Counter() + worker._reservations = manager + worker._transfer = transfer_engine + payload = { + "transfer_id": "replacement", + "mm_hash": "replacement-hash", + "nbytes": 64, + "shape": [16], + "dtype": "float32", + } + monotonic = ( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "reservation.time.monotonic" + ) + with patch(monotonic, side_effect=[1.0, 1.0, 2.0, 2.0]): + replacement = worker._reserve_push_destination(payload) + + assert replacement["write"] + assert manager.get("old").state is ConsumerReservationState.EXPIRED + assert manager.get("replacement").allocation.offset == old_offset == 0 + assert worker._consumer_worker_metrics["reservations_expired"] == 1 + assert worker._consumer_worker_metrics["cancellations_deferred"] == 0 + assert worker._consumer_worker_metrics["cancel_records_dropped"] == 0 + + def test_failed_allocation_still_accounts_inner_expiry(self): + pool = Mock() + pool.lock = threading.RLock() + pool.acquire_cached.return_value = None + ready_allocation = memory.MemoryAllocation(0, 64, torch.empty(16)) + writing_allocation = memory.MemoryAllocation(64, 64, torch.empty(16)) + pool.try_allocate.side_effect = [ready_allocation, writing_allocation] + pool.reclaim_and_allocate.return_value = None + manager = ConsumerReservationManager(pool, 300, 16) + ready, _, _, _ = manager.reserve( + "ready", "ready-hash", 64, (16,), "float32", torch.float32 + ) + writing, _, _, _ = manager.reserve( + "writing", "writing-hash", 64, (16,), "float32", torch.float32 + ) + assert ready is not None and writing is not None + manager.complete("ready", ready.reservation_id) + assert manager.cancel("stale", "") == ( + CancellationOutcome.PRE_RESERVED, + 0, + ) + ready.expires_at = writing.expires_at = 1.5 + manager.get("stale").expires_at = 1.5 + pool.try_allocate.side_effect = [None, None] + + worker = object.__new__(ECMooncakeWorker) + worker._consumer_worker_metrics = Counter() + worker._reservations = manager + payload = { + "transfer_id": "failed", + "mm_hash": "failed-hash", + "nbytes": 64, + "shape": [16], + "dtype": "float32", + } + monotonic = ( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "reservation.time.monotonic" + ) + with patch(monotonic, side_effect=[1.0, 1.0, 2.0, 2.0, 2.0]): + with pytest.raises(RuntimeError, match="^EC consumer buffer pool is full$"): + worker._reserve_push_destination(payload) + metrics = dict(worker._consumer_worker_metrics) + worker._expire_push_reservations() + + assert metrics == { + "reservations_expired": 1, + "cancellations_deferred": 1, + "cancel_records_dropped": 1, + } + assert dict(worker._consumer_worker_metrics) == metrics + assert ready.state is ConsumerReservationState.EXPIRED + assert writing.state is ConsumerReservationState.EXPIRE_PENDING + assert manager.get("stale") is None + pool.free.assert_called_once_with(ready_allocation) + + def test_stale_shards_are_abandoned_before_remote_re_reserve(self): + worker = object.__new__(ECMooncakeWorker) + worker._control_client = Mock() + events: list[tuple[str, str] | tuple[str]] = [] + + def request(addr, payload): + events.append(("abandon", payload["reservation_id"])) + assert payload["abandon"] and payload["refresh"] + return {"cancelled": True} + + worker._control_client.request.side_effect = request + replacement = [{"reservation_id": "new"}] + + def reserve_remote(spec): + events.append(("reserve",)) + return replacement + + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:19019", + transfer_id="transfer", + ) + shards = [ + { + "addr": f"tcp://consumer:{19019 + rank}", + "reservation_id": f"old-{rank}", + "ready": False, + } + for rank in range(2) + ] + + with ( + ThreadPoolExecutor(max_workers=2) as executor, + patch.object(worker, "_shard_executor", return_value=executor), + patch.object(worker, "_reserve_remote", side_effect=reserve_remote), + ): + assert worker._refresh_remote_reservations(spec, shards) is replacement + assert set(events[:2]) == { + ("abandon", "old-0"), + ("abandon", "old-1"), + } + assert events[2] == ("reserve",) + + def test_cancel_retry_only_retries_failed_shards(self): + worker = object.__new__(ECMooncakeWorker) + worker._control_client = Mock() + attempts: Counter[str] = Counter() + + def request(_addr, payload): + reservation_id = payload["reservation_id"] + attempts[reservation_id] += 1 + if reservation_id == "r0" and attempts[reservation_id] == 1: + raise RuntimeError("transient cancel failure") + return {"cancelled": True} + + worker._control_client.request.side_effect = request + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:19019", + transfer_id="transfer", + ) + reservations = [ + { + "addr": f"tcp://consumer:{19019 + rank}", + "reservation_id": f"r{rank}", + } + for rank in range(3) + ] + + with ( + ThreadPoolExecutor(max_workers=2) as executor, + patch.object(worker, "_shard_executor", return_value=executor), + ): + worker._retry_cancel_reservations(spec, reservations) + + assert attempts == Counter({"r0": 2, "r1": 1, "r2": 1}) + + def test_partial_refresh_cleans_only_the_failed_shard_and_keeps_first_error( + self, + ): + worker = object.__new__(ECMooncakeWorker) + worker._control_client = Mock() + calls: list[tuple[str, bool]] = [] + + def request(_addr, payload): + reservation_id = payload["reservation_id"] + refreshing = payload.get("refresh", False) + calls.append((reservation_id, refreshing)) + if refreshing and reservation_id == "r0": + raise RuntimeError("refresh shard failed") + return {"cancelled": True} + + worker._control_client.request.side_effect = request + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:19019", + transfer_id="transfer", + ) + reservations = [ + { + "addr": f"tcp://consumer:{19019 + rank}", + "reservation_id": f"r{rank}", + "ready": False, + } + for rank in range(3) + ] + + with ( + ThreadPoolExecutor(max_workers=2) as executor, + patch.object(worker, "_shard_executor", return_value=executor), + pytest.raises(RuntimeError, match="^refresh shard failed$"), + ): + worker._refresh_remote_reservations(spec, reservations) + + assert Counter(calls) == Counter( + {("r0", True): 1, ("r1", True): 1, ("r2", True): 1, ("r0", False): 1} + ) + + def test_reservation_snapshot_and_resident_retirement_are_atomic(self): + worker = object.__new__(ECMooncakeWorker) + worker._resolve_consumer_rank = Mock() + worker._is_receiving_rank = True + worker._transfer = Mock() + worker._buffer_device = "cpu" + worker._consumer_memory = ConsumerMemoryPool(256, Mock()) + worker._reservations = ConsumerReservationManager( + worker._consumer_memory, _LEASE_TTL_SECONDS, 16 + ) + retire_entered = threading.Event() + finish_retire = threading.Event() + lock_acquired = threading.Event() + + def retire_stale(*args): + retire_entered.set() + assert finish_retire.wait(2) + + def load(): + worker.start_load_caches(ECMooncakeConnectorMetadata(), {}) + + def update_reservations(): + with worker._consumer_memory.lock: + lock_acquired.set() + + with patch.object( + worker._consumer_memory, "retire_stale", side_effect=retire_stale + ): + load_thread = threading.Thread(target=load) + load_thread.start() + assert retire_entered.wait(2) + update_thread = threading.Thread(target=update_reservations) + update_thread.start() + assert not lock_acquired.wait(0.05) + finish_retire.set() + load_thread.join(2) + update_thread.join(2) + + assert not load_thread.is_alive() + assert not update_thread.is_alive() + assert lock_acquired.is_set() + + def test_control_server_start_failure_closes_server( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + + with ( + patch_ec_mooncake_deps(), + patch( + "vllm.distributed.ec_transfer.ec_connector.mooncake." + "worker.ConsumerControlServer" + ) as server_cls, + ): + server_cls.return_value.start.side_effect = RuntimeError("bind failed") + connector = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + try: + with pytest.raises(RuntimeError, match="bind failed"): + connector.start_worker_services() + server_cls.return_value.close.assert_called_once_with() + assert connector._worker._control_server is None + finally: + connector.shutdown() + + def test_abandon_retries_allocation_before_reclaiming_resident( + self, mock_vllm_config_consumer + ): + config = mock_vllm_config_consumer + config.ec_transfer_config.ec_buffer_device = "cpu" + config.ec_transfer_config.ec_buffer_size = 512 + config.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 512 + + def payload(transfer_id: str, mm_hash: str) -> dict[str, object]: + return { + "transfer_id": transfer_id, + "mm_hash": mm_hash, + "nbytes": 64, + "shape": [16], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + connector = ECMooncakeConnector(config, ECConnectorRole.WORKER) + worker = connector._worker + memory_pool = worker._consumer_memory + try: + memory_pool.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + resident = memory_pool.try_allocate(64, (16,), torch.float32) + assert resident is not None + memory_pool.publish("resident", resident) + retire_event = MagicMock() + retire_event.query.return_value = True + with ( + patch.object(memory.torch, "Event", return_value=retire_event), + patch.object(memory.torch.accelerator, "current_stream"), + ): + memory_pool.retire_stale({}, set()) + old = worker._reserve_push_destination(payload("old", "old")) + try_allocate = memory_pool.try_allocate + first_attempt = True + + def abandon_between_attempts(*args): + nonlocal first_attempt + if first_attempt: + first_attempt = False + worker._cancel_push("old", old["reservation_id"], abandon=True) + return None + return try_allocate(*args) + + with patch.object( + memory_pool, + "try_allocate", + side_effect=abandon_between_attempts, + ): + new = worker._reserve_push_destination(payload("new", "new")) + + assert new["dst_ptr"] == old["dst_ptr"] + assert memory_pool.drain_reclaimed() == set() + finally: + connector.shutdown() + + def test_producer_push_state_owns_source_until_every_future_is_terminal(self): + manager = ProducerPushManager() + reservation: Future[list[dict[str, Any]]] = Future() + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="transfer", + ) + record, created = manager.reserve(spec, lambda: reservation) + duplicate, duplicate_created = manager.reserve(spec, lambda: Future()) + assert created + assert duplicate is record + assert not duplicate_created + changed = copy.copy(spec) + changed.mm_hash = "other" + with pytest.raises(ValueError, match="changed identity"): + manager.reserve(changed, lambda: Future()) + + source = torch.empty(16) + manager.bind_source("hash", source, None) + assert record.source is not None + reservation.set_result([]) + assert manager.resolve_reservations(record) == [] + assert record.state is ProducerPushState.WAITING_SOURCE + manager.begin_writing(record) + manager.begin_notifying([record]) + + failed: Future[None] = Future() + failed.set_exception(RuntimeError("one shard failed")) + still_writing: Future[None] = Future() + manager.track_shard_futures([record], [failed, still_writing]) + with pytest.raises(RuntimeError, match="source too early"): + manager.fail([record], RuntimeError("write failed")) + assert record.state is ProducerPushState.NOTIFYING + assert record.source is not None + assert record.source.tensor is source + + still_writing.set_result(None) + manager.fail([record], RuntimeError("write failed")) + assert record.state is ProducerPushState.FAILED + assert record.source is None + manager.fail([record], RuntimeError("duplicate failure")) + with pytest.raises(RuntimeError, match="FAILED to NOTIFYING"): + manager.begin_notifying([record]) + + late, late_created = manager.reserve(spec, lambda: Future()) + assert late is record + assert not late_created + + def test_reservation_failure_after_source_binding_releases_the_lease(self): + manager = ProducerPushManager() + reservation: Future[list[dict[str, Any]]] = Future() + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="transfer", + ) + record, _ = manager.reserve(spec, lambda: reservation) + source = torch.empty(16) + manager.bind_source("hash", source, None) + reservation.set_exception(RuntimeError("reserve failed")) + + assert record.state is ProducerPushState.RESERVING + + def run(records) -> None: + try: + manager.resolve_reservations(records[0]) + except RuntimeError as exc: + manager.fail(records, exc) + + with ThreadPoolExecutor(max_workers=1) as executor: + manager.submit_batches(executor, run, lambda: None) + assert manager.poll() == [("hash", "reserve failed")] + assert manager.poll() == [] + assert record.state is ProducerPushState.FAILED + assert record.source is None + + def test_late_reservation_callback_cannot_replace_refreshed_results( + self, mock_vllm_config_producer + ): + manager = ProducerPushManager() + reservation: Future[list[dict[str, Any]]] = Future() + callback_started = threading.Event() + finish_callback = threading.Event() + + def block_callback(_future) -> None: + callback_started.set() + assert finish_callback.wait(2) + + reservation.add_done_callback(block_callback) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="transfer", + ) + record, _ = manager.reserve(spec, lambda: reservation) + old = [{"addr": "old", "reservation_id": "old"}] + refreshed = [{"addr": "new", "reservation_id": "new"}] + setter = threading.Thread(target=reservation.set_result, args=(old,)) + setter.start() + assert callback_started.wait(2) + assert manager.resolve_reservations(record) == old + manager.replace_reservations(record, refreshed) + finish_callback.set() + setter.join(2) + assert not setter.is_alive() + manager.settle_all([record]) + assert record.reservations == refreshed + + with patch_ec_mooncake_deps(): + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + try: + with patch.object( + connector._worker._control_client, + "request", + ) as request: + connector._worker._abandon_pushes([record]) + assert request.call_args.args[1]["reservation_id"] == "new" + finally: + connector.shutdown() + + def test_producer_hot_paths_do_not_scan_terminal_records(self): + manager = ProducerPushManager() + request_ids = set() + limit = 4096 + with patch.object(producer, "_TERMINAL_LIMIT", limit): + for index in range(limit + 2): + reservation: Future[list[dict[str, Any]]] = Future() + reservation.set_result([]) + request_id = f"request-{index}" + request_ids.add(request_id) + spec = ECMooncakePushSpec( + mm_hash=f"hash-{index}", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id=f"transfer-{index}", + request_id=request_id, + ) + manager.reserve(spec, lambda r=reservation: r) + cancelled = manager.cancel_requests(request_ids) + assert len(cancelled) == limit + 2 + for record in cancelled: + manager.finish_cancel(record) + + pinned = manager.get("transfer-0") + assert pinned is not None + batch_started = threading.Event() + finish_batch = threading.Event() + + def block_batch(_record) -> None: + batch_started.set() + assert finish_batch.wait(2) + + executor = ThreadPoolExecutor(max_workers=1) + manager.submit_cancel(pinned, executor, block_batch) + assert batch_started.wait(2) + + class NoScanRecords(OrderedDict): + """Fail if Producer hot paths scan every transfer record.""" + + def __iter__(self): + raise AssertionError("record table scanned") + + def items(self): + raise AssertionError("record table scanned") + + def values(self): + raise AssertionError("record table scanned") + + class NoScanIndex(OrderedDict): + """Fail if Producer hot paths scan every lifecycle index.""" + + def __iter__(self): + raise AssertionError("reapable index scanned") + + def items(self): + raise AssertionError("reapable index scanned") + + def values(self): + raise AssertionError("reapable index scanned") + + manager._records = NoScanRecords(manager._records) + manager._reapable_terminal_ids = NoScanIndex(manager._reapable_terminal_ids) + assert manager.pending + manager.submit_batches(MagicMock(), MagicMock(), MagicMock()) + assert manager.poll() == [] + assert manager.get("transfer-0") is pinned + assert manager.get("transfer-1") is None + assert manager.get(f"transfer-{limit}") is not None + assert len(manager._records) == limit + 1 + + finish_batch.set() + executor.shutdown(wait=True) + assert manager.poll() == [] + assert manager.get("transfer-0") is pinned + assert manager.get("transfer-2") is None + assert len(manager._records) == limit + + def test_producer_push_cancel_handles_pending_and_late_reservations(self): + manager = ProducerPushManager() + pending: Future[list[dict[str, Any]]] = Future() + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="pending", + request_id="request", + ) + record, _ = manager.reserve(spec, lambda: pending) + assert manager.pending + assert manager.cancel_requests({"request"}) == [record] + assert record.state is ProducerPushState.CANCEL_PENDING + manager.bind_source("hash", torch.empty(16), None) + assert record.source is None + + pending.set_result([]) + manager.resolve_reservations(record) + with ThreadPoolExecutor(max_workers=1) as executor: + manager.submit_cancel( + record, + executor, + lambda cancelled: manager.finish_cancel(cancelled), + ) + manager.finish_cancel(record) + manager.poll() + assert record.state is ProducerPushState.CANCELLED + assert not manager.pending + assert manager.cancel_requests({"request"}) == [] + + ready: Future[list[dict[str, Any]]] = Future() + ready.set_result([]) + later = ECMooncakePushSpec( + mm_hash="other", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="ready", + request_id="request-2", + ) + ready_record, _ = manager.reserve(later, lambda: ready) + manager.resolve_reservations(ready_record) + assert manager.cancel_requests({"request-2"}) == [ready_record] + assert ready_record.state is ProducerPushState.CANCEL_PENDING + manager.finish_cancel(ready_record) + assert ready_record.state is ProducerPushState.CANCELLED + + def test_same_source_has_one_lease_per_transfer(self): + manager = ProducerPushManager() + source = torch.empty(16) + records = [] + for transfer_id in ("first", "second"): + reservation: Future[list[dict[str, Any]]] = Future() + reservation.set_result([]) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id=transfer_id, + ) + record, _ = manager.reserve(spec, lambda r=reservation: r) + records.append(record) + manager.bind_source("hash", source, None) + assert all(record.source is not None for record in records) + for record in records: + manager.resolve_reservations(record) + manager.begin_writing(record) + manager.begin_notifying([record]) + manager.complete([records[0]]) + assert records[0].source is None + assert records[1].source is not None + manager.complete([records[1]]) + assert records[1].source is None + + def test_shard_submit_failure_waits_before_source_release( + self, mock_vllm_config_producer + ): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + source = torch.empty(16) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="transfer", + ) + slow_started = threading.Event() + finish_slow = threading.Event() + slow_finished = threading.Event() + released_after_slow: list[bool] = [] + + def request(addr, payload): + if payload["op"] == "reserve": + index = int(addr.rsplit(":", 1)[1]) + return { + "reservation_id": f"reservation-{index}", + "dst_session": f"session-{index}", + "dst_ptr": 1000 + index, + "nbytes": source.nbytes, + "write": True, + "ready": False, + } + return {} + + def write(session, sources, destinations, lengths): + if session == "session-1": + slow_started.set() + assert finish_slow.wait(2) + slow_finished.set() + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + worker = producer._worker + producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) + try: + with ( + patch.object( + worker._topology, + "shards", + return_value=[ + "tcp://consumer:0", + "tcp://consumer:1", + "tcp://consumer:2", + ], + ), + patch.object( + worker._control_client, "request", side_effect=request + ), + patch.object(worker._producer_memory, "stage", return_value=None), + patch.object( + worker._transfer, + "acquire_sources", + return_value=[source.data_ptr()], + ), + patch.object( + worker._transfer, + "release_sources", + side_effect=lambda _: released_after_slow.append( + slow_finished.is_set() + ), + ), + patch.object(worker._transfer, "write", side_effect=write), + ): + producer.start_save_caches(encoder_cache={"hash": source}) + record = worker._producer_pushes.get("transfer") + assert record is not None + record.reservation_futures[0].result(timeout=2) + with ThreadPoolExecutor(max_workers=1) as executor: + submit_count = 0 + + def submit(fn, *args): + nonlocal submit_count + submit_count += 1 + if submit_count == 1: + return executor.submit(fn, *args) + raise RuntimeError("second shard submit failed") + + shard_executor = MagicMock() + shard_executor.submit.side_effect = submit + with patch.object( + worker, + "_shard_executor", + return_value=shard_executor, + ): + assert producer.build_connector_worker_meta().pending_saves + assert slow_started.wait(2) + assert record.source is not None + assert released_after_slow == [] + finish_slow.set() + _wait_for_worker_io(producer) + + record = worker._producer_pushes.get("transfer") + assert record is not None + assert record.state is ProducerPushState.FAILED + assert record.source is None + assert released_after_slow == [True] + finally: + finish_slow.set() + producer.shutdown() + + def test_reserve_failure_waits_for_every_started_shard( + self, mock_vllm_config_producer + ): + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="transfer", + ) + slow_started = threading.Event() + finish_slow = threading.Event() + finished = threading.Event() + errors: list[Exception] = [] + + def reserve_one(addr, _spec): + if addr.endswith(":0"): + raise RuntimeError("first shard failed") + if addr.endswith(":2"): + slow_started.set() + assert finish_slow.wait(2) + return {"addr": addr} + + with patch_ec_mooncake_deps(): + producer = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + worker = producer._worker + + def reserve() -> None: + try: + worker._reserve_remote(spec) + except Exception as exc: + errors.append(exc) + finally: + finished.set() + + try: + with ( + patch.object( + worker._topology, + "shards", + return_value=["shard:0", "shard:1", "shard:2"], + ), + patch.object(worker, "_reserve_one", side_effect=reserve_one), + ): + thread = threading.Thread(target=reserve) + thread.start() + assert slow_started.wait(2) + assert not finished.wait(0.05) + finish_slow.set() + thread.join(2) + assert not thread.is_alive() + assert len(errors) == 1 + assert str(errors[0]) == "first shard failed" + finally: + finish_slow.set() + producer.shutdown() + + def test_reserve_submit_failure_drains_started_shards( + self, mock_vllm_config_producer + ): + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id="transfer", + ) + slow_started = threading.Event() + finish_slow = threading.Event() + finished = threading.Event() + errors: list[Exception] = [] + + def reserve_one(addr, _spec): + slow_started.set() + assert finish_slow.wait(2) + return {"addr": addr} + + with patch_ec_mooncake_deps(): + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER ) + worker = connector._worker - # The same request remains visible on a later scheduler step, but - # its worker push metadata must not be emitted a second time. - assert scheduler.ensure_cache_available(request, 0) - next_meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) - ) + def reserve() -> None: + try: + worker._reserve_remote(spec) + except Exception as exc: + errors.append(exc) + finally: + finished.set() - scheduler.request_finished(request) + try: + with ThreadPoolExecutor(max_workers=1) as executor: + submit_count = 0 + + def submit(fn, *args): + nonlocal submit_count + submit_count += 1 + if submit_count == 1: + return executor.submit(fn, *args) + raise RuntimeError("second shard submit failed") + + shard_executor = MagicMock() + shard_executor.submit.side_effect = submit + with ( + patch.object( + worker._topology, + "shards", + return_value=["shard:0", "shard:1", "shard:2"], + ), + patch.object(worker, "_reserve_one", side_effect=reserve_one), + patch.object( + worker, + "_shard_executor", + return_value=shard_executor, + ), + ): + thread = threading.Thread(target=reserve) + thread.start() + assert slow_started.wait(2) + assert not finished.wait(0.05) + finish_slow.set() + thread.join(2) + assert not thread.is_alive() + assert len(errors) == 1 + assert str(errors[0]) == "second shard submit failed" + finally: + finish_slow.set() + connector.shutdown() - assert meta.loads == [] - assert meta.pushes == [ - ECMooncakePushSpec( - mm_hash="img_hash_1", - nbytes=100 * 16 * 4, - shape=(100, 16), - dtype="float32", - consumer_zmq="tcp://decode:19019", - transfer_id="transfer-1", - request_id="test_req_123", + @pytest.mark.parametrize("source_before_failure", [False, True]) + def test_partial_reserve_is_compensated_before_its_future_fails( + self, mock_vllm_config_producer, source_before_failure + ): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + source = torch.empty(16) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq="tcp://consumer:0", + transfer_id="transfer", + ) + cancel_attempts = 0 + + def reserve_one(addr, _spec): + if addr.endswith(":1"): + raise RuntimeError("reserve shard failed") + return {"addr": addr, "reservation_id": "partial-r0"} + + def request(_addr, payload): + nonlocal cancel_attempts + assert payload == { + "op": "cancel", + "transfer_id": "transfer", + "reservation_id": "partial-r0", + "abandon": True, + } + cancel_attempts += 1 + if cancel_attempts == 1: + raise RuntimeError("transient cleanup failure") + return {"cancelled": True} + + with patch_ec_mooncake_deps(): + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER ) - ] - assert next_meta.pushes == [] - assert "transfer-1" not in scheduler._prepared_push_transfer_ids + worker = connector._worker + connector.bind_connector_metadata( + ECMooncakeConnectorMetadata(pushes=[spec]) + ) + try: + with ( + patch.object( + worker._topology, + "shards", + return_value=["tcp://consumer:0", "tcp://consumer:1"], + ), + patch.object(worker, "_reserve_one", side_effect=reserve_one), + patch.object( + worker._control_client, "request", side_effect=request + ), + ): + connector.start_save_caches( + encoder_cache={"hash": source} + if source_before_failure + else None + ) + record = worker._producer_pushes.get("transfer") + assert record is not None + with pytest.raises( + RuntimeError, match="^reserve shard failed$" + ) as e: + record.reservation_futures[0].result(timeout=2) + assert e.value.partial_reservations == [ + { + "addr": "tcp://consumer:0", + "reservation_id": "partial-r0", + } + ] + assert cancel_attempts == 2 + + if source_before_failure: + assert record.source is not None + connector.build_connector_worker_meta() + _wait_for_worker_io(connector) + else: + assert record.state is ProducerPushState.FAILED + connector.save_caches({"hash": source}, "hash") + assert record.state is ProducerPushState.FAILED + assert record.source is None + finally: + connector.shutdown() - def test_producer_uses_deepstack_encoder_cache_width( - self, mock_vllm_config_producer, mock_request_with_3_mm + def test_partial_complete_abandons_all_shards_before_releasing_source( + self, mock_vllm_config_producer ): - request = mock_request_with_3_mm - request.ec_transfer_params = { - "consumer_zmq": "tcp://decode:19019", - "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], - } - mock_vllm_config_producer.model_config.dtype = torch.bfloat16 - mock_vllm_config_producer.model_config.hf_config = SimpleNamespace( - vision_config=SimpleNamespace( - out_hidden_size=2560, - deepstack_visual_indexes=[5, 11, 17], - ) + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + source = torch.empty(16) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=tuple(source.shape), + dtype="float32", + consumer_zmq="tcp://consumer:0", + transfer_id="transfer", ) + reservations = [ + { + "addr": f"tcp://consumer:{rank}", + "reservation_id": f"r{rank}", + "dst_session": f"session-{rank}", + "dst_ptr": 1000 + rank, + "nbytes": source.nbytes, + "write": True, + "ready": False, + } + for rank in range(2) + ] + slow_started = threading.Event() + finish_slow = threading.Event() + cancelled: list[str] = [] + + def request(addr, payload): + if payload["op"] == "complete_batch": + if addr.endswith(":0"): + raise RuntimeError("complete shard failed") + slow_started.set() + assert finish_slow.wait(2) + return {"items": [{"completed": True}]} + assert payload["op"] == "cancel" and payload["abandon"] + cancelled.append(payload["reservation_id"]) + return {"cancelled": True} with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER ) - scheduler.update_state_after_alloc(request, 0) - meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) + worker = connector._worker + connector.bind_connector_metadata( + ECMooncakeConnectorMetadata(pushes=[spec]) ) + try: + with ( + patch.object(worker, "_reserve_remote", return_value=reservations), + patch.object( + worker._control_client, "request", side_effect=request + ), + patch.object(worker._producer_memory, "stage", return_value=None), + patch.object( + worker._transfer, + "acquire_sources", + return_value=[source.data_ptr()], + ), + patch.object(worker._transfer, "release_sources"), + patch.object(worker._transfer, "write"), + ): + connector.start_save_caches(encoder_cache={"hash": source}) + assert connector.build_connector_worker_meta().pending_saves + record = worker._producer_pushes.get("transfer") + assert record is not None and record.batch_future is not None + assert slow_started.wait(2) + assert record.source is not None + assert not record.batch_future.done() + finish_slow.set() + record.batch_future.result(timeout=2) + connector.build_connector_worker_meta() + + assert Counter(cancelled) == Counter({"r0": 1, "r1": 1}) + assert record.state is ProducerPushState.FAILED + assert record.source is None + assert record.error == "complete shard failed" + assert all(future.done() for future in record.shard_futures) + finally: + finish_slow.set() + connector.shutdown() - num_tokens = request.get_num_encoder_embeds(0) - spec = meta.pushes[0] - assert spec.shape == (num_tokens, 10240) - assert spec.nbytes == num_tokens * 10240 * torch.bfloat16.itemsize - - def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer): - feature = SimpleNamespace( - identifier="image_uuid", - modality="image", - data=SimpleNamespace( - get_data=lambda: { - "image_grid_thw": torch.tensor([1, 32, 48]), - "pixel_values": torch.ones(2), - } - ), + @pytest.mark.parametrize("permanent_failure", [False, True]) + def test_orphan_cancel_is_bounded_retryable_and_skips_cached_shards( + self, mock_vllm_config_producer, permanent_failure + ): + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:0", + transfer_id="transfer", + request_id="request", ) - request = SimpleNamespace(mm_features=[feature]) + reservation: Future[list[dict[str, Any]]] = Future() + reservation.set_result( + [ + { + "addr": "tcp://consumer:0", + "reservation_id": "active", + }, + { + "addr": "tcp://consumer:1", + "reservation_id": "cached", + "cached": True, + }, + { + "addr": "tcp://consumer:2", + "reservation_id": "cancelled", + "cancelled": True, + }, + ] + ) + attempts: Counter[str] = Counter() + + def request(_addr, payload): + reservation_id = payload["reservation_id"] + attempts[reservation_id] += 1 + assert reservation_id == "active" + if permanent_failure or attempts[reservation_id] == 1: + raise RuntimeError("orphan cancel failed") + return {"cancelled": True} with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER ) - with patch.object( - scheduler, - "_placeholder_metadata_fields", - return_value={"image_grid_thw"}, - ): - delay_free, params = scheduler.request_finished(request) + worker = connector._worker + record, _ = worker._producer_pushes.reserve(spec, lambda: reservation) + worker._producer_pushes.resolve_reservations(record) + assert worker._producer_pushes.cancel_requests({"request"}) == [record] + try: + with patch.object( + worker._control_client, "request", side_effect=request + ): + worker._producer_pushes.submit_cancel( + record, + worker._io_executor, + worker._cancel_orphaned_reservation, + ) + assert record.batch_future is not None + if permanent_failure: + with pytest.raises( + RuntimeError, match="^orphan cancel failed$" + ): + record.batch_future.result(timeout=2) + else: + record.batch_future.result(timeout=2) + failures = worker._producer_pushes.poll() + + assert attempts == Counter({"active": 2}) + assert record.state is ProducerPushState.CANCELLED + assert failures == ( + [("hash", "orphan cancel failed")] if permanent_failure else [] + ) + finally: + connector.shutdown() - assert not delay_free - assert params == { - "ec_items": [{"mm_hash": "image_uuid", "image_grid_thw": [1, 32, 48]}] - } + def test_source_contract_checks_shape_dtype_contiguity_and_size(self): + tensors_and_specs = [ + (torch.empty(2, 8), (16,), "float32", 64, "shape"), + (torch.empty(16, dtype=torch.float16), (16,), "float32", 32, "dtype"), + (torch.empty(4, 4).t(), (4, 4), "float32", 64, "contiguous"), + (torch.empty(16), (16,), "float32", 65, "size"), + ] + for index, (tensor, shape, dtype, nbytes, message) in enumerate( + tensors_and_specs + ): + spec = ECMooncakePushSpec( + mm_hash=f"hash-{index}", + nbytes=nbytes, + shape=shape, + dtype=dtype, + consumer_zmq="tcp://consumer:0", + transfer_id=f"transfer-{index}", + ) + reservation: Future[list[dict[str, Any]]] = Future() + reservation.set_result([]) + manager = ProducerPushManager() + record, _ = manager.reserve(spec, lambda future=reservation: future) + manager.bind_source(spec.mm_hash, tensor, None) + with pytest.raises(ValueError, match=message): + ECMooncakeWorker._validate_push_source(record) + + def test_invalid_source_fails_asynchronously_before_staging( + self, mock_vllm_config_producer + ): + mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" + source = torch.empty(2, 8) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=source.nbytes, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:0", + transfer_id="transfer", + ) + def request(_addr, payload): + if payload["op"] == "reserve": + return { + "reservation_id": "reservation", + "dst_session": "session", + "dst_ptr": 1000, + "nbytes": source.nbytes, + "write": True, + "ready": False, + } + assert payload["op"] == "cancel" + return {"cancelled": True} + + with patch_ec_mooncake_deps(): + connector = ECMooncakeConnector( + mock_vllm_config_producer, ECConnectorRole.WORKER + ) + worker = connector._worker + connector.bind_connector_metadata( + ECMooncakeConnectorMetadata(pushes=[spec]) + ) + try: + with ( + patch.object( + worker._topology, + "shards", + return_value=["tcp://consumer:0"], + ), + patch.object( + worker._control_client, "request", side_effect=request + ), + patch.object(worker._producer_memory, "stage") as stage, + patch.object(worker._transfer, "acquire_sources") as register, + ): + connector.start_save_caches(encoder_cache={"hash": source}) + assert connector.build_connector_worker_meta().pending_saves + record = worker._producer_pushes.get("transfer") + assert record is not None and record.batch_future is not None + record.batch_future.result(timeout=2) + connector.build_connector_worker_meta() + + assert record.state is ProducerPushState.FAILED + assert record.source is None + assert record.error == "EC source shape mismatch for mm_hash=hash" + stage.assert_not_called() + register.assert_not_called() + finally: + connector.shutdown() -class TestECMooncakeWorkerTransfer: def test_batches_pushes_from_one_model_step(self, mock_vllm_config_producer): port = _find_free_port() consumer_cfg = Mock(spec=VllmConfig) @@ -1079,15 +4368,15 @@ def test_batches_pushes_from_one_model_step(self, mock_vllm_config_producer): producer.start_save_caches(encoder_cache=sources) _wait_for_worker_io(producer) - engine = producer._engine + engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) assert len(engine.transfer_calls) == 1 assert sorted(engine.transfer_calls[0]) == sorted( tensor.nbytes for tensor in sources.values() ) assert all( - reservation.ready - for reservation in consumer._push_reservations.values() + reservation.state is ConsumerReservationState.READY + for reservation in consumer._worker._reservations.active_records() ) finally: producer.shutdown() @@ -1131,7 +4420,9 @@ def test_push_reserves_before_encoder_output_is_saved( producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) try: producer.start_save_caches(encoder_cache={}) - _, reservation = producer._pending_reservations["hash"][0] + push_record = producer._worker._producer_pushes.get("transfer-1") + assert push_record is not None + reservation = push_record.reservation_futures[0] shards = reservation.result(timeout=2) # One reservation per consumer shard; this consumer is single. assert len(shards) == 1 @@ -1139,11 +4430,11 @@ def test_push_reserves_before_encoder_output_is_saved( assert reservation_data["nbytes"] == source.nbytes old_reservation_id = reservation_data["reservation_id"] reservation_data["_received_at"] -= _LEASE_TTL_SECONDS - consumer._push_reservations["transfer-1"].expires_at = 0 + consumer._worker._reservations.get("transfer-1").expires_at = 0 with patch.object( - scheduler, - "_send_control", - wraps=scheduler._send_control, + scheduler._scheduler._control_client, + "request", + wraps=scheduler._scheduler._control_client.request, ) as send_control: assert not scheduler.has_cache_item("hash") assert not scheduler.has_cache_item("hash") @@ -1153,12 +4444,12 @@ def test_push_reserves_before_encoder_output_is_saved( {"op": "peers"}, {"op": "event_port"}, ] - assert "transfer-1" in consumer._push_reservations + assert consumer._worker._reservations.status("transfer-1") producer.save_caches({"hash": source}, "hash") _wait_for_worker_io(producer) assert ( - consumer._push_reservations["transfer-1"].reservation_id + consumer._worker._reservations.get("transfer-1").reservation_id != old_reservation_id ) deadline = time.monotonic() + 2 @@ -1168,7 +4459,9 @@ def test_push_reserves_before_encoder_output_is_saved( # Still just the two setup requests: polling for readiness # must not re-open the channel. assert send_control.call_count == 2 - load = scheduler._pending_specs["transfer-1"] + record = scheduler._scheduler._transfers.get("transfer-1") + assert record is not None and record.spec is not None + load = record.spec consumer.bind_connector_metadata( ECMooncakeConnectorMetadata(loads=[load]) ) @@ -1177,7 +4470,7 @@ def test_push_reserves_before_encoder_output_is_saved( first_meta = consumer.build_connector_worker_meta() assert first_meta.loaded == {"hash"} assert torch.equal(loaded["hash"], source) - consumer_engine = consumer._engine + consumer_engine = consumer._worker._transfer._engine assert isinstance(consumer_engine, CopyingFakeTransferEngine) assert consumer_engine.transfer_calls == [] finally: @@ -1224,14 +4517,18 @@ def test_finished_request_cancels_unbound_reservation( producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) try: producer.start_save_caches(encoder_cache={}) - _, reservation = producer._pending_reservations["hash"][0] + push_record = producer._worker._producer_pushes.get("transfer-1") + assert push_record is not None + reservation = push_record.reservation_futures[0] reservation.result(timeout=2) - assert "transfer-1" in consumer._push_reservations + assert consumer._worker._reservations.status("transfer-1") producer.get_finished({"request-1"}) _wait_for_worker_io(producer) - assert "hash" not in producer._pending_reservations - assert "transfer-1" not in consumer._push_reservations + push_record = producer._worker._producer_pushes.get("transfer-1") + assert push_record is not None + assert push_record.state is ProducerPushState.CANCELLED + assert consumer._worker._reservations.status("transfer-1") is None finally: producer.shutdown() consumer.shutdown() @@ -1278,20 +4575,27 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( producer.start_save_caches(encoder_cache={"hash": source}) _wait_for_worker_io(producer) - engine = producer._engine + engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) assert engine.transfer_calls == [[source.nbytes]] - reservation = consumer._push_reservations["transfer-1"] - assert reservation.ready - assert consumer._consumer_worker_metrics["completions_accepted"] == 1 - assert consumer._consumer_worker_metrics["completions_repeated"] == 0 + reservation = consumer._worker._reservations.get("transfer-1") + assert reservation.state is ConsumerReservationState.READY + assert ( + consumer._worker._consumer_worker_metrics["completions_accepted"] + == 1 + ) + assert ( + consumer._worker._consumer_worker_metrics["completions_repeated"] + == 0 + ) deadline = time.monotonic() + 2 while not scheduler.has_cache_item("hash"): assert time.monotonic() < deadline time.sleep(0.01) - load = scheduler._pop_pending_spec("transfer-1") - assert load is not None + record = scheduler._scheduler._transfers.get("transfer-1") + assert record is not None and record.spec is not None + load = record.spec load.num_token = 4 consumer.bind_connector_metadata( ECMooncakeConnectorMetadata(loads=[load]) @@ -1313,16 +4617,22 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( producer.start_save_caches(encoder_cache={"hash": source}) _wait_for_worker_io(producer) assert engine.transfer_calls == [[source.nbytes]] - cached = consumer._push_reservations["transfer-2"] - assert cached.ready and not cached.owns_allocation - assert consumer._consumer_worker_metrics["reservations_cached"] == 1 + cached = consumer._worker._reservations.get("transfer-2") + assert cached is not None + assert cached.state is ConsumerReservationState.READY + assert cached.lease is not None + assert ( + consumer._worker._consumer_worker_metrics["reservations_cached"] + == 1 + ) deadline = time.monotonic() + 2 while not scheduler.has_cache_item("hash"): assert time.monotonic() < deadline time.sleep(0.01) - cached_load = scheduler._pop_pending_spec("transfer-2") - assert cached_load is not None + record = scheduler._scheduler._transfers.get("transfer-2") + assert record is not None and record.spec is not None + cached_load = record.spec cached_load.num_token = 4 consumer.bind_connector_metadata( ECMooncakeConnectorMetadata(loads=[cached_load]) @@ -1330,7 +4640,7 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( consumer.start_load_caches(loaded) cached_meta = consumer.build_connector_worker_meta() assert cached_meta.loaded == {"hash"} - assert "transfer-2" not in consumer._push_reservations + assert consumer._worker._reservations.status("transfer-2") is None assert torch.equal(loaded["hash"], source) finally: producer.shutdown() @@ -1372,20 +4682,25 @@ def test_retired_item_reserved_again_still_serves_a_local_load( with patch_ec_mooncake_deps(): consumer = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) try: - consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) - pool = consumer._consumer_pool - allocator = consumer._consumer_pool_allocator - assert pool is not None and allocator is not None - offset, size = allocator.allocate(spec.nbytes) - tensor = ( - pool.narrow(0, offset, spec.nbytes).view(torch.float32).view(4, 4) - ) - allocation = _ConsumerPoolAllocation(offset, size, tensor) - consumer._consumer_residents.insert("hash", allocation, size) - consumer._consumer_residents.retire("hash") + consumer._worker._consumer_memory.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + allocation = consumer._worker._consumer_memory.try_allocate( + spec.nbytes, spec.shape, torch.float32 + ) + assert allocation is not None + tensor = allocation.tensor + consumer._worker._consumer_memory.publish("hash", allocation) + retire_event = MagicMock() + retire_event.query.return_value = True + with ( + patch.object(memory.torch, "Event", return_value=retire_event), + patch.object(memory.torch.accelerator, "current_stream"), + ): + consumer._worker._consumer_memory.retire_stale({}, set()) # A later push reserves the retired copy instead of transferring. - consumer._reserve_push_destination( + consumer._worker._reserve_push_destination( { "transfer_id": "t1", "mm_hash": "hash", @@ -1394,9 +4709,68 @@ def test_retired_item_reserved_again_still_serves_a_local_load( "dtype": spec.dtype, } ) - assert consumer._consumer_residents.num_evictable == 0 + assert consumer._worker._consumer_memory.stats()[2] == 0 + + assert ( + consumer._worker._consumer_memory.take_resident( + spec.mm_hash, spec.shape, spec.dtype + ) + is tensor + ) + finally: + consumer.shutdown() + + def test_cached_take_uses_newer_same_hash_canonical( + self, mock_vllm_config_consumer + ): + config = mock_vllm_config_consumer + config.ec_transfer_config.ec_buffer_device = "cpu" + config.ec_transfer_config.ec_buffer_size = 768 + config.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 768 + shape = (16,) + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector(config, ECConnectorRole.WORKER) + worker = consumer._worker + memory_pool = worker._consumer_memory + try: + memory_pool.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + first = memory_pool.try_allocate(64, shape, torch.float32) + replacement = memory_pool.try_allocate(64, shape, torch.float32) + assert first is not None and replacement is not None + memory_pool.publish("hash", first) + worker._reserve_push_destination( + { + "transfer_id": "cached", + "mm_hash": "hash", + "nbytes": 64, + "shape": list(shape), + "dtype": "float32", + } + ) + memory_pool.publish("hash", replacement) + memory_pool.retire_stale({}, {"hash"}) + spec = ECMooncakeLoadSpec( + mm_hash="hash", + num_token=1, + nbytes=64, + shape=shape, + dtype="float32", + pushed=True, + transfer_id="cached", + ) + + tensor, allocation = worker._take_pushed_tensor(spec) - assert consumer._take_resident_tensor(spec) is tensor + assert allocation is replacement + assert tensor is replacement.tensor + reused = memory_pool.try_allocate(64, shape, torch.float32) + assert reused is not None + assert reused.offset == first.offset finally: consumer.shutdown() @@ -1445,11 +4819,15 @@ def fake_send(addr: str, request: dict): ) producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) try: - with patch.object(producer, "_send_control", side_effect=fake_send): + with patch.object( + producer._worker._control_client, + "request", + side_effect=fake_send, + ): producer.start_save_caches(encoder_cache={"hash": source}) _wait_for_worker_io(producer) - engine = producer._engine + engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) # One write per rank, and every rank got the same bytes. assert len(engine.transfer_calls) == len(shard_ports) @@ -1501,18 +4879,18 @@ def test_pushes_stage_through_the_registered_pool(self, mock_vllm_config_produce producer.start_save_caches(encoder_cache={"hash": source}) _wait_for_worker_io(producer) - engine = producer._engine + engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) # The staging pool is registered once; a transfer registers # nothing of its own. - pool = producer._producer_pool + pool = producer._worker._producer_memory.tensor assert pool is not None assert engine.register_calls == [[pool.data_ptr()]] assert engine.batch_unregister_calls == [] assert engine.transfer_calls == [[source.nbytes, source.nbytes]] assert all( - reservation.ready - for reservation in consumer._push_reservations.values() + reservation.state is ConsumerReservationState.READY + for reservation in consumer._worker._reservations.active_records() ) finally: producer.shutdown() @@ -1524,9 +4902,6 @@ def test_push_falls_back_to_per_tensor_registration_without_a_pool( """A pool that cannot be created must not break pushes.""" port = _find_free_port() consumer_cfg = self._push_harness_config(mock_vllm_config_producer, port) - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "producer_buffer_pool_size" - ] = 0 source = torch.randn(4, 16) spec = ECMooncakePushSpec( mm_hash="hash", @@ -1545,11 +4920,16 @@ def test_push_falls_back_to_per_tensor_registration_without_a_pool( consumer.start_worker_services() producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) try: - producer.start_save_caches(encoder_cache={"hash": source}) - _wait_for_worker_io(producer) - engine = producer._engine + with patch( + "vllm.distributed.ec_transfer.ec_connector." + "mooncake.memory.torch.empty", + side_effect=torch.OutOfMemoryError, + ): + producer.start_save_caches(encoder_cache={"hash": source}) + _wait_for_worker_io(producer) + engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) - assert producer._producer_pool is None + assert producer._worker._producer_memory.tensor is None assert engine.register_calls == [[source.data_ptr()]] assert engine.batch_unregister_calls == [[source.data_ptr()]] assert engine.transfer_calls == [[source.nbytes]] @@ -1569,15 +4949,15 @@ def test_concurrent_pushes_hold_source_registration_until_last_release( mock_vllm_config_producer, ECConnectorRole.WORKER ) try: - first = producer._acquire_push_source_registrations([source]) - second = producer._acquire_push_source_registrations([source]) - engine = producer._engine + first = producer._worker._transfer.acquire_sources([source]) + second = producer._worker._transfer.acquire_sources([source]) + engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) assert len(engine.register_calls) == 1 - producer._release_push_source_registrations(first) + producer._worker._transfer.release_sources(first) assert engine.batch_unregister_calls == [] - producer._release_push_source_registrations(second) + producer._worker._transfer.release_sources(second) assert engine.batch_unregister_calls == [first] finally: producer.shutdown() @@ -1627,7 +5007,9 @@ def test_batch_completion_sends_one_control_message( producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=pushes)) try: with patch.object( - producer, "_send_control", wraps=producer._send_control + producer._worker._control_client, + "request", + wraps=producer._worker._control_client.request, ) as send_control: producer.start_save_caches(encoder_cache=sources) _wait_for_worker_io(producer) @@ -1635,8 +5017,8 @@ def test_batch_completion_sends_one_control_message( assert ops.count("complete_batch") == 1 assert "complete" not in ops assert all( - reservation.ready - for reservation in consumer._push_reservations.values() + reservation.state is ConsumerReservationState.READY + for reservation in consumer._worker._reservations.active_records() ) finally: producer.shutdown() @@ -1664,13 +5046,14 @@ def test_failed_push_is_reported_not_raised(self, mock_vllm_config_producer): consumer.start_worker_services() producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) try: - engine = producer._engine or producer._ensure_engine() + producer._worker._transfer.ensure_ready() + engine = producer._worker._transfer._engine with patch.object(engine, "batch_transfer_sync_write", return_value=1): producer.start_save_caches(encoder_cache={"hash": source}) # No raise: the batch reports itself and gives up the # consumer-side reservation. _wait_for_worker_io(producer) - assert "transfer" not in consumer._push_reservations + assert consumer._worker._reservations.status("transfer") is None finally: producer.shutdown() consumer.shutdown() @@ -1689,8 +5072,10 @@ def test_complete_is_idempotent_without_republishing( mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) - reservation = consumer._reserve_push_destination( + consumer._worker._consumer_memory.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + reservation = consumer._worker._reserve_push_destination( { "mm_hash": "hash", "transfer_id": "transfer-1", @@ -1701,14 +5086,65 @@ def test_complete_is_idempotent_without_republishing( ) reservation_id = reservation["reservation_id"] - first = consumer._complete_push("transfer-1", reservation_id) - repeated = consumer._complete_push("transfer-1", reservation_id) + first = consumer._worker._complete_push("transfer-1", reservation_id) + repeated = consumer._worker._complete_push("transfer-1", reservation_id) assert first.accepted and first.became_ready assert repeated.accepted and not repeated.became_ready finally: consumer.shutdown() + def test_cancel_pending_repeat_reserve_is_terminal_without_releasing( + self, mock_vllm_config_consumer + ): + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" + mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ + "consumer_buffer_pool_size" + ] = 4096 + payload = { + "mm_hash": "hash", + "transfer_id": "transfer", + "nbytes": 64, + "shape": [4, 4], + "dtype": "float32", + } + + with patch_ec_mooncake_deps(): + consumer = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.WORKER + ) + memory_pool = consumer._worker._consumer_memory + try: + memory_pool.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + first = consumer._worker._reserve_push_destination(payload) + record = consumer._worker._reservations.get("transfer") + assert record is not None and record.allocation is not None + allocation = record.allocation + with patch.object(memory_pool, "free", wraps=memory_pool.free) as free: + assert consumer._worker._cancel_push( + "transfer", first["reservation_id"] + ) + repeated = consumer._worker._reserve_push_destination(payload) + + assert repeated["cancelled"] + assert not repeated["write"] and not repeated["ready"] + assert record.state is ConsumerReservationState.CANCEL_PENDING + assert record.allocation is allocation + free.assert_not_called() + + completed = consumer._worker._complete_push( + "transfer", first["reservation_id"] + ) + assert completed.accepted and not completed.became_ready + assert record.state is ConsumerReservationState.CANCELLED + assert record.allocation is None + free.assert_called_once_with(allocation) + finally: + consumer.shutdown() + def test_same_hash_transfers_have_independent_lifecycles( self, mock_vllm_config_consumer ): @@ -1732,18 +5168,28 @@ def payload(transfer_id: str) -> dict: mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) - first = consumer._reserve_push_destination(payload("first")) - second = consumer._reserve_push_destination(payload("second")) - - consumer._complete_push("first", first["reservation_id"]) - assert consumer._push_reservations["first"].ready - assert not consumer._push_reservations["second"].ready - - assert consumer._cancel_push("first", first["reservation_id"]) - assert "first" not in consumer._push_reservations - assert "second" in consumer._push_reservations - assert consumer._complete_push("second", second["reservation_id"]) + consumer._worker._consumer_memory.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + first = consumer._worker._reserve_push_destination(payload("first")) + second = consumer._worker._reserve_push_destination(payload("second")) + + consumer._worker._complete_push("first", first["reservation_id"]) + assert ( + consumer._worker._reservations.get("first").state + is ConsumerReservationState.READY + ) + assert ( + consumer._worker._reservations.get("second").state + is ConsumerReservationState.WRITING + ) + + assert consumer._worker._cancel_push("first", first["reservation_id"]) + assert consumer._worker._reservations.status("first") is None + assert consumer._worker._reservations.status("second") + assert consumer._worker._complete_push( + "second", second["reservation_id"] + ) finally: consumer.shutdown() @@ -1768,16 +5214,35 @@ def test_late_completion_cannot_complete_new_reservation( mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) - old = consumer._reserve_push_destination(payload) - consumer._push_reservations["transfer"].expires_at = 0 - consumer._expire_push_reservations() - new = consumer._reserve_push_destination(payload) + consumer._worker._consumer_memory.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + old = consumer._worker._reserve_push_destination(payload) + consumer._worker._reservations.get("transfer").expires_at = 0 + consumer._worker._expire_push_reservations() + assert ( + consumer._worker._reservations.get("transfer").state + is ConsumerReservationState.EXPIRE_PENDING + ) + assert consumer._worker._cancel_push( + "transfer", + old["reservation_id"], + abandon=True, + refresh=True, + ) + new = consumer._worker._reserve_push_destination(payload) + new_record = consumer._worker._reservations.get("transfer") + assert new_record is not None and new_record.allocation is not None + new_allocation = new_record.allocation assert old["reservation_id"] != new["reservation_id"] - stale = consumer._complete_push("transfer", old["reservation_id"]) + stale = consumer._worker._complete_push( + "transfer", old["reservation_id"] + ) assert not stale.accepted - assert not consumer._push_reservations["transfer"].ready + assert consumer._worker._reservations.get("transfer") is new_record + assert new_record.allocation is new_allocation + assert new_record.state is ConsumerReservationState.WRITING finally: consumer.shutdown() @@ -1793,8 +5258,10 @@ def test_ready_reservation_has_a_terminal_expiry(self, mock_vllm_config_consumer mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) - reservation = consumer._reserve_push_destination( + consumer._worker._consumer_memory.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + reservation = consumer._worker._reserve_push_destination( { "mm_hash": "hash", "transfer_id": "transfer-1", @@ -1803,11 +5270,13 @@ def test_ready_reservation_has_a_terminal_expiry(self, mock_vllm_config_consumer "dtype": "float32", } ) - consumer._complete_push("transfer-1", reservation["reservation_id"]) - consumer._push_reservations["transfer-1"].expires_at = 0 + consumer._worker._complete_push( + "transfer-1", reservation["reservation_id"] + ) + consumer._worker._reservations.get("transfer-1").expires_at = 0 - assert consumer._expire_push_reservations() == 1 - assert "transfer-1" not in consumer._push_reservations + assert consumer._worker._expire_push_reservations() == 1 + assert consumer._worker._reservations.status("transfer-1") is None finally: consumer.shutdown() @@ -1832,15 +5301,19 @@ def test_cancel_before_reserve_creates_bounded_tombstone( mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) - assert consumer._cancel_push("cancelled-transfer", "") - cancelled = consumer._reserve_push_destination(payload) + consumer._worker._consumer_memory.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + assert consumer._worker._cancel_push("cancelled-transfer", "") + cancelled = consumer._worker._reserve_push_destination(payload) assert cancelled["cancelled"] and not cancelled["write"] - assert "cancelled-transfer" not in consumer._push_reservations + assert ( + consumer._worker._reservations.status("cancelled-transfer") is None + ) - consumer._cancelled_transfers["cancelled-transfer"] = 0 - consumer._expire_push_reservations() - replacement = consumer._reserve_push_destination(payload) + consumer._worker._reservations.get("cancelled-transfer").expires_at = 0 + consumer._worker._expire_push_reservations() + replacement = consumer._worker._reserve_push_destination(payload) assert replacement["write"] finally: consumer.shutdown() @@ -1865,16 +5338,25 @@ def test_repeated_cancel_does_not_strand_older_tombstones( mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._ensure_consumer_pool(torch.device("cpu"), allow_host=True) - assert consumer._cancel_push("refreshed-transfer", "") - assert consumer._cancel_push("stale-transfer", "") - consumer._cancelled_transfers["stale-transfer"] = 0.0 - assert consumer._cancel_push("refreshed-transfer", "") + consumer._worker._consumer_memory.prepare( + torch.device("cpu"), receiving_rank=True, allow_host=True + ) + assert consumer._worker._cancel_push("refreshed-transfer", "") + assert consumer._worker._cancel_push("stale-transfer", "") + consumer._worker._reservations.get("stale-transfer").expires_at = 0.0 + assert consumer._worker._cancel_push("refreshed-transfer", "") - consumer._expire_push_reservations() + consumer._worker._expire_push_reservations() - assert list(consumer._cancelled_transfers) == ["refreshed-transfer"] - assert consumer._consumer_worker_metrics["cancel_records_dropped"] == 1 + assert consumer._worker._reservations.get("stale-transfer") is None + assert ( + consumer._worker._reservations.get("refreshed-transfer").state + is ConsumerReservationState.CANCELLED + ) + assert ( + consumer._worker._consumer_worker_metrics["cancel_records_dropped"] + == 1 + ) finally: consumer.shutdown() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/__init__.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/__init__.py new file mode 100644 index 000000000000..5cd38b5017ce --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/__init__.py @@ -0,0 +1,9 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Internal building blocks for the Mooncake encoder-cache connector. + +The public connector remains in :mod:`mooncake_ec_connector`. This package +separates configuration, control-plane messaging, transfer state, registered +memory, and worker/scheduler orchestration so that each component has one +owner and can be tested independently. +""" diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py new file mode 100644 index 000000000000..c6f9546e9a4d --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py @@ -0,0 +1,28 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Guard access to the optional Mooncake TransferEngine dependency. + +Keeping the import check here lets metadata and configuration modules remain +importable in environments that do not install Mooncake. +""" + +_MOONCAKE_IMPORT_ERROR: ImportError | None +try: + from mooncake.engine import TransferEngine as _TransferEngine # noqa: F401 +except ImportError as e: + _MOONCAKE_IMPORT_ERROR = e +else: + _MOONCAKE_IMPORT_ERROR = None + + +def ensure_mooncake_available() -> None: + """Raise a user-facing error when Mooncake is unavailable. + + Raises: + ImportError: If ``mooncake-transfer-engine`` cannot be imported. + """ + if _MOONCAKE_IMPORT_ERROR is not None: + raise ImportError( + "Install mooncake-transfer-engine (see " + "https://github.com/kvcache-ai/Mooncake ) to use ECMooncakeConnector." + ) from _MOONCAKE_IMPORT_ERROR diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py new file mode 100644 index 000000000000..1d1ece14d424 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py @@ -0,0 +1,232 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Parse and validate configuration shared by Mooncake connector roles.""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import TYPE_CHECKING + +from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole + +if TYPE_CHECKING: + from vllm.config import VllmConfig + + +def _integer(name: str, value: object) -> int: + message = f"ECMooncakeConnector requires {name} to be an integer." + if isinstance(value, int) and not isinstance(value, bool): + return value + if isinstance(value, float) and value.is_integer(): + return int(value) + if isinstance(value, str): + try: + return int(value) + except ValueError as error: + raise ValueError(message) from error + raise ValueError(message) + + +def _positive_integer(name: str, value: object) -> int: + parsed = _integer(name, value) + if parsed <= 0: + raise ValueError(f"ECMooncakeConnector requires {name} > 0.") + return parsed + + +def _finite_float(name: str, value: object, allow_zero: bool) -> float: + requirement = ">= 0" if allow_zero else "> 0" + message = f"ECMooncakeConnector requires {name} {requirement}." + if isinstance(value, bool) or not isinstance(value, (str, int, float)): + raise ValueError(message) + try: + parsed = float(value) + except ValueError as error: + raise ValueError(message) from error + if not math.isfinite(parsed) or parsed < 0 or (parsed == 0 and not allow_zero): + raise ValueError(message) + return parsed + + +def _nonempty_string(name: str, value: object) -> str: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"ECMooncakeConnector requires non-empty {name}.") + return value.strip() + + +@dataclass(frozen=True) +class MooncakeECConfig: + """Validated runtime settings for one Scheduler or Worker instance. + + Attributes: + is_producer: Whether this instance can originate encoder-cache pushes. + is_consumer: Whether this instance can receive encoder-cache pushes. + protocol: Mooncake transport protocol passed to ``TransferEngine``. + buffer_device: Device used for registered staging and receive buffers. + reservation_port: Rank-adjusted Worker control-plane base port. + reservation_addr: Scheduler-visible Consumer control address. + control_timeout_s: Timeout for one ZMQ request/response exchange. + push_wait_timeout_s: Maximum Scheduler wait for a ready notification. + transfer_workers: Maximum concurrent data-plane transfer batches. + control_workers: Maximum concurrent control-plane operations. + producer_pool_size: Bytes reserved for the Producer staging pool. + consumer_pool_size: Bytes reserved for the Consumer receive pool. + transfer_metrics_log_interval: Producer transfer log interval in seconds. + consumer_metrics_log_interval: Consumer metrics log interval in seconds. + """ + + is_producer: bool + is_consumer: bool + protocol: str + buffer_device: str + reservation_port: int | None + reservation_addr: str | None + control_timeout_s: float + push_wait_timeout_s: float + transfer_workers: int + control_workers: int + producer_pool_size: int + consumer_pool_size: int + transfer_metrics_log_interval: float + consumer_metrics_log_interval: float + + @property + def control_timeout_ms(self) -> int: + return max(1, math.ceil(self.control_timeout_s * 1000)) + + @classmethod + def from_vllm_config( + cls, vllm_config: VllmConfig, role: ECConnectorRole + ) -> MooncakeECConfig: + """Build role-specific settings from the top-level vLLM config. + + Args: + vllm_config: Source vLLM configuration. + role: Connector process role being configured. + + Returns: + Validated, normalized Mooncake connector settings. + + Raises: + ValueError: If an option is invalid or the requested parallel + topology is unsupported for a Producer. + """ + parallel_config = vllm_config.parallel_config + ec_config = vllm_config.ec_transfer_config + assert ec_config is not None + + is_producer = ec_config.is_ec_producer + if is_producer: + if parallel_config.tensor_parallel_size > 1: + raise ValueError( + "ECMooncakeConnector producers require tensor_parallel_size=1." + ) + if parallel_config.pipeline_parallel_size > 1: + raise ValueError( + "ECMooncakeConnector producers do not support pipeline parallelism." + ) + if parallel_config.data_parallel_size > 1: + raise ValueError( + "ECMooncakeConnector producers require data_parallel_size=1." + ) + + registered_buffer_size = _positive_integer( + "ec_buffer_size", ec_config.ec_buffer_size + ) + + extra = ec_config.ec_connector_extra_config + raw_port = extra.get("reservation_zmq_port") + reservation_port = ( + _integer("reservation_zmq_port", raw_port) if raw_port is not None else None + ) + if reservation_port is not None and not 1 <= reservation_port <= 65535: + raise ValueError( + "ECMooncakeConnector requires reservation_zmq_port in 1..65535." + ) + + if reservation_port is not None: + reservation_port += ( + parallel_config.data_parallel_index + * parallel_config.tensor_parallel_size + ) + highest_port = reservation_port + parallel_config.tensor_parallel_size - 1 + if not 1 <= reservation_port <= highest_port <= 65535: + raise ValueError( + "ECMooncakeConnector reservation ports must be in 1..65535." + ) + + reservation_addr = ( + _nonempty_string("reservation_zmq_addr", extra["reservation_zmq_addr"]) + if "reservation_zmq_addr" in extra + else None + ) + if reservation_addr is None and reservation_port is not None: + reservation_addr = f"tcp://127.0.0.1:{reservation_port}" + + is_consumer = ec_config.is_ec_consumer + if is_consumer and role == ECConnectorRole.SCHEDULER and not reservation_addr: + raise ValueError( + "ec_consumer with ECMooncakeConnector requires " + "reservation_zmq_port or reservation_zmq_addr." + ) + if is_consumer and role == ECConnectorRole.WORKER and reservation_port is None: + raise ValueError( + "ec_consumer with ECMooncakeConnector workers require " + "reservation_zmq_port." + ) + + control_timeout_s = _finite_float( + "control_timeout_s", extra.get("control_timeout_s", 30), False + ) + if control_timeout_s > (2**31 - 1) / 1000: + raise ValueError("ECMooncakeConnector control_timeout_s is too large.") + push_wait_timeout_s = _finite_float( + "push_wait_timeout_s", extra.get("push_wait_timeout_s", 60), False + ) + transfer_workers = _positive_integer( + "transfer_max_workers", extra.get("transfer_max_workers", 4) + ) + control_workers = _positive_integer( + "control_max_workers", extra.get("control_max_workers", 8) + ) + producer_pool_size = _positive_integer( + "producer_buffer_pool_size", + extra.get("producer_buffer_pool_size", registered_buffer_size), + ) + consumer_pool_size = _positive_integer( + "consumer_buffer_pool_size", + extra.get("consumer_buffer_pool_size", registered_buffer_size), + ) + + protocol = _nonempty_string( + "mooncake_protocol", extra.get("mooncake_protocol", "rdma") + ) + raw_buffer_device = ec_config.ec_buffer_device + if raw_buffer_device is not None and not isinstance(raw_buffer_device, str): + raise ValueError("ECMooncakeConnector ec_buffer_device must be a string.") + + return cls( + is_producer=is_producer, + is_consumer=is_consumer, + protocol=protocol, + buffer_device=(raw_buffer_device or "cuda").strip() or "cuda", + reservation_port=reservation_port, + reservation_addr=reservation_addr, + control_timeout_s=control_timeout_s, + push_wait_timeout_s=push_wait_timeout_s, + transfer_workers=transfer_workers, + control_workers=control_workers, + producer_pool_size=producer_pool_size, + consumer_pool_size=consumer_pool_size, + transfer_metrics_log_interval=_finite_float( + "transfer_metrics_log_interval", + extra.get("transfer_metrics_log_interval", 10), + True, + ), + consumer_metrics_log_interval=_finite_float( + "consumer_metrics_log_interval", + extra.get("consumer_metrics_log_interval", 10), + True, + ), + ) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py new file mode 100644 index 000000000000..260b21d89295 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py @@ -0,0 +1,527 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""ZMQ control plane for Mooncake encoder-cache reservations and events. + +Tensor bytes never travel through this module. It coordinates destination +reservations, completion/cancellation, TP-shard discovery, and readiness +notifications while :mod:`transfer` owns the Mooncake data plane. +""" + +from __future__ import annotations + +import threading +import time +from collections import Counter, deque +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any, Literal, TypedDict, cast + +import torch +import zmq +from typing_extensions import NotRequired + +from vllm.logger import init_logger + +logger = init_logger(__name__) + +_MAX_PENDING_EVENTS = 4096 +_RESERVATION_REAP_INTERVAL_SECONDS = 1 + + +class ReservationItem(TypedDict): + """Identify one Consumer reservation in a batch control request.""" + + transfer_id: str + reservation_id: str + + +class PeersRequest(TypedDict): + """Request the control ports of every Consumer TP shard.""" + + op: Literal["peers"] + + +class EventPortRequest(TypedDict): + """Request the PUSH socket port used for readiness events.""" + + op: Literal["event_port"] + + +class StatusRequest(TypedDict): + """Request the current status of a transfer reservation.""" + + op: Literal["status"] + transfer_id: str + + +class ReserveRequest(TypedDict): + """Request destination memory for an encoder-cache tensor.""" + + op: Literal["reserve"] + transfer_id: str + mm_hash: str + nbytes: int + shape: list[int] + dtype: str + + +class CompleteBatchRequest(TypedDict): + """Mark several destination writes complete in one exchange.""" + + op: Literal["complete_batch"] + items: list[ReservationItem] + + +class ReservationActionRequest(ReservationItem): + """Complete or cancel one previously created reservation.""" + + op: Literal["complete", "cancel"] + abandon: NotRequired[bool] + refresh: NotRequired[bool] + + +ControlRequest = ( + PeersRequest + | EventPortRequest + | StatusRequest + | ReserveRequest + | ReservationActionRequest + | CompleteBatchRequest +) + + +class ControlSuccess(TypedDict): + """Successful wire response with an optional operation result.""" + + ok: Literal[True] + result: NotRequired[Any] + + +class ControlFailure(TypedDict): + """Failed wire response containing a user-facing error message.""" + + ok: Literal[False] + error: str + + +ControlResponse = ControlSuccess | ControlFailure + + +@dataclass(frozen=True) +class ControlCompletion: + """Summarize the effect of a Consumer completion request. + + Attributes: + accepted: Whether the reservation identity was valid. + became_ready: Whether this call newly made the tensor readable. + """ + + accepted: bool + became_ready: bool = False + + +class ControlClient: + """Send control requests through reusable, thread-local REQ sockets. + + Attributes: + _context: ZMQ context that owns all client sockets. + _timeout_ms: Send and receive timeout for each exchange. + _local: Thread-local mapping from address to REQ socket. + _closed: Whether the client context has been destroyed. + """ + + def __init__(self, timeout_ms: int) -> None: + self._context = zmq.Context() + self._timeout_ms = timeout_ms + self._local = threading.local() + self._closed = False + + def _sockets(self) -> dict[str, zmq.Socket]: + sockets = getattr(self._local, "sockets", None) + if sockets is None: + sockets = {} + self._local.sockets = sockets + return sockets + + def _discard(self, addr: str) -> None: + socket = self._sockets().pop(addr, None) + if socket is not None: + socket.close(linger=0) + + def _exchange(self, addr: str, payload: ControlRequest) -> ControlResponse: + sockets = self._sockets() + socket = sockets.get(addr) + if socket is None: + socket = self._context.socket(zmq.REQ) + socket.setsockopt(zmq.RCVTIMEO, self._timeout_ms) + socket.setsockopt(zmq.SNDTIMEO, self._timeout_ms) + socket.setsockopt(zmq.LINGER, 0) + socket.connect(addr) + sockets[addr] = socket + try: + socket.send_json(payload) + response = socket.recv_json() + except Exception: + self._discard(addr) + raise + assert isinstance(response, dict) + return cast(ControlResponse, response) + + def request(self, addr: str, payload: ControlRequest) -> Any: + response = self._exchange(addr, payload) + if not response.get("ok"): + raise RuntimeError(response.get("error", "EC control request failed")) + return response.get("result") + + def close(self) -> None: + if self._closed: + return + self._closed = True + self._context.destroy(linger=0) + + +def make_cancel_request( + transfer_id: str, + reservation_id: str, + *, + abandon: bool = False, + refresh: bool = False, +) -> ReservationActionRequest: + request: ReservationActionRequest = { + "op": "cancel", + "transfer_id": transfer_id, + "reservation_id": reservation_id, + } + if abandon: + request["abandon"] = True + if refresh: + request["refresh"] = True + return request + + +class ShardTopology: + """Discover and cache every control address for a Consumer. + + Attributes: + _client: Client used to query the Consumer's ``peers`` operation. + _cache: Base Consumer addresses mapped to all TP-shard addresses. + """ + + def __init__(self, client: ControlClient) -> None: + self._client = client + self._cache: dict[str, list[str]] = {} + + def discover(self, base_addr: str) -> list[str] | None: + """Return a confirmed complete topology, retrying transient failures.""" + cached = self._cache.get(base_addr) + if cached is not None: + return cached + try: + reply = self._client.request(base_addr, {"op": "peers"}) + ports = reply.get("ports") if isinstance(reply, dict) else None + if not isinstance(ports, list) or not ports: + raise ValueError("invalid or empty peer list") + prefix = base_addr.rsplit(":", 1)[0] + shards = [f"{prefix}:{int(port)}" for port in ports] + except Exception: + logger.warning( + "EC Mooncake consumer at %s did not report its shards; " + "using it directly for this attempt.", + base_addr, + exc_info=True, + ) + return None + self._cache[base_addr] = shards + if len(shards) > 1: + logger.info( + "EC Mooncake consumer at %s has %d shards", base_addr, len(shards) + ) + return shards + + def shards(self, base_addr: str) -> list[str]: + """Return confirmed shards or a one-attempt data-plane fallback.""" + return self.discover(base_addr) or [base_addr] + + +class EventInbox: + """Receive Consumer readiness events without blocking the Scheduler. + + Attributes: + _client: Control client used to discover event ports. + _topology: Source of Consumer TP-shard addresses. + _context: Lazily created context for the PULL socket. + _socket: PULL socket connected to every expected Consumer shard. + _closed: Whether event resources have been released. + shard_count: Number of event channels in the complete topology. + """ + + def __init__(self, client: ControlClient, topology: ShardTopology) -> None: + self._client = client + self._topology = topology + self._context: zmq.Context | None = None + self._socket: zmq.Socket | None = None + self._closed = False + self.shard_count = 1 + + def _connect(self, base_addr: str) -> None: + if self._socket is not None: + return + shards = self._topology.discover(base_addr) + if shards is None: + return + endpoints = [] + for addr in shards: + try: + event_port = self._client.request(addr, {"op": "event_port"}) + address, _ = addr.rsplit(":", 1) + endpoints.append(f"{address}:{int(event_port)}") + except Exception: + logger.warning( + "EC Mooncake could not subscribe to the event channel of " + "consumer shard %s; retrying the complete topology later.", + addr, + ) + return + context = zmq.Context() + socket = context.socket(zmq.PULL) + for endpoint in endpoints: + socket.connect(endpoint) + self._context = context + self._socket = socket + self.shard_count = len(shards) + + def drain(self, base_addr: str) -> list[dict[str, Any]]: + self._connect(base_addr) + if self._socket is None: + return [] + events = [] + while True: + try: + events.append(self._socket.recv_json(flags=zmq.DONTWAIT)) + except zmq.Again: + return events + + def close(self) -> None: + if self._closed: + return + self._closed = True + if self._socket is not None: + self._socket.close(linger=0) + if self._context is not None: + self._context.term() + + +class ConsumerControlServer: + """Expose Consumer reservations and readiness events over ZMQ. + + One server runs on every receiving TP rank. The REP channel handles + reservation operations, while a PUSH channel publishes newly ready items + to the Scheduler. + + Attributes: + host: Interface on which the control server listens. + port: Rank-local REP control port. + peer_ports: Control ports for every Consumer TP shard. + event_port: Dynamically allocated PUSH event port after startup. + _device: Device selected in the control thread when CUDA is used. + _reserve: Callback that allocates or reuses destination memory. + _status: Callback that reports active reservation state. + _complete: Callback that marks destination writes complete. + _cancel: Callback that cancels or abandons reservations. + _reap: Callback that expires stale reservations. + _metrics_log_interval: Interval for aggregate control-plane logs. + _stop: Signal requesting termination of the server loop. + _started: Signal indicating that socket binding has completed. + _thread: Background server thread. + _startup_error: Socket binding error captured from the server thread. + """ + + def __init__( + self, + host: str, + port: int, + reserve: Callable[[dict[str, Any]], dict[str, Any]], + status: Callable[[str], dict[str, Any] | None], + complete: Callable[[str, str], ControlCompletion], + cancel: Callable[[str, str, bool, bool], bool], + reap: Callable[[], int], + metrics_log_interval: float = 10, + peer_ports: list[int] | None = None, + device: torch.device | None = None, + ) -> None: + self.host = host + self.port = port + self.peer_ports = peer_ports or [port] + self._device = device + self.event_port: int | None = None + self._reserve = reserve + self._status = status + self._complete = complete + self._cancel = cancel + self._reap = reap + self._metrics_log_interval = metrics_log_interval + self._stop = threading.Event() + self._started = threading.Event() + self._thread: threading.Thread | None = None + self._startup_error: Exception | None = None + + def start(self) -> None: + def loop() -> None: + if self._device is not None and self._device.type == "cuda": + torch.accelerator.set_device_index(self._device.index or 0) + context = zmq.Context() + socket = context.socket(zmq.REP) + event_socket = context.socket(zmq.PUSH) + pending_events: deque[dict[str, Any]] = deque() + metrics: Counter[str] = Counter() + + def queue_event(event: dict[str, Any]) -> None: + event["shard"] = self.port + if len(pending_events) >= _MAX_PENDING_EVENTS: + pending_events.popleft() + metrics["events_dropped"] += 1 + pending_events.append(event) + metrics["events_queued"] += 1 + + metrics_started_at = time.monotonic() + last_reap_at = metrics_started_at + socket.setsockopt(zmq.RCVTIMEO, 100) + try: + socket.bind(f"tcp://{self.host}:{self.port}") + self.event_port = event_socket.bind_to_random_port(f"tcp://{self.host}") + except Exception as e: + self._startup_error = e + self._started.set() + socket.close(linger=0) + event_socket.close(linger=0) + context.term() + return + self._started.set() + try: + while not self._stop.is_set(): + while pending_events: + try: + event_socket.send_json( + pending_events[0], flags=zmq.DONTWAIT + ) + except zmq.Again: + break + pending_events.popleft() + metrics["events_sent"] += 1 + now = time.monotonic() + if now - last_reap_at >= _RESERVATION_REAP_INTERVAL_SECONDS: + metrics["reservations_reaped"] += self._reap() + last_reap_at = now + if ( + self._metrics_log_interval > 0 + and now - metrics_started_at >= self._metrics_log_interval + ): + logger.info( + "EC Mooncake consumer control: requests=%s, " + "events_queued=%d, events_sent=%d, events_dropped=%d, " + "event_backlog=%d, reservations_reaped=%d", + { + key.removeprefix("request_"): value + for key, value in metrics.items() + if key.startswith("request_") + }, + metrics["events_queued"], + metrics["events_sent"], + metrics["events_dropped"], + len(pending_events), + metrics["reservations_reaped"], + ) + metrics.clear() + metrics_started_at = now + try: + request = socket.recv_json() + except zmq.Again: + continue + try: + op = request.get("op") + result: Any = None + metrics[f"request_{op}"] += 1 + if op == "reserve": + result = self._reserve(request) + if result.get("ready"): + transfer_id = str(request["transfer_id"]) + status = self._status(transfer_id) + if status is not None: + queue_event({"transfer_id": transfer_id, **status}) + elif op == "status": + result = self._status(str(request["transfer_id"])) + elif op == "event_port": + result = self.event_port + elif op == "peers": + result = {"ports": self.peer_ports} + elif op in ("complete", "complete_batch"): + items = ( + request["items"] + if op == "complete_batch" + else [request] + ) + completions = [] + for item in items: + transfer_id = str(item["transfer_id"]) + completion = self._complete( + transfer_id, str(item["reservation_id"]) + ) + completions.append( + { + "completed": completion.accepted, + "became_ready": completion.became_ready, + } + ) + if not completion.became_ready: + continue + status = self._status(transfer_id) + if status is not None: + queue_event({"transfer_id": transfer_id, **status}) + result = ( + {"items": completions} + if op == "complete_batch" + else completions[0] + ) + elif op == "cancel": + result = { + "cancelled": self._cancel( + str(request["transfer_id"]), + str(request.get("reservation_id", "")), + bool(request.get("abandon", False)), + bool(request.get("refresh", False)), + ) + } + else: + raise ValueError(f"unknown control op: {op!r}") + socket.send_json({"ok": True, "result": result}) + except Exception as e: + socket.send_json({"ok": False, "error": str(e)}) + finally: + socket.close(linger=0) + event_socket.close(linger=0) + context.term() + + self._thread = threading.Thread( + target=loop, name="ec-mooncake-control", daemon=True + ) + self._thread.start() + if not self._started.wait(timeout=5): + raise RuntimeError("EC Mooncake control channel failed to start") + if self._startup_error is not None: + raise RuntimeError("EC Mooncake control channel failed to bind") from ( + self._startup_error + ) + logger.info( + "EC Mooncake control channel listening on tcp://%s:%d (events tcp://%s:%d)", + self.host, + self.port, + self.host, + self.event_port, + ) + + def close(self) -> None: + if self._thread is None: + return + self._stop.set() + self._thread.join() + self._thread = None diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py new file mode 100644 index 000000000000..1f09e19ebfdd --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py @@ -0,0 +1,616 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Registered-memory allocation and residency for Mooncake transfers. + +The Producer pool stages source tensors in a registered slab. The Consumer +pool owns destination allocations from reservation through publication, +resident reuse, CUDA-safe retirement, and pressure-driven reclamation. +""" + +from __future__ import annotations + +import bisect +import math +import threading +from collections import Counter, OrderedDict +from collections.abc import Callable +from dataclasses import dataclass +from typing import Generic, TypeVar + +import torch + +from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( + MooncakeTransfer, +) +from vllm.logger import init_logger + +logger = init_logger(__name__) + +_T = TypeVar("_T") + + +@dataclass +class MemoryAllocation: + """Describe a tensor view carved from the Consumer receive slab. + + Attributes: + offset: Byte offset of the allocation within the slab. + size: Aligned number of slab bytes owned by the allocation. + tensor: Typed tensor view exposed to transfer and cache code. + """ + + offset: int + size: int + tensor: torch.Tensor + + +@dataclass +class _ResidentEntry(Generic[_T]): + """Store one resident value and its ownership accounting. + + Attributes: + value: Resident value owned by the pool. + nbytes: Capacity charged to the resident pool. + pinned: Whether active cache state prevents LRU eviction. + leases: Number of in-flight reservations borrowing this value. + """ + + value: _T + nbytes: int + pinned: bool = True + leases: int = 0 + + +@dataclass +class ResidentLease(Generic[_T]): + """Represent a borrow of one resident entry. + + Attributes: + key: Cache identifier used to find the canonical resident entry. + _entry: Entry retained even if the canonical mapping is replaced. + _active: Whether this lease still contributes to the reference count. + """ + + key: str + _entry: _ResidentEntry[_T] + _active: bool = True + + @property + def value(self) -> _T: + return self._entry.value + + +class ContiguousAllocator: + """Allocate aligned regions from one contiguous byte range. + + Attributes: + capacity: Total number of bytes managed by the allocator. + alignment: Allocation granularity in bytes. + _free: Sorted free ranges represented as ``(offset, size)`` pairs. + """ + + def __init__(self, capacity: int, alignment: int = 256): + self.capacity = capacity + self.alignment = alignment + self._free = [(0, capacity)] + + def allocate(self, nbytes: int) -> tuple[int, int] | None: + size = math.ceil(nbytes / self.alignment) * self.alignment + for index, (offset, available) in enumerate(self._free): + if size > available: + continue + if size == available: + self._free.pop(index) + else: + self._free[index] = (offset + size, available - size) + return offset, size + return None + + def free(self, offset: int, size: int) -> None: + index = bisect.bisect_left(self._free, (offset, size)) + self._free.insert(index, (offset, size)) + if index + 1 < len(self._free): + next_offset, next_size = self._free[index + 1] + if offset + size == next_offset: + self._free[index] = (offset, size + next_size) + self._free.pop(index + 1) + if index > 0: + previous_offset, previous_size = self._free[index - 1] + current_offset, current_size = self._free[index] + if previous_offset + previous_size == current_offset: + self._free[index - 1] = ( + previous_offset, + previous_size + current_size, + ) + self._free.pop(index) + + +class ResidentPool(Generic[_T]): + """Track resident values, active leases, and LRU eviction eligibility. + + Attributes: + used: Total bytes charged by current and leased displaced entries. + _entries: Canonical resident entries keyed by cache identifier. + _evictable: Unpinned and unleased entries in LRU order. + """ + + def __init__(self): + self.used = 0 + self._entries: dict[str, _ResidentEntry[_T]] = {} + self._evictable: OrderedDict[str, None] = OrderedDict() + + def __len__(self) -> int: + return len(self._entries) + + @property + def num_evictable(self) -> int: + return len(self._evictable) + + def referenced(self) -> list[str]: + return [key for key, entry in self._entries.items() if entry.pinned] + + def get(self, key: str) -> _T | None: + entry = self._entries.get(key) + return entry.value if entry is not None else None + + def insert( + self, + key: str, + value: _T, + nbytes: int, + ) -> _T | None: + """Pin an entry and return a displaced value that has no owners.""" + previous = self._entries.get(key) + entry = _ResidentEntry(value, nbytes) + if previous is not None and previous.leases == 0: + self.used -= previous.nbytes + self._entries[key] = entry + self.used += nbytes + self._evictable.pop(key, None) + if previous is not None and previous.leases == 0: + return previous.value + return None + + def pin(self, key: str) -> _T | None: + entry = self._entries.get(key) + if entry is None: + return None + self._evictable.pop(key, None) + entry.pinned = True + return entry.value + + def acquire(self, key: str) -> ResidentLease[_T] | None: + entry = self._entries.get(key) + if entry is None: + return None + entry.leases += 1 + self._evictable.pop(key, None) + return ResidentLease(key, entry) + + def release(self, lease: ResidentLease[_T]) -> _T | None: + if not lease._active: + return None + lease._active = False + entry = lease._entry + entry.leases -= 1 + current = self._entries.get(lease.key) + if current is not entry: + if entry.leases == 0: + self.used -= entry.nbytes + return entry.value + return None + if not entry.pinned and entry.leases == 0: + self._evictable[lease.key] = None + return None + + def consume(self, lease: ResidentLease[_T]) -> tuple[_T, _T | None]: + current = self._entries[lease.key] + current.pinned = True + self._evictable.pop(lease.key, None) + released = self.release(lease) + return current.value, released + + def retire(self, key: str) -> None: + if key not in self._entries: + return + entry = self._entries[key] + entry.pinned = False + if entry.leases == 0: + self._evictable[key] = None + + def evict_lru(self, evict: Callable[[str, _T], bool]) -> str | None: + for key in list(self._evictable): + entry = self._entries[key] + if not evict(key, entry.value): + continue + self._evictable.pop(key, None) + del self._entries[key] + self.used -= entry.nbytes + return key + return None + + def clear(self) -> None: + self._entries.clear() + self._evictable.clear() + self.used = 0 + + +@dataclass +class StagedSources: + """Own Producer tensor views and their staging-slab regions. + + Attributes: + tensors: Registered tensor views used as Mooncake sources. + regions: Allocator regions released after the write finishes. + """ + + tensors: list[torch.Tensor] + regions: list[tuple[int, int]] + + +class ProducerMemoryPool: + """Own the Producer staging slab and regions carved from it. + + Attributes: + _capacity: Requested staging-slab size in bytes. + _transfer: Data-plane owner used to register the slab. + _pool: Lazily allocated registered byte tensor. + _allocator: Region allocator for the staging slab. + _disabled: Whether initialization failed and fallback is required. + _lock: Lock protecting initialization and region allocation. + """ + + def __init__(self, capacity: int, transfer: MooncakeTransfer) -> None: + self._capacity = capacity + self._transfer = transfer + self._pool: torch.Tensor | None = None + self._allocator: ContiguousAllocator | None = None + self._disabled = False + self._lock = threading.Lock() + + @property + def tensor(self) -> torch.Tensor | None: + return self._pool + + def _ensure_pool(self, device: torch.device) -> None: + if self._pool is not None or self._disabled: + return + with self._lock: + if self._pool is not None or self._disabled: + return + try: + pool = torch.empty(self._capacity, dtype=torch.uint8, device=device) + ret = self._transfer.register_memory(pool) + if ret != 0: + raise RuntimeError(f"Mooncake returned {ret}") + except (RuntimeError, torch.OutOfMemoryError) as error: + self._disabled = True + logger.warning( + "Could not initialize the EC producer staging pool; falling " + "back to per-transfer registration: %s", + error, + ) + return + self._pool = pool + self._allocator = ContiguousAllocator(pool.nbytes) + logger.info( + "Registered %d-byte staging pool for Mooncake EC pushes", + pool.nbytes, + ) + + def _free_regions(self, regions: list[tuple[int, int]]) -> None: + assert self._allocator is not None + for offset, size in regions: + self._allocator.free(offset, size) + + def stage(self, tensors: list[torch.Tensor]) -> StagedSources | None: + """Copy tensors into one registered slab, or return None for fallback.""" + if not tensors: + return StagedSources([], []) + self._ensure_pool(tensors[0].device) + pool = self._pool + allocator = self._allocator + if pool is None or allocator is None: + return None + staged: list[torch.Tensor] = [] + regions: list[tuple[int, int]] = [] + with self._lock: + for tensor in tensors: + region = allocator.allocate(tensor.nbytes) + if region is None: + self._free_regions(regions) + return None + regions.append(region) + staged.append( + pool.narrow(0, region[0], tensor.nbytes) + .view(tensor.dtype) + .view(tensor.shape) + ) + for destination, source in zip(staged, tensors): + destination.copy_(source, non_blocking=True) + return StagedSources(staged, regions) + + def release(self, staged: StagedSources) -> None: + if not staged.regions: + return + with self._lock: + self._free_regions(staged.regions) + + def close(self) -> None: + """Retain the registered slab until the full close phase owns it.""" + + +class ConsumerMemoryPool: + """Own the registered receive slab and resident allocation lifecycle. + + Attributes: + _capacity: Requested receive-slab size in bytes. + _transfer: Data-plane owner used to register the slab. + _metrics: Counters describing resident-cache behavior. + _pool: Registered byte tensor that receives Mooncake writes. + _allocator: Region allocator for the receive slab. + _residents: Published allocations available for local reuse. + _retire_events: CUDA events guarding retired resident entries. + _pending_frees: Allocations waiting for CUDA consumers to finish. + _reclaimed: Cache identifiers evicted under allocation pressure. + _disabled: Whether receive-slab initialization has failed. + lock: Reentrant lock shared with reservation state transitions. + """ + + def __init__( + self, + capacity: int, + transfer: MooncakeTransfer, + ) -> None: + self._capacity = capacity + self._transfer = transfer + self._metrics: Counter[str] = Counter() + self._pool: torch.Tensor | None = None + self._allocator: ContiguousAllocator | None = None + self._residents: ResidentPool[MemoryAllocation] = ResidentPool() + self._retire_events: dict[str, torch.Event] = {} + self._pending_frees: list[tuple[torch.Event, MemoryAllocation]] = [] + self._reclaimed: set[str] = set() + self._disabled = False + self.lock = threading.RLock() + + @property + def tensor(self) -> torch.Tensor | None: + return self._pool + + def prepare( + self, + device: torch.device, + *, + receiving_rank: bool, + allow_host: bool = False, + ) -> None: + if not receiving_rank: + return + if ( + self._pool is not None + or self._disabled + or (device.type != "cuda" and not allow_host) + ): + return + try: + pool = torch.empty(self._capacity, dtype=torch.uint8, device=device) + ret = self._transfer.register_memory(pool) + if ret != 0: + raise RuntimeError(f"Mooncake returned {ret}") + except (RuntimeError, torch.OutOfMemoryError) as error: + self._disabled = True + logger.warning( + "Could not initialize the EC consumer buffer pool; falling back " + "to per-tensor registration: %s", + error, + ) + return + self._pool = pool + self._allocator = ContiguousAllocator(pool.nbytes) + logger.info( + "Prepared %d-byte CUDA receive pool for Mooncake EC (registered=%s)", + pool.nbytes, + receiving_rank, + ) + + def _free(self, allocation: MemoryAllocation) -> None: + assert self._allocator is not None + self._allocator.free(allocation.offset, allocation.size) + + def free(self, allocation: MemoryAllocation) -> None: + with self.lock: + self._free(allocation) + + def _defer_or_free( + self, allocation: MemoryAllocation, event: torch.Event | None + ) -> None: + if event is None or event.query(): + self._free(allocation) + else: + self._pending_frees.append((event, allocation)) + + def _poll_frees_locked(self) -> None: + pending = [] + for event, allocation in self._pending_frees: + if event.query(): + self._free(allocation) + else: + pending.append((event, allocation)) + self._pending_frees = pending + + def _reclaim_locked(self, nbytes: int) -> tuple[int, int] | None: + assert self._allocator is not None + + def evict(mm_hash: str, allocation: MemoryAllocation) -> bool: + event = self._retire_events.pop(mm_hash, None) + self._defer_or_free(allocation, event) + self._reclaimed.add(mm_hash) + self._metrics["residents_reclaimed"] += 1 + return True + + while self._residents.evict_lru(evict) is not None: + region = self._allocator.allocate(nbytes) + if region is not None: + return region + return None + + def _make_allocation( + self, + region: tuple[int, int], + nbytes: int, + shape: tuple[int, ...], + dtype: torch.dtype, + ) -> MemoryAllocation: + assert self._pool is not None + offset, size = region + tensor = self._pool.narrow(0, offset, nbytes).view(dtype).view(shape) + return MemoryAllocation(offset, size, tensor) + + def try_allocate( + self, nbytes: int, shape: tuple[int, ...], dtype: torch.dtype + ) -> MemoryAllocation | None: + with self.lock: + allocator = self._allocator + assert self._pool is not None and allocator is not None + self._poll_frees_locked() + region = allocator.allocate(nbytes) + if region is None: + return None + return self._make_allocation(region, nbytes, shape, dtype) + + def reclaim_and_allocate( + self, nbytes: int, shape: tuple[int, ...], dtype: torch.dtype + ) -> MemoryAllocation | None: + with self.lock: + region = self._reclaim_locked(nbytes) + if region is None: + return None + return self._make_allocation(region, nbytes, shape, dtype) + + def acquire_cached( + self, mm_hash: str, shape: tuple[int, ...], dtype: torch.dtype + ) -> ResidentLease[MemoryAllocation] | None: + with self.lock: + allocation = self._residents.get(mm_hash) + if allocation is None: + return None + if ( + tuple(allocation.tensor.shape) != shape + or allocation.tensor.dtype != dtype + ): + raise ValueError("conflicting cached tensor for mm_hash") + return self._residents.acquire(mm_hash) + + def take_resident( + self, mm_hash: str, shape: tuple[int, ...], dtype_name: str + ) -> torch.Tensor | None: + with self.lock: + allocation = self._residents.get(mm_hash) + if allocation is None: + self._metrics["residents_missed"] += 1 + return None + tensor = allocation.tensor + if ( + tuple(tensor.shape) != shape + or str(tensor.dtype).split(".")[-1] != dtype_name + ): + self._metrics["residents_mismatched"] += 1 + return None + self._residents.pin(mm_hash) + self._retire_events.pop(mm_hash, None) + self._metrics["residents_promoted"] += 1 + return tensor + + def _record_release_event(self) -> torch.Event | None: + if self._pool is None or self._pool.device.type != "cuda": + return None + event = torch.Event() + event.record(torch.accelerator.current_stream(self._pool.device)) + return event + + def release_cached(self, lease: ResidentLease[MemoryAllocation]) -> None: + with self.lock: + released = self._residents.release(lease) + if released is not None: + self._defer_or_free(released, self._record_release_event()) + + def publish( + self, + mm_hash: str, + allocation: MemoryAllocation, + lease: ResidentLease[MemoryAllocation] | None = None, + ) -> MemoryAllocation: + with self.lock: + if lease is not None: + canonical, released = self._residents.consume(lease) + self._retire_events.pop(mm_hash, None) + if released is not None: + self._defer_or_free(released, self._record_release_event()) + return canonical + previous = self._residents.get(mm_hash) + displaced = self._residents.insert(mm_hash, allocation, allocation.size) + event = None + if previous is not None and previous is not allocation: + event = self._retire_events.pop(mm_hash, None) + if displaced is None or displaced is allocation: + return allocation + if event is None: + event = self._record_release_event() + self._defer_or_free(displaced, event) + return allocation + + def retire_stale( + self, + encoder_cache: dict[str, torch.Tensor], + reserved_hashes: set[str], + ) -> None: + if self._pool is None: + return + with self.lock: + for mm_hash in self._residents.referenced(): + allocation = self._residents.get(mm_hash) + if allocation is None: + continue + if encoder_cache.get(mm_hash) is allocation.tensor: + continue + if mm_hash in reserved_hashes: + continue + event = torch.Event() + event.record(torch.accelerator.current_stream(self._pool.device)) + self._retire_events[mm_hash] = event + self._residents.retire(mm_hash) + self._metrics["residents_retired"] += 1 + self._poll_frees_locked() + + def drain_reclaimed(self) -> set[str]: + with self.lock: + reclaimed = self._reclaimed + self._reclaimed = set() + return reclaimed + + def stats(self) -> tuple[int, int, int, int]: + with self.lock: + return ( + len(self._residents), + len(self._residents.referenced()), + self._residents.num_evictable, + len(self._pending_frees), + ) + + def take_metrics(self) -> dict[str, int]: + with self.lock: + metrics = dict(self._metrics) + self._metrics.clear() + return metrics + + def close(self) -> None: + with self.lock: + pool = self._pool + if pool is None or not self._transfer.unregister_memory(pool): + return + self._pool = None + self._allocator = None + self._residents.clear() + self._retire_events.clear() + self._pending_frees.clear() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py new file mode 100644 index 000000000000..96923e415856 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py @@ -0,0 +1,118 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Metadata exchanged by the Mooncake encoder-cache connector.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +from vllm.distributed.ec_transfer.ec_connector.base import ( + ECConnectorMetadata, + ECConnectorWorkerMetadata, +) + + +@dataclass +class ECMooncakeLoadSpec: + """Describe one Consumer-side cache load requested by the Scheduler. + + Attributes: + mm_hash: Stable identifier of the multimodal encoder item. + num_token: Number of encoder tokens expected by the request. + nbytes: Tensor payload size in bytes. + shape: Tensor shape reconstructed by the Consumer. + dtype: Unqualified ``torch.dtype`` name used for reconstruction. + pushed: Whether the tensor came from a remote Producer reservation. + transfer_id: Identity shared by Scheduler, Producer, and Consumer. + reservation_id: Consumer-issued capability for completing the write. + local: Whether to reuse a tensor already resident on the Consumer. + """ + + mm_hash: str + num_token: int + nbytes: int + shape: tuple[int, ...] + dtype: str + pushed: bool = False + transfer_id: str = "" + reservation_id: str = "" + # The consumer pool still holds this item, so the load is a local handoff: + # no transfer, no producer. + local: bool = False + + +@dataclass +class ECMooncakePushSpec: + """Describe a destination reservation prepared before a tensor is ready. + + Attributes: + mm_hash: Stable identifier of the multimodal encoder item. + nbytes: Number of bytes the Consumer must reserve. + shape: Shape of the tensor that will be written. + dtype: Unqualified ``torch.dtype`` name of the tensor. + consumer_zmq: Base control address of the destination Consumer. + transfer_id: Identity shared by Scheduler, Producer, and Consumer. + request_id: Request that owns the push and may cancel it. + """ + + mm_hash: str + nbytes: int + shape: tuple[int, ...] + dtype: str + consumer_zmq: str + transfer_id: str + request_id: str = "" + + +@dataclass +class ECMooncakeConnectorMetadata(ECConnectorMetadata): + """Worker operations emitted for one Scheduler step. + + Attributes: + loads: Consumer loads that should be attached to ``encoder_cache``. + pushes: Producer reservations that should begin before sources arrive. + """ + + loads: list[ECMooncakeLoadSpec] = field(default_factory=list) + pushes: list[ECMooncakePushSpec] = field(default_factory=list) + + def add_load(self, spec: ECMooncakeLoadSpec) -> None: + self.loads.append(spec) + + def add_push(self, spec: ECMooncakePushSpec) -> None: + self.pushes.append(spec) + + +@dataclass +class ECMooncakeWorkerMetadata(ECConnectorWorkerMetadata): + """Completion state reported from Workers to the Scheduler. + + Attributes: + loaded: Cache identifiers loaded successfully on this Worker. + failed_loads: Cache identifiers that could not be loaded. + reclaimed: Resident items evicted because the receive pool was full. + pending_loads: Whether this Worker still owns asynchronous load work. + pending_saves: Whether this Worker still owns asynchronous push work. + """ + + loaded: set[str] = field(default_factory=set) + failed_loads: set[str] = field(default_factory=set) + # Items the receive pool dropped under pressure. The scheduler assumes an + # evicted item stays resident until told otherwise. + reclaimed: set[str] = field(default_factory=set) + pending_loads: bool = False + pending_saves: bool = False + + def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadata: + assert isinstance(other, ECMooncakeWorkerMetadata) + return ECMooncakeWorkerMetadata( + # Every tensor-parallel rank gathers the embedding from its own + # cache, so an item counts as loaded only where all of them have + # it; one rank falling short must fail the load rather than leave + # the scheduler believing it is ready. + loaded=self.loaded & other.loaded, + failed_loads=self.failed_loads | other.failed_loads, + reclaimed=self.reclaimed | other.reclaimed, + pending_loads=self.pending_loads or other.pending_loads, + pending_saves=self.pending_saves or other.pending_saves, + ) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py new file mode 100644 index 000000000000..16949916870c --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py @@ -0,0 +1,427 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Producer-side push lifecycle, futures, and source-tensor ownership. + +This module records when a push may advance or release its source. Worker +orchestration performs the actual control exchanges and Mooncake writes. +""" + +from __future__ import annotations + +import threading +import time +from collections import OrderedDict +from collections.abc import Callable +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import suppress +from dataclasses import dataclass, field +from enum import Enum, auto +from typing import Any + +import torch + +from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakePushSpec, +) + + +class ProducerPushState(Enum): + """Lifecycle of one Producer push from reservation to terminal state.""" + + RESERVING = auto() + WAITING_SOURCE = auto() + WRITING = auto() + NOTIFYING = auto() + DONE = auto() + CANCEL_PENDING = auto() + CANCELLED = auto() + FAILED = auto() + + +@dataclass +class ProducerSourceLease: + """Keep an encoder tensor alive until every remote write is settled. + + Attributes: + tensor: Source encoder tensor owned by the push. + ready_event: CUDA event proving that production of the tensor finished. + """ + + tensor: torch.Tensor + ready_event: torch.Event | None + + +@dataclass +class ProducerPushRecord: + """Collect all asynchronous state for one Producer push. + + Attributes: + spec: Immutable identity and destination metadata for the push. + state: Current Producer lifecycle state. + reservation_futures: Futures resolving Consumer shard reservations. + reservations: Resolved destination descriptors for every shard. + shard_futures: Data-plane futures that may still read the source. + source: Source tensor lease once encoder computation has completed. + batch_future: Transfer or cancellation batch currently owning the push. + error: First asynchronous error retained for Worker reporting. + source_at: Time the source became available for queue metrics. + """ + + spec: ECMooncakePushSpec + state: ProducerPushState + reservation_futures: list[Future[list[dict[str, Any]]]] + reservations: list[dict[str, Any]] = field(default_factory=list) + shard_futures: list[Future[Any]] = field(default_factory=list) + source: ProducerSourceLease | None = None + batch_future: Future[None] | None = None + error: str | None = None + source_at: float | None = None + + +_ALLOWED_TRANSITIONS = { + ProducerPushState.RESERVING: { + ProducerPushState.WAITING_SOURCE, + ProducerPushState.CANCEL_PENDING, + ProducerPushState.FAILED, + }, + ProducerPushState.WAITING_SOURCE: { + ProducerPushState.WRITING, + ProducerPushState.CANCEL_PENDING, + ProducerPushState.FAILED, + }, + ProducerPushState.WRITING: { + ProducerPushState.NOTIFYING, + ProducerPushState.FAILED, + }, + ProducerPushState.NOTIFYING: { + ProducerPushState.DONE, + ProducerPushState.FAILED, + }, + ProducerPushState.CANCEL_PENDING: {ProducerPushState.CANCELLED}, + ProducerPushState.DONE: set(), + ProducerPushState.CANCELLED: set(), + ProducerPushState.FAILED: set(), +} + +_TERMINAL_STATES = { + ProducerPushState.DONE, + ProducerPushState.CANCELLED, + ProducerPushState.FAILED, +} +_SOURCE_WAIT_STATES = { + ProducerPushState.RESERVING, + ProducerPushState.WAITING_SOURCE, +} +_TERMINAL_LIMIT = 1 << 16 + + +class ProducerPushManager: + """Own Producer push records, transitions, and source tensor leases. + + Attributes: + _records: All active and retained terminal records by transfer ID. + _active_ids: Non-terminal transfer IDs in insertion order. + _reapable_terminal_ids: Terminal records safe to discard. + _unreported_ids: Failed records awaiting Worker error reporting. + _batch_ids: Records whose batch future has not been reaped. + _source_waiters: Transfer IDs waiting for each cache identifier. + _lock: Reentrant lock protecting lifecycle and ownership changes. + """ + + def __init__(self) -> None: + self._records: OrderedDict[str, ProducerPushRecord] = OrderedDict() + self._active_ids: OrderedDict[str, None] = OrderedDict() + self._reapable_terminal_ids: OrderedDict[str, None] = OrderedDict() + self._unreported_ids: OrderedDict[str, None] = OrderedDict() + self._batch_ids: OrderedDict[str, None] = OrderedDict() + self._source_waiters: dict[str, OrderedDict[str, None]] = {} + self._lock = threading.RLock() + + def get(self, transfer_id: str) -> ProducerPushRecord | None: + with self._lock: + return self._records.get(transfer_id) + + def reserve( + self, + spec: ECMooncakePushSpec, + submit: Callable[[], Future[list[dict[str, Any]]]], + ) -> tuple[ProducerPushRecord, bool]: + with self._lock: + existing = self._records.get(spec.transfer_id) + if existing is not None: + if existing.spec != spec: + raise ValueError( + f"Producer transfer {spec.transfer_id!r} changed identity" + ) + return existing, False + record = ProducerPushRecord( + spec=spec, + state=ProducerPushState.RESERVING, + reservation_futures=[submit()], + ) + self._records[spec.transfer_id] = record + self._active_ids[spec.transfer_id] = None + self._source_waiters.setdefault(spec.mm_hash, OrderedDict())[ + spec.transfer_id + ] = None + record.reservation_futures[0].add_done_callback( + lambda future: self._reservation_done(record, future) + ) + return record, True + + def bind_source( + self, + mm_hash: str, + tensor: torch.Tensor, + ready_event: torch.Event | None, + ) -> None: + with self._lock: + waiters = self._source_waiters.pop(mm_hash, OrderedDict()) + for transfer_id in waiters: + record = self._records[transfer_id] + if record.source is not None or record.state not in _SOURCE_WAIT_STATES: + continue + record.source = ProducerSourceLease(tensor, ready_event) + record.source_at = time.monotonic() + + def submit_batches( + self, + executor: ThreadPoolExecutor, + run_batch: Callable[[list[ProducerPushRecord]], None], + on_submit: Callable[[], None], + ) -> None: + with self._lock: + grouped: dict[str, list[ProducerPushRecord]] = {} + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if ( + record.source is not None + and record.batch_future is None + and record.state in _SOURCE_WAIT_STATES + ): + grouped.setdefault(record.spec.consumer_zmq, []).append(record) + batches = list(grouped.values()) + for records in batches: + on_submit() + future = executor.submit(run_batch, records) + for record in records: + record.batch_future = future + self._batch_ids[record.spec.transfer_id] = None + + def resolve_reservations(self, record: ProducerPushRecord) -> list[dict[str, Any]]: + results: list[dict[str, Any]] = [] + error: Exception | None = None + for future in record.reservation_futures: + try: + results.extend(future.result()) + except Exception as exc: + if error is None: + error = exc + if error is not None: + raise error + with self._lock: + record.reservations = results + if record.state is ProducerPushState.RESERVING: + self._transition(record, ProducerPushState.WAITING_SOURCE) + return list(results) + + def _reservation_done( + self, + record: ProducerPushRecord, + future: Future[list[dict[str, Any]]], + ) -> None: + try: + future.result() + except Exception as exc: + with self._lock: + self._set_error(record, exc) + if ( + record.state is ProducerPushState.RESERVING + and record.source is None + ): + self._transition(record, ProducerPushState.FAILED) + return + with self._lock: + if record.state is ProducerPushState.RESERVING: + self._transition(record, ProducerPushState.WAITING_SOURCE) + + def settle_all(self, records: list[ProducerPushRecord]) -> None: + for record in records: + for future in record.reservation_futures: + with suppress(Exception): + future.result() + + def replace_reservations( + self, + record: ProducerPushRecord, + reservations: list[dict[str, Any]], + ) -> None: + with self._lock: + record.reservations = reservations + + def track_shard_futures( + self, + records: list[ProducerPushRecord], + futures: list[Future[Any]], + ) -> None: + with self._lock: + for record in records: + record.shard_futures.extend(futures) + + def begin_writing(self, record: ProducerPushRecord) -> None: + with self._lock: + if record.source is None: + raise RuntimeError( + f"Producer push {record.spec.transfer_id!r} has no source tensor" + ) + self._transition(record, ProducerPushState.WRITING) + + def begin_notifying(self, records: list[ProducerPushRecord]) -> None: + with self._lock: + for record in records: + self._transition(record, ProducerPushState.NOTIFYING) + + def complete(self, records: list[ProducerPushRecord]) -> None: + with self._lock: + self._check_sources_releasable(records) + for record in records: + self._transition(record, ProducerPushState.DONE) + self._release_source(record) + + def fail(self, records: list[ProducerPushRecord], error: Exception) -> None: + with self._lock: + self._check_sources_releasable(records) + for record in records: + if record.state not in _TERMINAL_STATES: + self._set_error(record, error) + self._transition(record, ProducerPushState.FAILED) + self._release_source(record) + + def cancel_requests(self, request_ids: set[str]) -> list[ProducerPushRecord]: + cancelled = [] + with self._lock: + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if ( + record.spec.request_id not in request_ids + or record.source is not None + or record.batch_future is not None + ): + continue + if record.state not in _SOURCE_WAIT_STATES: + continue + self._transition(record, ProducerPushState.CANCEL_PENDING) + cancelled.append(record) + return cancelled + + def finish_cancel(self, record: ProducerPushRecord) -> None: + with self._lock: + if record.state is ProducerPushState.CANCEL_PENDING: + self._transition(record, ProducerPushState.CANCELLED) + + def submit_cancel( + self, + record: ProducerPushRecord, + executor: ThreadPoolExecutor, + run_cancel: Callable[[ProducerPushRecord], None], + ) -> None: + with self._lock: + future = executor.submit(run_cancel, record) + record.batch_future = future + transfer_id = record.spec.transfer_id + self._batch_ids[transfer_id] = None + if record.state in _TERMINAL_STATES: + self._reapable_terminal_ids.pop(transfer_id, None) + + def poll(self) -> list[tuple[str, str]]: + failures = [] + with self._lock: + for transfer_id in list(self._batch_ids): + record = self._records[transfer_id] + future = record.batch_future + assert future is not None + if not future.done(): + continue + try: + future.result() + except Exception as exc: + self._set_error(record, exc) + self._batch_ids.pop(transfer_id) + self._mark_reapable(record) + for transfer_id in list(self._unreported_ids): + if transfer_id in self._batch_ids: + continue + record = self._records[transfer_id] + assert record.error is not None + failures.append((record.spec.mm_hash, record.error)) + self._unreported_ids.pop(transfer_id) + self._mark_reapable(record) + self._reap_terminals() + return failures + + @property + def pending(self) -> bool: + with self._lock: + return bool(self._active_ids or self._batch_ids) + + def _release_source(self, record: ProducerPushRecord) -> None: + if record.source is not None: + record.source = None + + def _set_error(self, record: ProducerPushRecord, error: BaseException) -> None: + if record.error is None: + record.error = str(error) + if record.state in _TERMINAL_STATES: + transfer_id = record.spec.transfer_id + self._unreported_ids[transfer_id] = None + self._reapable_terminal_ids.pop(transfer_id, None) + + def _reap_terminals(self) -> None: + while len(self._reapable_terminal_ids) > _TERMINAL_LIMIT: + transfer_id, _ = self._reapable_terminal_ids.popitem(last=False) + self._records.pop(transfer_id) + + def _mark_reapable(self, record: ProducerPushRecord) -> None: + transfer_id = record.spec.transfer_id + if ( + record.state in _TERMINAL_STATES + and transfer_id not in self._batch_ids + and transfer_id not in self._unreported_ids + ): + self._reapable_terminal_ids[transfer_id] = None + + def _drop_source_waiter(self, record: ProducerPushRecord) -> None: + waiters = self._source_waiters.get(record.spec.mm_hash) + if waiters is None: + return + waiters.pop(record.spec.transfer_id, None) + if not waiters: + self._source_waiters.pop(record.spec.mm_hash) + + @staticmethod + def _check_sources_releasable(records: list[ProducerPushRecord]) -> None: + for record in records: + futures = [*record.reservation_futures, *record.shard_futures] + if record.source is not None and not all( + future.done() for future in futures + ): + raise RuntimeError( + f"Producer push {record.spec.transfer_id!r} released its " + "source too early" + ) + + def _transition(self, record: ProducerPushRecord, state: ProducerPushState) -> None: + if state not in _ALLOWED_TRANSITIONS[record.state]: + raise RuntimeError( + f"Cannot transition producer push {record.spec.transfer_id!r} from " + f"{record.state.name} to {state.name}" + ) + record.state = state + if state not in _SOURCE_WAIT_STATES: + self._drop_source_waiter(record) + if state in _TERMINAL_STATES: + transfer_id = record.spec.transfer_id + self._active_ids.pop(transfer_id, None) + if record.error is not None: + self._unreported_ids[transfer_id] = None + self._mark_reapable(record) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py new file mode 100644 index 000000000000..2873f55a89d1 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py @@ -0,0 +1,441 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Consumer-side reservation lifecycle and destination-memory ownership. + +Reservations make remote writes idempotent and ensure cancellation or expiry +cannot free a destination while Mooncake may still be writing into it. +""" + +from __future__ import annotations + +import time +import uuid +from collections import OrderedDict +from dataclasses import dataclass, field +from enum import Enum, auto + +import torch + +from vllm.distributed.ec_transfer.ec_connector.mooncake.memory import ( + ConsumerMemoryPool, + MemoryAllocation, + ResidentLease, +) + + +class ConsumerReservationState(Enum): + """Lifecycle of one Consumer destination reservation.""" + + RESERVED = auto() + WRITING = auto() + READY = auto() + TAKEN = auto() + RESIDENT = auto() + CANCEL_PENDING = auto() + EXPIRE_PENDING = auto() + CANCELLED = auto() + EXPIRED = auto() + + +_ALLOWED_TRANSITIONS = { + ConsumerReservationState.RESERVED: { + ConsumerReservationState.WRITING, + ConsumerReservationState.CANCELLED, + }, + ConsumerReservationState.WRITING: { + ConsumerReservationState.READY, + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, + ConsumerReservationState.CANCELLED, + }, + ConsumerReservationState.READY: { + ConsumerReservationState.TAKEN, + ConsumerReservationState.CANCELLED, + ConsumerReservationState.EXPIRED, + }, + ConsumerReservationState.TAKEN: {ConsumerReservationState.RESIDENT}, + ConsumerReservationState.RESIDENT: set(), + ConsumerReservationState.CANCEL_PENDING: {ConsumerReservationState.CANCELLED}, + ConsumerReservationState.EXPIRE_PENDING: { + ConsumerReservationState.CANCELLED, + ConsumerReservationState.EXPIRED, + }, + ConsumerReservationState.CANCELLED: set(), + ConsumerReservationState.EXPIRED: set(), +} + +_ACTIVE_STATES = { + ConsumerReservationState.WRITING, + ConsumerReservationState.READY, + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, +} +_DEFERRED_STATES = { + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, +} + + +@dataclass +class ConsumerReservation: + """Own the identity and destination allocation of one remote push. + + Attributes: + transfer_id: Cross-process identity of the transfer. + mm_hash: Stable identifier of the encoder-cache item. + reservation_id: Consumer-issued identity required for completion. + state: Current destination reservation state. + shape: Expected tensor shape. + dtype: Unqualified expected ``torch.dtype`` name. + allocation: Receive-slab allocation while this record owns it. + lease: Borrowed resident allocation used for a cache hit. + created_at: Monotonic creation time used for diagnostics. + expires_at: Deadline for the active record or terminal tombstone. + """ + + transfer_id: str + mm_hash: str + reservation_id: str + state: ConsumerReservationState + shape: tuple[int, ...] = () + dtype: str = "" + allocation: MemoryAllocation | None = None + lease: ResidentLease[MemoryAllocation] | None = None + created_at: float = field(default_factory=time.monotonic) + expires_at: float = 0 + + +@dataclass(frozen=True) +class CompletionResult: + """Describe how a completion request affected a reservation. + + Attributes: + accepted: Whether transfer and reservation identities matched. + became_ready: Whether the call transitioned WRITING to READY. + repeated: Whether the reservation was already ready. + discarded: Whether deferred cancellation consumed the completion. + """ + + accepted: bool + became_ready: bool = False + repeated: bool = False + discarded: bool = False + + +class CancellationOutcome(Enum): + """Outcome categories used for control responses and metrics.""" + + REJECTED = auto() + PRE_RESERVED = auto() + DEFERRED = auto() + CANCELLED = auto() + + +class ConsumerReservationManager: + """Own reservation transitions and destination allocation releases. + + Attributes: + _memory: Consumer memory pool that owns destination allocations. + _lease_ttl: Lifetime of active reservations and terminal tombstones. + _tombstone_limit: Maximum retained terminal cancellation records. + _records: Active and terminal records keyed by transfer ID. + _active_ids: Transfer IDs requiring expiry scans. + _tombstones: Terminal records retained in expiry order. + """ + + def __init__( + self, + memory: ConsumerMemoryPool, + lease_ttl: float, + tombstone_limit: int, + ) -> None: + self._memory = memory + self._lease_ttl = lease_ttl + self._tombstone_limit = tombstone_limit + self._records: dict[str, ConsumerReservation] = {} + self._active_ids: dict[str, None] = {} + self._tombstones: OrderedDict[str, None] = OrderedDict() + + def get(self, transfer_id: str) -> ConsumerReservation | None: + return self._records.get(transfer_id) + + def active_records(self) -> list[ConsumerReservation]: + return [self._records[transfer_id] for transfer_id in self._active_ids] + + def reserve( + self, + transfer_id: str, + mm_hash: str, + nbytes: int, + shape: tuple[int, ...], + dtype_name: str, + dtype: torch.dtype, + ) -> tuple[ConsumerReservation | None, bool, bool, tuple[int, int, int]]: + with self._memory.lock: + existing = self._records.get(transfer_id) + if ( + existing is not None + and existing.state is ConsumerReservationState.CANCELLED + ): + return existing, False, False, (0, 0, 0) + if ( + existing is not None + and existing.state is ConsumerReservationState.EXPIRED + ): + self._remove(transfer_id) + existing = None + if ( + existing is not None + and existing.state is ConsumerReservationState.EXPIRE_PENDING + ): + raise RuntimeError( + f"Transfer {transfer_id!r} still has an active writer" + ) + if existing is not None: + if ( + existing.mm_hash != mm_hash + or existing.shape != shape + or existing.dtype != dtype_name + ): + raise ValueError("conflicting reservation for transfer_id") + if existing.state not in _ACTIVE_STATES: + raise RuntimeError( + f"Cannot reserve transfer {transfer_id!r} in " + f"{existing.state.name}" + ) + if existing.state is ConsumerReservationState.WRITING: + existing.expires_at = time.monotonic() + self._lease_ttl + return existing, False, True, (0, 0, 0) + + lease = self._memory.acquire_cached(mm_hash, shape, dtype) + now = time.monotonic() + if lease is not None: + record = ConsumerReservation( + transfer_id, + mm_hash, + uuid.uuid4().hex, + ConsumerReservationState.READY, + shape, + dtype_name, + lease.value, + lease, + now, + now + self._lease_ttl, + ) + self._insert(record) + return record, False, False, (0, 0, 0) + expiry_counts = (0, 0, 0) + allocation = self._memory.try_allocate(nbytes, shape, dtype) + if allocation is None: + expiry_counts = self._expire_locked(time.monotonic()) + allocation = self._memory.try_allocate(nbytes, shape, dtype) + if allocation is None: + allocation = self._memory.reclaim_and_allocate(nbytes, shape, dtype) + if allocation is None: + return None, False, False, expiry_counts + record = ConsumerReservation( + transfer_id, + mm_hash, + uuid.uuid4().hex, + ConsumerReservationState.RESERVED, + shape, + dtype_name, + allocation, + None, + now, + now + self._lease_ttl, + ) + self._transition(record, ConsumerReservationState.WRITING) + self._insert(record) + return record, True, False, expiry_counts + + def status(self, transfer_id: str) -> ConsumerReservation | None: + with self._memory.lock: + record = self._records.get(transfer_id) + if record is None or record.state not in _ACTIVE_STATES: + return None + return record + + def complete(self, transfer_id: str, reservation_id: str) -> CompletionResult: + with self._memory.lock: + record = self._records.get(transfer_id) + if record is None or record.reservation_id != reservation_id: + return CompletionResult(False) + if record.state is ConsumerReservationState.READY: + return CompletionResult(True, repeated=True) + if record.state in _DEFERRED_STATES: + terminal = ( + ConsumerReservationState.CANCELLED + if record.state is ConsumerReservationState.CANCEL_PENDING + else ConsumerReservationState.EXPIRED + ) + self._terminate(record, terminal) + return CompletionResult(True, discarded=True) + if record.state is not ConsumerReservationState.WRITING: + return CompletionResult(False) + self._transition(record, ConsumerReservationState.READY) + record.expires_at = time.monotonic() + self._lease_ttl + return CompletionResult(True, became_ready=True) + + def cancel( + self, + transfer_id: str, + reservation_id: str, + abandon: bool = False, + refresh: bool = False, + ) -> tuple[CancellationOutcome, int]: + with self._memory.lock: + record = self._records.get(transfer_id) + if ( + record is not None + and reservation_id + and record.reservation_id != reservation_id + ): + return CancellationOutcome.REJECTED, 0 + if record is None: + now = time.monotonic() + record = ConsumerReservation( + transfer_id, + "", + "", + ConsumerReservationState.CANCELLED, + created_at=now, + ) + self._insert(record) + self._set_tombstone_deadline(record) + dropped = self._reap_tombstones(now) + return CancellationOutcome.PRE_RESERVED, dropped + if record.state is ConsumerReservationState.CANCELLED: + self._set_tombstone_deadline(record) + dropped = self._reap_tombstones(time.monotonic()) + return CancellationOutcome.PRE_RESERVED, dropped + if refresh: + if not abandon or record.state not in { + ConsumerReservationState.WRITING, + ConsumerReservationState.EXPIRE_PENDING, + }: + return CancellationOutcome.REJECTED, 0 + if record.state is ConsumerReservationState.WRITING: + self._transition(record, ConsumerReservationState.EXPIRE_PENDING) + self._terminate(record, ConsumerReservationState.EXPIRED) + dropped = self._reap_tombstones(time.monotonic()) + return CancellationOutcome.CANCELLED, dropped + if record.state in _DEFERRED_STATES and not abandon: + return CancellationOutcome.DEFERRED, 0 + if record.state is ConsumerReservationState.WRITING and not abandon: + self._transition(record, ConsumerReservationState.CANCEL_PENDING) + return CancellationOutcome.DEFERRED, 0 + if ( + record.state not in _ACTIVE_STATES + and record.state is not ConsumerReservationState.RESERVED + ): + return CancellationOutcome.REJECTED, 0 + self._terminate(record, ConsumerReservationState.CANCELLED) + dropped = self._reap_tombstones(time.monotonic()) + return CancellationOutcome.CANCELLED, dropped + + def take(self, transfer_id: str, mm_hash: str) -> MemoryAllocation: + with self._memory.lock: + record = self._records.get(transfer_id) + if ( + record is None + or record.state is not ConsumerReservationState.READY + or record.mm_hash != mm_hash + or record.allocation is None + ): + raise RuntimeError( + f"Pushed EC tensor is not ready for mm_hash={mm_hash}" + ) + self._transition(record, ConsumerReservationState.TAKEN) + allocation = self._memory.publish(mm_hash, record.allocation, record.lease) + record.allocation = None + record.lease = None + self._transition(record, ConsumerReservationState.RESIDENT) + self._remove(transfer_id) + return allocation + + def expire(self) -> tuple[int, int, int]: + with self._memory.lock: + return self._expire_locked(time.monotonic()) + + def retire_stale(self, encoder_cache: dict[str, torch.Tensor]) -> None: + with self._memory.lock: + reserved_hashes = {record.mm_hash for record in self.active_records()} + self._memory.retire_stale(encoder_cache, reserved_hashes) + + def _expire_locked(self, now: float) -> tuple[int, int, int]: + expired = 0 + deferred = 0 + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if record.expires_at > now: + continue + if record.state is ConsumerReservationState.READY: + self._terminate(record, ConsumerReservationState.EXPIRED) + expired += 1 + elif record.state is ConsumerReservationState.WRITING: + self._transition(record, ConsumerReservationState.EXPIRE_PENDING) + deferred += 1 + dropped = self._reap_tombstones(now) + return expired, deferred, dropped + + def _terminate( + self, record: ConsumerReservation, state: ConsumerReservationState + ) -> None: + self._transition(record, state) + self._release(record) + self._set_tombstone_deadline(record) + + def _release(self, record: ConsumerReservation) -> None: + allocation = record.allocation + if allocation is None: + return + if record.lease is not None: + self._memory.release_cached(record.lease) + else: + self._memory.free(allocation) + record.allocation = None + record.lease = None + + def _set_tombstone_deadline(self, record: ConsumerReservation) -> None: + record.expires_at = time.monotonic() + self._lease_ttl + self._tombstones[record.transfer_id] = None + self._tombstones.move_to_end(record.transfer_id) + + def _reap_tombstones(self, now: float) -> int: + dropped = 0 + while self._tombstones: + transfer_id = next(iter(self._tombstones)) + record = self._records[transfer_id] + if ( + record.expires_at > now + and len(self._tombstones) <= self._tombstone_limit + ): + break + self._remove(transfer_id) + dropped += 1 + return dropped + + def _insert(self, record: ConsumerReservation) -> None: + self._records[record.transfer_id] = record + if record.state in _ACTIVE_STATES: + self._active_ids[record.transfer_id] = None + + def _remove(self, transfer_id: str) -> None: + self._records.pop(transfer_id, None) + self._active_ids.pop(transfer_id, None) + self._tombstones.pop(transfer_id, None) + + def _transition( + self, record: ConsumerReservation, state: ConsumerReservationState + ) -> None: + if state not in _ALLOWED_TRANSITIONS[record.state]: + raise RuntimeError( + f"Cannot transition {record.transfer_id!r} from " + f"{record.state.name} to {state.name}" + ) + record.state = state + if state in _ACTIVE_STATES: + self._active_ids[record.transfer_id] = None + else: + self._active_ids.pop(record.transfer_id, None) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py new file mode 100644 index 000000000000..acc226a7d85d --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py @@ -0,0 +1,636 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Scheduler-side planning and observation for Mooncake cache transfers. + +The Scheduler prepares Producer reservations, consumes Consumer readiness +events, emits per-step Worker metadata, and converts Worker results back into +request availability without owning tensor memory or running data transfers. +""" + +from __future__ import annotations + +import math +import time +from collections import Counter, OrderedDict +from collections.abc import Collection +from concurrent.futures import Future, ThreadPoolExecutor +from typing import TYPE_CHECKING, Any + +import torch + +from vllm.distributed.ec_transfer.ec_connector.base import ( + ECConnectorMetadata, + ECConnectorRole, +) +from vllm.distributed.ec_transfer.ec_connector.cpu.common import ( + _get_encoder_cache_hidden_dim, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake._availability import ( + ensure_mooncake_available, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.config import MooncakeECConfig +from vllm.distributed.ec_transfer.ec_connector.mooncake.control import ( + ControlClient, + EventInbox, + ShardTopology, + make_cancel_request, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeConnectorMetadata, + ECMooncakeLoadSpec, + ECMooncakePushSpec, + ECMooncakeWorkerMetadata, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.state import ( + SchedulerTransferState, + SchedulerTransferTable, +) +from vllm.logger import init_logger +from vllm.v1.core.sched.output import SchedulerOutput +from vllm.v1.outputs import ECConnectorOutput + +if TYPE_CHECKING: + from vllm.config import ModelConfig, VllmConfig + +logger = init_logger(__name__) + +_LEASE_TTL_SECONDS = 300 +_DRAIN_MIN_INTERVAL = 0.005 +_MAX_PENDING_EVENTS = 4096 +_MAX_TERMINAL_TRANSFER_RECORDS = 1 << 16 +_CANCEL_ATTEMPTS = 2 + + +class ECMooncakeScheduler: + """Coordinate Mooncake transfers from the vLLM Scheduler process. + + Attributes: + _is_producer: Whether this Scheduler prepares outbound pushes. + _is_consumer: Whether this Scheduler waits for inbound pushes. + _reservation_zmq_addr: Base Consumer control-plane address. + _consumer_pool_capacity: Capacity mirrored for resident-state limits. + _push_wait_timeout: Maximum wait for a Consumer readiness event. + _consumer_metrics_log_interval: Interval for Scheduler metrics logs. + _encoder_cache_hidden_dim: Hidden width used to derive push shapes. + _model_config: Model metadata used for dtypes and multimodal fields. + _control_client: Client for reservation status and cancellation. + _topology: Discovery cache for Consumer TP shards. + _event_inbox: Non-blocking source of Consumer readiness events. + _control_executor: Executor for cancellation requests. + _metadata_fields_cache: Placeholder metadata fields by modality. + _consumer_metrics_started_at: Start time of the current metric window. + _consumer_scheduler_metrics: Counters for scheduling decisions. + _drain_pending: Whether the next scheduling pass should drain events. + _drained_at: Time of the most recent readiness-event drain. + _pending_cancels: Asynchronous cancellations by transfer ID. + _transfers: Scheduler-owned transfer lifecycle table. + _scheduler_pending_work: Whether Workers reported unfinished work. + _pushes_to_prepare: Push specs awaiting metadata emission. + _prepared_push_transfer_ids: Transfer IDs already prepared once. + _event_shard_count: Number of Consumer event channels in the topology. + _event_ready_shards: Ready shard IDs accumulated per transfer. + """ + + @classmethod + def from_vllm_config(cls, vllm_config: VllmConfig) -> ECMooncakeScheduler: + ensure_mooncake_available() + config = MooncakeECConfig.from_vllm_config( + vllm_config, ECConnectorRole.SCHEDULER + ) + + control_client = ControlClient(config.control_timeout_ms) + topology = ShardTopology(control_client) + event_inbox = EventInbox(control_client, topology) + control_executor = ThreadPoolExecutor( + max_workers=config.control_workers, + thread_name_prefix="ec-mooncake-control", + ) + encoder_cache_hidden_dim = ( + _get_encoder_cache_hidden_dim(vllm_config) if config.is_producer else None + ) + return cls( + config, + encoder_cache_hidden_dim, + model_config=vllm_config.model_config, + control_client=control_client, + topology=topology, + event_inbox=event_inbox, + control_executor=control_executor, + ) + + def __init__( + self, + config: MooncakeECConfig, + encoder_cache_hidden_dim: int | None, + model_config: ModelConfig, + control_client: ControlClient, + topology: ShardTopology, + event_inbox: EventInbox, + control_executor: ThreadPoolExecutor, + ) -> None: + self._is_producer = config.is_producer + self._is_consumer = config.is_consumer + self._reservation_zmq_addr = config.reservation_addr + self._consumer_pool_capacity = config.consumer_pool_size + self._push_wait_timeout = config.push_wait_timeout_s + self._consumer_metrics_log_interval = config.consumer_metrics_log_interval + self._encoder_cache_hidden_dim = encoder_cache_hidden_dim + self._model_config = model_config + self._control_client = control_client + self._topology = topology + self._event_inbox = event_inbox + self._control_executor = control_executor + + self._metadata_fields_cache: dict[str, set[str]] = {} + self._consumer_metrics_started_at = time.monotonic() + self._consumer_scheduler_metrics: Counter[str] = Counter() + self._drain_pending = True + self._drained_at = 0.0 + self._pending_cancels: dict[str, Future[Any]] = {} + self._transfers = SchedulerTransferTable( + self._consumer_pool_capacity, _LEASE_TTL_SECONDS + ) + self._scheduler_pending_work = False + self._pushes_to_prepare: dict[str, ECMooncakePushSpec] = {} + self._prepared_push_transfer_ids: set[str] = set() + self._event_shard_count = 1 + self._event_ready_shards: OrderedDict[str, set[int]] = OrderedDict() + + def _cancel_remote( + self, consumer_zmq: str, transfer_id: str, reservation_id: str + ) -> bool: + pending = None + for _ in range(_CANCEL_ATTEMPTS): + pending = self._topology.discover(consumer_zmq) + if pending is not None: + break + if pending is None: + raise RuntimeError( + f"Could not discover every EC consumer shard at {consumer_zmq}" + ) + + cancelled = False + error: BaseException | None = None + for _ in range(_CANCEL_ATTEMPTS): + failed = [] + for addr in pending: + try: + result = self._control_client.request( + addr, + make_cancel_request(transfer_id, reservation_id), + ) + except Exception as exc: + if error is None: + error = exc + failed.append(addr) + continue + cancelled |= isinstance(result, dict) and bool(result.get("cancelled")) + if not failed: + return cancelled + pending = failed + if error is not None: + raise error + return cancelled + + def _note_awaiting_push( + self, + mm_hash: str, + transfer_id: str, + request_id: str, + ) -> None: + now = time.monotonic() + record = self._transfers.wait_for_event( + transfer_id, + request_id, + mm_hash, + now + self._push_wait_timeout, + ) + self._consumer_scheduler_metrics["missing_event"] += 1 + if record.state is not SchedulerTransferState.WAITING_EVENT: + return + assert record.deadline is not None + if now < record.deadline: + return + elapsed = now - record.deadline + self._push_wait_timeout + self._transfers.mark_unavailable( + transfer_id, "push readiness event timed out", now + ) + self._consumer_scheduler_metrics["given_up"] += 1 + self._consumer_scheduler_metrics["stalled"] += 1 + reservation: Any = "unknown" + if self._reservation_zmq_addr is not None: + try: + reservation = self._control_client.request( + self._reservation_zmq_addr, + {"op": "status", "transfer_id": transfer_id}, + ) + except Exception as e: # noqa: BLE001 - diagnostic only + reservation = f"status failed: {e}" + logger.warning( + "EC Mooncake waited %.1fs for a push of mm_hash=%s " + "(transfer_id=%s) that never arrived; worker reservation=%s; " + "requests needing it fail with a retryable error.", + elapsed, + mm_hash, + transfer_id, + reservation, + ) + + def take_unavailable_requests(self) -> set[str]: + return self._transfers.take_unavailable_requests() + + def _maybe_log_consumer_scheduler_metrics(self) -> None: + now = time.monotonic() + if ( + self._consumer_metrics_log_interval <= 0 + or now - self._consumer_metrics_started_at + < self._consumer_metrics_log_interval + ): + return + missing = self._transfers.count(SchedulerTransferState.WAITING_EVENT) + loading = self._transfers.count(SchedulerTransferState.LOADING) + pending = self._transfers.count(SchedulerTransferState.AVAILABLE) + logger.info( + "EC Mooncake consumer scheduler: decisions=%s, ready=%d, loading=%d, " + "resident=%d, pending_specs=%d, needs_load=%d, missing=%d", + dict(self._consumer_scheduler_metrics), + self._transfers.count(SchedulerTransferState.READY), + loading, + self._transfers.count(SchedulerTransferState.RESIDENT), + pending, + loading, + missing, + ) + self._consumer_scheduler_metrics.clear() + self._consumer_metrics_started_at = now + + def _poll_pending_cancels(self) -> None: + pending = {} + for transfer_id, future in self._pending_cancels.items(): + if not future.done(): + pending[transfer_id] = future + continue + try: + cancelled = future.result() + except Exception: + self._consumer_scheduler_metrics["cancellations_failed"] += 1 + logger.warning( + "EC Mooncake reservation cancellation failed", exc_info=True + ) + else: + key = "cancellations_completed" if cancelled else "cancellations_stale" + self._consumer_scheduler_metrics[key] += 1 + self._pending_cancels = pending + + def _note_shard_ready(self, data: dict[str, Any]) -> bool: + if self._event_shard_count <= 1: + return True + transfer_id = str(data["transfer_id"]) + record = self._transfers.get(transfer_id) + if ( + record is not None + and record.state is not SchedulerTransferState.WAITING_EVENT + ): + return False + shard = data.get("shard") + shards = self._event_ready_shards.setdefault(transfer_id, set()) + self._event_ready_shards.move_to_end(transfer_id) + shards.add(int(shard) if shard is not None else len(shards)) + if len(shards) < self._event_shard_count: + self._consumer_scheduler_metrics["events_awaiting_shards"] += 1 + while len(self._event_ready_shards) > _MAX_PENDING_EVENTS: + self._event_ready_shards.popitem(last=False) + self._consumer_scheduler_metrics["events_partial_dropped"] += 1 + return False + self._event_ready_shards.pop(transfer_id, None) + self._consumer_scheduler_metrics["events_all_shards_ready"] += 1 + return True + + def _forget_shard_readiness(self, transfer_id: str) -> None: + self._event_ready_shards.pop(transfer_id, None) + + def _store_pushed_spec(self, data: dict[str, Any]) -> bool: + transfer_id = str(data["transfer_id"]) + identifier = str(data["mm_hash"]) + reservation_id = str(data["reservation_id"]) + _, accepted = self._transfers.observe_ready( + ECMooncakeLoadSpec( + mm_hash=identifier, + num_token=0, + nbytes=int(data["nbytes"]), + shape=tuple(int(value) for value in data["shape"]), + dtype=str(data["dtype"]), + pushed=True, + transfer_id=transfer_id, + reservation_id=reservation_id, + ), + time.monotonic() + _LEASE_TTL_SECONDS, + ) + return accepted + + def _queue_cancel( + self, + transfer_id: str, + reservation_id: str = "", + mm_hash: str = "", + request_id: str = "", + ) -> None: + if not self._transfers.cancel( + transfer_id, + time.monotonic(), + mm_hash=mm_hash, + request_id=request_id, + ): + return + self._forget_shard_readiness(transfer_id) + if self._reservation_zmq_addr is None: + return + self._pending_cancels[transfer_id] = self._control_executor.submit( + self._cancel_remote, + self._reservation_zmq_addr, + transfer_id, + reservation_id, + ) + + def _expire_transfers(self) -> None: + now = time.monotonic() + expired, dropped = self._transfers.expire(now, _MAX_TERMINAL_TRANSFER_RECORDS) + self._consumer_scheduler_metrics["cancel_records_dropped"] += dropped + for record in expired: + self._consumer_scheduler_metrics["pending_specs_expired"] += 1 + self._queue_cancel(record.transfer_id) + + def _drain_push_notifications(self) -> None: + now = time.monotonic() + if not self._drain_pending and now - self._drained_at < _DRAIN_MIN_INTERVAL: + return + self._drain_pending = False + self._drained_at = now + self._poll_pending_cancels() + self._expire_transfers() + if self._reservation_zmq_addr is None: + return + events = self._event_inbox.drain(self._reservation_zmq_addr) + self._event_shard_count = self._event_inbox.shard_count + for data in events: + identifier = str(data["mm_hash"]) + self._consumer_scheduler_metrics["events_received"] += 1 + if data.get("ready"): + self._consumer_scheduler_metrics["events_ready"] += 1 + transfer_id = str(data["transfer_id"]) + record = self._transfers.get(transfer_id) + if record is not None and record.state in { + SchedulerTransferState.CANCELLED, + SchedulerTransferState.UNAVAILABLE, + SchedulerTransferState.EXPIRED, + SchedulerTransferState.FAILED, + }: + self._consumer_scheduler_metrics["events_cancelled"] += 1 + continue + if self._transfers.has_state( + identifier, (SchedulerTransferState.READY,) + ): + self._consumer_scheduler_metrics["events_redundant"] += 1 + if not self._note_shard_ready(data): + continue + if not self._store_pushed_spec(data): + self._consumer_scheduler_metrics["events_duplicate"] += 1 + else: + self._consumer_scheduler_metrics["events_not_ready"] += 1 + + def has_cache_item(self, identifier: str) -> bool: + if not self._is_consumer: + return False + self._drain_push_notifications() + self._maybe_log_consumer_scheduler_metrics() + if self._transfers.has_state(identifier, (SchedulerTransferState.READY,)): + self._consumer_scheduler_metrics["ready"] += 1 + return True + if self._transfers.has_state(identifier, (SchedulerTransferState.LOADING,)): + self._consumer_scheduler_metrics["loading"] += 1 + return False + if self._transfers.has_state(identifier, (SchedulerTransferState.RESIDENT,)): + self._consumer_scheduler_metrics["resident"] += 1 + return True + if self._transfers.has_state(identifier, (SchedulerTransferState.AVAILABLE,)): + self._consumer_scheduler_metrics["pending_spec"] += 1 + return True + self._consumer_scheduler_metrics["missing_event"] += 1 + return False + + @staticmethod + def _request_transfer_id(request: Any, index: int) -> str | None: + params = getattr(request, "ec_transfer_params", None) or {} + items = params.get("ec_items") or [] + mm_hash = request.mm_features[index].identifier + if index < len(items): + item = items[index] + if item.get("mm_hash") in (None, mm_hash) and item.get("transfer_id"): + return str(item["transfer_id"]) + for item in items: + if item.get("mm_hash") == mm_hash and item.get("transfer_id"): + return str(item["transfer_id"]) + return None + + def ensure_cache_available( + self, + request: Any, + num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, + ) -> bool: + if self._is_producer: + for index, feature in enumerate(request.mm_features): + if ( + feature.mm_position.offset + feature.mm_position.length + > num_computed_tokens + ): + self._prepare_push_spec(request, index) + if not self._is_consumer: + return True + + self._drain_push_notifications() + local_cache_hashes = local_cache_hashes or set() + all_ready = True + for index, feature in enumerate(request.mm_features): + if ( + feature.mm_position.offset + feature.mm_position.length + <= num_computed_tokens + ): + continue + mm_hash = feature.identifier + transfer_id = self._request_transfer_id(request, index) + if transfer_id is not None: + self._transfers.touch_available( + transfer_id, time.monotonic() + _LEASE_TTL_SECONDS + ) + if mm_hash in local_cache_hashes: + continue + if self._transfers.has_state(mm_hash, (SchedulerTransferState.READY,)): + self._consumer_scheduler_metrics["ready"] += 1 + continue + if self._transfers.has_state(mm_hash, (SchedulerTransferState.LOADING,)): + self._consumer_scheduler_metrics["loading"] += 1 + all_ready = False + continue + if self._transfers.has_state(mm_hash, (SchedulerTransferState.RESIDENT,)): + self._consumer_scheduler_metrics["resident_hit"] += 1 + record = self._transfers.begin_load( + mm_hash, + request.get_num_encoder_embeds(index), + transfer_id, + request.request_id, + ) + if record is not None: + self._scheduler_pending_work = True + all_ready = False + else: + waiting_id = transfer_id or f"{request.request_id}:{index}" + self._note_awaiting_push(mm_hash, waiting_id, request.request_id) + all_ready = False + return all_ready + + def _prepare_push_spec(self, request: Any, index: int) -> None: + params = getattr(request, "ec_transfer_params", None) or {} + consumer_zmq = params.get("consumer_zmq") + mm_hash = request.mm_features[index].identifier + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + transfer_id = f"{request.request_id}:{index}" + if not consumer_zmq or transfer_id in self._prepared_push_transfer_ids: + return + num_tokens = request.get_num_encoder_embeds(index) + dtype = self._model_config.dtype + assert isinstance(dtype, torch.dtype) + assert self._encoder_cache_hidden_dim is not None + dtype_name = str(dtype).split(".")[-1] + shape = (num_tokens, self._encoder_cache_hidden_dim) + nbytes = math.prod(shape) * dtype.itemsize + self._pushes_to_prepare[transfer_id] = ECMooncakePushSpec( + mm_hash=mm_hash, + nbytes=nbytes, + shape=shape, + dtype=dtype_name, + consumer_zmq=str(consumer_zmq), + transfer_id=transfer_id, + request_id=request.request_id, + ) + self._prepared_push_transfer_ids.add(transfer_id) + + def update_state_after_alloc(self, request: Any, index: int) -> None: + if self._is_producer: + self._prepare_push_spec(request, index) + + def update_state_after_free(self, request: Any, index: int) -> None: + if not self._is_consumer: + return + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + return + self._queue_cancel( + transfer_id, + mm_hash=request.mm_features[index].identifier, + request_id=request.request_id, + ) + + def build_connector_meta( + self, scheduler_output: SchedulerOutput + ) -> ECConnectorMetadata: + for mm_hash in scheduler_output.free_encoder_mm_hashes: + self._transfers.release_ready(mm_hash, time.monotonic()) + meta = ECMooncakeConnectorMetadata() + for push_spec in self._pushes_to_prepare.values(): + meta.add_push(push_spec) + self._pushes_to_prepare.clear() + for record in self._transfers.take_loads_to_dispatch(): + assert record.spec is not None + meta.add_load(record.spec) + self._poll_pending_cancels() + self._maybe_log_consumer_scheduler_metrics() + self._drain_pending = True + return meta + + def update_connector_output(self, connector_output: ECConnectorOutput) -> None: + meta = connector_output.ec_connector_worker_meta + if not isinstance(meta, ECMooncakeWorkerMetadata): + return + for mm_hash in meta.loaded: + self._transfers.complete_load(mm_hash) + self._consumer_scheduler_metrics["loads_completed"] += 1 + for mm_hash in meta.failed_loads: + self._transfers.fail_load( + mm_hash, "worker failed to load", time.monotonic() + ) + self._consumer_scheduler_metrics["loads_failed"] += 1 + for mm_hash in meta.reclaimed: + self._transfers.reclaim(mm_hash, time.monotonic()) + self._consumer_scheduler_metrics["resident_reclaimed"] += 1 + self._scheduler_pending_work = meta.pending_loads or meta.pending_saves + + def has_pending_push_work(self) -> bool: + return self._scheduler_pending_work + + def _placeholder_metadata_fields(self, modality: str) -> set[str]: + if modality in self._metadata_fields_cache: + return self._metadata_fields_cache[modality] + + fields: set[str] = set() + try: + from vllm.multimodal import MULTIMODAL_REGISTRY + + info = MULTIMODAL_REGISTRY.create_processor(self._model_config).info + fields = info.data_parser.placeholder_metadata_fields(modality) + except Exception: + logger.warning( + "Could not determine the placeholder metadata fields for " + "modality %s; the consumer will preprocess the media itself.", + modality, + exc_info=True, + ) + + self._metadata_fields_cache[modality] = fields + return fields + + def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: + if self._is_consumer: + for index in range(len(request.mm_features)): + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + continue + self._queue_cancel( + transfer_id, + mm_hash=request.mm_features[index].identifier, + request_id=request.request_id, + ) + if self._is_producer and self._prepared_push_transfer_ids: + for index in range(len(request.mm_features)): + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + transfer_id = f"{request.request_id}:{index}" + self._prepared_push_transfer_ids.discard(transfer_id) + if not self._is_producer: + return False, None + + items = [] + for index, feature in enumerate(request.mm_features): + metadata = {} + if feature.data is not None: + wanted = self._placeholder_metadata_fields(feature.modality) + metadata = { + key: value.tolist() + for key, value in feature.data.get_data().items() + if key in wanted and isinstance(value, torch.Tensor) + } + transfer_id = self._request_transfer_id(request, index) + item = {"mm_hash": feature.identifier, **metadata} + if transfer_id is not None: + item["transfer_id"] = transfer_id + items.append(item) + + if not items: + return False, None + return False, {"ec_items": items} + + def close(self) -> None: + self._control_executor.shutdown(wait=True, cancel_futures=True) + self._control_client.close() + self._event_inbox.close() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py new file mode 100644 index 000000000000..3904bbb6b614 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py @@ -0,0 +1,432 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Scheduler-owned lifecycle and indexes for Consumer-bound transfers.""" + +from __future__ import annotations + +from collections import OrderedDict, deque +from collections.abc import Iterable +from dataclasses import dataclass, field, replace +from enum import Enum, auto + +from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeLoadSpec, +) + + +class SchedulerTransferState(Enum): + """Lifecycle states visible to the Scheduler. + + The states cover waiting for a Consumer event, dispatching a Worker load, + retaining a locally reusable tensor, and terminal failure conditions. + """ + + WAITING_EVENT = auto() + AVAILABLE = auto() + LOADING = auto() + READY = auto() + RESIDENT = auto() + UNAVAILABLE = auto() + EXPIRED = auto() + FAILED = auto() + CANCELLED = auto() + + +@dataclass +class SchedulerTransfer: + """Track one transfer as observed by the Scheduler. + + Attributes: + transfer_id: Cross-process identity of the transfer. + request_id: Request currently waiting for the transfer. + mm_hash: Stable identifier of the encoder-cache item. + state: Current Scheduler lifecycle state. + spec: Load metadata once the Consumer reports the tensor ready. + deadline: Expiry time for waiting, available, or terminal records. + last_error: Last terminal error associated with the transfer. + notified_requests: Requests already told that the item is unavailable. + """ + + transfer_id: str + request_id: str + mm_hash: str + state: SchedulerTransferState + spec: ECMooncakeLoadSpec | None + deadline: float | None + last_error: str | None = None + notified_requests: set[str] = field(default_factory=set, repr=False) + + +class InvalidSchedulerTransferTransition(RuntimeError): + """Raised when code attempts an unsupported Scheduler state transition.""" + + pass + + +_ALLOWED_TRANSITIONS = { + SchedulerTransferState.WAITING_EVENT: { + SchedulerTransferState.AVAILABLE, + SchedulerTransferState.UNAVAILABLE, + SchedulerTransferState.CANCELLED, + }, + SchedulerTransferState.AVAILABLE: { + SchedulerTransferState.LOADING, + SchedulerTransferState.EXPIRED, + SchedulerTransferState.CANCELLED, + }, + SchedulerTransferState.LOADING: { + SchedulerTransferState.READY, + SchedulerTransferState.FAILED, + SchedulerTransferState.CANCELLED, + }, + SchedulerTransferState.READY: { + SchedulerTransferState.RESIDENT, + SchedulerTransferState.EXPIRED, + }, + SchedulerTransferState.RESIDENT: { + SchedulerTransferState.LOADING, + SchedulerTransferState.EXPIRED, + }, + SchedulerTransferState.UNAVAILABLE: {SchedulerTransferState.CANCELLED}, + SchedulerTransferState.EXPIRED: {SchedulerTransferState.CANCELLED}, + SchedulerTransferState.FAILED: set(), + SchedulerTransferState.CANCELLED: set(), +} + +_TERMINAL_STATES = { + SchedulerTransferState.UNAVAILABLE, + SchedulerTransferState.EXPIRED, + SchedulerTransferState.FAILED, + SchedulerTransferState.CANCELLED, +} + + +class SchedulerTransferTable: + """Own Scheduler transfer state, lookup indexes, and dispatch queues. + + Attributes: + _resident_capacity: Maximum bytes represented by resident records. + _tombstone_ttl: Retention time for terminal records. + _records: Ordered transfer records keyed by transfer ID. + _hash_index: Transfer IDs grouped by encoder-cache hash. + _loads_to_dispatch: Ordered IDs awaiting Worker metadata emission. + _unavailable_requests: Requests awaiting retryable failure reporting. + """ + + def __init__(self, resident_capacity: int, tombstone_ttl: float) -> None: + self._resident_capacity = resident_capacity + self._tombstone_ttl = tombstone_ttl + self._records: OrderedDict[str, SchedulerTransfer] = OrderedDict() + self._hash_index: dict[str, deque[str]] = {} + self._loads_to_dispatch: OrderedDict[str, None] = OrderedDict() + self._unavailable_requests: set[str] = set() + + def get(self, transfer_id: str) -> SchedulerTransfer | None: + return self._records.get(transfer_id) + + def records_for_hash( + self, + mm_hash: str, + states: Iterable[SchedulerTransferState], + ) -> list[SchedulerTransfer]: + wanted = set(states) + return [ + record + for transfer_id in self._hash_index.get(mm_hash, ()) + if (record := self._records.get(transfer_id)) is not None + and record.state in wanted + ] + + def first_for_hash( + self, + mm_hash: str, + states: Iterable[SchedulerTransferState], + ) -> SchedulerTransfer | None: + return next(iter(self.records_for_hash(mm_hash, states)), None) + + def has_state(self, mm_hash: str, states: Iterable[SchedulerTransferState]) -> bool: + return self.first_for_hash(mm_hash, states) is not None + + def count(self, state: SchedulerTransferState) -> int: + return sum(record.state is state for record in self._records.values()) + + @property + def resident_bytes(self) -> int: + return sum( + record.spec.nbytes + for record in self._records.values() + if record.state is SchedulerTransferState.RESIDENT and record.spec + ) + + def wait_for_event( + self, + transfer_id: str, + request_id: str, + mm_hash: str, + deadline: float, + ) -> SchedulerTransfer: + record = self._records.get(transfer_id) + if record is None: + record = SchedulerTransfer( + transfer_id=transfer_id, + request_id=request_id, + mm_hash=mm_hash, + state=SchedulerTransferState.WAITING_EVENT, + spec=None, + deadline=deadline, + ) + self._insert(record) + else: + self._check_identity(record, mm_hash) + if not record.request_id: + record.request_id = request_id + if record.state in _TERMINAL_STATES and request_id: + self._notify_unavailable(record, request_id) + return record + + def observe_ready( + self, spec: ECMooncakeLoadSpec, deadline: float + ) -> tuple[SchedulerTransfer, bool]: + transfer_id = spec.transfer_id or spec.mm_hash + record = self._records.get(transfer_id) + if record is None: + record = SchedulerTransfer( + transfer_id=transfer_id, + request_id="", + mm_hash=spec.mm_hash, + state=SchedulerTransferState.WAITING_EVENT, + spec=None, + deadline=None, + ) + self._insert(record) + else: + self._check_identity(record, spec.mm_hash) + if record.state is not SchedulerTransferState.WAITING_EVENT: + return record, False + record.spec = spec + record.deadline = deadline + self._transition(record, SchedulerTransferState.AVAILABLE) + return record, True + + def touch_available(self, transfer_id: str, deadline: float) -> None: + record = self._records.get(transfer_id) + if record is not None and record.state is SchedulerTransferState.AVAILABLE: + record.deadline = deadline + + def begin_load( + self, + mm_hash: str, + num_token: int, + transfer_id: str | None = None, + request_id: str = "", + ) -> SchedulerTransfer | None: + record = self._records.get(transfer_id) if transfer_id else None + if record is not None and ( + record.mm_hash != mm_hash + or record.state is not SchedulerTransferState.AVAILABLE + ): + record = None + if record is None: + record = self.first_for_hash( + mm_hash, + ( + SchedulerTransferState.AVAILABLE, + SchedulerTransferState.RESIDENT, + ), + ) + if record is None or record.spec is None: + return None + if not record.request_id: + record.request_id = request_id + record.spec = replace(record.spec, num_token=num_token) + record.deadline = None + self._transition(record, SchedulerTransferState.LOADING) + self._loads_to_dispatch[record.transfer_id] = None + return record + + def take_loads_to_dispatch(self) -> list[SchedulerTransfer]: + records = [ + record + for transfer_id in self._loads_to_dispatch + if (record := self._records.get(transfer_id)) is not None + and record.state is SchedulerTransferState.LOADING + ] + self._loads_to_dispatch.clear() + return records + + def complete_load(self, mm_hash: str) -> bool: + record = self.first_for_hash(mm_hash, (SchedulerTransferState.LOADING,)) + if record is None: + return self.has_state(mm_hash, (SchedulerTransferState.READY,)) + self._transition(record, SchedulerTransferState.READY) + return True + + def fail_load(self, mm_hash: str, error: str, now: float) -> bool: + record = self.first_for_hash(mm_hash, (SchedulerTransferState.LOADING,)) + if record is None: + return self.has_state(mm_hash, (SchedulerTransferState.FAILED,)) + self._transition(record, SchedulerTransferState.FAILED, error, now=now) + return True + + def release_ready(self, mm_hash: str, now: float) -> None: + ready = [ + record + for record in self._records.values() + if record.mm_hash == mm_hash + and record.state is SchedulerTransferState.READY + ] + if ready: + canonical = ready[-1] + for record in self.records_for_hash( + mm_hash, + (SchedulerTransferState.READY, SchedulerTransferState.RESIDENT), + ): + if record is not canonical: + self._transition(record, SchedulerTransferState.EXPIRED, now=now) + if canonical.spec is None: + self._transition(canonical, SchedulerTransferState.EXPIRED, now=now) + else: + canonical.spec = replace(canonical.spec, num_token=0, local=True) + self._transition(canonical, SchedulerTransferState.RESIDENT) + self._evict_residents(now) + + def reclaim(self, mm_hash: str, now: float) -> None: + for record in self.records_for_hash( + mm_hash, + (SchedulerTransferState.READY, SchedulerTransferState.RESIDENT), + ): + if record.state is SchedulerTransferState.READY: + record.spec = None + else: + self._transition(record, SchedulerTransferState.EXPIRED, now=now) + + def mark_unavailable(self, transfer_id: str, error: str, now: float) -> None: + record = self._records[transfer_id] + self._transition(record, SchedulerTransferState.UNAVAILABLE, error, now=now) + if record.request_id: + self._notify_unavailable(record, record.request_id) + + def cancel( + self, + transfer_id: str, + now: float, + mm_hash: str = "", + request_id: str = "", + ) -> bool: + record = self._records.get(transfer_id) + if record is None: + record = SchedulerTransfer( + transfer_id=transfer_id, + request_id=request_id, + mm_hash=mm_hash, + state=SchedulerTransferState.WAITING_EVENT, + spec=None, + deadline=None, + ) + self._insert(record) + if record.state in { + SchedulerTransferState.CANCELLED, + SchedulerTransferState.READY, + SchedulerTransferState.RESIDENT, + SchedulerTransferState.FAILED, + }: + return False + self._transition(record, SchedulerTransferState.CANCELLED, now=now) + self._loads_to_dispatch.pop(transfer_id, None) + return True + + def expire( + self, now: float, terminal_limit: int + ) -> tuple[list[SchedulerTransfer], int]: + if terminal_limit < 0: + raise ValueError("terminal_limit must be non-negative") + expired = [] + dropped = 0 + for record in list(self._records.values()): + if record.deadline is None or record.deadline > now: + continue + if record.state is SchedulerTransferState.AVAILABLE: + self._transition( + record, + SchedulerTransferState.EXPIRED, + "lease expired", + now=now, + ) + expired.append(record) + elif record.state in _TERMINAL_STATES: + self._remove(record.transfer_id) + dropped += 1 + terminal_ids = [ + record.transfer_id + for record in self._records.values() + if record.state in _TERMINAL_STATES + ] + excess = max(0, len(terminal_ids) - terminal_limit) + for transfer_id in terminal_ids[:excess]: + self._remove(transfer_id) + dropped += 1 + return expired, dropped + + def take_unavailable_requests(self) -> set[str]: + unavailable = self._unavailable_requests + self._unavailable_requests = set() + return unavailable + + def _notify_unavailable(self, record: SchedulerTransfer, request_id: str) -> None: + if request_id not in record.notified_requests: + record.notified_requests.add(request_id) + self._unavailable_requests.add(request_id) + + def _insert(self, record: SchedulerTransfer) -> None: + self._records[record.transfer_id] = record + if record.mm_hash: + self._hash_index.setdefault(record.mm_hash, deque()).append( + record.transfer_id + ) + + @staticmethod + def _check_identity(record: SchedulerTransfer, mm_hash: str) -> None: + if record.mm_hash and record.mm_hash != mm_hash: + raise ValueError( + f"Transfer {record.transfer_id!r} changed mm_hash from " + f"{record.mm_hash!r} to {mm_hash!r}" + ) + + def _transition( + self, + record: SchedulerTransfer, + state: SchedulerTransferState, + error: str | None = None, + now: float | None = None, + ) -> None: + if state not in _ALLOWED_TRANSITIONS[record.state]: + raise InvalidSchedulerTransferTransition( + f"Cannot transition {record.transfer_id!r} from " + f"{record.state.name} to {state.name}" + ) + if state in _TERMINAL_STATES: + if now is None: + raise ValueError("Terminal transition requires a timestamp") + record.deadline = now + self._tombstone_ttl + record.state = state + record.last_error = error + self._records.move_to_end(record.transfer_id) + + def _evict_residents(self, now: float) -> None: + while self.resident_bytes > self._resident_capacity: + record = next( + record + for record in self._records.values() + if record.state is SchedulerTransferState.RESIDENT + ) + self._transition(record, SchedulerTransferState.EXPIRED, now=now) + + def _remove(self, transfer_id: str) -> None: + record = self._records.pop(transfer_id, None) + self._loads_to_dispatch.pop(transfer_id, None) + if record is None or not record.mm_hash: + return + transfer_ids = self._hash_index[record.mm_hash] + transfer_ids.remove(transfer_id) + if not transfer_ids: + self._hash_index.pop(record.mm_hash) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py new file mode 100644 index 000000000000..469b2a280bd5 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py @@ -0,0 +1,216 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Mooncake data-plane engine and memory-registration ownership. + +The wrapper isolates optional-engine initialization, session addressing, +synchronous batched writes, and the lifetime of transient source registrations. +""" + +from __future__ import annotations + +import threading +from dataclasses import dataclass + +import torch + +from vllm.logger import init_logger + +logger = init_logger(__name__) + +try: + from mooncake.engine import TransferEngine +except ImportError: + TransferEngine = None # type: ignore[misc, assignment] + + +@dataclass +class _SourceRegistration: + """Retain one transient source range while batches reference it. + + Attributes: + tensor: Tensor keeping the registered storage alive. + nbytes: Exact registered byte length. + users: Number of active acquisitions of this address. + """ + + tensor: torch.Tensor + nbytes: int + users: int = 1 + + +class MooncakeTransfer: + """Own a lazy Mooncake engine and transient memory registrations. + + Attributes: + _hostname: Address advertised in the Mooncake session identifier. + _protocol: Transport protocol used to initialize ``TransferEngine``. + _engine: Lazily initialized Mooncake engine. + _engine_lock: Lock serializing first engine initialization. + _source_registrations: Reference-counted transient source ranges. + _pending_unregister: Tensors retained after an unregister failure. + _registration_lock: Lock protecting registration ownership. + _closed: Whether final data-plane cleanup has begun. + """ + + def __init__(self, hostname: str, protocol: str) -> None: + self._hostname = hostname + self._protocol = protocol + self._engine: TransferEngine | None = None + self._engine_lock = threading.Lock() + self._source_registrations: dict[int, _SourceRegistration] = {} + self._pending_unregister: dict[int, torch.Tensor] = {} + self._registration_lock = threading.Lock() + self._closed = False + + def _ensure_engine(self) -> TransferEngine: + if self._engine is not None: + return self._engine + with self._engine_lock: + if self._engine is not None: + return self._engine + engine = TransferEngine() + ret = engine.initialize(self._hostname, "P2PHANDSHAKE", self._protocol, "") + if ret != 0: + raise RuntimeError("Mooncake TransferEngine initialization failed.") + self._engine = engine + logger.info( + "ECMooncakeConnector TransferEngine ready at %s:%d", + self._hostname, + engine.get_rpc_port(), + ) + return self._engine + + def ensure_ready(self) -> None: + self._ensure_engine() + + def local_session(self) -> str: + engine = self._ensure_engine() + return f"{self._hostname}:{engine.get_rpc_port()}" + + def register_memory(self, tensor: torch.Tensor) -> int: + return self._ensure_engine().batch_register_memory( + [tensor.data_ptr()], [tensor.nbytes] + ) + + def unregister_memory(self, tensor: torch.Tensor) -> bool: + engine = self._ensure_engine() + address = tensor.data_ptr() + ret = engine.unregister_memory(address) + if ret != 0: + logger.error( + "Mooncake EC memory unregistration failed for address %d: %d", + address, + ret, + ) + self._pending_unregister[address] = tensor + return False + self._pending_unregister.pop(address, None) + return True + + @staticmethod + def _source_range(tensor: torch.Tensor) -> tuple[int, int]: + # Encoder batches commonly split one storage into sibling tensor views. + # Registering each view's exact bytes avoids overlapping memory regions. + return tensor.data_ptr(), tensor.nbytes + + def acquire_sources(self, tensors: list[torch.Tensor]) -> list[int]: + ranges: dict[int, tuple[int, torch.Tensor]] = {} + for tensor in tensors: + address, nbytes = self._source_range(tensor) + ranges.setdefault(address, (nbytes, tensor)) + + engine = self._ensure_engine() + acquired: list[int] = [] + new_addresses: list[int] = [] + new_lengths: list[int] = [] + with self._registration_lock: + for address, (nbytes, tensor) in ranges.items(): + entry = self._source_registrations.get(address) + if entry is not None: + if entry.nbytes != nbytes: + raise RuntimeError( + "Mooncake EC source storage changed size while registered" + ) + entry.users += 1 + acquired.append(address) + continue + new_addresses.append(address) + new_lengths.append(nbytes) + self._source_registrations[address] = _SourceRegistration( + tensor=tensor, + nbytes=nbytes, + ) + acquired.append(address) + + if new_addresses: + ret = engine.batch_register_memory(new_addresses, new_lengths) + if ret != 0: + for address in acquired: + entry = self._source_registrations[address] + entry.users -= 1 + if entry.users == 0: + del self._source_registrations[address] + raise RuntimeError("Mooncake EC source registration failed") + return acquired + + def release_sources(self, addresses: list[int]) -> bool: + if not addresses: + return True + with self._registration_lock: + unused = [] + for address in addresses: + entry = self._source_registrations.get(address) + if entry is None: + continue + entry.users -= 1 + if entry.users == 0: + unused.append(address) + if not unused: + return True + ret = self._ensure_engine().batch_unregister_memory(unused) + if ret != 0: + logger.warning( + "Keeping %d EC source tensors registered after Mooncake " + "unregistration failure", + len(unused), + ) + return False + for address in unused: + del self._source_registrations[address] + self._pending_unregister.pop(address, None) + return True + + def write( + self, + session: str, + sources: list[int], + destinations: list[int], + lengths: list[int], + ) -> None: + """Write one synchronous batch, returning only at terminal status.""" + ret = self._ensure_engine().batch_transfer_sync_write( + session, sources, destinations, lengths + ) + if ret != 0: + raise RuntimeError( + f"Mooncake EC push to {session} failed with status {ret}" + ) + + def close(self) -> None: + if self._closed: + return + self._closed = True + engine = self._engine + if engine is None: + return + with self._registration_lock: + addresses = list(self._source_registrations) + addresses.extend(self._pending_unregister) + if not addresses: + return + ret = engine.batch_unregister_memory(list(dict.fromkeys(addresses))) + if ret != 0: + logger.error("Mooncake EC batch memory unregistration failed: %d", ret) + return + self._source_registrations.clear() + self._pending_unregister.clear() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py new file mode 100644 index 000000000000..bd85eddae1b6 --- /dev/null +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py @@ -0,0 +1,1154 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Worker-side orchestration of Mooncake control, memory, and data planes. + +Consumer Workers expose rank-local reservations and publish received tensors. +Producer Workers reserve every destination shard, bind computed sources, run +batched Mooncake writes, and report asynchronous completion to the Scheduler. +""" + +from __future__ import annotations + +import math +import threading +import time +from collections import Counter +from collections.abc import Callable +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import suppress +from dataclasses import dataclass, field +from functools import partial +from typing import TYPE_CHECKING, Any, TypeVar, cast + +import torch + +from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole +from vllm.distributed.ec_transfer.ec_connector.mooncake._availability import ( + ensure_mooncake_available, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.config import MooncakeECConfig +from vllm.distributed.ec_transfer.ec_connector.mooncake.control import ( + ConsumerControlServer, + ControlClient, + ControlCompletion, + ShardTopology, + make_cancel_request, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.memory import ( + ConsumerMemoryPool, + MemoryAllocation, + ProducerMemoryPool, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeConnectorMetadata, + ECMooncakeLoadSpec, + ECMooncakePushSpec, + ECMooncakeWorkerMetadata, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.producer import ( + ProducerPushManager, + ProducerPushRecord, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.reservation import ( + CancellationOutcome, + ConsumerReservationManager, + ConsumerReservationState, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( + MooncakeTransfer, +) +from vllm.logger import init_logger +from vllm.utils.network_utils import get_ip + +logger = init_logger(__name__) + +_T = TypeVar("_T") + +if TYPE_CHECKING: + from vllm.config import VllmConfig + +_LEASE_TTL_SECONDS = 300 +_RESERVATION_REFRESH_SECONDS = _LEASE_TTL_SECONDS / 2 +_MAX_CANCELLED_TRANSFER_IDS = 1 << 16 +_CANCEL_ATTEMPTS = 2 +_PUSH_STAGES = ( + "reserve", + "cuda", + "register", + "rdma", + "unregister", + "complete", +) + + +@dataclass +class _PushPerfWindow: + """Accumulate Producer batch metrics between periodic log messages. + + Attributes: + started_at: Monotonic start time of the aggregation window. + batches: Number of completed push batches. + items: Number of push records included in those batches. + bytes: Number of tensor bytes written over the data plane. + skipped_items: Items satisfied by cache or cancellation without a write. + failures: Number of batches that ended in failure. + stage_totals_ms: Accumulated time for every push stage. + stage_max_ms: Maximum observed time for every push stage. + """ + + started_at: float = field(default_factory=time.monotonic) + batches: int = 0 + items: int = 0 + bytes: int = 0 + skipped_items: int = 0 + failures: int = 0 + stage_totals_ms: dict[str, float] = field(default_factory=dict) + stage_max_ms: dict[str, float] = field(default_factory=dict) + + +class _FanoutError(RuntimeError): + """Retain shard outcomes after every started task settles.""" + + def __init__(self, error: BaseException, results: list[Any | None]) -> None: + self.results = results + super().__init__(str(error)) + + +class _ReservationFanoutError(RuntimeError): + """Expose partial reservations for precise idempotent cleanup retries.""" + + def __init__( + self, + error: BaseException, + partial_reservations: list[dict[str, Any]], + ) -> None: + super().__init__(str(error)) + self.partial_reservations = partial_reservations + + +class ECMooncakeWorker: + """Orchestrate consumer reservations and producer push batches. + + ``mooncake_protocol`` selects the transfer protocol. Consumer workers use + ``consumer_buffer_pool_size`` and ``reservation_zmq_port`` for their + registered receive arena and rank-local control endpoint. Producers use + ``producer_buffer_pool_size`` for staging. ``transfer_max_workers`` and + ``control_max_workers`` bound the two executor pools; the transfer and + consumer metrics intervals control aggregate logging. + + Consumers may use TP, PP, and DP. Only the first PP stage receives encoder + outputs, and each TP rank exposes a consecutive control port and receives + the same source concurrently. Producers remain unsharded and unreplicated. + With DP, the caller must route both halves of a request to the same replica + and pass that replica's control address to the producer. + + Attributes: + is_producer: Whether this Worker originates encoder-cache pushes. + is_consumer: Whether this Worker accepts encoder-cache pushes. + _buffer_device: Device requested for registered memory pools. + _reservation_zmq_port: Base Consumer control port for this DP replica. + _transfer: Owner of the Mooncake engine and memory registrations. + _consumer_worker_metrics: Consumer lifecycle metric counters. + _consumer_memory: Registered receive slab and resident cache. + _reservations: Consumer destination reservation state manager. + _consumer_rank_resolved: Whether TP/PP placement has been discovered. + _is_receiving_rank: Whether this PP stage owns encoder outputs. + _tp_rank: Tensor-parallel rank used to derive the local control port. + _tp_size: Number of Consumer tensor-parallel destination shards. + _control_server: Rank-local Consumer reservation server. + _consumer_metrics_log_interval: Consumer metrics log interval. + _consumer_metrics_started_at: Start time of the Consumer metric window. + _producer_memory: Registered Producer source staging slab. + _transfer_metrics_log_interval: Producer performance log interval. + _control_client: Client for remote Consumer control operations. + _topology: Discovery cache for remote Consumer TP shards. + _producer_metrics: Producer lifecycle metric counters. + _io_executor: Executor that owns transfer and cancellation batches. + _control_executor: Executor that creates remote reservations. + _shard_pool: Lazily created executor for concurrent TP-shard work. + _shard_pool_lock: Lock protecting shard-pool initialization. + _producer_pushes: Producer lifecycle and source-ownership manager. + _push_perf_lock: Lock protecting Producer performance counters. + _push_perf: Current Producer performance aggregation window. + _active_transfer_batches: Batches currently executing data-plane work. + _queued_transfer_batches: Batches submitted but not yet executing. + _completed_loads: Successful Consumer loads awaiting reporting. + _failed_loads: Failed Consumer loads awaiting reporting. + _shutdown: Whether Worker resource shutdown has started. + """ + + @classmethod + def from_vllm_config(cls, vllm_config: VllmConfig) -> ECMooncakeWorker: + ensure_mooncake_available() + config = MooncakeECConfig.from_vllm_config(vllm_config, ECConnectorRole.WORKER) + hostname = get_ip() + control_client = ControlClient(config.control_timeout_ms) + try: + return cls( + config, + hostname, + control_client, + ShardTopology(control_client), + ) + except Exception: + control_client.close() + raise + + def __init__( + self, + config: MooncakeECConfig, + hostname: str, + control_client: ControlClient, + topology: ShardTopology, + ) -> None: + self.is_producer = config.is_producer + self.is_consumer = config.is_consumer + self._buffer_device = config.buffer_device + self._reservation_zmq_port = config.reservation_port + self._transfer = MooncakeTransfer(hostname, config.protocol) + self._consumer_worker_metrics: Counter[str] = Counter() + self._consumer_memory = ConsumerMemoryPool( + config.consumer_pool_size, + self._transfer, + ) + self._reservations = ConsumerReservationManager( + self._consumer_memory, + _LEASE_TTL_SECONDS, + _MAX_CANCELLED_TRANSFER_IDS, + ) + self._consumer_rank_resolved = False + self._is_receiving_rank = True + self._tp_rank = 0 + self._tp_size = 1 + self._control_server: ConsumerControlServer | None = None + self._consumer_metrics_log_interval = config.consumer_metrics_log_interval + self._consumer_metrics_started_at = time.monotonic() + # Worker producer + self._producer_memory = ProducerMemoryPool( + config.producer_pool_size, + self._transfer, + ) + self._transfer_metrics_log_interval = config.transfer_metrics_log_interval + self._control_client = control_client + self._topology = topology + self._producer_metrics: Counter[str] = Counter() + self._io_executor = ThreadPoolExecutor( + max_workers=config.transfer_workers, + thread_name_prefix="ec-mooncake-transfer", + ) + self._control_executor = ThreadPoolExecutor( + max_workers=config.control_workers, + thread_name_prefix="ec-mooncake-control", + ) + self._shard_pool: ThreadPoolExecutor | None = None + self._shard_pool_lock = threading.Lock() + self._producer_pushes = ProducerPushManager() + self._push_perf_lock = threading.Lock() + self._push_perf = _PushPerfWindow() + self._active_transfer_batches = 0 + self._queued_transfer_batches = 0 + self._completed_loads: set[str] = set() + self._failed_loads: set[str] = set() + self._shutdown = False + + def _resolve_consumer_rank(self) -> None: + """Place this worker in the consumer receive topology.""" + if self._consumer_rank_resolved: + return + self._consumer_rank_resolved = True + try: + from vllm.distributed.parallel_state import get_pp_group, get_tp_group + + tp_group = get_tp_group() + self._tp_rank = tp_group.rank_in_group + self._tp_size = tp_group.world_size + self._is_receiving_rank = get_pp_group().is_first_rank + except AssertionError: + # Groups are only absent outside a distributed run, where this + # worker is the whole consumer. + self._tp_rank = 0 + self._tp_size = 1 + self._is_receiving_rank = True + + def start_services(self) -> None: + if ( + not self.is_consumer + or self._reservation_zmq_port is None + or self._control_server is not None + ): + return + self._resolve_consumer_rank() + if not self._is_receiving_rank: + # Later pipeline stages hold no encoder outputs, so they need + # neither a receive pool nor a control channel. + return + raw_device = self._buffer_device + device_name = ( + raw_device.lower() if isinstance(raw_device, str) and raw_device else "cuda" + ) + self._consumer_memory.prepare( + torch.device(device_name), + receiving_rank=self._is_receiving_rank, + allow_host=True, + ) + consumer_pool = self._consumer_memory.tensor + if consumer_pool is None: + raise RuntimeError( + "Mooncake push mode requires a registered consumer buffer pool." + ) + base_port = self._reservation_zmq_port + self._control_server = ConsumerControlServer( + "0.0.0.0", + base_port + self._tp_rank, + self._reserve_push_destination, + self._push_status, + self._complete_push, + self._cancel_push, + self._expire_push_reservations, + self._consumer_metrics_log_interval, + peer_ports=[base_port + rank for rank in range(self._tp_size)], + device=consumer_pool.device, + ) + try: + self._control_server.start() + except Exception: + self._control_server.close() + self._control_server = None + raise + + def _maybe_log_consumer_worker_metrics(self) -> None: + now = time.monotonic() + if ( + self._consumer_metrics_log_interval <= 0 + or now - self._consumer_metrics_started_at + < self._consumer_metrics_log_interval + ): + return + with self._consumer_memory.lock: + reservations = self._reservations.active_records() + ready = [ + record.mm_hash + for record in reservations + if record.state is ConsumerReservationState.READY + ] + pending = [ + record.mm_hash + for record in reservations + if record.state is not ConsumerReservationState.READY + ] + metrics = dict(self._consumer_worker_metrics) + self._consumer_worker_metrics.clear() + metrics.update(self._consumer_memory.take_metrics()) + residents, live, retired, pending_frees = self._consumer_memory.stats() + oldest_reservation_ms = max( + ((now - reservation.created_at) * 1000 for reservation in reservations), + default=0.0, + ) + logger.info( + "EC Mooncake consumer worker: lifecycle=%s, reservations_ready=%d, " + "reservations_pending=%d, residents=%d, live=%d, retired=%d, " + "pending_frees=%d, " + "oldest_reservation_ms=%.1f, ready_hashes=%s, pending_hashes=%s", + metrics, + len(ready), + len(pending), + residents, + live, + retired, + pending_frees, + oldest_reservation_ms, + [value[:16] for value in ready[:5]], + [value[:16] for value in pending[:5]], + ) + self._consumer_metrics_started_at = now + + def _expire_push_reservations(self) -> int: + return self._record_expiry_metrics(self._reservations.expire()) + + def _record_expiry_metrics(self, counts: tuple[int, int, int]) -> int: + expired, deferred, tombstones_dropped = counts + self._consumer_worker_metrics["reservations_expired"] += expired + self._consumer_worker_metrics["cancellations_deferred"] += deferred + self._consumer_worker_metrics["cancel_records_dropped"] += tombstones_dropped + return expired + + def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: + transfer_id = str(payload["transfer_id"]) + mm_hash = str(payload["mm_hash"]) + nbytes = int(payload["nbytes"]) + shape = tuple(int(value) for value in payload["shape"]) + dtype_name = str(payload["dtype"]) + dtype = getattr(torch, dtype_name, None) + if dtype is None: + raise ValueError(f"Unsupported torch dtype string: {dtype_name!r}") + expected_nbytes = math.prod(shape) * dtype.itemsize + if expected_nbytes != nbytes: + raise ValueError("shape and dtype do not match nbytes") + + self._expire_push_reservations() + reservation, should_write, reused, expiry_counts = self._reservations.reserve( + transfer_id, mm_hash, nbytes, shape, dtype_name, dtype + ) + self._record_expiry_metrics(expiry_counts) + if reservation is None: + raise RuntimeError("EC consumer buffer pool is full") + if reservation.state in { + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.CANCELLED, + }: + self._consumer_worker_metrics["reservations_cancelled_early"] += 1 + return { + "reservation_id": "", + "dst_session": "", + "dst_ptr": 0, + "nbytes": nbytes, + "write": False, + "ready": False, + "cancelled": True, + } + if reused: + key = ( + "reservations_reused_ready" + if reservation.state is ConsumerReservationState.READY + else "reservations_reused_pending" + ) + self._consumer_worker_metrics[key] += 1 + elif reservation.lease is not None: + self._consumer_worker_metrics["reservations_cached"] += 1 + else: + self._consumer_worker_metrics["reservations_created"] += 1 + assert reservation.allocation is not None + + return { + "reservation_id": reservation.reservation_id, + "dst_session": self._transfer.local_session(), + "dst_ptr": reservation.allocation.tensor.data_ptr(), + "nbytes": reservation.allocation.tensor.nbytes, + "write": should_write, + "ready": reservation.state is ConsumerReservationState.READY, + "cached": reservation.lease is not None, + } + + def _push_status(self, transfer_id: str) -> dict[str, Any] | None: + reservation = self._reservations.status(transfer_id) + if reservation is None: + return None + assert reservation.allocation is not None + return { + "mm_hash": reservation.mm_hash, + "ready": reservation.state is ConsumerReservationState.READY, + "reservation_id": reservation.reservation_id, + "nbytes": reservation.allocation.tensor.nbytes, + "shape": list(reservation.shape), + "dtype": reservation.dtype, + } + + def _complete_push( + self, transfer_id: str, reservation_id: str + ) -> ControlCompletion: + result = self._reservations.complete(transfer_id, reservation_id) + if not result.accepted: + self._consumer_worker_metrics["completions_rejected"] += 1 + elif result.repeated: + self._consumer_worker_metrics["completions_repeated"] += 1 + else: + self._consumer_worker_metrics["completions_accepted"] += 1 + if result.discarded: + self._consumer_worker_metrics["reservations_discarded"] += 1 + return ControlCompletion(result.accepted, result.became_ready) + + def _cancel_push( + self, + transfer_id: str, + reservation_id: str, + abandon: bool = False, + refresh: bool = False, + ) -> bool: + outcome, tombstones_dropped = self._reservations.cancel( + transfer_id, reservation_id, abandon, refresh + ) + metrics = { + CancellationOutcome.REJECTED: "cancellations_rejected", + CancellationOutcome.PRE_RESERVED: "cancellations_pre_reserved", + CancellationOutcome.DEFERRED: "cancellations_deferred", + CancellationOutcome.CANCELLED: "reservations_cancelled", + } + self._consumer_worker_metrics[metrics[outcome]] += 1 + self._consumer_worker_metrics["cancel_records_dropped"] += tombstones_dropped + return outcome is not CancellationOutcome.REJECTED + + def _take_pushed_tensor( + self, spec: ECMooncakeLoadSpec + ) -> tuple[torch.Tensor, MemoryAllocation]: + try: + allocation = self._reservations.take(spec.transfer_id, spec.mm_hash) + except RuntimeError: + self._consumer_worker_metrics["takes_rejected"] += 1 + raise + self._consumer_worker_metrics["reservations_taken"] += 1 + return allocation.tensor, allocation + + def _shard_executor(self) -> ThreadPoolExecutor: + """Use a separate pool so nested shard fan-out cannot deadlock.""" + with self._shard_pool_lock: + if self._shard_pool is None: + self._shard_pool = ThreadPoolExecutor( + max_workers=32, thread_name_prefix="ec-mooncake-shard" + ) + return self._shard_pool + + def _reserve_one(self, addr: str, spec: ECMooncakePushSpec) -> dict[str, Any]: + result = self._control_client.request( + addr, + { + "op": "reserve", + "transfer_id": spec.transfer_id, + "mm_hash": spec.mm_hash, + "nbytes": spec.nbytes, + "shape": list(spec.shape), + "dtype": spec.dtype, + }, + ) + if not isinstance(result, dict): + raise RuntimeError("Invalid EC reservation response") + result["_received_at"] = time.monotonic() + result["addr"] = addr + return result + + def _run_fanout( + self, + tasks: list[Callable[[], _T]], + on_submit: Callable[[int, Future[_T]], None] | None = None, + ) -> list[_T]: + if not tasks: + return [] + futures: list[tuple[int, Future[_T]]] = [] + results: list[_T | None] = [None] * len(tasks) + error: BaseException | None = None + for index, task in enumerate(tasks[1:], 1): + try: + future = self._shard_executor().submit(task) + except Exception as exc: + error = exc + break + futures.append((index, future)) + if on_submit is not None: + on_submit(index, future) + if error is None: + try: + results[0] = tasks[0]() + except Exception as exc: + error = exc + for index, future in futures: + try: + results[index] = future.result() + except Exception as exc: + if error is None: + error = exc + if error is not None: + raise _FanoutError(error, results) + return cast(list[_T], results) + + def _cancel_reservations( + self, + spec: ECMooncakePushSpec, + reservations: list[dict[str, Any]], + *, + refresh: bool = False, + record: ProducerPushRecord | None = None, + ) -> None: + reservations = [ + shard for shard in reservations if not shard.get("cancelled", False) + ] + if not reservations: + return + + def cancel(shard: dict[str, Any]) -> dict[str, Any]: + result = self._control_client.request( + str(shard.get("addr", spec.consumer_zmq)), + make_cancel_request( + spec.transfer_id, + str(shard.get("reservation_id", "")), + abandon=True, + refresh=refresh, + ), + ) + if not isinstance(result, dict) or not result.get("cancelled"): + raise RuntimeError( + f"Could not cancel EC reservation for mm_hash={spec.mm_hash}" + ) + return shard + + def track(_index: int, future: Future[dict[str, Any]]) -> None: + if record is not None: + self._producer_pushes.track_shard_futures([record], [future]) + + self._run_fanout([partial(cancel, shard) for shard in reservations], track) + + def _retry_cancel_reservations( + self, + spec: ECMooncakePushSpec, + reservations: list[dict[str, Any]], + *, + record: ProducerPushRecord | None = None, + ) -> None: + pending = [shard for shard in reservations if not shard.get("cancelled", False)] + error: _FanoutError | None = None + for _ in range(_CANCEL_ATTEMPTS): + try: + self._cancel_reservations(spec, pending, record=record) + except _FanoutError as exc: + error = exc + pending = [ + shard + for index, shard in enumerate(pending) + if exc.results[index] is None + ] + continue + return + assert error is not None + raise error + + def _reserve_remote(self, spec: ECMooncakePushSpec) -> list[dict[str, Any]]: + """Reserve a destination on every shard of the consumer.""" + shards = self._topology.shards(spec.consumer_zmq) + tasks: list[Callable[[], dict[str, Any]]] = [ + partial(self._reserve_one, addr, spec) for addr in shards + ] + try: + return self._run_fanout(tasks) + except _FanoutError as exc: + successful = [result for result in exc.results if isinstance(result, dict)] + try: + self._retry_cancel_reservations(spec, successful) + except _FanoutError as cleanup_error: + raise _ReservationFanoutError(exc, successful) from cleanup_error + raise _ReservationFanoutError(exc, successful) from exc + + def _refresh_remote_reservations( + self, + spec: ECMooncakePushSpec, + reservations: list[dict[str, Any]], + record: ProducerPushRecord | None = None, + ) -> list[dict[str, Any]]: + stale = [ + shard + for shard in reservations + if not shard.get("ready", False) + and not shard.get("cached", False) + and not shard.get("cancelled", False) + ] + try: + self._cancel_reservations(spec, stale, refresh=True, record=record) + except _FanoutError as exc: + pending = [ + shard for index, shard in enumerate(stale) if exc.results[index] is None + ] + try: + self._retry_cancel_reservations(spec, pending, record=record) + except _FanoutError as cleanup_error: + raise exc from cleanup_error + raise + return self._reserve_remote(spec) + + @staticmethod + def _validate_push_source(push: ProducerPushRecord) -> None: + source = push.source + assert source is not None + tensor = source.tensor + spec = push.spec + if tuple(tensor.shape) != tuple(spec.shape): + raise ValueError(f"EC source shape mismatch for mm_hash={spec.mm_hash}") + if str(tensor.dtype).split(".")[-1] != spec.dtype: + raise ValueError(f"EC source dtype mismatch for mm_hash={spec.mm_hash}") + if not tensor.is_contiguous(): + raise ValueError(f"EC source must be contiguous for mm_hash={spec.mm_hash}") + if tensor.nbytes != spec.nbytes: + raise ValueError(f"EC source size mismatch for mm_hash={spec.mm_hash}") + + def start_save_caches( + self, + metadata: ECMooncakeConnectorMetadata, + encoder_cache: dict[str, torch.Tensor] | None = None, + **kwargs: Any, + ) -> None: + for spec in metadata.pushes: + self._producer_pushes.reserve( + spec, + partial(self._submit_reservation, spec), + ) + if not isinstance(encoder_cache, dict): + return + for mm_hash in dict.fromkeys(spec.mm_hash for spec in metadata.pushes): + tensor = encoder_cache.get(mm_hash) + if tensor is not None: + self._bind_push_source(tensor, mm_hash) + + def _submit_reservation( + self, spec: ECMooncakePushSpec + ) -> Future[list[dict[str, Any]]]: + return self._control_executor.submit(self._reserve_remote, spec) + + def start_load_caches( + self, + metadata: ECMooncakeConnectorMetadata, + encoder_cache: dict[str, torch.Tensor], + **kwargs: Any, + ) -> None: + self._resolve_consumer_rank() + if not self._is_receiving_rank: + # Later pipeline stages never gather multimodal embeddings. + return + self._transfer.ensure_ready() + raw_buf = self._buffer_device + buf = raw_buf.lower() if isinstance(raw_buf, str) and raw_buf else "cuda" + if buf == "cuda" and not torch.accelerator.is_available(): + raise RuntimeError( + "ECMooncakeConnector requires CUDA for ec_buffer_device=cuda" + ) + self._reservations.retire_stale(encoder_cache) + + for spec in metadata.loads: + if spec.mm_hash in encoder_cache: + if spec.pushed: + # The spec's id is one shard's; cancel by transfer. + self._cancel_push(spec.transfer_id, "") + self._completed_loads.add(spec.mm_hash) + continue + if spec.local: + tensor = self._consumer_memory.take_resident( + spec.mm_hash, tuple(spec.shape), spec.dtype + ) + elif spec.pushed: + try: + tensor, _ = self._take_pushed_tensor(spec) + except RuntimeError as e: + logger.warning("EC Mooncake pushed load failed: %s", e) + tensor = None + else: + logger.warning( + "EC Mooncake load for mm_hash=%s has no transfer to take", + spec.mm_hash, + ) + tensor = None + if tensor is None: + self._failed_loads.add(spec.mm_hash) + else: + encoder_cache[spec.mm_hash] = tensor + self._completed_loads.add(spec.mm_hash) + + def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: + started_at = time.monotonic() + with self._push_perf_lock: + self._queued_transfer_batches -= 1 + self._active_transfer_batches += 1 + + queue_waits_ms = [] + for push in pushes: + assert push.source_at is not None + queue_waits_ms.append(max(0, started_at - push.source_at) * 1000) + stage_ms = {"queue": sum(queue_waits_ms), **dict.fromkeys(_PUSH_STAGES, 0.0)} + ready: list[tuple[ProducerPushRecord, dict[str, Any]]] = [] + written_pushes: dict[str, ProducerPushRecord] = {} + failed = False + failure: Exception | None = None + try: + for push in pushes: + self._validate_push_source(push) + stage_started_at = time.monotonic() + reservations = self._producer_pushes.resolve_reservations(push) + stale = [ + index + for index, shard in enumerate(reservations) + if not shard.get("ready", False) + and not shard.get("cancelled", False) + and time.monotonic() - float(shard.get("_received_at", started_at)) + >= _RESERVATION_REFRESH_SECONDS + ] + if stale: + reservations = self._refresh_remote_reservations( + push.spec, reservations, push + ) + self._producer_pushes.replace_reservations(push, reservations) + stage_ms["reserve"] += (time.monotonic() - stage_started_at) * 1000 + self._producer_pushes.begin_writing(push) + writable = [ + shard + for shard in reservations + if not shard.get("cached", False) + and not shard.get("cancelled", False) + and shard.get("write", True) + ] + source = push.source + assert source is not None + if writable and source.ready_event is not None: + stage_started_at = time.monotonic() + source.ready_event.synchronize() + stage_ms["cuda"] += (time.monotonic() - stage_started_at) * 1000 + for shard in writable: + if int(shard["nbytes"]) != source.tensor.nbytes: + raise RuntimeError( + "Reserved EC size does not match tensor for " + f"mm_hash={push.spec.mm_hash}" + ) + ready.append((push, shard)) + written_pushes.setdefault(push.spec.transfer_id, push) + if ready: + # Stage each source once, then write it to every destination. + source_index = { + push.spec.transfer_id: index + for index, push in enumerate(written_pushes.values()) + } + tensors = [ + push.source.tensor + for push in written_pushes.values() + if push.source + ] + lengths = [tensor.nbytes for tensor in tensors] + stage_started_at = time.monotonic() + staged = self._producer_memory.stage(tensors) + registered_sources: list[int] = [] + if staged is not None: + sources = staged.tensors + # The NIC reads outside the CUDA stream. + if sources and sources[0].device.type == "cuda": + torch.accelerator.current_stream( + sources[0].device + ).synchronize() + else: + sources = tensors + registered_sources = self._transfer.acquire_sources(tensors) + addresses = [tensor.data_ptr() for tensor in sources] + stage_ms["register"] = (time.monotonic() - stage_started_at) * 1000 + try: + by_session: dict[str, list[tuple[int, int]]] = {} + session_records: dict[str, dict[str, ProducerPushRecord]] = {} + for push, shard in ready: + session = str(shard["dst_session"]) + by_session.setdefault(session, []).append( + (source_index[push.spec.transfer_id], int(shard["dst_ptr"])) + ) + session_records.setdefault(session, {})[ + push.spec.transfer_id + ] = push + stage_started_at = time.monotonic() + + def write(session: str, items: list[tuple[int, int]]) -> None: + self._transfer.write( + session, + [addresses[index] for index, _ in items], + [dst for _, dst in items], + [lengths[index] for index, _ in items], + ) + + sessions = list(by_session.items()) + + # Write shards concurrently to avoid serial TP latency. + def track_write(index: int, future: Future[None]) -> None: + session = sessions[index][0] + self._producer_pushes.track_shard_futures( + list(session_records[session].values()), [future] + ) + + writes: list[Callable[[], None]] = [ + partial(write, *session) for session in sessions + ] + self._run_fanout(writes, track_write) + stage_ms["rdma"] = (time.monotonic() - stage_started_at) * 1000 + finally: + stage_started_at = time.monotonic() + if staged is not None: + self._producer_memory.release(staged) + self._transfer.release_sources(registered_sources) + stage_ms["unregister"] = ( + time.monotonic() - stage_started_at + ) * 1000 + + self._producer_pushes.begin_notifying(pushes) + stage_started_at = time.monotonic() + self._notify_completions(ready) + stage_ms["complete"] = (time.monotonic() - stage_started_at) * 1000 + self._producer_pushes.complete(pushes) + except Exception as exc: + # Report asynchronously; raising here would fail EngineCore. + failed = True + failure = exc + logger.exception( + "EC Mooncake push batch failed for mm_hashes=%s", + [push.spec.mm_hash for push in pushes], + ) + self._producer_pushes.settle_all(pushes) + self._abandon_pushes(pushes) + finally: + if failure is not None: + self._producer_pushes.fail(pushes, failure) + stage_ms["total"] = (time.monotonic() - started_at) * 1000 + self._record_push_perf( + stage_ms, + stage_max_ms={"queue": max(queue_waits_ms, default=0.0)}, + item_count=len(pushes), + byte_count=sum(push.spec.nbytes for push in written_pushes.values()), + skipped_items=len(pushes) - len(written_pushes), + failed=failed, + ) + + def _notify_completions( + self, notifications: list[tuple[ProducerPushRecord, dict[str, Any]]] + ) -> None: + """Tell the consumer, in one message per destination, what landed.""" + if not notifications: + return + by_destination: dict[str, list[tuple[ProducerPushRecord, dict[str, Any]]]] = {} + for push, reservation in notifications: + by_destination.setdefault( + str(reservation.get("addr", push.spec.consumer_zmq)), [] + ).append((push, reservation)) + destinations = list(by_destination.items()) + + def notify( + consumer_zmq: str, + items: list[tuple[ProducerPushRecord, dict[str, Any]]], + ) -> None: + result = self._control_client.request( + consumer_zmq, + { + "op": "complete_batch", + "items": [ + { + "transfer_id": push.spec.transfer_id, + "reservation_id": reservation["reservation_id"], + } + for push, reservation in items + ], + }, + ) + completions = result.get("items", []) if isinstance(result, dict) else [] + if len(completions) != len(items): + raise RuntimeError("Malformed EC completion response") + for (push, _), completion in zip(items, completions): + if not completion.get("completed"): + raise RuntimeError( + f"Unknown EC reservation for mm_hash={push.spec.mm_hash}" + ) + + def track(index: int, future: Future[None]) -> None: + records = { + push.spec.transfer_id: push for push, _ in destinations[index][1] + } + self._producer_pushes.track_shard_futures(list(records.values()), [future]) + + self._run_fanout( + [ + partial(notify, destination, items) + for destination, items in destinations + ], + track, + ) + + @staticmethod + def _known_reservations(record: ProducerPushRecord) -> list[dict[str, Any]]: + if record.reservations: + return list(record.reservations) + reservations: list[dict[str, Any]] = [] + for future in record.reservation_futures: + try: + reservations.extend(future.result()) + except _ReservationFanoutError as exc: + reservations.extend(exc.partial_reservations) + except Exception: + continue + return reservations + + def _abandon_pushes(self, pushes: list[ProducerPushRecord]) -> None: + """Release the consumer-side reservations of a batch that failed.""" + for push in pushes: + shards = self._known_reservations(push) + if not shards: + shards = [{"addr": push.spec.consumer_zmq, "reservation_id": ""}] + try: + self._retry_cancel_reservations(push.spec, shards, record=push) + except _FanoutError: + logger.exception( + "Failed to abandon EC reservations for transfer_id=%s", + push.spec.transfer_id, + ) + + def _record_push_perf( + self, + stage_ms: dict[str, float], + *, + stage_max_ms: dict[str, float], + item_count: int, + byte_count: int, + skipped_items: int, + failed: bool, + ) -> None: + now = time.monotonic() + report: tuple[_PushPerfWindow, int, int] | None = None + with self._push_perf_lock: + self._active_transfer_batches -= 1 + perf = self._push_perf + perf.batches += 1 + perf.items += item_count + perf.bytes += byte_count + perf.skipped_items += skipped_items + perf.failures += int(failed) + for stage, elapsed_ms in stage_ms.items(): + perf.stage_totals_ms[stage] = ( + perf.stage_totals_ms.get(stage, 0.0) + elapsed_ms + ) + perf.stage_max_ms[stage] = max( + perf.stage_max_ms.get(stage, 0.0), + stage_max_ms.get(stage, elapsed_ms), + ) + if ( + self._transfer_metrics_log_interval > 0 + and now - perf.started_at >= self._transfer_metrics_log_interval + ): + report = ( + perf, + self._active_transfer_batches, + self._queued_transfer_batches, + ) + self._push_perf = _PushPerfWindow(started_at=now) + if report is None: + return + perf, active_batches, queued_batches = report + batches = max(perf.batches, 1) + items = max(perf.items, 1) + stage_parts = [] + for stage in ("queue", *_PUSH_STAGES, "total"): + divisor = items if stage == "queue" else batches + average = perf.stage_totals_ms.get(stage, 0.0) / divisor + maximum = perf.stage_max_ms.get(stage, 0.0) + stage_parts.append(f"{stage}_ms={average:.1f}/{maximum:.1f}") + stage_summary = " ".join(stage_parts) + producer_metrics = dict(self._producer_metrics) + self._producer_metrics.clear() + logger.info( + "EC Mooncake push perf: batches=%d items=%d bytes=%d " + "batch_items=%.1f skipped=%d failures=%d active=%d queued=%d " + "producer=%s queue_item_avg/max and stage_batch_avg/max: %s", + perf.batches, + perf.items, + perf.bytes, + perf.items / batches, + perf.skipped_items, + perf.failures, + active_batches, + queued_batches, + producer_metrics, + stage_summary, + ) + + def _flush_pending_pushes(self) -> None: + self._producer_pushes.submit_batches( + self._io_executor, + self._push_batch, + self._note_push_batch_queued, + ) + + def _note_push_batch_queued(self) -> None: + with self._push_perf_lock: + self._queued_transfer_batches += 1 + + def _bind_push_source(self, tensor: torch.Tensor, mm_hash: str) -> None: + ready_event = None + if tensor.device.type == "cuda": + ready_event = torch.Event() + ready_event.record(torch.accelerator.current_stream(tensor.device)) + self._producer_pushes.bind_source(mm_hash, tensor, ready_event) + + def _cancel_orphaned_reservation(self, record: ProducerPushRecord) -> None: + try: + reservations = self._producer_pushes.resolve_reservations(record) + except Exception: + reservations = self._known_reservations(record) + known = bool(reservations) + reservations = [ + shard + for shard in reservations + if not shard.get("cached", False) and not shard.get("cancelled", False) + ] + if not known: + reservations = [{"addr": record.spec.consumer_zmq, "reservation_id": ""}] + error = None + try: + self._retry_cancel_reservations(record.spec, reservations, record=record) + except _FanoutError as exc: + error = exc + self._producer_pushes.finish_cancel(record) + if error is not None: + raise error + + def get_finished( + self, finished_req_ids: set[str] + ) -> tuple[set[str] | None, set[str] | None]: + if not self.is_producer: + return None, None + + for record in self._producer_pushes.cancel_requests(finished_req_ids): + self._producer_pushes.submit_cancel( + record, + self._io_executor, + self._cancel_orphaned_reservation, + ) + return None, None + + def save_caches( + self, encoder_cache: dict[str, torch.Tensor], mm_hash: str, **kwargs: Any + ) -> None: + if not self.is_producer: + return + tensor = encoder_cache[mm_hash] + self._bind_push_source(tensor, mm_hash) + + def build_connector_worker_meta(self) -> ECMooncakeWorkerMetadata | None: + if self.is_consumer and not self._is_receiving_rank: + # `loaded` is intersected across reporting ranks, so a stage that + # never loads must not report at all rather than report nothing. + return None + + self._flush_pending_pushes() + failures = self._producer_pushes.poll() + self._producer_metrics["saves_failed"] += len(failures) + for mm_hash, error in failures: + logger.error( + "EC Mooncake async save failed for mm_hash=%s: %s", + mm_hash, + error, + ) + reclaimed = self._consumer_memory.drain_reclaimed() + meta = ECMooncakeWorkerMetadata( + loaded=self._completed_loads, + failed_loads=self._failed_loads, + reclaimed=reclaimed, + pending_loads=False, + pending_saves=self._producer_pushes.pending, + ) + self._completed_loads = set() + self._failed_loads = set() + if self.is_consumer: + self._maybe_log_consumer_worker_metrics() + return meta + + def close(self) -> None: + if self._shutdown: + return + self._shutdown = True + self._flush_pending_pushes() + self._io_executor.shutdown(wait=True, cancel_futures=True) + if self._shard_pool is not None: + self._shard_pool.shutdown(wait=True, cancel_futures=True) + self._control_executor.shutdown(wait=True, cancel_futures=True) + # Every thread that could hold a control socket is stopped by now. + self._control_client.close() + if self._control_server is not None: + self._control_server.close() + self._consumer_memory.close() + self._producer_memory.close() + self._transfer.close() + + def __del__(self) -> None: + with suppress(Exception): + self.close() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 3d5a31a64358..8a271db69d46 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -1,2868 +1,157 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -""" -Encoder-cache (EC) connector backed by Mooncake TransferEngine. - -Used in disaggregated setups where an encoder / prefill instance produces -multimodal encoder outputs and a decode instance loads them over RDMA-capable -Mooncake transport instead of shared filesystem. -""" +"""Encoder-cache connector backed by Mooncake TransferEngine.""" from __future__ import annotations -import bisect -import math -import threading -import time -import uuid -from collections import Counter, OrderedDict, deque -from collections.abc import Callable, Collection -from concurrent.futures import Future, ThreadPoolExecutor +from collections.abc import Collection from contextlib import suppress -from dataclasses import dataclass, field -from typing import Any, Generic, TypeVar - -import torch -import zmq +from typing import TYPE_CHECKING, Any -from vllm.config import VllmConfig from vllm.distributed.ec_transfer.ec_connector.base import ( ECConnectorBase, ECConnectorMetadata, ECConnectorRole, ECConnectorWorkerMetadata, ) -from vllm.distributed.ec_transfer.ec_connector.cpu.common import ( - _get_encoder_cache_hidden_dim, +from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeConnectorMetadata, + ECMooncakeLoadSpec, + ECMooncakePushSpec, + ECMooncakeWorkerMetadata, ) -from vllm.logger import init_logger -from vllm.utils.network_utils import get_ip -from vllm.v1.core.sched.output import SchedulerOutput -from vllm.v1.outputs import ECConnectorOutput - -logger = init_logger(__name__) - -_T = TypeVar("_T") - -_LEASE_TTL_SECONDS = 300 -_RESERVATION_REFRESH_SECONDS = _LEASE_TTL_SECONDS / 2 -_RESERVATION_REAP_INTERVAL_SECONDS = 1 -_DRAIN_MIN_INTERVAL = 0.005 -# Readiness notifications are advisory: the scheduler also learns from the -# reserve reply. Cap the queue so a shard nobody subscribed to cannot grow -# without bound. -_MAX_PENDING_EVENTS = 4096 -# A cancelled transfer stays on the scheduler's ignore list for as long as the -# worker refuses to reserve it again. The count is a backstop for a rate that -# outruns that TTL; the race it guards is a single drain interval wide. -_MAX_CANCELLED_TRANSFER_IDS = 1 << 16 - -_MOONCAKE_IMPORT_ERROR: ImportError | None -try: - from mooncake.engine import TransferEngine -except ImportError as e: - TransferEngine = None # type: ignore[misc, assignment] - _MOONCAKE_IMPORT_ERROR = e -else: - _MOONCAKE_IMPORT_ERROR = None - - -@dataclass -class ECMooncakeLoadSpec: - """Per-item metadata shipped from scheduler to worker (pickle-friendly).""" - - mm_hash: str - num_token: int - nbytes: int - shape: tuple[int, ...] - dtype: str - pushed: bool = False - transfer_id: str = "" - reservation_id: str = "" - # The consumer pool still holds this item, so the load is a local handoff: - # no transfer, no producer. - local: bool = False - - -@dataclass -class ECMooncakePushSpec: - """Destination reservation requested before an encoder tensor is ready.""" - - mm_hash: str - nbytes: int - shape: tuple[int, ...] - dtype: str - consumer_zmq: str - transfer_id: str - request_id: str = "" - - -@dataclass -class ECMooncakeConnectorMetadata(ECConnectorMetadata): - """Worker-side metadata for one scheduler step.""" - - loads: list[ECMooncakeLoadSpec] = field(default_factory=list) - pushes: list[ECMooncakePushSpec] = field(default_factory=list) - - def add_load(self, spec: ECMooncakeLoadSpec) -> None: - self.loads.append(spec) - - def add_push(self, spec: ECMooncakePushSpec) -> None: - self.pushes.append(spec) - - -@dataclass -class ECMooncakeWorkerMetadata(ECConnectorWorkerMetadata): - """Completion state reported from workers to the scheduler.""" - - loaded: set[str] = field(default_factory=set) - failed_loads: set[str] = field(default_factory=set) - # Items the receive pool dropped under pressure. The scheduler assumes an - # evicted item stays resident until told otherwise. - reclaimed: set[str] = field(default_factory=set) - pending_loads: bool = False - pending_saves: bool = False - - def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadata: - assert isinstance(other, ECMooncakeWorkerMetadata) - return ECMooncakeWorkerMetadata( - # Every tensor-parallel rank gathers the embedding from its own - # cache, so an item counts as loaded only where all of them have - # it; one rank falling short must fail the load rather than leave - # the scheduler believing it is ready. - loaded=self.loaded & other.loaded, - failed_loads=self.failed_loads | other.failed_loads, - reclaimed=self.reclaimed | other.reclaimed, - pending_loads=self.pending_loads or other.pending_loads, - pending_saves=self.pending_saves or other.pending_saves, - ) - - -@dataclass -class _PushSourceRegistration: - tensor: torch.Tensor - nbytes: int - users: int = 1 - - -@dataclass -class _ConsumerPoolAllocation: - offset: int - size: int - tensor: torch.Tensor - - -@dataclass -class _PushReservation: - mm_hash: str - reservation_id: str - allocation: _ConsumerPoolAllocation - shape: tuple[int, ...] - dtype: str - ready: bool = False - owns_allocation: bool = True - discard_on_complete: bool = False - created_at: float = field(default_factory=time.monotonic) - expires_at: float = 0 - - -@dataclass(frozen=True) -class _PushCompletion: - accepted: bool - became_ready: bool = False - - -@dataclass -class _PendingPush: - tensor: torch.Tensor - spec: ECMooncakePushSpec - reservation: Future[list[dict[str, Any]]] - ready_event: torch.Event | None - enqueued_at: float - - -@dataclass -class _PushPerfWindow: - started_at: float = field(default_factory=time.monotonic) - batches: int = 0 - items: int = 0 - bytes: int = 0 - skipped_items: int = 0 - failures: int = 0 - stage_totals_ms: dict[str, float] = field(default_factory=dict) - stage_max_ms: dict[str, float] = field(default_factory=dict) - - -class _ControlChannel: - """Reusable REQ sockets for the ZMQ control plane. - - One context and one connection per message costs a thread spawn plus a - TCP handshake, and the push path sends one message per reserve, complete - and cancel. Sockets are cached per thread because a REQ socket is neither - thread-safe nor usable after a failed exchange. - """ - - def __init__(self, timeout_ms: int): - self._context = zmq.Context() - self._timeout_ms = timeout_ms - self._local = threading.local() - - def _sockets(self) -> dict[str, zmq.Socket]: - sockets = getattr(self._local, "sockets", None) - if sockets is None: - sockets = {} - self._local.sockets = sockets - return sockets - - def _discard(self, addr: str) -> None: - socket = self._sockets().pop(addr, None) - if socket is not None: - socket.close(linger=0) - - def send(self, addr: str, payload: dict[str, Any]) -> dict[str, Any]: - sockets = self._sockets() - socket = sockets.get(addr) - if socket is None: - socket = self._context.socket(zmq.REQ) - socket.setsockopt(zmq.RCVTIMEO, self._timeout_ms) - socket.setsockopt(zmq.SNDTIMEO, self._timeout_ms) - socket.setsockopt(zmq.LINGER, 0) - socket.connect(addr) - sockets[addr] = socket - try: - socket.send_json(payload) - response = socket.recv_json() - except Exception: - # A REQ socket cannot recover from a half-finished exchange. - self._discard(addr) - raise - assert isinstance(response, dict) - return response - - def request(self, addr: str, payload: dict[str, Any]) -> Any: - response = self.send(addr, payload) - if not response.get("ok"): - raise RuntimeError(response.get("error", "EC control request failed")) - return response.get("result") - - def close(self) -> None: - # Callers must have stopped every thread that used this channel. - self._context.destroy(linger=0) - - -class _ResidentPool(Generic[_T]): - """Content-addressed entries kept until their space is needed. - - Both sides of the connector hold the same thing under different names: a - map from mm_hash to a device resource, a count of who is using it, and an - eviction order over the rest. This is `BlockPool`'s accounting for - variable-sized entries: `acquire`/`release` mirror `touch`/`free_blocks`, - and `evict_lru` mirrors the reclaim inside `get_new_blocks`. - - An unreferenced entry stays resident. Eviction is driven by pressure, so - the entry serves whoever needs it next instead of being transferred again. - """ - - def __init__(self, capacity: int): - self.capacity = capacity - self.used = 0 - self._entries: dict[str, tuple[_T, int]] = {} - self._refs: Counter[str] = Counter() - # Unreferenced entries in eviction order, oldest first. - self._evictable: OrderedDict[str, None] = OrderedDict() - - def __len__(self) -> int: - return len(self._entries) - - def __contains__(self, key: str) -> bool: - return key in self._entries - - @property - def num_evictable(self) -> int: - return len(self._evictable) - - def referenced(self) -> list[str]: - """Keys that are in use. `_refs` only holds entries above zero.""" - return list(self._refs) - - def referenced_or_retired(self) -> list[str]: - """Every key held, in insertion order.""" - return list(self._entries) - - def get(self, key: str) -> _T | None: - entry = self._entries.get(key) - return entry[0] if entry is not None else None - - def insert(self, key: str, value: _T, nbytes: int) -> None: - """Add a referenced entry, replacing any previous one.""" - previous = self._entries.get(key) - if previous is not None: - self.used -= previous[1] - self._entries[key] = (value, nbytes) - self.used += nbytes - self.pin(key) - - def pin(self, key: str) -> _T | None: - """Mark an entry as in use without counting a new reference. - - For a holder whose references are discovered by scanning rather than - released in pairs, `pin`/`retire` are the matching operations. - """ - entry = self._entries.get(key) - if entry is None: - return None - self._evictable.pop(key, None) - self._refs[key] = max(1, self._refs[key]) - return entry[0] - - def retire(self, key: str) -> None: - """Drop every reference; the entry is evictable from now on.""" - if key not in self._entries: - return - self._refs.pop(key, None) - self._evictable[key] = None - - def refresh(self, key: str) -> None: - """Move an unreferenced entry to the back of the eviction order.""" - if key in self._evictable: - self._evictable.move_to_end(key) - - def acquire(self, key: str) -> _T | None: - """Take one reference so pressure cannot evict the entry.""" - entry = self._entries.get(key) - if entry is None: - return None - self._evictable.pop(key, None) - self._refs[key] += 1 - return entry[0] - - def release(self, key: str) -> None: - """Drop one reference; the entry becomes evictable at zero.""" - if key not in self._entries: - return - count = self._refs[key] - 1 - if count > 0: - self._refs[key] = count - return - self._refs.pop(key, None) - self._evictable[key] = None - - def evict_lru(self, evict: Callable[[str, _T], bool]) -> str | None: - """Drop the oldest entry `evict` accepts, and return its key. - - `evict` returns False for an entry that cannot go yet (a lease the - remote side still holds, a deregistration that failed). Those keep - their place in the order and the next candidate is tried. - """ - for key in list(self._evictable): - value, nbytes = self._entries[key] - if not evict(key, value): - continue - self._evictable.pop(key, None) - del self._entries[key] - self._refs.pop(key, None) - self.used -= nbytes - return key - return None - - def clear(self) -> None: - self._entries.clear() - self._refs.clear() - self._evictable.clear() - self.used = 0 - - -class _ContiguousAllocator: - def __init__(self, capacity: int, alignment: int = 256): - self.capacity = capacity - self.alignment = alignment - self._free = [(0, capacity)] - - def allocate(self, nbytes: int) -> tuple[int, int] | None: - size = math.ceil(nbytes / self.alignment) * self.alignment - for index, (offset, available) in enumerate(self._free): - if size > available: - continue - if size == available: - self._free.pop(index) - else: - self._free[index] = (offset + size, available - size) - return offset, size - return None - - def free(self, offset: int, size: int) -> None: - index = bisect.bisect_left(self._free, (offset, size)) - self._free.insert(index, (offset, size)) - # Coalesce with the neighbours only; the rest of the list is already - # merged, so a full re-scan per free is wasted work. - if index + 1 < len(self._free): - next_offset, next_size = self._free[index + 1] - if offset + size == next_offset: - self._free[index] = (offset, size + next_size) - self._free.pop(index + 1) - if index > 0: - previous_offset, previous_size = self._free[index - 1] - current_offset, current_size = self._free[index] - if previous_offset + previous_size == current_offset: - self._free[index - 1] = ( - previous_offset, - previous_size + current_size, - ) - self._free.pop(index) - - -class ECMooncakeControlServer: - """Expose consumer reservations over a lightweight ZMQ control channel.""" - - def __init__( - self, - host: str, - port: int, - reserve: Callable[[dict[str, Any]], dict[str, Any]], - status: Callable[[str], dict[str, Any] | None], - complete: Callable[[str, str], _PushCompletion], - cancel: Callable[[str, str, bool], bool], - reap: Callable[[], int], - metrics_log_interval: float = 10, - peer_ports: list[int] | None = None, - device: torch.device | None = None, - ): - self.host = host - self.port = port - self.peer_ports = peer_ports or [port] - self._device = device - self.event_port: int | None = None - self._reserve = reserve - self._status = status - self._complete = complete - self._cancel = cancel - self._reap = reap - self._metrics_log_interval = metrics_log_interval - self._stop = threading.Event() - self._started = threading.Event() - self._thread: threading.Thread | None = None - self._startup_error: Exception | None = None - - def start(self) -> None: - def loop() -> None: - if self._device is not None and self._device.type == "cuda": - # Reserving can retire an entry, and the event that orders its - # reuse is created on the recording thread's device rather - # than the stream's. A thread starts on device 0, which under - # a shard-local CUDA_VISIBLE_DEVICES is a peer's GPU, so - # without this every shard but the first strands a primary - # context there. The event orders correctly either way; what - # it costs is a few hundred MiB on someone else's card. - torch.accelerator.set_device_index(self._device.index or 0) - context = zmq.Context() - socket = context.socket(zmq.REP) - event_socket = context.socket(zmq.PUSH) - pending_events: deque[dict[str, Any]] = deque() - metrics: Counter[str] = Counter() - - def queue_event(event: dict[str, Any]) -> None: - # The shard tag lets the scheduler tell each rank's readiness - # apart; a transfer is only loadable once every rank has it. - event["shard"] = self.port - if len(pending_events) >= _MAX_PENDING_EVENTS: - pending_events.popleft() - metrics["events_dropped"] += 1 - pending_events.append(event) - metrics["events_queued"] += 1 +from vllm.distributed.ec_transfer.ec_connector.mooncake.scheduler import ( + ECMooncakeScheduler, +) +from vllm.distributed.ec_transfer.ec_connector.mooncake.worker import ECMooncakeWorker - metrics_started_at = time.monotonic() - last_reap_at = metrics_started_at - socket.setsockopt(zmq.RCVTIMEO, 100) - try: - socket.bind(f"tcp://{self.host}:{self.port}") - self.event_port = event_socket.bind_to_random_port(f"tcp://{self.host}") - except Exception as e: - self._startup_error = e - self._started.set() - socket.close(linger=0) - event_socket.close(linger=0) - context.term() - return - self._started.set() - try: - while not self._stop.is_set(): - while pending_events: - try: - event_socket.send_json( - pending_events[0], flags=zmq.DONTWAIT - ) - except zmq.Again: - break - pending_events.popleft() - metrics["events_sent"] += 1 - now = time.monotonic() - if now - last_reap_at >= _RESERVATION_REAP_INTERVAL_SECONDS: - metrics["reservations_reaped"] += self._reap() - last_reap_at = now - if ( - self._metrics_log_interval > 0 - and now - metrics_started_at >= self._metrics_log_interval - ): - logger.info( - "EC Mooncake consumer control: requests=%s, " - "events_queued=%d, events_sent=%d, events_dropped=%d, " - "event_backlog=%d, reservations_reaped=%d", - { - key.removeprefix("request_"): value - for key, value in metrics.items() - if key.startswith("request_") - }, - metrics["events_queued"], - metrics["events_sent"], - metrics["events_dropped"], - len(pending_events), - metrics["reservations_reaped"], - ) - metrics.clear() - metrics_started_at = now - try: - request = socket.recv_json() - except zmq.Again: - continue - try: - op = request.get("op") - result: Any = None - metrics[f"request_{op}"] += 1 - if op == "reserve": - result = self._reserve(request) - if result.get("ready"): - transfer_id = str(request["transfer_id"]) - status = self._status(transfer_id) - if status is not None: - queue_event({"transfer_id": transfer_id, **status}) - elif op == "status": - result = self._status(str(request["transfer_id"])) - elif op == "event_port": - result = self.event_port - elif op == "peers": - # Every consumer shard receives its own copy, so a - # producer holding one address needs the rest. - result = {"ports": self.peer_ports} - elif op in ("complete", "complete_batch"): - items = ( - request["items"] - if op == "complete_batch" - else [request] - ) - completions = [] - for item in items: - transfer_id = str(item["transfer_id"]) - completion = self._complete( - transfer_id, - str(item["reservation_id"]), - ) - completions.append( - { - "completed": completion.accepted, - "became_ready": completion.became_ready, - } - ) - if not completion.became_ready: - continue - status = self._status(transfer_id) - if status is not None: - queue_event({"transfer_id": transfer_id, **status}) - result = ( - {"items": completions} - if op == "complete_batch" - else completions[0] - ) - elif op == "cancel": - result = { - "cancelled": self._cancel( - str(request["transfer_id"]), - str(request.get("reservation_id", "")), - bool(request.get("abandon", False)), - ) - } - else: - raise ValueError(f"unknown control op: {op!r}") - socket.send_json({"ok": True, "result": result}) - except Exception as e: - socket.send_json({"ok": False, "error": str(e)}) - finally: - socket.close(linger=0) - event_socket.close(linger=0) - context.term() +if TYPE_CHECKING: + import torch - self._thread = threading.Thread( - target=loop, name="ec-mooncake-control", daemon=True - ) - self._thread.start() - if not self._started.wait(timeout=5): - raise RuntimeError("EC Mooncake control channel failed to start") - if self._startup_error is not None: - raise RuntimeError("EC Mooncake control channel failed to bind") from ( - self._startup_error - ) - logger.info( - "EC Mooncake control channel listening on tcp://%s:%d (events tcp://%s:%d)", - self.host, - self.port, - self.host, - self.event_port, - ) + from vllm.config import VllmConfig + from vllm.v1.core.sched.output import SchedulerOutput + from vllm.v1.outputs import ECConnectorOutput + from vllm.v1.request import Request - def shutdown(self) -> None: - if self._thread is None: - return - self._stop.set() - self._thread.join() +__all__ = [ + "ECMooncakeConnector", + "ECMooncakeConnectorMetadata", + "ECMooncakeLoadSpec", + "ECMooncakePushSpec", + "ECMooncakeWorkerMetadata", +] class ECMooncakeConnector(ECConnectorBase): - """ - EC connector using Mooncake TransferEngine for GPU tensor transport. + """Preserve the public API while delegating to one process-role component. - The producer pushes each encoder output into a receive buffer the consumer - reserved for it, so the transfer overlaps encoding instead of waiting for - the consumer to ask. An item the consumer's encoder cache evicted stays in - that pool and is handed back locally; when neither has it, the load fails - with a retryable error so the caller can re-issue the request. - - Extra config (``ec_connector_extra_config``): - - - ``mooncake_protocol`` (optional): Passed to ``TransferEngine.initialize`` - (default ``"rdma"``). - - ``consumer_buffer_pool_size`` (consumer, optional): Bytes reserved for a - long-lived registered CUDA receive arena (default ``ec_buffer_size``). - - ``reservation_zmq_port`` (consumer worker, required): Exposes registered - receive addresses over ZMQ. Replica ``d`` of the first pipeline stage owns - the block starting at ``port + d * tensor_parallel_size``; tensor-parallel - rank ``r`` in that block listens on ``block + r``, and rank 0 reports the - whole block, so a producer only needs the block's first address. - - ``reservation_zmq_addr`` (consumer scheduler, required): Address of the - consumer control channel. Defaults to ``tcp://127.0.0.1:``. - - ``transfer_max_workers`` (optional): Maximum concurrent Mooncake transfer - batches (default ``4``). - - ``control_max_workers`` (optional): Maximum concurrent reservation requests - issued by a producer (default ``8``). - - ``transfer_metrics_log_interval`` (optional): Seconds between aggregated - push-transfer performance logs (default ``10``; ``0`` disables them). - - ``consumer_metrics_log_interval`` (optional): Seconds between aggregated - consumer lifecycle logs (default ``10``; ``0`` disables them). - - Parallelism: consumers may use tensor, pipeline and data parallelism. - Producers must be unsharded and unreplicated: one copy of each encoder - output is held and addressed directly, so splitting the producer would only - duplicate the push. - - Only the first pipeline stage holds encoder outputs, and each tensor-parallel - rank there gathers from its own cache, so every rank exposes a control - channel and the producer writes into all of them concurrently from one - registered source. That costs bandwidth but not latency, and avoids the - second hop a receive-then-broadcast would add. - - Data parallelism additionally requires the caller to route both halves of a - request to the same replica, because a push has to land where the request - will run. The proxy is the only component that knows which replica it picked: - it names the replica to the consumer (``X-data-parallel-rank``) and passes - that replica's control address to the producer. Getting this wrong is loud - rather than silent -- the replica that runs the request never sees its - embedding and gives up after ``push_wait_timeout_s`` -- but it is the - caller's responsibility, not something this connector can detect. + Attributes: + _scheduler: Scheduler implementation when constructed for that role. + _worker: Worker implementation when constructed for that role. + _closed: Whether role-specific resources have already been released. """ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): super().__init__(vllm_config=vllm_config, role=role) - if _MOONCAKE_IMPORT_ERROR is not None or TransferEngine is None: - raise ImportError( - "Install mooncake-transfer-engine (see " - "https://github.com/kvcache-ai/Mooncake ) to use ECMooncakeConnector." - ) from _MOONCAKE_IMPORT_ERROR - parallel_config = vllm_config.parallel_config - ec_cfg_early = vllm_config.ec_transfer_config - assert ec_cfg_early is not None - if ec_cfg_early.is_ec_producer: - # The producer holds one copy of each encoder output and addresses - # consumers directly; sharding or replicating it would only - # duplicate the push. - if parallel_config.tensor_parallel_size > 1: - raise ValueError( - "ECMooncakeConnector producers require tensor_parallel_size=1." - ) - if parallel_config.pipeline_parallel_size > 1: - raise ValueError( - "ECMooncakeConnector producers do not support pipeline parallelism." - ) - if parallel_config.data_parallel_size > 1: - raise ValueError( - "ECMooncakeConnector producers require data_parallel_size=1." - ) + self._scheduler: ECMooncakeScheduler | None = None + self._worker: ECMooncakeWorker | None = None + self._closed = False - # Each data-parallel replica runs its own scheduler and its own control - # channels, so their ports must not overlap. `data_parallel_index` is the - # only field that identifies the replica in both cases: a non-MoE replica - # is reconfigured to look like DP=1, which resets `data_parallel_rank` - # and `data_parallel_size`. Deriving the offset from the config rather - # than from a process group keeps the scheduler, which has no groups, in - # agreement with its workers. - self._control_port_offset = ( - parallel_config.data_parallel_index * parallel_config.tensor_parallel_size - ) - - self._role = role - ec_cfg = vllm_config.ec_transfer_config - assert ec_cfg is not None - self._ec_cfg = ec_cfg - self._extra = self._ec_cfg.ec_connector_extra_config - self._protocol: str = self._extra.get("mooncake_protocol", "rdma") - reservation_port = self._extra.get("reservation_zmq_port") - self._reservation_zmq_port = ( - int(reservation_port) if reservation_port is not None else None - ) - self._reservation_zmq_addr: str | None = self._extra.get("reservation_zmq_addr") - if ( - self._reservation_zmq_addr is None - and self._reservation_zmq_port is not None - ): - base = self._reservation_zmq_port + self._control_port_offset - self._reservation_zmq_addr = f"tcp://127.0.0.1:{base}" - self._registered_capacity = int(self._ec_cfg.ec_buffer_size) - if self._registered_capacity <= 0: - raise ValueError("ECMooncakeConnector requires ec_buffer_size > 0.") - self._model_config = vllm_config.model_config - self._metadata_fields_cache: dict[str, set[str]] = {} - - pool_size = self._extra.get( - "consumer_buffer_pool_size", self._registered_capacity - ) - self._consumer_pool_capacity = int(pool_size) - self._consumer_pool: torch.Tensor | None = None - self._consumer_pool_allocator: _ContiguousAllocator | None = None - # The receive pool is orders of magnitude larger than the encoder - # cache, so an item the encoder cache evicted stays resident here and - # a later request gets it for a dict lookup instead of a transfer. - self._consumer_residents: _ResidentPool[_ConsumerPoolAllocation] = ( - _ResidentPool(self._consumer_pool_capacity) - ) - self._consumer_retire_events: dict[str, torch.Event] = {} - self._consumer_pending_frees: list[ - tuple[torch.Event, _ConsumerPoolAllocation] - ] = [] - self._consumer_reclaimed: set[str] = set() - self._consumer_rank_resolved = False - self._is_receiving_rank = True - self._tp_rank = 0 - self._tp_size = 1 - self._consumer_pool_disabled = self._consumer_pool_capacity <= 0 - self._consumer_lock = threading.Lock() - self._push_reservations: dict[str, _PushReservation] = {} - self._cancelled_transfers: OrderedDict[str, float] = OrderedDict() - self._control_server: ECMooncakeControlServer | None = None - self._consumer_metrics_log_interval = float( - self._extra.get("consumer_metrics_log_interval", 10) - ) - self._consumer_metrics_started_at = time.monotonic() - self._consumer_worker_metrics: Counter[str] = Counter() - self._consumer_scheduler_metrics: Counter[str] = Counter() - self._consumer_missing_since: dict[str, float] = {} - self._stalled_hashes: set[str] = set() - self._unavailable_requests: set[str] = set() - self._active_push_sources: Counter[tuple[str, int]] = Counter() - self._active_push_sources_lock = threading.Lock() - self._push_wait_timeout = float(self._extra.get("push_wait_timeout_s", 60)) - self._drain_pending = True - self._drained_at = 0.0 - self._consumer_loading_since: dict[str, float] = {} - self._consumer_pending_since: dict[str, float] = {} - self._pending_spec_deadlines: dict[str, float] = {} - self._pending_cancels: dict[str, Future[Any]] = {} - # Cancelled transfers, oldest deadline first, so the sweep can stop - # at the first live entry. - self._cancelled_transfer_ids: OrderedDict[str, float] = OrderedDict() - - # Scheduler (consumer): transfer_id -> pending tensor layout. - self._pending_specs: dict[str, ECMooncakeLoadSpec] = {} - self._pending_specs_by_hash: dict[str, deque[str]] = {} - self._load_specs: dict[str, ECMooncakeLoadSpec] = {} - self._mm_datas_need_loads: dict[str, int] = {} - self._loading_hashes: set[str] = set() - self._ready_hashes: set[str] = set() - # Scheduler-side mirror of the worker's receive pool, oldest first. An - # item stays here after the encoder cache evicts it, so the next - # request that needs it is served locally instead of consuming another - # transfer. The worker reports what it reclaims under pressure; the - # byte budget only guards against drift. - self._resident_specs: OrderedDict[str, ECMooncakeLoadSpec] = OrderedDict() - self._resident_bytes = 0 - self._scheduler_pending_work = False - self._pushes_to_prepare: dict[str, ECMooncakePushSpec] = {} - # A producer request may be revisited across scheduler steps. Queue its - # initial push metadata once; the worker owns reservation refreshes - # after the scheduler emits it. - self._prepared_push_transfer_ids: set[str] = set() - - # Worker producer - self._engine: TransferEngine | None = None - self._engine_lock = threading.Lock() - self._hostname = get_ip() - # Published encoder outputs, referenced while a pull is reading them. - self._pending_unregister: dict[int, torch.Tensor] = {} - self._push_source_registrations: dict[int, _PushSourceRegistration] = {} - self._push_source_registration_lock = threading.Lock() - producer_pool = self._extra.get( - "producer_buffer_pool_size", self._registered_capacity - ) - self._producer_pool_capacity = int(producer_pool) - self._producer_pool: torch.Tensor | None = None - self._producer_pool_allocator: _ContiguousAllocator | None = None - self._producer_pool_disabled = self._producer_pool_capacity <= 0 - self._producer_pool_lock = threading.Lock() - transfer_workers = int(self._extra.get("transfer_max_workers", 4)) - control_workers = int(self._extra.get("control_max_workers", 8)) - self._transfer_metrics_log_interval = float( - self._extra.get("transfer_metrics_log_interval", 10) - ) - self._control_channel = _ControlChannel( - int(float(self._extra.get("control_timeout_s", 30)) * 1000) - ) - self._producer_metrics: Counter[str] = Counter() - self._io_executor = ThreadPoolExecutor( - max_workers=transfer_workers, thread_name_prefix="ec-mooncake-transfer" - ) - self._control_executor = ThreadPoolExecutor( - max_workers=control_workers, thread_name_prefix="ec-mooncake-control" - ) - self._consumer_shard_cache: dict[str, list[str]] = {} - self._shard_pool: ThreadPoolExecutor | None = None - self._shard_pool_lock = threading.Lock() - self._pending_saves: list[tuple[str, Future[None]]] = [] - self._pending_reservations: dict[ - str, deque[tuple[ECMooncakePushSpec, Future[list[dict[str, Any]]]]] - ] = {} - self._pending_pushes: list[_PendingPush] = [] - self._push_perf_lock = threading.Lock() - self._push_perf = _PushPerfWindow() - self._active_transfer_batches = 0 - self._queued_transfer_batches = 0 - self._event_zmq_ctx: zmq.Context | None = None - self._event_zmq_socket: zmq.Socket | None = None - self._event_shard_count = 1 - # transfer_id -> shards that reported it ready, oldest first. A sharded - # consumer writes one copy per rank, so the item is only loadable once - # every rank has reported. Bounded: a transfer whose last rank never - # arrives is given up on by the push-wait timeout, not by this map. - self._event_ready_shards: OrderedDict[str, set[int]] = OrderedDict() - self._completed_loads: set[str] = set() - self._failed_loads: set[str] = set() - self._shutdown = False - - if ( - role == ECConnectorRole.SCHEDULER - and self.is_consumer - and not self._reservation_zmq_addr - ): - raise ValueError( - "ec_consumer with ECMooncakeConnector requires " - "reservation_zmq_port or reservation_zmq_addr." - ) - - def _ensure_engine(self) -> TransferEngine: - if self._engine is not None: - return self._engine - with self._engine_lock: - if self._engine is not None: - return self._engine - eng = TransferEngine() - ret = eng.initialize(self._hostname, "P2PHANDSHAKE", self._protocol, "") - if ret != 0: - raise RuntimeError("Mooncake TransferEngine initialization failed.") - self._engine = eng - logger.info( - "ECMooncakeConnector TransferEngine ready at %s:%d", - self._hostname, - eng.get_rpc_port(), - ) - return self._engine - - def _resolve_consumer_rank(self) -> None: - """Place this worker in the consumer's receive topology. - - Encoder outputs only exist on the first pipeline stage, and every - tensor-parallel rank there gathers from its own cache, so each of them - receives its own copy on its own control channel. Ports run - consecutively from the configured one so a producer holding the first - address can reach the rest. - """ - if self._consumer_rank_resolved: - return - self._consumer_rank_resolved = True - try: - from vllm.distributed.parallel_state import get_pp_group, get_tp_group - - tp_group = get_tp_group() - self._tp_rank = tp_group.rank_in_group - self._tp_size = tp_group.world_size - self._is_receiving_rank = get_pp_group().is_first_rank - except AssertionError: - # Groups are only absent outside a distributed run, where this - # worker is the whole consumer. - self._tp_rank = 0 - self._tp_size = 1 - self._is_receiving_rank = True + if role == ECConnectorRole.SCHEDULER: + self._scheduler = ECMooncakeScheduler.from_vllm_config(vllm_config) + elif role == ECConnectorRole.WORKER: + self._worker = ECMooncakeWorker.from_vllm_config(vllm_config) + else: + raise ValueError(f"Unknown EC connector role: {role}") def start_worker_services(self) -> None: - if ( - self._role != ECConnectorRole.WORKER - or not self.is_consumer - or self._reservation_zmq_port is None - or self._control_server is not None - ): - return - self._resolve_consumer_rank() - if not self._is_receiving_rank: - # Later pipeline stages hold no encoder outputs, so they need - # neither a receive pool nor a control channel. - return - raw_device = self._ec_cfg.ec_buffer_device - device_name = ( - raw_device.lower() if isinstance(raw_device, str) and raw_device else "cuda" - ) - self._ensure_consumer_pool(torch.device(device_name), allow_host=True) - if self._consumer_pool is None: - raise RuntimeError( - "Mooncake push mode requires a registered consumer buffer pool." - ) - base_port = self._reservation_zmq_port + self._control_port_offset - self._control_server = ECMooncakeControlServer( - "0.0.0.0", - base_port + self._tp_rank, - self._reserve_push_destination, - self._push_status, - self._complete_push, - self._cancel_push, - self._expire_push_reservations, - self._consumer_metrics_log_interval, - peer_ports=[base_port + rank for rank in range(self._tp_size)], - device=self._consumer_pool.device, - ) - self._control_server.start() - - def _unregister_memory(self, tensor: torch.Tensor) -> bool: - assert self._engine is not None - ret = self._engine.unregister_memory(tensor.data_ptr()) - if ret != 0: - logger.error( - "Mooncake EC memory unregistration failed for address %d: %d", - tensor.data_ptr(), - ret, - ) - self._pending_unregister[tensor.data_ptr()] = tensor - return False - self._pending_unregister.pop(tensor.data_ptr(), None) - return True - - def _unregister_memories(self, tensors: list[torch.Tensor]) -> None: - assert self._engine is not None - addresses = [tensor.data_ptr() for tensor in tensors] - ret = self._engine.batch_unregister_memory(addresses) - if ret != 0: - for tensor in tensors: - self._pending_unregister[tensor.data_ptr()] = tensor - logger.warning( - "Keeping %d EC tensors alive after Mooncake unregistration failure", - len(tensors), - ) - return - for address in addresses: - self._pending_unregister.pop(address, None) - - @staticmethod - def _push_source_range(tensor: torch.Tensor) -> tuple[int, int]: - # Register exactly the bytes that will be transferred. One encoder - # batch returns its items as views of a single storage (models split - # the batched embeddings, e.g. `image_embeds.split(sizes)`), so - # registering the whole storage would overlap the per-tensor - # registration a sibling item takes -- and Mooncake rejects - # overlapping memory regions. - return tensor.data_ptr(), tensor.nbytes - - def _acquire_push_source_registrations( - self, tensors: list[torch.Tensor] - ) -> list[int]: - ranges: dict[int, tuple[int, torch.Tensor]] = {} - for tensor in tensors: - address, nbytes = self._push_source_range(tensor) - ranges.setdefault(address, (nbytes, tensor)) - - eng = self._ensure_engine() - acquired: list[int] = [] - new_addresses: list[int] = [] - new_lengths: list[int] = [] - with self._push_source_registration_lock: - for address, (nbytes, tensor) in ranges.items(): - entry = self._push_source_registrations.get(address) - if entry is not None: - if entry.nbytes != nbytes: - raise RuntimeError( - "Mooncake EC source storage changed size while registered" - ) - entry.users += 1 - acquired.append(address) - continue - new_addresses.append(address) - new_lengths.append(nbytes) - self._push_source_registrations[address] = _PushSourceRegistration( - tensor=tensor, - nbytes=nbytes, - ) - acquired.append(address) - - if new_addresses: - ret = eng.batch_register_memory(new_addresses, new_lengths) - if ret != 0: - for address in acquired: - entry = self._push_source_registrations[address] - entry.users -= 1 - if entry.users == 0: - del self._push_source_registrations[address] - raise RuntimeError("Mooncake EC source registration failed") - return acquired - - def _release_push_source_registrations(self, addresses: list[int]) -> bool: - if not addresses: - return True - with self._push_source_registration_lock: - unused = [] - for address in addresses: - entry = self._push_source_registrations.get(address) - if entry is None: - continue - entry.users -= 1 - if entry.users == 0: - unused.append(address) - if not unused: - return True - ret = self._ensure_engine().batch_unregister_memory(unused) - if ret != 0: - logger.warning( - "Keeping %d EC source tensors registered after Mooncake " - "unregistration failure", - len(unused), - ) - return False - for address in unused: - del self._push_source_registrations[address] - self._pending_unregister.pop(address, None) - return True - - def _ensure_consumer_pool( - self, device: torch.device, *, allow_host: bool = False - ) -> None: - if ( - self._consumer_pool is not None - or self._consumer_pool_disabled - or (device.type != "cuda" and not allow_host) - ): - return - try: - pool = torch.empty( - self._consumer_pool_capacity, dtype=torch.uint8, device=device - ) - if self._is_receiving_rank: - # Producers write into this pool directly, so it needs a memory - # region. Later pipeline stages never receive and skip it. - ret = self._ensure_engine().batch_register_memory( - [pool.data_ptr()], [pool.nbytes] - ) - if ret != 0: - raise RuntimeError(f"Mooncake returned {ret}") - except (RuntimeError, torch.OutOfMemoryError) as e: - self._consumer_pool_disabled = True - logger.warning( - "Could not initialize the EC consumer buffer pool; falling back " - "to per-tensor registration: %s", - e, - ) - return - self._consumer_pool = pool - self._consumer_pool_allocator = _ContiguousAllocator(pool.nbytes) - logger.info( - "Prepared %d-byte CUDA receive pool for Mooncake EC (registered=%s)", - pool.nbytes, - self._is_receiving_rank, - ) - - def _ensure_producer_pool(self, device: torch.device) -> None: - """Register one staging slab so pushes never register per transfer. - - Registering the encoder output itself costs more than the transfer - (register+unregister dominated the push path); staging into a slab - that is registered once trades that for a device-to-device copy. - """ - if self._producer_pool is not None or self._producer_pool_disabled: - return - with self._producer_pool_lock: - if self._producer_pool is not None or self._producer_pool_disabled: - return - try: - pool = torch.empty( - self._producer_pool_capacity, dtype=torch.uint8, device=device - ) - ret = self._ensure_engine().batch_register_memory( - [pool.data_ptr()], [pool.nbytes] - ) - if ret != 0: - raise RuntimeError(f"Mooncake returned {ret}") - except (RuntimeError, torch.OutOfMemoryError) as e: - self._producer_pool_disabled = True - logger.warning( - "Could not initialize the EC producer staging pool; falling " - "back to per-transfer registration: %s", - e, - ) - return - self._producer_pool = pool - self._producer_pool_allocator = _ContiguousAllocator(pool.nbytes) - logger.info( - "Registered %d-byte staging pool for Mooncake EC pushes", - pool.nbytes, - ) - - def _stage_push_sources( - self, tensors: list[torch.Tensor] - ) -> tuple[list[torch.Tensor], list[tuple[int, int]]] | None: - """Copy the batch into the staging pool; None if it does not fit.""" - if not tensors: - return [], [] - self._ensure_producer_pool(tensors[0].device) - pool = self._producer_pool - allocator = self._producer_pool_allocator - if pool is None or allocator is None: - return None - staged: list[torch.Tensor] = [] - regions: list[tuple[int, int]] = [] - with self._producer_pool_lock: - for tensor in tensors: - region = allocator.allocate(tensor.nbytes) - if region is None: - for offset, size in regions: - allocator.free(offset, size) - return None - regions.append(region) - offset = region[0] - staged.append( - pool.narrow(0, offset, tensor.nbytes) - .view(tensor.dtype) - .view(tensor.shape) - ) - for destination, source in zip(staged, tensors): - destination.copy_(source, non_blocking=True) - return staged, regions - - def _release_push_staging(self, regions: list[tuple[int, int]]) -> None: - allocator = self._producer_pool_allocator - if allocator is None or not regions: - return - with self._producer_pool_lock: - for offset, size in regions: - allocator.free(offset, size) - - def _poll_consumer_pool_frees(self) -> None: - allocator = self._consumer_pool_allocator - if allocator is None: - return - with self._consumer_lock: - pending = [] - for event, allocation in self._consumer_pending_frees: - if event.query(): - allocator.free(allocation.offset, allocation.size) - else: - pending.append((event, allocation)) - self._consumer_pending_frees = pending - - def _reclaim_residents_locked( - self, allocator: _ContiguousAllocator, nbytes: int - ) -> tuple[int, int] | None: - """Give up retired items, oldest first, until `nbytes` fits. - - Called only when the pool cannot satisfy an allocation, so a retired - item survives until its memory is genuinely needed. - """ - - def evict(mm_hash: str, allocation: _ConsumerPoolAllocation) -> bool: - event = self._consumer_retire_events.pop(mm_hash, None) - if event is None or event.query(): - allocator.free(allocation.offset, allocation.size) - else: - self._consumer_pending_frees.append((event, allocation)) - self._consumer_reclaimed.add(mm_hash) - self._consumer_worker_metrics["residents_reclaimed"] += 1 - return True - - while self._consumer_residents.evict_lru(evict) is not None: - region = allocator.allocate(nbytes) - if region is not None: - return region - return None - - def _take_resident_tensor(self, spec: ECMooncakeLoadSpec) -> torch.Tensor | None: - """Hand back a copy the pool still holds. - - Retired and in-use entries live in the same map, so an item a later - push reserved again still serves this load. - """ - with self._consumer_lock: - allocation = self._consumer_residents.get(spec.mm_hash) - if allocation is None: - self._consumer_worker_metrics["residents_missed"] += 1 - return None - tensor = allocation.tensor - if ( - tuple(tensor.shape) != tuple(spec.shape) - or str(tensor.dtype).split(".")[-1] != spec.dtype - ): - self._consumer_worker_metrics["residents_mismatched"] += 1 - return None - self._consumer_residents.pin(spec.mm_hash) - self._consumer_retire_events.pop(spec.mm_hash, None) - self._consumer_worker_metrics["residents_promoted"] += 1 - return tensor - - def _release_stale_consumer_allocations( - self, encoder_cache: dict[str, torch.Tensor] - ) -> None: - if self._consumer_pool is None: - return - with self._consumer_lock: - reserved_allocations = { - id(reservation.allocation) - for reservation in self._push_reservations.values() - } - # Walk only the referenced entries: the retired set grows to - # thousands and none of it can change state here. - for mm_hash in self._consumer_residents.referenced(): - allocation = self._consumer_residents.get(mm_hash) - if allocation is None: - continue - if encoder_cache.get(mm_hash) is allocation.tensor: - continue - if id(allocation) in reserved_allocations: - continue - # Retire rather than free: the bytes stay valid and serve the - # next request that needs this item. The event orders the - # eventual reuse behind whatever still reads the tensor. - event = torch.Event() - event.record( - torch.accelerator.current_stream(self._consumer_pool.device) - ) - self._consumer_retire_events[mm_hash] = event - self._consumer_residents.retire(mm_hash) - self._consumer_worker_metrics["residents_retired"] += 1 - self._poll_consumer_pool_frees() - - def _clear_item_timers(self, mm_hash: str) -> None: - self._consumer_missing_since.pop(mm_hash, None) - self._consumer_loading_since.pop(mm_hash, None) - self._consumer_pending_since.pop(mm_hash, None) - self._stalled_hashes.discard(mm_hash) - - def _note_awaiting_push( - self, - mm_hash: str, - transfer_id: str | None = None, - request_id: str | None = None, - ) -> bool: - """Wait for an item with nothing in flight, and give up on timeout. - - Nothing on this side can produce the item, so a push that never - arrives would defer the request forever. Past the timeout the request - is reported unavailable instead: the scheduler fails it with a - retryable error and the caller can re-issue it, which re-runs the - encode and produces a fresh transfer. - - Returns: - True once this request has been given up on. - """ - now = time.monotonic() - since = self._consumer_missing_since.setdefault(mm_hash, now) - self._consumer_scheduler_metrics["missing_event"] += 1 - elapsed = now - since - if elapsed < self._push_wait_timeout: - return False - stale = mm_hash in self._stalled_hashes - if request_id is not None: - self._unavailable_requests.add(request_id) - self._consumer_scheduler_metrics["given_up"] += 1 - # Start a fresh window: a re-issued request pushes this item again, - # and it must be allowed to wait for that push rather than inherit - # this one's deadline and be given up on immediately. Only the - # deadline resets -- `_stalled_hashes` keeps the warning to one per - # hash, while `given_up` counts every occurrence. - self._consumer_missing_since.pop(mm_hash, None) - if stale: - return request_id is not None - self._stalled_hashes.add(mm_hash) - self._consumer_scheduler_metrics["stalled"] += 1 - # Ask the worker what it knows about this transfer: whether the - # reservation exists at all separates "the producer never sent it" - # from "it arrived and the scheduler missed it". - reservation: Any = "unknown" - if transfer_id and self._reservation_zmq_addr is not None: - try: - reservation = self._send_control( - self._reservation_zmq_addr, - {"op": "status", "transfer_id": transfer_id}, - ) - except Exception as e: # noqa: BLE001 - diagnostic only - reservation = f"status failed: {e}" - logger.warning( - "EC Mooncake waited %.1fs for a push of mm_hash=%s " - "(transfer_id=%s) that never arrived; worker reservation=%s; " - "requests needing it fail with a retryable error.", - elapsed, - mm_hash, - transfer_id, - reservation, - ) - return request_id is not None - - def take_unavailable_requests(self) -> set[str]: - given_up = self._unavailable_requests - self._unavailable_requests = set() - return given_up - - @staticmethod - def _hash_samples(values: list[str], limit: int = 5) -> list[str]: - return [value[:16] for value in values[:limit]] - - def _maybe_log_consumer_worker_metrics(self) -> None: - now = time.monotonic() - if ( - self._consumer_metrics_log_interval <= 0 - or now - self._consumer_metrics_started_at - < self._consumer_metrics_log_interval - ): - return - with self._consumer_lock: - ready = [ - mm_hash - for mm_hash, reservation in self._push_reservations.items() - if reservation.ready - ] - pending = [ - mm_hash - for mm_hash, reservation in self._push_reservations.items() - if not reservation.ready - ] - metrics = dict(self._consumer_worker_metrics) - self._consumer_worker_metrics.clear() - residents = len(self._consumer_residents) - live = len(self._consumer_residents.referenced()) - retired = self._consumer_residents.num_evictable - pending_frees = len(self._consumer_pending_frees) - oldest_reservation_ms = max( - ( - (now - reservation.created_at) * 1000 - for reservation in self._push_reservations.values() - ), - default=0.0, - ) - logger.info( - "EC Mooncake consumer worker: lifecycle=%s, reservations_ready=%d, " - "reservations_pending=%d, residents=%d, live=%d, retired=%d, " - "pending_frees=%d, " - "oldest_reservation_ms=%.1f, ready_hashes=%s, pending_hashes=%s", - metrics, - len(ready), - len(pending), - residents, - live, - retired, - pending_frees, - oldest_reservation_ms, - self._hash_samples(ready), - self._hash_samples(pending), - ) - self._consumer_metrics_started_at = now - - def _maybe_log_consumer_scheduler_metrics(self) -> None: - now = time.monotonic() - if ( - self._consumer_metrics_log_interval <= 0 - or now - self._consumer_metrics_started_at - < self._consumer_metrics_log_interval - ): - return - missing = sorted(self._consumer_missing_since.items(), key=lambda item: item[1]) - loading = sorted(self._consumer_loading_since.items(), key=lambda item: item[1]) - pending = sorted(self._consumer_pending_since.items(), key=lambda item: item[1]) - oldest_missing_ms = round((now - missing[0][1]) * 1000, 1) if missing else 0.0 - oldest_loading_ms = round((now - loading[0][1]) * 1000, 1) if loading else 0.0 - oldest_pending_ms = round((now - pending[0][1]) * 1000, 1) if pending else 0.0 - logger.info( - "EC Mooncake consumer scheduler: decisions=%s, ready=%d, loading=%d, " - "resident=%d, pending_specs=%d, needs_load=%d, missing=%d, " - "oldest_missing_ms=%.1f, oldest_loading_ms=%.1f, " - "oldest_pending_ms=%.1f, missing_hashes=%s, loading_hashes=%s, " - "pending_hashes=%s", - dict(self._consumer_scheduler_metrics), - len(self._ready_hashes), - len(self._loading_hashes), - len(self._resident_specs), - len(self._pending_specs), - len(self._mm_datas_need_loads), - len(missing), - oldest_missing_ms, - oldest_loading_ms, - oldest_pending_ms, - self._hash_samples([mm_hash for mm_hash, _ in missing]), - self._hash_samples([mm_hash for mm_hash, _ in loading]), - self._hash_samples([mm_hash for mm_hash, _ in pending]), - ) - self._consumer_scheduler_metrics.clear() - self._consumer_metrics_started_at = now - - @staticmethod - def _expire_cancel_records(records: OrderedDict[str, float], now: float) -> int: - """Drop the cancels that can no longer be told apart from unknown ids. - - Both roles keep one record per multimodal item they handle, and both - consult it on a per-item hot path, so a full rescan costs the square - of the item rate: at 53 items/s the worker's 300 s window is 16k - entries and its sweep ran under `_consumer_lock` on every - reservation. Callers append in deadline order -- `move_to_end` when - refreshing one -- so the front is always the oldest and the sweep - stops at the first live entry. - - Returns: - How many records were dropped. - """ - dropped = 0 - while records: - expires_at = next(iter(records.values())) - if expires_at > now and len(records) <= _MAX_CANCELLED_TRANSFER_IDS: - break - records.popitem(last=False) - dropped += 1 - return dropped - - def _expire_push_reservations_locked(self) -> None: - now = time.monotonic() - allocator = self._consumer_pool_allocator - assert allocator is not None - for transfer_id, reservation in list(self._push_reservations.items()): - if reservation.expires_at > now: - continue - if reservation.owns_allocation: - allocator.free( - reservation.allocation.offset, reservation.allocation.size - ) - self._push_reservations.pop(transfer_id) - self._consumer_worker_metrics["reservations_expired"] += 1 - self._consumer_worker_metrics["cancel_records_dropped"] += ( - self._expire_cancel_records(self._cancelled_transfers, now) - ) - - def _expire_push_reservations(self) -> int: - with self._consumer_lock: - before = len(self._push_reservations) - self._expire_push_reservations_locked() - return before - len(self._push_reservations) - - def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: - transfer_id = str(payload["transfer_id"]) - mm_hash = str(payload["mm_hash"]) - nbytes = int(payload["nbytes"]) - shape = tuple(int(value) for value in payload["shape"]) - dtype_name = str(payload["dtype"]) - dtype = getattr(torch, dtype_name, None) - if dtype is None: - raise ValueError(f"Unsupported torch dtype string: {dtype_name!r}") - expected_nbytes = math.prod(shape) * dtype.itemsize - if expected_nbytes != nbytes: - raise ValueError("shape and dtype do not match nbytes") - - with self._consumer_lock: - self._expire_push_reservations_locked() - if transfer_id in self._cancelled_transfers: - self._consumer_worker_metrics["reservations_cancelled_early"] += 1 - return { - "reservation_id": "", - "dst_session": "", - "dst_ptr": 0, - "nbytes": nbytes, - "write": False, - "ready": False, - "cancelled": True, - } - existing = self._push_reservations.get(transfer_id) - if existing is not None: - if ( - existing.mm_hash != mm_hash - or existing.shape != shape - or existing.dtype != dtype_name - ): - raise ValueError("conflicting reservation for transfer_id") - reservation = existing - should_write = False - key = ( - "reservations_reused_ready" - if existing.ready - else ("reservations_reused_pending") - ) - self._consumer_worker_metrics[key] += 1 - if not existing.ready: - existing.expires_at = time.monotonic() + _LEASE_TTL_SECONDS - else: - cached = self._consumer_residents.get(mm_hash) - if cached is not None: - if ( - tuple(cached.tensor.shape) != shape - or cached.tensor.dtype != dtype - ): - raise ValueError("conflicting cached tensor for mm_hash") - reservation = _PushReservation( - mm_hash=mm_hash, - reservation_id=uuid.uuid4().hex, - allocation=cached, - shape=shape, - dtype=dtype_name, - ready=True, - owns_allocation=False, - expires_at=time.monotonic() + _LEASE_TTL_SECONDS, - ) - should_write = False - # Live again: it must not be reclaimed under pressure. - self._consumer_residents.pin(mm_hash) - self._consumer_retire_events.pop(mm_hash, None) - self._consumer_worker_metrics["reservations_cached"] += 1 - else: - pool = self._consumer_pool - allocator = self._consumer_pool_allocator - assert pool is not None and allocator is not None - region = allocator.allocate(nbytes) - if region is None: - self._expire_push_reservations_locked() - region = allocator.allocate(nbytes) - if region is None: - region = self._reclaim_residents_locked(allocator, nbytes) - if region is None: - raise RuntimeError("EC consumer buffer pool is full") - offset, size = region - tensor = pool.narrow(0, offset, nbytes).view(dtype).view(shape) - allocation = _ConsumerPoolAllocation(offset, size, tensor) - reservation = _PushReservation( - mm_hash=mm_hash, - reservation_id=uuid.uuid4().hex, - allocation=allocation, - shape=shape, - dtype=dtype_name, - expires_at=time.monotonic() + _LEASE_TTL_SECONDS, - ) - should_write = True - self._consumer_worker_metrics["reservations_created"] += 1 - self._push_reservations[transfer_id] = reservation - - eng = self._ensure_engine() - return { - "reservation_id": reservation.reservation_id, - "dst_session": f"{self._hostname}:{eng.get_rpc_port()}", - "dst_ptr": reservation.allocation.tensor.data_ptr(), - "nbytes": reservation.allocation.tensor.nbytes, - "write": should_write, - "ready": reservation.ready, - "cached": not reservation.owns_allocation, - } - - def _push_status(self, transfer_id: str) -> dict[str, Any] | None: - with self._consumer_lock: - reservation = self._push_reservations.get(transfer_id) - if reservation is None: - return None - return { - "mm_hash": reservation.mm_hash, - "ready": reservation.ready, - "reservation_id": reservation.reservation_id, - "nbytes": reservation.allocation.tensor.nbytes, - "shape": list(reservation.shape), - "dtype": reservation.dtype, - } - - def _complete_push(self, transfer_id: str, reservation_id: str) -> _PushCompletion: - with self._consumer_lock: - reservation = self._push_reservations.get(transfer_id) - if reservation is None or reservation.reservation_id != reservation_id: - self._consumer_worker_metrics["completions_rejected"] += 1 - return _PushCompletion(False) - if reservation.ready: - self._consumer_worker_metrics["completions_repeated"] += 1 - return _PushCompletion(True) - self._consumer_worker_metrics["completions_accepted"] += 1 - if reservation.discard_on_complete: - allocator = self._consumer_pool_allocator - assert allocator is not None - self._push_reservations.pop(transfer_id) - if reservation.owns_allocation: - allocator.free( - reservation.allocation.offset, reservation.allocation.size - ) - self._consumer_worker_metrics["reservations_discarded"] += 1 - return _PushCompletion(True) - reservation.ready = True - reservation.expires_at = time.monotonic() + _LEASE_TTL_SECONDS - return _PushCompletion(True, became_ready=True) - - def _cancel_push( - self, transfer_id: str, reservation_id: str, abandon: bool = False - ) -> bool: - with self._consumer_lock: - reservation = self._push_reservations.get(transfer_id) - if ( - reservation is not None - and reservation_id - and reservation.reservation_id != reservation_id - ): - self._consumer_worker_metrics["cancellations_rejected"] += 1 - return False - self._cancelled_transfers[transfer_id] = ( - time.monotonic() + _LEASE_TTL_SECONDS - ) - self._cancelled_transfers.move_to_end(transfer_id) - if reservation is None: - self._consumer_worker_metrics["cancellations_pre_reserved"] += 1 - return True - allocator = self._consumer_pool_allocator - assert allocator is not None - if not reservation.ready and not abandon: - reservation.discard_on_complete = True - self._consumer_worker_metrics["cancellations_deferred"] += 1 - return True - self._push_reservations.pop(transfer_id) - if reservation.owns_allocation: - allocator.free( - reservation.allocation.offset, reservation.allocation.size - ) - self._consumer_worker_metrics["reservations_cancelled"] += 1 - return True - - def _take_pushed_tensor( - self, spec: ECMooncakeLoadSpec - ) -> tuple[torch.Tensor, _ConsumerPoolAllocation]: - with self._consumer_lock: - reservation = self._push_reservations.get(spec.transfer_id) - # Not compared against `spec.reservation_id`: each shard mints its - # own, while the spec carries the one from whichever shard's event - # the scheduler observed. `transfer_id` is assigned per request - # item and is already unique, and a stale reservation for a reused - # one is rejected by `_reserve_push_destination`. - if reservation is None or not reservation.ready: - self._consumer_worker_metrics["takes_rejected"] += 1 - raise RuntimeError( - f"Pushed EC tensor is not ready for mm_hash={spec.mm_hash}" - ) - self._push_reservations.pop(spec.transfer_id) - self._consumer_residents.insert( - spec.mm_hash, reservation.allocation, reservation.allocation.size - ) - self._consumer_worker_metrics["reservations_taken"] += 1 - return reservation.allocation.tensor, reservation.allocation - - def _send_control(self, addr: str, request: dict[str, Any]) -> Any: - return self._control_channel.request(addr, request) - - def _shard_executor(self) -> ThreadPoolExecutor: - """Threads for the extra shards of a sharded consumer. - - Reserving and writing both fan out from a task that already holds a - worker of the control or transfer pool, so the extra shards need a - pool of their own: queueing them behind their own caller deadlocks as - soon as every worker there is waiting. Nothing submitted here fans out - again, so this pool cannot deadlock on itself. - """ - with self._shard_pool_lock: - if self._shard_pool is None: - self._shard_pool = ThreadPoolExecutor( - max_workers=32, thread_name_prefix="ec-mooncake-shard" - ) - return self._shard_pool - - def _consumer_shards(self, base_addr: str) -> list[str]: - """Every control channel of the consumer reachable at `base_addr`. - - A tensor-parallel consumer gathers from each rank's own cache, so each - rank receives its own copy. Asking the first one for the roster keeps - the address list out of the request and the proxy configuration. - """ - cached = self._consumer_shard_cache.get(base_addr) - if cached is not None: - return cached - shards = [base_addr] - try: - reply = self._send_control(base_addr, {"op": "peers"}) - ports = reply.get("ports") if isinstance(reply, dict) else None - if ports: - prefix = base_addr.rsplit(":", 1)[0] - shards = [f"{prefix}:{int(port)}" for port in ports] - except Exception: - # An older consumer does not answer this, and it can only be - # unsharded, so its single address is the whole roster. - logger.warning( - "EC Mooncake consumer at %s did not report its shards; " - "assuming it is unsharded.", - base_addr, - exc_info=True, - ) - self._consumer_shard_cache[base_addr] = shards - if len(shards) > 1: - logger.info( - "EC Mooncake consumer at %s has %d shards", base_addr, len(shards) - ) - return shards - - def _reserve_one(self, addr: str, spec: ECMooncakePushSpec) -> dict[str, Any]: - result = self._send_control( - addr, - { - "op": "reserve", - "transfer_id": spec.transfer_id, - "mm_hash": spec.mm_hash, - "nbytes": spec.nbytes, - "shape": list(spec.shape), - "dtype": spec.dtype, - }, - ) - if not isinstance(result, dict): - raise RuntimeError("Invalid EC reservation response") - result["_received_at"] = time.monotonic() - result["addr"] = addr - return result - - def _reserve_remote(self, spec: ECMooncakePushSpec) -> list[dict[str, Any]]: - """Reserve a destination on every shard of the consumer.""" - shards = self._consumer_shards(spec.consumer_zmq) - if len(shards) == 1: - return [self._reserve_one(shards[0], spec)] - # This already runs on the control pool, so the extra shards go to the - # fan-out pool: queueing them behind their own caller would deadlock - # once every control worker is holding a reservation. - extra = [ - self._shard_executor().submit(self._reserve_one, addr, spec) - for addr in shards[1:] - ] - return [self._reserve_one(shards[0], spec)] + [f.result() for f in extra] - - def _cancel_remote( - self, consumer_zmq: str, transfer_id: str, reservation_id: str - ) -> bool: - """Release this transfer on every shard that reserved for it. - - A sharded consumer holds one reservation per rank, so cancelling only - the first would leave the rest pinning pool slots until they expire. - """ - cancelled = False - for addr in self._consumer_shards(consumer_zmq): - result = self._send_control( - addr, - { - "op": "cancel", - "transfer_id": transfer_id, - "reservation_id": reservation_id, - }, - ) - cancelled |= isinstance(result, dict) and bool(result.get("cancelled")) - return cancelled - - def _poll_pending_cancels(self) -> None: - pending = {} - for transfer_id, future in self._pending_cancels.items(): - if not future.done(): - pending[transfer_id] = future - continue - try: - cancelled = future.result() - except Exception: - self._cancelled_transfer_ids.pop(transfer_id, None) - self._consumer_scheduler_metrics["cancellations_failed"] += 1 - logger.warning( - "EC Mooncake reservation cancellation failed", exc_info=True - ) - else: - key = "cancellations_completed" if cancelled else "cancellations_stale" - self._consumer_scheduler_metrics[key] += 1 - self._pending_cancels = pending + assert self._worker is not None + self._worker.start_services() def start_save_caches(self, **kwargs: Any) -> None: + assert self._worker is not None metadata = self._get_connector_metadata() assert isinstance(metadata, ECMooncakeConnectorMetadata) - for spec in metadata.pushes: - reservation = self._control_executor.submit(self._reserve_remote, spec) - self._pending_reservations.setdefault(spec.mm_hash, deque()).append( - (spec, reservation) - ) - encoder_cache = kwargs.get("encoder_cache") - if not isinstance(encoder_cache, dict): - return - for mm_hash in dict.fromkeys(spec.mm_hash for spec in metadata.pushes): - tensor = encoder_cache.get(mm_hash) - if tensor is not None: - self._submit_reserved_pushes(tensor, mm_hash) + self._worker.start_save_caches(metadata, **kwargs) def start_load_caches( self, encoder_cache: dict[str, torch.Tensor], **kwargs: Any ) -> None: - self._resolve_consumer_rank() - if not self._is_receiving_rank: - # Reached on steps with no work, from a stage that never gathers - # multimodal embeddings. Taking a transfer here would fail for - # want of a reservation and fail the load for everyone. - return + assert self._worker is not None metadata = self._get_connector_metadata() assert isinstance(metadata, ECMooncakeConnectorMetadata) - self._ensure_engine() - raw_buf = self._ec_cfg.ec_buffer_device - buf = raw_buf.lower() if isinstance(raw_buf, str) and raw_buf else "cuda" - if buf == "cuda" and not torch.accelerator.is_available(): - raise RuntimeError( - "ECMooncakeConnector requires CUDA for ec_buffer_device=cuda" - ) - self._release_stale_consumer_allocations(encoder_cache) - - for spec in metadata.loads: - if spec.mm_hash in encoder_cache: - if spec.pushed: - # The spec's id is one shard's; cancel by transfer. - self._cancel_push(spec.transfer_id, "") - self._completed_loads.add(spec.mm_hash) - continue - if spec.local: - resident = self._take_resident_tensor(spec) - if resident is None: - # Reclaimed before the scheduler heard about it; the load - # falls back to a transfer on a later step. - self._failed_loads.add(spec.mm_hash) - else: - encoder_cache[spec.mm_hash] = resident - self._completed_loads.add(spec.mm_hash) - continue - if spec.pushed: - try: - pushed_tensor, _ = self._take_pushed_tensor(spec) - except RuntimeError as e: - logger.warning("EC Mooncake pushed load failed: %s", e) - self._failed_loads.add(spec.mm_hash) - continue - encoder_cache[spec.mm_hash] = pushed_tensor - self._completed_loads.add(spec.mm_hash) - continue - logger.warning( - "EC Mooncake load for mm_hash=%s has no transfer to take", - spec.mm_hash, - ) - self._failed_loads.add(spec.mm_hash) - - def _push_batch(self, pushes: list[_PendingPush]) -> None: - started_at = time.monotonic() - with self._push_perf_lock: - self._queued_transfer_batches -= 1 - self._active_transfer_batches += 1 - - queue_waits_ms = [ - max(0, started_at - push.enqueued_at) * 1000 for push in pushes - ] - stage_ms = { - "queue": sum(queue_waits_ms), - "reserve": 0.0, - "cuda": 0.0, - "register": 0.0, - "rdma": 0.0, - "unregister": 0.0, - "complete": 0.0, - } - ready: list[tuple[_PendingPush, dict[str, Any]]] = [] - notifications: list[tuple[_PendingPush, dict[str, Any]]] = [] - failed = False - try: - synchronized: set[int] = set() - for push in pushes: - stage_started_at = time.monotonic() - reservations = push.reservation.result() - stale = [ - index - for index, shard in enumerate(reservations) - if not shard.get("ready", False) - and time.monotonic() - float(shard.get("_received_at", started_at)) - >= _RESERVATION_REFRESH_SECONDS - ] - if stale: - reservations = self._reserve_remote(push.spec) - stage_ms["reserve"] += (time.monotonic() - stage_started_at) * 1000 - for shard in reservations: - if shard.get("cached", False) or shard.get("cancelled", False): - continue - if not shard.get("write", True): - continue - if push.ready_event is not None and id(push) not in synchronized: - stage_started_at = time.monotonic() - push.ready_event.synchronize() - stage_ms["cuda"] += (time.monotonic() - stage_started_at) * 1000 - synchronized.add(id(push)) - if int(shard["nbytes"]) != push.tensor.nbytes: - raise RuntimeError( - "Reserved EC size does not match tensor for " - f"mm_hash={push.spec.mm_hash}" - ) - ready.append((push, shard)) - notifications.append((push, shard)) - if not ready and not notifications: - return - - if ready: - eng = self._ensure_engine() - # One source per push: a sharded consumer reads the same bytes - # into each of its ranks, so staging and registration happen - # once however many destinations there are. - unique: list[_PendingPush] = [] - source_index: dict[int, int] = {} - for push, _ in ready: - if id(push) not in source_index: - source_index[id(push)] = len(unique) - unique.append(push) - tensors = [push.tensor for push in unique] - lengths = [tensor.nbytes for tensor in tensors] - stage_started_at = time.monotonic() - staged = self._stage_push_sources(tensors) - registered_sources: list[int] = [] - staged_regions: list[tuple[int, int]] = [] - if staged is not None: - sources, staged_regions = staged - # The NIC reads outside the CUDA stream, so the staging - # copies have to have landed before the transfer starts. - if sources and sources[0].device.type == "cuda": - torch.accelerator.current_stream( - sources[0].device - ).synchronize() - else: - sources = tensors - registered_sources = self._acquire_push_source_registrations( - tensors - ) - addresses = [tensor.data_ptr() for tensor in sources] - stage_ms["register"] = (time.monotonic() - stage_started_at) * 1000 - try: - by_session: dict[str, list[tuple[int, int]]] = {} - for push, shard in ready: - by_session.setdefault(str(shard["dst_session"]), []).append( - (source_index[id(push)], int(shard["dst_ptr"])) - ) - stage_started_at = time.monotonic() - - def write(session: str, items: list[tuple[int, int]]) -> None: - ret = eng.batch_transfer_sync_write( - session, - [addresses[index] for index, _ in items], - [dst for _, dst in items], - [lengths[index] for index, _ in items], - ) - if ret != 0: - raise RuntimeError( - f"Mooncake EC push to {session} failed with " - f"status {ret}" - ) - - sessions = list(by_session.items()) - # Shards are written concurrently: serialising them would - # make the transfer cost the sum of the ranks instead of - # the slowest one. - extra = [ - self._shard_executor().submit(write, session, items) - for session, items in sessions[1:] - ] - try: - write(*sessions[0]) - finally: - for future in extra: - future.result() - stage_ms["rdma"] = (time.monotonic() - stage_started_at) * 1000 - finally: - stage_started_at = time.monotonic() - self._release_push_staging(staged_regions) - self._release_push_source_registrations(registered_sources) - stage_ms["unregister"] = ( - time.monotonic() - stage_started_at - ) * 1000 - - stage_started_at = time.monotonic() - self._notify_completions(notifications) - stage_ms["complete"] = (time.monotonic() - stage_started_at) * 1000 - except Exception: - # A failed batch must not take the engine down with it: the - # consumer is told to drop its reservations and this item falls - # back to whatever the consumer can still do (pull, or a local - # re-encode). Raising here would surface in - # `build_connector_worker_meta` as a fatal EngineCore error. - failed = True - logger.exception( - "EC Mooncake push batch failed for mm_hashes=%s", - [push.spec.mm_hash for push in pushes], - ) - self._abandon_pushes(pushes) - finally: - with self._active_push_sources_lock: - for push in pushes: - key = (push.spec.mm_hash, id(push.tensor)) - self._active_push_sources[key] -= 1 - if self._active_push_sources[key] == 0: - del self._active_push_sources[key] - stage_ms["total"] = (time.monotonic() - started_at) * 1000 - self._record_push_perf( - stage_ms, - stage_max_ms={"queue": max(queue_waits_ms, default=0.0)}, - item_count=len(pushes), - # `ready` holds one entry per destination shard, so count the - # distinct items rather than the writes. - byte_count=sum( - push.tensor.nbytes for push in {id(p): p for p, _ in ready}.values() - ), - skipped_items=len(pushes) - len({id(push) for push, _ in ready}), - failed=failed, - ) - - def _notify_completions( - self, notifications: list[tuple[_PendingPush, dict[str, Any]]] - ) -> None: - """Tell the consumer, in one message per destination, what landed.""" - if not notifications: - return - by_destination: dict[str, list[tuple[_PendingPush, dict[str, Any]]]] = {} - for push, reservation in notifications: - by_destination.setdefault( - str(reservation.get("addr", push.spec.consumer_zmq)), [] - ).append((push, reservation)) - for consumer_zmq, items in by_destination.items(): - result = self._send_control( - consumer_zmq, - { - "op": "complete_batch", - "items": [ - { - "transfer_id": push.spec.transfer_id, - "reservation_id": reservation["reservation_id"], - } - for push, reservation in items - ], - }, - ) - completions = result.get("items", []) if isinstance(result, dict) else [] - if len(completions) != len(items): - raise RuntimeError("Malformed EC completion response") - for (push, _), completion in zip(items, completions): - if not completion.get("completed"): - raise RuntimeError( - f"Unknown EC reservation for mm_hash={push.spec.mm_hash}" - ) - - def _abandon_pushes(self, pushes: list[_PendingPush]) -> None: - """Release the consumer-side reservations of a batch that failed.""" - for push in pushes: - shards: list[dict[str, Any]] = [] - if push.reservation.done() and not push.reservation.cancelled(): - with suppress(Exception): - shards = push.reservation.result() - if not shards: - shards = [{"addr": push.spec.consumer_zmq, "reservation_id": ""}] - for shard in shards: - with suppress(Exception): - self._send_control( - str(shard.get("addr", push.spec.consumer_zmq)), - { - "op": "cancel", - "transfer_id": push.spec.transfer_id, - "reservation_id": str(shard.get("reservation_id", "")), - "abandon": True, - }, - ) - - def _record_push_perf( - self, - stage_ms: dict[str, float], - *, - stage_max_ms: dict[str, float], - item_count: int, - byte_count: int, - skipped_items: int, - failed: bool, - ) -> None: - now = time.monotonic() - report: tuple[_PushPerfWindow, int, int] | None = None - with self._push_perf_lock: - self._active_transfer_batches -= 1 - perf = self._push_perf - perf.batches += 1 - perf.items += item_count - perf.bytes += byte_count - perf.skipped_items += skipped_items - perf.failures += int(failed) - for stage, elapsed_ms in stage_ms.items(): - perf.stage_totals_ms[stage] = ( - perf.stage_totals_ms.get(stage, 0.0) + elapsed_ms - ) - perf.stage_max_ms[stage] = max( - perf.stage_max_ms.get(stage, 0.0), - stage_max_ms.get(stage, elapsed_ms), - ) - if ( - self._transfer_metrics_log_interval > 0 - and now - perf.started_at >= self._transfer_metrics_log_interval - ): - report = ( - perf, - self._active_transfer_batches, - self._queued_transfer_batches, - ) - self._push_perf = _PushPerfWindow(started_at=now) - if report is None: - return - perf, active_batches, queued_batches = report - batches = max(perf.batches, 1) - items = max(perf.items, 1) - stage_parts = [] - for stage in ( - "queue", - "reserve", - "cuda", - "register", - "rdma", - "unregister", - "complete", - "total", - ): - divisor = items if stage == "queue" else batches - average = perf.stage_totals_ms.get(stage, 0.0) / divisor - maximum = perf.stage_max_ms.get(stage, 0.0) - stage_parts.append(f"{stage}_ms={average:.1f}/{maximum:.1f}") - stage_summary = " ".join(stage_parts) - producer_metrics = dict(self._producer_metrics) - self._producer_metrics.clear() - logger.info( - "EC Mooncake push perf: batches=%d items=%d bytes=%d " - "batch_items=%.1f skipped=%d failures=%d active=%d queued=%d " - "producer=%s queue_item_avg/max and stage_batch_avg/max: %s", - perf.batches, - perf.items, - perf.bytes, - perf.items / batches, - perf.skipped_items, - perf.failures, - active_batches, - queued_batches, - producer_metrics, - stage_summary, - ) - - def _flush_pending_pushes(self) -> None: - if not self._pending_pushes: - return - grouped: dict[str, list[_PendingPush]] = {} - for push in self._pending_pushes: - grouped.setdefault(push.spec.consumer_zmq, []).append(push) - self._pending_pushes = [] - for pushes in grouped.values(): - with self._push_perf_lock: - self._queued_transfer_batches += 1 - future = self._io_executor.submit(self._push_batch, pushes) - hashes = ",".join(push.spec.mm_hash for push in pushes) - self._pending_saves.append((hashes, future)) - - def _submit_push( - self, - tensor: torch.Tensor, - spec: ECMooncakePushSpec, - reservation: Future[list[dict[str, Any]]], - ) -> None: - ready_event = None - if tensor.device.type == "cuda": - ready_event = torch.Event() - ready_event.record(torch.accelerator.current_stream(tensor.device)) - self._pending_pushes.append( - _PendingPush( - tensor=tensor, - spec=spec, - reservation=reservation, - ready_event=ready_event, - enqueued_at=time.monotonic(), - ) - ) - - def _submit_reserved_pushes(self, tensor: torch.Tensor, mm_hash: str) -> None: - reservations = self._pending_reservations.pop(mm_hash, deque()) - if reservations: - with self._active_push_sources_lock: - self._active_push_sources[(mm_hash, id(tensor))] += len(reservations) - for spec, reservation in reservations: - self._submit_push(tensor, spec, reservation) + self._worker.start_load_caches(metadata, encoder_cache, **kwargs) - def _cancel_orphaned_reservation( - self, - spec: ECMooncakePushSpec, - reservation: Future[list[dict[str, Any]]], + def save_caches( + self, encoder_cache: dict[str, torch.Tensor], mm_hash: str, **kwargs: Any ) -> None: - try: - for shard in reservation.result(): - if shard.get("cached", False) or shard.get("cancelled", False): - continue - self._send_control( - str(shard.get("addr", spec.consumer_zmq)), - { - "op": "cancel", - "transfer_id": spec.transfer_id, - "reservation_id": str(shard.get("reservation_id", "")), - "abandon": True, - }, - ) - except Exception: - logger.exception( - "Failed to cancel orphaned EC reservation for transfer_id=%s", - spec.transfer_id, - ) + assert self._worker is not None + self._worker.save_caches(encoder_cache, mm_hash, **kwargs) def get_finished( self, finished_req_ids: set[str] ) -> tuple[set[str] | None, set[str] | None]: - if not self.is_producer or self._role != ECConnectorRole.WORKER: - return None, None - - Reserved = tuple[ECMooncakePushSpec, Future[list[dict[str, Any]]]] - orphaned: list[Reserved] = [] - for mm_hash, reservations in list(self._pending_reservations.items()): - remaining: deque[Reserved] = deque() - for spec, reservation in reservations: - if spec.request_id in finished_req_ids: - orphaned.append((spec, reservation)) - else: - remaining.append((spec, reservation)) - if remaining: - self._pending_reservations[mm_hash] = remaining - else: - self._pending_reservations.pop(mm_hash) - - for spec, reservation in orphaned: - future = self._io_executor.submit( - self._cancel_orphaned_reservation, spec, reservation - ) - self._pending_saves.append((f"cancel:{spec.transfer_id}", future)) - return None, None - - def save_caches( - self, encoder_cache: dict[str, torch.Tensor], mm_hash: str, **kwargs: Any - ) -> None: - if not self.is_producer or self._role != ECConnectorRole.WORKER: - return - tensor = encoder_cache[mm_hash] - if mm_hash in self._pending_reservations: - self._submit_reserved_pushes(tensor, mm_hash) - - def _index_pending_spec(self, spec: ECMooncakeLoadSpec) -> None: - transfer_id = spec.transfer_id or spec.mm_hash - if transfer_id in self._pending_specs: - self._consumer_scheduler_metrics["events_duplicate"] += 1 - return - self._pending_specs[transfer_id] = spec - self._pending_specs_by_hash.setdefault(spec.mm_hash, deque()).append( - transfer_id - ) - self._pending_spec_deadlines[transfer_id] = ( - time.monotonic() + _LEASE_TTL_SECONDS - ) - self._consumer_missing_since.pop(spec.mm_hash, None) - self._consumer_pending_since.setdefault(spec.mm_hash, time.monotonic()) - - def _pop_pending_spec(self, transfer_id: str) -> ECMooncakeLoadSpec | None: - spec = self._pending_specs.pop(transfer_id, None) - self._pending_spec_deadlines.pop(transfer_id, None) - self._forget_shard_readiness(transfer_id) - if spec is not None: - if not self._pending_specs_by_hash.get(spec.mm_hash): - self._consumer_pending_since.pop(spec.mm_hash, None) - transfer_ids = self._pending_specs_by_hash.get(spec.mm_hash) - if transfer_ids is not None: - with suppress(ValueError): - transfer_ids.remove(transfer_id) - if not transfer_ids: - self._pending_specs_by_hash.pop(spec.mm_hash, None) - self._consumer_pending_since.pop(spec.mm_hash, None) - else: - self._consumer_pending_since.pop(spec.mm_hash, None) - return spec - - def _first_pending_spec(self, mm_hash: str) -> ECMooncakeLoadSpec | None: - transfer_ids = self._pending_specs_by_hash.get(mm_hash) - if transfer_ids is None: - return None - while transfer_ids: - spec = self._pending_specs.get(transfer_ids[0]) - if spec is not None: - return spec - transfer_ids.popleft() - self._pending_specs_by_hash.pop(mm_hash, None) - self._consumer_pending_since.pop(mm_hash, None) - return None - - def _note_shard_ready(self, data: dict[str, Any]) -> bool: - """Whether every consumer shard has now reported this transfer ready. + assert self._worker is not None + return self._worker.get_finished(finished_req_ids) - Loading before the last rank has its copy makes that rank miss, which - `ECMooncakeWorkerMetadata.aggregate` catches by intersecting `loaded` - across ranks -- at the cost of rescheduling the whole load. - """ - if self._event_shard_count <= 1: - return True - transfer_id = str(data["transfer_id"]) - if transfer_id in self._pending_specs: - # Already indexed; later shards are just confirmations. - return False - shard = data.get("shard") - shards = self._event_ready_shards.setdefault(transfer_id, set()) - self._event_ready_shards.move_to_end(transfer_id) - shards.add(int(shard) if shard is not None else len(shards)) - if len(shards) < self._event_shard_count: - self._consumer_scheduler_metrics["events_awaiting_shards"] += 1 - while len(self._event_ready_shards) > _MAX_PENDING_EVENTS: - self._event_ready_shards.popitem(last=False) - self._consumer_scheduler_metrics["events_partial_dropped"] += 1 - return False - self._event_ready_shards.pop(transfer_id, None) - self._consumer_scheduler_metrics["events_all_shards_ready"] += 1 - return True - - def _forget_shard_readiness(self, transfer_id: str) -> None: - self._event_ready_shards.pop(transfer_id, None) - - def _store_pushed_spec(self, data: dict[str, Any]) -> None: - transfer_id = str(data["transfer_id"]) - identifier = str(data["mm_hash"]) - reservation_id = str(data["reservation_id"]) - self._index_pending_spec( - ECMooncakeLoadSpec( - mm_hash=identifier, - num_token=0, - nbytes=int(data["nbytes"]), - shape=tuple(int(value) for value in data["shape"]), - dtype=str(data["dtype"]), - pushed=True, - transfer_id=transfer_id, - reservation_id=reservation_id, - ) - ) - - def _note_resident(self, spec: ECMooncakeLoadSpec) -> None: - """Record that the worker's receive pool now holds this item.""" - self._drop_resident(spec.mm_hash) - self._resident_specs[spec.mm_hash] = ECMooncakeLoadSpec( - mm_hash=spec.mm_hash, - num_token=0, - nbytes=spec.nbytes, - shape=spec.shape, - dtype=spec.dtype, - local=True, - ) - self._resident_bytes += spec.nbytes - while ( - self._resident_specs and self._resident_bytes > self._consumer_pool_capacity - ): - _, dropped = self._resident_specs.popitem(last=False) - self._resident_bytes -= dropped.nbytes - - def _drop_resident(self, mm_hash: str) -> None: - spec = self._resident_specs.pop(mm_hash, None) - if spec is not None: - self._resident_bytes -= spec.nbytes - - def _queue_cancel(self, transfer_id: str, reservation_id: str = "") -> None: - if ( - self._reservation_zmq_addr is None - or transfer_id in self._pending_cancels - or transfer_id in self._cancelled_transfer_ids - ): - return - self._cancelled_transfer_ids[transfer_id] = ( - time.monotonic() + _LEASE_TTL_SECONDS - ) - self._pending_cancels[transfer_id] = self._control_executor.submit( - self._cancel_remote, - self._reservation_zmq_addr, - transfer_id, - reservation_id, - ) - - def _expire_pending_specs(self) -> None: - now = time.monotonic() - for transfer_id, deadline in list(self._pending_spec_deadlines.items()): - if deadline > now: - continue - spec = self._pop_pending_spec(transfer_id) - if spec is not None: - self._consumer_pending_since.pop(spec.mm_hash, None) - self._consumer_scheduler_metrics["pending_specs_expired"] += 1 - self._queue_cancel(transfer_id) - - def _ensure_event_channel(self) -> None: - if self._event_zmq_socket is not None: - return - assert self._reservation_zmq_addr is not None - shards = self._consumer_shards(self._reservation_zmq_addr) - ctx = zmq.Context() - socket = ctx.socket(zmq.PULL) - # One PULL fair-queues across every shard's PUSH. Subscribing to the - # first shard alone leaves the others' notifications queued on their - # side forever, and hides their readiness from the scheduler. - connected = 0 - for addr in shards: - try: - event_port = self._send_control(addr, {"op": "event_port"}) - address, _ = addr.rsplit(":", 1) - socket.connect(f"{address}:{int(event_port)}") - except Exception: - logger.warning( - "EC Mooncake could not subscribe to the event channel of " - "consumer shard %s; its readiness will only be seen " - "through reserve replies.", - addr, - ) - continue - connected += 1 - if not connected: - socket.close(linger=0) - ctx.term() - return - self._event_zmq_ctx = ctx - self._event_zmq_socket = socket - self._event_shard_count = connected + def build_connector_worker_meta(self) -> ECConnectorWorkerMetadata | None: + assert self._worker is not None + return self._worker.build_connector_worker_meta() - def _drain_push_notifications(self) -> None: - # `has_cache_item` and `ensure_cache_available` run once per request - # per multimodal item, so draining on every call rescans the cancel - # and deadline tables thousands of times per step. Once per step is - # enough: `build_connector_meta` re-arms this at the end of each one. - now = time.monotonic() - if not self._drain_pending and now - self._drained_at < _DRAIN_MIN_INTERVAL: - return - self._drain_pending = False - self._drained_at = now - self._poll_pending_cancels() - self._expire_pending_specs() - self._consumer_scheduler_metrics["cancel_records_dropped"] += ( - self._expire_cancel_records(self._cancelled_transfer_ids, now) - ) - if self._reservation_zmq_addr is not None: - self._ensure_event_channel() - socket = self._event_zmq_socket - if socket is None: - return - while True: - try: - data = socket.recv_json(flags=zmq.DONTWAIT) - except zmq.Again: - return - identifier = str(data["mm_hash"]) - self._consumer_scheduler_metrics["events_received"] += 1 - if data.get("ready"): - self._consumer_scheduler_metrics["events_ready"] += 1 - transfer_id = str(data["transfer_id"]) - if transfer_id in self._cancelled_transfer_ids: - self._consumer_scheduler_metrics["events_cancelled"] += 1 - continue - if identifier in self._ready_hashes: - # Redundant only for as long as the hash stays ready; hold - # on to the spec so an eviction does not strand whoever - # this transfer belongs to. - self._consumer_scheduler_metrics["events_redundant"] += 1 - if not self._note_shard_ready(data): - continue - self._store_pushed_spec(data) - else: - self._consumer_scheduler_metrics["events_not_ready"] += 1 + def take_unavailable_requests(self) -> set[str]: + assert self._scheduler is not None + return self._scheduler.take_unavailable_requests() def has_cache_item(self, identifier: str) -> bool: - if not self.is_consumer or self._role != ECConnectorRole.SCHEDULER: - return False - self._drain_push_notifications() - self._maybe_log_consumer_scheduler_metrics() - if identifier in self._ready_hashes: - self._consumer_scheduler_metrics["ready"] += 1 - self._clear_item_timers(identifier) - return True - if identifier in self._loading_hashes: - self._consumer_scheduler_metrics["loading"] += 1 - return False - if identifier in self._resident_specs: - self._consumer_scheduler_metrics["resident"] += 1 - self._consumer_missing_since.pop(identifier, None) - return True - pending = self._first_pending_spec(identifier) - if pending is not None: - self._consumer_scheduler_metrics["pending_spec"] += 1 - self._consumer_missing_since.pop(identifier, None) - return True - self._consumer_scheduler_metrics["missing_event"] += 1 - self._consumer_missing_since.setdefault(identifier, time.monotonic()) - return False - - @staticmethod - def _request_transfer_id(request: Any, index: int) -> str | None: - params = getattr(request, "ec_transfer_params", None) or {} - items = params.get("ec_items") or [] - mm_hash = request.mm_features[index].identifier - if index < len(items): - item = items[index] - if item.get("mm_hash") in (None, mm_hash) and item.get("transfer_id"): - return str(item["transfer_id"]) - for item in items: - if item.get("mm_hash") == mm_hash and item.get("transfer_id"): - return str(item["transfer_id"]) - return None + assert self._scheduler is not None + return self._scheduler.has_cache_item(identifier) def ensure_cache_available( self, - request: Any, + request: Request, num_computed_tokens: int, local_cache_hashes: Collection[str] | None = None, ) -> bool: - if self.is_producer: - for index, feature in enumerate(request.mm_features): - if ( - feature.mm_position.offset + feature.mm_position.length - > num_computed_tokens - ): - self._prepare_push_spec(request, index) - if not self.is_consumer or self._role != ECConnectorRole.SCHEDULER: - return True - - self._drain_push_notifications() - local_cache_hashes = local_cache_hashes or set() - all_ready = True - for index, feature in enumerate(request.mm_features): - if ( - feature.mm_position.offset + feature.mm_position.length - <= num_computed_tokens - ): - continue - mm_hash = feature.identifier - transfer_id = self._request_transfer_id(request, index) - if transfer_id is not None and transfer_id in self._pending_spec_deadlines: - # A live request still references this transfer, so keep it out - # of the orphan sweep in `_expire_pending_specs`. - self._pending_spec_deadlines[transfer_id] = ( - time.monotonic() + _LEASE_TTL_SECONDS - ) - if mm_hash in local_cache_hashes: - # Keep the transfer: `local_cache_hashes` is a snapshot, and - # the entry can be evicted before this request is scheduled. - # Cancelling here used to strand the request with no way to - # get the item back. `request_finished` releases it instead. - continue - if mm_hash in self._ready_hashes: - self._consumer_scheduler_metrics["ready"] += 1 - self._clear_item_timers(mm_hash) - continue - if mm_hash in self._loading_hashes: - self._consumer_scheduler_metrics["loading"] += 1 - all_ready = False - continue - # A resident copy is preferred over a transfer: it is already in - # this instance's memory, and using it leaves the transfer for - # whoever has no copy at all. - spec = self._resident_specs.get(mm_hash) - if spec is not None: - self._consumer_scheduler_metrics["resident_hit"] += 1 - else: - spec = ( - self._pending_specs.get(transfer_id) - if transfer_id is not None - else None - ) - if spec is None: - spec = self._first_pending_spec(mm_hash) - if spec is not None: - self._loading_hashes.add(mm_hash) - self._load_specs[mm_hash] = spec - self._consumer_loading_since.setdefault(mm_hash, time.monotonic()) - self._consumer_pending_since.pop(mm_hash, None) - self._mm_datas_need_loads[mm_hash] = request.get_num_encoder_embeds( - index - ) - self._scheduler_pending_work = True - all_ready = False - else: - # Keep waiting until the timeout, then let the request fail - # rather than hold a scheduler slot forever. - self._note_awaiting_push(mm_hash, transfer_id, request.request_id) - all_ready = False - return all_ready - - def _prepare_push_spec(self, request: Any, index: int) -> None: - params = getattr(request, "ec_transfer_params", None) or {} - consumer_zmq = params.get("consumer_zmq") - mm_hash = request.mm_features[index].identifier - transfer_id = self._request_transfer_id(request, index) - if transfer_id is None: - transfer_id = f"{request.request_id}:{index}" - if not consumer_zmq or transfer_id in self._prepared_push_transfer_ids: - return - num_tokens = request.get_num_encoder_embeds(index) - dtype = self._model_config.dtype - assert isinstance(dtype, torch.dtype) - dtype_name = str(dtype).split(".")[-1] - shape = (num_tokens, _get_encoder_cache_hidden_dim(self._vllm_config)) - nbytes = math.prod(shape) * dtype.itemsize - self._pushes_to_prepare[transfer_id] = ECMooncakePushSpec( - mm_hash=mm_hash, - nbytes=nbytes, - shape=shape, - dtype=dtype_name, - consumer_zmq=str(consumer_zmq), - transfer_id=transfer_id, - request_id=request.request_id, + assert self._scheduler is not None + return self._scheduler.ensure_cache_available( + request, num_computed_tokens, local_cache_hashes ) - self._prepared_push_transfer_ids.add(transfer_id) - def update_state_after_alloc(self, request: Any, index: int) -> None: - mm_hash = request.mm_features[index].identifier - if self.is_producer: - self._prepare_push_spec(request, index) - if not self.is_consumer: - return - if mm_hash in self._ready_hashes: - return - if mm_hash in self._loading_hashes: - return - num_encoder_token = request.get_num_encoder_embeds(index) - self._mm_datas_need_loads[mm_hash] = num_encoder_token - - def update_state_after_free(self, request: Any, index: int) -> None: - """Release this request's transfer as soon as it consumed the item. + def update_state_after_alloc(self, request: Request, index: int) -> None: + assert self._scheduler is not None + self._scheduler.update_state_after_alloc(request, index) - Waiting for `request_finished` would keep a consumer buffer (and the - pool slot its reservation pins) alive for the whole generation. - """ - if not self.is_consumer or self._role != ECConnectorRole.SCHEDULER: - return - transfer_id = self._request_transfer_id(request, index) - if transfer_id is None: - return - self._pop_pending_spec(transfer_id) - self._queue_cancel(transfer_id) + def update_state_after_free(self, request: Request, index: int) -> None: + assert self._scheduler is not None + self._scheduler.update_state_after_free(request, index) def build_connector_meta( self, scheduler_output: SchedulerOutput ) -> ECConnectorMetadata: - for mm_hash in scheduler_output.free_encoder_mm_hashes: - self._ready_hashes.discard(mm_hash) - self._clear_item_timers(mm_hash) - meta = ECMooncakeConnectorMetadata() - for push_spec in self._pushes_to_prepare.values(): - meta.add_push(push_spec) - self._pushes_to_prepare.clear() - for mm_hash, num_token in self._mm_datas_need_loads.items(): - load_spec = self._load_specs.pop(mm_hash, None) - if load_spec is None: - logger.warning("Missing EC Mooncake spec for mm_hash=%s", mm_hash) - continue - meta.add_load( - ECMooncakeLoadSpec( - mm_hash=load_spec.mm_hash, - num_token=num_token, - nbytes=load_spec.nbytes, - shape=load_spec.shape, - dtype=load_spec.dtype, - pushed=load_spec.pushed, - transfer_id=load_spec.transfer_id, - reservation_id=load_spec.reservation_id, - local=load_spec.local, - ) - ) - # Either way the pool holds the item once this load lands, so it - # can serve the next request without another transfer. - self._note_resident(load_spec) - if not load_spec.local: - self._pop_pending_spec(load_spec.transfer_id or load_spec.mm_hash) - self._mm_datas_need_loads.clear() - self._poll_pending_cancels() - self._maybe_log_consumer_scheduler_metrics() - self._drain_pending = True - return meta - - def build_connector_worker_meta(self) -> ECConnectorWorkerMetadata | None: - if self._role != ECConnectorRole.WORKER: - return None - if self.is_consumer and not self._is_receiving_rank: - # `loaded` is intersected across reporting ranks, so a stage that - # never loads must not report at all rather than report nothing. - return None - - self._flush_pending_pushes() - saves = self._pending_saves - completed_saves = [] - self._pending_saves = [ - (mm_hash, future) for mm_hash, future in saves if not future.done() - ] - for mm_hash, future in saves: - if future.done(): - completed_saves.append((mm_hash, future)) - for mm_hash, future in completed_saves: - try: - future.result() - except Exception: - # Publishing is best-effort: a consumer that cannot fetch this - # item falls back to encoding it locally. Failing the step - # instead would take the whole engine down. - self._producer_metrics["saves_failed"] += 1 - logger.exception( - "EC Mooncake async save failed for mm_hash=%s", mm_hash - ) - with self._consumer_lock: - reclaimed = self._consumer_reclaimed - self._consumer_reclaimed = set() - meta = ECMooncakeWorkerMetadata( - loaded=self._completed_loads, - failed_loads=self._failed_loads, - reclaimed=reclaimed, - pending_loads=False, - pending_saves=bool(self._pending_saves), - ) - self._completed_loads = set() - self._failed_loads = set() - if self.is_consumer: - self._maybe_log_consumer_worker_metrics() - return meta + assert self._scheduler is not None + return self._scheduler.build_connector_meta(scheduler_output) def update_connector_output(self, connector_output: ECConnectorOutput) -> None: - meta = connector_output.ec_connector_worker_meta - if not isinstance(meta, ECMooncakeWorkerMetadata): - return - for mm_hash in meta.loaded: - self._loading_hashes.discard(mm_hash) - self._ready_hashes.add(mm_hash) - self._clear_item_timers(mm_hash) - self._consumer_scheduler_metrics["loads_completed"] += 1 - for mm_hash in meta.failed_loads: - self._loading_hashes.discard(mm_hash) - self._load_specs.pop(mm_hash, None) - self._drop_resident(mm_hash) - self._clear_item_timers(mm_hash) - self._consumer_scheduler_metrics["loads_failed"] += 1 - for mm_hash in meta.reclaimed: - self._drop_resident(mm_hash) - self._consumer_scheduler_metrics["resident_reclaimed"] += 1 - self._scheduler_pending_work = meta.pending_loads or meta.pending_saves + assert self._scheduler is not None + self._scheduler.update_connector_output(connector_output) def has_pending_push_work(self) -> bool: - return self._scheduler_pending_work - - def _placeholder_metadata_fields(self, modality: str) -> set[str]: - if modality in self._metadata_fields_cache: - return self._metadata_fields_cache[modality] - - fields: set[str] = set() - try: - from vllm.multimodal import MULTIMODAL_REGISTRY - - info = MULTIMODAL_REGISTRY.create_processor(self._model_config).info - fields = info.data_parser.placeholder_metadata_fields(modality) - except Exception: - logger.warning( - "Could not determine the placeholder metadata fields for " - "modality %s; the consumer will preprocess the media itself.", - modality, - exc_info=True, - ) + assert self._scheduler is not None + return self._scheduler.has_pending_push_work() - self._metadata_fields_cache[modality] = fields - return fields - - def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: - if self.is_consumer and self._role == ECConnectorRole.SCHEDULER: - for index in range(len(request.mm_features)): - transfer_id = self._request_transfer_id(request, index) - if transfer_id is None: - continue - self._pop_pending_spec(transfer_id) - self._queue_cancel(transfer_id) - if ( - self.is_producer - and self._role == ECConnectorRole.SCHEDULER - and self._prepared_push_transfer_ids - ): - for index in range(len(request.mm_features)): - transfer_id = self._request_transfer_id(request, index) - if transfer_id is None: - transfer_id = f"{request.request_id}:{index}" - self._prepared_push_transfer_ids.discard(transfer_id) - if not self.is_producer: - return False, None - - items = [] - for index, feature in enumerate(request.mm_features): - metadata = {} - if feature.data is not None: - wanted = self._placeholder_metadata_fields(feature.modality) - metadata = { - key: value.tolist() - for key, value in feature.data.get_data().items() - if key in wanted and isinstance(value, torch.Tensor) - } - transfer_id = self._request_transfer_id(request, index) - item = {"mm_hash": feature.identifier, **metadata} - if transfer_id is not None: - item["transfer_id"] = transfer_id - items.append(item) - - if not items: - return False, None - return False, {"ec_items": items} + def request_finished(self, request: Request) -> tuple[bool, dict[str, Any] | None]: + assert self._scheduler is not None + return self._scheduler.request_finished(request) def shutdown(self) -> None: - if self._shutdown: + if self._closed: return - self._shutdown = True - self._flush_pending_pushes() - self._io_executor.shutdown(wait=True, cancel_futures=True) - if self._shard_pool is not None: - self._shard_pool.shutdown(wait=True, cancel_futures=True) - self._control_executor.shutdown(wait=True, cancel_futures=True) - # Every thread that could hold a control socket is stopped by now. - self._control_channel.close() - if self._control_server is not None: - self._control_server.shutdown() - if self._event_zmq_socket is not None: - self._event_zmq_socket.close(linger=0) - if self._event_zmq_ctx is not None: - self._event_zmq_ctx.term() - - if self._engine is not None: - if self._consumer_pool is not None and self._unregister_memory( - self._consumer_pool - ): - self._consumer_pool = None - self._consumer_pool_allocator = None - self._consumer_residents.clear() - self._consumer_retire_events.clear() - self._consumer_pending_frees.clear() - # Published tensors and in-flight push sources share one refcounted - # registration table, so a single pass covers both. - with self._push_source_registration_lock: - addresses = list(self._push_source_registrations) - addresses.extend(self._pending_unregister) - unregistered = True - if addresses: - ret = self._engine.batch_unregister_memory( - list(dict.fromkeys(addresses)) - ) - if ret != 0: - unregistered = False - logger.error( - "Mooncake EC batch memory unregistration failed: %d", ret - ) - if unregistered: - self._push_source_registrations.clear() - self._pending_unregister.clear() + self._closed = True + if self._scheduler is not None: + self._scheduler.close() + if self._worker is not None: + self._worker.close() def __del__(self) -> None: with suppress(Exception): From 7d034327ceafa5fac0a7429bb71faa917105bc6e Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Thu, 3 Sep 2026 01:39:14 +0000 Subject: [PATCH 21/30] [EPD] Fix the EC transfer-id desync that leaks consumer reservations Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 21 ++++++- .../ec_connector/mooncake/scheduler.py | 55 +++++++++++++++++++ .../ec_connector/mooncake/state.py | 20 +++++++ 3 files changed, 94 insertions(+), 2 deletions(-) diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index ad6877ea0dfa..ff1b4dbc8c4c 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -155,6 +155,7 @@ def rewrite_for_decode(req_data: dict, item_meta: dict[int, dict]) -> dict: meta = dict(item_meta.get(idx) or {}) idx += 1 item_uuid = meta.pop("mm_hash", None) + ec_mm_hash = meta.pop("ec_mm_hash", None) or item_uuid transfer_id = meta.pop("transfer_id", None) # Whatever keys the encoder reported are the metadata its model # declared as needed to size the placeholder range; the proxy does @@ -176,7 +177,7 @@ def rewrite_for_decode(req_data: dict, item_meta: dict[int, dict]) -> dict: ) if transfer_id is not None: transfer_items.append( - {"mm_hash": item_uuid, "transfer_id": transfer_id} + {"mm_hash": ec_mm_hash, "transfer_id": transfer_id} ) rewritten += 1 new_messages.append({**msg, "content": new_content}) @@ -281,9 +282,17 @@ async def fanout_encoder_primer( "stream": False, } if consumer_zmq is not None: + # No mm_hash here on purpose. The encoder's own + # `mm_features[i].identifier` is derived from the uuid *and* the + # engine's media_io_kwargs / mm_processor_kwargs, so this proxy + # cannot know it before the encoder runs. Sending the bare uuid + # would fail the connector's hash match and make the producer + # invent its own transfer id, which the consumer could then never + # cancel. Omitting it lets the connector match by position, which + # is exact: one encoder request carries exactly one item. encoder_req["ec_transfer_params"] = { "consumer_zmq": consumer_zmq, - "ec_items": [{"mm_hash": item_uuid, "transfer_id": transfer_id}], + "ec_items": [{"transfer_id": transfer_id}], } tasks.append( encode_session.post( @@ -334,9 +343,17 @@ async def fanout_encoder_primer( reported = [] if reported and idx in item_uuids: # One item per encoder request, so the first entry is this item's. + # The uuid this proxy assigned is what makes both instances derive + # the same cache key, so it stays the item's `uuid`. It is NOT the + # key the connectors use: when media_io_kwargs or + # mm_processor_kwargs are set, the engine re-hashes the uuid + # together with them, so `mm_features[i].identifier` is a derived + # value. The encoder already reported that derived value; carry it + # through separately instead of overwriting it. item_meta[idx] = { **reported[0], "mm_hash": item_uuids[idx], + "ec_mm_hash": reported[0].get("mm_hash"), "transfer_id": item_transfer_ids[idx], } diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py index acc226a7d85d..a79c99927057 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py @@ -54,6 +54,10 @@ logger = init_logger(__name__) +# Unmatched `ec_items` entries are a protocol desync, not a per-request +# event: warn on the first few, then only once in a while. +_MAX_UNRESOLVED_TRANSFER_ID_WARNINGS = 5 + _LEASE_TTL_SECONDS = 300 _DRAIN_MIN_INTERVAL = 0.005 _MAX_PENDING_EVENTS = 4096 @@ -144,6 +148,7 @@ def __init__( self._metadata_fields_cache: dict[str, set[str]] = {} self._consumer_metrics_started_at = time.monotonic() self._consumer_scheduler_metrics: Counter[str] = Counter() + self._unresolved_transfer_ids = 0 self._drain_pending = True self._drained_at = 0.0 self._pending_cancels: dict[str, Future[Any]] = {} @@ -418,6 +423,44 @@ def has_cache_item(self, identifier: str) -> bool: self._consumer_scheduler_metrics["missing_event"] += 1 return False + def _warn_unresolved_transfer_id( + self, request: Any, index: int, where: str + ) -> None: + """Report an `ec_items` entry that cannot be matched to a feature. + + A caller that cannot resolve a transfer id has to fall back to a + locally invented one or skip its bookkeeping entirely, and either way + the two sides of a transfer stop agreeing on its name. That desyncs + silently -- reservations are never released and only surface minutes + later as a full buffer pool -- so say it out loud the first time. + + `mm_hash` in `ec_items` must be a value some engine reported, never one + the caller derived itself: `mm_features[i].identifier` folds in the + engine's media_io_kwargs and mm_processor_kwargs, so a media uuid alone + no longer equals it. Omit the field to match by position instead. + """ + self._unresolved_transfer_ids += 1 + count = self._unresolved_transfer_ids + if count > _MAX_UNRESOLVED_TRANSFER_ID_WARNINGS and count % 1000: + return + params = getattr(request, "ec_transfer_params", None) + items = (params or {}).get("ec_items") or [] + logger.warning( + "EC Mooncake could not resolve a transfer id at %s for req=%s " + "index=%d: feature identifier=%s does not match any of the %d " + "ec_items %s (ec_transfer_params present=%s). Occurrence %d. " + "ec_items[].mm_hash must echo an engine-reported identifier, or " + "be omitted so the entry is matched by position.", + where, + request.request_id, + index, + request.mm_features[index].identifier[:16], + len(items), + [str(item.get("mm_hash"))[:16] for item in items[:4]], + params is not None, + count, + ) + @staticmethod def _request_transfer_id(request: Any, index: int) -> str | None: params = getattr(request, "ec_transfer_params", None) or {} @@ -495,6 +538,13 @@ def _prepare_push_spec(self, request: Any, index: int) -> None: mm_hash = request.mm_features[index].identifier transfer_id = self._request_transfer_id(request, index) if transfer_id is None: + # Still push: after the proxy has rewritten the item to embeds the + # consumer has no media left to fall back on, so a nameless push + # beats none. But the consumer knows this transfer by the id it + # sent, not by the one invented here, so its cancel will never + # reach the reservation this push is about to take. + if consumer_zmq: + self._warn_unresolved_transfer_id(request, index, "push prepare") transfer_id = f"{request.request_id}:{index}" if not consumer_zmq or transfer_id in self._prepared_push_transfer_ids: return @@ -525,6 +575,7 @@ def update_state_after_free(self, request: Any, index: int) -> None: return transfer_id = self._request_transfer_id(request, index) if transfer_id is None: + self._warn_unresolved_transfer_id(request, index, "encoder-cache free") return self._queue_cancel( transfer_id, @@ -537,6 +588,9 @@ def build_connector_meta( ) -> ECConnectorMetadata: for mm_hash in scheduler_output.free_encoder_mm_hashes: self._transfers.release_ready(mm_hash, time.monotonic()) + for transfer_id in self._transfers.drain_orphaned(): + self._consumer_scheduler_metrics["reservations_orphaned"] += 1 + self._queue_cancel(transfer_id) meta = ECMooncakeConnectorMetadata() for push_spec in self._pushes_to_prepare.values(): meta.add_push(push_spec) @@ -595,6 +649,7 @@ def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: for index in range(len(request.mm_features)): transfer_id = self._request_transfer_id(request, index) if transfer_id is None: + self._warn_unresolved_transfer_id(request, index, "request finish") continue self._queue_cancel( transfer_id, diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py index 3904bbb6b614..18e547da7167 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py @@ -120,10 +120,25 @@ def __init__(self, resident_capacity: int, tombstone_ttl: float) -> None: self._hash_index: dict[str, deque[str]] = {} self._loads_to_dispatch: OrderedDict[str, None] = OrderedDict() self._unavailable_requests: set[str] = set() + # Records expired straight out of READY/RESIDENT: the consumer worker + # still holds their push reservation, and nothing else tells it to let + # go, so the destination buffer would sit pinned until the lease TTL. + self._orphaned: list[str] = [] def get(self, transfer_id: str) -> SchedulerTransfer | None: return self._records.get(transfer_id) + def drain_orphaned(self) -> list[str]: + """Transfer IDs expired out of READY/RESIDENT since the last drain. + + The caller must release the consumer worker's reservation for each: + `cancel()` still accepts an EXPIRED record, so the usual cancel path + applies. + """ + orphaned = self._orphaned + self._orphaned = [] + return orphaned + def records_for_hash( self, mm_hash: str, @@ -408,6 +423,11 @@ def _transition( if now is None: raise ValueError("Terminal transition requires a timestamp") record.deadline = now + self._tombstone_ttl + if state is SchedulerTransferState.EXPIRED and record.state in ( + SchedulerTransferState.READY, + SchedulerTransferState.RESIDENT, + ): + self._orphaned.append(record.transfer_id) record.state = state record.last_error = error self._records.move_to_end(record.transfer_id) From 682955d59db4a57c675f5a6f3d551a7f6f9683d1 Mon Sep 17 00:00:00 2001 From: Zhou ziheng Date: Fri, 4 Sep 2026 11:11:56 +0800 Subject: [PATCH 22/30] [EPD] Simplify Mooncake EC connector internals (#10) * [EPD] Simplify Mooncake EC connector internals Consolidate the Mooncake EC implementation, close transfer lifecycle gaps, preserve transfer-id cleanup behavior, and retain focused scheduler, connector, proxy, and integration coverage. Signed-off-by: Zhou ziheng * [EPD] Validate Mooncake EC numeric settings Assisted-by: Codex Signed-off-by: Zhou ziheng --------- Signed-off-by: Zhou ziheng --- .../disaggregated_encoder/disagg_epd_proxy.py | 21 +- tests/v1/core/test_scheduler.py | 46 +- .../run_epd_mooncake_ec_full_pipeline.sh | 58 +- .../unit/test_ec_mooncake_connector.py | 3502 ++++------------- .../unit/test_epd_proxy_round_robin.py | 50 +- .../ec_transfer/ec_connector/base.py | 12 +- .../ec_transfer/ec_connector/cpu/connector.py | 4 +- .../ec_connector/mooncake/_availability.py | 28 - .../ec_connector/mooncake/config.py | 228 +- .../ec_connector/mooncake/control.py | 270 +- .../ec_connector/mooncake/memory.py | 156 +- .../ec_connector/mooncake/metadata.py | 56 +- .../ec_connector/mooncake/producer.py | 139 +- .../ec_connector/mooncake/reservation.py | 225 +- .../ec_connector/mooncake/scheduler.py | 255 +- .../ec_connector/mooncake/state.py | 145 +- .../ec_connector/mooncake/transfer.py | 32 +- .../ec_connector/mooncake/worker.py | 579 +-- .../ec_connector/mooncake_ec_connector.py | 38 +- vllm/v1/core/sched/scheduler.py | 4 +- 20 files changed, 1420 insertions(+), 4428 deletions(-) delete mode 100644 vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index ff1b4dbc8c4c..65ff9d7b592b 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -86,6 +86,17 @@ def encoder_rr_assignment( return urls, next_start +def validate_ec_consumer_routing( + prefill_urls: list[str], consumer_addrs: list[str] +) -> None: + """Reject the topology whose EC destination cannot be routed safely.""" + if prefill_urls and consumer_addrs: + raise ValueError( + "Mooncake EC consumer routing supports E+PD only; disable independent " + "prefill or omit --ec-consumer-zmq-addrs." + ) + + # Diagnostic switch: forward the original request to the decoder so the # only difference from the rewrite path is the rewrite itself. NO_REWRITE = False @@ -914,9 +925,9 @@ async def stop_profile(request: Request): default="", help=( "Comma-separated Mooncake EC consumer control addresses, aligned " - "with --decode-servers-urls. Required when the consumers use the " - "Mooncake EC connector. With --ec-consumer-dp-size > 1, list each " - "server's replicas consecutively: s0r0,s0r1,s1r0,s1r1." + "with --decode-servers-urls. Required for Mooncake EC consumers and " + "supported only in E+PD mode. With --ec-consumer-dp-size > 1, list " + "each server's replicas consecutively: s0r0,s0r1,s1r0,s1r1." ), ) parser.add_argument( @@ -968,6 +979,10 @@ async def stop_profile(request: Request): u.strip() for u in args.prefill_servers_urls.split(",") if u.strip() ] logger.info("Disaggregated prefill phase is enabled. Running E + P + D...") + try: + validate_ec_consumer_routing(app.state.p_urls, app.state.d_ec_urls) + except ValueError as exc: + parser.error(str(exc)) logger.info("Proxy listening on %s:%s", args.host, args.port) logger.info("Encode servers: %s", app.state.e_urls) diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index fb426080379a..2370709f16e2 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -5837,11 +5837,9 @@ def test_ec_connector_ensure_cache_available_defers_request(use_kv_connector): scheduler.add_request(request_behind) output = scheduler.schedule() - # ensure_cache_available must have been called with (request, num_computed_tokens=0) - # for a brand-new request that has no cached tokens yet. + # The public connector API remains the legacy two-argument method. ensure_call = scheduler.ec_connector.ensure_cache_available.call_args - assert ensure_call.args[:2] == (request_deferred, 0) - assert not ensure_call.args[2] + assert ensure_call.args == (request_deferred, 0) # Deferred request must NOT be scheduled assert request_deferred.request_id not in output.num_scheduled_tokens _assert_right_encoder_cache_allocated(scheduler, expected_total_allocated=0) @@ -5898,6 +5896,46 @@ def test_ec_connector_defers_running_request_for_async_reload(): assert ensure_call.args[:2] == (request, 32) +@pytest.mark.skip_global_cleanup +def test_ec_connector_legacy_ensure_cache_available_signature_is_supported(tmp_path): + """An out-of-tree connector with the original method remains callable.""" + + from vllm.distributed.ec_transfer.ec_connector.example_connector import ( + ECExampleConnector, + ) + + calls = [] + + class LegacyConnector(ECExampleConnector): + def ensure_cache_available(self, request, num_computed_tokens): + calls.append((request, num_computed_tokens)) + return False + + (tmp_path / "config.json").write_text( + '{"architectures": ["OPTForCausalLM"], "model_type": "opt"}' + ) + scheduler = create_scheduler( + model=str(tmp_path), + skip_tokenizer_init=True, + use_ec_connector=True, + ec_role="ec_consumer", + ) + scheduler.ec_connector = LegacyConnector( + scheduler.vllm_config, scheduler.ec_connector.role + ) + request = create_requests( + num_requests=1, + num_tokens=128, + mm_positions=[[PlaceholderRange(offset=48, length=32)]], + )[0] + + scheduler.add_request(request) + output = scheduler.schedule() + + assert request.request_id not in output.num_scheduled_tokens + assert calls == [(request, 0)] + + def test_ec_connector_pending_prefetch_only_checks_future_mm_features(): """Test that future mm feature filtering only yields features beyond the computed token frontier. diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh index c419784a7374..cfd1d347d469 100755 --- a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -54,6 +54,7 @@ export WITH_NVIDIA_PEERMEM="${WITH_NVIDIA_PEERMEM:-0}" LOG_PATH="${LOG_PATH:-/tmp}" BASELINE_FILE="${BASELINE_FILE:-/tmp/vllm_epd_mooncake_baseline.txt}" TIMEOUT_SECONDS="${TIMEOUT_SECONDS:-1200}" +PIDS=() mkdir -p "$LOG_PATH" @@ -76,11 +77,10 @@ import json, os print(json.dumps({ "ec_connector": "ECMooncakeConnector", "ec_role": "ec_consumer", + "ec_ip": os.environ.get("EC_MOONCAKE_RESERVATION_HOST", "127.0.0.1"), + "ec_port": int(os.environ.get("EC_MOONCAKE_RESERVATION_PORT", "19019")), "ec_connector_extra_config": { "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "rdma"), - "reservation_zmq_port": int( - os.environ.get("EC_MOONCAKE_RESERVATION_PORT", "19019") - ), }, }, separators=(",", ":"))) PY @@ -89,20 +89,41 @@ PY wait_for_server() { local port=$1 timeout "$TIMEOUT_SECONDS" bash -c " - until curl -s -o /dev/null -w '' localhost:${port}/v1/chat/completions; do + until curl -fsS http://localhost:${port}/health >/dev/null 2>&1; do sleep 2 done" && return 0 || return 1 } cleanup_instances() { - echo "Cleaning up vLLM / proxy processes..." - pkill -f "vllm serve" 2>/dev/null || true - pkill -f "vllm.entrypoints.cli.main serve" 2>/dev/null || true - pkill -f "disagg_epd_proxy.py" 2>/dev/null || true - sleep 2 + if ((${#PIDS[@]} == 0)); then + return + fi + echo "Cleaning up tracked vLLM / proxy processes..." + for pid in "${PIDS[@]}"; do + kill "$pid" 2>/dev/null || true + done + for _ in {1..10}; do + local alive=0 + for pid in "${PIDS[@]}"; do + if kill -0 "$pid" 2>/dev/null; then + alive=1 + fi + done + ((alive == 0)) && break + sleep 1 + done + for pid in "${PIDS[@]}"; do + if kill -0 "$pid" 2>/dev/null; then + kill -KILL "$pid" 2>/dev/null || true + fi + wait "$pid" 2>/dev/null || true + done + PIDS=() } -trap 'cleanup_instances; kill $(jobs -pr) 2>/dev/null || true' EXIT INT TERM +trap cleanup_instances EXIT +trap 'exit 130' INT +trap 'exit 143' TERM run_baseline() { echo "================================" @@ -118,6 +139,7 @@ run_baseline() { --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ >"${LOG_PATH}/mooncake_epd_baseline.log" 2>&1 & local BASELINE_PID=$! + PIDS+=("$BASELINE_PID") echo "Waiting for baseline..." wait_for_server "$PORT" || { echo "Baseline failed to start; tail log:"; tail -80 "${LOG_PATH}/mooncake_epd_baseline.log"; return 1; } curl -s "http://127.0.0.1:${PORT}/v1/models" | head -c 200 || true @@ -128,8 +150,6 @@ run_baseline() { --mode baseline \ --baseline_file "$BASELINE_FILE" \ $MM_FLAG - kill "$BASELINE_PID" 2>/dev/null || true - sleep 2 cleanup_instances } @@ -141,8 +161,6 @@ run_epd_mooncake() { echo "================================" cleanup_instances - declare -a PIDS=() - echo "Starting ENCODER on GPU $GPU_E port $ENCODE_PORT" CUDA_VISIBLE_DEVICES="$GPU_E" "${VLLM_SERVE[@]}" "$MODEL" \ --port "$ENCODE_PORT" \ @@ -155,7 +173,7 @@ run_epd_mooncake() { --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ --ec-transfer-config "$ENC_EC_JSON" \ >"${LOG_PATH}/mooncake_epd_encoder.log" 2>&1 & - PIDS+=($!) + PIDS+=("$!") echo "Starting PD on GPU $GPU_PD port $PREFILL_DECODE_PORT" CUDA_VISIBLE_DEVICES="$GPU_PD" "${VLLM_SERVE[@]}" "$MODEL" \ @@ -168,7 +186,7 @@ run_epd_mooncake() { --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ --ec-transfer-config "$PD_EC_JSON" \ >"${LOG_PATH}/mooncake_epd_pd.log" 2>&1 & - PIDS+=($!) + PIDS+=("$!") echo "Waiting for encoder..." wait_for_server "$ENCODE_PORT" || { echo "Encoder log:"; tail -100 "${LOG_PATH}/mooncake_epd_encoder.log"; return 1; } @@ -185,7 +203,7 @@ run_epd_mooncake() { --ec-consumer-zmq-addrs \ "tcp://localhost:$EC_MOONCAKE_RESERVATION_PORT" \ >"${LOG_PATH}/mooncake_epd_proxy.log" 2>&1 & - PIDS+=($!) + PIDS+=("$!") echo "Waiting for proxy..." wait_for_server "$ENDPOINT_PORT" || { echo "Proxy log:"; tail -80 "${LOG_PATH}/mooncake_epd_proxy.log"; return 1; } @@ -199,15 +217,11 @@ run_epd_mooncake() { --baseline_file "$BASELINE_FILE" \ $MM_FLAG - for pid in "${PIDS[@]}"; do - kill "$pid" 2>/dev/null || true - done - sleep 2 cleanup_instances } echo "================================" -echo "EPD + ECMooncake full pipeline" +echo "1E + 1PD ECMooncake end-to-end correctness" echo "MODEL=$MODEL" echo "================================" diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 580936a49054..f0c5c1433b38 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -12,18 +12,16 @@ import copy import ctypes import gc -import importlib import socket -import sys import threading import time import weakref -from collections import Counter, OrderedDict +from collections import Counter from concurrent.futures import Future, ThreadPoolExecutor from contextlib import contextmanager from dataclasses import FrozenInstanceError from multiprocessing.reduction import ForkingPickler -from types import ModuleType, SimpleNamespace +from types import SimpleNamespace from typing import Any from unittest.mock import MagicMock, Mock, call, patch @@ -33,26 +31,21 @@ from vllm.config import ModelConfig, VllmConfig from vllm.distributed.ec_transfer.ec_connector import mooncake_ec_connector -from vllm.distributed.ec_transfer.ec_connector.base import ( - ECConnectorMetadata, - ECConnectorRole, -) +from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole from vllm.distributed.ec_transfer.ec_connector.factory import ECConnectorFactory from vllm.distributed.ec_transfer.ec_connector.mooncake import ( control, memory, - metadata, - producer, - state, transfer, ) -from vllm.distributed.ec_transfer.ec_connector.mooncake.config import MooncakeECConfig +from vllm.distributed.ec_transfer.ec_connector.mooncake.config import ( + _RESERVATION_TTL_SECONDS, + MooncakeECConfig, +) from vllm.distributed.ec_transfer.ec_connector.mooncake.control import ( ConsumerControlServer, ControlClient, - ControlCompletion, EventInbox, - ShardTopology, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.memory import ( ConsumerMemoryPool, @@ -60,12 +53,18 @@ ProducerMemoryPool, ResidentPool, ) +from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeConnectorMetadata, + ECMooncakeLoadSpec, + ECMooncakePushSpec, + ECMooncakeWorkerMetadata, +) from vllm.distributed.ec_transfer.ec_connector.mooncake.producer import ( ProducerPushManager, + ProducerPushRecord, ProducerPushState, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.reservation import ( - CancellationOutcome, ConsumerReservationManager, ConsumerReservationState, ) @@ -73,42 +72,24 @@ ECMooncakeScheduler, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.state import ( - InvalidSchedulerTransferTransition, SchedulerTransferState, SchedulerTransferTable, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( MooncakeTransfer, ) -from vllm.distributed.ec_transfer.ec_connector.mooncake.worker import ( - _LEASE_TTL_SECONDS, - ECMooncakeWorker, -) +from vllm.distributed.ec_transfer.ec_connector.mooncake.worker import ECMooncakeWorker from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( ECMooncakeConnector, - ECMooncakeConnectorMetadata, - ECMooncakeLoadSpec, - ECMooncakePushSpec, - ECMooncakeWorkerMetadata, ) from vllm.v1.core.sched.output import SchedulerOutput pytest_plugins = ("tests.v1.ec_connector.unit.test_ec_example_connector",) +pytestmark = pytest.mark.skip_global_cleanup class CopyingFakeTransferEngine: - """Model Mooncake registration rules while copying bytes in-process. - - Attributes: - registered: Base addresses of currently registered ranges. - regions: Registered byte lengths keyed by base address. - register_calls: Address batches passed to memory registration. - unregister_calls: Addresses passed to single-range unregistration. - batch_unregister_calls: Address batches passed to unregistration. - transfer_calls: Byte lengths recorded for each transfer batch. - transfer_batches: Complete source and destination transfer arguments. - initialize_calls: Arguments used to initialize the fake engine. - """ + """Model Mooncake registration rules while copying bytes in-process.""" def __init__(self, *args, **kwargs): self.registered: set[int] = set() @@ -190,12 +171,18 @@ def _wait_for_worker_io( while time.monotonic() < deadline: meta = connector.build_connector_worker_meta() assert isinstance(meta, ECMooncakeWorkerMetadata) - if not meta.pending_loads and not meta.pending_saves: + if not meta.pending_saves: return meta time.sleep(0.01) raise TimeoutError("EC Mooncake worker I/O did not finish") +def _bind_extra_config(config: VllmConfig) -> None: + config.ec_transfer_config.get_from_extra_config.side_effect = lambda key, default: ( + config.ec_transfer_config.ec_connector_extra_config.get(key, default) + ) + + class TestECMooncakeControlPlane: """Validate ZMQ client reuse, shard discovery, events, and server RPCs.""" @@ -218,24 +205,6 @@ def test_worker_get_ip_failure_does_not_construct_client( client_cls.assert_not_called() - def test_worker_constructor_failure_closes_client(self, mock_vllm_config_producer): - with ( - patch_ec_mooncake_deps(), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "worker.ControlClient" - ) as client_cls, - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "worker.ThreadPoolExecutor", - side_effect=RuntimeError("executor failed"), - ), - pytest.raises(RuntimeError, match="executor failed"), - ): - ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.WORKER) - - client_cls.return_value.close.assert_called_once_with() - def test_client_reuses_socket_and_discards_failed_exchange(self): context = MagicMock() socket = context.socket.return_value @@ -301,28 +270,28 @@ def request() -> None: control_socket.connect.assert_called_once_with("tcp://consumer:19019") def test_topology_retries_transient_discovery_failures(self): - client = Mock(spec=ControlClient) + client = object.__new__(ControlClient) + client._topologies = {} + client.request = Mock() client.request.side_effect = [ {"ports": [19019, 19020]}, RuntimeError("old consumer"), {"ports": [19029, 19030]}, ] - topology = ShardTopology(client) - - assert topology.shards("tcp://consumer:19019") == [ + assert client.discover_shards("tcp://consumer:19019") == [ "tcp://consumer:19019", "tcp://consumer:19020", ] - assert topology.shards("tcp://consumer:19019") == [ + assert client.discover_shards("tcp://consumer:19019") == [ "tcp://consumer:19019", "tcp://consumer:19020", ] - assert topology.shards("tcp://legacy:19019") == ["tcp://legacy:19019"] - assert topology.shards("tcp://legacy:19019") == [ + assert client.discover_shards("tcp://legacy:19019") is None + assert client.discover_shards("tcp://legacy:19019") == [ "tcp://legacy:19029", "tcp://legacy:19030", ] - assert topology.shards("tcp://legacy:19019") == [ + assert client.discover_shards("tcp://legacy:19019") == [ "tcp://legacy:19029", "tcp://legacy:19030", ] @@ -333,7 +302,9 @@ def test_topology_retries_transient_discovery_failures(self): ] def test_event_inbox_retries_until_every_shard_is_connected(self): - client = Mock(spec=ControlClient) + client = object.__new__(ControlClient) + client._topologies = {} + client.request = Mock() client.request.side_effect = [ RuntimeError("peers not ready"), {"ports": [19019, 19020]}, @@ -342,14 +313,13 @@ def test_event_inbox_retries_until_every_shard_is_connected(self): 20001, 20002, ] - topology = ShardTopology(client) context = MagicMock() socket = context.socket.return_value event = {"transfer_id": "transfer", "ready": True} socket.recv_json.side_effect = [event, zmq.Again()] with patch.object(control.zmq, "Context", return_value=context) as create: - inbox = EventInbox(client, topology) + inbox = EventInbox(client) assert inbox.drain("tcp://consumer:19019") == [] assert inbox.shard_count == 1 create.assert_not_called() @@ -391,7 +361,7 @@ def status(transfer_id: str): def complete(transfer_id: str, reservation_id: str): completed.append((transfer_id, reservation_id)) - return ControlCompletion(True, became_ready=True) + return True, True def cancel( transfer_id: str, @@ -483,9 +453,12 @@ def mock_vllm_config_producer(): config.ec_transfer_config.is_ec_consumer = False config.ec_transfer_config.ec_buffer_device = "cuda" config.ec_transfer_config.ec_buffer_size = 1e9 + config.ec_transfer_config.ec_ip = "127.0.0.1" + config.ec_transfer_config.ec_port = 19019 config.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", } + _bind_extra_config(config) return config @@ -502,10 +475,12 @@ def mock_vllm_config_consumer(): config.ec_transfer_config.is_ec_consumer = True config.ec_transfer_config.ec_buffer_device = "cuda" config.ec_transfer_config.ec_buffer_size = 1e9 + config.ec_transfer_config.ec_ip = "127.0.0.1" + config.ec_transfer_config.ec_port = 19019 config.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, } + _bind_extra_config(config) return config @@ -518,7 +493,7 @@ def patch_ec_mooncake_deps(): ), patch( "vllm.distributed.ec_transfer.ec_connector.mooncake." - "_availability._MOONCAKE_IMPORT_ERROR", + "transfer._MOONCAKE_IMPORT_ERROR", None, ), patch( @@ -666,7 +641,7 @@ def write(): class TestECMooncakeFactory: - """Validate factory registration and compatibility exports.""" + """Validate factory registration.""" def test_factory_registers_connector(self): cls = ECConnectorFactory.get_connector_class( @@ -678,15 +653,6 @@ def test_factory_registers_connector(self): == "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector" ) - def test_public_exports_are_compatible_and_narrow(self): - assert mooncake_ec_connector.__all__ == [ - "ECMooncakeConnector", - "ECMooncakeConnectorMetadata", - "ECMooncakeLoadSpec", - "ECMooncakePushSpec", - "ECMooncakeWorkerMetadata", - ] - class TestContiguousAllocator: """Validate aligned allocation, reuse, and range coalescing.""" @@ -717,8 +683,8 @@ class TestResidentPool: def test_lru_skips_rejected_entry_and_replaces_without_losing_owner(self): pool = ResidentPool[str]() - pool.insert("oldest", "first", 256) - pool.insert("next", "second", 256) + pool.insert("oldest", "first") + pool.insert("next", "second") pool.retire("oldest") pool.retire("next") @@ -726,22 +692,19 @@ def test_lru_skips_rejected_entry_and_replaces_without_losing_owner(self): assert evicted == "next" assert pool.get("oldest") == "first" - assert pool.insert("oldest", "replacement", 128) == "first" + assert pool.insert("oldest", "replacement") == "first" assert pool.get("oldest") == "replacement" - assert pool.used == 128 def test_displaced_entry_waits_for_every_lease(self): pool = ResidentPool[str]() - pool.insert("hash", "original", 256) + pool.insert("hash", "original") first = pool.acquire("hash") second = pool.acquire("hash") assert first is not None and second is not None - assert pool.insert("hash", "replacement", 256) is None - assert pool.used == 512 + assert pool.insert("hash", "replacement") is None assert pool.release(first) is None assert pool.release(second) == "original" - assert pool.used == 256 assert pool.release(second) is None @@ -765,7 +728,7 @@ def test_consumer_replacement_waits_for_cached_owner(self): mooncake_transfer.register_memory.return_value = 0 mooncake_transfer.unregister_memory.return_value = True pool = ConsumerMemoryPool(768, mooncake_transfer) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + pool.prepare(torch.device("cpu")) first = pool.try_allocate(64, (16,), torch.float32) replacement = pool.try_allocate(64, (16,), torch.float32) assert first is not None and replacement is not None @@ -788,7 +751,7 @@ def test_cached_consume_returns_newer_canonical_allocation(self): mooncake_transfer = MagicMock(spec=MooncakeTransfer) mooncake_transfer.register_memory.return_value = 0 pool = ConsumerMemoryPool(768, mooncake_transfer) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + pool.prepare(torch.device("cpu")) first = pool.try_allocate(64, (16,), torch.float32) replacement = pool.try_allocate(64, (16,), torch.float32) assert first is not None and replacement is not None @@ -808,16 +771,13 @@ def test_consumer_defers_retired_reuse_until_event_completes(self): mooncake_transfer = MagicMock(spec=MooncakeTransfer) mooncake_transfer.register_memory.return_value = 0 pool = ConsumerMemoryPool(256, mooncake_transfer) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + pool.prepare(torch.device("cpu")) allocation = pool.try_allocate(64, (16,), torch.float32) assert allocation is not None pool.publish("hash", allocation) event = self._Event(complete=False) - with ( - patch.object(memory.torch, "Event", return_value=event), - patch.object(memory.torch.accelerator, "current_stream"), - ): + with patch.object(pool, "_record_release_event", return_value=event): pool.retire_stale({}, set()) assert pool.reclaim_and_allocate(64, (16,), torch.float32) is None event.complete = True @@ -832,30 +792,18 @@ def test_consumer_registration_failure_disables_pool(self): mooncake_transfer.register_memory.return_value = 1 pool = ConsumerMemoryPool(256, mooncake_transfer) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + pool.prepare(torch.device("cpu")) + pool.prepare(torch.device("cpu")) assert pool.tensor is None mooncake_transfer.register_memory.assert_called_once() - def test_nonreceiving_consumer_never_registers_or_unregisters_pool(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - pool = ConsumerMemoryPool(256, mooncake_transfer) - - pool.prepare(torch.device("cpu"), receiving_rank=False, allow_host=True) - pool.close() - pool.close() - - assert pool.tensor is None - mooncake_transfer.register_memory.assert_not_called() - mooncake_transfer.unregister_memory.assert_not_called() - def test_consumer_close_unregisters_once_and_releases_parent(self): mooncake_transfer = MagicMock(spec=MooncakeTransfer) mooncake_transfer.register_memory.return_value = 0 mooncake_transfer.unregister_memory.return_value = True pool = ConsumerMemoryPool(256, mooncake_transfer) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) + pool.prepare(torch.device("cpu")) parent = pool.tensor pool.close() @@ -864,9 +812,10 @@ def test_consumer_close_unregisters_once_and_releases_parent(self): mooncake_transfer.unregister_memory.assert_called_once_with(parent) assert pool.tensor is None - def test_producer_reuses_staging_and_keeps_parent_for_later_close_phase(self): + def test_producer_reuses_staging_and_unregisters_parent_on_close(self): mooncake_transfer = MagicMock(spec=MooncakeTransfer) mooncake_transfer.register_memory.return_value = 0 + mooncake_transfer.unregister_memory.return_value = True pool = ProducerMemoryPool(256, mooncake_transfer) source = torch.arange(16, dtype=torch.float32) @@ -878,13 +827,13 @@ def test_producer_reuses_staging_and_keeps_parent_for_later_close_phase(self): assert second is not None assert second.regions == first.regions pool.release(second) - parent = pool.tensor + parent = pool._pool pool.close() pool.close() - assert pool.tensor is parent - mooncake_transfer.unregister_memory.assert_not_called() + assert pool._pool is None + mooncake_transfer.unregister_memory.assert_called_once_with(parent) def test_producer_falls_back_when_staging_pool_allocation_fails(self): mooncake_transfer = MagicMock(spec=MooncakeTransfer) @@ -898,106 +847,43 @@ def test_producer_falls_back_when_staging_pool_allocation_fails(self): class TestMooncakeECConfig: - """Validate normalization, defaults, immutability, and bounds.""" - - def test_defaults_are_an_immutable_snapshot(self, mock_vllm_config_producer): - config = MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - - assert config == MooncakeECConfig( - is_producer=True, - is_consumer=False, - protocol="tcp", - buffer_device="cuda", - reservation_port=None, - reservation_addr=None, - control_timeout_s=30, - push_wait_timeout_s=60, - transfer_workers=4, - control_workers=8, - producer_pool_size=1_000_000_000, - consumer_pool_size=1_000_000_000, - transfer_metrics_log_interval=10, - consumer_metrics_log_interval=10, - ) + def test_defaults_are_a_frozen_snapshot(self, mock_vllm_config_producer): + config = MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "mooncake_protocol" - ] = "rdma" - assert config.protocol == "tcp" + assert ( + config.protocol, + config.buffer_device, + config.control_timeout_ms, + config.push_wait_timeout_s, + config.pool_size, + ) == ("tcp", "cuda", 30_000, 60, 1_000_000_000) with pytest.raises(FrozenInstanceError): config.protocol = "rdma" # type: ignore[misc] - def test_custom_values_are_normalized(self, mock_vllm_config_consumer): + def test_derives_rank_local_port_and_custom_resources( + self, mock_vllm_config_consumer + ): source = mock_vllm_config_consumer source.parallel_config.tensor_parallel_size = 2 - source.parallel_config.data_parallel_size = 3 source.parallel_config.data_parallel_index = 1 - source.ec_transfer_config.ec_buffer_device = " cpu " source.ec_transfer_config.ec_buffer_size = 2048 - source.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": " tcp ", - "reservation_zmq_port": "5000", - "control_timeout_s": "1.5", - "push_wait_timeout_s": "2.5", - "transfer_max_workers": "3", - "control_max_workers": "5", - "producer_buffer_pool_size": "1024", - "consumer_buffer_pool_size": "1536", - "transfer_metrics_log_interval": "0", - "consumer_metrics_log_interval": "7", - } - - config = MooncakeECConfig.from_vllm_config(source, ECConnectorRole.WORKER) - - assert config.reservation_port == 5002 - assert config.reservation_addr == "tcp://127.0.0.1:5002" - assert config.control_timeout_s == 1.5 - assert config.push_wait_timeout_s == 2.5 - assert config.transfer_workers == 3 - assert config.control_workers == 5 - assert config.producer_pool_size == 1024 - assert config.consumer_pool_size == 1536 - assert config.transfer_metrics_log_interval == 0 - assert config.consumer_metrics_log_interval == 7 - assert config.buffer_device == "cpu" - assert config.protocol == "tcp" - - @pytest.mark.parametrize( - ("key", "value"), - [ - ("control_timeout_s", 0), - ("push_wait_timeout_s", -1), - ("transfer_max_workers", 0), - ("control_max_workers", -1), - ("producer_buffer_pool_size", 0), - ("consumer_buffer_pool_size", -1), - ], - ) - @pytest.mark.parametrize( - "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] - ) - def test_rejects_nonpositive_values( - self, mock_vllm_config_producer, key, value, role - ): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( - value + source.ec_transfer_config.ec_port = 5000 + source.ec_transfer_config.ec_connector_extra_config.update( + { + "control_timeout_s": 1.5, + "push_wait_timeout_s": 2.5, + } ) - with pytest.raises(ValueError, match=key): - MooncakeECConfig.from_vllm_config(mock_vllm_config_producer, role) - - @pytest.mark.parametrize( - "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] - ) - def test_rejects_nonpositive_registered_buffer( - self, mock_vllm_config_producer, role - ): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = 0 + config = MooncakeECConfig.from_vllm_config(source) - with pytest.raises(ValueError, match="ec_buffer_size > 0"): - MooncakeECConfig.from_vllm_config(mock_vllm_config_producer, role) + assert config.control_port == 5002 + assert config.control_addr == "tcp://127.0.0.1:5002" + assert ( + config.control_timeout_ms, + config.push_wait_timeout_s, + config.pool_size, + ) == (1500, 2.5, 2048) @pytest.mark.parametrize("key", ["control_timeout_s", "push_wait_timeout_s"]) @pytest.mark.parametrize("value", [float("nan"), float("inf"), float("-inf"), True]) @@ -1007,683 +893,185 @@ def test_rejects_invalid_timeouts(self, mock_vllm_config_producer, key, value): ) with pytest.raises(ValueError, match=key): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) + MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) @pytest.mark.parametrize( - "key", - [ - "reservation_zmq_port", - "transfer_max_workers", - "control_max_workers", - "producer_buffer_pool_size", - "consumer_buffer_pool_size", - ], + "value", [1.5, True, float("nan"), float("inf"), float("-inf")] ) - @pytest.mark.parametrize("value", [1.5, True]) - def test_rejects_noninteger_values(self, mock_vllm_config_producer, key, value): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( - value - ) - - with pytest.raises(ValueError, match=key): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - - @pytest.mark.parametrize("value", [1.5, True]) - def test_rejects_noninteger_registered_buffer( - self, mock_vllm_config_producer, value - ): + def test_rejects_invalid_registered_buffer(self, mock_vllm_config_producer, value): mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = value with pytest.raises(ValueError, match="ec_buffer_size"): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - - def test_accepts_integral_float_integer_values(self, mock_vllm_config_producer): - source = mock_vllm_config_producer - source.ec_transfer_config.ec_buffer_size = 8.0 - source.ec_transfer_config.ec_connector_extra_config.update( - { - "reservation_zmq_port": 5000.0, - "transfer_max_workers": 2.0, - "control_max_workers": 3.0, - "producer_buffer_pool_size": 4.0, - "consumer_buffer_pool_size": 5.0, - } - ) - - config = MooncakeECConfig.from_vllm_config(source, ECConnectorRole.WORKER) - - assert ( - config.reservation_port, - config.transfer_workers, - config.control_workers, - config.producer_pool_size, - config.consumer_pool_size, - ) == (5000, 2, 3, 4, 5) + MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) @pytest.mark.parametrize( - ("key", "value"), + ("attribute", "message"), [ - ("mooncake_protocol", ""), - ("mooncake_protocol", " "), - ("mooncake_protocol", None), - ("mooncake_protocol", 1), - ("reservation_zmq_addr", ""), - ("reservation_zmq_addr", " "), - ("reservation_zmq_addr", None), - ("reservation_zmq_addr", 1), + ("tensor_parallel_size", "tensor_parallel_size=1"), + ("pipeline_parallel_size", "pipeline parallelism"), + ("data_parallel_size", "data_parallel_size=1"), ], ) - def test_rejects_invalid_strings(self, mock_vllm_config_producer, key, value): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( - value - ) - - with pytest.raises(ValueError, match=key): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - - @pytest.mark.parametrize("buffer_device", [None, "", " \t"]) - def test_normalizes_default_buffer_device( - self, mock_vllm_config_producer, buffer_device + def test_rejects_sharded_producer( + self, mock_vllm_config_producer, attribute, message ): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = buffer_device - - config = MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) + setattr(mock_vllm_config_producer.parallel_config, attribute, 2) + with pytest.raises(ValueError, match=message): + MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) - assert config.buffer_device == "cuda" - - def test_strips_buffer_device(self, mock_vllm_config_producer): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = " cuda " + @pytest.mark.parametrize("port", [0, 65536]) + def test_rejects_out_of_range_port(self, mock_vllm_config_consumer, port): + mock_vllm_config_consumer.ec_transfer_config.ec_port = port + with pytest.raises(ValueError, match="1..65535"): + MooncakeECConfig.from_vllm_config(mock_vllm_config_consumer) - config = MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) + def test_uses_upstream_ip_and_port(self, mock_vllm_config_consumer): + mock_vllm_config_consumer.ec_transfer_config.ec_ip = "consumer" + mock_vllm_config_consumer.ec_transfer_config.ec_port = 19100 - assert config.buffer_device == "cuda" + config = MooncakeECConfig.from_vllm_config(mock_vllm_config_consumer) - @pytest.mark.parametrize("buffer_device", [1, True]) - def test_rejects_invalid_buffer_device( - self, mock_vllm_config_producer, buffer_device - ): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = buffer_device + assert config.control_addr == "tcp://consumer:19100" - with pytest.raises(ValueError, match="ec_buffer_device"): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) +class TestECMooncakeConnectorValidation: @pytest.mark.parametrize( - "key", ["transfer_metrics_log_interval", "consumer_metrics_log_interval"] + "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] ) + def test_requires_mooncake_dependency(self, mock_vllm_config_producer, role): + with ( + patch.object(transfer, "_MOONCAKE_IMPORT_ERROR", ImportError("missing")), + pytest.raises(ImportError, match="mooncake-transfer-engine"), + ): + ECMooncakeConnector(mock_vllm_config_producer, role) + @pytest.mark.parametrize( - "value", [-1, float("nan"), float("inf"), float("-inf"), True, None] + ("role", "active", "inactive"), + [ + (ECConnectorRole.SCHEDULER, "_scheduler", "_worker"), + (ECConnectorRole.WORKER, "_worker", "_scheduler"), + ], ) - def test_rejects_invalid_metrics_intervals( - self, mock_vllm_config_producer, key, value - ): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( - value - ) - - with pytest.raises(ValueError, match=key): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - - def test_zero_disables_metrics(self, mock_vllm_config_producer): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config.update( - { - "transfer_metrics_log_interval": 0, - "consumer_metrics_log_interval": 0, - } - ) - - config = MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - - assert config.transfer_metrics_log_interval == 0 - assert config.consumer_metrics_log_interval == 0 - - def test_submillisecond_control_timeout_uses_one_millisecond( - self, mock_vllm_config_producer + def test_constructs_one_delegate_and_closes_once( + self, mock_vllm_config_producer, role, active, inactive ): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "control_timeout_s" - ] = 0.0001 + scheduler = Mock() + worker = Mock() with ( - patch_ec_mooncake_deps(), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "scheduler.ControlClient" - ) as scheduler_client, - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "worker.ControlClient" - ) as worker_client, + patch.object( + mooncake_ec_connector, + "ECMooncakeScheduler", + return_value=scheduler, + ), + patch.object( + mooncake_ec_connector, + "ECMooncakeWorker", + return_value=worker, + ), ): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - worker = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - scheduler.shutdown() - worker.shutdown() - - scheduler_client.assert_called_once_with(1) - worker_client.assert_called_once_with(1) + connector = ECMooncakeConnector(mock_vllm_config_producer, role) + assert getattr(connector, active) is not None + assert getattr(connector, inactive) is None + connector.shutdown() + connector.shutdown() - @pytest.mark.parametrize("timeout", [1e308, sys.float_info.max]) - def test_rejects_control_timeout_too_large_for_zmq( - self, mock_vllm_config_producer, timeout - ): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[ - "control_timeout_s" - ] = timeout + if role == ECConnectorRole.SCHEDULER: + scheduler.close.assert_called_once_with() + worker.close.assert_not_called() + else: + worker.close.assert_called_once_with() + scheduler.close.assert_not_called() - with pytest.raises(ValueError, match="control_timeout_s"): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - def test_rejects_pipeline_parallel_producer(self, mock_vllm_config_producer): - mock_vllm_config_producer.parallel_config.pipeline_parallel_size = 2 +class TestECMooncakeMetadata: + @pytest.mark.parametrize( + "value", + [ + ECMooncakeConnectorMetadata( + loads=[ + ECMooncakeLoadSpec( + mm_hash="load", + nbytes=8, + shape=(2, 4), + dtype="float16", + transfer_id="transfer", + local=True, + ) + ], + pushes=[ + ECMooncakePushSpec( + mm_hash="push", + nbytes=8, + shape=(2, 4), + dtype="float16", + consumer_zmq="tcp://127.0.0.1:1234", + transfer_id="transfer", + request_id="request", + ) + ], + ), + ECMooncakeWorkerMetadata( + loaded={"loaded"}, + failed_loads={"failed"}, + reclaimed={"reclaimed"}, + pending_saves=True, + ), + ], + ) + def test_pickle_round_trip(self, value): + assert ForkingPickler.loads(ForkingPickler.dumps(value)) == value - with pytest.raises(ValueError, match="pipeline parallelism"): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - @pytest.mark.parametrize("port", [0, 65536]) - def test_rejects_out_of_range_base_port(self, mock_vllm_config_consumer, port): - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "reservation_zmq_port" - ] = port +class TestECMooncakeWorkerMetadataAggregation: + """Validate cross-rank success intersection and failure union rules.""" - with pytest.raises(ValueError, match="1..65535"): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) + def test_an_item_one_rank_missed_is_not_loaded(self): + """Each rank gathers from its own cache, so all of them must have it. - def test_rejects_topology_that_overflows_port_range( - self, mock_vllm_config_consumer - ): - source = mock_vllm_config_consumer - source.parallel_config.tensor_parallel_size = 4 - source.parallel_config.data_parallel_index = 1 - source.ec_transfer_config.ec_connector_extra_config["reservation_zmq_port"] = ( - 65530 - ) + Reporting it as loaded because one rank succeeded left the scheduler + marking the hash ready while another rank raised on the cache miss. + """ + rank0 = ECMooncakeWorkerMetadata(loaded={"a", "b"}) + rank1 = ECMooncakeWorkerMetadata(loaded={"a"}, failed_loads={"b"}) - with pytest.raises(ValueError, match="ports must be in 1..65535"): - MooncakeECConfig.from_vllm_config(source, ECConnectorRole.SCHEDULER) + merged = rank0.aggregate(rank1) - def test_consumer_role_requirements_differ_only_by_process_role( - self, mock_vllm_config_consumer - ): - extra = {"reservation_zmq_addr": "tcp://consumer:19019"} - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = extra + assert merged.loaded == {"a"} + assert merged.failed_loads == {"b"} - scheduler = MooncakeECConfig.from_vllm_config( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + def test_a_reclaim_on_any_rank_invalidates_residency(self): + """The scheduler mirrors one pool, so the weakest rank decides.""" + merged = ECMooncakeWorkerMetadata(loaded={"a"}).aggregate( + ECMooncakeWorkerMetadata(loaded={"a"}, reclaimed={"c"}) ) - assert scheduler.reservation_addr == "tcp://consumer:19019" - with pytest.raises(ValueError, match="workers require reservation_zmq_port"): - MooncakeECConfig.from_vllm_config( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) + assert merged.reclaimed == {"c"} -class TestECMooncakeConnectorValidation: - """Validate role construction, optional dependencies, and topology rules.""" +class TestSchedulerTransferTable: + @staticmethod + def pushed_spec(transfer_id: str, mm_hash: str = "hash") -> ECMooncakeLoadSpec: + return ECMooncakeLoadSpec( + mm_hash=mm_hash, + nbytes=16, + shape=(4,), + dtype="float32", + transfer_id=transfer_id, + ) - @pytest.mark.parametrize( - "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] - ) - def test_requires_transfer_engine_symbol_for_each_role( - self, mock_vllm_config_producer, monkeypatch, role - ): - from vllm.distributed.ec_transfer.ec_connector.mooncake import _availability - - fake_package = ModuleType("mooncake") - fake_package.__path__ = [] - fake_engine = ModuleType("mooncake.engine") - try: - with monkeypatch.context() as context: - context.setitem(sys.modules, "mooncake", fake_package) - context.setitem(sys.modules, "mooncake.engine", fake_engine) - importlib.reload(_availability) - with pytest.raises(ImportError, match="mooncake-transfer-engine"): - ECMooncakeConnector(mock_vllm_config_producer, role) - finally: - importlib.reload(_availability) - - def test_rejects_sharded_producer(self, mock_vllm_config_producer): - """One copy of each encoder output, so sharding only duplicates it.""" - mock_vllm_config_producer.parallel_config.tensor_parallel_size = 2 - with ( - patch_ec_mooncake_deps(), - pytest.raises(ValueError, match="tensor_parallel_size"), - ): - ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.WORKER) - - def test_accepts_sharded_consumer(self, mock_vllm_config_consumer): - """Consumers shard: each rank gathers from its own encoder cache.""" - mock_vllm_config_consumer.parallel_config.tensor_parallel_size = 4 - mock_vllm_config_consumer.parallel_config.pipeline_parallel_size = 2 - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - connector.shutdown() - - def test_replicated_consumer_addresses_its_own_block( - self, mock_vllm_config_consumer - ): - """Each replica owns a distinct block of control ports. - - Replicas run their own schedulers and control channels, so sharing a - port would collide at bind time and cross-subscribe their event - channels. The block is derived from `data_parallel_index` because a - non-MoE replica is reconfigured to look like DP=1, which resets - `data_parallel_rank`. - """ - cfg = mock_vllm_config_consumer - cfg.parallel_config.tensor_parallel_size = 2 - cfg.parallel_config.data_parallel_size = 3 - cfg.parallel_config.data_parallel_index = 2 - # What a non-MoE replica actually looks like: reconfigured to DP=1, so - # `data_parallel_rank` no longer identifies it but the index still does. - cfg.parallel_config.data_parallel_rank = 0 - cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - "reservation_zmq_port": 19500, - } - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector(cfg, ECConnectorRole.SCHEDULER) - try: - # Replica 2 of a TP=2 consumer starts after two 2-port blocks. - assert connector._scheduler is not None - assert ( - connector._scheduler._reservation_zmq_addr - == "tcp://127.0.0.1:19504" - ) - finally: - connector.shutdown() - - def test_rejects_replicated_producer(self, mock_vllm_config_producer): - """A producer holds one copy of each output, so replicating it only - duplicates the push.""" - mock_vllm_config_producer.parallel_config.data_parallel_size = 2 - with ( - patch_ec_mooncake_deps(), - pytest.raises(ValueError, match="data_parallel_size=1"), - ): - ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.SCHEDULER) - - def test_scheduler_hooks_route_exactly(self, mock_vllm_config_producer): - scheduler = Mock() - scheduler.take_unavailable_requests.return_value = {"unavailable"} - scheduler.has_cache_item.return_value = True - scheduler.ensure_cache_available.return_value = False - scheduler.build_connector_meta.return_value = "metadata" - scheduler.has_pending_push_work.return_value = True - scheduler.request_finished.return_value = (True, {"result": 1}) - with patch.object( - ECMooncakeScheduler, - "from_vllm_config", - return_value=scheduler, - ) as from_vllm_config: - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - request = Mock() - scheduler_output = Mock() - connector_output = Mock() - try: - assert connector.take_unavailable_requests() == {"unavailable"} - assert connector.has_cache_item("hash") is True - assert connector.ensure_cache_available(request, 7, {"local"}) is False - connector.update_state_after_alloc(request, 2) - connector.update_state_after_free(request, 3) - assert connector.build_connector_meta(scheduler_output) == "metadata" - connector.update_connector_output(connector_output) - assert connector.has_pending_push_work() is True - assert connector.request_finished(request) == (True, {"result": 1}) - finally: - connector.shutdown() - - from_vllm_config.assert_called_once_with(mock_vllm_config_producer) - scheduler.take_unavailable_requests.assert_called_once_with() - scheduler.has_cache_item.assert_called_once_with("hash") - scheduler.ensure_cache_available.assert_called_once_with(request, 7, {"local"}) - scheduler.update_state_after_alloc.assert_called_once_with(request, 2) - scheduler.update_state_after_free.assert_called_once_with(request, 3) - scheduler.build_connector_meta.assert_called_once_with(scheduler_output) - scheduler.update_connector_output.assert_called_once_with(connector_output) - scheduler.has_pending_push_work.assert_called_once_with() - scheduler.request_finished.assert_called_once_with(request) - scheduler.close.assert_called_once_with() - - def test_worker_hooks_route_exactly(self, mock_vllm_config_producer): - metadata = ECMooncakeConnectorMetadata() - worker = Mock() - worker.get_finished.return_value = ({"saved"}, {"loaded"}) - worker.build_connector_worker_meta.return_value = "worker-metadata" - with patch.object( - ECMooncakeWorker, - "from_vllm_config", - return_value=worker, - ) as from_vllm_config: - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - connector.bind_connector_metadata(metadata) - encoder_cache: dict[str, torch.Tensor] = {} - try: - connector.start_worker_services() - connector.start_save_caches(encoder_cache=encoder_cache, marker=1) - connector.start_load_caches(encoder_cache, marker=2) - connector.save_caches(encoder_cache, "hash", marker=3) - assert connector.get_finished({"finished"}) == ( - {"saved"}, - {"loaded"}, - ) - assert connector.build_connector_worker_meta() == "worker-metadata" - finally: - connector.shutdown() - - from_vllm_config.assert_called_once_with(mock_vllm_config_producer) - worker.start_services.assert_called_once_with() - worker.start_save_caches.assert_called_once_with( - metadata, encoder_cache=encoder_cache, marker=1 - ) - worker.start_load_caches.assert_called_once_with( - metadata, encoder_cache, marker=2 - ) - worker.save_caches.assert_called_once_with(encoder_cache, "hash", marker=3) - worker.get_finished.assert_called_once_with({"finished"}) - worker.build_connector_worker_meta.assert_called_once_with() - worker.close.assert_called_once_with() - - @pytest.mark.parametrize( - ("method", "args"), - [ - ("start_save_caches", ()), - ("start_load_caches", ({},)), - ], - ) - def test_worker_load_and_save_reject_wrong_metadata( - self, mock_vllm_config_producer, method, args - ): - class OtherMetadata(ECConnectorMetadata): - """Represent an incompatible connector metadata implementation.""" - - pass - - worker = Mock() - with patch.object( - ECMooncakeWorker, - "from_vllm_config", - return_value=worker, - ): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - connector.bind_connector_metadata(OtherMetadata()) - try: - with pytest.raises(AssertionError): - getattr(connector, method)(*args) - finally: - connector.shutdown() - - getattr(worker, method).assert_not_called() - - @pytest.mark.parametrize( - ("method", "args", "kwargs"), - [ - ("start_worker_services", (), {}), - ("start_save_caches", (), {}), - ("start_load_caches", ({},), {}), - ("save_caches", ({}, "hash"), {}), - ("get_finished", (set(),), {}), - ("build_connector_worker_meta", (), {}), - ], - ) - def test_scheduler_rejects_worker_hooks( - self, mock_vllm_config_producer, method, args, kwargs - ): - with patch.object( - ECMooncakeScheduler, - "from_vllm_config", - return_value=Mock(), - ): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - try: - with pytest.raises(AssertionError): - getattr(connector, method)(*args, **kwargs) - finally: - connector.shutdown() - - @pytest.mark.parametrize( - ("method", "args", "kwargs"), - [ - ("take_unavailable_requests", (), {}), - ("has_cache_item", ("hash",), {}), - ("ensure_cache_available", (Mock(), 0, set()), {}), - ("update_state_after_alloc", (Mock(), 0), {}), - ("update_state_after_free", (Mock(), 0), {}), - ("build_connector_meta", (Mock(),), {}), - ("update_connector_output", (Mock(),), {}), - ("has_pending_push_work", (), {}), - ("request_finished", (Mock(),), {}), - ], - ) - def test_worker_rejects_scheduler_hooks( - self, mock_vllm_config_producer, method, args, kwargs - ): - with patch.object( - ECMooncakeWorker, - "from_vllm_config", - return_value=Mock(), - ): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - try: - with pytest.raises(AssertionError): - getattr(connector, method)(*args, **kwargs) - finally: - connector.shutdown() - - @pytest.mark.parametrize( - ("role", "active", "inactive"), - [ - (ECConnectorRole.SCHEDULER, "_scheduler", "_worker"), - (ECConnectorRole.WORKER, "_worker", "_scheduler"), - ], - ) - def test_exactly_one_role_and_idempotent_shutdown( - self, mock_vllm_config_producer, role, active, inactive - ): - scheduler = Mock() - worker = Mock() - with ( - patch.object( - ECMooncakeScheduler, - "from_vllm_config", - return_value=scheduler, - ), - patch.object( - ECMooncakeWorker, - "from_vllm_config", - return_value=worker, - ), - ): - connector = ECMooncakeConnector(mock_vllm_config_producer, role) - assert getattr(connector, active) is not None - assert getattr(connector, inactive) is None - assert set(connector.__dict__) - { - "_connector_metadata", - "_vllm_config", - "_role", - "_is_producer", - "_is_consumer", - } == {"_scheduler", "_worker", "_closed"} - connector.shutdown() - connector.shutdown() - - if role == ECConnectorRole.SCHEDULER: - scheduler.close.assert_called_once_with() - worker.close.assert_not_called() - else: - worker.close.assert_called_once_with() - scheduler.close.assert_not_called() - - def test_rejects_unknown_role(self, mock_vllm_config_producer): - invalid_role = Mock(name="invalid_role") - with pytest.raises(ValueError, match="Unknown EC connector role"): - ECMooncakeConnector(mock_vllm_config_producer, invalid_role) - - def test_del_is_best_effort(self): - connector = object.__new__(ECMooncakeConnector) - with patch.object( - ECMooncakeConnector, - "shutdown", - side_effect=RuntimeError("shutdown failed"), - ) as shutdown: - connector.__del__() - shutdown.assert_called_once_with() - - -class TestECMooncakeMetadata: - """Validate metadata compatibility, pickling, and aggregation inputs.""" - - def test_old_imports_reexport_packaged_metadata(self): - assert ECMooncakeLoadSpec is metadata.ECMooncakeLoadSpec - assert ECMooncakePushSpec is metadata.ECMooncakePushSpec - assert ECMooncakeConnectorMetadata is metadata.ECMooncakeConnectorMetadata - assert ECMooncakeWorkerMetadata is metadata.ECMooncakeWorkerMetadata - - @pytest.mark.parametrize( - "metadata", - [ - ECMooncakeConnectorMetadata( - loads=[ - ECMooncakeLoadSpec( - mm_hash="load", - num_token=2, - nbytes=8, - shape=(2, 4), - dtype="float16", - pushed=True, - transfer_id="transfer", - reservation_id="reservation", - local=True, - ) - ], - pushes=[ - ECMooncakePushSpec( - mm_hash="push", - nbytes=8, - shape=(2, 4), - dtype="float16", - consumer_zmq="tcp://127.0.0.1:1234", - transfer_id="transfer", - request_id="request", - ) - ], - ), - ECMooncakeWorkerMetadata( - loaded={"loaded"}, - failed_loads={"failed"}, - reclaimed={"reclaimed"}, - pending_loads=True, - pending_saves=True, - ), - ], - ) - def test_metadata_pickle_round_trip(self, metadata): - assert ForkingPickler.loads(ForkingPickler.dumps(metadata)) == metadata - - -class TestECMooncakeWorkerMetadataAggregation: - """Validate cross-rank success intersection and failure union rules.""" - - def test_an_item_one_rank_missed_is_not_loaded(self): - """Each rank gathers from its own cache, so all of them must have it. - - Reporting it as loaded because one rank succeeded left the scheduler - marking the hash ready while another rank raised on the cache miss. - """ - rank0 = ECMooncakeWorkerMetadata(loaded={"a", "b"}) - rank1 = ECMooncakeWorkerMetadata(loaded={"a"}, failed_loads={"b"}) - - merged = rank0.aggregate(rank1) - - assert merged.loaded == {"a"} - assert merged.failed_loads == {"b"} - - def test_a_reclaim_on_any_rank_invalidates_residency(self): - """The scheduler mirrors one pool, so the weakest rank decides.""" - merged = ECMooncakeWorkerMetadata(loaded={"a"}).aggregate( - ECMooncakeWorkerMetadata(loaded={"a"}, reclaimed={"c"}) - ) - assert merged.reclaimed == {"c"} - - -class TestSchedulerTransferTable: - """Validate Scheduler transfer transitions, indexes, and retention.""" - - @staticmethod - def pushed_spec(transfer_id: str, mm_hash: str = "hash") -> ECMooncakeLoadSpec: - return ECMooncakeLoadSpec( - mm_hash=mm_hash, - num_token=0, - nbytes=16, - shape=(4,), - dtype="float32", - pushed=True, - transfer_id=transfer_id, - reservation_id=f"reservation-{transfer_id}", - ) - - def test_legal_load_and_resident_reload_use_authoritative_record(self): - assert SchedulerTransferTable is state.SchedulerTransferTable + def test_load_completion_and_resident_reload(self): table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) record, accepted = table.observe_ready(self.pushed_spec("transfer"), 10) - assert accepted and record.state is SchedulerTransferState.AVAILABLE - assert table.begin_load("hash", 7, "transfer", "request") is record + assert accepted + assert table.begin_load("hash", "transfer", "request") is record assert table.take_loads_to_dispatch() == [record] - assert record.spec is not None and record.spec.num_token == 7 assert table.complete_load("hash") table.release_ready("hash", 1) assert record.state is SchedulerTransferState.RESIDENT - assert table.begin_load("hash", 9) is record + assert table.begin_load("hash") is record assert record.spec is not None and record.spec.local - def test_illegal_transition_is_rejected(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record, _ = table.observe_ready(self.pushed_spec("transfer"), 10) - - with pytest.raises(InvalidSchedulerTransferTransition): - table.mark_unavailable("transfer", "late", 1) - assert record.state is SchedulerTransferState.AVAILABLE - - def test_same_hash_index_preserves_transfer_order_and_identity(self): + def test_same_hash_index_preserves_order_and_identity(self): table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) first, _ = table.observe_ready(self.pushed_spec("first"), 10) second, _ = table.observe_ready(self.pushed_spec("second"), 10) @@ -1692,133 +1080,69 @@ def test_same_hash_index_preserves_transfer_order_and_identity(self): first, second, ] - assert table.begin_load("hash", 3) is first + assert table.begin_load("hash") is first assert ( table.first_for_hash("hash", (SchedulerTransferState.AVAILABLE,)) is second ) with pytest.raises(ValueError): table.observe_ready(self.pushed_spec("first", "other-hash"), 10) - def test_unavailable_notification_drains_once_and_rejects_late_ready(self): + def test_unavailable_notification_is_drained_once_and_rejects_late_ready(self): table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record = table.wait_for_event("transfer", "request-r", "hash", 1) - table.mark_unavailable("transfer", "timed out", 2) + record = table.wait_for_event("transfer", "request", "hash", 1) + table.mark_unavailable("transfer", 2) - assert table.take_unavailable_requests() == {"request-r"} - assert table.take_unavailable_requests() == set() - table.wait_for_event("transfer", "request-r", "hash", 3) + assert table.take_unavailable_requests() == {"request"} assert table.take_unavailable_requests() == set() - table.wait_for_event("transfer", "request-n", "hash", 3) - assert table.take_unavailable_requests() == {"request-n"} - assert table.take_unavailable_requests() == set() - table.wait_for_event("transfer", "request-n", "hash", 3) - assert table.take_unavailable_requests() == set() - same, accepted = table.observe_ready(self.pushed_spec("transfer"), 40) - assert same is record and not accepted + _, accepted = table.observe_ready(self.pushed_spec("transfer"), 40) + assert not accepted assert record.state is SchedulerTransferState.UNAVAILABLE def test_cancel_and_duplicate_completion_are_idempotent(self): table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - cancelled = table.wait_for_event("cancelled", "request", "hash", 10) + table.wait_for_event("cancelled", "request", "hash", 10) assert table.cancel("cancelled", 1) assert not table.cancel("cancelled", 2) _, accepted = table.observe_ready(self.pushed_spec("cancelled"), 40) - assert not accepted and cancelled.state is SchedulerTransferState.CANCELLED + assert not accepted record, _ = table.observe_ready(self.pushed_spec("completed", "other"), 10) - table.begin_load("other", 4, "completed") + table.begin_load("other", "completed") assert table.complete_load("other") assert table.complete_load("other") assert record.state is SchedulerTransferState.READY - def test_failed_record_expires_from_record_and_hash_index(self): + def test_terminal_records_expire_and_are_bounded(self): table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) record, _ = table.observe_ready(self.pushed_spec("failed"), 10) - table.begin_load("hash", 4, "failed") - - assert table.fail_load("hash", "copy failed", 20) - assert record.deadline == 50 - _, dropped = table.expire(51, terminal_limit=100) - assert dropped == 1 - assert table.get("failed") is None - assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [] - - def test_reclaimed_resident_tombstone_expires(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record, _ = table.observe_ready(self.pushed_spec("reclaimed"), 10) - table.begin_load("hash", 4, "reclaimed") - table.complete_load("hash") - table.release_ready("hash", 20) - - table.reclaim("hash", 30) - assert record.state is SchedulerTransferState.EXPIRED - assert record.deadline == 60 - table.expire(61, terminal_limit=100) - assert table.get("reclaimed") is None - assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [] - - def test_capacity_eviction_tombstone_expires(self): - table = SchedulerTransferTable(resident_capacity=0, tombstone_ttl=30) - record, _ = table.observe_ready(self.pushed_spec("evicted"), 10) - table.begin_load("hash", 4, "evicted") - table.complete_load("hash") - - table.release_ready("hash", 20) - assert record.state is SchedulerTransferState.EXPIRED + table.begin_load("hash", "failed") + table.fail_load("hash", 20) assert record.deadline == 50 table.expire(51, terminal_limit=100) - assert table.get("evicted") is None - assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [] + assert table.get("failed") is None - def test_terminal_record_limit_prunes_oldest_records(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) for transfer_id in ("first", "second", "third"): - table.cancel(transfer_id, 1) - - _, dropped = table.expire(2, terminal_limit=1) - assert dropped == 2 + table.cancel(transfer_id, 60) + table.expire(61, terminal_limit=1) assert table.get("first") is None assert table.get("second") is None assert table.get("third") is not None - def test_zero_terminal_limit_prunes_every_record(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - table.cancel("first", 1) - table.cancel("second", 1) - - _, dropped = table.expire(2, terminal_limit=0) - assert dropped == 2 - assert table.get("first") is None - assert table.get("second") is None - - def test_negative_terminal_limit_is_rejected_without_mutation(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record = table.wait_for_event("transfer", "request", "hash", 1) - - with pytest.raises(ValueError, match="terminal_limit"): - table.expire(2, terminal_limit=-1) - assert table.get("transfer") is record - assert record.state is SchedulerTransferState.WAITING_EVENT - - def test_same_hash_residency_uses_only_the_latest_completed_record(self): + def test_same_hash_keeps_only_latest_resident(self): table = SchedulerTransferTable(resident_capacity=32, tombstone_ttl=30) first, _ = table.observe_ready(self.pushed_spec("first"), 10) - table.begin_load("hash", 4, "first") + table.begin_load("hash", "first") table.complete_load("hash") table.release_ready("hash", 20) second, _ = table.observe_ready(self.pushed_spec("second"), 30) - table.begin_load("hash", 4, "second") + table.begin_load("hash", "second") table.complete_load("hash") table.release_ready("hash", 40) - third, _ = table.observe_ready(self.pushed_spec("third", "other"), 50) - table.begin_load("other", 4, "third") - table.complete_load("other") - table.release_ready("other", 60) assert first.state is SchedulerTransferState.EXPIRED assert second.state is SchedulerTransferState.RESIDENT - assert third.state is SchedulerTransferState.RESIDENT - assert table.resident_bytes == 32 + assert table.drain_orphaned() == ["first"] + assert table.drain_orphaned() == [] class TestECMooncakeSchedulerMetadata: @@ -1826,12 +1150,11 @@ class TestECMooncakeSchedulerMetadata: def test_cancel_confirms_topology_and_retries_only_failed_shards(self): scheduler = object.__new__(ECMooncakeScheduler) - scheduler._topology = Mock(spec=ShardTopology) - scheduler._topology.discover.side_effect = [ + scheduler._control_client = Mock(spec=ControlClient) + scheduler._control_client.discover_shards.side_effect = [ None, ["shard-0", "shard-1", "shard-2"], ] - scheduler._control_client = Mock(spec=ControlClient) called = [] def request(addr, _payload): @@ -1841,8 +1164,8 @@ def request(addr, _payload): return {"cancelled": True} scheduler._control_client.request.side_effect = request - assert scheduler._cancel_remote("base", "transfer", "reservation") - assert scheduler._topology.discover.call_args_list == [ + assert scheduler._cancel_remote("base", "transfer") + assert scheduler._control_client.discover_shards.call_args_list == [ call("base"), call("base"), ] @@ -1850,51 +1173,22 @@ def request(addr, _payload): def test_cancel_rejects_unconfirmed_topology_without_sending(self): scheduler = object.__new__(ECMooncakeScheduler) - scheduler._topology = Mock(spec=ShardTopology) - scheduler._topology.discover.return_value = None scheduler._control_client = Mock(spec=ControlClient) + scheduler._control_client.discover_shards.return_value = None with pytest.raises(RuntimeError, match="discover every EC consumer shard"): - scheduler._cancel_remote("base", "transfer", "reservation") + scheduler._cancel_remote("base", "transfer") - assert scheduler._topology.discover.call_count == 2 + assert scheduler._control_client.discover_shards.call_count == 2 scheduler._control_client.request.assert_not_called() - def test_missing_push_event_is_tracked( + def test_item_with_no_transfer_in_flight_is_reported_as_stalled( self, mock_vllm_config_consumer, mock_request_with_3_mm ): + """A push that never arrives must not wait silently forever.""" mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, - } - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert not scheduler.ensure_cache_available(request, 0) - assert ( - scheduler._scheduler._consumer_scheduler_metrics["missing_event"] - == 1 - ) - record = scheduler._scheduler._transfers.get(f"{request.request_id}:0") - assert record is not None - assert record.state is SchedulerTransferState.WAITING_EVENT - finally: - scheduler.shutdown() - - def test_item_with_no_transfer_in_flight_is_reported_as_stalled( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """A push that never arrives must not wait silently forever.""" - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, - "push_wait_timeout_s": 0.001, + "push_wait_timeout_s": 0.001, } request = mock_request_with_3_mm request.mm_features = request.mm_features[:1] @@ -1905,6 +1199,11 @@ def test_item_with_no_transfer_in_flight_is_reported_as_stalled( try: with ( patch.object(scheduler._scheduler, "_drain_push_notifications"), + patch.object( + scheduler._scheduler._control_client, + "request", + return_value=None, + ) as control_request, patch( "vllm.distributed.ec_transfer.ec_connector." "mooncake.scheduler.time.monotonic", @@ -1913,72 +1212,15 @@ def test_item_with_no_transfer_in_flight_is_reported_as_stalled( ): assert not scheduler.ensure_cache_available(request, 0) assert not scheduler.ensure_cache_available(request, 0) - assert ( - scheduler._scheduler._consumer_scheduler_metrics["stalled"] == 1 - ) record = scheduler._scheduler._transfers.get( f"{request.request_id}:0" ) assert record is not None assert record.state is SchedulerTransferState.UNAVAILABLE - # The stall is reported once, not once per scheduling pass. + assert scheduler.take_unavailable_requests() == {request.request_id} assert not scheduler.ensure_cache_available(request, 0) - assert scheduler._scheduler._consumer_scheduler_metrics["stalled"] == 1 - finally: - scheduler.shutdown() - - def test_pending_observation_ends_with_last_spec(self, mock_vllm_config_consumer): - specs = [ - ECMooncakeLoadSpec( - mm_hash="hash", - num_token=1, - nbytes=32, - shape=(8,), - dtype="float32", - pushed=True, - transfer_id=f"transfer-{index}", - ) - for index in range(2) - ] - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - for spec in specs: - scheduler._scheduler._transfers.observe_ready(spec, 10) - scheduler._scheduler._transfers.cancel("transfer-0", 0) - available = scheduler._scheduler._transfers.records_for_hash( - "hash", (SchedulerTransferState.AVAILABLE,) - ) - assert [record.transfer_id for record in available] == ["transfer-1"] - scheduler._scheduler._transfers.cancel("transfer-1", 0) - assert ( - scheduler._scheduler._transfers.first_for_hash( - "hash", (SchedulerTransferState.AVAILABLE,) - ) - is None - ) - finally: - scheduler.shutdown() - - def test_available_expiry_is_cancelled_before_tombstone_cleanup( - self, mock_vllm_config_consumer - ): - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - record, _ = scheduler._scheduler._transfers.observe_ready( - TestSchedulerTransferTable.pushed_spec("expired"), 0 - ) - try: - with patch.object(scheduler._scheduler, "_queue_cancel") as cancel: - scheduler._scheduler._expire_transfers() - cancel.assert_called_once_with("expired") - assert record.state is SchedulerTransferState.EXPIRED - assert scheduler._scheduler._transfers.get("expired") is record + assert scheduler.take_unavailable_requests() == set() + control_request.assert_not_called() finally: scheduler.shutdown() @@ -2006,13 +1248,10 @@ def test_local_cache_hit_keeps_the_transfer( scheduler._scheduler._transfers.observe_ready( ECMooncakeLoadSpec( mm_hash=mm_hash, - num_token=0, nbytes=16, shape=(4,), dtype="float32", - pushed=True, transfer_id="request-transfer", - reservation_id="reservation", ), 10, ) @@ -2020,14 +1259,14 @@ def test_local_cache_hit_keeps_the_transfer( patch.object(scheduler._scheduler, "_drain_push_notifications"), patch.object(scheduler._scheduler, "_queue_cancel") as cancel, ): - assert scheduler.ensure_cache_available(request, 0, {mm_hash}) + assert scheduler._ensure_cache_available(request, 0, {mm_hash}) cancel.assert_not_called() record = scheduler._scheduler._transfers.get("request-transfer") assert record is not None assert record.state is SchedulerTransferState.AVAILABLE # Once the entry is evicted the request can still get it. - assert not scheduler.ensure_cache_available(request, 0, set()) + assert not scheduler._ensure_cache_available(request, 0, set()) assert ( scheduler._scheduler._transfers.first_for_hash( mm_hash, (SchedulerTransferState.LOADING,) @@ -2060,101 +1299,23 @@ def test_consumed_item_releases_its_transfer_immediately( scheduler._scheduler._transfers.observe_ready( ECMooncakeLoadSpec( mm_hash=mm_hash, - num_token=0, nbytes=16, shape=(4,), dtype="float32", - pushed=True, transfer_id="consumed-transfer", - reservation_id="reservation", ), 10, ) - scheduler.update_state_after_free(request, 0) + with patch.object( + scheduler._scheduler, "_cancel_remote", return_value=True + ): + scheduler.update_state_after_free(request, 0) record = scheduler._scheduler._transfers.get("consumed-transfer") assert record is not None assert record.state is SchedulerTransferState.CANCELLED finally: scheduler.shutdown() - def test_ready_hash_eviction_does_not_strand_a_later_transfer( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """The long-tail stall: an event arrives while the hash is still ready. - - Dropping it as redundant left the next request with no transfer at - all once the encoder cache entry was freed, and nothing could bring - the item back. - """ - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, - } - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - mm_hash = request.mm_features[0].identifier - request.ec_transfer_params = { - "ec_items": [{"mm_hash": mm_hash, "transfer_id": "later-transfer"}] - } - event = { - "mm_hash": mm_hash, - "transfer_id": "later-transfer", - "ready": True, - "reservation_id": "later", - "nbytes": 16, - "shape": [4], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - current = ECMooncakeLoadSpec( - mm_hash=mm_hash, - num_token=0, - nbytes=16, - shape=(4,), - dtype="float32", - pushed=True, - transfer_id="current-transfer", - reservation_id="current", - ) - scheduler._scheduler._transfers.observe_ready( - current, time.monotonic() + _LEASE_TTL_SECONDS - ) - scheduler._scheduler._transfers.begin_load( - mm_hash, 4, "current-transfer" - ) - scheduler._scheduler._transfers.take_loads_to_dispatch() - scheduler._scheduler._transfers.complete_load(mm_hash) - scheduler._scheduler._event_inbox.drain = Mock(return_value=[event]) - with patch.object(scheduler._scheduler, "_queue_cancel") as cancel: - scheduler._scheduler._drain_push_notifications() - cancel.assert_not_called() - later = scheduler._scheduler._transfers.get("later-transfer") - assert later is not None - assert later.state is SchedulerTransferState.AVAILABLE - - # The scheduler frees the encoder cache entry. - scheduler.build_connector_meta( - SimpleNamespace(free_encoder_mm_hashes=[mm_hash]) - ) - assert ( - scheduler._scheduler._transfers.first_for_hash( - mm_hash, (SchedulerTransferState.READY,) - ) - is None - ) - - # The request that owns the transfer can still pick it up. - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert not scheduler.ensure_cache_available(request, 0, set()) - assert later.state is SchedulerTransferState.LOADING - finally: - scheduler.shutdown() - def test_cancelled_transfer_ignores_late_ready_events( self, mock_vllm_config_consumer ): @@ -2190,199 +1351,6 @@ def test_cancelled_transfer_ignores_late_ready_events( assert record is not None assert record.state is SchedulerTransferState.CANCELLED assert transfer_id not in scheduler._scheduler._event_ready_shards - assert scheduler._scheduler._consumer_scheduler_metrics[ - "events_cancelled" - ] == len(ports) - finally: - scheduler.shutdown() - - def test_cancel_rpc_failure_keeps_tombstone_and_rejects_late_ready( - self, mock_vllm_config_consumer - ): - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - spec = TestSchedulerTransferTable.pushed_spec("transfer") - record, _ = scheduler._scheduler._transfers.observe_ready(spec, 10) - scheduler._scheduler._transfers.cancel("transfer", 1) - failed = Mock() - failed.done.return_value = True - failed.result.side_effect = RuntimeError("unknown remote result") - scheduler._scheduler._pending_cancels["transfer"] = failed - try: - scheduler._scheduler._poll_pending_cancels() - assert record.state is SchedulerTransferState.CANCELLED - assert record.spec is spec - same, accepted = scheduler._scheduler._transfers.observe_ready(spec, 20) - assert same is record and not accepted - finally: - scheduler.shutdown() - - def test_cancel_between_shards_drops_the_partial_readiness( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """A cancel mid-aggregation leaves nothing for the late shards to finish. - - The early shards are already counted when the request releases the - item. Only clearing them keeps the remaining notifications from - completing the set and rebuilding a spec for a buffer the worker - freed as it cancelled. - """ - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - mm_hash = request.mm_features[0].identifier - transfer_id = "half-reported-transfer" - request.ec_transfer_params = { - "ec_items": [{"mm_hash": mm_hash, "transfer_id": transfer_id}] - } - event = { - "mm_hash": mm_hash, - "transfer_id": transfer_id, - "ready": True, - "reservation_id": "reservation", - "nbytes": 16, - "shape": [4], - "dtype": "float32", - } - - def deliver(scheduler, *shards): - scheduler._scheduler._event_inbox.drain.return_value = [ - {**event, "shard": shard} for shard in shards - ] - scheduler._scheduler._drain_pending = True - scheduler._scheduler._drain_push_notifications() - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - scheduler._scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" - scheduler._scheduler._event_inbox.shard_count = 4 - scheduler._scheduler._event_inbox.drain = Mock() - - deliver(scheduler, 0, 1) - assert scheduler._scheduler._event_ready_shards[transfer_id] == {0, 1} - - with patch.object( - scheduler._scheduler, "_cancel_remote", return_value=True - ): - scheduler.update_state_after_free(request, 0) - record = scheduler._scheduler._transfers.get(transfer_id) - assert record is not None - assert record.state is SchedulerTransferState.CANCELLED - assert transfer_id not in scheduler._scheduler._event_ready_shards - - deliver(scheduler, 2, 3) - - assert record.state is SchedulerTransferState.CANCELLED - assert transfer_id not in scheduler._scheduler._event_ready_shards - assert ( - scheduler._scheduler._consumer_scheduler_metrics["events_cancelled"] - == 2 - ) - finally: - scheduler.shutdown() - - def test_cancelled_transfer_ids_stay_bounded(self, mock_vllm_config_consumer): - """The ignore list is swept, not accumulated. - - It is consulted for every readiness notification and grows by one - entry per multimodal item the instance serves, so retaining ids the - worker has itself forgotten leaks for the life of the process. - """ - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - scheduler._scheduler._reservation_zmq_addr = "tcp://127.0.0.1:19101" - scheduler._scheduler._event_inbox.drain = Mock(return_value=[]) - with patch.object( - scheduler._scheduler, "_cancel_remote", return_value=True - ): - for name in ("first", "second", "third"): - scheduler._scheduler._queue_cancel(name) - - now = time.monotonic() - third = scheduler._scheduler._transfers.get("third") - assert third is not None and third.deadline is not None - assert third.deadline > now - assert third.deadline <= now + _LEASE_TTL_SECONDS - - # Ignored for exactly as long as the worker refuses to reserve - # the id again, and no longer. The drain is what sweeps. - first = scheduler._scheduler._transfers.get("first") - assert first is not None - first.deadline = 0.0 - scheduler._scheduler._drain_pending = True - scheduler._scheduler._drain_push_notifications() - assert scheduler._scheduler._transfers.get("first") is None - assert scheduler._scheduler._transfers.get("second") is not None - assert scheduler._scheduler._transfers.get("third") is third - - # The count is the backstop for a rate that outruns the TTL. - with patch( - "vllm.distributed.ec_transfer.ec_connector." - "mooncake.scheduler._MAX_TERMINAL_TRANSFER_RECORDS", - 1, - ): - scheduler._scheduler._drain_pending = True - scheduler._scheduler._drain_push_notifications() - assert scheduler._scheduler._transfers.get("second") is None - assert scheduler._scheduler._transfers.get("third") is third - assert ( - scheduler._scheduler._consumer_scheduler_metrics[ - "cancel_records_dropped" - ] - == 2 - ) - finally: - scheduler.shutdown() - - def test_item_that_never_arrives_fails_the_request( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """A push that never lands must end the request, not hold it forever. - - The failure is retryable: the caller re-issues, the encode runs again - and produces a fresh transfer. Deferring instead left the request - parked until the client timed out. - """ - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, - "push_wait_timeout_s": 0.001, - } - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - with ( - patch.object(scheduler._scheduler, "_drain_push_notifications"), - patch.object( - scheduler._scheduler._control_client, - "request", - return_value=None, - ), - patch( - "vllm.distributed.ec_transfer.ec_connector." - "mooncake.scheduler.time.monotonic", - side_effect=[10, 10.002], - ), - ): - assert not scheduler.ensure_cache_available(request, 0, set()) - assert not scheduler.ensure_cache_available(request, 0, set()) - assert scheduler.take_unavailable_requests() == {request.request_id} - # Draining clears it: the scheduler acts on each id once. - assert scheduler.take_unavailable_requests() == set() - record = scheduler._scheduler._transfers.get(f"{request.request_id}:0") - assert record is not None - assert record.state is SchedulerTransferState.UNAVAILABLE finally: scheduler.shutdown() @@ -2410,9 +1378,7 @@ def fake_send(addr: str, request: dict): mock_vllm_config_consumer, ECConnectorRole.SCHEDULER ) try: - scheduler._scheduler._reservation_zmq_addr = ( - f"tcp://127.0.0.1:{ports[0]}" - ) + scheduler._scheduler._control_addr = f"tcp://127.0.0.1:{ports[0]}" with patch.object( scheduler._scheduler._control_client, "request", @@ -2426,7 +1392,7 @@ def fake_send(addr: str, request: dict): if call.args[1]["op"] == "event_port" ] assert len(subscribed) == len(ports) - assert scheduler._scheduler._event_shard_count == len(ports) + assert scheduler._scheduler._event_inbox.shard_count == len(ports) event = {"transfer_id": "transfer-0"} assert not scheduler._scheduler._note_shard_ready( @@ -2459,8 +1425,6 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( """ mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, - "consumer_buffer_pool_size": 1 << 20, } first = mock_request_with_3_mm first.mm_features = first.mm_features[:1] @@ -2492,7 +1456,7 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( ) scheduler._scheduler._drain_push_notifications() - assert not scheduler.ensure_cache_available(first, 0, set()) + assert not scheduler._ensure_cache_available(first, 0, set()) meta = scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[]) ) @@ -2517,7 +1481,7 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( # transfer is spent. It must still be served. with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert scheduler.has_cache_item(mm_hash) - assert not scheduler.ensure_cache_available(second, 0, set()) + assert not scheduler._ensure_cache_available(second, 0, set()) assert record.state is SchedulerTransferState.LOADING reload = scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[]) @@ -2532,8 +1496,6 @@ def test_reclaimed_item_stops_being_offered_as_resident( """Residency is a mirror of the worker's pool, not a promise.""" mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, - "consumer_buffer_pool_size": 1 << 20, } request = mock_request_with_3_mm request.mm_features = request.mm_features[:1] @@ -2546,15 +1508,14 @@ def test_reclaimed_item_stops_being_offered_as_resident( try: spec = ECMooncakeLoadSpec( mm_hash=mm_hash, - num_token=0, nbytes=16, shape=(4,), dtype="float32", transfer_id="transfer", ) table = scheduler._scheduler._transfers - table.observe_ready(spec, time.monotonic() + _LEASE_TTL_SECONDS) - table.begin_load(mm_hash, 4, "transfer") + table.observe_ready(spec, time.monotonic() + _RESERVATION_TTL_SECONDS) + table.begin_load(mm_hash, "transfer") table.take_loads_to_dispatch() table.complete_load(mm_hash) table.release_ready(mm_hash, time.monotonic()) @@ -2570,46 +1531,7 @@ def test_reclaimed_item_stops_being_offered_as_resident( ) with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert not scheduler.has_cache_item(mm_hash) - assert table.resident_bytes == 0 - finally: - scheduler.shutdown() - - def test_reclaim_keeps_ready_cache_visible_until_it_is_freed( - self, mock_vllm_config_consumer - ): - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - table = scheduler._scheduler._transfers - spec = ECMooncakeLoadSpec( - mm_hash="hash", - num_token=0, - nbytes=16, - shape=(4,), - dtype="float32", - transfer_id="transfer", - ) - try: - table.observe_ready(spec, time.monotonic() + _LEASE_TTL_SECONDS) - table.begin_load("hash", 4, "transfer") - table.take_loads_to_dispatch() - table.complete_load("hash") - - scheduler.update_connector_output( - SimpleNamespace( - ec_connector_worker_meta=ECMooncakeWorkerMetadata( - reclaimed={"hash"} - ) - ) - ) - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert scheduler.has_cache_item("hash") - scheduler.build_connector_meta( - SimpleNamespace(free_encoder_mm_hashes=["hash"]) - ) - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert not scheduler.has_cache_item("hash") + assert not table.has_state(mm_hash, (SchedulerTransferState.RESIDENT,)) finally: scheduler.shutdown() @@ -2618,7 +1540,6 @@ def test_retains_new_completion_while_same_hash_is_loading( ): mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": 19019, } event = { "mm_hash": "hash", @@ -2637,21 +1558,20 @@ def test_retains_new_completion_while_same_hash_is_loading( scheduler._scheduler._event_inbox.drain = Mock(return_value=[event]) current = ECMooncakeLoadSpec( mm_hash="hash", - num_token=0, nbytes=64, shape=(2, 8), dtype="float32", transfer_id="current-transfer", ) scheduler._scheduler._transfers.observe_ready(current, time.monotonic() + 1) - scheduler._scheduler._transfers.begin_load("hash", 2, "current-transfer") + scheduler._scheduler._transfers.begin_load("hash", "current-transfer") scheduler._scheduler._drain_push_notifications() pending = scheduler._scheduler._transfers.get("next-transfer") assert pending is not None and pending.spec is not None assert pending.state is SchedulerTransferState.AVAILABLE - assert pending.spec.reservation_id == "next" + assert pending.spec.transfer_id == "next-transfer" def test_build_connector_meta_clears_pending( self, mock_vllm_config_consumer, mock_request_with_3_mm @@ -2663,43 +1583,26 @@ def test_build_connector_meta_clears_pending( mm_hash = mock_request_with_3_mm.mm_features[0].identifier load_spec = ECMooncakeLoadSpec( mm_hash=mm_hash, - num_token=0, nbytes=32, shape=(2, 4), dtype="float32", transfer_id="transfer", ) scheduler._scheduler._transfers.observe_ready( - load_spec, time.monotonic() + _LEASE_TTL_SECONDS + load_spec, time.monotonic() + _RESERVATION_TTL_SECONDS ) - scheduler._scheduler._transfers.begin_load(mm_hash, 100, "transfer") + scheduler._scheduler._transfers.begin_load(mm_hash, "transfer") meta = scheduler.build_connector_meta( Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) ) assert isinstance(meta, ECMooncakeConnectorMetadata) assert len(meta.loads) == 1 assert meta.loads[0].mm_hash == mm_hash - assert meta.loads[0].num_token == 100 assert scheduler._scheduler._transfers.take_loads_to_dispatch() == [] record = scheduler._scheduler._transfers.get("transfer") assert record is not None assert record.state is SchedulerTransferState.LOADING - def test_producer_does_not_build_load_metadata( - self, mock_vllm_config_producer, mock_request_with_3_mm - ): - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - scheduler.update_state_after_alloc(mock_request_with_3_mm, 0) - meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) - ) - - assert isinstance(meta, ECMooncakeConnectorMetadata) - assert meta.loads == [] - def test_producer_builds_push_metadata_after_preprocessing( self, mock_vllm_config_producer, mock_request_with_3_mm ): @@ -2807,8 +1710,6 @@ def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer class TestConsumerReservationManager: - """Validate Consumer destination ownership and cancellation races.""" - @staticmethod def manager(): pool = Mock() @@ -2821,383 +1722,117 @@ def manager(): @staticmethod def reserve(manager: ConsumerReservationManager): - record, write, reused, _ = manager.reserve( + record, write = manager.reserve( "transfer", "hash", 64, (16,), "float32", torch.float32 ) assert record is not None - return record, write, reused + return record, write - def test_writing_ready_and_repeated_completion_use_one_state_record(self): + def test_complete_is_idempotent(self): manager, _, _ = self.manager() - record, write, _ = self.reserve(manager) + record, write = self.reserve(manager) - assert write - assert record.state is ConsumerReservationState.WRITING - assert manager.status("transfer") is record - completed = manager.complete("transfer", record.reservation_id) - repeated = manager.complete("transfer", record.reservation_id) - - assert completed.accepted and completed.became_ready - assert repeated.accepted and repeated.repeated + assert write and record.state is ConsumerReservationState.WRITING + assert manager.complete("transfer", record.reservation_id) == (True, True) + assert manager.complete("transfer", record.reservation_id) == (True, False) assert record.state is ConsumerReservationState.READY - with pytest.raises(RuntimeError): - manager._transition(record, ConsumerReservationState.WRITING) - def test_writing_cancel_defers_the_only_allocation_release(self): + def test_cancel_waits_for_an_active_writer(self): manager, pool, allocation = self.manager() - record, _, _ = self.reserve(manager) + record, _ = self.reserve(manager) - assert manager.cancel("transfer", "wrong-id") == ( - CancellationOutcome.REJECTED, - 0, - ) - outcome, dropped = manager.cancel("transfer", record.reservation_id) - assert outcome is CancellationOutcome.DEFERRED - assert dropped == 0 + assert not manager.cancel("transfer", "wrong-id") + assert manager.cancel("transfer", record.reservation_id) assert record.state is ConsumerReservationState.CANCEL_PENDING pool.free.assert_not_called() - completed = manager.complete("transfer", record.reservation_id) - repeated = manager.complete("transfer", record.reservation_id) - assert completed.accepted and completed.discarded - assert not repeated.accepted + assert manager.complete("transfer", record.reservation_id) == (True, False) assert record.state is ConsumerReservationState.CANCELLED assert record.allocation is None pool.free.assert_called_once_with(allocation) - def test_ready_expiry_releases_once_and_keeps_a_tombstone(self): - manager, pool, allocation = self.manager() - record, _, _ = self.reserve(manager) - manager.complete("transfer", record.reservation_id) - record.expires_at = 0 - - first_expired, _, _ = manager.expire() - second_expired, _, _ = manager.expire() - - assert first_expired == 1 - assert second_expired == 0 - assert manager.status("transfer") is None - assert record.state is ConsumerReservationState.EXPIRED - pool.free.assert_called_once_with(allocation) - - def test_failed_allocation_returns_deferred_and_tombstone_counts(self): - manager, pool, _ = self.manager() - writing, _, _ = self.reserve(manager) - writing.expires_at = 1.5 - assert manager.cancel("stale", "") == ( - CancellationOutcome.PRE_RESERVED, - 0, - ) - manager.get("stale").expires_at = 0 - pool.try_allocate.side_effect = [None, None] - - monotonic = ( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "reservation.time.monotonic" - ) - with patch(monotonic, side_effect=[1.0, 2.0]): - record, write, reused, counts = manager.reserve( - "new", "new-hash", 64, (16,), "float32", torch.float32 - ) + def test_shutdown_rejects_new_reservations(self): + manager, _, _ = self.manager() + manager.begin_shutdown() - assert record is None and not write and not reused - assert counts == (0, 1, 1) - assert writing.state is ConsumerReservationState.EXPIRE_PENDING - assert manager.get("stale") is None - pool.free.assert_not_called() + with pytest.raises(RuntimeError, match="shutting down"): + self.reserve(manager) - def test_expired_writer_refresh_precedes_re_reserve_and_old_completion(self): + def test_expired_writer_is_replaced_only_after_refresh_abandon(self): manager, pool, old_allocation = self.manager() new_allocation = memory.MemoryAllocation(256, 64, torch.ones(16)) pool.try_allocate.side_effect = [old_allocation, new_allocation] - old, _, _ = self.reserve(manager) + old, _ = self.reserve(manager) old.expires_at = 0 - _, deferred, _ = manager.expire() + manager.expire() - with pytest.raises(RuntimeError, match="still has an active writer"): + with pytest.raises(RuntimeError, match="active writer"): self.reserve(manager) - assert old.allocation is old_allocation - pool.free.assert_not_called() - - (refreshed, dropped) = manager.cancel( + assert manager.cancel( "transfer", old.reservation_id, abandon=True, refresh=True ) - new, write, _ = self.reserve(manager) - late = manager.complete("transfer", old.reservation_id) + new, write = self.reserve(manager) - assert deferred == 1 - assert refreshed is CancellationOutcome.CANCELLED - assert dropped == 0 assert write and new.reservation_id != old.reservation_id - assert new.state is ConsumerReservationState.WRITING assert new.allocation is new_allocation - assert not late.accepted + assert manager.complete("transfer", old.reservation_id) == (False, False) pool.free.assert_called_once_with(old_allocation) - def test_expired_writer_single_slot_is_reused_only_after_refresh_abandon(self): - transfer_engine = Mock() - transfer_engine.register_memory.return_value = 0 - pool = ConsumerMemoryPool(256, transfer_engine) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) - manager = ConsumerReservationManager(pool, 300, 16) - - old, _, _, _ = manager.reserve( - "transfer", "hash", 64, (16,), "float32", torch.float32 - ) - assert old is not None - assert old.allocation is not None - old_offset = old.allocation.offset - old.expires_at = 0 - manager.expire() - - with pytest.raises(RuntimeError, match="still has an active writer"): - manager.reserve("transfer", "hash", 64, (16,), "float32", torch.float32) - assert pool.try_allocate(64, (16,), torch.float32) is None - - outcome, dropped = manager.cancel( - "transfer", old.reservation_id, abandon=True, refresh=True - ) - new, write, _, _ = manager.reserve( - "transfer", "hash", 64, (16,), "float32", torch.float32 - ) - assert new is not None - assert outcome is CancellationOutcome.CANCELLED - assert dropped == 0 - assert write and new.reservation_id != old.reservation_id - assert old.allocation is None - assert new.allocation is not None - assert new.allocation.offset == old_offset == 0 - - def test_expired_writer_completion_releases_before_re_reserve(self): - manager, pool, old_allocation = self.manager() - new_allocation = memory.MemoryAllocation(256, 64, torch.ones(16)) - pool.try_allocate.side_effect = [old_allocation, new_allocation] - old, _, _ = self.reserve(manager) - old.expires_at = 0 - manager.expire() - - completed = manager.complete("transfer", old.reservation_id) - new, write, _ = self.reserve(manager) - - assert completed.accepted and completed.discarded - assert old.state is ConsumerReservationState.EXPIRED - assert old.allocation is None - assert write and new.allocation is new_allocation - assert new.reservation_id != old.reservation_id - pool.free.assert_called_once_with(old_allocation) - - def test_expired_writer_cancel_stays_deferred_until_completion(self): + def test_ready_expiry_releases_once(self): manager, pool, allocation = self.manager() - record, _, _ = self.reserve(manager) + record, _ = self.reserve(manager) + manager.complete("transfer", record.reservation_id) record.expires_at = 0 - _, deferred, _ = manager.expire() - - cancelled, dropped = manager.cancel("transfer", record.reservation_id) - assert deferred == 1 - assert cancelled is CancellationOutcome.DEFERRED - assert dropped == 0 - assert record.state is ConsumerReservationState.EXPIRE_PENDING - assert record.allocation is allocation - pool.free.assert_not_called() - completed = manager.complete("transfer", record.reservation_id) - assert completed.accepted and completed.discarded + assert manager.expire() == 1 + assert manager.expire() == 0 assert record.state is ConsumerReservationState.EXPIRED pool.free.assert_called_once_with(allocation) - def test_expired_writer_abandon_releases_once(self): - manager, pool, allocation = self.manager() - record, _, _ = self.reserve(manager) - record.expires_at = 0 - manager.expire() - - abandoned, first_dropped = manager.cancel( - "transfer", record.reservation_id, abandon=True - ) - repeated, second_dropped = manager.cancel( - "transfer", record.reservation_id, abandon=True - ) - - assert abandoned is CancellationOutcome.CANCELLED - assert repeated is CancellationOutcome.PRE_RESERVED - assert first_dropped == second_dropped == 0 - assert record.state is ConsumerReservationState.CANCELLED - assert record.allocation is None - pool.free.assert_called_once_with(allocation) - - def test_cached_take_returns_the_memory_pool_canonical_allocation(self): + def test_cached_take_uses_the_pool_canonical_allocation(self): manager, pool, cached = self.manager() lease = SimpleNamespace(value=cached) canonical = memory.MemoryAllocation(256, 64, torch.ones(16)) pool.acquire_cached.return_value = lease pool.publish.return_value = canonical - record, write, _ = self.reserve(manager) + record, write = self.reserve(manager) assert not write and record.lease is lease - taken = manager.take("transfer", "hash") - - assert taken is canonical - assert record.state is ConsumerReservationState.RESIDENT - assert record.allocation is None and record.lease is None + assert manager.take("transfer", "hash") is canonical pool.publish.assert_called_once_with("hash", cached, lease) - pool.free.assert_not_called() - pool.release_cached.assert_not_called() - def test_tombstone_indexes_reap_by_prefix_without_scanning_records(self): + def test_cancel_tombstones_are_bounded(self): manager, _, _ = self.manager() - - class NoScanDict(dict): - """Fail if reservation code scans the complete record mapping.""" - - def __iter__(self): - raise AssertionError("record table must not be scanned") - - def items(self): - raise AssertionError("record table must not be scanned") - - def values(self): - raise AssertionError("record table must not be scanned") - manager._tombstone_limit = 3 - manager._records = NoScanDict(manager._records) - for transfer_id in ("a", "b", "c"): - assert manager.cancel(transfer_id, "") == ( - CancellationOutcome.PRE_RESERVED, - 0, - ) - assert manager.cancel("a", "") == (CancellationOutcome.PRE_RESERVED, 0) - assert manager.cancel("d", "") == (CancellationOutcome.PRE_RESERVED, 1) - assert list(manager._tombstones) == ["c", "a", "d"] - assert set(manager._records.keys()) == {"c", "a", "d"} - manager.get("c").expires_at = 0 - _, _, dropped = manager.expire() - assert dropped == 1 - assert list(manager._tombstones) == ["a", "d"] - assert set(manager._records.keys()) == {"a", "d"} - assert not manager._active_ids - - def test_active_index_tracks_reserve_complete_take_and_expiry(self): - manager, pool, allocation = self.manager() - pool.publish.return_value = allocation - record, _, _ = self.reserve(manager) - assert list(manager._active_ids) == ["transfer"] - assert not manager._tombstones + for transfer_id in ("a", "b", "c", "a", "d"): + assert manager.cancel(transfer_id, "") - manager.complete("transfer", record.reservation_id) - manager.take("transfer", "hash") - assert manager.get("transfer") is None - assert not manager._active_ids and not manager._tombstones - - replacement, _, _ = self.reserve(manager) - manager.complete("transfer", replacement.reservation_id) - replacement.expires_at = 0 - manager.expire() - assert not manager._active_ids - assert list(manager._tombstones) == ["transfer"] - assert manager.get("transfer").state is ConsumerReservationState.EXPIRED + assert list(manager._tombstones) == ["c", "a", "d"] + assert set(manager._records) == {"c", "a", "d"} class TestECMooncakeWorkerTransfer: """Validate end-to-end Worker reservation, push, load, and cleanup flows.""" - def test_allocation_retry_accounts_for_expiry_after_outer_sweep(self): - transfer_engine = Mock() - transfer_engine.register_memory.return_value = 0 - transfer_engine.local_session.return_value = "local-session" - pool = ConsumerMemoryPool(256, transfer_engine) - pool.prepare(torch.device("cpu"), receiving_rank=True, allow_host=True) - manager = ConsumerReservationManager(pool, 300, 16) - old, _, _, counts = manager.reserve( - "old", "old-hash", 64, (16,), "float32", torch.float32 - ) - assert counts == (0, 0, 0) - assert old is not None - assert old.allocation is not None - old_offset = old.allocation.offset - manager.complete("old", old.reservation_id) - old.expires_at = 1.5 - + def test_reservation_requires_confirmed_topology_before_any_rpc(self): worker = object.__new__(ECMooncakeWorker) - worker._consumer_worker_metrics = Counter() - worker._reservations = manager - worker._transfer = transfer_engine - payload = { - "transfer_id": "replacement", - "mm_hash": "replacement-hash", - "nbytes": 64, - "shape": [16], - "dtype": "float32", - } - monotonic = ( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "reservation.time.monotonic" + worker._control_client = Mock() + worker._control_client.discover_shards.return_value = None + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:19019", + transfer_id="transfer", ) - with patch(monotonic, side_effect=[1.0, 1.0, 2.0, 2.0]): - replacement = worker._reserve_push_destination(payload) - assert replacement["write"] - assert manager.get("old").state is ConsumerReservationState.EXPIRED - assert manager.get("replacement").allocation.offset == old_offset == 0 - assert worker._consumer_worker_metrics["reservations_expired"] == 1 - assert worker._consumer_worker_metrics["cancellations_deferred"] == 0 - assert worker._consumer_worker_metrics["cancel_records_dropped"] == 0 - - def test_failed_allocation_still_accounts_inner_expiry(self): - pool = Mock() - pool.lock = threading.RLock() - pool.acquire_cached.return_value = None - ready_allocation = memory.MemoryAllocation(0, 64, torch.empty(16)) - writing_allocation = memory.MemoryAllocation(64, 64, torch.empty(16)) - pool.try_allocate.side_effect = [ready_allocation, writing_allocation] - pool.reclaim_and_allocate.return_value = None - manager = ConsumerReservationManager(pool, 300, 16) - ready, _, _, _ = manager.reserve( - "ready", "ready-hash", 64, (16,), "float32", torch.float32 - ) - writing, _, _, _ = manager.reserve( - "writing", "writing-hash", 64, (16,), "float32", torch.float32 - ) - assert ready is not None and writing is not None - manager.complete("ready", ready.reservation_id) - assert manager.cancel("stale", "") == ( - CancellationOutcome.PRE_RESERVED, - 0, - ) - ready.expires_at = writing.expires_at = 1.5 - manager.get("stale").expires_at = 1.5 - pool.try_allocate.side_effect = [None, None] + with pytest.raises(RuntimeError, match="discover every EC consumer shard"): + worker._reserve_remote(spec) - worker = object.__new__(ECMooncakeWorker) - worker._consumer_worker_metrics = Counter() - worker._reservations = manager - payload = { - "transfer_id": "failed", - "mm_hash": "failed-hash", - "nbytes": 64, - "shape": [16], - "dtype": "float32", - } - monotonic = ( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "reservation.time.monotonic" - ) - with patch(monotonic, side_effect=[1.0, 1.0, 2.0, 2.0, 2.0]): - with pytest.raises(RuntimeError, match="^EC consumer buffer pool is full$"): - worker._reserve_push_destination(payload) - metrics = dict(worker._consumer_worker_metrics) - worker._expire_push_reservations() - - assert metrics == { - "reservations_expired": 1, - "cancellations_deferred": 1, - "cancel_records_dropped": 1, - } - assert dict(worker._consumer_worker_metrics) == metrics - assert ready.state is ConsumerReservationState.EXPIRED - assert writing.state is ConsumerReservationState.EXPIRE_PENDING - assert manager.get("stale") is None - pool.free.assert_called_once_with(ready_allocation) + assert worker._control_client.discover_shards.call_count == 2 + worker._control_client.request.assert_not_called() def test_stale_shards_are_abandoned_before_remote_re_reserve(self): worker = object.__new__(ECMooncakeWorker) @@ -3245,218 +1880,6 @@ def reserve_remote(spec): } assert events[2] == ("reserve",) - def test_cancel_retry_only_retries_failed_shards(self): - worker = object.__new__(ECMooncakeWorker) - worker._control_client = Mock() - attempts: Counter[str] = Counter() - - def request(_addr, payload): - reservation_id = payload["reservation_id"] - attempts[reservation_id] += 1 - if reservation_id == "r0" and attempts[reservation_id] == 1: - raise RuntimeError("transient cancel failure") - return {"cancelled": True} - - worker._control_client.request.side_effect = request - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="transfer", - ) - reservations = [ - { - "addr": f"tcp://consumer:{19019 + rank}", - "reservation_id": f"r{rank}", - } - for rank in range(3) - ] - - with ( - ThreadPoolExecutor(max_workers=2) as executor, - patch.object(worker, "_shard_executor", return_value=executor), - ): - worker._retry_cancel_reservations(spec, reservations) - - assert attempts == Counter({"r0": 2, "r1": 1, "r2": 1}) - - def test_partial_refresh_cleans_only_the_failed_shard_and_keeps_first_error( - self, - ): - worker = object.__new__(ECMooncakeWorker) - worker._control_client = Mock() - calls: list[tuple[str, bool]] = [] - - def request(_addr, payload): - reservation_id = payload["reservation_id"] - refreshing = payload.get("refresh", False) - calls.append((reservation_id, refreshing)) - if refreshing and reservation_id == "r0": - raise RuntimeError("refresh shard failed") - return {"cancelled": True} - - worker._control_client.request.side_effect = request - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="transfer", - ) - reservations = [ - { - "addr": f"tcp://consumer:{19019 + rank}", - "reservation_id": f"r{rank}", - "ready": False, - } - for rank in range(3) - ] - - with ( - ThreadPoolExecutor(max_workers=2) as executor, - patch.object(worker, "_shard_executor", return_value=executor), - pytest.raises(RuntimeError, match="^refresh shard failed$"), - ): - worker._refresh_remote_reservations(spec, reservations) - - assert Counter(calls) == Counter( - {("r0", True): 1, ("r1", True): 1, ("r2", True): 1, ("r0", False): 1} - ) - - def test_reservation_snapshot_and_resident_retirement_are_atomic(self): - worker = object.__new__(ECMooncakeWorker) - worker._resolve_consumer_rank = Mock() - worker._is_receiving_rank = True - worker._transfer = Mock() - worker._buffer_device = "cpu" - worker._consumer_memory = ConsumerMemoryPool(256, Mock()) - worker._reservations = ConsumerReservationManager( - worker._consumer_memory, _LEASE_TTL_SECONDS, 16 - ) - retire_entered = threading.Event() - finish_retire = threading.Event() - lock_acquired = threading.Event() - - def retire_stale(*args): - retire_entered.set() - assert finish_retire.wait(2) - - def load(): - worker.start_load_caches(ECMooncakeConnectorMetadata(), {}) - - def update_reservations(): - with worker._consumer_memory.lock: - lock_acquired.set() - - with patch.object( - worker._consumer_memory, "retire_stale", side_effect=retire_stale - ): - load_thread = threading.Thread(target=load) - load_thread.start() - assert retire_entered.wait(2) - update_thread = threading.Thread(target=update_reservations) - update_thread.start() - assert not lock_acquired.wait(0.05) - finish_retire.set() - load_thread.join(2) - update_thread.join(2) - - assert not load_thread.is_alive() - assert not update_thread.is_alive() - assert lock_acquired.is_set() - - def test_control_server_start_failure_closes_server( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 - - with ( - patch_ec_mooncake_deps(), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "worker.ConsumerControlServer" - ) as server_cls, - ): - server_cls.return_value.start.side_effect = RuntimeError("bind failed") - connector = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - with pytest.raises(RuntimeError, match="bind failed"): - connector.start_worker_services() - server_cls.return_value.close.assert_called_once_with() - assert connector._worker._control_server is None - finally: - connector.shutdown() - - def test_abandon_retries_allocation_before_reclaiming_resident( - self, mock_vllm_config_consumer - ): - config = mock_vllm_config_consumer - config.ec_transfer_config.ec_buffer_device = "cpu" - config.ec_transfer_config.ec_buffer_size = 512 - config.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 512 - - def payload(transfer_id: str, mm_hash: str) -> dict[str, object]: - return { - "transfer_id": transfer_id, - "mm_hash": mm_hash, - "nbytes": 64, - "shape": [16], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector(config, ECConnectorRole.WORKER) - worker = connector._worker - memory_pool = worker._consumer_memory - try: - memory_pool.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) - resident = memory_pool.try_allocate(64, (16,), torch.float32) - assert resident is not None - memory_pool.publish("resident", resident) - retire_event = MagicMock() - retire_event.query.return_value = True - with ( - patch.object(memory.torch, "Event", return_value=retire_event), - patch.object(memory.torch.accelerator, "current_stream"), - ): - memory_pool.retire_stale({}, set()) - old = worker._reserve_push_destination(payload("old", "old")) - try_allocate = memory_pool.try_allocate - first_attempt = True - - def abandon_between_attempts(*args): - nonlocal first_attempt - if first_attempt: - first_attempt = False - worker._cancel_push("old", old["reservation_id"], abandon=True) - return None - return try_allocate(*args) - - with patch.object( - memory_pool, - "try_allocate", - side_effect=abandon_between_attempts, - ): - new = worker._reserve_push_destination(payload("new", "new")) - - assert new["dst_ptr"] == old["dst_ptr"] - assert memory_pool.drain_reclaimed() == set() - finally: - connector.shutdown() - def test_producer_push_state_owns_source_until_every_future_is_terminal(self): manager = ProducerPushManager() reservation: Future[list[dict[str, Any]]] = Future() @@ -3480,10 +1903,10 @@ def test_producer_push_state_owns_source_until_every_future_is_terminal(self): source = torch.empty(16) manager.bind_source("hash", source, None) - assert record.source is not None + assert record.source_tensor is source reservation.set_result([]) assert manager.resolve_reservations(record) == [] - assert record.state is ProducerPushState.WAITING_SOURCE + assert record.state is ProducerPushState.WAITING_INPUTS manager.begin_writing(record) manager.begin_notifying([record]) @@ -3494,13 +1917,12 @@ def test_producer_push_state_owns_source_until_every_future_is_terminal(self): with pytest.raises(RuntimeError, match="source too early"): manager.fail([record], RuntimeError("write failed")) assert record.state is ProducerPushState.NOTIFYING - assert record.source is not None - assert record.source.tensor is source + assert record.source_tensor is source still_writing.set_result(None) manager.fail([record], RuntimeError("write failed")) assert record.state is ProducerPushState.FAILED - assert record.source is None + assert record.source_tensor is None manager.fail([record], RuntimeError("duplicate failure")) with pytest.raises(RuntimeError, match="FAILED to NOTIFYING"): manager.begin_notifying([record]) @@ -3509,7 +1931,31 @@ def test_producer_push_state_owns_source_until_every_future_is_terminal(self): assert late is record assert not late_created - def test_reservation_failure_after_source_binding_releases_the_lease(self): + def test_cancel_requests_none_selects_every_source_less_waiter(self): + manager = ProducerPushManager() + + def reserve(transfer_id: str, mm_hash: str, request_id: str): + spec = ECMooncakePushSpec( + mm_hash=mm_hash, + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:1", + transfer_id=transfer_id, + request_id=request_id, + ) + return manager.reserve(spec, lambda: Future())[0] + + first = reserve("first", "hash-first", "request-first") + second = reserve("second", "hash-second", "request-second") + assert manager.cancel_requests({"request-first"}) == [first] + assert manager.cancel_requests(None) == [second] + assert first.state is ProducerPushState.CANCEL_PENDING + assert second.state is ProducerPushState.CANCEL_PENDING + + def test_worker_close_cancels_orphaned_reservations_before_executor_shutdown( + self, + ): manager = ProducerPushManager() reservation: Future[list[dict[str, Any]]] = Future() spec = ECMooncakePushSpec( @@ -3517,236 +1963,288 @@ def test_reservation_failure_after_source_binding_releases_the_lease(self): nbytes=64, shape=(16,), dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="transfer", + consumer_zmq="tcp://consumer:19019", + transfer_id="orphan", + request_id="request", ) record, _ = manager.reserve(spec, lambda: reservation) - source = torch.empty(16) - manager.bind_source("hash", source, None) - reservation.set_exception(RuntimeError("reserve failed")) + events: list[Any] = [] - assert record.state is ProducerPushState.RESERVING + class RecordingExecutor: + def __init__(self, name: str): + self.name = name - def run(records) -> None: - try: - manager.resolve_reservations(records[0]) - except RuntimeError as exc: - manager.fail(records, exc) + def submit(self, function, *args): + events.append(f"{self.name}.submit") + future: Future[Any] = Future() + try: + result = function(*args) + except BaseException as exc: + future.set_exception(exc) + else: + future.set_result(result) + return future - with ThreadPoolExecutor(max_workers=1) as executor: - manager.submit_batches(executor, run, lambda: None) - assert manager.poll() == [("hash", "reserve failed")] - assert manager.poll() == [] - assert record.state is ProducerPushState.FAILED - assert record.source is None + def shutdown(self, wait=True, **kwargs): + events.append((f"{self.name}.shutdown", wait, kwargs)) - def test_late_reservation_callback_cannot_replace_refreshed_results( - self, mock_vllm_config_producer + worker = object.__new__(ECMooncakeWorker) + worker._producer_pushes = manager + worker._io_executor = RecordingExecutor("io") + worker._control_executor = RecordingExecutor("control") + worker._shard_pool = RecordingExecutor("shard") + worker._shutdown = False + worker._control_client = Mock() + worker._control_server = None + worker._consumer_memory = Mock() + worker._producer_memory = Mock() + worker._transfer = Mock() + worker._flush_pending_pushes = lambda: events.append("flush") + + def finish_cancel(orphan: ProducerPushRecord): + events.append(f"cancel:{orphan.spec.transfer_id}") + manager.finish_cancel(orphan) + + worker._cancel_orphaned_reservation = finish_cancel + + worker.close() + worker.close() + + assert record.state is ProducerPushState.CANCELLED + assert events == [ + "flush", + "io.submit", + "cancel:orphan", + ("io.shutdown", True, {}), + ("control.shutdown", True, {}), + ("shard.shutdown", True, {}), + ] + worker._control_client.close.assert_called_once_with() + worker._consumer_memory.close.assert_called_once_with() + worker._producer_memory.close.assert_called_once_with() + worker._transfer.close.assert_called_once_with() + + def test_worker_close_drains_an_unresolved_reservation_before_control_shutdown( + self, ): manager = ProducerPushManager() - reservation: Future[list[dict[str, Any]]] = Future() - callback_started = threading.Event() - finish_callback = threading.Event() + control_executor = ThreadPoolExecutor(max_workers=1) - def block_callback(_future) -> None: - callback_started.set() - assert finish_callback.wait(2) + def resolve_reservation(): + time.sleep(0.02) + return [ + { + "addr": "tcp://consumer:19019", + "reservation_id": "r0", + }, + { + "addr": "tcp://consumer:19020", + "reservation_id": "r1", + }, + ] - reservation.add_done_callback(block_callback) + reservation = control_executor.submit(resolve_reservation) spec = ECMooncakePushSpec( mm_hash="hash", nbytes=64, shape=(16,), dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="transfer", + consumer_zmq="tcp://consumer:19019", + transfer_id="orphan", + request_id="request", ) record, _ = manager.reserve(spec, lambda: reservation) - old = [{"addr": "old", "reservation_id": "old"}] - refreshed = [{"addr": "new", "reservation_id": "new"}] - setter = threading.Thread(target=reservation.set_result, args=(old,)) - setter.start() - assert callback_started.wait(2) - assert manager.resolve_reservations(record) == old - manager.replace_reservations(record, refreshed) - finish_callback.set() - setter.join(2) - assert not setter.is_alive() - manager.settle_all([record]) - assert record.reservations == refreshed + io_executor = ThreadPoolExecutor(max_workers=1) + shard_pool = ThreadPoolExecutor(max_workers=1) + worker = object.__new__(ECMooncakeWorker) + worker._producer_pushes = manager + worker._io_executor = io_executor + worker._control_executor = control_executor + worker._shard_pool = shard_pool + worker._shard_pool_lock = threading.Lock() + worker._shutdown = False + worker._control_client = Mock() + worker._control_client.request.return_value = {"cancelled": True} + worker._control_server = None + worker._consumer_memory = Mock() + worker._producer_memory = Mock() + worker._transfer = Mock() + worker._flush_pending_pushes = Mock() - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - try: - with patch.object( - connector._worker._control_client, - "request", - ) as request: - connector._worker._abandon_pushes([record]) - assert request.call_args.args[1]["reservation_id"] == "new" - finally: - connector.shutdown() + worker.close() + + assert record.state is ProducerPushState.CANCELLED + assert worker._control_client.request.call_count == 2 + assert all( + request.args[1]["op"] == "cancel" and request.args[1]["abandon"] + for request in worker._control_client.request.call_args_list + ) + + def test_consumer_close_waits_for_remote_writer_before_releasing_pool(self): + reservations, consumer_memory, _ = TestConsumerReservationManager.manager() + record, _ = TestConsumerReservationManager.reserve(reservations) + shutdown_started = threading.Event() + original_begin_shutdown = reservations.begin_shutdown + + def begin_shutdown(): + original_begin_shutdown() + shutdown_started.set() + + reservations.begin_shutdown = begin_shutdown + worker = object.__new__(ECMooncakeWorker) + worker._producer_pushes = ProducerPushManager() + worker._io_executor = Mock() + worker._control_executor = Mock() + worker._shard_pool = None + worker._shutdown = False + worker._control_client = Mock() + worker._control_server = Mock() + worker._consumer_memory = consumer_memory + worker._producer_memory = Mock() + worker._transfer = Mock() + worker._reservations = reservations + worker._shutdown_drain_timeout_s = 1 + worker._flush_pending_pushes = Mock() + + close_thread = threading.Thread(target=worker.close) + close_thread.start() + assert shutdown_started.wait(1) + assert record.state is ConsumerReservationState.CANCEL_PENDING + consumer_memory.close.assert_not_called() + worker._control_server.close.assert_not_called() + + assert reservations.complete("transfer", record.reservation_id) == ( + True, + False, + ) + close_thread.join(1) + + assert not close_thread.is_alive() + consumer_memory.close.assert_called_once_with() + worker._control_server.close.assert_called_once_with() + + def test_consumer_close_timeout_keeps_receive_pool_registered(self): + reservations, consumer_memory, allocation = ( + TestConsumerReservationManager.manager() + ) + record, _ = TestConsumerReservationManager.reserve(reservations) + worker = object.__new__(ECMooncakeWorker) + worker._producer_pushes = ProducerPushManager() + worker._io_executor = Mock() + worker._control_executor = Mock() + worker._shard_pool = None + worker._shutdown = False + worker._control_client = Mock() + worker._control_server = Mock() + worker._consumer_memory = consumer_memory + worker._producer_memory = Mock() + worker._transfer = Mock() + worker._reservations = reservations + worker._shutdown_drain_timeout_s = 0 + worker._flush_pending_pushes = Mock() + + worker.close() + + assert record.state is ConsumerReservationState.CANCEL_PENDING + assert record.allocation is allocation + consumer_memory.close.assert_not_called() + worker._control_server.close.assert_called_once_with() - def test_producer_hot_paths_do_not_scan_terminal_records(self): + def test_permanent_orphan_cleanup_failure_marks_push_failed(self): manager = ProducerPushManager() - request_ids = set() - limit = 4096 - with patch.object(producer, "_TERMINAL_LIMIT", limit): - for index in range(limit + 2): - reservation: Future[list[dict[str, Any]]] = Future() - reservation.set_result([]) - request_id = f"request-{index}" - request_ids.add(request_id) - spec = ECMooncakePushSpec( - mm_hash=f"hash-{index}", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id=f"transfer-{index}", - request_id=request_id, - ) - manager.reserve(spec, lambda r=reservation: r) - cancelled = manager.cancel_requests(request_ids) - assert len(cancelled) == limit + 2 - for record in cancelled: - manager.finish_cancel(record) - - pinned = manager.get("transfer-0") - assert pinned is not None - batch_started = threading.Event() - finish_batch = threading.Event() - - def block_batch(_record) -> None: - batch_started.set() - assert finish_batch.wait(2) - - executor = ThreadPoolExecutor(max_workers=1) - manager.submit_cancel(pinned, executor, block_batch) - assert batch_started.wait(2) - - class NoScanRecords(OrderedDict): - """Fail if Producer hot paths scan every transfer record.""" - - def __iter__(self): - raise AssertionError("record table scanned") - - def items(self): - raise AssertionError("record table scanned") - - def values(self): - raise AssertionError("record table scanned") - - class NoScanIndex(OrderedDict): - """Fail if Producer hot paths scan every lifecycle index.""" - - def __iter__(self): - raise AssertionError("reapable index scanned") - - def items(self): - raise AssertionError("reapable index scanned") - - def values(self): - raise AssertionError("reapable index scanned") - - manager._records = NoScanRecords(manager._records) - manager._reapable_terminal_ids = NoScanIndex(manager._reapable_terminal_ids) - assert manager.pending - manager.submit_batches(MagicMock(), MagicMock(), MagicMock()) - assert manager.poll() == [] - assert manager.get("transfer-0") is pinned - assert manager.get("transfer-1") is None - assert manager.get(f"transfer-{limit}") is not None - assert len(manager._records) == limit + 1 - - finish_batch.set() - executor.shutdown(wait=True) - assert manager.poll() == [] - assert manager.get("transfer-0") is pinned - assert manager.get("transfer-2") is None - assert len(manager._records) == limit - - def test_producer_push_cancel_handles_pending_and_late_reservations(self): + reservation: Future[list[dict[str, Any]]] = Future() + reservation.set_result( + [ + { + "addr": "tcp://consumer:19019", + "reservation_id": "reservation", + "cached": True, + } + ] + ) + spec = ECMooncakePushSpec( + mm_hash="hash", + nbytes=64, + shape=(16,), + dtype="float32", + consumer_zmq="tcp://consumer:19019", + transfer_id="transfer", + request_id="request", + ) + record, _ = manager.reserve(spec, lambda: reservation) + assert manager.cancel_requests({"request"}) == [record] + + worker = object.__new__(ECMooncakeWorker) + worker._producer_pushes = manager + worker._control_client = Mock() + worker._control_client.request.side_effect = RuntimeError("cancel failed") + with ( + ThreadPoolExecutor(max_workers=1) as executor, + patch.object(worker, "_shard_executor", return_value=executor), + ): + worker._cancel_orphaned_reservation(record) + + assert record.state is ProducerPushState.FAILED + assert record.error == "cancel failed" + assert worker._control_client.request.call_count == 2 + assert manager.poll() == [("hash", "cancel failed")] + assert manager.poll() == [] + + def test_orphan_topology_failure_marks_push_failed_without_base_cancel(self): manager = ProducerPushManager() - pending: Future[list[dict[str, Any]]] = Future() + reservation: Future[list[dict[str, Any]]] = Future() spec = ECMooncakePushSpec( mm_hash="hash", nbytes=64, shape=(16,), dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="pending", + consumer_zmq="tcp://consumer:19019", + transfer_id="transfer", request_id="request", ) - record, _ = manager.reserve(spec, lambda: pending) - assert manager.pending + record, _ = manager.reserve(spec, lambda: reservation) assert manager.cancel_requests({"request"}) == [record] - assert record.state is ProducerPushState.CANCEL_PENDING - manager.bind_source("hash", torch.empty(16), None) - assert record.source is None + reservation.set_exception(RuntimeError("topology unavailable")) - pending.set_result([]) - manager.resolve_reservations(record) - with ThreadPoolExecutor(max_workers=1) as executor: - manager.submit_cancel( - record, - executor, - lambda cancelled: manager.finish_cancel(cancelled), - ) - manager.finish_cancel(record) - manager.poll() - assert record.state is ProducerPushState.CANCELLED - assert not manager.pending - assert manager.cancel_requests({"request"}) == [] + worker = object.__new__(ECMooncakeWorker) + worker._producer_pushes = manager + worker._control_client = Mock() + worker._cancel_orphaned_reservation(record) - ready: Future[list[dict[str, Any]]] = Future() - ready.set_result([]) - later = ECMooncakePushSpec( - mm_hash="other", + assert record.state is ProducerPushState.FAILED + assert record.error == "topology unavailable" + worker._control_client.request.assert_not_called() + + def test_reservation_failure_after_source_binding_releases_the_lease(self): + manager = ProducerPushManager() + reservation: Future[list[dict[str, Any]]] = Future() + spec = ECMooncakePushSpec( + mm_hash="hash", nbytes=64, shape=(16,), dtype="float32", consumer_zmq="tcp://consumer:1", - transfer_id="ready", - request_id="request-2", + transfer_id="transfer", ) - ready_record, _ = manager.reserve(later, lambda: ready) - manager.resolve_reservations(ready_record) - assert manager.cancel_requests({"request-2"}) == [ready_record] - assert ready_record.state is ProducerPushState.CANCEL_PENDING - manager.finish_cancel(ready_record) - assert ready_record.state is ProducerPushState.CANCELLED - - def test_same_source_has_one_lease_per_transfer(self): - manager = ProducerPushManager() + record, _ = manager.reserve(spec, lambda: reservation) source = torch.empty(16) - records = [] - for transfer_id in ("first", "second"): - reservation: Future[list[dict[str, Any]]] = Future() - reservation.set_result([]) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id=transfer_id, - ) - record, _ = manager.reserve(spec, lambda r=reservation: r) - records.append(record) manager.bind_source("hash", source, None) - assert all(record.source is not None for record in records) - for record in records: - manager.resolve_reservations(record) - manager.begin_writing(record) - manager.begin_notifying([record]) - manager.complete([records[0]]) - assert records[0].source is None - assert records[1].source is not None - manager.complete([records[1]]) - assert records[1].source is None + reservation.set_exception(RuntimeError("reserve failed")) + + assert record.state is ProducerPushState.WAITING_INPUTS + + def run(records) -> None: + try: + manager.resolve_reservations(records[0]) + except RuntimeError as exc: + manager.fail(records, exc) + + with ThreadPoolExecutor(max_workers=1) as executor: + manager.submit_batches(executor, run) + assert manager.poll() == [("hash", "reserve failed")] + assert manager.poll() == [] + assert record.state is ProducerPushState.FAILED + assert record.source_tensor is None def test_shard_submit_failure_waits_before_source_release( self, mock_vllm_config_producer @@ -3794,8 +2292,8 @@ def write(session, sources, destinations, lengths): try: with ( patch.object( - worker._topology, - "shards", + worker._control_client, + "discover_shards", return_value=[ "tcp://consumer:0", "tcp://consumer:1", @@ -3804,192 +2302,58 @@ def write(session, sources, destinations, lengths): ), patch.object( worker._control_client, "request", side_effect=request - ), - patch.object(worker._producer_memory, "stage", return_value=None), - patch.object( - worker._transfer, - "acquire_sources", - return_value=[source.data_ptr()], - ), - patch.object( - worker._transfer, - "release_sources", - side_effect=lambda _: released_after_slow.append( - slow_finished.is_set() - ), - ), - patch.object(worker._transfer, "write", side_effect=write), - ): - producer.start_save_caches(encoder_cache={"hash": source}) - record = worker._producer_pushes.get("transfer") - assert record is not None - record.reservation_futures[0].result(timeout=2) - with ThreadPoolExecutor(max_workers=1) as executor: - submit_count = 0 - - def submit(fn, *args): - nonlocal submit_count - submit_count += 1 - if submit_count == 1: - return executor.submit(fn, *args) - raise RuntimeError("second shard submit failed") - - shard_executor = MagicMock() - shard_executor.submit.side_effect = submit - with patch.object( - worker, - "_shard_executor", - return_value=shard_executor, - ): - assert producer.build_connector_worker_meta().pending_saves - assert slow_started.wait(2) - assert record.source is not None - assert released_after_slow == [] - finish_slow.set() - _wait_for_worker_io(producer) - - record = worker._producer_pushes.get("transfer") - assert record is not None - assert record.state is ProducerPushState.FAILED - assert record.source is None - assert released_after_slow == [True] - finally: - finish_slow.set() - producer.shutdown() - - def test_reserve_failure_waits_for_every_started_shard( - self, mock_vllm_config_producer - ): - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="transfer", - ) - slow_started = threading.Event() - finish_slow = threading.Event() - finished = threading.Event() - errors: list[Exception] = [] - - def reserve_one(addr, _spec): - if addr.endswith(":0"): - raise RuntimeError("first shard failed") - if addr.endswith(":2"): - slow_started.set() - assert finish_slow.wait(2) - return {"addr": addr} - - with patch_ec_mooncake_deps(): - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - worker = producer._worker - - def reserve() -> None: - try: - worker._reserve_remote(spec) - except Exception as exc: - errors.append(exc) - finally: - finished.set() - - try: - with ( - patch.object( - worker._topology, - "shards", - return_value=["shard:0", "shard:1", "shard:2"], - ), - patch.object(worker, "_reserve_one", side_effect=reserve_one), - ): - thread = threading.Thread(target=reserve) - thread.start() - assert slow_started.wait(2) - assert not finished.wait(0.05) - finish_slow.set() - thread.join(2) - assert not thread.is_alive() - assert len(errors) == 1 - assert str(errors[0]) == "first shard failed" - finally: - finish_slow.set() - producer.shutdown() - - def test_reserve_submit_failure_drains_started_shards( - self, mock_vllm_config_producer - ): - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="transfer", - ) - slow_started = threading.Event() - finish_slow = threading.Event() - finished = threading.Event() - errors: list[Exception] = [] - - def reserve_one(addr, _spec): - slow_started.set() - assert finish_slow.wait(2) - return {"addr": addr} - - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - worker = connector._worker + ), + patch.object(worker._producer_memory, "stage", return_value=None), + patch.object( + worker._transfer, + "acquire_sources", + return_value=[source.data_ptr()], + ), + patch.object( + worker._transfer, + "release_sources", + side_effect=lambda _: released_after_slow.append( + slow_finished.is_set() + ), + ), + patch.object(worker._transfer, "write", side_effect=write), + ): + producer.start_save_caches(encoder_cache={"hash": source}) + record = worker._producer_pushes._records.get("transfer") + assert record is not None + record.reservation_future.result(timeout=2) + with ThreadPoolExecutor(max_workers=1) as executor: + submit_count = 0 - def reserve() -> None: - try: - worker._reserve_remote(spec) - except Exception as exc: - errors.append(exc) - finally: - finished.set() + def submit(fn, *args): + nonlocal submit_count + submit_count += 1 + if submit_count == 1: + return executor.submit(fn, *args) + raise RuntimeError("second shard submit failed") - try: - with ThreadPoolExecutor(max_workers=1) as executor: - submit_count = 0 - - def submit(fn, *args): - nonlocal submit_count - submit_count += 1 - if submit_count == 1: - return executor.submit(fn, *args) - raise RuntimeError("second shard submit failed") - - shard_executor = MagicMock() - shard_executor.submit.side_effect = submit - with ( - patch.object( - worker._topology, - "shards", - return_value=["shard:0", "shard:1", "shard:2"], - ), - patch.object(worker, "_reserve_one", side_effect=reserve_one), - patch.object( + shard_executor = MagicMock() + shard_executor.submit.side_effect = submit + with patch.object( worker, "_shard_executor", return_value=shard_executor, - ), - ): - thread = threading.Thread(target=reserve) - thread.start() - assert slow_started.wait(2) - assert not finished.wait(0.05) - finish_slow.set() - thread.join(2) - assert not thread.is_alive() - assert len(errors) == 1 - assert str(errors[0]) == "second shard submit failed" + ): + assert producer.build_connector_worker_meta().pending_saves + assert slow_started.wait(2) + assert record.source_tensor is source + assert released_after_slow == [] + finish_slow.set() + _wait_for_worker_io(producer) + + record = worker._producer_pushes._records.get("transfer") + assert record is not None + assert record.state is ProducerPushState.FAILED + assert record.source_tensor is None + assert released_after_slow == [True] finally: finish_slow.set() - connector.shutdown() + producer.shutdown() @pytest.mark.parametrize("source_before_failure", [False, True]) def test_partial_reserve_is_compensated_before_its_future_fails( @@ -4005,23 +2369,23 @@ def test_partial_reserve_is_compensated_before_its_future_fails( consumer_zmq="tcp://consumer:0", transfer_id="transfer", ) - cancel_attempts = 0 + cancel_attempts: Counter[str] = Counter() def reserve_one(addr, _spec): if addr.endswith(":1"): raise RuntimeError("reserve shard failed") return {"addr": addr, "reservation_id": "partial-r0"} - def request(_addr, payload): - nonlocal cancel_attempts - assert payload == { - "op": "cancel", - "transfer_id": "transfer", - "reservation_id": "partial-r0", - "abandon": True, - } - cancel_attempts += 1 - if cancel_attempts == 1: + def request(addr, payload): + assert payload["op"] == "cancel" and payload["abandon"] + assert payload["transfer_id"] == "transfer" + reservation_id = str(payload["reservation_id"]) + cancel_attempts[reservation_id] += 1 + if ( + addr.endswith(":0") + and reservation_id == "partial-r0" + and cancel_attempts[reservation_id] == 1 + ): raise RuntimeError("transient cleanup failure") return {"cancelled": True} @@ -4036,8 +2400,8 @@ def request(_addr, payload): try: with ( patch.object( - worker._topology, - "shards", + worker._control_client, + "discover_shards", return_value=["tcp://consumer:0", "tcp://consumer:1"], ), patch.object(worker, "_reserve_one", side_effect=reserve_one), @@ -4050,29 +2414,30 @@ def request(_addr, payload): if source_before_failure else None ) - record = worker._producer_pushes.get("transfer") + record = worker._producer_pushes._records.get("transfer") assert record is not None with pytest.raises( RuntimeError, match="^reserve shard failed$" ) as e: - record.reservation_futures[0].result(timeout=2) - assert e.value.partial_reservations == [ + record.reservation_future.result(timeout=2) + assert e.value.results == [ { "addr": "tcp://consumer:0", "reservation_id": "partial-r0", - } + }, + {"addr": "tcp://consumer:1", "reservation_id": ""}, ] - assert cancel_attempts == 2 + assert cancel_attempts == Counter({"partial-r0": 2, "": 1}) if source_before_failure: - assert record.source is not None + assert record.source_tensor is source connector.build_connector_worker_meta() _wait_for_worker_io(connector) else: assert record.state is ProducerPushState.FAILED connector.save_caches({"hash": source}, "hash") assert record.state is ProducerPushState.FAILED - assert record.source is None + assert record.source_tensor is None finally: connector.shutdown() @@ -4141,10 +2506,10 @@ def request(addr, payload): ): connector.start_save_caches(encoder_cache={"hash": source}) assert connector.build_connector_worker_meta().pending_saves - record = worker._producer_pushes.get("transfer") + record = worker._producer_pushes._records.get("transfer") assert record is not None and record.batch_future is not None assert slow_started.wait(2) - assert record.source is not None + assert record.source_tensor is source assert not record.batch_future.done() finish_slow.set() record.batch_future.result(timeout=2) @@ -4152,116 +2517,13 @@ def request(addr, payload): assert Counter(cancelled) == Counter({"r0": 1, "r1": 1}) assert record.state is ProducerPushState.FAILED - assert record.source is None + assert record.source_tensor is None assert record.error == "complete shard failed" assert all(future.done() for future in record.shard_futures) finally: finish_slow.set() connector.shutdown() - @pytest.mark.parametrize("permanent_failure", [False, True]) - def test_orphan_cancel_is_bounded_retryable_and_skips_cached_shards( - self, mock_vllm_config_producer, permanent_failure - ): - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:0", - transfer_id="transfer", - request_id="request", - ) - reservation: Future[list[dict[str, Any]]] = Future() - reservation.set_result( - [ - { - "addr": "tcp://consumer:0", - "reservation_id": "active", - }, - { - "addr": "tcp://consumer:1", - "reservation_id": "cached", - "cached": True, - }, - { - "addr": "tcp://consumer:2", - "reservation_id": "cancelled", - "cancelled": True, - }, - ] - ) - attempts: Counter[str] = Counter() - - def request(_addr, payload): - reservation_id = payload["reservation_id"] - attempts[reservation_id] += 1 - assert reservation_id == "active" - if permanent_failure or attempts[reservation_id] == 1: - raise RuntimeError("orphan cancel failed") - return {"cancelled": True} - - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - worker = connector._worker - record, _ = worker._producer_pushes.reserve(spec, lambda: reservation) - worker._producer_pushes.resolve_reservations(record) - assert worker._producer_pushes.cancel_requests({"request"}) == [record] - try: - with patch.object( - worker._control_client, "request", side_effect=request - ): - worker._producer_pushes.submit_cancel( - record, - worker._io_executor, - worker._cancel_orphaned_reservation, - ) - assert record.batch_future is not None - if permanent_failure: - with pytest.raises( - RuntimeError, match="^orphan cancel failed$" - ): - record.batch_future.result(timeout=2) - else: - record.batch_future.result(timeout=2) - failures = worker._producer_pushes.poll() - - assert attempts == Counter({"active": 2}) - assert record.state is ProducerPushState.CANCELLED - assert failures == ( - [("hash", "orphan cancel failed")] if permanent_failure else [] - ) - finally: - connector.shutdown() - - def test_source_contract_checks_shape_dtype_contiguity_and_size(self): - tensors_and_specs = [ - (torch.empty(2, 8), (16,), "float32", 64, "shape"), - (torch.empty(16, dtype=torch.float16), (16,), "float32", 32, "dtype"), - (torch.empty(4, 4).t(), (4, 4), "float32", 64, "contiguous"), - (torch.empty(16), (16,), "float32", 65, "size"), - ] - for index, (tensor, shape, dtype, nbytes, message) in enumerate( - tensors_and_specs - ): - spec = ECMooncakePushSpec( - mm_hash=f"hash-{index}", - nbytes=nbytes, - shape=shape, - dtype=dtype, - consumer_zmq="tcp://consumer:0", - transfer_id=f"transfer-{index}", - ) - reservation: Future[list[dict[str, Any]]] = Future() - reservation.set_result([]) - manager = ProducerPushManager() - record, _ = manager.reserve(spec, lambda future=reservation: future) - manager.bind_source(spec.mm_hash, tensor, None) - with pytest.raises(ValueError, match=message): - ECMooncakeWorker._validate_push_source(record) - def test_invalid_source_fails_asynchronously_before_staging( self, mock_vllm_config_producer ): @@ -4300,8 +2562,8 @@ def request(_addr, payload): try: with ( patch.object( - worker._topology, - "shards", + worker._control_client, + "discover_shards", return_value=["tcp://consumer:0"], ), patch.object( @@ -4312,13 +2574,13 @@ def request(_addr, payload): ): connector.start_save_caches(encoder_cache={"hash": source}) assert connector.build_connector_worker_meta().pending_saves - record = worker._producer_pushes.get("transfer") + record = worker._producer_pushes._records.get("transfer") assert record is not None and record.batch_future is not None record.batch_future.result(timeout=2) connector.build_connector_worker_meta() assert record.state is ProducerPushState.FAILED - assert record.source is None + assert record.source_tensor is None assert record.error == "EC source shape mismatch for mm_hash=hash" stage.assert_not_called() register.assert_not_called() @@ -4335,11 +2597,12 @@ def test_batches_pushes_from_one_model_step(self, mock_vllm_config_producer): consumer_cfg.ec_transfer_config.is_ec_consumer = True consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" + consumer_cfg.ec_transfer_config.ec_port = port consumer_cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": port, - "consumer_buffer_pool_size": 4096, } + _bind_extra_config(consumer_cfg) mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" sources = { "first": torch.randn(4, 16), @@ -4376,7 +2639,7 @@ def test_batches_pushes_from_one_model_step(self, mock_vllm_config_producer): ) assert all( reservation.state is ConsumerReservationState.READY - for reservation in consumer._worker._reservations.active_records() + for reservation in consumer._worker._reservations._records.values() ) finally: producer.shutdown() @@ -4394,11 +2657,12 @@ def test_push_reserves_before_encoder_output_is_saved( consumer_cfg.ec_transfer_config.is_ec_consumer = True consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" + consumer_cfg.ec_transfer_config.ec_port = port consumer_cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": port, - "consumer_buffer_pool_size": 4096, } + _bind_extra_config(consumer_cfg) mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" source = torch.randn(4, 16) push = ECMooncakePushSpec( @@ -4420,17 +2684,19 @@ def test_push_reserves_before_encoder_output_is_saved( producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) try: producer.start_save_caches(encoder_cache={}) - push_record = producer._worker._producer_pushes.get("transfer-1") + push_record = producer._worker._producer_pushes._records.get( + "transfer-1" + ) assert push_record is not None - reservation = push_record.reservation_futures[0] + reservation = push_record.reservation_future shards = reservation.result(timeout=2) # One reservation per consumer shard; this consumer is single. assert len(shards) == 1 reservation_data = shards[0] assert reservation_data["nbytes"] == source.nbytes old_reservation_id = reservation_data["reservation_id"] - reservation_data["_received_at"] -= _LEASE_TTL_SECONDS - consumer._worker._reservations.get("transfer-1").expires_at = 0 + reservation_data["_received_at"] -= _RESERVATION_TTL_SECONDS + consumer._worker._reservations._records["transfer-1"].expires_at = 0 with patch.object( scheduler._scheduler._control_client, "request", @@ -4449,7 +2715,9 @@ def test_push_reserves_before_encoder_output_is_saved( producer.save_caches({"hash": source}, "hash") _wait_for_worker_io(producer) assert ( - consumer._worker._reservations.get("transfer-1").reservation_id + consumer._worker._reservations._records[ + "transfer-1" + ].reservation_id != old_reservation_id ) deadline = time.monotonic() + 2 @@ -4491,11 +2759,12 @@ def test_finished_request_cancels_unbound_reservation( consumer_cfg.ec_transfer_config.is_ec_consumer = True consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" + consumer_cfg.ec_transfer_config.ec_port = port consumer_cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": port, - "consumer_buffer_pool_size": 4096, } + _bind_extra_config(consumer_cfg) mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" source = torch.randn(4, 16) push = ECMooncakePushSpec( @@ -4517,15 +2786,19 @@ def test_finished_request_cancels_unbound_reservation( producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) try: producer.start_save_caches(encoder_cache={}) - push_record = producer._worker._producer_pushes.get("transfer-1") + push_record = producer._worker._producer_pushes._records.get( + "transfer-1" + ) assert push_record is not None - reservation = push_record.reservation_futures[0] + reservation = push_record.reservation_future reservation.result(timeout=2) assert consumer._worker._reservations.status("transfer-1") producer.get_finished({"request-1"}) _wait_for_worker_io(producer) - push_record = producer._worker._producer_pushes.get("transfer-1") + push_record = producer._worker._producer_pushes._records.get( + "transfer-1" + ) assert push_record is not None assert push_record.state is ProducerPushState.CANCELLED assert consumer._worker._reservations.status("transfer-1") is None @@ -4545,11 +2818,12 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( consumer_cfg.ec_transfer_config.is_ec_consumer = True consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" + consumer_cfg.ec_transfer_config.ec_port = port consumer_cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": port, - "consumer_buffer_pool_size": 4096, } + _bind_extra_config(consumer_cfg) mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" source = torch.randn(4, 16) push = ECMooncakePushSpec( @@ -4578,16 +2852,8 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) assert engine.transfer_calls == [[source.nbytes]] - reservation = consumer._worker._reservations.get("transfer-1") + reservation = consumer._worker._reservations._records.get("transfer-1") assert reservation.state is ConsumerReservationState.READY - assert ( - consumer._worker._consumer_worker_metrics["completions_accepted"] - == 1 - ) - assert ( - consumer._worker._consumer_worker_metrics["completions_repeated"] - == 0 - ) deadline = time.monotonic() + 2 while not scheduler.has_cache_item("hash"): @@ -4596,7 +2862,6 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( record = scheduler._scheduler._transfers.get("transfer-1") assert record is not None and record.spec is not None load = record.spec - load.num_token = 4 consumer.bind_connector_metadata( ECMooncakeConnectorMetadata(loads=[load]) ) @@ -4617,14 +2882,10 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( producer.start_save_caches(encoder_cache={"hash": source}) _wait_for_worker_io(producer) assert engine.transfer_calls == [[source.nbytes]] - cached = consumer._worker._reservations.get("transfer-2") + cached = consumer._worker._reservations._records.get("transfer-2") assert cached is not None assert cached.state is ConsumerReservationState.READY assert cached.lease is not None - assert ( - consumer._worker._consumer_worker_metrics["reservations_cached"] - == 1 - ) deadline = time.monotonic() + 2 while not scheduler.has_cache_item("hash"): @@ -4633,7 +2894,6 @@ def test_duplicate_pushes_share_one_transfer_per_reservation( record = scheduler._scheduler._transfers.get("transfer-2") assert record is not None and record.spec is not None cached_load = record.spec - cached_load.num_token = 4 consumer.bind_connector_metadata( ECMooncakeConnectorMetadata(loads=[cached_load]) ) @@ -4665,26 +2925,25 @@ def test_retired_item_reserved_again_still_serves_a_local_load( cfg.ec_transfer_config.is_ec_consumer = True cfg.ec_transfer_config.ec_buffer_device = "cpu" cfg.ec_transfer_config.ec_buffer_size = 4096 + cfg.ec_transfer_config.ec_ip = "127.0.0.1" + cfg.ec_transfer_config.ec_port = port cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": port, - "consumer_buffer_pool_size": 4096, } + _bind_extra_config(cfg) spec = ECMooncakeLoadSpec( mm_hash="hash", - num_token=0, nbytes=64, shape=(4, 4), dtype="float32", + transfer_id="local-transfer", local=True, ) with patch_ec_mooncake_deps(): consumer = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) try: - consumer._worker._consumer_memory.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) + consumer._worker._consumer_memory.prepare(torch.device("cpu")) allocation = consumer._worker._consumer_memory.try_allocate( spec.nbytes, spec.shape, torch.float32 ) @@ -4709,7 +2968,7 @@ def test_retired_item_reserved_again_still_serves_a_local_load( "dtype": spec.dtype, } ) - assert consumer._worker._consumer_memory.stats()[2] == 0 + assert not consumer._worker._consumer_memory._residents._evictable assert ( consumer._worker._consumer_memory.take_resident( @@ -4720,60 +2979,6 @@ def test_retired_item_reserved_again_still_serves_a_local_load( finally: consumer.shutdown() - def test_cached_take_uses_newer_same_hash_canonical( - self, mock_vllm_config_consumer - ): - config = mock_vllm_config_consumer - config.ec_transfer_config.ec_buffer_device = "cpu" - config.ec_transfer_config.ec_buffer_size = 768 - config.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 768 - shape = (16,) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(config, ECConnectorRole.WORKER) - worker = consumer._worker - memory_pool = worker._consumer_memory - try: - memory_pool.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) - first = memory_pool.try_allocate(64, shape, torch.float32) - replacement = memory_pool.try_allocate(64, shape, torch.float32) - assert first is not None and replacement is not None - memory_pool.publish("hash", first) - worker._reserve_push_destination( - { - "transfer_id": "cached", - "mm_hash": "hash", - "nbytes": 64, - "shape": list(shape), - "dtype": "float32", - } - ) - memory_pool.publish("hash", replacement) - memory_pool.retire_stale({}, {"hash"}) - spec = ECMooncakeLoadSpec( - mm_hash="hash", - num_token=1, - nbytes=64, - shape=shape, - dtype="float32", - pushed=True, - transfer_id="cached", - ) - - tensor, allocation = worker._take_pushed_tensor(spec) - - assert allocation is replacement - assert tensor is replacement.tensor - reused = memory_pool.try_allocate(64, shape, torch.float32) - assert reused is not None - assert reused.offset == first.offset - finally: - consumer.shutdown() - def test_push_reaches_every_consumer_shard(self, mock_vllm_config_producer): """A sharded consumer gets one copy per rank, from one source. @@ -4849,11 +3054,12 @@ def test_pushes_stage_through_the_registered_pool(self, mock_vllm_config_produce consumer_cfg.ec_transfer_config.is_ec_consumer = True consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" + consumer_cfg.ec_transfer_config.ec_port = port consumer_cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": port, - "consumer_buffer_pool_size": 4096, } + _bind_extra_config(consumer_cfg) mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" source = torch.randn(4, 16) pushes = [ @@ -4883,14 +3089,14 @@ def test_pushes_stage_through_the_registered_pool(self, mock_vllm_config_produce assert isinstance(engine, CopyingFakeTransferEngine) # The staging pool is registered once; a transfer registers # nothing of its own. - pool = producer._worker._producer_memory.tensor + pool = producer._worker._producer_memory._pool assert pool is not None assert engine.register_calls == [[pool.data_ptr()]] assert engine.batch_unregister_calls == [] assert engine.transfer_calls == [[source.nbytes, source.nbytes]] assert all( reservation.state is ConsumerReservationState.READY - for reservation in consumer._worker._reservations.active_records() + for reservation in consumer._worker._reservations._records.values() ) finally: producer.shutdown() @@ -4929,7 +3135,7 @@ def test_push_falls_back_to_per_tensor_registration_without_a_pool( _wait_for_worker_io(producer) engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) - assert producer._worker._producer_memory.tensor is None + assert producer._worker._producer_memory._pool is None assert engine.register_calls == [[source.data_ptr()]] assert engine.batch_unregister_calls == [[source.data_ptr()]] assert engine.transfer_calls == [[source.nbytes]] @@ -4971,11 +3177,12 @@ def _push_harness_config(self, producer_cfg, port: int): consumer_cfg.ec_transfer_config.is_ec_consumer = True consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 + consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" + consumer_cfg.ec_transfer_config.ec_port = port consumer_cfg.ec_transfer_config.ec_connector_extra_config = { "mooncake_protocol": "tcp", - "reservation_zmq_port": port, - "consumer_buffer_pool_size": 4096, } + _bind_extra_config(consumer_cfg) producer_cfg.ec_transfer_config.ec_buffer_device = "cpu" return consumer_cfg @@ -5018,7 +3225,7 @@ def test_batch_completion_sends_one_control_message( assert "complete" not in ops assert all( reservation.state is ConsumerReservationState.READY - for reservation in consumer._worker._reservations.active_records() + for reservation in consumer._worker._reservations._records.values() ) finally: producer.shutdown() @@ -5063,18 +3270,13 @@ def test_complete_is_idempotent_without_republishing( ): mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 with patch_ec_mooncake_deps(): consumer = ECMooncakeConnector( mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._worker._consumer_memory.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) + consumer._worker._consumer_memory.prepare(torch.device("cpu")) reservation = consumer._worker._reserve_push_destination( { "mm_hash": "hash", @@ -5086,110 +3288,15 @@ def test_complete_is_idempotent_without_republishing( ) reservation_id = reservation["reservation_id"] - first = consumer._worker._complete_push("transfer-1", reservation_id) - repeated = consumer._worker._complete_push("transfer-1", reservation_id) - - assert first.accepted and first.became_ready - assert repeated.accepted and not repeated.became_ready - finally: - consumer.shutdown() - - def test_cancel_pending_repeat_reserve_is_terminal_without_releasing( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 - payload = { - "mm_hash": "hash", - "transfer_id": "transfer", - "nbytes": 64, - "shape": [4, 4], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - memory_pool = consumer._worker._consumer_memory - try: - memory_pool.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) - first = consumer._worker._reserve_push_destination(payload) - record = consumer._worker._reservations.get("transfer") - assert record is not None and record.allocation is not None - allocation = record.allocation - with patch.object(memory_pool, "free", wraps=memory_pool.free) as free: - assert consumer._worker._cancel_push( - "transfer", first["reservation_id"] - ) - repeated = consumer._worker._reserve_push_destination(payload) - - assert repeated["cancelled"] - assert not repeated["write"] and not repeated["ready"] - assert record.state is ConsumerReservationState.CANCEL_PENDING - assert record.allocation is allocation - free.assert_not_called() - - completed = consumer._worker._complete_push( - "transfer", first["reservation_id"] - ) - assert completed.accepted and not completed.became_ready - assert record.state is ConsumerReservationState.CANCELLED - assert record.allocation is None - free.assert_called_once_with(allocation) - finally: - consumer.shutdown() - - def test_same_hash_transfers_have_independent_lifecycles( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 - - def payload(transfer_id: str) -> dict: - return { - "mm_hash": "shared-hash", - "transfer_id": transfer_id, - "nbytes": 64, - "shape": [4, 4], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - consumer._worker._consumer_memory.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) - first = consumer._worker._reserve_push_destination(payload("first")) - second = consumer._worker._reserve_push_destination(payload("second")) - - consumer._worker._complete_push("first", first["reservation_id"]) - assert ( - consumer._worker._reservations.get("first").state - is ConsumerReservationState.READY + first = consumer._worker._reservations.complete( + "transfer-1", reservation_id ) - assert ( - consumer._worker._reservations.get("second").state - is ConsumerReservationState.WRITING + repeated = consumer._worker._reservations.complete( + "transfer-1", reservation_id ) - assert consumer._worker._cancel_push("first", first["reservation_id"]) - assert consumer._worker._reservations.status("first") is None - assert consumer._worker._reservations.status("second") - assert consumer._worker._complete_push( - "second", second["reservation_id"] - ) + assert first == (True, True) + assert repeated == (True, False) finally: consumer.shutdown() @@ -5198,9 +3305,6 @@ def test_late_completion_cannot_complete_new_reservation( ): mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 payload = { "mm_hash": "hash", "transfer_id": "transfer", @@ -5214,149 +3318,36 @@ def test_late_completion_cannot_complete_new_reservation( mock_vllm_config_consumer, ECConnectorRole.WORKER ) try: - consumer._worker._consumer_memory.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) + consumer._worker._consumer_memory.prepare(torch.device("cpu")) old = consumer._worker._reserve_push_destination(payload) - consumer._worker._reservations.get("transfer").expires_at = 0 - consumer._worker._expire_push_reservations() + consumer._worker._reservations._records["transfer"].expires_at = 0 + consumer._worker._reservations.expire() assert ( - consumer._worker._reservations.get("transfer").state + consumer._worker._reservations._records["transfer"].state is ConsumerReservationState.EXPIRE_PENDING ) - assert consumer._worker._cancel_push( + assert consumer._worker._reservations.cancel( "transfer", old["reservation_id"], abandon=True, refresh=True, ) new = consumer._worker._reserve_push_destination(payload) - new_record = consumer._worker._reservations.get("transfer") + new_record = consumer._worker._reservations._records.get("transfer") assert new_record is not None and new_record.allocation is not None new_allocation = new_record.allocation assert old["reservation_id"] != new["reservation_id"] - stale = consumer._worker._complete_push( + stale = consumer._worker._reservations.complete( "transfer", old["reservation_id"] ) - assert not stale.accepted - assert consumer._worker._reservations.get("transfer") is new_record - assert new_record.allocation is new_allocation - assert new_record.state is ConsumerReservationState.WRITING - finally: - consumer.shutdown() - - def test_ready_reservation_has_a_terminal_expiry(self, mock_vllm_config_consumer): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - consumer._worker._consumer_memory.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) - reservation = consumer._worker._reserve_push_destination( - { - "mm_hash": "hash", - "transfer_id": "transfer-1", - "nbytes": 64, - "shape": [4, 4], - "dtype": "float32", - } - ) - consumer._worker._complete_push( - "transfer-1", reservation["reservation_id"] - ) - consumer._worker._reservations.get("transfer-1").expires_at = 0 - - assert consumer._worker._expire_push_reservations() == 1 - assert consumer._worker._reservations.status("transfer-1") is None - finally: - consumer.shutdown() - - def test_cancel_before_reserve_creates_bounded_tombstone( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 - payload = { - "mm_hash": "hash", - "transfer_id": "cancelled-transfer", - "nbytes": 64, - "shape": [4, 4], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - consumer._worker._consumer_memory.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) - assert consumer._worker._cancel_push("cancelled-transfer", "") - cancelled = consumer._worker._reserve_push_destination(payload) - assert cancelled["cancelled"] and not cancelled["write"] - assert ( - consumer._worker._reservations.status("cancelled-transfer") is None - ) - - consumer._worker._reservations.get("cancelled-transfer").expires_at = 0 - consumer._worker._expire_push_reservations() - replacement = consumer._worker._reserve_push_destination(payload) - assert replacement["write"] - finally: - consumer.shutdown() - - def test_repeated_cancel_does_not_strand_older_tombstones( - self, mock_vllm_config_consumer - ): - """Re-cancelling refreshes a tombstone without breaking the sweep order. - - The sweep stops at the first live record, so a refreshed one that kept - its original position would shield every older record behind it and - the table would grow for the life of the process. - """ - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config[ - "consumer_buffer_pool_size" - ] = 4096 - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - consumer._worker._consumer_memory.prepare( - torch.device("cpu"), receiving_rank=True, allow_host=True - ) - assert consumer._worker._cancel_push("refreshed-transfer", "") - assert consumer._worker._cancel_push("stale-transfer", "") - consumer._worker._reservations.get("stale-transfer").expires_at = 0.0 - assert consumer._worker._cancel_push("refreshed-transfer", "") - - consumer._worker._expire_push_reservations() - - assert consumer._worker._reservations.get("stale-transfer") is None - assert ( - consumer._worker._reservations.get("refreshed-transfer").state - is ConsumerReservationState.CANCELLED - ) + assert stale == (False, False) assert ( - consumer._worker._consumer_worker_metrics["cancel_records_dropped"] - == 1 + consumer._worker._reservations._records.get("transfer") + is new_record ) + assert new_record.allocation is new_allocation + assert new_record.state is ConsumerReservationState.WRITING finally: consumer.shutdown() @@ -5366,13 +3357,10 @@ def test_missing_push_reservation_reports_failed_load( mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" spec = ECMooncakeLoadSpec( mm_hash="hash", - num_token=1, nbytes=32, shape=(8,), dtype="float32", - pushed=True, transfer_id="missing-transfer", - reservation_id="missing-reservation", ) with patch_ec_mooncake_deps(): @@ -5390,13 +3378,3 @@ def test_missing_push_reservation_reports_failed_load( assert cache == {} finally: consumer.shutdown() - - def test_producer_scheduler_has_cache_item_false( - self, mock_vllm_config_producer, mock_request_with_3_mm - ): - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - mm_hash = mock_request_with_3_mm.mm_features[0].identifier - assert not scheduler.has_cache_item(mm_hash) diff --git a/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py b/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py index 48123235a72a..c7278ad55bdf 100644 --- a/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py +++ b/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py @@ -32,8 +32,13 @@ def _load_proxy_module(): @pytest.fixture(scope="module") -def assign(): - return _load_proxy_module().encoder_rr_assignment +def proxy(): + return _load_proxy_module() + + +@pytest.fixture(scope="module") +def assign(proxy): + return proxy.encoder_rr_assignment def _drive(assign, e_urls, counts): @@ -80,3 +85,44 @@ def test_single_encoder_always_resolves_to_it(assign): urls, next_cursor = assign(["E0"], 0, 4) assert urls == ["E0"] * 4 assert next_cursor == 0 + + +def test_generic_e_p_d_routing_without_mooncake_is_supported(proxy): + proxy.validate_ec_consumer_routing(["http://prefill"], []) + + +def test_mooncake_e_pd_routing_is_supported(proxy): + proxy.validate_ec_consumer_routing([], ["tcp://decode:19019"]) + + +def test_mooncake_independent_prefill_routing_fails_fast(proxy): + with pytest.raises(ValueError, match=r"supports E\+PD only"): + proxy.validate_ec_consumer_routing(["http://prefill"], ["tcp://decode:19019"]) + + +def test_decode_rewrite_preserves_engine_reported_ec_hash(proxy): + request = { + "messages": [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "image"}}], + } + ] + } + + rewritten = proxy.rewrite_for_decode( + request, + { + 0: { + "mm_hash": "proxy-uuid", + "ec_mm_hash": "engine-derived-hash", + "transfer_id": "transfer", + "image_grid_thw": [1, 2, 3], + } + }, + ) + + assert rewritten["messages"][0]["content"][0]["uuid"] == "proxy-uuid" + assert rewritten["ec_transfer_params"]["ec_items"] == [ + {"mm_hash": "engine-derived-hash", "transfer_id": "transfer"} + ] diff --git a/vllm/distributed/ec_transfer/ec_connector/base.py b/vllm/distributed/ec_transfer/ec_connector/base.py index 5a526def1487..196c4f32bc1d 100644 --- a/vllm/distributed/ec_transfer/ec_connector/base.py +++ b/vllm/distributed/ec_transfer/ec_connector/base.py @@ -268,7 +268,6 @@ def ensure_cache_available( self, request: "Request", num_computed_tokens: int, - local_cache_hashes: Collection[str] | None = None, ) -> bool: """ Ensure encoder cache items are available for the given request. @@ -277,14 +276,21 @@ def ensure_cache_available( Args: request: the request whose multimodal features to check. num_computed_tokens: tokens already covered by cached KV blocks. - local_cache_hashes: encoder outputs already cached locally. - Returns: True if all items are ready or no transfer is needed. False if any items are still in transit (request should be deferred). """ return True + def _ensure_cache_available( + self, + request: "Request", + num_computed_tokens: int, + local_cache_hashes: Collection[str], + ) -> bool: + """Core-only adapter that preserves the connector extension API.""" + return self.ensure_cache_available(request, num_computed_tokens) + @abstractmethod def update_state_after_alloc(self, request: "Request", index: int): """ diff --git a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py index aab35b2136a5..0de266be8eff 100644 --- a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py @@ -8,7 +8,6 @@ offloaded to CPU instead of recomputing them. """ -from collections.abc import Collection from typing import TYPE_CHECKING import torch @@ -95,11 +94,10 @@ def ensure_cache_available( self, request: "Request", num_computed_tokens: int, - local_cache_hashes: Collection[str] | None = None, ) -> bool: assert self.connector_scheduler is not None return self.connector_scheduler.ensure_cache_available( - request, num_computed_tokens, local_cache_hashes + request, num_computed_tokens ) def update_state_after_alloc(self, request: "Request", index: int) -> None: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py deleted file mode 100644 index c6f9546e9a4d..000000000000 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/_availability.py +++ /dev/null @@ -1,28 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Guard access to the optional Mooncake TransferEngine dependency. - -Keeping the import check here lets metadata and configuration modules remain -importable in environments that do not install Mooncake. -""" - -_MOONCAKE_IMPORT_ERROR: ImportError | None -try: - from mooncake.engine import TransferEngine as _TransferEngine # noqa: F401 -except ImportError as e: - _MOONCAKE_IMPORT_ERROR = e -else: - _MOONCAKE_IMPORT_ERROR = None - - -def ensure_mooncake_available() -> None: - """Raise a user-facing error when Mooncake is unavailable. - - Raises: - ImportError: If ``mooncake-transfer-engine`` cannot be imported. - """ - if _MOONCAKE_IMPORT_ERROR is not None: - raise ImportError( - "Install mooncake-transfer-engine (see " - "https://github.com/kvcache-ai/Mooncake ) to use ECMooncakeConnector." - ) from _MOONCAKE_IMPORT_ERROR diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py index 1d1ece14d424..e0bf88bb685e 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py @@ -1,123 +1,78 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Parse and validate configuration shared by Mooncake connector roles.""" - from __future__ import annotations import math from dataclasses import dataclass from typing import TYPE_CHECKING -from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole +from vllm.utils.network_utils import make_zmq_path if TYPE_CHECKING: from vllm.config import VllmConfig +# Scheduler and Worker must agree on when a remote reservation becomes stale. +_RESERVATION_TTL_SECONDS = 300 + -def _integer(name: str, value: object) -> int: - message = f"ECMooncakeConnector requires {name} to be an integer." - if isinstance(value, int) and not isinstance(value, bool): - return value - if isinstance(value, float) and value.is_integer(): - return int(value) - if isinstance(value, str): +def _positive_int(name: str, value: object) -> int: + message = f"ECMooncakeConnector requires {name} to be a positive integer." + if isinstance(value, bool): + raise ValueError(message) + if isinstance(value, int): + result = value + elif isinstance(value, float) and math.isfinite(value) and value.is_integer(): + result = int(value) + elif isinstance(value, str): try: - return int(value) + result = int(value) except ValueError as error: raise ValueError(message) from error - raise ValueError(message) - - -def _positive_integer(name: str, value: object) -> int: - parsed = _integer(name, value) - if parsed <= 0: + else: + raise ValueError(message) + if result <= 0: raise ValueError(f"ECMooncakeConnector requires {name} > 0.") - return parsed + return result -def _finite_float(name: str, value: object, allow_zero: bool) -> float: - requirement = ">= 0" if allow_zero else "> 0" - message = f"ECMooncakeConnector requires {name} {requirement}." - if isinstance(value, bool) or not isinstance(value, (str, int, float)): +def _positive_float(name: str, value: object) -> float: + message = f"ECMooncakeConnector requires {name} > 0." + if isinstance(value, bool): raise ValueError(message) try: - parsed = float(value) - except ValueError as error: + result = float(value) # type: ignore[arg-type] + except (TypeError, ValueError) as error: raise ValueError(message) from error - if not math.isfinite(parsed) or parsed < 0 or (parsed == 0 and not allow_zero): + if not math.isfinite(result) or result <= 0: raise ValueError(message) - return parsed - - -def _nonempty_string(name: str, value: object) -> str: - if not isinstance(value, str) or not value.strip(): - raise ValueError(f"ECMooncakeConnector requires non-empty {name}.") - return value.strip() + return result @dataclass(frozen=True) class MooncakeECConfig: - """Validated runtime settings for one Scheduler or Worker instance. - - Attributes: - is_producer: Whether this instance can originate encoder-cache pushes. - is_consumer: Whether this instance can receive encoder-cache pushes. - protocol: Mooncake transport protocol passed to ``TransferEngine``. - buffer_device: Device used for registered staging and receive buffers. - reservation_port: Rank-adjusted Worker control-plane base port. - reservation_addr: Scheduler-visible Consumer control address. - control_timeout_s: Timeout for one ZMQ request/response exchange. - push_wait_timeout_s: Maximum Scheduler wait for a ready notification. - transfer_workers: Maximum concurrent data-plane transfer batches. - control_workers: Maximum concurrent control-plane operations. - producer_pool_size: Bytes reserved for the Producer staging pool. - consumer_pool_size: Bytes reserved for the Consumer receive pool. - transfer_metrics_log_interval: Producer transfer log interval in seconds. - consumer_metrics_log_interval: Consumer metrics log interval in seconds. + """Validated settings shared by the Scheduler and Worker roles. + + ``control_port`` is the first TP-shard port after the DP offset; + ``control_addr`` targets that shard, which advertises the full topology. """ is_producer: bool is_consumer: bool protocol: str buffer_device: str - reservation_port: int | None - reservation_addr: str | None - control_timeout_s: float + control_port: int + control_addr: str + control_timeout_ms: int push_wait_timeout_s: float - transfer_workers: int - control_workers: int - producer_pool_size: int - consumer_pool_size: int - transfer_metrics_log_interval: float - consumer_metrics_log_interval: float - - @property - def control_timeout_ms(self) -> int: - return max(1, math.ceil(self.control_timeout_s * 1000)) + pool_size: int @classmethod - def from_vllm_config( - cls, vllm_config: VllmConfig, role: ECConnectorRole - ) -> MooncakeECConfig: - """Build role-specific settings from the top-level vLLM config. - - Args: - vllm_config: Source vLLM configuration. - role: Connector process role being configured. - - Returns: - Validated, normalized Mooncake connector settings. - - Raises: - ValueError: If an option is invalid or the requested parallel - topology is unsupported for a Producer. - """ + def from_vllm_config(cls, vllm_config: VllmConfig) -> MooncakeECConfig: parallel_config = vllm_config.parallel_config ec_config = vllm_config.ec_transfer_config assert ec_config is not None - is_producer = ec_config.is_ec_producer - if is_producer: + if ec_config.is_ec_producer: if parallel_config.tensor_parallel_size > 1: raise ValueError( "ECMooncakeConnector producers require tensor_parallel_size=1." @@ -131,102 +86,33 @@ def from_vllm_config( "ECMooncakeConnector producers require data_parallel_size=1." ) - registered_buffer_size = _positive_integer( + registered_buffer_size = _positive_int( "ec_buffer_size", ec_config.ec_buffer_size ) - - extra = ec_config.ec_connector_extra_config - raw_port = extra.get("reservation_zmq_port") - reservation_port = ( - _integer("reservation_zmq_port", raw_port) if raw_port is not None else None - ) - if reservation_port is not None and not 1 <= reservation_port <= 65535: - raise ValueError( - "ECMooncakeConnector requires reservation_zmq_port in 1..65535." - ) - - if reservation_port is not None: - reservation_port += ( - parallel_config.data_parallel_index - * parallel_config.tensor_parallel_size - ) - highest_port = reservation_port + parallel_config.tensor_parallel_size - 1 - if not 1 <= reservation_port <= highest_port <= 65535: - raise ValueError( - "ECMooncakeConnector reservation ports must be in 1..65535." - ) - - reservation_addr = ( - _nonempty_string("reservation_zmq_addr", extra["reservation_zmq_addr"]) - if "reservation_zmq_addr" in extra - else None - ) - if reservation_addr is None and reservation_port is not None: - reservation_addr = f"tcp://127.0.0.1:{reservation_port}" - - is_consumer = ec_config.is_ec_consumer - if is_consumer and role == ECConnectorRole.SCHEDULER and not reservation_addr: - raise ValueError( - "ec_consumer with ECMooncakeConnector requires " - "reservation_zmq_port or reservation_zmq_addr." - ) - if is_consumer and role == ECConnectorRole.WORKER and reservation_port is None: - raise ValueError( - "ec_consumer with ECMooncakeConnector workers require " - "reservation_zmq_port." - ) - - control_timeout_s = _finite_float( - "control_timeout_s", extra.get("control_timeout_s", 30), False - ) - if control_timeout_s > (2**31 - 1) / 1000: - raise ValueError("ECMooncakeConnector control_timeout_s is too large.") - push_wait_timeout_s = _finite_float( - "push_wait_timeout_s", extra.get("push_wait_timeout_s", 60), False - ) - transfer_workers = _positive_integer( - "transfer_max_workers", extra.get("transfer_max_workers", 4) - ) - control_workers = _positive_integer( - "control_max_workers", extra.get("control_max_workers", 8) - ) - producer_pool_size = _positive_integer( - "producer_buffer_pool_size", - extra.get("producer_buffer_pool_size", registered_buffer_size), - ) - consumer_pool_size = _positive_integer( - "consumer_buffer_pool_size", - extra.get("consumer_buffer_pool_size", registered_buffer_size), - ) - - protocol = _nonempty_string( - "mooncake_protocol", extra.get("mooncake_protocol", "rdma") + get = ec_config.get_from_extra_config + control_port = int(ec_config.ec_port) + ( + parallel_config.data_parallel_index * parallel_config.tensor_parallel_size ) - raw_buffer_device = ec_config.ec_buffer_device - if raw_buffer_device is not None and not isinstance(raw_buffer_device, str): - raise ValueError("ECMooncakeConnector ec_buffer_device must be a string.") + highest_port = control_port + parallel_config.tensor_parallel_size - 1 + if not 1 <= control_port <= highest_port <= 65535: + raise ValueError("ECMooncakeConnector ec_port must be in 1..65535.") return cls( - is_producer=is_producer, - is_consumer=is_consumer, - protocol=protocol, - buffer_device=(raw_buffer_device or "cuda").strip() or "cuda", - reservation_port=reservation_port, - reservation_addr=reservation_addr, - control_timeout_s=control_timeout_s, - push_wait_timeout_s=push_wait_timeout_s, - transfer_workers=transfer_workers, - control_workers=control_workers, - producer_pool_size=producer_pool_size, - consumer_pool_size=consumer_pool_size, - transfer_metrics_log_interval=_finite_float( - "transfer_metrics_log_interval", - extra.get("transfer_metrics_log_interval", 10), - True, + is_producer=ec_config.is_ec_producer, + is_consumer=ec_config.is_ec_consumer, + protocol=str(get("mooncake_protocol", "rdma")), + buffer_device=str(ec_config.ec_buffer_device or "cuda").lower(), + control_port=control_port, + control_addr=make_zmq_path("tcp", ec_config.ec_ip, control_port), + control_timeout_ms=max( + 1, + math.ceil( + _positive_float("control_timeout_s", get("control_timeout_s", 30)) + * 1000 + ), ), - consumer_metrics_log_interval=_finite_float( - "consumer_metrics_log_interval", - extra.get("consumer_metrics_log_interval", 10), - True, + push_wait_timeout_s=_positive_float( + "push_wait_timeout_s", get("push_wait_timeout_s", 60) ), + pool_size=registered_buffer_size, ) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py index 260b21d89295..daeb7646a707 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py @@ -11,14 +11,12 @@ import threading import time -from collections import Counter, deque +from collections import deque from collections.abc import Callable -from dataclasses import dataclass -from typing import Any, Literal, TypedDict, cast +from typing import Any import torch import zmq -from typing_extensions import NotRequired from vllm.logger import init_logger @@ -28,112 +26,18 @@ _RESERVATION_REAP_INTERVAL_SECONDS = 1 -class ReservationItem(TypedDict): - """Identify one Consumer reservation in a batch control request.""" - - transfer_id: str - reservation_id: str - - -class PeersRequest(TypedDict): - """Request the control ports of every Consumer TP shard.""" - - op: Literal["peers"] - - -class EventPortRequest(TypedDict): - """Request the PUSH socket port used for readiness events.""" - - op: Literal["event_port"] - - -class StatusRequest(TypedDict): - """Request the current status of a transfer reservation.""" - - op: Literal["status"] - transfer_id: str - - -class ReserveRequest(TypedDict): - """Request destination memory for an encoder-cache tensor.""" - - op: Literal["reserve"] - transfer_id: str - mm_hash: str - nbytes: int - shape: list[int] - dtype: str - - -class CompleteBatchRequest(TypedDict): - """Mark several destination writes complete in one exchange.""" - - op: Literal["complete_batch"] - items: list[ReservationItem] - - -class ReservationActionRequest(ReservationItem): - """Complete or cancel one previously created reservation.""" - - op: Literal["complete", "cancel"] - abandon: NotRequired[bool] - refresh: NotRequired[bool] - - -ControlRequest = ( - PeersRequest - | EventPortRequest - | StatusRequest - | ReserveRequest - | ReservationActionRequest - | CompleteBatchRequest -) - - -class ControlSuccess(TypedDict): - """Successful wire response with an optional operation result.""" - - ok: Literal[True] - result: NotRequired[Any] - - -class ControlFailure(TypedDict): - """Failed wire response containing a user-facing error message.""" - - ok: Literal[False] - error: str - - -ControlResponse = ControlSuccess | ControlFailure - - -@dataclass(frozen=True) -class ControlCompletion: - """Summarize the effect of a Consumer completion request. - - Attributes: - accepted: Whether the reservation identity was valid. - became_ready: Whether this call newly made the tensor readable. - """ - - accepted: bool - became_ready: bool = False +ControlRequest = dict[str, Any] +ControlResponse = dict[str, Any] class ControlClient: - """Send control requests through reusable, thread-local REQ sockets. - - Attributes: - _context: ZMQ context that owns all client sockets. - _timeout_ms: Send and receive timeout for each exchange. - _local: Thread-local mapping from address to REQ socket. - _closed: Whether the client context has been destroyed. - """ + """Send control requests through reusable, thread-local REQ sockets.""" def __init__(self, timeout_ms: int) -> None: self._context = zmq.Context() self._timeout_ms = timeout_ms self._local = threading.local() + self._topologies: dict[str, list[str]] = {} self._closed = False def _sockets(self) -> dict[str, zmq.Socket]: @@ -165,7 +69,7 @@ def _exchange(self, addr: str, payload: ControlRequest) -> ControlResponse: self._discard(addr) raise assert isinstance(response, dict) - return cast(ControlResponse, response) + return response def request(self, addr: str, payload: ControlRequest) -> Any: response = self._exchange(addr, payload) @@ -173,6 +77,26 @@ def request(self, addr: str, payload: ControlRequest) -> Any: raise RuntimeError(response.get("error", "EC control request failed")) return response.get("result") + def discover_shards(self, base_addr: str) -> list[str] | None: + if base_addr in self._topologies: + return self._topologies[base_addr] + try: + reply = self.request(base_addr, {"op": "peers"}) + ports = reply.get("ports") if isinstance(reply, dict) else None + if not isinstance(ports, list) or not ports: + raise ValueError("invalid or empty peer list") + prefix = base_addr.rsplit(":", 1)[0] + shards = [f"{prefix}:{int(port)}" for port in ports] + except Exception: + logger.warning( + "EC Mooncake consumer at %s did not report its shards", + base_addr, + exc_info=True, + ) + return None + self._topologies[base_addr] = shards + return shards + def close(self) -> None: if self._closed: return @@ -182,12 +106,12 @@ def close(self) -> None: def make_cancel_request( transfer_id: str, - reservation_id: str, + reservation_id: str = "", *, abandon: bool = False, refresh: bool = False, -) -> ReservationActionRequest: - request: ReservationActionRequest = { +) -> ControlRequest: + request: ControlRequest = { "op": "cancel", "transfer_id": transfer_id, "reservation_id": reservation_id, @@ -199,65 +123,11 @@ def make_cancel_request( return request -class ShardTopology: - """Discover and cache every control address for a Consumer. - - Attributes: - _client: Client used to query the Consumer's ``peers`` operation. - _cache: Base Consumer addresses mapped to all TP-shard addresses. - """ - - def __init__(self, client: ControlClient) -> None: - self._client = client - self._cache: dict[str, list[str]] = {} - - def discover(self, base_addr: str) -> list[str] | None: - """Return a confirmed complete topology, retrying transient failures.""" - cached = self._cache.get(base_addr) - if cached is not None: - return cached - try: - reply = self._client.request(base_addr, {"op": "peers"}) - ports = reply.get("ports") if isinstance(reply, dict) else None - if not isinstance(ports, list) or not ports: - raise ValueError("invalid or empty peer list") - prefix = base_addr.rsplit(":", 1)[0] - shards = [f"{prefix}:{int(port)}" for port in ports] - except Exception: - logger.warning( - "EC Mooncake consumer at %s did not report its shards; " - "using it directly for this attempt.", - base_addr, - exc_info=True, - ) - return None - self._cache[base_addr] = shards - if len(shards) > 1: - logger.info( - "EC Mooncake consumer at %s has %d shards", base_addr, len(shards) - ) - return shards - - def shards(self, base_addr: str) -> list[str]: - """Return confirmed shards or a one-attempt data-plane fallback.""" - return self.discover(base_addr) or [base_addr] - - class EventInbox: - """Receive Consumer readiness events without blocking the Scheduler. - - Attributes: - _client: Control client used to discover event ports. - _topology: Source of Consumer TP-shard addresses. - _context: Lazily created context for the PULL socket. - _socket: PULL socket connected to every expected Consumer shard. - _closed: Whether event resources have been released. - shard_count: Number of event channels in the complete topology. - """ + """Receive Consumer readiness events without blocking the Scheduler.""" - def __init__(self, client: ControlClient, topology: ShardTopology) -> None: + def __init__(self, client: ControlClient) -> None: self._client = client - self._topology = topology self._context: zmq.Context | None = None self._socket: zmq.Socket | None = None self._closed = False @@ -266,7 +136,7 @@ def __init__(self, client: ControlClient, topology: ShardTopology) -> None: def _connect(self, base_addr: str) -> None: if self._socket is not None: return - shards = self._topology.discover(base_addr) + shards = self._client.discover_shards(base_addr) if shards is None: return endpoints = [] @@ -317,23 +187,6 @@ class ConsumerControlServer: One server runs on every receiving TP rank. The REP channel handles reservation operations, while a PUSH channel publishes newly ready items to the Scheduler. - - Attributes: - host: Interface on which the control server listens. - port: Rank-local REP control port. - peer_ports: Control ports for every Consumer TP shard. - event_port: Dynamically allocated PUSH event port after startup. - _device: Device selected in the control thread when CUDA is used. - _reserve: Callback that allocates or reuses destination memory. - _status: Callback that reports active reservation state. - _complete: Callback that marks destination writes complete. - _cancel: Callback that cancels or abandons reservations. - _reap: Callback that expires stale reservations. - _metrics_log_interval: Interval for aggregate control-plane logs. - _stop: Signal requesting termination of the server loop. - _started: Signal indicating that socket binding has completed. - _thread: Background server thread. - _startup_error: Socket binding error captured from the server thread. """ def __init__( @@ -342,10 +195,9 @@ def __init__( port: int, reserve: Callable[[dict[str, Any]], dict[str, Any]], status: Callable[[str], dict[str, Any] | None], - complete: Callable[[str, str], ControlCompletion], + complete: Callable[[str, str], tuple[bool, bool]], cancel: Callable[[str, str, bool, bool], bool], reap: Callable[[], int], - metrics_log_interval: float = 10, peer_ports: list[int] | None = None, device: torch.device | None = None, ) -> None: @@ -359,7 +211,6 @@ def __init__( self._complete = complete self._cancel = cancel self._reap = reap - self._metrics_log_interval = metrics_log_interval self._stop = threading.Event() self._started = threading.Event() self._thread: threading.Thread | None = None @@ -373,18 +224,19 @@ def loop() -> None: socket = context.socket(zmq.REP) event_socket = context.socket(zmq.PUSH) pending_events: deque[dict[str, Any]] = deque() - metrics: Counter[str] = Counter() def queue_event(event: dict[str, Any]) -> None: event["shard"] = self.port if len(pending_events) >= _MAX_PENDING_EVENTS: pending_events.popleft() - metrics["events_dropped"] += 1 pending_events.append(event) - metrics["events_queued"] += 1 - metrics_started_at = time.monotonic() - last_reap_at = metrics_started_at + def queue_ready(transfer_id: str) -> None: + status = self._status(transfer_id) + if status is not None: + queue_event({"transfer_id": transfer_id, **status}) + + last_reap_at = time.monotonic() socket.setsockopt(zmq.RCVTIMEO, 100) try: socket.bind(f"tcp://{self.host}:{self.port}") @@ -407,32 +259,10 @@ def queue_event(event: dict[str, Any]) -> None: except zmq.Again: break pending_events.popleft() - metrics["events_sent"] += 1 now = time.monotonic() if now - last_reap_at >= _RESERVATION_REAP_INTERVAL_SECONDS: - metrics["reservations_reaped"] += self._reap() + self._reap() last_reap_at = now - if ( - self._metrics_log_interval > 0 - and now - metrics_started_at >= self._metrics_log_interval - ): - logger.info( - "EC Mooncake consumer control: requests=%s, " - "events_queued=%d, events_sent=%d, events_dropped=%d, " - "event_backlog=%d, reservations_reaped=%d", - { - key.removeprefix("request_"): value - for key, value in metrics.items() - if key.startswith("request_") - }, - metrics["events_queued"], - metrics["events_sent"], - metrics["events_dropped"], - len(pending_events), - metrics["reservations_reaped"], - ) - metrics.clear() - metrics_started_at = now try: request = socket.recv_json() except zmq.Again: @@ -440,14 +270,10 @@ def queue_event(event: dict[str, Any]) -> None: try: op = request.get("op") result: Any = None - metrics[f"request_{op}"] += 1 if op == "reserve": result = self._reserve(request) if result.get("ready"): - transfer_id = str(request["transfer_id"]) - status = self._status(transfer_id) - if status is not None: - queue_event({"transfer_id": transfer_id, **status}) + queue_ready(str(request["transfer_id"])) elif op == "status": result = self._status(str(request["transfer_id"])) elif op == "event_port": @@ -463,20 +289,18 @@ def queue_event(event: dict[str, Any]) -> None: completions = [] for item in items: transfer_id = str(item["transfer_id"]) - completion = self._complete( + accepted, became_ready = self._complete( transfer_id, str(item["reservation_id"]) ) completions.append( { - "completed": completion.accepted, - "became_ready": completion.became_ready, + "completed": accepted, + "became_ready": became_ready, } ) - if not completion.became_ready: + if not became_ready: continue - status = self._status(transfer_id) - if status is not None: - queue_event({"transfer_id": transfer_id, **status}) + queue_ready(transfer_id) result = ( {"items": completions} if op == "complete_batch" diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py index 1f09e19ebfdd..d7be4620dbc8 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py @@ -10,9 +10,8 @@ from __future__ import annotations import bisect -import math import threading -from collections import Counter, OrderedDict +from collections import OrderedDict from collections.abc import Callable from dataclasses import dataclass from typing import Generic, TypeVar @@ -31,13 +30,7 @@ @dataclass class MemoryAllocation: - """Describe a tensor view carved from the Consumer receive slab. - - Attributes: - offset: Byte offset of the allocation within the slab. - size: Aligned number of slab bytes owned by the allocation. - tensor: Typed tensor view exposed to transfer and cache code. - """ + """Describe a tensor view carved from the Consumer receive slab.""" offset: int size: int @@ -46,30 +39,16 @@ class MemoryAllocation: @dataclass class _ResidentEntry(Generic[_T]): - """Store one resident value and its ownership accounting. - - Attributes: - value: Resident value owned by the pool. - nbytes: Capacity charged to the resident pool. - pinned: Whether active cache state prevents LRU eviction. - leases: Number of in-flight reservations borrowing this value. - """ + """Store one resident value and its ownership accounting.""" value: _T - nbytes: int pinned: bool = True leases: int = 0 @dataclass class ResidentLease(Generic[_T]): - """Represent a borrow of one resident entry. - - Attributes: - key: Cache identifier used to find the canonical resident entry. - _entry: Entry retained even if the canonical mapping is replaced. - _active: Whether this lease still contributes to the reference count. - """ + """Represent a borrow of one resident entry.""" key: str _entry: _ResidentEntry[_T] @@ -81,21 +60,14 @@ def value(self) -> _T: class ContiguousAllocator: - """Allocate aligned regions from one contiguous byte range. - - Attributes: - capacity: Total number of bytes managed by the allocator. - alignment: Allocation granularity in bytes. - _free: Sorted free ranges represented as ``(offset, size)`` pairs. - """ + """Allocate aligned regions from one contiguous byte range.""" def __init__(self, capacity: int, alignment: int = 256): - self.capacity = capacity self.alignment = alignment self._free = [(0, capacity)] def allocate(self, nbytes: int) -> tuple[int, int] | None: - size = math.ceil(nbytes / self.alignment) * self.alignment + size = (nbytes + self.alignment - 1) // self.alignment * self.alignment for index, (offset, available) in enumerate(self._free): if size > available: continue @@ -126,26 +98,12 @@ def free(self, offset: int, size: int) -> None: class ResidentPool(Generic[_T]): - """Track resident values, active leases, and LRU eviction eligibility. - - Attributes: - used: Total bytes charged by current and leased displaced entries. - _entries: Canonical resident entries keyed by cache identifier. - _evictable: Unpinned and unleased entries in LRU order. - """ + """Track resident values, active leases, and LRU eviction eligibility.""" def __init__(self): - self.used = 0 self._entries: dict[str, _ResidentEntry[_T]] = {} self._evictable: OrderedDict[str, None] = OrderedDict() - def __len__(self) -> int: - return len(self._entries) - - @property - def num_evictable(self) -> int: - return len(self._evictable) - def referenced(self) -> list[str]: return [key for key, entry in self._entries.items() if entry.pinned] @@ -157,15 +115,11 @@ def insert( self, key: str, value: _T, - nbytes: int, ) -> _T | None: """Pin an entry and return a displaced value that has no owners.""" previous = self._entries.get(key) - entry = _ResidentEntry(value, nbytes) - if previous is not None and previous.leases == 0: - self.used -= previous.nbytes + entry = _ResidentEntry(value) self._entries[key] = entry - self.used += nbytes self._evictable.pop(key, None) if previous is not None and previous.leases == 0: return previous.value @@ -196,7 +150,6 @@ def release(self, lease: ResidentLease[_T]) -> _T | None: current = self._entries.get(lease.key) if current is not entry: if entry.leases == 0: - self.used -= entry.nbytes return entry.value return None if not entry.pinned and entry.leases == 0: @@ -225,40 +178,24 @@ def evict_lru(self, evict: Callable[[str, _T], bool]) -> str | None: continue self._evictable.pop(key, None) del self._entries[key] - self.used -= entry.nbytes return key return None def clear(self) -> None: self._entries.clear() self._evictable.clear() - self.used = 0 @dataclass class StagedSources: - """Own Producer tensor views and their staging-slab regions. - - Attributes: - tensors: Registered tensor views used as Mooncake sources. - regions: Allocator regions released after the write finishes. - """ + """Own Producer tensor views and their staging-slab regions.""" tensors: list[torch.Tensor] regions: list[tuple[int, int]] class ProducerMemoryPool: - """Own the Producer staging slab and regions carved from it. - - Attributes: - _capacity: Requested staging-slab size in bytes. - _transfer: Data-plane owner used to register the slab. - _pool: Lazily allocated registered byte tensor. - _allocator: Region allocator for the staging slab. - _disabled: Whether initialization failed and fallback is required. - _lock: Lock protecting initialization and region allocation. - """ + """Own the Producer staging slab and regions carved from it.""" def __init__(self, capacity: int, transfer: MooncakeTransfer) -> None: self._capacity = capacity @@ -268,10 +205,6 @@ def __init__(self, capacity: int, transfer: MooncakeTransfer) -> None: self._disabled = False self._lock = threading.Lock() - @property - def tensor(self) -> torch.Tensor | None: - return self._pool - def _ensure_pool(self, device: torch.device) -> None: if self._pool is not None or self._disabled: return @@ -337,25 +270,16 @@ def release(self, staged: StagedSources) -> None: self._free_regions(staged.regions) def close(self) -> None: - """Retain the registered slab until the full close phase owns it.""" + with self._lock: + pool = self._pool + if pool is None or not self._transfer.unregister_memory(pool): + return + self._pool = None + self._allocator = None class ConsumerMemoryPool: - """Own the registered receive slab and resident allocation lifecycle. - - Attributes: - _capacity: Requested receive-slab size in bytes. - _transfer: Data-plane owner used to register the slab. - _metrics: Counters describing resident-cache behavior. - _pool: Registered byte tensor that receives Mooncake writes. - _allocator: Region allocator for the receive slab. - _residents: Published allocations available for local reuse. - _retire_events: CUDA events guarding retired resident entries. - _pending_frees: Allocations waiting for CUDA consumers to finish. - _reclaimed: Cache identifiers evicted under allocation pressure. - _disabled: Whether receive-slab initialization has failed. - lock: Reentrant lock shared with reservation state transitions. - """ + """Own the registered receive slab and resident allocation lifecycle.""" def __init__( self, @@ -364,7 +288,6 @@ def __init__( ) -> None: self._capacity = capacity self._transfer = transfer - self._metrics: Counter[str] = Counter() self._pool: torch.Tensor | None = None self._allocator: ContiguousAllocator | None = None self._residents: ResidentPool[MemoryAllocation] = ResidentPool() @@ -381,17 +304,8 @@ def tensor(self) -> torch.Tensor | None: def prepare( self, device: torch.device, - *, - receiving_rank: bool, - allow_host: bool = False, ) -> None: - if not receiving_rank: - return - if ( - self._pool is not None - or self._disabled - or (device.type != "cuda" and not allow_host) - ): + if self._pool is not None or self._disabled: return try: pool = torch.empty(self._capacity, dtype=torch.uint8, device=device) @@ -401,17 +315,15 @@ def prepare( except (RuntimeError, torch.OutOfMemoryError) as error: self._disabled = True logger.warning( - "Could not initialize the EC consumer buffer pool; falling back " - "to per-tensor registration: %s", + "Could not initialize the EC consumer buffer pool: %s", error, ) return self._pool = pool self._allocator = ContiguousAllocator(pool.nbytes) logger.info( - "Prepared %d-byte CUDA receive pool for Mooncake EC (registered=%s)", + "Prepared %d-byte receive pool for Mooncake EC", pool.nbytes, - receiving_rank, ) def _free(self, allocation: MemoryAllocation) -> None: @@ -446,7 +358,6 @@ def evict(mm_hash: str, allocation: MemoryAllocation) -> bool: event = self._retire_events.pop(mm_hash, None) self._defer_or_free(allocation, event) self._reclaimed.add(mm_hash) - self._metrics["residents_reclaimed"] += 1 return True while self._residents.evict_lru(evict) is not None: @@ -508,18 +419,15 @@ def take_resident( with self.lock: allocation = self._residents.get(mm_hash) if allocation is None: - self._metrics["residents_missed"] += 1 return None tensor = allocation.tensor if ( tuple(tensor.shape) != shape or str(tensor.dtype).split(".")[-1] != dtype_name ): - self._metrics["residents_mismatched"] += 1 return None self._residents.pin(mm_hash) self._retire_events.pop(mm_hash, None) - self._metrics["residents_promoted"] += 1 return tensor def _record_release_event(self) -> torch.Event | None: @@ -549,7 +457,7 @@ def publish( self._defer_or_free(released, self._record_release_event()) return canonical previous = self._residents.get(mm_hash) - displaced = self._residents.insert(mm_hash, allocation, allocation.size) + displaced = self._residents.insert(mm_hash, allocation) event = None if previous is not None and previous is not allocation: event = self._retire_events.pop(mm_hash, None) @@ -576,11 +484,10 @@ def retire_stale( continue if mm_hash in reserved_hashes: continue - event = torch.Event() - event.record(torch.accelerator.current_stream(self._pool.device)) - self._retire_events[mm_hash] = event + event = self._record_release_event() + if event is not None: + self._retire_events[mm_hash] = event self._residents.retire(mm_hash) - self._metrics["residents_retired"] += 1 self._poll_frees_locked() def drain_reclaimed(self) -> set[str]: @@ -589,21 +496,6 @@ def drain_reclaimed(self) -> set[str]: self._reclaimed = set() return reclaimed - def stats(self) -> tuple[int, int, int, int]: - with self.lock: - return ( - len(self._residents), - len(self._residents.referenced()), - self._residents.num_evictable, - len(self._pending_frees), - ) - - def take_metrics(self) -> dict[str, int]: - with self.lock: - metrics = dict(self._metrics) - self._metrics.clear() - return metrics - def close(self) -> None: with self.lock: pool = self._pool diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py index 96923e415856..2ccf80911ef2 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py @@ -14,28 +14,13 @@ @dataclass class ECMooncakeLoadSpec: - """Describe one Consumer-side cache load requested by the Scheduler. - - Attributes: - mm_hash: Stable identifier of the multimodal encoder item. - num_token: Number of encoder tokens expected by the request. - nbytes: Tensor payload size in bytes. - shape: Tensor shape reconstructed by the Consumer. - dtype: Unqualified ``torch.dtype`` name used for reconstruction. - pushed: Whether the tensor came from a remote Producer reservation. - transfer_id: Identity shared by Scheduler, Producer, and Consumer. - reservation_id: Consumer-issued capability for completing the write. - local: Whether to reuse a tensor already resident on the Consumer. - """ + """Describe a remote reservation or resident allocation to load.""" mm_hash: str - num_token: int nbytes: int shape: tuple[int, ...] dtype: str - pushed: bool = False - transfer_id: str = "" - reservation_id: str = "" + transfer_id: str # The consumer pool still holds this item, so the load is a local handoff: # no transfer, no producer. local: bool = False @@ -43,17 +28,7 @@ class ECMooncakeLoadSpec: @dataclass class ECMooncakePushSpec: - """Describe a destination reservation prepared before a tensor is ready. - - Attributes: - mm_hash: Stable identifier of the multimodal encoder item. - nbytes: Number of bytes the Consumer must reserve. - shape: Shape of the tensor that will be written. - dtype: Unqualified ``torch.dtype`` name of the tensor. - consumer_zmq: Base control address of the destination Consumer. - transfer_id: Identity shared by Scheduler, Producer, and Consumer. - request_id: Request that owns the push and may cancel it. - """ + """Describe a destination reservation prepared before a tensor is ready.""" mm_hash: str nbytes: int @@ -66,41 +41,21 @@ class ECMooncakePushSpec: @dataclass class ECMooncakeConnectorMetadata(ECConnectorMetadata): - """Worker operations emitted for one Scheduler step. - - Attributes: - loads: Consumer loads that should be attached to ``encoder_cache``. - pushes: Producer reservations that should begin before sources arrive. - """ + """Worker operations emitted for one Scheduler step.""" loads: list[ECMooncakeLoadSpec] = field(default_factory=list) pushes: list[ECMooncakePushSpec] = field(default_factory=list) - def add_load(self, spec: ECMooncakeLoadSpec) -> None: - self.loads.append(spec) - - def add_push(self, spec: ECMooncakePushSpec) -> None: - self.pushes.append(spec) - @dataclass class ECMooncakeWorkerMetadata(ECConnectorWorkerMetadata): - """Completion state reported from Workers to the Scheduler. - - Attributes: - loaded: Cache identifiers loaded successfully on this Worker. - failed_loads: Cache identifiers that could not be loaded. - reclaimed: Resident items evicted because the receive pool was full. - pending_loads: Whether this Worker still owns asynchronous load work. - pending_saves: Whether this Worker still owns asynchronous push work. - """ + """Completion state reported from Workers to the Scheduler.""" loaded: set[str] = field(default_factory=set) failed_loads: set[str] = field(default_factory=set) # Items the receive pool dropped under pressure. The scheduler assumes an # evicted item stays resident until told otherwise. reclaimed: set[str] = field(default_factory=set) - pending_loads: bool = False pending_saves: bool = False def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadata: @@ -113,6 +68,5 @@ def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadat loaded=self.loaded & other.loaded, failed_loads=self.failed_loads | other.failed_loads, reclaimed=self.reclaimed | other.reclaimed, - pending_loads=self.pending_loads or other.pending_loads, pending_saves=self.pending_saves or other.pending_saves, ) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py index 16949916870c..947395b9997d 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py @@ -9,7 +9,6 @@ from __future__ import annotations import threading -import time from collections import OrderedDict from collections.abc import Callable from concurrent.futures import Future, ThreadPoolExecutor @@ -28,8 +27,8 @@ class ProducerPushState(Enum): """Lifecycle of one Producer push from reservation to terminal state.""" - RESERVING = auto() - WAITING_SOURCE = auto() + # The reservation reply and source tensor may arrive in either order. + WAITING_INPUTS = auto() WRITING = auto() NOTIFYING = auto() DONE = auto() @@ -38,53 +37,23 @@ class ProducerPushState(Enum): FAILED = auto() -@dataclass -class ProducerSourceLease: - """Keep an encoder tensor alive until every remote write is settled. - - Attributes: - tensor: Source encoder tensor owned by the push. - ready_event: CUDA event proving that production of the tensor finished. - """ - - tensor: torch.Tensor - ready_event: torch.Event | None - - @dataclass class ProducerPushRecord: - """Collect all asynchronous state for one Producer push. - - Attributes: - spec: Immutable identity and destination metadata for the push. - state: Current Producer lifecycle state. - reservation_futures: Futures resolving Consumer shard reservations. - reservations: Resolved destination descriptors for every shard. - shard_futures: Data-plane futures that may still read the source. - source: Source tensor lease once encoder computation has completed. - batch_future: Transfer or cancellation batch currently owning the push. - error: First asynchronous error retained for Worker reporting. - source_at: Time the source became available for queue metrics. - """ + """Collect all asynchronous state for one Producer push.""" spec: ECMooncakePushSpec state: ProducerPushState - reservation_futures: list[Future[list[dict[str, Any]]]] + reservation_future: Future[list[dict[str, Any]]] reservations: list[dict[str, Any]] = field(default_factory=list) shard_futures: list[Future[Any]] = field(default_factory=list) - source: ProducerSourceLease | None = None + source_tensor: torch.Tensor | None = None + source_event: torch.Event | None = None batch_future: Future[None] | None = None error: str | None = None - source_at: float | None = None _ALLOWED_TRANSITIONS = { - ProducerPushState.RESERVING: { - ProducerPushState.WAITING_SOURCE, - ProducerPushState.CANCEL_PENDING, - ProducerPushState.FAILED, - }, - ProducerPushState.WAITING_SOURCE: { + ProducerPushState.WAITING_INPUTS: { ProducerPushState.WRITING, ProducerPushState.CANCEL_PENDING, ProducerPushState.FAILED, @@ -97,7 +66,10 @@ class ProducerPushRecord: ProducerPushState.DONE, ProducerPushState.FAILED, }, - ProducerPushState.CANCEL_PENDING: {ProducerPushState.CANCELLED}, + ProducerPushState.CANCEL_PENDING: { + ProducerPushState.CANCELLED, + ProducerPushState.FAILED, + }, ProducerPushState.DONE: set(), ProducerPushState.CANCELLED: set(), ProducerPushState.FAILED: set(), @@ -108,24 +80,14 @@ class ProducerPushRecord: ProducerPushState.CANCELLED, ProducerPushState.FAILED, } -_SOURCE_WAIT_STATES = { - ProducerPushState.RESERVING, - ProducerPushState.WAITING_SOURCE, -} _TERMINAL_LIMIT = 1 << 16 class ProducerPushManager: """Own Producer push records, transitions, and source tensor leases. - Attributes: - _records: All active and retained terminal records by transfer ID. - _active_ids: Non-terminal transfer IDs in insertion order. - _reapable_terminal_ids: Terminal records safe to discard. - _unreported_ids: Failed records awaiting Worker error reporting. - _batch_ids: Records whose batch future has not been reaped. - _source_waiters: Transfer IDs waiting for each cache identifier. - _lock: Reentrant lock protecting lifecycle and ownership changes. + Terminal records are retained for duplicate-metadata idempotency. Separate + ordered indexes keep hot polling paths from scanning those tombstones. """ def __init__(self) -> None: @@ -137,10 +99,6 @@ def __init__(self) -> None: self._source_waiters: dict[str, OrderedDict[str, None]] = {} self._lock = threading.RLock() - def get(self, transfer_id: str) -> ProducerPushRecord | None: - with self._lock: - return self._records.get(transfer_id) - def reserve( self, spec: ECMooncakePushSpec, @@ -156,15 +114,15 @@ def reserve( return existing, False record = ProducerPushRecord( spec=spec, - state=ProducerPushState.RESERVING, - reservation_futures=[submit()], + state=ProducerPushState.WAITING_INPUTS, + reservation_future=submit(), ) self._records[spec.transfer_id] = record self._active_ids[spec.transfer_id] = None self._source_waiters.setdefault(spec.mm_hash, OrderedDict())[ spec.transfer_id ] = None - record.reservation_futures[0].add_done_callback( + record.reservation_future.add_done_callback( lambda future: self._reservation_done(record, future) ) return record, True @@ -179,50 +137,40 @@ def bind_source( waiters = self._source_waiters.pop(mm_hash, OrderedDict()) for transfer_id in waiters: record = self._records[transfer_id] - if record.source is not None or record.state not in _SOURCE_WAIT_STATES: + if ( + record.source_tensor is not None + or record.state is not ProducerPushState.WAITING_INPUTS + ): continue - record.source = ProducerSourceLease(tensor, ready_event) - record.source_at = time.monotonic() + record.source_tensor = tensor + record.source_event = ready_event def submit_batches( self, executor: ThreadPoolExecutor, run_batch: Callable[[list[ProducerPushRecord]], None], - on_submit: Callable[[], None], ) -> None: with self._lock: grouped: dict[str, list[ProducerPushRecord]] = {} for transfer_id in list(self._active_ids): record = self._records[transfer_id] if ( - record.source is not None + record.source_tensor is not None and record.batch_future is None - and record.state in _SOURCE_WAIT_STATES + and record.state is ProducerPushState.WAITING_INPUTS ): grouped.setdefault(record.spec.consumer_zmq, []).append(record) batches = list(grouped.values()) for records in batches: - on_submit() future = executor.submit(run_batch, records) for record in records: record.batch_future = future self._batch_ids[record.spec.transfer_id] = None def resolve_reservations(self, record: ProducerPushRecord) -> list[dict[str, Any]]: - results: list[dict[str, Any]] = [] - error: Exception | None = None - for future in record.reservation_futures: - try: - results.extend(future.result()) - except Exception as exc: - if error is None: - error = exc - if error is not None: - raise error + results = record.reservation_future.result() with self._lock: record.reservations = results - if record.state is ProducerPushState.RESERVING: - self._transition(record, ProducerPushState.WAITING_SOURCE) return list(results) def _reservation_done( @@ -236,20 +184,15 @@ def _reservation_done( with self._lock: self._set_error(record, exc) if ( - record.state is ProducerPushState.RESERVING - and record.source is None + record.state is ProducerPushState.WAITING_INPUTS + and record.source_tensor is None ): self._transition(record, ProducerPushState.FAILED) - return - with self._lock: - if record.state is ProducerPushState.RESERVING: - self._transition(record, ProducerPushState.WAITING_SOURCE) def settle_all(self, records: list[ProducerPushRecord]) -> None: for record in records: - for future in record.reservation_futures: - with suppress(Exception): - future.result() + with suppress(Exception): + record.reservation_future.result() def replace_reservations( self, @@ -270,7 +213,7 @@ def track_shard_futures( def begin_writing(self, record: ProducerPushRecord) -> None: with self._lock: - if record.source is None: + if record.source_tensor is None: raise RuntimeError( f"Producer push {record.spec.transfer_id!r} has no source tensor" ) @@ -297,18 +240,22 @@ def fail(self, records: list[ProducerPushRecord], error: Exception) -> None: self._transition(record, ProducerPushState.FAILED) self._release_source(record) - def cancel_requests(self, request_ids: set[str]) -> list[ProducerPushRecord]: + def cancel_requests(self, request_ids: set[str] | None) -> list[ProducerPushRecord]: + """Cancel source-less waiters for requests, or all waiters at shutdown.""" cancelled = [] with self._lock: for transfer_id in list(self._active_ids): record = self._records[transfer_id] if ( - record.spec.request_id not in request_ids - or record.source is not None + ( + request_ids is not None + and record.spec.request_id not in request_ids + ) + or record.source_tensor is not None or record.batch_future is not None ): continue - if record.state not in _SOURCE_WAIT_STATES: + if record.state is not ProducerPushState.WAITING_INPUTS: continue self._transition(record, ProducerPushState.CANCEL_PENDING) cancelled.append(record) @@ -365,8 +312,8 @@ def pending(self) -> bool: return bool(self._active_ids or self._batch_ids) def _release_source(self, record: ProducerPushRecord) -> None: - if record.source is not None: - record.source = None + record.source_tensor = None + record.source_event = None def _set_error(self, record: ProducerPushRecord, error: BaseException) -> None: if record.error is None: @@ -401,8 +348,8 @@ def _drop_source_waiter(self, record: ProducerPushRecord) -> None: @staticmethod def _check_sources_releasable(records: list[ProducerPushRecord]) -> None: for record in records: - futures = [*record.reservation_futures, *record.shard_futures] - if record.source is not None and not all( + futures = [record.reservation_future, *record.shard_futures] + if record.source_tensor is not None and not all( future.done() for future in futures ): raise RuntimeError( @@ -417,7 +364,7 @@ def _transition(self, record: ProducerPushRecord, state: ProducerPushState) -> N f"{record.state.name} to {state.name}" ) record.state = state - if state not in _SOURCE_WAIT_STATES: + if state is not ProducerPushState.WAITING_INPUTS: self._drop_source_waiter(record) if state in _TERMINAL_STATES: transfer_id = record.spec.transfer_id diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py index 2873f55a89d1..e0be41260953 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py @@ -8,10 +8,11 @@ from __future__ import annotations +import threading import time import uuid from collections import OrderedDict -from dataclasses import dataclass, field +from dataclasses import dataclass from enum import Enum, auto import torch @@ -26,11 +27,8 @@ class ConsumerReservationState(Enum): """Lifecycle of one Consumer destination reservation.""" - RESERVED = auto() WRITING = auto() READY = auto() - TAKEN = auto() - RESIDENT = auto() CANCEL_PENDING = auto() EXPIRE_PENDING = auto() CANCELLED = auto() @@ -38,10 +36,6 @@ class ConsumerReservationState(Enum): _ALLOWED_TRANSITIONS = { - ConsumerReservationState.RESERVED: { - ConsumerReservationState.WRITING, - ConsumerReservationState.CANCELLED, - }, ConsumerReservationState.WRITING: { ConsumerReservationState.READY, ConsumerReservationState.CANCEL_PENDING, @@ -49,12 +43,9 @@ class ConsumerReservationState(Enum): ConsumerReservationState.CANCELLED, }, ConsumerReservationState.READY: { - ConsumerReservationState.TAKEN, ConsumerReservationState.CANCELLED, ConsumerReservationState.EXPIRED, }, - ConsumerReservationState.TAKEN: {ConsumerReservationState.RESIDENT}, - ConsumerReservationState.RESIDENT: set(), ConsumerReservationState.CANCEL_PENDING: {ConsumerReservationState.CANCELLED}, ConsumerReservationState.EXPIRE_PENDING: { ConsumerReservationState.CANCELLED, @@ -74,24 +65,16 @@ class ConsumerReservationState(Enum): ConsumerReservationState.CANCEL_PENDING, ConsumerReservationState.EXPIRE_PENDING, } +_WRITER_OWNED_STATES = { + ConsumerReservationState.WRITING, + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, +} @dataclass class ConsumerReservation: - """Own the identity and destination allocation of one remote push. - - Attributes: - transfer_id: Cross-process identity of the transfer. - mm_hash: Stable identifier of the encoder-cache item. - reservation_id: Consumer-issued identity required for completion. - state: Current destination reservation state. - shape: Expected tensor shape. - dtype: Unqualified expected ``torch.dtype`` name. - allocation: Receive-slab allocation while this record owns it. - lease: Borrowed resident allocation used for a cache hit. - created_at: Monotonic creation time used for diagnostics. - expires_at: Deadline for the active record or terminal tombstone. - """ + """Own the identity and destination allocation of one remote push.""" transfer_id: str mm_hash: str @@ -101,46 +84,15 @@ class ConsumerReservation: dtype: str = "" allocation: MemoryAllocation | None = None lease: ResidentLease[MemoryAllocation] | None = None - created_at: float = field(default_factory=time.monotonic) expires_at: float = 0 -@dataclass(frozen=True) -class CompletionResult: - """Describe how a completion request affected a reservation. - - Attributes: - accepted: Whether transfer and reservation identities matched. - became_ready: Whether the call transitioned WRITING to READY. - repeated: Whether the reservation was already ready. - discarded: Whether deferred cancellation consumed the completion. - """ - - accepted: bool - became_ready: bool = False - repeated: bool = False - discarded: bool = False - - -class CancellationOutcome(Enum): - """Outcome categories used for control responses and metrics.""" - - REJECTED = auto() - PRE_RESERVED = auto() - DEFERRED = auto() - CANCELLED = auto() - - class ConsumerReservationManager: """Own reservation transitions and destination allocation releases. - Attributes: - _memory: Consumer memory pool that owns destination allocations. - _lease_ttl: Lifetime of active reservations and terminal tombstones. - _tombstone_limit: Maximum retained terminal cancellation records. - _records: Active and terminal records keyed by transfer ID. - _active_ids: Transfer IDs requiring expiry scans. - _tombstones: Terminal records retained in expiry order. + A remote writer owns a ``WRITING`` allocation. Cancellation and expiry are + therefore deferred until completion, unless refresh explicitly abandons + that writer before replacing its reservation. """ def __init__( @@ -155,12 +107,8 @@ def __init__( self._records: dict[str, ConsumerReservation] = {} self._active_ids: dict[str, None] = {} self._tombstones: OrderedDict[str, None] = OrderedDict() - - def get(self, transfer_id: str) -> ConsumerReservation | None: - return self._records.get(transfer_id) - - def active_records(self) -> list[ConsumerReservation]: - return [self._records[transfer_id] for transfer_id in self._active_ids] + self._condition = threading.Condition(memory.lock) + self._shutting_down = False def reserve( self, @@ -170,14 +118,16 @@ def reserve( shape: tuple[int, ...], dtype_name: str, dtype: torch.dtype, - ) -> tuple[ConsumerReservation | None, bool, bool, tuple[int, int, int]]: - with self._memory.lock: + ) -> tuple[ConsumerReservation | None, bool]: + with self._condition: + if self._shutting_down: + raise RuntimeError("Consumer reservation manager is shutting down") existing = self._records.get(transfer_id) if ( existing is not None and existing.state is ConsumerReservationState.CANCELLED ): - return existing, False, False, (0, 0, 0) + return existing, False if ( existing is not None and existing.state is ConsumerReservationState.EXPIRED @@ -205,7 +155,7 @@ def reserve( ) if existing.state is ConsumerReservationState.WRITING: existing.expires_at = time.monotonic() + self._lease_ttl - return existing, False, True, (0, 0, 0) + return existing, False lease = self._memory.acquire_cached(mm_hash, shape, dtype) now = time.monotonic() @@ -219,35 +169,31 @@ def reserve( dtype_name, lease.value, lease, - now, now + self._lease_ttl, ) self._insert(record) - return record, False, False, (0, 0, 0) - expiry_counts = (0, 0, 0) + return record, False allocation = self._memory.try_allocate(nbytes, shape, dtype) if allocation is None: - expiry_counts = self._expire_locked(time.monotonic()) + self._expire_locked(time.monotonic()) allocation = self._memory.try_allocate(nbytes, shape, dtype) if allocation is None: allocation = self._memory.reclaim_and_allocate(nbytes, shape, dtype) if allocation is None: - return None, False, False, expiry_counts + return None, False record = ConsumerReservation( transfer_id, mm_hash, uuid.uuid4().hex, - ConsumerReservationState.RESERVED, + ConsumerReservationState.WRITING, shape, dtype_name, allocation, None, - now, now + self._lease_ttl, ) - self._transition(record, ConsumerReservationState.WRITING) self._insert(record) - return record, True, False, expiry_counts + return record, True def status(self, transfer_id: str) -> ConsumerReservation | None: with self._memory.lock: @@ -256,26 +202,55 @@ def status(self, transfer_id: str) -> ConsumerReservation | None: return None return record - def complete(self, transfer_id: str, reservation_id: str) -> CompletionResult: - with self._memory.lock: - record = self._records.get(transfer_id) - if record is None or record.reservation_id != reservation_id: - return CompletionResult(False) - if record.state is ConsumerReservationState.READY: - return CompletionResult(True, repeated=True) - if record.state in _DEFERRED_STATES: - terminal = ( - ConsumerReservationState.CANCELLED - if record.state is ConsumerReservationState.CANCEL_PENDING - else ConsumerReservationState.EXPIRED - ) - self._terminate(record, terminal) - return CompletionResult(True, discarded=True) - if record.state is not ConsumerReservationState.WRITING: - return CompletionResult(False) - self._transition(record, ConsumerReservationState.READY) - record.expires_at = time.monotonic() + self._lease_ttl - return CompletionResult(True, became_ready=True) + def complete(self, transfer_id: str, reservation_id: str) -> tuple[bool, bool]: + with self._condition: + try: + record = self._records.get(transfer_id) + if record is None or record.reservation_id != reservation_id: + return False, False + if record.state is ConsumerReservationState.READY: + return True, False + if record.state in _DEFERRED_STATES: + terminal = ( + ConsumerReservationState.CANCELLED + if record.state is ConsumerReservationState.CANCEL_PENDING + else ConsumerReservationState.EXPIRED + ) + self._terminate(record, terminal) + return True, False + if record.state is not ConsumerReservationState.WRITING: + return False, False + self._transition(record, ConsumerReservationState.READY) + record.expires_at = time.monotonic() + self._lease_ttl + return True, True + finally: + self._condition.notify_all() + + def begin_shutdown(self) -> None: + """Stop new reservations and cancel everything without a remote writer.""" + with self._condition: + if self._shutting_down: + return + self._shutting_down = True + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if record.state is ConsumerReservationState.READY: + self._terminate(record, ConsumerReservationState.CANCELLED) + elif record.state is ConsumerReservationState.WRITING: + self._transition(record, ConsumerReservationState.CANCEL_PENDING) + self._condition.notify_all() + + def wait_for_writers(self, timeout: float) -> bool: + """Wait until no reservation is still owned by a remote writer.""" + + def writers_finished() -> bool: + return not any( + self._records[transfer_id].state in _WRITER_OWNED_STATES + for transfer_id in self._active_ids + ) + + with self._condition: + return self._condition.wait_for(writers_finished, timeout=max(0, timeout)) def cancel( self, @@ -283,7 +258,7 @@ def cancel( reservation_id: str, abandon: bool = False, refresh: bool = False, - ) -> tuple[CancellationOutcome, int]: + ) -> bool: with self._memory.lock: record = self._records.get(transfer_id) if ( @@ -291,48 +266,46 @@ def cancel( and reservation_id and record.reservation_id != reservation_id ): - return CancellationOutcome.REJECTED, 0 + return False if record is None: - now = time.monotonic() record = ConsumerReservation( transfer_id, "", "", ConsumerReservationState.CANCELLED, - created_at=now, ) self._insert(record) self._set_tombstone_deadline(record) - dropped = self._reap_tombstones(now) - return CancellationOutcome.PRE_RESERVED, dropped + self._reap_tombstones(time.monotonic()) + return True if record.state is ConsumerReservationState.CANCELLED: self._set_tombstone_deadline(record) - dropped = self._reap_tombstones(time.monotonic()) - return CancellationOutcome.PRE_RESERVED, dropped + self._reap_tombstones(time.monotonic()) + return True if refresh: if not abandon or record.state not in { ConsumerReservationState.WRITING, ConsumerReservationState.EXPIRE_PENDING, }: - return CancellationOutcome.REJECTED, 0 + return False if record.state is ConsumerReservationState.WRITING: self._transition(record, ConsumerReservationState.EXPIRE_PENDING) self._terminate(record, ConsumerReservationState.EXPIRED) - dropped = self._reap_tombstones(time.monotonic()) - return CancellationOutcome.CANCELLED, dropped + self._condition.notify_all() + self._reap_tombstones(time.monotonic()) + return True if record.state in _DEFERRED_STATES and not abandon: - return CancellationOutcome.DEFERRED, 0 + return True if record.state is ConsumerReservationState.WRITING and not abandon: self._transition(record, ConsumerReservationState.CANCEL_PENDING) - return CancellationOutcome.DEFERRED, 0 - if ( - record.state not in _ACTIVE_STATES - and record.state is not ConsumerReservationState.RESERVED - ): - return CancellationOutcome.REJECTED, 0 + return True + if record.state not in _ACTIVE_STATES: + return False self._terminate(record, ConsumerReservationState.CANCELLED) - dropped = self._reap_tombstones(time.monotonic()) - return CancellationOutcome.CANCELLED, dropped + if self._shutting_down: + self._condition.notify_all() + self._reap_tombstones(time.monotonic()) + return True def take(self, transfer_id: str, mm_hash: str) -> MemoryAllocation: with self._memory.lock: @@ -346,26 +319,25 @@ def take(self, transfer_id: str, mm_hash: str) -> MemoryAllocation: raise RuntimeError( f"Pushed EC tensor is not ready for mm_hash={mm_hash}" ) - self._transition(record, ConsumerReservationState.TAKEN) allocation = self._memory.publish(mm_hash, record.allocation, record.lease) record.allocation = None record.lease = None - self._transition(record, ConsumerReservationState.RESIDENT) self._remove(transfer_id) return allocation - def expire(self) -> tuple[int, int, int]: + def expire(self) -> int: with self._memory.lock: return self._expire_locked(time.monotonic()) def retire_stale(self, encoder_cache: dict[str, torch.Tensor]) -> None: with self._memory.lock: - reserved_hashes = {record.mm_hash for record in self.active_records()} + reserved_hashes = { + self._records[transfer_id].mm_hash for transfer_id in self._active_ids + } self._memory.retire_stale(encoder_cache, reserved_hashes) - def _expire_locked(self, now: float) -> tuple[int, int, int]: + def _expire_locked(self, now: float) -> int: expired = 0 - deferred = 0 for transfer_id in list(self._active_ids): record = self._records[transfer_id] if record.expires_at > now: @@ -375,9 +347,8 @@ def _expire_locked(self, now: float) -> tuple[int, int, int]: expired += 1 elif record.state is ConsumerReservationState.WRITING: self._transition(record, ConsumerReservationState.EXPIRE_PENDING) - deferred += 1 - dropped = self._reap_tombstones(now) - return expired, deferred, dropped + self._reap_tombstones(now) + return expired def _terminate( self, record: ConsumerReservation, state: ConsumerReservationState diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py index a79c99927057..201db2a97c13 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py @@ -11,7 +11,7 @@ import math import time -from collections import Counter, OrderedDict +from collections import OrderedDict from collections.abc import Collection from concurrent.futures import Future, ThreadPoolExecutor from typing import TYPE_CHECKING, Any @@ -20,19 +20,17 @@ from vllm.distributed.ec_transfer.ec_connector.base import ( ECConnectorMetadata, - ECConnectorRole, ) from vllm.distributed.ec_transfer.ec_connector.cpu.common import ( _get_encoder_cache_hidden_dim, ) -from vllm.distributed.ec_transfer.ec_connector.mooncake._availability import ( - ensure_mooncake_available, +from vllm.distributed.ec_transfer.ec_connector.mooncake.config import ( + _RESERVATION_TTL_SECONDS, + MooncakeECConfig, ) -from vllm.distributed.ec_transfer.ec_connector.mooncake.config import MooncakeECConfig from vllm.distributed.ec_transfer.ec_connector.mooncake.control import ( ControlClient, EventInbox, - ShardTopology, make_cancel_request, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( @@ -45,12 +43,15 @@ SchedulerTransferState, SchedulerTransferTable, ) +from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( + ensure_mooncake_available, +) from vllm.logger import init_logger from vllm.v1.core.sched.output import SchedulerOutput from vllm.v1.outputs import ECConnectorOutput if TYPE_CHECKING: - from vllm.config import ModelConfig, VllmConfig + from vllm.config import VllmConfig logger = init_logger(__name__) @@ -58,115 +59,52 @@ # event: warn on the first few, then only once in a while. _MAX_UNRESOLVED_TRANSFER_ID_WARNINGS = 5 -_LEASE_TTL_SECONDS = 300 _DRAIN_MIN_INTERVAL = 0.005 _MAX_PENDING_EVENTS = 4096 _MAX_TERMINAL_TRANSFER_RECORDS = 1 << 16 _CANCEL_ATTEMPTS = 2 +_CONTROL_WORKERS = 8 class ECMooncakeScheduler: - """Coordinate Mooncake transfers from the vLLM Scheduler process. - - Attributes: - _is_producer: Whether this Scheduler prepares outbound pushes. - _is_consumer: Whether this Scheduler waits for inbound pushes. - _reservation_zmq_addr: Base Consumer control-plane address. - _consumer_pool_capacity: Capacity mirrored for resident-state limits. - _push_wait_timeout: Maximum wait for a Consumer readiness event. - _consumer_metrics_log_interval: Interval for Scheduler metrics logs. - _encoder_cache_hidden_dim: Hidden width used to derive push shapes. - _model_config: Model metadata used for dtypes and multimodal fields. - _control_client: Client for reservation status and cancellation. - _topology: Discovery cache for Consumer TP shards. - _event_inbox: Non-blocking source of Consumer readiness events. - _control_executor: Executor for cancellation requests. - _metadata_fields_cache: Placeholder metadata fields by modality. - _consumer_metrics_started_at: Start time of the current metric window. - _consumer_scheduler_metrics: Counters for scheduling decisions. - _drain_pending: Whether the next scheduling pass should drain events. - _drained_at: Time of the most recent readiness-event drain. - _pending_cancels: Asynchronous cancellations by transfer ID. - _transfers: Scheduler-owned transfer lifecycle table. - _scheduler_pending_work: Whether Workers reported unfinished work. - _pushes_to_prepare: Push specs awaiting metadata emission. - _prepared_push_transfer_ids: Transfer IDs already prepared once. - _event_shard_count: Number of Consumer event channels in the topology. - _event_ready_shards: Ready shard IDs accumulated per transfer. - """ - - @classmethod - def from_vllm_config(cls, vllm_config: VllmConfig) -> ECMooncakeScheduler: - ensure_mooncake_available() - config = MooncakeECConfig.from_vllm_config( - vllm_config, ECConnectorRole.SCHEDULER - ) + """Coordinate Mooncake transfers from the vLLM Scheduler process.""" - control_client = ControlClient(config.control_timeout_ms) - topology = ShardTopology(control_client) - event_inbox = EventInbox(control_client, topology) - control_executor = ThreadPoolExecutor( - max_workers=config.control_workers, - thread_name_prefix="ec-mooncake-control", - ) - encoder_cache_hidden_dim = ( - _get_encoder_cache_hidden_dim(vllm_config) if config.is_producer else None - ) - return cls( - config, - encoder_cache_hidden_dim, - model_config=vllm_config.model_config, - control_client=control_client, - topology=topology, - event_inbox=event_inbox, - control_executor=control_executor, - ) + def __init__(self, vllm_config: VllmConfig) -> None: + ensure_mooncake_available() + config = MooncakeECConfig.from_vllm_config(vllm_config) - def __init__( - self, - config: MooncakeECConfig, - encoder_cache_hidden_dim: int | None, - model_config: ModelConfig, - control_client: ControlClient, - topology: ShardTopology, - event_inbox: EventInbox, - control_executor: ThreadPoolExecutor, - ) -> None: self._is_producer = config.is_producer self._is_consumer = config.is_consumer - self._reservation_zmq_addr = config.reservation_addr - self._consumer_pool_capacity = config.consumer_pool_size + self._control_addr = config.control_addr self._push_wait_timeout = config.push_wait_timeout_s - self._consumer_metrics_log_interval = config.consumer_metrics_log_interval - self._encoder_cache_hidden_dim = encoder_cache_hidden_dim - self._model_config = model_config - self._control_client = control_client - self._topology = topology - self._event_inbox = event_inbox - self._control_executor = control_executor + self._encoder_cache_hidden_dim = ( + _get_encoder_cache_hidden_dim(vllm_config) if config.is_producer else None + ) + self._model_config = vllm_config.model_config + self._control_client = ControlClient(config.control_timeout_ms) + self._event_inbox = EventInbox(self._control_client) + self._control_executor = ThreadPoolExecutor( + max_workers=_CONTROL_WORKERS, + thread_name_prefix="ec-mooncake-control", + ) self._metadata_fields_cache: dict[str, set[str]] = {} - self._consumer_metrics_started_at = time.monotonic() - self._consumer_scheduler_metrics: Counter[str] = Counter() self._unresolved_transfer_ids = 0 self._drain_pending = True self._drained_at = 0.0 self._pending_cancels: dict[str, Future[Any]] = {} self._transfers = SchedulerTransferTable( - self._consumer_pool_capacity, _LEASE_TTL_SECONDS + config.pool_size, _RESERVATION_TTL_SECONDS ) self._scheduler_pending_work = False self._pushes_to_prepare: dict[str, ECMooncakePushSpec] = {} self._prepared_push_transfer_ids: set[str] = set() - self._event_shard_count = 1 self._event_ready_shards: OrderedDict[str, set[int]] = OrderedDict() - def _cancel_remote( - self, consumer_zmq: str, transfer_id: str, reservation_id: str - ) -> bool: + def _cancel_remote(self, consumer_zmq: str, transfer_id: str) -> bool: pending = None for _ in range(_CANCEL_ATTEMPTS): - pending = self._topology.discover(consumer_zmq) + pending = self._control_client.discover_shards(consumer_zmq) if pending is not None: break if pending is None: @@ -182,7 +120,7 @@ def _cancel_remote( try: result = self._control_client.request( addr, - make_cancel_request(transfer_id, reservation_id), + make_cancel_request(transfer_id), ) except Exception as exc: if error is None: @@ -210,65 +148,25 @@ def _note_awaiting_push( mm_hash, now + self._push_wait_timeout, ) - self._consumer_scheduler_metrics["missing_event"] += 1 if record.state is not SchedulerTransferState.WAITING_EVENT: return assert record.deadline is not None if now < record.deadline: return elapsed = now - record.deadline + self._push_wait_timeout - self._transfers.mark_unavailable( - transfer_id, "push readiness event timed out", now - ) - self._consumer_scheduler_metrics["given_up"] += 1 - self._consumer_scheduler_metrics["stalled"] += 1 - reservation: Any = "unknown" - if self._reservation_zmq_addr is not None: - try: - reservation = self._control_client.request( - self._reservation_zmq_addr, - {"op": "status", "transfer_id": transfer_id}, - ) - except Exception as e: # noqa: BLE001 - diagnostic only - reservation = f"status failed: {e}" + self._transfers.mark_unavailable(transfer_id, now) logger.warning( "EC Mooncake waited %.1fs for a push of mm_hash=%s " - "(transfer_id=%s) that never arrived; worker reservation=%s; " + "(transfer_id=%s) that never arrived; " "requests needing it fail with a retryable error.", elapsed, mm_hash, transfer_id, - reservation, ) def take_unavailable_requests(self) -> set[str]: return self._transfers.take_unavailable_requests() - def _maybe_log_consumer_scheduler_metrics(self) -> None: - now = time.monotonic() - if ( - self._consumer_metrics_log_interval <= 0 - or now - self._consumer_metrics_started_at - < self._consumer_metrics_log_interval - ): - return - missing = self._transfers.count(SchedulerTransferState.WAITING_EVENT) - loading = self._transfers.count(SchedulerTransferState.LOADING) - pending = self._transfers.count(SchedulerTransferState.AVAILABLE) - logger.info( - "EC Mooncake consumer scheduler: decisions=%s, ready=%d, loading=%d, " - "resident=%d, pending_specs=%d, needs_load=%d, missing=%d", - dict(self._consumer_scheduler_metrics), - self._transfers.count(SchedulerTransferState.READY), - loading, - self._transfers.count(SchedulerTransferState.RESIDENT), - pending, - loading, - missing, - ) - self._consumer_scheduler_metrics.clear() - self._consumer_metrics_started_at = now - def _poll_pending_cancels(self) -> None: pending = {} for transfer_id, future in self._pending_cancels.items(): @@ -276,19 +174,16 @@ def _poll_pending_cancels(self) -> None: pending[transfer_id] = future continue try: - cancelled = future.result() + future.result() except Exception: - self._consumer_scheduler_metrics["cancellations_failed"] += 1 logger.warning( "EC Mooncake reservation cancellation failed", exc_info=True ) - else: - key = "cancellations_completed" if cancelled else "cancellations_stale" - self._consumer_scheduler_metrics[key] += 1 self._pending_cancels = pending def _note_shard_ready(self, data: dict[str, Any]) -> bool: - if self._event_shard_count <= 1: + shard_count = self._event_inbox.shard_count + if shard_count <= 1: return True transfer_id = str(data["transfer_id"]) record = self._transfers.get(transfer_id) @@ -301,14 +196,11 @@ def _note_shard_ready(self, data: dict[str, Any]) -> bool: shards = self._event_ready_shards.setdefault(transfer_id, set()) self._event_ready_shards.move_to_end(transfer_id) shards.add(int(shard) if shard is not None else len(shards)) - if len(shards) < self._event_shard_count: - self._consumer_scheduler_metrics["events_awaiting_shards"] += 1 + if len(shards) < shard_count: while len(self._event_ready_shards) > _MAX_PENDING_EVENTS: self._event_ready_shards.popitem(last=False) - self._consumer_scheduler_metrics["events_partial_dropped"] += 1 return False self._event_ready_shards.pop(transfer_id, None) - self._consumer_scheduler_metrics["events_all_shards_ready"] += 1 return True def _forget_shard_readiness(self, transfer_id: str) -> None: @@ -317,26 +209,21 @@ def _forget_shard_readiness(self, transfer_id: str) -> None: def _store_pushed_spec(self, data: dict[str, Any]) -> bool: transfer_id = str(data["transfer_id"]) identifier = str(data["mm_hash"]) - reservation_id = str(data["reservation_id"]) _, accepted = self._transfers.observe_ready( ECMooncakeLoadSpec( mm_hash=identifier, - num_token=0, nbytes=int(data["nbytes"]), shape=tuple(int(value) for value in data["shape"]), dtype=str(data["dtype"]), - pushed=True, transfer_id=transfer_id, - reservation_id=reservation_id, ), - time.monotonic() + _LEASE_TTL_SECONDS, + time.monotonic() + _RESERVATION_TTL_SECONDS, ) return accepted def _queue_cancel( self, transfer_id: str, - reservation_id: str = "", mm_hash: str = "", request_id: str = "", ) -> None: @@ -348,21 +235,16 @@ def _queue_cancel( ): return self._forget_shard_readiness(transfer_id) - if self._reservation_zmq_addr is None: - return self._pending_cancels[transfer_id] = self._control_executor.submit( self._cancel_remote, - self._reservation_zmq_addr, + self._control_addr, transfer_id, - reservation_id, ) def _expire_transfers(self) -> None: now = time.monotonic() - expired, dropped = self._transfers.expire(now, _MAX_TERMINAL_TRANSFER_RECORDS) - self._consumer_scheduler_metrics["cancel_records_dropped"] += dropped + expired = self._transfers.expire(now, _MAX_TERMINAL_TRANSFER_RECORDS) for record in expired: - self._consumer_scheduler_metrics["pending_specs_expired"] += 1 self._queue_cancel(record.transfer_id) def _drain_push_notifications(self) -> None: @@ -373,15 +255,9 @@ def _drain_push_notifications(self) -> None: self._drained_at = now self._poll_pending_cancels() self._expire_transfers() - if self._reservation_zmq_addr is None: - return - events = self._event_inbox.drain(self._reservation_zmq_addr) - self._event_shard_count = self._event_inbox.shard_count + events = self._event_inbox.drain(self._control_addr) for data in events: - identifier = str(data["mm_hash"]) - self._consumer_scheduler_metrics["events_received"] += 1 if data.get("ready"): - self._consumer_scheduler_metrics["events_ready"] += 1 transfer_id = str(data["transfer_id"]) record = self._transfers.get(transfer_id) if record is not None and record.state in { @@ -390,38 +266,23 @@ def _drain_push_notifications(self) -> None: SchedulerTransferState.EXPIRED, SchedulerTransferState.FAILED, }: - self._consumer_scheduler_metrics["events_cancelled"] += 1 continue - if self._transfers.has_state( - identifier, (SchedulerTransferState.READY,) - ): - self._consumer_scheduler_metrics["events_redundant"] += 1 if not self._note_shard_ready(data): continue - if not self._store_pushed_spec(data): - self._consumer_scheduler_metrics["events_duplicate"] += 1 - else: - self._consumer_scheduler_metrics["events_not_ready"] += 1 + self._store_pushed_spec(data) def has_cache_item(self, identifier: str) -> bool: if not self._is_consumer: return False self._drain_push_notifications() - self._maybe_log_consumer_scheduler_metrics() - if self._transfers.has_state(identifier, (SchedulerTransferState.READY,)): - self._consumer_scheduler_metrics["ready"] += 1 - return True - if self._transfers.has_state(identifier, (SchedulerTransferState.LOADING,)): - self._consumer_scheduler_metrics["loading"] += 1 - return False - if self._transfers.has_state(identifier, (SchedulerTransferState.RESIDENT,)): - self._consumer_scheduler_metrics["resident"] += 1 - return True - if self._transfers.has_state(identifier, (SchedulerTransferState.AVAILABLE,)): - self._consumer_scheduler_metrics["pending_spec"] += 1 - return True - self._consumer_scheduler_metrics["missing_event"] += 1 - return False + return self._transfers.has_state( + identifier, + ( + SchedulerTransferState.READY, + SchedulerTransferState.RESIDENT, + SchedulerTransferState.AVAILABLE, + ), + ) def _warn_unresolved_transfer_id( self, request: Any, index: int, where: str @@ -504,22 +365,17 @@ def ensure_cache_available( transfer_id = self._request_transfer_id(request, index) if transfer_id is not None: self._transfers.touch_available( - transfer_id, time.monotonic() + _LEASE_TTL_SECONDS + transfer_id, time.monotonic() + _RESERVATION_TTL_SECONDS ) if mm_hash in local_cache_hashes: continue if self._transfers.has_state(mm_hash, (SchedulerTransferState.READY,)): - self._consumer_scheduler_metrics["ready"] += 1 continue if self._transfers.has_state(mm_hash, (SchedulerTransferState.LOADING,)): - self._consumer_scheduler_metrics["loading"] += 1 all_ready = False continue - if self._transfers.has_state(mm_hash, (SchedulerTransferState.RESIDENT,)): - self._consumer_scheduler_metrics["resident_hit"] += 1 record = self._transfers.begin_load( mm_hash, - request.get_num_encoder_embeds(index), transfer_id, request.request_id, ) @@ -589,17 +445,15 @@ def build_connector_meta( for mm_hash in scheduler_output.free_encoder_mm_hashes: self._transfers.release_ready(mm_hash, time.monotonic()) for transfer_id in self._transfers.drain_orphaned(): - self._consumer_scheduler_metrics["reservations_orphaned"] += 1 self._queue_cancel(transfer_id) meta = ECMooncakeConnectorMetadata() for push_spec in self._pushes_to_prepare.values(): - meta.add_push(push_spec) + meta.pushes.append(push_spec) self._pushes_to_prepare.clear() for record in self._transfers.take_loads_to_dispatch(): assert record.spec is not None - meta.add_load(record.spec) + meta.loads.append(record.spec) self._poll_pending_cancels() - self._maybe_log_consumer_scheduler_metrics() self._drain_pending = True return meta @@ -609,16 +463,11 @@ def update_connector_output(self, connector_output: ECConnectorOutput) -> None: return for mm_hash in meta.loaded: self._transfers.complete_load(mm_hash) - self._consumer_scheduler_metrics["loads_completed"] += 1 for mm_hash in meta.failed_loads: - self._transfers.fail_load( - mm_hash, "worker failed to load", time.monotonic() - ) - self._consumer_scheduler_metrics["loads_failed"] += 1 + self._transfers.fail_load(mm_hash, time.monotonic()) for mm_hash in meta.reclaimed: self._transfers.reclaim(mm_hash, time.monotonic()) - self._consumer_scheduler_metrics["resident_reclaimed"] += 1 - self._scheduler_pending_work = meta.pending_loads or meta.pending_saves + self._scheduler_pending_work = meta.pending_saves def has_pending_push_work(self) -> bool: return self._scheduler_pending_work diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py index 18e547da7167..a8f48867ef1c 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py @@ -34,18 +34,7 @@ class SchedulerTransferState(Enum): @dataclass class SchedulerTransfer: - """Track one transfer as observed by the Scheduler. - - Attributes: - transfer_id: Cross-process identity of the transfer. - request_id: Request currently waiting for the transfer. - mm_hash: Stable identifier of the encoder-cache item. - state: Current Scheduler lifecycle state. - spec: Load metadata once the Consumer reports the tensor ready. - deadline: Expiry time for waiting, available, or terminal records. - last_error: Last terminal error associated with the transfer. - notified_requests: Requests already told that the item is unavailable. - """ + """Track one transfer as observed by the Scheduler.""" transfer_id: str request_id: str @@ -53,46 +42,9 @@ class SchedulerTransfer: state: SchedulerTransferState spec: ECMooncakeLoadSpec | None deadline: float | None - last_error: str | None = None notified_requests: set[str] = field(default_factory=set, repr=False) -class InvalidSchedulerTransferTransition(RuntimeError): - """Raised when code attempts an unsupported Scheduler state transition.""" - - pass - - -_ALLOWED_TRANSITIONS = { - SchedulerTransferState.WAITING_EVENT: { - SchedulerTransferState.AVAILABLE, - SchedulerTransferState.UNAVAILABLE, - SchedulerTransferState.CANCELLED, - }, - SchedulerTransferState.AVAILABLE: { - SchedulerTransferState.LOADING, - SchedulerTransferState.EXPIRED, - SchedulerTransferState.CANCELLED, - }, - SchedulerTransferState.LOADING: { - SchedulerTransferState.READY, - SchedulerTransferState.FAILED, - SchedulerTransferState.CANCELLED, - }, - SchedulerTransferState.READY: { - SchedulerTransferState.RESIDENT, - SchedulerTransferState.EXPIRED, - }, - SchedulerTransferState.RESIDENT: { - SchedulerTransferState.LOADING, - SchedulerTransferState.EXPIRED, - }, - SchedulerTransferState.UNAVAILABLE: {SchedulerTransferState.CANCELLED}, - SchedulerTransferState.EXPIRED: {SchedulerTransferState.CANCELLED}, - SchedulerTransferState.FAILED: set(), - SchedulerTransferState.CANCELLED: set(), -} - _TERMINAL_STATES = { SchedulerTransferState.UNAVAILABLE, SchedulerTransferState.EXPIRED, @@ -104,13 +56,9 @@ class InvalidSchedulerTransferTransition(RuntimeError): class SchedulerTransferTable: """Own Scheduler transfer state, lookup indexes, and dispatch queues. - Attributes: - _resident_capacity: Maximum bytes represented by resident records. - _tombstone_ttl: Retention time for terminal records. - _records: Ordered transfer records keyed by transfer ID. - _hash_index: Transfer IDs grouped by encoder-cache hash. - _loads_to_dispatch: Ordered IDs awaiting Worker metadata emission. - _unavailable_requests: Requests awaiting retryable failure reporting. + The Scheduler is single-threaded, so direct transitions are sufficient; + the Producer and Consumer managers retain stricter transition matrices + because callbacks and control requests can race there. """ def __init__(self, resident_capacity: int, tombstone_ttl: float) -> None: @@ -157,22 +105,24 @@ def first_for_hash( mm_hash: str, states: Iterable[SchedulerTransferState], ) -> SchedulerTransfer | None: - return next(iter(self.records_for_hash(mm_hash, states)), None) + states = tuple(states) + if len(states) == 1: + wanted_state = states[0] + for transfer_id in self._hash_index.get(mm_hash, ()): + record = self._records.get(transfer_id) + if record is not None and record.state is wanted_state: + return record + return None + wanted_states = set(states) + for transfer_id in self._hash_index.get(mm_hash, ()): + record = self._records.get(transfer_id) + if record is not None and record.state in wanted_states: + return record + return None def has_state(self, mm_hash: str, states: Iterable[SchedulerTransferState]) -> bool: return self.first_for_hash(mm_hash, states) is not None - def count(self, state: SchedulerTransferState) -> int: - return sum(record.state is state for record in self._records.values()) - - @property - def resident_bytes(self) -> int: - return sum( - record.spec.nbytes - for record in self._records.values() - if record.state is SchedulerTransferState.RESIDENT and record.spec - ) - def wait_for_event( self, transfer_id: str, @@ -202,7 +152,7 @@ def wait_for_event( def observe_ready( self, spec: ECMooncakeLoadSpec, deadline: float ) -> tuple[SchedulerTransfer, bool]: - transfer_id = spec.transfer_id or spec.mm_hash + transfer_id = spec.transfer_id record = self._records.get(transfer_id) if record is None: record = SchedulerTransfer( @@ -231,7 +181,6 @@ def touch_available(self, transfer_id: str, deadline: float) -> None: def begin_load( self, mm_hash: str, - num_token: int, transfer_id: str | None = None, request_id: str = "", ) -> SchedulerTransfer | None: @@ -253,7 +202,6 @@ def begin_load( return None if not record.request_id: record.request_id = request_id - record.spec = replace(record.spec, num_token=num_token) record.deadline = None self._transition(record, SchedulerTransferState.LOADING) self._loads_to_dispatch[record.transfer_id] = None @@ -276,20 +224,15 @@ def complete_load(self, mm_hash: str) -> bool: self._transition(record, SchedulerTransferState.READY) return True - def fail_load(self, mm_hash: str, error: str, now: float) -> bool: + def fail_load(self, mm_hash: str, now: float) -> bool: record = self.first_for_hash(mm_hash, (SchedulerTransferState.LOADING,)) if record is None: return self.has_state(mm_hash, (SchedulerTransferState.FAILED,)) - self._transition(record, SchedulerTransferState.FAILED, error, now=now) + self._transition(record, SchedulerTransferState.FAILED, now=now) return True def release_ready(self, mm_hash: str, now: float) -> None: - ready = [ - record - for record in self._records.values() - if record.mm_hash == mm_hash - and record.state is SchedulerTransferState.READY - ] + ready = self.records_for_hash(mm_hash, (SchedulerTransferState.READY,)) if ready: canonical = ready[-1] for record in self.records_for_hash( @@ -301,7 +244,7 @@ def release_ready(self, mm_hash: str, now: float) -> None: if canonical.spec is None: self._transition(canonical, SchedulerTransferState.EXPIRED, now=now) else: - canonical.spec = replace(canonical.spec, num_token=0, local=True) + canonical.spec = replace(canonical.spec, local=True) self._transition(canonical, SchedulerTransferState.RESIDENT) self._evict_residents(now) @@ -315,9 +258,9 @@ def reclaim(self, mm_hash: str, now: float) -> None: else: self._transition(record, SchedulerTransferState.EXPIRED, now=now) - def mark_unavailable(self, transfer_id: str, error: str, now: float) -> None: + def mark_unavailable(self, transfer_id: str, now: float) -> None: record = self._records[transfer_id] - self._transition(record, SchedulerTransferState.UNAVAILABLE, error, now=now) + self._transition(record, SchedulerTransferState.UNAVAILABLE, now=now) if record.request_id: self._notify_unavailable(record, record.request_id) @@ -350,13 +293,10 @@ def cancel( self._loads_to_dispatch.pop(transfer_id, None) return True - def expire( - self, now: float, terminal_limit: int - ) -> tuple[list[SchedulerTransfer], int]: + def expire(self, now: float, terminal_limit: int) -> list[SchedulerTransfer]: if terminal_limit < 0: raise ValueError("terminal_limit must be non-negative") expired = [] - dropped = 0 for record in list(self._records.values()): if record.deadline is None or record.deadline > now: continue @@ -364,13 +304,11 @@ def expire( self._transition( record, SchedulerTransferState.EXPIRED, - "lease expired", now=now, ) expired.append(record) elif record.state in _TERMINAL_STATES: self._remove(record.transfer_id) - dropped += 1 terminal_ids = [ record.transfer_id for record in self._records.values() @@ -379,8 +317,7 @@ def expire( excess = max(0, len(terminal_ids) - terminal_limit) for transfer_id in terminal_ids[:excess]: self._remove(transfer_id) - dropped += 1 - return expired, dropped + return expired def take_unavailable_requests(self) -> set[str]: unavailable = self._unavailable_requests @@ -411,17 +348,13 @@ def _transition( self, record: SchedulerTransfer, state: SchedulerTransferState, - error: str | None = None, now: float | None = None, ) -> None: - if state not in _ALLOWED_TRANSITIONS[record.state]: - raise InvalidSchedulerTransferTransition( - f"Cannot transition {record.transfer_id!r} from " - f"{record.state.name} to {state.name}" - ) if state in _TERMINAL_STATES: if now is None: raise ValueError("Terminal transition requires a timestamp") + # Keep a bounded tombstone so late events and repeat request IDs + # remain idempotent instead of reviving a finished transfer. record.deadline = now + self._tombstone_ttl if state is SchedulerTransferState.EXPIRED and record.state in ( SchedulerTransferState.READY, @@ -429,17 +362,23 @@ def _transition( ): self._orphaned.append(record.transfer_id) record.state = state - record.last_error = error self._records.move_to_end(record.transfer_id) def _evict_residents(self, now: float) -> None: - while self.resident_bytes > self._resident_capacity: - record = next( - record - for record in self._records.values() - if record.state is SchedulerTransferState.RESIDENT - ) + residents: list[SchedulerTransfer] = [] + resident_bytes = 0 + for record in self._records.values(): + if record.state is not SchedulerTransferState.RESIDENT: + continue + residents.append(record) + if record.spec is not None: + resident_bytes += record.spec.nbytes + for record in residents: + if resident_bytes <= self._resident_capacity: + break self._transition(record, SchedulerTransferState.EXPIRED, now=now) + if record.spec is not None: + resident_bytes -= record.spec.nbytes def _remove(self, transfer_id: str) -> None: record = self._records.pop(transfer_id, None) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py index 469b2a280bd5..f5adcb772c7c 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py @@ -17,21 +17,24 @@ logger = init_logger(__name__) +_MOONCAKE_IMPORT_ERROR: ImportError | None = None try: from mooncake.engine import TransferEngine -except ImportError: +except ImportError as error: TransferEngine = None # type: ignore[misc, assignment] + _MOONCAKE_IMPORT_ERROR = error + + +def ensure_mooncake_available() -> None: + if _MOONCAKE_IMPORT_ERROR is not None: + raise ImportError( + "Install mooncake-transfer-engine to use ECMooncakeConnector." + ) from _MOONCAKE_IMPORT_ERROR @dataclass class _SourceRegistration: - """Retain one transient source range while batches reference it. - - Attributes: - tensor: Tensor keeping the registered storage alive. - nbytes: Exact registered byte length. - users: Number of active acquisitions of this address. - """ + """Retain one transient source range while batches reference it.""" tensor: torch.Tensor nbytes: int @@ -39,18 +42,7 @@ class _SourceRegistration: class MooncakeTransfer: - """Own a lazy Mooncake engine and transient memory registrations. - - Attributes: - _hostname: Address advertised in the Mooncake session identifier. - _protocol: Transport protocol used to initialize ``TransferEngine``. - _engine: Lazily initialized Mooncake engine. - _engine_lock: Lock serializing first engine initialization. - _source_registrations: Reference-counted transient source ranges. - _pending_unregister: Tensors retained after an unregister failure. - _registration_lock: Lock protecting registration ownership. - _closed: Whether final data-plane cleanup has begun. - """ + """Own a lazy Mooncake engine and transient memory registrations.""" def __init__(self, hostname: str, protocol: str) -> None: self._hostname = hostname diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py index bd85eddae1b6..70a079c5b21e 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py @@ -12,36 +12,28 @@ import math import threading import time -from collections import Counter from collections.abc import Callable from concurrent.futures import Future, ThreadPoolExecutor -from contextlib import suppress -from dataclasses import dataclass, field from functools import partial from typing import TYPE_CHECKING, Any, TypeVar, cast import torch -from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole -from vllm.distributed.ec_transfer.ec_connector.mooncake._availability import ( - ensure_mooncake_available, +from vllm.distributed.ec_transfer.ec_connector.mooncake.config import ( + _RESERVATION_TTL_SECONDS, + MooncakeECConfig, ) -from vllm.distributed.ec_transfer.ec_connector.mooncake.config import MooncakeECConfig from vllm.distributed.ec_transfer.ec_connector.mooncake.control import ( ConsumerControlServer, ControlClient, - ControlCompletion, - ShardTopology, make_cancel_request, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.memory import ( ConsumerMemoryPool, - MemoryAllocation, ProducerMemoryPool, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( ECMooncakeConnectorMetadata, - ECMooncakeLoadSpec, ECMooncakePushSpec, ECMooncakeWorkerMetadata, ) @@ -50,12 +42,12 @@ ProducerPushRecord, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.reservation import ( - CancellationOutcome, ConsumerReservationManager, ConsumerReservationState, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( MooncakeTransfer, + ensure_mooncake_available, ) from vllm.logger import init_logger from vllm.utils.network_utils import get_ip @@ -67,43 +59,11 @@ if TYPE_CHECKING: from vllm.config import VllmConfig -_LEASE_TTL_SECONDS = 300 -_RESERVATION_REFRESH_SECONDS = _LEASE_TTL_SECONDS / 2 +_RESERVATION_REFRESH_SECONDS = _RESERVATION_TTL_SECONDS / 2 _MAX_CANCELLED_TRANSFER_IDS = 1 << 16 _CANCEL_ATTEMPTS = 2 -_PUSH_STAGES = ( - "reserve", - "cuda", - "register", - "rdma", - "unregister", - "complete", -) - - -@dataclass -class _PushPerfWindow: - """Accumulate Producer batch metrics between periodic log messages. - - Attributes: - started_at: Monotonic start time of the aggregation window. - batches: Number of completed push batches. - items: Number of push records included in those batches. - bytes: Number of tensor bytes written over the data plane. - skipped_items: Items satisfied by cache or cancellation without a write. - failures: Number of batches that ended in failure. - stage_totals_ms: Accumulated time for every push stage. - stage_max_ms: Maximum observed time for every push stage. - """ - - started_at: float = field(default_factory=time.monotonic) - batches: int = 0 - items: int = 0 - bytes: int = 0 - skipped_items: int = 0 - failures: int = 0 - stage_totals_ms: dict[str, float] = field(default_factory=dict) - stage_max_ms: dict[str, float] = field(default_factory=dict) +_TRANSFER_WORKERS = 4 +_CONTROL_WORKERS = 8 class _FanoutError(RuntimeError): @@ -114,106 +74,31 @@ def __init__(self, error: BaseException, results: list[Any | None]) -> None: super().__init__(str(error)) -class _ReservationFanoutError(RuntimeError): - """Expose partial reservations for precise idempotent cleanup retries.""" - - def __init__( - self, - error: BaseException, - partial_reservations: list[dict[str, Any]], - ) -> None: - super().__init__(str(error)) - self.partial_reservations = partial_reservations - - class ECMooncakeWorker: """Orchestrate consumer reservations and producer push batches. - ``mooncake_protocol`` selects the transfer protocol. Consumer workers use - ``consumer_buffer_pool_size`` and ``reservation_zmq_port`` for their - registered receive arena and rank-local control endpoint. Producers use - ``producer_buffer_pool_size`` for staging. ``transfer_max_workers`` and - ``control_max_workers`` bound the two executor pools; the transfer and - consumer metrics intervals control aggregate logging. - Consumers may use TP, PP, and DP. Only the first PP stage receives encoder outputs, and each TP rank exposes a consecutive control port and receives the same source concurrently. Producers remain unsharded and unreplicated. With DP, the caller must route both halves of a request to the same replica and pass that replica's control address to the producer. - - Attributes: - is_producer: Whether this Worker originates encoder-cache pushes. - is_consumer: Whether this Worker accepts encoder-cache pushes. - _buffer_device: Device requested for registered memory pools. - _reservation_zmq_port: Base Consumer control port for this DP replica. - _transfer: Owner of the Mooncake engine and memory registrations. - _consumer_worker_metrics: Consumer lifecycle metric counters. - _consumer_memory: Registered receive slab and resident cache. - _reservations: Consumer destination reservation state manager. - _consumer_rank_resolved: Whether TP/PP placement has been discovered. - _is_receiving_rank: Whether this PP stage owns encoder outputs. - _tp_rank: Tensor-parallel rank used to derive the local control port. - _tp_size: Number of Consumer tensor-parallel destination shards. - _control_server: Rank-local Consumer reservation server. - _consumer_metrics_log_interval: Consumer metrics log interval. - _consumer_metrics_started_at: Start time of the Consumer metric window. - _producer_memory: Registered Producer source staging slab. - _transfer_metrics_log_interval: Producer performance log interval. - _control_client: Client for remote Consumer control operations. - _topology: Discovery cache for remote Consumer TP shards. - _producer_metrics: Producer lifecycle metric counters. - _io_executor: Executor that owns transfer and cancellation batches. - _control_executor: Executor that creates remote reservations. - _shard_pool: Lazily created executor for concurrent TP-shard work. - _shard_pool_lock: Lock protecting shard-pool initialization. - _producer_pushes: Producer lifecycle and source-ownership manager. - _push_perf_lock: Lock protecting Producer performance counters. - _push_perf: Current Producer performance aggregation window. - _active_transfer_batches: Batches currently executing data-plane work. - _queued_transfer_batches: Batches submitted but not yet executing. - _completed_loads: Successful Consumer loads awaiting reporting. - _failed_loads: Failed Consumer loads awaiting reporting. - _shutdown: Whether Worker resource shutdown has started. """ - @classmethod - def from_vllm_config(cls, vllm_config: VllmConfig) -> ECMooncakeWorker: + def __init__(self, vllm_config: VllmConfig) -> None: ensure_mooncake_available() - config = MooncakeECConfig.from_vllm_config(vllm_config, ECConnectorRole.WORKER) - hostname = get_ip() - control_client = ControlClient(config.control_timeout_ms) - try: - return cls( - config, - hostname, - control_client, - ShardTopology(control_client), - ) - except Exception: - control_client.close() - raise - - def __init__( - self, - config: MooncakeECConfig, - hostname: str, - control_client: ControlClient, - topology: ShardTopology, - ) -> None: + config = MooncakeECConfig.from_vllm_config(vllm_config) self.is_producer = config.is_producer self.is_consumer = config.is_consumer self._buffer_device = config.buffer_device - self._reservation_zmq_port = config.reservation_port - self._transfer = MooncakeTransfer(hostname, config.protocol) - self._consumer_worker_metrics: Counter[str] = Counter() + self._control_port = config.control_port + self._transfer = MooncakeTransfer(get_ip(), config.protocol) self._consumer_memory = ConsumerMemoryPool( - config.consumer_pool_size, + config.pool_size, self._transfer, ) self._reservations = ConsumerReservationManager( self._consumer_memory, - _LEASE_TTL_SECONDS, + _RESERVATION_TTL_SECONDS, _MAX_CANCELLED_TRANSFER_IDS, ) self._consumer_rank_resolved = False @@ -221,32 +106,24 @@ def __init__( self._tp_rank = 0 self._tp_size = 1 self._control_server: ConsumerControlServer | None = None - self._consumer_metrics_log_interval = config.consumer_metrics_log_interval - self._consumer_metrics_started_at = time.monotonic() # Worker producer self._producer_memory = ProducerMemoryPool( - config.producer_pool_size, + config.pool_size, self._transfer, ) - self._transfer_metrics_log_interval = config.transfer_metrics_log_interval - self._control_client = control_client - self._topology = topology - self._producer_metrics: Counter[str] = Counter() + self._control_client = ControlClient(config.control_timeout_ms) + self._shutdown_drain_timeout_s = config.control_timeout_ms / 1000 self._io_executor = ThreadPoolExecutor( - max_workers=config.transfer_workers, + max_workers=_TRANSFER_WORKERS, thread_name_prefix="ec-mooncake-transfer", ) self._control_executor = ThreadPoolExecutor( - max_workers=config.control_workers, + max_workers=_CONTROL_WORKERS, thread_name_prefix="ec-mooncake-control", ) self._shard_pool: ThreadPoolExecutor | None = None self._shard_pool_lock = threading.Lock() self._producer_pushes = ProducerPushManager() - self._push_perf_lock = threading.Lock() - self._push_perf = _PushPerfWindow() - self._active_transfer_batches = 0 - self._queued_transfer_batches = 0 self._completed_loads: set[str] = set() self._failed_loads: set[str] = set() self._shutdown = False @@ -271,41 +148,28 @@ def _resolve_consumer_rank(self) -> None: self._is_receiving_rank = True def start_services(self) -> None: - if ( - not self.is_consumer - or self._reservation_zmq_port is None - or self._control_server is not None - ): + if not self.is_consumer or self._control_server is not None: return self._resolve_consumer_rank() if not self._is_receiving_rank: # Later pipeline stages hold no encoder outputs, so they need # neither a receive pool nor a control channel. return - raw_device = self._buffer_device - device_name = ( - raw_device.lower() if isinstance(raw_device, str) and raw_device else "cuda" - ) - self._consumer_memory.prepare( - torch.device(device_name), - receiving_rank=self._is_receiving_rank, - allow_host=True, - ) + self._consumer_memory.prepare(torch.device(self._buffer_device)) consumer_pool = self._consumer_memory.tensor if consumer_pool is None: raise RuntimeError( "Mooncake push mode requires a registered consumer buffer pool." ) - base_port = self._reservation_zmq_port + base_port = self._control_port self._control_server = ConsumerControlServer( "0.0.0.0", base_port + self._tp_rank, self._reserve_push_destination, self._push_status, - self._complete_push, - self._cancel_push, - self._expire_push_reservations, - self._consumer_metrics_log_interval, + self._reservations.complete, + self._reservations.cancel, + self._reservations.expire, peer_ports=[base_port + rank for rank in range(self._tp_size)], device=consumer_pool.device, ) @@ -316,62 +180,6 @@ def start_services(self) -> None: self._control_server = None raise - def _maybe_log_consumer_worker_metrics(self) -> None: - now = time.monotonic() - if ( - self._consumer_metrics_log_interval <= 0 - or now - self._consumer_metrics_started_at - < self._consumer_metrics_log_interval - ): - return - with self._consumer_memory.lock: - reservations = self._reservations.active_records() - ready = [ - record.mm_hash - for record in reservations - if record.state is ConsumerReservationState.READY - ] - pending = [ - record.mm_hash - for record in reservations - if record.state is not ConsumerReservationState.READY - ] - metrics = dict(self._consumer_worker_metrics) - self._consumer_worker_metrics.clear() - metrics.update(self._consumer_memory.take_metrics()) - residents, live, retired, pending_frees = self._consumer_memory.stats() - oldest_reservation_ms = max( - ((now - reservation.created_at) * 1000 for reservation in reservations), - default=0.0, - ) - logger.info( - "EC Mooncake consumer worker: lifecycle=%s, reservations_ready=%d, " - "reservations_pending=%d, residents=%d, live=%d, retired=%d, " - "pending_frees=%d, " - "oldest_reservation_ms=%.1f, ready_hashes=%s, pending_hashes=%s", - metrics, - len(ready), - len(pending), - residents, - live, - retired, - pending_frees, - oldest_reservation_ms, - [value[:16] for value in ready[:5]], - [value[:16] for value in pending[:5]], - ) - self._consumer_metrics_started_at = now - - def _expire_push_reservations(self) -> int: - return self._record_expiry_metrics(self._reservations.expire()) - - def _record_expiry_metrics(self, counts: tuple[int, int, int]) -> int: - expired, deferred, tombstones_dropped = counts - self._consumer_worker_metrics["reservations_expired"] += expired - self._consumer_worker_metrics["cancellations_deferred"] += deferred - self._consumer_worker_metrics["cancel_records_dropped"] += tombstones_dropped - return expired - def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: transfer_id = str(payload["transfer_id"]) mm_hash = str(payload["mm_hash"]) @@ -385,18 +193,16 @@ def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: if expected_nbytes != nbytes: raise ValueError("shape and dtype do not match nbytes") - self._expire_push_reservations() - reservation, should_write, reused, expiry_counts = self._reservations.reserve( + self._reservations.expire() + reservation, should_write = self._reservations.reserve( transfer_id, mm_hash, nbytes, shape, dtype_name, dtype ) - self._record_expiry_metrics(expiry_counts) if reservation is None: raise RuntimeError("EC consumer buffer pool is full") if reservation.state in { ConsumerReservationState.CANCEL_PENDING, ConsumerReservationState.CANCELLED, }: - self._consumer_worker_metrics["reservations_cancelled_early"] += 1 return { "reservation_id": "", "dst_session": "", @@ -406,17 +212,6 @@ def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: "ready": False, "cancelled": True, } - if reused: - key = ( - "reservations_reused_ready" - if reservation.state is ConsumerReservationState.READY - else "reservations_reused_pending" - ) - self._consumer_worker_metrics[key] += 1 - elif reservation.lease is not None: - self._consumer_worker_metrics["reservations_cached"] += 1 - else: - self._consumer_worker_metrics["reservations_created"] += 1 assert reservation.allocation is not None return { @@ -443,51 +238,6 @@ def _push_status(self, transfer_id: str) -> dict[str, Any] | None: "dtype": reservation.dtype, } - def _complete_push( - self, transfer_id: str, reservation_id: str - ) -> ControlCompletion: - result = self._reservations.complete(transfer_id, reservation_id) - if not result.accepted: - self._consumer_worker_metrics["completions_rejected"] += 1 - elif result.repeated: - self._consumer_worker_metrics["completions_repeated"] += 1 - else: - self._consumer_worker_metrics["completions_accepted"] += 1 - if result.discarded: - self._consumer_worker_metrics["reservations_discarded"] += 1 - return ControlCompletion(result.accepted, result.became_ready) - - def _cancel_push( - self, - transfer_id: str, - reservation_id: str, - abandon: bool = False, - refresh: bool = False, - ) -> bool: - outcome, tombstones_dropped = self._reservations.cancel( - transfer_id, reservation_id, abandon, refresh - ) - metrics = { - CancellationOutcome.REJECTED: "cancellations_rejected", - CancellationOutcome.PRE_RESERVED: "cancellations_pre_reserved", - CancellationOutcome.DEFERRED: "cancellations_deferred", - CancellationOutcome.CANCELLED: "reservations_cancelled", - } - self._consumer_worker_metrics[metrics[outcome]] += 1 - self._consumer_worker_metrics["cancel_records_dropped"] += tombstones_dropped - return outcome is not CancellationOutcome.REJECTED - - def _take_pushed_tensor( - self, spec: ECMooncakeLoadSpec - ) -> tuple[torch.Tensor, MemoryAllocation]: - try: - allocation = self._reservations.take(spec.transfer_id, spec.mm_hash) - except RuntimeError: - self._consumer_worker_metrics["takes_rejected"] += 1 - raise - self._consumer_worker_metrics["reservations_taken"] += 1 - return allocation.tensor, allocation - def _shard_executor(self) -> ThreadPoolExecutor: """Use a separate pool so nested shard fan-out cannot deadlock.""" with self._shard_pool_lock: @@ -520,6 +270,11 @@ def _run_fanout( tasks: list[Callable[[], _T]], on_submit: Callable[[int, Future[_T]], None] | None = None, ) -> list[_T]: + """Run every started shard task and retain partial results on failure. + + Waiting for all submitted tasks is what makes source-memory release and + partial-reservation cleanup safe after one shard fails. + """ if not tasks: return [] futures: list[tuple[int, Future[_T]]] = [] @@ -564,8 +319,13 @@ def _cancel_reservations( return def cancel(shard: dict[str, Any]) -> dict[str, Any]: + addr = shard.get("addr") + if not isinstance(addr, str) or not addr: + raise RuntimeError( + "EC reservation is missing a confirmed shard address" + ) result = self._control_client.request( - str(shard.get("addr", spec.consumer_zmq)), + addr, make_cancel_request( spec.transfer_id, str(shard.get("reservation_id", "")), @@ -611,19 +371,38 @@ def _retry_cancel_reservations( def _reserve_remote(self, spec: ECMooncakePushSpec) -> list[dict[str, Any]]: """Reserve a destination on every shard of the consumer.""" - shards = self._topology.shards(spec.consumer_zmq) + shards = None + for _ in range(_CANCEL_ATTEMPTS): + shards = self._control_client.discover_shards(spec.consumer_zmq) + if shards is not None: + break + if shards is None: + raise RuntimeError( + f"Could not discover every EC consumer shard at {spec.consumer_zmq}" + ) tasks: list[Callable[[], dict[str, Any]]] = [ partial(self._reserve_one, addr, spec) for addr in shards ] try: return self._run_fanout(tasks) except _FanoutError as exc: - successful = [result for result in exc.results if isinstance(result, dict)] + # Keep one cleanup entry per confirmed shard. A missing result + # means the reservation outcome is unknown, so transfer-level + # cancellation with an empty reservation ID is the only safe + # idempotent action for that exact address. + reservations = [] + for addr, result in zip(shards, exc.results): + if isinstance(result, dict): + result["addr"] = addr + reservations.append(result) + else: + reservations.append({"addr": addr, "reservation_id": ""}) + exc.results[:] = reservations try: - self._retry_cancel_reservations(spec, successful) + self._retry_cancel_reservations(spec, reservations) except _FanoutError as cleanup_error: - raise _ReservationFanoutError(exc, successful) from cleanup_error - raise _ReservationFanoutError(exc, successful) from exc + raise exc from cleanup_error + raise def _refresh_remote_reservations( self, @@ -653,9 +432,8 @@ def _refresh_remote_reservations( @staticmethod def _validate_push_source(push: ProducerPushRecord) -> None: - source = push.source - assert source is not None - tensor = source.tensor + tensor = push.source_tensor + assert tensor is not None spec = push.spec if tuple(tensor.shape) != tuple(spec.shape): raise ValueError(f"EC source shape mismatch for mm_hash={spec.mm_hash}") @@ -700,9 +478,7 @@ def start_load_caches( # Later pipeline stages never gather multimodal embeddings. return self._transfer.ensure_ready() - raw_buf = self._buffer_device - buf = raw_buf.lower() if isinstance(raw_buf, str) and raw_buf else "cuda" - if buf == "cuda" and not torch.accelerator.is_available(): + if self._buffer_device == "cuda" and not torch.accelerator.is_available(): raise RuntimeError( "ECMooncakeConnector requires CUDA for ec_buffer_device=cuda" ) @@ -710,27 +486,22 @@ def start_load_caches( for spec in metadata.loads: if spec.mm_hash in encoder_cache: - if spec.pushed: + if not spec.local: # The spec's id is one shard's; cancel by transfer. - self._cancel_push(spec.transfer_id, "") + self._reservations.cancel(spec.transfer_id, "") self._completed_loads.add(spec.mm_hash) continue if spec.local: tensor = self._consumer_memory.take_resident( spec.mm_hash, tuple(spec.shape), spec.dtype ) - elif spec.pushed: + else: try: - tensor, _ = self._take_pushed_tensor(spec) + allocation = self._reservations.take(spec.transfer_id, spec.mm_hash) + tensor = allocation.tensor except RuntimeError as e: logger.warning("EC Mooncake pushed load failed: %s", e) tensor = None - else: - logger.warning( - "EC Mooncake load for mm_hash=%s has no transfer to take", - spec.mm_hash, - ) - tensor = None if tensor is None: self._failed_loads.add(spec.mm_hash) else: @@ -739,23 +510,12 @@ def start_load_caches( def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: started_at = time.monotonic() - with self._push_perf_lock: - self._queued_transfer_batches -= 1 - self._active_transfer_batches += 1 - - queue_waits_ms = [] - for push in pushes: - assert push.source_at is not None - queue_waits_ms.append(max(0, started_at - push.source_at) * 1000) - stage_ms = {"queue": sum(queue_waits_ms), **dict.fromkeys(_PUSH_STAGES, 0.0)} ready: list[tuple[ProducerPushRecord, dict[str, Any]]] = [] written_pushes: dict[str, ProducerPushRecord] = {} - failed = False failure: Exception | None = None try: for push in pushes: self._validate_push_source(push) - stage_started_at = time.monotonic() reservations = self._producer_pushes.resolve_reservations(push) stale = [ index @@ -770,7 +530,6 @@ def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: push.spec, reservations, push ) self._producer_pushes.replace_reservations(push, reservations) - stage_ms["reserve"] += (time.monotonic() - stage_started_at) * 1000 self._producer_pushes.begin_writing(push) writable = [ shard @@ -779,14 +538,12 @@ def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: and not shard.get("cancelled", False) and shard.get("write", True) ] - source = push.source + source = push.source_tensor assert source is not None - if writable and source.ready_event is not None: - stage_started_at = time.monotonic() - source.ready_event.synchronize() - stage_ms["cuda"] += (time.monotonic() - stage_started_at) * 1000 + if writable and push.source_event is not None: + push.source_event.synchronize() for shard in writable: - if int(shard["nbytes"]) != source.tensor.nbytes: + if int(shard["nbytes"]) != source.nbytes: raise RuntimeError( "Reserved EC size does not match tensor for " f"mm_hash={push.spec.mm_hash}" @@ -800,12 +557,11 @@ def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: for index, push in enumerate(written_pushes.values()) } tensors = [ - push.source.tensor + cast(torch.Tensor, push.source_tensor) for push in written_pushes.values() - if push.source + if push.source_tensor is not None ] lengths = [tensor.nbytes for tensor in tensors] - stage_started_at = time.monotonic() staged = self._producer_memory.stage(tensors) registered_sources: list[int] = [] if staged is not None: @@ -819,7 +575,6 @@ def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: sources = tensors registered_sources = self._transfer.acquire_sources(tensors) addresses = [tensor.data_ptr() for tensor in sources] - stage_ms["register"] = (time.monotonic() - stage_started_at) * 1000 try: by_session: dict[str, list[tuple[int, int]]] = {} session_records: dict[str, dict[str, ProducerPushRecord]] = {} @@ -831,7 +586,6 @@ def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: session_records.setdefault(session, {})[ push.spec.transfer_id ] = push - stage_started_at = time.monotonic() def write(session: str, items: list[tuple[int, int]]) -> None: self._transfer.write( @@ -854,24 +608,16 @@ def track_write(index: int, future: Future[None]) -> None: partial(write, *session) for session in sessions ] self._run_fanout(writes, track_write) - stage_ms["rdma"] = (time.monotonic() - stage_started_at) * 1000 finally: - stage_started_at = time.monotonic() if staged is not None: self._producer_memory.release(staged) self._transfer.release_sources(registered_sources) - stage_ms["unregister"] = ( - time.monotonic() - stage_started_at - ) * 1000 self._producer_pushes.begin_notifying(pushes) - stage_started_at = time.monotonic() self._notify_completions(ready) - stage_ms["complete"] = (time.monotonic() - stage_started_at) * 1000 self._producer_pushes.complete(pushes) except Exception as exc: # Report asynchronously; raising here would fail EngineCore. - failed = True failure = exc logger.exception( "EC Mooncake push batch failed for mm_hashes=%s", @@ -882,15 +628,6 @@ def track_write(index: int, future: Future[None]) -> None: finally: if failure is not None: self._producer_pushes.fail(pushes, failure) - stage_ms["total"] = (time.monotonic() - started_at) * 1000 - self._record_push_perf( - stage_ms, - stage_max_ms={"queue": max(queue_waits_ms, default=0.0)}, - item_count=len(pushes), - byte_count=sum(push.spec.nbytes for push in written_pushes.values()), - skipped_items=len(pushes) - len(written_pushes), - failed=failed, - ) def _notify_completions( self, notifications: list[tuple[ProducerPushRecord, dict[str, Any]]] @@ -949,22 +686,22 @@ def track(index: int, future: Future[None]) -> None: def _known_reservations(record: ProducerPushRecord) -> list[dict[str, Any]]: if record.reservations: return list(record.reservations) - reservations: list[dict[str, Any]] = [] - for future in record.reservation_futures: - try: - reservations.extend(future.result()) - except _ReservationFanoutError as exc: - reservations.extend(exc.partial_reservations) - except Exception: - continue - return reservations + try: + return list(record.reservation_future.result()) + except _FanoutError as exc: + return [result for result in exc.results if isinstance(result, dict)] + except Exception: + return [] def _abandon_pushes(self, pushes: list[ProducerPushRecord]) -> None: """Release the consumer-side reservations of a batch that failed.""" for push in pushes: shards = self._known_reservations(push) if not shards: - shards = [{"addr": push.spec.consumer_zmq, "reservation_id": ""}] + # No confirmed topology means there is no safe address to + # cancel. The original push failure remains observable via + # ProducerPushManager.fail below. + continue try: self._retry_cancel_reservations(push.spec, shards, record=push) except _FanoutError: @@ -973,85 +710,12 @@ def _abandon_pushes(self, pushes: list[ProducerPushRecord]) -> None: push.spec.transfer_id, ) - def _record_push_perf( - self, - stage_ms: dict[str, float], - *, - stage_max_ms: dict[str, float], - item_count: int, - byte_count: int, - skipped_items: int, - failed: bool, - ) -> None: - now = time.monotonic() - report: tuple[_PushPerfWindow, int, int] | None = None - with self._push_perf_lock: - self._active_transfer_batches -= 1 - perf = self._push_perf - perf.batches += 1 - perf.items += item_count - perf.bytes += byte_count - perf.skipped_items += skipped_items - perf.failures += int(failed) - for stage, elapsed_ms in stage_ms.items(): - perf.stage_totals_ms[stage] = ( - perf.stage_totals_ms.get(stage, 0.0) + elapsed_ms - ) - perf.stage_max_ms[stage] = max( - perf.stage_max_ms.get(stage, 0.0), - stage_max_ms.get(stage, elapsed_ms), - ) - if ( - self._transfer_metrics_log_interval > 0 - and now - perf.started_at >= self._transfer_metrics_log_interval - ): - report = ( - perf, - self._active_transfer_batches, - self._queued_transfer_batches, - ) - self._push_perf = _PushPerfWindow(started_at=now) - if report is None: - return - perf, active_batches, queued_batches = report - batches = max(perf.batches, 1) - items = max(perf.items, 1) - stage_parts = [] - for stage in ("queue", *_PUSH_STAGES, "total"): - divisor = items if stage == "queue" else batches - average = perf.stage_totals_ms.get(stage, 0.0) / divisor - maximum = perf.stage_max_ms.get(stage, 0.0) - stage_parts.append(f"{stage}_ms={average:.1f}/{maximum:.1f}") - stage_summary = " ".join(stage_parts) - producer_metrics = dict(self._producer_metrics) - self._producer_metrics.clear() - logger.info( - "EC Mooncake push perf: batches=%d items=%d bytes=%d " - "batch_items=%.1f skipped=%d failures=%d active=%d queued=%d " - "producer=%s queue_item_avg/max and stage_batch_avg/max: %s", - perf.batches, - perf.items, - perf.bytes, - perf.items / batches, - perf.skipped_items, - perf.failures, - active_batches, - queued_batches, - producer_metrics, - stage_summary, - ) - def _flush_pending_pushes(self) -> None: self._producer_pushes.submit_batches( self._io_executor, self._push_batch, - self._note_push_batch_queued, ) - def _note_push_batch_queued(self) -> None: - with self._push_perf_lock: - self._queued_transfer_batches += 1 - def _bind_push_source(self, tensor: torch.Tensor, mm_hash: str) -> None: ready_event = None if tensor.device.type == "cuda": @@ -1060,26 +724,37 @@ def _bind_push_source(self, tensor: torch.Tensor, mm_hash: str) -> None: self._producer_pushes.bind_source(mm_hash, tensor, ready_event) def _cancel_orphaned_reservation(self, record: ProducerPushRecord) -> None: + resolution_error: Exception | None = None try: reservations = self._producer_pushes.resolve_reservations(record) - except Exception: + except Exception as exc: + resolution_error = exc reservations = self._known_reservations(record) - known = bool(reservations) reservations = [ - shard - for shard in reservations - if not shard.get("cached", False) and not shard.get("cancelled", False) + shard for shard in reservations if not shard.get("cancelled", False) ] - if not known: - reservations = [{"addr": record.spec.consumer_zmq, "reservation_id": ""}] - error = None + if not reservations: + if resolution_error is not None: + self._producer_pushes.fail([record], resolution_error) + else: + self._producer_pushes.finish_cancel(record) + return try: self._retry_cancel_reservations(record.spec, reservations, record=record) - except _FanoutError as exc: - error = exc - self._producer_pushes.finish_cancel(record) - if error is not None: - raise error + except Exception as cleanup_error: + if resolution_error is not None: + combined_error = RuntimeError( + f"EC reservation resolution failed ({resolution_error}); " + f"cleanup also failed ({cleanup_error})" + ) + combined_error.__cause__ = resolution_error + cleanup_error = combined_error + self._producer_pushes.fail([record], cleanup_error) + return + if resolution_error is not None: + self._producer_pushes.fail([record], resolution_error) + else: + self._producer_pushes.finish_cancel(record) def get_finished( self, finished_req_ids: set[str] @@ -1111,7 +786,6 @@ def build_connector_worker_meta(self) -> ECMooncakeWorkerMetadata | None: self._flush_pending_pushes() failures = self._producer_pushes.poll() - self._producer_metrics["saves_failed"] += len(failures) for mm_hash, error in failures: logger.error( "EC Mooncake async save failed for mm_hash=%s: %s", @@ -1123,32 +797,43 @@ def build_connector_worker_meta(self) -> ECMooncakeWorkerMetadata | None: loaded=self._completed_loads, failed_loads=self._failed_loads, reclaimed=reclaimed, - pending_loads=False, pending_saves=self._producer_pushes.pending, ) self._completed_loads = set() self._failed_loads = set() - if self.is_consumer: - self._maybe_log_consumer_worker_metrics() return meta def close(self) -> None: if self._shutdown: return self._shutdown = True + if self._control_server is not None: + self._reservations.begin_shutdown() self._flush_pending_pushes() - self._io_executor.shutdown(wait=True, cancel_futures=True) + for record in self._producer_pushes.cancel_requests(None): + self._producer_pushes.submit_cancel( + record, + self._io_executor, + self._cancel_orphaned_reservation, + ) + self._io_executor.shutdown(wait=True) + self._control_executor.shutdown(wait=True) if self._shard_pool is not None: - self._shard_pool.shutdown(wait=True, cancel_futures=True) - self._control_executor.shutdown(wait=True, cancel_futures=True) - # Every thread that could hold a control socket is stopped by now. + self._shard_pool.shutdown(wait=True) + # Every producer-side thread that could hold a control socket is stopped. self._control_client.close() + drained = True if self._control_server is not None: + drained = self._reservations.wait_for_writers( + self._shutdown_drain_timeout_s + ) self._control_server.close() - self._consumer_memory.close() + if drained: + self._consumer_memory.close() + else: + logger.error( + "Timed out waiting for Mooncake EC writers; keeping the consumer " + "receive pool registered" + ) self._producer_memory.close() self._transfer.close() - - def __del__(self) -> None: - with suppress(Exception): - self.close() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 8a271db69d46..8f375c39a6b3 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -5,7 +5,6 @@ from __future__ import annotations from collections.abc import Collection -from contextlib import suppress from typing import TYPE_CHECKING, Any from vllm.distributed.ec_transfer.ec_connector.base import ( @@ -16,9 +15,6 @@ ) from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( ECMooncakeConnectorMetadata, - ECMooncakeLoadSpec, - ECMooncakePushSpec, - ECMooncakeWorkerMetadata, ) from vllm.distributed.ec_transfer.ec_connector.mooncake.scheduler import ( ECMooncakeScheduler, @@ -33,23 +29,9 @@ from vllm.v1.outputs import ECConnectorOutput from vllm.v1.request import Request -__all__ = [ - "ECMooncakeConnector", - "ECMooncakeConnectorMetadata", - "ECMooncakeLoadSpec", - "ECMooncakePushSpec", - "ECMooncakeWorkerMetadata", -] - class ECMooncakeConnector(ECConnectorBase): - """Preserve the public API while delegating to one process-role component. - - Attributes: - _scheduler: Scheduler implementation when constructed for that role. - _worker: Worker implementation when constructed for that role. - _closed: Whether role-specific resources have already been released. - """ + """Preserve the public API while delegating to one process-role component.""" def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): super().__init__(vllm_config=vllm_config, role=role) @@ -59,9 +41,9 @@ def __init__(self, vllm_config: VllmConfig, role: ECConnectorRole): self._closed = False if role == ECConnectorRole.SCHEDULER: - self._scheduler = ECMooncakeScheduler.from_vllm_config(vllm_config) + self._scheduler = ECMooncakeScheduler(vllm_config) elif role == ECConnectorRole.WORKER: - self._worker = ECMooncakeWorker.from_vllm_config(vllm_config) + self._worker = ECMooncakeWorker(vllm_config) else: raise ValueError(f"Unknown EC connector role: {role}") @@ -111,7 +93,15 @@ def ensure_cache_available( self, request: Request, num_computed_tokens: int, - local_cache_hashes: Collection[str] | None = None, + ) -> bool: + assert self._scheduler is not None + return self._scheduler.ensure_cache_available(request, num_computed_tokens) + + def _ensure_cache_available( + self, + request: Request, + num_computed_tokens: int, + local_cache_hashes: Collection[str], ) -> bool: assert self._scheduler is not None return self._scheduler.ensure_cache_available( @@ -152,7 +142,3 @@ def shutdown(self) -> None: self._scheduler.close() if self._worker is not None: self._worker.close() - - def __del__(self) -> None: - with suppress(Exception): - self.shutdown() diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index b331aaae6033..1a26805b97ce 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -593,7 +593,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: if ( self.ec_connector is not None and request.mm_features - and not self.ec_connector.ensure_cache_available( + and not self.ec_connector._ensure_cache_available( request, request.num_computed_tokens - request.num_output_placeholders, self.encoder_cache_manager.cached.keys(), @@ -934,7 +934,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: if ( self.ec_connector is not None and request.mm_features - and not self.ec_connector.ensure_cache_available( + and not self.ec_connector._ensure_cache_available( request, num_computed_tokens, self.encoder_cache_manager.cached.keys(), From 1faf3eec8f3b556832e306dbbed57202a6b1c452 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Fri, 4 Sep 2026 03:33:21 +0000 Subject: [PATCH 23/30] [EPD] Keep local_cache_hashes on the public EC connector hook Signed-off-by: Tianyu Guo --- tests/v1/core/test_scheduler.py | 46 ++----------------- .../unit/test_ec_mooncake_connector.py | 8 ++-- .../ec_transfer/ec_connector/base.py | 12 ++--- .../ec_transfer/ec_connector/cpu/connector.py | 4 +- .../ec_connector/mooncake_ec_connector.py | 10 +--- vllm/v1/core/sched/scheduler.py | 4 +- 6 files changed, 17 insertions(+), 67 deletions(-) diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index 2370709f16e2..fb426080379a 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -5837,9 +5837,11 @@ def test_ec_connector_ensure_cache_available_defers_request(use_kv_connector): scheduler.add_request(request_behind) output = scheduler.schedule() - # The public connector API remains the legacy two-argument method. + # ensure_cache_available must have been called with (request, num_computed_tokens=0) + # for a brand-new request that has no cached tokens yet. ensure_call = scheduler.ec_connector.ensure_cache_available.call_args - assert ensure_call.args == (request_deferred, 0) + assert ensure_call.args[:2] == (request_deferred, 0) + assert not ensure_call.args[2] # Deferred request must NOT be scheduled assert request_deferred.request_id not in output.num_scheduled_tokens _assert_right_encoder_cache_allocated(scheduler, expected_total_allocated=0) @@ -5896,46 +5898,6 @@ def test_ec_connector_defers_running_request_for_async_reload(): assert ensure_call.args[:2] == (request, 32) -@pytest.mark.skip_global_cleanup -def test_ec_connector_legacy_ensure_cache_available_signature_is_supported(tmp_path): - """An out-of-tree connector with the original method remains callable.""" - - from vllm.distributed.ec_transfer.ec_connector.example_connector import ( - ECExampleConnector, - ) - - calls = [] - - class LegacyConnector(ECExampleConnector): - def ensure_cache_available(self, request, num_computed_tokens): - calls.append((request, num_computed_tokens)) - return False - - (tmp_path / "config.json").write_text( - '{"architectures": ["OPTForCausalLM"], "model_type": "opt"}' - ) - scheduler = create_scheduler( - model=str(tmp_path), - skip_tokenizer_init=True, - use_ec_connector=True, - ec_role="ec_consumer", - ) - scheduler.ec_connector = LegacyConnector( - scheduler.vllm_config, scheduler.ec_connector.role - ) - request = create_requests( - num_requests=1, - num_tokens=128, - mm_positions=[[PlaceholderRange(offset=48, length=32)]], - )[0] - - scheduler.add_request(request) - output = scheduler.schedule() - - assert request.request_id not in output.num_scheduled_tokens - assert calls == [(request, 0)] - - def test_ec_connector_pending_prefetch_only_checks_future_mm_features(): """Test that future mm feature filtering only yields features beyond the computed token frontier. diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index f0c5c1433b38..b7a79c8cdfc1 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -1259,14 +1259,14 @@ def test_local_cache_hit_keeps_the_transfer( patch.object(scheduler._scheduler, "_drain_push_notifications"), patch.object(scheduler._scheduler, "_queue_cancel") as cancel, ): - assert scheduler._ensure_cache_available(request, 0, {mm_hash}) + assert scheduler.ensure_cache_available(request, 0, {mm_hash}) cancel.assert_not_called() record = scheduler._scheduler._transfers.get("request-transfer") assert record is not None assert record.state is SchedulerTransferState.AVAILABLE # Once the entry is evicted the request can still get it. - assert not scheduler._ensure_cache_available(request, 0, set()) + assert not scheduler.ensure_cache_available(request, 0, set()) assert ( scheduler._scheduler._transfers.first_for_hash( mm_hash, (SchedulerTransferState.LOADING,) @@ -1456,7 +1456,7 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( ) scheduler._scheduler._drain_push_notifications() - assert not scheduler._ensure_cache_available(first, 0, set()) + assert not scheduler.ensure_cache_available(first, 0, set()) meta = scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[]) ) @@ -1481,7 +1481,7 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( # transfer is spent. It must still be served. with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert scheduler.has_cache_item(mm_hash) - assert not scheduler._ensure_cache_available(second, 0, set()) + assert not scheduler.ensure_cache_available(second, 0, set()) assert record.state is SchedulerTransferState.LOADING reload = scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[]) diff --git a/vllm/distributed/ec_transfer/ec_connector/base.py b/vllm/distributed/ec_transfer/ec_connector/base.py index 196c4f32bc1d..5a526def1487 100644 --- a/vllm/distributed/ec_transfer/ec_connector/base.py +++ b/vllm/distributed/ec_transfer/ec_connector/base.py @@ -268,6 +268,7 @@ def ensure_cache_available( self, request: "Request", num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, ) -> bool: """ Ensure encoder cache items are available for the given request. @@ -276,21 +277,14 @@ def ensure_cache_available( Args: request: the request whose multimodal features to check. num_computed_tokens: tokens already covered by cached KV blocks. + local_cache_hashes: encoder outputs already cached locally. + Returns: True if all items are ready or no transfer is needed. False if any items are still in transit (request should be deferred). """ return True - def _ensure_cache_available( - self, - request: "Request", - num_computed_tokens: int, - local_cache_hashes: Collection[str], - ) -> bool: - """Core-only adapter that preserves the connector extension API.""" - return self.ensure_cache_available(request, num_computed_tokens) - @abstractmethod def update_state_after_alloc(self, request: "Request", index: int): """ diff --git a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py index 0de266be8eff..aab35b2136a5 100644 --- a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py @@ -8,6 +8,7 @@ offloaded to CPU instead of recomputing them. """ +from collections.abc import Collection from typing import TYPE_CHECKING import torch @@ -94,10 +95,11 @@ def ensure_cache_available( self, request: "Request", num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, ) -> bool: assert self.connector_scheduler is not None return self.connector_scheduler.ensure_cache_available( - request, num_computed_tokens + request, num_computed_tokens, local_cache_hashes ) def update_state_after_alloc(self, request: "Request", index: int) -> None: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index 8f375c39a6b3..efa3a10fe5b8 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -93,15 +93,7 @@ def ensure_cache_available( self, request: Request, num_computed_tokens: int, - ) -> bool: - assert self._scheduler is not None - return self._scheduler.ensure_cache_available(request, num_computed_tokens) - - def _ensure_cache_available( - self, - request: Request, - num_computed_tokens: int, - local_cache_hashes: Collection[str], + local_cache_hashes: Collection[str] | None = None, ) -> bool: assert self._scheduler is not None return self._scheduler.ensure_cache_available( diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index 1a26805b97ce..b331aaae6033 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -593,7 +593,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: if ( self.ec_connector is not None and request.mm_features - and not self.ec_connector._ensure_cache_available( + and not self.ec_connector.ensure_cache_available( request, request.num_computed_tokens - request.num_output_placeholders, self.encoder_cache_manager.cached.keys(), @@ -934,7 +934,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: if ( self.ec_connector is not None and request.mm_features - and not self.ec_connector._ensure_cache_available( + and not self.ec_connector.ensure_cache_available( request, num_computed_tokens, self.encoder_cache_manager.cached.keys(), From f073b1dfc26af98c1659af657d2c1ea7279b722b Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Fri, 4 Sep 2026 04:19:13 +0000 Subject: [PATCH 24/30] [EPD] Mirror the local encoder cache in the Mooncake scheduler Signed-off-by: Tianyu Guo --- tests/v1/core/test_scheduler.py | 3 +- .../unit/test_ec_mooncake_connector.py | 12 ++-- .../ec_transfer/ec_connector/base.py | 57 +++++++------------ .../ec_transfer/ec_connector/cpu/connector.py | 8 +-- .../ec_connector/mooncake/scheduler.py | 17 +++--- .../ec_connector/mooncake_ec_connector.py | 10 +--- vllm/v1/core/sched/scheduler.py | 5 +- 7 files changed, 42 insertions(+), 70 deletions(-) diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index d3112b1f3890..7a6d9f78877c 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -5903,8 +5903,7 @@ def test_ec_connector_ensure_cache_available_defers_request(use_kv_connector): # ensure_cache_available must have been called with (request, num_computed_tokens=0) # for a brand-new request that has no cached tokens yet. ensure_call = scheduler.ec_connector.ensure_cache_available.call_args - assert ensure_call.args[:2] == (request_deferred, 0) - assert not ensure_call.args[2] + assert ensure_call.args == (request_deferred, 0) # Deferred request must NOT be scheduled assert request_deferred.request_id not in output.num_scheduled_tokens _assert_right_encoder_cache_allocated(scheduler, expected_total_allocated=0) diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 4a0da66e62fa..7b548879b08f 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -1259,14 +1259,18 @@ def test_local_cache_hit_keeps_the_transfer( patch.object(scheduler._scheduler, "_drain_push_notifications"), patch.object(scheduler._scheduler, "_queue_cancel") as cancel, ): - assert scheduler.ensure_cache_available(request, 0, {mm_hash}) + # The local encoder cache is mirrored from the Scheduler's + # own alloc/free notifications, not passed in per call. + scheduler.update_state_after_alloc(request, 0) + assert scheduler.ensure_cache_available(request, 0) cancel.assert_not_called() record = scheduler._scheduler._transfers.get("request-transfer") assert record is not None assert record.state is SchedulerTransferState.AVAILABLE # Once the entry is evicted the request can still get it. - assert not scheduler.ensure_cache_available(request, 0, set()) + scheduler._scheduler._local_cache.discard(mm_hash) + assert not scheduler.ensure_cache_available(request, 0) assert ( scheduler._scheduler._transfers.first_for_hash( mm_hash, (SchedulerTransferState.LOADING,) @@ -1456,7 +1460,7 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( ) scheduler._scheduler._drain_push_notifications() - assert not scheduler.ensure_cache_available(first, 0, set()) + assert not scheduler.ensure_cache_available(first, 0) meta = scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[]) ) @@ -1481,7 +1485,7 @@ def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( # transfer is spent. It must still be served. with patch.object(scheduler._scheduler, "_drain_push_notifications"): assert scheduler.has_cache_item(mm_hash) - assert not scheduler.ensure_cache_available(second, 0, set()) + assert not scheduler.ensure_cache_available(second, 0) assert record.state is SchedulerTransferState.LOADING reload = scheduler.build_connector_meta( SimpleNamespace(free_encoder_mm_hashes=[]) diff --git a/vllm/distributed/ec_transfer/ec_connector/base.py b/vllm/distributed/ec_transfer/ec_connector/base.py index 5a526def1487..ea44fb6f5beb 100644 --- a/vllm/distributed/ec_transfer/ec_connector/base.py +++ b/vllm/distributed/ec_transfer/ec_connector/base.py @@ -26,7 +26,6 @@ import enum from abc import ABC, abstractmethod -from collections.abc import Collection from typing import TYPE_CHECKING, Any import torch @@ -160,25 +159,6 @@ def register_caches( # TODO: Implement this later for P2P feature return - def start_save_caches(self, **kwargs: Any) -> None: - """Start work that can overlap encoder execution.""" - return None - - def start_worker_services(self) -> None: - """Start services that require the worker device to be initialized.""" - return None - - def take_unavailable_requests(self) -> set[str]: - """Requests whose encoder inputs can no longer be obtained. - - A connector that cannot always deliver an item reports the affected - requests here instead of deferring them forever. The scheduler fails - them with a retryable error, leaving the caller to decide whether to - re-issue the request. Called once per scheduling pass; the returned - ids are cleared. - """ - return set() - @abstractmethod def start_load_caches( self, encoder_cache: dict[str, torch.Tensor], **kwargs @@ -265,10 +245,7 @@ def has_cache_item( pass def ensure_cache_available( - self, - request: "Request", - num_computed_tokens: int, - local_cache_hashes: Collection[str] | None = None, + self, request: "Request", num_computed_tokens: int ) -> bool: """ Ensure encoder cache items are available for the given request. @@ -277,7 +254,6 @@ def ensure_cache_available( Args: request: the request whose multimodal features to check. num_computed_tokens: tokens already covered by cached KV blocks. - local_cache_hashes: encoder outputs already cached locally. Returns: True if all items are ready or no transfer is needed. @@ -285,27 +261,34 @@ def ensure_cache_available( """ return True - @abstractmethod - def update_state_after_alloc(self, request: "Request", index: int): - """ - Update ECConnector state to decide allocate cache for requests + def start_save_caches(self, **kwargs: Any) -> None: + """Prepare this step's outbound pushes before the model runs.""" + return - Args: - request (Request): the request object. + def start_worker_services(self) -> None: + """Start Worker-side services once the model is resident.""" + return + + def take_unavailable_requests(self) -> set[str]: + """Request IDs whose encoder inputs the connector can no longer obtain. + + The Scheduler fails these; re-issuing the request re-runs the encode. """ - pass + return set() def update_state_after_free(self, request: "Request", index: int): + """Notify the connector that an encoder cache entry was released.""" + return + + @abstractmethod + def update_state_after_alloc(self, request: "Request", index: int): """ - Called once the request has consumed the encoder input, well before it - finishes generating. Connectors that hold per-request transfer state - (buffers, reservations) should release this item's share here. + Update ECConnector state to decide allocate cache for requests Args: request (Request): the request object. - index (int): the multimodal item index within the request. """ - return + pass @abstractmethod def build_connector_meta( diff --git a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py index b296a7cafe2c..a0996d903515 100644 --- a/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/cpu/connector.py @@ -8,7 +8,6 @@ offloaded to CPU instead of recomputing them. """ -from collections.abc import Collection from typing import TYPE_CHECKING, Any import torch @@ -98,14 +97,11 @@ def has_cache_item(self, identifier: str) -> bool: return self.connector_scheduler.has_cache_item(identifier) def ensure_cache_available( - self, - request: "Request", - num_computed_tokens: int, - local_cache_hashes: Collection[str] | None = None, + self, request: "Request", num_computed_tokens: int ) -> bool: assert self.connector_scheduler is not None return self.connector_scheduler.ensure_cache_available( - request, num_computed_tokens, local_cache_hashes + request, num_computed_tokens ) def update_state_after_alloc(self, request: "Request", index: int) -> None: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py index 73420aa1aec5..9e52e55c5cab 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py @@ -12,7 +12,6 @@ import math import time from collections import OrderedDict -from collections.abc import Collection from concurrent.futures import Future, ThreadPoolExecutor from typing import TYPE_CHECKING, Any @@ -104,6 +103,10 @@ def __init__(self, vllm_config: VllmConfig) -> None: self._pushes_to_prepare: dict[str, ECMooncakePushSpec] = {} self._prepared_push_transfer_ids: set[str] = set() self._event_ready_shards: OrderedDict[str, set[int]] = OrderedDict() + # Mirror of the engine's encoder cache, maintained from the alloc and + # free notifications the Scheduler already sends. Tracking it here + # keeps `ensure_cache_available` on the upstream two-argument shape. + self._local_cache: set[str] = set() def _cancel_remote(self, consumer_zmq: str, transfer_id: str) -> bool: pending = None @@ -340,12 +343,7 @@ def _request_transfer_id(request: Any, index: int) -> str | None: return str(item["transfer_id"]) return None - def ensure_cache_available( - self, - request: Any, - num_computed_tokens: int, - local_cache_hashes: Collection[str] | None = None, - ) -> bool: + def ensure_cache_available(self, request: Any, num_computed_tokens: int) -> bool: if self._is_producer: for index, feature in enumerate(request.mm_features): if ( @@ -357,7 +355,6 @@ def ensure_cache_available( return True self._drain_push_notifications() - local_cache_hashes = local_cache_hashes or set() all_ready = True for index, feature in enumerate(request.mm_features): if ( @@ -371,7 +368,7 @@ def ensure_cache_available( self._transfers.touch_available( transfer_id, time.monotonic() + _RESERVATION_TTL_SECONDS ) - if mm_hash in local_cache_hashes: + if mm_hash in self._local_cache: continue if self._transfers.has_state(mm_hash, (SchedulerTransferState.READY,)): continue @@ -427,6 +424,7 @@ def _prepare_push_spec(self, request: Any, index: int) -> None: self._prepared_push_transfer_ids.add(transfer_id) def update_state_after_alloc(self, request: Any, index: int) -> None: + self._local_cache.add(request.mm_features[index].identifier) if self._is_producer: self._prepare_push_spec(request, index) @@ -447,6 +445,7 @@ def build_connector_meta( self, scheduler_output: SchedulerOutput ) -> ECConnectorMetadata: for mm_hash in scheduler_output.free_encoder_mm_hashes: + self._local_cache.discard(mm_hash) self._transfers.release_ready(mm_hash, time.monotonic()) for transfer_id in self._transfers.drain_orphaned(): self._queue_cancel(transfer_id) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py index efa3a10fe5b8..2949742220ec 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -4,7 +4,6 @@ from __future__ import annotations -from collections.abc import Collection from typing import TYPE_CHECKING, Any from vllm.distributed.ec_transfer.ec_connector.base import ( @@ -90,15 +89,10 @@ def has_cache_item(self, identifier: str) -> bool: return self._scheduler.has_cache_item(identifier) def ensure_cache_available( - self, - request: Request, - num_computed_tokens: int, - local_cache_hashes: Collection[str] | None = None, + self, request: Request, num_computed_tokens: int ) -> bool: assert self._scheduler is not None - return self._scheduler.ensure_cache_available( - request, num_computed_tokens, local_cache_hashes - ) + return self._scheduler.ensure_cache_available(request, num_computed_tokens) def update_state_after_alloc(self, request: Request, index: int) -> None: assert self._scheduler is not None diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index 8d85d277b5dd..fc1c44f926f8 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -596,7 +596,6 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: and not self.ec_connector.ensure_cache_available( request, request.num_computed_tokens - request.num_output_placeholders, - self.encoder_cache_manager.cached.keys(), ) ): req_index += 1 @@ -935,9 +934,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: self.ec_connector is not None and request.mm_features and not self.ec_connector.ensure_cache_available( - request, - num_computed_tokens, - self.encoder_cache_manager.cached.keys(), + request, num_computed_tokens ) ): request_queue.pop_request() From 8808bfbedc3e9b8204574a3389a2f44d7821a9a8 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Fri, 4 Sep 2026 07:26:56 +0000 Subject: [PATCH 25/30] [EPD] Address Mooncake EC connector review findings Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 11 ++-- .../unit/test_ec_mooncake_connector.py | 42 ++++++++++++++ .../ec_connector/unit/test_epd_proxy_retry.py | 57 +++++++++++++++++++ .../ec_connector/mooncake/control.py | 29 +++++++++- .../ec_connector/mooncake/memory.py | 6 +- .../ec_connector/mooncake/transfer.py | 21 +++---- .../ec_connector/mooncake/worker.py | 2 +- 7 files changed, 150 insertions(+), 18 deletions(-) create mode 100644 tests/v1/ec_connector/unit/test_epd_proxy_retry.py diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index 94051db80c8f..ac1742f8cc41 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -395,8 +395,11 @@ async def maybe_prefill( ) -> dict: """ - Do prefill-only task if p_url exist; - - Return modified request data with kv transfer params (for nixl connector) + - Return a new body carrying kv transfer params (for nixl connector) - Else, skip and return the original request data for decode + + `req_data` is never mutated: a decode retry re-enters this function with the + same body, and one attempt's `remote_block_ids` must not reach the next. """ if p_url: logger.info("[%s] Processing through prefill: %s", req_id, p_url) @@ -406,11 +409,9 @@ async def maybe_prefill( prefill_response_json = await prefill_response.json() kv_transfer_params = prefill_response_json.get("kv_transfer_params", {}) if kv_transfer_params: - req_data["kv_transfer_params"] = kv_transfer_params + return {**req_data, "kv_transfer_params": kv_transfer_params} - return req_data - else: - return req_data + return req_data async def process_prefill_stage( diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 7b548879b08f..61db3fef33a1 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -435,6 +435,48 @@ def cancel( assert completed == [("transfer", "r0"), ("transfer", "r0")] assert cancelled == [("transfer", "r0", True, False)] + def test_server_keeps_serving_after_an_undecodable_request(self): + """One bad frame must not take the shard's control channel down. + + The loop's `finally` closes both sockets, so an escaping exception ends + the thread silently and every later reserve against this shard surfaces + only as a control timeout. + """ + port = _find_free_port() + server = ConsumerControlServer( + "127.0.0.1", + port, + reserve=lambda request: {"nbytes": request["nbytes"], "ready": False}, + status=lambda transfer_id: None, + complete=lambda transfer_id, reservation_id: (True, True), + cancel=lambda transfer_id, reservation_id, abandon, refresh: True, + reap=lambda: 0, + peer_ports=[port], + ) + server.start() + addr = f"tcp://127.0.0.1:{port}" + context = zmq.Context() + raw = context.socket(zmq.REQ) + raw.setsockopt(zmq.RCVTIMEO, 2000) + raw.setsockopt(zmq.LINGER, 0) + raw.connect(addr) + client = ControlClient(2000) + try: + raw.send(b"{not json") + assert raw.recv_json()["ok"] is False + + # An unknown op raises inside the handler; both paths must leave + # the channel able to answer the next caller. + with pytest.raises(RuntimeError): + client.request(addr, {"op": "nonsense"}) + + assert client.request(addr, {"op": "peers"}) == {"ports": [port]} + finally: + raw.close(linger=0) + context.term() + client.close() + server.close() + @pytest.fixture def mock_vllm_config_producer(): diff --git a/tests/v1/ec_connector/unit/test_epd_proxy_retry.py b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py new file mode 100644 index 000000000000..86835da32a0d --- /dev/null +++ b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py @@ -0,0 +1,57 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Body handling across the EPD proxy's decode retries. + +Exercises the REAL helpers loaded from the ``examples/`` proxy, so a future +change to them is what these tests catch. +""" + +import asyncio +import importlib.util +from pathlib import Path + +import pytest + +PROXY_REL = "examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py" + + +@pytest.fixture(scope="module") +def proxy(): + path = Path(__file__).parents[4] / PROXY_REL + spec = importlib.util.spec_from_file_location("disagg_epd_proxy_retry", path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +class _Response: + def __init__(self, params): + self._params = params + + async def json(self): + return {"kv_transfer_params": self._params} + + +def test_maybe_prefill_leaves_the_caller_body_untouched(proxy, monkeypatch): + """A decode retry re-enters this function with the body it was given. + + Mutating that body in place let one attempt's `remote_block_ids` survive + into the next, so a retry whose prefill returns nothing sent decode blocks + the prefiller may already have freed. + """ + served = [{"remote_block_ids": [1, 2]}, {}] + + async def _stage(req_data, p_url, req_id): + assert "kv_transfer_params" not in req_data + return _Response(served.pop(0)) + + monkeypatch.setattr(proxy, "process_prefill_stage", _stage) + + body = {"messages": [], "stream": False} + first = asyncio.run(proxy.maybe_prefill(body, "http://prefill", "r1")) + assert first["kv_transfer_params"] == {"remote_block_ids": [1, 2]} + assert "kv_transfer_params" not in body + + second = asyncio.run(proxy.maybe_prefill(body, "http://prefill", "r1")) + assert "kv_transfer_params" not in second diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py index daeb7646a707..28d2a633e0df 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py @@ -228,7 +228,13 @@ def loop() -> None: def queue_event(event: dict[str, Any]) -> None: event["shard"] = self.port if len(pending_events) >= _MAX_PENDING_EVENTS: - pending_events.popleft() + dropped = pending_events.popleft() + logger.warning( + "EC Mooncake event backlog full on port %d; dropping " + "readiness for transfer_id=%s", + self.port, + dropped.get("transfer_id"), + ) pending_events.append(event) def queue_ready(transfer_id: str) -> None: @@ -267,6 +273,18 @@ def queue_ready(transfer_id: str) -> None: request = socket.recv_json() except zmq.Again: continue + except Exception: + # The frame arrived but did not decode. REP still owes a + # reply, so answer before returning to the loop. + logger.exception( + "EC Mooncake control channel on port %d received an " + "undecodable request", + self.port, + ) + socket.send_json( + {"ok": False, "error": "malformed control request"} + ) + continue try: op = request.get("op") result: Any = None @@ -320,6 +338,15 @@ def queue_ready(transfer_id: str) -> None: socket.send_json({"ok": True, "result": result}) except Exception as e: socket.send_json({"ok": False, "error": str(e)}) + except Exception: + # `finally` closes the sockets and the thread ends, so without + # this every later reserve against this shard would surface + # only as a control timeout with nothing to attribute it to. + logger.exception( + "EC Mooncake control channel on port %d stopped serving", + self.port, + ) + raise finally: socket.close(linger=0) event_socket.close(linger=0) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py index d7be4620dbc8..1dd812dd7666 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py @@ -361,10 +361,14 @@ def evict(mm_hash: str, allocation: MemoryAllocation) -> bool: return True while self._residents.evict_lru(evict) is not None: + # Eviction defers any free whose CUDA event is still pending, so + # those bytes reach the allocator only once the event is polled. + self._poll_frees_locked() region = self._allocator.allocate(nbytes) if region is not None: return region - return None + self._poll_frees_locked() + return self._allocator.allocate(nbytes) def _make_allocation( self, diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py index f5adcb772c7c..bd1651fb8fbf 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/transfer.py @@ -88,16 +88,17 @@ def unregister_memory(self, tensor: torch.Tensor) -> bool: engine = self._ensure_engine() address = tensor.data_ptr() ret = engine.unregister_memory(address) - if ret != 0: - logger.error( - "Mooncake EC memory unregistration failed for address %d: %d", - address, - ret, - ) - self._pending_unregister[address] = tensor - return False - self._pending_unregister.pop(address, None) - return True + with self._registration_lock: + if ret != 0: + logger.error( + "Mooncake EC memory unregistration failed for address %d: %d", + address, + ret, + ) + self._pending_unregister[address] = tensor + return False + self._pending_unregister.pop(address, None) + return True @staticmethod def _source_range(tensor: torch.Tensor) -> tuple[int, int]: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py index 70a079c5b21e..cb5d38fe2dd6 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py @@ -187,7 +187,7 @@ def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: shape = tuple(int(value) for value in payload["shape"]) dtype_name = str(payload["dtype"]) dtype = getattr(torch, dtype_name, None) - if dtype is None: + if not isinstance(dtype, torch.dtype): raise ValueError(f"Unsupported torch dtype string: {dtype_name!r}") expected_nbytes = math.prod(shape) * dtype.itemsize if expected_nbytes != nbytes: From bee2142a5498edf37fd3f0cc8da831fa6c78ad75 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Fri, 4 Sep 2026 14:09:28 +0000 Subject: [PATCH 26/30] [EPD] Harden the Mooncake EC connector against control-plane faults Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 38 ++- .../unit/test_ec_mooncake_connector.py | 276 +++++++++++++++++- .../ec_connector/unit/test_epd_proxy_retry.py | 88 ++++++ .../ec_connector/mooncake/config.py | 5 + .../ec_connector/mooncake/control.py | 65 ++++- .../ec_connector/mooncake/producer.py | 36 ++- .../ec_connector/mooncake/reservation.py | 43 ++- .../ec_connector/mooncake/scheduler.py | 52 ++-- .../ec_connector/mooncake/state.py | 61 +++- .../ec_connector/mooncake/worker.py | 6 +- 10 files changed, 611 insertions(+), 59 deletions(-) diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index ac1742f8cc41..409621235acf 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -228,7 +228,7 @@ async def fanout_encoder_primer( e_urls: list[str], req_id: str, consumer_zmq: str | None = None, -) -> dict[int, dict]: +) -> tuple[dict[int, dict], dict[str, dict]]: """ 1. Build one request *per MM item* with all text removed. 2. Send them concurrently to the encode cluster. @@ -238,13 +238,17 @@ async def fanout_encoder_primer( `ec_transfer_params`: its EC cache key and the grid its processor produced. The proxy still supplies the uuid so both sides key the cache the same way; the grid can only come from the encoder, which is the side that computed it. + + Also returns the connector handles to put on the decode body, as a fresh + mapping. `orig_request` is left untouched so a retry re-encodes from the + original request instead of carrying the previous attempt's handles. """ logger.info("[%s] Processing multimodal items...", req_id) mm_items = extract_mm_items(orig_request) if not mm_items: logger.info("[%s] No multimodal items, skipping encoder", req_id) - return {} # nothing to do + return {}, {} # nothing to do logger.info("[%s] got %d multimodal items...", req_id, len(mm_items)) @@ -252,6 +256,7 @@ async def fanout_encoder_primer( item_uuids: dict[int, str] = {} item_transfer_ids: dict[int, str] = {} item_meta: dict[int, dict] = {} + ec_params: dict[str, dict] = {} # Round-robin over encode servers to distribute load a bit. The cursor # persists across requests so fan-out doesn't restart at e_urls[0] every @@ -378,14 +383,12 @@ async def fanout_encoder_primer( # connector's own handle on the published embedding (for NIXL, # peer_host/peer_port/size_bytes). The decoder's connector # looks it up by mm_hash on the request, so carry it through. - orig_request.setdefault("ec_transfer_params", {})[item_uuids[idx]] = ( - reported - ) + ec_params[item_uuids[idx]] = reported logger.info( "[%s] All %d encoder requests completed successfully", req_id, len(mm_items) ) - return item_meta + return item_meta, ec_params async def maybe_prefill( @@ -522,7 +525,17 @@ async def log_requests(request: Request, call_next): async def on_startup() -> None: global encode_session, prefill_session, decode_session timeout = aiohttp.ClientTimeout(total=100_000) - connector = aiohttp.TCPConnector(limit=0, force_close=False) + # vLLM closes an idle keep-alive connection after + # VLLM_HTTP_TIMEOUT_KEEP_ALIVE seconds (5 by default), while aiohttp keeps + # pooling it for 15. Reusing one it has already closed fails the request + # with ServerDisconnectedError, and the server logs nothing at all: it + # closed the socket before the request arrived. Retire ours first. + server_keep_alive = float(os.getenv("VLLM_HTTP_TIMEOUT_KEEP_ALIVE", "5")) + connector = aiohttp.TCPConnector( + limit=0, + force_close=False, + keepalive_timeout=max(1.0, server_keep_alive - 1.0), + ) encode_session = aiohttp.ClientSession(timeout=timeout, connector=connector) if app.state.p_urls: # only setup if prefill instance(s) exist @@ -559,9 +572,18 @@ async def prepare_for_decode( rather than from a body whose images are already metadata references. """ _t0 = time.perf_counter() - item_meta = await fanout_encoder_primer(req_data, e_urls, req_id, consumer_zmq) + item_meta, ec_params = await fanout_encoder_primer( + req_data, e_urls, req_id, consumer_zmq + ) _t1 = time.perf_counter() prepared = req_data if NO_REWRITE else rewrite_for_decode(req_data, item_meta) + if ec_params: + # A fresh body every time: `rewrite_for_decode` hands back `req_data` + # itself when it rewrote nothing, and this attempt's handles must not + # outlive it into a retry. + handles = dict(prepared.get("ec_transfer_params") or {}) + handles.update(ec_params) + prepared = {**prepared, "ec_transfer_params": handles} _t2 = time.perf_counter() prepared = await maybe_prefill(prepared, p_url, req_id) return prepared, _t1 - _t0, _t2 - _t1 diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index 61db3fef33a1..c7179bdac833 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -19,7 +19,7 @@ from collections import Counter from concurrent.futures import Future, ThreadPoolExecutor from contextlib import contextmanager -from dataclasses import FrozenInstanceError +from dataclasses import FrozenInstanceError, replace from multiprocessing.reduction import ForkingPickler from types import SimpleNamespace from typing import Any @@ -177,6 +177,22 @@ def _wait_for_worker_io( raise TimeoutError("EC Mooncake worker I/O did not finish") +def _drain_until_subscribed(scheduler: Any, timeout: float = 5.0) -> None: + """Drain until the event-channel discovery lands. + + Subscribing runs on the control executor so an unreachable consumer costs + the Scheduler nothing, which means the first drain only starts it. + """ + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + scheduler._drain_pending = True + scheduler._drain_push_notifications() + if scheduler._event_inbox._socket is not None: + return + time.sleep(0.005) + raise TimeoutError("EC Mooncake event channel was never subscribed") + + def _bind_extra_config(config: VllmConfig) -> None: config.ec_transfer_config.get_from_extra_config.side_effect = lambda key, default: ( config.ec_transfer_config.ec_connector_extra_config.get(key, default) @@ -477,6 +493,89 @@ def test_server_keeps_serving_after_an_undecodable_request(self): client.close() server.close() + def test_a_dead_writer_gives_its_destination_back(self): + """Reserve, then vanish: the receive pool has to recover on its own. + + Driven over the wire the way a Producer and the Scheduler drive it -- + `reserve`, then the Scheduler's own cancel, which does not abandon the + writer. That cancel used to leave the destination pinned with no path + back, so a Producer crash permanently cost the Consumer a buffer. + """ + engine = MagicMock(spec=MooncakeTransfer) + engine.register_memory.return_value = 0 + engine.unregister_memory.return_value = True + pool = ConsumerMemoryPool(768, engine) + pool.prepare(torch.device("cpu")) + manager = ConsumerReservationManager(pool, 300.0, 16) + + def reserve(request: dict[str, Any]) -> dict[str, Any]: + """The Worker's own handler, minus the dtype plumbing.""" + manager.expire() + reservation, _ = manager.reserve( + str(request["transfer_id"]), + str(request["mm_hash"]), + int(request["nbytes"]), + tuple(int(value) for value in request["shape"]), + str(request["dtype"]), + torch.float32, + ) + if reservation is None: + raise RuntimeError("EC consumer buffer pool is full") + return {"ready": reservation.state is ConsumerReservationState.READY} + + port = _find_free_port() + server = ConsumerControlServer( + "127.0.0.1", + port, + reserve=reserve, + status=manager.status, + complete=manager.complete, + cancel=manager.cancel, + reap=manager.expire, + peer_ports=[port], + ) + server.start() + addr = f"tcp://127.0.0.1:{port}" + client = ControlClient(2000) + + def request_destination(transfer_id: str) -> None: + client.request( + addr, + { + "op": "reserve", + "transfer_id": transfer_id, + "mm_hash": f"hash-{transfer_id}", + "nbytes": 64, + "shape": [16], + "dtype": "float32", + }, + ) + + try: + abandoned = 0 + for index in range(64): + try: + request_destination(f"t{index}") + except RuntimeError as error: + assert "pool is full" in str(error) + break + # The writer dies here, so no completion ever arrives; the + # Scheduler times the transfer out and cancels it. + assert client.request( + addr, control.make_cancel_request(f"t{index}") + ) == {"cancelled": True} + abandoned += 1 + assert abandoned, "the pool should accept a writer before filling up" + + # No in-flight write can still land, so the buffers come back and + # an unrelated request gets a destination again. + for record in manager._records.values(): + record.expires_at = 0 + request_destination("after-the-grace") + finally: + client.close() + server.close() + @pytest.fixture def mock_vllm_config_producer(): @@ -1126,8 +1225,13 @@ def test_same_hash_index_preserves_order_and_identity(self): assert ( table.first_for_hash("hash", (SchedulerTransferState.AVAILABLE,)) is second ) - with pytest.raises(ValueError): - table.observe_ready(self.pushed_spec("first", "other-hash"), 10) + # A colliding transfer id is refused rather than raised: transfer ids + # arrive on the request, so the engine must not fail on one. + collided, accepted = table.observe_ready( + self.pushed_spec("first", "other-hash"), 10 + ) + assert collided is first and not accepted + assert first.mm_hash == "hash" def test_unavailable_notification_is_drained_once_and_rejects_late_ready(self): table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) @@ -1186,10 +1290,119 @@ def test_same_hash_keeps_only_latest_resident(self): assert table.drain_orphaned() == ["first"] assert table.drain_orphaned() == [] + def test_a_load_that_never_reports_back_fails_its_request(self): + """A dispatched load needs a deadline of its own. + + LOADING used to carry none, so a lost Worker report left the hash + loading for good and deferred every later request for it, silently + and without even a retriable failure. + """ + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + record, _ = table.observe_ready(self.pushed_spec("transfer"), 10) + + assert table.begin_load("hash", "transfer", "request", 20) is record + assert table.take_loads_to_dispatch() == [record] + + assert table.expire(15, 16) == [] + assert table.expire(21, 16) == [record] + assert record.state is SchedulerTransferState.UNAVAILABLE + assert table.take_unavailable_requests() == {"request"} + + def test_a_colliding_transfer_id_fails_only_the_request_that_named_it(self): + """Transfer ids arrive on the request, so a collision is reachable. + + Raising took the engine down with it; the transfer already under that + id has to survive and only the newcomer may fail. + """ + table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) + first = table.wait_for_event("shared", "req-a", "hash", 1) + + assert first is not None + assert table.wait_for_event("shared", "req-b", "other-hash", 1) is None + assert table.take_unavailable_requests() == {"req-b"} + assert first.mm_hash == "hash" + assert first.request_id == "req-a" + class TestECMooncakeSchedulerMetadata: """Validate Scheduler decisions and per-step Worker metadata.""" + def test_an_unreachable_consumer_does_not_block_the_scheduler( + self, mock_vllm_config_consumer + ): + """Subscribing must never cost the Scheduler a control timeout. + + `has_cache_item` runs inside `schedule()`, and discovery is one + blocking request per shard. Doing it inline froze every request in the + engine for `control_timeout_s` on each drain while a shard was + unreachable, and `discover_shards` caches only successes, so the cost + repeated for as long as the shard stayed down. + """ + timeout_s = 0.4 + mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { + "mooncake_protocol": "tcp", + "control_timeout_s": timeout_s, + } + mock_vllm_config_consumer.ec_transfer_config.ec_port = _find_free_port() + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + elapsed = [] + for _ in range(3): + started = time.monotonic() + assert scheduler.has_cache_item("hash") is False + elapsed.append(time.monotonic() - started) + finally: + scheduler.shutdown() + + assert max(elapsed) < timeout_s, elapsed + + def test_a_duplicate_transfer_id_fails_the_request_not_the_engine( + self, mock_vllm_config_consumer, mock_request_with_3_mm + ): + """`ec_transfer_params` is a request field, so ids can collide. + + `ensure_cache_available` runs inside `schedule()`: raising there took + EngineCore down on input any client could send. + """ + with patch_ec_mooncake_deps(): + scheduler = ECMooncakeConnector( + mock_vllm_config_consumer, ECConnectorRole.SCHEDULER + ) + try: + first = mock_request_with_3_mm + first.mm_features = first.mm_features[:1] + first.request_id = "req-a" + first.ec_transfer_params = { + "ec_items": [ + { + "mm_hash": first.mm_features[0].identifier, + "transfer_id": "shared", + } + ] + } + second = copy.copy(first) + second.request_id = "req-b" + second.mm_features = [ + replace(first.mm_features[0], identifier="another_hash") + ] + second.ec_transfer_params = { + "ec_items": [{"mm_hash": "another_hash", "transfer_id": "shared"}] + } + + with patch.object(scheduler._scheduler, "_drain_push_notifications"): + assert not scheduler.ensure_cache_available(first, 0) + assert not scheduler.ensure_cache_available(second, 0) + + assert scheduler.take_unavailable_requests() == {"req-b"} + record = scheduler._scheduler._transfers.get("shared") + assert record is not None + assert record.request_id == "req-a" + finally: + scheduler.shutdown() + def test_cancel_confirms_topology_and_retries_only_failed_shards(self): scheduler = object.__new__(ECMooncakeScheduler) scheduler._control_client = Mock(spec=ControlClient) @@ -1430,7 +1643,7 @@ def fake_send(addr: str, request: dict): "request", side_effect=fake_send, ) as send_control: - scheduler._scheduler._drain_push_notifications() + _drain_until_subscribed(scheduler._scheduler) subscribed = [ call.args[0] @@ -1858,6 +2071,45 @@ def test_cancel_tombstones_are_bounded(self): assert list(manager._tombstones) == ["c", "a", "d"] assert set(manager._records) == {"c", "a", "d"} + @pytest.mark.parametrize( + "deferred,terminal", + [ + ( + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.CANCELLED, + ), + ( + ConsumerReservationState.EXPIRE_PENDING, + ConsumerReservationState.EXPIRED, + ), + ], + ) + def test_a_deferred_release_is_reclaimed_when_no_writer_reports( + self, deferred, terminal + ): + """A writer that never reports must not pin its destination for good. + + Cancellation and expiry defer to the remote writer, so a Producer that + died between reserve and complete used to hold the buffer for the life + of the process and every later reserve saw a full pool. + """ + manager, pool, allocation = self.manager() + record, _ = self.reserve(manager) + + if deferred is ConsumerReservationState.CANCEL_PENDING: + assert manager.cancel("transfer", "") + else: + record.expires_at = 0 + manager.expire() + assert record.state is deferred + # Mooncake may still be writing, so the release waits first. + pool.free.assert_not_called() + + record.expires_at = 0 + assert manager.expire() == 1 + assert record.state is terminal + pool.free.assert_called_once_with(allocation) + class TestECMooncakeWorkerTransfer: """Validate end-to-end Worker reservation, push, load, and cleanup flows.""" @@ -1943,10 +2195,17 @@ def test_producer_push_state_owns_source_until_every_future_is_terminal(self): assert created assert duplicate is record assert not duplicate_created + # Another request may legitimately name the same encoding. + reasked = copy.copy(spec) + reasked.request_id = "another-request" + assert manager.reserve(reasked, lambda: Future()) == (record, False) + + # A different payload under the same id drops the newcomer instead of + # failing the engine; the push in flight keeps the id. changed = copy.copy(spec) changed.mm_hash = "other" - with pytest.raises(ValueError, match="changed identity"): - manager.reserve(changed, lambda: Future()) + assert manager.reserve(changed, lambda: Future()) == (record, False) + assert record.spec.mm_hash == "hash" source = torch.empty(16) manager.bind_source("hash", source, None) @@ -2751,8 +3010,9 @@ def test_push_reserves_before_encoder_output_is_saved( ) as send_control: assert not scheduler.has_cache_item("hash") assert not scheduler.has_cache_item("hash") - # The channel is built once, not per call: the roster is - # fetched and every shard subscribed to on the first one. + _drain_until_subscribed(scheduler._scheduler) + # The channel is built once, not per drain: the roster is + # fetched and every shard subscribed to exactly once. assert [call.args[1] for call in send_control.call_args_list] == [ {"op": "peers"}, {"op": "event_port"}, diff --git a/tests/v1/ec_connector/unit/test_epd_proxy_retry.py b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py index 86835da32a0d..3dc073888e77 100644 --- a/tests/v1/ec_connector/unit/test_epd_proxy_retry.py +++ b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py @@ -55,3 +55,91 @@ async def _stage(req_data, p_url, req_id): second = asyncio.run(proxy.maybe_prefill(body, "http://prefill", "r1")) assert "kv_transfer_params" not in second + + +class _EncoderResponse: + def __init__(self, params): + self.status = 200 + self._params = params + + async def json(self): + return {"ec_transfer_params": self._params} + + async def text(self): + return "" + + +class _EncoderSession: + """Serve one canned encoder reply per attempt.""" + + def __init__(self, replies): + self._replies = list(replies) + + async def post(self, url, json=None, headers=None): + return _EncoderResponse(self._replies.pop(0)) + + +def test_a_decode_retry_does_not_inherit_the_previous_handles(proxy, monkeypatch): + """The retry loop re-enters `prepare_for_decode` with the same body. + + Recording the encoder's connector handles on that body in place let + attempt 1's handle survive into attempt 2, so a second encode that + reported nothing still sent decode a handle on an embedding the encoder + no longer publishes -- the exact state the retry exists to leave behind. + """ + handle = {"metadata": {"image_grid_thw": [1, 2, 2]}, "peer_port": 1234} + monkeypatch.setattr( + proxy, + "encode_session", + _EncoderSession([{"encoder-side-hash": handle}, {}]), + ) + + async def _no_prefill(req_data, p_url, req_id): + return req_data + + monkeypatch.setattr(proxy, "maybe_prefill", _no_prefill) + + body = { + "messages": [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "image"}}], + } + ], + "stream": False, + } + args = ("r1", ["http://encoder"], "http://prefill", None) + + first, _, _ = asyncio.run(proxy.prepare_for_decode(body, *args)) + reported = first["ec_transfer_params"] + assert [handle] == [value for key, value in reported.items() if key != "ec_items"] + assert "ec_transfer_params" not in body + + # Attempt 2's encode reports nothing: decode must be told nothing. + second, _, _ = asyncio.run(proxy.prepare_for_decode(body, *args)) + assert "ec_transfer_params" not in second + + +@pytest.mark.parametrize("server_keep_alive", ["5", "2", "30"]) +def test_pooled_connections_are_retired_before_the_server_closes_them( + proxy, monkeypatch, server_keep_alive +): + """The proxy must not hand a request a connection the server has dropped. + + vLLM closes idle keep-alive connections at `VLLM_HTTP_TIMEOUT_KEEP_ALIVE` + seconds while aiohttp pools them for 15, so a slow hop leaves a dead + connection in the pool. The next request fails with + ServerDisconnectedError and the server logs nothing, because it closed the + socket before the request arrived. + """ + monkeypatch.setenv("VLLM_HTTP_TIMEOUT_KEEP_ALIVE", server_keep_alive) + monkeypatch.setattr(proxy.app.state, "p_urls", [], raising=False) + + asyncio.run(proxy.on_startup()) + try: + # No public accessor for the pool's idle timeout. + pooled_for = proxy.encode_session.connector._keepalive_timeout + assert pooled_for < float(server_keep_alive) + assert proxy.decode_session.connector._keepalive_timeout == pooled_for + finally: + asyncio.run(proxy.on_shutdown()) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py index e0bf88bb685e..10c9acae4676 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/config.py @@ -54,12 +54,16 @@ class MooncakeECConfig: ``control_port`` is the first TP-shard port after the DP offset; ``control_addr`` targets that shard, which advertises the full topology. + ``control_host`` is the address this instance both advertises and binds, + so reaching a Consumer means being routed to it rather than finding it on + every interface. """ is_producer: bool is_consumer: bool protocol: str buffer_device: str + control_host: str control_port: int control_addr: str control_timeout_ms: int @@ -102,6 +106,7 @@ def from_vllm_config(cls, vllm_config: VllmConfig) -> MooncakeECConfig: is_consumer=ec_config.is_ec_consumer, protocol=str(get("mooncake_protocol", "rdma")), buffer_device=str(ec_config.ec_buffer_device or "cuda").lower(), + control_host=str(ec_config.ec_ip), control_port=control_port, control_addr=make_zmq_path("tcp", ec_config.ec_ip, control_port), control_timeout_ms=max( diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py index 28d2a633e0df..67b870baa26f 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py @@ -13,6 +13,7 @@ import time from collections import deque from collections.abc import Callable +from concurrent.futures import Executor, Future from typing import Any import torch @@ -124,21 +125,30 @@ def make_cancel_request( class EventInbox: - """Receive Consumer readiness events without blocking the Scheduler.""" + """Receive Consumer readiness events without blocking the Scheduler. - def __init__(self, client: ControlClient) -> None: + Subscribing needs one blocking control request per shard, so it runs on + `executor` when one is given: an unreachable Consumer must cost the + Scheduler nothing, not one control timeout per drain. + """ + + def __init__(self, client: ControlClient, executor: Executor | None = None) -> None: self._client = client + self._executor = executor + self._discovery: Future[list[str] | None] | None = None self._context: zmq.Context | None = None self._socket: zmq.Socket | None = None self._closed = False self.shard_count = 1 - def _connect(self, base_addr: str) -> None: - if self._socket is not None: - return + def _endpoints(self, base_addr: str) -> list[str] | None: + """Ask every shard where it publishes readiness events. + + Blocking: called on `self._executor` unless there is none. + """ shards = self._client.discover_shards(base_addr) if shards is None: - return + return None endpoints = [] for addr in shards: try: @@ -151,14 +161,44 @@ def _connect(self, base_addr: str) -> None: "consumer shard %s; retrying the complete topology later.", addr, ) - return + return None + return endpoints + + def _connect(self, base_addr: str) -> None: + if self._socket is not None or self._closed: + return + if self._executor is None: + self._install(self._endpoints(base_addr)) + return + if self._discovery is None: + self._discovery = self._executor.submit(self._endpoints, base_addr) + return + if not self._discovery.done(): + return + discovery, self._discovery = self._discovery, None + try: + self._install(discovery.result()) + except Exception: + logger.warning( + "EC Mooncake event-channel discovery for %s failed; retrying.", + base_addr, + exc_info=True, + ) + + def _install(self, endpoints: list[str] | None) -> None: + """Adopt a discovered topology. + + Runs on the caller's thread so the socket is only ever touched there. + """ + if not endpoints or self._closed: + return context = zmq.Context() socket = context.socket(zmq.PULL) for endpoint in endpoints: socket.connect(endpoint) self._context = context self._socket = socket - self.shard_count = len(shards) + self.shard_count = len(endpoints) def drain(self, base_addr: str) -> list[dict[str, Any]]: self._connect(base_addr) @@ -170,11 +210,20 @@ def drain(self, base_addr: str) -> list[dict[str, Any]]: events.append(self._socket.recv_json(flags=zmq.DONTWAIT)) except zmq.Again: return events + except Exception: + # The event channel is a plain PULL socket: an undecodable + # frame must cost one frame, not the engine. + logger.warning( + "Discarding an undecodable EC Mooncake event.", exc_info=True + ) def close(self) -> None: if self._closed: return self._closed = True + if self._discovery is not None: + self._discovery.cancel() + self._discovery = None if self._socket is not None: self._socket.close(linger=0) if self._context is not None: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py index 947395b9997d..7e99779cd015 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py @@ -22,6 +22,26 @@ from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( ECMooncakePushSpec, ) +from vllm.logger import init_logger + +logger = init_logger(__name__) + + +def _same_destination( + existing: ECMooncakePushSpec, incoming: ECMooncakePushSpec +) -> bool: + """Whether two specs describe one push, ignoring which request asked. + + The same encoding requested twice shares a transfer id legitimately; only + a different payload or destination is a collision. + """ + return ( + existing.mm_hash == incoming.mm_hash + and existing.nbytes == incoming.nbytes + and existing.shape == incoming.shape + and existing.dtype == incoming.dtype + and existing.consumer_zmq == incoming.consumer_zmq + ) class ProducerPushState(Enum): @@ -107,9 +127,19 @@ def reserve( with self._lock: existing = self._records.get(spec.transfer_id) if existing is not None: - if existing.spec != spec: - raise ValueError( - f"Producer transfer {spec.transfer_id!r} changed identity" + if not _same_destination(existing.spec, spec): + # Transfer ids come in on the request, so two requests can + # name one id. Keep the push already in flight and drop + # the newcomer: the consumer times its own wait out and + # fails that request, which an engine-level raise would + # not. + logger.warning( + "EC Mooncake producer transfer_id=%s already pushes " + "mm_hash=%s to %s; dropping mm_hash=%s", + spec.transfer_id, + existing.spec.mm_hash[:16], + existing.spec.consumer_zmq, + spec.mm_hash[:16], ) return existing, False record = ProducerPushRecord( diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py index e0be41260953..0c6e05358d35 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py @@ -22,6 +22,9 @@ MemoryAllocation, ResidentLease, ) +from vllm.logger import init_logger + +logger = init_logger(__name__) class ConsumerReservationState(Enum): @@ -103,6 +106,9 @@ def __init__( ) -> None: self._memory = memory self._lease_ttl = lease_ttl + # A writer that has not reported for a full extra lease is gone; no + # RDMA write outlives that, so reclaiming its destination is safe. + self._writer_grace = lease_ttl self._tombstone_limit = tombstone_limit self._records: dict[str, ConsumerReservation] = {} self._active_ids: dict[str, None] = {} @@ -237,7 +243,7 @@ def begin_shutdown(self) -> None: if record.state is ConsumerReservationState.READY: self._terminate(record, ConsumerReservationState.CANCELLED) elif record.state is ConsumerReservationState.WRITING: - self._transition(record, ConsumerReservationState.CANCEL_PENDING) + self._defer(record, ConsumerReservationState.CANCEL_PENDING) self._condition.notify_all() def wait_for_writers(self, timeout: float) -> bool: @@ -297,7 +303,7 @@ def cancel( if record.state in _DEFERRED_STATES and not abandon: return True if record.state is ConsumerReservationState.WRITING and not abandon: - self._transition(record, ConsumerReservationState.CANCEL_PENDING) + self._defer(record, ConsumerReservationState.CANCEL_PENDING) return True if record.state not in _ACTIVE_STATES: return False @@ -346,10 +352,41 @@ def _expire_locked(self, now: float) -> int: self._terminate(record, ConsumerReservationState.EXPIRED) expired += 1 elif record.state is ConsumerReservationState.WRITING: - self._transition(record, ConsumerReservationState.EXPIRE_PENDING) + self._defer(record, ConsumerReservationState.EXPIRE_PENDING) + elif record.state in _DEFERRED_STATES: + # The writer that this release deferred to never reported + # back: it crashed, or its host is gone. Waiting forever + # pins the destination for the life of the process, so take + # the buffer back once no in-flight write could still land. + logger.warning( + "Reclaiming EC destination for transfer_id=%s after %.0fs " + "in %s: its writer never reported back", + record.transfer_id, + self._writer_grace, + record.state.name, + ) + self._terminate( + record, + ConsumerReservationState.CANCELLED + if record.state is ConsumerReservationState.CANCEL_PENDING + else ConsumerReservationState.EXPIRED, + ) + expired += 1 self._reap_tombstones(now) return expired + def _defer( + self, record: ConsumerReservation, state: ConsumerReservationState + ) -> None: + """Hand a release to the remote writer, but not indefinitely. + + Mooncake may still be writing into the destination, so the release + waits for the writer to report. `_expire_locked` reclaims the buffer + once the grace has passed without a report. + """ + self._transition(record, state) + record.expires_at = time.monotonic() + self._writer_grace + def _terminate( self, record: ConsumerReservation, state: ConsumerReservationState ) -> None: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py index 9e52e55c5cab..75e628711b1a 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py @@ -85,11 +85,11 @@ def __init__(self, vllm_config: VllmConfig) -> None: ) self._model_config = vllm_config.model_config self._control_client = ControlClient(config.control_timeout_ms) - self._event_inbox = EventInbox(self._control_client) self._control_executor = ThreadPoolExecutor( max_workers=_CONTROL_WORKERS, thread_name_prefix="ec-mooncake-control", ) + self._event_inbox = EventInbox(self._control_client, self._control_executor) self._metadata_resolver = PlaceholderMetadataResolver(vllm_config.model_config) self._unresolved_transfer_ids = 0 @@ -147,14 +147,17 @@ def _note_awaiting_push( mm_hash: str, transfer_id: str, request_id: str, + now: float, ) -> None: - now = time.monotonic() record = self._transfers.wait_for_event( transfer_id, request_id, mm_hash, now + self._push_wait_timeout, ) + if record is None: + # The id names another encoding; the request is already failed. + return if record.state is not SchedulerTransferState.WAITING_EVENT: return assert record.deadline is not None @@ -264,19 +267,30 @@ def _drain_push_notifications(self) -> None: self._expire_transfers() events = self._event_inbox.drain(self._control_addr) for data in events: - if data.get("ready"): - transfer_id = str(data["transfer_id"]) - record = self._transfers.get(transfer_id) - if record is not None and record.state in { - SchedulerTransferState.CANCELLED, - SchedulerTransferState.UNAVAILABLE, - SchedulerTransferState.EXPIRED, - SchedulerTransferState.FAILED, - }: - continue - if not self._note_shard_ready(data): - continue - self._store_pushed_spec(data) + if not data.get("ready"): + continue + try: + self._accept_ready_event(data) + except (KeyError, TypeError, ValueError): + # Readiness events cross a plain PULL socket, so a malformed + # one costs that event and not the engine. + logger.warning( + "Discarding a malformed EC readiness event.", exc_info=True + ) + + def _accept_ready_event(self, data: dict[str, Any]) -> None: + transfer_id = str(data["transfer_id"]) + record = self._transfers.get(transfer_id) + if record is not None and record.state in { + SchedulerTransferState.CANCELLED, + SchedulerTransferState.UNAVAILABLE, + SchedulerTransferState.EXPIRED, + SchedulerTransferState.FAILED, + }: + return + if not self._note_shard_ready(data): + return + self._store_pushed_spec(data) def has_cache_item(self, identifier: str) -> bool: if not self._is_consumer: @@ -355,6 +369,9 @@ def ensure_cache_available(self, request: Any, num_computed_tokens: int) -> bool return True self._drain_push_notifications() + # One timestamp for the whole decision, so the deadlines it hands out + # cannot disagree between features. + now = time.monotonic() all_ready = True for index, feature in enumerate(request.mm_features): if ( @@ -366,7 +383,7 @@ def ensure_cache_available(self, request: Any, num_computed_tokens: int) -> bool transfer_id = self._request_transfer_id(request, index) if transfer_id is not None: self._transfers.touch_available( - transfer_id, time.monotonic() + _RESERVATION_TTL_SECONDS + transfer_id, now + _RESERVATION_TTL_SECONDS ) if mm_hash in self._local_cache: continue @@ -379,13 +396,14 @@ def ensure_cache_available(self, request: Any, num_computed_tokens: int) -> bool mm_hash, transfer_id, request.request_id, + now + self._push_wait_timeout, ) if record is not None: self._scheduler_pending_work = True all_ready = False else: waiting_id = transfer_id or f"{request.request_id}:{index}" - self._note_awaiting_push(mm_hash, waiting_id, request.request_id) + self._note_awaiting_push(mm_hash, waiting_id, request.request_id, now) all_ready = False return all_ready diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py index a8f48867ef1c..1234cbccb339 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py @@ -12,6 +12,9 @@ from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( ECMooncakeLoadSpec, ) +from vllm.logger import init_logger + +logger = init_logger(__name__) class SchedulerTransferState(Enum): @@ -129,7 +132,13 @@ def wait_for_event( request_id: str, mm_hash: str, deadline: float, - ) -> SchedulerTransfer: + ) -> SchedulerTransfer | None: + """Track a request waiting for a push. + + Returns None when `transfer_id` already names a different encoding. + The id comes from the request, so a collision has to fail that one + request; re-issuing it re-runs the encode under a fresh id. + """ record = self._records.get(transfer_id) if record is None: record = SchedulerTransfer( @@ -141,8 +150,10 @@ def wait_for_event( deadline=deadline, ) self._insert(record) + elif not self._identity_matches(record, mm_hash): + self._refuse_colliding_id(record, mm_hash, request_id) + return None else: - self._check_identity(record, mm_hash) if not record.request_id: record.request_id = request_id if record.state in _TERMINAL_STATES and request_id: @@ -164,8 +175,9 @@ def observe_ready( deadline=None, ) self._insert(record) - else: - self._check_identity(record, spec.mm_hash) + elif not self._identity_matches(record, spec.mm_hash): + self._refuse_colliding_id(record, spec.mm_hash, record.request_id) + return record, False if record.state is not SchedulerTransferState.WAITING_EVENT: return record, False record.spec = spec @@ -183,6 +195,7 @@ def begin_load( mm_hash: str, transfer_id: str | None = None, request_id: str = "", + deadline: float | None = None, ) -> SchedulerTransfer | None: record = self._records.get(transfer_id) if transfer_id else None if record is not None and ( @@ -202,7 +215,9 @@ def begin_load( return None if not record.request_id: record.request_id = request_id - record.deadline = None + # A dispatched load that never reports back would otherwise hold this + # hash in LOADING for good, deferring every later request for it. + record.deadline = deadline self._transition(record, SchedulerTransferState.LOADING) self._loads_to_dispatch[record.transfer_id] = None return record @@ -307,6 +322,15 @@ def expire(self, now: float, terminal_limit: int) -> list[SchedulerTransfer]: now=now, ) expired.append(record) + elif record.state is SchedulerTransferState.LOADING: + self._transition( + record, + SchedulerTransferState.UNAVAILABLE, + now=now, + ) + if record.request_id: + self._notify_unavailable(record, record.request_id) + expired.append(record) elif record.state in _TERMINAL_STATES: self._remove(record.transfer_id) terminal_ids = [ @@ -337,12 +361,27 @@ def _insert(self, record: SchedulerTransfer) -> None: ) @staticmethod - def _check_identity(record: SchedulerTransfer, mm_hash: str) -> None: - if record.mm_hash and record.mm_hash != mm_hash: - raise ValueError( - f"Transfer {record.transfer_id!r} changed mm_hash from " - f"{record.mm_hash!r} to {mm_hash!r}" - ) + def _identity_matches(record: SchedulerTransfer, mm_hash: str) -> bool: + return not record.mm_hash or record.mm_hash == mm_hash + + def _refuse_colliding_id( + self, record: SchedulerTransfer, mm_hash: str, request_id: str + ) -> None: + """Report a transfer id that already names another encoding. + + Transfer ids arrive on the request, so a collision is reachable from + outside and must not reach the engine as an exception. + """ + logger.warning( + "EC Mooncake transfer_id=%s already names mm_hash=%s; refusing " + "mm_hash=%s for request %s", + record.transfer_id, + record.mm_hash[:16], + mm_hash[:16], + request_id or "", + ) + if request_id: + self._unavailable_requests.add(request_id) def _transition( self, diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py index cb5d38fe2dd6..e6d41c8c9621 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py @@ -90,6 +90,7 @@ def __init__(self, vllm_config: VllmConfig) -> None: self.is_producer = config.is_producer self.is_consumer = config.is_consumer self._buffer_device = config.buffer_device + self._control_host = config.control_host self._control_port = config.control_port self._transfer = MooncakeTransfer(get_ip(), config.protocol) self._consumer_memory = ConsumerMemoryPool( @@ -163,7 +164,10 @@ def start_services(self) -> None: ) base_port = self._control_port self._control_server = ConsumerControlServer( - "0.0.0.0", + # The reservation channel hands out registered-memory addresses + # and takes cancellations, so it listens where the Consumer + # advertises itself (`ec_ip`), not on every interface. + self._control_host, base_port + self._tp_rank, self._reserve_push_destination, self._push_status, From c1c954551d3c7cad53105729721b683c6639aeb5 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Fri, 4 Sep 2026 14:32:19 +0000 Subject: [PATCH 27/30] [EPD] Gate both WAITING paths on encoder-cache availability Signed-off-by: Tianyu Guo --- vllm/v1/core/sched/scheduler.py | 27 +++++++++++++++++++-------- 1 file changed, 19 insertions(+), 8 deletions(-) diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index fc1c44f926f8..9f0da7f49f97 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -930,13 +930,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: assert num_computed_tokens <= request.num_tokens # Skip request with pending mm encoding prefetches - if ( - self.ec_connector is not None - and request.mm_features - and not self.ec_connector.ensure_cache_available( - request, num_computed_tokens - ) - ): + if self._ec_transfer_pending(request, num_computed_tokens): request_queue.pop_request() step_skipped_waiting.prepend_request(request) continue @@ -951,11 +945,18 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: ) else: # KVTransfer: WAITING reqs have num_computed_tokens > 0 - # after async KV recvs are completed. + # after async KV recvs are completed. A streaming-input + # session resumes here too, carrying whatever media its + # latest chunk added, so this branch needs the same gate. new_computed_blocks = self.kv_cache_manager.empty_kv_cache_blocks num_new_local_computed_tokens = 0 num_computed_tokens = request.num_computed_tokens + if self._ec_transfer_pending(request, num_computed_tokens): + request_queue.pop_request() + step_skipped_waiting.prepend_request(request) + continue + encoder_inputs_to_schedule = None external_load_encoder_input = [] new_encoder_compute_budget = encoder_compute_budget @@ -2247,6 +2248,16 @@ def update_from_output( return engine_core_outputs + def _ec_transfer_pending(self, request: Request, num_computed_tokens: int) -> bool: + """Whether an encoder input this request needs is still in transit.""" + return ( + self.ec_connector is not None + and bool(request.mm_features) + and not self.ec_connector.ensure_cache_available( + request, num_computed_tokens + ) + ) + @staticmethod def _is_blocked_waiting_status(status: RequestStatus) -> bool: return status in ( From 875fe57da8a0f3dd1f8e8e86939bea4cccdb8b6d Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Sat, 5 Sep 2026 01:50:32 +0000 Subject: [PATCH 28/30] [EPD] Harden Mooncake EC transfers and batch asynchronous pushes Batch reservation RPCs, coalesce same-hash writes, and dispatch ready pushes without waiting for another model step. Share registered-slab allocation while retaining separate producer and consumer lifecycles. Harden control-plane failures, proxy metadata forwarding, and safe receive-buffer reclamation. Validation: 179 related tests passed; all applicable pre-commit hooks passed. Co-authored-by: OpenAI Codex Signed-off-by: Tianyu Guo --- .../disaggregated_encoder/disagg_epd_proxy.py | 24 +- .../unit/test_ec_mooncake_connector.py | 252 ++++++++++++++++-- .../ec_connector/unit/test_epd_proxy_retry.py | 29 +- .../ec_connector/mooncake/control.py | 47 +++- .../ec_connector/mooncake/memory.py | 164 ++++++------ .../ec_connector/mooncake/metadata.py | 3 + .../ec_connector/mooncake/producer.py | 48 ++-- .../ec_connector/mooncake/reservation.py | 116 +++++--- .../ec_connector/mooncake/scheduler.py | 15 +- .../ec_connector/mooncake/state.py | 64 +++-- .../ec_connector/mooncake/worker.py | 197 ++++++++++---- 11 files changed, 695 insertions(+), 264 deletions(-) diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py index 409621235acf..1581e21748d4 100644 --- a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -31,6 +31,7 @@ import time import uuid from collections.abc import AsyncIterator +from typing import Any import aiohttp import pybase64 as base64 @@ -228,7 +229,7 @@ async def fanout_encoder_primer( e_urls: list[str], req_id: str, consumer_zmq: str | None = None, -) -> tuple[dict[int, dict], dict[str, dict]]: +) -> tuple[dict[int, dict], dict[str, Any]]: """ 1. Build one request *per MM item* with all text removed. 2. Send them concurrently to the encode cluster. @@ -256,7 +257,7 @@ async def fanout_encoder_primer( item_uuids: dict[int, str] = {} item_transfer_ids: dict[int, str] = {} item_meta: dict[int, dict] = {} - ec_params: dict[str, dict] = {} + ec_params: dict[str, Any] = {} # Round-robin over encode servers to distribute load a bit. The cursor # persists across requests so fan-out doesn't restart at e_urls[0] every @@ -357,7 +358,7 @@ async def fanout_encoder_primer( except Exception: logger.warning("[%s] Could not read encoder metadata #%d", req_id, idx) params = {} - if idx in item_uuids: + if params: # One encoder request carries exactly one item, so there is a # single reported entry. Do not key it by this proxy's uuid: when # media_io_kwargs or mm_processor_kwargs are set the engine @@ -365,7 +366,7 @@ async def fanout_encoder_primer( # `mm_features[i].identifier` is a derived value this proxy cannot # predict. Fall back to the sole entry, and carry the key the # encoder actually used through as `ec_mm_hash`. - ec_mm_hash = item_uuids[idx] + ec_mm_hash = item_uuids.get(idx) reported = params.get(ec_mm_hash) if reported is None and len(params) == 1: ((ec_mm_hash, reported),) = params.items() @@ -374,7 +375,7 @@ async def fanout_encoder_primer( if metadata: item_meta[idx] = { **metadata, - "mm_hash": item_uuids[idx], + "mm_hash": item_uuids.get(idx, ec_mm_hash), "ec_mm_hash": ec_mm_hash, } if idx in item_transfer_ids: @@ -383,7 +384,11 @@ async def fanout_encoder_primer( # connector's own handle on the published embedding (for NIXL, # peer_host/peer_port/size_bytes). The decoder's connector # looks it up by mm_hash on the request, so carry it through. - ec_params[item_uuids[idx]] = reported + ec_params[item_uuids.get(idx, ec_mm_hash)] = reported + if NO_REWRITE and consumer_zmq is not None: + ec_params.setdefault("ec_items", []).append( + {"mm_hash": ec_mm_hash, "transfer_id": item_transfer_ids[idx]} + ) logger.info( "[%s] All %d encoder requests completed successfully", req_id, len(mm_items) @@ -533,8 +538,11 @@ async def on_startup() -> None: server_keep_alive = float(os.getenv("VLLM_HTTP_TIMEOUT_KEEP_ALIVE", "5")) connector = aiohttp.TCPConnector( limit=0, - force_close=False, - keepalive_timeout=max(1.0, server_keep_alive - 1.0), + **( + {"keepalive_timeout": server_keep_alive / 2} + if server_keep_alive > 0 + else {"force_close": True} + ), ) encode_session = aiohttp.ClientSession(timeout=timeout, connector=connector) if app.state.p_urls: diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py index c7179bdac833..568e4ac502d4 100644 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py @@ -78,7 +78,10 @@ from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( MooncakeTransfer, ) -from vllm.distributed.ec_transfer.ec_connector.mooncake.worker import ECMooncakeWorker +from vllm.distributed.ec_transfer.ec_connector.mooncake.worker import ( + ECMooncakeWorker, + _FanoutError, +) from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( ECMooncakeConnector, ) @@ -202,6 +205,16 @@ def _bind_extra_config(config: VllmConfig) -> None: class TestECMooncakeControlPlane: """Validate ZMQ client reuse, shard discovery, events, and server RPCs.""" + def test_event_transport_failure_stops_draining(self): + inbox = EventInbox(Mock()) + socket = Mock() + socket.recv_json.side_effect = zmq.ZMQError(zmq.ENOTSOCK) + inbox._socket = socket + assert inbox.drain("unused") == [] + socket.recv_json.assert_called_once() + inbox.close() + assert inbox._socket is None + def test_worker_get_ip_failure_does_not_construct_client( self, mock_vllm_config_producer ): @@ -249,6 +262,7 @@ def test_client_reuses_socket_and_discards_failed_exchange(self): context.socket.assert_called_once_with(zmq.REQ) assert socket.setsockopt.call_args_list == [ + call(zmq.IPV6, 1), call(zmq.RCVTIMEO, 17), call(zmq.SNDTIMEO, 17), call(zmq.LINGER, 0), @@ -493,14 +507,8 @@ def test_server_keeps_serving_after_an_undecodable_request(self): client.close() server.close() - def test_a_dead_writer_gives_its_destination_back(self): - """Reserve, then vanish: the receive pool has to recover on its own. - - Driven over the wire the way a Producer and the Scheduler drive it -- - `reserve`, then the Scheduler's own cancel, which does not abandon the - writer. That cancel used to leave the destination pinned with no path - back, so a Producer crash permanently cost the Consumer a buffer. - """ + def test_an_unconfirmed_writer_keeps_its_destination_reserved(self): + """Only writer confirmation makes a timed-out destination reusable.""" engine = MagicMock(spec=MooncakeTransfer) engine.register_memory.return_value = 0 engine.unregister_memory.return_value = True @@ -567,11 +575,15 @@ def request_destination(transfer_id: str) -> None: abandoned += 1 assert abandoned, "the pool should accept a writer before filling up" - # No in-flight write can still land, so the buffers come back and - # an unrelated request gets a destination again. for record in manager._records.values(): record.expires_at = 0 - request_destination("after-the-grace") + with pytest.raises(RuntimeError, match="pool is full"): + request_destination("after-the-grace") + for transfer_id in list(manager._records): + client.request( + addr, control.make_cancel_request(transfer_id, abandon=True) + ) + request_destination("after-writer-confirmation") finally: client.close() server.close() @@ -968,14 +980,53 @@ def test_producer_reuses_staging_and_unregisters_parent_on_close(self): assert second is not None assert second.regions == first.regions pool.release(second) - parent = pool._pool + parent = pool.tensor pool.close() pool.close() - assert pool._pool is None + assert pool.tensor is None mooncake_transfer.unregister_memory.assert_called_once_with(parent) + @pytest.mark.parametrize("pool_type", [ProducerMemoryPool, ConsumerMemoryPool]) + def test_failed_unregister_keeps_the_registered_buffer_until_retry(self, pool_type): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 0 + mooncake_transfer.unregister_memory.side_effect = [False, True] + pool = pool_type(256, mooncake_transfer) + if isinstance(pool, ProducerMemoryPool): + staged = pool.stage([torch.ones(16)]) + assert staged is not None + pool.release(staged) + else: + pool.prepare(torch.device("cpu")) + parent = pool.tensor + assert parent is not None + + pool.close() + assert pool.tensor is parent + pool.close() + assert pool.tensor is None + pool.close() + assert mooncake_transfer.unregister_memory.call_args_list == [ + call(parent), + call(parent), + ] + + def test_producer_failed_batch_returns_reserved_regions(self): + mooncake_transfer = MagicMock(spec=MooncakeTransfer) + mooncake_transfer.register_memory.return_value = 0 + pool = ProducerMemoryPool(256, mooncake_transfer) + source = torch.arange(16, dtype=torch.float32) + + assert pool.stage([source, source]) is None + staged = pool.stage([source]) + assert staged is not None + assert torch.equal(staged.tensors[0], source) + mooncake_transfer.register_memory.assert_called_once() + pool.release(staged) + pool.close() + def test_producer_falls_back_when_staging_pool_allocation_fails(self): mooncake_transfer = MagicMock(spec=MooncakeTransfer) pool = ProducerMemoryPool(256, mooncake_transfer) @@ -1322,11 +1373,26 @@ def test_a_colliding_transfer_id_fails_only_the_request_that_named_it(self): assert table.take_unavailable_requests() == {"req-b"} assert first.mm_hash == "hash" assert first.request_id == "req-a" + assert not table.cancel("shared", 2, "other-hash", "req-b") + assert first.state is SchedulerTransferState.WAITING_EVENT class TestECMooncakeSchedulerMetadata: """Validate Scheduler decisions and per-step Worker metadata.""" + @pytest.mark.parametrize("payload", [[], "ready", 1, None]) + def test_non_object_events_are_discarded(self, payload): + scheduler = object.__new__(ECMooncakeScheduler) + scheduler._drain_pending = True + scheduler._control_addr = "unused" + scheduler._poll_pending_cancels = Mock() + scheduler._expire_transfers = Mock() + scheduler._event_inbox = Mock() + scheduler._event_inbox.drain.return_value = [payload] + scheduler._accept_ready_event = Mock() + scheduler._drain_push_notifications() + scheduler._accept_ready_event.assert_not_called() + def test_an_unreachable_consumer_does_not_block_the_scheduler( self, mock_vllm_config_consumer ): @@ -1970,6 +2036,50 @@ def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer class TestConsumerReservationManager: + @pytest.mark.parametrize("cancelled", [None, "writer", "follower"]) + def test_inflight_duplicates_share_memory_and_cancel_independently(self, cancelled): + engine = Mock() + engine.register_memory.return_value = 0 + pool = ConsumerMemoryPool(256, engine) + pool.prepare(torch.device("cpu")) + manager = ConsumerReservationManager(pool, 300, 16) + writer, write = manager.reserve( + "writer", "hash", 64, (16,), "float32", torch.float32 + ) + follower, write_again = manager.reserve( + "follower", "hash", 64, (16,), "float32", torch.float32 + ) + assert write and not write_again + assert writer.allocation is follower.allocation + tensor = writer.allocation.tensor + tensor.fill_(7) + if cancelled is not None: + manager.cancel(cancelled, "") + assert manager.complete("writer", writer.reservation_id)[0] + for name in ("writer", "follower"): + if name != cancelled: + loaded = manager.take(name, "hash") + assert loaded.tensor.data_ptr() == tensor.data_ptr() + assert torch.all(loaded.tensor == 7) + assert pool.try_allocate(64, (16,), torch.float32) is None + + def test_follower_refresh_does_not_abandon_the_shared_writer(self): + manager, pool, _ = self.manager() + writer, _ = self.reserve(manager) + follower, _ = manager.reserve( + "follower", "hash", 64, (16,), "float32", torch.float32 + ) + assert manager.cancel( + "follower", follower.reservation_id, abandon=True, refresh=True + ) + renewed, write = manager.reserve( + "follower", "hash", 64, (16,), "float32", torch.float32 + ) + assert not write and renewed.writer_id == writer.transfer_id + assert renewed.reservation_id != follower.reservation_id + assert writer.state is ConsumerReservationState.WRITING + pool.free.assert_not_called() + @staticmethod def manager(): pool = Mock() @@ -2084,15 +2194,8 @@ def test_cancel_tombstones_are_bounded(self): ), ], ) - def test_a_deferred_release_is_reclaimed_when_no_writer_reports( - self, deferred, terminal - ): - """A writer that never reports must not pin its destination for good. - - Cancellation and expiry defer to the remote writer, so a Producer that - died between reserve and complete used to hold the buffer for the life - of the process and every later reserve saw a full pool. - """ + def test_a_deferred_release_requires_writer_completion(self, deferred, terminal): + """A timeout cannot prove that a remote writer stopped using its address.""" manager, pool, allocation = self.manager() record, _ = self.reserve(manager) @@ -2106,7 +2209,9 @@ def test_a_deferred_release_is_reclaimed_when_no_writer_reports( pool.free.assert_not_called() record.expires_at = 0 - assert manager.expire() == 1 + assert manager.expire() == 0 + pool.free.assert_not_called() + assert manager.complete("transfer", record.reservation_id) == (True, False) assert record.state is terminal pool.free.assert_called_once_with(allocation) @@ -2114,6 +2219,94 @@ def test_a_deferred_release_is_reclaimed_when_no_writer_reports( class TestECMooncakeWorkerTransfer: """Validate end-to-end Worker reservation, push, load, and cleanup flows.""" + def test_partial_batch_reservation_cleans_only_the_failed_item(self): + """An item rejected on one shard must not cancel its successful sibling.""" + worker = object.__new__(ECMooncakeWorker) + worker._control_client = Mock() + shards = ["tcp://consumer:0", "tcp://consumer:1"] + worker._control_client.discover_shards.return_value = shards + specs = [ + ECMooncakePushSpec(name, 64, (16,), "float32", shards[0], name) + for name in ("good", "bad") + ] + + def request(addr, payload): + assert payload["op"] == "reserve_batch" + return { + "items": [ + {"ok": True, "result": {"reservation_id": "good-" + addr}}, + {"ok": False, "error": "full"} + if addr == shards[1] + else {"ok": True, "result": {"reservation_id": "partial"}}, + ] + } + + worker._control_client.request.side_effect = request + worker._run_fanout = lambda tasks: [task() for task in tasks] + worker._retry_cancel_reservations = Mock() + good, bad = worker._reserve_remote_many(specs) + assert [item["addr"] for item in good] == shards + assert isinstance(bad, _FanoutError) and str(bad) == "full" + worker._retry_cancel_reservations.assert_called_once_with(specs[1], bad.results) + assert [item["reservation_id"] for item in bad.results] == ["partial", ""] + + def test_reservation_completion_dispatches_without_another_model_step( + self, mock_vllm_config_producer + ): + allow_reservation = threading.Event() + transferred = threading.Event() + source = torch.ones(16) + spec = ECMooncakePushSpec("hash", 64, (16,), "float32", "unused", "transfer") + with patch_ec_mooncake_deps(): + worker = ECMooncakeWorker(mock_vllm_config_producer) + + def reserve(_): + assert allow_reservation.wait(5) + return [] + + def write(records): + for record in records: + worker._producer_pushes.begin_writing(record) + worker._producer_pushes.begin_notifying(records) + worker._producer_pushes.complete(records) + transferred.set() + + worker._reserve_remote = reserve + worker._push_batch = write + try: + worker.start_save_caches( + ECMooncakeConnectorMetadata(pushes=[spec]), {"hash": source} + ) + worker.build_connector_worker_meta() + assert not transferred.is_set() + allow_reservation.set() + assert transferred.wait(5) + finally: + allow_reservation.set() + worker.close() + + @pytest.mark.parametrize("waiting_for", ["reservation", "encoder"]) + def test_a_ready_push_does_not_wait_for_another_item(self, waiting_for): + manager = ProducerPushManager() + records = [] + for name in ("slow", "fast"): + future: Future[list[dict[str, Any]]] = Future() + spec = ECMooncakePushSpec( + name, 64, (16,), "float32", "tcp://consumer:1", name + ) + record, _ = manager.reserve(spec, lambda future=future: future) + event = Mock() + event.query.return_value = name == "fast" or waiting_for == "reservation" + manager.bind_source(name, torch.empty(16), event) + if name == "fast" or waiting_for == "encoder": + future.set_result([]) + records.append(record) + executor = Mock() + executor.submit.return_value = Future() + run = Mock() + manager.submit_batches(executor, run) + executor.submit.assert_called_once_with(run, [records[1]]) + def test_reservation_requires_confirmed_topology_before_any_rpc(self): worker = object.__new__(ECMooncakeWorker) worker._control_client = Mock() @@ -2204,7 +2397,8 @@ def test_producer_push_state_owns_source_until_every_future_is_terminal(self): # failing the engine; the push in flight keeps the id. changed = copy.copy(spec) changed.mm_hash = "other" - assert manager.reserve(changed, lambda: Future()) == (record, False) + with pytest.raises(ValueError, match="Conflicting EC destination"): + manager.reserve(changed, lambda: Future()) assert record.spec.mm_hash == "hash" source = torch.empty(16) @@ -2321,8 +2515,8 @@ def finish_cancel(orphan: ProducerPushRecord): "flush", "io.submit", "cancel:orphan", - ("io.shutdown", True, {}), ("control.shutdown", True, {}), + ("io.shutdown", True, {}), ("shard.shutdown", True, {}), ] worker._control_client.close.assert_called_once_with() @@ -3396,11 +3590,11 @@ def test_pushes_stage_through_the_registered_pool(self, mock_vllm_config_produce assert isinstance(engine, CopyingFakeTransferEngine) # The staging pool is registered once; a transfer registers # nothing of its own. - pool = producer._worker._producer_memory._pool + pool = producer._worker._producer_memory.tensor assert pool is not None assert engine.register_calls == [[pool.data_ptr()]] assert engine.batch_unregister_calls == [] - assert engine.transfer_calls == [[source.nbytes, source.nbytes]] + assert engine.transfer_calls == [[source.nbytes]] assert all( reservation.state is ConsumerReservationState.READY for reservation in consumer._worker._reservations._records.values() @@ -3442,7 +3636,7 @@ def test_push_falls_back_to_per_tensor_registration_without_a_pool( _wait_for_worker_io(producer) engine = producer._worker._transfer._engine assert isinstance(engine, CopyingFakeTransferEngine) - assert producer._worker._producer_memory._pool is None + assert producer._worker._producer_memory.tensor is None assert engine.register_calls == [[source.data_ptr()]] assert engine.batch_unregister_calls == [[source.data_ptr()]] assert engine.transfer_calls == [[source.nbytes]] diff --git a/tests/v1/ec_connector/unit/test_epd_proxy_retry.py b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py index 3dc073888e77..b026b5537a20 100644 --- a/tests/v1/ec_connector/unit/test_epd_proxy_retry.py +++ b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py @@ -120,7 +120,34 @@ async def _no_prefill(req_data, p_url, req_id): assert "ec_transfer_params" not in second -@pytest.mark.parametrize("server_keep_alive", ["5", "2", "30"]) +def test_raw_media_keeps_encoder_transfer_identity(proxy, monkeypatch): + handle = {"metadata": {"image_grid_thw": [1, 2, 2]}} + monkeypatch.setattr(proxy, "NO_REWRITE", True) + monkeypatch.setattr( + proxy, "encode_session", _EncoderSession([{"encoded-hash": handle}]) + ) + body = { + "messages": [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "image"}}], + } + ] + } + prepared, _, _ = asyncio.run( + proxy.prepare_for_decode( + body, "request", ["http://encoder"], "", "tcp://consumer:1" + ) + ) + assert prepared["messages"] == body["messages"] + assert "ec_transfer_params" not in body + params = prepared["ec_transfer_params"] + assert params["encoded-hash"] == handle + assert params["ec_items"][0]["mm_hash"] == "encoded-hash" + assert params["ec_items"][0]["transfer_id"] + + +@pytest.mark.parametrize("server_keep_alive", ["0.1", "1", "5", "2", "30"]) def test_pooled_connections_are_retired_before_the_server_closes_them( proxy, monkeypatch, server_keep_alive ): diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py index 67b870baa26f..2f4a0ad6058e 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/control.py @@ -20,6 +20,7 @@ import zmq from vllm.logger import init_logger +from vllm.utils.network_utils import make_zmq_path logger = init_logger(__name__) @@ -58,6 +59,7 @@ def _exchange(self, addr: str, payload: ControlRequest) -> ControlResponse: socket = sockets.get(addr) if socket is None: socket = self._context.socket(zmq.REQ) + socket.setsockopt(zmq.IPV6, 1) socket.setsockopt(zmq.RCVTIMEO, self._timeout_ms) socket.setsockopt(zmq.SNDTIMEO, self._timeout_ms) socket.setsockopt(zmq.LINGER, 0) @@ -194,6 +196,7 @@ def _install(self, endpoints: list[str] | None) -> None: return context = zmq.Context() socket = context.socket(zmq.PULL) + socket.setsockopt(zmq.IPV6, 1) for endpoint in endpoints: socket.connect(endpoint) self._context = context @@ -210,7 +213,10 @@ def drain(self, base_addr: str) -> list[dict[str, Any]]: events.append(self._socket.recv_json(flags=zmq.DONTWAIT)) except zmq.Again: return events - except Exception: + except zmq.ZMQError: + logger.warning("EC Mooncake event channel failed", exc_info=True) + return events + except (ValueError, UnicodeError): # The event channel is a plain PULL socket: an undecodable # frame must cost one frame, not the engine. logger.warning( @@ -226,6 +232,7 @@ def close(self) -> None: self._discovery = None if self._socket is not None: self._socket.close(linger=0) + self._socket = None if self._context is not None: self._context.term() @@ -249,6 +256,7 @@ def __init__( reap: Callable[[], int], peer_ports: list[int] | None = None, device: torch.device | None = None, + drain_ready: Callable[[], list[str]] = lambda: [], ) -> None: self.host = host self.port = port @@ -260,6 +268,7 @@ def __init__( self._complete = complete self._cancel = cancel self._reap = reap + self._drain_ready = drain_ready self._stop = threading.Event() self._started = threading.Event() self._thread: threading.Thread | None = None @@ -272,6 +281,8 @@ def loop() -> None: context = zmq.Context() socket = context.socket(zmq.REP) event_socket = context.socket(zmq.PUSH) + socket.setsockopt(zmq.IPV6, 1) + event_socket.setsockopt(zmq.IPV6, 1) pending_events: deque[dict[str, Any]] = deque() def queue_event(event: dict[str, Any]) -> None: @@ -294,8 +305,13 @@ def queue_ready(transfer_id: str) -> None: last_reap_at = time.monotonic() socket.setsockopt(zmq.RCVTIMEO, 100) try: - socket.bind(f"tcp://{self.host}:{self.port}") - self.event_port = event_socket.bind_to_random_port(f"tcp://{self.host}") + socket.bind(make_zmq_path("tcp", self.host, self.port)) + event_socket.bind(make_zmq_path("tcp", self.host, 0)) + self.event_port = int( + event_socket.getsockopt(zmq.LAST_ENDPOINT) + .decode() + .rsplit(":", 1)[1] + ) except Exception as e: self._startup_error = e self._started.set() @@ -337,10 +353,25 @@ def queue_ready(transfer_id: str) -> None: try: op = request.get("op") result: Any = None - if op == "reserve": - result = self._reserve(request) - if result.get("ready"): - queue_ready(str(request["transfer_id"])) + if op in ("reserve", "reserve_batch"): + items = ( + request["items"] if op == "reserve_batch" else [request] + ) + results = [] + for item in items: + try: + reserved = self._reserve(item) + results.append({"ok": True, "result": reserved}) + if reserved.get("ready"): + queue_ready(str(item["transfer_id"])) + except Exception as exc: + results.append({"ok": False, "error": str(exc)}) + if op == "reserve_batch": + result = {"items": results} + elif not results[0]["ok"]: + raise RuntimeError(results[0]["error"]) + else: + result = results[0]["result"] elif op == "status": result = self._status(str(request["transfer_id"])) elif op == "event_port": @@ -359,6 +390,8 @@ def queue_ready(transfer_id: str) -> None: accepted, became_ready = self._complete( transfer_id, str(item["reservation_id"]) ) + for ready_id in self._drain_ready(): + queue_ready(ready_id) completions.append( { "completed": accepted, diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py index 1dd812dd7666..c6bc1c68515e 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/memory.py @@ -60,11 +60,43 @@ def value(self) -> _T: class ContiguousAllocator: - """Allocate aligned regions from one contiguous byte range.""" + """Own a registered slab and its aligned regions; callers serialize access.""" def __init__(self, capacity: int, alignment: int = 256): self.alignment = alignment + self._capacity = capacity self._free = [(0, capacity)] + self.tensor: torch.Tensor | None = None + self._disabled = False + + def prepare(self, device: torch.device, transfer: MooncakeTransfer) -> None: + if self.tensor is not None or self._disabled: + return + try: + tensor = torch.empty(self._capacity, dtype=torch.uint8, device=device) + ret = transfer.register_memory(tensor) + if ret != 0: + raise RuntimeError(f"Mooncake returned {ret}") + except (RuntimeError, torch.OutOfMemoryError) as error: + self._disabled = True + logger.warning("Could not initialize the EC registered buffer: %s", error) + return + self.tensor = tensor + self._free = [(0, tensor.nbytes)] + logger.info("Registered %d-byte buffer for Mooncake EC", tensor.nbytes) + + def view( + self, offset: int, nbytes: int, shape: tuple[int, ...], dtype: torch.dtype + ) -> torch.Tensor: + assert self.tensor is not None + return self.tensor.narrow(0, offset, nbytes).view(dtype).view(shape) + + def close(self, transfer: MooncakeTransfer) -> bool: + if self.tensor is not None and not transfer.unregister_memory(self.tensor): + return False + self.tensor = None + self._free.clear() + return True def allocate(self, nbytes: int) -> tuple[int, int] | None: size = (nbytes + self.alignment - 1) // self.alignment * self.alignment @@ -198,41 +230,16 @@ class ProducerMemoryPool: """Own the Producer staging slab and regions carved from it.""" def __init__(self, capacity: int, transfer: MooncakeTransfer) -> None: - self._capacity = capacity self._transfer = transfer - self._pool: torch.Tensor | None = None - self._allocator: ContiguousAllocator | None = None - self._disabled = False + self._allocator = ContiguousAllocator(capacity) self._lock = threading.Lock() + self._local = threading.local() - def _ensure_pool(self, device: torch.device) -> None: - if self._pool is not None or self._disabled: - return - with self._lock: - if self._pool is not None or self._disabled: - return - try: - pool = torch.empty(self._capacity, dtype=torch.uint8, device=device) - ret = self._transfer.register_memory(pool) - if ret != 0: - raise RuntimeError(f"Mooncake returned {ret}") - except (RuntimeError, torch.OutOfMemoryError) as error: - self._disabled = True - logger.warning( - "Could not initialize the EC producer staging pool; falling " - "back to per-transfer registration: %s", - error, - ) - return - self._pool = pool - self._allocator = ContiguousAllocator(pool.nbytes) - logger.info( - "Registered %d-byte staging pool for Mooncake EC pushes", - pool.nbytes, - ) + @property + def tensor(self) -> torch.Tensor | None: + return self._allocator.tensor def _free_regions(self, regions: list[tuple[int, int]]) -> None: - assert self._allocator is not None for offset, size in regions: self._allocator.free(offset, size) @@ -240,14 +247,14 @@ def stage(self, tensors: list[torch.Tensor]) -> StagedSources | None: """Copy tensors into one registered slab, or return None for fallback.""" if not tensors: return StagedSources([], []) - self._ensure_pool(tensors[0].device) - pool = self._pool allocator = self._allocator - if pool is None or allocator is None: - return None staged: list[torch.Tensor] = [] regions: list[tuple[int, int]] = [] with self._lock: + allocator.prepare(tensors[0].device, self._transfer) + pool = allocator.tensor + if pool is None: + return None for tensor in tensors: region = allocator.allocate(tensor.nbytes) if region is None: @@ -255,12 +262,22 @@ def stage(self, tensors: list[torch.Tensor]) -> StagedSources | None: return None regions.append(region) staged.append( - pool.narrow(0, region[0], tensor.nbytes) - .view(tensor.dtype) - .view(tensor.shape) + allocator.view( + region[0], tensor.nbytes, tuple(tensor.shape), tensor.dtype + ) ) - for destination, source in zip(staged, tensors): - destination.copy_(source, non_blocking=True) + if pool.device.type == "cuda": + stream = getattr(self._local, "stream", None) + if stream is None: + stream = torch.cuda.Stream(device=pool.device) + self._local.stream = stream + with torch.cuda.stream(stream): + for destination, source in zip(staged, tensors): + destination.copy_(source, non_blocking=True) + stream.synchronize() + else: + for destination, source in zip(staged, tensors): + destination.copy_(source) return StagedSources(staged, regions) def release(self, staged: StagedSources) -> None: @@ -271,11 +288,7 @@ def release(self, staged: StagedSources) -> None: def close(self) -> None: with self._lock: - pool = self._pool - if pool is None or not self._transfer.unregister_memory(pool): - return - self._pool = None - self._allocator = None + self._allocator.close(self._transfer) class ConsumerMemoryPool: @@ -286,48 +299,26 @@ def __init__( capacity: int, transfer: MooncakeTransfer, ) -> None: - self._capacity = capacity self._transfer = transfer - self._pool: torch.Tensor | None = None - self._allocator: ContiguousAllocator | None = None + self._allocator = ContiguousAllocator(capacity) self._residents: ResidentPool[MemoryAllocation] = ResidentPool() self._retire_events: dict[str, torch.Event] = {} self._pending_frees: list[tuple[torch.Event, MemoryAllocation]] = [] self._reclaimed: set[str] = set() - self._disabled = False self.lock = threading.RLock() @property def tensor(self) -> torch.Tensor | None: - return self._pool + return self._allocator.tensor def prepare( self, device: torch.device, ) -> None: - if self._pool is not None or self._disabled: - return - try: - pool = torch.empty(self._capacity, dtype=torch.uint8, device=device) - ret = self._transfer.register_memory(pool) - if ret != 0: - raise RuntimeError(f"Mooncake returned {ret}") - except (RuntimeError, torch.OutOfMemoryError) as error: - self._disabled = True - logger.warning( - "Could not initialize the EC consumer buffer pool: %s", - error, - ) - return - self._pool = pool - self._allocator = ContiguousAllocator(pool.nbytes) - logger.info( - "Prepared %d-byte receive pool for Mooncake EC", - pool.nbytes, - ) + with self.lock: + self._allocator.prepare(device, self._transfer) def _free(self, allocation: MemoryAllocation) -> None: - assert self._allocator is not None self._allocator.free(allocation.offset, allocation.size) def free(self, allocation: MemoryAllocation) -> None: @@ -352,8 +343,6 @@ def _poll_frees_locked(self) -> None: self._pending_frees = pending def _reclaim_locked(self, nbytes: int) -> tuple[int, int] | None: - assert self._allocator is not None - def evict(mm_hash: str, allocation: MemoryAllocation) -> bool: event = self._retire_events.pop(mm_hash, None) self._defer_or_free(allocation, event) @@ -377,9 +366,8 @@ def _make_allocation( shape: tuple[int, ...], dtype: torch.dtype, ) -> MemoryAllocation: - assert self._pool is not None offset, size = region - tensor = self._pool.narrow(0, offset, nbytes).view(dtype).view(shape) + tensor = self._allocator.view(offset, nbytes, shape, dtype) return MemoryAllocation(offset, size, tensor) def try_allocate( @@ -387,7 +375,7 @@ def try_allocate( ) -> MemoryAllocation | None: with self.lock: allocator = self._allocator - assert self._pool is not None and allocator is not None + assert allocator.tensor is not None self._poll_frees_locked() region = allocator.allocate(nbytes) if region is None: @@ -435,10 +423,11 @@ def take_resident( return tensor def _record_release_event(self) -> torch.Event | None: - if self._pool is None or self._pool.device.type != "cuda": + pool = self.tensor + if pool is None or pool.device.type != "cuda": return None event = torch.Event() - event.record(torch.accelerator.current_stream(self._pool.device)) + event.record(torch.accelerator.current_stream(pool.device)) return event def release_cached(self, lease: ResidentLease[MemoryAllocation]) -> None: @@ -452,6 +441,8 @@ def publish( mm_hash: str, allocation: MemoryAllocation, lease: ResidentLease[MemoryAllocation] | None = None, + *, + pin: bool = True, ) -> MemoryAllocation: with self.lock: if lease is not None: @@ -462,6 +453,8 @@ def publish( return canonical previous = self._residents.get(mm_hash) displaced = self._residents.insert(mm_hash, allocation) + if not pin: + self._residents.retire(mm_hash) event = None if previous is not None and previous is not allocation: event = self._retire_events.pop(mm_hash, None) @@ -475,18 +468,20 @@ def publish( def retire_stale( self, encoder_cache: dict[str, torch.Tensor], - reserved_hashes: set[str], + reserved_hashes: set[str] | None = None, + *, + freed: list[str] | None = None, ) -> None: - if self._pool is None: + if self.tensor is None: return with self.lock: - for mm_hash in self._residents.referenced(): + for mm_hash in self._residents.referenced() if freed is None else freed: allocation = self._residents.get(mm_hash) if allocation is None: continue if encoder_cache.get(mm_hash) is allocation.tensor: continue - if mm_hash in reserved_hashes: + if reserved_hashes and mm_hash in reserved_hashes: continue event = self._record_release_event() if event is not None: @@ -502,11 +497,8 @@ def drain_reclaimed(self) -> set[str]: def close(self) -> None: with self.lock: - pool = self._pool - if pool is None or not self._transfer.unregister_memory(pool): + if not self._allocator.close(self._transfer): return - self._pool = None - self._allocator = None self._residents.clear() self._retire_events.clear() self._pending_frees.clear() diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py index 2ccf80911ef2..a006d1bf2f80 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/metadata.py @@ -45,6 +45,7 @@ class ECMooncakeConnectorMetadata(ECConnectorMetadata): loads: list[ECMooncakeLoadSpec] = field(default_factory=list) pushes: list[ECMooncakePushSpec] = field(default_factory=list) + freed: list[str] | None = None @dataclass @@ -57,6 +58,7 @@ class ECMooncakeWorkerMetadata(ECConnectorWorkerMetadata): # evicted item stays resident until told otherwise. reclaimed: set[str] = field(default_factory=set) pending_saves: bool = False + failed_saves: set[str] = field(default_factory=set) def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadata: assert isinstance(other, ECMooncakeWorkerMetadata) @@ -69,4 +71,5 @@ def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadat failed_loads=self.failed_loads | other.failed_loads, reclaimed=self.reclaimed | other.reclaimed, pending_saves=self.pending_saves or other.pending_saves, + failed_saves=self.failed_saves | other.failed_saves, ) diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py index 7e99779cd015..48b4e4dd585e 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/producer.py @@ -22,9 +22,6 @@ from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( ECMooncakePushSpec, ) -from vllm.logger import init_logger - -logger = init_logger(__name__) def _same_destination( @@ -110,7 +107,7 @@ class ProducerPushManager: ordered indexes keep hot polling paths from scanning those tombstones. """ - def __init__(self) -> None: + def __init__(self, wake: Callable[[], None] = lambda: None) -> None: self._records: OrderedDict[str, ProducerPushRecord] = OrderedDict() self._active_ids: OrderedDict[str, None] = OrderedDict() self._reapable_terminal_ids: OrderedDict[str, None] = OrderedDict() @@ -118,6 +115,7 @@ def __init__(self) -> None: self._batch_ids: OrderedDict[str, None] = OrderedDict() self._source_waiters: dict[str, OrderedDict[str, None]] = {} self._lock = threading.RLock() + self._wake = wake def reserve( self, @@ -128,18 +126,8 @@ def reserve( existing = self._records.get(spec.transfer_id) if existing is not None: if not _same_destination(existing.spec, spec): - # Transfer ids come in on the request, so two requests can - # name one id. Keep the push already in flight and drop - # the newcomer: the consumer times its own wait out and - # fails that request, which an engine-level raise would - # not. - logger.warning( - "EC Mooncake producer transfer_id=%s already pushes " - "mm_hash=%s to %s; dropping mm_hash=%s", - spec.transfer_id, - existing.spec.mm_hash[:16], - existing.spec.consumer_zmq, - spec.mm_hash[:16], + raise ValueError( + f"Conflicting EC destination for transfer_id={spec.transfer_id}" ) return existing, False record = ProducerPushRecord( @@ -174,12 +162,16 @@ def bind_source( continue record.source_tensor = tensor record.source_event = ready_event + self._wake() def submit_batches( self, executor: ThreadPoolExecutor, run_batch: Callable[[list[ProducerPushRecord]], None], - ) -> None: + *, + wait: bool = False, + ) -> bool: + pending_event = False with self._lock: grouped: dict[str, list[ProducerPushRecord]] = {} for transfer_id in list(self._active_ids): @@ -189,6 +181,14 @@ def submit_batches( and record.batch_future is None and record.state is ProducerPushState.WAITING_INPUTS ): + if not record.reservation_future.done(): + continue + if record.source_event is not None: + if wait: + record.source_event.synchronize() + elif not record.source_event.query(): + pending_event = True + continue grouped.setdefault(record.spec.consumer_zmq, []).append(record) batches = list(grouped.values()) for records in batches: @@ -196,6 +196,7 @@ def submit_batches( for record in records: record.batch_future = future self._batch_ids[record.spec.transfer_id] = None + return pending_event def resolve_reservations(self, record: ProducerPushRecord) -> list[dict[str, Any]]: results = record.reservation_future.result() @@ -203,6 +204,17 @@ def resolve_reservations(self, record: ProducerPushRecord) -> list[dict[str, Any record.reservations = results return list(results) + def finish_reservations( + self, + outcomes: list[tuple[ProducerPushRecord, list[dict[str, Any]] | Exception]], + ) -> None: + with self._lock: + for record, result in outcomes: + if isinstance(result, Exception): + record.reservation_future.set_exception(result) + else: + record.reservation_future.set_result(result) + def _reservation_done( self, record: ProducerPushRecord, @@ -218,6 +230,8 @@ def _reservation_done( and record.source_tensor is None ): self._transition(record, ProducerPushState.FAILED) + finally: + self._wake() def settle_all(self, records: list[ProducerPushRecord]) -> None: for record in records: diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py index 0c6e05358d35..407b3911934d 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/reservation.py @@ -22,9 +22,6 @@ MemoryAllocation, ResidentLease, ) -from vllm.logger import init_logger - -logger = init_logger(__name__) class ConsumerReservationState(Enum): @@ -88,6 +85,7 @@ class ConsumerReservation: allocation: MemoryAllocation | None = None lease: ResidentLease[MemoryAllocation] | None = None expires_at: float = 0 + writer_id: str = "" class ConsumerReservationManager: @@ -106,13 +104,13 @@ def __init__( ) -> None: self._memory = memory self._lease_ttl = lease_ttl - # A writer that has not reported for a full extra lease is gone; no - # RDMA write outlives that, so reclaiming its destination is safe. - self._writer_grace = lease_ttl self._tombstone_limit = tombstone_limit self._records: dict[str, ConsumerReservation] = {} self._active_ids: dict[str, None] = {} self._tombstones: OrderedDict[str, None] = OrderedDict() + self._writers: dict[str, ConsumerReservation] = {} + self._followers: dict[str, dict[str, None]] = {} + self._ready: list[str] = [] self._condition = threading.Condition(memory.lock) self._shutting_down = False @@ -179,6 +177,24 @@ def reserve( ) self._insert(record) return record, False + writer = self._writers.get(mm_hash) + if writer is not None and writer.state is ConsumerReservationState.WRITING: + if writer.shape != shape or writer.dtype != dtype_name: + raise ValueError("conflicting in-flight tensor for mm_hash") + record = ConsumerReservation( + transfer_id, + mm_hash, + uuid.uuid4().hex, + ConsumerReservationState.WRITING, + shape, + dtype_name, + writer.allocation, + expires_at=now + self._lease_ttl, + writer_id=writer.transfer_id, + ) + self._insert(record) + self._followers.setdefault(writer.transfer_id, {})[transfer_id] = None + return record, False allocation = self._memory.try_allocate(nbytes, shape, dtype) if allocation is None: self._expire_locked(time.monotonic()) @@ -199,6 +215,7 @@ def reserve( now + self._lease_ttl, ) self._insert(record) + self._writers[mm_hash] = record return record, True def status(self, transfer_id: str) -> ConsumerReservation | None: @@ -216,6 +233,10 @@ def complete(self, transfer_id: str, reservation_id: str) -> tuple[bool, bool]: return False, False if record.state is ConsumerReservationState.READY: return True, False + if record.writer_id: + return False, False + if record.state in _WRITER_OWNED_STATES: + self._publish_followers(record) if record.state in _DEFERRED_STATES: terminal = ( ConsumerReservationState.CANCELLED @@ -232,6 +253,31 @@ def complete(self, transfer_id: str, reservation_id: str) -> tuple[bool, bool]: finally: self._condition.notify_all() + def _publish_followers(self, writer: ConsumerReservation) -> None: + if self._writers.get(writer.mm_hash) is writer: + self._writers.pop(writer.mm_hash) + followers = self._followers.pop(writer.transfer_id, {}) + if not followers: + return + assert writer.allocation is not None + allocation = self._memory.publish(writer.mm_hash, writer.allocation, pin=False) + for record in [writer, *(self._records[key] for key in followers)]: + record.allocation = allocation + record.lease = self._memory.acquire_cached( + record.mm_hash, record.shape, allocation.tensor.dtype + ) + assert record.lease is not None + if record is not writer: + record.writer_id = "" + self._transition(record, ConsumerReservationState.READY) + record.expires_at = time.monotonic() + self._lease_ttl + self._ready.append(record.transfer_id) + + def drain_ready(self) -> list[str]: + with self._memory.lock: + ready, self._ready = self._ready, [] + return ready + def begin_shutdown(self) -> None: """Stop new reservations and cancel everything without a remote writer.""" with self._condition: @@ -240,7 +286,7 @@ def begin_shutdown(self) -> None: self._shutting_down = True for transfer_id in list(self._active_ids): record = self._records[transfer_id] - if record.state is ConsumerReservationState.READY: + if record.writer_id or record.state is ConsumerReservationState.READY: self._terminate(record, ConsumerReservationState.CANCELLED) elif record.state is ConsumerReservationState.WRITING: self._defer(record, ConsumerReservationState.CANCEL_PENDING) @@ -300,6 +346,9 @@ def cancel( self._condition.notify_all() self._reap_tombstones(time.monotonic()) return True + if record.writer_id: + self._terminate(record, ConsumerReservationState.CANCELLED) + return True if record.state in _DEFERRED_STATES and not abandon: return True if record.state is ConsumerReservationState.WRITING and not abandon: @@ -335,12 +384,11 @@ def expire(self) -> int: with self._memory.lock: return self._expire_locked(time.monotonic()) - def retire_stale(self, encoder_cache: dict[str, torch.Tensor]) -> None: + def retire_stale( + self, encoder_cache: dict[str, torch.Tensor], freed: list[str] | None = None + ) -> None: with self._memory.lock: - reserved_hashes = { - self._records[transfer_id].mm_hash for transfer_id in self._active_ids - } - self._memory.retire_stale(encoder_cache, reserved_hashes) + self._memory.retire_stale(encoder_cache, freed=freed) def _expire_locked(self, now: float) -> int: expired = 0 @@ -352,40 +400,19 @@ def _expire_locked(self, now: float) -> int: self._terminate(record, ConsumerReservationState.EXPIRED) expired += 1 elif record.state is ConsumerReservationState.WRITING: - self._defer(record, ConsumerReservationState.EXPIRE_PENDING) - elif record.state in _DEFERRED_STATES: - # The writer that this release deferred to never reported - # back: it crashed, or its host is gone. Waiting forever - # pins the destination for the life of the process, so take - # the buffer back once no in-flight write could still land. - logger.warning( - "Reclaiming EC destination for transfer_id=%s after %.0fs " - "in %s: its writer never reported back", - record.transfer_id, - self._writer_grace, - record.state.name, - ) - self._terminate( - record, - ConsumerReservationState.CANCELLED - if record.state is ConsumerReservationState.CANCEL_PENDING - else ConsumerReservationState.EXPIRED, - ) - expired += 1 + if record.writer_id: + self._terminate(record, ConsumerReservationState.CANCELLED) + else: + self._defer(record, ConsumerReservationState.EXPIRE_PENDING) self._reap_tombstones(now) return expired def _defer( self, record: ConsumerReservation, state: ConsumerReservationState ) -> None: - """Hand a release to the remote writer, but not indefinitely. - - Mooncake may still be writing into the destination, so the release - waits for the writer to report. `_expire_locked` reclaims the buffer - once the grace has passed without a report. - """ + """Keep the destination until the writer completes or abandons it.""" self._transition(record, state) - record.expires_at = time.monotonic() + self._writer_grace + record.expires_at = float("inf") def _terminate( self, record: ConsumerReservation, state: ConsumerReservationState @@ -395,6 +422,17 @@ def _terminate( self._set_tombstone_deadline(record) def _release(self, record: ConsumerReservation) -> None: + if record.writer_id: + followers = self._followers.get(record.writer_id, {}) + followers.pop(record.transfer_id, None) + record.allocation = None + return + if self._writers.get(record.mm_hash) is record: + self._writers.pop(record.mm_hash) + for transfer_id in list(self._followers.pop(record.transfer_id, {})): + self._terminate( + self._records[transfer_id], ConsumerReservationState.CANCELLED + ) allocation = record.allocation if allocation is None: return diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py index 75e628711b1a..0b52dbe9c331 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/scheduler.py @@ -107,6 +107,7 @@ def __init__(self, vllm_config: VllmConfig) -> None: # free notifications the Scheduler already sends. Tracking it here # keeps `ensure_cache_available` on the upstream two-argument shape. self._local_cache: set[str] = set() + self._failed_saves: set[str] = set() def _cancel_remote(self, consumer_zmq: str, transfer_id: str) -> bool: pending = None @@ -175,7 +176,8 @@ def _note_awaiting_push( ) def take_unavailable_requests(self) -> set[str]: - return self._transfers.take_unavailable_requests() + failed, self._failed_saves = self._failed_saves, set() + return failed | self._transfers.take_unavailable_requests() def _poll_pending_cancels(self) -> None: pending = {} @@ -267,9 +269,11 @@ def _drain_push_notifications(self) -> None: self._expire_transfers() events = self._event_inbox.drain(self._control_addr) for data in events: - if not data.get("ready"): - continue try: + if not isinstance(data, dict): + raise TypeError("EC readiness event must be an object") + if not data.get("ready"): + continue self._accept_ready_event(data) except (KeyError, TypeError, ValueError): # Readiness events cross a plain PULL socket, so a malformed @@ -467,7 +471,9 @@ def build_connector_meta( self._transfers.release_ready(mm_hash, time.monotonic()) for transfer_id in self._transfers.drain_orphaned(): self._queue_cancel(transfer_id) - meta = ECMooncakeConnectorMetadata() + meta = ECMooncakeConnectorMetadata( + freed=scheduler_output.free_encoder_mm_hashes + ) for push_spec in self._pushes_to_prepare.values(): meta.pushes.append(push_spec) self._pushes_to_prepare.clear() @@ -489,6 +495,7 @@ def update_connector_output(self, connector_output: ECConnectorOutput) -> None: for mm_hash in meta.reclaimed: self._transfers.reclaim(mm_hash, time.monotonic()) self._scheduler_pending_work = meta.pending_saves + self._failed_saves.update(meta.failed_saves) def has_pending_push_work(self) -> bool: return self._scheduler_pending_work diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py index 1234cbccb339..a9991b20ac77 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/state.py @@ -68,6 +68,10 @@ def __init__(self, resident_capacity: int, tombstone_ttl: float) -> None: self._resident_capacity = resident_capacity self._tombstone_ttl = tombstone_ttl self._records: OrderedDict[str, SchedulerTransfer] = OrderedDict() + self._active_ids: dict[str, None] = {} + self._terminal_ids: OrderedDict[str, None] = OrderedDict() + self._resident_ids: OrderedDict[str, None] = OrderedDict() + self._resident_bytes = 0 self._hash_index: dict[str, deque[str]] = {} self._loads_to_dispatch: OrderedDict[str, None] = OrderedDict() self._unavailable_requests: set[str] = set() @@ -287,6 +291,12 @@ def cancel( request_id: str = "", ) -> bool: record = self._records.get(transfer_id) + if ( + record is not None + and mm_hash + and not self._identity_matches(record, mm_hash) + ): + return False if record is None: record = SchedulerTransfer( transfer_id=transfer_id, @@ -312,7 +322,8 @@ def expire(self, now: float, terminal_limit: int) -> list[SchedulerTransfer]: if terminal_limit < 0: raise ValueError("terminal_limit must be non-negative") expired = [] - for record in list(self._records.values()): + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] if record.deadline is None or record.deadline > now: continue if record.state is SchedulerTransferState.AVAILABLE: @@ -331,15 +342,15 @@ def expire(self, now: float, terminal_limit: int) -> list[SchedulerTransfer]: if record.request_id: self._notify_unavailable(record, record.request_id) expired.append(record) - elif record.state in _TERMINAL_STATES: - self._remove(record.transfer_id) - terminal_ids = [ - record.transfer_id - for record in self._records.values() - if record.state in _TERMINAL_STATES - ] - excess = max(0, len(terminal_ids) - terminal_limit) - for transfer_id in terminal_ids[:excess]: + while self._terminal_ids: + transfer_id = next(iter(self._terminal_ids)) + record = self._records[transfer_id] + if ( + record.deadline is not None + and record.deadline > now + and len(self._terminal_ids) <= terminal_limit + ): + break self._remove(transfer_id) return expired @@ -355,6 +366,7 @@ def _notify_unavailable(self, record: SchedulerTransfer, request_id: str) -> Non def _insert(self, record: SchedulerTransfer) -> None: self._records[record.transfer_id] = record + self._active_ids[record.transfer_id] = None if record.mm_hash: self._hash_index.setdefault(record.mm_hash, deque()).append( record.transfer_id @@ -389,12 +401,23 @@ def _transition( state: SchedulerTransferState, now: float | None = None, ) -> None: + if record.state is SchedulerTransferState.RESIDENT: + self._resident_ids.pop(record.transfer_id, None) + if record.spec is not None: + self._resident_bytes -= record.spec.nbytes + if state is SchedulerTransferState.RESIDENT: + self._resident_ids[record.transfer_id] = None + if record.spec is not None: + self._resident_bytes += record.spec.nbytes if state in _TERMINAL_STATES: if now is None: raise ValueError("Terminal transition requires a timestamp") # Keep a bounded tombstone so late events and repeat request IDs # remain idempotent instead of reviving a finished transfer. record.deadline = now + self._tombstone_ttl + self._active_ids.pop(record.transfer_id, None) + self._terminal_ids[record.transfer_id] = None + self._terminal_ids.move_to_end(record.transfer_id) if state is SchedulerTransferState.EXPIRED and record.state in ( SchedulerTransferState.READY, SchedulerTransferState.RESIDENT, @@ -404,23 +427,18 @@ def _transition( self._records.move_to_end(record.transfer_id) def _evict_residents(self, now: float) -> None: - residents: list[SchedulerTransfer] = [] - resident_bytes = 0 - for record in self._records.values(): - if record.state is not SchedulerTransferState.RESIDENT: - continue - residents.append(record) - if record.spec is not None: - resident_bytes += record.spec.nbytes - for record in residents: - if resident_bytes <= self._resident_capacity: - break + while self._resident_bytes > self._resident_capacity and self._resident_ids: + record = self._records[next(iter(self._resident_ids))] self._transition(record, SchedulerTransferState.EXPIRED, now=now) - if record.spec is not None: - resident_bytes -= record.spec.nbytes def _remove(self, transfer_id: str) -> None: record = self._records.pop(transfer_id, None) + self._active_ids.pop(transfer_id, None) + self._terminal_ids.pop(transfer_id, None) + if transfer_id in self._resident_ids: + self._resident_ids.pop(transfer_id) + if record is not None and record.spec is not None: + self._resident_bytes -= record.spec.nbytes self._loads_to_dispatch.pop(transfer_id, None) if record is None or not record.mm_hash: return diff --git a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py index e6d41c8c9621..05ce16b181e6 100644 --- a/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py +++ b/vllm/distributed/ec_transfer/ec_connector/mooncake/worker.py @@ -124,10 +124,34 @@ def __init__(self, vllm_config: VllmConfig) -> None: ) self._shard_pool: ThreadPoolExecutor | None = None self._shard_pool_lock = threading.Lock() - self._producer_pushes = ProducerPushManager() + self._push_ready = threading.Event() + self._producer_pushes = ProducerPushManager(self._push_ready.set) + self._dispatch_stop = threading.Event() + self._dispatcher: threading.Thread | None = None + self._failed_saves: set[str] = set() + self._collecting_sources = False self._completed_loads: set[str] = set() self._failed_loads: set[str] = set() self._shutdown = False + if self.is_producer: + self._dispatcher = threading.Thread( + target=self._dispatch_pushes, name="ec-mooncake-ready", daemon=True + ) + self._dispatcher.start() + + def _dispatch_pushes(self) -> None: + pending_event = False + while not self._dispatch_stop.is_set(): + self._push_ready.wait(timeout=0.001 if pending_event else None) + self._push_ready.clear() + if self._dispatch_stop.is_set(): + break + if self._collecting_sources: + pending_event = False + continue + pending_event = self._producer_pushes.submit_batches( + self._io_executor, self._push_batch + ) def _resolve_consumer_rank(self) -> None: """Place this worker in the consumer receive topology.""" @@ -176,6 +200,7 @@ def start_services(self) -> None: self._reservations.expire, peer_ports=[base_port + rank for rank in range(self._tp_size)], device=consumer_pool.device, + drain_ready=self._reservations.drain_ready, ) try: self._control_server.start() @@ -375,38 +400,10 @@ def _retry_cancel_reservations( def _reserve_remote(self, spec: ECMooncakePushSpec) -> list[dict[str, Any]]: """Reserve a destination on every shard of the consumer.""" - shards = None - for _ in range(_CANCEL_ATTEMPTS): - shards = self._control_client.discover_shards(spec.consumer_zmq) - if shards is not None: - break - if shards is None: - raise RuntimeError( - f"Could not discover every EC consumer shard at {spec.consumer_zmq}" - ) - tasks: list[Callable[[], dict[str, Any]]] = [ - partial(self._reserve_one, addr, spec) for addr in shards - ] - try: - return self._run_fanout(tasks) - except _FanoutError as exc: - # Keep one cleanup entry per confirmed shard. A missing result - # means the reservation outcome is unknown, so transfer-level - # cancellation with an empty reservation ID is the only safe - # idempotent action for that exact address. - reservations = [] - for addr, result in zip(shards, exc.results): - if isinstance(result, dict): - result["addr"] = addr - reservations.append(result) - else: - reservations.append({"addr": addr, "reservation_id": ""}) - exc.results[:] = reservations - try: - self._retry_cancel_reservations(spec, reservations) - except _FanoutError as cleanup_error: - raise exc from cleanup_error - raise + result = self._reserve_remote_many([spec])[0] + if isinstance(result, Exception): + raise result + return result def _refresh_remote_reservations( self, @@ -454,11 +451,19 @@ def start_save_caches( encoder_cache: dict[str, torch.Tensor] | None = None, **kwargs: Any, ) -> None: + self._collecting_sources = True + new: dict[str, list[ProducerPushRecord]] = {} for spec in metadata.pushes: - self._producer_pushes.reserve( - spec, - partial(self._submit_reservation, spec), - ) + try: + record, created = self._producer_pushes.reserve(spec, Future) + if created: + new.setdefault(spec.consumer_zmq, []).append(record) + except ValueError: + logger.warning("Rejected conflicting EC push", exc_info=True) + self._failed_saves.add(spec.request_id) + for records in new.values(): + future = self._control_executor.submit(self._reserve_batch, records) + future.add_done_callback(partial(self._reservation_batch_done, records)) if not isinstance(encoder_cache, dict): return for mm_hash in dict.fromkeys(spec.mm_hash for spec in metadata.pushes): @@ -466,10 +471,97 @@ def start_save_caches( if tensor is not None: self._bind_push_source(tensor, mm_hash) - def _submit_reservation( - self, spec: ECMooncakePushSpec - ) -> Future[list[dict[str, Any]]]: - return self._control_executor.submit(self._reserve_remote, spec) + def _reserve_batch(self, records: list[ProducerPushRecord]) -> None: + """Resolve independent item futures with one reserve RPC per shard.""" + if len(records) == 1: + record = records[0] + try: + record.reservation_future.set_result(self._reserve_remote(record.spec)) + except Exception as exc: + record.reservation_future.set_exception(exc) + return + outcomes = self._reserve_remote_many([record.spec for record in records]) + self._producer_pushes.finish_reservations(list(zip(records, outcomes))) + + def _reserve_remote_many( + self, specs: list[ECMooncakePushSpec] + ) -> list[list[dict[str, Any]] | Exception]: + shards = None + for _ in range(_CANCEL_ATTEMPTS): + shards = self._control_client.discover_shards(specs[0].consumer_zmq) + if shards is not None: + break + if shards is None: + raise RuntimeError( + f"Could not discover every EC consumer shard at {specs[0].consumer_zmq}" + ) + + def reserve(addr: str) -> list[dict[str, Any]]: + if len(specs) == 1: + return [{"ok": True, "result": self._reserve_one(addr, specs[0])}] + response = self._control_client.request( + addr, + { + "op": "reserve_batch", + "items": [ + { + "transfer_id": spec.transfer_id, + "mm_hash": spec.mm_hash, + "nbytes": spec.nbytes, + "shape": list(spec.shape), + "dtype": spec.dtype, + } + for spec in specs + ], + }, + ) + items = response["items"] + if len(items) != len(specs): + raise RuntimeError("Malformed EC reservation batch response") + for item in items: + if item.get("ok"): + item["result"]["_received_at"] = time.monotonic() + return items + + fanout_error = None + results: list[Any] + try: + results = self._run_fanout([partial(reserve, addr) for addr in shards]) + except _FanoutError as exc: + results = exc.results + fanout_error = exc + outcomes: list[list[dict[str, Any]] | Exception] = [] + for index, spec in enumerate(specs): + reservations: list[dict[str, Any]] = [] + error = None + for addr, items in zip(shards, results): + item = items[index] if items is not None else {} + if item.get("ok"): + reservations.append({**item["result"], "addr": addr}) + else: + error = fanout_error or RuntimeError(item["error"]) + reservations.append({"addr": addr, "reservation_id": ""}) + if error is None: + outcomes.append(reservations) + else: + failure = _FanoutError(error, list(reservations)) + try: + self._retry_cancel_reservations(spec, reservations) + except Exception as cleanup_error: + failure.__cause__ = cleanup_error + logger.exception("Failed to release partial EC reservation batch") + outcomes.append(failure) + return outcomes + + @staticmethod + def _reservation_batch_done( + records: list[ProducerPushRecord], future: Future[None] + ) -> None: + error = future.exception() + if error is not None: + for record in records: + if not record.reservation_future.done(): + record.reservation_future.set_exception(error) def start_load_caches( self, @@ -486,7 +578,7 @@ def start_load_caches( raise RuntimeError( "ECMooncakeConnector requires CUDA for ec_buffer_device=cuda" ) - self._reservations.retire_stale(encoder_cache) + self._reservations.retire_stale(encoder_cache, metadata.freed) for spec in metadata.loads: if spec.mm_hash in encoder_cache: @@ -544,8 +636,6 @@ def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: ] source = push.source_tensor assert source is not None - if writable and push.source_event is not None: - push.source_event.synchronize() for shard in writable: if int(shard["nbytes"]) != source.nbytes: raise RuntimeError( @@ -570,11 +660,6 @@ def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: registered_sources: list[int] = [] if staged is not None: sources = staged.tensors - # The NIC reads outside the CUDA stream. - if sources and sources[0].device.type == "cuda": - torch.accelerator.current_stream( - sources[0].device - ).synchronize() else: sources = tensors registered_sources = self._transfer.acquire_sources(tensors) @@ -788,6 +873,8 @@ def build_connector_worker_meta(self) -> ECMooncakeWorkerMetadata | None: # never loads must not report at all rather than report nothing. return None + self._collecting_sources = False + self._push_ready.set() self._flush_pending_pushes() failures = self._producer_pushes.poll() for mm_hash, error in failures: @@ -802,7 +889,9 @@ def build_connector_worker_meta(self) -> ECMooncakeWorkerMetadata | None: failed_loads=self._failed_loads, reclaimed=reclaimed, pending_saves=self._producer_pushes.pending, + failed_saves=self._failed_saves, ) + self._failed_saves = set() self._completed_loads = set() self._failed_loads = set() return meta @@ -811,6 +900,11 @@ def close(self) -> None: if self._shutdown: return self._shutdown = True + dispatcher = getattr(self, "_dispatcher", None) + if dispatcher is not None: + self._dispatch_stop.set() + self._push_ready.set() + dispatcher.join() if self._control_server is not None: self._reservations.begin_shutdown() self._flush_pending_pushes() @@ -820,8 +914,11 @@ def close(self) -> None: self._io_executor, self._cancel_orphaned_reservation, ) - self._io_executor.shutdown(wait=True) self._control_executor.shutdown(wait=True) + self._producer_pushes.submit_batches( + self._io_executor, self._push_batch, wait=True + ) + self._io_executor.shutdown(wait=True) if self._shard_pool is not None: self._shard_pool.shutdown(wait=True) # Every producer-side thread that could hold a control socket is stopped. From 50eb506df53227ea3815d1b292b7235674006a84 Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Sun, 6 Sep 2026 03:41:54 +0000 Subject: [PATCH 29/30] [EPD] Replace Mooncake EC unit tests with TCP E2E CI Replace implementation-coupled Mooncake unit tests with single-image, multi-image, and duplicate-image smoke cases. Run concurrent repeated requests against a colocated baseline and wire the TCP-only suite into Buildkite. Validation: 6/6 local TCP E2E responses matched the baseline; 19 shared proxy tests passed. Validated pipeline rendering, source filters, shell failure handling, and CUDA wheel selection. Remote L4 CI has not run yet. Co-authored-by: OpenAI Codex Signed-off-by: Tianyu Guo --- .../test_areas/disaggregated_mooncake.yaml | 54 + tests/v1/ec_connector/integration/README.md | 42 + .../run_epd_mooncake_ec_full_pipeline.sh | 94 +- .../integration/test_epd_correctness.py | 111 +- .../unit/test_ec_mooncake_connector.py | 3881 ----------------- 5 files changed, 248 insertions(+), 3934 deletions(-) delete mode 100644 tests/v1/ec_connector/unit/test_ec_mooncake_connector.py diff --git a/.buildkite/test_areas/disaggregated_mooncake.yaml b/.buildkite/test_areas/disaggregated_mooncake.yaml index 0911507eb22c..8c061a95ad8f 100644 --- a/.buildkite/test_areas/disaggregated_mooncake.yaml +++ b/.buildkite/test_areas/disaggregated_mooncake.yaml @@ -19,3 +19,57 @@ steps: commands: - bash /vllm-workspace/.buildkite/scripts/install-kv-connectors.sh - bash v1/kv_connector/mooncake_integration/config_sweep_accuracy_test.sh + +- label: ":nvidia: (L4) Mooncake EC TCP E2E" + key: mooncake-ec-tcp-e2e-2-gpus + timeout_in_minutes: 30 + working_dir: "/vllm-workspace" + device: l4 + num_devices: 2 + source_file_dependencies: + - vllm/distributed/ec_transfer/ + - vllm/config/ec_transfer.py + - vllm/config/multimodal.py + - vllm/config/vllm.py + - vllm/multimodal/ + - vllm/v1/core/encoder_cache_manager.py + - vllm/v1/core/sched/ + - vllm/v1/engine/ + - vllm/v1/worker/ec_connector_model_runner_mixin.py + - vllm/v1/worker/gpu_model_runner.py + - vllm/v1/worker/gpu_worker.py + - vllm/v1/worker/gpu/ec_connector.py + - vllm/v1/worker/gpu/model_runner.py + - vllm/v1/worker/gpu/mm/ + - examples/disaggregated/disaggregated_encoder/ + - tests/v1/ec_connector/ + - requirements/kv_connectors.txt + env: + PYTHON_BIN: "/vllm-workspace/.venv/bin/python" + PYTHONUNBUFFERED: "1" + VLLM_HOST_IP: "127.0.0.1" + VLLM_USE_V2_MODEL_RUNNER: "1" + MOONCAKE_EC_PROTOCOL: "tcp" + USE_MM_PROMPTS: "1" + SKIP_BASELINE: "0" + CONCURRENCY: "3" + REPEAT: "2" + LOG_PATH: "/tmp/mooncake-ec-e2e" + BASELINE_FILE: "/tmp/mooncake-ec-e2e/baseline.json" + commands: + - uv venv --system-site-packages --python 3.12 .venv + - | + uv pip install --python .venv/bin/python "$(.venv/bin/python - <<'PY' + from pathlib import Path + import torch + + requirement = next( + line for line in Path("requirements/kv_connectors.txt").read_text().splitlines() + if line.startswith("mooncake-transfer-engine ") + ) + if torch.version.cuda.split(".")[0] == "13": + requirement = requirement.replace("mooncake-transfer-engine", "mooncake-transfer-engine-cuda13") + print(requirement) + PY + )" + - bash tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh diff --git a/tests/v1/ec_connector/integration/README.md b/tests/v1/ec_connector/integration/README.md index 1b3bac1dd3e2..bf2dcd76d632 100644 --- a/tests/v1/ec_connector/integration/README.md +++ b/tests/v1/ec_connector/integration/README.md @@ -58,6 +58,48 @@ EC_SHARED_STORAGE_PATH="/tmp/my_ec_cache" bash ./tests/v1/ec_connector/integrati ## How It Works +### Mooncake EC (1E + 1PD) + +```bash +PYTHON_BIN="$PWD/.venv/bin/python" \ + bash tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +``` + +Requires two GPUs and Mooncake TransferEngine. TCP is the default transport; +no RDMA-capable network hardware is required. +The script sets `MC_FORCE_TCP=1` in TCP mode so Mooncake cannot auto-select RDMA. +The baseline runs first on GPU 0, followed by E on GPU 0 and PD on GPU 1. +`MOONCAKE_EC_PROTOCOL=rdma` selects RDMA instead; host-specific transport +environment variables should be set by the caller. + +Three black-box cases check fixed short answers and compare with the baseline: + +- One image: read the STOP sign. +- Two different images (including a local file): identify flowers and birds. +- The same image twice: read both STOP signs. + +By default, all three requests run concurrently for two rounds, exercising +shared hashes across requests and reuse after completion. Every response is +compared, not just the final round. Set `CONCURRENCY` and `REPEAT` to override. +Prefix caching is disabled and CUDA graphs remain enabled. This is a small +correctness suite, not a performance benchmark or failure-injection suite. + +`LOG_PATH` and `BASELINE_FILE` select the log directory and reference output. +`SKIP_BASELINE=1` reuses a reference generated with the same model and test +configuration. `USE_MM_PROMPTS=0` only checks text routing, not EC transfer. + +Buildkite runs this script as `mooncake-ec-tcp-e2e-2-gpus` on two L4 GPUs. +The job is defined in `.buildkite/test_areas/disaggregated_mooncake.yaml` and +selected for changes to EC, multimodal processing, scheduler/model-runner +integration, the proxy, or these tests (subject to normal PR CI approval). +It reuses the CI image's Python packages through a system-site-packages venv +and installs the CUDA-compatible Mooncake wheel. It uses loopback networking +and TCP only, without RDMA devices or peer-memory setup. + +The job fails on startup errors, request errors, or answer mismatches. On +failure, the script prints the last 100 lines of each server/proxy log to +the CI job log; full files remain under `LOG_PATH` while the container exists. + ### Step 1: Baseline 1. Start single vLLM instance on GPU diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh index cfd1d347d469..7b8b24cbecbb 100755 --- a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -11,30 +11,33 @@ # # Env: # MODEL HF model id (default: Qwen/Qwen2.5-VL-3B-Instruct) -# GPU_SINGLE / GPU_E / GPU_PD GPU ids (defaults 0 / 1 / 2) +# GPU_SINGLE / GPU_E / GPU_PD GPU ids (defaults 0 / 0 / 1) # ENDPOINT_PORT, ENCODE_PORT, PREFILL_DECODE_PORT -# MOONCAKE_EC_PROTOCOL rdma | tcp (default rdma) +# MOONCAKE_EC_PROTOCOL tcp | rdma (default tcp) # USE_MM_PROMPTS 1 (default) or 0 for text-only quick sanity # TIMEOUT_SECONDS wait_for_server timeout (default 1200) # SKIP_BASELINE set to 1 to reuse existing BASELINE_FILE +# CONCURRENCY / REPEAT concurrent requests / rounds (defaults 3 / 2) +# MAX_MODEL_LEN context length (default 16384) set -euo pipefail -GIT_ROOT=$(git rev-parse --show-toplevel) +GIT_ROOT=$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd) cd "$GIT_ROOT" || exit 1 export PYTHONPATH="${GIT_ROOT}:${PYTHONPATH:-}" PYTHON_BIN="${PYTHON_BIN:-${GIT_ROOT}/.venv/bin/python}" MODEL="${MODEL:-Qwen/Qwen2.5-VL-3B-Instruct}" USE_MM_PROMPTS="${USE_MM_PROMPTS:-1}" -MM_FLAG="" +TEST_ARGS=(--concurrency "${CONCURRENCY:-3}" --repeat "${REPEAT:-2}") if [[ "$USE_MM_PROMPTS" == "1" ]]; then - MM_FLAG="--use_mm_prompts" + TEST_ARGS+=(--use_mm_prompts --mm_smoke_test) fi +MAX_MODEL_LEN="${MAX_MODEL_LEN:-16384}" GPU_SINGLE="${GPU_SINGLE:-0}" -GPU_E="${GPU_E:-1}" -GPU_PD="${GPU_PD:-2}" +GPU_E="${GPU_E:-0}" +GPU_PD="${GPU_PD:-1}" ENCODE_PORT="${ENCODE_PORT:-19534}" PREFILL_DECODE_PORT="${PREFILL_DECODE_PORT:-19537}" @@ -42,14 +45,15 @@ ENDPOINT_PORT="${ENDPOINT_PORT:-10002}" BASELINE_PORT="${BASELINE_PORT:-10003}" EC_MOONCAKE_RESERVATION_PORT="${EC_MOONCAKE_RESERVATION_PORT:-19019}" -MOONCAKE_EC_PROTOCOL="${MOONCAKE_EC_PROTOCOL:-rdma}" +MOONCAKE_EC_PROTOCOL="${MOONCAKE_EC_PROTOCOL:-tcp}" export EC_MOONCAKE_RESERVATION_PORT export MOONCAKE_EC_PROTOCOL -# Mooncake cannot register CUDA memory through the peer-memory path on hosts -# whose kernel lacks the OFED peer-memory API; every transfer then fails at -# setup with -202. Opting out selects the path that works there and is a no-op -# where GPUDirect is available. -export WITH_NVIDIA_PEERMEM="${WITH_NVIDIA_PEERMEM:-0}" +if [[ "$MOONCAKE_EC_PROTOCOL" == "tcp" ]]; then + # TransferEngine may otherwise auto-select RDMA on hosts with an HCA. + export MC_FORCE_TCP=1 +else + unset MC_FORCE_TCP +fi LOG_PATH="${LOG_PATH:-/tmp}" BASELINE_FILE="${BASELINE_FILE:-/tmp/vllm_epd_mooncake_baseline.txt}" @@ -66,7 +70,7 @@ print(json.dumps({ "ec_connector": "ECMooncakeConnector", "ec_role": "ec_producer", "ec_connector_extra_config": { - "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "rdma"), + "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), }, }, separators=(",", ":"))) PY @@ -80,7 +84,7 @@ print(json.dumps({ "ec_ip": os.environ.get("EC_MOONCAKE_RESERVATION_HOST", "127.0.0.1"), "ec_port": int(os.environ.get("EC_MOONCAKE_RESERVATION_PORT", "19019")), "ec_connector_extra_config": { - "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "rdma"), + "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), }, }, separators=(",", ":"))) PY @@ -88,10 +92,16 @@ PY wait_for_server() { local port=$1 - timeout "$TIMEOUT_SECONDS" bash -c " - until curl -fsS http://localhost:${port}/health >/dev/null 2>&1; do - sleep 2 - done" && return 0 || return 1 + local pid=$2 + local deadline=$((SECONDS + TIMEOUT_SECONDS)) + while ((SECONDS < deadline)); do + kill -0 "$pid" 2>/dev/null || return 1 + if curl --max-time 2 -fsS "http://localhost:${port}/health" >/dev/null 2>&1; then + return 0 + fi + sleep 2 + done + return 1 } cleanup_instances() { @@ -121,7 +131,20 @@ cleanup_instances() { PIDS=() } -trap cleanup_instances EXIT +finish() { + local status=$? + cleanup_instances + if ((status != 0)); then + for log in "${LOG_PATH}"/mooncake_epd_*.log; do + [[ -f "$log" ]] || continue + echo "=== $log (last 100 lines) ===" + tail -100 "$log" + done + fi + exit "$status" +} + +trap finish EXIT trap 'exit 130' INT trap 'exit 143' TERM @@ -136,12 +159,15 @@ run_baseline() { --port "$PORT" \ --gpu-memory-utilization 0.75 \ --max-num-seqs 32 \ + --max-model-len "$MAX_MODEL_LEN" \ + --no-enable-prefix-caching \ + --limit-mm-per-prompt '{"image":2,"video":0}' \ --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ >"${LOG_PATH}/mooncake_epd_baseline.log" 2>&1 & local BASELINE_PID=$! PIDS+=("$BASELINE_PID") echo "Waiting for baseline..." - wait_for_server "$PORT" || { echo "Baseline failed to start; tail log:"; tail -80 "${LOG_PATH}/mooncake_epd_baseline.log"; return 1; } + wait_for_server "$PORT" "$BASELINE_PID" || { echo "Baseline failed to start"; return 1; } curl -s "http://127.0.0.1:${PORT}/v1/models" | head -c 200 || true echo "" "$PYTHON_BIN" "${GIT_ROOT}/tests/v1/ec_connector/integration/test_epd_correctness.py" \ @@ -149,7 +175,7 @@ run_baseline() { --model_name "$MODEL" \ --mode baseline \ --baseline_file "$BASELINE_FILE" \ - $MM_FLAG + "${TEST_ARGS[@]}" cleanup_instances } @@ -168,12 +194,15 @@ run_epd_mooncake() { --mm-tensor-ipc torch_shm \ --enable-request-id-headers \ --no-enable-prefix-caching \ - --max-num-batched-tokens 114688 \ + --max-num-batched-tokens "$MAX_MODEL_LEN" \ --max-num-seqs 32 \ + --max-model-len "$MAX_MODEL_LEN" \ + --limit-mm-per-prompt '{"image":2,"video":0}' \ --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ --ec-transfer-config "$ENC_EC_JSON" \ >"${LOG_PATH}/mooncake_epd_encoder.log" 2>&1 & - PIDS+=("$!") + local ENCODER_PID=$! + PIDS+=("$ENCODER_PID") echo "Starting PD on GPU $GPU_PD port $PREFILL_DECODE_PORT" CUDA_VISIBLE_DEVICES="$GPU_PD" "${VLLM_SERVE[@]}" "$MODEL" \ @@ -183,15 +212,19 @@ run_epd_mooncake() { --enable-mm-embeds \ --enable-request-id-headers \ --max-num-seqs 32 \ + --max-model-len "$MAX_MODEL_LEN" \ + --no-enable-prefix-caching \ + --limit-mm-per-prompt '{"image":2,"video":0}' \ --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ --ec-transfer-config "$PD_EC_JSON" \ >"${LOG_PATH}/mooncake_epd_pd.log" 2>&1 & - PIDS+=("$!") + local PD_PID=$! + PIDS+=("$PD_PID") echo "Waiting for encoder..." - wait_for_server "$ENCODE_PORT" || { echo "Encoder log:"; tail -100 "${LOG_PATH}/mooncake_epd_encoder.log"; return 1; } + wait_for_server "$ENCODE_PORT" "$ENCODER_PID" || { echo "Encoder failed to start"; return 1; } echo "Waiting for PD..." - wait_for_server "$PREFILL_DECODE_PORT" || { echo "PD log:"; tail -100 "${LOG_PATH}/mooncake_epd_pd.log"; return 1; } + wait_for_server "$PREFILL_DECODE_PORT" "$PD_PID" || { echo "PD failed to start"; return 1; } echo "Starting EPD proxy on $ENDPOINT_PORT" "$PYTHON_BIN" "${GIT_ROOT}/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py" \ @@ -203,10 +236,11 @@ run_epd_mooncake() { --ec-consumer-zmq-addrs \ "tcp://localhost:$EC_MOONCAKE_RESERVATION_PORT" \ >"${LOG_PATH}/mooncake_epd_proxy.log" 2>&1 & - PIDS+=("$!") + local PROXY_PID=$! + PIDS+=("$PROXY_PID") echo "Waiting for proxy..." - wait_for_server "$ENDPOINT_PORT" || { echo "Proxy log:"; tail -80 "${LOG_PATH}/mooncake_epd_proxy.log"; return 1; } + wait_for_server "$ENDPOINT_PORT" "$PROXY_PID" || { echo "Proxy failed to start"; return 1; } curl -s "http://127.0.0.1:${ENDPOINT_PORT}/health" || true echo "" @@ -215,7 +249,7 @@ run_epd_mooncake() { --model_name "$MODEL" \ --mode disagg \ --baseline_file "$BASELINE_FILE" \ - $MM_FLAG + "${TEST_ARGS[@]}" cleanup_instances } diff --git a/tests/v1/ec_connector/integration/test_epd_correctness.py b/tests/v1/ec_connector/integration/test_epd_correctness.py index ece73efa42d3..d3ae0f1d8792 100644 --- a/tests/v1/ec_connector/integration/test_epd_correctness.py +++ b/tests/v1/ec_connector/integration/test_epd_correctness.py @@ -26,6 +26,7 @@ import json import os import time +from concurrent.futures import ThreadPoolExecutor import openai import requests @@ -138,17 +139,20 @@ def run_chat_completion( Returns: Generated text content """ - client = openai.OpenAI(api_key="EMPTY", base_url=base_url) - - completion = client.chat.completions.create( - model=model_name, - messages=messages, - max_tokens=max_tokens, - temperature=0.0, - seed=42, - ) + with openai.OpenAI( + api_key="EMPTY", base_url=base_url, timeout=120, max_retries=0 + ) as client: + completion = client.chat.completions.create( + model=model_name, + messages=messages, + max_tokens=max_tokens, + temperature=0.0, + seed=42, + ) - return completion.choices[0].message.content + content = completion.choices[0].message.content + assert content, "Expected a nonempty completion" + return content def main(): @@ -198,7 +202,19 @@ def main(): help="Skip the two-image multimodal prompt", ) + parser.add_argument( + "--mm_smoke_test", + action="store_true", + help="Use three short, fixed-answer image cases, including duplicate images", + ) + parser.add_argument("--concurrency", type=int, default=1) + parser.add_argument("--repeat", type=int, default=1) + args = parser.parse_args() + if args.concurrency < 1 or args.repeat < 1: + parser.error("--concurrency and --repeat must be positive") + if args.mm_smoke_test and (not args.use_mm_prompts or args.skip_two_image_prompt): + parser.error("--mm_smoke_test requires all multimodal prompts") print(f"Service URL: {args.service_url}") print(f"Model: {args.model_name}") @@ -235,27 +251,76 @@ def main(): test_prompts = SAMPLE_PROMPTS_TEXT print("Using text-only prompts for quick testing") + if args.mm_smoke_test: + image = SAMPLE_PROMPTS_MM[0]["messages"][0]["content"][0] + image_pair = SAMPLE_PROMPTS_MM[1]["messages"][0]["content"][:2] + cases = [ + ( + "Single image", + [image], + "What word is on the red road sign? Reply with only that word " + "in uppercase.", + "STOP", + ), + ( + "Two different images", + image_pair, + "Do these pictures contain both flowers and birds? " + "Reply with only YES or NO.", + "YES", + ), + ( + "Same image twice", + [image, image], + "Read the red road sign in each image, in image order. Reply " + "with only the two uppercase words separated by a comma and a space.", + "STOP, STOP", + ), + ] + test_prompts = [ + { + "description": description, + "expected": expected, + "messages": [ + { + "role": "user", + "content": [*images, {"type": "text", "text": question}], + } + ], + } + for description, images, question, expected in cases + ] + # Run completions service_url = f"{args.service_url}/v1" output_strs = {} - for i, prompt_data in enumerate(test_prompts): - print( - f"\nRunning prompt {i + 1}/{len(test_prompts)}: " - f"{prompt_data['description']}" - ) - - output_str = run_chat_completion( + def complete(prompt_data): + output = run_chat_completion( base_url=service_url, model_name=args.model_name, messages=prompt_data["messages"], - max_tokens=MAX_OUTPUT_LEN, + max_tokens=16 if args.mm_smoke_test else MAX_OUTPUT_LEN, ) - - # Use description as key for comparison - key = prompt_data["description"] - output_strs[key] = output_str - print(f"Output: {output_str}") + if args.mm_smoke_test: + output = output.strip() + assert output == prompt_data["expected"], ( + f"{prompt_data['description']}: expected {prompt_data['expected']!r}, " + f"got {output!r}" + ) + return output + + # Each round includes concurrent requests sharing image hashes; later rounds + # exercise reuse after previous requests have finished. + with ThreadPoolExecutor(max_workers=args.concurrency) as executor: + for repeat in range(args.repeat): + outputs = executor.map(complete, test_prompts) + for prompt_data, output_str in zip(test_prompts, outputs): + key = prompt_data["description"] + if args.repeat > 1: + key = f"{key} (round {repeat + 1})" + output_strs[key] = output_str + print(f"{key}: {output_str}") if args.mode in ("baseline", "baseline_pd"): # Baseline mode: Save outputs diff --git a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py b/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py deleted file mode 100644 index 568e4ac502d4..000000000000 --- a/tests/v1/ec_connector/unit/test_ec_mooncake_connector.py +++ /dev/null @@ -1,3881 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Behavioral contract for the refactored Mooncake encoder-cache connector. - -The suite covers public compatibility, configuration, control and data planes, -memory ownership, the three role-specific lifecycle managers, Scheduler -metadata, Worker orchestration, and failure or cancellation races. -""" - -from __future__ import annotations - -import copy -import ctypes -import gc -import socket -import threading -import time -import weakref -from collections import Counter -from concurrent.futures import Future, ThreadPoolExecutor -from contextlib import contextmanager -from dataclasses import FrozenInstanceError, replace -from multiprocessing.reduction import ForkingPickler -from types import SimpleNamespace -from typing import Any -from unittest.mock import MagicMock, Mock, call, patch - -import pytest -import torch -import zmq - -from vllm.config import ModelConfig, VllmConfig -from vllm.distributed.ec_transfer.ec_connector import mooncake_ec_connector -from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorRole -from vllm.distributed.ec_transfer.ec_connector.factory import ECConnectorFactory -from vllm.distributed.ec_transfer.ec_connector.mooncake import ( - control, - memory, - transfer, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.config import ( - _RESERVATION_TTL_SECONDS, - MooncakeECConfig, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.control import ( - ConsumerControlServer, - ControlClient, - EventInbox, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.memory import ( - ConsumerMemoryPool, - ContiguousAllocator, - ProducerMemoryPool, - ResidentPool, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.metadata import ( - ECMooncakeConnectorMetadata, - ECMooncakeLoadSpec, - ECMooncakePushSpec, - ECMooncakeWorkerMetadata, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.producer import ( - ProducerPushManager, - ProducerPushRecord, - ProducerPushState, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.reservation import ( - ConsumerReservationManager, - ConsumerReservationState, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.scheduler import ( - ECMooncakeScheduler, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.state import ( - SchedulerTransferState, - SchedulerTransferTable, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.transfer import ( - MooncakeTransfer, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake.worker import ( - ECMooncakeWorker, - _FanoutError, -) -from vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector import ( - ECMooncakeConnector, -) -from vllm.v1.core.sched.output import SchedulerOutput - -pytest_plugins = ("tests.v1.ec_connector.unit.test_ec_example_connector",) -pytestmark = pytest.mark.skip_global_cleanup - - -class CopyingFakeTransferEngine: - """Model Mooncake registration rules while copying bytes in-process.""" - - def __init__(self, *args, **kwargs): - self.registered: set[int] = set() - self.regions: dict[int, int] = {} - self.register_calls: list[list[int]] = [] - self.unregister_calls: list[int] = [] - self.batch_unregister_calls: list[list[int]] = [] - self.transfer_calls: list[list[int]] = [] - self.transfer_batches: list[tuple[str, list[int], list[int], list[int]]] = [] - self.initialize_calls: list[tuple[str, str, str, str]] = [] - - def initialize(self, local_hostname, metadata_server, protocol, device_name) -> int: - self.initialize_calls.append( - (local_hostname, metadata_server, protocol, device_name) - ) - return 0 - - def get_rpc_port(self) -> int: - return 12345 - - def batch_transfer_sync_write( - self, target_hostname, buffers, peer_buffer_addresses, lengths - ) -> int: - sources = [int(address) for address in buffers] - destinations = [int(address) for address in peer_buffer_addresses] - sizes = [int(length) for length in lengths] - self.transfer_calls.append(sizes) - self.transfer_batches.append( - (str(target_hostname), sources, destinations, sizes) - ) - for src, dst, nbytes in zip(sources, destinations, sizes): - ctypes.memmove(int(dst), int(src), int(nbytes)) - return 0 - - def batch_register_memory(self, buffer_addresses, capacities) -> int: - addresses = [int(addr) for addr in buffer_addresses] - lengths = [int(length) for length in capacities] - # A real Transfer Engine refuses overlapping memory regions, so model - # that here: registering a range that intersects a live one fails. - regions = dict(self.regions) - for address, length in zip(addresses, lengths): - for other, other_length in regions.items(): - if address < other + other_length and other < address + length: - return 1 - regions[address] = length - self.register_calls.append(addresses) - self.registered.update(addresses) - self.regions = regions - return 0 - - def unregister_memory(self, buffer_address) -> int: - address = int(buffer_address) - self.unregister_calls.append(address) - self.registered.discard(address) - self.regions.pop(address, None) - return 0 - - def batch_unregister_memory(self, buffer_addresses) -> int: - addresses = [int(addr) for addr in buffer_addresses] - self.batch_unregister_calls.append(addresses) - self.registered.difference_update(addresses) - for address in addresses: - self.regions.pop(address, None) - return 0 - - -def _find_free_port() -> int: - s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - s.bind(("127.0.0.1", 0)) - _, port = s.getsockname() - s.close() - return int(port) - - -def _wait_for_worker_io( - connector: ECMooncakeConnector, timeout: float = 5.0 -) -> ECMooncakeWorkerMetadata: - deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - meta = connector.build_connector_worker_meta() - assert isinstance(meta, ECMooncakeWorkerMetadata) - if not meta.pending_saves: - return meta - time.sleep(0.01) - raise TimeoutError("EC Mooncake worker I/O did not finish") - - -def _drain_until_subscribed(scheduler: Any, timeout: float = 5.0) -> None: - """Drain until the event-channel discovery lands. - - Subscribing runs on the control executor so an unreachable consumer costs - the Scheduler nothing, which means the first drain only starts it. - """ - deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - scheduler._drain_pending = True - scheduler._drain_push_notifications() - if scheduler._event_inbox._socket is not None: - return - time.sleep(0.005) - raise TimeoutError("EC Mooncake event channel was never subscribed") - - -def _bind_extra_config(config: VllmConfig) -> None: - config.ec_transfer_config.get_from_extra_config.side_effect = lambda key, default: ( - config.ec_transfer_config.ec_connector_extra_config.get(key, default) - ) - - -class TestECMooncakeControlPlane: - """Validate ZMQ client reuse, shard discovery, events, and server RPCs.""" - - def test_event_transport_failure_stops_draining(self): - inbox = EventInbox(Mock()) - socket = Mock() - socket.recv_json.side_effect = zmq.ZMQError(zmq.ENOTSOCK) - inbox._socket = socket - assert inbox.drain("unused") == [] - socket.recv_json.assert_called_once() - inbox.close() - assert inbox._socket is None - - def test_worker_get_ip_failure_does_not_construct_client( - self, mock_vllm_config_producer - ): - with ( - patch_ec_mooncake_deps(), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake.worker.get_ip", - side_effect=RuntimeError("no address"), - ), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "worker.ControlClient" - ) as client_cls, - pytest.raises(RuntimeError, match="no address"), - ): - ECMooncakeConnector(mock_vllm_config_producer, ECConnectorRole.WORKER) - - client_cls.assert_not_called() - - def test_client_reuses_socket_and_discards_failed_exchange(self): - context = MagicMock() - socket = context.socket.return_value - socket.recv_json.side_effect = [ - {"ok": True, "result": {"ports": [19019]}}, - {"ok": True}, - {"ok": False, "error": "reservation rejected"}, - RuntimeError("timeout"), - ] - - with patch.object(control.zmq, "Context", return_value=context): - client = ControlClient(17) - assert client.request("tcp://consumer:19019", {"op": "peers"}) == { - "ports": [19019] - } - assert client.request("tcp://consumer:19019", {"op": "event_port"}) is None - with pytest.raises(RuntimeError, match="reservation rejected"): - client.request( - "tcp://consumer:19019", - {"op": "status", "transfer_id": "transfer"}, - ) - with pytest.raises(RuntimeError, match="timeout"): - client.request("tcp://consumer:19019", {"op": "event_port"}) - client.close() - client.close() - - context.socket.assert_called_once_with(zmq.REQ) - assert socket.setsockopt.call_args_list == [ - call(zmq.IPV6, 1), - call(zmq.RCVTIMEO, 17), - call(zmq.SNDTIMEO, 17), - call(zmq.LINGER, 0), - ] - socket.connect.assert_called_once_with("tcp://consumer:19019") - socket.close.assert_called_once_with(linger=0) - context.destroy.assert_called_once_with(linger=0) - - def test_client_uses_one_socket_per_thread(self): - context = MagicMock() - sockets = [MagicMock(), MagicMock()] - for index, control_socket in enumerate(sockets): - control_socket.recv_json.return_value = {"ok": True, "result": index} - context.socket.side_effect = sockets - barrier = threading.Barrier(2) - results: list[int] = [] - - with patch.object(control.zmq, "Context", return_value=context): - client = ControlClient(20) - - def request() -> None: - barrier.wait() - results.append(client.request("tcp://consumer:19019", {"op": "peers"})) - - threads = [threading.Thread(target=request) for _ in range(2)] - for thread in threads: - thread.start() - for thread in threads: - thread.join() - client.close() - - assert sorted(results) == [0, 1] - assert context.socket.call_args_list == [call(zmq.REQ), call(zmq.REQ)] - for control_socket in sockets: - control_socket.connect.assert_called_once_with("tcp://consumer:19019") - - def test_topology_retries_transient_discovery_failures(self): - client = object.__new__(ControlClient) - client._topologies = {} - client.request = Mock() - client.request.side_effect = [ - {"ports": [19019, 19020]}, - RuntimeError("old consumer"), - {"ports": [19029, 19030]}, - ] - assert client.discover_shards("tcp://consumer:19019") == [ - "tcp://consumer:19019", - "tcp://consumer:19020", - ] - assert client.discover_shards("tcp://consumer:19019") == [ - "tcp://consumer:19019", - "tcp://consumer:19020", - ] - assert client.discover_shards("tcp://legacy:19019") is None - assert client.discover_shards("tcp://legacy:19019") == [ - "tcp://legacy:19029", - "tcp://legacy:19030", - ] - assert client.discover_shards("tcp://legacy:19019") == [ - "tcp://legacy:19029", - "tcp://legacy:19030", - ] - assert client.request.call_args_list == [ - call("tcp://consumer:19019", {"op": "peers"}), - call("tcp://legacy:19019", {"op": "peers"}), - call("tcp://legacy:19019", {"op": "peers"}), - ] - - def test_event_inbox_retries_until_every_shard_is_connected(self): - client = object.__new__(ControlClient) - client._topologies = {} - client.request = Mock() - client.request.side_effect = [ - RuntimeError("peers not ready"), - {"ports": [19019, 19020]}, - 20001, - RuntimeError("event port not ready"), - 20001, - 20002, - ] - context = MagicMock() - socket = context.socket.return_value - event = {"transfer_id": "transfer", "ready": True} - socket.recv_json.side_effect = [event, zmq.Again()] - - with patch.object(control.zmq, "Context", return_value=context) as create: - inbox = EventInbox(client) - assert inbox.drain("tcp://consumer:19019") == [] - assert inbox.shard_count == 1 - create.assert_not_called() - assert inbox.drain("tcp://consumer:19019") == [] - assert inbox.shard_count == 1 - create.assert_not_called() - assert inbox.drain("tcp://consumer:19019") == [event] - assert inbox.shard_count == 2 - inbox.close() - inbox.close() - - assert client.request.call_args_list == [ - call("tcp://consumer:19019", {"op": "peers"}), - call("tcp://consumer:19019", {"op": "peers"}), - call("tcp://consumer:19019", {"op": "event_port"}), - call("tcp://consumer:19020", {"op": "event_port"}), - call("tcp://consumer:19019", {"op": "event_port"}), - call("tcp://consumer:19020", {"op": "event_port"}), - ] - create.assert_called_once_with() - assert socket.connect.call_args_list == [ - call("tcp://consumer:20001"), - call("tcp://consumer:20002"), - ] - assert socket.recv_json.call_args_list == [ - call(flags=zmq.DONTWAIT), - call(flags=zmq.DONTWAIT), - ] - socket.close.assert_called_once_with(linger=0) - context.term.assert_called_once_with() - - def test_server_preserves_wire_shapes_and_closes_twice(self): - port = _find_free_port() - completed: list[tuple[str, str]] = [] - cancelled: list[tuple[str, str, bool, bool]] = [] - - def status(transfer_id: str): - return {"transfer_id": transfer_id, "ready": False} - - def complete(transfer_id: str, reservation_id: str): - completed.append((transfer_id, reservation_id)) - return True, True - - def cancel( - transfer_id: str, - reservation_id: str, - abandon: bool, - refresh: bool, - ): - cancelled.append((transfer_id, reservation_id, abandon, refresh)) - return True - - server = ConsumerControlServer( - "127.0.0.1", - port, - reserve=lambda request: {"nbytes": request["nbytes"], "ready": False}, - status=status, - complete=complete, - cancel=cancel, - reap=lambda: 0, - peer_ports=[port, port + 1], - ) - client = ControlClient(1000) - server.start() - try: - addr = f"tcp://127.0.0.1:{port}" - assert client.request(addr, {"op": "peers"}) == {"ports": [port, port + 1]} - assert isinstance(client.request(addr, {"op": "event_port"}), int) - assert client.request( - addr, {"op": "status", "transfer_id": "transfer"} - ) == {"transfer_id": "transfer", "ready": False} - assert client.request( - addr, - { - "op": "reserve", - "transfer_id": "transfer", - "mm_hash": "hash", - "nbytes": 16, - "shape": [4], - "dtype": "float32", - }, - ) == {"nbytes": 16, "ready": False} - assert client.request( - addr, - { - "op": "complete", - "transfer_id": "transfer", - "reservation_id": "r0", - }, - ) == {"completed": True, "became_ready": True} - assert client.request( - addr, - { - "op": "complete_batch", - "items": [{"transfer_id": "transfer", "reservation_id": "r0"}], - }, - ) == {"items": [{"completed": True, "became_ready": True}]} - assert client.request( - addr, - { - "op": "cancel", - "transfer_id": "transfer", - "reservation_id": "r0", - "abandon": True, - }, - ) == {"cancelled": True} - finally: - client.close() - client.close() - server.close() - server.close() - - assert completed == [("transfer", "r0"), ("transfer", "r0")] - assert cancelled == [("transfer", "r0", True, False)] - - def test_server_keeps_serving_after_an_undecodable_request(self): - """One bad frame must not take the shard's control channel down. - - The loop's `finally` closes both sockets, so an escaping exception ends - the thread silently and every later reserve against this shard surfaces - only as a control timeout. - """ - port = _find_free_port() - server = ConsumerControlServer( - "127.0.0.1", - port, - reserve=lambda request: {"nbytes": request["nbytes"], "ready": False}, - status=lambda transfer_id: None, - complete=lambda transfer_id, reservation_id: (True, True), - cancel=lambda transfer_id, reservation_id, abandon, refresh: True, - reap=lambda: 0, - peer_ports=[port], - ) - server.start() - addr = f"tcp://127.0.0.1:{port}" - context = zmq.Context() - raw = context.socket(zmq.REQ) - raw.setsockopt(zmq.RCVTIMEO, 2000) - raw.setsockopt(zmq.LINGER, 0) - raw.connect(addr) - client = ControlClient(2000) - try: - raw.send(b"{not json") - assert raw.recv_json()["ok"] is False - - # An unknown op raises inside the handler; both paths must leave - # the channel able to answer the next caller. - with pytest.raises(RuntimeError): - client.request(addr, {"op": "nonsense"}) - - assert client.request(addr, {"op": "peers"}) == {"ports": [port]} - finally: - raw.close(linger=0) - context.term() - client.close() - server.close() - - def test_an_unconfirmed_writer_keeps_its_destination_reserved(self): - """Only writer confirmation makes a timed-out destination reusable.""" - engine = MagicMock(spec=MooncakeTransfer) - engine.register_memory.return_value = 0 - engine.unregister_memory.return_value = True - pool = ConsumerMemoryPool(768, engine) - pool.prepare(torch.device("cpu")) - manager = ConsumerReservationManager(pool, 300.0, 16) - - def reserve(request: dict[str, Any]) -> dict[str, Any]: - """The Worker's own handler, minus the dtype plumbing.""" - manager.expire() - reservation, _ = manager.reserve( - str(request["transfer_id"]), - str(request["mm_hash"]), - int(request["nbytes"]), - tuple(int(value) for value in request["shape"]), - str(request["dtype"]), - torch.float32, - ) - if reservation is None: - raise RuntimeError("EC consumer buffer pool is full") - return {"ready": reservation.state is ConsumerReservationState.READY} - - port = _find_free_port() - server = ConsumerControlServer( - "127.0.0.1", - port, - reserve=reserve, - status=manager.status, - complete=manager.complete, - cancel=manager.cancel, - reap=manager.expire, - peer_ports=[port], - ) - server.start() - addr = f"tcp://127.0.0.1:{port}" - client = ControlClient(2000) - - def request_destination(transfer_id: str) -> None: - client.request( - addr, - { - "op": "reserve", - "transfer_id": transfer_id, - "mm_hash": f"hash-{transfer_id}", - "nbytes": 64, - "shape": [16], - "dtype": "float32", - }, - ) - - try: - abandoned = 0 - for index in range(64): - try: - request_destination(f"t{index}") - except RuntimeError as error: - assert "pool is full" in str(error) - break - # The writer dies here, so no completion ever arrives; the - # Scheduler times the transfer out and cancels it. - assert client.request( - addr, control.make_cancel_request(f"t{index}") - ) == {"cancelled": True} - abandoned += 1 - assert abandoned, "the pool should accept a writer before filling up" - - for record in manager._records.values(): - record.expires_at = 0 - with pytest.raises(RuntimeError, match="pool is full"): - request_destination("after-the-grace") - for transfer_id in list(manager._records): - client.request( - addr, control.make_cancel_request(transfer_id, abandon=True) - ) - request_destination("after-writer-confirmation") - finally: - client.close() - server.close() - - -@pytest.fixture -def mock_vllm_config_producer(): - config = Mock(spec=VllmConfig) - config.model_config = Mock(spec=ModelConfig) - config.model_config.dtype = torch.float16 - config.model_config.hf_config = None - config.model_config.get_inputs_embeds_size.return_value = 16 - config.parallel_config = Mock() - config.parallel_config.tensor_parallel_size = 1 - config.parallel_config.pipeline_parallel_size = 1 - config.parallel_config.data_parallel_size = 1 - config.parallel_config.data_parallel_index = 0 - config.ec_transfer_config = Mock() - config.ec_transfer_config.is_ec_producer = True - config.ec_transfer_config.is_ec_consumer = False - config.ec_transfer_config.ec_buffer_device = "cuda" - config.ec_transfer_config.ec_buffer_size = 1e9 - config.ec_transfer_config.ec_ip = "127.0.0.1" - config.ec_transfer_config.ec_port = 19019 - config.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(config) - return config - - -@pytest.fixture -def mock_vllm_config_consumer(): - config = Mock(spec=VllmConfig) - config.parallel_config = Mock() - config.parallel_config.tensor_parallel_size = 1 - config.parallel_config.pipeline_parallel_size = 1 - config.parallel_config.data_parallel_size = 1 - config.parallel_config.data_parallel_index = 0 - config.ec_transfer_config = Mock() - config.ec_transfer_config.is_ec_producer = False - config.ec_transfer_config.is_ec_consumer = True - config.ec_transfer_config.ec_buffer_device = "cuda" - config.ec_transfer_config.ec_buffer_size = 1e9 - config.ec_transfer_config.ec_ip = "127.0.0.1" - config.ec_transfer_config.ec_port = 19019 - config.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(config) - return config - - -@contextmanager -def patch_ec_mooncake_deps(): - with ( - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake.transfer.TransferEngine", - CopyingFakeTransferEngine, - ), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake." - "transfer._MOONCAKE_IMPORT_ERROR", - None, - ), - patch( - "vllm.distributed.ec_transfer.ec_connector.mooncake.worker.get_ip", - return_value="127.0.0.1", - ), - ): - yield - - -class TestMooncakeTransfer: - """Validate lazy engine setup and source registration ownership.""" - - def test_initializes_engine_once_on_first_use(self): - engine = CopyingFakeTransferEngine() - with patch.object( - transfer, "TransferEngine", return_value=engine - ) as engine_cls: - data_plane = MooncakeTransfer("host", "tcp") - assert engine_cls.call_count == 0 - - assert data_plane.local_session() == "host:12345" - data_plane.ensure_ready() - - engine_cls.assert_called_once_with() - assert engine.initialize_calls == [("host", "P2PHANDSHAKE", "tcp", "")] - data_plane.close() - - def test_source_registration_is_refcounted_and_failed_release_keeps_owner(self): - engine = CopyingFakeTransferEngine() - data_plane = MooncakeTransfer("host", "tcp") - source = torch.randn(4, 4) - source_ref = weakref.ref(source) - with patch.object(transfer, "TransferEngine", return_value=engine): - first = data_plane.acquire_sources([source]) - second = data_plane.acquire_sources([source]) - assert engine.register_calls == [[source.data_ptr()]] - - assert data_plane.release_sources(first) - assert engine.batch_unregister_calls == [] - with patch.object( - engine, "batch_unregister_memory", return_value=1 - ) as unregister: - assert not data_plane.release_sources(second) - unregister.assert_called_once_with(second) - - del source - gc.collect() - assert source_ref() is not None - with patch.object( - engine, "batch_unregister_memory", return_value=0 - ) as unregister: - data_plane.close() - data_plane.close() - unregister.assert_called_once_with(second) - gc.collect() - assert source_ref() is None - - def test_failed_destination_unregister_is_retried_on_close(self): - engine = CopyingFakeTransferEngine() - data_plane = MooncakeTransfer("host", "tcp") - destination = torch.zeros(4, 4) - destination_ref = weakref.ref(destination) - with patch.object(transfer, "TransferEngine", return_value=engine): - assert data_plane.register_memory(destination) == 0 - with patch.object(engine, "unregister_memory", return_value=2): - assert not data_plane.unregister_memory(destination) - del destination - gc.collect() - assert destination_ref() is not None - - with patch.object( - engine, "batch_unregister_memory", return_value=0 - ) as unregister: - data_plane.close() - unregister.assert_called_once() - gc.collect() - assert destination_ref() is None - - def test_write_preserves_segments_and_reports_terminal_failure(self): - engine = CopyingFakeTransferEngine() - data_plane = MooncakeTransfer("host", "tcp") - sources = [torch.tensor([1, 2]), torch.tensor([3, 4])] - destinations = [torch.zeros_like(source) for source in sources] - source_addresses = [source.data_ptr() for source in sources] - destination_addresses = [tensor.data_ptr() for tensor in destinations] - lengths = [source.nbytes for source in sources] - with patch.object(transfer, "TransferEngine", return_value=engine): - data_plane.write("peer:1", source_addresses, destination_addresses, lengths) - assert engine.transfer_batches == [ - ("peer:1", source_addresses, destination_addresses, lengths) - ] - assert all( - torch.equal(source, destination) - for source, destination in zip(sources, destinations) - ) - - with ( - patch.object(engine, "batch_transfer_sync_write", return_value=9), - pytest.raises(RuntimeError, match="peer:2 failed with status 9"), - ): - data_plane.write( - "peer:2", source_addresses, destination_addresses, lengths - ) - data_plane.close() - - def test_write_returns_only_after_sync_engine_call_finishes(self): - engine = CopyingFakeTransferEngine() - data_plane = MooncakeTransfer("host", "tcp") - source = torch.ones(1, dtype=torch.uint8) - entered = threading.Event() - finish = threading.Event() - - def blocking_write(*args): - entered.set() - assert finish.wait(timeout=2) - return 0 - - with ( - patch.object(transfer, "TransferEngine", return_value=engine), - patch.object( - engine, "batch_transfer_sync_write", side_effect=blocking_write - ), - ): - addresses = data_plane.acquire_sources([source]) - completed = threading.Event() - - def write(): - try: - data_plane.write("peer:1", addresses, [2], [source.nbytes]) - finally: - data_plane.release_sources(addresses) - completed.set() - - thread = threading.Thread(target=write) - thread.start() - assert entered.wait(timeout=2) - assert not completed.is_set() - assert engine.batch_unregister_calls == [] - finish.set() - thread.join(timeout=2) - assert completed.is_set() - assert engine.batch_unregister_calls == [addresses] - data_plane.close() - - -class TestECMooncakeFactory: - """Validate factory registration.""" - - def test_factory_registers_connector(self): - cls = ECConnectorFactory.get_connector_class( - Mock(ec_connector="ECMooncakeConnector") - ) - assert cls is ECMooncakeConnector - assert ( - cls.__module__ - == "vllm.distributed.ec_transfer.ec_connector.mooncake_ec_connector" - ) - - -class TestContiguousAllocator: - """Validate aligned allocation, reuse, and range coalescing.""" - - def test_reuses_and_coalesces_contiguous_regions(self): - allocator = ContiguousAllocator(1024, alignment=256) - - first = allocator.allocate(1) - second = allocator.allocate(300) - assert first == (0, 256) - assert second == (256, 512) - assert allocator.allocate(300) is None - - allocator.free(*first) - allocator.free(*second) - assert allocator.allocate(1024) == (0, 1024) - - def test_splits_until_exhausted(self): - allocator = ContiguousAllocator(768, alignment=256) - - assert allocator.allocate(257) == (0, 512) - assert allocator.allocate(1) == (512, 256) - assert allocator.allocate(1) is None - - -class TestResidentPool: - """Validate resident pin, lease, replacement, and LRU semantics.""" - - def test_lru_skips_rejected_entry_and_replaces_without_losing_owner(self): - pool = ResidentPool[str]() - pool.insert("oldest", "first") - pool.insert("next", "second") - pool.retire("oldest") - pool.retire("next") - - evicted = pool.evict_lru(lambda key, _: key != "oldest") - - assert evicted == "next" - assert pool.get("oldest") == "first" - assert pool.insert("oldest", "replacement") == "first" - assert pool.get("oldest") == "replacement" - - def test_displaced_entry_waits_for_every_lease(self): - pool = ResidentPool[str]() - pool.insert("hash", "original") - first = pool.acquire("hash") - second = pool.acquire("hash") - assert first is not None and second is not None - - assert pool.insert("hash", "replacement") is None - assert pool.release(first) is None - assert pool.release(second) == "original" - assert pool.release(second) is None - - -class TestMooncakeMemoryPools: - """Validate Producer staging and Consumer residency ownership.""" - - class _Event: - """Minimal CUDA-event substitute controlling deferred frees.""" - - def __init__(self, complete: bool): - self.complete = complete - - def record(self, stream): - pass - - def query(self): - return self.complete - - def test_consumer_replacement_waits_for_cached_owner(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 0 - mooncake_transfer.unregister_memory.return_value = True - pool = ConsumerMemoryPool(768, mooncake_transfer) - pool.prepare(torch.device("cpu")) - first = pool.try_allocate(64, (16,), torch.float32) - replacement = pool.try_allocate(64, (16,), torch.float32) - assert first is not None and replacement is not None - pool.publish("hash", first) - held = pool.acquire_cached("hash", (16,), torch.float32) - assert held is not None - pool.publish("hash", replacement) - - third = pool.try_allocate(64, (16,), torch.float32) - assert third is not None - assert third.offset != first.offset - - pool.release_cached(held) - reused = pool.try_allocate(64, (16,), torch.float32) - assert reused is not None - assert reused.offset == first.offset - assert pool.take_resident("hash", (16,), "float32") is replacement.tensor - - def test_cached_consume_returns_newer_canonical_allocation(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 0 - pool = ConsumerMemoryPool(768, mooncake_transfer) - pool.prepare(torch.device("cpu")) - first = pool.try_allocate(64, (16,), torch.float32) - replacement = pool.try_allocate(64, (16,), torch.float32) - assert first is not None and replacement is not None - pool.publish("hash", first) - held = pool.acquire_cached("hash", (16,), torch.float32) - assert held is not None - pool.publish("hash", replacement) - - canonical = pool.publish("hash", held.value, held) - - assert canonical is replacement - reused = pool.try_allocate(64, (16,), torch.float32) - assert reused is not None - assert reused.offset == first.offset - - def test_consumer_defers_retired_reuse_until_event_completes(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 0 - pool = ConsumerMemoryPool(256, mooncake_transfer) - pool.prepare(torch.device("cpu")) - allocation = pool.try_allocate(64, (16,), torch.float32) - assert allocation is not None - pool.publish("hash", allocation) - event = self._Event(complete=False) - - with patch.object(pool, "_record_release_event", return_value=event): - pool.retire_stale({}, set()) - assert pool.reclaim_and_allocate(64, (16,), torch.float32) is None - event.complete = True - reused = pool.try_allocate(64, (16,), torch.float32) - - assert reused is not None - assert reused.offset == allocation.offset - assert pool.drain_reclaimed() == {"hash"} - - def test_consumer_registration_failure_disables_pool(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 1 - pool = ConsumerMemoryPool(256, mooncake_transfer) - - pool.prepare(torch.device("cpu")) - pool.prepare(torch.device("cpu")) - - assert pool.tensor is None - mooncake_transfer.register_memory.assert_called_once() - - def test_consumer_close_unregisters_once_and_releases_parent(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 0 - mooncake_transfer.unregister_memory.return_value = True - pool = ConsumerMemoryPool(256, mooncake_transfer) - pool.prepare(torch.device("cpu")) - parent = pool.tensor - - pool.close() - pool.close() - - mooncake_transfer.unregister_memory.assert_called_once_with(parent) - assert pool.tensor is None - - def test_producer_reuses_staging_and_unregisters_parent_on_close(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 0 - mooncake_transfer.unregister_memory.return_value = True - pool = ProducerMemoryPool(256, mooncake_transfer) - source = torch.arange(16, dtype=torch.float32) - - first = pool.stage([source]) - assert first is not None - assert torch.equal(first.tensors[0], source) - pool.release(first) - second = pool.stage([source]) - assert second is not None - assert second.regions == first.regions - pool.release(second) - parent = pool.tensor - - pool.close() - pool.close() - - assert pool.tensor is None - mooncake_transfer.unregister_memory.assert_called_once_with(parent) - - @pytest.mark.parametrize("pool_type", [ProducerMemoryPool, ConsumerMemoryPool]) - def test_failed_unregister_keeps_the_registered_buffer_until_retry(self, pool_type): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 0 - mooncake_transfer.unregister_memory.side_effect = [False, True] - pool = pool_type(256, mooncake_transfer) - if isinstance(pool, ProducerMemoryPool): - staged = pool.stage([torch.ones(16)]) - assert staged is not None - pool.release(staged) - else: - pool.prepare(torch.device("cpu")) - parent = pool.tensor - assert parent is not None - - pool.close() - assert pool.tensor is parent - pool.close() - assert pool.tensor is None - pool.close() - assert mooncake_transfer.unregister_memory.call_args_list == [ - call(parent), - call(parent), - ] - - def test_producer_failed_batch_returns_reserved_regions(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - mooncake_transfer.register_memory.return_value = 0 - pool = ProducerMemoryPool(256, mooncake_transfer) - source = torch.arange(16, dtype=torch.float32) - - assert pool.stage([source, source]) is None - staged = pool.stage([source]) - assert staged is not None - assert torch.equal(staged.tensors[0], source) - mooncake_transfer.register_memory.assert_called_once() - pool.release(staged) - pool.close() - - def test_producer_falls_back_when_staging_pool_allocation_fails(self): - mooncake_transfer = MagicMock(spec=MooncakeTransfer) - pool = ProducerMemoryPool(256, mooncake_transfer) - - with patch.object(memory.torch, "empty", side_effect=torch.OutOfMemoryError): - assert pool.stage([torch.ones(16)]) is None - assert pool.stage([torch.ones(16)]) is None - - mooncake_transfer.register_memory.assert_not_called() - - -class TestMooncakeECConfig: - def test_defaults_are_a_frozen_snapshot(self, mock_vllm_config_producer): - config = MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) - - assert ( - config.protocol, - config.buffer_device, - config.control_timeout_ms, - config.push_wait_timeout_s, - config.pool_size, - ) == ("tcp", "cuda", 30_000, 60, 1_000_000_000) - with pytest.raises(FrozenInstanceError): - config.protocol = "rdma" # type: ignore[misc] - - def test_derives_rank_local_port_and_custom_resources( - self, mock_vllm_config_consumer - ): - source = mock_vllm_config_consumer - source.parallel_config.tensor_parallel_size = 2 - source.parallel_config.data_parallel_index = 1 - source.ec_transfer_config.ec_buffer_size = 2048 - source.ec_transfer_config.ec_port = 5000 - source.ec_transfer_config.ec_connector_extra_config.update( - { - "control_timeout_s": 1.5, - "push_wait_timeout_s": 2.5, - } - ) - - config = MooncakeECConfig.from_vllm_config(source) - - assert config.control_port == 5002 - assert config.control_addr == "tcp://127.0.0.1:5002" - assert ( - config.control_timeout_ms, - config.push_wait_timeout_s, - config.pool_size, - ) == (1500, 2.5, 2048) - - @pytest.mark.parametrize("key", ["control_timeout_s", "push_wait_timeout_s"]) - @pytest.mark.parametrize("value", [float("nan"), float("inf"), float("-inf"), True]) - def test_rejects_invalid_timeouts(self, mock_vllm_config_producer, key, value): - mock_vllm_config_producer.ec_transfer_config.ec_connector_extra_config[key] = ( - value - ) - - with pytest.raises(ValueError, match=key): - MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) - - @pytest.mark.parametrize( - "value", [1.5, True, float("nan"), float("inf"), float("-inf")] - ) - def test_rejects_invalid_registered_buffer(self, mock_vllm_config_producer, value): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_size = value - - with pytest.raises(ValueError, match="ec_buffer_size"): - MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) - - @pytest.mark.parametrize( - ("attribute", "message"), - [ - ("tensor_parallel_size", "tensor_parallel_size=1"), - ("pipeline_parallel_size", "pipeline parallelism"), - ("data_parallel_size", "data_parallel_size=1"), - ], - ) - def test_rejects_sharded_producer( - self, mock_vllm_config_producer, attribute, message - ): - setattr(mock_vllm_config_producer.parallel_config, attribute, 2) - with pytest.raises(ValueError, match=message): - MooncakeECConfig.from_vllm_config(mock_vllm_config_producer) - - @pytest.mark.parametrize("port", [0, 65536]) - def test_rejects_out_of_range_port(self, mock_vllm_config_consumer, port): - mock_vllm_config_consumer.ec_transfer_config.ec_port = port - with pytest.raises(ValueError, match="1..65535"): - MooncakeECConfig.from_vllm_config(mock_vllm_config_consumer) - - def test_uses_upstream_ip_and_port(self, mock_vllm_config_consumer): - mock_vllm_config_consumer.ec_transfer_config.ec_ip = "consumer" - mock_vllm_config_consumer.ec_transfer_config.ec_port = 19100 - - config = MooncakeECConfig.from_vllm_config(mock_vllm_config_consumer) - - assert config.control_addr == "tcp://consumer:19100" - - -class TestECMooncakeConnectorValidation: - @pytest.mark.parametrize( - "role", [ECConnectorRole.SCHEDULER, ECConnectorRole.WORKER] - ) - def test_requires_mooncake_dependency(self, mock_vllm_config_producer, role): - with ( - patch.object(transfer, "_MOONCAKE_IMPORT_ERROR", ImportError("missing")), - pytest.raises(ImportError, match="mooncake-transfer-engine"), - ): - ECMooncakeConnector(mock_vllm_config_producer, role) - - @pytest.mark.parametrize( - ("role", "active", "inactive"), - [ - (ECConnectorRole.SCHEDULER, "_scheduler", "_worker"), - (ECConnectorRole.WORKER, "_worker", "_scheduler"), - ], - ) - def test_constructs_one_delegate_and_closes_once( - self, mock_vllm_config_producer, role, active, inactive - ): - scheduler = Mock() - worker = Mock() - with ( - patch.object( - mooncake_ec_connector, - "ECMooncakeScheduler", - return_value=scheduler, - ), - patch.object( - mooncake_ec_connector, - "ECMooncakeWorker", - return_value=worker, - ), - ): - connector = ECMooncakeConnector(mock_vllm_config_producer, role) - assert getattr(connector, active) is not None - assert getattr(connector, inactive) is None - connector.shutdown() - connector.shutdown() - - if role == ECConnectorRole.SCHEDULER: - scheduler.close.assert_called_once_with() - worker.close.assert_not_called() - else: - worker.close.assert_called_once_with() - scheduler.close.assert_not_called() - - -class TestECMooncakeMetadata: - @pytest.mark.parametrize( - "value", - [ - ECMooncakeConnectorMetadata( - loads=[ - ECMooncakeLoadSpec( - mm_hash="load", - nbytes=8, - shape=(2, 4), - dtype="float16", - transfer_id="transfer", - local=True, - ) - ], - pushes=[ - ECMooncakePushSpec( - mm_hash="push", - nbytes=8, - shape=(2, 4), - dtype="float16", - consumer_zmq="tcp://127.0.0.1:1234", - transfer_id="transfer", - request_id="request", - ) - ], - ), - ECMooncakeWorkerMetadata( - loaded={"loaded"}, - failed_loads={"failed"}, - reclaimed={"reclaimed"}, - pending_saves=True, - ), - ], - ) - def test_pickle_round_trip(self, value): - assert ForkingPickler.loads(ForkingPickler.dumps(value)) == value - - -class TestECMooncakeWorkerMetadataAggregation: - """Validate cross-rank success intersection and failure union rules.""" - - def test_an_item_one_rank_missed_is_not_loaded(self): - """Each rank gathers from its own cache, so all of them must have it. - - Reporting it as loaded because one rank succeeded left the scheduler - marking the hash ready while another rank raised on the cache miss. - """ - rank0 = ECMooncakeWorkerMetadata(loaded={"a", "b"}) - rank1 = ECMooncakeWorkerMetadata(loaded={"a"}, failed_loads={"b"}) - - merged = rank0.aggregate(rank1) - - assert merged.loaded == {"a"} - assert merged.failed_loads == {"b"} - - def test_a_reclaim_on_any_rank_invalidates_residency(self): - """The scheduler mirrors one pool, so the weakest rank decides.""" - merged = ECMooncakeWorkerMetadata(loaded={"a"}).aggregate( - ECMooncakeWorkerMetadata(loaded={"a"}, reclaimed={"c"}) - ) - assert merged.reclaimed == {"c"} - - -class TestSchedulerTransferTable: - @staticmethod - def pushed_spec(transfer_id: str, mm_hash: str = "hash") -> ECMooncakeLoadSpec: - return ECMooncakeLoadSpec( - mm_hash=mm_hash, - nbytes=16, - shape=(4,), - dtype="float32", - transfer_id=transfer_id, - ) - - def test_load_completion_and_resident_reload(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record, accepted = table.observe_ready(self.pushed_spec("transfer"), 10) - - assert accepted - assert table.begin_load("hash", "transfer", "request") is record - assert table.take_loads_to_dispatch() == [record] - assert table.complete_load("hash") - table.release_ready("hash", 1) - assert record.state is SchedulerTransferState.RESIDENT - assert table.begin_load("hash") is record - assert record.spec is not None and record.spec.local - - def test_same_hash_index_preserves_order_and_identity(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - first, _ = table.observe_ready(self.pushed_spec("first"), 10) - second, _ = table.observe_ready(self.pushed_spec("second"), 10) - - assert table.records_for_hash("hash", tuple(SchedulerTransferState)) == [ - first, - second, - ] - assert table.begin_load("hash") is first - assert ( - table.first_for_hash("hash", (SchedulerTransferState.AVAILABLE,)) is second - ) - # A colliding transfer id is refused rather than raised: transfer ids - # arrive on the request, so the engine must not fail on one. - collided, accepted = table.observe_ready( - self.pushed_spec("first", "other-hash"), 10 - ) - assert collided is first and not accepted - assert first.mm_hash == "hash" - - def test_unavailable_notification_is_drained_once_and_rejects_late_ready(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record = table.wait_for_event("transfer", "request", "hash", 1) - table.mark_unavailable("transfer", 2) - - assert table.take_unavailable_requests() == {"request"} - assert table.take_unavailable_requests() == set() - _, accepted = table.observe_ready(self.pushed_spec("transfer"), 40) - assert not accepted - assert record.state is SchedulerTransferState.UNAVAILABLE - - def test_cancel_and_duplicate_completion_are_idempotent(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - table.wait_for_event("cancelled", "request", "hash", 10) - assert table.cancel("cancelled", 1) - assert not table.cancel("cancelled", 2) - _, accepted = table.observe_ready(self.pushed_spec("cancelled"), 40) - assert not accepted - - record, _ = table.observe_ready(self.pushed_spec("completed", "other"), 10) - table.begin_load("other", "completed") - assert table.complete_load("other") - assert table.complete_load("other") - assert record.state is SchedulerTransferState.READY - - def test_terminal_records_expire_and_are_bounded(self): - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record, _ = table.observe_ready(self.pushed_spec("failed"), 10) - table.begin_load("hash", "failed") - table.fail_load("hash", 20) - assert record.deadline == 50 - table.expire(51, terminal_limit=100) - assert table.get("failed") is None - - for transfer_id in ("first", "second", "third"): - table.cancel(transfer_id, 60) - table.expire(61, terminal_limit=1) - assert table.get("first") is None - assert table.get("second") is None - assert table.get("third") is not None - - def test_same_hash_keeps_only_latest_resident(self): - table = SchedulerTransferTable(resident_capacity=32, tombstone_ttl=30) - first, _ = table.observe_ready(self.pushed_spec("first"), 10) - table.begin_load("hash", "first") - table.complete_load("hash") - table.release_ready("hash", 20) - second, _ = table.observe_ready(self.pushed_spec("second"), 30) - table.begin_load("hash", "second") - table.complete_load("hash") - table.release_ready("hash", 40) - - assert first.state is SchedulerTransferState.EXPIRED - assert second.state is SchedulerTransferState.RESIDENT - assert table.drain_orphaned() == ["first"] - assert table.drain_orphaned() == [] - - def test_a_load_that_never_reports_back_fails_its_request(self): - """A dispatched load needs a deadline of its own. - - LOADING used to carry none, so a lost Worker report left the hash - loading for good and deferred every later request for it, silently - and without even a retriable failure. - """ - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - record, _ = table.observe_ready(self.pushed_spec("transfer"), 10) - - assert table.begin_load("hash", "transfer", "request", 20) is record - assert table.take_loads_to_dispatch() == [record] - - assert table.expire(15, 16) == [] - assert table.expire(21, 16) == [record] - assert record.state is SchedulerTransferState.UNAVAILABLE - assert table.take_unavailable_requests() == {"request"} - - def test_a_colliding_transfer_id_fails_only_the_request_that_named_it(self): - """Transfer ids arrive on the request, so a collision is reachable. - - Raising took the engine down with it; the transfer already under that - id has to survive and only the newcomer may fail. - """ - table = SchedulerTransferTable(resident_capacity=64, tombstone_ttl=30) - first = table.wait_for_event("shared", "req-a", "hash", 1) - - assert first is not None - assert table.wait_for_event("shared", "req-b", "other-hash", 1) is None - assert table.take_unavailable_requests() == {"req-b"} - assert first.mm_hash == "hash" - assert first.request_id == "req-a" - assert not table.cancel("shared", 2, "other-hash", "req-b") - assert first.state is SchedulerTransferState.WAITING_EVENT - - -class TestECMooncakeSchedulerMetadata: - """Validate Scheduler decisions and per-step Worker metadata.""" - - @pytest.mark.parametrize("payload", [[], "ready", 1, None]) - def test_non_object_events_are_discarded(self, payload): - scheduler = object.__new__(ECMooncakeScheduler) - scheduler._drain_pending = True - scheduler._control_addr = "unused" - scheduler._poll_pending_cancels = Mock() - scheduler._expire_transfers = Mock() - scheduler._event_inbox = Mock() - scheduler._event_inbox.drain.return_value = [payload] - scheduler._accept_ready_event = Mock() - scheduler._drain_push_notifications() - scheduler._accept_ready_event.assert_not_called() - - def test_an_unreachable_consumer_does_not_block_the_scheduler( - self, mock_vllm_config_consumer - ): - """Subscribing must never cost the Scheduler a control timeout. - - `has_cache_item` runs inside `schedule()`, and discovery is one - blocking request per shard. Doing it inline froze every request in the - engine for `control_timeout_s` on each drain while a shard was - unreachable, and `discover_shards` caches only successes, so the cost - repeated for as long as the shard stayed down. - """ - timeout_s = 0.4 - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - "control_timeout_s": timeout_s, - } - mock_vllm_config_consumer.ec_transfer_config.ec_port = _find_free_port() - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - elapsed = [] - for _ in range(3): - started = time.monotonic() - assert scheduler.has_cache_item("hash") is False - elapsed.append(time.monotonic() - started) - finally: - scheduler.shutdown() - - assert max(elapsed) < timeout_s, elapsed - - def test_a_duplicate_transfer_id_fails_the_request_not_the_engine( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """`ec_transfer_params` is a request field, so ids can collide. - - `ensure_cache_available` runs inside `schedule()`: raising there took - EngineCore down on input any client could send. - """ - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - first = mock_request_with_3_mm - first.mm_features = first.mm_features[:1] - first.request_id = "req-a" - first.ec_transfer_params = { - "ec_items": [ - { - "mm_hash": first.mm_features[0].identifier, - "transfer_id": "shared", - } - ] - } - second = copy.copy(first) - second.request_id = "req-b" - second.mm_features = [ - replace(first.mm_features[0], identifier="another_hash") - ] - second.ec_transfer_params = { - "ec_items": [{"mm_hash": "another_hash", "transfer_id": "shared"}] - } - - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert not scheduler.ensure_cache_available(first, 0) - assert not scheduler.ensure_cache_available(second, 0) - - assert scheduler.take_unavailable_requests() == {"req-b"} - record = scheduler._scheduler._transfers.get("shared") - assert record is not None - assert record.request_id == "req-a" - finally: - scheduler.shutdown() - - def test_cancel_confirms_topology_and_retries_only_failed_shards(self): - scheduler = object.__new__(ECMooncakeScheduler) - scheduler._control_client = Mock(spec=ControlClient) - scheduler._control_client.discover_shards.side_effect = [ - None, - ["shard-0", "shard-1", "shard-2"], - ] - called = [] - - def request(addr, _payload): - called.append(addr) - if addr == "shard-0" and called.count(addr) == 1: - raise RuntimeError("cancel shard failed") - return {"cancelled": True} - - scheduler._control_client.request.side_effect = request - assert scheduler._cancel_remote("base", "transfer") - assert scheduler._control_client.discover_shards.call_args_list == [ - call("base"), - call("base"), - ] - assert called == ["shard-0", "shard-1", "shard-2", "shard-0"] - - def test_cancel_rejects_unconfirmed_topology_without_sending(self): - scheduler = object.__new__(ECMooncakeScheduler) - scheduler._control_client = Mock(spec=ControlClient) - scheduler._control_client.discover_shards.return_value = None - - with pytest.raises(RuntimeError, match="discover every EC consumer shard"): - scheduler._cancel_remote("base", "transfer") - - assert scheduler._control_client.discover_shards.call_count == 2 - scheduler._control_client.request.assert_not_called() - - def test_item_with_no_transfer_in_flight_is_reported_as_stalled( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """A push that never arrives must not wait silently forever.""" - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - "push_wait_timeout_s": 0.001, - } - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - with ( - patch.object(scheduler._scheduler, "_drain_push_notifications"), - patch.object( - scheduler._scheduler._control_client, - "request", - return_value=None, - ) as control_request, - patch( - "vllm.distributed.ec_transfer.ec_connector." - "mooncake.scheduler.time.monotonic", - side_effect=[10, 10.002, 10.003], - ), - ): - assert not scheduler.ensure_cache_available(request, 0) - assert not scheduler.ensure_cache_available(request, 0) - record = scheduler._scheduler._transfers.get( - f"{request.request_id}:0" - ) - assert record is not None - assert record.state is SchedulerTransferState.UNAVAILABLE - assert scheduler.take_unavailable_requests() == {request.request_id} - assert not scheduler.ensure_cache_available(request, 0) - assert scheduler.take_unavailable_requests() == set() - control_request.assert_not_called() - finally: - scheduler.shutdown() - - def test_local_cache_hit_keeps_the_transfer( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """`local_cache_hashes` is a snapshot, so the transfer must survive it. - - Cancelling on a cache hit strands the request when the entry is - evicted before it is scheduled: the item is then unreachable. - """ - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - mm_hash = request.mm_features[0].identifier - request.ec_transfer_params = { - "ec_items": [ - {"mm_hash": mm_hash, "transfer_id": "request-transfer"} - ] - } - scheduler._scheduler._transfers.observe_ready( - ECMooncakeLoadSpec( - mm_hash=mm_hash, - nbytes=16, - shape=(4,), - dtype="float32", - transfer_id="request-transfer", - ), - 10, - ) - with ( - patch.object(scheduler._scheduler, "_drain_push_notifications"), - patch.object(scheduler._scheduler, "_queue_cancel") as cancel, - ): - # The local encoder cache is mirrored from the Scheduler's - # own alloc/free notifications, not passed in per call. - scheduler.update_state_after_alloc(request, 0) - assert scheduler.ensure_cache_available(request, 0) - cancel.assert_not_called() - record = scheduler._scheduler._transfers.get("request-transfer") - assert record is not None - assert record.state is SchedulerTransferState.AVAILABLE - - # Once the entry is evicted the request can still get it. - scheduler._scheduler._local_cache.discard(mm_hash) - assert not scheduler.ensure_cache_available(request, 0) - assert ( - scheduler._scheduler._transfers.first_for_hash( - mm_hash, (SchedulerTransferState.LOADING,) - ) - is not None - ) - finally: - scheduler.shutdown() - - def test_consumed_item_releases_its_transfer_immediately( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """The buffer goes back as soon as the item is consumed. - - Holding it until `request_finished` would pin a pool slot for the - whole generation, long after the embedding was used. - """ - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - mm_hash = request.mm_features[0].identifier - request.ec_transfer_params = { - "ec_items": [{"mm_hash": mm_hash, "transfer_id": "consumed-transfer"}] - } - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - scheduler._scheduler._transfers.observe_ready( - ECMooncakeLoadSpec( - mm_hash=mm_hash, - nbytes=16, - shape=(4,), - dtype="float32", - transfer_id="consumed-transfer", - ), - 10, - ) - with patch.object( - scheduler._scheduler, "_cancel_remote", return_value=True - ): - scheduler.update_state_after_free(request, 0) - record = scheduler._scheduler._transfers.get("consumed-transfer") - assert record is not None - assert record.state is SchedulerTransferState.CANCELLED - finally: - scheduler.shutdown() - - def test_cancelled_transfer_ignores_late_ready_events( - self, mock_vllm_config_consumer - ): - """Cancelled is terminal even when a ready event was already queued.""" - transfer_id = "cancelled-transfer" - ports = [19101, 19102, 19103, 19104] - event = { - "mm_hash": "hash", - "transfer_id": transfer_id, - "ready": True, - "reservation_id": "reservation", - "nbytes": 16, - "shape": [4], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - scheduler._scheduler._event_inbox.shard_count = len(ports) - scheduler._scheduler._transfers.cancel( - transfer_id, time.monotonic(), mm_hash="hash" - ) - scheduler._scheduler._event_inbox.drain = Mock( - return_value=[{**event, "shard": port} for port in ports] - ) - - scheduler._scheduler._drain_push_notifications() - - record = scheduler._scheduler._transfers.get(transfer_id) - assert record is not None - assert record.state is SchedulerTransferState.CANCELLED - assert transfer_id not in scheduler._scheduler._event_ready_shards - finally: - scheduler.shutdown() - - def test_readiness_needs_every_consumer_shard(self, mock_vllm_config_consumer): - """A sharded consumer is only ready once every rank reports. - - Each rank runs its own control channel and pushes its own readiness - notifications. Subscribing to the first rank alone strands the other - ranks' queues and lets a load be scheduled that the last rank cannot - serve, which only `aggregate`'s `loaded` intersection then catches. - """ - ports = [19101, 19102, 19103] - event_ports = {port: 19201 + index for index, port in enumerate(ports)} - - def fake_send(addr: str, request: dict): - port = int(addr.rsplit(":", 1)[1]) - if request["op"] == "peers": - return {"ports": ports} - if request["op"] == "event_port": - return event_ports[port] - return {} - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - scheduler._scheduler._control_addr = f"tcp://127.0.0.1:{ports[0]}" - with patch.object( - scheduler._scheduler._control_client, - "request", - side_effect=fake_send, - ) as send_control: - _drain_until_subscribed(scheduler._scheduler) - - subscribed = [ - call.args[0] - for call in send_control.call_args_list - if call.args[1]["op"] == "event_port" - ] - assert len(subscribed) == len(ports) - assert scheduler._scheduler._event_inbox.shard_count == len(ports) - - event = {"transfer_id": "transfer-0"} - assert not scheduler._scheduler._note_shard_ready( - {**event, "shard": ports[0]} - ) - # The same rank reporting twice is not two ranks. - assert not scheduler._scheduler._note_shard_ready( - {**event, "shard": ports[0]} - ) - assert not scheduler._scheduler._note_shard_ready( - {**event, "shard": ports[1]} - ) - assert scheduler._scheduler._note_shard_ready( - {**event, "shard": ports[2]} - ) - # Nothing is retained once the transfer is handed on. - assert "transfer-0" not in scheduler._scheduler._event_ready_shards - finally: - scheduler.shutdown() - - def test_evicted_item_is_reloaded_from_the_pool_without_a_transfer( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """The stall at high concurrency: more needs than transfers. - - Requests sharing an image get one transfer each, but a load consumes - one spec and serves everyone at once. After an eviction the remaining - requests need the item again with no spec left. The receive pool still - holds it, so the reload must come from there. - """ - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - first = mock_request_with_3_mm - first.mm_features = first.mm_features[:1] - mm_hash = first.mm_features[0].identifier - first.ec_transfer_params = { - "ec_items": [{"mm_hash": mm_hash, "transfer_id": "only-transfer"}] - } - second = copy.copy(first) - second.request_id = "second-request" - second.ec_transfer_params = None - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - scheduler._scheduler._event_inbox.drain = Mock( - return_value=[ - { - "mm_hash": mm_hash, - "transfer_id": "only-transfer", - "ready": True, - "reservation_id": "r0", - "nbytes": 16, - "shape": [4], - "dtype": "float32", - } - ] - ) - scheduler._scheduler._drain_push_notifications() - - assert not scheduler.ensure_cache_available(first, 0) - meta = scheduler.build_connector_meta( - SimpleNamespace(free_encoder_mm_hashes=[]) - ) - assert [spec.transfer_id for spec in meta.loads] == ["only-transfer"] - scheduler.update_connector_output( - SimpleNamespace( - ec_connector_worker_meta=ECMooncakeWorkerMetadata( - loaded={mm_hash} - ) - ) - ) - record = scheduler._scheduler._transfers.get("only-transfer") - assert record is not None - assert record.state is SchedulerTransferState.READY - - # The encoder cache evicts the entry. - scheduler.build_connector_meta( - SimpleNamespace(free_encoder_mm_hashes=[mm_hash]) - ) - - # The second request has no transfer of its own, and the only - # transfer is spent. It must still be served. - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert scheduler.has_cache_item(mm_hash) - assert not scheduler.ensure_cache_available(second, 0) - assert record.state is SchedulerTransferState.LOADING - reload = scheduler.build_connector_meta( - SimpleNamespace(free_encoder_mm_hashes=[]) - ) - assert [spec.local for spec in reload.loads] == [True] - finally: - scheduler.shutdown() - - def test_reclaimed_item_stops_being_offered_as_resident( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - """Residency is a mirror of the worker's pool, not a promise.""" - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - mm_hash = request.mm_features[0].identifier - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - try: - spec = ECMooncakeLoadSpec( - mm_hash=mm_hash, - nbytes=16, - shape=(4,), - dtype="float32", - transfer_id="transfer", - ) - table = scheduler._scheduler._transfers - table.observe_ready(spec, time.monotonic() + _RESERVATION_TTL_SECONDS) - table.begin_load(mm_hash, "transfer") - table.take_loads_to_dispatch() - table.complete_load(mm_hash) - table.release_ready(mm_hash, time.monotonic()) - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert scheduler.has_cache_item(mm_hash) - - scheduler.update_connector_output( - SimpleNamespace( - ec_connector_worker_meta=ECMooncakeWorkerMetadata( - reclaimed={mm_hash} - ) - ) - ) - with patch.object(scheduler._scheduler, "_drain_push_notifications"): - assert not scheduler.has_cache_item(mm_hash) - assert not table.has_state(mm_hash, (SchedulerTransferState.RESIDENT,)) - finally: - scheduler.shutdown() - - def test_retains_new_completion_while_same_hash_is_loading( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - event = { - "mm_hash": "hash", - "transfer_id": "next-transfer", - "ready": True, - "reservation_id": "next", - "nbytes": 64, - "shape": [2, 8], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - scheduler._scheduler._event_inbox.drain = Mock(return_value=[event]) - current = ECMooncakeLoadSpec( - mm_hash="hash", - nbytes=64, - shape=(2, 8), - dtype="float32", - transfer_id="current-transfer", - ) - scheduler._scheduler._transfers.observe_ready(current, time.monotonic() + 1) - scheduler._scheduler._transfers.begin_load("hash", "current-transfer") - - scheduler._scheduler._drain_push_notifications() - - pending = scheduler._scheduler._transfers.get("next-transfer") - assert pending is not None and pending.spec is not None - assert pending.state is SchedulerTransferState.AVAILABLE - assert pending.spec.transfer_id == "next-transfer" - - def test_build_connector_meta_clears_pending( - self, mock_vllm_config_consumer, mock_request_with_3_mm - ): - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.SCHEDULER - ) - mm_hash = mock_request_with_3_mm.mm_features[0].identifier - load_spec = ECMooncakeLoadSpec( - mm_hash=mm_hash, - nbytes=32, - shape=(2, 4), - dtype="float32", - transfer_id="transfer", - ) - scheduler._scheduler._transfers.observe_ready( - load_spec, time.monotonic() + _RESERVATION_TTL_SECONDS - ) - scheduler._scheduler._transfers.begin_load(mm_hash, "transfer") - meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) - ) - assert isinstance(meta, ECMooncakeConnectorMetadata) - assert len(meta.loads) == 1 - assert meta.loads[0].mm_hash == mm_hash - assert scheduler._scheduler._transfers.take_loads_to_dispatch() == [] - record = scheduler._scheduler._transfers.get("transfer") - assert record is not None - assert record.state is SchedulerTransferState.LOADING - - def test_producer_builds_push_metadata_after_preprocessing( - self, mock_vllm_config_producer, mock_request_with_3_mm - ): - request = mock_request_with_3_mm - request.mm_features = request.mm_features[:1] - request.ec_transfer_params = { - "consumer_zmq": "tcp://decode:19019", - "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], - } - mock_vllm_config_producer.model_config.dtype = torch.float32 - mock_vllm_config_producer.model_config.hf_config = None - mock_vllm_config_producer.model_config.get_inputs_embeds_size.return_value = 16 - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - scheduler.update_state_after_alloc(request, 0) - meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) - ) - - # The same request remains visible on a later scheduler step, but - # its worker push metadata must not be emitted a second time. - assert scheduler.ensure_cache_available(request, 0) - next_meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) - ) - - scheduler.request_finished(request) - - assert meta.loads == [] - assert meta.pushes == [ - ECMooncakePushSpec( - mm_hash="img_hash_1", - nbytes=100 * 16 * 4, - shape=(100, 16), - dtype="float32", - consumer_zmq="tcp://decode:19019", - transfer_id="transfer-1", - request_id="test_req_123", - ) - ] - assert next_meta.pushes == [] - assert "transfer-1" not in scheduler._scheduler._prepared_push_transfer_ids - - def test_producer_uses_deepstack_encoder_cache_width( - self, mock_vllm_config_producer, mock_request_with_3_mm - ): - request = mock_request_with_3_mm - request.ec_transfer_params = { - "consumer_zmq": "tcp://decode:19019", - "ec_items": [{"mm_hash": "img_hash_1", "transfer_id": "transfer-1"}], - } - mock_vllm_config_producer.model_config.dtype = torch.bfloat16 - mock_vllm_config_producer.model_config.hf_config = SimpleNamespace( - vision_config=SimpleNamespace( - out_hidden_size=2560, - deepstack_visual_indexes=[5, 11, 17], - ) - ) - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - scheduler.update_state_after_alloc(request, 0) - meta = scheduler.build_connector_meta( - Mock(spec=SchedulerOutput, free_encoder_mm_hashes=[]) - ) - - num_tokens = request.get_num_encoder_embeds(0) - spec = meta.pushes[0] - assert spec.shape == (num_tokens, 10240) - assert spec.nbytes == num_tokens * 10240 * torch.bfloat16.itemsize - - def test_producer_reports_proxy_rewrite_metadata(self, mock_vllm_config_producer): - feature = SimpleNamespace( - identifier="image_uuid", - modality="image", - data=SimpleNamespace( - get_data=lambda: { - "image_grid_thw": torch.tensor([1, 32, 48]), - "pixel_values": torch.ones(2), - } - ), - ) - request = SimpleNamespace(mm_features=[feature]) - - with patch_ec_mooncake_deps(): - scheduler = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.SCHEDULER - ) - with patch.object( - scheduler._scheduler._metadata_resolver, - "fields_for", - return_value={"image_grid_thw"}, - ): - delay_free, params = scheduler.request_finished(request) - - assert not delay_free - # Upstream shape: keyed by the engine's own identifier, with the - # placeholder metadata nested so a connector can report its own - # transfer coordinates alongside it. - assert params == {"image_uuid": {"metadata": {"image_grid_thw": [1, 32, 48]}}} - - -class TestConsumerReservationManager: - @pytest.mark.parametrize("cancelled", [None, "writer", "follower"]) - def test_inflight_duplicates_share_memory_and_cancel_independently(self, cancelled): - engine = Mock() - engine.register_memory.return_value = 0 - pool = ConsumerMemoryPool(256, engine) - pool.prepare(torch.device("cpu")) - manager = ConsumerReservationManager(pool, 300, 16) - writer, write = manager.reserve( - "writer", "hash", 64, (16,), "float32", torch.float32 - ) - follower, write_again = manager.reserve( - "follower", "hash", 64, (16,), "float32", torch.float32 - ) - assert write and not write_again - assert writer.allocation is follower.allocation - tensor = writer.allocation.tensor - tensor.fill_(7) - if cancelled is not None: - manager.cancel(cancelled, "") - assert manager.complete("writer", writer.reservation_id)[0] - for name in ("writer", "follower"): - if name != cancelled: - loaded = manager.take(name, "hash") - assert loaded.tensor.data_ptr() == tensor.data_ptr() - assert torch.all(loaded.tensor == 7) - assert pool.try_allocate(64, (16,), torch.float32) is None - - def test_follower_refresh_does_not_abandon_the_shared_writer(self): - manager, pool, _ = self.manager() - writer, _ = self.reserve(manager) - follower, _ = manager.reserve( - "follower", "hash", 64, (16,), "float32", torch.float32 - ) - assert manager.cancel( - "follower", follower.reservation_id, abandon=True, refresh=True - ) - renewed, write = manager.reserve( - "follower", "hash", 64, (16,), "float32", torch.float32 - ) - assert not write and renewed.writer_id == writer.transfer_id - assert renewed.reservation_id != follower.reservation_id - assert writer.state is ConsumerReservationState.WRITING - pool.free.assert_not_called() - - @staticmethod - def manager(): - pool = Mock() - pool.lock = threading.RLock() - pool.acquire_cached.return_value = None - allocation = memory.MemoryAllocation(0, 64, torch.empty(16)) - pool.try_allocate.return_value = allocation - pool.reclaim_and_allocate.return_value = None - return ConsumerReservationManager(pool, 300, 16), pool, allocation - - @staticmethod - def reserve(manager: ConsumerReservationManager): - record, write = manager.reserve( - "transfer", "hash", 64, (16,), "float32", torch.float32 - ) - assert record is not None - return record, write - - def test_complete_is_idempotent(self): - manager, _, _ = self.manager() - record, write = self.reserve(manager) - - assert write and record.state is ConsumerReservationState.WRITING - assert manager.complete("transfer", record.reservation_id) == (True, True) - assert manager.complete("transfer", record.reservation_id) == (True, False) - assert record.state is ConsumerReservationState.READY - - def test_cancel_waits_for_an_active_writer(self): - manager, pool, allocation = self.manager() - record, _ = self.reserve(manager) - - assert not manager.cancel("transfer", "wrong-id") - assert manager.cancel("transfer", record.reservation_id) - assert record.state is ConsumerReservationState.CANCEL_PENDING - pool.free.assert_not_called() - - assert manager.complete("transfer", record.reservation_id) == (True, False) - assert record.state is ConsumerReservationState.CANCELLED - assert record.allocation is None - pool.free.assert_called_once_with(allocation) - - def test_shutdown_rejects_new_reservations(self): - manager, _, _ = self.manager() - manager.begin_shutdown() - - with pytest.raises(RuntimeError, match="shutting down"): - self.reserve(manager) - - def test_expired_writer_is_replaced_only_after_refresh_abandon(self): - manager, pool, old_allocation = self.manager() - new_allocation = memory.MemoryAllocation(256, 64, torch.ones(16)) - pool.try_allocate.side_effect = [old_allocation, new_allocation] - old, _ = self.reserve(manager) - old.expires_at = 0 - manager.expire() - - with pytest.raises(RuntimeError, match="active writer"): - self.reserve(manager) - assert manager.cancel( - "transfer", old.reservation_id, abandon=True, refresh=True - ) - new, write = self.reserve(manager) - - assert write and new.reservation_id != old.reservation_id - assert new.allocation is new_allocation - assert manager.complete("transfer", old.reservation_id) == (False, False) - pool.free.assert_called_once_with(old_allocation) - - def test_ready_expiry_releases_once(self): - manager, pool, allocation = self.manager() - record, _ = self.reserve(manager) - manager.complete("transfer", record.reservation_id) - record.expires_at = 0 - - assert manager.expire() == 1 - assert manager.expire() == 0 - assert record.state is ConsumerReservationState.EXPIRED - pool.free.assert_called_once_with(allocation) - - def test_cached_take_uses_the_pool_canonical_allocation(self): - manager, pool, cached = self.manager() - lease = SimpleNamespace(value=cached) - canonical = memory.MemoryAllocation(256, 64, torch.ones(16)) - pool.acquire_cached.return_value = lease - pool.publish.return_value = canonical - - record, write = self.reserve(manager) - assert not write and record.lease is lease - assert manager.take("transfer", "hash") is canonical - pool.publish.assert_called_once_with("hash", cached, lease) - - def test_cancel_tombstones_are_bounded(self): - manager, _, _ = self.manager() - manager._tombstone_limit = 3 - - for transfer_id in ("a", "b", "c", "a", "d"): - assert manager.cancel(transfer_id, "") - - assert list(manager._tombstones) == ["c", "a", "d"] - assert set(manager._records) == {"c", "a", "d"} - - @pytest.mark.parametrize( - "deferred,terminal", - [ - ( - ConsumerReservationState.CANCEL_PENDING, - ConsumerReservationState.CANCELLED, - ), - ( - ConsumerReservationState.EXPIRE_PENDING, - ConsumerReservationState.EXPIRED, - ), - ], - ) - def test_a_deferred_release_requires_writer_completion(self, deferred, terminal): - """A timeout cannot prove that a remote writer stopped using its address.""" - manager, pool, allocation = self.manager() - record, _ = self.reserve(manager) - - if deferred is ConsumerReservationState.CANCEL_PENDING: - assert manager.cancel("transfer", "") - else: - record.expires_at = 0 - manager.expire() - assert record.state is deferred - # Mooncake may still be writing, so the release waits first. - pool.free.assert_not_called() - - record.expires_at = 0 - assert manager.expire() == 0 - pool.free.assert_not_called() - assert manager.complete("transfer", record.reservation_id) == (True, False) - assert record.state is terminal - pool.free.assert_called_once_with(allocation) - - -class TestECMooncakeWorkerTransfer: - """Validate end-to-end Worker reservation, push, load, and cleanup flows.""" - - def test_partial_batch_reservation_cleans_only_the_failed_item(self): - """An item rejected on one shard must not cancel its successful sibling.""" - worker = object.__new__(ECMooncakeWorker) - worker._control_client = Mock() - shards = ["tcp://consumer:0", "tcp://consumer:1"] - worker._control_client.discover_shards.return_value = shards - specs = [ - ECMooncakePushSpec(name, 64, (16,), "float32", shards[0], name) - for name in ("good", "bad") - ] - - def request(addr, payload): - assert payload["op"] == "reserve_batch" - return { - "items": [ - {"ok": True, "result": {"reservation_id": "good-" + addr}}, - {"ok": False, "error": "full"} - if addr == shards[1] - else {"ok": True, "result": {"reservation_id": "partial"}}, - ] - } - - worker._control_client.request.side_effect = request - worker._run_fanout = lambda tasks: [task() for task in tasks] - worker._retry_cancel_reservations = Mock() - good, bad = worker._reserve_remote_many(specs) - assert [item["addr"] for item in good] == shards - assert isinstance(bad, _FanoutError) and str(bad) == "full" - worker._retry_cancel_reservations.assert_called_once_with(specs[1], bad.results) - assert [item["reservation_id"] for item in bad.results] == ["partial", ""] - - def test_reservation_completion_dispatches_without_another_model_step( - self, mock_vllm_config_producer - ): - allow_reservation = threading.Event() - transferred = threading.Event() - source = torch.ones(16) - spec = ECMooncakePushSpec("hash", 64, (16,), "float32", "unused", "transfer") - with patch_ec_mooncake_deps(): - worker = ECMooncakeWorker(mock_vllm_config_producer) - - def reserve(_): - assert allow_reservation.wait(5) - return [] - - def write(records): - for record in records: - worker._producer_pushes.begin_writing(record) - worker._producer_pushes.begin_notifying(records) - worker._producer_pushes.complete(records) - transferred.set() - - worker._reserve_remote = reserve - worker._push_batch = write - try: - worker.start_save_caches( - ECMooncakeConnectorMetadata(pushes=[spec]), {"hash": source} - ) - worker.build_connector_worker_meta() - assert not transferred.is_set() - allow_reservation.set() - assert transferred.wait(5) - finally: - allow_reservation.set() - worker.close() - - @pytest.mark.parametrize("waiting_for", ["reservation", "encoder"]) - def test_a_ready_push_does_not_wait_for_another_item(self, waiting_for): - manager = ProducerPushManager() - records = [] - for name in ("slow", "fast"): - future: Future[list[dict[str, Any]]] = Future() - spec = ECMooncakePushSpec( - name, 64, (16,), "float32", "tcp://consumer:1", name - ) - record, _ = manager.reserve(spec, lambda future=future: future) - event = Mock() - event.query.return_value = name == "fast" or waiting_for == "reservation" - manager.bind_source(name, torch.empty(16), event) - if name == "fast" or waiting_for == "encoder": - future.set_result([]) - records.append(record) - executor = Mock() - executor.submit.return_value = Future() - run = Mock() - manager.submit_batches(executor, run) - executor.submit.assert_called_once_with(run, [records[1]]) - - def test_reservation_requires_confirmed_topology_before_any_rpc(self): - worker = object.__new__(ECMooncakeWorker) - worker._control_client = Mock() - worker._control_client.discover_shards.return_value = None - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="transfer", - ) - - with pytest.raises(RuntimeError, match="discover every EC consumer shard"): - worker._reserve_remote(spec) - - assert worker._control_client.discover_shards.call_count == 2 - worker._control_client.request.assert_not_called() - - def test_stale_shards_are_abandoned_before_remote_re_reserve(self): - worker = object.__new__(ECMooncakeWorker) - worker._control_client = Mock() - events: list[tuple[str, str] | tuple[str]] = [] - - def request(addr, payload): - events.append(("abandon", payload["reservation_id"])) - assert payload["abandon"] and payload["refresh"] - return {"cancelled": True} - - worker._control_client.request.side_effect = request - replacement = [{"reservation_id": "new"}] - - def reserve_remote(spec): - events.append(("reserve",)) - return replacement - - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="transfer", - ) - shards = [ - { - "addr": f"tcp://consumer:{19019 + rank}", - "reservation_id": f"old-{rank}", - "ready": False, - } - for rank in range(2) - ] - - with ( - ThreadPoolExecutor(max_workers=2) as executor, - patch.object(worker, "_shard_executor", return_value=executor), - patch.object(worker, "_reserve_remote", side_effect=reserve_remote), - ): - assert worker._refresh_remote_reservations(spec, shards) is replacement - assert set(events[:2]) == { - ("abandon", "old-0"), - ("abandon", "old-1"), - } - assert events[2] == ("reserve",) - - def test_producer_push_state_owns_source_until_every_future_is_terminal(self): - manager = ProducerPushManager() - reservation: Future[list[dict[str, Any]]] = Future() - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="transfer", - ) - record, created = manager.reserve(spec, lambda: reservation) - duplicate, duplicate_created = manager.reserve(spec, lambda: Future()) - assert created - assert duplicate is record - assert not duplicate_created - # Another request may legitimately name the same encoding. - reasked = copy.copy(spec) - reasked.request_id = "another-request" - assert manager.reserve(reasked, lambda: Future()) == (record, False) - - # A different payload under the same id drops the newcomer instead of - # failing the engine; the push in flight keeps the id. - changed = copy.copy(spec) - changed.mm_hash = "other" - with pytest.raises(ValueError, match="Conflicting EC destination"): - manager.reserve(changed, lambda: Future()) - assert record.spec.mm_hash == "hash" - - source = torch.empty(16) - manager.bind_source("hash", source, None) - assert record.source_tensor is source - reservation.set_result([]) - assert manager.resolve_reservations(record) == [] - assert record.state is ProducerPushState.WAITING_INPUTS - manager.begin_writing(record) - manager.begin_notifying([record]) - - failed: Future[None] = Future() - failed.set_exception(RuntimeError("one shard failed")) - still_writing: Future[None] = Future() - manager.track_shard_futures([record], [failed, still_writing]) - with pytest.raises(RuntimeError, match="source too early"): - manager.fail([record], RuntimeError("write failed")) - assert record.state is ProducerPushState.NOTIFYING - assert record.source_tensor is source - - still_writing.set_result(None) - manager.fail([record], RuntimeError("write failed")) - assert record.state is ProducerPushState.FAILED - assert record.source_tensor is None - manager.fail([record], RuntimeError("duplicate failure")) - with pytest.raises(RuntimeError, match="FAILED to NOTIFYING"): - manager.begin_notifying([record]) - - late, late_created = manager.reserve(spec, lambda: Future()) - assert late is record - assert not late_created - - def test_cancel_requests_none_selects_every_source_less_waiter(self): - manager = ProducerPushManager() - - def reserve(transfer_id: str, mm_hash: str, request_id: str): - spec = ECMooncakePushSpec( - mm_hash=mm_hash, - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id=transfer_id, - request_id=request_id, - ) - return manager.reserve(spec, lambda: Future())[0] - - first = reserve("first", "hash-first", "request-first") - second = reserve("second", "hash-second", "request-second") - assert manager.cancel_requests({"request-first"}) == [first] - assert manager.cancel_requests(None) == [second] - assert first.state is ProducerPushState.CANCEL_PENDING - assert second.state is ProducerPushState.CANCEL_PENDING - - def test_worker_close_cancels_orphaned_reservations_before_executor_shutdown( - self, - ): - manager = ProducerPushManager() - reservation: Future[list[dict[str, Any]]] = Future() - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="orphan", - request_id="request", - ) - record, _ = manager.reserve(spec, lambda: reservation) - events: list[Any] = [] - - class RecordingExecutor: - def __init__(self, name: str): - self.name = name - - def submit(self, function, *args): - events.append(f"{self.name}.submit") - future: Future[Any] = Future() - try: - result = function(*args) - except BaseException as exc: - future.set_exception(exc) - else: - future.set_result(result) - return future - - def shutdown(self, wait=True, **kwargs): - events.append((f"{self.name}.shutdown", wait, kwargs)) - - worker = object.__new__(ECMooncakeWorker) - worker._producer_pushes = manager - worker._io_executor = RecordingExecutor("io") - worker._control_executor = RecordingExecutor("control") - worker._shard_pool = RecordingExecutor("shard") - worker._shutdown = False - worker._control_client = Mock() - worker._control_server = None - worker._consumer_memory = Mock() - worker._producer_memory = Mock() - worker._transfer = Mock() - worker._flush_pending_pushes = lambda: events.append("flush") - - def finish_cancel(orphan: ProducerPushRecord): - events.append(f"cancel:{orphan.spec.transfer_id}") - manager.finish_cancel(orphan) - - worker._cancel_orphaned_reservation = finish_cancel - - worker.close() - worker.close() - - assert record.state is ProducerPushState.CANCELLED - assert events == [ - "flush", - "io.submit", - "cancel:orphan", - ("control.shutdown", True, {}), - ("io.shutdown", True, {}), - ("shard.shutdown", True, {}), - ] - worker._control_client.close.assert_called_once_with() - worker._consumer_memory.close.assert_called_once_with() - worker._producer_memory.close.assert_called_once_with() - worker._transfer.close.assert_called_once_with() - - def test_worker_close_drains_an_unresolved_reservation_before_control_shutdown( - self, - ): - manager = ProducerPushManager() - control_executor = ThreadPoolExecutor(max_workers=1) - - def resolve_reservation(): - time.sleep(0.02) - return [ - { - "addr": "tcp://consumer:19019", - "reservation_id": "r0", - }, - { - "addr": "tcp://consumer:19020", - "reservation_id": "r1", - }, - ] - - reservation = control_executor.submit(resolve_reservation) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="orphan", - request_id="request", - ) - record, _ = manager.reserve(spec, lambda: reservation) - io_executor = ThreadPoolExecutor(max_workers=1) - shard_pool = ThreadPoolExecutor(max_workers=1) - worker = object.__new__(ECMooncakeWorker) - worker._producer_pushes = manager - worker._io_executor = io_executor - worker._control_executor = control_executor - worker._shard_pool = shard_pool - worker._shard_pool_lock = threading.Lock() - worker._shutdown = False - worker._control_client = Mock() - worker._control_client.request.return_value = {"cancelled": True} - worker._control_server = None - worker._consumer_memory = Mock() - worker._producer_memory = Mock() - worker._transfer = Mock() - worker._flush_pending_pushes = Mock() - - worker.close() - - assert record.state is ProducerPushState.CANCELLED - assert worker._control_client.request.call_count == 2 - assert all( - request.args[1]["op"] == "cancel" and request.args[1]["abandon"] - for request in worker._control_client.request.call_args_list - ) - - def test_consumer_close_waits_for_remote_writer_before_releasing_pool(self): - reservations, consumer_memory, _ = TestConsumerReservationManager.manager() - record, _ = TestConsumerReservationManager.reserve(reservations) - shutdown_started = threading.Event() - original_begin_shutdown = reservations.begin_shutdown - - def begin_shutdown(): - original_begin_shutdown() - shutdown_started.set() - - reservations.begin_shutdown = begin_shutdown - worker = object.__new__(ECMooncakeWorker) - worker._producer_pushes = ProducerPushManager() - worker._io_executor = Mock() - worker._control_executor = Mock() - worker._shard_pool = None - worker._shutdown = False - worker._control_client = Mock() - worker._control_server = Mock() - worker._consumer_memory = consumer_memory - worker._producer_memory = Mock() - worker._transfer = Mock() - worker._reservations = reservations - worker._shutdown_drain_timeout_s = 1 - worker._flush_pending_pushes = Mock() - - close_thread = threading.Thread(target=worker.close) - close_thread.start() - assert shutdown_started.wait(1) - assert record.state is ConsumerReservationState.CANCEL_PENDING - consumer_memory.close.assert_not_called() - worker._control_server.close.assert_not_called() - - assert reservations.complete("transfer", record.reservation_id) == ( - True, - False, - ) - close_thread.join(1) - - assert not close_thread.is_alive() - consumer_memory.close.assert_called_once_with() - worker._control_server.close.assert_called_once_with() - - def test_consumer_close_timeout_keeps_receive_pool_registered(self): - reservations, consumer_memory, allocation = ( - TestConsumerReservationManager.manager() - ) - record, _ = TestConsumerReservationManager.reserve(reservations) - worker = object.__new__(ECMooncakeWorker) - worker._producer_pushes = ProducerPushManager() - worker._io_executor = Mock() - worker._control_executor = Mock() - worker._shard_pool = None - worker._shutdown = False - worker._control_client = Mock() - worker._control_server = Mock() - worker._consumer_memory = consumer_memory - worker._producer_memory = Mock() - worker._transfer = Mock() - worker._reservations = reservations - worker._shutdown_drain_timeout_s = 0 - worker._flush_pending_pushes = Mock() - - worker.close() - - assert record.state is ConsumerReservationState.CANCEL_PENDING - assert record.allocation is allocation - consumer_memory.close.assert_not_called() - worker._control_server.close.assert_called_once_with() - - def test_permanent_orphan_cleanup_failure_marks_push_failed(self): - manager = ProducerPushManager() - reservation: Future[list[dict[str, Any]]] = Future() - reservation.set_result( - [ - { - "addr": "tcp://consumer:19019", - "reservation_id": "reservation", - "cached": True, - } - ] - ) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="transfer", - request_id="request", - ) - record, _ = manager.reserve(spec, lambda: reservation) - assert manager.cancel_requests({"request"}) == [record] - - worker = object.__new__(ECMooncakeWorker) - worker._producer_pushes = manager - worker._control_client = Mock() - worker._control_client.request.side_effect = RuntimeError("cancel failed") - with ( - ThreadPoolExecutor(max_workers=1) as executor, - patch.object(worker, "_shard_executor", return_value=executor), - ): - worker._cancel_orphaned_reservation(record) - - assert record.state is ProducerPushState.FAILED - assert record.error == "cancel failed" - assert worker._control_client.request.call_count == 2 - assert manager.poll() == [("hash", "cancel failed")] - assert manager.poll() == [] - - def test_orphan_topology_failure_marks_push_failed_without_base_cancel(self): - manager = ProducerPushManager() - reservation: Future[list[dict[str, Any]]] = Future() - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:19019", - transfer_id="transfer", - request_id="request", - ) - record, _ = manager.reserve(spec, lambda: reservation) - assert manager.cancel_requests({"request"}) == [record] - reservation.set_exception(RuntimeError("topology unavailable")) - - worker = object.__new__(ECMooncakeWorker) - worker._producer_pushes = manager - worker._control_client = Mock() - worker._cancel_orphaned_reservation(record) - - assert record.state is ProducerPushState.FAILED - assert record.error == "topology unavailable" - worker._control_client.request.assert_not_called() - - def test_reservation_failure_after_source_binding_releases_the_lease(self): - manager = ProducerPushManager() - reservation: Future[list[dict[str, Any]]] = Future() - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=64, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="transfer", - ) - record, _ = manager.reserve(spec, lambda: reservation) - source = torch.empty(16) - manager.bind_source("hash", source, None) - reservation.set_exception(RuntimeError("reserve failed")) - - assert record.state is ProducerPushState.WAITING_INPUTS - - def run(records) -> None: - try: - manager.resolve_reservations(records[0]) - except RuntimeError as exc: - manager.fail(records, exc) - - with ThreadPoolExecutor(max_workers=1) as executor: - manager.submit_batches(executor, run) - assert manager.poll() == [("hash", "reserve failed")] - assert manager.poll() == [] - assert record.state is ProducerPushState.FAILED - assert record.source_tensor is None - - def test_shard_submit_failure_waits_before_source_release( - self, mock_vllm_config_producer - ): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.empty(16) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq="tcp://consumer:1", - transfer_id="transfer", - ) - slow_started = threading.Event() - finish_slow = threading.Event() - slow_finished = threading.Event() - released_after_slow: list[bool] = [] - - def request(addr, payload): - if payload["op"] == "reserve": - index = int(addr.rsplit(":", 1)[1]) - return { - "reservation_id": f"reservation-{index}", - "dst_session": f"session-{index}", - "dst_ptr": 1000 + index, - "nbytes": source.nbytes, - "write": True, - "ready": False, - } - return {} - - def write(session, sources, destinations, lengths): - if session == "session-1": - slow_started.set() - assert finish_slow.wait(2) - slow_finished.set() - - with patch_ec_mooncake_deps(): - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - worker = producer._worker - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) - try: - with ( - patch.object( - worker._control_client, - "discover_shards", - return_value=[ - "tcp://consumer:0", - "tcp://consumer:1", - "tcp://consumer:2", - ], - ), - patch.object( - worker._control_client, "request", side_effect=request - ), - patch.object(worker._producer_memory, "stage", return_value=None), - patch.object( - worker._transfer, - "acquire_sources", - return_value=[source.data_ptr()], - ), - patch.object( - worker._transfer, - "release_sources", - side_effect=lambda _: released_after_slow.append( - slow_finished.is_set() - ), - ), - patch.object(worker._transfer, "write", side_effect=write), - ): - producer.start_save_caches(encoder_cache={"hash": source}) - record = worker._producer_pushes._records.get("transfer") - assert record is not None - record.reservation_future.result(timeout=2) - with ThreadPoolExecutor(max_workers=1) as executor: - submit_count = 0 - - def submit(fn, *args): - nonlocal submit_count - submit_count += 1 - if submit_count == 1: - return executor.submit(fn, *args) - raise RuntimeError("second shard submit failed") - - shard_executor = MagicMock() - shard_executor.submit.side_effect = submit - with patch.object( - worker, - "_shard_executor", - return_value=shard_executor, - ): - assert producer.build_connector_worker_meta().pending_saves - assert slow_started.wait(2) - assert record.source_tensor is source - assert released_after_slow == [] - finish_slow.set() - _wait_for_worker_io(producer) - - record = worker._producer_pushes._records.get("transfer") - assert record is not None - assert record.state is ProducerPushState.FAILED - assert record.source_tensor is None - assert released_after_slow == [True] - finally: - finish_slow.set() - producer.shutdown() - - @pytest.mark.parametrize("source_before_failure", [False, True]) - def test_partial_reserve_is_compensated_before_its_future_fails( - self, mock_vllm_config_producer, source_before_failure - ): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.empty(16) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq="tcp://consumer:0", - transfer_id="transfer", - ) - cancel_attempts: Counter[str] = Counter() - - def reserve_one(addr, _spec): - if addr.endswith(":1"): - raise RuntimeError("reserve shard failed") - return {"addr": addr, "reservation_id": "partial-r0"} - - def request(addr, payload): - assert payload["op"] == "cancel" and payload["abandon"] - assert payload["transfer_id"] == "transfer" - reservation_id = str(payload["reservation_id"]) - cancel_attempts[reservation_id] += 1 - if ( - addr.endswith(":0") - and reservation_id == "partial-r0" - and cancel_attempts[reservation_id] == 1 - ): - raise RuntimeError("transient cleanup failure") - return {"cancelled": True} - - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - worker = connector._worker - connector.bind_connector_metadata( - ECMooncakeConnectorMetadata(pushes=[spec]) - ) - try: - with ( - patch.object( - worker._control_client, - "discover_shards", - return_value=["tcp://consumer:0", "tcp://consumer:1"], - ), - patch.object(worker, "_reserve_one", side_effect=reserve_one), - patch.object( - worker._control_client, "request", side_effect=request - ), - ): - connector.start_save_caches( - encoder_cache={"hash": source} - if source_before_failure - else None - ) - record = worker._producer_pushes._records.get("transfer") - assert record is not None - with pytest.raises( - RuntimeError, match="^reserve shard failed$" - ) as e: - record.reservation_future.result(timeout=2) - assert e.value.results == [ - { - "addr": "tcp://consumer:0", - "reservation_id": "partial-r0", - }, - {"addr": "tcp://consumer:1", "reservation_id": ""}, - ] - assert cancel_attempts == Counter({"partial-r0": 2, "": 1}) - - if source_before_failure: - assert record.source_tensor is source - connector.build_connector_worker_meta() - _wait_for_worker_io(connector) - else: - assert record.state is ProducerPushState.FAILED - connector.save_caches({"hash": source}, "hash") - assert record.state is ProducerPushState.FAILED - assert record.source_tensor is None - finally: - connector.shutdown() - - def test_partial_complete_abandons_all_shards_before_releasing_source( - self, mock_vllm_config_producer - ): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.empty(16) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq="tcp://consumer:0", - transfer_id="transfer", - ) - reservations = [ - { - "addr": f"tcp://consumer:{rank}", - "reservation_id": f"r{rank}", - "dst_session": f"session-{rank}", - "dst_ptr": 1000 + rank, - "nbytes": source.nbytes, - "write": True, - "ready": False, - } - for rank in range(2) - ] - slow_started = threading.Event() - finish_slow = threading.Event() - cancelled: list[str] = [] - - def request(addr, payload): - if payload["op"] == "complete_batch": - if addr.endswith(":0"): - raise RuntimeError("complete shard failed") - slow_started.set() - assert finish_slow.wait(2) - return {"items": [{"completed": True}]} - assert payload["op"] == "cancel" and payload["abandon"] - cancelled.append(payload["reservation_id"]) - return {"cancelled": True} - - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - worker = connector._worker - connector.bind_connector_metadata( - ECMooncakeConnectorMetadata(pushes=[spec]) - ) - try: - with ( - patch.object(worker, "_reserve_remote", return_value=reservations), - patch.object( - worker._control_client, "request", side_effect=request - ), - patch.object(worker._producer_memory, "stage", return_value=None), - patch.object( - worker._transfer, - "acquire_sources", - return_value=[source.data_ptr()], - ), - patch.object(worker._transfer, "release_sources"), - patch.object(worker._transfer, "write"), - ): - connector.start_save_caches(encoder_cache={"hash": source}) - assert connector.build_connector_worker_meta().pending_saves - record = worker._producer_pushes._records.get("transfer") - assert record is not None and record.batch_future is not None - assert slow_started.wait(2) - assert record.source_tensor is source - assert not record.batch_future.done() - finish_slow.set() - record.batch_future.result(timeout=2) - connector.build_connector_worker_meta() - - assert Counter(cancelled) == Counter({"r0": 1, "r1": 1}) - assert record.state is ProducerPushState.FAILED - assert record.source_tensor is None - assert record.error == "complete shard failed" - assert all(future.done() for future in record.shard_futures) - finally: - finish_slow.set() - connector.shutdown() - - def test_invalid_source_fails_asynchronously_before_staging( - self, mock_vllm_config_producer - ): - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.empty(2, 8) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=(16,), - dtype="float32", - consumer_zmq="tcp://consumer:0", - transfer_id="transfer", - ) - - def request(_addr, payload): - if payload["op"] == "reserve": - return { - "reservation_id": "reservation", - "dst_session": "session", - "dst_ptr": 1000, - "nbytes": source.nbytes, - "write": True, - "ready": False, - } - assert payload["op"] == "cancel" - return {"cancelled": True} - - with patch_ec_mooncake_deps(): - connector = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - worker = connector._worker - connector.bind_connector_metadata( - ECMooncakeConnectorMetadata(pushes=[spec]) - ) - try: - with ( - patch.object( - worker._control_client, - "discover_shards", - return_value=["tcp://consumer:0"], - ), - patch.object( - worker._control_client, "request", side_effect=request - ), - patch.object(worker._producer_memory, "stage") as stage, - patch.object(worker._transfer, "acquire_sources") as register, - ): - connector.start_save_caches(encoder_cache={"hash": source}) - assert connector.build_connector_worker_meta().pending_saves - record = worker._producer_pushes._records.get("transfer") - assert record is not None and record.batch_future is not None - record.batch_future.result(timeout=2) - connector.build_connector_worker_meta() - - assert record.state is ProducerPushState.FAILED - assert record.source_tensor is None - assert record.error == "EC source shape mismatch for mm_hash=hash" - stage.assert_not_called() - register.assert_not_called() - finally: - connector.shutdown() - - def test_batches_pushes_from_one_model_step(self, mock_vllm_config_producer): - port = _find_free_port() - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config - consumer_cfg.model_config = Mock() - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 - consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" - consumer_cfg.ec_transfer_config.ec_port = port - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(consumer_cfg) - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - sources = { - "first": torch.randn(4, 16), - "second": torch.randn(8, 16), - } - pushes = [ - ECMooncakePushSpec( - mm_hash=mm_hash, - nbytes=tensor.nbytes, - shape=tuple(tensor.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id=f"transfer-{mm_hash}", - ) - for mm_hash, tensor in sources.items() - ] - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - consumer.start_worker_services() - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=pushes)) - try: - producer.start_save_caches(encoder_cache=sources) - _wait_for_worker_io(producer) - - engine = producer._worker._transfer._engine - assert isinstance(engine, CopyingFakeTransferEngine) - assert len(engine.transfer_calls) == 1 - assert sorted(engine.transfer_calls[0]) == sorted( - tensor.nbytes for tensor in sources.values() - ) - assert all( - reservation.state is ConsumerReservationState.READY - for reservation in consumer._worker._reservations._records.values() - ) - finally: - producer.shutdown() - consumer.shutdown() - - def test_push_reserves_before_encoder_output_is_saved( - self, mock_vllm_config_producer - ): - port = _find_free_port() - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config - consumer_cfg.model_config = Mock() - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 - consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" - consumer_cfg.ec_transfer_config.ec_port = port - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(consumer_cfg) - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.randn(4, 16) - push = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id="transfer-1", - ) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - consumer.start_worker_services() - scheduler = ECMooncakeConnector(consumer_cfg, ECConnectorRole.SCHEDULER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) - try: - producer.start_save_caches(encoder_cache={}) - push_record = producer._worker._producer_pushes._records.get( - "transfer-1" - ) - assert push_record is not None - reservation = push_record.reservation_future - shards = reservation.result(timeout=2) - # One reservation per consumer shard; this consumer is single. - assert len(shards) == 1 - reservation_data = shards[0] - assert reservation_data["nbytes"] == source.nbytes - old_reservation_id = reservation_data["reservation_id"] - reservation_data["_received_at"] -= _RESERVATION_TTL_SECONDS - consumer._worker._reservations._records["transfer-1"].expires_at = 0 - with patch.object( - scheduler._scheduler._control_client, - "request", - wraps=scheduler._scheduler._control_client.request, - ) as send_control: - assert not scheduler.has_cache_item("hash") - assert not scheduler.has_cache_item("hash") - _drain_until_subscribed(scheduler._scheduler) - # The channel is built once, not per drain: the roster is - # fetched and every shard subscribed to exactly once. - assert [call.args[1] for call in send_control.call_args_list] == [ - {"op": "peers"}, - {"op": "event_port"}, - ] - assert consumer._worker._reservations.status("transfer-1") - - producer.save_caches({"hash": source}, "hash") - _wait_for_worker_io(producer) - assert ( - consumer._worker._reservations._records[ - "transfer-1" - ].reservation_id - != old_reservation_id - ) - deadline = time.monotonic() + 2 - while not scheduler.has_cache_item("hash"): - assert time.monotonic() < deadline - time.sleep(0.01) - # Still just the two setup requests: polling for readiness - # must not re-open the channel. - assert send_control.call_count == 2 - record = scheduler._scheduler._transfers.get("transfer-1") - assert record is not None and record.spec is not None - load = record.spec - consumer.bind_connector_metadata( - ECMooncakeConnectorMetadata(loads=[load]) - ) - loaded: dict[str, torch.Tensor] = {} - consumer.start_load_caches(loaded) - first_meta = consumer.build_connector_worker_meta() - assert first_meta.loaded == {"hash"} - assert torch.equal(loaded["hash"], source) - consumer_engine = consumer._worker._transfer._engine - assert isinstance(consumer_engine, CopyingFakeTransferEngine) - assert consumer_engine.transfer_calls == [] - finally: - producer.shutdown() - scheduler.shutdown() - consumer.shutdown() - - def test_finished_request_cancels_unbound_reservation( - self, mock_vllm_config_producer - ): - """A pre-reservation without an encoder tensor must not outlive its request.""" - port = _find_free_port() - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config - consumer_cfg.model_config = Mock() - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 - consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" - consumer_cfg.ec_transfer_config.ec_port = port - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(consumer_cfg) - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.randn(4, 16) - push = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id="transfer-1", - request_id="request-1", - ) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - consumer.start_worker_services() - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[push])) - try: - producer.start_save_caches(encoder_cache={}) - push_record = producer._worker._producer_pushes._records.get( - "transfer-1" - ) - assert push_record is not None - reservation = push_record.reservation_future - reservation.result(timeout=2) - assert consumer._worker._reservations.status("transfer-1") - - producer.get_finished({"request-1"}) - _wait_for_worker_io(producer) - push_record = producer._worker._producer_pushes._records.get( - "transfer-1" - ) - assert push_record is not None - assert push_record.state is ProducerPushState.CANCELLED - assert consumer._worker._reservations.status("transfer-1") is None - finally: - producer.shutdown() - consumer.shutdown() - - def test_duplicate_pushes_share_one_transfer_per_reservation( - self, mock_vllm_config_producer - ): - port = _find_free_port() - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config - consumer_cfg.model_config = Mock() - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 - consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" - consumer_cfg.ec_transfer_config.ec_port = port - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(consumer_cfg) - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.randn(4, 16) - push = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id="transfer-1", - ) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - scheduler = ECMooncakeConnector(consumer_cfg, ECConnectorRole.SCHEDULER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - consumer.start_worker_services() - producer.bind_connector_metadata( - ECMooncakeConnectorMetadata(pushes=[push, push]) - ) - try: - producer.start_save_caches(encoder_cache={"hash": source}) - _wait_for_worker_io(producer) - - engine = producer._worker._transfer._engine - assert isinstance(engine, CopyingFakeTransferEngine) - assert engine.transfer_calls == [[source.nbytes]] - reservation = consumer._worker._reservations._records.get("transfer-1") - assert reservation.state is ConsumerReservationState.READY - - deadline = time.monotonic() + 2 - while not scheduler.has_cache_item("hash"): - assert time.monotonic() < deadline - time.sleep(0.01) - record = scheduler._scheduler._transfers.get("transfer-1") - assert record is not None and record.spec is not None - load = record.spec - consumer.bind_connector_metadata( - ECMooncakeConnectorMetadata(loads=[load]) - ) - loaded: dict[str, torch.Tensor] = {} - consumer.start_load_caches(loaded) - - cached_push = ECMooncakePushSpec( - mm_hash=push.mm_hash, - nbytes=push.nbytes, - shape=push.shape, - dtype=push.dtype, - consumer_zmq=push.consumer_zmq, - transfer_id="transfer-2", - ) - producer.bind_connector_metadata( - ECMooncakeConnectorMetadata(pushes=[cached_push]) - ) - producer.start_save_caches(encoder_cache={"hash": source}) - _wait_for_worker_io(producer) - assert engine.transfer_calls == [[source.nbytes]] - cached = consumer._worker._reservations._records.get("transfer-2") - assert cached is not None - assert cached.state is ConsumerReservationState.READY - assert cached.lease is not None - - deadline = time.monotonic() + 2 - while not scheduler.has_cache_item("hash"): - assert time.monotonic() < deadline - time.sleep(0.01) - record = scheduler._scheduler._transfers.get("transfer-2") - assert record is not None and record.spec is not None - cached_load = record.spec - consumer.bind_connector_metadata( - ECMooncakeConnectorMetadata(loads=[cached_load]) - ) - consumer.start_load_caches(loaded) - cached_meta = consumer.build_connector_worker_meta() - assert cached_meta.loaded == {"hash"} - assert consumer._worker._reservations.status("transfer-2") is None - assert torch.equal(loaded["hash"], source) - finally: - producer.shutdown() - scheduler.shutdown() - consumer.shutdown() - - def test_retired_item_reserved_again_still_serves_a_local_load( - self, mock_vllm_config_producer - ): - """A push for a retired item makes it live, not gone. - - Reusing the allocation for a new reservation takes it out of the - reclaim order. Looking the load up there instead of in the residency - map failed it, and the request fell back to waiting for a transfer. - """ - port = _find_free_port() - cfg = Mock(spec=VllmConfig) - cfg.parallel_config = mock_vllm_config_producer.parallel_config - cfg.model_config = Mock() - cfg.ec_transfer_config = Mock() - cfg.ec_transfer_config.is_ec_producer = False - cfg.ec_transfer_config.is_ec_consumer = True - cfg.ec_transfer_config.ec_buffer_device = "cpu" - cfg.ec_transfer_config.ec_buffer_size = 4096 - cfg.ec_transfer_config.ec_ip = "127.0.0.1" - cfg.ec_transfer_config.ec_port = port - cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(cfg) - spec = ECMooncakeLoadSpec( - mm_hash="hash", - nbytes=64, - shape=(4, 4), - dtype="float32", - transfer_id="local-transfer", - local=True, - ) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(cfg, ECConnectorRole.WORKER) - try: - consumer._worker._consumer_memory.prepare(torch.device("cpu")) - allocation = consumer._worker._consumer_memory.try_allocate( - spec.nbytes, spec.shape, torch.float32 - ) - assert allocation is not None - tensor = allocation.tensor - consumer._worker._consumer_memory.publish("hash", allocation) - retire_event = MagicMock() - retire_event.query.return_value = True - with ( - patch.object(memory.torch, "Event", return_value=retire_event), - patch.object(memory.torch.accelerator, "current_stream"), - ): - consumer._worker._consumer_memory.retire_stale({}, set()) - - # A later push reserves the retired copy instead of transferring. - consumer._worker._reserve_push_destination( - { - "transfer_id": "t1", - "mm_hash": "hash", - "nbytes": spec.nbytes, - "shape": list(spec.shape), - "dtype": spec.dtype, - } - ) - assert not consumer._worker._consumer_memory._residents._evictable - - assert ( - consumer._worker._consumer_memory.take_resident( - spec.mm_hash, spec.shape, spec.dtype - ) - is tensor - ) - finally: - consumer.shutdown() - - def test_push_reaches_every_consumer_shard(self, mock_vllm_config_producer): - """A sharded consumer gets one copy per rank, from one source. - - Each rank gathers from its own encoder cache, so the push has to land - on all of them; the bytes are identical, so staging and registration - happen once however many destinations there are. - """ - shard_ports = [_find_free_port() for _ in range(3)] - base = f"tcp://127.0.0.1:{shard_ports[0]}" - source = torch.randn(4, 16) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=base, - transfer_id="transfer-0", - ) - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - destinations = [torch.zeros_like(source) for _ in shard_ports] - - def fake_send(addr: str, request: dict): - if request["op"] == "peers": - return {"ports": shard_ports} - index = shard_ports.index(int(addr.rsplit(":", 1)[1])) - if request["op"] == "reserve": - return { - "reservation_id": f"r{index}", - "dst_session": f"session-{index}", - "dst_ptr": destinations[index].data_ptr(), - "nbytes": source.nbytes, - "write": True, - "ready": True, - "addr": addr, - } - if request["op"] == "complete_batch": - return {"items": [{"completed": True} for _ in request["items"]]} - return {} - - with patch_ec_mooncake_deps(): - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) - try: - with patch.object( - producer._worker._control_client, - "request", - side_effect=fake_send, - ): - producer.start_save_caches(encoder_cache={"hash": source}) - _wait_for_worker_io(producer) - - engine = producer._worker._transfer._engine - assert isinstance(engine, CopyingFakeTransferEngine) - # One write per rank, and every rank got the same bytes. - assert len(engine.transfer_calls) == len(shard_ports) - for destination in destinations: - assert torch.equal(destination, source) - # The source is staged once, not once per destination. - assert len(engine.register_calls) == 1 - finally: - producer.shutdown() - - def test_pushes_stage_through_the_registered_pool(self, mock_vllm_config_producer): - """Repeated content must not register overlapping source storage.""" - port = _find_free_port() - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = mock_vllm_config_producer.parallel_config - consumer_cfg.model_config = Mock() - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 - consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" - consumer_cfg.ec_transfer_config.ec_port = port - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(consumer_cfg) - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.randn(4, 16) - pushes = [ - ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id=f"transfer-{index}", - ) - for index in range(2) - ] - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - consumer.start_worker_services() - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=pushes)) - try: - producer.start_save_caches(encoder_cache={"hash": source}) - _wait_for_worker_io(producer) - - engine = producer._worker._transfer._engine - assert isinstance(engine, CopyingFakeTransferEngine) - # The staging pool is registered once; a transfer registers - # nothing of its own. - pool = producer._worker._producer_memory.tensor - assert pool is not None - assert engine.register_calls == [[pool.data_ptr()]] - assert engine.batch_unregister_calls == [] - assert engine.transfer_calls == [[source.nbytes]] - assert all( - reservation.state is ConsumerReservationState.READY - for reservation in consumer._worker._reservations._records.values() - ) - finally: - producer.shutdown() - consumer.shutdown() - - def test_push_falls_back_to_per_tensor_registration_without_a_pool( - self, mock_vllm_config_producer - ): - """A pool that cannot be created must not break pushes.""" - port = _find_free_port() - consumer_cfg = self._push_harness_config(mock_vllm_config_producer, port) - source = torch.randn(4, 16) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id="transfer", - ) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - consumer.start_worker_services() - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) - try: - with patch( - "vllm.distributed.ec_transfer.ec_connector." - "mooncake.memory.torch.empty", - side_effect=torch.OutOfMemoryError, - ): - producer.start_save_caches(encoder_cache={"hash": source}) - _wait_for_worker_io(producer) - engine = producer._worker._transfer._engine - assert isinstance(engine, CopyingFakeTransferEngine) - assert producer._worker._producer_memory.tensor is None - assert engine.register_calls == [[source.data_ptr()]] - assert engine.batch_unregister_calls == [[source.data_ptr()]] - assert engine.transfer_calls == [[source.nbytes]] - finally: - producer.shutdown() - consumer.shutdown() - - def test_concurrent_pushes_hold_source_registration_until_last_release( - self, mock_vllm_config_producer - ): - """Concurrent transfers share one MR until every user releases it.""" - mock_vllm_config_producer.ec_transfer_config.ec_buffer_device = "cpu" - source = torch.randn(4, 16) - - with patch_ec_mooncake_deps(): - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - try: - first = producer._worker._transfer.acquire_sources([source]) - second = producer._worker._transfer.acquire_sources([source]) - engine = producer._worker._transfer._engine - assert isinstance(engine, CopyingFakeTransferEngine) - assert len(engine.register_calls) == 1 - - producer._worker._transfer.release_sources(first) - assert engine.batch_unregister_calls == [] - producer._worker._transfer.release_sources(second) - assert engine.batch_unregister_calls == [first] - finally: - producer.shutdown() - - def _push_harness_config(self, producer_cfg, port: int): - consumer_cfg = Mock(spec=VllmConfig) - consumer_cfg.parallel_config = producer_cfg.parallel_config - consumer_cfg.model_config = Mock() - consumer_cfg.ec_transfer_config = Mock() - consumer_cfg.ec_transfer_config.is_ec_producer = False - consumer_cfg.ec_transfer_config.is_ec_consumer = True - consumer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - consumer_cfg.ec_transfer_config.ec_buffer_size = 4096 - consumer_cfg.ec_transfer_config.ec_ip = "127.0.0.1" - consumer_cfg.ec_transfer_config.ec_port = port - consumer_cfg.ec_transfer_config.ec_connector_extra_config = { - "mooncake_protocol": "tcp", - } - _bind_extra_config(consumer_cfg) - producer_cfg.ec_transfer_config.ec_buffer_device = "cpu" - return consumer_cfg - - def test_batch_completion_sends_one_control_message( - self, mock_vllm_config_producer - ): - """Completion is per batch, not per item: k items used to cost k RTTs.""" - port = _find_free_port() - consumer_cfg = self._push_harness_config(mock_vllm_config_producer, port) - sources = {"a": torch.randn(4, 16), "b": torch.randn(4, 16)} - pushes = [ - ECMooncakePushSpec( - mm_hash=mm_hash, - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id=f"transfer-{mm_hash}", - ) - for mm_hash, source in sources.items() - ] - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - consumer.start_worker_services() - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=pushes)) - try: - with patch.object( - producer._worker._control_client, - "request", - wraps=producer._worker._control_client.request, - ) as send_control: - producer.start_save_caches(encoder_cache=sources) - _wait_for_worker_io(producer) - ops = [call.args[1]["op"] for call in send_control.call_args_list] - assert ops.count("complete_batch") == 1 - assert "complete" not in ops - assert all( - reservation.state is ConsumerReservationState.READY - for reservation in consumer._worker._reservations._records.values() - ) - finally: - producer.shutdown() - consumer.shutdown() - - def test_failed_push_is_reported_not_raised(self, mock_vllm_config_producer): - """A transfer failure must not surface as a fatal engine error.""" - port = _find_free_port() - consumer_cfg = self._push_harness_config(mock_vllm_config_producer, port) - source = torch.randn(4, 16) - spec = ECMooncakePushSpec( - mm_hash="hash", - nbytes=source.nbytes, - shape=tuple(source.shape), - dtype="float32", - consumer_zmq=f"tcp://127.0.0.1:{port}", - transfer_id="transfer", - ) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector(consumer_cfg, ECConnectorRole.WORKER) - producer = ECMooncakeConnector( - mock_vllm_config_producer, ECConnectorRole.WORKER - ) - consumer.start_worker_services() - producer.bind_connector_metadata(ECMooncakeConnectorMetadata(pushes=[spec])) - try: - producer._worker._transfer.ensure_ready() - engine = producer._worker._transfer._engine - with patch.object(engine, "batch_transfer_sync_write", return_value=1): - producer.start_save_caches(encoder_cache={"hash": source}) - # No raise: the batch reports itself and gives up the - # consumer-side reservation. - _wait_for_worker_io(producer) - assert consumer._worker._reservations.status("transfer") is None - finally: - producer.shutdown() - consumer.shutdown() - - def test_complete_is_idempotent_without_republishing( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - consumer._worker._consumer_memory.prepare(torch.device("cpu")) - reservation = consumer._worker._reserve_push_destination( - { - "mm_hash": "hash", - "transfer_id": "transfer-1", - "nbytes": 64, - "shape": [4, 4], - "dtype": "float32", - } - ) - reservation_id = reservation["reservation_id"] - - first = consumer._worker._reservations.complete( - "transfer-1", reservation_id - ) - repeated = consumer._worker._reservations.complete( - "transfer-1", reservation_id - ) - - assert first == (True, True) - assert repeated == (True, False) - finally: - consumer.shutdown() - - def test_late_completion_cannot_complete_new_reservation( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_size = 4096 - payload = { - "mm_hash": "hash", - "transfer_id": "transfer", - "nbytes": 64, - "shape": [4, 4], - "dtype": "float32", - } - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - consumer._worker._consumer_memory.prepare(torch.device("cpu")) - old = consumer._worker._reserve_push_destination(payload) - consumer._worker._reservations._records["transfer"].expires_at = 0 - consumer._worker._reservations.expire() - assert ( - consumer._worker._reservations._records["transfer"].state - is ConsumerReservationState.EXPIRE_PENDING - ) - assert consumer._worker._reservations.cancel( - "transfer", - old["reservation_id"], - abandon=True, - refresh=True, - ) - new = consumer._worker._reserve_push_destination(payload) - new_record = consumer._worker._reservations._records.get("transfer") - assert new_record is not None and new_record.allocation is not None - new_allocation = new_record.allocation - - assert old["reservation_id"] != new["reservation_id"] - stale = consumer._worker._reservations.complete( - "transfer", old["reservation_id"] - ) - assert stale == (False, False) - assert ( - consumer._worker._reservations._records.get("transfer") - is new_record - ) - assert new_record.allocation is new_allocation - assert new_record.state is ConsumerReservationState.WRITING - finally: - consumer.shutdown() - - def test_missing_push_reservation_reports_failed_load( - self, mock_vllm_config_consumer - ): - mock_vllm_config_consumer.ec_transfer_config.ec_buffer_device = "cpu" - spec = ECMooncakeLoadSpec( - mm_hash="hash", - nbytes=32, - shape=(8,), - dtype="float32", - transfer_id="missing-transfer", - ) - - with patch_ec_mooncake_deps(): - consumer = ECMooncakeConnector( - mock_vllm_config_consumer, ECConnectorRole.WORKER - ) - try: - consumer.bind_connector_metadata( - ECMooncakeConnectorMetadata(loads=[spec]) - ) - cache: dict[str, torch.Tensor] = {} - consumer.start_load_caches(cache) - meta = consumer.build_connector_worker_meta() - assert meta.failed_loads == {"hash"} - assert cache == {} - finally: - consumer.shutdown() From a88a755bf7744a1f8187ad9be2455d18a92c9a4b Mon Sep 17 00:00:00 2001 From: Tianyu Guo Date: Mon, 7 Sep 2026 03:49:35 +0000 Subject: [PATCH 30/30] [EPD] Align Mooncake EC test control addresses Use the same configured host for the consumer bind address and proxy ZMQ endpoint, defaulting to IPv4 loopback. Format endpoints with make_zmq_path to preserve explicit IPv6 support. Validation: bash syntax and applicable pre-commit checks passed. Script argument capture reproduces the previous mismatch; default IPv4, custom IPv4, and IPv6 pass real ZMQ peers/event_port handshakes. Full Buildkite E2E rerun is pending. Co-authored-by: OpenAI Codex Signed-off-by: Tianyu Guo --- .../run_epd_mooncake_ec_full_pipeline.sh | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh index 7b8b24cbecbb..05af1bb6fdbc 100755 --- a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -13,6 +13,7 @@ # MODEL HF model id (default: Qwen/Qwen2.5-VL-3B-Instruct) # GPU_SINGLE / GPU_E / GPU_PD GPU ids (defaults 0 / 0 / 1) # ENDPOINT_PORT, ENCODE_PORT, PREFILL_DECODE_PORT +# EC_MOONCAKE_RESERVATION_HOST consumer control host (default 127.0.0.1) # MOONCAKE_EC_PROTOCOL tcp | rdma (default tcp) # USE_MM_PROMPTS 1 (default) or 0 for text-only quick sanity # TIMEOUT_SECONDS wait_for_server timeout (default 1200) @@ -44,8 +45,10 @@ PREFILL_DECODE_PORT="${PREFILL_DECODE_PORT:-19537}" ENDPOINT_PORT="${ENDPOINT_PORT:-10002}" BASELINE_PORT="${BASELINE_PORT:-10003}" +EC_MOONCAKE_RESERVATION_HOST="${EC_MOONCAKE_RESERVATION_HOST:-127.0.0.1}" EC_MOONCAKE_RESERVATION_PORT="${EC_MOONCAKE_RESERVATION_PORT:-19019}" MOONCAKE_EC_PROTOCOL="${MOONCAKE_EC_PROTOCOL:-tcp}" +export EC_MOONCAKE_RESERVATION_HOST export EC_MOONCAKE_RESERVATION_PORT export MOONCAKE_EC_PROTOCOL if [[ "$MOONCAKE_EC_PROTOCOL" == "tcp" ]]; then @@ -81,7 +84,7 @@ import json, os print(json.dumps({ "ec_connector": "ECMooncakeConnector", "ec_role": "ec_consumer", - "ec_ip": os.environ.get("EC_MOONCAKE_RESERVATION_HOST", "127.0.0.1"), + "ec_ip": os.environ["EC_MOONCAKE_RESERVATION_HOST"], "ec_port": int(os.environ.get("EC_MOONCAKE_RESERVATION_PORT", "19019")), "ec_connector_extra_config": { "mooncake_protocol": os.environ.get("MOONCAKE_EC_PROTOCOL", "tcp"), @@ -90,6 +93,18 @@ print(json.dumps({ PY ) +EC_MOONCAKE_RESERVATION_ADDR=$("$PYTHON_BIN" <<'PY' +import os + +from vllm.utils.network_utils import make_zmq_path + +print(make_zmq_path( + "tcp", os.environ["EC_MOONCAKE_RESERVATION_HOST"], + int(os.environ["EC_MOONCAKE_RESERVATION_PORT"]), +)) +PY +) + wait_for_server() { local port=$1 local pid=$2 @@ -234,7 +249,7 @@ run_epd_mooncake() { --prefill-servers-urls "disable" \ --decode-servers-urls "http://localhost:$PREFILL_DECODE_PORT" \ --ec-consumer-zmq-addrs \ - "tcp://localhost:$EC_MOONCAKE_RESERVATION_PORT" \ + "$EC_MOONCAKE_RESERVATION_ADDR" \ >"${LOG_PATH}/mooncake_epd_proxy.log" 2>&1 & local PROXY_PID=$! PIDS+=("$PROXY_PID")