diff --git a/.buildkite/scripts/ci-bake-rocm.sh b/.buildkite/scripts/ci-bake-rocm.sh index 3b9a9a1d7136..31d56bdc2996 100644 --- a/.buildkite/scripts/ci-bake-rocm.sh +++ b/.buildkite/scripts/ci-bake-rocm.sh @@ -17,7 +17,7 @@ DEFAULT_REPO_SLUG="vllm-project/vllm" DEFAULT_CI_HCL_SOURCE="docker/ci-rocm.hcl" DEFAULT_CI_BASE_CONTENT_FILES="requirements/common.txt requirements/rocm.txt requirements/test/rocm.txt docker/Dockerfile.rocm_base docker/ci-rocm.hcl docker/docker-bake-rocm.hcl tools/install_torchcodec_rocm.sh tools/install_protoc.sh rust-toolchain.toml tests/vllm_test_utils .buildkite/scripts/ci-bake-rocm.sh .buildkite/scripts/rocm/build-ci-base.sh" DEFAULT_CI_BASE_DOCKERFILE="docker/Dockerfile.rocm" -DEFAULT_CI_BASE_DOCKERFILE_STAGES="base rust_toolchain_input_0 rust_toolchain_input_1 rust-toolchain-input rust-toolchain build_rixl build_rocshmem build_deepep mori_base ci_base" +DEFAULT_CI_BASE_DOCKERFILE_STAGES="base rust_toolchain_input_0 rust_toolchain_input_1 rust-toolchain-input rust-toolchain build_nixl build_rocshmem build_deepep mori_base ci_base" DEFAULT_CI_BASE_METADATA_VERSION="1" IMAGE_EXISTED_BEFORE_BUILD=0 @@ -1159,8 +1159,8 @@ ci_base_metadata_pairs() { metadata_pair "vllm.rocm.nic_backend" "$(resolve_dockerfile_arg_value "${dockerfile}" "NIC_BACKEND")" metadata_pair "vllm.rocm.ainic_version" "$(resolve_dockerfile_arg_value "${dockerfile}" "AINIC_VERSION")" metadata_pair "vllm.rocm.ubuntu_codename" "$(resolve_dockerfile_arg_value "${dockerfile}" "UBUNTU_CODENAME")" - metadata_pair "vllm.rocm.rixl_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "RIXL_REPO")" - metadata_pair "vllm.rocm.rixl_commit" "${RIXL_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "RIXL_BRANCH")}" + metadata_pair "vllm.rocm.nixl_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "NIXL_REPO")" + metadata_pair "vllm.rocm.nixl_commit" "${NIXL_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "NIXL_BRANCH")}" metadata_pair "vllm.rocm.ucx_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "UCX_REPO")" metadata_pair "vllm.rocm.ucx_commit" "${UCX_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "UCX_BRANCH")}" metadata_pair "vllm.rocm.rocshmem_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "ROCSHMEM_REPO")" @@ -1169,7 +1169,7 @@ ci_base_metadata_pairs() { metadata_pair "vllm.rocm.deepep_commit" "${DEEPEP_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "DEEPEP_BRANCH")}" metadata_pair "vllm.rocm.deepep_nic" "$(resolve_dockerfile_arg_value "${dockerfile}" "DEEPEP_NIC")" metadata_pair "vllm.rocm.deepep_rocm_arch" "$(resolve_dockerfile_arg_value "${dockerfile}" "DEEPEP_ROCM_ARCH")" - metadata_pair "vllm.rocm.rixl_cache_key" "${RIXL_CACHE_KEY:-}" + metadata_pair "vllm.rocm.nixl_cache_key" "${NIXL_CACHE_KEY:-}" metadata_pair "vllm.rocm.rocshmem_cache_key" "${ROCSHMEM_CACHE_KEY:-}" metadata_pair "vllm.rocm.deepep_cache_key" "${DEEPEP_CACHE_KEY:-}" @@ -1686,7 +1686,7 @@ extract_dependency_pins() { return 0 fi - for var in RIXL_BRANCH UCX_BRANCH ROCSHMEM_BRANCH DEEPEP_BRANCH; do + for var in NIXL_BRANCH UCX_BRANCH ROCSHMEM_BRANCH DEEPEP_BRANCH; do if [[ -n "${!var:-}" ]]; then echo "Using provided ${var}: ${!var}" continue @@ -1706,30 +1706,30 @@ extract_dependency_pins() { compute_dependency_cache_keys() { local bake_dir="" local dockerfile_rocm="" - local rixl_branch="" + local nixl_branch="" local ucx_branch="" local rocshmem_branch="" local deepep_branch="" - local rixl_material="" + local nixl_material="" local rocshmem_material="" local deepep_material="" bake_dir=$(dirname "${VLLM_BAKE_FILE}") dockerfile_rocm="${bake_dir}/Dockerfile.rocm" - rixl_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "RIXL_BRANCH") + nixl_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "NIXL_BRANCH") ucx_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "UCX_BRANCH") rocshmem_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "ROCSHMEM_BRANCH") deepep_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "DEEPEP_BRANCH") - if [[ -n "${rixl_branch}" && -n "${ucx_branch}" ]]; then - rixl_material=$(compose_stage_cache_material "${dockerfile_rocm}" "base build_rixl") - RIXL_CACHE_KEY=$( + if [[ -n "${nixl_branch}" && -n "${ucx_branch}" ]]; then + nixl_material=$(compose_stage_cache_material "${dockerfile_rocm}" "base build_nixl") + NIXL_CACHE_KEY=$( compose_dependency_cache_key \ - "${rixl_branch}-ucx-${ucx_branch}" \ - "${rixl_material}" + "${nixl_branch}-ucx-${ucx_branch}" \ + "${nixl_material}" ) - export RIXL_CACHE_KEY - echo "RIXL dependency cache key: ${RIXL_CACHE_KEY}" + export NIXL_CACHE_KEY + echo "NIXL dependency cache key: ${NIXL_CACHE_KEY}" fi if [[ -n "${rocshmem_branch}" ]]; then @@ -1780,11 +1780,11 @@ dependency_cache_ref_for_target() { local cache_repo="${DOCKERHUB_CACHE_REPO:-rocm/vllm-ci-cache}" case "${target}" in - rixl-rocm-ci) - if [[ -n "${RIXL_CACHE_KEY:-}" ]]; then - printf '%s\n' "${cache_repo}:rixl-rocm-${RIXL_CACHE_KEY}" - elif [[ -n "${RIXL_BRANCH:-}" ]]; then - printf '%s\n' "${cache_repo}:rixl-rocm-${RIXL_BRANCH}-ucx-${UCX_BRANCH:-}" + nixl-rocm-ci) + if [[ -n "${NIXL_CACHE_KEY:-}" ]]; then + printf '%s\n' "${cache_repo}:nixl-rocm-${NIXL_CACHE_KEY}" + elif [[ -n "${NIXL_BRANCH:-}" ]]; then + printf '%s\n' "${cache_repo}:nixl-rocm-${NIXL_BRANCH}-ucx-${UCX_BRANCH:-}" fi ;; rocshmem-rocm-ci) @@ -1815,7 +1815,7 @@ add_dependency_cache_target() { resolve_ci_base_dependency_targets() { local mode="${ROCM_DEP_CACHE_EXPORT_MODE:-missing}" - local rixl_ref="" + local nixl_ref="" local rocshmem_ref="" local deepep_ref="" @@ -1824,7 +1824,7 @@ resolve_ci_base_dependency_targets() { case "${mode}" in always) echo "ROCM_DEP_CACHE_EXPORT_MODE=always; exporting all dependency caches serially" - for target in rixl-rocm-ci rocshmem-rocm-ci deepep-rocm-ci; do + for target in nixl-rocm-ci rocshmem-rocm-ci deepep-rocm-ci; do if [[ -n "$(dependency_cache_ref_for_target "${target}")" ]]; then add_dependency_cache_target "${target}" fi @@ -1844,13 +1844,13 @@ resolve_ci_base_dependency_targets() { ;; esac - if [[ "${mode}" != "always" && -n "${RIXL_CACHE_KEY:-}" ]]; then - rixl_ref=$(dependency_cache_ref_for_target "rixl-rocm-ci") - if dependency_cache_ref_exists "${rixl_ref}"; then - echo "RIXL dependency cache exists: ${rixl_ref}" + if [[ "${mode}" != "always" && -n "${NIXL_CACHE_KEY:-}" ]]; then + nixl_ref=$(dependency_cache_ref_for_target "nixl-rocm-ci") + if dependency_cache_ref_exists "${nixl_ref}"; then + echo "NIXL dependency cache exists: ${nixl_ref}" else - echo "RIXL dependency cache missing; will seed: ${rixl_ref}" - add_dependency_cache_target "rixl-rocm-ci" + echo "NIXL dependency cache missing; will seed: ${nixl_ref}" + add_dependency_cache_target "nixl-rocm-ci" fi fi diff --git a/.buildkite/test_areas/basic_correctness.yaml b/.buildkite/test_areas/basic_correctness.yaml index 1b92babb3ce7..88f3e0548250 100644 --- a/.buildkite/test_areas/basic_correctness.yaml +++ b/.buildkite/test_areas/basic_correctness.yaml @@ -4,7 +4,7 @@ depends_on: steps: - label: Basic Correctness key: basic-correctness - timeout_in_minutes: 45 + timeout_in_minutes: 68 device: h200_18gb source_file_dependencies: - vllm/ diff --git a/.buildkite/test_areas/benchmarks.yaml b/.buildkite/test_areas/benchmarks.yaml index 4b709f3962fd..8aac6545cfa5 100644 --- a/.buildkite/test_areas/benchmarks.yaml +++ b/.buildkite/test_areas/benchmarks.yaml @@ -4,7 +4,7 @@ depends_on: steps: - label: Benchmarks CLI Test key: benchmarks-cli-test - timeout_in_minutes: 30 + timeout_in_minutes: 45 device: h200_18gb source_file_dependencies: - vllm/ diff --git a/.buildkite/test_areas/engine.yaml b/.buildkite/test_areas/engine.yaml index 4557aa0fcbec..6e88451af861 100644 --- a/.buildkite/test_areas/engine.yaml +++ b/.buildkite/test_areas/engine.yaml @@ -51,7 +51,7 @@ steps: - label: e2e Scheduling (1 GPU) key: e2e-scheduling-1-gpu - timeout_in_minutes: 35 + timeout_in_minutes: 53 device: h200_18gb source_file_dependencies: - vllm/v1/ diff --git a/.buildkite/test_areas/entrypoints.yaml b/.buildkite/test_areas/entrypoints.yaml index f5668e40b9da..9d571ff0af53 100644 --- a/.buildkite/test_areas/entrypoints.yaml +++ b/.buildkite/test_areas/entrypoints.yaml @@ -39,7 +39,7 @@ steps: - label: Entrypoints Integration (API Server) key: entrypoints-integration-api-server device: h200_35gb - timeout_in_minutes: 50 + timeout_in_minutes: 75 working_dir: "/vllm-workspace/tests" source_file_dependencies: - vllm/ @@ -59,7 +59,7 @@ steps: - label: Entrypoints Integration (API Server OpenAI - Part 1) device: h200_35gb key: entrypoints-integration-api-server-openai-part-1 - timeout_in_minutes: 45 + timeout_in_minutes: 68 working_dir: "/vllm-workspace/tests" source_file_dependencies: - vllm/ @@ -78,7 +78,7 @@ steps: - label: Entrypoints Integration (API Server OpenAI - Part 2) device: h200_35gb key: entrypoints-integration-api-server-openai-part-2 - timeout_in_minutes: 55 + timeout_in_minutes: 83 working_dir: "/vllm-workspace/tests" source_file_dependencies: - vllm/ @@ -156,7 +156,7 @@ steps: - label: Entrypoints Integration (Pooling) device: h200_35gb key: entrypoints-integration-pooling - timeout_in_minutes: 50 + timeout_in_minutes: 75 working_dir: "/vllm-workspace/tests" source_file_dependencies: - vllm/ diff --git a/.buildkite/test_areas/misc.yaml b/.buildkite/test_areas/misc.yaml index a427ac5e7733..d61a6a816fb3 100644 --- a/.buildkite/test_areas/misc.yaml +++ b/.buildkite/test_areas/misc.yaml @@ -31,7 +31,7 @@ steps: - label: V1 Sample + Logits key: v1-sample-logits - timeout_in_minutes: 55 + timeout_in_minutes: 83 device: h200_18gb source_file_dependencies: - vllm/config/ @@ -90,6 +90,7 @@ steps: - tests/v1/kv_offload - tests/v1/simple_kv_offload - tests/v1/worker + - tests/v1/streaming_input - tests/v1/kv_connector/unit - tests/v1/ec_connector/unit - tests/v1/metrics @@ -103,6 +104,7 @@ steps: - pytest -v -s v1/kv_offload - pytest -v -s v1/simple_kv_offload - pytest -v -s v1/worker + - pytest -v -s v1/streaming_input - pytest -v -s -m 'not cpu_test' v1/kv_connector/unit - pytest -v -s -m 'not cpu_test' v1/ec_connector/unit - pytest -v -s -m 'not cpu_test' v1/metrics diff --git a/.buildkite/test_areas/model_executor.yaml b/.buildkite/test_areas/model_executor.yaml index 689d19488b56..e2f90079291e 100644 --- a/.buildkite/test_areas/model_executor.yaml +++ b/.buildkite/test_areas/model_executor.yaml @@ -5,7 +5,7 @@ steps: - label: Model Executor device: h200_35gb key: model-executor - timeout_in_minutes: 45 + timeout_in_minutes: 60 source_file_dependencies: - vllm/engine/arg_utils.py - vllm/config/model.py diff --git a/.buildkite/test_areas/models_language.yaml b/.buildkite/test_areas/models_language.yaml index b37f2dbd20f9..2aea721e56c5 100644 --- a/.buildkite/test_areas/models_language.yaml +++ b/.buildkite/test_areas/models_language.yaml @@ -137,7 +137,7 @@ steps: - label: Language Models Test (MTEB) key: language-models-test-mteb - timeout_in_minutes: 45 + timeout_in_minutes: 68 device: h200_18gb optional: true source_file_dependencies: diff --git a/.buildkite/test_areas/models_multimodal.yaml b/.buildkite/test_areas/models_multimodal.yaml index 2a73eb4a47ea..57f559c59fb6 100644 --- a/.buildkite/test_areas/models_multimodal.yaml +++ b/.buildkite/test_areas/models_multimodal.yaml @@ -4,7 +4,7 @@ depends_on: steps: - label: "Multi-Modal Models (Standard) 1: qwen2" key: multi-modal-models-standard-1-qwen2 - timeout_in_minutes: 45 + timeout_in_minutes: 68 device: h200_18gb source_file_dependencies: - vllm/ @@ -20,7 +20,7 @@ steps: - label: "Multi-Modal Models (Standard) 2: qwen3 + gemma" key: multi-modal-models-standard-2-qwen3-gemma - timeout_in_minutes: 50 + timeout_in_minutes: 75 device: h200_18gb source_file_dependencies: - vllm/ @@ -54,7 +54,7 @@ steps: - label: "Multi-Modal Models (Standard) 4: other + whisper" device: h200_35gb key: multi-modal-models-standard-4-other-whisper - timeout_in_minutes: 50 + timeout_in_minutes: 75 source_file_dependencies: - vllm/ - tests/models/multimodal @@ -85,7 +85,7 @@ steps: - label: Multi-Modal Processor # 44min key: multi-modal-processor - timeout_in_minutes: 65 + timeout_in_minutes: 98 device: h200_18gb source_file_dependencies: - vllm/ diff --git a/.github/workflows/issue_autolabel.yml b/.github/workflows/issue_autolabel.yml index e5d1accf477c..fadbb9c58537 100644 --- a/.github/workflows/issue_autolabel.yml +++ b/.github/workflows/issue_autolabel.yml @@ -130,6 +130,35 @@ jobs: }, ], }, + "intel-gpu": { + // Keyword search - matches whole words only (with word boundaries) + keywords: [ + { + term: "B50", + searchIn: "both" + }, + { + term: "B60", + searchIn: "both" + }, + { + term: "B70", + searchIn: "both" + }, + { + term: "intel gpu", + searchIn: "both" + }, + { + term: "Arc GPU", + searchIn: "both" + }, + { + term: "BMG", + searchIn: "both" + }, + ], + }, // Add more label configurations here as needed // example: { // keywords: [...], diff --git a/cmake/external_projects/vllm_flash_attn.cmake b/cmake/external_projects/vllm_flash_attn.cmake index 89ef9192e846..6b559d3725c0 100644 --- a/cmake/external_projects/vllm_flash_attn.cmake +++ b/cmake/external_projects/vllm_flash_attn.cmake @@ -39,7 +39,7 @@ else() FetchContent_Declare( vllm-flash-attn GIT_REPOSITORY https://github.com/vllm-project/flash-attention.git - GIT_TAG 168920233059c48de6199e2cda74003b2ce3d199 + GIT_TAG ed4b7342bc8f0489dd9b649d5288867e35fc6a32 GIT_PROGRESS TRUE # Don't share the vllm-flash-attn build between build types BINARY_DIR ${CMAKE_BINARY_DIR}/vllm-flash-attn diff --git a/docker/Dockerfile.rocm b/docker/Dockerfile.rocm index a7cc8c0870bc..47438b06b52c 100644 --- a/docker/Dockerfile.rocm +++ b/docker/Dockerfile.rocm @@ -339,18 +339,17 @@ COPY --from=build_vllm ${COMMON_WORKDIR}/vllm/rust /rust COPY --from=build_vllm ${COMMON_WORKDIR}/vllm/rust-toolchain.toml /rust-toolchain.toml COPY --from=build_vllm ${COMMON_WORKDIR}/vllm/vllm/v1 /vllm_v1 -# RIXL/UCX build stages -FROM base AS build_rixl -ARG RIXL_BRANCH="39be1de8" -ARG RIXL_REPO="https://github.com/ROCm/RIXL.git" -ARG UCX_BRANCH="bfb51733" +# NIXL/UCX build stages +FROM base AS build_nixl +ARG NIXL_BRANCH="231d56753047c989062a5cb2ac703a1ad761c7d2" +ARG NIXL_REPO="https://github.com/ai-dynamo/nixl.git" +ARG UCX_BRANCH="96e58a16039f6d7d213bc967b8069238742c5194" ARG UCX_REPO="https://github.com/openucx/ucx.git" ENV ROCM_PATH=/opt/rocm ENV UCX_HOME=/usr/local/ucx -ENV RIXL_HOME=/usr/local/rixl -ENV RIXL_BENCH_HOME=/usr/local/rixl_bench +ENV NIXL_HOME=/usr/local/nixl -# RIXL build system dependences and RDMA support +# NIXL build system dependencies and RDMA support RUN apt-get -y update && apt-get -y install autoconf libtool pkg-config \ libgrpc-dev \ libgrpc++-dev \ @@ -368,7 +367,8 @@ RUN apt-get -y update && apt-get -y install autoconf libtool pkg-config \ && rm -rf /var/lib/apt/lists/* RUN --mount=type=cache,target=/root/.cache/uv \ - uv pip install --system meson auditwheel patchelf tomlkit + uv pip install --system meson meson-python pybind11 pyyaml types-PyYAML \ + auditwheel build patchelf pytest tomlkit "setuptools>=80.9.0" RUN --mount=type=cache,target=/root/.cache/ccache \ cd /usr/local/src && \ @@ -396,30 +396,50 @@ ENV PATH=/usr/local/ucx/bin:$PATH ENV LD_LIBRARY_PATH=${UCX_HOME}/lib:${LD_LIBRARY_PATH} RUN --mount=type=cache,target=/root/.cache/ccache \ - git clone ${RIXL_REPO} /opt/rixl && \ - cd /opt/rixl && \ - git checkout ${RIXL_BRANCH} && \ + git clone ${NIXL_REPO} /opt/nixl && \ + cd /opt/nixl && \ + git checkout ${NIXL_BRANCH} && \ CC="ccache gcc" CXX="ccache g++" \ - meson setup build --prefix=${RIXL_HOME} \ + meson setup build --prefix=${NIXL_HOME} \ -Ducx_path=${UCX_HOME} \ - -Drocm_path=${ROCM_PATH} && \ + -Dwheel_variant=rocm \ + -Dbuild_tests=false \ + -Dbuild_examples=false && \ cd build && \ ninja -j$(nproc) && \ - ninja install - -# Generate RIXL wheel + ninja install && \ + echo "${NIXL_HOME}/lib/$(uname -m)-linux-gnu" \ + > /etc/ld.so.conf.d/nixl.conf && \ + echo "${NIXL_HOME}/lib/$(uname -m)-linux-gnu/plugins" \ + >> /etc/ld.so.conf.d/nixl.conf && \ + ldconfig + +# Generate the ROCm NIXL wheel. Upstream's generic wheel helper detects CUDA, +# so configure the ROCm wheel variant directly through Meson. # Exclude libcore and libpull from auditwheel: transitive dependencies # that are not shipped in the wheel and vary across base images. -RUN cd /opt/rixl && \ - sed -i "s/--exclude 'libamdhip64\*'/--exclude 'libamdhip64*' --exclude 'libcore*' --exclude 'libpull*'/" \ - contrib/build-wheel.sh && \ - mkdir -p /app/install && \ - _ucx_install_dir=${UCX_HOME} \ - ./contrib/build-wheel.sh \ - --output-dir /app/install \ - --rocm-dir ${ROCM_PATH} \ +RUN cd /opt/nixl && \ + ./contrib/tomlutil.py --wheel-name nixl-rocm pyproject.toml && \ + CC="ccache gcc" CXX="ccache g++" \ + uv build --wheel --no-build-isolation --out-dir /tmp/nixl_wheels \ + --python ${PYTHON_VERSION} \ + -Csetup-args=-Ducx_path=${UCX_HOME} \ + -Csetup-args=-Dwheel_variant=rocm \ + -Csetup-args=-Dbuild_tests=false \ + -Csetup-args=-Dbuild_examples=false && \ + mkdir -p /tmp/nixl_wheels/repaired /app/install && \ + auditwheel repair \ + --exclude 'libamdhip64*' \ + --exclude 'libcore*' \ + --exclude 'libpull*' \ + /tmp/nixl_wheels/nixl_rocm*.whl \ + --plat manylinux_2_34_$(uname -m) \ + --wheel-dir /tmp/nixl_wheels/repaired && \ + ./contrib/wheel_add_ucx_plugins.py \ --ucx-plugins-dir ${UCX_HOME}/lib/ucx \ - --nixl-plugins-dir ${RIXL_HOME}/lib/x86_64-linux-gnu/plugins + --nixl-plugins-dir ${NIXL_HOME}/lib/$(uname -m)-linux-gnu/plugins \ + /tmp/nixl_wheels/repaired/*.whl && \ + cp /tmp/nixl_wheels/repaired/*.whl /app/install # ROCShmem build stage - split from DeepEP so changing DEEPEP_BRANCH does not # invalidate the slow ROCShmem build. @@ -660,10 +680,10 @@ RUN if [ "${DEEPEP_NIC}" = "cx7" ] || [ "${DEEPEP_NIC}" = "io" ]; then \ ninja && ninja install && ldconfig && rm -rf /tmp/rdma-core; \ fi -# Install RIXL + DeepEP wheels. -RUN --mount=type=bind,from=build_rixl,src=/app/install,target=/rixl_install \ +# Install NIXL + DeepEP wheels. +RUN --mount=type=bind,from=build_nixl,src=/app/install,target=/nixl_install \ --mount=type=bind,from=build_deepep,src=/app/deep_install,target=/deep_install \ - uv pip install --system /rixl_install/*.whl /deep_install/*.whl + uv pip install --system /nixl_install/*.whl /deep_install/*.whl # Copy ROCShmem runtime libraries. COPY --from=build_rocshmem /opt/rocshmem /opt/rocshmem @@ -724,6 +744,7 @@ ENV MIOPEN_DEBUG_CONV_GEMM=0 # Use legacy IPC mode for HSA to avoid GPU memory pinning issues with UCX rocm_ipc. # See: https://github.com/ROCm/rocm-libraries/issues/6266 ENV HSA_ENABLE_IPC_MODE_LEGACY=1 +ENV UCX_RMA_PPLN_ENABLE=y # ROCm profiler limits workaround. RUN echo "ROCTRACER_MAX_EVENTS=10000000" > ${COMMON_WORKDIR}/libkineto.conf @@ -796,9 +817,9 @@ RUN --mount=type=bind,from=export_vllm,src=/,target=/install \ && pip uninstall -y vllm \ && uv pip install --system *.whl -# Install RIXL wheel -RUN --mount=type=bind,from=build_rixl,src=/app/install,target=/rixl_install \ - uv pip install --system /rixl_install/*.whl +# Install NIXL ROCm wheel +RUN --mount=type=bind,from=build_nixl,src=/app/install,target=/nixl_install \ + uv pip install --system /nixl_install/*.whl ARG COMMON_WORKDIR ARG BASE_IMAGE @@ -813,6 +834,7 @@ COPY --from=export_vllm /docker ${COMMON_WORKDIR}/vllm/docker # Use legacy IPC mode for HSA to avoid GPU memory pinning issues with UCX rocm_ipc # See: https://github.com/ROCm/rocm-libraries/issues/6266 ENV HSA_ENABLE_IPC_MODE_LEGACY=1 +ENV UCX_RMA_PPLN_ENABLE=y ENV TOKENIZERS_PARALLELISM=false diff --git a/docker/ci-rocm.hcl b/docker/ci-rocm.hcl index 0ae991bd52d8..7062a0be9346 100644 --- a/docker/ci-rocm.hcl +++ b/docker/ci-rocm.hcl @@ -59,7 +59,7 @@ variable "PYTORCH_ROCM_ARCH" { } # Pre-built CI base image (Tier 1). Per-PR builds pull this instead of -# rebuilding RIXL/DeepEP/torchcodec from scratch. The ci_base stage in +# rebuilding NIXL/DeepEP/torchcodec from scratch. The ci_base stage in # Dockerfile.rocm inherits from base, so CI_BASE_IMAGE only affects the test # stage and is irrelevant when building --target ci_base itself. variable "CI_BASE_IMAGE" { @@ -75,7 +75,7 @@ variable "CI_MAX_JOBS" { # Upstream dependency commit pins -- extracted from Dockerfile.rocm by # ci-bake-rocm.sh at build time. Empty defaults are safe: the cache # functions produce no entries when the variable is empty. -variable "RIXL_BRANCH" { +variable "NIXL_BRANCH" { default = "" } @@ -91,7 +91,7 @@ variable "DEEPEP_BRANCH" { default = "" } -variable "RIXL_CACHE_KEY" { +variable "NIXL_CACHE_KEY" { default = "" } @@ -236,7 +236,7 @@ function "get_cache_to_rocm_rust" { ]) } -# Cache functions for upstream dependency stages (RIXL/UCX, ROCShmem, DeepEP). +# Cache functions for upstream dependency stages (NIXL/UCX, ROCShmem, DeepEP). # These stages are pinned to specific upstream commit hashes, so cache keys use # those hashes rather than the Buildkite commit. This means the cache persists # across all vLLM commits as long as the upstream dependency pins don't change. @@ -244,16 +244,16 @@ function "get_cache_to_rocm_rust" { function "get_cache_from_rocm_deps" { params = [] result = compact([ - RIXL_CACHE_KEY != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:rixl-rocm-${RIXL_CACHE_KEY}" : (RIXL_BRANCH != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:rixl-rocm-${RIXL_BRANCH}-ucx-${UCX_BRANCH}" : ""), + NIXL_CACHE_KEY != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:nixl-rocm-${NIXL_CACHE_KEY}" : (NIXL_BRANCH != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:nixl-rocm-${NIXL_BRANCH}-ucx-${UCX_BRANCH}" : ""), ROCSHMEM_CACHE_KEY != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:rocshmem-rocm-${ROCSHMEM_CACHE_KEY}" : (ROCSHMEM_BRANCH != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:rocshmem-rocm-${ROCSHMEM_BRANCH}" : ""), DEEPEP_CACHE_KEY != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:deepep-rocm-${DEEPEP_CACHE_KEY}" : (DEEPEP_BRANCH != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:deepep-rocm-${DEEPEP_BRANCH}-rocshmem-${ROCSHMEM_BRANCH}" : ""), ]) } -function "get_cache_to_rocm_rixl" { +function "get_cache_to_rocm_nixl" { params = [] result = compact([ - RIXL_CACHE_KEY != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:rixl-rocm-${RIXL_CACHE_KEY},mode=min" : (RIXL_BRANCH != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:rixl-rocm-${RIXL_BRANCH}-ucx-${UCX_BRANCH},mode=min" : ""), + NIXL_CACHE_KEY != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:nixl-rocm-${NIXL_CACHE_KEY},mode=min" : (NIXL_BRANCH != "" ? "type=registry,ref=${DOCKERHUB_CACHE_REPO}:nixl-rocm-${NIXL_BRANCH}-ucx-${UCX_BRANCH},mode=min" : ""), ]) } @@ -372,11 +372,11 @@ variable "CI_BASE_IMAGE_TAG_STABLE" { # in the registry cache keyed by its upstream commit hash. When ci_base rebuilds # (e.g., requirements change), these stages are cache hits if their upstream # pins haven't changed -- saving ~35min of compilation. -target "rixl-rocm-ci" { +target "nixl-rocm-ci" { inherits = ["_common-rocm", "_ci-rocm"] - target = "build_rixl" + target = "build_nixl" cache-from = get_cache_from_rocm_deps() - cache-to = get_cache_to_rocm_rixl() + cache-to = get_cache_to_rocm_nixl() output = ["type=cacheonly"] } @@ -396,7 +396,7 @@ target "deepep-rocm-ci" { output = ["type=cacheonly"] } -# Builds only the ci_base stage (RIXL, DeepEP, torchcodec, etc.) +# Builds only the ci_base stage (NIXL, DeepEP, torchcodec, etc.) # Invoked by the ensure-ci-base step when the content hash of ci_base-affecting # files drifts from the remote image label. Per-PR builds then pull the result # as CI_BASE_IMAGE instead of rebuilding those slow layers on every commit. @@ -412,7 +412,7 @@ target "ci-base-rocm-ci" { CI_BASE_IMAGE_TAG_CONTENT_EXTRA != "" ? "type=registry,ref=${CI_BASE_IMAGE_TAG_CONTENT_EXTRA}" : "", CI_BASE_IMAGE_TAG_STABLE != "" ? "type=registry,ref=${CI_BASE_IMAGE_TAG_STABLE}" : "", ]), - # Import upstream dependency caches so RIXL/ROCShmem/DeepEP stages + # Import upstream dependency caches so NIXL/ROCShmem/DeepEP stages # are cache hits even when ci_base itself needs rebuilding. get_cache_from_rocm_deps(), ) @@ -424,5 +424,5 @@ target "ci-base-rocm-ci" { # Group for ci_base builds -- exports dependency stage caches alongside the # ci_base image so future rebuilds can reuse them independently. group "ci-base-rocm-ci-with-deps" { - targets = ["rixl-rocm-ci", "rocshmem-rocm-ci", "deepep-rocm-ci", "ci-base-rocm-ci"] + targets = ["nixl-rocm-ci", "rocshmem-rocm-ci", "deepep-rocm-ci", "ci-base-rocm-ci"] } diff --git a/docker/docker-bake-rocm.hcl b/docker/docker-bake-rocm.hcl index 6b51781834b2..9d73a6505eab 100644 --- a/docker/docker-bake-rocm.hcl +++ b/docker/docker-bake-rocm.hcl @@ -53,7 +53,7 @@ variable "CI_BASE_IMAGE" { # Upstream dependency commit pins. Plain local bake builds use the Dockerfile # ARG defaults. ci-bake-rocm.sh resolves those defaults (plus any env # overrides) and writes a small HCL override before invoking CI targets. -variable "RIXL_BRANCH" { +variable "NIXL_BRANCH" { default = "" } @@ -106,7 +106,7 @@ target "test-rocm" { output = ["type=docker"] } -# CI base image target - builds only the ci_base stage (RIXL, DeepEP, +# CI base image target - builds only the ci_base stage (NIXL, DeepEP, # torchcodec, requirements, etc.). Used by the weekly scheduled build and # the auto-rebuild trigger when requirements change in a PR. target "ci-base-rocm" { diff --git a/docs/features/nixl_connector_usage.md b/docs/features/nixl_connector_usage.md index 03b05751c14d..a84e3de15b1a 100644 --- a/docs/features/nixl_connector_usage.md +++ b/docs/features/nixl_connector_usage.md @@ -13,11 +13,7 @@ Install the NIXL library: `uv pip install nixl`, as a quick start on Nvidia plat - Refer to [NIXL official repository](https://github.com/ai-dynamo/nixl) for more installation instructions - The specified required NIXL version can be found in [requirements/kv_connectors.txt](../../requirements/kv_connectors.txt) and other relevant config files -For ROCm platform, the [ROCm docker file](../../docker/Dockerfile.rocm) includes RIXL and ucx already. - -- Refer to [RIXL official repository](https://github.com/rocm/rixl) for more information -- The supportive libraries for RIXL can be found in [requirements/kv_connectors_rocm.txt](../../requirements/kv_connectors_rocm.txt) -- In the future we may remove RIXL from docker image file and users will be able to install from pre-compiled binary packages +For ROCm, the [ROCm Dockerfile](../../docker/Dockerfile.rocm) builds NIXL and UCX with ROCm support from source. For non-cuda platform, please install nixl with ucx build from source, instructed as below. diff --git a/requirements/tpu.txt b/requirements/tpu.txt index da477a68461c..38d91c44968a 100644 --- a/requirements/tpu.txt +++ b/requirements/tpu.txt @@ -12,4 +12,4 @@ ray[data] setuptools==78.1.0 setuptools-rust>=1.9.0 nixl==0.3.0 -tpu-inference==0.24.0 +tpu-inference==0.25.0 diff --git a/rust/proto/vllm_grpc.proto b/rust/proto/vllm_grpc.proto index 2509d5071b63..c10da4818e3b 100644 --- a/rust/proto/vllm_grpc.proto +++ b/rust/proto/vllm_grpc.proto @@ -14,6 +14,10 @@ service Generate { rpc GenerateStream (GenerateRequest) returns (stream GenerateResponse) {} } +service Control { + rpc Abort (AbortRequest) returns (AbortResponse) {} +} + // ====================================================================================== // Generate Request // ====================================================================================== @@ -201,3 +205,12 @@ message TokenIds { repeated uint32 ids = 1; } +// ====================================================================================== +// Control +// ====================================================================================== + +message AbortRequest { + repeated string request_ids = 1; +} + +message AbortResponse {} diff --git a/rust/src/chat/src/lib.rs b/rust/src/chat/src/lib.rs index 3188657b9096..e584257eaa43 100644 --- a/rust/src/chat/src/lib.rs +++ b/rust/src/chat/src/lib.rs @@ -54,6 +54,7 @@ mod stream; use vllm_engine_core_client::EngineCoreClient; use vllm_engine_core_client::protocol::dtype::ModelDtype; +use vllm_engine_core_client::protocol::multimodal::MmFeatures; use vllm_engine_core_client::protocol::request::ReasoningParserKwargs; use vllm_llm::Llm; use vllm_text::{Prompt, TextLlm, TextRequest}; @@ -88,6 +89,92 @@ pub fn validate_parser_overrides( Ok(()) } +/// Chat request preparation shared by inference and render-only frontends. +pub struct ChatRequestProcessor { + backend: DynChatBackend, + /// Effective model dtype reported by the engine. + /// Absent for text-only frontends without an engine handshake. + model_dtype: Option, +} + +impl ChatRequestProcessor { + /// Create a processor with multimodal support using the effective model + /// dtype reported by the engine. + fn new(backend: DynChatBackend, model_dtype: ModelDtype) -> Self { + Self { + backend, + model_dtype: Some(model_dtype), + } + } + + /// Create a render-only processor that rejects multimodal requests. + pub fn render_only(backend: DynChatBackend) -> Self { + Self { + backend, + model_dtype: None, + } + } + + async fn finalize_rendered_prompt( + &self, + request: &ChatRequest, + rendered: RenderedPrompt, + ) -> Result<(Prompt, Option)> { + match self.model_dtype { + Some(model_dtype) => { + multimodal::finalize_rendered_prompt( + request, + rendered, + self.backend.multimodal_model_info(), + model_dtype, + ) + .await + } + None if !request.has_multimodal() => Ok((rendered.prompt, None)), + None => Err(Error::UnsupportedMultimodalRenderer), + } + } + + /// Prepare one chat request without submitting it to an engine. + pub async fn prepare( + &self, + mut request: ChatRequest, + options: NewChatOutputProcessorOptions<'_>, + ) -> Result<(TextRequest, DynChatOutputProcessor)> { + request.validate()?; + + // Stamp before rendering so render and tokenize count toward TTFT/e2e. + let arrival_time = vllm_llm::current_unix_timestamp_secs(); + let output_processor = self.backend.new_chat_output_processor(&mut request, options)?; + let rendered = self.backend.chat_renderer().render(&request)?; + let reasoning_parser_kwargs = + request + .sampling_params + .structured_outputs + .is_some() + .then(|| ReasoningParserKwargs { + chat_template_kwargs: rendered.effective_template_kwargs.clone(), + }); + let (prompt, mm_features) = self.finalize_rendered_prompt(&request, rendered).await?; + let text_request = TextRequest { + request_id: request.request_id, + prompt, + mm_features, + sampling_params: request.sampling_params, + decode_options: request.decode_options, + intermediate: request.intermediate, + priority: request.priority, + cache_salt: request.cache_salt, + add_special_tokens: request.add_special_tokens, + data_parallel_rank: request.data_parallel_rank, + reasoning_parser_kwargs, + lora_request: request.lora_request, + arrival_time: Some(arrival_time), + }; + Ok((text_request, output_processor)) + } +} + /// Structured chat facade above [`TextLlm`]. /// /// This layer stays above raw text semantics: it takes care of chat-template @@ -95,9 +182,7 @@ pub fn validate_parser_overrides( /// request semantics such as tool calls. pub struct ChatLlm { text: TextLlm, - backend: DynChatBackend, - /// Effective model dtype reported by the engine. - model_dtype: ModelDtype, + processor: ChatRequestProcessor, /// Tool-call parser selection. tool_call_parser: ParserSelection, /// Reasoning parser selection. @@ -112,8 +197,7 @@ impl ChatLlm { Self { text, - backend, - model_dtype, + processor: ChatRequestProcessor::new(backend, model_dtype), tool_call_parser: ParserSelection::Auto, reasoning_parser: ParserSelection::Auto, } @@ -140,7 +224,7 @@ impl ChatLlm { /// Override the effective model dtype used for multimodal tensor encoding. pub fn with_model_dtype(mut self, model_dtype: ModelDtype) -> Self { - self.model_dtype = model_dtype; + self.processor.model_dtype = Some(model_dtype); self } @@ -172,57 +256,23 @@ impl ChatLlm { } /// Render, tokenize, and submit one chat request. - pub async fn chat(&self, mut request: ChatRequest) -> Result { - request.validate()?; - - // Stamp before rendering so render and tokenize count toward TTFT/e2e. - let arrival_time = vllm_llm::current_unix_timestamp_secs(); - - let output_processor = self.backend.new_chat_output_processor( - &mut request, - NewChatOutputProcessorOptions { - tool_call_parser: &self.tool_call_parser, - reasoning_parser: &self.reasoning_parser, - }, - )?; - let rendered = self.backend.chat_renderer().render(&request)?; - let reasoning_parser_kwargs = - request - .sampling_params - .structured_outputs - .is_some() - .then(|| ReasoningParserKwargs { - chat_template_kwargs: rendered.effective_template_kwargs.clone(), - }); - - let (prompt, mm_features) = multimodal::finalize_rendered_prompt( - &request, - rendered, - self.backend.multimodal_model_info(), - self.model_dtype, - ) - .await?; - - let text_request = TextRequest { - request_id: request.request_id.clone(), - prompt, - mm_features, - sampling_params: request.sampling_params, - decode_options: request.decode_options, - intermediate: request.intermediate, - priority: request.priority, - cache_salt: request.cache_salt, - add_special_tokens: request.add_special_tokens, - data_parallel_rank: request.data_parallel_rank, - reasoning_parser_kwargs, - lora_request: request.lora_request, - arrival_time: Some(arrival_time), - }; + pub async fn chat(&self, request: ChatRequest) -> Result { + let (text_request, output_processor) = self + .processor + .prepare( + request, + NewChatOutputProcessorOptions { + tool_call_parser: &self.tool_call_parser, + reasoning_parser: &self.reasoning_parser, + }, + ) + .await?; + let request_id = text_request.request_id.clone(); let decoded_stream = self.text.generate(text_request).await?.map_err(Error::from).boxed(); let structured_stream = output_processor.process(decoded_stream)?; - Ok(ChatEventStream::new(request.request_id, structured_stream)) + Ok(ChatEventStream::new(request_id, structured_stream)) } /// Render through the chat template and tokenize, without submitting to the engine. @@ -233,14 +283,9 @@ impl ChatLlm { pub async fn tokenize_chat(&self, request: ChatRequest) -> Result> { request.validate()?; - let rendered = self.backend.chat_renderer().render(&request)?; - let (prompt, _mm_features) = multimodal::finalize_rendered_prompt( - &request, - rendered, - self.backend.multimodal_model_info(), - self.model_dtype, - ) - .await?; + let rendered = self.processor.backend.chat_renderer().render(&request)?; + let (prompt, _mm_features) = + self.processor.finalize_rendered_prompt(&request, rendered).await?; let tokenizer = self.text.tokenizer(); let token_ids = match prompt { diff --git a/rust/src/server/src/grpc/health.rs b/rust/src/server/src/grpc/health.rs index 70bb81da9f9f..554dfae675f5 100644 --- a/rust/src/server/src/grpc/health.rs +++ b/rust/src/server/src/grpc/health.rs @@ -8,7 +8,7 @@ use tonic_health::ServingStatus; use tonic_health::server::HealthReporter; use tracing::{info, warn}; -use super::GenerateGrpcService; +use super::{ControlGrpcService, GenerateGrpcService}; pub(crate) async fn monitor_health( mut health_reporter: HealthReporter, @@ -16,21 +16,18 @@ pub(crate) async fn monitor_health( shutdown: CancellationToken, ) { let generate_service = GenerateGrpcService::NAME; + let control_service = ControlGrpcService::NAME; let status = ServingStatus::NotServing; let health_event_first = tokio::select! { result = engine_health.wait_for(|healthy| !*healthy) => { match result { Ok(_) => warn!( - generate_service, - overall_service = true, status = ?status, reason = "engine_unhealthy", "marking gRPC health services as not serving" ), Err(error) => warn!( %error, - generate_service, - overall_service = true, status = ?status, reason = "health_channel_closed", "engine health channel closed; marking gRPC health services as not serving" @@ -40,8 +37,6 @@ pub(crate) async fn monitor_health( } _ = shutdown.cancelled() => { info!( - generate_service, - overall_service = true, status = ?status, reason = "server_shutdown", "server shutting down; marking gRPC health services as not serving" @@ -51,20 +46,20 @@ pub(crate) async fn monitor_health( }; health_reporter.set_not_serving::().await; - // Generate is currently the only engine-backed gRPC service, so overall - // server health intentionally mirrors it. + health_reporter.set_not_serving::().await; + // Both gRPC services use the same engine client, so overall server health + // mirrors their shared engine health. health_reporter.set_service_status("", status).await; if health_event_first { shutdown.cancelled().await; info!( - generate_service, - overall_service = true, reason = "server_shutdown", "server shutting down; closing gRPC health watches" ); } health_reporter.clear_service_status(generate_service).await; + health_reporter.clear_service_status(control_service).await; health_reporter.clear_service_status("").await; } diff --git a/rust/src/server/src/grpc/mod.rs b/rust/src/server/src/grpc/mod.rs index 4170e5bb895f..f5d2f347bb3c 100644 --- a/rust/src/server/src/grpc/mod.rs +++ b/rust/src/server/src/grpc/mod.rs @@ -26,8 +26,10 @@ pub mod pb { } pub(crate) use health::monitor_health; +pub use pb::control_server::ControlServer; pub use pb::generate_server::GenerateServer; +pub(crate) type ControlGrpcService = ControlServer; pub(crate) type GenerateGrpcService = GenerateServer; #[cfg(test)] @@ -44,6 +46,36 @@ impl GenerateServiceImpl { } } +/// gRPC control service backed by the shared application state. +pub struct ControlServiceImpl { + state: Arc, +} + +impl ControlServiceImpl { + pub fn new(state: Arc) -> Self { + Self { state } + } +} + +#[tonic::async_trait] +impl pb::control_server::Control for ControlServiceImpl { + async fn abort( + &self, + request: Request, + ) -> Result, Status> { + let request_ids = request.into_inner().request_ids; + if request_ids.is_empty() { + return Ok(Response::new(pb::AbortResponse {})); + } + self.state + .chat + .abort(&request_ids) + .await + .map_err(|error| Status::internal(error.to_report_string()))?; + Ok(Response::new(pb::AbortResponse {})) + } +} + #[tonic::async_trait] impl pb::generate_server::Generate for GenerateServiceImpl { type GenerateStreamStream = diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index 32412dd9ab7e..4cd67f4a4b36 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -38,8 +38,9 @@ use vllm_tokenizer::test_utils::TestTokenizer; use zeromq::prelude::{SocketRecv, SocketSend}; use zeromq::{DealerSocket, PushSocket, ZmqMessage}; +use super::pb::control_client::ControlClient; use super::pb::generate_client::GenerateClient; -use super::{GenerateServer, GenerateServiceImpl, pb}; +use super::{ControlServer, ControlServiceImpl, GenerateServer, GenerateServiceImpl, pb}; use crate::listener::{Listener, MaybeTlsListener}; use crate::state::AppState; use crate::tls; @@ -153,10 +154,6 @@ async fn recv_engine_message(dealer: &mut DealerSocket) -> Vec { dealer.recv().await.expect("recv engine message").into_vec() } -fn test_llm(client: EngineCoreClient) -> Llm { - Llm::new(client).with_request_id_randomization(false) -} - #[derive(Clone, Debug)] struct FakeTextBackend; @@ -206,6 +203,7 @@ async fn setup_grpc_service( output_specs: Vec<(Vec, Option)>, ) -> ( GenerateServer, + ControlServer, tokio::sync::watch::Receiver, MockEngineTask, ) { @@ -243,12 +241,13 @@ async fn setup_grpc_service( let engine_health = client.subscribe_health(); let chat = ChatLlm::from_shared_backend( - test_llm(client), + Llm::new(client), Arc::new(FakeTextBackend) as Arc, ); let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat)); ( - GenerateServer::new(GenerateServiceImpl::new(state)), + GenerateServer::new(GenerateServiceImpl::new(state.clone())), + ControlServer::new(ControlServiceImpl::new(state)), engine_health, engine_task, ) @@ -264,9 +263,11 @@ async fn grpc_test_server( tokio::task::JoinHandle<()>, MockEngineTask, ) { - let (svc, engine_health, engine_task) = setup_grpc_service(engine_id, output_specs).await; + let (generate_service, control_service, engine_health, engine_task) = + setup_grpc_service(engine_id, output_specs).await; let (channel, server_task) = start_grpc_test_server( - svc, + generate_service, + control_service, engine_health, tokio_util::sync::CancellationToken::new(), ) @@ -276,11 +277,13 @@ async fn grpc_test_server( async fn start_grpc_test_server( generate_service: GenerateServer, + control_service: ControlServer, engine_health: tokio::sync::watch::Receiver, shutdown: tokio_util::sync::CancellationToken, ) -> (Channel, tokio::task::JoinHandle<()>) { let (health_reporter, health_service) = health_reporter(); health_reporter.set_serving::>().await; + health_reporter.set_serving::>().await; let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener"); let addr = listener.local_addr().expect("local addr"); @@ -289,6 +292,7 @@ async fn start_grpc_test_server( let incoming = MaybeTlsListener::plain(Listener::Tcp(listener)); let server = TonicServer::builder() .add_service(health_service) + .add_service(control_service) .add_service(generate_service) .serve_with_incoming_shutdown(incoming, shutdown.clone().cancelled_owned()); let health_monitor = @@ -319,7 +323,8 @@ async fn grpc_tls_test_server( certs: &TestCerts, cert_reqs: i32, ) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) { - let (svc, _engine_health, engine_task) = setup_grpc_service(engine_id, output_specs).await; + let (generate_service, control_service, _engine_health, engine_task) = + setup_grpc_service(engine_id, output_specs).await; let context = tls::build_grpc_server_config(&server_tls(certs, cert_reqs)) .expect("build grpc tls config"); @@ -329,7 +334,8 @@ async fn grpc_tls_test_server( let server_task = tokio::spawn(async move { let incoming = MaybeTlsListener::tls(Listener::Tcp(listener), context); TonicServer::builder() - .add_service(svc) + .add_service(control_service) + .add_service(generate_service) .serve_with_incoming(incoming) .await .expect("grpc tls server"); @@ -409,7 +415,7 @@ async fn grpc_server_with_keepalive( engine_id: impl Into, keepalive: Option, ) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) { - let (svc, _engine_health, engine_task) = + let (generate_service, control_service, _engine_health, engine_task) = setup_grpc_service(engine_id, default_stream_output_specs()).await; let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener"); @@ -425,7 +431,8 @@ async fn grpc_server_with_keepalive( let server_task = tokio::spawn(async move { let incoming = MaybeTlsListener::plain(Listener::Tcp(listener)); builder - .add_service(svc) + .add_service(control_service) + .add_service(generate_service) .serve_with_incoming(incoming) .await .expect("grpc server"); @@ -1073,14 +1080,106 @@ async fn grpc_without_keepalive_keeps_unresponsive_connection_open() { server_task.abort(); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn control_abort_resolves_external_id_and_empty_is_noop() { + let (generate_service, control_service, engine_health, engine_task) = + setup_grpc_service(b"engine-grpc-abort-active", vec![(vec![b'h' as u32], None)]).await; + let (channel, server_task) = start_grpc_test_server( + generate_service, + control_service, + engine_health, + tokio_util::sync::CancellationToken::new(), + ) + .await; + let mut generate_client = GenerateClient::new(channel.clone()); + let mut control_client = ControlClient::new(channel); + let request_id = "test-abort-active"; + + let mut stream = generate_client + .generate_stream(pb::GenerateRequest { + request_id: request_id.to_string(), + model: "test-model".to_string(), + prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())), + stopping: Some(pb::StoppingCriteria { + max_new_tokens: 10, + ..Default::default() + }), + ..Default::default() + }) + .await + .expect("start generation") + .into_inner(); + + loop { + let response = tokio::time::timeout(Duration::from_secs(2), stream.message()) + .await + .expect("timed out waiting for active generation output") + .expect("read active generation output") + .expect("generation ended before producing output"); + if let Some(output) = response.outputs { + assert!( + output.finish_info.is_none(), + "generation finished before abort behavior was exercised" + ); + break; + } + } + + control_client + .abort(pb::AbortRequest::default()) + .await + .expect("empty abort should be a no-op"); + assert!( + tokio::time::timeout(Duration::from_millis(100), stream.message()) + .await + .is_err(), + "empty abort unexpectedly ended the active generation" + ); + + control_client + .abort(pb::AbortRequest { + request_ids: vec![ + request_id.to_string(), + request_id.to_string(), + "unknown".to_string(), + ], + }) + .await + .expect("abort active generation"); + + let finish_reason = loop { + let response = tokio::time::timeout(Duration::from_secs(2), stream.message()) + .await + .expect("timed out waiting for aborted generation") + .expect("read aborted generation") + .expect("generation ended without an aborted response"); + if let Some(finish_info) = response.outputs.and_then(|output| output.finish_info) { + break finish_info.finish_reason; + } + }; + assert_eq!(finish_reason, pb::finish_info::FinishReason::Aborted as i32); + + control_client + .abort(pb::AbortRequest { + request_ids: vec![request_id.to_string()], + }) + .await + .expect("repeated abort should be idempotent"); + + engine_task.await.expect("mock engine task"); + server_task.abort(); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() { - let (generate_service, _connected_engine_health, _engine_task) = + let (generate_service, control_service, _connected_engine_health, _engine_task) = setup_grpc_service(b"engine-grpc-health-failure", default_stream_output_specs()).await; let (engine_health_tx, engine_health) = tokio::sync::watch::channel(true); let (channel, server_task) = start_grpc_test_server( generate_service, + control_service, engine_health, tokio_util::sync::CancellationToken::new(), ) @@ -1088,7 +1187,7 @@ async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() let mut health_client = HealthClient::new(channel); let mut health_streams = Vec::new(); - for service in ["vllm.Generate", ""] { + for service in ["vllm.Generate", "vllm.Control", ""] { let service_label = if service.is_empty() { "overall" } else { @@ -1143,14 +1242,19 @@ async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn grpc_health_watch_closes_on_graceful_shutdown() { - let (generate_service, engine_health, _engine_task) = setup_grpc_service( + let (generate_service, control_service, engine_health, _engine_task) = setup_grpc_service( b"engine-grpc-health-shutdown", default_stream_output_specs(), ) .await; let shutdown = tokio_util::sync::CancellationToken::new(); - let (channel, server_task) = - start_grpc_test_server(generate_service, engine_health, shutdown.clone()).await; + let (channel, server_task) = start_grpc_test_server( + generate_service, + control_service, + engine_health, + shutdown.clone(), + ) + .await; let mut health_client = HealthClient::new(channel); let mut stream = health_client .watch(HealthCheckRequest { diff --git a/rust/src/server/src/lib.rs b/rust/src/server/src/lib.rs index 220da778af8e..1f3468fe46d6 100644 --- a/rust/src/server/src/lib.rs +++ b/rust/src/server/src/lib.rs @@ -207,6 +207,9 @@ where let (health_reporter, health_service) = health_reporter(); let engine_health = state.engine_core_client().subscribe_health(); health_reporter.set_serving::().await; + health_reporter.set_serving::().await; + let control_service = + grpc::ControlGrpcService::new(grpc::ControlServiceImpl::new(state.clone())); let generate_service = grpc::GenerateGrpcService::new(grpc::GenerateServiceImpl::new(state.clone())); let svc = TonicServer::builder() @@ -214,6 +217,7 @@ where .http2_keepalive_timeout(Some(GRPC_KEEPALIVE_TIMEOUT)) .layer(middleware::request_runtime_layer(state.clone())) .add_service(health_service) + .add_service(control_service) .add_service(generate_service); info!(%addr, tls = grpc_tls.is_some(), "starting gRPC server"); Some((grpc_listener, svc, grpc_tls, health_reporter, engine_health)) diff --git a/rust/src/text/src/lib.rs b/rust/src/text/src/lib.rs index 1ccd53127acd..b00155999e8b 100644 --- a/rust/src/text/src/lib.rs +++ b/rust/src/text/src/lib.rs @@ -38,22 +38,90 @@ trait_set! { pub trait TextOutputStream = Stream> + Send + 'static; } -/// Raw text facade above [`Llm`]. -/// -/// This layer stays below chat semantics: prompt text or prompt token IDs flow -/// in, decoded text deltas and terminal metadata flow out. -pub struct TextLlm { - /// Generate-only client owned by this text facade. - llm: Llm, +/// Text request preparation shared by inference and render-only frontends. +pub struct TextRequestProcessor { /// Tokenizer/model metadata backend responsible for prompt encode/decode /// and sampling hints. backend: DynTextBackend, /// Runtime context window size reported by the engine startup handshake. + /// Render-only frontends supply the downstream engine's effective value. max_model_len: u32, /// Maximum number of top log probabilities accepted by this text facade. max_logprobs: i32, } +impl TextRequestProcessor { + /// Create a processor with the effective model context length. + pub fn new(backend: DynTextBackend, max_model_len: u32) -> Self { + Self { + backend, + max_model_len, + max_logprobs: SamplingLimits::DEFAULT_MAX_LOGPROBS, + } + } + + /// Override the maximum accepted logprobs count. + pub fn with_max_logprobs(mut self, max_logprobs: Option) -> Self { + if let Some(max_logprobs) = max_logprobs { + self.max_logprobs = max_logprobs; + } + self + } + + /// Return the tokenizer used by this processor. + pub fn tokenizer(&self) -> DynTokenizer { + self.backend.tokenizer() + } + + /// Return the effective model context length. + pub fn max_model_len(&self) -> u32 { + self.max_model_len + } + + /// Tokenize and lower one request without submitting it to an engine. + pub fn prepare(&self, mut request: TextRequest) -> Result { + request.validate()?; + + if request.arrival_time.is_none() { + request.arrival_time = Some(vllm_llm::current_unix_timestamp_secs()); + } + + let tokenizer = self.backend.tokenizer(); + let prompt_token_ids = match take(&mut request.prompt) { + Prompt::Text(text) => tokenizer.encode(&text, request.add_special_tokens)?, + // Pre-tokenized prompts are the main completions-side escape hatch that lets benchmark + // and infra workloads bypass chat rendering and tokenizer overhead entirely. + Prompt::TokenIds(token_ids) => token_ids, + }; + let sampling_hints = self.backend.sampling_hints()?; + let sampling_limits = SamplingLimits { + max_model_len: self.max_model_len, + max_logprobs: self.max_logprobs, + model_vocab_size: self.backend.model_vocab_size(), + tokenizer_vocab_size: self.backend.tokenizer_vocab_size(), + }; + + lower_text_request( + request, + prompt_token_ids, + sampling_hints, + sampling_limits, + tokenizer.as_ref(), + ) + } +} + +/// Raw text facade above [`Llm`]. +/// +/// This layer stays below chat semantics: prompt text or prompt token IDs flow +/// in, decoded text deltas and terminal metadata flow out. +pub struct TextLlm { + /// Generate-only client owned by this text facade. + llm: Llm, + /// Shared engine-free request preparation. + processor: TextRequestProcessor, +} + impl TextLlm { /// Create a new text-generation facade from a shared LLM client plus a text /// backend. @@ -64,23 +132,19 @@ impl TextLlm { Self { llm, - backend, - max_model_len, - max_logprobs: SamplingLimits::DEFAULT_MAX_LOGPROBS, + processor: TextRequestProcessor::new(backend, max_model_len), } } /// Override the maximum accepted logprobs count. pub fn with_max_logprobs(mut self, max_logprobs: Option) -> Self { - if let Some(max_logprobs) = max_logprobs { - self.max_logprobs = max_logprobs; - } + self.processor = self.processor.with_max_logprobs(max_logprobs); self } /// Return the backend model ID. pub fn model_id(&self) -> &str { - self.backend.model_id() + self.processor.backend.model_id() } /// Expose the underlying engine-core client for low-level utility/admin @@ -91,19 +155,19 @@ impl TextLlm { /// Return the tokenizer used by this text backend. pub fn tokenizer(&self) -> DynTokenizer { - self.backend.tokenizer() + self.processor.tokenizer() } /// Tokenizer vocabulary size (the number of tokens the tokenizer knows), /// used to bound `allowed_token_ids` like the Python frontend `len(tokenizer)`. pub fn tokenizer_vocab_size(&self) -> usize { - self.backend.tokenizer_vocab_size() + self.processor.backend.tokenizer_vocab_size() } /// Model vocabulary size from the model config, used to bound generated /// token IDs and logits-domain sampling controls. pub fn model_vocab_size(&self) -> usize { - self.backend.model_vocab_size() + self.processor.backend.model_vocab_size() } /// Tokenize if needed, lower to a generate request, and return the raw @@ -117,7 +181,7 @@ impl TextLlm { /// incrementally decoded text. pub async fn generate(&self, request: TextRequest) -> Result { let (text_request, raw_stream) = self.generate_inner(request).await?; - let tokenizer = self.backend.tokenizer(); + let tokenizer = self.processor.tokenizer(); let decoded_stream = output::decoded_text_event_stream( text_request.request_id, tokenizer, @@ -131,40 +195,12 @@ impl TextLlm { async fn generate_inner( &self, - mut request: TextRequest, + request: TextRequest, ) -> Result<(TextRequest, GenerateOutputStream)> { - request.validate()?; - - if request.arrival_time.is_none() { - request.arrival_time = Some(vllm_llm::current_unix_timestamp_secs()); - } - - let tokenizer = self.backend.tokenizer(); - let prompt_token_ids = match take(&mut request.prompt) { - Prompt::Text(text) => tokenizer.encode(&text, request.add_special_tokens)?, - // Pre-tokenized prompts are the main completions-side escape hatch that lets benchmark - // and infra workloads bypass chat rendering and tokenizer overhead entirely. - Prompt::TokenIds(token_ids) => token_ids, - }; - - let sampling_hints = self.backend.sampling_hints()?; - let sampling_limits = SamplingLimits { - max_model_len: self.max_model_len, - max_logprobs: self.max_logprobs, - model_vocab_size: self.backend.model_vocab_size(), - tokenizer_vocab_size: self.backend.tokenizer_vocab_size(), - }; - let PreparedTextRequest { text_request, generate_request, - } = lower_text_request( - request, - prompt_token_ids, - sampling_hints, - sampling_limits, - &*tokenizer, - )?; + } = self.processor.prepare(request)?; let raw_stream = self.llm.generate(generate_request).await?; Ok((text_request, raw_stream)) diff --git a/tests/config/test_multimodal_config.py b/tests/config/test_multimodal_config.py index 9720d84672fd..5260d7a40f7c 100644 --- a/tests/config/test_multimodal_config.py +++ b/tests/config/test_multimodal_config.py @@ -43,6 +43,15 @@ def test_language_model_only_affects_model_hash(): assert base_hash != lm_only_hash +@pytest.mark.parametrize("backend_arg", ["video_backend", "backend"]) +def test_use_gpu_video_backend_from_media_io_kwargs(backend_arg: str): + config = MultiModalConfig( + media_io_kwargs={"video": {backend_arg: "pynvvideocodec"}} + ) + + assert config.use_gpu_video_backend() + + def test_mm_encoder_fp8_scale_path_requires_fp8(): with pytest.raises(ValueError, match="mm_encoder_attn_dtype"): MultiModalConfig(mm_encoder_fp8_scale_path="/tmp/scales.json") diff --git a/tests/entrypoints/pooling/embed/test_io_processor.py b/tests/entrypoints/pooling/embed/test_io_processor.py index 5a7a8aab2a60..f06e43d143ea 100644 --- a/tests/entrypoints/pooling/embed/test_io_processor.py +++ b/tests/entrypoints/pooling/embed/test_io_processor.py @@ -9,8 +9,6 @@ from vllm import PoolingParams from vllm.entrypoints.pooling.embed.io_processor import EmbedIOProcessor from vllm.entrypoints.pooling.embed.protocol import ( - CohereEmbedContent, - CohereEmbedInput, CohereEmbedRequest, EmbeddingBatchChatInputRequest, EmbeddingBatchChatRequest, @@ -19,7 +17,10 @@ EmbeddingCompletionRequest, EmbeddingRequest, ) -from vllm.entrypoints.pooling.typing import PoolingServeContext +from vllm.entrypoints.pooling.typing import ( + PoolingEngineInput, + PoolingServeContext, +) from vllm.outputs import PoolingOutput, PoolingRequestOutput @@ -410,6 +411,7 @@ class _FakeModelConfig: def _make_handler(cls): handler = object.__new__(EmbedIOProcessor) handler.model_config = cls._FakeModelConfig() + handler.enable_chunked_processing = True return handler @staticmethod @@ -421,15 +423,29 @@ def _make_context() -> PoolingServeContext[EmbeddingCompletionRequest]: } ) assert isinstance(request, EmbeddingCompletionRequest) + pooling_params = PoolingParams() return PoolingServeContext( request=request, - pooling_params=PoolingParams(), + pooling_params=pooling_params, model_name="test", request_id="embd-client-prompt-999-chunk-888", engine_inputs=[ - {"prompt_token_ids": [0, 1, 2, 3, 4]}, - {"prompt_token_ids": [10, 11]}, + PoolingEngineInput( + prompts={"prompt_token_ids": [0, 1, 2, 3, 4]}, + params=pooling_params, + lora_requests=None, + priorities=0, + ), + PoolingEngineInput( + prompts={"prompt_token_ids": [10, 11]}, + params=pooling_params, + lora_requests=None, + priorities=0, + ), ], + lora_request=None, + priorities=0, + prompt_extras=None, ) @staticmethod @@ -450,7 +466,7 @@ def test_aggregation_uses_metadata_not_request_id_parsing(self): handler = self._make_handler() ctx = self._make_context() - handler._pre_process_chunked(ctx) + handler.maybe_pre_process_chunked(ctx) assert ctx.prompt_request_ids == [ "embd-client-prompt-999-chunk-888-prompt-0-chunk-0", @@ -488,227 +504,3 @@ def test_aggregation_uses_metadata_not_request_id_parsing(self): ctx.final_res_batch[1].outputs.data, torch.tensor([9.0, 9.0]), ) - - -class TestPreProcessCohereOnline: - """Unit tests for EmbedIOProcessor._pre_process_cohere_online.""" - - @staticmethod - def _make_context(**request_kwargs) -> PoolingServeContext[CohereEmbedRequest]: - return PoolingServeContext( - request=CohereEmbedRequest(model="test", **request_kwargs), - pooling_params=PoolingParams(), - model_name="test", - request_id="embd-test", - ) - - @staticmethod - def _make_handler(): - handler = object.__new__(EmbedIOProcessor) - handler._validate_input_type = lambda _input_type: None - return handler - - def test_text_only_without_task_prefix_uses_completion_path(self): - handler = self._make_handler() - ctx = self._make_context(texts=["hello"]) - calls: list[tuple[str, object]] = [] - - def preprocess_cmpl_online(request, prompt_input, prompt_embeds): - calls.append(("completion", prompt_input)) - return ["completion"] - - handler._get_task_instruction_prefix = lambda _input_type: None - handler._has_chat_template = lambda: False - handler._preprocess_cmpl_online = preprocess_cmpl_online - handler._batch_render_chat = lambda *_args, **_kwargs: pytest.fail( - "text-only request should not require chat rendering" - ) - - handler._pre_process_cohere_online(ctx) - - assert ctx.engine_inputs == ["completion"] - assert calls == [("completion", ["hello"])] - - def test_text_only_falls_back_to_prefixed_completion_without_template(self): - handler = self._make_handler() - ctx = self._make_context(texts=["hello"], input_type="query") - calls: list[tuple[str, object]] = [] - - def preprocess_cmpl(request, prompt_input, prompt_embeds): - calls.append(("completion", prompt_input)) - return ["fallback"] - - handler._get_task_instruction_prefix = lambda _input_type: "query: " - handler._has_chat_template = lambda: False - handler._batch_render_chat = lambda *_args, **_kwargs: pytest.fail( - "chat rendering should be skipped without a template" - ) - handler._preprocess_cmpl_online = preprocess_cmpl - - handler._pre_process_cohere_online(ctx) - - assert ctx.engine_inputs == ["fallback"] - assert calls == [("completion", ["query: hello"])] - - def test_text_only_with_template_uses_chat_path(self): - handler = self._make_handler() - ctx = self._make_context(texts=["hello"], input_type="query") - calls: list[tuple[str, object]] = [] - - def batch_render_chat( - request, - all_messages, - truncate_prompt_tokens, - truncation_side, - ): - calls.append( - ( - "chat", - { - "request": request, - "all_messages": all_messages, - "truncate_prompt_tokens": truncate_prompt_tokens, - "truncation_side": truncation_side, - }, - ) - ) - return ["chat"] - - handler._get_task_instruction_prefix = lambda _input_type: "query: " - handler._has_chat_template = lambda: True - handler._batch_render_chat = batch_render_chat - handler._preprocess_cmpl_online = lambda *_args, **_kwargs: pytest.fail( - "completion path should be skipped when a template exists" - ) - - handler._pre_process_cohere_online(ctx) - - assert ctx.engine_inputs == ["chat"] - assert calls == [ - ( - "chat", - { - "request": ctx.request, - "all_messages": [ - handler._mixed_input_to_messages( - CohereEmbedInput( - content=[CohereEmbedContent(type="text", text="hello")] - ), - task_prefix="query: ", - ) - ], - "truncate_prompt_tokens": -1, - "truncation_side": None, - }, - ) - ] - - -class TestPreProcessOpenAIEmbeddingChatOnline: - """Unit tests for OpenAI embedding chat preprocessing.""" - - class _FakeModelConfig: - max_model_len = 128 - encoder_config: dict[str, object] = {} - pooler_config = None - multimodal_config = None - is_encoder_decoder = False - - class _FakeRenderer: - tokenizer = object() - - def __init__(self): - self.calls = [] - - def render_chat( - self, - all_messages, - chat_params, - tok_params, - prompt_extras=None, - ): - self.calls.append( - { - "all_messages": all_messages, - "chat_params": chat_params, - "tok_params": tok_params, - "prompt_extras": prompt_extras, - } - ) - return all_messages, [ - {"prompt_token_ids": [index]} for index, _ in enumerate(all_messages) - ] - - @classmethod - def _make_handler(cls, renderer): - handler = object.__new__(EmbedIOProcessor) - handler.renderer = renderer - handler.model_config = cls._FakeModelConfig() - handler.chat_template = "template" - handler.chat_template_content_format = "auto" - handler.trust_request_chat_template = False - handler.enable_chunked_processing = False - return handler - - @staticmethod - def _make_context( - request: ( - EmbeddingChatRequest - | EmbeddingBatchChatRequest - | EmbeddingChatInputRequest - | EmbeddingBatchChatInputRequest - ), - ) -> PoolingServeContext[ - EmbeddingChatRequest - | EmbeddingBatchChatRequest - | EmbeddingChatInputRequest - | EmbeddingBatchChatInputRequest - ]: - return PoolingServeContext( - request=request, - pooling_params=PoolingParams(), - model_name="test", - request_id="embd-test", - ) - - def test_chat_template_kwargs_forwarded_for_batched_input_messages(self): - request = TypeAdapter(EmbeddingRequest).validate_python( - { - "model": "test", - "input": [ - [{"role": "user", "content": "hello"}], - [{"role": "user", "content": "goodbye"}], - ], - "add_generation_prompt": True, - "chat_template_kwargs": {"instruction": "Represent the query: "}, - "mm_processor_kwargs": {"max_pixels": 1}, - "cache_salt": "salt", - } - ) - assert isinstance(request, EmbeddingBatchChatInputRequest) - - renderer = self._FakeRenderer() - handler = self._make_handler(renderer) - ctx = self._make_context(request) - - handler.pre_process_online(ctx) - - assert ctx.engine_inputs == [ - {"prompt_token_ids": [0]}, - {"prompt_token_ids": [1]}, - ] - assert len(renderer.calls) == 1 - - call = renderer.calls[0] - assert call["all_messages"] == request.messages - assert call["prompt_extras"] == { - "mm_processor_kwargs": {"max_pixels": 1}, - "cache_salt": "salt", - } - - chat_template_kwargs = call["chat_params"].chat_template_kwargs - assert chat_template_kwargs["instruction"] == "Represent the query: " - assert chat_template_kwargs["add_generation_prompt"] is True - assert chat_template_kwargs["continue_final_message"] is False - assert "tools" not in chat_template_kwargs - assert chat_template_kwargs["tokenize"] is False diff --git a/tests/entrypoints/pooling/scoring/test_cross_encoder_offline.py b/tests/entrypoints/pooling/scoring/test_cross_encoder_offline.py index df79a387afd2..56e83de3f74f 100644 --- a/tests/entrypoints/pooling/scoring/test_cross_encoder_offline.py +++ b/tests/entrypoints/pooling/scoring/test_cross_encoder_offline.py @@ -2,7 +2,6 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import weakref -from types import SimpleNamespace import pytest import torch @@ -10,10 +9,7 @@ from tests.models.utils import softmax from vllm import LLM, PoolingParams from vllm.distributed import cleanup_dist_env_and_memory -from vllm.entrypoints.pooling.scoring.io_processor import CrossEncoderIOProcessor -from vllm.entrypoints.pooling.scoring.typing import ScoringData from vllm.platforms import current_platform -from vllm.renderers import TokenizeParams MODEL_NAME = "tomaarsen/Qwen3-Reranker-0.6B-seq-cls" PROMPT = "The chef prepared a delicious meal." @@ -145,45 +141,6 @@ def test_max_tokens_per_doc(llm: LLM): assert with_limit_tokens < no_limit_tokens -def test_token_type_ids_follow_post_tokenization(): - processor = object.__new__(CrossEncoderIOProcessor) - processor.tokenizer = SimpleNamespace(truncation_side="right", pad_token_id=-1) - processor.renderer = SimpleNamespace(process_for_engine=lambda prompt, _: prompt) - processor.model_config = None - processor.get_score_prompt = lambda **_: ( - "", - { - "prompt_token_ids": list(range(32)), - "token_type_ids": [0] * 16 + [1] * 16, - }, - ) - - engine_inputs, pooling_params = processor._pre_process( - ScoringData(data_1=["query"], data_2=["document"]), - TokenizeParams( - max_total_tokens=None, - truncate_prompt_tokens=16, - truncation_side="left", - ), - PoolingParams(task="classify", extra_kwargs={"cache_salt": "salt"}), - ) - - assert engine_inputs[0]["prompt_token_ids"] == list(range(16, 32)) - assert pooling_params[0].extra_kwargs == { - "cache_salt": "salt", - "compressed_token_type_ids": 0, - } - - engine_inputs, pooling_params = processor._pre_process( - ScoringData(data_1=["query"], data_2=["document"]), - TokenizeParams(max_total_tokens=None, pad_prompt_tokens=40), - PoolingParams(task="classify"), - ) - - assert engine_inputs[0]["prompt_token_ids"] == list(range(32)) + [-1] * 8 - assert pooling_params[0].extra_kwargs == {"compressed_token_type_ids": 16} - - def test_pooling_params(llm: LLM): def get_outputs(use_activation): outputs = llm.score( diff --git a/tests/entrypoints/serve/lora/test_serving_models.py b/tests/entrypoints/serve/lora/test_serving_models.py index 658d004580a2..e75cbe908b1d 100644 --- a/tests/entrypoints/serve/lora/test_serving_models.py +++ b/tests/entrypoints/serve/lora/test_serving_models.py @@ -161,7 +161,9 @@ def _make_pooling_serving(lora_name: str) -> _ConcretePoolingServing: return serving -def _make_pooling_ctx(model_name: str) -> PoolingServeContext: +def _make_pooling_ctx( + model_name: str, serving: PoolingBaseServing +) -> PoolingServeContext: mock_request = MagicMock() mock_request.model = model_name return PoolingServeContext( @@ -169,6 +171,9 @@ def _make_pooling_ctx(model_name: str) -> PoolingServeContext: model_name=MODEL_NAME, request_id="test-id", pooling_params=PoolingParams(), + lora_request=serving._maybe_get_adapters(mock_request), + priorities=0, + prompt_extras=None, ) @@ -176,9 +181,7 @@ def test_pooling_maybe_get_adapters_lora_name_sets_lora_request(): """LoRA adapter name must populate ctx.lora_request without raising.""" lora_name = "bot-embed-lora" serving = _make_pooling_serving(lora_name) - ctx = _make_pooling_ctx(lora_name) - - ctx.lora_request = serving._maybe_get_adapters(ctx.request) + ctx = _make_pooling_ctx(lora_name, serving) assert ctx.lora_request is not None assert ctx.lora_request.lora_name == lora_name @@ -187,7 +190,6 @@ def test_pooling_maybe_get_adapters_lora_name_sets_lora_request(): def test_pooling_maybe_get_adapters_unknown_model_raises(): """An unrecognised model name must still raise VLLMNotFoundError.""" serving = _make_pooling_serving("some-lora") - ctx = _make_pooling_ctx("unknown-model") with pytest.raises(VLLMNotFoundError): - serving._maybe_get_adapters(ctx.request) + _make_pooling_ctx("unknown-model", serving) diff --git a/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml index 13d71af20e6e..14b4f2c0d1da 100644 --- a/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml +++ b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml @@ -6,6 +6,7 @@ max_concurrency: 100 server_args: >- --enforce-eager --max-model-len 4096 + --max-num-batched-tokens 32768 --safetensors-load-strategy prefetch --moe-backend flashinfer_cutlass --prefill-context-parallel-size 4 diff --git a/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml index b21d0b12f02f..1a96afe46d71 100644 --- a/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml +++ b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml @@ -6,6 +6,7 @@ max_concurrency: 100 server_args: >- --enforce-eager --max-model-len 4096 + --max-num-batched-tokens 32768 --safetensors-load-strategy prefetch --moe-backend flashinfer_cutlass --tensor-parallel-size 2 diff --git a/tests/kernels/attention/test_attention_selector.py b/tests/kernels/attention/test_attention_selector.py index 1e85f76b64c3..587f75a38d6d 100644 --- a/tests/kernels/attention/test_attention_selector.py +++ b/tests/kernels/attention/test_attention_selector.py @@ -546,8 +546,11 @@ def test_flash_attn_accepts_handled_fp8_variants( ): """FlashAttentionBackend must accept the two fp8 dtypes it can actually handle: 'fp8' (alias for fp8_e4m3fn) and 'fp8_e4m3'.""" - import vllm.v1.attention.backends.flash_attn as fa_mod + import vllm.v1.attention.backends.fa_utils as fa_utils_mod from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend - monkeypatch.setattr(fa_mod.current_platform, "is_xpu", lambda: True) + # The fp8 decision is made in fa_utils, using its own current_platform + # binding, so patch is_xpu there (not on flash_attn's) to stay robust to + # import order across earlier tests that patch vllm.platforms.current_platform. + monkeypatch.setattr(fa_utils_mod.current_platform, "is_xpu", lambda: True) assert FlashAttentionBackend.supports_kv_cache_dtype(kv_cache_dtype) diff --git a/tests/kernels/attention/test_merge_attn_states.py b/tests/kernels/attention/test_merge_attn_states.py index 40af84887a99..d5c9a1f47471 100644 --- a/tests/kernels/attention/test_merge_attn_states.py +++ b/tests/kernels/attention/test_merge_attn_states.py @@ -11,6 +11,9 @@ scaled_fp8_quant, ) from vllm.platforms import current_platform +from vllm.v1.attention.ops.triton_merge_attn_states import ( + mask_empty_context, +) from vllm.v1.attention.ops.triton_merge_attn_states import ( merge_attn_states as merge_attn_states_triton, ) @@ -73,6 +76,59 @@ def merge_attn_states_torch( all_case_info: list[tuple] = [] +def test_mask_empty_context() -> None: + query_lens = torch.tensor([2] + [1] * 31 + [131, 1], dtype=torch.int32) + query_start_loc = torch.cat( + (torch.zeros(1, dtype=torch.int32), query_lens.cumsum(0)) + ).cuda() + context_lens = torch.tensor([4] * 32 + [0, 3], dtype=torch.int32) + context_start_loc = torch.cat( + (torch.zeros(1, dtype=torch.int32), context_lens.cumsum(0)) + ).cuda() + num_heads, num_tokens, head_dim = 4, 165, 16 + lse = torch.randn(num_heads, num_tokens, device="cuda") + output = torch.randn(num_tokens, num_heads, head_dim, device="cuda") + # Empty-context rows carry undefined (possibly non-finite) attention output. + output[33:164] = float("nan") + + expected_lse = lse.clone() + expected_lse[:, 33:164] = float("-inf") + expected_output = output.clone() + expected_output[33:164] = 0.0 + + mask_empty_context(lse, output, query_start_loc, context_start_loc) + + torch.testing.assert_close(lse, expected_lse) + torch.testing.assert_close(output, expected_output) + + +@pytest.mark.parametrize("merge_fn", [merge_attn_states_cuda, merge_attn_states_triton]) +@pytest.mark.parametrize("output_dtype", [torch.float32, torch.half, torch.bfloat16]) +def test_merge_attn_states_both_empty(merge_fn, output_dtype) -> None: + """When a token is empty on both sides (both LSE -inf), the 0/0 softmax + scales must not surface as NaN in the merged output.""" + num_tokens, num_heads, head_size = 6, 8, 128 + prefix_output = torch.zeros( + num_tokens, num_heads, head_size, device="cuda", dtype=output_dtype + ) + prefix_lse = torch.randn(num_heads, num_tokens, device="cuda") + suffix_output = torch.zeros( + num_tokens, num_heads, head_size, device="cuda", dtype=output_dtype + ) + suffix_lse = torch.randn(num_heads, num_tokens, device="cuda") + + # Tokens 2 and 3 are empty on both sides (mask_empty_context already zeroed + # their outputs and set both LSEs to -inf). + empty = slice(2, 4) + prefix_lse[:, empty] = float("-inf") + suffix_lse[:, empty] = float("-inf") + + output = torch.empty_like(prefix_output) + merge_fn(output, prefix_output, prefix_lse, suffix_output, suffix_lse) + + assert not output.isnan().any() + + def generate_markdown_table(): global all_case_info table_header = ( diff --git a/tests/model_executor/model_loader/test_mtp_validation.py b/tests/model_executor/model_loader/test_mtp_validation.py new file mode 100644 index 000000000000..ecccaa0cd74b --- /dev/null +++ b/tests/model_executor/model_loader/test_mtp_validation.py @@ -0,0 +1,19 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest + +from vllm.model_executor.model_loader.mtp_validation import ( + disable_mtp_completeness_check, + is_mtp_completeness_check_enabled, +) + + +def test_disable_mtp_completeness_check_is_scoped(): + assert is_mtp_completeness_check_enabled() + + with pytest.raises(RuntimeError), disable_mtp_completeness_check(): + assert not is_mtp_completeness_check_enabled() + raise RuntimeError + + assert is_mtp_completeness_check_enabled() diff --git a/tests/multimodal/test_gpu_ipc_memory.py b/tests/multimodal/test_gpu_ipc_memory.py index bc6bf031fec0..5c0f263af901 100644 --- a/tests/multimodal/test_gpu_ipc_memory.py +++ b/tests/multimodal/test_gpu_ipc_memory.py @@ -6,15 +6,43 @@ import pytest +import vllm.config.multimodal as multimodal_config_module +from vllm.config.multimodal import MultiModalConfig from vllm.multimodal.gpu_ipc_memory import ( MultiModalGPUMemoryPool, get_mm_gpu_ipc_pool, maybe_init_mm_gpu_ipc_pool, + reserve_mm_ipc_gpu_memory, set_mm_gpu_ipc_pool, ) +from vllm.multimodal.video import ( + PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES, + PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES, + PYNVVIDEOCODEC_MAX_RETAINED_DECODERS, + PYNVVIDEOCODEC_VIDEO_BACKEND, +) from vllm.utils.mem_constants import GiB_bytes +def _mm_config( + *, + mm_ipc_gpu_memory_gb: float = 0, + video_backend: str | None = None, +) -> MultiModalConfig: + video_kwargs = {} if video_backend is None else {"video_backend": video_backend} + return MultiModalConfig( + mm_ipc_gpu_memory_gb=mm_ipc_gpu_memory_gb, + media_io_kwargs={"video": video_kwargs} if video_kwargs else {}, + ) + + +def _pynvvideocodec_decoder_budget(api_process_count: int = 1) -> int: + return api_process_count * ( + PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES * PYNVVIDEOCODEC_MAX_RETAINED_DECODERS + + PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES + ) + + def test_acquire_release_accounting(): pool = MultiModalGPUMemoryPool(total_bytes=100) assert pool.available_bytes == 100 @@ -143,3 +171,72 @@ def test_global_pool_splits_budget_across_api_processes(): def test_global_pool_rejects_invalid_api_process_count(): with pytest.raises(ValueError): maybe_init_mm_gpu_ipc_pool(2, api_process_count=0) + + +@pytest.mark.parametrize("video_backend", [None, "opencv"]) +def test_reserve_mm_ipc_gpu_memory_raw_frame_budget_only( + monkeypatch: pytest.MonkeyPatch, + video_backend: str | None, +): + monkeypatch.setattr( + multimodal_config_module.envs, + "VLLM_VIDEO_LOADER_BACKEND", + "opencv", + ) + mm_config = _mm_config( + mm_ipc_gpu_memory_gb=0.25, + video_backend=video_backend, + ) + + assert reserve_mm_ipc_gpu_memory(GiB_bytes, mm_config) == int(0.75 * GiB_bytes) + + +def test_reserve_mm_ipc_gpu_memory_includes_pynvvideocodec_decoder_budget( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr( + multimodal_config_module.envs, + "VLLM_VIDEO_LOADER_BACKEND", + "opencv", + ) + mm_config = _mm_config( + mm_ipc_gpu_memory_gb=0.25, + video_backend=PYNVVIDEOCODEC_VIDEO_BACKEND, + ) + available_bytes = 4 * GiB_bytes + + assert reserve_mm_ipc_gpu_memory(available_bytes, mm_config) == ( + available_bytes - int(0.25 * GiB_bytes) - _pynvvideocodec_decoder_budget() + ) + + +def test_reserve_mm_ipc_gpu_memory_uses_env_video_backend( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr( + multimodal_config_module.envs, + "VLLM_VIDEO_LOADER_BACKEND", + PYNVVIDEOCODEC_VIDEO_BACKEND, + ) + available_bytes = 4 * GiB_bytes + + assert reserve_mm_ipc_gpu_memory(available_bytes, _mm_config()) == ( + available_bytes - _pynvvideocodec_decoder_budget() + ) + + +def test_reserve_mm_ipc_gpu_memory_scales_decoder_budget_by_api_servers( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr( + multimodal_config_module.envs, + "VLLM_VIDEO_LOADER_BACKEND", + PYNVVIDEOCODEC_VIDEO_BACKEND, + ) + available_bytes = 8 * GiB_bytes + + assert reserve_mm_ipc_gpu_memory( + available_bytes, + _mm_config(), + api_process_count=3, + ) == available_bytes - _pynvvideocodec_decoder_budget(api_process_count=3) diff --git a/tests/quantization/test_auto_gptq.py b/tests/quantization/test_auto_gptq.py index b733ee486c18..9eefd88e915e 100644 --- a/tests/quantization/test_auto_gptq.py +++ b/tests/quantization/test_auto_gptq.py @@ -5,13 +5,17 @@ Run `pytest tests/quantization/test_auto_gptq.py -v -s`. """ +from types import SimpleNamespace + import pytest import torch from tests.quantization.utils import is_quant_method_supported +from vllm.model_executor.layers.fused_moe import RoutedExperts from vllm.model_executor.layers.quantization.auto_gptq import ( AutoGPTQConfig, AutoGPTQLinearMethod, + AutoGPTQMoEMethod, ) PROMPT = "On the surface of Mars, we found" @@ -54,3 +58,78 @@ def check_model(model): def test_auto_gptq_config_get_name(): """Test that AutoGPTQConfig.get_name() returns 'auto_gptq'.""" assert AutoGPTQConfig.get_name() == "auto_gptq" + + +def test_auto_gptq_moe_creates_zero_initialized_expert_biases(): + method = object.__new__(AutoGPTQMoEMethod) + method.quant_config = AutoGPTQConfig(4, 128, False, True, False, {}, {}) + method.input_dtype = None + method.experts_cls = None + layer = torch.nn.Module() + + method.create_weights( + layer=layer, + num_experts=2, + hidden_size=8, + intermediate_size_per_partition=4, + params_dtype=torch.float16, + intermediate_size_full=4, + weight_loader=lambda *args, **kwargs: None, + ) + + assert layer.w13_bias.shape == (2, 8) + assert layer.w2_bias.shape == (2, 8) + assert torch.count_nonzero(layer.w13_bias) == 0 + assert torch.count_nonzero(layer.w2_bias) == 0 + + +def test_routed_experts_loads_per_expert_biases(): + class Loader: + quant_config = None + quant_method = object() + moe_config = SimpleNamespace( + is_act_and_mul=True, + tp_rank=0, + moe_parallel_config=SimpleNamespace(tp_size=1), + ) + _get_hidden_dim = staticmethod(RoutedExperts._get_hidden_dim) + _narrow_expert_data_for_padding = staticmethod( + RoutedExperts._narrow_expert_data_for_padding + ) + _load_w13 = RoutedExperts._load_w13 + _loaded_expert_biases = set() + + @staticmethod + def _map_global_expert_id_to_local_expert_id(expert_id): + return expert_id + + loader = Loader() + w13_bias = torch.nn.Parameter(torch.zeros(1, 8), requires_grad=False) + w2_bias = torch.nn.Parameter(torch.zeros(1, 4), requires_grad=False) + + for shard_id, loaded in ( + ("w1", torch.tensor([1.0, 2.0, 3.0, 4.0])), + ("w3", torch.tensor([5.0, 6.0, 7.0, 8.0])), + ): + assert RoutedExperts.weight_loader( + loader, + w13_bias, + loaded, + weight_name="model.layers.0.mlp.experts.w13_bias", + shard_id=shard_id, + expert_id=0, + return_success=True, + ) + + assert RoutedExperts.weight_loader( + loader, + w2_bias, + torch.tensor([9.0, 10.0, 11.0, 12.0]), + weight_name="model.layers.0.mlp.experts.w2_bias", + shard_id="w2", + expert_id=0, + return_success=True, + ) + assert torch.equal(w13_bias, torch.arange(1, 9, dtype=torch.float32).reshape(1, 8)) + assert torch.equal(w2_bias, torch.arange(9, 13, dtype=torch.float32).reshape(1, 4)) + assert loader._loaded_expert_biases == {"w13_bias", "w2_bias"} diff --git a/tests/quantization/test_moe_wna16.py b/tests/quantization/test_moe_wna16.py index c4b0ab5a8464..16c68fd1c79f 100644 --- a/tests/quantization/test_moe_wna16.py +++ b/tests/quantization/test_moe_wna16.py @@ -5,45 +5,194 @@ import pytest import torch +from compressed_tensors.quantization import ( + ActivationOrdering, + QuantizationArgs, + QuantizationStrategy, + QuantizationType, +) -from vllm.model_executor.layers.fused_moe.activation import MoEActivation -from vllm.model_executor.layers.quantization.moe_wna16 import MoeWNA16Method -from vllm.platforms import current_platform +from vllm.model_executor.layers.fused_moe.oracle.int_wna16 import ( + WNA16MoEBackend, + _backend_incompatibility_reason, + _convert_moe_wna16_humming_tensors, + convert_to_wna16_moe_kernel_format, + map_wna16_backend, +) +from vllm.model_executor.layers.quantization import moe_wna16 +from vllm.model_executor.layers.quantization.auto_awq import AutoAWQConfig +from vllm.model_executor.layers.quantization.auto_gptq import AutoGPTQConfig +from vllm.model_executor.layers.quantization.moe_wna16 import ( + MoeWNA16Config, + MoeWNA16Method, +) -@pytest.mark.skipif(not current_platform.is_cuda(), reason="Only test on CUDA") -def test_moe_wna16_apply_passes_layer_activation(monkeypatch): - captured_kwargs = {} +def test_map_wna16_backend_supports_triton(): + assert map_wna16_backend("triton") == WNA16MoEBackend.TRITON - def fake_fused_experts(*args, **kwargs): - captured_kwargs.update(kwargs) - return torch.empty(1, 2) - monkeypatch.setattr( - "vllm.model_executor.layers.fused_moe.fused_experts", - fake_fused_experts, +@pytest.mark.parametrize( + ("backend", "quant_config", "may_have_zp", "may_have_bias", "expected"), + [ + ( + WNA16MoEBackend.TRITON, + AutoAWQConfig(4, 128, True, False), + True, + False, + "AutoAWQ weight layout", + ), + ( + WNA16MoEBackend.TRITON, + AutoGPTQConfig(4, 128, True, True, False, {}, {}), + False, + False, + "activation ordering", + ), + ( + WNA16MoEBackend.TRITON, + QuantizationArgs( + num_bits=4, + type=QuantizationType.INT, + strategy=QuantizationStrategy.GROUP, + symmetric=True, + dynamic=False, + group_size=128, + actorder=ActivationOrdering.GROUP, + ), + False, + False, + "activation ordering", + ), + ( + WNA16MoEBackend.TRITON, + AutoGPTQConfig(4, 128, False, True, False, {}, {}), + False, + True, + "bias", + ), + ( + WNA16MoEBackend.MARLIN, + MoeWNA16Config( + linear_quant_method="gptq", + weight_bits=4, + group_size=128, + has_zp=False, + lm_head_quantized=False, + modules_to_not_convert=None, + full_config={}, + ), + False, + False, + "MoeWNA16 checkpoint layout", + ), + ], +) +def test_wna16_oracle_rejects_incompatible_quant_structures( + backend, quant_config, may_have_zp, may_have_bias, expected +): + reason = _backend_incompatibility_reason( + backend=backend, + quant_config=quant_config, + may_have_zp=may_have_zp, + may_have_bias=may_have_bias, + ) + + assert reason is not None + assert expected in reason + + +def test_compressed_tensors_weights_are_transposed_for_triton(): + quant_config = QuantizationArgs( + num_bits=4, + type=QuantizationType.INT, + strategy=QuantizationStrategy.GROUP, + symmetric=True, + dynamic=False, + group_size=32, ) + w13 = torch.arange(16, dtype=torch.int32).reshape(1, 2, 8) + w2 = torch.arange(12, dtype=torch.int32).reshape(1, 2, 6) + w13_scale = torch.arange(32, dtype=torch.float16).reshape(1, 4, 8) + w2_scale = torch.arange(18, dtype=torch.float16).reshape(1, 3, 6) + converted = convert_to_wna16_moe_kernel_format( + backend=WNA16MoEBackend.TRITON, + layer=torch.nn.Module(), + quant_config=quant_config, + input_dtype=None, + w13=w13, + w2=w2, + w13_scale=w13_scale, + w2_scale=w2_scale, + ) + + assert converted is not None + assert torch.equal(converted[0], w13.transpose(1, 2).contiguous().view(torch.uint8)) + assert torch.equal(converted[1], w2.transpose(1, 2).contiguous().view(torch.uint8)) + assert torch.equal(converted[2], w13_scale.transpose(1, 2).contiguous()) + assert torch.equal(converted[3], w2_scale.transpose(1, 2).contiguous()) + + +def test_moe_wna16_setup_forwards_selected_backend(monkeypatch): method = object.__new__(MoeWNA16Method) - method.moe = SimpleNamespace(disable_inplace=False) - method.moe_quant_config = object() - layer = SimpleNamespace( - w13_qweight=torch.empty(1, 2), - w2_qweight=torch.empty(1, 2), - activation=MoEActivation.GELU_TANH, - apply_router_weight_on_input=False, - global_num_experts=1, - expert_map=None, + method.experts_cls = object + method.wna16_backend = WNA16MoEBackend.HUMMING + method.moe = object() + quant_config = object() + method.get_fused_moe_quant_config = lambda layer: quant_config + layer = SimpleNamespace(_expert_routing_tables=lambda: (None, None, None)) + captured = {} + kernel = object() + + def fake_make_wna16_moe_kernel(**kwargs): + captured.update(kwargs) + return kernel + + monkeypatch.setattr(moe_wna16, "make_wna16_moe_kernel", fake_make_wna16_moe_kernel) + + method._setup_kernel(layer) + + assert method.moe_kernel is kernel + assert captured["backend"] == WNA16MoEBackend.HUMMING + assert captured["layer"] is layer + + +def test_moe_wna16_humming_adapter_repacks_uint8_tensors(): + qweight = torch.arange(32, dtype=torch.uint8).reshape(1, 4, 8) + scales = torch.arange(16, dtype=torch.float16).reshape(1, 4, 4) + qzeros = torch.arange(16, dtype=torch.uint8).reshape(1, 8, 2) + + converted = _convert_moe_wna16_humming_tensors( + {"qweight": qweight, "scales": scales, "qzeros": qzeros}, + has_zero_point=True, ) - output = method.apply( - layer, - x=torch.empty(1, 2), - topk_weights=torch.empty(1, 1), - topk_ids=torch.empty(1, 1, dtype=torch.int32), - shared_experts=None, - shared_experts_input=None, + assert torch.equal(converted["weight"], qweight.view(torch.int32)) + assert converted["weight"].shape == (1, 4, 2) + assert torch.equal(converted["weight_scale"], scales) + expected_qzeros = ( + qzeros.transpose(-1, -2) + .contiguous() + .view(torch.int32) + .transpose(-1, -2) + .contiguous() + ) + assert torch.equal(converted["zero_point"], expected_qzeros) + assert converted["zero_point"].shape == (1, 2, 2) + + +def test_moe_wna16_uses_humming_quant_config(monkeypatch): + from vllm.model_executor.layers.quantization.utils import humming_utils + + method = object.__new__(MoeWNA16Method) + method.wna16_backend = WNA16MoEBackend.HUMMING + layer = object() + quant_config = object() + monkeypatch.setattr( + humming_utils, + "get_humming_moe_quant_config", + lambda actual_layer: quant_config if actual_layer is layer else None, ) - assert output.shape == (1, 2) - assert captured_kwargs["activation"] is MoEActivation.GELU_TANH + assert method.get_fused_moe_quant_config(layer) is quant_config diff --git a/tests/v1/attention/test_attention_backends.py b/tests/v1/attention/test_attention_backends.py index b190433b254a..87a9c80942ee 100644 --- a/tests/v1/attention/test_attention_backends.py +++ b/tests/v1/attention/test_attention_backends.py @@ -46,9 +46,12 @@ DEVICE_TYPE = current_platform.device_type +# Use the platform's preferred FP8 type so the stored cache matches what the +# backends reinterpret at runtime. On ROCm gfx94x this is e4m3fnuz, not e4m3fn; +# storing e4m3fn bytes there would be re-read as fnuz and produce NaNs. FP8_KV_CACHE_DTYPES = { - "fp8": torch.float8_e4m3fn, - "fp8_e4m3": torch.float8_e4m3fn, + "fp8": current_platform.fp8_dtype(), + "fp8_e4m3": current_platform.fp8_dtype(), } # Remove flashinfer from the list if it's not available diff --git a/tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh b/tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh index 0a45b6f54b87..244399a13f4a 100755 --- a/tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh +++ b/tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh @@ -75,6 +75,7 @@ fi NUM_PREFILL_INSTANCES=${NUM_PREFILL_INSTANCES:-1} # Default to 1 NUM_DECODE_INSTANCES=${NUM_DECODE_INSTANCES:-1} # Default to 1 PREFILLER_TP_SIZE=${PREFILLER_TP_SIZE:-1} +PREFILLER_PCP_SIZE=${PREFILLER_PCP_SIZE:-1} PREFILLER_PP_SIZE=${PREFILLER_PP_SIZE:-1} # >1 requires NixlPushConnector DECODER_TP_SIZE=${DECODER_TP_SIZE:-1} GPU_MEMORY_UTILIZATION=${GPU_MEMORY_UTILIZATION:-0.2} @@ -148,8 +149,8 @@ run_tests_for_model() { GPU_START=$((i % $(get_num_gpus))) GPU_ID=$GPU_START NEXT_GPU=$GPU_START - # Reserve TP*PP GPUs for the prefiller (TP shards across PP stages). - PREFILLER_WORLD_SIZE=$((PREFILLER_TP_SIZE * PREFILLER_PP_SIZE)) + # Reserve TP*PCP*PP GPUs for the prefiller. + PREFILLER_WORLD_SIZE=$((PREFILLER_TP_SIZE * PREFILLER_PCP_SIZE * PREFILLER_PP_SIZE)) for (( j=1; j < PREFILLER_WORLD_SIZE; j++ )); do NEXT_GPU=$(((GPU_START + j) % $(get_num_gpus))) GPU_ID="${GPU_ID},${NEXT_GPU}" @@ -174,6 +175,7 @@ run_tests_for_model() { --block-size ${PREFILL_BLOCK_SIZE} \ --gpu-memory-utilization $GPU_MEMORY_UTILIZATION \ --tensor-parallel-size $PREFILLER_TP_SIZE \ + --prefill-context-parallel-size $PREFILLER_PCP_SIZE \ --pipeline-parallel-size $PREFILLER_PP_SIZE \ --kv-transfer-config '$KV_CONFIG_P'" if [[ "$ENFORCE_EAGER" == "1" ]]; then diff --git a/tests/v1/kv_connector/unit/test_nixl_connector.py b/tests/v1/kv_connector/unit/test_nixl_connector.py index 31d34c69b3fb..ddeffe6cdc27 100644 --- a/tests/v1/kv_connector/unit/test_nixl_connector.py +++ b/tests/v1/kv_connector/unit/test_nixl_connector.py @@ -204,13 +204,12 @@ def get_xfer_telemetry(self, handle: int) -> dict: def _make_fake_nixl_pkg(): """Context manager that creates a temporary package making `from nixl._api import nixl_agent` resolve to our FakeNixlWrapper. - Also creates rixl package for ROCm compatibility. + Also creates the ROCm NIXL packages. Automatically cleans up the temporary directory when done. """ with tempfile.TemporaryDirectory() as td: - # Create both nixl and rixl packages for cross-platform compatibility - for pkg_name in ["nixl", "rixl"]: + for pkg_name in ["nixl", "nixl_rocm"]: pkg_root = os.path.join(td, pkg_name, "_api") os.makedirs(pkg_root, exist_ok=True) @@ -566,6 +565,45 @@ def _nixl_handshake( class TestNixlHandshake: + @patch( + "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper", + FakeNixlWrapper, + ) + def test_pcp_producer_publishes_one_kv_replica( + self, default_vllm_config, dist_init + ): + vllm_config = create_vllm_config(kv_role="kv_producer") + connector = NixlConnector( + vllm_config, + KVConnectorRole.WORKER, + make_kv_cache_config(block_size=16), + ) + payload = MagicMock(spec=NixlHandshakePayload) + worker = connector.connector_worker + assert worker is not None + worker.xfer_handshake_metadata = payload + + worker.pcp_rank = 0 + assert connector.get_handshake_metadata() is payload + + worker.pcp_rank = 1 + assert connector.get_handshake_metadata() is None + + def test_pcp_producer_waits_only_for_published_replicas( + self, default_vllm_config, dist_init + ): + vllm_config = create_vllm_config(kv_role="kv_producer") + vllm_config.parallel_config.tensor_parallel_size = 2 + vllm_config.parallel_config.prefill_context_parallel_size = 2 + vllm_config.parallel_config.pipeline_parallel_size = 2 + connector = NixlConnector( + vllm_config, + KVConnectorRole.SCHEDULER, + make_kv_cache_config(block_size=16), + ) + + assert connector.get_finished_count() == 4 + @patch( "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper", FakeNixlWrapper, @@ -3130,10 +3168,12 @@ def test_mla_broadcast_notif_uses_remote_request_id( ) +@patch( + "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper", + FakeNixlWrapper, +) def test_kv_both_deprecation_warning(default_vllm_config, dist_init): """kv_role='kv_both' should emit a deprecation log warning.""" - from unittest.mock import patch - from vllm.logger import _print_warning_once _print_warning_once.cache_clear() @@ -3156,10 +3196,12 @@ def test_kv_both_deprecation_warning(default_vllm_config, dist_init): assert "deprecated" in msg +@patch( + "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper", + FakeNixlWrapper, +) def test_explicit_kv_role_no_deprecation_warning(default_vllm_config, dist_init): """kv_role='kv_consumer' or 'kv_producer' should NOT emit a warning.""" - from unittest.mock import patch - for role in ("kv_consumer", "kv_producer"): vllm_config = create_vllm_config(kv_role=role) with patch( diff --git a/tests/v1/kv_connector/unit/test_rixl_gpu_mem_diag.py b/tests/v1/kv_connector/unit/test_nixl_rocm_gpu_mem_diag.py similarity index 94% rename from tests/v1/kv_connector/unit/test_rixl_gpu_mem_diag.py rename to tests/v1/kv_connector/unit/test_nixl_rocm_gpu_mem_diag.py index 2371b5555bce..da8be07157d8 100644 --- a/tests/v1/kv_connector/unit/test_rixl_gpu_mem_diag.py +++ b/tests/v1/kv_connector/unit/test_nixl_rocm_gpu_mem_diag.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Verify that GPU memory is fully released after RixlConnector shutdown on ROCm. +"""Verify that GPU memory is released after NixlConnector shutdown on ROCm. Regression test for ROCm/ucx#33: UCX rocm_ipc transport permanently pinned GPU memory via hsa_amd_ipc_memory_create during ucp_mem_map, causing @@ -62,7 +62,7 @@ def _full_gpu_cleanup(): @pytest.mark.parametrize("model_name, sw_size", [("google/gemma-3-1b-it", 512)]) -def test_gpu_memory_rixl_hma(model_name, sw_size): +def test_gpu_memory_nixl_hma(model_name, sw_size): """Track GPU memory through NixlConnector create/infer/shutdown cycle.""" from vllm import LLM, SamplingParams from vllm.config import KVTransferConfig @@ -84,7 +84,7 @@ def test_gpu_memory_rixl_hma(model_name, sw_size): } print("\n" + "=" * 90) - print("GPU MEMORY -- RIXL NixlConnector HMA (ROCm)") + print("GPU MEMORY -- NIXL NixlConnector HMA (ROCm)") print("=" * 90) gc.collect() torch.accelerator.empty_cache() @@ -169,14 +169,14 @@ def test_gpu_memory_rixl_hma(model_name, sw_size): @pytest.mark.parametrize("model_name", ["google/gemma-3-1b-it"]) -def test_gpu_memory_no_rixl_baseline(model_name): +def test_gpu_memory_no_nixl_baseline(model_name): """Same workload without NixlConnector. Comparing driver-level memory - between this and test_gpu_memory_rixl_hma isolates UCX/RIXL impact.""" + between this and test_gpu_memory_nixl_hma isolates UCX/NIXL impact.""" from vllm import LLM, SamplingParams from vllm.distributed.parallel_state import cleanup_dist_env_and_memory print("\n" + "=" * 90) - print("CONTROL -- same model, no RIXL connector") + print("CONTROL -- same model, no NIXL connector") print("=" * 90) gc.collect() torch.accelerator.empty_cache() @@ -209,7 +209,7 @@ def test_gpu_memory_no_rixl_baseline(model_name): drv_base = snap0["drv_used_mb"] drv_leaked = snap_final["drv_used_mb"] - drv_base drv_peak = snap_peak["drv_used_mb"] - drv_base - print(f"\n Driver leaked (no rixl): {drv_leaked:.0f} MB") + print(f"\n Driver leaked (no NIXL): {drv_leaked:.0f} MB") print("=" * 90) leak_pct = (drv_leaked / drv_peak * 100) if drv_peak > 0 else 0 diff --git a/tests/v1/streaming_input/test_gpu_model_runner_streaming.py b/tests/v1/streaming_input/test_gpu_model_runner_streaming.py index fd619610b767..a8f7221da81b 100644 --- a/tests/v1/streaming_input/test_gpu_model_runner_streaming.py +++ b/tests/v1/streaming_input/test_gpu_model_runner_streaming.py @@ -16,7 +16,8 @@ from vllm.v1.worker.gpu_input_batch import CachedRequestState, InputBatch from vllm.v1.worker.gpu_model_runner import GPUModelRunner -pytestmark = pytest.mark.cpu_test +# Not cpu_test: InputBatch allocates pinned (UVA) memory, which requires a +# CUDA device even though the batch tensors live on the CPU. @pytest.fixture @@ -28,6 +29,7 @@ def mock_model_runner_with_input_batch(): runner.requests = {} runner.max_num_reqs = 10 runner.max_model_len = 1024 + runner.late_interaction_runner = Mock() # Create a real InputBatch for e2e testing runner.input_batch = InputBatch( diff --git a/tests/v1/streaming_input/test_gpu_model_runner_v2_streaming.py b/tests/v1/streaming_input/test_gpu_model_runner_v2_streaming.py index 8fde0f117ca2..f598d0547f04 100644 --- a/tests/v1/streaming_input/test_gpu_model_runner_v2_streaming.py +++ b/tests/v1/streaming_input/test_gpu_model_runner_v2_streaming.py @@ -16,7 +16,8 @@ from vllm.v1.worker.gpu.model_runner import GPUModelRunner from vllm.v1.worker.gpu.states import RequestState -pytestmark = pytest.mark.cpu_test +# Not cpu_test: RequestState allocates pinned (UVA) memory, which requires a +# CUDA device even though the request state itself lives on the CPU. @pytest.fixture @@ -31,13 +32,12 @@ def mock_model_runner_with_req_states(): num_speculative_steps=0, vocab_size=32000, device=torch.device("cpu"), - model_dtype=torch.float32, - cache_draft_logits=False, ) runner.encoder_cache = None runner.model_state = Mock() runner.block_tables = Mock() runner.lora_state = Mock() + runner.pp_handler = None runner.sampler = None runner.prompt_logprobs_worker = None runner.is_last_pp_rank = False diff --git a/tests/v1/streaming_input/test_scheduler_streaming.py b/tests/v1/streaming_input/test_scheduler_streaming.py index 7d680895b836..822b4e49da34 100644 --- a/tests/v1/streaming_input/test_scheduler_streaming.py +++ b/tests/v1/streaming_input/test_scheduler_streaming.py @@ -53,6 +53,8 @@ def create_scheduler() -> Scheduler: vllm_config.model_config = MagicMock() vllm_config.model_config.skip_tokenizer_init = True vllm_config.model_config.is_multimodal_model = False + vllm_config.model_config.is_encoder_decoder = False + vllm_config.model_config.is_diffusion = False vllm_config.model_config.max_model_len = 1024 vllm_config.model_config.enable_return_routed_experts = False vllm_config.cache_config = MagicMock() @@ -496,7 +498,9 @@ def test_streaming_e2e_lifecycle(self): eco_cycle2 = eco_dict_cycle2[session.client_index].outputs[0] assert eco_cycle2.finish_reason == FinishReason.STOP assert session.status == RequestStatus.WAITING_FOR_STREAMING_REQ - assert session in scheduler.waiting + # Sessions paused for streaming input are blocked-waiting, so they + # live in the skipped_waiting queue rather than the main waiting queue. + assert session in scheduler.skipped_waiting assert session._all_token_ids == [1, 2, 3, 10, STOP_TOKEN] # CRITICAL ASSERTION: Cached prompt_token_ids STILL must not have changed diff --git a/tests/v1/worker/test_gpu_worker.py b/tests/v1/worker/test_gpu_worker.py index cdaa644b62ec..43232ca6be40 100644 --- a/tests/v1/worker/test_gpu_worker.py +++ b/tests/v1/worker/test_gpu_worker.py @@ -6,122 +6,13 @@ import pytest -import vllm.v1.worker.gpu_worker as gpu_worker_module -from vllm.multimodal.video import ( - PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES, - PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES, - PYNVVIDEOCODEC_MAX_RETAINED_DECODERS, - PYNVVIDEOCODEC_VIDEO_BACKEND, -) from vllm.utils.mem_constants import GiB_bytes from vllm.v1.worker import startup_plan -from vllm.v1.worker.gpu_worker import Worker from vllm.v1.worker.startup_plan import ( maybe_apply_startup_plan, maybe_save_startup_plan, ) - -def _worker_with_mm_config( - mm_config: SimpleNamespace, - *, - api_process_count: int = 1, -) -> Worker: - worker = object.__new__(Worker) - worker.model_config = SimpleNamespace(multimodal_config=mm_config) - worker.parallel_config = SimpleNamespace(_api_process_count=api_process_count) - return worker - - -def _mm_config( - *, - mm_ipc_gpu_memory_gb: float = 0, - video_backend: str | None = None, -) -> SimpleNamespace: - video_kwargs = {} if video_backend is None else {"video_backend": video_backend} - return SimpleNamespace( - mm_ipc_gpu_memory_gb=mm_ipc_gpu_memory_gb, - media_io_kwargs={"video": video_kwargs} if video_kwargs else {}, - ) - - -def _pynvvideocodec_decoder_budget(api_process_count: int = 1) -> int: - return api_process_count * ( - PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES * PYNVVIDEOCODEC_MAX_RETAINED_DECODERS - + PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES - ) - - -@pytest.mark.parametrize("video_backend", [None, "opencv"]) -def test_reserve_mm_ipc_gpu_memory_raw_frame_budget_only( - monkeypatch: pytest.MonkeyPatch, - video_backend: str | None, -): - monkeypatch.setattr( - gpu_worker_module.envs, - "VLLM_VIDEO_LOADER_BACKEND", - "opencv", - ) - worker = _worker_with_mm_config( - _mm_config(mm_ipc_gpu_memory_gb=0.25, video_backend=video_backend) - ) - - assert worker._reserve_mm_ipc_gpu_memory(GiB_bytes) == int(0.75 * GiB_bytes) - - -def test_reserve_mm_ipc_gpu_memory_includes_pynvvideocodec_decoder_budget( - monkeypatch: pytest.MonkeyPatch, -): - monkeypatch.setattr( - gpu_worker_module.envs, - "VLLM_VIDEO_LOADER_BACKEND", - "opencv", - ) - worker = _worker_with_mm_config( - _mm_config( - mm_ipc_gpu_memory_gb=0.25, - video_backend=PYNVVIDEOCODEC_VIDEO_BACKEND, - ) - ) - available_bytes = 4 * GiB_bytes - - assert worker._reserve_mm_ipc_gpu_memory(available_bytes) == ( - available_bytes - int(0.25 * GiB_bytes) - _pynvvideocodec_decoder_budget() - ) - - -def test_reserve_mm_ipc_gpu_memory_uses_env_video_backend( - monkeypatch: pytest.MonkeyPatch, -): - monkeypatch.setattr( - gpu_worker_module.envs, - "VLLM_VIDEO_LOADER_BACKEND", - PYNVVIDEOCODEC_VIDEO_BACKEND, - ) - worker = _worker_with_mm_config(_mm_config()) - available_bytes = 4 * GiB_bytes - - assert worker._reserve_mm_ipc_gpu_memory(available_bytes) == ( - available_bytes - _pynvvideocodec_decoder_budget() - ) - - -def test_reserve_mm_ipc_gpu_memory_scales_pynvvideocodec_budget_by_api_servers( - monkeypatch: pytest.MonkeyPatch, -): - monkeypatch.setattr( - gpu_worker_module.envs, - "VLLM_VIDEO_LOADER_BACKEND", - PYNVVIDEOCODEC_VIDEO_BACKEND, - ) - worker = _worker_with_mm_config(_mm_config(), api_process_count=3) - available_bytes = 8 * GiB_bytes - - assert worker._reserve_mm_ipc_gpu_memory(available_bytes) == ( - available_bytes - _pynvvideocodec_decoder_budget(api_process_count=3) - ) - - # Startup-plan persistence (vllm/v1/worker/startup_plan.py), applied and # saved by Worker.determine_available_memory / compile_or_warm_up_model. diff --git a/vllm/_aiter_ops.py b/vllm/_aiter_ops.py index 8ab06d446bf4..4a915c0dd68e 100644 --- a/vllm/_aiter_ops.py +++ b/vllm/_aiter_ops.py @@ -409,37 +409,6 @@ def _rocm_aiter_fused_topk_fake( # Cache whether aiter supports FP8 MLA parameters _AITER_MLA_SUPPORTS_FP8: bool | None = None -_AITER_HAS_FUSED_QK_RMSNORM: bool | None = None - - -def check_aiter_fused_qk_rmsnorm() -> bool: - """Check if aiter provides fused_qk_rmsnorm. - - Supports both the new private name ``_fused_qk_rmsnorm`` - (AITER >= PR #2958) and the old public name ``fused_qk_rmsnorm`` - (AITER >= PR #2442). - - TODO(rbrugaro-amd): remove the legacy fused_qk_rmsnorm path once - AITER stabilizes the API (https://github.com/ROCm/aiter/issues/3207). - """ - global _AITER_HAS_FUSED_QK_RMSNORM - if _AITER_HAS_FUSED_QK_RMSNORM is None: - try: - from aiter.ops.fused_qk_norm_rope_cache_quant import ( # noqa: F401 - _fused_qk_rmsnorm, - ) - - _AITER_HAS_FUSED_QK_RMSNORM = True - except (ImportError, ModuleNotFoundError, AttributeError): - try: - from aiter.ops.fused_qk_norm_rope_cache_quant import ( # noqa: F401 - fused_qk_rmsnorm, - ) - - _AITER_HAS_FUSED_QK_RMSNORM = True - except (ImportError, ModuleNotFoundError, AttributeError): - _AITER_HAS_FUSED_QK_RMSNORM = False - return _AITER_HAS_FUSED_QK_RMSNORM def _check_aiter_mla_fp8_support() -> bool: @@ -1267,43 +1236,17 @@ def _fused_mla_dual_rms_norm_impl( x1_epsilon: float, x2_epsilon: float, ) -> tuple[torch.Tensor, torch.Tensor]: - try: - import aiter.ops.fused_qk_norm_rope_cache_quant as aiter_ops - except (ImportError, ModuleNotFoundError, AttributeError) as exc: - raise ImportError( - "fused_qk_rmsnorm requires AITer >= PR #2442. " - "Please upgrade aiter or disable the " - "fuse_mla_dual_rms_norm pass." - ) from exc - - if hasattr(aiter_ops, "_fused_qk_rmsnorm"): - return aiter_ops._fused_qk_rmsnorm( - q_out=None, - q=x1, - q_weight=x1_weight, - q_eps=x1_epsilon, - k_out=None, - k=x2, - k_weight=x2_weight, - k_eps=x2_epsilon, - ) - - # TODO(rbrugaro-amd): remove the legacy fused_qk_rmsnorm path once - # AITER stabilizes the API (https://github.com/ROCm/aiter/issues/3207). - if hasattr(aiter_ops, "fused_qk_rmsnorm"): - return aiter_ops.fused_qk_rmsnorm( - q=x1, - q_weight=x1_weight, - q_eps=x1_epsilon, - k=x2, - k_weight=x2_weight, - k_eps=x2_epsilon, - ) - - raise ImportError( - "fused_qk_rmsnorm requires AITer >= PR #2442. " - "Please upgrade aiter or disable the " - "fuse_mla_dual_rms_norm pass." + from aiter.ops.fused_qk_norm_rope_cache_quant import _fused_qk_rmsnorm + + return _fused_qk_rmsnorm( + q_out=None, + q=x1, + q_weight=x1_weight, + q_eps=x1_epsilon, + k_out=None, + k=x2, + k_weight=x2_weight, + k_eps=x2_epsilon, ) diff --git a/vllm/_custom_ops.py b/vllm/_custom_ops.py index f635077a1d50..b22c8f6551d4 100644 --- a/vllm/_custom_ops.py +++ b/vllm/_custom_ops.py @@ -2391,6 +2391,17 @@ def topk_softmax( e_score_correction_bias: torch.Tensor | None = None, is_padding: torch.Tensor | None = None, ) -> None: + if current_platform.is_xpu(): + # TODO: Remove after vllm-xpu-kernels supports is_padding. + torch.ops._moe_C.topk_softmax( + topk_weights, + topk_ids, + token_expert_indices, + gating_output, + renormalize, + e_score_correction_bias, + ) + return torch.ops._moe_C.topk_softmax( topk_weights, topk_ids, @@ -2436,6 +2447,21 @@ def topk_hash_softplus_sqrt( hash_indices_table: torch.Tensor | None = None, is_padding: torch.Tensor | None = None, ) -> None: + if current_platform.is_xpu(): + # TODO: Remove after vllm-xpu-kernels supports is_padding. + torch.ops._moe_C.topk_softplus_sqrt( + topk_weights, + topk_indices, + token_expert_indices, + gating_output, + renormalize, + routed_scaling_factor, + e_score_correction_bias, + input_tokens, + hash_indices_table, + ) + + return torch.ops._moe_C.topk_softplus_sqrt( topk_weights, topk_indices, diff --git a/vllm/compilation/passes/pass_manager.py b/vllm/compilation/passes/pass_manager.py index 67a3e4cbae5c..2fe863189c1f 100644 --- a/vllm/compilation/passes/pass_manager.py +++ b/vllm/compilation/passes/pass_manager.py @@ -7,7 +7,7 @@ from torch import fx as fx from vllm import envs -from vllm._aiter_ops import check_aiter_fused_qk_rmsnorm, rocm_aiter_ops +from vllm._aiter_ops import rocm_aiter_ops from vllm.compilation.passes.utility.post_cleanup import PostCleanupPass from vllm.config import VllmConfig, set_current_vllm_config from vllm.logger import init_logger @@ -177,11 +177,7 @@ def configure(self, config: VllmConfig) -> None: self.passes += [ScatterSplitReplacementPass(config)] self.passes += [QkNormRopeKvCacheFusionPass(config)] - if ( - self.pass_config.fuse_mla_dual_rms_norm - and rocm_aiter_ops.is_enabled() - and check_aiter_fused_qk_rmsnorm() - ): + if self.pass_config.fuse_mla_dual_rms_norm and rocm_aiter_ops.is_enabled(): self.passes += [MLADualRMSNormFusionPass(config)] if self.pass_config.fuse_rope_kvcache: diff --git a/vllm/config/multimodal.py b/vllm/config/multimodal.py index 150d58cf1f74..865615dff799 100644 --- a/vllm/config/multimodal.py +++ b/vllm/config/multimodal.py @@ -8,6 +8,7 @@ from pydantic import ConfigDict, Field, field_validator, model_validator from pydantic.dataclasses import dataclass +import vllm.envs as envs from vllm.config.utils import config from vllm.utils.hashing import safe_hash from vllm.v1.attention.backends.registry import AttentionBackendEnum @@ -344,5 +345,19 @@ def merge_mm_processor_kwargs( kwargs = self.mm_processor_kwargs or {} return kwargs | dict(inference_kwargs) + def use_gpu_video_backend(self) -> bool: + """Return whether the configured video loader or codec uses the GPU.""" + from vllm.multimodal.video import VIDEO_LOADER_REGISTRY + + video_kwargs = self.media_io_kwargs.get("video", {}) + video_loader_backend = ( + video_kwargs.get("video_backend") or envs.VLLM_VIDEO_LOADER_BACKEND + ) + codec_backend = video_kwargs.get("backend") + return VIDEO_LOADER_REGISTRY.backend_requires_gpu(video_loader_backend) or ( + codec_backend is not None + and VIDEO_LOADER_REGISTRY.backend_requires_gpu(codec_backend) + ) + def is_multimodal_pruning_enabled(self): return self.video_pruning_rate is not None and self.video_pruning_rate > 0 diff --git a/vllm/config/parallel.py b/vllm/config/parallel.py index 53688c05d92d..ce038a99ea9a 100644 --- a/vllm/config/parallel.py +++ b/vllm/config/parallel.py @@ -92,7 +92,7 @@ class EPLBConfig: Backend for EPLB expert weight communication: - "torch_nccl": Use torch.distributed on the device process group - "torch_gloo": Use torch.distributed gloo with CPU staging - - "nixl": Use NIXL/ RIXL with staged send/recv buffers + - "nixl": Use NIXL with staged send/recv buffers - "pynccl": Use PyNccl send/recv - None: Auto-select backend (prefers "nixl", falls back to "torch_gloo") """ diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 4b4a97f41b6b..bda5013a5ff6 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -187,10 +187,10 @@ def enable_norm_pad_fusion(cfg: "VllmConfig") -> bool: def enable_mla_dual_rms_norm_fusion(cfg: "VllmConfig") -> bool: - """Enable MLA dual RMS norm fusion when AITer has fused_qk_rmsnorm.""" - from vllm._aiter_ops import check_aiter_fused_qk_rmsnorm, rocm_aiter_ops + """Enable MLA dual RMS norm fusion on ROCm with AITER.""" + from vllm._aiter_ops import rocm_aiter_ops - return rocm_aiter_ops.is_enabled() and check_aiter_fused_qk_rmsnorm() + return rocm_aiter_ops.is_enabled() def enable_qk_norm_rope_kvcache(cfg: "VllmConfig") -> bool: diff --git a/vllm/distributed/eplb/eplb_communicator.py b/vllm/distributed/eplb/eplb_communicator.py index 891b57bcf18f..f9a9a8a90a81 100644 --- a/vllm/distributed/eplb/eplb_communicator.py +++ b/vllm/distributed/eplb/eplb_communicator.py @@ -38,7 +38,7 @@ def has_nixl() -> bool: - """Whether the optional NIXL / RIXL package is available.""" + """Whether the optional NIXL package is available.""" return nixl_utils.NixlWrapper is not None @@ -266,7 +266,7 @@ def __init__( assert expert_buffer, "NixlEplbCommunicator requires non-empty expert_buffer." nixl_wrapper_cls = nixl_utils.NixlWrapper if nixl_wrapper_cls is None: - raise RuntimeError("NIXL/ RIXL is unavailable.") + raise RuntimeError("NIXL is unavailable.") self._cpu_group = cpu_group self._world_size = cpu_group.size() diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py index e9315f6776c3..18936993ba8f 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py @@ -61,6 +61,7 @@ ) from vllm.distributed.nixl_utils import NixlWrapper, nixl_agent_config from vllm.distributed.parallel_state import ( + get_pcp_group, get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, ) @@ -354,6 +355,7 @@ def __init__( self.engine_id: EngineId = engine_id self.tp_rank = get_tensor_model_parallel_rank() self.world_size = get_tensor_model_parallel_world_size() + self.pcp_rank = get_pcp_group().rank_in_group self.num_blocks = kv_cache_config.num_blocks self.enable_permute_local_kv = False @@ -1598,12 +1600,14 @@ def add_remote_agent( ### (Optional) Register local agent memory regions. MLA is not split. if ( tp_ratio < 0 - and not self.use_mla + and (not self.use_mla or len(plan.all_source_ranks) > 1) and tp_ratio not in self.src_xfer_handles_by_tp_ratio ): # Remote tp_size > local tp_size: read from multiple remote ranks. - # Logically "split" own regions into |tp_ratio| chunks. Mind that - # we only do this once per remote tp_size (replica-friendly). + # Logically "split" own regions into per-source chunks. Hybrid + # MLA+SSM also needs this path: MLA is replicated and read once, + # while the SSM state is sharded across every remote TP rank. + # We only do this once per remote tp_size (replica-friendly). self.src_xfer_handles_by_tp_ratio[tp_ratio] = [] for handle_data in self._build_local_splits_from_plan( diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py index 7d9a34406b7d..938e5f2c9e16 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py @@ -160,6 +160,18 @@ def get_required_kvcache_layout(cls, vllm_config: VllmConfig): # Scheduler Side Methods ############################################################ + def get_finished_count(self) -> int | None: + parallel_config = self._vllm_config.parallel_config + if ( + self.kv_transfer_config.kv_role == "kv_producer" + and parallel_config.prefill_context_parallel_size > 1 + ): + return ( + parallel_config.tensor_parallel_size + * parallel_config.pipeline_parallel_size + ) + return None + def get_num_new_matched_tokens( self, request: "Request", num_computed_tokens: int ) -> tuple[int | None, bool]: @@ -240,7 +252,13 @@ def set_host_xfer_buffer_ops(self, copy_operation: CopyBlocksOp): def get_finished(self, finished_req_ids: set[str]) -> tuple[set[str], set[str]]: """Get the finished recving and sending requests.""" assert self.connector_worker is not None - return self.connector_worker.get_finished() + done_sending, done_recving = self.connector_worker.get_finished() + if ( + self.kv_transfer_config.kv_role == "kv_producer" + and self.connector_worker.pcp_rank != 0 + ): + done_sending.clear() + return done_sending, done_recving def get_block_ids_with_load_errors(self) -> set[int]: """Get block IDs that failed to load via NIXL.""" @@ -316,6 +334,11 @@ def get_handshake_metadata(self) -> KVConnectorHandshakeMetadata | None: None if no handshake metadata is available. """ assert self.connector_worker is not None + if ( + self.kv_transfer_config.kv_role == "kv_producer" + and self.connector_worker.pcp_rank != 0 + ): + return None return self.connector_worker.xfer_handshake_metadata diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py index 63969382dff8..40d6851769c9 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py @@ -161,9 +161,9 @@ def _read_blocks_for_req(self, req_id: str, meta: ReqMeta): ] # D may have to perform multiple reads from different remote ranks. - # MLA opt: when P TP > D TP, only a single read is executed for - # the first remote rank (cache is duplicated).. - if self.use_mla and tp_ratio < 0: + # Pure MLA reads once because its cache is replicated. Hybrid + # MLA+SSM still needs one read per SSM source rank. + if self.use_mla and tp_ratio < 0 and not self._has_mamba: assert len(read_specs) == 1 for i, spec in enumerate(read_specs): @@ -177,7 +177,7 @@ def _read_blocks_for_req(self, req_id: str, meta: ReqMeta): req_id, ) # Get side handles. - if tp_ratio < 0 and not self.use_mla: + if tp_ratio < 0 and (not self.use_mla or len(read_specs) > 1): assert remote_block_size == self.block_size # Remote tp_size > local tp_size: we must perform multiple # reads. Get the memory chunk onto which we will write to. @@ -203,7 +203,7 @@ def _read_blocks_for_req(self, req_id: str, meta: ReqMeta): remote_xfer_side_handle=remote_xfer_side_handle, ) - if self.use_mla and tp_ratio < 0 and read_specs: + if self.use_mla and tp_ratio < 0 and len(read_specs) == 1: # ..but we still need to notify the other remote ranks that we # have the blocks we need so they can update the request state. notif_id = f"{meta.remote.request_id}:{self.world_size}".encode() diff --git a/vllm/distributed/nixl_utils.py b/vllm/distributed/nixl_utils.py index 634d59976f70..3879bee0ce29 100644 --- a/vllm/distributed/nixl_utils.py +++ b/vllm/distributed/nixl_utils.py @@ -21,7 +21,7 @@ def _maybe_set_ucx_rcache_limit() -> None: if "UCX_RCACHE_MAX_UNRELEASED" in os.environ: return - if "nixl" in sys.modules or "rixl" in sys.modules: + if "nixl" in sys.modules or "nixl_rocm" in sys.modules: logger.warning_once( "NIXL was already imported, we can't reset " "UCX_RCACHE_MAX_UNRELEASED. " @@ -36,8 +36,12 @@ def _maybe_set_ucx_rcache_limit() -> None: os.environ["UCX_RCACHE_MAX_UNRELEASED"] = "1024" +def _get_nixl_package_name() -> str: + return "nixl_rocm" if current_platform.is_rocm() else "nixl" + + def _get_nixl_module_name(name: str) -> str: - package_name = "rixl" if current_platform.is_rocm() else "nixl" + package_name = _get_nixl_package_name() if name == "nixlXferTelemetry": return f"{package_name}._bindings" return f"{package_name}._api" @@ -80,11 +84,11 @@ def __getattr__(name: str) -> Any: def is_nixl_available() -> bool: - """Lightweight check for nixl/rixl package without importing it.""" + """Lightweight check for the platform's NIXL package without importing it.""" import importlib.util - pkg = "rixl" if current_platform.is_rocm() else "nixl" - return importlib.util.find_spec(pkg) is not None + pkg = _get_nixl_package_name() + return pkg in sys.modules or importlib.util.find_spec(pkg) is not None __all__ = [ diff --git a/vllm/distributed/weight_transfer/ipc_engine.py b/vllm/distributed/weight_transfer/ipc_engine.py index f1b6070893bf..b191e8c50dba 100644 --- a/vllm/distributed/weight_transfer/ipc_engine.py +++ b/vllm/distributed/weight_transfer/ipc_engine.py @@ -274,7 +274,12 @@ def receive_weights(self, update_info: IPCWeightTransferUpdateInfo) -> None: weight = rebuild_cuda_tensor(*list_args) weights.append((name, weight)) - self.model.load_weights(weights) + from vllm.model_executor.model_loader.mtp_validation import ( + disable_mtp_completeness_check, + ) + + with disable_mtp_completeness_check(): + self.model.load_weights(weights) def shutdown(self) -> None: pass diff --git a/vllm/distributed/weight_transfer/nccl_engine.py b/vllm/distributed/weight_transfer/nccl_engine.py index 5838a8ba8a1d..0ba44f3486e8 100644 --- a/vllm/distributed/weight_transfer/nccl_engine.py +++ b/vllm/distributed/weight_transfer/nccl_engine.py @@ -168,36 +168,41 @@ def receive_weights(self, update_info: NCCLWeightTransferUpdateInfo) -> None: "Call init_transfer_engine() first." ) - if update_info.packed: - # Build iterator of (name, (shape, dtype)) from update_info - def state_dict_info_iterator(): + from vllm.model_executor.model_loader.mtp_validation import ( + disable_mtp_completeness_check, + ) + + with disable_mtp_completeness_check(): + if update_info.packed: + # Build iterator of (name, (shape, dtype)) from update_info + def state_dict_info_iterator(): + for name, dtype_name, shape in zip( + update_info.names, update_info.dtype_names, update_info.shapes + ): + dtype = getattr(torch, dtype_name) + yield (name, (shape, dtype)) + + packed_nccl_broadcast_consumer( + iterator=state_dict_info_iterator(), + group=self.model_update_group, + src=0, + post_unpack_func=self.model.load_weights, + buffer_size_bytes=update_info.packed_buffer_size_bytes, + num_buffers=update_info.packed_num_buffers, + device=self.device, + ) + else: + # Use simple one-by-one broadcasting for name, dtype_name, shape in zip( update_info.names, update_info.dtype_names, update_info.shapes ): dtype = getattr(torch, dtype_name) - yield (name, (shape, dtype)) - - packed_nccl_broadcast_consumer( - iterator=state_dict_info_iterator(), - group=self.model_update_group, - src=0, - post_unpack_func=self.model.load_weights, - buffer_size_bytes=update_info.packed_buffer_size_bytes, - num_buffers=update_info.packed_num_buffers, - device=self.device, - ) - else: - # Use simple one-by-one broadcasting - for name, dtype_name, shape in zip( - update_info.names, update_info.dtype_names, update_info.shapes - ): - dtype = getattr(torch, dtype_name) - weight = torch.empty(shape, dtype=dtype, device=self.device) - self.model_update_group.broadcast( - weight, src=0, stream=torch.cuda.current_stream() - ) - self.model.load_weights([(name, weight)]) - del weight + weight = torch.empty(shape, dtype=dtype, device=self.device) + self.model_update_group.broadcast( + weight, src=0, stream=torch.cuda.current_stream() + ) + self.model.load_weights([(name, weight)]) + del weight def shutdown(self) -> None: if self.model_update_group is not None: diff --git a/vllm/entrypoints/pooling/base/io_processor.py b/vllm/entrypoints/pooling/base/io_processor.py index cdc4c16cacb1..f74b740276e8 100644 --- a/vllm/entrypoints/pooling/base/io_processor.py +++ b/vllm/entrypoints/pooling/base/io_processor.py @@ -2,6 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from collections.abc import Sequence +from concurrent.futures import Executor from typing import Any, Final, cast from vllm import ( @@ -10,21 +11,17 @@ ) from vllm.config import VllmConfig from vllm.entrypoints.chat_utils import ( - ChatCompletionMessageParam, ChatTemplateConfig, - ChatTemplateContentFormatOption, - ConversationMessage, ) -from vllm.entrypoints.serve.engine.typing import RendererChatRequest, RendererRequest -from vllm.inputs import EngineInput, SingletonPrompt from vllm.lora.request import LoRARequest from vllm.renderers import BaseRenderer, merge_kwargs from vllm.renderers.inputs.preprocess import parse_model_prompt, prompt_to_seq -from vllm.tool_parsers import ToolParser +from vllm.utils.async_utils import make_async from vllm.utils.mistral import is_mistral_tokenizer from ..typing import ( - ALLOfflineInputsContext, + AnyOfflineInputsContext, + AnyRenderParam, EncodeChatRenderParams, EncodeCMPLRenderParams, OfflineEncodeInputsContext, @@ -66,14 +63,25 @@ def __init__( chat_template_config.trust_request_chat_template ) + self.template_kwargs = None + self.tool_dicts = None + + # Shared thread pool executor for preprocessing + self._executor: Executor = self.renderer._executor + self.render_async = make_async(self.render, executor=self._executor) + ####################################### # online APIs def create_pooling_params(self, request): return request.to_pooling_params() - def pre_process_online(self, ctx: PoolingServeContext): + def get_request_factory_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: request = ctx.request + renderer = self.renderer + requests: Sequence[AnyRenderParam] if isinstance(request, PoolingChatLikeRequest): self._validate_chat_template( @@ -81,24 +89,87 @@ def pre_process_online(self, ctx: PoolingServeContext): chat_template_kwargs=request.chat_template_kwargs, trust_request_chat_template=self.trust_request_chat_template, ) - _, engine_inputs = self._preprocess_chat_online( - request, - request.messages, - default_template=self.chat_template, - default_template_content_format=self.chat_template_content_format, - default_template_kwargs=None, + + num_requests = 1 + default_template_kwargs = merge_kwargs( + self.template_kwargs, + dict( + tools=self.tool_dicts, + tokenize=is_mistral_tokenizer(renderer.tokenizer), + ), + ) + + mm_config = self.model_config.multimodal_config + tok_params = request.build_tok_params(self.model_config) + chat_params = request.build_chat_params( + self.chat_template, self.chat_template_content_format + ).with_defaults( + default_template_kwargs, + default_media_io_kwargs=( + mm_config.media_io_kwargs if mm_config else None + ), + ) + + params_seq = self._params_to_seq(ctx.pooling_params, num_requests) + seq_lora_requests = self._lora_request_to_seq( + ctx.lora_request, num_requests ) + seq_priority = self._priority_to_seq(ctx.priorities, num_requests) + + requests = [ + EncodeChatRenderParams( + conversations=request.messages, + chat_params=chat_params, + tok_params=tok_params, + prompt_extras=ctx.prompt_extras, + skip_mm_cache=False, + params=params_seq[i], + lora_requests=seq_lora_requests[i], + priorities=seq_priority[i], + ) + for i in range(num_requests) + ] + + return requests + elif isinstance(request, PoolingCompletionLikeRequest): - engine_inputs = self._preprocess_cmpl_online( - request, - prompt_input=request.input, - prompt_embeds=None, + model_config = self.model_config + prompts_seq = prompt_to_seq(request.input) + num_requests = len(prompts_seq) + + parsed_prompts = [ + ( + prompt + if isinstance(prompt, bytes) + else parse_model_prompt(model_config, prompt) + ) + for prompt in prompts_seq + ] + tok_params = request.build_tok_params(model_config) + + params_seq = self._params_to_seq(ctx.pooling_params, num_requests) + seq_lora_requests = self._lora_request_to_seq( + ctx.lora_request, num_requests ) + seq_priority = self._priority_to_seq(ctx.priorities, num_requests) + + requests = [ + EncodeCMPLRenderParams( + prompts=parsed_prompts[i], + tok_params=tok_params, + prompt_extras=ctx.prompt_extras, + skip_mm_cache=False, + params=params_seq[i], + lora_requests=seq_lora_requests[i], + priorities=seq_priority[i], + ) + for i in range(num_requests) + ] + + return requests else: raise ValueError(f"Invalid {self.name} request type") - ctx.engine_inputs = engine_inputs - def post_process_online( self, ctx: PoolingServeContext, @@ -109,7 +180,7 @@ def post_process_online( # offline APIs def get_request_factory_offline( - self, ctx: ALLOfflineInputsContext + self, ctx: AnyOfflineInputsContext ) -> tuple[RequestFactory, int]: assert isinstance(ctx, OfflineEncodeInputsContext) @@ -209,84 +280,6 @@ def render( priorities=render_params["priorities"], ) - def _preprocess_cmpl_online( - self, - request: RendererRequest, - prompt_input: str | list[str] | list[int] | list[list[int]] | None, - prompt_embeds: bytes | list[bytes] | None, - ) -> list[EngineInput]: - renderer = self.renderer - model_config = self.model_config - - prompts = list[SingletonPrompt | bytes]() - if prompt_embeds is not None: # embeds take higher priority - prompts.extend(prompt_to_seq(prompt_embeds)) - if prompt_input is not None: - prompts.extend(prompt_to_seq(prompt_input)) - - parsed_prompts = [ - ( - prompt - if isinstance(prompt, bytes) - else parse_model_prompt(model_config, prompt) - ) - for prompt in prompts - ] - tok_params = request.build_tok_params(model_config) - - return renderer.render_cmpl( - parsed_prompts, - tok_params, - prompt_extras={ - k: v - for k in ("mm_processor_kwargs", "cache_salt") - if (v := getattr(request, k, None)) is not None - }, - ) - - def _preprocess_chat_online( - self, - request: RendererChatRequest, - messages: list[ChatCompletionMessageParam], - default_template: str | None, - default_template_content_format: ChatTemplateContentFormatOption, - default_template_kwargs: dict[str, Any] | None, - tool_dicts: list[dict[str, Any]] | None = None, - tool_parser: type[ToolParser] | None = None, - ) -> tuple[list[ConversationMessage], list[EngineInput]]: - renderer = self.renderer - - default_template_kwargs = merge_kwargs( - default_template_kwargs, - dict( - tools=tool_dicts, - tokenize=is_mistral_tokenizer(renderer.tokenizer), - ), - ) - - mm_config = self.model_config.multimodal_config - - tok_params = request.build_tok_params(self.model_config) - chat_params = request.build_chat_params( - default_template, default_template_content_format - ).with_defaults( - default_template_kwargs, - default_media_io_kwargs=(mm_config.media_io_kwargs if mm_config else None), - ) - - (conversation,), (engine_input,) = renderer.render_chat( - [messages], - chat_params, - tok_params, - prompt_extras={ - k: v - for k in ("mm_processor_kwargs", "cache_salt") - if (v := getattr(request, k, None)) is not None - }, - ) - - return conversation, [engine_input] - def _validate_chat_template( self, request_chat_template: str | None, @@ -341,10 +334,13 @@ def _lora_request_to_seq( def _priority_to_seq( self, - priority: Sequence[int] | None, + priority: int | Sequence[int] | None, num_requests: int, ) -> Sequence[int]: if priority is not None: + if isinstance(priority, int): + return [priority] * num_requests + if len(priority) != num_requests: raise ValueError( f"The lengths of prompts ({num_requests}) " diff --git a/vllm/entrypoints/pooling/base/serving.py b/vllm/entrypoints/pooling/base/serving.py index 79f93a148049..36b1c22e61a5 100644 --- a/vllm/entrypoints/pooling/base/serving.py +++ b/vllm/entrypoints/pooling/base/serving.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project - +import asyncio from abc import ABC, abstractmethod from collections.abc import AsyncGenerator, Mapping from concurrent.futures import Executor @@ -63,9 +63,6 @@ def __init__( # Shared thread pool executor for preprocessing and postprocessing. self._executor: Executor = self.renderer._executor - self._preprocessing_async = make_async( - self._preprocessing, executor=self._executor - ) self._postprocessing_async = make_async( self._postprocessing, executor=self._executor ) @@ -77,7 +74,7 @@ async def __call__( ) -> Response: io_processor = self.get_io_processor(request) ctx = await self._init_ctx(io_processor, request, raw_request) - await self._preprocessing_async(io_processor, ctx) + await self._preprocessing(io_processor, ctx) await self._prepare_generators(ctx) await self._collect_batch(ctx) return await self._postprocessing_async(io_processor, ctx) @@ -86,11 +83,17 @@ async def __call__( def get_io_processor(self, request: AnyPoolingRequest) -> PoolingIOProcessor: raise NotImplementedError - @torch.inference_mode() - def _preprocessing( + async def _preprocessing( self, io_processor: PoolingIOProcessor, ctx: PoolingServeContext ): - return io_processor.pre_process_online(ctx) + requests = io_processor.get_request_factory_online(ctx) + + if len(requests) == 0: + raise ValueError("You must pass at least one prompt") + + ctx.engine_inputs = await asyncio.gather( + *[io_processor.render_async(request) for request in requests] + ) @torch.inference_mode() def _postprocessing( @@ -110,16 +113,26 @@ async def _init_ctx( await self._check_model(request) pooling_params = io_processor.create_pooling_params(request) + lora_request = self._maybe_get_adapters(request) + priorities = getattr(request, "priority", 0) + prompt_extras = { + k: v + for k in ("mm_processor_kwargs", "cache_salt", "chat_template_kwargs") + if (v := getattr(request, k, None)) is not None + } + ctx = PoolingServeContext( request=request, raw_request=raw_request, model_name=model_name, pooling_params=pooling_params, request_id=request_id, + lora_request=lora_request, + priorities=priorities, + prompt_extras=prompt_extras, ) self._validate_request(ctx) - ctx.lora_request = self._maybe_get_adapters(ctx.request) return ctx async def _prepare_generators( @@ -127,7 +140,7 @@ async def _prepare_generators( ctx: PoolingServeContext, ): if ctx.engine_inputs is None: - raise ValueError("Engine prompts not available") + raise ValueError("Engine inputs not available") generators: list[AsyncGenerator[PoolingRequestOutput, None]] = [] @@ -139,12 +152,7 @@ async def _prepare_generators( assert ctx.pooling_params is not None pooling_params = ctx.pooling_params - - if isinstance(pooling_params, list): - for params in pooling_params: - params.verify(self.model_config) - else: - pooling_params.verify(self.model_config) + pooling_params.verify(self.model_config) for i, engine_input in enumerate(ctx.engine_inputs): prompt_request_id = ( @@ -153,26 +161,20 @@ async def _prepare_generators( else ctx.prompt_request_ids[i] ) - params = ( - pooling_params[i] - if isinstance(pooling_params, list) - else pooling_params - ) - self._log_inputs( prompt_request_id, - engine_input, - params=params, + engine_input["prompts"], + params=engine_input["params"], lora_request=ctx.lora_request, ) generator = self.engine_client.encode( - engine_input, - params, - prompt_request_id, - lora_request=ctx.lora_request, + request_id=prompt_request_id, + prompt=engine_input["prompts"], + pooling_params=engine_input["params"], + lora_request=engine_input["lora_requests"], + priority=engine_input["priorities"], trace_headers=trace_headers, - priority=getattr(ctx.request, "priority", 0), ) generators.append(generator) diff --git a/vllm/entrypoints/pooling/classify/api_router.py b/vllm/entrypoints/pooling/classify/api_router.py index 9e016a72e843..fc86f49d6fd1 100644 --- a/vllm/entrypoints/pooling/classify/api_router.py +++ b/vllm/entrypoints/pooling/classify/api_router.py @@ -16,8 +16,11 @@ router = APIRouter() -def classify(request: Request) -> ServingClassification | None: - return request.app.state.serving_classification +def classify(request: Request) -> ServingClassification: + handler = getattr(request.app.state, "serving_classification", None) + if handler is None: + raise NotImplementedError("The model does not support Classification API") + return handler @router.post("/classify", dependencies=[Depends(validate_json_request)]) @@ -27,7 +30,4 @@ async def create_classify( request: ClassificationRequest, raw_request: Request ) -> Response: handler = classify(raw_request) - if handler is None: - raise NotImplementedError("The model does not support Classification API") - return await handler(request, raw_request) diff --git a/vllm/entrypoints/pooling/embed/api_router.py b/vllm/entrypoints/pooling/embed/api_router.py index 7ffb5840d5bd..adee8a94310b 100644 --- a/vllm/entrypoints/pooling/embed/api_router.py +++ b/vllm/entrypoints/pooling/embed/api_router.py @@ -18,8 +18,11 @@ router = APIRouter() -def embedding(request: Request) -> ServingEmbedding | None: - return request.app.state.serving_embedding +def embedding(request: Request) -> ServingEmbedding: + handler = getattr(request.app.state, "serving_embedding", None) + if handler is None: + raise NotImplementedError("The model does not support Embeddings API") + return handler @router.post( @@ -37,9 +40,6 @@ async def create_embedding( raw_request: Request, ): handler = embedding(raw_request) - if handler is None: - raise NotImplementedError("The model does not support Embeddings API") - return await handler(request, raw_request) @@ -58,7 +58,4 @@ async def create_cohere_embedding( raw_request: Request, ): handler = embedding(raw_request) - if handler is None: - raise NotImplementedError("The model does not support Embeddings API") - return await handler(request, raw_request) diff --git a/vllm/entrypoints/pooling/embed/io_processor.py b/vllm/entrypoints/pooling/embed/io_processor.py index 5c2f825a9675..57728c88377c 100644 --- a/vllm/entrypoints/pooling/embed/io_processor.py +++ b/vllm/entrypoints/pooling/embed/io_processor.py @@ -17,7 +17,7 @@ ChatCompletionMessageParam, CustomChatCompletionMessageParam, ) -from vllm.inputs import EngineInput, tokens_input +from vllm.inputs import tokens_input from vllm.logger import init_logger from vllm.outputs import PoolingOutput, PoolingRequestOutput from vllm.renderers import merge_kwargs @@ -28,11 +28,14 @@ from ..base.io_processor import PoolingIOProcessor from ..scoring.io_processor import JinaRankingIOProcessorMixin from ..typing import ( - ALLOfflineInputsContext, + AnyOfflineInputsContext, + AnyRenderParam, ChunkedEmbeddingMetadata, + EncodeChatRenderParams, OfflineEncodeInputsContext, PoolingChatLikeRequest, PoolingCompletionLikeRequest, + PoolingEngineInput, PoolingServeContext, RequestFactory, ) @@ -76,9 +79,11 @@ def __init__(self, *args, **kwargs): list(self.task_instructions.keys()), ) - def pre_process_online(self, ctx: PoolingServeContext): + def get_request_factory_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: if isinstance(ctx.request, CohereEmbedRequest): - self._pre_process_cohere_online(ctx) + requests = self._get_request_factory_cohere_online(ctx) elif isinstance( ctx.request, ( @@ -88,12 +93,11 @@ def pre_process_online(self, ctx: PoolingServeContext): EmbeddingBatchChatInputRequest, ), ): - self._pre_process_openai_chat_online(ctx) + requests = self._get_request_factory_chat_input_online(ctx) else: - super().pre_process_online(ctx) + requests = super().get_request_factory_online(ctx) - if self.enable_chunked_processing: - self._pre_process_chunked(ctx) + return requests def post_process_online( self, @@ -114,18 +118,21 @@ def post_process_online( # PTAL: examples/pooling/embed/openai_embedding_long_text ################################################################# - def _pre_process_chunked(self, ctx: PoolingServeContext) -> None: + def maybe_pre_process_chunked(self, ctx: PoolingServeContext) -> None: + if not self.enable_chunked_processing: + return None + if ctx.engine_inputs is None: raise ValueError("Engine prompts not available") ctx.original_engine_inputs = ctx.engine_inputs request_id = ctx.request_id max_model_len = self.model_config.max_model_len - chunked_engine_inputs: list[EngineInput] = [] + chunked_engine_inputs: list[PoolingEngineInput] = [] prompt_request_ids: list[str] = [] chunked_embedding_metadata: list[ChunkedEmbeddingMetadata] = [] for prompt_idx, engine_input in enumerate(ctx.engine_inputs): - token_ids = engine_input.get("prompt_token_ids", None) + token_ids = engine_input["prompts"].get("prompt_token_ids", None) if token_ids is None: raise NotImplementedError( "Long Text Embedding with Chunked Processing does " @@ -138,7 +145,12 @@ def _pre_process_chunked(self, ctx: PoolingServeContext) -> None: chunk_list(prompt_token_ids, max_model_len) ): chunked_engine_inputs.append( - tokens_input(prompt_token_ids=chunk_tokens) + PoolingEngineInput( + prompts=tokens_input(prompt_token_ids=chunk_tokens), + params=engine_input["params"], + lora_requests=engine_input["lora_requests"], + priorities=engine_input["priorities"], + ) ) prompt_request_ids.append( f"{request_id}-prompt-{prompt_idx}-chunk-{chunk_idx}" @@ -234,7 +246,7 @@ def _post_process_chunked(self, ctx: PoolingServeContext) -> None: # Get original prompt token IDs for this prompt original_prompt = original_engine_inputs[prompt_idx] - token_ids = original_prompt.get("prompt_token_ids", None) + token_ids = original_prompt["prompts"].get("prompt_token_ids", None) if token_ids is None: raise NotImplementedError( "Long Text Embedding with Chunked Processing does " @@ -262,6 +274,68 @@ def _post_process_chunked(self, ctx: PoolingServeContext) -> None: return None + ################################################################# + # Chat input Request Preprocessing & Postprocessing + ################################################################# + + def _get_request_factory_chat_input_online( + self, + ctx: PoolingServeContext, + ) -> Sequence[AnyRenderParam]: + request = ctx.request + renderer = self.renderer + + self._validate_chat_template( + request_chat_template=request.chat_template, + chat_template_kwargs=request.chat_template_kwargs, + trust_request_chat_template=self.trust_request_chat_template, + ) + + if isinstance( + request, (EmbeddingBatchChatRequest, EmbeddingBatchChatInputRequest) + ): + all_messages = request.messages + else: + all_messages = [request.messages] + num_requests = len(all_messages) + + default_template_kwargs = merge_kwargs( + self.template_kwargs, + dict( + tools=self.tool_dicts, + tokenize=is_mistral_tokenizer(renderer.tokenizer), + ), + ) + + mm_config = self.model_config.multimodal_config + tok_params = request.build_tok_params(self.model_config) + chat_params = request.build_chat_params( + self.chat_template, self.chat_template_content_format + ).with_defaults( + default_template_kwargs, + default_media_io_kwargs=(mm_config.media_io_kwargs if mm_config else None), + ) + + params_seq = self._params_to_seq(ctx.pooling_params, num_requests) + seq_lora_requests = self._lora_request_to_seq(ctx.lora_request, num_requests) + seq_priority = self._priority_to_seq(ctx.priorities, num_requests) + + requests = [ + EncodeChatRenderParams( + conversations=all_messages[i], + chat_params=chat_params, + tok_params=tok_params, + prompt_extras=ctx.prompt_extras, + skip_mm_cache=False, + params=params_seq[i], + lora_requests=seq_lora_requests[i], + priorities=seq_priority[i], + ) + for i in range(num_requests) + ] + + return requests + ################################################################# # Cohere Request Preprocessing & Postprocessing ################################################################# @@ -378,71 +452,9 @@ def create_pooling_params(self, request): ) return super().create_pooling_params(request) - def _pre_process_openai_chat_online( - self, - ctx: PoolingServeContext[ - EmbeddingChatRequest - | EmbeddingBatchChatRequest - | EmbeddingChatInputRequest - | EmbeddingBatchChatInputRequest - ], - ) -> None: - request = ctx.request - self._validate_chat_template( - request_chat_template=request.chat_template, - chat_template_kwargs=request.chat_template_kwargs, - trust_request_chat_template=self.trust_request_chat_template, - ) - - if isinstance( - request, (EmbeddingBatchChatRequest, EmbeddingBatchChatInputRequest) - ): - all_messages = request.messages - else: - all_messages = [request.messages] - ctx.engine_inputs = self._batch_render_openai_chat(request, all_messages) - - def _batch_render_openai_chat( - self, - request: ( - EmbeddingChatRequest - | EmbeddingBatchChatRequest - | EmbeddingChatInputRequest - | EmbeddingBatchChatInputRequest - ), - all_messages: Sequence[list[ChatCompletionMessageParam]], - ) -> list[EngineInput]: - renderer = self.renderer - mm_config = self.model_config.multimodal_config - - tok_params = request.build_tok_params(self.model_config) - chat_params = request.build_chat_params( - self.chat_template, - self.chat_template_content_format, - ).with_defaults( - merge_kwargs( - None, - dict( - tools=None, - tokenize=is_mistral_tokenizer(renderer.tokenizer), - ), - ), - default_media_io_kwargs=(mm_config.media_io_kwargs if mm_config else None), - ) - - _, engine_inputs = renderer.render_chat( - all_messages, - chat_params, - tok_params, - prompt_extras={ - k: v - for k in ("mm_processor_kwargs", "cache_salt") - if (v := getattr(request, k, None)) is not None - }, - ) - return engine_inputs - - def _pre_process_cohere_online(self, ctx: PoolingServeContext) -> None: + def _get_request_factory_cohere_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: """Convert a ``CohereEmbedRequest`` into engine prompts. If a model has a chat template the task instruction are rendered @@ -479,13 +491,12 @@ def _pre_process_cohere_online(self, ctx: PoolingServeContext) -> None: task_prefix = self._get_task_instruction_prefix(input_type) if task_prefix is None: - ctx.engine_inputs = self._preprocess_cohere_text_completion( - request, + return self._get_request_factory_cohere_text_completion( + ctx, texts, truncate_prompt_tokens, truncation_side, ) - return all_messages = [ self._mixed_input_to_messages( @@ -497,27 +508,30 @@ def _pre_process_cohere_online(self, ctx: PoolingServeContext) -> None: for text in texts ] if self._has_chat_template(): - ctx.engine_inputs = self._batch_render_chat( - request, + return self._batch_render_chat( + ctx, all_messages, truncate_prompt_tokens, truncation_side, ) else: - ctx.engine_inputs = self._preprocess_cohere_text_completion( - request, + return self._get_request_factory_cohere_text_completion( + ctx, self._apply_task_instruction(texts, input_type), truncate_prompt_tokens, truncation_side, ) - return task_prefix = self._get_task_instruction_prefix(input_type) all_messages = [ self._mixed_input_to_messages(inp, task_prefix=task_prefix) for inp in input ] - ctx.engine_inputs = self._batch_render_chat( - request, all_messages, truncate_prompt_tokens, truncation_side + + return self._batch_render_chat( + ctx, + all_messages, + truncate_prompt_tokens, + truncation_side, ) def _has_chat_template(self) -> bool: @@ -531,13 +545,14 @@ def _has_chat_template(self) -> bool: is not None ) - def _preprocess_cohere_text_completion( + def _get_request_factory_cohere_text_completion( self, - request: CohereEmbedRequest, + ctx: PoolingServeContext, texts: list[str], truncate_prompt_tokens: int | None, truncation_side: Literal["left", "right"] | None, - ) -> list[EngineInput]: + ) -> Sequence[AnyRenderParam]: + request = ctx.request proxy = EmbeddingCompletionRequest( model=request.model, input=texts, @@ -546,50 +561,32 @@ def _preprocess_cohere_text_completion( truncate_prompt_tokens=truncate_prompt_tokens, truncation_side=truncation_side, ) - return self._preprocess_cmpl_online( - proxy, prompt_input=proxy.input, prompt_embeds=None - ) + ctx.request = proxy + requests = super().get_request_factory_online(ctx) + ctx.request = request + return requests def _batch_render_chat( self, - request: CohereEmbedRequest, + ctx: PoolingServeContext, all_messages: Sequence[list[ChatCompletionMessageParam]], truncate_prompt_tokens: int | None, truncation_side: Literal["left", "right"] | None, - ) -> list[EngineInput]: + ) -> Sequence[AnyRenderParam]: """Batch-render multiple conversations through the chat template.""" - if not all_messages: - return [] - - proxy = EmbeddingChatRequest( + request = ctx.request + proxy = EmbeddingBatchChatRequest( model=request.model, - messages=list(all_messages[0]), + messages=all_messages, dimensions=request.output_dimension, encoding_format="float", truncate_prompt_tokens=truncate_prompt_tokens, truncation_side=truncation_side, ) - - renderer = self.renderer - mm_config = self.model_config.multimodal_config - - tok_params = proxy.build_tok_params(self.model_config) - chat_params = proxy.build_chat_params( - self.chat_template, - self.chat_template_content_format, - ).with_defaults( - merge_kwargs( - None, - dict( - tools=None, - tokenize=is_mistral_tokenizer(renderer.tokenizer), - ), - ), - default_media_io_kwargs=(mm_config.media_io_kwargs if mm_config else None), - ) - - _, engine_inputs = renderer.render_chat(all_messages, chat_params, tok_params) - return engine_inputs + ctx.request = proxy + requests = self._get_request_factory_chat_input_online(ctx) + ctx.request = request + return requests def _validate_input_type(self, input_type: str | None) -> None: """Raise if *input_type* is not supported by this model.""" @@ -640,7 +637,9 @@ class TokenEmbedIOProcessor(PoolingIOProcessor): class JinaRankingTokenEmbedIOProcessor( TokenEmbedIOProcessor, JinaRankingIOProcessorMixin ): - def pre_process_online(self, ctx: PoolingServeContext): + def get_request_factory_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: request = ctx.request if isinstance(request, PoolingCompletionLikeRequest): prompts = request.input @@ -655,20 +654,15 @@ def pre_process_online(self, ctx: PoolingServeContext): query=text_prompts[-1], docs=text_prompts[:-1] ) - engine_inputs = self._preprocess_cmpl_online( - request, - prompt_input=prompt_input, - prompt_embeds=None, - ) + request.input = prompt_input + return super().get_request_factory_online(ctx) elif isinstance(request, PoolingChatLikeRequest): raise ValueError("The JinaForRanking does not support chat Request.") else: raise ValueError(f"Invalid {self.name} request type") - ctx.engine_inputs = engine_inputs - def get_request_factory_offline( - self, ctx: ALLOfflineInputsContext + self, ctx: AnyOfflineInputsContext ) -> tuple[RequestFactory, int]: assert isinstance(ctx, OfflineEncodeInputsContext) if not isinstance(ctx.prompts, Sequence) or len(ctx.prompts) < 2: diff --git a/vllm/entrypoints/pooling/embed/protocol.py b/vllm/entrypoints/pooling/embed/protocol.py index b2912e544f28..96a87fea2018 100644 --- a/vllm/entrypoints/pooling/embed/protocol.py +++ b/vllm/entrypoints/pooling/embed/protocol.py @@ -93,9 +93,9 @@ class EmbeddingBatchChatRequest( ``messages`` instead of introducing a separate batch-specific field. """ - messages: list[Annotated[list[ChatCompletionMessageParam], Field(min_length=1)]] = ( - Field(..., min_length=1) - ) + messages: Sequence[ + Annotated[list[ChatCompletionMessageParam], Field(min_length=1)] + ] = Field(..., min_length=1) def to_pooling_params(self): return PoolingParams( @@ -133,9 +133,9 @@ def normalize_input_messages(cls, data): class EmbeddingBatchChatInputRequest(EmbeddingBatchChatRequest): """OpenAI embeddings request with batched chat conversations in ``input``.""" - input: list[Annotated[list[ChatCompletionMessageParam], Field(min_length=1)]] = ( - Field(..., min_length=1) - ) + input: Sequence[ + Annotated[list[ChatCompletionMessageParam], Field(min_length=1)] + ] = Field(..., min_length=1) @model_validator(mode="before") @classmethod diff --git a/vllm/entrypoints/pooling/embed/serving.py b/vllm/entrypoints/pooling/embed/serving.py index fd8140982ecf..a9f6dd310ec7 100644 --- a/vllm/entrypoints/pooling/embed/serving.py +++ b/vllm/entrypoints/pooling/embed/serving.py @@ -9,6 +9,7 @@ from vllm.outputs import PoolingRequestOutput from vllm.utils.serial_utils import EmbedDType, Endianness +from ..base.io_processor import PoolingIOProcessor from ..base.serving import PoolingServing from ..typing import PoolingServeContext from ..utils import ( @@ -53,6 +54,12 @@ def __init__(self, *args, **kwargs): def init_io_processor(self, *args, **kwargs) -> EmbedIOProcessor: return EmbedIOProcessor(*args, **kwargs) + async def _preprocessing( + self, io_processor: PoolingIOProcessor, ctx: PoolingServeContext + ): + await super()._preprocessing(io_processor, ctx) + self.io_processor.maybe_pre_process_chunked(ctx) + def _build_response( self, ctx: PoolingServeContext, diff --git a/vllm/entrypoints/pooling/offline.py b/vllm/entrypoints/pooling/offline.py index bb3be812c6b0..ade3c039849f 100644 --- a/vllm/entrypoints/pooling/offline.py +++ b/vllm/entrypoints/pooling/offline.py @@ -25,7 +25,7 @@ from .scoring.io_processor import ScoringIOProcessor from .scoring.typing import ScoreInput from .typing import ( - ALLOfflineInputsContext, + AnyOfflineInputsContext, OfflineEncodeInputsContext, OfflineOutputsContext, OfflinePluginInputsContext, @@ -104,7 +104,7 @@ def encode( io_processor = self.pooling_io_processors[pooling_task] - ctx: ALLOfflineInputsContext + ctx: AnyOfflineInputsContext if isinstance(prompts, dict) and "data" in prompts: ctx = OfflinePluginInputsContext( pooling_task=pooling_task, @@ -399,6 +399,9 @@ def _run_tiling_engine( num_requests: int, use_tqdm: bool | Callable[..., tqdm] = True, ): + if num_requests == 0: + raise ValueError("You must pass at least one prompt") + # Keeping max_num_seqs * 2 requests in the core can already saturate the core. # Therefore, keep most requests waiting outside the core. max_requests_in_core = ( diff --git a/vllm/entrypoints/pooling/pooling/api_router.py b/vllm/entrypoints/pooling/pooling/api_router.py index 653a36f699ac..40f7b9eb28a8 100644 --- a/vllm/entrypoints/pooling/pooling/api_router.py +++ b/vllm/entrypoints/pooling/pooling/api_router.py @@ -17,8 +17,11 @@ router = APIRouter() -def pooling(request: Request) -> ServingPooling | None: - return request.app.state.serving_pooling +def pooling(request: Request) -> ServingPooling: + handler = getattr(request.app.state, "serving_pooling", None) + if handler is None: + raise NotImplementedError("The model does not support Pooling API") + return handler @router.post( @@ -33,7 +36,4 @@ def pooling(request: Request) -> ServingPooling | None: @load_aware_call async def create_pooling(request: PoolingRequest, raw_request: Request): handler = pooling(raw_request) - if handler is None: - raise NotImplementedError("The model does not support Pooling API") - return await handler(request, raw_request) diff --git a/vllm/entrypoints/pooling/pooling/io_processor.py b/vllm/entrypoints/pooling/pooling/io_processor.py index b0c38b63a32c..ecc356c477a8 100644 --- a/vllm/entrypoints/pooling/pooling/io_processor.py +++ b/vllm/entrypoints/pooling/pooling/io_processor.py @@ -10,7 +10,9 @@ from ..base.io_processor import PoolingIOProcessor from ..typing import ( - ALLOfflineInputsContext, + AnyOfflineInputsContext, + AnyRenderParam, + EncodeCMPLRenderParams, OfflineEncodeInputsContext, OfflineOutputsContext, OfflinePluginInputsContext, @@ -49,7 +51,9 @@ def __init__(self, *args, **kwargs): ####################################### # online APIs - def pre_process_online(self, ctx: PoolingServeContext): + def get_request_factory_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: assert isinstance(ctx.request, IOProcessorRequest) validated_prompt = self.io_processor.parse_data(ctx.request.data) @@ -66,24 +70,33 @@ def pre_process_online(self, ctx: PoolingServeContext): ) for prompt in prompt_to_seq(raw_prompts) ] + num_requests = len(parsed_prompts) tok_params = ctx.request.build_tok_params(self.model_config) - ctx.engine_inputs = self.renderer.render_cmpl( - parsed_prompts, - tok_params, - prompt_extras={ - k: v - for k in ("mm_processor_kwargs", "cache_salt") - if (v := getattr(ctx.request, k, None)) is not None - }, - ) - pooling_params = self.io_processor.merge_pooling_params() if pooling_params.task is None: pooling_params.task = "plugin" ctx.pooling_params = pooling_params + params_seq = self._params_to_seq(ctx.pooling_params, num_requests) + seq_lora_requests = self._lora_request_to_seq(ctx.lora_request, num_requests) + seq_priority = self._priority_to_seq(ctx.priorities, num_requests) + + requests = [ + EncodeCMPLRenderParams( + prompts=parsed_prompts[i], + tok_params=tok_params, + prompt_extras=ctx.prompt_extras, + skip_mm_cache=False, + params=params_seq[i], + lora_requests=seq_lora_requests[i], + priorities=seq_priority[i], + ) + for i in range(num_requests) + ] + return requests + def post_process_online( self, ctx: PoolingServeContext, @@ -114,7 +127,7 @@ def post_process_online( # offline APIs def get_request_factory_offline( - self, ctx: ALLOfflineInputsContext + self, ctx: AnyOfflineInputsContext ) -> tuple[RequestFactory, int]: assert isinstance(ctx, OfflinePluginInputsContext) assert isinstance(ctx.prompts, dict) and "data" in ctx.prompts diff --git a/vllm/entrypoints/pooling/scoring/api_router.py b/vllm/entrypoints/pooling/scoring/api_router.py index f67b5e912f30..a26b3ca6d0ac 100644 --- a/vllm/entrypoints/pooling/scoring/api_router.py +++ b/vllm/entrypoints/pooling/scoring/api_router.py @@ -20,12 +20,18 @@ logger = init_logger(__name__) -def score(request: Request) -> ServingScores | None: - return request.app.state.serving_scores +def score(request: Request) -> ServingScores: + handler = getattr(request.app.state, "serving_scores", None) + if handler is None: + raise NotImplementedError("The model does not support Score API") + return handler -def rerank(request: Request) -> ServingScores | None: - return request.app.state.serving_scores +def rerank(request: Request) -> ServingScores: + handler = getattr(request.app.state, "serving_scores", None) + if handler is None: + raise NotImplementedError("The model does not support Rerank (Score) API") + return handler @router.post( @@ -40,9 +46,6 @@ def rerank(request: Request) -> ServingScores | None: @load_aware_call async def create_score(request: ScoreRequest, raw_request: Request): handler = score(raw_request) - if handler is None: - raise NotImplementedError("The model does not support Score API") - return await handler(request, raw_request) @@ -77,9 +80,6 @@ async def create_score_v1(request: ScoreRequest, raw_request: Request): @load_aware_call async def do_rerank(request: RerankRequest, raw_request: Request): handler = rerank(raw_request) - if handler is None: - raise NotImplementedError("The model does not support Rerank (Score) API") - return await handler(request, raw_request) diff --git a/vllm/entrypoints/pooling/scoring/io_processor.py b/vllm/entrypoints/pooling/scoring/io_processor.py index e8b3efc32eae..8ebabca57cb0 100644 --- a/vllm/entrypoints/pooling/scoring/io_processor.py +++ b/vllm/entrypoints/pooling/scoring/io_processor.py @@ -6,8 +6,7 @@ import torch.nn.functional as F -from vllm import PoolingParams, PoolingRequestOutput, PromptType, TokensPrompt -from vllm.inputs import EngineInput +from vllm import PoolingParams, PoolingRequestOutput, TokensPrompt from vllm.renderers import TokenizeParams from vllm.renderers.hf import safe_apply_chat_template from vllm.renderers.inputs.preprocess import ( @@ -20,8 +19,11 @@ from ...chat_utils import ChatTemplateResolutionError from ..base.io_processor import PoolingIOProcessor +from ..pooling.protocol import PoolingCompletionRequest from ..typing import ( - ALLOfflineInputsContext, + AnyOfflineInputsContext, + AnyPoolingRequest, + AnyRenderParam, EncodeChatRenderParams, EncodeCMPLRenderParams, OfflineEncodeInputsContext, @@ -167,17 +169,7 @@ def valid_inputs( ) return scoring_data - -class BiEncoderIOProcessor(ScoringIOProcessor): - name = "bi-encoder" - pooling_task: PoolingTask = "embed" - - ####################################### - # online APIs - - def pre_process_online(self, ctx: ScoringServeContext): - request = ctx.request - + def valid_inputs_online(self, request: AnyPoolingRequest): if isinstance(request, ScoreRequest): data_1 = request.data_1 data_2 = request.data_2 @@ -185,9 +177,24 @@ def pre_process_online(self, ctx: ScoringServeContext): data_1 = request.query data_2 = request.documents else: - raise ValueError(f"Invalid {self.name} request type") + raise ValueError(f"Invalid {request.__class__.__name__} request type") scoring_data = self.valid_inputs(data_1, data_2) + return scoring_data + + +class BiEncoderIOProcessor(ScoringIOProcessor): + name = "bi-encoder" + pooling_task: PoolingTask = "embed" + + ####################################### + # online APIs + + def get_request_factory_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: + request = ctx.request + scoring_data = self.valid_inputs_online(request) max_tokens_per_query, max_tokens_per_doc = self._get_token_limits( request=request @@ -197,19 +204,37 @@ def pre_process_online(self, ctx: ScoringServeContext): scoring_data, max_tokens_per_query, max_tokens_per_doc ) - tok_params = request.build_tok_params(self.model_config) - engine_inputs = self._pre_process( - scoring_data, - tok_params, - prompt_extras={ - k: v - for k in ("mm_processor_kwargs", "cache_salt", "chat_template_kwargs") - if (v := getattr(request, k, None)) is not None - }, + data_1 = score_data_to_prompts(scoring_data.data_1, "query", self.model_config) + data_2 = score_data_to_prompts( + scoring_data.data_2, "document", self.model_config ) + prompts = data_1 + data_2 + ctx.n_queries = len(data_1) - ctx.engine_inputs = engine_inputs - ctx.n_queries = len(scoring_data.data_1) + prompts_seq = prompt_to_seq(prompts) + parsed_prompts = [ + parse_model_prompt(self.model_config, prompt) for prompt in prompts_seq + ] + num_requests = len(parsed_prompts) + + tok_params = request.build_tok_params(self.model_config) + params_seq = self._params_to_seq(ctx.pooling_params, num_requests) + seq_lora_requests = self._lora_request_to_seq(ctx.lora_request, num_requests) + seq_priority = self._priority_to_seq(ctx.priorities, num_requests) + + requests = [ + EncodeCMPLRenderParams( + prompts=parsed_prompts[i], + tok_params=tok_params, + prompt_extras=ctx.prompt_extras, + skip_mm_cache=False, + params=params_seq[i], + lora_requests=seq_lora_requests[i], + priorities=seq_priority[i], + ) + for i in range(num_requests) + ] + return requests def post_process_online( self, @@ -226,7 +251,7 @@ def post_process_online( # offline APIs def get_request_factory_offline( - self, ctx: ALLOfflineInputsContext + self, ctx: AnyOfflineInputsContext ) -> tuple[RequestFactory, int]: assert isinstance(ctx, OfflineScoringInputsContext) @@ -267,41 +292,6 @@ def post_process_offline( ####################################### # helpers - def _pre_process( - self, - scoring_data: ScoringData, - tok_params: TokenizeParams, - prompt_extras: dict[str, Any] | None = None, - ) -> Sequence[EngineInput]: - data_1 = score_data_to_prompts(scoring_data.data_1, "query", self.model_config) - data_2 = score_data_to_prompts( - scoring_data.data_2, "document", self.model_config - ) - - return self._preprocess_cmpl_offline( - prompts=data_1 + data_2, tok_params=tok_params, prompt_extras=prompt_extras - ) - - def _preprocess_cmpl_offline( - self, - prompts: PromptType | Sequence[PromptType], - tok_params: TokenizeParams, - prompt_extras: dict[str, Any] | None = None, - ) -> Sequence[EngineInput]: - prompts = prompt_to_seq(prompts) - parsed_prompts = [ - ( - prompt - if isinstance(prompt, bytes) - else parse_model_prompt(self.model_config, prompt) - ) - for prompt in prompts - ] - - return self.renderer.render_cmpl( - parsed_prompts, tok_params, prompt_extras=prompt_extras - ) - def _post_process(self, outputs: list[PoolingRequestOutput], n_queries: int): emb_data_1 = outputs[:n_queries] emb_data_2 = outputs[n_queries:] @@ -432,49 +422,50 @@ def __init__(self, *args, **kwargs): ####################################### # online APIs - def pre_process_online(self, ctx: ScoringServeContext): + def get_request_factory_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: request = ctx.request + scoring_data = self.valid_inputs_online(request) + data_1 = scoring_data.data_1 + data_2 = scoring_data.data_2 + num_requests = len(data_2) - if isinstance(request, ScoreRequest): - data_1 = request.data_1 - data_2 = request.data_2 - elif isinstance(request, RerankRequest): - data_1 = request.query - data_2 = request.documents - else: - raise ValueError(f"Invalid {self.name} request type") - - scoring_data = self.valid_inputs(data_1, data_2) + if len(data_1) == 1: + data_1 = data_1 * num_requests max_tokens_per_query, max_tokens_per_doc = self._get_token_limits( request=request ) tok_params = request.build_tok_params(self.model_config) - pooling_params = self.create_pooling_params(request) - - engine_inputs, pooling_params_list = self._pre_process( - scoring_data, - tok_params, - pooling_params, - chat_template=self.chat_template, - max_tokens_per_query=max_tokens_per_query, - max_tokens_per_doc=max_tokens_per_doc, - prompt_extras={ - k: v - for k in ("mm_processor_kwargs", "cache_salt", "chat_template_kwargs") - if (v := getattr(request, k, None)) is not None - }, - ) + seq_lora_requests = self._lora_request_to_seq(ctx.lora_request, num_requests) + seq_priority = self._priority_to_seq(ctx.priorities, num_requests) + + requests = [ + ScoringRenderParams( + data_1=data_1[i], + data_2=data_2[i], + chat_template=self.chat_template, + max_tokens_per_query=max_tokens_per_query, + max_tokens_per_doc=max_tokens_per_doc, + tok_params=tok_params, + prompt_extras=ctx.prompt_extras, + skip_mm_cache=False, + params=ctx.pooling_params, + lora_requests=seq_lora_requests[i], + priorities=seq_priority[i], + ) + for i in range(num_requests) + ] - ctx.engine_inputs = engine_inputs - ctx.pooling_params = pooling_params_list + return requests ####################################### # offline APIs def get_request_factory_offline( - self, ctx: ALLOfflineInputsContext + self, ctx: AnyOfflineInputsContext ) -> tuple[RequestFactory, int]: assert isinstance(ctx, OfflineScoringInputsContext) @@ -485,10 +476,12 @@ def get_request_factory_offline( if len(data_1) == 1: data_1 = data_1 * num_requests + max_tokens_per_query, max_tokens_per_doc = self._get_token_limits( + pooling_params=ctx.pooling_params + ) tok_params = self.renderer.default_cmpl_tok_params.with_kwargs( **(ctx.tokenization_kwargs or {}) ) - prompt_extras = ctx.pooling_params.extra_kwargs seq_lora_requests = self._lora_request_to_seq(ctx.lora_request, num_requests) @@ -500,6 +493,8 @@ def request_factory() -> RequestGenerator: data_1=data_1[i], data_2=data_2[i], chat_template=ctx.chat_template, + max_tokens_per_query=max_tokens_per_query, + max_tokens_per_doc=max_tokens_per_doc, tok_params=tok_params, prompt_extras=prompt_extras, skip_mm_cache=False, @@ -531,17 +526,13 @@ def render( params = render_params["params"] prompt_extras = render_params["prompt_extras"] - max_tokens_per_query, max_tokens_per_doc = self._get_token_limits( - pooling_params=params - ) - _, engine_prompt = self.get_score_prompt( data_1=render_params["data_1"], data_2=render_params["data_2"], encode_kwargs=tok_params.get_encode_kwargs(), chat_template=render_params["chat_template"], - max_tokens_per_query=max_tokens_per_query, - max_tokens_per_doc=max_tokens_per_doc, + max_tokens_per_query=render_params["max_tokens_per_query"], + max_tokens_per_doc=render_params["max_tokens_per_doc"], chat_template_kwargs=prompt_extras.get("chat_template_kwargs") if prompt_extras else None, @@ -551,29 +542,16 @@ def render( if token_type_ids := engine_prompt.pop("token_type_ids", None): params = params.clone() - compressed = compress_token_type_ids(token_type_ids) - params.extra_kwargs = {"compressed_token_type_ids": compressed} - - engine_input = self.renderer.process_for_engine(engine_prompt, arrival_time) - - return PoolingEngineInput( - prompts=engine_input, - params=params, - lora_requests=render_params["lora_requests"], - priorities=render_params["priorities"], - ) + compressed = compress_token_type_ids( + _apply_post_tokenization_to_token_type_ids( + self.tokenizer, tok_params, token_type_ids + ) + ) + params.extra_kwargs = { + **(params.extra_kwargs or {}), + "compressed_token_type_ids": compressed, + } - def _pre_process( - self, - scoring_data: ScoringData, - tok_params: TokenizeParams, - pooling_params: PoolingParams | None, - chat_template: str | None = None, - max_tokens_per_query: int = 0, - max_tokens_per_doc: int = 0, - prompt_extras: dict[str, Any] | None = None, - ) -> tuple[Sequence[EngineInput], list[PoolingParams]]: - arrival_time = time.time() engine_prompt_extras = ( { k: v @@ -584,55 +562,18 @@ def _pre_process( else None ) - data_1 = scoring_data.data_1 - data_2 = scoring_data.data_2 + if engine_prompt_extras: + target_prompt = extract_target_prompt(self.model_config, engine_prompt) + target_prompt.update(engine_prompt_extras) - if len(data_1) == 1: - data_1 = data_1 * len(data_2) - - if pooling_params is None: - pooling_params = PoolingParams(task="classify") - - pooling_params_list = list[PoolingParams]() - engine_inputs = list[EngineInput]() - for q, d in zip(data_1, data_2): - _, engine_prompt = self.get_score_prompt( - data_1=q, - data_2=d, - encode_kwargs=tok_params.get_encode_kwargs(), - chat_template=chat_template, - max_tokens_per_query=max_tokens_per_query, - max_tokens_per_doc=max_tokens_per_doc, - chat_template_kwargs=prompt_extras.get("chat_template_kwargs") - if prompt_extras - else None, - ) - - token_type_ids = engine_prompt.pop("token_type_ids", None) - tok_params.apply_post_tokenization(self.tokenizer, engine_prompt) - - if token_type_ids is not None: - params = pooling_params.clone() - compressed = compress_token_type_ids( - _apply_post_tokenization_to_token_type_ids( - self.tokenizer, tok_params, token_type_ids - ) - ) - params.extra_kwargs = { - **(params.extra_kwargs or {}), - "compressed_token_type_ids": compressed, - } - pooling_params_list.append(params) - else: - pooling_params_list.append(pooling_params) + engine_input = self.renderer.process_for_engine(engine_prompt, arrival_time) - if engine_prompt_extras: - target_prompt = extract_target_prompt(self.model_config, engine_prompt) - target_prompt.update(engine_prompt_extras) - engine_inputs.append( - self.renderer.process_for_engine(engine_prompt, arrival_time) - ) - return engine_inputs, pooling_params_list + return PoolingEngineInput( + prompts=engine_input, + params=params, + lora_requests=render_params["lora_requests"], + priorities=render_params["priorities"], + ) def get_score_prompt( self, @@ -837,16 +778,27 @@ class JinaRankingIOProcessor(LateInteractionIOProcessor, JinaRankingIOProcessorM name = "jina-reranking-scoring" pooling_task: PoolingTask = "token_embed" - def get_request_factory_offline( - self, ctx: ALLOfflineInputsContext - ) -> tuple[RequestFactory, int]: - assert isinstance(ctx, OfflineScoringInputsContext) + def get_request_factory_online( + self, ctx: PoolingServeContext + ) -> Sequence[AnyRenderParam]: + request = ctx.request + ctx.n_queries = 1 - scoring_data = ctx.scoring_data - prompt_extras = ctx.pooling_params.extra_kwargs + prompt_extras = ctx.prompt_extras + scoring_data = self.valid_inputs_online(request) + + max_tokens_per_query, max_tokens_per_doc = self._get_token_limits( + request=request + ) + + if max_tokens_per_query > 0 or max_tokens_per_doc > 0: + scoring_data = self._truncate_scoring_data( + scoring_data, max_tokens_per_query, max_tokens_per_doc + ) queries = self.ensure_str(scoring_data.data_1) docs = self.ensure_str(scoring_data.data_2) + chat_template_kwargs = ( prompt_extras.get("chat_template_kwargs") if prompt_extras else None ) @@ -868,24 +820,19 @@ def get_request_factory_offline( for q, d in zip(queries, docs) ] - return PoolingIOProcessor.get_request_factory_offline( - self, - OfflineEncodeInputsContext( - pooling_task=self.pooling_task, - prompts=prompts, - tokenization_kwargs=ctx.tokenization_kwargs, - pooling_params=ctx.pooling_params, - lora_request=ctx.lora_request, - priorities=ctx.priorities, - ), - ) + ctx.request = PoolingCompletionRequest(task="token_embed", input=prompts) + requests = PoolingIOProcessor.get_request_factory_online(self, ctx) + ctx.request = request + return requests + + def get_request_factory_offline( + self, ctx: AnyOfflineInputsContext + ) -> tuple[RequestFactory, int]: + assert isinstance(ctx, OfflineScoringInputsContext) + + scoring_data = ctx.scoring_data + prompt_extras = ctx.pooling_params.extra_kwargs - def _pre_process( - self, - scoring_data: ScoringData, - tok_params: TokenizeParams, - prompt_extras: dict[str, Any] | None = None, - ) -> Sequence[EngineInput]: queries = self.ensure_str(scoring_data.data_1) docs = self.ensure_str(scoring_data.data_2) chat_template_kwargs = ( @@ -909,8 +856,16 @@ def _pre_process( for q, d in zip(queries, docs) ] - return self._preprocess_cmpl_offline( - prompts=prompts, tok_params=tok_params, prompt_extras=prompt_extras + return PoolingIOProcessor.get_request_factory_offline( + self, + OfflineEncodeInputsContext( + pooling_task=self.pooling_task, + prompts=prompts, + tokenization_kwargs=ctx.tokenization_kwargs, + pooling_params=ctx.pooling_params, + lora_request=ctx.lora_request, + priorities=ctx.priorities, + ), ) def _post_process(self, outputs: list[PoolingRequestOutput], n_queries: int): diff --git a/vllm/entrypoints/pooling/scoring/serving.py b/vllm/entrypoints/pooling/scoring/serving.py index 5937664d5687..34a2887d5442 100644 --- a/vllm/entrypoints/pooling/scoring/serving.py +++ b/vllm/entrypoints/pooling/scoring/serving.py @@ -190,7 +190,7 @@ def _request_output_to_rerank_response( async def flash_late_interaction(self, *args, **kwargs) -> Response: ctx = await self._init_ctx(self.io_processor, *args, **kwargs) - await self._preprocessing_async(self.io_processor, ctx) + await self._preprocessing(self.io_processor, ctx) # stage 1: encode queries and cache token embeddings on workers. await self._flash_late_interaction_encode_queries(ctx) @@ -211,7 +211,6 @@ async def _flash_late_interaction_encode_queries(self, ctx: ScoringServeContext) query_keys = [f"{ctx.request_id}-query-{i}" for i in range(n_queries)] query_uses = [n_docs if n_queries == 1 else 1] * n_queries - query_pooling_params_list = [] for i in range(n_queries): pooling_params = ctx.pooling_params.clone() pooling_params.late_interaction_params = ( @@ -220,23 +219,21 @@ async def _flash_late_interaction_encode_queries(self, ctx: ScoringServeContext) query_uses=query_uses[i], ) ) - query_pooling_params_list.append(pooling_params) + query_engine_inputs[i]["params"] = pooling_params - assert ( - n_queries - == len(query_pooling_params_list) - == len(query_engine_inputs) - == len(query_keys) - ) + assert n_queries == len(query_engine_inputs) == len(query_keys) query_ctx = ScoringServeContext( request=ctx.request, raw_request=ctx.raw_request, model_name=ctx.model_name, request_id=ctx.request_id, - pooling_params=query_pooling_params_list, + pooling_params=ctx.pooling_params, prompt_request_ids=query_keys, engine_inputs=query_engine_inputs, + lora_request=ctx.lora_request, + priorities=ctx.priorities, + prompt_extras=ctx.prompt_extras, ) await self._prepare_generators(query_ctx) @@ -255,30 +252,27 @@ async def _flash_late_interaction_encode_docs(self, ctx: ScoringServeContext): query_keys = [f"{ctx.request_id}-query-{i}" for i in range(n_queries)] doc_keys = [f"{ctx.request_id}-doc-{i}" for i in range(n_docs)] - doc_pooling_params_list = [] for i in range(n_docs): query_idx = 0 if n_queries == 1 else i pooling_params = ctx.pooling_params.clone() pooling_params.late_interaction_params = build_late_interaction_doc_params( query_key=query_keys[query_idx] ) - doc_pooling_params_list.append(pooling_params) + doc_engine_inputs[i]["params"] = pooling_params - assert ( - n_docs - == len(doc_pooling_params_list) - == len(doc_engine_inputs) - == len(doc_keys) - ) + assert n_docs == len(doc_engine_inputs) == len(doc_keys) doc_ctx = ScoringServeContext( request=ctx.request, raw_request=ctx.raw_request, model_name=ctx.model_name, request_id=ctx.request_id, - pooling_params=doc_pooling_params_list, + pooling_params=ctx.pooling_params, prompt_request_ids=doc_keys, engine_inputs=doc_engine_inputs, + lora_request=ctx.lora_request, + priorities=ctx.priorities, + prompt_extras=ctx.prompt_extras, ) await self._prepare_generators(doc_ctx) diff --git a/vllm/entrypoints/pooling/typing.py b/vllm/entrypoints/pooling/typing.py index b44c9476e789..288f329d09a3 100644 --- a/vllm/entrypoints/pooling/typing.py +++ b/vllm/entrypoints/pooling/typing.py @@ -88,10 +88,13 @@ class PoolingServeContext(Generic[PoolingRequestT]): raw_request: Request | None = None model_name: str request_id: str - pooling_params: PoolingParams | list[PoolingParams] + pooling_params: PoolingParams + lora_request: LoRARequest | None + priorities: int | Sequence[int] | None + prompt_extras: dict[str, Any] | None + created_time: int = field(default_factory=lambda: int(time.time())) - lora_request: LoRARequest | None = None - engine_inputs: Sequence[EngineInput] | None = None + engine_inputs: Sequence["PoolingEngineInput"] | None = None prompt_request_ids: list[str] | None = None result_generator: AsyncGenerator[tuple[int, PoolingRequestOutput], None] | None = ( @@ -100,7 +103,7 @@ class PoolingServeContext(Generic[PoolingRequestT]): final_res_batch: list[PoolingRequestOutput] = field(default_factory=list) ## for Long Text Embedding with Chunked Processing - original_engine_inputs: Sequence[EngineInput] | None = None + original_engine_inputs: Sequence["PoolingEngineInput"] | None = None chunked_embedding_metadata: list[ChunkedEmbeddingMetadata] | None = None ## for bi-encoder & late-interaction @@ -118,7 +121,7 @@ class OfflineInputsContext: pooling_task: PoolingTask tokenization_kwargs: dict[str, Any] | None lora_request: Sequence[LoRARequest | None] | None - priorities: Sequence[int] | None + priorities: int | Sequence[int] | None @dataclass @@ -140,7 +143,7 @@ class OfflinePluginInputsContext(OfflineInputsContext): pooling_params: PoolingParams | Sequence[PoolingParams] | None -ALLOfflineInputsContext: TypeAlias = ( +AnyOfflineInputsContext: TypeAlias = ( OfflineEncodeInputsContext | OfflineScoringInputsContext | OfflinePluginInputsContext @@ -178,6 +181,15 @@ class ScoringRenderParams(RenderParams): data_1: ScoreData data_2: ScoreData chat_template: str | None + max_tokens_per_query: int + max_tokens_per_doc: int + + +AnyRenderParam: TypeAlias = ( + EncodeCMPLRenderParams | EncodeChatRenderParams | ScoringRenderParams +) +RequestGenerator: TypeAlias = Generator[AnyRenderParam] +RequestFactory: TypeAlias = Callable[[], RequestGenerator] class PoolingEngineInput(TypedDict): @@ -185,9 +197,3 @@ class PoolingEngineInput(TypedDict): params: PoolingParams lora_requests: LoRARequest | None priorities: int - - -RequestGenerator: TypeAlias = Generator[ - EncodeCMPLRenderParams | EncodeChatRenderParams | ScoringRenderParams -] -RequestFactory: TypeAlias = Callable[[], RequestGenerator] diff --git a/vllm/entrypoints/serve/utils/api_utils.py b/vllm/entrypoints/serve/utils/api_utils.py index 3b9aeb381226..5fc855d9b5a9 100644 --- a/vllm/entrypoints/serve/utils/api_utils.py +++ b/vllm/entrypoints/serve/utils/api_utils.py @@ -325,12 +325,13 @@ def log_version_and_model(lgr: Logger, version: str, model_name: str) -> None: message = "vLLM server version %s, serving model %s" else: logo_template = Template( - "\n ${b}█ █ █▄ ▄█${r}\n" - " ${o}▄▄${r} ${b}▄█${r} ${b}█ █ █ ▀▄▀ █${r} version ${b}%s${r}\n" - " ${o}█${r}${b}▄█▀${r} ${b}█ █ █ █${r} model ${b}%s${r}\n" - " ${b}▀▀${r} ${b}▀▀▀▀▀ ▀▀▀▀▀ ▀ ▀${r}\n" + "\n ${w}█ █ █▄ ▄█${r}\n" + " ${o}▄▄${r} ${b}▄█${r} ${w}█ █ █ ▀▄▀ █${r} version ${w}%s${r}\n" + " ${o}█${r}${b}▄█▀${r} ${w}█ █ █ █${r} model ${w}%s${r}\n" + " ${b}▀▀${r} ${w}▀▀▀▀▀ ▀▀▀▀▀ ▀ ▀${r}\n" ) colors = { + "w": "\033[1m", # bold, default foreground "o": "\033[93m", # orange "b": "\033[94m", # blue "r": "\033[0m", # reset diff --git a/vllm/model_executor/layers/attention/mla_attention.py b/vllm/model_executor/layers/attention/mla_attention.py index 14662e9fc03d..884c4ca4e687 100644 --- a/vllm/model_executor/layers/attention/mla_attention.py +++ b/vllm/model_executor/layers/attention/mla_attention.py @@ -277,6 +277,7 @@ from vllm.v1.attention.ops.common import cp_lse_ag_out_ar, cp_lse_ag_out_rs from vllm.v1.attention.ops.dcp_alltoall import dcp_a2a_lse_reduce from vllm.v1.attention.ops.merge_attn_states import merge_attn_states +from vllm.v1.attention.ops.triton_merge_attn_states import mask_empty_context from vllm.v1.attention.selector import get_attn_backend from vllm.v1.kv_cache_interface import ( AttentionSpec, @@ -1383,6 +1384,7 @@ class ChunkedContextMetadata: workspace: torch.Tensor token_to_seq: torch.Tensor chunk_total_token: list[int] + has_empty_context: list[bool] # for mla DCP padded_local_chunk_seq_lens: list[list[int]] | None = None @@ -1592,6 +1594,7 @@ def build_mla_chunked_context_metadata( ) chunk_seq_lens = chunk_ends - chunk_starts chunk_seq_lens.clamp_(min=0) + has_empty_context = torch.any(chunk_seq_lens == 0, dim=1).tolist() cu_seq_lens_cpu = torch.zeros( num_chunks, num_prefills + 1, dtype=torch.int32, pin_memory=True @@ -1670,6 +1673,7 @@ def build_mla_chunked_context_metadata( token_to_seq=token_to_seq_cpu.to(device, non_blocking=True), chunk_total_token=chunk_total_token.tolist(), workspace=chunked_prefill_workspace, + has_empty_context=has_empty_context, prefill_tokens_with_context=prefill_tokens_with_context, padded_local_chunk_seq_lens=padded_local_chunk_seq_lens.tolist(), local_context_lens_allranks=local_context_lens_allranks.tolist(), @@ -1692,6 +1696,7 @@ def build_mla_chunked_context_metadata( token_to_seq=token_to_seq_cpu.to(device, non_blocking=True), chunk_total_token=chunk_total_token, workspace=chunked_prefill_workspace, + has_empty_context=has_empty_context, prefill_tokens_with_context=prefill_tokens_with_context, ) @@ -2279,6 +2284,13 @@ def _compute_prefill_context( v=v, ) ) + if prefill_metadata.chunked_context.has_empty_context[i]: + mask_empty_context( + attn_softmax_lse, + attn_output, + prefill_metadata.query_start_loc, + prefill_metadata.chunked_context.cu_seq_lens[i], + ) if output is None: output = attn_output @@ -2429,6 +2441,13 @@ def _context_parallel_compute_prefill_context( v=v, ) ) + if prefill_metadata.chunked_context.has_empty_context[i]: + mask_empty_context( + attn_softmax_lse, + attn_output, + prefill_metadata.query_start_loc, + prefill_metadata.chunked_context.cu_seq_lens[i], + ) if output is None: output = attn_output diff --git a/vllm/model_executor/layers/fused_moe/experts/triton_moe.py b/vllm/model_executor/layers/fused_moe/experts/triton_moe.py index 3196667b3f7a..8610e6b1a70b 100644 --- a/vllm/model_executor/layers/fused_moe/experts/triton_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/triton_moe.py @@ -43,7 +43,10 @@ kFp8Static128BlockSym, kFp8StaticChannelSym, kFp8StaticTensorSym, + kInt4Static, + kInt4Static32, kInt8DynamicTokenSym, + kInt8Static, kInt8StaticChannelSym, ) from vllm.platforms import current_platform @@ -542,40 +545,45 @@ def moe_sum(self, input: torch.Tensor, output: torch.Tensor) -> None: class TritonWNA16Experts(TritonExperts): @staticmethod def _supports_current_device() -> bool: - raise NotImplementedError( - "TritonWNA16Experts is not yet used by an Oracle. " - "This method should not be called." - ) + return current_platform.is_cuda_alike() or current_platform.is_xpu() @staticmethod def _supports_no_act_and_mul() -> bool: - raise NotImplementedError( - "TritonWNA16Experts is not yet used by an Oracle. " - "This method should not be called." - ) + return True @staticmethod def _supports_quant_scheme( weight_key: QuantKey | None, activation_key: QuantKey | None, ) -> bool: - raise NotImplementedError( - "TritonWNA16Experts is not yet used by an Oracle. " - "This method should not be called." - ) + SUPPORTED_W = [ + kInt4Static, + kInt8Static, + kInt4Static32, + # other group sizes? + ] + return weight_key in SUPPORTED_W @staticmethod def _supports_activation(activation: MoEActivation) -> bool: - raise NotImplementedError( - "TritonWNA16Experts is not yet used by an Oracle. " - "This method should not be called." - ) + return activation in [ + MoEActivation.SILU, + MoEActivation.GELU, + MoEActivation.GELU_TANH, + MoEActivation.SWIGLUOAI, + MoEActivation.SWIGLUSTEP, + MoEActivation.SILU_NO_MUL, + MoEActivation.GELU_NO_MUL, + MoEActivation.GELU_TANH_NO_MUL, + MoEActivation.RELU2_NO_MUL, + ] @staticmethod def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool: - raise NotImplementedError( - "TritonWNA16Experts is not yet used by an Oracle. " - "This method should not be called." + # Why? + return not ( + moe_parallel_config.use_fi_nvl_two_sided_kernels + or moe_parallel_config.use_fi_nvl_one_sided_kernels ) def apply( @@ -598,7 +606,9 @@ def apply( ): # Check constraints. if self.quant_config.use_int4_w4a16: - assert hidden_states.size(-1) // 2 == w1.size(2), "Hidden size mismatch" + assert hidden_states.size(-1) // 2 == w1.size(2), ( + f"Hidden size mismatch {hidden_states.size(-1) // 2} == {w1.size(2)}" + ) else: assert hidden_states.size(-1) == w1.size(2), ( f"Hidden size mismatch {hidden_states.size(-1)} != {w1.size(2)}" diff --git a/vllm/model_executor/layers/fused_moe/fused_moe.py b/vllm/model_executor/layers/fused_moe/fused_moe.py index 62cf12ea24e2..0d357fcbf456 100644 --- a/vllm/model_executor/layers/fused_moe/fused_moe.py +++ b/vllm/model_executor/layers/fused_moe/fused_moe.py @@ -31,6 +31,7 @@ ) from vllm.platforms import current_platform from vllm.triton_utils import tl, triton +from vllm.utils.math_utils import next_power_of_2 from vllm.utils.platform_utils import get_device_name_as_file_name from vllm.utils.torch_utils import direct_register_custom_op @@ -1298,7 +1299,11 @@ def get_default_config( bit = 4 if dtype == "int4_w4a16" else 8 use_moe_wna16_cuda = should_moe_wna16_use_cuda(M * topk, block_shape[1], E, bit) if use_moe_wna16_cuda: - config = {"BLOCK_SIZE_M": min(16, M), "SPLIT_K": 1} + config = { + "BLOCK_SIZE_M": min(16, next_power_of_2(M)), + "GROUP_SIZE_M": 1, + "SPLIT_K": 1, + } elif M <= 20: config = {"BLOCK_SIZE_M": 16, "GROUP_SIZE_M": 1, "SPLIT_K": 1} elif M <= 40: diff --git a/vllm/model_executor/layers/fused_moe/oracle/int_wna16.py b/vllm/model_executor/layers/fused_moe/oracle/int_wna16.py index a80013f65baa..9d9d3aa76569 100644 --- a/vllm/model_executor/layers/fused_moe/oracle/int_wna16.py +++ b/vllm/model_executor/layers/fused_moe/oracle/int_wna16.py @@ -24,6 +24,9 @@ MarlinExperts, MarlinExpertsBase, ) +from vllm.model_executor.layers.fused_moe.experts.triton_moe import ( + TritonWNA16Experts, +) from vllm.model_executor.layers.fused_moe.experts.trtllm_mxint4_moe import ( TrtLlmMxint4ExpertsMonolithic, ) @@ -50,6 +53,7 @@ class WNA16MoEBackend(Enum): HUMMING = "HUMMING" CPU = "CPU" FLASHINFER_TRTLLM = "FLASHINFER_TRTLLM" + TRITON = "TRITON" XPU = "XPU" EMULATION = "EMULATION" @@ -76,6 +80,8 @@ def backend_to_kernel_cls( return [BatchedMarlinExperts] elif backend == WNA16MoEBackend.FLASHINFER_TRTLLM: return [TrtLlmMxint4ExpertsMonolithic] + elif backend == WNA16MoEBackend.TRITON: + return [TritonWNA16Experts] elif backend == WNA16MoEBackend.XPU: from vllm.model_executor.layers.fused_moe.experts.xpu_moe import ( XPUExpertsWNA16, @@ -107,19 +113,56 @@ def _get_priority_backends() -> list[WNA16MoEBackend]: if current_platform.is_xpu(): return [WNA16MoEBackend.XPU] - _AVAILABLE_BACKENDS = [ + return [ WNA16MoEBackend.FLASHINFER_TRTLLM, WNA16MoEBackend.MARLIN, WNA16MoEBackend.BATCHED_MARLIN, + WNA16MoEBackend.TRITON, WNA16MoEBackend.HUMMING, WNA16MoEBackend.EMULATION, ] - return _AVAILABLE_BACKENDS + + +def _backend_incompatibility_reason( + backend: WNA16MoEBackend, + quant_config: QuantizationConfig | QuantizationArgs, + may_have_zp: bool, + may_have_bias: bool, +) -> str | None: + if backend == WNA16MoEBackend.FLASHINFER_TRTLLM and (may_have_zp or may_have_bias): + return "zero points and bias are not supported" + + from vllm.model_executor.layers.quantization.auto_awq import AutoAWQConfig + from vllm.model_executor.layers.quantization.auto_gptq import AutoGPTQConfig + from vllm.model_executor.layers.quantization.moe_wna16 import MoeWNA16Config + + if backend == WNA16MoEBackend.TRITON: + if may_have_bias: + return "expert bias is not supported" + if isinstance(quant_config, AutoAWQConfig): + return "the AutoAWQ weight layout is not supported" + if isinstance(quant_config, AutoGPTQConfig) and quant_config.desc_act: + return "GPTQ activation ordering is not supported" + if ( + isinstance(quant_config, QuantizationArgs) + and quant_config.actorder == "group" + ): + return "group activation ordering is not supported" + + if isinstance(quant_config, MoeWNA16Config) and backend in ( + WNA16MoEBackend.MARLIN, + WNA16MoEBackend.BATCHED_MARLIN, + WNA16MoEBackend.EMULATION, + ): + return "the MoeWNA16 checkpoint layout is not supported" + + return None def map_wna16_backend(runner_backend: MoEBackend) -> WNA16MoEBackend: """Map user's MoEBackend to WNA16MoEBackend.""" mapping = { + "triton": WNA16MoEBackend.TRITON, "marlin": WNA16MoEBackend.MARLIN, "humming": WNA16MoEBackend.HUMMING, "flashinfer_trtllm": WNA16MoEBackend.FLASHINFER_TRTLLM, @@ -136,6 +179,9 @@ def map_wna16_backend(runner_backend: MoEBackend) -> WNA16MoEBackend: def select_wna16_moe_backend( config: FusedMoEConfig, weight_key: QuantKey, + quant_config: QuantizationConfig | QuantizationArgs, + may_have_zp: bool, + may_have_bias: bool, ) -> tuple[WNA16MoEBackend, type[mk.FusedMoEExperts]]: """Select the WNA16 MoE backend. @@ -143,6 +189,9 @@ def select_wna16_moe_backend( config: the shared ``FusedMoEConfig`` for this layer. weight_key: The QuantKey describing the weight quantization. Must have int4 or int8 type. + quant_config: Quantization structure and checkpoint format description. + may_have_zp: Whether the integration can provide weight zero points. + may_have_bias: Whether the integration can provide expert bias. Returns: A tuple of (``WNA16MoEBackend``, experts class or ``None``). @@ -189,6 +238,11 @@ def _return_or_raise( runner_backend = config.moe_backend if runner_backend != "auto": requested_backend = map_wna16_backend(runner_backend) + reason = _backend_incompatibility_reason( + requested_backend, quant_config, may_have_zp, may_have_bias + ) + if reason is not None: + raise ValueError(_make_log_unsupported(requested_backend, reason)) return _return_or_raise( requested_backend, config, weight_key, None, activation_format ) @@ -197,6 +251,12 @@ def _return_or_raise( AVAILABLE_BACKENDS = _get_priority_backends() for backend in AVAILABLE_BACKENDS: + reason = _backend_incompatibility_reason( + backend, quant_config, may_have_zp, may_have_bias + ) + if reason is not None: + logger.debug_once(_make_log_unsupported(backend, reason), scope="local") + continue activation_key = None # always BF16 activation for WNA16 MoE for k_cls in backend_to_kernel_cls(backend): supported, reason = k_cls.is_supported_config( @@ -294,6 +354,7 @@ def make_wna16_moe_kernel( allowed_experts: tuple[type[mk.FusedMoEExperts], ...] = ( MarlinExperts, BatchedMarlinExperts, + TritonWNA16Experts, TrtLlmMxint4ExpertsMonolithic, XPUExpertsWNA16, CPUExpertsInt4, @@ -315,6 +376,7 @@ def make_wna16_moe_kernel( assert prepare_finalize is not None logger.info_once("Using %s", prepare_finalize.__class__.__name__, scope="local") + logger.info_once("Using %s", experts_cls.__name__, scope="local") extra_args: dict[str, Any] = {} if backend == WNA16MoEBackend.HUMMING: @@ -329,35 +391,17 @@ def make_wna16_moe_kernel( "is_k_full": is_k_full, } - if experts_cls is XPUExpertsWNA16: - assert ( - prepare_finalize.activation_format == mk.FusedMoEActivationFormat.Standard - ), ( - "XPUExpertsWNA16 only supports the Standard activation format; " - "xpu_fused_moe(is_int4=True) does not implement BatchedExperts." - ) - experts: mk.FusedMoEExperts = XPUExpertsWNA16( - moe_config=moe_config, - quant_config=moe_quant_config, - ) - elif ( - prepare_finalize.activation_format == mk.FusedMoEActivationFormat.BatchedExperts - ): + if prepare_finalize.activation_format == mk.FusedMoEActivationFormat.BatchedExperts: max_num_tokens = prepare_finalize.max_num_tokens_per_rank() assert max_num_tokens is not None - experts = experts_cls( - max_num_tokens=max_num_tokens, - num_dispatchers=prepare_finalize.num_dispatchers(), - moe_config=moe_config, - quant_config=moe_quant_config, - **extra_args, - ) - else: - experts = experts_cls( - moe_config=moe_config, - quant_config=moe_quant_config, - **extra_args, - ) + extra_args["max_num_tokens"] = max_num_tokens + extra_args["num_dispatchers"] = prepare_finalize.num_dispatchers() + + experts = experts_cls( + moe_config=moe_config, + quant_config=moe_quant_config, + **extra_args, + ) return mk.FusedMoEKernel( prepare_finalize, @@ -1034,11 +1078,68 @@ def _humming_wna16_weight_schema( "sym": quant_config.is_sym, } raise TypeError( - "Humming WNA16 MoE requires AutoAWQConfig or AutoGPTQConfig, " + "Humming WNA16 checkpoint schema requires AutoAWQConfig or " + "AutoGPTQConfig, " f"got {type(quant_config).__name__}." ) +def _convert_moe_wna16_humming_tensors( + tensors: dict[str, torch.Tensor], has_zero_point: bool +) -> dict[str, torch.Tensor]: + """Convert MoeWNA16's N-first uint8 packing to Humming's int32 packing.""" + if sys.byteorder != "little": + raise NotImplementedError( + "MoeWNA16 to Humming conversion requires a little-endian host." + ) + + output = { + "weight": tensors["qweight"].contiguous().view(torch.int32), + "weight_scale": tensors["scales"], + } + if has_zero_point: + qzeros = tensors["qzeros"] + output["zero_point"] = ( + qzeros.transpose(-1, -2) + .contiguous() + .view(torch.int32) + .transpose(-1, -2) + .contiguous() + ) + return output + + +class _MoeWNA16HummingWeightSchema: + """Adapter from MoeWNA16's generic packed layout to Humming's layout.""" + + def __init__(self, bits: int, group_size: int, has_zero_point: bool) -> None: + self.bits = bits + self.group_size = group_size + self.has_zero_point = has_zero_point + + def convert_humming( + self, + tensors: dict[str, torch.Tensor], + shape_n_stacks: list[int], + shape_k_stacks: list[int], + param_dtype: torch.dtype, + num_experts: int | None = None, + ) -> tuple[Any, dict[str, torch.Tensor]]: + del shape_n_stacks, shape_k_stacks, num_experts + from vllm.utils.humming import HummingWeightSchema, dtypes + + output = _convert_moe_wna16_humming_tensors( + tensors, has_zero_point=self.has_zero_point + ) + output["weight_scale"] = output["weight_scale"].to(param_dtype) + schema = HummingWeightSchema( + b_dtype=dtypes.DataType.from_str(f"uint{self.bits}"), + weight_scale_group_size=self.group_size, + has_zero_point=self.has_zero_point, + ) + return schema, output + + def _unpack_and_dequant_int4_gptq( w_int32: torch.Tensor, scale: torch.Tensor, @@ -1323,13 +1424,27 @@ def convert_to_wna16_moe_kernel_format( input_dtype: optional activation dtype, usually should be 16 bit. """ if backend == WNA16MoEBackend.HUMMING: + from vllm.model_executor.layers.quantization.moe_wna16 import MoeWNA16Config from vllm.model_executor.layers.quantization.utils.humming_utils import ( convert_to_humming_moe_kernel_format, ) - convert_to_humming_moe_kernel_format( - layer, quant_config=_humming_wna16_weight_schema(quant_config) - ) + if isinstance(quant_config, MoeWNA16Config): + from vllm.utils.humming import HummingInputSchema + + convert_to_humming_moe_kernel_format( + layer, + weight_schema=_MoeWNA16HummingWeightSchema( + bits=quant_config.weight_bits, + group_size=layer.group_size, + has_zero_point=quant_config.has_zp, + ), + input_schema=HummingInputSchema(), + ) + else: + convert_to_humming_moe_kernel_format( + layer, quant_config=_humming_wna16_weight_schema(quant_config) + ) return None if backend in ( @@ -1344,16 +1459,32 @@ def convert_to_wna16_moe_kernel_format( ) if isinstance(quant_config, AutoAWQConfig): - if w13_qzeros is None or w2_qzeros is None: - raise ValueError("AWQ Marlin MoE requires zero-point tensors.") - - weight_bits = quant_config.weight_bits + num_bits = quant_config.weight_bits pack_factor = quant_config.pack_factor group_size = quant_config.group_size + elif isinstance(quant_config, AutoGPTQConfig): + num_bits = quant_config.quant_type.size_bits + pack_factor = quant_config.pack_factor + group_size = quant_config.group_size + actorder = "group" if quant_config.desc_act else None + elif isinstance(quant_config, QuantizationArgs): + num_bits = quant_config.num_bits + pack_factor = 32 // quant_config.num_bits + group_size = quant_config.group_size + actorder = quant_config.actorder + else: + raise TypeError( + "Marlin WNA16 MoE backend requires AutoGPTQConfig, AutoAWQConfig or " + f"QuantizationArgs, got {type(quant_config).__name__}." + ) + + if isinstance(quant_config, AutoAWQConfig): + if w13_qzeros is None or w2_qzeros is None: + raise ValueError("AWQ Marlin MoE requires zero-point tensors.") return _process_awq_weights_marlin( layer, - weight_bits, + num_bits, pack_factor, group_size, input_dtype, @@ -1366,41 +1497,28 @@ def convert_to_wna16_moe_kernel_format( w13_bias, w2_bias, ) - elif isinstance(quant_config, AutoGPTQConfig): - num_bits = quant_config.quant_type.size_bits - pack_factor = quant_config.pack_factor - group_size = quant_config.group_size - actorder = "group" if quant_config.desc_act else None - elif isinstance(quant_config, QuantizationArgs): - num_bits = quant_config.num_bits - pack_factor = 32 // quant_config.num_bits - group_size = quant_config.group_size - actorder = quant_config.actorder else: - raise TypeError( - "Marlin WNA16 MoE backend requires AutoAWQConfig, AutoGPTQConfig or " - f"QuantizationArgs, got {type(quant_config).__name__}." + if w13_g_idx is None or w2_g_idx is None: + raise ValueError("GPTQ Marlin MoE requires g_idx tensors.") + + return _process_weights_marlin( + layer, + input_dtype, + num_bits, + pack_factor, + group_size, + actorder, + w13, + w2, + w13_scale, + w2_scale, + w13_g_idx, + w2_g_idx, + w13_qzeros, + w2_qzeros, + w13_bias, + w2_bias, ) - if w13_g_idx is None or w2_g_idx is None: - raise ValueError("GPTQ Marlin MoE requires g_idx tensors.") - return _process_weights_marlin( - layer, - input_dtype, - num_bits, - pack_factor, - group_size, - actorder, - w13, - w2, - w13_scale, - w2_scale, - w13_g_idx, - w2_g_idx, - w13_qzeros, - w2_qzeros, - w13_bias, - w2_bias, - ) elif backend == WNA16MoEBackend.CPU: return _process_weights_cpu( quant_config, @@ -1482,5 +1600,47 @@ def convert_to_wna16_moe_kernel_format( w13_qzeros, w2_qzeros, ) + elif backend == WNA16MoEBackend.TRITON: + # Two possible input layouts depending on the quantization source: + # + # MoeWNA16 (uint8): (E, N_out, K // bit8_pack) — N-first + # → just view as uint8 (no-op) + # + # AutoGPTQ/compressed-tensors (int32, K-first): + # (E, K // pack32, N_out) + # → transpose to N-first, then view as uint8 to get + # (E, N_out, K // bit8_pack) [int32 = 4 bytes → 4 uint8s] + # Scales: (E, K // gs, N_out) → transpose → (E, N_out, K // gs) + from vllm.model_executor.layers.quantization.auto_gptq import ( + AutoGPTQConfig, + ) + + if isinstance(quant_config, (AutoGPTQConfig, QuantizationArgs)): + # These integrations build in K-first format even when the Triton + # backend is selected. Transpose to N-first first. + w13_uint8 = w13.transpose(1, 2).contiguous().view(torch.uint8) + w2_uint8 = w2.transpose(1, 2).contiguous().view(torch.uint8) + w13_scale = w13_scale.transpose(1, 2).contiguous() + w2_scale = w2_scale.transpose(1, 2).contiguous() + else: + # MoeWNA16 uses N-first uint8 weights and scales. + w13_uint8 = w13.view(torch.uint8) + w2_uint8 = w2.view(torch.uint8) + return ( + w13_uint8, + w2_uint8, + w13_scale, + w2_scale, + None, + None, + None, + None, + w13_qzeros, + w2_qzeros, + None, + None, + w13_bias, + w2_bias, + ) else: raise ValueError(f"Unsupported wna16 MoE backend: {backend.value}") diff --git a/vllm/model_executor/layers/fused_moe/routed_experts.py b/vllm/model_executor/layers/fused_moe/routed_experts.py index 80e4064ea6b6..8bf0123e57fe 100644 --- a/vllm/model_executor/layers/fused_moe/routed_experts.py +++ b/vllm/model_executor/layers/fused_moe/routed_experts.py @@ -114,6 +114,7 @@ def __init__( self.e_score_correction_bias = e_score_correction_bias self.apply_router_weight_on_input = apply_router_weight_on_input # End random parameters + self._loaded_expert_biases: set[str] = set() self.quant_method = self._get_quant_method( self.layer_name, @@ -691,6 +692,25 @@ def weight_loader( expert_data = param.data if full_load else param.data[expert_id] + if "bias" in weight_name: + self._loaded_expert_biases.add(weight_name.rsplit(".", 1)[-1]) + if shard_id == "w2": + expert_data = self._narrow_expert_data_for_padding( + expert_data, + loaded_weight, + hidden_dim=0, + ) + expert_data.copy_(loaded_weight) + else: + self._load_w13( + shard_id=shard_id, + shard_dim=0, + loaded_weight=loaded_weight, + expert_data=expert_data, + tp_rank=self.moe_config.tp_rank, + ) + return True if return_success else None + # Case input scale: input_scale loading is only supported for fp8 if "input_scale" in weight_name: # this is needed for compressed-tensors only diff --git a/vllm/model_executor/layers/quantization/auto_awq.py b/vllm/model_executor/layers/quantization/auto_awq.py index 1e49ec3387e4..a77fef05254e 100644 --- a/vllm/model_executor/layers/quantization/auto_awq.py +++ b/vllm/model_executor/layers/quantization/auto_awq.py @@ -560,6 +560,9 @@ def __init__( self.wna16_moe_backend, self.experts_cls = select_wna16_moe_backend( moe, kInt4Static, + quant_config=self.quant_config, + may_have_zp=self.quant_config.zero_point, + may_have_bias=True, ) def create_weights( diff --git a/vllm/model_executor/layers/quantization/auto_gptq.py b/vllm/model_executor/layers/quantization/auto_gptq.py index b056f5222af8..f94bbb25b3ca 100644 --- a/vllm/model_executor/layers/quantization/auto_gptq.py +++ b/vllm/model_executor/layers/quantization/auto_gptq.py @@ -489,6 +489,9 @@ def __init__( self.wna16_moe_backend, self.experts_cls = select_wna16_moe_backend( moe, weight_key, + quant_config=self.quant_config, + may_have_zp=True, + may_have_bias=True, ) def create_weights( @@ -577,25 +580,25 @@ def create_weights( set_weight_attrs(w2_scales, extra_weight_attrs) # don't shard the w2 scales when running act order set_weight_attrs(w2_scales, {"load_full_w2": self.quant_config.desc_act}) - # up_proj scales + # up_proj zero points w13_qzeros = torch.nn.Parameter( torch.empty( num_experts, scales_size13, 2 * intermediate_size_per_partition // self.quant_config.pack_factor, - dtype=params_dtype, + dtype=torch.int32, ), requires_grad=False, ) layer.register_parameter("w13_qzeros", w13_qzeros) set_weight_attrs(w13_qzeros, extra_weight_attrs) - # down_proj scales + # down_proj zero points w2_qzeros = torch.nn.Parameter( torch.empty( num_experts, scales_size2, hidden_size // self.quant_config.pack_factor, - dtype=params_dtype, + dtype=torch.int32, ), requires_grad=False, ) @@ -644,6 +647,26 @@ def create_weights( layer.register_parameter("w2_g_idx_sort_indices", w2_g_idx_sort_indices) set_weight_attrs(w2_g_idx_sort_indices, extra_weight_attrs) + # Some GPTQ checkpoints contain expert biases even when the model + # architecture does not declare them. Zero initialization keeps + # checkpoints without biases equivalent to the bias-free path. + w13_bias = torch.nn.Parameter( + torch.zeros( + num_experts, + 2 * intermediate_size_per_partition, + dtype=params_dtype, + ), + requires_grad=False, + ) + layer.register_parameter("w13_bias", w13_bias) + set_weight_attrs(w13_bias, extra_weight_attrs) + w2_bias = torch.nn.Parameter( + torch.zeros(num_experts, hidden_size, dtype=params_dtype), + requires_grad=False, + ) + layer.register_parameter("w2_bias", w2_bias) + set_weight_attrs(w2_bias, extra_weight_attrs) + if self.experts_cls is not None and issubclass( self.experts_cls, FusedMoEExpertsModular ): @@ -651,12 +674,31 @@ def create_weights( layer.workspace = marlin_make_workspace_new(device, 4) def process_weights_after_loading(self, layer: RoutedExperts) -> None: + def replace_or_register(name: str, val: torch.Tensor | None): + if val is None: + return + + if hasattr(layer, name): + replace_parameter(layer, name, val) + else: + layer.register_parameter( + name, torch.nn.Parameter(val, requires_grad=False) + ) + is_a_8bit = self.input_dtype is not None and self.input_dtype.itemsize == 1 - if is_a_8bit: - assert self.quant_config.quant_type.size_bits == 8, ( - "W8A8-INT8 is not supported by marlin kernel." - ) + assert not is_a_8bit or self.quant_config.quant_type.size_bits == 8, ( + "W8A8-INT8 is not supported by marlin kernel." + ) + + w13_bias = getattr(layer, "w13_bias", None) + if "w13_bias" not in layer._loaded_expert_biases: + layer.register_parameter("w13_bias", None) + w13_bias = None + w2_bias = getattr(layer, "w2_bias", None) + if "w2_bias" not in layer._loaded_expert_biases: + layer.register_parameter("w2_bias", None) + w2_bias = None converted = convert_to_wna16_moe_kernel_format( backend=self.wna16_moe_backend, @@ -669,8 +711,10 @@ def process_weights_after_loading(self, layer: RoutedExperts) -> None: w2_scale=layer.w2_scales, w13_g_idx=layer.w13_g_idx, w2_g_idx=layer.w2_g_idx, - w13_bias=getattr(layer, "w13_bias", None), - w2_bias=getattr(layer, "w2_bias", None), + w13_bias=w13_bias, + w2_bias=w2_bias, + w13_qzeros=getattr(layer, "w13_qzeros", None), + w2_qzeros=getattr(layer, "w2_qzeros", None), ) if converted is None: @@ -703,42 +747,12 @@ def process_weights_after_loading(self, layer: RoutedExperts) -> None: replace_parameter(layer, "w2_g_idx", w2_g_idx) replace_parameter(layer, "w13_g_idx_sort_indices", w13_g_idx_sort_indices) replace_parameter(layer, "w2_g_idx_sort_indices", w2_g_idx_sort_indices) - if w13_qzeros is not None: - replace_parameter(layer, "w13_qzeros", w13_qzeros) - if w2_qzeros is not None: - replace_parameter(layer, "w2_qzeros", w2_qzeros) - if w13_input_global_scale is not None: - if hasattr(layer, "w13_input_global_scale"): - replace_parameter( - layer, "w13_input_global_scale", w13_input_global_scale - ) - else: - layer.register_parameter( - "w13_input_global_scale", - torch.nn.Parameter(w13_input_global_scale, requires_grad=False), - ) - if w2_input_global_scale is not None: - if hasattr(layer, "w2_input_global_scale"): - replace_parameter(layer, "w2_input_global_scale", w2_input_global_scale) - else: - layer.register_parameter( - "w2_input_global_scale", - torch.nn.Parameter(w2_input_global_scale, requires_grad=False), - ) - if w13_bias is not None: - if hasattr(layer, "w13_bias"): - replace_parameter(layer, "w13_bias", w13_bias) - else: - layer.register_parameter( - "w13_bias", torch.nn.Parameter(w13_bias, requires_grad=False) - ) - if w2_bias is not None: - if hasattr(layer, "w2_bias"): - replace_parameter(layer, "w2_bias", w2_bias) - else: - layer.register_parameter( - "w2_bias", torch.nn.Parameter(w2_bias, requires_grad=False) - ) + replace_or_register("w13_input_global_scale", w13_input_global_scale) + replace_or_register("w2_input_global_scale", w2_input_global_scale) + replace_or_register("w13_bias", w13_bias) + replace_or_register("w2_bias", w2_bias) + replace_or_register("w13_qzeros", w13_qzeros) + replace_or_register("w2_qzeros", w2_qzeros) # The modular kernel reads w13_weight/w2_weight; marlin keeps *_qweight. layer.w13_weight = layer.w13_qweight diff --git a/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors_moe/compressed_tensors_moe_wna16_marlin.py b/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors_moe/compressed_tensors_moe_wna16_marlin.py index 2b3317d00f3c..516d94fae6e9 100644 --- a/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors_moe/compressed_tensors_moe_wna16_marlin.py +++ b/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors_moe/compressed_tensors_moe_wna16_marlin.py @@ -92,7 +92,15 @@ def __init__( self.wna16_backend, self.experts_cls = select_wna16_moe_backend( config=self.moe, weight_key=weight_key, + quant_config=self.weight_quant, + may_have_zp=not self.symmetric, + may_have_bias=False, ) + self.is_marlin = self.wna16_backend in [ + WNA16MoEBackend.MARLIN, + WNA16MoEBackend.BATCHED_MARLIN, + ] + self.is_transposed = self.wna16_backend != WNA16MoEBackend.FLASHINFER_TRTLLM def get_weight_shape( self, @@ -118,7 +126,6 @@ def get_weight_shape( "num_groups_w2 must be provided for weight scales/zero_points" ) w13_num_shards = 2 if self.moe.is_act_and_mul else 1 - is_flashinfer = self.wna16_backend == WNA16MoEBackend.FLASHINFER_TRTLLM shape_map = { "w13_weight": { "Flashinfer": ( @@ -177,7 +184,7 @@ def get_weight_shape( ), }, } - backend_key = "Flashinfer" if is_flashinfer else "Marlin" + backend_key = "Marlin" if self.is_transposed else "Flashinfer" return shape_map[weight_name][backend_key] @staticmethod @@ -217,9 +224,8 @@ def create_weights( # Will transpose the loaded weight along the # intermediate and hidden dim sizes. Will # shard for TP along the transposed dims - is_transposed = self.wna16_backend != WNA16MoEBackend.FLASHINFER_TRTLLM extra_weight_attrs.update( - {"is_transposed": is_transposed, "quant_method": self.strategy} + {"is_transposed": self.is_transposed, "quant_method": self.strategy} ) w13_weight = torch.nn.Parameter( @@ -415,7 +421,6 @@ def create_weights( def process_weights_after_loading(self, layer: torch.nn.Module) -> None: # Process weights using the shared oracle infrastructure - is_flashinfer = self.wna16_backend == WNA16MoEBackend.FLASHINFER_TRTLLM converted = convert_to_wna16_moe_kernel_format( backend=self.wna16_backend, layer=layer, @@ -430,6 +435,7 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: w13_qzeros=getattr(layer, "w13_weight_zero_point", None), w2_qzeros=getattr(layer, "w2_weight_zero_point", None), ) + if converted is None: # In-place backends (e.g. Humming) are not wired through this # marlin-only method; fail clearly rather than unpacking None. @@ -437,6 +443,7 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: f"{type(self).__name__} does not support the " f"{self.wna16_backend.value} MoE backend." ) + ( w13_qweight, w2_qweight, @@ -466,7 +473,7 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: replace_parameter(layer, "w2_weight_zero_point", w2_qzeros) # Marlin-specific parameters (not needed for Flashinfer) - if not is_flashinfer: + if self.is_marlin: if w13_g_idx_processed is not None: replace_parameter(layer, "w13_weight_g_idx", w13_g_idx_processed) if w2_g_idx_processed is not None: @@ -510,7 +517,7 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None: # Add Marlin-specific arguments marlin_args: dict[str, Any] = {} - if not is_flashinfer: + if self.is_marlin: marlin_args = { "w13_g_idx": layer.w13_weight_g_idx, "w2_g_idx": layer.w2_weight_g_idx, diff --git a/vllm/model_executor/layers/quantization/moe_wna16.py b/vllm/model_executor/layers/quantization/moe_wna16.py index 23e175fd6249..76f2d740bc73 100644 --- a/vllm/model_executor/layers/quantization/moe_wna16.py +++ b/vllm/model_executor/layers/quantization/moe_wna16.py @@ -15,8 +15,13 @@ ) from vllm.model_executor.layers.fused_moe.config import ( FusedMoEQuantConfig, - int4_w4a16_moe_quant_config, - int8_w8a16_moe_quant_config, +) +from vllm.model_executor.layers.fused_moe.oracle.int_wna16 import ( + WNA16MoEBackend, + convert_to_wna16_moe_kernel_format, + make_wna16_moe_kernel, + make_wna16_moe_quant_config, + select_wna16_moe_backend, ) from vllm.model_executor.layers.fused_moe.unquantized_fused_moe_method import ( UnquantizedFusedMoEMethod, @@ -27,7 +32,15 @@ QuantizationConfig, QuantizeMethodBase, ) -from vllm.model_executor.utils import set_weight_attrs +from vllm.model_executor.layers.quantization.utils.quant_utils import ( + INT4_DTYPE, + INT8_DTYPE, + QuantKey, + kInt4Static32GroupScale, + kInt4StaticGroupScale, + kInt8StaticGroupScale, +) +from vllm.model_executor.utils import replace_parameter, set_weight_attrs from vllm.platforms import current_platform @@ -199,6 +212,34 @@ def __init__(self, quant_config: MoeWNA16Config, moe: "FusedMoEConfig") -> None: super().__init__(moe) self.quant_config = quant_config + num_bits = self.quant_config.weight_bits + group_size = self.quant_config.group_size + + if num_bits == 4: + quant_type = INT4_DTYPE + if group_size == 32: + scale = kInt4Static32GroupScale + else: + scale = kInt4StaticGroupScale + elif num_bits == 8: + assert group_size == -1 + quant_type = INT8_DTYPE + scale = kInt8StaticGroupScale + else: + raise ValueError("MoeWNA16Method only supports int4 and int8 now.") + + weight_key = QuantKey(quant_type, scale) + + # Select WNA16 MoE backend via oracle. + # handle ZP? + self.wna16_backend, self.experts_cls = select_wna16_moe_backend( + config=self.moe, + weight_key=weight_key, + quant_config=self.quant_config, + may_have_zp=self.quant_config.has_zp, + may_have_bias=False, + ) + def create_weights( self, layer: RoutedExperts, @@ -322,23 +363,104 @@ def create_weights( def get_fused_moe_quant_config( self, layer: RoutedExperts ) -> FusedMoEQuantConfig | None: - weight_bits = self.quant_config.weight_bits - has_zp = self.quant_config.has_zp - assert weight_bits == 4 or weight_bits == 8 - config_builder = ( - int4_w4a16_moe_quant_config - if weight_bits == 4 - else int8_w8a16_moe_quant_config - ) + if self.wna16_backend == WNA16MoEBackend.HUMMING: + from vllm.model_executor.layers.quantization.utils.humming_utils import ( + get_humming_moe_quant_config, + ) + + return get_humming_moe_quant_config(layer) - return config_builder( + has_zp = self.quant_config.has_zp + return make_wna16_moe_quant_config( w1_scale=layer.w13_scales, w2_scale=layer.w2_scales, w1_zp=layer.w13_qzeros if has_zp else None, w2_zp=layer.w2_qzeros if has_zp else None, - block_shape=[0, layer.group_size], + group_size=layer.group_size, + num_bits=self.quant_config.weight_bits, ) + def _setup_kernel(self, layer: RoutedExperts): + assert self.experts_cls is not None + self.moe_quant_config = self.get_fused_moe_quant_config(layer) + assert self.moe_quant_config is not None + self.moe_kernel = make_wna16_moe_kernel( + moe_quant_config=self.moe_quant_config, + moe_config=self.moe, + experts_cls=self.experts_cls, + backend=self.wna16_backend, + layer=layer, + routing_tables=layer._expert_routing_tables(), + ) + + def process_weights_after_loading(self, layer: RoutedExperts) -> None: + has_zp = self.quant_config.has_zp + converted = convert_to_wna16_moe_kernel_format( + backend=self.wna16_backend, + layer=layer, + quant_config=self.quant_config, + input_dtype=None, + w13=layer.w13_qweight, + w2=layer.w2_qweight, + w13_scale=layer.w13_scales, + w2_scale=layer.w2_scales, + w13_g_idx=None, + w2_g_idx=None, + w13_qzeros=layer.w13_qzeros if has_zp else None, + w2_qzeros=layer.w2_qzeros if has_zp else None, + ) + + if converted is None: + # Backend rewrote the layer's params in place (e.g. Humming). + self._setup_kernel(layer) + return + + ( + w13_qweight, + w2_qweight, + w13_scales, + w2_scales, + _, + _, + _, + _, + w13_qzeros, + w2_qzeros, + w13_input_global_scale, + w2_input_global_scale, + _, # w13_bias + _, # w2_bias + ) = converted + + # Replace common parameters + replace_parameter(layer, "w13_qweight", w13_qweight) + replace_parameter(layer, "w2_qweight", w2_qweight) + replace_parameter(layer, "w13_scales", w13_scales) + replace_parameter(layer, "w2_scales", w2_scales) + layer.w13_weight = layer.w13_qweight + layer.w2_weight = layer.w2_qweight + + if has_zp: + assert w13_qzeros is not None and w2_qzeros is not None + replace_parameter(layer, "w13_qzeros", w13_qzeros) + replace_parameter(layer, "w2_qzeros", w2_qzeros) + + # Marlin-specific parameters (not needed for Flashinfer) + if self.wna16_backend != WNA16MoEBackend.FLASHINFER_TRTLLM: + # Register input global scales if present + if w13_input_global_scale is not None: + layer.register_parameter( + "w13_input_global_scale", + torch.nn.Parameter(w13_input_global_scale, requires_grad=False), + ) + if w2_input_global_scale is not None: + layer.register_parameter( + "w2_input_global_scale", + torch.nn.Parameter(w2_input_global_scale, requires_grad=False), + ) + + self._setup_kernel(layer) + def apply( self, layer: RoutedExperts, @@ -348,19 +470,44 @@ def apply( shared_experts: SharedExperts | None, shared_experts_input: torch.Tensor | None, ) -> torch.Tensor: - from vllm.model_executor.layers.fused_moe import fused_experts - - return fused_experts( + assert not self.is_monolithic + assert self.moe_kernel is not None + return self.moe_kernel.apply( x, - layer.w13_qweight, - layer.w2_qweight, + layer.w13_weight, + layer.w2_weight, topk_weights=topk_weights, topk_ids=topk_ids, activation=layer.activation, + global_num_experts=layer.global_num_experts, + expert_map=layer.expert_map, apply_router_weight_on_input=layer.apply_router_weight_on_input, + shared_experts=shared_experts, + shared_experts_input=shared_experts_input, + ) + + def apply_monolithic( + self, + layer: RoutedExperts, + x: torch.Tensor, + router_logits: torch.Tensor, + input_ids: torch.Tensor | None = None, + ) -> torch.Tensor: + assert self.is_monolithic + assert self.moe_kernel is not None + return self.moe_kernel.apply_monolithic( + x, + layer.w13_weight, + layer.w2_weight, + router_logits, + activation=layer.activation, global_num_experts=layer.global_num_experts, expert_map=layer.expert_map, - quant_config=self.moe_quant_config, + apply_router_weight_on_input=layer.apply_router_weight_on_input, + num_expert_group=layer.num_expert_group, + topk_group=layer.topk_group, + e_score_correction_bias=layer.e_score_correction_bias, + routed_scaling_factor=layer.routed_scaling_factor, ) @staticmethod diff --git a/vllm/model_executor/layers/quantization/utils/gptq_utils.py b/vllm/model_executor/layers/quantization/utils/gptq_utils.py index 691d80b0b747..c73c42f4ff02 100644 --- a/vllm/model_executor/layers/quantization/utils/gptq_utils.py +++ b/vllm/model_executor/layers/quantization/utils/gptq_utils.py @@ -3,7 +3,7 @@ from collections.abc import Mapping from copy import deepcopy from types import MappingProxyType -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any import regex as re import torch @@ -69,6 +69,20 @@ def get_dynamic_override( return default_value +def flatten_list(lst: list[Any]) -> list[Any]: + output = [] + + def _flatten(lst: list[Any]): + for i in lst: + if isinstance(i, list): + _flatten(i) + else: + output.append(i) + + _flatten(lst) + return output + + def is_layer_gptq_quantized( prefix: str, quantized_layers: list[str], @@ -83,6 +97,8 @@ def is_layer_gptq_quantized( proj_name = prefix.split(".")[-1] + quantized_layers = flatten_list(quantized_layers) + # Fused layers like gate_up_proj or qkv_proj will not be fused # in the safetensors checkpoint. So, we convert the name # from the fused version to unfused + check to make sure that diff --git a/vllm/model_executor/model_loader/mtp_validation.py b/vllm/model_executor/model_loader/mtp_validation.py new file mode 100644 index 000000000000..3f20756abce2 --- /dev/null +++ b/vllm/model_executor/model_loader/mtp_validation.py @@ -0,0 +1,26 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Scoped controls for MTP checkpoint completeness validation.""" + +from collections.abc import Iterator +from contextlib import contextmanager +from contextvars import ContextVar + +_mtp_completeness_check_enabled: ContextVar[bool] = ContextVar( + "mtp_completeness_check_enabled", default=True +) + + +def is_mtp_completeness_check_enabled() -> bool: + """Return whether MTP completeness validation is enabled in this scope.""" + return _mtp_completeness_check_enabled.get() + + +@contextmanager +def disable_mtp_completeness_check() -> Iterator[None]: + """Temporarily disable MTP completeness validation for one weight load.""" + token = _mtp_completeness_check_enabled.set(False) + try: + yield + finally: + _mtp_completeness_check_enabled.reset(token) diff --git a/vllm/model_executor/models/bailing_moe_mtp.py b/vllm/model_executor/models/bailing_moe_mtp.py index da6b1ddb8b66..53679edf89b9 100644 --- a/vllm/model_executor/models/bailing_moe_mtp.py +++ b/vllm/model_executor/models/bailing_moe_mtp.py @@ -20,6 +20,9 @@ ParallelLMHead, VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, @@ -370,7 +373,10 @@ def load_lm_head(loaded_weight: torch.Tensor) -> None: self.model.mtp_start_layer_idx, self.model.mtp_start_layer_idx + self.model.num_mtp_layers, ): - if layer_idx not in loaded_mtp_layers: + if ( + layer_idx not in loaded_mtp_layers + and is_mtp_completeness_check_enabled() + ): raise ValueError( f"Bailing MTP speculative decoding layer {layer_idx} " "weights are missing from checkpoint. Use a checkpoint " diff --git a/vllm/model_executor/models/deepseek_mtp.py b/vllm/model_executor/models/deepseek_mtp.py index 3a0c21fe7d29..65f860a4a3ff 100644 --- a/vllm/model_executor/models/deepseek_mtp.py +++ b/vllm/model_executor/models/deepseek_mtp.py @@ -21,6 +21,9 @@ ParallelLMHead, VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, @@ -516,7 +519,7 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: self.model.mtp_start_layer_idx, self.model.mtp_start_layer_idx + self.model.num_mtp_layers, ): - if layer_idx not in loaded_layers: + if layer_idx not in loaded_layers and is_mtp_completeness_check_enabled(): raise ValueError( f"MTP speculative decoding layer {layer_idx} weights " f"missing from checkpoint. The checkpoint may have " diff --git a/vllm/model_executor/models/step3p5_mtp.py b/vllm/model_executor/models/step3p5_mtp.py index b533a9111dcf..7e95274771da 100644 --- a/vllm/model_executor/models/step3p5_mtp.py +++ b/vllm/model_executor/models/step3p5_mtp.py @@ -15,6 +15,9 @@ ParallelLMHead, VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors @@ -283,7 +286,7 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: and getattr(param, "requires_grad", False) is False } params_need_to_load -= optional_params - if params_need_to_load != loaded_params: + if params_need_to_load != loaded_params and is_mtp_completeness_check_enabled(): missing_params = list(params_need_to_load - loaded_params) param_name_example = missing_params[0] raise RuntimeError( diff --git a/vllm/models/deepseek_v32/nvidia/mtp.py b/vllm/models/deepseek_v32/nvidia/mtp.py index d3d8e5aae9c5..c8f2fcc5ffef 100644 --- a/vllm/models/deepseek_v32/nvidia/mtp.py +++ b/vllm/models/deepseek_v32/nvidia/mtp.py @@ -17,6 +17,9 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, @@ -416,7 +419,7 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: self.model.mtp_start_layer_idx, self.model.mtp_start_layer_idx + self.model.num_mtp_layers, ): - if layer_idx not in loaded_layers: + if layer_idx not in loaded_layers and is_mtp_completeness_check_enabled(): raise ValueError( f"MTP speculative decoding layer {layer_idx} weights " f"missing from checkpoint." diff --git a/vllm/models/deepseek_v4/amd/mtp.py b/vllm/models/deepseek_v4/amd/mtp.py index f5ef4bb06746..a12b401f0d46 100644 --- a/vllm/models/deepseek_v4/amd/mtp.py +++ b/vllm/models/deepseek_v4/amd/mtp.py @@ -32,6 +32,9 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.deepseek_mtp import SharedHead from vllm.model_executor.models.deepseek_v2 import get_spec_layer_idx_from_weight_name @@ -461,7 +464,7 @@ def _find_mtp_layer_idx(name: str) -> int: self.model.mtp_start_layer_idx, self.model.mtp_start_layer_idx + self.model.num_mtp_layers, ): - if layer_idx not in loaded_layers: + if layer_idx not in loaded_layers and is_mtp_completeness_check_enabled(): raise ValueError( f"MTP speculative decoding layer {layer_idx} weights " f"missing from checkpoint. The checkpoint may have " diff --git a/vllm/models/deepseek_v4/nvidia/mtp.py b/vllm/models/deepseek_v4/nvidia/mtp.py index 64715deae99b..8aa3d2c9a299 100644 --- a/vllm/models/deepseek_v4/nvidia/mtp.py +++ b/vllm/models/deepseek_v4/nvidia/mtp.py @@ -37,6 +37,9 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.deepseek_mtp import SharedHead from vllm.model_executor.models.deepseek_v2 import get_spec_layer_idx_from_weight_name @@ -467,7 +470,7 @@ def _find_mtp_layer_idx(name: str) -> int: self.model.mtp_start_layer_idx, self.model.mtp_start_layer_idx + self.model.num_mtp_layers, ): - if layer_idx not in loaded_layers: + if layer_idx not in loaded_layers and is_mtp_completeness_check_enabled(): raise ValueError( f"MTP speculative decoding layer {layer_idx} weights " f"missing from checkpoint. The checkpoint may have " diff --git a/vllm/models/deepseek_v4/xpu/mtp.py b/vllm/models/deepseek_v4/xpu/mtp.py index 8baca78b8ba1..35687ccf169e 100644 --- a/vllm/models/deepseek_v4/xpu/mtp.py +++ b/vllm/models/deepseek_v4/xpu/mtp.py @@ -32,6 +32,9 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.deepseek_mtp import SharedHead from vllm.model_executor.models.deepseek_v2 import get_spec_layer_idx_from_weight_name @@ -469,7 +472,7 @@ def _find_mtp_layer_idx(name: str) -> int: self.model.mtp_start_layer_idx, self.model.mtp_start_layer_idx + self.model.num_mtp_layers, ): - if layer_idx not in loaded_layers: + if layer_idx not in loaded_layers and is_mtp_completeness_check_enabled(): raise ValueError( f"MTP speculative decoding layer {layer_idx} weights " f"missing from checkpoint. The checkpoint may have " diff --git a/vllm/models/inkling/nvidia/mtp.py b/vllm/models/inkling/nvidia/mtp.py index a2559d4cf057..34cde020ed38 100644 --- a/vllm/models/inkling/nvidia/mtp.py +++ b/vllm/models/inkling/nvidia/mtp.py @@ -25,6 +25,9 @@ from vllm.model_executor.layers.linear import ReplicatedLinear from vllm.model_executor.layers.logits_processor import LogitsProcessor from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.utils import maybe_prefix from vllm.sequence import IntermediateTensors @@ -396,7 +399,7 @@ def _load(name: str, weight: torch.Tensor, shard_id: object = None) -> bool: for name in params if name.startswith("model.layers.") or name.startswith("model.chain_norm.") } - if missing := sorted(required - loaded): + if (missing := sorted(required - loaded)) and is_mtp_completeness_check_enabled(): raise ValueError( "Inkling MTP checkpoint is missing required parameters: " + ", ".join(missing) diff --git a/vllm/models/minimax_m3/amd/mtp.py b/vllm/models/minimax_m3/amd/mtp.py index f62face1d2e1..cfb26d7948df 100644 --- a/vllm/models/minimax_m3/amd/mtp.py +++ b/vllm/models/minimax_m3/amd/mtp.py @@ -35,6 +35,9 @@ ParallelLMHead, VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, @@ -322,7 +325,10 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: # Validate that weights were loaded for each MTP layer. for layer_idx in range(self.model.num_mtp_layers): - if layer_idx not in loaded_mtp_layers: + if ( + layer_idx not in loaded_mtp_layers + and is_mtp_completeness_check_enabled() + ): raise ValueError( f"Failed to load MTP layer {layer_idx} weights from checkpoint." ) diff --git a/vllm/models/minimax_m3/nvidia/mtp.py b/vllm/models/minimax_m3/nvidia/mtp.py index e2c7f8821d96..832c872a1006 100644 --- a/vllm/models/minimax_m3/nvidia/mtp.py +++ b/vllm/models/minimax_m3/nvidia/mtp.py @@ -19,6 +19,9 @@ ParallelLMHead, VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.mtp_validation import ( + is_mtp_completeness_check_enabled, +) from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, @@ -304,7 +307,10 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: # Validate that weights were loaded for each MTP layer. for layer_idx in range(self.model.num_mtp_layers): - if layer_idx not in loaded_mtp_layers: + if ( + layer_idx not in loaded_mtp_layers + and is_mtp_completeness_check_enabled() + ): raise ValueError( f"Failed to load MTP layer {layer_idx} weights from checkpoint." ) diff --git a/vllm/multimodal/gpu_ipc_memory.py b/vllm/multimodal/gpu_ipc_memory.py index 15b912064a85..02317490a84b 100644 --- a/vllm/multimodal/gpu_ipc_memory.py +++ b/vllm/multimodal/gpu_ipc_memory.py @@ -17,9 +17,14 @@ """ import threading +from typing import TYPE_CHECKING from vllm.logger import init_logger from vllm.utils.mem_constants import GiB_bytes +from vllm.utils.mem_utils import format_gib + +if TYPE_CHECKING: + from vllm.config.multimodal import MultiModalConfig logger = init_logger(__name__) @@ -145,3 +150,67 @@ def maybe_init_mm_gpu_ipc_pool( api_process_count, ) return pool + + +def reserve_mm_ipc_gpu_memory( + available_kv_cache_memory_bytes: int, + mm_config: "MultiModalConfig | None", + api_process_count: int = 1, +) -> int: + """Carve frontend multimodal GPU memory out of the KV cache. + + Raw decoded frames are bounded by ``mm_ipc_gpu_memory_gb`` and acquired by + the frontend semaphore. Some decoders also keep persistent surfaces around; + reserve a fixed upper bound for those when a GPU backend is configured. + """ + if mm_config is None: + return available_kv_cache_memory_bytes + + from vllm.multimodal.video import ( + PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES, + PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES, + PYNVVIDEOCODEC_MAX_RETAINED_DECODERS, + ) + + raw_frame_reserved_bytes = int(mm_config.mm_ipc_gpu_memory_gb * GiB_bytes) + # Each API server process runs its own decoder surfaces and NVDEC/CUVID CUDA + # context on the GPU, outside the worker memory pool. Reserve that footprint + # per process so gpu_memory_utilization bounds total GPU usage across them. + num_api_servers = max(1, api_process_count) + per_server_decoder_bytes = ( + PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES * PYNVVIDEOCODEC_MAX_RETAINED_DECODERS + + PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES + ) + decoder_reserved_bytes = ( + num_api_servers * per_server_decoder_bytes + if mm_config.use_gpu_video_backend() + else 0 + ) + reserved_bytes = raw_frame_reserved_bytes + decoder_reserved_bytes + if reserved_bytes <= 0: + return available_kv_cache_memory_bytes + + remaining = available_kv_cache_memory_bytes - reserved_bytes + if remaining <= 0: + raise ValueError( + f"frontend multimodal GPU decoding reserves " + f"{format_gib(reserved_bytes)} GiB " + f"({format_gib(raw_frame_reserved_bytes)} GiB raw-frame budget, " + f"{format_gib(decoder_reserved_bytes)} GiB decoder cache budget), " + f"but only {format_gib(available_kv_cache_memory_bytes)} GiB is " + "available for the KV cache. Reduce mm_ipc_gpu_memory_gb, use a " + "different video backend, or increase gpu_memory_utilization." + ) + logger.info_once( + "Reserving %s GiB of GPU memory for frontend multimodal decoding " + "(%s GiB raw-frame semaphore budget, %s GiB decoder+CUDA-context " + "across %d API server(s) @ %s GiB/server); " + "KV cache memory reduced to %s GiB.", + format_gib(reserved_bytes), + format_gib(raw_frame_reserved_bytes), + format_gib(decoder_reserved_bytes), + num_api_servers, + format_gib(per_server_decoder_bytes), + format_gib(remaining), + ) + return remaining diff --git a/vllm/renderers/online_derenderer.py b/vllm/renderers/online_derenderer.py index 3ce4f74d9d82..22e0b45790e3 100644 --- a/vllm/renderers/online_derenderer.py +++ b/vllm/renderers/online_derenderer.py @@ -24,6 +24,7 @@ from vllm.renderers import BaseRenderer from vllm.tokenizers import TokenizerLike from vllm.utils import random_uuid +from vllm.utils.async_utils import make_async logger = init_logger(__name__) @@ -73,10 +74,26 @@ def __init__( self.supports_browsing = False self.supports_code_interpreter = False + # Detokenization, logprob resolution and parsing are CPU-bound; + # offload them in one hop to keep the event loop responsive. + self._derender_chat_async = make_async( + self._derender_chat, executor=renderer._executor + ) + self._derender_completion_async = make_async( + self._derender_completion, executor=renderer._executor + ) + async def derender_chat( self, generate_response: GenerateResponse, chat_request: ChatCompletionRequest | None = None, + ) -> list[ChatCompletionResponseChoice]: + return await self._derender_chat_async(generate_response, chat_request) + + def _derender_chat( + self, + generate_response: GenerateResponse, + chat_request: ChatCompletionRequest | None = None, ) -> list[ChatCompletionResponseChoice]: tokenizer = self.renderer.get_tokenizer() choices: list[ChatCompletionResponseChoice] = [] @@ -172,6 +189,13 @@ async def derender_completion( self, generate_responses: list[GenerateResponse], prompt_tokens: list[int] | None = None, + ) -> tuple[list[CompletionResponseChoice], int, int]: + return await self._derender_completion_async(generate_responses, prompt_tokens) + + def _derender_completion( + self, + generate_responses: list[GenerateResponse], + prompt_tokens: list[int] | None = None, ) -> tuple[list[CompletionResponseChoice], int, int]: n = len(generate_responses) prompt_tokens_list: list[int] = ( diff --git a/vllm/v1/attention/ops/triton_merge_attn_states.py b/vllm/v1/attention/ops/triton_merge_attn_states.py index 14a52ada97fd..ca06c2970b59 100644 --- a/vllm/v1/attention/ops/triton_merge_attn_states.py +++ b/vllm/v1/attention/ops/triton_merge_attn_states.py @@ -9,6 +9,118 @@ float8_info = torch.finfo(current_platform.fp8_dtype()) +def mask_empty_context( + lse: torch.Tensor, + output: torch.Tensor, + query_start_loc: torch.Tensor, + context_start_loc: torch.Tensor, +) -> None: + """Neutralize context chunks that cover no keys before merging. + + A prefill query whose context chunk is empty attended to no keys, so its + partial attention is undefined: the backend leaves the output rows as + uninitialized scratch (which may hold NaN/Inf) even when it reports an LSE + of -inf. Sanitize both here so ``merge_attn_states`` can stay generic: + force the LSE to -inf (zero softmax weight) and zero the undefined output + rows (so a zero weight cannot combine with NaN/Inf). Emptiness is derived + from the context offsets, not from the -inf LSE, so no merge kernel has to + reason about undefined partials. + + Args: + lse: Chunk log-sum-exp, shape [num_heads, num_tokens]. + output: Chunk attention output, shape [num_tokens, num_heads, ...]. + query_start_loc: Prefill query cumulative offsets, shape [num_reqs + 1]. + context_start_loc: Chunk context cumulative offsets, + shape [num_reqs + 1]; an empty chunk has a zero-length span. + """ + num_heads, num_tokens = lse.shape + num_reqs = query_start_loc.shape[0] - 1 + block_size = 128 + # Reserve the worst-case number of request-local blocks. + num_query_blocks = num_tokens // block_size + num_reqs + is_empty = torch.zeros(num_tokens, dtype=torch.bool, device=lse.device) + mask_empty_context_kernel[(num_query_blocks,)]( + lse, + is_empty, + query_start_loc, + context_start_loc, + lse.stride(0), + lse.stride(1), + num_reqs, + NUM_HEADS=num_heads, + BLOCK_SIZE=block_size, + BLOCK_HEADS=8, + num_warps=8, + ) + output.masked_fill_(is_empty[:, None, None], 0.0) + + +@triton.jit +def mask_empty_context_kernel( + lse, + is_empty, + query_start_loc, + context_start_loc, + lse_head_stride, + lse_token_stride, + num_reqs, + NUM_HEADS: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + BLOCK_HEADS: tl.constexpr, +): + query_block_idx = tl.program_id(0) + + lanes = tl.arange(0, 32) + chunk_start = 0 + req_idx = 0 + req_idx_found = False + while (chunk_start < num_reqs) & (not req_idx_found): + req_offsets = chunk_start + lanes + req_mask = req_offsets < num_reqs + query_starts = tl.load(query_start_loc + req_offsets, mask=req_mask) + # Assume the worst-case number of blocks for each request. + req_block_starts = query_starts // BLOCK_SIZE + req_offsets + matched_idx = tl.sum( + (req_mask & (req_block_starts <= query_block_idx)).to(tl.int32) + ) + # matched_idx == 32 means the match is past this warp chunk. + req_idx = chunk_start + matched_idx - 1 + req_idx_found = matched_idx < 32 + chunk_start += 32 + + query_start = tl.load(query_start_loc + req_idx) + query_end = tl.load(query_start_loc + req_idx + 1) + query_len = query_end - query_start + req_first_block = query_start // BLOCK_SIZE + req_idx + block_in_req = query_block_idx - req_first_block + token_offset = block_in_req * BLOCK_SIZE + if token_offset >= query_len: + return + + context_start = tl.load(context_start_loc + req_idx) + context_end = tl.load(context_start_loc + req_idx + 1) + if context_start != context_end: + return + + token_offsets = token_offset + tl.arange(0, BLOCK_SIZE) + token_indices = query_start + token_offsets + token_lse_offsets = token_indices * lse_token_stride + valid_tokens = token_offsets < query_len + tl.store(is_empty + token_indices, True, mask=valid_tokens) + head_offsets = tl.arange(0, BLOCK_HEADS) + for head_start in range(0, NUM_HEADS, BLOCK_HEADS): + head_indices = head_start + head_offsets + lse_ptrs = ( + lse + head_indices[:, None] * lse_head_stride + token_lse_offsets[None, :] + ) + valid_heads = head_indices < NUM_HEADS + tl.store( + lse_ptrs, + float("-inf"), + mask=valid_heads[:, None] & valid_tokens[None, :], + ) + + # Implements section 2.2 of https://www.arxiv.org/pdf/2501.01005 # can be used to combine partial attention results (in the split-KV case) def merge_attn_states( @@ -136,6 +248,9 @@ def merge_attn_states_kernel( if OUTPUT_LSE: out_lse = tl.log(out_se) + max_lse + # Both sides empty (max_lse == -inf) => undefined merge; keep -inf so + # downstream merges continue to treat the token as empty. + out_lse = tl.where(max_lse == float("-inf"), float("-inf"), out_lse) tl.store(output_lse + head_idx * num_tokens + token_idx, out_lse) p_out = tl.load( @@ -159,6 +274,10 @@ def merge_attn_states_kernel( p_scale = p_se / out_se s_scale = s_se / out_se out = p_out * p_scale + s_out * s_scale + # If both sides are empty (max_lse == -inf) the scales are 0/0 = NaN; emit + # zeros rather than NaN. Callers with empty chunks (see mask_empty_context) + # zero those inputs, so this only guards the fully-undefined corner. + out = tl.where(max_lse == float("-inf"), 0.0, out) if USE_FP8: out = out * (1.0 / tl.load(output_scale)) diff --git a/vllm/v1/worker/gpu/cudagraph_utils.py b/vllm/v1/worker/gpu/cudagraph_utils.py index 26f331fe843d..3a174cba80d7 100644 --- a/vllm/v1/worker/gpu/cudagraph_utils.py +++ b/vllm/v1/worker/gpu/cudagraph_utils.py @@ -498,10 +498,7 @@ def create_forward_fn( block_tables, attn_groups, kv_cache_config, - skip_attn=( - desc.cg_mode == CUDAGraphMode.PIECEWISE - and not self.use_breakable_cg - ), + full_cudagraph=desc.cg_mode == CUDAGraphMode.FULL, ) # Capture with dummy rows marked as padding. @@ -510,7 +507,6 @@ def create_forward_fn( def forward_fn(cg_mode: CUDAGraphMode) -> None: batch_descriptor = None if cg_mode == CUDAGraphMode.PIECEWISE: - assert (attn_metadata is not None) == self.use_breakable_cg batch_descriptor = BatchDescriptor( num_tokens=num_tokens, has_lora=has_lora, @@ -593,7 +589,7 @@ def prepare_inputs_to_capture( block_tables: BlockTables, attn_groups: list[list[AttentionGroup]], kv_cache_config: KVCacheConfig, - skip_attn: bool = False, + full_cudagraph: bool, ) -> AttentionState: input_batch = InputBatch.make_dummy(num_reqs, num_tokens, input_buffers) input_block_tables = block_tables.get_dummy_block_tables(num_reqs) @@ -614,15 +610,36 @@ def prepare_inputs_to_capture( ) input_batch.dcp_local_seq_lens = input_buffers.dcp_local_seq_lens[:num_reqs] - attn_metadata = None - if not skip_attn: - attn_metadata = model_state.prepare_attn( - input_batch, - CUDAGraphMode.NONE, - input_block_tables, - slot_mappings, - attn_groups, - kv_cache_config, - for_capture=True, - ) + # NOTE(woosuk): Attention metadata is required not just by standard attention + # kernels, but also by specialized attention-like operations (e.g., Inkling's sconv, + # DSV4 compressor), which maintain their own states and require special metadata + # such as block tables. + # During CUDA graph capture: + # - For FULL CUDA graphs: We set for_capture=True so that both attention and + # attention-like ops produce capturable metadata compatible with CUDA graphs. + # - For PIECEWISE CUDA graphs: We still build attention metadata, but set + # for_capture=False. This is because: + # * Attention-like ops (such as sconv or DSV4 compressor) may not be used as + # breakpoints in PIECEWISE CUDA graphs, so we must generate their attention + # metadata so they can execute and be captured during graph capture. + # * Standard attention ops that are treated as breakpoints will be executed + # eagerly at capture time (not included in the graph itself), and for these, + # setting for_capture=False is essential. Some attention backends + # (like linear attention) cannot generate capturable metadata for prefill, + # so for_capture=False ensures they execute without issue. + # * We assume that attention-like operations intended for capture will still + # produce capturable metadata, even when for_capture=False. While this + # assumption is brittle, it currently works in practice. + # In summary: We always generate attention metadata for both FULL and PIECEWISE + # CUDA graphs, setting for_capture=True for FULL graphs, and for_capture=False + # for PIECEWISE graphs, to ensure correct execution and capture. + attn_metadata = model_state.prepare_attn( + input_batch, + CUDAGraphMode.NONE, + input_block_tables, + slot_mappings, + attn_groups, + kv_cache_config, + for_capture=full_cudagraph, + ) return AttentionState(attn_metadata, slot_mappings_by_layer) diff --git a/vllm/v1/worker/gpu/spec_decode/autoregressive/cudagraph_utils.py b/vllm/v1/worker/gpu/spec_decode/autoregressive/cudagraph_utils.py index 19919043c831..ef3b6e2ed535 100644 --- a/vllm/v1/worker/gpu/spec_decode/autoregressive/cudagraph_utils.py +++ b/vllm/v1/worker/gpu/spec_decode/autoregressive/cudagraph_utils.py @@ -56,10 +56,7 @@ def create_forward_fn( block_tables, attn_groups, kv_cache_config, - skip_attn=( - desc.cg_mode == CUDAGraphMode.PIECEWISE - and not self.use_breakable_cg - ), + full_cudagraph=desc.cg_mode == CUDAGraphMode.FULL, ) return lambda cg_mode: forward_fn( diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index 0c89b2fb13ef..ca63e0a117c7 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -51,12 +51,7 @@ from vllm.logger import init_logger from vllm.lora.request import LoRARequest from vllm.model_executor.warmup.kernel_warmup import kernel_warmup -from vllm.multimodal.video import ( - PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES, - PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES, - PYNVVIDEOCODEC_MAX_RETAINED_DECODERS, - VIDEO_LOADER_REGISTRY, -) +from vllm.multimodal.gpu_ipc_memory import reserve_mm_ipc_gpu_memory from vllm.platforms import current_platform from vllm.profiler.wrapper import CudaProfilerWrapper, TorchProfilerWrapper from vllm.sequence import IntermediateTensors @@ -478,7 +473,11 @@ def determine_available_memory(self) -> int: "correspondingly." ) logger.info(msg) - return self._reserve_mm_ipc_gpu_memory(kv_cache_memory_bytes) + return reserve_mm_ipc_gpu_memory( + kv_cache_memory_bytes, + self.model_config.multimodal_config, + getattr(self.parallel_config, "_api_process_count", 1), + ) # Execute a forward pass with dummy inputs to profile the memory usage # of the model. @@ -589,81 +588,11 @@ def determine_available_memory(self) -> int: suggested_util, ) - return self._reserve_mm_ipc_gpu_memory( - int(self.available_kv_cache_memory_bytes) - ) - - @staticmethod - def _uses_gpu_video_backend(mm_config) -> bool: - video_kwargs = mm_config.media_io_kwargs.get("video", {}) - video_loader_backend = ( - video_kwargs.get("video_backend") or envs.VLLM_VIDEO_LOADER_BACKEND - ) - codec_backend = video_kwargs.get("backend") - return VIDEO_LOADER_REGISTRY.backend_requires_gpu(video_loader_backend) or ( - codec_backend is not None - and VIDEO_LOADER_REGISTRY.backend_requires_gpu(codec_backend) - ) - - def _reserve_mm_ipc_gpu_memory(self, available_kv_cache_memory_bytes: int) -> int: - """Carve frontend multimodal GPU memory out of the KV cache. - - The frontend (API-server) process allocates GPU memory for hardware - multimodal decoding. Raw decoded frames are bounded by - ``mm_ipc_gpu_memory_gb`` and acquired by the frontend semaphore. Some - decoders also keep persistent surfaces around; reserve a fixed upper - bound for those when the corresponding backend is configured. - """ - mm_config = self.model_config.multimodal_config - if mm_config is None: - return available_kv_cache_memory_bytes - - raw_frame_reserved_bytes = int(mm_config.mm_ipc_gpu_memory_gb * GiB_bytes) - # Each api_server_count process runs its OWN decoder surfaces + NVDEC/CUVID - # CUDA context on the GPU, outside this (worker) memory pool. Reserve that - # per-server footprint x api_server_count so gpu_memory_utilization bounds - # TOTAL GPU usage across all API-server processes. Without the multiply, - # HW decode overshoots the budget by ~(api_server_count-1) x per-server and - # OOMs at high gmu, while SW decode (no per-server GPU allocation) does not. - num_api_servers = max(1, getattr(self.parallel_config, "_api_process_count", 1)) - per_server_decoder_bytes = ( - PYNVVIDEOCODEC_DECODER_GPU_MEMORY_BYTES - * PYNVVIDEOCODEC_MAX_RETAINED_DECODERS - + PYNVVIDEOCODEC_CUDA_CONTEXT_BYTES - ) - decoder_reserved_bytes = ( - num_api_servers * per_server_decoder_bytes - if self._uses_gpu_video_backend(mm_config) - else 0 - ) - reserved_bytes = raw_frame_reserved_bytes + decoder_reserved_bytes - if reserved_bytes <= 0: - return available_kv_cache_memory_bytes - - remaining = available_kv_cache_memory_bytes - reserved_bytes - if remaining <= 0: - raise ValueError( - f"frontend multimodal GPU decoding reserves " - f"{format_gib(reserved_bytes)} GiB " - f"({format_gib(raw_frame_reserved_bytes)} GiB raw-frame budget, " - f"{format_gib(decoder_reserved_bytes)} GiB decoder cache budget), " - f"but only {format_gib(available_kv_cache_memory_bytes)} GiB is " - "available for the KV cache. Reduce mm_ipc_gpu_memory_gb, use a " - "different video backend, or increase gpu_memory_utilization." - ) - logger.info_once( - "Reserving %s GiB of GPU memory for frontend multimodal decoding " - "(%s GiB raw-frame semaphore budget, %s GiB decoder+CUDA-context " - "across %d API server(s) @ %s GiB/server); " - "KV cache memory reduced to %s GiB.", - format_gib(reserved_bytes), - format_gib(raw_frame_reserved_bytes), - format_gib(decoder_reserved_bytes), - num_api_servers, - format_gib(per_server_decoder_bytes), - format_gib(remaining), + return reserve_mm_ipc_gpu_memory( + int(self.available_kv_cache_memory_bytes), + self.model_config.multimodal_config, + getattr(self.parallel_config, "_api_process_count", 1), ) - return remaining def get_kv_connector_handshake_metadata( self,