Compare commits

...
Author SHA1 Message Date
Bugen Zhao 7fc60bb26d m3 engine parser (reasoning-only) stage 2026-06-22 20:55:24 +08:00
Palaiologos1453andGitHub 1f2c614c27 Merge branch 'main' into minimax-m3-streaming-reasoning-text-markers-main 2026-06-21 13:07:53 +08:00
Ting SUNandGitHub 183a430c13 [Bugfix][Model Runner V2] Fix min_tokens off-by-one in the V2 GPU sampler (#46243)
Signed-off-by: Ting Sun <suntcrick@gmail.com>
2026-06-21 05:06:49 +00:00
MattandGitHub a346d589f5 [Bugfix] Fix NVFP4/OCP MX MoE emulation (#46254)
Signed-off-by: Matthew Wong <Matthew.Wong2@amd.com>
2026-06-20 23:13:10 -05:00
Nick HillandGitHub 7df3d7dada [Core] Ensure memory is pinned prior to async h2d copy (#45424)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-06-20 20:02:24 -07:00
8dd1b702f2 [Misc] Fix stale doc URL and docstring module path (#35530)
Signed-off-by: umut-polat <52835619+umut-polat@users.noreply.github.com>
Co-authored-by: Flora Feng <4florafeng@gmail.com>
2026-06-20 23:57:01 +00:00
f57ac274b2 [Render] Add reasoning/tool parsing to /derender + fix byte-fallback FFFD (#45919)
Signed-off-by: aoshen524 <aoshen524@gmail.com>
Co-authored-by: Martin Hickey <martin.hickey@ie.ibm.com>
2026-06-20 19:43:32 -04:00
6e919960af [Perf] Skip/shrink all_token_ids copy in scheduler for non-async and V2 runner (#45840)
Signed-off-by: amanchugh89 <amanchugh.89@gmail.com>
Signed-off-by: Nick Hill <nickhill123@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
2026-06-20 22:36:57 +00:00
Jonathan ChenandGitHub c88d3d4775 [SimpleCPUOffloadConnector] PCP + DCP support (#39831)
Signed-off-by: Jonathan Chen <chenleejonathan@gmail.com>
2026-06-20 15:01:06 -07:00
Yifan QiaoandGitHub ab7fcbdd5d [Perf][KVConnector][Mooncake] Compact chunk-hash keys and zero-copy lookup wire format (#45969) 2026-06-20 15:00:11 -07:00
3b4a76b63f [KV-Offloading] : Expose CPU cache usage metric (#45737)
Signed-off-by: Varun Sundar Rabindranath <varun-sundar-rabindranath@h100-01.nemg-001.lab.rdu2.dc.redhat.com>
Signed-off-by: <>
Co-authored-by: Varun Sundar Rabindranath <varun-sundar-rabindranath@h100-01.nemg-001.lab.rdu2.dc.redhat.com>
2026-06-20 21:21:55 +00:00
cc22621b51 [KV Offload] Support packed HMA KV cache layout (#46205)
Signed-off-by: Lucas Wilkinson <lwilkins@redhat.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
Co-authored-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
2026-06-20 21:19:40 +00:00
77148992cf [Bugfix] Move extract_layer_index back inside is_v32 guard (#46199)
Signed-off-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-06-20 21:19:10 +00:00
891cc4b9c5 [Frontend] Report cache usage in Anthropic /v1/messages API (#40912)
Signed-off-by: mistral0105 <zhangshuoming17@mails.ucas.ac.cn>
Signed-off-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
Co-authored-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
2026-06-20 21:12:48 +00:00
TJianandGitHub 1bdf9810aa [ROCm] [Bugfix] Bugfix ROCm Sparse Indexer (#46222)
Signed-off-by: tjtanaa <tunjian.tan@embeddedllm.com>
2026-06-20 13:38:42 -07:00
Tyler Michael SmithGitHubClaudeCodexmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>kourosh hakhamaneshi
ebfbcfe46a Stop setting CUDA_VISIBLE_DEVICES internally in vLLM, add device_ids arg (#45026)
Signed-off-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Codex <codex@openai.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Co-authored-by: kourosh hakhamaneshi <kouroshHakha@users.noreply.github.com>
2026-06-20 13:38:10 -07:00
e9de72fe6c [Bugfix] Guard model_config access in _log_compilation_config (#46198)
Signed-off-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-06-20 19:26:38 +00:00
L丶GitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
d272418f45 [Perf] Optimize Qwen3-VL multi-video prompt processing (#46026)
Signed-off-by: Sirius29 <422058530@qq.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-20 07:09:18 -07:00
Sumanth R HegdeandGitHub 7ff7f5c8eb Revert "Fix Stale Encoder Cache After Weight Update" (#46125) 2026-06-20 07:09:09 -07:00
Palaiologos1453andGitHub f24d8d5bb4 Merge branch 'main' into minimax-m3-streaming-reasoning-text-markers-main 2026-06-20 20:44:14 +08:00
MattGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
dced290769 [Hardware][AMD][CI] Fix e2e core test group (#46024)
Signed-off-by: Matthew Wong <Matthew.Wong2@amd.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-20 02:04:35 -05:00
JasonLi314andGitHub 93bad11912 [Bugfix] Fix gridDim.y overflow for large row counts (#45255)
Signed-off-by: Jason Li <li.jason.cs@gmail.com>
2026-06-19 23:27:45 -04:00
djramicandGitHub 0fbf42af84 [ROCm] Fix VRAM not freed in test_phi3v (#46046)
Signed-off-by: Djordje Ramic <djoramic@amd.com>
2026-06-19 17:20:59 -05:00
Charlie FuandGitHub e6cd8913dd [ROCm][CI] Skip Qwen3.5-35B-A3B-MXFP4-AITER-TP2 for non gfx950 (#46109)
Signed-off-by: charlifu <charlifu@amd.com>
2026-06-19 17:20:10 -05:00
Ben BrowningandGitHub 859e4d436b [Bugfix][Parser] Fix U+FFFD leak at reasoning-to-content transition in engine parsers (#46159)
Signed-off-by: Ben Browning <bbrownin@redhat.com>
2026-06-19 22:09:28 +00:00
Micah WilliamsonandGitHub 4a083cc858 [ROCm][CI] Pin test_rocm_compressed_tensors_w8a8 to TRITON_ATTN (#46180)
Signed-off-by: Micah Williamson <micah.williamson@amd.com>
2026-06-19 15:20:06 -05:00
Vadim GimpelsonandGitHub ca7e1f2c43 Move CI failure diagnosis docs into ci-fails-buildkite skill (#45975)
Signed-off-by: Vadim Gimpelson <vadim.gimpelson@gmail.com>
2026-06-19 20:12:40 +00:00
djramicandGitHub dec860fb19 [ROCm] Use vLLM's fp8 quant max in AITER hipBLASLt accuracy test (#46176)
Signed-off-by: Djordje Ramic <djoramic@amd.com>
2026-06-19 13:24:02 -05:00
Harry MellorandGitHub 0a49fb2b13 Fix dead link in docs (#46181)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-06-19 18:16:09 +00:00
Ben BrowningandGitHub 4a8abf37c7 [Test] Migrate test_openai_schema.py to schemathesis 4.x (#46173)
Signed-off-by: Ben Browning <bbrownin@redhat.com>
2026-06-19 18:05:18 +00:00
01192139bf [DSv4] Pack KV caches into contiguous per-block allocations for DeepSeek V4 (#44577)
Signed-off-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
Signed-off-by: Lucas Wilkinson <lwilkins@redhat.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Co-authored-by: Matthew Bonanni <mbonanni@redhat.com>
Co-authored-by: Lucas Wilkinson <LucasWilkinson@users.noreply.github.com>
Co-authored-by: Lucas Wilkinson <lwilkins@redhat.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
2026-06-19 12:55:42 -04:00
Chris LeonardandGitHub b9a7cd464c [12/n] final _C library kernel migration (#45415) 2026-06-19 06:57:26 -07:00
test test 928e13af5f Fix MiniMax M3 streaming reasoning markers
Signed-off-by: test test <2260891073@qq.com>
2026-06-16 01:29:52 +08:00
159 changed files with 4032 additions and 1187 deletions
+1 -14
View File
@@ -647,7 +647,7 @@ steps:
- pytest -v -s v1/cudagraph/test_cudagraph_mode.py
- label: e2e Core (1 GPU) # TBD
timeout_in_minutes: 180
timeout_in_minutes: 35
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
agent_pool: mi250_1
optional: true
@@ -2075,19 +2075,6 @@ steps:
- export VLLM_ALLOW_INSECURE_SERIALIZATION=1
- pytest -v -s v1/spec_decode/test_acceptance_length.py -m slow_test
- label: e2e Core (1 GPU) # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
agent_pool: mi300_1
optional: true
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- vllm/v1/
- tests/v1/e2e/
- vllm/platforms/rocm.py
commands:
- pytest -v -s v1/e2e/general --ignore v1/e2e/general/test_async_scheduling.py
- label: e2e Scheduling (1 GPU) # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
+10
View File
@@ -74,6 +74,16 @@ steps:
- tests/v1/e2e/general/
commands:
- pytest -v -s v1/e2e/general --ignore v1/e2e/general/test_async_scheduling.py
mirror:
amd:
device: mi250_1
timeout_in_minutes: 35
depends_on:
- image-build-amd
source_file_dependencies:
- vllm/v1/
- tests/v1/e2e/general/
- vllm/platforms/rocm.py
- label: V1 e2e (2 GPUs)
key: v1-e2e-2-gpus
+1
View File
@@ -104,6 +104,7 @@ steps:
source_file_dependencies:
- csrc/quantization/
- vllm/model_executor/layers/quantization
- vllm/config/
- tests/kernels/quantization
- tests/kernels/quantization/test_rocm_skinny_gemms.py
- vllm/_aiter_ops.py
+10
View File
@@ -101,6 +101,16 @@ steps:
num_devices: 8
commands:
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-h200.txt
mirror:
amd:
device: mi300_8
timeout_in_minutes: 180
depends_on:
- image-build-amd
commands:
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
- export PYTORCH_ROCM_ARCH=gfx942 # Limit Quark compilation to save time
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-mi3xx.txt
- label: MoE Refactor Integration Test (H100 - TEMPORARY)
key: moe-refactor-integration-test-h100-temporary
@@ -68,7 +68,6 @@ steps:
- cd .. && VLLM_WORKER_MULTIPROC_METHOD=spawn pytest -v -s tests/models/multimodal/generation/test_whisper.py -m core_model # Otherwise, mp_method="spawn" doesn't work
mirror:
amd:
soft_fail: true
device: mi325_1
depends_on:
- image-build-amd
@@ -0,0 +1,35 @@
---
name: ci-fails-buildkite
description: Fetch and diagnose vLLM Buildkite CI failure logs. Use when investigating failing CI jobs on a PR or build, when the user pastes a buildkite.com URL, or asks to fetch/diagnose CI logs.
---
# Diagnosing vLLM Buildkite CI Failures
Buildkite logs are public; no login needed.
`.buildkite/scripts/ci-fetch-log.sh` saves each log as `ci-<build>-<job-name>.log`, stripped of timestamps and ANSI codes. Existing files are kept; set `CI_FETCH_LOG_FORCE=1` to refetch.
## Fetching logs
```bash
# All failed jobs in a PR's latest build (current branch's PR if omitted):
.buildkite/scripts/ci-fetch-log.sh --pr <PR>
# All failed jobs in a build (--soft also includes soft-failed jobs;
# --all fetches every finished job):
.buildkite/scripts/ci-fetch-log.sh "https://buildkite.com/vllm/ci/builds/<N>"
# One job — `gh pr checks` URLs (#<job_uuid>) and web UI URLs (?sid=) both
# work; pass "-" as a second argument to stream to stdout:
.buildkite/scripts/ci-fetch-log.sh "https://buildkite.com/vllm/ci/builds/<N>#<job_uuid>"
```
To clean an already-downloaded log with `.buildkite/scripts/ci-clean-log.sh`:
```bash
./ci-clean-log.sh ci.log
```
## Reference
See [docs/contributing/ci/failures.md](../../../docs/contributing/ci/failures.md) for the full guide: filing CI failure issues, investigating/bisecting, reproducing flaky tests, and daily triage.
+2 -3
View File
@@ -2,15 +2,14 @@
# for more info about CODEOWNERS file
# This lists cover the "core" components of vLLM that require careful review
/vllm/compilation @zou3519 @youkaichao @ProExpertProg @BoyuanFeng @vadiklyutiy
/vllm/compilation @zou3519 @youkaichao @ProExpertProg @BoyuanFeng
/vllm/distributed/kv_transfer @NickLucche @ApostaC @orozery @xuechendi
/vllm/lora @jeejeelee
/vllm/model_executor/layers/attention @LucasWilkinson @MatthewBonanni
/vllm/model_executor/layers/fused_moe @mgoin @pavanimajety @zyongye
/vllm/model_executor/layers/quantization @mgoin @robertgshaw2-redhat @tlrmchlsmth @yewentao256 @pavanimajety @zyongye
/vllm/model_executor/layers/mamba @tdoublep @tomeras91
/vllm/model_executor/layers/mamba/gdn_linear_attn.py @tdoublep @ZJY0516 @vadiklyutiy
/vllm/model_executor/layers/rotary_embedding.py @vadiklyutiy
/vllm/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py @tdoublep @ZJY0516 @vadiklyutiy
/vllm/model_executor/model_loader @22quinn
/vllm/model_executor/layers/batch_invariant.py @yewentao256
/vllm/ir @ProExpertProg
+3 -1
View File
@@ -199,7 +199,9 @@ cython_debug/
.vscode/
# Claude
.claude/
.claude/*
!.claude/skills/
!.claude/skills/**
# Codex
.codex/
-11
View File
@@ -114,17 +114,6 @@ Follow these rules for all code changes in this repository:
- Keep comments and docstrings minimal and concise.
- Assume the reader is familiar with vLLM.
### Diagnosing CI failures
Buildkite logs are public; no login needed. Details: [docs/contributing/ci/failures.md](docs/contributing/ci/failures.md).
```bash
# All failed-job logs for a PR's latest build (current branch's PR if omitted):
.buildkite/scripts/ci-fetch-log.sh --pr <PR>
# Any Buildkite build or job URL also works:
.buildkite/scripts/ci-fetch-log.sh "<buildkite_url>"
```
### Commit messages
Add attribution using commit trailers such as `Co-authored-by:` (other projects use `Assisted-by:` or `Generated-by:`). For example:
+56 -72
View File
@@ -319,82 +319,35 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
endif()
#
# _C extension
# Legacy _C extension (ROCm only — CUDA ops migrated to _C_stable_libtorch)
#
set(VLLM_EXT_SRC
"csrc/quantization/activation_kernels.cu"
"csrc/torch_bindings.cpp")
if(VLLM_GPU_LANG STREQUAL "CUDA")
SET(CUTLASS_ENABLE_HEADERS_ONLY ON CACHE BOOL "Enable only the header library")
# Set CUTLASS_REVISION. Used for FetchContent. Also fixes some bogus messages when building.
set(CUTLASS_REVISION "v4.4.2")
# Use the specified CUTLASS source directory for compilation if VLLM_CUTLASS_SRC_DIR is provided
if (DEFINED ENV{VLLM_CUTLASS_SRC_DIR})
set(VLLM_CUTLASS_SRC_DIR $ENV{VLLM_CUTLASS_SRC_DIR})
endif()
if(VLLM_CUTLASS_SRC_DIR)
if(NOT IS_ABSOLUTE VLLM_CUTLASS_SRC_DIR)
get_filename_component(VLLM_CUTLASS_SRC_DIR "${VLLM_CUTLASS_SRC_DIR}" ABSOLUTE)
endif()
message(STATUS "The VLLM_CUTLASS_SRC_DIR is set, using ${VLLM_CUTLASS_SRC_DIR} for compilation")
FetchContent_Declare(cutlass SOURCE_DIR ${VLLM_CUTLASS_SRC_DIR})
else()
FetchContent_Declare(
cutlass
GIT_REPOSITORY https://github.com/nvidia/cutlass.git
# Please keep this in sync with CUTLASS_REVISION line above.
GIT_TAG ${CUTLASS_REVISION}
GIT_PROGRESS TRUE
# Speed up CUTLASS download by retrieving only the specified GIT_TAG instead of the history.
# Important: If GIT_SHALLOW is enabled then GIT_TAG works only with branch names and tags.
# So if the GIT_TAG above is updated to a commit hash, GIT_SHALLOW must be set to FALSE
GIT_SHALLOW TRUE
)
endif()
FetchContent_MakeAvailable(cutlass)
set_gencode_flags_for_srcs(
SRCS "${VLLM_EXT_SRC}"
CUDA_ARCHS "${CUDA_ARCHS}")
# if CUDA endif
endif()
if (VLLM_GPU_LANG STREQUAL "HIP")
# Add QuickReduce kernels (ROCm-only; not part of stable ABI migration).
# TODO: Remove the cuda_view when ROCm upgrade to torch 2.11.
list(APPEND VLLM_EXT_SRC
if(VLLM_GPU_LANG STREQUAL "HIP")
set(VLLM_EXT_SRC
"csrc/torch_bindings.cpp"
"csrc/custom_quickreduce.cu"
"csrc/cuda_view.cu"
"csrc/libtorch_stable/cuda_utils_kernels.cu"
)
# if ROCM endif
endif()
"csrc/libtorch_stable/cuda_utils_kernels.cu")
message(STATUS "Enabling C extension.")
define_extension_target(
_C
DESTINATION vllm
LANGUAGE ${VLLM_GPU_LANG}
SOURCES ${VLLM_EXT_SRC}
COMPILE_FLAGS ${VLLM_GPU_FLAGS}
ARCHITECTURES ${VLLM_GPU_ARCHES}
INCLUDE_DIRECTORIES ${CUTLASS_INCLUDE_DIR}
INCLUDE_DIRECTORIES ${CUTLASS_TOOLS_UTIL_INCLUDE_DIR}
USE_SABI 3
WITH_SOABI)
message(STATUS "Enabling C extension.")
define_extension_target(
_C
DESTINATION vllm
LANGUAGE ${VLLM_GPU_LANG}
SOURCES ${VLLM_EXT_SRC}
COMPILE_FLAGS ${VLLM_GPU_FLAGS}
ARCHITECTURES ${VLLM_GPU_ARCHES}
INCLUDE_DIRECTORIES ${CUTLASS_INCLUDE_DIR}
INCLUDE_DIRECTORIES ${CUTLASS_TOOLS_UTIL_INCLUDE_DIR}
USE_SABI 3
WITH_SOABI)
# If CUTLASS is compiled on NVCC >= 12.5, it by default uses
# cudaGetDriverEntryPointByVersion as a wrapper to avoid directly calling the
# driver API. This causes problems when linking with earlier versions of CUDA.
# Setting this variable sidesteps the issue by calling the driver directly.
target_compile_definitions(_C PRIVATE CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1)
# If CUTLASS is compiled on NVCC >= 12.5, it by default uses
# cudaGetDriverEntryPointByVersion as a wrapper to avoid directly calling the
# driver API. This causes problems when linking with earlier versions of CUDA.
# Setting this variable sidesteps the issue by calling the driver directly.
target_compile_definitions(_C PRIVATE CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1)
endif() # _C HIP endif
if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
#
@@ -403,6 +356,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
set(VLLM_STABLE_EXT_SRC
"csrc/libtorch_stable/torch_bindings.cpp"
"csrc/libtorch_stable/activation_kernels.cu"
"csrc/libtorch_stable/quantization/activation_kernels.cu"
"csrc/libtorch_stable/quantization/w8a8/int8/scaled_quant.cu"
"csrc/libtorch_stable/quantization/w8a8/fp8/common.cu"
"csrc/libtorch_stable/quantization/w8a8/fp8/per_token_group_quant.cu"
@@ -429,6 +383,38 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
"csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu")
if(VLLM_GPU_LANG STREQUAL "CUDA")
SET(CUTLASS_ENABLE_HEADERS_ONLY ON CACHE BOOL "Enable only the header library")
# Set CUTLASS_REVISION. Used for FetchContent. Also fixes some bogus messages when building.
set(CUTLASS_REVISION "v4.4.2")
# Use the specified CUTLASS source directory for compilation if VLLM_CUTLASS_SRC_DIR is provided
if (DEFINED ENV{VLLM_CUTLASS_SRC_DIR})
set(VLLM_CUTLASS_SRC_DIR $ENV{VLLM_CUTLASS_SRC_DIR})
endif()
if(VLLM_CUTLASS_SRC_DIR)
if(NOT IS_ABSOLUTE VLLM_CUTLASS_SRC_DIR)
get_filename_component(VLLM_CUTLASS_SRC_DIR "${VLLM_CUTLASS_SRC_DIR}" ABSOLUTE)
endif()
message(STATUS "The VLLM_CUTLASS_SRC_DIR is set, using ${VLLM_CUTLASS_SRC_DIR} for compilation")
FetchContent_Declare(cutlass SOURCE_DIR ${VLLM_CUTLASS_SRC_DIR})
else()
FetchContent_Declare(
cutlass
GIT_REPOSITORY https://github.com/nvidia/cutlass.git
# Please keep this in sync with CUTLASS_REVISION line above.
GIT_TAG ${CUTLASS_REVISION}
GIT_PROGRESS TRUE
# Speed up CUTLASS download by retrieving only the specified GIT_TAG instead of the history.
# Important: If GIT_SHALLOW is enabled then GIT_TAG works only with branch names and tags.
# So if the GIT_TAG above is updated to a commit hash, GIT_SHALLOW must be set to FALSE
GIT_SHALLOW TRUE
)
endif()
FetchContent_MakeAvailable(cutlass)
list(APPEND VLLM_STABLE_EXT_SRC
"csrc/libtorch_stable/cuda_view.cu"
"csrc/libtorch_stable/cuda_utils_kernels.cu"
@@ -929,7 +915,6 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
SRCS "${FP4_SM120_SRCS}"
CUDA_ARCHS "${FP4_SM120_ARCHS}")
list(APPEND VLLM_STABLE_EXT_SRC "${FP4_SM120_SRCS}")
target_compile_definitions(_C PRIVATE ENABLE_NVFP4_SM120=1)
list(APPEND VLLM_GPU_FLAGS "-DENABLE_NVFP4_SM120=1")
list(APPEND VLLM_GPU_FLAGS "-DENABLE_CUTLASS_MOE_SM120=1")
message(STATUS "Building SM12x NVFP4 for archs: ${FP4_SM120_ARCHS}")
@@ -962,7 +947,6 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
SRCS "${FP4_SM100_SRCS}"
CUDA_ARCHS "${FP4_SM100_ARCHS}")
list(APPEND VLLM_STABLE_EXT_SRC "${FP4_SM100_SRCS}")
target_compile_definitions(_C PRIVATE ENABLE_NVFP4_SM100=1)
list(APPEND VLLM_GPU_FLAGS "-DENABLE_NVFP4_SM100=1")
list(APPEND VLLM_GPU_FLAGS "-DENABLE_CUTLASS_MOE_SM100=1")
message(STATUS "Building SM10x/11x NVFP4/MXFP4 for archs: ${FP4_SM100_ARCHS}")
+29 -5
View File
@@ -60,6 +60,7 @@ endif()
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND QUTLASS_ARCHS)
set(QUTLASS_SOURCES
csrc/qutlass_registration.cpp
${qutlass_SOURCE_DIR}/qutlass/csrc/bindings.cpp
${qutlass_SOURCE_DIR}/qutlass/csrc/gemm.cu
${qutlass_SOURCE_DIR}/qutlass/csrc/gemm_ada.cu
@@ -78,8 +79,19 @@ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND QUTLASS_ARCHS)
if(CUTLASS_INCLUDE_DIR AND EXISTS "${CUTLASS_INCLUDE_DIR}/cutlass/cutlass.h")
list(APPEND QUTLASS_INCLUDES "${CUTLASS_INCLUDE_DIR}")
if(CUTLASS_TOOLS_UTIL_INCLUDE_DIR AND
EXISTS "${CUTLASS_TOOLS_UTIL_INCLUDE_DIR}/cutlass/util/packed_stride.hpp")
list(APPEND QUTLASS_INCLUDES "${CUTLASS_TOOLS_UTIL_INCLUDE_DIR}")
else()
get_filename_component(_qutlass_cutlass_root "${CUTLASS_INCLUDE_DIR}" DIRECTORY)
if(EXISTS "${_qutlass_cutlass_root}/tools/util/include/cutlass/util/packed_stride.hpp")
list(APPEND QUTLASS_INCLUDES "${_qutlass_cutlass_root}/tools/util/include")
endif()
endif()
elseif(EXISTS "${qutlass_SOURCE_DIR}/qutlass/third_party/cutlass/include/cutlass/cutlass.h")
list(APPEND QUTLASS_INCLUDES "${qutlass_SOURCE_DIR}/qutlass/third_party/cutlass/include")
list(APPEND QUTLASS_INCLUDES
"${qutlass_SOURCE_DIR}/qutlass/third_party/cutlass/include"
"${qutlass_SOURCE_DIR}/qutlass/third_party/cutlass/tools/util/include")
message(STATUS "[QUTLASS] Using QuTLASS vendored CUTLASS headers (no vLLM CUTLASS detected).")
else()
message(FATAL_ERROR "[QUTLASS] CUTLASS headers not found. "
@@ -91,12 +103,23 @@ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND QUTLASS_ARCHS)
CUDA_ARCHS "${QUTLASS_ARCHS}"
)
target_sources(_C PRIVATE ${QUTLASS_SOURCES})
target_include_directories(_C PRIVATE ${QUTLASS_INCLUDES})
target_compile_definitions(_C PRIVATE
# QuTLASS uses legacy ATen headers and cannot be built with TORCH_TARGET_VERSION.
# Keep it as its own extension (registers torch.ops._qutlass_C).
define_extension_target(
_qutlass_C
DESTINATION vllm
LANGUAGE ${VLLM_GPU_LANG}
SOURCES ${QUTLASS_SOURCES}
COMPILE_FLAGS ${VLLM_GPU_FLAGS}
ARCHITECTURES ${VLLM_GPU_ARCHES}
INCLUDE_DIRECTORIES ${QUTLASS_INCLUDES}
USE_SABI 3
WITH_SOABI)
target_compile_definitions(_qutlass_C PRIVATE
QUTLASS_DISABLE_PYBIND=1
TARGET_CUDA_ARCH=${QUTLASS_TARGET_CC}
)
CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1)
set_property(SOURCE ${QUTLASS_SOURCES} APPEND PROPERTY COMPILE_OPTIONS
$<$<COMPILE_LANGUAGE:CUDA>:--expt-relaxed-constexpr --use_fast_math -O3>
@@ -111,4 +134,5 @@ else()
"[QUTLASS] Skipping build: no supported arch (12.0f / 10.0f) found in "
"CUDA_ARCHS='${CUDA_ARCHS}'.")
endif()
add_custom_target(_qutlass_C)
endif()
@@ -268,9 +268,14 @@ int64_t sm100_cutlass_mla_get_workspace_size(int64_t max_seq_len, int64_t num_ba
using TileShapeD = typename MlaSm100Type::TileShapeD;
arguments.problem_shape =
cute::make_tuple(TileShapeH{}, static_cast<int>(max_seq_len), TileShapeD{}, static_cast<int>(num_batches));
// Assumes device 0 when getting sm_count.
arguments.hw_info.sm_count =
sm_count <= 0 ? cutlass::KernelHardwareInfo::query_device_multiprocessor_count(/*device_id=*/0) : sm_count;
if (sm_count <= 0) {
int current_device = 0;
cudaGetDevice(&current_device);
arguments.hw_info.sm_count =
cutlass::KernelHardwareInfo::query_device_multiprocessor_count(current_device);
} else {
arguments.hw_info.sm_count = sm_count;
}
arguments.split_kv = static_cast<int>(num_kv_splits);
MlaSm100Type::Fmha::set_split_kv(arguments);
@@ -9,7 +9,7 @@
#include <torch/headeronly/core/ScalarType.h>
#include "../../cuda_compat.h"
#include "core/math.hpp"
#include "libtorch_stable/core/math.hpp"
#include "libtorch_stable/dispatch_utils.h"
#include "libtorch_stable/torch_utils.h"
+28
View File
@@ -2,9 +2,25 @@
#include <torch/csrc/stable/library.h>
#include <torch/csrc/stable/tensor.h>
#include <torch/headeronly/util/Exception.h>
#include <optional>
#include <string>
#include <vector>
#include <torch/csrc/stable/ops.h>
inline torch::stable::Tensor weak_ref_tensor(torch::stable::Tensor& tensor) {
// Ensure tensor is on CUDA
STD_TORCH_CHECK(tensor.device().is_cuda(), "Tensor must be on CUDA device");
// Get the raw data pointer
void* data_ptr = tensor.mutable_data_ptr();
/// Create a new tensor from the raw data pointer
return torch::stable::from_blob(data_ptr, tensor.sizes(), tensor.strides(),
tensor.device(), tensor.scalar_type());
}
void per_token_group_quant_fp8(const torch::stable::Tensor& input,
torch::stable::Tensor& output_q,
@@ -371,6 +387,18 @@ void silu_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input);
void silu_and_mul_clamp(torch::stable::Tensor& out,
torch::stable::Tensor& input, double limit,
double alpha = 1.0, double beta = 0.0);
void silu_and_mul_quant(torch::stable::Tensor& out,
torch::stable::Tensor& input,
torch::stable::Tensor& scale);
void persistent_masked_m_silu_mul_quant(
const torch::stable::Tensor& input, // (E, T, 2*H)
const torch::stable::Tensor& tokens_per_expert, // (E)
torch::stable::Tensor& y_q, // (E, T, H) [OUT]
torch::stable::Tensor& y_s, // (E, T, H//group_size) [OUT]
bool use_ue8m0);
void mul_and_silu(torch::stable::Tensor& out, torch::stable::Tensor& input);
void gelu_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input);
void gelu_tanh_and_mul(torch::stable::Tensor& out,
@@ -1,16 +1,12 @@
#include <ATen/cuda/CUDAContext.h>
#include <torch/all.h>
#include <c10/cuda/CUDAGuard.h>
#include "libtorch_stable/torch_utils.h"
#include <cmath>
#include "core/math.hpp"
#include "../cuda_compat.h"
#include "dispatch_utils.h"
#include "libtorch_stable/core/math.hpp"
#include "cuda_compat.h"
#include "libtorch_stable/dispatch_utils.h"
#include "quantization/w8a8/fp8/common.cuh"
#include <c10/util/Float8_e4m3fn.h>
#ifndef USE_ROCM
#include <cuda_bf16.h>
#include <cuda_fp16.h>
@@ -33,7 +29,6 @@ typedef __hip_fp8x4_e4m3_fnuz __nv_fp8x4_e4m3;
#endif
#endif
#include "core/registration.h"
namespace vllm {
template <typename T>
@@ -564,41 +559,47 @@ __global__ void silu_mul_fp8_quant_deep_gemm_kernel(
} // namespace vllm
// Launch activation, gating, and quantize kernel.
#define LAUNCH_ACTIVATION_GATE_KERNEL(KERNEL) \
int d = input.size(-1) / 2; \
int64_t num_tokens = input.numel() / input.size(-1); \
dim3 grid(num_tokens, num_tokens > 16 ? num_tokens > 32 ? 1 : 2 : 4); \
dim3 block(std::min(d, 512)); \
const at::cuda::OptionalCUDAGuard device_guard(device_of(input)); \
const cudaStream_t stream = at::cuda::getCurrentCUDAStream(); \
VLLM_DISPATCH_FLOATING_TYPES( \
input.scalar_type(), "act_and_mul_kernel", [&] { \
VLLM_DISPATCH_FP8_TYPES( \
out.scalar_type(), "fused_add_rms_norm_kernel_fp8_type", [&] { \
vllm::act_and_mul_quant_kernel<scalar_t, KERNEL<scalar_t>, \
fp8_t> \
<<<grid, block, 0, stream>>>(out.data_ptr<fp8_t>(), \
input.data_ptr<scalar_t>(), \
scale.data_ptr<float>(), d); \
}); \
#define LAUNCH_ACTIVATION_GATE_KERNEL(KERNEL) \
int d = input.size(-1) / 2; \
int64_t num_tokens = input.numel() / input.size(-1); \
dim3 grid(num_tokens, num_tokens > 16 ? num_tokens > 32 ? 1 : 2 : 4); \
dim3 block(std::min(d, 512)); \
const torch::stable::accelerator::DeviceGuard device_guard( \
input.get_device_index()); \
const cudaStream_t stream = \
get_current_cuda_stream(input.get_device_index()); \
VLLM_STABLE_DISPATCH_FLOATING_TYPES( \
input.scalar_type(), "act_and_mul_kernel", [&] { \
VLLM_STABLE_DISPATCH_FP8_TYPES( \
out.scalar_type(), "act_and_mul_quant_kernel_fp8_type", [&] { \
vllm::act_and_mul_quant_kernel<scalar_t, KERNEL<scalar_t>, \
fp8_t> \
<<<grid, block, 0, stream>>>( \
out.mutable_data_ptr<fp8_t>(), \
input.const_data_ptr<scalar_t>(), \
scale.const_data_ptr<float>(), d); \
}); \
});
void silu_and_mul_quant(torch::Tensor& out, // [..., d]
torch::Tensor& input, // [..., 2 * d]
torch::Tensor& scale) {
TORCH_CHECK(out.dtype() == torch::kFloat8_e4m3fn ||
out.dtype() == torch::kFloat8_e4m3fnuz);
TORCH_CHECK(input.dtype() == torch::kFloat16 ||
input.dtype() == torch::kBFloat16);
TORCH_CHECK(input.size(-1) % 2 == 0);
void silu_and_mul_quant(torch::stable::Tensor& out, // [..., d]
torch::stable::Tensor& input, // [..., 2 * d]
torch::stable::Tensor& scale) {
STD_TORCH_CHECK(
out.scalar_type() == torch::headeronly::ScalarType::Float8_e4m3fn ||
out.scalar_type() == torch::headeronly::ScalarType::Float8_e4m3fnuz);
STD_TORCH_CHECK(
input.scalar_type() == torch::headeronly::ScalarType::Half ||
input.scalar_type() == torch::headeronly::ScalarType::BFloat16,
"Input must be FP16 or BF16");
STD_TORCH_CHECK(input.size(-1) % 2 == 0);
LAUNCH_ACTIVATION_GATE_KERNEL(vllm::silu_kernel);
}
void persistent_masked_m_silu_mul_quant(
const at::Tensor& input, // (E, T, 2*H)
const at::Tensor& tokens_per_expert, // (E)
at::Tensor& y_q, // (E, T, H) [OUT]
at::Tensor& y_s, // (E, T, H//group_size) [OUT]
const torch::stable::Tensor& input, // (E, T, 2*H)
const torch::stable::Tensor& tokens_per_expert, // (E)
torch::stable::Tensor& y_q, // (E, T, H) [OUT]
torch::stable::Tensor& y_s, // (E, T, H//group_size) [OUT]
bool cast_scale_ue8m0) {
#ifndef USE_ROCM
@@ -606,14 +607,18 @@ void persistent_masked_m_silu_mul_quant(
// fixed GROUP_SIZE of 128.
static constexpr int GROUP_SIZE = 128;
TORCH_CHECK(input.dtype() == torch::kBFloat16);
TORCH_CHECK(y_q.dtype() == torch::kFloat8_e4m3fn ||
y_q.dtype() == torch::kFloat8_e4m3fnuz);
TORCH_CHECK(input.size(-1) % (GROUP_SIZE * 2) == 0);
STD_TORCH_CHECK(input.scalar_type() ==
torch::headeronly::ScalarType::BFloat16);
STD_TORCH_CHECK(
y_q.scalar_type() == torch::headeronly::ScalarType::Float8_e4m3fn ||
y_q.scalar_type() == torch::headeronly::ScalarType::Float8_e4m3fnuz);
STD_TORCH_CHECK(input.size(-1) % (GROUP_SIZE * 2) == 0);
bool const is_packed_ue8m0 =
(y_s.dtype() == torch::kInt32 && cast_scale_ue8m0);
TORCH_CHECK(y_s.dtype() == torch::kFloat32 || is_packed_ue8m0);
(y_s.scalar_type() == torch::headeronly::ScalarType::Int &&
cast_scale_ue8m0);
STD_TORCH_CHECK(y_s.scalar_type() == torch::headeronly::ScalarType::Float ||
is_packed_ue8m0);
using Idx_t = int64_t;
@@ -631,7 +636,7 @@ void persistent_masked_m_silu_mul_quant(
int const NUM_GROUPS = H / GROUP_SIZE;
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
const cudaStream_t stream = get_current_cuda_stream(input.get_device_index());
// TODO: Get this from cuda_arch ?
static constexpr int SILU_V2_BLOCK_COUNT = 132 * 32;
@@ -643,18 +648,21 @@ void persistent_masked_m_silu_mul_quant(
static constexpr int max_shared_mem_bytes = \
GROUP_SIZE * 2 * STAGES * NUM_WARPS * 2; \
dim3 grid(sms), block(THREAD_COUNT); \
const at::cuda::OptionalCUDAGuard device_guard(device_of(input)); \
VLLM_DISPATCH_FP8_TYPES( \
const torch::stable::accelerator::DeviceGuard device_guard( \
input.get_device_index()); \
VLLM_STABLE_DISPATCH_FP8_TYPES( \
y_q.scalar_type(), "silu_mul_fp8_quant_deep_gemm_kernel", [&] { \
vllm::silu_mul_fp8_quant_deep_gemm_kernel< \
BLOCK_COUNT, max_shared_mem_bytes, fp8_t, scale_t, THREAD_COUNT, \
Idx_t, CEIL_UE8M0, GROUP_SIZE, STAGES> \
<<<grid, block, max_shared_mem_bytes + (E + 1) * 16, stream>>>( \
reinterpret_cast<__nv_bfloat16*>(input.data_ptr()), \
(fp8_t*)y_q.data_ptr(), \
reinterpret_cast<scale_t*>(y_s.data_ptr()), \
reinterpret_cast<int32_t*>(tokens_per_expert.data_ptr()), E, \
T, H, stride_i_e, stride_i_t, stride_i_h, stride_yq_e, \
reinterpret_cast<const __nv_bfloat16*>( \
input.const_data_ptr()), \
y_q.mutable_data_ptr<fp8_t>(), \
reinterpret_cast<scale_t*>(y_s.mutable_data_ptr()), \
reinterpret_cast<const int32_t*>( \
tokens_per_expert.const_data_ptr()), \
E, T, H, stride_i_e, stride_i_t, stride_i_h, stride_yq_e, \
stride_yq_t, stride_yq_h, STRIDE_YS_E, STRIDE_YS_T, \
STRIDE_YS_G, STRIDE_YS_P, stride_counts_e); \
});
@@ -679,7 +687,7 @@ void persistent_masked_m_silu_mul_quant(
Idx_t stride_ys_g = y_s.stride(2);
Idx_t stride_ys_p = 0;
if (!cast_scale_ue8m0) {
TORCH_CHECK(!is_packed_ue8m0);
STD_TORCH_CHECK(!is_packed_ue8m0);
LAUNCH_ON_H(float, stride_ys_e, stride_ys_t, stride_ys_g, stride_ys_p,
false);
return;
@@ -692,8 +700,8 @@ void persistent_masked_m_silu_mul_quant(
return;
}
TORCH_CHECK(cast_scale_ue8m0 && is_packed_ue8m0);
TORCH_CHECK(y_s.dtype() == torch::kInt32);
STD_TORCH_CHECK(cast_scale_ue8m0 && is_packed_ue8m0);
STD_TORCH_CHECK(y_s.scalar_type() == torch::headeronly::ScalarType::Int);
// Int32 packed ue8m0 scales tensor.
// Let E, T, G be the number to experts, number of tokens and number of groups
@@ -31,7 +31,7 @@
#include "cutlass/util/packed_stride.hpp"
#include "core/math.hpp"
#include "libtorch_stable/core/math.hpp"
#include "core/batch_invariant.hpp"
using namespace cute;
@@ -31,7 +31,7 @@
#include "cutlass/util/packed_stride.hpp"
#include "core/math.hpp"
#include "libtorch_stable/core/math.hpp"
#include "core/batch_invariant.hpp"
using namespace cute;
@@ -19,7 +19,7 @@
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/util/packed_stride.hpp"
#include "core/math.hpp"
#include "libtorch_stable/core/math.hpp"
#include "libtorch_stable/cutlass_extensions/common.hpp"
// clang-format on
@@ -14,7 +14,7 @@
#include "cutlass/epilogue/collective/collective_builder.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "core/math.hpp"
#include "libtorch_stable/core/math.hpp"
#include "libtorch_stable/cutlass_extensions/common.hpp"
// clang-format on
@@ -22,7 +22,7 @@
#include "cutlass/epilogue/threadblock/fusion/visitors.hpp"
#include "cutlass/gemm/kernel/default_gemm_universal_with_visitor.h"
#include "core/math.hpp"
#include "libtorch_stable/core/math.hpp"
#include "libtorch_stable/cutlass_extensions/common.hpp"
// clang-format on
@@ -301,8 +301,9 @@ __global__ void per_token_group_quant_8bit_packed_register_kernel(
const int sf_k_local = local_group_id % kGroupsPerBlockX;
const int row_local = local_group_id / kGroupsPerBlockX;
const int sf_k_idx = blockIdx.x * kGroupsPerBlockX + sf_k_local;
const int mn_idx = blockIdx.y * kRowsPerBlock + row_local;
// Rows on grid.x: mn scales with tokens and can exceed the 65535 grid.y cap.
const int sf_k_idx = blockIdx.y * kGroupsPerBlockX + sf_k_local;
const int mn_idx = blockIdx.x * kRowsPerBlock + row_local;
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
asm volatile("griddepcontrol.wait;");
@@ -496,14 +497,15 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
" is not a multiple of 4.");
const int kx = GetGroupsPerBlockX(padded_groups_per_row);
const int ry = 16 / kx;
const int64_t blocks_x = padded_groups_per_row / kx;
const int64_t blocks_y = (tma_aligned_mn + ry - 1) / ry;
const int64_t row_blocks = (tma_aligned_mn + ry - 1) / ry;
const int64_t sf_k_blocks = padded_groups_per_row / kx;
const int num_threads = (kx * ry) * THREADS_PER_GROUP;
// CUDA caps grid.x and grid.y at 2^31 - 1; guard against pathological inputs.
STD_TORCH_CHECK(blocks_x <= static_cast<int64_t>(INT32_MAX) &&
blocks_y <= static_cast<int64_t>(INT32_MAX),
// CUDA caps grid.x at 2^31 - 1 and grid.y at 2^16 - 1 (65535).
constexpr int64_t kMaxGridDimYZ = 65535;
STD_TORCH_CHECK(row_blocks <= static_cast<int64_t>(INT32_MAX) &&
sf_k_blocks <= kMaxGridDimYZ,
"per_token_group_quant_8bit_packed grid too large: (",
blocks_x, ", ", blocks_y, ").");
row_blocks, ", ", sf_k_blocks, ").");
auto dst_type = output_q.scalar_type();
@@ -513,8 +515,8 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
#define LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, KX, RY) \
do { \
cudaLaunchConfig_t config = {}; \
config.gridDim = dim3(static_cast<unsigned int>(blocks_x), \
static_cast<unsigned int>(blocks_y)); \
config.gridDim = dim3(static_cast<unsigned int>(row_blocks), \
static_cast<unsigned int>(sf_k_blocks)); \
config.blockDim = dim3(num_threads); \
config.dynamicSmemBytes = 0; \
config.stream = stream; \
@@ -539,8 +541,8 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
#else
#define LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, KX, RY) \
do { \
dim3 grid(static_cast<unsigned int>(blocks_x), \
static_cast<unsigned int>(blocks_y)); \
dim3 grid(static_cast<unsigned int>(row_blocks), \
static_cast<unsigned int>(sf_k_blocks)); \
dim3 block(num_threads); \
per_token_group_quant_8bit_packed_register_kernel<T, DST_DTYPE, 128, KX, \
RY> \
+27
View File
@@ -34,6 +34,20 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
// TODO: Remove this once ROCm upgrade to torch 2.11.
ops.def("get_cuda_view_from_cpu_tensor(Tensor cpu_tensor) -> Tensor");
// Note about marlin kernel 'workspace' arguments:
// Technically these should be mutable since they are modified by the kernel.
// But since they are set back to zero once the kernel is finished we can
// hand wave and say that they have no net effect.
//
// The reason to mark 'workspace' as immutable is so that they don't interfere
// with using ScalarType arguments in the ops. If they are marked as mutable,
// pytorch throws an assert in
// 'torch._higher_order_ops._register_effectful_op' that prevents these
// kernels from being torch.compile'd.
// See the following document for more info on custom types and ops that use
// custom types:
// https://docs.google.com/document/d/18fBMPuOJ0fY5ZQ6YyrHUppw9FA332CpNtgB6SOIgyuA
// Machete (Dense) Optimized Mixed Precision GEMM for Hopper.
ops.def(
"machete_supported_schedules("
@@ -480,6 +494,11 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
"Tensor workspace, int k, int max_seq_len) -> ()");
// Activation ops
ops.def(
"persistent_masked_m_silu_mul_quant(Tensor input, Tensor counts, Tensor! "
"y_q, Tensor! y_s, bool use_ue8m0) -> ()");
ops.def("weak_ref_tensor(Tensor input) -> Tensor");
// Activation function used in SwiGLU.
ops.def("silu_and_mul(Tensor! result, Tensor input) -> ()");
@@ -492,6 +511,10 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
"silu_and_mul_with_clamp(Tensor! result, Tensor input, float limit, "
"float alpha=1.0, float beta=0.0) -> ()");
// SwiGLU activation with FP8 quantization.
ops.def(
"silu_and_mul_quant(Tensor! result, Tensor input, Tensor scale) -> ()");
// Activation function used in GeGLU with `none` approximation.
ops.def("gelu_and_mul(Tensor! out, Tensor input) -> ()");
@@ -690,6 +713,10 @@ STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, ops) {
ops.impl("persistent_topk", TORCH_BOX(&persistent_topk));
// Activation kernels (shared CUDA/ROCm)
ops.impl("persistent_masked_m_silu_mul_quant",
TORCH_BOX(&persistent_masked_m_silu_mul_quant));
ops.impl("weak_ref_tensor", TORCH_BOX(&weak_ref_tensor));
ops.impl("silu_and_mul_quant", TORCH_BOX(&silu_and_mul_quant));
ops.impl("silu_and_mul", TORCH_BOX(&silu_and_mul));
ops.impl("mul_and_silu", TORCH_BOX(&mul_and_silu));
ops.impl("gelu_and_mul", TORCH_BOX(&gelu_and_mul));
-32
View File
@@ -9,28 +9,6 @@
#include <vector>
torch::Tensor weak_ref_tensor(torch::Tensor& tensor) {
// Ensure tensor is on CUDA
if (!tensor.is_cuda()) {
throw std::runtime_error("Tensor must be on CUDA device");
}
// Get the raw data pointer
void* data_ptr = tensor.data_ptr();
// Get tensor sizes and strides
std::vector<int64_t> sizes = tensor.sizes().vec();
std::vector<int64_t> strides = tensor.strides().vec();
// Get tensor options (dtype, device)
auto options = tensor.options();
// Create a new tensor from the raw data pointer
auto new_tensor = torch::from_blob(data_ptr, sizes, strides, options);
return new_tensor;
}
// rms_norm and fused_add_rms_norm declarations also exist in
// csrc/libtorch_stable/ops.h (torch::stable ABI for CUDA). They remain here
// because the CPU build still uses these torch::Tensor declarations.
@@ -53,16 +31,6 @@ void silu_and_mul(torch::Tensor& out, torch::Tensor& input);
void silu_and_mul_clamp(torch::Tensor& out, torch::Tensor& input, double limit,
double alpha = 1.0, double beta = 0.0);
void silu_and_mul_quant(torch::Tensor& out, torch::Tensor& input,
torch::Tensor& scale);
void persistent_masked_m_silu_mul_quant(
const at::Tensor& input, // (E, T, 2*H)
const at::Tensor& counts, // (E)
at::Tensor& y_q, // (E, T, H) [OUT]
at::Tensor& y_s, // (E, T, H//group_size) [OUT]
bool use_ue8m0);
void gelu_and_mul(torch::Tensor& out, torch::Tensor& input);
void gelu_tanh_and_mul(torch::Tensor& out, torch::Tensor& input);
+5
View File
@@ -0,0 +1,5 @@
#include "core/registration.h"
// QuTLASS registers torch.ops._qutlass_C via TORCH_LIBRARY in bindings.cpp.
// This stub lets Python import vllm._qutlass_C to trigger op registration.
REGISTER_EXTENSION(_qutlass_C)
-40
View File
@@ -20,17 +20,6 @@
TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
// vLLM custom ops
//
ops.def(
"persistent_masked_m_silu_mul_quant(Tensor input, Tensor counts, Tensor! "
"y_q, Tensor! y_s,"
"bool use_ue8m0) -> ()");
ops.impl("persistent_masked_m_silu_mul_quant", torch::kCUDA,
&persistent_masked_m_silu_mul_quant);
ops.def("weak_ref_tensor(Tensor input) -> Tensor");
ops.impl("weak_ref_tensor", torch::kCUDA, &weak_ref_tensor);
#ifdef USE_ROCM
// TODO: Remove this once we upgrade to torch 2.11.
@@ -39,35 +28,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
ops.def("get_cuda_view_from_cpu_tensor(Tensor cpu_tensor) -> Tensor");
ops.impl("get_cuda_view_from_cpu_tensor", torch::kCPU,
&get_cuda_view_from_cpu_tensor);
#endif
// Activation ops (quantized only — basic ops moved to _C_stable_libtorch)
ops.def(
"silu_and_mul_quant(Tensor! result, Tensor input, Tensor scale) -> ()");
ops.impl("silu_and_mul_quant", torch::kCUDA, &silu_and_mul_quant);
// Horizontally-fused DeepseekV4-MLA: per-head RMSNorm + GPT-J RoPE for Q, and
// GPT-J RoPE + UE8M0 FP8 quant + paged cache insert for KV, all in one
// kernel launch. Registered in _C_stable_libtorch (incl. the FlashInfer V4
// full-cache bf16/fp8 variants).
// Quantization ops
#ifndef USE_ROCM
// Note about marlin kernel 'workspace' arguments:
// Technically these should be mutable since they are modified by the kernel.
// But since they are set back to zero once the kernel is finished we can
// hand wave and say that they have no net effect.
//
// The reason to mark 'workspace' as immutable is so that they don't interfere
// with using ScalarType arguments in the ops. If they are marked as mutable,
// pytorch throws an assert in
// 'torch._higher_order_ops._register_effectful_op' that prevents these
// kernels from being torch.compile'd.
// See the following document for more info on custom types and ops that use
// custom types:
// https://docs.google.com/document/d/18fBMPuOJ0fY5ZQ6YyrHUppw9FA332CpNtgB6SOIgyuA
#endif
}
+1 -1
View File
@@ -133,7 +133,7 @@ The model should inherit protocol `IsAttentionFree` and also implement class met
For the mamba layers themselves, please use the [`MambaMixer`](../../../vllm/model_executor/layers/mamba/mamba_mixer.py) (for Mamba-1) or [`MambaMixer2`](../../../vllm/model_executor/layers/mamba/mamba_mixer2.py) (for Mamba-2) classes.
The model should also be added to the `MODELS_CONFIG_MAP` dictionary in [vllm/model_executor/models/config.py](../../../vllm/model_executor/models/config.py) to ensure that the runtime defaults are optimized.
For case (2), we recommend using as a reference the implementation of [`JambaForCausalLM`](../../../vllm/model_executor/models/jamba.py) (for an example of a model that uses Mamba-1 and attention together) or [`BambaForCausalLM`](../../../vllm/model_executor/models/bamba.py) (for an example of a model that uses Mamba-2 and attention together).
For case (2), we recommend using as a reference the implementation of [`JambaForCausalLM`](../../../vllm/model_executor/models/jamba.py) (for an example of a model that uses Mamba-1 and attention together) or [`NemotronHForCausalLM`](../../../vllm/model_executor/models/nemotron_h.py) (for an example of a model that uses Mamba-2 and attention together).
These models should follow the same instructions as case (1), but they should inherit protocol `IsHybrid` (instead of `IsAttentionFree`) and it is *not* necessary to add them to the `MODELS_CONFIG_MAP` (their runtime defaults will be inferred from the protocol).
For case (3), we recommend looking at the implementation of [`MiniMaxText01ForCausalLM`](../../../vllm/model_executor/models/minimax_text_01.py) or [`Lfm2ForCausalLM`](../../../vllm/model_executor/models/lfm2.py) as a reference, which use custom "mamba-like" layers `MiniMaxText01LinearAttention` and `ShortConv` respectively.
+1 -1
View File
@@ -40,7 +40,7 @@ lm-eval[api]>=0.4.12 # required for model evaluation test
mteb[bm25s]>=2, <3 # required for mteb test
transformers==5.5.3
tokenizers==0.22.2
schemathesis>=3.39.15 # Required for openai schema test.
schemathesis>=4.0.0 # Required for openai schema test.
# quantization
bitsandbytes==0.49.2
buildkite-test-collector==0.1.9
+1 -1
View File
@@ -31,7 +31,7 @@ lm-eval[api]>=0.4.12 # required for model evaluation test
mteb[bm25s]>=2, <3 # required for mteb test
transformers==5.5.3
tokenizers==0.22.2
schemathesis>=3.39.15 # Required for openai schema test.
schemathesis>=4.0.0 # Required for openai schema test.
# quantization
bitsandbytes>=0.49.2
buildkite-test-collector==0.1.9
+1 -1
View File
@@ -39,7 +39,7 @@ lm-eval[api]>=0.4.12 # required for model evaluation test
mteb[bm25s]>=2, <3 # required for mteb test
transformers==5.5.3
tokenizers==0.22.2
schemathesis>=3.39.15 # Required for openai schema test
schemathesis>=4.0.0 # Required for openai schema test
# quantization
bitsandbytes==0.49.2
buildkite-test-collector==0.1.9
+4 -1
View File
@@ -769,6 +769,7 @@ class precompiled_wheel_utils:
"vllm/_C.abi3.so",
"vllm/_C_stable_libtorch.abi3.so",
"vllm/_moe_C_stable_libtorch.abi3.so",
"vllm/_qutlass_C.abi3.so",
"vllm/_flashmla_C.abi3.so",
"vllm/_flashmla_extension_C.abi3.so",
"vllm/_sparse_flashmla_C.abi3.so",
@@ -1135,6 +1136,7 @@ if _is_cuda():
# DeepGEMM requires CUDA 12.3+ (SM90/SM100)
# Optional since it won't build on unsupported architectures
ext_modules.append(CMakeExtension(name="vllm._deep_gemm_C", optional=True))
ext_modules.append(CMakeExtension(name="vllm._qutlass_C", optional=True))
# fmha_sm100 is a Python/CuTe-DSL package installed into vllm.third_party.
ext_modules.append(CMakeExtension(name="vllm.fmha_sm100", optional=True))
@@ -1149,7 +1151,8 @@ if _is_cpu():
ext_modules.append(CMakeExtension(name="vllm._C"))
if _build_custom_ops():
ext_modules.append(CMakeExtension(name="vllm._C"))
if _is_hip():
ext_modules.append(CMakeExtension(name="vllm._C"))
if _is_cuda() or _is_hip():
ext_modules.append(CMakeExtension(name="vllm._C_stable_libtorch"))
ext_modules.append(CMakeExtension(name="vllm._moe_C_stable_libtorch"))
+193
View File
@@ -649,3 +649,196 @@ def test_cloud_storage_tokenizer_skips_get_model_path(monkeypatch):
args = EngineArgs(model="s3://bucket/model", tokenizer="s3://bucket/tokenizer")
assert args.model == "s3://bucket/model"
assert args.tokenizer == "s3://bucket/tokenizer"
class TestDeviceIds:
def test_device_ids_with_cvd_out_of_range(self, monkeypatch):
"""--device-ids index beyond the CVD set raises ValueError."""
from vllm.platforms import current_platform
key = current_platform.device_control_env_var
monkeypatch.setenv(key, "4,5")
args = EngineArgs(model="m", device_ids=[0, 2])
with pytest.raises(ValueError, match="out of range"):
args._resolve_device_ids()
def test_device_ids_with_cvd_resolve_to_physical_ids(self, monkeypatch):
"""--device-ids are CVD-local indices resolved to physical ids."""
from vllm.platforms import current_platform
key = current_platform.device_control_env_var
monkeypatch.setenv(key, "4,5")
args = EngineArgs(model="m", device_ids=[0, 1])
assert args._resolve_device_ids() == [4, 5]
def test_device_ids_with_uuid_cvd_resolve_to_physical_ids(self, monkeypatch):
"""--device-ids support UUID CVD values resolved by the platform."""
from vllm.platforms import current_platform
key = current_platform.device_control_env_var
monkeypatch.setenv(key, "GPU-abcd1234,GPU-ef567890")
monkeypatch.setattr(
type(current_platform),
"device_control_id_to_physical_device_id",
classmethod(
lambda cls, device_id: {"GPU-abcd1234": 4, "GPU-ef567890": 5}[device_id]
),
)
args = EngineArgs(model="m", device_ids=[0, 1])
assert args._resolve_device_ids() == [4, 5]
def test_device_ids_with_uuid_args_resolve_to_physical_ids(self, monkeypatch):
"""UUID --device-ids are resolved to physical IDs immediately."""
from vllm.platforms import current_platform
monkeypatch.setattr(
type(current_platform),
"device_control_id_to_physical_device_id",
classmethod(lambda cls, device_id: {"GPU-abcd1234": 4}[device_id]),
)
args = EngineArgs(model="m", device_ids=["GPU-abcd1234"])
assert args._resolve_device_ids() == [4]
def test_device_ids_reject_mixed_integer_and_uuid_args(self):
"""--device-ids must not mix CVD indices and UUIDs."""
args = EngineArgs(model="m", device_ids=[0, "GPU-abcd1234"])
with pytest.raises(ValueError, match="must not mix"):
args._resolve_device_ids()
def test_no_device_ids(self):
"""No --device-ids returns None."""
args = EngineArgs(model="m")
assert args._resolve_device_ids() is None
def test_cli_parsing(self):
"""--device-ids parses comma-separated string from CLI."""
parser = FlexibleArgumentParser()
EngineArgs.add_cli_args(parser)
parsed = parser.parse_args(["--model", "m", "--device-ids", "0,2,4"])
assert parsed.device_ids == [0, 2, 4]
def test_cli_parsing_uuid(self):
"""--device-ids parses comma-separated UUID strings from CLI."""
parser = FlexibleArgumentParser()
EngineArgs.add_cli_args(parser)
parsed = parser.parse_args(
["--model", "m", "--device-ids", "GPU-abcd1234,GPU-ef567890"]
)
assert parsed.device_ids == ["GPU-abcd1234", "GPU-ef567890"]
def test_assigned_physical_gpu_ids_are_physical_with_cvd(self, monkeypatch):
"""assigned_physical_gpu_ids are already physical and not composed with CVD."""
import vllm.platforms.interface as platform_interface
from vllm.platforms import current_platform
monkeypatch.setattr(platform_interface, "_assigned_physical_gpu_ids", [4, 5])
monkeypatch.setenv(current_platform.device_control_env_var, "4,5")
assert current_platform.device_id_to_physical_device_id(0) == 4
assert current_platform.device_id_to_physical_device_id(1) == 5
assert current_platform.logical_device_id_to_visible_device_id(0) == 0
assert current_platform.logical_device_id_to_visible_device_id(1) == 1
def test_assigned_physical_gpu_ids_map_to_visible_uuid_cvd(self, monkeypatch):
"""Physical IDs map back to visible ordinals when CVD uses UUIDs."""
import vllm.platforms.interface as platform_interface
from vllm.platforms import current_platform
monkeypatch.setattr(platform_interface, "_assigned_physical_gpu_ids", [5])
monkeypatch.setenv(
current_platform.device_control_env_var,
"GPU-abcd1234,GPU-ef567890",
)
monkeypatch.setattr(
type(current_platform),
"device_control_id_to_physical_device_id",
classmethod(
lambda cls, device_id: {"GPU-abcd1234": 4, "GPU-ef567890": 5}[device_id]
),
)
assert current_platform.logical_device_id_to_visible_device_id(0) == 1
def test_device_ids_reject_duplicates(self):
"""--device-ids must not contain duplicate entries."""
args = EngineArgs(model="m", device_ids=[2, 2])
with pytest.raises(ValueError, match="duplicates"):
args._resolve_device_ids()
def test_cli_parsing_strips_whitespace(self):
"""--device-ids tolerates whitespace around commas."""
parser = FlexibleArgumentParser()
EngineArgs.add_cli_args(parser)
parsed = parser.parse_args(["--model", "m", "--device-ids", "0, 2, 4"])
assert parsed.device_ids == [0, 2, 4]
def test_visible_ordinal_to_physical_ignores_assigned_ids(self, monkeypatch):
"""visible_device_id_to_physical_device_id maps torch device ordinals,
independent of the logical-to-physical mapping.
Regression test: CustomAllreduce passes device.index (a visible
ordinal) and must not index into assigned_physical_gpu_ids, which
raised IndexError for non-identity --device-ids like [2, 3].
"""
import vllm.platforms.interface as platform_interface
from vllm.platforms import current_platform
monkeypatch.setattr(platform_interface, "_assigned_physical_gpu_ids", [2, 3])
monkeypatch.delenv(current_platform.device_control_env_var, raising=False)
# CVD unset: visible ordinal == physical ID, even beyond the
# assigned list's length.
assert current_platform.visible_device_id_to_physical_device_id(2) == 2
assert current_platform.visible_device_id_to_physical_device_id(3) == 3
monkeypatch.setenv(current_platform.device_control_env_var, "4,5")
assert current_platform.visible_device_id_to_physical_device_id(1) == 5
with pytest.raises(IndexError, match="out of range"):
current_platform.visible_device_id_to_physical_device_id(2)
class TestDpDeviceIdSharding:
def test_dp_supervisor_device_ids_stay_env_relative(self):
"""Regression test: the DP supervisor must pass env-relative indices,
not physical IDs, because each child re-resolves --device-ids
against its inherited device-control env var."""
import argparse
from vllm.entrypoints.openai.dp_supervisor import _build_device_ids
args = argparse.Namespace(
tensor_parallel_size=2, pipeline_parallel_size=1, device_ids=None
)
assert _build_device_ids(args, local_rank=0) == [0, 1]
assert _build_device_ids(args, local_rank=1) == [2, 3]
def test_dp_supervisor_shards_user_device_ids(self):
"""User-provided --device-ids are sharded across DP children."""
import argparse
from vllm.entrypoints.openai.dp_supervisor import _build_device_ids
args = argparse.Namespace(
tensor_parallel_size=2, pipeline_parallel_size=1, device_ids=[4, 5, 6, 7]
)
assert _build_device_ids(args, local_rank=0) == [4, 5]
assert _build_device_ids(args, local_rank=1) == [6, 7]
with pytest.raises(ValueError, match="needs devices"):
_build_device_ids(args, local_rank=2)
def test_dp_rank_shards_user_assigned_gpu_ids(self):
"""get_physical_gpu_ids_for_local_dp_rank slices the user-provided
--device-ids list instead of recomputing from the env var."""
from vllm.platforms import current_platform
from vllm.v1.engine.utils import get_physical_gpu_ids_for_local_dp_rank
evar = current_platform.device_control_env_var
assert get_physical_gpu_ids_for_local_dp_rank(
evar, local_dp_rank=1, world_size=2, user_assigned_gpu_ids=[4, 5, 6, 7]
) == [6, 7]
with pytest.raises(ValueError, match="needs devices"):
get_physical_gpu_ids_for_local_dp_rank(
evar, local_dp_rank=2, world_size=2, user_assigned_gpu_ids=[4, 5, 6, 7]
)
@@ -8,6 +8,8 @@ AnthropicServingMessages._convert_anthropic_to_openai_request().
Also covers extended-thinking edge cases such as ``redacted_thinking``
blocks echoed back by Anthropic clients, and streaming conversion in
``message_stream_converter``.
Also covers cache usage computation in ``_build_anthropic_usage``.
"""
import json
@@ -18,7 +20,11 @@ import pytest
from vllm.entrypoints.anthropic.protocol import (
AnthropicMessagesRequest,
)
from vllm.entrypoints.anthropic.serving import AnthropicServingMessages
from vllm.entrypoints.anthropic.serving import (
AnthropicServingMessages,
_build_anthropic_usage,
_get_cached_tokens,
)
from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionResponseStreamChoice,
ChatCompletionStreamResponse,
@@ -27,6 +33,7 @@ from vllm.entrypoints.openai.engine.protocol import (
DeltaFunctionCall,
DeltaMessage,
DeltaToolCall,
PromptTokenUsageInfo,
UsageInfo,
)
@@ -653,6 +660,108 @@ class TestThinkingBlockConversion:
assert asst.get("content") == "Hi!"
# ======================================================================
# Cache usage computation
# ======================================================================
class TestGetCachedTokens:
"""Tests for _get_cached_tokens helper."""
def test_none_usage(self):
assert _get_cached_tokens(None) is None
def test_no_prompt_tokens_details(self):
usage = UsageInfo(prompt_tokens=100, completion_tokens=10)
assert _get_cached_tokens(usage) is None
def test_cached_tokens_present(self):
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
)
assert _get_cached_tokens(usage) == 80
def test_cached_tokens_zero(self):
"""Zero cached tokens should return 0, not None."""
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
)
assert _get_cached_tokens(usage) == 0
def test_cached_tokens_none_in_details(self):
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=None),
)
assert _get_cached_tokens(usage) is None
class TestBuildAnthropicUsage:
"""Tests for _build_anthropic_usage helper.
Anthropic defines: total_input = input_tokens + cache_read + cache_creation
vLLM's prompt_tokens is the total.
"""
def test_no_cache_info(self):
"""When cache info is unavailable, return raw prompt_tokens."""
result = _build_anthropic_usage(100, 10, None)
assert result.input_tokens == 100
assert result.output_tokens == 10
assert result.cache_read_input_tokens is None
assert result.cache_creation_input_tokens is None
def test_cache_hit(self):
"""When cache is hit, input_tokens excludes cached tokens."""
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
)
result = _build_anthropic_usage(100, 10, usage)
assert result.input_tokens == 20 # 100 - 80
assert result.output_tokens == 10
assert result.cache_read_input_tokens == 80
assert result.cache_creation_input_tokens == 0
def test_zero_cached_tokens(self):
"""Zero cached tokens should still set cache_creation to 0."""
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
)
result = _build_anthropic_usage(100, 10, usage)
assert result.input_tokens == 100 # 100 - 0
assert result.cache_read_input_tokens == 0
assert result.cache_creation_input_tokens == 0
def test_all_tokens_cached(self):
"""When all tokens are cached, input_tokens should be 0."""
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=100),
)
result = _build_anthropic_usage(100, 10, usage)
assert result.input_tokens == 0
assert result.cache_read_input_tokens == 100
assert result.cache_creation_input_tokens == 0
def test_no_prompt_tokens_details(self):
"""UsageInfo without prompt_tokens_details returns no cache info."""
usage = UsageInfo(prompt_tokens=100, completion_tokens=10)
result = _build_anthropic_usage(100, 10, usage)
assert result.input_tokens == 100
assert result.cache_read_input_tokens is None
assert result.cache_creation_input_tokens is None
class TestInlineSystemMessageInMessagesArray:
"""Verify that ``role: system`` messages embedded inside the ``messages``
array are preserved in their original position.
@@ -1098,6 +1207,135 @@ class TestMessageStartIncludesTypeAndRole:
assert message["role"] == "assistant"
class TestStreamingCacheUsageSemantics:
"""Locks in the documented streaming behavior of cache usage fields.
vLLM's OpenAI chat completion streaming only attaches
``prompt_tokens_details`` to the terminal usage chunk. The Anthropic layer
mirrors that contract: cache fields are omitted on ``message_start`` (key
absence signals "unknown") and populated on ``message_delta`` (the final
cumulative count). This is intentionally consistent with vLLM's OpenAI
behavior, even though Anthropic's upstream API populates cache fields on
``message_start``; closing that gap requires plumbing cache info into the
first chunk at the OpenAI layer, which is out of scope here.
"""
@pytest.mark.asyncio
async def test_streaming_cache_fields_absent_then_populated(self):
"""First chunk lacks prompt_tokens_details (vLLM contract);
message_start omits cache fields. The final chunk carries
prompt_tokens_details, so message_delta carries resolved values."""
async def sse_input():
yield _make_stream_chunk(
delta=DeltaMessage(role="assistant", content="hi"),
usage=UsageInfo(prompt_tokens=100, total_tokens=100),
)
yield _make_stream_chunk(finish_reason="stop")
yield _make_stream_chunk(
choices=[],
usage=UsageInfo(
prompt_tokens=100,
completion_tokens=5,
total_tokens=105,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
),
)
yield "data: [DONE]"
converter = _make_stream_converter()
output = []
async for event in converter.message_stream_converter(sse_input()):
output.append(event)
events = _parse_sse_events(output)
# message_start: cache fields unknown → omitted from JSON entirely.
start_usage = events[0][1]["message"]["usage"]
assert events[0][0] == "message_start"
assert start_usage["input_tokens"] == 100
assert "cache_read_input_tokens" not in start_usage
assert "cache_creation_input_tokens" not in start_usage
# message_delta: authoritative usage with cache fields populated.
delta_usage = next(
data["usage"] for ev, data in events if ev == "message_delta"
)
assert delta_usage["input_tokens"] == 20 # 100 - 80
assert delta_usage["cache_read_input_tokens"] == 80
assert delta_usage["cache_creation_input_tokens"] == 0
@pytest.mark.asyncio
async def test_streaming_no_cache_hit(self):
"""When the final chunk reports cached_tokens=0, message_delta carries
cache fields = 0 (cache miss); message_start still omits them."""
async def sse_input():
yield _make_stream_chunk(
delta=DeltaMessage(role="assistant"),
usage=UsageInfo(prompt_tokens=50, total_tokens=50),
)
yield _make_stream_chunk(finish_reason="stop")
yield _make_stream_chunk(
choices=[],
usage=UsageInfo(
prompt_tokens=50,
completion_tokens=5,
total_tokens=55,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
),
)
yield "data: [DONE]"
converter = _make_stream_converter()
output = []
async for event in converter.message_stream_converter(sse_input()):
output.append(event)
events = _parse_sse_events(output)
start_usage = events[0][1]["message"]["usage"]
delta_usage = next(
data["usage"] for ev, data in events if ev == "message_delta"
)
assert start_usage["input_tokens"] == 50
assert "cache_read_input_tokens" not in start_usage
assert "cache_creation_input_tokens" not in start_usage
assert delta_usage["input_tokens"] == 50 # 50 - 0
assert delta_usage["cache_read_input_tokens"] == 0
assert delta_usage["cache_creation_input_tokens"] == 0
@pytest.mark.asyncio
async def test_streaming_no_prompt_tokens_details_at_all(self):
"""If --enable-prompt-tokens-details is off, no chunk carries cache
info; both message_start and message_delta omit cache fields."""
async def sse_input():
yield _make_stream_chunk(
delta=DeltaMessage(role="assistant"),
usage=UsageInfo(prompt_tokens=30, total_tokens=30),
)
yield _make_stream_chunk(finish_reason="stop")
yield _make_stream_chunk(
choices=[],
usage=UsageInfo(prompt_tokens=30, completion_tokens=2, total_tokens=32),
)
yield "data: [DONE]"
converter = _make_stream_converter()
output = []
async for event in converter.message_stream_converter(sse_input()):
output.append(event)
events = _parse_sse_events(output)
start_usage = events[0][1]["message"]["usage"]
delta_usage = next(
data["usage"] for ev, data in events if ev == "message_delta"
)
assert "cache_read_input_tokens" not in start_usage
assert "cache_creation_input_tokens" not in start_usage
assert "cache_read_input_tokens" not in delta_usage
assert "cache_creation_input_tokens" not in delta_usage
# ======================================================================
# Auto-detection of system-first template requirement
# ======================================================================
@@ -364,7 +364,7 @@ class MockVLLMServer:
await self._serve_task
def launch_mock_vllm(child_args: argparse.Namespace, env_updates: dict[str, str]):
def launch_mock_vllm(child_args: argparse.Namespace):
logger.info("Launching mock vLLM on port %s", child_args.port)
mock_vllm = MockVLLMServer(
port=child_args.port,
@@ -375,7 +375,7 @@ def launch_mock_vllm(child_args: argparse.Namespace, env_updates: dict[str, str]
def launch_mock_vllm_with_drain(
child_args: argparse.Namespace, env_updates: dict[str, str]
child_args: argparse.Namespace,
):
logger.info("Launching mock vLLM with 15s drain on port %s", child_args.port)
mock_vllm = MockVLLMServer(
+57 -44
View File
@@ -6,15 +6,22 @@ from typing import Final
import pytest
import schemathesis
from hypothesis import HealthCheck, settings
from schemathesis import GenerationConfig
from schemathesis.models import Case
from schemathesis import GenerationMode
from schemathesis.config import (
ChecksConfig,
CoveragePhaseConfig,
GenerationConfig,
PhasesConfig,
PositiveDataAcceptanceConfig,
ProjectConfig,
ProjectsConfig,
SchemathesisConfig,
)
from vllm.platforms import current_platform
from ...utils import RemoteOpenAIServer
schemathesis.experimental.OPEN_API_3_1.enable()
MODEL_NAME = "HuggingFaceTB/SmolVLM-256M-Instruct"
MAXIMUM_IMAGES = 2
_ROCM_TIMEOUT_MULTIPLIER = 3 if current_platform.is_rocm() else 1
@@ -44,21 +51,38 @@ def server():
@pytest.fixture(scope="module")
def get_schema(server):
# avoid generating null (\x00) bytes in strings during test case generation
return schemathesis.openapi.from_uri(
return schemathesis.openapi.from_url(
f"{server.url_root}/openapi.json",
generation_config=GenerationConfig(allow_x00=False),
config=SchemathesisConfig(
projects=ProjectsConfig(
default=ProjectConfig(
generation=GenerationConfig(
allow_x00=False,
modes=[GenerationMode.POSITIVE],
),
checks=ChecksConfig(
positive_data_acceptance=PositiveDataAcceptanceConfig(
enabled=False,
),
),
phases=PhasesConfig(
coverage=CoveragePhaseConfig(enabled=False),
),
),
),
),
)
schema = schemathesis.from_pytest_fixture("get_schema")
schema = schemathesis.pytest.from_fixture("get_schema")
@schemathesis.hook
def before_generate_case(context: schemathesis.hooks.HookContext, strategy):
def before_generate_case(context: schemathesis.HookContext, strategy):
op = context.operation
assert op is not None
def no_invalid_types(case: schemathesis.models.Case):
def no_invalid_types(case: schemathesis.Case):
"""
Skips tool_calls with `"type": "custom"` which schemathesis incorrectly
generates instead of the valid `"type": "function"`.
@@ -68,39 +92,25 @@ def before_generate_case(context: schemathesis.hooks.HookContext, strategy):
-d '{"messages": [{"role": "assistant", "tool_calls": [{"custom": {"input": "", "name": ""}, "id": "", "type": "custom"}]}]}' \
http://localhost:8000/v1/chat/completions
""" # noqa: E501
if hasattr(case, "body") and isinstance(case.body, dict):
if (
"messages" in case.body
and isinstance(case.body["messages"], list)
and len(case.body["messages"]) > 0
):
for message in case.body["messages"]:
if not isinstance(message, dict):
continue
if (
hasattr(case, "body")
and isinstance(case.body, dict)
and "messages" in case.body
and isinstance(case.body["messages"], list)
and len(case.body["messages"]) > 0
):
for message in case.body["messages"]:
if not isinstance(message, dict):
continue
tool_calls = message.get("tool_calls", [])
if isinstance(tool_calls, list):
for tool_call in tool_calls:
if isinstance(tool_call, dict):
if tool_call.get("type") != "function":
return False
if "custom" in tool_call:
return False
# Sometimes structured_outputs.grammar is generated to be empty
# Causing a server error in EBNF grammar parsing
# https://github.com/vllm-project/vllm/pull/22587#issuecomment-3195253421
structured_outputs = case.body.get("structured_outputs", {})
grammar = (
structured_outputs.get("grammar")
if isinstance(structured_outputs, dict)
else None
)
if grammar == "":
# Allow None (will be handled as no grammar)
# But skip empty strings
return False
tool_calls = message.get("tool_calls", [])
if isinstance(tool_calls, list):
for tool_call in tool_calls:
if isinstance(tool_call, dict):
if tool_call.get("type") != "function":
return False
if "custom" in tool_call:
return False
return True
@@ -108,7 +118,6 @@ def before_generate_case(context: schemathesis.hooks.HookContext, strategy):
@schema.parametrize()
@schema.override(headers={"Content-Type": "application/json"})
@settings(
deadline=LONG_TIMEOUT_SECONDS * 1000,
max_examples=50,
@@ -122,7 +131,7 @@ def before_generate_case(context: schemathesis.hooks.HookContext, strategy):
# generating large-but-valid request bodies before vLLM is called.
suppress_health_check=[HealthCheck.filter_too_much, HealthCheck.data_too_large],
)
def test_openapi_stateless(case: Case):
def test_openapi_stateless(case: schemathesis.Case):
key = (
case.operation.method.upper(),
case.operation.path,
@@ -151,4 +160,8 @@ def test_openapi_stateless(case: Case):
}.get(key, DEFAULT_TIMEOUT_SECONDS)
# No need to verify SSL certificate for localhost
case.call_and_validate(verify=False, timeout=timeout)
case.call_and_validate(
verify=False,
timeout=timeout,
headers={"Content-Type": "application/json"},
)
@@ -8,6 +8,7 @@ import pytest
import pytest_asyncio
from tests.utils import RemoteLaunchRenderServer
from vllm.tokenizers import get_tokenizer
MODEL_NAME = "hmellor/tiny-random-LlamaForCausalLM"
@@ -486,3 +487,438 @@ async def test_derender_completion_kv_transfer_params_passthrough(client):
)
assert response.status_code == 200
assert response.json()["kv_transfer_params"] == kv
# ---------------------------------------------------------------------------
# E2E: render -> derender roundtrip with parser (reasoning + tool calls)
# ---------------------------------------------------------------------------
PARSER_MODEL = "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B"
_E2E_TOOLS = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a city",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
}
]
@pytest.fixture(scope="module")
def parser_server():
args = [
"--enable-auto-tool-choice",
"--tool-call-parser",
"hermes",
"--reasoning-parser",
"deepseek_r1",
]
with RemoteLaunchRenderServer(PARSER_MODEL, args) as remote_server:
yield remote_server
@pytest_asyncio.fixture
async def parser_client(parser_server):
async with httpx.AsyncClient(
base_url=parser_server.url_for(""), timeout=60.0
) as http_client:
yield http_client
@pytest.fixture(scope="module")
def parser_tokenizer():
return get_tokenizer(PARSER_MODEL)
def _encode(tokenizer, text: str) -> list[int]:
return tokenizer.encode(text, add_special_tokens=False)
def _decoded(tokenizer, token_ids: list[int]) -> str:
return tokenizer.decode(token_ids, skip_special_tokens=True)
def _require_markers_survive(tokenizer, text: str, *markers: str) -> list[int]:
"""Encode text and skip the test if any marker is lost in roundtrip."""
ids = _encode(tokenizer, text)
decoded = tokenizer.decode(ids, skip_special_tokens=False)
for m in markers:
if m not in decoded:
pytest.skip(f"Marker {m!r} lost in encode->decode roundtrip")
return ids
async def _e2e_render_chat(
client: httpx.AsyncClient,
model: str,
messages: list[dict],
) -> dict:
resp = await client.post(
"/v1/chat/completions/render",
json={"model": model, "messages": messages},
)
assert resp.status_code == 200, resp.text
return resp.json()
def _e2e_generate_response(
token_ids: list[int],
request_id: str = "chatcmpl-e2e-test",
) -> dict:
return {
"request_id": request_id,
"choices": [
{
"index": 0,
"token_ids": token_ids,
"finish_reason": "stop",
}
],
}
@pytest.mark.asyncio
async def test_e2e_plain_roundtrip(parser_client, parser_tokenizer):
"""Plain text without reasoning markers roundtrips correctly."""
messages = [{"role": "user", "content": "What is 2+2?"}]
gen_req = await _e2e_render_chat(parser_client, PARSER_MODEL, messages)
answer = "The answer is four."
output_ids = _encode(parser_tokenizer, answer)
expected = _decoded(parser_tokenizer, output_ids)
resp = await parser_client.post(
"/v1/chat/completions/derender",
json={
"model": PARSER_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
},
)
assert resp.status_code == 200, resp.text
content = resp.json()["choices"][0]["message"]["content"]
assert content == expected
@pytest.mark.asyncio
async def test_e2e_token_identity(parser_client, parser_tokenizer):
"""encode(derender(token_ids)) == token_ids (RL invariant)."""
messages = [{"role": "user", "content": "Hi"}]
gen_req = await _e2e_render_chat(parser_client, PARSER_MODEL, messages)
answer = "Hello! How can I help?"
output_ids = _encode(parser_tokenizer, answer)
resp = await parser_client.post(
"/v1/chat/completions/derender",
json={
"model": PARSER_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
},
)
assert resp.status_code == 200
content = resp.json()["choices"][0]["message"]["content"]
re_encoded = _encode(parser_tokenizer, content)
assert output_ids == re_encoded
@pytest.mark.asyncio
async def test_e2e_non_ascii_roundtrip(parser_client, parser_tokenizer):
"""CJK + emoji roundtrip without U+FFFD."""
messages = [{"role": "user", "content": "Reply in Chinese"}]
gen_req = await _e2e_render_chat(parser_client, PARSER_MODEL, messages)
answer = "你好世界 😀"
output_ids = _encode(parser_tokenizer, answer)
resp = await parser_client.post(
"/v1/chat/completions/derender",
json={
"model": PARSER_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
},
)
assert resp.status_code == 200
content = resp.json()["choices"][0]["message"]["content"]
assert "" not in content
@pytest.mark.asyncio
async def test_e2e_parsed_reasoning(parser_client, parser_tokenizer):
"""<think>...</think> splits into reasoning + content."""
messages = [{"role": "user", "content": "What is 2+3?"}]
gen_req = await _e2e_render_chat(parser_client, PARSER_MODEL, messages)
reasoning_text = "The user wants 2 plus 3. That is 5."
answer_text = "The answer is 5."
output_text = f"<think>{reasoning_text}</think>{answer_text}"
output_ids = _require_markers_survive(parser_tokenizer, output_text, "</think>")
resp = await parser_client.post(
"/v1/chat/completions/derender",
json={
"model": PARSER_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
"chat_request": {
"model": PARSER_MODEL,
"messages": messages,
"include_reasoning": True,
},
},
)
assert resp.status_code == 200, resp.text
msg = resp.json()["choices"][0]["message"]
assert msg["reasoning"] is not None
assert reasoning_text in msg["reasoning"]
assert answer_text in msg["content"]
assert "<think>" not in msg["content"]
@pytest.mark.asyncio
async def test_e2e_parsed_tool_call(parser_client, parser_tokenizer):
"""<tool_call> extracted into tool_calls field."""
messages = [{"role": "user", "content": "Weather in Paris?"}]
gen_req = await _e2e_render_chat(parser_client, PARSER_MODEL, messages)
output_text = (
"<think>Let me check the weather.</think>"
'<tool_call>\n{"name": "get_weather", '
'"arguments": {"city": "Paris"}}\n</tool_call>'
)
output_ids = _require_markers_survive(
parser_tokenizer,
output_text,
"</think>",
"<tool_call>",
"</tool_call>",
)
resp = await parser_client.post(
"/v1/chat/completions/derender",
json={
"model": PARSER_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
"chat_request": {
"model": PARSER_MODEL,
"messages": messages,
"tools": _E2E_TOOLS,
"tool_choice": "auto",
},
},
)
assert resp.status_code == 200, resp.text
choice = resp.json()["choices"][0]
assert choice["message"]["tool_calls"]
assert choice["message"]["tool_calls"][0]["function"]["name"] == "get_weather"
@pytest.mark.asyncio
async def test_e2e_parsed_reasoning_and_tool_call(parser_client, parser_tokenizer):
"""Reasoning + tool call in the same output."""
messages = [{"role": "user", "content": "Weather in Paris?"}]
gen_req = await _e2e_render_chat(parser_client, PARSER_MODEL, messages)
reasoning_text = "I should look up the weather."
tool_text = (
'<tool_call>\n{"name": "get_weather", '
'"arguments": {"city": "Paris"}}\n</tool_call>'
)
output_text = f"<think>{reasoning_text}</think>{tool_text}"
output_ids = _require_markers_survive(
parser_tokenizer, output_text, "</think>", "<tool_call>"
)
resp = await parser_client.post(
"/v1/chat/completions/derender",
json={
"model": PARSER_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
"chat_request": {
"model": PARSER_MODEL,
"messages": messages,
"tools": _E2E_TOOLS,
"tool_choice": "auto",
"include_reasoning": True,
},
},
)
assert resp.status_code == 200, resp.text
choice = resp.json()["choices"][0]
assert choice["message"]["reasoning"] is not None
assert reasoning_text in choice["message"]["reasoning"]
assert choice["message"]["tool_calls"]
@pytest.mark.asyncio
async def test_e2e_no_chat_request_fallback(parser_client, parser_tokenizer):
"""Without chat_request, derender falls back to plain detokenization."""
messages = [{"role": "user", "content": "Hello"}]
gen_req = await _e2e_render_chat(parser_client, PARSER_MODEL, messages)
answer = "Hi there!"
output_ids = _encode(parser_tokenizer, answer)
resp = await parser_client.post(
"/v1/chat/completions/derender",
json={
"model": PARSER_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
},
)
assert resp.status_code == 200
content = resp.json()["choices"][0]["message"]["content"]
assert "Hi" in content
# ---------------------------------------------------------------------------
# E2E: HarmonyParser + GPT-OSS
# ---------------------------------------------------------------------------
HARMONY_MODEL = "openai/gpt-oss-20b"
def _ensure_harmony_vocab():
"""Pre-cache the o200k_base BPE file needed by openai-harmony.
The Rust tiktoken-rs backend downloads from Azure Blob Storage, which
may be unreachable in some environments. When the cache is cold we
fetch the file ourselves and place it in ``/tmp/tiktoken-rs-cache/``
using the SHA-1(URL) filename that tiktoken-rs expects.
"""
import hashlib
import urllib.request
from pathlib import Path
url = "https://openaipublic.blob.core.windows.net/encodings/o200k_base.tiktoken"
cache_dir = Path("/tmp/tiktoken-rs-cache")
cache_key = hashlib.sha1(url.encode()).hexdigest()
cache_file = cache_dir / cache_key
if not cache_file.exists():
cache_dir.mkdir(parents=True, exist_ok=True)
urllib.request.urlretrieve(url, cache_file)
@pytest.fixture(scope="module")
def harmony_server():
_ensure_harmony_vocab()
args = [
"--trust-remote-code",
"--enable-auto-tool-choice",
"--tool-call-parser",
"openai",
"--reasoning-parser",
"openai_gptoss",
]
with RemoteLaunchRenderServer(HARMONY_MODEL, args) as remote_server:
yield remote_server
@pytest_asyncio.fixture
async def harmony_client(harmony_server):
async with httpx.AsyncClient(
base_url=harmony_server.url_for(""), timeout=60.0
) as http_client:
yield http_client
@pytest.fixture(scope="module")
def harmony_tokenizer():
return get_tokenizer(HARMONY_MODEL, trust_remote_code=True)
def _harmony_extract_assistant_ids(
tokenizer, assistant_msg: dict, user_content: str = "test"
) -> list[int]:
"""Extract assistant token IDs via apply_chat_template diff."""
prompt = [{"role": "user", "content": user_content}]
full = prompt + [assistant_msg]
text_prompt = tokenizer.apply_chat_template(
prompt, add_generation_prompt=True, tokenize=False
)
text_full = tokenizer.apply_chat_template(
full, add_generation_prompt=False, tokenize=False
)
prompt_ids = tokenizer.encode(text_prompt)
full_ids = tokenizer.encode(text_full)
assistant_ids = list(full_ids[len(prompt_ids) :])
if not assistant_ids:
pytest.skip("Could not extract assistant tokens for Harmony")
return assistant_ids
@pytest.mark.asyncio
async def test_e2e_harmony_plain_roundtrip(harmony_client, harmony_tokenizer):
"""GPT-OSS content-only roundtrip."""
messages = [{"role": "user", "content": "What is 2+2?"}]
gen_req = await _e2e_render_chat(harmony_client, HARMONY_MODEL, messages)
assistant_msg = {"role": "assistant", "content": "Four."}
output_ids = _harmony_extract_assistant_ids(harmony_tokenizer, assistant_msg)
resp = await harmony_client.post(
"/v1/chat/completions/derender",
json={
"model": HARMONY_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
"chat_request": {
"model": HARMONY_MODEL,
"messages": messages,
},
},
)
assert resp.status_code == 200, resp.text
content = resp.json()["choices"][0]["message"]["content"]
assert content is not None and len(content) > 0
assert "Four" in content
@pytest.mark.asyncio
async def test_e2e_harmony_reasoning(harmony_client, harmony_tokenizer):
"""GPT-OSS reasoning: analysis channel extracted."""
messages = [{"role": "user", "content": "Add 2 and 3."}]
gen_req = await _e2e_render_chat(harmony_client, HARMONY_MODEL, messages)
reasoning_text = "The user wants 2 plus 3."
answer_text = "The answer is 5."
assistant_msg = {
"role": "assistant",
"thinking": reasoning_text,
"content": answer_text,
}
output_ids = _harmony_extract_assistant_ids(harmony_tokenizer, assistant_msg)
decoded = harmony_tokenizer.decode(output_ids)
if reasoning_text not in decoded:
pytest.skip("Harmony template did not render thinking")
resp = await harmony_client.post(
"/v1/chat/completions/derender",
json={
"model": HARMONY_MODEL,
"generate_response": _e2e_generate_response(output_ids),
"prompt_tokens": len(gen_req["token_ids"]),
"chat_request": {
"model": HARMONY_MODEL,
"messages": messages,
"include_reasoning": True,
},
},
)
assert resp.status_code == 200, resp.text
msg = resp.json()["choices"][0]["message"]
assert msg["reasoning"] is not None
assert reasoning_text in msg["reasoning"]
assert answer_text in (msg["content"] or "")
@@ -78,7 +78,16 @@ def test_gsm8k_correctness(config_filename):
"Skipping DeepSeek-V3.2 and DeepSeek-R1 on ROCm platforms "
"due to agent pool disk space issues and pod evictions."
)
if current_platform.is_rocm() and (
"Qwen3.5-35B-A3B-MXFP4-AITER-TP2" in config_filename.name
):
from vllm.platforms.rocm import on_gfx950
if not on_gfx950():
pytest.skip(
"Skipping Qwen3.5-35B-A3B-MXFP4-AITER-TP2 on non-GFX950 platforms. "
"The quantization scheme is not supported on non-GFX950 platforms."
)
# Parse server arguments from config (use shlex to handle quoted strings)
server_args_str = eval_config.get("server_args", "")
server_args = shlex.split(server_args_str) if server_args_str else []
@@ -345,6 +345,63 @@ def test_per_token_group_quant_fp8_packed_zero_fills_padded_output_q(
)
@pytest.mark.skipif(
not current_platform.is_cuda_alike(),
reason="packed FP8 per-token-group quant kernel requires a CUDA-alike GPU",
)
def test_per_token_group_quant_fp8_packed_large_mn():
"""Regression test for https://github.com/vllm-project/vllm/issues/45099.
Some background: gridDim.x and gridDim.y have different limits of 2^31 - 1 and
2^16 - 1, respectively.
Prior code introduced a bug where it incorrectly assumed grid.x and y both have
2^31 - 1 limits and mixed them up, which doesn't surface until the kernel is
launched with a large mn that exceeds grid.y limit (2^16 - 1).
This issue doesn't surface often because each forward pass only processes a
bounded token batch, not the full context.
Quantizing tensors with more rows than that will fail at launch with
"CUDA error: invalid argument".
This is a differential test that compares fp8 output against Triton output
reference when token size sits just above the gridDim.y 2^16 - 1 limit.
"""
device = "cuda"
group_size = 128
# hidden 2048 -> 2048/128 = 16 groups per row -> kx=16, ry=1: one grid row per mn
# row, so any mn > 65535 overflowed grid.y before the fix.
num_tokens, hidden_dim = 65537, 2048
torch.manual_seed(42)
x = torch.randn((num_tokens, hidden_dim), device=device, dtype=torch.bfloat16) * 8
out_q, out_s_packed = fp8_utils.per_token_group_quant_fp8_packed_for_deepgemm(
x,
group_size=group_size,
use_ue8m0=True,
)
with patch("vllm.platforms.current_platform.is_cuda_alike", return_value=False):
ref_q, ref_s = fp8_utils.per_token_group_quant_fp8(
x, group_size, use_ue8m0=True
)
assert torch.equal(out_q, ref_q), "Quantized output mismatch"
# Vectorized packed-scale check; the per-element loop used by the smaller
# tests is too slow at this size. groups_per_row is a multiple of 4 here,
# so there is no K padding and the packed view lines up.
mn = num_tokens
groups_per_row = hidden_dim // group_size
k_num_packed = (groups_per_row + 3) // 4
assert groups_per_row % 4 == 0
ref_exponents = (ref_s.reshape(mn, groups_per_row).view(torch.int32) >> 23) & 0xFF
exp = ref_exponents.view(mn, k_num_packed, 4)
expected = (
exp[..., 0] | (exp[..., 1] << 8) | (exp[..., 2] << 16) | (exp[..., 3] << 24)
)
assert torch.equal(out_s_packed.cpu(), expected.cpu()), "Packed scale mismatch"
@pytest.mark.parametrize("shape", [(32, 128), (64, 256), (16, 512)])
@pytest.mark.parametrize("group_size", [64, 128])
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
@@ -60,8 +60,10 @@ def test_rocm_compressed_tensors_w8a8(
vllm_runner, example_prompts, model_path, max_tokens, num_logprobs
):
dtype = "bfloat16"
with vllm_runner(model_path, dtype=dtype) as vllm_model:
# Pin to TRITON_ATTN, see https://github.com/vllm-project/vllm/issues/46179
with vllm_runner(
model_path, dtype=dtype, attention_backend="TRITON_ATTN"
) as vllm_model:
vllm_model.generate_greedy_logprobs(example_prompts, max_tokens, num_logprobs)
@@ -2,6 +2,7 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
import torch
import torch.nn.functional as F
import transformers.utils
from PIL import Image
@@ -52,6 +53,7 @@ def _get_cherry_blossom_image() -> Image.Image:
)
@torch.inference_mode()
def _run_test(
hf_runner: type[HfRunner],
vllm_runner: type[VllmRunner],
@@ -92,3 +92,49 @@ def test_processor_num_frames_timestamp(
assert len(video_phs) == 1, (
f"Expected exactly 1 video placeholder, got {len(video_phs)}"
)
@pytest.mark.parametrize("model_id", [MODEL_ID])
@pytest.mark.parametrize("num_videos", [2, 4])
def test_processor_multi_video(
model_id: str,
num_videos: int,
) -> None:
"""Verify that multi-video processing produces correct placeholders.
This exercises the token-level replacement path in
``_call_hf_processor`` which avoids the quadratic text-level
prompt expansion.
"""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 0, "video": num_videos},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
prompt = "<|vision_start|><|video_pad|><|vision_end|>" * num_videos
mm_data = {"video": [_build_video_mm_data(num_frames=8)["video"][0]] * num_videos}
processed = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs={"num_frames": 8},
)
token_ids = processed["prompt_token_ids"]
assert len(token_ids) > 0
video_phs = processed["mm_placeholders"].get("video", [])
assert len(video_phs) == num_videos, (
f"Expected {num_videos} video placeholders, got {len(video_phs)}"
)
# All placeholders should have the same length (same video params)
# and must not overlap.
lengths = {ph.length for ph in video_phs}
assert len(lengths) == 1, f"Placeholder lengths differ: {lengths}"
for i in range(1, len(video_phs)):
prev_end = video_phs[i - 1].offset + video_phs[i - 1].length
assert video_phs[i].offset >= prev_end, (
f"Placeholder {i} overlaps with placeholder {i - 1}"
)
+3
View File
@@ -96,6 +96,9 @@ class MockTokenizer:
return "".join(parts)
CHUNK_SIZES = [1, 2, 3, 5, 11, 23, None]
def make_mock_tokenizer(sample: Sample) -> MockTokenizer:
"""Build a mock tokenizer from a sample's vocab and token data."""
return MockTokenizer(
@@ -19,6 +19,7 @@ import pytest
from pydantic import TypeAdapter
from tests.parser.engine.replay_harness import (
CHUNK_SIZES,
MockTokenizer,
assert_parse_output,
collect_output,
@@ -113,8 +114,6 @@ _PAIRINGS = _discover_pairings()
_ALL_SAMPLES = [(p.parser_cls, s) for p in _PAIRINGS for s in p.samples]
CHUNK_SIZES = [1, 2, 3, 5, 11, 23, None]
@pytest.mark.parametrize("chunk_size", CHUNK_SIZES, ids=lambda c: f"chunk={c}")
@pytest.mark.parametrize(
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Regression test for U+FFFD leak at reasoning→content transition.
When byte-fallback tokens span the reasoning/content boundary,
decoding isolated content-side token IDs via tokenizer.decode()
produces U+FFFD (Unicode replacement character). The fix flushes
the reasoning parser's engine lexer instead.
Reproduces the bug at various chunk sizes and validates that the
fix prevents U+FFFD from leaking into streamed content.
"""
from __future__ import annotations
import pytest
from tests.parser.engine.replay_harness import (
CHUNK_SIZES,
MockTokenizer,
collect_output,
replay_streaming,
)
from vllm.parser.abstract_parser import DelegatingParser
from vllm.parser.engine.registered_adapters import (
Glm47MoeParserReasoningAdapter,
Glm47MoeParserToolAdapter,
Qwen3ParserReasoningAdapter,
Qwen3ParserToolAdapter,
)
class ByteFallbackMockTokenizer(MockTokenizer):
"""MockTokenizer that returns U+FFFD for specified token IDs.
Simulates byte-fallback tokenizer behavior where isolated
partial-byte tokens decode to the Unicode replacement character.
"""
def __init__(
self,
vocab: dict[str, int],
tokens: list[tuple[int, str]],
ufffd_token_ids: set[int],
) -> None:
super().__init__(vocab, tokens)
self._ufffd_token_ids = frozenset(ufffd_token_ids)
def decode(self, ids: list[int], skip_special_tokens: bool = False) -> str:
parts: list[str] = []
for tid in ids:
if skip_special_tokens and tid in self._special_ids:
continue
if tid in self._ufffd_token_ids:
parts.append("")
else:
text = self._token_decode_map.get(tid, f"?{tid}?")
parts.append(text)
return "".join(parts)
# ── Model-specific DelegatingParser subclasses ───────────────────────
class _Glm47Delegating(DelegatingParser):
reasoning_parser_cls = Glm47MoeParserReasoningAdapter
tool_parser_cls = Glm47MoeParserToolAdapter
class _Qwen3Delegating(DelegatingParser):
reasoning_parser_cls = Qwen3ParserReasoningAdapter
tool_parser_cls = Qwen3ParserToolAdapter
# ── Shared test data ─────────────────────────────────────────────────
_SHARED_TOKENS: list[tuple[int, str]] = [
(100, "Let me"),
(101, " think"),
(102, " about"),
(103, " Samsung."),
(51, "</think>"),
(200, "삼성"),
(201, "전자의"),
(202, " 주가를"),
(203, " 분석합니다."),
]
_SHARED_UFFFD_IDS: set[int] = {200}
EXPECTED_REASONING = "Let me think about Samsung."
EXPECTED_CONTENT = "삼성전자의 주가를 분석합니다."
_MODEL_CONFIGS = [
pytest.param(
{
"<think>": 50,
"</think>": 51,
"<tool_call>": 60,
"</tool_call>": 61,
"<arg_key>": 62,
"</arg_key>": 63,
"<arg_value>": 64,
"</arg_value>": 65,
},
_Glm47Delegating,
id="glm47",
),
pytest.param(
{
"<think>": 50,
"</think>": 51,
"<tool_call>": 60,
"</tool_call>": 61,
},
_Qwen3Delegating,
id="qwen3",
),
]
# ── Tests ────────────────────────────────────────────────────────────
class TestUfffdReasoningTransition:
"""U+FFFD must not appear at the reasoning→content transition."""
@pytest.mark.parametrize("vocab,delegating_cls", _MODEL_CONFIGS)
@pytest.mark.parametrize("chunk_size", CHUNK_SIZES, ids=lambda c: f"chunk={c}")
def test_no_ufffd(self, chunk_size, vocab, delegating_cls):
tokenizer = ByteFallbackMockTokenizer(vocab, _SHARED_TOKENS, _SHARED_UFFFD_IDS)
parser = delegating_cls(tokenizer)
deltas = replay_streaming(
parser,
_SHARED_TOKENS,
chunk_size=chunk_size,
finished_on_last=True,
)
output = collect_output(deltas)
assert "" not in output.content, (
f"U+FFFD leaked into content: {output.content!r}"
)
assert output.content == EXPECTED_CONTENT
assert output.reasoning == EXPECTED_REASONING
def test_byte_fallback_tokenizer_produces_ufffd(self):
"""Validate the fixture: decode() returns U+FFFD for isolated
byte-fallback token IDs, proving the old code path would leak."""
vocab = dict(_MODEL_CONFIGS[0].values[0])
tokenizer = ByteFallbackMockTokenizer(vocab, _SHARED_TOKENS, _SHARED_UFFFD_IDS)
assert tokenizer.decode([200]) == ""
@pytest.mark.parametrize("chunk_size", CHUNK_SIZES, ids=lambda c: f"chunk={c}")
def test_multiple_ufffd_tokens_at_boundary(self, chunk_size):
"""Multiple consecutive byte-fallback tokens at the boundary."""
tokens: list[tuple[int, str]] = [
(100, "Reasoning."),
(51, "</think>"),
(200, ""),
(201, ""),
(202, "전자"),
]
ufffd_ids: set[int] = {200, 201}
vocab = dict(_MODEL_CONFIGS[0].values[0])
tokenizer = ByteFallbackMockTokenizer(vocab, tokens, ufffd_ids)
parser = _Glm47Delegating(tokenizer)
deltas = replay_streaming(
parser,
tokens,
chunk_size=chunk_size,
finished_on_last=True,
)
output = collect_output(deltas)
assert "" not in output.content, (
f"U+FFFD leaked into content: {output.content!r}"
)
assert output.content == "삼성전자"
assert output.reasoning == "Reasoning."
@@ -83,6 +83,20 @@ class MiniMaxM3Tokenizer:
return "".join(tokens)
class SplitMiniMaxM3Tokenizer(MiniMaxM3Tokenizer):
"""Tokenizer that exposes marker vocab entries but encodes them as text."""
def tokenize(self, text: str) -> list[str]:
return list(text)
class RuntimeSplitMiniMaxM3Tokenizer(MiniMaxM3Tokenizer):
"""Tokenizer whose runtime output splits markers despite atomic encodes."""
def encode_runtime(self, text: str) -> list[int]:
return [self._add_token(token) for token in list(text)]
def make_parser(
chat_template_kwargs: dict[str, str] | None = None,
) -> tuple[MiniMaxM3ReasoningParser, MiniMaxM3Tokenizer]:
@@ -105,7 +119,8 @@ def run_streaming(
reasoning_end_states: list[bool] = []
for chunk in chunks:
delta_token_ids = tokenizer.encode(chunk, add_special_tokens=False)
encode_runtime = getattr(tokenizer, "encode_runtime", tokenizer.encode)
delta_token_ids = encode_runtime(chunk)
current_text = previous_text + chunk
current_token_ids = previous_token_ids + delta_token_ids
delta = parser.extract_reasoning_streaming(
@@ -174,14 +189,14 @@ def test_nonstreaming_drops_leading_end_tag():
assert content == "answer"
def test_nonstreaming_non_leading_end_tag_is_content():
def test_nonstreaming_end_tag_in_content_state_is_dropped():
parser, _ = make_parser()
request = ChatCompletionRequest(messages=[], model="test-model")
reasoning, content = parser.extract_reasoning("XXX</mm:think>YYY", request)
assert reasoning is None
assert content == "XXX</mm:think>YYY"
assert content == "XXXYYY"
def test_nonstreaming_enabled_mode_starts_in_reasoning():
@@ -246,7 +261,7 @@ def test_streaming_drops_leading_end_tag():
assert end_states == [True, True]
def test_streaming_non_leading_end_tag_is_content():
def test_streaming_end_tag_in_content_state_is_dropped():
parser, tokenizer = make_parser()
reasoning, content, end_states = run_streaming(
@@ -256,7 +271,7 @@ def test_streaming_non_leading_end_tag_is_content():
)
assert reasoning is None
assert content == "XXX</mm:think>YYY"
assert content == "XXXYYY"
assert end_states == [True]
@@ -288,6 +303,110 @@ def test_streaming_plain_content_ends_reasoning_phase():
assert end_states == [True, True]
def test_streaming_split_marker_tokens_are_not_returned():
tokenizer = RuntimeSplitMiniMaxM3Tokenizer()
parser = MiniMaxM3ReasoningParser(tokenizer)
reasoning, content, end_states = run_streaming(
parser,
tokenizer,
["<mm:think>", "Reasoning", " content", "</mm:think>", "content"],
)
assert reasoning == "Reasoning content"
assert content == "content"
assert end_states == [False, False, False, True, True]
def test_streaming_split_marker_text_drives_end_state():
tokenizer = RuntimeSplitMiniMaxM3Tokenizer()
parser = MiniMaxM3ReasoningParser(tokenizer)
previous_text = ""
previous_token_ids: list[int] = []
for chunk in ["<mm:think>", "Reasoning", " content", "</mm:think>"]:
delta_token_ids = tokenizer.encode_runtime(chunk)
current_text = previous_text + chunk
current_token_ids = previous_token_ids + delta_token_ids
parser.extract_reasoning_streaming(
previous_text=previous_text,
current_text=current_text,
delta_text=chunk,
previous_token_ids=previous_token_ids,
current_token_ids=current_token_ids,
delta_token_ids=delta_token_ids,
)
previous_text = current_text
previous_token_ids = current_token_ids
assert parser.is_reasoning_end_streaming(previous_token_ids, []) is True
def test_streaming_split_marker_tokens_enabled_mode():
tokenizer = RuntimeSplitMiniMaxM3Tokenizer()
parser = MiniMaxM3ReasoningParser(
tokenizer, chat_template_kwargs={"thinking_mode": "enabled"}
)
reasoning, content, end_states = run_streaming(
parser,
tokenizer,
["Reasoning", " content", "</mm:think>", "content"],
)
assert reasoning == "Reasoning content"
assert content == "content"
assert end_states == [False, False, True, True]
def test_streaming_split_marker_text_across_deltas():
tokenizer = RuntimeSplitMiniMaxM3Tokenizer()
parser = MiniMaxM3ReasoningParser(tokenizer)
reasoning, content, end_states = run_streaming(
parser,
tokenizer,
["<mm:", "think>", "Reasoning", " content", "</mm:", "think>", "content"],
)
assert reasoning == "Reasoning content"
assert content == "content"
assert end_states == [False, False, False, False, False, True, True]
def test_streaming_split_leading_end_marker_text_across_deltas():
tokenizer = RuntimeSplitMiniMaxM3Tokenizer()
parser = MiniMaxM3ReasoningParser(tokenizer)
reasoning, content, end_states = run_streaming(
parser,
tokenizer,
["</mm:", "think>", "content"],
)
assert reasoning is None
assert content == "content"
assert end_states == [False, True, True]
def test_token_id_helpers_with_split_marker_tokens():
tokenizer = SplitMiniMaxM3Tokenizer()
parser = MiniMaxM3ReasoningParser(tokenizer)
output_ids = tokenizer.encode(
"<mm:think>abc</mm:think>def", add_special_tokens=False
)
open_reasoning_ids = tokenizer.encode("<mm:think>abc", add_special_tokens=False)
content_ids = tokenizer.encode("plain", add_special_tokens=False)
assert parser.is_reasoning_end(output_ids)
assert not parser.is_reasoning_end(open_reasoning_ids)
assert not parser.is_reasoning_end(content_ids)
assert tokenizer.decode(parser.extract_content_ids(output_ids)) == "def"
assert parser.extract_content_ids(open_reasoning_ids) == []
assert parser.extract_content_ids(content_ids) == content_ids
assert parser.count_reasoning_tokens(output_ids) == len(tokenizer.encode("abc"))
def test_token_id_helpers():
parser, tokenizer = make_parser()
output_ids = tokenizer.encode(
@@ -18,6 +18,7 @@ from vllm.model_executor.kernels.linear.scaled_mm.ScaledMMLinearKernel import (
FP8ScaledMMLinearLayerConfig,
)
from vllm.model_executor.layers.quantization.utils.quant_utils import (
get_fp8_min_max,
kFp8DynamicTokenSym,
kFp8StaticChannelSym,
kFp8StaticTensorSym,
@@ -309,7 +310,7 @@ def test_hipb_mm_kernel_forward_accuracy(enable_hipb_mm_kernel):
_check_bpreshuffle_runtime_support(weight_shape, num_tokens=num_tokens)
fp8_dtype = current_platform.fp8_dtype()
fp8_max = torch.finfo(fp8_dtype).max
fp8_max = get_fp8_min_max()[1]
device = torch.device("cuda")
# Build a bf16 weight and quantize per output channel (one scale per row).
+228
View File
@@ -0,0 +1,228 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for contiguous KV cache packing."""
from unittest.mock import MagicMock
import pytest
import torch
from vllm import envs
from vllm.v1.core.kv_cache_utils import (
_get_kv_cache_config_deepseek_v4,
get_kv_cache_config_from_groups,
)
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheGroupSpec,
KVCacheTensor,
MLAAttentionSpec,
SlidingWindowSpec,
UniformTypeKVCacheSpecs,
)
def _make_mla_spec(page_size: int, block_size: int = 256) -> MLAAttentionSpec:
return MLAAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=512,
dtype=torch.uint8,
page_size_padded=page_size,
cache_dtype_str="fp8_ds_mla",
model_version="deepseek_v4",
alignment=576,
)
def _make_full_spec() -> FullAttentionSpec:
return FullAttentionSpec(
block_size=16,
num_kv_heads=2,
head_size=64,
dtype=torch.float16,
)
def _make_sw_spec() -> SlidingWindowSpec:
return SlidingWindowSpec(
block_size=16,
num_kv_heads=2,
head_size=64,
dtype=torch.float16,
sliding_window=128,
)
def _make_groups(n_c4, n_c128, n_swa):
PS_C4_MLA = 37440
PS_C4_IDX = 8640
PS_C128 = 1728
PS_SWA = 37440
mla_specs = {}
for i in range(n_c4):
mla_specs[f"c4_mla.{i}"] = _make_mla_spec(PS_C4_MLA)
mla_specs[f"c4_idx.{i}"] = _make_mla_spec(PS_C4_IDX)
for i in range(n_c128):
mla_specs[f"c128_mla.{i}"] = _make_mla_spec(PS_C128)
mla_group = KVCacheGroupSpec(
layer_names=list(mla_specs.keys()),
kv_cache_spec=UniformTypeKVCacheSpecs(block_size=256, kv_cache_specs=mla_specs),
)
swa_specs = {}
for i in range(n_swa):
swa_specs[f"swa.{i}"] = _make_mla_spec(PS_SWA)
swa_group = KVCacheGroupSpec(
layer_names=list(swa_specs.keys()),
kv_cache_spec=UniformTypeKVCacheSpecs(block_size=256, kv_cache_specs=swa_specs),
)
return [mla_group, swa_group]
def _mock_vllm_config():
config = MagicMock()
config.cache_config.num_gpu_blocks_override = None
return config
def _run(n_c4=3, n_c128=2, n_swa=5, mem=100 * 1024 * 1024):
groups = _make_groups(n_c4, n_c128, n_swa)
return _get_kv_cache_config_deepseek_v4(_mock_vllm_config(), groups, mem)
def _page_sizes_by_layer(
groups: list[KVCacheGroupSpec],
) -> dict[str, int]:
page_sizes = {}
for group in groups:
specs = group.kv_cache_spec.kv_cache_specs
for layer_name in group.layer_names:
page_sizes[layer_name] = specs[layer_name].page_size_bytes
return page_sizes
class TestInterleavedPacking:
def test_all_tensors_have_block_stride(self):
_, tensors = _run()
for t in tensors:
assert t.block_stride > 0
def test_all_tensors_share_same_size(self):
_, tensors = _run()
sizes = set(t.size for t in tensors)
assert len(sizes) == 1
assert sizes.pop() > 0
def test_offsets_within_one_block(self):
_, tensors = _run()
for t in tensors:
assert t.offset < t.block_stride
def test_all_layers_accounted_for(self):
n_c4, n_c128, n_swa = 5, 4, 7
_, tensors = _run(n_c4=n_c4, n_c128=n_c128, n_swa=n_swa)
all_names = set()
for t in tensors:
all_names.update(t.shared_by)
expected = n_c4 * 2 + n_c128 + n_swa
assert len(all_names) == expected
def test_strided_views_are_independent(self):
groups = _make_groups(n_c4=3, n_c128=2, n_swa=5)
page_sizes = _page_sizes_by_layer(groups)
num_blocks, tensors = _get_kv_cache_config_deepseek_v4(
_mock_vllm_config(), groups, 100 * 1024 * 1024
)
backing = torch.zeros(tensors[0].size, dtype=torch.uint8)
views = []
for t in tensors:
page_size = page_sizes[t.shared_by[0]]
v = torch.as_strided(
backing,
size=(num_blocks, page_size),
stride=(t.block_stride, 1),
storage_offset=t.offset,
)
views.append(v)
for i, v in enumerate(views):
v.fill_(i + 1)
for i, v in enumerate(views):
assert (v == i + 1).all(), f"View {i} was corrupted"
def test_hma_attention_groups_keep_default_backing(self, monkeypatch):
monkeypatch.setattr(envs, "VLLM_USE_PACKED_HMA_KV_CACHE", False, raising=False)
full = _make_full_spec()
sw = _make_sw_spec()
page_size = full.page_size_bytes
groups = [
KVCacheGroupSpec(["full.0", "full.1"], full),
KVCacheGroupSpec(["sw.0", "sw.2"], sw),
KVCacheGroupSpec(["sw.1", "sw.3"], sw),
]
config = get_kv_cache_config_from_groups(
_mock_vllm_config(), groups, available_memory=page_size * 2 * 32
)
assert config.num_blocks == 32
assert sum(t.size for t in config.kv_cache_tensors) == page_size * 2 * 32
assert config.kv_cache_tensors == [
KVCacheTensor(size=page_size * 32, shared_by=["full.0", "sw.0", "sw.1"]),
KVCacheTensor(size=page_size * 32, shared_by=["full.1", "sw.2", "sw.3"]),
]
def test_hma_attention_groups_use_packed_backing_with_flag(self, monkeypatch):
monkeypatch.setattr(envs, "VLLM_USE_PACKED_HMA_KV_CACHE", True, raising=False)
full = _make_full_spec()
sw = _make_sw_spec()
page_size = full.page_size_bytes
groups = [
KVCacheGroupSpec(["full.0", "full.1"], full),
KVCacheGroupSpec(["sw.0", "sw.2"], sw),
KVCacheGroupSpec(["sw.1", "sw.3"], sw),
]
config = get_kv_cache_config_from_groups(
_mock_vllm_config(), groups, available_memory=page_size * 2 * 32
)
assert config.num_blocks == 32
assert {t.size for t in config.kv_cache_tensors} == {page_size * 2 * 32}
assert config.kv_cache_tensors == [
KVCacheTensor(
size=page_size * 2 * 32,
shared_by=["full.0", "sw.0", "sw.1"],
offset=0,
block_stride=page_size * 2,
),
KVCacheTensor(
size=page_size * 2 * 32,
shared_by=["full.1", "sw.2", "sw.3"],
offset=page_size,
block_stride=page_size * 2,
),
]
def test_single_group_attention_keeps_unpacked_layout(self):
spec = _make_full_spec()
groups = [KVCacheGroupSpec(["full.0", "full.1"], spec)]
config = get_kv_cache_config_from_groups(
_mock_vllm_config(), groups, available_memory=spec.page_size_bytes * 2 * 32
)
assert sum(t.size for t in config.kv_cache_tensors) == (
spec.page_size_bytes * 2 * 32
)
assert [t.block_stride for t in config.kv_cache_tensors] == [0, 0]
if __name__ == "__main__":
pytest.main([__file__, "-v"])
+37
View File
@@ -144,6 +144,43 @@ def test_async_scheduling_pp_allows_rescheduling_with_output_placeholders():
assert req.request_id in output.num_scheduled_tokens
def test_cached_request_data_resumed_all_token_ids_mrv1_only():
"""all_token_ids carries a resumed request's token ids to the connector
for the V1 model runner, but is skipped entirely for the V2 model runner.
"""
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
scheduler = create_scheduler()
(req,) = create_requests(num_requests=1, num_tokens=8)
req.append_output_token_ids([101, 102, 103])
# A resumed request was not scheduled in the previous step.
assert req.request_id not in scheduler.prev_step_scheduled_req_ids
empty_blocks = KVCacheBlocks(blocks=((),))
def make_cached():
return scheduler._make_cached_request_data(
running_reqs=[],
resumed_reqs=[req],
num_scheduled_tokens={req.request_id: 1},
spec_decode_tokens={},
req_to_new_blocks={req.request_id: empty_blocks},
)
# V1 model runner: the full token id list is propagated.
assert not scheduler.use_v2_model_runner
cached = make_cached()
assert req.request_id in cached.resumed_req_ids
assert cached.all_token_ids[req.request_id] == list(req.all_token_ids)
# V2 model runner: all_token_ids is skipped entirely.
scheduler.use_v2_model_runner = True
cached = make_cached()
assert req.request_id in cached.resumed_req_ids
assert cached.all_token_ids == {}
def test_schedule_partial_requests():
"""Test scheduling behavior with partial requests.
@@ -4,9 +4,16 @@
import pytest
from vllm import LLM, SamplingParams
from vllm.platforms import current_platform
from ....utils import create_new_process_for_each_test
if current_platform.is_rocm():
pytest.skip(
"Cascade attention backends FLASH_ATTN and FLASHINFER are notsupported on ROCm",
allow_module_level=True,
)
@create_new_process_for_each_test()
@pytest.mark.parametrize("attn_backend", ["FLASH_ATTN", "FLASHINFER"])
@@ -7,7 +7,10 @@ from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.coordinator imp
ExternalCachedBlockPool,
MooncakeStoreCoordinator,
)
from vllm.v1.core.kv_cache_utils import BlockHash, BlockHashListWithBlockSize
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import (
chunk_hashes_for_block_size,
)
from vllm.v1.core.kv_cache_utils import BlockHash
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheGroupSpec,
@@ -182,7 +185,7 @@ def test_coordinator_group_block_size_double_hash():
]
coord = _make_coord(groups, hash_block_size=16)
hs = _hashes(4)
big_hashes = list(BlockHashListWithBlockSize(hs, 16, 32))
big_hashes = list(chunk_hashes_for_block_size(hs, 16, 32))
exists = {(0, bytes(h)) for h in hs}
exists |= {(1, bytes(bh)) for bh in big_hashes}
cmap = ExternalCachedBlockPool(exists)
@@ -323,8 +323,8 @@ def test_recv_skips_swa_blocks_before_window():
def test_chunked_token_database_hash_block_size_smaller_than_block_size():
"""DSv4-style: hash_block_size=4, group block_size=16 — process_tokens
must merge every 4 fine hashes into one chunk hash via
BlockHashListWithBlockSize."""
keys each 16-token chunk by its last fine hash, keeping the Mooncake key
at one digest instead of concatenating all 4 fine hashes."""
md = KeyMetadata("m", 0, 0, 0, 0, group_id=3)
db = ChunkedTokenDatabase(md, block_size=16, hash_block_size=4)
db.set_kv_caches_base_addr([0])
@@ -335,8 +335,7 @@ def test_chunked_token_database_hash_block_size_smaller_than_block_size():
assert len(out) == 2
assert out[0][0] == 0 and out[0][1] == 16
assert out[1][0] == 16 and out[1][1] == 32
# Each chunk's hash is the concatenation of 4 fine hashes.
expected0 = b"".join(fine_hashes[0:4]).hex()
expected1 = b"".join(fine_hashes[4:8]).hex()
assert out[0][2].chunk_hash == expected0
assert out[1][2].chunk_hash == expected1
# Each chunk's hash is its last (4th) fine hash, which already chains the
# prior three.
assert out[0][2].chunk_hash == fine_hashes[3].hex()
assert out[1][2].chunk_hash == fine_hashes[7].hex()
@@ -23,6 +23,7 @@ from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import (
worker as mooncake_store_worker,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import (
BlobBlockHashes,
ChunkedTokenDatabase,
KeyMetadata,
LoadSpec,
@@ -32,6 +33,7 @@ from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import (
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.metrics import (
MooncakeStoreConnectorStats,
)
from vllm.v1.core.kv_cache_utils import BlockHash
def _default_send_coord() -> mooncake_store_worker.MooncakeStoreCoordinator:
@@ -1179,9 +1181,9 @@ def test_store_sending_thread_kv_events_use_group_chunk_metadata():
assert full_event.group_idx == 0
assert full_event.block_size == 32
assert full_event.token_ids == list(range(32))
assert full_event.block_hashes == [
maybe_convert_block_hash(BlockHash(b"".join(hs)))
]
# block_size=32 over hash_block_size=8 (scale 4): the chunk is keyed by its
# last sub-hash, not the concatenation of all four.
assert full_event.block_hashes == [maybe_convert_block_hash(BlockHash(hs[3]))]
assert swa_event.group_idx == 1
assert swa_event.block_size == 8
@@ -1749,3 +1751,33 @@ def test_store_worker_close_swallows_store_errors():
worker.close()
assert worker.store is None
def test_blob_block_hashes_wire_roundtrip():
"""The lookup wire format sends a ``hash_len`` frame plus the raw hashes
concatenated back-to-back; the server rebuilds them through a zero-copy
``BlobBlockHashes`` view over the frame buffer."""
hashes = [BlockHash(bytes([i]) * 16) for i in range(5)]
hash_len = len(hashes[0])
# Client side (LookupKeyClient._lookup): flat payload frame.
blob = b"".join(hashes)
# Server side (LookupKeyServer): view over the frame buffer (a memoryview),
# never materializing the full hash list upfront.
view = BlobBlockHashes(memoryview(blob), hash_len)
assert len(view) == 5
assert list(view) == hashes # default Sequence iter terminates via IndexError
assert [bytes(h) for h in view] == hashes
assert bytes(view[-1]) == hashes[-1]
assert [bytes(h) for h in view[1:3]] == hashes[1:3]
with pytest.raises(IndexError):
_ = view[5]
def test_blob_block_hashes_empty():
"""Empty lookups send hash_len=0 and an empty payload."""
view = BlobBlockHashes(memoryview(b""), 0)
assert len(view) == 0
assert list(view) == []
+41 -5
View File
@@ -14,12 +14,13 @@ from vllm.v1.kv_offload.base import (
ReqContext,
make_offload_key,
)
from vllm.v1.kv_offload.cpu.common import CPULoadStoreSpec
from vllm.v1.kv_offload.cpu.common import (
CPULoadStoreSpec,
CPUOffloadingMetrics,
)
from vllm.v1.kv_offload.cpu.manager import CPUOffloadingManager
from vllm.v1.kv_offload.cpu.policies.arc import ARCCachePolicy
STORES_SKIPPED = "vllm:kv_offload_stores_skipped"
def make_req_context(
req_id: str = "", kv_transfer_params: dict | None = None
@@ -181,10 +182,45 @@ def test_filter_reused_manager_reports_stores_skipped_counter():
)
stats = manager.get_stats()
assert stats is not None
assert stats.reduce()[STORES_SKIPPED] == 3
assert stats.reduce()[CPUOffloadingMetrics.STORES_SKIPPED] == 3
stats = manager.get_stats()
assert stats is not None
assert stats.reduce()[STORES_SKIPPED] == 0
assert stats.reduce()[CPUOffloadingMetrics.STORES_SKIPPED] == 0
def test_cpu_manager_reports_cache_usage_gauge():
def check_usage_stats(manager: CPUOffloadingManager, value: float):
stats = manager.get_stats()
assert stats is not None
assert stats.reduce()[
CPUOffloadingMetrics.CPU_CACHE_USAGE_PERC
] == pytest.approx(value)
# Zero-capacity manager always reports 0.0
manager = make_cpu_manager(num_blocks=0)
check_usage_stats(manager, 0.0)
# Empty manager (4 blocks, none allocated): usage = 0.0
manager = make_cpu_manager(num_blocks=4)
check_usage_stats(manager, 0.0)
# After allocating 2 of 4 blocks: usage = 0.5
manager.prepare_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
check_usage_stats(manager, 0.5)
# After filling all 4 blocks: usage = 1.0
manager.prepare_store(to_keys([3, 4]), _EMPTY_REQ_CTX)
check_usage_stats(manager, 1.0)
# After completing store, the blocks becomes evictable as it is not actively used
# and usage drops.
manager.complete_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
check_usage_stats(manager, 0.5)
# After completing store, the blocks becomes evictable as it is not actively used
# and usage drops.
manager.complete_store(to_keys([3, 4]), _EMPTY_REQ_CTX)
check_usage_stats(manager, 0.0)
def test_cpu_manager():
@@ -145,7 +145,6 @@ def _generate_fake_sampling_metadata(
vllm_config.scheduler_config.max_num_seqs,
num_spec,
device,
PIN_MEMORY_AVAILABLE,
)
fake_sampling_metadata = SamplingMetadata(
temperature=torch.full((batch_size,), 0.0),
@@ -880,7 +879,6 @@ def test_maybe_create_thinking_budget_holder_without_reasoning():
cfg.scheduler_config.max_num_seqs,
0,
torch.device("cpu"),
False,
)
is None
)
@@ -6,6 +6,7 @@ from __future__ import annotations
from dataclasses import dataclass
import pytest
import torch
from vllm import SamplingParams
@@ -1528,3 +1529,232 @@ def test_reset_pending_loads() -> None:
# All GPU blocks free
num_used = gpu_pool.num_gpu_blocks - gpu_pool.get_num_free_blocks()
assert num_used == 1, f"Expected only null block in use, got {num_used}"
def _make_cp_vllm_config(
dcp_world_size: int = 1,
pcp_world_size: int = 1,
) -> VllmConfig:
"""VllmConfig with context-parallel sizes set for scheduler-only tests."""
cfg = _make_vllm_config()
cfg.parallel_config.decode_context_parallel_size = dcp_world_size
cfg.parallel_config.prefill_context_parallel_size = pcp_world_size
return cfg
def _make_cp_scheduler(
*,
dcp_world_size: int = 1,
pcp_world_size: int = 1,
num_cpu_blocks: int = 8,
num_gpu_blocks: int = 16,
lazy: bool = False,
) -> SchedulerFixture:
"""Build a SimpleCPUOffloadScheduler with CP-scaled virtual block size."""
cp_world_size = dcp_world_size * pcp_world_size
virtual_block_size = BLOCK_SIZE * cp_world_size
kv_cache_config = _make_kv_cache_config(num_gpu_blocks)
vllm_config = _make_cp_vllm_config(dcp_world_size, pcp_world_size)
cpu_capacity_bytes = _BYTES_PER_BLOCK * num_cpu_blocks
sched = SimpleCPUOffloadScheduler(
vllm_config=vllm_config,
kv_cache_config=kv_cache_config,
cpu_capacity_bytes=cpu_capacity_bytes,
scheduler_block_size=virtual_block_size,
hash_block_size=virtual_block_size,
lazy_offload=lazy,
)
gpu_block_pool = BlockPool(
num_gpu_blocks=num_gpu_blocks,
enable_caching=True,
hash_block_size=virtual_block_size,
)
sched.bind_gpu_block_pool(gpu_block_pool)
return SchedulerFixture(
scheduler=sched,
gpu_block_pool=gpu_block_pool,
vllm_config=vllm_config,
kv_cache_config=kv_cache_config,
)
def _make_cp_request(
num_blocks: int,
virtual_block_size: int,
request_id: str | None = None,
) -> Request:
"""Create a request whose block hashes are computed at the virtual
(CP-scaled) block size, matching what the real scheduler does.
"""
global _req_counter
_req_counter += 1
if request_id is None:
request_id = f"req-cp-{_req_counter}"
num_tokens = num_blocks * virtual_block_size + 1
start = _req_counter * 10000
prompt_token_ids = list(range(start, start + num_tokens))
sampling_params = SamplingParams(max_tokens=1)
return Request(
request_id=request_id,
prompt_token_ids=prompt_token_ids,
sampling_params=sampling_params,
pooling_params=None,
mm_features=None,
block_hasher=get_request_block_hasher(virtual_block_size, sha256),
)
def _allocate_cp_gpu_blocks(
gpu_block_pool: BlockPool,
request: Request,
num_blocks: int,
virtual_block_size: int,
group_id: int = 0,
) -> list:
"""Allocate GPU blocks and cache them using the CP-scaled block size."""
blocks = gpu_block_pool.get_new_blocks(num_blocks)
num_full = min(num_blocks, len(request.block_hashes))
if num_full > 0:
gpu_block_pool.cache_full_blocks(
request=request,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=num_full,
block_size=virtual_block_size,
kv_cache_group_id=group_id,
)
return blocks
# ---------------------------------------------------------------------------
# Test 15: CP block size scaling is correct
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"dcp_world_size, pcp_world_size",
[
(2, 1), # DCP only
(1, 2), # PCP only
(2, 2), # DCP + PCP
],
)
def test_cp_block_size_scaling(dcp_world_size: int, pcp_world_size: int) -> None:
"""Verify that the scheduler's block_size and cp_world_size are correctly
scaled when context parallelism is enabled."""
fix = _make_cp_scheduler(
dcp_world_size=dcp_world_size, pcp_world_size=pcp_world_size
)
sched = fix.scheduler
expected_cp = dcp_world_size * pcp_world_size
assert sched.cp_world_size == expected_cp
assert sched.block_size == BLOCK_SIZE * expected_cp
# ---------------------------------------------------------------------------
# Test 16: CP eager store-and-load roundtrip
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"dcp_world_size, pcp_world_size",
[
(2, 1),
(1, 2),
],
)
def test_cp_eager_store_and_load_roundtrip(
dcp_world_size: int, pcp_world_size: int
) -> None:
"""With CP enabled, store blocks to CPU and reload them for a new request
with matching tokens. Verifies that hash matching and transfer-pair
construction work with the virtual block size."""
fix = _make_cp_scheduler(
dcp_world_size=dcp_world_size,
pcp_world_size=pcp_world_size,
num_cpu_blocks=8,
num_gpu_blocks=16,
lazy=False,
)
sched = fix.scheduler
cp = dcp_world_size * pcp_world_size
vbs = BLOCK_SIZE * cp
num_blocks = 2
req = _make_cp_request(num_blocks, vbs)
# Allocate GPU blocks and register hashes
gpu_blocks = _allocate_cp_gpu_blocks(fix.gpu_block_pool, req, num_blocks, vbs)
kv_blocks = KVCacheBlocks(blocks=(gpu_blocks,))
req.num_computed_tokens = num_blocks * vbs
sched.update_state_after_alloc(req, kv_blocks, num_external_tokens=0)
block_ids = kv_blocks.get_block_ids()
sched_out = make_scheduler_output(
{req.request_id: num_blocks * vbs},
new_reqs={req.request_id: block_ids},
)
meta = sched.build_connector_meta(sched_out)
assert meta.store_event >= 0, "Expected a store event"
assert len(meta.store_gpu_blocks) == num_blocks
assert len(meta.store_cpu_blocks) == num_blocks
simulate_store_completion(sched, meta.store_event)
# New request with same tokens — should get a full CPU cache hit.
req2 = Request(
request_id="req-cp-load",
prompt_token_ids=req.prompt_token_ids,
sampling_params=req.sampling_params,
pooling_params=None,
mm_features=None,
block_hasher=req._block_hasher,
)
hit_tokens, is_async = sched.get_num_new_matched_tokens(req2, num_computed_tokens=0)
assert hit_tokens == num_blocks * vbs
assert is_async is True
# Allocate fresh GPU blocks for the load.
gpu_blocks2 = fix.gpu_block_pool.get_new_blocks(num_blocks)
kv_blocks2 = KVCacheBlocks(blocks=(gpu_blocks2,))
sched.update_state_after_alloc(req2, kv_blocks2, num_external_tokens=hit_tokens)
sched_out2 = make_scheduler_output(
{req2.request_id: 1},
new_reqs={req2.request_id: kv_blocks2.get_block_ids()},
)
meta2 = sched.build_connector_meta(sched_out2)
assert meta2.load_event >= 0, "Expected a load event"
assert len(meta2.load_gpu_blocks) == num_blocks
assert len(meta2.load_cpu_blocks) == num_blocks
# ---------------------------------------------------------------------------
# Test 17: CP lazy target blocks are scaled correctly
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("cp_world_size", [1, 2, 4])
def test_cp_lazy_target_blocks_scaling(cp_world_size: int) -> None:
"""_estimate_lazy_target_blocks returns fewer blocks when cp_world_size > 1
because each virtual block covers more tokens."""
kv_cache_config = _make_kv_cache_config(num_blocks=16)
max_batched = 64
target_base = SimpleCPUOffloadScheduler._estimate_lazy_target_blocks(
kv_cache_config, max_batched, cp_world_size=1
)
target_cp = SimpleCPUOffloadScheduler._estimate_lazy_target_blocks(
kv_cache_config, max_batched, cp_world_size=cp_world_size
)
if cp_world_size == 1:
assert target_cp == target_base
else:
assert target_cp < target_base, (
f"cp_world_size={cp_world_size}: target_cp={target_cp} should be "
f"less than target_base={target_base}"
)
@@ -35,7 +35,6 @@ def mock_model_runner_with_input_batch():
max_model_len=1024,
max_num_batched_tokens=1024,
device="cpu",
pin_memory=False,
vocab_size=32000,
block_sizes=[16],
kernel_block_sizes=[16],
-6
View File
@@ -10,7 +10,6 @@ import torch
from vllm.platforms import current_platform
from vllm.sampling_params import SamplingParams
from vllm.utils.platform_utils import is_pin_memory_available
from vllm.utils.torch_utils import make_tensor_with_pad
from vllm.v1.pool.metadata import PoolingMetadata
from vllm.v1.sample.logits_processor import LogitsProcessors
@@ -236,7 +235,6 @@ def test_sampling_metadata_in_input_batch(device: str, batch_size: int):
max_model_len=1024,
max_num_batched_tokens=1024,
device=torch.device(device),
pin_memory=is_pin_memory_available(),
vocab_size=1024,
block_sizes=[1],
kernel_block_sizes=[1],
@@ -331,7 +329,6 @@ def test_swap_states_in_input_batch(device: str, batch_size: int, swap_list: lis
max_model_len=1024,
max_num_batched_tokens=1024,
device=torch.device(device),
pin_memory=is_pin_memory_available(),
vocab_size=1024,
block_sizes=[1],
kernel_block_sizes=[1],
@@ -341,7 +338,6 @@ def test_swap_states_in_input_batch(device: str, batch_size: int, swap_list: lis
max_model_len=1024,
max_num_batched_tokens=1024,
device=torch.device(device),
pin_memory=is_pin_memory_available(),
vocab_size=1024,
block_sizes=[1],
kernel_block_sizes=[1],
@@ -410,7 +406,6 @@ def test_pooling_prompt_lens_not_aliased(device: str):
max_model_len=MAX_PROMPT_SIZE + NUM_OUTPUT_TOKENS,
max_num_batched_tokens=batch_size * (MAX_PROMPT_SIZE + NUM_OUTPUT_TOKENS),
device=torch.device(device),
pin_memory=is_pin_memory_available(),
vocab_size=VOCAB_SIZE,
block_sizes=[16],
kernel_block_sizes=[16],
@@ -459,7 +454,6 @@ def test_pooling_metadata_token_id_buffers(
max_model_len=MAX_PROMPT_SIZE + NUM_OUTPUT_TOKENS,
max_num_batched_tokens=MAX_PROMPT_SIZE + NUM_OUTPUT_TOKENS,
device=torch.device("cpu"),
pin_memory=False,
vocab_size=VOCAB_SIZE,
block_sizes=[16],
kernel_block_sizes=[16],
-4
View File
@@ -85,7 +85,6 @@ def initialize_kv_cache(runner: GPUModelRunner):
max_model_len=runner.max_model_len,
max_num_batched_tokens=runner.max_num_tokens,
device=runner.device,
pin_memory=runner.pin_memory,
vocab_size=runner.model_config.get_vocab_size(),
block_sizes=[kv_cache_config.kv_cache_groups[0].kv_cache_spec.block_size],
kernel_block_sizes=[
@@ -1405,7 +1404,6 @@ def test_input_batch_with_kernel_block_sizes():
max_model_len = 512
max_num_batched_tokens = 512
device = torch.device(DEVICE_TYPE)
pin_memory = False
vocab_size = 50272
# Test with different kernel block sizes
@@ -1417,7 +1415,6 @@ def test_input_batch_with_kernel_block_sizes():
max_model_len=max_model_len,
max_num_batched_tokens=max_num_batched_tokens,
device=device,
pin_memory=pin_memory,
vocab_size=vocab_size,
block_sizes=block_sizes,
kernel_block_sizes=kernel_block_sizes,
@@ -1478,7 +1475,6 @@ def test_hybrid_cache_integration(default_vllm_config, dist_init):
max_model_len=runner.max_model_len,
max_num_batched_tokens=runner.max_num_tokens,
device=runner.device,
pin_memory=runner.pin_memory,
vocab_size=runner.model_config.get_vocab_size(),
block_sizes=[kv_cache_config.kv_cache_groups[0].kv_cache_spec.block_size],
kernel_block_sizes=[16],
+1 -1
View File
@@ -991,7 +991,7 @@ class VllmBackend:
},
payload_fn=lambda: json.dumps(
{
"model": self.vllm_config.model_config.model,
"model": getattr(self.vllm_config.model_config, "model", "unknown"),
"prefix": self.prefix,
"mode": str(cc.mode),
"backend": cc.backend,
+9
View File
@@ -302,6 +302,14 @@ class ParallelConfig:
Each entry must use `numactl --physcpubind` CPU-list syntax, for example
`"0-3"` or `"0,2,4-7"`.
"""
assigned_physical_gpu_ids: list[int] | None = None
"""Mapping from vLLM-local logical GPU IDs to physical GPU IDs.
For example, ``[2, 3]`` means logical GPU 0 maps to physical GPU 2,
and logical GPU 1 maps to physical GPU 3. Physical IDs are used only
at platform/topology boundaries such as NVML, NIC affinity, P2P
checks, and final CUDA device selection when needed. When None,
logical IDs map to visible device IDs in order."""
distributed_timeout_seconds: int | None = None
"""Timeout in seconds for distributed operations (e.g., init_process_group).
@@ -772,6 +780,7 @@ class ParallelConfig:
"numa_bind",
"numa_bind_nodes",
"numa_bind_cpus",
"assigned_physical_gpu_ids",
}
from vllm.config.utils import get_hash_factors, hash_factors
+2 -2
View File
@@ -18,8 +18,8 @@ import torch
from vllm.device_allocator import AllocationData, HandleType
from vllm.logger import init_logger
from vllm.utils.platform_utils import is_pin_memory_available
from vllm.utils.system_utils import find_loaded_library
from vllm.utils.torch_utils import PIN_MEMORY
logger = init_logger(__name__)
@@ -196,7 +196,7 @@ class CuMemAllocator:
size_in_bytes,
dtype=torch.uint8,
device="cpu",
pin_memory=is_pin_memory_available(),
pin_memory=PIN_MEMORY,
)
cpu_ptr = cpu_backup_tensor.data_ptr()
libcudart.cudaMemcpy(cpu_ptr, ptr, size_in_bytes)
+2 -2
View File
@@ -11,7 +11,7 @@ import torch
from vllm.device_allocator import AllocationData, HandleType
from vllm.logger import init_logger
from vllm.utils.platform_utils import is_pin_memory_available
from vllm.utils.torch_utils import PIN_MEMORY
logger = init_logger(__name__)
@@ -188,7 +188,7 @@ class XpuMemAllocator:
size_in_bytes,
dtype=torch.uint8,
device="cpu",
pin_memory=is_pin_memory_available(),
pin_memory=PIN_MEMORY,
)
cpu_ptr = cpu_backup_tensor.data_ptr()
_xpu_memcpy_sync(
@@ -704,7 +704,14 @@ class FlashInferNVLinkOneSidedManager(All2AllManagerBase):
self.num_experts = num_experts
self.cleanup()
gpus_per_node = torch.accelerator.device_count()
from vllm.platforms.interface import get_assigned_physical_gpu_ids
assigned_physical_gpu_ids = get_assigned_physical_gpu_ids()
gpus_per_node = (
len(assigned_physical_gpu_ids)
if assigned_physical_gpu_ids is not None
else torch.accelerator.device_count()
)
logger.debug(
"Making One-sided NVLink mapping: rank=%d, world size=%d",
self.rank,
@@ -320,13 +320,21 @@ def gpu_p2p_access_check(src: int, tgt: int) -> bool:
is_distributed = dist.is_initialized()
num_dev = current_platform.device_count()
cuda_visible_devices = envs.CUDA_VISIBLE_DEVICES
if cuda_visible_devices is None:
cuda_visible_devices = ",".join(str(i) for i in range(num_dev))
from vllm.platforms.interface import get_assigned_physical_gpu_ids
assigned_physical_gpu_ids = get_assigned_physical_gpu_ids()
if assigned_physical_gpu_ids is not None:
# Key by the ordered list: the cache stores directed local-index
# pairs, so permutations of the same set are distinct mappings.
cache_key = ",".join(str(i) for i in assigned_physical_gpu_ids)
num_dev = len(assigned_physical_gpu_ids)
else:
num_dev = current_platform.device_count()
cuda_visible_devices = envs.CUDA_VISIBLE_DEVICES
cache_key = cuda_visible_devices or ",".join(str(i) for i in range(num_dev))
path = os.path.join(
envs.VLLM_CACHE_ROOT, f"gpu_p2p_access_cache_for_{cuda_visible_devices}.json"
envs.VLLM_CACHE_ROOT, f"gpu_p2p_access_cache_for_{cache_key}.json"
)
os.makedirs(os.path.dirname(path), exist_ok=True)
from vllm.distributed.parallel_state import get_world_group
@@ -338,7 +346,15 @@ def gpu_p2p_access_check(src: int, tgt: int) -> bool:
# enter this block to calculate the cache
logger.info("generating GPU P2P access cache in %s", path)
cache: dict[str, bool] = {}
ids = list(range(num_dev))
# The probe subprocesses inherit this process's device-control env
# var, so they must be given visible ordinals, not physical IDs.
if assigned_physical_gpu_ids is not None:
ids = [
current_platform.logical_device_id_to_visible_device_id(local)
for local in range(num_dev)
]
else:
ids = list(range(num_dev))
# batch of all pairs of GPUs
batch_src, batch_tgt = zip(*list(product(ids, ids)))
# NOTE: we use `subprocess` rather than `multiprocessing` here
@@ -368,8 +384,11 @@ def gpu_p2p_access_check(src: int, tgt: int) -> bool:
) from e
with open(output_file.name, "rb") as f:
result = pickle.load(f)
# Cache entries must be keyed by local indices (0..N-1) because
# gpu_p2p_access_check() is called with local ranks.
id_to_local = {device_id: local for local, device_id in enumerate(ids)}
for _i, _j, r in zip(batch_src, batch_tgt, result):
cache[f"{_i}->{_j}"] = r
cache[f"{id_to_local[_i]}->{id_to_local[_j]}"] = r
with open(path, "w") as f:
json.dump(cache, f, indent=4)
if is_distributed:
@@ -34,7 +34,12 @@ def _can_p2p(rank: int, world_size: int) -> bool:
continue
if envs.VLLM_SKIP_P2P_CHECK:
logger.debug("Skipping P2P check and trusting the driver's P2P report.")
return torch.cuda.can_device_access_peer(rank, i)
# can_device_access_peer takes visible device ordinals, while
# rank and i are logical local IDs.
return torch.cuda.can_device_access_peer(
current_platform.logical_device_id_to_visible_device_id(rank),
current_platform.logical_device_id_to_visible_device_id(i),
)
if not gpu_p2p_access_check(rank, i):
return False
return True
@@ -126,13 +131,10 @@ class CustomAllreduce:
CUSTOM_ALL_REDUCE_MAX_SIZES[device_capability_str][world_size],
max_size,
)
cuda_visible_devices = envs.CUDA_VISIBLE_DEVICES
if cuda_visible_devices:
device_ids = list(map(int, cuda_visible_devices.split(",")))
else:
device_ids = list(range(current_platform.device_count()))
physical_device_id = device_ids[device.index]
# device.index is a visible ordinal, not a logical local ID.
physical_device_id = current_platform.visible_device_id_to_physical_device_id(
device.index
)
tensor = torch.tensor([physical_device_id], dtype=torch.int, device="cpu")
gather_list = [
torch.tensor([0], dtype=torch.int, device="cpu") for _ in range(world_size)
@@ -129,12 +129,10 @@ class QuickAllReduce:
assert isinstance(device, torch.device)
self.device = device
cuda_visible_devices = envs.CUDA_VISIBLE_DEVICES
if cuda_visible_devices:
device_ids = list(map(int, cuda_visible_devices.split(",")))
else:
device_ids = list(range(current_platform.device_count()))
physical_device_id = device_ids[device.index]
# device.index is a visible ordinal, not a logical local ID.
physical_device_id = current_platform.visible_device_id_to_physical_device_id(
device.index
)
tensor = torch.tensor([physical_device_id], dtype=torch.int, device="cpu")
gather_list = [
torch.tensor([0], dtype=torch.int, device="cpu")
@@ -840,7 +840,13 @@ class MessageQueue:
The MessageQueue instance for the calling process,
and a list of handles (only non-empty for the reader process).
"""
local_size = current_platform.device_count()
from vllm.platforms.interface import get_assigned_physical_gpu_ids
assigned_physical_gpu_ids = get_assigned_physical_gpu_ids()
if assigned_physical_gpu_ids is not None:
local_size = len(assigned_physical_gpu_ids)
else:
local_size = current_platform.device_count()
rank = dist.get_rank()
same_node = rank // local_size == reader_rank // local_size
buffer_io = MessageQueue(
@@ -482,10 +482,11 @@ def _init_lmcache_engine(
)
# Change current device.
num_gpus = torch.accelerator.device_count()
local_rank = parallel_config.rank % num_gpus
torch.accelerator.set_device_index(local_rank)
device = torch.device(f"cuda:{local_rank}")
from vllm.distributed.parallel_state import get_world_group
device_index = get_world_group().device_index
torch.accelerator.set_device_index(device_index)
device = torch.device(f"cuda:{device_index}")
metadata = LMCacheEngineMetadata(
model_config.model,
parallel_config.world_size,
@@ -2,13 +2,15 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""External-store cache-hit coordinator for MooncakeStoreConnector."""
from collections.abc import Sequence
from typing import cast
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import (
chunk_hashes_for_block_size,
)
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_utils import (
BlockHash,
BlockHashList,
BlockHashListWithBlockSize,
KVCacheBlock,
)
from vllm.v1.core.single_type_kv_cache_manager import (
@@ -120,7 +122,7 @@ class MooncakeStoreCoordinator:
def find_longest_cache_hit(
self,
block_hashes: list[BlockHash],
block_hashes: Sequence[BlockHash],
max_length: int,
cached_block_pool: ExternalCachedBlockPool,
*,
@@ -147,7 +149,7 @@ class MooncakeStoreCoordinator:
def load_mask(
self,
block_hashes: list[BlockHash],
block_hashes: Sequence[BlockHash],
token_len: int,
) -> tuple[list[bool], ...]:
"""Per-group load masks: ``mask[g][i]`` is True iff group ``g``'s
@@ -236,17 +238,15 @@ class MooncakeStoreCoordinator:
return tuple(masks)
def block_hashes_for_spec(
self, block_hashes: list[BlockHash], spec: KVCacheSpec
) -> BlockHashList:
if spec.block_size == self.hash_block_size:
return block_hashes
return BlockHashListWithBlockSize(
self, block_hashes: Sequence[BlockHash], spec: KVCacheSpec
) -> Sequence[BlockHash]:
return chunk_hashes_for_block_size(
block_hashes, self.hash_block_size, spec.block_size
)
def _find_hit_blocks(
self,
block_hashes: list[BlockHash],
block_hashes: Sequence[BlockHash],
max_length: int,
cached_block_pool: ExternalCachedBlockPool,
*,
@@ -264,7 +264,7 @@ class MooncakeStoreCoordinator:
spec, group_ids, manager_cls = self.attention_groups[0]
hashes = self.block_hashes_for_spec(block_hashes, spec)
hit_blocks = manager_cls.find_longest_cache_hit(
block_hashes=hashes,
block_hashes=hashes, # type: ignore[arg-type]
max_length=max_length,
kv_cache_group_ids=group_ids,
block_pool=cast(BlockPool, cached_block_pool),
@@ -304,7 +304,7 @@ class MooncakeStoreCoordinator:
_max_length = min(curr_hit_length + spec.block_size, max_length)
hashes = self.block_hashes_for_spec(block_hashes, spec)
hit_blocks = manager_cls.find_longest_cache_hit(
block_hashes=hashes,
block_hashes=hashes, # type: ignore[arg-type]
max_length=_max_length,
kv_cache_group_ids=group_ids,
block_pool=cast(BlockPool, cached_block_pool),
@@ -5,8 +5,9 @@
# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/).
"""Data classes for MooncakeStoreConnector."""
from collections.abc import Iterable
from collections.abc import Iterable, Sequence
from dataclasses import dataclass
from typing import cast
import torch
@@ -23,6 +24,77 @@ from vllm.v1.core.kv_cache_utils import (
logger = init_logger(__name__)
class BlobBlockHashes(Sequence[BlockHash]):
"""Lazy view over a flat buffer of fixed-size block hashes to avoid the overhead
of materializing all hashes upfront.
"""
def __init__(self, blob: memoryview, hash_len: int):
self._blob = blob
self._hash_len = hash_len
self._n = len(blob) // hash_len if hash_len else 0
def __len__(self) -> int:
return self._n
def __getitem__(self, idx):
if isinstance(idx, slice):
return [self[i] for i in range(*idx.indices(self._n))]
if idx < 0:
idx += self._n
if not 0 <= idx < self._n:
raise IndexError(idx)
off = idx * self._hash_len
return BlockHash(self._blob[off : off + self._hash_len])
class _CompactChunkHashList(BlockHashListWithBlockSize):
"""View that keys each ``block_size`` chunk by the last constituent
``hash_block_size`` hash instead of concatenating all of them.
The engine chains block hashes (each hash folds in the previous one), so the
final sub-block hash of a chunk already uniquely identifies the whole chunk
and its prefix. Using it keeps a Mooncake key at a single hash digest
regardless of the ``block_size`` / ``hash_block_size`` ratio, instead of
growing the key linearly with it (e.g. 64x for ``block_size=256``,
``hash_block_size=4``).
"""
def __init__(
self,
block_hashes: Sequence[BlockHash],
hash_block_size: int,
target_block_size: int,
):
# Accept any indexable sequence (e.g. the lazy ``BlobBlockHashes``), not
# just ``list``; the base only indexes/sizes it.
assert target_block_size % hash_block_size == 0
self.block_hashes = block_hashes # type: ignore[assignment]
self.scale_factor = target_block_size // hash_block_size
def _get_value_at(self, idx: int) -> BlockHash:
return self.block_hashes[idx * self.scale_factor + self.scale_factor - 1]
def chunk_hashes_for_block_size(
block_hashes: Sequence[BlockHash],
hash_block_size: int,
block_size: int,
) -> Sequence[BlockHash]:
"""Map ``hash_block_size``-granular block hashes to one compact hash per
``block_size`` chunk (the chunk's last sub-hash). Returns ``block_hashes``
unchanged when the two sizes are equal.
"""
if block_size == hash_block_size:
return block_hashes
# Structurally a Sequence[BlockHash] (indexable + sized); the base class
# just isn't declared as one.
return cast(
"Sequence[BlockHash]",
_CompactChunkHashList(block_hashes, hash_block_size, block_size),
)
@dataclass
class KeyMetadata:
"""Metadata for constructing pool keys."""
@@ -138,18 +210,15 @@ class ChunkedTokenDatabase:
Args:
token_len: Total number of tokens.
block_hashes: Block hashes computed at ``hash_block_size`` granularity.
When ``block_size > hash_block_size`` consecutive hashes are merged
up to the group's ``block_size`` via ``BlockHashListWithBlockSize``.
When ``block_size > hash_block_size`` each group's ``block_size`` chunk
is keyed by its last sub-hash via ``chunk_hashes_for_block_size``.
mask_num: Number of tokens to skip from the beginning.
"""
if not block_hashes:
return
if self.block_size == self.hash_block_size:
chunk_hashes: Iterable[BlockHash] = block_hashes
else:
chunk_hashes = BlockHashListWithBlockSize(
block_hashes, self.hash_block_size, self.block_size
)
chunk_hashes: Iterable[BlockHash] = chunk_hashes_for_block_size(
block_hashes, self.hash_block_size, self.block_size
)
for chunk_id, h in enumerate(chunk_hashes):
start_idx = chunk_id * self.block_size
if start_idx >= token_len:
@@ -11,7 +11,10 @@ Wire format (REQ/REP over IPC):
msg_type == LOOKUP_MSG:
frame 1: token_len (u32 big-endian, 4 bytes)
frame 2..n: msgpack-encoded list[str] of block-hash hex digests
frame 2: hash_len (u16 big-endian, 2 bytes) byte length of each
fixed-size block hash (0 when there are no hashes)
frame 3: raw block hashes concatenated back-to-back (each hash_len
bytes); the server splits on hash_len
Response: [hit_count: u32 big-endian, 4 bytes]
msg_type == RESET_MSG:
@@ -18,7 +18,7 @@ import socket
import threading
import time
from collections import defaultdict
from collections.abc import Callable
from collections.abc import Callable, Sequence
from concurrent.futures import Future, ThreadPoolExecutor
from dataclasses import dataclass
from typing import Any, Literal, TypeVar
@@ -45,6 +45,7 @@ from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.coordinator imp
MooncakeStoreCoordinator,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501
BlobBlockHashes,
ChunkedTokenDatabase,
KeyMetadata,
MooncakeStoreConnectorMetadata,
@@ -65,7 +66,6 @@ from vllm.v1.core.kv_cache_utils import (
resolve_kv_cache_block_sizes,
)
from vllm.v1.kv_cache_interface import KVCacheConfig, KVCacheGroupSpec
from vllm.v1.serial_utils import MsgpackDecoder, MsgpackEncoder
from .metrics import MooncakeStoreConnectorStats
@@ -1372,7 +1372,7 @@ class MooncakeStoreWorker:
return finished_sending
def lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
def lookup(self, token_len: int, block_hashes: Sequence[BlockHash]) -> int:
"""Check how many prefix tokens exist in the store.
Checks across all TP ranks and PP ranks.
@@ -1392,6 +1392,11 @@ class MooncakeStoreWorker:
group_hashes = self.coord.block_hashes_for_spec(
block_hashes, self._kv_cache_groups[g_idx].kv_cache_spec
)
metadata_templates = [
dataclasses.replace(db.metadata, tp_rank=tp, pp_rank=pp)
for tp in range(tp_count)
for pp in range(self.pp_size)
]
for chunk_id, h in enumerate(group_hashes):
start_idx = chunk_id * spec_block_size
if start_idx >= token_len:
@@ -1400,11 +1405,11 @@ class MooncakeStoreWorker:
chunk_id >= len(lookup_mask) or not lookup_mask[chunk_id]
):
continue
for tp in range(tp_count):
for pp in range(self.pp_size):
md = dataclasses.replace(db.metadata, tp_rank=tp, pp_rank=pp)
candidate_keys.append(PoolKey(md, h.hex()).to_string())
candidate_meta.append((g_idx, bytes(h)))
h_hex = h.hex()
h_bytes = bytes(h)
for md in metadata_templates:
candidate_keys.append(PoolKey(md, h_hex).to_string())
candidate_meta.append((g_idx, h_bytes))
if not candidate_keys:
return 0
@@ -1483,7 +1488,6 @@ class LookupKeyServer:
store_worker: MooncakeStoreWorker,
vllm_config: VllmConfig,
):
self.decoder = MsgpackDecoder()
self.ctx = zmq.Context() # type: ignore[attr-defined]
socket_path = get_zmq_rpc_path_lookup(vllm_config)
self._ipc_path = socket_path.removeprefix("ipc://")
@@ -1506,9 +1510,9 @@ class LookupKeyServer:
if msg_type == LOOKUP_MSG:
token_len = int.from_bytes(all_frames[1], byteorder="big")
hash_frames = all_frames[2:]
hashes_str = self.decoder.decode(hash_frames)
block_hashes = [BlockHash(bytes.fromhex(s)) for s in hashes_str]
hash_len = int.from_bytes(all_frames[2], byteorder="big")
blob = all_frames[3].buffer
block_hashes = BlobBlockHashes(blob, hash_len)
result = self.store_worker.lookup(token_len, block_hashes)
self.socket.send(result.to_bytes(4, "big"))
@@ -1557,7 +1561,6 @@ class LookupKeyClient:
"""
def __init__(self, vllm_config: VllmConfig):
self.encoder = MsgpackEncoder()
self.ctx = zmq.Context() # type: ignore[attr-defined]
socket_path = get_zmq_rpc_path_lookup(vllm_config)
self.socket = make_zmq_socket(
@@ -1574,14 +1577,16 @@ class LookupKeyClient:
self.futures: dict[str, Future[int]] = {}
def _lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
hash_strs = [h.hex() for h in block_hashes]
hash_frames = self.encoder.encode(hash_strs)
token_len_bytes = token_len.to_bytes(4, byteorder="big")
all_frames = [LOOKUP_MSG, token_len_bytes] + list(hash_frames)
hash_len = len(block_hashes[0]) if block_hashes else 0
all_frames = (
LOOKUP_MSG,
token_len.to_bytes(4, byteorder="big"),
hash_len.to_bytes(2, byteorder="big"),
b"".join(block_hashes),
)
self.socket.send_multipart(all_frames, copy=False)
resp = self.socket.recv()
result = int.from_bytes(resp, "big")
return result
return int.from_bytes(resp, "big")
def lookup(
self,
@@ -841,8 +841,106 @@ class NixlBaseConnectorWorker:
# Forwarding a real layer name rather than a synthetic key
self.register_kv_caches({first_layer: kv_cache})
def _register_packed_kv_cache(
self,
storage: torch.UntypedStorage,
) -> None:
"""Register a packed KV cache as a single NIXL region.
The packed allocation interleaves all layers per block, so each
block_stride-byte chunk is one logical block. We register 1
NIXL region and create 1 descriptor per block.
"""
self.transfer_topo = TransferTopology(
tp_rank=self.tp_rank,
tp_size=self.world_size,
block_size=self.block_size,
engine_id=self.engine_id,
is_mla=self.use_mla,
total_num_kv_heads=self.model_config.get_total_num_kv_heads(),
attn_backends=self.attn_backends,
tensor_shape=None,
is_mamba=self._has_mamba,
)
self.compat_hash = compute_nixl_compatibility_hash(
self.vllm_config,
self.backend_name,
self.transfer_topo.cross_layers_blocks,
)
total_size = storage.nbytes()
block_stride = total_size // self.num_blocks
base_addr = storage.data_ptr()
device_id = storage.device.index
assert device_id is not None
logger.info(
"Registering packed KV cache: total_size=%s, block_stride=%s, "
"num_blocks=%s, num_regions=1",
total_size,
block_stride,
self.num_blocks,
)
self.device_id = device_id
caches_data = [(base_addr, total_size, self.device_id, "")]
self.block_len_per_layer = [block_stride]
self.num_regions = 1
self.num_descs = self.num_blocks
self.kv_caches_base_addr[self.engine_id][self.tp_rank] = [base_addr]
descs = self.nixl_wrapper.get_reg_descs(caches_data, self.nixl_memory_type)
self.nixl_wrapper.register_memory(descs, backends=self.nixl_backends)
self._registered_descs.append(descs)
self.dst_num_blocks[self.engine_id] = self.num_blocks
self.src_xfer_handles_by_block_size[self.block_size], (self.src_blocks_data) = (
self.register_local_xfer_handler(self.block_size)
)
agent_metadata = NixlAgentMetadata(
engine_id=self.engine_id,
agent_metadata=self.nixl_wrapper.get_agent_metadata(),
device_id=self.device_id,
kv_caches_base_addr=(
self.kv_caches_base_addr[self.engine_id][self.tp_rank]
),
num_blocks=self.num_blocks,
block_lens=self.block_len_per_layer,
kv_cache_layout=self.kv_cache_layout,
block_size=self.block_size,
ssm_sizes=self._mamba_ssm_size,
attn_backend_name=self.backend_name,
physical_blocks_per_logical_kv_block=(
self._physical_blocks_per_logical_kv_block
),
)
assert self.compat_hash is not None
encoder = msgspec.msgpack.Encoder()
self.xfer_handshake_metadata = NixlHandshakePayload(
compatibility_hash=self.compat_hash,
agent_metadata_bytes=encoder.encode(agent_metadata),
)
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
"""Register the KV Cache data in nixl."""
# Detect packed allocation: all tensors are strided views into the
# same backing storage (different data_ptr but same storage).
# This happens with DSv4-style contiguous per-block packing.
if len(kv_caches) > 1 and not self._has_mamba:
storage = next(iter(kv_caches.values())).untyped_storage()
storage_ptrs = {
cache.untyped_storage().data_ptr() for cache in kv_caches.values()
}
data_ptrs = {cache.data_ptr() for cache in kv_caches.values()}
if len(storage_ptrs) == 1 and len(data_ptrs) > 1:
self._register_packed_kv_cache(storage)
self.device_kv_caches = kv_caches
return
self.transfer_topo = TransferTopology(
tp_rank=self.tp_rank,
tp_size=self.world_size,
@@ -50,7 +50,8 @@ class OffloadingConnectorWorker:
def register_kv_caches(
self, kv_caches: dict[str, torch.Tensor | list[torch.Tensor]]
):
num_blocks = self.spec.kv_cache_config.num_blocks
kv_cache_config = self.spec.kv_cache_config
num_blocks = kv_cache_config.num_blocks
# layer_name -> (num_blocks, page_size_bytes) tensor
tensors_per_block: dict[str, tuple[torch.Tensor, ...]] = {}
@@ -58,7 +59,7 @@ class OffloadingConnectorWorker:
unpadded_page_size_bytes: dict[str, int] = {}
# layer_name -> size of page in bytes
page_size_bytes: dict[str, int] = {}
for kv_cache_group in self.spec.kv_cache_config.kv_cache_groups:
for kv_cache_group in kv_cache_config.kv_cache_groups:
group_layer_names = kv_cache_group.layer_names
group_kv_cache_spec = kv_cache_group.kv_cache_spec
if isinstance(group_kv_cache_spec, UniformTypeKVCacheSpecs):
@@ -72,18 +73,22 @@ class OffloadingConnectorWorker:
if isinstance(layer_kv_cache_spec, AttentionSpec):
layer_kv_cache = kv_caches[layer_name]
assert isinstance(layer_kv_cache, torch.Tensor)
assert layer_kv_cache.storage_offset() == 0
storage = layer_kv_cache.untyped_storage()
page = layer_kv_cache_spec.page_size_bytes
elem_size = layer_kv_cache.element_size()
byte_offset = layer_kv_cache.storage_offset() * elem_size
block_stride_bytes = layer_kv_cache.stride(0) * elem_size
tensors_per_block[layer_name] = (
torch.tensor(
[],
dtype=torch.int8,
device=layer_kv_cache.device,
)
.set_(storage)
.view(num_blocks, page),
).set_(
layer_kv_cache.untyped_storage(),
byte_offset,
(num_blocks, page),
(block_stride_bytes, 1),
),
)
page_size_bytes[layer_name] = layer_kv_cache_spec.page_size_bytes
unpadded_page_size_bytes[layer_name] = (
@@ -118,9 +123,35 @@ class OffloadingConnectorWorker:
else:
raise NotImplementedError
packed_kv_cache_tensor = next(
(t for t in kv_cache_config.kv_cache_tensors if t.block_stride), None
)
is_dsv4 = all(
isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs)
for group in kv_cache_config.kv_cache_groups
)
if packed_kv_cache_tensor is not None and not is_dsv4:
(tensor,) = tensors_per_block[packed_kv_cache_tensor.shared_by[0]]
block_stride = tensor.stride(0)
packed_tensor = tensor.as_strided(
(num_blocks, block_stride),
(block_stride, 1),
storage_offset=0,
)
self._register_handlers(
CanonicalKVCaches(
[CanonicalKVCacheTensor(packed_tensor, block_stride)],
[
[CanonicalKVCacheRef(0, block_stride)]
for _ in kv_cache_config.kv_cache_groups
],
)
)
return
block_tensors: list[CanonicalKVCacheTensor] = []
block_data_refs: dict[str, list[CanonicalKVCacheRef]] = defaultdict(list)
for kv_cache_tensor in self.spec.kv_cache_config.kv_cache_tensors:
for kv_cache_tensor in kv_cache_config.kv_cache_tensors:
# Filter to layers that were actually processed above.
# _get_kv_cache_config_deepseek_v4 emits KVCacheTensor entries for
# every (tuple_idx, page_size) slot; slots where no group has a
@@ -162,7 +193,7 @@ class OffloadingConnectorWorker:
)
group_data_refs: list[list[CanonicalKVCacheRef]] = []
for kv_cache_group in self.spec.kv_cache_config.kv_cache_groups:
for kv_cache_group in kv_cache_config.kv_cache_groups:
group_refs: list[CanonicalKVCacheRef] = []
for layer_name in kv_cache_group.layer_names:
group_refs += block_data_refs[layer_name]
+24 -4
View File
@@ -392,6 +392,14 @@ class GroupCoordinator:
self.rank = torch.distributed.get_rank()
self.local_rank = local_rank
self.device_index: int
if _WORLD is not None:
self.device_index = _WORLD.device_index
else:
assert local_rank >= 0, (
"local_rank must be provided when creating the world group"
)
self.device_index = local_rank
self_device_group = None
self_cpu_group = None
@@ -442,11 +450,18 @@ class GroupCoordinator:
from vllm.platforms import current_platform
if current_platform.is_cuda_alike():
self.device = torch.device(f"cuda:{local_rank}")
visible_device_index = (
current_platform.logical_device_id_to_visible_device_id(
self.device_index
)
)
self.device = torch.device(f"cuda:{visible_device_index}")
elif current_platform.is_xpu():
self.device = torch.device(f"xpu:{local_rank}")
self.device = torch.device(f"xpu:{self.device_index}")
elif current_platform.is_out_of_tree():
self.device = torch.device(f"{current_platform.device_name}:{local_rank}")
self.device = torch.device(
f"{current_platform.device_name}:{self.device_index}"
)
else:
self.device = torch.device("cpu")
@@ -1438,7 +1453,12 @@ def _init_process_group_for_split_group(
"""
if torch.accelerator.is_available() and backend != "gloo":
init_backend = "cpu:gloo,cuda:nccl"
device_id: torch.device | None = torch.device(f"cuda:{local_rank}")
from vllm.platforms import current_platform
visible_device_index = current_platform.logical_device_id_to_visible_device_id(
local_rank
)
device_id: torch.device | None = torch.device(f"cuda:{visible_device_index}")
else:
init_backend = "gloo"
device_id = None
+19 -3
View File
@@ -86,6 +86,15 @@ class StatelessGroupCoordinator(GroupCoordinator):
self.rank = global_rank
self.local_rank = local_rank
from vllm.distributed.parallel_state import _WORLD
if _WORLD is not None:
self.device_index = _WORLD.device_index
else:
assert local_rank >= 0, (
"local_rank must be provided when creating the world group"
)
self.device_index = local_rank
self_device_group = None
self_cpu_group = None
@@ -152,11 +161,18 @@ class StatelessGroupCoordinator(GroupCoordinator):
self.tcp_store_group = self_tcp_store_group
if current_platform.is_cuda_alike():
self.device = torch.device(f"cuda:{local_rank}")
visible_device_index = (
current_platform.logical_device_id_to_visible_device_id(
self.device_index
)
)
self.device = torch.device(f"cuda:{visible_device_index}")
elif current_platform.is_xpu():
self.device = torch.device(f"xpu:{local_rank}")
self.device = torch.device(f"xpu:{self.device_index}")
elif current_platform.is_out_of_tree():
self.device = torch.device(f"{current_platform.device_name}:{local_rank}")
self.device = torch.device(
f"{current_platform.device_name}:{self.device_index}"
)
else:
self.device = torch.device("cpu")
+58
View File
@@ -6,6 +6,7 @@ import copy
import dataclasses
import functools
import json
import os
import sys
from collections.abc import Callable
from dataclasses import MISSING, asdict, dataclass, fields, is_dataclass
@@ -465,6 +466,7 @@ class EngineArgs:
numa_bind: bool = ParallelConfig.numa_bind
numa_bind_nodes: list[int] | None = ParallelConfig.numa_bind_nodes
numa_bind_cpus: list[str] | None = ParallelConfig.numa_bind_cpus
device_ids: list[int | str] | None = None
tensor_parallel_size: int = ParallelConfig.tensor_parallel_size
prefill_context_parallel_size: int = ParallelConfig.prefill_context_parallel_size
decode_context_parallel_size: int = ParallelConfig.decode_context_parallel_size
@@ -979,6 +981,20 @@ class EngineArgs:
parallel_group.add_argument(
"--numa-bind-cpus", **parallel_kwargs["numa_bind_cpus"]
)
parallel_group.add_argument(
"--device-ids",
type=lambda s: [
int(device_id) if device_id.isdigit() else device_id
for device_id in (part.strip() for part in s.split(","))
],
default=None,
help="Comma-separated physical GPU device IDs or UUIDs to use "
'(e.g. --device-ids "2,3,5,7"). Avoids setting '
"CUDA_VISIBLE_DEVICES, preserving full GPU topology "
"visibility for GPU-NIC affinity and DeepGEMM. "
"Note: has no effect with Ray executors; use Ray "
"placement groups for GPU selection instead.",
)
parallel_group.add_argument(
"--tensor-parallel-size", "-tp", **parallel_kwargs["tensor_parallel_size"]
)
@@ -1716,6 +1732,47 @@ class EngineArgs:
)
return SpeculativeConfig(**self.speculative_config)
def _resolve_device_ids(self) -> list[int] | None:
if not self.device_ids:
return None
if self.distributed_executor_backend == "ray":
logger.warning(
"--device-ids has no effect when using the Ray executor. "
"Use Ray placement groups for GPU selection instead."
)
ids = self.device_ids
if len(set(ids)) != len(ids):
raise ValueError(f"--device-ids must not contain duplicates: {ids}")
if all(isinstance(i, str) for i in ids):
return [
current_platform.device_control_id_to_physical_device_id(i)
for i in cast(list[str], ids)
]
if any(isinstance(i, str) for i in ids):
raise ValueError("--device-ids must not mix integer IDs and UUIDs")
int_ids = cast(list[int], ids)
# Compose with CUDA_VISIBLE_DEVICES: if CVD is set, treat
# --device-ids values as indices into the CVD-visible set.
cvd = getattr(
envs,
current_platform.device_control_env_var,
os.environ.get(current_platform.device_control_env_var),
)
if cvd:
cvd_ids = [
current_platform.device_control_id_to_physical_device_id(x)
for x in cvd.split(",")
]
for i in int_ids:
if i >= len(cvd_ids):
raise ValueError(
f"--device-ids index {i} is out of range for "
f"{current_platform.device_control_env_var}"
f"={cvd} ({len(cvd_ids)} devices visible)"
)
return [cvd_ids[i] for i in int_ids]
return int_ids
def create_diffusion_config(self) -> DiffusionConfig | None:
if self.diffusion_config is None:
return None
@@ -2029,6 +2086,7 @@ class EngineArgs:
cp_kv_cache_interleave_size=self.cp_kv_cache_interleave_size,
_api_process_count=self._api_process_count,
_api_process_rank=self._api_process_rank,
assigned_physical_gpu_ids=self._resolve_device_ids(),
numa_bind=self.numa_bind,
numa_bind_nodes=self.numa_bind_nodes,
numa_bind_cpus=self.numa_bind_cpus,
+56 -9
View File
@@ -43,6 +43,7 @@ from vllm.entrypoints.openai.engine.protocol import (
JsonSchemaResponseFormat,
ResponseFormat,
StreamOptions,
UsageInfo,
)
from vllm.entrypoints.openai.models.serving import OpenAIServingModels
from vllm.entrypoints.serve.utils.api_utils import sanitize_message
@@ -54,6 +55,49 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def _get_cached_tokens(usage: UsageInfo | None) -> int | None:
"""Extract cached token count from OpenAI UsageInfo."""
if usage is None or usage.prompt_tokens_details is None:
return None
return usage.prompt_tokens_details.cached_tokens
def _build_anthropic_usage(
prompt_tokens: int,
completion_tokens: int | None,
usage: UsageInfo | None,
) -> AnthropicUsage:
"""Build an AnthropicUsage from OpenAI-style token counts.
Anthropic defines ``total_input == input_tokens + cache_read +
cache_creation``. vLLM's ``prompt_tokens`` is the total, so
``input_tokens = prompt_tokens - cached_tokens``.
OpenAI usage only exposes ``cached_tokens`` (hits); there is no
cache-creation analog, so ``cache_creation_input_tokens`` is ``0``
when cache info is present. When cache info is absent (e.g.
``--enable-prompt-tokens-details`` off, or a streaming chunk that
hasn't carried it yet), cache fields are left **unset** so
``exclude_unset=True`` serialization omits them entirely.
``completion_tokens`` follows ``UsageInfo`` and may be ``None`` on
intermediate stream chunks; we coerce to ``0`` for the wire format.
"""
output_tokens = completion_tokens or 0
cached = _get_cached_tokens(usage)
if cached is not None:
return AnthropicUsage(
input_tokens=prompt_tokens - cached,
output_tokens=output_tokens,
cache_read_input_tokens=cached,
cache_creation_input_tokens=0,
)
return AnthropicUsage(
input_tokens=prompt_tokens,
output_tokens=output_tokens,
)
def wrap_data_with_event(data: str, event: str):
return f"event: {event}\ndata: {data}\n\n"
@@ -582,9 +626,10 @@ class AnthropicServingMessages(OpenAIServingChat):
id=generator.id,
content=[],
model=generator.model,
usage=AnthropicUsage(
input_tokens=generator.usage.prompt_tokens,
output_tokens=generator.usage.completion_tokens,
usage=_build_anthropic_usage(
generator.usage.prompt_tokens,
generator.usage.completion_tokens,
generator.usage,
),
kv_transfer_params=generator.kv_transfer_params,
)
@@ -765,11 +810,12 @@ class AnthropicServingMessages(OpenAIServingChat):
model=origin_chunk.model,
stop_reason=None,
stop_sequence=None,
usage=AnthropicUsage(
input_tokens=origin_chunk.usage.prompt_tokens
usage=_build_anthropic_usage(
origin_chunk.usage.prompt_tokens
if origin_chunk.usage
else 0,
output_tokens=0,
0,
origin_chunk.usage,
),
),
)
@@ -788,13 +834,14 @@ class AnthropicServingMessages(OpenAIServingChat):
chunk = AnthropicStreamEvent(
type="message_delta",
delta=AnthropicDelta(stop_reason=stop_reason),
usage=AnthropicUsage(
input_tokens=origin_chunk.usage.prompt_tokens
usage=_build_anthropic_usage(
origin_chunk.usage.prompt_tokens
if origin_chunk.usage
else 0,
output_tokens=origin_chunk.usage.completion_tokens
origin_chunk.usage.completion_tokens
if origin_chunk.usage
else 0,
origin_chunk.usage,
),
)
data = chunk.model_dump_json(exclude_unset=True)
-6
View File
@@ -898,12 +898,6 @@ class LLM(BeamSearchOfflineMixin, PoolingOfflineMixin, OfflineInferenceMixin):
def finish_weight_update(self) -> None:
"""Finish the current weight update."""
self.llm_engine.collective_rpc("finish_weight_update")
# Invalidate cached state computed with the old weights so it isn't
# reused for subsequent requests:
# - prefix cache: KV blocks computed with the old weights
# - encoder cache: multimodal embeddings keyed only by mm_hash
self.llm_engine.reset_prefix_cache()
self.llm_engine.reset_encoder_cache()
def __repr__(self) -> str:
"""Return a transformers-style hierarchical view of the model."""
+1 -1
View File
@@ -455,7 +455,7 @@ async def init_render_app_state(
enable_auto_tools=args.enable_auto_tool_choice,
exclude_tools_when_tool_choice_none=args.exclude_tools_when_tool_choice_none,
tool_parser=args.tool_call_parser,
reasoning_parser=args.structured_outputs_config.reasoning_parser,
reasoning_parser=args.reasoning_parser,
default_chat_template_kwargs=args.default_chat_template_kwargs,
log_error_stack=args.log_error_stack,
)
+19 -18
View File
@@ -23,12 +23,10 @@ import uvloop
from fastapi import FastAPI, Response
from vllm.logger import init_logger
from vllm.platforms import current_platform
from vllm.utils.system_utils import (
decorate_logs,
kill_process_tree,
set_process_title,
update_environment_variables,
)
logger = init_logger(__name__)
@@ -127,22 +125,29 @@ def _build_vllm_dp_server_args(
child_args.data_parallel_multi_port_external_lb = False
child_args.data_parallel_supervisor_port = None
child_args.api_server_count = 1
child_args.device_ids = _build_device_ids(args, local_rank)
return child_args
def _build_vllm_dp_server_env(
args: argparse.Namespace, local_rank: int
) -> dict[str, str]:
# set visible devices for the child process
def _build_device_ids(args: argparse.Namespace, local_rank: int) -> list[int | str]:
"""Build the --device-ids value for a DP child process.
The child resolves these against its own inherited device-control env
var (e.g. CUDA_VISIBLE_DEVICES), so integer IDs must stay env-relative
here rather than being translated to physical IDs.
"""
devices_per_rank = args.tensor_parallel_size * args.pipeline_parallel_size
start = local_rank * devices_per_rank
stop = start + devices_per_rank
device_env = current_platform.device_control_env_var
visible_devices = ",".join(
str(current_platform.device_id_to_physical_device_id(idx))
for idx in range(start, stop)
)
return {device_env: visible_devices}
device_ids = getattr(args, "device_ids", None)
if device_ids is not None:
if stop > len(device_ids):
raise ValueError(
f"--device-ids has {len(device_ids)} entries, but DP rank "
f"{local_rank} needs devices [{start}, {stop})"
)
return device_ids[start:stop]
return list(range(start, stop))
def _child_base_url(args: argparse.Namespace, port: int) -> str:
@@ -228,9 +233,7 @@ def _build_dp_supervisor_app(supervisor: DPSupervisor) -> FastAPI:
return app
def _run_vllm_dp_server(
child_args: argparse.Namespace, env_updates: dict[str, str]
) -> None:
def _run_vllm_dp_server(child_args: argparse.Namespace) -> None:
"""
Entrypoint function for the vLLM DP Server.
"""
@@ -241,7 +244,6 @@ def _run_vllm_dp_server(
os.setpgrp()
name = f"APIServer_DP{child_args.data_parallel_rank}"
update_environment_variables(env_updates)
set_process_title(name)
decorate_logs(name)
uvloop.run(run_server(child_args))
@@ -345,11 +347,10 @@ class DPSupervisor:
context = multiprocessing.get_context("spawn")
for local_rank in range(self.args.data_parallel_size_local):
child_args = _build_vllm_dp_server_args(self.args, local_rank)
child_env = _build_vllm_dp_server_env(self.args, local_rank)
process = context.Process(
target=_run_vllm_dp_server,
name=f"APIServer_DPRank_{child_args.data_parallel_rank}",
args=(child_args, child_env),
args=(child_args,),
)
process.start()
self._processes.append(process)
+8 -4
View File
@@ -219,10 +219,14 @@ class GenerateResponse(BaseModel):
class DerenderChatRequest(BaseModel):
"""Request for the /v1/chat/completions/derender endpoint.
"""Request for the /v1/chat/completions/derender endpoint (non-streaming).
Wraps a GenerateResponse and caller-supplied metadata needed to produce
a fully-formed ChatCompletionResponse without a GPU.
Wraps a complete GenerateResponse and caller-supplied metadata needed to
produce a fully-formed ChatCompletionResponse without a GPU.
Streaming derender would require a separate endpoint design with
incremental token delivery, ``OutputProcessor``-based detokenization,
and ``parser.parse_delta()`` instead of ``parser.parse()``.
"""
model: str
@@ -244,7 +248,7 @@ class DerenderChatRequest(BaseModel):
class DerenderCompletionRequest(BaseModel):
"""Request for the /v1/completions/derender endpoint.
"""Request for the /v1/completions/derender endpoint (non-streaming).
Parallel to DerenderChatRequest but handles the multi-prompt completions
case: one GenerateResponse per prompt, mirroring the list[GenerateRequest]
+156 -32
View File
@@ -27,6 +27,7 @@ from vllm.entrypoints.openai.completion.protocol import (
)
from vllm.entrypoints.openai.engine.protocol import (
ErrorResponse,
ToolCall,
UsageInfo,
)
from vllm.entrypoints.openai.engine.serving import resolve_token_id_placeholder
@@ -43,7 +44,6 @@ from vllm.entrypoints.serve.disagg.protocol import (
DerenderChatRequest,
DerenderCompletionRequest,
GenerateRequest,
GenerateResponseChoice,
MultiModalFeatures,
PlaceholderRangeInfo,
)
@@ -76,21 +76,83 @@ from vllm.utils.mistral import mt as _mt
logger = init_logger(__name__)
def _parse_token_id_placeholder(token: str) -> int | None:
"""Extract token ID from a 'token_id:N' placeholder string."""
if not token.startswith("token_id:"):
return None
try:
return int(token[len("token_id:") :])
except ValueError:
return None
def _correct_decoded_token(
token_id: int, context_token_ids: list[int], tokenizer: TokenizerLike
) -> str:
"""Use preceding tokens as context to fix U+FFFD from byte-fallback.
Mirrors LogprobsProcessor._correct_decoded_token in v1/engine/logprobs.py.
"""
max_ctx = min(len(context_token_ids), 4)
for num_ctx in range(1, max_ctx + 1):
context = context_token_ids[-num_ctx:]
full_decoded = tokenizer.decode(context + [token_id])
if full_decoded.endswith(""):
continue
clean_end = len(context)
for j in range(len(context) - 1, -1, -1):
if tokenizer.decode([context[j]]).endswith(""):
clean_end = j
else:
break
clean_prefix = tokenizer.decode(context[:clean_end]) if clean_end > 0 else ""
if full_decoded.startswith(clean_prefix):
return full_decoded[len(clean_prefix) :]
common_len = 0
for a, b in zip(clean_prefix, full_decoded):
if a != b:
break
common_len += 1
return full_decoded[common_len:]
return ""
def _resolve_logprobs(
logprobs: ChatCompletionLogProbs, tokenizer: TokenizerLike
) -> ChatCompletionLogProbs:
"""Resolve all token_id:N placeholders in a ChatCompletionLogProbs object."""
"""Resolve token_id:N placeholders in a ChatCompletionLogProbs object."""
if logprobs.content is None:
return logprobs
context_token_ids: list[int] = []
resolved_content = []
for entry in logprobs.content:
token_str, token_bytes = resolve_token_id_placeholder(entry.token, tokenizer)
sampled_id = _parse_token_id_placeholder(entry.token)
if token_str.endswith("") and sampled_id is not None:
token_str = _correct_decoded_token(sampled_id, context_token_ids, tokenizer)
token_bytes = list(token_str.encode("utf-8"))
resolved_top = []
for top in entry.top_logprobs:
top_str, top_bytes = resolve_token_id_placeholder(top.token, tokenizer)
top_id = _parse_token_id_placeholder(top.token)
if top_str.endswith("") and top_id is not None:
top_str = _correct_decoded_token(top_id, context_token_ids, tokenizer)
top_bytes = list(top_str.encode("utf-8"))
resolved_top.append(
top.model_copy(update={"token": top_str, "bytes": top_bytes})
)
resolved_content.append(
entry.model_copy(
update={
@@ -100,6 +162,10 @@ def _resolve_logprobs(
}
)
)
if sampled_id is not None:
context_token_ids.append(sampled_id)
return ChatCompletionLogProbs(content=resolved_content)
@@ -136,30 +202,6 @@ def _convert_chat_logprobs_to_completion_logprobs(
)
def _build_chat_choice(
choice: GenerateResponseChoice, tokenizer: TokenizerLike
) -> ChatCompletionResponseChoice:
"""Detokenize and resolve logprobs for a single GenerateResponseChoice.
Raises:
ValueError: if choice.token_ids is empty or None.
"""
if not choice.token_ids:
raise ValueError(f"choice {choice.index} has empty or null token_ids")
decoded_text = tokenizer.decode(choice.token_ids, skip_special_tokens=True)
resolved_logprobs = (
_resolve_logprobs(choice.logprobs, tokenizer)
if choice.logprobs is not None
else None
)
return ChatCompletionResponseChoice(
index=choice.index,
message=ChatMessage(role="assistant", content=decoded_text),
logprobs=resolved_logprobs,
finish_reason=choice.finish_reason,
)
class OpenAIServingRender:
def __init__(
self,
@@ -536,9 +578,12 @@ class OpenAIServingRender:
) -> ChatCompletionResponse | ErrorResponse:
"""Postprocess a GenerateResponse into a ChatCompletionResponse.
This is the symmetric inverse of render_chat_request: it detokenizes
output token IDs, resolves token_id:N logprob placeholders, and
formats the result as an OpenAI-compatible chat completion response.
Non-streaming only: expects the complete GenerateResponse with all
token IDs present. Uses ``parser.parse()`` for one-shot extraction.
When ``request.chat_request`` is provided, the parser splits the
output into (reasoning, content, tool_calls). Otherwise falls
back to plain detokenization.
"""
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
@@ -546,11 +591,89 @@ class OpenAIServingRender:
tokenizer = self.renderer.get_tokenizer()
gen = request.generate_response
chat_request = request.chat_request
choices: list[ChatCompletionResponseChoice] = []
try:
for choice in gen.choices:
choices.append(_build_chat_choice(choice, tokenizer))
if not choice.token_ids:
raise ValueError(
f"choice {choice.index} has empty or null token_ids"
)
resolved_logprobs = (
_resolve_logprobs(choice.logprobs, tokenizer)
if choice.logprobs is not None
else None
)
if self.parser is not None and chat_request is not None:
# Parser path: decode with special tokens preserved
# so the parser can see markers like </think>,
# <tool_call>, or Harmony channel tokens.
decoded_text = tokenizer.decode(
choice.token_ids, skip_special_tokens=False
)
chat_template_kwargs: dict[str, Any] = {}
if not self.use_harmony:
chat_template_kwargs = (
chat_request.build_chat_params(
self.chat_template,
self.chat_template_content_format,
)
.with_defaults(self.default_chat_template_kwargs)
.chat_template_kwargs
)
parser = self.parser(
tokenizer,
chat_request.tools,
chat_template_kwargs=chat_template_kwargs,
)
reasoning, content, tool_calls = parser.parse(
decoded_text,
chat_request,
enable_auto_tools=self.enable_auto_tools,
model_output_token_ids=choice.token_ids,
)
if not getattr(chat_request, "include_reasoning", True):
reasoning = None
tc_items = (
[
ToolCall(
id=random_uuid(),
function=tc,
)
for tc in tool_calls
]
if tool_calls
else []
)
message = ChatMessage(
role="assistant",
reasoning=reasoning,
content=content,
tool_calls=tc_items,
)
else:
# No parser: plain detokenization.
decoded_text = tokenizer.decode(
choice.token_ids, skip_special_tokens=True
)
message = ChatMessage(role="assistant", content=decoded_text)
choices.append(
ChatCompletionResponseChoice(
index=choice.index,
message=message,
logprobs=resolved_logprobs,
finish_reason=choice.finish_reason,
)
)
except ValueError as exc:
return self.create_error_response(str(exc))
@@ -587,8 +710,9 @@ class OpenAIServingRender:
) -> CompletionResponse | ErrorResponse:
"""Postprocess a list of GenerateResponses into a CompletionResponse.
Mirrors the multi-prompt completions case: one GenerateResponse per
prompt, parallel to the list[GenerateRequest] from /v1/completions/render.
Non-streaming only. Mirrors the multi-prompt completions case: one
GenerateResponse per prompt, parallel to the list[GenerateRequest]
from /v1/completions/render.
"""
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
+7 -1
View File
@@ -209,6 +209,7 @@ if TYPE_CHECKING:
VLLM_EXECUTE_MODEL_TIMEOUT_SECONDS: int = 300
VLLM_WORKER_SHUTDOWN_TIMEOUT_SECONDS: int = 5
VLLM_KV_CACHE_LAYOUT: Literal["NHD", "HND"] | None = None
VLLM_USE_PACKED_HMA_KV_CACHE: bool = False
VLLM_SSM_CONV_STATE_LAYOUT: Literal["SD", "DS"] | None = None
VLLM_COMPUTE_NANS_IN_LOGITS: bool = False
VLLM_ROCM_QUICK_REDUCE_QUANTIZATION: Literal[
@@ -485,7 +486,7 @@ def get_vllm_port() -> int | None:
raise ValueError(
f"VLLM_PORT '{port}' appears to be a URI. "
"This may be caused by a Kubernetes service discovery issue,"
"check the warning in: https://docs.vllm.ai/en/stable/serving/env_vars.html"
"check the warning in: https://docs.vllm.ai/en/latest/configuration/env_vars.html"
) from None
raise ValueError(f"VLLM_PORT '{port}' must be a valid integer") from err
@@ -1608,6 +1609,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
"VLLM_KV_CACHE_LAYOUT": env_with_choices(
"VLLM_KV_CACHE_LAYOUT", None, ["NHD", "HND"]
),
# Opt into packed per-block KV cache allocation for multi-group
# attention-only HMA models (e.g. gpt-oss, Gemma 3/4).
"VLLM_USE_PACKED_HMA_KV_CACHE": lambda: bool(
int(os.getenv("VLLM_USE_PACKED_HMA_KV_CACHE", "0"))
),
# SSM conv state layout used for Mamba models.
# - SD: (state_len, dim) — dim contiguous (default)
# - DS: (dim, state_len) — TP-sharded dim on dim1,
+2 -2
View File
@@ -17,7 +17,7 @@ from vllm.lora.utils import (
)
from vllm.model_executor.model_loader.tensorizer import TensorizerConfig
from vllm.model_executor.models.utils import WeightsMapper
from vllm.utils.platform_utils import is_pin_memory_available
from vllm.utils.torch_utils import PIN_MEMORY
logger = init_logger(__name__)
@@ -126,7 +126,7 @@ class LoRAModel:
skip_prefixes: list[str] | None = None,
) -> "LoRAModel":
"""Create a LoRAModel from a dictionary of tensors."""
pin_memory = str(device) == "cpu" and is_pin_memory_available()
pin_memory = str(device) == "cpu" and PIN_MEMORY
loras: dict[str, LoRALayerWeights] = {}
for tensor_name, tensor in tensors.items():
if is_base_embedding_weights(tensor_name):
+2 -2
View File
@@ -7,7 +7,7 @@ import torch
import torch.types
from vllm.lora.peft_helper import PEFTHelper
from vllm.utils.platform_utils import is_pin_memory_available
from vllm.utils.torch_utils import PIN_MEMORY
class LoRALayerWeights:
@@ -79,7 +79,7 @@ class LoRALayerWeights:
dtype: torch.dtype,
device: torch.types.Device,
) -> "LoRALayerWeights":
pin_memory = str(device) == "cpu" and is_pin_memory_available()
pin_memory = str(device) == "cpu" and PIN_MEMORY
lora_a = torch.zeros(
[rank, input_dim], dtype=dtype, device=device, pin_memory=pin_memory
)
+2 -2
View File
@@ -42,7 +42,7 @@ from vllm.model_executor.models.utils import PPMissingLayer
from vllm.multimodal import MULTIMODAL_REGISTRY
from vllm.multimodal.encoder_budget import MultiModalBudget
from vllm.utils.cache import LRUCache
from vllm.utils.platform_utils import is_pin_memory_available
from vllm.utils.torch_utils import PIN_MEMORY
logger = init_logger(__name__)
@@ -801,7 +801,7 @@ class LoRAModelManager:
# 2. The weight packing above (e.g., pack_moe) may invalidate the
# pin_memory allocation, so we execute it after packing.
pin_memory = str(lora_device) == "cpu" and is_pin_memory_available()
pin_memory = str(lora_device) == "cpu" and PIN_MEMORY
if pin_memory:
for lora in lora_model.loras.values():
if isinstance(lora.lora_a, list):
@@ -1684,12 +1684,13 @@ class MLACommonMetadataBuilder(AttentionMetadataBuilder[M]):
# [[0, 0, 0, 0], [256, 256, 256, 256], [512, 512, 512, 512]]
# Note(simon): this is done in CPU because of downstream's
# of `to_list`.
chunk_starts = (
chunk_starts = torch.empty(
num_chunks, num_prefills, dtype=torch.int32, pin_memory=True
).copy_(
torch.arange(num_chunks, dtype=torch.int32)
.multiply_(max_context_chunk)
.unsqueeze(1)
.expand(-1, num_prefills)
* max_context_chunk
).pin_memory()
)
chunk_ends = torch.min(
context_lens_cpu.unsqueeze(0), chunk_starts + max_context_chunk
)
@@ -1746,12 +1747,13 @@ class MLACommonMetadataBuilder(AttentionMetadataBuilder[M]):
)
* self.dcp_local_block_size
)
local_chunk_starts = (
local_chunk_starts = torch.empty(
num_chunks, num_prefills, dtype=torch.int32, pin_memory=True
).copy_(
torch.arange(num_chunks, dtype=torch.int32)
.multiply_(padded_local_max_context_chunk_across_ranks)
.unsqueeze(1)
.expand(-1, num_prefills)
* padded_local_max_context_chunk_across_ranks
).pin_memory()
)
local_chunk_ends = torch.min(
padded_local_context_lens_cpu.unsqueeze(0),
local_chunk_starts
@@ -28,6 +28,7 @@ from vllm.utils.flashinfer import (
is_flashinfer_cudnn_fp8_prefill_attn_supported,
)
from vllm.utils.math_utils import round_up
from vllm.utils.torch_utils import async_tensor_h2d
from vllm.v1.attention.backends.fa_utils import get_flash_attn_version
from vllm.v1.attention.backends.registry import AttentionBackendEnum
from vllm.v1.attention.ops.vit_attn_wrappers import (
@@ -311,7 +312,7 @@ class MMEncoderAttention(CustomOp):
)
cu_seqlens = np.concatenate([cu_seqlens_qko, cu_seqlens_v])
cu_seqlens = torch.from_numpy(cu_seqlens).to(device, non_blocking=True)
cu_seqlens = async_tensor_h2d(cu_seqlens, device=device)
return cu_seqlens
def __init__(
@@ -21,7 +21,6 @@ from vllm.model_executor.layers.fused_moe.config import (
FusedMoEQuantConfig,
)
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
from vllm.model_executor.layers.fused_moe.utils import moe_kernel_quantize_input
from vllm.model_executor.layers.quantization.utils.nvfp4_emulation_utils import (
dequantize_to_dtype,
)
@@ -135,14 +134,6 @@ class Nvfp4QuantizationEmulationTritonExperts(TritonExperts):
swizzle=False,
)
hidden_states, _ = moe_kernel_quantize_input(
A=hidden_states,
A_scale=self.quant_config.a1_gscale,
quant_dtype="nvfp4",
per_act_token_quant=False,
quantization_emulation=True,
)
# Activation quantization/dequantization is deferred to
# `moe_kernel_quantize_input` in TritonExperts.apply.
super().apply(
@@ -21,7 +21,6 @@ from vllm.model_executor.layers.fused_moe.config import (
FusedMoEQuantConfig,
)
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
from vllm.model_executor.layers.fused_moe.utils import moe_kernel_quantize_input
from vllm.model_executor.layers.quantization.utils.mxfp4_utils import dequant_mxfp4
from vllm.model_executor.layers.quantization.utils.mxfp6_utils import dequant_mxfp6
from vllm.model_executor.layers.quantization.utils.ocp_mx_utils import (
@@ -155,16 +154,6 @@ class OCP_MXQuantizationEmulationTritonExperts(TritonExperts):
w2, self.w2_scale_val, hidden_states.dtype
)
# Apply activation QDQ if needed by the OCP MX scheme
hidden_states, _ = moe_kernel_quantize_input(
A=hidden_states,
A_scale=None,
quant_dtype=self.quant_config.quant_dtype,
per_act_token_quant=False,
ocp_mx_scheme=self.ocp_mx_scheme,
quantization_emulation=True,
)
# Activation quantization/dequantization is deferred to
# `moe_kernel_quantize_input` in TritonExperts.apply.
super().apply(
@@ -245,7 +245,7 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
lora_unquantized_hidden_states = hidden_states
hidden_states, a1q_scale = moe_kernel_quantize_input(
hidden_states,
self.a1_scale,
self.a1_scale or self.a1_gscale,
self.quant_dtype,
self.per_act_token_quant,
self.block_shape,
@@ -296,6 +296,7 @@ def moe_kernel_quantize_input(
if not quantization_emulation:
return _nvfp4_quantize(A, A_scale, is_sf_swizzled_layout=is_scale_swizzled)
else:
assert A_scale is not None
A = ref_nvfp4_quant_dequant(A, A_scale, block_size=16)
return A, None
elif quant_dtype == "mxfp4":
@@ -10,6 +10,7 @@ import torch.nn as nn
from vllm.config.pooler import SequencePoolingType
from vllm.model_executor.layers.pooler import PoolingParamsUpdate
from vllm.tasks import PoolingTask
from vllm.utils.torch_utils import async_tensor_h2d
from vllm.v1.pool.metadata import PoolingMetadata
SequencePoolingMethodOutput: TypeAlias = torch.Tensor | list[torch.Tensor]
@@ -74,15 +75,14 @@ class MeanPool(SequencePoolingMethod):
# early return for empty batch
return hidden_states.new_empty((0, hidden_size), dtype=torch.float32)
# Build segment_ids on CPU so repeat_interleave doesn't need to sync
# GPU->CPU to learn its data-dependent output length, then upload
# non-blocking. eg. [2, 1, 3] -> [0, 0, 1, 2, 2, 2]
prompt_lens = async_tensor_h2d(
prompt_lens_cpu, device=hidden_states.device, dtype=torch.int64
)
# eg. [2, 1, 3] -> [0, 0, 1, 2, 2, 2]
segment_ids = torch.repeat_interleave(
torch.arange(num_seqs, dtype=torch.long),
prompt_lens_cpu,
).to(hidden_states.device, non_blocking=True)
prompt_lens = prompt_lens_cpu.to(
hidden_states.device, dtype=torch.int64, non_blocking=True
torch.arange(num_seqs, device=hidden_states.device, dtype=torch.long),
prompt_lens,
output_size=int(prompt_lens_cpu.sum()),
)
segment_sums = torch.zeros(
(num_seqs, hidden_size),
+22 -16
View File
@@ -1001,24 +1001,30 @@ class DeepseekV2MLAAttention(nn.Module):
# IndexCache config
# Refer: https://arxiv.org/abs/2603.12201 for more details.
_skip_topk = False
_index_topk_freq = getattr(config, "index_topk_freq", 1)
_index_topk_pattern = getattr(config, "index_topk_pattern", None)
_index_skip_topk_offset = getattr(config, "index_skip_topk_offset", 2)
layer_id = extract_layer_index(prefix)
is_mtp_layer = False
if self.is_v32:
_index_topk_freq = getattr(config, "index_topk_freq", 1)
_index_topk_pattern = getattr(config, "index_topk_pattern", None)
_index_skip_topk_offset = getattr(config, "index_skip_topk_offset", 2)
layer_id = extract_layer_index(prefix)
if _index_topk_pattern is None:
_skip_topk = (
max(layer_id - _index_skip_topk_offset + 1, 0) % _index_topk_freq != 0
if _index_topk_pattern is None:
_skip_topk = (
max(layer_id - _index_skip_topk_offset + 1, 0) % _index_topk_freq
!= 0
)
elif 0 <= layer_id < len(_index_topk_pattern):
_skip_topk = _index_topk_pattern[layer_id] == "S"
# The skip pattern only governs backbone layers. MTP/nextn
# layers (layer_id >= num_hidden_layers) always build a full
# indexer: they compute indices at draft step 0 and toggle
# at runtime via set_skip_topk
# (index_share_for_mtp_iteration).
_num_hidden_layers = getattr(config, "num_hidden_layers", None)
is_mtp_layer = (
_num_hidden_layers is not None and layer_id >= _num_hidden_layers
)
elif 0 <= layer_id < len(_index_topk_pattern):
_skip_topk = _index_topk_pattern[layer_id] == "S"
# The skip pattern only governs backbone layers. MTP/nextn layers
# (layer_id >= num_hidden_layers) always build a full indexer: they
# compute indices at draft step 0 and toggle at runtime via
# set_skip_topk (index_share_for_mtp_iteration).
_num_hidden_layers = getattr(config, "num_hidden_layers", None)
is_mtp_layer = _num_hidden_layers is not None and layer_id >= _num_hidden_layers
if self.is_v32 and (not _skip_topk or is_mtp_layer):
self.indexer_rope_emb = get_rope(
+3 -2
View File
@@ -66,6 +66,7 @@ from vllm.model_executor.models.utils import maybe_prefix
from vllm.model_executor.models.vision import is_vit_use_data_parallel
from vllm.platforms import current_platform
from vllm.transformers_utils.configs.moonvit import MoonViTConfig
from vllm.utils.torch_utils import async_tensor_h2d
def _apply_rope_input_validation(x, freqs_cis):
@@ -758,7 +759,7 @@ class MoonVitPretrainedModel(PreTrainedModel):
),
]
)
metadata["cu_seqlens"] = torch.from_numpy(cu_seqlens_np).to(device)
metadata["cu_seqlens"] = async_tensor_h2d(cu_seqlens_np, device=device)
if max_seqlen_override is not None:
max_seqlen_val = int(max_seqlen_override)
@@ -770,7 +771,7 @@ class MoonVitPretrainedModel(PreTrainedModel):
metadata["max_seqlen"] = torch.tensor(max_seqlen_val, dtype=torch.int32)
gather_idx_np = _build_merge_gather_idx(grid_pairs, self.merge_kernel_size)
metadata["merge_gather_idx"] = torch.from_numpy(gather_idx_np).to(device)
metadata["merge_gather_idx"] = async_tensor_h2d(gather_idx_np, device=device)
return metadata
+2 -3
View File
@@ -83,9 +83,8 @@ from vllm.multimodal.parse import MultiModalDataItems
from vllm.multimodal.processing import PromptReplacement, PromptUpdate
from vllm.platforms import current_platform
from vllm.sequence import IntermediateTensors
from vllm.utils.platform_utils import is_pin_memory_available
from vllm.utils.tensor_schema import TensorSchema, TensorShape
from vllm.utils.torch_utils import async_tensor_h2d
from vllm.utils.torch_utils import PIN_MEMORY, async_tensor_h2d
from vllm.v1.attention.backends.registry import AttentionBackendEnum
from vllm.v1.worker.encoder_cudagraph_defs import EncoderCudaGraphReplayBuffers
@@ -825,7 +824,7 @@ class Qwen2_5_VisionTransformer(nn.Module):
@staticmethod
def invert_permutation(perm: torch.Tensor) -> torch.Tensor:
# building the inverse permutation in O(n) time
inv = torch.empty_like(perm, pin_memory=is_pin_memory_available())
inv = torch.empty_like(perm, pin_memory=PIN_MEMORY)
inv[perm] = torch.arange(perm.numel(), device=perm.device, dtype=perm.dtype)
return inv
+78 -26
View File
@@ -1202,6 +1202,49 @@ class Qwen3VLDummyInputsBuilder(BaseDummyInputsBuilder[Qwen3VLProcessingInfo]):
return video_items
def _replace_video_token_placeholders(
prompt_ids: list[int],
target: list[int],
replacements: list[list[int]],
) -> list[int]:
"""Replace each 3-token video placeholder with its expanded sequence.
Args:
prompt_ids: Token IDs of the original (unexpanded) prompt.
target: 3-element list ``[vision_start_id, video_pad_id,
vision_end_id]`` to search for.
replacements: Per-video expanded token sequences, in prompt order.
Returns:
Token IDs with every placeholder triplet replaced.
"""
result: list[int] = []
repl_idx = 0
i = 0
n = len(prompt_ids)
t0, t1, t2 = target
num_repl = len(replacements)
while i < n:
if (
i + 2 < n
and prompt_ids[i] == t0
and prompt_ids[i + 1] == t1
and prompt_ids[i + 2] == t2
):
result.extend(replacements[repl_idx])
repl_idx += 1
i += 3
else:
result.append(prompt_ids[i])
i += 1
assert repl_idx == num_repl, (
f"Found {repl_idx} video placeholders but expected {num_repl}"
)
return result
class Qwen3VLMultiModalProcessor(BaseMultiModalProcessor[Qwen3VLProcessingInfo]):
def _call_hf_processor(
self,
@@ -1211,15 +1254,23 @@ class Qwen3VLMultiModalProcessor(BaseMultiModalProcessor[Qwen3VLProcessingInfo])
tok_kwargs: Mapping[str, object],
) -> BatchFeature:
mm_data = dict(mm_data)
processor = self.info.get_hf_processor(**mm_kwargs)
# Separate video processing from image processing. Because the videos
# are processed into several image patches
video_input_ids_lst: list[list[int]] = []
if videos := mm_data.pop("videos", []):
video_grid_thw_lst = []
pixel_values_videos_lst = []
timestamps_per_video = []
hf_config = self.info.get_hf_config()
tokenizer = self.info.get_tokenizer()
merge_size = hf_config.vision_config.spatial_merge_size
video_pruning_rate = self.info.ctx.get_mm_config().video_pruning_rate
vision_start_token_id = hf_config.vision_start_token_id
vision_end_token_id = hf_config.vision_end_token_id
video_token_id = hf_config.video_token_id
for item in videos:
video_array, metadata = item
@@ -1269,55 +1320,38 @@ class Qwen3VLMultiModalProcessor(BaseMultiModalProcessor[Qwen3VLProcessingInfo])
tok_kwargs=tok_kwargs,
)
merge_size = processor.video_processor.merge_size
# Get video grid info for EVS calculation.
# Discard HF output input_ids — we use get_video_repl below
# to generate the correct (EVS-adjusted) token sequence.
video_outputs.pop("input_ids", None)
video_grid_thw = video_outputs["video_grid_thw"]
num_frames = int(video_grid_thw[0, 0])
tokens_per_frame_base = int(video_grid_thw[0, 1:].prod()) // (
merge_size**2
)
# Apply EVS if enabled.
video_pruning_rate = self.info.ctx.get_mm_config().video_pruning_rate
if video_pruning_rate is not None and video_pruning_rate > 0.0:
num_tokens = compute_retained_tokens_count(
tokens_per_frame=tokens_per_frame_base,
num_frames=num_frames,
q=video_pruning_rate,
)
# Here we just need placeholders that won't actually be replaced -
# we just need to make sure the total number of tokens is correct
# assign all tokens to the first frame.
tokens_per_frame = [num_tokens] + [0] * (num_frames - 1)
select_token_id = False
else:
tokens_per_frame = [tokens_per_frame_base] * num_frames
select_token_id = True
# Generate the video replacement with EVS-adjusted token counts
tokenizer = self.info.get_tokenizer()
hf_config = self.info.get_hf_config()
video_repl = Qwen3VLMultiModalProcessor.get_video_repl(
tokens_per_frame=tokens_per_frame,
timestamps=timestamps,
tokenizer=tokenizer,
vision_start_token_id=hf_config.vision_start_token_id,
vision_end_token_id=hf_config.vision_end_token_id,
video_token_id=hf_config.video_token_id,
vision_start_token_id=vision_start_token_id,
vision_end_token_id=vision_end_token_id,
video_token_id=video_token_id,
select_token_id=select_token_id,
)
# Convert token IDs to text for the HF processor flow
video_placeholder = tokenizer.decode(
video_repl.full, skip_special_tokens=False
)
input_ids = video_outputs.pop("input_ids")
video_placeholder = processor.tokenizer.batch_decode(input_ids)[0]
prompt = prompt.replace(
"<|vision_start|><|video_pad|><|vision_end|>",
video_placeholder,
1,
)
video_input_ids_lst.append(list(video_repl.full))
video_grid_thw_lst.append(video_outputs["video_grid_thw"])
pixel_values_videos_lst.append(video_outputs["pixel_values_videos"])
@@ -1335,6 +1369,24 @@ class Qwen3VLMultiModalProcessor(BaseMultiModalProcessor[Qwen3VLProcessingInfo])
mm_kwargs=mm_kwargs,
tok_kwargs=tok_kwargs,
)
# Replace each placeholder triplet with pre-computed video tokens.
if video_input_ids_lst:
hf_config = self.info.get_hf_config()
video_target = [
hf_config.vision_start_token_id,
hf_config.video_token_id,
hf_config.vision_end_token_id,
]
input_ids = processed_outputs.pop("input_ids")
if not isinstance(input_ids, list):
input_ids = input_ids.tolist()
(prompt_ids,) = input_ids
expanded_ids = _replace_video_token_placeholders(
prompt_ids, video_target, video_input_ids_lst
)
processed_outputs["input_ids"] = [expanded_ids]
combined_outputs = dict(
processed_outputs,
**video_outputs,
+2 -1
View File
@@ -13,6 +13,7 @@ from vllm.config.cache import CacheDType
from vllm.platforms.interface import DeviceCapability
from vllm.triton_utils import tl, triton
from vllm.utils.math_utils import cdiv
from vllm.utils.torch_utils import np_to_pinned_tensor
from vllm.v1.attention.backend import (
AttentionBackend,
AttentionCGSupport,
@@ -207,7 +208,7 @@ class DeepseekV4FlashMLAMetadataBuilder(
# Zero-fill for cudagraphs
self.req_id_per_token_buffer.fill_(0)
self.req_id_per_token_buffer[: req_id_per_token.shape[0]].copy_(
torch.from_numpy(req_id_per_token), non_blocking=True
np_to_pinned_tensor(req_id_per_token), non_blocking=True
)
req_id_per_token = self.req_id_per_token_buffer[:num_tokens]
+14 -2
View File
@@ -488,7 +488,13 @@ class MultiModalBatchedField(BaseMultiModalField):
# An optimization when `batch` contains only one tensor:
# - produce exactly same result as `torch.stack(batch)`
# - will achieve zero-copy if the tensor is contiguous
return batch[0].unsqueeze(0).contiguous()
out = batch[0].unsqueeze(0)
if not pin_memory:
return out.contiguous()
# Avoid extra copy - pinning unpinned memory will make it contiguous
if not out.is_contiguous() and out.is_pinned():
out = out.contiguous()
return out.pin_memory()
first_shape = batch[0].shape
if all(elem.shape == first_shape for elem in batch):
out = torch.empty(
@@ -538,7 +544,13 @@ class MultiModalFlatField(BaseMultiModalField):
# An optimization when `batch` contains only one tensor:
# - produce exactly same result as `torch.concat(batch)`
# - will achieve zero-copy if the tensor is contiguous
return batch[0].contiguous()
out = batch[0]
if not pin_memory:
return out.contiguous()
# Avoid extra copy - pinning unpinned memory will make it contiguous
if not out.is_contiguous() and out.is_pinned():
out = out.contiguous()
return out.pin_memory()
dim = self.dim + (self.dim < 0) * len(batch[0].shape)

Some files were not shown because too many files have changed in this diff Show More