diff --git a/.github/workflows/test-vllm-e2e.yml b/.github/workflows/test-vllm-e2e.yml new file mode 100644 index 000000000..6afc42c2e --- /dev/null +++ b/.github/workflows/test-vllm-e2e.yml @@ -0,0 +1,111 @@ +name: test-vllm-e2e +permissions: + contents: read +on: + pull_request: + branches: ["main"] + workflow_dispatch: + inputs: + runs-on: + description: "GPU runner label (needs >= 2 GPUs, e.g. A10)" + type: string + default: "gpu-a10-x2" + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: ${{ github.event_name == 'pull_request' }} + +jobs: + vllm-e2e: + name: vllm-e2e (${{ matrix.kind }}) + # GPU runners are self-hosted; the default label can be overridden per-repo + # via the VLLM_E2E_RUNS_ON variable or the workflow_dispatch input. + runs-on: ${{ inputs.runs-on || vars.VLLM_E2E_RUNS_ON || 'gpu-a10-x2' }} + timeout-minutes: 180 + strategy: + fail-fast: false + matrix: + include: + # Full-attention model: single FullAttentionSpec kv cache group. + - kind: full-attention + model_var: VLLM_E2E_MODEL_FULL_ATTN + # Hybrid model: MambaSpec groups + FullAttentionSpec group + # (mamba_cache_mode="align"). + - kind: hybrid-attention + model_var: VLLM_E2E_MODEL_HYBRID + steps: + - uses: actions/checkout@v4 + + - name: check_gpus + run: | + nvidia-smi --query-gpu=index,name,memory.total --format=csv,noheader + GPU_COUNT=$(nvidia-smi --query-gpu=name --format=csv,noheader | wc -l) + if [ "$GPU_COUNT" -lt 2 ]; then + echo "::error::This test needs at least 2 GPUs, found $GPU_COUNT" + exit 1 + fi + + - name: resolve_runner_config + # Self-hosted runner provisioning: a vLLM (>= 0.26.0) venv and the model + # checkouts are pre-cached on the runner and exposed via repo variables: + # VLLM_E2E_PYTHON venv python with vllm installed + # VLLM_E2E_MODEL_FULL_ATTN e.g. a local Qwen2.5-7B-Instruct checkout + # VLLM_E2E_MODEL_HYBRID e.g. a local Qwen3.5-4B checkout + env: + VENV_PY: ${{ vars.VLLM_E2E_PYTHON }} + MODEL: ${{ vars[matrix.model_var] }} + run: | + if [ ! -x "$VENV_PY" ]; then + echo "::error::VLLM_E2E_PYTHON ($VENV_PY) not found; provision the runner" + exit 1 + fi + "$VENV_PY" -c 'import vllm; v = vllm.__version__; print("vllm", v)' + if [ ! -f "$MODEL/config.json" ]; then + echo "::error::model not found at $MODEL; provision the runner" + exit 1 + fi + echo "KVCM_E2E_PYTHON=$VENV_PY" >> "$GITHUB_ENV" + echo "KVCM_E2E_MODEL=$MODEL" >> "$GITHUB_ENV" + + - name: build_binaries + run: | + set -x + bazelisk build //kv_cache_manager:kv_cache_manager_bin \ + //kv_cache_manager/client/pybind:kvcm_py_client_lib_wheel \ + //kv_cache_manager/py_connector/vllm:kvcm_vllm_connector_wheel \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' + + - name: install_wheels + # Bazel wheel filenames contain unstamped {STABLE_*} template variables; + # read the real version from the wheel METADATA and rename before install. + run: | + set -x + mkdir -p /tmp/kvcm_whl && rm -f /tmp/kvcm_whl/*.whl + for whl in bazel-bin/kv_cache_manager/client/pybind/kvcm_py_client-*.whl \ + bazel-bin/kv_cache_manager/py_connector/vllm/kvcm_vllm_connector-*.whl; do + pkg=$(basename "$whl" | sed 's/-{STABLE.*//') + ver=$(unzip -p "$whl" "*.dist-info/METADATA" | awk '/^Version:/{print $2; exit}') + cp "$whl" "/tmp/kvcm_whl/${pkg}-${ver}-cp312-cp312-manylinux_2_32_x86_64.whl" + done + "$KVCM_E2E_PYTHON" -m pip install --no-deps --force-reinstall /tmp/kvcm_whl/*.whl || \ + uv pip install --python "$KVCM_E2E_PYTHON" --no-deps --force-reinstall /tmp/kvcm_whl/*.whl + + - name: run_e2e_tests + run: | + set -x + bazelisk test //integration_test/vllm_e2e/... \ + --cache_test_results=no --test_output=errors \ + --test_env=KVCM_E2E_PYTHON="$KVCM_E2E_PYTHON" \ + --test_env=KVCM_E2E_MODEL="$KVCM_E2E_MODEL" \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' + + - name: upload_logs + if: failure() + uses: actions/upload-artifact@v6 + with: + name: vllm-e2e-logs-${{ matrix.kind }} + path: | + /tmp/kvcm_vllm_e2e/**/*.stdout + /tmp/kvcm_vllm_e2e/**/*.stderr + bazel-out/*-opt/testlogs/integration_test/vllm_e2e/** + if-no-files-found: ignore diff --git a/integration_test/vllm_e2e/BUILD b/integration_test/vllm_e2e/BUILD new file mode 100644 index 000000000..f8ceef86b --- /dev/null +++ b/integration_test/vllm_e2e/BUILD @@ -0,0 +1,143 @@ +package(default_visibility = ["//integration_test:__subpackages__"]) + +# Shared library: orchestration (manager + vLLM + driver + comparison) and the +# verifying connector injected into vLLM via kv_connector_module_path. +py_library( + name = "e2e_lib", + srcs = [ + "e2e_lib.py", + "test_connector.py", + ], + imports = ["."], + tags = ["no-remote-exec"], +) + +py_test( + name = "test_basic", + srcs = ["test_basic.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_concurrent", + srcs = ["test_concurrent.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_tp", + srcs = ["test_tp.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_full_hit", + srcs = ["test_full_hit.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_partial_hit", + srcs = ["test_partial_hit.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_load_failure", + srcs = ["test_load_failure.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_multi_turn", + srcs = ["test_multi_turn.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +# Meta-test: injects an off-by-one into the connector's token translation and +# asserts the KV verification FAILS -- proof the harness is not vacuous. +py_test( + name = "test_mutation", + srcs = ["test_mutation.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) diff --git a/integration_test/vllm_e2e/README.md b/integration_test/vllm_e2e/README.md new file mode 100644 index 000000000..40c82f624 --- /dev/null +++ b/integration_test/vllm_e2e/README.md @@ -0,0 +1,102 @@ +# vLLM <-> KVCM End-to-End KV Cache Verification + +End-to-end integration tests for the KVCM vLLM connector +(`kv_cache_manager/py_connector/vllm`). Each test starts a real KVCM manager +(local-file storage backend) and a real vLLM OpenAI server, drives prompts +through the OpenAI API and verifies that the KV cache data saved to / loaded +from KVCM is correct. + +Requires 1-2 GPUs and vLLM >= 0.26.0. + +## What is verified + +The connector translates between three block spaces per `kv_cache_group`: + +``` +KVCM manager block idx -> global token idx -> group logical block + (step 1, connector-only) (step 2/3, shared with vLLM) +``` + +A bug in step 1 is *symmetric*: save gathers from the wrong slots and load +scatters back to the same wrong slots, so a transport round trip alone cannot +detect it. The test breaks the symmetry with `VerifyingConnector` +(`test_connector.py`), a subclass of the production connector that +independently captures KV data from vLLM's paged cache using only vLLM's own +block-table mapping: + +1. **Phase 1** — fresh prompts: prefill -> connector saves to KVCM. The saved + token ranges are captured from the paged cache (**reference** captures). +2. **Phase 2** — same prompts + suffix: connector reports an external match and + loads from KVCM. The loaded blocks are captured (**loaded** captures). +3. The driver (`e2e_lib.py`) matches loaded captures against references by + token content and compares per layer: bit-exact preferred, cosine + similarity > 99.99% as fallback. + +## Model coverage + +The same test targets run against either model kind, selected by +`KVCM_E2E_MODEL`: + +| Kind | Example | Groups | Orchestration | +|---|---|---|---| +| Full attention | Qwen2.5-7B-Instruct | 1 `FullAttentionSpec` | prefix caching off, one server for both phases | +| Hybrid | Qwen3.5-4B | 3 `MambaSpec` + 1 `FullAttentionSpec` | prefix caching on (`mamba_cache_mode="align"`), server restarted between phases so phase 2 loads from KVCM instead of the local prefix cache | + +Hybrid specifics verified: + +* Per-group location specs (`tp{rank}_g{group}`) and per-group block tables. +* Attention groups: token-granular gather/scatter through the Triton kernel. +* Mamba/linear groups: per-block opaque state copy, where a manager block's + *last* token selects the state block (`_state_block_ids`). + +## Scenarios + +| Test | TP | Prompts | Notes | +|---|---|---|---| +| `test_basic` | 1 | 1 | Minimal save -> load round trip | +| `test_concurrent` | 1 | 4 | Concurrent requests: ReqState tracking, per-request block attribution | +| `test_tp` | 2 | 2 | TP coordination; for full-attention models also `preferred_block_size=32` != vLLM block size (16), forcing real cross-block translation | + +## Running + +Build prerequisites (from the repo root): + +```bash +bazelisk build //kv_cache_manager:kv_cache_manager_bin \ + //kv_cache_manager/client/pybind:kvcm_py_client_lib_wheel \ + //kv_cache_manager/py_connector/vllm:kvcm_vllm_connector_wheel \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' +``` + +Install both wheels into the vLLM venv (rename them first: the Bazel output +name contains unstamped `{STABLE_*}` template variables; read the real version +from the wheel's `METADATA`). + +Run (tagged `exclusive`, so they execute serially): + +```bash +bazelisk test //integration_test/vllm_e2e/... \ + --cache_test_results=no --test_output=errors \ + --test_env=KVCM_E2E_PYTHON=/path/to/vllm-venv/bin/python \ + --test_env=KVCM_E2E_MODEL=/path/to/model \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' +``` + +Environment variables: + +| Variable | Meaning | +|---|---| +| `KVCM_E2E_PYTHON` | Python interpreter with vLLM + both KVCM wheels installed | +| `KVCM_E2E_MODEL` | Model path; hybrid models are auto-detected from `config.json` | + +## Debugging + +Bazel's `test.log` only shows the driver's view (e.g. HTTP 500). The real +tracebacks live in the scenario workdir under `$TEST_TMPDIR`: + +``` +/kvcm_vllm_e2e// + manager/manager.stdout|stderr # KVCM manager + vllm/vllm*.stdout|stderr # vLLM (EngineCore tracebacks are here) + captures/{ref|loaded}_tp{rank}_{token_hash}.pt +``` diff --git a/integration_test/vllm_e2e/e2e_lib.py b/integration_test/vllm_e2e/e2e_lib.py new file mode 100644 index 000000000..29190eda8 --- /dev/null +++ b/integration_test/vllm_e2e/e2e_lib.py @@ -0,0 +1,833 @@ +"""Orchestration for the KVCM <-> vLLM end-to-end KV cache verification test. + +This module is imported by the Bazel ``py_test`` targets. It: + +1. Starts a KVCM manager (``kv_cache_manager_bin``) with a local-file storage + backend. +2. Starts a vLLM OpenAI server configured with the ``VerifyingConnector`` + (injected via ``kv_connector_module_path`` -- no vLLM files are modified). +3. Drives prompts through the OpenAI API in two phases and compares the + independently-captured KV data (reference from the save path vs loaded from + the load path). + +The comparison is done on KV data captured from vLLM's paged cache using vLLM's +own block-table mapping (independent of the connector's per-group translation), +which is what makes the test able to detect symmetric save/load translation +bugs. See ``test_connector.py`` for the capture-side details. + +Full-attention vs hybrid models +------------------------------- +The same test targets run against either model, selected by ``$KVCM_E2E_MODEL``: + +* Full-attention (e.g. Qwen2.5): a single ``FullAttentionSpec`` group. Prefix + caching is disabled so every phase-2 request is served through the connector + (no local prefix hit); the two phases share one server. +* Hybrid (e.g. Qwen3.5): several ``MambaSpec`` groups plus a ``FullAttentionSpec`` + group. Prefix caching must be enabled for vLLM to produce per-group block + tables (``mamba_cache_mode="align"``). Because that also populates the local + prefix cache, the vLLM server is restarted between the two phases so phase 2 + genuinely loads from KVCM instead of hitting the local cache. +""" + +import glob +import json +import logging +import os +import shutil +import socket +import subprocess +import time +import uuid +from typing import Optional + +import requests + +logger = logging.getLogger("vllm_e2e") +# This module only runs inside test drivers; make the orchestration evidence +# (block counts, verification report, log-scan results) visible in test.log. +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(name)s %(levelname)s %(message)s") + +MODEL_PATH = os.environ.get("KVCM_E2E_MODEL", "/root/ws/resources/models/Qwen2.5-7B-Instruct") +COSINE_THRESHOLD = 0.9999 +# Bit-exact comparison is the default (empirically all scenarios achieve it). +# Cosine fallback must be explicitly requested. +ALLOW_COSINE = os.environ.get("KVCM_E2E_ALLOW_COSINE", "0") == "1" + + +def is_hybrid_model(model_path: str) -> bool: + """Detect a hybrid (mamba/linear + full attention) model from its config.""" + try: + with open(os.path.join(model_path, "config.json")) as f: + cfg = json.load(f) + except Exception: + return False + text_cfg = cfg.get("text_config", cfg) + # Hybrid models interleave linear/mamba layers with full attention and + # expose a full_attention_interval / linear_* knob. + return ( + "full_attention_interval" in text_cfg + or "linear_conv_kernel_dim" in text_cfg + or cfg.get("model_type", "").startswith("qwen3_5") + ) + + +# --------------------------------------------------------------------------- # +# Paths / binaries +# --------------------------------------------------------------------------- # +def _runfiles_root() -> Optional[str]: + return os.environ.get("RUNFILES_DIR") or os.environ.get("TEST_SRCDIR") + + +def find_repo_root() -> str: + """Locate the KVCM repository root (works under Bazel runfiles and plain).""" + here = os.path.dirname(os.path.abspath(__file__)) + # integration_test/vllm_e2e/e2e_lib.py -> repo root is two levels up. + candidate = os.path.abspath(os.path.join(here, "..", "..")) + if os.path.exists(os.path.join(candidate, "WORKSPACE")): + return candidate + runfiles = _runfiles_root() + if runfiles: + cand = os.path.join(runfiles, "kv_cache_manager") + if os.path.exists(os.path.join(cand, "WORKSPACE")): + return cand + return candidate + + +def find_manager_binary(repo_root: str) -> str: + candidates = [ + os.path.join(repo_root, "bazel-bin/kv_cache_manager/kv_cache_manager_bin"), + os.path.join(repo_root, "bazel-out/k8-opt/bin/kv_cache_manager/kv_cache_manager_bin"), + ] + runfiles = _runfiles_root() + if runfiles: + candidates.append( + os.path.join(runfiles, "kv_cache_manager", "kv_cache_manager", + "kv_cache_manager_bin") + ) + for c in candidates: + if os.path.exists(c): + return c + raise RuntimeError( + "kv_cache_manager_bin not found; build it with: " + "bazelisk build //kv_cache_manager:kv_cache_manager_bin" + ) + + +def find_python() -> str: + return os.environ.get("KVCM_E2E_PYTHON", "/root/ws/env/global_vllm/.venv/bin/python") + + +def free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("", 0)) + return s.getsockname()[1] + + +def wait_http(url: str, timeout: float, post_body: Optional[dict] = None) -> bool: + deadline = time.time() + timeout + while time.time() < deadline: + try: + if post_body is not None: + r = requests.post(url, json=post_body, timeout=3) + else: + r = requests.get(url, timeout=3) + if r.status_code < 500: + return True + except Exception: + pass + time.sleep(1.0) + return False + + +# --------------------------------------------------------------------------- # +# KVCM manager +# --------------------------------------------------------------------------- # +class ManagerProcess: + def __init__(self, workdir: str, storage_root: str, key_count_per_file: int = 8): + self.workdir = workdir + os.makedirs(workdir, exist_ok=True) + self.rpc_port = free_port() + self.http_port = free_port() + self.admin_rpc_port = free_port() + self.admin_http_port = free_port() + self.storage_root = storage_root + self.key_count_per_file = key_count_per_file + self.proc: Optional[subprocess.Popen] = None + self.config_path = os.path.join(workdir, "startup_config.json") + + def manager_uri(self) -> str: + return f"http://127.0.0.1:{self.http_port}" + + def _write_config(self): + cfg = { + "storage_config": { + "type": "file", + "global_unique_name": "nfs_01", + "storage_spec": { + # The backend concatenates root_path + key with no + # separator; the trailing slash keeps files inside the dir. + "root_path": self.storage_root.rstrip("/") + "/", + "key_count_per_file": self.key_count_per_file, + }, + }, + "instance_group": { + "name": "default", + "storage_candidates": ["nfs_01"], + "global_quota_group_name": "default_quota_group", + "max_instance_count": 100, + "quota": { + "capacity": 30000000000, + "quota_config": [ + {"storage_type": "file", "capacity": 10000000000}, + {"storage_type": "hf3fs", "capacity": 10000000000}, + {"storage_type": "pace", "capacity": 10000000000}, + ], + }, + "cache_config": { + "reclaim_strategy": { + "reclaim_policy": 1, + "trigger_strategy": {"used_percentage": 0.8}, + "delay_before_delete_ms": 1000, + }, + "cache_prefer_strategy": 2, + "meta_indexer_config": { + "max_key_count": 1000000, + "mutex_shard_num": 16, + "batch_key_size": 16, + "meta_storage_backend_config": { + "storage_type": "local", + "storage_uri": "", + }, + "meta_cache_policy_config": { + "type": "LRU", + "capacity": 10000, + "cache_shard_bits": 0, + "high_pri_pool_ratio": 0.0, + }, + }, + }, + "user_data": '{"description": "vllm e2e test instance group"}', + "version": 1, + }, + } + with open(self.config_path, "w") as f: + json.dump(cfg, f, indent=2) + + def start(self, repo_root: str): + self._write_config() + binary = find_manager_binary(repo_root) + cmd = [ + binary, + "--env", f"kvcm.service.rpc_port={self.rpc_port}", + "--env", f"kvcm.service.http_port={self.http_port}", + "--env", f"kvcm.service.admin_rpc_port={self.admin_rpc_port}", + "--env", f"kvcm.service.admin_http_port={self.admin_http_port}", + "--env", f"kvcm.startup_config={self.config_path}", + "--env", "kvcm.logger.log_level=5", + ] + logger.info("starting manager: %s (cwd=%s)", " ".join(cmd), self.workdir) + self.proc = subprocess.Popen( + cmd, + cwd=self.workdir, + stdout=open(os.path.join(self.workdir, "manager.stdout"), "w"), + stderr=open(os.path.join(self.workdir, "manager.stderr"), "w"), + ) + if not wait_http( + f"{self.manager_uri()}/api/getClusterInfo", + timeout=60, + post_body={"trace_id": "probe", "instance_id": "probe"}, + ): + raise RuntimeError("manager did not become ready; see manager.stderr") + logger.info("manager ready at %s", self.manager_uri()) + + def stop(self): + if self.proc and self.proc.poll() is None: + self.proc.terminate() + try: + self.proc.wait(timeout=10) + except subprocess.TimeoutExpired: + self.proc.kill() + + +# --------------------------------------------------------------------------- # +# vLLM server +# --------------------------------------------------------------------------- # +class VllmServer: + def __init__(self, workdir: str, capture_dir: str, manager_uri: str, + tp_size: int, coordinator_base_port: int, + instance_id: str, preferred_block_size: int, + enable_prefix_caching: bool, + connector_name: str = "VerifyingConnector", + log_level: str = "INFO", + extra_config_overrides: Optional[dict] = None, + kv_load_failure_policy: Optional[str] = None): + self.workdir = workdir + os.makedirs(workdir, exist_ok=True) + self.capture_dir = capture_dir + os.makedirs(capture_dir, exist_ok=True) + self.port = free_port() + self.manager_uri = manager_uri + self.tp_size = tp_size + self.coordinator_base_port = coordinator_base_port + self.instance_id = instance_id + self.preferred_block_size = preferred_block_size + self.enable_prefix_caching = enable_prefix_caching + self.connector_name = connector_name + self.log_level = log_level + self.extra_config_overrides = extra_config_overrides or {} + self.kv_load_failure_policy = kv_load_failure_policy + self.proc: Optional[subprocess.Popen] = None + + def base_url(self) -> str: + return f"http://127.0.0.1:{self.port}" + + def start(self, repo_root: str, log_suffix: str = ""): + extra_config = { + "manager_uri": self.manager_uri, + "coordinator_base_port": self.coordinator_base_port, + "instance_group": "default", + "instance_id": self.instance_id, + "preferred_block_size": self.preferred_block_size, + "log_level": self.log_level, + } + extra_config.update(self.extra_config_overrides) + kv_transfer_config = { + "kv_connector": self.connector_name, + "kv_role": "kv_both", + "kv_connector_module_path": "test_connector", + "kv_connector_extra_config": extra_config, + } + if self.kv_load_failure_policy: + kv_transfer_config["kv_load_failure_policy"] = self.kv_load_failure_policy + cmd = [ + find_python(), "-m", "vllm.entrypoints.openai.api_server", + "--model", MODEL_PATH, + "--served-model-name", "qwen", + "--port", str(self.port), + "--tensor-parallel-size", str(self.tp_size), + "--max-model-len", "4096", + "--gpu-memory-utilization", "0.85", + "--enforce-eager", + "--max-num-seqs", "16", + "--kv-transfer-config", json.dumps(kv_transfer_config), + ] + if self.enable_prefix_caching: + # Hybrid models need prefix caching to expose per-group block tables + # (mamba_cache_mode="align"); align mode requires chunked prefill. + cmd += ["--enable-prefix-caching", "--enable-chunked-prefill"] + else: + cmd += ["--no-enable-prefix-caching"] + env = os.environ.copy() + env["PYTHONPATH"] = os.path.dirname(os.path.abspath(__file__)) + os.pathsep + env.get("PYTHONPATH", "") + env["KVCM_E2E_CAPTURE_DIR"] = self.capture_dir + # Keep the connector's KV cache layout matching its expected + # [2, num_blocks, block_size, num_kv_heads, head_size] shape. + env.setdefault("VLLM_KV_CACHE_LAYOUT", "NHD") + # Force FlashAttention for the full-attention layers: it produces the + # [2, num_blocks, block_size, num_kv_heads, head_size] layout the + # connector expects, and avoids the flashinfer backend entirely. + env.setdefault("VLLM_ATTENTION_BACKEND", "FLASH_ATTN") + # Use the PyTorch-native sampler; the flashinfer sampler JIT-compiles + # with ninja, which is not available in the test environment. + env.setdefault("VLLM_USE_FLASHINFER_SAMPLER", "0") + # The test venv ships mismatched flashinfer / flashinfer-cubin wheels; + # skip the version check so importing vLLM's attention registry does not + # crash before the FlashAttention backend is selected. + env.setdefault("FLASHINFER_DISABLE_VERSION_CHECK", "1") + logger.info("starting vllm: %s", " ".join(cmd)) + self.proc = subprocess.Popen( + cmd, + cwd=self.workdir, + env=env, + stdout=open(os.path.join(self.workdir, f"vllm{log_suffix}.stdout"), "w"), + stderr=open(os.path.join(self.workdir, f"vllm{log_suffix}.stderr"), "w"), + ) + if not wait_http(f"{self.base_url()}/health", timeout=600): + raise RuntimeError("vllm did not become ready; see vllm.stderr") + logger.info("vllm ready at %s", self.base_url()) + + def stop(self): + if self.proc and self.proc.poll() is None: + self.proc.terminate() + try: + self.proc.wait(timeout=15) + except subprocess.TimeoutExpired: + self.proc.kill() + + +# --------------------------------------------------------------------------- # +# Request driver +# --------------------------------------------------------------------------- # +def tokenize(base_url: str, prompt: str) -> list[int]: + r = requests.post(f"{base_url}/tokenize", + json={"model": "qwen", "prompt": prompt}, timeout=60) + r.raise_for_status() + return r.json()["tokens"] + + +def get_manager_block_size(manager_uri: str, instance_id: str) -> int: + """Ask the manager for the registered instance's manager block size.""" + r = requests.post(f"{manager_uri}/api/getInstanceInfo", + json={"trace_id": "e2e_bs", "instance_id": instance_id}, + timeout=10) + r.raise_for_status() + block_size = r.json()["instance_info"]["block_size"] + assert block_size > 0, f"bad manager block size: {block_size}" + return block_size + + +def block_token_hash(token_ids: list[int]) -> str: + """Token-content hash used in capture file names (mirrors test_connector).""" + import hashlib + import torch + return hashlib.sha256( + torch.tensor(token_ids, dtype=torch.int64).numpy().tobytes() + ).hexdigest()[:16] + + +def full_block_hashes(token_ids: list[int], manager_block_size: int) -> list[str]: + """Per-manager-block capture hashes for the full blocks of a token stream.""" + n = len(token_ids) // manager_block_size + return [ + block_token_hash(token_ids[i * manager_block_size:(i + 1) * manager_block_size]) + for i in range(n) + ] + + +def wait_for_prefix_cached(manager_uri: str, instance_id: str, + token_ids: list[int], min_blocks: int, + timeout: float = 120.0) -> bool: + """Poll the manager until at least min_blocks of the prefix are committed. + + Mirrors what the connector's get_num_new_matched_tokens queries + (query_type=QT_PREFIX_MATCH, block_mask offset=0 for a fresh request), so it + guarantees phase 2 will actually hit the external cache for all min_blocks. + """ + deadline = time.time() + timeout + payload = { + "trace_id": "e2e_probe", + "token_ids": token_ids, + "instance_id": instance_id, + "query_type": "QT_PREFIX_MATCH", + "block_mask": {"offset": 0}, + } + while time.time() < deadline: + try: + r = requests.post(f"{manager_uri}/api/getCacheLocation", + json=payload, timeout=10) + if r.status_code == 200: + data = r.json() + if data.get("header", {}).get("status", {}).get("code") == "OK": + locs = data.get("locations", []) + if len(locs) >= min_blocks: + logger.info("prefix cached: %d location(s)", len(locs)) + return True + except Exception: + pass + time.sleep(1.0) + logger.warning("timed out waiting for prefix to be cached") + return False + + +def send_completions(base_url: str, prompts: list, max_tokens: int = 4, + temperature: float = 0.0, **extra_payload) -> list[dict]: + """Send prompts concurrently and return the OpenAI responses. + + Each prompt may be a string or a list of token ids (the completions API + accepts both). extra_payload is merged into the request body (e.g. + return_token_ids=True).""" + from concurrent.futures import ThreadPoolExecutor + + client_url = f"{base_url}/v1/completions" + + def _one(prompt) -> dict: + payload = { + "model": "qwen", + "prompt": prompt, + "max_tokens": max_tokens, + "temperature": temperature, + **extra_payload, + } + r = requests.post(client_url, json=payload, timeout=300) + r.raise_for_status() + return r.json() + + with ThreadPoolExecutor(max_workers=max(1, len(prompts))) as ex: + return list(ex.map(_one, prompts)) + + +# --------------------------------------------------------------------------- # +# Capture comparison +# --------------------------------------------------------------------------- # +def count_captures(capture_dir: str, kind: str) -> int: + return len(glob.glob(os.path.join(capture_dir, f"{kind}_*.pt"))) + + +def wait_for_captures(capture_dir: str, kind: str, expected: int, + timeout: float = 120.0) -> int: + """Wait until at least ``expected`` captures of ``kind`` exist. + + Raises AssertionError on timeout: a missing capture means the connector + never exercised the code path under test, so the scenario must fail rather + than silently verify fewer blocks. + """ + deadline = time.time() + timeout + while time.time() < deadline: + n = count_captures(capture_dir, kind) + if n >= expected: + logger.info("saw %d/%d %s captures", n, expected, kind) + return n + time.sleep(1.0) + n = count_captures(capture_dir, kind) + raise AssertionError( + f"timed out waiting for {kind} captures: got {n}, want {expected}") + + +def _cosine(a, b) -> float: + import torch + a = a.reshape(-1).float() + b = b.reshape(-1).float() + denom = (a.norm() * b.norm()).clamp_min(1e-12) + return float((a @ b) / denom) + + +def compare_captures(capture_dir: str, tp_size: int) -> dict: + """Compare loaded captures against reference captures. + + Every block that was *loaded* from KVCM must correspond to a *reference* + capture (same tp rank + token content) with matching KV data. The direction + matters: saves are incremental, so some saved blocks may legitimately not be + reloaded (e.g. the tokenization boundary block) -- but every loaded block + must match something that was saved. + + Returns a report dict; the caller asserts on it. + """ + import torch + + refs = {} + loaded = {} + for path in glob.glob(os.path.join(capture_dir, "*.pt")): + name = os.path.basename(path)[:-3] # strip .pt + parts = name.split("_") + kind, tp, token_hash = parts[0], parts[1], "_".join(parts[2:]) + key = (tp, token_hash) + (refs if kind == "ref" else loaded)[key] = path + + report = { + "num_refs": len(refs), + "num_loaded": len(loaded), + "matched": 0, + "bit_exact": 0, + "cosine_pass": 0, + "failures": [], + "loaded_without_ref": [], + "matched_keys": [], + } + + for key, loaded_path in sorted(loaded.items()): + if key not in refs: + report["loaded_without_ref"].append(key) + continue + ref = torch.load(refs[key], map_location="cpu", weights_only=True) + got = torch.load(loaded_path, map_location="cpu", weights_only=True) + + assert ref["token_ids"] == got["token_ids"], f"token id mismatch for {key}" + + all_bit_exact = True + worst_cosine = 1.0 + # Compare every layer present in the *loaded* capture: each one was + # actually written by the connector and must match its reference. + # Mamba "align" state layers can legitimately be absent on either side + # (vLLM materializes states only at segment boundaries; interior blocks + # get the null block and the connector transfers them vacuously) -- but + # a loaded layer without a reference is a hard error. + for layer_name, got_kv in got["kv"].items(): + assert layer_name in ref["kv"], ( + f"loaded layer {layer_name} of {key} has no reference capture") + ref_kv = ref["kv"][layer_name] + # Attention groups are a single Tensor; mamba/linear/gdn groups are a + # list[Tensor] (e.g. [conv_state, ssm_state]). Compare uniformly. + if isinstance(ref_kv, (list, tuple)): + ref_parts = list(ref_kv) + got_parts = list(got_kv) + assert len(ref_parts) == len(got_parts), ( + f"state count mismatch {layer_name}: " + f"{len(ref_parts)} vs {len(got_parts)}" + ) + else: + ref_parts = [ref_kv] + got_parts = [got_kv] + + for si, (ref_t, got_t) in enumerate(zip(ref_parts, got_parts)): + assert ref_t.shape == got_t.shape, ( + f"shape mismatch {layer_name}[{si}]: {ref_t.shape} vs {got_t.shape}" + ) + if not torch.equal(ref_t, got_t): + all_bit_exact = False + cos = _cosine(ref_t, got_t) + worst_cosine = min(worst_cosine, cos) + if not ALLOW_COSINE or cos < COSINE_THRESHOLD: + report["failures"].append({ + "key": key, + "layer": f"{layer_name}[{si}]", + "cosine": cos, + }) + + report["matched"] += 1 + report["matched_keys"].append(key) + if all_bit_exact: + report["bit_exact"] += 1 + else: + report["cosine_pass"] += 1 + logger.warning("capture %s not bit-exact (worst cosine=%.6f)", + key, worst_cosine) + + return report + + +def assert_report_ok(report: dict, min_matched: int = 1): + """Assert the comparison succeeded. + + min_matched is the exact lower bound of (ref, loaded) capture pairs computed + from the prompts' tokenization (num full manager blocks x tp ranks); a lower + count means some blocks were silently never saved or never loaded. + """ + problems = [] + if report["loaded_without_ref"]: + problems.append( + f"loaded captures with no matching reference: {report['loaded_without_ref']}" + ) + if report["failures"]: + kind = "cosine" if ALLOW_COSINE else "bit-exact" + problems.append(f"{kind} failures: {report['failures']}") + if report["matched"] < min_matched: + problems.append( + f"matched {report['matched']} loaded captures, expected >= {min_matched}" + ) + if problems: + raise AssertionError("KV verification failed: " + "; ".join(problems)) + logger.info( + "KV verification OK: matched=%d bit_exact=%d cosine_pass=%d (refs=%d loaded=%d)", + report["matched"], report["bit_exact"], report["cosine_pass"], + report["num_refs"], report["num_loaded"], + ) + + +# --------------------------------------------------------------------------- # +# Scenario runner +# --------------------------------------------------------------------------- # +class ScenarioEnv: + """Owns one scenario's manager + vLLM server lifecycle and scratch dirs. + + Custom scenarios (partial-hit, full-hit, load-failure, multi-turn) share + this; run_e2e keeps its own two-phase flow on top of the same pieces. + """ + + def __init__(self, scenario: str, tp_size: int = 1, + preferred_block_size: int = 0, + enable_prefix_caching: Optional[bool] = None, + connector_name: str = "VerifyingConnector", + log_level: str = "INFO", + extra_config_overrides: Optional[dict] = None, + key_count_per_file: int = 8, + kv_load_failure_policy: Optional[str] = None): + self.scenario = scenario + self.tp_size = tp_size + self.hybrid = is_hybrid_model(MODEL_PATH) + self.preferred_block_size = 0 if self.hybrid else preferred_block_size + self.enable_prefix_caching = (self.hybrid if enable_prefix_caching is None + else enable_prefix_caching) + self.connector_name = connector_name + self.log_level = log_level + self.extra_config_overrides = extra_config_overrides + self.kv_load_failure_policy = kv_load_failure_policy + + self.repo_root = find_repo_root() + scratch_root = (os.environ.get("TEST_TMPDIR") + or os.environ.get("TMPDIR") or "/tmp") + self.base_workdir = os.path.join(scratch_root, "kvcm_vllm_e2e", scenario) + if os.path.exists(self.base_workdir): + shutil.rmtree(self.base_workdir) + self.storage_root = os.path.join(self.base_workdir, "nfs") + self.capture_dir = os.path.join(self.base_workdir, "captures") + self.vllm_dir = os.path.join(self.base_workdir, "vllm") + os.makedirs(self.storage_root, exist_ok=True) + + self.instance_id = f"e2e-{scenario}-{uuid.uuid4().hex[:8]}" + self.manager = ManagerProcess( + os.path.join(self.base_workdir, "manager"), self.storage_root, + key_count_per_file=key_count_per_file) + self.vllm: Optional[VllmServer] = None + + def start_manager(self): + self.manager.start(self.repo_root) + + def start_vllm(self, log_suffix: str = "") -> VllmServer: + self.vllm = VllmServer( + self.vllm_dir, self.capture_dir, self.manager.manager_uri(), + self.tp_size, coordinator_base_port=free_port(), + instance_id=self.instance_id, + preferred_block_size=self.preferred_block_size, + enable_prefix_caching=self.enable_prefix_caching, + connector_name=self.connector_name, + log_level=self.log_level, + extra_config_overrides=self.extra_config_overrides, + kv_load_failure_policy=self.kv_load_failure_policy, + ) + self.vllm.start(self.repo_root, log_suffix=log_suffix) + return self.vllm + + def restart_vllm(self, log_suffix: str = "") -> VllmServer: + """Restart vLLM to drop its local prefix cache (KVCM state persists).""" + if self.vllm: + self.vllm.stop() + return self.start_vllm(log_suffix) + + def manager_block_size(self) -> int: + return get_manager_block_size(self.manager.manager_uri(), self.instance_id) + + def scan_connector_logs(self, pattern: str) -> list: + """Regex-scan all vLLM std streams; returns list of match groups.""" + import re + out = [] + for path in glob.glob(os.path.join(self.vllm_dir, "vllm*.std*")): + with open(path, errors="replace") as f: + for line in f: + m = re.search(pattern, line) + if m: + out.append(m.groups() if m.groups() else m.group(0)) + return out + + def stop(self): + if self.vllm: + self.vllm.stop() + self.manager.stop() + + +def make_base_prompts(num_prompts: int, hybrid: bool) -> list[str]: + """Distinct, deterministic prompts. Each sentence carries a unique counter + so every manager block has unique token content -- this avoids hash + collisions between blocks with identical text but different KV (RoPE is + position-dependent). + + Hybrid models pin the manager block size to the scheduler block size (528), + so their prompts must be much longer to span multiple manager blocks + (empirically 140 sentences ~ 2100 tokens > 3 x 528). Full-attention models + use manager blocks of 16/32 tokens, where 40 sentences (~580 tokens) already + span dozens of blocks. + """ + num_sentences = 140 if hybrid else 40 + return [ + f"Prompt number {i}. " + " ".join( + f"Sentence {j} of prompt {i} has value {j * 7 + i * 131}." + for j in range(num_sentences) + ) + for i in range(num_prompts) + ] + + +def shared_token_prefix_len(a: list[int], b: list[int]) -> int: + n = 0 + for x, y in zip(a, b): + if x != y: + break + n += 1 + return n + + +def run_e2e(scenario: str, tp_size: int, num_prompts: int, + preferred_block_size: int, connector_name: str = "VerifyingConnector", + expect_verification_failure: bool = False): + """Run one full save-then-load verification scenario. + + Full-attention models: prefix caching off, one server across both phases. + Hybrid models: prefix caching on (align mode -> per-group block tables), the + vLLM server is restarted between phases so phase 2 loads from KVCM instead of + hitting the local prefix cache. + + connector_name selects the connector class inside test_connector.py; the + mutation meta-test passes "MutatedConnector" and sets + expect_verification_failure=True to prove the harness detects an injected + off-by-one in the token translation. + """ + import torch # noqa: F401 (ensure torch importable early for clear errors) + + env = ScenarioEnv(scenario, tp_size=tp_size, + preferred_block_size=preferred_block_size, + connector_name=connector_name) + hybrid = env.hybrid + capture_dir = env.capture_dir + logger.info("scenario=%s model=%s hybrid=%s tp=%d prompts=%d preferred_bs=%d", + scenario, MODEL_PATH, hybrid, tp_size, num_prompts, + env.preferred_block_size) + + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="" if not hybrid else "_p1") + + base_prompts = make_base_prompts(num_prompts, hybrid) + suffixes = [f" Now answer question {i}: what is 2+2?" for i in range(num_prompts)] + phase2_prompts = [p + s for p, s in zip(base_prompts, suffixes)] + + # Compute per-prompt expected block counts from the actual tokenization + # so a silently dropped prompt (or block) fails the run. + mbs = env.manager_block_size() + base_tokens = [tokenize(vllm.base_url(), p) for p in base_prompts] + phase2_tokens = [tokenize(vllm.base_url(), p) for p in phase2_prompts] + expected_save_blocks = [len(t) // mbs for t in base_tokens] + # Loads cover the shared token prefix (the base/suffix boundary token may + # re-merge under tokenization, shortening the shared prefix by one). + expected_load_blocks = [ + min(shared_token_prefix_len(b, p2) // mbs, s) + for b, p2, s in zip(base_tokens, phase2_tokens, expected_save_blocks) + ] + assert all(n >= 1 for n in expected_load_blocks), ( + f"prompts too short to span a manager block (mbs={mbs}): " + f"{expected_load_blocks}") + logger.info("mbs=%d expected save blocks=%s load blocks=%s", + mbs, expected_save_blocks, expected_load_blocks) + + # ---- Phase 1: fresh prefill -> connector saves -> reference capture. + logger.info("phase 1: sending %d fresh prompts", num_prompts) + send_completions(vllm.base_url(), base_prompts) + wait_for_captures(capture_dir, "ref", + expected=tp_size * sum(expected_save_blocks), timeout=180) + + # The save is committed to the manager asynchronously after the ref + # capture (which fires when the save is submitted). Wait until the manager + # actually has the prefix, otherwise phase 2 would find no match. + for toks, blocks in zip(base_tokens, expected_save_blocks): + if not wait_for_prefix_cached(env.manager.manager_uri(), env.instance_id, + toks, min_blocks=blocks): + raise AssertionError("save was not committed to the manager in time") + + # Hybrid models keep prefix caching on, which also populates the local + # prefix cache; restart vLLM so phase 2 loads from KVCM, not locally. + if hybrid: + logger.info("restarting vLLM before phase 2 (clear local prefix cache)") + vllm = env.restart_vllm(log_suffix="_p2") + + # ---- Phase 2: same prefix + suffix -> connector loads -> loaded capture. + logger.info("phase 2: sending %d prefix+suffix prompts", num_prompts) + send_completions(vllm.base_url(), phase2_prompts) + wait_for_captures(capture_dir, "loaded", + expected=tp_size * sum(expected_load_blocks), timeout=180) + + report = compare_captures(capture_dir, tp_size) + min_matched = tp_size * sum(expected_load_blocks) + if expect_verification_failure: + try: + assert_report_ok(report, min_matched=min_matched) + except AssertionError as e: + logger.info("verification failed as expected: %s", e) + return + raise AssertionError( + "mutated connector passed KV verification; the harness is blind") + assert_report_ok(report, min_matched=min_matched) + logger.info("scenario %s PASSED: %s", scenario, json.dumps( + {k: v for k, v in report.items() if k not in ("failures", "matched_keys")}, + default=str)) + finally: + env.stop() diff --git a/integration_test/vllm_e2e/test_basic.py b/integration_test/vllm_e2e/test_basic.py new file mode 100644 index 000000000..a409d12f6 --- /dev/null +++ b/integration_test/vllm_e2e/test_basic.py @@ -0,0 +1,27 @@ +"""test_basic: single-request save/load KV verification (TP=1). + +Sends one prompt (prefill + save -> reference capture), then the same prompt +with a suffix (load + prefill -> loaded capture), and verifies the prefix KV +data matches (bit-exact preferred, cosine > 99.99% as fallback). + +Works for both full-attention and hybrid models (selected via KVCM_E2E_MODEL); +see e2e_lib.run_e2e for the per-model orchestration differences. +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestBasic(unittest.TestCase): + def test_basic(self): + run_e2e( + scenario="basic", + tp_size=1, + num_prompts=1, + preferred_block_size=0, # manager block size == vllm block size + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_concurrent.py b/integration_test/vllm_e2e/test_concurrent.py new file mode 100644 index 000000000..4ea7501de --- /dev/null +++ b/integration_test/vllm_e2e/test_concurrent.py @@ -0,0 +1,28 @@ +"""test_concurrent: multiple concurrent requests save/load KV verification. + +Sends several distinct prompts concurrently (all prefill + save -> reference +captures), then the same prompts each with their own suffix concurrently (all +load + prefill -> loaded captures), and verifies each request's prefix KV data +matches. This exercises ReqState tracking, per-request block attribution and +async task races. + +Works for both full-attention and hybrid models (selected via KVCM_E2E_MODEL). +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestConcurrent(unittest.TestCase): + def test_concurrent(self): + run_e2e( + scenario="concurrent", + tp_size=1, + num_prompts=4, + preferred_block_size=0, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_connector.py b/integration_test/vllm_e2e/test_connector.py new file mode 100644 index 000000000..4d5055306 --- /dev/null +++ b/integration_test/vllm_e2e/test_connector.py @@ -0,0 +1,315 @@ +"""A verification wrapper around the production KVCM vLLM connector. + +This connector is injected via vLLM's ``kv_connector_module_path`` and subclasses +the production ``TairKvCacheConnector`` without modifying it. Its purpose is to +independently capture the KV data that lives in vLLM's paged KV cache so that the +test driver can verify the connector's save/load translation layer. + +Why this catches translation bugs +--------------------------------- +The production connector is built around *per-group* transfer. Every +``kv_cache_group`` (a ``FullAttentionSpec`` group for pure-attention models, or +several ``MambaSpec`` groups plus one ``FullAttentionSpec`` group for hybrid +models) is a self-contained transfer unit with its own block table and its own +translation: + + KVCM manager block idx -> global token idx -> group logical block + (step 1, connector-only) (step 2/3, shared with vLLM) + +Step 1 is connector-only logic. A bug there makes *save* gather from the wrong +physical slots and *load* scatter to the wrong physical slots. Because save and +load share the same translation, a transport-level round trip still "matches" +(the bug is symmetric). + +To break the symmetry we capture KV data using ONLY the position -> physical +slot mapping that vLLM itself uses (its slot_mapping kernel), a pure function of +the group's block table and the token position. This reference is independent of +the connector's step-1 logic, so a step-1 bug makes the captured data diverge +from what the connector saved/loaded. + +Group kinds +----------- +* Attention groups -> ``torch.Tensor`` of shape + ``[2, num_blocks, kernel_block_size, num_kv_heads, head_size]``. Captured + per-token using the three-tier mapping (group logical block -> kernel physical + block) because the scheduler's group block size may exceed the kernel block + size. +* Mamba/linear/gdn groups -> ``list[Tensor]`` (e.g. ``[conv_state, ssm_state]``). + The state is stored **per group block**; we capture the whole state slice for + the group block that each manager block maps to (mirroring the connector's + ``_state_block_ids``: a manager block's *last* token selects the block). + +Capture points +-------------- +* Reference (save path): in ``wait_for_save`` we read the KV of the saved token + range straight out of the paged cache (the forward pass has completed and the + slots are not modified by the parent's async gather). +* Loaded (load path): loads are async, so the load step has no forward pass and + the worker does not yet know the request's token ids. We record the load's + per-group block tables in ``start_load_kv`` and emit the capture in a later + ``wait_for_save`` once the token ids have arrived. The loaded KV persists in + the paged cache (its blocks are allocated to the request). + +Captures are written to ``$KVCM_E2E_CAPTURE_DIR`` as ``.pt`` files named +``{ref|loaded}_tp{rank}_{token_hash}.pt`` so the out-of-process driver can match +reference vs loaded by content (the captured token ids). +""" + +import hashlib +import os +import threading +import typing + +import torch + +from kv_cache_manager.py_connector.common.logger import logger +from kv_cache_manager.py_connector.vllm.metadata import TairKvCacheConnectorMetadata +from kv_cache_manager.py_connector.vllm.v1_connector import TairKvCacheConnector + +CAPTURE_DIR_ENV = "KVCM_E2E_CAPTURE_DIR" + + +class VerifyingConnector(TairKvCacheConnector): + """Production connector + independent per-group KV capture for e2e.""" + + # ------------------------------------------------------------------ # + # Setup + # ------------------------------------------------------------------ # + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): + super().register_kv_caches(kv_caches) + + self._capture_dir = os.environ.get(CAPTURE_DIR_ENV, "") + if self._capture_dir: + os.makedirs(self._capture_dir, exist_ok=True) + + # Snapshot the static per-group description into a capture-friendly form. + # Each entry: (group_idx, is_attention, layer_names, group_block_size, + # kernel_block_size). kernel_block_size is read straight off the tensor. + self._cap_groups = [] + for meta in self._group_metas: + if meta.is_attention: + ref = kv_caches[meta.layer_names[0]] + kernel_bs = ref.shape[2] + else: + kernel_bs = 0 + self._cap_groups.append( + (meta.group_idx, meta.is_attention, list(meta.layer_names), + meta.block_size, kernel_bs)) + + # Track completion of async load scatters. The parent's load task already + # CPU-synchronizes its own scatter before reporting the task result, so a + # threading.Event set from the done callback is sufficient to know the + # scatter is globally visible. + self._load_done_events: dict[str, list[threading.Event]] = {} + self._load_events_lock = threading.Lock() + + # Loads are async and their step has no forward pass, so the worker does + # not yet have the request's token ids. Record the load's per-group block + # tables here and emit the capture once the token ids arrive. + # req_id -> (manager_block_idxes, block_ids_per_group) + self._pending_loaded: dict[str, tuple[list, list]] = {} + + orig_factory = self._data_transfer.create_load_done_callback + + def tracking_factory(req_id, *args, **kwargs): + orig_cb = orig_factory(req_id, *args, **kwargs) + evt = threading.Event() + with self._load_events_lock: + self._load_done_events.setdefault(req_id, []).append(evt) + + def cb(task_results): + try: + orig_cb(task_results) + finally: + evt.set() + + return cb + + self._data_transfer.create_load_done_callback = tracking_factory + logger.warning( + "VerifyingConnector enabled, capture_dir=%s tp_rank=%s groups=%s " + "vllm_bs=%s manager_bs=%s", + self._capture_dir, self._tp_rank, + [(g[0], "attn" if g[1] else "state", g[3], g[4]) for g in self._cap_groups], + self._vllm_block_size, self._manager_block_size, + ) + + # ------------------------------------------------------------------ # + # Load hook: record pending loaded captures + # ------------------------------------------------------------------ # + def start_load_kv(self, forward_context, **kwargs) -> None: + meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) + load_reqs = [ + (lr.req_id, list(lr.manager_block_idxes), + [list(b) for b in lr.all_block_ids]) + for lr in meta.to_load_requests + if lr.all_block_ids and lr.need_load_locations + ] + + super().start_load_kv(forward_context, **kwargs) + + if getattr(self, "_capture_dir", "") and load_reqs: + for req_id, mbis, bpg in load_reqs: + self._pending_loaded[req_id] = (mbis, bpg) + logger.warning( + "VerifyingConnector recorded %d pending loaded capture(s)", + len(load_reqs)) + + # ------------------------------------------------------------------ # + # Save hook: reference captures + emit pending loaded captures + # ------------------------------------------------------------------ # + def wait_for_save(self): + meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) + + if getattr(self, "_capture_dir", "") and getattr(self, "_kv_caches", None): + try: + self._capture_refs(meta) + self._capture_pending_loaded() + except Exception as e: # never break inference for a capture error + logger.warning("VerifyingConnector capture failed: %s", e, exc_info=True) + + super().wait_for_save() + + def _capture_refs(self, meta: TairKvCacheConnectorMetadata): + if not meta.to_save_requests: + return + # Make all forward-pass KV writes visible before reading the paged cache. + self._device_mod.synchronize() + for save_req in meta.to_save_requests: + req = self._alive_requests.get(save_req.req_id) + if req is None or not req.block_ids_per_group: + continue + self._capture_range( + kind="ref", + token_ids=req.token_ids, + block_ids_per_group=req.block_ids_per_group, + manager_block_idxes=save_req.manager_block_idxes, + ) + + def _capture_pending_loaded(self): + if not self._pending_loaded: + return + done = [] + for req_id, (mbis, bpg) in self._pending_loaded.items(): + req = self._alive_requests.get(req_id) + if req is None: + # token ids have not arrived on this worker yet; wait for a + # later step in which the request is scheduled. + continue + with self._load_events_lock: + evts = list(self._load_done_events.get(req_id, [])) + for evt in evts: + evt.wait(timeout=120) + self._device_mod.synchronize() + self._capture_range( + kind="loaded", + token_ids=req.token_ids, + block_ids_per_group=bpg, + manager_block_idxes=mbis, + ) + done.append(req_id) + for req_id in done: + del self._pending_loaded[req_id] + + # ------------------------------------------------------------------ # + # Capture helpers + # ------------------------------------------------------------------ # + def _capture_range(self, kind, token_ids, block_ids_per_group, manager_block_idxes): + if not manager_block_idxes or not block_ids_per_group: + return + # One record per manager block. Saves are batched incrementally while + # loads arrive all-at-once, so per-block records let the driver match + # reference vs loaded captures by each block's token content. + for b in manager_block_idxes: + self._capture_block(kind, token_ids, block_ids_per_group, b) + + def _attn_token_slot(self, pos, block_table, group_bs, kernel_bs): + """Map a global token position to its flat slot in an attention group. + + Mirrors vLLM's own slot_mapping kernel expressed with the three-tier + block hierarchy (group logical block -> kernel physical block). Works for + pure-attention groups (group_bs == kernel_bs, ratio 1) and hybrid + attention groups (group block larger than kernel block). Independent of + the connector's step-1 (manager-block) logic, which is what we verify. + """ + ratio = group_bs // kernel_bs + logical = pos // group_bs + off = pos % group_bs + physical = block_table[logical] * ratio + off // kernel_bs + return physical * kernel_bs + off % kernel_bs + + def _capture_block(self, kind, token_ids, block_ids_per_group, manager_block_idx): + mbs = self._manager_block_size + + # Global token positions covered by this manager block. + positions = list(range(manager_block_idx * mbs, (manager_block_idx + 1) * mbs)) + if positions[-1] >= len(token_ids): + positions = [p for p in positions if p < len(token_ids)] + if not positions: + return + + captured_token_ids = [token_ids[p] for p in positions] + kv_by_layer = {} + + for group_idx, is_attention, layer_names, group_bs, kernel_bs in self._cap_groups: + block_table = block_ids_per_group[group_idx] + if is_attention: + slots = [self._attn_token_slot(p, block_table, group_bs, kernel_bs) + for p in positions] + slot_tensor = torch.tensor(slots, dtype=torch.long, device=self._device) + for layer_name in layer_names: + kv_cache = self._kv_caches[layer_name] + # vLLM >= 0.26.0: (num_blocks, num_kv_heads, kernel_bs, + # 2*head_size) packed, NHD memory order is token-major. Flatten + # the (block, token) dims and gather the whole per-token vector + # (K and V packed) -- the packing is opaque to verification. + per_token = kv_cache.shape[1] * kv_cache.shape[3] + flat = kv_cache.permute(0, 2, 1, 3).reshape(-1, per_token) + gathered = flat[slot_tensor, :].contiguous() # [n_tok, per_token] + kv_by_layer[layer_name] = gathered.cpu() + else: + # State stored once per group block; the manager block's last + # token selects the block (mirrors _state_block_ids). vLLM's + # mamba "align" mode materializes states only at segment + # boundaries -- interior blocks hold the null block (id 0) and + # carry no state to capture (the connector skips them too). + logical = ((manager_block_idx + 1) * mbs - 1) // group_bs + block_id = block_table[logical] + if block_id == 0: + continue + for layer_name in layer_names: + states = self._kv_caches[layer_name] # list[Tensor] + kv_by_layer[layer_name] = [s[block_id].detach().cpu() for s in states] + + token_hash = hashlib.sha256( + torch.tensor(captured_token_ids, dtype=torch.int64).numpy().tobytes() + ).hexdigest()[:16] + path = os.path.join(self._capture_dir, f"{kind}_tp{self._tp_rank}_{token_hash}.pt") + torch.save({"token_ids": captured_token_ids, "kv": kv_by_layer}, path) + logger.warning( + "VerifyingConnector captured %s block=%d tokens=%d..%d tp=%s -> %s", + kind, manager_block_idx, positions[0], positions[-1], self._tp_rank, path) + + +class MutatedConnector(VerifyingConnector): + """Meta-test connector: injects an off-by-one into the attention token + translation (every gathered/scattered slot shifted by -1). + + The shift is symmetric between save and load, so with contiguous block + tables a transport round trip cancels it in the interior of the loaded + range (slot(t)-1 == slot(t-1)); the leak is at the boundary: the last + loaded token's true slot is never written and keeps stale (uninitialized) + data. The capture-based verification reads the cache through vLLM's own + slot mapping and must observe that divergence -- the mutation e2e test + asserts that verification FAILS with this connector. + + -1 (not +1) keeps every shifted slot in bounds: vLLM reserves physical + block 0 as the null block, so real slots are >= kernel_block_size and + slot-1 >= 0, while slot+1 of the cache's last block would read/write out + of bounds. Only reachable through the test-side ``kv_connector_module_path`` + injection; never part of the production wheel. + """ + + def _attn_token_indices(self, group, manager_block_idxes, block_table): + out = super()._attn_token_indices(group, manager_block_idxes, block_table) + return [[slot - 1 for slot in block] for block in out] diff --git a/integration_test/vllm_e2e/test_full_hit.py b/integration_test/vllm_e2e/test_full_hit.py new file mode 100644 index 000000000..39c214ecb --- /dev/null +++ b/integration_test/vllm_e2e/test_full_hit.py @@ -0,0 +1,82 @@ +"""test_full_hit: full-prompt external hit must not crash the engine. + +Regression test for the synchronous-load full-hit bug: this connector reports +external matches with load_kv_async=False, so vLLM schedules +``num_tokens - num_computed_tokens`` new tokens and asserts that count is > 0 +(vllm/v1/core/sched/scheduler.py, waiting-queue loop: ``assert num_new_tokens +> 0``). Without capping, a prompt whose token count is an exact multiple of the +manager block size and whose blocks are all externally cached would make the +count 0 and kill the engine. + +Phase 1 saves a prompt of exactly N manager blocks; phase 2 resends the very +same prompt (as explicit token ids, so tokenization cannot shift the length). +Asserts: the engine survives, the completion is well-formed, and the connector +reports 0 < matched < prompt tokens (the cap dropped at least the last block). + +Runs against both full-attention and hybrid models via $KVCM_E2E_MODEL. +""" + +import logging +import unittest + +from e2e_lib import ( + ScenarioEnv, make_base_prompts, send_completions, tokenize, + wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +class TestFullHit(unittest.TestCase): + def test_full_hit(self): + env = ScenarioEnv("full_hit") + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_p1" if env.hybrid else "") + + mbs = env.manager_block_size() + # Trim a long-enough prompt's token ids to an exact multiple of the + # manager block size (>= 2 blocks so the cap has room to drop one). + toks = tokenize(vllm.base_url(), make_base_prompts(1, env.hybrid)[0]) + num_blocks = len(toks) // mbs + self.assertGreaterEqual( + num_blocks, 2, f"prompt too short: {len(toks)} tokens, mbs={mbs}") + prompt_ids = toks[:num_blocks * mbs] + logger.info("full-hit prompt: %d tokens = %d x %d", + len(prompt_ids), num_blocks, mbs) + + # Phase 1: fresh prefill -> all blocks saved. + resp1 = send_completions(vllm.base_url(), [prompt_ids])[0] + self.assertTrue(resp1["choices"][0]["text"]) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, prompt_ids, + min_blocks=num_blocks)) + + if env.hybrid: + # Clear the local prefix cache so phase 2 goes external. + vllm = env.restart_vllm(log_suffix="_p2") + + # Phase 2: the exact same prompt -> full external hit. Without the + # cap this crashes the engine (assert num_new_tokens > 0). + resp2 = send_completions(vllm.base_url(), [prompt_ids])[0] + self.assertTrue(resp2["choices"][0]["text"]) + + # The engine must still be alive and serving. + resp3 = send_completions(vllm.base_url(), ["sanity check prompt"])[0] + self.assertTrue(resp3["choices"][0]["text"]) + + # Connector-side evidence: matched > 0 (external hit happened) and + # matched < prompt tokens (the cap left tokens to recompute). + matched = [int(g[0]) for g in + env.scan_connector_logs(r"matched (\d+) external tokens")] + self.assertTrue(matched, "no 'matched N external tokens' log found") + hit = [m for m in matched if m > 0] + self.assertTrue(hit, f"no positive external match in {matched}") + self.assertTrue(all(m < len(prompt_ids) for m in hit), + f"match not capped below prompt len: {matched}") + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_load_failure.py b/integration_test/vllm_e2e/test_load_failure.py new file mode 100644 index 000000000..e353dbfe7 --- /dev/null +++ b/integration_test/vllm_e2e/test_load_failure.py @@ -0,0 +1,161 @@ +"""test_load_failure: storage loss between save and load must degrade, not kill. + +Phase 1 saves a long prompt; the test then deletes the storage files of the +*tail half* of the manager blocks (resolved through the manager's ordered +getCacheLocation response, one file per block via key_count_per_file=1). +Phase 2 reloads the same prefix with ``block_per_load_task=1`` so every block +fails or succeeds independently. + +Full-attention models (single group, report_failures=True): the connector +reports the failed blocks' vLLM block ids; with +``kv_load_failure_policy="recompute"`` (vLLM 0.26.0 defaults to "fail", which +turns any load failure into a 500) vLLM truncates the computed-token count at +the first invalid block and recomputes from there +(vllm/v1/core/sched/scheduler.py::_handle_invalid_blocks / +_update_requests_with_invalid_blocks). Asserts: the request returns a normal +completion, the failures were logged and reported, every *surviving* head +block's loaded KV is bit-exact, and any mismatching capture belongs to a +deleted (recomputed) block. Recomputed blocks are not held to bit-exactness: +they contain freshly recomputed KV whose numerics depend on prefill chunking, +which is vLLM's business, not the connector's. + +Hybrid models (multiple groups, report_failures=False): vLLM's invalid-block +recovery only supports single-group block tables, so the connector only logs +the failure. Asserts: the request still returns (no hang, no crash) and the +failure was logged. KV content is NOT verified: with the failure swallowed +the affected blocks keep garbage by design. + +This scenario also regression-tests the fail-reschedule loop fix: a request +whose external load failed must not re-match external blocks on requeue +(v1_connector.get_num_new_matched_tokens retry guard), otherwise the engine +loops load-fail-reschedule forever and the request hangs. +""" + +import logging +import os +import unittest +from urllib.parse import urlparse + +import requests + +from e2e_lib import ( + ScenarioEnv, compare_captures, full_block_hashes, make_base_prompts, + send_completions, tokenize, wait_for_captures, wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +def get_block_files(manager_uri: str, instance_id: str, token_ids: list[int], + spec_name: str = "tp0_g0") -> list[str]: + """Per-manager-block storage file paths, in block order, from the manager's + getCacheLocation response.""" + r = requests.post(f"{manager_uri}/api/getCacheLocation", json={ + "trace_id": "e2e_block_files", + "token_ids": token_ids, + "instance_id": instance_id, + "query_type": "QT_PREFIX_MATCH", + "block_mask": {"offset": 0}, + }, timeout=10) + r.raise_for_status() + files = [] + for location in r.json().get("locations", []): + for spec in location.get("location_specs", []): + if spec["name"] == spec_name: + # uri: file://?size=... + files.append(urlparse(spec["uri"]).path) + return files + + +class TestLoadFailure(unittest.TestCase): + def test_load_failure(self): + env = ScenarioEnv( + "load_failure", + extra_config_overrides={"block_per_load_task": 1}, + key_count_per_file=1, # one file per block -> per-block failures + kv_load_failure_policy="recompute", + ) + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_p1" if env.hybrid else "") + mbs = env.manager_block_size() + + prompt = make_base_prompts(1, env.hybrid)[0] + suffix = " Now answer: what is 2+2?" + toks = tokenize(vllm.base_url(), prompt) + save_blocks = len(toks) // mbs + self.assertGreaterEqual(save_blocks, 2) + + # ---- Phase 1: save everything. + send_completions(vllm.base_url(), [prompt]) + wait_for_captures(env.capture_dir, "ref", expected=save_blocks, + timeout=180) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, toks, + min_blocks=save_blocks)) + + # ---- Sabotage: delete the tail half of the blocks' files. The + # head blocks stay loadable, so vLLM truncates at the first deleted + # block and the surviving loads remain verifiable. + files = get_block_files(env.manager.manager_uri(), env.instance_id, + toks) + self.assertEqual(len(files), save_blocks) + keep = save_blocks // 2 + for path in files[keep:]: + os.remove(path) + logger.info("deleted %d/%d block files (kept blocks 0..%d)", + save_blocks - keep, save_blocks, keep - 1) + + if env.hybrid: + vllm = env.restart_vllm(log_suffix="_p2") + + # ---- Phase 2: load with holes. The request must return normally. + resp = send_completions(vllm.base_url(), [prompt + suffix])[0] + self.assertTrue(resp["choices"][0]["text"]) + + # The engine must survive and keep serving. + resp2 = send_completions(vllm.base_url(), ["engine alive?"])[0] + self.assertTrue(resp2["choices"][0]["text"]) + + failed_tasks = env.scan_connector_logs(r"load task failed") + self.assertTrue(failed_tasks, "no load failure was logged; the " + "sabotage did not break any loaded block") + + if env.hybrid: + # report_failures=False path: swallowed but logged. + swallowed = env.scan_connector_logs( + r"load failed for \d+/\d+ blocks .*hybrid") + self.assertTrue(swallowed, + "hybrid load failure was not logged") + return + + # Full-attention: vLLM was told about the invalid blocks... + reported = env.scan_connector_logs(r"block_ids_with_load_errors") + self.assertTrue(reported, "failed loads were not reported to vLLM") + + # ...and every surviving head block's loaded KV is bit-exact, + # while any mismatch belongs to a deleted (recomputed) block. + wait_for_captures(env.capture_dir, "loaded", expected=keep, + timeout=180) + report = compare_captures(env.capture_dir, tp_size=1) + hashes = full_block_hashes(toks, mbs) + kept_keys = {("tp0", h) for h in hashes[:keep]} + deleted_keys = {("tp0", h) for h in hashes[keep:]} + failed_keys = {f["key"] for f in report["failures"]} + self.assertFalse( + failed_keys & kept_keys, + f"surviving loaded blocks mismatched: {failed_keys & kept_keys}") + self.assertTrue( + failed_keys <= deleted_keys, + f"mismatches outside the deleted blocks: " + f"{failed_keys - deleted_keys}") + matched_kept = kept_keys & set(report["matched_keys"]) + self.assertEqual( + len(matched_kept), keep, + f"only {len(matched_kept)}/{keep} surviving blocks verified") + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_multi_turn.py b/integration_test/vllm_e2e/test_multi_turn.py new file mode 100644 index 000000000..5912b0895 --- /dev/null +++ b/integration_test/vllm_e2e/test_multi_turn.py @@ -0,0 +1,112 @@ +"""test_multi_turn: decode-time incremental save feeds the next turn's hit. + +Turn 1 sends a prompt and generates enough output tokens (max_tokens crossing +at least one manager block boundary) that the save threshold in +``build_connector_meta`` fires again during decode: blocks composed of +generated tokens are saved incrementally. Turn 2 sends prompt + turn-1 output +as its prompt (a real multi-turn conversation) and must externally match +*more* blocks than the turn-1 prompt alone covers -- proving decode-produced +blocks were saved -- and their loaded KV must verify against the references +captured during decode. + +Full-attention (mbs=16): three decode blocks, same server both turns (prefix +caching off, the external hit is directly observable). +Hybrid (mbs=528): one decode block (528+ generated tokens); the server is +restarted before turn 2 because prefix caching must stay on for hybrid models +and would otherwise mask the external hit with a local one. +""" + +import logging +import unittest + +from e2e_lib import ( + ScenarioEnv, assert_report_ok, compare_captures, send_completions, + shared_token_prefix_len, tokenize, wait_for_captures, + wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +class TestMultiTurn(unittest.TestCase): + def test_multi_turn(self): + env = ScenarioEnv("multi_turn") + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_t1" if env.hybrid else "") + mbs = env.manager_block_size() + + turn1_prompt = ("A short story request. Please write a long, " + "detailed story about a robot that learns to paint.") + prompt_tokens = tokenize(vllm.base_url(), turn1_prompt) + prompt_blocks = len(prompt_tokens) // mbs + + # ---- Turn 1: generate output crossing >= 1 manager block + # boundary (3 blocks for full-attn's mbs=16; 1 block for hybrid's + # mbs=528 to keep decode time bounded). + max_tokens = (mbs + 32) if env.hybrid else (mbs * 3 + 5) + resp = send_completions( + vllm.base_url(), [turn1_prompt], max_tokens=max_tokens, + ignore_eos=True, return_token_ids=True)[0] + choice = resp["choices"][0] + output_ids = choice["token_ids"] + self.assertEqual(len(output_ids), max_tokens) + turn1_ids = choice["prompt_token_ids"] + output_ids + # The connector tracks tokens when they are *scheduled as input*; + # the very last sampled token never re-enters a step, so at most + # (len - 1) // mbs blocks can have been committed. + turn1_blocks = (len(turn1_ids) - 1) // mbs + self.assertGreater( + turn1_blocks, prompt_blocks, + "turn 1 output did not cross a manager block boundary") + logger.info("turn1: %d prompt + %d output tokens = %d blocks " + "(prompt alone: %d)", len(choice["prompt_token_ids"]), + len(output_ids), turn1_blocks, prompt_blocks) + + # Decode-produced blocks must be committed: the manager holds the + # full prompt+output prefix, more blocks than the prompt covers. + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, turn1_ids, + min_blocks=turn1_blocks)) + wait_for_captures(env.capture_dir, "ref", expected=turn1_blocks, + timeout=180) + + if env.hybrid: + # Prefix caching is on for hybrid; restart so turn 2's hit + # comes from KVCM, not the local prefix cache. + vllm = env.restart_vllm(log_suffix="_t2") + + # ---- Turn 2: conversation continues; prompt embeds turn 1's + # prompt + output as token ids (immune to detokenization drift). + turn2_suffix = tokenize(vllm.base_url(), + " Now summarize the story in one word.") + turn2_ids = turn1_ids + turn2_suffix + shared_blocks = min( + shared_token_prefix_len(turn1_ids, turn2_ids) // mbs, + turn1_blocks) + self.assertGreater(shared_blocks, prompt_blocks, + "turn 2 shares no decode-produced block") + resp2 = send_completions(vllm.base_url(), [turn2_ids])[0] + self.assertTrue(resp2["choices"][0]["text"]) + + # The external hit must cover decode-produced blocks. + matched = [int(g[0]) for g in + env.scan_connector_logs(r"matched (\d+) external tokens")] + best = max(matched, default=0) + self.assertGreater( + best, prompt_blocks * mbs, + f"external hit ({best} tokens) does not exceed the prompt-only " + f"coverage ({prompt_blocks * mbs} tokens): decode-time saves " + f"were not used") + + # And the loaded decode-block KV must verify. + wait_for_captures(env.capture_dir, "loaded", + expected=shared_blocks, timeout=180) + report = compare_captures(env.capture_dir, tp_size=1) + assert_report_ok(report, min_matched=shared_blocks) + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_mutation.py b/integration_test/vllm_e2e/test_mutation.py new file mode 100644 index 000000000..12c4e1e18 --- /dev/null +++ b/integration_test/vllm_e2e/test_mutation.py @@ -0,0 +1,35 @@ +"""test_mutation: meta-test proving the e2e KV verification is not vacuous. + +Runs the standard basic scenario with ``MutatedConnector`` (defined in +test_connector.py, injected only through the test-side +``kv_connector_module_path``), which shifts every attention slot produced by +``_attn_token_indices`` by one -- a symmetric off-by-one: save gathers token +t's KV from the shifted slot and load scatters it back there, so with +contiguous block tables a transport round trip cancels the bug in the interior +of the loaded range. It cannot cancel at the range boundary: one loaded +token's true slot is never written and keeps stale uninitialized data. The +capture comparison reads the cache through vLLM's own slot mapping and must +observe the divergence; run_e2e(expect_verification_failure=True) asserts the +verification FAILS. If the mutated run verifies clean, the harness is blind +and this test fails. +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestMutation(unittest.TestCase): + def test_mutated_connector_is_caught(self): + run_e2e( + scenario="mutation", + tp_size=1, + num_prompts=1, + preferred_block_size=0, + connector_name="MutatedConnector", + expect_verification_failure=True, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_partial_hit.py b/integration_test/vllm_e2e/test_partial_hit.py new file mode 100644 index 000000000..399bd5eb9 --- /dev/null +++ b/integration_test/vllm_e2e/test_partial_hit.py @@ -0,0 +1,96 @@ +"""test_partial_hit: incremental query + incremental save on a non-zero prefix. + +The standard scenarios always start requests from zero computed blocks, so the +``computed_blocks > 0`` branch of ``get_num_new_matched_tokens`` (manager query +with a non-zero ``block_mask.offset``) and the ``block_mask.offset`` branch of +the save path (manager skipping already-stored blocks) are never exercised. +This scenario forces both, with prefix caching enabled for all model types: + +* Stage 1: save prefix A (fresh prefill). +* Stage 2 (same server): send A+B. vLLM locally hits A, so the connector + queries the manager with offset = locally computed blocks (> 0), then + extends the save; the manager's start_write_cache response skips the + already-stored A blocks via a non-zero block_mask offset. +* Stage 3 (restarted server, local cache cleared): send A+B+C. The connector + externally matches A+B -- proving stage 2's incremental save committed -- + loads it, and the loaded KV is verified against the reference captures. + +Log evidence asserted: a getCacheLocation request carrying a non-zero offset +(requires connector DEBUG logging, enabled via log_level). +""" + +import logging +import unittest + +from e2e_lib import ( + ScenarioEnv, assert_report_ok, compare_captures, make_base_prompts, + send_completions, shared_token_prefix_len, tokenize, wait_for_captures, + wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +class TestPartialHit(unittest.TestCase): + def test_partial_hit(self): + env = ScenarioEnv("partial_hit", enable_prefix_caching=True, + log_level="DEBUG") + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_s12") + mbs = env.manager_block_size() + + base = make_base_prompts(1, env.hybrid)[0] + # Three nested prompts: A < A+B < A+B+C. + prompt_a = base + prompt_ab = base + " Continuation section B. " + " ".join( + f"Extra sentence {j} carries value {j * 13 + 7}." + for j in range(90 if env.hybrid else 30)) + prompt_abc = prompt_ab + " Final question: what is 2+2?" + + toks_a = tokenize(vllm.base_url(), prompt_a) + toks_ab = tokenize(vllm.base_url(), prompt_ab) + toks_abc = tokenize(vllm.base_url(), prompt_abc) + blocks_a = len(toks_a) // mbs + blocks_ab = len(toks_ab) // mbs + shared_abc = shared_token_prefix_len(toks_ab, toks_abc) // mbs + self.assertGreaterEqual(blocks_a, 1) + self.assertGreater(blocks_ab, blocks_a, + "B must add at least one manager block") + logger.info("blocks: A=%d AB=%d shared(AB,ABC)=%d", + blocks_a, blocks_ab, shared_abc) + + # ---- Stage 1: save A. + send_completions(vllm.base_url(), [prompt_a]) + wait_for_captures(env.capture_dir, "ref", expected=blocks_a, timeout=180) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, toks_a, + min_blocks=blocks_a)) + + # ---- Stage 2: A hits the local prefix cache -> incremental + # external query (non-zero offset) + incremental save of B. + send_completions(vllm.base_url(), [prompt_ab]) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, toks_ab, + min_blocks=blocks_ab)) + wait_for_captures(env.capture_dir, "ref", expected=blocks_ab, timeout=180) + + offsets = [int(g[0]) for g in env.scan_connector_logs( + r"get_kvcache_location request:.*'offset': (\d+)")] + self.assertTrue(any(o > 0 for o in offsets), + f"no incremental query with non-zero offset: {offsets}") + + # ---- Stage 3: restart (clear local cache) and load A+B. + vllm = env.restart_vllm(log_suffix="_s3") + send_completions(vllm.base_url(), [prompt_abc]) + wait_for_captures(env.capture_dir, "loaded", + expected=shared_abc, timeout=180) + + report = compare_captures(env.capture_dir, tp_size=1) + assert_report_ok(report, min_matched=shared_abc) + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_tp.py b/integration_test/vllm_e2e/test_tp.py new file mode 100644 index 000000000..a44c572f4 --- /dev/null +++ b/integration_test/vllm_e2e/test_tp.py @@ -0,0 +1,33 @@ +"""test_tp: TP=2 save/load KV verification with a non-trivial block translation. + +Runs the save/load verification under tensor parallelism (TP=2), where each rank +has an independent forward context, slot mapping and capture, and the connector's +ZMQ-based TP coordination is fully exercised. + +For full-attention models it also sets ``preferred_block_size=32`` while vLLM +uses its default block size (16), forcing the connector's manager-block <-> +group-block translation (``_attn_token_indices``) to do real cross-block +mapping -- the code path most prone to symmetric save/load bugs. For hybrid +models the manager block size is pinned to the scheduler block size (mamba state +is per scheduler block), so run_e2e ignores preferred_block_size there. + +Works for both full-attention and hybrid models (selected via KVCM_E2E_MODEL). +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestTp(unittest.TestCase): + def test_tp(self): + run_e2e( + scenario="tp", + tp_size=2, + num_prompts=2, + preferred_block_size=32, # != vllm block size (16) for full-attn models + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/common/types.py b/kv_cache_manager/py_connector/common/types.py index bc341e7e7..19a546db2 100644 --- a/kv_cache_manager/py_connector/common/types.py +++ b/kv_cache_manager/py_connector/common/types.py @@ -1,22 +1,54 @@ -from enum import Enum -from typing import Tuple, Dict, Optional, Any +from dataclasses import dataclass, field +from typing import List, Optional -import attrs import torch -@attrs.define(frozen=True) +@dataclass +class TransferGroup: + """One KV cache group = one independent transfer unit. + + A vLLM model exposes ``kv_cache_config.kv_cache_groups``: for pure attention + models there is a single ``FullAttentionSpec`` group; hybrid models expose + several ``MambaSpec`` groups plus one ``FullAttentionSpec`` group. Every + group has its own block table (``block_ids`` is a tuple indexed by group) + in units of its own ``block_size``, and its own storage strategy, so the + connector treats each group as a self-contained transfer unit. + """ + + group_idx: int + # Location spec name registered with the manager, e.g. "tp0_g3". + spec_name: str + # True for attention layers (token-granular strided gather/scatter); + # False for mamba/linear/gdn state layers (per-block opaque byte copy). + is_attention: bool + layer_names: List[str] + # The group's own block table granularity in tokens (spec.block_size). + block_size: int + # Bytes stored per manager block for this whole group (all its layers). + per_block_bytes: int + layer_num: int = 0 + + # --- Attention-only fields (is_attention == True) --- + # int64 tensor of [K0, V0, K1, V1, ...] data ptrs on the compute device. + kvcache_ptr_tensor_gpu: Optional[torch.Tensor] = None + per_token_dim: int = 0 # num_kv_heads * head_size + kernel_block_size: int = 0 # tensor.shape[2] + kv_stride: int = 0 # tensor.stride(0), 0 => contiguous flat layout + block_stride: int = 0 # tensor.stride(1), 0 => contiguous flat layout + + # --- State-only fields (is_attention == False) --- + # Per layer (num_blocks, page_size_bytes) uint8 views into the state storage. + block_view_tensors: List[torch.Tensor] = field(default_factory=list) + page_size_bytes: int = 0 # bytes per block per state layer + + +@dataclass class KVCacheInfo: + """Worker-side registered KV cache description (all groups).""" + tp_rank: int world_size: int - kvcaches: Dict[str, torch.Tensor] - kvcache_ptr_tensor_cpu: torch.Tensor - kvcache_ptr_tensor_gpu: torch.Tensor - all_kvcache_ptr_tensor_gpu: torch.Tensor - layer_num: int - local_token_num: int - per_manager_block_shape: Tuple[int, ...] - per_manager_block_byte_size: int - per_token_per_layer_dim_size: int + groups: List[TransferGroup] device: torch.device dtype: torch.dtype diff --git a/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py b/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py index ebb297505..0f986a633 100644 --- a/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py +++ b/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py @@ -42,8 +42,15 @@ def kv_cache_batch_gather_kernel( NUM_KVCACHE_PTRS: tl.constexpr, # num_layers * kv_count BLOCK_SIZE: tl.constexpr, # 隐藏维度分块大小 DTYPE: tl.constexpr = tl.float16, + kv_stride: tl.constexpr = 0, # stride between K and V (for V pointers) + block_stride: tl.constexpr = 0, # stride between blocks (0 = use flat indexing) + local_block_size: tl.constexpr = 0, # actual block size in tensor (0 = use NUM_TOKENS_PER_BLOCK) ): NUM_DIMS_PER_BLOCK = NUM_TOKENS_PER_BLOCK * NUM_DIMS_PER_TOKEN + + # Determine if using strided layout + USE_STRIDED: tl.constexpr = (block_stride != 0) + EFFECTIVE_LOCAL_BLOCK_SIZE: tl.constexpr = local_block_size if local_block_size > 0 else NUM_TOKENS_PER_BLOCK pid = tl.program_id(0) grid_size = tl.num_programs(0) # 实际grid大小 (如3) @@ -64,6 +71,8 @@ def kv_cache_batch_gather_kernel( # 3. 遍历所有KV缓存指针 (k/v for each layer) for ptr_idx in tl.range(NUM_KVCACHE_PTRS): # 3.1 加载当前层的KV缓存基地址 + # Note: For non-MLA, pointer array is [K0, V0, K1, V1, ...] + # V pointer is already V's base (tensor[1].data_ptr()), no need to add kv_stride kvcache_ptr = tl.load(kv_cache_ptrs_ptr + ptr_idx).to(tl.pointer_type(DTYPE)) # 3.2 计算当前层在dst中的基础偏移 @@ -90,7 +99,16 @@ def kv_cache_batch_gather_kernel( # 从HBM的KV缓存加载数据 # 计算源指针: [BLOCK_SIZE] - src_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token + if USE_STRIDED: + # Strided layout: convert flat token index to strided offset + # V pointer already includes kv_stride offset, so no need to add it again + kv_block_idx = global_token_idx // EFFECTIVE_LOCAL_BLOCK_SIZE + token_in_kv_block = global_token_idx % EFFECTIVE_LOCAL_BLOCK_SIZE + strided_offset = kv_block_idx * block_stride + token_in_kv_block * NUM_DIMS_PER_TOKEN + src_ptrs = kvcache_ptr + strided_offset + dim_idx_in_token + else: + # Contiguous layout: flat indexing + src_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token load_mask = mask & token_gather_mask data = tl.load(src_ptrs, mask=load_mask, other=0.0) # 大块连续写入 host memory (PCIe优化) @@ -107,7 +125,10 @@ def batch_gather_kv_caches( dst_block_indices: List[int], # List of dst block indices num_tokens_per_block: int, dim_size_per_token_per_layer: int, - sm_count: int = 3 + sm_count: int = 3, + kv_stride: int = 0, # stride between K and V (for V pointers) + block_stride: int = 0, # stride between blocks (0 = use flat indexing) + local_block_size: int = 0, # actual block size in tensor (0 = use num_tokens_per_block) ): # 配置参数 total_blocks = len(dst_block_indices) @@ -132,6 +153,9 @@ def batch_gather_kv_caches( BLOCK_SIZE=2048, DTYPE=pytorch_dtype_to_triton_dtype(dst_tensor.dtype), num_warps=32, + kv_stride=kv_stride, + block_stride=block_stride, + local_block_size=local_block_size if local_block_size > 0 else num_tokens_per_block, ) # TODO autotune num_warps and BLOCK_SIZE @@ -148,8 +172,15 @@ def kv_cache_batch_scatter_kernel( NUM_KVCACHE_PTRS: tl.constexpr, # num_layers * kv_count BLOCK_SIZE: tl.constexpr, # 隐藏维度分块大小 DTYPE: tl.constexpr = tl.float16, + kv_stride: tl.constexpr = 0, # stride between K and V (for V pointers) + block_stride: tl.constexpr = 0, # stride between blocks (0 = use flat indexing) + local_block_size: tl.constexpr = 0, # actual block size in tensor (0 = use NUM_TOKENS_PER_BLOCK) ): NUM_DIMS_PER_BLOCK = NUM_TOKENS_PER_BLOCK * NUM_DIMS_PER_TOKEN + + # Determine if using strided layout + USE_STRIDED: tl.constexpr = (block_stride != 0) + EFFECTIVE_LOCAL_BLOCK_SIZE: tl.constexpr = local_block_size if local_block_size > 0 else NUM_TOKENS_PER_BLOCK pid = tl.program_id(0) grid_size = tl.num_programs(0) # 实际grid大小 (如3) @@ -170,6 +201,8 @@ def kv_cache_batch_scatter_kernel( # 3. 遍历所有KV缓存指针 (k/v for each layer) for ptr_idx in range(NUM_KVCACHE_PTRS): # 3.1 加载当前层的KV缓存基地址 + # Note: For non-MLA, pointer array is [K0, V0, K1, V1, ...] + # V pointer is already V's base (tensor[1].data_ptr()), no need to add kv_stride kvcache_ptr = tl.load(kv_cache_ptrs_ptr + ptr_idx).to(tl.pointer_type(DTYPE)) # 3.2 计算当前层在src中的基础偏移 @@ -201,7 +234,16 @@ def kv_cache_batch_scatter_kernel( # 向HBM的KV缓存写入数据 # 计算目的指针: [BLOCK_SIZE] - dst_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token + if USE_STRIDED: + # Strided layout: convert flat token index to strided offset + # V pointer already includes kv_stride offset, so no need to add it again + kv_block_idx = global_token_idx // EFFECTIVE_LOCAL_BLOCK_SIZE + token_in_kv_block = global_token_idx % EFFECTIVE_LOCAL_BLOCK_SIZE + strided_offset = kv_block_idx * block_stride + token_in_kv_block * NUM_DIMS_PER_TOKEN + dst_ptrs = kvcache_ptr + strided_offset + dim_idx_in_token + else: + # Contiguous layout: flat indexing + dst_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token tl.store(dst_ptrs, data, mask=load_mask) @@ -214,7 +256,10 @@ def batch_scatter_kv_caches( src_block_indices: List[int], # List of src block indices num_tokens_per_block: int, dim_size_per_token_per_layer: int, - sm_count: int = 3 + sm_count: int = 3, + kv_stride: int = 0, # stride between K and V (for V pointers) + block_stride: int = 0, # stride between blocks (0 = use flat indexing) + local_block_size: int = 0, # actual block size in tensor (0 = use num_tokens_per_block) ): # 配置参数 total_blocks = len(src_block_indices) @@ -245,4 +290,7 @@ def batch_scatter_kv_caches( BLOCK_SIZE=2048, DTYPE=pytorch_dtype_to_triton_dtype(src_tensor.dtype), num_warps=32, + kv_stride=kv_stride, + block_stride=block_stride, + local_block_size=local_block_size if local_block_size > 0 else num_tokens_per_block, ) diff --git a/kv_cache_manager/py_connector/test/BUILD b/kv_cache_manager/py_connector/test/BUILD new file mode 100644 index 000000000..4d11111b5 --- /dev/null +++ b/kv_cache_manager/py_connector/test/BUILD @@ -0,0 +1,36 @@ +load("@rules_python//python:py_library.bzl", "py_library") +load("@rules_python//python:py_test.bzl", "py_test") + +# Stubs that make v1_connector importable without vLLM / CUDA / the compiled +# kvcm_py_client. Tests import this module before anything under vllm/. +py_library( + name = "vllm_stubs", + srcs = [ + "__init__.py", + "vllm_stubs.py", + ], + deps = [ + "//kv_cache_manager/py_connector/vllm:vllm_connector", + ], +) + +py_test( + name = "test_block_translation", + srcs = ["test_block_translation.py"], + tags = ["no-remote-exec"], + deps = [":vllm_stubs"], +) + +py_test( + name = "test_data_transfer_results", + srcs = ["test_data_transfer_results.py"], + tags = ["no-remote-exec"], + deps = [":vllm_stubs"], +) + +py_test( + name = "test_scheduler_state", + srcs = ["test_scheduler_state.py"], + tags = ["no-remote-exec"], + deps = [":vllm_stubs"], +) diff --git a/kv_cache_manager/py_connector/test/kernel/BUILD b/kv_cache_manager/py_connector/test/kernel/BUILD new file mode 100644 index 000000000..bbea02aeb --- /dev/null +++ b/kv_cache_manager/py_connector/test/kernel/BUILD @@ -0,0 +1,18 @@ +load("@rules_python//python:py_test.bzl", "py_test") + +py_test( + name = "test_strided_gather_scatter", + srcs = [ + "__init__.py", + "test_strided_gather_scatter.py", + ], + main = "test_strided_gather_scatter.py", + tags = [ + "no-remote-exec", + "gpu", # requires 1 GPU + "exclusive", # GPU tests run serially to avoid CUDA contention + ], + deps = [ + "//kv_cache_manager/py_connector/kernel", + ], +) diff --git a/kv_cache_manager/py_connector/test/kernel/test_strided_gather_scatter.py b/kv_cache_manager/py_connector/test/kernel/test_strided_gather_scatter.py new file mode 100644 index 000000000..b1ca1caa2 --- /dev/null +++ b/kv_cache_manager/py_connector/test/kernel/test_strided_gather_scatter.py @@ -0,0 +1,186 @@ +"""GPU tests for the strided path of the batch gather/scatter Triton kernel. + +The flat path (block_stride=0) is covered by test_batch_gather_scatter.py. +Here we cover the strided path added for vLLM's paged layout, where the flat +token index is decomposed as (kv_block, token_in_block) and the block starts +``block_stride`` elements apart -- including padded pages where +``block_stride > local_block_size * dims_per_token`` leaves a gap between +blocks that must be skipped, not walked. + +Every case is checked element-wise against a naive torch reference that +performs the same (kv_block, token) decomposition with plain indexing. +""" + +import unittest + +import torch + +from kv_cache_manager.py_connector.kernel.batch_gather_scatter_helper import ( + batch_gather_kv_caches, + batch_scatter_kv_caches, +) + + +def _make_paged_caches(num_layers, num_blocks, local_block_size, dims_per_token, + pad_tokens, device, dtype, fill_random=True): + """Per-layer paged caches shaped (num_blocks, padded_tokens, dims) where + padded_tokens = local_block_size + pad_tokens. block_stride (in elements) + is padded_tokens * dims_per_token.""" + caches = [] + for _ in range(num_layers): + t = torch.randn(num_blocks, local_block_size + pad_tokens, dims_per_token, + device=device, dtype=dtype) if fill_random else \ + torch.zeros(num_blocks, local_block_size + pad_tokens, dims_per_token, + device=device, dtype=dtype) + caches.append(t) + return caches + + +def _ref_slot(cache, flat_token_idx, local_block_size): + blk = flat_token_idx // local_block_size + tok = flat_token_idx % local_block_size + return cache[blk, tok, :] + + +class TestStridedGatherScatter(unittest.TestCase): + # (local_block_size, pad_tokens, tokens_per_manager_block) + CASES = [ + (16, 0, 16), # strided == flat geometry (stride still exercised) + (16, 4, 16), # padded pages: gap between blocks + (64, 0, 528), # hybrid attention: manager block spans many kv blocks + (64, 8, 48), # padded + manager block not aligned to kv block + ] + + def setUp(self): + if not torch.cuda.is_available(): + self.skipTest("requires a GPU") + torch.manual_seed(7) + self.device = "cuda" + self.dtype = torch.bfloat16 + self.num_layers = 3 + self.dims = 128 + self.num_kv_blocks = 64 + + def _indices(self, num_manager_blocks, tokens_per_block, local_block_size): + total_tokens = self.num_kv_blocks * local_block_size + need = num_manager_blocks * tokens_per_block + assert need <= total_tokens, "test setup: not enough kv slots" + perm = torch.randperm(total_tokens)[:need] + return perm.tolist() + + def test_gather_strided_matches_reference(self): + for local_bs, pad, tokens_per_block in self.CASES: + with self.subTest(local_bs=local_bs, pad=pad, tpb=tokens_per_block): + caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype) + block_stride = caches[0].stride(0) + self.assertEqual(block_stride, (local_bs + pad) * self.dims) + ptrs = torch.tensor([c.data_ptr() for c in caches], + device=self.device, dtype=torch.int64) + num_mb = 4 + token_indices = self._indices(num_mb, tokens_per_block, local_bs) + dst_block_indices = [2, 0, 3, 1] + dst = torch.zeros(num_mb, self.num_layers, tokens_per_block, + self.dims, device="cpu", dtype=self.dtype, + pin_memory=True) + batch_gather_kv_caches( + ptrs, dst, token_indices, dst_block_indices, + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + + caches_cpu = [c.cpu() for c in caches] + for mb in range(num_mb): + for pos in range(tokens_per_block): + flat_idx = token_indices[mb * tokens_per_block + pos] + for layer in range(self.num_layers): + want = _ref_slot(caches_cpu[layer], flat_idx, local_bs) + got = dst[dst_block_indices[mb], layer, pos, :] + torch.testing.assert_close( + got, want, + msg=f"gather mismatch mb={mb} pos={pos} " + f"layer={layer} flat={flat_idx}") + + def test_scatter_strided_matches_reference(self): + for local_bs, pad, tokens_per_block in self.CASES: + with self.subTest(local_bs=local_bs, pad=pad, tpb=tokens_per_block): + caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype, fill_random=False) + # Sentinel in the padding region: scatter must never touch it. + sentinel = 123.0 + if pad: + for c in caches: + c[:, local_bs:, :] = sentinel + block_stride = caches[0].stride(0) + ptrs = torch.tensor([c.data_ptr() for c in caches], + device=self.device, dtype=torch.int64) + num_mb = 4 + token_indices = self._indices(num_mb, tokens_per_block, local_bs) + src_block_indices = [1, 3, 0, 2] + src = torch.randn(num_mb, self.num_layers, tokens_per_block, + self.dims, dtype=self.dtype).pin_memory() + batch_scatter_kv_caches( + ptrs, src, token_indices, src_block_indices, + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + + caches_cpu = [c.cpu() for c in caches] + for mb in range(num_mb): + for pos in range(tokens_per_block): + flat_idx = token_indices[mb * tokens_per_block + pos] + for layer in range(self.num_layers): + got = _ref_slot(caches_cpu[layer], flat_idx, local_bs) + want = src[src_block_indices[mb], layer, pos, :] + torch.testing.assert_close( + got, want, + msg=f"scatter mismatch mb={mb} pos={pos} " + f"layer={layer} flat={flat_idx}") + if pad: + for layer, c in enumerate(caches_cpu): + self.assertTrue( + bool((c[:, local_bs:, :] == sentinel).all()), + f"scatter wrote into the padding of layer {layer}") + + def test_gather_scatter_roundtrip_strided(self): + """Scattering gathered data into zeroed caches must reproduce exactly + the gathered slots (and only them).""" + local_bs, pad, tokens_per_block = 64, 8, 48 + src_caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype) + dst_caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype, fill_random=False) + block_stride = src_caches[0].stride(0) + src_ptrs = torch.tensor([c.data_ptr() for c in src_caches], + device=self.device, dtype=torch.int64) + dst_ptrs = torch.tensor([c.data_ptr() for c in dst_caches], + device=self.device, dtype=torch.int64) + num_mb = 3 + token_indices = self._indices(num_mb, tokens_per_block, local_bs) + buf = torch.zeros(num_mb, self.num_layers, tokens_per_block, self.dims, + device="cpu", dtype=self.dtype, pin_memory=True) + batch_gather_kv_caches( + src_ptrs, buf, token_indices, list(range(num_mb)), + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + batch_scatter_kv_caches( + dst_ptrs, buf, token_indices, list(range(num_mb)), + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + src_cpu = [c.cpu() for c in src_caches] + dst_cpu = [c.cpu() for c in dst_caches] + for flat_idx in token_indices: + for layer in range(self.num_layers): + torch.testing.assert_close( + _ref_slot(dst_cpu[layer], flat_idx, local_bs), + _ref_slot(src_cpu[layer], flat_idx, local_bs)) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/test_block_translation.py b/kv_cache_manager/py_connector/test/test_block_translation.py new file mode 100644 index 000000000..a75af6484 --- /dev/null +++ b/kv_cache_manager/py_connector/test/test_block_translation.py @@ -0,0 +1,125 @@ +"""Unit tests for the connector's manager-block -> physical-slot translation. + +Covers ``_attn_token_indices`` (attention groups: token-granular three-tier +mapping) and ``_state_block_ids`` (mamba/state groups: manager block's last +token selects the group block), verifying against an independent brute-force +reference implementation, token by token. +""" + +import unittest + +from kv_cache_manager.py_connector.test.vllm_stubs import make_connector +from kv_cache_manager.py_connector.common.types import TransferGroup + + +def _make_group(group_bs, kernel_bs=0, is_attention=True): + return TransferGroup( + group_idx=0, + spec_name="tp0_g0", + is_attention=is_attention, + layer_names=["layer0"], + block_size=group_bs, + per_block_bytes=0, + kernel_block_size=kernel_bs, + ) + + +def _ref_attn_token_indices(manager_bs, group_bs, kernel_bs, manager_block_idxes, + block_table): + """Brute-force reference: walk every token of every manager block and map it + through the block hierarchy step by step.""" + out = [] + for mb in manager_block_idxes: + slots = [] + for tok in range(mb * manager_bs, (mb + 1) * manager_bs): + group_block = tok // group_bs # logical block in group table + tok_in_group = tok - group_block * group_bs + kernel_in_group = tok_in_group // kernel_bs + tok_in_kernel = tok_in_group - kernel_in_group * kernel_bs + physical = block_table[group_block] * (group_bs // kernel_bs) + kernel_in_group + slots.append(physical * kernel_bs + tok_in_kernel) + out.append(slots) + return out + + +def _ref_state_block_ids(manager_bs, group_bs, manager_block_idxes, block_table): + """Brute-force reference: the state covering a manager block is the state of + the group block containing the manager block's last token.""" + out = [] + for mb in manager_block_idxes: + last_token = (mb + 1) * manager_bs - 1 + out.append(block_table[last_token // group_bs]) + return out + + +class TestAttnTokenIndices(unittest.TestCase): + # (manager_bs, group_bs, kernel_bs): ratio=1, ratio>1, manager != group. + CASES = [ + (16, 16, 16), # full attention default: all equal + (32, 16, 16), # preferred_block_size > vllm block size + (528, 528, 64), # hybrid: group block spans several kernel blocks + (528, 528, 528), # hybrid with kernel == group + (48, 16, 8), # manager > group > kernel + ] + + def test_against_reference(self): + for manager_bs, group_bs, kernel_bs in self.CASES: + with self.subTest(manager_bs=manager_bs, group_bs=group_bs, + kernel_bs=kernel_bs): + conn = make_connector(manager_block_size=manager_bs) + group = _make_group(group_bs, kernel_bs) + # Enough non-trivially permuted blocks for 4 manager blocks. + needed = 4 * manager_bs // group_bs + 1 + block_table = [(i * 7 + 3) % 97 for i in range(needed)] + mbis = [0, 1, 3] + got = conn._attn_token_indices(group, mbis, block_table) + want = _ref_attn_token_indices( + manager_bs, group_bs, kernel_bs, mbis, block_table) + self.assertEqual(got, want) + + def test_manual_example(self): + # manager_bs=4, group_bs=2, kernel_bs=2; block_table maps logical + # blocks 0..3 -> physical 5,2,9,0. Manager block 1 covers tokens 4..7 -> + # logical blocks 2,3 -> physical 9,0 -> slots 18,19,0,1. + conn = make_connector(manager_block_size=4) + group = _make_group(group_bs=2, kernel_bs=2) + got = conn._attn_token_indices(group, [1], [5, 2, 9, 0]) + self.assertEqual(got, [[18, 19, 0, 1]]) + + def test_out_of_range_asserts(self): + conn = make_connector(manager_block_size=16) + group = _make_group(group_bs=16, kernel_bs=16) + with self.assertRaises(AssertionError): + conn._attn_token_indices(group, [1], [0]) # table too short + + +class TestStateBlockIds(unittest.TestCase): + def test_against_reference(self): + for manager_bs, group_bs in [(528, 528), (16, 16), (16, 32), (48, 16)]: + with self.subTest(manager_bs=manager_bs, group_bs=group_bs): + conn = make_connector(manager_block_size=manager_bs) + group = _make_group(group_bs, is_attention=False) + needed = 4 * manager_bs // group_bs + 1 + block_table = [(i * 11 + 5) % 89 for i in range(needed)] + mbis = [0, 1, 3] + got = conn._state_block_ids(group, mbis, block_table) + want = _ref_state_block_ids(manager_bs, group_bs, mbis, block_table) + self.assertEqual(got, want) + + def test_manual_example(self): + # manager_bs=4, group_bs=8: manager blocks 0 and 1 both end inside group + # block 0; manager block 2 ends in group block 1. + conn = make_connector(manager_block_size=4) + group = _make_group(group_bs=8, is_attention=False) + got = conn._state_block_ids(group, [0, 1, 2], [7, 3]) + self.assertEqual(got, [7, 7, 3]) + + def test_out_of_range_asserts(self): + conn = make_connector(manager_block_size=16) + group = _make_group(group_bs=16, is_attention=False) + with self.assertRaises(AssertionError): + conn._state_block_ids(group, [2], [0, 1]) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/test_data_transfer_results.py b/kv_cache_manager/py_connector/test/test_data_transfer_results.py new file mode 100644 index 000000000..7e43e9f7e --- /dev/null +++ b/kv_cache_manager/py_connector/test/test_data_transfer_results.py @@ -0,0 +1,173 @@ +"""Unit tests for MultiResult flattening and the save/load done callbacks. + +The done callbacks decode a flat result list whose layout is an implicit +contract with ``_submit_group_tasks``: tasks are submitted group-major +(group0's blocks, then group1's blocks, ...), so a manager block's success is +the stride-AND ``flat[i % num_blocks]``. These tests pin that contract with +hand-computed expectations. +""" + +import threading +import unittest +from unittest.mock import MagicMock + +from kv_cache_manager.py_connector.test import vllm_stubs # noqa: F401 (stubs) +from kv_cache_manager.py_connector.vllm.data_transfer import ( + DataTransferManager, MultiResult) +from kv_cache_manager.py_connector.common.tp_coordinator import ( + CoordinateMsgSerializer) + + +class TestMultiResult(unittest.TestCase): + def test_flatten_in_submission_order(self): + got = [] + mr = MultiResult(3, got.extend) + mr.submit_result(0, [True, False]) + mr.submit_result(1, [False]) + mr.submit_result(2, [True, True, True]) + self.assertEqual(got, [True, False, False, True, True, True]) + + def test_out_of_order_submit(self): + got = [] + mr = MultiResult(3, got.extend) + mr.submit_result(2, ["c"]) + mr.submit_result(0, ["a"]) + self.assertEqual(got, []) # callback must not fire early + mr.submit_result(1, ["b"]) + self.assertEqual(got, ["a", "b", "c"]) + + def test_duplicate_submit_asserts(self): + mr = MultiResult(2, lambda flat: None) + mr.submit_result(0, [True]) + with self.assertRaises(AssertionError): + mr.submit_result(0, [True]) + + def test_concurrent_submit(self): + n = 64 + results = [] + done = threading.Event() + + def cb(flat): + results.append(flat) + done.set() + + mr = MultiResult(n, cb) + barrier = threading.Barrier(n) + + def worker(i): + barrier.wait() + mr.submit_result(i, [i]) + + threads = [threading.Thread(target=worker, args=(i,)) for i in range(n)] + for t in threads: + t.start() + for t in threads: + t.join() + self.assertTrue(done.wait(timeout=5)) + self.assertEqual(len(results), 1) # callback fires exactly once + self.assertEqual(results[0], list(range(n))) + + +def _make_dtm(): + """DataTransferManager with only the state the callbacks touch.""" + dtm = DataTransferManager.__new__(DataTransferManager) + dtm._coordinator_client = MagicMock() + return dtm + + +def _sent_event(dtm): + (payload,), _ = dtm._coordinator_client.send.call_args + return CoordinateMsgSerializer.loads(payload).content + + +class TestSaveDoneCallback(unittest.TestCase): + def test_multi_group_stride_and(self): + # 3 blocks x 2 groups, flat = group0[b0,b1,b2] + group1[b0,b1,b2]. + # Block b is saved only if both groups succeeded for b. + dtm = _make_dtm() + cb = dtm.create_save_done_callback("req", 0, "sess", num_blocks=3) + cb([True, True, False, # group 0 + True, False, True]) # group 1 + evt = _sent_event(dtm) + self.assertEqual(evt.type, "SendBlockFinishedEvent") + self.assertEqual(evt.write_session_id, "sess") + self.assertEqual(evt.is_success_list, [True, False, False]) + + def test_single_group_passthrough(self): + dtm = _make_dtm() + cb = dtm.create_save_done_callback("req", 1, "sess", num_blocks=2) + cb([False, True]) + self.assertEqual(_sent_event(dtm).is_success_list, [False, True]) + + +class TestLoadDoneCallback(unittest.TestCase): + def test_multi_group_failure_merge(self): + dtm = _make_dtm() + cb = dtm.create_load_done_callback( + "req", 0, epoch=7, block_ids=[10, 20, 30], num_blocks=3) + cb([True, False, True, # group 0 + True, True, False]) # group 1 + evt = _sent_event(dtm) + self.assertEqual(evt.type, "LoadBlockFinishedEvent") + self.assertEqual(evt.epoch, 7) + # blocks 1 and 2 each failed in one group -> report their table ids. + self.assertEqual(evt.failed_block_idxs, [20, 30]) + + def test_all_success_reports_empty(self): + dtm = _make_dtm() + cb = dtm.create_load_done_callback( + "req", 0, epoch=0, block_ids=[10, 20], num_blocks=2) + cb([True, True, True, True]) + self.assertEqual(_sent_event(dtm).failed_block_idxs, []) + + def test_report_failures_false_hybrid(self): + # Hybrid models cannot report invalid block ids to vLLM: the failure + # must be swallowed (empty failed list) but the finished event still sent. + dtm = _make_dtm() + cb = dtm.create_load_done_callback( + "req", 0, epoch=1, block_ids=[], num_blocks=2, report_failures=False) + cb([False, True]) + evt = _sent_event(dtm) + self.assertEqual(evt.type, "LoadBlockFinishedEvent") + self.assertEqual(evt.failed_block_idxs, []) + + +class TestNullStateBlocks(unittest.TestCase): + """Mamba 'align' mode: null (id 0) state targets carry no state by design + (vLLM materializes states only at segment boundaries), so save/load must + treat them as vacuous successes -- failing them would stride-AND whole + manager blocks out of the manager's prefix chain and kill multi-block + hybrid caching. The all-null path takes no GPU work, so it runs on CPU.""" + + @staticmethod + def _state_group(): + from kv_cache_manager.py_connector.common.types import TransferGroup + return TransferGroup( + group_idx=0, spec_name="tp0_g0", is_attention=False, + layer_names=["m0"], block_size=528, per_block_bytes=1024, + kernel_block_size=528) + + def test_save_all_null_blocks_vacuously_succeed(self): + dtm = _make_dtm() + results = {} + mr = MultiResult(1, lambda flat: results.setdefault("flat", flat)) + dtm.save_task(mr, 0, self._state_group(), + remote_uris=["u0", "u1"], + block_token_indices=None, + block_ids=[0, 0], + ready_event=None) + self.assertEqual(results["flat"], [True, True]) + + def test_load_all_null_blocks_vacuously_succeed(self): + dtm = _make_dtm() + results = {} + mr = MultiResult(1, lambda flat: results.setdefault("flat", flat)) + dtm.load_task(mr, 0, self._state_group(), + remote_uris=["u0", "u1", "u2"], + block_token_indices=None, + block_ids=[0, 0, 0]) + self.assertEqual(results["flat"], [True, True, True]) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/test_scheduler_state.py b/kv_cache_manager/py_connector/test/test_scheduler_state.py new file mode 100644 index 000000000..c4f70ead6 --- /dev/null +++ b/kv_cache_manager/py_connector/test/test_scheduler_state.py @@ -0,0 +1,413 @@ +"""Unit tests for the connector's scheduler-side logic. + +Covers, against fake vLLM SchedulerOutput / Request objects: + +* ``get_num_new_matched_tokens`` -- including the full-prompt external hit cap + (a fully cached prompt must leave >= 1 token to recompute, otherwise vLLM's + synchronous-load scheduling path asserts ``num_new_tokens > 0``); +* ``parse_block_mask_to_save_indices`` -- ``offset`` and ``bool_masks`` forms; +* ``_parse_groups`` -- full-attention single group, hybrid multi group, eagle + group skip, unsupported spec error; +* ``build_connector_meta`` -- new request, cached deltas (``new_block_ids`` + None / non-None), preemption resume via both the 0.26 ``resumed_req_ids`` and + the legacy ``resumed_from_preemption`` interfaces, save-threshold trigger, + and the two ``request_finished`` paths (saves landed / in-flight). +""" + +import unittest +from dataclasses import dataclass, field +from types import SimpleNamespace +from unittest.mock import MagicMock + +from kv_cache_manager.py_connector.test.vllm_stubs import ( + make_connector, ReqState, GroupMeta) +from kv_cache_manager.py_connector.vllm.v1_connector import TairKvCacheConnector +from kv_cache_manager.py_connector.vllm.metadata import SaveRequest + + +# --------------------------------------------------------------------------- # +# Fakes +# --------------------------------------------------------------------------- # +@dataclass +class FakeRequest: + request_id: str + prompt_token_ids: list + output_token_ids: list = field(default_factory=list) + + @property + def num_tokens(self): + return len(self.prompt_token_ids) + len(self.output_token_ids) + + @property + def all_token_ids(self): + return self.prompt_token_ids + self.output_token_ids + + +def make_scheduler_connector(mbs=16, vllm_bs=None, locations=None): + """Connector with the scheduler-side state build_connector_meta needs.""" + conn = make_connector(manager_block_size=mbs, vllm_block_size=vllm_bs) + conn._epoch = 0 + conn._alive_requests = {} + conn._waiting_to_load_requests = [] + import threading + conn._waiting_to_save_requests_lock = threading.Lock() + conn._waiting_to_save_requests = [] + conn._waiting_to_finish_requests = [] + conn._canceled_save_request_ids_lock = threading.Lock() + conn._canceled_save_request_ids = [] + conn._http_executor = MagicMock() + conn._location_query_manager = MagicMock() + conn._location_query_manager.get_locations_for_query.return_value = ( + True, locations if locations is not None else []) + return conn + + +def fake_scheduler_output(new_reqs=(), cached_req_ids=(), num_scheduled=None, + new_block_ids=(), resumed_req_ids=frozenset(), + legacy_resumed=None): + """Build a fake SchedulerOutput. legacy_resumed switches the cached-reqs + container to the pre-0.26 interface (resumed_from_preemption list, no + resumed_req_ids attribute).""" + if legacy_resumed is not None: + cached = SimpleNamespace( + req_ids=list(cached_req_ids), + resumed_from_preemption=list(legacy_resumed), + new_block_ids=list(new_block_ids), + ) + else: + cached = SimpleNamespace( + req_ids=list(cached_req_ids), + resumed_req_ids=set(resumed_req_ids), + new_block_ids=list(new_block_ids), + ) + return SimpleNamespace( + scheduled_new_reqs=list(new_reqs), + scheduled_cached_reqs=cached, + num_scheduled_tokens=dict(num_scheduled or {}), + ) + + +def make_locations(n): + return [{"location_specs": [{"name": "tp0_g0", "uri": f"file://blk{i}"}]} + for i in range(n)] + + +# --------------------------------------------------------------------------- # +# get_num_new_matched_tokens +# --------------------------------------------------------------------------- # +class TestGetNumNewMatchedTokens(unittest.TestCase): + MBS = 16 + + def _run(self, prompt_len, num_computed, num_locations): + conn = make_scheduler_connector( + mbs=self.MBS, locations=make_locations(num_locations)) + req = FakeRequest("r0", list(range(prompt_len))) + matched, async_load = conn.get_num_new_matched_tokens(req, num_computed) + return conn, matched, async_load + + def test_partial_hit_no_cap(self): + conn, matched, async_load = self._run(4 * self.MBS + 5, 0, 4) + self.assertEqual(matched, 4 * self.MBS) + self.assertTrue(async_load) # a pending load is reported as async + self.assertEqual(conn._waiting_to_load_requests[0].manager_block_idxes, + [0, 1, 2, 3]) + + def test_full_hit_capped_to_leave_one_token(self): + # Prompt is exactly 4 manager blocks, all externally cached: the last + # block must be dropped so vLLM still schedules >= 1 new token. + conn, matched, _ = self._run(4 * self.MBS, 0, 4) + self.assertEqual(matched, 3 * self.MBS) + self.assertEqual(conn._waiting_to_load_requests[0].manager_block_idxes, + [0, 1, 2]) + # has_saved_block_num counts only the blocks actually treated as hit. + self.assertEqual(conn._alive_requests["r0"].has_saved_block_num, 3) + + def test_full_hit_with_local_prefix(self): + # 2 blocks locally computed + 2 remote = whole prompt -> drop one remote. + conn, matched, _ = self._run(4 * self.MBS, 2 * self.MBS, 2) + self.assertEqual(matched, self.MBS) + self.assertEqual(conn._waiting_to_load_requests[0].manager_block_idxes, [2]) + + def test_single_block_full_hit_degrades_to_zero(self): + conn, matched, async_load = self._run(self.MBS, 0, 1) + self.assertEqual(matched, 0) + self.assertFalse(async_load) + self.assertEqual(conn._waiting_to_load_requests, []) + + def test_no_locations(self): + conn, matched, async_load = self._run(100, 0, 0) + self.assertEqual(matched, 0) + self.assertFalse(async_load) + + def test_requery_after_load_attempt_skips_external(self): + # A request that already went through an external load (blocks were + # allocated) and returned to WAITING -- KV load failure with + # policy=recompute, or preemption -- must not re-match: the manager may + # still advertise blocks whose storage is gone, and re-matching loops + # fail -> reschedule forever. + conn = make_scheduler_connector(mbs=self.MBS, locations=make_locations(2)) + req = FakeRequest("r0", list(range(4 * self.MBS + 5))) + matched, _ = conn.get_num_new_matched_tokens(req, 0) + self.assertEqual(matched, 2 * self.MBS) + # vLLM allocates blocks for the load attempt. + conn.update_state_after_alloc( + req, SimpleNamespace(get_block_ids=lambda: [[100, 101, 102]]), matched) + # Retry: same request re-enters the waiting queue with 0 computed. + matched2, async2 = conn.get_num_new_matched_tokens(req, 0) + self.assertEqual(matched2, 0) + self.assertFalse(async2) + self.assertEqual(len(conn._waiting_to_load_requests), 1) # no new load + state = conn._alive_requests["r0"] + self.assertEqual(state.remote_matched_token_num, 0) + self.assertEqual(state.has_saved_block_num, 0) + + +# --------------------------------------------------------------------------- # +# parse_block_mask_to_save_indices +# --------------------------------------------------------------------------- # +class TestParseBlockMask(unittest.TestCase): + def setUp(self): + self.conn = make_connector() + + def test_offset_branch(self): + resp = {"block_mask": {"offset": 2}} + self.assertEqual( + self.conn.parse_block_mask_to_save_indices(resp, 5), [2, 3, 4]) + + def test_offset_zero(self): + resp = {"block_mask": {"offset": 0}} + self.assertEqual( + self.conn.parse_block_mask_to_save_indices(resp, 3), [0, 1, 2]) + + def test_bool_masks_branch(self): + resp = {"block_mask": {"bool_masks": {"values": [True, False, True, False]}}} + self.assertEqual( + self.conn.parse_block_mask_to_save_indices(resp, 4), [1, 3]) + + def test_missing_mask(self): + self.assertEqual(self.conn.parse_block_mask_to_save_indices({}, 3), []) + + +# --------------------------------------------------------------------------- # +# _parse_groups +# --------------------------------------------------------------------------- # +class TestParseGroups(unittest.TestCase): + def _kv_cache_config(self, groups): + return SimpleNamespace(kv_cache_groups=groups) + + def _attn_group(self, layers, block_size=16, page_size_bytes=32768): + from vllm.v1.kv_cache_interface import FullAttentionSpec + return SimpleNamespace( + layer_names=layers, + kv_cache_spec=FullAttentionSpec(block_size, page_size_bytes)) + + def _mamba_group(self, layers, block_size=528, page_size_bytes=1024): + from vllm.v1.kv_cache_interface import MambaSpec + return SimpleNamespace( + layer_names=layers, + kv_cache_spec=MambaSpec(block_size, page_size_bytes)) + + def test_full_attention_single_group(self): + conn = make_connector(manager_block_size=32) + metas = conn._parse_groups(self._kv_cache_config( + [self._attn_group(["l0", "l1"], block_size=16, page_size_bytes=32768)])) + self.assertEqual(len(metas), 1) + m = metas[0] + self.assertTrue(m.is_attention) + self.assertEqual(m.group_idx, 0) + self.assertEqual(m.block_size, 16) + # per_token = 32768 // 16 = 2048; per_block = 2048 * 32 (manager) * 2 layers + self.assertEqual(m.per_block_bytes, 2048 * 32 * 2) + + def test_hybrid_multi_group(self): + conn = make_connector(manager_block_size=528) + metas = conn._parse_groups(self._kv_cache_config([ + self._mamba_group(["m0", "m1"], page_size_bytes=1000), + self._mamba_group(["m2"], page_size_bytes=2000), + self._attn_group(["a0"], block_size=528, page_size_bytes=528 * 64), + ])) + self.assertEqual([m.group_idx for m in metas], [0, 1, 2]) + self.assertEqual([m.is_attention for m in metas], [False, False, True]) + self.assertEqual(metas[0].per_block_bytes, 1000 * 2) # page * layers + self.assertEqual(metas[1].per_block_bytes, 2000) + self.assertEqual(metas[2].per_block_bytes, 64 * 528) # per_token * mbs + + def test_eagle_group_skipped(self): + conn = make_connector() + eagle = self._attn_group(["drafter"]) + eagle.is_eagle_group = True + metas = conn._parse_groups(self._kv_cache_config( + [eagle, self._attn_group(["a0"])])) + self.assertEqual(len(metas), 1) + self.assertEqual(metas[0].layer_names, ["a0"]) + self.assertEqual(metas[0].group_idx, 1) # group_idx keeps vLLM numbering + + def test_unsupported_spec_raises(self): + conn = make_connector() + bad = SimpleNamespace(layer_names=["x"], kv_cache_spec=object()) + with self.assertRaises(NotImplementedError): + conn._parse_groups(self._kv_cache_config([bad])) + + def test_no_usable_groups_asserts(self): + conn = make_connector() + with self.assertRaises(AssertionError): + conn._parse_groups(self._kv_cache_config([])) + + +# --------------------------------------------------------------------------- # +# build_connector_meta +# --------------------------------------------------------------------------- # +class TestBuildConnectorMeta(unittest.TestCase): + MBS = 16 + + def _new_request(self, conn, req_id, num_tokens, num_blocks, + num_locations=0): + """Simulate the scheduler flow for a fresh request: query, alloc, then + one build_connector_meta step.""" + conn._location_query_manager.get_locations_for_query.return_value = ( + True, make_locations(num_locations)) + req = FakeRequest(req_id, list(range(num_tokens))) + conn.get_num_new_matched_tokens(req, 0) + block_ids = [list(range(100, 100 + num_blocks))] + conn.update_state_after_alloc( + req, SimpleNamespace(get_block_ids=lambda: block_ids), 0) + out = fake_scheduler_output( + new_reqs=[SimpleNamespace(req_id=req_id, block_ids=block_ids)]) + return req, conn.build_connector_meta(out) + + def test_new_request_full_state(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, meta = self._new_request(conn, "r0", 40, 3) + self.assertEqual(len(meta.requests), 1) + delta = meta.requests[0] + self.assertFalse(delta.is_delta) + self.assertEqual(delta.new_tokens_ids, list(range(40))) + self.assertEqual(delta.new_block_ids_per_group, [[100, 101, 102]]) + # 40 tokens / 3 blocks -> min(40, 48)//16 = 2 blocks to save. + conn._http_executor.submit.assert_called_once() + args = conn._http_executor.submit.call_args[0] + self.assertEqual(args[1:], ("r0", list(range(32)), 2)) + self.assertEqual(conn._alive_requests["r0"].has_saved_block_num, 2) + + def test_load_request_emitted_after_alloc(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, meta = self._new_request(conn, "r0", 40, 3, num_locations=2) + self.assertEqual(len(meta.to_load_requests), 1) + lr = meta.to_load_requests[0] + self.assertEqual(lr.manager_block_idxes, [0, 1]) + self.assertEqual(lr.all_block_ids, [[100, 101, 102]]) + # Externally hit blocks are not re-saved. + conn._http_executor.submit.assert_not_called() + + def test_cached_delta_with_and_without_new_blocks(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + # Step 2: 8 decode tokens, no new blocks (PR #23262: may be None). + req.output_token_ids = list(range(1000, 1008)) + out = fake_scheduler_output( + cached_req_ids=["r0"], num_scheduled={"r0": 8}, new_block_ids=[None]) + meta = conn.build_connector_meta(out) + delta = meta.requests[0] + self.assertTrue(delta.is_delta) + self.assertEqual(delta.new_tokens_ids, list(range(1000, 1008))) + self.assertEqual(delta.new_block_ids_per_group, []) + # Step 3: 2 more tokens with a new block -> table grows. + req.output_token_ids = list(range(1000, 1010)) + out = fake_scheduler_output( + cached_req_ids=["r0"], num_scheduled={"r0": 2}, + new_block_ids=[[[103]]]) + meta = conn.build_connector_meta(out) + self.assertEqual(meta.requests[0].new_block_ids_per_group, [[103]]) + self.assertEqual(conn._alive_requests["r0"].block_ids_per_group, + [[100, 101, 102, 103]]) + + def _preempted_step(self, conn, req, use_legacy): + kwargs = dict(cached_req_ids=["r0"], num_scheduled={"r0": 0}, + new_block_ids=[[[200, 201]]]) + if use_legacy: + kwargs["legacy_resumed"] = [True] + else: + kwargs["resumed_req_ids"] = {"r0"} + return conn.build_connector_meta(fake_scheduler_output(**kwargs)) + + def test_resumed_from_preemption_both_interfaces(self): + for use_legacy in (False, True): + with self.subTest(legacy=use_legacy): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + meta = self._preempted_step(conn, req, use_legacy) + delta = meta.requests[0] + self.assertTrue(delta.resumed_from_preemption) + # Resume replaces (not extends) the block table. + self.assertEqual(conn._alive_requests["r0"].block_ids_per_group, + [[200, 201]]) + + def test_save_threshold_grows_incrementally(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) # saved 2 blocks + conn._http_executor.submit.reset_mock() + # 8 more tokens -> 48 total, table full at 3 blocks -> third block saves. + req.output_token_ids = list(range(1000, 1008)) + out = fake_scheduler_output( + cached_req_ids=["r0"], num_scheduled={"r0": 8}, new_block_ids=[[[103]]]) + conn.build_connector_meta(out) + args = conn._http_executor.submit.call_args[0] + self.assertEqual(args[3], 3) # target_save_num + self.assertEqual(conn._alive_requests["r0"].has_saved_block_num, 3) + + def test_save_request_drain_and_finish_paths(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + state = conn._alive_requests["r0"] + self.assertEqual(state.scheduled_saving_count, 1) + + # Finish while the save is still in flight: request must stay alive. + keep, extra = conn.request_finished(req, []) + self.assertTrue(keep) + self.assertTrue(state.need_report_after_saving_finished) + self.assertIn("r0", conn._alive_requests) + + # The async save lands: drained into to_save_requests and, because the + # request already finished, a FinishRequest is emitted and state dropped. + with conn._waiting_to_save_requests_lock: + conn._waiting_to_save_requests.append( + SaveRequest("r0", make_locations(2), [0, 1], "sess")) + meta = conn.build_connector_meta(fake_scheduler_output()) + self.assertEqual(len(meta.to_save_requests), 1) + self.assertEqual([f.req_id for f in meta.to_finish_requests], ["r0"]) + self.assertNotIn("r0", conn._alive_requests) + + def test_request_finished_when_saves_landed(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + with conn._waiting_to_save_requests_lock: + conn._waiting_to_save_requests.append( + SaveRequest("r0", make_locations(2), [0, 1], "sess")) + conn.build_connector_meta(fake_scheduler_output()) + keep, extra = conn.request_finished(req, []) + self.assertTrue(keep) + self.assertNotIn("r0", conn._alive_requests) + meta = conn.build_connector_meta(fake_scheduler_output()) + self.assertEqual([f.req_id for f in meta.to_finish_requests], ["r0"]) + + def test_canceled_save_unknown_request_no_crash(self): + # Cancellations arrive from http_executor threads and may race request + # teardown; an unknown req_id must be skipped, not KeyError. + conn = make_scheduler_connector(mbs=self.MBS) + with conn._canceled_save_request_ids_lock: + conn._canceled_save_request_ids.append("ghost") + conn.build_connector_meta(fake_scheduler_output()) # must not raise + + def test_canceled_save_finishes_request(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + conn.request_finished(req, []) # save in flight -> delayed finish + with conn._canceled_save_request_ids_lock: + conn._canceled_save_request_ids.append("r0") + meta = conn.build_connector_meta(fake_scheduler_output()) + self.assertEqual([f.req_id for f in meta.to_finish_requests], ["r0"]) + self.assertNotIn("r0", conn._alive_requests) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/vllm_stubs.py b/kv_cache_manager/py_connector/test/vllm_stubs.py new file mode 100644 index 000000000..8c6644c17 --- /dev/null +++ b/kv_cache_manager/py_connector/test/vllm_stubs.py @@ -0,0 +1,133 @@ +"""Shared test stubs: make ``v1_connector`` importable without vLLM/CUDA/pybind. + +``v1_connector`` imports vLLM and the compiled ``kvcm_py_client`` at module +level. For pure-logic unit tests we register lightweight stand-ins in +``sys.modules`` *before* the first import, then build connector instances via +``__new__`` with only the attributes the code under test reads. No production +module is modified. +""" + +import sys +import types +from typing import Optional +from unittest.mock import MagicMock + + +def _module(name: str) -> types.ModuleType: + mod = sys.modules.get(name) + if mod is None: + mod = types.ModuleType(name) + sys.modules[name] = mod + return mod + + +def _install_stubs(): + existing = sys.modules.get("vllm") + if existing is not None: + # Either our stub is already in place or the real vLLM is importable; + # in both cases the connector import will succeed as-is. + return + + # ---- kv_cache_manager.client.pybind (compiled extension) ---- + pybind = _module("kv_cache_manager.client.pybind") + kvcm_py_client = MagicMock() + kvcm_py_client.ClientErrorCode.ER_OK = 0 + pybind.kvcm_py_client = kvcm_py_client + + # ---- kv_cache_manager.py_connector.common._version_info (generated) ---- + version = _module("kv_cache_manager.py_connector.common._version_info") + version.FULL_VERSION = "0.0.0-test" + version.GIT_COMMIT = "test" + version.BUILD_TIME = "test" + + # ---- vllm ---- + vllm = _module("vllm") + vllm._kvcm_test_stub = True + + config = _module("vllm.config") + config.VllmConfig = MagicMock + vllm.config = config + + distributed = _module("vllm.distributed") + distributed.get_tensor_model_parallel_rank = lambda: 0 + vllm.distributed = distributed + _module("vllm.distributed.kv_transfer") + _module("vllm.distributed.kv_transfer.kv_connector") + _module("vllm.distributed.kv_transfer.kv_connector.v1") + base = _module("vllm.distributed.kv_transfer.kv_connector.v1.base") + + class KVConnectorRole: + SCHEDULER = 0 + WORKER = 1 + + class KVConnectorMetadata: + pass + + class KVConnectorBase_V1: + def __init__(self, vllm_config, role, kv_cache_config=None): + self._connector_metadata = None + + def _get_connector_metadata(self): + return self._connector_metadata + + class SupportsHMA: + pass + + base.KVConnectorBase_V1 = KVConnectorBase_V1 + base.KVConnectorMetadata = KVConnectorMetadata + base.KVConnectorRole = KVConnectorRole + base.SupportsHMA = SupportsHMA + + utils = _module("vllm.utils") + torch_utils = _module("vllm.utils.torch_utils") + torch_utils.get_kv_cache_torch_dtype = MagicMock() + network_utils = _module("vllm.utils.network_utils") + network_utils.get_ip = lambda: "127.0.0.1" + utils.torch_utils = torch_utils + utils.network_utils = network_utils + + v1 = _module("vllm.v1") + kv_cache_interface = _module("vllm.v1.kv_cache_interface") + + class FullAttentionSpec: + def __init__(self, block_size, page_size_bytes): + self.block_size = block_size + self.page_size_bytes = page_size_bytes + + class MambaSpec: + def __init__(self, block_size, page_size_bytes): + self.block_size = block_size + self.page_size_bytes = page_size_bytes + + kv_cache_interface.FullAttentionSpec = FullAttentionSpec + kv_cache_interface.MambaSpec = MambaSpec + + _module("vllm.v1.core") + sched = _module("vllm.v1.core.sched") + output = _module("vllm.v1.core.sched.output") + output.SchedulerOutput = MagicMock + sched.output = output + + outputs = _module("vllm.v1.outputs") + outputs.KVConnectorOutput = MagicMock + v1.kv_cache_interface = kv_cache_interface + v1.outputs = outputs + + +_install_stubs() + +# Import after stubs are in place. +from kv_cache_manager.py_connector.vllm.v1_connector import ( # noqa: E402 + TairKvCacheConnector, GroupMeta, ReqState) + + +def make_connector(manager_block_size: int = 16, + vllm_block_size: Optional[int] = None, + num_groups: int = 1) -> TairKvCacheConnector: + """Build a bare TairKvCacheConnector (no __init__) with the minimal state + used by the pure translation / scheduler-side logic under test.""" + conn = TairKvCacheConnector.__new__(TairKvCacheConnector) + conn._manager_block_size = manager_block_size + conn._vllm_block_size = vllm_block_size or manager_block_size + conn._num_groups = num_groups + return conn diff --git a/kv_cache_manager/py_connector/vllm/BUILD b/kv_cache_manager/py_connector/vllm/BUILD index 2fc64b26c..acf24672e 100644 --- a/kv_cache_manager/py_connector/vllm/BUILD +++ b/kv_cache_manager/py_connector/vllm/BUILD @@ -7,6 +7,7 @@ load("@python_platform//:platform.bzl", "python_platform") py_library( name = "vllm_connector", srcs = glob(["*.py"]), + visibility = ["//kv_cache_manager/py_connector:__subpackages__"], deps = ["//kv_cache_manager/py_connector/common:common", "//kv_cache_manager/py_connector/kernel:kernel", "//kv_cache_manager/client/pybind:kvcm_py_client_lib"] diff --git a/kv_cache_manager/py_connector/vllm/data_transfer.py b/kv_cache_manager/py_connector/vllm/data_transfer.py index 36a16323b..0642d1fe3 100644 --- a/kv_cache_manager/py_connector/vllm/data_transfer.py +++ b/kv_cache_manager/py_connector/vllm/data_transfer.py @@ -1,15 +1,36 @@ -import time -import threading -from concurrent.futures.thread import ThreadPoolExecutor +"""Per-group KV cache transfer between vLLM's paged cache and KVCM storage. + +Each ``TransferGroup`` is an independent transfer unit: -from typing import Any +* Attention groups store token-granular KV; a manager block is gathered/scattered + through the strided Triton kernel (``batch_gather_scatter_helper``) which handles + both the contiguous full-attention layout and the block-strided hybrid layout. +* Mamba/linear/gdn groups store per-block opaque state; a manager block maps to a + single logical block whose raw bytes are copied verbatim. + +The transport itself is layout-agnostic: for every manager block we hand the SDK a +``BlockBuffer`` (a pinned CPU region) and the block's remote URI. Save gathers HBM +-> CPU then ``SaveKvCaches``; load ``LoadKvCaches`` -> CPU then scatters CPU -> HBM. +""" + +import threading +import time +from concurrent.futures import ThreadPoolExecutor import torch from kv_cache_manager.client.pybind import kvcm_py_client +from kv_cache_manager.py_connector.common.tp_coordinator import ( + CoordinateMsgSerializer, TpCoordinatorClient, CoordinateMessage, + SendBlockFinishedEvent, LoadBlockFinishedEvent, +) +from kv_cache_manager.py_connector.common.logger import logger +from kv_cache_manager.py_connector.common.types import KVCacheInfo, TransferGroup +from kv_cache_manager.py_connector.kernel import batch_gather_scatter_helper + def _get_device_module(device=None): - """Return the device module matching the runtime device.""" + """Return the torch device module matching the runtime device.""" if device is not None and hasattr(torch, "get_device_module"): return torch.get_device_module(device) try: @@ -20,25 +41,18 @@ def _get_device_module(device=None): pass return torch.cuda -from kv_cache_manager.py_connector.common.tp_coordinator import CoordinateMsgSerializer, TpCoordinatorClient, \ - CoordinateMessage, SendBlockFinishedEvent, LoadBlockFinishedEvent -from kv_cache_manager.py_connector.common.logger import logger -from kv_cache_manager.py_connector.common.types import KVCacheInfo -from kv_cache_manager.py_connector.kernel import batch_gather_scatter_helper -from kv_cache_manager.py_connector.kernel.gather_scatter_helper import CopyBufferAllocator - class MultiResult: - """多任务结果管理类 - - 用于管理多个异步任务的结果, 当所有任务完成时触发回调 - """ + """Collect the per-block success flags of several async tasks and fire a + callback once every task has reported. Each result is a list[bool] aligned + with the manager blocks the task handled (in submission order).""" + def __init__(self, size: int, callback): - self._size: int = size + self._size = size self._results = [None] * size self._lock = threading.Lock() - self._finished_num: int = 0 - self._finished_callback = callback + self._finished_num = 0 + self._callback = callback def submit_result(self, idx: int, result): with self._lock: @@ -46,247 +60,215 @@ def submit_result(self, idx: int, result): self._results[idx] = result self._finished_num += 1 if self._finished_num == self._size: - self._finished_callback(self._results) + # Flatten in submission order. + flat = [ok for part in self._results for ok in part] + self._callback(flat) class DataTransferManager: - """KVCache数据传输核心类 - - 负责实际的KV缓存保存和加载操作, 包括: - 1. 保存任务(save_task) - 2. 加载任务(load_task) - 3. 回调创建(_create_save_done_callback, _create_load_done_callback) - """ - - def __init__(self, - kvcache_info: KVCacheInfo, - manager_block_size: int, - copy_buffer_allocator: CopyBufferAllocator, - transfer_client: Any, - coordinator_client: TpCoordinatorClient, - extra_config: Any): - """ - 初始化KV数据传输器 - - Args: - kvcache_info: KV缓存信息 - manager_block_size: instance的block_size - copy_buffer_allocator: 复制缓冲区分配器 - transfer_client: 传输客户端 - coordinator_client: 协调器客户端 - extra_config: 额外配置 - """ - self._kvcache_info = kvcache_info + def __init__(self, kvcache_info: KVCacheInfo, manager_block_size: int, + transfer_client, coordinator_client: TpCoordinatorClient, extra_config): + self._info = kvcache_info self._manager_block_size = manager_block_size - self._copy_buffer_allocator = copy_buffer_allocator self._transfer_client = transfer_client self._coordinator_client = coordinator_client self._extra_config = extra_config - self._device_mod = _get_device_module(self._kvcache_info.device) - - # 创建内部线程池执行器 - self._io_executor = self._create_io_executor() - - # 保存和加载流 + self._device = kvcache_info.device + self._device_mod = _get_device_module(self._device) self._save_stream = self._device_mod.Stream() self._load_stream = self._device_mod.Stream() - - def _create_io_executor(self) -> ThreadPoolExecutor: - """创建IO线程池执行器""" - from concurrent.futures import ThreadPoolExecutor - - # 初始化线程池,设置线程名和初始化函数 - def init_worker(): - import torch - self._device_mod.set_device(self._kvcache_info.device) - - return ThreadPoolExecutor( - max_workers=32, - thread_name_prefix="kvcm_io_", - initializer=init_worker - ) - - def submit_task(self, func, *args, **kwargs): - """提交任务到内部线程池 - - Args: - func: 要执行的函数 - *args: 函数参数 - **kwargs: 函数关键字参数 - - Returns: - Future对象 - """ - return self._io_executor.submit(func, *args, **kwargs) - def load_task(self, multi_result: MultiResult, task_idx, remote_uris, block_token_indices): - """加载任务 - - Args: - multi_result: 多任务结果管理器 - task_idx: 任务索引 - remote_uris: 远程URI列表 - block_token_indices: 块令牌索引列表 - """ - logger.debug("load remote_uris:%s, block_token_indices:%s", remote_uris, block_token_indices) + def _init_worker(): + self._device_mod.set_device(self._device) - copy_buffer_indices = self._copy_buffer_allocator.alloc_buffer_idx_blocking(len(remote_uris)) - copy_buffers = self._copy_buffer_allocator.get_buffer_by_idx(copy_buffer_indices) + self._io_executor = ThreadPoolExecutor( + max_workers=32, thread_name_prefix="kvcm_io_", initializer=_init_worker) + + def submit_task(self, func, *args, **kwargs): + return self._io_executor.submit(func, *args, **kwargs) + # ------------------------------------------------------------------ # + # BlockBuffer helper + # ------------------------------------------------------------------ # + @staticmethod + def _make_block_buffers(base_ptr: int, per_block_bytes: int, count: int): buffers = [] - for copy_buffer in copy_buffers: - buffer = kvcm_py_client.BlockBuffer() - iovs = [] + for i in range(count): + buf = kvcm_py_client.BlockBuffer() iov = kvcm_py_client.Iov() iov.type = kvcm_py_client.MemoryType.CPU - iov.base = copy_buffer.data_ptr() - iov.size = copy_buffer.nbytes + iov.base = base_ptr + i * per_block_bytes + iov.size = per_block_bytes iov.ignore = False - iovs.append(iov) - buffer.iovs = iovs - buffers.append(buffer) - logger.debug("start transfer") - transfer_result = self._transfer_client.LoadKvCaches(remote_uris, buffers) - logger.debug("done transfer,result:%s", transfer_result) - if transfer_result == kvcm_py_client.ClientErrorCode.ER_OK: - with self._device_mod.stream(self._load_stream): - batch_gather_scatter_helper.batch_scatter_kv_caches( - self._kvcache_info.all_kvcache_ptr_tensor_gpu, - self._copy_buffer_allocator._raw_buffer, - block_token_indices, - copy_buffer_indices, - self._manager_block_size, - self._kvcache_info.per_token_per_layer_dim_size, - ) + buf.iovs = [iov] + buffers.append(buf) + return buffers - copy_done_event = self._device_mod.Event() - copy_done_event.record(self._load_stream) - copy_done_event.synchronize() + # ------------------------------------------------------------------ # + # Save + # ------------------------------------------------------------------ # + def save_task(self, multi_result: MultiResult, task_idx, group: TransferGroup, + remote_uris, block_token_indices, block_ids, ready_event): + """Gather one group's manager blocks from HBM and save them. - logger.debug("done scatter") + block_token_indices: attention -> list[list[int]] flat token slots per block. + block_ids: state -> list[int] block id per manager block; + id 0 is vLLM's null block: no state exists at that + boundary by design (mamba "align" sparse states), + the block is reported saved vacuously. + """ + n = len(remote_uris) + if group.is_attention: + valid = list(range(n)) else: - logger.warning("load task failed, remote_uris:%s, block_token_indices:%s, transfer_result:%s", - remote_uris, - block_token_indices, transfer_result) - self._copy_buffer_allocator.free_buffer(copy_buffer_indices) - multi_result.submit_result(task_idx, [transfer_result] * len(remote_uris)) + # vLLM's mamba "align" mode only materializes the state block at + # segment boundaries (single_type_kv_cache_manager.MambaManager + # allocates the null block for intermediate positions; a hit only + # ever consumes the state ending the matched region). A null + # (id 0) state target therefore means "no state exists by design", + # not a failure: report it saved vacuously. Failing it would + # stride-AND the whole manager block out of the manager's prefix + # chain and kill multi-block hybrid caching entirely. + valid = [i for i in range(n) if block_ids[i] != 0] + if len(valid) < n: + logger.info("save group %s: %d/%d blocks have no materialized " + "state, saving vacuously", group.spec_name, + n - len(valid), n) + # Vacuous (skipped) blocks succeed; transferred blocks start False and + # are flipped by the transfer result below. + valid_set = set(valid) + ok_mask = [i not in valid_set for i in range(n)] + if valid: + cpu_buffer = torch.empty(len(valid) * group.per_block_bytes, dtype=torch.uint8, + device="cpu", pin_memory=True) + with self._device_mod.stream(self._save_stream): + ready_event.wait() + gpu_buffer = torch.empty(len(valid) * group.per_block_bytes, + dtype=torch.uint8, device=self._device) + if group.is_attention: + view = gpu_buffer.view(self._info.dtype).view( + len(valid), group.layer_num, + self._manager_block_size, group.per_token_dim) + batch_gather_scatter_helper.batch_gather_kv_caches( + group.kvcache_ptr_tensor_gpu, view, block_token_indices, + list(range(len(valid))), self._manager_block_size, + group.per_token_dim, + kv_stride=group.kv_stride, block_stride=group.block_stride, + local_block_size=group.kernel_block_size) + else: + for out_i, i in enumerate(valid): + for layer_idx in range(group.layer_num): + dst = (out_i * group.layer_num + layer_idx) * group.page_size_bytes + gpu_buffer[dst:dst + group.page_size_bytes].copy_( + group.block_view_tensors[layer_idx][block_ids[i]]) + cpu_buffer.copy_(gpu_buffer, non_blocking=True) + done = self._device_mod.Event() + done.record(self._save_stream) + done.synchronize() - def create_load_done_callback(self, req_id, tp_rank, epoch, local_block_ids): - """创建加载完成回调函数 - - Args: - req_id: 请求ID - tp_rank: TP rank - epoch - local_block_ids: 本地块ID列表 - - Returns: - 回调函数 - """ - def generate_message(task_results): - failed_block_idxs = [] - idx = 0 - for task_result in task_results: - for block_result in task_result: - if block_result != kvcm_py_client.ClientErrorCode.ER_OK: - failed_block_idxs.append(local_block_ids[idx]) - idx += 1 + buffers = self._make_block_buffers( + cpu_buffer.data_ptr(), group.per_block_bytes, len(valid)) + uris = [remote_uris[i] for i in valid] + result = self._transfer_client.SaveKvCaches(uris, buffers) + ok = (result[0] == kvcm_py_client.ClientErrorCode.ER_OK) + if not ok: + logger.warning("save task failed group=%s uris=%d result=%s", + group.spec_name, len(uris), result) + for i in valid: + ok_mask[i] = ok + multi_result.submit_result(task_idx, ok_mask) - msg = CoordinateMessage( - time.time(), - LoadBlockFinishedEvent(request_id=req_id, tp_rank=tp_rank, - epoch=epoch, failed_block_idxs=failed_block_idxs) - ) + def create_save_done_callback(self, req_id, tp_rank, write_session_id, num_blocks): + """block success = AND across all groups. task results are ordered + group0[blocks], group1[blocks], ... so we AND stride-wise.""" + def cb(flat): + is_success = [True] * num_blocks + for i, ok in enumerate(flat): + is_success[i % num_blocks] = is_success[i % num_blocks] and ok + msg = CoordinateMessage(time.time(), SendBlockFinishedEvent( + request_id=req_id, tp_rank=tp_rank, + write_session_id=write_session_id, is_success_list=is_success)) self._coordinator_client.send(CoordinateMsgSerializer.dumps(msg)) + return cb - return generate_message - - def save_task(self, multi_result: MultiResult, task_idx, remote_uris, block_token_indices, - kvcache_ready_event): - """保存任务 - - Args: - multi_result: 多任务结果管理器 - task_idx: 任务索引 - remote_uris: 远程URI列表 - block_token_indices: 块令牌索引列表 - kvcache_ready_event: KV缓存就绪事件 - """ - logger.debug("save remote_uris:%s, block_token_indices:%s", remote_uris, block_token_indices) - - with self._device_mod.stream(self._save_stream): - kvcache_ready_event.wait() - copy_buffer_indices = self._copy_buffer_allocator.alloc_buffer_idx_blocking(len(remote_uris)) - batch_gather_scatter_helper.batch_gather_kv_caches( - self._kvcache_info.all_kvcache_ptr_tensor_gpu, - self._copy_buffer_allocator._raw_buffer, - block_token_indices, - copy_buffer_indices, - self._manager_block_size, - self._kvcache_info.per_token_per_layer_dim_size, - ) - copy_done_event = self._device_mod.Event() - copy_done_event.record(self._save_stream) - - copy_done_event.synchronize() - - logger.debug("done gather") - - copy_buffers = self._copy_buffer_allocator.get_buffer_by_idx(copy_buffer_indices) - buffers = [] - for copy_buffer in copy_buffers: - buffer = kvcm_py_client.BlockBuffer() - iovs = [] - iov = kvcm_py_client.Iov() - iov.type = kvcm_py_client.MemoryType.CPU - iov.base = copy_buffer.data_ptr() - iov.size = copy_buffer.nbytes - iov.ignore = False - iovs.append(iov) - buffer.iovs = iovs - buffers.append(buffer) - logger.debug("start transfer") - - transfer_result = self._transfer_client.SaveKvCaches(remote_uris, buffers) - logger.debug("done transfer,result:%s", transfer_result) - if transfer_result[0] != kvcm_py_client.ClientErrorCode.ER_OK: - logger.warning("save task failed, remote_uris:%s, block_token_indices:%s, transfer_result:%s", remote_uris, - block_token_indices, transfer_result) - - self._copy_buffer_allocator.free_buffer(copy_buffer_indices) - # TODO: submit uri when enable local alloc - multi_result.submit_result(task_idx, [transfer_result[0]] * len(remote_uris)) + # ------------------------------------------------------------------ # + # Load + # ------------------------------------------------------------------ # + def load_task(self, multi_result: MultiResult, task_idx, group: TransferGroup, + remote_uris, block_token_indices, block_ids): + n = len(remote_uris) + if group.is_attention: + valid = list(range(n)) + else: + # Mirror of save_task: in mamba "align" mode vLLM only allocates a + # real state block for the final matched boundary; intermediate + # manager blocks get the null block (their state is not needed to + # resume). Skip them vacuously and load only materialized targets. + valid = [i for i in range(n) if block_ids[i] != 0] + if len(valid) < n: + logger.info("load group %s: %d/%d blocks have null state " + "targets, skipping them", group.spec_name, + n - len(valid), n) + valid_set = set(valid) + ok_mask = [i not in valid_set for i in range(n)] + if not valid: + multi_result.submit_result(task_idx, ok_mask) + return + cpu_buffer = torch.empty(len(valid) * group.per_block_bytes, dtype=torch.uint8, + device="cpu", pin_memory=True) + buffers = self._make_block_buffers(cpu_buffer.data_ptr(), + group.per_block_bytes, len(valid)) + uris = [remote_uris[i] for i in valid] + result = self._transfer_client.LoadKvCaches(uris, buffers) + ok = (result == kvcm_py_client.ClientErrorCode.ER_OK) + if ok: + with self._device_mod.stream(self._load_stream): + gpu_buffer = cpu_buffer.to(self._device, non_blocking=True) + if group.is_attention: + view = gpu_buffer.view(self._info.dtype).view( + n, group.layer_num, self._manager_block_size, group.per_token_dim) + batch_gather_scatter_helper.batch_scatter_kv_caches( + group.kvcache_ptr_tensor_gpu, view, block_token_indices, + list(range(n)), self._manager_block_size, group.per_token_dim, + kv_stride=group.kv_stride, block_stride=group.block_stride, + local_block_size=group.kernel_block_size) + else: + for out_i, i in enumerate(valid): + for layer_idx in range(group.layer_num): + src = (out_i * group.layer_num + layer_idx) * group.page_size_bytes + group.block_view_tensors[layer_idx][block_ids[i]].copy_( + gpu_buffer[src:src + group.page_size_bytes]) + done = self._device_mod.Event() + done.record(self._load_stream) + done.synchronize() + else: + logger.warning("load task failed group=%s uris=%d result=%s", + group.spec_name, len(uris), result) + for i in valid: + ok_mask[i] = ok + multi_result.submit_result(task_idx, ok_mask) - def create_save_done_callback(self, req_id, tp_rank, write_session_id): - """创建保存完成回调函数 - - Args: - req_id: 请求ID - tp_rank: TP rank - write_session_id: 写入会话ID - - Returns: - 回调函数 - """ - def generate_message(task_results): - is_successes = [] - # TODO: report uri when enable local alloc - # remote_uris = [] - for task_result in task_results: - for block_result in task_result: - if block_result != kvcm_py_client.ClientErrorCode.ER_OK: - is_successes.append(False) - # remote_uris.append(None) - else: - is_successes.append(True) - # remote_uris.extend(future_result[1]) + def create_load_done_callback(self, req_id, tp_rank, epoch, block_ids, num_blocks, + report_failures=True): + """A manager block is loaded only if every group succeeded for it. - msg = CoordinateMessage( - time.time(), - SendBlockFinishedEvent(request_id=req_id, tp_rank=tp_rank, - write_session_id=write_session_id, - is_success_list=is_successes) - ) + block_ids is the block table used to report vLLM-visible invalid block + ids. vLLM's invalid-block recovery only understands single-group block + tables, so multi-group (hybrid) connectors pass report_failures=False + and rely on request rescheduling instead.""" + def cb(flat): + merged = [True] * num_blocks + for i, ok in enumerate(flat): + merged[i % num_blocks] = merged[i % num_blocks] and ok + failed = [] + if report_failures: + failed = [block_ids[i] for i in range(min(num_blocks, len(block_ids))) + if not merged[i]] + elif not all(merged): + logger.warning("load failed for %d/%d blocks of req %s (hybrid: " + "not reporting invalid block ids)", + merged.count(False), num_blocks, req_id) + msg = CoordinateMessage(time.time(), LoadBlockFinishedEvent( + request_id=req_id, tp_rank=tp_rank, epoch=epoch, failed_block_idxs=failed)) self._coordinator_client.send(CoordinateMsgSerializer.dumps(msg)) - - return generate_message + return cb diff --git a/kv_cache_manager/py_connector/vllm/metadata.py b/kv_cache_manager/py_connector/vllm/metadata.py index d89d872b8..b63e1bf8c 100644 --- a/kv_cache_manager/py_connector/vllm/metadata.py +++ b/kv_cache_manager/py_connector/vllm/metadata.py @@ -1,56 +1,58 @@ from dataclasses import dataclass, field +from typing import List + from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorMetadata @dataclass class SaveRequest: req_id: str - target_locations: list[dict] - manager_block_idxes: list + # CacheLocation dicts returned by the manager (one per manager block to save), + # each carrying location_specs for every registered spec name. + target_locations: List[dict] + # Manager block indices (into the request's token stream) being saved. + manager_block_idxes: List[int] write_session_id: str -@dataclass() +@dataclass class LoadRequest: req_id: str - manager_block_idxes: list - need_load_locations: list[dict] - local_block_ids: list = field(default_factory=list) + manager_block_idxes: List[int] + need_load_locations: List[dict] + # Per-group block tables: all_block_ids[group_idx] is the list of local block + # ids for that kv_cache_group. Length 1 for pure-attention models. + all_block_ids: List[List[int]] = field(default_factory=list) -@dataclass() +@dataclass class FinishRequest: req_id: str @dataclass class ReqStateToWorker: - """发送给工作节点的请求状态数据结构""" + """Scheduler -> worker per-request state delta.""" req_id: str has_saved_block_num: int new_tokens_ids: list = field(default_factory=list) - new_local_block_ids: list = field(default_factory=list) + # Per-group new local block ids (indexed by kv_cache_group). + new_block_ids_per_group: List[List[int]] = field(default_factory=list) resumed_from_preemption: bool = False is_delta: bool = True + @dataclass class TairKvCacheConnectorMetadata(KVConnectorMetadata): - """TairKvCacheConnector的元数据类,用于在调度器和工作节点之间传递状态""" - requests: list[ReqStateToWorker] + """Scheduler -> worker metadata for one engine step.""" def __init__(self, epoch: int): - """ - 初始化元数据 - - Args: - epoch: 当前epoch编号 - """ self.epoch = epoch - self.requests: list[ReqStateToWorker] = [] - self.to_load_requests: list[LoadRequest] = [] - self.to_save_requests: list[SaveRequest] = [] - self.to_finish_requests: list[FinishRequest] = [] + self.requests: List[ReqStateToWorker] = [] + self.to_load_requests: List[LoadRequest] = [] + self.to_save_requests: List[SaveRequest] = [] + self.to_finish_requests: List[FinishRequest] = [] def add_req_state_to_worker(self, request: ReqStateToWorker): self.requests.append(request) @@ -65,5 +67,6 @@ def add_finish_request(self, finish_request: FinishRequest): self.to_finish_requests.append(finish_request) def __repr__(self): - return f"TairKvCacheConnectorMetadata(requests={self.requests})" - + return (f"TairKvCacheConnectorMetadata(epoch={self.epoch}, " + f"requests={len(self.requests)}, load={len(self.to_load_requests)}, " + f"save={len(self.to_save_requests)}, finish={len(self.to_finish_requests)})") diff --git a/kv_cache_manager/py_connector/vllm/v1_connector.py b/kv_cache_manager/py_connector/vllm/v1_connector.py index 3a0de4d98..8eeb4eb5f 100644 --- a/kv_cache_manager/py_connector/vllm/v1_connector.py +++ b/kv_cache_manager/py_connector/vllm/v1_connector.py @@ -1,13 +1,32 @@ +"""KVCM vLLM connector (v1), built around per-group transfer. + +vLLM models expose one or more ``kv_cache_groups`` (``KVCacheConfig``): + +* Pure-attention models: a single ``FullAttentionSpec`` group. +* Hybrid models (e.g. Qwen3.5): several ``MambaSpec`` groups plus one (or more) + ``FullAttentionSpec`` group. With ``mamba_cache_mode="align"`` every group has + its own block table (``block_ids`` is a tuple indexed by group) but all groups + share the scheduler block size. + +The connector treats every group as an independent transfer unit with its own +KVCM location spec (``tp{rank}_g{group}``), its own block table and its own data +access strategy (token-granular gather/scatter for attention, per-block byte +copy for mamba state). There is no separate "hybrid path": a full-attention +model is simply the one-group case. + +A manager block covers the same token range in every group, so one KVCM cache +key (hashed from token ids) owns the location specs of all groups of all ranks. +""" + import copy import json import math import time import typing -import inspect import threading from dataclasses import dataclass, field -from typing import Any, Optional, List, Dict, Tuple +from typing import Any, List, Optional, Tuple from concurrent.futures import ThreadPoolExecutor from kv_cache_manager.client.pybind import kvcm_py_client @@ -20,6 +39,7 @@ KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole, + SupportsHMA, ) try: @@ -30,6 +50,7 @@ # vllm <= v0.11.0 from vllm.utils import get_kv_cache_torch_dtype, get_ip +from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec from vllm.v1.core.sched.output import SchedulerOutput from vllm.v1.outputs import KVConnectorOutput @@ -39,8 +60,7 @@ from kv_cache_manager.py_connector.common.logger import logger, configure_log_level from kv_cache_manager.py_connector.common._version_info import FULL_VERSION, GIT_COMMIT, BUILD_TIME -from kv_cache_manager.py_connector.common.types import KVCacheInfo -from kv_cache_manager.py_connector.kernel.gather_scatter_helper import CopyBufferAllocator +from kv_cache_manager.py_connector.common.types import KVCacheInfo, TransferGroup from kv_cache_manager.py_connector.vllm.metadata import SaveRequest, LoadRequest, FinishRequest, ReqStateToWorker, \ TairKvCacheConnectorMetadata from kv_cache_manager.py_connector.vllm.config import TairKvCacheConnectorExtraConfig @@ -52,197 +72,190 @@ from vllm.attention import AttentionMetadata from vllm.v1.request import Request from vllm.v1.core.kv_cache_manager import KVCacheBlocks + from vllm.v1.kv_cache_interface import KVCacheConfig + + +@dataclass +class GroupMeta: + """Static description of one kv_cache_group, derived from KVCacheConfig. + + Available in both scheduler and worker roles (before tensors exist).""" + + group_idx: int + is_attention: bool + layer_names: List[str] + # The group's block table granularity in tokens (spec.block_size). + block_size: int + # Bytes stored per manager block for the whole group. + per_block_bytes: int + # Mamba only: bytes per block per layer (page_size_bytes of the spec). + page_size_bytes: int = 0 @dataclass class ReqState: - """请求状态类,跟踪单个请求的状态信息""" + """Tracks one request. Lives in the scheduler and (mirrored) in workers.""" - # TODO: split this class to ReqStateInScheduler and ReqStateInWorker req_id: str - token_ids: list[int] - local_block_ids: list[int] + token_ids: list + # Per kv_cache_group block table (same length across groups). + block_ids_per_group: List[List[int]] has_saved_block_num: int local_matched_token_num: int remote_matched_token_num: int - # vllm_request only avail in scheduler + # vllm_request only available in scheduler vllm_request: Optional["Request"] - # scheduled_saving_count, sent_saving_count, need_report_after_saving_finished: - # not sync between scheduler and worker and have different meaning - # only available in scheduler and tp0 worker + # Saving progress counters; only meaningful in scheduler and tp0 worker. scheduled_saving_count: int = 0 sent_saving_count: int = 0 need_report_after_saving_finished: bool = False + @property + def num_allocated_blocks(self) -> int: + if not self.block_ids_per_group: + return 0 + return min(len(b) for b in self.block_ids_per_group) + @staticmethod - def create_from_delta(req_state_delta: 'ReqStateToWorker'): - """从ReqStateToWorker创建ReqState实例""" + def create_from_delta(delta: "ReqStateToWorker") -> "ReqState": return ReqState( - req_id=req_state_delta.req_id, - token_ids=req_state_delta.new_tokens_ids, - local_block_ids=req_state_delta.new_local_block_ids, - has_saved_block_num=req_state_delta.has_saved_block_num, + req_id=delta.req_id, + token_ids=list(delta.new_tokens_ids), + block_ids_per_group=[list(b) for b in delta.new_block_ids_per_group], + has_saved_block_num=delta.has_saved_block_num, local_matched_token_num=0, remote_matched_token_num=0, - vllm_request=None + vllm_request=None, ) - def update_from_delta(self, req_state_delta: 'ReqStateToWorker'): - """使用ReqStateToWorker更新当前状态""" - self.token_ids.extend(req_state_delta.new_tokens_ids) - - if req_state_delta.resumed_from_preemption: - self.local_block_ids = req_state_delta.new_local_block_ids + def update_from_delta(self, delta: "ReqStateToWorker"): + self.token_ids.extend(delta.new_tokens_ids) + if not delta.new_block_ids_per_group: + return + if delta.resumed_from_preemption: + self.block_ids_per_group = [list(b) for b in delta.new_block_ids_per_group] else: - self.local_block_ids.extend(req_state_delta.new_local_block_ids) + if not self.block_ids_per_group: + self.block_ids_per_group = [[] for _ in delta.new_block_ids_per_group] + for group_ids, new_ids in zip(self.block_ids_per_group, delta.new_block_ids_per_group): + group_ids.extend(new_ids) -@dataclass -class TransferTaskArgs: - blocks_idx: List[List[int]] = field(default_factory=list) - remote_uris: List[str] = field(default_factory=list) - - -class TairKvCacheConnector(KVConnectorBase_V1): - def _tp_rank_to_spec_name(self, tp_rank: int) -> str: - """Convert TP rank to location spec name.""" - return f"tp{tp_rank}" - - def __init__(self, - vllm_config: "VllmConfig", - role: KVConnectorRole, - kv_cache_config: Optional["KVCacheConfig"] = None, - ): - - init_params = inspect.signature(KVConnectorBase_V1.__init__).parameters - if len(init_params) == 3: - # vllm <= 0.11.0 - super().__init__(vllm_config, role) - else: - # vllm >= 0.11.1 - super().__init__(vllm_config, role, kv_cache_config) +class TairKvCacheConnector(KVConnectorBase_V1, SupportsHMA): - logger.warning("KVCM vllm connector version: %s (commit: %s, build: %s)", FULL_VERSION, GIT_COMMIT, BUILD_TIME) + # ------------------------------------------------------------------ # + # Init / registration + # ------------------------------------------------------------------ # + def __init__(self, vllm_config: "VllmConfig", role: KVConnectorRole, + kv_cache_config: Optional["KVCacheConfig"] = None): + super().__init__(vllm_config, role, kv_cache_config) + assert kv_cache_config is not None, \ + "TairKvCacheConnector requires vLLM to pass kv_cache_config (vllm >= 0.11.1)" - connector_extra_config = vllm_config.kv_transfer_config.kv_connector_extra_config - self._extra_config = TairKvCacheConnectorExtraConfig(connector_extra_config) + logger.warning("KVCM vllm connector version: %s (commit: %s, build: %s)", + FULL_VERSION, GIT_COMMIT, BUILD_TIME) - # Apply log level with priority: env var > startup param > default + self._extra_config = TairKvCacheConnectorExtraConfig( + vllm_config.kv_transfer_config.kv_connector_extra_config) configure_log_level(self._extra_config.log_level) - self._kv_caches: Optional[dict[str, torch.Tensor]] = None - self._local_block_size = vllm_config.cache_config.block_size - model_config = vllm_config.model_config + assert vllm_config.parallel_config.pipeline_parallel_size == 1 + if getattr(model_config, "use_mla", False): + raise NotImplementedError("MLA models are not supported by TairKvCacheConnector") - self._use_mla = (hasattr(model_config, "use_mla") and - isinstance(model_config.use_mla, bool) and - model_config.use_mla) - - manager_block_size = self._local_block_size + self._vllm_block_size = vllm_config.cache_config.block_size + self._tp_size = vllm_config.parallel_config.tensor_parallel_size + self._kv_dtype = get_kv_cache_torch_dtype( + vllm_config.cache_config.cache_dtype, model_config.dtype) + + # Manager block size: attention KV is token-granular and can be re-blocked, + # but mamba state exists once per scheduler block, so hybrid models must + # keep manager block == scheduler block. + manager_block_size = self._vllm_block_size + self._has_state_groups = any( + isinstance(g.kv_cache_spec, MambaSpec) for g in kv_cache_config.kv_cache_groups) if self._extra_config.preferred_block_size != 0: - manager_block_size = self._extra_config.preferred_block_size + if self._has_state_groups: + if self._extra_config.preferred_block_size != self._vllm_block_size: + logger.warning( + "preferred_block_size=%d ignored for hybrid model: mamba state is " + "per scheduler block (%d)", self._extra_config.preferred_block_size, + self._vllm_block_size) + else: + manager_block_size = self._extra_config.preferred_block_size + self._manager_block_size = manager_block_size - self._tp_size = vllm_config.parallel_config.tensor_parallel_size - kv_dtype = get_kv_cache_torch_dtype(vllm_config.cache_config.cache_dtype, model_config.dtype) - num_layer = model_config.get_num_layers(vllm_config.parallel_config) - per_tp_rank_kv_head_num = model_config.get_num_kv_heads(vllm_config.parallel_config) - head_size = model_config.get_head_size() - per_manager_location_spec_shape = [num_layer, 1 if self._use_mla else 2, manager_block_size, - per_tp_rank_kv_head_num, - head_size] + self._group_metas = self._parse_groups(kv_cache_config) + self._num_groups = len(self._group_metas) - assert vllm_config.parallel_config.pipeline_parallel_size == 1 deployment = { "model_name": model_config.served_model_name, - "dtype": str(kv_dtype)[6:], # remove "torch." - "use_mla": self._use_mla, - "tp_size": vllm_config.parallel_config.tensor_parallel_size, + "dtype": str(self._kv_dtype)[6:], # strip "torch." + "use_mla": False, + "tp_size": self._tp_size, "dp_size": vllm_config.parallel_config.data_parallel_size, "pp_size": vllm_config.parallel_config.pipeline_parallel_size, } - logger.info(deployment) + logger.info("deployment: %s, groups: %s", deployment, self._group_metas) - self._manager_client = KvCacheManagerClient.from_connector_config( - vars(self._extra_config) - ) - self._manager_block_size = manager_block_size + self._manager_client = KvCacheManagerClient.from_connector_config(vars(self._extra_config)) self._alive_requests: dict[str, ReqState] = {} self._waiting_to_load_requests: List[LoadRequest] = [] self._waiting_to_save_requests_lock = threading.Lock() self._waiting_to_save_requests: List[SaveRequest] = [] self._waiting_to_finish_requests: List[FinishRequest] = [] - self._canceled_save_request_ids_lock = threading.Lock() self._canceled_save_request_ids: List[str] = [] - # TODO: add coordinator host auto detection, maybe use data parallel host - # TODO: add DP support self._host_ip = get_ip() port = self._extra_config.coordinator_base_port register_response = self._manager_client.register_instance({ - "trace_id": "trace_trace", + "trace_id": "register_%s" % self._extra_config.instance_id, "instance_group": self._extra_config.instance_group, "instance_id": self._extra_config.instance_id, "model_deployment": deployment, "block_size": manager_block_size, - "location_spec_infos": [{ - "name": self._tp_rank_to_spec_name(rank), - "size": math.prod(per_manager_location_spec_shape) * kv_dtype.itemsize - } for rank in range(self._tp_size)], + "location_spec_infos": [ + {"name": self._spec_name(rank, meta.group_idx), "size": meta.per_block_bytes} + for rank in range(self._tp_size) for meta in self._group_metas + ], }) - # TODO: check conflict and update - self._iov_size = math.prod( - per_manager_location_spec_shape) * kv_dtype.itemsize * self._extra_config.hf3fs_concurrent_io_block_count + + max_group_bytes = max(m.per_block_bytes for m in self._group_metas) + self._iov_size = max_group_bytes * self._extra_config.hf3fs_concurrent_io_block_count if role == KVConnectorRole.SCHEDULER: self._epoch = 0 self._coordinator_client = TpCoordinatorClient(self._host_ip, port) self._http_executor = ThreadPoolExecutor(max_workers=4, thread_name_prefix="kvcm_http_") - self._location_query_manager = LocationQueryManager(self._manager_client, self._http_executor, - self._extra_config.instance_id, - self._extra_config.async_get_cache_location) - + self._location_query_manager = LocationQueryManager( + self._manager_client, self._http_executor, self._extra_config.instance_id, + self._extra_config.async_get_cache_location) logger.warning( - "TairKvCacheConnector in scheduler inited, kv_connector_extra_config: %r," - " server block size: %d, vllm block size: %d,", - self._extra_config.__dict__, - self._manager_block_size, - self._local_block_size, - ) + "TairKvCacheConnector scheduler inited, extra_config: %r, manager block size: %d, " + "vllm block size: %d, groups: %d", + self._extra_config.__dict__, self._manager_block_size, + self._vllm_block_size, self._num_groups) elif role == KVConnectorRole.WORKER: self._tp_rank = get_tensor_model_parallel_rank() self._device_mod = None - self._save_stream = None - self._load_stream = None - - logger.warning( - "TairKvCacheConnector in worker inited, tp rank: %d, tp size: %d, host_ip: %s, port: %d" % ( - self._tp_rank, self._tp_size, self._host_ip, port) - ) - if self._tp_rank == 0: - # start coordinator - self._coordinator_server = TpCoordinatorServer(self._host_ip, port, self._tp_size, - self.on_save_finished) - + self._coordinator_server = TpCoordinatorServer( + self._host_ip, port, self._tp_size, self.on_save_finished) self._coordinator_client = TpCoordinatorClient(self._host_ip, port) self._storage_configs = register_response["storage_configs"] - # data transfer setup - self._location_spec_name = self._tp_rank_to_spec_name(self._tp_rank) - self._write_timeout_seconds = self._extra_config.write_timeout_seconds - - sdk_backend_configs = [] - - hf3fs_configs = self.parse_hf3fs_configs(self._storage_configs) - sdk_backend_configs.extend(hf3fs_configs) - logger.debug(sdk_backend_configs) + sdk_backend_configs = self.parse_hf3fs_configs(self._storage_configs) + self._self_spec_names = { + meta.group_idx: self._spec_name(self._tp_rank, meta.group_idx) + for meta in self._group_metas + } transfer_client_json = { "instance_group": self._extra_config.instance_group, "instance_id": self._extra_config.instance_id, @@ -257,26 +270,63 @@ def __init__(self, }, }, "location_spec_infos": { - self._location_spec_name: math.prod(per_manager_location_spec_shape) * kv_dtype.itemsize, + self._self_spec_names[meta.group_idx]: meta.per_block_bytes + for meta in self._group_metas }, } - self._transfer_client_config = json.dumps(transfer_client_json) - - self._init_params = kvcm_py_client.InitParams() - self._init_params.role_type = kvcm_py_client.RoleType.WORKER - self._init_params.self_location_spec_name = self._location_spec_name - self._init_params.storage_configs = f"{self._storage_configs}" - - logger.info("_transfer_client_config:%s, _init_params:%s", self._transfer_client_config, self._init_params) - + init_params = kvcm_py_client.InitParams() + init_params.role_type = kvcm_py_client.RoleType.WORKER + init_params.self_location_spec_name = self._self_spec_names[self._group_metas[0].group_idx] + init_params.storage_configs = f"{self._storage_configs}" + transfer_client_config = json.dumps(transfer_client_json) + logger.info("transfer_client_config: %s", transfer_client_config) self._transfer_client = kvcm_py_client.TransferClient.Create( - self._transfer_client_config, self._init_params - ) + transfer_client_config, init_params) assert self._transfer_client is not None, "kvcm_py_client.TransferClient.Create failed" + logger.warning( + "TairKvCacheConnector worker inited, tp rank: %d/%d, host: %s:%d, groups: %d", + self._tp_rank, self._tp_size, self._host_ip, port, self._num_groups) + + def _spec_name(self, tp_rank: int, group_idx: int) -> str: + return f"tp{tp_rank}_g{group_idx}" + + def _parse_groups(self, kv_cache_config: "KVCacheConfig") -> List[GroupMeta]: + metas = [] + for idx, group in enumerate(kv_cache_config.kv_cache_groups): + if getattr(group, "is_eagle_group", False): + logger.warning("skip eagle group %d (%d layers)", idx, len(group.layer_names)) + continue + spec = group.kv_cache_spec + if isinstance(spec, MambaSpec): + metas.append(GroupMeta( + group_idx=idx, + is_attention=False, + layer_names=list(group.layer_names), + block_size=spec.block_size, + per_block_bytes=spec.page_size_bytes * len(group.layer_names), + page_size_bytes=spec.page_size_bytes, + )) + elif isinstance(spec, FullAttentionSpec): + # Attention KV is token-granular; scale from the spec's page size + # to the manager block size. + per_token_bytes = spec.page_size_bytes // spec.block_size + metas.append(GroupMeta( + group_idx=idx, + is_attention=True, + layer_names=list(group.layer_names), + block_size=spec.block_size, + per_block_bytes=per_token_bytes * self._manager_block_size * len(group.layer_names), + )) + else: + raise NotImplementedError( + f"Unsupported kv cache spec {type(spec).__name__} in group {idx}") + assert metas, "no usable kv cache groups" + return metas def shutdown(self): - # TODO: stop background threads and cleanup transfer client self._manager_client.close() + if hasattr(self, "_location_query_manager"): + self._location_query_manager.shutdown() return None def parse_hf3fs_configs(self, storage_configs): @@ -286,7 +336,7 @@ def parse_hf3fs_configs(self, storage_configs): if storage_config["type"] == "vcns_hf3fs": storage_config["type"] = "hf3fs" if storage_config["type"] == "hf3fs" and storage_config["is_available"]: - hf3fs_config = { + hf3fs_configs.append({ "type": storage_config["type"], "mountpoint": storage_config["storage_spec"]["mountpoint"], "root_dir": storage_config["storage_spec"]["root_dir"], @@ -294,371 +344,422 @@ def parse_hf3fs_configs(self, storage_configs): "read_iov_size": self._iov_size, "write_iov_block_size": self._extra_config.write_iov_block_size, "write_iov_size": self._iov_size, - } - hf3fs_configs.append(hf3fs_config) + }) self._storage_configs = json.dumps(storage_configs_json) return hf3fs_configs - def generate_blocks(self, token_ids, block_size, max_token_length) -> list[dict[str, Any]]: - results = [] - token_length = min(len(token_ids), max_token_length) - # token_length = len(token_ids) - for i in range(0, token_length, block_size): - if i + block_size > token_length: - break - results.append({ - "token_ids": token_ids[i:i + block_size], - "unique_id": None, - "location": None - }) - return results - - # ============================== - # Worker-side methods - # ============================== - - def generate_blocks_idx(self, manager_block_idxes, local_block_ids): - blocks_idx = [] - for manager_block_idx in manager_block_idxes: - # get kvcache index list - block_idx = [] - for i in range(self._manager_block_size): - now_token_idx = manager_block_idx * self._manager_block_size + i - assert now_token_idx // self._local_block_size < len(local_block_ids) - local_block_id = local_block_ids[now_token_idx // self._local_block_size] - token_offset = now_token_idx % self._local_block_size - block_idx.append(local_block_id * self._local_block_size + token_offset) - blocks_idx.append(block_idx) - return blocks_idx - - def on_save_finished(self, write_session_id: str, save_context: SaveContext): - logger.debug(save_context.result_per_rank) - for block_idx in range(len(save_context.locations)): - # TODO: report uri when enable local alloc - # location_specs = [] - is_fully_saved = True - for rank in range(self._tp_size): - is_success = save_context.result_per_rank[rank][block_idx] - if not is_success: - # this spec is not fully saved, report failed - is_fully_saved = False - # else: - # # Convert the spec to include name field instead of tp_rank - # location_specs.append({ - # "name": self._tp_rank_to_spec_name(rank), - # "uri": spec - # }) - if is_fully_saved: - # save_context.locations[block_idx]["location_specs"] = location_specs - save_context.success_mask.append(True) - else: - save_context.success_mask.append(False) - logger.debug("finish_write_cache blocks:%s mask:%s write_session_id:%s", save_context.locations, - save_context.success_mask, write_session_id) - try: - self._manager_client.finish_write_cache({ - "trace_id": "test_test", - "instance_id": self._extra_config.instance_id, - "write_session_id": write_session_id, - "success_blocks": { - "bool_masks": { - "values": save_context.success_mask - } - } - }) - except Exception as e: - logger.warning("finish_write_cache failed, write_session_id: %s, error: %s", write_session_id, e) - + # ------------------------------------------------------------------ # + # Worker side: KV cache registration + # ------------------------------------------------------------------ # def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): - _, first_layer_kvcache = next(iter(kv_caches.items())) self._kv_caches = kv_caches - # TODO: support MLA - - assert self._local_block_size == first_layer_kvcache.shape[2], "kv cache shape error" - for layer_name, kvcache in kv_caches.items(): - assert kvcache.is_contiguous(), "kv cache must be contiguous" - - # torch.Size([2, block_num, block_size, kv_head_num, kv_dim]) - # 2 -> key, value - self._local_block_num = first_layer_kvcache.shape[1] - self._local_token_num = self._local_block_num * self._local_block_size - - self._dtype = first_layer_kvcache.dtype - self._device = first_layer_kvcache.device + first_attn = next(kv_caches[name] + for meta in self._group_metas if meta.is_attention + for name in meta.layer_names) + self._dtype = first_attn.dtype + self._device = first_attn.device self._device_mod = _get_device_module(self._device) - self._save_stream = self._device_mod.Stream() - self._load_stream = self._device_mod.Stream() - self._per_manager_location_spec_layer_shape = [first_layer_kvcache.shape[0], - self._manager_block_size, - first_layer_kvcache.shape[3] * first_layer_kvcache.shape[4]] - self._per_manager_location_spec_layer_byte_size = math.prod( - self._per_manager_location_spec_layer_shape) * self._dtype.itemsize - self._per_layer_token_key_dim_size = first_layer_kvcache.shape[3] * first_layer_kvcache.shape[4] - self._per_layer_token_key_byte_size = (first_layer_kvcache.shape[3] * - first_layer_kvcache.shape[4] * self._dtype.itemsize) - assert self._per_layer_token_key_byte_size == first_layer_kvcache[0][0][1].data_ptr() - \ - first_layer_kvcache[0][0][0].data_ptr(), "kv cache shape error" - assert self._per_manager_location_spec_layer_byte_size == 2 * self._manager_block_size * self._per_layer_token_key_byte_size - - self._per_manager_location_spec_shape = [len(self._kv_caches)] + self._per_manager_location_spec_layer_shape - self._per_manager_location_spec_byte_size = math.prod( - self._per_manager_location_spec_shape) * self._dtype.itemsize - - self._kvcache_ptr_tensor_cpu = torch.tensor( - [self._kv_caches[name].data_ptr() for name in self._kv_caches], - dtype=torch.int64, - device="cpu" - ) - self._kvcache_ptr_tensor_gpu = self._kvcache_ptr_tensor_cpu.to(self._device) - if self._use_mla: - self._all_kvcache_ptr_tensor_cpu = torch.tensor( - [self._kv_caches[name].data_ptr() for name in self._kv_caches], - dtype=torch.int64, - device="cpu" - ) - else: - kvcache_ptrs = [] - for name in self._kv_caches: - kvcache_ptrs.append(self._kv_caches[name][0].data_ptr()) - kvcache_ptrs.append(self._kv_caches[name][1].data_ptr()) - self._all_kvcache_ptr_tensor_cpu = torch.tensor( - kvcache_ptrs, - dtype=torch.int64, - device="cpu" - ) - self._all_kvcache_ptr_tensor_gpu = self._all_kvcache_ptr_tensor_cpu.to(self._device) + + groups = [self._build_transfer_group(meta, kv_caches) for meta in self._group_metas] self._kvcache_info = KVCacheInfo( - self._tp_rank, - self._tp_size, - self._kv_caches, - self._kvcache_ptr_tensor_cpu, - self._kvcache_ptr_tensor_gpu, - self._all_kvcache_ptr_tensor_gpu, - len(self._kv_caches), - self._local_token_num, - tuple(self._per_manager_location_spec_shape), - self._per_manager_location_spec_byte_size, - self._per_layer_token_key_dim_size, - self._device, - self._dtype + tp_rank=self._tp_rank, + world_size=self._tp_size, + groups=groups, + device=self._device, + dtype=self._dtype, ) - self._copy_buffer_allocator = CopyBufferAllocator(torch.device("cpu"), self._dtype, - self._per_manager_location_spec_shape, 1024) - - # 初始化DataTransferManager实例 self._data_transfer = DataTransferManager( - self._kvcache_info, - self._manager_block_size, - self._copy_buffer_allocator, - self._transfer_client, - self._coordinator_client, - self._extra_config, - ) + self._kvcache_info, self._manager_block_size, + self._transfer_client, self._coordinator_client, self._extra_config) + + logger.warning("register_kv_caches done: %s", [ + (g.spec_name, "attn" if g.is_attention else "state", + g.layer_num, g.per_block_bytes) for g in groups]) + + def _build_transfer_group(self, meta: GroupMeta, kv_caches) -> TransferGroup: + spec_name = self._self_spec_names[meta.group_idx] + if meta.is_attention: + tensors = [kv_caches[name] for name in meta.layer_names] + ref = tensors[0] + # vLLM >= 0.26.0 packs K and V into the content dim: logical shape + # (num_blocks, num_kv_heads, kernel_block_size, 2*head_size). With the + # default NHD stride order the memory is laid out token-major as + # (num_blocks, kernel_block_size, num_kv_heads, 2*head_size), so a flat + # index (global_token * per_token_dim + dim) walks the storage + # correctly. K/V packing is opaque to the byte-exact transport. + assert ref.dim() == 4, f"unexpected kv layout {ref.shape}" + for t in tensors: + assert t.shape == ref.shape and t.stride() == ref.stride(), \ + "attention layers in one group must share shape/stride" + kernel_block_size = ref.shape[2] + assert meta.block_size % kernel_block_size == 0, \ + f"group block size {meta.block_size} not a multiple of kernel " \ + f"block size {kernel_block_size}" + per_token_dim = ref.shape[1] * ref.shape[3] # num_kv_heads * 2*head_size + # The gather/scatter kernel needs token-major memory inside a page: + # dims (blk, head, tok, dim) laid out as (blk, tok, head, dim). This + # is vLLM's NHD order; HND would interleave heads across tokens. + assert ref.stride()[1:] == (ref.shape[3], per_token_dim, 1), \ + f"kv cache page not token-major: shape={ref.shape} " \ + f"stride={ref.stride()}; set VLLM_KV_CACHE_LAYOUT=NHD" + # Padded pages (page_size_padded) leave gaps between blocks; the + # kernel's strided path skips them. Stride 0 = fast flat indexing. + flat = ref.stride(0) == kernel_block_size * per_token_dim + block_stride = 0 if flat else ref.stride(0) + # One pointer per layer: K and V are packed in the content dim, and + # data_ptr() of the permuted view is the storage base. + ptrs = [t.data_ptr() for t in tensors] + ptr_tensor = torch.tensor(ptrs, dtype=torch.int64, device="cpu").to(self._device) + return TransferGroup( + group_idx=meta.group_idx, + spec_name=spec_name, + is_attention=True, + layer_names=meta.layer_names, + block_size=meta.block_size, + per_block_bytes=meta.per_block_bytes, + kvcache_ptr_tensor_gpu=ptr_tensor, + layer_num=len(meta.layer_names), + per_token_dim=per_token_dim, + kernel_block_size=kernel_block_size, + kv_stride=0, + block_stride=block_stride, + ) - logger.warning("register_kv_caches, _per_manager_location_spec_layer_shape: %s", - self._per_manager_location_spec_layer_shape) + # Mamba/state group: each layer is a list[Tensor] sharing one storage; + # rebuild a (num_blocks, page_size_bytes) byte view for opaque copy. + block_views = [] + for name in meta.layer_names: + states = kv_caches[name] + assert isinstance(states, (list, tuple)) and len(states) > 0, \ + f"state layer {name} should be a list of tensors" + storage = states[0].untyped_storage() + for st in states[1:]: + assert st.untyped_storage().data_ptr() == storage.data_ptr(), \ + f"state layer {name}: tensors do not share storage" + num_blocks = states[0].shape[0] + need = num_blocks * meta.page_size_bytes + assert storage.nbytes() >= need, \ + f"state layer {name}: storage {storage.nbytes()} < {need}" + byte_view = torch.tensor([], dtype=torch.uint8, device=self._device).set_(storage) + block_views.append(byte_view[:need].view(num_blocks, meta.page_size_bytes)) + return TransferGroup( + group_idx=meta.group_idx, + spec_name=spec_name, + is_attention=False, + layer_names=meta.layer_names, + block_size=meta.block_size, + per_block_bytes=meta.per_block_bytes, + layer_num=len(meta.layer_names), + block_view_tensors=block_views, + page_size_bytes=meta.page_size_bytes, + ) + # ------------------------------------------------------------------ # + # Block index translation + # ------------------------------------------------------------------ # + def _attn_token_indices(self, group: TransferGroup, manager_block_idxes, + block_table) -> List[List[int]]: + """Map manager blocks to flat token slots of one attention group. + + Three-tier hierarchy: + manager block (KVCM unit) -> global token idx + -> group block (block_table unit, group.block_size tokens) + -> kernel physical block (tensor unit; ratio physical per group block). + """ + mbs = self._manager_block_size + gbs = group.block_size + kbs = group.kernel_block_size + ratio = gbs // kbs + out = [] + for mb in manager_block_idxes: + idxs = [] + base = mb * mbs + for i in range(mbs): + tok = base + i + logical = tok // gbs + assert logical < len(block_table), ( + f"group block {logical} out of range (len={len(block_table)})") + off = tok % gbs + phys = block_table[logical] * ratio + off // kbs + idxs.append(phys * kbs + off % kbs) + out.append(idxs) + return out + + def _state_block_ids(self, group: TransferGroup, manager_block_idxes, + block_table) -> List[int]: + """Map manager blocks to block ids of a state (mamba) group. + + State is stored once per group block and covers the whole prefix up to + that block, so the manager block's last token selects the block.""" + mbs = self._manager_block_size + gbs = group.block_size + out = [] + for mb in manager_block_idxes: + logical = ((mb + 1) * mbs - 1) // gbs + assert logical < len(block_table), ( + f"group block {logical} out of range (len={len(block_table)})") + out.append(block_table[logical]) + return out + + def _self_uris(self, locations, spec_name: str) -> List[str]: + uris = [] + for location in locations: + for spec in location.get("location_specs", []): + if spec["name"] == spec_name: + uris.append(spec["uri"]) + return uris + + # ------------------------------------------------------------------ # + # Worker side: load / save + # ------------------------------------------------------------------ # + def _submit_group_tasks(self, task_fn, multi_result, task_idx, group, + uris, token_indices, block_ids, per_task_size, *extra): + for i in range(0, len(uris), per_task_size): + end = min(len(uris), i + per_task_size) + self._data_transfer.submit_task( + task_fn, multi_result, task_idx, group, uris[i:end], + token_indices[i:end] if token_indices is not None else None, + block_ids[i:end] if block_ids is not None else None, *extra) + task_idx += 1 + return task_idx + + def _plan_group_transfers(self, locations, manager_block_idxes, block_ids_per_group): + """Build (group, uris, token_indices, block_ids) for every group. + + Returns None if any group's URI list does not cover all blocks.""" + num_blocks = len(manager_block_idxes) + plans = [] + for group in self._kvcache_info.groups: + uris = self._self_uris(locations, group.spec_name) + if len(uris) != num_blocks: + logger.warning("group %s: %d uris for %d blocks, skip transfer", + group.spec_name, len(uris), num_blocks) + return None + # block_ids_per_group is indexed by the vLLM group index. + block_table = block_ids_per_group[group.group_idx] + if group.is_attention: + plans.append((group, uris, + self._attn_token_indices(group, manager_block_idxes, block_table), + None)) + else: + plans.append((group, uris, None, + self._state_block_ids(group, manager_block_idxes, block_table))) + return plans def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None: meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) - for load_req in meta.to_load_requests: - if len(load_req.need_load_locations) == 0: + if not load_req.need_load_locations: + continue + num_blocks = len(load_req.manager_block_idxes) + plans = self._plan_group_transfers( + load_req.need_load_locations, load_req.manager_block_idxes, + load_req.all_block_ids) + + # Report failures against the block table vLLM can act on: map each + # manager block to the logical block holding its first token. vLLM + # truncates computed tokens at the first invalid block, so this is + # sufficient for recovery. vLLM's invalid-block handling only + # supports single-group models; for hybrid models a failed load can + # only be logged. + report_ids = [] + if self._num_groups == 1: + table = load_req.all_block_ids[0] + gbs = self._group_metas[0].block_size + report_ids = [table[(mb * self._manager_block_size) // gbs] + for mb in load_req.manager_block_idxes] + done_cb = self._data_transfer.create_load_done_callback( + load_req.req_id, self._tp_rank, meta.epoch, + copy.copy(report_ids), num_blocks, + report_failures=self._num_groups == 1) + + if plans is None: + # Nothing submitted; report the whole load as failed. + mr = MultiResult(1, done_cb) + mr.submit_result(0, [False] * num_blocks * self._num_groups) continue - block_token_indices = self.generate_blocks_idx(load_req.manager_block_idxes, load_req.local_block_ids) - all_remote_uris = self.get_self_uris(load_req.need_load_locations) - - per_task_size = self._extra_config.block_per_load_task - task_num = math.ceil(len(block_token_indices) / per_task_size) - done_callback = self._data_transfer.create_load_done_callback( - load_req.req_id, - self._kvcache_info.tp_rank, - meta.epoch, - copy.copy(load_req.local_block_ids) - ) - multi_result = MultiResult(task_num, done_callback) - + per_task = self._extra_config.block_per_load_task + task_num = sum(math.ceil(num_blocks / per_task) for _ in plans) + multi_result = MultiResult(task_num, done_cb) task_idx = 0 - for i in range(0, len(block_token_indices), per_task_size): - end_idx = min(len(block_token_indices), i + per_task_size) - task_remote_uris = all_remote_uris[i:end_idx] - task_block_token_indices = block_token_indices[i:end_idx] - self._data_transfer.submit_task(self._data_transfer.load_task, multi_result, task_idx, task_remote_uris, - task_block_token_indices) - task_idx += 1 + for group, uris, token_indices, block_ids in plans: + task_idx = self._submit_group_tasks( + self._data_transfer.load_task, multi_result, task_idx, + group, uris, token_indices, block_ids, per_task) def wait_for_layer_load(self, layer_name: str) -> None: - # logger.warning("wait_for_layer_load, layer_name: %s", layer_name) pass - def save_kv_layer(self, layer_name: str, kv_layer: torch.Tensor, attn_metadata: "AttentionMetadata", - **kwargs) -> None: - # logger.warning("save_kv_layer, layer_name: %s", layer_name) + def save_kv_layer(self, layer_name: str, kv_layer: torch.Tensor, + attn_metadata: "AttentionMetadata", **kwargs) -> None: pass def wait_for_save(self): meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) - # logger.warning("wait_for_save, meta: %r", meta) - - kvcache_ready_event = None - if len(meta.to_save_requests) > 0: - kvcache_ready_event = self._device_mod.Event() - kvcache_ready_event.record(self._device_mod.current_stream()) - - for req_save in meta.to_save_requests: - req = self._alive_requests[req_save.req_id] - - # get idx - blocks_idx = self.generate_blocks_idx(req_save.manager_block_idxes, req.local_block_ids) - all_remote_uris = self.get_self_uris(req_save.target_locations) - - per_task_size = self._extra_config.block_per_save_task - task_num = math.ceil(len(blocks_idx) / per_task_size) - done_callback = self._data_transfer.create_save_done_callback( - req.req_id, - self._kvcache_info.tp_rank, - req_save.write_session_id - ) - multi_result = MultiResult(task_num, done_callback) + if not meta.to_save_requests: + return + ready_event = self._device_mod.Event() + ready_event.record(self._device_mod.current_stream()) + + for save_req in meta.to_save_requests: + req = self._alive_requests[save_req.req_id] + num_blocks = len(save_req.manager_block_idxes) + plans = self._plan_group_transfers( + save_req.target_locations, save_req.manager_block_idxes, + req.block_ids_per_group) + + done_cb = self._data_transfer.create_save_done_callback( + req.req_id, self._tp_rank, save_req.write_session_id, num_blocks) + + if plans is None: + mr = MultiResult(1, done_cb) + mr.submit_result(0, [False] * num_blocks * self._num_groups) + continue + per_task = self._extra_config.block_per_save_task + task_num = sum(math.ceil(num_blocks / per_task) for _ in plans) + multi_result = MultiResult(task_num, done_cb) task_idx = 0 - for i in range(0, len(blocks_idx), per_task_size): - end_idx = min(len(blocks_idx), i + per_task_size) - task_remote_uris = all_remote_uris[i:end_idx] - task_block_token_indices = blocks_idx[i:end_idx] - self._data_transfer.submit_task(self._data_transfer.save_task, multi_result, task_idx, task_remote_uris, - task_block_token_indices, - kvcache_ready_event) - task_idx += 1 + for group, uris, token_indices, block_ids in plans: + task_idx = self._submit_group_tasks( + self._data_transfer.save_task, multi_result, task_idx, + group, uris, token_indices, block_ids, per_task, ready_event) if self._tp_rank == 0: req.scheduled_saving_count += 1 - def get_self_uris(self, locations): - all_remote_uris = [] - for idx, location in enumerate(locations): - for location_spec in location["location_specs"]: - # Match by location spec name instead of tp_rank - if self._tp_rank_to_spec_name(self._kvcache_info.tp_rank) == location_spec["name"]: - all_remote_uris.append(location_spec["uri"]) - return all_remote_uris - - def get_finished( - self, finished_req_ids: set[str] - ) -> Tuple[Optional[set[str]], Optional[set[str]]]: + def on_save_finished(self, write_session_id: str, save_context: SaveContext): + for block_idx in range(len(save_context.locations)): + fully_saved = all(save_context.result_per_rank[rank][block_idx] + for rank in range(self._tp_size)) + save_context.success_mask.append(fully_saved) + logger.debug("finish_write_cache mask:%s session:%s", + save_context.success_mask, write_session_id) + try: + self._manager_client.finish_write_cache({ + "trace_id": "finish_%s" % write_session_id[:8], + "instance_id": self._extra_config.instance_id, + "write_session_id": write_session_id, + "success_blocks": {"bool_masks": {"values": save_context.success_mask}}, + }) + except Exception as e: + logger.warning("finish_write_cache failed, session: %s, error: %s", + write_session_id, e) + + def get_finished(self, finished_req_ids: set) -> Tuple[Optional[set], Optional[set]]: meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) if self._tp_rank != 0: for finish_req in meta.to_finish_requests: - req_id = finish_req.req_id - if req_id in self._alive_requests: - self._alive_requests.pop(req_id) + self._alive_requests.pop(finish_req.req_id, None) return None, None - # self._tp_rank == 0 - finished_saving_reqs = [] - # check if any request is saving kvcache - (finished_saving_tasks, finished_loading_tasks) = self._coordinator_server.get_finished_tasks() + finished_saving = [] + finished_saving_tasks, finished_loading_tasks = self._coordinator_server.get_finished_tasks() for req_id in finished_saving_tasks: req = self._alive_requests[req_id] req.sent_saving_count += 1 - assert req.sent_saving_count <= req.scheduled_saving_count if (req.need_report_after_saving_finished and req.sent_saving_count == req.scheduled_saving_count): - finished_saving_reqs.append(req_id) + finished_saving.append(req_id) self._alive_requests.pop(req_id) for finish_req in meta.to_finish_requests: - req_id = finish_req.req_id - if req_id not in self._alive_requests: - # called get_num_new_matched_tokens but never scheduled + req = self._alive_requests.get(finish_req.req_id) + if req is None: continue - req = self._alive_requests[req_id] if req.sent_saving_count == req.scheduled_saving_count: - finished_saving_reqs.append(req_id) - self._alive_requests.pop(req_id) + finished_saving.append(req.req_id) + self._alive_requests.pop(req.req_id) else: - self._alive_requests[req_id].need_report_after_saving_finished = True - return set(finished_saving_reqs), set(finished_loading_tasks) + req.need_report_after_saving_finished = True + return set(finished_saving), set(finished_loading_tasks) - def get_block_ids_with_load_errors(self) -> set[int]: + def get_block_ids_with_load_errors(self) -> set: if self._tp_rank != 0: return set() - failed_set = self._coordinator_server.get_failed_loading_block_idxs() - if len(failed_set) > 0: - logger.warning("block_ids_with_load_errors: %s", failed_set) - return failed_set + failed = self._coordinator_server.get_failed_loading_block_idxs() + if failed: + logger.warning("block_ids_with_load_errors: %s", failed) + return failed - def bind_connector_metadata( - self, connector_metadata: KVConnectorMetadata) -> None: + def bind_connector_metadata(self, connector_metadata: KVConnectorMetadata) -> None: self._connector_metadata = connector_metadata meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) - - for req_state_delta in meta.requests: - if req_state_delta.req_id not in self._alive_requests: - assert not req_state_delta.is_delta - if not req_state_delta.is_delta: - self._alive_requests[req_state_delta.req_id] = ReqState.create_from_delta(req_state_delta) + for delta in meta.requests: + if not delta.is_delta: + self._alive_requests[delta.req_id] = ReqState.create_from_delta(delta) else: - self._alive_requests[req_state_delta.req_id].update_from_delta(req_state_delta) - - # ============================== - # Scheduler-side methods - # ============================== - def get_num_new_matched_tokens(self, request: "Request", num_computed_tokens: int) -> Tuple[int, bool]: - # logger.warning("get matched token ids: %s, id: %s", request.prompt_token_ids, request.request_id) - - bypass_match = False - # TODO: add arrival_time to req_id in order to handle same request id - if request.request_id in self._alive_requests: - # bypass remote match for alive requests - # possible cases: - # 1. reschedule when all kvcache loading failed - # 2. TODO: no enough hbm to schedule the request - # bypass_match = True - # logger.warning("bypass match for alive request, req_id: %s", request.request_id) - pass - - computed_manager_block_size = num_computed_tokens // self._manager_block_size - all_calced_remote_block_num = computed_manager_block_size - new_matched_count = 0 - - if not bypass_match: - is_query_done, need_load_locations = ( - self._location_query_manager.get_locations_for_query(request, computed_manager_block_size)) - if not is_query_done: - # async get_cache_location - return None, False - new_matched_count = len(need_load_locations) * self._manager_block_size - logger.info("req:%s, new_matched_count:%d", request.request_id, new_matched_count) - - all_calced_remote_block_num = computed_manager_block_size + len(need_load_locations) - - if new_matched_count != 0: - self._waiting_to_load_requests.append(LoadRequest( - req_id=request.request_id, - manager_block_idxes=[i for i in range(computed_manager_block_size, all_calced_remote_block_num)], - need_load_locations=need_load_locations, - )) - - new_req_meta = ReqState(request.request_id, copy.copy(request.prompt_token_ids), [], - all_calced_remote_block_num, - num_computed_tokens, - new_matched_count, - request) + assert delta.req_id in self._alive_requests + self._alive_requests[delta.req_id].update_from_delta(delta) + + # ------------------------------------------------------------------ # + # Scheduler side + # ------------------------------------------------------------------ # + def get_num_new_matched_tokens(self, request: "Request", + num_computed_tokens: int) -> Tuple[Optional[int], bool]: + prev = self._alive_requests.get(request.request_id) + if (prev is not None and prev.remote_matched_token_num + and prev.block_ids_per_group): + # The request already went through an external load (blocks were + # allocated) and returned to WAITING -- a KV load failure or a + # preemption. The manager may still advertise blocks whose storage + # is gone, so re-matching risks an endless fail-reschedule loop; + # recompute locally instead. (A pending re-query after a failed + # allocation has empty block_ids_per_group and is not affected.) + logger.warning("req:%s re-queried after an external load attempt, " + "skip external match", request.request_id) + prev.local_matched_token_num = num_computed_tokens + prev.remote_matched_token_num = 0 + prev.has_saved_block_num = num_computed_tokens // self._manager_block_size + return 0, False + + computed_blocks = num_computed_tokens // self._manager_block_size + + is_query_done, need_load_locations = ( + self._location_query_manager.get_locations_for_query(request, computed_blocks)) + if not is_query_done: + # async query in flight; vLLM will ask again + return None, False + + new_matched_count = len(need_load_locations) * self._manager_block_size + # This connector loads synchronously (load_kv_async=False), so vLLM will + # schedule num_tokens - num_computed_tokens new tokens and asserts that + # count is > 0 (vllm/v1/core/sched/scheduler.py). If the whole prompt is + # externally cached, drop trailing blocks so at least one token is + # recomputed locally. + while new_matched_count and num_computed_tokens + new_matched_count >= request.num_tokens: + need_load_locations = need_load_locations[:-1] + new_matched_count -= self._manager_block_size + total_remote_blocks = computed_blocks + len(need_load_locations) + logger.info("req:%s matched %d external tokens", request.request_id, new_matched_count) + + if new_matched_count: + self._waiting_to_load_requests.append(LoadRequest( + req_id=request.request_id, + manager_block_idxes=list(range(computed_blocks, total_remote_blocks)), + need_load_locations=need_load_locations, + )) - self._alive_requests[request.request_id] = new_req_meta + self._alive_requests[request.request_id] = ReqState( + req_id=request.request_id, + token_ids=copy.copy(request.prompt_token_ids), + block_ids_per_group=[], + has_saved_block_num=total_remote_blocks, + local_matched_token_num=num_computed_tokens, + remote_matched_token_num=new_matched_count, + vllm_request=request, + ) return new_matched_count, new_matched_count > 0 - def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", num_external_tokens: int): - if request.request_id not in self._alive_requests: + def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", + num_external_tokens: int): + req_state = self._alive_requests.get(request.request_id) + if req_state is None: return - req_state = self._alive_requests[request.request_id] - # blocks_ids[0]: only one KV cache groups for now - # refer to vllm/v1/core/kv_cache_manager.py:35 - req_state.local_block_ids = copy.copy(blocks.get_block_ids()[0]) + req_state.block_ids_per_group = [list(b) for b in blocks.get_block_ids()] def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnectorMetadata: meta = TairKvCacheConnectorMetadata(self._epoch) @@ -666,89 +767,78 @@ def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnector for load_req in self._waiting_to_load_requests: request = self._alive_requests[load_req.req_id] - if len(request.local_block_ids) == 0: - # ignore load_req if vllm has not called update_state_after_alloc, - # vllm will call get_num_new_matched_tokens again + if not request.block_ids_per_group: + # update_state_after_alloc was never called; vLLM will re-query. continue - load_req.local_block_ids = request.local_block_ids + load_req.all_block_ids = [list(b) for b in request.block_ids_per_group] meta.add_load_request(load_req) self._waiting_to_load_requests = [] for vllm_req in scheduler_output.scheduled_new_reqs: request = self._alive_requests[vllm_req.req_id] - request.local_block_ids = copy.copy(vllm_req.block_ids[0]) - - state_to_worker = ReqStateToWorker(req_id=request.req_id, - has_saved_block_num=request.has_saved_block_num, - new_tokens_ids=request.token_ids, - new_local_block_ids=request.local_block_ids, - is_delta=False - ) - meta.add_req_state_to_worker(state_to_worker) - logger.info("new request: %s, block_ids_len: %d", vllm_req.req_id, len(vllm_req.block_ids[0])) + request.block_ids_per_group = [list(b) for b in vllm_req.block_ids] + meta.add_req_state_to_worker(ReqStateToWorker( + req_id=request.req_id, + has_saved_block_num=request.has_saved_block_num, + new_tokens_ids=request.token_ids, + new_block_ids_per_group=request.block_ids_per_group, + is_delta=False, + )) cached_reqs = scheduler_output.scheduled_cached_reqs for idx, req_id in enumerate(cached_reqs.req_ids): request = self._alive_requests[req_id] - vllm_req = request.vllm_request num_new_tokens = scheduler_output.num_scheduled_tokens[req_id] num_current_tokens = len(request.token_ids) - - new_token_ids = vllm_req.all_token_ids[ - num_current_tokens: num_current_tokens + num_new_tokens - ] - state_to_worker = ReqStateToWorker(req_id=request.req_id, - has_saved_block_num=request.has_saved_block_num) - + new_token_ids = request.vllm_request.all_token_ids[ + num_current_tokens:num_current_tokens + num_new_tokens] request.token_ids.extend(new_token_ids) - state_to_worker.new_tokens_ids = new_token_ids - resumed_from_preemption = False - if hasattr(cached_reqs, "resumed_req_ids"): - # vllm >= 0.11.1 - resumed_from_preemption = req_id in cached_reqs.resumed_req_ids - else: - # vllm <= 0.11.0 - resumed_from_preemption = cached_reqs.resumed_from_preemption[idx] + delta = ReqStateToWorker( + req_id=request.req_id, + has_saved_block_num=request.has_saved_block_num, + new_tokens_ids=new_token_ids, + ) - if resumed_from_preemption: - request.local_block_ids = copy.copy(cached_reqs.new_block_ids[idx][0]) - state_to_worker.resumed_from_preemption = True - state_to_worker.new_local_block_ids = request.local_block_ids + if hasattr(cached_reqs, "resumed_req_ids"): + resumed = req_id in cached_reqs.resumed_req_ids else: - if cached_reqs.new_block_ids[idx] is None: - # https://github.com/vllm-project/vllm/pull/23262 - continue - new_block_ids = cached_reqs.new_block_ids[idx][0] - request.local_block_ids.extend(new_block_ids) - state_to_worker.new_local_block_ids = new_block_ids - meta.add_req_state_to_worker(state_to_worker) + resumed = cached_reqs.resumed_from_preemption[idx] + + new_block_ids = cached_reqs.new_block_ids[idx] + if resumed: + request.block_ids_per_group = [list(b) for b in new_block_ids] + delta.resumed_from_preemption = True + delta.new_block_ids_per_group = request.block_ids_per_group + elif new_block_ids is not None: + # https://github.com/vllm-project/vllm/pull/23262: may be None + delta.new_block_ids_per_group = [list(b) for b in new_block_ids] + for group_ids, new_ids in zip(request.block_ids_per_group, + delta.new_block_ids_per_group): + group_ids.extend(new_ids) + meta.add_req_state_to_worker(delta) for req in self._alive_requests.values(): - target_save_num = min(len(req.token_ids), - len(req.local_block_ids) * self._local_block_size) // self._manager_block_size + target_save_num = min( + len(req.token_ids), + req.num_allocated_blocks * self._vllm_block_size) // self._manager_block_size if target_save_num > req.has_saved_block_num: req.scheduled_saving_count += 1 self._http_executor.submit( - self.start_save_kvcache_async, - req.req_id, + self.start_save_kvcache_async, req.req_id, req.token_ids[:target_save_num * self._manager_block_size], - target_save_num - ) + target_save_num) req.has_saved_block_num = target_save_num - new_save_reqs: List[SaveRequest] = [] with self._waiting_to_save_requests_lock: new_save_reqs = self._waiting_to_save_requests self._waiting_to_save_requests = [] for save_req in new_save_reqs: - if save_req.req_id not in self._alive_requests: - # TODO: should not happen anymore + req = self._alive_requests.get(save_req.req_id) + if req is None: logger.warning("request %s is not alive, skip saving", save_req.req_id) continue - req = self._alive_requests[save_req.req_id] meta.add_save_request(save_req) - req.sent_saving_count += 1 if (req.need_report_after_saving_finished and req.scheduled_saving_count == req.sent_saving_count): @@ -760,8 +850,6 @@ def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnector for finish_req in self._waiting_to_finish_requests: meta.add_finish_request(finish_req) self._waiting_to_finish_requests = [] - - # logger.warning("build_connector_meta: %r", meta) return meta def start_save_kvcache_async(self, req_id, token_ids, target_save_num): @@ -770,62 +858,54 @@ def start_save_kvcache_async(self, req_id, token_ids, target_save_num): "instance_id": self._extra_config.instance_id, "block_keys": [], "token_ids": token_ids, - "write_timeout_seconds": 30 + "write_timeout_seconds": self._extra_config.write_timeout_seconds, } - logger.debug("start_write_cache req: %s", request) try: response = self._manager_client.start_write_cache(request) except Exception as e: - logger.warning("start_write_cache error, skip this saving, exception: %s", e) + logger.warning("start_write_cache error, skip saving: %s", e) with self._canceled_save_request_ids_lock: self._canceled_save_request_ids.append(req_id) return - # call manager start write - logger.debug("start_write_cache resp: %s", response) + locations = response["locations"] write_session_id = response["write_session_id"] - # check if success - if len(locations) == 0: + if not locations: try: self._manager_client.finish_write_cache({ - "trace_id": "test_test", + "trace_id": "finish_%s" % write_session_id[:8], "instance_id": self._extra_config.instance_id, "write_session_id": write_session_id, - "success_blocks": { - "bool_masks": { - "offset": 0 - } - } + "success_blocks": {"bool_masks": {"offset": 0}}, }) except Exception as e: - logger.warning("finish_write_cache failed, write_session_id: %s, error: %s", write_session_id, e) + logger.warning("finish_write_cache failed, session: %s, error: %s", + write_session_id, e) with self._canceled_save_request_ids_lock: self._canceled_save_request_ids.append(req_id) return need_block_idx = self.parse_block_mask_to_save_indices(response, target_save_num) - logger.debug("target_save_num: %s, need_block_idx: %s", target_save_num, need_block_idx) - message = CoordinateMessage(time.time(), SendBlockStartEvent(request_id=req_id, - write_session_id=write_session_id, - locations=locations)) + message = CoordinateMessage(time.time(), SendBlockStartEvent( + request_id=req_id, write_session_id=write_session_id, locations=locations)) self._coordinator_client.send(CoordinateMsgSerializer.dumps(message)) with self._waiting_to_save_requests_lock: self._waiting_to_save_requests.append(SaveRequest( - req_id, - locations, - need_block_idx, - write_session_id - )) + req_id, locations, need_block_idx, write_session_id)) def handle_canceled_save_req(self): - canceled_save_req_ids = [] with self._canceled_save_request_ids_lock: - canceled_save_req_ids = self._canceled_save_request_ids + canceled = self._canceled_save_request_ids self._canceled_save_request_ids = [] - for canceled_req_id in canceled_save_req_ids: - req = self._alive_requests[canceled_req_id] + for req_id in canceled: + # Cancellations come from http_executor threads; the request may + # already have been finished and removed by the scheduler loop. + req = self._alive_requests.get(req_id) + if req is None: + logger.warning("canceled save for unknown request %s, skip", req_id) + continue req.sent_saving_count += 1 if (req.need_report_after_saving_finished and req.scheduled_saving_count == req.sent_saving_count): @@ -833,47 +913,37 @@ def handle_canceled_save_req(self): self._alive_requests.pop(req.req_id) def get_finished_count(self): - # only rank0 will return finished + # Only rank0 reports finished requests. return 1 def update_connector_output(self, connector_output: KVConnectorOutput): - """ - Update KVConnector state from worker-side connectors output. - - Args: - connector_output (KVConnectorOutput): the worker-side - connectors output. - """ - return - def parse_block_mask_to_save_indices(self, response: dict, target_save_num: int) -> list[int]: - # 从response中提取block_mask + def parse_block_mask_to_save_indices(self, response: dict, target_save_num: int) -> List[int]: block_mask = response.get("block_mask", {}) - save_indices = [] if "offset" in block_mask: - offset = block_mask["offset"] - for idx in range(offset, target_save_num): - save_indices.append(idx) - else: - bool_masks = block_mask.get("bool_masks", {}).get("values", []) - # 找出所有为False的索引(需要保存的block) - for idx, is_saved in enumerate(bool_masks): - if not is_saved: # False表示需要保存 - save_indices.append(idx) - - return save_indices - - def request_finished( - self, - request: "Request", - block_ids: list[int], - ) -> Tuple[bool, Optional[dict[str, Any]]]: - if request.request_id not in self._alive_requests: - logger.info("request_finished not alive request: %s", request.request_id) + return list(range(block_mask["offset"], target_save_num)) + values = block_mask.get("bool_masks", {}).get("values", []) + return [idx for idx, saved in enumerate(values) if not saved] + + # ------------------------------------------------------------------ # + # Request finish + # ------------------------------------------------------------------ # + def request_finished_all_groups( + self, request: "Request", + block_ids: Tuple[List[int], ...]) -> Tuple[bool, Optional[dict]]: + return self._finish_request(request) + + def request_finished(self, request: "Request", + block_ids: List[int]) -> Tuple[bool, Optional[dict]]: + return self._finish_request(request) + + def _finish_request(self, request: "Request") -> Tuple[bool, Optional[dict]]: + req = self._alive_requests.get(request.request_id) + if req is None: + logger.info("request_finished for unknown request: %s", request.request_id) return False, {} - req = self._alive_requests[request.request_id] extra_info = {"local_matched_token_num": req.local_matched_token_num, "remote_matched_token_num": req.remote_matched_token_num} @@ -882,8 +952,6 @@ def request_finished( self._alive_requests.pop(req.req_id) return True, extra_info - # This request still has some save requests waiting to be issued or canceled, - # delay finishing this request + # Saves still in flight; delay freeing the blocks until they land. req.need_report_after_saving_finished = True - return True, extra_info