diff --git a/.buildkite/gpu_suites.py b/.buildkite/gpu_suites.py index 0aec7c898..6a564ff27 100644 --- a/.buildkite/gpu_suites.py +++ b/.buildkite/gpu_suites.py @@ -43,24 +43,24 @@ ], "megatron": [ ("test_full_disk_weight_update.py", 4, "", {}), - ("test_quick_start_glm4_9B.py", 8, "", {}), + ("test_quick_start_glm4_9B.py", 8, "", {"ENABLE_EVAL": "0"}), ("test_glm4.7_30B_A3B_pd_mooncake.py", 8, "", {}), ( "test_qwen3_30B_A3B.py", 8, "", - {"USE_DEEPEP": "1", "USE_FP8_ROLLOUT": "1"}, + {"USE_DEEPEP": "1", "USE_FP8_ROLLOUT": "1", "ENABLE_EVAL": "0"}, ), ("test_qwen3.6_35B_A3B_pd_mooncake.py", 8, "", {"USE_DEEPEP": "1"}), ("test_qwen3_30B_A3B_r3.py", 8, "", {"USE_DEEPEP": "1", "USE_FP8_ROLLOUT": "1", "ENABLE_EVAL": "0"}), ("test_qwen3_30B_A3B_r3.py", 8, "", {"ENABLE_EVAL": "0"}), - ("test_qwen3_4B_ppo.py", 8, "", {}), - ("test_qwen3_4B_ppo_disaggregate.py", 8, "", {}), - ("test_qwen3_4B_ppo_train_critic_only.py", 8, "", {}), + ("test_qwen3_4B_ppo.py", 8, "", {"ENABLE_EVAL": "0"}), + ("test_qwen3_4B_ppo_disaggregate.py", 8, "", {"ENABLE_EVAL": "0"}), + ("test_qwen3_4B_ppo_train_critic_only.py", 8, "", {"ENABLE_EVAL": "0"}), ("test_ppo_logprob_entropy_gpu.py", 2, "", {}), ("test_release_train.py", 4, "", {}), ("test_qwen3_4B_streaming_partial_rollout.py", 8, "", {}), - ("test_moonlight_16B_A3B.py", 8, "", {}), + ("test_moonlight_16B_A3B.py", 8, "", {"ENABLE_EVAL": "0"}), ("test_moonlight_16B_A3B_r3.py", 8, "", {"ENABLE_EVAL": "0"}), ("test_mimo_7B_mtp_only_grad.py", 8, "", {}), ("test_qwen2.5_0.5B_debug_rollout_then_train.py", 8, "", {}), diff --git a/.claude/skills/add-tests-and-ci/SKILL.md b/.claude/skills/add-tests-and-ci/SKILL.md index a729a36d3..8a84af486 100644 --- a/.claude/skills/add-tests-and-ci/SKILL.md +++ b/.claude/skills/add-tests-and-ci/SKILL.md @@ -37,23 +37,16 @@ if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) ``` -- `run-ci-changed` extracts a top-level `NUM_GPUS = ` constant from added/modified `tests/test_*.py` and `tests/plugin_contracts/test_*.py`; if missing, it defaults to 8 GPUs. Set `NUM_GPUS = 0` for CPU-only tests. +- Set `NUM_GPUS = 0` for CPU-only tests, following the existing test metadata convention. - For GPU/e2e tests, follow the nearby file pattern (`prepare()`, `execute()`, `NUM_GPUS`, and any model/dataset constants). -### Step 3: Register Tests in GitHub CI +### Step 3: Register Tests in Buildkite CI -Whenever adding, moving, or renaming a test file, update the GitHub workflow template before finishing: +Whenever adding, moving, or renaming a test file, update its Buildkite registration before finishing: -1. Add the test to the appropriate matrix in `.github/workflows/pr-test.yml.j2`. - - CPU-only pytest/unit tests usually belong in `cpu-unittest` with `num_gpus: 0`. - - GPU/e2e tests should be placed beside the nearest similar model/path test with the matching `num_gpus` and environment fields. -2. Regenerate workflows: - -```bash -python .github/workflows/generate_github_workflows.py -``` - -3. Include both `.github/workflows/pr-test.yml.j2` and the generated `.github/workflows/pr-test.yml` in the change set. +1. Register CPU test files in the appropriate command list in `.buildkite/pipeline.yml`, beside similar tests. Agent CPU tests belong in `agent-adapter`. +2. Register GPU/e2e tests in `.buildkite/gpu_suites.py`, with the matching GPU count and environment settings. Update `.buildkite/pipeline.yml` when changing suite selection or wiring. +3. Include the registration changes with the tests. These files are the source of truth; there is no GitHub workflow regeneration step. Only skip fixed matrix registration when the test is intentionally helper-only or manually invoked; state that reason in the final response. @@ -64,34 +57,29 @@ Only skip fixed matrix registration when the test is intentionally helper-only o - Run repository-wide checks only when they are already part of the task or workflow. - Avoid documenting placeholder test commands that may not exist in the current tree. -### Step 5: Keep Workflow Template as Source of Truth +### Step 5: Keep Buildkite Sources in Sync For CI workflow changes unrelated to a new, moved, or renamed test: -1. Edit `.github/workflows/pr-test.yml.j2` -2. Regenerate workflows: - -```bash -python .github/workflows/generate_github_workflows.py -``` - -3. Include both the template and generated workflow file in the change set (`.j2` and `.yml`). If the user asked for a commit, commit both. +1. Edit `.buildkite/pipeline.yml` for always-on CPU commands and pipeline wiring. +2. Edit `.buildkite/gpu_suites.py` for generated GPU jobs rather than editing its generated output. +3. Keep suite definitions, selection, and `.buildkite/README.md` consistent when changing suites. ### Step 6: Provide Verifiable PR Notes Include: - Which tests were added/changed -- Where each new/renamed test was registered in `.github/workflows/pr-test.yml.j2` +- Where each new/renamed test was registered in `.buildkite/pipeline.yml` or `.buildkite/gpu_suites.py` - Exact commands executed - GPU assumptions for each test path - Why this coverage protects against regression ## Common Mistakes -- Editing generated workflow file only -- Relying on `run-ci-changed` discovery for a new test that should run in the regular PR matrix -- Forgetting `NUM_GPUS = 0` on a CPU-only changed test, causing `run-ci-changed` to default to 8 GPUs +- Editing generated GPU jobs instead of their source +- Relying on pytest discovery for a new test in a suite with an explicit file list +- Treating a green CPU build as GPU validation; GPU suites require the manual Buildkite gate - Adding a CPU pytest file that passes under `pytest tests/foo.py` but fails under CI's `python tests/foo.py` - Adding tests without following existing constants/conventions - Making tests too large or non-deterministic @@ -101,5 +89,6 @@ Include: - Pytest config: `pyproject.toml` - Tests: `tests/` -- CI template: `.github/workflows/pr-test.yml.j2` +- CI sources: `.buildkite/pipeline.yml`, `.buildkite/gpu_suites.py` +- Buildkite guide: `.buildkite/README.md` - CI guide: `docs/en/developer_guide/ci.md` diff --git a/.claude/skills/release/SKILL.md b/.claude/skills/release/SKILL.md new file mode 100644 index 000000000..8dd30e2bc --- /dev/null +++ b/.claude/skills/release/SKILL.md @@ -0,0 +1,40 @@ +--- +name: release +description: Prepare and verify a Vime release, including version metadata, Docker patch-stack validation, and release-specific checks. Use when cutting or auditing a Vime release. +--- + +# Release Vime + +Prepare a release without creating Git tags, GitHub releases, or publishing +images unless the user explicitly requests those external actions. + +## Establish the release baseline + +- Preserve unrelated changes and compare the previous Vime release tag. +- Confirm the package version, the pinned `BASE_IMAGE`, and the Docker patch + stack in `docker/patch/latest/`. +- Do not upgrade the vLLM base image as part of a release unless Slime has + upgraded its corresponding inference-image baseline. + +## Prepare the release PR + +- Update `setup.py` and `docs/conf.py` to the requested package version. +- Give `docker/version.txt` a new unique dated image tag. +- Review every remaining occurrence of the old Vime version rather than making + a blind repository-wide replacement. +- Verify every patch under `docker/patch/latest/` is consumed in Dockerfile + application order and applies to its target in separate clean checkouts of + the pinned vLLM and Megatron revisions. Do not validate patch application + against a dirty developer checkout. + +## Validate and publish + +- Run `python .claude/skills/release/scripts/check_release.py --repo . + --expected-version `, `python setup.py --version`, and + `git diff --check`. +- Build a candidate image from the release commit and run the required E2E + tests before promoting an image tag. +- Merge the green release PR, then create the matching Git tag and GitHub + release at its merge commit. +- Publish the versioned image first. Only update `vllm/vime:latest` after the + candidate has passed and every required vLLM patch has merged upstream. diff --git a/.claude/skills/release/scripts/check_release.py b/.claude/skills/release/scripts/check_release.py new file mode 100644 index 000000000..192a65f4f --- /dev/null +++ b/.claude/skills/release/scripts/check_release.py @@ -0,0 +1,100 @@ +#!/usr/bin/env python3 +"""Check Vime release metadata and its Docker patch stack.""" + +import argparse +import ast +import re +import sys +from pathlib import Path + + +def setup_version(path: Path) -> str: + tree = ast.parse(path.read_text()) + for node in ast.walk(tree): + if not isinstance(node, ast.Call) or getattr(node.func, "id", None) != "setup": + continue + for keyword in node.keywords: + if keyword.arg == "version": + return ast.literal_eval(keyword.value) + raise ValueError(f"setup version not found in {path}") + + +def assigned_string(path: Path, name: str) -> str: + tree = ast.parse(path.read_text()) + for node in tree.body: + if not isinstance(node, ast.Assign): + continue + if any(isinstance(target, ast.Name) and target.id == name for target in node.targets): + return ast.literal_eval(node.value) + raise ValueError(f"{name} not found in {path}") + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--repo", type=Path, default=Path.cwd()) + parser.add_argument("--expected-version") + args = parser.parse_args() + + repo = args.repo.resolve() + errors: list[str] = [] + package_version = setup_version(repo / "setup.py") + docs_version = assigned_string(repo / "docs/conf.py", "__version__") + if package_version != docs_version: + errors.append(f"setup.py={package_version} but docs/conf.py={docs_version}") + if args.expected_version and package_version != args.expected_version: + errors.append(f"release version is {package_version}, expected {args.expected_version}") + + dockerfile = (repo / "docker/Dockerfile").read_text() + image_tag = (repo / "docker/version.txt").read_text().strip() + if not re.fullmatch(r"nightly-dev-\d{8}[a-z]", image_tag): + errors.append(f"unexpected docker/version.txt format: {image_tag}") + if not re.search(r"^ARG BASE_IMAGE=", dockerfile, re.MULTILINE): + errors.append("docker/Dockerfile does not pin BASE_IMAGE") + if not re.search(r"^ARG PATCH_VERSION=latest$", dockerfile, re.MULTILINE): + errors.append("docker/Dockerfile must build from docker/patch/latest") + + patch_dir = repo / "docker/patch/latest" + patches = {path.name for path in patch_dir.glob("*.patch")} + copied = { + name + for name in re.findall(r"COPY docker/patch/\$\{PATCH_VERSION\}/([^\s]+\.patch)", dockerfile) + if "*" not in name + } + if "megatron*.patch" in dockerfile: + copied.add("megatron.patch") + if patches != copied: + errors.append( + "Dockerfile patch set differs from docker/patch/latest: " + f"only_patches={sorted(patches - copied)}, " + f"only_dockerfile={sorted(copied - patches)}" + ) + applied = set( + re.findall( + r"git apply(?:\s+--?[\w-]+)*\s+(?:/tmp/)?([^ \\]+\.patch)", + dockerfile, + ) + ) + if patches != applied: + errors.append( + "Dockerfile does not apply every patch: " + f"not_applied={sorted(patches - applied)}, " + f"unknown={sorted(applied - patches)}" + ) + for patch in sorted(patches): + if not (patch_dir / patch).read_text().startswith("diff --git "): + errors.append(f"invalid git patch: {patch}") + + justfile = (repo / "docker/justfile").read_text() + if 'VERSION="$(cat docker/version.txt | tr -d' not in justfile: + errors.append("docker/justfile does not source docker/version.txt") + + if errors: + print(*[f"ERROR: {error}" for error in errors], sep="\n", file=sys.stderr) + return 1 + + print(f"release={package_version}, image={image_tag}, " f"patches={','.join(sorted(patches))}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/README.md b/README.md index 275531f0f..a5b2b141f 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,8 @@ The vLLM community horizontally supports many LLM post-training frameworks, incl - [Quick Start](#quick-start) - [Agentic RL examples](#agentic-rl-examples) - [Arguments Walkthrough](#arguments-walkthrough) + - [Engine Deployment](#engine-deployment) + - [Correctness, Stability, and CI](#correctness-stability-and-ci) - [Code Reading Path](#code-reading-path) - [Developer Guide](#developer-guide) - [slime doc](#slime-doc) @@ -81,6 +83,29 @@ Arguments in Vime are divided into three categories: For complete usage instructions, please refer to the [Usage Documentation](docs/en/get_started/usage.md). +## Engine Deployment + +Vime keeps the Megatron and vLLM control surfaces close to the upstream engines while adding the RL dataflow around them. Beyond the argument pass-through described above, see: + +- [vLLM Config](docs/en/advanced/vllm-config.md) for optional YAML topology configuration, heterogeneous server groups, multi-model serving, and per-group overrides; +- [PD Disaggregation](docs/en/advanced/pd-disaggregation.md) for multi-turn and agentic workloads with different prefill/decode resource needs; +- router policies such as session affinity for multi-turn agents (see [vLLM Config](docs/en/advanced/vllm-config.md)); +- [Delta Weight Sync](docs/en/advanced/delta-weight-sync.md) for disk-based updates of disaggregated rollout engines; +- [External Rollout Engines](docs/en/advanced/external-rollout-engines.md) for serving managed outside the training job. Serving can use an independent environment; disk transport avoids an NCCL group between training and serving. Different GPU models or vendors still require compatible model formats, precision, and vLLM hardware support. + +## Correctness, Stability, and CI + +RL bugs can be silent. Vime keeps the dataflow explicit and supports separate rollout-only and train-only debugging paths. CPU unit tests, customization-hook contract tests, and GPU end-to-end suites protect different parts of this workflow. Buildkite runs always-on CPU checks; GPU suites require the manual gate, so a green CPU build is not GPU validation. + +Useful engineering docs: + +- [CI](docs/en/developer_guide/ci.md) +- [Debugging](docs/en/developer_guide/debug.md) +- [Reproducibility](docs/en/advanced/reproducibility.md) +- [Fault Tolerance](docs/en/advanced/fault-tolerance.md) +- [Trace Viewer](docs/en/developer_guide/trace.md) +- [Profiling](docs/en/developer_guide/profiling.md) + ## Code Reading Path Start from the training loop and follow the calls only as deep as needed: diff --git a/README_zh.md b/README_zh.md index 3b7a55b4d..18927a796 100644 --- a/README_zh.md +++ b/README_zh.md @@ -34,6 +34,8 @@ vLLM 社区横向支持许多 LLM post-training 框架,包括(按字母顺 - [快速开始](#快速开始) - [Agentic RL 示例](#agentic-rl-示例) - [参数说明](#参数说明) + - [Engine 部署](#engine-部署) + - [正确性、稳定性与 CI](#正确性稳定性与-ci) - [代码阅读路径](#代码阅读路径) - [开发指南](#开发指南) - [slime doc](#slime-doc) @@ -81,6 +83,29 @@ Vime 的参数分为三类: 完整使用说明请查阅 [使用文档](docs/zh/get_started/usage.md)。 +## Engine 部署 + +Vime 在 Megatron 与 vLLM 原生控制接口外组织 RL 数据流。除上述参数透传外,请参阅: + +- [vLLM Config](docs/zh/advanced/vllm-config.md):可选的 YAML 拓扑配置、异构 server group、多模型 serving 和 per-group override; +- [PD Disaggregation](docs/zh/advanced/pd-disaggregation.md):面向 prefill/decode 资源需求不同的多轮和 agentic 工作负载; +- 面向多轮 agent 的 session affinity 等 router policy,见 [vLLM Config](docs/zh/advanced/vllm-config.md); +- [Delta Weight Sync](docs/zh/advanced/delta-weight-sync.md):分离部署 rollout engine 的磁盘增量更新; +- [External Rollout Engines](docs/zh/advanced/external-rollout-engines.md):由训练任务外部管理 serving。Serving 可以使用独立环境;disk transport 无需训练端和 serving 端组成 NCCL group。不同 GPU 型号或厂商仍需满足模型格式、精度和 vLLM 硬件支持的兼容要求。 + +## 正确性、稳定性与 CI + +RL bug 可能不会立即报错。Vime 保持显式数据流,支持 rollout-only 和 train-only 分离调试。CPU 单测、customization hook contract test 和 GPU 端到端测试分别保护这条链路的不同部分。Buildkite 自动运行 CPU 检查;GPU suite 需要手动开启 gate,因此 CPU 构建通过不代表 GPU 验证通过。 + +相关工程文档: + +- [CI](docs/zh/developer_guide/ci.md) +- [Debugging](docs/zh/developer_guide/debug.md) +- [Reproducibility](docs/zh/advanced/reproducibility.md) +- [Fault Tolerance](docs/zh/advanced/fault-tolerance.md) +- [Trace Viewer](docs/zh/developer_guide/trace.md) +- [Profiling](docs/zh/developer_guide/profiling.md) + ## 代码阅读路径 建议从训练主循环开始,只在需要时继续深入: diff --git a/docker/Dockerfile b/docker/Dockerfile index a6acb47b9..8617855f4 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -12,6 +12,7 @@ ARG FLASH_QLA_COMMIT=821fd9d37ede18fdc2a4e707fefe3770bfc32e58 ARG TRANSFORMER_ENGINE_COMMIT=c9877beb87ad7e711e1869dd0b5062167ede447a ARG TRANSFORMER_ENGINE_CUDA_ARCHS=90;100a;103a ARG TMS_COMMIT=8d30c59ca12a68d9deccbc9c6599076a1218cbc5 +ARG TMS_CUDA_MAJOR= ARG ENABLE_CUDA_13=1 ARG FA2_MAX_JOBS=64 @@ -97,7 +98,7 @@ RUN git clone https://github.com/NVIDIA/Megatron-LM.git --recursive && \ # zhuzilin fork builds, grouped together right after Megatron-LM: # torch_memory_saver, plus the GLM-5 train/rollout alignment kernels. -RUN TMS_CUDA_MAJOR="$(python -c 'import torch; print(torch.version.cuda.split(".")[0])')" && \ +RUN TMS_CUDA_MAJOR="${TMS_CUDA_MAJOR:-$(python -c 'import torch; print(torch.version.cuda.split(".")[0])')}" && \ export TMS_CUDA_MAJOR && \ pip install git+https://github.com/zhuzilin/torch_memory_saver.git@${TMS_COMMIT} --no-cache-dir --force-reinstall @@ -118,7 +119,7 @@ RUN git clone https://github.com/zhuzilin/DeepEP.git /root/DeepEP && \ TORCH_CUDA_ARCH_LIST="${CUDA_ARCHS}" MAX_JOBS=64 python setup.py bdist_wheel && \ pip install --force-reinstall --no-deps dist/deep_ep-*.whl && \ cd /root/ && rm -rf DeepEP -RUN pip install nvidia-modelopt[torch]>=0.37.0 --no-build-isolation +RUN pip install "nvidia-modelopt[torch]>=0.37.0" --no-build-isolation COPY requirements.txt /tmp/requirements.txt RUN pip install --ignore-installed PyJWT && \ @@ -148,16 +149,21 @@ RUN cd Megatron-LM && \ rm -f megatron*.patch && \ pip install -e . -# Patch vLLM with vime's local fixes. vLLM is a pip install (not a git checkout) -# so apply with plain `git apply` (no --3way). Pull-weights lands first because -# the general patch also updates gpu_worker.py against the resulting line layout. +# Patch vLLM with vime's local fixes. vLLM is a pip install (not a git checkout), +# so apply the independently maintained patches in their validated order. COPY docker/patch/${PATCH_VERSION}/vllm-pull_weights.patch /tmp/vllm-pull_weights.patch COPY docker/patch/${PATCH_VERSION}/vllm.patch /tmp/vllm.patch +COPY docker/patch/${PATCH_VERSION}/vllm-pd-request-metrics.patch /tmp/vllm-pd-request-metrics.patch +COPY docker/patch/${PATCH_VERSION}/vllm-inflight-queue-diagnostics.patch /tmp/vllm-inflight-queue-diagnostics.patch RUN VLLM_SITE="$(python3 -c 'import os, vllm; print(os.path.dirname(os.path.dirname(vllm.__file__)))')" && \ cd "$VLLM_SITE" && \ git apply -v /tmp/vllm-pull_weights.patch && \ git apply -v --allow-empty /tmp/vllm.patch && \ - rm /tmp/vllm-pull_weights.patch /tmp/vllm.patch + git apply -v /tmp/vllm-pd-request-metrics.patch && \ + git apply -v /tmp/vllm-inflight-queue-diagnostics.patch && \ + rm /tmp/vllm-pull_weights.patch /tmp/vllm.patch \ + /tmp/vllm-pd-request-metrics.patch \ + /tmp/vllm-inflight-queue-diagnostics.patch # ====================================== Install main package ============================================ diff --git a/docker/patch/latest/vllm-inflight-queue-diagnostics.patch b/docker/patch/latest/vllm-inflight-queue-diagnostics.patch new file mode 100644 index 000000000..833cf85c7 --- /dev/null +++ b/docker/patch/latest/vllm-inflight-queue-diagnostics.patch @@ -0,0 +1,304 @@ +diff --git a/tests/entrypoints/serve/instrumentator/test_basic.py b/tests/entrypoints/serve/instrumentator/test_basic.py +index 73a97c4fa26..d238bfb3de9 100644 +--- a/tests/entrypoints/serve/instrumentator/test_basic.py ++++ b/tests/entrypoints/serve/instrumentator/test_basic.py +@@ -205,6 +205,29 @@ async def test_server_load(server: RemoteOpenAIServer): + assert response.json().get("server_load") == 0 + + ++@pytest.mark.asyncio ++async def test_server_load_with_inflight_diagnostics(): ++ from vllm.entrypoints.serve.instrumentator.basic import get_server_load_metrics ++ ++ mock_request = Mock(spec=Request) ++ mock_request.app.state.server_load_metrics = 2 ++ mock_request.app.state.engine_client = AsyncMock() ++ mock_request.app.state.engine_client.get_inflight_queue_diagnostics.return_value = [ ++ {"data_parallel_rank": 0, "queues": []} ++ ] ++ ++ response = await get_server_load_metrics( ++ mock_request, include_inflight=True, inflight_limit=10 ++ ) ++ ++ assert response.body == ( ++ b'{"server_load":2,"inflight":[{"data_parallel_rank":0,"queues":[]}]}' ++ ) ++ mock_request.app.state.engine_client.get_inflight_queue_diagnostics.assert_awaited_once_with( ++ 10 ++ ) ++ ++ + @pytest.mark.asyncio + async def test_health_check_engine_dead_error(): + # Import the health function directly to test it in isolation +diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py +index 1ef63cbb042..506b75039fb 100644 +--- a/tests/v1/core/test_scheduler.py ++++ b/tests/v1/core/test_scheduler.py +@@ -1,6 +1,7 @@ + # SPDX-License-Identifier: Apache-2.0 + # SPDX-FileCopyrightText: Copyright contributors to the vLLM project + import dataclasses ++import time + from concurrent.futures import Future + from unittest.mock import Mock + +@@ -147,6 +148,29 @@ def test_get_num_unfinished_requests(): + assert scheduler.get_num_unfinished_requests() == len(requests) - i - 1 + + ++def test_get_inflight_queue_diagnostics(): ++ scheduler = create_scheduler() ++ requests = create_requests(num_requests=3) ++ for request in requests: ++ request.arrival_time = time.time() - 1 ++ scheduler.add_request(request) ++ ++ scheduler.running.append(scheduler.waiting.pop_request()) ++ scheduler.running[0].status = RequestStatus.RUNNING ++ ++ diagnostics = scheduler.get_inflight_queue_diagnostics(limit=2) ++ ++ assert diagnostics["data_parallel_rank"] == 0 ++ assert [queue["name"] for queue in diagnostics["queues"]] == [ ++ "running", ++ "waiting", ++ "skipped_waiting", ++ ] ++ assert [len(queue["requests"]) for queue in diagnostics["queues"]] == [1, 1, 0] ++ assert diagnostics["queues"][0]["requests"][0]["status"] == "RUNNING" ++ assert diagnostics["queues"][0]["requests"][0]["age_seconds"] >= 1 ++ ++ + @pytest.mark.parametrize( + "enable_prefix_caching, prompt_logprobs", + [ +diff --git a/vllm/engine/protocol.py b/vllm/engine/protocol.py +index 63f6970e056..13a97bfe64b 100644 +--- a/vllm/engine/protocol.py ++++ b/vllm/engine/protocol.py +@@ -296,3 +296,7 @@ class EngineClient(ABC): + async def get_weight_version(self) -> str: + """Return the latest committed weight version.""" + raise NotImplementedError ++ ++ async def get_inflight_queue_diagnostics(self, limit: int) -> list[dict[str, Any]]: ++ """Return bounded snapshots of in-flight request queues.""" ++ raise NotImplementedError +diff --git a/vllm/entrypoints/serve/instrumentator/basic.py b/vllm/entrypoints/serve/instrumentator/basic.py +index be091a1f433..73f5bbcf9b2 100644 +--- a/vllm/entrypoints/serve/instrumentator/basic.py ++++ b/vllm/entrypoints/serve/instrumentator/basic.py +@@ -28,7 +28,9 @@ def engine_client(request: Request) -> EngineClient: + + + @router.get("/load") +-async def get_server_load_metrics(request: Request): ++async def get_server_load_metrics( ++ request: Request, include_inflight: bool = False, inflight_limit: int = 100 ++): + # This endpoint returns the current server load metrics. + # It tracks requests utilizing the GPU from the following routes: + # - /v1/responses +@@ -47,7 +49,12 @@ async def get_server_load_metrics(request: Request): + # - /rerank + # - /v1/rerank + # - /v2/rerank +- return JSONResponse(content={"server_load": request.app.state.server_load_metrics}) ++ content = {"server_load": request.app.state.server_load_metrics} ++ if include_inflight: ++ content["inflight"] = await engine_client( ++ request ++ ).get_inflight_queue_diagnostics(inflight_limit) ++ return JSONResponse(content=content) + + + @router.get("/version") +diff --git a/vllm/v1/core/sched/interface.py b/vllm/v1/core/sched/interface.py +index c15296dd051..273fda18929 100644 +--- a/vllm/v1/core/sched/interface.py ++++ b/vllm/v1/core/sched/interface.py +@@ -3,7 +3,7 @@ + import enum + from abc import ABC, abstractmethod + from collections.abc import Iterable +-from typing import TYPE_CHECKING ++from typing import TYPE_CHECKING, Any + + from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalRegistry + +@@ -235,6 +235,11 @@ class SchedulerInterface(ABC): + """Returns (num_running_reqs, num_waiting_reqs).""" + raise NotImplementedError + ++ @abstractmethod ++ def get_inflight_queue_diagnostics(self, limit: int) -> dict[str, Any]: ++ """Return a bounded snapshot of in-flight request queues.""" ++ raise NotImplementedError ++ + def get_kv_cache_usage(self) -> float: + """Returns the fraction of the KV cache currently in use (0.0-1.0).""" + return 0.0 +diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py +index d9d86668a62..8fda1fd42f5 100644 +--- a/vllm/v1/core/sched/scheduler.py ++++ b/vllm/v1/core/sched/scheduler.py +@@ -2391,6 +2391,40 @@ class Scheduler(SchedulerInterface): + """Returns (num_running_reqs, num_waiting_reqs).""" + return len(self.running), len(self.waiting) + len(self.skipped_waiting) + ++ def get_inflight_queue_diagnostics(self, limit: int) -> dict[str, Any]: ++ """Return a bounded snapshot of in-flight request queues.""" ++ now = time.time() ++ remaining = max(0, min(limit, 1000)) ++ queues = [] ++ ++ for name, requests in ( ++ ("running", self.running), ++ ("waiting", self.waiting), ++ ("skipped_waiting", self.skipped_waiting), ++ ): ++ entries = [] ++ for request in requests: ++ if remaining == 0: ++ break ++ entries.append( ++ { ++ "request_id": request.request_id, ++ "status": request.status.name, ++ "age_seconds": round(max(0.0, now - request.arrival_time), 3), ++ "prompt_tokens": request.num_prompt_tokens, ++ "output_tokens": request.num_output_tokens, ++ } ++ ) ++ remaining -= 1 ++ queues.append( ++ {"name": name, "num_requests": len(requests), "requests": entries} ++ ) ++ ++ return { ++ "data_parallel_rank": self.parallel_config.data_parallel_rank, ++ "queues": queues, ++ } ++ + def get_kv_cache_usage(self) -> float: + """Returns the fraction of the KV cache currently in use (0.0-1.0).""" + return self.kv_cache_manager.usage +diff --git a/vllm/v1/engine/async_llm.py b/vllm/v1/engine/async_llm.py +index fdeebdcceb5..f4815325186 100644 +--- a/vllm/v1/engine/async_llm.py ++++ b/vllm/v1/engine/async_llm.py +@@ -1240,3 +1240,6 @@ class AsyncLLM(EngineClient): + async def get_weight_version(self) -> str: + """Return the latest committed weight version.""" + return await self.engine_core.get_weight_version_async() ++ ++ async def get_inflight_queue_diagnostics(self, limit: int) -> list[dict[str, Any]]: ++ return await self.engine_core.get_inflight_queue_diagnostics_async(limit) +diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py +index 06ebfb73051..4572dd9ec1e 100644 +--- a/vllm/v1/engine/core.py ++++ b/vllm/v1/engine/core.py +@@ -996,6 +996,9 @@ class EngineCore: + """Return the latest committed weight version.""" + return self._weight_version + ++ def get_inflight_queue_diagnostics(self, limit: int) -> dict[str, Any]: ++ return self.scheduler.get_inflight_queue_diagnostics(limit) ++ + def preprocess_add_request(self, request: EngineCoreRequest) -> tuple[Request, int]: + """Preprocess the request. + +diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py +index fa856461326..9627f1d79b2 100644 +--- a/vllm/v1/engine/core_client.py ++++ b/vllm/v1/engine/core_client.py +@@ -209,6 +209,9 @@ class EngineCoreClient(ABC): + def get_weight_version(self) -> str: + raise NotImplementedError + ++ def get_inflight_queue_diagnostics(self, limit: int) -> list[dict[str, Any]]: ++ raise NotImplementedError ++ + async def execute_dummy_batch_async(self) -> None: + raise NotImplementedError + +@@ -218,6 +221,11 @@ class EngineCoreClient(ABC): + async def get_weight_version_async(self) -> str: + raise NotImplementedError + ++ async def get_inflight_queue_diagnostics_async( ++ self, limit: int ++ ) -> list[dict[str, Any]]: ++ raise NotImplementedError ++ + def abort_requests(self, request_ids: list[str]) -> None: + raise NotImplementedError + +@@ -414,6 +422,9 @@ class InprocClient(EngineCoreClient): + def get_weight_version(self) -> str: + return self.engine_core.get_weight_version() + ++ def get_inflight_queue_diagnostics(self, limit: int) -> list[dict[str, Any]]: ++ return [self.engine_core.get_inflight_queue_diagnostics(limit)] ++ + def add_lora(self, lora_request: LoRARequest) -> bool: + return self.engine_core.add_lora(lora_request) + +@@ -1028,6 +1039,9 @@ class SyncMPClient(MPClient): + def get_weight_version(self) -> str: + return self.call_utility("get_weight_version") + ++ def get_inflight_queue_diagnostics(self, limit: int) -> list[dict[str, Any]]: ++ return [self.call_utility("get_inflight_queue_diagnostics", limit)] ++ + def collective_rpc( + self, + method: str | Callable[..., _R], +@@ -1272,6 +1286,12 @@ class AsyncMPClient(MPClient): + async def get_weight_version_async(self) -> str: + return await self.call_utility_async("get_weight_version") + ++ async def get_inflight_queue_diagnostics_async( ++ self, limit: int ++ ) -> list[dict[str, Any]]: ++ result = await self.call_utility_async("get_inflight_queue_diagnostics", limit) ++ return [result] ++ + async def add_lora_async(self, lora_request: LoRARequest) -> bool: + return await self.call_utility_async("add_lora", lora_request) + +@@ -1607,6 +1627,18 @@ class DPLBAsyncMPClient(DPAsyncMPClient): + ) + )[0] + ++ async def get_inflight_queue_diagnostics_async( ++ self, limit: int ++ ) -> list[dict[str, Any]]: ++ return await asyncio.gather( ++ *[ ++ self._call_utility_async( ++ "get_inflight_queue_diagnostics", limit, engine=engine ++ ) ++ for engine in self.core_engines ++ ] ++ ) ++ + @staticmethod + async def process_engine_outputs( + self: "DPLBAsyncMPClient", outputs: EngineCoreOutputs +diff --git a/vllm/v1/engine/llm_engine.py b/vllm/v1/engine/llm_engine.py +index 32096079e06..994cda54a53 100644 +--- a/vllm/v1/engine/llm_engine.py ++++ b/vllm/v1/engine/llm_engine.py +@@ -443,6 +443,9 @@ class LLMEngine: + """Return the latest committed weight version.""" + return self.engine_core.get_weight_version() + ++ def get_inflight_queue_diagnostics(self, limit: int) -> list[dict[str, Any]]: ++ return [self.engine_core.get_inflight_queue_diagnostics(limit)] ++ + def apply_model(self, func: Callable[[nn.Module], _R]) -> list[_R]: + return self.collective_rpc("apply_model", args=(func,)) + diff --git a/docker/patch/latest/vllm-pd-request-metrics.patch b/docker/patch/latest/vllm-pd-request-metrics.patch new file mode 100644 index 000000000..1a878e024 --- /dev/null +++ b/docker/patch/latest/vllm-pd-request-metrics.patch @@ -0,0 +1,479 @@ +diff --git a/rust/src/engine-core-client/src/protocol/output.rs b/rust/src/engine-core-client/src/protocol/output.rs +index cc7541eae1b..6c83765f9a1 100644 +--- a/rust/src/engine-core-client/src/protocol/output.rs ++++ b/rust/src/engine-core-client/src/protocol/output.rs +@@ -130,6 +130,8 @@ pub struct EngineCoreOutput { + /// the Rust frontend does not yet surface it in responses. + #[serde(default)] + pub spec_decode_metrics: Option, ++ #[serde(default)] ++ pub remote_kv_wait_time: Option, + } + + impl EngineCoreOutput { +@@ -449,6 +451,7 @@ mod tests { + mm_cache_miss_hashes: None, + new_sampling_mask: None, + spec_decode_metrics: None, ++ remote_kv_wait_time: None, + }, + ], + scheduler_stats: None, +diff --git a/rust/src/engine-core-client/src/tests/client.rs b/rust/src/engine-core-client/src/tests/client.rs +index 6f4e53ff1b4..04b4d49a561 100644 +--- a/rust/src/engine-core-client/src/tests/client.rs ++++ b/rust/src/engine-core-client/src/tests/client.rs +@@ -2744,6 +2744,7 @@ fn python_msgpack_fixtures_match_rust_encoding() { + mm_cache_miss_hashes: None, + new_sampling_mask: None, + spec_decode_metrics: None, ++ remote_kv_wait_time: None, + }, + ], + scheduler_stats: None, +diff --git a/rust/src/engine-core-client/src/tests/python_compat.py b/rust/src/engine-core-client/src/tests/python_compat.py +index d3705cc3918..498f8678397 100755 +--- a/rust/src/engine-core-client/src/tests/python_compat.py ++++ b/rust/src/engine-core-client/src/tests/python_compat.py +@@ -100,6 +100,8 @@ class EngineCoreOutput( + num_nans_in_logits: int = 0 + mm_cache_miss_hashes: list[str] | None = None + new_sampling_mask: object | None = None ++ spec_decode_metrics: object | None = None ++ remote_kv_wait_time: float | None = None + + + class EngineCoreOutputs( +diff --git a/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py b/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py +index 99cf457935f..a05ffdb2233 100644 +--- a/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py ++++ b/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py +@@ -23,6 +23,7 @@ from vllm.renderers import renderer_from_config + from vllm.renderers.online_renderer import OnlineRenderer + from vllm.sampling_params import SamplingParams + from vllm.v1.engine.async_llm import AsyncLLM ++from vllm.v1.metrics.stats import RequestStateStats + + MODEL_NAME = "openai-community/gpt2" + BASE_MODEL_PATHS = [ +@@ -125,6 +126,7 @@ def _make_request_output( + logprobs: list[dict[int, Any] | None] | None = None, + num_cached_tokens: int | None = None, + index: int = 0, ++ metrics: RequestStateStats | None = None, + ) -> RequestOutput: + return RequestOutput( + request_id=request_id, +@@ -142,7 +144,7 @@ def _make_request_output( + ) + ], + finished=finished, +- metrics=None, ++ metrics=metrics, + lora_request=None, + encoder_prompt=None, + encoder_prompt_token_ids=None, +@@ -197,12 +199,54 @@ async def test_serve_tokens_skips_mm_cache_for_remote_engine_execution(): + response = await serving.serve_tokens(request) + + assert isinstance(response, GenerateResponse) ++ assert response.request_metrics is None + assert ( + serving.online_renderer.preprocess_completion.call_args.kwargs["skip_mm_cache"] + is True + ) + + ++@pytest.mark.asyncio ++async def test_serve_tokens_returns_enabled_request_metrics(): ++ engine = _mock_engine() ++ engine.get_weight_version = AsyncMock(return_value="v1") ++ metrics = RequestStateStats( ++ queued_ts=1.0, ++ scheduled_ts=2.0, ++ first_token_ts=5.0, ++ last_token_ts=9.0, ++ first_token_latency=6.0, ++ remote_kv_wait_time=0.75, ++ ) ++ ++ async def mock_generate(*args, **kwargs): ++ yield _make_request_output( ++ "req-1", ++ token_ids=[10], ++ finish_reason="stop", ++ finished=True, ++ metrics=metrics, ++ ) ++ ++ engine.generate = MagicMock(side_effect=mock_generate) ++ serving = _build_serving_tokens(engine, enable_per_request_metrics=True) ++ request = GenerateRequest( ++ token_ids=[1, 2, 3], ++ sampling_params=SamplingParams(max_tokens=1), ++ model=MODEL_NAME, ++ stream=False, ++ ) ++ ++ response = await serving.serve_tokens(request) ++ ++ assert isinstance(response, GenerateResponse) ++ assert response.request_metrics is not None ++ assert response.request_metrics.queue_time_ms == 1000.0 ++ assert response.request_metrics.time_to_first_token_ms == 3000.0 ++ assert response.request_metrics.generation_time_ms == 4000.0 ++ assert response.request_metrics.remote_kv_wait_time_ms == 750.0 ++ ++ + @pytest.mark.asyncio + async def test_serve_tokens_threads_session_id_header_to_engine(): + engine = _mock_engine() +diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py +index 920823baeb8..11e487e5ae4 100644 +--- a/tests/v1/core/test_scheduler.py ++++ b/tests/v1/core/test_scheduler.py +@@ -2026,6 +2026,12 @@ def test_kv_connector_basic(is_async: bool): + + # Ensure ScheduleOutput is correct. + output = scheduler.schedule() ++ for request in requests: ++ if is_async: ++ assert request.remote_kv_wait_time > 0 ++ assert request.remote_kv_wait_started_at is None ++ else: ++ assert request.remote_kv_wait_time == 0 + _assert_right_scheduler_output( + output=output, + num_requests=NUM_REQUESTS, +diff --git a/tests/v1/engine/test_output_processor.py b/tests/v1/engine/test_output_processor.py +index f578aae7f01..6fcc491bd10 100644 +--- a/tests/v1/engine/test_output_processor.py ++++ b/tests/v1/engine/test_output_processor.py +@@ -23,6 +23,7 @@ from vllm.tokenizers import TokenizerLike + from vllm.v1.engine import ( + EngineCoreEvent, + EngineCoreEventType, ++ EngineCoreOutput, + EngineCoreOutputs, + EngineCoreRequest, + FinishReason, +@@ -32,7 +33,33 @@ from vllm.v1.engine.output_processor import ( + RequestOutputCollector, + RequestState, + ) +-from vllm.v1.metrics.stats import IterationStats, SchedulerStats ++from vllm.v1.metrics.stats import IterationStats, RequestStateStats, SchedulerStats ++ ++ ++def test_remote_kv_wait_time_is_added_to_request_stats(): ++ output_processor = object.__new__(OutputProcessor) ++ request_stats = RequestStateStats() ++ request_state = MagicMock( ++ stats=request_stats, ++ routed_experts_chunks=[], ++ is_prefilling=False, ++ ) ++ request_state.make_request_output.return_value = None ++ output_processor.request_states = {"request": request_state} ++ output_processor._update_stats_from_output = MagicMock() ++ ++ output_processor.process_outputs( ++ [ ++ EngineCoreOutput( ++ request_id="request", ++ new_token_ids=[], ++ pooling_output=MagicMock(), ++ remote_kv_wait_time=0.75, ++ ) ++ ] ++ ) ++ ++ assert request_stats.remote_kv_wait_time == 0.75 + + + @pytest.mark.parametrize("flat_logprobs", [False, True]) +diff --git a/tests/v1/test_request.py b/tests/v1/test_request.py +index be417b9b2ff..3b2d33e09db 100644 +--- a/tests/v1/test_request.py ++++ b/tests/v1/test_request.py +@@ -40,3 +40,28 @@ def test_request_copies_session_id_from_engine_core_request(): + request = Request.from_engine_core_request(engine_request, block_hasher=None) + + assert request.session_id == "session-1" ++ ++ ++def test_request_accumulates_remote_kv_waits(monkeypatch): ++ engine_request = EngineCoreRequest( ++ request_id="request-1", ++ prompt_token_ids=[1, 2, 3], ++ mm_features=None, ++ sampling_params=SamplingParams(max_tokens=1), ++ pooling_params=None, ++ arrival_time=0.0, ++ lora_request=None, ++ cache_salt=None, ++ data_parallel_rank=None, ++ ) ++ request = Request.from_engine_core_request(engine_request, block_hasher=None) ++ timestamps = iter([1.0, 3.0, 5.0, 9.0]) ++ monkeypatch.setattr("vllm.v1.request.time.monotonic", lambda: next(timestamps)) ++ ++ request.start_remote_kv_wait() ++ request.stop_remote_kv_wait() ++ request.start_remote_kv_wait() ++ request.stop_remote_kv_wait() ++ ++ assert request.remote_kv_wait_time == 6.0 ++ assert request.remote_kv_wait_started_at is None +diff --git a/vllm/entrypoints/generate/base/serving.py b/vllm/entrypoints/generate/base/serving.py +index d6ae0a20906..0d6746ddeba 100644 +--- a/vllm/entrypoints/generate/base/serving.py ++++ b/vllm/entrypoints/generate/base/serving.py +@@ -101,6 +101,11 @@ def build_per_request_timing_metrics( + queue_time_ms=queue_time_ms, + mean_itl_ms=mean_itl_ms, + tokens_per_second=tokens_per_second, ++ remote_kv_wait_time_ms=( ++ metrics.remote_kv_wait_time * 1000 ++ if metrics.remote_kv_wait_time is not None ++ else None ++ ), + ) + + +diff --git a/vllm/entrypoints/openai/engine/protocol.py b/vllm/entrypoints/openai/engine/protocol.py +index 9635ece46b0..f67434e2a43 100644 +--- a/vllm/entrypoints/openai/engine/protocol.py ++++ b/vllm/entrypoints/openai/engine/protocol.py +@@ -159,6 +159,7 @@ class PerRequestMetrics(OpenAIBaseModel): + tokens_per_second: float | None = None + # Experimental, subject to change. + speculative_decoding: SpeculativeDecodingMetrics | None = None ++ remote_kv_wait_time_ms: float | None = None + + + class RequestResponseMetadata(BaseModel): +diff --git a/vllm/entrypoints/scale_out/factories.py b/vllm/entrypoints/scale_out/factories.py +index 341dad13a86..dc31e1536fc 100644 +--- a/vllm/entrypoints/scale_out/factories.py ++++ b/vllm/entrypoints/scale_out/factories.py +@@ -53,6 +53,7 @@ def init_scale_out_state( + request_logger=request_logger, + return_tokens_as_token_ids=args.return_tokens_as_token_ids, + enable_prompt_tokens_details=args.enable_prompt_tokens_details, ++ enable_per_request_metrics=args.enable_per_request_metrics, + enable_log_outputs=args.enable_log_outputs, + force_no_detokenize=args.tokens_only, + ) +diff --git a/vllm/entrypoints/scale_out/token_in_token_out/protocol.py b/vllm/entrypoints/scale_out/token_in_token_out/protocol.py +index 71a17e9363d..1614d6c76b2 100644 +--- a/vllm/entrypoints/scale_out/token_in_token_out/protocol.py ++++ b/vllm/entrypoints/scale_out/token_in_token_out/protocol.py +@@ -20,7 +20,11 @@ from vllm.entrypoints.openai.completion.protocol import ( + CompletionRequest, + CompletionStreamResponse, + ) +-from vllm.entrypoints.openai.engine.protocol import StreamOptions, UsageInfo ++from vllm.entrypoints.openai.engine.protocol import ( ++ PerRequestMetrics, ++ StreamOptions, ++ UsageInfo, ++) + from vllm.logprobs import Logprob + from vllm.renderers import TokenizeParams + from vllm.sampling_params import SamplingParams +@@ -271,6 +275,10 @@ class GenerateResponse(BaseModel): + "ECTransfer parameters used for encoder-cache disaggregated serving." + ), + ) ++ request_metrics: PerRequestMetrics | None = Field( ++ default=None, ++ description="Per-request generation and remote KV wait timings.", ++ ) + + + class DerenderChatRequest(BaseModel): +diff --git a/vllm/entrypoints/scale_out/token_in_token_out/serving.py b/vllm/entrypoints/scale_out/token_in_token_out/serving.py +index 809f1a66ea5..4095f41f643 100644 +--- a/vllm/entrypoints/scale_out/token_in_token_out/serving.py ++++ b/vllm/entrypoints/scale_out/token_in_token_out/serving.py +@@ -14,6 +14,7 @@ from vllm.engine.protocol import EngineClient + from vllm.entrypoints.chat_utils import AsyncMultiModalItemTracker + from vllm.entrypoints.generate.base.serving import ( + GenerateBaseServing, ++ build_per_request_timing_metrics, + build_spec_decoding_metrics, + clamp_prompt_logprobs, + ) +@@ -71,6 +72,7 @@ class ServingTokens(GenerateBaseServing): + force_no_detokenize: bool = False, + return_tokens_as_token_ids: bool = False, + enable_prompt_tokens_details: bool = False, ++ enable_per_request_metrics: bool = False, + enable_log_outputs: bool = False, + ): + super().__init__( +@@ -81,6 +83,7 @@ class ServingTokens(GenerateBaseServing): + ) + self.online_renderer = online_renderer + self.enable_prompt_tokens_details = enable_prompt_tokens_details ++ self.enable_per_request_metrics = enable_per_request_metrics + self.enable_log_outputs = enable_log_outputs + self.force_no_detokenize = force_no_detokenize + if force_no_detokenize: +@@ -371,6 +374,15 @@ class ServingTokens(GenerateBaseServing): + + request_metadata.final_usage_info = usage + ++ request_metrics = ( ++ build_per_request_timing_metrics( ++ final_res.metrics, ++ num_generated_tokens, ++ ) ++ if self.enable_per_request_metrics ++ else None ++ ) ++ + response = GenerateResponse( + request_id=request_id, + created=created_time, +@@ -382,6 +394,7 @@ class ServingTokens(GenerateBaseServing): + prompt_logprobs=clamp_prompt_logprobs(final_res.prompt_logprobs), + kv_transfer_params=final_res.kv_transfer_params, + ec_transfer_params=final_res.ec_transfer_params, ++ request_metrics=request_metrics, + ) + + # Log complete response if output logging is enabled +diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py +index 51f75a63d6b..af5e1d73061 100644 +--- a/vllm/v1/core/sched/scheduler.py ++++ b/vllm/v1/core/sched/scheduler.py +@@ -1112,6 +1112,7 @@ class Scheduler(SchedulerInterface): + if load_kv_async: + # If loading async, allocate memory and put request + # into the WAITING_FOR_REMOTE_KV state. ++ request.start_remote_kv_wait() + request.status = RequestStatus.WAITING_FOR_REMOTE_KVS + step_skipped_waiting.prepend_request(request) + # Set num_computed_tokens even though KVs are not yet loaded. +@@ -1913,6 +1914,7 @@ class Scheduler(SchedulerInterface): + pooler_output = pooler_outputs[req_index] if pooler_outputs else None + kv_transfer_params = None + ec_transfer_params = None ++ remote_kv_wait_time = None + prefill_stats = None + status_before_stop = request.status + num_output_tokens_before = len(request._output_token_ids) +@@ -2024,6 +2026,8 @@ class Scheduler(SchedulerInterface): + finished = self._handle_stopped_request(request) + if finished: + kv_transfer_params, ec_transfer_params = self._free_request(request) ++ if request.remote_kv_wait_time: ++ remote_kv_wait_time = request.remote_kv_wait_time + + if status_before_stop == RequestStatus.RUNNING: + stopped_running_reqs.add(request) +@@ -2071,6 +2075,7 @@ class Scheduler(SchedulerInterface): + ), + kv_transfer_params=kv_transfer_params, + ec_transfer_params=ec_transfer_params, ++ remote_kv_wait_time=remote_kv_wait_time, + trace_headers=request.trace_headers, + routed_experts=routed_experts, + num_nans_in_logits=request.num_nans_in_logits, +@@ -2447,6 +2452,8 @@ class Scheduler(SchedulerInterface): + ) -> tuple[dict[str, Any] | None, dict[str, Any] | None]: + assert request.is_finished() + ++ request.stop_remote_kv_wait() ++ + self._inflight_prefills.discard(request) + connector_delay_free_blocks, kv_xfer_params = self._connector_finished(request) + +@@ -2832,6 +2839,7 @@ class Scheduler(SchedulerInterface): + if request.request_id not in self.finished_recving_kv_req_ids: + return False + self._update_waiting_for_remote_kv(request) ++ request.stop_remote_kv_wait() + if request.num_preemptions: + request.status = RequestStatus.PREEMPTED + else: +diff --git a/vllm/v1/engine/__init__.py b/vllm/v1/engine/__init__.py +index 5ae9ee0cac8..6161f0f704e 100644 +--- a/vllm/v1/engine/__init__.py ++++ b/vllm/v1/engine/__init__.py +@@ -230,9 +230,11 @@ class EngineCoreOutput( + new_sampling_mask: SamplingMaskLists | None = None + + # Per-request spec-decode acceptance; attached only on the final output. +- # Appended last so `array_like` positional serialization stays compatible. + spec_decode_metrics: RequestSpecDecodeMetrics | None = None + ++ # Appended last so `array_like` positional serialization stays compatible. ++ remote_kv_wait_time: float | None = None ++ + @property + def finished(self) -> bool: + return self.finish_reason is not None +diff --git a/vllm/v1/engine/output_processor.py b/vllm/v1/engine/output_processor.py +index 6238fbc8175..22d32d6b16d 100644 +--- a/vllm/v1/engine/output_processor.py ++++ b/vllm/v1/engine/output_processor.py +@@ -652,6 +652,13 @@ class OutputProcessor: + stop_reason = engine_core_output.stop_reason + kv_transfer_params = engine_core_output.kv_transfer_params + ec_transfer_params = engine_core_output.ec_transfer_params ++ if ( ++ engine_core_output.remote_kv_wait_time is not None ++ and req_state.stats is not None ++ ): ++ req_state.stats.remote_kv_wait_time = ( ++ engine_core_output.remote_kv_wait_time ++ ) + if engine_core_output.routed_experts is not None: + req_state.routed_experts_chunks.append( + engine_core_output.routed_experts +diff --git a/vllm/v1/metrics/stats.py b/vllm/v1/metrics/stats.py +index 3dbc5206ca9..2e2c2371d2c 100644 +--- a/vllm/v1/metrics/stats.py ++++ b/vllm/v1/metrics/stats.py +@@ -232,6 +232,8 @@ class RequestStateStats: + # first token latency + first_token_latency: float = 0.0 + ++ remote_kv_wait_time: float | None = None ++ + # Track if this request is corrupted (NaNs in logits) + is_corrupted: bool = False + +diff --git a/vllm/v1/request.py b/vllm/v1/request.py +index 8b453a09069..59f5d83881a 100644 +--- a/vllm/v1/request.py ++++ b/vllm/v1/request.py +@@ -101,6 +101,8 @@ class Request: + + # P/D: Connector-specific KV transfer parameters. + self.kv_transfer_params: dict[str, Any] | None = None ++ self.remote_kv_wait_started_at: float | None = None ++ self.remote_kv_wait_time = 0.0 + # E/P/D: Connector-specific encoder-cache transfer parameters. + self.ec_transfer_params: dict[str, Any] | None = None + +@@ -334,6 +336,16 @@ class Request: + ) -> None: + self.events.append(EngineCoreEvent.new_event(event_type, timestamp)) + ++ def start_remote_kv_wait(self) -> None: ++ assert self.remote_kv_wait_started_at is None ++ self.remote_kv_wait_started_at = time.monotonic() ++ ++ def stop_remote_kv_wait(self) -> None: ++ if self.remote_kv_wait_started_at is None: ++ return ++ self.remote_kv_wait_time += time.monotonic() - self.remote_kv_wait_started_at ++ self.remote_kv_wait_started_at = None ++ + def take_events(self) -> list[EngineCoreEvent] | None: + if not self.events: + return None diff --git a/docs/en/index.rst b/docs/en/index.rst index c8c552acf..c4b9167c1 100644 --- a/docs/en/index.rst +++ b/docs/en/index.rst @@ -34,6 +34,7 @@ Start by Use Case get_started/quick_start.md get_started/usage.md get_started/customization.md + get_started/agent.md get_started/qa.md .. toctree:: diff --git a/docs/zh/index.rst b/docs/zh/index.rst index 08f305224..69d7e479f 100644 --- a/docs/zh/index.rst +++ b/docs/zh/index.rst @@ -34,6 +34,7 @@ vime 构建于 `slime `_ 之上,slime 正是 G get_started/quick_start.md get_started/usage.md get_started/customization.md + get_started/agent.md get_started/qa.md .. toctree:: diff --git a/examples/delta_weight_sync/README.md b/examples/delta_weight_sync/README.md index 477dd84e2..51de88baa 100644 --- a/examples/delta_weight_sync/README.md +++ b/examples/delta_weight_sync/README.md @@ -7,7 +7,7 @@ directory; each engine's `/pull_weights` applies them into a host-local checkpoi host it spans, and the engines reload through the ordinary `update_weights_from_disk` path — vime only ever talks to one endpoint per engine. -See [Delta Weight Sync](../../docs/en/advanced/delta-weight-sync.md) for the full mechanism, +See [Delta Weight Sync](https://github.com/vllm-project/vime/blob/main/docs/en/advanced/delta-weight-sync.md) for the full mechanism, encodings, integrity checks, and shared-filesystem visibility hooks. ## Try it @@ -37,5 +37,5 @@ at `--update-weight-disk-dir`): For object-store-backed volumes that need an explicit commit/refresh to make writes visible across hosts, supply `--custom-update-weight-post-write-path` (trainer side) / -`--vllm-custom-pull-weights-pre-read-hook` (engine side) — no vendor-specific code lives in vime +`--custom-update-weight-pre-read-path` (engine side) — no vendor-specific code lives in vime or vllm; see the doc. diff --git a/examples/multi_agent/agent_system.py b/examples/multi_agent/agent_system.py index ebd521900..309d4c892 100644 --- a/examples/multi_agent/agent_system.py +++ b/examples/multi_agent/agent_system.py @@ -5,7 +5,11 @@ from copy import deepcopy from vime.rollout.rm_hub import batched_async_rm -from vime.rollout.vllm_rollout import _build_inference_sampling_params, _inference_generate_tokens_and_logprobs +from vime.rollout.vllm_rollout import ( + _build_inference_sampling_params, + _inference_generate_meta_info, + _inference_generate_tokens_and_logprobs, +) from vime.utils.http_utils import post from vime.utils.types import Sample @@ -50,6 +54,7 @@ async def generate_response(args, prompt, key): tokens=new_response_tokens, log_probs=new_response_log_probs, trainable=True, + meta_info=_inference_generate_meta_info(output), ) assert len(sample.rollout_log_probs) == sample.response_length, ( f"rollout logprob length mismatch: {len(sample.rollout_log_probs)} logprobs " @@ -201,6 +206,19 @@ async def run_agent_system(args, sample): args = deepcopy(args) # Deep copy args because rollout_with_multi_agents mutates them. args.sample = sample args.results_dict = {"solver": [], "rewriter": [], "selector": []} + # Every sample emitted below is a training sample split out of this one + # rollout execution (the input ``sample``). Stamp the shared rollout id on + # every collected sample at each return point so the per-rollout loss + # reducer aggregates the solver / rewriter / selector siblings as one + # rollout instead of N, and the by-rollout step splitter keeps them in + # the same step. Captured here because ``sample`` gets shadowed by zip- + # loop variables further down. + input_rollout_id = sample.index + + def _emit(samples_list): + for emitted_sample in samples_list: + emitted_sample.rollout_id = input_rollout_id + return samples_list problem_statement = sample.prompt tasks = [solver_worker(args, problem_statement, worker_id) for worker_id in range(args.num_parallel)] @@ -219,7 +237,7 @@ def reward_adjustment(samples, reward_weight): if len(previous_solutions) == 0: reward_adjustment(args.results_dict["solver"], args.incorrect_reward_weight) - return args.results_dict["solver"] + return _emit(args.results_dict["solver"]) # Rewriting tasks = [ @@ -241,7 +259,7 @@ def reward_adjustment(samples, reward_weight): if len(rewrited_solutions) == 0: reward_adjustment(args.results_dict["solver"], args.incorrect_reward_weight) reward_adjustment(args.results_dict["rewriter"], args.incorrect_reward_weight) - return args.results_dict["solver"] + args.results_dict["rewriter"] + return _emit(args.results_dict["solver"] + args.results_dict["rewriter"]) # Selection selector = SelectorAgent() @@ -249,7 +267,7 @@ def reward_adjustment(samples, reward_weight): if len(args.results_dict["selector"]) == 0: reward_adjustment(args.results_dict["solver"], args.incorrect_reward_weight) reward_adjustment(args.results_dict["rewriter"], args.incorrect_reward_weight) - return args.results_dict["solver"] + args.results_dict["rewriter"] + return _emit(args.results_dict["solver"] + args.results_dict["rewriter"]) assert ( len(args.results_dict["selector"]) == 1 @@ -278,4 +296,4 @@ def reward_adjustment(samples, reward_weight): reward_adjustment(args.results_dict["rewriter"], args.incorrect_reward_weight) reward_adjustment(args.results_dict["selector"], args.incorrect_reward_weight) - return args.results_dict["solver"] + args.results_dict["rewriter"] + args.results_dict["selector"] + return _emit(args.results_dict["solver"] + args.results_dict["rewriter"] + args.results_dict["selector"]) diff --git a/scripts/run-deepseek-r1.sh b/scripts/run-deepseek-r1.sh index 76cb982c2..2dd1a1541 100755 --- a/scripts/run-deepseek-r1.sh +++ b/scripts/run-deepseek-r1.sh @@ -153,6 +153,7 @@ ray job submit --address="http://127.0.0.1:8265" \ "MASTER_ADDR": "${MASTER_ADDR}", "PYTHONPATH": "/root/Megatron-LM/", "CUDA_DEVICE_MAX_CONNECTIONS": "1", + "NVSHMEM_DISABLE_NCCL": "1" } }' \ -- python3 train.py \ diff --git a/scripts/run-glm5-744B-A40B.sh b/scripts/run-glm5-744B-A40B.sh index a80c12de5..cb21eb58e 100755 --- a/scripts/run-glm5-744B-A40B.sh +++ b/scripts/run-glm5-744B-A40B.sh @@ -115,7 +115,6 @@ VLLM_ARGS=( # mtp # dsa - --vllm-attention-backend nsa --vllm-max-cudagraph-capture-size 40 --vllm-max-num-seqs 512 diff --git a/scripts/run-glm5.2-744B-A40B.sh b/scripts/run-glm5.2-744B-A40B.sh index 28bffba2f..5c446bda5 100644 --- a/scripts/run-glm5.2-744B-A40B.sh +++ b/scripts/run-glm5.2-744B-A40B.sh @@ -134,7 +134,6 @@ vllm: num_gpus: 64 num_gpus_per_engine: 64 overrides: - # Prefill uses data/expert parallelism with the high-throughput DeepEP backend. data_parallel_size: 64 enable_expert_parallel: true max_num_batched_tokens: 131072 @@ -149,7 +148,6 @@ vllm: num_gpus: 192 num_gpus_per_engine: 64 overrides: - # Decode uses the low-latency DeepEP backend. data_parallel_size: 64 enable_expert_parallel: true max_num_seqs: 768 diff --git a/scripts/run-mimo-7B-rl-eagle.sh b/scripts/run-mimo-7B-rl-eagle.sh index 0bcd3fca8..f26f262e7 100755 --- a/scripts/run-mimo-7B-rl-eagle.sh +++ b/scripts/run-mimo-7B-rl-eagle.sh @@ -111,7 +111,8 @@ VLLM_ARGS=( # for speculative decoding # sometimes flashinfer has IMA bugs. Use fa3 as instead - --vllm-attention-backend fa3 + --vllm-attention-backend FLASH_ATTN + --vllm-attention-config '{"flash_attn_version":3}' --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":4}' ) diff --git a/tests/_unit_stubs.py b/tests/_unit_stubs.py index 2f0bb62ba..64779531a 100644 --- a/tests/_unit_stubs.py +++ b/tests/_unit_stubs.py @@ -111,6 +111,12 @@ def add_cli_args( ): # noqa: ARG003 prefix = "router-" if use_router_prefix else "" dprefix = "router_" if use_router_prefix else "" + parser.add_argument( + f"--{prefix}log-level", + dest=f"{dprefix}log_level", + default=None, + choices=["debug", "info", "warning", "error", "critical"], + ) parser.add_argument( f"--{prefix}policy", dest=f"{dprefix}policy", diff --git a/tests/observability/test_trace_utils.py b/tests/observability/test_trace_utils.py index 4bcba095b..9f2e771fe 100644 --- a/tests/observability/test_trace_utils.py +++ b/tests/observability/test_trace_utils.py @@ -55,6 +55,68 @@ def test_build_vllm_meta_trace_attrs_keeps_standard_and_pd_fields(): } +@pytest.mark.unit +def test_build_vllm_meta_trace_attrs_normalizes_request_metrics(): + attrs = build_vllm_meta_trace_attrs( + { + "request_metrics": { + "queue_time_ms": 100, + "time_to_first_token_ms": 200, + "generation_time_ms": 300, + "tokens_per_second": 20, + "remote_kv_wait_time_ms": 50, + } + } + ) + trace_children = attrs.pop(TRACE_CHILDREN_KEY) + + assert attrs == { + "queue_time": pytest.approx(0.1), + "e2e_latency": pytest.approx(0.6), + "decode_throughput": pytest.approx(20), + } + assert trace_children == [ + { + "type": "span", + "name": "vllm_pd_decode", + "start_offset": 0.0, + "end_offset": pytest.approx(0.05), + "attrs": {"phase": "decode", "duration_s": pytest.approx(0.05)}, + "children": [ + { + "type": "span", + "name": "vllm_pd_decode_transfer", + "start_offset": 0.0, + "end_offset": pytest.approx(0.05), + "attrs": {"pd_decode_transfer_duration": pytest.approx(0.05)}, + } + ], + } + ] + + +@pytest.mark.unit +def test_build_vllm_meta_trace_attrs_reads_tito_response(): + attrs = build_vllm_meta_trace_attrs( + { + "request_id": "request-7", + "choices": [{"finish_reason": "length"}], + "usage": { + "prompt_tokens": 12, + "completion_tokens": 7, + "prompt_tokens_details": {"cached_tokens": 3}, + }, + } + ) + assert attrs == { + "vllm_request_id": "request-7", + "finish_reason": "length", + "prompt_tokens": 12, + "completion_tokens": 7, + "cached_tokens": 3, + } + + @pytest.mark.unit def test_trace_timeline_viewer_omits_virtual_pd_lanes_without_pd_attrs(tmp_path: Path): viewer = _load_trace_timeline_viewer_module() diff --git a/tests/plugin_contracts/test_plugin_generate_contracts.py b/tests/plugin_contracts/test_plugin_generate_contracts.py index 60a2afb1e..e6c673ec5 100644 --- a/tests/plugin_contracts/test_plugin_generate_contracts.py +++ b/tests/plugin_contracts/test_plugin_generate_contracts.py @@ -61,6 +61,8 @@ def __init__(self, args) -> None: self.pendings = set() self.remaining_batch_size = 0 self.aborted = False + self.cancellable_tasks = set() + self.active_server_generations = 0 self.group_sampling_seeds = None if getattr(args, "vllm_enable_deterministic_inference", False): self.group_sampling_seeds = [args.rollout_seed + i for i in range(args.n_samples_per_prompt)] diff --git a/tests/test_agent/test_adapters.py b/tests/test_agent/test_adapters.py index e7c6ec601..d59f3f636 100644 --- a/tests/test_agent/test_adapters.py +++ b/tests/test_agent/test_adapters.py @@ -170,7 +170,7 @@ async def run_case(): async with FakeVLLMServer([[(-0.1, 101), (-0.2, 102)]]) as vllm: tok = FakeTokenizer(outputs={(101, 102): "done now"}) adapter = anthropic.AnthropicAdapter(tokenizer=tok, vllm_url=vllm.url) - adapter.open_session("sid-a") + adapter.open_session("sid-a", sampling_defaults={"min_new_tokens": 2, "repetition_penalty": 1.2}) client = TestClient(TestServer(adapter.app)) await client.start_server() try: @@ -189,6 +189,8 @@ async def run_case(): assert data["content"] == [{"type": "text", "text": "done now"}] # adapter posted the rendered prompt ids and capped max_tokens at the request cap. assert vllm.requests[0]["sampling_params"]["max_tokens"] == 7 + assert vllm.requests[0]["sampling_params"]["min_tokens"] == 2 + assert vllm.requests[0]["sampling_params"]["repetition_penalty"] == 1.2 assert vllm.routing_keys == ["sid-a"] # one trained turn: the two response ids carry loss=1 + real logprobs. assert len(samples) == 1 @@ -201,6 +203,60 @@ async def run_case(): asyncio.run(run_case()) +@pytest.mark.parametrize("protocol", ["anthropic", "openai"]) +@pytest.mark.parametrize("enabled", [False, True]) +def test_session_sampling_defaults_reach_vllm(protocol, enabled): + defaults = { + "max_new_tokens": 20, + "min_new_tokens": 2, + "repetition_penalty": 1.2, + "seed": 37 if enabled else 0, + "min_p": 0.1 if enabled else 0.0, + "presence_penalty": 0.5 if enabled else 0.0, + "frequency_penalty": -0.5 if enabled else 0.0, + "ignore_eos": enabled, + "spaces_between_special_tokens": enabled, + "no_stop_trim": enabled, + "logit_bias": {"101": 0.5} if enabled else {}, + "stop": ["END"], + "stop_token_ids": [99], + "skip_special_tokens": enabled, + "temperature": 0.8, + "top_p": 0.9, + "top_k": -1, + } + + async def run_case(): + async with FakeVLLMServer([[(-0.1, 101)]]) as vllm: + adapter_cls = anthropic.AnthropicAdapter if protocol == "anthropic" else openai.OpenAIAdapter + adapter = adapter_cls(tokenizer=FakeTokenizer(outputs={(101,): "done"}), vllm_url=vllm.url) + adapter.open_session("sampling", sampling_defaults=defaults) + client = TestClient(TestServer(adapter.app)) + await client.start_server() + try: + response = await client.post( + "/v1/messages" if protocol == "anthropic" else "/v1/chat/completions", + headers={"Authorization": "Bearer sampling"}, + json={"model": "m", "max_tokens": 7, "messages": [{"role": "user", "content": "hi"}]}, + ) + await response.json() + assert response.status == 200 + finally: + await client.close() + await _drain(adapter, "sampling") + + expected = dict(defaults) + expected.pop("max_new_tokens") + expected["max_tokens"] = 7 + expected["min_tokens"] = expected.pop("min_new_tokens") + expected["include_stop_str_in_output"] = expected.pop("no_stop_trim") + expected["logprobs"] = 1 + assert vllm.requests[0]["sampling_params"] == expected + assert defaults["max_new_tokens"] == 20 + + asyncio.run(run_case()) + + def test_openai_chat_completions_nonstream_records_token_segments(): async def run_case(): async with FakeVLLMServer([[(-0.3, 201)]]) as vllm: diff --git a/tests/test_agent/test_agent_rollout_cpu.py b/tests/test_agent/test_agent_rollout_cpu.py index 66bda4c99..414a29cf1 100644 --- a/tests/test_agent/test_agent_rollout_cpu.py +++ b/tests/test_agent/test_agent_rollout_cpu.py @@ -23,11 +23,13 @@ from __future__ import annotations +import ast import asyncio import contextlib import dataclasses import sys import types +from copy import deepcopy from pathlib import Path from types import SimpleNamespace @@ -75,6 +77,49 @@ async def _timeout_shim(_delay): NUM_GPUS = 0 + +@pytest.mark.parametrize( + "exit_stage, expected_count", [("solver", 2), ("rewriter", 4), ("selector", 4), ("complete", 5)] +) +def test_multi_agent_emits_shared_rollout_id(exit_stage, expected_count): + source_path = REPO_ROOT / "examples/multi_agent/agent_system.py" + source = ast.parse(source_path.read_text()) + function = next( + node for node in source.body if isinstance(node, ast.AsyncFunctionDef) and node.name == "run_agent_system" + ) + + async def solver_worker(args, problem_statement, worker_id): + args.results_dict["solver"].append(Sample(index=worker_id, prompt=problem_statement)) + return None if exit_stage == "solver" else "solution" + + async def rewrite_worker(args, previous_solutions, problem_statement, worker_id): + args.results_dict["rewriter"].append(Sample(index=worker_id, prompt=problem_statement)) + return None if exit_stage == "rewriter" else "rewritten" + + async def batched_async_rm(args, samples): + return [1.0] * len(samples) + + class SelectorAgent: + async def select(self, args, problem_statement, solutions): + if exit_stage != "selector": + args.results_dict["selector"].append(Sample(index=0, prompt=problem_statement)) + return None + + namespace = { + "asyncio": asyncio, + "deepcopy": deepcopy, + "solver_worker": solver_worker, + "rewrite_worker": rewrite_worker, + "batched_async_rm": batched_async_rm, + "SelectorAgent": SelectorAgent, + } + exec(compile(ast.Module(body=[function], type_ignores=[]), str(source_path), "exec"), namespace) + args = SimpleNamespace(num_parallel=2, incorrect_reward_weight=0.5, correct_reward_weight=1.0) + samples = asyncio.run(namespace["run_agent_system"](args, Sample(index=37, prompt="problem"))) + assert len(samples) == expected_count + assert all(sample.rollout_id == 37 for sample in samples) + + _REAL_SLEEP = asyncio.sleep diff --git a/tests/test_docs_consistency.py b/tests/test_docs_consistency.py index 217a29884..103fe16ba 100644 --- a/tests/test_docs_consistency.py +++ b/tests/test_docs_consistency.py @@ -86,5 +86,38 @@ def test_customization_anchor_links_exist(language): assert not missing, f"Local anchors without matching headings: {missing}" +@pytest.mark.parametrize("language", ["en", "zh"]) +def test_agent_guide_is_in_get_started_toctree(language): + text = (ROOT / "docs" / language / "index.rst").read_text(encoding="utf-8") + assert re.search(r"^ get_started/agent\.md$", text, re.MULTILINE) + + +def test_ci_skill_references_current_buildkite_sources(): + text = (ROOT / ".claude/skills/add-tests-and-ci/SKILL.md").read_text(encoding="utf-8") + for source in (".buildkite/pipeline.yml", ".buildkite/gpu_suites.py"): + assert source in text + assert (ROOT / source).is_file() + assert ".github/workflows/pr-test" not in text + assert "generate_github_workflows.py" not in text + + +@pytest.mark.parametrize("language, filename", [("en", "README.md"), ("zh", "README_zh.md")]) +def test_readme_links_deployment_and_correctness_guides(language, filename): + text = (ROOT / filename).read_text(encoding="utf-8") + for guide in ( + "advanced/vllm-config.md", + "advanced/pd-disaggregation.md", + "advanced/delta-weight-sync.md", + "advanced/external-rollout-engines.md", + "advanced/reproducibility.md", + "advanced/fault-tolerance.md", + "developer_guide/ci.md", + "developer_guide/debug.md", + "developer_guide/trace.md", + "developer_guide/profiling.md", + ): + assert f"docs/{language}/{guide}" in text + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/tests/test_external_vllm_engines.py b/tests/test_external_vllm_engines.py index 5e417ef9c..785e99306 100644 --- a/tests/test_external_vllm_engines.py +++ b/tests/test_external_vllm_engines.py @@ -74,6 +74,32 @@ def fake_get(url, timeout): } +@pytest.mark.parametrize("nested", [False, True]) +@pytest.mark.parametrize( + "ec_connector, ec_role, kv_role, expected", + [ + ("ECExampleConnector", "ec_producer", None, "encoder"), + ("ECExampleConnector", "ec_consumer", "kv_producer", "prefill"), + ("ECExampleConnector", "ec_consumer", "kv_consumer", "decode"), + ("ECExampleConnector", "ec_both", None, "regular"), + (None, "ec_producer", None, "regular"), + ], +) +def test_discover_external_native_ec_roles(monkeypatch, nested, ec_connector, ec_role, kv_role, expected): + config = { + "ec_transfer_config": {"ec_connector": ec_connector, "ec_role": ec_role}, + "kv_transfer_config": {"kv_role": kv_role}, + "parallel_config": {"tensor_parallel_size": 1}, + } + payload = {"vllm_config": config} if nested else config + monkeypatch.setattr("vime.backends.vllm_utils.external.requests.get", lambda *args, **kwargs: _Response(payload)) + + info = discover_external_engines(["host1:10090"])[0] + + assert info.worker_type == expected + assert info.server_info["ec_transfer_config"] == config["ec_transfer_config"] + + def test_start_external_rollout_servers_exposes_parallel_configs(monkeypatch): class FakeActor: init = Namespace(remote=lambda **kwargs: kwargs) diff --git a/tests/test_rollout_metrics.py b/tests/test_rollout_metrics.py index 273802d67..fb40def3f 100644 --- a/tests/test_rollout_metrics.py +++ b/tests/test_rollout_metrics.py @@ -5,7 +5,7 @@ import pytest import torch -from vime.observability.rollout_metrics import _compute_top_p_kept_vocab_metrics +from vime.observability.rollout_metrics import _compute_spec_metrics, _compute_top_p_kept_vocab_metrics from vime.utils.misc import decode_int32_meta_array from vime.utils.types import Sample @@ -13,7 +13,20 @@ def _make_args(): - return Namespace(vllm_speculative_algorithm=False, num_layers=2, moe_router_topk=2) + return Namespace(vllm_speculative_config=None, num_layers=2, moe_router_topk=2) + + +@pytest.mark.unit +def test_spec_metrics_use_vllm_speculative_config(): + sample = Sample() + sample.spec_info.spec_accept_token_num = 6 + sample.spec_info.spec_draft_token_num = 8 + sample.spec_info.spec_verify_ct = 2 + args = Namespace(vllm_speculative_config={"method": "mtp", "num_speculative_tokens": 4}) + metrics = _compute_spec_metrics(args, [sample]) + assert metrics["spec_accept_rate"] == sample.spec_info.spec_accept_rate + assert metrics["spec_accept_length"] == sample.spec_info.spec_accept_length + assert _compute_spec_metrics(_make_args(), [sample]) == {} @pytest.mark.unit diff --git a/tests/test_sample.py b/tests/test_sample.py index bc9d3ff6d..32d7d647d 100644 --- a/tests/test_sample.py +++ b/tests/test_sample.py @@ -178,7 +178,7 @@ def test_round_trip_through_default_constructed_sample(): def _make_args(speculative: bool = False) -> argparse.Namespace: - """``append_response_tokens`` only consults ``args.vllm_speculative_algorithm`` + """``append_response_tokens`` only consults ``args.vllm_speculative_config`` — minimal stub is enough.""" return argparse.Namespace(vllm_speculative_config=speculative) @@ -193,7 +193,7 @@ def _make_args(speculative: bool = False) -> argparse.Namespace: ], ) def test_status_mapping_for_each_finish_reason(finish_reason, expected_status): - """The match statement at types.py:176-182 is the one place the engine's + """The match statement at types.py:176-182 is the one place vllm's finish_reason ever gets translated. Each branch must hit the right enum; a typo in the enum name would crash later in unrelated places.""" sample = Sample() @@ -266,7 +266,7 @@ def test_prefix_cache_info_is_accumulated_across_calls(): def test_spec_info_only_updated_when_speculative_enabled(): """``spec_info.add`` is gated on ``args.vllm_speculative_config`` (types.py:166-168). Without the flag, spec stats stay at zero even - if the engine sends them.""" + if vllm sends them.""" meta_info = { "finish_reason": {"type": "stop"}, "spec_accept_token_num": 7, diff --git a/tests/test_vllm_rollout.py b/tests/test_vllm_rollout.py index 7b0bd392e..6da675fcf 100644 --- a/tests/test_vllm_rollout.py +++ b/tests/test_vllm_rollout.py @@ -67,6 +67,8 @@ def __init__(self, args: Namespace) -> None: self.aborted = False self.remaining_batch_size = 0 self.pendings: set = set() + self.cancellable_tasks: set = set() + self.active_server_generations = 0 self.dp_counts = [0] self.dp_rank = 0 self.group_sampling_seeds = None @@ -84,6 +86,8 @@ def dp_rank_context(self): def reset(self) -> None: self.remaining_batch_size = 0 self.pendings = set() + self.cancellable_tasks = set() + self.active_server_generations = 0 self.aborted = False @@ -133,6 +137,7 @@ def _generate_response( weight_version: str | None = None, request_spec_decode_stats: dict[str, int] | None = None, sampling_mask: list[list[int]] | None = None, + request_metrics: dict[str, float] | None = None, ) -> dict: tids = token_ids or [50, 51] response = { @@ -151,6 +156,8 @@ def _generate_response( response["request_spec_decode_stats"] = request_spec_decode_stats if sampling_mask is not None: response["choices"][0]["sampling_mask"] = sampling_mask + if request_metrics is not None: + response["request_metrics"] = request_metrics return response @@ -220,6 +227,8 @@ def test_build_inference_sampling_params_maps_rollout_fields(): sp = mod._build_inference_sampling_params( { "max_new_tokens": 16, + "min_new_tokens": 4, + "repetition_penalty": 1.2, "temperature": 0.7, "top_p": 0.9, "top_k": 40, @@ -230,6 +239,8 @@ def test_build_inference_sampling_params_maps_rollout_fields(): } ) assert sp["max_tokens"] == 16 + assert sp["min_tokens"] == 4 + assert sp["repetition_penalty"] == 1.2 assert sp["temperature"] == 0.7 assert sp["top_p"] == 0.9 assert sp["top_k"] == 40 @@ -351,6 +362,13 @@ def test_generate_text_path_updates_sample(patch_generate_state, monkeypatch): "num_draft_tokens": 8, "num_spec_steps": 2, }, + request_metrics={ + "queue_time_ms": 100, + "time_to_first_token_ms": 200, + "generation_time_ms": 300, + "tokens_per_second": 20, + "remote_kv_wait_time_ms": 50, + }, ) ) monkeypatch.setattr(mod, "post", post_mock) @@ -360,7 +378,7 @@ def test_generate_text_path_updates_sample(patch_generate_state, monkeypatch): mod.generate( _rollout_args(vllm_speculative_config={"method": "mtp"}), sample, - _default_sampling_params(max_new_tokens=8), + _default_sampling_params(max_new_tokens=8, min_new_tokens=2, repetition_penalty=1.2), ) ) @@ -374,9 +392,29 @@ def test_generate_text_path_updates_sample(patch_generate_state, monkeypatch): assert result.spec_info.spec_draft_token_num == 8 assert result.spec_info.spec_verify_ct == 2 assert result.status == Sample.Status.COMPLETED + generate_span = next( + event for event in result.trace["events"] if event["type"] == "span_end" and event["name"] == "vllm_generate" + ) + assert generate_span["attrs"] == { + "prompt_tokens": 3, + "completion_tokens": 2, + "cached_tokens": 0, + "finish_reason": "stop", + "queue_time": pytest.approx(0.1), + "e2e_latency": pytest.approx(0.6), + "decode_throughput": pytest.approx(20), + } + decode_transfer_span = next( + event + for event in result.trace["events"] + if event["type"] == "span_end" and event["name"] == "vllm_pd_decode_transfer" + ) + assert decode_transfer_span["attrs"] == {"pd_decode_transfer_duration": pytest.approx(0.05)} body = post_mock.await_args_list[0].args[1] assert body["token_ids"] == [97, 98, 99] assert body["sampling_params"]["max_tokens"] == 8 + assert body["sampling_params"]["min_tokens"] == 2 + assert body["sampling_params"]["repetition_penalty"] == 1.2 @pytest.mark.unit @@ -396,6 +434,7 @@ def raise_for_status(self): async def aiter_lines(self): chunks = [ { + "request_id": "stream-7", "weight_version": "step-7", "request_spec_decode_stats": { "num_accepted_tokens": 6, @@ -423,6 +462,12 @@ async def aiter_lines(self): ], "usage": {"prompt_tokens": 3, "completion_tokens": 2}, }, + { + "request_id": "stream-7", + "choices": [], + "usage": {"prompt_tokens": 3, "completion_tokens": 2}, + "request_metrics": {"queue_time_ms": 100}, + }, ] for chunk in chunks: yield f"data: {json.dumps(chunk)}" @@ -453,6 +498,66 @@ def stream(self, *args, **kwargs): assert result.spec_info.spec_verify_ct == 2 assert result.status == Sample.Status.COMPLETED + event = next( + event + for event in result.trace["events"] + if event["type"] == "span_end" and event["name"] == "vllm_inference_generate_stream" + ) + assert event["attrs"]["vllm_request_id"] == "stream-7" + assert event["attrs"]["finish_reason"] == "stop" + assert event["attrs"]["completion_tokens"] == 2 + assert event["attrs"]["queue_time"] == pytest.approx(0.1) + + +@pytest.mark.unit +def test_generate_streaming_stops_at_partial_token_budget(patch_generate_state, monkeypatch): + from vime.rollout import vllm_streaming_rollout as streaming + + monkeypatch.setattr(streaming, "GenerateState", _PatchedGenerateState) + monkeypatch.setattr(streaming.http_utils, "_http_client", None) + sample = Sample(index=0, prompt="abc", tokens=[97, 98, 99, 0], response_length=1) + sample.status = Sample.Status.ABORTED + result = asyncio.run( + streaming.generate_streaming(_rollout_args(), sample, _default_sampling_params(max_new_tokens=1)) + ) + assert result.status == Sample.Status.TRUNCATED + assert result.tokens == [97, 98, 99, 0] + + +@pytest.mark.unit +def test_generate_streaming_rejects_unexpected_eof(patch_generate_state, monkeypatch): + from vime.rollout import vllm_streaming_rollout as streaming + + class FakeStreamResponse: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, traceback): + return False + + def raise_for_status(self): + return None + + async def aiter_lines(self): + yield 'data: {"choices": [{"token_ids": [50], "finish_reason": null}]}' + yield "data: [DONE]" + + class FakeClient: + def stream(self, *args, **kwargs): + return FakeStreamResponse() + + monkeypatch.setattr(streaming, "GenerateState", _PatchedGenerateState) + monkeypatch.setattr(streaming.http_utils, "_http_client", FakeClient()) + + with pytest.raises(RuntimeError, match="without a terminal finish_reason"): + asyncio.run( + streaming.generate_streaming( + _rollout_args(), + Sample(index=0, prompt="abc"), + _default_sampling_params(max_new_tokens=8), + ) + ) + @pytest.mark.unit def test_generate_consistent_hash_header(patch_generate_state, monkeypatch): @@ -809,6 +914,7 @@ def test_abort_deletes_inflight_without_pause_resume(patch_generate_state, monke from vime.backends.vllm_utils import server_control state = _PatchedGenerateState(_rollout_args()) + state.active_server_generations = 1 monkeypatch.setattr(mod, "GenerateState", lambda args: state) aborted = asyncio.Event() @@ -854,6 +960,7 @@ def test_abort_collects_partial_samples_when_partial_rollout(patch_generate_stat args = _rollout_args(partial_rollout=True) state = _PatchedGenerateState(args) + state.active_server_generations = 1 monkeypatch.setattr(mod, "GenerateState", lambda a: state) aborted = asyncio.Event() @@ -871,9 +978,11 @@ async def fake_post(url, payload, max_retries=60, headers=None): sample = Sample(index=0, prompt="p") sample.response = "partial" + sample.response_length = 1 async def pending_group(): await aborted.wait() + sample.status = Sample.Status.ABORTED return [sample] async def run_abort(): @@ -886,5 +995,247 @@ async def run_abort(): assert sample.metadata["start_rollout_id"] == 7 +@pytest.mark.unit +def test_abort_cancels_request_without_server_abort(patch_generate_state, monkeypatch): + args = _rollout_args(partial_rollout=True) + state = _PatchedGenerateState(args) + monkeypatch.setattr(mod, "GenerateState", lambda a: state) + get_mock = AsyncMock() + monkeypatch.setattr(mod, "get", get_mock) + + sample = Sample(index=0, prompt="p") + sample.response = "partial" + sample.response_length = 1 + + async def request(): + await asyncio.Future() + + async def run(): + generate_task = asyncio.create_task(mod._run_request_abortable_generate(state, sample, request())) + await asyncio.sleep(0) + + async def group(): + return [await generate_task] + + state.pendings = {asyncio.create_task(group())} + return await asyncio.wait_for(mod.abort(args, rollout_id=9), timeout=5.0) + + assert asyncio.run(run()) == [[sample]] + assert sample.status == Sample.Status.ABORTED + assert sample.metadata["start_rollout_id"] == 9 + get_mock.assert_not_awaited() + + +@pytest.mark.unit +@pytest.mark.parametrize("unrelated_cancel", [False, True]) +def test_stream_cancellation_closes_http_and_preserves_prefix(patch_generate_state, monkeypatch, unrelated_cancel): + from vime.rollout import vllm_streaming_rollout as streaming + + args = _rollout_args() + state = _PatchedGenerateState(args) + monkeypatch.setattr(mod, "GenerateState", lambda args: state) + monkeypatch.setattr(streaming, "GenerateState", lambda args: state) + monkeypatch.setattr(mod, "load_function", lambda path: streaming.generate_streaming) + prefix_seen = asyncio.Event() + request_closed = asyncio.Event() + server_started = asyncio.Event() + server_released = asyncio.Event() + + class FakeResponse: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, traceback): + request_closed.set() + return False + + def raise_for_status(self): + return None + + async def aiter_lines(self): + chunk = _generate_response([120], sampling_mask=[[12, 120]]) + chunk["choices"][0]["finish_reason"] = None + yield f"data: {json.dumps(chunk)}" + prefix_seen.set() + await asyncio.Future() + + class FakeClient: + def stream(self, *args, **kwargs): + return FakeResponse() + + async def server_generate(args, sample, sampling_params): + server_started.set() + await server_released.wait() + sample.status = Sample.Status.ABORTED + return sample + + async def abort_servers(urls): + assert urls == ["http://worker:9000"] + server_released.set() + + monkeypatch.setattr(streaming.http_utils, "_http_client", FakeClient()) + monkeypatch.setattr(mod, "generate", server_generate) + get_mock = AsyncMock(return_value={"workers": [{"url": "http://worker:9000"}]}) + monkeypatch.setattr(mod, "get", get_mock) + monkeypatch.setattr(mod, "abort_inflight_requests", abort_servers) + sample = Sample(prompt="abc", generate_function_path="streaming") + + async def exercise(): + request_task = asyncio.create_task(mod.generate_and_rm(args, sample, _default_sampling_params())) + await prefix_seen.wait() + if unrelated_cancel: + state.aborted = True + request_task.cancel() + with pytest.raises(asyncio.CancelledError): + await request_task + get_mock.assert_not_awaited() + else: + server_task = asyncio.create_task(mod.generate_and_rm(args, Sample(prompt="abc"), {})) + await server_started.wait() + await mod.abort(args, rollout_id=7) + result, server_result = await asyncio.gather(request_task, server_task) + assert result is sample + assert result.status == server_result.status == Sample.Status.ABORTED + get_mock.assert_awaited_once() + + async def run(): + await asyncio.wait_for(exercise(), timeout=5) + + asyncio.run(run()) + assert request_closed.is_set() + assert sample.tokens == [97, 98, 99, 120] + assert sample.response == "x" + assert sample.response_length == 1 + assert sample.rollout_log_probs == [-0.1] + assert sample.rollout_top_p_token_ids.tolist() == [12, 120] + assert sample.rollout_top_p_token_offsets.tolist() == [0, 2] + assert not state.cancellable_tasks + assert state.active_server_generations == 0 + + +@pytest.mark.unit +def test_partial_abort_resumes_only_aborted_siblings(patch_generate_state, monkeypatch): + from vime.rollout import vllm_streaming_rollout as streaming + + args = _rollout_args( + partial_rollout=True, + mask_offpolicy_in_partial_rollout=True, + custom_generate_function_path="streaming", + ) + state = _PatchedGenerateState(args) + monkeypatch.setattr(mod, "GenerateState", lambda args: state) + monkeypatch.setattr(streaming, "GenerateState", lambda args: state) + monkeypatch.setattr(mod, "load_function", lambda path: streaming.generate_streaming) + partial = Sample( + prompt="abc", + tokens=[97, 98, 99, 120], + response="x", + response_length=1, + rollout_log_probs=[-0.1], + loss_mask=[1], + status=Sample.Status.ABORTED, + ) + terminal = Sample( + prompt="abc", + tokens=[97, 98, 99, 122], + response="z", + response_length=1, + rollout_log_probs=[-0.3], + reward=1.0, + status=Sample.Status.COMPLETED, + ) + empty = Sample(status=Sample.Status.ABORTED) + mixed_group = [terminal, partial] + payloads = [] + + class FakeResponse: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, traceback): + return False + + def raise_for_status(self): + return None + + async def aiter_lines(self): + yield f"data: {json.dumps(_generate_response([121]))}" + yield "data: [DONE]" + + class FakeClient: + def stream(self, method, url, *, json, headers): + payloads.append(json) + return FakeResponse() + + monkeypatch.setattr(streaming.http_utils, "_http_client", FakeClient()) + reward = AsyncMock(return_value=2.0) + monkeypatch.setattr(mod, "async_rm", reward) + + async def exercise(): + state.pendings = { + asyncio.create_task(asyncio.sleep(0, result=group)) for group in (mixed_group, [terminal], [empty]) + } + buffered = await mod.abort(args, rollout_id=7) + assert buffered == [mixed_group] + assert not state.pendings + state.aborted = False + return await mod.generate_and_rm_group(args, buffered[0], _default_sampling_params(max_new_tokens=8)) + + async def run(): + return await asyncio.wait_for(exercise(), timeout=5) + + assert asyncio.run(run()) == mixed_group + assert len(payloads) == 1 + assert payloads[0]["token_ids"] == [97, 98, 99, 120] + assert payloads[0]["sampling_params"]["max_tokens"] == 7 + assert partial.tokens == [97, 98, 99, 120, 121] + assert partial.response == "xy" + assert partial.response_length == 2 + assert partial.rollout_log_probs == [-0.1, -0.1] + assert partial.loss_mask == [0, 1] + assert partial.status == terminal.status == Sample.Status.COMPLETED + assert partial.reward == 2.0 + assert terminal.tokens == [97, 98, 99, 122] + assert terminal.response == "z" + assert terminal.reward == 1.0 + assert partial.metadata["start_rollout_id"] == terminal.metadata["start_rollout_id"] == 7 + reward.assert_awaited_once_with(args, partial) + + +@pytest.mark.unit +def test_multi_agent_generate_response_preserves_request_metadata(monkeypatch): + from examples.multi_agent import agent_system + + class CallableTokenizer(_FakeTokenizer): + def __call__(self, prompt, add_special_tokens=False): + return {"input_ids": self.encode(prompt, add_special_tokens=add_special_tokens)} + + args = _rollout_args( + tokenizer=CallableTokenizer(), + sampling_params=_default_sampling_params(), + rollout_max_context_len=32, + sample=Sample(), + results_dict={"solver": []}, + vllm_speculative_config={"method": "mtp"}, + ) + post_mock = AsyncMock( + return_value=_generate_response( + weight_version="step-7", + request_spec_decode_stats={"num_accepted_tokens": 6, "num_draft_tokens": 8, "num_verify_steps": 2}, + ) + ) + monkeypatch.setattr(agent_system, "post", post_mock) + + asyncio.run(agent_system.generate_response(args, "abc", "solver")) + + post_mock.assert_awaited_once() + assert len(args.results_dict["solver"]) == 1 + sample = args.results_dict["solver"][0] + assert sample.weight_versions == ["step-7"] + assert sample.spec_info.spec_accept_token_num == 6 + assert sample.spec_info.spec_draft_token_num == 8 + assert sample.spec_info.spec_verify_ct == 2 + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/tests/utils/test_megatron_role_config.py b/tests/utils/test_megatron_role_config.py index fe1c847db..21f2891df 100644 --- a/tests/utils/test_megatron_role_config.py +++ b/tests/utils/test_megatron_role_config.py @@ -16,6 +16,8 @@ _unit_stubs.install_rollout_optional_stubs() +NUM_GPUS = 0 + def _write_yaml(data: dict) -> str: handle = tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) diff --git a/tests/utils/test_vllm_arguments.py b/tests/utils/test_vllm_arguments.py index 743f73831..0d22eef01 100644 --- a/tests/utils/test_vllm_arguments.py +++ b/tests/utils/test_vllm_arguments.py @@ -3,6 +3,9 @@ from __future__ import annotations import argparse +import ast +import logging +import random import sys from pathlib import Path from types import SimpleNamespace @@ -161,6 +164,38 @@ def test_add_vllm_router_arguments_defaults_to_cache_aware(args_mod): assert parsed.router_policy == "cache_aware" +@pytest.mark.unit +@pytest.mark.parametrize("flags, expected", [([], "warning"), (["--router-log-level", "debug"], "debug")]) +def test_router_log_level_survives_launch(args_mod, monkeypatch, flags, expected): + from vllm_router.router_args import RouterArgs + + parser = argparse.ArgumentParser(add_help=False) + args_mod.add_vllm_router_arguments(parser) + args = parser.parse_args(flags) + assert args.router_log_level == expected + router_args = SimpleNamespace(log_level=args.router_log_level) + monkeypatch.setattr(RouterArgs, "from_cli_args", lambda *args, **kwargs: router_args) + launches = [] + + def make_process(*, target, args): + launches.append(args[0]) + return SimpleNamespace(start=lambda: None, is_alive=lambda: True) + + source_path = Path(args_mod.__file__).with_name("deployment.py") + source = ast.parse(source_path.read_text()) + function = next(node for node in source.body if isinstance(node, ast.FunctionDef) and node.name == "_start_router") + namespace = { + "random": random, + "logger": logging.getLogger(__name__), + "find_available_port": lambda port: port, + "time": SimpleNamespace(sleep=lambda seconds: None), + "multiprocessing": SimpleNamespace(Process=make_process), + } + exec(compile(ast.Module(body=[function], type_ignores=[]), str(source_path), "exec"), namespace) + namespace["_start_router"](args, bind=("127.0.0.1", 30000)) + assert launches[0].log_level == expected + + @pytest.mark.unit def test_add_vllm_arguments_overrides_router_balance_threshold_defaults(args_mod, monkeypatch): _patch_device_config(monkeypatch) diff --git a/tests/utils/test_vllm_engine.py b/tests/utils/test_vllm_engine.py index 5f599b397..e33cc3518 100644 --- a/tests/utils/test_vllm_engine.py +++ b/tests/utils/test_vllm_engine.py @@ -91,6 +91,63 @@ def json(self) -> dict: return self._json_data +@pytest.mark.unit +def test_flush_cache_retries_unsuccessful_reset(vllm_engine, monkeypatch, caplog): + responses = iter( + [ + _MockResponse(json_data={"success": False}, text='{"success": false}'), + _MockResponse(json_data={"success": True}), + ] + ) + calls = [] + sleeps = [] + + def fake_post(url, *, params): + calls.append((url, params)) + return next(responses) + + monkeypatch.setattr(mod.requests, "post", fake_post) + monkeypatch.setattr(mod.time, "sleep", sleeps.append) + with caplog.at_level("INFO", logger=mod.__name__): + vllm_engine.flush_cache() + + assert calls == [("http://127.0.0.1:8765/reset_prefix_cache", {"reset_running_requests": False})] * 2 + assert sleeps == [1] + assert "Error flushing cache: HTTP 200" in caplog.text + assert '{"success": false}' in caplog.text + + +@pytest.mark.unit +def test_flush_cache_retries_http_error(vllm_engine, monkeypatch, caplog): + responses = iter( + [ + _MockResponse(status_code=503, text="busy"), + _MockResponse(json_data={"success": True}), + ] + ) + sleeps = [] + monkeypatch.setattr(mod.requests, "post", lambda *args, **kwargs: next(responses)) + monkeypatch.setattr(mod.time, "sleep", sleeps.append) + with caplog.at_level("INFO", logger=mod.__name__): + vllm_engine.flush_cache() + + assert sleeps == [1] + assert "Error flushing cache: HTTP 503 'busy'" in caplog.text + assert next(responses, None) is None + + +@pytest.mark.unit +def test_flush_cache_times_out_after_unsuccessful_resets(vllm_engine, monkeypatch): + sleeps = [] + monkeypatch.setattr(mod.requests, "post", lambda *args, **kwargs: _MockResponse(json_data={"success": False})) + monkeypatch.setattr(mod.time, "sleep", sleeps.append) + + with pytest.raises(TimeoutError, match="Timeout while flushing cache"): + vllm_engine.flush_cache() + + assert sleeps == [1] * 60 + + @pytest.mark.unit def test_normalize_vllm_wake_tags_drops_unsupported(): assert mod._normalize_vllm_wake_tags(["weights", "cuda_graph", "kv_cache"]) == ["weights", "kv_cache"] @@ -113,6 +170,7 @@ def test_launch_config_single_node(vllm_args): assert sa["_pp_size"] == 1 assert sa["_pcp_size"] == 1 assert sa["_dp_size"] == 1 + assert sa["enable_per_request_metrics"] is True @pytest.mark.unit @@ -784,6 +842,7 @@ def test_pull_weights_posts_collective_rpc(vllm_engine, monkeypatch): vllm_engine.args.update_weight_local_checkpoint_dir = "/local/checkpoint" vllm_engine.args.update_weight_disk_dir = "/shared/checkpoints" vllm_engine.args.custom_update_weight_pre_read_path = "hooks.refresh" + vllm_engine._weight_version = "old" seen = [] def fake_post(url, *, json=None): @@ -806,16 +865,13 @@ def fake_post(url, *, json=None): }, }, ), - ( - "http://127.0.0.1:8765/update_weight_version", - {"new_version": "8"}, - ), ] - assert vllm_engine._weight_version == "8" + assert vllm_engine._weight_version == "old" @pytest.mark.unit -def test_pull_weights_does_not_advance_version_when_pull_fails(vllm_engine, monkeypatch): +@pytest.mark.parametrize("operation", ["pull", "reload"]) +def test_disk_update_does_not_advance_version_on_failure(vllm_engine, monkeypatch, operation): vllm_engine.args.update_weight_local_checkpoint_dir = "/local/checkpoint" vllm_engine.args.update_weight_disk_dir = "/shared/checkpoints" vllm_engine.args.custom_update_weight_pre_read_path = None @@ -828,7 +884,10 @@ def fake_post(url, *, json=None): monkeypatch.setattr(mod.requests, "post", fake_post) with pytest.raises(requests.exceptions.HTTPError): - vllm_engine.pull_weights(8) + if operation == "pull": + vllm_engine.pull_weights(8) + else: + vllm_engine.update_weights_from_disk("/local/checkpoint", weight_version="8") assert vllm_engine._weight_version == "old" diff --git a/tools/convert_hf_to_torch_dist.py b/tools/convert_hf_to_torch_dist.py index 798e79ef6..5d189caf2 100644 --- a/tools/convert_hf_to_torch_dist.py +++ b/tools/convert_hf_to_torch_dist.py @@ -1,4 +1,3 @@ -import argparse import gc import os import shutil @@ -29,10 +28,6 @@ def add_convertion_args(parser): help="Path to a custom model provider function.", ) parser.add_argument("--allgather-cp", action="store_true", default=False) - try: - parser.add_argument("--use-gated-attention", action="store_true", default=False) - except argparse.ArgumentError: - pass try: parser.add_argument("--padded-vocab-size", type=int, default=None) except Exception: diff --git a/vime/agent/adapters/common.py b/vime/agent/adapters/common.py index 15899d067..9b9b784f3 100644 --- a/vime/agent/adapters/common.py +++ b/vime/agent/adapters/common.py @@ -449,6 +449,21 @@ def _vllm_sampling_body(sp: dict) -> dict: } if "temperature" in sp: body["temperature"] = sp["temperature"] + if sp.get("min_new_tokens") is not None: + body["min_tokens"] = sp["min_new_tokens"] + if sp.get("repetition_penalty") is not None: + body["repetition_penalty"] = sp["repetition_penalty"] + for key in ( + "seed", + "min_p", + "presence_penalty", + "frequency_penalty", + "ignore_eos", + "logit_bias", + "spaces_between_special_tokens", + ): + if sp.get(key) is not None: + body[key] = sp[key] if "top_p" in sp: body["top_p"] = sp["top_p"] tk = sp.get("top_k") @@ -456,6 +471,8 @@ def _vllm_sampling_body(sp: dict) -> dict: body["top_k"] = tk if sp.get("stop"): body["stop"] = sp["stop"] + if sp.get("no_stop_trim") is not None: + body["include_stop_str_in_output"] = sp["no_stop_trim"] if sp.get("stop_token_ids"): body["stop_token_ids"] = sp["stop_token_ids"] if sp.get("skip_special_tokens") is not None: @@ -545,9 +562,8 @@ async def call_vllm_generate( fr = choice.get("finish_reason") finish = fr if isinstance(fr, str) and fr else "stop" except (asyncio.CancelledError, aiohttp.ClientError, asyncio.TimeoutError) as e: - # vLLM ``/inference/v1/generate`` has no per-request HTTP abort endpoint. - # Cancelling the in-flight task tears down the aiohttp request, which drops - # the streaming connection so vLLM stops generating. + # vLLM has no per-request abort endpoint. Closing this router request also + # closes its selected worker request, so vLLM cancels the engine request. logger.debug("[%s] sid=%s turn aborted: %s", adapter.log_prefix, session_id, type(e).__name__) if task is not None: task.cancel() diff --git a/vime/backends/megatron_utils/alignment/deepgemm_forward.py b/vime/backends/megatron_utils/alignment/deepgemm_forward.py index 34f89b68b..338bc3b24 100644 --- a/vime/backends/megatron_utils/alignment/deepgemm_forward.py +++ b/vime/backends/megatron_utils/alignment/deepgemm_forward.py @@ -176,7 +176,7 @@ def _norm_forward( if normalization == "RMSNorm" and os.environ.get("MEGATRON_USE_VLLM_FUSED_RESIDUAL_RMS", "0") == "1": if norm_bias is not None: raise RuntimeError("VLLM RMSNorm alignment does not support a norm bias") - from vllm.model_executor.layers.batch_invariant import rms_norm_batch_invariant + from vllm.model_executor.determinism.batch_invariant import rms_norm_batch_invariant weight = norm_weight if zero_centered_gamma: @@ -725,7 +725,7 @@ def enable_vllm_global_batch_invariant_ops() -> None: }: return - from vllm.model_executor.layers import batch_invariant + from vllm.model_executor.determinism import batch_invariant batch_invariant.enable_batch_invariant_mode() @@ -735,7 +735,7 @@ def _vllm_batch_invariant_rmsnorm( weight: torch.Tensor, eps: float, ) -> torch.Tensor: - from vllm.model_executor.layers.batch_invariant import rms_norm_batch_invariant + from vllm.model_executor.determinism.batch_invariant import rms_norm_batch_invariant return rms_norm_batch_invariant(value, weight, eps) diff --git a/vime/backends/vllm_utils/arguments.py b/vime/backends/vllm_utils/arguments.py index ccb1c7863..800fdf22f 100644 --- a/vime/backends/vllm_utils/arguments.py +++ b/vime/backends/vllm_utils/arguments.py @@ -30,6 +30,7 @@ def add_vllm_router_arguments(parser): help="Timeout for requests to the vllm router in seconds", ) RouterArgs.add_cli_args(parser, use_router_prefix=True, exclude_host_port=True) + parser.set_defaults(router_log_level="warning") return parser diff --git a/vime/backends/vllm_utils/deployment.py b/vime/backends/vllm_utils/deployment.py index 3fa1e4e87..e87b8dec6 100644 --- a/vime/backends/vllm_utils/deployment.py +++ b/vime/backends/vllm_utils/deployment.py @@ -44,7 +44,6 @@ def _start_router( router_args.host = router_ip router_args.port = router_port router_args.prometheus_port = find_available_port(random.randint(4000, 5000)) - router_args.log_level = "warning" router_args.request_timeout_secs = args.vllm_router_request_timeout_secs if has_pd_disaggregation: diff --git a/vime/backends/vllm_utils/external.py b/vime/backends/vllm_utils/external.py index bb78dd034..6d3940677 100644 --- a/vime/backends/vllm_utils/external.py +++ b/vime/backends/vllm_utils/external.py @@ -116,6 +116,9 @@ def find_config_value(config, name): weight_transfer_config = find_config_value(vllm_config, "weight_transfer_config") if weight_transfer_config is not None: normalized["weight_transfer_config"] = weight_transfer_config + ec_transfer_config = find_config_value(vllm_config, "ec_transfer_config") + if ec_transfer_config is not None: + normalized["ec_transfer_config"] = ec_transfer_config if isinstance(kv_transfer_config, dict): role = kv_transfer_config.get("kv_role") if role == "kv_producer": @@ -135,6 +138,13 @@ def find_config_value(config, name): def _infer_worker_type(server_info: dict) -> str: if server_info.get("encoder_only"): return "encoder" + ec_transfer_config = server_info.get("ec_transfer_config") + if ( + isinstance(ec_transfer_config, dict) + and ec_transfer_config.get("ec_connector") is not None + and ec_transfer_config.get("ec_role") == "ec_producer" + ): + return "encoder" kv_transfer_config = server_info.get("kv_transfer_config") if isinstance(kv_transfer_config, dict): role = kv_transfer_config.get("kv_role") diff --git a/vime/backends/vllm_utils/vllm_engine.py b/vime/backends/vllm_utils/vllm_engine.py index 82c07e6ff..b831a117c 100644 --- a/vime/backends/vllm_utils/vllm_engine.py +++ b/vime/backends/vllm_utils/vllm_engine.py @@ -10,6 +10,7 @@ import cloudpickle import requests +from urllib3.exceptions import NewConnectionError from vllm.utils.system_utils import kill_process_tree from vime.backends.vllm_utils.external import get_server_info @@ -286,9 +287,23 @@ def flush_cache(self): if self.node_rank != 0: return params = {"reset_running_requests": False} - requests.post( - f"http://{self.server_host}:{self.server_port}/reset_prefix_cache", params=params - ).raise_for_status() + for _ in range(60): + try: + response = requests.post( + f"http://{self.server_host}:{self.server_port}/reset_prefix_cache", params=params + ) + if response.status_code == 200 and response.json()["success"]: + break + logger.info(f"Error flushing cache: HTTP {response.status_code} {response.text!r}") + time.sleep(1) + except NewConnectionError as e: + raise e + except Exception as e: + logger.info(f"Error flushing cache: {e}") + time.sleep(1) + continue + else: + raise TimeoutError("Timeout while flushing cache.") def get_url(self): if self.node_rank != 0: @@ -390,9 +405,7 @@ def pull_weights(self, target_version: int): }, ) response.raise_for_status() - result = response.json() - self.set_weight_version(str(target_version)) - return result + return response.json() def update_weights_from_disk( self, @@ -632,6 +645,7 @@ def _compute_server_args( "tensor_parallel_size": tp, "logprobs_mode": "processed_logprobs", "enable_prompt_tokens_details": True, + "enable_per_request_metrics": True, "enable_server_load_tracking": True, } @@ -764,5 +778,6 @@ def _vllm_server_field_names() -> frozenset[str]: "tensor_parallel_size", "logprobs_mode", "enable_prompt_tokens_details", + "enable_per_request_metrics", "enable_server_load_tracking", ] diff --git a/vime/observability/logging_utils.py b/vime/observability/logging_utils.py index 428b798f1..74a75b9bd 100644 --- a/vime/observability/logging_utils.py +++ b/vime/observability/logging_utils.py @@ -8,7 +8,7 @@ _LOGGER_CONFIGURED = False -# ref: vLLM +# ref: SGLang def configure_logger(prefix: str = ""): global _LOGGER_CONFIGURED if _LOGGER_CONFIGURED: diff --git a/vime/observability/rollout_metrics.py b/vime/observability/rollout_metrics.py index 0e31bdfec..7ae2bff59 100644 --- a/vime/observability/rollout_metrics.py +++ b/vime/observability/rollout_metrics.py @@ -190,7 +190,7 @@ def _compute_top_p_kept_vocab_metrics(all_samples: list[Sample]): def _compute_spec_metrics(args, all_samples: list[Sample]): - if getattr(args, "vllm_speculative_algorithm", None) is None: + if getattr(args, "vllm_speculative_config", None) is None: return {} num_samples = len(all_samples) metrics = {} diff --git a/vime/observability/trace_utils.py b/vime/observability/trace_utils.py index 34b4a3e93..de613f638 100644 --- a/vime/observability/trace_utils.py +++ b/vime/observability/trace_utils.py @@ -140,6 +140,31 @@ def _new_span_id() -> str: def build_vllm_meta_trace_attrs(meta: dict[str, Any]) -> dict[str, Any]: attrs: dict[str, Any] = {} try: + if meta.get("choices"): + meta = dict(meta) + meta["finish_reason"] = meta["choices"][0].get("finish_reason") + if meta.get("usage"): + meta = dict(meta) + usage = meta["usage"] + meta["prompt_tokens"] = usage.get("prompt_tokens", 0) + meta["completion_tokens"] = usage.get("completion_tokens", 0) + meta["cached_tokens"] = (usage.get("prompt_tokens_details") or {}).get("cached_tokens", 0) + request_metrics = meta.get("request_metrics") + if isinstance(request_metrics, dict): + meta = dict(meta) + for target, source, scale in ( + ("queue_time", "queue_time_ms", 0.001), + ("decode_throughput", "tokens_per_second", 1.0), + ("pd_decode_transfer_duration", "remote_kv_wait_time_ms", 0.001), + ): + if request_metrics.get(source) is not None: + meta[target] = request_metrics[source] * scale + latency_parts = [ + request_metrics.get(key) for key in ("queue_time_ms", "time_to_first_token_ms", "generation_time_ms") + ] + if all(value is not None for value in latency_parts): + meta["e2e_latency"] = sum(latency_parts) / 1000 + attrs.update({key: meta[key] for key in VLLM_TRACE_META_KEYS if key in meta and meta[key] is not None}) finish_reason = meta.get("finish_reason") if isinstance(finish_reason, dict) and finish_reason.get("type") is not None: @@ -147,8 +172,9 @@ def build_vllm_meta_trace_attrs(meta: dict[str, Any]) -> dict[str, Any]: elif finish_reason is not None: attrs["finish_reason"] = finish_reason - if meta.get("id") is not None: - attrs["vllm_request_id"] = meta["id"] + request_id = meta.get("request_id", meta.get("id")) + if request_id is not None: + attrs["vllm_request_id"] = request_id trace_children = _build_vllm_pd_trace_children(meta) if trace_children: diff --git a/vime/ray/rollout.py b/vime/ray/rollout.py index 78d21d15f..c4cdae5ea 100644 --- a/vime/ray/rollout.py +++ b/vime/ray/rollout.py @@ -408,7 +408,7 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl if samples[0].rollout_log_probs is not None: train_data["rollout_log_probs"] = [sample.rollout_log_probs for sample in samples] - if getattr(self.args, "rollout_top_p", 1.0) != 1.0 and samples[0].rollout_top_p_token_ids is not None: + if getattr(self.args, "rollout_top_p", 1.0) != 1.0: for sample in samples: assert sample.rollout_top_p_token_ids is not None assert sample.rollout_top_p_token_offsets is not None diff --git a/vime/rollout/vllm_rollout.py b/vime/rollout/vllm_rollout.py index 4629ef0fd..8372d5edf 100644 --- a/vime/rollout/vllm_rollout.py +++ b/vime/rollout/vllm_rollout.py @@ -7,7 +7,7 @@ import logging import uuid from argparse import Namespace -from collections.abc import Callable +from collections.abc import Awaitable, Callable from contextlib import contextmanager from typing import Any @@ -191,6 +191,8 @@ def reset(self) -> None: self.remaining_batch_size = 0 self.pendings = set() self.aborted = False + self.cancellable_tasks = set() + self.active_server_generations = 0 def submit_generate_tasks(self, samples: list[list[Sample]]) -> None: for group in samples: @@ -227,6 +229,10 @@ def _build_inference_sampling_params(sampling_params: dict[str, Any]) -> dict[st sp["stop_token_ids"] = sampling_params["stop_token_ids"] if sampling_params.get("seed") is not None: sp["seed"] = sampling_params["seed"] + if sampling_params.get("min_new_tokens") is not None: + sp["min_tokens"] = sampling_params["min_new_tokens"] + if sampling_params.get("repetition_penalty") is not None: + sp["repetition_penalty"] = sampling_params["repetition_penalty"] if sampling_params.get("skip_special_tokens") is not None: sp["skip_special_tokens"] = bool(sampling_params["skip_special_tokens"]) return sp @@ -413,6 +419,20 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A skip_decode = True if skip_sp is None else bool(skip_sp) text = state.tokenizer.decode(new_response_tokens, skip_special_tokens=skip_decode) if new_response_tokens else "" + sample.append_response_tokens( + args, + tokens=new_response_tokens, + log_probs=new_response_log_probs, + trainable=True, + meta_info=_inference_generate_meta_info(output), + text=text, + ) + + return sample + + +def _inference_generate_meta_info(output: dict[str, Any]) -> dict[str, Any]: + choice = output["choices"][0] # Build meta_info from the vLLM `choices` response format. fr = choice.get("finish_reason") or "stop" if isinstance(fr, dict): @@ -453,17 +473,37 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A for token_ids in sampling_mask: offsets.append(offsets[-1] + len(token_ids)) meta["top_p_token_offsets"] = offsets + return meta - sample.append_response_tokens( - args, - tokens=new_response_tokens, - log_probs=new_response_log_probs, - trainable=True, - meta_info=meta, - text=text, - ) - return sample +async def _run_request_abortable_generate( + state: GenerateState, + sample: Sample, + generate_call: Awaitable[Sample | list[Sample]], +) -> Sample | list[Sample]: + task = asyncio.current_task() + assert task is not None + state.cancellable_tasks.add(task) + try: + return await generate_call + except asyncio.CancelledError: + if task in state.cancellable_tasks: + raise + sample.status = Sample.Status.ABORTED + return sample + finally: + state.cancellable_tasks.discard(task) + + +async def _run_server_abort_generate( + state: GenerateState, + generate_call: Awaitable[Sample | list[Sample]], +) -> Sample | list[Sample]: + state.active_server_generations += 1 + try: + return await generate_call + finally: + state.active_server_generations -= 1 @trace_function("generate_and_rm", target="sample") @@ -493,18 +533,18 @@ async def generate_and_rm( return sample with state.dp_rank_context() as _: - # Check sample.generate_function_path for per-sample custom_generate_function_path (e.g., from eval dataset config) - custom_func_path = getattr(sample, "generate_function_path", None) or args.custom_generate_function_path - - if custom_func_path is not None: - custom_generate_func = load_function(custom_func_path) - # if signature has evaluation, pass evaluation - if "evaluation" in inspect.signature(custom_generate_func).parameters: - sample = await custom_generate_func(args, sample, sampling_params, evaluation=evaluation) - else: - sample = await custom_generate_func(args, sample, sampling_params) + custom_func_path = sample.generate_function_path or args.custom_generate_function_path + generate_func = load_function(custom_func_path) if custom_func_path is not None else generate + + if custom_func_path is not None and "evaluation" in inspect.signature(generate_func).parameters: + generate_call = generate_func(args, sample, sampling_params, evaluation=evaluation) + else: + generate_call = generate_func(args, sample, sampling_params) + + if getattr(generate_func, "abort_mode", None) == "request": + sample = await _run_request_abortable_generate(state, sample, generate_call) else: - sample = await generate(args, sample, sampling_params) + sample = await _run_server_abort_generate(state, generate_call) sample = await apply_rollout_sample_hooks(args, sample, evaluation=evaluation) @@ -588,8 +628,14 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: assert not state.aborted state.aborted = True + cancellable_tasks = list(state.cancellable_tasks) + state.cancellable_tasks.difference_update(cancellable_tasks) + for task in cancellable_tasks: + task.cancel() + loop = asyncio.get_running_loop() - if state.pendings: + server_abort = state.active_server_generations > 0 + if server_abort: base = f"http://{args.vllm_router_ip}:{args.vllm_router_port}" response = await get(f"{base}/workers") urls = [worker["url"] for worker in response["workers"]] @@ -598,6 +644,8 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: await abort_inflight_requests(urls) last_sweep = loop.time() + await asyncio.gather(*cancellable_tasks, return_exceptions=True) + # make sure all the pending tasks are finished count = 0 while state.pendings: @@ -609,7 +657,7 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: # Re-sweep on a fixed interval to truncate late stragglers (e.g. a # multi-turn turn-2 fired after the initial abort), regardless of drain. - if loop.time() - last_sweep >= _ABORT_RESWEEP_INTERVAL_S: + if server_abort and loop.time() - last_sweep >= _ABORT_RESWEEP_INTERVAL_S: await abort_inflight_requests(urls) last_sweep = loop.time() @@ -619,6 +667,8 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: # for partial rollout, collect the partial samples into the data buffer for task in done: group = task.result() + if not any(sample.status == Sample.Status.ABORTED and sample.response_length > 0 for sample in group): + continue for sample in group: if sample.response and "start_rollout_id" not in sample.metadata: sample.metadata["start_rollout_id"] = rollout_id diff --git a/vime/rollout/vllm_streaming_rollout.py b/vime/rollout/vllm_streaming_rollout.py index 9007b37dc..6eac7d7ed 100644 --- a/vime/rollout/vllm_streaming_rollout.py +++ b/vime/rollout/vllm_streaming_rollout.py @@ -17,6 +17,12 @@ partial-rollout buffer hand-off) is still owned by ``vllm_rollout``; this file only replaces the inner HTTP call. +This generator selects request-level abort, so Vime cancels each active HTTP +stream instead of aborting every request on its vLLM server. + +Request cancellation preserves only metadata received before disconnect; +terminal-only data such as routed-expert replay is unavailable after abort. + vLLM's ``/inference/v1/generate`` SSE chunks carry **delta** ``token_ids`` + ``logprobs`` per ``GenerateResponseStreamChoice`` — so we *accumulate* the per-chunk deltas (``+=``) rather than overwriting from each chunk. Each delta @@ -40,12 +46,13 @@ _align_mm_feature_placeholders_to_tokens, _build_inference_sampling_params, _coerce_flat_int_token_ids, + _inference_generate_meta_info, _mm_render_response_to_generate_body, _prepare_prompt_ids, prime_encoder, ) from vime.utils import http_utils -from vime.utils.processing_utils import build_multimodal_messages, build_processor_kwargs +from vime.utils.processing_utils import build_multimodal_messages from vime.utils.types import Sample __all__ = ["generate_streaming"] @@ -53,24 +60,6 @@ logger = logging.getLogger(__name__) -def _base_dataset_prompt_ids(sample: Sample, tokenizer, processor: Any) -> list[int]: - """Token ids for the dataset prompt only (never reuse ``sample.tokens``). - - Used for partial-continuation budgeting: ``max_new_tokens -= len(sample.tokens) - - len(base_prompt_ids)`` when ``sample.response`` is non-empty. vLLM's - ``/inference/v1/generate`` is token-only, so on a partial resume we re-send the - full prefix and must subtract the already-generated tokens from the budget. - This lives here (not in ``vllm_rollout``) because it is specific to the - streaming path's partial-continuation handling. - """ - raw_multimodal_inputs = sample.multimodal_inputs or {} - has_multimodal_inputs = any(value is not None for value in raw_multimodal_inputs.values()) - if processor and has_multimodal_inputs: - processor_output = processor(text=sample.prompt, **build_processor_kwargs(raw_multimodal_inputs)) - return _coerce_flat_int_token_ids(processor_output["input_ids"][0]) - return _coerce_flat_int_token_ids(tokenizer.encode(sample.prompt, add_special_tokens=False)) - - async def generate_streaming(args: Namespace, sample: Sample, sampling_params: dict[str, Any]) -> Sample: """Streaming counterpart to :func:`vime.rollout.vllm_rollout.generate`. @@ -90,17 +79,13 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d ), f"Sample status is {sample.status}" prompt_ids = _prepare_prompt_ids(sample, state.tokenizer, state.processor) - base_prompt_ids = _base_dataset_prompt_ids(sample, state.tokenizer, state.processor) messages = build_multimodal_messages(sample.prompt, sample.multimodal_inputs) params = dict(sampling_params) - if len(sample.response) > 0: - params["max_new_tokens"] -= len(sample.tokens) - len(base_prompt_ids) + params["max_new_tokens"] -= sample.response_length - assert ( - params["max_new_tokens"] >= 0 - ), f"max_new_tokens: {params['max_new_tokens']} should not be less than 0 (after partial continuation adjustment; tokens={len(sample.tokens)}, base_prompt={len(base_prompt_ids)})" + assert params["max_new_tokens"] >= 0, f"max_new_tokens: {params['max_new_tokens']} should not be less than 0" if params["max_new_tokens"] == 0: sample.status = Sample.Status.TRUNCATED return sample @@ -159,7 +144,7 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d last_usage: dict[str, Any] | None = None weight_version: str | None = None request_spec_decode_stats: dict[str, int] | None = None - sampling_mask: list[list[int]] | None = None + trace_metadata: dict[str, Any] = {} finish_reason: Any = None client = http_utils._http_client @@ -186,6 +171,9 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d weight_version = str(chunk["weight_version"]) if chunk.get("request_spec_decode_stats") is not None: request_spec_decode_stats = chunk["request_spec_decode_stats"] + for key in ("request_id", "request_metrics"): + if chunk.get(key) is not None: + trace_metadata[key] = chunk[key] choices = chunk.get("choices") or [] if not choices: @@ -195,10 +183,6 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d continue choice = choices[0] last_choice = choice - if choice.get("sampling_mask") is not None: - if sampling_mask is None: - sampling_mask = [] - sampling_mask.extend(choice["sampling_mask"]) if chunk.get("usage"): last_usage = chunk["usage"] if choice.get("finish_reason"): @@ -231,12 +215,18 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d if base_loss_mask is not None: assert args.partial_rollout and args.mask_offpolicy_in_partial_rollout sample.loss_mask = base_loss_mask + [1] * len(call_tokens) + sample._apply_meta_info( + args, + _inference_generate_meta_info(chunk), + new_token_count=len(delta_tokens), + update_terminal_info=False, + ) if state.aborted: break if finish_reason and last_choice is not None: - span.update(build_vllm_meta_trace_attrs({"choices": [last_choice], "usage": last_usage})) + span.update(build_vllm_meta_trace_attrs({**trace_metadata, "choices": [last_choice], "usage": last_usage})) if finish_reason and last_choice is not None: new_response_tokens = call_tokens @@ -282,23 +272,16 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d if last_choice.get("routed_experts") is not None: raw = base64.b64decode(last_choice["routed_experts"].encode("ascii"), validate=True) meta["routed_experts"] = np.load(io.BytesIO(raw), allow_pickle=False) - if sampling_mask is not None: - top_p_meta = {"top_p_token_ids": [token_id for token_ids in sampling_mask for token_id in token_ids]} - offsets = [0] - for token_ids in sampling_mask: - offsets.append(offsets[-1] + len(token_ids)) - top_p_meta["top_p_token_offsets"] = offsets - sample._apply_meta_info( - args, - top_p_meta, - new_token_count=len(new_response_tokens), - update_terminal_info=False, - ) # tokens already accumulated above; finalize metadata only (no token re-append). sample.append_response_tokens(args, meta_info=meta) elif state.aborted: if weight_version is not None: sample.weight_versions.append(weight_version) sample.status = Sample.Status.ABORTED + else: + raise RuntimeError("vLLM streaming response ended without a terminal finish_reason.") return sample + + +generate_streaming.abort_mode = "request" diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index 4eefc41fb..59e6eb630 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -499,7 +499,9 @@ def add_rollout_arguments(parser): default=None, help=( "Only substitue the `def generate(args, sample, sampling_params)` function within the example rollout function. " - "This should be useful if you need to implement some special rollout logic, e.g. multi-turn, function calling." + "This should be useful if you need to implement some special rollout logic, e.g. multi-turn, function calling. " + "Set `abort_mode = 'request'` on the function when cancelling its task aborts only that request; " + "otherwise Vime aborts all in-flight requests on the server." ), ) parser.add_argument( @@ -2074,8 +2076,8 @@ def vime_validate_args(args): "debug_rollout_only and debug_train_only cannot be set at the same time, " "please set only one of them." ) - # Colocate normally offloads Megatron between rollout and train. Release-train - # destroys Megatron actors instead, so only rollout needs memory-saver offload. + # Colocate normally offloads Megatron between rollout and train. Release-train mode + # releases Megatron actors instead, so only rollout needs memory-saver offload. if args.colocate: if args.release_train: if args.offload_train: diff --git a/vime_plugins/models/glm5/glm5.py b/vime_plugins/models/glm5/glm5.py index 5847ea6f6..f7807fcd1 100644 --- a/vime_plugins/models/glm5/glm5.py +++ b/vime_plugins/models/glm5/glm5.py @@ -743,7 +743,7 @@ def _get_indexer_q_input(self, q_compressed: torch.Tensor) -> torch.Tensor: if self.config.layernorm_zero_centered_gamma: norm_weight = norm_weight + 1.0 if os.getenv("MEGATRON_USE_VLLM_FUSED_RESIDUAL_RMS", "0") == "1": - from vllm.model_executor.layers.batch_invariant import rms_norm_batch_invariant + from vllm.model_executor.determinism.batch_invariant import rms_norm_batch_invariant return rms_norm_batch_invariant( q_compressed,