forked from Karylab-cklius/vllm
Compare commits
43
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
95dcefaaa5 | ||
|
|
536047755e | ||
|
|
1907d3854a | ||
|
|
ea9ddf59fc | ||
|
|
8cf7c4d8ad | ||
|
|
8e9d70fdd5 | ||
|
|
364ee36af1 | ||
|
|
06fae69114 | ||
|
|
14f8660a18 | ||
|
|
aed541def4 | ||
|
|
2bc20e8aba | ||
|
|
8cc242335d | ||
|
|
ba22cb6765 | ||
|
|
81bcced482 | ||
|
|
fb42e5219e | ||
|
|
0feca7ffa8 | ||
|
|
97b5ce5c39 | ||
|
|
4236514098 | ||
|
|
e45c8a9f4b | ||
|
|
b153dd3f28 | ||
|
|
930f8dc0a1 | ||
|
|
a16dbd5b85 | ||
|
|
bec232a914 | ||
|
|
b5c9e1ac33 | ||
|
|
ae2c4f3db7 | ||
|
|
fca432e60a | ||
|
|
af1ee8c475 | ||
|
|
5b4cb69523 | ||
|
|
9fc0c08026 | ||
|
|
f2b5fabb23 | ||
|
|
b8cb75b149 | ||
|
|
43916891b2 | ||
|
|
cda05ee8c4 | ||
|
|
77654d080c | ||
|
|
75698e60b3 | ||
|
|
8632c884dc | ||
|
|
c3734e8334 | ||
|
|
53f7553f09 | ||
|
|
4eb227992a | ||
|
|
ebcf511ec3 | ||
|
|
8fc1b2d046 | ||
|
|
5316638a5e | ||
|
|
61ab70ec3b |
@@ -23,4 +23,5 @@ steps:
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
|
||||
pytest -v -s basic_correctness/test_cpu_offload.py &&
|
||||
pytest -v -s basic_correctness/test_mem.py::test_end_to_end'
|
||||
|
||||
@@ -128,10 +128,10 @@ steps:
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
|
||||
(pytest -v -s lora/test_mixtral.py --deselect="tests/lora/test_mixtral.py::test_mixtral_lora[4]" || true) &&
|
||||
pytest -v -s lora/test_quant_model.py --deselect="tests/lora/test_quant_model.py::test_quant_model_lora[model0]" --deselect="tests/lora/test_quant_model.py::test_quant_model_lora[model1]" --deselect="tests/lora/test_quant_model.py::test_quant_model_tp_equality[model0]" &&
|
||||
pytest -v -s lora/test_transformers_model.py &&
|
||||
pytest -v -s lora/test_chatglm3_tp.py &&
|
||||
pytest -v -s lora/test_llama_tp.py::test_llama_lora &&
|
||||
pytest -s -v lora/test_minicpmv_tp.py'
|
||||
|
||||
- label: LoRA Multimodal
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
group: Models - Distributed
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
steps:
|
||||
- label: Distributed Model Tests (2 GPUs)
|
||||
key: distributed-model-tests-2-gpus
|
||||
timeout_in_minutes: 50
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 2+
|
||||
mem: 24+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/model_loader/sharded_state_loader.py
|
||||
- vllm/model_executor/models/
|
||||
- tests/model_executor/model_loader/test_sharded_state_loader.py
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
pytest -v -s model_executor/model_loader/test_sharded_state_loader.py -m "not slow_test"'
|
||||
+21
-21
@@ -1198,6 +1198,27 @@ steps:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-mi3xx.txt
|
||||
|
||||
- label: ROCm LM Eval Large Models (8 GPUs) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
agent_pool: mi300_8
|
||||
optional: true
|
||||
num_gpus: 8
|
||||
working_dir: "/vllm-workspace/.buildkite/lm-eval-harness"
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/models/
|
||||
- vllm/model_executor/model_loader/
|
||||
- vllm/model_executor/layers/quantization/
|
||||
- vllm/v1/attention/backends/
|
||||
- vllm/v1/attention/selector.py
|
||||
- vllm/model_executor/layers/layernorm.py
|
||||
- csrc/
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large-rocm.txt --tp-size=8
|
||||
|
||||
#--------------------------------------------------------- mi300 · examples ----------------------------------------------------------#
|
||||
|
||||
- label: Examples # TBD
|
||||
@@ -2392,27 +2413,6 @@ steps:
|
||||
- export VLLM_USE_DEEP_GEMM=0
|
||||
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large-rocm-fp8.txt --tp-size=4
|
||||
|
||||
- label: ROCm LM Eval Large Models (8 GPUs) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi325]
|
||||
agent_pool: mi325_8
|
||||
optional: true
|
||||
num_gpus: 8
|
||||
working_dir: "/vllm-workspace/.buildkite/lm-eval-harness"
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/models/
|
||||
- vllm/model_executor/model_loader/
|
||||
- vllm/model_executor/layers/quantization/
|
||||
- vllm/v1/attention/backends/
|
||||
- vllm/v1/attention/selector.py
|
||||
- vllm/model_executor/layers/layernorm.py
|
||||
- csrc/
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large-rocm.txt --tp-size=8
|
||||
|
||||
#----------------------------------------------------- mi325 · models / language -----------------------------------------------------#
|
||||
|
||||
- label: Language Models Test (Extended Generation) # TBD
|
||||
|
||||
@@ -29,6 +29,8 @@ steps:
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
# TODO(akaratza): Test after Torch >= 2.12 bump
|
||||
soft_fail: true
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
|
||||
@@ -94,6 +94,8 @@ steps:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 65
|
||||
# TODO(akaratza): Test after Torch >= 2.12 bump
|
||||
soft_fail: true
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -327,7 +327,7 @@ jobs:
|
||||
message: 'CC {users} for ROCm-related issue',
|
||||
},
|
||||
mistral: {
|
||||
users: ['patrickvonplaten', 'juliendenize', 'andylolu2'],
|
||||
users: ['patrickvonplaten', 'juliendenize', 'andylolu2', 'NickLucche'],
|
||||
message: 'CC {users} for Mistral-related issue',
|
||||
},
|
||||
// Add more label -> user mappings here
|
||||
|
||||
@@ -27,7 +27,7 @@ jobs:
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@8e8c483db84b4bee98b60c0593521ed34d9990e8 # v6.0.1
|
||||
- uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
|
||||
@@ -48,8 +48,8 @@ jobs:
|
||||
if: always() && (needs.pre-run-check.result == 'success' || needs.pre-run-check.result == 'skipped')
|
||||
runs-on: [self-hosted, linux, x64, vllm-runners]
|
||||
steps:
|
||||
- uses: actions/checkout@8e8c483db84b4bee98b60c0593521ed34d9990e8 # v6.0.1
|
||||
- uses: actions/setup-python@83679a892e2d95755f2dac6acb0bfd1e9ac5d548 # v6.1.0
|
||||
- uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
# Provide shellcheck on PATH so tools/pre_commit/shellcheck.sh skips its
|
||||
|
||||
@@ -131,6 +131,19 @@ repos:
|
||||
--python-version, "3.12",
|
||||
]
|
||||
files: ^requirements/(common|xpu|test/xpu)\.(in|txt)$
|
||||
- id: pip-compile
|
||||
alias: pip-compile-cpu
|
||||
name: pip-compile-cpu
|
||||
args: [
|
||||
requirements/test/cuda.in,
|
||||
-o, requirements/test/cpu.txt,
|
||||
--index-strategy, unsafe-best-match,
|
||||
--torch-backend, cpu,
|
||||
--python-platform, x86_64-manylinux_2_28,
|
||||
--python-version, "3.12",
|
||||
]
|
||||
files: ^requirements/(common|cpu|test/(cuda|cpu))\.(in|txt)$
|
||||
exclude: ^requirements/test/cuda\.txt$
|
||||
- id: pip-compile
|
||||
alias: pip-compile-docs
|
||||
name: pip-compile-docs
|
||||
|
||||
+9
-19
@@ -193,26 +193,16 @@ FROM base AS vllm-test-deps
|
||||
|
||||
WORKDIR /vllm-workspace
|
||||
|
||||
# Copy test requirements
|
||||
COPY requirements/test/cuda.in requirements/test/cpu.in
|
||||
# Test requirements are compiled from requirements/test/cuda.in into
|
||||
# requirements/test/cpu.txt by the pip-compile-cpu pre-commit hook, which
|
||||
# resolves CPU wheels via uv's --torch-backend cpu.
|
||||
COPY requirements/test/cpu.txt requirements/test/cpu.txt
|
||||
|
||||
RUN \
|
||||
sed -i '/mamba_ssm/d' requirements/test/cpu.in && \
|
||||
remove_packages_not_supported_on_aarch64() { \
|
||||
case "$(uname -m)" in \
|
||||
aarch64|arm64) \
|
||||
sed -i '/decord/d' requirements/test/cpu.in; \
|
||||
sed -i '/terratorch/d' requirements/test/cpu.in; \
|
||||
;; \
|
||||
esac; \
|
||||
}; \
|
||||
remove_packages_not_supported_on_aarch64 && \
|
||||
sed -i 's/^torch==.*/torch==2.11.0/g' requirements/test/cpu.in && \
|
||||
sed -i 's/torchaudio.*/torchaudio/g' requirements/test/cpu.in && \
|
||||
sed -i 's/torchvision.*/torchvision/g' requirements/test/cpu.in && \
|
||||
# Related issue: https://github.com/vllm-project/vllm/pull/38800#issuecomment-4228314305
|
||||
sed -i 's/^sentence-transformers.*/sentence-transformers==5.3.0/g' requirements/test/cpu.in && \
|
||||
uv pip compile requirements/test/cpu.in -o requirements/test/cpu.txt --index-strategy unsafe-best-match --torch-backend cpu
|
||||
# cpu.txt is compiled for x86_64, so platform markers are resolved away. Drop
|
||||
# packages unavailable on aarch64 (decord, terratorch) for arm builds.
|
||||
RUN case "$(uname -m)" in \
|
||||
aarch64|arm64) sed -i '/^decord==/d; /^terratorch==/d' requirements/test/cpu.txt ;; \
|
||||
esac
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install -r requirements/test/cpu.txt
|
||||
|
||||
@@ -167,6 +167,7 @@ Priority is **1 = highest** (tried first).
|
||||
| `FLASH_ATTN` | FA4* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | ≥10.0 |
|
||||
| `FLASH_ATTN_DIFFKV` | | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
|
||||
| `FLEX_ATTENTION` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder Only | Any |
|
||||
| `HPC_ATTN` | | fp16, bf16 | `auto`, `fp8_e4m3` | 64 | 128 | ❌ | ❌ | ❌ | ❌ | Decoder | ≥9.0 |
|
||||
| `ROCM_AITER_FA` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32 | 64, 128, 256 | ✅ | ✅ | ❌ | ❌ | Decoder | N/A |
|
||||
| `ROCM_AITER_UNIFIED_ATTN` | | bf16 | `auto`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ✅ | ❌ | ✅ | ❌ | All | N/A |
|
||||
| `ROCM_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 128, 160, 192, 224, 256 | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder, Encoder Only | N/A |
|
||||
|
||||
@@ -120,6 +120,20 @@ To enable KV cache sharing between multiple vLLM instances using the same `root_
|
||||
PYTHONHASHSEED=0 vllm serve ...
|
||||
```
|
||||
|
||||
### P2P (Including P/D)
|
||||
|
||||
The P2P tier (`type: "p2p"`) shares completed KV blocks between vLLM instances over RDMA via NIXL. Each instance binds a control socket on `host:port` and exchanges blocks directly with peers — no shared filesystem required.
|
||||
|
||||
| Key | Required | Default | Notes |
|
||||
| --- | --- | --- | --- |
|
||||
| `type` | yes | — | Must be `p2p`. |
|
||||
| `host` | no | `0.0.0.0` | Address the control socket binds to. |
|
||||
| `port` | no | `7777` | Port for the control socket. Must be reachable from peers. |
|
||||
| `backends` | no | `["UCX"]` | NIXL transport backends. See [NixlConnector Usage Guide](nixl_connector_usage.md#selecting-a-nixl-transport-backend-plugin) for available backends and selection guidance. |
|
||||
| `num_threads` | no | `4` | NIXL agent worker threads. Only used when `backends` is UCX-only; ignored when any non-UCX backend is requested. |
|
||||
|
||||
The `backends` and `num_threads` options mirror the conditional logic used by [`NixlConnector`](nixl_connector_usage.md#selecting-a-nixl-transport-backend-plugin): when any non-UCX backend is configured, NIXL is initialised with `backends=...`; otherwise it falls back to a UCX-only agent with the configured `num_threads`. This lets the P2P tier use a different transport (e.g. `MOONCAKE`, `GDS_MT`, `LIBFABRIC`) than the main `NixlConnector` running in the same process.
|
||||
|
||||
## Tuning Tips
|
||||
|
||||
- `cpu_bytes_to_use`: a bigger CPU tier means fewer trips to slower secondary tiers and a higher hit rate. The value is total across all workers, not per-worker. Leave headroom for the rest of the host workload.
|
||||
|
||||
@@ -586,7 +586,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
|
||||
| `MiDashengLMModel` | MiDashengLM | T + A<sup>+</sup> | `mispeech/midashenglm-7b` | | ✅︎ |
|
||||
| `MiMoV2OmniForCausalLM` | MiMo-V2.5-Omni | T + I<sup>E+</sup> + V<sup>E+</sup> + A<sup>+</sup> | `XiaomiMiMo/MiMo-V2.5-Omni` | | ✅︎ |
|
||||
| `MiniCPMO` | MiniCPM-O | T + I<sup>E+</sup> + V<sup>E+</sup> + A<sup>E+</sup> | `openbmb/MiniCPM-o-2_6`, etc. | ✅︎ | ✅︎ |
|
||||
| `MiniCPMV` | MiniCPM-V | T + I<sup>E+</sup> + V<sup>E+</sup> | `openbmb/MiniCPM-V-2` (see note), `openbmb/MiniCPM-Llama3-V-2_5`, `openbmb/MiniCPM-V-2_6`, `openbmb/MiniCPM-V-4`, `openbmb/MiniCPM-V-4_5`, etc. | ✅︎ | |
|
||||
| `MiniCPMV` | MiniCPM-V | T + I<sup>E+</sup> + V<sup>E+</sup> | `openbmb/MiniCPM-V-2` (see note), `openbmb/MiniCPM-Llama3-V-2_5`, `openbmb/MiniCPM-V-2_6`, `openbmb/MiniCPM-V-4`, `openbmb/MiniCPM-V-4_5`, `openbmb/MiniCPM-V-4_6`, etc. | ✅︎ | |
|
||||
| `MiniMaxM3SparseForConditionalGeneration` | MiniMax-M3 | T + I<sup>+</sup> + V<sup>+</sup> | `MiniMaxAI/MiniMax-M3`, `MiniMaxAI/MiniMax-M3-MXFP8`, etc. | | ✅︎ |
|
||||
| `MiniMaxVL01ForConditionalGeneration` | MiniMax-VL | T + I<sup>E+</sup> | `MiniMaxAI/MiniMax-VL-01`, etc. | | ✅︎ |
|
||||
| `Mistral3ForConditionalGeneration` | Mistral3 (HF Transformers) | T + I<sup>+</sup> | `mistralai/Mistral-Small-3.1-24B-Instruct-2503`, etc. | ✅︎ | ✅︎ |
|
||||
|
||||
@@ -85,6 +85,21 @@ significantly reduce the attack surface for these types of abuse.
|
||||
Also, consider setting `VLLM_MEDIA_URL_ALLOW_REDIRECTS=0` to prevent HTTP
|
||||
redirects from being followed to bypass domain restrictions.
|
||||
|
||||
### 5. **Restrict Media Decode Sizes:**
|
||||
|
||||
Compressed media files can expand into gigabytes of memory during decoding. vLLM
|
||||
enforces decode-size limits to prevent out-of-memory denial of service:
|
||||
|
||||
| Environment Variable | Default | Description |
|
||||
| --- | --- | --- |
|
||||
| `VLLM_MAX_IMAGE_PIXELS` | `178956970` (~179M pixels) | Maximum decoded image size in pixels. Images exceeding this are rejected before raster memory is allocated. Default matches PIL's built-in 2x decompression-bomb threshold (~680 MB for RGB). |
|
||||
| `VLLM_MAX_AUDIO_CLIP_FILESIZE_MB` | `25` | Maximum filesize in MB for a single audio file. |
|
||||
| `VLLM_MAX_AUDIO_DECODE_DURATION_S` | `600` | Maximum decoded audio duration in seconds. Prevents compressed audio from expanding into gigabytes of float32 PCM. |
|
||||
|
||||
Setting any of these to `0` disables the corresponding limit. This is **not
|
||||
recommended** for deployments exposed to untrusted users, as it removes the
|
||||
protection against resource-exhaustion attacks.
|
||||
|
||||
## Security and Firewalls: Protecting Exposed vLLM Systems
|
||||
|
||||
While vLLM is designed to allow unsafe network services to be isolated to
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,3 +1,5 @@
|
||||
-r ../common.txt
|
||||
|
||||
# testing
|
||||
pytest
|
||||
tensorizer==2.10.1
|
||||
@@ -13,7 +15,6 @@ albumentations # required for Nemotron Parse in test_common.py
|
||||
av # required for audio_in_video tests
|
||||
backoff # required for phi4mm test
|
||||
blobfile # required for kimi-vl test
|
||||
einops # required for MPT, qwen-vl
|
||||
httpx
|
||||
librosa # required for audio tests
|
||||
vector_quantize_pytorch # required for minicpmo_26 test
|
||||
@@ -34,7 +35,6 @@ matplotlib # required for qwen-vl test
|
||||
mistral_common[image,audio] >= 1.11.5 # required for voxtral test
|
||||
num2words # required for smolvlm test
|
||||
open_clip_torch==2.32.0 # Required for nemotron_vl test, Nemotron Parse in test_common.py
|
||||
opencv-python-headless >= 4.13.0 # required for video test
|
||||
datamodel_code_generator # required for minicpm3 test
|
||||
lm-eval[api]>=0.4.12 # required for model evaluation test
|
||||
mteb[bm25s]>=2, <3 # required for mteb test
|
||||
@@ -55,11 +55,9 @@ grpcio-reflection==1.78.0
|
||||
|
||||
arctic-inference == 0.1.1; platform_machine == "x86_64" # Required for suffix decoding test
|
||||
numba == 0.65.0 # Required for N-gram speculative decoding
|
||||
numpy
|
||||
runai-model-streamer[s3,gcs,azure]==0.15.7
|
||||
fastsafetensors>=0.3.2
|
||||
instanttensor>=0.1.5; platform_machine == "x86_64"
|
||||
pydantic>=2.12 # 2.11 leads to error on python 3.13
|
||||
decord==0.6.0; platform_machine == "x86_64"
|
||||
# terratorch is temporarily disabled while PyPI has the `lightning` package
|
||||
# in `quarantined` status (every published terratorch version transitively
|
||||
|
||||
+288
-18
@@ -9,6 +9,7 @@ aiohappyeyeballs==2.6.1
|
||||
aiohttp==3.13.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# aiohttp-cors
|
||||
# datasets
|
||||
# fsspec
|
||||
@@ -24,17 +25,34 @@ albumentations==1.4.6
|
||||
alembic==1.16.4
|
||||
# via optuna
|
||||
annotated-doc==0.0.4
|
||||
# via fastapi
|
||||
# via
|
||||
# fastapi
|
||||
# typer
|
||||
annotated-types==0.7.0
|
||||
# via pydantic
|
||||
anyio==4.6.2.post1
|
||||
anthropic==0.112.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
anyio==4.14.1
|
||||
# via
|
||||
# anthropic
|
||||
# httpx
|
||||
# mcp
|
||||
# openai
|
||||
# sse-starlette
|
||||
# starlette
|
||||
# watchfiles
|
||||
apache-tvm-ffi==0.1.9
|
||||
# via
|
||||
# -c requirements/cuda.txt
|
||||
# xgrammar
|
||||
arctic-inference==0.1.1
|
||||
# via -r requirements/test/cuda.in
|
||||
argcomplete==3.5.1
|
||||
# via datamodel-code-generator
|
||||
astor==0.8.1
|
||||
# via depyf
|
||||
attrs==24.2.0
|
||||
# via
|
||||
# aiohttp
|
||||
@@ -59,6 +77,8 @@ bitsandbytes==0.49.2
|
||||
# via -r requirements/test/cuda.in
|
||||
black==24.10.0
|
||||
# via datamodel-code-generator
|
||||
blake3==1.0.9
|
||||
# via -r requirements/test/../common.txt
|
||||
blobfile==3.0.0
|
||||
# via -r requirements/test/cuda.in
|
||||
bm25s==0.2.13
|
||||
@@ -76,12 +96,17 @@ bounded-pool-executor==0.0.3
|
||||
buildkite-test-collector==0.1.9
|
||||
# via -r requirements/test/cuda.in
|
||||
cachetools==5.5.2
|
||||
# via google-auth
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# google-auth
|
||||
cbor2==6.1.2
|
||||
# via -r requirements/test/../common.txt
|
||||
certifi==2024.8.30
|
||||
# via
|
||||
# httpcore
|
||||
# httpx
|
||||
# requests
|
||||
# sentry-sdk
|
||||
cffi==2.0.0
|
||||
# via
|
||||
# cryptography
|
||||
@@ -98,9 +123,11 @@ click==8.1.7
|
||||
# jiwer
|
||||
# nltk
|
||||
# ray
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# typer
|
||||
# uvicorn
|
||||
cloudpickle==3.1.2
|
||||
# via -r requirements/test/../common.txt
|
||||
cohere-melody==0.9.0
|
||||
# via -r requirements/test/cuda.in
|
||||
colorama==0.4.6
|
||||
@@ -111,6 +138,10 @@ colorful==0.5.6
|
||||
# via ray
|
||||
colorlog==6.10.1
|
||||
# via optuna
|
||||
compressed-tensors==0.17.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
contourpy==1.3.0
|
||||
# via matplotlib
|
||||
coverage==7.10.6
|
||||
@@ -149,30 +180,49 @@ decorator==5.1.1
|
||||
# via librosa
|
||||
decord==0.6.0
|
||||
# via -r requirements/test/cuda.in
|
||||
depyf==0.20.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
detect-installer==0.1.0
|
||||
# via fastapi-cloud-cli
|
||||
dill==0.3.8
|
||||
# via
|
||||
# datasets
|
||||
# depyf
|
||||
# evaluate
|
||||
# lm-eval
|
||||
# multiprocess
|
||||
diskcache==5.6.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
distlib==0.3.9
|
||||
# via virtualenv
|
||||
distro==1.9.0
|
||||
# via
|
||||
# anthropic
|
||||
# openai
|
||||
dnspython==2.7.0
|
||||
# via email-validator
|
||||
docker==7.1.0
|
||||
# via gpt-oss
|
||||
docopt==0.6.2
|
||||
# via num2words
|
||||
docstring-parser==0.18.0
|
||||
# via anthropic
|
||||
einops==0.8.1
|
||||
# via
|
||||
# -r requirements/test/cuda.in
|
||||
# -r requirements/test/../common.txt
|
||||
# encodec
|
||||
# vector-quantize-pytorch
|
||||
# vocos
|
||||
einx==0.3.0
|
||||
# via vector-quantize-pytorch
|
||||
email-validator==2.2.0
|
||||
# via pydantic
|
||||
# via
|
||||
# fastapi
|
||||
# pydantic
|
||||
encodec==0.1.1
|
||||
# via vocos
|
||||
et-xmlfile==2.0.0
|
||||
@@ -182,7 +232,17 @@ evaluate==0.4.3
|
||||
fastapi==0.136.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# gpt-oss
|
||||
# model-hosting-container-standards
|
||||
fastapi-cli==0.0.27
|
||||
# via fastapi
|
||||
fastapi-cloud-cli==0.21.0
|
||||
# via fastapi-cli
|
||||
fastar==0.11.0
|
||||
# via
|
||||
# fastapi
|
||||
# fastapi-cloud-cli
|
||||
fastparquet==2024.11.0
|
||||
# via genai-perf
|
||||
fastrlock==0.8.2
|
||||
@@ -194,6 +254,7 @@ fastsafetensors==0.3.2
|
||||
filelock==3.16.1
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# blobfile
|
||||
# datasets
|
||||
# huggingface-hub
|
||||
@@ -243,7 +304,10 @@ google-crc32c==1.7.1
|
||||
google-resumable-media==2.7.2
|
||||
# via google-cloud-storage
|
||||
googleapis-common-protos==1.70.0
|
||||
# via google-api-core
|
||||
# via
|
||||
# google-api-core
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
gpt-oss==0.0.8
|
||||
# via -r requirements/test/cuda.in
|
||||
graphql-core==3.2.6
|
||||
@@ -254,6 +318,7 @@ grpcio==1.78.0
|
||||
# via
|
||||
# -r requirements/test/cuda.in
|
||||
# grpcio-reflection
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# ray
|
||||
grpcio-reflection==1.78.0
|
||||
# via -r requirements/test/cuda.in
|
||||
@@ -275,12 +340,22 @@ html2text==2025.4.15
|
||||
# via gpt-oss
|
||||
httpcore==1.0.6
|
||||
# via httpx
|
||||
httptools==0.8.0
|
||||
# via uvicorn
|
||||
httpx==0.27.2
|
||||
# via
|
||||
# -r requirements/test/cuda.in
|
||||
# anthropic
|
||||
# fastapi
|
||||
# fastapi-cloud-cli
|
||||
# huggingface-hub
|
||||
# mcp
|
||||
# model-hosting-container-standards
|
||||
# openai
|
||||
# perceptron
|
||||
# schemathesis
|
||||
httpx-sse==0.4.3
|
||||
# via mcp
|
||||
huggingface-hub==1.10.2
|
||||
# via
|
||||
# accelerate
|
||||
@@ -314,6 +389,8 @@ idna==3.10
|
||||
# httpx
|
||||
# requests
|
||||
# yarl
|
||||
ijson==3.5.0
|
||||
# via -r requirements/test/../common.txt
|
||||
imagehash==4.3.2
|
||||
# via -r requirements/test/cuda.in
|
||||
imageio==2.37.0
|
||||
@@ -326,6 +403,8 @@ iniconfig==2.0.0
|
||||
# via pytest
|
||||
instanttensor==0.1.5
|
||||
# via -r requirements/test/cuda.in
|
||||
interegular==0.3.3
|
||||
# via lm-format-enforcer
|
||||
isodate==0.7.2
|
||||
# via azure-storage-blob
|
||||
isort==5.13.2
|
||||
@@ -333,15 +412,21 @@ isort==5.13.2
|
||||
jinja2==3.1.6
|
||||
# via
|
||||
# datamodel-code-generator
|
||||
# fastapi
|
||||
# genai-perf
|
||||
# lm-eval
|
||||
# torch
|
||||
jiter==0.15.0
|
||||
# via
|
||||
# anthropic
|
||||
# openai
|
||||
jiwer==3.0.5
|
||||
# via -r requirements/test/cuda.in
|
||||
jmespath==1.0.1
|
||||
# via
|
||||
# boto3
|
||||
# botocore
|
||||
# model-hosting-container-standards
|
||||
joblib==1.4.2
|
||||
# via
|
||||
# librosa
|
||||
@@ -350,7 +435,9 @@ joblib==1.4.2
|
||||
jsonschema==4.23.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# hypothesis-jsonschema
|
||||
# mcp
|
||||
# mistral-common
|
||||
# ray
|
||||
jsonschema-rs==0.46.5
|
||||
@@ -365,6 +452,10 @@ kaleido==0.2.1
|
||||
# via genai-perf
|
||||
kiwisolver==1.4.7
|
||||
# via matplotlib
|
||||
lark==1.2.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
lazy-loader==0.4
|
||||
# via
|
||||
# librosa
|
||||
@@ -373,10 +464,20 @@ libnacl==2.1.0
|
||||
# via tensorizer
|
||||
librosa==0.10.2.post1
|
||||
# via -r requirements/test/cuda.in
|
||||
llguidance==1.7.6
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
llvmlite==0.47.0
|
||||
# via numba
|
||||
lm-eval==0.4.12
|
||||
# via -r requirements/test/cuda.in
|
||||
lm-format-enforcer==0.11.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
loguru==0.7.3
|
||||
# via compressed-tensors
|
||||
lxml==5.3.0
|
||||
# via
|
||||
# blobfile
|
||||
@@ -398,12 +499,19 @@ mbstrdecoder==1.1.3
|
||||
# dataproperty
|
||||
# pytablewriter
|
||||
# typepy
|
||||
mcp==1.28.1
|
||||
# via -r requirements/test/../common.txt
|
||||
mdurl==0.1.2
|
||||
# via markdown-it-py
|
||||
mistral-common==1.11.5
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/cuda.in
|
||||
model-hosting-container-standards==0.1.16
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
more-itertools==10.5.0
|
||||
# via lm-eval
|
||||
mpmath==1.3.0
|
||||
@@ -418,6 +526,8 @@ msgpack==1.1.0
|
||||
# via
|
||||
# librosa
|
||||
# ray
|
||||
msgspec==0.21.1
|
||||
# via -r requirements/test/../common.txt
|
||||
mteb==2.8.3
|
||||
# via -r requirements/test/cuda.in
|
||||
multidict==6.1.0
|
||||
@@ -434,6 +544,8 @@ networkx==3.2.1
|
||||
# via
|
||||
# scikit-image
|
||||
# torch
|
||||
ninja==1.13.0
|
||||
# via -r requirements/test/../common.txt
|
||||
nltk==3.9.1
|
||||
# via rouge-score
|
||||
num2words==0.5.14
|
||||
@@ -445,7 +557,7 @@ numba==0.65.0
|
||||
# librosa
|
||||
numpy==2.2.6
|
||||
# via
|
||||
# -r requirements/test/cuda.in
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
# albumentations
|
||||
# bitsandbytes
|
||||
@@ -489,6 +601,7 @@ numpy==2.2.6
|
||||
# transformers
|
||||
# tritonclient
|
||||
# vocos
|
||||
# xgrammar
|
||||
nvidia-cublas==13.1.0.3
|
||||
# via
|
||||
# cuda-toolkit
|
||||
@@ -530,9 +643,14 @@ nvidia-nvtx==13.0.85
|
||||
# via cuda-toolkit
|
||||
open-clip-torch==2.32.0
|
||||
# via -r requirements/test/cuda.in
|
||||
openai==2.44.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
openai-harmony==0.0.4
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# gpt-oss
|
||||
opencensus==0.11.4
|
||||
# via ray
|
||||
@@ -541,7 +659,7 @@ opencensus-context==0.1.3
|
||||
opencv-python-headless==4.13.0.90
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/cuda.in
|
||||
# -r requirements/test/../common.txt
|
||||
# albumentations
|
||||
# mistral-common
|
||||
openpyxl==3.1.5
|
||||
@@ -549,24 +667,54 @@ openpyxl==3.1.5
|
||||
opentelemetry-api==1.35.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# opentelemetry-exporter-prometheus
|
||||
# opentelemetry-sdk
|
||||
# opentelemetry-semantic-conventions
|
||||
opentelemetry-exporter-otlp==1.35.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
opentelemetry-exporter-otlp-proto-common==1.35.0
|
||||
# via
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
opentelemetry-exporter-otlp-proto-grpc==1.35.0
|
||||
# via opentelemetry-exporter-otlp
|
||||
opentelemetry-exporter-otlp-proto-http==1.35.0
|
||||
# via opentelemetry-exporter-otlp
|
||||
opentelemetry-exporter-prometheus==0.56b0
|
||||
# via ray
|
||||
opentelemetry-proto==1.35.0
|
||||
# via ray
|
||||
# via
|
||||
# opentelemetry-exporter-otlp-proto-common
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# ray
|
||||
opentelemetry-sdk==1.35.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# opentelemetry-exporter-prometheus
|
||||
# ray
|
||||
opentelemetry-semantic-conventions==0.56b0
|
||||
# via opentelemetry-sdk
|
||||
opentelemetry-semantic-conventions-ai==0.4.13
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
optuna==3.6.1
|
||||
# via genai-perf
|
||||
orjson==3.11.5
|
||||
# via genai-perf
|
||||
outlines-core==0.2.14
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
packaging==24.2
|
||||
# via
|
||||
# accelerate
|
||||
@@ -578,6 +726,7 @@ packaging==24.2
|
||||
# fastparquet
|
||||
# huggingface-hub
|
||||
# lazy-loader
|
||||
# lm-format-enforcer
|
||||
# matplotlib
|
||||
# optuna
|
||||
# peft
|
||||
@@ -597,6 +746,8 @@ pandas==2.2.3
|
||||
# fastparquet
|
||||
# genai-perf
|
||||
# statsmodels
|
||||
partial-json-parser==0.2.1.1.post7
|
||||
# via -r requirements/test/../common.txt
|
||||
pathspec==0.12.1
|
||||
# via black
|
||||
pathvalidate==3.2.1
|
||||
@@ -611,6 +762,7 @@ perf-analyzer==0.1.0
|
||||
# via genai-perf
|
||||
pillow==10.4.0
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# genai-perf
|
||||
# imagehash
|
||||
# imageio
|
||||
@@ -644,8 +796,14 @@ pqdm==0.2.0
|
||||
prometheus-client==0.22.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# opentelemetry-exporter-prometheus
|
||||
# prometheus-fastapi-instrumentator
|
||||
# ray
|
||||
prometheus-fastapi-instrumentator==8.0.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
propcache==0.2.0
|
||||
# via
|
||||
# aiohttp
|
||||
@@ -655,6 +813,7 @@ proto-plus==1.26.1
|
||||
protobuf==6.33.6
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# google-api-core
|
||||
# googleapis-common-protos
|
||||
# grpcio-reflection
|
||||
@@ -664,11 +823,14 @@ protobuf==6.33.6
|
||||
# tensorizer
|
||||
psutil==6.1.0
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
# peft
|
||||
# tensorizer
|
||||
py==1.11.0
|
||||
# via pytest-forked
|
||||
py-cpuinfo==9.0.0
|
||||
# via -r requirements/test/../common.txt
|
||||
py-spy==0.4.0
|
||||
# via ray
|
||||
pyarrow==23.0.0
|
||||
@@ -681,6 +843,8 @@ pyasn1==0.6.1
|
||||
# rsa
|
||||
pyasn1-modules==0.4.2
|
||||
# via google-auth
|
||||
pybase64==1.4.3
|
||||
# via -r requirements/test/../common.txt
|
||||
pycountry==24.6.1
|
||||
# via pydantic-extra-types
|
||||
pycparser==2.22
|
||||
@@ -690,26 +854,43 @@ pycryptodomex==3.22.0
|
||||
pydantic==2.12.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/cuda.in
|
||||
# -r requirements/test/../common.txt
|
||||
# albumentations
|
||||
# anthropic
|
||||
# compressed-tensors
|
||||
# datamodel-code-generator
|
||||
# fastapi
|
||||
# fastapi-cloud-cli
|
||||
# gpt-oss
|
||||
# lm-format-enforcer
|
||||
# mcp
|
||||
# mistral-common
|
||||
# model-hosting-container-standards
|
||||
# mteb
|
||||
# openai
|
||||
# openai-harmony
|
||||
# pydantic-extra-types
|
||||
# pydantic-settings
|
||||
# ray
|
||||
# xgrammar
|
||||
pydantic-core==2.41.1
|
||||
# via pydantic
|
||||
pydantic-extra-types==2.10.5
|
||||
# via mistral-common
|
||||
# via
|
||||
# fastapi
|
||||
# mistral-common
|
||||
pydantic-settings==2.14.2
|
||||
# via
|
||||
# fastapi
|
||||
# mcp
|
||||
pygments==2.18.0
|
||||
# via
|
||||
# pytest
|
||||
# rich
|
||||
pyjwt==2.11.0
|
||||
# via msal
|
||||
# via
|
||||
# mcp
|
||||
# msal
|
||||
pyparsing==3.2.0
|
||||
# via matplotlib
|
||||
pyrate-limiter==4.4.0
|
||||
@@ -751,6 +932,16 @@ python-dateutil==2.9.0.post0
|
||||
# matplotlib
|
||||
# pandas
|
||||
# typepy
|
||||
python-dotenv==1.2.2
|
||||
# via
|
||||
# pydantic-settings
|
||||
# uvicorn
|
||||
python-json-logger==4.1.0
|
||||
# via -r requirements/test/../common.txt
|
||||
python-multipart==0.0.32
|
||||
# via
|
||||
# fastapi
|
||||
# mcp
|
||||
python-rapidjson==1.20
|
||||
# via tritonclient
|
||||
pytrec-eval-terrier==0.5.7
|
||||
@@ -763,12 +954,14 @@ pywavelets==1.9.0
|
||||
# via imagehash
|
||||
pyyaml==6.0.2
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
# albumentations
|
||||
# datamodel-code-generator
|
||||
# datasets
|
||||
# genai-perf
|
||||
# huggingface-hub
|
||||
# lm-format-enforcer
|
||||
# optuna
|
||||
# peft
|
||||
# ray
|
||||
@@ -776,7 +969,12 @@ pyyaml==6.0.2
|
||||
# schemathesis
|
||||
# timm
|
||||
# transformers
|
||||
# uvicorn
|
||||
# vocos
|
||||
pyzmq==27.1.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
rapidfuzz==3.12.1
|
||||
# via jiwer
|
||||
ray==2.48.0
|
||||
@@ -789,6 +987,7 @@ referencing==0.35.1
|
||||
# jsonschema-specifications
|
||||
regex==2026.2.28
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# nltk
|
||||
# open-clip-torch
|
||||
# sacrebleu
|
||||
@@ -797,6 +996,7 @@ regex==2026.2.28
|
||||
requests==2.32.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# azure-core
|
||||
# buildkite-test-collector
|
||||
# datasets
|
||||
@@ -809,6 +1009,7 @@ requests==2.32.3
|
||||
# mistral-common
|
||||
# msal
|
||||
# mteb
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# pooch
|
||||
# ray
|
||||
# responses
|
||||
@@ -822,8 +1023,15 @@ rich==13.9.4
|
||||
# genai-perf
|
||||
# mteb
|
||||
# perceptron
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# typer
|
||||
rich-toolkit==0.20.1
|
||||
# via
|
||||
# fastapi-cli
|
||||
# fastapi-cloud-cli
|
||||
rignore==0.7.6
|
||||
# via fastapi-cloud-cli
|
||||
rouge-score==0.1.2
|
||||
# via lm-eval
|
||||
rpds-py==0.20.1
|
||||
@@ -847,6 +1055,7 @@ sacrebleu==2.4.3
|
||||
safetensors==0.7.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
# open-clip-torch
|
||||
# peft
|
||||
@@ -882,9 +1091,17 @@ sentence-transformers==5.2.0
|
||||
# via
|
||||
# -r requirements/test/cuda.in
|
||||
# mteb
|
||||
sentencepiece==0.2.1
|
||||
# via -r requirements/test/../common.txt
|
||||
sentry-sdk==2.63.0
|
||||
# via fastapi-cloud-cli
|
||||
setproctitle==1.3.7
|
||||
# via -r requirements/test/../common.txt
|
||||
setuptools==77.0.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# model-hosting-container-standards
|
||||
# pytablewriter
|
||||
# torch
|
||||
shellingham==1.5.4
|
||||
@@ -894,6 +1111,7 @@ shellingham==1.5.4
|
||||
six==1.16.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# junit-xml
|
||||
# opencensus
|
||||
# python-dateutil
|
||||
@@ -902,8 +1120,9 @@ smart-open==7.1.0
|
||||
# via ray
|
||||
sniffio==1.3.1
|
||||
# via
|
||||
# anyio
|
||||
# anthropic
|
||||
# httpx
|
||||
# openai
|
||||
sortedcontainers==2.4.0
|
||||
# via hypothesis
|
||||
soundfile==0.12.1
|
||||
@@ -922,10 +1141,17 @@ sqlalchemy==2.0.41
|
||||
# optuna
|
||||
sqlitedict==2.1.0
|
||||
# via lm-eval
|
||||
sse-starlette==3.4.5
|
||||
# via mcp
|
||||
starlette==1.3.1
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# fastapi
|
||||
# mcp
|
||||
# model-hosting-container-standards
|
||||
# prometheus-fastapi-instrumentator
|
||||
# sse-starlette
|
||||
# starlette-testclient
|
||||
starlette-testclient==0.4.1
|
||||
# via schemathesis
|
||||
@@ -933,6 +1159,8 @@ statsmodels==0.14.4
|
||||
# via genai-perf
|
||||
structlog==25.4.0
|
||||
# via gpt-oss
|
||||
supervisor==4.3.0
|
||||
# via model-hosting-container-standards
|
||||
sympy==1.13.3
|
||||
# via
|
||||
# einx
|
||||
@@ -962,6 +1190,7 @@ tifffile==2025.3.30
|
||||
tiktoken==0.12.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# gpt-oss
|
||||
# lm-eval
|
||||
# mistral-common
|
||||
@@ -973,6 +1202,7 @@ timm==1.0.17
|
||||
tokenizers==0.22.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/cuda.in
|
||||
# transformers
|
||||
torch==2.11.0+cu130
|
||||
@@ -981,6 +1211,7 @@ torch==2.11.0+cu130
|
||||
# -r requirements/test/cuda.in
|
||||
# accelerate
|
||||
# bitsandbytes
|
||||
# compressed-tensors
|
||||
# encodec
|
||||
# instanttensor
|
||||
# mteb
|
||||
@@ -994,6 +1225,7 @@ torch==2.11.0+cu130
|
||||
# torchvision
|
||||
# vector-quantize-pytorch
|
||||
# vocos
|
||||
# xgrammar
|
||||
torchaudio==2.11.0+cu130
|
||||
# via
|
||||
# -c requirements/cuda.txt
|
||||
@@ -1009,6 +1241,7 @@ torchvision==0.26.0+cu130
|
||||
# timm
|
||||
tqdm==4.67.3
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# datasets
|
||||
# evaluate
|
||||
# huggingface-hub
|
||||
@@ -1016,6 +1249,7 @@ tqdm==4.67.3
|
||||
# mteb
|
||||
# nltk
|
||||
# open-clip-torch
|
||||
# openai
|
||||
# optuna
|
||||
# peft
|
||||
# pqdm
|
||||
@@ -1025,15 +1259,20 @@ tqdm==4.67.3
|
||||
transformers==5.5.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/cuda.in
|
||||
# compressed-tensors
|
||||
# genai-perf
|
||||
# peft
|
||||
# sentence-transformers
|
||||
# transformers-stream-generator
|
||||
# xgrammar
|
||||
transformers-stream-generator==0.0.5
|
||||
# via -r requirements/test/cuda.in
|
||||
triton==3.6.0
|
||||
# via torch
|
||||
# via
|
||||
# torch
|
||||
# xgrammar
|
||||
tritonclient==2.64.0
|
||||
# via -r requirements/test/cuda.in
|
||||
typepy==1.3.2
|
||||
@@ -1041,8 +1280,10 @@ typepy==1.3.2
|
||||
# dataproperty
|
||||
# pytablewriter
|
||||
# tabledata
|
||||
typer==0.15.2
|
||||
typer==0.26.8
|
||||
# via
|
||||
# fastapi-cli
|
||||
# fastapi-cloud-cli
|
||||
# fastsafetensors
|
||||
# huggingface-hub
|
||||
# perceptron
|
||||
@@ -1050,9 +1291,13 @@ typer==0.15.2
|
||||
typing-extensions==4.15.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# aiosignal
|
||||
# albumentations
|
||||
# alembic
|
||||
# anthropic
|
||||
# anyio
|
||||
# apache-tvm-ffi
|
||||
# azure-core
|
||||
# azure-identity
|
||||
# azure-storage-blob
|
||||
@@ -1062,9 +1307,13 @@ typing-extensions==4.15.0
|
||||
# huggingface-hub
|
||||
# librosa
|
||||
# lm-eval
|
||||
# mcp
|
||||
# mistral-common
|
||||
# mteb
|
||||
# openai
|
||||
# opentelemetry-api
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# opentelemetry-sdk
|
||||
# opentelemetry-semantic-conventions
|
||||
# pqdm
|
||||
@@ -1072,17 +1321,20 @@ typing-extensions==4.15.0
|
||||
# pydantic-core
|
||||
# pydantic-extra-types
|
||||
# pytest-asyncio
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# sentence-transformers
|
||||
# sqlalchemy
|
||||
# starlette
|
||||
# torch
|
||||
# typer
|
||||
# typing-inspection
|
||||
# xgrammar
|
||||
typing-inspection==0.4.2
|
||||
# via
|
||||
# fastapi
|
||||
# mcp
|
||||
# pydantic
|
||||
# pydantic-settings
|
||||
tzdata==2024.2
|
||||
# via pandas
|
||||
urllib3==2.2.3
|
||||
@@ -1092,23 +1344,41 @@ urllib3==2.2.3
|
||||
# docker
|
||||
# requests
|
||||
# responses
|
||||
# sentry-sdk
|
||||
# tritonclient
|
||||
uvicorn==0.35.0
|
||||
# via gpt-oss
|
||||
# via
|
||||
# fastapi
|
||||
# fastapi-cli
|
||||
# fastapi-cloud-cli
|
||||
# gpt-oss
|
||||
# mcp
|
||||
uvloop==0.22.1
|
||||
# via uvicorn
|
||||
vector-quantize-pytorch==1.21.2
|
||||
# via -r requirements/test/cuda.in
|
||||
virtualenv==20.31.2
|
||||
# via ray
|
||||
vocos==0.1.0
|
||||
# via -r requirements/test/cuda.in
|
||||
watchfiles==1.2.0
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# uvicorn
|
||||
wcwidth==0.2.13
|
||||
# via ftfy
|
||||
websockets==16.0
|
||||
# via uvicorn
|
||||
werkzeug==3.1.3
|
||||
# via schemathesis
|
||||
word2number==1.1
|
||||
# via lm-eval
|
||||
wrapt==1.17.2
|
||||
# via smart-open
|
||||
xgrammar==0.2.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
xxhash==3.5.0
|
||||
# via
|
||||
# datasets
|
||||
|
||||
@@ -15,7 +15,6 @@ albumentations # required for Nemotron Parse in test_common.py
|
||||
av # required for audio_in_video tests
|
||||
backoff # required for phi4mm test
|
||||
blobfile # required for kimi-vl test
|
||||
einops # required for MPT, qwen-vl
|
||||
httpx
|
||||
librosa # required for audio tests
|
||||
vector_quantize_pytorch # required for minicpmo_26 test
|
||||
@@ -33,7 +32,6 @@ matplotlib # required for qwen-vl test
|
||||
mistral_common[image,audio]>=1.11.5 # required for voxtral test
|
||||
num2words # required for smolvlm test
|
||||
open_clip_torch==2.32.0 # Required for nemotron_vl test, Nemotron Parse in test_common.py
|
||||
opencv-python-headless>=4.13.0 # required for video test
|
||||
datamodel_code_generator # required for minicpm3 test
|
||||
lm-eval[api]>=0.4.12 # required for model evaluation test
|
||||
mteb[bm25s]>=2, <3 # required for mteb test
|
||||
@@ -54,11 +52,9 @@ grpcio-reflection==1.78.0
|
||||
|
||||
arctic-inference==0.1.1 # Required for suffix decoding test
|
||||
numba==0.65.0 # Required for N-gram speculative decoding
|
||||
numpy
|
||||
runai-model-streamer[s3,gcs,azure]==0.15.7
|
||||
fastsafetensors>=0.3.2
|
||||
instanttensor>=0.1.5
|
||||
pydantic>=2.12 # 2.11 leads to error on python 3.13
|
||||
decord==0.6.0
|
||||
|
||||
# Prithvi tests
|
||||
@@ -74,6 +70,7 @@ gpt-oss>=0.0.7; python_version > '3.11'
|
||||
|
||||
perceptron # required for isaac test
|
||||
kaldi-native-fbank>=1.18.7 # required for fireredasr2 test
|
||||
cohere_melody>=0.9.0 # required for cohere command reasoning parser test
|
||||
|
||||
# Newer versions of datasets require torchcoded, that makes the tests fail in CI because of a missing library.
|
||||
# Older versions are in conflict with terratorch requirements.
|
||||
|
||||
@@ -130,6 +130,8 @@ cloudpickle==3.1.2
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# tilelang
|
||||
cohere-melody==0.9.0
|
||||
# via -r requirements/test/rocm.in
|
||||
colorama==0.4.6
|
||||
# via
|
||||
# perceptron
|
||||
@@ -205,7 +207,6 @@ docstring-parser==0.17.0
|
||||
einops==0.8.2
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/rocm.in
|
||||
# encodec
|
||||
# vector-quantize-pytorch
|
||||
# vocos
|
||||
@@ -561,7 +562,6 @@ numba==0.65.0
|
||||
numpy==2.2.6
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/rocm.in
|
||||
# accelerate
|
||||
# albumentations
|
||||
# bitsandbytes
|
||||
@@ -630,7 +630,6 @@ opencv-python-headless==4.13.0.92
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/rocm.in
|
||||
# albumentations
|
||||
# mistral-common
|
||||
openpyxl==3.1.5
|
||||
@@ -834,7 +833,6 @@ pydantic==2.12.5
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/rocm.in
|
||||
# albumentations
|
||||
# anthropic
|
||||
# compressed-tensors
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
-r ../common.txt
|
||||
|
||||
# --- Test Infrastructure ---
|
||||
tblib
|
||||
pytest
|
||||
|
||||
+316
-4
@@ -11,6 +11,7 @@ aiohappyeyeballs==2.6.1
|
||||
aiohttp==3.13.4
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# fsspec
|
||||
# gpt-oss
|
||||
# lm-eval
|
||||
@@ -24,12 +25,25 @@ annotated-doc==0.0.4
|
||||
# typer
|
||||
annotated-types==0.7.0
|
||||
# via pydantic
|
||||
anthropic==0.112.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
anyio==4.13.0
|
||||
# via
|
||||
# anthropic
|
||||
# httpx
|
||||
# mcp
|
||||
# openai
|
||||
# sse-starlette
|
||||
# starlette
|
||||
# watchfiles
|
||||
apache-tvm-ffi==0.1.12
|
||||
# via xgrammar
|
||||
arctic-inference==0.1.1
|
||||
# via -r requirements/test/xpu.in
|
||||
astor==0.8.1
|
||||
# via depyf
|
||||
attrs==26.1.0
|
||||
# via
|
||||
# aiohttp
|
||||
@@ -39,6 +53,8 @@ audioread==3.0.1
|
||||
# via
|
||||
# -r requirements/test/xpu.in
|
||||
# librosa
|
||||
blake3==1.0.9
|
||||
# via -r requirements/test/../common.txt
|
||||
blobfile==3.0.0
|
||||
# via -r requirements/test/xpu.in
|
||||
bm25s==0.2.13
|
||||
@@ -47,13 +63,20 @@ bm25s==0.2.13
|
||||
# mteb
|
||||
bounded-pool-executor==0.0.3
|
||||
# via pqdm
|
||||
cachetools==7.1.4
|
||||
# via -r requirements/test/../common.txt
|
||||
cbor2==6.1.2
|
||||
# via -r requirements/test/../common.txt
|
||||
certifi==2026.2.25
|
||||
# via
|
||||
# httpcore
|
||||
# httpx
|
||||
# requests
|
||||
# sentry-sdk
|
||||
cffi==2.0.0
|
||||
# via soundfile
|
||||
# via
|
||||
# cryptography
|
||||
# soundfile
|
||||
chardet==5.2.0
|
||||
# via mbstrdecoder
|
||||
charset-normalizer==3.4.6
|
||||
@@ -64,13 +87,22 @@ click==8.3.1
|
||||
# via
|
||||
# jiwer
|
||||
# nltk
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# typer
|
||||
# uvicorn
|
||||
cloudpickle==3.1.2
|
||||
# via -r requirements/test/../common.txt
|
||||
colorama==0.4.6
|
||||
# via sacrebleu
|
||||
compressed-tensors==0.17.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
coverage==7.13.5
|
||||
# via pytest-cov
|
||||
cryptography==49.0.0
|
||||
# via pyjwt
|
||||
dataproperty==1.1.0
|
||||
# via
|
||||
# pytablewriter
|
||||
@@ -82,16 +114,35 @@ datasets==4.8.4
|
||||
# mteb
|
||||
decorator==5.2.1
|
||||
# via librosa
|
||||
depyf==0.20.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
detect-installer==0.1.0
|
||||
# via fastapi-cloud-cli
|
||||
dill==0.4.1
|
||||
# via
|
||||
# datasets
|
||||
# depyf
|
||||
# evaluate
|
||||
# lm-eval
|
||||
# multiprocess
|
||||
diskcache==5.6.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
distro==1.9.0
|
||||
# via
|
||||
# anthropic
|
||||
# openai
|
||||
dnspython==2.8.0
|
||||
# via email-validator
|
||||
docker==7.1.0
|
||||
# via gpt-oss
|
||||
docopt==0.6.2
|
||||
# via num2words
|
||||
docstring-parser==0.18.0
|
||||
# via anthropic
|
||||
dpcpp-cpp-rt==2025.3.2
|
||||
# via
|
||||
# onemkl-sycl-blas
|
||||
@@ -100,15 +151,30 @@ dpcpp-cpp-rt==2025.3.2
|
||||
# onemkl-sycl-rng
|
||||
# onemkl-sycl-sparse
|
||||
# torch
|
||||
einops==0.8.2
|
||||
# via -r requirements/test/../common.txt
|
||||
email-validator==2.3.0
|
||||
# via
|
||||
# fastapi
|
||||
# pydantic
|
||||
evaluate==0.4.6
|
||||
# via lm-eval
|
||||
fastapi==0.135.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# gpt-oss
|
||||
# model-hosting-container-standards
|
||||
fastapi-cli==0.0.27
|
||||
# via fastapi
|
||||
fastapi-cloud-cli==0.21.0
|
||||
# via fastapi-cli
|
||||
fastar==0.11.0
|
||||
# via fastapi-cloud-cli
|
||||
filelock==3.25.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# blobfile
|
||||
# datasets
|
||||
# huggingface-hub
|
||||
@@ -124,10 +190,16 @@ fsspec==2026.2.0
|
||||
# evaluate
|
||||
# huggingface-hub
|
||||
# torch
|
||||
googleapis-common-protos==1.75.0
|
||||
# via
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
gpt-oss==0.0.8
|
||||
# via -r requirements/test/xpu.in
|
||||
graphql-core==3.2.8
|
||||
# via hypothesis-graphql
|
||||
grpcio==1.81.1
|
||||
# via opentelemetry-exporter-otlp-proto-grpc
|
||||
h11==0.16.0
|
||||
# via
|
||||
# httpcore
|
||||
@@ -140,11 +212,21 @@ html2text==2025.4.15
|
||||
# via gpt-oss
|
||||
httpcore==1.0.9
|
||||
# via httpx
|
||||
httptools==0.8.0
|
||||
# via uvicorn
|
||||
httpx==0.28.1
|
||||
# via
|
||||
# anthropic
|
||||
# datasets
|
||||
# fastapi
|
||||
# fastapi-cloud-cli
|
||||
# huggingface-hub
|
||||
# mcp
|
||||
# model-hosting-container-standards
|
||||
# openai
|
||||
# schemathesis
|
||||
httpx-sse==0.4.3
|
||||
# via mcp
|
||||
huggingface-hub==1.10.2
|
||||
# via
|
||||
# accelerate
|
||||
@@ -166,9 +248,12 @@ hypothesis-jsonschema==0.23.1
|
||||
idna==3.11
|
||||
# via
|
||||
# anyio
|
||||
# email-validator
|
||||
# httpx
|
||||
# requests
|
||||
# yarl
|
||||
ijson==3.5.0
|
||||
# via -r requirements/test/../common.txt
|
||||
imageio==2.37.3
|
||||
# via scikit-image
|
||||
impi-rt==2021.17.2
|
||||
@@ -212,13 +297,22 @@ intel-sycl-rt==2025.3.2
|
||||
# dpcpp-cpp-rt
|
||||
# oneccl
|
||||
# torch
|
||||
interegular==0.3.3
|
||||
# via lm-format-enforcer
|
||||
jinja2==3.1.6
|
||||
# via
|
||||
# -c requirements/xpu.txt
|
||||
# fastapi
|
||||
# lm-eval
|
||||
# torch
|
||||
jiter==0.15.0
|
||||
# via
|
||||
# anthropic
|
||||
# openai
|
||||
jiwer==4.0.0
|
||||
# via -r requirements/test/xpu.in
|
||||
jmespath==1.1.0
|
||||
# via model-hosting-container-standards
|
||||
joblib==1.5.3
|
||||
# via
|
||||
# librosa
|
||||
@@ -227,7 +321,9 @@ joblib==1.5.3
|
||||
jsonschema==4.26.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# hypothesis-jsonschema
|
||||
# mcp
|
||||
# mistral-common
|
||||
# schemathesis
|
||||
jsonschema-rs==0.45.0
|
||||
@@ -236,16 +332,30 @@ jsonschema-specifications==2025.9.1
|
||||
# via jsonschema
|
||||
junit-xml==1.9
|
||||
# via schemathesis
|
||||
lark==1.2.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
lazy-loader==0.5
|
||||
# via
|
||||
# librosa
|
||||
# scikit-image
|
||||
librosa==0.10.2.post1
|
||||
# via -r requirements/test/xpu.in
|
||||
llguidance==1.7.6
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
llvmlite==0.47.0
|
||||
# via numba
|
||||
lm-eval==0.4.12
|
||||
# via -r requirements/test/xpu.in
|
||||
lm-format-enforcer==0.11.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
loguru==0.7.3
|
||||
# via compressed-tensors
|
||||
lxml==6.0.2
|
||||
# via
|
||||
# blobfile
|
||||
@@ -262,11 +372,14 @@ mbstrdecoder==1.1.4
|
||||
# dataproperty
|
||||
# pytablewriter
|
||||
# typepy
|
||||
mcp==1.28.1
|
||||
# via -r requirements/test/../common.txt
|
||||
mdurl==0.1.2
|
||||
# via markdown-it-py
|
||||
mistral-common==1.11.5
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# -r requirements/test/xpu.in
|
||||
mkl==2025.3.1
|
||||
# via
|
||||
@@ -276,6 +389,10 @@ mkl==2025.3.1
|
||||
# onemkl-sycl-rng
|
||||
# onemkl-sycl-sparse
|
||||
# torch
|
||||
model-hosting-container-standards==0.1.16
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
modelscope==1.35.3
|
||||
# via -r requirements/test/xpu.in
|
||||
more-itertools==10.8.0
|
||||
@@ -284,6 +401,8 @@ mpmath==1.3.0
|
||||
# via sympy
|
||||
msgpack==1.1.2
|
||||
# via librosa
|
||||
msgspec==0.21.1
|
||||
# via -r requirements/test/../common.txt
|
||||
mteb==2.12.7
|
||||
# via -r requirements/test/xpu.in
|
||||
multidict==6.7.1
|
||||
@@ -298,6 +417,8 @@ networkx==3.6.1
|
||||
# via
|
||||
# scikit-image
|
||||
# torch
|
||||
ninja==1.13.0
|
||||
# via -r requirements/test/../common.txt
|
||||
nltk==3.9.4
|
||||
# via rouge-score
|
||||
num2words==0.5.14
|
||||
@@ -308,6 +429,7 @@ numba==0.65.0
|
||||
# librosa
|
||||
numpy==2.2.6
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
# albumentations
|
||||
# bm25s
|
||||
@@ -333,6 +455,7 @@ numpy==2.2.6
|
||||
# tifffile
|
||||
# torchvision
|
||||
# transformers
|
||||
# xgrammar
|
||||
oneccl==2021.17.2
|
||||
# via
|
||||
# oneccl-devel
|
||||
@@ -356,15 +479,65 @@ onemkl-sycl-rng==2025.3.1
|
||||
# via torch
|
||||
onemkl-sycl-sparse==2025.3.1
|
||||
# via torch
|
||||
openai==2.44.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
openai-harmony==0.0.8
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# gpt-oss
|
||||
opencv-python-headless==4.13.0.92
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# albumentations
|
||||
# mistral-common
|
||||
opentelemetry-api==1.43.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# opentelemetry-sdk
|
||||
# opentelemetry-semantic-conventions
|
||||
opentelemetry-exporter-otlp==1.43.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
opentelemetry-exporter-otlp-proto-common==1.43.0
|
||||
# via
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
opentelemetry-exporter-otlp-proto-grpc==1.43.0
|
||||
# via opentelemetry-exporter-otlp
|
||||
opentelemetry-exporter-otlp-proto-http==1.43.0
|
||||
# via opentelemetry-exporter-otlp
|
||||
opentelemetry-proto==1.43.0
|
||||
# via
|
||||
# opentelemetry-exporter-otlp-proto-common
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
opentelemetry-sdk==1.43.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# opentelemetry-semantic-conventions-ai
|
||||
opentelemetry-semantic-conventions==0.64b0
|
||||
# via
|
||||
# opentelemetry-sdk
|
||||
# opentelemetry-semantic-conventions-ai
|
||||
opentelemetry-semantic-conventions-ai==0.5.1
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
outlines-core==0.2.14
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
packaging==26.0
|
||||
# via
|
||||
# -c requirements/xpu.txt
|
||||
@@ -373,6 +546,7 @@ packaging==26.0
|
||||
# evaluate
|
||||
# huggingface-hub
|
||||
# lazy-loader
|
||||
# lm-format-enforcer
|
||||
# modelscope
|
||||
# pooch
|
||||
# pytest
|
||||
@@ -384,10 +558,13 @@ pandas==3.0.1
|
||||
# via
|
||||
# datasets
|
||||
# evaluate
|
||||
partial-json-parser==0.2.1.1.post7
|
||||
# via -r requirements/test/../common.txt
|
||||
pathvalidate==3.3.1
|
||||
# via pytablewriter
|
||||
pillow==12.1.1
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# imageio
|
||||
# mistral-common
|
||||
# scikit-image
|
||||
@@ -410,16 +587,37 @@ portalocker==3.2.0
|
||||
# via sacrebleu
|
||||
pqdm==0.2.0
|
||||
# via -r requirements/test/xpu.in
|
||||
prometheus-client==0.25.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# prometheus-fastapi-instrumentator
|
||||
prometheus-fastapi-instrumentator==8.0.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
propcache==0.4.1
|
||||
# via
|
||||
# aiohttp
|
||||
# yarl
|
||||
protobuf==7.35.1
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# googleapis-common-protos
|
||||
# opentelemetry-proto
|
||||
psutil==7.2.2
|
||||
# via accelerate
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
py==1.11.0
|
||||
# via pytest-forked
|
||||
py-cpuinfo==9.0.0
|
||||
# via -r requirements/test/../common.txt
|
||||
pyarrow==23.0.1
|
||||
# via datasets
|
||||
pybase64==1.4.3
|
||||
# via -r requirements/test/../common.txt
|
||||
pycountry==26.2.16
|
||||
# via pydantic-extra-types
|
||||
pycparser==3.0
|
||||
@@ -429,23 +627,41 @@ pycryptodomex==3.23.0
|
||||
pydantic==2.12.5
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# albumentations
|
||||
# anthropic
|
||||
# compressed-tensors
|
||||
# fastapi
|
||||
# fastapi-cloud-cli
|
||||
# gpt-oss
|
||||
# lm-format-enforcer
|
||||
# mcp
|
||||
# mistral-common
|
||||
# model-hosting-container-standards
|
||||
# mteb
|
||||
# openai
|
||||
# openai-harmony
|
||||
# pydantic-extra-types
|
||||
# pydantic-settings
|
||||
# xgrammar
|
||||
pydantic-core==2.41.5
|
||||
# via pydantic
|
||||
pydantic-extra-types==2.11.1
|
||||
# via mistral-common
|
||||
# via
|
||||
# fastapi
|
||||
# mistral-common
|
||||
pydantic-settings==2.14.2
|
||||
# via
|
||||
# fastapi
|
||||
# mcp
|
||||
pyelftools==0.32
|
||||
# via triton-xpu
|
||||
pygments==2.20.0
|
||||
# via
|
||||
# pytest
|
||||
# rich
|
||||
pyjwt==2.13.0
|
||||
# via mcp
|
||||
pyrate-limiter==4.1.0
|
||||
# via schemathesis
|
||||
pystemmer==3.0.0
|
||||
@@ -480,19 +696,36 @@ python-dateutil==2.9.0.post0
|
||||
# via
|
||||
# pandas
|
||||
# typepy
|
||||
python-dotenv==1.2.2
|
||||
# via
|
||||
# pydantic-settings
|
||||
# uvicorn
|
||||
python-json-logger==4.1.0
|
||||
# via -r requirements/test/../common.txt
|
||||
python-multipart==0.0.32
|
||||
# via
|
||||
# fastapi
|
||||
# mcp
|
||||
pytrec-eval-terrier==0.5.10
|
||||
# via mteb
|
||||
pytz==2026.1.post1
|
||||
# via typepy
|
||||
pyyaml==6.0.3
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
# albumentations
|
||||
# datasets
|
||||
# huggingface-hub
|
||||
# lm-format-enforcer
|
||||
# schemathesis
|
||||
# timm
|
||||
# transformers
|
||||
# uvicorn
|
||||
pyzmq==27.1.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
rapidfuzz==3.12.1
|
||||
# via
|
||||
# -r requirements/test/xpu.in
|
||||
@@ -503,6 +736,7 @@ referencing==0.37.0
|
||||
# jsonschema-specifications
|
||||
regex==2026.3.32
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# nltk
|
||||
# sacrebleu
|
||||
# tiktoken
|
||||
@@ -510,6 +744,7 @@ regex==2026.3.32
|
||||
requests==2.33.1
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# datasets
|
||||
# docker
|
||||
# evaluate
|
||||
@@ -518,6 +753,7 @@ requests==2.33.1
|
||||
# mistral-common
|
||||
# modelscope
|
||||
# mteb
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# pooch
|
||||
# schemathesis
|
||||
# starlette-testclient
|
||||
@@ -525,8 +761,15 @@ requests==2.33.1
|
||||
rich==14.3.3
|
||||
# via
|
||||
# mteb
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# typer
|
||||
rich-toolkit==0.20.1
|
||||
# via
|
||||
# fastapi-cli
|
||||
# fastapi-cloud-cli
|
||||
rignore==0.7.6
|
||||
# via fastapi-cloud-cli
|
||||
rouge-score==0.1.2
|
||||
# via lm-eval
|
||||
rpds-py==0.30.0
|
||||
@@ -538,6 +781,7 @@ sacrebleu==2.6.0
|
||||
safetensors==0.7.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# accelerate
|
||||
# timm
|
||||
# transformers
|
||||
@@ -564,10 +808,18 @@ scipy==1.17.1
|
||||
# sentence-transformers
|
||||
sentence-transformers==5.3.0
|
||||
# via mteb
|
||||
sentencepiece==0.2.1
|
||||
# via -r requirements/test/../common.txt
|
||||
sentry-sdk==2.63.0
|
||||
# via fastapi-cloud-cli
|
||||
setproctitle==1.3.7
|
||||
# via -r requirements/test/../common.txt
|
||||
setuptools==80.10.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -c requirements/xpu.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# model-hosting-container-standards
|
||||
# modelscope
|
||||
# pytablewriter
|
||||
# torch
|
||||
@@ -576,9 +828,14 @@ shellingham==1.5.4
|
||||
six==1.17.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# junit-xml
|
||||
# python-dateutil
|
||||
# rouge-score
|
||||
sniffio==1.3.1
|
||||
# via
|
||||
# anthropic
|
||||
# openai
|
||||
sortedcontainers==2.4.0
|
||||
# via hypothesis
|
||||
soundfile==0.13.1
|
||||
@@ -593,15 +850,24 @@ soxr==0.5.0.post1
|
||||
# mistral-common
|
||||
sqlitedict==2.1.0
|
||||
# via lm-eval
|
||||
sse-starlette==3.4.5
|
||||
# via mcp
|
||||
starlette==1.3.1
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# fastapi
|
||||
# mcp
|
||||
# model-hosting-container-standards
|
||||
# prometheus-fastapi-instrumentator
|
||||
# sse-starlette
|
||||
# starlette-testclient
|
||||
starlette-testclient==0.4.1
|
||||
# via schemathesis
|
||||
structlog==25.5.0
|
||||
# via gpt-oss
|
||||
supervisor==4.3.0
|
||||
# via model-hosting-container-standards
|
||||
sympy==1.14.0
|
||||
# via torch
|
||||
tabledata==1.3.4
|
||||
@@ -636,6 +902,7 @@ tifffile==2026.3.3
|
||||
tiktoken==0.12.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# gpt-oss
|
||||
# lm-eval
|
||||
# mistral-common
|
||||
@@ -644,19 +911,23 @@ timm==1.0.17
|
||||
tokenizers==0.22.2
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# transformers
|
||||
torch==2.12.0+xpu
|
||||
# via
|
||||
# -c requirements/xpu.txt
|
||||
# accelerate
|
||||
# compressed-tensors
|
||||
# mteb
|
||||
# sentence-transformers
|
||||
# timm
|
||||
# torchvision
|
||||
# xgrammar
|
||||
torchvision==0.27.0+xpu
|
||||
# via timm
|
||||
tqdm==4.67.3
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# datasets
|
||||
# evaluate
|
||||
# huggingface-hub
|
||||
@@ -664,13 +935,19 @@ tqdm==4.67.3
|
||||
# modelscope
|
||||
# mteb
|
||||
# nltk
|
||||
# openai
|
||||
# pqdm
|
||||
# sentence-transformers
|
||||
# transformers
|
||||
transformers==5.5.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# compressed-tensors
|
||||
# sentence-transformers
|
||||
# xgrammar
|
||||
triton==3.7.1
|
||||
# via xgrammar
|
||||
triton-xpu==3.7.1
|
||||
# via torch
|
||||
typepy==1.3.4
|
||||
@@ -680,36 +957,53 @@ typepy==1.3.4
|
||||
# tabledata
|
||||
typer==0.24.1
|
||||
# via
|
||||
# fastapi-cli
|
||||
# fastapi-cloud-cli
|
||||
# huggingface-hub
|
||||
# transformers
|
||||
typing-extensions==4.15.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
# aiosignal
|
||||
# albumentations
|
||||
# anthropic
|
||||
# anyio
|
||||
# apache-tvm-ffi
|
||||
# chz
|
||||
# fastapi
|
||||
# grpcio
|
||||
# huggingface-hub
|
||||
# librosa
|
||||
# lm-eval
|
||||
# mcp
|
||||
# mistral-common
|
||||
# mteb
|
||||
# openai
|
||||
# opentelemetry-api
|
||||
# opentelemetry-exporter-otlp-proto-grpc
|
||||
# opentelemetry-exporter-otlp-proto-http
|
||||
# opentelemetry-sdk
|
||||
# opentelemetry-semantic-conventions
|
||||
# pqdm
|
||||
# pydantic
|
||||
# pydantic-core
|
||||
# pydantic-extra-types
|
||||
# pytest-asyncio
|
||||
# referencing
|
||||
# rich-toolkit
|
||||
# schemathesis
|
||||
# sentence-transformers
|
||||
# starlette
|
||||
# torch
|
||||
# typing-inspection
|
||||
# xgrammar
|
||||
typing-inspection==0.4.2
|
||||
# via
|
||||
# fastapi
|
||||
# mcp
|
||||
# pydantic
|
||||
# pydantic-settings
|
||||
umf==1.0.3
|
||||
# via
|
||||
# intel-cmplr-lib-ur
|
||||
@@ -720,12 +1014,30 @@ urllib3==2.6.3
|
||||
# docker
|
||||
# modelscope
|
||||
# requests
|
||||
# sentry-sdk
|
||||
uvicorn==0.42.0
|
||||
# via gpt-oss
|
||||
# via
|
||||
# fastapi
|
||||
# fastapi-cli
|
||||
# fastapi-cloud-cli
|
||||
# gpt-oss
|
||||
# mcp
|
||||
uvloop==0.22.1
|
||||
# via uvicorn
|
||||
watchfiles==1.2.0
|
||||
# via
|
||||
# -r requirements/test/../common.txt
|
||||
# uvicorn
|
||||
websockets==16.0
|
||||
# via uvicorn
|
||||
werkzeug==3.1.7
|
||||
# via schemathesis
|
||||
word2number==1.1
|
||||
# via lm-eval
|
||||
xgrammar==0.2.3
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
# -r requirements/test/../common.txt
|
||||
xxhash==3.6.0
|
||||
# via
|
||||
# datasets
|
||||
|
||||
Generated
+69
-20
@@ -272,6 +272,18 @@ version = "1.1.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0"
|
||||
|
||||
[[package]]
|
||||
name = "auto_enums"
|
||||
version = "0.8.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2e4487600931c9a89f8db7ffbdf3fbdd45bb7bd85e26861f659a463cd0dff966"
|
||||
dependencies = [
|
||||
"derive_utils",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "auto_impl"
|
||||
version = "1.3.0"
|
||||
@@ -938,6 +950,17 @@ dependencies = [
|
||||
"unicode-xid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "derive_utils"
|
||||
version = "0.15.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "362f47930db19fe7735f527e6595e4900316b893ebf6d48ad3d31be928d57dd6"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "digest"
|
||||
version = "0.10.7"
|
||||
@@ -1478,9 +1501,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "h2"
|
||||
version = "0.4.13"
|
||||
version = "0.4.15"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2f44da3a8150a6703ed5d34e164b875fd14c2cdab9af1252a9a1020bde2bdc54"
|
||||
checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155"
|
||||
dependencies = [
|
||||
"atomic-waker",
|
||||
"bytes",
|
||||
@@ -1638,9 +1661,9 @@ checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9"
|
||||
|
||||
[[package]]
|
||||
name = "hyper"
|
||||
version = "1.8.1"
|
||||
version = "1.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2ab2d4f250c3d7b1c9fcdff1cece94ea4e2dfbec68614f7b87cb205f24ca9d11"
|
||||
checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498"
|
||||
dependencies = [
|
||||
"atomic-waker",
|
||||
"bytes",
|
||||
@@ -1653,7 +1676,6 @@ dependencies = [
|
||||
"httpdate",
|
||||
"itoa",
|
||||
"pin-project-lite",
|
||||
"pin-utils",
|
||||
"smallvec",
|
||||
"tokio",
|
||||
"want",
|
||||
@@ -2569,15 +2591,14 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openssl"
|
||||
version = "0.10.76"
|
||||
version = "0.10.81"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "951c002c75e16ea2c65b8c7e4d3d51d5530d8dfa7d060b4776828c88cfb18ecf"
|
||||
checksum = "77823a27f0babb03091cb9ed9ef80af3b39dbc82f97e8fa530374b7dafd87a45"
|
||||
dependencies = [
|
||||
"bitflags",
|
||||
"cfg-if",
|
||||
"foreign-types",
|
||||
"libc",
|
||||
"once_cell",
|
||||
"openssl-macros",
|
||||
"openssl-sys",
|
||||
]
|
||||
@@ -2610,9 +2631,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openssl-sys"
|
||||
version = "0.9.112"
|
||||
version = "0.9.117"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "57d55af3b3e226502be1526dfdba67ab0e9c96fc293004e79576b2b9edb0dbdb"
|
||||
checksum = "b47e7e6bb2c38cd930d25a23b40fa52e068c10e85f3e03a7f5ba5aaca5713695"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"libc",
|
||||
@@ -2783,12 +2804,6 @@ version = "0.2.17"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd"
|
||||
|
||||
[[package]]
|
||||
name = "pin-utils"
|
||||
version = "0.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184"
|
||||
|
||||
[[package]]
|
||||
name = "pkg-config"
|
||||
version = "0.3.32"
|
||||
@@ -2988,7 +3003,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "343d3bd7056eda839b03204e68deff7d1b13aba7af2b2fd16890697274262ee7"
|
||||
dependencies = [
|
||||
"heck",
|
||||
"itertools 0.10.5",
|
||||
"itertools 0.14.0",
|
||||
"log",
|
||||
"multimap",
|
||||
"petgraph",
|
||||
@@ -3009,7 +3024,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "27c6023962132f4b30eb4c172c91ce92d933da334c59c23cddee82358ddafb0b"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"itertools 0.10.5",
|
||||
"itertools 0.14.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
@@ -3503,9 +3518,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "rustls-pki-types"
|
||||
version = "1.14.0"
|
||||
version = "1.14.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "be040f8b0a225e40375822a563fa9524378b9d63112f53e19ffff34df5d33fdd"
|
||||
checksum = "30a7197ae7eb376e574fe940d068c30fe0462554a3ddbe4eca7838e049c937a9"
|
||||
dependencies = [
|
||||
"zeroize",
|
||||
]
|
||||
@@ -4385,6 +4400,22 @@ dependencies = [
|
||||
"serde_json",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tls-listener"
|
||||
version = "0.11.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1461056cc1ef47003f7ee16e4cef3741068d4c7f6b627bfce49b7c00c120a530"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"futures-util",
|
||||
"openssl",
|
||||
"pin-project-lite",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"tokio-openssl",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokenizers"
|
||||
version = "0.22.2"
|
||||
@@ -4457,6 +4488,17 @@ dependencies = [
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-openssl"
|
||||
version = "0.6.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "59df6849caa43bb7567f9a36f863c447d95a11d5903c9cc334ba32576a27eadd"
|
||||
dependencies = [
|
||||
"openssl",
|
||||
"openssl-sys",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-rustls"
|
||||
version = "0.26.4"
|
||||
@@ -5220,6 +5262,7 @@ dependencies = [
|
||||
"anyhow",
|
||||
"async-openai",
|
||||
"asynk-strim-attr",
|
||||
"auto_enums",
|
||||
"axum",
|
||||
"bytes",
|
||||
"clap",
|
||||
@@ -5227,10 +5270,13 @@ dependencies = [
|
||||
"expect-test",
|
||||
"futures",
|
||||
"http-body",
|
||||
"hyper",
|
||||
"hyper-util",
|
||||
"indexmap 2.13.0",
|
||||
"itertools 0.14.0",
|
||||
"libc",
|
||||
"llm-multimodal",
|
||||
"openssl",
|
||||
"prost",
|
||||
"prost-types",
|
||||
"rmp-serde",
|
||||
@@ -5242,8 +5288,11 @@ dependencies = [
|
||||
"sha2",
|
||||
"socket2",
|
||||
"subtle",
|
||||
"tempfile",
|
||||
"thiserror-ext",
|
||||
"tls-listener",
|
||||
"tokio",
|
||||
"tokio-openssl",
|
||||
"tokio-stream",
|
||||
"tokio-util",
|
||||
"tonic",
|
||||
|
||||
@@ -26,6 +26,7 @@ arc-swap = "1.9.0"
|
||||
async-openai = { version = "0.33.1", default-features = false, features = ["native-tls"] }
|
||||
async-trait = "0.1.89"
|
||||
asynk-strim-attr = "0.1.0"
|
||||
auto_enums = { version = "0.8.9", features = ["tokio1"] }
|
||||
axum = "0.8.8"
|
||||
base64 = "0.22.1"
|
||||
bytemuck = { version = "1.25.0", features = ["extern_crate_alloc"] }
|
||||
@@ -43,6 +44,12 @@ half = { version = "2.7.1", features = ["bytemuck"] }
|
||||
hex = "0.4.3"
|
||||
hf-hub = { version = "0.5.0", default-features = false, features = ["tokio"] }
|
||||
http-body = "1.0.1"
|
||||
hyper = { version = "1.10.1", features = ["http1", "server"] }
|
||||
hyper-util = { version = "0.1.20", features = [
|
||||
"server-graceful",
|
||||
"service",
|
||||
"tokio",
|
||||
] }
|
||||
indexmap = "2.13.0"
|
||||
itertools = "0.14.0"
|
||||
libc = "0.2.177"
|
||||
@@ -54,6 +61,7 @@ native-tls-vendored = { package = "native-tls", version = "0.2.18", features = [
|
||||
ndarray = { version = "0.16.1", features = ["serde"] }
|
||||
openai-harmony = { package = "oss-harmony", git = "https://github.com/oss-harmony/harmony", tag = "v0.0.11", default-features = false }
|
||||
openai-protocol = "1.6.0"
|
||||
openssl = "0.10"
|
||||
parking_lot = "0.12.5"
|
||||
paste = "1.0.15"
|
||||
prometheus-client = "0.24.0"
|
||||
@@ -89,6 +97,7 @@ thiserror = "2.0.16"
|
||||
thiserror-ext = "0.3.0"
|
||||
tiktoken-rs = "0.9.1"
|
||||
time = { version = "0.3.47", features = ["formatting", "local-offset", "macros"] }
|
||||
tls-listener = { version = "0.11.2", default-features = false, features = ["openssl", "tokio-net", "axum"] }
|
||||
tokenizers = "0.22.0"
|
||||
tokio = { version = "1.47.1", features = [
|
||||
"macros",
|
||||
@@ -97,6 +106,7 @@ tokio = { version = "1.47.1", features = [
|
||||
"sync",
|
||||
"time",
|
||||
] }
|
||||
tokio-openssl = "0.6"
|
||||
tokio-stream = "0.1"
|
||||
tokio-util = { version = "0.7.18", features = ["rt"] }
|
||||
tonic = "0.14.5"
|
||||
|
||||
+66
-1
@@ -25,7 +25,7 @@ use vllm_managed_engine::ManagedEngineConfig;
|
||||
use vllm_managed_engine::cli::{ManagedEngineArgs, repartition_managed_engine_args};
|
||||
use vllm_server::{
|
||||
ApiServerOptions, ChatTemplateContentFormatOption, Config, CoordinatorMode, CorsConfig,
|
||||
HttpListenerMode, ParserSelection, RendererSelection,
|
||||
DEFAULT_KEEP_ALIVE_TIMEOUT, HttpListenerMode, ParserSelection, RendererSelection, TlsConfig,
|
||||
};
|
||||
|
||||
use crate::cli::unsupported::UnsupportedArgs;
|
||||
@@ -154,6 +154,11 @@ pub struct SharedRuntimeArgs {
|
||||
#[arg(long, default_value_t = 0)]
|
||||
#[serde(default)]
|
||||
pub shutdown_timeout: u64,
|
||||
/// Maximum idle time (seconds) on a keep-alive HTTP connection before the
|
||||
/// server closes it (default 5).
|
||||
#[arg(long = "http-timeout-keep-alive", env = "VLLM_HTTP_TIMEOUT_KEEP_ALIVE")]
|
||||
#[serde(default)]
|
||||
pub http_timeout_keep_alive: Option<u64>,
|
||||
|
||||
/// The file path to the chat template, or the template in single-line form
|
||||
/// for the specified model.
|
||||
@@ -257,6 +262,34 @@ pub struct SharedRuntimeArgs {
|
||||
#[serde(default)]
|
||||
pub allow_credentials: bool,
|
||||
|
||||
/// The file path to the SSL key file. When omitted, the key is read from
|
||||
/// `--ssl-certfile` (combined PEM).
|
||||
#[arg(long)]
|
||||
#[serde(default)]
|
||||
pub ssl_keyfile: Option<String>,
|
||||
|
||||
/// The file path to the SSL cert file. Enables TLS when set.
|
||||
#[arg(long)]
|
||||
#[serde(default)]
|
||||
pub ssl_certfile: Option<String>,
|
||||
|
||||
/// The CA certificates file used to verify client certificates (mTLS).
|
||||
#[arg(long)]
|
||||
#[serde(default)]
|
||||
pub ssl_ca_certs: Option<String>,
|
||||
|
||||
/// Whether a client certificate is required: 0 = none, 1 = optional,
|
||||
/// 2 = required (mirrors Python's `ssl.CERT_*`).
|
||||
#[arg(long, default_value_t = 0, value_parser = clap::value_parser!(i32).range(0..=2))]
|
||||
#[serde(default)]
|
||||
pub ssl_cert_reqs: i32,
|
||||
|
||||
/// OpenSSL cipher string for HTTPS (TLS 1.2 and below).
|
||||
/// When unset, the linked OpenSSL's default suites are used.
|
||||
#[arg(long)]
|
||||
#[serde(default)]
|
||||
pub ssl_ciphers: Option<String>,
|
||||
|
||||
/// Unsupported Python vLLM frontend arguments recognized but not yet
|
||||
/// implemented in Rust.
|
||||
#[educe(Debug(ignore))]
|
||||
@@ -277,6 +310,13 @@ impl SharedRuntimeArgs {
|
||||
Duration::from_secs(self.shutdown_timeout)
|
||||
}
|
||||
|
||||
/// Maximum idle time on a keep-alive HTTP connection before the server
|
||||
/// closes it.
|
||||
pub fn keep_alive_timeout(&self) -> Duration {
|
||||
self.http_timeout_keep_alive
|
||||
.map_or(DEFAULT_KEEP_ALIVE_TIMEOUT, Duration::from_secs)
|
||||
}
|
||||
|
||||
/// Apply fallback logic for API key configuration from env variables.
|
||||
fn apply_env_api_key_fallback(&mut self) {
|
||||
if self.api_key.is_empty()
|
||||
@@ -301,8 +341,10 @@ impl SharedRuntimeArgs {
|
||||
) -> Config {
|
||||
let ready_timeout = self.ready_timeout();
|
||||
let shutdown_timeout = self.shutdown_timeout();
|
||||
let keep_alive_timeout = self.keep_alive_timeout();
|
||||
let api_server_options = self.api_server_options();
|
||||
let cors = self.cors_config();
|
||||
let tls = self.tls_config();
|
||||
|
||||
Config {
|
||||
transport_mode: TransportMode::Bootstrapped {
|
||||
@@ -329,10 +371,12 @@ impl SharedRuntimeArgs {
|
||||
max_logprobs: self.max_logprobs,
|
||||
api_server_options,
|
||||
cors,
|
||||
tls,
|
||||
api_keys: self.api_key,
|
||||
disable_log_stats: self.disable_log_stats,
|
||||
grpc_port: self.grpc_port,
|
||||
shutdown_timeout,
|
||||
keep_alive_timeout,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -349,8 +393,10 @@ impl SharedRuntimeArgs {
|
||||
) -> Config {
|
||||
let ready_timeout = self.ready_timeout();
|
||||
let shutdown_timeout = self.shutdown_timeout();
|
||||
let keep_alive_timeout = self.keep_alive_timeout();
|
||||
let api_server_options = self.api_server_options();
|
||||
let cors = self.cors_config();
|
||||
let tls = self.tls_config();
|
||||
|
||||
Config {
|
||||
transport_mode: TransportMode::HandshakeOwner {
|
||||
@@ -375,10 +421,12 @@ impl SharedRuntimeArgs {
|
||||
max_logprobs: self.max_logprobs,
|
||||
api_server_options,
|
||||
cors,
|
||||
tls,
|
||||
api_keys: self.api_key,
|
||||
disable_log_stats: self.disable_log_stats,
|
||||
grpc_port: self.grpc_port,
|
||||
shutdown_timeout,
|
||||
keep_alive_timeout,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -398,6 +446,23 @@ impl SharedRuntimeArgs {
|
||||
allow_credentials: self.allow_credentials,
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the TLS config: `Some` when any `ssl_*` argument is set, else
|
||||
/// `None` (plaintext). The combination is validated in [`Config::validate`].
|
||||
fn tls_config(&self) -> Option<TlsConfig> {
|
||||
let tls_requested = self.ssl_certfile.is_some()
|
||||
|| self.ssl_keyfile.is_some()
|
||||
|| self.ssl_ca_certs.is_some()
|
||||
|| self.ssl_cert_reqs != 0
|
||||
|| self.ssl_ciphers.is_some();
|
||||
tls_requested.then(|| TlsConfig {
|
||||
cert_file: self.ssl_certfile.clone(),
|
||||
key_file: self.ssl_keyfile.clone(),
|
||||
ca_certs: self.ssl_ca_certs.clone(),
|
||||
cert_reqs: self.ssl_cert_reqs,
|
||||
ciphers: self.ssl_ciphers.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn default_engine_ready_timeout_secs() -> u64 {
|
||||
|
||||
@@ -41,6 +41,7 @@ fn serve_args_forward_python_flags_with_separator() {
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
http_timeout_keep_alive: None,
|
||||
chat_template: None,
|
||||
default_chat_template_kwargs: None,
|
||||
chat_template_content_format: Auto,
|
||||
@@ -65,6 +66,11 @@ fn serve_args_forward_python_flags_with_separator() {
|
||||
],
|
||||
),
|
||||
allow_credentials: false,
|
||||
ssl_keyfile: None,
|
||||
ssl_certfile: None,
|
||||
ssl_ca_certs: None,
|
||||
ssl_cert_reqs: 0,
|
||||
ssl_ciphers: None,
|
||||
},
|
||||
managed_engine: ManagedEngineArgs {
|
||||
python: "../vllm/.venv/bin/python",
|
||||
@@ -363,6 +369,140 @@ fn serve_passes_enable_prompt_tokens_details_into_config() {
|
||||
assert!(config.api_server_options.enable_prompt_tokens_details);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_passes_tls_into_config() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--ssl-certfile",
|
||||
"/tmp/cert.pem",
|
||||
"--ssl-keyfile",
|
||||
"/tmp/key.pem",
|
||||
"--ssl-ca-certs",
|
||||
"/tmp/ca.pem",
|
||||
"--ssl-cert-reqs",
|
||||
"2",
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Serve(args) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
|
||||
let tls = config.tls.expect("tls configured");
|
||||
assert_eq!(tls.cert_file.as_deref(), Some("/tmp/cert.pem"));
|
||||
assert_eq!(tls.key_file.as_deref(), Some("/tmp/key.pem"));
|
||||
assert_eq!(tls.ca_certs.as_deref(), Some("/tmp/ca.pem"));
|
||||
assert_eq!(tls.cert_reqs, 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_without_ssl_flags_has_no_tls() {
|
||||
let cli = Cli::try_parse_from(["vllm-rs", "serve", "Qwen/Qwen3-0.6B"]).unwrap();
|
||||
|
||||
let Command::Serve(args) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
|
||||
assert!(config.tls.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_ssl_keyfile_without_certfile_fails_validation() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--ssl-keyfile",
|
||||
"/tmp/key.pem",
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Serve(args) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
|
||||
// TLS is requested (a key was given) but there is no certificate, so
|
||||
// validation fails loud rather than silently serving plaintext.
|
||||
assert_eq!(config.tls.as_ref().expect("tls requested").cert_file, None);
|
||||
let err = config.validate().unwrap_err().to_string();
|
||||
assert!(err.contains("--ssl-certfile is required"), "{err}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_mtls_without_ca_certs_fails_validation() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--ssl-certfile",
|
||||
"/tmp/cert.pem",
|
||||
"--ssl-cert-reqs",
|
||||
"2",
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Serve(args) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
|
||||
// Client-cert verification without a CA bundle has nothing to verify
|
||||
// against, so it fails loud at startup.
|
||||
let err = config.validate().unwrap_err().to_string();
|
||||
assert!(err.contains("--ssl-ca-certs is required"), "{err}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn frontend_args_json_passes_tls_into_config() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"frontend",
|
||||
"--listen-fd",
|
||||
"3",
|
||||
"--input-address",
|
||||
"ipc:///tmp/input.sock",
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","ssl_certfile":"/tmp/cert.pem","ssl_keyfile":"/tmp/key.pem"}"#,
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Frontend(args) = cli.command else {
|
||||
panic!("expected frontend args");
|
||||
};
|
||||
let config = args.into_config();
|
||||
let tls = config.tls.expect("tls configured");
|
||||
assert_eq!(tls.cert_file.as_deref(), Some("/tmp/cert.pem"));
|
||||
assert_eq!(tls.key_file.as_deref(), Some("/tmp/key.pem"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn frontend_args_json_rejects_out_of_range_cert_reqs() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"frontend",
|
||||
"--listen-fd",
|
||||
"3",
|
||||
"--input-address",
|
||||
"ipc:///tmp/input.sock",
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","ssl_certfile":"/tmp/cert.pem","ssl_cert_reqs":5}"#,
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Frontend(args) = cli.command else {
|
||||
panic!("expected frontend args");
|
||||
};
|
||||
// The JSON path bypasses clap's range check, so validate() is the only guard.
|
||||
let config = args.into_config();
|
||||
let err = config.validate().unwrap_err().to_string();
|
||||
assert!(err.contains("--ssl-cert-reqs"), "{err}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn frontend_args_json_passes_enable_request_id_headers_into_config() {
|
||||
let cli = Cli::try_parse_from([
|
||||
@@ -481,13 +621,13 @@ fn serve_args_reject_unsupported_flag_arg() {
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--ssl-keyfile",
|
||||
"/tmp/key.pem",
|
||||
"--root-path",
|
||||
"/prefix",
|
||||
])
|
||||
.unwrap_err();
|
||||
|
||||
expect![[r#"
|
||||
error: invalid value '/tmp/key.pem' for '--ssl-keyfile <SSL_KEYFILE>': argument is not implemented in Rust frontend yet
|
||||
error: invalid value '/prefix' for '--root-path <ROOT_PATH>': argument is not implemented in Rust frontend yet
|
||||
|
||||
Remove this unsupported argument to continue.
|
||||
|
||||
@@ -562,6 +702,7 @@ fn frontend_args_accept_json() {
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
http_timeout_keep_alive: None,
|
||||
chat_template: None,
|
||||
default_chat_template_kwargs: None,
|
||||
chat_template_content_format: Auto,
|
||||
@@ -586,6 +727,11 @@ fn frontend_args_accept_json() {
|
||||
],
|
||||
),
|
||||
allow_credentials: false,
|
||||
ssl_keyfile: None,
|
||||
ssl_certfile: None,
|
||||
ssl_ca_certs: None,
|
||||
ssl_cert_reqs: 0,
|
||||
ssl_ciphers: None,
|
||||
},
|
||||
},
|
||||
),
|
||||
@@ -798,14 +944,14 @@ fn frontend_args_json_rejects_unsupported_fields() {
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","ssl_keyfile":"/tmp/key.pem"}"#,
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","root_path":"/prefix"}"#,
|
||||
])
|
||||
.unwrap_err();
|
||||
|
||||
expect![[r#"
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","ssl_keyfile":"/tmp/key.pem"}' for '--args-json <JSON>':
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","root_path":"/prefix"}' for '--args-json <JSON>':
|
||||
The following arguments are not implemented in Rust frontend yet:
|
||||
- ssl_keyfile
|
||||
- root_path
|
||||
|
||||
Remove these arguments to continue.
|
||||
|
||||
@@ -825,16 +971,16 @@ fn frontend_args_json_aggregates_multiple_unsupported_fields() {
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","ssl_keyfile":"/tmp/key.pem"}"#,
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","root_path":"/prefix"}"#,
|
||||
])
|
||||
.unwrap_err();
|
||||
|
||||
let actual = error.to_string().replace(": \n", ":\n");
|
||||
expect![[r#"
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","ssl_keyfile":"/tmp/key.pem"}' for '--args-json <JSON>':
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","root_path":"/prefix"}' for '--args-json <JSON>':
|
||||
The following arguments are not implemented in Rust frontend yet:
|
||||
- response_role
|
||||
- ssl_keyfile
|
||||
- root_path
|
||||
|
||||
Remove these arguments to continue.
|
||||
|
||||
@@ -1077,6 +1223,7 @@ fn serve_args_accept_handshake_aliases() {
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
http_timeout_keep_alive: None,
|
||||
chat_template: None,
|
||||
default_chat_template_kwargs: None,
|
||||
chat_template_content_format: Auto,
|
||||
@@ -1101,6 +1248,11 @@ fn serve_args_accept_handshake_aliases() {
|
||||
],
|
||||
),
|
||||
allow_credentials: false,
|
||||
ssl_keyfile: None,
|
||||
ssl_certfile: None,
|
||||
ssl_ca_certs: None,
|
||||
ssl_cert_reqs: 0,
|
||||
ssl_ciphers: None,
|
||||
},
|
||||
managed_engine: ManagedEngineArgs {
|
||||
python: "python3",
|
||||
@@ -1234,10 +1386,12 @@ fn serve_frontend_config_uses_dp_address_as_advertised_host() {
|
||||
],
|
||||
allow_credentials: false,
|
||||
},
|
||||
tls: None,
|
||||
api_keys: [],
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0ns,
|
||||
keep_alive_timeout: 5s,
|
||||
}
|
||||
"#]]
|
||||
.assert_debug_eq(&Config {
|
||||
@@ -1315,10 +1469,12 @@ fn serve_frontend_config_keeps_tcp_transport_for_non_local_only_topology() {
|
||||
],
|
||||
allow_credentials: false,
|
||||
},
|
||||
tls: None,
|
||||
api_keys: [],
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0ns,
|
||||
keep_alive_timeout: 5s,
|
||||
}
|
||||
"#]]
|
||||
.assert_debug_eq(&config);
|
||||
@@ -1414,10 +1570,12 @@ fn frontend_config_uses_external_coordinator_when_coordinator_address_is_present
|
||||
],
|
||||
allow_credentials: false,
|
||||
},
|
||||
tls: None,
|
||||
api_keys: [],
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0ns,
|
||||
keep_alive_timeout: 5s,
|
||||
}
|
||||
"#]]
|
||||
.assert_debug_eq(&config);
|
||||
|
||||
@@ -526,18 +526,6 @@ pub struct ServerUnsupportedArgs {
|
||||
#[arg(long)]
|
||||
pub disable_access_log_for_endpoints: Option<Noop>,
|
||||
|
||||
/// The file path to the SSL key file.
|
||||
#[arg(long)]
|
||||
pub ssl_keyfile: Option<Unsupported>,
|
||||
|
||||
/// The file path to the SSL cert file.
|
||||
#[arg(long)]
|
||||
pub ssl_certfile: Option<Unsupported>,
|
||||
|
||||
/// The CA certificates file.
|
||||
#[arg(long)]
|
||||
pub ssl_ca_certs: Option<Unsupported>,
|
||||
|
||||
/// Refresh SSL Context when SSL certificate files change
|
||||
#[arg(
|
||||
long,
|
||||
@@ -547,15 +535,6 @@ pub struct ServerUnsupportedArgs {
|
||||
)]
|
||||
pub enable_ssl_refresh: Option<Unsupported>,
|
||||
|
||||
/// Whether client certificate is required (see stdlib ssl module's).
|
||||
#[arg(long)]
|
||||
pub ssl_cert_reqs: Option<Unsupported>,
|
||||
|
||||
/// SSL cipher suites for HTTPS (TLS 1.2 and below only).
|
||||
/// Example: 'ECDHE-RSA-AES256-GCM-SHA384:ECDHE-RSA-CHACHA20-POLY1305'
|
||||
#[arg(long)]
|
||||
pub ssl_ciphers: Option<Unsupported>,
|
||||
|
||||
/// FastAPI root_path when app is behind a path based routing proxy.
|
||||
#[arg(long)]
|
||||
pub root_path: Option<Unsupported>,
|
||||
|
||||
@@ -100,6 +100,7 @@ impl EngineRoutingState {
|
||||
pub struct RequestRegistry {
|
||||
closed: bool,
|
||||
requests: HashMap<String, TrackedRequest>,
|
||||
active_lora_requests: usize,
|
||||
routing_per_engine: BTreeMap<EngineId, EngineRoutingState>,
|
||||
}
|
||||
|
||||
@@ -108,6 +109,7 @@ impl RequestRegistry {
|
||||
Self {
|
||||
closed: false,
|
||||
requests: HashMap::default(),
|
||||
active_lora_requests: 0,
|
||||
routing_per_engine: engines
|
||||
.iter()
|
||||
.map(|engine| (engine.engine_id.clone(), EngineRoutingState::default()))
|
||||
@@ -133,15 +135,19 @@ impl RequestRegistry {
|
||||
|
||||
let engine_id = self.choose_engine_for_request(data_parallel_rank)?;
|
||||
let (tx, rx) = mpsc::unbounded_channel();
|
||||
let lora = lora_name.map(|adapter_name| LoraRequestState {
|
||||
adapter_name,
|
||||
phase: LoraPhase::Waiting,
|
||||
});
|
||||
if lora.is_some() {
|
||||
self.active_lora_requests += 1;
|
||||
}
|
||||
self.requests.insert(
|
||||
request_id,
|
||||
TrackedRequest {
|
||||
sender: tx,
|
||||
engine_id: engine_id.clone(),
|
||||
lora: lora_name.map(|adapter_name| LoraRequestState {
|
||||
adapter_name,
|
||||
phase: LoraPhase::Waiting,
|
||||
}),
|
||||
lora,
|
||||
},
|
||||
);
|
||||
|
||||
@@ -230,6 +236,10 @@ impl RequestRegistry {
|
||||
/// Snapshot the adapter names of tracked LoRA requests as
|
||||
/// (running, waiting) sets. Feeds the `vllm:lora_requests_info` gauge.
|
||||
pub fn lora_adapter_states(&self) -> (BTreeSet<String>, BTreeSet<String>) {
|
||||
if self.active_lora_requests == 0 {
|
||||
return (BTreeSet::new(), BTreeSet::new());
|
||||
}
|
||||
|
||||
let mut running = BTreeSet::new();
|
||||
let mut waiting = BTreeSet::new();
|
||||
for lora in self.requests.values().filter_map(|tracked| tracked.lora.as_ref()) {
|
||||
@@ -283,6 +293,7 @@ impl RequestRegistry {
|
||||
}
|
||||
|
||||
self.closed = true;
|
||||
self.active_lora_requests = 0;
|
||||
std::mem::take(&mut self.requests)
|
||||
.into_values()
|
||||
.map(|tracked| tracked.sender)
|
||||
@@ -322,6 +333,9 @@ impl RequestRegistry {
|
||||
#[must_use]
|
||||
pub fn remove(&mut self, request_id: &str) -> Option<(OutputSender, EngineId)> {
|
||||
let tracked = self.requests.remove(request_id)?;
|
||||
if tracked.lora.is_some() {
|
||||
self.active_lora_requests -= 1;
|
||||
}
|
||||
self.routing_per_engine
|
||||
.get_mut(&tracked.engine_id)
|
||||
.expect("request registry must track all known engines")
|
||||
@@ -359,6 +373,11 @@ impl RequestRegistry {
|
||||
pub fn is_closed(&self) -> bool {
|
||||
self.closed
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn active_lora_requests(&self) -> usize {
|
||||
self.active_lora_requests
|
||||
}
|
||||
}
|
||||
|
||||
/// Internal registry for tracking active utility calls and their waiting
|
||||
@@ -574,6 +593,63 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registry_counts_only_active_lora_requests() {
|
||||
let mut registry = RequestRegistry::new(&[connected_engine(EngineId::from(b"engine-0"))]);
|
||||
|
||||
registry.register("req-plain".to_string(), None, None).unwrap();
|
||||
assert_eq!(registry.active_lora_requests(), 0);
|
||||
assert_eq!(
|
||||
registry.lora_adapter_states(),
|
||||
(adapter_names(&[]), adapter_names(&[]))
|
||||
);
|
||||
|
||||
registry
|
||||
.register(
|
||||
"req-lora-a".to_string(),
|
||||
Some("adapter-a".to_string()),
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
registry
|
||||
.register(
|
||||
"req-lora-b".to_string(),
|
||||
Some("adapter-b".to_string()),
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(registry.active_lora_requests(), 2);
|
||||
|
||||
drop(registry.remove("req-plain"));
|
||||
assert_eq!(registry.active_lora_requests(), 2);
|
||||
|
||||
drop(registry.finish_many(&["req-lora-a".to_string()]));
|
||||
assert_eq!(registry.active_lora_requests(), 1);
|
||||
|
||||
drop(registry.abort_many(&["req-lora-b".to_string()], 0.0));
|
||||
assert_eq!(registry.active_lora_requests(), 0);
|
||||
assert_eq!(
|
||||
registry.lora_adapter_states(),
|
||||
(adapter_names(&[]), adapter_names(&[]))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registry_clears_lora_count_on_close() {
|
||||
let mut registry = RequestRegistry::new(&[connected_engine(EngineId::from(b"engine-0"))]);
|
||||
registry
|
||||
.register("req-lora".to_string(), Some("adapter-a".to_string()), None)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(registry.active_lora_requests(), 1);
|
||||
drop(registry.close());
|
||||
assert_eq!(registry.active_lora_requests(), 0);
|
||||
assert_eq!(
|
||||
registry.lora_adapter_states(),
|
||||
(adapter_names(&[]), adapter_names(&[]))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registry_drops_lora_tracking_on_abort() {
|
||||
let mut registry = RequestRegistry::new(&[connected_engine(EngineId::from(b"engine-0"))]);
|
||||
|
||||
@@ -20,6 +20,27 @@ pub(crate) struct CoordinatorStateSnapshot {
|
||||
pub engines_running: bool,
|
||||
}
|
||||
|
||||
impl CoordinatorStateSnapshot {
|
||||
/// Resume the engines for a `FirstRequest` and return the wave to broadcast
|
||||
/// and the engine to exclude from the wakeup.
|
||||
///
|
||||
/// The request may have been stamped with a `request_wave` older than
|
||||
/// `current_wave` if a `WaveComplete` advanced it after the command was
|
||||
/// enqueued. Such a request still needs serving, so the current wave is
|
||||
/// broadcast to every engine (`exclude = None`); the wave is never rewound.
|
||||
/// A non-stale request excludes the engine that already received it. Mirrors
|
||||
/// the Python coordinator's front-end path.
|
||||
pub(crate) fn start_wave_for_first_request(
|
||||
&mut self,
|
||||
request_wave: u32,
|
||||
target_engine_index: u32,
|
||||
) -> (u32, Option<u32>) {
|
||||
self.engines_running = true;
|
||||
let exclude = (request_wave >= self.current_wave).then_some(target_engine_index);
|
||||
(self.current_wave, exclude)
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared in-process coordinator state.
|
||||
pub(crate) type CoordinatorState = Mutex<CoordinatorStateSnapshot>;
|
||||
|
||||
|
||||
@@ -27,9 +27,10 @@ use crate::protocol::{
|
||||
struct StartDpWaveMessage {
|
||||
/// DP wave number that all engines should start processing.
|
||||
wave: u32,
|
||||
/// Engine index that already received the triggering request and should not
|
||||
/// receive an extra wakeup notification.
|
||||
exclude_engine_index: u32,
|
||||
/// Engine index that already received the triggering request and so does not
|
||||
/// need an extra wakeup. `None` wakes every engine (used when the triggering
|
||||
/// request was for a stale wave).
|
||||
exclude_engine_index: Option<u32>,
|
||||
}
|
||||
|
||||
/// Background half of the in-process coordinator.
|
||||
@@ -57,7 +58,11 @@ impl InProcCoordinatorRunner {
|
||||
}
|
||||
|
||||
/// Broadcast Python-compatible `START_DP_WAVE` to all connected engines.
|
||||
async fn broadcast_start_wave(&mut self, wave: u32, exclude_engine_index: u32) -> Result<()> {
|
||||
async fn broadcast_start_wave(
|
||||
&mut self,
|
||||
wave: u32,
|
||||
exclude_engine_index: Option<u32>,
|
||||
) -> Result<()> {
|
||||
let payload = encode_msgpack(&StartDpWaveMessage {
|
||||
wave,
|
||||
exclude_engine_index,
|
||||
@@ -86,13 +91,17 @@ impl InProcCoordinatorRunner {
|
||||
engine_id: target_engine_id.to_vec(),
|
||||
}
|
||||
})?;
|
||||
self.state.lock().current_wave = wave;
|
||||
let (current_wave, exclude) = {
|
||||
let mut state = self.state.lock();
|
||||
state.start_wave_for_first_request(wave, target_engine_index)
|
||||
};
|
||||
debug!(
|
||||
wave,
|
||||
exclude_engine_index = target_engine_index,
|
||||
current_wave,
|
||||
request_wave = wave,
|
||||
?exclude,
|
||||
"starting DP wave after first request while engines were paused"
|
||||
);
|
||||
self.broadcast_start_wave(wave, target_engine_index).await?;
|
||||
self.broadcast_start_wave(current_wave, exclude).await?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
@@ -150,7 +159,7 @@ impl InProcCoordinatorRunner {
|
||||
exclude_engine_index = engine_index,
|
||||
"starting DP wave after stale-wave notification from engine"
|
||||
);
|
||||
self.broadcast_start_wave(wave, engine_index).await?;
|
||||
self.broadcast_start_wave(wave, Some(engine_index)).await?;
|
||||
}
|
||||
}
|
||||
},
|
||||
@@ -202,3 +211,48 @@ impl InProcCoordinatorRunner {
|
||||
inner.close_registries(Arc::new(error));
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::coordinator::handle::CoordinatorStateSnapshot;
|
||||
|
||||
/// A `FirstRequest` for the current wave starts that wave and excludes the
|
||||
/// engine that already received the triggering request.
|
||||
#[test]
|
||||
fn first_request_for_current_wave_excludes_target() {
|
||||
let mut state = CoordinatorStateSnapshot {
|
||||
current_wave: 3,
|
||||
engines_running: false,
|
||||
};
|
||||
|
||||
let (wave, exclude) = state.start_wave_for_first_request(3, 2);
|
||||
|
||||
assert_eq!(wave, 3);
|
||||
assert_eq!(exclude, Some(2));
|
||||
assert!(state.engines_running);
|
||||
assert_eq!(state.current_wave, 3);
|
||||
}
|
||||
|
||||
/// A `FirstRequest` whose wave was superseded by a racing `WaveComplete`
|
||||
/// (`request_wave < current_wave`) must still start the request's wave: it
|
||||
/// broadcasts the current wave and wakes every engine (`exclude = None`)
|
||||
/// rather than rewinding the wave or dropping the request.
|
||||
#[test]
|
||||
fn stale_first_request_starts_current_wave_for_all_engines() {
|
||||
let mut state = CoordinatorStateSnapshot {
|
||||
current_wave: 4,
|
||||
engines_running: false,
|
||||
};
|
||||
|
||||
// Request stamped with wave 3 while the coordinator already advanced to 4.
|
||||
let (wave, exclude) = state.start_wave_for_first_request(3, 2);
|
||||
|
||||
assert_eq!(
|
||||
wave, 4,
|
||||
"must broadcast the current wave, not the stale one"
|
||||
);
|
||||
assert_eq!(exclude, None, "a stale request must wake every engine");
|
||||
assert!(state.engines_running);
|
||||
assert_eq!(state.current_wave, 4, "wave must not be rewound");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1285,18 +1285,24 @@ async fn dropping_multiple_live_streams_aborts_all_in_a_burst() {
|
||||
)
|
||||
.await;
|
||||
|
||||
let abort =
|
||||
timeout(Duration::from_secs(1), recv_engine_message(dealer)).await.unwrap();
|
||||
assert_eq!(abort[0].as_ref(), &[0x01]);
|
||||
let ids: Vec<String> = rmp_serde::from_slice(&abort[1]).unwrap();
|
||||
// Aborts may coalesce into one burst or split across several.
|
||||
let mut aborted = BTreeSet::new();
|
||||
while aborted.len() < 3 {
|
||||
let abort =
|
||||
timeout(Duration::from_secs(1), recv_engine_message(dealer)).await.unwrap();
|
||||
assert_eq!(abort[0].as_ref(), &[0x01]);
|
||||
let ids: Vec<String> = rmp_serde::from_slice(&abort[1]).unwrap();
|
||||
aborted.extend(ids);
|
||||
}
|
||||
assert_eq!(
|
||||
ids,
|
||||
vec![
|
||||
aborted,
|
||||
BTreeSet::from([
|
||||
"req-1".to_string(),
|
||||
"req-2".to_string(),
|
||||
"req-3".to_string()
|
||||
]
|
||||
])
|
||||
);
|
||||
// No spurious extra aborts.
|
||||
assert!(
|
||||
timeout(Duration::from_millis(100), recv_engine_message(dealer)).await.is_err()
|
||||
);
|
||||
|
||||
@@ -7,14 +7,18 @@ license.workspace = true
|
||||
[dependencies]
|
||||
anyhow.workspace = true
|
||||
asynk-strim-attr.workspace = true
|
||||
auto_enums.workspace = true
|
||||
axum.workspace = true
|
||||
educe.workspace = true
|
||||
futures.workspace = true
|
||||
http-body.workspace = true
|
||||
hyper.workspace = true
|
||||
hyper-util.workspace = true
|
||||
indexmap.workspace = true
|
||||
itertools.workspace = true
|
||||
libc.workspace = true
|
||||
llm-multimodal.workspace = true
|
||||
openssl.workspace = true
|
||||
prost.workspace = true
|
||||
prost-types.workspace = true
|
||||
rmpv.workspace = true
|
||||
@@ -25,7 +29,9 @@ sha2.workspace = true
|
||||
socket2.workspace = true
|
||||
subtle.workspace = true
|
||||
thiserror-ext.workspace = true
|
||||
tls-listener.workspace = true
|
||||
tokio.workspace = true
|
||||
tokio-openssl.workspace = true
|
||||
tokio-stream.workspace = true
|
||||
tokio-util.workspace = true
|
||||
tonic.workspace = true
|
||||
@@ -54,6 +60,7 @@ clap.workspace = true
|
||||
expect-test.workspace = true
|
||||
rmp-serde.workspace = true
|
||||
serial_test.workspace = true
|
||||
tempfile.workspace = true
|
||||
tower.workspace = true
|
||||
vllm-engine-core-client = { workspace = true, features = ["test-util"] }
|
||||
zeromq.workspace = true
|
||||
|
||||
@@ -71,10 +71,12 @@ async fn main() -> Result<()> {
|
||||
max_logprobs: None,
|
||||
api_server_options: ApiServerOptions::default(),
|
||||
cors: CorsConfig::default(),
|
||||
tls: None,
|
||||
api_keys: Vec::new(),
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: Duration::ZERO,
|
||||
keep_alive_timeout: Duration::from_secs(5),
|
||||
};
|
||||
|
||||
let bind_address = format!("127.0.0.1:{port}");
|
||||
|
||||
@@ -10,6 +10,10 @@ use serde_json::Value;
|
||||
use vllm_chat::{ChatTemplateContentFormatOption, ParserSelection, RendererSelection};
|
||||
use vllm_engine_core_client::{CoordinatorMode as EngineCoreCoordinatorMode, TransportMode};
|
||||
|
||||
/// Default keep-alive idle timeout (seconds); also the head-read bound
|
||||
/// when keep-alive is disabled (`0`).
|
||||
pub const DEFAULT_KEEP_ALIVE_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
|
||||
/// How the HTTP server obtains its listening socket.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
|
||||
pub enum HttpListenerMode {
|
||||
@@ -99,6 +103,54 @@ impl CorsConfig {
|
||||
}
|
||||
}
|
||||
|
||||
/// TLS settings mirroring Python's uvicorn `ssl_*` arguments.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
|
||||
pub struct TlsConfig {
|
||||
/// PEM certificate chain file. Required when TLS is configured; may also
|
||||
/// hold the private key (combined PEM) when `key_file` is unset.
|
||||
pub cert_file: Option<String>,
|
||||
/// PEM private key file. When `None`, the key is read from `cert_file`
|
||||
/// (combined PEM).
|
||||
pub key_file: Option<String>,
|
||||
/// PEM CA bundle used to verify client certificates (mTLS). Required when
|
||||
/// `cert_reqs` is non-zero.
|
||||
pub ca_certs: Option<String>,
|
||||
/// Client-certificate requirement, mirroring Python's `ssl.CERT_*`:
|
||||
/// 0 = none, 1 = optional, 2 = required.
|
||||
pub cert_reqs: i32,
|
||||
/// OpenSSL cipher string for TLS 1.2 and below, mirroring Python's
|
||||
/// `ssl.set_ciphers`. `None` keeps the forward-secret AEAD default.
|
||||
pub ciphers: Option<String>,
|
||||
}
|
||||
|
||||
impl TlsConfig {
|
||||
/// Structurally validate the TLS arguments; the cert/key material is parsed
|
||||
/// later, when the OpenSSL context is built.
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
if self.cert_file.is_none() {
|
||||
bail!(
|
||||
"--ssl-certfile is required to enable TLS; \
|
||||
--ssl-keyfile/--ssl-ca-certs/--ssl-cert-reqs/--ssl-ciphers \
|
||||
cannot be used without it"
|
||||
);
|
||||
}
|
||||
if !matches!(self.cert_reqs, 0..=2) {
|
||||
bail!(
|
||||
"--ssl-cert-reqs must be 0 (none), 1 (optional), or 2 (required), got {}",
|
||||
self.cert_reqs
|
||||
);
|
||||
}
|
||||
if self.cert_reqs != 0 && self.ca_certs.is_none() {
|
||||
bail!(
|
||||
"--ssl-ca-certs is required when --ssl-cert-reqs is {} \
|
||||
(client certificate verification)",
|
||||
self.cert_reqs
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Normalized runtime configuration for the minimal OpenAI-compatible server.
|
||||
#[derive(Educe, Clone, PartialEq, Eq, Serialize)]
|
||||
#[educe(Debug)]
|
||||
@@ -138,6 +190,9 @@ pub struct Config {
|
||||
pub api_server_options: ApiServerOptions,
|
||||
/// CORS settings applied to every HTTP response.
|
||||
pub cors: CorsConfig,
|
||||
/// TLS settings. `None` serves plaintext HTTP; `Some` terminates TLS at the
|
||||
/// listener.
|
||||
pub tls: Option<TlsConfig>,
|
||||
/// API keys accepted as bearer tokens for guarded routes.
|
||||
#[serde(skip_serializing)]
|
||||
#[educe(Debug(method(fmt_redacted_api_keys)))]
|
||||
@@ -150,6 +205,9 @@ pub struct Config {
|
||||
pub grpc_port: Option<u16>,
|
||||
/// Maximum time to wait for active HTTP/gRPC requests to drain on shutdown.
|
||||
pub shutdown_timeout: Duration,
|
||||
/// Maximum idle time on a keep-alive HTTP connection before the server
|
||||
/// closes it (`VLLM_HTTP_TIMEOUT_KEEP_ALIVE`, default 5s).
|
||||
pub keep_alive_timeout: Duration,
|
||||
}
|
||||
|
||||
impl Config {
|
||||
@@ -158,6 +216,9 @@ impl Config {
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
vllm_chat::validate_parser_overrides(&self.tool_call_parser, &self.reasoning_parser)?;
|
||||
self.cors.validate()?;
|
||||
if let Some(tls) = &self.tls {
|
||||
tls.validate()?;
|
||||
}
|
||||
if let Some(max_logprobs) = self.max_logprobs
|
||||
&& max_logprobs < -1
|
||||
{
|
||||
|
||||
@@ -4,16 +4,21 @@ mod convert;
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use futures::{Stream, StreamExt as _};
|
||||
use futures::{Stream, StreamExt as _, stream};
|
||||
use thiserror_ext::AsReport as _;
|
||||
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_openssl::SslStream;
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use tonic::transport::server::{Connected, TcpConnectInfo};
|
||||
use tonic::{Request, Response, Status};
|
||||
use tracing::info;
|
||||
use vllm_text::{DecodedTextEvent, TextOutputStreamExt as _};
|
||||
|
||||
use self::convert::ResponseOpts;
|
||||
use crate::listener::{Listener, ListenerIo};
|
||||
use crate::state::AppState;
|
||||
|
||||
/// Generated protobuf/gRPC types for the `vllm` package.
|
||||
@@ -26,6 +31,78 @@ pub use pb::generate_server::GenerateServer;
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
/// Newtype over `tokio-openssl`'s `SslStream` so we can implement tonic's
|
||||
/// [`Connected`] on it (the orphan rule blocks doing so on the foreign type).
|
||||
pub(crate) struct GrpcTlsStream {
|
||||
inner: SslStream<ListenerIo>,
|
||||
}
|
||||
|
||||
impl GrpcTlsStream {
|
||||
pub(crate) fn new(inner: SslStream<ListenerIo>) -> Self {
|
||||
Self { inner }
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncRead for GrpcTlsStream {
|
||||
fn poll_read(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &mut ReadBuf<'_>,
|
||||
) -> Poll<std::io::Result<()>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_read(cx, buf)
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncWrite for GrpcTlsStream {
|
||||
fn poll_write(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<std::io::Result<usize>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_write(cx, buf)
|
||||
}
|
||||
|
||||
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_flush(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
Pin::new(&mut self.get_mut().inner).poll_shutdown(cx)
|
||||
}
|
||||
}
|
||||
|
||||
impl Connected for GrpcTlsStream {
|
||||
type ConnectInfo = TcpConnectInfo;
|
||||
|
||||
fn connect_info(&self) -> TcpConnectInfo {
|
||||
self.inner.get_ref().connect_info()
|
||||
}
|
||||
}
|
||||
|
||||
/// Adapt the shared server listener into tonic's incoming stream shape.
|
||||
pub(crate) fn incoming(listener: Listener) -> impl Stream<Item = std::io::Result<ListenerIo>> {
|
||||
stream::unfold(listener, |mut listener| async move {
|
||||
let (io, _) = axum::serve::Listener::accept(&mut listener).await;
|
||||
Some((Ok(io), listener))
|
||||
})
|
||||
}
|
||||
|
||||
/// Wrap the gRPC listener so each accepted connection completes a TLS handshake
|
||||
/// before tonic serves it.
|
||||
pub(crate) fn tls_incoming(
|
||||
listener: Listener,
|
||||
context: openssl::ssl::SslContext,
|
||||
handshake_timeout: std::time::Duration,
|
||||
) -> impl Stream<Item = std::io::Result<GrpcTlsStream>> {
|
||||
tls_listener::builder(context)
|
||||
.handshake_timeout(handshake_timeout)
|
||||
.listen(listener)
|
||||
.map(|res| {
|
||||
res.map(|(inner, _addr)| GrpcTlsStream::new(inner))
|
||||
.map_err(std::io::Error::other)
|
||||
})
|
||||
}
|
||||
|
||||
/// gRPC Generate service implementation backed by the shared application state.
|
||||
pub struct GenerateServiceImpl {
|
||||
state: Arc<AppState>,
|
||||
|
||||
@@ -1,11 +1,19 @@
|
||||
use std::future::Future;
|
||||
use std::io;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
|
||||
use futures::StreamExt as _;
|
||||
use hyper_util::rt::TokioIo;
|
||||
use openssl::ssl::{SslConnector, SslFiletype, SslMethod};
|
||||
use serial_test::serial;
|
||||
use tonic::transport::Server as TonicServer;
|
||||
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_openssl::SslStream;
|
||||
use tonic::transport::{Channel, Endpoint, Server as TonicServer, Uri};
|
||||
use tower::service_fn;
|
||||
use vllm_chat::{
|
||||
ChatBackend, ChatLlm, ChatRenderer, ChatRequest, ChatTextBackend, DefaultChatOutputProcessor,
|
||||
DynChatOutputProcessor, DynChatRenderer, NewChatOutputProcessorOptions, RenderedPrompt,
|
||||
@@ -22,8 +30,11 @@ use zeromq::prelude::{SocketRecv, SocketSend};
|
||||
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
|
||||
|
||||
use super::pb::generate_client::GenerateClient;
|
||||
use super::{GenerateServer, GenerateServiceImpl, pb};
|
||||
use super::{GenerateServer, GenerateServiceImpl, incoming, pb, tls_incoming};
|
||||
use crate::listener::Listener;
|
||||
use crate::state::AppState;
|
||||
use crate::tls;
|
||||
use crate::tls_tests::{TestCerts, server_tls};
|
||||
|
||||
// ========================================================================================
|
||||
// Helpers (mirrors the patterns in routes/tests.rs)
|
||||
@@ -211,17 +222,12 @@ impl ChatRenderer for FakeTextBackend {
|
||||
}
|
||||
}
|
||||
|
||||
/// Spin up a gRPC server backed by a mock engine that serves a single request
|
||||
/// with the given output specs. Returns the client, the gRPC server task, and
|
||||
/// the mock engine task.
|
||||
async fn grpc_test_server(
|
||||
/// Build the gRPC service + mock engine that serves a single request with the
|
||||
/// given output specs. Shared by the plaintext and TLS server fixtures.
|
||||
async fn setup_grpc_service(
|
||||
engine_id: impl Into<EngineId>,
|
||||
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
|
||||
) -> (
|
||||
GenerateClient<tonic::transport::Channel>,
|
||||
tokio::task::JoinHandle<()>,
|
||||
MockEngineTask,
|
||||
) {
|
||||
) -> (GenerateServer<GenerateServiceImpl>, MockEngineTask) {
|
||||
let ipc = IpcNamespace::new().expect("create ipc namespace");
|
||||
let handshake_address = ipc.handshake_endpoint();
|
||||
let engine_id = engine_id.into();
|
||||
@@ -259,14 +265,29 @@ async fn grpc_test_server(
|
||||
Arc::new(FakeTextBackend) as Arc<dyn ChatTextBackend>,
|
||||
);
|
||||
let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat));
|
||||
let svc = GenerateServer::new(GenerateServiceImpl::new(state));
|
||||
(
|
||||
GenerateServer::new(GenerateServiceImpl::new(state)),
|
||||
engine_task,
|
||||
)
|
||||
}
|
||||
|
||||
/// Spin up a plaintext gRPC server backed by a mock engine. Returns the client,
|
||||
/// the gRPC server task, and the mock engine task.
|
||||
async fn grpc_test_server(
|
||||
engine_id: impl Into<EngineId>,
|
||||
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
|
||||
) -> (
|
||||
GenerateClient<tonic::transport::Channel>,
|
||||
tokio::task::JoinHandle<()>,
|
||||
MockEngineTask,
|
||||
) {
|
||||
let (svc, engine_task) = setup_grpc_service(engine_id, output_specs).await;
|
||||
|
||||
// Bind to an OS-assigned port.
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
|
||||
let addr = listener.local_addr().expect("local addr");
|
||||
|
||||
let server_task = tokio::spawn(async move {
|
||||
let incoming = tokio_stream::wrappers::TcpListenerStream::new(listener);
|
||||
let incoming = incoming(Listener::Tcp(listener));
|
||||
TonicServer::builder()
|
||||
.add_service(svc)
|
||||
.serve_with_incoming(incoming)
|
||||
@@ -274,7 +295,6 @@ async fn grpc_test_server(
|
||||
.expect("grpc server");
|
||||
});
|
||||
|
||||
// Connect the client.
|
||||
let grpc_client = GenerateClient::connect(format!("http://{addr}"))
|
||||
.await
|
||||
.expect("connect grpc client");
|
||||
@@ -282,6 +302,158 @@ async fn grpc_test_server(
|
||||
(grpc_client, server_task, engine_task)
|
||||
}
|
||||
|
||||
/// Spin up a TLS gRPC server (server cert from `certs`, `cert_reqs` mTLS mode).
|
||||
/// Returns the address, the server task, and the mock engine task.
|
||||
async fn grpc_tls_test_server(
|
||||
engine_id: impl Into<EngineId>,
|
||||
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
|
||||
certs: &TestCerts,
|
||||
cert_reqs: i32,
|
||||
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
|
||||
let (svc, engine_task) = setup_grpc_service(engine_id, output_specs).await;
|
||||
let context = tls::build_grpc_server_config(&server_tls(certs, cert_reqs))
|
||||
.expect("build grpc tls config");
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
|
||||
let addr = listener.local_addr().expect("local addr").to_string();
|
||||
|
||||
let server_task = tokio::spawn(async move {
|
||||
let incoming = tls_incoming(Listener::Tcp(listener), context, tls::TLS_HANDSHAKE_TIMEOUT);
|
||||
TonicServer::builder()
|
||||
.add_service(svc)
|
||||
.serve_with_incoming(incoming)
|
||||
.await
|
||||
.expect("grpc tls server");
|
||||
});
|
||||
|
||||
(addr, server_task, engine_task)
|
||||
}
|
||||
|
||||
/// Build a tonic `Generate` client over a tokio-openssl connector, optionally
|
||||
/// with a client identity for mTLS. Hand-rolled because tonic 0.14 ships no
|
||||
/// OpenSSL transport.
|
||||
async fn grpc_tls_client(
|
||||
certs: &TestCerts,
|
||||
addr: &str,
|
||||
identity: Option<&str>,
|
||||
) -> Result<GenerateClient<Channel>, tonic::transport::Error> {
|
||||
let ca = certs.path("ca.pem");
|
||||
let identity = identity.map(|name| {
|
||||
(
|
||||
certs.path(&format!("{name}.pem")),
|
||||
certs.path(&format!("{name}.key")),
|
||||
)
|
||||
});
|
||||
let target = addr.to_string();
|
||||
|
||||
let connector = service_fn(move |_: Uri| {
|
||||
let ca = ca.clone();
|
||||
let identity = identity.clone();
|
||||
let target = target.clone();
|
||||
async move {
|
||||
let tcp = TcpStream::connect(&target).await?;
|
||||
let mut builder =
|
||||
SslConnector::builder(SslMethod::tls_client()).map_err(io::Error::other)?;
|
||||
builder.set_ca_file(&ca).map_err(io::Error::other)?;
|
||||
if let Some((cert, key)) = &identity {
|
||||
builder.set_certificate_chain_file(cert).map_err(io::Error::other)?;
|
||||
builder.set_private_key_file(key, SslFiletype::PEM).map_err(io::Error::other)?;
|
||||
}
|
||||
let mut config = builder.build().configure().map_err(io::Error::other)?;
|
||||
config.set_verify_hostname(false);
|
||||
config.set_alpn_protos(b"\x02h2").map_err(io::Error::other)?;
|
||||
let ssl = config.into_ssl("127.0.0.1").map_err(io::Error::other)?;
|
||||
let mut stream = SslStream::new(ssl, tcp).map_err(io::Error::other)?;
|
||||
Pin::new(&mut stream).connect().await.map_err(io::Error::other)?;
|
||||
Ok::<_, io::Error>(TokioIo::new(stream))
|
||||
}
|
||||
});
|
||||
|
||||
let channel = Endpoint::from_shared(format!("https://{addr}"))
|
||||
.expect("grpc endpoint")
|
||||
.connect_with_connector(connector)
|
||||
.await?;
|
||||
Ok(GenerateClient::new(channel))
|
||||
}
|
||||
|
||||
/// Complete a raw TLS handshake against the gRPC port (offering ALPN `h2`) for
|
||||
/// the ALPN-negotiation assertion.
|
||||
async fn grpc_tls_handshake(
|
||||
certs: &TestCerts,
|
||||
addr: &str,
|
||||
) -> io::Result<Pin<Box<SslStream<TcpStream>>>> {
|
||||
let tcp = TcpStream::connect(addr).await?;
|
||||
let mut builder = SslConnector::builder(SslMethod::tls_client()).map_err(io::Error::other)?;
|
||||
builder.set_ca_file(certs.path("ca.pem")).map_err(io::Error::other)?;
|
||||
let mut config = builder.build().configure().map_err(io::Error::other)?;
|
||||
config.set_verify_hostname(false);
|
||||
config.set_alpn_protos(b"\x02h2").map_err(io::Error::other)?;
|
||||
let ssl = config.into_ssl("127.0.0.1").map_err(io::Error::other)?;
|
||||
let mut stream = Box::pin(SslStream::new(ssl, tcp).map_err(io::Error::other)?);
|
||||
stream.as_mut().connect().await.map_err(io::Error::other)?;
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
/// Spin up a plaintext gRPC server, optionally with HTTP/2 keepalive set to
|
||||
/// `keepalive` for both the PING interval and the unanswered-PING timeout.
|
||||
async fn grpc_server_with_keepalive(
|
||||
engine_id: impl Into<EngineId>,
|
||||
keepalive: Option<Duration>,
|
||||
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
|
||||
let (svc, engine_task) = setup_grpc_service(engine_id, default_stream_output_specs()).await;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
|
||||
let addr = listener.local_addr().expect("local addr").to_string();
|
||||
|
||||
let mut builder = TonicServer::builder();
|
||||
if let Some(interval) = keepalive {
|
||||
builder = builder
|
||||
.http2_keepalive_interval(Some(interval))
|
||||
.http2_keepalive_timeout(Some(interval));
|
||||
}
|
||||
|
||||
let server_task = tokio::spawn(async move {
|
||||
let incoming = incoming(Listener::Tcp(listener));
|
||||
builder
|
||||
.add_service(svc)
|
||||
.serve_with_incoming(incoming)
|
||||
.await
|
||||
.expect("grpc server");
|
||||
});
|
||||
|
||||
(addr, server_task, engine_task)
|
||||
}
|
||||
|
||||
/// Establish an HTTP/2 connection (preface + SETTINGS exchange) then go silent,
|
||||
/// ACKing the server's SETTINGS but never its keepalive PINGs. Returns whether
|
||||
/// the SERVER closes the connection within `wait`. A minimal hand-rolled h2 peer
|
||||
/// because a real client auto-ACKs PINGs and so can never be kept-alive-evicted.
|
||||
async fn h2_unresponsive_peer_closed_within(addr: &str, wait: Duration) -> bool {
|
||||
let mut tcp = TcpStream::connect(addr).await.expect("connect");
|
||||
tcp.write_all(b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n").await.expect("preface");
|
||||
tcp.write_all(&[0, 0, 0, 0x4, 0, 0, 0, 0, 0]).await.expect("client settings");
|
||||
|
||||
let closed = tokio::time::timeout(wait, async {
|
||||
let mut header = [0u8; 9];
|
||||
while tcp.read_exact(&mut header).await.is_ok() {
|
||||
let len = u32::from_be_bytes([0, header[0], header[1], header[2]]) as usize;
|
||||
let frame_type = header[3];
|
||||
let flags = header[4];
|
||||
let mut payload = vec![0u8; len];
|
||||
if tcp.read_exact(&mut payload).await.is_err() {
|
||||
return;
|
||||
}
|
||||
// ACK the server's SETTINGS so the only thing left unanswered is PINGs.
|
||||
if frame_type == 0x4 && flags & 0x1 == 0 {
|
||||
let _ = tcp.write_all(&[0, 0, 0, 0x4, 0x1, 0, 0, 0, 0]).await;
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
|
||||
closed.is_ok()
|
||||
}
|
||||
|
||||
// ========================================================================================
|
||||
// Tests
|
||||
// ========================================================================================
|
||||
@@ -720,3 +892,173 @@ async fn unary_generate_output_text_defaults_to_true() {
|
||||
engine_task.await.expect("mock engine task");
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_generate_succeeds_over_tls() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, server_task, engine_task) = grpc_tls_test_server(
|
||||
b"engine-grpc-tls-unary",
|
||||
default_stream_output_specs(),
|
||||
&certs,
|
||||
0,
|
||||
)
|
||||
.await;
|
||||
|
||||
let mut client = grpc_tls_client(&certs, &addr, None).await.expect("tls client");
|
||||
let response = client
|
||||
.generate(pb::GenerateRequest {
|
||||
request_id: "test-tls-unary".to_string(),
|
||||
model: "test-model".to_string(),
|
||||
prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())),
|
||||
stopping: Some(pb::StoppingCriteria {
|
||||
max_new_tokens: 10,
|
||||
..Default::default()
|
||||
}),
|
||||
response: Some(pb::ResponseOptions {
|
||||
output_text: Some(true),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.expect("unary generate over tls")
|
||||
.into_inner();
|
||||
|
||||
assert_eq!(response.outputs.expect("outputs present").text, "hi");
|
||||
|
||||
engine_task.await.expect("mock engine task");
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_tls_negotiates_h2_alpn() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, server_task, _engine_task) = grpc_tls_test_server(
|
||||
b"engine-grpc-tls-alpn",
|
||||
default_stream_output_specs(),
|
||||
&certs,
|
||||
0,
|
||||
)
|
||||
.await;
|
||||
|
||||
let stream = grpc_tls_handshake(&certs, &addr).await.expect("handshake");
|
||||
assert_eq!(
|
||||
stream.ssl().selected_alpn_protocol(),
|
||||
Some(&b"h2"[..]),
|
||||
"server must negotiate h2 ALPN"
|
||||
);
|
||||
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_mtls_required_rejects_client_without_certificate() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, server_task, _engine_task) = grpc_tls_test_server(
|
||||
b"engine-grpc-tls-mtls-reject",
|
||||
default_stream_output_specs(),
|
||||
&certs,
|
||||
2,
|
||||
)
|
||||
.await;
|
||||
|
||||
// With TLS 1.3 the missing-client-cert rejection surfaces on first use, not
|
||||
// at the handshake, so drive an RPC and assert the call fails.
|
||||
let outcome = match grpc_tls_client(&certs, &addr, None).await {
|
||||
Err(_) => Err(()),
|
||||
Ok(mut client) => client
|
||||
.generate(pb::GenerateRequest {
|
||||
request_id: "test-tls-mtls-reject".to_string(),
|
||||
model: "test-model".to_string(),
|
||||
prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())),
|
||||
stopping: Some(pb::StoppingCriteria {
|
||||
max_new_tokens: 10,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.map(|_| ())
|
||||
.map_err(|_| ()),
|
||||
};
|
||||
assert!(
|
||||
outcome.is_err(),
|
||||
"mTLS-required gRPC must reject a client without a certificate"
|
||||
);
|
||||
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_mtls_required_accepts_valid_client_certificate() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, server_task, engine_task) = grpc_tls_test_server(
|
||||
b"engine-grpc-tls-mtls-accept",
|
||||
default_stream_output_specs(),
|
||||
&certs,
|
||||
2,
|
||||
)
|
||||
.await;
|
||||
|
||||
let mut client = grpc_tls_client(&certs, &addr, Some("client")).await.expect("mtls client");
|
||||
let response = client
|
||||
.generate(pb::GenerateRequest {
|
||||
request_id: "test-tls-mtls".to_string(),
|
||||
model: "test-model".to_string(),
|
||||
prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())),
|
||||
stopping: Some(pb::StoppingCriteria {
|
||||
max_new_tokens: 10,
|
||||
..Default::default()
|
||||
}),
|
||||
response: Some(pb::ResponseOptions {
|
||||
output_text: Some(true),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.expect("mtls generate over tls")
|
||||
.into_inner();
|
||||
|
||||
assert_eq!(response.outputs.expect("outputs present").text, "hi");
|
||||
|
||||
engine_task.await.expect("mock engine task");
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_keepalive_closes_unresponsive_connection() {
|
||||
let (addr, server_task, _engine_task) =
|
||||
grpc_server_with_keepalive(b"engine-grpc-keepalive", Some(Duration::from_millis(150)))
|
||||
.await;
|
||||
|
||||
let closed = h2_unresponsive_peer_closed_within(&addr, Duration::from_secs(5)).await;
|
||||
assert!(
|
||||
closed,
|
||||
"keepalive must close a peer that stops answering PINGs"
|
||||
);
|
||||
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_without_keepalive_keeps_unresponsive_connection_open() {
|
||||
// Without keepalive the same unresponsive peer is NOT
|
||||
// closed, proving the close above is attributable to keepalive.
|
||||
let (addr, server_task, _engine_task) =
|
||||
grpc_server_with_keepalive(b"engine-grpc-no-keepalive", None).await;
|
||||
|
||||
let closed = h2_unresponsive_peer_closed_within(&addr, Duration::from_secs(1)).await;
|
||||
assert!(
|
||||
!closed,
|
||||
"without keepalive an idle h2 connection must stay open"
|
||||
);
|
||||
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
+205
-35
@@ -10,20 +10,34 @@ mod routes;
|
||||
mod runtime;
|
||||
mod server_info;
|
||||
mod state;
|
||||
mod tls;
|
||||
#[cfg(test)]
|
||||
mod tls_tests;
|
||||
mod utils;
|
||||
|
||||
use std::future::Future;
|
||||
use std::sync::{Arc, OnceLock};
|
||||
use std::time::Duration;
|
||||
|
||||
use anyhow::{Context as _, Result};
|
||||
use axum::Router;
|
||||
use axum::serve::ListenerExt as _;
|
||||
pub use config::{ApiServerOptions, Config, CoordinatorMode, CorsConfig, HttpListenerMode};
|
||||
use axum::body::Body;
|
||||
use axum::http::Request;
|
||||
pub use config::{
|
||||
ApiServerOptions, Config, CoordinatorMode, CorsConfig, DEFAULT_KEEP_ALIVE_TIMEOUT,
|
||||
HttpListenerMode, TlsConfig,
|
||||
};
|
||||
use futures::FutureExt as _;
|
||||
use hyper::body::Incoming;
|
||||
use hyper::server::conn::http1;
|
||||
use hyper_util::rt::{TokioIo, TokioTimer};
|
||||
use hyper_util::server::graceful::GracefulShutdown;
|
||||
use hyper_util::service::TowerToHyperService;
|
||||
use tokio::net::TcpListener;
|
||||
use tokio::time::{Instant, sleep_until};
|
||||
use tokio_stream::wrappers::TcpListenerStream;
|
||||
use tokio_util::either::Either;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use tonic::transport::Server as TonicServer;
|
||||
use tower::ServiceExt as _;
|
||||
use tracing::{info, trace, warn};
|
||||
use vllm_chat::{ChatLlm, LoadModelBackendsOptions, load_model_backends};
|
||||
pub use vllm_chat::{ChatTemplateContentFormatOption, ParserSelection, RendererSelection};
|
||||
@@ -36,6 +50,13 @@ use crate::routes::build_router;
|
||||
use crate::server_info::ServerInfoSnapshot;
|
||||
use crate::state::AppState;
|
||||
|
||||
/// How often the server PINGs an idle gRPC connection to reap a dead peer;
|
||||
/// tonic enables no keepalive by default. 2h matches the gRPC-core default.
|
||||
const GRPC_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(7200);
|
||||
/// How long the server waits for a keepalive PING reply before dropping the gRPC
|
||||
/// connection. 20s matches the gRPC-core default.
|
||||
const GRPC_KEEPALIVE_TIMEOUT: Duration = Duration::from_secs(20);
|
||||
|
||||
/// Resolve the public model names accepted by the frontend.
|
||||
fn effective_served_model_names(model: &str, served_model_name: &[String]) -> Vec<String> {
|
||||
if served_model_name.is_empty() {
|
||||
@@ -45,6 +66,17 @@ fn effective_served_model_names(model: &str, served_model_name: &[String]) -> Ve
|
||||
}
|
||||
}
|
||||
|
||||
/// Choose the gRPC listener host. It follows the HTTP TCP host when there is
|
||||
/// one; otherwise (unix socket or inherited fd) it defaults to IPv4 loopback
|
||||
/// rather than all interfaces, so the side-car is never accidentally
|
||||
/// network-exposed.
|
||||
fn grpc_bind_host(listener_mode: &HttpListenerMode) -> &str {
|
||||
match listener_mode {
|
||||
HttpListenerMode::BindTcp { host, .. } => host.as_str(),
|
||||
HttpListenerMode::BindUnix { .. } | HttpListenerMode::InheritedFd { .. } => "127.0.0.1",
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the shared application state for one configured model and one engine
|
||||
/// client.
|
||||
async fn build_state(config: &Config) -> Result<Arc<AppState>> {
|
||||
@@ -130,6 +162,15 @@ where
|
||||
{
|
||||
config.validate().context("invalid OpenAI frontend configuration")?;
|
||||
|
||||
// Build the TLS server config once, up front, so a bad cert/key fails fast
|
||||
// before the (potentially long) engine handshake.
|
||||
let tls_config = config
|
||||
.tls
|
||||
.as_ref()
|
||||
.map(tls::build_server_config)
|
||||
.transpose()
|
||||
.context("invalid TLS configuration")?;
|
||||
|
||||
// Also check shutdown during the (potentially long) startup handshake.
|
||||
let state = tokio::select! {
|
||||
result = build_state(&config) => result?,
|
||||
@@ -144,40 +185,39 @@ where
|
||||
|
||||
// Optionally bind the gRPC Generate server on a separate port. Bind
|
||||
// synchronously here so bind errors (port in use, permission denied, ...)
|
||||
// surface before we start serving, rather than being deferred until
|
||||
// shutdown. The gRPC listener follows the same host as the HTTP listener so
|
||||
// that enabling --grpc-port does not accidentally expose the service on all
|
||||
// interfaces when HTTP is intentionally local-only.
|
||||
// surface before serving rather than being deferred until shutdown.
|
||||
let grpc_setup = if let Some(grpc_port) = config.grpc_port {
|
||||
let grpc_host = match &config.listener_mode {
|
||||
HttpListenerMode::BindTcp { host, .. } => host.as_str(),
|
||||
HttpListenerMode::BindUnix { .. } | HttpListenerMode::InheritedFd { .. } => "0.0.0.0",
|
||||
};
|
||||
let grpc_host = grpc_bind_host(&config.listener_mode);
|
||||
let grpc_listener = TcpListener::bind((grpc_host, grpc_port))
|
||||
.await
|
||||
.with_context(|| format!("failed to bind gRPC listener on {grpc_host}:{grpc_port}"))?;
|
||||
let addr = grpc_listener.local_addr()?;
|
||||
let grpc_listener = Listener::Tcp(grpc_listener);
|
||||
// gRPC reuses the HTTP TLS config (same SslContext) plus ALPN h2.
|
||||
let grpc_tls = config
|
||||
.tls
|
||||
.as_ref()
|
||||
.map(tls::build_grpc_server_config)
|
||||
.transpose()
|
||||
.context("invalid gRPC TLS configuration")?;
|
||||
let svc = grpc::GenerateServer::new(grpc::GenerateServiceImpl::new(state.clone()));
|
||||
let svc = TonicServer::builder()
|
||||
.http2_keepalive_interval(Some(GRPC_KEEPALIVE_INTERVAL))
|
||||
.http2_keepalive_timeout(Some(GRPC_KEEPALIVE_TIMEOUT))
|
||||
.layer(middleware::request_runtime_layer(state.clone()))
|
||||
.add_service(svc);
|
||||
info!(%addr, "starting gRPC server");
|
||||
Some((grpc_listener, svc))
|
||||
info!(%addr, tls = grpc_tls.is_some(), "starting gRPC server");
|
||||
Some((grpc_listener, svc, grpc_tls))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
info!(%bind_address, %model, "starting OpenAI server");
|
||||
|
||||
// Set TCP_NODELAY on accepted connections to reduce latency.
|
||||
// By `tap_io` we will do this on every accepted connection.
|
||||
let listener = listener.tap_io(|io| {
|
||||
if let Either::Left(tcp_stream) = io
|
||||
&& let Err(err) = tcp_stream.set_nodelay(true)
|
||||
{
|
||||
trace!(error = %err, "failed to enable TCP_NODELAY on accepted HTTP connection");
|
||||
}
|
||||
});
|
||||
let scheme = if tls_config.is_some() {
|
||||
"https"
|
||||
} else {
|
||||
"http"
|
||||
};
|
||||
info!(%bind_address, %scheme, %model, "starting OpenAI server");
|
||||
|
||||
// Run HTTP and gRPC concurrently under a child token of the caller's shutdown
|
||||
// token. Caller cancellation propagates into both protocols; if either
|
||||
@@ -208,17 +248,27 @@ where
|
||||
}
|
||||
});
|
||||
|
||||
// 0 disables keep-alive but still bounds the head read (default), so a
|
||||
// silent client cannot hold the connection open.
|
||||
let keep_alive_timeout = config.keep_alive_timeout;
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: tls::TLS_HANDSHAKE_TIMEOUT,
|
||||
header_read: if keep_alive_timeout.is_zero() {
|
||||
DEFAULT_KEEP_ALIVE_TIMEOUT
|
||||
} else {
|
||||
keep_alive_timeout
|
||||
},
|
||||
keep_alive_enabled: !keep_alive_timeout.is_zero(),
|
||||
};
|
||||
|
||||
let http_fut = {
|
||||
let shutdown = server_shutdown.child_token();
|
||||
let server_shutdown = server_shutdown.clone();
|
||||
let force_shutdown = force_shutdown.clone();
|
||||
async move {
|
||||
let server =
|
||||
axum::serve(listener, app).with_graceful_shutdown(shutdown.cancelled_owned());
|
||||
|
||||
let result = tokio::select! {
|
||||
result = server => {
|
||||
result.context("HTTP server failed")
|
||||
result = serve_listener(listener, tls_config, app, shutdown.cancelled_owned(), timeouts) => {
|
||||
result
|
||||
}
|
||||
_ = force_shutdown.cancelled() => {
|
||||
warn!("HTTP graceful shutdown deadline elapsed; aborting server");
|
||||
@@ -236,16 +286,24 @@ where
|
||||
let server_shutdown = server_shutdown.clone();
|
||||
let force_shutdown = force_shutdown.clone();
|
||||
async move {
|
||||
let Some((grpc_listener, svc)) = grpc_setup else {
|
||||
let Some((grpc_listener, svc, grpc_tls)) = grpc_setup else {
|
||||
// No gRPC configured: just wait for shutdown so we do not race the
|
||||
// join! by resolving early and tripping the cancellation token.
|
||||
shutdown.cancelled().await;
|
||||
return Ok(());
|
||||
};
|
||||
let server = svc.serve_with_incoming_shutdown(
|
||||
TcpListenerStream::new(grpc_listener),
|
||||
shutdown.cancelled_owned(),
|
||||
);
|
||||
// Box to unify the TLS and plaintext arms' different stream types.
|
||||
let server = match grpc_tls {
|
||||
Some(context) => {
|
||||
let incoming =
|
||||
grpc::tls_incoming(grpc_listener, context, tls::TLS_HANDSHAKE_TIMEOUT);
|
||||
svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned()).boxed()
|
||||
}
|
||||
None => {
|
||||
let incoming = grpc::incoming(grpc_listener);
|
||||
svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned()).boxed()
|
||||
}
|
||||
};
|
||||
|
||||
let result = tokio::select! {
|
||||
result = server => {
|
||||
@@ -272,6 +330,99 @@ where
|
||||
state.shutdown(shutdown_deadline).await
|
||||
}
|
||||
|
||||
/// Per-connection timeouts applied while serving HTTP/HTTPS.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct ConnectionTimeouts {
|
||||
/// Max time for a client to complete the TLS handshake (TLS path only).
|
||||
pub(crate) handshake: Duration,
|
||||
/// HTTP/1 header-read timeout (bounds idle keep-alive and the head read).
|
||||
pub(crate) header_read: Duration,
|
||||
/// Whether HTTP/1 keep-alive is enabled; `false` closes after each response.
|
||||
pub(crate) keep_alive_enabled: bool,
|
||||
}
|
||||
|
||||
/// Apply optional TLS termination and per-connection HTTP timeouts, then serve
|
||||
/// `app`. Shared by [`serve_with_router_extension`] and the TLS tests.
|
||||
async fn serve_listener(
|
||||
listener: Listener,
|
||||
tls: Option<openssl::ssl::SslContext>,
|
||||
app: Router,
|
||||
shutdown: impl Future<Output = ()> + Send + 'static,
|
||||
timeouts: ConnectionTimeouts,
|
||||
) -> Result<()> {
|
||||
match tls {
|
||||
Some(context) => {
|
||||
// tls-listener terminates TLS (handshake + timeout); serve_connections
|
||||
// owns the HTTP keep-alive/idle bound that axum::serve cannot express.
|
||||
// Failed handshakes (incl. timeouts) log at ERROR via tls-listener.
|
||||
let listener = tls_listener::builder(context)
|
||||
.handshake_timeout(timeouts.handshake)
|
||||
.listen(listener);
|
||||
serve_connections(
|
||||
listener,
|
||||
app,
|
||||
shutdown,
|
||||
timeouts.header_read,
|
||||
timeouts.keep_alive_enabled,
|
||||
)
|
||||
.await
|
||||
.context("HTTPS server failed")
|
||||
}
|
||||
None => serve_connections(
|
||||
listener,
|
||||
app,
|
||||
shutdown,
|
||||
timeouts.header_read,
|
||||
timeouts.keep_alive_enabled,
|
||||
)
|
||||
.await
|
||||
.context("HTTP server failed"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Serve `app` per connection (HTTP/1) with a keep-alive idle timeout and
|
||||
/// graceful drain. Hand-rolled on hyper because [`axum::serve()`] takes no config.
|
||||
async fn serve_connections<L>(
|
||||
mut listener: L,
|
||||
app: Router,
|
||||
shutdown: impl Future<Output = ()> + Send,
|
||||
header_read: Duration,
|
||||
keep_alive_enabled: bool,
|
||||
) -> Result<()>
|
||||
where
|
||||
L: axum::serve::Listener,
|
||||
{
|
||||
let graceful = GracefulShutdown::new();
|
||||
let mut shutdown = std::pin::pin!(shutdown);
|
||||
loop {
|
||||
let (io, _addr) = tokio::select! {
|
||||
conn = listener.accept() => conn,
|
||||
() = &mut shutdown => break,
|
||||
};
|
||||
|
||||
let service = TowerToHyperService::new(
|
||||
app.clone().map_request(|req: Request<Incoming>| req.map(Body::new)),
|
||||
);
|
||||
let mut builder = http1::Builder::new();
|
||||
builder.timer(TokioTimer::new()).header_read_timeout(header_read);
|
||||
if !keep_alive_enabled {
|
||||
builder.keep_alive(false);
|
||||
}
|
||||
let connection = builder.serve_connection(TokioIo::new(io), service);
|
||||
let connection = graceful.watch(connection);
|
||||
|
||||
tokio::spawn(async move {
|
||||
if let Err(err) = connection.await {
|
||||
trace!(error = %err, "failed to serve connection");
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
drop(listener);
|
||||
graceful.shutdown().await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -293,4 +444,23 @@ mod tests {
|
||||
served_names
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grpc_bind_host_follows_http_tcp_host() {
|
||||
let mode = HttpListenerMode::BindTcp {
|
||||
host: "0.0.0.0".to_string(),
|
||||
port: 8000,
|
||||
};
|
||||
assert_eq!(grpc_bind_host(&mode), "0.0.0.0");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grpc_bind_host_defaults_to_loopback_without_tcp_host() {
|
||||
let unix = HttpListenerMode::BindUnix {
|
||||
path: "/tmp/vllm.sock".to_string(),
|
||||
};
|
||||
let inherited = HttpListenerMode::InheritedFd { fd: 3 };
|
||||
assert_eq!(grpc_bind_host(&unix), "127.0.0.1");
|
||||
assert_eq!(grpc_bind_host(&inherited), "127.0.0.1");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,28 +1,49 @@
|
||||
//! Unified HTTP listener wrapper for the Rust frontend.
|
||||
//! Unified listener wrapper for the Rust frontend.
|
||||
//!
|
||||
//! This module hides the difference between TCP and Unix-domain listeners so
|
||||
//! the rest of the server can bind or inherit one socket and pass it to
|
||||
//! `axum::serve(...)` through a single type.
|
||||
|
||||
use std::io::Result;
|
||||
use std::net::TcpListener as StdTcpListener;
|
||||
use std::net::{SocketAddr, TcpListener as StdTcpListener};
|
||||
use std::os::fd::{FromRawFd, IntoRawFd, OwnedFd};
|
||||
use std::os::unix::net::UnixListener as StdUnixListener;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll, ready};
|
||||
|
||||
use auto_enums::enum_derive;
|
||||
use socket2::Socket;
|
||||
use tls_listener::{AsyncAccept, AsyncListener};
|
||||
use tokio::net::{TcpListener, TcpStream, UnixListener, UnixStream};
|
||||
use tokio_util::either::Either;
|
||||
use tonic::transport::server::{Connected, TcpConnectInfo};
|
||||
use tracing::trace;
|
||||
|
||||
use crate::HttpListenerMode;
|
||||
|
||||
/// Runtime listener type used by the OpenAI-compatible HTTP server, which is
|
||||
/// either a TCP listener or a Unix-domain listener.
|
||||
/// Runtime listener type used by the OpenAI-compatible HTTP or gRPC server,
|
||||
/// which is either a TCP listener or a Unix-domain listener.
|
||||
#[derive(Debug)]
|
||||
pub enum Listener {
|
||||
Tcp(TcpListener),
|
||||
Unix(UnixListener),
|
||||
}
|
||||
|
||||
/// Runtime listener I/O type which is either a TCP stream or a Unix-domain stream.
|
||||
#[derive(Debug)]
|
||||
#[enum_derive(tokio1::AsyncRead, tokio1::AsyncWrite)]
|
||||
pub enum ListenerIo {
|
||||
Tcp(TcpStream),
|
||||
Unix(UnixStream),
|
||||
}
|
||||
|
||||
/// Runtime listener address type which is either a TCP address or a Unix-domain address.
|
||||
#[derive(Debug)]
|
||||
#[allow(dead_code)]
|
||||
pub enum ListenerAddr {
|
||||
Tcp(SocketAddr),
|
||||
Unix(tokio::net::unix::SocketAddr),
|
||||
}
|
||||
|
||||
impl Listener {
|
||||
/// Bind or adopt the listener described by the frontend configuration.
|
||||
///
|
||||
@@ -70,34 +91,95 @@ impl Listener {
|
||||
Ok(Self::Tcp(TcpListener::from_std(std_listener)?))
|
||||
}
|
||||
}
|
||||
|
||||
fn listener_addr(&self) -> Result<ListenerAddr> {
|
||||
match self {
|
||||
Self::Tcp(listener) => listener.local_addr().map(ListenerAddr::Tcp),
|
||||
Self::Unix(listener) => listener.local_addr().map(ListenerAddr::Unix),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Connected for ListenerIo {
|
||||
type ConnectInfo = TcpConnectInfo;
|
||||
|
||||
fn connect_info(&self) -> TcpConnectInfo {
|
||||
match self {
|
||||
Self::Tcp(stream) => stream.connect_info(),
|
||||
Self::Unix(_) => TcpConnectInfo {
|
||||
local_addr: None,
|
||||
remote_addr: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Attempt to set `TCP_NODELAY` on the accepted TCP stream.
|
||||
fn enable_tcp_nodelay(stream: TcpStream) -> TcpStream {
|
||||
if let Err(err) = stream.set_nodelay(true) {
|
||||
trace!(error = %err, "failed to enable TCP_NODELAY on accepted TCP connection");
|
||||
}
|
||||
stream
|
||||
}
|
||||
|
||||
/// Allow the unified listener to plug directly into `axum::serve(...)`.
|
||||
impl axum::serve::Listener for Listener {
|
||||
type Addr = Either<std::net::SocketAddr, tokio::net::unix::SocketAddr>;
|
||||
type Io = Either<TcpStream, UnixStream>;
|
||||
type Addr = ListenerAddr;
|
||||
type Io = ListenerIo;
|
||||
|
||||
async fn accept(&mut self) -> (Self::Io, Self::Addr) {
|
||||
match self {
|
||||
Self::Tcp(listener) => {
|
||||
let (io, addr) = listener.accept().await;
|
||||
(Either::Left(io), Either::Left(addr))
|
||||
let (io, addr) = axum::serve::Listener::accept(listener).await;
|
||||
(
|
||||
ListenerIo::Tcp(enable_tcp_nodelay(io)),
|
||||
ListenerAddr::Tcp(addr),
|
||||
)
|
||||
}
|
||||
Self::Unix(listener) => {
|
||||
let (io, addr) = listener.accept().await;
|
||||
(Either::Right(io), Either::Right(addr))
|
||||
let (io, addr) = axum::serve::Listener::accept(listener).await;
|
||||
(ListenerIo::Unix(io), ListenerAddr::Unix(addr))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn local_addr(&self) -> Result<Self::Addr> {
|
||||
match self {
|
||||
Self::Tcp(listener) => listener.local_addr().map(Either::Left),
|
||||
Self::Unix(listener) => listener.local_addr().map(Either::Right),
|
||||
self.listener_addr()
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow the unified listener to be adaptable to `tls_listener`.
|
||||
impl AsyncAccept for Listener {
|
||||
type Connection = ListenerIo;
|
||||
type Address = ListenerAddr;
|
||||
type Error = std::io::Error;
|
||||
|
||||
fn poll_accept(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Result<(Self::Connection, Self::Address)>> {
|
||||
match self.get_mut() {
|
||||
Self::Tcp(listener) => {
|
||||
let (io, addr) = ready!(listener.poll_accept(cx))?;
|
||||
Poll::Ready(Ok((
|
||||
ListenerIo::Tcp(enable_tcp_nodelay(io)),
|
||||
ListenerAddr::Tcp(addr),
|
||||
)))
|
||||
}
|
||||
Self::Unix(listener) => {
|
||||
let (io, addr) = ready!(listener.poll_accept(cx))?;
|
||||
Poll::Ready(Ok((ListenerIo::Unix(io), ListenerAddr::Unix(addr))))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncListener for Listener {
|
||||
fn local_addr(&self) -> Result<Self::Address> {
|
||||
self.listener_addr()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::net::{Ipv4Addr, SocketAddrV4};
|
||||
|
||||
@@ -34,14 +34,20 @@ pub(super) fn validate_request_compat(
|
||||
);
|
||||
}
|
||||
|
||||
if let Some(prompt_logprobs) = request.sampling_params.prompt_logprobs
|
||||
&& prompt_logprobs < 0
|
||||
&& prompt_logprobs != -1
|
||||
{
|
||||
bail_invalid_request!(
|
||||
param = "sampling_params",
|
||||
"`prompt_logprobs` must be a non-negative value or -1."
|
||||
);
|
||||
if let Some(prompt_logprobs) = request.sampling_params.prompt_logprobs {
|
||||
if prompt_logprobs < 0 && prompt_logprobs != -1 {
|
||||
bail_invalid_request!(
|
||||
param = "sampling_params",
|
||||
"`prompt_logprobs` must be a non-negative value or -1."
|
||||
);
|
||||
}
|
||||
|
||||
if request.stream {
|
||||
bail_invalid_request!(
|
||||
param = "sampling_params",
|
||||
"`prompt_logprobs` are not available when `stream=true`."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -97,4 +103,54 @@ mod tests {
|
||||
};
|
||||
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_request_compat_rejects_streaming_prompt_logprobs() {
|
||||
let request: GenerateRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"token_ids": [11, 22],
|
||||
"stream": true,
|
||||
"sampling_params": {
|
||||
"prompt_logprobs": 0
|
||||
}
|
||||
}))
|
||||
.expect("parse request");
|
||||
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
|
||||
|
||||
let request: GenerateRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"token_ids": [11, 22],
|
||||
"stream": true,
|
||||
"sampling_params": {
|
||||
"prompt_logprobs": 1
|
||||
}
|
||||
}))
|
||||
.expect("parse request");
|
||||
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
|
||||
|
||||
let request: GenerateRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"token_ids": [11, 22],
|
||||
"stream": true,
|
||||
"sampling_params": {
|
||||
"prompt_logprobs": -1
|
||||
}
|
||||
}))
|
||||
.expect("parse request");
|
||||
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_request_compat_accepts_non_stream_prompt_logprobs() {
|
||||
let request: GenerateRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"token_ids": [11, 22],
|
||||
"stream": false,
|
||||
"sampling_params": {
|
||||
"prompt_logprobs": 1
|
||||
}
|
||||
}))
|
||||
.expect("parse request");
|
||||
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_ok());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4129,6 +4129,45 @@ async fn raw_generate_rejects_empty_token_ids() {
|
||||
assert_eq!(json["error"]["param"], "token_ids");
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn raw_generate_rejects_streaming_prompt_logprobs() {
|
||||
let mut app = test_app().await;
|
||||
|
||||
for prompt_logprobs in [0, 1] {
|
||||
let response = app
|
||||
.call(
|
||||
Request::builder()
|
||||
.method("POST")
|
||||
.uri("/inference/v1/generate")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(
|
||||
json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"token_ids": [11, 22],
|
||||
"stream": true,
|
||||
"sampling_params": {
|
||||
"prompt_logprobs": prompt_logprobs
|
||||
}
|
||||
})
|
||||
.to_string(),
|
||||
))
|
||||
.expect("build request"),
|
||||
)
|
||||
.await
|
||||
.expect("call app");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.expect("read body");
|
||||
let json: serde_json::Value = serde_json::from_slice(&body).expect("decode json");
|
||||
assert_eq!(json["error"]["param"], "sampling_params");
|
||||
assert_eq!(
|
||||
json["error"]["message"],
|
||||
"`prompt_logprobs` are not available when `stream=true`."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn raw_generate_rejects_wrong_model() {
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
//! OpenSSL server-config construction for TLS termination.
|
||||
//!
|
||||
//! Builds an OpenSSL [`SslContext`] from the uvicorn-style `ssl_*` arguments
|
||||
//! (certificate chain, private key, mTLS client verifier, optional cipher list).
|
||||
//! The `tls-listener` crate drives the handshake on each accepted connection.
|
||||
//!
|
||||
//! Crypto runs through whichever OpenSSL the binary links (system by default,
|
||||
//! vendored when built with that feature).
|
||||
|
||||
use std::path::Path;
|
||||
use std::time::Duration;
|
||||
|
||||
use anyhow::{Context as _, Result};
|
||||
use openssl::ssl::{
|
||||
AlpnError, SslAcceptor, SslAcceptorBuilder, SslContext, SslContextBuilder, SslFiletype,
|
||||
SslMethod, SslOptions, SslVerifyMode, select_next_proto,
|
||||
};
|
||||
|
||||
use crate::config::TlsConfig;
|
||||
|
||||
/// Time a client has to complete the TLS handshake before the connection is dropped.
|
||||
pub(crate) const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(60);
|
||||
|
||||
/// ALPN wire bytes for HTTP/2 (length-prefixed).
|
||||
const ALPN_H2: &[u8] = b"\x02h2";
|
||||
|
||||
/// Build the shared OpenSSL acceptor from validated [`TlsConfig`]: the full
|
||||
/// certificate chain, the private key (`key_file`, or the certificate file when
|
||||
/// unset), the mTLS client verifier, and an optional cipher list.
|
||||
///
|
||||
/// Starts from the Mozilla intermediate baseline (forward-secret AEAD suites,
|
||||
/// TLS 1.2 floor, server cipher preference, no compression), a slightly
|
||||
/// stricter subset of the Python frontend's default suites; `--ssl-ciphers`
|
||||
/// overrides it.
|
||||
fn build_server_builder(tls: &TlsConfig) -> Result<SslAcceptorBuilder> {
|
||||
let cert_file = tls.cert_file.as_deref().context("--ssl-certfile is required to enable TLS")?;
|
||||
|
||||
let mut builder = SslAcceptor::mozilla_intermediate_v5(SslMethod::tls_server())
|
||||
.context("failed to initialize TLS")?;
|
||||
builder.set_options(SslOptions::CIPHER_SERVER_PREFERENCE);
|
||||
|
||||
// Load the whole chain (leaf + intermediates), not just the leaf, so
|
||||
// deployments behind an intermediate CA serve a complete chain.
|
||||
ensure_exists(cert_file, "--ssl-certfile")?;
|
||||
builder.set_certificate_chain_file(cert_file).with_context(|| {
|
||||
format!("failed to parse certificate chain in --ssl-certfile {cert_file:?}")
|
||||
})?;
|
||||
|
||||
// When `key_file` is unset the key is read from the certificate file
|
||||
// (combined PEM).
|
||||
let key_file = tls.key_file.as_deref().unwrap_or(cert_file);
|
||||
ensure_exists(key_file, "private key file")?;
|
||||
builder
|
||||
.set_private_key_file(key_file, SslFiletype::PEM)
|
||||
.with_context(|| format!("failed to parse private key in {key_file:?}"))?;
|
||||
builder
|
||||
.check_private_key()
|
||||
.context("the certificate and private key do not match")?;
|
||||
|
||||
configure_client_auth(&mut builder, tls)?;
|
||||
|
||||
if let Some(ciphers) = tls.ciphers.as_deref().filter(|c| !c.is_empty()) {
|
||||
builder
|
||||
.set_cipher_list(ciphers)
|
||||
.with_context(|| format!("invalid --ssl-ciphers {ciphers:?}"))?;
|
||||
}
|
||||
|
||||
Ok(builder)
|
||||
}
|
||||
|
||||
/// Build the HTTP [`SslContext`] (HTTP/1.1; no ALPN, matching uvicorn).
|
||||
pub(crate) fn build_server_config(tls: &TlsConfig) -> Result<SslContext> {
|
||||
Ok(build_server_builder(tls)?.build().into_context())
|
||||
}
|
||||
|
||||
/// Build the gRPC [`SslContext`]: identical to [`build_server_config`] but
|
||||
/// negotiates ALPN `h2`, which HTTP/2 over TLS requires.
|
||||
pub(crate) fn build_grpc_server_config(tls: &TlsConfig) -> Result<SslContext> {
|
||||
let mut builder = build_server_builder(tls)?;
|
||||
builder.set_alpn_select_callback(|_ssl, client| {
|
||||
select_next_proto(ALPN_H2, client).ok_or(AlpnError::NOACK)
|
||||
});
|
||||
Ok(builder.build().into_context())
|
||||
}
|
||||
|
||||
/// Fail loudly with a flag-named message when a configured file is missing,
|
||||
/// distinguishing it from a malformed-PEM error raised later by OpenSSL (whose
|
||||
/// `ErrorStack` does not name the offending file).
|
||||
fn ensure_exists(path: &str, what: &str) -> Result<()> {
|
||||
std::fs::metadata(Path::new(path))
|
||||
.map(drop)
|
||||
.with_context(|| format!("failed to read {what} {path:?}"))
|
||||
}
|
||||
|
||||
/// Apply the `cert_reqs` client-certificate policy: 0 = none, 1 = optional
|
||||
/// (verify if presented, allow anonymous), 2 = required. `PEER` without a custom
|
||||
/// verify callback still rejects a presented-but-untrusted certificate.
|
||||
fn configure_client_auth(builder: &mut SslContextBuilder, tls: &TlsConfig) -> Result<()> {
|
||||
if tls.cert_reqs == 0 {
|
||||
builder.set_verify(SslVerifyMode::NONE);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let ca_file = tls
|
||||
.ca_certs
|
||||
.as_deref()
|
||||
.context("--ssl-ca-certs is required for client certificate verification")?;
|
||||
ensure_exists(ca_file, "--ssl-ca-certs")?;
|
||||
builder
|
||||
.set_ca_file(ca_file)
|
||||
.with_context(|| format!("failed to parse --ssl-ca-certs {ca_file:?}"))?;
|
||||
|
||||
let mut mode = SslVerifyMode::PEER;
|
||||
if tls.cert_reqs == 2 {
|
||||
mode |= SslVerifyMode::FAIL_IF_NO_PEER_CERT;
|
||||
}
|
||||
builder.set_verify(mode);
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,688 @@
|
||||
//! TLS tests: `build_server_config` unit checks plus end-to-end OpenSSL handshakes
|
||||
//! through the production `serve_listener` path, with a trivial router since TLS
|
||||
//! terminates below the app.
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::Router;
|
||||
use axum::routing::get;
|
||||
use openssl::asn1::Asn1Time;
|
||||
use openssl::bn::{BigNum, MsbOption};
|
||||
use openssl::ec::{EcGroup, EcKey};
|
||||
use openssl::hash::MessageDigest;
|
||||
use openssl::nid::Nid;
|
||||
use openssl::pkey::{PKey, Private};
|
||||
use openssl::ssl::{SslConnector, SslFiletype, SslMethod, SslVersion};
|
||||
use openssl::x509::extension::{BasicConstraints, KeyUsage, SubjectAlternativeName};
|
||||
use openssl::x509::{X509, X509NameBuilder};
|
||||
use tempfile::TempDir;
|
||||
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_openssl::SslStream;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::config::{HttpListenerMode, TlsConfig};
|
||||
use crate::listener::Listener;
|
||||
use crate::{ConnectionTimeouts, serve_listener, tls};
|
||||
|
||||
// ============================================================================
|
||||
// Test infrastructure
|
||||
// ============================================================================
|
||||
|
||||
/// A throwaway CA + server/client/untrusted/chain cert set as PEM files in a
|
||||
/// temp dir; dropping it deletes them.
|
||||
pub(crate) struct TestCerts {
|
||||
dir: TempDir,
|
||||
}
|
||||
|
||||
impl TestCerts {
|
||||
pub(crate) fn generate() -> Self {
|
||||
let dir = tempfile::tempdir().expect("tempdir");
|
||||
|
||||
let (ca, ca_key) = build_ca();
|
||||
let (server, server_key) = build_leaf("server", &["127.0.0.1", "localhost"], &ca, &ca_key);
|
||||
let (client, client_key) = build_leaf("client", &[], &ca, &ca_key);
|
||||
let (untrusted, untrusted_key) = build_self_signed("untrusted client");
|
||||
|
||||
// Leaf signed by an intermediate (itself signed by the root); the cert
|
||||
// file holds leaf + intermediate, for the chain-serving test.
|
||||
let (intermediate, intermediate_key) = build_intermediate(&ca, &ca_key);
|
||||
let (chain_leaf, chain_leaf_key) = build_leaf(
|
||||
"chain",
|
||||
&["127.0.0.1", "localhost"],
|
||||
&intermediate,
|
||||
&intermediate_key,
|
||||
);
|
||||
|
||||
let server_pem = pem(&server);
|
||||
let server_key_pem = key_pem(&server_key);
|
||||
let files = [
|
||||
("ca.pem", pem(&ca)),
|
||||
("server.pem", server_pem.clone()),
|
||||
("server.key", server_key_pem.clone()),
|
||||
("client.pem", pem(&client)),
|
||||
("client.key", key_pem(&client_key)),
|
||||
("untrusted_client.pem", pem(&untrusted)),
|
||||
("untrusted_client.key", key_pem(&untrusted_key)),
|
||||
(
|
||||
"server_combined.pem",
|
||||
format!("{server_pem}{server_key_pem}"),
|
||||
),
|
||||
(
|
||||
"server_chain.pem",
|
||||
format!("{}{}", pem(&chain_leaf), pem(&intermediate)),
|
||||
),
|
||||
("server_chain.key", key_pem(&chain_leaf_key)),
|
||||
];
|
||||
for (name, contents) in files {
|
||||
std::fs::write(dir.path().join(name), contents).expect("write fixture");
|
||||
}
|
||||
Self { dir }
|
||||
}
|
||||
|
||||
/// Absolute path to a fixture by name; the file need not exist.
|
||||
pub(crate) fn path(&self, name: &str) -> String {
|
||||
self.dir.path().join(name).to_str().expect("utf-8 path").to_string()
|
||||
}
|
||||
}
|
||||
|
||||
fn gen_key() -> PKey<Private> {
|
||||
let group = EcGroup::from_curve_name(Nid::X9_62_PRIME256V1).expect("ec group");
|
||||
let ec = EcKey::generate(&group).expect("ec key");
|
||||
PKey::from_ec_key(ec).expect("pkey")
|
||||
}
|
||||
|
||||
fn serial() -> openssl::asn1::Asn1Integer {
|
||||
let mut bn = BigNum::new().expect("bignum");
|
||||
bn.rand(159, MsbOption::MAYBE_ZERO, false).expect("rand serial");
|
||||
bn.to_asn1_integer().expect("asn1 serial")
|
||||
}
|
||||
|
||||
fn x509_name(cn: &str) -> openssl::x509::X509Name {
|
||||
let mut builder = X509NameBuilder::new().expect("name builder");
|
||||
builder.append_entry_by_text("CN", cn).expect("cn");
|
||||
builder.build()
|
||||
}
|
||||
|
||||
fn pem(cert: &X509) -> String {
|
||||
String::from_utf8(cert.to_pem().expect("cert pem")).expect("utf-8 cert")
|
||||
}
|
||||
|
||||
fn key_pem(key: &PKey<Private>) -> String {
|
||||
String::from_utf8(key.private_key_to_pem_pkcs8().expect("key pem")).expect("utf-8 key")
|
||||
}
|
||||
|
||||
/// A self-signed CA used to sign the server/client leaf certs.
|
||||
fn build_ca() -> (X509, PKey<Private>) {
|
||||
let key = gen_key();
|
||||
let name = x509_name("vLLM Test CA");
|
||||
let mut builder = X509::builder().expect("x509 builder");
|
||||
builder.set_version(2).expect("version");
|
||||
builder.set_serial_number(&serial()).expect("serial");
|
||||
builder.set_subject_name(&name).expect("subject");
|
||||
builder.set_issuer_name(&name).expect("issuer");
|
||||
builder.set_pubkey(&key).expect("pubkey");
|
||||
builder
|
||||
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
|
||||
.expect("set nb");
|
||||
builder
|
||||
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
|
||||
.expect("set na");
|
||||
builder
|
||||
.append_extension(BasicConstraints::new().critical().ca().build().expect("bc"))
|
||||
.expect("ext bc");
|
||||
builder
|
||||
.append_extension(
|
||||
KeyUsage::new().critical().key_cert_sign().crl_sign().build().expect("ku"),
|
||||
)
|
||||
.expect("ext ku");
|
||||
builder.sign(&key, MessageDigest::sha256()).expect("sign ca");
|
||||
(builder.build(), key)
|
||||
}
|
||||
|
||||
/// A CA-signed leaf cert with optional subject-alternative names (IP or DNS).
|
||||
fn build_leaf(cn: &str, sans: &[&str], ca: &X509, ca_key: &PKey<Private>) -> (X509, PKey<Private>) {
|
||||
let key = gen_key();
|
||||
let mut builder = X509::builder().expect("x509 builder");
|
||||
builder.set_version(2).expect("version");
|
||||
builder.set_serial_number(&serial()).expect("serial");
|
||||
builder.set_subject_name(&x509_name(cn)).expect("subject");
|
||||
builder.set_issuer_name(ca.subject_name()).expect("issuer");
|
||||
builder.set_pubkey(&key).expect("pubkey");
|
||||
builder
|
||||
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
|
||||
.expect("set nb");
|
||||
builder
|
||||
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
|
||||
.expect("set na");
|
||||
builder
|
||||
.append_extension(BasicConstraints::new().build().expect("bc"))
|
||||
.expect("ext bc");
|
||||
if !sans.is_empty() {
|
||||
let mut san = SubjectAlternativeName::new();
|
||||
for entry in sans {
|
||||
if entry.parse::<std::net::IpAddr>().is_ok() {
|
||||
san.ip(entry);
|
||||
} else {
|
||||
san.dns(entry);
|
||||
}
|
||||
}
|
||||
let ext = san.build(&builder.x509v3_context(Some(ca), None)).expect("san");
|
||||
builder.append_extension(ext).expect("ext san");
|
||||
}
|
||||
builder.sign(ca_key, MessageDigest::sha256()).expect("sign leaf");
|
||||
(builder.build(), key)
|
||||
}
|
||||
|
||||
/// A self-signed leaf not chained to the CA, for the untrusted-client test.
|
||||
fn build_self_signed(cn: &str) -> (X509, PKey<Private>) {
|
||||
let key = gen_key();
|
||||
let name = x509_name(cn);
|
||||
let mut builder = X509::builder().expect("x509 builder");
|
||||
builder.set_version(2).expect("version");
|
||||
builder.set_serial_number(&serial()).expect("serial");
|
||||
builder.set_subject_name(&name).expect("subject");
|
||||
builder.set_issuer_name(&name).expect("issuer");
|
||||
builder.set_pubkey(&key).expect("pubkey");
|
||||
builder
|
||||
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
|
||||
.expect("set nb");
|
||||
builder
|
||||
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
|
||||
.expect("set na");
|
||||
builder
|
||||
.append_extension(BasicConstraints::new().build().expect("bc"))
|
||||
.expect("ext bc");
|
||||
builder.sign(&key, MessageDigest::sha256()).expect("sign self");
|
||||
(builder.build(), key)
|
||||
}
|
||||
|
||||
/// A CA-capable intermediate signed by the root, for the full-chain test.
|
||||
fn build_intermediate(ca: &X509, ca_key: &PKey<Private>) -> (X509, PKey<Private>) {
|
||||
let key = gen_key();
|
||||
let mut builder = X509::builder().expect("x509 builder");
|
||||
builder.set_version(2).expect("version");
|
||||
builder.set_serial_number(&serial()).expect("serial");
|
||||
builder
|
||||
.set_subject_name(&x509_name("vLLM Test Intermediate CA"))
|
||||
.expect("subject");
|
||||
builder.set_issuer_name(ca.subject_name()).expect("issuer");
|
||||
builder.set_pubkey(&key).expect("pubkey");
|
||||
builder
|
||||
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
|
||||
.expect("set nb");
|
||||
builder
|
||||
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
|
||||
.expect("set na");
|
||||
builder
|
||||
.append_extension(BasicConstraints::new().critical().ca().build().expect("bc"))
|
||||
.expect("ext bc");
|
||||
builder
|
||||
.append_extension(
|
||||
KeyUsage::new().critical().key_cert_sign().crl_sign().build().expect("ku"),
|
||||
)
|
||||
.expect("ext ku");
|
||||
builder.sign(ca_key, MessageDigest::sha256()).expect("sign intermediate");
|
||||
(builder.build(), key)
|
||||
}
|
||||
|
||||
pub(crate) fn server_tls(certs: &TestCerts, cert_reqs: i32) -> TlsConfig {
|
||||
TlsConfig {
|
||||
cert_file: Some(certs.path("server.pem")),
|
||||
key_file: Some(certs.path("server.key")),
|
||||
ca_certs: (cert_reqs != 0).then(|| certs.path("ca.pem")),
|
||||
cert_reqs,
|
||||
ciphers: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// A plaintext-listener TLS config for `build_server_config` checks (`cert_reqs`
|
||||
/// 0, no client auth), with the cert/key files chosen by the caller.
|
||||
fn build_tls(certs: &TestCerts, cert: &str, key: Option<&str>) -> TlsConfig {
|
||||
TlsConfig {
|
||||
cert_file: Some(certs.path(cert)),
|
||||
key_file: key.map(|k| certs.path(k)),
|
||||
ca_certs: None,
|
||||
cert_reqs: 0,
|
||||
ciphers: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Generous per-connection timeouts that never fire during the fast tests.
|
||||
const TEST_TIMEOUTS: ConnectionTimeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_secs(5),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
|
||||
async fn spawn_server(tls_config: Option<TlsConfig>) -> (String, CancellationToken) {
|
||||
spawn_server_with_timeouts(tls_config, TEST_TIMEOUTS).await
|
||||
}
|
||||
|
||||
/// Bind an ephemeral listener and serve a trivial router via the production
|
||||
/// `serve_listener`, optionally with TLS. The listener is bound (and thus
|
||||
/// accepting into the backlog) before returning, so a client may connect
|
||||
/// immediately without a sleep.
|
||||
async fn spawn_server_with_timeouts(
|
||||
tls_config: Option<TlsConfig>,
|
||||
timeouts: ConnectionTimeouts,
|
||||
) -> (String, CancellationToken) {
|
||||
let listener = Listener::bind(&HttpListenerMode::BindTcp {
|
||||
host: "127.0.0.1".to_string(),
|
||||
port: 0,
|
||||
})
|
||||
.await
|
||||
.expect("bind listener");
|
||||
let addr = listener.local_addr().expect("local addr");
|
||||
|
||||
let server_config =
|
||||
tls_config.map(|cfg| tls::build_server_config(&cfg).expect("build server config"));
|
||||
let app = Router::new().route("/health", get(|| async { "ok" }));
|
||||
let shutdown = CancellationToken::new();
|
||||
let server_shutdown = shutdown.clone();
|
||||
tokio::spawn(async move {
|
||||
let _ = serve_listener(
|
||||
listener,
|
||||
server_config,
|
||||
app,
|
||||
server_shutdown.cancelled_owned(),
|
||||
timeouts,
|
||||
)
|
||||
.await;
|
||||
});
|
||||
(addr, shutdown)
|
||||
}
|
||||
|
||||
/// Open a TLS connection trusting the test CA and finish the handshake,
|
||||
/// optionally presenting a client identity (`<name>.pem` + `<name>.key`) for
|
||||
/// mTLS. Hostname verification is disabled (the IP-SAN match is not under test);
|
||||
/// chain verification stays on, so an untrusted server cert is still rejected.
|
||||
async fn connect_tls(
|
||||
certs: &TestCerts,
|
||||
addr: &str,
|
||||
identity: Option<&str>,
|
||||
) -> std::io::Result<Pin<Box<SslStream<TcpStream>>>> {
|
||||
let tcp = TcpStream::connect(addr).await?;
|
||||
|
||||
let mut builder = SslConnector::builder(SslMethod::tls_client()).expect("connector builder");
|
||||
builder.set_ca_file(certs.path("ca.pem")).expect("trust ca");
|
||||
if let Some(name) = identity {
|
||||
builder
|
||||
.set_certificate_chain_file(certs.path(&format!("{name}.pem")))
|
||||
.expect("client cert");
|
||||
builder
|
||||
.set_private_key_file(certs.path(&format!("{name}.key")), SslFiletype::PEM)
|
||||
.expect("client key");
|
||||
}
|
||||
let connector = builder.build();
|
||||
let mut config = connector.configure().expect("configure");
|
||||
config.set_verify_hostname(false);
|
||||
let ssl = config.into_ssl("127.0.0.1").expect("ssl");
|
||||
|
||||
let mut stream = Box::pin(SslStream::new(ssl, tcp).expect("client ssl stream"));
|
||||
stream.as_mut().connect().await.map_err(std::io::Error::other)?;
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
/// Issue an HTTPS GET (with `Connection: close`), optionally with an mTLS identity.
|
||||
async fn https_get(
|
||||
certs: &TestCerts,
|
||||
addr: &str,
|
||||
identity: Option<&str>,
|
||||
) -> std::io::Result<String> {
|
||||
let mut stream = connect_tls(certs, addr, identity).await?;
|
||||
stream
|
||||
.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n")
|
||||
.await?;
|
||||
let mut response = String::new();
|
||||
stream.read_to_string(&mut response).await?;
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
/// Attempt a handshake offering only a legacy CBC+SHA1 suite over TLS 1.2,
|
||||
/// capping the version so TLS 1.3 cannot rescue the negotiation.
|
||||
async fn legacy_suite_handshake(certs: &TestCerts, addr: &str) -> std::io::Result<()> {
|
||||
let tcp = TcpStream::connect(addr).await?;
|
||||
|
||||
let mut builder = SslConnector::builder(SslMethod::tls_client()).expect("connector builder");
|
||||
builder.set_ca_file(certs.path("ca.pem")).expect("trust ca");
|
||||
builder.set_max_proto_version(Some(SslVersion::TLS1_2)).expect("cap tls1.2");
|
||||
builder
|
||||
.set_cipher_list("ECDHE-ECDSA-AES256-SHA:@SECLEVEL=0")
|
||||
.expect("legacy cipher");
|
||||
let connector = builder.build();
|
||||
let mut config = connector.configure().expect("configure");
|
||||
config.set_verify_hostname(false);
|
||||
let ssl = config.into_ssl("127.0.0.1").expect("ssl");
|
||||
|
||||
let stream = SslStream::new(ssl, tcp).expect("client ssl stream");
|
||||
tokio::pin!(stream);
|
||||
stream.as_mut().connect().await.map_err(std::io::Error::other)
|
||||
}
|
||||
|
||||
async fn plain_get(addr: &str) -> std::io::Result<String> {
|
||||
let mut tcp = TcpStream::connect(addr).await?;
|
||||
tcp.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n")
|
||||
.await?;
|
||||
let mut response = String::new();
|
||||
tcp.read_to_string(&mut response).await?;
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Tests
|
||||
// ============================================================================
|
||||
|
||||
#[test]
|
||||
fn builds_from_combined_pem() {
|
||||
// Key omitted: it is read from the combined cert+key file.
|
||||
let certs = TestCerts::generate();
|
||||
assert!(tls::build_server_config(&build_tls(&certs, "server_combined.pem", None)).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_missing_cert_file() {
|
||||
let certs = TestCerts::generate();
|
||||
assert!(tls::build_server_config(&build_tls(&certs, "does_not_exist.pem", None)).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_valid_cipher_list() {
|
||||
let certs = TestCerts::generate();
|
||||
let mut cfg = build_tls(&certs, "server.pem", Some("server.key"));
|
||||
cfg.ciphers = Some("ECDHE-ECDSA-AES256-GCM-SHA384".to_string());
|
||||
assert!(tls::build_server_config(&cfg).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_cipher_list() {
|
||||
let certs = TestCerts::generate();
|
||||
let mut cfg = build_tls(&certs, "server.pem", Some("server.key"));
|
||||
cfg.ciphers = Some("THIS-IS-NOT-A-CIPHER".to_string());
|
||||
assert!(tls::build_server_config(&cfg).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_mismatched_cert_and_key() {
|
||||
// check_private_key must reject a key that does not match the certificate.
|
||||
let certs = TestCerts::generate();
|
||||
let tls = build_tls(&certs, "client.pem", Some("server.key"));
|
||||
assert!(tls::build_server_config(&tls).is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn https_request_succeeds_over_tls() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 0))).await;
|
||||
let response = https_get(&certs, &addr, None).await.expect("https request");
|
||||
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn serves_full_certificate_chain() {
|
||||
// Cert file holds leaf + intermediate; a client trusting only the root can
|
||||
// verify only if the server sends the intermediate, guarding against a
|
||||
// leaf-only load.
|
||||
let certs = TestCerts::generate();
|
||||
let tls = TlsConfig {
|
||||
cert_file: Some(certs.path("server_chain.pem")),
|
||||
key_file: Some(certs.path("server_chain.key")),
|
||||
ca_certs: None,
|
||||
cert_reqs: 0,
|
||||
ciphers: None,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server(Some(tls)).await;
|
||||
let response = https_get(&certs, &addr, None).await.expect("chained https request");
|
||||
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_legacy_cipher_only_client() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 0))).await;
|
||||
let result = legacy_suite_handshake(&certs, &addr).await;
|
||||
assert!(result.is_err(), "legacy-only client must be rejected");
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ssl_ciphers_override_widens_past_preset() {
|
||||
// Counterpart to rejects_legacy_cipher_only_client: --ssl-ciphers set to that
|
||||
// same legacy suite lets the client through, proving the override beats the preset.
|
||||
let certs = TestCerts::generate();
|
||||
let mut tls = server_tls(&certs, 0);
|
||||
tls.ciphers = Some("ECDHE-ECDSA-AES256-SHA:@SECLEVEL=0".to_string());
|
||||
let (addr, shutdown) = spawn_server(Some(tls)).await;
|
||||
let result = legacy_suite_handshake(&certs, &addr).await;
|
||||
assert!(
|
||||
result.is_ok(),
|
||||
"override must allow the legacy suite: {result:?}"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mtls_required_rejects_client_without_certificate() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 2))).await;
|
||||
let result = https_get(&certs, &addr, None).await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"handshake must fail without a client certificate"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mtls_required_accepts_valid_client_certificate() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 2))).await;
|
||||
let response = https_get(&certs, &addr, Some("client")).await.expect("mtls request");
|
||||
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mtls_optional_allows_anonymous_and_authenticated() {
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 1))).await;
|
||||
let anonymous = https_get(&certs, &addr, None).await.expect("anonymous request");
|
||||
assert!(anonymous.starts_with("HTTP/1.1 200"), "{anonymous}");
|
||||
let authenticated =
|
||||
https_get(&certs, &addr, Some("client")).await.expect("authenticated request");
|
||||
assert!(authenticated.starts_with("HTTP/1.1 200"), "{authenticated}");
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mtls_rejects_untrusted_client_certificate() {
|
||||
// Optional (1) still verifies a presented cert, so a self-signed cert not
|
||||
// chained to the CA is rejected in both modes, not just required (2).
|
||||
let certs = TestCerts::generate();
|
||||
for cert_reqs in [1, 2] {
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, cert_reqs))).await;
|
||||
let result = https_get(&certs, &addr, Some("untrusted_client")).await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"cert_reqs={cert_reqs}: untrusted client cert must be rejected"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn plain_http_serves_when_tls_is_disabled() {
|
||||
let (addr, shutdown) = spawn_server(None).await;
|
||||
let response = plain_get(&addr).await.expect("http request");
|
||||
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tls_handshake_timeout_drops_silent_client() {
|
||||
// Silent client (no ClientHello) must be dropped at the handshake deadline.
|
||||
let certs = TestCerts::generate();
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_millis(150),
|
||||
header_read: Duration::from_secs(5),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server_with_timeouts(Some(server_tls(&certs, 0)), timeouts).await;
|
||||
|
||||
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
|
||||
let mut buf = [0u8; 1];
|
||||
let read = tokio::time::timeout(Duration::from_secs(5), tcp.read(&mut buf)).await;
|
||||
assert!(
|
||||
matches!(read, Ok(Ok(0)) | Ok(Err(_))),
|
||||
"server must drop a stalled TLS handshake (expected close, got {read:?})"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn keep_alive_timeout_closes_idle_connection() {
|
||||
// Idle keep-alive connection must be closed at the deadline.
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
|
||||
|
||||
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
|
||||
// No `Connection: close`, so it stays alive until the idle deadline.
|
||||
tcp.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n")
|
||||
.await
|
||||
.expect("write request");
|
||||
|
||||
let drained = tokio::time::timeout(Duration::from_secs(5), async {
|
||||
let mut buf = [0u8; 1024];
|
||||
loop {
|
||||
match tcp.read(&mut buf).await {
|
||||
Ok(0) => return Ok(()),
|
||||
Ok(_) => continue,
|
||||
Err(err) => return Err(err),
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
assert!(
|
||||
matches!(drained, Ok(Ok(()))),
|
||||
"server must close an idle keep-alive connection (got {drained:?})"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn keep_alive_timeout_closes_idle_tls_connection() {
|
||||
// The keep-alive idle bound lives in serve_connections, below TLS; assert it
|
||||
// still fires through tls-listener's post-handshake SslStream, not just plaintext.
|
||||
let certs = TestCerts::generate();
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server_with_timeouts(Some(server_tls(&certs, 0)), timeouts).await;
|
||||
|
||||
let mut stream = connect_tls(&certs, &addr, None).await.expect("handshake");
|
||||
// No `Connection: close`, so the connection stays alive until the idle deadline.
|
||||
stream
|
||||
.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n")
|
||||
.await
|
||||
.expect("write request");
|
||||
|
||||
let closed = tokio::time::timeout(Duration::from_secs(5), async {
|
||||
let mut buf = [0u8; 1024];
|
||||
loop {
|
||||
// A clean close_notify (Ok(0)) or an abrupt TLS EOF both mean the
|
||||
// server closed; only the outer timeout (still open) is a failure.
|
||||
match stream.read(&mut buf).await {
|
||||
Ok(0) | Err(_) => break,
|
||||
Ok(_) => continue,
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
assert!(
|
||||
closed.is_ok(),
|
||||
"server must close an idle keep-alive TLS connection at the deadline"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn idle_timeout_closes_silent_client() {
|
||||
// Silent client closed by the header-read timeout (http1-only arms it from byte 0).
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: true,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
|
||||
|
||||
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
|
||||
let mut buf = [0u8; 1];
|
||||
let read = tokio::time::timeout(Duration::from_secs(5), tcp.read(&mut buf)).await;
|
||||
assert!(
|
||||
matches!(read, Ok(Ok(0)) | Ok(Err(_))),
|
||||
"server must close a silent client (expected close, got {read:?})"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn keep_alive_zero_disables_keep_alive() {
|
||||
// 0 disables keep-alive (serve, then close), like uvicorn's timeout_keep_alive=0.
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_secs(5),
|
||||
keep_alive_enabled: false,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
|
||||
|
||||
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
|
||||
tcp.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n")
|
||||
.await
|
||||
.expect("write request");
|
||||
|
||||
let mut response = String::new();
|
||||
let read =
|
||||
tokio::time::timeout(Duration::from_secs(5), tcp.read_to_string(&mut response)).await;
|
||||
assert!(
|
||||
read.is_ok(),
|
||||
"server must close after one response, not hang"
|
||||
);
|
||||
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
|
||||
// Assert `Connection: close`, not just 200: a 0 header-read timeout would also
|
||||
// serve an immediate request, so 200 alone wouldn't prove keep-alive is off.
|
||||
assert!(
|
||||
response.to_ascii_lowercase().contains("connection: close"),
|
||||
"keep-alive must be disabled (expected Connection: close): {response}"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn disabled_keep_alive_still_closes_silent_client() {
|
||||
// Even with keep-alive off, the head read stays bounded, so a silent client
|
||||
// is dropped rather than held open.
|
||||
let timeouts = ConnectionTimeouts {
|
||||
handshake: Duration::from_secs(60),
|
||||
header_read: Duration::from_millis(150),
|
||||
keep_alive_enabled: false,
|
||||
};
|
||||
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
|
||||
|
||||
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
|
||||
let mut buf = [0u8; 1];
|
||||
let read = tokio::time::timeout(Duration::from_secs(5), tcp.read(&mut buf)).await;
|
||||
assert!(
|
||||
matches!(read, Ok(Ok(0)) | Ok(Err(_))),
|
||||
"disabled keep-alive must still close a silent client (got {read:?})"
|
||||
);
|
||||
shutdown.cancel();
|
||||
}
|
||||
@@ -23,7 +23,7 @@ def test_python_error():
|
||||
error happening from the C++ side.
|
||||
"""
|
||||
allocator = get_mem_allocator_instance()
|
||||
total_bytes = current_platform.mem_get_info()[1]
|
||||
total_bytes = torch.accelerator.get_memory_info()[1]
|
||||
alloc_bytes = int(total_bytes * 0.7)
|
||||
tensors = []
|
||||
with allocator.use_memory_pool():
|
||||
@@ -64,9 +64,9 @@ def test_basic_cumem():
|
||||
output = x + y + z
|
||||
assert torch.allclose(output, torch.ones_like(output) * 3)
|
||||
|
||||
free_bytes = current_platform.mem_get_info()[0]
|
||||
free_bytes = torch.accelerator.get_memory_info()[0]
|
||||
allocator.sleep()
|
||||
free_bytes_after_sleep = current_platform.mem_get_info()[0]
|
||||
free_bytes_after_sleep = torch.accelerator.get_memory_info()[0]
|
||||
assert free_bytes_after_sleep > free_bytes
|
||||
allocator.wake_up()
|
||||
|
||||
@@ -99,9 +99,9 @@ def test_cumem_with_cudagraph():
|
||||
with torch.cuda.graph(model_graph):
|
||||
y = model(x)
|
||||
|
||||
free_bytes = current_platform.mem_get_info()[0]
|
||||
free_bytes = torch.accelerator.get_memory_info()[0]
|
||||
allocator.sleep()
|
||||
free_bytes_after_sleep = current_platform.mem_get_info()[0]
|
||||
free_bytes_after_sleep = torch.accelerator.get_memory_info()[0]
|
||||
assert free_bytes_after_sleep > free_bytes
|
||||
allocator.wake_up()
|
||||
|
||||
@@ -132,7 +132,7 @@ def test_cumem_with_cudagraph():
|
||||
],
|
||||
)
|
||||
def test_end_to_end(model: str):
|
||||
free, total = current_platform.mem_get_info()
|
||||
free, total = torch.accelerator.get_memory_info()
|
||||
used_bytes_baseline = total - free # in case other process is running
|
||||
llm = LLM(model, enable_sleep_mode=True)
|
||||
prompt = "How are you?"
|
||||
@@ -144,7 +144,7 @@ def test_end_to_end(model: str):
|
||||
# test sleep level 1 here.
|
||||
llm.sleep(level=1)
|
||||
|
||||
free_gpu_bytes_after_sleep, total = current_platform.mem_get_info()
|
||||
free_gpu_bytes_after_sleep, total = torch.accelerator.get_memory_info()
|
||||
used_bytes = total - free_gpu_bytes_after_sleep - used_bytes_baseline
|
||||
# now the memory usage is mostly cudagraph memory pool,
|
||||
# and it should be less than the model weights (1B model, 2GiB weights)
|
||||
@@ -164,7 +164,7 @@ def test_end_to_end(model: str):
|
||||
llm.sleep(level=1)
|
||||
llm.wake_up(tags=["weights"])
|
||||
|
||||
free_gpu_bytes_wake_up_w, total = current_platform.mem_get_info()
|
||||
free_gpu_bytes_wake_up_w, total = torch.accelerator.get_memory_info()
|
||||
used_bytes = total - free_gpu_bytes_wake_up_w - used_bytes_baseline
|
||||
|
||||
# should just reallocate memory for weights (1B model, ~2GiB weights)
|
||||
@@ -181,7 +181,7 @@ def test_end_to_end(model: str):
|
||||
@create_new_process_for_each_test()
|
||||
def test_deep_sleep():
|
||||
model = "hmellor/tiny-random-LlamaForCausalLM"
|
||||
free, total = current_platform.mem_get_info()
|
||||
free, total = torch.accelerator.get_memory_info()
|
||||
used_bytes_baseline = total - free # in case other process is running
|
||||
llm = LLM(model, enable_sleep_mode=True)
|
||||
prompt = "How are you?"
|
||||
@@ -191,13 +191,13 @@ def test_deep_sleep():
|
||||
# Put the engine to deep sleep
|
||||
llm.sleep(level=2)
|
||||
|
||||
free_gpu_bytes_after_sleep, total = current_platform.mem_get_info()
|
||||
free_gpu_bytes_after_sleep, total = torch.accelerator.get_memory_info()
|
||||
used_bytes = total - free_gpu_bytes_after_sleep - used_bytes_baseline
|
||||
assert used_bytes < 3 * GiB_bytes
|
||||
|
||||
llm.wake_up(tags=["weights"])
|
||||
llm.collective_rpc("reload_weights")
|
||||
free_gpu_bytes_wake_up_w, total = current_platform.mem_get_info()
|
||||
free_gpu_bytes_wake_up_w, total = torch.accelerator.get_memory_info()
|
||||
used_bytes = total - free_gpu_bytes_wake_up_w - used_bytes_baseline
|
||||
assert used_bytes < 4 * GiB_bytes
|
||||
|
||||
@@ -213,7 +213,7 @@ def test_deep_sleep():
|
||||
def test_deep_sleep_async():
|
||||
async def test():
|
||||
model = "hmellor/tiny-random-LlamaForCausalLM"
|
||||
free, total = current_platform.mem_get_info()
|
||||
free, total = torch.accelerator.get_memory_info()
|
||||
used_bytes_baseline = total - free # in case other process is running
|
||||
engine_args = AsyncEngineArgs(
|
||||
model=model,
|
||||
@@ -232,7 +232,7 @@ def test_deep_sleep_async():
|
||||
|
||||
await llm.wake_up(tags=["weights"])
|
||||
await llm.collective_rpc("reload_weights")
|
||||
free_gpu_bytes_wake_up_w, total = current_platform.mem_get_info()
|
||||
free_gpu_bytes_wake_up_w, total = torch.accelerator.get_memory_info()
|
||||
used_bytes = total - free_gpu_bytes_wake_up_w - used_bytes_baseline
|
||||
assert used_bytes < 4 * GiB_bytes
|
||||
|
||||
|
||||
@@ -29,9 +29,43 @@ from vllm.distributed.weight_transfer.nccl_engine import (
|
||||
NCCLWeightTransferInitInfo,
|
||||
NCCLWeightTransferUpdateInfo,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.network_utils import get_open_port
|
||||
|
||||
|
||||
def _weight_transfer_ray_env_vars() -> dict[str, str]:
|
||||
if not current_platform.is_rocm():
|
||||
return {}
|
||||
|
||||
return {
|
||||
"RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES": "1",
|
||||
"RAY_EXPERIMENTAL_NOSET_HIP_VISIBLE_DEVICES": "1",
|
||||
"RAY_EXPERIMENTAL_NOSET_ROCR_VISIBLE_DEVICES": "1",
|
||||
}
|
||||
|
||||
|
||||
def _init_ray_for_weight_transfer() -> None:
|
||||
if ray.is_initialized():
|
||||
return
|
||||
ray.init(
|
||||
ignore_reinit_error=True,
|
||||
runtime_env={"env_vars": _weight_transfer_ray_env_vars()},
|
||||
)
|
||||
|
||||
|
||||
def _get_ray_assigned_device() -> torch.device:
|
||||
gpu_ids = ray.get_gpu_ids()
|
||||
if not gpu_ids:
|
||||
return torch.device("cuda:0")
|
||||
return torch.device(f"cuda:{int(gpu_ids[0])}")
|
||||
|
||||
|
||||
def _set_ray_assigned_device() -> torch.device:
|
||||
device = _get_ray_assigned_device()
|
||||
torch.accelerator.set_device(device)
|
||||
return device
|
||||
|
||||
|
||||
def create_mock_parallel_config(
|
||||
rank: int = 0,
|
||||
world_size: int = 1,
|
||||
@@ -321,6 +355,8 @@ def trainer_broadcast_tensor(
|
||||
"""Trainer task that broadcasts a tensor via NCCL."""
|
||||
import torch
|
||||
|
||||
device = _set_ray_assigned_device()
|
||||
|
||||
from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator
|
||||
from vllm.distributed.utils import StatelessProcessGroup
|
||||
|
||||
@@ -331,12 +367,11 @@ def trainer_broadcast_tensor(
|
||||
rank=0,
|
||||
world_size=world_size,
|
||||
)
|
||||
# Ray sets CUDA_VISIBLE_DEVICES, so device 0 is the assigned GPU
|
||||
comm = PyNcclCommunicator(pg, device=0)
|
||||
comm = PyNcclCommunicator(pg, device=device.index)
|
||||
|
||||
# Create and broadcast the tensor
|
||||
dtype = getattr(torch, tensor_dtype)
|
||||
tensor_to_send = torch.ones(tensor_shape, dtype=dtype, device="cuda:0")
|
||||
tensor_to_send = torch.ones(tensor_shape, dtype=dtype, device=device)
|
||||
comm.broadcast(tensor_to_send, src=0, stream=torch.cuda.current_stream())
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
@@ -356,6 +391,8 @@ def inference_receive_tensor(
|
||||
|
||||
import torch
|
||||
|
||||
_set_ray_assigned_device()
|
||||
|
||||
from vllm.config.parallel import ParallelConfig
|
||||
from vllm.config.weight_transfer import WeightTransferConfig
|
||||
from vllm.distributed.weight_transfer.nccl_engine import (
|
||||
@@ -435,7 +472,7 @@ def test_nccl_weight_transfer_between_processes():
|
||||
This test verifies that the NCCLWeightTransferEngine can receive
|
||||
tensors broadcast by a trainer process via NCCL.
|
||||
"""
|
||||
ray.init(ignore_reinit_error=True)
|
||||
_init_ray_for_weight_transfer()
|
||||
|
||||
master_address = "127.0.0.1"
|
||||
master_port = get_open_port()
|
||||
@@ -473,6 +510,8 @@ def trainer_broadcast_sparse_tensor(
|
||||
"""Trainer task that broadcasts sparse patches via NCCL."""
|
||||
import torch
|
||||
|
||||
device = _set_ray_assigned_device()
|
||||
|
||||
from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator
|
||||
from vllm.distributed.utils import StatelessProcessGroup
|
||||
from vllm.distributed.weight_transfer.base import SparseWeightPatch
|
||||
@@ -487,12 +526,12 @@ def trainer_broadcast_sparse_tensor(
|
||||
rank=0,
|
||||
world_size=world_size,
|
||||
)
|
||||
comm = PyNcclCommunicator(pg, device=0)
|
||||
comm = PyNcclCommunicator(pg, device=device.index)
|
||||
|
||||
patch = SparseWeightPatch(
|
||||
name="test.weight",
|
||||
indices=torch.tensor([1, 7, 25], dtype=torch.int32, device="cuda:0"),
|
||||
values=torch.tensor([10.0, 20.0, 30.0], dtype=torch.float32, device="cuda:0"),
|
||||
indices=torch.tensor([1, 7, 25], dtype=torch.int32, device=device),
|
||||
values=torch.tensor([10.0, 20.0, 30.0], dtype=torch.float32, device=device),
|
||||
)
|
||||
NCCLWeightTransferEngine.trainer_send_sparse_weights(
|
||||
iter([patch]),
|
||||
@@ -513,6 +552,8 @@ def inference_receive_sparse_tensor(
|
||||
|
||||
import torch
|
||||
|
||||
device = _set_ray_assigned_device()
|
||||
|
||||
from vllm.config.parallel import ParallelConfig
|
||||
from vllm.config.weight_transfer import WeightTransferConfig
|
||||
from vllm.distributed.weight_transfer.nccl_engine import (
|
||||
@@ -540,7 +581,7 @@ def inference_receive_sparse_tensor(
|
||||
)
|
||||
)
|
||||
|
||||
target = torch.zeros(30, dtype=torch.float32, device="cuda")
|
||||
target = torch.zeros(30, dtype=torch.float32, device=device)
|
||||
|
||||
def apply_sparse_patches(patches: list[SparseWeightPatch]):
|
||||
for patch in patches:
|
||||
@@ -556,9 +597,9 @@ def inference_receive_sparse_tensor(
|
||||
engine.receive_sparse_weights(update_info, apply_sparse_patches)
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
expected = torch.zeros(30, dtype=torch.float32, device="cuda")
|
||||
expected = torch.zeros(30, dtype=torch.float32, device=device)
|
||||
expected[[1, 7, 25]] = torch.tensor(
|
||||
[10.0, 20.0, 30.0], dtype=torch.float32, device="cuda"
|
||||
[10.0, 20.0, 30.0], dtype=torch.float32, device=device
|
||||
)
|
||||
success = torch.equal(target, expected)
|
||||
engine.shutdown()
|
||||
@@ -574,7 +615,7 @@ def inference_receive_sparse_tensor(
|
||||
)
|
||||
def test_nccl_sparse_weight_transfer_between_processes():
|
||||
"""Test NCCL sparse weight transfer from trainer to inference process."""
|
||||
ray.init(ignore_reinit_error=True)
|
||||
_init_ray_for_weight_transfer()
|
||||
|
||||
master_address = "127.0.0.1"
|
||||
master_port = get_open_port()
|
||||
@@ -933,16 +974,18 @@ class TrainerActor:
|
||||
"""Trainer actor that creates and holds CUDA IPC handles."""
|
||||
|
||||
def __init__(self, tensor_shape: list[int], tensor_dtype: str):
|
||||
device = _set_ray_assigned_device()
|
||||
|
||||
# Create tensor on GPU and keep it alive
|
||||
dtype = getattr(torch, tensor_dtype)
|
||||
self.tensor = torch.ones(tensor_shape, dtype=dtype, device="cuda:0")
|
||||
self.tensor = torch.ones(tensor_shape, dtype=dtype, device=device)
|
||||
self.tensor.fill_(42.0) # Fill with 42 to verify correct transfer
|
||||
|
||||
# Create IPC handle (tensor must stay alive for IPC to work)
|
||||
# reduce_tensor returns (rebuild_func, args); we only send args
|
||||
# since the receiver imports rebuild_cuda_tensor directly.
|
||||
_, ipc_args = reduce_tensor(self.tensor)
|
||||
gpu_uuid = get_physical_gpu_id(0)
|
||||
gpu_uuid = get_physical_gpu_id(device.index)
|
||||
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
@@ -974,6 +1017,8 @@ def inference_receive_ipc_tensor(
|
||||
|
||||
import torch
|
||||
|
||||
_set_ray_assigned_device()
|
||||
|
||||
from vllm.config.parallel import ParallelConfig
|
||||
from vllm.config.weight_transfer import WeightTransferConfig
|
||||
from vllm.distributed.weight_transfer.ipc_engine import (
|
||||
@@ -1072,7 +1117,7 @@ def test_ipc_weight_transfer_between_processes(mode: str):
|
||||
from ray.util.placement_group import placement_group
|
||||
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
|
||||
|
||||
ray.init(ignore_reinit_error=True)
|
||||
_init_ray_for_weight_transfer()
|
||||
|
||||
# Create a placement group to ensure both processes are on the same GPU
|
||||
# Use fractional GPUs so both tasks can share the same GPU bundle
|
||||
|
||||
@@ -76,3 +76,60 @@ async def test_chat_logit_bias_invalid(client):
|
||||
assert error.status_code == 400
|
||||
assert str(invalid_token_id) in error_message
|
||||
assert str(vocab_size) in error_message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_logit_bias_non_integer_key(client):
|
||||
"""Test that a non-integer logit_bias key is rejected with a clean,
|
||||
informative error instead of a raw 'invalid literal for int()' message."""
|
||||
with pytest.raises(openai.BadRequestError) as excinfo:
|
||||
await client.chat.completions.create(
|
||||
model=MODEL_NAME,
|
||||
messages=[{"role": "user", "content": "Testing invalid logit bias key"}],
|
||||
max_tokens=5,
|
||||
logit_bias={"not_a_token_id": 50},
|
||||
)
|
||||
|
||||
error = excinfo.value
|
||||
error_message = str(error)
|
||||
|
||||
assert error.status_code == 400
|
||||
assert "not_a_token_id" in error_message
|
||||
assert "logit_bias" in error_message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_logit_bias_non_numeric_value(client):
|
||||
"""Test that a non-numeric logit_bias value is rejected with a message
|
||||
that names the specific offending token, not just a generic TypeError."""
|
||||
with pytest.raises(openai.BadRequestError) as excinfo:
|
||||
await client.chat.completions.create(
|
||||
model=MODEL_NAME,
|
||||
messages=[{"role": "user", "content": "Testing invalid logit bias value"}],
|
||||
max_tokens=5,
|
||||
logit_bias={"1": "not_a_number"},
|
||||
)
|
||||
|
||||
error = excinfo.value
|
||||
error_message = str(error)
|
||||
|
||||
assert error.status_code == 400
|
||||
assert "logit_bias" in error_message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_logit_bias_multiple_non_integer_keys(client):
|
||||
"""Test that ALL invalid logit_bias keys are reported together,
|
||||
not just the first one encountered."""
|
||||
with pytest.raises(openai.BadRequestError) as excinfo:
|
||||
await client.chat.completions.create(
|
||||
model=MODEL_NAME,
|
||||
messages=[{"role": "user", "content": "Testing multiple bad keys"}],
|
||||
max_tokens=5,
|
||||
logit_bias={"bad1": 50.0, "bad2": 20.0},
|
||||
)
|
||||
|
||||
error_message = str(excinfo.value)
|
||||
assert excinfo.value.status_code == 400
|
||||
assert "bad1" in error_message
|
||||
assert "bad2" in error_message
|
||||
|
||||
@@ -142,6 +142,7 @@ class TestHarmonyToResponseOutput:
|
||||
)
|
||||
assert output_items[0].call_id.startswith("call_")
|
||||
assert output_items[0].id.startswith("fc_")
|
||||
assert output_items[0].status == "completed"
|
||||
|
||||
def test_commentary_with_python_recipient_creates_reasoning(self):
|
||||
"""Test that commentary with recipient='python' creates reasoning items."""
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
"""
|
||||
Unit tests for stop_token_ids propagation from default_sampling_params
|
||||
to SamplingParams in ChatCompletionRequest and CompletionRequest.
|
||||
|
||||
Regression test for https://github.com/vllm-project/vllm/issues/22519
|
||||
where gpt-oss model stop tokens (e.g., </call> = 200012) were loaded into
|
||||
default_sampling_params at server startup but silently discarded on every
|
||||
request because to_sampling_params() never fell back to defaults.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.completion.protocol import (
|
||||
CompletionRequest,
|
||||
)
|
||||
|
||||
|
||||
class TestChatCompletionStopTokenIds:
|
||||
"""Test stop_token_ids merging in ChatCompletionRequest.to_sampling_params()."""
|
||||
|
||||
@pytest.fixture
|
||||
def minimal_chat_request(self):
|
||||
return ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
)
|
||||
|
||||
def test_default_stop_token_ids_applied(self, minimal_chat_request):
|
||||
"""Server-default stop_token_ids are applied when client sends none."""
|
||||
default_sampling_params = {
|
||||
"stop_token_ids": [200012, 200002],
|
||||
}
|
||||
|
||||
sampling_params = minimal_chat_request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params=default_sampling_params,
|
||||
)
|
||||
|
||||
assert set(sampling_params.stop_token_ids) == {200012, 200002}
|
||||
|
||||
def test_client_stop_token_ids_merged_with_defaults(self):
|
||||
"""Client-specified stop_token_ids are merged with server defaults."""
|
||||
request = ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
stop_token_ids=[99999],
|
||||
)
|
||||
default_sampling_params = {
|
||||
"stop_token_ids": [200012, 200002],
|
||||
}
|
||||
|
||||
sampling_params = request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params=default_sampling_params,
|
||||
)
|
||||
|
||||
assert set(sampling_params.stop_token_ids) == {200012, 200002, 99999}
|
||||
assert sampling_params.stop_token_ids == [99999, 200012, 200002]
|
||||
|
||||
def test_no_stop_token_ids_anywhere(self, minimal_chat_request):
|
||||
"""When neither client nor server specifies stop_token_ids, result is empty."""
|
||||
sampling_params = minimal_chat_request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params={},
|
||||
)
|
||||
|
||||
assert not sampling_params.stop_token_ids
|
||||
|
||||
def test_only_client_stop_token_ids(self):
|
||||
"""Client stop_token_ids work when no server defaults exist."""
|
||||
request = ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
stop_token_ids=[42, 43],
|
||||
)
|
||||
|
||||
sampling_params = request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params={},
|
||||
)
|
||||
|
||||
assert set(sampling_params.stop_token_ids) == {42, 43}
|
||||
|
||||
def test_duplicate_stop_token_ids_deduplicated(self):
|
||||
"""Overlapping stop_token_ids between client and server are deduplicated."""
|
||||
request = ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
stop_token_ids=[200012, 55555],
|
||||
)
|
||||
default_sampling_params = {
|
||||
"stop_token_ids": [200012, 200002],
|
||||
}
|
||||
|
||||
sampling_params = request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params=default_sampling_params,
|
||||
)
|
||||
|
||||
assert set(sampling_params.stop_token_ids) == {200012, 200002, 55555}
|
||||
assert sampling_params.stop_token_ids == [200012, 55555, 200002]
|
||||
assert len(sampling_params.stop_token_ids) == 3
|
||||
|
||||
|
||||
class TestCompletionStopTokenIds:
|
||||
"""Test stop_token_ids merging in CompletionRequest.to_sampling_params()."""
|
||||
|
||||
@pytest.fixture
|
||||
def minimal_completion_request(self):
|
||||
return CompletionRequest(
|
||||
model="test-model",
|
||||
prompt="hello",
|
||||
)
|
||||
|
||||
def test_default_stop_token_ids_applied(self, minimal_completion_request):
|
||||
"""Server-default stop_token_ids are applied when client sends none."""
|
||||
default_sampling_params = {
|
||||
"stop_token_ids": [200012, 200002],
|
||||
}
|
||||
|
||||
sampling_params = minimal_completion_request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params=default_sampling_params,
|
||||
)
|
||||
|
||||
assert set(sampling_params.stop_token_ids) == {200012, 200002}
|
||||
|
||||
def test_client_stop_token_ids_merged_with_defaults(self):
|
||||
"""Client-specified stop_token_ids are merged with server defaults."""
|
||||
request = CompletionRequest(
|
||||
model="test-model",
|
||||
prompt="hello",
|
||||
stop_token_ids=[99999],
|
||||
)
|
||||
default_sampling_params = {
|
||||
"stop_token_ids": [200012, 200002],
|
||||
}
|
||||
|
||||
sampling_params = request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params=default_sampling_params,
|
||||
)
|
||||
|
||||
assert set(sampling_params.stop_token_ids) == {200012, 200002, 99999}
|
||||
assert sampling_params.stop_token_ids == [99999, 200012, 200002]
|
||||
|
||||
def test_no_stop_token_ids_anywhere(self, minimal_completion_request):
|
||||
"""When neither client nor server specifies stop_token_ids, result is empty."""
|
||||
sampling_params = minimal_completion_request.to_sampling_params(
|
||||
max_tokens=100,
|
||||
default_sampling_params={},
|
||||
)
|
||||
|
||||
assert not sampling_params.stop_token_ids
|
||||
@@ -347,7 +347,7 @@ def test_selective_state_update(dim, dstate, has_z, itype):
|
||||
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
|
||||
if itype == torch.bfloat16:
|
||||
rtol, atol = 1e-2, 5e-2
|
||||
if torch.version.hip:
|
||||
if current_platform.is_rocm() or current_platform.is_xpu():
|
||||
atol *= 2
|
||||
# set seed
|
||||
set_random_seed(0)
|
||||
@@ -437,7 +437,7 @@ def test_selective_state_update_varlen(dim, dstate, has_z, itype, max_seq_len):
|
||||
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
|
||||
if itype == torch.bfloat16:
|
||||
rtol, atol = 5e-2, 1.5e-1
|
||||
if torch.version.hip:
|
||||
if current_platform.is_rocm() or current_platform.is_xpu():
|
||||
atol *= 2
|
||||
# set seed
|
||||
set_random_seed(0)
|
||||
@@ -700,7 +700,7 @@ def test_selective_state_update_with_batch_indices(
|
||||
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
|
||||
if itype == torch.bfloat16:
|
||||
rtol, atol = 1e-1, 1e-1
|
||||
if torch.version.hip:
|
||||
if current_platform.is_rocm() or current_platform.is_xpu():
|
||||
atol *= 2
|
||||
# set seed
|
||||
torch.random.manual_seed(0)
|
||||
@@ -865,7 +865,7 @@ def test_selective_state_update_with_num_accepted_tokens(
|
||||
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
|
||||
if itype == torch.bfloat16:
|
||||
rtol, atol = 5e-2, 1.5e-1
|
||||
if torch.version.hip:
|
||||
if current_platform.is_rocm() or current_platform.is_xpu():
|
||||
atol *= 2
|
||||
|
||||
set_random_seed(0)
|
||||
@@ -991,7 +991,7 @@ def test_selective_state_update_varlen_with_num_accepted(
|
||||
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
|
||||
if itype == torch.bfloat16:
|
||||
rtol, atol = 5e-2, 1.5e-1
|
||||
if torch.version.hip:
|
||||
if current_platform.is_rocm() or current_platform.is_xpu():
|
||||
atol *= 2
|
||||
|
||||
set_random_seed(0)
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Equivalence test for ``precopy_mamba_align_fused_kernel``.
|
||||
|
||||
The V2 "align" pre-copy must migrate mamba state across block boundaries with
|
||||
byte-identical semantics to the V1 copy specs (``get_conv_copy_spec`` /
|
||||
``get_temporal_copy_spec``):
|
||||
|
||||
* conv state (SD layout, conv_width > 0): shift the sliding window by
|
||||
``token_bias`` tokens -- ``state[bt[src_col], token_bias:]`` ->
|
||||
``state[bt[dst_col], :conv_width - token_bias]``.
|
||||
* temporal state (conv_width == 0): ``token_bias`` selects the accepted
|
||||
speculative column -- ``state[bt[src_col + token_bias]]`` ->
|
||||
``state[bt[dst_col]]``.
|
||||
|
||||
The kernel must also no-op when ``src_col < 0`` (fresh request) or
|
||||
``src_col == dst_col`` (no boundary crossed).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.v1.worker.mamba_utils import precopy_mamba_align_fused_kernel
|
||||
|
||||
try:
|
||||
import pytest
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not current_platform.is_cuda(),
|
||||
reason="precopy_mamba_align_fused_kernel needs CUDA/Triton",
|
||||
)
|
||||
_parametrize = pytest.mark.parametrize
|
||||
except ModuleNotFoundError: # allow running directly as ``python <thisfile>``
|
||||
pytest = None
|
||||
|
||||
def _parametrize(_name, _values):
|
||||
def _deco(fn):
|
||||
return fn
|
||||
|
||||
return _deco
|
||||
|
||||
|
||||
NUM_LAYERS = 3
|
||||
CONV_WIDTH = 4 # conv_kernel - 1 + num_spec
|
||||
CONV_DIM = 96
|
||||
SSM_SHAPE = (4, 16, 16)
|
||||
MAX_COLS = 8
|
||||
|
||||
|
||||
def _build_state(num_blocks, device):
|
||||
"""Per-layer (conv SD [nb, width, dim] bf16, ssm [nb, *shape] fp32) pools."""
|
||||
convs, ssms = [], []
|
||||
for _ in range(NUM_LAYERS):
|
||||
convs.append(
|
||||
torch.randn(
|
||||
num_blocks, CONV_WIDTH, CONV_DIM, dtype=torch.bfloat16, device=device
|
||||
)
|
||||
)
|
||||
ssms.append(
|
||||
torch.randn(num_blocks, *SSM_SHAPE, dtype=torch.float32, device=device)
|
||||
)
|
||||
return convs, ssms
|
||||
|
||||
|
||||
def _build_meta(convs, ssms, device):
|
||||
"""Flattened per-(layer, state-type) metadata, ordered conv, ssm per layer."""
|
||||
n = NUM_LAYERS * 2
|
||||
base = torch.zeros(n, dtype=torch.int64, device=device)
|
||||
blk_stride = torch.zeros(n, dtype=torch.int64, device=device)
|
||||
elem = torch.zeros(n, dtype=torch.int32, device=device)
|
||||
inner = torch.zeros(n, dtype=torch.int64, device=device)
|
||||
width = torch.zeros(n, dtype=torch.int32, device=device)
|
||||
group = torch.zeros(n, dtype=torch.int32, device=device)
|
||||
drc = torch.zeros(n, dtype=torch.int32, device=device) # DS rows (unused, SD)
|
||||
drs = torch.zeros(n, dtype=torch.int64, device=device)
|
||||
i = 0
|
||||
for layer in range(NUM_LAYERS):
|
||||
conv, ssm = convs[layer], ssms[layer]
|
||||
# conv (SD): width = size(1), inner = stride(1)
|
||||
base[i] = conv.data_ptr()
|
||||
blk_stride[i] = conv.stride(0) * conv.element_size()
|
||||
elem[i] = conv.element_size()
|
||||
width[i] = conv.size(1)
|
||||
inner[i] = conv.stride(1)
|
||||
i += 1
|
||||
# ssm (temporal): width = 0, inner = elems per block
|
||||
base[i] = ssm.data_ptr()
|
||||
blk_stride[i] = ssm.stride(0) * ssm.element_size()
|
||||
elem[i] = ssm.element_size()
|
||||
width[i] = 0
|
||||
inner[i] = ssm[0].numel()
|
||||
i += 1
|
||||
return base, blk_stride, elem, inner, width, group, drc, drs
|
||||
|
||||
|
||||
def _reference(convs, ssms, bt, src_col, dst_col, bias, num_reqs):
|
||||
"""Apply the V1 copy semantics on clones, reading from the pre-copy state."""
|
||||
conv_pre = [c.clone() for c in convs]
|
||||
ssm_pre = [s.clone() for s in ssms]
|
||||
conv_ref = [c.clone() for c in convs]
|
||||
ssm_ref = [s.clone() for s in ssms]
|
||||
for r in range(num_reqs):
|
||||
sc, dc, tb = int(src_col[r]), int(dst_col[r]), int(bias[r])
|
||||
if sc < 0 or sc == dc:
|
||||
continue
|
||||
sblk, dblk = int(bt[r, sc]), int(bt[r, dc])
|
||||
tblk = int(bt[r, sc + tb]) # temporal src column shifted by bias
|
||||
for layer in range(NUM_LAYERS):
|
||||
conv_ref[layer][dblk, : CONV_WIDTH - tb] = conv_pre[layer][sblk, tb:]
|
||||
ssm_ref[layer][dblk] = ssm_pre[layer][tblk]
|
||||
return conv_ref, ssm_ref
|
||||
|
||||
|
||||
@_parametrize("num_reqs", [1, 4, 16])
|
||||
@_parametrize("token_bias", [0, 1, 2])
|
||||
def test_precopy_matches_v1_copy_specs(num_reqs, token_bias):
|
||||
device = torch.device("cuda")
|
||||
torch.manual_seed(0)
|
||||
# Distinct physical block per (req, col) so copies never alias.
|
||||
num_blocks = num_reqs * MAX_COLS + 1
|
||||
bt = torch.empty(num_reqs, MAX_COLS, dtype=torch.int32, device=device)
|
||||
for r in range(num_reqs):
|
||||
bt[r] = torch.arange(
|
||||
1 + r * MAX_COLS, 1 + (r + 1) * MAX_COLS, dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
# Per-req columns: req 0 fresh (src=-1, skip), req 1 same block (skip),
|
||||
# the rest cross from col 1 -> col 0 with the given spec token bias.
|
||||
src_col = torch.full((num_reqs,), 1, dtype=torch.int32, device=device)
|
||||
dst_col = torch.zeros(num_reqs, dtype=torch.int32, device=device)
|
||||
bias = torch.full((num_reqs,), token_bias, dtype=torch.int32, device=device)
|
||||
if num_reqs >= 1:
|
||||
src_col[0] = -1 # fresh -> no copy
|
||||
if num_reqs >= 2:
|
||||
dst_col[1] = 1 # src_col == dst_col -> no copy
|
||||
|
||||
convs, ssms = _build_state(num_blocks, device)
|
||||
conv_ref, ssm_ref = _reference(
|
||||
convs, ssms, bt.cpu(), src_col.cpu(), dst_col.cpu(), bias.cpu(), num_reqs
|
||||
)
|
||||
|
||||
base, blk_stride, elem, inner, width, group, drc, drs = _build_meta(
|
||||
convs, ssms, device
|
||||
)
|
||||
bt_ptrs = torch.tensor([bt.data_ptr()], dtype=torch.int64, device=device)
|
||||
idx_mapping = torch.arange(num_reqs, dtype=torch.int32, device=device)
|
||||
grid = (num_reqs, NUM_LAYERS * 2)
|
||||
precopy_mamba_align_fused_kernel[grid](
|
||||
dst_col,
|
||||
src_col,
|
||||
bias,
|
||||
bt_ptrs,
|
||||
bt.stride(0),
|
||||
base,
|
||||
blk_stride,
|
||||
elem,
|
||||
inner,
|
||||
width,
|
||||
group,
|
||||
drc,
|
||||
drs,
|
||||
idx_mapping,
|
||||
num_reqs,
|
||||
COPY_BLOCK_SIZE=1024,
|
||||
CONV_STATE_DIM_FIRST=False,
|
||||
)
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
for layer in range(NUM_LAYERS):
|
||||
torch.testing.assert_close(convs[layer], conv_ref[layer], rtol=0, atol=0)
|
||||
torch.testing.assert_close(ssms[layer], ssm_ref[layer], rtol=0, atol=0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
for nr in (1, 4, 16):
|
||||
for tb in (0, 1, 2):
|
||||
test_precopy_matches_v1_copy_specs(nr, tb)
|
||||
print(f"OK num_reqs={nr} token_bias={tb}")
|
||||
@@ -406,7 +406,7 @@ def test_fused_moe_int64_overflow(workspace_init):
|
||||
Reproduces the scenario from PR #34279.
|
||||
"""
|
||||
# ~12 GB GPU memory needed for intermediate caches
|
||||
free_mem = torch.cuda.mem_get_info()[0]
|
||||
free_mem = torch.accelerator.get_memory_info()[0]
|
||||
if free_mem < 12 * 1024**3:
|
||||
pytest.skip("Insufficient GPU memory for overflow test")
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ def test_sharded_state_loader(
|
||||
ctx = mp.get_context("spawn")
|
||||
|
||||
platform_args = {}
|
||||
if current_platform.is_rocm():
|
||||
if current_platform.is_rocm() or current_platform.is_xpu():
|
||||
platform_args["max_num_seqs"] = 1
|
||||
|
||||
# Run in separate processes for memory & CUDA isolation
|
||||
|
||||
@@ -83,7 +83,7 @@ def _ru_maxrss_bytes() -> int | None:
|
||||
|
||||
def _gpu_used_bytes() -> int:
|
||||
torch.accelerator.synchronize()
|
||||
free_bytes, total_bytes = current_platform.mem_get_info()
|
||||
free_bytes, total_bytes = torch.accelerator.get_memory_info()
|
||||
return int(total_bytes - free_bytes)
|
||||
|
||||
|
||||
|
||||
@@ -205,3 +205,57 @@ def test_image_media_io_load_file(tmp_path):
|
||||
|
||||
with pytest.raises(ValueError, match="Failed to load image"):
|
||||
image_io.load_file(truncated_real_file)
|
||||
|
||||
|
||||
def test_image_pixel_limit_respected():
|
||||
"""A small image within the pixel limit loads successfully."""
|
||||
import vllm.envs as envs
|
||||
|
||||
image = Image.new("RGB", (100, 100), (255, 0, 0))
|
||||
from io import BytesIO
|
||||
|
||||
buf = BytesIO()
|
||||
image.save(buf, format="PNG")
|
||||
data = buf.getvalue()
|
||||
|
||||
assert envs.VLLM_MAX_IMAGE_PIXELS >= 100 * 100
|
||||
|
||||
image_io = ImageMediaIO()
|
||||
result = image_io.load_bytes(data)
|
||||
assert result.media.size == (100, 100)
|
||||
|
||||
|
||||
def test_image_pixel_limit_rejected(monkeypatch):
|
||||
"""An image exceeding the pixel limit is rejected before raster decode."""
|
||||
import vllm.envs as envs
|
||||
|
||||
monkeypatch.setattr(envs, "VLLM_MAX_IMAGE_PIXELS", 100)
|
||||
|
||||
image = Image.new("RGB", (20, 20), (0, 255, 0))
|
||||
from io import BytesIO
|
||||
|
||||
buf = BytesIO()
|
||||
image.save(buf, format="PNG")
|
||||
data = buf.getvalue()
|
||||
|
||||
image_io = ImageMediaIO()
|
||||
with pytest.raises(ValueError, match="exceed"):
|
||||
image_io.load_bytes(data)
|
||||
|
||||
|
||||
def test_image_pixel_limit_disabled(monkeypatch):
|
||||
"""Setting VLLM_MAX_IMAGE_PIXELS=0 disables the pixel limit."""
|
||||
import vllm.envs as envs
|
||||
|
||||
monkeypatch.setattr(envs, "VLLM_MAX_IMAGE_PIXELS", 0)
|
||||
|
||||
image = Image.new("RGB", (1000, 1000), (0, 0, 255))
|
||||
from io import BytesIO
|
||||
|
||||
buf = BytesIO()
|
||||
image.save(buf, format="PNG")
|
||||
data = buf.getvalue()
|
||||
|
||||
image_io = ImageMediaIO()
|
||||
result = image_io.load_bytes(data)
|
||||
assert result.media.size == (1000, 1000)
|
||||
|
||||
@@ -31,6 +31,7 @@ from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
from vllm.parser.engine.registered_adapters import (
|
||||
Gemma4Parser,
|
||||
Glm47MoeParser,
|
||||
KimiK2Parser,
|
||||
MinimaxM2Parser,
|
||||
NemotronV3Parser,
|
||||
Qwen3Parser,
|
||||
@@ -717,6 +718,96 @@ def _build_glm47_moe(scenario: Scenario, validate: bool = True) -> Sample:
|
||||
return sample
|
||||
|
||||
|
||||
# ── Kimi K2 (native tool-call section, starts in REASONING) ──────────
|
||||
|
||||
_KIMI_K2_VOCAB: dict[str, int] = {
|
||||
"<think>": 50,
|
||||
"</think>": 51,
|
||||
"<|tool_calls_section_begin|>": 60,
|
||||
"<|tool_calls_section_end|>": 61,
|
||||
"<|tool_call_begin|>": 62,
|
||||
"<|tool_call_end|>": 63,
|
||||
"<|tool_call_argument_begin|>": 64,
|
||||
}
|
||||
|
||||
|
||||
def _kimi_k2_tool_segments(
|
||||
tool_calls: list[ToolCallSpec],
|
||||
) -> list[tuple[str, bool]]:
|
||||
segs: list[tuple[str, bool]] = [("<|tool_calls_section_begin|>", True)]
|
||||
for index, tc in enumerate(tool_calls):
|
||||
args = json.dumps(tc.arguments, ensure_ascii=False, separators=(",", ":"))
|
||||
segs.extend(
|
||||
[
|
||||
("<|tool_call_begin|>", True),
|
||||
(f"functions.{tc.name}:{index}\n", False),
|
||||
("<|tool_call_argument_begin|>", True),
|
||||
(args, False),
|
||||
("<|tool_call_end|>", True),
|
||||
]
|
||||
)
|
||||
segs.append(("<|tool_calls_section_end|>", True))
|
||||
return segs
|
||||
|
||||
|
||||
def _kimi_k2_segments(scenario: Scenario) -> list[tuple[str, bool]]:
|
||||
segs: list[tuple[str, bool]] = []
|
||||
if scenario.reasoning is not None:
|
||||
segs.append(("<think>", True))
|
||||
segs.append((scenario.reasoning, False))
|
||||
if scenario.content is not None or scenario.tool_calls is not None:
|
||||
segs.append(("</think>", True))
|
||||
if scenario.content is not None:
|
||||
segs.append((scenario.content, False))
|
||||
if scenario.tool_calls is not None:
|
||||
segs.extend(_kimi_k2_tool_segments(scenario.tool_calls))
|
||||
return segs
|
||||
|
||||
|
||||
def _build_kimi_k2(
|
||||
scenario: Scenario,
|
||||
validate: bool = True,
|
||||
thinking: bool = True,
|
||||
) -> Sample:
|
||||
expected_reasoning = (
|
||||
scenario.reasoning.rstrip()
|
||||
if (thinking and scenario.reasoning is not None)
|
||||
else None
|
||||
)
|
||||
if thinking and scenario.reasoning is None:
|
||||
expected_reasoning = ""
|
||||
|
||||
sample = _make_sample(
|
||||
sample_id=f"kimi_k2-{scenario.id}",
|
||||
description=scenario.description,
|
||||
vocab=_KIMI_K2_VOCAB,
|
||||
segments=_kimi_k2_segments(scenario),
|
||||
expected_reasoning=expected_reasoning,
|
||||
expected_content=_qwen3_expected_content(scenario),
|
||||
expected_tool_calls=_expected_tc(scenario),
|
||||
tools=_expected_tools(scenario),
|
||||
chat_template_kwargs=None if thinking else {"thinking": False},
|
||||
)
|
||||
if validate:
|
||||
_validate_sample(
|
||||
sample,
|
||||
KimiK2Parser,
|
||||
chat_template_kwargs=sample.chat_template_kwargs,
|
||||
)
|
||||
return sample
|
||||
|
||||
|
||||
_KIMI_K2_SCENARIOS = [
|
||||
*SCENARIOS,
|
||||
Scenario(
|
||||
id="trailing-reasoning-whitespace",
|
||||
description="Reasoning trailing whitespace is stripped",
|
||||
reasoning="Reasoning with trailing whitespace. \n\t",
|
||||
content="Done.",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
# ── Registry and public API ──────────────────────────────────────────
|
||||
|
||||
_BUILDERS: dict[str, Any] = {
|
||||
@@ -726,6 +817,7 @@ _BUILDERS: dict[str, Any] = {
|
||||
"nemotron_v3": _build_nemotron_v3,
|
||||
"seed_oss": _build_seed_oss,
|
||||
"glm47_moe": _build_glm47_moe,
|
||||
"kimi_k2": _build_kimi_k2,
|
||||
}
|
||||
|
||||
|
||||
@@ -733,7 +825,8 @@ _BUILDERS: dict[str, Any] = {
|
||||
def build_samples(model: str) -> tuple[Sample, ...]:
|
||||
"""Build all scenario samples for a model, self-validated."""
|
||||
builder = _BUILDERS[model]
|
||||
return tuple(builder(s) for s in SCENARIOS)
|
||||
scenarios = _KIMI_K2_SCENARIOS if model == "kimi_k2" else SCENARIOS
|
||||
return tuple(builder(s) for s in scenarios)
|
||||
|
||||
|
||||
def build_sample(model: str, scenario: Scenario) -> Sample:
|
||||
|
||||
@@ -7,7 +7,6 @@ import pytest
|
||||
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.reasoning.identity_reasoning_parser import IdentityReasoningParser
|
||||
from vllm.reasoning.kimi_k2_reasoning_parser import KimiK2ReasoningParser
|
||||
from vllm.tokenizers import get_tokenizer
|
||||
|
||||
@@ -33,20 +32,6 @@ def kimi_k2_tokenizer():
|
||||
return get_tokenizer(tokenizer_name=REASONING_MODEL_NAME, trust_remote_code=True)
|
||||
|
||||
|
||||
def test_parser_selection_thinking_enabled(kimi_k2_tokenizer):
|
||||
parser = KimiK2ReasoningParser(
|
||||
kimi_k2_tokenizer, chat_template_kwargs={"thinking": True}
|
||||
)
|
||||
assert parser._identity_parser is None
|
||||
|
||||
|
||||
def test_parser_selection_thinking_disabled(kimi_k2_tokenizer):
|
||||
parser = KimiK2ReasoningParser(
|
||||
kimi_k2_tokenizer, chat_template_kwargs={"thinking": False}
|
||||
)
|
||||
assert isinstance(parser._identity_parser, IdentityReasoningParser)
|
||||
|
||||
|
||||
def test_extract_reasoning_with_think_tags(kimi_k2_tokenizer):
|
||||
parser = KimiK2ReasoningParser(kimi_k2_tokenizer)
|
||||
request = ChatCompletionRequest(model="test-model", messages=[], temperature=1.0)
|
||||
@@ -65,7 +50,7 @@ def test_extract_reasoning_empty_thinking(kimi_k2_tokenizer):
|
||||
reasoning, content = parser.extract_reasoning(
|
||||
"<think></think>final answer", request
|
||||
)
|
||||
assert reasoning == ""
|
||||
assert reasoning is None
|
||||
assert content == "final answer"
|
||||
|
||||
|
||||
@@ -96,8 +81,8 @@ def test_streaming_reasoning_then_content(kimi_k2_tokenizer):
|
||||
"""Token-by-token streaming: reasoning tokens then content after </think>."""
|
||||
parser = KimiK2ReasoningParser(kimi_k2_tokenizer)
|
||||
|
||||
think_id = parser._start_token_id
|
||||
end_think_id = parser._end_token_id
|
||||
think_id = parser._parser_engine._start_token_id
|
||||
end_think_id = parser._parser_engine._end_token_id
|
||||
# Use a real token ID from the tokenizer for regular content
|
||||
regular_id = kimi_k2_tokenizer.encode("hello", add_special_tokens=False)[0]
|
||||
|
||||
@@ -154,8 +139,8 @@ def test_streaming_tool_section_ends_reasoning(kimi_k2_tokenizer):
|
||||
"""<|tool_calls_section_begin|> in delta ends reasoning during streaming."""
|
||||
parser = KimiK2ReasoningParser(kimi_k2_tokenizer)
|
||||
|
||||
think_id = parser._start_token_id
|
||||
tool_begin_id = parser._tool_section_start_token_id
|
||||
think_id = parser._parser_engine._start_token_id
|
||||
tool_begin_id = parser._parser_engine._tool_section_start_token_id
|
||||
regular_id = kimi_k2_tokenizer.encode("hello", add_special_tokens=False)[0]
|
||||
|
||||
# Tool section token arrives — should transition from reasoning to content
|
||||
@@ -169,50 +154,3 @@ def test_streaming_tool_section_ends_reasoning(kimi_k2_tokenizer):
|
||||
)
|
||||
assert isinstance(result, DeltaMessage)
|
||||
assert result.content == "<|tool_calls_section_begin|>"
|
||||
|
||||
|
||||
def test_streaming_end_token_id_buffered(mock_kimi_k2_tokenizer):
|
||||
"""When stop sequences buffer text, </think> ID arrives before its text.
|
||||
|
||||
The token ID is present in delta_token_ids but the actual string is not
|
||||
yet in delta_text (still buffered). The parser must return None to wait
|
||||
for the next delta, instead of calling find() which returns -1 and
|
||||
silently corrupting the text split.
|
||||
"""
|
||||
parser = KimiK2ReasoningParser(mock_kimi_k2_tokenizer)
|
||||
think_id = parser._start_token_id
|
||||
end_think_id = parser._end_token_id
|
||||
|
||||
# Simulate: </think> ID arrived but text not yet flushed.
|
||||
# Two token IDs in delta to bypass the single-special-token guard.
|
||||
result = parser.extract_reasoning_streaming(
|
||||
previous_text="some reasoning",
|
||||
current_text="some reasoning extra",
|
||||
delta_text="extra", # </think> text not yet flushed
|
||||
previous_token_ids=[think_id],
|
||||
current_token_ids=[think_id, end_think_id, 999],
|
||||
delta_token_ids=[end_think_id, 999],
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_streaming_tool_section_id_buffered(mock_kimi_k2_tokenizer):
|
||||
"""When stop sequences buffer text, tool section start ID arrives before its text.
|
||||
|
||||
Same buffering scenario as above but for <|tool_calls_section_begin|>.
|
||||
Without the guard, find() returns -1 and delta_text[:tool_index] silently
|
||||
drops the last character of reasoning.
|
||||
"""
|
||||
parser = KimiK2ReasoningParser(mock_kimi_k2_tokenizer)
|
||||
think_id = parser._start_token_id
|
||||
tool_begin_id = parser._tool_section_start_token_id
|
||||
|
||||
result = parser.extract_reasoning_streaming(
|
||||
previous_text="some reasoning",
|
||||
current_text="some reasoning extra",
|
||||
delta_text="extra", # tool section text not yet flushed
|
||||
previous_token_ids=[think_id],
|
||||
current_token_ids=[think_id, tool_begin_id, 999],
|
||||
delta_token_ids=[tool_begin_id, 999],
|
||||
)
|
||||
assert result is None
|
||||
|
||||
@@ -1292,7 +1292,9 @@ def test_vllm_config_explicit_overrides():
|
||||
compilation_config=compilation_config,
|
||||
)
|
||||
assert config.compilation_config.cudagraph_mode == CUDAGraphMode.NONE
|
||||
assert config.compilation_config.pass_config.enable_qk_norm_rope_fusion is True
|
||||
assert config.compilation_config.pass_config.enable_qk_norm_rope_fusion is (
|
||||
current_platform.is_cuda_alike() or current_platform.is_xpu()
|
||||
)
|
||||
# Mode should still use default for O2
|
||||
assert config.compilation_config.mode == CompilationMode.VLLM_COMPILE
|
||||
|
||||
|
||||
@@ -102,45 +102,36 @@ def test_get_model_structural_tag_supports_vllm_hermes(
|
||||
)
|
||||
|
||||
assert isinstance(tag, StructuralTag)
|
||||
assert tag.model_dump() == {
|
||||
"type": "structural_tag",
|
||||
"format": {
|
||||
"type": "tags_with_separator",
|
||||
"tags": [
|
||||
{
|
||||
"type": "tag",
|
||||
"begin": '<tool_call>\n{"name": "get_weather", "arguments": ',
|
||||
"content": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
"style": "json",
|
||||
},
|
||||
"end": "}\n</tool_call>",
|
||||
},
|
||||
{
|
||||
"type": "tag",
|
||||
"begin": '<tool_call>{"name": "get_weather", "arguments": ',
|
||||
"content": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
"style": "json",
|
||||
},
|
||||
"end": "}</tool_call>",
|
||||
},
|
||||
],
|
||||
"separator": "",
|
||||
"at_least_one": True,
|
||||
"stop_after_first": False,
|
||||
},
|
||||
|
||||
# Assert the semantically meaningful structure rather than the full
|
||||
# model_dump(), which gains version-specific keys across xgrammar releases
|
||||
# (e.g. "any_order" was added to json_schema content in 0.2.3).
|
||||
dump = tag.model_dump()
|
||||
assert dump["type"] == "structural_tag"
|
||||
|
||||
fmt = dump["format"]
|
||||
assert fmt["type"] == "tags_with_separator"
|
||||
assert fmt["separator"] == ""
|
||||
assert fmt["at_least_one"] is True
|
||||
assert fmt["stop_after_first"] is False
|
||||
|
||||
expected_schema = {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
}
|
||||
expected_tags = [
|
||||
('<tool_call>\n{"name": "get_weather", "arguments": ', "}\n</tool_call>"),
|
||||
('<tool_call>{"name": "get_weather", "arguments": ', "}</tool_call>"),
|
||||
]
|
||||
assert len(fmt["tags"]) == len(expected_tags)
|
||||
for tag_dump, (begin, end) in zip(fmt["tags"], expected_tags):
|
||||
assert tag_dump["type"] == "tag"
|
||||
assert tag_dump["begin"] == begin
|
||||
assert tag_dump["end"] == end
|
||||
content = tag_dump["content"]
|
||||
assert content["type"] == "json_schema"
|
||||
assert content["json_schema"] == expected_schema
|
||||
|
||||
|
||||
def test_hermes_required_tool_calls_use_empty_separator():
|
||||
|
||||
+53
-4
@@ -604,6 +604,13 @@ class RemoteVLLMServer:
|
||||
mem_info = nvmlDeviceGetMemoryInfo(handle)
|
||||
total_used += mem_info.used
|
||||
return total_used
|
||||
elif current_platform.is_xpu():
|
||||
total_used = 0
|
||||
device_count = current_platform.device_count()
|
||||
for i in range(device_count):
|
||||
free, total = torch.xpu.mem_get_info(i)
|
||||
total_used += total - free
|
||||
return total_used
|
||||
except Exception as e:
|
||||
print(f"[RemoteOpenAIServer] Could not query GPU memory: {e}")
|
||||
return None
|
||||
@@ -1501,6 +1508,9 @@ def wait_for_gpu_memory_to_clear(
|
||||
threshold_bytes: int | dict[int, int] | None = None,
|
||||
threshold_ratio: float | dict[int, float] | None = None,
|
||||
timeout_s: float = 120,
|
||||
stable_duration_s: float = 0,
|
||||
stable_tolerance_bytes: int = 512 * 1024**2,
|
||||
poll_interval_s: float = 5,
|
||||
) -> None:
|
||||
assert threshold_bytes is not None or threshold_ratio is not None
|
||||
devices = get_physical_device_indices(devices)
|
||||
@@ -1528,8 +1538,13 @@ def wait_for_gpu_memory_to_clear(
|
||||
# Use nvml instead of pytorch to reduce measurement error from torch cuda
|
||||
# context.
|
||||
start_time = time.time()
|
||||
stable_since: float | None = None
|
||||
stable_used_bytes: dict[int, int] | None = None
|
||||
while True:
|
||||
output_raw = record_gpu_memory_usage_stats(devices=devices)
|
||||
used_bytes_by_device = {
|
||||
device: int(gb_used * 2**30) for device, (gb_used, _) in output_raw.items()
|
||||
}
|
||||
output = {
|
||||
device: f"{gb_used:.02f}/{gb_total:.02f}"
|
||||
for device, (gb_used, gb_total) in output_raw.items()
|
||||
@@ -1577,15 +1592,45 @@ def wait_for_gpu_memory_to_clear(
|
||||
|
||||
dur_s = time.time() - start_time
|
||||
if all_free:
|
||||
print(f"Done waiting for free GPU memory on ({threshold=}) {dur_s=:.02f}")
|
||||
break
|
||||
if stable_duration_s <= 0:
|
||||
print(
|
||||
f"Done waiting for free GPU memory on devices {devices=} "
|
||||
f"({threshold=}) {dur_s=:.02f}"
|
||||
)
|
||||
break
|
||||
|
||||
now = time.time()
|
||||
if stable_used_bytes is None:
|
||||
stable_since = now
|
||||
stable_used_bytes = used_bytes_by_device
|
||||
else:
|
||||
memory_changed = any(
|
||||
abs(used_bytes_by_device[device] - stable_used_bytes[device])
|
||||
> stable_tolerance_bytes
|
||||
for device in devices
|
||||
)
|
||||
if memory_changed:
|
||||
stable_since = now
|
||||
stable_used_bytes = used_bytes_by_device
|
||||
elif (
|
||||
stable_since is not None and now - stable_since >= stable_duration_s
|
||||
):
|
||||
print(
|
||||
f"Done waiting for stable free GPU memory on devices "
|
||||
f"{devices=} ({threshold=}) {dur_s=:.02f}"
|
||||
)
|
||||
break
|
||||
else:
|
||||
stable_since = None
|
||||
stable_used_bytes = None
|
||||
|
||||
if dur_s >= timeout_s:
|
||||
raise ValueError(
|
||||
f"Memory of devices not free after {dur_s=:.02f} ({threshold=})"
|
||||
f"Memory of devices {devices=} not free after "
|
||||
f"{dur_s=:.02f} ({threshold=})"
|
||||
)
|
||||
|
||||
time.sleep(5)
|
||||
time.sleep(poll_interval_s)
|
||||
|
||||
|
||||
def wait_for_rocm_memory_to_settle(
|
||||
@@ -1606,11 +1651,15 @@ def wait_for_rocm_memory_to_settle(
|
||||
num_gpus = current_platform.device_count()
|
||||
if num_gpus == 0:
|
||||
return
|
||||
if threshold_ratio is None:
|
||||
threshold_ratio = 0.1
|
||||
|
||||
wait_for_gpu_memory_to_clear(
|
||||
devices=list(range(num_gpus)),
|
||||
threshold_ratio=threshold_ratio,
|
||||
timeout_s=timeout_s,
|
||||
stable_duration_s=2.0,
|
||||
poll_interval_s=1.0,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -36,7 +36,7 @@ def test_memory_profiling():
|
||||
weights_memory = 128 * 1024 * 1024 * 4 # 512 MiB
|
||||
|
||||
def measure_current_non_torch():
|
||||
free, total = torch.cuda.mem_get_info()
|
||||
free, total = torch.accelerator.get_memory_info()
|
||||
current_used = total - free
|
||||
current_torch = torch.accelerator.memory_reserved()
|
||||
current_non_torch = current_used - current_torch
|
||||
@@ -81,8 +81,9 @@ def test_memory_snapshot_uses_psutil_on_integrated_gpu():
|
||||
with (
|
||||
patch("vllm.utils.mem_utils.current_platform") as mock_platform,
|
||||
patch("vllm.utils.mem_utils.psutil") as mock_psutil,
|
||||
patch("torch.accelerator") as mock_accelerator,
|
||||
):
|
||||
mock_platform.mem_get_info.return_value = (
|
||||
mock_accelerator.get_memory_info.return_value = (
|
||||
mock_cuda_free,
|
||||
mock_cuda_total,
|
||||
)
|
||||
@@ -90,8 +91,8 @@ def test_memory_snapshot_uses_psutil_on_integrated_gpu():
|
||||
mock_platform.memory_stats.return_value = {
|
||||
"allocated_bytes.all.peak": 0,
|
||||
}
|
||||
mock_platform.memory_reserved.return_value = 0
|
||||
mock_platform.current_device = lambda: "cuda:0"
|
||||
mock_accelerator.memory_reserved.return_value = 0
|
||||
mock_accelerator.current_device = lambda: "cuda:0"
|
||||
|
||||
mock_vmem = MagicMock()
|
||||
mock_vmem.available = mock_psutil_available
|
||||
@@ -105,24 +106,25 @@ def test_memory_snapshot_uses_psutil_on_integrated_gpu():
|
||||
|
||||
|
||||
def test_memory_snapshot_uses_cuda_on_discrete_gpu():
|
||||
"""On discrete GPUs, free_memory should come from CUDA mem_get_info."""
|
||||
"""On discrete GPUs, free_memory should come from accelerator get_memory_info."""
|
||||
mock_cuda_free = 70 * 1024**3
|
||||
mock_cuda_total = 80 * 1024**3
|
||||
|
||||
with (
|
||||
patch("vllm.utils.mem_utils.current_platform") as mock_platform,
|
||||
patch("vllm.utils.mem_utils.psutil") as mock_psutil,
|
||||
patch("torch.accelerator") as mock_accelerator,
|
||||
):
|
||||
mock_platform.mem_get_info.return_value = (
|
||||
mock_accelerator.get_memory_info.return_value = (
|
||||
mock_cuda_free,
|
||||
mock_cuda_total,
|
||||
)
|
||||
mock_platform.is_integrated_gpu.return_value = False
|
||||
mock_platform.memory_stats.return_value = {
|
||||
mock_accelerator.memory_stats.return_value = {
|
||||
"allocated_bytes.all.peak": 0,
|
||||
}
|
||||
mock_platform.memory_reserved.return_value = 0
|
||||
mock_platform.current_device = lambda: "cuda:0"
|
||||
mock_accelerator.memory_reserved.return_value = 0
|
||||
mock_accelerator.current_device = lambda: "cuda:0"
|
||||
|
||||
snapshot = MemorySnapshot(device="cuda:0")
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ import datasets
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import vllm.envs as envs
|
||||
from tests.utils import create_new_process_for_each_test
|
||||
from vllm import LLM, SamplingParams, TokensPrompt
|
||||
from vllm.config import CacheConfig
|
||||
@@ -494,12 +495,7 @@ def apply_patch(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(mamba_utils, "do_mamba_copy_block", fake_copy_fn)
|
||||
|
||||
|
||||
@create_new_process_for_each_test()
|
||||
def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
|
||||
run_ref_mamba_state_in_subprocess()
|
||||
apply_patch(monkeypatch)
|
||||
prompt_dataset = datasets.load_dataset("heheda/a_long_article")
|
||||
full_prompt = prompt_dataset["train"][0]["text"]
|
||||
def get_mamba_prefix_cache_step_configs() -> dict[str, TestConfig]:
|
||||
tests = {
|
||||
"accept_1": TestConfig(
|
||||
num_prompt_tokens=554,
|
||||
@@ -731,6 +727,27 @@ def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
|
||||
),
|
||||
}
|
||||
|
||||
return tests
|
||||
|
||||
|
||||
def fill_following_kv_cache_block_ids(test_config: TestConfig) -> None:
|
||||
for step_action_prev, step_action_next in zip(
|
||||
test_config.step_actions[:-1], test_config.step_actions[1:]
|
||||
):
|
||||
if len(step_action_next.kv_cache_block_ids) == 0:
|
||||
step_action_next.kv_cache_block_ids = (
|
||||
step_action_prev.kv_cache_block_ids.copy()
|
||||
)
|
||||
|
||||
|
||||
@create_new_process_for_each_test()
|
||||
def test_mamba_prefix_cache_mrv1(monkeypatch: pytest.MonkeyPatch):
|
||||
run_ref_mamba_state_in_subprocess()
|
||||
apply_patch(monkeypatch)
|
||||
prompt_dataset = datasets.load_dataset("heheda/a_long_article")
|
||||
full_prompt = prompt_dataset["train"][0]["text"]
|
||||
tests = get_mamba_prefix_cache_step_configs()
|
||||
|
||||
engine = LLM(
|
||||
model=MODEL,
|
||||
enable_prefix_caching=True,
|
||||
@@ -758,16 +775,7 @@ def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
|
||||
)
|
||||
global cur_step_action_idx
|
||||
cur_step_action_idx = 0
|
||||
for step_action_prev, step_action_next in zip(
|
||||
test_config.step_actions[:-1], test_config.step_actions[1:]
|
||||
):
|
||||
if (
|
||||
step_action_next.kv_cache_block_ids is not None
|
||||
and len(step_action_next.kv_cache_block_ids) == 0
|
||||
):
|
||||
prev_block_ids = step_action_prev.kv_cache_block_ids
|
||||
if prev_block_ids is not None:
|
||||
step_action_next.kv_cache_block_ids = prev_block_ids.copy()
|
||||
fill_following_kv_cache_block_ids(test_config)
|
||||
global step_actions
|
||||
step_actions = test_config.step_actions
|
||||
_ = engine.generate(
|
||||
@@ -787,3 +795,259 @@ def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
|
||||
del engine
|
||||
torch.accelerator.empty_cache()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
@create_new_process_for_each_test()
|
||||
def test_mamba_prefix_cache_mrv2(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
|
||||
monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "1")
|
||||
envs.disable_envs_cache()
|
||||
|
||||
from vllm.v1.worker.gpu.model_runner import GPUModelRunner as MRV2GPUModelRunner
|
||||
from vllm.v1.worker.gpu.model_states.mamba_hybrid import (
|
||||
MambaHybridModelState,
|
||||
)
|
||||
from vllm.v1.worker.gpu.sample.output import SamplerOutput as MRV2SamplerOutput
|
||||
|
||||
events: list[int] = []
|
||||
original_execute_model = MRV2GPUModelRunner.execute_model
|
||||
original_sample = MRV2GPUModelRunner.sample
|
||||
original_preprocess_state = MambaHybridModelState.preprocess_state
|
||||
original_postprocess_state = MambaHybridModelState.postprocess_state
|
||||
original_step_action_fn = InprocClient.get_output
|
||||
original_allocate_slots = KVCacheManager.allocate_slots
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def temporal_states(model_state, block_tables, kv_cache_config):
|
||||
# Qwen3-Next keeps the temporal (ssm) state as the last Mamba cache.
|
||||
forward_context = (
|
||||
model_state.vllm_config.compilation_config.static_forward_context
|
||||
)
|
||||
group_ids, _ = get_mamba_groups(kv_cache_config)
|
||||
for group_id in group_ids:
|
||||
block_table = block_tables[group_id]
|
||||
for layer_name in kv_cache_config.kv_cache_groups[group_id].layer_names:
|
||||
yield forward_context[layer_name].kv_cache[-1], block_table
|
||||
|
||||
def temporal_block(temporal_state, block_table, col):
|
||||
return temporal_state[int(block_table[0, col].item())]
|
||||
|
||||
def wrapped_preprocess_state(
|
||||
self: MambaHybridModelState,
|
||||
input_batch: Any,
|
||||
block_tables: tuple[torch.Tensor, ...],
|
||||
kv_cache_config: KVCacheConfig,
|
||||
num_computed_tokens: torch.Tensor,
|
||||
) -> None:
|
||||
captured["block_tables"] = block_tables
|
||||
captured["kv_cache_config"] = kv_cache_config
|
||||
expected = (
|
||||
None if cur_step_action is None else cur_step_action.preprocess_copy_idx
|
||||
)
|
||||
snapshots = []
|
||||
if expected is not None and expected != (-1, -1):
|
||||
for temporal, bt in temporal_states(self, block_tables, kv_cache_config):
|
||||
snapshots.append(
|
||||
(temporal, bt, temporal_block(temporal, bt, expected[0]).clone())
|
||||
)
|
||||
ret = original_preprocess_state(
|
||||
self, input_batch, block_tables, kv_cache_config, num_computed_tokens
|
||||
)
|
||||
if cur_step_action is not None:
|
||||
req_idx = int(input_batch.idx_mapping[0].item())
|
||||
src_col = int(self._mamba_src_col_gpu[req_idx].item())
|
||||
off = int(self._mamba_src_off_gpu[req_idx].item())
|
||||
dst = int(self._mamba_state_idx_gpu[req_idx].item())
|
||||
actual = (-1, -1) if src_col < 0 or src_col == dst else (src_col + off, dst)
|
||||
assert actual == expected, (
|
||||
f"V2 align preprocess copy: expected={expected}, "
|
||||
f"actual={actual}, {cur_step_action=}"
|
||||
)
|
||||
for temporal, bt, src_state in snapshots:
|
||||
torch.testing.assert_close(
|
||||
temporal_block(temporal, bt, expected[1]), src_state
|
||||
)
|
||||
return ret
|
||||
|
||||
def wrapped_postprocess_state(
|
||||
self: MambaHybridModelState,
|
||||
idx_mapping: torch.Tensor,
|
||||
num_sampled: torch.Tensor | int,
|
||||
num_computed_tokens: torch.Tensor | None = None,
|
||||
) -> None:
|
||||
action = cur_step_action
|
||||
block_tables = captured.get("block_tables")
|
||||
kv_cache_config = captured.get("kv_cache_config")
|
||||
# The postprocess kernel does not expose its indices, so only the copy
|
||||
# case is checked, by effect: snapshot the src block, expect dst == src.
|
||||
if (
|
||||
action is None
|
||||
or num_computed_tokens is None
|
||||
or block_tables is None
|
||||
or action.postprocess_copy_idx == (-1, -1)
|
||||
):
|
||||
return original_postprocess_state(
|
||||
self, idx_mapping, num_sampled, num_computed_tokens
|
||||
)
|
||||
expected = action.postprocess_copy_idx
|
||||
snapshots = [
|
||||
(temporal, bt, temporal_block(temporal, bt, expected[0]).clone())
|
||||
for temporal, bt in temporal_states(self, block_tables, kv_cache_config)
|
||||
]
|
||||
ret = original_postprocess_state(
|
||||
self, idx_mapping, num_sampled, num_computed_tokens
|
||||
)
|
||||
for temporal, bt, src_state in snapshots:
|
||||
torch.testing.assert_close(
|
||||
temporal_block(temporal, bt, expected[1]), src_state
|
||||
)
|
||||
return ret
|
||||
|
||||
def wrapped_execute_model(
|
||||
self: MRV2GPUModelRunner,
|
||||
scheduler_output: SchedulerOutput,
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
):
|
||||
events.extend(
|
||||
req.num_computed_tokens for req in scheduler_output.scheduled_new_reqs
|
||||
)
|
||||
events.extend(scheduler_output.scheduled_cached_reqs.num_computed_tokens)
|
||||
if cur_step_action is not None:
|
||||
num_scheduled_tokens = next(
|
||||
iter(scheduler_output.num_scheduled_tokens.values())
|
||||
)
|
||||
assert num_scheduled_tokens == cur_step_action.num_scheduled_tokens
|
||||
ret = original_execute_model(self, scheduler_output, *args, **kwargs)
|
||||
if cur_step_action is not None and self.execute_model_state is not None:
|
||||
input_batch = self.execute_model_state.input_batch
|
||||
assert (
|
||||
cur_step_action.num_computed_tokens_start
|
||||
== input_batch.positions[input_batch.query_start_loc[0]].item()
|
||||
)
|
||||
return ret
|
||||
|
||||
def fake_sample(
|
||||
self: MRV2GPUModelRunner,
|
||||
hidden_states: torch.Tensor,
|
||||
input_batch: Any,
|
||||
grammar_output: Any,
|
||||
):
|
||||
if cur_step_action is None:
|
||||
return original_sample(self, hidden_states, input_batch, grammar_output)
|
||||
|
||||
num_reqs = input_batch.num_reqs
|
||||
sampled_token_ids = torch.ones(
|
||||
(num_reqs, self.num_speculative_steps + 1),
|
||||
device=hidden_states.device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
num_logits = torch.tensor(
|
||||
input_batch.cu_num_logits_np[1 : num_reqs + 1]
|
||||
- input_batch.cu_num_logits_np[:num_reqs],
|
||||
device=hidden_states.device,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
accepted = torch.full_like(num_logits, num_accepted_tokens)
|
||||
num_sampled = torch.minimum(accepted, num_logits)
|
||||
prefill_lens = self.req_states.prefill_len.gpu[input_batch.idx_mapping]
|
||||
is_chunked_prefill = input_batch.seq_lens[:num_reqs] < prefill_lens
|
||||
num_sampled = torch.where(is_chunked_prefill, 0, num_sampled)
|
||||
num_rejected = torch.where(is_chunked_prefill, 0, num_logits - num_sampled)
|
||||
sampler_output = MRV2SamplerOutput(
|
||||
sampled_token_ids=sampled_token_ids,
|
||||
logprobs_tensors=None,
|
||||
num_nans=None,
|
||||
num_sampled=num_sampled,
|
||||
)
|
||||
return sampler_output, num_sampled, num_rejected
|
||||
|
||||
monkeypatch.setattr(
|
||||
InprocClient,
|
||||
"get_output",
|
||||
get_fake_step_action_fn(original_step_action_fn),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
KVCacheManager,
|
||||
"allocate_slots",
|
||||
get_fake_allocate_slots_fn(original_allocate_slots),
|
||||
)
|
||||
monkeypatch.setattr(MRV2GPUModelRunner, "execute_model", wrapped_execute_model)
|
||||
monkeypatch.setattr(MRV2GPUModelRunner, "sample", fake_sample)
|
||||
monkeypatch.setattr(
|
||||
MambaHybridModelState, "preprocess_state", wrapped_preprocess_state
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
MambaHybridModelState, "postprocess_state", wrapped_postprocess_state
|
||||
)
|
||||
|
||||
engine = LLM(
|
||||
model=MODEL,
|
||||
load_format="dummy",
|
||||
enforce_eager=True,
|
||||
skip_tokenizer_init=True,
|
||||
enable_prefix_caching=True,
|
||||
block_size=BLOCK_SIZE,
|
||||
mamba_cache_mode="align",
|
||||
speculative_config={
|
||||
"method": "qwen3_next_mtp",
|
||||
"num_speculative_tokens": num_speculative_tokens,
|
||||
},
|
||||
max_num_batched_tokens=3072,
|
||||
max_model_len=BLOCK_SIZE * 12,
|
||||
hf_overrides={"num_hidden_layers": NUM_HIDDEN_LAYERS},
|
||||
seed=42,
|
||||
)
|
||||
|
||||
try:
|
||||
tests = get_mamba_prefix_cache_step_configs()
|
||||
|
||||
global step_actions
|
||||
global cur_step_action_idx
|
||||
global num_accepted_tokens
|
||||
for test_name, test_config in tests.items():
|
||||
num_accepted_tokens = test_config.num_accepted_tokens
|
||||
cur_step_action_idx = 0
|
||||
fill_following_kv_cache_block_ids(test_config)
|
||||
step_actions = test_config.step_actions
|
||||
sampling_params = SamplingParams(
|
||||
temperature=0.0,
|
||||
max_tokens=test_config.num_generated_tokens,
|
||||
ignore_eos=True,
|
||||
)
|
||||
_ = engine.generate(
|
||||
[TokensPrompt(prompt_token_ids=[1] * test_config.num_prompt_tokens)],
|
||||
sampling_params=sampling_params,
|
||||
)
|
||||
assert cur_step_action_idx == len(test_config.step_actions), test_name
|
||||
assert (
|
||||
engine.llm_engine.engine_core.engine_core.scheduler.reset_prefix_cache()
|
||||
)
|
||||
|
||||
step_actions = []
|
||||
cur_step_action_idx = 0
|
||||
num_accepted_tokens = 1
|
||||
prompt = TokensPrompt(prompt_token_ids=[1] * (BLOCK_SIZE * 2))
|
||||
sampling_params = SamplingParams(
|
||||
temperature=0.0,
|
||||
max_tokens=1,
|
||||
ignore_eos=True,
|
||||
)
|
||||
_ = engine.generate([prompt], sampling_params=sampling_params)
|
||||
first_event_count = len(events)
|
||||
_ = engine.generate([prompt], sampling_params=sampling_params)
|
||||
second_events = events[first_event_count:]
|
||||
prefix_hits = [
|
||||
num_computed_tokens
|
||||
for num_computed_tokens in second_events
|
||||
if num_computed_tokens >= BLOCK_SIZE
|
||||
]
|
||||
assert prefix_hits, (
|
||||
"Expected the second identical prompt to hit prefix cache, "
|
||||
f"got events={second_events!r}"
|
||||
)
|
||||
assert engine.llm_engine.engine_core.engine_core.scheduler.reset_prefix_cache()
|
||||
finally:
|
||||
del engine
|
||||
torch.accelerator.empty_cache()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -33,14 +33,37 @@ from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import
|
||||
)
|
||||
from vllm.utils.network_utils import get_open_port
|
||||
from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig
|
||||
from vllm.v1.kv_cache_interface import (
|
||||
FullAttentionSpec,
|
||||
KVCacheConfig,
|
||||
KVCacheGroupSpec,
|
||||
)
|
||||
from vllm.v1.request import RequestStatus
|
||||
|
||||
from .utils import create_request, create_scheduler, create_vllm_config
|
||||
|
||||
|
||||
def _make_test_kv_cache_config() -> KVCacheConfig:
|
||||
return KVCacheConfig(num_blocks=0, kv_cache_tensors=[], kv_cache_groups=[])
|
||||
return KVCacheConfig(
|
||||
num_blocks=0,
|
||||
kv_cache_tensors=[],
|
||||
kv_cache_groups=[
|
||||
KVCacheGroupSpec(
|
||||
[
|
||||
"model.layers.0.self_attn",
|
||||
"model.layers.1.self_attn",
|
||||
"model.layers.0.mla_attn",
|
||||
"model.layers.1.eagle_attn",
|
||||
],
|
||||
FullAttentionSpec(
|
||||
block_size=16,
|
||||
num_kv_heads=4,
|
||||
head_size=64,
|
||||
dtype=torch.float16,
|
||||
),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class FakeMooncakeWrapper:
|
||||
@@ -126,6 +149,8 @@ async def test_build_transfer_params_separates_prefill_pp_layers():
|
||||
worker.is_kv_producer = True
|
||||
worker.tp_rank = 0
|
||||
worker.tp_size = 1
|
||||
worker.kv_cache_config = _make_test_kv_cache_config()
|
||||
worker._physical_blocks_per_logical_kv_block = 1
|
||||
worker.transfer_topo = SimpleNamespace(local_replicates_kv_cache=False)
|
||||
|
||||
block_len = 256
|
||||
@@ -206,6 +231,7 @@ async def test_build_transfer_params_separates_prefill_pp_layers():
|
||||
req_blocks={"d-req-pp": (transfer_id, [[20, 21]])},
|
||||
kv_caches_base_addr=[region.base_addr for region in remote_regions],
|
||||
block_lens=[region.block_len for region in remote_regions],
|
||||
kv_block_lens=[region.kv_block_len for region in remote_regions],
|
||||
registered_layer_names=[region.layer_name for region in remote_regions],
|
||||
registered_layer_indices=[region.layer_index for region in remote_regions],
|
||||
)
|
||||
@@ -266,6 +292,7 @@ async def test_send_kv_to_decode_aligns_consumer_regions_by_layer_metadata(
|
||||
kv_half = block_len // 2
|
||||
prefill_worker.kv_caches_base_addr = [0x1000]
|
||||
prefill_worker.block_len_per_layer = [block_len]
|
||||
prefill_worker.kv_block_len_per_layer = [kv_half]
|
||||
prefill_worker.registered_layer_names = ["model.layers.1.self_attn"]
|
||||
prefill_worker.registered_layer_indices = [1]
|
||||
|
||||
@@ -294,6 +321,7 @@ async def test_send_kv_to_decode_aligns_consumer_regions_by_layer_metadata(
|
||||
req_blocks={"d-req-layer-align": (transfer_id, [[20]])},
|
||||
kv_caches_base_addr=[0xA000, 0xB000],
|
||||
block_lens=[block_len, block_len],
|
||||
kv_block_lens=[kv_half, kv_half],
|
||||
registered_layer_names=[
|
||||
"model.layers.0.self_attn",
|
||||
"model.layers.1.self_attn",
|
||||
@@ -804,7 +832,9 @@ async def test_kv_producer(monkeypatch):
|
||||
prefill_worker = prefill_connector.connector_worker
|
||||
prefill_worker.kv_caches_base_addr = [0x1000]
|
||||
block_len = 4096
|
||||
kv_half = block_len // 2
|
||||
prefill_worker.block_len_per_layer = [block_len]
|
||||
prefill_worker.kv_block_len_per_layer = [kv_half]
|
||||
prefill_worker.registered_layer_names = ["model.layers.0.self_attn"]
|
||||
prefill_worker.registered_layer_indices = [0]
|
||||
|
||||
@@ -832,6 +862,7 @@ async def test_kv_producer(monkeypatch):
|
||||
req_blocks={"d-req-1": (transfer_id, [[20, 21]])},
|
||||
kv_caches_base_addr=[0x2000],
|
||||
block_lens=[block_len],
|
||||
kv_block_lens=[kv_half],
|
||||
registered_layer_names=["model.layers.0.self_attn"],
|
||||
registered_layer_indices=[0],
|
||||
)
|
||||
@@ -845,8 +876,6 @@ async def test_kv_producer(monkeypatch):
|
||||
) as mock_send_blocks:
|
||||
# With blocks-first layout, each block is virtually split
|
||||
# into K and V halves, producing non-coalesced transfers.
|
||||
kv_half = block_len // 2
|
||||
|
||||
def expected_split_transfers(src_base, dst_base, src_blocks, dst_blocks):
|
||||
"""Build expected (src_ptrs, dst_ptrs, lengths) for
|
||||
virtual-split K/V transfers."""
|
||||
@@ -981,6 +1010,7 @@ async def test_kv_consumuer(monkeypatch):
|
||||
decode_worker = decode_connector.connector_worker
|
||||
decode_worker.kv_caches_base_addr = [0x1000]
|
||||
decode_worker.block_len_per_layer = [4096]
|
||||
decode_worker.kv_block_len_per_layer = [4096]
|
||||
decode_worker.registered_layer_names = ["model.layers.0.self_attn"]
|
||||
decode_worker.registered_layer_indices = [0]
|
||||
decode_worker.rpc_port = 54321
|
||||
@@ -1236,6 +1266,7 @@ async def test_kv_producer_heterogeneous_tp(monkeypatch, d_tp_size):
|
||||
|
||||
prefill_worker.kv_caches_base_addr = [0x1000]
|
||||
prefill_worker.block_len_per_layer = [local_block_len]
|
||||
prefill_worker.kv_block_len_per_layer = [local_block_len // 2]
|
||||
prefill_worker.registered_layer_names = ["model.layers.0.self_attn"]
|
||||
prefill_worker.registered_layer_indices = [0]
|
||||
|
||||
@@ -1283,6 +1314,7 @@ async def test_kv_producer_heterogeneous_tp(monkeypatch, d_tp_size):
|
||||
},
|
||||
kv_caches_base_addr=[0x2000],
|
||||
block_lens=[remote_block_len],
|
||||
kv_block_lens=[remote_block_len // 2],
|
||||
registered_layer_names=["model.layers.0.self_attn"],
|
||||
registered_layer_indices=[0],
|
||||
)
|
||||
|
||||
@@ -257,6 +257,7 @@ async def test_build_transfer_params_multi_group_trimming(monkeypatch):
|
||||
},
|
||||
kv_caches_base_addr=[0x2000],
|
||||
block_lens=[block_len],
|
||||
kv_block_lens=[block_len],
|
||||
)
|
||||
|
||||
local_regions = [
|
||||
@@ -348,6 +349,7 @@ async def test_build_transfer_params_group_count_mismatch(monkeypatch):
|
||||
},
|
||||
kv_caches_base_addr=[0x2000],
|
||||
block_lens=[block_len],
|
||||
kv_block_lens=[block_len],
|
||||
)
|
||||
|
||||
local_regions = [
|
||||
|
||||
@@ -0,0 +1,386 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Unit tests for MooncakeConnector hybrid FA + GDN support.
|
||||
|
||||
GDN is represented as a MambaSpec in vLLM, so these tests exercise the
|
||||
Mooncake MambaSpec path with mamba_type=GDN_ATTN. Mamba2 is intentionally not
|
||||
validated by this test module.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm.config import set_current_vllm_config
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_connector import (
|
||||
KVConnectorRole,
|
||||
MooncakeConnector,
|
||||
MooncakeConnectorScheduler,
|
||||
MooncakeConnectorWorker,
|
||||
MooncakeXferMetadata,
|
||||
SendBlockMeta,
|
||||
TransferRegion,
|
||||
)
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
|
||||
from vllm.v1.kv_cache_interface import (
|
||||
FullAttentionSpec,
|
||||
KVCacheConfig,
|
||||
KVCacheGroupSpec,
|
||||
MambaSpec,
|
||||
)
|
||||
|
||||
from .test_mooncake_connector import patch_worker_dependencies
|
||||
from .utils import create_request, create_vllm_config
|
||||
|
||||
|
||||
def noop_shutdown():
|
||||
pass
|
||||
|
||||
|
||||
def make_hybrid_gdn_kv_cache_config(block_size: int) -> KVCacheConfig:
|
||||
gdn_spec = MambaSpec(
|
||||
block_size=block_size,
|
||||
shapes=((6, 3), (1, 2, 2)),
|
||||
dtypes=(torch.float16, torch.float16),
|
||||
mamba_type=MambaAttentionBackendEnum.GDN_ATTN,
|
||||
)
|
||||
assert gdn_spec.mamba_type == MambaAttentionBackendEnum.GDN_ATTN
|
||||
return KVCacheConfig(
|
||||
num_blocks=16,
|
||||
kv_cache_tensors=[],
|
||||
kv_cache_groups=[
|
||||
KVCacheGroupSpec(
|
||||
["model.layers.0.self_attn"],
|
||||
FullAttentionSpec(
|
||||
block_size=block_size,
|
||||
num_kv_heads=1,
|
||||
head_size=1,
|
||||
dtype=torch.float16,
|
||||
),
|
||||
),
|
||||
KVCacheGroupSpec(
|
||||
["model.layers.1.linear_attn"],
|
||||
gdn_spec,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def make_hybrid_gdn_scheduler(kv_role: str) -> MooncakeConnectorScheduler:
|
||||
vllm_config = create_vllm_config(
|
||||
kv_connector="MooncakeConnector",
|
||||
kv_role=kv_role,
|
||||
)
|
||||
vllm_config.scheduler_config.disable_hybrid_kv_cache_manager = False
|
||||
return MooncakeConnectorScheduler(
|
||||
vllm_config=vllm_config,
|
||||
engine_id="test-engine",
|
||||
kv_cache_config=make_hybrid_gdn_kv_cache_config(
|
||||
vllm_config.cache_config.block_size
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.cpu_test
|
||||
def test_hybrid_gdn_remote_prefill_uses_mamba_n_minus_one():
|
||||
scheduler = make_hybrid_gdn_scheduler(kv_role="kv_consumer")
|
||||
request = create_request(num_tokens=10, do_remote_prefill=True)
|
||||
|
||||
num_new_tokens, is_async = scheduler.get_num_new_matched_tokens(
|
||||
request, num_computed_tokens=0
|
||||
)
|
||||
|
||||
assert num_new_tokens == request.num_prompt_tokens - 1
|
||||
assert is_async is True
|
||||
|
||||
|
||||
@pytest.mark.cpu_test
|
||||
def test_hybrid_gdn_remote_decode_truncates_prefill_once():
|
||||
scheduler = make_hybrid_gdn_scheduler(kv_role="kv_producer")
|
||||
request = create_request(num_tokens=10, do_remote_decode=True)
|
||||
original_tokens = list(request.prompt_token_ids)
|
||||
|
||||
num_new_tokens, is_async = scheduler.get_num_new_matched_tokens(
|
||||
request, num_computed_tokens=0
|
||||
)
|
||||
|
||||
assert num_new_tokens == 0
|
||||
assert is_async is False
|
||||
assert request.prompt_token_ids == original_tokens[:-1]
|
||||
assert request._all_token_ids == original_tokens[:-1]
|
||||
assert request.num_prompt_tokens == len(original_tokens) - 1
|
||||
assert request.max_tokens == 1
|
||||
assert request.kv_transfer_params["_p_side_truncated"] is True
|
||||
|
||||
scheduler.get_num_new_matched_tokens(request, num_computed_tokens=0)
|
||||
assert request.prompt_token_ids == original_tokens[:-1]
|
||||
|
||||
|
||||
def test_register_kv_caches_emits_fa_and_gdn_regions(monkeypatch):
|
||||
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
|
||||
vllm_config = create_vllm_config(
|
||||
kv_connector="MooncakeConnector",
|
||||
kv_role="kv_consumer",
|
||||
)
|
||||
kv_cache_config = make_hybrid_gdn_kv_cache_config(
|
||||
vllm_config.cache_config.block_size
|
||||
)
|
||||
|
||||
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
|
||||
connector = MooncakeConnector(
|
||||
vllm_config,
|
||||
KVConnectorRole.WORKER,
|
||||
kv_cache_config,
|
||||
)
|
||||
worker = connector.connector_worker
|
||||
|
||||
fa_cache = torch.empty((2, 2, 11), dtype=torch.float16)
|
||||
gdn_conv_state = torch.empty((2, 22), dtype=torch.float16)
|
||||
gdn_ssm_state = torch.empty((2, 4), dtype=torch.float16)
|
||||
|
||||
worker.register_kv_caches(
|
||||
{
|
||||
"model.layers.0.self_attn": fa_cache,
|
||||
"model.layers.1.linear_attn": (gdn_conv_state, gdn_ssm_state),
|
||||
}
|
||||
)
|
||||
|
||||
assert worker.transfer_topo.is_mamba is True
|
||||
assert worker.registered_layer_names == [
|
||||
"model.layers.0.self_attn",
|
||||
"model.layers.1.linear_attn",
|
||||
]
|
||||
assert worker.registered_group_indices == [0, 1]
|
||||
assert worker.kv_caches_base_addr == [
|
||||
fa_cache.data_ptr(),
|
||||
gdn_conv_state.data_ptr(),
|
||||
]
|
||||
|
||||
worker.shutdown()
|
||||
worker.shutdown = noop_shutdown
|
||||
connector.connector_worker = None
|
||||
|
||||
|
||||
def test_register_kv_caches_deduplicates_shared_backing_memory(monkeypatch):
|
||||
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
|
||||
vllm_config = create_vllm_config(
|
||||
kv_connector="MooncakeConnector",
|
||||
kv_role="kv_consumer",
|
||||
)
|
||||
kv_cache_config = make_hybrid_gdn_kv_cache_config(
|
||||
vllm_config.cache_config.block_size
|
||||
)
|
||||
|
||||
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
|
||||
connector = MooncakeConnector(
|
||||
vllm_config,
|
||||
KVConnectorRole.WORKER,
|
||||
kv_cache_config,
|
||||
)
|
||||
worker = connector.connector_worker
|
||||
|
||||
backing = torch.empty((4, 64), dtype=torch.float16)
|
||||
fa_cache = backing[:2, :16]
|
||||
gdn_conv_state = backing[:3]
|
||||
gdn_ssm_state = torch.empty((3, 4), dtype=torch.float16)
|
||||
|
||||
with patch.object(
|
||||
worker.engine, "batch_register_memory", return_value=0
|
||||
) as batch_register_memory:
|
||||
worker.register_kv_caches(
|
||||
{
|
||||
"model.layers.0.self_attn": fa_cache,
|
||||
"model.layers.1.linear_attn": (gdn_conv_state, gdn_ssm_state),
|
||||
}
|
||||
)
|
||||
|
||||
assert worker.kv_caches_base_addr == [
|
||||
fa_cache.data_ptr(),
|
||||
gdn_conv_state.data_ptr(),
|
||||
]
|
||||
batch_register_memory.assert_called_once()
|
||||
registered_ptrs, registered_lens = batch_register_memory.call_args[0]
|
||||
assert registered_ptrs == [backing.data_ptr()]
|
||||
assert registered_lens == [backing.untyped_storage().nbytes()]
|
||||
|
||||
worker.shutdown()
|
||||
worker.shutdown = noop_shutdown
|
||||
connector.connector_worker = None
|
||||
|
||||
|
||||
def test_hybrid_gdn_transfer_params_preserve_group_identity(monkeypatch):
|
||||
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
|
||||
vllm_config = create_vllm_config(
|
||||
kv_connector="MooncakeConnector",
|
||||
kv_role="kv_producer",
|
||||
)
|
||||
kv_cache_config = make_hybrid_gdn_kv_cache_config(
|
||||
vllm_config.cache_config.block_size
|
||||
)
|
||||
|
||||
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
|
||||
connector = MooncakeConnector(
|
||||
vllm_config,
|
||||
KVConnectorRole.WORKER,
|
||||
kv_cache_config,
|
||||
)
|
||||
worker = connector.connector_worker
|
||||
|
||||
block_len = 0x100
|
||||
transfer_id = "xfer-hybrid-gdn"
|
||||
|
||||
async def build_transfer_params():
|
||||
send_meta = SendBlockMeta(
|
||||
p_req_id="p-hybrid-gdn",
|
||||
transfer_id=transfer_id,
|
||||
local_block_ids=[
|
||||
[10, 11],
|
||||
[NULL_BLOCK_ID, 4],
|
||||
],
|
||||
ready=asyncio.Event(),
|
||||
)
|
||||
return await worker._build_transfer_params(
|
||||
[("d-hybrid-gdn", send_meta)],
|
||||
xfer_meta,
|
||||
local_regions,
|
||||
remote_regions,
|
||||
)
|
||||
|
||||
xfer_meta = MooncakeXferMetadata(
|
||||
remote_hostname="consumer-host",
|
||||
remote_port=54321,
|
||||
remote_tp_size=1,
|
||||
remote_tp_rank=0,
|
||||
req_blocks={
|
||||
"d-hybrid-gdn": (
|
||||
transfer_id,
|
||||
[
|
||||
[30, 31],
|
||||
[NULL_BLOCK_ID, 7],
|
||||
],
|
||||
)
|
||||
},
|
||||
kv_caches_base_addr=[],
|
||||
block_lens=[],
|
||||
kv_block_lens=[],
|
||||
)
|
||||
|
||||
local_regions = [
|
||||
TransferRegion(
|
||||
layer_name="model.layers.1.linear_attn",
|
||||
layer_index=1,
|
||||
base_addr=0x5000,
|
||||
block_len=block_len,
|
||||
kv_block_len=block_len,
|
||||
group_index=1,
|
||||
),
|
||||
TransferRegion(
|
||||
layer_name="model.layers.0.self_attn",
|
||||
layer_index=0,
|
||||
base_addr=0x1000,
|
||||
block_len=block_len,
|
||||
kv_block_len=block_len,
|
||||
group_index=0,
|
||||
),
|
||||
]
|
||||
remote_regions = [
|
||||
TransferRegion(
|
||||
layer_name="model.layers.1.linear_attn",
|
||||
layer_index=1,
|
||||
base_addr=0x6000,
|
||||
block_len=block_len,
|
||||
kv_block_len=block_len,
|
||||
group_index=1,
|
||||
),
|
||||
TransferRegion(
|
||||
layer_name="model.layers.0.self_attn",
|
||||
layer_index=0,
|
||||
base_addr=0x2000,
|
||||
block_len=block_len,
|
||||
kv_block_len=block_len,
|
||||
group_index=0,
|
||||
),
|
||||
]
|
||||
|
||||
src_ptrs, dst_ptrs, lengths, err_reqs, err_msg = asyncio.run(
|
||||
build_transfer_params()
|
||||
)
|
||||
|
||||
assert err_reqs == []
|
||||
assert err_msg is None
|
||||
assert src_ptrs == [
|
||||
0x5000 + 4 * block_len,
|
||||
0x1000 + 10 * block_len,
|
||||
]
|
||||
assert dst_ptrs == [
|
||||
0x6000 + 7 * block_len,
|
||||
0x2000 + 30 * block_len,
|
||||
]
|
||||
assert lengths == [block_len, 2 * block_len]
|
||||
|
||||
worker.shutdown()
|
||||
worker.shutdown = noop_shutdown
|
||||
connector.connector_worker = None
|
||||
|
||||
|
||||
def test_logical_to_kernel_block_ids_expands_fa_not_gdn():
|
||||
worker = object.__new__(MooncakeConnectorWorker)
|
||||
worker.shutdown = noop_shutdown
|
||||
worker._physical_blocks_per_logical_kv_block = 17
|
||||
worker.kv_cache_config = make_hybrid_gdn_kv_cache_config(block_size=544)
|
||||
|
||||
block_ids = [[2], [2]]
|
||||
kernel_block_ids = worker._logical_to_kernel_block_ids(block_ids)
|
||||
|
||||
assert kernel_block_ids == [list(range(34, 51)), [2]]
|
||||
|
||||
|
||||
def test_hybrid_gdn_splits_fa_regions_but_keeps_gdn_state_whole(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
|
||||
vllm_config = create_vllm_config(
|
||||
kv_connector="MooncakeConnector",
|
||||
kv_role="kv_producer",
|
||||
)
|
||||
kv_cache_config = make_hybrid_gdn_kv_cache_config(
|
||||
vllm_config.cache_config.block_size
|
||||
)
|
||||
|
||||
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
|
||||
connector = MooncakeConnector(
|
||||
vllm_config,
|
||||
KVConnectorRole.WORKER,
|
||||
kv_cache_config,
|
||||
)
|
||||
worker = connector.connector_worker
|
||||
|
||||
worker.transfer_topo = SimpleNamespace(virtually_split_kv_in_blocks=True)
|
||||
regions = worker._get_transfer_regions(
|
||||
base_addrs=[0x1000, 0x2000],
|
||||
block_lens=[0x100, 0x100],
|
||||
kv_block_lens=[0x40, 0x100],
|
||||
layer_names=[
|
||||
"model.layers.0.self_attn",
|
||||
"model.layers.1.linear_attn",
|
||||
],
|
||||
layer_indices=[0, 1],
|
||||
group_indices=[0, 1],
|
||||
)
|
||||
|
||||
assert [
|
||||
(region.group_index, region.base_addr, region.kv_block_len)
|
||||
for region in regions
|
||||
] == [
|
||||
(0, 0x1000, 0x40),
|
||||
(0, 0x1040, 0x40),
|
||||
(1, 0x2000, 0x100),
|
||||
]
|
||||
|
||||
worker.shutdown()
|
||||
worker.shutdown = noop_shutdown
|
||||
connector.connector_worker = None
|
||||
@@ -29,9 +29,8 @@ def _gpu_snapshot(tag: str, prev_alloc: float = 0.0) -> dict:
|
||||
torch.accelerator.synchronize()
|
||||
alloc = torch.accelerator.memory_allocated()
|
||||
reserved = torch.accelerator.memory_reserved()
|
||||
# mem_get_info is not available on torch.accelerator
|
||||
try:
|
||||
drv_free, drv_total = torch.cuda.mem_get_info()
|
||||
drv_free, drv_total = torch.accelerator.get_memory_info()
|
||||
drv_used = drv_total - drv_free
|
||||
drv_pct = drv_used / drv_total * 100
|
||||
except Exception:
|
||||
|
||||
@@ -0,0 +1,331 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
P2PConnector proxy server for OffloadingConnector + TieringOffloadingSpec.
|
||||
|
||||
Unlike NixlConnector (which returns remote_host/remote_port in the prefill
|
||||
response), OffloadingConnector does not embed connector coordinates in its
|
||||
response. This proxy injects the prefiller's P2PConnector address into
|
||||
kv_transfer_params before forwarding the decode request so the decoder knows
|
||||
where to pull KV blocks from.
|
||||
|
||||
Usage:
|
||||
.venv/bin/python p2p_connector_proxy.py \
|
||||
--port 8192 \
|
||||
--prefiller-host 127.0.0.1 --prefiller-port 8100 \
|
||||
--decoder-host 127.0.0.1 --decoder-port 8200 \
|
||||
--p2p-connector-host 127.0.0.1 --p2p-connector-port 7777
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import itertools
|
||||
import logging
|
||||
import os
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.setLevel(logging.DEBUG)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
app.state.prefill_clients = []
|
||||
app.state.decode_clients = []
|
||||
|
||||
for i, (host, port) in enumerate(global_args.prefiller_instances):
|
||||
app.state.prefill_clients.append(
|
||||
{
|
||||
"client": httpx.AsyncClient(
|
||||
timeout=None,
|
||||
base_url=f"http://{host}:{port}/v1",
|
||||
limits=httpx.Limits(
|
||||
max_connections=None, max_keepalive_connections=None
|
||||
),
|
||||
),
|
||||
"host": host,
|
||||
"port": port,
|
||||
"id": i,
|
||||
}
|
||||
)
|
||||
|
||||
for i, (host, port) in enumerate(global_args.decoder_instances):
|
||||
app.state.decode_clients.append(
|
||||
{
|
||||
"client": httpx.AsyncClient(
|
||||
timeout=None,
|
||||
base_url=f"http://{host}:{port}/v1",
|
||||
limits=httpx.Limits(
|
||||
max_connections=None, max_keepalive_connections=None
|
||||
),
|
||||
),
|
||||
"host": host,
|
||||
"port": port,
|
||||
"id": i,
|
||||
}
|
||||
)
|
||||
|
||||
app.state.prefill_iterator = itertools.cycle(range(len(app.state.prefill_clients)))
|
||||
app.state.decode_iterator = itertools.cycle(range(len(app.state.decode_clients)))
|
||||
|
||||
mode = "decoder-first" if global_args.decoder_first else "prefiller-first"
|
||||
pd_host = global_args.p2p_connector_host
|
||||
pd_port = global_args.p2p_connector_port
|
||||
print(
|
||||
f"Proxy ready [{mode}]: "
|
||||
f"{len(app.state.prefill_clients)} prefiller(s), "
|
||||
f"{len(app.state.decode_clients)} decoder(s). "
|
||||
f"P2PConnector at {pd_host}:{pd_port}"
|
||||
)
|
||||
yield
|
||||
|
||||
for ci in app.state.prefill_clients:
|
||||
await ci["client"].aclose()
|
||||
for ci in app.state.decode_clients:
|
||||
await ci["client"].aclose()
|
||||
|
||||
|
||||
app = FastAPI(lifespan=lifespan)
|
||||
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--port", type=int, default=8192)
|
||||
p.add_argument("--host", type=str, default="127.0.0.1")
|
||||
p.add_argument("--prefiller-hosts", type=str, nargs="+", default=["127.0.0.1"])
|
||||
p.add_argument("--prefiller-ports", type=int, nargs="+", default=[8100])
|
||||
p.add_argument("--decoder-hosts", type=str, nargs="+", default=["127.0.0.1"])
|
||||
p.add_argument("--decoder-ports", type=int, nargs="+", default=[8200])
|
||||
# P2PConnector coordinates of the prefiller — injected into decode requests.
|
||||
p.add_argument(
|
||||
"--p2p-connector-host",
|
||||
type=str,
|
||||
default="127.0.0.1",
|
||||
help="Host of the prefiller's P2PConnector ZMQ socket",
|
||||
)
|
||||
p.add_argument(
|
||||
"--p2p-connector-port",
|
||||
type=int,
|
||||
default=7777,
|
||||
help="Port of the prefiller's P2PConnector ZMQ socket",
|
||||
)
|
||||
# P2PConnector coordinates of the decoder — injected into prefill requests
|
||||
# so the prefiller's submit_store can resolve the peer to push KV to.
|
||||
p.add_argument(
|
||||
"--decoder-p2p-connector-host",
|
||||
type=str,
|
||||
default="127.0.0.1",
|
||||
help="Host of the decoder's P2PConnector ZMQ socket",
|
||||
)
|
||||
p.add_argument(
|
||||
"--decoder-p2p-connector-port",
|
||||
type=int,
|
||||
default=7778,
|
||||
help="Port of the decoder's P2PConnector ZMQ socket",
|
||||
)
|
||||
p.add_argument(
|
||||
"--decoder-first",
|
||||
action="store_true",
|
||||
help="Send decode request before prefill so decoder is already "
|
||||
"waiting when KV blocks arrive (decoder-first mode)",
|
||||
)
|
||||
args = p.parse_args()
|
||||
if len(args.prefiller_hosts) != len(args.prefiller_ports):
|
||||
raise ValueError("Prefiller host/port count mismatch")
|
||||
if len(args.decoder_hosts) != len(args.decoder_ports):
|
||||
raise ValueError("Decoder host/port count mismatch")
|
||||
args.prefiller_instances = list(zip(args.prefiller_hosts, args.prefiller_ports))
|
||||
args.decoder_instances = list(zip(args.decoder_hosts, args.decoder_ports))
|
||||
return args
|
||||
|
||||
|
||||
def _get_next(app, service: str):
|
||||
if service == "prefill":
|
||||
return app.state.prefill_clients[next(app.state.prefill_iterator)]
|
||||
return app.state.decode_clients[next(app.state.decode_iterator)]
|
||||
|
||||
|
||||
def _auth_headers(request_id: str) -> dict:
|
||||
headers: dict = {"X-Request-Id": request_id}
|
||||
api_key = os.environ.get("OPENAI_API_KEY", "")
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return headers
|
||||
|
||||
|
||||
async def _prefill(client_info, endpoint, req_data, request_id):
|
||||
"""Send a prefill-only request (max_tokens=1) to the prefiller."""
|
||||
data = req_data.copy()
|
||||
data["kv_transfer_params"] = {
|
||||
"decode": {
|
||||
"kv_request_id": request_id,
|
||||
},
|
||||
}
|
||||
data["stream"] = False
|
||||
data["max_tokens"] = 1
|
||||
data.pop("max_completion_tokens", None)
|
||||
data.pop("stream_options", None)
|
||||
data.pop("min_tokens", None)
|
||||
data.pop("min_completion_tokens", None)
|
||||
|
||||
headers = _auth_headers(request_id)
|
||||
resp = await client_info["client"].post(endpoint, json=data, headers=headers)
|
||||
resp.raise_for_status()
|
||||
await resp.aread()
|
||||
return resp
|
||||
|
||||
|
||||
async def _stream_decode(client_info, endpoint, req_data, request_id):
|
||||
headers = _auth_headers(request_id)
|
||||
async with client_info["client"].stream(
|
||||
"POST", endpoint, json=req_data, headers=headers
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
async for chunk in resp.aiter_bytes():
|
||||
yield chunk
|
||||
|
||||
|
||||
async def _handle_completions(api: str, request: Request):
|
||||
try:
|
||||
req_data = await request.json()
|
||||
request_id = str(uuid.uuid4())
|
||||
|
||||
prefill_client = _get_next(request.app, "prefill")
|
||||
await _prefill(prefill_client, api, req_data, request_id)
|
||||
|
||||
# Inject the prefiller's P2PConnector address so the decoder can pull
|
||||
# KV blocks from it via the P2PConnector transport.
|
||||
req_data["kv_transfer_params"] = {
|
||||
"prefill": {
|
||||
"kv_request_id": request_id,
|
||||
"remote_host": global_args.p2p_connector_host,
|
||||
"remote_port": global_args.p2p_connector_port,
|
||||
},
|
||||
}
|
||||
|
||||
decode_client = _get_next(request.app, "decode")
|
||||
logger.debug("prefill=%s decode=%s", prefill_client, decode_client)
|
||||
|
||||
async def generate():
|
||||
async for chunk in _stream_decode(decode_client, api, req_data, request_id):
|
||||
yield chunk
|
||||
|
||||
return StreamingResponse(generate(), media_type="application/json")
|
||||
|
||||
except Exception as e:
|
||||
import sys
|
||||
import traceback
|
||||
|
||||
print(f"Proxy error on {api}: {e}")
|
||||
print("".join(traceback.format_exception(*sys.exc_info())))
|
||||
raise
|
||||
|
||||
|
||||
async def _handle_completions_decoder_first(api: str, request: Request):
|
||||
"""Decoder-first mode: send decode request before prefill.
|
||||
|
||||
The decoder establishes its request and starts polling for KV blocks
|
||||
immediately. The prefill is then sent so the prefiller computes and
|
||||
pushes blocks to the already-waiting decoder.
|
||||
"""
|
||||
try:
|
||||
req_data = await request.json()
|
||||
request_id = str(uuid.uuid4())
|
||||
|
||||
prefill_client = _get_next(request.app, "prefill")
|
||||
decode_client = _get_next(request.app, "decode")
|
||||
|
||||
decode_data = req_data.copy()
|
||||
decode_data["kv_transfer_params"] = {
|
||||
"prefill": {
|
||||
"kv_request_id": request_id,
|
||||
"remote_host": global_args.p2p_connector_host,
|
||||
"remote_port": global_args.p2p_connector_port,
|
||||
},
|
||||
}
|
||||
|
||||
async def generate():
|
||||
queue: asyncio.Queue = asyncio.Queue()
|
||||
|
||||
async def _run_decode():
|
||||
try:
|
||||
async for chunk in _stream_decode(
|
||||
decode_client, api, decode_data, request_id
|
||||
):
|
||||
await queue.put(("data", chunk))
|
||||
except Exception as exc:
|
||||
await queue.put(("error", exc))
|
||||
finally:
|
||||
await queue.put(("done", None))
|
||||
|
||||
# 1. Start decode request — decoder is now waiting for KV blocks
|
||||
asyncio.create_task(_run_decode())
|
||||
|
||||
# 2. Send prefill — blocks are computed and pushed to the decoder
|
||||
try:
|
||||
await _prefill(prefill_client, api, req_data, request_id)
|
||||
except Exception as exc:
|
||||
logger.warning("decoder-first: prefill failed: %s", exc)
|
||||
|
||||
logger.debug(
|
||||
"decoder-first: prefill done, streaming decode prefill=%s decode=%s",
|
||||
prefill_client,
|
||||
decode_client,
|
||||
)
|
||||
|
||||
# 3. Stream the decode response
|
||||
while True:
|
||||
kind, value = await queue.get()
|
||||
if kind == "done":
|
||||
break
|
||||
if kind == "error":
|
||||
raise value # type: ignore[misc]
|
||||
yield value
|
||||
|
||||
return StreamingResponse(generate(), media_type="application/json")
|
||||
|
||||
except Exception as e:
|
||||
import sys
|
||||
import traceback
|
||||
|
||||
print(f"Proxy error on {api}: {e}")
|
||||
print("".join(traceback.format_exception(*sys.exc_info())))
|
||||
raise
|
||||
|
||||
|
||||
def _route_handler(api: str):
|
||||
if global_args.decoder_first:
|
||||
return lambda req: _handle_completions_decoder_first(api, req)
|
||||
return lambda req: _handle_completions(api, req)
|
||||
|
||||
|
||||
@app.post("/v1/completions")
|
||||
async def completions(request: Request):
|
||||
return await _route_handler("/completions")(request)
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
async def chat_completions(request: Request):
|
||||
return await _route_handler("/chat/completions")(request)
|
||||
|
||||
|
||||
@app.get("/healthcheck")
|
||||
async def healthcheck():
|
||||
return {
|
||||
"status": "ok",
|
||||
"prefill_instances": len(app.state.prefill_clients),
|
||||
"decode_instances": len(app.state.decode_clients),
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
global global_args
|
||||
global_args = parse_args()
|
||||
import uvicorn
|
||||
|
||||
uvicorn.run(app, host=global_args.host, port=global_args.port)
|
||||
+312
@@ -0,0 +1,312 @@
|
||||
#!/bin/bash
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# Accuracy test driver for the p2p connector
|
||||
# (OffloadingConnector + TieringOffloadingSpec + p2p tier).
|
||||
#
|
||||
# Mirrors tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh:
|
||||
# brings up N prefillers + M decoders on the local host, fronts them
|
||||
# with p2p_connector_proxy.py, then runs the connector-agnostic
|
||||
# test_accuracy.py (gsm8k via lm_eval) against the proxy.
|
||||
#
|
||||
# Knobs (env vars unless flagged otherwise):
|
||||
# MODEL_NAMES space-separated model list (default: Llama-3.2-1B-Instruct)
|
||||
# NUM_PREFILL_INSTANCES default 1
|
||||
# NUM_DECODE_INSTANCES default 1
|
||||
# PREFILLER_TP_SIZE default 1
|
||||
# DECODER_TP_SIZE default 1
|
||||
# GPU_MEMORY_UTILIZATION default 0.45
|
||||
# MAX_MODEL_LEN default 512
|
||||
# PREFILL_BLOCK_SIZE default 128
|
||||
# DECODE_BLOCK_SIZE default 128
|
||||
# CPU_BYTES default 209715200 (200 MB)
|
||||
# VLLM_SERVE_EXTRA_ARGS comma-separated extra args for vllm serve
|
||||
# --decoder-first toggle decoder-first proxy mode
|
||||
#
|
||||
# Examples:
|
||||
# bash tests/v1/kv_offload/tiering/p2p/run_accuracy_test.sh
|
||||
# NUM_PREFILL_INSTANCES=2 NUM_DECODE_INSTANCES=2 \
|
||||
# bash tests/v1/kv_offload/tiering/p2p/run_accuracy_test.sh
|
||||
# bash tests/v1/kv_offload/tiering/p2p/run_accuracy_test.sh --decoder-first
|
||||
|
||||
set -xe
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Args
|
||||
# ---------------------------------------------------------------------------
|
||||
DECODER_FIRST="false"
|
||||
|
||||
while [[ $# -gt 0 ]]; do
|
||||
case $1 in
|
||||
--decoder-first)
|
||||
DECODER_FIRST="true"
|
||||
shift 1
|
||||
;;
|
||||
*)
|
||||
echo "Unknown option $1"
|
||||
echo "Usage: $0 [--decoder-first]"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
done
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Models
|
||||
# ---------------------------------------------------------------------------
|
||||
MODEL_NAMES=${MODEL_NAMES:-}
|
||||
if [[ -n "$MODEL_NAMES" ]]; then
|
||||
# shellcheck disable=SC2206
|
||||
MODELS=($MODEL_NAMES)
|
||||
else
|
||||
MODELS=(
|
||||
"meta-llama/Llama-3.2-1B-Instruct"
|
||||
)
|
||||
fi
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Defaults
|
||||
# ---------------------------------------------------------------------------
|
||||
NUM_PREFILL_INSTANCES=${NUM_PREFILL_INSTANCES:-1}
|
||||
NUM_DECODE_INSTANCES=${NUM_DECODE_INSTANCES:-1}
|
||||
PREFILLER_TP_SIZE=${PREFILLER_TP_SIZE:-1}
|
||||
DECODER_TP_SIZE=${DECODER_TP_SIZE:-1}
|
||||
GPU_MEMORY_UTILIZATION=${GPU_MEMORY_UTILIZATION:-0.45}
|
||||
MAX_MODEL_LEN=${MAX_MODEL_LEN:-512}
|
||||
PREFILL_BLOCK_SIZE=${PREFILL_BLOCK_SIZE:-128}
|
||||
DECODE_BLOCK_SIZE=${DECODE_BLOCK_SIZE:-128}
|
||||
CPU_BYTES=${CPU_BYTES:-209715200}
|
||||
VLLM_SERVE_EXTRA_ARGS=${VLLM_SERVE_EXTRA_ARGS:-}
|
||||
|
||||
# Base ports — per-instance offsets layered on top.
|
||||
PREFILL_HTTP_BASE=8100
|
||||
DECODE_HTTP_BASE=8200
|
||||
PREFILL_PD_BASE=7777
|
||||
DECODE_PD_BASE=$((PREFILL_PD_BASE + NUM_PREFILL_INSTANCES))
|
||||
PROXY_PORT=8192
|
||||
P2P_HOST=127.0.0.1
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Resolve repo root + venv (works in .venv and /workspace/venv pods)
|
||||
# ---------------------------------------------------------------------------
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../../.." && pwd -P)}"
|
||||
|
||||
if [[ -z "${VLLM_BIN:-}" ]]; then
|
||||
if [[ -x "${GIT_ROOT}/.venv/bin/vllm" ]]; then
|
||||
VLLM_BIN="${GIT_ROOT}/.venv/bin/vllm"
|
||||
elif [[ -x "/workspace/venv/bin/vllm" ]]; then
|
||||
VLLM_BIN="/workspace/venv/bin/vllm"
|
||||
else
|
||||
VLLM_BIN="$(command -v vllm)"
|
||||
fi
|
||||
fi
|
||||
if [[ -z "${PYTHON_BIN:-}" ]]; then
|
||||
if [[ -x "${GIT_ROOT}/.venv/bin/python" ]]; then
|
||||
PYTHON_BIN="${GIT_ROOT}/.venv/bin/python"
|
||||
elif [[ -x "/workspace/venv/bin/python" ]]; then
|
||||
PYTHON_BIN="/workspace/venv/bin/python"
|
||||
else
|
||||
PYTHON_BIN="$(command -v python3 || command -v python)"
|
||||
fi
|
||||
fi
|
||||
echo "Using vllm: ${VLLM_BIN}"
|
||||
echo "Using python: ${PYTHON_BIN}"
|
||||
|
||||
SMI_BIN=$(command -v nvidia-smi || command -v rocm-smi || echo "")
|
||||
|
||||
# Trap SIGINT/SIGTERM/EXIT to kill background jobs.
|
||||
trap 'kill $(jobs -pr) 2>/dev/null || true' SIGINT SIGTERM EXIT
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
wait_for_server() {
|
||||
local port=$1
|
||||
timeout 1200 bash -c "
|
||||
until curl -s localhost:${port}/v1/completions > /dev/null; do
|
||||
sleep 1
|
||||
done" && return 0 || return 1
|
||||
}
|
||||
|
||||
cleanup_instances() {
|
||||
echo "Cleaning up any running vLLM / proxy instances..."
|
||||
pkill -f "vllm serve" || true
|
||||
pkill -f "p2p_connector_proxy.py" || true
|
||||
sleep 2
|
||||
}
|
||||
|
||||
get_num_gpus() {
|
||||
if [[ "$SMI_BIN" == *"nvidia"* ]]; then
|
||||
$SMI_BIN --query-gpu=name --format=csv,noheader | wc -l
|
||||
elif [[ "$SMI_BIN" == *"rocm"* ]]; then
|
||||
$SMI_BIN -l | grep -c GPU
|
||||
else
|
||||
echo "1"
|
||||
fi
|
||||
}
|
||||
|
||||
# Build the OffloadingConnector kv-transfer-config for a given PD port.
|
||||
# Mirrors deploy_local.sh:131.
|
||||
build_kv_config() {
|
||||
local pd_port=$1
|
||||
printf '{"kv_connector":"OffloadingConnector","kv_role":"kv_both","kv_connector_extra_config":{"spec_name":"TieringOffloadingSpec","cpu_bytes_to_use":%s,"secondary_tiers":[{"type":"p2p","host":"%s","port":%s}]}}' \
|
||||
"${CPU_BYTES}" "${P2P_HOST}" "${pd_port}"
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Per-model run
|
||||
# ---------------------------------------------------------------------------
|
||||
run_tests_for_model() {
|
||||
local model_name=$1
|
||||
echo "================================"
|
||||
echo "Testing model: $model_name"
|
||||
echo " prefillers=${NUM_PREFILL_INSTANCES} (tp=${PREFILLER_TP_SIZE})"
|
||||
echo " decoders=${NUM_DECODE_INSTANCES} (tp=${DECODER_TP_SIZE})"
|
||||
echo " decoder_first=${DECODER_FIRST}"
|
||||
echo "================================"
|
||||
|
||||
PREFILL_HOSTS=()
|
||||
PREFILL_PORTS=()
|
||||
PREFILL_PD_PORTS=()
|
||||
DECODE_HOSTS=()
|
||||
DECODE_PORTS=()
|
||||
DECODE_PD_PORTS=()
|
||||
|
||||
local num_gpus
|
||||
num_gpus=$(get_num_gpus)
|
||||
local next_gpu=0
|
||||
|
||||
# ---- Prefillers ----
|
||||
for i in $(seq 0 $((NUM_PREFILL_INSTANCES-1))); do
|
||||
local gpu_id=$((i * PREFILLER_TP_SIZE % num_gpus))
|
||||
local cuda_devs="${gpu_id}"
|
||||
for (( j=1; j < PREFILLER_TP_SIZE; j++ )); do
|
||||
cuda_devs="${cuda_devs},$(((gpu_id + j) % num_gpus))"
|
||||
done
|
||||
next_gpu=$(((gpu_id + PREFILLER_TP_SIZE) % num_gpus))
|
||||
|
||||
local http_port=$((PREFILL_HTTP_BASE + i))
|
||||
local pd_port=$((PREFILL_PD_BASE + i))
|
||||
local kv_cfg
|
||||
kv_cfg=$(build_kv_config "${pd_port}")
|
||||
|
||||
echo "Prefiller $i: gpu=[${cuda_devs}] http=${http_port} pd=${pd_port}"
|
||||
|
||||
BASE_CMD="CUDA_VISIBLE_DEVICES=${cuda_devs} \
|
||||
PYTHONHASHSEED=42 \
|
||||
${VLLM_BIN} serve ${model_name} \
|
||||
--port ${http_port} \
|
||||
--enforce-eager \
|
||||
--block-size ${PREFILL_BLOCK_SIZE} \
|
||||
--gpu-memory-utilization ${GPU_MEMORY_UTILIZATION} \
|
||||
--max-model-len ${MAX_MODEL_LEN} \
|
||||
--tensor-parallel-size ${PREFILLER_TP_SIZE} \
|
||||
--kv-transfer-config '${kv_cfg}'"
|
||||
|
||||
if [[ -n "$VLLM_SERVE_EXTRA_ARGS" ]]; then
|
||||
IFS=',' read -r -a extra_args <<< "$VLLM_SERVE_EXTRA_ARGS"
|
||||
for arg in "${extra_args[@]}"; do
|
||||
BASE_CMD="${BASE_CMD} $arg"
|
||||
done
|
||||
fi
|
||||
|
||||
eval "${BASE_CMD} &"
|
||||
|
||||
PREFILL_HOSTS+=("${P2P_HOST}")
|
||||
PREFILL_PORTS+=("${http_port}")
|
||||
PREFILL_PD_PORTS+=("${pd_port}")
|
||||
done
|
||||
|
||||
# ---- Decoders ----
|
||||
for i in $(seq 0 $((NUM_DECODE_INSTANCES-1))); do
|
||||
local gpu_id=$(((next_gpu + i * DECODER_TP_SIZE) % num_gpus))
|
||||
local cuda_devs="${gpu_id}"
|
||||
for (( j=1; j < DECODER_TP_SIZE; j++ )); do
|
||||
cuda_devs="${cuda_devs},$(((gpu_id + j) % num_gpus))"
|
||||
done
|
||||
|
||||
local http_port=$((DECODE_HTTP_BASE + i))
|
||||
local pd_port=$((DECODE_PD_BASE + i))
|
||||
local kv_cfg
|
||||
kv_cfg=$(build_kv_config "${pd_port}")
|
||||
|
||||
echo "Decoder $i: gpu=[${cuda_devs}] http=${http_port} pd=${pd_port}"
|
||||
|
||||
BASE_CMD="CUDA_VISIBLE_DEVICES=${cuda_devs} \
|
||||
PYTHONHASHSEED=42 \
|
||||
${VLLM_BIN} serve ${model_name} \
|
||||
--port ${http_port} \
|
||||
--enforce-eager \
|
||||
--block-size ${DECODE_BLOCK_SIZE} \
|
||||
--gpu-memory-utilization ${GPU_MEMORY_UTILIZATION} \
|
||||
--max-model-len ${MAX_MODEL_LEN} \
|
||||
--tensor-parallel-size ${DECODER_TP_SIZE} \
|
||||
--kv-transfer-config '${kv_cfg}'"
|
||||
|
||||
if [[ -n "$VLLM_SERVE_EXTRA_ARGS" ]]; then
|
||||
IFS=',' read -r -a extra_args <<< "$VLLM_SERVE_EXTRA_ARGS"
|
||||
for arg in "${extra_args[@]}"; do
|
||||
BASE_CMD="${BASE_CMD} $arg"
|
||||
done
|
||||
fi
|
||||
|
||||
eval "${BASE_CMD} &"
|
||||
|
||||
DECODE_HOSTS+=("${P2P_HOST}")
|
||||
DECODE_PORTS+=("${http_port}")
|
||||
DECODE_PD_PORTS+=("${pd_port}")
|
||||
done
|
||||
|
||||
# ---- Wait for HTTP readiness ----
|
||||
for port in "${PREFILL_PORTS[@]}"; do
|
||||
echo "Waiting for prefill instance on port $port to start..."
|
||||
wait_for_server "$port"
|
||||
done
|
||||
for port in "${DECODE_PORTS[@]}"; do
|
||||
echo "Waiting for decode instance on port $port to start..."
|
||||
wait_for_server "$port"
|
||||
done
|
||||
|
||||
# ---- Proxy ----
|
||||
# The proxy currently advertises a single prefiller PD address to decoders.
|
||||
# For the 1xM and matched NxM common cases the first prefiller's PD coords
|
||||
# are the right pick; multi-prefiller PD round-robin is a follow-up.
|
||||
PROXY_CMD="${PYTHON_BIN} ${SCRIPT_DIR}/p2p_connector_proxy.py \
|
||||
--port ${PROXY_PORT} \
|
||||
--host ${P2P_HOST} \
|
||||
--prefiller-hosts ${PREFILL_HOSTS[*]} \
|
||||
--prefiller-ports ${PREFILL_PORTS[*]} \
|
||||
--decoder-hosts ${DECODE_HOSTS[*]} \
|
||||
--decoder-ports ${DECODE_PORTS[*]} \
|
||||
--p2p-connector-host ${P2P_HOST} \
|
||||
--p2p-connector-port ${PREFILL_PD_PORTS[0]} \
|
||||
--decoder-p2p-connector-host ${P2P_HOST} \
|
||||
--decoder-p2p-connector-port ${DECODE_PD_PORTS[0]}"
|
||||
|
||||
if [[ "${DECODER_FIRST}" == "true" ]]; then
|
||||
PROXY_CMD="${PROXY_CMD} --decoder-first"
|
||||
fi
|
||||
|
||||
echo "Starting proxy: ${PROXY_CMD}"
|
||||
eval "${PROXY_CMD} &"
|
||||
|
||||
sleep 5
|
||||
|
||||
# ---- Run accuracy test (reused from nixl_integration) ----
|
||||
echo "Running tests for $model_name"
|
||||
TEST_MODEL=$model_name "${PYTHON_BIN}" -m pytest -s -x \
|
||||
"${GIT_ROOT}/tests/v1/kv_connector/nixl_integration/test_accuracy.py"
|
||||
|
||||
cleanup_instances
|
||||
sleep 3
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Drive
|
||||
# ---------------------------------------------------------------------------
|
||||
for model in "${MODELS[@]}"; do
|
||||
run_tests_for_model "$model"
|
||||
done
|
||||
|
||||
echo "All tests completed!"
|
||||
@@ -0,0 +1,330 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Unit tests for DataTransport base class and NixlTransport."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ctypes
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
|
||||
from vllm.v1.kv_offload.tiering.p2p.data.base import PollResult
|
||||
from vllm.v1.kv_offload.tiering.p2p.data.nixl import NixlTransport
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DataTransport base class tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDataTransportBase:
|
||||
"""Tests for the DataTransport abstract base properties."""
|
||||
|
||||
def _make_view(self, num_blocks: int = 8, block_len: int = 1024) -> memoryview:
|
||||
"""Create a memoryview with the given shape."""
|
||||
buf = np.zeros((num_blocks, block_len), dtype=np.uint8)
|
||||
return memoryview(buf)
|
||||
|
||||
def test_properties(self):
|
||||
"""base_addr, num_blocks, block_len are set from memoryview shape."""
|
||||
view = self._make_view(num_blocks=4, block_len=2048)
|
||||
|
||||
# Use NixlTransport (concrete) with NIXL mocked away
|
||||
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
|
||||
transport = NixlTransport("test:1", view)
|
||||
|
||||
assert transport.num_blocks == 4
|
||||
assert transport.block_len == 2048
|
||||
assert transport.base_addr == ctypes.addressof(ctypes.c_char.from_buffer(view))
|
||||
|
||||
def test_config_fingerprint_empty_when_no_fields(self):
|
||||
"""No config fields → empty fingerprint."""
|
||||
view = self._make_view()
|
||||
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
|
||||
transport = NixlTransport("test:1", view, config_fields=None)
|
||||
assert transport.config_fingerprint == ""
|
||||
|
||||
def test_config_fingerprint_deterministic(self):
|
||||
"""Same config fields → same fingerprint."""
|
||||
view = self._make_view()
|
||||
fields = {"model": "llama", "dtype": "float16", "block_size_factor": 1}
|
||||
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
|
||||
t1 = NixlTransport("test:1", view, config_fields=fields)
|
||||
t2 = NixlTransport("test:2", view, config_fields=fields)
|
||||
assert t1.config_fingerprint == t2.config_fingerprint
|
||||
assert len(t1.config_fingerprint) == 16
|
||||
|
||||
def test_config_fingerprint_differs_for_different_fields(self):
|
||||
"""Different config fields → different fingerprint."""
|
||||
view = self._make_view()
|
||||
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
|
||||
t1 = NixlTransport("test:1", view, config_fields={"model": "a"})
|
||||
t2 = NixlTransport("test:2", view, config_fields={"model": "b"})
|
||||
assert t1.config_fingerprint != t2.config_fingerprint
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# NixlTransport tests (with mocked NIXL agent)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestNixlTransportWithMockedAgent:
|
||||
"""Tests for NixlTransport logic with a mocked NIXL agent."""
|
||||
|
||||
def _make_transport(self) -> NixlTransport:
|
||||
"""Create a NixlTransport with mocked NIXL internals."""
|
||||
view = memoryview(np.zeros((8, 1024), dtype=np.uint8))
|
||||
|
||||
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
|
||||
transport = NixlTransport("test:1", view)
|
||||
|
||||
# Manually set up a mock agent after construction
|
||||
agent = MagicMock()
|
||||
agent.add_remote_agent.return_value = "nixl-peer-name"
|
||||
agent.get_xfer_descs.return_value = MagicMock()
|
||||
agent.prep_xfer_dlist.return_value = MagicMock()
|
||||
agent.make_prepped_xfer.return_value = MagicMock(name="handle")
|
||||
agent.transfer.return_value = None
|
||||
agent.check_xfer_state.return_value = "PROC"
|
||||
agent.get_agent_metadata.return_value = b"test-metadata"
|
||||
|
||||
transport._agent = agent
|
||||
transport._local_dlist = MagicMock()
|
||||
return transport
|
||||
|
||||
def test_available_false_without_nixl(self):
|
||||
"""Without NIXL installed, available is False."""
|
||||
view = memoryview(np.zeros((4, 512), dtype=np.uint8))
|
||||
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
|
||||
transport = NixlTransport("test:1", view)
|
||||
assert transport.available is False
|
||||
|
||||
def test_available_true_with_agent(self):
|
||||
transport = self._make_transport()
|
||||
assert transport.available is True
|
||||
|
||||
def test_get_agent_metadata(self):
|
||||
transport = self._make_transport()
|
||||
assert transport.get_agent_metadata() == b"test-metadata"
|
||||
|
||||
def test_write_blocks_returns_none_for_unknown_peer(self):
|
||||
"""write_blocks returns None if peer not registered."""
|
||||
transport = self._make_transport()
|
||||
result = transport.write_blocks("unknown:1", [0, 1], [2, 3])
|
||||
assert result is None
|
||||
|
||||
def test_write_blocks_returns_transfer_id(self):
|
||||
"""write_blocks returns an integer transfer_id on success."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid = transport.write_blocks("peer:1", [0, 1], [2, 3])
|
||||
assert tid is not None
|
||||
assert isinstance(tid, int)
|
||||
|
||||
def test_write_blocks_increments_transfer_id(self):
|
||||
"""Each write_blocks call gets a unique transfer_id."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid1 = transport.write_blocks("peer:1", [0], [1])
|
||||
tid2 = transport.write_blocks("peer:1", [2], [3])
|
||||
assert tid1 != tid2
|
||||
|
||||
def test_poll_empty_when_no_inflight(self):
|
||||
"""poll returns empty when nothing is inflight."""
|
||||
transport = self._make_transport()
|
||||
result = transport.poll()
|
||||
assert result == PollResult(done=(), failed=())
|
||||
|
||||
def test_poll_returns_done_when_transfer_completes(self):
|
||||
"""Completed transfer appears in poll().done."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid = transport.write_blocks("peer:1", [0], [1])
|
||||
|
||||
# Simulate completion
|
||||
transport._agent.check_xfer_state.return_value = "DONE"
|
||||
result = transport.poll()
|
||||
|
||||
assert tid in result.done
|
||||
assert result.failed == ()
|
||||
# Handle released
|
||||
transport._agent.release_xfer_handle.assert_called()
|
||||
|
||||
def test_poll_returns_failed_for_error_state(self):
|
||||
"""Transfer in error state appears in poll().failed."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid = transport.write_blocks("peer:1", [0], [1])
|
||||
|
||||
transport._agent.check_xfer_state.return_value = "ERR"
|
||||
result = transport.poll()
|
||||
|
||||
assert result.done == ()
|
||||
assert tid in result.failed
|
||||
|
||||
def test_poll_ignores_in_progress(self):
|
||||
"""Transfers in PROC/PEND state stay inflight."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
transport.write_blocks("peer:1", [0], [1])
|
||||
|
||||
transport._agent.check_xfer_state.return_value = "PROC"
|
||||
result = transport.poll()
|
||||
assert result.done == ()
|
||||
assert result.failed == ()
|
||||
|
||||
transport._agent.check_xfer_state.return_value = "PEND"
|
||||
result = transport.poll()
|
||||
assert result.done == ()
|
||||
assert result.failed == ()
|
||||
|
||||
def test_cancel_removes_inflight(self):
|
||||
"""cancel removes transfers and releases handles."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid = transport.write_blocks("peer:1", [0], [1])
|
||||
assert tid in transport._inflight
|
||||
|
||||
result = transport.cancel([tid])
|
||||
assert result == []
|
||||
assert tid not in transport._inflight
|
||||
transport._agent.release_xfer_handle.assert_called()
|
||||
|
||||
def test_cancel_ignores_unknown_ids(self):
|
||||
"""cancel with unknown IDs doesn't crash."""
|
||||
transport = self._make_transport()
|
||||
assert transport.cancel([999, 1000]) == []
|
||||
assert transport.cancel([999, 1000], mode="wait") == []
|
||||
|
||||
def test_cancel_wait_release_succeeds(self):
|
||||
"""wait-mode cancel that succeeds pops the entry and returns []."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid = transport.write_blocks("peer:1", [0], [1])
|
||||
assert tid in transport._inflight
|
||||
|
||||
result = transport.cancel([tid], mode="wait")
|
||||
assert result == []
|
||||
assert tid not in transport._inflight
|
||||
transport._agent.release_xfer_handle.assert_called_once()
|
||||
|
||||
def test_cancel_wait_release_raises(self):
|
||||
"""wait-mode cancel keeps the entry and returns the tid on raise."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid = transport.write_blocks("peer:1", [0], [1])
|
||||
transport._agent.release_xfer_handle.side_effect = RuntimeError(
|
||||
"NIXL_ERR_REPOST_ACTIVE"
|
||||
)
|
||||
|
||||
result = transport.cancel([tid], mode="wait")
|
||||
assert result == [tid]
|
||||
assert tid in transport._inflight
|
||||
|
||||
def test_cancel_wait_then_poll_completes(self):
|
||||
"""A wait-cancel that left a tid pending later completes via poll."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
|
||||
tid = transport.write_blocks("peer:1", [0], [1])
|
||||
transport._agent.release_xfer_handle.side_effect = RuntimeError("busy")
|
||||
assert transport.cancel([tid], mode="wait") == [tid]
|
||||
assert tid in transport._inflight
|
||||
|
||||
transport._agent.release_xfer_handle.side_effect = None
|
||||
transport._agent.check_xfer_state.return_value = "DONE"
|
||||
|
||||
result = transport.poll()
|
||||
assert tid in result.done
|
||||
assert tid not in transport._inflight
|
||||
|
||||
def test_add_and_remove_remote_peer(self):
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
assert "peer:1" in transport._remote_dlists
|
||||
|
||||
transport.remove_remote_peer("peer:1")
|
||||
assert "peer:1" not in transport._remote_dlists
|
||||
transport._agent.release_dlist_handle.assert_called()
|
||||
transport._agent.remove_remote_agent.assert_called()
|
||||
|
||||
def test_close_releases_everything(self):
|
||||
"""close releases all handles and clears state."""
|
||||
transport = self._make_transport()
|
||||
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
|
||||
transport.write_blocks("peer:1", [0], [1])
|
||||
|
||||
transport.close()
|
||||
assert transport._agent is None
|
||||
assert transport._inflight == {}
|
||||
assert transport._remote_dlists == {}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# NIXL agent-config selection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestNixlAgentConfigSelection:
|
||||
"""Tests that backends/num_threads pick the right nixl_agent_config call.
|
||||
|
||||
Mirrors the conditional in
|
||||
vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py:325-329.
|
||||
"""
|
||||
|
||||
def _make_view(self) -> memoryview:
|
||||
return memoryview(np.zeros((4, 512), dtype=np.uint8))
|
||||
|
||||
def test_non_ucx_backends_passes_backends_kwarg(self):
|
||||
"""When any non-UCX backend is requested, pass backends + telemetry."""
|
||||
agent_cls = MagicMock()
|
||||
config_fn = MagicMock(return_value=MagicMock(name="cfg"))
|
||||
with (
|
||||
patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", agent_cls),
|
||||
patch(
|
||||
"vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgentConfig", config_fn
|
||||
),
|
||||
):
|
||||
NixlTransport("test:1", self._make_view(), backends=["MOONCAKE"])
|
||||
|
||||
config_fn.assert_called_once_with(backends=["MOONCAKE"], capture_telemetry=True)
|
||||
# num_threads must NOT be passed on the non-UCX branch.
|
||||
assert "num_threads" not in config_fn.call_args.kwargs
|
||||
|
||||
def test_ucx_only_passes_num_threads(self):
|
||||
"""UCX-only configuration passes num_threads + telemetry, no backends."""
|
||||
agent_cls = MagicMock()
|
||||
config_fn = MagicMock(return_value=MagicMock(name="cfg"))
|
||||
with (
|
||||
patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", agent_cls),
|
||||
patch(
|
||||
"vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgentConfig", config_fn
|
||||
),
|
||||
):
|
||||
NixlTransport("test:1", self._make_view(), num_threads=8)
|
||||
|
||||
config_fn.assert_called_once_with(num_threads=8, capture_telemetry=True)
|
||||
assert "backends" not in config_fn.call_args.kwargs
|
||||
|
||||
def test_default_backends_is_ucx_only(self):
|
||||
"""No backends arg → defaults to UCX-only branch."""
|
||||
agent_cls = MagicMock()
|
||||
config_fn = MagicMock(return_value=MagicMock(name="cfg"))
|
||||
with (
|
||||
patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", agent_cls),
|
||||
patch(
|
||||
"vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgentConfig", config_fn
|
||||
),
|
||||
):
|
||||
NixlTransport("test:1", self._make_view())
|
||||
|
||||
# Default num_threads=4, no backends kwarg.
|
||||
config_fn.assert_called_once_with(num_threads=4, capture_telemetry=True)
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,230 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Unit tests for vllm.v1.kv_offload.tiering.p2p.control.zmq."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import socket
|
||||
import time
|
||||
|
||||
import pytest
|
||||
import zmq
|
||||
|
||||
from vllm.v1.kv_offload.tiering.p2p.control.zmq import (
|
||||
ZmqConnection,
|
||||
ZmqTransport,
|
||||
_Sockets,
|
||||
)
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
"""Find a free TCP port."""
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(("127.0.0.1", 0))
|
||||
return s.getsockname()[1]
|
||||
|
||||
|
||||
def _make_transport(host: str = "127.0.0.1", attempts: int = 8):
|
||||
"""Construct a ZmqTransport on a fresh port, retrying on bind collisions.
|
||||
|
||||
Why: _free_port() releases the probe socket before ZmqTransport binds the
|
||||
same port — a parallel test run can steal it in between. Retrying on
|
||||
ZMQError/OSError closes that race without a production change.
|
||||
"""
|
||||
last_err: Exception | None = None
|
||||
for _ in range(attempts):
|
||||
port = _free_port()
|
||||
try:
|
||||
return ZmqTransport(f"{host}:{port}", host, port), port
|
||||
except (zmq.ZMQError, OSError) as e:
|
||||
last_err = e
|
||||
assert last_err is not None
|
||||
raise last_err
|
||||
|
||||
|
||||
def _wait_for_inbound(transport: ZmqTransport, deadline: float = 2.0):
|
||||
"""Poll until at least one new inbound connection is accepted, or fail."""
|
||||
end = time.monotonic() + deadline
|
||||
while time.monotonic() < end:
|
||||
new = transport.poll()
|
||||
if new:
|
||||
return new
|
||||
time.sleep(0.005)
|
||||
raise AssertionError(f"no inbound connection within {deadline}s")
|
||||
|
||||
|
||||
def _wait_for_messages(
|
||||
transport: ZmqTransport,
|
||||
conn: ZmqConnection,
|
||||
n: int,
|
||||
deadline: float = 2.0,
|
||||
) -> list[dict]:
|
||||
"""Poll until `conn` has received at least `n` messages, then return them."""
|
||||
end = time.monotonic() + deadline
|
||||
msgs: list[dict] = []
|
||||
while time.monotonic() < end:
|
||||
transport.poll()
|
||||
msgs.extend(conn.recv())
|
||||
if len(msgs) >= n:
|
||||
return msgs
|
||||
time.sleep(0.005)
|
||||
raise AssertionError(f"got {len(msgs)}/{n} messages within {deadline}s")
|
||||
|
||||
|
||||
def _make_mock_connection(peer_id: str = "test:1234") -> ZmqConnection:
|
||||
"""Create a ZmqConnection with mock sockets for unit testing."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
sockets = _Sockets(dealer=MagicMock(), monitor=MagicMock())
|
||||
return ZmqConnection(peer_id, sockets)
|
||||
|
||||
|
||||
class TestZmqConnection:
|
||||
"""Tests for ZmqConnection in isolation (no real sockets)."""
|
||||
|
||||
def test_enqueue_and_recv(self):
|
||||
"""Messages enqueued are returned by recv() in order."""
|
||||
conn = _make_mock_connection()
|
||||
|
||||
conn.enqueue({"type": "a"})
|
||||
conn.enqueue({"type": "b"})
|
||||
|
||||
msgs = conn.recv()
|
||||
assert list(msgs) == [{"type": "a"}, {"type": "b"}]
|
||||
# Second recv is empty
|
||||
assert not conn.recv()
|
||||
|
||||
def test_recv_returns_empty_initially(self):
|
||||
conn = _make_mock_connection()
|
||||
assert not conn.recv()
|
||||
|
||||
def test_alive_initially_true(self):
|
||||
conn = _make_mock_connection()
|
||||
assert conn.alive is True
|
||||
|
||||
def test_mark_dead(self):
|
||||
conn = _make_mock_connection()
|
||||
conn.mark_dead()
|
||||
assert conn.alive is False
|
||||
|
||||
def test_send_raises_when_closed(self):
|
||||
conn = _make_mock_connection()
|
||||
conn.mark_dead()
|
||||
|
||||
with pytest.raises(RuntimeError, match="closed connection"):
|
||||
conn.send({"type": "test"})
|
||||
|
||||
|
||||
class TestZmqTransportConnectivity:
|
||||
"""Integration tests for ZmqTransport with real ZMQ sockets."""
|
||||
|
||||
def test_connect_and_send_message(self):
|
||||
"""Two transports can connect and exchange messages."""
|
||||
transport_a, port_a = _make_transport()
|
||||
transport_b, port_b = _make_transport()
|
||||
|
||||
try:
|
||||
peer_a_id = f"127.0.0.1:{port_a}"
|
||||
conn_b_to_a = transport_b.connect(peer_a_id)
|
||||
conn_b_to_a.send({"type": "hello", "data": 42})
|
||||
|
||||
new_conns = _wait_for_inbound(transport_a)
|
||||
assert len(new_conns) == 1
|
||||
|
||||
conn_a_from_b = new_conns[0]
|
||||
assert conn_a_from_b.peer_id == f"127.0.0.1:{port_b}"
|
||||
|
||||
msgs = _wait_for_messages(transport_a, conn_a_from_b, 1)
|
||||
assert msgs == [{"type": "hello", "data": 42}]
|
||||
finally:
|
||||
transport_a.close()
|
||||
transport_b.close()
|
||||
|
||||
def test_bidirectional_messaging(self):
|
||||
"""Both sides can send and receive after connection."""
|
||||
transport_a, port_a = _make_transport()
|
||||
transport_b, _ = _make_transport()
|
||||
|
||||
try:
|
||||
conn_b = transport_b.connect(f"127.0.0.1:{port_a}")
|
||||
conn_b.send({"type": "connect", "from": "b"})
|
||||
|
||||
new_conns = _wait_for_inbound(transport_a)
|
||||
assert len(new_conns) == 1
|
||||
conn_a = new_conns[0]
|
||||
|
||||
conn_a.send({"type": "reply", "from": "a"})
|
||||
|
||||
msgs = _wait_for_messages(transport_b, conn_b, 1)
|
||||
assert msgs == [{"type": "reply", "from": "a"}]
|
||||
finally:
|
||||
transport_a.close()
|
||||
transport_b.close()
|
||||
|
||||
def test_poll_returns_empty_when_no_connections(self):
|
||||
transport, _ = _make_transport()
|
||||
try:
|
||||
assert not transport.poll()
|
||||
finally:
|
||||
transport.close()
|
||||
|
||||
def test_multiple_messages(self):
|
||||
"""Multiple messages are buffered and returned together."""
|
||||
transport_a, port_a = _make_transport()
|
||||
transport_b, _ = _make_transport()
|
||||
|
||||
try:
|
||||
conn_b = transport_b.connect(f"127.0.0.1:{port_a}")
|
||||
conn_b.send({"seq": 1})
|
||||
conn_b.send({"seq": 2})
|
||||
conn_b.send({"seq": 3})
|
||||
|
||||
new_conns = _wait_for_inbound(transport_a)
|
||||
assert len(new_conns) == 1
|
||||
conn_a = new_conns[0]
|
||||
|
||||
msgs = _wait_for_messages(transport_a, conn_a, 3)
|
||||
assert [m["seq"] for m in msgs] == [1, 2, 3]
|
||||
finally:
|
||||
transport_a.close()
|
||||
transport_b.close()
|
||||
|
||||
def test_duplicate_connect_asserts(self):
|
||||
"""Connecting to the same peer twice raises AssertionError."""
|
||||
# port_a is never bound — we just need a syntactically-valid peer id.
|
||||
port_a = _free_port()
|
||||
transport_b, _ = _make_transport()
|
||||
try:
|
||||
transport_b.connect(f"127.0.0.1:{port_a}")
|
||||
with pytest.raises(AssertionError, match="already exists"):
|
||||
transport_b.connect(f"127.0.0.1:{port_a}")
|
||||
finally:
|
||||
transport_b.close()
|
||||
|
||||
def test_dead_connection_removed_on_poll(self):
|
||||
"""Dead connections are cleaned up during poll."""
|
||||
transport_a, port_a = _make_transport()
|
||||
transport_b, _ = _make_transport()
|
||||
|
||||
try:
|
||||
conn_b = transport_b.connect(f"127.0.0.1:{port_a}")
|
||||
conn_b.send({"type": "hello"})
|
||||
|
||||
new_conns = _wait_for_inbound(transport_a)
|
||||
assert len(new_conns) == 1
|
||||
|
||||
# Mark the inbound connection dead manually.
|
||||
new_conns[0].mark_dead()
|
||||
|
||||
# Pruning is synchronous within poll().
|
||||
transport_a.poll()
|
||||
assert len(transport_a._connections) == 0
|
||||
finally:
|
||||
transport_a.close()
|
||||
transport_b.close()
|
||||
|
||||
def test_close_is_idempotent(self):
|
||||
"""Calling close() twice doesn't raise."""
|
||||
transport, _ = _make_transport()
|
||||
transport.close()
|
||||
transport.close() # should not raise
|
||||
@@ -1257,3 +1257,127 @@ def test_thinking_budget_long_thinking_section_end_marker_found_at_correct_index
|
||||
|
||||
assert h._state[0]["start_thinking"] == 0
|
||||
assert h._state[0]["end_thinking"] == expected_end_idx
|
||||
|
||||
|
||||
# --- Thinking budget re-entry tests (issue #43708) ---
|
||||
# Regression tests: after budget forces end-of-thinking token sequence,
|
||||
# the state machine must detect and enforce budget on subsequent blocks.
|
||||
|
||||
|
||||
class TestThinkingBudgetReentry:
|
||||
THINK_START = 100
|
||||
THINK_END_SINGLE = [200]
|
||||
THINK_END_MULTI = [200, 201, 202]
|
||||
BUDGET = 5
|
||||
CONTENT_TOKEN = 50
|
||||
THINK_TOKEN = 60
|
||||
|
||||
@staticmethod
|
||||
def _make_holder(end_token_ids: list[int]) -> ThinkingBudgetStateHolder:
|
||||
class FakeReasoningConfig:
|
||||
reasoning_start_token_ids = [TestThinkingBudgetReentry.THINK_START]
|
||||
reasoning_end_token_ids: list[int] = []
|
||||
enabled = True
|
||||
|
||||
cfg = FakeReasoningConfig()
|
||||
cfg.reasoning_end_token_ids = end_token_ids
|
||||
return ThinkingBudgetStateHolder(
|
||||
reasoning_config=cfg,
|
||||
max_num_seqs=8,
|
||||
num_spec_tokens=0,
|
||||
device=torch.device("cpu"),
|
||||
is_pin_memory=False,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _sync_batch(holder: ThinkingBudgetStateHolder, budget: int) -> None:
|
||||
holder.sync_batch(
|
||||
BatchUpdate(
|
||||
batch_size=1,
|
||||
removed=(),
|
||||
added=[(0, SamplingParams(thinking_token_budget=budget), None, [])],
|
||||
moved=(),
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _step(holder: ThinkingBudgetStateHolder, output_tok_ids: list[int]) -> None:
|
||||
holder.update_state(
|
||||
output_token_ids=[output_tok_ids],
|
||||
spec_token_ids=None,
|
||||
repeat_indices=None,
|
||||
)
|
||||
|
||||
def _exhaust_budget(self, holder: ThinkingBudgetStateHolder) -> list[int]:
|
||||
output = [self.THINK_START]
|
||||
self._step(holder, list(output))
|
||||
for _ in range(self.BUDGET):
|
||||
output.append(self.THINK_TOKEN)
|
||||
self._step(holder, list(output))
|
||||
|
||||
assert holder._state[0]["in_end"]
|
||||
return output
|
||||
|
||||
def _accept_end_tokens(
|
||||
self,
|
||||
holder: ThinkingBudgetStateHolder,
|
||||
output: list[int],
|
||||
end_token_ids: list[int],
|
||||
) -> None:
|
||||
for tok in end_token_ids:
|
||||
output.append(tok)
|
||||
self._step(holder, list(output))
|
||||
|
||||
def test_single_token_end_reentry(self):
|
||||
holder = self._make_holder(self.THINK_END_SINGLE)
|
||||
self._sync_batch(holder, self.BUDGET)
|
||||
|
||||
output = self._exhaust_budget(holder)
|
||||
self._accept_end_tokens(holder, output, self.THINK_END_SINGLE)
|
||||
|
||||
for _ in range(3):
|
||||
output.append(self.CONTENT_TOKEN)
|
||||
self._step(holder, list(output))
|
||||
|
||||
output.append(self.THINK_START)
|
||||
self._step(holder, list(output))
|
||||
for _ in range(self.BUDGET):
|
||||
output.append(self.THINK_TOKEN)
|
||||
self._step(holder, list(output))
|
||||
|
||||
assert holder._state[0]["in_end"], (
|
||||
"Second thinking block must also be budget-enforced"
|
||||
)
|
||||
|
||||
def test_multi_token_end_reentry(self):
|
||||
holder = self._make_holder(self.THINK_END_MULTI)
|
||||
self._sync_batch(holder, self.BUDGET)
|
||||
|
||||
output = self._exhaust_budget(holder)
|
||||
self._accept_end_tokens(holder, output, self.THINK_END_MULTI)
|
||||
|
||||
assert not holder._state[0]["in_end"]
|
||||
|
||||
output.append(self.THINK_START)
|
||||
self._step(holder, list(output))
|
||||
for _ in range(self.BUDGET):
|
||||
output.append(self.THINK_TOKEN)
|
||||
self._step(holder, list(output))
|
||||
|
||||
assert holder._state[0]["in_end"], (
|
||||
"Immediate re-entry after multi-token end must be enforced"
|
||||
)
|
||||
|
||||
def test_single_block_not_broken(self):
|
||||
holder = self._make_holder(self.THINK_END_SINGLE)
|
||||
self._sync_batch(holder, self.BUDGET)
|
||||
|
||||
output = self._exhaust_budget(holder)
|
||||
self._accept_end_tokens(holder, output, self.THINK_END_SINGLE)
|
||||
|
||||
for _ in range(20):
|
||||
output.append(self.CONTENT_TOKEN)
|
||||
self._step(holder, list(output))
|
||||
|
||||
assert not holder._state[0]["in_end"]
|
||||
assert not holder._state[0]["in_think"]
|
||||
|
||||
@@ -1285,7 +1285,7 @@ def test_token_logprobs_large_batch_int64_row_offset():
|
||||
batch_size = 2**31 // vocab_size + 64 # batch_size * vocab_size > 2**31
|
||||
# logits (the large input) plus small logprob/rank outputs; ~1 GB headroom.
|
||||
required_bytes = batch_size * vocab_size * 4 + (1 << 30)
|
||||
if torch.cuda.mem_get_info()[0] < required_bytes:
|
||||
if torch.accelerator.get_memory_info()[0] < required_bytes:
|
||||
pytest.skip(f"needs ~{required_bytes / 1e9:.0f} GB of free GPU memory")
|
||||
|
||||
logits = torch.randn(batch_size, vocab_size, device=device, dtype=torch.float32)
|
||||
|
||||
@@ -426,7 +426,7 @@ class TestTritonTopkTopp:
|
||||
# logits is modified in place; the only extra device memory is the
|
||||
# per-SM scratch buffer (~num_sm * vocab), so allow ~1 GB of headroom.
|
||||
required_bytes = batch_size * vocab_size * 4 + (1 << 30)
|
||||
if torch.cuda.mem_get_info()[0] < required_bytes:
|
||||
if torch.accelerator.get_memory_info()[0] < required_bytes:
|
||||
pytest.skip(f"needs ~{required_bytes / 1e9:.0f} GB of free GPU memory")
|
||||
|
||||
logits = torch.randn(
|
||||
|
||||
@@ -7,6 +7,8 @@ import time
|
||||
from opentelemetry.sdk.environment_variables import OTEL_EXPORTER_OTLP_TRACES_INSECURE
|
||||
|
||||
from vllm import LLM, SamplingParams
|
||||
from vllm.distributed import cleanup_dist_env_and_memory
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.tracing import SpanAttributes
|
||||
|
||||
# Import shared fixtures from the tracing conftest
|
||||
@@ -23,6 +25,11 @@ def test_traces(
|
||||
):
|
||||
with monkeypatch.context() as m:
|
||||
m.setenv(OTEL_EXPORTER_OTLP_TRACES_INSECURE, "true")
|
||||
if current_platform.is_rocm():
|
||||
# The fake OTLP server starts gRPC worker threads before the engine
|
||||
# core is launched. On ROCm CI, forking while those threads are
|
||||
# active can segfault in gRPC during engine startup or teardown.
|
||||
m.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn")
|
||||
|
||||
sampling_params = SamplingParams(
|
||||
temperature=0.01,
|
||||
@@ -30,58 +37,70 @@ def test_traces(
|
||||
max_tokens=256,
|
||||
)
|
||||
model = "facebook/opt-125m"
|
||||
llm = LLM(
|
||||
model=model,
|
||||
otlp_traces_endpoint=FAKE_TRACE_SERVER_ADDRESS,
|
||||
gpu_memory_utilization=0.3,
|
||||
disable_log_stats=False,
|
||||
)
|
||||
prompts = ["This is a short prompt"]
|
||||
outputs = llm.generate(prompts, sampling_params=sampling_params)
|
||||
print(f"test_traces outputs is : {outputs}")
|
||||
llm = None
|
||||
try:
|
||||
llm = LLM(
|
||||
model=model,
|
||||
otlp_traces_endpoint=FAKE_TRACE_SERVER_ADDRESS,
|
||||
gpu_memory_utilization=0.3,
|
||||
disable_log_stats=False,
|
||||
)
|
||||
prompts = ["This is a short prompt"]
|
||||
outputs = llm.generate(prompts, sampling_params=sampling_params)
|
||||
print(f"test_traces outputs is : {outputs}")
|
||||
|
||||
# Wait for the "llm_request" span to be exported.
|
||||
# The BatchSpanProcessor batches spans and exports them periodically,
|
||||
# so we need to wait specifically for the llm_request span to appear.
|
||||
timeout = 15
|
||||
deadline = time.time() + timeout
|
||||
llm_request_spans = []
|
||||
while time.time() < deadline:
|
||||
all_spans = trace_service.get_all_spans()
|
||||
llm_request_spans = [s for s in all_spans if s["name"] == "llm_request"]
|
||||
if llm_request_spans:
|
||||
break
|
||||
time.sleep(0.5)
|
||||
# Wait for the "llm_request" span to be exported.
|
||||
# The BatchSpanProcessor batches spans and exports them periodically,
|
||||
# so we need to wait specifically for the llm_request span to appear.
|
||||
timeout = 15
|
||||
deadline = time.time() + timeout
|
||||
llm_request_spans = []
|
||||
while time.time() < deadline:
|
||||
all_spans = trace_service.get_all_spans()
|
||||
llm_request_spans = [s for s in all_spans if s["name"] == "llm_request"]
|
||||
if llm_request_spans:
|
||||
break
|
||||
time.sleep(0.5)
|
||||
|
||||
assert len(llm_request_spans) == 1, (
|
||||
f"Expected exactly 1 'llm_request' span, but got {len(llm_request_spans)}. "
|
||||
f"All span names: {[s['name'] for s in all_spans]}"
|
||||
)
|
||||
assert len(llm_request_spans) == 1, (
|
||||
f"Expected exactly 1 'llm_request' span, but got "
|
||||
f"{len(llm_request_spans)}. "
|
||||
f"All span names: {[s['name'] for s in all_spans]}"
|
||||
)
|
||||
|
||||
attributes = llm_request_spans[0]["attributes"]
|
||||
# assert attributes.get(SpanAttributes.GEN_AI_RESPONSE_MODEL) == model
|
||||
assert attributes.get(SpanAttributes.GEN_AI_REQUEST_ID) == outputs[0].request_id
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_REQUEST_TEMPERATURE)
|
||||
== sampling_params.temperature
|
||||
)
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_REQUEST_TOP_P) == sampling_params.top_p
|
||||
)
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_REQUEST_MAX_TOKENS)
|
||||
== sampling_params.max_tokens
|
||||
)
|
||||
assert attributes.get(SpanAttributes.GEN_AI_REQUEST_N) == sampling_params.n
|
||||
assert attributes.get(SpanAttributes.GEN_AI_USAGE_PROMPT_TOKENS) == len(
|
||||
outputs[0].prompt_token_ids
|
||||
)
|
||||
completion_tokens = sum(len(o.token_ids) for o in outputs[0].outputs)
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_USAGE_COMPLETION_TOKENS)
|
||||
== completion_tokens
|
||||
)
|
||||
attributes = llm_request_spans[0]["attributes"]
|
||||
# assert attributes.get(SpanAttributes.GEN_AI_RESPONSE_MODEL) == model
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_REQUEST_ID)
|
||||
== outputs[0].request_id
|
||||
)
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_REQUEST_TEMPERATURE)
|
||||
== sampling_params.temperature
|
||||
)
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_REQUEST_TOP_P)
|
||||
== sampling_params.top_p
|
||||
)
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_REQUEST_MAX_TOKENS)
|
||||
== sampling_params.max_tokens
|
||||
)
|
||||
assert attributes.get(SpanAttributes.GEN_AI_REQUEST_N) == sampling_params.n
|
||||
assert attributes.get(SpanAttributes.GEN_AI_USAGE_PROMPT_TOKENS) == len(
|
||||
outputs[0].prompt_token_ids
|
||||
)
|
||||
completion_tokens = sum(len(o.token_ids) for o in outputs[0].outputs)
|
||||
assert (
|
||||
attributes.get(SpanAttributes.GEN_AI_USAGE_COMPLETION_TOKENS)
|
||||
== completion_tokens
|
||||
)
|
||||
|
||||
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_IN_QUEUE) > 0
|
||||
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_TO_FIRST_TOKEN) > 0
|
||||
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_E2E) > 0
|
||||
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_IN_QUEUE) > 0
|
||||
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_TO_FIRST_TOKEN) > 0
|
||||
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_E2E) > 0
|
||||
finally:
|
||||
if llm is not None:
|
||||
shutdown_timeout = 60.0 if current_platform.is_rocm() else 5.0
|
||||
llm.llm_engine.engine_core.shutdown(timeout=shutdown_timeout)
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -8,11 +8,12 @@ import regex as re
|
||||
# Regex: match `torch.cuda.xxx` but allow `torch.accelerator.xxx`
|
||||
# --------------------------------------------------------------------------- #
|
||||
_TORCH_CUDA_PATTERNS = [
|
||||
r"\btorch\.cuda\.(empty_cache|synchronize|device_count|current_device|memory_reserved|memory_allocated|max_memory_allocated|max_memory_reserved|reset_peak_memory_stats|memory_stats|set_device|device\()\b",
|
||||
r"\btorch\.cuda\.(empty_cache|synchronize|device_count|current_device|memory_reserved|memory_allocated|max_memory_allocated|max_memory_reserved|reset_peak_memory_stats|memory_stats|mem_get_info|set_device|device\()\b",
|
||||
r"\btorch\.cuda\.(manual_seed|manual_seed_all)\b",
|
||||
r"\bwith\storch\.cuda\.device\b",
|
||||
# Calls torch.cuda.{_is_compiled/_device_count_amdsmi/_device_count_nvml} internally
|
||||
r"\bcuda_device_count_stateless\(\)\b",
|
||||
r"\bcurrent_platform\.mem_get_info\(\)\b",
|
||||
]
|
||||
|
||||
ALLOWED_FILES = {
|
||||
|
||||
@@ -262,10 +262,12 @@ class PassConfig:
|
||||
"Fusion enabled but reshape elimination disabled. "
|
||||
"RMSNorm + padding fusion might not work"
|
||||
)
|
||||
if self.enable_qk_norm_rope_fusion and not current_platform.is_cuda_alike():
|
||||
if self.enable_qk_norm_rope_fusion and not (
|
||||
current_platform.is_cuda_alike() or current_platform.is_xpu()
|
||||
):
|
||||
logger.warning_once(
|
||||
"QK Norm + RoPE fusion enabled but the current platform is not "
|
||||
"CUDA or ROCm. The fusion will be disabled."
|
||||
"CUDA, ROCm or XPU. The fusion will be disabled."
|
||||
)
|
||||
self.enable_qk_norm_rope_fusion = False
|
||||
if self.fuse_act_padding and not current_platform.is_rocm():
|
||||
@@ -757,6 +759,7 @@ class CompilationConfig:
|
||||
"vllm::sparse_attn_indexer",
|
||||
"vllm::rocm_aiter_sparse_attn_indexer",
|
||||
"vllm::deepseek_v4_attention",
|
||||
"vllm::hpc_rope_norm_forward",
|
||||
]
|
||||
|
||||
def compute_hash(self) -> str:
|
||||
|
||||
@@ -213,7 +213,7 @@ class ModelConfig:
|
||||
flexibility."""
|
||||
enable_return_routed_experts: bool = False
|
||||
"""Whether to return routed experts."""
|
||||
max_logprobs: int = 20
|
||||
max_logprobs: int = Field(default=20, ge=-1)
|
||||
"""Maximum number of log probabilities to return when `logprobs` is
|
||||
specified in `SamplingParams`. The default value comes the default for the
|
||||
OpenAI Chat Completions API. -1 means no cap, i.e. all (output_length *
|
||||
|
||||
@@ -77,9 +77,9 @@ class SchedulerConfig:
|
||||
this less than max_num_partial_prefills will allow shorter prompts to jump
|
||||
the queue in front of longer prompts in some cases, improving latency."""
|
||||
|
||||
long_prefill_token_threshold: int = 0
|
||||
long_prefill_token_threshold: int = Field(default=0, ge=0)
|
||||
"""For chunked prefill, a request is considered long if the prompt is
|
||||
longer than this number of tokens."""
|
||||
longer than this number of tokens. 0 disables the cap (default)."""
|
||||
|
||||
enable_chunked_prefill: bool = True
|
||||
"""If True, prefill requests can be chunked based
|
||||
|
||||
+19
-13
@@ -931,16 +931,28 @@ class VllmConfig:
|
||||
model_type,
|
||||
)
|
||||
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.v1.executor.abstract import Executor
|
||||
|
||||
executor_backend = self.parallel_config.distributed_executor_backend
|
||||
executor_class = Executor.get_class(self)
|
||||
executor_supports_async_sched = executor_class.supports_async_scheduling()
|
||||
uses_rocm_deepep_ht_dbo = (
|
||||
current_platform.is_rocm()
|
||||
and self.parallel_config.enable_dbo
|
||||
and self.parallel_config.all2all_backend == "deepep_high_throughput"
|
||||
)
|
||||
|
||||
if self.scheduler_config.async_scheduling:
|
||||
# Async scheduling explicitly enabled, hard fail any incompatibilities.
|
||||
# Currently, async scheduling only support eagle speculative
|
||||
# decoding.
|
||||
if uses_rocm_deepep_ht_dbo:
|
||||
raise ValueError(
|
||||
"Async scheduling is not compatible with ROCm DeepEP "
|
||||
"high-throughput DBO. Please use --no-async-scheduling or "
|
||||
"select a different all2all backend."
|
||||
)
|
||||
if self.speculative_config is not None:
|
||||
if (
|
||||
self.speculative_config.method not in get_args(EagleModelTypes)
|
||||
@@ -1000,6 +1012,13 @@ class VllmConfig:
|
||||
executor_backend,
|
||||
)
|
||||
self.scheduler_config.async_scheduling = False
|
||||
elif uses_rocm_deepep_ht_dbo:
|
||||
logger.warning_once(
|
||||
"Async scheduling is disabled for ROCm DeepEP "
|
||||
"high-throughput DBO because that combination can corrupt "
|
||||
"DP+EP generation accuracy."
|
||||
)
|
||||
self.scheduler_config.async_scheduling = False
|
||||
else:
|
||||
self.scheduler_config.async_scheduling = True
|
||||
|
||||
@@ -1044,8 +1063,6 @@ class VllmConfig:
|
||||
"VLLM_WORKER_MULTIPROC_METHOD set to spawn"
|
||||
)
|
||||
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
if (
|
||||
self.model_config is not None
|
||||
and self.scheduler_config.enable_chunked_prefill
|
||||
@@ -1997,13 +2014,6 @@ class VllmConfig:
|
||||
model_config = self.model_config
|
||||
speculative_config = self.speculative_config
|
||||
|
||||
if (
|
||||
model_config is not None
|
||||
and model_config.has_inner_state
|
||||
and self.cache_config.mamba_cache_mode == "align"
|
||||
):
|
||||
unsupported.append("hybrid/mamba models with align cache mode")
|
||||
|
||||
if self.parallel_config.prefill_context_parallel_size > 1:
|
||||
unsupported.append("prefill context parallelism")
|
||||
|
||||
@@ -2152,10 +2162,6 @@ class VllmConfig:
|
||||
"to schedule a multiple of block_size tokens even if they are "
|
||||
"in the middle of a mm input"
|
||||
)
|
||||
# TODO: support align mamba cache mode for model runner v2
|
||||
assert not envs.VLLM_USE_V2_MODEL_RUNNER, (
|
||||
"Model Runner V2 has not yet supported mamba_cache_mode='align'. "
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_nvfp4_kv_cache_with_mla(self) -> "VllmConfig":
|
||||
|
||||
+35
-21
@@ -95,40 +95,54 @@ def mma_bf16(
|
||||
return cute.TensorSSA(vec, 4, Float32)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
def _bf16x2_abs(a: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
def _bf16x2_unary(asm: str, a: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
out = llvm.inline_asm(
|
||||
T.i32(),
|
||||
[a.ir_value(loc=loc, ip=ip)],
|
||||
"abs.bf16x2 $0, $1;",
|
||||
f"{asm}.bf16x2 $0, $1;",
|
||||
"=r,r",
|
||||
has_side_effects=False,
|
||||
is_align_stack=False,
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
return Uint32(out)
|
||||
|
||||
|
||||
def _bf16x2_binary(asm: str, a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
out = llvm.inline_asm(
|
||||
T.i32(),
|
||||
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
|
||||
f"{asm}.bf16x2 $0, $1, $2;",
|
||||
"=r,r,r",
|
||||
has_side_effects=False,
|
||||
is_align_stack=False,
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
return Uint32(out)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
def _bf16x2_abs(a: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
return _bf16x2_unary("abs", a, loc=loc, ip=ip)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
def _bf16x2_neg(a: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
return _bf16x2_unary("neg", a, loc=loc, ip=ip)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
def _bf16x2_max(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
out = llvm.inline_asm(
|
||||
T.i32(),
|
||||
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
|
||||
"max.bf16x2 $0, $1, $2;",
|
||||
"=r,r,r",
|
||||
has_side_effects=False,
|
||||
is_align_stack=False,
|
||||
)
|
||||
return Uint32(out)
|
||||
return _bf16x2_binary("max", a, b, loc=loc, ip=ip)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
def _bf16x2_mul(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
out = llvm.inline_asm(
|
||||
T.i32(),
|
||||
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
|
||||
"mul.rn.bf16x2 $0, $1, $2;",
|
||||
"=r,r,r",
|
||||
has_side_effects=False,
|
||||
is_align_stack=False,
|
||||
)
|
||||
return Uint32(out)
|
||||
return _bf16x2_binary("mul.rn", a, b, loc=loc, ip=ip)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
def _bf16x2_sub(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
|
||||
return _bf16x2_binary("sub.rn", a, b, loc=loc, ip=ip)
|
||||
|
||||
+34
-30
@@ -90,17 +90,18 @@ def mma_f16(
|
||||
loc=None,
|
||||
ip=None,
|
||||
) -> None:
|
||||
nvvm.tcgen05_mma(
|
||||
nvvm.Tcgen05MMAKind.F16,
|
||||
NVVM_CTA_GROUP_MAP[cta_group],
|
||||
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
|
||||
Uint64(a_desc).ir_value(loc=loc, ip=ip),
|
||||
Uint64(b_desc).ir_value(loc=loc, ip=ip),
|
||||
Int32(idesc).ir_value(loc=loc, ip=ip),
|
||||
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
with cute.arch.elect_one():
|
||||
nvvm.tcgen05_mma(
|
||||
nvvm.Tcgen05MMAKind.F16,
|
||||
NVVM_CTA_GROUP_MAP[cta_group],
|
||||
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
|
||||
Uint64(a_desc).ir_value(loc=loc, ip=ip),
|
||||
Uint64(b_desc).ir_value(loc=loc, ip=ip),
|
||||
Int32(idesc).ir_value(loc=loc, ip=ip),
|
||||
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
@@ -115,17 +116,18 @@ def mma_ts_f16(
|
||||
loc=None,
|
||||
ip=None,
|
||||
) -> None:
|
||||
nvvm.tcgen05_mma(
|
||||
nvvm.Tcgen05MMAKind.F16,
|
||||
NVVM_CTA_GROUP_MAP[cta_group],
|
||||
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
|
||||
_make_tmem_llvm_ptr(a_tmem, loc=loc, ip=ip),
|
||||
Uint64(b_desc).ir_value(loc=loc, ip=ip),
|
||||
Int32(idesc).ir_value(loc=loc, ip=ip),
|
||||
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
with cute.arch.elect_one():
|
||||
nvvm.tcgen05_mma(
|
||||
nvvm.Tcgen05MMAKind.F16,
|
||||
NVVM_CTA_GROUP_MAP[cta_group],
|
||||
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
|
||||
_make_tmem_llvm_ptr(a_tmem, loc=loc, ip=ip),
|
||||
Uint64(b_desc).ir_value(loc=loc, ip=ip),
|
||||
Int32(idesc).ir_value(loc=loc, ip=ip),
|
||||
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
@@ -133,15 +135,17 @@ def commit(mbar, cta_mask=None, cta_group: int = 1, *, loc=None, ip=None):
|
||||
mbar_llvm = mbar.to_llvm_ptr(loc=loc, ip=ip)
|
||||
group = NVVM_CTA_GROUP_MAP[cta_group]
|
||||
if cutlass.const_expr(cta_mask is not None):
|
||||
nvvm.tcgen05_commit_arrive(
|
||||
mbar_llvm,
|
||||
multicast_mask=cta_mask.ir_value(loc=loc, ip=ip),
|
||||
group=group,
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
with cute.arch.elect_one():
|
||||
nvvm.tcgen05_commit_arrive(
|
||||
mbar_llvm,
|
||||
multicast_mask=cta_mask.ir_value(loc=loc, ip=ip),
|
||||
group=group,
|
||||
loc=loc,
|
||||
ip=ip,
|
||||
)
|
||||
else:
|
||||
nvvm.tcgen05_commit_arrive(mbar_llvm, group=group, loc=loc, ip=ip)
|
||||
with cute.arch.elect_one():
|
||||
nvvm.tcgen05_commit_arrive(mbar_llvm, group=group, loc=loc, ip=ip)
|
||||
|
||||
|
||||
@dsl_user_op
|
||||
|
||||
@@ -22,7 +22,6 @@ from vllm.config import VllmConfig
|
||||
from vllm.distributed.kv_transfer.kv_connector.utils import (
|
||||
EngineId,
|
||||
TransferTopology,
|
||||
get_current_attn_backend,
|
||||
get_current_attn_backends,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
|
||||
@@ -51,10 +50,18 @@ from vllm.platforms import current_platform
|
||||
from vllm.utils.math_utils import cdiv
|
||||
from vllm.utils.network_utils import get_ip, make_zmq_path, make_zmq_socket
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.utils import get_kv_cache_layout
|
||||
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID, get_kv_cache_layout
|
||||
from vllm.v1.core.sched.output import SchedulerOutput
|
||||
from vllm.v1.kv_cache_interface import FullAttentionSpec, SlidingWindowSpec
|
||||
from vllm.v1.kv_cache_interface import (
|
||||
FullAttentionSpec,
|
||||
KVCacheSpec,
|
||||
MambaSpec,
|
||||
MLAAttentionSpec,
|
||||
SlidingWindowMLASpec,
|
||||
SlidingWindowSpec,
|
||||
)
|
||||
from vllm.v1.request import RequestStatus
|
||||
from vllm.v1.worker.block_table import BlockTable
|
||||
from vllm.v1.worker.utils import select_common_block_size
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -85,6 +92,7 @@ class TransferRegion:
|
||||
base_addr: int
|
||||
block_len: int
|
||||
kv_block_len: int
|
||||
group_index: int = 0
|
||||
|
||||
|
||||
def _get_tp_ratio(local_tp_size: int, remote_tp_size: int) -> int:
|
||||
@@ -111,24 +119,58 @@ def _get_tp_ratio(local_tp_size: int, remote_tp_size: int) -> int:
|
||||
def _expand_transfer_regions(
|
||||
base_addrs: list[int],
|
||||
block_lens: list[int],
|
||||
kv_block_lens: list[int],
|
||||
layer_names: list[str],
|
||||
layer_indices: list[int],
|
||||
is_kv_layout_blocks_first: bool,
|
||||
group_indices: list[int] | None = None,
|
||||
split_kv_regions: list[bool] | None = None,
|
||||
) -> list[TransferRegion]:
|
||||
"""Expand registered KV tensors into the regions transferred by Mooncake."""
|
||||
assert (
|
||||
len(base_addrs) == len(block_lens) == len(layer_names) == len(layer_indices)
|
||||
len(base_addrs)
|
||||
== len(block_lens)
|
||||
== len(kv_block_lens)
|
||||
== len(layer_names)
|
||||
== len(layer_indices)
|
||||
), (
|
||||
"Mooncake transfer regions require matching metadata lengths, got "
|
||||
f"base_addrs={len(base_addrs)}, block_lens={len(block_lens)}, "
|
||||
f"kv_block_lens={len(kv_block_lens)}, "
|
||||
f"layer_names={len(layer_names)}, "
|
||||
f"layer_indices={len(layer_indices)}."
|
||||
)
|
||||
if group_indices is None:
|
||||
group_indices = [0] * len(layer_names)
|
||||
assert len(group_indices) == len(layer_names), (
|
||||
"Mooncake transfer regions require matching group metadata lengths, "
|
||||
f"got group_indices={len(group_indices)}, layer_names={len(layer_names)}."
|
||||
)
|
||||
if split_kv_regions is None:
|
||||
split_kv_regions = [is_kv_layout_blocks_first] * len(layer_names)
|
||||
assert len(split_kv_regions) == len(layer_names), (
|
||||
"Mooncake transfer regions require matching split metadata, "
|
||||
f"got split_kv_regions={len(split_kv_regions)}, "
|
||||
f"layer_names={len(layer_names)}."
|
||||
)
|
||||
regions: list[TransferRegion] = []
|
||||
for base_addr, block_len, layer_name, layer_index in zip(
|
||||
base_addrs, block_lens, layer_names, layer_indices
|
||||
for (
|
||||
base_addr,
|
||||
block_len,
|
||||
kv_block_len,
|
||||
layer_name,
|
||||
layer_index,
|
||||
group_index,
|
||||
split_kv_region,
|
||||
) in zip(
|
||||
base_addrs,
|
||||
block_lens,
|
||||
kv_block_lens,
|
||||
layer_names,
|
||||
layer_indices,
|
||||
group_indices,
|
||||
split_kv_regions,
|
||||
):
|
||||
kv_block_len = block_len // 2 if is_kv_layout_blocks_first else block_len
|
||||
regions.append(
|
||||
TransferRegion(
|
||||
layer_name=layer_name,
|
||||
@@ -136,9 +178,10 @@ def _expand_transfer_regions(
|
||||
base_addr=base_addr,
|
||||
block_len=block_len,
|
||||
kv_block_len=kv_block_len,
|
||||
group_index=group_index,
|
||||
)
|
||||
)
|
||||
if is_kv_layout_blocks_first:
|
||||
if split_kv_region:
|
||||
regions.append(
|
||||
TransferRegion(
|
||||
layer_name=layer_name,
|
||||
@@ -146,6 +189,7 @@ def _expand_transfer_regions(
|
||||
base_addr=base_addr + kv_block_len,
|
||||
block_len=block_len,
|
||||
kv_block_len=kv_block_len,
|
||||
group_index=group_index,
|
||||
)
|
||||
)
|
||||
return regions
|
||||
@@ -308,6 +352,17 @@ def _align_transfer_regions(
|
||||
f"{remote_region.layer_index}."
|
||||
),
|
||||
)
|
||||
if local_region.group_index != remote_region.group_index:
|
||||
return (
|
||||
[],
|
||||
[],
|
||||
(
|
||||
"Mooncake registered group index mismatch for "
|
||||
f"{local_region.layer_name}: producer="
|
||||
f"{local_region.group_index}, consumer="
|
||||
f"{remote_region.group_index}."
|
||||
),
|
||||
)
|
||||
aligned_local.append(local_region)
|
||||
aligned_remote.append(remote_region)
|
||||
|
||||
@@ -332,8 +387,10 @@ class MooncakeXferMetadata(
|
||||
req_blocks: dict[ReqId, tuple[TransferId, list[list[int]]]]
|
||||
kv_caches_base_addr: list[int]
|
||||
block_lens: list[int]
|
||||
kv_block_lens: list[int]
|
||||
registered_layer_names: list[str] = msgspec.field(default_factory=list)
|
||||
registered_layer_indices: list[int] = msgspec.field(default_factory=list)
|
||||
registered_group_indices: list[int] = msgspec.field(default_factory=list)
|
||||
|
||||
|
||||
class MooncakeXferResponseStatus(IntEnum):
|
||||
@@ -581,6 +638,9 @@ class MooncakeConnectorScheduler:
|
||||
for g in kv_cache_config.kv_cache_groups
|
||||
)
|
||||
)
|
||||
# GDN is represented as a MambaSpec in vLLM. This Mooncake MambaSpec
|
||||
# path is currently tested with GDN; Mamba2 is not validated yet.
|
||||
self._has_mamba = kv_cache_config.has_mamba_layers
|
||||
|
||||
# Requests that need to start recv/send.
|
||||
# New requests are added by update_state_after_alloc in
|
||||
@@ -617,6 +677,38 @@ class MooncakeConnectorScheduler:
|
||||
for i, blocks in enumerate(block_ids)
|
||||
]
|
||||
|
||||
def _get_remote_prefill_token_count(self, num_prompt_tokens: int) -> int:
|
||||
"""D-side only. Returns N-1 for Mamba models since the decoder
|
||||
always recomputes the last token and must start from h(N-1)."""
|
||||
if self._has_mamba and num_prompt_tokens > 1:
|
||||
return num_prompt_tokens - 1
|
||||
return num_prompt_tokens
|
||||
|
||||
def _truncate_mamba_request_for_prefill(self, request: "Request") -> None:
|
||||
"""P-side only: drop the last prompt token so the prefiller computes
|
||||
h(N-1) instead of h(N). The decoder recomputes the last token to
|
||||
derive h(N) correctly.
|
||||
|
||||
Guarded by ``_p_side_truncated`` to avoid repeated truncation if the
|
||||
request is preempted and rescheduled."""
|
||||
params = request.kv_transfer_params
|
||||
if (
|
||||
params is not None
|
||||
and not params.get("_p_side_truncated")
|
||||
and request.num_prompt_tokens > 1
|
||||
):
|
||||
if request.prompt_token_ids is not None:
|
||||
request.prompt_token_ids.pop()
|
||||
elif request.prompt_embeds is not None:
|
||||
request.prompt_embeds = request.prompt_embeds[:-1]
|
||||
else:
|
||||
return
|
||||
|
||||
request._all_token_ids.pop()
|
||||
request.num_prompt_tokens -= 1
|
||||
request.max_tokens = 1
|
||||
params["_p_side_truncated"] = True
|
||||
|
||||
def get_num_new_matched_tokens(
|
||||
self, request: "Request", num_computed_tokens: int
|
||||
) -> tuple[int, bool]:
|
||||
@@ -650,10 +742,15 @@ class MooncakeConnectorScheduler:
|
||||
# Remote prefill: get all prompt blocks from remote.
|
||||
assert not self.is_kv_producer
|
||||
token_ids = request.prompt_token_ids or []
|
||||
count = len(token_ids) - num_computed_tokens
|
||||
count = self._get_remote_prefill_token_count(len(token_ids)) - (
|
||||
num_computed_tokens
|
||||
)
|
||||
if count > 0:
|
||||
return count, True
|
||||
|
||||
if params.get("do_remote_decode") and self._has_mamba:
|
||||
self._truncate_mamba_request_for_prefill(request)
|
||||
|
||||
# No remote prefill for this request.
|
||||
return 0, False
|
||||
|
||||
@@ -802,7 +899,7 @@ class MooncakeConnectorWorker:
|
||||
self,
|
||||
vllm_config: VllmConfig,
|
||||
engine_id: str,
|
||||
kv_cache_config: "KVCacheConfig | None" = None,
|
||||
kv_cache_config: "KVCacheConfig",
|
||||
):
|
||||
if TransferEngine is None:
|
||||
logger.error("Mooncake is not available")
|
||||
@@ -831,10 +928,15 @@ class MooncakeConnectorWorker:
|
||||
protocol = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr]
|
||||
"mooncake_protocol", "rdma"
|
||||
)
|
||||
device_name = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr]
|
||||
"device_name", ""
|
||||
)
|
||||
logger.info(
|
||||
"The Mooncake Transfer Engine is using %s as its protocol.", protocol
|
||||
)
|
||||
ret_value = self.engine.initialize(self.hostname, "P2PHANDSHAKE", protocol, "")
|
||||
ret_value = self.engine.initialize(
|
||||
self.hostname, "P2PHANDSHAKE", protocol, device_name
|
||||
)
|
||||
if ret_value != 0:
|
||||
raise RuntimeError("Mooncake Transfer Engine initialization failed.")
|
||||
|
||||
@@ -852,10 +954,11 @@ class MooncakeConnectorWorker:
|
||||
self.engine_id: EngineId = engine_id
|
||||
self.tp_rank = get_tensor_model_parallel_rank()
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.num_blocks = 0
|
||||
self.block_len_per_layer: list[int] = []
|
||||
self.kv_block_len_per_layer: list[int] = []
|
||||
self.registered_layer_names: list[str] = []
|
||||
self.registered_layer_indices: list[int] = []
|
||||
self.registered_group_indices: list[int] = []
|
||||
self.seen_base_addresses: list[int] = []
|
||||
|
||||
assert (parallel_config := vllm_config.parallel_config)
|
||||
@@ -916,26 +1019,40 @@ class MooncakeConnectorWorker:
|
||||
self.cache_config = vllm_config.cache_config
|
||||
self.kv_cache_config = kv_cache_config
|
||||
self.use_mla = self.model_config.use_mla
|
||||
self._physical_blocks_per_logical_kv_block = 1
|
||||
self._sync_block_size_with_kernel()
|
||||
|
||||
# Get the attention backend from the first layer
|
||||
# NOTE (NickLucche) models with multiple backends are not supported yet
|
||||
backend = get_current_attn_backend(vllm_config)
|
||||
self.backend_name = backend.get_name()
|
||||
self.attn_backends = get_current_attn_backends(vllm_config)
|
||||
self.kv_cache_layout = get_kv_cache_layout()
|
||||
logger.debug("Detected attention backend %s", self.backend_name)
|
||||
logger.debug(
|
||||
"Detected attention backends %s",
|
||||
[backend.get_name() for backend in self.attn_backends],
|
||||
)
|
||||
logger.debug("Detected kv cache layout %s", self.kv_cache_layout)
|
||||
|
||||
self._tp_size: dict[EngineId, int] = {self.engine_id: self.tp_size}
|
||||
self._layer_specs: dict[str, KVCacheSpec] = {}
|
||||
for group in kv_cache_config.kv_cache_groups:
|
||||
group_spec = group.kv_cache_spec
|
||||
specs_by_layer = getattr(group_spec, "kv_cache_specs", {})
|
||||
for layer_name in group.layer_names:
|
||||
self._layer_specs[layer_name] = specs_by_layer.get(
|
||||
layer_name, group_spec
|
||||
)
|
||||
self._layer_group_indices: dict[str, int] = {
|
||||
layer: group_index
|
||||
for group_index, group in enumerate(kv_cache_config.kv_cache_groups)
|
||||
for layer in group.layer_names
|
||||
}
|
||||
self.transfer_topo = TransferTopology(
|
||||
tp_rank=self.tp_rank,
|
||||
tp_size=self.tp_size,
|
||||
block_size=self.block_size,
|
||||
engine_id=self.engine_id,
|
||||
is_mla=self.use_mla,
|
||||
is_mamba=False,
|
||||
is_mamba=kv_cache_config.has_mamba_layers,
|
||||
total_num_kv_heads=self.model_config.get_total_num_kv_heads(),
|
||||
attn_backends=[backend],
|
||||
attn_backends=self.attn_backends,
|
||||
)
|
||||
|
||||
self.async_zmq_ctx = zmq.asyncio.Context()
|
||||
@@ -958,6 +1075,9 @@ class MooncakeConnectorWorker:
|
||||
kernel_block_size,
|
||||
)
|
||||
assert self.block_size > kernel_block_size
|
||||
self._physical_blocks_per_logical_kv_block = (
|
||||
self.block_size // kernel_block_size
|
||||
)
|
||||
self.block_size = kernel_block_size
|
||||
|
||||
def __del__(self):
|
||||
@@ -1092,14 +1212,18 @@ class MooncakeConnectorWorker:
|
||||
local_regions = self._get_transfer_regions(
|
||||
self.kv_caches_base_addr,
|
||||
self.block_len_per_layer,
|
||||
self.kv_block_len_per_layer,
|
||||
self.registered_layer_names,
|
||||
self.registered_layer_indices,
|
||||
self.registered_group_indices,
|
||||
)
|
||||
remote_regions = self._get_transfer_regions(
|
||||
meta.kv_caches_base_addr,
|
||||
meta.block_lens,
|
||||
meta.kv_block_lens,
|
||||
meta.registered_layer_names,
|
||||
meta.registered_layer_indices,
|
||||
meta.registered_group_indices,
|
||||
)
|
||||
local_regions, remote_regions, align_err = _align_transfer_regions(
|
||||
local_regions, remote_regions
|
||||
@@ -1271,6 +1395,32 @@ class MooncakeConnectorWorker:
|
||||
remote_tp_ranks,
|
||||
)
|
||||
|
||||
def _logical_to_kernel_block_ids(
|
||||
self, block_ids: list[list[int]]
|
||||
) -> list[list[int]]:
|
||||
# For example, if a 544-token logical block is served by 32-token
|
||||
# FA kernel blocks, FA block id k expands to [17k, ..., 17k + 16],
|
||||
# while the matching Mamba/GDN state block remains k. Only attention
|
||||
# groups need logical block ids expanded to kernel block ids; Mamba/GDN
|
||||
# state block ids stay in the logical/page-id space.
|
||||
if self._physical_blocks_per_logical_kv_block == 1:
|
||||
return block_ids
|
||||
|
||||
block_arange = np.arange(self._physical_blocks_per_logical_kv_block).reshape(
|
||||
1, -1
|
||||
)
|
||||
group_specs = self.kv_cache_config.kv_cache_groups
|
||||
return [
|
||||
BlockTable.map_to_kernel_blocks(
|
||||
np.array(group),
|
||||
self._physical_blocks_per_logical_kv_block,
|
||||
block_arange,
|
||||
).tolist()
|
||||
if not isinstance(group_specs[i].kv_cache_spec, MambaSpec)
|
||||
else group
|
||||
for i, group in enumerate(block_ids)
|
||||
]
|
||||
|
||||
async def _build_transfer_params(
|
||||
self,
|
||||
ready_reqs: list[tuple[ReqId, SendBlockMeta]],
|
||||
@@ -1293,14 +1443,6 @@ class MooncakeConnectorWorker:
|
||||
):
|
||||
continue
|
||||
|
||||
# Per-group partial hit trimming, then flatten.
|
||||
# With HMA, groups share the same KV tensor but use different
|
||||
# block ranges. We trim and concatenate so the coalescer and
|
||||
# address math see one flat block list — same as non-HMA, but
|
||||
# now including blocks from every group.
|
||||
local_block_ids: list[int] = []
|
||||
remote_block_ids: list[int] = []
|
||||
has_block_error = False
|
||||
if len(send_meta.local_block_ids) != len(remote_block_ids_per_group):
|
||||
logger.error(
|
||||
"req %s: KV group count mismatch: local=%d, remote=%d",
|
||||
@@ -1312,26 +1454,55 @@ class MooncakeConnectorWorker:
|
||||
if err_msg is None:
|
||||
err_msg = "KV group count mismatch"
|
||||
continue
|
||||
for local_group, remote_group in zip(
|
||||
send_meta.local_block_ids, remote_block_ids_per_group
|
||||
|
||||
# Keep KV-cache group identity. Hybrid/HMA groups can carry
|
||||
# different semantics (e.g. full-attention KV pages vs GDN/Mamba
|
||||
# inner-state slots), so their block IDs must not be flattened and
|
||||
# reused for every registered region.
|
||||
local_block_ids_by_group: list[list[int]] = []
|
||||
remote_block_ids_by_group: list[list[int]] = []
|
||||
has_block_error = False
|
||||
group_specs = self.kv_cache_config.kv_cache_groups
|
||||
for group_index, (local_group, remote_group) in enumerate(
|
||||
zip(send_meta.local_block_ids, remote_block_ids_per_group)
|
||||
):
|
||||
is_mamba_group = isinstance(
|
||||
group_specs[group_index].kv_cache_spec,
|
||||
MambaSpec,
|
||||
)
|
||||
if is_mamba_group:
|
||||
# Mamba/GDN prefix caching can use null blocks only as
|
||||
# align-mode placeholders. They do not carry transferable
|
||||
# state, so skip them on both producer and consumer sides.
|
||||
local_group = [
|
||||
block_id
|
||||
for block_id in local_group
|
||||
if block_id != NULL_BLOCK_ID
|
||||
]
|
||||
remote_group = [
|
||||
block_id
|
||||
for block_id in remote_group
|
||||
if block_id != NULL_BLOCK_ID
|
||||
]
|
||||
|
||||
n_local = len(local_group)
|
||||
n_remote = len(remote_group)
|
||||
if n_local < n_remote:
|
||||
logger.error(
|
||||
"req %s: local blocks(%d) < remote blocks(%d) "
|
||||
"in a KV cache group",
|
||||
"in a KV cache group (is_mamba_group=%s)",
|
||||
d_req_id,
|
||||
n_local,
|
||||
n_remote,
|
||||
is_mamba_group,
|
||||
)
|
||||
has_block_error = True
|
||||
break
|
||||
if n_local > n_remote:
|
||||
elif n_local > n_remote:
|
||||
# Partial prefix cache hit: just read uncomputed blocks.
|
||||
local_group = local_group[-n_remote:]
|
||||
local_block_ids.extend(local_group)
|
||||
remote_block_ids.extend(remote_group)
|
||||
local_group = local_group[-n_remote:] if n_remote > 0 else []
|
||||
local_block_ids_by_group.append(local_group)
|
||||
remote_block_ids_by_group.append(remote_group)
|
||||
|
||||
if has_block_error:
|
||||
err_reqs.append(d_req_id)
|
||||
@@ -1339,22 +1510,44 @@ class MooncakeConnectorWorker:
|
||||
err_msg = "P num blocks less than D"
|
||||
continue
|
||||
|
||||
if not local_block_ids:
|
||||
if not any(local_block_ids_by_group):
|
||||
continue
|
||||
|
||||
# Group by indices
|
||||
group_local_block_ids, group_remote_block_ids = group_concurrent_contiguous(
|
||||
local_block_ids, remote_block_ids
|
||||
local_block_ids_by_group = self._logical_to_kernel_block_ids(
|
||||
local_block_ids_by_group
|
||||
)
|
||||
remote_block_ids_by_group = self._logical_to_kernel_block_ids(
|
||||
remote_block_ids_by_group
|
||||
)
|
||||
|
||||
for local_region, remote_region in zip(local_regions, remote_regions):
|
||||
should_transfer, src_region_offset, dst_region_offset, transfer_len = (
|
||||
self._get_sender_transfer_plan(
|
||||
local_kv_block_len=local_region.kv_block_len,
|
||||
remote_kv_block_len=remote_region.kv_block_len,
|
||||
remote_tp_rank=agent_meta.remote_tp_rank,
|
||||
remote_tp_size=agent_meta.remote_tp_size,
|
||||
)
|
||||
assert local_region.group_index == remote_region.group_index, (
|
||||
"Aligned Mooncake transfer regions must belong to the same "
|
||||
"KV group."
|
||||
)
|
||||
group_index = local_region.group_index
|
||||
assert group_index < len(local_block_ids_by_group), (
|
||||
"Transfer region references a missing KV group."
|
||||
)
|
||||
local_block_ids = local_block_ids_by_group[group_index]
|
||||
remote_block_ids = remote_block_ids_by_group[group_index]
|
||||
if not local_block_ids:
|
||||
continue
|
||||
|
||||
# Group by indices within this region's KV-cache group only.
|
||||
group_local_block_ids, group_remote_block_ids = (
|
||||
group_concurrent_contiguous(local_block_ids, remote_block_ids)
|
||||
)
|
||||
(
|
||||
should_transfer,
|
||||
src_region_offset,
|
||||
dst_region_offset,
|
||||
transfer_len,
|
||||
) = self._get_sender_transfer_plan(
|
||||
local_kv_block_len=local_region.kv_block_len,
|
||||
remote_kv_block_len=remote_region.kv_block_len,
|
||||
remote_tp_rank=agent_meta.remote_tp_rank,
|
||||
remote_tp_size=agent_meta.remote_tp_size,
|
||||
)
|
||||
if not should_transfer:
|
||||
# Replicated KV cache: only one producer rank in the TP group
|
||||
@@ -1368,7 +1561,7 @@ class MooncakeConnectorWorker:
|
||||
"Computed source transfer region exceeds local KV block size."
|
||||
)
|
||||
assert dst_region_offset + transfer_len <= remote_region.kv_block_len, (
|
||||
"Computed destination transfer region exceeds remote KV block size."
|
||||
"Destination transfer region exceeds remote KV block size."
|
||||
)
|
||||
# Collapse one contiguous block group into a single larger
|
||||
# transfer descriptor when the per-block copy is identical.
|
||||
@@ -1411,28 +1604,10 @@ class MooncakeConnectorWorker:
|
||||
)
|
||||
lengths.append(transfer_len)
|
||||
|
||||
if local_region is local_regions[0]:
|
||||
logger.debug(
|
||||
"Mooncake transfer plan for request %s: local_tp=%d "
|
||||
"remote_tp=%d remote_tp_rank=%d local_block_len=%d "
|
||||
"remote_block_len=%d src_offset=%d dst_offset=%d "
|
||||
"transfer_len=%d coalesce=%s",
|
||||
d_req_id,
|
||||
self.tp_size,
|
||||
agent_meta.remote_tp_size,
|
||||
agent_meta.remote_tp_rank,
|
||||
local_region.block_len,
|
||||
remote_region.block_len,
|
||||
src_region_offset,
|
||||
dst_region_offset,
|
||||
transfer_len,
|
||||
can_coalesce,
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"Sending kv_caches for request %s (%d blocks) to %s",
|
||||
d_req_id,
|
||||
len(local_block_ids),
|
||||
sum(len(group) for group in local_block_ids_by_group),
|
||||
remote_session,
|
||||
)
|
||||
|
||||
@@ -1480,18 +1655,33 @@ class MooncakeConnectorWorker:
|
||||
|
||||
logger.info("Registering KV_Caches. use_mla: %s", self.use_mla)
|
||||
|
||||
kv_data_ptrs = []
|
||||
kv_data_lens = []
|
||||
seen_base_addresses = []
|
||||
kv_data_ptrs: list[int] = []
|
||||
kv_data_lens: list[int] = []
|
||||
region_base_addresses: list[int] = []
|
||||
seen_storage_ptrs: set[int] = set()
|
||||
self.block_len_per_layer = []
|
||||
self.kv_block_len_per_layer = []
|
||||
self.registered_layer_names = []
|
||||
self.registered_layer_indices = []
|
||||
self.registered_group_indices = []
|
||||
|
||||
split_k_and_v = self.transfer_topo.split_k_and_v
|
||||
tensor_size_bytes = None
|
||||
for layer_name, cache_or_caches in kv_caches.items():
|
||||
layer_index = extract_layer_index(layer_name)
|
||||
cache_list = cache_or_caches if split_k_and_v else [cache_or_caches]
|
||||
layer_spec = self._layer_specs.get(layer_name)
|
||||
if layer_spec is None:
|
||||
logger.debug(
|
||||
"Skipping layer %s because no KV cache spec is present.",
|
||||
layer_name,
|
||||
)
|
||||
continue
|
||||
if isinstance(layer_spec, MambaSpec):
|
||||
conv, _ = cache_or_caches
|
||||
cache_list = [conv]
|
||||
else:
|
||||
cache_list = self.transfer_topo.get_transfer_cache_regions(
|
||||
cache_or_caches, layer_spec
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"registering layer %s with %d cache tensor(s)",
|
||||
layer_name,
|
||||
@@ -1501,45 +1691,46 @@ class MooncakeConnectorWorker:
|
||||
for cache in cache_list:
|
||||
self._log_debug_cache_registration(layer_name, cache)
|
||||
base_addr = cache.data_ptr()
|
||||
if base_addr in seen_base_addresses:
|
||||
continue
|
||||
|
||||
seen_base_addresses.append(base_addr)
|
||||
|
||||
if tensor_size_bytes is None:
|
||||
tensor_size_bytes = cache.nbytes
|
||||
self.num_blocks = cache.shape[0]
|
||||
assert cache.shape[0] == self.num_blocks, (
|
||||
"All kv cache tensors must have the same number of blocks"
|
||||
)
|
||||
|
||||
# Use stride-based block length so RDMA reaches the last
|
||||
# block's padding (e.g. DeepseekV4 MLA alignment). stride(0)
|
||||
# reflects the actual byte distance between consecutive
|
||||
# blocks in GPU memory, which matches or exceeds the
|
||||
# shape-based size.
|
||||
block_len = cache.stride(0) * cache.element_size()
|
||||
region_base_addresses.append(base_addr)
|
||||
|
||||
if isinstance(layer_spec, (MLAAttentionSpec, SlidingWindowMLASpec)):
|
||||
kv_block_len = layer_spec.page_size_bytes
|
||||
elif self.transfer_topo.virtually_split_kv_in_blocks and not isinstance(
|
||||
layer_spec, MambaSpec
|
||||
):
|
||||
kv_block_len = block_len // 2
|
||||
else:
|
||||
kv_block_len = block_len
|
||||
self.block_len_per_layer.append(block_len)
|
||||
self.kv_block_len_per_layer.append(kv_block_len)
|
||||
self.registered_layer_names.append(layer_name)
|
||||
self.registered_layer_indices.append(layer_index)
|
||||
kv_data_ptrs.append(base_addr)
|
||||
kv_data_lens.append(self.num_blocks * block_len)
|
||||
self.registered_group_indices.append(
|
||||
self._layer_group_indices[layer_name]
|
||||
)
|
||||
storage = cache.untyped_storage()
|
||||
storage_addr = storage.data_ptr()
|
||||
if storage_addr not in seen_storage_ptrs:
|
||||
seen_storage_ptrs.add(storage_addr)
|
||||
kv_data_ptrs.append(storage_addr)
|
||||
kv_data_lens.append(storage.nbytes())
|
||||
|
||||
self.kv_caches_base_addr = seen_base_addresses
|
||||
self.seen_base_addresses = seen_base_addresses
|
||||
self.kv_caches_base_addr = region_base_addresses
|
||||
self.seen_base_addresses = kv_data_ptrs
|
||||
|
||||
if not kv_data_ptrs:
|
||||
raise RuntimeError("No KV cache tensors were registered with Mooncake.")
|
||||
|
||||
ret_value = self.engine.batch_register_memory(kv_data_ptrs, kv_data_lens)
|
||||
if ret_value != 0:
|
||||
raise RuntimeError("Mooncake batch memory registration failed.")
|
||||
|
||||
assert tensor_size_bytes is not None
|
||||
assert self.num_blocks != 0
|
||||
self.device_kv_caches = kv_caches
|
||||
logger.debug(
|
||||
"registered num_blocks=%d block_lens=%s",
|
||||
self.num_blocks,
|
||||
"registered block_lens=%s kv_block_lens=%s",
|
||||
self.block_len_per_layer,
|
||||
self.kv_block_len_per_layer,
|
||||
)
|
||||
|
||||
# No need to launch server for D node.
|
||||
@@ -1642,8 +1833,10 @@ class MooncakeConnectorWorker:
|
||||
},
|
||||
kv_caches_base_addr=self.kv_caches_base_addr,
|
||||
block_lens=self.block_len_per_layer,
|
||||
kv_block_lens=self.kv_block_len_per_layer,
|
||||
registered_layer_names=self.registered_layer_names,
|
||||
registered_layer_indices=self.registered_layer_indices,
|
||||
registered_group_indices=self.registered_group_indices,
|
||||
)
|
||||
|
||||
encoded_data = self._encoder.encode(metadata)
|
||||
@@ -1852,15 +2045,34 @@ class MooncakeConnectorWorker:
|
||||
self,
|
||||
base_addrs: list[int],
|
||||
block_lens: list[int],
|
||||
kv_block_lens: list[int],
|
||||
layer_names: list[str],
|
||||
layer_indices: list[int],
|
||||
group_indices: list[int] | None = None,
|
||||
) -> list[TransferRegion]:
|
||||
if not group_indices:
|
||||
group_indices = [
|
||||
self._layer_group_indices.get(layer_name, 0)
|
||||
for layer_name in layer_names
|
||||
]
|
||||
split_kv_regions = None
|
||||
if self.transfer_topo.virtually_split_kv_in_blocks:
|
||||
split_kv_regions = [
|
||||
not isinstance(
|
||||
self._layer_specs[layer_name],
|
||||
(MambaSpec, MLAAttentionSpec, SlidingWindowMLASpec),
|
||||
)
|
||||
for layer_name in layer_names
|
||||
]
|
||||
return _expand_transfer_regions(
|
||||
base_addrs=base_addrs,
|
||||
block_lens=block_lens,
|
||||
kv_block_lens=kv_block_lens,
|
||||
layer_names=layer_names,
|
||||
layer_indices=layer_indices,
|
||||
is_kv_layout_blocks_first=self.transfer_topo.virtually_split_kv_in_blocks,
|
||||
group_indices=group_indices,
|
||||
split_kv_regions=split_kv_regions,
|
||||
)
|
||||
|
||||
def _get_sender_transfer_plan(
|
||||
|
||||
@@ -492,12 +492,15 @@ class MultiConnector(KVConnectorBase_V1, SupportsHMA):
|
||||
async_saves += 1
|
||||
if txfer_params is not None:
|
||||
if kv_txfer_params is not None:
|
||||
# TODO we can probably change this to merge the dicts here,
|
||||
# checking for key clashes.
|
||||
raise RuntimeError(
|
||||
"Only one connector can produce KV transfer params"
|
||||
)
|
||||
kv_txfer_params = txfer_params
|
||||
clashes = set(kv_txfer_params) & set(txfer_params)
|
||||
if clashes:
|
||||
raise RuntimeError(
|
||||
"Key clash in kv_transfer_params from multiple "
|
||||
f"connectors: {clashes}"
|
||||
)
|
||||
kv_txfer_params.update(txfer_params)
|
||||
else:
|
||||
kv_txfer_params = txfer_params
|
||||
if async_saves > 1:
|
||||
self._extra_async_saves[request.request_id] = async_saves - 1
|
||||
|
||||
|
||||
@@ -320,7 +320,7 @@ class NixlBaseConnectorScheduler:
|
||||
logger.warning("Connection listener got unexpected message %s", msg)
|
||||
sock.send_multipart((identity, b"", encoded_data[target_tp_rank]))
|
||||
|
||||
def _mamba_prefill_token_count(self, num_prompt_tokens: int) -> int:
|
||||
def _get_remote_prefill_token_count(self, num_prompt_tokens: int) -> int:
|
||||
"""D-side only. Returns N-1 for Mamba models since the decoder
|
||||
always recomputes the last token and must start from h(N-1)."""
|
||||
if self._has_mamba and num_prompt_tokens > 1:
|
||||
|
||||
@@ -60,7 +60,7 @@ class NixlPullConnectorScheduler(NixlBaseConnectorScheduler):
|
||||
if params is not None and params.get("do_remote_prefill"):
|
||||
# Remote prefill: get all prompt blocks from remote.
|
||||
token_ids = request.prompt_token_ids or []
|
||||
actual = self._mamba_prefill_token_count(len(token_ids))
|
||||
actual = self._get_remote_prefill_token_count(len(token_ids))
|
||||
count = actual - num_computed_tokens
|
||||
if count > 0:
|
||||
return count, True
|
||||
|
||||
@@ -116,7 +116,7 @@ class NixlPushConnectorScheduler(NixlBaseConnectorScheduler):
|
||||
|
||||
if params is not None and params.get("do_remote_prefill"):
|
||||
token_ids = request.prompt_token_ids or []
|
||||
actual = self._mamba_prefill_token_count(len(token_ids))
|
||||
actual = self._get_remote_prefill_token_count(len(token_ids))
|
||||
count = actual - num_computed_tokens
|
||||
if count > 0:
|
||||
return count, True
|
||||
|
||||
@@ -610,6 +610,18 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
"min_p", self._DEFAULT_SAMPLING_PARAMS["min_p"]
|
||||
)
|
||||
|
||||
# Merge server-default stop_token_ids (e.g., model-specific tokens
|
||||
# like </call> for gpt-oss) with any request-specified ones
|
||||
stop_token_ids = self.stop_token_ids
|
||||
default_stop_ids = default_sampling_params.get("stop_token_ids")
|
||||
if default_stop_ids:
|
||||
if not stop_token_ids:
|
||||
stop_token_ids = list(default_stop_ids)
|
||||
else:
|
||||
stop_token_ids = list(
|
||||
dict.fromkeys([*stop_token_ids, *default_stop_ids])
|
||||
)
|
||||
|
||||
prompt_logprobs = self.prompt_logprobs
|
||||
if prompt_logprobs is None and self.echo:
|
||||
prompt_logprobs = self.top_logprobs
|
||||
@@ -661,7 +673,7 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
min_p=min_p,
|
||||
seed=self.seed,
|
||||
stop=self.stop,
|
||||
stop_token_ids=self.stop_token_ids,
|
||||
stop_token_ids=stop_token_ids,
|
||||
logprobs=self.top_logprobs if self.logprobs else None,
|
||||
prompt_logprobs=prompt_logprobs,
|
||||
ignore_eos=self.ignore_eos,
|
||||
|
||||
@@ -288,6 +288,18 @@ class CompletionRequest(OpenAIBaseModel):
|
||||
"min_p", self._DEFAULT_SAMPLING_PARAMS["min_p"]
|
||||
)
|
||||
|
||||
# Merge server-default stop_token_ids (e.g., model-specific tokens
|
||||
# like </call> for gpt-oss) with any request-specified ones
|
||||
stop_token_ids = self.stop_token_ids
|
||||
default_stop_ids = default_sampling_params.get("stop_token_ids")
|
||||
if default_stop_ids:
|
||||
if not stop_token_ids:
|
||||
stop_token_ids = list(default_stop_ids)
|
||||
else:
|
||||
stop_token_ids = list(
|
||||
dict.fromkeys([*stop_token_ids, *default_stop_ids])
|
||||
)
|
||||
|
||||
prompt_logprobs = self.prompt_logprobs
|
||||
if prompt_logprobs is None and self.echo:
|
||||
prompt_logprobs = self.logprobs
|
||||
@@ -341,7 +353,7 @@ class CompletionRequest(OpenAIBaseModel):
|
||||
min_p=min_p,
|
||||
seed=self.seed,
|
||||
stop=self.stop,
|
||||
stop_token_ids=self.stop_token_ids,
|
||||
stop_token_ids=stop_token_ids,
|
||||
logprobs=self.logprobs,
|
||||
ignore_eos=self.ignore_eos,
|
||||
max_tokens=max_tokens if not echo_without_generation else 1,
|
||||
|
||||
@@ -318,6 +318,7 @@ def _parse_function_call(message: Message, recipient: str) -> list[ResponseOutpu
|
||||
type="function_call",
|
||||
name=function_name,
|
||||
id=f"fc_{random_id}",
|
||||
status="completed",
|
||||
)
|
||||
output_items.append(response_item)
|
||||
return output_items
|
||||
|
||||
@@ -79,6 +79,7 @@ if TYPE_CHECKING:
|
||||
VLLM_MAX_AUDIO_CLIP_FILESIZE_MB: int = 25
|
||||
VLLM_MAX_AUDIO_DECODE_DURATION_S: int = 600
|
||||
VLLM_MAX_AUDIO_PREPROCESS_WORKERS: int = max(1, min(os.cpu_count() or 1, 2))
|
||||
VLLM_MAX_IMAGE_PIXELS: int = 178_956_970
|
||||
VLLM_VIDEO_LOADER_BACKEND: str = "opencv"
|
||||
VLLM_MEDIA_CONNECTOR: str = "http"
|
||||
VLLM_MM_HASHER_ALGORITHM: str = "blake3"
|
||||
@@ -954,6 +955,13 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
str(max(1, min(os.cpu_count() or 1, 2))),
|
||||
)
|
||||
),
|
||||
# Maximum decoded image size in pixels. Small compressed images can
|
||||
# expand into gigabytes of raster memory. This limit is enforced before
|
||||
# decoding so the memory is never allocated. Default matches PIL's
|
||||
# built-in 2x decompression-bomb threshold (~179M pixels, ~680 MB RGB).
|
||||
"VLLM_MAX_IMAGE_PIXELS": lambda: int(
|
||||
os.getenv("VLLM_MAX_IMAGE_PIXELS", "178956970")
|
||||
),
|
||||
# Backend for Video IO — selects the frame-sampling algorithm.
|
||||
# - "opencv": uniform sampling.
|
||||
# - "opencv_dynamic": duration-aware dynamic sampling.
|
||||
@@ -2083,6 +2091,7 @@ def compile_factors() -> dict[str, object]:
|
||||
"VLLM_MAX_AUDIO_CLIP_FILESIZE_MB",
|
||||
"VLLM_MAX_AUDIO_DECODE_DURATION_S",
|
||||
"VLLM_MAX_AUDIO_PREPROCESS_WORKERS",
|
||||
"VLLM_MAX_IMAGE_PIXELS",
|
||||
"VLLM_VIDEO_LOADER_BACKEND",
|
||||
"VLLM_MEDIA_CONNECTOR",
|
||||
"VLLM_OBJECT_STORAGE_SHM_BUFFER_NAME",
|
||||
|
||||
@@ -458,6 +458,7 @@ class Attention(nn.Module, AttentionLayerBase):
|
||||
# shape does not match the query shape, so we optionally let the model
|
||||
# definition specify the output tensor shape.
|
||||
output_shape: torch.Size | None = None,
|
||||
output_dtype: torch.dtype | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
The KV cache is stored inside this class and is accessed via
|
||||
@@ -472,7 +473,8 @@ class Attention(nn.Module, AttentionLayerBase):
|
||||
torch.ops.vllm.maybe_calc_kv_scales(
|
||||
query, key, value, _encode_layer_name(self.layer_name)
|
||||
)
|
||||
output_dtype = query.dtype
|
||||
if output_dtype is None:
|
||||
output_dtype = query.dtype
|
||||
if self.query_quant is not None:
|
||||
# quantizing with a simple torch operation enables
|
||||
# torch.compile to fuse this into previous ops
|
||||
|
||||
@@ -818,14 +818,14 @@ class MLAAttention(nn.Module, AttentionLayerBase):
|
||||
attn_out,
|
||||
lse,
|
||||
get_dcp_group(),
|
||||
is_lse_base_on_e=True,
|
||||
is_lse_base_on_e=self.impl.lse_base_on_e,
|
||||
)
|
||||
else:
|
||||
attn_out = cp_lse_ag_out_rs(
|
||||
attn_out,
|
||||
lse,
|
||||
get_dcp_group(),
|
||||
is_lse_base_on_e=True,
|
||||
is_lse_base_on_e=self.impl.lse_base_on_e,
|
||||
)
|
||||
|
||||
# v_up projection
|
||||
|
||||
@@ -301,6 +301,13 @@ def convert_to_unquantized_kernel_format(
|
||||
is_gated_act_gemm=is_act_and_mul,
|
||||
)
|
||||
|
||||
if (
|
||||
unquantized_backend == UnquantizedMoeBackend.TRITON
|
||||
and current_platform.is_rocm()
|
||||
and envs.VLLM_ROCM_MOE_PADDING
|
||||
):
|
||||
# Skip .contiguous(): it would undo the ROCm MoE weight padding.
|
||||
return w13_weight, w2_weight
|
||||
return w13_weight.contiguous(), w2_weight.contiguous()
|
||||
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
||||
TopKWeightAndReduceDelegate,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.utils import moe_kernel_quantize_input
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.math_utils import round_up
|
||||
from vllm.v1.worker.ubatching import (
|
||||
dbo_current_ubatch_id,
|
||||
@@ -59,6 +60,7 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
|
||||
self.dp_size = dp_size
|
||||
self.rank_expert_offset = rank_expert_offset
|
||||
self.async_prepare = True
|
||||
self.sync_dbo_comm = current_platform.is_rocm()
|
||||
|
||||
# The dispatch function returns a handle that the combine function
|
||||
# requires. Under DBO microbatching we must track one handle per
|
||||
@@ -68,6 +70,13 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
|
||||
# From https://github.com/deepseek-ai/DeepEP/blob/9fe9021f29c9083cd1808ab36b740208524d9f63/deep_ep/buffer.py#L164
|
||||
self.available_rank_configs = [2, 4, 8, 16, 24, 32, 64, 128, 144, 160]
|
||||
|
||||
def _sync_dbo_comm_if_needed(self) -> None:
|
||||
if self.sync_dbo_comm and dbo_enabled():
|
||||
# ROCm DeepEP HT dispatch/combine reuse Buffer-owned communication
|
||||
# workspace. Do not let the next DBO ubatch reuse that workspace
|
||||
# before this ubatch's HT kernel has completed.
|
||||
torch.cuda.current_stream().synchronize()
|
||||
|
||||
def num_dispatchers(self) -> int:
|
||||
return self.num_dispatchers_
|
||||
|
||||
@@ -161,6 +170,8 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
|
||||
allocate_on_comm_stream=False,
|
||||
)
|
||||
|
||||
self._sync_dbo_comm_if_needed()
|
||||
|
||||
# record the handle for this ubatch
|
||||
a2a_idx = dbo_current_ubatch_id()
|
||||
self.handles[a2a_idx] = handle
|
||||
@@ -375,6 +386,8 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
|
||||
allocate_on_comm_stream=False,
|
||||
)
|
||||
|
||||
self._sync_dbo_comm_if_needed()
|
||||
|
||||
dbo_switch_to_compute()
|
||||
|
||||
if do_async:
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from vllm.model_executor.layers.hpc.hpc_module import HpcModule
|
||||
from vllm.model_executor.layers.hpc.rope_norm import HpcRopeNorm, QkNormPolicy
|
||||
|
||||
__all__ = [
|
||||
"HpcModule",
|
||||
"HpcRopeNorm",
|
||||
"QkNormPolicy",
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class HpcModule(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
@classmethod
|
||||
def support(cls, *args, **kwargs):
|
||||
return True
|
||||
|
||||
def process_weights_after_loading(self, model):
|
||||
pass
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
pass
|
||||
@@ -0,0 +1,408 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""HPC fused RoPE + QK-Norm + KV-Cache-Write (+ optional FP8 Q quant).
|
||||
|
||||
Decoupled from HpcAttentionImpl; extra params are passed via layer attrs.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import IntEnum
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.config import get_current_vllm_config_or_none
|
||||
from vllm.forward_context import ForwardContext, get_forward_context
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
from vllm.model_executor.layers.hpc.hpc_module import HpcModule
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
from vllm.v1.attention.backends.hpc_attn import HpcAttnMetadata
|
||||
from vllm.v1.attention.backends.registry import AttentionBackendEnum
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_hpc_rope_norm_instances: dict[str, HpcRopeNorm] = {}
|
||||
|
||||
|
||||
class QkNormPolicy(IntEnum):
|
||||
"""Order of QK-RMSNorm relative to RoPE in the fused HPC rope_norm kernel.
|
||||
|
||||
The values are part of the HPC kernel ABI (passed through as ints), so they
|
||||
must stay in sync with the kernel's expectations.
|
||||
"""
|
||||
|
||||
# No QK-Norm: apply RoPE only.
|
||||
NONE = 0
|
||||
# Apply RoPE first, then QK-RMSNorm.
|
||||
ROPE_THEN_NORM = 1
|
||||
# Apply QK-RMSNorm first, then RoPE (e.g. HunYuan V3).
|
||||
NORM_THEN_ROPE = 2
|
||||
|
||||
|
||||
def hpc_rope_norm_forward(
|
||||
qkv: torch.Tensor,
|
||||
output: torch.Tensor,
|
||||
layer_name: str,
|
||||
) -> None:
|
||||
"""Top-level custom op: RoPE + QK-Norm + KV-Cache-Write + FP8 Q quant.
|
||||
|
||||
Fully opaque to torch.compile (dynamo).
|
||||
"""
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
attn_metadata: Any = forward_context.attn_metadata
|
||||
if isinstance(attn_metadata, dict):
|
||||
attn_metadata = attn_metadata[layer_name]
|
||||
|
||||
if attn_metadata is None:
|
||||
output.zero_()
|
||||
return
|
||||
|
||||
attn_layer = forward_context.no_compile_layers[layer_name]
|
||||
# bind_kv_cache stores the per-layer KV cache as a single 5D tensor
|
||||
# (num_blocks, 2, block_size, num_kv_heads, head_size), so use it directly.
|
||||
kv_cache = attn_layer.kv_cache
|
||||
|
||||
if kv_cache.numel() == 0:
|
||||
output.zero_()
|
||||
return
|
||||
|
||||
assert kv_cache.dim() == 5, (
|
||||
f"Expected kv_cache to have 5 dims, got {tuple(kv_cache.shape)}"
|
||||
)
|
||||
|
||||
rope_norm = _hpc_rope_norm_instances[layer_name]
|
||||
rope_norm._forward_impl(qkv, kv_cache, attn_metadata, attn_layer, output)
|
||||
|
||||
|
||||
def hpc_rope_norm_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
output: torch.Tensor,
|
||||
layer_name: str,
|
||||
) -> None:
|
||||
"""Fake impl for torch.compile trace; output is a mutated arg."""
|
||||
return
|
||||
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="hpc_rope_norm_forward",
|
||||
op_func=hpc_rope_norm_forward,
|
||||
mutates_args=["output"],
|
||||
fake_impl=hpc_rope_norm_forward_fake,
|
||||
)
|
||||
|
||||
|
||||
@CustomOp.register("hpc_rope_norm")
|
||||
class HpcRopeNorm(CustomOp, HpcModule):
|
||||
"""HPC fused RoPE + QK-Norm + KV-Cache-Write (+ optional FP8 Q quant).
|
||||
|
||||
Registered as a sub-module in model layers (e.g. HunYuanAttention).
|
||||
Norm weights are extracted from fallback norm modules via
|
||||
process_weights_after_loading() after all weights are loaded.
|
||||
|
||||
forward() is dispatched by CustomOp framework:
|
||||
- In compiled mode: forward_cuda() calls torch.ops.vllm.hpc_rope_norm_forward
|
||||
as a splitting point — internal Python control flow is opaque
|
||||
to torch.compile and not captured by CUDA Graph.
|
||||
- In eager/native mode: forward_native() falls back to forward_cuda().
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
head_dim: int,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
use_qk_norm: bool,
|
||||
fallback_qnorm: torch.nn.Module | None,
|
||||
fallback_knorm: torch.nn.Module | None,
|
||||
kv_cache_dtype: str,
|
||||
layer_name: str,
|
||||
qk_norm_policy: QkNormPolicy = QkNormPolicy.ROPE_THEN_NORM,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.num_kv_heads = num_kv_heads
|
||||
self.head_dim = head_dim
|
||||
|
||||
self.use_qk_norm = use_qk_norm
|
||||
|
||||
self.q_size = num_heads * head_dim
|
||||
self.kv_size = num_kv_heads * head_dim
|
||||
|
||||
# Register as a non-persistent buffer so it participates in sleep
|
||||
# level-2 save/restore (CuMemAllocator) but is excluded from the
|
||||
# checkpoint state_dict.
|
||||
self.register_buffer("cos_sin_cache", cos_sin_cache.float(), persistent=False)
|
||||
|
||||
self.fallback_qnorm = fallback_qnorm
|
||||
self.fallback_knorm = fallback_knorm
|
||||
|
||||
self.head_per_group = num_heads // num_kv_heads
|
||||
|
||||
# Pre-allocate norm weight tensors as Parameters so they are tracked by
|
||||
# CuMemAllocator (for sleep/wake_up) and have stable addresses for CUDA
|
||||
# Graph replay. process_weights_after_loading() updates them inplace via
|
||||
# copy_() so refit does not invalidate captured graph tensor pointers.
|
||||
# Shape is [head_dim] to match the HPC kernel's q/k_norm_weight layout.
|
||||
if use_qk_norm and fallback_qnorm is not None:
|
||||
self.qnorm_weight: torch.nn.Parameter | None = torch.nn.Parameter(
|
||||
torch.empty(head_dim, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
else:
|
||||
self.qnorm_weight = None
|
||||
if use_qk_norm and fallback_knorm is not None:
|
||||
self.knorm_weight: torch.nn.Parameter | None = torch.nn.Parameter(
|
||||
torch.empty(head_dim, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
else:
|
||||
self.knorm_weight = None
|
||||
|
||||
self.use_fp8 = "fp8" in kv_cache_dtype
|
||||
# The RMSNorm/RoPE ordering is model dependent (e.g. HunYuan V3 applies
|
||||
# QK-Norm before RoPE -> NORM_THEN_ROPE), so it is supplied by the
|
||||
# caller. When QK-Norm is disabled the policy is forced to NONE.
|
||||
self.qk_norm_policy = qk_norm_policy if use_qk_norm else QkNormPolicy.NONE
|
||||
|
||||
# Register layer_name + add self to the global instance registry so the
|
||||
# module-level custom op (hpc_rope_norm_forward) can route back here.
|
||||
self.layer_name: str | None = None
|
||||
self.register_layer_name(layer_name)
|
||||
|
||||
@classmethod
|
||||
def support(
|
||||
cls,
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
head_dim: int,
|
||||
kv_cache_dtype: str,
|
||||
) -> bool:
|
||||
"""Check whether HpcRopeNorm is supported for the given config."""
|
||||
# HpcRopeNorm is only enabled together with the HPC attention backend.
|
||||
vllm_config = get_current_vllm_config_or_none()
|
||||
if (
|
||||
vllm_config is None
|
||||
or vllm_config.attention_config.backend != AttentionBackendEnum.HPC_ATTN
|
||||
):
|
||||
return False
|
||||
|
||||
if kv_cache_dtype not in ("fp8_e4m3", "auto"):
|
||||
logger.warning_once(
|
||||
f"hpc rope_norm not support kv_cache_dtype:{kv_cache_dtype}, "
|
||||
"only support fp8_e4m3, bfloat16"
|
||||
)
|
||||
return False
|
||||
|
||||
if head_dim not in (128,):
|
||||
logger.warning_once("hpc rope_norm only support head_dim == 128.")
|
||||
return False
|
||||
|
||||
head_per_group = num_heads // num_kv_heads
|
||||
if head_per_group not in (4, 8):
|
||||
logger.warning_once("hpc rope_norm only support head_per_group in [4, 8].")
|
||||
return False
|
||||
|
||||
logger.info_once("enable hpc rope_norm")
|
||||
return True
|
||||
|
||||
def process_weights_after_loading(self, model: torch.nn.Module = None) -> None:
|
||||
"""Copy norm weights (float32) from fallback norm modules inplace.
|
||||
|
||||
Uses copy_() to preserve tensor addresses for CUDA Graph / refit
|
||||
compatibility. Called by the model's load_weights() after all weights
|
||||
are loaded (and generically from the model loader for DummyModelLoader
|
||||
/ sleep-wake_up reload paths).
|
||||
"""
|
||||
if self.use_qk_norm:
|
||||
if self.fallback_qnorm is not None and self.qnorm_weight is not None:
|
||||
self.qnorm_weight.data.copy_(self.fallback_qnorm.weight.data.float())
|
||||
if self.fallback_knorm is not None and self.knorm_weight is not None:
|
||||
self.knorm_weight.data.copy_(self.fallback_knorm.weight.data.float())
|
||||
|
||||
def register_layer_name(self, layer_name: str) -> None:
|
||||
"""Register layer_name and add self to the global registry.
|
||||
|
||||
The global registry is needed because the bottom-level torch op
|
||||
(hpc_rope_norm_forward) is a module-level function and needs to
|
||||
route back to the correct instance via layer_name.
|
||||
"""
|
||||
self.layer_name = layer_name
|
||||
_hpc_rope_norm_instances[layer_name] = self
|
||||
logger.debug(
|
||||
"[rope_norm] registered HpcRopeNorm for layer: %s",
|
||||
layer_name,
|
||||
)
|
||||
|
||||
def forward_native(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
layer_name: str,
|
||||
) -> torch.Tensor:
|
||||
"""Native fallback path: delegates to forward_cuda().
|
||||
|
||||
For now, the default native path will use CUDA backend path.
|
||||
Other platforms may override via OOT registration.
|
||||
"""
|
||||
return self.forward_cuda(qkv, layer_name)
|
||||
|
||||
def forward_cuda(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
layer_name: str,
|
||||
) -> torch.Tensor:
|
||||
"""CUDA path: invoke the torch custom op as a compile splitting point."""
|
||||
num_tokens = qkv.shape[0]
|
||||
output = torch.empty(
|
||||
(num_tokens, self.num_heads, self.head_dim),
|
||||
dtype=torch.float8_e4m3fn if self.use_fp8 else qkv.dtype,
|
||||
device=qkv.device,
|
||||
)
|
||||
|
||||
torch.ops.vllm.hpc_rope_norm_forward(qkv, output, layer_name)
|
||||
return output
|
||||
|
||||
def _forward_impl(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: HpcAttnMetadata,
|
||||
attn_layer: torch.nn.Module,
|
||||
output: torch.Tensor,
|
||||
) -> None:
|
||||
"""Actual forward logic called by the custom op.
|
||||
|
||||
Writes processed q into *output* and attaches extra params
|
||||
(e.g. FP8 scales) to *attn_layer* as attributes.
|
||||
"""
|
||||
import hpc
|
||||
|
||||
num_actual_tokens = attn_metadata.num_actual_tokens
|
||||
num_prefill_reqs = attn_metadata.num_prefills
|
||||
num_decode_reqs = attn_metadata.num_decodes
|
||||
num_decode_tokens = attn_metadata.num_decode_tokens
|
||||
|
||||
qkv = qkv[:num_actual_tokens]
|
||||
|
||||
num_prefill_tokens = num_actual_tokens - num_decode_tokens
|
||||
|
||||
# KV cache for the FP8 path is stored as uint8; view it as fp8 so the
|
||||
# rope_norm_store_kv_fp8 kernel can write quantized K/V in-place.
|
||||
if self.use_fp8:
|
||||
kv_cache = kv_cache.view(torch.float8_e4m3fn)
|
||||
|
||||
# Per-tensor K/V scales (shape [1]) used by the FP8 kernel.
|
||||
k_scale = attn_layer._k_scale.reshape(1)
|
||||
v_scale = attn_layer._v_scale.reshape(1)
|
||||
|
||||
q_norm_weight = (
|
||||
self.qnorm_weight if self.qk_norm_policy != QkNormPolicy.NONE else None
|
||||
)
|
||||
k_norm_weight = (
|
||||
self.knorm_weight if self.qk_norm_policy != QkNormPolicy.NONE else None
|
||||
)
|
||||
|
||||
# Dynamic per-token-per-head Q quant + per-tensor K/V (dqskv).
|
||||
# rope_norm_store_kv_fp8 is registered as a torch op whose ``quant_policy``
|
||||
# argument is typed as ``int``; pybind cannot cast the hpc.QuantType enum
|
||||
# automatically, so pass its integer ``.value``.
|
||||
QUANT_POLICY_DQSKV = hpc.QuantType.QPERTOKEN_PERHEAD_KPERTENSOR_VPERTENSOR.value
|
||||
|
||||
# --- Prefill ---
|
||||
if num_prefill_reqs > 0:
|
||||
seq_lens_prefill = attn_metadata.seq_lens[num_decode_reqs:]
|
||||
cu_seqlens_prefill = attn_metadata.qo_indptr
|
||||
max_seqlens = attn_metadata.max_query_len
|
||||
block_table_prefill = attn_metadata.block_table_tensor[num_decode_reqs:]
|
||||
qkv_prefill = qkv[num_decode_tokens:]
|
||||
out_q_prefill = output[
|
||||
num_decode_tokens : num_decode_tokens + num_prefill_tokens
|
||||
]
|
||||
|
||||
if self.use_fp8:
|
||||
_, q_scale, split_k_flag = hpc.rope_norm_store_kv_fp8(
|
||||
key_cache=kv_cache[:, 0],
|
||||
value_cache=kv_cache[:, 1],
|
||||
qkv=qkv_prefill,
|
||||
cos_sin=self.cos_sin_cache,
|
||||
num_seqlen_per_req=seq_lens_prefill,
|
||||
q_index=cu_seqlens_prefill,
|
||||
kvcache_indices=block_table_prefill,
|
||||
is_prefill=True,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
quant_policy=QUANT_POLICY_DQSKV,
|
||||
max_seqlens=max_seqlens,
|
||||
q_norm_weight=q_norm_weight,
|
||||
k_norm_weight=k_norm_weight,
|
||||
qk_norm_policy=self.qk_norm_policy,
|
||||
out_q=out_q_prefill,
|
||||
)
|
||||
attn_metadata.hpc_prefill_q_scale = q_scale
|
||||
else:
|
||||
hpc.rope_norm_store_kv(
|
||||
kv_cache[:, 0],
|
||||
kv_cache[:, 1],
|
||||
qkv_prefill,
|
||||
self.cos_sin_cache,
|
||||
seq_lens_prefill,
|
||||
cu_seqlens_prefill,
|
||||
block_table_prefill,
|
||||
True, # is_prefill
|
||||
q_norm_weight=q_norm_weight,
|
||||
k_norm_weight=k_norm_weight,
|
||||
out_q=out_q_prefill,
|
||||
qk_norm_policy=self.qk_norm_policy,
|
||||
)
|
||||
|
||||
# --- Decode ---
|
||||
if num_decode_reqs > 0:
|
||||
num_seq_kvcache = attn_metadata.seq_lens[:num_decode_reqs]
|
||||
block_table_decode = attn_metadata.block_table_tensor[:num_decode_reqs]
|
||||
qkv_decode = qkv[:num_decode_tokens]
|
||||
# Single-token decode: q_index is the per-request prefix sum
|
||||
# [0, 1, ..., num_decode_reqs].
|
||||
qo_indptr_decode = torch.arange(
|
||||
num_decode_reqs + 1, dtype=torch.int32, device=qkv.device
|
||||
)
|
||||
out_q_decode = output[:num_decode_tokens]
|
||||
|
||||
if self.use_fp8:
|
||||
_, q_scale, split_k_flag = hpc.rope_norm_store_kv_fp8(
|
||||
key_cache=kv_cache[:, 0],
|
||||
value_cache=kv_cache[:, 1],
|
||||
qkv=qkv_decode,
|
||||
cos_sin=self.cos_sin_cache,
|
||||
num_seqlen_per_req=num_seq_kvcache,
|
||||
q_index=qo_indptr_decode,
|
||||
kvcache_indices=block_table_decode,
|
||||
is_prefill=False,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
quant_policy=QUANT_POLICY_DQSKV,
|
||||
max_seqlens=1,
|
||||
q_norm_weight=q_norm_weight,
|
||||
k_norm_weight=k_norm_weight,
|
||||
qk_norm_policy=self.qk_norm_policy,
|
||||
out_q=out_q_decode,
|
||||
)
|
||||
attn_metadata.hpc_decode_q_scale = q_scale
|
||||
if split_k_flag is not None:
|
||||
attn_metadata.hpc_split_k_flag = split_k_flag
|
||||
else:
|
||||
hpc.rope_norm_store_kv(
|
||||
kv_cache[:, 0],
|
||||
kv_cache[:, 1],
|
||||
qkv_decode,
|
||||
self.cos_sin_cache,
|
||||
num_seq_kvcache,
|
||||
qo_indptr_decode,
|
||||
block_table_decode,
|
||||
False, # is_prefill
|
||||
q_norm_weight=q_norm_weight,
|
||||
k_norm_weight=k_norm_weight,
|
||||
out_q=out_q_decode,
|
||||
qk_norm_policy=self.qk_norm_policy,
|
||||
)
|
||||
@@ -174,31 +174,31 @@ def chunk_gated_delta_rule_cutedsl(
|
||||
When ``core_attn_out`` is provided, ``output`` is an unsqueezed view of
|
||||
that buffer.
|
||||
"""
|
||||
q_3d = q.squeeze(0)
|
||||
k_3d = k.squeeze(0)
|
||||
v_3d = v.squeeze(0)
|
||||
g_2d = g.squeeze(0)
|
||||
beta_2d = beta.squeeze(0)
|
||||
q = q.squeeze(0)
|
||||
k = k.squeeze(0)
|
||||
v = v.squeeze(0)
|
||||
g = g.squeeze(0)
|
||||
beta = beta.squeeze(0)
|
||||
|
||||
_, _, head_k_dim = k_3d.shape
|
||||
_, num_v_heads, head_v_dim = v_3d.shape
|
||||
_, _, K_dim = k.shape
|
||||
_, num_v_heads, V_dim = v.shape
|
||||
chunk_size = 64
|
||||
upper_bound_chunks = chunk_indices.shape[0]
|
||||
pad_t = upper_bound_chunks * chunk_size
|
||||
total_chunks_ptr = chunk_offsets[-1:]
|
||||
|
||||
g_cu = torch.empty_like(g_2d, dtype=torch.float32)
|
||||
u = q_3d.new_empty(pad_t, num_v_heads, head_v_dim)
|
||||
w = q_3d.new_empty(pad_t, num_v_heads, head_k_dim)
|
||||
g_cu = torch.empty_like(g, dtype=torch.float32)
|
||||
u = q.new_empty(pad_t, num_v_heads, V_dim)
|
||||
w = q.new_empty(pad_t, num_v_heads, K_dim)
|
||||
|
||||
num_sms = torch.cuda.get_device_properties(q.device).multi_processor_count
|
||||
kkt_inv_uw_cutedsl(
|
||||
k_3d,
|
||||
v_3d,
|
||||
k,
|
||||
v,
|
||||
u,
|
||||
w,
|
||||
g_2d,
|
||||
beta_2d,
|
||||
g,
|
||||
beta,
|
||||
g_cu,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
@@ -206,16 +206,11 @@ def chunk_gated_delta_rule_cutedsl(
|
||||
num_sms=num_sms,
|
||||
)
|
||||
|
||||
h = k_3d.new_empty(
|
||||
upper_bound_chunks,
|
||||
num_v_heads,
|
||||
head_v_dim,
|
||||
head_k_dim,
|
||||
)
|
||||
v_new = q_3d.new_empty(pad_t, num_v_heads, head_v_dim)
|
||||
h = k.new_empty(upper_bound_chunks, num_v_heads, V_dim, K_dim)
|
||||
v_new = q.new_empty(pad_t, num_v_heads, V_dim)
|
||||
final_state = torch.empty_like(initial_state)
|
||||
h_cutedsl(
|
||||
k_3d,
|
||||
k,
|
||||
u,
|
||||
w,
|
||||
v_new,
|
||||
@@ -227,12 +222,12 @@ def chunk_gated_delta_rule_cutedsl(
|
||||
chunk_offsets,
|
||||
)
|
||||
|
||||
output = core_attn_out if core_attn_out is not None else torch.empty_like(v_3d)
|
||||
scale = head_k_dim**-0.5
|
||||
output = core_attn_out if core_attn_out is not None else torch.empty_like(v)
|
||||
scale = K_dim**-0.5
|
||||
o_cutedsl(
|
||||
q_3d,
|
||||
k_3d,
|
||||
v_new.view(upper_bound_chunks, chunk_size, num_v_heads, head_v_dim),
|
||||
q,
|
||||
k,
|
||||
v_new,
|
||||
h,
|
||||
g_cu,
|
||||
output,
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user