Merge branch 'main' into wentao-fix-mypy-models-a-b

This commit is contained in:
yewentao256
2026-07-23 15:20:32 +00:00
89 changed files with 2476 additions and 556 deletions
+1 -1
View File
@@ -5,7 +5,7 @@ steps:
- label: PyTorch Compilation Unit Tests
device: h200_35gb
key: pytorch-compilation-unit-tests
timeout_in_minutes: 90
timeout_in_minutes: 110
source_file_dependencies:
- vllm/__init__.py
- vllm/_aiter_ops.py
+12
View File
@@ -181,6 +181,18 @@ pull_request_rules:
add:
- performance
- name: label-quantization
description: Automatically apply quantization label
conditions:
- label != stale
- or:
- files~=^vllm/model_executor/layers/quantization/
- title~=(?i)quant
actions:
label:
add:
- quantization
- name: label-qwen
description: Automatically apply qwen label
conditions:
+13 -1
View File
@@ -130,6 +130,18 @@ jobs:
},
],
},
quantization: {
keywords: [
{
term: "quantization",
searchIn: "both"
},
{
term: "quantized",
searchIn: "both"
},
],
},
"intel-gpu": {
// Keyword search - matches whole words only (with word boundaries)
keywords: [
@@ -520,4 +532,4 @@ jobs:
issue_number: context.issue.number,
body: message,
});
core.notice(`Requested missing ROCm info from @${author}: ${missing.map(m => m.name).join(', ')}`);
core.notice(`Requested missing ROCm info from @${author}: ${missing.map(m => m.name).join(', ')}`);
+3
View File
@@ -1025,6 +1025,9 @@ __global__ void gather_and_maybe_dequant_cache(
batch_offset += offset;
int32_t block_table_id = batch_offset / block_size;
int32_t slot_id = batch_offset % block_size;
// seq_starts may push the block index past the end of the batch's block
// table row.
if (block_table_id >= block_table_stride) continue;
int32_t block_table_offset = batch_id * block_table_stride + block_table_id;
int32_t block_id = block_table[block_table_offset];
int64_t cache_offset =
+1 -1
View File
@@ -793,7 +793,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
# Install FlashInfer JIT cache (requires CUDA-version-specific index URL)
# https://docs.flashinfer.ai/installation.html
# From versions.json: .flashinfer.version
ARG FLASHINFER_VERSION=0.6.14
ARG FLASHINFER_VERSION=0.6.15.post1
RUN --mount=type=cache,target=/opt/uv/cache \
uv pip install --system flashinfer-jit-cache==${FLASHINFER_VERSION} \
--index-url https://flashinfer.ai/whl/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.')
+25 -1
View File
@@ -86,6 +86,29 @@ RUN --mount=type=cache,target=/root/.cache/uv \
mkdir -p /tmp/hf-xet/dist && \
cp dist/*.whl /tmp/hf-xet/dist/
# Build LLVM 20 from source for llvmlite (system repos ship LLVM 21 which
# llvmlite v0.47 does not support; only SystemZ target is needed).
FROM base AS llvm20-build
ARG LLVM_VERSION=20.1.8
WORKDIR /tmp
RUN microdnf install -y ninja-build gcc gcc-c++ python3 xz && \
curl -LO https://github.com/llvm/llvm-project/releases/download/llvmorg-${LLVM_VERSION}/llvm-project-${LLVM_VERSION}.src.tar.xz && \
tar -xf llvm-project-${LLVM_VERSION}.src.tar.xz && \
cmake -G Ninja -S llvm-project-${LLVM_VERSION}.src/llvm -B build \
-DCMAKE_BUILD_TYPE=Release \
-DCMAKE_INSTALL_PREFIX=/opt/llvm20 \
-DLLVM_TARGETS_TO_BUILD="SystemZ" \
-DLLVM_ENABLE_RTTI=ON \
-DLLVM_BUILD_TOOLS=OFF \
-DLLVM_BUILD_UTILS=ON \
-DLLVM_BUILD_EXAMPLES=OFF \
-DLLVM_BUILD_TESTS=OFF \
-DLLVM_INCLUDE_TESTS=OFF \
-DLLVM_INCLUDE_EXAMPLES=OFF \
-DLLVM_INCLUDE_BENCHMARKS=OFF && \
ninja -C build install && \
rm -rf build llvm-project-${LLVM_VERSION}.src*
# Build numba
FROM python-install AS numba-builder
@@ -96,11 +119,13 @@ WORKDIR /tmp
# Clone all required dependencies
RUN --mount=type=cache,target=/root/.cache/uv \
--mount=type=bind,from=llvm20-build,source=/opt/llvm20,target=/opt/llvm20 \
microdnf install ninja-build gcc gcc-c++ -y && \
git clone --recursive https://github.com/numba/llvmlite.git -b v0.47.0 && \
git clone --recursive https://github.com/numba/numba.git -b ${NUMBA_VERSION} && \
cd llvmlite && \
uv pip install 'cmake<4' 'setuptools<70' numpy && \
CMAKE_PREFIX_PATH=/opt/llvm20 LLVM_CONFIG=/opt/llvm20/bin/llvm-config \
python setup.py bdist_wheel && \
cd ../numba && \
if ! grep '#include "dynamic_annotations.h"' numba/_dispatcher.cpp; then \
@@ -158,7 +183,6 @@ RUN --mount=type=cache,target=/root/.cache/uv \
NUMBA_WHL_FILE=$(ls /tmp/numba-wheels/*.whl) && \
OPENCV_WHL_FILE=$(ls /tmp/opencv-wheels/*.whl) && \
uv pip install -v \
$ARROW_WHL_FILE \
$VISION_WHL_FILE \
$HF_XET_WHL_FILE \
$LLVM_WHL_FILE \
+1 -1
View File
@@ -68,7 +68,7 @@
"default": "true"
},
"FLASHINFER_VERSION": {
"default": "0.6.14"
"default": "0.6.15.post1"
},
"GDRCOPY_CUDA_VERSION": {
"default": "12.8"
+1 -1
View File
@@ -315,7 +315,7 @@ vLLM CPU supports data parallel (DP), tensor parallel (TP) and pipeline parallel
- vLLM CPU supports quantizations:
- AWQ (x86 only)
- GPTQ (x86 only)
- compressed-tensor INT8 W8A8 (x86, s390x)
- compressed-tensor INT8 W8A8 (x86 only)
### Why do I see `get_mempolicy: Operation not permitted` when running in Docker?
@@ -11,7 +11,7 @@ Currently, the CPU implementation for s390x architecture supports FP32, BF16 and
- OS: `Linux`
- SDK: `gcc/g++ >= 14.0.0` or later with Command Line Tools
- Instruction Set Architecture (ISA): VXE support is required. Works with Z14 and above.
- Build install python packages: `torchvision`, `llvmlite`, `numba`, `pyarrow (for testing)`, `opencv-headless`
- Build from source python packages (no pre-built s390x wheels): `torchvision`, `llvmlite`, `numba`, `opencv-python-headless`, `hf-xet`
--8<-- [end:requirements]
--8<-- [start:set-up-using-python]
@@ -28,13 +28,24 @@ Install the following packages from the package manager before building the vLLM
```bash
dnf install -y \
which procps findutils tar vim git gcc-toolset-14 gcc-toolset-14-binutils gcc-toolset-14-libatomic-devel zlib-devel \
which procps findutils tar vim git patch xz ninja-build \
gcc-toolset-14 gcc-toolset-14-binutils gcc-toolset-14-libatomic-devel zlib-devel \
libjpeg-turbo-devel libtiff-devel libpng-devel libwebp-devel freetype-devel harfbuzz-devel \
openssl-devel openblas openblas-devel autoconf automake libtool cmake numpy libsndfile \
clang llvm-devel llvm-static clang-devel
```
Install rust>=1.80 which is needed for `outlines-core` and `uvloop` python packages installation.
Build and install `numactl` from source:
```bash
curl -LO https://github.com/numactl/numactl/archive/refs/tags/v2.0.19.tar.gz
tar -xvzf v2.0.19.tar.gz
cd numactl-2.0.19
./autogen.sh && ./configure && make && make install
cd ..
```
Install rust>=1.80 which is needed for `outlines-core`, `uvloop`, and `hf-xet` python packages installation.
```bash
curl https://sh.rustup.rs -sSf | sh -s -- -y && \
@@ -44,26 +55,79 @@ curl https://sh.rustup.rs -sSf | sh -s -- -y && \
Execute the following commands to build and install vLLM from source.
!!! tip
Please build the following dependencies, `torchvision`, `llvmlite`, `numba`, `llguidance`, `pyarrow`, `opencv-headless` from source before building vLLM.
Pre-built wheels are not available for s390x for the following packages. Build them from source before building vLLM: `torchvision`, `llvmlite`, `numba`, `opencv-python-headless`, `hf-xet`.
See `docker/Dockerfile.s390x` for exact versions and build commands used in each multi-stage build.
!!! note "LLVM 20 required for llvmlite"
`llvmlite v0.47` requires LLVM 20, but UBI 9.6 repos ship LLVM 21 which is
not compatible. You must build LLVM 20 from source before building `llvmlite`:
```bash
curl -LO https://github.com/llvm/llvm-project/releases/download/llvmorg-20.1.8/llvm-project-20.1.8.src.tar.xz
tar -xf llvm-project-20.1.8.src.tar.xz
cmake -G Ninja -S llvm-project-20.1.8.src/llvm -B llvm-build \
-DCMAKE_BUILD_TYPE=Release \
-DCMAKE_INSTALL_PREFIX=/opt/llvm20 \
-DLLVM_TARGETS_TO_BUILD="SystemZ" \
-DLLVM_ENABLE_RTTI=ON \
-DLLVM_BUILD_TOOLS=OFF \
-DLLVM_BUILD_UTILS=ON \
-DLLVM_BUILD_EXAMPLES=OFF \
-DLLVM_BUILD_TESTS=OFF \
-DLLVM_INCLUDE_TESTS=OFF \
-DLLVM_INCLUDE_EXAMPLES=OFF \
-DLLVM_INCLUDE_BENCHMARKS=OFF
ninja -C llvm-build install
```
Then build `llvmlite` pointing to LLVM 20:
```bash
CMAKE_PREFIX_PATH=/opt/llvm20 LLVM_CONFIG=/opt/llvm20/bin/llvm-config \
python setup.py bdist_wheel
```
```bash
uv pip install -v \
-r requirements/build/cpu.txt \
-r requirements/cpu.txt \
--torch-backend cpu \
--index-strategy unsafe-best-match && \
VLLM_TARGET_DEVICE=cpu python setup.py bdist_wheel && \
uv pip install dist/*.whl
uv pip install -v \
/path/to/torchvision.whl \
/path/to/llvmlite.whl \
/path/to/numba.whl \
/path/to/opencv_python_headless.whl \
/path/to/hf_xet.whl \
-r requirements/build/cpu.txt \
-r requirements/cpu.txt \
--torch-backend cpu \
--index-strategy unsafe-best-match && \
VLLM_TARGET_DEVICE=cpu VLLM_CPU_MOE_PREPACK=0 python setup.py bdist_wheel && \
uv pip install dist/*.whl
```
??? console "pip"
```bash
pip install -v \
--extra-index-url https://download.pytorch.org/whl/cpu \
-r requirements/build/cpu.txt \
-r requirements/cpu.txt \
VLLM_TARGET_DEVICE=cpu python setup.py bdist_wheel && \
pip install dist/*.whl
pip install -v \
--extra-index-url https://download.pytorch.org/whl/cpu \
/path/to/torchvision.whl \
/path/to/llvmlite.whl \
/path/to/numba.whl \
/path/to/opencv_python_headless.whl \
/path/to/hf_xet.whl \
-r requirements/build/cpu.txt \
-r requirements/cpu.txt && \
VLLM_TARGET_DEVICE=cpu VLLM_CPU_MOE_PREPACK=0 python setup.py bdist_wheel && \
pip install dist/*.whl
```
!!! warning "Protobuf workaround for s390x"
The C++ protobuf extension crashes on s390x. After installation, set the
following environment variable and remove the C++ extensions:
```bash
export PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=python
# Remove C++ protobuf extensions that crash on s390x
SITE_PKGS=$(python -c "import site; print(site.getsitepackages()[0])")
rm -rf "$SITE_PKGS/google/_upb/"*.so \
"$SITE_PKGS/google/protobuf/pyext/"*.so 2>/dev/null || true
```
--8<-- [end:build-wheel-from-source]
@@ -80,19 +144,20 @@ docker build -f docker/Dockerfile.s390x \
# Launch OpenAI server
docker run --rm \
--privileged true \
--security-opt seccomp=unconfined \
--cap-add SYS_NICE \
--shm-size 4g \
-p 8000:8000 \
-e VLLM_CPU_KVCACHE_SPACE=<KV cache space> \
-e VLLM_CPU_OMP_THREADS_BIND=<CPU cores for inference> \
vllm-cpu-env \
--model meta-llama/Llama-3.2-1B-Instruct \
--dtype float \
--dtype bfloat16 \
other vLLM OpenAI server arguments
```
!!! tip
An alternative of `--privileged true` is `--cap-add SYS_NICE --security-opt seccomp=unconfined`.
Alternatively, `--privileged=true` also works but is broader and not generally recommended.
--8<-- [end:build-image-from-source]
--8<-- [start:extra-information]
+3
View File
@@ -0,0 +1,3 @@
// Reo.Dev documentation tracking
// https://docs.reo.dev/integrations/input-sources/developer-insights/documentation
!function(){var e,t,n;e="d5c4337961ef0ac",t=function(){Reo.init({clientID:"d5c4337961ef0ac", enableThirdPartyTracking: true})},(n=document.createElement("script")).src="https://static.reo.dev/"+e+"/reo.js",n.defer=!0,n.onload=t,document.head.appendChild(n)}();
+1
View File
@@ -160,3 +160,4 @@ extra_javascript:
- https://unpkg.com/mathjax@3.2.2/es5/tex-mml-chtml.js
- mkdocs/javascript/edit_and_feedback.js
- mkdocs/javascript/slack_and_forum.js
- mkdocs/javascript/reo.js
+2 -2
View File
@@ -14,8 +14,8 @@ PyNvVideoCodec==2.0.4
# flashinfer-cubin is not on PyPI since 0.6.14; setup.py excludes it from
# install_requires so the published wheel does not carry an unresolvable pin
--extra-index-url https://flashinfer.ai/whl/
flashinfer-python==0.6.14
flashinfer-cubin==0.6.14
flashinfer-python==0.6.15.post1
flashinfer-cubin==0.6.15.post1
apache-tvm-ffi==0.1.10
tilelang==0.1.9
nvidia-cudnn-frontend>=1.19.1
+62
View File
@@ -214,3 +214,65 @@ def test_cache_config_hash_ignores_kv_cache_sizing_knobs():
base_hash = CacheConfig().compute_hash()
assert CacheConfig(kv_cache_memory_bytes=1 << 30).compute_hash() == base_hash
assert CacheConfig(gpu_memory_utilization=0.5).compute_hash() == base_hash
def test_envs_compile_factors_relocation_invariant(tmp_path):
"""Relocating HOME or the XDG roots must not change the compile-cache
env hash.
Location-derived env vars (VLLM_XLA_CACHE_PATH from XDG_CACHE_HOME,
VLLM_CONFIG_ROOT from XDG_CONFIG_HOME/HOME) carry no information about
compiled artifacts, only about where directories live. When they leak
into compile_factors(), a cache produced under one HOME/XDG layout
silently misses under another - which defeats copying or pre-baking a
compile cache into a container image.
"""
import os
import subprocess
import sys
code = """
import sys
import logging
logging.disable(logging.CRITICAL)
from vllm import envs
from vllm.config.utils import hash_factors
print(hash_factors(envs.compile_factors()))
"""
def hash_with(extra_env):
env = {**dict(os.environ), "VLLM_LOGGING_LEVEL": "ERROR"}
# Drop explicit overrides so the derived defaults are what is
# exercised, then apply the relocation under test.
for key in ("VLLM_XLA_CACHE_PATH", "VLLM_CONFIG_ROOT", "VLLM_CACHE_ROOT"):
env.pop(key, None)
env.update(extra_env)
result = subprocess.run(
[sys.executable, "-c", code],
capture_output=True,
text=True,
check=True,
env=env,
)
return result.stdout.strip()
xdg_cache = tmp_path / "relocated-xdg-cache"
xdg_config = tmp_path / "relocated-xdg-config"
new_home = tmp_path / "relocated-home"
for d in (xdg_cache, xdg_config, new_home):
d.mkdir()
base = hash_with({})
relocated_xdg = hash_with(
{"XDG_CACHE_HOME": str(xdg_cache), "XDG_CONFIG_HOME": str(xdg_config)}
)
relocated_home = hash_with({"HOME": str(new_home)})
assert relocated_xdg == base, (
"XDG_CACHE_HOME/XDG_CONFIG_HOME relocation changed the compile-cache "
"env hash - a location-only derived var is leaking into the key"
)
assert relocated_home == base, (
"HOME relocation changed the compile-cache env hash - a "
"location-only derived var is leaking into the key"
)
+30
View File
@@ -349,6 +349,20 @@ _T = TypeVar("_T", nn.Module, torch.Tensor, BatchEncoding, BatchFeature, dict)
_R = TypeVar("_R")
def _fix_v4_tied_weights_keys(model_cls: type) -> None:
"""Convert a v4 list-format _tied_weights_keys to the transformers v5 dict form."""
tied = getattr(model_cls, "_tied_weights_keys", None)
if not isinstance(tied, list) or not tied:
return
result = {
k: "model.embed_tokens.weight"
for k in tied
if "lm_head" in k and k.endswith(".weight")
}
if result:
setattr(model_cls, "_tied_weights_keys", result)
class HfRunner:
def get_default_device(self):
from vllm.platforms import current_platform
@@ -474,6 +488,22 @@ class HfRunner:
trust_remote_code=trust_remote_code,
)
else:
if trust_remote_code and hasattr(self.config, "auto_map"):
cls_ref = self.config.auto_map.get(auto_cls.__name__)
if cls_ref is not None:
from vllm.transformers_utils.dynamic_module import (
try_get_class_from_dynamic_module,
)
model_cls = try_get_class_from_dynamic_module(
cls_ref,
model_name,
trust_remote_code=trust_remote_code,
warn_on_fail=False,
)
if model_cls is not None:
_fix_v4_tied_weights_keys(model_cls)
model = cast(
nn.Module,
auto_cls.from_pretrained(
@@ -0,0 +1,76 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for check_stop_strings.
These are pure-function tests (no model / GPU). They pin down which stop
string is selected when several stop strings match within the text that was
appended in a single step -- which happens under speculative decoding, where
multiple tokens (and therefore multiple stop strings) can be appended at once.
"""
import pytest
from vllm.v1.engine.detokenizer import check_stop_strings
@pytest.mark.parametrize("stop", [["a", "is"], ["is", "a"]])
def test_earliest_completing_stop_wins_regardless_of_list_order(stop):
# " The user is a": " is a" (5 chars) was appended in one step. Both "is"
# (index 10) and " a" (index 13) land in the same window. "is" completes
# earlier in the text, so it must win over list order.
text = " The user is a"
new_char_count = len(" is a")
assert check_stop_strings(text, new_char_count, stop, include_in_output=False) == (
"is",
10,
)
@pytest.mark.parametrize("stop", [["a", "is"], ["is", "a"]])
def test_earliest_completing_stop_include_in_output(stop):
text = " The user is a"
new_char_count = len(" is a")
# Truncate to the end of "is" (index 12) -> " The user is".
assert check_stop_strings(text, new_char_count, stop, include_in_output=True) == (
"is",
12,
)
def test_completion_position_not_start_position():
# "b" starts later than "abc" but completes earlier, so it must win.
text = "abc"
assert check_stop_strings(
text, len(text), ["abc", "b"], include_in_output=False
) == ("b", 1)
@pytest.mark.parametrize(
"stop,expected",
[
(["ab", "b"], ("ab", 0)),
(["b", "ab"], ("b", 1)),
],
)
def test_ties_broken_by_list_order(stop, expected):
# "ab" and "b" both complete at index 2; list order decides the winner.
text = "ab"
assert (
check_stop_strings(text, len(text), stop, include_in_output=False) == expected
)
def test_single_stop_in_window_unchanged():
# The common case (one stop in the window) is unaffected by the change.
text = "hello world."
assert check_stop_strings(text, 1, ["."], include_in_output=False) == (".", 11)
# Stop completes at the very end -> no truncation needed (-1).
assert check_stop_strings(text, 1, ["."], include_in_output=True) == (".", -1)
def test_no_match_and_empty_inputs_return_none():
assert check_stop_strings("hello", 5, ["zzz"], include_in_output=False) is None
assert check_stop_strings("hello", 0, ["h"], include_in_output=False) is None
assert check_stop_strings("hello", 5, [], include_in_output=False) is None
+130
View File
@@ -348,6 +348,136 @@ def test_message_queue_busy_to_idle():
distributed_run(worker_fn_test_busy_to_idle, 4)
@pytest.mark.parametrize("should_warn", [False, True])
def test_reader_timeout_caps_indefinite_waits(should_warn):
with (
mock.patch(
"vllm.distributed.device_communicators.shm_broadcast."
"SHM_READER_RECHECK_INTERVAL_MS",
new=7,
),
mock.patch(
"vllm.distributed.device_communicators.shm_broadcast."
"VLLM_RINGBUFFER_WARNING_INTERVAL",
new=60,
),
):
timeout = MessageQueue.ReadTimeoutWithWarnings(
timeout=None, should_warn=should_warn
)
assert timeout.timeout_ms() == 7
def test_reader_rechecks_shm_after_idle_wait_timeout_without_notify():
writer = MessageQueue(
n_reader=1,
n_local_reader=1,
max_chunk_bytes=1024 * 1024,
max_chunks=1,
)
reader = MessageQueue.create_from_handle(writer.export_handle(), rank=0)
payload = 123
poll_started = threading.Event()
allow_timeout = threading.Event()
result = {}
def acquire_read_in_thread():
try:
with reader.acquire_read(indefinite=True) as buf:
result["value"] = buf[0]
except Exception as exc:
result["exc"] = exc
def poll_timeout(*, timeout: int | None = None):
poll_started.set()
assert allow_timeout.wait(timeout=5)
return []
try:
writer.wait_until_ready()
reader.wait_until_ready()
reader._spin_condition.last_read = 0
reader._spin_condition.busy_loop_s = 0
with (
mock.patch(
"vllm.distributed.device_communicators.shm_broadcast."
"SHM_READER_RECHECK_INTERVAL_MS",
new=50,
),
mock.patch(
"vllm.distributed.device_communicators.shm_broadcast."
"VLLM_RINGBUFFER_WARNING_INTERVAL",
new=60,
),
mock.patch.object(
reader._spin_condition.poller,
"poll",
side_effect=poll_timeout,
) as poll,
):
read_thread = threading.Thread(target=acquire_read_in_thread, daemon=True)
read_thread.start()
assert poll_started.wait(timeout=5)
with writer.acquire_write(timeout=0.1) as buf:
buf[0] = payload
allow_timeout.set()
read_thread.join(timeout=5)
assert not read_thread.is_alive()
poll.assert_called_once_with(timeout=50)
if "exc" in result:
raise result["exc"]
assert result["value"] == payload
with writer.buffer.get_metadata(0) as metadata_buffer:
assert metadata_buffer[0] == 1
assert metadata_buffer[1] == 1
finally:
writer.shutdown()
reader.shutdown()
for socket in (
writer.local_socket,
writer._spin_condition.local_notify_socket,
reader.local_socket,
reader._spin_condition.local_notify_socket,
reader._spin_condition.read_cancel_socket,
reader._spin_condition.write_cancel_socket,
):
socket.close(linger=0)
def test_acquire_read_releases_slot_when_reader_raises():
writer = MessageQueue(
n_reader=1,
n_local_reader=1,
max_chunk_bytes=1024 * 1024,
max_chunks=1,
)
reader = MessageQueue.create_from_handle(writer.export_handle(), rank=0)
try:
writer.wait_until_ready()
reader.wait_until_ready()
writer.enqueue({"payload": "first"})
with (
pytest.raises(RuntimeError, match="reader failed"),
reader.acquire_read(timeout=0.1),
):
raise RuntimeError("reader failed")
with writer.buffer.get_metadata(0) as metadata_buffer:
assert metadata_buffer[0] == 1
assert metadata_buffer[1] == 1
with writer.acquire_write(timeout=0.1) as buf:
buf[0] = 0
finally:
writer.shutdown()
reader.shutdown()
def test_warning_logs(caplog_vllm):
"""
Test that warning logs are emitted at VLLM_RINGBUFFER_WARNING_INTERVAL intervals
+15
View File
@@ -0,0 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
import torch
@pytest.fixture(autouse=True)
def reset_default_torch_device():
"""Several kernel tests call torch.set_default_device without restoring
it, which poisons subsequent tests in the same pytest run (e.g. CPU
tensors silently created on CUDA). Restore the factory default after
every test.
"""
yield
torch.set_default_device(None)
+14 -19
View File
@@ -33,12 +33,10 @@ from vllm.third_party.flash_linear_attention.ops.index import ( # noqa: E402
@pytest.mark.parametrize("num_seqs", [1, 5, 257])
@pytest.mark.parametrize("state_dtype", [torch.bfloat16, torch.float32])
def test_gdn_chunk_cutedsl_correctness(num_seqs: int, state_dtype: torch.dtype):
seq_lens = torch.randint(
1,
130,
(num_seqs,),
dtype=torch.int32,
)
rng_cpu = torch.Generator("cpu").manual_seed(1234)
rng = torch.Generator("cuda").manual_seed(2345)
seq_lens = torch.randint(1, 130, (num_seqs,), dtype=torch.int32, generator=rng_cpu)
cu_seqlens = torch.zeros(num_seqs + 1, device="cuda", dtype=torch.int32)
cu_seqlens[1:] = seq_lens.to(device="cuda").cumsum(0)
total_tokens = int(cu_seqlens[-1].item())
@@ -56,8 +54,9 @@ def test_gdn_chunk_cutedsl_correctness(num_seqs: int, state_dtype: torch.dtype):
head_k_dim,
device="cuda",
dtype=dtype,
generator=rng,
)
k = torch.randn_like(q)
k = torch.randn_like(q, generator=rng)
v = torch.randn(
1,
total_tokens,
@@ -65,29 +64,24 @@ def test_gdn_chunk_cutedsl_correctness(num_seqs: int, state_dtype: torch.dtype):
head_v_dim,
device="cuda",
dtype=dtype,
generator=rng,
)
q = F.normalize(q.float(), p=2, dim=-1).to(dtype)
k = F.normalize(k.float(), p=2, dim=-1).to(dtype)
a = torch.randn(
1,
total_tokens,
num_v_heads,
device="cuda",
dtype=dtype,
1, total_tokens, num_v_heads, device="cuda", dtype=dtype, generator=rng
)
b = torch.randn(
1,
total_tokens,
num_v_heads,
device="cuda",
dtype=dtype,
1, total_tokens, num_v_heads, device="cuda", dtype=dtype, generator=rng
)
# Match upstream FLA GatedDeltaNet synthetic initialization:
# https://github.com/fla-org/flash-linear-attention/blob/main/fla/layers/gated_deltanet.py
A = torch.empty(num_v_heads, device="cuda", dtype=torch.float32).uniform_(0, 16)
A = torch.empty(num_v_heads, device="cuda", dtype=torch.float32).uniform_(
0, 16, generator=rng
)
A_log = torch.log(A)
dt = torch.exp(
torch.rand(num_v_heads, device="cuda", dtype=torch.float32)
torch.rand(num_v_heads, device="cuda", dtype=torch.float32, generator=rng)
* (math.log(0.1) - math.log(0.001))
+ math.log(0.001)
)
@@ -105,6 +99,7 @@ def test_gdn_chunk_cutedsl_correctness(num_seqs: int, state_dtype: torch.dtype):
head_k_dim,
device="cuda",
dtype=state_dtype,
generator=rng,
)
* 0.05
)
+5 -3
View File
@@ -21,9 +21,9 @@ def test_gather_cache_oob():
seq_starts causes the block_table offset to read out of bounds.
"""
batch_size = 1
block_size = 64
entry_size = 128
# The kernel only supports the MLA entry sizes.
entry_size = 576
block_table = torch.tensor([[1, 2]], dtype=torch.int32, device="cuda")
@@ -34,6 +34,7 @@ def test_gather_cache_oob():
seq_len = 65
cu_seq_lens = torch.tensor([0, seq_len], dtype=torch.int32, device="cuda")
token_to_seq = torch.zeros(seq_len, dtype=torch.int32, device="cuda")
# src_cache: [num_blocks, block_size, entry_size]
num_blocks = 5
@@ -51,7 +52,8 @@ def test_gather_cache_oob():
dst,
block_table,
cu_seq_lens,
batch_size,
token_to_seq,
seq_len,
"auto", # kv_cache_dtype
scale,
seq_starts,
+21
View File
@@ -13,6 +13,7 @@ These tests cover:
"""
import math
from types import SimpleNamespace
import pytest
import torch
@@ -27,6 +28,7 @@ from vllm.models.deepseek_v4.common.ops.fused_compress_quant_cache import (
_fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn,
_launch_two_stage_sparse_attn_compressor,
)
from vllm.models.deepseek_v4.compressor import _get_c128_boundary
from vllm.platforms import current_platform
from .test_fused_indexer_q_rope_quant import quantize_to_mxfp4
@@ -58,6 +60,25 @@ def _ue8m0_reference(x: torch.Tensor, block_size: int, fp8_max: float):
return x_fp8, scales
@pytest.mark.parametrize(
("starts", "query_start_loc", "expected"),
[
([0], [0, 127], False),
([0], [0, 128], True),
([127], [0, 1], True),
([128], [0, 127], False),
([1, 255], [0, 1, 2], True),
(None, [0, 1], None),
],
)
def test_get_c128_boundary(starts, query_start_loc, expected):
metadata = SimpleNamespace(
_num_computed_tokens_cpu=None if starts is None else torch.tensor(starts),
query_start_loc_cpu=torch.tensor(query_start_loc),
)
assert _get_c128_boundary(metadata) is expected
# ── Test A: DeepseekV4 Attention path ──────────────────────────────────────────────
+1 -1
View File
@@ -264,7 +264,7 @@ def test_block_mask_direct_vs_slow_path():
device = torch.device("cuda")
vllm_config = create_vllm_config(
model_name="meta-llama/Meta-Llama-3-8B", block_size=16, max_model_len=1024
model_name="Qwen/Qwen2.5-1.5B-Instruct", block_size=16, max_model_len=1024
)
kv_cache_spec = create_standard_kv_cache_spec(vllm_config)
+13 -4
View File
@@ -726,14 +726,19 @@ def test_einsum_end_to_end(num_tokens, num_heads, n_groups):
This catches stride/layout bugs that only manifest when the einsum
kernel actually consumes the quantized activations.
"""
from deep_gemm.utils.math import ceil_div
from vllm.utils.deep_gemm import (
fp8_einsum,
is_deep_gemm_supported,
per_block_cast_to_fp8,
transform_sf_into_required_layout,
)
if not is_deep_gemm_supported():
pytest.skip("DeepGEMM not supported on this platform")
def ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
heads_per_group = num_heads // n_groups
d = heads_per_group * HEAD_DIM
o_lora_rank = 1024
@@ -809,8 +814,12 @@ def test_einsum_end_to_end(num_tokens, num_heads, n_groups):
# -- Checks --
# Einsum output: Triton and CUDA both rotate in fp32 now, so diffs
# come from fp32 ordering and UE8M0 boundary shifts only.
# Use relative diff (same metric as test_fp8_einsum.py).
from deep_gemm.testing import calc_diff
# Use relative diff (same metric as deep_gemm.testing.calc_diff).
def calc_diff(x, y):
x, y = x.double(), y.double()
denominator = (x * x + y * y).sum()
sim = 2 * (x * y).sum() / denominator
return 1 - sim
z_diff = calc_diff(z_fused, z_ref)
assert z_diff < 0.01, (
@@ -84,6 +84,22 @@ def norm_rope_ref(x, weight, positions, cos_sin_cache, eps):
return roped
def assert_fp8_cache_close(kv_cache, expected_kv_cache):
"""Compare two e4m3 caches allowing 1 ulp.
On CUDA the fused kernel quantizes K from its fp32 intermediate, while the
reshape_and_cache_flash reference quantizes the bf16-materialized value, so
rounding-boundary values may differ by one e4m3 code.
"""
byte_diff = (kv_cache.int() - expected_kv_cache.int()).abs()
got = kv_cache.view(torch.float8_e4m3fn).float()
exp = expected_kv_cache.view(torch.float8_e4m3fn).float()
ok = (byte_diff <= 1) | ((got == 0) & (exp == 0))
assert bool(ok.all()), (
f"fp8 cache differs by more than 1 ulp in {int((~ok).sum())} elements"
)
# ── Test 1: dense mode (norm+rope only, no index, no insert) ─────────────────
@@ -265,7 +281,7 @@ def test_sparse_full(num_tokens, block_size, kv_cache_dtype):
scale,
scale,
)
torch.testing.assert_close(kv_cache, expected_kv_cache, rtol=0, atol=0)
assert_fp8_cache_close(kv_cache, expected_kv_cache)
else:
for t in range(num_tokens):
s = slot_mapping[t].item()
@@ -383,7 +399,7 @@ def test_sparse_skip_index_branch(num_tokens, block_size, kv_cache_dtype):
scale,
scale,
)
torch.testing.assert_close(kv_cache, expected_kv_cache, rtol=0, atol=0)
assert_fp8_cache_close(kv_cache, expected_kv_cache)
else:
k_ref_h = k_ref.view(num_tokens, num_kv_heads, HEAD_DIM)
v_ref_h = v_in.view(num_tokens, num_kv_heads, HEAD_DIM)
@@ -41,11 +41,12 @@ def test_fused_recurrent_packed_decode_matches_reference(
A_log = torch.randn((HV,), device=device, dtype=dtype)
dt_bias = torch.randn((HV,), device=device, dtype=dtype)
# Continuous batching indices (include PAD_SLOT_ID=-1 cases).
ssm_state_indices = torch.arange(B, device=device, dtype=torch.int32)
# Continuous batching indices (include PAD_SLOT_ID=-1 cases). Index 0 is
# reserved as NULL_BLOCK_ID (CUDA graph padding), so valid slots start at 1.
ssm_state_indices = torch.arange(1, B + 1, device=device, dtype=torch.int32)
ssm_state_indices[-3:] = -1
state0 = torch.randn((B, HV, V, K), device=device, dtype=dtype)
state0 = torch.randn((B + 1, HV, V, K), device=device, dtype=dtype)
state_ref = state0.clone()
state_packed = state0.clone()
@@ -94,5 +95,8 @@ def test_fused_recurrent_packed_decode_matches_reference(
atol = 2e-2 if dtype != torch.float32 else 1e-4
rtol = 1e-2 if dtype != torch.float32 else 1e-4
torch.testing.assert_close(out_packed, out_ref, rtol=rtol, atol=atol)
# Output rows for PAD_SLOT_ID entries are never written (uninitialized in
# both paths), so compare only the valid rows.
valid = ssm_state_indices > 0
torch.testing.assert_close(out_packed[valid], out_ref[valid], rtol=rtol, atol=atol)
torch.testing.assert_close(state_packed, state_ref, rtol=rtol, atol=atol)
@@ -58,10 +58,12 @@ def test_fused_sigmoid_gating_delta_rule_update_non_spec(
dt_bias = torch.rand(num_v_heads // tp_size, dtype=dtype)
a = torch.rand(num_tokens, num_v_heads, dtype=dtype)
b = torch.rand(num_tokens, num_v_heads, dtype=dtype)
# Entry 0 is reserved as NULL_BLOCK_ID (CUDA graph padding), so valid
# state indices start at 1.
ssm_state = torch.rand(
total_entries, num_v_heads, head_k_dim, head_v_dim, dtype=dtype
total_entries + 1, num_v_heads, head_k_dim, head_v_dim, dtype=dtype
)
state_indices = torch.randperm(total_entries, dtype=torch.int32)[:num_tokens]
state_indices = (torch.randperm(total_entries, dtype=torch.int32) + 1)[:num_tokens]
cu_seqlens = torch.arange(0, num_tokens + 1, dtype=torch.int32)
beta = b.sigmoid()
@@ -144,13 +146,14 @@ def test_fused_sigmoid_gating_delta_rule_update_spec(
dt_bias = torch.rand(num_v_heads // tp_size, dtype=dtype)
a = torch.rand(num_tokens, num_v_heads, dtype=dtype)
b = torch.rand(num_tokens, num_v_heads, dtype=dtype)
# Entry 0 is reserved as NULL_BLOCK_ID (CUDA graph padding), so valid
# state indices start at 1.
ssm_state = torch.rand(
total_entries, num_v_heads, head_k_dim, head_v_dim, dtype=dtype
total_entries + 1, num_v_heads, head_k_dim, head_v_dim, dtype=dtype
)
state_indices = torch.randperm(
total_entries,
dtype=torch.int32,
)[:num_tokens].view(num_reqs, num_speculative_tokens + 1)
state_indices = (torch.randperm(total_entries, dtype=torch.int32) + 1)[
:num_tokens
].view(num_reqs, num_speculative_tokens + 1)
num_accepted_tokens = torch.randint(
1, num_speculative_tokens + 1, (num_reqs,), dtype=torch.int32
)
@@ -6,6 +6,7 @@ import pytest
from vllm.config.compilation import CompilationConfig, CUDAGraphMode
from vllm.models.inkling.common.mm_preprocess import InklingMultiModalDataParser
from vllm.models.inkling.common.towers import plan_out_scales
from vllm.models.inkling.configs import (
InklingAudioConfig,
InklingModelConfig,
@@ -17,6 +18,22 @@ from vllm.models.inkling.nvidia.sconv_swa_attn import (
from vllm.v1.attention.backend import AttentionCGSupport
def test_vision_scale_plan_matches_released_config():
assert plan_out_scales(2, 40, 4) == [
(1, 1, 1, 3),
(1, 5, 5, 128),
(1, 10, 10, 320),
(1, 40, 40, 4800),
(2, 40, 40, 9600),
]
def test_vision_scale_plan_breaks_assignment_ties_in_order():
reductions = [np.prod(scale[:-1]) for scale in plan_out_scales(2, 52, 4)]
assert reductions == sorted(set(reductions))
@pytest.mark.parametrize(
("config_cls", "kwargs", "missing"),
[
-7
View File
@@ -487,13 +487,6 @@ _TEXT_GENERATION_EXAMPLE_MODELS = {
"Plamo2ForCausalLM": _HfExamplesInfo(
"pfnet/plamo-2-1b",
trust_remote_code=True,
max_transformers_version="4.57",
transformers_version_reason={
"hf": (
"Custom model code uses `_tied_weight_keys: list[str]` but "
"Transformers v5 now expects `_tied_weight_keys: dict[str, str]`"
)
},
),
"Plamo3ForCausalLM": _HfExamplesInfo(
"pfnet/plamo-3-nict-2b-base",
+42
View File
@@ -920,3 +920,45 @@ class TestDelegatingParserLargeDelta:
assert output.tool_calls[0]["name"] == "get_weather"
args = json.loads(output.tool_calls[0]["arguments"])
assert args == {"location": "Berlin"}
@pytest.mark.parametrize(
"chunk_size",
[1, 2, 3, 5, None],
ids=lambda c: f"chunk={c}",
)
def test_eos_not_leaked_when_reasoning_never_ends(self, chunk_size):
"""EOS must not leak into reasoning_content when the model never
emits </think> (generation ends while still in REASONING state)."""
eos_text = "<end▁of▁sentence>"
eos_id = 128801
vocab = {
**_DSV4_FULL_VOCAB,
eos_text: eos_id,
}
reasoning_text = "Good morning! How can I help you today?"
tokens: list[tuple[int, str]] = []
tid = 100
for word in reasoning_text.split(" "):
prefix = " " if tokens else ""
tokens.append((tid, prefix + word))
tid += 1
tokens.append((eos_id, eos_text))
tokenizer = MockTokenizer(vocab=vocab, tokens=tokens)
parser = _DeepSeekV4Delegating(
tokenizer,
chat_template_kwargs={"thinking": True},
)
deltas = replay_streaming(
parser,
tokens,
chunk_size=chunk_size,
finished_on_last=True,
)
output = collect_output(deltas)
assert reasoning_text in output.reasoning
assert eos_text not in output.reasoning
assert output.content == ""
assert output.tool_calls == []
+29 -7
View File
@@ -1624,19 +1624,41 @@ class TestDropSpecialTokens:
assert delta is not None
assert "<bos>" in delta.reasoning
def test_drops_suppressed_with_skip_tool_parsing(self):
"""When skip_tool_parsing is active, drop tokens are preserved
as content so a later tool-call pass can see them."""
def test_drops_applied_with_skip_tool_parsing(self):
"""Drop tokens are always dropped, even with skip_tool_parsing.
DROP_TERMINALs have no transitions by construction, so no parser
pass can use them."""
for initial_state in (ParserState.REASONING, ParserState.CONTENT):
engine = _make_engine(
vocab=_DROP_VOCAB,
special_tokens=list(_DROP_VOCAB.keys()),
)
engine._engine.skip_tool_parsing = True
engine._engine.reset(initial_state=initial_state)
events = engine._engine.feed("hello<bos>world", [72, 204, 73])
delta = engine._events_to_delta(events)
assert delta is not None
output = (delta.reasoning or "") + (delta.content or "")
assert "<bos>" not in output, f"<bos> leaked in state {initial_state}"
def test_transitions_unaffected_by_drop_in_reasoning_with_skip_tool_parsing(self):
"""With skip_tool_parsing in REASONING state, drop tokens are
removed but configured terminals still fire their transitions."""
engine = _make_engine(
vocab=_DROP_VOCAB,
special_tokens=list(_DROP_VOCAB.keys()),
)
engine._engine.skip_tool_parsing = True
engine._engine.reset()
events = engine._engine.feed("hello<bos>world", [72, 204, 73])
delta = engine._events_to_delta(events)
assert delta is not None
assert "<bos>" in delta.reasoning
events = engine._engine.feed("thought<bos></think>answer", [72, 204, 201, 73])
types = [e.type for e in events]
assert EventType.REASONING_CHUNK in types
assert EventType.REASONING_END in types
assert EventType.TEXT_CHUNK in types
reasoning_text = "".join(
e.value for e in events if e.type == EventType.REASONING_CHUNK
)
assert "<bos>" not in reasoning_text
def test_drops_in_tool_args_state(self):
"""Drop tokens in TOOL_ARGS state are silently discarded."""
@@ -1,13 +1,55 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from types import SimpleNamespace
import torch
from vllm.model_executor.layers.fused_moe.oracle.nvfp4 import NvFp4MoeBackend
from vllm.model_executor.layers.quantization.utils import flashinfer_fp4_moe
from vllm.model_executor.layers.quantization.utils.flashinfer_fp4_moe import (
prepare_nvfp4_moe_layer_for_fi_or_cutlass,
)
from vllm.model_executor.layers.quantization.utils.flashinfer_utils import (
align_trtllm_fp4_moe_hidden_dim_for_fi,
)
def test_shared_nvfp4_input_scales_have_writable_storage(monkeypatch):
monkeypatch.setattr(flashinfer_fp4_moe, "swizzle_blockscale", lambda x: x)
num_experts = 3
layer = SimpleNamespace(activation=SimpleNamespace(is_gated=False))
w13 = torch.zeros((num_experts, 2, 1), dtype=torch.uint8)
w2 = torch.zeros((num_experts, 2, 1), dtype=torch.uint8)
w13_scale = torch.zeros((num_experts, 2, 1), dtype=torch.float8_e4m3fn)
w2_scale = torch.zeros((num_experts, 2, 1), dtype=torch.float8_e4m3fn)
weight_scale = torch.ones(num_experts)
outputs = prepare_nvfp4_moe_layer_for_fi_or_cutlass(
backend=NvFp4MoeBackend.FLASHINFER_CUTLASS,
layer=layer,
w13=w13,
w13_scale=w13_scale,
w13_scale_2=weight_scale,
a13_scale=torch.tensor([1.0, 2.0, 3.0]),
w2=w2,
w2_scale=w2_scale,
w2_scale_2=weight_scale,
a2_scale=torch.tensor([4.0, 5.0, 6.0]),
is_act_and_mul=False,
)
a13_scale, a2_scale = outputs[3], outputs[7]
torch.testing.assert_close(a13_scale, torch.full((num_experts,), 3.0))
torch.testing.assert_close(a2_scale, torch.full((num_experts,), 6.0))
distinct_values = torch.arange(num_experts, dtype=torch.float32)
a13_scale.copy_(distinct_values)
a2_scale.copy_(distinct_values)
torch.testing.assert_close(a13_scale, distinct_values)
torch.testing.assert_close(a2_scale, distinct_values)
def test_align_trtllm_fp4_moe_hidden_dim_noop():
w13 = torch.arange(2 * 8 * 256, dtype=torch.uint8).reshape(2, 8, 256)
w13_scale = torch.arange(2 * 8 * 32, dtype=torch.uint8).reshape(2, 8, 32)
+3 -1
View File
@@ -9,6 +9,7 @@ from vllm.v1.core.sched.async_scheduler import AsyncScheduler
from vllm.v1.core.sched.output import CachedRequestData, SchedulerOutput
from vllm.v1.outputs import ModelRunnerOutput
from vllm.v1.request import RequestStatus
from vllm.v1.structured_output import StructuredOutputGrammar
from vllm.v1.utils import ConstantList
from .utils import create_requests, create_scheduler
@@ -262,7 +263,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler = object.__new__(AsyncScheduler)
request = create_requests(num_requests=1, num_tokens=1)[0]
request.structured_output_request = Mock()
request.structured_output_request.grammar = Mock()
request.structured_output_request.grammar = Mock(spec=StructuredOutputGrammar)
request.structured_output_request.grammar.accept_tokens.return_value = False
request.status = RequestStatus.RUNNING
request.num_computed_tokens = request.num_tokens
@@ -284,6 +285,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler.kv_event_publisher = Mock()
scheduler.finished_req_ids = set()
scheduler.finished_req_ids_dict = None
scheduler.grammar_compile_error_reqs = set()
scheduler.vllm_config = Mock()
scheduler.vllm_config.model_config.enable_return_routed_experts = False
scheduler.enable_return_routed_experts = False
+168 -20
View File
@@ -9,6 +9,7 @@ import torch
from vllm.v1.core.kv_cache_utils import (
_get_kv_cache_config_packed,
_get_kv_cache_groups_uniform_groups,
get_kv_cache_config_from_groups,
)
from vllm.v1.kv_cache_interface import (
@@ -16,6 +17,7 @@ from vllm.v1.kv_cache_interface import (
KVCacheGroupSpec,
KVCacheTensor,
MLAAttentionSpec,
SlidingWindowMLASpec,
SlidingWindowSpec,
UniformTypeKVCacheSpecs,
)
@@ -109,7 +111,132 @@ def _page_sizes_by_layer(
return page_sizes
def _packing_by_layer(
tensors: list[KVCacheTensor],
) -> dict[str, tuple[int, int]]:
return {
layer_name: (tensor.offset, tensor.block_stride)
for tensor in tensors
for layer_name in tensor.shared_by
}
def _make_views(
groups: list[KVCacheGroupSpec],
num_blocks: int,
tensors: list[KVCacheTensor],
) -> dict[str, torch.Tensor]:
page_sizes = _page_sizes_by_layer(groups)
packing = _packing_by_layer(tensors)
backing = torch.zeros(tensors[0].size, dtype=torch.uint8)
return {
layer_name: torch.as_strided(
backing,
size=(num_blocks, page_size),
stride=(packing[layer_name][1], 1),
storage_offset=packing[layer_name][0],
)
for layer_name, page_size in page_sizes.items()
}
def _make_page_group(prefix: str, page_sizes: list[int]) -> KVCacheGroupSpec:
specs = {
f"{prefix}.{i}": MagicMock(page_size_bytes=page_size)
for i, page_size in enumerate(page_sizes)
}
return KVCacheGroupSpec(
layer_names=list(specs),
kv_cache_spec=UniformTypeKVCacheSpecs(block_size=256, kv_cache_specs=specs),
)
class TestInterleavedPacking:
def test_compact_cache_overlays_fp32_state_group(self):
full_specs = {}
state_specs = {}
for i in range(2):
full_specs[f"mla.{i}"] = MLAAttentionSpec(
block_size=256,
num_kv_heads=1,
head_size=512,
dtype=torch.uint8,
page_size_padded=32768,
indexes_kv_by_block_stride=True,
compress_ratio=4,
)
full_specs[f"indexer.{i}"] = MLAAttentionSpec(
block_size=256,
num_kv_heads=1,
head_size=68,
dtype=torch.uint8,
page_size_padded=4608,
compress_ratio=4,
)
state_specs[f"mla_state.{i}"] = SlidingWindowMLASpec(
block_size=4,
num_kv_heads=1,
head_size=2048,
dtype=torch.float32,
sliding_window=8,
indexes_kv_by_block_stride=True,
)
state_specs[f"indexer_state.{i}"] = SlidingWindowMLASpec(
block_size=4,
num_kv_heads=1,
head_size=512,
dtype=torch.float32,
sliding_window=8,
indexes_kv_by_block_stride=True,
)
grouped_specs = [
UniformTypeKVCacheSpecs(block_size=256, kv_cache_specs=full_specs),
UniformTypeKVCacheSpecs(block_size=4, kv_cache_specs=state_specs),
]
groups = _get_kv_cache_groups_uniform_groups(grouped_specs)
assert len(groups) == 2
assert {full_specs[f"indexer.{i}"].page_size_bytes for i in range(2)} == {4608}
assert {full_specs[f"indexer.{i}"].real_page_size_bytes for i in range(2)} == {
4352
}
assert {
state_specs[f"indexer_state.{i}"].page_size_bytes for i in range(2)
} == {8192}
full_group_bytes = 2 * (32768 + 4608)
state_group_bytes = 2 * (32768 + 8192)
bytes_per_block = max(full_group_bytes, state_group_bytes)
num_blocks, tensors = _get_kv_cache_config_packed(
_mock_vllm_config(), groups, bytes_per_block * 32
)
assert num_blocks == 32
assert {tensor.block_stride for tensor in tensors} == {bytes_per_block}
packing = _packing_by_layer(tensors)
assert packing["mla.0"][0] == packing["mla_state.0"][0] == 0
assert packing["indexer.0"][0] == 32768
assert packing["indexer_state.0"][0] == 32768
def test_deepseek_v4_pro_stride(self):
groups = [
_make_page_group("full", [32768, 4608] * 30 + [1024] * 31),
_make_page_group("c4_state", [32768, 8192] * 30),
_make_page_group("c128_state", [32768] * 31),
_make_page_group("swa.0", [32768] * 31),
_make_page_group("swa.1", [32768] * 30),
]
expected_stride = 1_228_800
num_blocks, tensors = _get_kv_cache_config_packed(
_mock_vllm_config(), groups, expected_stride * 32
)
assert num_blocks == 32
assert {tensor.block_stride for tensor in tensors} == {expected_stride}
assert {tensor.size for tensor in tensors} == {expected_stride * 32}
def test_all_tensors_have_block_stride(self):
_, tensors = _run()
for t in tensors:
@@ -122,9 +249,30 @@ class TestInterleavedPacking:
assert sizes.pop() > 0
def test_offsets_within_one_block(self):
_, tensors = _run()
for t in tensors:
assert t.offset < t.block_stride
groups = _make_groups(n_c4=3, n_c128=2, n_swa=5)
_, tensors = _get_kv_cache_config_packed(
_mock_vllm_config(), groups, 100 * 1024 * 1024
)
page_sizes = _page_sizes_by_layer(groups)
packing = _packing_by_layer(tensors)
for layer_name, page_size in page_sizes.items():
offset, block_stride = packing[layer_name]
assert offset + page_size <= block_stride
def test_layouts_are_disjoint_within_each_group(self):
groups = _make_groups(n_c4=3, n_c128=2, n_swa=5)
_, tensors = _get_kv_cache_config_packed(
_mock_vllm_config(), groups, 100 * 1024 * 1024
)
page_sizes = _page_sizes_by_layer(groups)
packing = _packing_by_layer(tensors)
for group in groups:
ranges = sorted(
(packing[name][0], packing[name][0] + page_sizes[name])
for name in group.layer_names
)
assert all(left[1] <= right[0] for left, right in zip(ranges, ranges[1:]))
def test_all_layers_accounted_for(self):
n_c4, n_c128, n_swa = 5, 4, 7
@@ -135,29 +283,29 @@ class TestInterleavedPacking:
expected = n_c4 * 2 + n_c128 + n_swa
assert len(all_names) == expected
def test_strided_views_are_independent(self):
def test_group_owned_blocks_do_not_alias(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_packed(
_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)
views = _make_views(groups, num_blocks, tensors)
for i, v in enumerate(views):
v.fill_(i + 1)
expected = {}
value = 1
for block_id, group in enumerate(groups):
for layer_name in group.layer_names:
views[layer_name][block_id].fill_(value)
expected[layer_name] = (block_id, value)
value += 1
for i, v in enumerate(views):
assert (v == i + 1).all(), f"View {i} was corrupted"
for layer_name, (block_id, value) in expected.items():
assert (views[layer_name][block_id] == value).all()
# Once the first group releases its block, another group may reuse it.
for layer_name in groups[1].layer_names:
views[layer_name][0].fill_(255)
for layer_name in groups[1].layer_names:
assert (views[layer_name][0] == 255).all()
def test_hma_attention_groups_keep_default_backing(self):
full = _make_full_spec()
+253 -2
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import dataclasses
from concurrent.futures import Future
from unittest.mock import Mock
import pytest
@@ -25,18 +26,21 @@ from vllm.multimodal.inputs import (
from vllm.sampling_params import SamplingParams, StructuredOutputsParams
from vllm.utils.hashing import sha256
from vllm.v1.core.encoder_cache_manager import EncoderCacheManager
from vllm.v1.core.kv_cache_coordinator import HybridKVCacheCoordinator
from vllm.v1.core.kv_cache_utils import get_request_block_hasher, init_none_hash
from vllm.v1.core.sched.output import CachedRequestData, SchedulerOutput
from vllm.v1.core.sched.scheduler import Scheduler
from vllm.v1.core.single_type_kv_cache_manager import register_all_kvcache_specs
from vllm.v1.engine import FinishReason
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheConfig,
KVCacheGroupSpec,
MambaSpec,
)
from vllm.v1.outputs import DraftTokenIds, KVConnectorOutput, ModelRunnerOutput
from vllm.v1.request import Request, RequestStatus
from vllm.v1.structured_output import StructuredOutputManager
from vllm.v1.structured_output import StructuredOutputGrammar, StructuredOutputManager
from .utils import EOS_TOKEN_ID, create_requests, create_scheduler, mock_kv
@@ -3141,6 +3145,58 @@ def test_schedule_skip_tokenizer_init_structured_output_request():
assert len(scheduler.skipped_waiting) == 1
@pytest.mark.parametrize("async_grammar", [True, False])
def test_grammar_compile_error_finishes_only_request(async_grammar: bool):
scheduler = create_scheduler()
manager = scheduler.structured_output_manager
manager.backend = Mock()
manager.backend.compile_grammar.side_effect = RuntimeError(
"forced FSM compilation error"
)
manager._use_async_grammar_compilation = async_grammar
sampling_params = SamplingParams(
max_tokens=16,
structured_outputs=StructuredOutputsParams(json='{"type": "object"}'),
)
sampling_params.update_from_generation_config({}, EOS_TOKEN_ID)
request = Request(
request_id="grammar-error",
prompt_token_ids=[0, 1],
sampling_params=sampling_params,
pooling_params=None,
)
manager.grammar_init(request)
assert request.structured_output_request is not None
grammar_future = request.structured_output_request._grammar
assert isinstance(grammar_future, Future)
assert isinstance(grammar_future.exception(timeout=5), RuntimeError)
scheduler.add_request(request)
scheduler_output = scheduler.schedule()
assert not scheduler_output.num_scheduled_tokens
engine_core_outputs = scheduler.update_from_output(
scheduler_output,
ModelRunnerOutput(req_ids=[], req_id_to_index={}),
)
assert request.status == RequestStatus.FINISHED_ERROR
assert request.request_id not in scheduler.requests
output = engine_core_outputs[0].outputs[0]
assert output.request_id == request.request_id
assert output.finish_reason == FinishReason.ERROR
assert output.stop_reason is None
healthy_request = create_requests(num_requests=1, req_ids=["healthy-request"])[0]
scheduler.add_request(healthy_request)
next_output = scheduler.schedule()
assert [req.req_id for req in next_output.scheduled_new_reqs] == [
healthy_request.request_id
]
def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler = object.__new__(Scheduler)
sampling_params = SamplingParams(ignore_eos=True, max_tokens=4)
@@ -3154,7 +3210,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
pooling_params=None,
)
request.structured_output_request = Mock()
request.structured_output_request.grammar = Mock()
request.structured_output_request.grammar = Mock(spec=StructuredOutputGrammar)
request.structured_output_request.grammar.accept_tokens.return_value = False
request.status = RequestStatus.RUNNING
request.num_computed_tokens = request.num_tokens
@@ -3175,6 +3231,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler.kv_event_publisher = Mock()
scheduler.finished_req_ids = set()
scheduler.finished_req_ids_dict = None
scheduler.grammar_compile_error_reqs = set()
scheduler.vllm_config = Mock()
scheduler.vllm_config.model_config.enable_return_routed_experts = False
scheduler.enable_return_routed_experts = False
@@ -5464,3 +5521,197 @@ def test_async_load_reservation_prevents_wedge_e2e():
assert b.status == RequestStatus.WAITING
assert b.num_preemptions == 0
assert b.request_id not in req_to_blocks
def _create_hybrid_mamba_connector_scheduler(
matched_tokens: int,
block_size: int = 16,
num_blocks: int = 100,
) -> Scheduler:
"""FA + Mamba ("all" cache mode) scheduler with a MockKVConnector."""
model_config = ModelConfig(
model="facebook/opt-125m",
trust_remote_code=True,
dtype="float16",
seed=42,
skip_tokenizer_init=True,
)
vllm_config = VllmConfig(
scheduler_config=SchedulerConfig(
max_num_seqs=4,
max_num_batched_tokens=8192,
max_model_len=8192,
enable_chunked_prefill=True,
is_encoder_decoder=False,
watermark=0.0,
),
model_config=model_config,
cache_config=CacheConfig(
block_size=block_size,
enable_prefix_caching=True,
mamba_cache_mode="all",
),
kv_transfer_config=KVTransferConfig(
kv_connector="MockKVConnector",
kv_role="kv_both",
kv_connector_extra_config={
"matched_tokens": matched_tokens,
"is_async": False,
},
),
)
vllm_config.cache_config.num_gpu_blocks = num_blocks
kv_cache_config = KVCacheConfig(
num_blocks=num_blocks,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["fa"],
FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["mamba"],
MambaSpec(
block_size=block_size,
shapes=((1, 1),),
dtypes=(torch.float32,),
mamba_cache_mode="all",
),
),
],
)
register_all_kvcache_specs(vllm_config)
return Scheduler(
vllm_config=vllm_config,
kv_cache_config=kv_cache_config,
structured_output_manager=StructuredOutputManager(vllm_config),
block_size=block_size,
hash_block_size=block_size,
log_stats=True,
)
@pytest.mark.parametrize(
"matched_tokens,expected_num_computed",
[
# No external hit: resume on the deepest locally-consistent boundary
# (block 0's state survives for both groups).
(0, 16),
# One external block on top of the reconciled local boundary.
(16, 32),
],
)
def test_hybrid_per_group_hit_divergence_with_connector(
matched_tokens: int, expected_num_computed: int
):
"""Per-group prefix hits can diverge for hybrid models with a connector
(#46453): under block pressure the FA prefix tail is evicted while a
deeper Mamba state block survives. The scheduler must not report the
deeper hit as locally computed (evicted FA blocks are not resident ->
engine crash / dirty KV); it falls back to the reconciled boundary that
every group is consistent at.
"""
block_size = 16
scheduler = _create_hybrid_mamba_connector_scheduler(matched_tokens)
manager = scheduler.kv_cache_manager
assert isinstance(manager.coordinator, HybridKVCacheCoordinator)
# Seed a 4-block prefix so both groups cache all four boundaries
# (mamba cache mode "all" caches every block's state densely).
[fill] = create_requests(
num_requests=1,
num_tokens=4 * block_size,
max_tokens=1,
same_prompt=True,
block_size=block_size,
req_ids=["fill"],
)
computed_blocks, num_computed, _ = manager.get_computed_blocks(fill)
blocks = manager.allocate_slots(
fill, fill.num_tokens, num_computed, computed_blocks
)
fa_ids = [b.block_id for b in blocks.blocks[0]]
mamba_ids = [b.block_id for b in blocks.blocks[1]]
manager.free(fill)
# Evict the FA tail and the middle mamba states; block 0 (both groups)
# and the deep mamba state at block 3 survive.
manager.block_pool.evict_blocks({fa_ids[2], fa_ids[3], mamba_ids[1], mamba_ids[2]})
# A replay of the prefix plus one extra block now sees diverged
# per-group hits: FA stops at the evicted tail, while the mamba lookup
# finds the deeper surviving state.
[replay] = create_requests(
num_requests=1,
num_tokens=5 * block_size,
max_tokens=1,
same_prompt=True,
block_size=block_size,
req_ids=["replay"],
)
_, per_group_hits = manager.coordinator.find_longest_cache_hit_per_group(
replay.block_hashes, replay.num_tokens - 1
)
assert per_group_hits == (2 * block_size, 4 * block_size) # diverged
scheduler.add_request(replay)
output = scheduler.schedule()
num_scheduled = output.num_scheduled_tokens[replay.request_id]
assert replay.num_tokens - num_scheduled == expected_num_computed
def test_hybrid_per_group_hit_divergence_fa_deeper_no_external():
"""The opposite divergence: the FA prefix survives deeper than the Mamba
state and the connector supplies nothing (ext == 0). Reporting the deep FA
hit as locally computed would resume with no valid Mamba state at that
boundary (silent bad output). The scheduler must fall back to the
convergent boundary that every group agrees on (block 0's surviving state).
"""
block_size = 16
scheduler = _create_hybrid_mamba_connector_scheduler(matched_tokens=0)
manager = scheduler.kv_cache_manager
assert isinstance(manager.coordinator, HybridKVCacheCoordinator)
# Seed a 4-block prefix in both groups.
[fill] = create_requests(
num_requests=1,
num_tokens=4 * block_size,
max_tokens=1,
same_prompt=True,
block_size=block_size,
req_ids=["fill"],
)
computed_blocks, num_computed, _ = manager.get_computed_blocks(fill)
blocks = manager.allocate_slots(
fill, fill.num_tokens, num_computed, computed_blocks
)
mamba_ids = [b.block_id for b in blocks.blocks[1]]
manager.free(fill)
# Keep all FA blocks; evict every mamba state but block 0. FA reaches 4
# blocks, the mamba hit only reaches 1 -> diverged (FA > Mamba).
manager.block_pool.evict_blocks({mamba_ids[1], mamba_ids[2], mamba_ids[3]})
[replay] = create_requests(
num_requests=1,
num_tokens=5 * block_size,
max_tokens=1,
same_prompt=True,
block_size=block_size,
req_ids=["replay"],
)
_, per_group_hits = manager.coordinator.find_longest_cache_hit_per_group(
replay.block_hashes, replay.num_tokens - 1
)
assert per_group_hits == (4 * block_size, 1 * block_size) # FA deeper
scheduler.add_request(replay)
output = scheduler.schedule()
num_scheduled = output.num_scheduled_tokens[replay.request_id]
# Must resume at the convergent boundary (block 0), not the deep FA hit.
assert replay.num_tokens - num_scheduled == block_size
@@ -410,7 +410,7 @@ def test_lookup_key_client_lookup_prepends_typed_tag():
# Blocking lookup (non_block defaults to False) runs on the executor and
# returns the resolved hit length.
assert client.lookup("req0", token_len=128, block_hashes=[]) == 5
assert client.lookup("req0", num_tokens=128, block_hashes=[]) == 5
sent_frames = fake_socket.send_multipart.call_args[0][0]
assert sent_frames[0] == protocol.LOOKUP_MSG
@@ -439,11 +439,11 @@ def test_lookup_key_client_reset_uses_typed_protocol():
assert client.reset() is False
def _poll_lookup(client, req_id, token_len=128, block_hashes=(), timeout=5.0):
def _poll_lookup(client, req_id, num_tokens=128, block_hashes=(), timeout=5.0):
"""Drive non-blocking lookup until the executor completes it."""
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
result = client.lookup(req_id, token_len, list(block_hashes), non_block=True)
result = client.lookup(req_id, num_tokens, list(block_hashes), non_block=True)
if result is not None:
return result
time.sleep(0.005)
@@ -240,7 +240,10 @@ def test_e2e_swa_plus_full_save_then_lookup_hits():
worker.store = store
# Both groups stored all 4 blocks -> full hit.
assert worker.lookup(token_len=64, block_hashes=hs) == 64
assert worker.lookup(num_tokens=65, block_hashes=hs) == 64
# Exact-multiple prompt: the full hit is re-derived one block lower,
# where both groups' stored blocks still cover the SWA window.
assert worker.lookup(num_tokens=64, block_hashes=hs) == 48
# Evict SWA's first two blocks (outside its window of 32 tokens = 2 blocks).
swa_keys_outside_window = [
@@ -253,7 +256,12 @@ def test_e2e_swa_plus_full_save_then_lookup_hits():
# SWA window=32 -> only last 2 blocks must be present in SWA group.
# Full has all 4. Coordinator should still return 64.
assert worker.lookup(token_len=64, block_hashes=hs) == 64
assert worker.lookup(num_tokens=65, block_hashes=hs) == 64
# Exact-multiple prompt after eviction: the boundary one block lower
# needs SWA block 1, which is gone — no usable stored boundary remains
# (the pre-fix arithmetic clamp would have returned 48 and livelocked
# on load failure -> recompute -> same lookup).
assert worker.lookup(num_tokens=64, block_hashes=hs) == 0
def test_recv_skips_swa_blocks_before_window():
@@ -0,0 +1,90 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for ChunkedTokenDatabase.prepare_values."""
import random
import pytest
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import (
ChunkedTokenDatabase,
KeyMetadata,
)
from vllm.utils.math_utils import cdiv
BLOCK_SIZE = 128
def _reference_prepare_value(
db: ChunkedTokenDatabase, start: int, end: int, block_ids: list[int]
) -> tuple[list[int], list[int], int]:
"""Compute a token range with the original scalar implementation."""
addr_list = []
size_list = []
block_id = block_ids[start // db.block_size]
length = len(db.block_len)
for index, base_addr in enumerate(db.kv_caches_base_addr):
addr = base_addr + block_id * db.block_len[index % length]
assert (end - start) % db.block_size == 0
size = db.block_len[index % length] * cdiv(end - start, db.block_size)
addr_list.append(addr)
size_list.append(size)
return addr_list, size_list, block_id
def _make_db(num_regions: int, num_block_lens: int) -> ChunkedTokenDatabase:
md = KeyMetadata(model_name="t", tp_rank=1, pcp_rank=0, dcp_rank=0, pp_rank=0)
db = ChunkedTokenDatabase(md, BLOCK_SIZE)
db.set_kv_caches_base_addr(
[0x7F00_0000_0000 + i * (1 << 30) for i in range(num_regions)]
)
# Exercise repeated block lengths when there are more cache regions.
db.set_block_len([30_208 + 512 * i for i in range(num_block_lens)])
return db
@pytest.mark.parametrize("num_regions,num_block_lens", [(96, 96), (96, 2), (1, 1)])
def test_prepare_values_matches_reference(num_regions: int, num_block_lens: int):
db = _make_db(num_regions, num_block_lens)
rng = random.Random(0)
n_blocks = 300
block_ids = [rng.randrange(0, 1 << 20) for _ in range(n_blocks)]
chunks = []
b = 0
while b < n_blocks - 4:
span = rng.choice([1, 1, 1, 2, 4])
chunks.append((b * BLOCK_SIZE, (b + span) * BLOCK_SIZE))
b += span + rng.choice([0, 1])
addrs, sizes, bids = db.prepare_values(chunks, block_ids)
assert len(addrs) == len(sizes) == len(bids) == len(chunks)
for (start, end), addr, size, bid in zip(chunks, addrs, sizes, bids):
ref_addr, ref_size, ref_bid = _reference_prepare_value(
db, start, end, block_ids
)
assert addr == ref_addr
assert size == ref_size
assert bid == ref_bid
# Native bindings require Python ints rather than numpy scalars.
assert all(type(a) is int for a in addr)
assert type(bid) is int
def test_prepare_value_single_matches_reference():
db = _make_db(8, 8)
block_ids = list(range(64))
got = db.prepare_value(5 * BLOCK_SIZE, 7 * BLOCK_SIZE, block_ids)
assert got == _reference_prepare_value(
db, 5 * BLOCK_SIZE, 7 * BLOCK_SIZE, block_ids
)
def test_prepare_values_empty():
db = _make_db(4, 4)
assert db.prepare_values([], [1, 2, 3]) == ([], [], [])
def test_prepare_values_rejects_unaligned_chunk():
db = _make_db(4, 4)
with pytest.raises(AssertionError):
db.prepare_values([(0, BLOCK_SIZE + 1)], [0, 1])
@@ -473,27 +473,25 @@ def test_from_request_tracker_no_load_saves_normally():
class _StubLookupClient:
def __init__(self, hit_tokens: int) -> None:
self._hit_tokens = hit_tokens
self.num_tokens: list[int] = []
def lookup(
self,
req_id: str,
token_len: int,
num_tokens: int,
block_hashes: list[bytes],
non_block: bool = False,
) -> int:
self.num_tokens.append(num_tokens)
return self._hit_tokens
def test_full_external_hit_keeps_kvpool_cached_tokens_block_aligned():
# When the external store hits the entire prompt, scheduler must leave at
# least one token uncomputed for sampling but stay on a block boundary.
# Otherwise the recv-side load mask floors token_len to
# (num_tokens-1)//block_size, the tail partial chunk is dropped, and -- if
# the local cache covers the aligned prefix -- key_list ends up empty
# (ZeroDivisionError in the recv thread's `tp_rank % len(key_list)`).
# The worker re-derives a full external hit below the request end on an
# existing boundary, so the scheduler receives the usable aligned hit.
scheduler = _make_bare_scheduler()
scheduler.load_async = True
scheduler.client = _StubLookupClient(hit_tokens=48) # full hit on 48-token prompt
scheduler.client = _StubLookupClient(hit_tokens=32)
request = SimpleNamespace(
request_id="req-0",
@@ -510,6 +508,7 @@ def test_full_external_hit_keeps_kvpool_cached_tokens_block_aligned():
assert need_to_allocate == 16
assert load_async is True
load_spec = scheduler.load_specs["req-0"]
assert scheduler.client.num_tokens == [48]
assert load_spec.vllm_cached_tokens == 16
assert load_spec.kvpool_cached_tokens == 32
assert load_spec.kvpool_cached_tokens % 16 == 0
@@ -522,7 +521,7 @@ def test_full_external_hit_with_full_local_hit_skips_load():
# into any block-aligned key.
scheduler = _make_bare_scheduler()
scheduler.load_async = True
scheduler.client = _StubLookupClient(hit_tokens=48)
scheduler.client = _StubLookupClient(hit_tokens=32)
request = SimpleNamespace(
request_id="req-0",
@@ -655,6 +655,68 @@ def test_store_sending_thread_delta_saves_only_new_masked_chunks():
assert masked_hashes == [b"a2".hex()]
def test_store_sending_thread_prepares_missing_chunks_once_per_group():
store = MagicMock()
store.batch_is_exist.return_value = [0, 1, 0, 1, 0, 0]
store.batch_put_from_multi_buffers.return_value = [256, 256, 512, 512]
coord = SimpleNamespace(
lcm_block_size=16,
store_mask=lambda token_len, start_token, num_prompt_tokens=None: (
None,
None,
),
)
db0 = ChunkedTokenDatabase(
KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
block_size=16,
)
db0.set_kv_caches_base_addr([0x1000])
db0.set_block_len([256])
db0.prepare_values = MagicMock(wraps=db0.prepare_values)
db0.prepare_value = MagicMock(side_effect=AssertionError("scalar path called"))
db1 = ChunkedTokenDatabase(
KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
block_size=16,
)
db1.set_kv_caches_base_addr([0x2000])
db1.set_block_len([512])
db1.prepare_values = MagicMock(wraps=db1.prepare_values)
db1.prepare_value = MagicMock(side_effect=AssertionError("scalar path called"))
thread = _make_store_sending_thread(
store,
coord=coord,
token_databases=[db0, db1],
)
thread.add_stored_request("req-a")
thread._handle_request(
ReqMeta(
req_id="req-a",
token_len_chunk=48,
block_ids=([0, 1, 2], [2, 1, 0]),
block_hashes=[b"a0", b"a1", b"a2"],
can_save=True,
)
)
db0.prepare_value.assert_not_called()
db1.prepare_value.assert_not_called()
db0.prepare_values.assert_called_once_with([(0, 16), (32, 48)], [0, 1, 2])
db1.prepare_values.assert_called_once_with([(16, 32), (32, 48)], [2, 1, 0])
keys, addrs, sizes, _ = store.batch_put_from_multi_buffers.call_args.args
assert [key.rsplit("@", 1)[-1] for key in keys] == [
"6130",
"6132",
"6131",
"6132",
]
assert addrs == [[0x1000], [0x1200], [0x2200], [0x2000]]
assert sizes == [[256], [256], [512], [512]]
def test_store_sending_thread_only_skips_on_no_available_handle():
store = MagicMock()
store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
@@ -1675,6 +1737,61 @@ def test_lookup_partial_prefix_returns_first_hit_length():
assert worker.lookup(48, [b"a0", b"a1", b"a2"]) == 32
def test_lookup_full_hit_reuses_existing_boundary():
"""A full hit is re-derived below the request end without another RPC."""
worker = _make_bare_worker(block_size=16)
worker.store.batch_is_exist.return_value = [1, 1]
assert worker.lookup(32, [b"h0", b"h1"]) == 16
assert worker.store.batch_is_exist.call_count == 1
def test_lookup_full_hit_with_eagle_pops_once_not_twice():
"""Eagle already leaves the last block for the drafter, so a
full-prompt re-derivation must never fire for eagle-governed hits:
firing would anchor the search one block lower and pop a second
block, regressing the hit by an extra producer boundary."""
worker = _make_bare_worker(block_size=16)
worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
worker._kv_cache_groups,
scheduler_block_size=16,
hash_block_size=16,
use_eagle=True,
)
worker.store.batch_is_exist.return_value = [1, 1, 1, 1]
# 64-token exact-multiple prompt, all 4 blocks stored: one eagle pop
# gives 48; a spurious re-derivation (anchored at 48) would pop again
# and return 32.
assert worker.lookup(64, [b"h0", b"h1", b"h2", b"h3"]) == 48
assert worker.store.batch_is_exist.call_count == 1
def test_lookup_full_hit_swa_degrades_when_no_stored_boundary_is_usable():
"""The motivating livelock: the producer of a 64-token prompt stored
only its SWA tail window (blocks 2-3). The old arithmetic clamp turned
the full hit into 48, whose SWA window needs the never-written block 1,
so every load failed and the recompute re-entered the same lookup. The
re-derivation must report that no stored boundary below the request end
is usable."""
from vllm.v1.kv_cache_interface import KVCacheGroupSpec, SlidingWindowSpec
worker = _make_bare_worker(block_size=16)
swa = SlidingWindowSpec(
block_size=16, num_kv_heads=8, head_size=64, dtype=None, sliding_window=32
)
worker._kv_cache_groups = [KVCacheGroupSpec(["layer0"], swa)]
worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
worker._kv_cache_groups,
scheduler_block_size=worker.hash_block_size,
hash_block_size=worker.hash_block_size,
)
worker.store.batch_is_exist.return_value = [0, 0, 1, 1]
assert worker.lookup(64, [b"h0", b"h1", b"h2", b"h3"]) == 0
assert worker.store.batch_is_exist.call_count == 1
def test_lookup_swa_single_group_returns_full_when_tail_window_present():
"""Single-SWA, sliding_window=32 (= 2 blocks): producer stored only the
tail. Coordinator-driven lookup returns full prefix even though the
@@ -1692,7 +1809,7 @@ def test_lookup_swa_single_group_returns_full_when_tail_window_present():
hash_block_size=worker.hash_block_size,
)
worker.store.batch_is_exist.return_value = [0, 0, 1, 1]
assert worker.lookup(64, [b"h0", b"h1", b"h2", b"h3"]) == 64
assert worker.lookup(65, [b"h0", b"h1", b"h2", b"h3"]) == 64
def test_lookup_checks_all_potential_swa_hit_boundaries():
@@ -2157,7 +2274,7 @@ def test_lookup_records_mooncake_metrics():
worker = _make_bare_worker()
worker.store.batch_is_exist.return_value = [1, 1]
result = worker.lookup(32, [b"a0", b"a1"])
result = worker.lookup(33, [b"a0", b"a1"])
stats = worker.get_kv_connector_stats()
assert result == 32
@@ -412,3 +412,65 @@ def test_block_verification_accepts_at_least_as_many(num_speculative_steps: int)
f"Block verification mean accepted length {mean_block:.4f} is worse "
f"than standard {mean_standard:.4f}."
)
@pytest.mark.parametrize("has_draft_logits", [True, False])
def test_chunked_requests_match_full_batch(has_draft_logits: bool):
torch.manual_seed(7)
device = "cuda"
num_reqs = 5
num_speculative_steps = 3
vocab_size = 257
target_logits = torch.randn(vocab_size, device=device)
draft_logits = torch.randn(vocab_size, device=device)
inputs = _build_rejection_sample_inputs(
target_logits,
draft_logits,
num_speculative_steps,
temperature=0.6,
num_trials=num_reqs,
)
padded_target_logits = torch.empty(
inputs["target_logits"].shape[0], vocab_size + 3, device=device
)
padded_target_logits[:, :vocab_size].copy_(inputs["target_logits"])
inputs["target_logits"] = padded_target_logits[:, :vocab_size]
assert inputs["target_logits"].stride(-1) == 1
assert not inputs["target_logits"].is_contiguous()
if not has_draft_logits:
inputs["draft_logits"] = None
sampled, num_sampled = rejection_sample(
**inputs, num_speculative_steps=num_speculative_steps
)
sampled_chunks = []
num_sampled_chunks = []
for start, end in ((0, 2), (2, 5)):
lo = start * (num_speculative_steps + 1)
hi = end * (num_speculative_steps + 1)
chunk_inputs = dict(inputs)
for name in (
"target_logits",
"draft_sampled",
"pos",
"expanded_idx_mapping",
"expanded_local_pos",
):
chunk_inputs[name] = inputs[name][lo:hi]
chunk_inputs["cu_num_logits"] = inputs["cu_num_logits"][start : end + 1] - lo
chunk_inputs["idx_mapping"] = inputs["idx_mapping"][start:end]
chunk_sampled, chunk_num_sampled = rejection_sample(
**chunk_inputs, num_speculative_steps=num_speculative_steps
)
sampled_chunks.append(chunk_sampled)
num_sampled_chunks.append(chunk_num_sampled)
chunked_sampled = torch.cat(sampled_chunks)
chunked_num_sampled = torch.cat(num_sampled_chunks)
assert torch.equal(chunked_num_sampled, num_sampled)
steps = torch.arange(num_speculative_steps + 1, device=device)
valid = steps.unsqueeze(0) < num_sampled.unsqueeze(1)
assert torch.equal(chunked_sampled[valid], sampled[valid])
+26 -1
View File
@@ -2,7 +2,32 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from unittest import TestCase
from vllm.v1.outputs import LogprobsLists
import torch
from vllm.v1.outputs import LogprobsLists, LogprobsTensors
def test_logprobs_tensors_cat():
first = LogprobsTensors(
torch.tensor([[1, 2]]),
torch.tensor([[0.1, 0.2]]),
torch.tensor([1]),
)
second = LogprobsTensors(
torch.tensor([[3, 4]]),
torch.tensor([[0.3, 0.4]]),
torch.tensor([2]),
)
result = LogprobsTensors.cat([first, second], [0, 1, 2])
assert result.logprob_token_ids.tolist() == [[1, 2], [3, 4]]
assert result.logprobs.tolist() == (
first.logprobs.tolist() + second.logprobs.tolist()
)
assert result.selected_token_ranks.tolist() == [1, 2]
assert result.cu_num_generated_tokens == [0, 1, 2]
assert LogprobsTensors.cat([first]) is first
class TestLogprobsLists(TestCase):
@@ -0,0 +1,109 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from types import MethodType, SimpleNamespace
from typing import get_args
import numpy as np
import pytest
import torch
from vllm.config.model import PROCESSED_LOGPROBS_MODES, LogprobsMode
from vllm.platforms import current_platform
from vllm.v1.worker.gpu.spec_decode.rejection_sampler import (
RejectionSampler,
_iter_request_chunks,
)
def test_iter_request_chunks_preserves_request_boundaries():
cu_num_logits = np.array([0, 3, 4, 11, 13], dtype=np.int32)
assert list(_iter_request_chunks(cu_num_logits, max_chunk_logits=5)) == [
(0, 2),
(2, 3),
(3, 4),
]
@pytest.mark.skipif(not current_platform.is_cuda(), reason="Requires CUDA")
@pytest.mark.parametrize("logprobs_mode", get_args(LogprobsMode))
def test_chunked_scores_match_full_batch(logprobs_mode: str):
device = torch.device("cuda")
cu_num_logits_np = np.array([0, 3, 4, 8, 10], dtype=np.int32)
num_logits_per_req = np.diff(cu_num_logits_np)
idx_mapping_np = np.array([7, 2, 9, 1], dtype=np.int32)
input_batch = SimpleNamespace(
num_reqs=4,
cu_num_logits_np=cu_num_logits_np,
cu_num_logits=torch.from_numpy(cu_num_logits_np).to(device),
idx_mapping_np=idx_mapping_np,
idx_mapping=torch.from_numpy(idx_mapping_np).to(device),
expanded_idx_mapping=torch.from_numpy(
np.repeat(idx_mapping_np, num_logits_per_req)
).to(device),
expanded_local_pos=torch.from_numpy(
np.concatenate(
[np.arange(count, dtype=np.int32) for count in num_logits_per_req]
)
).to(device),
)
rejection_sampler = object.__new__(RejectionSampler)
rejection_sampler.sampler = SimpleNamespace(logprobs_mode=logprobs_mode)
rejection_sampler.num_speculative_steps = 3
def fake_verify(
self,
logits,
_draft_logits,
_draft_sampled,
_pos,
cu_num_logits,
idx_mapping,
*_mappings,
):
num_sampled = torch.diff(cu_num_logits).to(torch.int32)
sampled = (
idx_mapping.to(torch.int64).unsqueeze(1) + torch.arange(4, device=device)
) % logits.shape[1]
return logits.float() + 1, sampled, num_sampled
rejection_sampler._verify = MethodType(fake_verify, rejection_sampler)
logits = torch.arange(170, dtype=torch.float32, device=device).view(10, 17)
sampled, num_sampled, chunked_logprobs = rejection_sampler._verify_in_chunks(
logits,
input_batch,
draft_logits=None,
draft_sampled=torch.arange(10, device=device),
pos=torch.arange(10, device=device),
max_chunk_logits=5,
max_num_logprobs=2,
)
score_logits = logits + 1 if logprobs_mode in PROCESSED_LOGPROBS_MODES else logits
full_logprobs = rejection_sampler._get_logprobs_tensors(
sampled,
num_sampled,
score_logits,
input_batch.cu_num_logits,
input_batch.cu_num_logits_np,
max_num_logprobs=2,
)
assert sampled[:, 0].tolist() == idx_mapping_np.tolist()
assert num_sampled.tolist() == num_logits_per_req.tolist()
assert chunked_logprobs is not None
assert full_logprobs is not None
assert torch.equal(
chunked_logprobs.logprob_token_ids,
full_logprobs.logprob_token_ids,
)
assert torch.equal(chunked_logprobs.logprobs, full_logprobs.logprobs)
assert torch.equal(
chunked_logprobs.selected_token_ranks,
full_logprobs.selected_token_ranks,
)
assert (
chunked_logprobs.cu_num_generated_tokens
== full_logprobs.cu_num_generated_tokens
)
+4 -7
View File
@@ -14,14 +14,10 @@ def test_block_ids_are_not_overwritten_while_copy_is_in_flight():
page_size_el = 4
storage = torch.ones((num_blocks, page_size_el), dtype=torch.int32, device=device)
# Build the minimal zeroer state directly so the test can focus on ID-buffer
# lifetime without constructing model attention groups.
# Build the minimal zeroer state directly so the test can focus on the
# in-flight copy behavior without constructing model attention groups.
zeroer = KVBlockZeroer.__new__(KVBlockZeroer)
zeroer.device = device
zeroer.pin_memory = True
zeroer.max_concurrency = 2
zeroer._id_cap = 8
zeroer._allocate_id_buffers()
zeroer._meta = (
torch.tensor([storage.data_ptr()], dtype=torch.uint64, device=device),
page_size_el,
@@ -32,7 +28,8 @@ def test_block_ids_are_not_overwritten_while_copy_is_in_flight():
stream = torch.cuda.Stream()
with torch.cuda.stream(stream):
# Keep the first nonblocking H2D copy pending while the host submits the
# second call. A single shared pinned source would be overwritten here.
# second call. Each call must stage from its own pinned source so the
# first copy is not corrupted before it runs.
torch.cuda._sleep(10_000_000)
zeroer.zero_block_ids([1])
zeroer.zero_block_ids([2])
+1 -1
View File
@@ -2460,8 +2460,8 @@ def topk_hash_softplus_sqrt(
input_tokens,
hash_indices_table,
)
return
return
torch.ops._moe_C.topk_softplus_sqrt(
topk_weights,
topk_indices,
+4
View File
@@ -90,6 +90,10 @@ ModelDType = Literal["auto", "half", "float16", "bfloat16", "float", "float32"]
LogprobsMode = Literal[
"raw_logits", "raw_logprobs", "processed_logits", "processed_logprobs"
]
PROCESSED_LOGPROBS_MODES: tuple[LogprobsMode, ...] = (
"processed_logits",
"processed_logprobs",
)
HfOverrides = dict[str, Any] | Callable[[PretrainedConfig], PretrainedConfig]
ModelImpl = Literal["auto", "vllm", "transformers", "terratorch"]
LayerBlockType = Literal["attention", "linear_attention", "mamba"]
@@ -57,6 +57,11 @@ if TYPE_CHECKING:
from _typeshed import SizedBuffer
VLLM_RINGBUFFER_WARNING_INTERVAL = envs.VLLM_RINGBUFFER_WARNING_INTERVAL
# Cap on how long an idle reader parks before re-reading the authoritative SHM
# written-flag. Bounds lost-notify recovery latency to ~5s while the periodic
# wakeup stays negligible (one flag check per reader every 5s).
SHM_READER_RECHECK_INTERVAL_MS = 5000
from_bytes_big = functools.partial(int.from_bytes, byteorder="big")
@@ -631,25 +636,22 @@ class MessageQueue:
self.n_warning = 1
self.timeout = timeout
def timeout_ms(self) -> int | None:
"""Returns a timeout that is:
def timeout_ms(self) -> int:
"""Returns a timeout, capped at the recheck interval, that is:
- min(time to deadline, time to next warning) if we're logging warnings
- time to deadline, if we're not logging warnings
- None if the timeout is None and we're not logging warnings
- recheck interval if the timeout is None and we're not logging warnings
- raise TimeoutError if we are past the deadline
"""
warning_wait_time = self.warning_wait_time_ms
wait_ms = SHM_READER_RECHECK_INTERVAL_MS
if self.warning_wait_time_ms is not None:
wait_ms = min(wait_ms, self.warning_wait_time_ms)
if self.timeout is None:
return warning_wait_time
return wait_ms
time_left_ms = int((self.deadline - time.monotonic()) * 1000)
if time_left_ms <= 0:
raise TimeoutError
if warning_wait_time and warning_wait_time < time_left_ms:
return warning_wait_time
return time_left_ms
return min(wait_ms, time_left_ms)
def should_warn(self) -> bool:
"""Returns true if it's time to log a warning for a timeout that is not
@@ -710,18 +712,18 @@ class MessageQueue:
# found a block that is not read by this reader
# let caller read from the buffer
with self.buffer.get_data(self.current_idx) as buf:
yield buf
# caller has read from the buffer
# set the read flag
metadata_buffer[self.local_reader_rank + 1] = 1
# Memory fence ensures the read flag is visible to the writer.
# Without this, writer may not see our read completion and
# could wait indefinitely for all readers to finish.
memory_fence()
self.current_idx = (self.current_idx + 1) % self.buffer.max_chunks
self._spin_condition.record_read()
try:
yield buf
finally:
# caller has read from the buffer; set the read flag.
metadata_buffer[self.local_reader_rank + 1] = 1
# Memory fence ensures the read flag is visible to the writer.
# Without this, writer may not see our read completion and
# could wait indefinitely for all readers to finish.
memory_fence()
next_idx = self.current_idx + 1
self.current_idx = next_idx % self.buffer.max_chunks
self._spin_condition.record_read()
break
def enqueue(self, obj, timeout: float | None = None):
@@ -93,6 +93,9 @@ class MooncakeStoreCoordinator:
self.eagle_group_ids = set(range(len(kv_cache_groups)))
self._verify_and_split_kv_cache_groups()
def align_lookup_length(self, length: int) -> int:
return length // self.lcm_block_size * self.lcm_block_size
def _verify_and_split_kv_cache_groups(self) -> None:
"""Mirrors KVCacheCoordinator.verify_and_split_kv_cache_groups but
dispatches via spec_manager_map (we don't allocate managers).
@@ -9,6 +9,7 @@ from collections.abc import Iterable, Sequence
from dataclasses import dataclass
from typing import cast
import numpy as np
import torch
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
@@ -196,22 +197,46 @@ class ChunkedTokenDatabase:
def prepare_value(
self, start: int, end: int, block_ids: list[int]
) -> tuple[list[int], list[int], int]:
"""Compute memory addresses and sizes for a token range.
"""Compute memory addresses and sizes for a single token range.
Returns:
(addr_list, size_list, block_id)
"""
addr_list = []
size_list = []
block_id = block_ids[start // self.block_size]
addr_lists, size_lists, chunk_block_ids = self.prepare_values(
((start, end),), block_ids
)
return addr_lists[0], size_lists[0], chunk_block_ids[0]
def prepare_values(
self,
chunks: Sequence[tuple[int, int]],
block_ids: list[int],
) -> tuple[list[list[int]], list[list[int]], list[int]]:
"""Compute memory addresses and sizes for multiple token ranges.
Returns:
(addr_lists, size_lists, chunk_block_ids), one entry per chunk.
"""
if not chunks:
return [], [], []
base = np.asarray(self.kv_caches_base_addr, dtype=np.int64)
length = len(self.block_len)
for index, base_addr in enumerate(self.kv_caches_base_addr):
addr = base_addr + block_id * self.block_len[index % length]
assert (end - start) % self.block_size == 0
size = self.block_len[index % length] * cdiv(end - start, self.block_size)
addr_list.append(addr)
size_list.append(size)
return addr_list, size_list, block_id
blen = np.asarray(
[self.block_len[i % length] for i in range(base.shape[0])],
dtype=np.int64,
)
n = len(chunks)
starts = np.fromiter((c[0] for c in chunks), dtype=np.int64, count=n)
spans = np.fromiter((c[1] for c in chunks), dtype=np.int64, count=n) - starts
assert not (spans % self.block_size).any()
bids = np.fromiter(
(block_ids[i] for i in (starts // self.block_size).tolist()),
dtype=np.int64,
count=n,
)
addrs = base[None, :] + bids[:, None] * blen[None, :]
sizes = blen[None, :] * (spans // self.block_size)[:, None]
return addrs.tolist(), sizes.tolist(), bids.tolist()
def process_tokens(
self,
@@ -10,7 +10,8 @@ Wire format (REQ/REP over IPC):
Request: [msg_type: bytes] [payload_frames...]
msg_type == LOOKUP_MSG:
frame 1: token_len (u32 big-endian, 4 bytes)
frame 1: num_tokens (u32 big-endian, 4 bytes); the worker derives
the aligned lookup length
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
@@ -80,14 +80,12 @@ class MooncakeStoreScheduler:
Returns ``(None, False)`` when an async lookup is still in flight,
signaling the scheduler to retry this request on a later step.
"""
# Look up against the full prefill range, not just the prompt.
token_len = request.num_tokens // self._block_size * self._block_size
if token_len < self._block_size:
if request.num_tokens < self._block_size:
return 0, False
num_external_hit_tokens = self.client.lookup(
request.request_id,
token_len,
request.num_tokens,
request.block_hashes,
non_block=self.lookup_async,
)
@@ -95,14 +93,6 @@ class MooncakeStoreScheduler:
# Lookup not ready yet; scheduler will retry on a later step.
return None, False
if num_external_hit_tokens == request.num_tokens:
# Leave a sub-block tail uncomputed for sampling, on a block
# boundary so the recv-side load mask covers every yielded chunk.
num_external_hit_tokens = max(
0,
(request.num_tokens - 1) // self._block_size * self._block_size,
)
if num_external_hit_tokens < num_computed_tokens:
need_to_allocate = 0
else:
@@ -634,6 +634,21 @@ class KVCacheStoreSendingThread(KVTransferThread):
addrs: list[list[int]] = []
sizes: list[list[int]] = []
stored_events: list[BlockStored] = []
chunks_per_group: list[list[tuple[int, int]]] = [
[] for _ in self.token_databases
]
for start, end, g_idx in zip(starts, ends, group_indices, strict=True):
chunks_per_group[g_idx].append((start, end))
for g_idx, chunks in enumerate(chunks_per_group):
if not chunks:
continue
db = self.token_databases[g_idx]
group_addrs, group_sizes, _ = db.prepare_values(
chunks, block_ids_per_group[g_idx]
)
addrs.extend(group_addrs)
sizes.extend(group_sizes)
# parent_block_hash chains live within a group, not across.
if self.enable_kv_event:
prev_key_per_group: dict[int, Any] = {}
@@ -645,10 +660,6 @@ class KVCacheStoreSendingThread(KVTransferThread):
zip(starts, ends, group_indices, strict=True)
):
db = self.token_databases[g_idx]
addr, size, _ = db.prepare_value(s, e, block_ids_per_group[g_idx])
addrs.append(addr)
sizes.append(size)
if self.enable_kv_event:
token_ids = (
req_meta.token_ids[s:e]
@@ -805,19 +816,21 @@ class KVCacheStoreRecvingThread(KVTransferThread):
block_id_list: list[int] = []
for g_idx, db in enumerate(self.token_databases):
mask = load_mask_per_group[g_idx]
chunks: list[tuple[int, int]] = []
for start, end, block_hash in db.process_tokens(
token_len, req_meta.block_hashes, mask_num
):
chunk_idx = start // db.block_size
if chunk_idx >= len(mask) or not mask[chunk_idx]:
continue
addr, size, block_id = db.prepare_value(
start, end, req_meta.block_ids[g_idx]
)
key_list.append(db.key_for(block_hash))
addr_list.append(addr)
size_list.append(size)
block_id_list.append(block_id)
chunks.append((start, end))
g_addrs, g_sizes, g_block_ids = db.prepare_values(
chunks, req_meta.block_ids[g_idx]
)
addr_list.extend(g_addrs)
size_list.extend(g_sizes)
block_id_list.extend(g_block_ids)
# Rotate aligned lists by tp_rank for load balancing.
rotation = self.tp_rank % len(key_list)
@@ -1454,11 +1467,14 @@ class MooncakeStoreWorker:
return finished_sending
def lookup(self, token_len: int, block_hashes: Sequence[BlockHash]) -> int:
def lookup(self, num_tokens: int, block_hashes: Sequence[BlockHash]) -> int:
"""Check how many prefix tokens exist in the store.
Checks across all rank-specific key namespaces that may be loaded.
Checks across all rank-specific key namespaces that may be loaded. A
hit covering all ``num_tokens`` is re-derived below the request end so
the last token is recomputed for sampling.
"""
token_len = self.coord.align_lookup_length(num_tokens)
if not block_hashes or token_len <= 0:
return 0
@@ -1522,11 +1538,24 @@ class MooncakeStoreWorker:
)
}
cached_block_pool = ExternalCachedBlockPool(
self.hash_block_size,
exists_set,
)
_masks, hit_length = self.coord.find_longest_cache_hit(
block_hashes,
token_len,
ExternalCachedBlockPool(self.hash_block_size, exists_set),
cached_block_pool,
)
if hit_length >= num_tokens:
usable_length = self.coord.align_lookup_length(num_tokens - 1)
if usable_length <= 0:
return 0
_masks, hit_length = self.coord.find_longest_cache_hit(
block_hashes,
usable_length,
cached_block_pool,
)
return hit_length
def get_kv_events(self) -> list[BlockStored]:
@@ -1592,11 +1621,11 @@ class LookupKeyServer:
msg_type = bytes(all_frames[0])
if msg_type == LOOKUP_MSG:
token_len = int.from_bytes(all_frames[1], byteorder="big")
num_tokens = int.from_bytes(all_frames[1], byteorder="big")
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)
result = self.store_worker.lookup(num_tokens, block_hashes)
self.socket.send(result.to_bytes(4, "big"))
elif msg_type == RESET_MSG:
@@ -1659,11 +1688,11 @@ class LookupKeyClient:
)
self.futures: dict[str, Future[int]] = {}
def _lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
def _lookup(self, num_tokens: int, block_hashes: list[BlockHash]) -> int:
hash_len = len(block_hashes[0]) if block_hashes else 0
all_frames = (
LOOKUP_MSG,
token_len.to_bytes(4, byteorder="big"),
num_tokens.to_bytes(4, byteorder="big"),
hash_len.to_bytes(2, byteorder="big"),
b"".join(block_hashes),
)
@@ -1674,7 +1703,7 @@ class LookupKeyClient:
def lookup(
self,
req_id: str,
token_len: int,
num_tokens: int,
block_hashes: list[BlockHash],
non_block: bool = False,
) -> int | None:
@@ -1682,7 +1711,7 @@ class LookupKeyClient:
so the caller retries on a later step."""
future = self.futures.get(req_id)
if future is None:
future = self.executor.submit(self._lookup, token_len, list(block_hashes))
future = self.executor.submit(self._lookup, num_tokens, list(block_hashes))
self.futures[req_id] = future
if non_block and not future.done():
return None
+69 -37
View File
@@ -968,7 +968,9 @@ class MultiModalContentParser(BaseMultiModalContentParser):
`tensor.shape[0]` placeholder tokens after tokenization.
"""
if not self.model_config.enable_prompt_embeds:
raise ValueError(_ENABLE_PROMPT_EMBEDS_ERROR)
raise VLLMValidationError(
_ENABLE_PROMPT_EMBEDS_ERROR, parameter="prompt_embeds"
)
tensor = safe_load_prompt_embeds(self.model_config, data.encode())
self._tracker.add("prompt_embeds", (tensor, None))
@@ -987,8 +989,9 @@ class MultiModalContentParser(BaseMultiModalContentParser):
) -> None:
mm_config = self.model_config.get_multimodal_config()
if not mm_config.enable_mm_embeds:
raise ValueError(
"You must set `--enable-mm-embeds` to input `image_embeds`"
raise VLLMValidationError(
"You must set `--enable-mm-embeds` to input `image_embeds`",
parameter="image_embeds",
)
if isinstance(image_embeds, dict):
@@ -1014,8 +1017,9 @@ class MultiModalContentParser(BaseMultiModalContentParser):
) -> None:
mm_config = self.model_config.get_multimodal_config()
if not mm_config.enable_mm_embeds:
raise ValueError(
"You must set `--enable-mm-embeds` to input `audio_embeds`"
raise VLLMValidationError(
"You must set `--enable-mm-embeds` to input `audio_embeds`",
parameter="audio_embeds",
)
if isinstance(audio_embeds, dict):
@@ -1117,7 +1121,9 @@ class AsyncMultiModalContentParser(BaseMultiModalContentParser):
thread-pool executor via `safe_load_prompt_embeds_async`.
"""
if not self.model_config.enable_prompt_embeds:
raise ValueError(_ENABLE_PROMPT_EMBEDS_ERROR)
raise VLLMValidationError(
_ENABLE_PROMPT_EMBEDS_ERROR, parameter="prompt_embeds"
)
self._tracker.add(
"prompt_embeds", partial(self._load_prompt_embeds_async, data.encode())
@@ -1151,26 +1157,36 @@ class AsyncMultiModalContentParser(BaseMultiModalContentParser):
) -> None:
mm_config = self.model_config.get_multimodal_config()
if not mm_config.enable_mm_embeds:
raise ValueError(
"You must set `--enable-mm-embeds` to input `image_embeds`"
raise VLLMValidationError(
"You must set `--enable-mm-embeds` to input `image_embeds`",
parameter="image_embeds",
)
if isinstance(image_embeds, dict):
embeds = {
k: self._connector.fetch_image_embedding(v)
for k, v in image_embeds.items()
}
elif isinstance(image_embeds, str):
embedding = self._connector.fetch_image_embedding(image_embeds)
embeds = embedding
else:
embeds = None
placeholder = self._tracker.add(
"image_embeds", partial(self._item_with_uuid_async, embeds, uuid)
"image_embeds",
partial(self._image_embeds_with_uuid_async, image_embeds, uuid),
)
self._add_placeholder("image", placeholder)
async def _image_embeds_with_uuid_async(
self,
image_embeds: str | dict[str, str] | None,
uuid: str | None,
):
if isinstance(image_embeds, dict):
tensors = await asyncio.gather(
*(
self._connector.fetch_image_embedding_async(v)
for v in image_embeds.values()
)
)
embeds = dict(zip(image_embeds, tensors))
elif isinstance(image_embeds, str):
embeds = await self._connector.fetch_image_embedding_async(image_embeds)
else:
embeds = None
return embeds, uuid
def parse_audio_embeds(
self,
audio_embeds: str | dict[str, str] | None,
@@ -1178,26 +1194,36 @@ class AsyncMultiModalContentParser(BaseMultiModalContentParser):
) -> None:
mm_config = self.model_config.get_multimodal_config()
if not mm_config.enable_mm_embeds:
raise ValueError(
"You must set `--enable-mm-embeds` to input `audio_embeds`"
raise VLLMValidationError(
"You must set `--enable-mm-embeds` to input `audio_embeds`",
parameter="audio_embeds",
)
if isinstance(audio_embeds, dict):
embeds = {
k: self._connector.fetch_audio_embedding(v)
for k, v in audio_embeds.items()
}
elif isinstance(audio_embeds, str):
embedding = self._connector.fetch_audio_embedding(audio_embeds)
embeds = embedding
else:
embeds = None
placeholder = self._tracker.add(
"audio_embeds", partial(self._item_with_uuid_async, embeds, uuid)
"audio_embeds",
partial(self._audio_embeds_with_uuid_async, audio_embeds, uuid),
)
self._add_placeholder("audio", placeholder)
async def _audio_embeds_with_uuid_async(
self,
audio_embeds: str | dict[str, str] | None,
uuid: str | None,
):
if isinstance(audio_embeds, dict):
tensors = await asyncio.gather(
*(
self._connector.fetch_audio_embedding_async(v)
for v in audio_embeds.values()
)
)
embeds = dict(zip(audio_embeds, tensors))
elif isinstance(audio_embeds, str):
embeds = await self._connector.fetch_audio_embedding_async(audio_embeds)
else:
embeds = None
return embeds, uuid
def parse_image_pil(
self,
image_pil: Image.Image | None,
@@ -1592,10 +1618,14 @@ def _parse_chat_message_content_mm_part(
tool_reference = tool_reference_params.get("name", None)
return "tool_reference", tool_reference
# Raise an error if no 'type' or direct URL is found.
raise ValueError("Missing 'type' field in multimodal part.")
raise VLLMValidationError(
"Missing 'type' field in multimodal part.", parameter="type"
)
if not isinstance(part_type, str):
raise ValueError("Invalid 'type' field in multimodal part.")
raise VLLMValidationError(
"Invalid 'type' field in multimodal part.", parameter="type"
)
return part_type, "unknown part_type content"
@@ -1732,7 +1762,9 @@ def _parse_chat_message_content_part(
modality = "audio"
elif part_type == "prompt_embeds":
if not content:
raise ValueError(_PROMPT_EMBEDS_MISSING_DATA_ERROR)
raise VLLMValidationError(
_PROMPT_EMBEDS_MISSING_DATA_ERROR, parameter="prompt_embeds"
)
mm_parser.parse_prompt_embeds(cast(str, content))
modality = "prompt_embeds"
elif part_type == "audio_url":
+7
View File
@@ -2114,6 +2114,13 @@ def compile_factors() -> dict[str, object]:
"VLLM_CACHE_ROOT",
# Runtime memory-plan persistence; does not affect compiled graphs.
"VLLM_ENABLE_STARTUP_PLAN",
# Location-only derived paths: where a cache/config directory lives
# cannot affect compiled artifacts, and hashing them means relocating
# HOME or the XDG roots silently invalidates every compile cache
# (VLLM_CACHE_ROOT above and VLLM_FLASHINFER_AUTOTUNE_CACHE_DIR below
# are already ignored for the same reason).
"VLLM_XLA_CACHE_PATH",
"VLLM_CONFIG_ROOT",
"LD_LIBRARY_PATH",
"VLLM_SERVER_DEV_MODE",
"VLLM_DP_MASTER_IP",
@@ -139,10 +139,13 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, CustomOp):
def _maybe_pad_weight(self, weight: torch.Tensor) -> torch.Tensor:
# Pad the weight tensor. This is an optimization on ROCm platform, which
# can benefit from tensors located far enough from one another in memory
# can benefit from tensors located far enough from one another in memory.
# Skip padding when EPLB is enabled because EPLB requires contiguous
# weights for the view/rearrangement operations.
if (
envs.VLLM_ROCM_MOE_PADDING
and current_platform.is_rocm()
and not self.moe.moe_parallel_config.enable_eplb
and weight.stride(-1) == 1
and (weight.stride(-2) * weight.element_size()) % 512 == 0
):
@@ -109,8 +109,8 @@ def prepare_nvfp4_moe_layer_for_flashinfer_cutedsl(
# Global scaling factors (same as other FlashInfer backends).
num_experts = w13.shape[0]
a13_scale = a13_scale.max().to(torch.float32).expand(num_experts)
a2_scale = a2_scale.max().to(torch.float32).expand(num_experts)
a13_scale = a13_scale.max().to(torch.float32).repeat(num_experts)
a2_scale = a2_scale.max().to(torch.float32).repeat(num_experts)
half = w13.shape[1] // 2
w13 = torch.cat([w13[:, half:], w13[:, :half]], dim=1)
@@ -338,8 +338,8 @@ def prepare_nvfp4_moe_layer_for_fi_or_cutlass(
# For some FI kernels, the input scales are shared by all experts.
if is_global_sf_supported_for_nvfp4_backend(backend):
num_experts = w13.shape[0]
a13_scale = a13_scale.max().to(torch.float32).expand(num_experts)
a2_scale = a2_scale.max().to(torch.float32).expand(num_experts)
a13_scale = a13_scale.max().to(torch.float32).repeat(num_experts)
a2_scale = a2_scale.max().to(torch.float32).repeat(num_experts)
else:
a13_scale = a13_scale.max(dim=1).values.to(torch.float32)
+12 -2
View File
@@ -103,6 +103,16 @@ class Qwen3_5MultiTokenPredictor(nn.Module):
prefix=f"{prefix}.fc",
)
# GPTQ: quantized checkpoints may exclude MTP from quantization via
# quantization_config.dynamic with "-:pattern" entries. When detected,
# disable quantization for MTP layers so they use unquantized params.
original_quant = vllm_config.quant_config
if quant_config and quant_config.get_name() not in ("modelopt_fp4",):
hf_qc = getattr(model_config.hf_config, "quantization_config", None)
if isinstance(hf_qc, dict):
dynamic = hf_qc.get("dynamic", {})
if any(k.startswith("-:") and "mtp" in k for k in dynamic):
vllm_config.quant_config = None
self.layers = torch.nn.ModuleList(
Qwen3_5DecoderLayer(
vllm_config,
@@ -111,11 +121,10 @@ class Qwen3_5MultiTokenPredictor(nn.Module):
)
for idx in range(self.num_mtp_layers)
)
vllm_config.quant_config = original_quant
self.make_empty_intermediate_tensors = make_empty_intermediate_tensors_factory(
["hidden_states", "residual"], config.hidden_size
)
self.norm = Qwen3_5RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.pre_fc_norm_hidden = Qwen3_5RMSNorm(
config.hidden_size, eps=config.rms_norm_eps
@@ -170,6 +179,7 @@ class Qwen3_5MultiTokenPredictor(nn.Module):
positions.shape[-1],
self.config.hidden_size,
)
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
+14 -4
View File
@@ -133,6 +133,18 @@ def _extract_data_from_fused_moe_module(
return w13, w13_s, w2, w2_s, num_topk
def _is_deep_gemm_backed_kernel(fp8_linear: object) -> bool:
"""
Return True if the selected linear kernel dispatches to DeepGEMM, either
directly or as the fallback branch of a dynamic wrapper.
"""
if isinstance(fp8_linear, DeepGemmFp8BlockScaledMMKernel):
return True
return isinstance(
getattr(fp8_linear, "fallback", None), DeepGemmFp8BlockScaledMMKernel
)
def _fp8_linear_may_use_deep_gemm(module: torch.nn.Module) -> bool:
"""
Return True if the input module/layer could be processed with DeepGEMM.
@@ -147,10 +159,8 @@ def _fp8_linear_may_use_deep_gemm(module: torch.nn.Module) -> bool:
):
return False
if not isinstance(
getattr(module.quant_method, "fp8_linear", None),
DeepGemmFp8BlockScaledMMKernel,
):
fp8_linear = getattr(module.quant_method, "fp8_linear", None)
if not _is_deep_gemm_backed_kernel(fp8_linear):
return False
block_size = get_mk_alignment_for_contiguous_layout()[0]
+121 -11
View File
@@ -8,7 +8,8 @@ import regex as re
import torch
import torch.nn as nn
from vllm.config import VllmConfig
import vllm.envs as envs
from vllm.config import VllmConfig, get_current_vllm_config
from vllm.distributed import (
get_pp_group,
get_tensor_model_parallel_rank,
@@ -110,6 +111,51 @@ class DeepseekV4MLP(nn.Module):
return x
def _shared_experts_are_fp4(config, layer_idx: int | None = None) -> bool:
"""Whether the shared experts are MXFP4 and thus fusable.
``layer_idx=None`` resolves the model-wide default (global scheme), used by
the main-model weight loader / mapper callers that operate per-model.
"""
quant_cfg = getattr(config, "quantization_config", None)
if quant_cfg is None:
return False
if layer_idx is None:
base = None
elif layer_idx >= config.num_hidden_layers:
base = f"mtp.{layer_idx - config.num_hidden_layers}.ffn.shared_experts"
else:
base = f"layers.{layer_idx}.ffn.shared_experts"
if base and any(e.startswith(base) for e in (quant_cfg.get("exclude") or [])):
return False
entry = (
(quant_cfg.get("layer_quant_config") or {}).get(f"{base}.w1") if base else None
)
if entry is None:
entry = quant_cfg.get("global_quant_config")
return ((entry or {}).get("weight") or {}).get("dtype") == "fp4"
def _fuse_shared_experts_enabled(config, prefix: str = "") -> bool:
"""Whether to fuse the shared expert into the routed MXFP4 grouped GEMM.
Fusion fuses the shared expert into the routed experts' MXFP4 grouped GEMM,
so it only applies where the shared expert is the same precision as the
routed experts. Some layers may carry a shared expert in a different quantization
than the routed experts; when so, it runs as its own linear and must not be fused.
"""
if not (
current_platform.is_rocm()
and getattr(config, "n_shared_experts", None)
and envs.VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS
and not get_current_vllm_config().parallel_config.enable_expert_parallel
):
return False
return _shared_experts_are_fp4(
config, extract_layer_index(prefix) if prefix else None
)
class DeepseekV4MoE(nn.Module):
def __init__(
self,
@@ -164,7 +210,11 @@ class DeepseekV4MoE(nn.Module):
requires_grad=False,
)
if config.n_shared_experts is None:
self.n_shared_experts = config.n_shared_experts
self.fuse_shared_experts = _fuse_shared_experts_enabled(config, prefix)
if config.n_shared_experts is None or self.fuse_shared_experts:
self.shared_experts = None
else:
intermediate_size = config.moe_intermediate_size * config.n_shared_experts
@@ -188,6 +238,9 @@ class DeepseekV4MoE(nn.Module):
self.experts = FusedMoE(
shared_experts=self.shared_experts,
n_shared_experts=(
config.n_shared_experts if self.fuse_shared_experts else None
),
gate=self.gate,
num_experts=config.n_routed_experts,
top_k=config.num_experts_per_tok,
@@ -667,7 +720,38 @@ class DeepseekV4Model(nn.Module, EagleModelMixin):
# Pre-compute expert mapping ONCE.
expert_mapping = self.get_expert_mapping()
# Use each MoE's own per-layer fusion decision (computed with its prefix
# at init) as the single source of truth, so the redirect below cannot
# diverge from how the module was built if per-layer quantization ever
# mixes fused and non-fused layers.
fuse_by_layer = {
extract_layer_index(mod_name): mod.fuse_shared_experts
for mod_name, mod in self.named_modules()
if isinstance(mod, DeepseekV4MoE)
}
n_routed = self.config.n_routed_experts
# The redirect below maps the single shared-expert tensor group to one
# appended slot; multiple shared experts would need per-expert slicing
# (see deepseek_v2.py). DeepSeek-V4 has n_shared_experts == 1.
if any(fuse_by_layer.values()) and self.config.n_shared_experts != 1:
raise NotImplementedError(
"deepseek-v4 fused shared-expert loading supports only "
f"n_shared_experts == 1, got {self.config.n_shared_experts}"
)
for name, loaded_weight in weights:
# Shared-expert fusion: redirect ``.ffn.shared_experts.w{1,2,3}``
# into appended routed-expert slot ``.ffn.experts.{n_routed}``
# so the MXFP4-quantized shared expert loads through the routed
# expert loader (grouped GEMM). Single shared expert only.
if ".ffn.shared_experts.w" in name and fuse_by_layer.get(
extract_layer_index(name), False
):
name = name.replace(
".ffn.shared_experts.w",
f".ffn.experts.{n_routed}.w",
)
for param_name, weight_name, shard_id in stacked_params_mapping:
# Skip non-stacked layers and experts (experts handled below).
if ".experts." in name:
@@ -745,24 +829,41 @@ class DeepseekV4Model(nn.Module, EagleModelMixin):
def get_expert_mapping(self) -> list[tuple[str, str, int, str]]:
# Params for weights, fp8 weight scales, fp8 activation scales
# (param_name, weight_name, expert_id, shard_id)
# When fusing shared experts, include the appended slots
# (ids n_routed_experts .. n_routed_experts + n_shared - 1) so the
# redirected shared-expert weights route through the expert loader.
n_shared = getattr(self.config, "n_shared_experts", 0) or 0
num_experts = self.config.n_routed_experts + (
n_shared if _fuse_shared_experts_enabled(self.config) else 0
)
return fused_moe_make_expert_params_mapping(
self,
ckpt_gate_proj_name="w1",
ckpt_down_proj_name="w2",
ckpt_up_proj_name="w3",
num_experts=self.config.n_routed_experts,
num_experts=num_experts,
)
def _make_deepseek_v4_weights_mapper(expert_dtype: str) -> WeightsMapper:
def _make_deepseek_v4_weights_mapper(
expert_dtype: str, fuse_shared_experts: bool = False
) -> WeightsMapper:
if expert_dtype == "fp4":
# MXFP4 experts use Mxfp4MoEMethod, which registers scales as
# ``w{1,2,3}_weight_scale`` (no _inv suffix). FP8 linear and
# shared experts use Fp8LinearMethod's block scales, which
# register as ``weight_scale_inv``.
# (non-fused) shared experts use Fp8LinearMethod's block scales,
# which register as ``weight_scale_inv``.
#
# - DeepSeek native ``.scale``: expert scales -> ``.weight_scale``,
# everything else -> ``.weight_scale_inv``.
# - AMD-Quark ``.weight_scale``: linear/attn scales ->
# ``.weight_scale_inv``. Expert and shared-expert
# ``w{1,2,3}.weight_scale`` are left untouched (consumed as-is by
# the MXFP4 expert loader, which produces ``w{13,2}_weight_scale``);
scale_regex = {
re.compile(r"(\.experts\.\d+\.w[123])\.scale$"): r"\1.weight_scale",
re.compile(r"\.scale$"): ".weight_scale_inv",
re.compile(r"(?<!\.w[123])\.weight_scale$"): ".weight_scale_inv",
}
else:
# FP8 experts use Fp8MoEMethod (block_quant=True), which registers
@@ -771,6 +872,14 @@ def _make_deepseek_v4_weights_mapper(expert_dtype: str) -> WeightsMapper:
scale_regex = {
re.compile(r"\.scale$"): ".weight_scale_inv",
}
# When shared experts are fused into the routed MXFP4 grouped GEMM, the
# shared_experts tensors are redirected to routed expert slots ; leave
# their names untouched here.
substr_map = (
{}
if fuse_shared_experts
else {".shared_experts.w2": ".shared_experts.down_proj"}
)
return WeightsMapper(
orig_to_new_prefix={
"layers.": "model.layers.",
@@ -785,9 +894,7 @@ def _make_deepseek_v4_weights_mapper(expert_dtype: str) -> WeightsMapper:
"embed.weight": "embed_tokens.weight",
".ffn.gate.bias": ".ffn.gate.e_score_correction_bias",
},
orig_to_new_substr={
".shared_experts.w2": ".shared_experts.down_proj",
},
orig_to_new_substr=substr_map,
)
@@ -804,8 +911,11 @@ class DeepseekV4ForCausalLM(nn.Module, SupportsPP, SupportsEagle3):
config = vllm_config.model_config.hf_config
self.config = config
expert_dtype = getattr(config, "expert_dtype", "fp4")
if expert_dtype != "fp4":
self.hf_to_vllm_mapper = _make_deepseek_v4_weights_mapper(expert_dtype)
fuse_shared_experts = _fuse_shared_experts_enabled(config)
if expert_dtype != "fp4" or fuse_shared_experts:
self.hf_to_vllm_mapper = _make_deepseek_v4_weights_mapper(
expert_dtype, fuse_shared_experts=fuse_shared_experts
)
self.model = self.model_cls(
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
+17
View File
@@ -334,6 +334,21 @@ class DeepSeekV4MTP(nn.Module):
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
def _resolve_scale_name(name: str) -> str:
# Quark checkpoints name FP8 block scales ``.weight_scale``,
# but block-FP8 layers register them as ``.weight_scale_inv``
# while MXFP4 experts register ``.weight_scale``. Auto-detect:
# rename to ``_inv`` only when that variant exists and the plain
# one does not.
if name.endswith(".weight_scale") and name not in params_dict:
inv = name.removesuffix(".weight_scale") + ".weight_scale_inv"
if inv in params_dict:
return inv
# Otherwise leave the name unchanged: either it already matches a
# param, or it is genuinely unknown and should surface the normal
# KeyError downstream rather than be silently rewritten.
return name
# TP for attention
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tensor_model_parallel_rank()
@@ -393,6 +408,7 @@ class DeepSeekV4MTP(nn.Module):
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
name = _resolve_scale_name(name)
param = params_dict[name]
weight_loader = param.weight_loader
@@ -447,6 +463,7 @@ class DeepSeekV4MTP(nn.Module):
)
if name.endswith(".ffn.gate.bias"):
name = name.replace(".bias", ".e_score_correction_bias")
name = _resolve_scale_name(name)
param = params_dict[name]
weight_loader = getattr(
param, "weight_loader", default_weight_loader
+11 -5
View File
@@ -26,6 +26,7 @@ from vllm.models.deepseek_v4.common.ops import (
fused_indexer_q_rope_quant,
fused_q_kv_rmsnorm,
)
from vllm.models.deepseek_v4.common.ops.fused_indexer_q import MXFP4_BLOCK_SIZE
if TYPE_CHECKING:
from vllm.v1.attention.backends.mla.sparse_swa import (
@@ -727,11 +728,16 @@ class DeepseekV4Indexer(nn.Module):
)
assert cache_config is not None, "Deepseek V4 indexer requires cache_config"
# NOTE(yifan): FP8 indxer cache use the same layout as V3.2:
# head_dim bytes = 128 fp8 + 4 fp32 scale = 132.
# For FP4 indexer cache, we still allocate the same amount of memory as FP8,
# but only use the first half of the memory.
k_cache_head_dim = self.head_dim + self.head_dim // self.quant_block_size * 4
if self.use_fp4_kv:
# MXFP4 stores two values per byte plus one UE8M0 byte per 32 values.
# head_dim bytes = 64 packed values + 4 UE8M0 scales = 68.
k_cache_head_dim = self.head_dim // 2 + self.head_dim // MXFP4_BLOCK_SIZE
else:
# NOTE(yifan): FP8 indexer cache uses the same layout as V3.2:
# head_dim bytes = 128 fp8 + 4 fp32 scale = 132.
k_cache_head_dim = (
self.head_dim + self.head_dim // self.quant_block_size * 4
)
self.k_cache = DeepseekV4IndexerCache(
head_dim=k_cache_head_dim,
dtype=torch.uint8,
+32 -2
View File
@@ -7,7 +7,7 @@ from typing import Any, ClassVar, cast
import torch
from torch import nn
from vllm.config import VllmConfig, get_current_vllm_config
from vllm.config import CUDAGraphMode, VllmConfig, get_current_vllm_config
from vllm.forward_context import get_forward_context
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.model_executor.layers.layernorm import RMSNorm
@@ -42,6 +42,19 @@ def _prefer_two_stage_compressor() -> bool:
return current_platform.is_rocm()
def _get_c128_boundary(metadata: CommonAttentionMetadata) -> bool | None:
starts = metadata._num_computed_tokens_cpu
if starts is None:
return None
starts_list = starts.tolist()
query_start_loc = metadata.query_start_loc_cpu.tolist()
return any(
start % 128 + query_start_loc[i + 1] - query_start_loc[i] >= 128
for i, start in enumerate(starts_list)
)
class CompressorBackend(AttentionBackend):
def __init__(self):
super().__init__()
@@ -90,6 +103,7 @@ class CompressorMetadata:
token_to_req_indices: torch.Tensor | None = None # [num_tokens]
num_decode_tokens: int | None = None
c128_boundary: bool | None = None
class CompressorMetadataBuilder(AttentionMetadataBuilder):
@@ -127,6 +141,11 @@ class CompressorMetadataBuilder(AttentionMetadataBuilder):
block_size=self.block_size,
token_to_req_indices=token_to_req_indices,
num_decode_tokens=num_decode_tokens,
c128_boundary=(
_get_c128_boundary(common_attn_metadata)
if self.block_size == 8
else None
),
)
@@ -317,7 +336,8 @@ class DeepseekCompressor(nn.Module):
)
# Get the metadata and handle dummy profiling run.
attn_metadata = get_forward_context().attn_metadata
forward_context = get_forward_context()
attn_metadata = forward_context.attn_metadata
if not isinstance(attn_metadata, dict):
return
@@ -359,6 +379,16 @@ class DeepseekCompressor(nn.Module):
pdl_kwargs=pdl_kwargs,
)
# full graph cannot branch on per-step CPU metadata after capture
if (
current_platform.is_cuda()
and self.head_dim == 512
and self.compress_ratio == 128
and forward_context.cudagraph_runtime_mode != CUDAGraphMode.FULL
and state_metadata.c128_boundary is False
):
return
# Fused: compress → RMSNorm → RoPE → FP8 quant → KV cache write.
# RoPE requirements (kernel applies forward GPT-J style rotation):
# - is_neox_style=False (interleaved pairs, NOT split-half)
+12 -4
View File
@@ -43,6 +43,7 @@ from vllm.model_executor.models.utils import maybe_prefix
from .model import (
DeepseekV4DecoderLayer,
DeepseekV4Model,
make_deepseek_v4_expert_params_mapping,
)
@@ -277,6 +278,11 @@ class DSparkDeepseekV4ForCausalLM(nn.Module):
assert vllm_config.speculative_config is not None
self.draft_model_config = vllm_config.speculative_config.draft_model_config
self.config = self.draft_model_config.hf_config
self.quant_config = vllm_config.quant_config
self.pad_shared_expert = (
getattr(self.quant_config, "weight_block_size", None) is not None
and not vllm_config.parallel_config.use_sequence_parallel_moe
)
self.model = DSparkDeepseekV4Model(
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
)
@@ -396,6 +402,12 @@ class DSparkDeepseekV4ForCausalLM(nn.Module):
else ".weight_scale_inv"
)
name = name.removesuffix(".scale") + suffix
if ".shared_experts.w2" in name:
name = name.replace(".shared_experts.w2", ".shared_experts.down_proj")
if self.pad_shared_expert and ".shared_experts." in name:
loaded_weight = DeepseekV4Model._pad_shared_expert_weight(
self.quant_config, name, loaded_weight
)
# E8M0 expert scales: keep raw exponent bytes.
if ".experts." in name:
@@ -440,10 +452,6 @@ class DSparkDeepseekV4ForCausalLM(nn.Module):
params_dict[name][: narrow.shape[0]].copy_(narrow)
loaded_params.add(name)
continue
if ".shared_experts.w2" in name:
name = name.replace(
".shared_experts.w2", ".shared_experts.down_proj"
)
if name.endswith(".ffn.gate.bias"):
name = name.replace(
".ffn.gate.bias", ".ffn.gate.e_score_correction_bias"
+8 -3
View File
@@ -1181,7 +1181,9 @@ class DeepseekV4Model(nn.Module, EagleModelMixin):
for name, loaded_weight in weights:
if pad_shared_expert and ".shared_experts." in name:
loaded_weight = self._pad_shared_expert_weight(name, loaded_weight)
loaded_weight = self._pad_shared_expert_weight(
self.quant_config, name, loaded_weight
)
for param_name, weight_name, shard_id in stacked_params_mapping:
# Skip non-stacked layers and experts (experts handled below).
if ".experts." in name:
@@ -1256,15 +1258,18 @@ class DeepseekV4Model(nn.Module, EagleModelMixin):
return loaded_params
@staticmethod
def _pad_shared_expert_weight(
self, name: str, loaded_weight: torch.Tensor
quant_config: QuantizationConfig | None,
name: str,
loaded_weight: torch.Tensor,
) -> torch.Tensor:
"""Zero-pad a block-FP8 shared-expert weight/scale on its intermediate
axis so the standard TP loaders split it into even, block-aligned shards
(trailing ranks get the zero pad). gate (w1)/up (w3) [I, H] pad dim 0;
down (w2 -> down_proj) [H, I] pads dim 1.
"""
block_size = getattr(self.quant_config, "weight_block_size", None)
block_size = getattr(quant_config, "weight_block_size", None)
assert block_size is not None
# Round the intermediate axis up to a whole number of TP shards. The axis
# is in elements for weights (step = block) and in blocks for scales.
+11 -4
View File
@@ -52,6 +52,7 @@ from vllm.sequence import IntermediateTensors
from .model import (
DeepseekV4DecoderLayer,
DeepseekV4Model,
make_deepseek_v4_expert_params_mapping,
)
@@ -265,6 +266,10 @@ class DeepSeekV4MTP(nn.Module):
super().__init__()
self.config = vllm_config.model_config.hf_config
self.quant_config = vllm_config.quant_config
self.pad_shared_expert = (
getattr(self.quant_config, "weight_block_size", None) is not None
and not vllm_config.parallel_config.use_sequence_parallel_moe
)
self.model = DeepSeekV4MultiTokenPredictor(
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
)
@@ -387,6 +392,12 @@ class DeepSeekV4MTP(nn.Module):
else ".weight_scale_inv"
)
name = name.removesuffix(".scale") + suffix
if ".shared_experts.w2" in name:
name = name.replace(".shared_experts.w2", ".shared_experts.down_proj")
if self.pad_shared_expert and ".shared_experts." in name:
loaded_weight = DeepseekV4Model._pad_shared_expert_weight(
self.quant_config, name, loaded_weight
)
for param_name, weight_name, shard_id in stacked_params_mapping:
# Skip non-stacked layers and experts (experts handled below).
if ".experts." in name:
@@ -442,10 +453,6 @@ class DeepSeekV4MTP(nn.Module):
loaded_params.add(name)
continue
else:
if ".shared_experts.w2" in name:
name = name.replace(
".shared_experts.w2", ".shared_experts.down_proj"
)
if name.endswith(".ffn.gate.bias"):
# ``e_score_correction_bias`` lives on the gate
# under a different attribute name.
+37 -2
View File
@@ -4,7 +4,7 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, cast
from vllm.config import get_current_vllm_config
from vllm.model_executor.layers.fused_moe import (
@@ -117,13 +117,29 @@ class DeepseekV4FP8Config(Fp8Config):
def get_name(cls) -> QuantizationMethods:
return "deepseek_v4_fp8"
@staticmethod
def _is_quark_mxfp4_ocp(hf_quant_cfg: dict) -> bool:
"""True for AMD-Quark exports whose global scheme is MXFP4."""
weight = (hf_quant_cfg.get("global_quant_config") or {}).get("weight") or {}
return (
weight.get("dtype") == "fp4"
and weight.get("qscheme") == "per_group"
and weight.get("group_size") == 32
)
@classmethod
def override_quantization_method(
cls, hf_quant_cfg, user_quant, hf_config=None
) -> QuantizationMethods | None:
if not (
isinstance(hf_quant_cfg, dict)
and hf_quant_cfg.get("quant_method") in ("fp8", "deepseek_v4_fp8")
and (
hf_quant_cfg.get("quant_method") in ("fp8", "deepseek_v4_fp8")
or (
hf_quant_cfg.get("quant_method") == "quark"
and cls._is_quark_mxfp4_ocp(hf_quant_cfg)
)
)
):
return None
model_type = getattr(hf_config, "model_type", None)
@@ -131,6 +147,25 @@ class DeepseekV4FP8Config(Fp8Config):
return "deepseek_v4_fp8"
return None
@classmethod
def from_config(cls, config: dict) -> DeepseekV4FP8Config:
# Reroute AMD-Quark fused shared expert MXFP4 checkpoints onto the fp8
# path: the runtime layout matches the DeepSeek-native fp8 checkpoint,
# so translate the schema into format Fp8Config.from_config expects.
if config.get("quant_method") == "quark":
quark_exclude = config.get("exclude") or []
config = {
"quant_method": "fp8",
"activation_scheme": "dynamic",
"fmt": "e4m3",
"scale_fmt": "ue8m0",
"weight_block_size": [128, 128],
"ignored_layers": [
name for name in quark_exclude if isinstance(name, str)
],
}
return cast("DeepseekV4FP8Config", super().from_config(config))
def get_quant_method(self, layer, prefix):
if isinstance(layer, RoutedExperts):
if is_layer_skipped(
+15 -2
View File
@@ -9,6 +9,7 @@ Both use vLLM's standard ``RMSNorm`` (CPU-friendly, with a native fallback).
from __future__ import annotations
from itertools import combinations
from typing import cast
import numpy as np
@@ -45,6 +46,20 @@ def _prime_factors(n: int) -> list[int]:
return factors
def linear_sum_assignment(
cost_matrix: np.ndarray,
) -> tuple[np.ndarray, np.ndarray]:
"""Implement SciPy's assignment for Inkling's ordered L1 cost matrix."""
rows = np.arange(cost_matrix.shape[0])
cols = np.array(
min(
combinations(range(cost_matrix.shape[1]), len(rows)),
key=lambda candidate: cost_matrix[rows, candidate].sum(),
)
)
return rows, cols
def plan_out_scales(
temporal_patch_size: int, patch_size: int, n_layers: int, n_channels: int = 3
) -> list[tuple[int, int, int, int]]:
@@ -97,8 +112,6 @@ def plan_out_scales(
if n_layers >= len(scales):
idxs = np.argmin(cost_matrix, axis=1)
else:
from scipy.optimize import linear_sum_assignment
idxs = linear_sum_assignment(cost_matrix)[1]
assert len(idxs) >= 2
+1 -1
View File
@@ -58,6 +58,6 @@ class InklingDenseMLP(nn.Module):
x = silu_and_mul_triton(gate_up)
x, _ = self.down_proj(x)
if self.global_scale is not None:
x = x * self.global_scale
x.mul_(self.global_scale)
# TP-partial output: the layer's reduce-scatter fallback consumes it.
return x
+1 -1
View File
@@ -528,7 +528,7 @@ class InklingMoE(nn.Module):
)
self._routed_sel = None
return out + sink_out
return out.add_(sink_out)
# -- weight loading ----------------------------------------------------
+28
View File
@@ -593,6 +593,20 @@ class MediaConnector:
return image_embedding_io.load_base64("", data)
async def fetch_image_embedding_async(
self,
data: str,
) -> torch.Tensor:
"""
Asynchronously load image embedding from a URL.
"""
image_embedding_io = ImageEmbeddingMediaIO()
loop = asyncio.get_running_loop()
return await loop.run_in_executor(
global_thread_pool, image_embedding_io.load_base64, "", data
)
def fetch_audio_embedding(
self,
data: str,
@@ -603,3 +617,17 @@ class MediaConnector:
audio_embedding_io = AudioEmbeddingMediaIO()
return audio_embedding_io.load_base64("", data)
async def fetch_audio_embedding_async(
self,
data: str,
) -> torch.Tensor:
"""
Asynchronously load audio embedding from a URL.
"""
audio_embedding_io = AudioEmbeddingMediaIO()
loop = asyncio.get_running_loop()
return await loop.run_in_executor(
global_thread_pool, audio_embedding_io.load_base64, "", data
)
@@ -304,14 +304,7 @@ class StreamingParserEngine:
transition = self.config.transitions.get(key)
if transition is None:
if (
self._has_drops
and terminal == DROP_TERMINAL
# Preserve drop tokens when skip_tool_parsing is active so
# the reasoning pass doesn't silently remove tokens that a
# later tool-call pass might need to see.
and not self.skip_tool_parsing
):
if self._has_drops and terminal == DROP_TERMINAL:
return []
return self._emit_for_state(value)
+11 -7
View File
@@ -731,13 +731,17 @@ def get_config(
raise ValueError(error_message) from e
config_parser = get_config_parser(config_format)
config_dict, config = config_parser.parse(
model,
trust_remote_code=trust_remote_code,
revision=revision,
code_revision=code_revision,
hf_overrides=hf_overrides_kw or hf_overrides_fn,
**kwargs,
# Retry to tolerate a concurrent HF cache refresh briefly hiding config.json.
config_dict, config = with_retry(
lambda: config_parser.parse(
model,
trust_remote_code=trust_remote_code,
revision=revision,
code_revision=code_revision,
hf_overrides=hf_overrides_kw or hf_overrides_fn,
**kwargs,
),
f"Error parsing config for {model}",
)
# Architecture mapping for models without explicit architectures field
+9
View File
@@ -634,6 +634,15 @@ class HybridKVCacheCoordinator(KVCacheCoordinator):
key=lambda g: not isinstance(g.spec, FullAttentionSpec)
)
# Dense reference group for per-group lookups (None when the model
# has no full-attention layers): full attention is downward-closed,
# so any group reporting a longer per-group hit implies the union of
# per-group hits is not consistent at a single boundary (#46453).
first = self.attention_groups[0]
self.full_attention_group_id: int | None = (
first.group_ids[0] if isinstance(first.spec, FullAttentionSpec) else None
)
# Propagate the eagle bit to each manager (default to ``use_eagle=False``).
for group in self.attention_groups:
if group.use_eagle:
+51 -1
View File
@@ -9,7 +9,10 @@ from typing import Literal, overload
from vllm.distributed.kv_events import BlockStored, KVCacheEvent
from vllm.logger import init_logger
from vllm.utils.math_utils import cdiv
from vllm.v1.core.kv_cache_coordinator import get_kv_cache_coordinator
from vllm.v1.core.kv_cache_coordinator import (
HybridKVCacheCoordinator,
get_kv_cache_coordinator,
)
from vllm.v1.core.kv_cache_metrics import KVCacheMetricsCollector
from vllm.v1.core.kv_cache_utils import KVCacheBlock, KVCacheBlockCopy
from vllm.v1.kv_cache_interface import (
@@ -287,6 +290,53 @@ class KVCacheManager:
blocks = self.create_kv_cache_blocks(computed_blocks)
return blocks, num_new_computed_tokens, shared_prefix_boundary
def get_computed_blocks_for_connector(
self, request: Request
) -> tuple[KVCacheBlocks, int, int, bool]:
"""Local prefix-cache lookup for a request scheduled with a KV connector.
Hybrid (Mamba + full-attention) models can have per-group prefix hits
diverge under block pressure: the full-attention tail may be evicted
while a deeper Mamba state survives, or vice versa. Report the
full-attention hit as the local prefix - the connector transfers the
remaining suffix and the Mamba state is transferred unconditionally by
nixl's ``_apply_prefix_caching`` - and flag when that hit ran deeper
than a lagging group. Such a hit only has a valid Mamba state at its
boundary if the connector supplies it, so the caller must fall back to
``get_computed_blocks`` to reconcile when no external tokens are found.
Non-hybrid models and already-convergent hits use ``get_computed_blocks``.
Returns:
The ``get_computed_blocks`` triple (blocks, number of local computed
tokens, shared-prefix boundary) plus ``hit_diverged``.
"""
coordinator = self.coordinator
if not (
self.kv_cache_config.has_mamba_layers
and isinstance(coordinator, HybridKVCacheCoordinator)
and coordinator.full_attention_group_id is not None
):
return *self.get_computed_blocks(request), False
if not self.prefix_cache_lookup_enabled(request):
return self.empty_kv_cache_blocks, 0, 0, False
fa_group_id = coordinator.full_attention_group_id
computed, per_group_hits = coordinator.find_longest_cache_hit_per_group(
request.block_hashes, request.num_tokens - 1
)
if any(hit > per_group_hits[fa_group_id] for hit in per_group_hits):
# A lagging group hit deeper than full attention means its
# full-attention blocks were evicted; use the reconciled boundary
# that every group agrees on.
return *self.get_computed_blocks(request), False
num_local = per_group_hits[fa_group_id]
blocks = self.create_kv_cache_blocks(computed)
# Per-group lookups do not detect an uncached shared prefix (boundary 0).
return blocks, num_local, 0, min(per_group_hits) < num_local
def allocate_slots(
self,
request: Request,
+34 -56
View File
@@ -983,9 +983,8 @@ def _pool_bytes_per_block(
):
return kv_cache_groups[0].kv_cache_spec.page_size_bytes
if _use_packed_kv_cache_config(vllm_config, kv_cache_groups):
# buckets = {page_size: [[layer_names], [layer_names], ...]}
buckets = _bucket_layers_by_page_size(kv_cache_groups)
return sum(ps * len(slots) for ps, slots in buckets.items())
block_stride, _ = _get_packed_kv_cache_layout(kv_cache_groups)
return block_stride
group_size = max(len(g.layer_names) for g in kv_cache_groups)
page_size = get_uniform_page_size([g.kv_cache_spec for g in kv_cache_groups])
return page_size * group_size
@@ -1260,29 +1259,29 @@ def _get_kv_cache_groups_uniform_page_size(
return create_kv_cache_group_specs(kv_cache_spec, grouped_layers)
def _bucket_layers_by_page_size(
def _get_packed_kv_cache_layout(
kv_cache_groups: list[KVCacheGroupSpec],
) -> dict[int, list[list[str]]]:
"""Bucket layers by page size: ``result[ps][slot_idx] = [layer_names]``.
) -> tuple[int, dict[int, list[str]]]:
"""Lay out each cache group densely in one shared block slab.
Layers from different groups at the same ``slot_idx`` share an underlying tensor
(they have independent block tables so block-id namespaces never collide).
A block ID is owned by one cache group at a time, so layouts from different
groups may overlap. Layers within a group remain disjoint.
"""
buckets: dict[int, list[list[str]]] = defaultdict(list)
layers_by_offset: dict[int, list[str]] = defaultdict(list)
block_stride = 0
for group in kv_cache_groups:
spec = group.kv_cache_spec
slot_count: dict[int, int] = defaultdict(int)
byte_offset = 0
for layer_name in group.layer_names:
if isinstance(spec, UniformTypeKVCacheSpecs):
ps = spec.kv_cache_specs[layer_name].page_size_bytes
page_size = spec.kv_cache_specs[layer_name].page_size_bytes
else:
ps = spec.page_size_bytes
slot_idx = slot_count[ps]
slot_count[ps] += 1
if slot_idx == len(buckets[ps]):
buckets[ps].append([])
buckets[ps][slot_idx].append(layer_name)
return buckets
page_size = spec.page_size_bytes
layers_by_offset[byte_offset].append(layer_name)
byte_offset += page_size
block_stride = max(block_stride, byte_offset)
assert block_stride > 0
return block_stride, layers_by_offset
def _use_packed_kv_cache_config(
@@ -1314,33 +1313,26 @@ def _get_kv_cache_config_packed(
) -> tuple[int, list[KVCacheTensor]]:
"""Plan a packed per-block KV cache tensor layout.
Emit one KVCacheTensor per (slot_idx, page_size). Layers from different
groups at the same slot share a tensor (they have independent block
tables so block-id namespaces never collide). Each emitted tensor aliases
one physical backing allocation, with per-block data laid out contiguously.
Cache groups use dense, overlapping layouts within one block slab. Each
emitted tensor aliases the same physical backing allocation.
"""
# buckets = {page_size: [[layer_names], [layer_names], ...]}
buckets = _bucket_layers_by_page_size(kv_cache_groups)
total_num_bytes_per_block = sum(ps * len(slots) for ps, slots in buckets.items())
block_stride, layers_by_offset = _get_packed_kv_cache_layout(kv_cache_groups)
num_blocks = available_memory // total_num_bytes_per_block
num_blocks = available_memory // block_stride
num_blocks = may_override_num_blocks(vllm_config, num_blocks)
total_size = total_num_bytes_per_block * num_blocks
total_size = block_stride * num_blocks
kv_cache_tensors: list[KVCacheTensor] = []
byte_offset = 0
for ps, slots in buckets.items():
for slot in slots:
kv_cache_tensors.append(
KVCacheTensor(
size=total_size,
shared_by=slot,
offset=byte_offset,
block_stride=total_num_bytes_per_block,
)
for byte_offset in sorted(layers_by_offset):
kv_cache_tensors.append(
KVCacheTensor(
size=total_size,
shared_by=layers_by_offset[byte_offset],
offset=byte_offset,
block_stride=block_stride,
)
byte_offset += ps
)
return num_blocks, kv_cache_tensors
@@ -1651,29 +1643,15 @@ def _get_kv_cache_groups_uniform_groups(
for spec in group.kv_cache_specs.values()
)
# Split each SWA UniformKV group into smaller groups to align their #(layer tuples)
# Possibly padding layer tuples for this.
# Additionally, we also pad KV blocks in each SWA layer, to align the page size
# with the corresponding layer in the full-MLA group.
all_page_sizes = full_mla_spec.get_page_sizes()
# Split each SWA UniformKV group into smaller groups to align their
# numbers of layer tuples. The packed block planner overlays groups, so
# their page sizes do not need to match.
swa_mla_groups = []
for sm_spec in swa_mla_specs:
sm_page_sizes = sm_spec.get_page_sizes()
layers_per_size: dict[int, list[str]] = defaultdict(list)
assert max(sm_page_sizes) <= max(all_page_sizes)
# Unify page size by padding layers' page_size to the nearest larger page_size.
# Compute candidate (nearest larger page_size) for each unique page size.
size_to_candidate: dict[int, int] = {}
for ps in sm_page_sizes:
size_to_candidate[ps] = min(x for x in all_page_sizes if x >= ps)
# Pad and collect layer names per page size.
for layer_name, layer_spec in sm_spec.kv_cache_specs.items():
current_size = layer_spec.page_size_bytes
candidate = size_to_candidate[current_size]
if current_size < candidate:
object.__setattr__(layer_spec, "page_size_padded", candidate)
layers_per_size[candidate].append(layer_name)
layers_per_size[layer_spec.page_size_bytes].append(layer_name)
# NOTE(yifan): for now, inside a UniformKV group, each page_size should
# have the same number of layers. This also means we don't need to pad layers
# inside a partial-full layer tuple.
+3 -3
View File
@@ -145,7 +145,7 @@ class SchedulerInterface(ABC):
self,
request_ids: str | Iterable[str] | None,
finished_status: "RequestStatus",
) -> list[tuple[str, int]]:
) -> "list[Request]":
"""Finish the requests in the scheduler's internal queue. If the request
is not in the queue, this method will do nothing for that request.
@@ -159,8 +159,8 @@ class SchedulerInterface(ABC):
finished_status: The finished status of the given requests.
Returns:
Tuple of (req_id, client_index) for requests that were aborted. Will not
include any that were already finished.
List of requests that were aborted. Will not include any that were
already finished.
"""
raise NotImplementedError
+46 -51
View File
@@ -35,7 +35,6 @@ from vllm.v1.core.encoder_cache_manager import (
EncoderCacheManager,
EncoderDecoderCacheManager,
)
from vllm.v1.core.kv_cache_coordinator import HybridKVCacheCoordinator
from vllm.v1.core.kv_cache_manager import KVCacheBlocks, KVCacheManager
from vllm.v1.core.kv_cache_metrics import KVCacheMetricsCollector
from vllm.v1.core.kv_cache_utils import KVCacheBlock
@@ -61,7 +60,7 @@ from vllm.v1.outputs import DraftTokenIds, KVConnectorOutput, ModelRunnerOutput
from vllm.v1.request import Request, RequestStatus, StreamingUpdate
from vllm.v1.spec_decode.dynamic.utils import build_dynamic_sd_schedule_lookup
from vllm.v1.spec_decode.metrics import SpecDecodingStats
from vllm.v1.structured_output import StructuredOutputManager
from vllm.v1.structured_output import StructuredOutputGrammar, StructuredOutputManager
from vllm.v1.utils import record_function_or_nullcontext
logger = init_logger(__name__)
@@ -202,6 +201,10 @@ class Scheduler(SchedulerInterface):
self.finished_recving_kv_req_ids: set[str] = set()
self.failed_recving_kv_req_ids: set[str] = set()
# Grammar compilation failures to finish as per-request errors in
# update_from_output.
self.grammar_compile_error_reqs: set[str] = set()
# Encoder-related.
# Calculate encoder cache size if applicable
supports_mm_inputs = mm_registry.supports_multimodal_inputs(
@@ -711,50 +714,24 @@ class Scheduler(SchedulerInterface):
# Get already-cached tokens.
if request.num_computed_tokens == 0:
did_prefix_cache_lookup = True
hit_diverged = False
# Get locally-cached tokens.
if (
self.connector is not None
and self.has_mamba_layers
and isinstance(
self.kv_cache_manager.coordinator, HybridKVCacheCoordinator
if self.connector is not None:
# A KV connector transfers the missing suffix, which needs a
# hybrid-aware lookup that can diverge across groups.
(
new_computed_blocks,
num_new_local_computed_tokens,
request.shared_prefix_boundary,
hit_diverged,
) = self.kv_cache_manager.get_computed_blocks_for_connector(
request
)
):
# The per-group lookup does not detect an uncached shared
# prefix, so there is no junction to pin in this path.
request.shared_prefix_boundary = 0
kv_cache_manager = self.kv_cache_manager
if not kv_cache_manager.prefix_cache_lookup_enabled(request):
# Mirror the get_computed_blocks() early-out: the
# request must recompute its prompt.
new_computed_blocks = kv_cache_manager.empty_kv_cache_blocks
num_new_local_computed_tokens = 0
else:
computed, per_group_hits = (
self.kv_cache_manager.coordinator.find_longest_cache_hit_per_group(
request.block_hashes, request.num_tokens - 1
)
)
new_computed_blocks = (
self.kv_cache_manager.create_kv_cache_blocks(computed)
)
# NOTE(ZhanqiuHu): For Mamba hybrid models,
# num_new_local_computed_tokens should be the FA hit
# length. This value is passed to the connector's
# get_num_new_matched_tokens which computes:
# external = total - local_computed.
# Using the FA hit skips re-transferring FA blocks
# already cached on D-side. The Mamba state (always
# the last block) is transferred unconditionally by
# _apply_prefix_caching in nixl/worker.py.
num_new_local_computed_tokens = max(per_group_hits)
else:
(
new_computed_blocks,
num_new_local_computed_tokens,
# Junction to pin (Marconi-style APC) so its
# sparse-retention state (Mamba block / sliding-window
# tail) survives retention and serves a later hit; 0
# if no uncached shared prefix was detected.
# Marconi shared-prefix junction to pin; 0 if none.
request.shared_prefix_boundary,
) = self.kv_cache_manager.get_computed_blocks(request)
@@ -776,6 +753,16 @@ class Scheduler(SchedulerInterface):
num_external_computed_tokens = ext_tokens
if hit_diverged and num_external_computed_tokens == 0:
# No external tokens back the deeper local hit, so its
# resume boundary would have no valid Mamba state.
# Reconcile to the boundary every group agrees on.
(
new_computed_blocks,
num_new_local_computed_tokens,
request.shared_prefix_boundary,
) = self.kv_cache_manager.get_computed_blocks(request)
connector_prefix_cache_queries = (
request.num_tokens - num_new_local_computed_tokens
)
@@ -1733,7 +1720,7 @@ class Scheduler(SchedulerInterface):
struct_output_request = request.structured_output_request
assert struct_output_request is not None
grammar = struct_output_request.grammar
assert grammar is not None
assert isinstance(grammar, StructuredOutputGrammar)
# new_token_ids can be a mixed block of reasoning content, then
# the reasoning end marker, then the start of the grammar content.
# Trim the reasoning content so the grammar only sees grammar content.
@@ -1863,10 +1850,16 @@ class Scheduler(SchedulerInterface):
# This is a rare case and unlikely to impact performance.
self.waiting.remove_requests(stopped_preempted_reqs)
error_req_ids = set(self.grammar_compile_error_reqs)
self.grammar_compile_error_reqs.clear()
if failed_kv_load_req_ids and not self.recompute_kv_load_failures:
requests = [self.requests[req_id] for req_id in failed_kv_load_req_ids]
self.finish_requests(failed_kv_load_req_ids, RequestStatus.FINISHED_ERROR)
for request in requests:
error_req_ids.update(failed_kv_load_req_ids)
if error_req_ids:
error_reqs = self.finish_requests(
error_req_ids, RequestStatus.FINISHED_ERROR
)
for request in error_reqs:
outputs[request.client_index].append(
EngineCoreOutput(
request_id=request.request_id,
@@ -2096,8 +2089,7 @@ class Scheduler(SchedulerInterface):
# Filter out spec tokens which do not adhere to the grammar.
if self.structured_output_manager.should_advance(request):
metadata = request.structured_output_request
assert metadata is not None and metadata.grammar is not None
spec_token_ids = metadata.grammar.validate_tokens(spec_token_ids)
spec_token_ids = metadata.grammar.validate_tokens(spec_token_ids) # type: ignore[union-attr]
# Pad to original number of spec tokens.
num_invalid_tokens = orig_num_spec_tokens - len(spec_token_ids)
if num_invalid_tokens:
@@ -2138,7 +2130,7 @@ class Scheduler(SchedulerInterface):
def finish_requests(
self, request_ids: str | Iterable[str] | None, finished_status: RequestStatus
) -> list[tuple[str, int]]:
) -> list[Request]:
"""Handles the finish signal from outside the scheduler.
For example, the API server can abort a request when the client
@@ -2147,8 +2139,8 @@ class Scheduler(SchedulerInterface):
If request_ids is None, all requests will be finished.
Returns:
Tuple of (req_id, client_index) for requests that were aborted. Will not
include any that were already finished.
List of requests that were aborted. Will not include any that were
already finished.
"""
assert RequestStatus.is_finished(finished_status)
if isinstance(request_ids, str):
@@ -2197,7 +2189,7 @@ class Scheduler(SchedulerInterface):
request.status = finished_status
self._free_request(request, delay_free_blocks=delay_free_blocks)
return [(r.request_id, r.client_index) for r in valid_requests]
return valid_requests
def _free_request(
self, request: Request, delay_free_blocks: bool = False
@@ -2597,7 +2589,10 @@ class Scheduler(SchedulerInterface):
if request.status == RequestStatus.WAITING_FOR_STRUCTURED_OUTPUT_GRAMMAR:
structured_output_req = request.structured_output_request
if not (structured_output_req and structured_output_req.grammar):
if not structured_output_req or structured_output_req.grammar is None:
return False
if isinstance(structured_output_req.grammar, Exception):
self.grammar_compile_error_reqs.add(request.request_id)
return False
request.status = RequestStatus.WAITING
return True
+3 -3
View File
@@ -1830,13 +1830,13 @@ class EngineCoreProc(EngineCore):
) -> None:
self._send_finish_outputs_to_client(req_ids, client_index, FinishReason.ERROR)
def _send_abort_outputs(self, aborted_reqs: list[tuple[str, int]]) -> None:
def _send_abort_outputs(self, aborted_reqs: list[Request]) -> None:
# TODO(nick) this will be moved inside the scheduler
if aborted_reqs:
# Map client_index to list of request_ids that belong to that client.
by_client = defaultdict[int, set[str]](set)
for req_id, client_index in aborted_reqs:
by_client[client_index].add(req_id)
for request in aborted_reqs:
by_client[request.client_index].add(request.request_id)
for client_index, req_ids in by_client.items():
self._send_abort_outputs_to_client(list(req_ids), client_index)
+28 -10
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import sys
from abc import ABC, abstractmethod
import tokenizers
@@ -320,10 +321,19 @@ def check_stop_strings(
Where stop_string is the matched stop string and offset is the
length to which output_text should be truncated, or -1 for no
truncation.
When several stop strings match within the newly generated text (for
example when speculative decoding appends multiple tokens in a single
step), the stop string that completes earliest in the text is selected,
so the result matches appending one token at a time. Ties are broken by
stop-list order.
"""
if not new_char_count or not stop:
return None
best_stop_str: str | None = None
best_stop_index = 0
best_end = sys.maxsize
for stop_str in stop:
stop_string_len = len(stop_str)
# Avoid searching already-searched text.
@@ -331,14 +341,22 @@ def check_stop_strings(
if stop_index == -1:
continue
if include_in_output:
# Truncate to end of stop string.
stop_index += stop_string_len
if stop_index >= len(output_text):
# No truncation required.
return stop_str, -1
# Prefer the stop string that completes earliest in the text.
end = stop_index + stop_string_len
if end < best_end:
best_stop_str = stop_str
best_stop_index = stop_index
best_end = end
# Truncate the output text to either the beginning
# or end of the stop string.
return stop_str, stop_index
return None
if best_stop_str is None:
return None
if include_in_output:
# Truncate to end of stop string.
if best_end >= len(output_text):
# No truncation required.
return best_stop_str, -1
return best_stop_str, best_end
# Truncate the output text to the beginning of the stop string.
return best_stop_str, best_stop_index
+27
View File
@@ -2,6 +2,7 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from abc import ABC, abstractmethod
from collections.abc import Sequence
from copy import copy
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, NamedTuple, TypeAlias
@@ -90,6 +91,32 @@ class LogprobsTensors(NamedTuple):
self.selected_token_ranks[mask],
)
@staticmethod
def cat(
tensors: Sequence["LogprobsTensors"],
cu_num_generated_tokens: list[int] | None = None,
) -> "LogprobsTensors":
"""Concatenate flattened logprob tensors."""
assert tensors
assert cu_num_generated_tokens is not None or all(
tensor.cu_num_generated_tokens is None for tensor in tensors
)
if len(tensors) == 1:
tensor = tensors[0]
if cu_num_generated_tokens is None:
return tensor
return tensor._replace(cu_num_generated_tokens=cu_num_generated_tokens)
return LogprobsTensors(
logprob_token_ids=torch.cat(
[tensor.logprob_token_ids for tensor in tensors]
),
logprobs=torch.cat([tensor.logprobs for tensor in tensors]),
selected_token_ranks=torch.cat(
[tensor.selected_token_ranks for tensor in tensors]
),
cu_num_generated_tokens=cu_num_generated_tokens,
)
@staticmethod
def empty_cpu(
num_positions: int, num_tokens_per_position: int
+8 -12
View File
@@ -7,7 +7,7 @@ import torch.nn as nn
from vllm import envs
from vllm._aiter_ops import rocm_aiter_ops
from vllm.config.model import LogprobsMode
from vllm.config.model import PROCESSED_LOGPROBS_MODES, LogprobsMode
from vllm.logger import init_logger
from vllm.platforms import CpuArchEnum, current_platform
from vllm.triton_utils import HAS_TRITON
@@ -87,7 +87,7 @@ class TopKTopPSampler(nn.Module):
# FlashInfer doesn't expose post-top-k/top-p logits/logprobs,
# so it can't be used when the configured mode requires them.
can_use_flashinfer = (
logprobs_mode not in ("processed_logits", "processed_logprobs")
logprobs_mode not in PROCESSED_LOGPROBS_MODES
and flashinfer_sampler_supported()
)
self.forward = (
@@ -108,7 +108,7 @@ class TopKTopPSampler(nn.Module):
else:
self.forward = self.forward_native
elif (
logprobs_mode not in ("processed_logits", "processed_logprobs")
logprobs_mode not in PROCESSED_LOGPROBS_MODES
and rocm_aiter_ops.is_enabled()
):
self.aiter_ops = None
@@ -165,7 +165,7 @@ class TopKTopPSampler(nn.Module):
return self.forward_native(logits, generators, k, p)
if self.use_fp64_gumbel:
return self.forward_native(logits, generators, k, p)
assert self.logprobs_mode not in ("processed_logits", "processed_logprobs"), (
assert self.logprobs_mode not in PROCESSED_LOGPROBS_MODES, (
"FlashInfer does not support returning logits/logprobs"
)
# flashinfer sampling functions expect contiguous logits.
@@ -236,10 +236,9 @@ class TopKTopPSampler(nn.Module):
return self.forward_native(logits, generators, k, p)
if self.use_fp64_gumbel:
return self.forward_native(logits, generators, k, p)
assert self.logprobs_mode not in (
"processed_logits",
"processed_logprobs",
), "aiter sampler does not support returning logits/logprobs."
assert self.logprobs_mode not in PROCESSED_LOGPROBS_MODES, (
"aiter sampler does not support returning logits/logprobs."
)
if self.aiter_ops is None and not self._init_aiter_ops():
return self.forward_native(logits, generators, k, p)
return self.aiter_sample(logits, k, p, generators), None
@@ -300,10 +299,7 @@ class TopKTopPSampler(nn.Module):
logits.shape[0], dtype=torch.int64, device=logits.device
)
logits_to_return = None
if (
self.logprobs_mode == "processed_logits"
or self.logprobs_mode == "processed_logprobs"
):
if self.logprobs_mode in PROCESSED_LOGPROBS_MODES:
logits_to_return = torch.empty_like(logits)
assert len(generators) != logits.shape[0], (
+2 -4
View File
@@ -10,6 +10,7 @@ from typing import TYPE_CHECKING
import torch
import torch.nn as nn
from vllm.config.model import PROCESSED_LOGPROBS_MODES
from vllm.logger import init_logger
from vllm.triton_utils import tl, triton
from vllm.v1.outputs import LogprobsLists, LogprobsTensors, SamplerOutput
@@ -67,10 +68,7 @@ class RejectionSampler(nn.Module):
self.sampler = sampler
self.use_fp64_gumbel = getattr(sampler, "use_fp64_gumbel", False)
logprobs_mode = self.sampler.logprobs_mode
self.is_processed_logprobs_mode = logprobs_mode in (
"processed_logprobs",
"processed_logits",
)
self.is_processed_logprobs_mode = logprobs_mode in PROCESSED_LOGPROBS_MODES
self.is_logits_logprobs_mode = logprobs_mode in (
"raw_logits",
"processed_logits",
+25 -14
View File
@@ -164,24 +164,33 @@ class StructuredOutputManager:
else:
raise ValueError(f"Unsupported structured output backend: {backend}")
grammar: Future[StructuredOutputGrammar] | StructuredOutputGrammar
if self._use_async_grammar_compilation:
grammar = self.executor.submit(self._create_grammar, request)
else:
grammar = self._create_grammar(request) # type: ignore[assignment]
request.structured_output_request.grammar = grammar # type: ignore[assignment]
try:
grammar = self._create_grammar(request)
except Exception as e:
grammar = Future()
grammar.set_exception(e)
request.structured_output_request.grammar = grammar
def _create_grammar(self, request: "Request") -> StructuredOutputGrammar:
key = request.structured_output_request.structured_output_key # type: ignore[union-attr]
struct_request = request.structured_output_request
assert struct_request is not None
# Note that the request was validated in the engine core client,
# so at this point we know it is a supported type of request.
#
# TODO: we still need to handle xgrammar compilation failures,
# though it should be unlikely as we test that up front as well.
request_type, grammar_spec = key
assert self.backend is not None
return self.backend.compile_grammar(request_type, grammar_spec)
# so at this point we know it is a supported type of request. Grammar
# compilation may still fail; the Future carries that error to the
# scheduler so it can fail only this request.
try:
request_type, grammar_spec = struct_request.structured_output_key
assert self.backend is not None
return self.backend.compile_grammar(request_type, grammar_spec)
except Exception:
logger.exception(
"Failed to compile grammar for request %s", request.request_id
)
raise
def _fill_bitmasks(
self, batch: Iterable[tuple[StructuredOutputGrammar, int, bool]]
@@ -244,8 +253,9 @@ class StructuredOutputManager:
structured_output_request = request.structured_output_request
if TYPE_CHECKING:
assert structured_output_request is not None
assert structured_output_request.grammar is not None
grammar = structured_output_request.grammar
if TYPE_CHECKING:
assert isinstance(grammar, StructuredOutputGrammar)
apply_bitmask = self.should_fill_bitmask(request)
batch.append((grammar, cumulative_index, apply_bitmask))
@@ -268,8 +278,9 @@ class StructuredOutputManager:
if TYPE_CHECKING:
assert structured_output_request is not None
assert structured_output_request.grammar is not None
grammar = structured_output_request.grammar
if TYPE_CHECKING:
assert isinstance(grammar, StructuredOutputGrammar)
apply_bitmask = self.should_fill_bitmask(request)
reasoner = self._get_reasoner(request)
+9 -6
View File
@@ -21,7 +21,9 @@ if TYPE_CHECKING:
@dataclasses.dataclass
class StructuredOutputRequest:
params: StructuredOutputsParams
_grammar: Future[StructuredOutputGrammar] | StructuredOutputGrammar | None = None
_grammar: (
Future[StructuredOutputGrammar] | StructuredOutputGrammar | Exception | None
) = None
reasoning_ended: bool | None = None
# Absolute index into the request's all_token_ids of the last reasoning
# token (the reasoning-end marker). Tokens at or before this index are
@@ -52,6 +54,8 @@ class StructuredOutputRequest:
self._grammar = self._grammar.result(timeout=0.0001)
except TimeoutError:
return False
except Exception as e:
self._grammar = e
return True
@property
@@ -59,11 +63,10 @@ class StructuredOutputRequest:
return self._check_grammar_completion()
@property
def grammar(self) -> StructuredOutputGrammar | None:
completed = self._check_grammar_completion()
return (
cast(StructuredOutputGrammar | None, self._grammar) if completed else None
)
def grammar(self) -> StructuredOutputGrammar | Exception | None:
if not self._check_grammar_completion():
return None
return cast(StructuredOutputGrammar | Exception | None, self._grammar)
@grammar.setter
def grammar(
+1 -3
View File
@@ -51,7 +51,7 @@ from vllm.sequence import IntermediateTensors
from vllm.tasks import SupportedTask
from vllm.utils.math_utils import cdiv
from vllm.utils.mem_utils import DeviceMemoryProfiler, format_gib
from vllm.utils.torch_utils import PIN_MEMORY, STR_DTYPE_TO_TORCH_DTYPE
from vllm.utils.torch_utils import STR_DTYPE_TO_TORCH_DTYPE
from vllm.v1.core.sched.output import GrammarOutput, SchedulerOutput
from vllm.v1.kv_cache_interface import KVCacheConfig, MambaSpec
from vllm.v1.outputs import DraftTokenIds, ModelRunnerOutput
@@ -521,12 +521,10 @@ class GPUModelRunner(LoRAModelRunnerMixin):
"""Build KV-block zeroing metadata; invoked from gpu_worker."""
self.kv_block_zeroer = KVBlockZeroer(
self.device,
pin_memory=PIN_MEMORY,
attn_groups_iter=(g for groups in self.attn_groups for g in groups),
kernel_block_sizes=self.kernel_block_sizes,
cache_dtype=self.cache_config.cache_dtype,
static_forward_context=self.compilation_config.static_forward_context,
max_concurrency=self.vllm_config.max_concurrent_batches,
)
@torch.inference_mode()
+1 -9
View File
@@ -132,15 +132,7 @@ class PromptLogprobsWorker:
if prompt_logprobs_list:
# Merge the in-progress logprobs.
logprobs = LogprobsTensors(
logprob_token_ids=torch.cat(
[x.logprob_token_ids for x in prompt_logprobs_list]
),
logprobs=torch.cat([x.logprobs for x in prompt_logprobs_list]),
selected_token_ranks=torch.cat(
[x.selected_token_ranks for x in prompt_logprobs_list]
),
)
logprobs = LogprobsTensors.cat(prompt_logprobs_list)
prompt_logprobs_list.clear()
if logprobs is None:
+3 -6
View File
@@ -5,7 +5,7 @@ import numpy as np
import torch
import vllm.envs as envs
from vllm.config.model import LogprobsMode
from vllm.config.model import PROCESSED_LOGPROBS_MODES, LogprobsMode
from vllm.sampling_params import SamplingParams
from vllm.v1.sample.ops.topk_topp_sampler import (
apply_top_k_top_p,
@@ -100,7 +100,7 @@ class Sampler:
)
if return_logprobs:
if self.logprobs_mode in ("processed_logprobs", "processed_logits"):
if self.logprobs_mode in PROCESSED_LOGPROBS_MODES:
logits = processed_logits
expanded_logits = logits.shape[0] != idx_mapping_np.shape[0]
cu_num_logits = cu_num_logits_np.tolist() if expanded_logits else None
@@ -221,10 +221,7 @@ class Sampler:
# any greedy requests or per-request seeds, or if post-processed
# logprobs need to be returned for any requests.
(top_k is None and top_p is None)
or (
return_logprobs
and self.logprobs_mode in ("processed_logprobs", "processed_logits")
)
or (return_logprobs and self.logprobs_mode in PROCESSED_LOGPROBS_MODES)
or self.sampling_states.any_greedy(idx_mapping_np)
or self.sampling_states.any_explicit_seed(idx_mapping_np)
)
@@ -1,8 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Iterator
import numpy as np
import torch
from vllm.config import SpeculativeConfig
from vllm.config.model import PROCESSED_LOGPROBS_MODES
from vllm.triton_utils import tl, triton
from vllm.v1.outputs import LogprobsTensors
from vllm.v1.spec_decode.utils import unconditional_to_conditional_rates
@@ -19,6 +23,29 @@ from vllm.v1.worker.gpu.spec_decode.rejection_sampler_utils import (
rejection_sample,
)
# Cap on the FP32 target-logits buffer materialized by apply_sampling_params.
# TODO(mgoin): Chunking is a workaround. The rejection kernels already upcast
# per vocab block on load and apply ops like temperature and gumbel, so folding
# sampling-param application into those kernels would remove this buffer and
# its traffic entirely.
MAX_CHUNK_BYTES = 2**30 # 1GB
_FP32_BYTES = 4
def _iter_request_chunks(
cu_num_logits: np.ndarray, max_chunk_logits: int
) -> Iterator[tuple[int, int]]:
"""Yield maximally packed request ranges without splitting requests."""
assert max_chunk_logits > 0
num_reqs = cu_num_logits.size - 1
start = 0
while start < num_reqs:
max_logit = int(cu_num_logits[start]) + max_chunk_logits
end = int(np.searchsorted(cu_num_logits, max_logit, side="right") - 1)
end = min(num_reqs, max(start + 1, end))
yield start, end
start = end
@triton.jit
def _flatten_sampled_kernel(
@@ -66,18 +93,17 @@ class RejectionSampler:
def _get_logprobs_tensors(
self,
input_batch: InputBatch,
sampled: torch.Tensor,
num_sampled: torch.Tensor,
logits: torch.Tensor,
cu_num_logits: torch.Tensor,
cu_num_logits_np: np.ndarray,
max_num_logprobs: int,
) -> LogprobsTensors | None:
max_num_logprobs = self.sampler.sampling_states.max_num_logprobs(
input_batch.idx_mapping_np
)
if max_num_logprobs == NO_LOGPROBS:
return None
num_reqs = input_batch.cu_num_logits.shape[0] - 1
num_reqs = cu_num_logits.shape[0] - 1
num_logits = logits.shape[0]
flat_sampled = torch.zeros(
num_logits, dtype=sampled.dtype, device=sampled.device
@@ -87,19 +113,122 @@ class RejectionSampler:
sampled,
sampled.stride(0),
num_sampled,
input_batch.cu_num_logits,
cu_num_logits,
num_warps=1,
)
expanded_logits = num_logits != input_batch.idx_mapping.shape[0]
expanded_logits = num_logits != num_reqs
return compute_topk_scores(
logits,
max_num_logprobs,
flat_sampled,
input_batch.cu_num_logits_np.tolist() if expanded_logits else None,
cu_num_logits_np.tolist() if expanded_logits else None,
logits_mode=self.sampler.logprobs_mode
in ("raw_logits", "processed_logits"),
)
def _verify(
self,
logits: torch.Tensor,
draft_logits: torch.Tensor | None,
draft_sampled: torch.Tensor,
pos: torch.Tensor,
cu_num_logits: torch.Tensor,
idx_mapping: torch.Tensor,
idx_mapping_np: np.ndarray,
expanded_idx_mapping: torch.Tensor,
expanded_local_pos: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
processed_logits = self.sampler.apply_sampling_params(
logits,
expanded_idx_mapping,
idx_mapping_np,
pos,
draft_sampled,
expanded_local_pos,
)
sampled, num_sampled = rejection_sample(
processed_logits,
draft_logits,
draft_sampled,
cu_num_logits,
pos,
idx_mapping,
expanded_idx_mapping,
expanded_local_pos,
self.sampler.sampling_states.temperature.gpu,
self.sampler.sampling_states.seeds.gpu,
self.num_speculative_steps,
self.synthetic_conditional_rates,
use_fp64=self.sampler.use_fp64_gumbel,
use_block_verification=self.use_block_verification,
)
return processed_logits, sampled, num_sampled
def _verify_in_chunks(
self,
logits: torch.Tensor,
input_batch: InputBatch,
draft_logits: torch.Tensor | None,
draft_sampled: torch.Tensor,
pos: torch.Tensor,
max_chunk_logits: int,
max_num_logprobs: int,
) -> tuple[torch.Tensor, torch.Tensor, LogprobsTensors | None]:
cu_num_logits_np = input_batch.cu_num_logits_np
use_processed_logits = self.sampler.logprobs_mode in PROCESSED_LOGPROBS_MODES
sampled_chunks: list[torch.Tensor] = []
num_sampled_chunks: list[torch.Tensor] = []
logprobs_chunks: list[LogprobsTensors] = []
for start, end in _iter_request_chunks(cu_num_logits_np, max_chunk_logits):
lo = int(cu_num_logits_np[start])
hi = int(cu_num_logits_np[end])
chunk_cu_num_logits_np = cu_num_logits_np[start : end + 1] - lo
chunk_cu_num_logits = input_batch.cu_num_logits[start : end + 1] - lo
# draft_logits uses persistent request-state indices and stays global.
processed_logits, sampled, num_sampled = self._verify(
logits[lo:hi],
draft_logits,
draft_sampled[lo:hi],
pos[lo:hi],
chunk_cu_num_logits,
input_batch.idx_mapping[start:end],
input_batch.idx_mapping_np[start:end],
input_batch.expanded_idx_mapping[lo:hi],
input_batch.expanded_local_pos[lo:hi],
)
chunk_logprobs = self._get_logprobs_tensors(
sampled,
num_sampled,
processed_logits if use_processed_logits else logits[lo:hi],
chunk_cu_num_logits,
chunk_cu_num_logits_np,
max_num_logprobs,
)
if chunk_logprobs is not None:
logprobs_chunks.append(chunk_logprobs)
del processed_logits
sampled_chunks.append(sampled)
num_sampled_chunks.append(num_sampled)
if len(sampled_chunks) == 1:
logprobs_tensors = logprobs_chunks[0] if logprobs_chunks else None
return sampled_chunks[0], num_sampled_chunks[0], logprobs_tensors
logprobs_tensors = None
if logprobs_chunks:
expanded_logits = logits.shape[0] != input_batch.num_reqs
logprobs_tensors = LogprobsTensors.cat(
logprobs_chunks,
cu_num_generated_tokens=(
cu_num_logits_np.tolist() if expanded_logits else None
),
)
sampled = torch.cat(sampled_chunks)
num_sampled = torch.cat(num_sampled_chunks)
return sampled, num_sampled, logprobs_tensors
def __call__(
self,
logits: torch.Tensor,
@@ -112,37 +241,19 @@ class RejectionSampler:
draft_sampled = input_batch.input_ids[input_batch.logits_indices]
pos = input_batch.positions[input_batch.logits_indices]
processed_logits = self.sampler.apply_sampling_params(
logits,
input_batch.expanded_idx_mapping,
input_batch.idx_mapping_np,
pos,
draft_sampled,
input_batch.expanded_local_pos,
max_num_logprobs = self.sampler.sampling_states.max_num_logprobs(
input_batch.idx_mapping_np
)
sampled, num_sampled = rejection_sample(
processed_logits,
max_chunk_logits = max(1, MAX_CHUNK_BYTES // (logits.shape[1] * _FP32_BYTES))
sampled, num_sampled, logprobs_tensors = self._verify_in_chunks(
logits,
input_batch,
draft_logits,
draft_sampled,
input_batch.cu_num_logits,
pos,
input_batch.idx_mapping,
input_batch.expanded_idx_mapping,
input_batch.expanded_local_pos,
self.sampler.sampling_states.temperature.gpu,
self.sampler.sampling_states.seeds.gpu,
self.num_speculative_steps,
self.synthetic_conditional_rates,
use_fp64=self.sampler.use_fp64_gumbel,
use_block_verification=self.use_block_verification,
)
logprobs_tensors = self._get_logprobs_tensors(
input_batch,
sampled,
num_sampled,
processed_logits
if self.sampler.logprobs_mode in ("processed_logprobs", "processed_logits")
else logits,
max_chunk_logits,
max_num_logprobs,
)
num_sampled, num_rejected = get_num_sampled_and_rejected(
@@ -888,6 +888,10 @@ def rejection_sample(
use_fp64: bool = False,
use_block_verification: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
assert target_logits.ndim == 2 and target_logits.stride(-1) == 1
assert draft_logits is None or (
draft_logits.ndim == 3 and draft_logits.stride(-1) == 1
)
num_reqs = cu_num_logits.shape[0] - 1
num_logits, vocab_size = target_logits.shape
draft_logits_stride_0 = 0
+2 -6
View File
@@ -37,6 +37,7 @@ from vllm.config import (
update_config,
)
from vllm.config.cache import CacheConfig
from vllm.config.model import PROCESSED_LOGPROBS_MODES
from vllm.distributed.ec_transfer import get_ec_transfer, has_ec_transfer
from vllm.distributed.eplb.eplb_state import EplbState
from vllm.distributed.kv_transfer import get_kv_transfer_group, has_kv_transfer_group
@@ -1134,13 +1135,11 @@ class GPUModelRunner(
"""
self._kv_block_zeroer = KVBlockZeroer(
self.device,
pin_memory=PIN_MEMORY,
attn_groups_iter=self._kv_cache_spec_attn_group_iterator(),
kernel_block_sizes=self._kernel_block_sizes,
cache_dtype=self.cache_config.cache_dtype,
runner_only_attn_layers=self.runner_only_attn_layers,
static_forward_context=self.compilation_config.static_forward_context,
max_concurrency=self.vllm_config.max_concurrent_batches,
)
def _zero_block_ids(self, block_ids: list[int]) -> None:
@@ -6217,10 +6216,7 @@ class GPUModelRunner(
# memory during profile_run.
# No .clone() of logits: warmup output is discarded, so any in-place
# mutation by forward_native does not affect correctness.
if self.sampler.logprobs_mode not in (
"processed_logits",
"processed_logprobs",
):
if self.sampler.logprobs_mode not in PROCESSED_LOGPROBS_MODES:
self.sampler(
logits=logits,
sampling_metadata=replace(
+1 -43
View File
@@ -91,13 +91,11 @@ class KVBlockZeroer:
def __init__(
self,
device: torch.device,
pin_memory: bool,
attn_groups_iter: Iterable["AttentionGroup"],
kernel_block_sizes: list[int],
cache_dtype: str,
static_forward_context: dict[str, Any],
runner_only_attn_layers: set[str] | None = None,
max_concurrency: int = 1,
) -> None:
"""Precompute the absolute-address table for the Triton zeroing kernel.
@@ -112,15 +110,7 @@ class KVBlockZeroer:
Only AttentionSpec layers are processed; Mamba layers are skipped.
"""
self.device = device
self.pin_memory = pin_memory
if max_concurrency < 1:
raise ValueError("max_concurrency must be at least 1")
self.max_concurrency = max_concurrency
self._meta: tuple[torch.Tensor, int, int, int] | None = None
self._id_cap: int = 0
self._ids_pinned: list[torch.Tensor] = []
self._ids_gpu: list[torch.Tensor] = []
self._id_buffer_index = 0
if runner_only_attn_layers is None:
runner_only_attn_layers = set()
@@ -182,8 +172,6 @@ class KVBlockZeroer:
return
blk_size = min(largest_power_of_2_divisor(page_size_el), 1024)
self._id_cap = 8192
self._allocate_id_buffers()
self._meta = (
torch.tensor(seg_addrs, dtype=torch.uint64, device=self.device),
page_size_el,
@@ -191,43 +179,13 @@ class KVBlockZeroer:
len(seg_addrs),
)
def _allocate_id_buffers(self) -> None:
self._ids_pinned = [
torch.empty(
self._id_cap,
dtype=torch.int64,
pin_memory=self.pin_memory,
)
for _ in range(self.max_concurrency)
]
self._ids_gpu = [
torch.empty(self._id_cap, dtype=torch.int64, device=self.device)
for _ in range(self.max_concurrency)
]
self._id_buffer_index = 0
def zero_block_ids(self, block_ids: list[int]) -> None:
"""Zero the KV cache memory for the given block IDs."""
if not block_ids or self._meta is None:
return
seg_addrs, page_size_el, blk_size, n_segs = self._meta
n_blocks = len(block_ids)
if n_blocks > self._id_cap:
# The old pinned buffers may still be the source of an in-flight
# nonblocking copy. Growing is rare, so we don't mind the sync overhead
torch.accelerator.synchronize()
self._id_cap = n_blocks * 2
self._allocate_id_buffers()
# The H2D copy is nonblocking, so its pinned source must not be mutated
# while this batch is in flight. Rotate through as many buffers as concurrent
# in-flight batches, to avoid collisions.
buffer_index = self._id_buffer_index
self._id_buffer_index = (buffer_index + 1) % self.max_concurrency
ids_pinned = self._ids_pinned[buffer_index]
ids_pinned[:n_blocks].numpy()[:] = block_ids
idx = self._ids_gpu[buffer_index][:n_blocks]
idx.copy_(ids_pinned[:n_blocks], non_blocking=True)
idx = async_tensor_h2d(block_ids, device=self.device, dtype=torch.int64)
grid = (n_blocks * n_segs * (page_size_el // blk_size),)
_zero_kv_blocks_kernel[grid](
seg_addrs,