Compare commits

...
Author SHA1 Message Date
Bugen Zhao cf30ef60cf use newtype for logprobs count (-1 for all)
Signed-off-by: Bugen Zhao <i@bugenzhao.com>
2026-06-22 19:30:10 +08:00
Tuukka SarviandGitHub 89accad2cc [ROCm][DSV4] Disable TileLang MHC dispatch on gfx942 (#45931)
Signed-off-by: Tuukka Sarvi <tuukka.sarvi@amd.com>
2026-06-22 09:26:54 +00:00
3c8e49596c [Model] ColQwen3.5: fix retrieval correctness (bias + bidirectional) (#46108)
Signed-off-by: Athrael Soju <athrael.soju@gmail.com>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-22 17:25:54 +08:00
Weiwei SunGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>Jiangyun Zhu
cec2ec1176 [Bugfix] Avoid racy accepted counts in async spec decode (#45100)
Signed-off-by: Weiwei Sun <68775773+sunnweiwei@users.noreply.github.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Co-authored-by: Jiangyun Zhu <riverclouds.zhu@qq.com>
2026-06-22 08:53:16 +00:00
liuzhenweiandGitHub 435f82d61a [Bugfix] Fix Llama4ForCausalLM initialization test failure (#46341)
Signed-off-by: zhenwei-intel <zhenwei.liu@intel.com>
2026-06-22 08:40:43 +00:00
Roger WangandGitHub 1c4b51b990 Temporarily skip M3 on CI (#46352)
Signed-off-by: Roger Wang <hey@rogerw.io>
2026-06-22 01:35:31 -07:00
2e2c47928b [Doc] Update MiniMax-M3 (#45940)
Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
Signed-off-by: Roger Wang <hey@rogerw.io>
Co-authored-by: Jiangyun Zhu <riverclouds.zhu@qq.com>
Co-authored-by: Roger Wang <hey@rogerw.io>
2026-06-22 01:23:27 -07:00
Chao-Ju ChenGitHubBugen Zhaomergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
80abe0de7d [Rust Frontend] Support thinking_token_budget for chat and completions (#46137)
Co-authored-by: Bugen Zhao <i@bugenzhao.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Signed-off-by: RickyChen / 陳昭儒 <ricky.chen@infinirc.com>
Signed-off-by: Bugen Zhao <i@bugenzhao.com>
2026-06-22 16:00:02 +08:00
a9f7b2d41c [feature][kv_offload] Self-describing KV events for OffloadingConnector (#43468)
Signed-off-by: Change72 <changg@nvidia.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-06-22 07:27:46 +00:00
Tiezhen WANGGitHubClaudemergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
d14e551a53 [Model] Remove MiniMaxText01, MiniMaxVL01, MiniMaxForCausalLM (#45993)
Signed-off-by: Xianbao QIAN <xianbao.qian@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-22 15:20:46 +08:00
68567ef2df [CPUOffloadingManager] Maintain evictable list in LRUCachePolicy (#46216)
Signed-off-by: <>
Co-authored-by: Varun Sundar Rabindranath <varun-sundar-rabindranath@h100-01.nemg-001.lab.rdu2.dc.redhat.com>
2026-06-22 06:54:44 +00:00
6bc6f2d86d [1/N][Core] add partial prefix cache primitives (#45939)
Signed-off-by: zjy0516 <riverclouds.zhu@qq.com>
Co-authored-by: Yifan Qiao <yifanqiao@inferact.ai>
2026-06-21 23:43:10 -07:00
wang.yuqiandGitHub 1eb2cc961e [Frontend] Refactor ServingTokenization entrypoint. (#46022)
Signed-off-by: wang.yuqi <yuqi.wang@daocloud.io>
2026-06-22 06:27:58 +00:00
ReidGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
31124749d1 [Bugfix] [Rust Frontend] Fix stop string truncation with repeated matches (#46113)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Signed-off-by: reidliu41 <reid201711@gmail.com>
2026-06-22 14:11:29 +08:00
Ma JianandGitHub 9037498c22 [DSV4][XPU] Pass gemm1_clamp_limit to XpuFusedMoe (#44517)
Signed-off-by: Ma Jian <jian1.ma@intel.com>
2026-06-22 12:57:10 +08:00
db32b53e30 [SpecDecode] Support DFlash with FlashInfer (#43081)
Signed-off-by: gss <2783977641@qq.com>
Co-authored-by: gss <2783977641@qq.com>
2026-06-22 04:55:30 +00:00
xiangdongandGitHub b529bfd6c5 [XPU][CI] Add agent_tags for Intel GPU CI (#45768)
Signed-off-by: zengxian <xiangdong.zeng@intel.com>
2026-06-22 10:33:17 +08:00
Micah WilliamsonandGitHub f3df7a7231 [ROCm][CI] Enable kv_connector unit tests on ROCm (#45955)
Signed-off-by: Micah Williamson <micah.williamson@amd.com>
2026-06-22 05:08:44 +03:00
485bbe1c6f [CI] Fix missing tp_size attribute on RoutedExperts (#46163)
Signed-off-by: Felix Marty <Felix.Marty@amd.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
2026-06-21 18:46:49 -06:00
MattandGitHub a19ff2218a [Hardware][AMD][CI] Fix Spec Decode Eagle test group (#46018)
Signed-off-by: Matthew Wong <Matthew.Wong2@amd.com>
2026-06-21 17:40:02 -05:00
MattandGitHub 4f0d0049a0 [Hardware][AMD][CI] Fix Kernels Attention test groups (#46080)
Signed-off-by: Matthew Wong <Matthew.Wong2@amd.com>
2026-06-21 17:10:51 -05:00
13b83d77ad [ROCm][CI] skip test_double_aiter_rms_quant_fusion (#45967)
Signed-off-by: charlifu <charlifu@amd.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
2026-06-21 16:53:11 -05:00
MattandGitHub 50241602fd [Hardware][AMD][CI] Fix gfx942 Kernels MoE test group (#46298)
Signed-off-by: Matthew Wong <Matthew.Wong2@amd.com>
2026-06-21 16:45:37 -05:00
Ting SUNandGitHub 12fe2a9aac [Bugfix][Qwen3-VL] Fix multi-video crash with list-valued fps/num_frames (#46305)
Signed-off-by: Ting Sun <suntcrick@gmail.com>
2026-06-21 14:31:23 -07:00
Benjamin ChislettandGitHub 89bd2c14d3 [Spec Decode] Add Qwen3 architecture support for EAGLE3 (#43132)
Signed-off-by: Benjamin Chislett <bchislett@nvidia.com>
2026-06-21 13:55:26 -07:00
ZedongLiuandGitHub 9c450b1027 [Kernel][Bugfix] Fix INT8 per-token-head KV cache rounding in Triton reshape-and-cache (#45361)
Signed-off-by: ZedongLiu <113341356+Zedong-Liu@users.noreply.github.com>
2026-06-21 15:59:40 -04:00
RanranGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>Isotr0py
635c38338a [Multimodal] Add Qwen2-VL/Qwen2.5-VL processor-mapped video loader (#45555)
Signed-off-by: Ranran <hzz5361@psu.edu>
Signed-off-by: Ranran Haoran Zhang <ranzhang@redhat.com>
Signed-off-by: Isotr0py <mozf@mail2.sysu.edu.cn>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Co-authored-by: Isotr0py <mozf@mail2.sysu.edu.cn>
2026-06-21 18:56:50 +00:00
c441ad1c07 [KV Offloading] Add labeled metrics support (#45957)
Signed-off-by: srinivas_oo7 <sklinkedin0120@gmail.com>
Co-authored-by: srinivas_oo7 <sklinkedin0120@gmail.com>
2026-06-21 18:04:01 +00:00
Jee Jee LiandGitHub 745bba5ea8 [Model]Fix MiniMaxM2ForCausalLM perf regression (#45935)
Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
2026-06-22 00:28:52 +08:00
2cac89f9da [Spec Decode] Support mixed KV page sizes for DFlash (#45181)
Signed-off-by: Alex Steiner <asteiner@nvidia.com>
Signed-off-by: Giancarlo Delfin <gdelfin@inferact.ai>
Signed-off-by: Yifan Qiao <yifanqiao@inferact.ai>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
Co-authored-by: Giancarlo Delfin <gdelfin@inferact.ai>
Co-authored-by: Yifan Qiao <yifanqiao@inferact.ai>
2026-06-21 22:45:14 +08:00
3e6e33526d [Disagg] return routed_experts on streaming generate responses (#44638)
Signed-off-by: aoshen02 <aoshen@inferact.ai>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: Roger Wang <hey@rogerw.io>
2026-06-21 07:37:10 -07:00
junkang1991GitHubHongxia YangTan Pin SiangvllmellmChun FangTianDi101functionstackxtjtanaamergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
b91b7726e0 [ROCm][P/D] Support MiniMax-M3 mixed KV layouts in MoRIIO READ mode (#46039)
Signed-off-by: Jun Kang Chow <junkangchow@gmail.com>
Signed-off-by: tjtanaa <tunjian.tan@embeddedllm.com>
Co-authored-by: Hongxia Yang <hongxia.yang@amd.com>
Co-authored-by: Tan Pin Siang <tanpinsiang@gmail.com>
Co-authored-by: vllmellm <vllm.ellm@embeddedllm.com>
Co-authored-by: Chun Fang <chun.fang@amd.com>
Co-authored-by: TianDi101 <ditian12@amd.com>
Co-authored-by: functionstackx <47992694+functionstackx@users.noreply.github.com>
Co-authored-by: tjtanaa <tunjian.tan@embeddedllm.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-21 12:55:19 +00:00
Palaiologos1453andGitHub d3ad8e8bcd [Bugfix] Defer offload reads while transfers are pending (#46231)
Signed-off-by: test test <2260891073@qq.com>
2026-06-21 14:30:13 +03:00
b80ce9dd2f [CI][test] Replace InternVL2-1B with InternVL3-1B in test_pipeline_parallel.py (#46241)
Signed-off-by: wentian-byte <192079369+wentian-byte@users.noreply.github.com>
Co-authored-by: wentian-byte <192079369+wentian-byte@users.noreply.github.com>
2026-06-21 15:11:19 +08:00
b5495cc5f9 Fix memory pointer overflow in Mamba state buffers (#44665)
Signed-off-by: Shifani Rajabose <shifani.rajabose@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-06-21 14:00:50 +08:00
153 changed files with 5398 additions and 4968 deletions
@@ -21,6 +21,10 @@ steps:
timeout_in_minutes: 30
optional: true
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 24+
no_plugin: true
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
@@ -38,6 +42,10 @@ steps:
timeout_in_minutes: 30
optional: true
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
@@ -55,6 +63,10 @@ steps:
timeout_in_minutes: 30
optional: true
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
@@ -5,6 +5,10 @@ steps:
- label: XPU Sleep Mode
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
+4
View File
@@ -5,6 +5,10 @@ steps:
- label: Engine (1 GPU)
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -6,6 +6,10 @@ steps:
key: eplb-algorithm
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
+4
View File
@@ -5,6 +5,10 @@ steps:
- label: vLLM IR Tests
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
+24
View File
@@ -5,6 +5,10 @@ steps:
- label: LoRA Runtime + Utils
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 24+
no_plugin: true
working_dir: "."
env:
@@ -34,6 +38,10 @@ steps:
- label: LoRA Fused/MoE Kernels
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -54,6 +62,10 @@ steps:
- label: LoRA Punica Kernels
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -74,6 +86,10 @@ steps:
- label: LoRA Punica FP8/XPU Ops
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -94,6 +110,10 @@ steps:
- label: LoRA Models
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 24+
no_plugin: true
working_dir: "."
env:
@@ -117,6 +137,10 @@ steps:
- label: LoRA Multimodal
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
+24
View File
@@ -5,6 +5,10 @@ steps:
- label: V1 Core + KV + Metrics
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -31,6 +35,10 @@ steps:
- label: V1 Sample + Logits
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -71,6 +79,10 @@ steps:
- label: XPU CPU Offload
timeout_in_minutes: 60
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -95,6 +107,10 @@ steps:
key: regression
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -126,6 +142,10 @@ steps:
timeout_in_minutes: 30
num_devices: 2
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -157,6 +177,10 @@ steps:
key: async-engine-inputs-utils-worker
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 24+
no_plugin: true
working_dir: "."
env:
@@ -5,6 +5,10 @@ steps:
- label: Model Runner V2 Core Tests (Intel)
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -30,6 +34,10 @@ steps:
- label: Model Runner V2 Examples (Intel)
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 24+
no_plugin: true
working_dir: "."
env:
@@ -6,6 +6,10 @@ steps:
key: multi-modal-models-standard-1-qwen2
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -27,6 +31,10 @@ steps:
key: multi-modal-models-standard-2-qwen3-gemma
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -47,6 +55,10 @@ steps:
key: multi-modal-models-standard-3-llava-qwen2-vl
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 24+
no_plugin: true
working_dir: "."
env:
@@ -68,6 +80,10 @@ steps:
key: multi-modal-models-standard-4-other-whisper
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
@@ -88,6 +104,10 @@ steps:
key: multi-modal-processor
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
+16
View File
@@ -19,6 +19,10 @@ steps:
- image-build-xpu
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 24+
no_plugin: true
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
@@ -49,6 +53,10 @@ steps:
- image-build-xpu
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
@@ -74,6 +82,10 @@ steps:
- image-build-xpu
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
@@ -93,6 +105,10 @@ steps:
- image-build-xpu
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
@@ -4,6 +4,11 @@
set -euo pipefail
if python3 -c "import torch; raise SystemExit(0 if torch.version.hip is not None else 1)"; then
uv pip install --system -r /vllm-workspace/requirements/kv_connectors_rocm.txt
exit 0
fi
REQUIREMENTS_FILE="${KV_CONNECTORS_REQUIREMENTS:-/vllm-workspace/requirements/kv_connectors.txt}"
uv pip install --system -r "${REQUIREMENTS_FILE}"
+11 -9
View File
@@ -1594,9 +1594,10 @@ steps:
#---------------------------------------------------------- mi300 · kernels ----------------------------------------------------------#
- label: Kernels Attention Test %N # TBD
timeout_in_minutes: 180
timeout_in_minutes: 55
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
agent_pool: mi300_1
optional: true
parallelism: 2
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -1627,10 +1628,11 @@ steps:
- pytest -v -s kernels/core --ignore=kernels/core/test_minimax_reduce_rms.py kernels/test_concat_mla_q.py kernels/test_top_k_per_row.py
- label: Kernels MoE Test %N # TBD
timeout_in_minutes: 180
timeout_in_minutes: 50
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
agent_pool: mi300_1
parallelism: 4
optional: true
parallelism: 5
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- csrc/quantization/cutlass_w8a8/moe/
@@ -2120,9 +2122,10 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "draft_model or no_sync or batch_inference"
- label: Spec Decode Eagle # TBD
timeout_in_minutes: 180
timeout_in_minutes: 45
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
agent_pool: mi300_1
optional: true
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- vllm/v1/spec_decode/
@@ -3040,7 +3043,7 @@ steps:
#---------------------------------------------------------- mi355 · kernels ----------------------------------------------------------#
- label: Kernels (B200-MI355) # TBD
timeout_in_minutes: 180
timeout_in_minutes: 15
mirror_hardwares: [amdexperimental, amdproduction, amdgfx950nightly, amdmi355]
agent_pool: mi355_1
working_dir: "/vllm-workspace/"
@@ -3064,11 +3067,10 @@ steps:
- pytest -v -s tests/kernels/attention/test_attention_selector.py
- label: Kernels Attention Test %N # TBD
timeout_in_minutes: 180
timeout_in_minutes: 60
mirror_hardwares: [amdexperimental, amdproduction, amdgfx950nightly, amdmi355]
agent_pool: mi355_1
parallelism: 2
optional: true
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- csrc/attention/
@@ -3082,10 +3084,10 @@ steps:
- pytest -v -s kernels/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
- label: Kernels MoE Test %N # TBD
timeout_in_minutes: 180
timeout_in_minutes: 50
mirror_hardwares: [amdexperimental, amdproduction, amdgfx950nightly, amdmi355]
agent_pool: mi355_1
parallelism: 4
parallelism: 5
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- csrc/quantization/cutlass_w8a8/moe/
+30
View File
@@ -74,6 +74,20 @@ steps:
commands:
- pytest -v -s kernels/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
parallelism: 2
mirror:
amd:
device: mi325_1
timeout_in_minutes: 55
depends_on:
- image-build-amd
source_file_dependencies:
- csrc/attention/
- vllm/v1/attention
- vllm/model_executor/layers/attention
- tests/kernels/attention
- vllm/_aiter_ops.py
- vllm/envs.py
- vllm/platforms/rocm.py
- label: Kernels Attention DiffKV Test (H100)
key: kernels-attention-diffkv-test-h100
@@ -128,6 +142,22 @@ steps:
- pytest -v -s kernels/moe --ignore=kernels/moe/test_modular_oai_triton_moe.py --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
- pytest -v -s kernels/moe/test_modular_oai_triton_moe.py --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
parallelism: 5
mirror:
amd:
device: mi325_1
timeout_in_minutes: 50
source_file_dependencies:
- csrc/quantization/cutlass_w8a8/moe/
- csrc/moe/
- tests/kernels/moe
- vllm/model_executor/layers/fused_moe/
- vllm/distributed/device_communicators/
- vllm/envs.py
- vllm/config
- vllm/_aiter_ops.py
- vllm/platforms/rocm.py
depends_on:
- image-build-amd
- label: Kernels Mamba Test
key: kernels-mamba-test
+6
View File
@@ -105,6 +105,12 @@ steps:
# Integration test for streaming correctness (requires special branch).
- pip install -U git+https://github.com/robertgshaw2-redhat/lm-evaluation-harness.git@streaming-api
- pytest -v -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine
mirror:
amd:
device: mi325_1
timeout_in_minutes: 60
depends_on:
- image-build-amd
- label: V1 Others (CPU)
key: v1-others-cpu
+6
View File
@@ -107,6 +107,12 @@ steps:
- tests/compile/passes
commands:
- pytest -s -v compile/passes --ignore compile/passes/distributed
mirror:
amd:
device: mi300_1
timeout_in_minutes: 180
depends_on:
- image-build-amd
- label: PyTorch Fullgraph Smoke Test
key: pytorch-fullgraph-smoke-test
+14
View File
@@ -12,6 +12,20 @@ steps:
- tests/v1/e2e/spec_decode/
commands:
- pytest -v -s v1/e2e/spec_decode -k "eagle_correctness"
mirror:
amd:
device: mi325_1
timeout_in_minutes: 45
depends_on:
- image-build-amd
source_file_dependencies:
- vllm/v1/spec_decode/
- vllm/v1/worker/gpu/spec_decode/
- vllm/model_executor/model_loader/
- vllm/v1/sample/
- vllm/model_executor/layers/
- tests/v1/e2e/spec_decode/
- vllm/platforms/rocm.py
- label: Spec Decode Eagle Nightly B200
key: spec-decode-eagle-nightly-b200
+2 -2
View File
@@ -136,7 +136,7 @@ The model should also be added to the `MODELS_CONFIG_MAP` dictionary in [vllm/mo
For case (2), we recommend using as a reference the implementation of [`JambaForCausalLM`](../../../vllm/model_executor/models/jamba.py) (for an example of a model that uses Mamba-1 and attention together) or [`NemotronHForCausalLM`](../../../vllm/model_executor/models/nemotron_h.py) (for an example of a model that uses Mamba-2 and attention together).
These models should follow the same instructions as case (1), but they should inherit protocol `IsHybrid` (instead of `IsAttentionFree`) and it is *not* necessary to add them to the `MODELS_CONFIG_MAP` (their runtime defaults will be inferred from the protocol).
For case (3), we recommend looking at the implementation of [`MiniMaxText01ForCausalLM`](../../../vllm/model_executor/models/minimax_text_01.py) or [`Lfm2ForCausalLM`](../../../vllm/model_executor/models/lfm2.py) as a reference, which use custom "mamba-like" layers `MiniMaxText01LinearAttention` and `ShortConv` respectively.
For case (3), we recommend looking at the implementation of [`Lfm2ForCausalLM`](../../../vllm/model_executor/models/lfm2.py) as a reference, which uses a custom "mamba-like" layer `ShortConv`.
Please follow the same guidelines as case (2) for implementing these models.
We use "mamba-like" to refer to layers that possess a state that is updated in-place, rather than being appended-to (like KV cache for attention).
For implementing new custom mamba-like layers, one should inherit from `MambaBase` and implement the methods `get_state_dtype`, `get_state_shape` to calculate the data types and state shapes at runtime, as well as `mamba_type` and `get_attn_backend`.
@@ -144,5 +144,5 @@ It is also necessary to implement the "attention meta-data" class which handles
Please see [`LinearAttentionMetadata`](../../../vllm/v1/attention/backends/linear_attn.py) or [`ShortConvAttentionMetadata`](../../../vllm/v1/attention/backends/short_conv_attn.py) for examples of this.
It is also worth noting that we should update `MambaAttentionBackendEnum` in [`registry.py`](../../../vllm/v1/attention/backends/registry.py) when adding a new mamba backend.
Finally, if one wants to support torch compile and CUDA graphs, it necessary to wrap the call to the mamba-like layer inside a custom op and register it.
Please see the calls to `direct_register_custom_op` in [vllm/model_executor/models/minimax_text_01.py](../../../vllm/model_executor/models/minimax_text_01.py) or [vllm/model_executor/layers/mamba/short_conv.py](../../../vllm/model_executor/layers/mamba/short_conv.py) for examples of this.
Please see the calls to `direct_register_custom_op` in [vllm/model_executor/layers/mamba/linear/minimax_linear_attn.py](../../../vllm/model_executor/layers/mamba/linear/minimax_linear_attn.py) or [vllm/model_executor/layers/mamba/short_conv.py](../../../vllm/model_executor/layers/mamba/short_conv.py) for examples of this.
The new custom op should then be added to the list `_attention_ops` in [vllm/config/compilation.py](../../../vllm/config/compilation.py) to ensure that piecewise CUDA graphs works as intended.
+2 -2
View File
@@ -170,8 +170,8 @@ Priority is **1 = highest** (tried first).
| Backend | Version | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | --------- | --- | --------------- | ------------ |
| `CPU_ATTN` | | fp16, bf16, fp32 | `auto`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 112, 128, 160, 192, 224, 256, 512 | ❌ | ❌ | ❌ | ❌ | All | N/A |
| `FLASHINFER` | Native† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ❌ | | ❌ | ✅ | Decoder | 7.x-9.x |
| `FLASHINFER` | TRTLLM† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `nvfp4` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ✅ | | ❌ | ✅ | Decoder | 10.x |
| `FLASHINFER` | Native† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ❌ | | ❌ | ✅ | Decoder | 7.x-9.x |
| `FLASHINFER` | TRTLLM† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `nvfp4` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ✅ | | ❌ | ✅ | Decoder | 10.x |
| `FLASH_ATTN` | FA2* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ❌ | ✅ | All | ≥8.0 |
| `FLASH_ATTN` | FA3* | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | 9.x |
| `FLASH_ATTN` | FA4* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | ≥10.0 |
+1
View File
@@ -74,6 +74,7 @@ vllm serve <model> \
| `max_tracker_size` | no | `64000` | single-tier | Max entries in the lookup tracker. |
| `secondary_tiers` | no | `[]` | multi-tier | List of secondary tier configs (see below). |
| `offload_prompt_only` | no | `true` | both | If `true`, only prompt (prefill) blocks are offloaded; decode blocks are skipped. |
| `self_describing_kv_events` | no | `false` | single-tier | Opt-in. When `true` *and* KV cache events are enabled (`--kv-events-config` with `enable_kv_cache_events`), the connector emits self-describing block-granular `BlockStored`/`BlockRemoved` payloads (constituent block hashes, whole-chunk `token_ids`, per-block `block_size`, parent hash, LoRA + group/cache-spec metadata) instead of the placeholder fallback, so external KV-event consumers can index offloaded blocks. Inert unless events are enabled. Currently rejected by `TieringOffloadingSpec`. Full-attention groups only; sliding-window/SSM groups keep the placeholder fallback. In chunk mode (`block_size` > GPU block size), overlapping chunks re-announce shared per-block hashes, so consumers must reference-count (deduplicate) repeated store/remove announcements. |
| `spec_module_path` | no | — | both | Python import path for a custom `OffloadingSpec` not in the built-in registry. Required only when `spec_name` is not built-in (advanced). |
## Secondary Tiers
-9
View File
@@ -321,15 +321,6 @@ For Qwen2.5, the chat template in tokenizer_config.json has already included sup
Flags: `--tool-call-parser hermes`
### MiniMax Models (`minimax_m1`)
Supported models:
* `MiniMaxAi/MiniMax-M1-40k` (use with [examples/tool_chat_template_minimax_m1.jinja](../../examples/tool_chat_template_minimax_m1.jinja))
* `MiniMaxAi/MiniMax-M1-80k` (use with [examples/tool_chat_template_minimax_m1.jinja](../../examples/tool_chat_template_minimax_m1.jinja))
Flags: `--tool-call-parser minimax --chat-template examples/tool_chat_template_minimax_m1.jinja`
### DeepSeek-V3 Models (`deepseek_v3`)
Supported models:
+1 -1
View File
@@ -61,7 +61,7 @@ Models of any architecture can be converted into embedding models using `--conve
| `ColModernVBertForRetrieval` | ColModernVBERT | T / I | `ModernVBERT/colmodernvbert-merged` | | |
| `ColPaliForRetrieval` | ColPali | T / I | `vidore/colpali-v1.3-hf` | | |
| `ColQwen3` | Qwen3-VL | T / I | `TomoroAI/tomoro-colqwen3-embed-4b`, `TomoroAI/tomoro-colqwen3-embed-8b` | | |
| `ColQwen3_5` | ColQwen3.5 | T + I + V | `athrael-soju/colqwen3.5-4.5B-v3` | | |
| `ColQwen3_5` | ColQwen3.5 | T + I + V | `athrael-soju/colqwen3.5-4.5B-v3`, `vultr/VultronRetrieverPrime-Qwen3.5-8B` | | |
| `OpsColQwen3Model` | Qwen3-VL | T / I | `OpenSearch-AI/Ops-Colqwen3-4B`, `OpenSearch-AI/Ops-Colqwen3-8B` | | |
| `Qwen3VLNemotronEmbedModel` | Qwen3-VL | T / I | `nvidia/nemotron-colembed-vl-4b-v2`, `nvidia/nemotron-colembed-vl-8b-v2` | ✅︎ | ✅︎ |
| `*ForConditionalGeneration`<sup>C</sup>, `*ForCausalLM`<sup>C</sup>, etc. | Generative models | \* | N/A | \* | \* |
+1 -3
View File
@@ -441,7 +441,6 @@ th {
| `MiMoV2ForCausalLM` | MiMoV2Pro | `XiaomiMiMo/MiMo-V2.5-Pro`, etc. | | ✅︎ |
| `MiniCPMForCausalLM` | MiniCPM | `openbmb/MiniCPM-2B-sft-bf16`, `openbmb/MiniCPM-2B-dpo-bf16`, `openbmb/MiniCPM-S-1B-sft`, etc. | ✅︎ | ✅︎ |
| `MiniCPM3ForCausalLM` | MiniCPM3 | `openbmb/MiniCPM3-4B`, etc. | ✅︎ | ✅︎ |
| `MiniMaxForCausalLM` | MiniMax-Text | `MiniMaxAI/MiniMax-Text-01-hf`, etc. | | |
| `MiniMaxM2ForCausalLM` | MiniMax-M2, MiniMax-M2.1 | `MiniMaxAI/MiniMax-M2`, etc. | ✅︎ | ✅︎ |
| `MistralForCausalLM` | Ministral-3, Mistral, Mistral-Instruct | `mistralai/Ministral-3-3B-Instruct-2512`, `mistralai/Mistral-7B-v0.1`, `mistralai/Mistral-7B-Instruct-v0.1`, etc. | ✅︎ | ✅︎ |
| `MistralLarge3ForCausalLM` | Mistral-Large-3-675B-Base-2512, Mistral-Large-3-675B-Instruct-2512 | `mistralai/Mistral-Large-3-675B-Base-2512`, `mistralai/Mistral-Large-3-675B-Instruct-2512`, etc. | ✅︎ | ✅︎ |
@@ -487,8 +486,6 @@ th {
| `TeleChat2ForCausalLM` | TeleChat2 | `Tele-AI/TeleChat2-3B`, `Tele-AI/TeleChat2-7B`, `Tele-AI/TeleChat2-35B`, etc. | ✅︎ | ✅︎ |
| `TeleChat3ForCausalLM` | TeleChat3 | `Tele-AI/TeleChat3-36B-Thinking`, `Tele-AI/TeleChat3-Coder-36B-Thinking`, etc. | ✅︎ | ✅︎ |
| `TeleFLMForCausalLM` | TeleFLM | `CofeAI/FLM-2-52B-Instruct-2407`, `CofeAI/Tele-FLM`, etc. | ✅︎ | ✅︎ |
| `MiniMaxM1ForCausalLM` | MiniMax-Text | `MiniMaxAI/MiniMax-M1-40k`, `MiniMaxAI/MiniMax-M1-80k`, etc. | | |
| `MiniMaxText01ForCausalLM` | MiniMax-Text | `MiniMaxAI/MiniMax-Text-01`, etc. | | |
| `Zamba2ForCausalLM` | Zamba2 | `Zyphra/Zamba2-7B-instruct`, `Zyphra/Zamba2-2.7B-instruct`, `Zyphra/Zamba2-1.2B-instruct`, etc. | | |
!!! note
@@ -595,6 +592,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
| `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. | ✅︎ | |
| `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. | ✅︎ | ✅︎ |
| `MolmoForCausalLM` | Molmo | T + I<sup>+</sup> | `allenai/Molmo-7B-D-0924`, `allenai/Molmo-7B-O-0924`, etc. | ✅︎ | ✅︎ |
+1 -1
View File
@@ -128,7 +128,7 @@ Models that use Mamba-2 and Mamba-1 layers (e.g., `Mamba2ForCausalLM`, `MambaFor
Hybrid models that combine Mamba-2 and Mamba-1 layers with standard attention layers are also supported (e.g., `BambaForCausalLM`,
`Zamba2ForCausalLM`, `NemotronHForCausalLM`, `FalconH1ForCausalLM` and `GraniteMoeHybridForCausalLM`, `JambaForCausalLM`, `Plamo2ForCausalLM`).
Hybrid models with mechanisms different to Mamba are also supported (e.g, `MiniMaxText01ForCausalLM`, `MiniMaxM1ForCausalLM`, `Lfm2ForCausalLM`).
Hybrid models with mechanisms different to Mamba are also supported (e.g, `Lfm2ForCausalLM`).
Please note that prefix caching is not yet supported for any of the above models.
@@ -1481,39 +1481,6 @@ def run_minicpmv(questions: list[str], modality: str) -> ModelRequestData:
return run_minicpmv_base(questions, modality, "openbmb/MiniCPM-V-2_6")
def run_minimax_vl_01(questions: list[str], modality: str) -> ModelRequestData:
assert modality == "image"
model_name = "MiniMaxAI/MiniMax-VL-01"
engine_args = EngineArgs(
model=model_name,
max_num_seqs=2,
limit_mm_per_prompt={modality: 1},
trust_remote_code=True,
tensor_parallel_size=8,
)
tokenizer = AutoTokenizer.from_pretrained(model_name)
messages = [
[
{
"role": "user",
"content": [{"type": "image"}, {"type": "text", "text": question}],
}
]
for question in questions
]
prompts = tokenizer.apply_chat_template(
messages, add_generation_prompt=True, tokenize=False
)
return ModelRequestData(
engine_args=engine_args,
prompts=prompts,
)
# Mistral-3 HF-format
def run_mistral3(questions: list[str], modality: str) -> ModelRequestData:
assert modality == "image"
@@ -2485,7 +2452,6 @@ model_example_map = {
"mantis": run_mantis,
"minicpmo": run_minicpmo,
"minicpmv": run_minicpmv,
"minimax_vl_01": run_minimax_vl_01,
"mistral3": run_mistral3,
"molmo": run_molmo,
"molmo2": run_molmo2,
@@ -7,11 +7,27 @@ ColQwen3.5 is a multi-modal ColBERT-style model based on Qwen3.5.
It produces per-token embeddings and uses MaxSim scoring for retrieval
and reranking. Supports both text and image inputs.
Works for any ColQwen3.5 checkpoint, e.g. `athrael-soju/colqwen3.5-4.5B-v3`
or `vultr/VultronRetrieverPrime-Qwen3.5-8B`.
Start the server with:
vllm serve athrael-soju/colqwen3.5-4.5B --max-model-len 4096
vllm serve athrael-soju/colqwen3.5-4.5B-v3 --max-model-len 4096 \
--mm-processor-kwargs '{"min_pixels": 65536, "max_pixels": 1835008}'
Then run this script:
python colqwen3_5_rerank_online.py
Parity note (matching the native colpali ColQwen3_5Processor pipeline):
- Visual-token budget: ColQwen3_5Processor uses max_num_visual_tokens=1792,
i.e. max_pixels = 1792 * (patch_size*merge_size)^2 = 1792 * 32^2 = 1835008
(with min_pixels = shortest_edge = 65536). Pass these via --mm-processor-kwargs
as above; the default budget gives fewer visual tokens and lower retrieval ndcg.
- When you build prompts yourself (token_embed), reproduce the processor exactly:
image (document): wrap in the instruction template
"<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>"
"Describe the image.<|im_end|><|endoftext|>"
query: append the augmentation suffix <text> + "<|endoftext|>" * 10
Omitting these reproduces a silent ~2.5 ndcg@10 drop vs the native pipeline.
"""
import requests
@@ -1,91 +0,0 @@
{{ '<begin_of_document>' -}}
{%- if custom_tools is defined %}
{%- set tools = custom_tools %}
{%- endif %}
{%- if not tools is defined %}
{%- set tools = none %}
{%- endif %}
{#- Extract system message #}
{% set ns = namespace(system_prompt='') -%}
{%- if messages[0]['role'] == 'system' %}
{%- if messages[0]['content'] is string %}
{%- set ns.system_prompt = messages[0]['content']|trim %}
{%- else %}
{%- set ns.system_prompt = messages[0]['content'][0]['text']|trim %}
{%- endif %}
{%- set messages = messages[1:] %}
{%- else %}
{%- if tools is not none %}
{%- set ns.system_prompt = "You are a helpful assistant created by Minimax based on MiniMax-M1 model." %}
{%- else %}
{%- set ns.system_prompt = "You are a helpful assistant created by Minimax based on MiniMax-M1 model." %}
{%- endif %}
{%- endif %}
{#- System message #}
{%- if ns.system_prompt != '' %}
{{ '<beginning_of_sentence>system ai_setting=assistant\n' + ns.system_prompt + '<end_of_sentence>\n' -}}
{%- endif %}
{#- Tools configuration #}
{%- if tools is not none %}
{{ '<beginning_of_sentence>system tool_setting=tools\nYou are provided with these tools:\n<tools>\n' -}}
{%- for tool in tools %}
{{ tool | tojson ~ '\n' -}}
{%- endfor %}
{{ '</tools>\n\nIf you need to call tools, please respond with <tool_calls></tool_calls> XML tags, and provide tool-name and json-object of arguments, following the format below:\n<tool_calls>\n{"name": <tool-name>, "arguments": <args-json-object>}\n...\n</tool_calls><end_of_sentence>\n' -}}
{%- endif %}
{#- Process messages #}
{%- for message in messages %}
{%- if not (message.role == 'ipython' or message.role == 'tool' or 'tool_calls' in message) %}
{%- if message['role'] == 'user' %}
{{ '<beginning_of_sentence>user name=user\n' -}}
{%- if message['content'] is string %}
{{ message['content']|trim -}}
{%- else %}
{%- for content in message['content'] %}
{%- if content['type'] == 'text' %}
{{ content['text']|trim -}}
{%- endif %}
{%- endfor %}
{%- endif %}
{{ '<end_of_sentence>\n' -}}
{%- elif message['role'] == 'assistant' %}
{{ '<beginning_of_sentence>ai name=assistant\n' -}}
{%- if message['content'] is string %}
{{ message['content']|trim -}}
{%- else %}
{%- for content in message['content'] | selectattr('type', 'equalto', 'text') %}
{{ content['text']|trim -}}
{%- endfor %}
{%- endif %}
{{ '<end_of_sentence>\n' -}}
{%- endif %}
{%- elif 'tool_calls' in message %}
{{ '<beginning_of_sentence>ai name=assistant\n<tool_calls>\n' -}}
{%- for tool_call in message.tool_calls %}
{{ '{"name": "' + tool_call.function.name + '", "arguments": ' + tool_call.function.arguments | tojson + '}\n' -}}
{%- endfor %}
{{ '</tool_calls><end_of_sentence>\n' -}}
{%- elif message.role == "tool" or message.role == "ipython" %}
{{ '<beginning_of_sentence>tool name=tools\n' -}}
{%- if message.content is string %}
{{ 'tool result: ' + message.content + '\n\n' -}}
{%- else %}
{%- for content in message['content'] %}
{%- if content['type'] == 'text' %}
{{ 'tool result: ' + content['text'] + '\n\n' -}}
{%- elif content.get('name') %}
{{ 'tool name: ' + content['name'] + '\ntool result: ' + content['text'] + '\n\n' -}}
{%- endif %}
{%- endfor %}
{%- endif %}
{{ '<end_of_sentence>\n' -}}
{%- endif %}
{%- endfor %}
{%- if add_generation_prompt %}
{{ '<beginning_of_sentence>ai name=assistant\n' -}}
{%- endif %}
-1
View File
@@ -386,7 +386,6 @@ mod tests {
tool_chat_template_llama3.2_pythonic.jinja => String
tool_chat_template_llama4_json.jinja => OpenAi
tool_chat_template_llama4_pythonic.jinja => OpenAi
tool_chat_template_minimax_m1.jinja => OpenAi
tool_chat_template_mistral.jinja => String
tool_chat_template_mistral3.jinja => OpenAi
tool_chat_template_mistral_parallel.jinja => String
+4 -3
View File
@@ -16,7 +16,8 @@ use vllm_engine_core_client::protocol::logprobs::{
Logprobs, MaybeWireLogprobs, PositionLogprobs, TokenLogprob,
};
use vllm_engine_core_client::protocol::{
EngineCoreFinishReason, EngineCoreOutput, EngineCoreOutputs, EngineCoreRequest, StopReason,
EngineCoreFinishReason, EngineCoreOutput, EngineCoreOutputs, EngineCoreRequest, LogprobsCount,
StopReason,
};
use vllm_engine_core_client::test_utils::{IpcNamespace, spawn_mock_engine_task};
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig};
@@ -1387,8 +1388,8 @@ async fn chat_stream_and_collect_preserve_prompt_and_sample_logprobs() {
.await;
let mut request = sample_request("chat-logprobs");
request.sampling_params.logprobs = Some(1);
request.sampling_params.prompt_logprobs = Some(1);
request.sampling_params.logprobs = Some(LogprobsCount::Top(1));
request.sampling_params.prompt_logprobs = Some(LogprobsCount::Top(1));
let mut stream = chat.chat(request.clone()).await.unwrap();
match next_semantic(&mut stream).await.unwrap().unwrap() {
@@ -1,91 +0,0 @@
{{ '<begin_of_document>' -}}
{%- if custom_tools is defined %}
{%- set tools = custom_tools %}
{%- endif %}
{%- if not tools is defined %}
{%- set tools = none %}
{%- endif %}
{#- Extract system message #}
{% set ns = namespace(system_prompt='') -%}
{%- if messages[0]['role'] == 'system' %}
{%- if messages[0]['content'] is string %}
{%- set ns.system_prompt = messages[0]['content']|trim %}
{%- else %}
{%- set ns.system_prompt = messages[0]['content'][0]['text']|trim %}
{%- endif %}
{%- set messages = messages[1:] %}
{%- else %}
{%- if tools is not none %}
{%- set ns.system_prompt = "You are a helpful assistant created by Minimax based on MiniMax-M1 model." %}
{%- else %}
{%- set ns.system_prompt = "You are a helpful assistant created by Minimax based on MiniMax-M1 model." %}
{%- endif %}
{%- endif %}
{#- System message #}
{%- if ns.system_prompt != '' %}
{{ '<beginning_of_sentence>system ai_setting=assistant\n' + ns.system_prompt + '<end_of_sentence>\n' -}}
{%- endif %}
{#- Tools configuration #}
{%- if tools is not none %}
{{ '<beginning_of_sentence>system tool_setting=tools\nYou are provided with these tools:\n<tools>\n' -}}
{%- for tool in tools %}
{{ tool | tojson ~ '\n' -}}
{%- endfor %}
{{ '</tools>\n\nIf you need to call tools, please respond with <tool_calls></tool_calls> XML tags, and provide tool-name and json-object of arguments, following the format below:\n<tool_calls>\n{"name": <tool-name>, "arguments": <args-json-object>}\n...\n</tool_calls><end_of_sentence>\n' -}}
{%- endif %}
{#- Process messages #}
{%- for message in messages %}
{%- if not (message.role == 'ipython' or message.role == 'tool' or 'tool_calls' in message) %}
{%- if message['role'] == 'user' %}
{{ '<beginning_of_sentence>user name=user\n' -}}
{%- if message['content'] is string %}
{{ message['content']|trim -}}
{%- else %}
{%- for content in message['content'] %}
{%- if content['type'] == 'text' %}
{{ content['text']|trim -}}
{%- endif %}
{%- endfor %}
{%- endif %}
{{ '<end_of_sentence>\n' -}}
{%- elif message['role'] == 'assistant' %}
{{ '<beginning_of_sentence>ai name=assistant\n' -}}
{%- if message['content'] is string %}
{{ message['content']|trim -}}
{%- else %}
{%- for content in message['content'] | selectattr('type', 'equalto', 'text') %}
{{ content['text']|trim -}}
{%- endfor %}
{%- endif %}
{{ '<end_of_sentence>\n' -}}
{%- endif %}
{%- elif 'tool_calls' in message %}
{{ '<beginning_of_sentence>ai name=assistant\n<tool_calls>\n' -}}
{%- for tool_call in message.tool_calls %}
{{ '{"name": "' + tool_call.function.name + '", "arguments": ' + tool_call.function.arguments | tojson + '}\n' -}}
{%- endfor %}
{{ '</tool_calls><end_of_sentence>\n' -}}
{%- elif message.role == "tool" or message.role == "ipython" %}
{{ '<beginning_of_sentence>tool name=tools\n' -}}
{%- if message.content is string %}
{{ 'tool result: ' + message.content + '\n\n' -}}
{%- else %}
{%- for content in message['content'] %}
{%- if content['type'] == 'text' %}
{{ 'tool result: ' + content['text'] + '\n\n' -}}
{%- elif content.get('name') %}
{{ 'tool name: ' + content['name'] + '\ntool result: ' + content['text'] + '\n\n' -}}
{%- endif %}
{%- endfor %}
{%- endif %}
{{ '<end_of_sentence>\n' -}}
{%- endif %}
{%- endfor %}
{%- if add_generation_prompt %}
{{ '<beginning_of_sentence>ai name=assistant\n' -}}
{%- endif %}
+8 -3
View File
@@ -20,6 +20,7 @@ use serde_with::{DefaultOnNull, OneOrMany, serde_as};
use thiserror_ext::AsReport as _;
use uuid::Uuid;
use vllm_engine_core_client::TransportMode;
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_managed_engine::ManagedEngineConfig;
use vllm_managed_engine::cli::{ManagedEngineArgs, repartition_managed_engine_args};
use vllm_server::{
@@ -136,9 +137,9 @@ pub struct SharedRuntimeArgs {
pub max_model_len: Option<u32>,
/// Maximum number of log probabilities to return when `logprobs` is
/// specified in sampling parameters. `-1` means no cap.
#[arg(long, value_parser = clap::value_parser!(i32).range(-1..), allow_negative_numbers = true)]
#[arg(long, allow_negative_numbers = true)]
#[serde(default)]
pub max_logprobs: Option<i32>,
pub max_logprobs: Option<LogprobsCount>,
/// TCP port for the gRPC Generate service. When not set, no gRPC server is
/// started.
#[arg(long)]
@@ -529,7 +530,7 @@ impl ServeArgs {
self.managed_engine.clone().into_config(
self.runtime.model.clone(),
self.runtime.max_model_len,
self.runtime.max_logprobs,
self.runtime.max_logprobs.map(managed_max_logprobs_to_i32),
self.runtime.language_model_only,
self.runtime.disable_log_stats,
self.runtime.shutdown_timeout,
@@ -555,5 +556,9 @@ fn frontend_ipc_addresses() -> (String, String) {
)
}
fn managed_max_logprobs_to_i32(count: LogprobsCount) -> i32 {
i32::try_from(count).expect("max_logprobs is parsed through i32")
}
#[cfg(test)]
mod tests;
+4 -3
View File
@@ -1,5 +1,6 @@
use expect_test::expect;
use vllm_engine_core_client::TransportMode;
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_server::{Config, HttpListenerMode, ParserSelection, RendererSelection};
use super::{Cli, Command};
@@ -165,10 +166,10 @@ fn serve_args_forward_max_logprobs_to_frontend_and_managed_engine() {
let Command::Serve(args) = cli.command else {
panic!("expected serve args");
};
assert_eq!(args.runtime.max_logprobs, Some(-1));
assert_eq!(args.runtime.max_logprobs, Some(LogprobsCount::All));
let frontend_config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
assert_eq!(frontend_config.max_logprobs, Some(-1));
assert_eq!(frontend_config.max_logprobs, Some(LogprobsCount::All));
let engine_config = args.to_managed_engine_config(5555);
assert_eq!(engine_config.python_args, vec!["--max-logprobs", "-1"]);
@@ -529,7 +530,7 @@ fn frontend_args_json_accepts_supported_non_default_fields() {
assert_eq!(args.runtime.renderer, RendererSelection::DeepSeekV32);
assert!(args.runtime.language_model_only);
assert_eq!(args.runtime.max_model_len, Some(8192));
assert_eq!(args.runtime.max_logprobs, Some(-1));
assert_eq!(args.runtime.max_logprobs, Some(LogprobsCount::All));
assert_eq!(args.runtime.shutdown_timeout, 3);
}
@@ -6,7 +6,7 @@ use futures::StreamExt as _;
use tokio::time::timeout;
use tracing_subscriber::EnvFilter;
use vllm_engine_core_client::protocol::{
EngineCoreFinishReason, EngineCoreRequest, EngineCoreSamplingParams,
EngineCoreFinishReason, EngineCoreRequest, EngineCoreSamplingParams, LogprobsCount,
};
use vllm_engine_core_client::{
EngineCoreClient, EngineCoreClientConfig, EngineCoreStreamOutput, TransportMode,
@@ -33,10 +33,10 @@ struct Args {
output_timeout_secs: u64,
#[arg(long, default_value_t = 1)]
max_tokens: u32,
#[arg(long, default_value_t = 2)]
logprobs: i32,
#[arg(long, default_value_t = 1)]
prompt_logprobs: i32,
#[arg(long, default_value_t = LogprobsCount::Top(2), allow_negative_numbers = true)]
logprobs: LogprobsCount,
#[arg(long, default_value_t = LogprobsCount::Top(1), allow_negative_numbers = true)]
prompt_logprobs: LogprobsCount,
#[arg(long, default_value_t = 96)]
prompt_repeats: usize,
}
@@ -64,8 +64,8 @@ fn build_request(
request_id: String,
prompt_token_ids: Vec<u32>,
max_tokens: u32,
logprobs: i32,
prompt_logprobs: i32,
logprobs: LogprobsCount,
prompt_logprobs: LogprobsCount,
client_index: u32,
) -> EngineCoreRequest {
EngineCoreRequest {
@@ -0,0 +1,135 @@
use std::fmt;
use std::str::FromStr;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
/// Number of log probabilities requested for a token position.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LogprobsCount {
/// Return the full model vocabulary.
All,
/// Return the top-N tokens by probability.
Top(u32),
}
impl LogprobsCount {
/// Expands the count to the actual number of logprobs to return, given the vocabulary size.
pub fn expanded(self, vocab_size: usize) -> usize {
match self {
Self::All => vocab_size,
Self::Top(count) => count as usize,
}
}
}
impl TryFrom<i32> for LogprobsCount {
type Error = String;
fn try_from(value: i32) -> Result<Self, Self::Error> {
match value {
-1 => Ok(Self::All),
value if value < -1 => Err(format!("must be non-negative or -1, got {value}")),
value => Ok(Self::Top(value as u32)),
}
}
}
impl TryFrom<LogprobsCount> for i32 {
type Error = String;
fn try_from(value: LogprobsCount) -> Result<Self, Self::Error> {
match value {
LogprobsCount::All => Ok(-1),
LogprobsCount::Top(count) => {
i32::try_from(count).map_err(|_| format!("must fit within i32, got {count}"))
}
}
}
}
impl FromStr for LogprobsCount {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let value = s
.parse::<i32>()
.map_err(|e| format!("must be an i32 integer, got {s:?}: {e}"))?;
Self::try_from(value)
}
}
impl fmt::Display for LogprobsCount {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::All => (-1).fmt(f),
Self::Top(count) => count.fmt(f),
}
}
}
impl Serialize for LogprobsCount {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let value: i32 = (*self).try_into().map_err(serde::ser::Error::custom)?;
value.serialize(serializer)
}
}
impl<'de> Deserialize<'de> for LogprobsCount {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let value = i32::deserialize(deserializer)?;
Self::try_from(value).map_err(serde::de::Error::custom)
}
}
#[cfg(test)]
mod tests {
use rmpv::Value;
use super::*;
use crate::protocol::{decode_msgpack, encode_msgpack};
#[test]
fn logprobs_count_serializes_as_wire_integer() {
assert_eq!(serde_json::to_value(LogprobsCount::All).unwrap(), -1);
assert_eq!(serde_json::to_value(LogprobsCount::Top(3)).unwrap(), 3);
}
#[test]
fn logprobs_count_deserializes_wire_integer() {
assert_eq!(
serde_json::from_value::<LogprobsCount>(serde_json::json!(-1)).unwrap(),
LogprobsCount::All
);
assert_eq!(
serde_json::from_value::<LogprobsCount>(serde_json::json!(3)).unwrap(),
LogprobsCount::Top(3)
);
assert!(serde_json::from_value::<LogprobsCount>(serde_json::json!(-2)).is_err());
assert!(
serde_json::from_value::<LogprobsCount>(serde_json::json!(i64::from(i32::MAX) + 1))
.is_err()
);
}
#[test]
fn logprobs_count_decodes_msgpack_signed_and_unsigned() {
let mut encoded = Vec::new();
rmpv::encode::write_value(&mut encoded, &Value::from(-1)).unwrap();
assert_eq!(
decode_msgpack::<LogprobsCount>(&encoded).unwrap(),
LogprobsCount::All
);
let encoded = encode_msgpack(&LogprobsCount::Top(7)).unwrap();
assert_eq!(
decode_msgpack::<LogprobsCount>(&encoded).unwrap(),
LogprobsCount::Top(7)
);
}
}
@@ -56,6 +56,7 @@ mod classified_outputs;
pub mod dtype;
pub mod handshake;
pub mod logprobs;
mod logprobs_count;
pub mod lora;
pub mod multimodal;
pub mod stats;
@@ -66,6 +67,7 @@ pub use classified_outputs::{
};
pub use dtype::ModelDtype;
pub use logprobs::decode_engine_core_outputs;
pub use logprobs_count::LogprobsCount;
/// Request types are encoded as single-byte protocol constants so they can be
/// sent over the ZMQ socket without an extra encoding step.
@@ -277,14 +279,20 @@ pub struct EngineCoreSamplingParams {
pub max_tokens: u32,
/// Minimum number of tokens to generate before EOS or stop-token handling.
pub min_tokens: u32,
/// Maximum number of reasoning ("thinking") tokens to emit before the
/// reasoning section is force-closed. `None` means unlimited; the
/// user-facing `-1` sentinel is normalized to `None` by the frontend before
/// reaching this DTO, so only non-negative values are sent. Enforced
/// engine-side (and only when a reasoning parser is configured).
pub thinking_token_budget: Option<u64>,
/// Number of log probabilities to return per generated token.
///
/// `None` disables sample logprobs. `-1` requests the full vocabulary.
pub logprobs: Option<i32>,
/// `None` disables sample logprobs.
pub logprobs: Option<LogprobsCount>,
/// Number of log probabilities to return per prompt token.
///
/// `None` disables prompt logprobs. `-1` requests the full vocabulary.
pub prompt_logprobs: Option<i32>,
/// `None` disables prompt logprobs.
pub prompt_logprobs: Option<LogprobsCount>,
/// Minimum probability threshold for token sampling.
pub min_p: f32,
/// Frequency penalty applied by the sampler.
@@ -345,6 +353,7 @@ impl EngineCoreSamplingParams {
seed: None,
max_tokens: 65536,
min_tokens: 0,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.0,
@@ -150,6 +150,7 @@ fn sample_request_with_id(request_id: &str) -> EngineCoreRequest {
top_k: 8,
max_tokens: 32,
min_tokens: 1,
thinking_token_budget: Some(256),
stop_token_ids: vec![151643],
eos_token_id: Some(151645),
all_stop_token_ids: BTreeSet::from([151643, 151645]),
@@ -2502,6 +2503,7 @@ fn python_msgpack_fixtures_match_rust_encoding() {
seed: None,
max_tokens: 16,
min_tokens: 0,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.0,
@@ -39,6 +39,7 @@ class EngineCoreSamplingParams(msgspec.Struct, dict=True, omit_defaults=True):
seed: int | None = None
max_tokens: int = 16
min_tokens: int = 0
thinking_token_budget: int | None = None
min_p: float = 0.0
frequency_penalty: float = 0.0
presence_penalty: float = 0.0
@@ -122,6 +123,7 @@ request = EngineCoreRequest(
seed=None,
max_tokens=32,
min_tokens=1,
thinking_token_budget=256,
min_p=0.0,
frequency_penalty=0.0,
presence_penalty=0.0,
+3 -11
View File
@@ -2,12 +2,13 @@ use std::collections::HashMap;
use std::fmt;
use std::time::Duration;
use anyhow::{Result, bail};
use anyhow::Result;
use axum::http::{HeaderName, HeaderValue, Method};
use educe::Educe;
use serde::Serialize;
use serde_json::Value;
use vllm_chat::{ChatTemplateContentFormatOption, ParserSelection, RendererSelection};
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_engine_core_client::{CoordinatorMode as EngineCoreCoordinatorMode, TransportMode};
/// How the HTTP server obtains its listening socket.
@@ -133,7 +134,7 @@ pub struct Config {
pub chat_template_content_format: ChatTemplateContentFormatOption,
/// Optional maximum number of top log probabilities accepted by the
/// frontend. `None` delegates to the text layer default.
pub max_logprobs: Option<i32>,
pub max_logprobs: Option<LogprobsCount>,
/// HTTP/API-server behavior switches.
pub api_server_options: ApiServerOptions,
/// CORS settings applied to every HTTP response.
@@ -158,15 +159,6 @@ 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(max_logprobs) = self.max_logprobs
&& max_logprobs < -1
{
bail!(
"max_logprobs must be non-negative or -1, got {}",
max_logprobs
);
}
Ok(())
}
+13
View File
@@ -103,6 +103,7 @@ fn is_request_validation_error(error: &vllm_text::Error) -> bool {
| vllm_text::Error::EmptyPromptTokenIds { .. }
| vllm_text::Error::Logprobs(_)
| vllm_text::Error::OutOfVocab(_)
| vllm_text::Error::InvalidThinkingTokenBudget
// An empty tokenized prompt detected later, at request prepare
// time, surfaces through the transparent Llm wrapper.
| vllm_text::Error::Llm(vllm_llm::Error::EmptyPromptTokenIds { .. })
@@ -127,6 +128,18 @@ mod tests {
assert!(response.error.message.contains("9000"));
}
#[test]
fn invalid_thinking_token_budget_maps_to_invalid_request() {
let api_error = text_submit_error(
"failed to submit completion request",
vllm_text::Error::InvalidThinkingTokenBudget,
);
assert_eq!(api_error.status_code(), StatusCode::BAD_REQUEST);
let response = api_error.to_error_response();
assert_eq!(response.error.error_type, "invalid_request_error");
assert!(response.error.message.contains("thinking_token_budget"));
}
#[test]
fn chat_wrapped_prompt_too_long_maps_to_invalid_request() {
let error = vllm_chat::Error::Text(vllm_text::Error::PromptTooLong {
+13 -9
View File
@@ -3,7 +3,7 @@
use tonic::Status;
use uuid::Uuid;
use vllm_engine_core_client::protocol::{StopReason, StructuredOutputsParams};
use vllm_engine_core_client::protocol::{LogprobsCount, StopReason, StructuredOutputsParams};
use vllm_text::{
DecodedLogprobs, DecodedPromptLogprobs, FinishReason, Finished, Prompt, SamplingParams,
TextDecodeOptions, TextRequest,
@@ -202,18 +202,22 @@ fn build_sampling_params(
/// Map the proto `CandidateTokens` selector to a `(logprobs_count,
/// logprob_token_ids)` pair.
///
/// - `top_n(k)` → `(k, None)` — return top-k candidates by probability
/// - `all` → `(-1, None)` — return the full vocabulary
/// - `top_n(k)` → `(Top(k), None)` — return top-k candidates by probability
/// - `all` → `(All, None)` — return the full vocabulary
/// - `token_ids(n)` → `(1, Some(vec of n token ids))` — return logprobs for specific tokens (the
/// count `n` is stored in the proto as the number of token IDs that follow, but the actual IDs
/// are carried via `logprob_token_ids` on `SamplingParams`)
/// - absent → `(1, None)` — just the sampled/scored token
fn candidate_logprob_spec(candidates: Option<&pb::CandidateTokens>) -> (i32, Option<Vec<u32>>) {
/// - absent → `(Top(1), None)` — just the sampled/scored token
fn candidate_logprob_spec(
candidates: Option<&pb::CandidateTokens>,
) -> (LogprobsCount, Option<Vec<u32>>) {
match candidates.and_then(|c| c.select.as_ref()) {
Some(pb::candidate_tokens::Select::TopN(n)) => (*n as i32, None),
Some(pb::candidate_tokens::Select::All(true)) => (-1, None),
Some(pb::candidate_tokens::Select::TokenIds(ids)) => (1, Some(ids.ids.clone())),
_ => (1, None),
Some(pb::candidate_tokens::Select::TopN(n)) => (LogprobsCount::Top(*n), None),
Some(pb::candidate_tokens::Select::All(true)) => (LogprobsCount::All, None),
Some(pb::candidate_tokens::Select::TokenIds(ids)) => {
(LogprobsCount::Top(1), Some(ids.ids.clone()))
}
_ => (LogprobsCount::Top(1), None),
}
}
@@ -87,6 +87,7 @@ pub(super) fn prepare_generate_request(
#[cfg(test)]
mod tests {
use serde_json::json;
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_text::Prompt;
use super::prepare_generate_request;
@@ -132,10 +133,13 @@ mod tests {
Prompt::TokenIds(vec![11, 22, 33])
);
assert_eq!(prepared.text_request.sampling_params.max_tokens, Some(7));
assert_eq!(prepared.text_request.sampling_params.logprobs, Some(2));
assert_eq!(
prepared.text_request.sampling_params.logprobs,
Some(LogprobsCount::Top(2))
);
assert_eq!(
prepared.text_request.sampling_params.prompt_logprobs,
Some(1)
Some(LogprobsCount::Top(1))
);
assert!(prepared.text_request.sampling_params.ignore_eos);
assert_eq!(prepared.text_request.priority, -3);
@@ -150,6 +154,33 @@ mod tests {
);
}
#[test]
fn prepare_generate_request_forwards_thinking_token_budget() {
let request: GenerateRequest = serde_json::from_value(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"token_ids": [11, 22, 33],
"sampling_params": {
"thinking_token_budget": 64
}
}))
.expect("parse request");
let prepared = prepare_generate_request(
request,
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
ResolvedRequestContext::default(),
)
.expect("prepare");
// The raw inference route shares `vllm_text::SamplingParams`, so the
// field is carried through to lowering exactly like the OpenAI routes
// (normalization/validation then happens in `lower_sampling_params`).
assert_eq!(
prepared.text_request.sampling_params.thinking_token_budget,
Some(64)
);
}
#[test]
fn prepare_generate_request_gates_continuous_usage_on_include_usage() {
let request: GenerateRequest = serde_json::from_value(json!({
@@ -34,16 +34,6 @@ 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."
);
}
Ok(())
}
@@ -4,6 +4,7 @@ use vllm_chat::{
ChatMessage as VllmChatMessage, ChatOptions, ChatRequest, ChatTool, ChatToolChoice,
GenerationPromptMode, SamplingParams,
};
use vllm_engine_core_client::protocol::LogprobsCount;
use super::types::ChatCompletionRequest;
use super::validate;
@@ -94,7 +95,7 @@ pub(super) fn prepare_chat_request(
// Auto-enable prompt logprobs for non-streaming echo, matching Python vLLM's
// behavior.
let top_logprobs = request.top_logprobs.unwrap_or(0);
let top_logprobs = request.top_logprobs.unwrap_or(LogprobsCount::Top(0));
let prompt_logprobs = request
.prompt_logprobs
.or((request.echo && !request.stream).then_some(top_logprobs));
@@ -115,6 +116,7 @@ pub(super) fn prepare_chat_request(
seed: request.seed,
max_tokens: request.max_completion_tokens,
min_tokens: request.min_tokens,
thinking_token_budget: request.thinking_token_budget,
logprobs: request.logprobs.then_some(top_logprobs),
prompt_logprobs,
min_p: request.min_p,
@@ -377,6 +379,7 @@ mod tests {
ChatTool as VllmChatTool, ChatToolChoice, GenerationPromptMode,
SamplingParams as VllmSamplingParams,
};
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_text::output::TextDecodeOptions;
use super::prepare_chat_request;
@@ -613,6 +616,31 @@ mod tests {
assert_eq!(prepared.chat_request.sampling_params, expected);
}
#[test]
fn prepare_chat_request_passes_through_thinking_token_budget() {
let prepare = |budget: Option<i64>| {
prepare_chat_request(
ChatCompletionRequest {
thinking_token_budget: budget,
..base_request()
},
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
ResolvedRequestContext::default(),
)
.expect("request is valid")
.chat_request
.sampling_params
.thinking_token_budget
};
// The convert layer forwards the raw value verbatim (including the `-1`
// "unlimited" sentinel); normalization/validation happens during
// lowering (see `vllm_text::lower`).
assert_eq!(prepare(Some(64)), Some(64));
assert_eq!(prepare(Some(-1)), Some(-1));
assert_eq!(prepare(None), None);
}
#[test]
fn prepare_chat_request_accepts_developer_messages() {
let request = ChatCompletionRequest {
@@ -941,7 +969,7 @@ mod tests {
let request = ChatCompletionRequest {
stream: false,
logprobs: true,
prompt_logprobs: Some(2),
prompt_logprobs: Some(LogprobsCount::Top(2)),
..base_request()
};
@@ -954,10 +982,13 @@ mod tests {
assert!(prepared.options.requested_logprobs);
assert!(prepared.options.include_prompt_logprobs);
assert_eq!(prepared.chat_request.sampling_params.logprobs, Some(0));
assert_eq!(
prepared.chat_request.sampling_params.logprobs,
Some(LogprobsCount::Top(0))
);
assert_eq!(
prepared.chat_request.sampling_params.prompt_logprobs,
Some(2)
Some(LogprobsCount::Top(2))
);
}
@@ -965,7 +996,7 @@ mod tests {
fn prepare_chat_request_keeps_prompt_logprobs_independent_from_echo() {
let request = ChatCompletionRequest {
logprobs: true,
top_logprobs: Some(3),
top_logprobs: Some(LogprobsCount::Top(3)),
echo: true,
..base_request()
};
@@ -977,7 +1008,10 @@ mod tests {
)
.expect("request is valid");
assert_eq!(prepared.chat_request.sampling_params.logprobs, Some(3));
assert_eq!(
prepared.chat_request.sampling_params.logprobs,
Some(LogprobsCount::Top(3))
);
assert_eq!(prepared.chat_request.sampling_params.prompt_logprobs, None);
assert!(!prepared.options.include_prompt_logprobs);
}
@@ -6,6 +6,7 @@ use serde_json::Value;
use serde_with::SerializeDisplay;
use validator::Validate;
use vllm_chat::ReasoningEffort;
use vllm_engine_core_client::protocol::LogprobsCount;
use crate::routes::openai::utils::structured_outputs::ResponseFormat;
use crate::routes::openai::utils::types::{
@@ -44,10 +45,8 @@ pub struct ChatCompletionRequest {
#[serde(default)]
pub logprobs: bool,
/// An integer specifying the number of most likely tokens to return
/// -1 means return all
#[validate(range(min = -1))]
pub top_logprobs: Option<i32>,
/// Number of most likely tokens to return. `-1` means return full vocab.
pub top_logprobs: Option<LogprobsCount>,
/// Deprecated: Replaced by max_completion_tokens
#[deprecated(note = "Use max_completion_tokens instead")]
@@ -155,8 +154,8 @@ pub struct ChatCompletionRequest {
/// Truncate prompt tokens to this length
pub truncate_prompt_tokens: Option<i64>,
/// Number of prompt logprobs to return
pub prompt_logprobs: Option<i32>,
/// Number of prompt logprobs to return. `-1` means return full vocab.
pub prompt_logprobs: Option<LogprobsCount>,
/// Restrict output to these token IDs only
pub allowed_token_ids: Option<Vec<u32>>,
@@ -165,8 +164,10 @@ pub struct ChatCompletionRequest {
pub bad_words: Option<Vec<String>>,
// -------- Extra vLLM Parameters --------
/// Token budget for reasoning/thinking
pub thinking_token_budget: Option<u32>,
/// Token budget for reasoning/thinking. Accepts a non-negative integer, or
/// `-1` for unlimited (mirroring the Python frontend, which normalizes `-1`
/// to "no budget").
pub thinking_token_budget: Option<i64>,
/// Whether to include reasoning content in the response
#[serde(default = "default_true")]
@@ -1,6 +1,7 @@
use super::types::ChatCompletionRequest;
use crate::error::{ApiError, bail_invalid_request};
use crate::routes::openai::utils::types::{ChatMessage, Tool, ToolChoice, ToolChoiceValue};
use vllm_engine_core_client::protocol::LogprobsCount;
/// Enforce the minimal compatibility contract for the Rust OpenAI server.
pub(super) fn validate_request_compat(
@@ -30,14 +31,12 @@ pub(super) fn validate_request_compat(
}
if let Some(prompt_logprobs) = request.prompt_logprobs {
if prompt_logprobs < 0 && prompt_logprobs != -1 {
bail_invalid_request!(
param = "prompt_logprobs",
"prompt_logprobs must be a non-negative value or -1."
);
}
if request.stream && (prompt_logprobs > 0 || prompt_logprobs == -1) {
if request.stream
&& matches!(
prompt_logprobs,
LogprobsCount::All | LogprobsCount::Top(1..)
)
{
bail_invalid_request!(
param = "prompt_logprobs",
"prompt_logprobs are not available when stream=true."
@@ -108,11 +107,6 @@ pub(super) fn validate_request_compat(
"truncate_prompt_tokens",
"truncate_prompt_tokens is not supported.",
)?;
reject_non_default(
request.thinking_token_budget.as_ref(),
"thinking_token_budget",
"thinking_token_budget is not supported.",
)?;
reject_non_default(
request.media_io_kwargs.as_ref(),
"media_io_kwargs",
@@ -159,6 +153,7 @@ mod tests {
use serde_json::json;
use vllm_chat::ReasoningEffort;
use vllm_engine_core_client::protocol::LogprobsCount;
use super::validate_request_compat;
use crate::routes::openai::chat_completions::types::ChatCompletionRequest;
@@ -304,7 +299,7 @@ mod tests {
#[test]
fn validate_request_compat_rejects_top_logprobs_without_logprobs() {
let request = ChatCompletionRequest {
top_logprobs: Some(0),
top_logprobs: Some(LogprobsCount::Top(0)),
..base_request()
};
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
@@ -313,26 +308,26 @@ mod tests {
#[test]
fn validate_request_compat_rejects_streaming_prompt_logprobs_requests() {
let request = ChatCompletionRequest {
prompt_logprobs: Some(1),
prompt_logprobs: Some(LogprobsCount::Top(1)),
..base_request()
};
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
let request = ChatCompletionRequest {
prompt_logprobs: Some(-1),
prompt_logprobs: Some(LogprobsCount::All),
..base_request()
};
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
}
#[test]
fn validate_request_compat_rejects_invalid_prompt_logprobs_value() {
let request = ChatCompletionRequest {
stream: false,
prompt_logprobs: Some(-2),
..base_request()
};
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
fn chat_request_deserialization_rejects_invalid_prompt_logprobs_value() {
let result = serde_json::from_value::<ChatCompletionRequest>(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"messages": [{"role": "user", "content": "hello"}],
"prompt_logprobs": -2
}));
assert!(result.is_err());
}
#[test]
@@ -1,3 +1,4 @@
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_text::{SamplingParams, TextDecodeOptions, TextRequest};
use super::types::CompletionRequest;
@@ -61,15 +62,7 @@ pub(super) fn prepare_completion_request(
.map(|request| request.lora_name.clone())
.unwrap_or_else(|| lora_resolution.model_names.first().cloned().unwrap_or_default());
let logprobs = match request.logprobs {
Some(logprobs) => Some(i32::try_from(logprobs).map_err(|_| {
ApiError::invalid_request(
"`logprobs` must fit within a signed 32-bit integer.".to_string(),
Some("logprobs"),
)
})?),
None => None,
};
let logprobs = request.logprobs.map(LogprobsCount::Top);
let prompt_only = request.echo && request.max_tokens == Some(0);
let prompt_logprobs =
request.prompt_logprobs.or(if request.echo && (!request.stream || prompt_only) {
@@ -108,6 +101,7 @@ pub(super) fn prepare_completion_request(
seed: request.seed,
max_tokens,
min_tokens: request.min_tokens,
thinking_token_budget: request.thinking_token_budget,
logprobs,
prompt_logprobs,
min_p: request.min_p,
@@ -162,6 +156,7 @@ pub(super) fn prepare_completion_request(
mod tests {
use axum::http::HeaderMap;
use serde_json::json;
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_text::Prompt;
use super::prepare_completion_request;
@@ -246,7 +241,10 @@ mod tests {
Prompt::TokenIds(vec![11, 22, 33])
);
assert_eq!(prepared.text_request.sampling_params.max_tokens, Some(7));
assert_eq!(prepared.text_request.sampling_params.logprobs, Some(2));
assert_eq!(
prepared.text_request.sampling_params.logprobs,
Some(LogprobsCount::Top(2))
);
assert_eq!(prepared.text_request.sampling_params.top_p, Some(0.9));
assert_eq!(prepared.text_request.sampling_params.top_k, Some(42));
assert_eq!(prepared.text_request.sampling_params.min_p, Some(0.1));
@@ -266,6 +264,34 @@ mod tests {
assert!(!prepared.text_request.decode_options.skip_special_tokens);
}
#[test]
fn prepare_completion_request_passes_through_thinking_token_budget() {
let prepare = |budget: serde_json::Value| {
let request: CompletionRequest = serde_json::from_value(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"prompt": "hello",
"thinking_token_budget": budget,
}))
.expect("parse request");
prepare_completion_request(
request,
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
ResolvedRequestContext::default(),
)
.expect("prepare")
.text_request
.sampling_params
.thinking_token_budget
};
// The convert layer forwards the raw value verbatim (including the `-1`
// "unlimited" sentinel); normalization/validation happens during
// lowering (see `vllm_text::lower`).
assert_eq!(prepare(json!(64)), Some(64));
assert_eq!(prepare(json!(-1)), Some(-1));
assert_eq!(prepare(json!(null)), None);
}
#[test]
fn prepare_completion_request_maps_stream_usage_and_token_format_options() {
let request: CompletionRequest = serde_json::from_value(json!({
@@ -381,10 +407,13 @@ mod tests {
.expect("prepare");
assert!(prepared.options.prompt_only);
assert_eq!(prepared.text_request.sampling_params.logprobs, Some(3));
assert_eq!(
prepared.text_request.sampling_params.logprobs,
Some(LogprobsCount::Top(3))
);
assert_eq!(
prepared.text_request.sampling_params.prompt_logprobs,
Some(3)
Some(LogprobsCount::Top(3))
);
}
@@ -406,10 +435,13 @@ mod tests {
)
.expect("prepare");
assert_eq!(prepared.text_request.sampling_params.logprobs, Some(3));
assert_eq!(
prepared.text_request.sampling_params.logprobs,
Some(LogprobsCount::Top(3))
);
assert_eq!(
prepared.text_request.sampling_params.prompt_logprobs,
Some(3)
Some(LogprobsCount::Top(3))
);
}
@@ -450,10 +482,13 @@ mod tests {
ResolvedRequestContext::default(),
)
.expect("prepare");
assert_eq!(prepared.text_request.sampling_params.logprobs, Some(1));
assert_eq!(
prepared.text_request.sampling_params.logprobs,
Some(LogprobsCount::Top(1))
);
assert_eq!(
prepared.text_request.sampling_params.prompt_logprobs,
Some(2)
Some(LogprobsCount::Top(2))
);
}
@@ -3,6 +3,7 @@ use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use validator::Validate;
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_text::Prompt;
use crate::routes::openai::utils::types::{
@@ -131,8 +132,8 @@ pub struct CompletionRequest {
/// Restrict output to these token IDs only
pub allowed_token_ids: Option<Vec<u32>>,
/// Number of prompt logprobs to return
pub prompt_logprobs: Option<i32>,
/// Number of prompt logprobs to return. `-1` means return full vocab.
pub prompt_logprobs: Option<LogprobsCount>,
// -------- Extra vLLM Parameters --------
/// Whether to add special tokens (e.g. BOS) to the prompt
@@ -146,6 +147,11 @@ pub struct CompletionRequest {
/// Additional kwargs for structured outputs
pub structured_outputs: Option<Value>,
/// Token budget for reasoning/thinking. Accepts a non-negative integer, or
/// `-1` for unlimited (mirroring the Python frontend, which normalizes `-1`
/// to "no budget").
pub thinking_token_budget: Option<i64>,
/// Request scheduling priority (lower means earlier; default 0)
pub priority: Option<i32>,
@@ -1,3 +1,4 @@
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_text::Prompt;
use super::types::CompletionRequest;
@@ -44,29 +45,18 @@ pub(super) fn validate_request_compat(
bail_invalid_request!(param = "suffix", "suffix is not supported.");
}
if let Some(logprobs) = request.logprobs
&& logprobs > i32::MAX as u32
{
bail_invalid_request!(
param = "logprobs",
"`logprobs` must fit within a signed 32-bit integer."
);
}
if let Some(prompt_logprobs) = request.prompt_logprobs {
if request.stream && (prompt_logprobs > 0 || prompt_logprobs == -1) {
if request.stream
&& matches!(
prompt_logprobs,
LogprobsCount::All | LogprobsCount::Top(1..)
)
{
bail_invalid_request!(
param = "prompt_logprobs",
"`prompt_logprobs` are not available when `stream=true`."
);
}
if prompt_logprobs < 0 && prompt_logprobs != -1 {
bail_invalid_request!(
param = "prompt_logprobs",
"`prompt_logprobs` must be a non-negative value or -1."
);
}
}
if request.use_beam_search {
@@ -101,6 +91,7 @@ pub(super) fn validate_request_compat(
#[cfg(test)]
mod tests {
use serde_json::json;
use vllm_engine_core_client::protocol::LogprobsCount;
use super::validate_request_compat;
use crate::routes::openai::completions::types::CompletionRequest;
@@ -150,7 +141,7 @@ mod tests {
#[test]
fn validate_request_compat_rejects_streaming_prompt_logprobs() {
let request = CompletionRequest {
prompt_logprobs: Some(1),
prompt_logprobs: Some(LogprobsCount::Top(1)),
..base_request()
};
assert!(
@@ -162,7 +153,7 @@ mod tests {
fn validate_request_compat_accepts_non_stream_prompt_logprobs() {
let request = CompletionRequest {
stream: false,
prompt_logprobs: Some(-1),
prompt_logprobs: Some(LogprobsCount::All),
..base_request()
};
assert!(
+3 -4
View File
@@ -2,6 +2,7 @@ pub mod hf;
use std::sync::Arc;
use vllm_engine_core_client::protocol::LogprobsCount;
use vllm_tokenizer::DynTokenizer;
use crate::error::Result;
@@ -26,9 +27,7 @@ pub struct SamplingLimits {
/// Runtime context window size reported by the engine startup handshake.
pub max_model_len: u32,
/// Maximum number of top log probabilities accepted by this frontend.
///
/// `-1` means allowing requests up to the model vocabulary size.
pub max_logprobs: i32,
pub max_logprobs: LogprobsCount,
/// Model vocabulary size from the model config, used to bound generated
/// token IDs and logits-domain sampling controls.
@@ -41,7 +40,7 @@ pub struct SamplingLimits {
impl SamplingLimits {
/// Original Python definition:
/// <https://github.com/vllm-project/vllm/blob/b5adb027ad03c29b46181752ba3b1cb84eff1dd4/vllm/config/model.py#L216-L220>
pub const DEFAULT_MAX_LOGPROBS: i32 = 20;
pub const DEFAULT_MAX_LOGPROBS: LogprobsCount = LogprobsCount::Top(20);
/// Original Python definition:
/// <https://github.com/vllm-project/vllm/blob/b5adb027ad03c29b46181752ba3b1cb84eff1dd4/vllm/sampling_params.py#L30-L32>
pub const MAX_LOGPROB_TOKEN_IDS: usize = 128;
+2
View File
@@ -20,6 +20,8 @@ pub enum Error {
Logprobs(#[from] LogprobsError),
#[error(transparent)]
OutOfVocab(#[from] OutOfVocabError),
#[error("`thinking_token_budget` must be a non-negative integer or -1 for unlimited.")]
InvalidThinkingTokenBudget,
#[error("text request stream `{request_id}` closed before terminal output")]
StreamClosedBeforeTerminalOutput { request_id: String },
#[error(transparent)]
+3 -2
View File
@@ -19,6 +19,7 @@ pub use output::{
pub use request::{Prompt, SamplingParams, TextRequest};
use trait_set::trait_set;
use vllm_engine_core_client::EngineCoreClient;
use vllm_engine_core_client::protocol::LogprobsCount;
pub use vllm_llm::FinishReason;
use vllm_llm::{GenerateOutputStream, Llm};
use vllm_tokenizer::DynTokenizer;
@@ -48,7 +49,7 @@ pub struct TextLlm {
/// Runtime context window size reported by the engine startup handshake.
max_model_len: u32,
/// Maximum number of top log probabilities accepted by this text facade.
max_logprobs: i32,
max_logprobs: LogprobsCount,
}
impl TextLlm {
@@ -68,7 +69,7 @@ impl TextLlm {
}
/// Override the maximum accepted logprobs count.
pub fn with_max_logprobs(mut self, max_logprobs: Option<i32>) -> Self {
pub fn with_max_logprobs(mut self, max_logprobs: Option<LogprobsCount>) -> Self {
if let Some(max_logprobs) = max_logprobs {
self.max_logprobs = max_logprobs;
}
+65 -10
View File
@@ -87,6 +87,7 @@ pub fn lower_sampling_params(
seed,
max_tokens,
min_tokens,
thinking_token_budget,
logprobs,
prompt_logprobs,
min_p,
@@ -128,6 +129,7 @@ pub fn lower_sampling_params(
prompt_len,
)?;
let min_tokens = min_tokens.unwrap_or(0);
let thinking_token_budget = normalize_thinking_token_budget(thinking_token_budget)?;
let frequency_penalty = frequency_penalty.unwrap_or(0.0);
let presence_penalty = presence_penalty.unwrap_or(0.0);
@@ -149,6 +151,7 @@ pub fn lower_sampling_params(
seed,
max_tokens,
min_tokens,
thinking_token_budget,
logprobs,
prompt_logprobs,
min_p,
@@ -170,6 +173,21 @@ pub fn lower_sampling_params(
Ok(params)
}
/// Normalize the user-facing `thinking_token_budget` into the engine value.
///
/// Mirrors Python's `validate_thinking_token_budget`
/// (<https://github.com/vllm-project/vllm/blob/ecf9d83520eb217401b47d8a5451a27c5231b8c2/vllm/sampling_params.py#L35-L55>):
/// `None` and the `-1` "unlimited" sentinel both map to `None`; any other
/// negative value is rejected; non-negative values pass through unchanged. Like
/// Python's `int`, no upper bound is imposed.
fn normalize_thinking_token_budget(value: Option<i64>) -> Result<Option<u64>> {
match value {
None | Some(-1) => Ok(None),
Some(budget) if budget >= 0 => Ok(Some(budget as u64)),
Some(_) => Err(Error::InvalidThinkingTokenBudget),
}
}
/// Convert bad-word strings into token-ID sequences, following the Python vLLM
/// logic in `SamplingParams.update_from_tokenizer()`.
///
@@ -251,6 +269,7 @@ mod tests {
use std::collections::{BTreeSet, HashMap};
use serial_test::file_serial;
use vllm_engine_core_client::protocol::LogprobsCount;
use super::*;
use crate::backend::hf::HfTextBackend;
@@ -366,6 +385,36 @@ mod tests {
)
}
#[test]
fn lower_sampling_params_normalizes_thinking_token_budget() {
let lower = |budget: Option<i64>| {
lower_sampling_params_with_limits(
SamplingParams {
thinking_token_budget: budget,
..SamplingParams::default()
},
sample_sampling_limits(),
)
};
// Non-negative budgets (including 0) pass through unchanged.
assert_eq!(lower(Some(256)).unwrap().thinking_token_budget, Some(256));
assert_eq!(lower(Some(0)).unwrap().thinking_token_budget, Some(0));
// `None` and the `-1` "unlimited" sentinel both disable the budget.
assert_eq!(lower(None).unwrap().thinking_token_budget, None);
assert_eq!(lower(Some(-1)).unwrap().thinking_token_budget, None);
// No upper bound is imposed, matching Python's `int`.
assert_eq!(
lower(Some(i64::from(u32::MAX) + 1)).unwrap().thinking_token_budget,
Some(u64::from(u32::MAX) + 1)
);
// Other negatives are rejected.
assert!(matches!(
lower(Some(-2)),
Err(Error::InvalidThinkingTokenBudget)
));
}
#[test]
fn lower_text_request_applies_python_style_eos_hints() {
let prepared = lower_text_request(
@@ -386,6 +435,7 @@ mod tests {
seed: None,
max_tokens: 999997,
min_tokens: 0,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.0,
@@ -437,6 +487,7 @@ mod tests {
seed: None,
max_tokens: 999997,
min_tokens: 0,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.0,
@@ -567,6 +618,7 @@ mod tests {
seed: None,
max_tokens: 40957,
min_tokens: 0,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.0,
@@ -628,6 +680,7 @@ mod tests {
seed: None,
max_tokens: 999997,
min_tokens: 0,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.0,
@@ -697,6 +750,7 @@ mod tests {
seed: None,
max_tokens: 32,
min_tokens: 2,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.1,
@@ -721,8 +775,8 @@ mod tests {
#[test]
fn lower_sampling_params_passes_logprobs_fields_through() {
let sampling_params = SamplingParams {
logprobs: Some(3),
prompt_logprobs: Some(-1),
logprobs: Some(LogprobsCount::Top(3)),
prompt_logprobs: Some(LogprobsCount::All),
..Default::default()
};
@@ -739,7 +793,7 @@ mod tests {
default_max_tokens: None,
},
SamplingLimits {
max_logprobs: -1,
max_logprobs: LogprobsCount::All,
..sample_sampling_limits()
},
3,
@@ -747,15 +801,15 @@ mod tests {
)
.unwrap();
assert_eq!(params.logprobs, Some(3));
assert_eq!(params.prompt_logprobs, Some(-1));
assert_eq!(params.logprobs, Some(LogprobsCount::Top(3)));
assert_eq!(params.prompt_logprobs, Some(LogprobsCount::All));
}
#[test]
fn lower_sampling_params_rejects_full_vocab_logprobs_over_default_cap() {
let error = lower_sampling_params_with_limits(
SamplingParams {
logprobs: Some(-1),
logprobs: Some(LogprobsCount::All),
..Default::default()
},
sample_sampling_limits(),
@@ -776,24 +830,24 @@ mod tests {
fn lower_sampling_params_expands_full_vocab_logprobs_from_model_vocab() {
let params = lower_sampling_params_with_limits(
SamplingParams {
logprobs: Some(-1),
logprobs: Some(LogprobsCount::All),
..Default::default()
},
SamplingLimits {
max_logprobs: 1500,
max_logprobs: LogprobsCount::Top(1500),
..sample_sampling_limits()
},
)
.unwrap();
assert_eq!(params.logprobs, Some(-1));
assert_eq!(params.logprobs, Some(LogprobsCount::All));
}
#[test]
fn lower_sampling_params_rejects_invalid_logprob_token_ids() {
let error = lower_sampling_params_with_limits(
SamplingParams {
logprobs: Some(1),
logprobs: Some(LogprobsCount::Top(1)),
logprob_token_ids: Some(vec![1000]),
..Default::default()
},
@@ -929,6 +983,7 @@ mod tests {
seed: None,
max_tokens: 128,
min_tokens: 0,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: 0.1,
+13 -24
View File
@@ -1,15 +1,14 @@
//! Python-compatible validation for logprobs sampling params.
//!
//! `-1` is expanded only for bounds checks. The original request values are
//! `All` is expanded only for bounds checks. The original request values are
//! passed through to engine-core.
use crate::backend::SamplingLimits;
use thiserror::Error;
use vllm_engine_core_client::protocol::LogprobsCount;
#[derive(Debug, Error)]
pub enum LogprobsError {
#[error("{parameter} must be non-negative or -1, got {value}")]
InvalidCount { parameter: &'static str, value: i32 },
#[error(
"requested {parameter} of {requested}, which is greater than max allowed: {max_allowed}"
)]
@@ -30,19 +29,21 @@ pub enum LogprobsError {
"when both logprobs and logprob_token_ids are set, logprobs must equal \
len(logprob_token_ids). Got logprobs={logprobs}, len(logprob_token_ids)={num_token_ids}."
)]
TokenIdsMismatch { logprobs: i32, num_token_ids: usize },
TokenIdsMismatch {
logprobs: LogprobsCount,
num_token_ids: usize,
},
}
/// Validate logprobs count sampling parameters.
pub(super) fn validate_logprobs(
logprobs: Option<i32>,
prompt_logprobs: Option<i32>,
logprobs: Option<LogprobsCount>,
prompt_logprobs: Option<LogprobsCount>,
logprob_token_ids: Option<&[u32]>,
sampling_limits: SamplingLimits,
) -> Result<(), LogprobsError> {
let vocab_size = sampling_limits.model_vocab_size;
let max_logprobs =
normalize_logprobs_count(sampling_limits.max_logprobs, vocab_size, "max_logprobs")?;
let max_logprobs = sampling_limits.max_logprobs.expanded(vocab_size);
validate_logprobs_count(logprobs, max_logprobs, vocab_size, "logprobs")?;
validate_logprobs_count(prompt_logprobs, max_logprobs, vocab_size, "prompt_logprobs")?;
@@ -50,7 +51,7 @@ pub(super) fn validate_logprobs(
}
fn validate_logprobs_count(
requested: Option<i32>,
requested: Option<LogprobsCount>,
max_logprobs: usize,
vocab_size: usize,
parameter: &'static str,
@@ -59,7 +60,7 @@ fn validate_logprobs_count(
return Ok(());
};
let requested = normalize_logprobs_count(requested, vocab_size, parameter)?;
let requested = requested.expanded(vocab_size);
if requested > max_logprobs {
return Err(LogprobsError::TooManyCount {
parameter,
@@ -72,7 +73,7 @@ fn validate_logprobs_count(
}
pub(super) fn validate_logprob_token_ids(
logprobs: Option<i32>,
logprobs: Option<LogprobsCount>,
logprob_token_ids: Option<&[u32]>,
) -> Result<(), LogprobsError> {
let Some(logprob_token_ids) = logprob_token_ids else {
@@ -88,7 +89,7 @@ pub(super) fn validate_logprob_token_ids(
}
if let Some(logprobs) = logprobs
&& logprobs != n as i32
&& logprobs != LogprobsCount::Top(n as u32)
{
return Err(LogprobsError::TokenIdsMismatch {
logprobs,
@@ -98,15 +99,3 @@ pub(super) fn validate_logprob_token_ids(
Ok(())
}
fn normalize_logprobs_count(
value: i32,
vocab_size: usize,
parameter: &'static str,
) -> Result<usize, LogprobsError> {
match value {
-1 => Ok(vocab_size),
value if value < 0 => Err(LogprobsError::InvalidCount { parameter, value }),
value => Ok(value as usize),
}
}
+8 -1
View File
@@ -309,7 +309,7 @@ fn matches_stop_string(stops: &[String], output: &str, new_bytes: usize) -> Opti
.find_map(|(ss_idx, (ss, len, start_off))| {
output[start_off..]
.windows(len)
.rposition(|w| w == ss)
.position(|w| w == ss)
.map(|pos| (ss_idx, start_off + pos))
})
}
@@ -562,6 +562,13 @@ mod tests {
assert_eq!(result, Some((0, 4)));
}
#[test]
fn stop_string_matches_leftmost_with_multiple_new_bytes() {
let stops = vec!["\n".to_string()];
let result = matches_stop_string(&stops, "Answer\n\n", 2);
assert_eq!(result, Some((0, 6)));
}
#[test]
fn stop_string_matches_at_beginning() {
let stops = vec!["say".to_string()];
+12 -5
View File
@@ -3,9 +3,9 @@ use std::collections::HashMap;
use enum_as_inner::EnumAsInner;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use vllm_engine_core_client::protocol::StructuredOutputsParams;
use vllm_engine_core_client::protocol::lora::LoraRequest;
use vllm_engine_core_client::protocol::multimodal::MmFeatures;
use vllm_engine_core_client::protocol::{LogprobsCount, StructuredOutputsParams};
use crate::error::{Error, Result};
use crate::output::TextDecodeOptions;
@@ -56,14 +56,20 @@ pub struct SamplingParams {
pub max_tokens: Option<u32>,
/// Minimum number of tokens to generate before EOS or stop-token handling.
pub min_tokens: Option<u32>,
/// Maximum number of reasoning ("thinking") tokens to emit before the
/// reasoning section is force-closed. `None` or the user-facing `-1`
/// "unlimited" sentinel both disable the budget. The raw value is carried
/// here; `-1` is normalized to `None` (and other negatives rejected) during
/// lowering (see `lower_sampling_params`).
pub thinking_token_budget: Option<i64>,
/// Number of log probabilities to return per generated token.
///
/// `None` disables sample logprobs. `-1` requests the full vocabulary.
pub logprobs: Option<i32>,
/// `None` disables sample logprobs.
pub logprobs: Option<LogprobsCount>,
/// Number of log probabilities to return per prompt token.
///
/// `None` disables prompt logprobs. `-1` requests the full vocabulary.
pub prompt_logprobs: Option<i32>,
/// `None` disables prompt logprobs.
pub prompt_logprobs: Option<LogprobsCount>,
/// Minimum probability threshold for token sampling. `None` means no
/// explicit user override.
pub min_p: Option<f32>,
@@ -116,6 +122,7 @@ impl Default for SamplingParams {
seed: None,
max_tokens: None,
min_tokens: None,
thinking_token_budget: None,
logprobs: None,
prompt_logprobs: None,
min_p: None,
@@ -22,7 +22,7 @@ import torch
import vllm.config
from tests.compile.backend import TestBackend
from vllm._aiter_ops import is_aiter_found_and_supported, rocm_aiter_ops
from vllm._aiter_ops import rocm_aiter_ops
from vllm.compilation.passes.utility.noop_elimination import NoOpEliminationPass
from vllm.compilation.passes.utility.post_cleanup import PostCleanupPass
from vllm.config import (
@@ -83,9 +83,8 @@ class _ViewDoubleQuantModel(torch.nn.Module):
[_NoViewDoubleQuantModel, _ViewDoubleQuantModel],
ids=["no_view", "with_view"],
)
@pytest.mark.skipif(
not is_aiter_found_and_supported(),
reason="Only test on ROCm with AITER installed and supported",
@pytest.mark.skip(
reason="Skipping for now because pytorch compiler removes one the two quant ops"
)
def test_double_aiter_rms_fp8_group_quant_fusion(
model_cls: type[torch.nn.Module],
+2 -2
View File
@@ -175,7 +175,7 @@ MULTIMODAL_MODELS = {
"facebook/chameleon-7b": PPTestSettings.fast(),
"adept/fuyu-8b": PPTestSettings.fast(),
"zai-org/glm-4v-9b": PPTestSettings.fast(),
"OpenGVLab/InternVL2-1B": PPTestSettings.fast(),
"OpenGVLab/InternVL3-1B": PPTestSettings.fast(),
"llava-hf/llava-1.5-7b-hf": PPTestSettings.fast(),
"llava-hf/llava-v1.6-mistral-7b-hf": PPTestSettings.fast(),
"llava-hf/LLaVA-NeXT-Video-7B-hf": PPTestSettings.fast(),
@@ -203,7 +203,7 @@ TEST_MODELS = [
"intfloat/e5-mistral-7b-instruct",
"BAAI/bge-multilingual-gemma2",
# [MULTIMODAL GENERATION]
"OpenGVLab/InternVL2-1B",
"OpenGVLab/InternVL3-1B",
"microsoft/Phi-3.5-vision-instruct",
"fixie-ai/ultravox-v0_5-llama-3_2-1b",
# [LANGUAGE GENERATION - HYBRID ARCH]
@@ -25,7 +25,7 @@ def server():
"--runner",
"pooling",
"--max-model-len",
"5000",
"16384",
"--enforce-eager",
"--limit-mm-per-prompt",
json.dumps({"video": MAXIMUM_VIDEOS}),
@@ -143,4 +143,4 @@ def test_chat_video_url_request(server: RemoteOpenAIServer, model_name: str):
assert output.model == model_name
assert len(output.data) == 1
assert len(output.data[0].probs) == 2
assert output.usage.prompt_tokens == 4807
assert output.usage.prompt_tokens == 8993
@@ -15,7 +15,7 @@ from vllm.entrypoints.serve.tokenize.protocol import (
TokenizeChatRequest,
TokenizeCompletionRequest,
)
from vllm.entrypoints.serve.tokenize.serving import OpenAIServingTokenization
from vllm.entrypoints.serve.tokenize.serving import ServingTokenization
from vllm.v1.engine.async_llm import AsyncLLM
MODEL_NAME = "openai-community/gpt2"
@@ -58,7 +58,7 @@ class MockModelConfig:
return self.diff_sampling_param or {}
def _build_serving_tokenization(engine: AsyncLLM) -> OpenAIServingTokenization:
def _build_serving_tokenization(engine: AsyncLLM) -> ServingTokenization:
models = OpenAIServingModels(
engine_client=engine,
base_model_paths=BASE_MODEL_PATHS,
@@ -71,8 +71,7 @@ def _build_serving_tokenization(engine: AsyncLLM) -> OpenAIServingTokenization:
chat_template=None,
chat_template_content_format="auto",
)
return OpenAIServingTokenization(
engine,
return ServingTokenization(
models,
openai_serving_render=serving_render,
request_logger=None,
@@ -15,16 +15,14 @@ from vllm.config import (
from vllm.platforms import current_platform
from vllm.platforms.cpu import CpuPlatform
# CudaPlatform and RocmPlatform import their respective compiled C extensions
# at module level, raising ModuleNotFoundError on incompatible builds.
try:
if current_platform.is_cuda():
from vllm.platforms.cuda import CudaPlatform
except (ImportError, ModuleNotFoundError):
else:
CudaPlatform = None
try:
if current_platform.is_rocm():
from vllm.platforms.rocm import RocmPlatform
except (ImportError, ModuleNotFoundError):
else:
RocmPlatform = None
from vllm.v1.attention.backends.registry import AttentionBackendEnum
@@ -434,9 +432,15 @@ def test_per_head_quant_scales_backend_selection(
[
("FLASH_ATTN", True, True), # FlashAttn supports non-causal
("FLASH_ATTN", False, True), # FlashAttn also works with causal
("FLASHINFER", True, False), # FlashInfer does not support non-causal
("FLASHINFER", False, True), # FlashInfer works with causal
],
]
+ (
[
("FLASHINFER", True, True), # FlashInfer supports non-causal
("FLASHINFER", False, True), # FlashInfer works with causal
]
if CudaPlatform is not None
else []
),
)
def test_non_causal_backend_selection(
backend_name: str, use_non_causal: bool, should_succeed: bool
@@ -459,11 +463,12 @@ def test_non_causal_backend_selection(
attention_config=attention_config, cache_config=cache_config
)
if CudaPlatform is None:
pytest.skip("CudaPlatform not available")
platform = CudaPlatform or RocmPlatform
if platform is None:
pytest.skip("CudaPlatform and RocmPlatform are not available")
with (
set_current_vllm_config(vllm_config),
patch("vllm.platforms.current_platform", CudaPlatform()),
patch("vllm.platforms.current_platform", platform()),
):
if should_succeed:
backend = get_attn_backend(
+17 -9
View File
@@ -5,10 +5,12 @@ import math
import random
import time
from collections.abc import Callable
from contextlib import nullcontext
import pytest
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
from vllm.platforms import current_platform
from vllm.utils.torch_utils import STR_DTYPE_TO_TORCH_DTYPE, set_random_seed
@@ -557,15 +559,21 @@ def test_contexted_kv_attention_alibi(
query_len, seq_len, alibi_slopes, device, dtype
)
# Compute attention
out = F.scaled_dot_product_attention(
q_sdpa,
k_sdpa,
v_sdpa,
attn_mask=alibi_mask,
dropout_p=0.0,
scale=scale,
)
# Compute attention. On ROCm we force use of the Math SDPA backend rather than
# the Flash or Mem-Efficient backends for increased numerical accuracy
if current_platform.is_rocm():
sdpa_context = sdpa_kernel(SDPBackend.MATH)
else:
sdpa_context = nullcontext()
with sdpa_context:
out = F.scaled_dot_product_attention(
q_sdpa,
k_sdpa,
v_sdpa,
attn_mask=alibi_mask,
dropout_p=0.0,
scale=scale,
)
# Reshape output back to [query_len, num_heads, head_size]
out = out.view(num_heads, query_len, head_size).permute(1, 0, 2)
@@ -90,7 +90,9 @@ def _ref_sparse_prefill_ragged(
return out.to(torch.bfloat16)
def _pack_fp8_ds_mla_cache(kv: torch.Tensor, block_size: int) -> torch.Tensor:
def _pack_fp8_ds_mla_cache(
kv: torch.Tensor, block_size: int, is_extra: bool = False
) -> torch.Tensor:
assert kv.shape[-1] == HEAD_DIM
num_tokens = kv.shape[0]
num_blocks = (num_tokens + block_size - 1) // block_size
@@ -101,7 +103,9 @@ def _pack_fp8_ds_mla_cache(kv: torch.Tensor, block_size: int) -> torch.Tensor:
)
cache_flat = cache.view(torch.uint8).flatten()
kv_nope_fp8 = (
kv[:, :NOPE_HEAD_DIM].to(current_platform.fp8_dtype()).view(torch.uint8)
kv[:, :NOPE_HEAD_DIM]
.to(torch.float8_e4m3fn if is_extra else current_platform.fp8_dtype())
.view(torch.uint8)
)
kv_rope_u8 = kv[:, NOPE_HEAD_DIM:].contiguous().view(torch.uint8)
@@ -120,7 +124,7 @@ def _pack_fp8_ds_mla_cache(kv: torch.Tensor, block_size: int) -> torch.Tensor:
def _read_fp8_ds_mla_cache(
cache: torch.Tensor, slot: int, block_size: int
cache: torch.Tensor, slot: int, block_size: int, is_extra: bool = False
) -> torch.Tensor:
cache_flat = cache.view(torch.uint8).flatten()
block_idx = slot // block_size
@@ -129,7 +133,9 @@ def _read_fp8_ds_mla_cache(
token_base = block_base + pos * 576
nope_u8 = cache_flat[token_base : token_base + NOPE_HEAD_DIM]
nope = nope_u8.view(current_platform.fp8_dtype()).to(torch.float32)
nope = nope_u8.view(
torch.float8_e4m3fn if is_extra else current_platform.fp8_dtype()
).to(torch.float32)
rope_u8 = cache_flat[
token_base + NOPE_HEAD_DIM : token_base + NOPE_HEAD_DIM + ROPE_HEAD_DIM * 2
]
@@ -157,7 +163,9 @@ def _ref_sparse_decode_ragged(
]
if extra_cache is not None and extra_rows is not None:
row_kv.extend(
_read_fp8_ds_mla_cache(extra_cache, int(slot), block_size)
_read_fp8_ds_mla_cache(
extra_cache, int(slot), block_size, is_extra=True
)
for slot in extra_rows[query_idx]
)
@@ -326,7 +334,7 @@ def test_sparse_attn_decode_ragged_kernel() -> None:
main_kv = torch.randn(6, HEAD_DIM, dtype=torch.bfloat16, device=device) * 0.125
extra_kv = torch.randn(5, HEAD_DIM, dtype=torch.bfloat16, device=device) * 0.125
main_cache = _pack_fp8_ds_mla_cache(main_kv, block_size)
extra_cache = _pack_fp8_ds_mla_cache(extra_kv, block_size)
extra_cache = _pack_fp8_ds_mla_cache(extra_kv, block_size, is_extra=True)
main_indices = torch.tensor([0, 2, 4, 1], dtype=torch.int32, device=device)
main_indptr = torch.tensor([0, 2, 4], dtype=torch.int32, device=device)
extra_indices = torch.tensor([1, 3, 0], dtype=torch.int32, device=device)
@@ -477,7 +485,7 @@ def test_sparse_attn_decode_split_k_kernel(
rows = [[1, 3, 0, 5, 2, 4], [3, 0, 6]]
extra_kv = torch.randn(7, HEAD_DIM, dtype=torch.bfloat16, device=device) * 0.125
extra_rows = rows
extra_cache = _pack_fp8_ds_mla_cache(extra_kv, block_size)
extra_cache = _pack_fp8_ds_mla_cache(extra_kv, block_size, is_extra=True)
extra_indices, extra_indptr = _ragged_from_rows(rows, device)
attn_sink = (
@@ -18,11 +18,7 @@ HEAD_SIZES = [128, 256]
BLOCK_SIZES = [16]
DTYPES = [torch.bfloat16]
QDTYPES = (
[None, torch.float8_e4m3fn]
if not current_platform.is_rocm()
else [None, torch.float8_e4m3fnuz]
)
QDTYPES = [None, current_platform.fp8_dtype()]
FP8_DTYPE = current_platform.fp8_dtype()
# one value large enough to test overflow in index calculation.
+59 -3
View File
@@ -10,8 +10,12 @@ from torch.multiprocessing import spawn
from tests.kernels.utils import opcheck
from tests.utils import ensure_current_vllm_config, init_test_distributed_environment
from vllm.distributed import cleanup_dist_env_and_memory
from vllm.model_executor.layers.minimax_rms_norm import MiniMaxText01RMSNormTP
from vllm.model_executor.layers.minimax_rms_norm import (
MiniMaxText01RMSNormTP,
rms_norm_tp,
)
from vllm.platforms import current_platform
from vllm.triton_utils import HAS_TRITON
from vllm.utils.network_utils import get_open_port
from vllm.utils.torch_utils import set_random_seed
@@ -54,8 +58,19 @@ def _worker_forward_qk(
torch.manual_seed(seed + 1000 + local_rank)
qkv = torch.randn(num_tokens, hq + hk + hk, dtype=dtype, device="cuda")
q_ref, k_ref, v_ref = qkv.clone().split([hq, hk, hk], dim=-1)
ref_q, ref_k = MiniMaxText01RMSNormTP.forward_qk(q_norm, k_norm, q_ref, k_ref)
# Reference: eager all-reduce path. ``forward_qk`` no longer all-reduces
# the variance (it is the tp==1 / already-reduced building block), so the
# multi-rank reference must use the eager path that performs the global
# variance all-reduce, matching the fused kernel below.
ref_q, ref_k = rms_norm_tp._minimax_qk_norm_tp_eager(
qkv.clone(),
q_norm.weight,
k_norm.weight,
hq,
hk,
world_size,
eps,
)
# Set up Lamport workspace.
from vllm.distributed.parallel_state import get_tp_group
@@ -150,3 +165,44 @@ def test_minimax_reduce_rms_qk(
nprocs=world_size,
join=True,
)
@pytest.mark.skipif(
not current_platform.is_cuda() or not HAS_TRITON,
reason="CUDA and Triton required",
)
@pytest.mark.parametrize("num_tokens", [1, 7, 128, 333, 2049])
@pytest.mark.parametrize("hidden_dims", [(3072, 512), (768, 256), (3000, 500)])
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize("tp_world", [1, 4, 8])
@pytest.mark.parametrize("eps", [1e-6])
@pytest.mark.parametrize("seed", [42])
def test_minimax_qk_norm_triton_fallback(
monkeypatch, num_tokens, hidden_dims, dtype, tp_world, eps, seed
):
"""Single-GPU check: Triton fallback kernels vs the pure-torch reference.
The all-reduce is a TP communication barrier, so it is monkeypatched to
identity here; both the Triton path and the reference see the same
(patched) reduction. This validates the kernel math and the folded
``/ tp_world`` scaling without needing multiple ranks -- ``hidden_dims``
are the per-rank q/k segment widths.
"""
monkeypatch.setattr(rms_norm_tp, "_all_reduce_variance", lambda v: v)
q_size, kv_size = hidden_dims
device = "cuda"
torch.manual_seed(seed)
qkv = torch.randn(num_tokens, q_size + 2 * kv_size, dtype=dtype, device=device)
q_weight = torch.randn(q_size, dtype=dtype, device=device)
k_weight = torch.randn(kv_size, dtype=dtype, device=device)
q_triton, k_triton = rms_norm_tp._minimax_qk_norm_tp_fallback(
qkv, q_weight, k_weight, q_size, kv_size, 0, tp_world, eps
)
q_ref, k_ref = rms_norm_tp._minimax_qk_norm_tp_eager(
qkv, q_weight, k_weight, q_size, kv_size, tp_world, eps
)
torch.testing.assert_close(q_triton, q_ref, atol=3e-2, rtol=3e-2)
torch.testing.assert_close(k_triton, k_ref, atol=3e-2, rtol=3e-2)
+8 -8
View File
@@ -9,6 +9,7 @@ import pytest
import torch
from packaging import version
from vllm._aiter_ops import is_aiter_found
from vllm.platforms import current_platform
from vllm.utils.flashinfer import has_flashinfer
@@ -31,17 +32,15 @@ HOPPER_MXFP4_BF16_AVAILABLE = (
# ROCm platform and dependencies
ROCM_AVAILABLE = current_platform.is_rocm()
ROCM_TRITON_KERNELS_AVAILABLE = False
ROCM_AITER_AVAILABLE = False
ROCM_AITER_AVAILABLE = is_aiter_found()
ROCM_GFX950 = False
if ROCM_AVAILABLE:
from vllm._aiter_ops import rocm_aiter_ops
from vllm.platforms.rocm import on_gfx950
from vllm.utils.import_utils import has_triton_kernels
ROCM_TRITON_KERNELS_AVAILABLE = has_triton_kernels()
ROCM_GFX950 = on_gfx950()
ROCM_AITER_AVAILABLE = rocm_aiter_ops.is_enabled()
if ROCM_AITER_AVAILABLE:
from aiter.ops.triton.moe.quant_moe import upcast_from_mxfp
@@ -83,7 +82,7 @@ def enable_pickle(monkeypatch):
[
ModelCase("fxmarty/qwen_1.5-moe-a2.7b-mxfp4", tp=2),
ModelCase("fxmarty/deepseek_r1_3_layers_mxfp4", tp=8),
ModelCase("fxmarty/Llama-4-Scout-17B-16E-Instruct-2-layers-mxfp4", tp=1),
ModelCase("mawong-amd/Llama-4-Scout-17B-16E-Instruct-2-layers-mxfp4", tp=1),
ModelCase("fxmarty/Llama-3.1-70B-Instruct-2-layers-mxfp6", tp=1),
ModelCase("fxmarty/Llama-3.1-70B-Instruct-2-layers-mxfp6", tp=4),
],
@@ -102,6 +101,7 @@ def test_mxfp4_loading_and_execution_moe(vllm_runner, model_case: ModelCase):
tensor_parallel_size=model_case.tp,
load_format="dummy",
compilation_config={"cudagraph_capture_sizes": [16]},
gpu_memory_utilization=0.8, # mxfp6 models use more scratch space
) as llm:
# Disabled as check_model is broken: https://github.com/vllm-project/vllm/pull/18465#issuecomment-3329880562
# def check_model(model):
@@ -1267,7 +1267,7 @@ def test_rocm_mxfp4_moe_oracle(
This test validates that the oracle functions work end-to-end:
- select_mxfp4_moe_backend() selects a valid backend
- convert_to_mxfp4_moe_kernel_format() converts weights without error
- convert_gpt_oss_weight_to_mxfp4_moe_kernel_format() converts weights without error
- make_mxfp4_moe_quant_config() builds a valid quant config
- make_mxfp4_moe_kernel() creates a kernel that runs without error
- The kernel output is within accuracy tolerance of reference
@@ -1287,7 +1287,7 @@ def test_rocm_mxfp4_moe_oracle(
from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
Mxfp4MoeBackend,
backend_to_kernel_cls,
convert_to_mxfp4_moe_kernel_format,
convert_gpt_oss_weight_to_mxfp4_moe_kernel_format,
make_mxfp4_moe_kernel,
make_mxfp4_moe_quant_config,
)
@@ -1387,7 +1387,7 @@ def test_rocm_mxfp4_moe_oracle(
# Convert weights using oracle
w13_conv, w2_conv, w13_scale_conv, w2_scale_conv, w13_bias_conv, w2_bias_conv = (
convert_to_mxfp4_moe_kernel_format(
convert_gpt_oss_weight_to_mxfp4_moe_kernel_format(
mxfp4_backend=backend,
layer=layer, # type: ignore[arg-type]
w13_weight=w13_quant,
@@ -1423,7 +1423,7 @@ def test_rocm_mxfp4_moe_oracle(
mxfp4_backend=backend,
experts_cls=experts_cls,
routing_tables=None,
shared_experts=None,
layer=None,
)
# Create inputs
+11 -11
View File
@@ -8,8 +8,8 @@ from vllm.model_executor.kernels.mhc.tilelang import (
_tilelang_hc_prenorm_gemm,
_torch_hc_prenorm_gemm,
)
from vllm.model_executor.layers.mhc import HAS_TILELANG_MHC
from vllm.platforms import current_platform
from vllm.utils.import_utils import has_tilelang
from vllm.utils.torch_utils import set_random_seed
DEVICE = current_platform.device_type
@@ -97,8 +97,8 @@ def hc_head_ref(
@pytest.mark.skipif(
not (current_platform.is_cuda_alike() and has_tilelang()),
reason="CUDA or ROCm and tilelang required",
not HAS_TILELANG_MHC,
reason="TileLang MHC support required",
)
@pytest.mark.parametrize("num_tokens", [1, 4, 8, 128])
@pytest.mark.parametrize("hidden_size", [4096, 7168])
@@ -150,8 +150,8 @@ def test_mhc_pre_tilelang(num_tokens, hidden_size, hc_mult):
@pytest.mark.skipif(
not (current_platform.is_cuda_alike() and has_tilelang()),
reason="CUDA or ROCm and tilelang required",
not HAS_TILELANG_MHC,
reason="TileLang MHC support required",
)
@pytest.mark.parametrize(
("num_tokens", "hidden_size"),
@@ -190,8 +190,8 @@ def test_hc_prenorm_gemm_tilelang(num_tokens, hidden_size):
@pytest.mark.skipif(
not (current_platform.is_cuda_alike() and has_tilelang()),
reason="CUDA or ROCm and tilelang required",
not HAS_TILELANG_MHC,
reason="TileLang MHC support required",
)
@pytest.mark.parametrize("num_tokens", [1, 4, 8, 128])
@pytest.mark.parametrize("hidden_size", [4096, 7168])
@@ -217,8 +217,8 @@ def test_mhc_post_tilelang(num_tokens, hidden_size, hc_mult):
@pytest.mark.skipif(
not (current_platform.is_cuda_alike() and has_tilelang()),
reason="CUDA or ROCm and tilelang required",
not HAS_TILELANG_MHC,
reason="TileLang MHC support required",
)
@pytest.mark.parametrize("num_tokens", [1, 4, 8, 128])
@pytest.mark.parametrize("hidden_size", [4096, 7168])
@@ -324,8 +324,8 @@ def test_hc_head_triton(num_tokens, hidden_size, hc_mult):
@pytest.mark.skipif(
not (current_platform.is_cuda_alike() and has_tilelang()),
reason="CUDA or ROCm and tilelang required",
not HAS_TILELANG_MHC,
reason="TileLang MHC support required",
)
@pytest.mark.parametrize("num_tokens", [1, 4, 8, 128])
@pytest.mark.parametrize("hidden_size", [4096, 7168])
@@ -810,29 +810,6 @@ VLM_TEST_SETTINGS = {
hf_output_post_proc=model_utils.minicpmv_trunc_hf_output,
patch_hf_runner=model_utils.minicpmv_26_patch_hf_runner,
),
"minimax_vl_01": VLMTestInfo(
models=["MiniMaxAI/MiniMax-VL-01"],
prompt_formatter=lambda img_prompt: f"<beginning_of_sentence>user: {img_prompt} assistant:<end_of_sentence>", # noqa: E501
img_idx_to_prompt=lambda _: "<image>",
test_type=(VLMTestType.IMAGE, VLMTestType.MULTI_IMAGE),
max_model_len=8192,
max_num_seqs=4,
dtype="bfloat16",
hf_output_post_proc=model_utils.minimax_vl_01_hf_output,
patch_hf_runner=model_utils.minimax_vl_01_patch_hf_runner,
auto_cls=AutoModelForImageTextToText,
marks=[
large_gpu_mark(min_gb=80),
# TODO: [ROCm] Fix pickle issue with ROCm spawn and tp>1
pytest.mark.skipif(
current_platform.is_rocm(),
reason=(
"ROCm: Model too large for single GPU; "
"multi-GPU blocked by HF _LazyConfigMapping pickle issue with spawn"
),
),
],
),
"molmo": VLMTestInfo(
models=["allenai/Molmo-7B-D-0924"],
test_type=(VLMTestType.IMAGE, VLMTestType.MULTI_IMAGE),
@@ -245,13 +245,6 @@ def minicpmv_trunc_hf_output(hf_output: RunnerOutput, model: str) -> RunnerOutpu
return output_ids, output_str, out_logprobs
def minimax_vl_01_hf_output(hf_output: RunnerOutput, model: str) -> RunnerOutput:
output_ids, output_str, out_logprobs = hf_output
if output_str.endswith("<end_of_sentence>"):
output_str = output_str.split("<end_of_sentence>")[0]
return output_ids, output_str, out_logprobs
def ultravox_trunc_hf_output(hf_output: RunnerOutput, model: str) -> RunnerOutput:
output_ids, output_str, out_logprobs = hf_output
@@ -1023,17 +1016,6 @@ def minicpmv_26_patch_hf_runner(hf_model: HfRunner) -> HfRunner:
return hf_model
def minimax_vl_01_patch_hf_runner(hf_model: HfRunner) -> HfRunner:
orig_generate = hf_model.model.generate
def _generate(self, *args, image_sizes=None, **kwargs):
return orig_generate(*args, decode_text=False, **kwargs)
hf_model.model.generate = types.MethodType(_generate, hf_model.model)
return hf_model
def molmo_patch_hf_runner(hf_model: HfRunner) -> HfRunner:
"""Patches and returns an instance of the HfRunner to use for Molmo."""
hf_processor = hf_model.processor
@@ -152,3 +152,21 @@ def test_colqwen3_5_relevance_ordering(
dtype: str,
) -> None:
_run_relevance_test(vllm_runner, model, dtype=dtype)
def test_colqwen3_5_config_enables_bidirectional_attention() -> None:
"""ColQwen3.5 retrieval must be served BIDIRECTIONAL (is_causal=False) so the
full_attention layers build with AttentionType.ENCODER_ONLY. This guards the
silent-causal regression (no GPU / model load needed)."""
from types import SimpleNamespace
from vllm.model_executor.models.config import (
MODELS_CONFIG_MAP,
ColQwen3_5Config,
)
assert MODELS_CONFIG_MAP["ColQwen3_5"] is ColQwen3_5Config
model_config = SimpleNamespace(hf_config=SimpleNamespace())
ColQwen3_5Config.verify_and_update_model_config(model_config)
assert model_config.hf_config.is_causal is False
@@ -1,113 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
from PIL import Image
from vllm.multimodal import MULTIMODAL_REGISTRY
from vllm.multimodal.parse import ImageSize
from vllm.multimodal.processing import BaseMultiModalProcessor
from ....conftest import ImageTestAssets
from ...utils import build_model_context
@pytest.mark.parametrize("model_id", ["MiniMaxAI/MiniMax-VL-01"])
@pytest.mark.parametrize("num_imgs", [1, 2])
def test_processor_override(
image_assets: ImageTestAssets,
model_id: str,
num_imgs: int,
):
ctx = build_model_context(
model_id,
mm_processor_kwargs=None,
limit_mm_per_prompt={"image": num_imgs},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
prompt = "<image>" * num_imgs
image = Image.new("RGB", size=(364, 364))
mm_data = {"image": [image] * num_imgs}
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs={},
)
image_placeholders = processed_inputs["mm_placeholders"]["image"]
assert len(image_placeholders) == num_imgs
def _validate_image_prompt_replacements_one(
processor: BaseMultiModalProcessor,
num_imgs: int,
failed_size_excs: list[tuple[ImageSize, Exception]],
image_size: ImageSize,
) -> None:
prompt = "<image>" * num_imgs
image = Image.new("RGB", size=image_size)
mm_data = {"image": [image] * num_imgs}
try:
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs={},
)
image_placeholders = processed_inputs["mm_placeholders"]["image"]
assert len(image_placeholders) == num_imgs
except Exception as exc:
failed_size_excs.append((image_size, exc))
def _test_image_prompt_replacements(
processor,
*,
num_imgs: int,
image_sizes: list[ImageSize],
) -> None:
failed_size_excs = list[tuple[ImageSize, Exception]]()
for size in image_sizes:
_validate_image_prompt_replacements_one(
processor, num_imgs, failed_size_excs, size
)
if failed_size_excs:
msg = "Found failing image sizes:" + "\n========\n".join(
f"[{size}]\n{exc}" for size, exc in failed_size_excs
)
raise AssertionError(msg)
@pytest.mark.parametrize("model_id", ["MiniMaxAI/MiniMax-VL-01"])
@pytest.mark.parametrize("num_imgs", [1, 2])
def test_processor_prompt_replacements_regression(model_id, num_imgs):
ctx = build_model_context(
model_id,
mm_processor_kwargs=None,
limit_mm_per_prompt={"image": num_imgs},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
image_ratios = [
(171, 152),
(184, 161),
(198, 176),
(333, 296),
(369, 328),
(488, 183),
(2560, 1669),
]
image_sizes = [
size for w, h in image_ratios for size in [ImageSize(w, h), ImageSize(h, w)]
]
_test_image_prompt_replacements(
processor,
num_imgs=num_imgs,
image_sizes=image_sizes,
)
@@ -138,3 +138,49 @@ def test_processor_multi_video(
assert video_phs[i].offset >= prev_end, (
f"Placeholder {i} overlaps with placeholder {i - 1}"
)
@pytest.mark.parametrize("model_id", [MODEL_ID])
@pytest.mark.parametrize(
"hf_mm_kwargs",
[{"num_frames": [8, 16]}, {"fps": [2.0, 4.0]}],
)
def test_processor_multi_video_list_kwargs(
model_id: str,
hf_mm_kwargs: dict[str, Any],
) -> None:
"""Regression test: a multi-video request with list-valued per-video
``mm_processor_kwargs`` (one ``fps``/``num_frames`` per video) must not
crash.
Before the fix, ``_call_hf_processor`` copied the whole kwargs to every
video without slicing, so ``_get_video_second_idx`` received the list
where a scalar was expected and raised ``TypeError``.
"""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 0, "video": 2},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
prompt = (
"<|vision_start|><|video_pad|><|vision_end|>"
"<|vision_start|><|video_pad|><|vision_end|>"
)
mm_data = {
"video": [
_build_video_mm_data(num_frames=16)["video"][0],
_build_video_mm_data(num_frames=32)["video"][0],
]
}
processed = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs=hf_mm_kwargs,
)
video_phs = processed["mm_placeholders"].get("video", [])
assert len(video_phs) == 2, (
f"Expected exactly 2 video placeholders, got {len(video_phs)}"
)
+19 -13
View File
@@ -421,15 +421,6 @@ _TEXT_GENERATION_EXAMPLE_MODELS = {
},
trust_remote_code=True,
),
"MiniMaxForCausalLM": _HfExamplesInfo("MiniMaxAI/MiniMax-Text-01-hf"),
"MiniMaxText01ForCausalLM": _HfExamplesInfo(
"MiniMaxAI/MiniMax-Text-01",
trust_remote_code=True,
revision="a59aa9cbc53b9fb8742ca4e9e1531b9802b6fdc3",
),
"MiniMaxM1ForCausalLM": _HfExamplesInfo(
"MiniMaxAI/MiniMax-M1-40k", trust_remote_code=True
),
"MiniMaxM2ForCausalLM": _HfExamplesInfo(
"MiniMaxAI/MiniMax-M2",
trust_remote_code=True,
@@ -1113,10 +1104,6 @@ _MULTIMODAL_EXAMPLE_MODELS = {
"openbmb/MiniCPM-V-4_6",
min_transformers_version="5.7.0",
),
"MiniMaxVL01ForConditionalGeneration": _HfExamplesInfo(
"MiniMaxAI/MiniMax-VL-01",
trust_remote_code=True,
),
"MiniMaxM3SparseForConditionalGeneration": _HfExamplesInfo(
"MiniMaxAI/MiniMax-M3",
trust_remote_code=True,
@@ -1530,6 +1517,16 @@ _SPECULATIVE_DECODING_EXAMPLE_MODELS = {
"Qwen/Qwen3-VL-8B-Instruct",
speculative_model="taobao-mnn/Qwen3-VL-8B-Instruct-Eagle3",
),
"Eagle3Qwen3ForCausalLM": _HfExamplesInfo(
"Qwen/Qwen3-8B",
trust_remote_code=True,
speculative_model=(
"inference-optimization/"
"Qwen3-8B-from-Qwen3-8B_regen-speculators.eagle3-qwen3arch-ckpt1"
),
tokenizer="Qwen/Qwen3-8B",
use_original_num_layers=True,
),
# [PEagle]
"PEagleDraftModel": _HfExamplesInfo(
"Qwen/Qwen3-8B",
@@ -1545,6 +1542,15 @@ _SPECULATIVE_DECODING_EXAMPLE_MODELS = {
tokenizer="Qwen/Qwen3-8B",
use_original_num_layers=True,
),
"PeagleQwen3ForCausalLM": _HfExamplesInfo(
"Qwen/Qwen3-8B",
trust_remote_code=True,
speculative_model=(
"inference-optimization/Qwen3-8B-speculators.peagle-qwen3arch-ckpt4"
),
tokenizer="Qwen/Qwen3-8B",
use_original_num_layers=True,
),
# [MTP]
"DeepSeekMTPModel": _HfExamplesInfo(
"luccafong/deepseek_mtp_main_random",
-5
View File
@@ -98,11 +98,6 @@ def can_initialize(
vllm_config.validate_block_size()
return scheduler_kv_cache_config
if model_arch == "MiniMaxVL01ForConditionalGeneration":
pytest.skip(
"pickle error when loading `transformers.models.auto.CONFIG_MAPPING`"
)
if model_arch == "MoonshotKimiaForCausalLM":
pytest.skip(
"Kimi-Audio requires SpeechToTextConfig "
+7 -1
View File
@@ -507,7 +507,13 @@ def dummy_hf_overrides(
# Only set MoE related config when the model has MoE layers.
# Otherwise all models detected as MoE by _get_transformers_backend_cls.
if model_arch_config.num_experts > 0:
num_experts_per_tok = 1 if model_arch == "Llama4ForConditionalGeneration" else 2
num_experts_per_tok = 2
if model_arch in (
"Llama4ForConditionalGeneration",
"Llama4ForCausalLM",
"EagleLlama4ForCausalLM",
):
num_experts_per_tok = 1
update_dict.update(
{
"num_experts": num_experts,
+25 -2
View File
@@ -15,6 +15,7 @@ from vllm.multimodal.video import (
DynamicVideoBackend,
GLM46VVideoBackend,
Molmo2VideoBackend,
Qwen2VLVideoBackend,
Qwen3VLVideoBackend,
VideoLoader,
VideoSourceMetadata,
@@ -70,11 +71,12 @@ def test_video_loader_type_doesnt_exist():
@pytest.mark.parametrize(
"model_repo, expected_loader_cls",
"model_repo, expected_loader_cls, hf_sample_kwargs",
[
pytest.param(
"allenai/Molmo2-4B",
Molmo2VideoBackend,
None,
marks=pytest.mark.skip(
reason="Video processor not aligned, investigate later.",
),
@@ -83,23 +85,44 @@ def test_video_loader_type_doesnt_exist():
pytest.param(
"zai-org/GLM-4.1V-9B-Thinking",
DynamicVideoBackend,
None,
id="glm4v",
),
pytest.param(
"zai-org/GLM-4.6V-Flash",
GLM46VVideoBackend,
None,
id="glm46v",
),
pytest.param(
"Qwen/Qwen3-VL-4B-Instruct",
Qwen3VLVideoBackend,
None,
id="qwen3vl",
),
# Qwen2-VL/Qwen2.5-VL ship no ``video_processor_type`` in their
# preprocessor config, so resolution relies on the model_type ->
# video processor fallback in get_video_processor_cls_name_from_config.
# They also ship no default fps/num_frames, so the HF sampler needs an
# explicit target rate; pass fps=2 to match the loader default.
pytest.param(
"Qwen/Qwen2-VL-7B-Instruct",
Qwen2VLVideoBackend,
{"fps": 2},
id="qwen2vl",
),
pytest.param(
"Qwen/Qwen2.5-VL-7B-Instruct",
Qwen2VLVideoBackend,
{"fps": 2},
id="qwen2_5_vl",
),
],
)
def test_video_processor_from_model_repo(
model_repo: str,
expected_loader_cls: type,
hf_sample_kwargs: dict[str, int | float] | None,
):
"""Test that a model repo resolves to the correct video loader backend.
@@ -143,7 +166,7 @@ def test_video_processor_from_model_repo(
fps=vllm_meta["fps"],
duration=vllm_meta["duration"],
)
hf_indices = processor.sample_frames(hf_metadata)
hf_indices = processor.sample_frames(hf_metadata, **(hf_sample_kwargs or {}))
vllm_indices = np.array(vllm_meta["frames_indices"])
np.testing.assert_array_equal(
hf_indices,
+53 -7
View File
@@ -61,8 +61,8 @@ class QuantConfig:
quant_max: float
quant_min: float
kv_quant_mode: KVQuantMode
# INT8 Triton stores truncate; FP8 hardware casts round.
uses_trunc: bool
# INT8 rounds explicitly; FP8 relies on dtype cast rounding.
rounds_before_store: bool
INT8_CONFIG = QuantConfig(
@@ -71,7 +71,7 @@ INT8_CONFIG = QuantConfig(
quant_max=127.0,
quant_min=-128.0,
kv_quant_mode=KVQuantMode.INT8_PER_TOKEN_HEAD,
uses_trunc=True,
rounds_before_store=True,
)
FP8_CONFIG = QuantConfig(
cache_dtype=FP8_DTYPE,
@@ -79,7 +79,7 @@ FP8_CONFIG = QuantConfig(
quant_max=FP8_MAX,
quant_min=FP8_MIN,
kv_quant_mode=KVQuantMode.FP8_PER_TOKEN_HEAD,
uses_trunc=False,
rounds_before_store=False,
)
QUANT_CONFIGS = [INT8_CONFIG, FP8_CONFIG]
@@ -104,7 +104,7 @@ def _quantize_per_token_head_ref(
absmax = data.float().abs().amax(dim=2) # [num_tokens, num_heads]
scales = (absmax / cfg.quant_max).clamp(min=1e-6)
scaled = data.float() * (1.0 / scales[:, :, None])
if cfg.uses_trunc:
if cfg.rounds_before_store:
q = scaled.round().clamp(cfg.quant_min, cfg.quant_max).to(cfg.cache_dtype)
else:
q = scaled.clamp(cfg.quant_min, cfg.quant_max).to(cfg.cache_dtype)
@@ -255,7 +255,7 @@ def test_per_token_head_round_trip_accuracy(
):
"""Verify per-token-head round-trip: kernel dequant matches reference.
INT8: Triton truncates on float->int8 store.
INT8: round-to-nearest before int8 store.
FP8: hardware cast (clamp then cast).
"""
from vllm.v1.attention.ops.triton_reshape_and_cache_flash import (
@@ -315,6 +315,52 @@ def test_per_token_head_round_trip_accuracy(
)
@torch.inference_mode()
def test_int8_per_token_head_raw_cache_matches_round_reference():
"""INT8 cache writes should match round-to-nearest quantization exactly."""
from vllm.v1.attention.ops.triton_reshape_and_cache_flash import (
triton_reshape_and_cache_flash_per_token_head_quant,
)
torch.set_default_device(DEVICE_TYPE)
head_size = 8
block_size = 4
key = torch.tensor(
[[[-127.0, -2.6, -2.4, -1.6, -1.4, -0.6, -0.4, 127.0]]],
dtype=torch.bfloat16,
)
value = -key
key_cache = torch.zeros(1, block_size, 1, head_size, dtype=torch.int8)
value_cache = torch.zeros_like(key_cache)
k_scale_cache = torch.ones(1, block_size, 1, dtype=torch.float32)
v_scale_cache = torch.ones_like(k_scale_cache)
slot_mapping = torch.tensor([2], dtype=torch.long)
triton_reshape_and_cache_flash_per_token_head_quant(
key,
value,
key_cache,
value_cache,
k_scale_cache,
v_scale_cache,
slot_mapping,
)
ref_k_quant, ref_k_scales = _quantize_per_token_head_ref(key, INT8_CONFIG)
ref_v_quant, ref_v_scales = _quantize_per_token_head_ref(value, INT8_CONFIG)
slot = slot_mapping.item()
blk = slot // block_size
off = slot % block_size
assert torch.equal(key_cache[blk, off], ref_k_quant[0])
assert torch.equal(value_cache[blk, off], ref_v_quant[0])
torch.testing.assert_close(k_scale_cache[blk, off], ref_k_scales[0])
torch.testing.assert_close(v_scale_cache[blk, off], ref_v_scales[0])
# ===========================================================================
# 4. Negative slot mapping (padding tokens should be skipped)
# ===========================================================================
@@ -461,7 +507,7 @@ def test_triton_unified_attention_per_token_head_scale(
scaled_k = key_cache_bf16.float() / k_scale_cache[:, :, :, None]
scaled_v = value_cache_bf16.float() / v_scale_cache[:, :, :, None]
if qcfg.uses_trunc:
if qcfg.rounds_before_store:
key_cache_q = (
scaled_k.round().clamp(qcfg.quant_min, qcfg.quant_max).to(qcfg.cache_dtype)
)
File diff suppressed because it is too large Load Diff
@@ -6,10 +6,12 @@ from types import SimpleNamespace
import pytest
from vllm.model_executor.layers.mamba.linear.minimax_linear_attn import (
MiniMaxText01LinearAttention,
)
from vllm.model_executor.layers.mamba.mamba_mixer import MambaMixer
from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2
from vllm.model_executor.layers.mamba.short_conv import ShortConv
from vllm.model_executor.models.minimax_text_01 import MiniMaxText01LinearAttention
from vllm.v1.attention.backends.linear_attn import LinearAttentionBackend
from vllm.v1.attention.backends.mamba1_attn import Mamba1AttentionBackend
from vllm.v1.attention.backends.mamba2_attn import Mamba2AttentionBackend
@@ -16,6 +16,7 @@ from tests.v1.attention.utils import (
create_vllm_config,
)
from vllm.config import SpeculativeConfig
from vllm.config.compilation import CUDAGraphMode
from vllm.v1.attention.backends.gdn_attn import (
GDNAttentionMetadata,
GDNAttentionMetadataBuilder,
@@ -123,9 +124,15 @@ GDN_BUILD_TEST_CASES = {
def _create_gdn_builder(
num_speculative_tokens: int = 0,
full_cuda_graph: bool = False,
) -> GDNAttentionMetadataBuilder:
"""Create a GDNAttentionMetadataBuilder with minimal config."""
vllm_config = create_vllm_config(block_size=BLOCK_SIZE)
vllm_config = create_vllm_config(
model_name="Qwen/Qwen3.5-0.8B",
block_size=BLOCK_SIZE,
)
if full_cuda_graph:
vllm_config.compilation_config.cudagraph_mode = CUDAGraphMode.FULL_AND_PIECEWISE
if num_speculative_tokens > 0:
vllm_config.speculative_config = SpeculativeConfig(
method="ngram",
@@ -189,3 +196,28 @@ def test_has_initial_state_after_reclassification():
assert meta.has_initial_state is not None
# req0 has context_lens = 65 - 1 = 64 > 0, so has_initial_state[0] = True
assert meta.has_initial_state[0].item() is True
def test_full_cudagraph_spec_metadata_uses_request_count():
"""FULL cudagraph token padding must not pad request-indexed metadata."""
num_speculative_tokens = 3
builder = _create_gdn_builder(
num_speculative_tokens=num_speculative_tokens,
full_cuda_graph=True,
)
batch = BatchSpec(seq_lens=[80, 96], query_lens=[4, 4])
meta = _build(builder, batch, num_decode_draft_tokens=[3, 3])
assert meta.num_spec_decodes == batch.batch_size
assert meta.num_spec_decode_tokens == batch.compute_num_tokens()
assert meta.spec_state_indices_tensor is not None
assert meta.spec_state_indices_tensor.shape == (
batch.batch_size,
num_speculative_tokens + 1,
)
assert meta.spec_sequence_masks is not None
assert meta.spec_sequence_masks.shape == (batch.batch_size,)
assert meta.spec_query_start_loc is not None
assert meta.spec_query_start_loc.shape == (batch.batch_size + 1,)
assert meta.num_accepted_tokens is not None
assert meta.num_accepted_tokens.shape == (batch.batch_size,)
@@ -0,0 +1,460 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Callable
import pytest
import vllm.v1.core.kv_cache_utils as kv_cache_utils
from vllm.distributed.kv_events import BlockRemoved, BlockStored
from vllm.sampling_params import SamplingParams
from vllm.utils.hashing import sha256
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_utils import (
BlockHash,
BlockHashListWithBlockSize,
KVCacheBlock,
get_request_block_hasher,
hash_block_tokens,
init_none_hash,
)
from vllm.v1.request import Request
pytestmark = pytest.mark.cpu_test
@pytest.fixture(autouse=True)
def _auto_init_hash_fn():
init_none_hash(sha256)
def make_request(
request_id: str,
prompt_token_ids: list[int],
hash_block_size: int,
hash_fn: Callable,
) -> Request:
sampling_params = SamplingParams(max_tokens=17)
sampling_params.update_from_generation_config({}, eos_token_id=100)
return Request(
request_id=request_id,
prompt_token_ids=prompt_token_ids,
sampling_params=sampling_params,
pooling_params=None,
block_hasher=get_request_block_hasher(hash_block_size, hash_fn),
)
def boundary_hash(req: Request, hash_block_size: int, num_tokens: int) -> BlockHash:
# Every boundary at a hash_block_size multiple is just the fine-grained
# chain hash ending there.
return req.block_hashes[num_tokens // hash_block_size - 1]
def cache_full_block_and_partial_tail(
token_ids: list[int],
*,
enable_kv_cache_events: bool = False,
) -> tuple[BlockPool, Request, list[KVCacheBlock], BlockHash]:
hash_block_size = 2
block_size = 6
kv_cache_group_id = 0
req = make_request("0", token_ids, hash_block_size, sha256)
pool = BlockPool(
num_gpu_blocks=3,
enable_caching=True,
hash_block_size=hash_block_size,
enable_kv_cache_events=enable_kv_cache_events,
)
blocks = pool.get_new_blocks(2)
pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=1,
block_size=block_size,
kv_cache_group_id=kv_cache_group_id,
)
partial_hash = boundary_hash(req, hash_block_size, len(token_ids))
assert pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=len(token_ids),
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
return pool, req, blocks, partial_hash
def test_boundary_hashes_reuse_fine_grained_chain():
hash_block_size = 2
block_size = 6
token_ids = [0, 0, 1, 1, 2, 2, 3, 3, 4, 4]
req = make_request("0", token_ids, hash_block_size, sha256)
coarse = BlockHashListWithBlockSize(req.block_hashes, hash_block_size, block_size)
# The block_size=6 full-block hash is the fine hash at the 6-token boundary,
# not a concatenation of the three fine hashes inside the block.
assert coarse[0] == req.block_hashes[6 // hash_block_size - 1]
assert coarse[0] != BlockHash(
req.block_hashes[0] + req.block_hashes[1] + req.block_hashes[2]
)
# A partial tail at 10 tokens is the fine hash at the 10-token boundary,
# which chains over the entire prefix.
tail_hash = boundary_hash(req, hash_block_size, 10)
assert tail_hash == req.block_hashes[4]
assert tail_hash == hash_block_tokens(sha256, req.block_hashes[3], token_ids[8:10])
def test_cache_partial_block_kv_cache_events():
hash_block_size = 4
block_size = 12
kv_cache_group_id = 2
pool = BlockPool(
num_gpu_blocks=2,
enable_caching=True,
hash_block_size=hash_block_size,
enable_kv_cache_events=True,
)
req = make_request(
"req_partial_events",
prompt_token_ids=list(range(hash_block_size * 2)),
hash_block_size=hash_block_size,
hash_fn=sha256,
)
block = pool.get_new_blocks(1)[0]
partial_entry_hash = pool.cache_partial_block(
request=req,
block=block,
num_tokens=hash_block_size * 2,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
events = pool.take_events()
assert len(events) == 1
stored_event = events[0]
assert isinstance(stored_event, BlockStored)
assert partial_entry_hash is not None
assert stored_event.block_hashes == [
kv_cache_utils.maybe_convert_block_hash(req.block_hashes[1])
]
assert stored_event.parent_block_hash == kv_cache_utils.maybe_convert_block_hash(
req.block_hashes[0]
)
assert stored_event.token_ids == req.all_token_ids[hash_block_size:]
assert stored_event.block_size == 4
assert stored_event.group_idx == kv_cache_group_id
duplicate_entry_hash = pool.cache_partial_block(
request=req,
block=block,
num_tokens=hash_block_size * 2,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
assert duplicate_entry_hash == partial_entry_hash
assert pool.take_events() == []
pool.free_blocks([block])
pool.get_new_blocks(1)
events = pool.take_events()
assert len(events) == 1
removed_event = events[0]
assert isinstance(removed_event, BlockRemoved)
assert removed_event.block_hashes == stored_event.block_hashes
assert removed_event.group_idx == kv_cache_group_id
def test_partial_block_replacement_emits_remove_then_store_events():
hash_block_size = 2
block_size = 6
kv_cache_group_id = 0
req = make_request("0", [0, 0, 1, 1, 2, 2, 3, 3], hash_block_size, sha256)
pool = BlockPool(
num_gpu_blocks=3,
enable_caching=True,
hash_block_size=hash_block_size,
enable_kv_cache_events=True,
)
blocks = pool.get_new_blocks(2)
pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=1,
block_size=block_size,
kv_cache_group_id=kv_cache_group_id,
)
partial_hash_8 = boundary_hash(req, hash_block_size, 8)
assert pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=8,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
assert pool.get_cached_block(partial_hash_8, [kv_cache_group_id]) == [blocks[1]]
pool.take_events()
req.append_output_token_ids([4, 4])
partial_hash_10 = boundary_hash(req, hash_block_size, 10)
assert pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=10,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
events = pool.take_events()
assert len(events) == 2
removed_event, stored_event = events
assert isinstance(removed_event, BlockRemoved)
assert removed_event.block_hashes == [
kv_cache_utils.maybe_convert_block_hash(partial_hash_8)
]
assert removed_event.group_idx == kv_cache_group_id
assert isinstance(stored_event, BlockStored)
assert stored_event.block_hashes == [
kv_cache_utils.maybe_convert_block_hash(partial_hash_10)
]
assert stored_event.parent_block_hash == kv_cache_utils.maybe_convert_block_hash(
boundary_hash(req, hash_block_size, 8)
)
assert stored_event.token_ids == req.all_token_ids[8:10]
assert stored_event.block_size == hash_block_size
assert stored_event.group_idx == kv_cache_group_id
assert pool.get_cached_block(partial_hash_8, [kv_cache_group_id]) is None
assert pool.get_cached_block(partial_hash_10, [kv_cache_group_id]) == [blocks[1]]
def test_later_request_hits_cached_partial_tail():
hash_block_size = 2
block_size = 6
kv_cache_group_id = 0
cached_token_ids = [0, 0, 1, 1, 2, 2, 3, 3, 4, 4]
req = make_request("0", cached_token_ids, hash_block_size, sha256)
pool = BlockPool(
num_gpu_blocks=3,
enable_caching=True,
hash_block_size=hash_block_size,
)
blocks = pool.get_new_blocks(2)
pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=1,
block_size=block_size,
kv_cache_group_id=kv_cache_group_id,
)
partial_hash_10 = boundary_hash(req, hash_block_size, 10)
assert pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=10,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
replay = make_request("1", cached_token_ids, hash_block_size, sha256)
replay_hash_10 = boundary_hash(replay, hash_block_size, 10)
assert replay_hash_10 == partial_hash_10
assert pool.get_cached_block(replay_hash_10, [kv_cache_group_id]) == [blocks[1]]
extended = make_request("2", cached_token_ids + [10], hash_block_size, sha256)
extended_hash_10 = boundary_hash(extended, hash_block_size, 10)
assert extended_hash_10 == partial_hash_10
assert pool.get_cached_block(extended_hash_10, [kv_cache_group_id]) == [blocks[1]]
def test_cache_partial_block_uses_fine_grained_boundary_hash():
hash_block_size = 2
block_size = 6
kv_cache_group_id = 0
token_ids = [0, 0, 1, 1, 2, 2, 3, 3, 4, 4]
req = make_request("0", token_ids, hash_block_size, sha256)
pool = BlockPool(
num_gpu_blocks=3,
enable_caching=True,
hash_block_size=hash_block_size,
)
blocks = pool.get_new_blocks(2)
pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=1,
block_size=block_size,
kv_cache_group_id=kv_cache_group_id,
)
partial_entry_hash = pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=10,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
# The partial entry is keyed by the fine-grained hash at the 10-token
# boundary, regardless of the owning group's block_size.
expected = boundary_hash(req, hash_block_size, 10)
assert partial_entry_hash == kv_cache_utils.make_block_hash_with_group_id(
expected, kv_cache_group_id
)
assert pool.get_cached_block(expected, [kv_cache_group_id]) == [blocks[1]]
def test_cache_partial_block_requires_hash_boundary():
hash_block_size = 2
block_size = 4
req = make_request("0", [0, 0, 1, 1], hash_block_size, sha256)
pool = BlockPool(
num_gpu_blocks=2,
enable_caching=True,
hash_block_size=hash_block_size,
)
block = pool.get_new_blocks(1)[0]
with pytest.raises(AssertionError):
pool.cache_partial_block(
request=req,
block=block,
num_tokens=3,
kv_cache_group_id=0,
block_size=block_size,
)
def test_cache_partial_block_duplicate_checks_all_blocks_for_hash():
hash_block_size = 2
block_size = 4
kv_cache_group_id = 0
req = make_request("0", [0, 0, 1, 1], hash_block_size, sha256)
pool = BlockPool(
num_gpu_blocks=4,
enable_caching=True,
hash_block_size=hash_block_size,
)
blocks = pool.get_new_blocks(2)
first_entry_hash = pool.cache_partial_block(
request=req,
block=blocks[0],
num_tokens=2,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
second_entry_hash = pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=2,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
assert first_entry_hash == second_entry_hash
duplicate_entry_hash = pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=2,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
assert duplicate_entry_hash == second_entry_hash
assert pool.cached_block_hashes_by_block == {}
def test_reset_prefix_cache_clears_partial_entry_metadata():
pool, req, blocks, partial_hash_10 = cache_full_block_and_partial_tail(
[0, 0, 1, 1, 2, 2, 3, 3, 4, 4]
)
full_hash = BlockHashListWithBlockSize(req.block_hashes, 2, 6)[0]
assert pool.get_cached_block(full_hash, [0]) == [blocks[0]]
assert pool.get_cached_block(partial_hash_10, [0]) == [blocks[1]]
pool.free_blocks(blocks)
assert pool.reset_prefix_cache()
assert pool.get_cached_block(full_hash, [0]) is None
assert pool.get_cached_block(partial_hash_10, [0]) is None
assert pool.cached_block_hashes_by_block == {}
def test_evict_cached_block_removes_full_hash_and_partial_entry():
pool, req, blocks, partial_hash_10 = cache_full_block_and_partial_tail(
[0, 0, 1, 1, 2, 2, 3, 3, 4, 4]
)
full_hash = BlockHashListWithBlockSize(req.block_hashes, 2, 6)[0]
assert pool.get_cached_block(full_hash, [0]) == [blocks[0]]
assert pool.get_cached_block(partial_hash_10, [0]) == [blocks[1]]
pool.evict_blocks({blocks[0].block_id, blocks[1].block_id})
assert pool.get_cached_block(full_hash, [0]) is None
assert pool.get_cached_block(partial_hash_10, [0]) is None
assert pool.cached_block_hashes_by_block == {}
def test_partial_block_promotes_to_direct_full_block_hash():
hash_block_size = 2
block_size = 6
kv_cache_group_id = 0
token_ids = [0, 0, 1, 1, 2, 2, 3, 3, 4, 4]
req = make_request("0", token_ids, hash_block_size, sha256)
pool = BlockPool(
num_gpu_blocks=3,
enable_caching=True,
hash_block_size=hash_block_size,
)
blocks = pool.get_new_blocks(2)
pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=1,
block_size=block_size,
kv_cache_group_id=kv_cache_group_id,
)
partial_hash_10 = boundary_hash(req, hash_block_size, 10)
assert pool.cache_partial_block(
request=req,
block=blocks[1],
num_tokens=10,
kv_cache_group_id=kv_cache_group_id,
block_size=block_size,
)
assert pool.get_cached_block(partial_hash_10, [kv_cache_group_id]) == [blocks[1]]
req.append_output_token_ids([5, 5])
full_hashes = BlockHashListWithBlockSize(
req.block_hashes, hash_block_size, block_size
)
promoted_full_hash = full_hashes[1]
# The promoted full-block hash is the fine hash at the 12-token boundary,
# not a concatenation of the fine hashes inside the block.
assert promoted_full_hash == req.block_hashes[12 // hash_block_size - 1]
assert promoted_full_hash != BlockHash(
req.block_hashes[3] + req.block_hashes[4] + req.block_hashes[5]
)
pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=1,
num_full_blocks=2,
block_size=block_size,
kv_cache_group_id=kv_cache_group_id,
)
assert pool.get_cached_block(promoted_full_hash, [kv_cache_group_id]) == [blocks[1]]
assert pool.get_cached_block(partial_hash_10, [kv_cache_group_id]) is None
+103 -8
View File
@@ -117,6 +117,7 @@ def new_kv_cache_spec(
page_size_padded=None,
sliding_window=None,
attention_chunk_size=None,
indexes_kv_by_block_stride=False,
):
return FullAttentionSpec(
block_size=block_size,
@@ -126,6 +127,7 @@ def new_kv_cache_spec(
page_size_padded=page_size_padded,
sliding_window=sliding_window,
attention_chunk_size=attention_chunk_size,
indexes_kv_by_block_stride=indexes_kv_by_block_stride,
)
@@ -136,6 +138,7 @@ def new_sliding_window_spec(
dtype=torch.float32,
page_size_padded=None,
sliding_window=1,
indexes_kv_by_block_stride=False,
):
return SlidingWindowSpec(
block_size=block_size,
@@ -144,6 +147,7 @@ def new_sliding_window_spec(
dtype=dtype,
page_size_padded=page_size_padded,
sliding_window=sliding_window,
indexes_kv_by_block_stride=indexes_kv_by_block_stride,
)
@@ -221,7 +225,7 @@ def test_kv_cache_block():
# Test block hash setting and resetting
block_hash = make_block_hash_with_group_id(BlockHash(b"abc"), 0)
block.block_hash = block_hash
block.set_block_hash(block_hash)
assert block.block_hash == block_hash
block.reset_hash()
@@ -1799,16 +1803,38 @@ def test_get_kv_cache_config_one_worker():
],
)
# different hidden size that cannot be aligned by using different block size
# different hidden size that cannot be aligned by using different block size,
# but can be aligned by padding the smaller physical page.
swa_spec = new_sliding_window_spec(head_size=96, indexes_kv_by_block_stride=True)
kv_cache_specs_hybrid = {
"layer_1": new_kv_cache_spec(head_size=64),
"layer_2": new_sliding_window_spec(head_size=96),
"layer_1": new_kv_cache_spec(head_size=64, indexes_kv_by_block_stride=True),
"layer_2": swa_spec,
}
with pytest.raises(NotImplementedError):
get_kv_cache_configs(
vllm_config, [kv_cache_specs_hybrid], [mem_per_block_per_layer * 2 * 32]
)[0]
kv_cache_config_hybrid = get_kv_cache_configs(
vllm_config, [kv_cache_specs_hybrid], [mem_per_block_per_layer * 2 * 32]
)[0]
padded_page_size = swa_spec.page_size_bytes
assert kv_cache_config_hybrid == KVCacheConfig(
num_blocks=42,
kv_cache_tensors=[
KVCacheTensor(size=padded_page_size * 42, shared_by=["layer_1", "layer_2"]),
],
kv_cache_groups=[
KVCacheGroupSpec(
["layer_1"],
new_kv_cache_spec(
head_size=64,
page_size_padded=padded_page_size,
indexes_kv_by_block_stride=True,
),
),
KVCacheGroupSpec(
["layer_2"],
new_sliding_window_spec(head_size=96, indexes_kv_by_block_stride=True),
),
],
)
# Test num_gpu_blocks_override
vllm_config.cache_config.num_gpu_blocks_override = 16
@@ -2322,6 +2348,75 @@ def test_check_enough_kv_cache_memory_respects_num_gpu_blocks_override():
get_kv_cache_configs(vllm_config, [kv_cache_specs], [large_available_memory])
def test_unify_kv_cache_page_size_uses_padding_for_non_divisible_sizes():
"""DFlash drafters can have a smaller head size than the target model.
For example, MiMo uses 192-dim target KV heads while its DFlash draft uses
128-dim KV heads. The resulting page sizes are 3:2 rather than an integer
block-size multiple, so the smaller page must be padded instead.
"""
# Both layers' backends opt into the padded-page strided view (e.g.
# FlashAttention / its DiffKV subclass), so padding is allowed.
target_spec = new_kv_cache_spec(
block_size=16,
num_kv_heads=1,
head_size=192,
dtype=torch.bfloat16,
indexes_kv_by_block_stride=True,
)
draft_spec = new_sliding_window_spec(
block_size=16,
num_kv_heads=1,
head_size=128,
dtype=torch.bfloat16,
sliding_window=1024,
indexes_kv_by_block_stride=True,
)
unified_specs = kv_cache_utils.unify_kv_cache_spec_page_size(
{
"target_attn": target_spec,
"draft_attn": draft_spec,
}
)
assert unified_specs["target_attn"] == target_spec
unified_draft_spec = unified_specs["draft_attn"]
assert unified_draft_spec.block_size == draft_spec.block_size
assert unified_draft_spec.real_page_size_bytes == draft_spec.real_page_size_bytes
assert unified_draft_spec.page_size_padded == target_spec.page_size_bytes
assert unified_draft_spec.page_size_bytes == target_spec.page_size_bytes
def test_unify_kv_cache_page_size_padding_requires_backend_support():
"""Padding is gated on the backend declaring ``indexes_kv_by_block_stride``.
A backend that does not support the strided padded-page view must raise
rather than silently padding (and misreading KV at runtime).
"""
target_spec = new_kv_cache_spec(
block_size=16,
num_kv_heads=1,
head_size=192,
dtype=torch.bfloat16,
indexes_kv_by_block_stride=True,
)
# The non-divisible draft layer needs padding but its backend does not
# support the strided padded-page view -> must raise, not silently pad.
draft_spec = new_sliding_window_spec(
block_size=16,
num_kv_heads=1,
head_size=128,
dtype=torch.bfloat16,
sliding_window=1024,
indexes_kv_by_block_stride=False,
)
specs = {"target_attn": target_spec, "draft_attn": draft_spec}
with pytest.raises(NotImplementedError):
kv_cache_utils.unify_kv_cache_spec_page_size(specs)
def test_unify_hybrid_kv_cache_specs():
# 1. has_full_attention and has_sliding_window
before_spec_1 = new_kv_cache_spec()
+1 -1
View File
@@ -2003,7 +2003,7 @@ def test_maybe_evict_cached_block():
assert len(pool.blocks) == len(block_hashes)
# Manually add all blocks to cached_blocks
for block, block_hash in zip(pool.blocks, block_hashes):
block.block_hash = block_hash
block.set_block_hash(block_hash)
pool.cached_block_hash_to_block.insert(block_hash, block)
block0, block1, block2, block3 = pool.blocks
+1 -1
View File
@@ -425,7 +425,7 @@ def _run_eagle_correctness(
if "deepseek" in model_setup[1].lower():
m.setenv("VLLM_ROCM_USE_AITER", "1")
m.delenv("VLLM_MLA_DISABLE", raising=False)
attention_config = {"backend": "TRITON_MLA"}
attention_config = {"backend": "ROCM_AITER_MLA"}
else:
m.setenv("VLLM_ROCM_USE_AITER", "1")
@@ -0,0 +1,355 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from unittest.mock import MagicMock
import pytest
import torch
from tests.v1.kv_connector.unit.utils import create_vllm_config
from vllm.config import KVEventsConfig, KVTransferConfig
from vllm.distributed.kv_events import BlockRemoved, BlockStored
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.events import (
OffloadingEventGroupSpec,
OffloadingEventsTracker,
)
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.scheduler import (
GroupOffloadConfig,
)
from vllm.v1.core.kv_cache_utils import BlockHash, maybe_convert_block_hash
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheConfig,
KVCacheGroupSpec,
KVCacheSpecKind,
)
from vllm.v1.kv_offload.base import (
OffloadingEvent,
OffloadingKVEventsConfig,
OffloadKey,
make_offload_key,
)
from vllm.v1.kv_offload.cpu.common import CPULoadStoreSpec
from vllm.v1.kv_offload.tiering.spec import TieringOffloadingSpec
_CPU_MEDIUM = CPULoadStoreSpec.medium()
_FULL_ATTENTION_EVENT_SPEC = OffloadingEventGroupSpec(
kv_cache_spec_kind=KVCacheSpecKind.FULL_ATTENTION.value,
kv_cache_spec_sliding_window=None,
)
def _tracker(
*,
enable_kv_cache_events: bool = True,
self_describing_kv_events: bool = True,
) -> OffloadingEventsTracker:
return OffloadingEventsTracker(
OffloadingKVEventsConfig(
enable_kv_cache_events=enable_kv_cache_events,
self_describing_kv_events=self_describing_kv_events,
)
)
def _hash(i: int) -> BlockHash:
return BlockHash(str(i).encode())
def _wire_hash(block_hash: BlockHash):
return maybe_convert_block_hash(block_hash)
def _request(*, block_hashes: list[BlockHash], token_count: int):
req = MagicMock()
req.block_hashes = block_hashes
req.all_token_ids = list(range(1, token_count + 1))
req.lora_request = None
return req
def _group_config(
*,
group_idx: int = 0,
block_size: int = 4,
block_size_factor: int = 1,
sliding_window_size_in_blocks: int | None = None,
) -> GroupOffloadConfig:
return GroupOffloadConfig(
group_idx=group_idx,
gpu_block_size=block_size,
offloaded_block_size=block_size * block_size_factor,
hash_block_size_factor=block_size_factor,
sliding_window_size_in_blocks=sliding_window_size_in_blocks,
kv_event_group_spec=_FULL_ATTENTION_EVENT_SPEC,
)
def _record_chunks(
tracker: OffloadingEventsTracker,
req,
group_config: GroupOffloadConfig,
num_chunks: int,
) -> list[OffloadKey]:
keys: list[OffloadKey] = []
hbf = group_config.hash_block_size_factor
for chunk_idx in range(num_chunks):
tail_hash = req.block_hashes[(chunk_idx + 1) * hbf - 1]
assert tail_hash is not None
key = make_offload_key(tail_hash, group_config.group_idx)
tracker.record_store(req, group_config, chunk_idx, key)
keys.append(key)
return keys
def _stored_event(keys: list[OffloadKey]) -> OffloadingEvent:
return OffloadingEvent(keys=keys, medium=_CPU_MEDIUM, removed=False)
def _removed_event(keys: list[OffloadKey]) -> OffloadingEvent:
return OffloadingEvent(keys=keys, medium=_CPU_MEDIUM, removed=True)
def test_take_events_publishes_routable_block_stored():
block_size = 4
tracker = _tracker()
group_config = _group_config(block_size=block_size)
req = _request(
block_hashes=[_hash(i) for i in range(6)],
token_count=block_size * 6,
)
keys = _record_chunks(tracker, req, group_config, num_chunks=6)
batch1 = list(tracker.take_events([_stored_event(keys[:3])]))
assert len(batch1) == 3
for i, event in enumerate(batch1):
assert isinstance(event, BlockStored)
assert event.medium == _CPU_MEDIUM
assert event.block_hashes == [_wire_hash(_hash(i))]
assert event.block_size == block_size
assert event.token_ids == list(
range(i * block_size + 1, (i + 1) * block_size + 1)
)
if i == 0:
assert event.parent_block_hash is None
else:
assert event.parent_block_hash == _wire_hash(_hash(i - 1))
assert event.lora_id is None
assert event.lora_name is None
assert event.extra_keys is None
assert event.group_idx == 0
assert event.kv_cache_spec_kind == KVCacheSpecKind.FULL_ATTENTION.value
assert event.kv_cache_spec_sliding_window is None
batch2 = list(tracker.take_events([_stored_event(keys[3:])]))
assert len(batch2) == 3
assert batch2[0].parent_block_hash == batch1[-1].block_hashes[-1]
assert len(tracker._pending_event_metadata) == 6
def test_take_events_factor_gt_1_chunk_store_and_remove():
block_size = 4
block_size_factor = 3
tracker = _tracker()
group_config = _group_config(
block_size=block_size, block_size_factor=block_size_factor
)
req = _request(
block_hashes=[_hash(i) for i in range(6)],
token_count=block_size * block_size_factor * 2,
)
keys = _record_chunks(tracker, req, group_config, num_chunks=2)
stored = list(tracker.take_events([_stored_event(keys)]))
assert len(stored) == 2
expected_hashes = []
for chunk_idx, event in enumerate(stored):
assert isinstance(event, BlockStored)
expected_chunk_hashes = [
_wire_hash(_hash(i))
for i in range(
chunk_idx * block_size_factor,
(chunk_idx + 1) * block_size_factor,
)
]
assert event.block_hashes == expected_chunk_hashes
assert event.block_size == block_size
assert len(event.token_ids) == block_size * block_size_factor
if chunk_idx == 0:
assert event.parent_block_hash is None
else:
assert event.parent_block_hash == _wire_hash(_hash(block_size_factor - 1))
expected_hashes.extend(expected_chunk_hashes)
assert len(tracker._pending_event_metadata) == 2
removed = list(tracker.take_events([_removed_event(keys)]))
assert len(removed) == 1
assert isinstance(removed[0], BlockRemoved)
assert removed[0].block_hashes == expected_hashes
assert removed[0].medium == _CPU_MEDIUM
assert removed[0].group_idx == 0
assert not tracker._pending_event_metadata
def test_take_events_factor_gt_1_store_is_order_independent():
block_size_factor = 3
tracker = _tracker()
group_config = _group_config(block_size_factor=block_size_factor)
req = _request(
block_hashes=[_hash(i) for i in range(6)],
token_count=4 * block_size_factor * 2,
)
keys = _record_chunks(tracker, req, group_config, num_chunks=2)
unknown_key = make_offload_key(_hash(12345), 0)
events = list(tracker.take_events([_stored_event([keys[1], unknown_key, keys[0]])]))
assert len(events) == 3
chunk1, placeholder, chunk0 = events
assert [len(event.block_hashes) for event in events] == [3, 1, 3]
assert placeholder.block_size == 0
assert placeholder.token_ids == []
assert chunk0.parent_block_hash is None
assert chunk1.parent_block_hash == chunk0.block_hashes[-1]
def test_take_events_opt_out_keeps_placeholders():
tracker = _tracker(self_describing_kv_events=False)
group_config = _group_config()
req = _request(block_hashes=[_hash(i) for i in range(3)], token_count=12)
keys = _record_chunks(tracker, req, group_config, num_chunks=3)
assert not tracker.self_describing_enabled
assert not tracker._pending_event_metadata
events = list(
tracker.take_events(
[
_stored_event(keys),
_removed_event(keys),
]
)
)
assert len(events) == 4
for event in events[:3]:
assert isinstance(event, BlockStored)
assert event.block_size == 0
assert event.token_ids == []
assert event.parent_block_hash is None
assert isinstance(events[3], BlockRemoved)
assert len(events[3].block_hashes) == 3
def test_record_store_skips_sliding_window_group():
tracker = _tracker()
group_config = _group_config(sliding_window_size_in_blocks=2)
req = _request(block_hashes=[_hash(i) for i in range(3)], token_count=12)
keys = _record_chunks(tracker, req, group_config, num_chunks=3)
assert not tracker._pending_event_metadata
events = list(tracker.take_events([_stored_event(keys[:1])]))
assert len(events) == 1
assert isinstance(events[0], BlockStored)
assert events[0].block_size == 0
def test_take_events_groups_removed_hashes_by_kv_group():
tracker = _tracker()
group0_config = _group_config(group_idx=0, block_size_factor=2)
group1_config = _group_config(group_idx=1, block_size_factor=2)
req0 = _request(block_hashes=[_hash(0), _hash(1)], token_count=8)
req1 = _request(block_hashes=[_hash(10), _hash(11)], token_count=8)
key0 = _record_chunks(tracker, req0, group0_config, num_chunks=1)[0]
key1 = _record_chunks(tracker, req1, group1_config, num_chunks=1)[0]
removed = list(tracker.take_events([_removed_event([key0, key1])]))
assert len(removed) == 2
by_group = {event.group_idx: event.block_hashes for event in removed}
assert by_group == {
0: [_wire_hash(_hash(0)), _wire_hash(_hash(1))],
1: [_wire_hash(_hash(10)), _wire_hash(_hash(11))],
}
def test_take_events_supports_restore_after_eviction():
block_size = 4
tracker = _tracker()
group_config = _group_config(block_size=block_size)
req = _request(block_hashes=[_hash(0)], token_count=block_size)
key = _record_chunks(tracker, req, group_config, num_chunks=1)[0]
first_store = list(tracker.take_events([_stored_event([key])]))
assert len(first_store) == 1
assert isinstance(first_store[0], BlockStored)
assert first_store[0].token_ids == [1, 2, 3, 4]
removed = list(tracker.take_events([_removed_event([key])]))
assert len(removed) == 1
assert isinstance(removed[0], BlockRemoved)
assert not tracker._pending_event_metadata
req.all_token_ids = [5, 6, 7, 8]
tracker.record_store(req, group_config, offload_block_idx=0, offload_key=key)
second_store = list(tracker.take_events([_stored_event([key])]))
assert len(second_store) == 1
assert isinstance(second_store[0], BlockStored)
assert second_store[0].token_ids == [5, 6, 7, 8]
def test_reset_cache_clears_side_table():
tracker = _tracker()
group_config = _group_config()
req = _request(block_hashes=[_hash(i) for i in range(3)], token_count=12)
_record_chunks(tracker, req, group_config, num_chunks=3)
assert tracker._pending_event_metadata
tracker.reset()
assert not tracker._pending_event_metadata
def test_tiering_rejects_self_describing_kv_events():
vllm_config = create_vllm_config(
block_size=4,
max_num_batched_tokens=16,
disable_hybrid_kv_cache_manager=False,
)
vllm_config.kv_transfer_config = KVTransferConfig(
kv_connector="OffloadingConnector",
kv_role="kv_both",
kv_connector_extra_config={
"spec_name": "TieringOffloadingSpec",
"cpu_bytes_to_use": 1 << 20,
"self_describing_kv_events": True,
"secondary_tiers": [{"type": "example"}],
},
)
vllm_config.kv_events_config = KVEventsConfig(
enable_kv_cache_events=True,
publisher="null",
)
kv_cache_config = KVCacheConfig(
num_blocks=0,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["layer"],
FullAttentionSpec(
block_size=4,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
)
],
)
with pytest.raises(ValueError, match="TieringOffloadingSpec"):
TieringOffloadingSpec(vllm_config, kv_cache_config)
@@ -22,7 +22,7 @@ from vllm.v1.kv_offload.base import (
OffloadingGaugeMetadata,
OffloadingHistogramMetadata,
)
from vllm.v1.kv_offload.cpu.spec import CPUOffloadingSpec
from vllm.v1.kv_offload.factory import OffloadingSpecFactory
LOAD_BYTES = _TransferMetricName.LOAD_BYTES
LOAD_TIME = _TransferMetricName.LOAD_TIME
@@ -33,6 +33,8 @@ STORE_SIZE = _TransferMetricName.STORE_SIZE
STORES_SKIPPED = "vllm:kv_offload_stores_skipped"
PENDING_STORES = "vllm:kv_offload_pending_stores"
LOOKUP_LATENCY = "vllm:kv_offload_lookup_latency_seconds"
MY_COUNTER = "my_counter"
MY_LABEL = "my_label"
class _FakeMetric:
@@ -67,6 +69,20 @@ class _FakeVllmConfig:
)
def _spec_cls_with_metric_definitions(
metric_definitions: dict[str, Any],
) -> type:
"""Build a fake offloading spec class reporting the given metric
definitions, so tests don't need to patch the real CPU spec."""
class _FakeOffloadingSpec:
@staticmethod
def build_metric_definitions(extra_config):
return metric_definitions
return _FakeOffloadingSpec
def _metric_metadata():
return {
LOAD_BYTES: OffloadingCounterMetadata(
@@ -96,9 +112,17 @@ def _metric_metadata():
LOOKUP_LATENCY: OffloadingHistogramMetadata(
documentation="lookup latency",
),
MY_COUNTER: OffloadingCounterMetadata(
documentation="counter with a label",
labelnames=(MY_LABEL,),
),
}
def _unlabeled(values: dict[str, Any], metric_name: str) -> Any:
return values[metric_name][()]
def test_build_kv_connector_stats_with_none():
"""Test that build_kv_connector_stats returns empty stats when given None."""
stats = OffloadingConnector.build_kv_connector_stats(data=None)
@@ -131,13 +155,13 @@ def test_build_kv_connector_stats_reconstructs_offload_stats():
STORES_SKIPPED: _MetricType.COUNTER,
},
_StatsKey.DATA: {
LOAD_BYTES: 24,
LOAD_TIME: 1.5,
LOAD_SIZE: [16, 8],
STORE_BYTES: 3,
STORE_TIME: 0.3,
STORE_SIZE: [1, 2],
STORES_SKIPPED: 5,
LOAD_BYTES: {(): 24},
LOAD_TIME: {(): 1.5},
LOAD_SIZE: {(): [16, 8]},
STORE_BYTES: {(): 3},
STORE_TIME: {(): 0.3},
STORE_SIZE: {(): [1, 2]},
STORES_SKIPPED: {(): 5},
},
}
@@ -145,22 +169,28 @@ def test_build_kv_connector_stats_reconstructs_offload_stats():
assert isinstance(stats, OffloadingConnectorStats)
values = stats.data[_StatsKey.DATA]
assert values[LOAD_BYTES] == 24
assert values[LOAD_TIME] == 1.5
assert values[LOAD_SIZE] == [16, 8]
assert values[STORE_BYTES] == 3
assert values[STORE_TIME] == 0.3
assert values[STORE_SIZE] == [1, 2]
assert values[STORES_SKIPPED] == 5
assert _unlabeled(values, LOAD_BYTES) == 24
assert _unlabeled(values, LOAD_TIME) == 1.5
assert _unlabeled(values, LOAD_SIZE) == [16, 8]
assert _unlabeled(values, STORE_BYTES) == 3
assert _unlabeled(values, STORE_TIME) == 0.3
assert _unlabeled(values, STORE_SIZE) == [1, 2]
assert _unlabeled(values, STORES_SKIPPED) == 5
def _make_stats_data(
metric_data: dict[str, Any],
metric_metadata: dict[str, Any],
) -> dict[str, Any]:
"""Build a structured data dict from flat metric data and metadata."""
"""Build a structured data dict from flat metric data and metadata.
Values for unlabeled metrics may be passed flat (wrapped here under the
empty label tuple); values for labeled metrics must already be passed as
a ``{labelvalues: value}`` map.
"""
metric_types = {}
for key in metric_data:
data = {}
for key, value in metric_data.items():
md = metric_metadata[key]
if isinstance(md, OffloadingCounterMetadata):
metric_types[key] = _MetricType.COUNTER
@@ -168,9 +198,10 @@ def _make_stats_data(
metric_types[key] = _MetricType.GAUGE
elif isinstance(md, OffloadingHistogramMetadata):
metric_types[key] = _MetricType.HISTOGRAM
data[key] = value if md.labelnames else {(): value}
return {
_StatsKey.TYPES: metric_types,
_StatsKey.DATA: metric_data,
_StatsKey.DATA: data,
}
@@ -215,34 +246,106 @@ def test_aggregate_same_connector():
assert result is stats1 # Should return self
values = result.data[_StatsKey.DATA]
assert values[LOAD_BYTES] == 34
assert values[LOAD_TIME] == 2.6
assert values[LOAD_SIZE] == [16, 8, 3, 7]
assert values[STORE_BYTES] == 19
assert values[STORE_TIME] == 2.3
assert values[STORE_SIZE] == [1, 2, 16]
assert values[STORES_SKIPPED] == 4
assert values[PENDING_STORES] == 1
assert values[LOOKUP_LATENCY] == [0.1, 0.2, 0.3]
assert _unlabeled(values, LOAD_BYTES) == 34
assert _unlabeled(values, LOAD_TIME) == 2.6
assert _unlabeled(values, LOAD_SIZE) == [16, 8, 3, 7]
assert _unlabeled(values, STORE_BYTES) == 19
assert _unlabeled(values, STORE_TIME) == 2.3
assert _unlabeled(values, STORE_SIZE) == [1, 2, 16]
assert _unlabeled(values, STORES_SKIPPED) == 4
assert _unlabeled(values, PENDING_STORES) == 1
assert _unlabeled(values, LOOKUP_LATENCY) == [0.1, 0.2, 0.3]
def test_aggregate_labeled_metrics():
metadata = _metric_metadata()
stats1 = OffloadingConnectorStats(
data=_make_stats_data(
{
MY_COUNTER: {
("a",): 10,
("b",): 3,
},
},
metadata,
),
)
stats2 = OffloadingConnectorStats(
data=_make_stats_data(
{
MY_COUNTER: {
("a",): 7,
("c",): 5,
},
},
metadata,
),
)
stats1.aggregate(stats2)
values = stats1.data[_StatsKey.DATA][MY_COUNTER]
assert values[("a",)] == 17
assert values[("b",)] == 3
assert values[("c",)] == 5
def test_aggregate_labeled_metric_missing_from_self():
"""Aggregating a labeled metric that self doesn't have at all yet."""
metadata = _metric_metadata()
stats1 = OffloadingConnectorStats()
stats2 = OffloadingConnectorStats(
data=_make_stats_data(
{
MY_COUNTER: {
("a",): 7,
("b",): 5,
},
},
metadata,
),
)
stats1.aggregate(stats2)
values = stats1.data[_StatsKey.DATA][MY_COUNTER]
assert values[("a",)] == 7
assert values[("b",)] == 5
assert stats1.data[_StatsKey.TYPES][MY_COUNTER] == _MetricType.COUNTER
def test_helper_methods_accept_labeled_metrics():
stats = OffloadingConnectorStats()
stats.increase_counter(MY_COUNTER, 3, ("a",))
stats.increase_counter(MY_COUNTER, 4, ("a",))
stats.set_gauge(PENDING_STORES, 2, ("b",))
stats.observe_histogram(LOOKUP_LATENCY, 0.1, ("b",))
stats.observe_histogram(LOOKUP_LATENCY, 0.2, ("b",))
values = stats.data[_StatsKey.DATA]
assert values[MY_COUNTER][("a",)] == 7
assert values[PENDING_STORES][("b",)] == 2
assert values[LOOKUP_LATENCY][("b",)] == [0.1, 0.2]
def test_aggregate_merges_types():
stats1 = OffloadingConnectorStats(
data={
_StatsKey.TYPES: {LOAD_BYTES: _MetricType.COUNTER},
_StatsKey.DATA: {LOAD_BYTES: 1},
_StatsKey.DATA: {LOAD_BYTES: {(): 1}},
},
)
stats2 = OffloadingConnectorStats(
data={
_StatsKey.TYPES: {PENDING_STORES: _MetricType.GAUGE},
_StatsKey.DATA: {PENDING_STORES: 2},
_StatsKey.DATA: {PENDING_STORES: {(): 2}},
},
)
result = stats1.aggregate(stats2)
assert result.data[_StatsKey.DATA][PENDING_STORES] == 2
assert _unlabeled(result.data[_StatsKey.DATA], PENDING_STORES) == 2
assert result.data[_StatsKey.TYPES][PENDING_STORES] == _MetricType.GAUGE
@@ -283,6 +386,26 @@ def test_reduce():
assert reduced[f"{LOOKUP_LATENCY}_sum"] == sum([0.1, 0.2, 0.3])
def test_reduce_labeled_metrics():
metadata = _metric_metadata()
stats = OffloadingConnectorStats(
data=_make_stats_data(
{
MY_COUNTER: {
("a",): 17,
("b",): 3,
},
},
metadata,
),
)
reduced = stats.reduce()
assert reduced[f"{MY_COUNTER}:{('a',)}"] == 17
assert reduced[f"{MY_COUNTER}:{('b',)}"] == 3
def test_reset():
"""Test that reset() resets all connector stats."""
metadata = _metric_metadata()
@@ -326,11 +449,11 @@ def test_prom_metrics_observes_manager_counter():
prom_metrics.observe(
{
_StatsKey.TYPES: {STORES_SKIPPED: _MetricType.COUNTER},
_StatsKey.DATA: {STORES_SKIPPED: 7},
_StatsKey.DATA: {STORES_SKIPPED: {(): 7}},
}
)
counter = prom_metrics.offloading_metrics[(0, STORES_SKIPPED)]
counter = prom_metrics.offloading_metrics[(0, STORES_SKIPPED, ())]
assert counter.increments == [7]
counter_def = prom_metrics._offloading_metric_defs[STORES_SKIPPED]
assert counter_def.kwargs["name"] == "vllm:kv_offload_stores_skipped"
@@ -360,22 +483,22 @@ def test_prom_metrics_observes_flat_transfer_metrics_and_legacy_metrics():
STORE_SIZE: _MetricType.HISTOGRAM,
},
_StatsKey.DATA: {
LOAD_BYTES: 24,
LOAD_TIME: 1.5,
LOAD_SIZE: [16, 8],
STORE_BYTES: 3,
STORE_TIME: 0.3,
STORE_SIZE: [1, 2],
LOAD_BYTES: {(): 24},
LOAD_TIME: {(): 1.5},
LOAD_SIZE: {(): [16, 8]},
STORE_BYTES: {(): 3},
STORE_TIME: {(): 0.3},
STORE_SIZE: {(): [1, 2]},
},
}
)
assert prom_metrics.offloading_metrics[(0, LOAD_BYTES)].increments == [24]
assert prom_metrics.offloading_metrics[(0, LOAD_TIME)].increments == [1.5]
assert prom_metrics.offloading_metrics[(0, LOAD_SIZE)].observed == [16, 8]
assert prom_metrics.offloading_metrics[(0, STORE_BYTES)].increments == [3]
assert prom_metrics.offloading_metrics[(0, STORE_TIME)].increments == [0.3]
assert prom_metrics.offloading_metrics[(0, STORE_SIZE)].observed == [1, 2]
assert prom_metrics.offloading_metrics[(0, LOAD_BYTES, ())].increments == [24]
assert prom_metrics.offloading_metrics[(0, LOAD_TIME, ())].increments == [1.5]
assert prom_metrics.offloading_metrics[(0, LOAD_SIZE, ())].observed == [16, 8]
assert prom_metrics.offloading_metrics[(0, STORE_BYTES, ())].increments == [3]
assert prom_metrics.offloading_metrics[(0, STORE_TIME, ())].increments == [0.3]
assert prom_metrics.offloading_metrics[(0, STORE_SIZE, ())].observed == [1, 2]
assert prom_metrics.counter_kv_bytes[(0, "CPU_to_GPU")].increments == [24]
assert prom_metrics.counter_kv_transfer_time[(0, "CPU_to_GPU")].increments == [1.5]
@@ -396,7 +519,9 @@ def test_prom_metrics_observes_manager_gauge_and_histogram():
),
}
with patch.object(
CPUOffloadingSpec, "build_metric_definitions", return_value=metric_definitions
OffloadingSpecFactory,
"get_spec_cls",
return_value=_spec_cls_with_metric_definitions(metric_definitions),
):
prom_metrics = OffloadPromMetrics(
vllm_config=_FakeVllmConfig(store_threshold=0), # type: ignore[arg-type]
@@ -416,20 +541,91 @@ def test_prom_metrics_observes_manager_gauge_and_histogram():
LOOKUP_LATENCY: _MetricType.HISTOGRAM,
},
_StatsKey.DATA: {
PENDING_STORES: 5,
LOOKUP_LATENCY: [0.2, 0.4],
PENDING_STORES: {(): 5},
LOOKUP_LATENCY: {(): [0.2, 0.4]},
},
}
)
gauge = prom_metrics.offloading_metrics[(0, PENDING_STORES)]
histogram = prom_metrics.offloading_metrics[(0, LOOKUP_LATENCY)]
gauge = prom_metrics.offloading_metrics[(0, PENDING_STORES, ())]
histogram = prom_metrics.offloading_metrics[(0, LOOKUP_LATENCY, ())]
assert gauge.set_values == [5]
assert histogram.observed == [0.2, 0.4]
histogram_def = prom_metrics._offloading_metric_defs[LOOKUP_LATENCY]
assert histogram_def.kwargs["buckets"] == (0.1, 1.0)
def test_prom_metrics_lazily_observes_labeled_metric():
metric_definitions = {
MY_COUNTER: OffloadingCounterMetadata(
documentation="counter with a label",
labelnames=(MY_LABEL,),
),
}
with patch.object(
OffloadingSpecFactory,
"get_spec_cls",
return_value=_spec_cls_with_metric_definitions(metric_definitions),
):
prom_metrics = OffloadPromMetrics(
vllm_config=_FakeVllmConfig(store_threshold=0), # type: ignore[arg-type]
metric_types={
Gauge: _FakeMetric,
Counter: _FakeMetric,
Histogram: _FakeMetric,
},
labelnames=["model_name", "engine"],
per_engine_labelvalues={0: ["model", "0"]},
)
assert (0, MY_COUNTER, ("a",)) not in prom_metrics.offloading_metrics
prom_metrics.observe(
{
_StatsKey.TYPES: {MY_COUNTER: _MetricType.COUNTER},
_StatsKey.DATA: {MY_COUNTER: {("a",): 7}},
}
)
counter = prom_metrics.offloading_metrics[(0, MY_COUNTER, ("a",))]
assert counter.increments == [7]
assert counter.labelvalues == ("model", "0", "a")
counter_def = prom_metrics._offloading_metric_defs[MY_COUNTER]
assert counter_def.kwargs["labelnames"] == ["model_name", "engine", MY_LABEL]
def test_prom_metrics_rejects_wrong_label_count():
metric_definitions = {
MY_COUNTER: OffloadingCounterMetadata(
documentation="counter with a label",
labelnames=(MY_LABEL,),
),
}
with patch.object(
OffloadingSpecFactory,
"get_spec_cls",
return_value=_spec_cls_with_metric_definitions(metric_definitions),
):
prom_metrics = OffloadPromMetrics(
vllm_config=_FakeVllmConfig(store_threshold=0), # type: ignore[arg-type]
metric_types={
Gauge: _FakeMetric,
Counter: _FakeMetric,
Histogram: _FakeMetric,
},
labelnames=["model_name", "engine"],
per_engine_labelvalues={0: ["model", "0"]},
)
with pytest.raises(AssertionError, match="expects 1 labels"):
prom_metrics.observe(
{
_StatsKey.TYPES: {MY_COUNTER: _MetricType.COUNTER},
_StatsKey.DATA: {MY_COUNTER: {("a", "extra"): 7}},
}
)
def test_prom_metrics_uses_configured_manager_metrics():
prom_metrics = OffloadPromMetrics(
vllm_config=_FakeVllmConfig(store_threshold=0), # type: ignore[arg-type]
@@ -458,9 +654,9 @@ def test_aggregate_into_empty_stats():
PENDING_STORES: _MetricType.GAUGE,
},
_StatsKey.DATA: {
LOAD_BYTES: 42,
LOAD_SIZE: [10, 20],
PENDING_STORES: 3,
LOAD_BYTES: {(): 42},
LOAD_SIZE: {(): [10, 20]},
PENDING_STORES: {(): 3},
},
},
)
@@ -469,9 +665,9 @@ def test_aggregate_into_empty_stats():
assert result is empty
values = result.data[_StatsKey.DATA]
assert values[LOAD_BYTES] == 42
assert values[LOAD_SIZE] == [10, 20]
assert values[PENDING_STORES] == 3
assert _unlabeled(values, LOAD_BYTES) == 42
assert _unlabeled(values, LOAD_SIZE) == [10, 20]
assert _unlabeled(values, PENDING_STORES) == 3
def test_prom_metrics_multi_engine_routing():
@@ -490,14 +686,13 @@ def test_prom_metrics_multi_engine_routing():
prom_metrics.observe(
{
_StatsKey.TYPES: {LOAD_BYTES: _MetricType.COUNTER},
_StatsKey.DATA: {LOAD_BYTES: 100},
_StatsKey.DATA: {LOAD_BYTES: {(): 100}},
},
engine_idx=1,
)
engine0 = prom_metrics.offloading_metrics[(0, LOAD_BYTES)]
engine1 = prom_metrics.offloading_metrics[(1, LOAD_BYTES)]
assert engine0.increments == []
assert (0, LOAD_BYTES, ()) not in prom_metrics.offloading_metrics
engine1 = prom_metrics.offloading_metrics[(1, LOAD_BYTES, ())]
assert engine1.increments == [100]
@@ -518,6 +713,6 @@ def test_prom_metrics_rejects_undeclared_metric():
prom_metrics.observe(
{
_StatsKey.TYPES: {"unknown:metric": _MetricType.COUNTER},
_StatsKey.DATA: {"unknown:metric": 1},
_StatsKey.DATA: {"unknown:metric": {(): 1}},
}
)
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Iterable
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
@@ -11,19 +11,16 @@ from tests.v1.kv_connector.unit.offloading_connector.utils import (
to_keys,
)
from tests.v1.kv_connector.unit.utils import EOS_TOKEN_ID
from vllm.distributed.kv_events import BlockRemoved, BlockStored
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.scheduler import (
OffloadingConnectorScheduler,
RequestOffloadState,
)
from vllm.v1.core.kv_cache_utils import BlockHash
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheGroupSpec,
SlidingWindowSpec,
)
from vllm.v1.kv_offload.base import (
OffloadingEvent,
OffloadingManager,
OffloadPolicy,
ReqContext,
@@ -145,31 +142,6 @@ def test_offloading_connector(request_runner, async_scheduling: bool):
runner.connector_scheduler._maximal_prefix_lookup = lambda key, req_context: 1
runner.run(decoded_tokens=[EOS_TOKEN_ID], expected_loaded=(3, 4, 5))
# test take_events
def to_hashes(int_hashes: list[int]) -> list[BlockHash]:
return [BlockHash(str(i).encode()) for i in int_hashes]
def take_events() -> Iterable[OffloadingEvent]:
yield OffloadingEvent(keys=to_keys([1, 2, 3]), medium="A", removed=False)
yield OffloadingEvent(keys=to_keys([4, 5, 6]), medium="B", removed=True)
runner.manager.take_events.side_effect = take_events
events = list(runner.scheduler_connector.take_events())
assert len(events) == 2
event = events[0]
assert isinstance(event, BlockStored)
assert event.block_hashes == to_hashes([1, 2, 3])
assert event.block_size == 0
assert event.medium == "A"
assert event.token_ids == []
assert event.parent_block_hash is None
assert event.lora_id is None
assert event.lora_name is None
event = events[1]
assert isinstance(event, BlockRemoved)
assert event.block_hashes == to_hashes([4, 5, 6])
assert event.medium == "B"
@pytest.mark.parametrize("async_scheduling", [True, False])
def test_request_preemption(request_runner, async_scheduling: bool):
@@ -1278,11 +1250,11 @@ def test_reset_cache_finalizes_finished_request_with_pending_store(
)
finalized: list[str] = []
runner.manager.on_request_finished.side_effect = (
lambda req_context: finalized.append(req_context.req_id)
runner.manager.on_request_finished.side_effect = lambda req_context: (
finalized.append(req_context.req_id)
)
runner.manager.prepare_store.side_effect = (
lambda keys, req_context: generate_store_output(keys)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
# Decode a couple of blocks and keep every transfer in flight, so the
@@ -1314,6 +1286,100 @@ def test_reset_cache_finalizes_finished_request_with_pending_store(
assert req_id not in cs._req_status
def test_pending_transfer_defers_prefix_lookup():
"""A request with an in-flight store must not issue a load on re-admission.
With async scheduling, a preempted request's store can be flushed by the
worker before the scheduler consumes its completion. If the request is
re-admitted in that window, the connector should defer it instead of
looking up offloaded blocks and later asserting when a load is queued while
the store job is still tracked.
"""
scheduler = object.__new__(OffloadingConnectorScheduler)
scheduler.manager = MagicMock(spec=OffloadingManager)
request = SimpleNamespace(request_id="req-0")
group_state = SimpleNamespace(block_ids=[1, 2, 3])
req_status = SimpleNamespace(
group_states=[group_state],
transfer_jobs={123},
)
scheduler._req_status = {request.request_id: req_status}
matched_tokens, is_async = scheduler.get_num_new_matched_tokens(
request,
num_computed_tokens=0,
)
assert matched_tokens is None
assert is_async is False
assert group_state.block_ids == []
scheduler.manager.lookup.assert_not_called()
def test_async_preempt_readmit_before_transfer_output_is_deferred(request_runner):
"""A preempted request can be scheduled again before flush output is read.
EngineCore.step_with_batch_queue() may schedule a new batch while a prior
preemption batch is still queued. The store completion from jobs_to_flush is
only cleared when that queued output reaches update_from_output(), so the
re-admission path must defer while the scheduler still tracks the store.
"""
block_size = 4
block_size_factor = 3
offloaded_block_size = block_size * block_size_factor
runner = request_runner(
block_size=block_size,
num_gpu_blocks=100,
async_scheduling=True,
block_size_factor=block_size_factor,
)
free_block_queue = runner.scheduler.kv_cache_manager.block_pool.free_block_queue
num_free_blocks_empty = free_block_queue.num_free_blocks
req_id = "0"
runner.new_request(token_ids=[0] * offloaded_block_size * 2)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.run(decoded_tokens=[0], complete_transfers=False)
runner.run(
decoded_tokens=[0] * (2 * offloaded_block_size - block_size),
complete_transfers=False,
)
req_status = runner.connector_scheduler._req_status[req_id]
pending_store_jobs = set(req_status.transfer_jobs)
assert pending_store_jobs
assert all(
runner.connector_scheduler._jobs[jid].is_store for jid in pending_store_jobs
)
free_block_queue.num_free_blocks = 0
preempt_output = runner.scheduler.schedule()
assert preempt_output.preempted_req_ids == {req_id}
assert preempt_output.kv_connector_metadata is not None
assert pending_store_jobs <= preempt_output.kv_connector_metadata.jobs_to_flush
assert req_status.transfer_jobs == pending_store_jobs
# Simulate the async batch-queue window: schedule again before the
# preemption batch's ModelRunnerOutput is consumed by update_from_output().
free_block_queue.num_free_blocks = num_free_blocks_empty
assert runner.scheduler.reset_prefix_cache()
runner.connector_scheduler._maximal_prefix_lookup = lambda key, req_context: len(
key
)
readmit_output = runner.scheduler.schedule()
assert readmit_output.num_scheduled_tokens == {}
assert readmit_output.kv_connector_metadata is not None
assert readmit_output.kv_connector_metadata.load_jobs == {}
assert req_status.transfer_jobs == pending_store_jobs
@pytest.mark.parametrize("async_scheduling", [True, False])
def test_swa_alignment_skip(request_runner, async_scheduling: bool):
"""SWA blocks unreachable by the load path are skipped during store.
@@ -14,7 +14,12 @@ from tests.v1.kv_connector.unit.utils import (
create_vllm_config,
)
from vllm import SamplingParams
from vllm.config import KVTransferConfig, VllmConfig, set_current_vllm_config
from vllm.config import (
KVEventsConfig,
KVTransferConfig,
VllmConfig,
set_current_vllm_config,
)
from vllm.distributed.kv_transfer.kv_connector.v1 import KVConnectorRole
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.common import (
OffloadingConnectorMetadata,
@@ -198,6 +203,9 @@ class RequestRunner:
"spec_module_path": "tests.v1.kv_connector.unit.offloading_connector.utils", # noqa: E501
# Preserve legacy behavior for tests; new opt-in tests override.
"offload_prompt_only": False,
# Exercise the self-describing KV events path by default;
# opt-out tests override this to cover the legacy placeholders.
"self_describing_kv_events": True,
}
if block_size_factor > 1:
extra_config["block_size"] = block_size * block_size_factor
@@ -209,6 +217,13 @@ class RequestRunner:
kv_role="kv_both",
kv_connector_extra_config=extra_config,
)
vllm_config.kv_events_config = KVEventsConfig(
# Enable so the offloading events tracker is active, but use the
# null publisher: these tests drain take_events directly and a
# real ZMQ publisher would bind a port per test.
enable_kv_cache_events=True,
publisher="null",
)
if kv_cache_groups is None:
kv_cache_groups = [
@@ -36,13 +36,33 @@ from vllm.utils.network_utils import (
get_ip,
make_zmq_path,
)
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheConfig,
KVCacheGroupSpec,
KVCacheTensor,
)
from .utils import create_request, create_scheduler
def _make_test_kv_cache_config() -> KVCacheConfig:
return KVCacheConfig(num_blocks=0, kv_cache_tensors=[], kv_cache_groups=[])
layer_names = ["layer0", "layer1", "layer2"]
return KVCacheConfig(
num_blocks=2,
kv_cache_tensors=[KVCacheTensor(size=0, shared_by=layer_names)],
kv_cache_groups=[
KVCacheGroupSpec(
layer_names=layer_names,
kv_cache_spec=FullAttentionSpec(
block_size=16,
num_kv_heads=4,
head_size=64,
dtype=torch.float16,
),
)
],
)
aiter_available = importlib.util.find_spec("aiter") is not None
@@ -175,9 +195,18 @@ class FakeMoRIIOConnectorWorker(MoRIIOConnectorWorker):
REMOTE_ENGINE_ID = "remote_engine"
def __init__(
self, *args, hand_shake_latency: float = 1.8, kv_cache_layout="HND", **kwargs
self,
vllm_config,
engine_id,
*args,
hand_shake_latency: float = 1.8,
kv_cache_layout="HND",
kv_cache_config=None,
**kwargs,
):
super().__init__(*args, **kwargs)
super().__init__(
vllm_config, engine_id, kv_cache_config or _make_test_kv_cache_config()
)
def create_vllm_config(
@@ -0,0 +1,228 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import importlib.util
from types import SimpleNamespace
import pytest
import torch
from vllm.platforms import current_platform
from vllm.v1.kv_cache_interface import FullAttentionSpec, MLAAttentionSpec
aiter_available = importlib.util.find_spec("aiter") is not None
mori_available = importlib.util.find_spec("mori") is not None
if not (current_platform.is_rocm() and mori_available):
pytest.skip(
"MoRIIOs are only available on ROCm with mori package installed",
allow_module_level=True,
)
moriio_layout = importlib.import_module(
"vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_layout"
)
def _full_spec(block_size: int = 4) -> FullAttentionSpec:
return FullAttentionSpec(
block_size=block_size,
num_kv_heads=2,
head_size=3,
dtype=torch.bfloat16,
)
def _mla_spec(block_size: int = 4) -> MLAAttentionSpec:
return MLAAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=3,
dtype=torch.bfloat16,
)
def _worker(
kv_caches: dict[str, torch.Tensor],
layer_to_spec: dict[str, object],
num_blocks: int = 8,
) -> SimpleNamespace:
return SimpleNamespace(
kv_caches=kv_caches,
layer_to_spec=layer_to_spec,
num_blocks=num_blocks,
block_size=4,
)
def _remote_meta(num_blocks: int = 16) -> SimpleNamespace:
return SimpleNamespace(num_blocks=num_blocks)
def test_separated_kv_layout_uses_kv_axis_zero_and_block_axis_one():
cache = torch.empty((2, 8, 4, 2, 3), dtype=torch.bfloat16)
worker = _worker({"layer": cache}, {"layer": _full_spec()})
geometry = moriio_layout.get_layer_transfer_geometry(
"layer", cache, worker.layer_to_spec, remote_num_blocks=16
)
assert geometry.block_stride == 24
assert geometry.local_kv_stride == 192
assert geometry.remote_kv_stride == 384
assert geometry.split_kv_regions
assert moriio_layout.compute_block_transfer_offsets(
"layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks
) == ([48, 144, 432, 528], [192, 240, 960, 1008], [48, 48, 48, 48])
def test_interleaved_kv_layout_uses_block_axis_zero_and_kv_axis_one():
cache = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16)
worker = _worker({"layer": cache}, {"layer": _full_spec()})
geometry = moriio_layout.get_layer_transfer_geometry(
"layer", cache, worker.layer_to_spec, remote_num_blocks=16
)
assert geometry.block_stride == 48
assert geometry.local_kv_stride == 24
assert geometry.remote_kv_stride == 24
assert not geometry.split_kv_regions
assert moriio_layout.compute_block_transfer_offsets(
"layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks
) == ([96, 288], [384, 480], [96, 96])
def test_mla_key_only_layout_transfers_one_slab_per_block():
cache = torch.empty((8, 4, 3), dtype=torch.bfloat16)
worker = _worker({"layer": cache}, {"layer": _mla_spec()})
geometry = moriio_layout.get_layer_transfer_geometry(
"layer", cache, worker.layer_to_spec, remote_num_blocks=16
)
assert geometry.block_stride == 12
assert geometry.local_kv_stride is None
assert geometry.remote_kv_stride is None
assert geometry.transfers_per_block == 1
assert moriio_layout.compute_block_transfer_offsets(
"layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks
) == ([24, 72], [96, 120], [24, 24])
def test_mixed_layers_compute_distinct_offsets_per_layer():
kv_caches = {
"separated": torch.empty((2, 8, 4, 2, 3), dtype=torch.bfloat16),
"interleaved": torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16),
"indexer": torch.empty((8, 4, 3), dtype=torch.bfloat16),
}
worker = _worker(
kv_caches,
{
"separated": _full_spec(),
"interleaved": _full_spec(),
"indexer": _mla_spec(),
},
)
separated = moriio_layout.compute_block_transfer_offsets(
"separated",
kv_caches["separated"],
worker.layer_to_spec,
[1, 3],
[4, 5],
_remote_meta().num_blocks,
)
interleaved = moriio_layout.compute_block_transfer_offsets(
"interleaved",
kv_caches["interleaved"],
worker.layer_to_spec,
[1, 3],
[4, 5],
_remote_meta().num_blocks,
)
indexer = moriio_layout.compute_block_transfer_offsets(
"indexer",
kv_caches["indexer"],
worker.layer_to_spec,
[1, 3],
[4, 5],
_remote_meta().num_blocks,
)
assert separated != interleaved
assert separated != indexer
assert interleaved != indexer
def test_block_id_length_mismatch_raises_value_error():
cache = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16)
worker = _worker({"layer": cache}, {"layer": _full_spec()})
with pytest.raises(ValueError, match="must have the same length"):
moriio_layout.compute_block_transfer_offsets(
"layer", cache, worker.layer_to_spec, [1, 3], [4], _remote_meta().num_blocks
)
def test_registration_regions_do_not_split_interleaved_or_mla_cache():
separated = torch.empty((2, 8, 4, 2, 3), dtype=torch.bfloat16)
interleaved = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16)
indexer = torch.empty((8, 4, 3), dtype=torch.bfloat16)
worker = _worker(
{
"separated": separated,
"interleaved": interleaved,
"indexer": indexer,
},
{
"separated": _full_spec(),
"interleaved": _full_spec(),
"indexer": _mla_spec(),
},
)
separated_regions = moriio_layout.iter_layer_registration_regions(
"separated", separated, worker.layer_to_spec
)
interleaved_regions = moriio_layout.iter_layer_registration_regions(
"interleaved", interleaved, worker.layer_to_spec
)
indexer_regions = moriio_layout.iter_layer_registration_regions(
"indexer", indexer, worker.layer_to_spec
)
assert [region[0].data_ptr() for region in separated_regions] == [
separated[0].data_ptr(),
separated[1].data_ptr(),
]
assert separated_regions[0][1] == 8 * 48
assert separated_regions[1][1] == 8 * 48
assert len(interleaved_regions) == 1
assert interleaved_regions[0][0].data_ptr() == interleaved.data_ptr()
assert interleaved_regions[0][1] == 8 * 2 * 48
assert len(indexer_regions) == 1
assert indexer_regions[0][0].data_ptr() == indexer.data_ptr()
assert indexer_regions[0][1] == 8 * 24
def test_registration_regions_use_layer_num_blocks():
cache = torch.empty((4, 2, 4, 2, 3), dtype=torch.bfloat16)
worker = _worker({"layer": cache}, {"layer": _full_spec()}, num_blocks=8)
regions = moriio_layout.iter_layer_registration_regions(
"layer", cache, worker.layer_to_spec
)
assert len(regions) == 1
assert regions[0][1] == 4 * 2 * 48
def test_unsupported_shape_raises_value_error():
cache = torch.empty((8, 4, 2, 3), dtype=torch.bfloat16)
worker = _worker({"layer": cache}, {"layer": _full_spec()})
with pytest.raises(ValueError, match="Unsupported MoRIIO K/V cache shape"):
moriio_layout.get_layer_transfer_geometry("layer", cache, worker.layer_to_spec)
@@ -219,6 +219,7 @@ def test_multi_example_connector_consistency():
enforce_eager=True,
gpu_memory_utilization=0.5,
kv_transfer_config=kv_transfer_config,
async_scheduling=False,
)
# Run generation - this should trigger saving KV cache
# Use a single prompt to avoid race conditions depending on the order of scheduling
@@ -138,6 +138,10 @@ def _wait_for_prefix_cache_reset(llm: LLM) -> None:
def _latency_test(llm: LLM, subscriber: MockSubscriber | None):
# TODO: Reintroduce latency test on ROCm once MRV2 supports cross
# layer KV Cache. See https://github.com/vllm-project/vllm/pull/45947
if current_platform.is_rocm():
return
sampling_params = SamplingParams(max_tokens=1)
num_times_cpu_better_than_cold = 0
+8 -8
View File
@@ -294,25 +294,25 @@ def test_cpu_manager():
# prepare store with no space ([2, 3] is being loaded)
assert cpu_manager.prepare_store(to_keys([6, 7, 8]), _EMPTY_REQ_CTX) is None
# complete load [2, 3]
# complete load [2, 3]. Load changes the eviction list, making 2, 3 recent.
cpu_manager.complete_load(to_keys([2, 3]), _EMPTY_REQ_CTX)
# prepare store [6, 7, 8] -> evicts [2, 3, 4] (oldest)
# prepare store [6, 7, 8] -> evicts [4, 5, 2] (oldest)
prepare_store_output = cpu_manager.prepare_store(to_keys([6, 7, 8]), _EMPTY_REQ_CTX)
verify_store_output(
prepare_store_output,
ExpectedPrepareStoreOutput(
keys_to_store=[6, 7, 8],
store_block_ids=[3, 2, 1],
evicted_keys=[2, 3, 4],
store_block_ids=[1, 0, 3],
evicted_keys=[4, 5, 2],
),
)
# complete store [6, 7, 8]
cpu_manager.complete_store(to_keys([6, 7, 8]), _EMPTY_REQ_CTX)
# touch [5, 6, 7] (move to end of LRU order)
cpu_manager.touch(to_keys([5, 6, 7]), _EMPTY_REQ_CTX)
# touch [3, 6, 7] (move to end of LRU order)
cpu_manager.touch(to_keys([3, 6, 7]), _EMPTY_REQ_CTX)
# prepare store [7, 9] -> evicts [8] (oldest following previous touch)
prepare_store_output = cpu_manager.prepare_store(to_keys([9]), _EMPTY_REQ_CTX)
@@ -320,7 +320,7 @@ def test_cpu_manager():
prepare_store_output,
ExpectedPrepareStoreOutput(
keys_to_store=[9],
store_block_ids=[1],
store_block_ids=[3],
evicted_keys=[8],
),
)
@@ -335,7 +335,7 @@ def test_cpu_manager():
verify_events(
cpu_manager.take_events(),
expected_stores=({3, 4, 5}, {6, 7, 8}),
expected_evictions=({2, 3, 4}, {8}),
expected_evictions=({4, 5, 2}, {8}),
)
@@ -295,6 +295,8 @@ class TestTieringOffloadingManager:
self.manager.prepare_store(blocks, _CTX)
self.manager.complete_store(blocks, _CTX, success=True)
self._simulate_on_schedule_end()
# for secondary tiers to drain jobs, so primary tier's blocks are evictable.
self._simulate_on_schedule_end()
self.secondary_tier1.touch = MagicMock(wraps=self.secondary_tier1.touch)
self.secondary_tier2.touch = MagicMock(wraps=self.secondary_tier2.touch)
@@ -303,7 +305,7 @@ class TestTieringOffloadingManager:
self.manager.touch(blocks, _CTX)
# Verify touch was called on primary tier (check LRU order)
primary_keys = list(self.primary_tier._policy.blocks.keys())
primary_keys = list(self.primary_tier._policy.evictable_blocks.keys())
assert primary_keys[-3:] == list(reversed(blocks))
# Verify touch was propagated to all secondary tiers
@@ -53,9 +53,39 @@ PEAGLE_CONFIG = SpeculatorTestConfig(
parallel_drafting=True,
)
QWEN3_EAGLE3_CONFIG = SpeculatorTestConfig(
model_path=(
"inference-optimization/"
"Qwen3-8B-from-Qwen3-8B_regen-speculators.eagle3-qwen3arch-ckpt1"
),
method="eagle3",
display_name="Qwen3 Eagle3",
expected_gsm8k_accuracy=0.88,
accuracy_rtol=0.05,
expected_acceptance_len=2.67,
acceptance_len_rtol=0.10,
expected_per_pos_acceptance_rates=(0.76, 0.55, 0.36),
per_pos_rtol=0.10,
)
QWEN3_PEAGLE_CONFIG = SpeculatorTestConfig(
model_path="inference-optimization/Qwen3-8B-speculators.peagle-qwen3arch-ckpt4",
method="eagle3",
display_name="Qwen3 PEagle",
expected_gsm8k_accuracy=0.88,
accuracy_rtol=0.05,
expected_acceptance_len=3.42,
acceptance_len_rtol=0.15,
expected_per_pos_acceptance_rates=(0.78, 0.59, 0.43, 0.29, 0.18, 0.10, 0.05),
per_pos_rtol=0.10,
parallel_drafting=True,
)
SPECULATOR_CONFIGS = [
pytest.param(DFLASH_CONFIG, id="dflash"),
pytest.param(PEAGLE_CONFIG, id="peagle"),
pytest.param(QWEN3_EAGLE3_CONFIG, id="qwen3arch_eagle3"),
pytest.param(QWEN3_PEAGLE_CONFIG, id="qwen3arch_peagle"),
]
@@ -176,6 +206,7 @@ def test_speculators_correctness(monkeypatch, config):
results = evaluate_gsm8k_offline(spec_llm)
accuracy = results["accuracy"]
print(f"GSM8K Accuracy: {accuracy:.4f}")
accuracy_threshold = config.expected_gsm8k_accuracy * (1 - config.accuracy_rtol)
assert accuracy >= accuracy_threshold, (
f"Expected GSM8K accuracy >= {accuracy_threshold:.3f}, got {accuracy:.3f}"
+242
View File
@@ -0,0 +1,242 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
from vllm.v1.kv_cache_interface import FullAttentionSpec, KVQuantMode
from vllm.v1.worker.gpu.attn_utils import _reshape_kv_cache
from vllm.v1.worker.utils import AttentionGroup
class FakeFlashAttentionBackend:
@staticmethod
def get_kv_cache_shape(
num_blocks: int,
block_size: int,
num_kv_heads: int,
head_size: int,
cache_dtype_str: str = "auto",
) -> tuple[int, ...]:
return (num_blocks, 2, block_size, num_kv_heads, head_size)
@staticmethod
def get_kv_cache_stride_order(
include_num_layers_dimension: bool = False,
) -> tuple[int, ...]:
assert not include_num_layers_dimension
return (0, 1, 2, 3, 4)
class FakeHNDFlashAttentionBackend(FakeFlashAttentionBackend):
@staticmethod
def get_kv_cache_stride_order(
include_num_layers_dimension: bool = False,
) -> tuple[int, ...]:
assert not include_num_layers_dimension
return (0, 1, 3, 2, 4)
def test_reshape_padded_flash_attention_kv_cache_strides_by_page():
num_blocks = 3
spec = FullAttentionSpec(
block_size=16,
num_kv_heads=1,
head_size=2,
dtype=torch.float32,
page_size_padded=384,
)
assert spec.real_page_size_bytes == 256
raw_tensors = {
"layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8)
}
attn_groups = [
AttentionGroup(
backend=FakeFlashAttentionBackend,
layer_names=["layer"],
kv_cache_spec=spec,
kv_cache_group_id=0,
)
]
kv_cache = _reshape_kv_cache(
attn_groups,
raw_tensors,
"auto",
[spec.block_size],
{},
)["layer"]
assert kv_cache.shape == (num_blocks, 2, 16, 1, 2)
assert kv_cache.stride(0) == spec.page_size_bytes // 4
assert kv_cache.stride(1) == spec.real_page_size_bytes // 2 // 4
assert kv_cache[1, 0].storage_offset() == spec.page_size_bytes // 4
assert (
kv_cache[1, 1].storage_offset()
== (spec.page_size_bytes + spec.real_page_size_bytes // 2) // 4
)
def test_reshape_padded_hnd_flash_attention_kv_cache_strides_by_page():
num_blocks = 3
spec = FullAttentionSpec(
block_size=16,
num_kv_heads=3,
head_size=2,
dtype=torch.float32,
page_size_padded=1024,
)
assert spec.real_page_size_bytes == 768
raw_tensors = {
"layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8)
}
attn_groups = [
AttentionGroup(
backend=FakeHNDFlashAttentionBackend,
layer_names=["layer"],
kv_cache_spec=spec,
kv_cache_group_id=0,
)
]
kv_cache = _reshape_kv_cache(
attn_groups,
raw_tensors,
"auto",
[spec.block_size],
{},
)["layer"]
assert kv_cache.shape == (num_blocks, 2, 16, 3, 2)
assert kv_cache.stride(0) == spec.page_size_bytes // 4
assert kv_cache.stride(1) == spec.real_page_size_bytes // 2 // 4
assert kv_cache.stride(2) == 2
assert kv_cache.stride(3) == spec.block_size * spec.head_size
assert kv_cache[1, 0].storage_offset() == spec.page_size_bytes // 4
assert (
kv_cache[1, 1].storage_offset()
== (spec.page_size_bytes + spec.real_page_size_bytes // 2) // 4
)
assert (
kv_cache[1, 1, 3, 2].storage_offset()
== (
spec.page_size_bytes
+ spec.real_page_size_bytes // 2
+ 3 * spec.head_size * 4
+ 2 * spec.block_size * spec.head_size * 4
)
// 4
)
class FakeDiffKVBackend:
@staticmethod
def get_kv_cache_shape(
num_blocks: int,
block_size: int,
num_kv_heads: int,
head_size: int,
cache_dtype_str: str = "auto",
) -> tuple[int, ...]:
return (num_blocks, block_size, num_kv_heads, head_size * 2)
@staticmethod
def get_kv_cache_stride_order(
include_num_layers_dimension: bool = False,
) -> tuple[int, ...]:
assert not include_num_layers_dimension
return (0, 1, 2, 3)
def test_reshape_padded_diff_kv_cache_does_not_infer_kv_dim():
num_blocks = 3
spec = FullAttentionSpec(
block_size=16,
num_kv_heads=1,
head_size=2,
dtype=torch.float32,
page_size_padded=384,
)
raw_tensors = {
"layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8)
}
attn_groups = [
AttentionGroup(
backend=FakeDiffKVBackend,
layer_names=["layer"],
kv_cache_spec=spec,
kv_cache_group_id=0,
)
]
kv_cache = _reshape_kv_cache(
attn_groups,
raw_tensors,
"auto",
[spec.block_size],
{},
)["layer"]
assert kv_cache.shape == (num_blocks, 16, 1, 4)
assert kv_cache.stride(0) == spec.page_size_bytes // 4
assert kv_cache.stride(1) == 4
class FakePerTokenScaleBackend:
@staticmethod
def get_kv_cache_shape(
num_blocks: int,
block_size: int,
num_kv_heads: int,
head_size: int,
cache_dtype_str: str = "auto",
) -> tuple[int, ...]:
return (num_blocks, 2, block_size, num_kv_heads, head_size + 4)
@staticmethod
def get_kv_cache_stride_order(
include_num_layers_dimension: bool = False,
) -> tuple[int, ...]:
assert not include_num_layers_dimension
return (0, 1, 2, 3, 4)
def test_reshape_padded_quantized_kv_cache_preserves_scale_stride():
num_blocks = 3
spec = FullAttentionSpec(
block_size=16,
num_kv_heads=1,
head_size=4,
dtype=torch.int8,
kv_quant_mode=KVQuantMode.INT8_PER_TOKEN_HEAD,
page_size_padded=384,
)
assert spec.real_page_size_bytes == 128
assert spec.page_size_bytes == 384
raw_tensors = {
"layer": torch.zeros(spec.page_size_bytes * num_blocks, dtype=torch.int8)
}
attn_groups = [
AttentionGroup(
backend=FakePerTokenScaleBackend,
layer_names=["layer"],
kv_cache_spec=spec,
kv_cache_group_id=0,
)
]
kv_cache = _reshape_kv_cache(
attn_groups,
raw_tensors,
"int8_per_token_head",
[spec.block_size],
{},
)["layer"]
assert kv_cache.shape == (num_blocks, 2, 16, 1, 8)
assert kv_cache.stride(0) == spec.page_size_bytes
assert kv_cache.stride(1) == 16 * 1 * 8
assert kv_cache[1, 1].storage_offset() == spec.page_size_bytes + 16 * 1 * 8
@@ -47,6 +47,14 @@ from vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_engine import (
MoRIIOWrapper,
MoRIIOWriter,
)
from vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_layout import (
LayerTransferGeometry,
build_layer_to_spec,
compute_block_transfer_offsets,
get_layer_transfer_geometry,
is_mla_cache_layer,
iter_layer_registration_regions,
)
from vllm.distributed.parallel_state import (
get_tensor_model_parallel_world_size,
get_tp_group,
@@ -71,6 +79,7 @@ if TYPE_CHECKING:
logger = init_logger(__name__)
try:
from mori.io import (
BackendType,
@@ -117,7 +126,9 @@ class MoRIIOConnector(KVConnectorBase_V1):
self.connector_worker: MoRIIOConnectorWorker | None = None
elif role == KVConnectorRole.WORKER:
self.connector_scheduler = None
self.connector_worker = MoRIIOConnectorWorker(vllm_config, self.engine_id)
self.connector_worker = MoRIIOConnectorWorker(
vllm_config, self.engine_id, kv_cache_config
)
logger.info(
"Initialized MoRIIO Connector,engine_id:%s,role: %s",
self.engine_id,
@@ -683,7 +694,12 @@ class MoRIIOConnectorScheduler:
class MoRIIOConnectorWorker:
"""Implementation of Worker side methods"""
def __init__(self, vllm_config: VllmConfig, engine_id: str):
def __init__(
self,
vllm_config: VllmConfig,
engine_id: str,
kv_cache_config: "KVCacheConfig",
):
if not is_moriio_available():
raise RuntimeError(
"MoRIIO is not available. Please ensure the 'mori' package "
@@ -707,6 +723,7 @@ class MoRIIOConnectorWorker:
)
self.kv_transfer_config = vllm_config.kv_transfer_config
self.is_producer = self.kv_transfer_config.is_kv_producer
self.layer_to_spec = build_layer_to_spec(kv_cache_config)
if self.is_producer:
set_role(ROLE.PRODUCER)
@@ -809,6 +826,8 @@ class MoRIIOConnectorWorker:
self.kv_cache_shape = None
self.block_shape = None
self.kv_element_size = 0
self.kv_cache_shapes: dict[str, torch.Size] = {}
self.block_lens: dict[str, int] = {}
# Map of engine_id -> {agent_name0, agent_name1..}.
self._remote_agents: dict[EngineId, set[str]] = {}
@@ -1218,51 +1237,86 @@ class MoRIIOConnectorWorker:
all_done_future = self._handshake_initiation_executor.submit(wait_all_dp)
all_done_future.add_done_callback(request_ready)
def _is_mla_cache_layer(self, layer_name: str) -> bool:
return is_mla_cache_layer(self.layer_to_spec, layer_name)
def _get_layer_transfer_geometry(
self, layer_name: str, remote_num_blocks: int | None = None
) -> LayerTransferGeometry:
return get_layer_transfer_geometry(
layer_name,
self.kv_caches[layer_name],
self.layer_to_spec,
remote_num_blocks,
)
def _iter_layer_registration_regions(
self, layer_name: str
) -> list[tuple[torch.Tensor, int]]:
return iter_layer_registration_regions(
layer_name,
self.kv_caches[layer_name],
self.layer_to_spec,
)
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
"""Register the KV Cache data in moriio."""
_, first_kv_cache = next(iter(kv_caches.items()))
self.kv_caches = kv_caches # layer name to kv cache
self.kv_cache_shapes = {
layer_name: kv_cache.shape for layer_name, kv_cache in kv_caches.items()
}
first_layer_name, first_kv_cache = next(
(
(layer_name, kv_cache)
for layer_name, kv_cache in kv_caches.items()
if (
not self._is_mla_cache_layer(layer_name)
and len(kv_cache.shape) == 5
and (kv_cache.shape[0] == 2 or kv_cache.shape[1] == 2)
)
),
next(iter(kv_caches.items())),
)
kv_elem_size = first_kv_cache.element_size()
use_mla = len(first_kv_cache.shape) == 3
assert use_mla == self.use_mla
use_mla = self._is_mla_cache_layer(first_layer_name)
first_geometry = self._get_layer_transfer_geometry(first_layer_name)
if use_mla:
# MLA case.
self.num_blocks = first_kv_cache.shape[0]
block_rank = 2 # [block_size, latent_dim]
block_shape = first_kv_cache.shape[-block_rank:]
block_size, kv_latent_dim = block_shape
self.slot_size_bytes = kv_elem_size * kv_latent_dim
else:
# [2 (k and v), num_blocks, ...]
self.num_blocks = first_kv_cache.shape[1]
# [2, num_blocks, ...] or [num_blocks, 2, ...]
block_rank = 3 # [block_size, kv_heads, head_dim]
block_shape = first_kv_cache.shape[-block_rank:]
block_size, n_kv_heads, head_dim = block_shape[-3:]
# head size in bytes.
self.slot_size_bytes = (
kv_elem_size * n_kv_heads * head_dim
) # 1 token 1 layer size , slot size
assert block_size == self.block_size
self.num_blocks = first_geometry.num_blocks
self.slot_size_bytes = first_geometry.slot_size_bytes
assert first_geometry.block_size == self.block_size
# TODO(tms): self.block_len needs to be per-layer for sliding window,
# hybrid attn, etc
# block size in bytes
self.block_len = kv_elem_size * math.prod(block_shape)
self.block_len = first_geometry.block_len
self.kv_cache_shape = first_kv_cache.shape
self.block_shape = block_shape
self.kv_element_size = kv_elem_size
self.dst_num_blocks[self.engine_id] = self.num_blocks
self.kv_caches = kv_caches # layer name to kv cache
kv_caches_base_addr = []
caches_data = []
for cache_or_caches in kv_caches.values():
cache_list = [cache_or_caches] if use_mla else cache_or_caches
for cache in cache_list:
for layer_name in kv_caches:
geometry = self._get_layer_transfer_geometry(layer_name)
if geometry.block_size != self.block_size:
raise ValueError(
"MoRIIO KV cache block size mismatch for layer "
f"{layer_name}: {geometry.block_size} != {self.block_size}"
)
self.block_lens[layer_name] = geometry.block_len
for cache, region_len in self._iter_layer_registration_regions(layer_name):
base_addr = cache.data_ptr()
region_len = self.num_blocks * self.block_len
caches_data.append((base_addr, region_len, cache.device.index, ""))
kv_caches_base_addr.append(base_addr)
@@ -1275,7 +1329,9 @@ class MoRIIOConnectorWorker:
moriio_mem_metadata
)
self.local_kv_cache_size.append(cache.nelement() * cache.element_size())
self.local_kv_cache_size.append(
kv_cache.nelement() * kv_cache.element_size()
)
self.kv_caches_base_addr[self.engine_id] = kv_caches_base_addr
self.num_regions = len(caches_data)
@@ -1666,47 +1722,17 @@ class MoRIIOConnectorWorker:
Returns:
Tuple of (local_offsets, remote_offsets, transfer_sizes)
"""
assert self.kv_cache_shape is not None, "KV caches shape not initialized"
is_mla = len(self.kv_cache_shape) == 3
stride = self.kv_caches[layer_name].stride()
sz = self.kv_caches[layer_name].element_size()
if is_mla:
blknum, blksize, hs = self.kv_cache_shape
hn = 1
block_stride = stride[0]
else:
_, blknum, blksize, hn, hs = self.kv_cache_shape
local_ktov_stride = stride[0]
block_stride = stride[1]
remote_ktov_stride = block_stride * remote_moriio_meta.num_blocks
transfer_size_byte = blksize * hn * hs * sz
per_block = 1 if is_mla else 2
total = len(local_block_ids) * per_block
offset_local = [0] * total
offset_remote = [0] * total
sizes = [transfer_size_byte] * total
w = 0
for i, lb in enumerate(local_block_ids):
rb = remote_block_ids[i]
# K
offset_local[w] = sz * (lb * block_stride)
offset_remote[w] = sz * (rb * block_stride)
w += 1
if not is_mla:
# V
# Handle num_block variations originating from PD (different kv strides)
# TODO: address block_sz differences in heterogeneous TP scenarios
# In MLA, we don't need to consider these two cases.
offset_local[w] = sz * (1 * local_ktov_stride + lb * block_stride)
offset_remote[w] = sz * (1 * remote_ktov_stride + rb * block_stride)
w += 1
merged_l, merged_r, merged_s = self.merge_contiguous_blocks(
offset_local, offset_remote, sizes, assume_sorted=False
return compute_block_transfer_offsets(
layer_name=layer_name,
kv_cache=self.kv_caches[layer_name],
layer_to_spec=self.layer_to_spec,
local_block_ids=local_block_ids,
remote_block_ids=remote_block_ids,
remote_num_blocks=remote_moriio_meta.num_blocks,
merge_fn=lambda local, remote, sizes: self.merge_contiguous_blocks(
local, remote, sizes, assume_sorted=False
),
)
return merged_l, merged_r, merged_s
def _read_blocks(
self,
@@ -1724,15 +1750,13 @@ class MoRIIOConnectorWorker:
dp0_engine_id = self.get_engine_name_with_dp(dst_engine_id, 0)
sessions, remote_moriio_meta = self._get_built_session(dp0_engine_id)
first_layer = list(self.layer_name_to_local_kv_cache_metadata.keys())[0]
offs = self._compute_block_transfer_offsets(
first_layer, local_block_ids, remote_block_ids, remote_moriio_meta
)
for layer_name in self.layer_name_to_local_kv_cache_metadata:
sess_idx = list(self.layer_name_to_local_kv_cache_metadata.keys()).index(
layer_name
)
offs = self._compute_block_transfer_offsets(
layer_name, local_block_ids, remote_block_ids, remote_moriio_meta
)
# TODO : apply multi-session batch-read when moriio support it
transfer_status = self.moriio_wrapper.read_remote_data(
offs[2], offs[0], offs[1], sessions[sess_idx]
@@ -0,0 +1,213 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Callable, Mapping
from typing import NamedTuple
import torch
from vllm.v1.kv_cache_interface import (
KVCacheConfig,
KVCacheSpec,
MLAAttentionSpec,
SlidingWindowMLASpec,
UniformTypeKVCacheSpecs,
)
class LayerTransferGeometry(NamedTuple):
num_blocks: int
block_size: int
block_len: int
slot_size_bytes: int
block_stride: int
local_kv_stride: int | None
remote_kv_stride: int | None
transfers_per_block: int
regions_per_block: int
split_kv_regions: bool
def build_layer_to_spec(kv_cache_config: KVCacheConfig) -> dict[str, KVCacheSpec]:
layer_to_spec: dict[str, KVCacheSpec] = {}
for group in kv_cache_config.kv_cache_groups:
group_spec = group.kv_cache_spec
if isinstance(group_spec, UniformTypeKVCacheSpecs):
layer_to_spec.update(
{
layer_name: group_spec.kv_cache_specs[layer_name]
for layer_name in group.layer_names
}
)
else:
layer_to_spec.update(
{layer_name: group_spec for layer_name in group.layer_names}
)
return layer_to_spec
def is_mla_cache_layer(
layer_to_spec: Mapping[str, KVCacheSpec], layer_name: str
) -> bool:
try:
spec = layer_to_spec[layer_name]
except KeyError as e:
raise ValueError(f"Missing KV cache spec for layer {layer_name}") from e
return isinstance(spec, (MLAAttentionSpec, SlidingWindowMLASpec))
def get_layer_transfer_geometry(
layer_name: str,
kv_cache: torch.Tensor,
layer_to_spec: Mapping[str, KVCacheSpec],
remote_num_blocks: int | None = None,
) -> LayerTransferGeometry:
shape = kv_cache.shape
stride = kv_cache.stride()
element_size = kv_cache.element_size()
is_mla_cache = is_mla_cache_layer(layer_to_spec, layer_name)
if is_mla_cache and len(shape) == 3:
num_blocks, block_size, latent_dim = shape
slot_size_bytes = latent_dim * element_size
block_len = block_size * slot_size_bytes
return LayerTransferGeometry(
num_blocks=num_blocks,
block_size=block_size,
block_len=block_len,
slot_size_bytes=slot_size_bytes,
block_stride=stride[0],
local_kv_stride=None,
remote_kv_stride=None,
transfers_per_block=1,
regions_per_block=1,
split_kv_regions=False,
)
if not is_mla_cache and len(shape) == 5 and shape[0] == 2:
_, num_blocks, block_size, num_kv_heads, head_dim = shape
slot_size_bytes = num_kv_heads * head_dim * element_size
block_len = block_size * slot_size_bytes
remote_kv_stride = stride[1] * (remote_num_blocks or num_blocks)
return LayerTransferGeometry(
num_blocks=num_blocks,
block_size=block_size,
block_len=block_len,
slot_size_bytes=slot_size_bytes,
block_stride=stride[1],
local_kv_stride=stride[0],
remote_kv_stride=remote_kv_stride,
transfers_per_block=2,
regions_per_block=1,
split_kv_regions=True,
)
if not is_mla_cache and len(shape) == 5 and shape[1] == 2:
num_blocks, _, block_size, num_kv_heads, head_dim = shape
slot_size_bytes = num_kv_heads * head_dim * element_size
block_len = block_size * slot_size_bytes
return LayerTransferGeometry(
num_blocks=num_blocks,
block_size=block_size,
block_len=block_len,
slot_size_bytes=slot_size_bytes,
block_stride=stride[0],
local_kv_stride=stride[1],
remote_kv_stride=stride[1],
transfers_per_block=2,
regions_per_block=2,
split_kv_regions=False,
)
cache_kind = "MLA" if is_mla_cache else "K/V"
raise ValueError(
f"Unsupported MoRIIO {cache_kind} cache shape for layer "
f"{layer_name}: {tuple(shape)}"
)
def iter_layer_registration_regions(
layer_name: str,
kv_cache: torch.Tensor,
layer_to_spec: Mapping[str, KVCacheSpec],
) -> list[tuple[torch.Tensor, int]]:
geometry = get_layer_transfer_geometry(layer_name, kv_cache, layer_to_spec)
region_len = geometry.num_blocks * geometry.regions_per_block * geometry.block_len
if geometry.split_kv_regions:
return [(cache, region_len) for cache in kv_cache]
return [(kv_cache, region_len)]
def merge_contiguous_offsets(
offsets_local: list[int],
offsets_remote: list[int],
sizes: list[int],
) -> tuple[list[int], list[int], list[int]]:
if not offsets_local:
return [], [], []
if not (len(offsets_local) == len(offsets_remote) == len(sizes)):
raise ValueError("Input list lengths mismatch")
rows = sorted(zip(offsets_local, offsets_remote, sizes), key=lambda row: row[0])
merged: list[list[int]] = []
for local, remote, size in rows:
if (
merged
and local == merged[-1][0] + merged[-1][2]
and remote == merged[-1][1] + merged[-1][2]
):
merged[-1][2] += size
else:
merged.append([local, remote, size])
return (
[row[0] for row in merged],
[row[1] for row in merged],
[row[2] for row in merged],
)
def compute_block_transfer_offsets(
layer_name: str,
kv_cache: torch.Tensor,
layer_to_spec: Mapping[str, KVCacheSpec],
local_block_ids: list[int],
remote_block_ids: list[int],
remote_num_blocks: int,
merge_fn: Callable[
[list[int], list[int], list[int]], tuple[list[int], list[int], list[int]]
] = merge_contiguous_offsets,
) -> tuple[list[int], list[int], list[int]]:
if len(local_block_ids) != len(remote_block_ids):
raise ValueError(
"local_block_ids and remote_block_ids must have the same length: "
f"{len(local_block_ids)} != {len(remote_block_ids)}"
)
geometry = get_layer_transfer_geometry(
layer_name, kv_cache, layer_to_spec, remote_num_blocks
)
element_size = kv_cache.element_size()
transfer_size_byte = geometry.block_len
per_block = geometry.transfers_per_block
total = len(local_block_ids) * per_block
offset_local = [0] * total
offset_remote = [0] * total
sizes = [transfer_size_byte] * total
w = 0
for lb, rb in zip(local_block_ids, remote_block_ids):
offset_local[w] = element_size * (lb * geometry.block_stride)
offset_remote[w] = element_size * (rb * geometry.block_stride)
w += 1
if per_block == 2:
assert geometry.local_kv_stride is not None
assert geometry.remote_kv_stride is not None
offset_local[w] = element_size * (
geometry.local_kv_stride + lb * geometry.block_stride
)
offset_remote[w] = element_size * (
geometry.remote_kv_stride + rb * geometry.block_stride
)
w += 1
return merge_fn(offset_local, offset_remote, sizes)
@@ -0,0 +1,286 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Self-describing KV cache events for the offloading connector.
The OffloadingManager identifies an offloaded chunk only by its OffloadKey,
so its raw events carry no token ids, parent hash, or block size.
:class:`OffloadingEventsTracker` snapshots each chunk's full ``BlockStored``
payload while the ``Request`` is alive and publishes stores as block-granular
payloads: a chunk event may carry multiple constituent per-block hashes, and
evictions fan out to the same hashes. Chunks overlapping a non-chunk-aligned
shared prefix re-announce the shared hashes once per chunk; consumers are
expected to deduplicate (reference-count) repeated store/remove announcements
of the same hash. Opt-in via
``kv_connector_extra_config["self_describing_kv_events"]``; inert unless
KV cache events are enabled. See the PR description for the full design.
"""
from collections.abc import Iterable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, NamedTuple
from vllm.distributed.kv_events import BlockRemoved, BlockStored, KVCacheEvent
from vllm.logger import init_logger
from vllm.v1.core.kv_cache_utils import BlockHash, maybe_convert_block_hash
from vllm.v1.kv_cache_interface import (
KVCacheGroupSpec,
get_kv_cache_spec_kind,
get_kv_cache_spec_sliding_window,
)
from vllm.v1.kv_offload.base import (
OffloadingEvent,
OffloadingKVEventsConfig,
OffloadKey,
get_offload_block_hash,
get_offload_group_idx,
)
from vllm.v1.request import Request
if TYPE_CHECKING:
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.scheduler import (
GroupOffloadConfig,
)
logger = init_logger(__name__)
class OffloadingEventGroupSpec(NamedTuple):
kv_cache_spec_kind: str | None
kv_cache_spec_sliding_window: int | None
def get_offloading_event_group_spec(
kv_cache_group: KVCacheGroupSpec,
) -> OffloadingEventGroupSpec:
kv_cache_spec = kv_cache_group.kv_cache_spec
return OffloadingEventGroupSpec(
kv_cache_spec_kind=get_kv_cache_spec_kind(kv_cache_spec).value,
kv_cache_spec_sliding_window=get_kv_cache_spec_sliding_window(kv_cache_spec),
)
@dataclass(slots=True)
class _OffloadEventMetadata:
"""BlockStored payload snapshot for one OffloadKey, captured at store
time and kept until the matching eviction event. ``medium`` is forwarded
from the OffloadingEvent."""
# The chunk's constituent block hashes; the last one is the OffloadKey.
block_hashes: tuple[BlockHash, ...]
parent_block_hash: BlockHash | None
token_ids: tuple[int, ...]
block_size: int
lora_id: int | None
lora_name: str | None
# Deferred: needs the same incremental curr_mm_idx handling as GPU events.
extra_keys: tuple[tuple[Any, ...] | None, ...] | None
group_idx: int
kv_cache_spec: OffloadingEventGroupSpec
class OffloadingEventsTracker:
"""Tracks offloaded chunks' KV event payloads from store to eviction.
The scheduler calls :meth:`record_store` from ``_build_store_jobs``
while the ``Request`` is available, and routes the manager's raw
:class:`OffloadingEvent` stream through :meth:`take_events`. All state
is bounded by the CPU pool capacity and cleared by :meth:`reset`.
"""
def __init__(self, config: OffloadingKVEventsConfig):
self.config = config
self.self_describing_enabled = (
config.enable_kv_cache_events and config.self_describing_kv_events
)
# OffloadKey -> payload snapshot, kept until the eviction event so
# BlockRemoved can fan out. Bounded: one entry per offloaded chunk.
self._pending_event_metadata: dict[OffloadKey, _OffloadEventMetadata] = {}
def record_store(
self,
req: Request,
group_config: "GroupOffloadConfig",
offload_block_idx: int,
offload_key: OffloadKey,
) -> None:
"""Snapshot the KV cache event payload for one offloaded chunk.
No-op when self-describing event capture is disabled or for
sliding-window / SSM groups, which keep the legacy placeholder payload.
"""
if not self.self_describing_enabled:
return
if group_config.sliding_window_size_in_blocks is not None:
return
meta = self._build_event_metadata(req, group_config, offload_block_idx)
self._pending_event_metadata[offload_key] = meta
def take_events(self, events: Iterable[OffloadingEvent]) -> Iterable[KVCacheEvent]:
"""Translate raw OffloadingEvents into self-describing KV events.
Complete metadata is available only for full-attention groups when
the tracker is enabled. Other shapes retain the legacy placeholder
payload so consumers can ignore them.
Yields:
``BlockStored`` or ``BlockRemoved`` events corresponding to
the underlying :class:`OffloadingEvent` stream.
"""
for event in events:
if event.removed:
yield from self._take_removed_event(event)
else:
yield from self._take_stored_event(event)
def reset(self) -> None:
"""Drop all tracked state; pending payloads are stale after a
manager cache reset."""
self._pending_event_metadata.clear()
def _build_event_metadata(
self,
req: Request,
group_config: "GroupOffloadConfig",
offload_block_idx: int,
) -> _OffloadEventMetadata:
"""Build the payload snapshot for one offloaded chunk: its
constituent per-block hashes, the whole chunk's tokens, and the
per-block ``block_size``."""
hbf = group_config.hash_block_size_factor
assert hbf > 0
assert offload_block_idx >= 0
# per-block token count (= the GPU/hash block size)
sub_block_size = group_config.offloaded_block_size // hbf
# chunk c covers hash-blocks [c*hbf, (c+1)*hbf); its tail block's hash
# is the chunk's OffloadKey.
first_hash_idx = offload_block_idx * hbf
last_hash_idx = first_hash_idx + hbf
assert first_hash_idx >= 0
assert last_hash_idx <= len(req.block_hashes)
chunk_hashes: list[BlockHash] = []
for block_hash in req.block_hashes[first_hash_idx:last_hash_idx]:
assert block_hash is not None
chunk_hashes.append(block_hash)
assert len(chunk_hashes) == hbf
if group_config.sliding_window_size_in_blocks is not None:
# record_store filters these out before calling this helper.
raise AssertionError("self-describing events only support full attention")
parent_block_hash: BlockHash | None
if first_hash_idx == 0:
parent_block_hash = None
else:
parent_block_hash = req.block_hashes[first_hash_idx - 1]
assert parent_block_hash is not None
tok_start = offload_block_idx * group_config.offloaded_block_size
tok_end = tok_start + group_config.offloaded_block_size
assert tok_end <= len(req.all_token_ids)
token_ids = tuple(req.all_token_ids[tok_start:tok_end])
lora_id: int | None = None
lora_name: str | None = None
if req.lora_request is not None:
lora_id = req.lora_request.adapter_id
lora_name = req.lora_request.name
return _OffloadEventMetadata(
block_hashes=tuple(chunk_hashes),
parent_block_hash=parent_block_hash,
token_ids=token_ids,
block_size=sub_block_size,
lora_id=lora_id,
lora_name=lora_name,
extra_keys=None,
group_idx=group_config.group_idx,
kv_cache_spec=group_config.kv_event_group_spec,
)
def _placeholder_stored(self, key: OffloadKey, medium: str) -> BlockStored:
return BlockStored(
block_hashes=[
maybe_convert_block_hash(BlockHash(get_offload_block_hash(key)))
],
parent_block_hash=None,
token_ids=[],
lora_id=None,
block_size=0,
medium=medium,
lora_name=None,
group_idx=get_offload_group_idx(key),
)
def _take_stored_event(self, event: OffloadingEvent) -> Iterable[KVCacheEvent]:
# Metadata is read, NOT popped: the entry must survive until the
# eviction event so BlockRemoved can fan out to the same hashes.
# Events are self-contained (own parent), so key order is free.
for key in event.keys:
meta = self._pending_event_metadata.get(key)
if meta is None:
if self.self_describing_enabled:
# Expected for unsupported shapes; warn once only.
logger.warning_once(
"OffloadingEventsTracker: no event metadata for "
"offload key during BlockStored emission; emitting a "
"placeholder payload. Expected for non-full-attention "
"groups; otherwise indicates a missing populate path."
)
yield self._placeholder_stored(key, event.medium)
continue
yield BlockStored(
block_hashes=list(
maybe_convert_block_hash(h) for h in meta.block_hashes
),
parent_block_hash=(
maybe_convert_block_hash(meta.parent_block_hash)
if meta.parent_block_hash is not None
else None
),
token_ids=list(meta.token_ids),
block_size=meta.block_size,
lora_id=meta.lora_id,
medium=event.medium,
lora_name=meta.lora_name,
extra_keys=(
list(meta.extra_keys) if meta.extra_keys is not None else None
),
group_idx=meta.group_idx,
kv_cache_spec_kind=meta.kv_cache_spec.kv_cache_spec_kind,
kv_cache_spec_sliding_window=(
meta.kv_cache_spec.kv_cache_spec_sliding_window
),
)
def _take_removed_event(self, event: OffloadingEvent) -> Iterable[KVCacheEvent]:
# Keep group_idx unambiguous if a manager batch spans groups.
by_group: dict[int, list] = {}
for key in event.keys:
meta = self._pending_event_metadata.pop(key, None)
if meta is not None:
group_idx = meta.group_idx
by_group.setdefault(group_idx, []).extend(
maybe_convert_block_hash(h) for h in meta.block_hashes
)
else:
if self.self_describing_enabled:
logger.warning_once(
"OffloadingEventsTracker: no event metadata for "
"offload key during BlockRemoved emission; emitting a "
"placeholder removal. Expected if the matching store "
"used the legacy placeholder payload; otherwise "
"indicates missing store metadata."
)
group_idx = get_offload_group_idx(key)
by_group.setdefault(group_idx, []).append(
maybe_convert_block_hash(BlockHash(get_offload_block_hash(key)))
)
for group_idx, hashes in by_group.items():
yield BlockRemoved(
block_hashes=hashes,
medium=event.medium,
group_idx=group_idx,
)
@@ -112,7 +112,7 @@ class _StatsKey:
# Maps metric name -> _MetricType value
TYPES = "types"
# Maps metric name -> observed value (number or list)
# Maps metric name -> {label values tuple -> observed value (number or list)}
DATA = "data"
@@ -125,15 +125,17 @@ class OffloadingConnectorStats(KVConnectorStats):
{
_StatsKey.TYPES: {name: _MetricType.*, ...},
_StatsKey.DATA: {name: value, ...},
_StatsKey.DATA: {name: {labelvalues: value, ...}, ...},
}
This structure is self-describing: it survives IPC serialization
without needing the full ``OffloadingMetricMetadata`` objects on the
receiving side.
Counter values are aggregated by summing, gauge values use the latest
snapshot, and histogram values are lists of observed samples.
Counter values are aggregated by summing per-label-tuple, gauge values
use the latest snapshot per-label-tuple, and histogram values are lists of
observed samples per-label-tuple. Unlabeled metrics use ``()`` as their
labelvalues tuple.
"""
def __post_init__(self):
@@ -160,26 +162,32 @@ class OffloadingConnectorStats(KVConnectorStats):
assert isinstance(other, OffloadingConnectorStats)
other_types = other._types
other_values = other._values
for key, value in other_values.items():
for key, other_label_values in other_values.items():
type_str = other_types.get(key)
if type_str is None:
raise AssertionError(f"Unknown offloading stats key: {key}")
self._types.setdefault(key, type_str)
if type_str == _MetricType.HISTOGRAM:
assert isinstance(value, list)
if key not in self._values:
self._values[key] = value
current_label_values = self._values.setdefault(key, {})
for labelvalues, value in other_label_values.items():
if type_str == _MetricType.HISTOGRAM:
assert isinstance(value, list)
if labelvalues not in current_label_values:
current_label_values[labelvalues] = list(value)
else:
assert isinstance(current_label_values[labelvalues], list)
current_label_values[labelvalues].extend(value)
elif type_str == _MetricType.COUNTER:
assert isinstance(value, int | float)
current_label_values[labelvalues] = (
current_label_values.get(labelvalues, 0) + value
)
elif type_str == _MetricType.GAUGE:
assert isinstance(value, int | float)
current_label_values[labelvalues] = value
else:
assert isinstance(self._values[key], list)
self._values[key].extend(value)
elif type_str == _MetricType.COUNTER:
assert isinstance(value, int | float)
self._values[key] = self._values.get(key, 0) + value
elif type_str == _MetricType.GAUGE:
assert isinstance(value, int | float)
self._values[key] = value
else:
raise AssertionError(f"Unknown metric type '{type_str}' for key: {key}")
raise AssertionError(
f"Unknown metric type '{type_str}' for key: {key}"
)
return self
def reduce(self) -> dict[str, int | float]:
@@ -190,44 +198,62 @@ class OffloadingConnectorStats(KVConnectorStats):
stats for the last time interval.
"""
return_dict: dict[str, int | float] = {}
for key, value in self._values.items():
for key, label_value_map in self._values.items():
type_str = self._types.get(key)
if type_str is None:
raise AssertionError(f"Unknown offloading stats key: {key}")
if type_str == _MetricType.HISTOGRAM:
assert isinstance(value, list)
return_dict[f"{key}_count"] = len(value)
return_dict[f"{key}_sum"] = sum(value)
elif type_str in (_MetricType.COUNTER, _MetricType.GAUGE):
assert isinstance(value, int | float)
return_dict[key] = value
else:
raise AssertionError(f"Unknown metric type '{type_str}' for key: {key}")
for labelvalues, value in label_value_map.items():
key_with_labels = f"{key}:{labelvalues}" if labelvalues else key
if type_str == _MetricType.HISTOGRAM:
assert isinstance(value, list)
return_dict[f"{key_with_labels}_count"] = len(value)
return_dict[f"{key_with_labels}_sum"] = sum(value)
elif type_str in (_MetricType.COUNTER, _MetricType.GAUGE):
assert isinstance(value, int | float)
return_dict[key_with_labels] = value
else:
raise AssertionError(
f"Unknown metric type '{type_str}' for key: {key}"
)
return return_dict
def is_empty(self) -> bool:
return not self.data.get(_StatsKey.DATA)
def increase_counter(
self, counter_name: str, counter_increase_value: int | float
self,
counter_name: str,
counter_increase_value: int | float,
labelvalues: tuple[str, ...] = (),
) -> None:
"""Increase a counter on the stats payload."""
self._types.setdefault(counter_name, _MetricType.COUNTER)
self._values[counter_name] = (
self._values.get(counter_name, 0) + counter_increase_value
counter_values = self._values.setdefault(counter_name, {})
counter_values[labelvalues] = (
counter_values.get(labelvalues, 0) + counter_increase_value
)
def set_gauge(self, gauge_name: str, gauge_value: int | float) -> None:
def set_gauge(
self,
gauge_name: str,
gauge_value: int | float,
labelvalues: tuple[str, ...] = (),
) -> None:
"""Set a gauge snapshot on the stats payload."""
self._types.setdefault(gauge_name, _MetricType.GAUGE)
self._values[gauge_name] = gauge_value
gauge_values = self._values.setdefault(gauge_name, {})
gauge_values[labelvalues] = gauge_value
def observe_histogram(
self, histogram_name: str, histogram_value: int | float
self,
histogram_name: str,
histogram_value: int | float,
labelvalues: tuple[str, ...] = (),
) -> None:
"""Record a histogram observation on the stats payload."""
self._types.setdefault(histogram_name, _MetricType.HISTOGRAM)
self._values.setdefault(histogram_name, []).append(histogram_value)
histogram_values = self._values.setdefault(histogram_name, {})
histogram_values.setdefault(labelvalues, []).append(histogram_value)
class OffloadPromMetrics(KVConnectorPromMetrics):
@@ -255,7 +281,10 @@ class OffloadPromMetrics(KVConnectorPromMetrics):
self._observe_deprecated_metrics = issubclass(spec_cls, CPUOffloadingSpec)
self._offloading_metric_defs: dict[str, PromMetricT] = {}
self.offloading_metrics: dict[tuple[int, str], PromMetricT] = {}
# (engine_idx, metric_name, labelvalues) -> metric with bound labels
self.offloading_metrics: dict[
tuple[int, str, tuple[str, ...]], PromMetricT
] = {}
self._counter_kv_bytes = self._counter_cls(
name=_DEPRECATED_TOTAL_BYTES,
@@ -301,10 +330,6 @@ class OffloadPromMetrics(KVConnectorPromMetrics):
self._offloading_metric_defs[metric_name] = self._create_metric(
metric_name, metadata
)
for engine_idx, labelvalues in per_engine_labelvalues.items():
self.offloading_metrics[(engine_idx, metric_name)] = (
self._offloading_metric_defs[metric_name].labels(*labelvalues)
)
def _create_metric(
self, metric_name: str, metadata: OffloadingMetricMetadata
@@ -312,7 +337,7 @@ class OffloadPromMetrics(KVConnectorPromMetrics):
kwargs: dict[str, Any] = {
"name": metric_name,
"documentation": metadata.documentation,
"labelnames": self._labelnames,
"labelnames": self._labelnames + list(metadata.labelnames),
}
if isinstance(metadata, OffloadingCounterMetadata):
metric_cls = self._counter_cls
@@ -326,11 +351,37 @@ class OffloadPromMetrics(KVConnectorPromMetrics):
raise AssertionError(f"Unknown offloading metric metadata: {metadata}")
return metric_cls(**kwargs)
def _get_prometheus_metric(
self,
metric_name: str,
labelvalues: tuple[str, ...],
engine_idx: int,
) -> PromMetric:
metadata = self._offloading_metric_metadata[metric_name]
if len(labelvalues) != len(metadata.labelnames):
raise AssertionError(
f"Metric {metric_name} expects {len(metadata.labelnames)} labels, "
f"got {len(labelvalues)}"
)
key = (engine_idx, metric_name, labelvalues)
prom_metric = self.offloading_metrics.get(key)
if prom_metric is None:
engine_labelvalues = self.per_engine_labelvalues[engine_idx]
prom_metric = self._offloading_metric_defs[metric_name].labels(
*(engine_labelvalues + list(labelvalues))
)
self.offloading_metrics[key] = prom_metric
return prom_metric
def _increase_counter(
self, metric_name: str, value: int | float, engine_idx: int
self,
metric_name: str,
value: int | float,
labelvalues: tuple[str, ...],
engine_idx: int,
) -> None:
self.offloading_metrics[(engine_idx, metric_name)].inc(value)
if not self._observe_deprecated_metrics:
self._get_prometheus_metric(metric_name, labelvalues, engine_idx).inc(value)
if labelvalues or not self._observe_deprecated_metrics:
return
# Keep deprecated CPU offload transfer metrics updated during the
# transition to flat metric names.
@@ -343,15 +394,26 @@ class OffloadPromMetrics(KVConnectorPromMetrics):
elif metric_name == _TransferMetricName.STORE_TIME:
self.counter_kv_transfer_time[(engine_idx, _TransferType.STORE)].inc(value)
def _set_gauge(self, metric_name: str, value: int | float, engine_idx: int) -> None:
self.offloading_metrics[(engine_idx, metric_name)].set(value)
def _set_gauge(
self,
metric_name: str,
value: int | float,
labelvalues: tuple[str, ...],
engine_idx: int,
) -> None:
self._get_prometheus_metric(metric_name, labelvalues, engine_idx).set(value)
def _observe_histogram(
self, metric_name: str, value: list[int | float], engine_idx: int
self,
metric_name: str,
value: list[int | float],
labelvalues: tuple[str, ...],
engine_idx: int,
) -> None:
prom_metric = self._get_prometheus_metric(metric_name, labelvalues, engine_idx)
for observation in value:
self.offloading_metrics[(engine_idx, metric_name)].observe(observation)
if not self._observe_deprecated_metrics:
prom_metric.observe(observation)
if labelvalues or not self._observe_deprecated_metrics:
continue
# Keep deprecated CPU offload transfer metrics updated during the
# transition to flat metric names.
@@ -368,20 +430,23 @@ class OffloadPromMetrics(KVConnectorPromMetrics):
"""Observe transfer statistics."""
metric_types = transfer_stats_data.get(_StatsKey.TYPES, {})
metric_data = transfer_stats_data.get(_StatsKey.DATA, {})
for key, value in metric_data.items():
for key, label_value_map in metric_data.items():
type_str = metric_types.get(key)
if type_str is None:
raise AssertionError(f"Unknown offloading stats key: {key}")
assert key in self._offloading_metric_defs
if type_str == _MetricType.COUNTER:
assert isinstance(value, int | float)
self._increase_counter(key, value, engine_idx)
elif type_str == _MetricType.GAUGE:
assert isinstance(value, int | float)
self._set_gauge(key, value, engine_idx)
elif type_str == _MetricType.HISTOGRAM:
assert isinstance(value, list)
assert all(isinstance(v, int | float) for v in value)
self._observe_histogram(key, value, engine_idx)
else:
raise AssertionError(f"Unknown metric type '{type_str}' for key: {key}")
for labelvalues, value in label_value_map.items():
if type_str == _MetricType.COUNTER:
assert isinstance(value, int | float)
self._increase_counter(key, value, labelvalues, engine_idx)
elif type_str == _MetricType.GAUGE:
assert isinstance(value, int | float)
self._set_gauge(key, value, labelvalues, engine_idx)
elif type_str == _MetricType.HISTOGRAM:
assert isinstance(value, list)
assert all(isinstance(v, int | float) for v in value)
self._observe_histogram(key, value, labelvalues, engine_idx)
else:
raise AssertionError(
f"Unknown metric type '{type_str}' for key: {key}"
)
@@ -5,7 +5,7 @@ from dataclasses import dataclass, field
from itertools import islice
from typing import Any, NamedTuple
from vllm.distributed.kv_events import BlockRemoved, BlockStored, KVCacheEvent
from vllm.distributed.kv_events import KVCacheEvent
from vllm.distributed.kv_transfer.kv_connector.utils import yield_req_data
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorMetadata
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.common import (
@@ -14,6 +14,11 @@ from vllm.distributed.kv_transfer.kv_connector.v1.offloading.common import (
ReqId,
TransferJob,
)
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.events import (
OffloadingEventGroupSpec,
OffloadingEventsTracker,
get_offloading_event_group_spec,
)
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.metrics import (
OffloadingConnectorStats,
_TransferMetricName,
@@ -36,7 +41,6 @@ from vllm.v1.kv_offload.base import (
OffloadPolicy,
ReqContext,
RequestOffloadingContext,
get_offload_block_hash,
make_offload_key,
)
from vllm.v1.outputs import KVConnectorOutput
@@ -69,6 +73,9 @@ class GroupOffloadConfig(NamedTuple):
gpu_block_size: int
offloaded_block_size: int
hash_block_size_factor: int
# KV cache spec metadata propagated onto emitted BlockStored events so
# KV-aware consumers can classify and filter the group.
kv_event_group_spec: OffloadingEventGroupSpec
# None below means full attention
sliding_window_size_in_blocks: int | None
# Number of this group's offloaded blocks per full-attention alignment
@@ -200,6 +207,9 @@ class SchedulerOffloadConfig(NamedTuple):
alignment_block_count=_alignment_block_count(
gpu_block_size * spec.block_size_factor, sw
),
kv_event_group_spec=get_offloading_event_group_spec(
spec.kv_cache_config.kv_cache_groups[idx]
),
is_eagle_group=idx in eagle_groups,
)
for idx, gpu_block_size in enumerate(spec.gpu_block_size)
@@ -361,6 +371,8 @@ class OffloadingConnectorScheduler:
# be freed before a request finishes).
self._block_id_to_pending_jobs: dict[int, set[int]] = {}
self._events_tracker = OffloadingEventsTracker(spec.kv_events_config)
def _generate_job_id(self) -> int:
job_id = self._job_counter
self._job_counter += 1
@@ -647,6 +659,13 @@ class OffloadingConnectorScheduler:
for group_state in req_status.group_states:
group_state.block_ids.clear()
if req_status.transfer_jobs:
logger.debug(
"Delaying request %s since it still has in-flight transfers",
request.request_id,
)
return None, False
req_status.update_offload_keys()
req_status.num_locally_computed_tokens = num_computed_tokens
@@ -927,6 +946,11 @@ class OffloadingConnectorScheduler:
continue
offloaded_block_idx = start_block_idx + idx
self._events_tracker.record_store(
req, group_config, offloaded_block_idx, offload_key
)
gpu_block_idx = offloaded_block_idx * block_size_factor
for i in range(block_size_factor):
block_id = block_ids[gpu_block_idx + i]
@@ -1177,25 +1201,17 @@ class OffloadingConnectorScheduler:
return False, None
def take_events(self) -> Iterable[KVCacheEvent]:
"""Take the KV cache events from the connector.
"""Drain pending KV cache events.
Returns:
A list of KV cache events.
Complete metadata is available only when self-describing KV events
are enabled, and only for full-attention groups. Other shapes retain
the previous placeholder payload so consumers can ignore them.
Yields:
``BlockStored`` or ``BlockRemoved`` events corresponding to
the underlying :class:`OffloadingEvent` stream.
"""
for event in self.manager.take_events():
block_hashes = [get_offload_block_hash(key) for key in event.keys]
if event.removed:
yield BlockRemoved(block_hashes=block_hashes, medium=event.medium)
else:
yield BlockStored(
block_hashes=block_hashes,
parent_block_hash=None,
token_ids=[],
lora_id=None,
block_size=0,
medium=event.medium,
lora_name=None,
)
yield from self._events_tracker.take_events(self.manager.take_events())
def reset_cache(self) -> None:
"""Reset the offloading manager cache, evicting all stored blocks."""
@@ -1231,6 +1247,10 @@ class OffloadingConnectorScheduler:
self._jobs.clear()
self._block_id_to_pending_jobs.clear()
# The manager pool is empty; pending event payloads and announced
# reference counts are stale.
self._events_tracker.reset()
# Note: _current_batch_jobs_to_flush is intentionally NOT cleared.
# The load flush IDs collected above must be delivered to workers.
if self._blocks_being_loaded is not None:
+2 -2
View File
@@ -61,7 +61,7 @@ def translate_error_response(response: ErrorResponse) -> JSONResponse:
async def create_messages(request: AnthropicMessagesRequest, raw_request: Request):
handler = messages(raw_request)
if handler is None:
base_server = raw_request.app.state.openai_serving_tokenization
base_server = raw_request.app.state.serving_tokenization
error = base_server.create_error_response(
NotImplementedError("The model does not support Messages API")
)
@@ -107,7 +107,7 @@ async def create_messages(request: AnthropicMessagesRequest, raw_request: Reques
async def count_tokens(request: AnthropicCountTokensRequest, raw_request: Request):
handler = messages(raw_request)
if handler is None:
base_server = raw_request.app.state.openai_serving_tokenization
base_server = raw_request.app.state.serving_tokenization
error = base_server.create_error_response(
NotImplementedError("The model does not support Messages API")
)

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