forked from Karylab-cklius/vllm
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
78d2334aab | ||
|
|
7b375c8502 |
@@ -67,7 +67,7 @@ steps:
|
||||
pytest -v -s v1/worker --ignore=v1/worker/test_gpu_model_runner.py --ignore=v1/worker/test_worker_memory_snapshot.py &&
|
||||
pytest -v -s v1/structured_output &&
|
||||
pytest -v -s v1/test_serial_utils.py &&
|
||||
pytest -v -s v1/spec_decode --ignore=v1/spec_decode/test_max_len.py --ignore=v1/spec_decode/test_speculators_eagle3.py --ignore=v1/spec_decode/test_acceptance_length.py --ignore=v1/spec_decode/test_speculators_correctness.py &&
|
||||
pytest -v -s v1/spec_decode --ignore=v1/spec_decode/test_max_len.py --ignore=v1/spec_decode/test_speculators_eagle3.py --ignore=v1/spec_decode/test_acceptance_length.py &&
|
||||
pytest -v -s v1/kv_connector/unit --ignore=v1/kv_connector/unit/test_multi_connector.py --ignore=v1/kv_connector/unit/test_example_connector.py --ignore=v1/kv_connector/unit/test_lmcache_integration.py --ignore=v1/kv_connector/unit/test_hf3fs_client.py --ignore=v1/kv_connector/unit/test_hf3fs_connector.py --ignore=v1/kv_connector/unit/test_hf3fs_metadata_server.py --ignore=v1/kv_connector/unit/test_offloading_connector.py'
|
||||
- label: "XPU server test"
|
||||
depends_on:
|
||||
|
||||
@@ -81,11 +81,11 @@ __global__ void rms_norm_kernel(
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC_SIZE; j++) {
|
||||
float x = static_cast<float>(src1.val[j]);
|
||||
scalar_t normalized = static_cast<scalar_t>(x * s_variance);
|
||||
if constexpr (HasWeight) {
|
||||
float w = static_cast<float>(src2.val[j]);
|
||||
dst.val[j] = static_cast<scalar_t>(x * s_variance * w);
|
||||
dst.val[j] = normalized * src2.val[j];
|
||||
} else {
|
||||
dst.val[j] = static_cast<scalar_t>(x * s_variance);
|
||||
dst.val[j] = normalized;
|
||||
}
|
||||
}
|
||||
v_out[i] = dst;
|
||||
@@ -151,8 +151,7 @@ fused_add_rms_norm_kernel(
|
||||
#pragma unroll
|
||||
for (int j = 0; j < width; ++j) {
|
||||
float x = Converter::convert(res.data[j]);
|
||||
float wf = Converter::convert(w.data[j]);
|
||||
out.data[j] = Converter::convert(x * s_variance * wf);
|
||||
out.data[j] = Converter::convert(x * s_variance) * w.data[j];
|
||||
}
|
||||
} else {
|
||||
#pragma unroll
|
||||
@@ -199,8 +198,8 @@ fused_add_rms_norm_kernel(
|
||||
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
|
||||
float x = (float)residual[blockIdx.x * hidden_size + idx];
|
||||
if constexpr (HasWeight) {
|
||||
float w = (float)weight[idx];
|
||||
input[blockIdx.x * input_stride + idx] = (scalar_t)(x * s_variance * w);
|
||||
input[blockIdx.x * input_stride + idx] =
|
||||
(scalar_t)(x * s_variance) * weight[idx];
|
||||
} else {
|
||||
input[blockIdx.x * input_stride + idx] = (scalar_t)(x * s_variance);
|
||||
}
|
||||
|
||||
@@ -66,13 +66,8 @@ __global__ void rms_norm_static_fp8_quant_kernel(
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC_SIZE; j++) {
|
||||
float x = static_cast<float>(src1.val[j]);
|
||||
float w = static_cast<float>(src2.val[j]);
|
||||
// Round normalized result through scalar_t to match the precision of the
|
||||
// unfused composite (rms_norm writes scalar_t, then
|
||||
// static_scaled_fp8_quant re-loads it as float before FP8 conversion).
|
||||
// Without this round, the fused path is strictly more accurate and
|
||||
// disagrees with the composite at exact E4M3 quantization tie boundaries.
|
||||
scalar_t out_norm = static_cast<scalar_t>(x * s_variance * w);
|
||||
// Multiply in weight's native dtype to match rms_norm_kernel.
|
||||
scalar_t out_norm = static_cast<scalar_t>(x * s_variance) * src2.val[j];
|
||||
out[blockIdx.x * hidden_size + idx * VEC_SIZE + j] =
|
||||
scaled_fp8_conversion<true, fp8_type>(static_cast<float>(out_norm),
|
||||
scale_inv);
|
||||
@@ -142,12 +137,8 @@ fused_add_rms_norm_static_fp8_quant_kernel(
|
||||
#pragma unroll
|
||||
for (int i = 0; i < width; ++i) {
|
||||
float x = Converter::convert(res.data[i]);
|
||||
float wf = Converter::convert(w.data[i]);
|
||||
// See note in rms_norm_static_fp8_quant_kernel: round through scalar_t
|
||||
// to match the unfused composite path at FP8 boundaries. We use the
|
||||
// backend's hip_type for the intermediate since c10::Half/BFloat16 has
|
||||
// ambiguous conversions on CUDA and no implicit conversion on ROCm.
|
||||
HipT out_norm_h = Converter::convert(x * s_variance * wf);
|
||||
// Multiply in weight's native dtype to match fused_add_rms_norm_kernel.
|
||||
HipT out_norm_h = Converter::convert(x * s_variance) * w.data[i];
|
||||
out[id * width + i] = scaled_fp8_conversion<true, fp8_type>(
|
||||
Converter::convert(out_norm_h), scale_inv);
|
||||
}
|
||||
@@ -192,10 +183,8 @@ fused_add_rms_norm_static_fp8_quant_kernel(
|
||||
|
||||
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
|
||||
float x = (float)residual[blockIdx.x * hidden_size + idx];
|
||||
float w = (float)weight[idx];
|
||||
// See note in rms_norm_static_fp8_quant_kernel: round through scalar_t
|
||||
// to match the unfused composite path at FP8 boundaries.
|
||||
scalar_t out_norm = static_cast<scalar_t>(x * s_variance * w);
|
||||
// Multiply in weight's native dtype to match fused_add_rms_norm_kernel.
|
||||
scalar_t out_norm = static_cast<scalar_t>(x * s_variance) * weight[idx];
|
||||
out[blockIdx.x * hidden_size + idx] = scaled_fp8_conversion<true, fp8_type>(
|
||||
static_cast<float>(out_norm), scale_inv);
|
||||
}
|
||||
|
||||
@@ -75,13 +75,13 @@ RUN wget -O- https://apt.repos.intel.com/intel-gpg-keys/GPG-PUB-KEY-INTEL-SW-PRO
|
||||
# Install UMD
|
||||
RUN mkdir neo && \
|
||||
cd neo && \
|
||||
wget https://github.com/intel/intel-graphics-compiler/releases/download/v2.34.4/intel-igc-core-2_2.34.4+21428_amd64.deb && \
|
||||
wget https://github.com/intel/intel-graphics-compiler/releases/download/v2.34.4/intel-igc-opencl-2_2.34.4+21428_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/26.18.38308.1/intel-ocloc_26.18.38308.1-0_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/26.18.38308.1/intel-opencl-icd_26.18.38308.1-0_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/26.18.38308.1/libigdgmm12_22.10.0_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/26.18.38308.1/libze-intel-gpu1_26.18.38308.1-0_amd64.deb && \
|
||||
wget https://github.com/oneapi-src/level-zero/releases/download/v1.28.2/level-zero_1.28.2+u24.04_amd64.deb && \
|
||||
wget https://github.com/intel/intel-graphics-compiler/releases/download/v2.24.8/intel-igc-core-2_2.24.8+20344_amd64.deb && \
|
||||
wget https://github.com/intel/intel-graphics-compiler/releases/download/v2.24.8/intel-igc-opencl-2_2.24.8+20344_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/25.48.36300.8/intel-ocloc_25.48.36300.8-0_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/25.48.36300.8/intel-opencl-icd_25.48.36300.8-0_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/25.48.36300.8/libigdgmm12_22.8.2_amd64.deb && \
|
||||
wget https://github.com/intel/compute-runtime/releases/download/25.48.36300.8/libze-intel-gpu1_25.48.36300.8-0_amd64.deb && \
|
||||
wget https://github.com/oneapi-src/level-zero/releases/download/v1.26.0/level-zero_1.26.0+u24.04_amd64.deb && \
|
||||
dpkg -i *.deb && \
|
||||
cd .. && \
|
||||
rm -rf neo
|
||||
|
||||
@@ -4,7 +4,7 @@ Deploying vLLM on Kubernetes is a scalable and efficient way to serve machine le
|
||||
|
||||
* **Upstream vLLM compatibility** – It wraps around upstream vLLM without modifying its code.
|
||||
* **Ease of use** – Simplified deployment via Helm charts and observability through Grafana dashboards.
|
||||
* **High performance** – Optimized for LLM workloads with features like multimodel support, model-aware and prefix-aware routing, fast vLLM bootstrapping, and KV cache offloading with [LMCache](https://github.com/LMCache/LMCache) (wired up in vLLM via `--kv-offloading-backend lmcache`; see the [LMCache examples](https://github.com/vllm-project/vllm/tree/main/examples/disaggregated/lmcache) and [docs.lmcache.ai](https://docs.lmcache.ai)), among others.
|
||||
* **High performance** – Optimized for LLM workloads with features like multimodel support, model-aware and prefix-aware routing, fast vLLM bootstrapping, and KV cache offloading with [LMCache](https://github.com/LMCache/LMCache), among others.
|
||||
|
||||
If you are new to Kubernetes, don't worry: in the vLLM production stack [repo](https://github.com/vllm-project/production-stack), we provide a step-by-step [guide](https://github.com/vllm-project/production-stack/blob/main/tutorials/00-install-kubernetes-env.md) and a [short video](https://www.youtube.com/watch?v=EsTJbQtzj0g) to set up everything and get started in **4 minutes**!
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ Two main reasons:
|
||||
Now supports 9 types of connectors:
|
||||
|
||||
- **ExampleConnector**: refer to [examples/disaggregated/example_connector/run.sh](../../examples/disaggregated/example_connector/run.sh) for the example usage of ExampleConnector disaggregated prefilling.
|
||||
- **LMCacheConnectorV1**: refer to [examples/disaggregated/lmcache/disagg_prefill_lmcache_v1/disagg_example_nixl.sh](../../examples/disaggregated/lmcache/disagg_prefill_lmcache_v1/disagg_example_nixl.sh) for the example usage of LMCacheConnectorV1 disaggregated prefilling which uses NIXL as the underlying KV transmission. LMCache also offers a multi-process (MP) mode via `LMCacheMPConnector`, where a standalone `lmcache server` holds the KV cache shared by one or more vLLM instances; see the [LMCache examples](../../examples/disaggregated/lmcache/README.md) and the [LMCache docs](https://docs.lmcache.ai) for setup.
|
||||
- **LMCacheConnectorV1**: refer to [examples/disaggregated/lmcache/disagg_prefill_lmcache_v1/disagg_example_nixl.sh](../../examples/disaggregated/lmcache/disagg_prefill_lmcache_v1/disagg_example_nixl.sh) for the example usage of LMCacheConnectorV1 disaggregated prefilling which uses NIXL as the underlying KV transmission.
|
||||
- **NixlConnector**: refer to [tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh](../../tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh) for the example usage of NixlConnector disaggregated prefilling which support fully async send/recv. For detailed usage guide, see [NixlConnector Usage Guide](nixl_connector_usage.md). For feature compatibility details, see [NixlConnector Compatibility Matrix](nixl_connector_compatibility.md). You may specify one or multiple NIXL transfer backends, such as:
|
||||
|
||||
```bash
|
||||
|
||||
@@ -203,7 +203,6 @@ the vLLM JSON config.
|
||||
### kv_connector_extra_config
|
||||
|
||||
- `load_async` (bool): Enable asynchronous loading for better compute-I/O overlap. Default: `true`.
|
||||
- `lookup_async` (bool): Run the external prefix-cache lookup on a background thread so it never blocks the scheduler step. The request is held until the in-flight lookup completes, then resumed on a later step. Default: `false`.
|
||||
- `enable_cross_layers_blocks` (bool): Enable cross-layer block packing for reduced store operations. Default: `false`.
|
||||
- `lookup_rpc_port` (int): Custom port for the ZMQ lookup RPC socket. Default: `0`.
|
||||
- `cache_prefix` (str): Namespace prepended to every store key. Lets separate deployments share one Mooncake master without polluting each other — instances configured with different prefixes never see each other's cached blocks, even for identical prompts. All instances that should share a prefix cache must use the same value. Default: `""` (no prefix; keys are byte-identical to the unprefixed format).
|
||||
|
||||
@@ -27,7 +27,6 @@ Currently, there are no pre-built XPU wheels.
|
||||
|
||||
- First, install required [driver](https://dgpu-docs.intel.com/driver/installation.html#installing-gpu-drivers).
|
||||
- Second, install Python packages for vLLM XPU backend building (Intel OneAPI dependencies are installed automatically as part of `torch-xpu`, see [PyTorch XPU get started](https://docs.pytorch.org/docs/stable/notes/get_start_xpu.html)):
|
||||
- Start from vllm-xpu-kernels v0.1.10, we recommend user upgrade driver to [compute runtime 26.18](https://github.com/intel/compute-runtime/releases/tag/26.14.37833.4) release, to avoid potential compatibility issue.
|
||||
|
||||
```bash
|
||||
git clone https://github.com/vllm-project/vllm.git
|
||||
|
||||
@@ -1,38 +1,10 @@
|
||||
# LMCache Examples
|
||||
|
||||
This folder demonstrates how to use LMCache with vLLM v1 for KV cache
|
||||
offloading, disaggregated prefilling, and KV cache sharing.
|
||||
This folder demonstrates how to use LMCache for disaggregated prefilling, CPU offloading and KV cache sharing.
|
||||
|
||||
## Integration modes
|
||||
## 1. Disaggregated Prefill in vLLM v1
|
||||
|
||||
LMCache integrates with vLLM v1 in two ways:
|
||||
|
||||
- **In-process mode** (`LMCacheConnectorV1`): LMCache runs inside the vLLM
|
||||
process and is configured through environment variables or a YAML config
|
||||
file (`LMCACHE_CONFIG_FILE`). This is the simplest way to add single-node
|
||||
CPU/disk offloading.
|
||||
- **Multi-process (MP) mode** (`LMCacheMPConnector`): LMCache runs as a
|
||||
standalone server (`lmcache server`) that owns the KV cache storage; one or
|
||||
more vLLM instances connect to it. This is the recommended mode for
|
||||
distributed KV storage and for sharing KV cache across instances. See the
|
||||
[LMCache docs](https://docs.lmcache.ai) for the full MP setup.
|
||||
|
||||
## 1. CPU offload (in-process)
|
||||
|
||||
- `python cpu_offload_lmcache.py` - CPU offloading with `LMCacheConnectorV1`
|
||||
for vLLM v1.
|
||||
|
||||
## 2. CPU offload (multi-process)
|
||||
|
||||
- `bash cpu_offload_lmcache_mp.sh` - CPU offloading with `LMCacheMPConnector`,
|
||||
using a standalone `lmcache server`. vLLM provides a built-in shortcut for
|
||||
this setup via `--kv-offloading-backend lmcache` and
|
||||
`--kv-offloading-size <GiB>`.
|
||||
|
||||
## 3. Disaggregated Prefill in vLLM v1
|
||||
|
||||
This example demonstrates how to run LMCache with disaggregated prefill using
|
||||
NIXL on a single node.
|
||||
This example demonstrates how to run LMCache with disaggregated prefill using NIXL on a single node.
|
||||
|
||||
### Prerequisites
|
||||
|
||||
@@ -74,7 +46,15 @@ The main script generates several log files:
|
||||
- `decoder.log` - Logs from the decode server
|
||||
- `proxy.log` - Logs from the proxy server
|
||||
|
||||
## 4. KV Cache Sharing
|
||||
## 2. CPU Offload Examples
|
||||
|
||||
The `kv_cache_sharing_lmcache_v1.py` example demonstrates how to share KV
|
||||
caches between vLLM v1 instances through a centralized LMCache server.
|
||||
- `python cpu_offload_lmcache.py -v v0` - CPU offloading implementation for vLLM v0
|
||||
- `python cpu_offload_lmcache.py -v v1` - CPU offloading implementation for vLLM v1
|
||||
|
||||
## 3. KV Cache Sharing
|
||||
|
||||
The `kv_cache_sharing_lmcache_v1.py` example demonstrates how to share KV caches between vLLM v1 instances.
|
||||
|
||||
## 4. Disaggregated Prefill in vLLM v0
|
||||
|
||||
The `disaggregated_prefill_lmcache_v0.py` provides an example of how to run disaggregated prefill in vLLM v0.
|
||||
|
||||
@@ -1,8 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
This file demonstrates the example usage of CPU offloading
|
||||
with LMCache in vLLM v1.
|
||||
This file demonstrates the example usage of cpu offloading
|
||||
with LMCache in vLLM v1 or v0.
|
||||
|
||||
Usage:
|
||||
|
||||
Specify vLLM version
|
||||
|
||||
-v v0 : Use LMCacheConnector
|
||||
model = mistralai/Mistral-7B-Instruct-v0.2
|
||||
(Includes enable_chunked_prefill = True)
|
||||
|
||||
-v v1 : Use LMCacheConnectorV1 (default)
|
||||
model = meta-llama/Meta-Llama-3.1-8B-Instruct
|
||||
(Without enable_chunked_prefill)
|
||||
|
||||
Note that `lmcache` is needed to run this example.
|
||||
Requirements:
|
||||
@@ -11,6 +23,7 @@ Learn more about LMCache environment setup, please refer to:
|
||||
https://docs.lmcache.ai/getting_started/installation.html
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import os
|
||||
import time
|
||||
@@ -26,6 +39,8 @@ from vllm.engine.arg_utils import EngineArgs
|
||||
|
||||
def setup_environment_variables():
|
||||
# LMCache-related environment variables
|
||||
# Use experimental features in LMCache
|
||||
os.environ["LMCACHE_USE_EXPERIMENTAL"] = "True"
|
||||
# LMCache is set to use 256 tokens per chunk
|
||||
os.environ["LMCACHE_CHUNK_SIZE"] = "256"
|
||||
# Enable local CPU backend in LMCache
|
||||
@@ -35,9 +50,9 @@ def setup_environment_variables():
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def build_llm_with_lmcache(model: str):
|
||||
def build_llm_with_lmcache(lmcache_connector: str, model: str):
|
||||
ktc = KVTransferConfig(
|
||||
kv_connector="LMCacheConnectorV1",
|
||||
kv_connector=lmcache_connector,
|
||||
kv_role="kv_both",
|
||||
)
|
||||
# Set GPU memory utilization to 0.8 for an A40 GPU with 40GB
|
||||
@@ -77,10 +92,23 @@ def print_output(
|
||||
print("-" * 50)
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"-v",
|
||||
"--version",
|
||||
choices=["v0", "v1"],
|
||||
default="v1",
|
||||
help="Specify vLLM version (default: v1)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
lmcache_connector = "LMCacheConnectorV1"
|
||||
model = "meta-llama/Meta-Llama-3.1-8B-Instruct"
|
||||
setup_environment_variables()
|
||||
with build_llm_with_lmcache(model) as llm:
|
||||
with build_llm_with_lmcache(lmcache_connector, model) as llm:
|
||||
# This example script runs two requests with a shared prefix.
|
||||
# Define the shared prompt and specific prompts
|
||||
shared_prompt = "Hello, how are you?" * 1000
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
#!/bin/bash
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# CPU offloading with LMCache in multi-process (MP) mode.
|
||||
#
|
||||
# In MP mode, LMCache runs as a standalone server process (`lmcache server`)
|
||||
# that owns the KV cache storage. One or more vLLM instances connect to it via
|
||||
# the `LMCacheMPConnector`. This is the recommended way to run LMCache for
|
||||
# distributed KV storage and for sharing KV cache across vLLM instances.
|
||||
#
|
||||
# vLLM ships a built-in shortcut for this setup: pass `--kv-offloading-backend
|
||||
# lmcache` together with `--kv-offloading-size <GiB>` and vLLM wires up the
|
||||
# `LMCacheMPConnector` for you (it defaults to the LMCache server at
|
||||
# tcp://localhost:5555, matching the `lmcache server` default).
|
||||
#
|
||||
# Requires `lmcache` to be installed (`pip install lmcache`).
|
||||
# Learn more: https://docs.lmcache.ai
|
||||
set -euo pipefail
|
||||
|
||||
MODEL=${MODEL:-meta-llama/Meta-Llama-3.1-8B-Instruct}
|
||||
|
||||
# 1. Launch the standalone LMCache server (binds tcp://localhost:5555 by
|
||||
# default). `--l1-size-gb` sets the CPU memory budget for the L1 cache.
|
||||
echo "Starting LMCache server..."
|
||||
lmcache server --host localhost --port 5555 --l1-size-gb 5 &
|
||||
LMCACHE_SERVER_PID=$!
|
||||
trap 'kill $LMCACHE_SERVER_PID 2>/dev/null || true' EXIT
|
||||
|
||||
# 2. Launch vLLM and offload KV cache to the LMCache server.
|
||||
# The MP connector currently requires the non-hybrid KV cache manager.
|
||||
echo "Starting vLLM server with LMCache MP offloading..."
|
||||
vllm serve "$MODEL" \
|
||||
--port 8000 \
|
||||
--kv-offloading-size 5 \
|
||||
--kv-offloading-backend lmcache \
|
||||
--disable-hybrid-kv-cache-manager
|
||||
|
||||
# Equivalent explicit configuration (instead of the two flags above):
|
||||
# --kv-transfer-config \
|
||||
# '{"kv_connector":"LMCacheMPConnector","kv_role":"kv_both",
|
||||
# "kv_connector_extra_config":{"lmcache.mp.host":"tcp://localhost",
|
||||
# "lmcache.mp.port":5555}}'
|
||||
@@ -0,0 +1,144 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
This file demonstrates the example usage of disaggregated prefilling
|
||||
with LMCache.
|
||||
We will launch 2 vllm instances (GPU 0 for prefill and GPU 1 for decode),
|
||||
and launch an additional LMCache server.
|
||||
KV cache is transferred in the following manner:
|
||||
vLLM prefill node -> LMCache server -> vLLM decode node.
|
||||
|
||||
Note that `pip install lmcache` is needed to run this example.
|
||||
Learn more about LMCache in https://github.com/LMCache/LMCache.
|
||||
"""
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import time
|
||||
from multiprocessing import Event, Process
|
||||
|
||||
from lmcache.experimental.cache_engine import LMCacheEngineBuilder
|
||||
from lmcache.integration.vllm.utils import ENGINE_NAME
|
||||
|
||||
from vllm import LLM, SamplingParams
|
||||
from vllm.config import KVTransferConfig
|
||||
|
||||
# LMCache-related environment variables
|
||||
# The port to start LMCache server
|
||||
port = 8100
|
||||
# Use experimental features in LMCache
|
||||
os.environ["LMCACHE_USE_EXPERIMENTAL"] = "True"
|
||||
# LMCache is set to use 256 tokens per chunk
|
||||
os.environ["LMCACHE_CHUNK_SIZE"] = "256"
|
||||
# Disable local CPU backend in LMCache
|
||||
os.environ["LMCACHE_LOCAL_CPU"] = "False"
|
||||
# Set local CPU memory buffer limit to 5.0 GB
|
||||
os.environ["LMCACHE_MAX_LOCAL_CPU_SIZE"] = "5.0"
|
||||
# Set the remote URL for LMCache server
|
||||
os.environ["LMCACHE_REMOTE_URL"] = f"lm://localhost:{port}"
|
||||
# Set the serializer/deserializer between vllm and LMCache server
|
||||
# `naive` indicates using raw bytes of the tensor without any compression
|
||||
os.environ["LMCACHE_REMOTE_SERDE"] = "naive"
|
||||
|
||||
prompts = [
|
||||
"Hello, how are you?" * 1000,
|
||||
]
|
||||
|
||||
|
||||
def run_prefill(prefill_done, prompts):
|
||||
# We use GPU 0 for prefill node.
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
|
||||
|
||||
sampling_params = SamplingParams(temperature=0, top_p=0.95, max_tokens=1)
|
||||
|
||||
ktc = KVTransferConfig(
|
||||
kv_connector="LMCacheConnector",
|
||||
kv_role="kv_producer",
|
||||
kv_rank=0,
|
||||
kv_parallel_size=2,
|
||||
)
|
||||
# Set GPU memory utilization to 0.8 for an A40 GPU with 40GB
|
||||
# memory. Reduce the value if your GPU has less memory.
|
||||
llm = LLM(
|
||||
model="mistralai/Mistral-7B-Instruct-v0.2",
|
||||
kv_transfer_config=ktc,
|
||||
max_model_len=8000,
|
||||
gpu_memory_utilization=0.8,
|
||||
enforce_eager=True,
|
||||
)
|
||||
|
||||
# llm.generate(prompts, sampling_params)
|
||||
outputs = llm.generate(prompts, sampling_params)
|
||||
for output in outputs:
|
||||
generated_text = output.outputs[0].text
|
||||
print(f"Generated text: {generated_text!r}")
|
||||
print("Prefill node is finished.")
|
||||
prefill_done.set()
|
||||
|
||||
# Clean up lmcache backend
|
||||
LMCacheEngineBuilder.destroy(ENGINE_NAME)
|
||||
|
||||
|
||||
def run_decode(prefill_done, prompts, timeout=1):
|
||||
# We use GPU 1 for decode node.
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = "1"
|
||||
|
||||
sampling_params = SamplingParams(temperature=0, top_p=0.95, max_tokens=10)
|
||||
|
||||
ktc = KVTransferConfig(
|
||||
kv_connector="LMCacheConnector",
|
||||
kv_role="kv_consumer",
|
||||
kv_rank=1,
|
||||
kv_parallel_size=2,
|
||||
)
|
||||
# Set GPU memory utilization to 0.8 for an A40 GPU with 40GB
|
||||
# of memory. Reduce the value if your GPU has less memory.
|
||||
llm = LLM(
|
||||
model="mistralai/Mistral-7B-Instruct-v0.2",
|
||||
kv_transfer_config=ktc,
|
||||
max_model_len=8000,
|
||||
gpu_memory_utilization=0.8,
|
||||
enforce_eager=True,
|
||||
)
|
||||
|
||||
print("Waiting for prefill node to finish...")
|
||||
prefill_done.wait()
|
||||
time.sleep(timeout)
|
||||
|
||||
outputs = llm.generate(prompts, sampling_params)
|
||||
for output in outputs:
|
||||
generated_text = output.outputs[0].text
|
||||
print(f"Generated text: {generated_text!r}")
|
||||
|
||||
# Clean up lmcache backend
|
||||
LMCacheEngineBuilder.destroy(ENGINE_NAME)
|
||||
|
||||
|
||||
def run_lmcache_server(port):
|
||||
server_proc = subprocess.Popen(
|
||||
["python", "-m", "lmcache.experimental.server", "localhost", str(port)]
|
||||
)
|
||||
return server_proc
|
||||
|
||||
|
||||
def main():
|
||||
prefill_done = Event()
|
||||
prefill_process = Process(target=run_prefill, args=(prefill_done, prompts))
|
||||
decode_process = Process(target=run_decode, args=(prefill_done, prompts))
|
||||
lmcache_server_process = run_lmcache_server(port)
|
||||
|
||||
# Start prefill node
|
||||
prefill_process.start()
|
||||
|
||||
# Start decode node
|
||||
decode_process.start()
|
||||
|
||||
# Clean up the processes
|
||||
decode_process.join()
|
||||
prefill_process.terminate()
|
||||
lmcache_server_process.terminate()
|
||||
lmcache_server_process.wait()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -30,6 +30,7 @@ if [[ $1 == "prefiller" ]]; then
|
||||
|
||||
UCX_TLS=cuda_ipc,cuda_copy,tcp \
|
||||
LMCACHE_CONFIG_FILE=$prefill_config_file \
|
||||
LMCACHE_USE_EXPERIMENTAL=True \
|
||||
VLLM_ENABLE_V1_MULTIPROCESSING=1 \
|
||||
VLLM_WORKER_MULTIPROC_METHOD=spawn \
|
||||
CUDA_VISIBLE_DEVICES=0 \
|
||||
@@ -46,6 +47,7 @@ elif [[ $1 == "decoder" ]]; then
|
||||
|
||||
UCX_TLS=cuda_ipc,cuda_copy,tcp \
|
||||
LMCACHE_CONFIG_FILE=$decode_config_file \
|
||||
LMCACHE_USE_EXPERIMENTAL=True \
|
||||
VLLM_ENABLE_V1_MULTIPROCESSING=1 \
|
||||
VLLM_WORKER_MULTIPROC_METHOD=spawn \
|
||||
CUDA_VISIBLE_DEVICES=1 \
|
||||
|
||||
@@ -26,6 +26,8 @@ from vllm.config import KVTransferConfig
|
||||
# LMCache-related environment variables
|
||||
# The port to start LMCache server
|
||||
port = 8100
|
||||
# Use experimental features in LMCache
|
||||
os.environ["LMCACHE_USE_EXPERIMENTAL"] = "True"
|
||||
# LMCache is set to use 256 tokens per chunk
|
||||
os.environ["LMCACHE_CHUNK_SIZE"] = "256"
|
||||
# Disable local CPU backend in LMCache
|
||||
|
||||
@@ -11,14 +11,13 @@ transformers >= 5.5.3
|
||||
tokenizers >= 0.21.1 # Required for fast incremental detokenization.
|
||||
safetensors >= 0.6.2 # MXFP4/MXFP6 dtype support (F8_E8M0, F4) added in 0.6.0: https://github.com/huggingface/safetensors/pull/611
|
||||
protobuf >= 5.29.6, !=6.30.*, !=6.31.*, !=6.32.*, !=6.33.0.*, !=6.33.1.*, !=6.33.2.*, !=6.33.3.*, !=6.33.4.* # Required by LlamaTokenizer, gRPC. CVE-2026-0994
|
||||
fastapi[standard] >= 0.133.0, < 0.137.0 # First version supporting Starlette 1.0; < 0.137.0 avoids route-tree change that breaks model-hosting-container-standards handler overrides.
|
||||
starlette >= 1.0.1 # CVE-2026-48710: Host header injection in < 1.0.1
|
||||
fastapi[standard] >= 0.115.0 # Required by FastAPI's form models in the OpenAI API server's audio transcriptions endpoint.
|
||||
aiohttp >= 3.13.3
|
||||
openai >= 2.0.0 # For Responses API with reasoning content
|
||||
pydantic >= 2.12.0
|
||||
prometheus_client >= 0.18.0
|
||||
pillow # Required for image processing
|
||||
prometheus-fastapi-instrumentator >= 8.0.0 # v8 unblocks starlette >= 1.0
|
||||
prometheus-fastapi-instrumentator >= 7.0.0
|
||||
tiktoken >= 0.6.0 # Required for DBRX tokenizer
|
||||
lm-format-enforcer == 0.11.3
|
||||
llguidance >= 1.7.0, < 1.8.0; platform_machine == "x86_64" or platform_machine == "arm64" or platform_machine == "aarch64" or platform_machine == "ppc64le"
|
||||
|
||||
@@ -11,7 +11,7 @@ numba == 0.65.0 # Required for N-gram speculative decoding
|
||||
datasets
|
||||
peft
|
||||
pytest-asyncio
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
packaging>=24.2
|
||||
setuptools>=77.0.3,<80.0.0
|
||||
setuptools-scm>=8
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# testing
|
||||
pytest
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
pytest-forked
|
||||
pytest-asyncio
|
||||
pytest-rerunfailures
|
||||
|
||||
+49
-24
@@ -35,11 +35,14 @@ arctic-inference==0.1.1
|
||||
# via -r requirements/test/cuda.in
|
||||
argcomplete==3.5.1
|
||||
# via datamodel-code-generator
|
||||
arrow==1.3.0
|
||||
# via isoduration
|
||||
attrs==24.2.0
|
||||
# via
|
||||
# aiohttp
|
||||
# hypothesis
|
||||
# jsonschema
|
||||
# pytest-subtests
|
||||
# referencing
|
||||
audioread==3.0.1
|
||||
# via librosa
|
||||
@@ -54,7 +57,9 @@ azure-identity==1.25.2
|
||||
azure-storage-blob==12.28.0
|
||||
# via runai-model-streamer-azure
|
||||
backoff==2.2.1
|
||||
# via -r requirements/test/cuda.in
|
||||
# via
|
||||
# -r requirements/test/cuda.in
|
||||
# schemathesis
|
||||
bitsandbytes==0.49.2
|
||||
# via -r requirements/test/cuda.in
|
||||
black==24.10.0
|
||||
@@ -105,6 +110,7 @@ colorama==0.4.6
|
||||
# via
|
||||
# perceptron
|
||||
# sacrebleu
|
||||
# schemathesis
|
||||
colorful==0.5.6
|
||||
# via ray
|
||||
colorlog==6.10.1
|
||||
@@ -177,7 +183,7 @@ et-xmlfile==2.0.0
|
||||
# via openpyxl
|
||||
evaluate==0.4.3
|
||||
# via lm-eval
|
||||
fastapi==0.136.3
|
||||
fastapi==0.128.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# gpt-oss
|
||||
@@ -200,6 +206,8 @@ filelock==3.16.1
|
||||
# virtualenv
|
||||
fonttools==4.55.0
|
||||
# via matplotlib
|
||||
fqdn==1.5.1
|
||||
# via jsonschema
|
||||
frozendict==2.4.6
|
||||
# via einx
|
||||
frozenlist==1.5.0
|
||||
@@ -261,7 +269,7 @@ h11==0.14.0
|
||||
# uvicorn
|
||||
h2==4.3.0
|
||||
# via httpx
|
||||
harfile==0.5.0
|
||||
harfile==0.3.0
|
||||
# via schemathesis
|
||||
hf-xet==1.4.3
|
||||
# via huggingface-hub
|
||||
@@ -301,7 +309,7 @@ hypothesis==6.131.0
|
||||
# hypothesis-graphql
|
||||
# hypothesis-jsonschema
|
||||
# schemathesis
|
||||
hypothesis-graphql==0.13.0
|
||||
hypothesis-graphql==0.11.1
|
||||
# via schemathesis
|
||||
hypothesis-jsonschema==0.23.1
|
||||
# via schemathesis
|
||||
@@ -310,6 +318,7 @@ idna==3.10
|
||||
# anyio
|
||||
# email-validator
|
||||
# httpx
|
||||
# jsonschema
|
||||
# requests
|
||||
# yarl
|
||||
imagehash==4.3.2
|
||||
@@ -326,6 +335,8 @@ instanttensor==0.1.5
|
||||
# via -r requirements/test/cuda.in
|
||||
isodate==0.7.2
|
||||
# via azure-storage-blob
|
||||
isoduration==20.11.0
|
||||
# via jsonschema
|
||||
isort==5.13.2
|
||||
# via datamodel-code-generator
|
||||
jinja2==3.1.6
|
||||
@@ -345,14 +356,15 @@ joblib==1.4.2
|
||||
# librosa
|
||||
# nltk
|
||||
# scikit-learn
|
||||
jsonpointer==3.0.0
|
||||
# via jsonschema
|
||||
jsonschema==4.23.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# hypothesis-jsonschema
|
||||
# mistral-common
|
||||
# ray
|
||||
jsonschema-rs==0.46.5
|
||||
# via schemathesis
|
||||
# schemathesis
|
||||
jsonschema-specifications==2024.10.1
|
||||
# via jsonschema
|
||||
junit-xml==1.9
|
||||
@@ -703,20 +715,18 @@ pydantic-core==2.41.1
|
||||
pydantic-extra-types==2.10.5
|
||||
# via mistral-common
|
||||
pygments==2.18.0
|
||||
# via
|
||||
# pytest
|
||||
# rich
|
||||
# via rich
|
||||
pyjwt==2.11.0
|
||||
# via msal
|
||||
pyparsing==3.2.0
|
||||
# via matplotlib
|
||||
pyrate-limiter==4.4.0
|
||||
pyrate-limiter==3.7.0
|
||||
# via schemathesis
|
||||
pystemmer==3.0.0
|
||||
# via mteb
|
||||
pytablewriter==1.2.0
|
||||
# via lm-eval
|
||||
pytest==9.1.0
|
||||
pytest==8.3.5
|
||||
# via
|
||||
# -r requirements/test/cuda.in
|
||||
# buildkite-test-collector
|
||||
@@ -727,9 +737,10 @@ pytest==9.1.0
|
||||
# pytest-mock
|
||||
# pytest-rerunfailures
|
||||
# pytest-shard
|
||||
# pytest-subtests
|
||||
# pytest-timeout
|
||||
# schemathesis
|
||||
pytest-asyncio==1.4.0
|
||||
pytest-asyncio==0.24.0
|
||||
# via -r requirements/test/cuda.in
|
||||
pytest-cov==6.3.0
|
||||
# via -r requirements/test/cuda.in
|
||||
@@ -741,10 +752,13 @@ pytest-rerunfailures==14.0
|
||||
# via -r requirements/test/cuda.in
|
||||
pytest-shard==0.1.2
|
||||
# via -r requirements/test/cuda.in
|
||||
pytest-subtests==0.14.1
|
||||
# via schemathesis
|
||||
pytest-timeout==2.3.1
|
||||
# via -r requirements/test/cuda.in
|
||||
python-dateutil==2.9.0.post0
|
||||
# via
|
||||
# arrow
|
||||
# botocore
|
||||
# matplotlib
|
||||
# pandas
|
||||
@@ -815,12 +829,15 @@ requests==2.32.3
|
||||
# tiktoken
|
||||
responses==0.25.3
|
||||
# via genai-perf
|
||||
rfc3339-validator==0.1.4
|
||||
# via jsonschema
|
||||
rfc3987==1.3.8
|
||||
# via jsonschema
|
||||
rich==13.9.4
|
||||
# via
|
||||
# genai-perf
|
||||
# mteb
|
||||
# perceptron
|
||||
# schemathesis
|
||||
# typer
|
||||
rouge-score==0.1.2
|
||||
# via lm-eval
|
||||
@@ -851,7 +868,7 @@ safetensors==0.7.0
|
||||
# segmentation-models-pytorch
|
||||
# timm
|
||||
# transformers
|
||||
schemathesis==4.21.6
|
||||
schemathesis==3.39.15
|
||||
# via -r requirements/test/cuda.in
|
||||
scikit-image==0.25.2
|
||||
# via albumentations
|
||||
@@ -895,6 +912,7 @@ six==1.16.0
|
||||
# junit-xml
|
||||
# opencensus
|
||||
# python-dateutil
|
||||
# rfc3339-validator
|
||||
# rouge-score
|
||||
smart-open==7.1.0
|
||||
# via ray
|
||||
@@ -920,10 +938,10 @@ sqlalchemy==2.0.41
|
||||
# optuna
|
||||
sqlitedict==2.1.0
|
||||
# via lm-eval
|
||||
starlette==1.3.1
|
||||
starlette==0.50.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# fastapi
|
||||
# schemathesis
|
||||
# starlette-testclient
|
||||
starlette-testclient==0.4.1
|
||||
# via schemathesis
|
||||
@@ -948,8 +966,7 @@ tenacity==9.1.2
|
||||
# gpt-oss
|
||||
# lm-eval
|
||||
# plotly
|
||||
# schemathesis
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
# via -r requirements/test/cuda.in
|
||||
termcolor==3.1.0
|
||||
# via gpt-oss
|
||||
@@ -973,6 +990,10 @@ tokenizers==0.22.2
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/cuda.in
|
||||
# transformers
|
||||
tomli==2.2.1
|
||||
# via schemathesis
|
||||
tomli-w==1.2.0
|
||||
# via schemathesis
|
||||
torch==2.11.0+cu130
|
||||
# via
|
||||
# -c requirements/cuda.txt
|
||||
@@ -1045,6 +1066,8 @@ typer==0.15.2
|
||||
# huggingface-hub
|
||||
# perceptron
|
||||
# transformers
|
||||
types-python-dateutil==2.9.0.20241206
|
||||
# via arrow
|
||||
typing-extensions==4.15.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
@@ -1069,8 +1092,6 @@ typing-extensions==4.15.0
|
||||
# pydantic
|
||||
# pydantic-core
|
||||
# pydantic-extra-types
|
||||
# pytest-asyncio
|
||||
# schemathesis
|
||||
# sentence-transformers
|
||||
# sqlalchemy
|
||||
# starlette
|
||||
@@ -1078,11 +1099,11 @@ typing-extensions==4.15.0
|
||||
# typer
|
||||
# typing-inspection
|
||||
typing-inspection==0.4.2
|
||||
# via
|
||||
# fastapi
|
||||
# pydantic
|
||||
# via pydantic
|
||||
tzdata==2024.2
|
||||
# via pandas
|
||||
uri-template==1.3.0
|
||||
# via jsonschema
|
||||
urllib3==2.2.3
|
||||
# via
|
||||
# blobfile
|
||||
@@ -1101,6 +1122,8 @@ vocos==0.1.0
|
||||
# via -r requirements/test/cuda.in
|
||||
wcwidth==0.2.13
|
||||
# via ftfy
|
||||
webcolors==24.11.1
|
||||
# via jsonschema
|
||||
werkzeug==3.1.3
|
||||
# via schemathesis
|
||||
word2number==1.1
|
||||
@@ -1112,6 +1135,8 @@ xxhash==3.5.0
|
||||
# datasets
|
||||
# evaluate
|
||||
yarl==1.17.1
|
||||
# via aiohttp
|
||||
# via
|
||||
# aiohttp
|
||||
# schemathesis
|
||||
zipp==3.23.0
|
||||
# via importlib-metadata
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# testing
|
||||
pytest
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
pytest-forked
|
||||
pytest-asyncio
|
||||
pytest-rerunfailures
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
# testing
|
||||
pytest
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
pytest-forked
|
||||
pytest-asyncio
|
||||
pytest-rerunfailures
|
||||
|
||||
+48
-22
@@ -51,12 +51,15 @@ arctic-inference==0.1.1
|
||||
# via -r requirements/test/rocm.in
|
||||
argcomplete==3.6.3
|
||||
# via datamodel-code-generator
|
||||
arrow==1.4.0
|
||||
# via isoduration
|
||||
astor==0.8.1
|
||||
# via depyf
|
||||
attrs==26.1.0
|
||||
# via
|
||||
# aiohttp
|
||||
# jsonschema
|
||||
# pytest-subtests
|
||||
# referencing
|
||||
audioread==3.0.1
|
||||
# via librosa
|
||||
@@ -71,7 +74,9 @@ azure-identity==1.25.3
|
||||
azure-storage-blob==12.28.0
|
||||
# via runai-model-streamer-azure
|
||||
backoff==2.2.1
|
||||
# via -r requirements/test/rocm.in
|
||||
# via
|
||||
# -r requirements/test/rocm.in
|
||||
# schemathesis
|
||||
bitsandbytes==0.49.2
|
||||
# via -r requirements/test/rocm.in
|
||||
black==26.3.1
|
||||
@@ -134,6 +139,7 @@ colorama==0.4.6
|
||||
# via
|
||||
# perceptron
|
||||
# sacrebleu
|
||||
# schemathesis
|
||||
colorful==0.5.8
|
||||
# via ray
|
||||
colorlog==6.10.1
|
||||
@@ -252,6 +258,8 @@ filelock==3.25.2
|
||||
# virtualenv
|
||||
fonttools==4.62.1
|
||||
# via matplotlib
|
||||
fqdn==1.5.1
|
||||
# via jsonschema
|
||||
frozendict==2.4.7
|
||||
# via einx
|
||||
frozenlist==1.8.0
|
||||
@@ -320,7 +328,7 @@ h11==0.16.0
|
||||
# uvicorn
|
||||
h2==4.3.0
|
||||
# via httpx
|
||||
harfile==0.5.0
|
||||
harfile==0.4.0
|
||||
# via schemathesis
|
||||
hf-xet==1.4.3
|
||||
# via huggingface-hub
|
||||
@@ -370,7 +378,7 @@ hypothesis==6.151.9
|
||||
# hypothesis-graphql
|
||||
# hypothesis-jsonschema
|
||||
# schemathesis
|
||||
hypothesis-graphql==0.13.0
|
||||
hypothesis-graphql==0.12.0
|
||||
# via schemathesis
|
||||
hypothesis-jsonschema==0.23.1
|
||||
# via schemathesis
|
||||
@@ -379,6 +387,7 @@ idna==3.11
|
||||
# anyio
|
||||
# email-validator
|
||||
# httpx
|
||||
# jsonschema
|
||||
# requests
|
||||
# yarl
|
||||
ijson==3.5.0
|
||||
@@ -399,6 +408,8 @@ interegular==0.3.3
|
||||
# via lm-format-enforcer
|
||||
isodate==0.7.2
|
||||
# via azure-storage-blob
|
||||
isoduration==20.11.0
|
||||
# via jsonschema
|
||||
isort==8.0.1
|
||||
# via datamodel-code-generator
|
||||
jinja2==3.1.6
|
||||
@@ -424,6 +435,8 @@ joblib==1.5.3
|
||||
# librosa
|
||||
# nltk
|
||||
# scikit-learn
|
||||
jsonpointer==3.1.0
|
||||
# via jsonschema
|
||||
jsonschema==4.26.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
@@ -432,8 +445,7 @@ jsonschema==4.26.0
|
||||
# mcp
|
||||
# mistral-common
|
||||
# ray
|
||||
jsonschema-rs==0.46.5
|
||||
# via schemathesis
|
||||
# schemathesis
|
||||
jsonschema-specifications==2025.9.1
|
||||
# via jsonschema
|
||||
junit-xml==1.9
|
||||
@@ -780,7 +792,7 @@ prometheus-client==0.24.1
|
||||
# opentelemetry-exporter-prometheus
|
||||
# prometheus-fastapi-instrumentator
|
||||
# ray
|
||||
prometheus-fastapi-instrumentator==8.0.0
|
||||
prometheus-fastapi-instrumentator==7.1.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
@@ -864,22 +876,20 @@ pydantic-settings==2.13.1
|
||||
# fastapi
|
||||
# mcp
|
||||
pygments==2.19.2
|
||||
# via
|
||||
# pytest
|
||||
# rich
|
||||
# via rich
|
||||
pyjwt==2.12.1
|
||||
# via
|
||||
# mcp
|
||||
# msal
|
||||
pyparsing==3.3.2
|
||||
# via matplotlib
|
||||
pyrate-limiter==4.4.0
|
||||
pyrate-limiter==3.9.0
|
||||
# via schemathesis
|
||||
pystemmer==3.0.0
|
||||
# via mteb
|
||||
pytablewriter==1.2.1
|
||||
# via lm-eval
|
||||
pytest==9.1.0
|
||||
pytest==8.3.5
|
||||
# via
|
||||
# -r requirements/test/rocm.in
|
||||
# buildkite-test-collector
|
||||
@@ -890,9 +900,10 @@ pytest==9.1.0
|
||||
# pytest-mock
|
||||
# pytest-rerunfailures
|
||||
# pytest-shard
|
||||
# pytest-subtests
|
||||
# pytest-timeout
|
||||
# schemathesis
|
||||
pytest-asyncio==1.4.0
|
||||
pytest-asyncio==0.24.0
|
||||
# via -r requirements/test/rocm.in
|
||||
pytest-cov==6.3.0
|
||||
# via -r requirements/test/rocm.in
|
||||
@@ -904,10 +915,13 @@ pytest-rerunfailures==14.0
|
||||
# via -r requirements/test/rocm.in
|
||||
pytest-shard==0.1.2
|
||||
# via -r requirements/test/rocm.in
|
||||
pytest-subtests==0.14.2
|
||||
# via schemathesis
|
||||
pytest-timeout==2.3.1
|
||||
# via -r requirements/test/rocm.in
|
||||
python-dateutil==2.9.0.post0
|
||||
# via
|
||||
# arrow
|
||||
# botocore
|
||||
# matplotlib
|
||||
# pandas
|
||||
@@ -1002,13 +1016,16 @@ requests==2.32.5
|
||||
# tiktoken
|
||||
responses==0.26.0
|
||||
# via genai-perf
|
||||
rfc3339-validator==0.1.4
|
||||
# via jsonschema
|
||||
rfc3987==1.3.8
|
||||
# via jsonschema
|
||||
rich==14.3.3
|
||||
# via
|
||||
# genai-perf
|
||||
# mteb
|
||||
# perceptron
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# typer
|
||||
rich-toolkit==0.19.7
|
||||
# via
|
||||
@@ -1046,7 +1063,7 @@ safetensors==0.7.0
|
||||
# segmentation-models-pytorch
|
||||
# timm
|
||||
# transformers
|
||||
schemathesis==4.21.6
|
||||
schemathesis==3.39.15
|
||||
# via -r requirements/test/rocm.in
|
||||
scikit-image==0.26.0
|
||||
# via albumentations
|
||||
@@ -1103,6 +1120,7 @@ six==1.17.0
|
||||
# junit-xml
|
||||
# opencensus
|
||||
# python-dateutil
|
||||
# rfc3339-validator
|
||||
# rouge-score
|
||||
smart-open==7.5.1
|
||||
# via ray
|
||||
@@ -1131,14 +1149,13 @@ sqlitedict==2.1.0
|
||||
# via lm-eval
|
||||
sse-starlette==3.3.4
|
||||
# via mcp
|
||||
starlette==1.3.1
|
||||
starlette==0.52.1
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# fastapi
|
||||
# mcp
|
||||
# model-hosting-container-standards
|
||||
# prometheus-fastapi-instrumentator
|
||||
# schemathesis
|
||||
# sse-starlette
|
||||
# starlette-testclient
|
||||
starlette-testclient==0.4.1
|
||||
@@ -1165,8 +1182,7 @@ tenacity==9.1.4
|
||||
# via
|
||||
# gpt-oss
|
||||
# lm-eval
|
||||
# schemathesis
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
# via
|
||||
# -c requirements/rocm.txt
|
||||
# -r requirements/test/rocm.in
|
||||
@@ -1199,6 +1215,10 @@ tokenizers==0.22.2
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/rocm.in
|
||||
# transformers
|
||||
tomli==2.4.0
|
||||
# via schemathesis
|
||||
tomli-w==1.2.0
|
||||
# via schemathesis
|
||||
torch-c-dlpack-ext==0.1.5
|
||||
# via tilelang
|
||||
tqdm==4.67.3
|
||||
@@ -1281,10 +1301,8 @@ typing-extensions==4.15.0
|
||||
# pydantic
|
||||
# pydantic-core
|
||||
# pydantic-extra-types
|
||||
# pytest-asyncio
|
||||
# referencing
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# sentence-transformers
|
||||
# sqlalchemy
|
||||
# starlette
|
||||
@@ -1299,6 +1317,10 @@ typing-inspection==0.4.2
|
||||
# mcp
|
||||
# pydantic
|
||||
# pydantic-settings
|
||||
tzdata==2025.3
|
||||
# via arrow
|
||||
uri-template==1.3.0
|
||||
# via jsonschema
|
||||
urllib3==2.6.3
|
||||
# via
|
||||
# blobfile
|
||||
@@ -1329,6 +1351,8 @@ watchfiles==1.1.1
|
||||
# uvicorn
|
||||
wcwidth==0.6.0
|
||||
# via ftfy
|
||||
webcolors==25.10.0
|
||||
# via jsonschema
|
||||
websockets==16.0
|
||||
# via uvicorn
|
||||
werkzeug==3.1.6
|
||||
@@ -1346,7 +1370,9 @@ xxhash==3.6.0
|
||||
# datasets
|
||||
# evaluate
|
||||
yarl==1.23.0
|
||||
# via aiohttp
|
||||
# via
|
||||
# aiohttp
|
||||
# schemathesis
|
||||
z3-solver==4.15.4.0
|
||||
# via tilelang
|
||||
zipp==3.23.0
|
||||
|
||||
@@ -593,9 +593,8 @@ soxr==0.5.0.post1
|
||||
# mistral-common
|
||||
sqlitedict==2.1.0
|
||||
# via lm-eval
|
||||
starlette==1.3.1
|
||||
starlette==1.0.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# fastapi
|
||||
# starlette-testclient
|
||||
starlette-testclient==0.4.1
|
||||
|
||||
@@ -17,4 +17,4 @@ torchaudio
|
||||
torchvision
|
||||
|
||||
auto_round_lib>=0.13.3
|
||||
vllm_xpu_kernels @ https://github.com/vllm-project/vllm-xpu-kernels/releases/download/v0.1.10/vllm_xpu_kernels-0.1.10-cp38-abi3-manylinux_2_28_x86_64.whl
|
||||
vllm_xpu_kernels @ https://github.com/vllm-project/vllm-xpu-kernels/releases/download/v0.1.9.1/vllm_xpu_kernels-0.1.9.1-cp38-abi3-manylinux_2_28_x86_64.whl
|
||||
|
||||
@@ -53,6 +53,38 @@ class SPTestSettings:
|
||||
runner: RunnerOption
|
||||
test_options: SPTestOptions
|
||||
|
||||
@staticmethod
|
||||
def detailed(
|
||||
*,
|
||||
tp_base: int = 2,
|
||||
pp_base: int = 1,
|
||||
multi_node_only: bool = False,
|
||||
runner: RunnerOption = "auto",
|
||||
load_format: str | None = None,
|
||||
):
|
||||
parallel_setups = []
|
||||
for eager_mode_val in [False, True]:
|
||||
for pp_multiplier in [1, 2]:
|
||||
for chunked_prefill_val in [False, True]:
|
||||
parallel_setups.append(
|
||||
ParallelSetup(
|
||||
tp_size=tp_base,
|
||||
pp_size=pp_multiplier * pp_base,
|
||||
fuse_norm_quant=False,
|
||||
fuse_act_quant=False,
|
||||
eager_mode=eager_mode_val,
|
||||
chunked_prefill=chunked_prefill_val,
|
||||
)
|
||||
)
|
||||
return SPTestSettings(
|
||||
parallel_setups=parallel_setups,
|
||||
distributed_backends=["mp", "ray"],
|
||||
runner=runner,
|
||||
test_options=SPTestOptions(
|
||||
multi_node_only=multi_node_only, load_format=load_format
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def fast(
|
||||
*,
|
||||
@@ -62,26 +94,23 @@ class SPTestSettings:
|
||||
multi_node_only: bool = False,
|
||||
load_format: str | None = None,
|
||||
):
|
||||
parallel_setups = []
|
||||
for eager_mode_val in [False, True]:
|
||||
for pp_multiplier in [1, 2]:
|
||||
for chunked_prefill_val in [False, True]:
|
||||
parallel_setups.append(
|
||||
ParallelSetup(
|
||||
tp_size=tp_base,
|
||||
pp_size=pp_multiplier * pp_base,
|
||||
fuse_norm_quant=False,
|
||||
fuse_act_quant=False,
|
||||
eager_mode=eager_mode_val,
|
||||
chunked_prefill=chunked_prefill_val,
|
||||
)
|
||||
)
|
||||
return SPTestSettings(
|
||||
parallel_setups=[
|
||||
ParallelSetup(
|
||||
tp_size=tp_base,
|
||||
pp_size=pp_base,
|
||||
fuse_norm_quant=False,
|
||||
fuse_act_quant=False,
|
||||
eager_mode=False,
|
||||
chunked_prefill=True,
|
||||
),
|
||||
ParallelSetup(
|
||||
tp_size=tp_base,
|
||||
pp_size=2 * pp_base,
|
||||
fuse_norm_quant=False,
|
||||
fuse_act_quant=False,
|
||||
eager_mode=False,
|
||||
chunked_prefill=True,
|
||||
),
|
||||
],
|
||||
distributed_backends=["mp"],
|
||||
parallel_setups=parallel_setups,
|
||||
distributed_backends=["mp", "ray"],
|
||||
runner=runner,
|
||||
test_options=SPTestOptions(
|
||||
multi_node_only=multi_node_only, load_format=load_format
|
||||
|
||||
@@ -8,8 +8,6 @@ AnthropicServingMessages._convert_anthropic_to_openai_request().
|
||||
Also covers extended-thinking edge cases such as ``redacted_thinking``
|
||||
blocks echoed back by Anthropic clients, and streaming conversion in
|
||||
``message_stream_converter``.
|
||||
|
||||
Also covers cache usage computation in ``_build_anthropic_usage``.
|
||||
"""
|
||||
|
||||
import json
|
||||
@@ -20,11 +18,7 @@ import pytest
|
||||
from vllm.entrypoints.anthropic.protocol import (
|
||||
AnthropicMessagesRequest,
|
||||
)
|
||||
from vllm.entrypoints.anthropic.serving import (
|
||||
AnthropicServingMessages,
|
||||
_build_anthropic_usage,
|
||||
_get_cached_tokens,
|
||||
)
|
||||
from vllm.entrypoints.anthropic.serving import AnthropicServingMessages
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionResponseStreamChoice,
|
||||
ChatCompletionStreamResponse,
|
||||
@@ -33,7 +27,6 @@ from vllm.entrypoints.openai.engine.protocol import (
|
||||
DeltaFunctionCall,
|
||||
DeltaMessage,
|
||||
DeltaToolCall,
|
||||
PromptTokenUsageInfo,
|
||||
UsageInfo,
|
||||
)
|
||||
|
||||
@@ -660,108 +653,6 @@ class TestThinkingBlockConversion:
|
||||
assert asst.get("content") == "Hi!"
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Cache usage computation
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestGetCachedTokens:
|
||||
"""Tests for _get_cached_tokens helper."""
|
||||
|
||||
def test_none_usage(self):
|
||||
assert _get_cached_tokens(None) is None
|
||||
|
||||
def test_no_prompt_tokens_details(self):
|
||||
usage = UsageInfo(prompt_tokens=100, completion_tokens=10)
|
||||
assert _get_cached_tokens(usage) is None
|
||||
|
||||
def test_cached_tokens_present(self):
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=10,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
|
||||
)
|
||||
assert _get_cached_tokens(usage) == 80
|
||||
|
||||
def test_cached_tokens_zero(self):
|
||||
"""Zero cached tokens should return 0, not None."""
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=10,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
|
||||
)
|
||||
assert _get_cached_tokens(usage) == 0
|
||||
|
||||
def test_cached_tokens_none_in_details(self):
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=10,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=None),
|
||||
)
|
||||
assert _get_cached_tokens(usage) is None
|
||||
|
||||
|
||||
class TestBuildAnthropicUsage:
|
||||
"""Tests for _build_anthropic_usage helper.
|
||||
|
||||
Anthropic defines: total_input = input_tokens + cache_read + cache_creation
|
||||
vLLM's prompt_tokens is the total.
|
||||
"""
|
||||
|
||||
def test_no_cache_info(self):
|
||||
"""When cache info is unavailable, return raw prompt_tokens."""
|
||||
result = _build_anthropic_usage(100, 10, None)
|
||||
assert result.input_tokens == 100
|
||||
assert result.output_tokens == 10
|
||||
assert result.cache_read_input_tokens is None
|
||||
assert result.cache_creation_input_tokens is None
|
||||
|
||||
def test_cache_hit(self):
|
||||
"""When cache is hit, input_tokens excludes cached tokens."""
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=10,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
|
||||
)
|
||||
result = _build_anthropic_usage(100, 10, usage)
|
||||
assert result.input_tokens == 20 # 100 - 80
|
||||
assert result.output_tokens == 10
|
||||
assert result.cache_read_input_tokens == 80
|
||||
assert result.cache_creation_input_tokens == 0
|
||||
|
||||
def test_zero_cached_tokens(self):
|
||||
"""Zero cached tokens should still set cache_creation to 0."""
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=10,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
|
||||
)
|
||||
result = _build_anthropic_usage(100, 10, usage)
|
||||
assert result.input_tokens == 100 # 100 - 0
|
||||
assert result.cache_read_input_tokens == 0
|
||||
assert result.cache_creation_input_tokens == 0
|
||||
|
||||
def test_all_tokens_cached(self):
|
||||
"""When all tokens are cached, input_tokens should be 0."""
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=10,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=100),
|
||||
)
|
||||
result = _build_anthropic_usage(100, 10, usage)
|
||||
assert result.input_tokens == 0
|
||||
assert result.cache_read_input_tokens == 100
|
||||
assert result.cache_creation_input_tokens == 0
|
||||
|
||||
def test_no_prompt_tokens_details(self):
|
||||
"""UsageInfo without prompt_tokens_details returns no cache info."""
|
||||
usage = UsageInfo(prompt_tokens=100, completion_tokens=10)
|
||||
result = _build_anthropic_usage(100, 10, usage)
|
||||
assert result.input_tokens == 100
|
||||
assert result.cache_read_input_tokens is None
|
||||
assert result.cache_creation_input_tokens is None
|
||||
|
||||
|
||||
class TestInlineSystemMessageInMessagesArray:
|
||||
"""Verify that ``role: system`` messages embedded inside the ``messages``
|
||||
array are preserved in their original position.
|
||||
@@ -1205,179 +1096,3 @@ class TestMessageStartIncludesTypeAndRole:
|
||||
message = events[0][1]["message"]
|
||||
assert message["type"] == "message"
|
||||
assert message["role"] == "assistant"
|
||||
|
||||
|
||||
class TestStreamingCacheUsageSemantics:
|
||||
"""Locks in the documented streaming behavior of cache usage fields.
|
||||
|
||||
vLLM's OpenAI chat completion streaming only attaches
|
||||
``prompt_tokens_details`` to the terminal usage chunk. The Anthropic layer
|
||||
mirrors that contract: cache fields are omitted on ``message_start`` (key
|
||||
absence signals "unknown") and populated on ``message_delta`` (the final
|
||||
cumulative count). This is intentionally consistent with vLLM's OpenAI
|
||||
behavior, even though Anthropic's upstream API populates cache fields on
|
||||
``message_start``; closing that gap requires plumbing cache info into the
|
||||
first chunk at the OpenAI layer, which is out of scope here.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_cache_fields_absent_then_populated(self):
|
||||
"""First chunk lacks prompt_tokens_details (vLLM contract);
|
||||
message_start omits cache fields. The final chunk carries
|
||||
prompt_tokens_details, so message_delta carries resolved values."""
|
||||
|
||||
async def sse_input():
|
||||
yield _make_stream_chunk(
|
||||
delta=DeltaMessage(role="assistant", content="hi"),
|
||||
usage=UsageInfo(prompt_tokens=100, total_tokens=100),
|
||||
)
|
||||
yield _make_stream_chunk(finish_reason="stop")
|
||||
yield _make_stream_chunk(
|
||||
choices=[],
|
||||
usage=UsageInfo(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=5,
|
||||
total_tokens=105,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
|
||||
),
|
||||
)
|
||||
yield "data: [DONE]"
|
||||
|
||||
converter = _make_stream_converter()
|
||||
output = []
|
||||
async for event in converter.message_stream_converter(sse_input()):
|
||||
output.append(event)
|
||||
events = _parse_sse_events(output)
|
||||
|
||||
# message_start: cache fields unknown → omitted from JSON entirely.
|
||||
start_usage = events[0][1]["message"]["usage"]
|
||||
assert events[0][0] == "message_start"
|
||||
assert start_usage["input_tokens"] == 100
|
||||
assert "cache_read_input_tokens" not in start_usage
|
||||
assert "cache_creation_input_tokens" not in start_usage
|
||||
|
||||
# message_delta: authoritative usage with cache fields populated.
|
||||
delta_usage = next(
|
||||
data["usage"] for ev, data in events if ev == "message_delta"
|
||||
)
|
||||
assert delta_usage["input_tokens"] == 20 # 100 - 80
|
||||
assert delta_usage["cache_read_input_tokens"] == 80
|
||||
assert delta_usage["cache_creation_input_tokens"] == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_no_cache_hit(self):
|
||||
"""When the final chunk reports cached_tokens=0, message_delta carries
|
||||
cache fields = 0 (cache miss); message_start still omits them."""
|
||||
|
||||
async def sse_input():
|
||||
yield _make_stream_chunk(
|
||||
delta=DeltaMessage(role="assistant"),
|
||||
usage=UsageInfo(prompt_tokens=50, total_tokens=50),
|
||||
)
|
||||
yield _make_stream_chunk(finish_reason="stop")
|
||||
yield _make_stream_chunk(
|
||||
choices=[],
|
||||
usage=UsageInfo(
|
||||
prompt_tokens=50,
|
||||
completion_tokens=5,
|
||||
total_tokens=55,
|
||||
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
|
||||
),
|
||||
)
|
||||
yield "data: [DONE]"
|
||||
|
||||
converter = _make_stream_converter()
|
||||
output = []
|
||||
async for event in converter.message_stream_converter(sse_input()):
|
||||
output.append(event)
|
||||
events = _parse_sse_events(output)
|
||||
|
||||
start_usage = events[0][1]["message"]["usage"]
|
||||
delta_usage = next(
|
||||
data["usage"] for ev, data in events if ev == "message_delta"
|
||||
)
|
||||
assert start_usage["input_tokens"] == 50
|
||||
assert "cache_read_input_tokens" not in start_usage
|
||||
assert "cache_creation_input_tokens" not in start_usage
|
||||
assert delta_usage["input_tokens"] == 50 # 50 - 0
|
||||
assert delta_usage["cache_read_input_tokens"] == 0
|
||||
assert delta_usage["cache_creation_input_tokens"] == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_no_prompt_tokens_details_at_all(self):
|
||||
"""If --enable-prompt-tokens-details is off, no chunk carries cache
|
||||
info; both message_start and message_delta omit cache fields."""
|
||||
|
||||
async def sse_input():
|
||||
yield _make_stream_chunk(
|
||||
delta=DeltaMessage(role="assistant"),
|
||||
usage=UsageInfo(prompt_tokens=30, total_tokens=30),
|
||||
)
|
||||
yield _make_stream_chunk(finish_reason="stop")
|
||||
yield _make_stream_chunk(
|
||||
choices=[],
|
||||
usage=UsageInfo(prompt_tokens=30, completion_tokens=2, total_tokens=32),
|
||||
)
|
||||
yield "data: [DONE]"
|
||||
|
||||
converter = _make_stream_converter()
|
||||
output = []
|
||||
async for event in converter.message_stream_converter(sse_input()):
|
||||
output.append(event)
|
||||
events = _parse_sse_events(output)
|
||||
|
||||
start_usage = events[0][1]["message"]["usage"]
|
||||
delta_usage = next(
|
||||
data["usage"] for ev, data in events if ev == "message_delta"
|
||||
)
|
||||
assert "cache_read_input_tokens" not in start_usage
|
||||
assert "cache_creation_input_tokens" not in start_usage
|
||||
assert "cache_read_input_tokens" not in delta_usage
|
||||
assert "cache_creation_input_tokens" not in delta_usage
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Auto-detection of system-first template requirement
|
||||
# ======================================================================
|
||||
|
||||
|
||||
Q35_TEMPLATE = (
|
||||
"{%- for message in messages %}"
|
||||
"{%- if message.role == 'system' %}"
|
||||
"{%- if not loop.first %}"
|
||||
"{{- raise_exception('System message must be at the beginning.') }}"
|
||||
"{%- endif %}"
|
||||
"{%- endif %}"
|
||||
"{%- endfor %}"
|
||||
)
|
||||
|
||||
|
||||
class TestDetectMergeInlineSystem:
|
||||
"""Verify _detect_merge_inline_system auto-detection.
|
||||
|
||||
Tests three scenarios:
|
||||
1. Template with system-first guard (e.g. Qwen) → merge needed
|
||||
2. Template without restrictions → no merge, cache-friendly
|
||||
3. No template provided → safe default: merge
|
||||
"""
|
||||
|
||||
def test_qwen_template_requires_merge(self):
|
||||
"""Template with loop.first guard rejects mid-conversation system."""
|
||||
assert (
|
||||
AnthropicServingMessages._detect_merge_inline_system(Q35_TEMPLATE) is True
|
||||
)
|
||||
|
||||
def test_no_restriction_no_merge(self):
|
||||
"""Template without restriction accepts mid-conversation system."""
|
||||
assert (
|
||||
AnthropicServingMessages._detect_merge_inline_system(
|
||||
"{%- for message in messages %}"
|
||||
"{{- message.role }}: {{ message.content }}\n"
|
||||
"{%- endfor %}"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_no_template_defaults_merge(self):
|
||||
"""No chat_template → conservative default: merge."""
|
||||
assert AnthropicServingMessages._detect_merge_inline_system(None) is True
|
||||
|
||||
@@ -1,65 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Tests that validation_exception_handler populates the `param` field
|
||||
in its error response using the Pydantic error's `loc`, even when no
|
||||
custom VLLMValidationError context is present.
|
||||
|
||||
Previously, `param` was only populated for errors carrying a custom
|
||||
VLLMValidationError in their Pydantic `ctx`. Plain validation failures
|
||||
(missing fields, wrong types) left `param` as None, even though the
|
||||
field name was readily available from `error['loc']`.
|
||||
"""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
|
||||
from vllm.entrypoints.serve.utils.server_utils import validation_exception_handler
|
||||
|
||||
|
||||
def _fake_request(log_error_stack: bool = False) -> SimpleNamespace:
|
||||
"""Minimal stand-in for a FastAPI Request - just enough for the
|
||||
handler to read req.app.state.args.log_error_stack."""
|
||||
return SimpleNamespace(
|
||||
app=SimpleNamespace(
|
||||
state=SimpleNamespace(args=SimpleNamespace(log_error_stack=log_error_stack))
|
||||
),
|
||||
state=SimpleNamespace(), # no request_metadata -> hasattr(...) is False
|
||||
)
|
||||
|
||||
|
||||
class TestValidationErrorParamFallback:
|
||||
"""Ensure `param` falls back to the Pydantic error's `loc` when no
|
||||
custom VLLMValidationError context is present."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("error_type", "msg"),
|
||||
[
|
||||
("missing", "Field required"),
|
||||
("list_type", "Input should be a valid list"),
|
||||
],
|
||||
ids=["missing-field", "wrong-type"],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_param_falls_back_to_loc(self, error_type: str, msg: str):
|
||||
errors = [{"type": error_type, "loc": ("body", "messages"), "msg": msg}]
|
||||
exc = RequestValidationError(errors)
|
||||
|
||||
response = await validation_exception_handler(_fake_request(), exc)
|
||||
body = json.loads(response.body)
|
||||
|
||||
assert body["error"]["param"] == "body.messages"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_param_fallback_does_not_crash_on_non_dict_error(self):
|
||||
"""Schemathesis fuzzing found that errors[0] isn't always a dict.
|
||||
The fallback must not crash in that case - it should just leave
|
||||
param as None instead of raising."""
|
||||
exc = RequestValidationError(["some unexpected non-dict error"])
|
||||
|
||||
response = await validation_exception_handler(_fake_request(), exc)
|
||||
body = json.loads(response.body)
|
||||
|
||||
assert body["error"]["param"] is None
|
||||
@@ -482,127 +482,3 @@ def test_kernels_hidden_size(
|
||||
seq_length=128,
|
||||
add_inputs=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("device", DEVICES)
|
||||
def test_add_lora_fused_moe_early_exit(device):
|
||||
"""
|
||||
Ensures add_lora_fused_moe does not invoke the LoRA kernel or
|
||||
modify the output tensor when no_lora_flag_cpu is True
|
||||
"""
|
||||
from types import SimpleNamespace
|
||||
|
||||
from vllm.lora.punica_wrapper.punica_gpu import PunicaWrapperGPU
|
||||
|
||||
torch.set_default_device(device)
|
||||
torch.accelerator.set_device_index(device)
|
||||
|
||||
max_loras, num_tokens = 4, 16
|
||||
num_experts, top_k, max_lora_rank = 8, 2, 16
|
||||
K, N = 256, 128
|
||||
|
||||
# build PunicaWrapperGPU with minimal lora_config mock
|
||||
lora_config = SimpleNamespace(
|
||||
max_loras=max_loras,
|
||||
specialize_active_lora=False,
|
||||
)
|
||||
wrapper = PunicaWrapperGPU(
|
||||
max_num_batched_tokens=num_tokens,
|
||||
max_batches=num_tokens,
|
||||
device=device,
|
||||
lora_config=lora_config,
|
||||
)
|
||||
|
||||
# simulate a prior LoRA batch so the internal mapping is
|
||||
# populated with stale LoRA IDs
|
||||
lora_mapping = torch.zeros(
|
||||
num_tokens,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
lora_mapping[:8] = 1
|
||||
lora_mapping[8:] = 2
|
||||
wrapper.token_mapping_meta.prepare_tensors(lora_mapping)
|
||||
|
||||
# simulate a base-model batch (all -1)
|
||||
base_mapping = torch.full(
|
||||
(num_tokens,),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
wrapper.token_mapping_meta.prepare_tensors(base_mapping)
|
||||
|
||||
assert wrapper.token_mapping_meta.no_lora_flag_cpu[0].item() is True
|
||||
|
||||
# dummy tensors for add_lora_fused_moe
|
||||
y = torch.rand(num_tokens, top_k, N, dtype=torch.bfloat16, device=device)
|
||||
y_snapshot = y.clone()
|
||||
x = torch.rand(num_tokens, K, dtype=torch.bfloat16, device=device)
|
||||
|
||||
lora_a_stacked = (
|
||||
torch.rand(
|
||||
max_loras,
|
||||
num_experts,
|
||||
max_lora_rank,
|
||||
K,
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
),
|
||||
)
|
||||
lora_b_stacked = (
|
||||
torch.rand(
|
||||
max_loras,
|
||||
num_experts,
|
||||
N,
|
||||
max_lora_rank,
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
),
|
||||
)
|
||||
topk_weights = torch.ones(
|
||||
num_tokens,
|
||||
top_k,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
adapter_enabled = torch.ones(
|
||||
max_loras + 1,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
shrink_config = expand_config = {
|
||||
"BLOCK_SIZE_M": 16,
|
||||
"BLOCK_SIZE_N": 32,
|
||||
"BLOCK_SIZE_K": 64,
|
||||
"GROUP_SIZE_M": 1,
|
||||
"NUM_WARPS": 4,
|
||||
"NUM_STAGES": 3,
|
||||
"SPLIT_K": 1,
|
||||
}
|
||||
|
||||
# call add_lora_fused_moe - the early exit should prevent any
|
||||
# modification to the output
|
||||
wrapper.add_lora_fused_moe(
|
||||
y=y,
|
||||
x=x,
|
||||
lora_a_stacked=lora_a_stacked,
|
||||
lora_b_stacked=lora_b_stacked,
|
||||
topk_weights=topk_weights,
|
||||
sorted_token_ids=None,
|
||||
expert_ids=torch.zeros(
|
||||
num_tokens * top_k,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
),
|
||||
num_tokens_post_padded=None,
|
||||
max_lora_rank=max_lora_rank,
|
||||
top_k_num=top_k,
|
||||
shrink_config=shrink_config,
|
||||
expand_config=expand_config,
|
||||
adapter_enabled=adapter_enabled,
|
||||
)
|
||||
|
||||
assert torch.equal(y, y_snapshot), (
|
||||
"add_lora_fused_moe modified output tensor despite no_lora_flag_cpu=True"
|
||||
)
|
||||
|
||||
@@ -130,12 +130,8 @@ def test_models(
|
||||
monkeypatch.setenv("VLLM_ROCM_USE_AITER", "1")
|
||||
if model == "TitanML/tiny-mixtral":
|
||||
# Untrained model: near-uniform logits make argmax sensitive to
|
||||
# AITER's bfloat16 rounding error. Route the plain rms_norm and the
|
||||
# fused MoE (whose near-uniform router logits flip expert selection
|
||||
# under ~1 ULP drift) through the native kernels for this model.
|
||||
# See ROCm/aiter#3806 for the tracking issue and minimal repro.
|
||||
# AITER's bfloat16 rounding error in plain rms_norm.
|
||||
monkeypatch.setenv("VLLM_ROCM_USE_AITER_RMSNORM", "0")
|
||||
monkeypatch.setenv("VLLM_ROCM_USE_AITER_MOE", "0")
|
||||
elif use_rocm_aiter and model not in AITER_MODEL_LIST:
|
||||
# Skip model that are not using AITER tests.
|
||||
# When more AITER kernels are added, this list will not be
|
||||
|
||||
@@ -25,7 +25,7 @@ TEST_IMAGE_NAMES = [
|
||||
]
|
||||
MAX_MODEL_LEN = 8192
|
||||
REQUESTS_PER_ROUND = 4
|
||||
WARMUP_ROUNDS = 2
|
||||
WARMUP_ROUNDS = 1
|
||||
MEASURED_ROUNDS = 16
|
||||
GPU_GROWTH_THRESHOLD_MIB = 0
|
||||
CPU_PEAK_GROWTH_THRESHOLD_MIB = 0
|
||||
|
||||
@@ -143,12 +143,6 @@ SCENARIOS: list[Scenario] = [
|
||||
tool_calls=[_READ_TOOL],
|
||||
after_tool_response=True,
|
||||
),
|
||||
Scenario(
|
||||
id="empty-tool-block",
|
||||
description="Empty tool block followed by content (edge case recovery)",
|
||||
content="Content after empty tools.",
|
||||
tool_calls=[],
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@@ -350,11 +344,8 @@ def _qwen3_segments(scenario: Scenario) -> list[tuple[str, bool]]:
|
||||
segs: list[tuple[str, bool]] = []
|
||||
if scenario.reasoning is not None:
|
||||
segs.append((scenario.reasoning, False))
|
||||
if scenario.content is not None or scenario.tool_calls is not None:
|
||||
if scenario.content is not None or scenario.tool_calls:
|
||||
segs.append(("</think>", True))
|
||||
if scenario.tool_calls is not None and not scenario.tool_calls:
|
||||
segs.append(("<tool_call>", True))
|
||||
segs.append(("</tool_call>", True))
|
||||
if scenario.content is not None:
|
||||
segs.append((scenario.content, False))
|
||||
if scenario.tool_calls:
|
||||
@@ -446,11 +437,8 @@ def _minimax_m2_segments(scenario: Scenario) -> list[tuple[str, bool]]:
|
||||
segs: list[tuple[str, bool]] = []
|
||||
if scenario.reasoning is not None:
|
||||
segs.append((scenario.reasoning, False))
|
||||
if scenario.content is not None or scenario.tool_calls is not None:
|
||||
if scenario.content is not None or scenario.tool_calls:
|
||||
segs.append(("</think>", True))
|
||||
if scenario.tool_calls is not None and not scenario.tool_calls:
|
||||
segs.append(("<minimax:tool_call>", True))
|
||||
segs.append(("</minimax:tool_call>", True))
|
||||
if scenario.content is not None:
|
||||
segs.append((scenario.content, False))
|
||||
if scenario.tool_calls:
|
||||
@@ -546,9 +534,6 @@ def _gemma4_segments(scenario: Scenario) -> list[tuple[str, bool]]:
|
||||
segs.append((_GEMMA4_THOUGHT_PREFIX, False))
|
||||
segs.append((scenario.reasoning, False))
|
||||
segs.append(("<channel|>", True))
|
||||
if scenario.tool_calls is not None and not scenario.tool_calls:
|
||||
segs.append(("<|tool_call>", True))
|
||||
segs.append(("<tool_call|>", True))
|
||||
if scenario.content is not None:
|
||||
segs.append((scenario.content, False))
|
||||
if scenario.tool_calls:
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import threading
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from vllm.config import set_current_vllm_config
|
||||
@@ -408,9 +406,7 @@ def test_lookup_key_client_lookup_prepends_typed_tag():
|
||||
fake_socket = mock_make_socket.return_value
|
||||
fake_socket.recv.return_value = (5).to_bytes(4, "big")
|
||||
|
||||
# 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(token_len=128, block_hashes=[]) == 5
|
||||
|
||||
sent_frames = fake_socket.send_multipart.call_args[0][0]
|
||||
assert sent_frames[0] == protocol.LOOKUP_MSG
|
||||
@@ -439,127 +435,6 @@ 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):
|
||||
"""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)
|
||||
if result is not None:
|
||||
return result
|
||||
time.sleep(0.005)
|
||||
return None
|
||||
|
||||
|
||||
def _gated_recv(gate: threading.Event, value: int):
|
||||
"""Mock recv side-effect that blocks until ``gate`` is set, so the
|
||||
executor's lookup can be held pending deterministically."""
|
||||
|
||||
def recv():
|
||||
gate.wait()
|
||||
return value.to_bytes(4, "big")
|
||||
|
||||
return recv
|
||||
|
||||
|
||||
def test_lookup_key_client_non_block_lookup_async():
|
||||
"""Non-blocking lookup defers to the executor: None first, hit once the
|
||||
Future resolves."""
|
||||
vllm_config = _make_vllm_config()
|
||||
|
||||
with patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.make_zmq_socket"
|
||||
) as mock_make_socket:
|
||||
client = worker.LookupKeyClient(vllm_config)
|
||||
|
||||
fake_socket = mock_make_socket.return_value
|
||||
# Hold the executor's lookup pending until we release the gate.
|
||||
gate = threading.Event()
|
||||
fake_socket.recv.side_effect = _gated_recv(gate, 7)
|
||||
|
||||
# First query submits the lookup and returns None while it is in flight.
|
||||
assert client.lookup("req1", 128, [], non_block=True) is None
|
||||
# Release the executor; a later poll returns the hit length.
|
||||
gate.set()
|
||||
assert _poll_lookup(client, "req1") == 7
|
||||
# Future is consumed (popped) on read.
|
||||
assert "req1" not in client.futures
|
||||
|
||||
|
||||
def test_lookup_key_client_discard_clears_state():
|
||||
"""discard() drops a completed lookup Future so it is not served stale."""
|
||||
vllm_config = _make_vllm_config()
|
||||
|
||||
with patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.make_zmq_socket"
|
||||
) as mock_make_socket:
|
||||
client = worker.LookupKeyClient(vllm_config)
|
||||
|
||||
fake_socket = mock_make_socket.return_value
|
||||
gate = threading.Event()
|
||||
fake_socket.recv.side_effect = _gated_recv(gate, 9)
|
||||
|
||||
# Submit while gated so the call returns None and the Future stays in
|
||||
# `futures` (unconsumed) once it resolves.
|
||||
assert client.lookup("req2", 128, [], non_block=True) is None
|
||||
gate.set()
|
||||
deadline = time.monotonic() + 5.0
|
||||
while time.monotonic() < deadline:
|
||||
if client.futures["req2"].done():
|
||||
break
|
||||
time.sleep(0.005)
|
||||
# discard() drops the completed result before any lookup consumes it.
|
||||
client.discard("req2")
|
||||
assert "req2" not in client.futures
|
||||
# A fresh query re-submits rather than returning a stale value: hold the
|
||||
# gate so the resubmitted lookup stays in flight.
|
||||
gate.clear()
|
||||
assert client.lookup("req2", 128, [], non_block=True) is None
|
||||
gate.set() # release the executor so the worker thread can drain
|
||||
|
||||
|
||||
def test_get_num_new_matched_tokens_async_defers_then_reports():
|
||||
"""Async lookup returns (None, False) until ready, then the hit count."""
|
||||
vllm_config = create_vllm_config(
|
||||
kv_connector="MooncakeStoreConnector",
|
||||
kv_role="kv_both",
|
||||
kv_connector_extra_config={"lookup_async": True},
|
||||
)
|
||||
kv_cache_config = _make_kv_cache_config()
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"scheduler.LookupKeyClient"
|
||||
) as mock_client_cls,
|
||||
):
|
||||
sched = scheduler.MooncakeStoreScheduler(vllm_config, kv_cache_config)
|
||||
|
||||
assert sched.lookup_async is True
|
||||
mock_client = mock_client_cls.return_value
|
||||
|
||||
block_size = sched._block_size
|
||||
request = MagicMock()
|
||||
request.request_id = "r1"
|
||||
request.num_tokens = 4 * block_size
|
||||
request.block_hashes = []
|
||||
|
||||
# Lookup not ready -> defer.
|
||||
mock_client.lookup.return_value = None
|
||||
assert sched.get_num_new_matched_tokens(request, 0) == (None, False)
|
||||
assert "r1" not in sched.load_specs
|
||||
|
||||
# Lookup ready with a hit -> report need_to_allocate + async-load flag.
|
||||
hit = 3 * block_size
|
||||
mock_client.lookup.return_value = hit
|
||||
need, load_async = sched.get_num_new_matched_tokens(request, 0)
|
||||
assert need == hit
|
||||
assert load_async == sched.load_async
|
||||
assert sched.load_specs["r1"].kvpool_cached_tokens == hit
|
||||
|
||||
|
||||
def test_protocol_tags_are_distinct_and_non_empty():
|
||||
"""Protocol tags must be unique and non-empty to avoid collision."""
|
||||
tags = {protocol.LOOKUP_MSG, protocol.RESET_MSG}
|
||||
|
||||
@@ -16,7 +16,6 @@ from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.scheduler impor
|
||||
def _make_bare_scheduler() -> MooncakeStoreScheduler:
|
||||
scheduler = object.__new__(MooncakeStoreScheduler)
|
||||
scheduler.kv_role = "kv_both"
|
||||
scheduler.lookup_async = False
|
||||
scheduler._block_size = 16
|
||||
scheduler.load_specs = {}
|
||||
scheduler._preempted_req_ids = set()
|
||||
@@ -406,13 +405,7 @@ class _StubLookupClient:
|
||||
def __init__(self, hit_tokens: int) -> None:
|
||||
self._hit_tokens = hit_tokens
|
||||
|
||||
def lookup(
|
||||
self,
|
||||
req_id: str,
|
||||
token_len: int,
|
||||
block_hashes: list[bytes],
|
||||
non_block: bool = False,
|
||||
) -> int:
|
||||
def lookup(self, token_len: int, block_hashes: list[bytes]) -> int:
|
||||
return self._hit_tokens
|
||||
|
||||
|
||||
|
||||
@@ -175,17 +175,14 @@ class _FakeModelConfig:
|
||||
|
||||
|
||||
def _make_vllm_config(
|
||||
*,
|
||||
extra_config: dict[str, object] | None = None,
|
||||
rank: int = 0,
|
||||
decode_context_parallel_size: int = 1,
|
||||
*, extra_config: dict[str, object] | None = None
|
||||
) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
model_config=_FakeModelConfig(),
|
||||
parallel_config=SimpleNamespace(
|
||||
pipeline_parallel_size=1,
|
||||
rank=rank,
|
||||
decode_context_parallel_size=decode_context_parallel_size,
|
||||
rank=0,
|
||||
decode_context_parallel_size=1,
|
||||
prefill_context_parallel_size=1,
|
||||
),
|
||||
kv_transfer_config=_FakeKVTransferConfig(extra_config=extra_config),
|
||||
@@ -234,23 +231,13 @@ def _install_fake_mooncake(monkeypatch, store_instance: MagicMock):
|
||||
return FakeReplicateConfig
|
||||
|
||||
|
||||
def _patch_worker_runtime(
|
||||
monkeypatch,
|
||||
*,
|
||||
local_ip: str = "10.0.0.7",
|
||||
tp_rank: int = 0,
|
||||
tp_size: int = 1,
|
||||
dcp_size: int = 1,
|
||||
) -> None:
|
||||
def _patch_worker_runtime(monkeypatch, *, local_ip: str = "10.0.0.7") -> None:
|
||||
single_rank_group = SimpleNamespace(world_size=1, rank_in_group=0)
|
||||
# DCP groups are contiguous splits of the TP group (see
|
||||
# parallel_state.py), so dcp_rank == tp_rank % dcp_size.
|
||||
dcp_group = SimpleNamespace(world_size=dcp_size, rank_in_group=tp_rank % dcp_size)
|
||||
monkeypatch.setattr(worker, "get_mooncake_dp_engine_index", lambda _: 0)
|
||||
monkeypatch.setattr(worker, "get_tensor_model_parallel_rank", lambda: tp_rank)
|
||||
monkeypatch.setattr(worker, "get_tensor_model_parallel_world_size", lambda: tp_size)
|
||||
monkeypatch.setattr(worker, "get_tensor_model_parallel_rank", lambda: 0)
|
||||
monkeypatch.setattr(worker, "get_tensor_model_parallel_world_size", lambda: 1)
|
||||
monkeypatch.setattr(worker, "get_pcp_group", lambda: single_rank_group)
|
||||
monkeypatch.setattr(worker, "get_dcp_group", lambda: dcp_group)
|
||||
monkeypatch.setattr(worker, "get_dcp_group", lambda: single_rank_group)
|
||||
monkeypatch.setattr(worker, "get_ip", lambda: local_ip)
|
||||
|
||||
|
||||
@@ -897,66 +884,6 @@ def test_requester_worker_init_builds_replicate_config_for_preferred_segment(
|
||||
assert w.store_replicate_config.preferred_segment == "10.0.0.7:50053"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dcp_size", [1, 4])
|
||||
def test_worker_put_striding_covers_every_rank_get_namespace(
|
||||
tmp_path, monkeypatch, dcp_size
|
||||
):
|
||||
"""Every key a rank GETs must have been PUT by some rank.
|
||||
|
||||
When num_kv_head < tp_size, ranks holding the same KV heads stripe
|
||||
their PUTs across one shared key namespace. That dedup is only valid
|
||||
when those ranks really share a namespace: with DCP > 1 each rank GETs
|
||||
every key from its own ``@dcpN`` namespace, so striding must be
|
||||
disabled.
|
||||
"""
|
||||
tp_size = 4
|
||||
store = MagicMock()
|
||||
store.setup.return_value = 0
|
||||
_install_fake_mooncake(monkeypatch, store)
|
||||
monkeypatch.setenv(
|
||||
"MOONCAKE_CONFIG_PATH",
|
||||
_write_mooncake_config(
|
||||
tmp_path,
|
||||
{
|
||||
"metadata_server": "http://metadata/endpoint",
|
||||
"protocol": "tcp",
|
||||
"device_name": "",
|
||||
"master_server_address": "10.0.0.7:50051",
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# _FakeModelConfig has num_kv_head=1 < tp_size, which enables striding.
|
||||
block_hashes = [f"hash-{i}".encode() for i in range(4)]
|
||||
put_keys: set[str] = set()
|
||||
get_keys_per_rank: dict[int, set[str]] = {}
|
||||
for tp_rank in range(tp_size):
|
||||
_patch_worker_runtime(
|
||||
monkeypatch, tp_rank=tp_rank, tp_size=tp_size, dcp_size=dcp_size
|
||||
)
|
||||
w = worker.MooncakeStoreWorker(
|
||||
_make_vllm_config(rank=tp_rank, decode_context_parallel_size=dcp_size),
|
||||
_make_kv_cache_config(),
|
||||
)
|
||||
db = w.token_dbs[0]
|
||||
token_len = len(block_hashes) * db.block_size
|
||||
keys = [
|
||||
key.to_string() for _, _, key in db.process_tokens(token_len, block_hashes)
|
||||
]
|
||||
assert len(keys) == len(block_hashes)
|
||||
# PUT side: mirrors KVCacheStoreSendingThread's striding slice.
|
||||
put_keys.update(keys[w.tp_rank % w.put_step :: w.put_step])
|
||||
# GET side: KVCacheStoreRecvingThread fetches every key.
|
||||
get_keys_per_rank[tp_rank] = set(keys)
|
||||
|
||||
for tp_rank, rank_keys in get_keys_per_rank.items():
|
||||
missing = rank_keys - put_keys
|
||||
assert not missing, (
|
||||
f"tp_rank={tp_rank} would GET {len(missing)}/{len(rank_keys)} keys "
|
||||
f"that no rank PUT (Mooncake OBJECT_NOT_FOUND): {sorted(missing)}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers for register_kv_caches tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -1,227 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Tests for the Model Runner V2 Gumbel-max sampling kernel.
|
||||
|
||||
Accuracy: define a target categorical distribution as a non-negative int64
|
||||
count tensor summing to N, turn it into logits (= log(count)), sample many
|
||||
times with `gumbel_sample`, and check the empirical distribution matches.
|
||||
|
||||
The count tensor is deliberately heavy-tailed (one dominant token, the rest
|
||||
~18 logits below). That tail is the sensitive part: the fp32 Gumbel noise must
|
||||
reach ~18 to ever sample it. A flat distribution would keep every token within
|
||||
a few logits of the top and would not exercise the noise tail at all.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
pytest.importorskip("triton")
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required for Gumbel sampler tests", allow_module_level=True)
|
||||
|
||||
from vllm.v1.worker.gpu.sample.gumbel import gumbel_sample
|
||||
|
||||
DEVICE = "cuda"
|
||||
VOCAB_SIZE = 200_000
|
||||
NUM_SAMPLES = 500_000
|
||||
# Dominant token is exp(HEAD_LOG_GAP)x larger than the unit-count tail, so the
|
||||
# tail sits ~HEAD_LOG_GAP logits below the top.
|
||||
HEAD_LOG_GAP = 18.0
|
||||
# 10-sigma band: a correct sampler effectively never trips it.
|
||||
Z_TOLERANCE = 10.0
|
||||
|
||||
|
||||
def _make_heavy_tailed_counts(seed: int = 1234) -> torch.Tensor:
|
||||
"""Non-negative int64 counts of shape [VOCAB_SIZE]; target prob = counts/N."""
|
||||
gen = torch.Generator(device=DEVICE).manual_seed(seed)
|
||||
counts = torch.randint(
|
||||
1, 4, (VOCAB_SIZE,), generator=gen, dtype=torch.int64, device=DEVICE
|
||||
)
|
||||
counts[0] = round(math.exp(HEAD_LOG_GAP)) # dominant token
|
||||
return counts
|
||||
|
||||
|
||||
def _counts_to_logits(counts: torch.Tensor) -> torch.Tensor:
|
||||
# softmax(log(count)) == count / sum(count); count 0 -> logit -inf -> prob 0.
|
||||
return counts.double().log().to(torch.float32)
|
||||
|
||||
|
||||
def _sample(
|
||||
logits_1d: torch.Tensor,
|
||||
num_samples: int,
|
||||
*,
|
||||
use_fp64: bool = False,
|
||||
temperature: float = 1.0,
|
||||
) -> torch.Tensor:
|
||||
"""Sample `num_samples` tokens from one logit vector.
|
||||
|
||||
Fixed seed with a distinct `pos` per sample gives independent draws; the
|
||||
logits are broadcast with a 0-stride view to avoid materializing
|
||||
[num_samples, vocab_size].
|
||||
"""
|
||||
vocab_size = logits_1d.shape[0]
|
||||
logits = logits_1d.unsqueeze(0).expand(num_samples, vocab_size)
|
||||
idx_mapping = torch.zeros(num_samples, dtype=torch.int32, device=DEVICE)
|
||||
temp = torch.tensor([temperature], dtype=torch.float32, device=DEVICE)
|
||||
seed = torch.tensor([0xABCD], dtype=torch.int64, device=DEVICE)
|
||||
pos = torch.arange(num_samples, dtype=torch.int64, device=DEVICE)
|
||||
return gumbel_sample(
|
||||
logits,
|
||||
idx_mapping,
|
||||
temp,
|
||||
seed,
|
||||
pos,
|
||||
apply_temperature=True,
|
||||
use_fp64=use_fp64,
|
||||
)
|
||||
|
||||
|
||||
def _z_score(observed: int, expected: float, num_trials: int) -> float:
|
||||
p = expected / num_trials
|
||||
return (observed - expected) / math.sqrt(num_trials * p * (1 - p))
|
||||
|
||||
|
||||
def _sample_histogram(
|
||||
logits_1d: torch.Tensor, num_samples: int, *, chunk: int = 1_000_000
|
||||
) -> torch.Tensor:
|
||||
"""Histogram of `num_samples` draws, accumulated in chunks.
|
||||
|
||||
Chunking keeps the kernel's per-sample scratch ([chunk, num_blocks]) bounded
|
||||
so a large sample count does not blow up memory.
|
||||
"""
|
||||
vocab_size = logits_1d.shape[0]
|
||||
hist = torch.zeros(vocab_size, dtype=torch.float64, device=DEVICE)
|
||||
for start in range(0, num_samples, chunk):
|
||||
size = min(chunk, num_samples - start)
|
||||
logits = logits_1d.unsqueeze(0).expand(size, vocab_size)
|
||||
idx_mapping = torch.zeros(size, dtype=torch.int32, device=DEVICE)
|
||||
temp = torch.tensor([1.0], dtype=torch.float32, device=DEVICE)
|
||||
seed = torch.tensor([0xABCD], dtype=torch.int64, device=DEVICE)
|
||||
pos = torch.arange(start, start + size, dtype=torch.int64, device=DEVICE)
|
||||
out = gumbel_sample(
|
||||
logits, idx_mapping, temp, seed, pos, apply_temperature=True
|
||||
)
|
||||
hist += torch.bincount(out, minlength=vocab_size).double()
|
||||
return hist
|
||||
|
||||
|
||||
# ----------------------------- Accuracy ------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize("use_fp64", [False, True])
|
||||
def test_sampling_matches_target_distribution(use_fp64: bool):
|
||||
counts = _make_heavy_tailed_counts()
|
||||
total = counts.sum().item()
|
||||
logits = _counts_to_logits(counts)
|
||||
|
||||
sampled = _sample(logits, NUM_SAMPLES, use_fp64=use_fp64)
|
||||
assert sampled.min() >= 0 and sampled.max() < VOCAB_SIZE
|
||||
|
||||
# The dominant token (index 0) and the aggregate tail are the two
|
||||
# statistically resolvable bins (individual tail tokens are far below the
|
||||
# ~5/N detectability floor). The tail mass is small but well above noise,
|
||||
# and it lives beyond the fp32 Gumbel cap -- the regime sensitive to noise
|
||||
# precision -- so matching it is the meaningful check.
|
||||
tail_prob = (total - counts[0].item()) / total
|
||||
tail_count = (sampled != 0).sum().item()
|
||||
z = _z_score(tail_count, NUM_SAMPLES * tail_prob, NUM_SAMPLES)
|
||||
assert abs(z) < Z_TOLERANCE, (
|
||||
f"sampled tail mass {tail_count / NUM_SAMPLES:.3e} != target "
|
||||
f"{tail_prob:.3e} (z={z:.2f})"
|
||||
)
|
||||
|
||||
|
||||
def test_full_vocab_distribution_fidelity():
|
||||
"""The sampled distribution matches the target across the WHOLE vocab.
|
||||
|
||||
A near-flat count tensor makes every one of the 200K bins individually
|
||||
measurable. With ~20 samples/bin, a goodness-of-fit over all bins checks
|
||||
that no part of the vocab is over- or under-represented (the heavy-tailed
|
||||
test above only resolves head vs aggregate tail). Empirically the fp32
|
||||
sampler is as faithful here as torch.multinomial; the residual error is the
|
||||
multinomial sampling-noise floor, not the kernel.
|
||||
"""
|
||||
gen = torch.Generator(device=DEVICE).manual_seed(2024)
|
||||
counts = torch.randint(
|
||||
500, 1500, (VOCAB_SIZE,), generator=gen, dtype=torch.int64, device=DEVICE
|
||||
)
|
||||
total = counts.sum().item()
|
||||
logits = _counts_to_logits(counts)
|
||||
|
||||
num_samples = 4_000_000
|
||||
hist = _sample_histogram(logits, num_samples)
|
||||
|
||||
# Diversity: essentially every token must be reachable (no starved region).
|
||||
coverage = (hist > 0).sum().item() / VOCAB_SIZE
|
||||
assert coverage > 0.99, f"only {coverage:.4f} of the vocab was ever sampled"
|
||||
|
||||
# Goodness-of-fit across all bins (each has expected count >= ~10).
|
||||
expected = (counts.double() / total) * num_samples
|
||||
chi2 = (((hist - expected) ** 2) / expected).sum().item()
|
||||
df = VOCAB_SIZE - 1
|
||||
assert chi2 < df + 10 * math.sqrt(2 * df), f"chi2={chi2:.0f}, df={df}"
|
||||
|
||||
|
||||
# ----------------------------- Edge cases ----------------------------------
|
||||
|
||||
|
||||
def test_greedy_temperature_zero_returns_argmax():
|
||||
"""temperature == 0 skips Gumbel noise and returns the exact argmax."""
|
||||
torch.manual_seed(0)
|
||||
num_reqs = 128
|
||||
logits = torch.randn(num_reqs, VOCAB_SIZE, device=DEVICE, dtype=torch.float32)
|
||||
idx_mapping = torch.arange(num_reqs, dtype=torch.int32, device=DEVICE)
|
||||
temp = torch.zeros(num_reqs, dtype=torch.float32, device=DEVICE)
|
||||
seed = torch.arange(num_reqs, dtype=torch.int64, device=DEVICE)
|
||||
pos = torch.arange(num_reqs, dtype=torch.int64, device=DEVICE)
|
||||
|
||||
sampled = gumbel_sample(
|
||||
logits, idx_mapping, temp, seed, pos, apply_temperature=True
|
||||
)
|
||||
assert torch.equal(sampled, logits.argmax(dim=-1))
|
||||
|
||||
|
||||
def test_zero_count_tokens_are_never_sampled():
|
||||
"""Count 0 -> -inf logit -> probability 0; must never be selected."""
|
||||
counts = _make_heavy_tailed_counts(seed=7)
|
||||
zeroed = torch.arange(1, VOCAB_SIZE, 2, device=DEVICE) # odd indices (not head)
|
||||
counts[zeroed] = 0
|
||||
logits = _counts_to_logits(counts)
|
||||
|
||||
sampled = _sample(logits, NUM_SAMPLES)
|
||||
assert sampled.min() >= 0 and sampled.max() < VOCAB_SIZE
|
||||
assert not torch.isin(sampled, zeroed).any(), "sampled a zero-probability token"
|
||||
|
||||
|
||||
def test_single_nonzero_token_is_always_sampled():
|
||||
"""A lone finite logit must win every draw, regardless of its index."""
|
||||
counts = torch.zeros(VOCAB_SIZE, dtype=torch.int64, device=DEVICE)
|
||||
counts[123_456] = 1000
|
||||
logits = _counts_to_logits(counts)
|
||||
|
||||
sampled = _sample(logits, 10_000)
|
||||
assert (sampled == 123_456).all()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("vocab_size", [1, 999, 1024, 4097])
|
||||
def test_vocab_size_not_multiple_of_block(vocab_size: int):
|
||||
"""Per-block tail masking for non-block-aligned vocab; all bins measurable."""
|
||||
gen = torch.Generator(device=DEVICE).manual_seed(vocab_size)
|
||||
counts = torch.randint(
|
||||
20, 200, (vocab_size,), generator=gen, dtype=torch.int64, device=DEVICE
|
||||
)
|
||||
total = counts.sum().item()
|
||||
logits = _counts_to_logits(counts)
|
||||
num_samples = max(40 * vocab_size, 50_000)
|
||||
|
||||
sampled = _sample(logits, num_samples)
|
||||
assert sampled.min() >= 0 and sampled.max() < vocab_size
|
||||
|
||||
observed = torch.bincount(sampled, minlength=vocab_size).double()
|
||||
expected = (counts.double() / total) * num_samples
|
||||
chi2 = (((observed - expected) ** 2) / expected).sum().item()
|
||||
df = vocab_size - 1
|
||||
if df >= 1:
|
||||
assert chi2 < df + 10 * math.sqrt(2 * df), f"chi2={chi2:.1f}, df={df}"
|
||||
@@ -176,7 +176,7 @@ class MooncakeStoreConnector(KVConnectorBase_V1, SupportsHMA):
|
||||
self,
|
||||
request: Request,
|
||||
num_computed_tokens: int,
|
||||
) -> tuple[int | None, bool]:
|
||||
) -> tuple[int, bool]:
|
||||
assert self.connector_scheduler is not None
|
||||
return self.connector_scheduler.get_num_new_matched_tokens(
|
||||
request, num_computed_tokens
|
||||
|
||||
@@ -54,9 +54,9 @@ class MooncakeStoreScheduler:
|
||||
):
|
||||
assert vllm_config.kv_transfer_config is not None
|
||||
self.kv_role = vllm_config.kv_transfer_config.kv_role
|
||||
kvc_extra_config = vllm_config.kv_transfer_config.kv_connector_extra_config
|
||||
self.load_async = kvc_extra_config.get("load_async", True)
|
||||
self.lookup_async = kvc_extra_config.get("lookup_async", False)
|
||||
self.load_async = vllm_config.kv_transfer_config.kv_connector_extra_config.get(
|
||||
"load_async", True
|
||||
)
|
||||
self.client = LookupKeyClient(vllm_config)
|
||||
|
||||
# Align with the engine's own scheduler_block_size and hash_block_size.
|
||||
@@ -75,26 +75,14 @@ class MooncakeStoreScheduler:
|
||||
self,
|
||||
request: Request,
|
||||
num_computed_tokens: int,
|
||||
) -> tuple[int | None, bool]:
|
||||
"""Check for external KV cache hit.
|
||||
|
||||
Returns ``(None, False)`` when an async lookup is still in flight,
|
||||
signaling the scheduler to retry this request on a later step.
|
||||
"""
|
||||
) -> tuple[int, bool]:
|
||||
"""Check for external KV cache hit."""
|
||||
# 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:
|
||||
return 0, False
|
||||
|
||||
num_external_hit_tokens = self.client.lookup(
|
||||
request.request_id,
|
||||
token_len,
|
||||
request.block_hashes,
|
||||
non_block=self.lookup_async,
|
||||
)
|
||||
if num_external_hit_tokens is None:
|
||||
# Lookup not ready yet; scheduler will retry on a later step.
|
||||
return None, False
|
||||
num_external_hit_tokens = self.client.lookup(token_len, request.block_hashes)
|
||||
|
||||
if num_external_hit_tokens == request.num_tokens:
|
||||
# Leave a sub-block tail uncomputed for sampling, on a block
|
||||
@@ -170,7 +158,6 @@ class MooncakeStoreScheduler:
|
||||
force_skip_save = self.kv_role == "kv_consumer"
|
||||
|
||||
for finished_req_id in scheduler_output.finished_req_ids:
|
||||
self.client.discard(finished_req_id)
|
||||
self.load_specs.pop(finished_req_id, None)
|
||||
self._request_trackers.pop(finished_req_id, None)
|
||||
self._unfinished_requests.pop(finished_req_id, None)
|
||||
|
||||
@@ -19,7 +19,6 @@ import threading
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import Future, ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal, TypeVar
|
||||
|
||||
@@ -972,13 +971,7 @@ class MooncakeStoreWorker:
|
||||
else:
|
||||
self.num_kv_head = model_config.get_total_num_kv_heads()
|
||||
|
||||
if self.num_kv_head < self.tp_size and self.dcp_size <= 1:
|
||||
# Dedup: TP ranks holding the same KV heads stripe PUTs across
|
||||
# one shared key namespace. DCP splits the TP group, so with
|
||||
# DCP>1 those ranks have different `@dcpN` namespaces and
|
||||
# striping would leave keys unwritten (OBJECT_NOT_FOUND on
|
||||
# GET). PCP is outer to TP (pcp_rank is constant within a TP
|
||||
# group), so it needs no guard.
|
||||
if self.num_kv_head < self.tp_size:
|
||||
self.put_step = self.tp_size // self.num_kv_head
|
||||
self.head_or_tp_rank = self.tp_rank // self.put_step
|
||||
else:
|
||||
@@ -1567,13 +1560,7 @@ class LookupKeyClient:
|
||||
bind=False,
|
||||
)
|
||||
|
||||
# Async lookup support
|
||||
self.executor = ThreadPoolExecutor(
|
||||
max_workers=1, thread_name_prefix="MooncakeLookupClient"
|
||||
)
|
||||
self.futures: dict[str, Future[int]] = {}
|
||||
|
||||
def _lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
|
||||
def lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
|
||||
hash_strs = [h.hex() for h in block_hashes]
|
||||
hash_frames = self.encoder.encode(hash_strs)
|
||||
token_len_bytes = token_len.to_bytes(4, byteorder="big")
|
||||
@@ -1583,36 +1570,7 @@ class LookupKeyClient:
|
||||
result = int.from_bytes(resp, "big")
|
||||
return result
|
||||
|
||||
def lookup(
|
||||
self,
|
||||
req_id: str,
|
||||
token_len: int,
|
||||
block_hashes: list[BlockHash],
|
||||
non_block: bool = False,
|
||||
) -> int | None:
|
||||
"""If non_block is True, will return None until the result is ready,
|
||||
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))
|
||||
self.futures[req_id] = future
|
||||
if non_block and not future.done():
|
||||
return None
|
||||
try:
|
||||
return future.result()
|
||||
except Exception as e:
|
||||
logger.error("Async Mooncake lookup failed for %s: %s", req_id, e)
|
||||
return 0
|
||||
finally:
|
||||
del self.futures[req_id]
|
||||
|
||||
def discard(self, req_id: str) -> None:
|
||||
"""Drop any cached/in-flight lookup for ``req_id`` (e.g. on abort)."""
|
||||
future = self.futures.pop(req_id, None)
|
||||
if future is not None:
|
||||
future.cancel()
|
||||
|
||||
def _reset(self) -> bool:
|
||||
def reset(self) -> bool:
|
||||
"""Trigger ``store.remove_all(force=True)`` on worker rank 0.
|
||||
|
||||
Ordering assumption: caller MUST ensure no in-flight Mooncake
|
||||
@@ -1624,11 +1582,7 @@ class LookupKeyClient:
|
||||
resp = self.socket.recv()
|
||||
return bytes(resp) == RESP_OK
|
||||
|
||||
def reset(self) -> bool:
|
||||
return self.executor.submit(self._reset).result()
|
||||
|
||||
def close(self):
|
||||
self.executor.shutdown(wait=False, cancel_futures=True)
|
||||
self.socket.close(linger=0)
|
||||
|
||||
|
||||
|
||||
@@ -12,7 +12,6 @@ import uuid
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import jinja2
|
||||
from fastapi import Request
|
||||
|
||||
from vllm.engine.protocol import EngineClient
|
||||
@@ -43,7 +42,6 @@ from vllm.entrypoints.openai.engine.protocol import (
|
||||
JsonSchemaResponseFormat,
|
||||
ResponseFormat,
|
||||
StreamOptions,
|
||||
UsageInfo,
|
||||
)
|
||||
from vllm.entrypoints.openai.models.serving import OpenAIServingModels
|
||||
from vllm.entrypoints.serve.utils.api_utils import sanitize_message
|
||||
@@ -55,45 +53,6 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_cached_tokens(usage: UsageInfo | None) -> int | None:
|
||||
"""Extract cached token count from OpenAI UsageInfo."""
|
||||
if usage is None or usage.prompt_tokens_details is None:
|
||||
return None
|
||||
return usage.prompt_tokens_details.cached_tokens
|
||||
|
||||
|
||||
def _build_anthropic_usage(
|
||||
prompt_tokens: int,
|
||||
completion_tokens: int,
|
||||
usage: UsageInfo | None,
|
||||
) -> AnthropicUsage:
|
||||
"""Build an AnthropicUsage from OpenAI-style token counts.
|
||||
|
||||
Anthropic defines ``total_input == input_tokens + cache_read +
|
||||
cache_creation``. vLLM's ``prompt_tokens`` is the total, so
|
||||
``input_tokens = prompt_tokens - cached_tokens``.
|
||||
|
||||
OpenAI usage only exposes ``cached_tokens`` (hits); there is no
|
||||
cache-creation analog, so ``cache_creation_input_tokens`` is ``0``
|
||||
when cache info is present. When cache info is absent (e.g.
|
||||
``--enable-prompt-tokens-details`` off, or a streaming chunk that
|
||||
hasn't carried it yet), cache fields are left **unset** so
|
||||
``exclude_unset=True`` serialization omits them entirely.
|
||||
"""
|
||||
cached = _get_cached_tokens(usage)
|
||||
if cached is not None:
|
||||
return AnthropicUsage(
|
||||
input_tokens=prompt_tokens - cached,
|
||||
output_tokens=completion_tokens,
|
||||
cache_read_input_tokens=cached,
|
||||
cache_creation_input_tokens=0,
|
||||
)
|
||||
return AnthropicUsage(
|
||||
input_tokens=prompt_tokens,
|
||||
output_tokens=completion_tokens,
|
||||
)
|
||||
|
||||
|
||||
def wrap_data_with_event(data: str, event: str):
|
||||
return f"event: {event}\ndata: {data}\n\n"
|
||||
|
||||
@@ -140,36 +99,6 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
"length": "max_tokens",
|
||||
"tool_calls": "tool_use",
|
||||
}
|
||||
self._merge_inline_system = self._detect_merge_inline_system(chat_template)
|
||||
|
||||
@staticmethod
|
||||
def _detect_merge_inline_system(chat_template: str | None) -> bool:
|
||||
"""Auto-detect whether the chat template requires system-first ordering.
|
||||
|
||||
Renders a [system, user, system, user] conversation against the
|
||||
template; if it raises (e.g. Qwen's ``loop.first`` guard), the
|
||||
model needs inline system messages merged into the leading block.
|
||||
"""
|
||||
if not chat_template:
|
||||
return True
|
||||
try:
|
||||
env = jinja2.sandbox.ImmutableSandboxedEnvironment(
|
||||
trim_blocks=True,
|
||||
lstrip_blocks=True,
|
||||
extensions=[jinja2.ext.loopcontrols],
|
||||
)
|
||||
env.from_string(chat_template).render(
|
||||
messages=[
|
||||
{"role": "system", "content": "t"},
|
||||
{"role": "user", "content": "t"},
|
||||
{"role": "system", "content": "t"},
|
||||
{"role": "user", "content": "t"},
|
||||
],
|
||||
add_generation_prompt=False,
|
||||
)
|
||||
return False
|
||||
except jinja2.TemplateError:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _convert_image_source_to_url(source: dict[str, Any]) -> str:
|
||||
@@ -194,24 +123,13 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
|
||||
@classmethod
|
||||
def _convert_anthropic_to_openai_request(
|
||||
cls,
|
||||
anthropic_request: AnthropicMessagesRequest | AnthropicCountTokensRequest,
|
||||
*,
|
||||
merge_inline_system: bool = False,
|
||||
cls, anthropic_request: AnthropicMessagesRequest | AnthropicCountTokensRequest
|
||||
) -> ChatCompletionRequest:
|
||||
"""Convert Anthropic message format to OpenAI format"""
|
||||
openai_messages: list[dict[str, Any]] = []
|
||||
|
||||
cls._convert_system_message(
|
||||
anthropic_request,
|
||||
openai_messages,
|
||||
merge_inline_system=merge_inline_system,
|
||||
)
|
||||
cls._convert_messages(
|
||||
anthropic_request.messages,
|
||||
openai_messages,
|
||||
merge_inline_system=merge_inline_system,
|
||||
)
|
||||
cls._convert_system_message(anthropic_request, openai_messages)
|
||||
cls._convert_messages(anthropic_request.messages, openai_messages)
|
||||
req = cls._build_base_request(anthropic_request, openai_messages)
|
||||
cls._handle_streaming_options(req, anthropic_request)
|
||||
cls._handle_output_config(req, anthropic_request)
|
||||
@@ -224,8 +142,6 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
cls,
|
||||
anthropic_request: AnthropicMessagesRequest | AnthropicCountTokensRequest,
|
||||
openai_messages: list[dict[str, Any]],
|
||||
*,
|
||||
merge_inline_system: bool = False,
|
||||
) -> None:
|
||||
"""Convert Anthropic system message to OpenAI format"""
|
||||
system_parts: list[str] = []
|
||||
@@ -243,17 +159,6 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
continue
|
||||
system_parts.append(block.text)
|
||||
|
||||
# When the template requires system-first ordering, extract inline
|
||||
# system messages from the messages array and merge them into the
|
||||
# top-level block so the template doesn't reject them.
|
||||
if merge_inline_system:
|
||||
for msg in anthropic_request.messages:
|
||||
if msg.role != "system":
|
||||
continue
|
||||
text = cls._extract_system_text(msg)
|
||||
if text:
|
||||
system_parts.append(text)
|
||||
|
||||
if system_parts:
|
||||
openai_messages.append({"role": "system", "content": "".join(system_parts)})
|
||||
|
||||
@@ -275,11 +180,7 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
|
||||
@classmethod
|
||||
def _convert_messages(
|
||||
cls,
|
||||
messages: list,
|
||||
openai_messages: list[dict[str, Any]],
|
||||
*,
|
||||
merge_inline_system: bool = False,
|
||||
cls, messages: list, openai_messages: list[dict[str, Any]]
|
||||
) -> None:
|
||||
"""Convert Anthropic messages to OpenAI format"""
|
||||
for msg in messages:
|
||||
@@ -289,8 +190,6 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
# doesn't strip billing headers and may produce messages with
|
||||
# no "content" key.
|
||||
if msg.role == "system":
|
||||
if merge_inline_system:
|
||||
continue # already merged into top-level by _convert_system_message
|
||||
text = cls._extract_system_text(msg)
|
||||
if text:
|
||||
openai_messages.append({"role": "system", "content": text})
|
||||
@@ -598,10 +497,7 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
"""
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("Received messages request %s", request.model_dump_json())
|
||||
chat_req = self._convert_anthropic_to_openai_request(
|
||||
request,
|
||||
merge_inline_system=self._merge_inline_system,
|
||||
)
|
||||
chat_req = self._convert_anthropic_to_openai_request(request)
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("Convert to OpenAI request %s", chat_req.model_dump_json())
|
||||
generator = await self.create_chat_completion(chat_req, raw_request)
|
||||
@@ -622,10 +518,9 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
id=generator.id,
|
||||
content=[],
|
||||
model=generator.model,
|
||||
usage=_build_anthropic_usage(
|
||||
generator.usage.prompt_tokens,
|
||||
generator.usage.completion_tokens,
|
||||
generator.usage,
|
||||
usage=AnthropicUsage(
|
||||
input_tokens=generator.usage.prompt_tokens,
|
||||
output_tokens=generator.usage.completion_tokens,
|
||||
),
|
||||
kv_transfer_params=generator.kv_transfer_params,
|
||||
)
|
||||
@@ -806,12 +701,11 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
model=origin_chunk.model,
|
||||
stop_reason=None,
|
||||
stop_sequence=None,
|
||||
usage=_build_anthropic_usage(
|
||||
origin_chunk.usage.prompt_tokens
|
||||
usage=AnthropicUsage(
|
||||
input_tokens=origin_chunk.usage.prompt_tokens
|
||||
if origin_chunk.usage
|
||||
else 0,
|
||||
0,
|
||||
origin_chunk.usage,
|
||||
output_tokens=0,
|
||||
),
|
||||
),
|
||||
)
|
||||
@@ -830,14 +724,13 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
chunk = AnthropicStreamEvent(
|
||||
type="message_delta",
|
||||
delta=AnthropicDelta(stop_reason=stop_reason),
|
||||
usage=_build_anthropic_usage(
|
||||
origin_chunk.usage.prompt_tokens
|
||||
usage=AnthropicUsage(
|
||||
input_tokens=origin_chunk.usage.prompt_tokens
|
||||
if origin_chunk.usage
|
||||
else 0,
|
||||
origin_chunk.usage.completion_tokens
|
||||
output_tokens=origin_chunk.usage.completion_tokens
|
||||
if origin_chunk.usage
|
||||
else 0,
|
||||
origin_chunk.usage,
|
||||
),
|
||||
)
|
||||
data = chunk.model_dump_json(exclude_unset=True)
|
||||
@@ -1012,10 +905,7 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
raw_request: Request | None = None,
|
||||
) -> AnthropicCountTokensResponse | ErrorResponse:
|
||||
"""Implements Anthropic's messages.count_tokens endpoint."""
|
||||
chat_req = self._convert_anthropic_to_openai_request(
|
||||
request,
|
||||
merge_inline_system=self._merge_inline_system,
|
||||
)
|
||||
chat_req = self._convert_anthropic_to_openai_request(request)
|
||||
result = await self.render_chat_request(chat_req)
|
||||
if isinstance(result, ErrorResponse):
|
||||
return result
|
||||
|
||||
@@ -427,12 +427,6 @@ async def validation_exception_handler(req: Request, exc: RequestValidationError
|
||||
param = ctx_error.parameter
|
||||
break
|
||||
|
||||
if param is None and errors:
|
||||
first_error = errors[0]
|
||||
loc = first_error.get("loc") if isinstance(first_error, dict) else None
|
||||
if loc:
|
||||
param = ".".join(str(part) for part in loc)
|
||||
|
||||
exc_str = str(exc)
|
||||
errors_str = str(errors)
|
||||
|
||||
|
||||
@@ -446,17 +446,11 @@ class PunicaWrapperGPU(PunicaWrapperBase):
|
||||
_,
|
||||
_,
|
||||
lora_ids,
|
||||
no_lora_flag,
|
||||
_,
|
||||
num_active_loras,
|
||||
) = self.token_mapping_meta.meta_args(
|
||||
x.size(0), self.lora_config.specialize_active_lora
|
||||
)
|
||||
|
||||
assert no_lora_flag.numel() == 1
|
||||
if no_lora_flag.item():
|
||||
# None of the inputs require LoRA.
|
||||
return
|
||||
|
||||
if token_lora_mapping is None:
|
||||
token_lora_mapping = token_lora_mapping_meta
|
||||
fused_moe_lora(
|
||||
|
||||
@@ -349,7 +349,6 @@ class MLAAttention(nn.Module, AttentionLayerBase):
|
||||
attn_backend: type[AttentionBackend] | None = None,
|
||||
use_sparse: bool = False,
|
||||
indexer: object | None = None,
|
||||
topk_indices_buffer: torch.Tensor | None = None,
|
||||
**extra_impl_args,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -438,11 +437,6 @@ class MLAAttention(nn.Module, AttentionLayerBase):
|
||||
)
|
||||
cache_config.enable_prefix_caching = False
|
||||
|
||||
# Sparse MLA reads top-k indices from a shared buffer. Pass it
|
||||
# explicitly so backbone "skip" layers (indexer=None) still find it.
|
||||
if use_sparse:
|
||||
extra_impl_args["topk_indices_buffer"] = topk_indices_buffer
|
||||
|
||||
impl_cls = cast(type[MLAAttentionImpl], self.attn_backend.get_impl_cls())
|
||||
self.impl = impl_cls( # type: ignore[assignment] # impl_cls always returns an MLAAttentionImpl subclass
|
||||
num_heads=self.num_heads,
|
||||
|
||||
@@ -59,10 +59,3 @@ class MoELoRAContext:
|
||||
# None means no dispatch happened (non-EP path), in which case callers
|
||||
# fall back to punica_wrapper.token_mapping_meta.
|
||||
local_token_lora_mapping: torch.Tensor | None = None
|
||||
|
||||
# Original unquantized hidden states, stashed by the modular kernel
|
||||
# before the prepare step potentially quantizes them. Used by
|
||||
# apply_w13_lora so the LoRA kernel sees correct-magnitude activations
|
||||
# instead of raw quantized values that are missing the activation scale.
|
||||
# Set per forward pass; None until the modular kernel writes it.
|
||||
original_hidden_states: torch.Tensor | None = None
|
||||
|
||||
@@ -77,16 +77,6 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
|
||||
def activation_format() -> mk.FusedMoEActivationFormat:
|
||||
return mk.FusedMoEActivationFormat.Standard
|
||||
|
||||
@property
|
||||
def expects_unquantized_inputs(self) -> bool:
|
||||
# Defer activation quantization to apply() only when LoRA is active AND
|
||||
# tokens are dispatched across ranks (DP+EP all2all).
|
||||
return (
|
||||
self._lora_context is not None
|
||||
and self.quant_dtype is not None
|
||||
and self.moe_config.moe_parallel_config.use_all2all_kernels
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_current_device() -> bool:
|
||||
return current_platform.is_cuda_alike() or current_platform.is_xpu()
|
||||
@@ -233,25 +223,6 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
|
||||
torch.float8_e4m3fnuz,
|
||||
]
|
||||
|
||||
# We declared expects_unquantized_inputs (LoRA + DP/EP all2all), so the
|
||||
# prepare step deferred activation quantization to this kernel:
|
||||
# `hidden_states` arrives unquantized. Keep the unquantized tensor for
|
||||
# the LoRA shrink input and quantize a copy here for the base GEMM
|
||||
# (mirrors what the prepare step would have done, but after the
|
||||
# all-gather so the layout matches the gathered topk_ids / token map).
|
||||
lora_unquantized_hidden_states: torch.Tensor | None = None
|
||||
if self.expects_unquantized_inputs:
|
||||
assert a1q_scale is None
|
||||
lora_unquantized_hidden_states = hidden_states
|
||||
hidden_states, a1q_scale = moe_kernel_quantize_input(
|
||||
hidden_states,
|
||||
self.a1_scale,
|
||||
self.quant_dtype,
|
||||
self.per_act_token_quant,
|
||||
self.block_shape,
|
||||
quantization_emulation=self.quantization_emulation,
|
||||
)
|
||||
|
||||
E, num_tokens, N, K, top_k_num = self.moe_problem_size(
|
||||
hidden_states, w1, w2, topk_ids
|
||||
)
|
||||
@@ -309,28 +280,12 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
|
||||
# GEMM on the default stream and the LoRA fast-path on aux_stream;
|
||||
# the LoRA writes its delta into a fresh zero buffer (add_inputs=
|
||||
# False) and we sum it into intermediate_cache1 after both finish.
|
||||
#
|
||||
# The LoRA shrink kernel needs unquantized, gathered-layout
|
||||
# activations. When activation quant was deferred to this kernel
|
||||
# (expects_unquantized_inputs), the input we quantized above is exactly
|
||||
# that, so use it directly. Otherwise fall back to the context stash
|
||||
# (e.g. weight-only quant), guarding on a row-count match so a
|
||||
# DP-gathered layout never indexes a local stash out of bounds.
|
||||
|
||||
sorted_token_ids_lora = None
|
||||
expert_ids_lora = None
|
||||
num_tokens_post_padded_lora = None
|
||||
token_lora_mapping = None
|
||||
lora_context = self._lora_context
|
||||
if lora_unquantized_hidden_states is not None:
|
||||
lora_x = lora_unquantized_hidden_states
|
||||
elif (
|
||||
lora_context is not None
|
||||
and lora_context.original_hidden_states is not None
|
||||
and lora_context.original_hidden_states.shape[0] == hidden_states.shape[0]
|
||||
):
|
||||
lora_x = lora_context.original_hidden_states
|
||||
else:
|
||||
lora_x = hidden_states
|
||||
|
||||
def _base_w13_fn():
|
||||
invoke_fused_moe_triton_kernel(
|
||||
@@ -367,7 +322,7 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
|
||||
return self.apply_w13_lora(
|
||||
lora_context,
|
||||
y=lora_delta_w13,
|
||||
x=lora_x,
|
||||
x=hidden_states,
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
expert_map=expert_map,
|
||||
@@ -404,7 +359,7 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
|
||||
) = self.apply_w13_lora(
|
||||
lora_context,
|
||||
y=intermediate_cache1,
|
||||
x=lora_x,
|
||||
x=hidden_states,
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
expert_map=expert_map,
|
||||
|
||||
@@ -1407,13 +1407,6 @@ class FusedMoEKernelModularImpl:
|
||||
apply_router_weight_on_input,
|
||||
)
|
||||
|
||||
# Stash the original unquantized hidden states on the LoRA context
|
||||
# so apply_w13_lora sees correct-magnitude activations instead of
|
||||
# the potentially quantized values produced by _prepare().
|
||||
lora_ctx = getattr(self.fused_experts, "_lora_context", None)
|
||||
if lora_ctx is not None:
|
||||
lora_ctx.original_hidden_states = hidden_states
|
||||
|
||||
fused_out = self._fused_experts(
|
||||
in_dtype=hidden_states.dtype,
|
||||
a1q=a1q,
|
||||
@@ -1431,9 +1424,6 @@ class FusedMoEKernelModularImpl:
|
||||
output_alias=output,
|
||||
)
|
||||
|
||||
if lora_ctx is not None:
|
||||
lora_ctx.original_hidden_states = None
|
||||
|
||||
return self._finalize(
|
||||
output,
|
||||
fused_out,
|
||||
|
||||
@@ -112,7 +112,6 @@ class MultiHeadLatentAttentionWrapper(PluggableLayer):
|
||||
kv_b_proj=self.kv_b_proj,
|
||||
use_sparse=self.is_sparse,
|
||||
indexer=self.indexer,
|
||||
topk_indices_buffer=mla_modules.topk_indices_buffer,
|
||||
)
|
||||
|
||||
self.prefix = prefix
|
||||
|
||||
@@ -119,12 +119,8 @@ class DeepSeekMultiTokenPredictorLayer(nn.Module):
|
||||
hidden_states=hidden_states,
|
||||
residual=None,
|
||||
)
|
||||
hidden_states = residual + hidden_states # pre-final-norm (logits hidden)
|
||||
# Recycle the post-final-norm hidden into the next draft step.
|
||||
# compute_logits applies shared_head (== final norm) to the pre-norm
|
||||
# element, so logits and the recycle each get exactly one final-norm.
|
||||
# Matches SGLang's deepseek_nextn.
|
||||
return hidden_states, self.shared_head(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DeepSeekMultiTokenPredictor(nn.Module):
|
||||
|
||||
@@ -998,29 +998,8 @@ class DeepseekV2MLAAttention(nn.Module):
|
||||
|
||||
self.is_v32 = hasattr(config, "index_topk")
|
||||
|
||||
# IndexCache config
|
||||
# Refer: https://arxiv.org/abs/2603.12201 for more details.
|
||||
_skip_topk = False
|
||||
_index_topk_freq = getattr(config, "index_topk_freq", 1)
|
||||
_index_topk_pattern = getattr(config, "index_topk_pattern", None)
|
||||
_index_skip_topk_offset = getattr(config, "index_skip_topk_offset", 2)
|
||||
layer_id = extract_layer_index(prefix)
|
||||
|
||||
if _index_topk_pattern is None:
|
||||
_skip_topk = (
|
||||
max(layer_id - _index_skip_topk_offset + 1, 0) % _index_topk_freq != 0
|
||||
)
|
||||
elif 0 <= layer_id < len(_index_topk_pattern):
|
||||
_skip_topk = _index_topk_pattern[layer_id] == "S"
|
||||
|
||||
# The skip pattern only governs backbone layers. MTP/nextn layers
|
||||
# (layer_id >= num_hidden_layers) always build a full indexer: they
|
||||
# compute indices at draft step 0 and toggle at runtime via
|
||||
# set_skip_topk (index_share_for_mtp_iteration).
|
||||
_num_hidden_layers = getattr(config, "num_hidden_layers", None)
|
||||
is_mtp_layer = _num_hidden_layers is not None and layer_id >= _num_hidden_layers
|
||||
|
||||
if self.is_v32 and (not _skip_topk or is_mtp_layer):
|
||||
if self.is_v32:
|
||||
self.indexer_rope_emb = get_rope(
|
||||
qk_rope_head_dim,
|
||||
max_position=max_position_embeddings,
|
||||
@@ -1038,6 +1017,22 @@ class DeepseekV2MLAAttention(nn.Module):
|
||||
f"{prefix}.indexer",
|
||||
is_inplace_rope=self.indexer_rope_emb.enabled(),
|
||||
)
|
||||
|
||||
# IndexCache config
|
||||
# Refer: https://arxiv.org/abs/2603.12201 for more details.
|
||||
_index_topk_freq = getattr(config, "index_topk_freq", 1)
|
||||
_index_topk_pattern = getattr(config, "index_topk_pattern", None)
|
||||
_index_skip_topk_offset = getattr(config, "index_skip_topk_offset", 2)
|
||||
layer_id = extract_layer_index(prefix)
|
||||
|
||||
if _index_topk_pattern is None:
|
||||
_skip_topk = (
|
||||
max(layer_id - _index_skip_topk_offset + 1, 0) % _index_topk_freq
|
||||
!= 0
|
||||
)
|
||||
elif 0 <= layer_id < len(_index_topk_pattern):
|
||||
_skip_topk = _index_topk_pattern[layer_id] == "S"
|
||||
|
||||
else:
|
||||
self.indexer_rope_emb = None
|
||||
self.indexer = None
|
||||
|
||||
@@ -64,7 +64,6 @@ from vllm.models.deepseek_v4.nvidia.flashinfer_sparse import (
|
||||
from vllm.models.deepseek_v4.nvidia.flashmla import DeepseekV4FlashMLAAttention
|
||||
from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import prepare_megamoe_inputs
|
||||
from vllm.sequence import IntermediateTensors
|
||||
from vllm.utils.math_utils import cdiv
|
||||
from vllm.v1.attention.backends.registry import AttentionBackendEnum
|
||||
|
||||
|
||||
@@ -86,15 +85,6 @@ class DeepseekV4MLP(nn.Module):
|
||||
# across the ranks within the tp_group. In this case the weights are
|
||||
# replicated and no collective ops are needed.
|
||||
# Otherwise we use standard TP with an allreduce at the end.
|
||||
#
|
||||
# Block-FP8 shards in whole 128-blocks; cdiv rounds the per-rank block
|
||||
# count up so the linear's even TP split stays block-aligned, with the
|
||||
# trailing ranks zero-filled by load_weights.
|
||||
block_size = getattr(quant_config, "weight_block_size", None)
|
||||
if block_size is not None and not is_sequence_parallel:
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
n_local = cdiv(intermediate_size // block_size[0], tp_size)
|
||||
intermediate_size = n_local * block_size[0] * tp_size
|
||||
self.gate_up_proj = MergedColumnParallelLinear(
|
||||
hidden_size,
|
||||
[intermediate_size] * 2,
|
||||
@@ -902,8 +892,6 @@ class DeepseekV4Model(nn.Module):
|
||||
config = vllm_config.model_config.hf_config
|
||||
quant_config = vllm_config.quant_config
|
||||
self.config = config
|
||||
self.quant_config = quant_config
|
||||
self.parallel_config = vllm_config.parallel_config
|
||||
self.use_mega_moe = (
|
||||
vllm_config.kernel_config.moe_backend == "deep_gemm_mega_moe"
|
||||
)
|
||||
@@ -1092,17 +1080,7 @@ class DeepseekV4Model(nn.Module):
|
||||
# Pre-compute expert mapping ONCE.
|
||||
expert_mapping = self.get_expert_mapping()
|
||||
|
||||
# Block-FP8 shared experts: pad the intermediate up to the TP-uniform
|
||||
# block count so the standard loaders below slice it evenly (trailing
|
||||
# ranks land on the zero pad). SP / unquantized ones need no padding.
|
||||
pad_shared_expert = (
|
||||
getattr(self.quant_config, "weight_block_size", None) is not None
|
||||
and not self.parallel_config.use_sequence_parallel_moe
|
||||
)
|
||||
|
||||
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)
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
# Skip non-stacked layers and experts (experts handled below).
|
||||
if ".experts." in name:
|
||||
@@ -1177,28 +1155,6 @@ class DeepseekV4Model(nn.Module):
|
||||
|
||||
return loaded_params
|
||||
|
||||
def _pad_shared_expert_weight(
|
||||
self, 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)
|
||||
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.
|
||||
step = 1 if name.endswith("weight_scale_inv") else block_size[0]
|
||||
dim = 1 if ".down_proj." in name else 0
|
||||
mult = get_tensor_model_parallel_world_size() * step
|
||||
pad = cdiv(loaded_weight.shape[dim], mult) * mult - loaded_weight.shape[dim]
|
||||
if pad == 0:
|
||||
return loaded_weight
|
||||
pad_shape = list(loaded_weight.shape)
|
||||
pad_shape[dim] = pad
|
||||
return torch.cat([loaded_weight, loaded_weight.new_zeros(pad_shape)], dim=dim)
|
||||
|
||||
def get_expert_mapping(self) -> list[tuple[str, str, int, str]]:
|
||||
first_layer = next(iter(islice(self.layers, self.start_layer, self.end_layer)))
|
||||
if first_layer.ffn.use_mega_moe:
|
||||
|
||||
@@ -672,7 +672,7 @@ class ParserEngine(Parser):
|
||||
if len(tool_call_deltas) > 1:
|
||||
tool_call_deltas = self._coalesce_tool_call_deltas(tool_call_deltas)
|
||||
|
||||
if self._deferred_content and (not seen_tool_event or not tool_call_deltas):
|
||||
if self._deferred_content and not seen_tool_event:
|
||||
content_parts.insert(0, self._deferred_content)
|
||||
self._deferred_content = ""
|
||||
|
||||
|
||||
@@ -375,10 +375,6 @@ def gemma4_config() -> ParserEngineConfig:
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END, EventType.TOOL_CALL_START),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "CALL_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(),
|
||||
|
||||
@@ -125,10 +125,6 @@ def qwen3_config(thinking: bool = True) -> ParserEngineConfig:
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "FUNC_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(),
|
||||
|
||||
@@ -20,6 +20,7 @@ except ImportError as e:
|
||||
) from e
|
||||
|
||||
|
||||
from vllm.entrypoints.mcp.tool_server import ToolServer
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
@@ -480,6 +481,15 @@ class BaseCohereCommandReasoningParser(ReasoningParser):
|
||||
def is_reasoning_end(self, input_ids: Sequence[int]) -> bool:
|
||||
return any(tid == self.end_token_id for tid in reversed(input_ids))
|
||||
|
||||
def prepare_structured_tag(
|
||||
self, original_tag: str | None, tool_server: ToolServer | None
|
||||
) -> str | None:
|
||||
# Responses API replaces ``structural_tag`` via the reasoning parser.
|
||||
# Default ``ReasoningParser.prepare_structured_tag`` returns None, which
|
||||
# would clear a Cohere tag produced in ``adjust_request`` and break
|
||||
# ``StructuredOutputsParams`` validation. Preserve the existing tag.
|
||||
return original_tag
|
||||
|
||||
def adjust_request(
|
||||
self, request: ChatCompletionRequest | ResponsesRequest
|
||||
) -> ChatCompletionRequest | ResponsesRequest:
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING
|
||||
|
||||
from transformers import PreTrainedTokenizerBase
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning import ReasoningParser
|
||||
from vllm.reasoning.deepseek_r1_reasoning_parser import DeepSeekR1ReasoningParser
|
||||
|
||||
@@ -16,6 +17,8 @@ if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class DeepSeekV3ReasoningParser(ReasoningParser):
|
||||
"""
|
||||
|
||||
@@ -7,12 +7,15 @@ from typing import TYPE_CHECKING
|
||||
from transformers import PreTrainedTokenizerBase
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning.basic_parsers import BaseThinkingReasoningParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Ernie45ReasoningParser(BaseThinkingReasoningParser):
|
||||
"""
|
||||
|
||||
@@ -8,12 +8,15 @@ import regex as re
|
||||
from transformers import PreTrainedTokenizerBase
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning import ReasoningParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class GraniteReasoningParser(ReasoningParser):
|
||||
"""
|
||||
|
||||
@@ -8,12 +8,15 @@ import regex as re
|
||||
from transformers import PreTrainedTokenizerBase
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning import ReasoningParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class HunyuanA13BReasoningParser(ReasoningParser):
|
||||
"""
|
||||
|
||||
@@ -7,12 +7,15 @@ from typing import TYPE_CHECKING
|
||||
from transformers import PreTrainedTokenizerBase
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning import ReasoningParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class IdentityReasoningParser(ReasoningParser):
|
||||
"""
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING
|
||||
from vllm.entrypoints.openai.engine.protocol import (
|
||||
DeltaMessage,
|
||||
)
|
||||
from vllm.logger import init_logger
|
||||
from vllm.parser.engine.registered_adapters import MinimaxM2ParserReasoningAdapter
|
||||
from vllm.reasoning.abs_reasoning_parsers import ReasoningParser
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
@@ -15,6 +16,8 @@ if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MiniMaxM2ReasoningParser(MinimaxM2ParserReasoningAdapter): # type: ignore[valid-type, misc]
|
||||
"""
|
||||
|
||||
@@ -5,6 +5,7 @@ from collections.abc import Iterable, Sequence
|
||||
from functools import cached_property
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning import ReasoningParser
|
||||
from vllm.reasoning.basic_parsers import BaseThinkingReasoningParser
|
||||
from vllm.tokenizers.mistral import MistralTokenizer
|
||||
@@ -13,6 +14,8 @@ if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MistralReasoningParser(BaseThinkingReasoningParser):
|
||||
"""
|
||||
|
||||
@@ -9,6 +9,7 @@ from typing import TYPE_CHECKING
|
||||
import regex as re
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning import ReasoningParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -16,6 +17,8 @@ if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Olmo3ReasoningState(enum.Enum):
|
||||
REASONING = 1
|
||||
|
||||
@@ -9,12 +9,15 @@ import regex as re
|
||||
from transformers import PreTrainedTokenizerBase
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.logger import init_logger
|
||||
from vllm.reasoning import ReasoningParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Step3ReasoningParser(ReasoningParser):
|
||||
"""
|
||||
|
||||
@@ -271,7 +271,7 @@ class FlashInferMLASparseImpl(SparseMLAAttentionImpl[FlashInferMLASparseMetadata
|
||||
attn_type: str,
|
||||
kv_sharing_target_layer_name: str | None,
|
||||
# MLA Specific Arguments
|
||||
topk_indices_buffer: torch.Tensor | None = None,
|
||||
topk_indice_buffer: torch.Tensor | None = None,
|
||||
indexer: "Indexer | None" = None,
|
||||
**mla_args,
|
||||
) -> None:
|
||||
@@ -301,12 +301,8 @@ class FlashInferMLASparseImpl(SparseMLAAttentionImpl[FlashInferMLASparseMetadata
|
||||
self.qk_nope_head_dim: int = mla_args["qk_nope_head_dim"]
|
||||
self.qk_rope_head_dim: int = mla_args["qk_rope_head_dim"]
|
||||
|
||||
# The indexer carries the shared buffer for normal layers and tests;
|
||||
# the explicitly-passed buffer covers backbone skip layers, whose
|
||||
# indexer is not constructed (see deepseek_v2.py).
|
||||
self.topk_indices_buffer: torch.Tensor | None = (
|
||||
indexer.topk_indices_buffer if indexer is not None else topk_indices_buffer
|
||||
)
|
||||
assert indexer is not None, "Indexer required for sparse MLA"
|
||||
self.topk_indices_buffer: torch.Tensor | None = indexer.topk_indices_buffer
|
||||
|
||||
self._workspace_buffer: torch.Tensor | None = None
|
||||
self.bmm1_scale: float | None = None
|
||||
|
||||
@@ -568,12 +568,8 @@ class FlashMLASparseImpl(SparseMLAAttentionImpl[FlashMLASparseMetadata]):
|
||||
self.kv_cache_dtype = kv_cache_dtype
|
||||
self.kv_lora_rank: int = mla_args["kv_lora_rank"]
|
||||
self.softmax_scale = scale
|
||||
# The indexer carries the shared buffer for normal layers and tests;
|
||||
# the explicitly-passed buffer covers backbone skip layers, whose
|
||||
# indexer is not constructed (see deepseek_v2.py).
|
||||
self.topk_indices_buffer: torch.Tensor | None = (
|
||||
indexer.topk_indices_buffer if indexer is not None else topk_indices_buffer
|
||||
)
|
||||
assert indexer is not None
|
||||
self.topk_indices_buffer: torch.Tensor | None = indexer.topk_indices_buffer
|
||||
# Prefill BF16 kernel requires 64 on Hopper, 128 on Blackwell
|
||||
self.prefill_padding = (
|
||||
128 if current_platform.is_device_capability_family(100) else 64
|
||||
|
||||
@@ -629,7 +629,7 @@ class ROCMAiterMLASparseImpl(SparseMLAAttentionImpl[ROCMAiterMLASparseMetadata])
|
||||
attn_type: str,
|
||||
kv_sharing_target_layer_name: str | None,
|
||||
# MLA Specific Arguments
|
||||
topk_indices_buffer: torch.Tensor | None = None,
|
||||
topk_indice_buffer: torch.Tensor | None = None,
|
||||
indexer: "Indexer | None" = None,
|
||||
**mla_args,
|
||||
) -> None:
|
||||
@@ -642,12 +642,8 @@ class ROCMAiterMLASparseImpl(SparseMLAAttentionImpl[ROCMAiterMLASparseMetadata])
|
||||
self.kv_cache_dtype = kv_cache_dtype
|
||||
self.kv_lora_rank: int = mla_args["kv_lora_rank"]
|
||||
self.softmax_scale = scale
|
||||
# The indexer carries the shared buffer for normal layers and tests;
|
||||
# the explicitly-passed buffer covers backbone skip layers, whose
|
||||
# indexer is not constructed (see deepseek_v2.py).
|
||||
self.topk_indices_buffer: torch.Tensor | None = (
|
||||
indexer.topk_indices_buffer if indexer is not None else topk_indices_buffer
|
||||
)
|
||||
assert indexer is not None
|
||||
self.topk_indices_buffer: torch.Tensor | None = indexer.topk_indices_buffer
|
||||
|
||||
vllm_config = get_current_vllm_config()
|
||||
max_tokens = vllm_config.scheduler_config.max_num_batched_tokens
|
||||
|
||||
@@ -184,7 +184,7 @@ class XPUMLASparseImpl(SparseMLAAttentionImpl[XPUMLASparseMetadata]):
|
||||
attn_type: str,
|
||||
kv_sharing_target_layer_name: str | None,
|
||||
# MLA Specific Arguments
|
||||
topk_indices_buffer: torch.Tensor | None = None,
|
||||
topk_indice_buffer: torch.Tensor | None = None,
|
||||
indexer: Optional["Indexer"] = None,
|
||||
**mla_args,
|
||||
) -> None:
|
||||
@@ -195,12 +195,8 @@ class XPUMLASparseImpl(SparseMLAAttentionImpl[XPUMLASparseMetadata]):
|
||||
self.kv_cache_dtype = kv_cache_dtype
|
||||
self.kv_lora_rank: int = mla_args["kv_lora_rank"]
|
||||
self.softmax_scale = scale
|
||||
# The indexer carries the shared buffer for normal layers and tests;
|
||||
# the explicitly-passed buffer covers backbone skip layers, whose
|
||||
# indexer is not constructed (see deepseek_v2.py).
|
||||
self.topk_indices_buffer: torch.Tensor | None = (
|
||||
indexer.topk_indices_buffer if indexer is not None else topk_indices_buffer
|
||||
)
|
||||
assert indexer is not None
|
||||
self.topk_indices_buffer: torch.Tensor | None = indexer.topk_indices_buffer
|
||||
|
||||
def _forward_bf16_kv(
|
||||
self,
|
||||
|
||||
@@ -918,13 +918,6 @@ class SpecDecodeBaseProposer:
|
||||
return per_group_attn_metadata, per_layer_attn_metadata
|
||||
|
||||
def model_returns_tuple(self) -> bool:
|
||||
if self.method == "mtp":
|
||||
# DeepSeek-family MTP (deepseek_mtp.py) recycles the post-final-
|
||||
# norm hidden, so its forward returns (logit_hidden,
|
||||
# recycle_hidden). Other MTP families return a single tensor.
|
||||
return "DeepSeekMTPModel" in (
|
||||
self.draft_model_config.hf_config.architectures or []
|
||||
)
|
||||
return self.method not in ("mtp", "draft_model", "dflash")
|
||||
|
||||
def prepare_next_token_ids_cpu(
|
||||
|
||||
@@ -2,16 +2,18 @@
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import torch
|
||||
|
||||
from vllm.triton_utils import HAS_TRITON, tl, tldevice, triton
|
||||
from vllm.triton_utils import HAS_TRITON, tl, triton
|
||||
|
||||
# Smallest positive value produced by Triton's fp32 `tl.rand`. Used to clamp
|
||||
# zero draws before the flipped Gumbel transform below.
|
||||
# Smallest positive normal fp32 value. Used to clamp the uniform draw so that
|
||||
# `log(u)` cannot produce -inf (and thus `-log(-log(u))` stays finite).
|
||||
#
|
||||
# Triton requires globals accessed from `@triton.jit` functions to be wrapped
|
||||
# in `tl.constexpr(...)`. We can only do that when Triton is actually
|
||||
# available — on the CPU worker path `tl` is a placeholder whose `constexpr`
|
||||
# attribute is `None`, and `tl.constexpr(...)` would crash at import time.
|
||||
_TL_RAND_MIN = tl.constexpr(4.6566127342e-10) if HAS_TRITON else 4.6566127342e-10
|
||||
_FP32_TINY = (
|
||||
tl.constexpr(float.fromhex("0x1p-126")) if HAS_TRITON else float.fromhex("0x1p-126")
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
@@ -129,17 +131,10 @@ def gumbel_block_argmax(
|
||||
|
||||
if USE_FP64:
|
||||
u = tl_rand64(gumbel_seed, block, includes_zero=False)
|
||||
gumbel_noise = -tl.log(-tl.log(u))
|
||||
else:
|
||||
u = tl.rand(gumbel_seed, block)
|
||||
u = tl.maximum(u, _TL_RAND_MIN)
|
||||
# Draw the large-noise tail (which decides the argmax winner) from u -> 0,
|
||||
# where fp32 has fine resolution, instead of u -> 1, where fp32 spacing is
|
||||
# ~2**-24. The naive `-log(-log(u))` puts the winning tail at u -> 1,
|
||||
# hard-capping the noise at ~16.6 and coarsely quantizing it; using
|
||||
# `log1p(-u)` == `log(1 - u)` keeps the tail in the well-resolved region.
|
||||
# Note `1 - u` would lose precision for small u, so `log1p` is required.
|
||||
gumbel_noise = -tl.log(-tldevice.log1p(-u))
|
||||
u = tl.maximum(u, _FP32_TINY)
|
||||
gumbel_noise = -tl.log(-tl.log(u))
|
||||
|
||||
# Apply gumbel noise.
|
||||
logits = tl.where(mask, logits + gumbel_noise, float("-inf"))
|
||||
|
||||
Reference in New Issue
Block a user