Compare commits

...
Author SHA1 Message Date
yewentao256 b870c8edb4 optimize allpool forward
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-05-04 22:47:43 +00:00
Matthew BonanniandGitHub be5983b874 [Docs] Add non-causal support to attention backend docs (#41643)
Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
2026-05-04 20:35:15 +00:00
fxmarty-amdandGitHub 9c07342fdc [NVFP4][fix] Fix layer.weight -> w13 typo in NVFP4 MOE emulation kernel preparation (#41630)
Signed-off-by: Felix Marty <Felix.Marty@amd.com>
2026-05-04 20:13:37 +00:00
844df54269 feat: update xgrammar==0.2.0 to use structural tags for strict tool calling + reasoning for more models (#40894)
Signed-off-by: Yuchuan <yuchuan.7streams@gmail.com>
Signed-off-by: Michael Goin <mgoin64@gmail.com>
Signed-off-by: mgoin <mgoin64@gmail.com>
Signed-off-by: Ubospica <ubospica@gmail.com>
Signed-off-by: sfeng33 <4florafeng@gmail.com>
Co-authored-by: Michael Goin <mgoin64@gmail.com>
Co-authored-by: Ubospica <ubospica@gmail.com>
Co-authored-by: sfeng33 <4florafeng@gmail.com>
2026-05-04 12:45:24 -07:00
422dd02598 [bugfix] Fix prompt logprobs on request eviction during chunked prefill (#41411)
Signed-off-by: Joachim Studnia <joachim@mistral.ai>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-05-04 11:46:00 -07:00
8c780943b4 Fix Nano Nemotron text-only weight loading (#41205)
Signed-off-by: sunghoon.baek <sunghoon.baek@connectfy.cloud>
Signed-off-by: Baekpica <35071468+Baekpica@users.noreply.github.com>
Signed-off-by: sunghoon.baek <seanbb93@gmail.com>
Co-authored-by: sunghoon.baek <sunghoon.baek@connectfy.cloud>
Co-authored-by: OpenAI Codex <codex@openai.com>
Co-authored-by: Netanel Haber <58652339+netanel-haber@users.noreply.github.com>
2026-05-04 21:43:07 +03:00
e724b0ea8d [ROCm] ROCm7.2.2 + profiler fix + AITER 0.1.12.post2 (#41386)
Signed-off-by: Rohan138 <rohanpotdar138@gmail.com>
Signed-off-by: Gregory Shtrasberg <Gregory.Shtrasberg@amd.com>
Co-authored-by: Rohan138 <rohanpotdar138@gmail.com>
2026-05-04 13:07:19 -05:00
712ad0286c [Bugfix] KimiK2ReasoningParser: guard against buffered end-token in streaming (#41068)
Signed-off-by: Keyi Li <likey6688@gmail.com>
Co-authored-by: Keyi Li <likey6688@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Flora Feng <4florafeng@gmail.com>
2026-05-04 17:42:05 +00:00
Ekagra RanjanandGitHub 321fa2d6d1 Limit gpu utils and lower max BS on test_transcription_api_correctness.py (#41649)
Signed-off-by: Ekagra Ranjan <3116519+ekagra-ranjan@users.noreply.github.com>
2026-05-04 10:30:02 -07:00
Wentao YeandGitHub 3e1ad4435f [Bug] Fix tests/compile/test_config.py AttributeError: 'NoneType' object has no attribute 'dtype' (#41288)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-05-05 00:22:07 +08:00
Netanel HaberandGitHub 8decbfa02c Test nemotron nano-v2 and nemotron nano-v3 separately, disable super-omni redundant tests (#41616)
Signed-off-by: Netanel Haber <58652339+netanel-haber@users.noreply.github.com>
2026-05-04 16:31:37 +03:00
Stefano CastagnettaandGitHub 62ba7516e8 Revert "[Doc] Fix RTD build: pytorch.org/docs/stable/objects.inv returns 404" (#41618)
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
2026-05-04 04:47:42 -07:00
Stefano CastagnettaandGitHub 6f53753fc9 [Bugfix] Apply ruff-format to hyperclovax.py (#41620)
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
2026-05-04 03:37:16 -07:00
Andreas KaratzasandGitHub 6ec9bbec38 [CI] Stabilize cpu offload compressed tensors test (#41102)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-05-04 05:22:42 +00:00
Andreas KaratzasandGitHub 01d4d1ad37 [ROCm][CI] Align spec decode logprob test prefill settings (#41335)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-05-04 04:33:29 +00:00
c103c02a1a [Transformers v5] Vendor HCXVisionConfig for compatibility (#38447)
Signed-off-by: Fang Han <fhan0520@gmail.com>
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-05-04 04:19:52 +00:00
Andreas KaratzasandGitHub 67058ca326 [CI] Clean up remote servers on pytest parent exit (#41570)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-05-04 03:11:22 +00:00
Akim TsvigunandGitHub 894a02500b [Bench] Forward --seed to CustomDataset and CustomMMDataset shuffle (#40788)
Signed-off-by: akimtsvigun <akimtsvigun@gmail.com>
2026-05-04 00:39:10 +00:00
66dfee7121 [Bugfix] Fix degenerate KV cache stride causing TMA cudaErrorIllegalInstruction (#40737)
Signed-off-by: David Oy <david@baseten.co>
Signed-off-by: David Oy <58150256+the-david-oy@users.noreply.github.com>
Signed-off-by: David Oy <david.oy@baseten.co>
Co-authored-by: David Oy <david@baseten.co>
Co-authored-by: Claude <claude@anthropic.com>
Co-authored-by: Vadim Gimpelson <156319763+vadiklyutiy@users.noreply.github.com>
2026-05-03 23:52:18 +00:00
Alex BrooksandGitHub db9a84e0cd [Bugfix] Fix FP8 Bias Loading (#41424)
Signed-off-by: Alex Brooks <albrooks@redhat.com>
2026-05-03 20:30:04 +00:00
tomeras91andGitHub cb03fee32b [Bugfix][Ray] Fix RayExecutorV2 actor name collision with DP > 1 (#40398)
Signed-off-by: Tomer Asida <57313761+tomeras91@users.noreply.github.com>
2026-05-03 13:00:41 -07:00
Wei ZhaoandGitHub c51df43005 Disable flashinfer autotune temporarily due to correctness issues (#41524)
Signed-off-by: wzhao18 <wzhao18.sz@gmail.com>
2026-05-03 16:19:59 +00:00
Taneem IbrahimandGitHub 54dc64d5d3 [Doc] Add Qwen3-30B-A3B-Thinking-2507-FP8 to batch invariance verified models (#41513)
Signed-off-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
2026-05-03 08:47:55 -04:00
Woosuk KwonandGitHub e6ff3e9c83 [MRV2] Add shutdown() method (#41297)
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
2026-05-02 21:06:30 -07:00
Jinzhen LinandGitHub 08834cc3ce [Quantization] add humming mxfp4 moe backend (#41083)
Signed-off-by: Jinzhen Lin <jinzhen.ljz@antgroup.com>
2026-05-02 18:36:03 -07:00
856ec4804a [DSv4] Tune default value of VLLM_MULTI_STREAM_GEMM_TOKEN_THRESHOLD (#41526)
Co-authored-by: Copilot <copilot@github.com>
2026-05-02 18:32:09 -07:00
Yongye ZhuandGitHub 1c607d7b2c [DSV4] Guard megamoe flag with Pure TP (#41522)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
2026-05-02 16:41:40 -07:00
4f7309fcc0 [CI] Add ci-fetch-log.sh helper for Buildkite job logs (#41517)
Signed-off-by: mgoin <mgoin64@gmail.com>
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-05-02 15:23:59 -07:00
Michael GoinandGitHub 0a9362d6ab Revert "[Build] Make bundled DeepGEMM wheel portable across Python versions" (#41512) 2026-05-02 09:42:41 -07:00
cfd2573f23 [Build] Switch CUDA 13.0 wheel builds to PyTorch manylinux_2_28 base (#41416)
Signed-off-by: mgoin <mgoin64@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-05-02 05:51:28 -07:00
Hoang NguyenGitHubClaudemergify[bot] <37929162+mergify[bot]@users.noreply.github.com>Isotr0py
c3ad791e1a [Bugfix][Gemma 4] Clamp soft-token estimate to max_soft_tokens (#40796)
Signed-off-by: Hoang Nguyen <118159510+hnt2601@users.noreply.github.com>
Signed-off-by: Isotr0py <mozf@mail2.sysu.edu.cn>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Co-authored-by: Isotr0py <mozf@mail2.sysu.edu.cn>
2026-05-02 06:34:59 +00:00
Matthew SantiagoandGitHub 8586369f61 Refactor Step3Text loading to use AutoWeightsLoader (#41492)
Signed-off-by: Matthew Santiago <carag.matthew@gmail.com>
2026-05-02 06:22:14 +00:00
ChaunceyandGitHub ae3b4deb8a [Doc] Add Codex usage example (#41358)
Signed-off-by: chaunceyjiang <chaunceyjiang@gmail.com>
2026-05-01 22:27:43 -07:00
Rita BrugarolasandGitHub c293ccc58e [ROCm][Bugfix] Fix init-time bias dtype cast when gate.out_dtype is None (#41405)
Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
2026-05-02 00:13:15 -04:00
Luka GovedičandGitHub d58c42e19c [vLLM IR] 2/N fused_add_rms_norm and maybe_inplace overload (#36823)
Signed-off-by: Luka Govedič <lgovedic@redhat.com>
Signed-off-by: Luka Govedič <ProExpertProg@users.noreply.github.com>
2026-05-01 23:41:15 -04:00
Ekagra RanjanandGitHub 3e49479c4b Limit concurrency on test_transcription_api_correctness.py (#41478)
Signed-off-by: Ekagra Ranjan <3116519+ekagra-ranjan@users.noreply.github.com>
2026-05-02 03:19:07 +00:00
John CalderonandGitHub 964a4bc2a5 [MM][CG] Support ViT CG for Qwen2.5-VL (#40830)
Signed-off-by: John Calderon <jcalderon@nvidia.com>
2026-05-02 11:10:14 +08:00
FredericOdermattandGitHub c408fdd663 [Fix] Sync gemma4 chat template from hf (#39570)
Signed-off-by: Frederic Odermatt <frederic.odermatt@44ai.ch>
2026-05-02 03:06:54 +00:00
Andy LoandGitHub 5737770c6c Re-enable allreduce rms fusion for DP / PP (#41458)
Signed-off-by: Andy Lo <andy@mistral.ai>
2026-05-01 19:01:37 -04:00
Michael GoinandGitHub 0c99629ede [Build] Make bundled DeepGEMM wheel portable across Python versions (#41476)
Signed-off-by: mgoin <mgoin64@gmail.com>
2026-05-01 14:45:03 -07:00
Yongye ZhuandGitHub edd60ac93a [Bugfix] Fix persistent_topk inter-CTA init race on RadixRowState (#41444)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
2026-05-01 14:42:52 -07:00
Yongye ZhuandGitHub bcf5cac9fb [DSV4] Add knob to enable pre-attn gemm (#41443)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
2026-05-01 12:23:17 -07:00
a9484dac7b [Perf] Intergrate Tile Kernels head_compute_mix_kernel for Deepseek-V4 (#41255)
Signed-off-by: Isotr0py <Isotr0py@outlook.com>
Co-authored-by: Roger Wang <hey@rogerw.io>
2026-05-01 15:01:17 -04:00
Matthew BonanniGitHubMichael GoinLucas Wilkinsonmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
f3fef12350 [Attention] Abstract the MLA prefill backends and eliminate cuDNN (#32623)
Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
Signed-off-by: Lucas Wilkinson <lwilkins@redhat.com>
Co-authored-by: Michael Goin <mgoin64@gmail.com>
Co-authored-by: Lucas Wilkinson <lwilkins@redhat.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-05-01 13:36:20 -04:00
51295793a2 [Model Runner V2] Add logprob_token_ids support (#40559)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
Signed-off-by: Nick Hill <nickhill123@gmail.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
2026-05-01 10:02:03 -07:00
Michael GoinandGitHub 3ccc1ff495 [Eval][CI] Add basic mrcr eval to tests/evals/ (#40164)
Signed-off-by: mgoin <mgoin64@gmail.com>
2026-05-01 12:00:38 -04:00
529c671e80 [ROCm][FEAT] AITER Fused Allreduce + RMSNorm (#37646)
Signed-off-by: vllmellm <vllm.ellm@embeddedllm.com>
Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
Signed-off-by: junkang1991 <junkangchow@gmail.com>
Co-authored-by: Rita Brugarolas <Rita.BrugarolasBrufau@amd.com>
Co-authored-by: junkang1991 <junkangchow@gmail.com>
Co-authored-by: Luka Govedič <ProExpertProg@users.noreply.github.com>
Co-authored-by: TJian <tunjian.tan@embeddedllm.com>
2026-05-01 23:07:18 +08:00
bc635fad23 [ROCm][Deepseek] dsv3.2 further optimization (#41217)
Signed-off-by: ganyi <ygan@amd.com>
Signed-off-by: Matthew Wong <Matthew.Wong2@amd.com>
Co-authored-by: Matthew Wong <Matthew.Wong2@amd.com>
2026-05-01 23:06:00 +09:00
Artem PerevedentsevGitHubgemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
c3e64696cd [Perf] Warmup forward_native sampler kernel (#41375)
Signed-off-by: Artem Perevedentsev <aperevedents@nvidia.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-05-01 18:04:11 +04:00
sungsoo haandGitHub 4f7bde572a [Kernel] Pack output and LSE in DCP A2A (#41160) 2026-05-01 09:01:17 -04:00
Or OzeriandGitHub 2fa1f8ec00 [kv_offload+HMA][13/N]: Enable HMA support (#41445)
This is the final PR in a series to enables HMA support for the
offloading connector. The connector advertises `SupportsHMA`
and is validated with unit tests and e2e tests.

Signed-off-by: Or Ozeri <oro@il.ibm.com>
2026-05-01 12:30:03 +01:00
raviguptaamdandGitHub 7075df79b3 [ROCm] Enable DBO (Dynamic Batch Optimization) on ROCm (#34726)
Signed-off-by: raviguptaamd <ravi.gupta@amd.com>
2026-05-01 09:18:30 +00:00
Yuyi AoandGitHub 0dbaf9daad Refractor longcat loading to use AutoWeightsLoader (#41448)
Signed-off-by: George-ao <yuyiao772@gmail.com>
2026-05-01 09:07:23 +00:00
a3ec4a35f5 [Bugfix][Metrics] Fix RayPrometheusMetric.labels() returning shared labeled child (#40840)
When vLLM runs with Ray Prometheus `vllm:request_success{finished_reason=...}`
only ever increments the repetition bucket regardless of the request's actual finish
reason; stop, length, abort, and error stay at zero. Root cause was `labels()` mutated
the wrapped Ray metric's default tags in place and returned self, so every `.labels(...)`
call on a given wrapper returned the same object. 

Co-authored-by: Marwan Sarieddine <sarieddine.marwan@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Signed-off-by: Marwan Sarieddine <sarieddine.marwan@gmail.com>
Signed-off-by: Seiji Eicher <seiji@anyscale.com>
2026-05-01 08:43:39 +01:00
Andreas KaratzasandGitHub 32964e7700 [ROCm][CI] Upgraded UCX and RIXL (#41210)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-05-01 16:40:47 +09:00
a07642667d [Bugfix] Pass reasoning parser kwargs to structured output (#41199)
Signed-off-by: Bugen Zhao <i@bugenzhao.com>
Co-authored-by: wang.yuqi <yuqi.wang@daocloud.io>
2026-04-30 23:38:02 -07:00
baonudesifeizhaiandGitHub c3868bbbe4 [compile] Add FlashInfer FP8 async TP fusion and preserve allreduce fusion ordering #27893 (#39505)
Signed-off-by: baonudesifeizhai <baonudesifeizhai@gmail.com>
Signed-off-by: baonudesifeizhai <85092850+baonudesifeizhai@users.noreply.github.com>
Signed-off-by: roG0d <baonudesifeizhai@gmail.com>
2026-05-01 05:08:34 +00:00
sychen52andGitHub 947138b6c2 Add nvfp4 kv cache support (#40177)
Signed-off-by: Shiyang Chen <shiychen@nvidia.com>
2026-05-01 04:55:16 +00:00
Or OzeriandGitHub 941fb50835 [kv_offload+HMA][12/N]: Scheduler-side support for sliding window groups (#41228)
Signed-off-by: Or Ozeri <oro@il.ibm.com>
2026-05-01 06:59:17 +03:00
Juhi MittalGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
6b6ac6c3c7 [Kernel][MoE] Support GELU on TRT-LLM NvFP4 fused MoE for Gemma4 (#41050)
Signed-off-by: Juhi Mittal <juhim@nvidia.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-05-01 03:37:43 +00:00
Stefano CastagnettaandGitHub b542bdf7fb [Bugfix] Disable FlashInfer CUTLASS MoE on SM110 (Jetson Thor AGX) (#40808)
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
2026-04-30 20:08:49 -07:00
Ronen SchafferandGitHub 415a879899 [KV Offload] Use Collection instead of Sequence/Iterable for OffloadingManager key parameters (#41361)
Signed-off-by: Ronen Schaffer <ronen.schaffer@ibm.com>
2026-05-01 05:18:38 +03:00
Dong WandGitHub 7198940b39 [Model] Add Moondream3 model support(only query and caption skills) (#32325)
Signed-off-by: Dong Wang <dongw2019@gmail.com>
2026-05-01 10:06:48 +08:00
14043dfecd feat: Enable prompt_embeds Content Part Support in vLLM Chat Completions API (#40720)
Signed-off-by: Luis Robaina <luis@protopia.ai>
Signed-off-by: Luis Robaina 🚀 <luisfabian1545@gmail.com>
Signed-off-by: LuisRobaina <luis@protopia.ai>
Co-authored-by: Andrew Sansom <qthequartermasterman@gmail.com>
2026-05-01 10:05:55 +08:00
Andreas KaratzasandGitHub 1adaa5056b [ROCm][CI] Add ROCm score absolute tolerance floor (#41341)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-04-30 18:59:35 -07:00
4d5c89295b (bugfix): block_size check for flex attn (#41363)
Co-authored-by: Matthew Bonanni <mbonanni@redhat.com>
2026-04-30 18:59:26 -07:00
Nick HillandGitHub dd5506a157 [Core] Simplify handling of scheduler_reserve_full_isl option (#41064)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-04-30 18:10:00 -07:00
a3c83ff2fd Faster per-token fp8 group quant packed kernel for blackwell (#41326)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
Co-authored-by: Roger Wang <hey@rogerw.io>
2026-04-30 18:09:55 -07:00
Woosuk KwonandGitHub 9c61864bf8 [DeepSeek] Use torch.mm for bf16xbf16->fp32 gemm (#41300)
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
2026-04-30 16:28:57 -07:00
Tran LeandGitHub 71725f6730 [Bugfix] Fix RoutedExpertsCapturer for Gemma 4 MoE (top_k_experts) (#41401)
Signed-off-by: Tran Le <tranle@fireworks.ai>
2026-04-30 16:19:59 -07:00
b4806c8ee1 [DSV4] Add BF16 and MXFP8 A2A support for flashinfer a2a one sided (#40960)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
Signed-off-by: Zijing Liu <liuzijing2014@gmail.com>
Co-authored-by: Zijing Liu <liuzijing2014@users.noreply.github.com>
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-30 15:33:12 -07:00
Wentao YeandGitHub 526927be94 [Model Runner v2] Fix v2 compile counter num_gpu_runner_capture_triggers and num_cudagraph_captured (#41285)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-04-30 15:20:11 -07:00
Michael GoinandGitHub 75a4c166f2 Fix typo in log message for indexer cache (#41419)
Signed-off-by: Michael Goin <mgoin64@gmail.com>
2026-04-30 15:02:14 -07:00
2917d6363a [NVFP4][Hopper/AMD Instinct] Add Triton kernels for NVFP4 dequantization and QDQ emulation (#40033)
Signed-off-by: Felix Marty <Felix.Marty@amd.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-04-30 17:35:48 -04:00
263 changed files with 18831 additions and 3331 deletions
+3 -3
View File
@@ -37,7 +37,7 @@ steps:
agents:
queue: arm64_cpu_queue_release
commands:
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_AARCH64}\" --build-arg BUILD_BASE_IMAGE=nvidia/cuda:13.0.2-devel-ubuntu22.04 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_AARCH64}\" --build-arg BUILD_OS=manylinux --build-arg BUILD_BASE_IMAGE=pytorch/manylinuxaarch64-builder:cuda13.0 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
- "mkdir artifacts"
- "docker run --rm -v $(pwd)/artifacts:/artifacts_host vllm-ci:build-image bash -c 'cp -r dist /artifacts_host && chmod -R a+rw /artifacts_host'"
- "bash .buildkite/scripts/upload-nightly-wheels.sh"
@@ -76,7 +76,7 @@ steps:
agents:
queue: cpu_queue_release
commands:
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_X86}\" --build-arg BUILD_BASE_IMAGE=nvidia/cuda:13.0.2-devel-ubuntu22.04 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_X86}\" --build-arg BUILD_OS=manylinux --build-arg BUILD_BASE_IMAGE=pytorch/manylinux2_28-builder:cuda13.0 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
- "mkdir artifacts"
- "docker run --rm -v $(pwd)/artifacts:/artifacts_host vllm-ci:build-image bash -c 'cp -r dist /artifacts_host && chmod -R a+rw /artifacts_host'"
- "bash .buildkite/scripts/upload-nightly-wheels.sh"
@@ -723,7 +723,7 @@ steps:
- "bash tools/vllm-rocm/generate-rocm-wheels-root-index.sh"
env:
S3_BUCKET: "vllm-wheels"
VARIANT: "rocm721"
VARIANT: "rocm722"
# ROCm Job 6: Build ROCm Release Docker Image
- label: ":docker: Build release image - x86_64 - ROCm"
+55
View File
@@ -0,0 +1,55 @@
#!/bin/bash
# Usage: ./ci-fetch-log.sh <buildkite_job_url> [output_file]
# ./ci-fetch-log.sh <build_number> <job_uuid> [output_file]
#
# Downloads the raw log for a Buildkite job from the public, unauthenticated
# /organizations/<org>/pipelines/<pipeline>/builds/<n>/jobs/<uuid>/download
# endpoint, then strips ANSI/timestamps via ci-clean-log.sh.
#
# Find <build_number> and <job_uuid> via:
# gh pr checks <PR> --repo vllm-project/vllm
# Each failing row's URL is .../builds/<build_number>#<job_uuid>.
set -euo pipefail
ORG="vllm"
PIPELINE="ci"
usage() {
echo "Usage: $0 <buildkite_job_url> [output_file]"
echo " $0 <build_number> <job_uuid> [output_file]"
exit 1
}
if [ $# -lt 1 ]; then usage; fi
if [[ "$1" == https://* ]]; then
BUILD=$(echo "$1" | sed -nE 's#.*/builds/([0-9]+).*#\1#p')
JOB=$(echo "$1" | grep -oE '[0-9a-f]{8}-[0-9a-f-]+' | head -n 1)
OUT="${2:-ci-${BUILD}-${JOB:0:8}.log}"
else
if [ $# -lt 2 ]; then usage; fi
BUILD="$1"
JOB="$2"
OUT="${3:-ci-${BUILD}-${JOB:0:8}.log}"
fi
if [ -z "$BUILD" ] || [ -z "$JOB" ]; then
echo "Could not parse build number or job UUID from: $1" >&2
usage
fi
COOKIES=$(mktemp)
trap 'rm -f "$COOKIES"' EXIT
# Buildkite issues a session cookie on first hit; subsequent /download needs it.
curl -fsSL -c "$COOKIES" -A "vllm-ci-fetch-log" \
"https://buildkite.com/${ORG}/${PIPELINE}/builds/${BUILD}" -o /dev/null
curl -fsSL -b "$COOKIES" -A "vllm-ci-fetch-log" \
"https://buildkite.com/organizations/${ORG}/pipelines/${PIPELINE}/builds/${BUILD}/jobs/${JOB}/download" \
-o "$OUT"
bash "$(dirname "$0")/ci-clean-log.sh" "$OUT"
echo "$OUT"
+1
View File
@@ -1108,6 +1108,7 @@ steps:
- export VLLM_TEST_CLEAN_GPU_MEMORY=1
- VLLM_TEST_CLEAN_GPU_MEMORY=1 pytest -v -s tests/compile/passes/distributed/test_async_tp.py
- pytest -v -s tests/compile/passes/distributed/test_sequence_parallelism.py
- pytest -v -s tests/compile/passes/distributed/test_tp2_ar_rms.py::test_tp2_ar_rms_fusions
#----------------------------------------------------------- mi300 · cuda ------------------------------------------------------------#
+7
View File
@@ -137,3 +137,10 @@ steps:
commands:
- uv pip install --system 'gpt-oss[eval]==0.0.5'
- pytest -s -v evals/gpt_oss/test_gpqa_correctness.py --config-list-file=configs/models-b200.txt
- label: MRCR Eval Small Models
timeout_in_minutes: 30
source_file_dependencies:
- tests/evals/mrcr/
commands:
- pytest -s -v evals/mrcr/test_mrcr_correctness.py --config-list-file=evals/mrcr/configs/models-small.txt
-1
View File
@@ -1053,7 +1053,6 @@ if(VLLM_GPU_LANG STREQUAL "CUDA")
list(APPEND VLLM_MOE_EXT_SRC
"csrc/moe/moe_wna16.cu"
"csrc/moe/grouped_topk_kernels.cu"
"csrc/moe/router_gemm.cu"
"csrc/moe/topk_softplus_sqrt_kernels.cu")
endif()
@@ -236,17 +236,41 @@ void per_token_group_quant_8bit(const torch::stable::Tensor& input,
#undef LAUNCH_KERNEL
}
template <typename T, typename DST_DTYPE>
__global__ void per_token_group_quant_8bit_packed_kernel(
// Register-resident fast path for group_size==128.
//
// Each thread holds 16 source elements (32 B = uint4 x 2) in registers across
// the absmax reduce -> scale compute -> quantize pipeline. No shared memory.
// UE8M0 scale extracted via bit math (bit-exact with exp2f(ceilf(log2f))).
//
// Loads two contiguous uint4s (16 B + 16 B = 32 B) per thread; on Blackwell
// nvcc fuses these into a single 256-bit LDG.E.256.
//
// Constraints: GROUP_SIZE % (THREADS_PER_GROUP * VEC_SIZE) == 0; for
// THREADS_PER_GROUP=8 and bf16/fp16 (VEC_SIZE=16), this means GROUP_SIZE=128.
template <typename T, typename DST_DTYPE, int GROUP_SIZE>
__global__ void per_token_group_quant_8bit_packed_register_kernel(
const T* __restrict__ input, void* __restrict__ output_q,
unsigned int* __restrict__ output_s_packed, const int group_size,
const int num_groups_padded, const int groups_per_block,
const int padded_groups_per_row, const int groups_per_row, const int mn,
const int tma_aligned_mn, const int num_scale_elems, const float eps,
unsigned int* __restrict__ output_s_packed, const int64_t num_groups_padded,
const int groups_per_block, const int padded_groups_per_row,
const int groups_per_row, const int mn, const int output_q_mn_extent,
const int tma_aligned_mn, const int64_t num_scale_elems, const float eps,
const float min_8bit, const float max_8bit) {
const int threads_per_group = 16;
const int64_t local_group_id = threadIdx.x / threads_per_group;
const int lane_id = threadIdx.x % threads_per_group;
static_assert(GROUP_SIZE == 128, "fast path supports GROUP_SIZE==128");
constexpr int THREADS_PER_GROUP = 8;
constexpr int VEC_SIZE = 32 / sizeof(T); // 16 for bf16/fp16
static_assert(GROUP_SIZE == THREADS_PER_GROUP * VEC_SIZE,
"GROUP_SIZE must equal THREADS_PER_GROUP * VEC_SIZE");
// Each group's 8 threads must live in a single warp octet so the
// 0xffu << (threadIdx.x & 24u) shuffle mask selects exactly the lanes
// that share a group. Requires 32 % THREADS_PER_GROUP == 0 and the host
// to launch num_threads as a multiple of THREADS_PER_GROUP (which it does
// via num_threads = groups_per_block * THREADS_PER_GROUP).
static_assert(32 % THREADS_PER_GROUP == 0,
"THREADS_PER_GROUP must divide warp size for the shuffle "
"mask to be valid");
const int local_group_id = threadIdx.x / THREADS_PER_GROUP;
const int lane_id = threadIdx.x % THREADS_PER_GROUP;
const int64_t block_group_id = blockIdx.x * groups_per_block;
const int64_t global_group_id = block_group_id + local_group_id;
@@ -254,141 +278,207 @@ __global__ void per_token_group_quant_8bit_packed_kernel(
return;
}
// map flat group id to 2D indices (mn_idx, sf_k_idx)
const int sf_k_idx =
static_cast<int>(global_group_id % padded_groups_per_row);
const int mn_idx = static_cast<int>(global_group_id / padded_groups_per_row);
// whether it is a valid group (not padding)
const bool is_valid_group = (mn_idx < mn) && (sf_k_idx < groups_per_row);
// shared memory to cache each group's data to avoid double DRAM reads.
extern __shared__ __align__(16) char smem_raw[];
T* smem = reinterpret_cast<T*>(smem_raw);
T* smem_group = smem + local_group_id * group_size;
// compute scale for valid groups
float y_s = 0.f;
// Load 16 input elements (32 B) into registers as two adjacent uint4
// loads. nvcc keeps these as 2x LDG.E.128 on sm_100; the per-thread cost
// is dominated by HBM bandwidth at large MN, so a fused 256-bit load via
// inline PTX gave no measurable speedup.
// alignas(16) is required so the uint4* reinterpret_cast below is
// well-defined for T == bf16/fp16 (default alignof is 2).
alignas(16) T regs[VEC_SIZE];
float local_absmax = eps;
if (is_valid_group) {
const T* group_input =
input + static_cast<int64_t>(mn_idx) * groups_per_row * group_size +
sf_k_idx * group_size;
y_s = ComputeGroupScale<T, true>(group_input, smem_group, group_size,
lane_id, threads_per_group, eps, max_8bit);
input + static_cast<int64_t>(mn_idx) * groups_per_row * GROUP_SIZE +
sf_k_idx * GROUP_SIZE + lane_id * VEC_SIZE;
uint4* dst = reinterpret_cast<uint4*>(&regs[0]);
const uint4* src = reinterpret_cast<const uint4*>(group_input);
dst[0] = src[0];
dst[1] = src[1];
#pragma unroll
for (int i = 0; i < VEC_SIZE; ++i) {
float v = fabsf(static_cast<float>(regs[i]));
local_absmax = fmaxf(local_absmax, v);
}
}
// pack 4 scales into a uint32 exponent
// 8-lane subgroup shuffle reduce (octet of the warp). The mask selects the
// 8 lanes within the warp that share a group.
unsigned mask = 0xffu << (threadIdx.x & 24u);
local_absmax = fmaxf(local_absmax, __shfl_xor_sync(mask, local_absmax, 4));
local_absmax = fmaxf(local_absmax, __shfl_xor_sync(mask, local_absmax, 2));
local_absmax = fmaxf(local_absmax, __shfl_xor_sync(mask, local_absmax, 1));
float y_s = local_absmax / max_8bit;
y_s = fmaxf(y_s, 1e-10f);
uint32_t bits = __float_as_uint(y_s);
uint32_t exp_bits = (bits >> 23) & 0xffu;
uint32_t mant_bits = bits & 0x7fffffu;
uint8_t exp_byte =
static_cast<uint8_t>(exp_bits + (mant_bits != 0u ? 1u : 0u));
// Lane 0 writes the packed scale byte.
if (lane_id == 0) {
// each uint32 in output_s_packed stores 4 packed scales
const int sf_k_pack_idx = sf_k_idx / 4;
const int pos = sf_k_idx % 4;
const int out_idx = sf_k_pack_idx * tma_aligned_mn + mn_idx;
if (is_valid_group) {
// reinterpret the UE8M0 scale y_s as IEEE bits, extract the 8-bit
// exponent, and place it into the correct byte of the 32-bit word.
const unsigned int bits = __float_as_uint(y_s);
const uint8_t exponent = static_cast<uint8_t>((bits >> 23u) & 0xffu);
reinterpret_cast<uint8_t*>(output_s_packed)[out_idx * 4 + pos] = exponent;
reinterpret_cast<uint8_t*>(output_s_packed)[out_idx * 4 + pos] = exp_byte;
} else if (out_idx < num_scale_elems) {
// write zero for padding groups if within bounds of output_s_packed
reinterpret_cast<uint8_t*>(output_s_packed)[out_idx * 4 + pos] = 0;
}
}
__syncthreads();
if (is_valid_group) {
DST_DTYPE* group_output =
static_cast<DST_DTYPE*>(output_q) +
static_cast<int64_t>(mn_idx) * groups_per_row * group_size +
sf_k_idx * group_size;
QuantizeGroup<T, DST_DTYPE>(smem_group, group_output, group_size, lane_id,
threads_per_group, y_s, min_8bit, max_8bit);
// For padded mn rows that fall within output_q's allocated extent, write
// a uint4 of zeros to keep the buffer clean for downstream TMA loads.
// Skip writes for sf_k padding (those positions don't exist in output_q).
if (!is_valid_group) {
if (sf_k_idx < groups_per_row && mn_idx >= mn &&
mn_idx < output_q_mn_extent) {
DST_DTYPE* group_output =
static_cast<DST_DTYPE*>(output_q) +
static_cast<int64_t>(mn_idx) * groups_per_row * GROUP_SIZE +
sf_k_idx * GROUP_SIZE + lane_id * VEC_SIZE;
*reinterpret_cast<uint4*>(group_output) = make_uint4(0, 0, 0, 0);
}
return;
}
// Reconstruct y_s as a power-of-2 float and use its reciprocal.
float y_s_q = __uint_as_float(static_cast<uint32_t>(exp_byte) << 23);
float inv_y = 1.0f / y_s_q;
// Quantize and pack into 16 fp8/int8 bytes (= uint4). VEC_SIZE==16 so we
// fill four 32-bit words, four bytes each.
uint32_t packed_lo = 0;
uint32_t packed_lo_hi = 0;
uint32_t packed_hi_lo = 0;
uint32_t packed_hi = 0;
#pragma unroll
for (int i = 0; i < VEC_SIZE; ++i) {
float q =
fminf(fmaxf(static_cast<float>(regs[i]) * inv_y, min_8bit), max_8bit);
DST_DTYPE qb = DST_DTYPE(q);
uint8_t byte = *reinterpret_cast<uint8_t*>(&qb);
const int shift = (i & 3) * 8;
if (i < 4) {
packed_lo |= static_cast<uint32_t>(byte) << shift;
} else if (i < 8) {
packed_lo_hi |= static_cast<uint32_t>(byte) << shift;
} else if (i < 12) {
packed_hi_lo |= static_cast<uint32_t>(byte) << shift;
} else {
packed_hi |= static_cast<uint32_t>(byte) << shift;
}
}
uint4 packed_out =
make_uint4(packed_lo, packed_lo_hi, packed_hi_lo, packed_hi);
DST_DTYPE* group_output =
static_cast<DST_DTYPE*>(output_q) +
static_cast<int64_t>(mn_idx) * groups_per_row * GROUP_SIZE +
sf_k_idx * GROUP_SIZE + lane_id * VEC_SIZE;
*reinterpret_cast<uint4*>(group_output) = packed_out;
}
// Public entry point: register-resident packed quant kernel.
// Constraints: group_size == 128 and bf16/fp16 input.
void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
torch::stable::Tensor& output_q,
torch::stable::Tensor& output_s_packed,
int64_t group_size, double eps,
double min_8bit, double max_8bit) {
STD_TORCH_CHECK(group_size == 128,
"per_token_group_quant_8bit_packed only supports "
"group_size==128, got ",
group_size, ".");
const auto in_dtype = input.scalar_type();
STD_TORCH_CHECK(
in_dtype == torch::headeronly::ScalarType::Half ||
in_dtype == torch::headeronly::ScalarType::BFloat16,
"per_token_group_quant_8bit_packed only supports bf16/fp16 input.");
STD_TORCH_CHECK(input.is_contiguous());
STD_TORCH_CHECK(output_q.is_contiguous());
const int64_t k = input.size(-1);
STD_TORCH_CHECK(k % group_size == 0, "Last dimension (", k,
") must be divisible by group_size (", group_size, ").");
STD_TORCH_CHECK(k % group_size == 0, "input last dim k=", k,
" is not divisible by group_size=", group_size, ".");
const int64_t mn = input.numel() / k;
const int64_t groups_per_row = k / group_size;
STD_TORCH_CHECK(output_s_packed.dim() == 2,
"output_s_packed must be 2D, got dim=", output_s_packed.dim(),
".");
const int64_t k_num_packed_sfk = (groups_per_row + 3) / 4;
const int64_t tma_aligned_mn = ((mn + 3) / 4) * 4;
// output_q may be allocated with extra padded mn rows (e.g.,
// (tma_aligned_mn, k)) so the kernel can zero-fill them in-line and the
// caller can use torch.empty instead of torch.zeros. The grid only covers
// up to tma_aligned_mn, so we cap the extent there.
const int64_t output_q_mn_actual = output_q.numel() / k;
STD_TORCH_CHECK(output_q_mn_actual >= mn,
"output_q must have at least mn rows; got ",
output_q_mn_actual, " rows for mn=", mn, ".");
const int64_t output_q_mn_extent =
output_q_mn_actual < tma_aligned_mn ? output_q_mn_actual : tma_aligned_mn;
STD_TORCH_CHECK(
output_s_packed.scalar_type() == torch::headeronly::ScalarType::Int,
"output_s_packed must have dtype int32 for UE8M0-packed scales.");
// DeepGEMM expects SFA scales in MN-major form with shape
// [mn, ceil_div(K, 128 * 4)] and TMA-aligned stride on the last
// dimension.
"output_s_packed must be int32 for UE8M0-packed scales.");
STD_TORCH_CHECK(output_s_packed.size(0) == mn &&
output_s_packed.size(1) == k_num_packed_sfk,
"output_s_packed shape must be [", mn, ", ", k_num_packed_sfk,
"], but got [", output_s_packed.size(0), ", ",
"]; got [", output_s_packed.size(0), ", ",
output_s_packed.size(1), "].");
// Verify column-major TMA-aligned layout
STD_TORCH_CHECK(output_s_packed.stride(0) == 1 &&
output_s_packed.stride(1) == tma_aligned_mn,
"output_s_packed must have strides [1, ", tma_aligned_mn,
"], but got [", output_s_packed.stride(0), ", ",
"output_s_packed strides must be [1, ", tma_aligned_mn,
"]; got [", output_s_packed.stride(0), ", ",
output_s_packed.stride(1), "].");
cudaStream_t stream = get_current_cuda_stream();
constexpr int THREADS_PER_GROUP = 16;
// Expand the grid to cover MN and K padding so every byte in
// output_s_packed is written (padding bytes get zeroed by the kernel).
constexpr int THREADS_PER_GROUP = 8;
const int64_t padded_groups_per_row = k_num_packed_sfk * 4;
const int64_t num_groups_padded = tma_aligned_mn * padded_groups_per_row;
// Number of elements in output_s_packed.
const int64_t num_scale_elems = mn + (k_num_packed_sfk - 1) * tma_aligned_mn;
const int groups_per_block = GetGroupsPerBlock(num_groups_padded);
auto dst_type = output_q.scalar_type();
const int num_blocks = num_groups_padded / groups_per_block;
const int64_t num_blocks = num_groups_padded / groups_per_block;
const int num_threads = groups_per_block * THREADS_PER_GROUP;
// CUDA caps grid.x at 2^31 - 1; this fits any realistic shape but guard
// against pathological inputs.
STD_TORCH_CHECK(num_blocks <= static_cast<int64_t>(INT32_MAX),
"per_token_group_quant_8bit_packed grid too large: ",
num_blocks, " blocks (max ", INT32_MAX, ").");
#define LAUNCH_PACKED_KERNEL(T, DST_DTYPE) \
do { \
dim3 grid(num_blocks); \
dim3 block(num_threads); \
size_t smem_bytes = \
static_cast<size_t>(groups_per_block) * group_size * sizeof(T); \
per_token_group_quant_8bit_packed_kernel<T, DST_DTYPE> \
<<<grid, block, smem_bytes, stream>>>( \
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
static_cast<int>(group_size), static_cast<int>(num_groups_padded), \
groups_per_block, static_cast<int>(padded_groups_per_row), \
static_cast<int>(groups_per_row), static_cast<int>(mn), \
static_cast<int>(tma_aligned_mn), \
static_cast<int>(num_scale_elems), static_cast<float>(eps), \
static_cast<float>(min_8bit), static_cast<float>(max_8bit)); \
#define LAUNCH_REG_KERNEL(T, DST_DTYPE) \
do { \
dim3 grid(static_cast<unsigned int>(num_blocks)); \
dim3 block(num_threads); \
per_token_group_quant_8bit_packed_register_kernel<T, DST_DTYPE, 128> \
<<<grid, block, 0, stream>>>( \
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
num_groups_padded, groups_per_block, \
static_cast<int>(padded_groups_per_row), \
static_cast<int>(groups_per_row), static_cast<int>(mn), \
static_cast<int>(output_q_mn_extent), \
static_cast<int>(tma_aligned_mn), num_scale_elems, \
static_cast<float>(eps), static_cast<float>(min_8bit), \
static_cast<float>(max_8bit)); \
} while (0)
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
input.scalar_type(), "per_token_group_quant_8bit_packed", ([&] {
VLLM_STABLE_DISPATCH_HALF_TYPES(
input.scalar_type(), "per_token_group_quant_8bit_packed_register", ([&] {
if (dst_type == torch::headeronly::ScalarType::Float8_e4m3fn) {
LAUNCH_PACKED_KERNEL(scalar_t, __nv_fp8_e4m3);
LAUNCH_REG_KERNEL(scalar_t, __nv_fp8_e4m3);
} else if (dst_type == torch::headeronly::ScalarType::Char) {
LAUNCH_PACKED_KERNEL(scalar_t, int8_t);
LAUNCH_REG_KERNEL(scalar_t, int8_t);
} else {
STD_TORCH_CHECK(
false,
@@ -397,7 +487,7 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
}
}));
#undef LAUNCH_PACKED_KERNEL
#undef LAUNCH_REG_KERNEL
}
void per_token_group_quant_fp8(const torch::stable::Tensor& input,
@@ -8,3 +8,13 @@ void per_token_group_quant_8bit(const torch::stable::Tensor& input,
torch::stable::Tensor& output_s,
int64_t group_size, double eps, double min_8bit,
double max_8bit, bool scale_ue8m0 = false);
// Public op: register-resident packed quant for the DeepGEMM Blackwell path.
// Restricted to group_size == 128 and bf16/fp16 input; other configurations
// raise STD_TORCH_CHECK. The legacy shared-memory fallback was removed because
// no production caller (deep_gemm_moe / input_quant_fp8) uses other shapes.
void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
torch::stable::Tensor& output_q,
torch::stable::Tensor& output_s_packed,
int64_t group_size, double eps,
double min_8bit, double max_8bit);
-4
View File
@@ -67,10 +67,6 @@ void shuffle_rows(const torch::Tensor& input_tensor,
torch::Tensor& output_tensor);
#ifndef USE_ROCM
// cuBLAS bf16 x bf16 -> fp32 router GEMM (fallback for non-SM90 / batch > 16)
torch::Tensor router_gemm_bf16_fp32(torch::Tensor const& input,
torch::Tensor const& weight);
// DeepSeek V3 optimized router GEMM kernel for SM90+
// Computes output = mat_a @ mat_b.T where:
// mat_a: [num_tokens, hidden_dim] in bf16
-52
View File
@@ -1,52 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
// bf16 x bf16 -> fp32 router GEMM via cuBLAS.
// Uses CUBLAS_COMPUTE_32F so bf16 operands accumulate into fp32,
// matching TRT-LLM's cuBLAS fallback behaviour in dsv3RouterGemmOp.
#include <torch/all.h>
#include <ATen/cuda/CUDAContext.h>
#include <cublas_v2.h>
// cuBLAS column-major math for row-major PyTorch tensors:
// weight[N,K]_row lda=K -> cuBLAS sees (K,N) col-major; CUBLAS_OP_T ->
// (N,K) input[M,K]_row ldb=K -> cuBLAS sees (K,M) col-major; CUBLAS_OP_N
// -> (K,M) out[M,N]_row ldc=N -> cuBLAS sees (N,M) col-major (written as
// output^T)
// cuBLAS: C(N,M) = weight(N,K) @ input(K,M) => C^T = output[M,N]
// params: m=N, n=M, k=K, lda=K (weight), ldb=K (input), ldc=N (output)
torch::Tensor router_gemm_bf16_fp32(torch::Tensor const& input,
torch::Tensor const& weight) {
TORCH_CHECK(input.dtype() == torch::kBFloat16,
"router_gemm_bf16_fp32: input must be bfloat16");
TORCH_CHECK(weight.dtype() == torch::kBFloat16,
"router_gemm_bf16_fp32: weight must be bfloat16");
TORCH_CHECK(input.dim() == 2 && weight.dim() == 2,
"router_gemm_bf16_fp32: input and weight must be 2-D");
TORCH_CHECK(input.size(1) == weight.size(1),
"router_gemm_bf16_fp32: inner dimensions must match");
int64_t const M = input.size(0);
int64_t const N = weight.size(0);
int64_t const K = input.size(1);
auto out = torch::empty({M, N}, input.options().dtype(torch::kFloat32));
cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
TORCH_CUDABLAS_CHECK(
cublasSetStream(handle, at::cuda::getCurrentCUDAStream()));
float const alpha = 1.0f;
float const beta = 0.0f;
TORCH_CUDABLAS_CHECK(cublasGemmEx(
handle, CUBLAS_OP_T, CUBLAS_OP_N, static_cast<int>(N),
static_cast<int>(M), static_cast<int>(K), &alpha, weight.data_ptr(),
CUDA_R_16BF, static_cast<int>(K), input.data_ptr(), CUDA_R_16BF,
static_cast<int>(K), &beta, out.data_ptr(), CUDA_R_32F,
static_cast<int>(N), CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
return out;
}
-4
View File
@@ -133,10 +133,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, m) {
"Tensor)");
m.impl("grouped_topk", torch::kCUDA, &grouped_topk);
// cuBLAS bf16 x bf16 -> fp32 router GEMM (fallback for non-SM90 / batch > 16)
m.def("router_gemm_bf16_fp32(Tensor input, Tensor weight) -> Tensor");
m.impl("router_gemm_bf16_fp32", torch::kCUDA, &router_gemm_bf16_fp32);
// DeepSeek V3 optimized router GEMM for SM90+
m.def("dsv3_router_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()");
// conditionally compiled so impl registration is in source file
+6 -19
View File
@@ -887,27 +887,14 @@ __global__ void __launch_bounds__(kThreadsPerBlock, 2)
uint32_t* shared_ordered =
reinterpret_cast<uint32_t*>(smem_raw + kFixedSmemLarge);
// RadixRowState for multi-CTA cooperative radix
// RadixRowState for multi-CTA cooperative radix.
// Zero-initialization is done host-side via cudaMemsetAsync in topk.cu
// before launch — that gives a stream-ordered happens-before edge for all
// CTAs, which the previous in-kernel init (CTA-0 only + intra-CTA
// __syncthreads) did not provide and which manifested as a race against
// CTA-1+'s first red_release on arrival_counter.
RadixRowState* state = &params.row_states[group_id];
// -- Initialize RadixRowState (only needed if large rows exist) --
if (params.max_seq_len > RADIX_THRESHOLD) {
if (cta_in_group == 0) {
for (uint32_t buf = 0; buf < 3; buf++) {
for (uint32_t i = tx; i < RADIX; i += kThreadsPerBlock) {
state->histogram[buf][i] = 0;
}
}
if (tx == 0) {
state->remaining_k = 0;
state->prefix = 0;
state->arrival_counter = 0;
state->output_counter = 0;
}
}
__syncthreads();
}
int barrier_phase = 0;
const uint32_t total_iters = (params.num_rows + num_groups - 1) / num_groups;
+23
View File
@@ -153,6 +153,29 @@ void launch_persistent_topk(const torch::Tensor& logits,
TORCH_CHECK(workspace.size(0) >= static_cast<int64_t>(state_bytes),
"workspace too small, need ", state_bytes, " bytes");
// Zero the per-group RadixRowState region before launch — only when the
// radix path will actually run (max_seq_len > RADIX_THRESHOLD). The
// RadixRowState fields (arrival_counter, histograms) are only touched by
// radix_topk; the decode/medium paths inside the persistent kernel
// operate purely in shared memory and never read these globals, so a
// stale workspace is harmless for them.
//
// Why we need the memset (when needs_cooperative is true):
// 1. arrival_counter accumulates within a launch and is never reset,
// so a prior call leaves it at a large positive value. Without this
// reset, the very first wait_ge in the next call sees counter >>
// target and returns instantly, breaking the barrier.
// 2. The previous in-kernel init only ran in CTA-0 with intra-CTA
// __syncthreads(), so it had no happens-before edge to CTA-1+'s
// first red_release. cudaMemsetAsync is stream-ordered: the zero
// is globally visible before any CTA runs.
if (needs_cooperative) {
cudaError_t mz_err = cudaMemsetAsync(workspace.data_ptr<uint8_t>(), 0,
state_bytes, stream);
TORCH_CHECK(mz_err == cudaSuccess,
"row_states memset failed: ", cudaGetErrorString(mz_err));
}
P::PersistentTopKParams params;
params.input = logits.data_ptr<float>();
params.output = output.data_ptr<int32_t>();
+67 -26
View File
@@ -41,6 +41,13 @@ ARG BUILD_BASE_IMAGE=nvidia/cuda:${CUDA_VERSION}-devel-ubuntu22.04
# Using cuda base image with minimal dependencies necessary for JIT compilation (FlashInfer, DeepGEMM, EP kernels)
ARG FINAL_BASE_IMAGE=nvidia/cuda:${CUDA_VERSION}-base-ubuntu${UBUNTU_VERSION}
# OS family of BUILD_BASE_IMAGE. Controls package manager (apt vs dnf) and
# Python bootstrap. Set to "manylinux" alongside a manylinux build base such
# as pytorch/manylinux2_28-builder:cuda13.0 to produce wheels with a glibc
# 2.28 floor (matches PyTorch's own published wheels). Default stays on
# Ubuntu for backwards compatibility.
ARG BUILD_OS=ubuntu
# By parameterizing the Deadsnakes repository URL, we allow third-party to use
# their own mirror. When doing so, we don't benefit from the transparent
# installation of the GPG key of the PPA, as done by add-apt-repository, so we
@@ -94,35 +101,64 @@ FROM ${BUILD_BASE_IMAGE} AS base
ARG CUDA_VERSION
ARG PYTHON_VERSION
ARG BUILD_OS
ENV DEBIAN_FRONTEND=noninteractive
# Install system dependencies including build tools
RUN apt-get update -y \
&& apt-get install -y --no-install-recommends \
ccache \
software-properties-common \
git \
curl \
sudo \
python3-pip \
libibverbs-dev \
# Upgrade to GCC 10 to avoid https://gcc.gnu.org/bugzilla/show_bug.cgi?id=92519
# as it was causing spam when compiling the CUTLASS kernels
gcc-10 \
g++-10 \
&& update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-10 110 --slave /usr/bin/g++ g++ /usr/bin/g++-10 \
# Install python dev headers if available (needed for cmake FindPython on Ubuntu 24.04
# which ships cmake 3.28 and requires Development.SABIModule; silently skipped on
# Ubuntu 20.04/22.04 where python3.x-dev is not available without a PPA)
&& (apt-get install -y --no-install-recommends python${PYTHON_VERSION}-dev 2>/dev/null || true) \
&& rm -rf /var/lib/apt/lists/* \
&& curl -LsSf https://astral.sh/uv/install.sh | sh \
&& $HOME/.local/bin/uv venv /opt/venv --python ${PYTHON_VERSION} \
# Install system dependencies including build tools.
# The Ubuntu path uses apt + deadsnakes-via-uv for Python; the manylinux path
# (AlmaLinux 8, e.g. pytorch/manylinux2_28-builder) uses dnf and the Python
# interpreters pre-installed at /opt/python/cpXY-cpXY/.
RUN if [ "${BUILD_OS}" = "manylinux" ]; then \
# rdma-core-devel provides libibverbs headers; ccache lives in EPEL,
# which the pytorch manylinux image already enables. git/curl/sudo
# are typically pre-installed but listed defensively.
dnf install -y --setopt=install_weak_deps=False \
ccache \
git \
curl \
sudo \
rdma-core-devel \
&& dnf clean all \
&& rm -rf /var/cache/dnf; \
else \
apt-get update -y \
&& apt-get install -y --no-install-recommends \
ccache \
software-properties-common \
git \
curl \
sudo \
python3-pip \
libibverbs-dev \
# Upgrade to GCC 10 to avoid https://gcc.gnu.org/bugzilla/show_bug.cgi?id=92519
# as it was causing spam when compiling the CUTLASS kernels
gcc-10 \
g++-10 \
&& update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-10 110 --slave /usr/bin/g++ g++ /usr/bin/g++-10 \
# Install python dev headers if available (needed for cmake FindPython on Ubuntu 24.04
# which ships cmake 3.28 and requires Development.SABIModule; silently skipped on
# Ubuntu 20.04/22.04 where python3.x-dev is not available without a PPA)
&& (apt-get install -y --no-install-recommends python${PYTHON_VERSION}-dev 2>/dev/null || true) \
&& rm -rf /var/lib/apt/lists/*; \
fi
# Install uv and bootstrap /opt/venv. Both paths converge on /opt/venv so all
# downstream stages stay distro-agnostic.
RUN curl -LsSf https://astral.sh/uv/install.sh | sh \
&& if [ "${BUILD_OS}" = "manylinux" ]; then \
# manylinux images ship Python at /opt/python/cpXY-cpXY/; point uv
# at the matching interpreter rather than letting it fetch one.
PYV_NODOT=$(echo ${PYTHON_VERSION} | tr -d '.') \
&& MANYLINUX_PY=/opt/python/cp${PYV_NODOT}-cp${PYV_NODOT}/bin/python${PYTHON_VERSION} \
&& $HOME/.local/bin/uv venv /opt/venv --python "$MANYLINUX_PY"; \
else \
$HOME/.local/bin/uv venv /opt/venv --python ${PYTHON_VERSION}; \
fi \
&& rm -f /usr/bin/python3 /usr/bin/python3-config /usr/bin/pip \
&& ln -s /opt/venv/bin/python3 /usr/bin/python3 \
&& ln -s /opt/venv/bin/python3-config /usr/bin/python3-config \
&& ln -s /opt/venv/bin/pip /usr/bin/pip \
&& ln -sf /opt/venv/bin/python3 /usr/bin/python3 \
&& ln -sf /opt/venv/bin/python3-config /usr/bin/python3-config \
&& ln -sf /opt/venv/bin/pip /usr/bin/pip \
&& python3 --version && python3 -m pip --version
# Activate virtual environment and add uv to PATH
@@ -433,6 +469,7 @@ FROM base AS dev
ARG PIP_INDEX_URL UV_INDEX_URL
ARG PIP_EXTRA_INDEX_URL UV_EXTRA_INDEX_URL
ARG PYTORCH_CUDA_INDEX_BASE_URL
ARG BUILD_OS
# This timeout (in seconds) is necessary when installing some dependencies via uv since it's likely to time out
# Reference: https://github.com/astral-sh/uv/pull/1694
@@ -442,7 +479,11 @@ ENV UV_INDEX_STRATEGY="unsafe-best-match"
ENV UV_LINK_MODE=copy
# Install libnuma-dev, required by fastsafetensors (fixes #20384)
RUN apt-get update && apt-get install -y --no-install-recommends libnuma-dev && rm -rf /var/lib/apt/lists/*
RUN if [ "${BUILD_OS}" = "manylinux" ]; then \
dnf install -y numactl-devel && dnf clean all && rm -rf /var/cache/dnf; \
else \
apt-get update && apt-get install -y --no-install-recommends libnuma-dev && rm -rf /var/lib/apt/lists/*; \
fi
# We can specify the standard or nightly build of PyTorch
+3 -2
View File
@@ -124,9 +124,9 @@ COPY --from=build_vllm ${COMMON_WORKDIR}/vllm/vllm/v1 /vllm_v1
# RIXL/UCX build stages
FROM base AS build_rixl
ARG RIXL_BRANCH="bf4a7214"
ARG RIXL_BRANCH="39be1de8"
ARG RIXL_REPO="https://github.com/ROCm/RIXL.git"
ARG UCX_BRANCH="7009d7a1"
ARG UCX_BRANCH="bfb51733"
ARG UCX_REPO="https://github.com/openucx/ucx.git"
ENV ROCM_PATH=/opt/rocm
ENV UCX_HOME=/usr/local/ucx
@@ -192,6 +192,7 @@ RUN cd /opt/rixl && \
sed -i "s/--exclude 'libamdhip64\*'/--exclude 'libamdhip64*' --exclude 'libcore*' --exclude 'libpull*'/" \
contrib/build-wheel.sh && \
mkdir -p /app/install && \
_ucx_install_dir=${UCX_HOME} \
./contrib/build-wheel.sh \
--output-dir /app/install \
--rocm-dir ${ROCM_PATH} \
+24 -4
View File
@@ -1,4 +1,4 @@
ARG BASE_IMAGE=rocm/dev-ubuntu-22.04:7.2.1-complete
ARG BASE_IMAGE=rocm/dev-ubuntu-22.04:7.2.2-complete
ARG TRITON_BRANCH="ba5c1517"
ARG TRITON_REPO="https://github.com/ROCm/triton.git"
ARG PYTORCH_BRANCH="8514f051" # release/2.10 as of 3/17
@@ -9,7 +9,7 @@ ARG PYTORCH_AUDIO_BRANCH="v2.9.0"
ARG PYTORCH_AUDIO_REPO="https://github.com/pytorch/audio.git"
ARG FA_BRANCH="0e60e394"
ARG FA_REPO="https://github.com/Dao-AILab/flash-attention.git"
ARG AITER_BRANCH="v0.1.10.post3"
ARG AITER_BRANCH="v0.1.12.post2"
ARG AITER_REPO="https://github.com/ROCm/aiter.git"
ARG MORI_BRANCH="v1.1.0"
ARG MORI_REPO="https://github.com/ROCm/mori.git"
@@ -104,6 +104,28 @@ ENV SCCACHE_REGION=${USE_SCCACHE:+${SCCACHE_REGION_NAME}}
ENV SCCACHE_S3_NO_CREDENTIALS=${USE_SCCACHE:+${SCCACHE_S3_NO_CREDENTIALS}}
ENV SCCACHE_IDLE_TIMEOUT=${USE_SCCACHE:+0}
# torch profiler hotfix for 7.2.2: rebuild CLR with https://github.com/ROCm/rocm-systems/pull/5062
# will be removed once we move to ROCm 7.2.3
RUN apt-get update && apt-get install -y rocm-llvm-dev
RUN pip install CppHeaderParser
RUN git clone --no-checkout --filter=blob:none https://github.com/ROCm/rocm-systems /tmp/rocm-systems \
&& cd /tmp/rocm-systems \
&& git sparse-checkout init --cone \
&& git sparse-checkout set projects/hip projects/clr \
&& git checkout 35e8c7bf8911862e5389509800e65fdf125412b3 \
&& export CLR_DIR=/tmp/rocm-systems/projects/clr \
&& export HIP_DIR=/tmp/rocm-systems/projects/hip \
&& mkdir -p $CLR_DIR/build && cd $CLR_DIR/build \
&& cmake \
-DHIP_COMMON_DIR=$HIP_DIR \
-DCMAKE_PREFIX_PATH="/opt/rocm/" \
-DCLR_BUILD_HIP=ON \
-DCLR_BUILD_OCL=OFF \
-DHIP_PLATFORM=amd \
.. \
&& make -j$(nproc) \
&& make install \
&& rm -rf /tmp/rocm-systems
###
### Triton Build
@@ -153,8 +175,6 @@ RUN git clone ${PYTORCH_REPO} pytorch
RUN cd pytorch && git checkout ${PYTORCH_BRANCH}
RUN cd pytorch \
&& pip install -r requirements.txt && git submodule update --init --recursive
RUN cd pytorch/third_party/kineto \
&& git remote add rocm https://github.com/ROCm/kineto && git fetch rocm && git checkout 2d73be3
RUN cd pytorch && python3 tools/amd_build/build_amd.py \
&& if [ "$USE_SCCACHE" = "1" ]; then \
export HIP_CLANG_PATH=/opt/sccache-wrappers \
+3
View File
@@ -16,6 +16,9 @@
"FINAL_BASE_IMAGE": {
"default": "nvidia/cuda:13.0.2-base-ubuntu22.04"
},
"BUILD_OS": {
"default": "ubuntu"
},
"GET_PIP_URL": {
"default": "https://bootstrap.pypa.io/get-pip.py"
},
+12 -2
View File
@@ -60,9 +60,19 @@ the failure?
## Logs Wrangling
Download the full log file from Buildkite locally.
Download a job's log (no Buildkite login required):
Strip timestamps and colorization:
[.buildkite/scripts/ci-fetch-log.sh](../../../.buildkite/scripts/ci-fetch-log.sh)
```bash
# Find the failing job. Each row's URL is .../builds/<N>#<job_uuid>:
gh pr checks <PR> --repo vllm-project/vllm
# Download + strip timestamps/ANSI in one step:
.buildkite/scripts/ci-fetch-log.sh "https://buildkite.com/vllm/ci/builds/<N>#<job_uuid>"
```
To clean an already-downloaded log:
[.buildkite/scripts/ci-clean-log.sh](../../../.buildkite/scripts/ci-clean-log.sh)
+42 -37
View File
@@ -155,6 +155,7 @@ Priority is **1 = highest** (tried first).
| **Block Sizes** | Supported KV cache block sizes (%N means multiples of N) |
| **Head Sizes** | Supported attention head sizes |
| **Sink** | Attention sink support (for StreamingLLM) |
| **Non-Causal** | Non-causal (bidirectional) attention support for decoder models |
| **Sparse** | Sparse attention support (MLA only) |
| **MM Prefix** | Multimodal prefix full attention support |
| **DCP** | Decode Context Parallelism support (`--decode-context-parallel-size`) |
@@ -165,22 +166,22 @@ Priority is **1 = highest** (tried first).
## Standard Attention (MHA, MQA, GQA) Backends
| Backend | Version | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------- | ------ | --------- | ----------- | ---------- | ---- | --------- | --- | --------------- | ------------ |
| `CPU_ATTN` | | fp16, bf16, fp32 | `auto`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | Any | 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 | 64, 128, 256 | ❌ | ❌ | ✅ | Decoder | 7.x-9.x |
| `FLASHINFER` | TRTLLM† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64 | 64, 128, 256 | ✅ | ❌ | ✅ | 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 |
| `FLASH_ATTN_DIFFKV` | | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ✅ | Decoder | Any |
| `FLEX_ATTENTION` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16` | Any | Any | ❌ | ✅ | ❌ | Decoder, Encoder Only | Any |
| `ROCM_AITER_FA` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32 | 64, 128, 256 | ❌ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_UNIFIED_ATTN` | | fp16, bf16 | `auto` | %16 | Any | ✅ | ✅ | ❌ | All | N/A |
| `ROCM_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 128, 160, 192, 224, 256 | ❌ | ✅ | ❌ | Decoder, Encoder, Encoder Only | N/A |
| `TREE_ATTN` | | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | 32, 64, 96, 128, 160, 192, 224, 256 | ❌ | ❌ | ❌ | Decoder | Any |
| `TRITON_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `int8_per_token_head`, `fp8_per_token_head` | %16 | Any | ✅ | ✅ | ❌ | All | Any |
| `TURBOQUANT` | | fp16, bf16 | `turboquant_k8v4`, `turboquant_4bit_nc`, `turboquant_k3v4_nc`, `turboquant_3bit_nc` | 16, 32, 64, 128 | Any | ❌ | ❌ | ❌ | Decoder | Any |
| 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` | Any | 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 | 64, 128, 256 | ❌ | ❌ | ❌ | ✅ | Decoder | 7.x-9.x |
| `FLASHINFER` | TRTLLM† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `nvfp4` | 16, 32, 64 | 64, 128, 256 | ✅ | ❌ | ❌ | ✅ | 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 |
| `FLASH_ATTN_DIFFKV` | | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
| `FLEX_ATTENTION` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder Only | Any |
| `ROCM_AITER_FA` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32 | 64, 128, 256 | ❌ | ✅ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_UNIFIED_ATTN` | | fp16, bf16 | `auto` | %16 | Any | ✅ | ❌ | ✅ | ❌ | All | N/A |
| `ROCM_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 128, 160, 192, 224, 256 | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder, Encoder Only | N/A |
| `TREE_ATTN` | | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | 32, 64, 96, 128, 160, 192, 224, 256 | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
| `TRITON_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `int8_per_token_head`, `fp8_per_token_head` | %16 | Any | ✅ | ❌ | ✅ | ❌ | All | Any |
| `TURBOQUANT` | | fp16, bf16 | `turboquant_k8v4`, `turboquant_4bit_nc`, `turboquant_k3v4_nc`, `turboquant_3bit_nc` | 16, 32, 64, 128 | Any | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
> **†** FlashInfer uses TRTLLM attention on Blackwell (SM100), which supports sinks. Disable via `--attention-config.use_trtllm_attention=0`.
>
@@ -192,31 +193,35 @@ MLA uses separate backends for prefill and decode phases.
### Prefill Backends
The prefill backend is selected at runtime based on hardware and
configuration.
To explicitly select a prefill backend, use
`-ac.mla_prefill_backend=<BACKEND>` (e.g., `FLASH_ATTN`, `FLASHINFER`).
Otherwise, the prefill backend is selected automatically at runtime based on
hardware and configuration.
| Backend | Description | Compute Cap. | Enable | Disable | Notes |
| ------- | ----------- | ------------ | ------ | ------- | ----- |
| TRT-LLM Ragged‡ | TensorRT-LLM ragged attention | 10.x | Default on SM100 | `-ac.use_trtllm_ragged_deepseek_prefill=0` | DeepSeek R1 dims only |
| FlashInfer | FlashInfer CUTLASS backend | 10.x | `-ac.disable_flashinfer_prefill=0` | `-ac.disable_flashinfer_prefill=1` | DeepSeek R1 dims only |
| cuDNN | cuDNN-based attention | 10.x | `-ac.use_cudnn_prefill=1` | `-ac.use_cudnn_prefill=0` | |
| FlashAttention | FlashAttention varlen (FA2/FA3) | Any | Default fallback | Use other backends | FA3 on SM90, FA2 otherwise |
| Backend | Description | Dtypes | Compute Cap. | Notes |
| ------- | ----------- | ------ | ------------ | ----- |
| `FLASH_ATTN`‡ | FlashAttention varlen (FA2/FA3/FA4) | fp16, bf16 | Any | FA4 on SM100+, FA3 on SM90, FA2 otherwise |
| `TRTLLM_RAGGED` | TensorRT-LLM ragged attention | fp16, bf16 | 10.x | DeepSeek R1 dims only |
| `FLASHINFER` | FlashInfer CUTLASS backend | fp16, bf16 | 10.x | DeepSeek R1 dims only |
> **‡** TRT-LLM Ragged is the default on Blackwell (SM100).
> On other GPUs, FlashAttention is used as the default.
### Decode Backends
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------ | --------- | ----------- | ---------- | ---- | ------ | --------- | --- | --------------- | ------------ |
| `CUTLASS_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 128 | Any | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
| `FLASHINFER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | Decoder | 10.x |
| `FLASHINFER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | 576 | ❌ | ✅ | ❌ | ❌ | Decoder | 10.x |
| `FLASHMLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 64 | Any | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x-10.x |
| `FLASHMLA_SPARSE` | bf16 | `auto`, `bfloat16`, `fp8_ds_mla` | 64 | 512, 576 | ❌ | | ❌ | ❌ | Decoder | 9.x-10.x |
| `FLASH_ATTN_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ❌ | ❌ | | Decoder | 9.x |
| `ROCM_AITER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %1 | Any | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | 1 | Any | ❌ | ✅ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_TRITON_MLA` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
| `TRITON_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ❌ | ❌ | ❌ | | Decoder | Any |
| `XPU_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | Any | 576 | ❌ | ✅ | ❌ | ❌ | Decoder | Any |
MLA decode backends are selected using the standard
`-ac.backend=<BACKEND>` argument (e.g., `FLASHMLA`, `TRITON_MLA`).
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ |
| `CUTLASS_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 128 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
| `FLASHINFER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | | ❌ | ❌ | Decoder | 10.x |
| `FLASHINFER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | 576 | ❌ | ❌ | ✅ | ❌ | | Decoder | 10.x |
| `FLASHMLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x-10.x |
| `FLASHMLA_SPARSE` | bf16 | `auto`, `bfloat16`, `fp8_ds_mla` | 64 | 512, 576 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x-10.x |
| `FLASH_ATTN_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x |
| `ROCM_AITER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %1 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 1, 64 | Any | | ❌ | ✅ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_TRITON_MLA` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
| `TRITON_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
| `XPU_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | Any | 576 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | Any |
+2
View File
@@ -86,9 +86,11 @@ Models opt-in to encoder CUDA Graphs by implementing the [SupportsEncoderCudaGra
| Architecture | Models | CG for Image | CG for Video |
| ------------ | ------ | ------------ | ------------ |
| `Qwen3VLForConditionalGeneration` | `Qwen3-VL` | ✅︎ | ✅︎ |
| `Qwen2_5_VLForConditionalGeneration` | `Qwen2.5-VL` | ✅︎ | ✅︎ |
!!! note
Encoder CUDA Graphs have currently been tested with `--mm-encoder-attn-backend=FLASH_ATTN` and `--mm-encoder-attn-backend=FLASHINFER` on Blackwell GPUs.
For Qwen2.5-VL only FA2 and FA3 has been tested.
## Configuration
+25 -7
View File
@@ -5,12 +5,14 @@ TL;DR:
- use tlparse to acquire torch.compile logs. Include these logs in bug reports and/or support asks.
- The vLLM-torch.compile integration is multiple pieces. vLLM exposes flags to turn off each piece:
| Online Flag | Offline Flag | Result |
| ----------- | ------------ | ------ |
| --enforce-eager | enforce_eager=True | Turn off torch.compile and CUDAGraphs |
| -cc.mode=0 | mode=CompilationMode.NONE | Turn off torch.compile only |
| -cc.cudagraph_mode=NONE | compilation_config=CompilationConfig(cudagraph_mode=CUDAGraphMode.NONE) | Turn off CUDAGraphs only |
| -cc.backend=eager | compilation_config=CompilationConfig(backend='eager') | Turn off TorchInductor |
| Online Flag | Offline Flag | Result |
|--------------------------------|--------------------------------------------------------------------------------|------------------------------------------------------|
| --enforce-eager | enforce_eager=True | Turn off torch.compile and CUDAGraphs |
| -cc.mode=0 | compilation_config=CompilationConfig(mode=CompilationMode.NONE) | Turn off torch.compile only |
| -cc.mode=1 | compilation_config=CompilationConfig(mode=CompilationMode.STOCK_TORCH_COMPILE) | Turn off vLLM-compile modifications to torch.compile |
| -cc.cudagraph_mode=NONE | compilation_config=CompilationConfig(cudagraph_mode=CUDAGraphMode.NONE) | Turn off CUDAGraphs only |
| -cc.backend=eager | compilation_config=CompilationConfig(backend='eager') | Turn off TorchInductor |
| -cc.ir_enable_torch_wrap=False | compilation_config=CompilationConfig(ir_enable_torch_wrap=False) | Turn off vLLM IR wrapping |
## vLLM-torch.compile overview
@@ -22,7 +24,7 @@ Most notably, vLLM-compile is NOT torch.compile, it is a custom compiler built u
- Given a model, we do a full graph capture via TorchDynamo that is dynamic on the batch size (number of tokens)
- vLLM then optionally splits and/or specializes this graph and then uses TorchInductor to compile each graph into a compiled artifact.
This step may use vLLM custom Inductor passes to further optimize the graph.
This step may use vLLM custom Inductor passes to further optimize the graph. This includes vLLM IR lowering to remove dispatch overhead.
- The compiled artifact is saved to vLLM's compile cache so that it can be loaded in the future.
- vLLM applies CUDAGraphs to reduce CPU overheads.
@@ -34,6 +36,7 @@ For more details on the design, please see the following resources:
- [Introduction to vLLM-torch.compile blogpost](https://blog.vllm.ai/2025/08/20/torch-compile.html)
- [vLLM-torch.compile integration design](./torch_compile.md)
- [vLLM IR design](./vllm_ir.md)
- [vLLM Office Hours #26](https://www.youtube.com/live/xLyxc7hxCJc?si=Xulo9pe53C6ywf0V&t=561)
- [Talk at PyTorch Conference 2025](https://youtu.be/1wV1ESbGrVQ?si=s1GqymUfwiwOrDTg&t=725)
@@ -117,6 +120,21 @@ from vllm.config.compilation import CompilationConfig, CUDAGraphMode
LLM(model, compilation_config=CompilationConfig(cudagraph_mode=CUDAGraphMode.NONE))
```
vLLM IR makes heavy use of the compilation pipeline, from functionalization, custom fusions, and lowering.
To turn that off and capture eager-mode dispatching behavior of vLLM IR, run with `ir_enable_torch_wrap=False`.
IR torch wrap is only enabled by default when using `mode=VLLM_COMPILE` and `backend="inductor"` (default).
```sh
# Online
vllm serve -cc.ir_enable_torch_wrap=False
```
```py
# Offline
from vllm.config.compilation import CompilationConfig
LLM(model, compilation_config=CompilationConfig(ir_enable_torch_wrap=False))
```
## Debugging TorchDynamo
vLLM requires model code be capturable into a full graph via TorchDynamo (torch.compile's frontend).
+1 -1
View File
@@ -36,7 +36,7 @@ th {
| deepep_high_throughput | standard | fp8 | G(128),A,T<sup>2</sup> | Y | Y | [`DeepEPHTPrepareAndFinalize`][vllm.model_executor.layers.fused_moe.prepare_finalize.deepep_ht.DeepEPHTPrepareAndFinalize] |
| deepep_low_latency | batched | fp8 | G(128),A,T<sup>3</sup> | Y | Y | [`DeepEPLLPrepareAndFinalize`][vllm.model_executor.layers.fused_moe.prepare_finalize.deepep_ll.DeepEPLLPrepareAndFinalize] |
| flashinfer_nvlink_two_sided | standard | nvfp4,fp8 | G,A,T | N | N | [`FlashInferNVLinkTwoSidedPrepareAndFinalize`][vllm.model_executor.layers.fused_moe.prepare_finalize.flashinfer_nvlink_two_sided.FlashInferNVLinkTwoSidedPrepareAndFinalize] |
| flashinfer_nvlink_one_sided | standard | nvfp4 | G,A,T | N | N | [`FlashInferNVLinkOneSidedPrepareAndFinalize`][vllm.model_executor.layers.fused_moe.prepare_finalize.flashinfer_nvlink_one_sided.FlashInferNVLinkOneSidedPrepareAndFinalize] |
| flashinfer_nvlink_one_sided | standard | nvfp4,bf16,mxfp8 | G,A,T | N | N | [`FlashInferNVLinkOneSidedPrepareAndFinalize`][vllm.model_executor.layers.fused_moe.prepare_finalize.flashinfer_nvlink_one_sided.FlashInferNVLinkOneSidedPrepareAndFinalize] |
!!! info "Table key"
1. All types: mxfp4, nvfp4, int4, int8, fp8
+615
View File
@@ -0,0 +1,615 @@
# vLLM IR: Functional Intermediate Representation
## Motivation
vLLM IR is a **functional intermediate representation (IR)** that fills the gap between
low-level `torch` ops and vLLM layers like `RMSNorm` and quantization operators,
By separating operator **semantics** from the **implementation** and **dispatching**,
vLLM IR simplifies both compilation and kernel registration & dispatching simultaneously.
It operates as a **dialect** in the torch FX representation, allowing full interoperability
with “regular” torch ops & custom torch ops/kernels, as well as a piecewise migration from
the previous `CustomOp` approach.
Key design principles:
- **Eager-compile consistency**: identical behavior (barring minor numerics) in eager and compiled modes
- **Simple, transparent, yet powerful kernel selection**: good visibility and control allowing easy debugging
- **Convention over configuration**: near-zero boilerplate required to register ops and implementations
- **Extensibility**: ops and implementations can be registered anywhere, in-tree or out-of-tree
- **Interoperability**: fully compatible with “regular” torch ops & custom torch ops/kernels,
reducing developer friction and allowing piecewise migration
The clean semantics/implementation separation enables a unified and extensible dispatching mechanism,
allowing multiple kernels per-platform and powerful kernel selection. The separation also facilitates
cleaner testing and benchmarking, removing much of the boilerplate standard for legacy approaches.
By delaying kernel selection until late in the compilation process, the compiler can operate on
a higher-level representation, which has the following main benefits:
- Pattern matching in fusion/transformation passes only requires a single, simple pattern per op
- OOT compiler backends can lower from the higher-level representation (in-progress)
- The compiler can autotune over available implementations (future feature)
## Quick Overview
### Declaring an IR Operation
IR operations are declared using the `@register_op` decorator with a native PyTorch implementation that defines the op's semantics:
```python
# vllm/ir/ops/layernorm.py
from torch import Tensor
from vllm.ir import register_op
@register_op
def rms_norm(x: Tensor, weight: Tensor | None, epsilon: float, variance_size: int | None = None) -> Tensor:
"""Weighted root-mean-square layer normalization"""
orig_dtype = x.dtype
x = x.to(torch.float32)
x_var = x if variance_size is None else x[..., :variance_size]
variance = x_var.pow(2).mean(dim=-1, keepdim=True)
x = x * torch.rsqrt(variance + epsilon)
x = x.to(orig_dtype)
if weight is not None:
x = x * weight
return x
```
The native implementation serves three purposes:
1. **Semantic definition**: Specifies the exact semantics of the operation, including shapes and strides
2. **Default implementation**: Used when no other (better) implementation is available
3. **Reference for testing**: Other implementations must match these semantics
### Registering Implementations
Kernel implementations are registered using the `register_impl` decorator on the IR op object:
```python
# vllm/kernels/vllm_c.py
from vllm import ir
rms_norm_no_var = lambda x, weight, epsilon, variance_size=None: variance_size is None
@ir.ops.rms_norm.register_impl("vllm_c", supports_args=rms_norm_no_var, supported=current_platform.is_cuda_alike())
def rms_norm(x: Tensor, weight: Tensor | None, epsilon: float, variance_size: int | None = None) -> Tensor:
output = torch.empty_like(x)
torch.ops._C.rms_norm(output, x, weight, epsilon)
return output
```
Implementations can specify:
- `supported`: Static boolean indicating if this implementation is available
- `supports_args`: Function checking if the implementation supports specific arguments
- `inplace`: Whether this implementation reuses input memory for outputs
### Using IR Operations in Models
IR operations are imported and called directly in model code:
```python
# vllm/model_executor/layers/layernorm.py
from vllm import ir
class RMSNorm(nn.Module):
def __init__(self, hidden_size: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
self.variance_epsilon = eps
def forward(self, x: Tensor, residual: Tensor | None = None):
if residual is None:
return ir.ops.rms_norm(x, self.weight, self.variance_epsilon)
# Use maybe_inplace overload to allow implementation to reuse input memory for outputs
# (using x or residual after this call is undefined behavior)
return ir.ops.fused_add_rms_norm.maybe_inplace(
x, residual, self.weight, self.variance_epsilon
)
```
### Configuring Kernel Selection
Kernel selection is controlled via priority lists in the configuration.
Priority lists specify the order in which implementations are considered,
with the first supported implementation being selected.
This includes the static support check (`supported=...`) and
the dynamic arg support check (`supports_args=...`).
#### Command Line Configuration
Use `--ir-op-priority.<op_name>=<provider1>,<provider2>,...`:
```bash
# CUDA: Use vllm_c implementation for rms_norm
vllm serve meta-llama/Llama-3.2-1B \
--ir-op-priority.rms_norm=vllm_c
# ROCm: Try aiter first, fall back to vllm_c, then native
vllm serve meta-llama/Llama-3.2-1B \
--ir-op-priority.rms_norm=aiter,vllm_c,native
# Configure multiple operations
vllm serve meta-llama/Llama-3.2-1B \
--ir-op-priority.rms_norm=vllm_c \
--ir-op-priority.fused_add_rms_norm=vllm_c
```
#### Python Configuration
```python
from vllm import LLM
from vllm.config import VllmConfig, KernelConfig
llm = LLM(
model="meta-llama/Llama-3.2-1B",
vllm_config=VllmConfig(
kernel_config=KernelConfig(
ir_op_priority={
"rms_norm": ["vllm_c", "native"],
"fused_add_rms_norm": ["vllm_c", "native"],
}
)
)
)
```
#### Platform Defaults
Each platform provides default priority lists that are automatically applied:
```python
# CUDA/XPU/ROCm platform defaults (when compiling with Inductor)
{
"rms_norm": ["native"], # Native torch is default
"fused_add_rms_norm": ["native"],
}
# CUDA platform defaults (eager or Dynamo-only)
{
"rms_norm": ["vllm_c", "native"],
"fused_add_rms_norm": ["vllm_c", "native"],
}
# ROCm platform defaults (future - currently same as CUDA)
{
"rms_norm": ["aiter", "vllm_c", "native"],
"fused_add_rms_norm": ["aiter", "vllm_c", "native"],
}
# XPU platform defaults (eager or Dynamo-only)
{
"rms_norm": ["xpu_kernels", "native"],
"fused_add_rms_norm": ["xpu_kernels", "native"],
}
```
User-specified priorities are prepended to platform defaults,
so you only need to specify the out-of-order implementations,
other implementations are appended automatically.
## Compilation Pipeline
vLLM IR heavily customizes the `torch.compile`-based compilation process to allow custom compile
passes to operate on high-level IR while still producing efficient low-level code at the end.
The compilation pipeline consists of several stages:
### 1. Dynamo Tracing
When `torch.compile` traces the model's forward pass, vLLM IR operations appear as custom operations
in the `vllm_ir` torch library. These operations are opaque to Dynamo, meaning they appear directly
in the FX graph without decomposition:
```python
# Python code (epsilon=1e-5)
x1 = ir.ops.rms_norm(x, weight, epsilon)
x2, residual_out = ir.ops.fused_add_rms_norm.maybe_inplace(x1, residual, weight, epsilon)
# FX graph after Dynamo tracing
x1 = torch.ops.vllm_ir.rms_norm.default(x, weight, 1e-5); x = None
out = torch.ops.vllm_ir.fused_add_rms_norm.maybe_inplace(x1, residual, weight, 1e-5); x1 = residual = None
x2 = out[0]
residual_out = out[1]
```
### 2. AOTAutograd and Functionalization
AOTAutograd functionalizes the graph, converting any mutating operations to functional equivalents.
For vLLM IR operations with `maybe_inplace` overloads, we perform this manually before AOTAutograd,
converting them to the functional `default` overload using the pre-grad custom pass hook.
```python
# After functionalization
x1 = torch.ops.vllm_ir.rms_norm.default(x, weight, 1e-5); x = None
out = torch.ops.vllm_ir.fused_add_rms_norm.default(x1, residual, weight, 1e-5); x1 = residual = None
x2 = out[0]
residual_out = out[1]
```
The pass also tracks which inputs were "donated" (passed to `maybe_inplace`),
storing this information in vLLM's `PassContext` for later use in clone elimination.
### 3. IR Fusion and Transformation Passes
After functionalization, custom vLLM passes operate on the functional FX graph containing high-level IR operations.
These passes can perform fusion, distribute operations for sequence parallelism, and other transformations:
```python
# Example: Sequence Parallelism (see SequenceParallelismPass)
# Before SP pass
all_reduce = torch.ops.vllm.all_reduce(x, "tp:0")
rms_norm = torch.ops.vllm_ir.rms_norm(all_reduce, weight, 1e-5)
# after SP pass
reduce_scatter = torch.ops.vllm.reduce_scatter(x, "tp:0")
rms_norm = torch.ops.vllm_ir.rms_norm(all_reduce, weight, 1e-5)
all_gather = torch.ops.vllm.all_gather(x, "tp:0")
```
Fusion passes benefit from the high-level representation: they don't need to match against low-level PyTorch operations,
handle different kernel implementations separately, or deal with functionalization of custom kernels.
### 4. IR Lowering
The lowering pass (`VllmIRLoweringPass`) replaces each vLLM IR operation with its selected implementation.
The implementation is chosen based on the priority list and support predicates,
using the **fake tensors** in the graph's metadata in place of op arguments:
```python
# Implementation selection, same in eager dispatch and compile lowering
def dispatch(*args) -> IrOpImpl:
for provider in priority_list: # e.g., ["vllm_c", "native"]
impl = ir_op.impls[provider]
if not impl.supported:
continue
if impl.supports_args and not impl.supports_args(*args):
continue
return impl
# make_fx uses torch.fx.symbolic_trace
impl_graph = make_fx(selected_impl.impl_fn)
# Replace IR op node with impl_graph's nodes
match.replace_by_example(selected_impl.impl_fn, node.args)
```
For example, lowering `rms_norm` with the `vllm_c` implementation:
```python
# Before lowering (IR op)
rms_norm = torch.ops.vllm_ir.rms_norm.default(x, weight, 1e-5)
# After lowering (vllm_c implementation traced)
# Note: Lowering does not currently functionalize, this will likely change in the future.
empty = torch.ops.aten.empty.memory_format(x.shape, ...)
rms_norm = torch.ops._C.rms_norm(empty, x, weight, 1e-5)
```
When lowering an implementation that mutates inputs (`inplace=True`),
the lowering pass inserts clones to preserve functional semantics:
```python
# vllm_c implementation for fused_add_rms_norm mutates its first two arguments
# Lowered with clones for safety
clone_default = torch.ops.aten.clone.default(x)
clone_default_1 = torch.ops.aten.clone.default(residual)
fused_add_rms_norm = torch.ops._C.fused_add_rms_norm.default(clone_default, clone_default_1, weight, 1e-5)
```
### 5. Clone Cleanup
After lowering, the clone elimination pass (`UnsafeCloneEliminationPass`) removes unnecessary clones introduced during lowering.
This pass is essential for achieving zero-copy behavior when using in-place kernels with `maybe_inplace`.
The pass removes a clone if:
- the cloned input is created in the graph and not used again in the graph
- the cloned input is a graph parameter, marked as donated
```python
# After cleanup (donated inputs, no subsequent uses)
fused_add_rms_norm = torch.ops._C.fused_add_rms_norm.default(x, residual, weight, 1e-5)
```
The combination of inplace functionalization (tracking donated inputs) and clone cleanup enables the compiler to safely
use in-place kernels without adding redundant copies or increasing the memory usage.
### 6. Inductor Optimization and Codegen
After IR lowering and cleanup, the graph contains only standard PyTorch operations and platform-specific custom ops.
Inductor then performs its standard codegen:
- **Inductor lowering and pointwise fusion**: Fusing element-wise operations, reductions, etc.
- **Memory planning**: Determining buffer allocation and reuse
- **Kernel generation**: Generating Triton or C++ code for fused operations
- **Autotuning**: Selecting the best kernel configurations
### Pipeline Summary
```text
Model Forward Pass
[Dynamo Tracing] → FX Graph with vllm_ir.* ops
[Pre-grad: Inplace Functionalization] → maybe_inplace → default, track donated inputs
[AOTAutograd] → Functionalization
[Post-grad: IR Fusion Passes] → Fuse high-level IR ops (e.g., rms_norm + quant)
[Post-grad: IR Lowering] → vllm_ir.* ops → impl ops (with clones if needed)
[Post-grad: Clone Cleanup] → Remove unnecessary clones using donated input info
[Inductor] → Pattern matching, fusion, memory planning, codegen
Compiled Code
```
## Core vLLM IR Concepts
### Operation Declaration
Operations are declared with the `@register_op` decorator, which creates an `IrOp` object:
```python
@register_op(
name=None, # Operation name (defaults to function name)
activations=None, # List of activation parameters (defaults to params starting with 'x')
allow_inplace=False, # Whether to create a maybe_inplace overload
)
def op_name(...):
...
```
**Parameters:**
- `activations`: List of parameter names considered "activations" (typically consumed by `maybe_inplace`). Defaults to parameters starting with `x`.
- `allow_inplace`: Creates a `maybe_inplace` overload for memory-efficient execution (see below).
### The `maybe_inplace` Overload
The `maybe_inplace` overload is a critical feature for memory efficiency in LLM inference.
It signals that the caller doesn't need to preserve the activation inputs after the operation,
allowing in-place implementations to reuse input memory for outputs.
#### Semantics and Usage
```python
# Standard usage: inputs are preserved
out, res_out = ir.ops.fused_add_rms_norm(x, residual, weight, epsilon)
# x and residual are unchanged, out and res_out are new tensors
# maybe_inplace: inputs may be modified
out, res_out = ir.ops.fused_add_rms_norm.maybe_inplace(x, residual, weight, epsilon)
# x and residual may be modified (undefined behavior to use them after this)
# out and res_out may alias x and residual
```
Using an activation input after passing it to `maybe_inplace` is **undefined behavior**:
```python
# WRONG: Using x after donating it
out, res_out = ir.ops.fused_add_rms_norm.maybe_inplace(x, residual, weight, epsilon)
result = out + x # ERROR: x was donated!
```
If you need to preserve an input, either use the default overload or clone manually:
```python
# Option 1: Use default overload
out, res_out = ir.ops.fused_add_rms_norm(x, residual, weight, epsilon)
result = out + x # OK: x is preserved
# Option 2: Clone before maybe_inplace
out, res_out = ir.ops.fused_add_rms_norm.maybe_inplace(x.clone(), residual, weight, epsilon)
result = out + x # OK: x is preserved, clone was donated
```
#### Compilation Behavior
During compilation, the inplace functionalization pass validates that donated inputs are
not used again and converts `maybe_inplace` to the functional `default` overload:
```python
# Inplace functionalization pass (pre-grad)
for node in graph.nodes:
if node.target == torch.ops.vllm_ir.fused_add_rms_norm.maybe_inplace:
# Check that activation inputs aren't used after this node
for activation_arg in activation_inputs:
for user in activation_arg.users:
if user appears after node:
raise ValueError(f"Input {activation_arg} donated but used again")
# Convert to default overload
node.target = torch.ops.vllm_ir.fused_add_rms_norm.default
# Track donated graph inputs for later clone elimination
for i, arg in enumerate(node.args):
if arg.op == "placeholder" and i in activation_indices:
pass_context.donated_input_ids.add(node_to_idx[arg])
```
The donated input information is then used by the clone cleanup pass to eliminate
unnecessary copies when in-place kernels are lowered.
#### Eager Mode Behavior
In eager mode (without `torch.compile`), `maybe_inplace` enables **maximally memory-efficient**
execution by allowing the IR operation to dispatch directly to in-place implementations:
```python
# Eager dispatch logic for maybe_inplace
impl: IrOpImpl = ir_op.dispatch(*args)
return impl.impl_fn(*args)
# Eager dispatch logic for default:
impl: IrOpImpl = ir_op.dispatch(*args)
if impl.inplace:
args = [
arg.clone() if i in ir_op.activations else arg
for i, arg in enumerate(args)
]
return impl.impl_fn(*args)
```
The combination of `maybe_inplace` in model code and in-place kernel implementations provides optimal memory efficiency
in both eager and compiled modes, with identical semantics in both cases.
#### Memory Savings Example
Consider a transformer layer with residual connections:
```python
# Without maybe_inplace (2 allocations per layer)
hidden_states = self.attention(input)
normed, residual = ir.ops.fused_add_rms_norm(hidden_states, input, weight, eps)
# Memory: input (preserved), hidden_states (preserved), normed (new), residual (new)
# With maybe_inplace (0 allocations per layer when using in-place kernel)
hidden_states = self.attention(input)
normed, residual = ir.ops.fused_add_rms_norm.maybe_inplace(hidden_states, input, weight, eps)
# Memory: normed (reuses hidden_states), residual (reuses input)
```
### Implementation Registration
Implementations are registered using the `register_impl` method:
```python
@ir.ops.op_name.register_impl(
provider="provider_name", # Unique identifier (e.g., "vllm_c", "aiter", "triton")
supported=True, # Static availability check
supports_args=None, # Dynamic argument support check
)
def impl_fn(...):
...
```
**Provider naming conventions:**
- `native`: Reserved for the native torch implementation (declared with `@register_op`)
- `vllm_c`: C++/CUDA kernels via `torch.ops._C`
- `aiter`: AMD AITER library
- `xpu_kernels`: SYCL/SYCLTLA kernels implemented in `vllm-xpu-kernels`
- `triton_*`: Triton kernels
- Platform/library names for other implementations
**Support checking:**
- `supported`: Static boolean, checked once at import time (e.g., `HAS_TRITON`, `is_cuda_alike()`)
- `supports_args`: Function `(*args, **kwargs) -> bool` checking argument compatibility
- Called with **fake tensors** during compilation for zero-cost checking
- Called with **real tensors** during eager mode dispatch
- Should NOT check batch sizes or add guards based on values
Example support predicate:
```python
def aiter_rms_norm_supports(x, weight, epsilon, variance_size=None):
# Check dtype (OK: doesn't depend on batch size)
if x.dtype not in [torch.float16, torch.bfloat16]:
return False
# Check optional parameter (OK: static check)
if variance_size is not None:
return False
return True
@ir.ops.rms_norm.register_impl("aiter", supports_args=aiter_rms_norm_supports)
def rms_norm(...):
...
```
Batch-invariant kernels are automatically selected when `VLLM_BATCH_INVARIANT=1` is set.
### Eager Mode vs Compile Mode
vLLM IR operations behave identically in eager and compile modes:
**Eager mode:**
- Direct dispatch to implementation based on priority list
- Support checked with real tensor arguments
- Minimal overhead (can be optimized further if needed)
**Compile mode:**
- IR ops appear in FX graph as `torch.ops.vllm_ir.*` custom ops
- Lowering selects implementation using fake tensors
- Full integration with Inductor optimizations
This consistency enables:
- Prototyping in eager mode with confidence
- Debugging by disabling compilation
- Gradual migration from eager to compiled execution
## Other Topics
### Out-of-Tree Implementations
External platforms can register implementations without modifying vLLM:
```python
# In external package
from vllm import ir
@ir.ops.rms_norm.register_impl("my_platform", supported=is_my_platform())
def rms_norm(x, weight, epsilon, variance_size=None):
return my_platform.rms_norm(x, weight, epsilon)
```
Then configure priority to use your implementation:
```python
class MyPlatform(Platform):
def get_default_ir_op_priority(self):
return IrOpPriorityConfig(rms_norm=['my_platform', 'native'])
# Users can still override priority in the same way
llm = LLM(ir_op_priority=IrOpPriorityConfig(rms_norm=['custom_oot_kernel']))
```
### Debugging and Observability
!!! note
Please let us know how observability can be improved for your use-case!
Enable debug logging to see kernel selection:
```bash
VLLM_LOGGING_LEVEL=DEBUG vllm serve ...
```
This logs:
- Which implementations are selected for each operation
- Why implementations were rejected (unsupported, args not supported)
- Compilation cache hits/misses
- IR lowering statistics
Check selected implementations in compiled graphs:
```python
# After compilation, inspect the lowering pass
lowering_pass = backend.lowering_pass
print(lowering_pass.selected_impls)
# Output: {'rms_norm': {'node_123': 'vllm_c', 'node_456': 'vllm_c'}}
```
## Migration from CustomOp
vLLM IR is designed to coexist with and gradually replace `CustomOp`:
1. **Op declaration**: Convert `CustomOp` class `PluggableLayer` and move `forward_native` to `@register_op` function
2. **Implementation registration**: Use `@ir.ops.op_name.register_impl` instead of overriding methods
3. **Layer usage**: Replace `self.op(...)` with `ir.ops.op_name(...)`
4. **Configuration**: Migrate `--compilation-config.custom-ops` to `--ir-op-priority`
The migration can be done incrementally, one operation at a time.
## See Also
- [torch.compile Integration](torch_compile.md) - General compilation infrastructure
- [Fusions](fusions.md) - Custom fusion and transformation passes in vLLM
- [Custom Operations](custom_op.md) - Legacy custom op system
+2 -2
View File
@@ -52,10 +52,10 @@ th:not(:first-child) {
| [mm](multimodal_inputs.md) | ✅ | ✅ | [🟠](https://github.com/vllm-project/vllm/pull/4194)<sup>^</sup> | ❔ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❔ | ✅ | | | |
| best-of | ✅ | ✅ | ✅ | [](https://github.com/vllm-project/vllm/issues/6137) | ✅ | ❌ | ✅ | ✅ | ✅ | ❔ | [](https://github.com/vllm-project/vllm/issues/7968) | ✅ | ✅ | | |
| beam-search | ✅ | ✅ | ✅ | [](https://github.com/vllm-project/vllm/issues/6137) | ✅ | ❌ | ✅ | ✅ | ✅ | ❔ | [](https://github.com/vllm-project/vllm/issues/7968) | ❔ | ✅ | ✅ | |
| [prompt-embeds](prompt_embeds.md) | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❔ | ❔ | | ❔ | ❔ | ✅ |
| [prompt-embeds](prompt_embeds.md) | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❔ | ❔ | | ❔ | ❔ | ✅ |
\* Chunked prefill and prefix caching are only applicable to last-token or all pooling with causal attention.
<sup>^</sup> LoRA is only applicable to the language backbone of multimodal models.
<sup>^</sup> LoRA is only applicable to the language backbone of multimodal models.
### Feature x Hardware
+1 -1
View File
@@ -105,7 +105,7 @@ Batch invariance has been tested and verified on the following models:
- **DeepSeek series**: `deepseek-ai/DeepSeek-V3`, `deepseek-ai/DeepSeek-V3-0324`, `deepseek-ai/DeepSeek-R1`, `deepseek-ai/DeepSeek-V3.1`
- **Qwen3 (Dense)**: `Qwen/Qwen3-1.7B`, `Qwen/Qwen3-8B`, `Qwen/Qwen3-4B-AWQ`, `Qwen/Qwen3-8B-AWQ`
- **Qwen3 (MoE)**: `Qwen/Qwen3-30B-A3B`, `Qwen/Qwen3-Next-80B-A3B-Instruct`
- **Qwen3 (MoE)**: `Qwen/Qwen3-30B-A3B`, `Qwen/Qwen3-Next-80B-A3B-Instruct`, `Qwen/Qwen3-30B-A3B-Thinking-2507-FP8`
- **Qwen2.5**: `Qwen/Qwen2.5-0.5B-Instruct`, `Qwen/Qwen2.5-1.5B-Instruct`, `Qwen/Qwen2.5-3B-Instruct`, `Qwen/Qwen2.5-7B-Instruct`, `Qwen/Qwen2.5-14B-Instruct`, `Qwen/Qwen2.5-32B-Instruct`
- **Llama 3**: `meta-llama/Llama-3.1-8B-Instruct`, `meta-llama/Llama-3.2-1B-Instruct`
- **GPT-OSS**: `openai/gpt-oss-20b`, `openai/gpt-oss-120b`
+61
View File
@@ -215,6 +215,67 @@ When loading RGBA images (images with transparency), vLLM converts them to RGB f
- This setting only affects RGBA images with transparency; RGB images are unchanged
- If not specified, the default white background `(255, 255, 255)` is used for backward compatibility
#### Moondream3 Prompt Recipes { #moondream3-prompt-recipes }
`Moondream3ForCausalLM` supports two task-specific prompt formats:
- `query`: ask a question about the image.
- `caption`: generate a caption for the image.
```python
from vllm import LLM, SamplingParams
from vllm.assets.image import ImageAsset
llm = LLM(
model="moondream/moondream3-preview",
tokenizer="moondream/starmie-v1",
trust_remote_code=True,
max_model_len=2048,
limit_mm_per_prompt={"image": 1},
)
image = ImageAsset("stop_sign").pil_image
def make_query_prompt(question: str) -> str:
return (
"<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>"
f"{question}<|md_reserved_2|>"
)
def make_caption_prompt(length: str = "normal") -> str:
return (
"<|endoftext|><image><|md_reserved_0|>"
f"describe<|md_reserved_1|>{length}<|md_reserved_2|>"
)
query_out = llm.generate(
{
"prompt": make_query_prompt("What is shown in this image?"),
"multi_modal_data": {"image": image},
},
SamplingParams(max_tokens=64, temperature=0),
)[0].outputs[0].text
caption_out = llm.generate(
{
"prompt": make_caption_prompt(),
"multi_modal_data": {"image": image},
},
SamplingParams(max_tokens=100, temperature=0),
)[0].outputs[0].text
print("query:", query_out)
print("caption:", caption_out)
```
!!! note
The native Moondream3 model also has `detect` and `point` skills. Those
require custom coordinate decoding and are not exposed by this vLLM
implementation.
### Video Inputs
You can pass a list of NumPy arrays directly to the `'video'` field of the multi-modal dictionary
+36 -1
View File
@@ -20,12 +20,47 @@ You can pass prompt embeddings from Hugging Face Transformers models to the `'p
## Online Serving
Our OpenAI-compatible server accepts prompt embeddings inputs via the [Completions API](https://platform.openai.com/docs/api-reference/completions). Prompt embeddings inputs are added via a new `'prompt_embeds'` key in the JSON package and are enabled by the `--enable-prompt-embeds` flag in `vllm serve`.
Our OpenAI-compatible server accepts prompt embeddings inputs via both the [Completions API](https://platform.openai.com/docs/api-reference/completions) and the [Chat Completions API](https://platform.openai.com/docs/api-reference/chat). Both are enabled by the `--enable-prompt-embeds` flag in `vllm serve`.
### Completions API
Prompt embeddings inputs are added via a `'prompt_embeds'` key in the JSON request body.
When a mixture of `'prompt_embeds'` and `'prompt'` inputs are provided in a single request, the prompt embeds are always returned first.
Prompt embeddings are passed in as base64 encoded torch tensors.
The Completions endpoint does **not** apply a chat template to `prompt_embeds`. If the model assumes some chat template, the caller is responsible for producing embeddings for the full, already-templated prompt: apply the chat template, then embed the resulting token IDs. Anything the model would normally need (system prompt, role markers, generation prompt, etc.) must already be baked into the embedded tokens.
### Chat Completions API
Prompt embeddings can be included as content parts in chat messages, interleaved with text:
```json
{
"messages": [
{
"role": "system",
"content": [
{"type": "text", "text": "You are a helpful assistant."},
{"type": "prompt_embeds", "data": "<base64_encoded_tensor>"}
]
},
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": "<base64_encoded_tensor>"},
{"type": "text", "text": "Summarize the above."}
]
}
]
}
```
Each `prompt_embeds` content part contains a `data` field with a base64-encoded `torch.Tensor` of shape `(num_tokens, hidden_size)`. Multiple `prompt_embeds` parts can appear in any message, in any position relative to text parts. The server expands each part into the correct number of placeholder tokens during chat template rendering, then splices the pre-computed embeddings into the model's input at the corresponding positions.
Unlike the Completions API, a `prompt_embeds` content part should encode **only** the content, not a templated conversation. The server wraps the chat template around the embedded content at request time, the same way it would for a plain text `content` string. Embedding a full templated conversation here would double-apply the template and produce incorrect inputs to the model.
!!! warning
The vLLM engine may crash if incorrect shape of embeddings is passed.
Only enable this flag for trusted users!
+27 -1
View File
@@ -7,7 +7,7 @@ import sys
import textwrap
import traceback
from argparse import SUPPRESS, Action, HelpFormatter
from collections.abc import Iterable
from collections.abc import Callable, Iterable
from importlib.machinery import ModuleSpec
from pathlib import Path
from typing import TYPE_CHECKING, Literal
@@ -48,6 +48,7 @@ class MockPluggableLayer:
mock_if_no_torch("vllm._C", MagicMock())
mock_if_no_torch("vllm._C_stable_libtorch", MagicMock())
mock_if_no_torch(
"vllm.model_executor.custom_op",
MagicMock(CustomOp=MockCustomOp, PluggableLayer=MockPluggableLayer),
@@ -67,6 +68,31 @@ importlib.metadata.version = lambda name: VERSIONS.get(name) or "0.0.0"
mock_if_no_torch("torch.nn", MagicMock(Parameter=object))
# Mock torch.library.infer_schema for vllm.ir.ops.IrOpInplaceOverload.__init__
# We need to return the corresponding number of inputs, as IR infra will assert it
def get_outputs(native_fn: Callable) -> str:
"""
Extract output schema from function's return type annotation,
e.g. 'Tensor' or 'Tensor, Tensor'.
"""
import typing
return_type = typing.get_type_hints(native_fn)["return"]
origin = typing.get_origin(return_type)
arg_name = lambda a: a.__name__ if hasattr(a, "__name__") else str(a)
if origin is tuple:
args = typing.get_args(return_type)
return ", ".join(arg_name(arg) for arg in args)
else:
return f"{arg_name(return_type)}"
mock_if_no_torch(
"torch.library",
MagicMock(infer_schema=lambda fn, **k: f"(Tensor x) -> {get_outputs(fn)}"),
)
class PydanticMagicMock(MagicMock):
"""`MagicMock` that's able to generate pydantic-core schemas."""
+7
View File
@@ -599,6 +599,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
| `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. | ✅︎ | ✅︎ |
| `Molmo2ForConditionalGeneration` | Molmo2 | T + I<sup>+</sup> / V | `allenai/Molmo2-4B`, `allenai/Molmo2-8B`, `allenai/Molmo2-O-7B` | ✅︎ | ✅︎ |
| `Moondream3ForCausalLM` | Moondream3 | T + I | `moondream/moondream3-preview` | | ✅︎ |
| `MusicFlamingoForConditionalGeneration` | MusicFlamingo | T + A | `nvidia/music-flamingo-2601-hf`, `nvidia/music-flamingo-think-2601-hf` | ✅︎ | ✅︎ |
| `NVLM_D_Model` | NVLM-D 1.0 | T + I<sup>+</sup> | `nvidia/NVLM-D-72B`, etc. | | ✅︎ |
| `OpenCUAForConditionalGeneration` | OpenCUA-7B | T + I<sup>E+</sup> | `xlangai/OpenCUA-7B` | ✅︎ | ✅︎ |
@@ -661,6 +662,12 @@ Some models are supported only via the [Transformers modeling backend](#transfor
!!! note
For `InternVLChatModel`, only InternVL2.5 with Qwen2.5 text backbone (`OpenGVLab/InternVL2.5-1B` etc.), InternVL3 and InternVL3.5 have video inputs support currently.
!!! note
`Moondream3ForCausalLM` uses task-specific prompt templates for `query`
and `caption`. The native `detect` and `point` skills require custom
coordinate decoding and are not exposed by this vLLM implementation.
See [Moondream3 prompt recipes](../features/multimodal_inputs.md#moondream3-prompt-recipes).
!!! note
To use `TIGER-Lab/Mantis-8B-siglip-llama3`, you have to pass `--hf_overrides '{"architectures": ["MantisForConditionalGeneration"]}'` when running vLLM.
+88
View File
@@ -0,0 +1,88 @@
# Codex
[Codex](https://github.com/openai/codex) is OpenAI's official agentic coding tool that lives in your terminal. It can understand your codebase, edit files, run commands, and help you write code more efficiently.
By pointing Codex at a vLLM server, you can use your own models as the backend instead of the OpenAI API. This is useful for:
- Running fully local/private coding assistance
- Using open-weight models with tool calling capabilities
- Testing and developing with custom models
## How It Works
vLLM implements the OpenAI-Responses API, which is the same API that Codex uses to communicate with OpenAI's servers. By configuring Codex to point at your vLLM server, Codex sends its requests to vLLM instead of OpenAI. vLLM then translates these requests to work with your local model and returns responses in the format Codex expects.
This means any model served by vLLM with proper tool calling support can act as a drop-in replacement for OpenAI models in Codex.
## Requirements
Codex requires a model with strong tool calling capabilities. The model must support the OpenAI-Responses tool calling API. See [Tool Calling](../../features/tool_calling.md) for details on enabling tool calling for your model.
## Installation
First, install Codex by following the [official installation guide](https://github.com/openai/codex).
## Starting the vLLM Server
Start vLLM with a tool-calling capable model - here's an example using `Qwen/Qwen3-27B`:
```bash
vllm serve Qwen/Qwen3.6-27B --port 8000 --tensor-parallel-size 8 --max-model-len 262144 --reasoning-parser qwen3 --enable-auto-tool-choice --tool-call-parser qwen3_coder
```
For other models, you'll need to enable tool calling explicitly with `--enable-auto-tool-choice` and the right `--tool-call-parser`. Refer to the [Tool Calling documentation](../../features/tool_calling.md) for the correct flags for your model.
## Configuring Codex
Codex is configured via a TOML file located at `~/.codex/config.toml`. Create or edit this file to point Codex at your vLLM server:
```toml
model = "my-model"
model_provider = "vllm"
[model_providers.vllm]
name = "vLLM"
env_key = "VLLM_API_KEY"
base_url = "http://localhost:8000/v1"
wire_api = "responses"
```
The configuration fields:
| Field | Description |
| ----- | ----------- |
| `model` | The model name to use. Must match the `--served-model-name` you passed to vLLM. |
| `model_provider` | Set to `"vllm"` to use your local vLLM server. |
| `[model_providers.vllm]` | Configuration section for the vLLM provider. |
| `name` | A display name for your vLLM provider. |
| `env_key` | The name of an environment variable that Codex will read for the API key. vLLM does not require authentication by default, so this can be any value. |
| `base_url` | The URL of your vLLM server's OpenAI-compatible API endpoint (default is `http://localhost:8000/v1`). |
| `wire_api` | The API style to use. Set to `"responses"` for the OpenAI Responses API |
!!! tip
You can set the `env_key` to any dummy environment variable since vLLM doesn't require authentication by default:
```bash
export VLLM_API_KEY=dummy
```
!!! warning
When using the `responses` API, ensure your vLLM version supports the OpenAI Responses API.
## Testing the Setup
Once Codex is configured, launch it in your project directory:
```bash
codex
```
Try a simple prompt to verify the connection, such as asking it to explain a file in your project. If the model responds correctly, your setup is working. You can now use Codex with your vLLM-served model for coding tasks.
## Troubleshooting
**Connection refused**: Ensure vLLM is running and accessible at the specified URL. Check that the port matches and that `base_url` includes the `/v1` path suffix.
**Tool calls not working**: Verify that your model supports tool calling and that you've enabled it with the correct `--tool-call-parser` flag. See [Tool Calling](../../features/tool_calling.md).
**Model not found**: Ensure the `model` field in `~/.codex/config.toml` matches the `--served-model-name` you passed to vLLM.
@@ -1,12 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
vLLM OpenAI-Compatible Client with Prompt Embeddings
"""vLLM OpenAI-Compatible Client with Prompt Embeddings.
This script demonstrates how to:
1. Generate prompt embeddings using Hugging Face Transformers
2. Encode them in base64 format
3. Send them to a vLLM server via the OpenAI-compatible Completions API
1. Generate prompt embeddings using Hugging Face Transformers.
2. Encode them in base64 format.
3. Send them to a vLLM server for inference via both:
- OpenAI-compatible Chat Completions API
- OpenAI-compatible Completions API
Important distinction between the two APIs:
- Chat Completions API: `prompt_embeds` content parts should encode ONLY
the user-provided content, not a templated conversation. The server
renders the surrounding chat template around the embedded content at
request time, the same way it would for a plain text `content` string.
Embedding a full templated conversation here would double-apply the
template and likely produce undesirable results.
- Completions API: the server does NOT apply a chat template to
`prompt_embeds`. The caller is responsible for producing embeddings for
the full, already-templated prompt (i.e. apply the chat template first,
then embed the resulting token IDs). Anything the model would normally
need (system prompt, role markers, generation prompt, etc.) must already
be baked into the embedded tokens.
Run the vLLM server first:
vllm serve meta-llama/Llama-3.2-1B-Instruct \
@@ -34,34 +51,68 @@ from openai import OpenAI
from vllm.utils.serial_utils import tensor2base64
def main():
client = OpenAI(
api_key="EMPTY",
base_url="http://localhost:8000/v1",
def run_chat_completion_prompt_embeds(
client: OpenAI,
model_name: str,
tokenizer: transformers.PreTrainedTokenizerBase,
embedding_layer,
messages: list[dict],
) -> None:
"""Run a Chat Completions API request using prompt_embeds content parts.
This example embeds ONLY the user-provided content of the final user turn, the
vLLM server applies the chat template around it at request time.
"""
user_content = messages[-1]["content"]
content_token_ids = tokenizer(
user_content, return_tensors="pt", add_special_tokens=False
).input_ids
content_prompt_embeds = embedding_layer(content_token_ids).squeeze(0)
encoded_embeds = tensor2base64(content_prompt_embeds)
api_messages = [
*messages[:-1],
{
"role": messages[-1]["role"],
"content": [{"type": "prompt_embeds", "data": encoded_embeds}],
},
]
chat_completion = client.chat.completions.create(
model=model_name,
max_tokens=6,
temperature=0.0,
messages=api_messages,
)
model_name = "meta-llama/Llama-3.2-1B-Instruct"
print("-" * 30)
print("Chat Completions API")
print(chat_completion.choices[0].message.content)
print("-" * 30)
# Transformers
tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
transformers_model = transformers.AutoModelForCausalLM.from_pretrained(model_name)
# Refer to the HuggingFace repo for the correct format to use
chat = [{"role": "user", "content": "Please tell me about the capital of France."}]
token_ids = tokenizer.apply_chat_template(
chat, add_generation_prompt=True, return_tensors="pt", return_dict=True
def run_completion_prompt_embeds(
client: OpenAI,
model_name: str,
tokenizer: transformers.PreTrainedTokenizerBase,
embedding_layer,
messages: list[dict],
) -> None:
"""Run a Completions API request using prompt embeddings.
The Completions endpoint does not apply a chat template,
so the caller must apply it and embed the full templated prompt.
"""
templated_token_ids = tokenizer.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="pt", return_dict=True
).input_ids
embedding_layer = transformers_model.get_input_embeddings()
prompt_embeds = embedding_layer(token_ids).squeeze(0)
# Prompt embeddings
encoded_embeds = tensor2base64(prompt_embeds)
templated_prompt_embeds = embedding_layer(templated_token_ids).squeeze(0)
encoded_embeds = tensor2base64(templated_prompt_embeds)
completion = client.completions.create(
model=model_name,
prompt=None,
max_tokens=5,
max_tokens=6,
temperature=0.0,
# NOTE: The OpenAI client allows passing in extra JSON body via the
# `extra_body` argument.
@@ -69,9 +120,39 @@ def main():
)
print("-" * 30)
print("Completions API")
print(completion.choices[0].text)
print("-" * 30)
def main() -> None:
client = OpenAI(
api_key="EMPTY",
base_url="http://localhost:8000/v1",
)
model_name = "meta-llama/Llama-3.2-1B-Instruct"
tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
transformers_model = transformers.AutoModelForCausalLM.from_pretrained(model_name)
embedding_layer = transformers_model.get_input_embeddings()
messages = [
{"role": "user", "content": "Please tell me about the capital of France."}
]
# Chat Completions API: embed ONLY the user content. The server wraps
# the embedding in the chat template when it renders the messages.
run_chat_completion_prompt_embeds(
client, model_name, tokenizer, embedding_layer, messages
)
# Completions API: embed the FULL templated prompt. The caller must
# apply the chat template up-front.
run_completion_prompt_embeds(
client, model_name, tokenizer, embedding_layer, messages
)
if __name__ == "__main__":
main()
@@ -2466,6 +2466,7 @@ MODELS_NEED_VIDEO_METADATA = [
MODELS_SUPPORT_VIT_CUDA_GRAPH = [
"qwen3_vl",
"qwen3_vl_moe",
"qwen2_5_vl",
]
+67 -44
View File
@@ -1,9 +1,9 @@
{%- macro format_parameters(properties, required) -%}
{%- macro format_parameters(properties, required, filter_keys=false) -%}
{%- set standard_keys = ['description', 'type', 'properties', 'required', 'nullable'] -%}
{%- set ns = namespace(found_first=false) -%}
{%- for key, value in properties | dictsort -%}
{%- set add_comma = false -%}
{%- if key not in standard_keys -%}
{%- if not filter_keys or key not in standard_keys -%}
{%- if ns.found_first %},{% endif -%}
{%- set ns.found_first = true -%}
{{ key }}:{
@@ -11,34 +11,15 @@
description:<|"|>{{ value['description'] }}<|"|>
{%- set add_comma = true -%}
{%- endif -%}
{%- if value['nullable'] %}
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
nullable:true
{%- endif -%}
{%- if value['type'] | upper == 'STRING' -%}
{%- if value['enum'] -%}
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
enum:{{ format_argument(value['enum']) }}
{%- endif -%}
{%- elif value['type'] | upper == 'OBJECT' -%}
,properties:{
{%- if value['properties'] is defined and value['properties'] is mapping -%}
{{- format_parameters(value['properties'], value['required'] | default([])) -}}
{%- elif value is mapping -%}
{{- format_parameters(value, value['required'] | default([])) -}}
{%- endif -%}
}
{%- if value['required'] -%}
,required:[
{%- for item in value['required'] | default([]) -%}
<|"|>{{- item -}}<|"|>
{%- if not loop.last %},{% endif -%}
{%- endfor -%}
]
{%- endif -%}
{%- elif value['type'] | upper == 'ARRAY' -%}
{%- if value['items'] is mapping and value['items'] -%}
,items:{
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
items:{
{%- set ns_items = namespace(found_first=false) -%}
{%- for item_key, item_value in value['items'] | dictsort -%}
{%- if item_value is not none -%}
@@ -71,6 +52,32 @@
}
{%- endif -%}
{%- endif -%}
{%- if value['nullable'] %}
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
nullable:true
{%- endif -%}
{%- if value['type'] | upper == 'OBJECT' -%}
{%- if value['properties'] is defined and value['properties'] is mapping -%}
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
properties:{
{{- format_parameters(value['properties'], value['required'] | default([])) -}}
}
{%- elif value is mapping -%}
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
properties:{
{{- format_parameters(value, value['required'] | default([]), filter_keys=true) -}}
}
{%- endif -%}
{%- if value['required'] -%}
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
required:[
{%- for item in value['required'] | default([]) -%}
<|"|>{{- item -}}<|"|>
{%- if not loop.last %},{% endif -%}
{%- endfor -%}
]
{%- endif -%}
{%- endif -%}
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
type:<|"|>{{ value['type'] | upper }}<|"|>}
{%- endif -%}
@@ -167,20 +174,25 @@
{%- set ns = namespace(prev_message_type=None) -%}
{%- set loop_messages = messages -%}
{{ bos_token }}
{{- bos_token -}}
{#- Handle System/Tool Definitions Block -#}
{%- if (enable_thinking is defined and enable_thinking) or tools or messages[0]['role'] in ['system', 'developer'] -%}
{{- '<|turn>system\n' -}}
{#- Inject Thinking token at the very top of the FIRST system turn -#}
{%- if enable_thinking is defined and enable_thinking -%}
{{- '<|think|>' -}}
{{- '<|think|>\n' -}}
{%- set ns.prev_message_type = 'think' -%}
{%- endif -%}
{%- if messages[0]['role'] in ['system', 'developer'] -%}
{{- messages[0]['content'] | trim -}}
{%- if messages[0]['content'] is string -%}
{{- messages[0]['content'] | trim -}}
{%- elif messages[0]['content'] is sequence -%}
{%- for item in messages[0]['content'] -%}
{{- item['text'] | trim + ' '-}}
{%- endfor -%}
{%- endif -%}
{%- set loop_messages = messages[1:] -%}
{%- endif -%}
{%- if tools -%}
{%- for tool in tools %}
{{- '<|tool>' -}}
@@ -189,10 +201,10 @@
{%- endfor %}
{%- set ns.prev_message_type = 'tool' -%}
{%- endif -%}
{{- '<turn|>\n' -}}
{%- endif %}
{#- Pre-scan: find last user message index for reasoning guard -#}
{%- set ns_turn = namespace(last_user_idx=-1) -%}
{%- for i in range(loop_messages | length) -%}
{%- if loop_messages[i]['role'] == 'user' -%}
@@ -200,12 +212,12 @@
{%- endif -%}
{%- endfor -%}
{#- Loop through messages -#}
{%- for message in loop_messages -%}
{%- if message['role'] != 'tool' -%}
{%- set ns.prev_message_type = None -%}
{%- set role = 'model' if message['role'] == 'assistant' else message['role'] -%}
{#- OpenAI may emit multiple assistant messages in one tool loop (user → asst → tool → asst → tool).
Only the first of those should open <|turn>model; later ones continue the same model turn. -#}
{#- Detect continuation: suppress duplicate <|turn>model when previous non-tool message was also assistant -#}
{%- set prev_nt = namespace(role=None, found=false) -%}
{%- if loop.index0 > 0 -%}
{%- for j in range(loop.index0 - 1, -1, -1) -%}
@@ -222,8 +234,10 @@
{{- '<|turn>' + role + '\n' }}
{%- endif -%}
{%- if message.get('reasoning') and loop.index0 > ns_turn.last_user_idx and message.get('tool_calls') -%}
{{- '<|channel>thought\n' + message['reasoning'] + '\n<channel|>'}}
{#- Render reasoning/reasoning_content as thinking channel -#}
{%- set thinking_text = message.get('reasoning') or message.get('reasoning_content') -%}
{%- if thinking_text and loop.index0 > ns_turn.last_user_idx and message.get('tool_calls') -%}
{{- '<|channel>thought\n' + thinking_text + '\n<channel|>' -}}
{%- endif -%}
{%- if message['tool_calls'] -%}
@@ -247,14 +261,14 @@
{%- set ns_tr_out = namespace(flag=false) -%}
{%- if message.get('tool_responses') -%}
{#- Legacy: tool_responses embedded on the assistant message -#}
{#- Legacy: tool_responses embedded on the assistant message (Google/Gemma native) -#}
{%- for tool_response in message['tool_responses'] -%}
{{- format_tool_response_block(tool_response['name'] | default('unknown'), tool_response['response']) -}}
{%- set ns_tr_out.flag = true -%}
{%- set ns.prev_message_type = 'tool_response' -%}
{%- endfor -%}
{%- elif message.get('tool_calls') -%}
{#- OpenAI Chat Completions: consecutive following messages with role "tool" (no break/continue; range scan) -#}
{#- OpenAI Chat Completions: forward-scan consecutive role:tool messages -#}
{%- set ns_tool_scan = namespace(stopped=false) -%}
{%- for k in range(loop.index0 + 1, loop_messages | length) -%}
{%- if ns_tool_scan.stopped -%}
@@ -262,12 +276,14 @@
{%- set ns_tool_scan.stopped = true -%}
{%- else -%}
{%- set follow = loop_messages[k] -%}
{#- Resolve tool_call_id to function name -#}
{%- set ns_tname = namespace(name=follow.get('name') | default('unknown')) -%}
{%- for tc in message['tool_calls'] -%}
{%- if tc.get('id') == follow.get('tool_call_id') -%}
{%- set ns_tname.name = tc['function']['name'] -%}
{%- endif -%}
{%- endfor -%}
{#- Handle content as string or content-parts array -#}
{%- set tool_body = follow.get('content') -%}
{%- if tool_body is string -%}
{{- format_tool_response_block(ns_tname.name, tool_body) -}}
@@ -288,6 +304,7 @@
{%- endfor -%}
{%- endif -%}
{%- set captured_content -%}
{%- if message['content'] is string -%}
{%- if role == 'model' -%}
{{- strip_thinking(message['content']) -}}
@@ -303,29 +320,35 @@
{{- item['text'] | trim -}}
{%- endif -%}
{%- elif item['type'] == 'image' -%}
{{- '\n\n<|image|>\n\n' -}}
{{- '<|image|>' -}}
{%- set ns.prev_message_type = 'image' -%}
{%- elif item['type'] == 'audio' -%}
{{- '<|audio|>' -}}
{%- set ns.prev_message_type = 'audio' -%}
{%- elif item['type'] == 'video' -%}
{{- '\n\n<|video|>\n\n' -}}
{{- '<|video|>' -}}
{%- set ns.prev_message_type = 'video' -%}
{%- endif -%}
{%- endfor -%}
{%- endif -%}
{%- endset -%}
{%- if not (ns_tr_out.flag and not message.get('content')) -%}
{{- captured_content -}}
{%- set has_content = captured_content | trim | length > 0 -%}
{%- if ns.prev_message_type == 'tool_call' and not ns_tr_out.flag -%}
{{- '<|tool_response>' -}}
{%- elif not (ns_tr_out.flag and not has_content) -%}
{{- '<turn|>\n' -}}
{%- endif -%}
{%- endif -%}
{%- endfor -%}
{%- if add_generation_prompt -%}
{%- if ns.prev_message_type != 'tool_response' -%}
{%- if ns.prev_message_type != 'tool_response' and ns.prev_message_type != 'tool_call' -%}
{{- '<|turn>model\n' -}}
{%- if not enable_thinking | default(false) -%}
{{- '<|channel>thought\n<channel|>' -}}
{%- endif -%}
{%- endif -%}
{%- if not enable_thinking | default(false) -%}
{{- '<|channel>thought\n<channel|>' -}}
{%- endif -%}
{%- endif -%}
{%- endif -%}
+1 -2
View File
@@ -105,8 +105,7 @@ plugins:
- https://docs.aiohttp.org/en/stable/objects.inv
- https://pillow.readthedocs.io/en/stable/objects.inv
- https://numpy.org/doc/stable/objects.inv
# TODO revert to stable once https://github.com/pytorch/pytorch/issues/182007 is fixed
- https://pytorch.org/docs/2.11/objects.inv
- https://pytorch.org/docs/stable/objects.inv
- redirects:
redirect_maps:
features/spec_decode/README.md: features/speculative_decoding/README.md
+1 -1
View File
@@ -24,7 +24,7 @@ outlines_core == 0.2.14
# required for outlines backend disk cache
diskcache == 5.6.3
lark == 1.2.2
xgrammar >= 0.1.32, < 1.0.0; platform_machine == "x86_64" or platform_machine == "aarch64" or platform_machine == "arm64" or platform_machine == "s390x" or platform_machine == "ppc64le"
xgrammar >= 0.2.0, < 1.0.0; platform_machine == "x86_64" or platform_machine == "aarch64" or platform_machine == "arm64" or platform_machine == "s390x" or platform_machine == "ppc64le"
typing_extensions >= 4.10
filelock >= 3.16.1 # need to contain https://github.com/tox-dev/filelock/pull/317
partial-json-parser # used for parsing partial JSON outputs
+4 -1
View File
@@ -42,6 +42,8 @@ anyio==4.13.0
# sse-starlette
# starlette
# watchfiles
apache-tvm-ffi==0.1.10
# via xgrammar
arctic-inference==0.1.1
# via -r requirements/test/rocm.in
argcomplete==3.6.3
@@ -1264,6 +1266,7 @@ typing-extensions==4.15.0
# alembic
# anthropic
# anyio
# apache-tvm-ffi
# azure-core
# azure-identity
# azure-storage-blob
@@ -1345,7 +1348,7 @@ word2number==1.1
# via lm-eval
wrapt==2.1.2
# via smart-open
xgrammar==0.1.33
xgrammar==0.2.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
@@ -0,0 +1,77 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import argparse
import json
from pathlib import Path
import pytest
from transformers import AutoTokenizer, PreTrainedTokenizerBase
from vllm.benchmarks.datasets import get_samples
@pytest.fixture(scope="session")
def hf_tokenizer() -> PreTrainedTokenizerBase:
return AutoTokenizer.from_pretrained("gpt2")
def _write_jsonl(path: Path, n_rows: int) -> None:
with path.open("w") as f:
for i in range(n_rows):
f.write(json.dumps({"prompt": f"row {i}: unique prompt content."}) + "\n")
def _args_for_custom(dataset_path: str, seed: int) -> argparse.Namespace:
return argparse.Namespace(
dataset_name="custom",
dataset_path=dataset_path,
disable_shuffle=False,
num_prompts=30,
custom_output_len=32,
skip_chat_template=True,
no_oversample=False,
seed=seed,
request_id_prefix="",
)
@pytest.mark.benchmark
def test_custom_dataset_seed_propagates(
hf_tokenizer: PreTrainedTokenizerBase, tmp_path: Path
) -> None:
"""--seed must control the CustomDataset shuffle used by get_samples.
Without the fix, CustomDataset was instantiated without random_seed,
so its load-time shuffle always used DEFAULT_SEED=0 regardless of
args.seed, causing every run with --dataset-name custom to pick the
same subset of rows from a larger file.
"""
jsonl = tmp_path / "data.jsonl"
_write_jsonl(jsonl, n_rows=60)
samples_a = get_samples(_args_for_custom(str(jsonl), seed=0), hf_tokenizer)
samples_b = get_samples(_args_for_custom(str(jsonl), seed=42), hf_tokenizer)
prompts_a = {s.prompt for s in samples_a}
prompts_b = {s.prompt for s in samples_b}
assert len(prompts_a) == 30
assert len(prompts_b) == 30
assert prompts_a != prompts_b
@pytest.mark.benchmark
def test_custom_dataset_same_seed_is_deterministic(
hf_tokenizer: PreTrainedTokenizerBase, tmp_path: Path
) -> None:
"""Same --seed must yield the same CustomDataset subset."""
jsonl = tmp_path / "data.jsonl"
_write_jsonl(jsonl, n_rows=60)
samples_a = get_samples(_args_for_custom(str(jsonl), seed=7), hf_tokenizer)
samples_b = get_samples(_args_for_custom(str(jsonl), seed=7), hf_tokenizer)
prompts_a = [s.prompt for s in samples_a]
prompts_b = [s.prompt for s in samples_b]
assert prompts_a == prompts_b
+15 -2
View File
@@ -12,10 +12,17 @@ from torch._ops import OpOverload, OpOverloadPacket
from torch.fx._utils import lazy_format_graph_code
from vllm.compilation.passes.fx_utils import find_op_nodes
from vllm.compilation.passes.inductor_pass import InductorPass
from vllm.compilation.passes.inductor_pass import (
InductorPass,
pass_context,
)
from vllm.compilation.passes.ir.inplace_functionalization import (
VllmIRInplaceFunctionalizationPass,
)
from vllm.compilation.passes.pass_manager import with_pattern_match_debug
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
from vllm.config import VllmConfig, get_current_vllm_config
from vllm.config.utils import Range
from vllm.logger import init_logger
logger = init_logger("vllm.tests.compile.backend")
@@ -53,11 +60,17 @@ class TestBackend:
self.custom_passes = list(passes)
vllm_config = get_current_vllm_config()
compile_config = vllm_config.compilation_config
self.range = Range(1, vllm_config.scheduler_config.max_num_batched_tokens)
# Deepcopy to allow multiple TestBackend instances to use the same VllmConfig
self.inductor_config = deepcopy(compile_config.inductor_compile_config)
self.inductor_config["force_disable_caches"] = True
self.inductor_config["post_grad_custom_post_pass"] = self.post_pass
# Add VllmIRInplaceFunctionalizationPass as pre-grad pass by default
self.inductor_config["pre_grad_custom_pass"] = (
VllmIRInplaceFunctionalizationPass(vllm_config)
)
if debug_dump_path := vllm_config.compile_debug_dump_path():
logger.debug("Dumping depyf output to %s", debug_dump_path)
self.debug_ctx = depyf.prepare_debug(debug_dump_path.as_posix())
@@ -68,7 +81,7 @@ class TestBackend:
self.graph_pre_compile = deepcopy(graph)
from torch._inductor.compile_fx import compile_fx
with self.debug_ctx:
with self.debug_ctx, pass_context(self.range):
return compile_fx(
graph, example_inputs, config_patches=self.inductor_config
)
+17 -3
View File
@@ -24,10 +24,24 @@ def mock_cuda_platform():
def _mock_platform(is_cuda: bool = True, capability: tuple[int, int] | None = None):
mock_platform = MagicMock()
mock_platform.is_cuda.return_value = is_cuda
if capability is not None:
mock_platform.get_device_capability.return_value = DeviceCapability(
*capability
device_capability = (
DeviceCapability(*capability) if capability is not None else None
)
mock_platform.get_device_capability.return_value = device_capability
def is_device_capability_family(
requested_capability: int, device_id: int = 0
) -> bool:
current_capability = mock_platform.get_device_capability(
device_id=device_id
)
if current_capability is None:
return False
return current_capability.major == (requested_capability // 10)
mock_platform.is_device_capability_family.side_effect = (
is_device_capability_family
)
with patch("vllm.platforms.current_platform", mock_platform):
yield mock_platform
+6
View File
@@ -97,6 +97,12 @@ def run_e2e_fusion_test(monkeypatch, caplog_mp_spawn):
f"attention backend '{attn_backend.backend.name}'"
)
if attn_backend.backend.name == "FLASHINFER":
from vllm.utils.flashinfer import supports_trtllm_attention
if not supports_trtllm_attention():
matches = matches._replace(attn_quant_fusion=0)
# TODO: remove this after finishing migration from envs to model kwargs
if model_name == "openai/gpt-oss-20b":
from .common import is_blackwell
+19 -3
View File
@@ -19,6 +19,8 @@ from .models import (
FLASHINFER_ATTN,
FLASHINFER_MLA_ATTN,
FLASHMLA_SPARSE_ATTN,
ROCM_AITER_UNIFIED_ATTN,
ROCM_ATTN,
TRITON_ATTN,
deepseek_coder_v2_lite_fp8,
deepseek_r1_fp4,
@@ -34,7 +36,9 @@ from .models import (
qwen3_a3b_fp8,
)
pytestmark = pytest.mark.skipif(not current_platform.is_cuda(), reason="Only test CUDA")
pytestmark = pytest.mark.skipif(
not current_platform.is_cuda_alike(), reason="Only test CUDA/ROCm"
)
@multi_gpu_test(num_gpus=2)
@@ -55,6 +59,7 @@ pytestmark = pytest.mark.skipif(not current_platform.is_cuda(), reason="Only tes
@pytest.mark.parametrize("n_layers", [4])
@pytest.mark.parametrize("custom_ops", custom_ops_combos("quant_fp8", "rms_norm"))
@pytest.mark.parametrize("inductor_graph_partition", INDUCTOR_GRAPH_PARTITION)
@pytest.mark.skipif(not current_platform.is_cuda(), reason="Only test CUDA")
def test_tp2_ar_rms_fp8_fusions(
model_name: str,
matches_fn: Callable[[int], Matches],
@@ -124,6 +129,7 @@ def test_tp2_ar_rms_fp8_fusions(
@pytest.mark.parametrize("custom_ops", custom_ops_combos("rms_norm"))
@pytest.mark.parametrize("inductor_graph_partition", INDUCTOR_GRAPH_PARTITION)
@pytest.mark.skipif(not is_blackwell(), reason="Blackwell required for fp4")
@pytest.mark.skipif(not current_platform.is_cuda(), reason="Only test CUDA")
def test_tp2_ar_rms_fp4_fusions(
model_name: str,
matches_fn: Callable[[int], Matches],
@@ -176,10 +182,19 @@ def test_tp2_ar_rms_fp4_fusions(
"model_name, matches_fn, model_kwargs, hf_overrides",
[llama3_8b, qwen3_a3b, gpt_oss_20b],
)
@pytest.mark.parametrize("attn_backend", [TRITON_ATTN])
@pytest.mark.parametrize(
"attn_backend",
[
TRITON_ATTN,
FLASHINFER_ATTN,
ROCM_ATTN,
ROCM_AITER_UNIFIED_ATTN,
],
)
@pytest.mark.parametrize("n_layers", [4])
@pytest.mark.parametrize("custom_ops", custom_ops_combos("rms_norm"))
@pytest.mark.parametrize("custom_ops", tuple(custom_ops_combos("rms_norm")))
@pytest.mark.parametrize("inductor_graph_partition", INDUCTOR_GRAPH_PARTITION)
@pytest.mark.skipif(not current_platform.is_cuda_alike(), reason="Only test CUDA/ROCm")
def test_tp2_ar_rms_fusions(
model_name: str,
matches_fn: Callable[[int], Matches],
@@ -221,4 +236,5 @@ def test_tp2_ar_rms_fusions(
compilation_config,
matches_check,
tp_size=2,
use_aiter=current_platform.is_rocm(),
)
@@ -13,7 +13,6 @@ from .common import (
AttentionBackendCase,
Matches,
custom_ops_combos,
is_blackwell,
)
from .models import (
FLASHINFER_ATTN,
@@ -46,14 +45,9 @@ def test_tp2_async_tp_fp8_fusions(
custom_ops: str,
inductor_graph_partition: bool,
run_e2e_fusion_test,
monkeypatch,
):
matches = matches_fn(n_layers)
if is_blackwell():
# Disable FlashInfer scaled_mm FP8 as it's not supported in async tp patterns
monkeypatch.setenv("VLLM_DISABLED_KERNELS", "FlashInferFP8ScaledMMLinearKernel")
# Reduce size of model and skip weight loading time
model_kwargs["hf_overrides"] = hf_overrides(n_layers)
model_kwargs["load_format"] = "dummy"
@@ -173,14 +167,9 @@ def test_tp2_sp_ar_rms_fp8_fusions(
custom_ops: str,
inductor_graph_partition: bool,
run_e2e_fusion_test,
monkeypatch,
):
matches = matches_fn(n_layers)
if is_blackwell():
# Disable FlashInfer scaled_mm FP8 as it's not supported in async tp patterns
monkeypatch.setenv("VLLM_DISABLED_KERNELS", "FlashInferFP8ScaledMMLinearKernel")
# Reduce size of model and skip weight loading time
model_kwargs["hf_overrides"] = hf_overrides(n_layers)
model_kwargs["load_format"] = "dummy"
@@ -8,8 +8,12 @@ import torch
import vllm.envs as envs
from tests.compile.backend import TestBackend
from tests.utils import TestFP8Layer, has_module_attribute, multi_gpu_test
from vllm._aiter_ops import IS_AITER_FOUND, rocm_aiter_ops
from vllm._custom_ops import cutlass_scaled_fp4_mm, scaled_fp4_quant
from vllm.compilation.passes.fusion.allreduce_rms_fusion import AllReduceFusionPass
from vllm.compilation.passes.fusion.allreduce_rms_fusion import (
AllReduceFusionPass,
RocmAiterAllReduceFusionPass,
)
from vllm.compilation.passes.utility.fix_functionalization import (
FixFunctionalizationPass,
)
@@ -42,13 +46,19 @@ DEVICE_TYPE = current_platform.device_type
class TestAllReduceRMSNormModel(torch.nn.Module):
def __init__(
self, hidden_size=16, token_num=16, eps=1e-6, dtype: torch.dtype = torch.float16
self,
hidden_size=16,
token_num=16,
eps=1e-6,
dtype: torch.dtype = torch.float16,
use_aiter: bool = False,
):
super().__init__()
self.hidden_size = hidden_size
self.eps = eps
self.norm = [RMSNorm(hidden_size, eps) for i in range(4)]
self.w = [torch.rand(hidden_size, hidden_size) for _ in range(3)]
self.use_aiter = use_aiter
def forward(self, x):
# avoid having graph input be an arg to a pattern directly
@@ -76,6 +86,8 @@ class TestAllReduceRMSNormModel(torch.nn.Module):
return [torch.ops.vllm.all_reduce.default]
def ops_in_model_after(self):
if self.use_aiter:
return [rocm_aiter_ops.get_fused_allreduce_rmsnorm_op()]
return [torch.ops.vllm.flashinfer_trtllm_fused_allreduce_norm.default]
@@ -194,12 +206,36 @@ class TestAllReduceFusedAddRMSNormStaticQuantFP4Model(torch.nn.Module):
@multi_gpu_test(num_gpus=2)
@pytest.mark.parametrize(
"test_model, enable_quant_fp8_custom_op",
"test_model, enable_quant_fp8_custom_op, use_aiter",
[
(TestAllReduceRMSNormModel, False),
(TestAllReduceRMSNormStaticQuantFP8Model, True),
(TestAllReduceRMSNormStaticQuantFP8Model, False),
(TestAllReduceFusedAddRMSNormStaticQuantFP4Model, False),
(TestAllReduceRMSNormModel, False, IS_AITER_FOUND),
pytest.param(
TestAllReduceRMSNormStaticQuantFP8Model,
True,
False,
marks=pytest.mark.skipif(
current_platform.is_rocm(),
reason="Not supported on ROCm platform",
),
),
pytest.param(
TestAllReduceRMSNormStaticQuantFP8Model,
False,
False,
marks=pytest.mark.skipif(
current_platform.is_rocm(),
reason="Not supported on ROCm platform",
),
),
pytest.param(
TestAllReduceFusedAddRMSNormStaticQuantFP4Model,
False,
False,
marks=pytest.mark.skipif(
current_platform.is_rocm(),
reason="Not supported on ROCm platform",
),
),
],
)
@pytest.mark.parametrize("batch_size", [8])
@@ -210,9 +246,18 @@ class TestAllReduceFusedAddRMSNormStaticQuantFP4Model(torch.nn.Module):
@pytest.mark.parametrize("flashinfer_allreduce_backend", ["trtllm", "mnnvl"])
@pytest.mark.skipif(envs.VLLM_TARGET_DEVICE not in ["cuda"], reason="Only test on CUDA")
@pytest.mark.skipif(
not find_spec("flashinfer")
or not has_module_attribute("flashinfer.comm", "allreduce_fusion")
or not has_module_attribute("flashinfer.comm", "create_allreduce_fusion_workspace"),
current_platform.is_rocm() and not IS_AITER_FOUND,
reason="aiter is not found",
)
@pytest.mark.skipif(
current_platform.is_cuda()
and (
not find_spec("flashinfer")
or not has_module_attribute("flashinfer.comm", "allreduce_fusion")
or not has_module_attribute(
"flashinfer.comm", "create_allreduce_fusion_workspace"
)
),
reason="flashinfer is not found or flashinfer "
"is not compiled with allreduce_fusion",
)
@@ -225,7 +270,14 @@ def test_all_reduce_fusion_pass_replace(
enable_rms_norm_custom_op,
enable_quant_fp8_custom_op,
flashinfer_allreduce_backend,
use_aiter: bool,
monkeypatch: pytest.MonkeyPatch,
):
if use_aiter:
with monkeypatch.context() as m:
m.setenv("VLLM_ROCM_USE_AITER", str(use_aiter))
rocm_aiter_ops.refresh_env_variables()
num_processes = 2
if (
test_model == TestAllReduceFusedAddRMSNormStaticQuantFP4Model
@@ -249,6 +301,8 @@ def test_all_reduce_fusion_pass_replace(
enable_rms_norm_custom_op,
enable_quant_fp8_custom_op,
flashinfer_allreduce_backend,
use_aiter,
monkeypatch,
),
nprocs=nprocs,
)
@@ -267,6 +321,8 @@ def all_reduce_fusion_pass_on_test_model(
enable_rms_norm_custom_op,
enable_quant_fp8_custom_op,
flashinfer_allreduce_backend,
use_aiter: bool,
monkeypatch: pytest.MonkeyPatch,
):
set_random_seed(0)
@@ -313,7 +369,11 @@ def all_reduce_fusion_pass_on_test_model(
)
with set_current_vllm_config(vllm_config):
initialize_model_parallel(tensor_model_parallel_size=world_size)
all_reduce_fusion_pass = AllReduceFusionPass(vllm_config)
all_reduce_fusion_pass = (
RocmAiterAllReduceFusionPass(vllm_config)
if use_aiter
else AllReduceFusionPass(vllm_config)
)
noop_pass = NoOpEliminationPass(vllm_config)
func_pass = FixFunctionalizationPass(vllm_config)
cleanup_pass = PostCleanupPass(vllm_config)
@@ -323,7 +383,12 @@ def all_reduce_fusion_pass_on_test_model(
)
token_num = batch_size * seq_len
model = test_model_cls(hidden_size, token_num, dtype=dtype)
if test_model_cls is TestAllReduceRMSNormModel:
model = test_model_cls(
hidden_size, token_num, dtype=dtype, use_aiter=use_aiter
)
else:
model = test_model_cls(hidden_size, token_num, dtype=dtype)
hidden_states = torch.randn((token_num, hidden_size), requires_grad=False)
@@ -88,14 +88,10 @@ class TestAllReduceRMSNormModel(torch.nn.Module):
]
def ops_in_model(self):
return (
[torch.ops.vllm_ir.rms_norm]
+ [
torch.ops._C.fused_add_rms_norm.default,
]
if RMSNorm.enabled()
else []
)
return [
torch.ops.vllm_ir.rms_norm,
torch.ops.vllm_ir.fused_add_rms_norm,
]
class TestAllReduceRMSNormStaticQuantFP8Model(torch.nn.Module):
@@ -152,16 +148,17 @@ class TestAllReduceRMSNormStaticQuantFP8Model(torch.nn.Module):
def ops_in_model(self):
if self.vllm_config.compilation_config.pass_config.fuse_norm_quant:
return [torch.ops._C.fused_add_rms_norm_static_fp8_quant.default]
elif RMSNorm.enabled():
return [
torch.ops._C.fused_add_rms_norm.default,
]
elif any(layer.is_quant_fp8_enabled() for layer in self.fp8_linear_layers):
return [
torch.ops._C.static_scaled_fp8_quant.default,
]
else:
return []
quant_ops = (
[torch.ops._C.static_scaled_fp8_quant.default]
if any(layer.is_quant_fp8_enabled() for layer in self.fp8_linear_layers)
else [torch.ops.aten.reciprocal]
)
return [
torch.ops.vllm_ir.rms_norm,
torch.ops.vllm_ir.fused_add_rms_norm,
*quant_ops,
]
@multi_gpu_test(num_gpus=2)
@@ -0,0 +1,412 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Comprehensive tests for UnsafeCloneEliminationPass.
This test suite exercises all possible valid FX graph patterns involving clones:
1. Clone with no users (dead code)
2. Clone with read-only users
3. Clone with mutation users
4. Clone of graph input
5. Clone with original used after mutation
6. Clone chains
"""
import pytest
import torch
from torch import fx
from torch.fx.experimental.proxy_tensor import make_fx
from vllm.compilation.passes.fx_utils import find_op_nodes
from vllm.compilation.passes.inductor_pass import get_pass_context, pass_context
from vllm.compilation.passes.ir.clone_elimination import (
UnsafeCloneEliminationPass,
user_writes_to_node,
)
from vllm.config import VllmConfig
from vllm.config.utils import Range
def count_clones(graph: fx.Graph) -> int:
"""Count clone nodes in a graph."""
return len(list(find_op_nodes(torch.ops.aten.clone.default, graph)))
@pytest.fixture(scope="function")
def clone_cleanup_pass():
return UnsafeCloneEliminationPass(VllmConfig())
@pytest.fixture(autouse=True)
def setup_pass_context():
"""Set up pass context for each test."""
with pass_context(compile_range=Range(1, 8192)):
yield
class TestCloneCleanup:
"""Test UnsafeCloneEliminationPass behavior on various graph patterns."""
def test_remove_clone_readonly_users(self, clone_cleanup_pass):
"""Clone with only read-only users should be removed."""
def f(x: torch.Tensor) -> torch.Tensor:
x_clone = x.clone()
return x_clone + 1
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 1
expected = graph_module(inp)
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
actual = graph_module(inp)
assert count_clones(graph_module.graph) == 0
torch.testing.assert_close(actual, expected)
def test_keep_clone_with_mutation_and_original_used_after(self, clone_cleanup_pass):
"""Clone must be kept if it's mutated AND original is used after mutation."""
def f(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
x = x.relu() # not a graph param
x_clone = x.clone()
x_clone.add_(1)
return x, x_clone
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 1
expected = graph_module(inp)
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
actual = graph_module(inp)
# Clone should be KEPT because original is used after mutation
assert count_clones(graph_module.graph) == 1
torch.testing.assert_close(actual[0], expected[0])
torch.testing.assert_close(actual[1], expected[1])
def test_remove_clone_with_mutation_no_original_use(self, clone_cleanup_pass):
"""Clone can be removed if it's mutated but original is not used after."""
def f(x: torch.Tensor) -> torch.Tensor:
x = x.relu() # not a graph param
x_clone = x.clone()
x_clone.add_(1)
return x_clone
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 1
expected = graph_module(inp)
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
actual = graph_module(inp)
assert count_clones(graph_module.graph) == 0
torch.testing.assert_close(actual, expected)
def test_clone_chain(self, clone_cleanup_pass):
"""Test handling of clone chains: x -> clone1 -> clone2."""
def f(x: torch.Tensor) -> torch.Tensor:
x = x.relu() # not a graph param
x1 = x.clone()
x2 = x1.clone()
return x2 + 1
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 2
expected = graph_module(inp)
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
actual = graph_module(inp)
# Both clones should be removed
assert count_clones(graph_module.graph) == 0
torch.testing.assert_close(actual, expected)
def test_multiple_clones_of_same_input(self, clone_cleanup_pass):
"""Test multiple independent clones of the same input."""
def f(x: torch.Tensor) -> torch.Tensor:
x1 = x.clone()
x2 = x.clone()
return x1 + x2
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 2
expected = graph_module(inp)
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
actual = graph_module(inp)
# Both clones should be removed (only readonly uses)
assert count_clones(graph_module.graph) == 0
torch.testing.assert_close(actual, expected)
def test_no_clones_in_graph(self, clone_cleanup_pass):
"""Test pass behavior when graph has no clones."""
def f(x: torch.Tensor) -> torch.Tensor:
return x + 1
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 0
expected = graph_module(inp)
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
actual = graph_module(inp)
assert count_clones(graph_module.graph) == 0
torch.testing.assert_close(actual, expected)
def test_multiple_passes(self, clone_cleanup_pass):
"""Test running the pass multiple times (should be idempotent)."""
def f(x: torch.Tensor) -> torch.Tensor:
x1 = x.clone()
return x1 + 1
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 1
expected = graph_module(inp)
clone_cleanup_pass(graph_module.graph)
assert count_clones(graph_module.graph) == 0
graph_module.recompile()
actual = graph_module(inp)
torch.testing.assert_close(actual, expected)
clone_cleanup_pass(graph_module.graph)
assert count_clones(graph_module.graph) == 0
graph_module.recompile()
actual = graph_module(inp)
torch.testing.assert_close(actual, expected)
def test_output_node_no_write(self):
"""Output nodes never write to their inputs."""
def f(x: torch.Tensor) -> torch.Tensor:
return x
graph_module = make_fx(f)(torch.randn(2, 3))
x_node = [n for n in graph_module.graph.nodes if n.op == "placeholder"][0]
output_node = [n for n in graph_module.graph.nodes if n.op == "output"][0]
assert not user_writes_to_node(output_node, x_node)
def test_readonly_op_no_write(self):
"""Readonly operations don't write to inputs."""
def f(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
return x + y
graph_module = make_fx(f)(torch.randn(2, 3), torch.randn(2, 3))
placeholders = [n for n in graph_module.graph.nodes if n.op == "placeholder"]
add_node = [
n
for n in graph_module.graph.nodes
if n.op == "call_function" and n.target == torch.ops.aten.add.Tensor
][0]
assert not user_writes_to_node(add_node, placeholders[0])
assert not user_writes_to_node(add_node, placeholders[1])
def test_inplace_op_writes(self):
"""Inplace operations write to first argument."""
def f(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
x.add_(y)
return x
graph_module = make_fx(f)(torch.randn(2, 3), torch.randn(2, 3))
placeholders = [n for n in graph_module.graph.nodes if n.op == "placeholder"]
add_node = [
n
for n in graph_module.graph.nodes
if n.op == "call_function" and "add_" in str(n.target)
][0]
# add_ writes to first arg but not second
assert user_writes_to_node(add_node, placeholders[0])
assert not user_writes_to_node(add_node, placeholders[1])
def test_copy_writes(self):
"""copy_ operation writes to first argument."""
def f(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
x.copy_(y)
return x
graph_module = make_fx(f)(torch.randn(2, 3), torch.randn(2, 3))
placeholders = [n for n in graph_module.graph.nodes if n.op == "placeholder"]
copy_node = [
n
for n in graph_module.graph.nodes
if n.op == "call_function" and "copy_" in str(n.target)
][0]
assert user_writes_to_node(copy_node, placeholders[0])
assert not user_writes_to_node(copy_node, placeholders[1])
def test_auto_functionalized_not_a_write(self):
"""auto_functionalized ops are follow-up uses, not writes."""
from torch._higher_order_ops.auto_functionalize import auto_functionalized
def f(x: torch.Tensor) -> torch.Tensor:
return x
graph_module = make_fx(f)(torch.randn(2, 3))
x_node = [n for n in graph_module.graph.nodes if n.op == "placeholder"][0]
# Create an auto_functionalized node in the graph
with graph_module.graph.inserting_before(None):
af_node = graph_module.graph.call_function(
auto_functionalized, kwargs={"input": x_node}
)
# auto_functionalized should not be treated as a write
assert not user_writes_to_node(af_node, x_node)
def test_higher_order_op_conservatively_writes(self):
"""Other higher-order operators are conservatively treated as writes."""
from torch._ops import HigherOrderOperator
def f(x: torch.Tensor) -> torch.Tensor:
return x
graph_module = make_fx(f)(torch.randn(2, 3))
x_node = [n for n in graph_module.graph.nodes if n.op == "placeholder"][0]
# Create a concrete higher-order operator subclass
class MockHigherOrderOp(HigherOrderOperator):
def __call__(self, *args, **kwargs):
return args[0] if args else None
mock_hoo = MockHigherOrderOp("mock_higher_order_op")
with graph_module.graph.inserting_before(None):
hoo_node = graph_module.graph.call_function(mock_hoo, args=(x_node,))
# Should be conservative and assume it could write
assert user_writes_to_node(hoo_node, x_node)
class TestCloneCleanupWithDonatedInputs:
"""Test UnsafeCloneEliminationPass with donated input tracking via PassContext."""
@pytest.fixture(autouse=True)
def setup_pass_context(self):
"""Set up pass context for each test."""
with pass_context(compile_range=Range(1, 8192)):
yield
def test_donated_input_clone_removed(self, clone_cleanup_pass):
"""Clone of donated input should be removed."""
def f(x: torch.Tensor) -> torch.Tensor:
x_clone = x.clone()
x_clone.add_(1)
return x_clone
inp = torch.randn(2, 3)
graph_module = make_fx(f)(inp)
assert count_clones(graph_module.graph) == 1
# Mark first parameter as donated
get_pass_context().donated_input_ids = {0}
expected = graph_module(inp.clone())
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
# Clone should be removed since input is donated
assert count_clones(graph_module.graph) == 0
# Input can be mutated (donated)
inp_copy = inp.clone()
actual = graph_module(inp_copy)
torch.testing.assert_close(actual, expected)
def test_non_donated_input_clone_kept(self, clone_cleanup_pass):
"""Clone of non-donated input with mutation should be kept."""
def f(x: torch.Tensor, y: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
x_clone = x.clone()
x_clone.add_(1)
return x, x_clone
inp_x = torch.randn(2, 3)
inp_y = torch.randn(2, 3)
graph_module = make_fx(f)(inp_x, inp_y)
assert count_clones(graph_module.graph) == 1
# No donated inputs
get_pass_context().donated_input_ids = set()
expected = graph_module(inp_x.clone(), inp_y.clone())
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
# Clone should be kept since input is not donated and original is used
assert count_clones(graph_module.graph) == 1
# Verify inputs are not mutated
inp_x_before = inp_x.clone()
inp_y_before = inp_y.clone()
actual = graph_module(inp_x, inp_y)
torch.testing.assert_close(
inp_x, inp_x_before, msg="Input x should not be mutated"
)
torch.testing.assert_close(
inp_y, inp_y_before, msg="Input y should not be mutated"
)
torch.testing.assert_close(actual[0], expected[0])
torch.testing.assert_close(actual[1], expected[1])
def test_mixed_donated_inputs(self, clone_cleanup_pass):
"""Test with some inputs donated and some not."""
def f(x: torch.Tensor, y: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
x_clone = x.clone()
x_clone.add_(1)
y_clone = y.clone()
y_clone.add_(2)
return x_clone, y_clone
inp_x = torch.randn(2, 3)
inp_y = torch.randn(2, 3)
graph_module = make_fx(f)(inp_x, inp_y)
assert count_clones(graph_module.graph) == 2
# Only x is donated
get_pass_context().donated_input_ids = {0}
expected = graph_module(inp_x.clone(), inp_y.clone())
clone_cleanup_pass(graph_module.graph)
graph_module.recompile()
# x_clone removed (x is donated), y_clone kept (y is not donated)
assert count_clones(graph_module.graph) == 1
# Verify y is not mutated (x can be mutated since it's donated)
inp_y_before = inp_y.clone()
actual = graph_module(inp_x.clone(), inp_y)
torch.testing.assert_close(
inp_y, inp_y_before, msg="Input y should not be mutated"
)
torch.testing.assert_close(actual[0], expected[0])
torch.testing.assert_close(actual[1], expected[1])
@@ -0,0 +1,465 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Tests for IR inplace functionalization pass integration.
This test suite verifies that the inplace functionalization pass, lowering pass,
and clone cleanup pass work together correctly with donated buffer tracking.
"""
from collections.abc import Callable
import pytest
import torch
import torch._dynamo.exc
from torch import nn
import vllm.kernels # noqa: F401 to register kernels
from vllm.compilation.passes.inductor_pass import InductorPass, get_pass_context
from vllm.compilation.passes.ir.clone_elimination import (
UnsafeCloneEliminationPass,
)
from vllm.compilation.passes.ir.inplace_functionalization import (
VllmIRInplaceFunctionalizationPass,
)
from vllm.compilation.passes.ir.lowering_pass import VllmIRLoweringPass
from vllm.config import VllmConfig
from vllm.ir import ops
from vllm.platforms import current_platform
from vllm.triton_utils import HAS_TRITON, tl, triton
from ...backend import TestBackend
class StoreDonationInfoPass(InductorPass):
def __init__(self):
self.donated_input_ids_sets: list[set[int]] = []
def __call__(self, *args, **kwargs):
ctx = get_pass_context()
self.donated_input_ids_sets += [ctx.donated_input_ids]
class MaybeInplaceModel(nn.Module):
"""Model using only maybe_inplace variants."""
def __init__(self, hidden_size=16):
super().__init__()
self.weight1 = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
self.weight2 = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
def forward(
self, x: torch.Tensor, residual1: torch.Tensor, residual2: torch.Tensor
):
# First maybe_inplace - x & residual1 are donated
x_normed1, residual_out1 = ops.fused_add_rms_norm.maybe_inplace(
x, residual1, self.weight1, 1e-5
)
# Second maybe_inplace - residual2 is donated
x_normed2, residual_out2 = ops.fused_add_rms_norm.maybe_inplace(
x_normed1, residual2, self.weight2, 1e-5
)
return x_normed2, residual_out1, residual_out2
class FunctionalModel(nn.Module):
"""Model using only functional (default) variants."""
def __init__(self, hidden_size=16):
super().__init__()
self.weight1 = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
self.weight2 = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
def forward(
self, x: torch.Tensor, residual1: torch.Tensor, residual2: torch.Tensor
):
# First functional - no donation
x_normed1, residual_out1 = ops.fused_add_rms_norm(
x, residual1, self.weight1, 1e-5
)
# Second functional - no donation
x_normed2, residual_out2 = ops.fused_add_rms_norm(
x_normed1, residual2, self.weight2, 1e-5
)
return x_normed2, residual_out1, residual_out2
class MixedModel(nn.Module):
"""Model mixing maybe_inplace and functional variants."""
def __init__(self, hidden_size=16):
super().__init__()
self.weight1 = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
self.weight2 = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
def forward(
self, x: torch.Tensor, residual1: torch.Tensor, residual2: torch.Tensor
):
# First maybe_inplace - x & residual1 are donated
x_normed1, residual_out1 = ops.fused_add_rms_norm.maybe_inplace(
x, residual1, self.weight1, 1e-5
)
# Second functional - no donation, x_normed1 must be preserved as it's returned
x_normed2, residual_out2 = ops.fused_add_rms_norm(
x_normed1, residual2, self.weight2, 1e-5
)
# Return both to prevent x_normed1 from being optimized away
return x_normed1, x_normed2, residual_out1, residual_out2
class ModelWithTritonAfterMaybeInplace(nn.Module):
"""
Model using maybe_inplace followed by a Triton kernel.
Test clone elimination can handle Triton in the graph
"""
def __init__(self, hidden_size=16):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
@triton.jit
def _triton_add_kernel(
x_ptr,
y_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(axis=0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
x = tl.load(x_ptr + offsets, mask=mask)
y = x + 0.1
tl.store(y_ptr + offsets, y, mask=mask)
def triton_add(x: torch.Tensor) -> torch.Tensor:
"""Simple Triton add kernel."""
y = torch.empty_like(x)
n_elements = x.numel()
grid = (triton.cdiv(n_elements, 256),)
_triton_add_kernel[grid](x, y, n_elements, BLOCK_SIZE=256)
return y
self.triton_add = triton_add
def forward(self, x: torch.Tensor, residual: torch.Tensor, residual2: torch.Tensor):
x_normed, residual_out = ops.fused_add_rms_norm.maybe_inplace(
x, residual, self.weight, 1e-5
)
x_processed = self.triton_add(x_normed)
# x_processed does not need to be cloned, residual2 does
x_normed2, residual_out2 = ops.fused_add_rms_norm(
x_processed, residual2, self.weight, 1e-5
)
return x_normed2, residual_out2
skipif_no_triton = pytest.mark.skipif(not HAS_TRITON, reason="Requires Triton")
@pytest.mark.skipif(
not current_platform.is_cuda_alike(),
reason="Only test on cuda and rocm platform",
)
@pytest.mark.parametrize(
"model_class,expected_functionalized,expected_donated,expected_clones",
[
# 2 inplace calls, all activations donated, all clones eliminated
(MaybeInplaceModel, 2, 3, 0),
# No inplace calls, no donations, 3 clones (one eliminated)
(FunctionalModel, 0, 0, 3),
# One inplace call, two donated activations, 2 clones
(MixedModel, 1, 2, 2),
# One inplace call, two donated, 1 clone remaining
pytest.param(ModelWithTritonAfterMaybeInplace, 1, 2, 1, marks=skipif_no_triton),
],
)
def test_inplace_functionalization(
default_vllm_config: VllmConfig,
model_class,
expected_functionalized: int,
expected_clones: int,
expected_donated: int,
):
"""Test inplace functionalization, lowering, and clone cleanup."""
torch.set_default_device(current_platform.device_type)
# Use vllm_c so inplace path is triggered
default_vllm_config.kernel_config.ir_op_priority.fused_add_rms_norm = [
"vllm_c",
"native",
]
# Create passes in order they run during compilation
functionalization_pass = VllmIRInplaceFunctionalizationPass(default_vllm_config)
lowering_pass = VllmIRLoweringPass(default_vllm_config)
donated_info_pass = StoreDonationInfoPass()
cleanup_pass = UnsafeCloneEliminationPass(default_vllm_config)
# Set up backend with pre-grad pass
backend = TestBackend(lowering_pass, donated_info_pass, cleanup_pass)
backend.inductor_config["pre_grad_custom_pass"] = functionalization_pass
model = model_class()
x = torch.randn(8, 16, dtype=torch.bfloat16)
residual1 = torch.randn(8, 16, dtype=torch.bfloat16)
residual2 = torch.randn(8, 16, dtype=torch.bfloat16)
with default_vllm_config.kernel_config.ir_op_priority.set_priority():
# Reference output without optimization
ref_output = model(x.clone(), residual1.clone(), residual2.clone())
# Compile with inplace optimization
compiled_model = torch.compile(model, backend=backend, fullgraph=True)
output = compiled_model(x.clone(), residual1.clone(), residual2.clone())
# Verify correctness (relaxed tolerance for bfloat16)
for i in range(len(ref_output)):
torch.testing.assert_close(output[i], ref_output[i], rtol=1e-2, atol=1e-2)
# Verify expected number of ops were functionalized
func_ops = functionalization_pass.functionalized_ops
assert len(func_ops) == int(bool(expected_functionalized))
if expected_functionalized > 0:
assert "fused_add_rms_norm" in func_ops
assert func_ops["fused_add_rms_norm"] == expected_functionalized
# Verify lowering happened (2 ops in all cases)
assert "fused_add_rms_norm" in lowering_pass.selected_impls
assert len(lowering_pass.selected_impls["fused_add_rms_norm"]) == 2
assert all(
provider == "vllm_c"
for node, provider in lowering_pass.selected_impls["fused_add_rms_norm"].items()
), lowering_pass.selected_impls
# Verify correct number of donated IDs
assert len(donated_info_pass.donated_input_ids_sets) == 1
assert len(donated_info_pass.donated_input_ids_sets[0]) == expected_donated
# Verify expected number of clones after cleanup
actual_clones = backend.op_count(torch.ops.aten.clone.default, before=False)
assert actual_clones == expected_clones, (
f"Expected {expected_clones} clones, got {actual_clones}:"
f"{backend.print_graphs()}"
)
@pytest.mark.skipif(
not current_platform.is_cuda_alike(),
reason="Only test on cuda and rocm platform",
)
def test_donated_buffer_context_propagation(default_vllm_config):
"""Test that donated_input_ids propagates correctly through pass_context."""
torch.set_default_device(current_platform.device_type)
# Create a custom backend that inspects pass_context in cleanup pass
functionalization_pass = VllmIRInplaceFunctionalizationPass(default_vllm_config)
lowering_pass = VllmIRLoweringPass(default_vllm_config)
donation_info_pass = StoreDonationInfoPass()
cleanup_pass = UnsafeCloneEliminationPass(default_vllm_config)
backend = TestBackend(lowering_pass, donation_info_pass, cleanup_pass)
backend.inductor_config["pre_grad_custom_pass"] = functionalization_pass
model = MaybeInplaceModel()
x = torch.randn(8, 16, dtype=torch.bfloat16)
residual1 = torch.randn(8, 16, dtype=torch.bfloat16)
residual2 = torch.randn(8, 16, dtype=torch.bfloat16)
compiled_model = torch.compile(model, backend=backend, fullgraph=True)
compiled_model(x.clone(), residual1.clone(), residual2.clone())
donated_ids_seen = donation_info_pass.donated_input_ids_sets
# Verify donated_input_ids was set and propagated
assert len(donated_ids_seen) == 1
# Should have donated inputs (exact indices depend on AOTAutograd)
assert len(donated_ids_seen[0]) == 3
# All donated ids should be valid non-negative integers
for idx in donated_ids_seen[0]:
assert isinstance(idx, int) and idx >= 0, f"Invalid donated index: {idx}"
@pytest.mark.skipif(
not current_platform.is_cuda_alike(),
reason="Only test on cuda and rocm platform",
)
def test_maybe_inplace_reuse_error(default_vllm_config):
"""Test that reusing a donated activation input raises ValueError."""
torch.set_default_device(current_platform.device_type)
class ReuseModel(nn.Module):
"""Model that incorrectly reuses a donated activation input."""
def __init__(self, hidden_size=16):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
def forward(self, x: torch.Tensor, residual: torch.Tensor):
# x is donated to maybe_inplace
x_normed, residual_out = ops.fused_add_rms_norm.maybe_inplace(
x, residual, self.weight, 1e-5
)
# ERROR: x is used again after being donated
return x_normed + x # This should raise ValueError
functionalization_pass = VllmIRInplaceFunctionalizationPass(default_vllm_config)
lowering_pass = VllmIRLoweringPass(default_vllm_config)
cleanup_pass = UnsafeCloneEliminationPass(default_vllm_config)
backend = TestBackend(lowering_pass, cleanup_pass)
backend.inductor_config["pre_grad_custom_pass"] = functionalization_pass
model = ReuseModel()
x = torch.randn(8, 16, dtype=torch.bfloat16)
residual = torch.randn(8, 16, dtype=torch.bfloat16)
# Compilation should raise BackendCompilerFailed wrapping ValueError
with pytest.raises(
torch._dynamo.exc.BackendCompilerFailed,
match="is used again after the node",
):
compiled_model = torch.compile(model, backend=backend, fullgraph=True)
compiled_model(x.clone(), residual.clone())
# Piecewise compilation tests with graph splitting
@torch.library.custom_op("vllm::test_split_marker", mutates_args=())
def test_split_marker(x: torch.Tensor) -> torch.Tensor:
"""Identity op that marks a split point for piecewise compilation."""
return x.clone()
@test_split_marker.register_fake
def _fake_split_marker(x: torch.Tensor) -> torch.Tensor:
return torch.empty_like(x)
class TransformerBlockWithSplits(nn.Module):
"""Transformer block with explicit split points for piecewise compilation."""
def __init__(self, hidden_size=32, intermediate_size=128):
super().__init__()
self.hidden_size = hidden_size
self.intermediate_size = intermediate_size
# Attention-like projection
self.attn_proj = nn.Linear(
hidden_size, hidden_size, bias=False, dtype=torch.bfloat16
)
# Post-attention norm
self.post_attn_norm = nn.Parameter(
torch.ones(hidden_size, dtype=torch.bfloat16)
)
# MLP
self.gate_proj = nn.Linear(
hidden_size, intermediate_size, bias=False, dtype=torch.bfloat16
)
self.up_proj = nn.Linear(
hidden_size, intermediate_size, bias=False, dtype=torch.bfloat16
)
self.down_proj = nn.Linear(
intermediate_size, hidden_size, bias=False, dtype=torch.bfloat16
)
# Post-MLP norm
self.post_mlp_norm = nn.Parameter(torch.ones(hidden_size, dtype=torch.bfloat16))
def forward(self, x: torch.Tensor):
# Attention block with residual
residual1 = x
attn_out = self.attn_proj(x)
# Fused add + norm (maybe_inplace: residual1 is donated)
normed1, residual1 = ops.fused_add_rms_norm.maybe_inplace(
attn_out, residual1, self.post_attn_norm, 1e-5
)
# Force a graph split here
normed1 = torch.ops.vllm.test_split_marker(normed1)
# MLP block
gate = self.gate_proj(normed1)
up = self.up_proj(normed1)
mlp_out = self.down_proj(gate * torch.nn.functional.silu(up))
# Fused add + norm (maybe_inplace: residual1 is donated)
normed2, residual2 = ops.fused_add_rms_norm.maybe_inplace(
mlp_out, residual1, self.post_mlp_norm, 1e-5
)
return normed2, residual2
def with_dyn_arg(fn: Callable, arg_index: int, dim_index: int):
def inner(*args):
torch._dynamo.mark_dynamic(args[arg_index], dim_index)
return fn(*args)
return inner
@pytest.mark.skipif(
not current_platform.is_cuda_alike(),
reason="Only test on cuda and rocm platform",
)
def test_piecewise_compilation_with_donated_buffers(monkeypatch, fresh_vllm_cache):
"""
Test piecewise compilation with donated buffers across graph splits.
Utilizes a custom splitting op. Uses fresh cache to avoid compilation caching.
"""
torch.set_default_device(current_platform.device_type)
# Disable compilation cache to avoid serialization issues
monkeypatch.setenv("VLLM_DISABLE_COMPILE_CACHE", "1")
from vllm.compilation.backends import VllmBackend
from vllm.config import CompilationConfig, VllmConfig
# Create config with custom splitting op
store_donation_info = StoreDonationInfoPass()
vllm_config = VllmConfig(
compilation_config=CompilationConfig(
custom_ops=["all"],
splitting_ops=["vllm::test_split_marker"],
inductor_compile_config={"post_grad_custom_post_pass": store_donation_info},
)
)
backend = VllmBackend(vllm_config)
model = TransformerBlockWithSplits()
x = torch.randn(8, 32, dtype=torch.bfloat16)
# Reference output
ref_output = with_dyn_arg(model, 0, 0)(x.clone())
# Compile with piecewise compilation (graph will split at split_marker)
compiled_model = torch.compile(model, backend=backend, fullgraph=False)
output = with_dyn_arg(compiled_model, 0, 0)(x.clone())
# Verify correctness (relaxed tolerance for bfloat16)
torch.testing.assert_close(output[0], ref_output[0], rtol=1e-2, atol=1e-2)
torch.testing.assert_close(output[1], ref_output[1], rtol=1e-2, atol=1e-2)
# Verify the model was split into multiple submodules
assert hasattr(backend, "split_gm"), "Backend should have split graph module"
# Should have at least 2 submodules (split by test_split_marker op)
submodules = list(backend.split_gm.named_children())
num_submodules = len(submodules)
assert num_submodules >= 2, (
f"Expected at least 2 submodules (split), got {num_submodules}"
)
# Check that donation info was propagated correctly
donated_inputs_sets = store_donation_info.donated_input_ids_sets
assert len(donated_inputs_sets) == 2
assert len(donated_inputs_sets[0]) == 1
assert len(donated_inputs_sets[1]) == 1
@@ -126,7 +126,7 @@ class TestFusedAddRMSNorm(torch.nn.Module):
if TEST_FP8 and do_fusion:
return [torch.ops._C.fused_add_rms_norm_static_fp8_quant.default]
else:
return [torch.ops._C.fused_add_rms_norm.default]
return []
def ops_not_in_model(self):
return []
@@ -59,7 +59,7 @@ class TestModel(torch.nn.Module):
def ops_in_model_before(self):
return [
rocm_aiter_ops.get_rmsnorm_fused_add_op(),
torch.ops.vllm_ir.fused_add_rms_norm,
torch.ops.aten.constant_pad_nd,
]
+4 -15
View File
@@ -17,7 +17,6 @@ from vllm.compilation.passes.fusion.rms_quant_fusion import (
FusedRMSQuantKey,
RMSNormQuantFusionPass,
)
from vllm.compilation.passes.fx_utils import find_op_nodes
from vllm.compilation.passes.utility.noop_elimination import NoOpEliminationPass
from vllm.compilation.passes.utility.post_cleanup import PostCleanupPass
from vllm.config import (
@@ -243,9 +242,10 @@ class TestModel(torch.nn.Module):
]
def ops_in_model_before_partial(self):
return [torch.ops.vllm_ir.rms_norm] + (
[RMS_ADD_OP] if self.enable_rms_norm_custom_op else [torch.ops.aten.rsqrt]
)
return [
torch.ops.vllm_ir.rms_norm,
torch.ops.vllm_ir.fused_add_rms_norm.default,
]
def _run_fusion_test(
@@ -383,17 +383,6 @@ def test_fusion_rmsnorm_quant(
model.ops_in_model_before_partial(), fully_replaced=False
)
# If RMSNorm custom op is disabled (native/torch impl used),
# there's a risk that the fused add doesn't get included in the
# replacement and only the rms part gets fused with quant.
# Hence, we check only 2 add nodes are left (final fused rmsnorm add).
if not enable_rms_norm_custom_op:
n_add_nodes = lambda g: sum(1 for _ in find_op_nodes(torch.ops.aten.add, g))
# rms_norm is IR, not included
# 6 = 3x2 (3xRMS_ADD, 2 each)
assert n_add_nodes(backend.graph_pre_pass) == 6
assert n_add_nodes(backend.graph_post_pass) == 2
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("hidden_size", [256])
+33 -8
View File
@@ -371,6 +371,8 @@ class HfRunner:
is_cross_encoder: bool = False,
skip_tokenizer_init: bool = False,
auto_cls: type[_BaseAutoModelClass] = AutoModelForCausalLM,
tokenizer_name: str | None = None,
processor: Any | None = None,
# Set this to avoid hanging issue
default_torch_num_threads: int | None = None,
) -> None:
@@ -391,6 +393,8 @@ class HfRunner:
is_cross_encoder=is_cross_encoder,
skip_tokenizer_init=skip_tokenizer_init,
auto_cls=auto_cls,
tokenizer_name=tokenizer_name,
processor=processor,
)
def _init(
@@ -405,6 +409,8 @@ class HfRunner:
is_cross_encoder: bool = False,
skip_tokenizer_init: bool = False,
auto_cls: type[_BaseAutoModelClass] = AutoModelForCausalLM,
tokenizer_name: str | None = None,
processor: Any | None = None,
) -> None:
model_name = maybe_model_redirect(model_name)
self.model_name = model_name
@@ -484,20 +490,27 @@ class HfRunner:
if not skip_tokenizer_init:
self.tokenizer: "PreTrainedTokenizer | PreTrainedTokenizerFast" = (
AutoTokenizer.from_pretrained(
model_name,
tokenizer_name or model_name,
trust_remote_code=trust_remote_code,
)
)
# don't put this import at the top level
# it will call torch.accelerator.device_count()
from transformers import AutoProcessor
if processor is not None:
self.processor = processor
else:
# don't put this import at the top level
# it will call torch.accelerator.device_count()
from transformers import AutoProcessor
self.processor = AutoProcessor.from_pretrained(
model_name,
trust_remote_code=trust_remote_code,
)
self.processor = AutoProcessor.from_pretrained(
model_name,
trust_remote_code=trust_remote_code,
)
if skip_tokenizer_init:
if self.processor is None:
raise ValueError(
"skip_tokenizer_init=True requires processor initialization."
)
self.tokenizer = self.processor.tokenizer
def get_inputs(
@@ -520,6 +533,12 @@ class HfRunner:
all_inputs: list[BatchFeature | BatchEncoding | dict[str, torch.Tensor]] = []
for i, prompt in enumerate(prompts):
if isinstance(prompt, str):
if self.processor is None:
raise RuntimeError(
"HfRunner.processor is not initialized. "
"Pass processor=... to HfRunner or set "
"hf_model.processor before generation."
)
# Create a copy to avoid modifying the original dict
processor_kwargs = (
tokenization_kwargs.copy()
@@ -617,6 +636,10 @@ class HfRunner:
use_cache=True,
**kwargs,
)
if self.processor is None:
raise RuntimeError(
"HfRunner.processor is not initialized; cannot decode output."
)
output_str = self.processor.batch_decode(
output_ids,
skip_special_tokens=True,
@@ -973,6 +996,8 @@ class VllmRunner:
req_sample_output_ids: list[list[int]] = []
req_sample_output_strs: list[str] = []
req_logprobs = []
if req_output.prompt_logprobs:
req_logprobs.extend(req_output.prompt_logprobs)
for sample in req_output.outputs:
output_str = sample.text
output_ids = list(sample.token_ids)
+306 -3
View File
@@ -10,10 +10,95 @@ Tests cover:
import math
import multiprocess as mp
import pytest
import torch
import torch.distributed as dist
from vllm.config.parallel import ParallelConfig
from vllm.utils.network_utils import get_open_port
from vllm.utils.system_utils import update_environment_variables
mp.set_start_method("spawn", force=True)
class _FakeCPGroup:
def __init__(self, world_size: int, device_group: dist.ProcessGroup):
self.world_size = world_size
self.device_group = device_group
def _dtype_from_name(dtype_name: str) -> torch.dtype:
return {
"float16": torch.float16,
"bfloat16": torch.bfloat16,
"float32": torch.float32,
}[dtype_name]
def _packed_a2a_reference(
cp_attn_out: torch.Tensor,
cp_attn_lse: torch.Tensor,
world_size: int,
h_per_rank: int,
is_lse_base_on_e: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
from vllm.v1.attention.ops.dcp_alltoall import _lse_weighted_combine
B, _H, D = cp_attn_out.shape
outputs = (
cp_attn_out.view(B, world_size, h_per_rank, D)
.permute(1, 0, 2, 3)
.contiguous()
.float()
)
lses = cp_attn_lse.view(B, world_size, h_per_rank).permute(1, 0, 2).contiguous()
return _lse_weighted_combine(
outputs,
lses,
return_lse=True,
is_lse_base_on_e=is_lse_base_on_e,
)
def _assert_packed_a2a_close(
actual: torch.Tensor,
expected: torch.Tensor,
dtype: torch.dtype,
) -> None:
if dtype == torch.float32:
torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-5)
else:
torch.testing.assert_close(
actual.float(), expected.float(), rtol=3e-2, atol=3e-2
)
def _distributed_run(fn, world_size: int, extra_env: dict[str, str]) -> None:
port = str(get_open_port())
processes: list[mp.Process] = []
for rank in range(world_size):
env = {
"RANK": str(rank),
"LOCAL_RANK": str(rank),
"WORLD_SIZE": str(world_size),
"LOCAL_WORLD_SIZE": str(world_size),
"MASTER_ADDR": "localhost",
"MASTER_PORT": port,
**extra_env,
}
process = mp.Process(target=fn, args=(env,))
processes.append(process)
process.start()
for process in processes:
process.join(timeout=120)
for process in processes:
if process.is_alive():
process.kill()
process.join()
assert process.exitcode == 0
class TestDCPCommBackendConfig:
@@ -38,14 +123,14 @@ class TestDCPCommBackendConfig:
"""A2A backend is valid when DCP > 1."""
config = ParallelConfig(
dcp_comm_backend="a2a",
tensor_parallel_size=8,
tensor_parallel_size=4,
decode_context_parallel_size=4,
)
assert config.dcp_comm_backend == "a2a"
def test_invalid_backend_rejected(self):
"""Invalid backend values are rejected."""
with pytest.raises(ValueError, match="must be one of"):
with pytest.raises(ValueError, match="must be one of|Input should be"):
ParallelConfig(
dcp_comm_backend="invalid",
)
@@ -134,7 +219,7 @@ class TestLSEWeightedCombine:
result = _lse_weighted_combine(outputs, lses)
assert result.shape == (B, H, D)
torch.testing.assert_close(result, outputs[1].squeeze(0), atol=1e-5, rtol=1e-5)
torch.testing.assert_close(result, outputs[1], atol=1e-5, rtol=1e-5)
def test_mathematically_correct(self):
"""Verify mathematical correctness of LSE combination."""
@@ -187,6 +272,224 @@ class TestLSEWeightedCombine:
assert global_lse.shape == (B, H)
assert abs(global_lse.item() - expected_global_lse) < 1e-5
def test_base2_return_lse(self):
"""Base-2 LSE mode returns log2-sum-exp2 global LSE."""
from vllm.v1.attention.ops.dcp_alltoall import _lse_weighted_combine
outputs = torch.tensor(
[
[[[1.0, 2.0]]],
[[[3.0, 4.0]]],
]
)
lses = torch.tensor(
[
[[1.0]],
[[2.0]],
]
)
result, global_lse = _lse_weighted_combine(
outputs,
lses,
return_lse=True,
is_lse_base_on_e=False,
)
expected_global_lse = math.log2(2**1 + 2**2)
w0 = 2**1 / (2**1 + 2**2)
w1 = 2**2 / (2**1 + 2**2)
expected = torch.tensor([[[w0 * 1.0 + w1 * 3.0, w0 * 2.0 + w1 * 4.0]]])
torch.testing.assert_close(result, expected, rtol=1e-5, atol=1e-5)
torch.testing.assert_close(
global_lse,
torch.tensor([[expected_global_lse]]),
rtol=1e-5,
atol=1e-5,
)
def test_lse_pack_dim(self):
"""Packed A2A stores one fp32 LSE in output-dtype lanes."""
from vllm.v1.attention.ops.dcp_alltoall import _dcp_a2a_lse_pack_dim
assert _dcp_a2a_lse_pack_dim(torch.bfloat16) == 2
assert _dcp_a2a_lse_pack_dim(torch.float16) == 2
assert _dcp_a2a_lse_pack_dim(torch.float32) == 1
class TestPackedA2AKernels:
@pytest.mark.skipif(
torch.accelerator.device_count() < 1, reason="CUDA is required."
)
@pytest.mark.parametrize("dtype_name", ["float16", "bfloat16", "float32"])
@pytest.mark.parametrize("return_lse", [False, True])
@pytest.mark.parametrize("is_lse_base_on_e", [False, True])
def test_pack_unpack_combine_matches_reference(
self,
dtype_name: str,
return_lse: bool,
is_lse_base_on_e: bool,
):
from vllm.v1.attention.ops.dcp_alltoall import (
_dcp_a2a_lse_pack_dim,
_dcp_a2a_pack_send,
_dcp_a2a_unpack_combine,
)
torch.manual_seed(0)
dtype = _dtype_from_name(dtype_name)
device = torch.device("cuda")
world_size, B, h_per_rank, D = 4, 7, 2, 32
H = world_size * h_per_rank
cp_attn_out = torch.randn(B, H, D, device=device, dtype=dtype)
cp_attn_lse = torch.randn(B, H, device=device, dtype=torch.float32)
lse_pack_dim = _dcp_a2a_lse_pack_dim(dtype)
send_buffer = torch.empty(
(world_size, B, h_per_rank, D + lse_pack_dim),
device=device,
dtype=dtype,
)
_dcp_a2a_pack_send(
cp_attn_out,
cp_attn_lse,
send_buffer,
world_size,
h_per_rank,
D,
lse_pack_dim,
)
actual = _dcp_a2a_unpack_combine(
send_buffer, D, lse_pack_dim, return_lse, is_lse_base_on_e
)
expected_out, expected_lse = _packed_a2a_reference(
cp_attn_out, cp_attn_lse, world_size, h_per_rank, is_lse_base_on_e
)
if return_lse:
actual_out, actual_lse = actual
_assert_packed_a2a_close(actual_out, expected_out, dtype)
torch.testing.assert_close(actual_lse, expected_lse, rtol=1e-4, atol=1e-4)
else:
_assert_packed_a2a_close(actual, expected_out, dtype)
def _distributed_packed_a2a_worker(env: dict[str, str]) -> None:
update_environment_variables(env)
local_rank = int(env["LOCAL_RANK"])
torch.accelerator.set_device_index(local_rank)
dist.init_process_group(backend="nccl")
use_workspace = env.get("USE_WORKSPACE") == "1"
if use_workspace:
from vllm.v1.worker.workspace import init_workspace_manager
init_workspace_manager(torch.device(f"cuda:{local_rank}"))
try:
from vllm.v1.attention.ops.dcp_alltoall import dcp_a2a_lse_reduce
dtype = _dtype_from_name(env["TEST_DTYPE"])
return_lse = env["RETURN_LSE"] == "1"
is_lse_base_on_e = env["LSE_BASE_E"] == "1"
rank = dist.get_rank()
world_size = dist.get_world_size()
B, h_per_rank, D = 5, 2, 32
H = world_size * h_per_rank
generator = torch.Generator(device=f"cuda:{local_rank}")
generator.manual_seed(1234 + rank)
cp_attn_out = torch.randn(
B,
H,
D,
device=f"cuda:{local_rank}",
dtype=dtype,
generator=generator,
)
cp_attn_lse = torch.randn(
B,
H,
device=f"cuda:{local_rank}",
dtype=torch.float32,
generator=generator,
)
actual = dcp_a2a_lse_reduce(
cp_attn_out,
cp_attn_lse,
_FakeCPGroup(world_size, dist.group.WORLD),
return_lse=return_lse,
is_lse_base_on_e=is_lse_base_on_e,
)
gathered_out = [torch.empty_like(cp_attn_out) for _ in range(world_size)]
gathered_lse = [torch.empty_like(cp_attn_lse) for _ in range(world_size)]
dist.all_gather(gathered_out, cp_attn_out)
dist.all_gather(gathered_lse, cp_attn_lse)
outputs = torch.stack(
[
t[:, rank * h_per_rank : (rank + 1) * h_per_rank, :]
for t in gathered_out
],
dim=0,
).float()
lses = torch.stack(
[t[:, rank * h_per_rank : (rank + 1) * h_per_rank] for t in gathered_lse],
dim=0,
)
from vllm.v1.attention.ops.dcp_alltoall import _lse_weighted_combine
expected_out, expected_lse = _lse_weighted_combine(
outputs,
lses,
return_lse=True,
is_lse_base_on_e=is_lse_base_on_e,
)
if return_lse:
actual_out, actual_lse = actual
_assert_packed_a2a_close(actual_out, expected_out, dtype)
torch.testing.assert_close(actual_lse, expected_lse, rtol=1e-4, atol=1e-4)
else:
_assert_packed_a2a_close(actual, expected_out, dtype)
finally:
if use_workspace:
from vllm.v1.worker.workspace import reset_workspace_manager
reset_workspace_manager()
dist.destroy_process_group()
@pytest.mark.skipif(
torch.accelerator.device_count() < 4, reason="Need at least 4 GPUs."
)
@pytest.mark.parametrize("dtype_name", ["float16", "bfloat16", "float32"])
def test_distributed_packed_a2a_matches_reference(dtype_name: str):
_distributed_run(
_distributed_packed_a2a_worker,
world_size=4,
extra_env={
"TEST_DTYPE": dtype_name,
"RETURN_LSE": "1",
"LSE_BASE_E": "1",
},
)
@pytest.mark.skipif(
torch.accelerator.device_count() < 4, reason="Need at least 4 GPUs."
)
def test_distributed_packed_a2a_with_workspace_matches_reference():
_distributed_run(
_distributed_packed_a2a_worker,
world_size=4,
extra_env={
"TEST_DTYPE": "bfloat16",
"RETURN_LSE": "1",
"LSE_BASE_E": "1",
"USE_WORKSPACE": "1",
},
)
if __name__ == "__main__":
pytest.main([__file__, "-v"])
-3
View File
@@ -333,8 +333,6 @@ def test_attention_config():
"true",
"--attention-config.flash_attn_max_num_splits_for_cuda_graph",
"16",
"--attention-config.use_cudnn_prefill",
"true",
"--attention-config.use_trtllm_ragged_deepseek_prefill",
"true",
"--attention-config.use_trtllm_attention",
@@ -352,7 +350,6 @@ def test_attention_config():
assert engine_args.attention_config.flash_attn_version == 3
assert engine_args.attention_config.use_prefill_decode_attention is True
assert engine_args.attention_config.flash_attn_max_num_splits_for_cuda_graph == 16
assert engine_args.attention_config.use_cudnn_prefill is True
assert engine_args.attention_config.use_trtllm_ragged_deepseek_prefill is True
assert engine_args.attention_config.use_trtllm_attention is True
assert engine_args.attention_config.disable_flashinfer_prefill is True
@@ -11,7 +11,9 @@ from vllm import LLM, SamplingParams
def _make_mock_llm() -> LLM:
llm = object.__new__(LLM)
llm.model_config = SimpleNamespace(runner_type="generate")
llm.model_config = SimpleNamespace(
runner_type="generate", enable_prompt_embeds=False
)
return llm
@@ -0,0 +1,190 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""E2E test for mixing `prompt_embeds` with `audio_embeds` in a single
Chat Completions request."""
import json
import openai
import pytest
import pytest_asyncio
import safetensors
import torch
import torch.nn as nn
from huggingface_hub import hf_hub_download
from transformers import AutoConfig, AutoTokenizer
from tests.utils import RemoteOpenAIServer
from vllm.utils.serial_utils import tensor2base64
QWEN2AUDIO_MODEL = "Qwen/Qwen2-Audio-7B-Instruct"
# Use the model's native dtype to avoid an implicit cast inside
# `safe_load_prompt_embeds` (mismatched floating-point dtypes are cast to the
# model's dtype automatically, matching here just skips the conversion).
QWEN2AUDIO_DTYPE = torch.bfloat16
@pytest.fixture(scope="module")
def qwen2audio_server_args() -> list[str]:
return [
"--dtype",
"bfloat16",
"--max-model-len",
"2048",
"--max-num-seqs",
"4",
"--enforce-eager",
"--trust-remote-code",
"--gpu-memory-utilization",
"0.85",
"--limit-mm-per-prompt",
json.dumps({"audio": 1}),
"--enable-prompt-embeds",
"--enable-mm-embeds",
]
@pytest.fixture(scope="module")
def qwen2audio_server(qwen2audio_server_args):
with RemoteOpenAIServer(
QWEN2AUDIO_MODEL,
qwen2audio_server_args,
max_wait_seconds=600,
) as remote_server:
yield remote_server
@pytest_asyncio.fixture
async def qwen2audio_client(qwen2audio_server):
async with qwen2audio_server.get_async_client() as async_client:
yield async_client
@pytest.fixture(scope="module")
def qwen2audio_hidden_size() -> int:
config = AutoConfig.from_pretrained(QWEN2AUDIO_MODEL, trust_remote_code=True)
return config.text_config.hidden_size
@pytest.fixture(scope="module")
def qwen2audio_prompt_embeds_b64(qwen2audio_hidden_size: int) -> str:
tensor = torch.randn(4, qwen2audio_hidden_size, dtype=QWEN2AUDIO_DTYPE)
return tensor2base64(tensor)
@pytest.fixture(scope="module")
def qwen2audio_audio_embeds_b64(qwen2audio_hidden_size: int) -> str:
# Shape matches the `audio_embeds` unit-test fixture.
torch.manual_seed(0)
tensor = torch.randn(1, 128, qwen2audio_hidden_size, dtype=QWEN2AUDIO_DTYPE)
return tensor2base64(tensor)
@pytest.mark.asyncio
async def test_prompt_embeds_plus_audio_embeds(
qwen2audio_client: openai.AsyncOpenAI,
qwen2audio_prompt_embeds_b64: str,
qwen2audio_audio_embeds_b64: str,
):
"""Single user message carrying both prompt_embeds and audio_embeds parts."""
chat = await qwen2audio_client.chat.completions.create(
model=QWEN2AUDIO_MODEL,
max_tokens=5,
temperature=0.0,
messages=[
{
"role": "user",
"content": [
{
"type": "prompt_embeds",
"data": qwen2audio_prompt_embeds_b64,
},
{
"type": "audio_embeds",
"audio_embeds": qwen2audio_audio_embeds_b64,
},
{"type": "text", "text": "Continue."},
],
}
],
)
assert chat.choices[0].message.content is not None
assert len(chat.choices[0].message.content) > 0
@pytest.fixture(scope="module")
def qwen2audio_aligned_content_and_embeds_b64() -> tuple[str, str]:
"""Return `(content, base64_embeds)` where the embeddings are the model's
embedding of `content` tokenized WITHOUT special tokens.
Loads only the `embed_tokens` shard from disk on CPU (~1.1 GB of host
RAM) instead of the full 7B model on GPU.
"""
content = "Describe this audio."
tokenizer = AutoTokenizer.from_pretrained(QWEN2AUDIO_MODEL, trust_remote_code=True)
index_path = hf_hub_download(QWEN2AUDIO_MODEL, "model.safetensors.index.json")
with open(index_path) as f:
weight_map = json.load(f)["weight_map"]
embed_key = next(k for k in weight_map if k.endswith("embed_tokens.weight"))
shard_path = hf_hub_download(QWEN2AUDIO_MODEL, weight_map[embed_key])
with safetensors.safe_open(shard_path, framework="pt", device="cpu") as f:
embed_weight = f.get_tensor(embed_key)
embed_layer = nn.Embedding.from_pretrained(embed_weight.to(QWEN2AUDIO_DTYPE))
ids = tokenizer(content, add_special_tokens=False, return_tensors="pt").input_ids
embeds = embed_layer(ids).squeeze(0)
return content, tensor2base64(embeds)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"audio_first",
[True, False],
ids=["audio_embeds-then-text", "text-then-audio_embeds"],
)
async def test_text_content_and_prompt_embeds_match_with_audio_embeds(
qwen2audio_client: openai.AsyncOpenAI,
qwen2audio_audio_embeds_b64: str,
qwen2audio_aligned_content_and_embeds_b64: tuple[str, str],
audio_first: bool,
):
"""Same content as text vs `prompt_embeds` should yield identical Chat
Completions output when mixed with `audio_embeds` in the same message.
"""
content, encoded_text_embeds = qwen2audio_aligned_content_and_embeds_b64
audio_part = {
"type": "audio_embeds",
"audio_embeds": qwen2audio_audio_embeds_b64,
}
text_part = {"type": "text", "text": content}
embeds_part = {"type": "prompt_embeds", "data": encoded_text_embeds}
if audio_first:
text_content = [audio_part, text_part]
embeds_content = [audio_part, embeds_part]
else:
text_content = [text_part, audio_part]
embeds_content = [embeds_part, audio_part]
text_resp = await qwen2audio_client.chat.completions.create(
model=QWEN2AUDIO_MODEL,
max_tokens=10,
temperature=0.0,
messages=[{"role": "user", "content": text_content}],
)
embeds_resp = await qwen2audio_client.chat.completions.create(
model=QWEN2AUDIO_MODEL,
max_tokens=10,
temperature=0.0,
messages=[{"role": "user", "content": embeds_content}],
)
text_out = text_resp.choices[0].message.content
embeds_out = embeds_resp.choices[0].message.content
assert text_out is not None and len(text_out) > 0
assert embeds_out is not None and len(embeds_out) > 0
assert text_out == embeds_out
@@ -0,0 +1,212 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""E2E tests for mixing `prompt_embeds` with image content parts in a single
Chat Completions request.
"""
import json
import openai
import pytest
import pytest_asyncio
import safetensors
import torch
import torch.nn as nn
from huggingface_hub import hf_hub_download
from transformers import AutoTokenizer
from tests.utils import RemoteOpenAIServer
from vllm.assets.image import ImageAsset
from vllm.multimodal.utils import encode_image_url
from vllm.utils.serial_utils import tensor2base64
MODEL_NAME = "Qwen/Qwen2-VL-2B-Instruct"
# Use the model's native dtype to skip the implicit cast inside
# `safe_load_prompt_embeds` (mismatched floating-point dtypes are cast to the
# model's dtype automatically).
MODEL_DTYPE = torch.bfloat16
@pytest.fixture(scope="module")
def server_args() -> list[str]:
return [
"--dtype",
"bfloat16",
"--max-model-len",
"2048",
"--max-num-seqs",
"4",
"--enforce-eager",
"--gpu-memory-utilization",
"0.4",
"--limit-mm-per-prompt",
json.dumps({"image": 1}),
"--enable-prompt-embeds",
"--enable-mm-embeds",
]
@pytest.fixture(scope="module")
def server(server_args):
with RemoteOpenAIServer(
MODEL_NAME,
server_args,
max_wait_seconds=600,
) as remote_server:
yield remote_server
@pytest_asyncio.fixture
async def client(server):
async with server.get_async_client() as async_client:
yield async_client
@pytest.fixture(scope="module")
def image_url() -> str:
"""Stable real image as a data URL, kept identical across both the
text and prompt_embeds requests so any output difference must come from
how the text content is delivered."""
return encode_image_url(ImageAsset("stop_sign").pil_image)
@pytest.fixture(scope="module")
def aligned_content_and_embeds_b64() -> tuple[str, str]:
"""`(content, base64_embeds)` where the embeddings are the model's
embedding of `content` tokenized WITHOUT special tokens.
Loads only the `embed_tokens` shard from disk on CPU instead of the full
model on GPU, so the fixture has zero VRAM footprint and won't contend
with the running vLLM server.
"""
content = "Describe this image."
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, trust_remote_code=True)
index_path = hf_hub_download(MODEL_NAME, "model.safetensors.index.json")
with open(index_path) as f:
weight_map = json.load(f)["weight_map"]
embed_key = next(k for k in weight_map if k.endswith("embed_tokens.weight"))
shard_path = hf_hub_download(MODEL_NAME, weight_map[embed_key])
with safetensors.safe_open(shard_path, framework="pt", device="cpu") as f:
embed_weight = f.get_tensor(embed_key)
embed_layer = nn.Embedding.from_pretrained(embed_weight.to(MODEL_DTYPE))
ids = tokenizer(content, add_special_tokens=False, return_tensors="pt").input_ids
embeds = embed_layer(ids).squeeze(0)
return content, tensor2base64(embeds)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"image_first",
[True, False],
ids=["image_url-then-text", "text-then-image_url"],
)
async def test_text_content_and_prompt_embeds_match_with_image_url(
client: openai.AsyncOpenAI,
image_url: str,
aligned_content_and_embeds_b64: tuple[str, str],
image_first: bool,
):
"""Same content as text vs `prompt_embeds` should yield identical Chat
Completions output when mixed with an `image_url` part in the same
message under greedy decoding.
"""
content, encoded_text_embeds = aligned_content_and_embeds_b64
image_part = {"type": "image_url", "image_url": {"url": image_url}}
text_part = {"type": "text", "text": content}
embeds_part = {"type": "prompt_embeds", "data": encoded_text_embeds}
if image_first:
text_content = [image_part, text_part]
embeds_content = [image_part, embeds_part]
else:
text_content = [text_part, image_part]
embeds_content = [embeds_part, image_part]
text_resp = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=10,
temperature=0.0,
messages=[{"role": "user", "content": text_content}],
)
embeds_resp = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=10,
temperature=0.0,
messages=[{"role": "user", "content": embeds_content}],
)
text_out = text_resp.choices[0].message.content
embeds_out = embeds_resp.choices[0].message.content
assert text_out is not None and len(text_out) > 0
assert embeds_out is not None and len(embeds_out) > 0
assert text_out == embeds_out
@pytest.fixture(scope="module")
def image_embeds_b64() -> dict[str, str]:
"""Synthetic but stable `image_embeds` for Qwen2-VL."""
grid = (1, 4, 4)
spatial_merge_size = 2
num_patches = (grid[1] // spatial_merge_size) * (grid[2] // spatial_merge_size)
text_hidden_size = 1536 # Qwen2-VL-2B
torch.manual_seed(0)
return {
"image_embeds": tensor2base64(
torch.randn(num_patches, text_hidden_size, dtype=MODEL_DTYPE)
),
"image_grid_thw": tensor2base64(torch.tensor(grid)),
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"image_first",
[True, False],
ids=["image_embeds-then-text", "text-then-image_embeds"],
)
async def test_text_content_and_prompt_embeds_match_with_image_embeds(
client: openai.AsyncOpenAI,
image_embeds_b64: dict[str, str],
aligned_content_and_embeds_b64: tuple[str, str],
image_first: bool,
):
"""Same content as text vs `prompt_embeds` should yield identical Chat
Completions output when mixed with a precomputed `image_embeds` part in
the same message under greedy decoding.
"""
content, encoded_text_embeds = aligned_content_and_embeds_b64
image_part = {"type": "image_embeds", "image_embeds": image_embeds_b64}
text_part = {"type": "text", "text": content}
embeds_part = {"type": "prompt_embeds", "data": encoded_text_embeds}
if image_first:
text_content = [image_part, text_part]
embeds_content = [image_part, embeds_part]
else:
text_content = [text_part, image_part]
embeds_content = [embeds_part, image_part]
text_resp = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=10,
temperature=0.0,
messages=[{"role": "user", "content": text_content}],
)
embeds_resp = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=10,
temperature=0.0,
messages=[{"role": "user", "content": embeds_content}],
)
text_out = text_resp.choices[0].message.content
embeds_out = embeds_resp.choices[0].message.content
assert text_out is not None and len(text_out) > 0
assert embeds_out is not None and len(embeds_out) > 0
assert text_out == embeds_out
@@ -0,0 +1,293 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""E2E tests for `prompt_embeds` content parts in the Chat Completions API."""
import asyncio
import io
import openai
import pybase64 as base64
import pytest
import pytest_asyncio
import torch
from openai import BadRequestError
from tests.utils import VLLM_PATH, RemoteOpenAIServer
MODEL_NAME = "facebook/opt-125m"
CHAT_TEMPLATE = VLLM_PATH / "examples/template_chatml.jinja"
# Matches `--dtype` in `server_args` to avoid an implicit cast in
# `safe_load_prompt_embeds` (mismatched floating-point dtypes are cast to the
# model's dtype automatically, we match here just to skip the conversion).
SERVER_DTYPE: torch.dtype = torch.bfloat16
@pytest.fixture(scope="module")
def server_args() -> list[str]:
return [
"--dtype",
"bfloat16",
"--max-model-len",
"2048",
"--max-num-seqs",
"128",
"--enforce-eager",
"--chat-template",
str(CHAT_TEMPLATE),
# Prompt Embeds server args
"--enable-prompt-embeds",
]
@pytest.fixture(scope="module")
def server(server_args):
with RemoteOpenAIServer(MODEL_NAME, server_args) as remote_server:
yield remote_server
@pytest_asyncio.fixture
async def client(server):
async with server.get_async_client() as async_client:
yield async_client
def _encode_embeds(embeds: torch.Tensor) -> str:
buf = io.BytesIO()
torch.save(embeds, buf)
return base64.b64encode(buf.getvalue()).decode("utf-8")
@pytest.fixture(scope="module")
def prompt_embeds_b64(hf_runner) -> list[str]:
"""Pre-compute embeddings for two short prompts and return as base64."""
prompts = ["Hello, my name is", "What is an LLM?"]
with hf_runner(MODEL_NAME) as hf_model:
embeddings = hf_model.get_prompt_embeddings(prompts)
# Cast to the server's dtype so `safe_load_prompt_embeds` doesn't need to
# convert on its own, the function accepts any floating-point dtype and
# will cast to the model's dtype, but matching up front skips the work.
return [_encode_embeds(e.to(SERVER_DTYPE)) for e in embeddings]
@pytest.mark.asyncio
async def test_single_prompt_embeds_part(
client: openai.AsyncOpenAI,
prompt_embeds_b64: list[str],
):
"""A user message with one prompt_embeds part + text."""
b64 = prompt_embeds_b64[0]
chat = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=5,
temperature=0.0,
messages=[
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": b64},
{"type": "text", "text": "Continue:"},
],
}
],
)
assert chat.choices[0].message.content is not None
assert len(chat.choices[0].message.content) > 0
@pytest.mark.asyncio
async def test_multiple_prompt_embeds_parts(
client: openai.AsyncOpenAI,
prompt_embeds_b64: list[str],
):
"""Multiple prompt_embeds parts in a single message."""
b64_a, b64_b = prompt_embeds_b64
chat = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=5,
temperature=0.0,
messages=[
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": b64_a},
{"type": "text", "text": " and "},
{"type": "prompt_embeds", "data": b64_b},
],
}
],
)
assert chat.choices[0].message.content is not None
assert len(chat.choices[0].message.content) > 0
@pytest.mark.asyncio
async def test_multi_message_conversation(
client: openai.AsyncOpenAI,
prompt_embeds_b64: list[str],
):
"""prompt_embeds in both system and user messages."""
b64_sys, b64_usr = prompt_embeds_b64
chat = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=5,
temperature=0.0,
messages=[
{
"role": "system",
"content": [
{"type": "text", "text": "You are helpful."},
{"type": "prompt_embeds", "data": b64_sys},
],
},
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": b64_usr},
{"type": "text", "text": "Summarize."},
],
},
],
)
assert chat.choices[0].message.content is not None
assert len(chat.choices[0].message.content) > 0
@pytest.mark.asyncio
async def test_streaming(
client: openai.AsyncOpenAI,
prompt_embeds_b64: list[str],
):
"""Streaming chat completion with prompt_embeds."""
b64 = prompt_embeds_b64[0]
# Non-streaming baseline.
baseline = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=5,
temperature=0.0,
messages=[
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": b64},
{"type": "text", "text": "Continue:"},
],
}
],
)
expected = baseline.choices[0].message.content
# Streaming.
stream = await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=5,
temperature=0.0,
stream=True,
messages=[
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": b64},
{"type": "text", "text": "Continue:"},
],
}
],
)
chunks: list[str] = []
async for chunk in stream:
delta = chunk.choices[0].delta.content
if delta:
chunks.append(delta)
assert "".join(chunks) == expected
@pytest.fixture(scope="module")
def aligned_content_and_embeds_b64(hf_runner) -> tuple[str, str]:
"""Return `(content, base64_embeds)` where the embeddings are the model's
embedding of `content` tokenized WITHOUT special tokens.
"""
content = "Hello, my name is"
with hf_runner(MODEL_NAME) as hf_model:
ids = hf_model.tokenizer(
content, add_special_tokens=False, return_tensors="pt"
).input_ids
ids = hf_model.wrap_device({"input_ids": ids})["input_ids"]
embed_layer = hf_model.model.get_input_embeddings()
embeds = embed_layer(ids).squeeze(0).to(SERVER_DTYPE).cpu()
return content, _encode_embeds(embeds)
@pytest.mark.asyncio
async def test_text_content_and_prompt_embeds_match(
client: openai.AsyncOpenAI,
aligned_content_and_embeds_b64: tuple[str, str],
):
"""Equal content in text and `prompt_embeds` should yield identical
Chat Completions output under greedy decoding.
"""
content, encoded_embeds = aligned_content_and_embeds_b64
text_resp, embeds_resp = await asyncio.gather(
client.chat.completions.create(
model=MODEL_NAME,
max_tokens=10,
temperature=0.0,
messages=[{"role": "user", "content": content}],
),
client.chat.completions.create(
model=MODEL_NAME,
max_tokens=10,
temperature=0.0,
messages=[
{
"role": "user",
"content": [{"type": "prompt_embeds", "data": encoded_embeds}],
}
],
),
)
text_out = text_resp.choices[0].message.content
embeds_out = embeds_resp.choices[0].message.content
assert text_out is not None and len(text_out) > 0
assert embeds_out is not None and len(embeds_out) > 0
assert text_out == embeds_out
@pytest.mark.asyncio
async def test_missing_data_field(
client: openai.AsyncOpenAI,
):
"""A prompt_embeds part without `data` should return a clear error."""
with pytest.raises(BadRequestError):
await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=5,
messages=[
{
"role": "user",
"content": [{"type": "prompt_embeds"}],
}
],
)
@pytest.mark.asyncio
async def test_invalid_base64(
client: openai.AsyncOpenAI,
):
"""Invalid base64 in the `data` field should return a clear error."""
with pytest.raises(BadRequestError):
await client.chat.completions.create(
model=MODEL_NAME,
max_tokens=5,
messages=[
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": "not_valid_base64!!"},
],
}
],
)
@@ -538,6 +538,7 @@ class MockModelConfig:
is_encoder_decoder: bool = False
is_multimodal_model: bool = False
renderer_num_workers: int = 1
enable_prompt_embeds: bool = False
def get_diff_sampling_param(self):
return self.diff_sampling_param or {}
@@ -62,6 +62,8 @@ def test_load_prompt_embeds(
):
model_config = Mock(spec=ModelConfig)
model_config.enable_prompt_embeds = True
model_config.get_hidden_size.return_value = hidden_size
model_config.dtype = dtype
# construct arbitrary tensors of various dtypes, layouts, and sizes.
# We need to check against different layouts to make sure that if a user
@@ -26,6 +26,10 @@ from vllm.tokenizers import get_tokenizer
from ....models.registry import HF_EXAMPLE_MODELS
from ....utils import RemoteOpenAIServer
# Tuned to prevent OOM on 18GB GPUs in transcription correctness tests.
MAX_SEQS_FOR_TRANSCRIPTION_TEST = 8
GPU_UTIL_FOR_TRANSCRIPTION_TEST = 0.5
def to_bytes(y, sr):
buffer = io.BytesIO()
@@ -184,6 +188,8 @@ def test_wer_correctness(
server_args = [
"--enforce-eager",
f"--tokenizer_mode={model_info.tokenizer_mode}",
f"--max_num_seqs={MAX_SEQS_FOR_TRANSCRIPTION_TEST}",
f"--gpu_memory_utilization={GPU_UTIL_FOR_TRANSCRIPTION_TEST}",
]
if model_info.trust_remote_code:
server_args.append("--trust-remote-code")
@@ -38,6 +38,15 @@ BACKEND_TOL: dict[str, float] = {
"FLEX_ATTENTION": 0.045, # gfx950:~3.25%, gfx942:~1.10%
}
# ROCm 7.2/gfx950 shows small absolute drift on the low text-vs-text
# probability even though larger scores remain well inside the relative
# tolerance. Keep the relative tolerances tight and add only a small floor.
BACKEND_ABS_TOL: dict[str, float] = {
"default": 0.0,
"ROCM_AITER_FA": 0.005,
"FLEX_ATTENTION": 0.006,
}
# ROCm: disable skinny GEMM to avoid non-deterministic results from
# atomic reductions in wvSplitKrc kernel.
# See: https://github.com/vllm-project/vllm/pull/33493#issuecomment-3906083975
@@ -57,18 +66,23 @@ def get_tol(backend: str) -> float:
return BACKEND_TOL.get(backend, BACKEND_TOL["default"])
def get_abs_tol(backend: str) -> float:
return BACKEND_ABS_TOL.get(backend, BACKEND_ABS_TOL["default"])
def assert_score(actual: float, expected: float, backend: str, label: str):
tol = get_tol(backend)
abs_tol = get_abs_tol(backend)
diff = abs(actual - expected)
rel_diff = diff / abs(expected) if expected != 0 else diff
print(
f"[{backend}] {label}: actual={actual:.6f} expected={expected:.6f} "
f"diff={diff:.6f} rel_diff={rel_diff:.4f} tol={tol}"
f"diff={diff:.6f} rel_diff={rel_diff:.4f} tol={tol} abs_tol={abs_tol}"
)
assert actual == pytest.approx(expected, rel=tol), (
assert actual == pytest.approx(expected, rel=tol, abs=abs_tol), (
f"[{backend}] {label}: score mismatch — "
f"actual={actual:.6f}, expected={expected:.6f}, "
f"rel_diff={rel_diff:.4f}, tol={tol}"
f"rel_diff={rel_diff:.4f}, tol={tol}, abs_tol={abs_tol}"
)
+44
View File
@@ -0,0 +1,44 @@
# MRCR Long-Context Accuracy Evaluation
Smoke test for long-context behavior using OpenAI's public [`openai/mrcr`](https://huggingface.co/datasets/openai/mrcr) dataset. The model sees a long chat with several near-duplicate "needles" and must reproduce a specific earlier assistant turn verbatim, prepended with a random anti-guessing string.
**Scoring:** if the response doesn't start with `random_string_to_prepend`, score is 0; otherwise the prefix is stripped and the mean `SequenceMatcher.ratio()` against the reference answer is reported.
## Usage
```bash
# Pytest (spawns the server)
pytest -s -v tests/evals/mrcr/test_mrcr_correctness.py \
--config-list-file=configs/models-small.txt
# Standalone (server already running; model and context auto-discovered)
vllm serve Qwen/Qwen3-0.6B --reasoning-parser qwen3 --port 8000
python tests/evals/mrcr/mrcr_eval.py --port 8000
```
## Configuration
```yaml
model_name: "Qwen/Qwen3-0.6B"
# Per-needle thresholds catch bucket-specific regressions (sliding window,
# chunked prefill, prefix cache) that an aggregate can hide. A scalar
# (e.g. `match_ratio_threshold: 0.20`) is also accepted and checked against
# the mean match ratio.
match_ratio_threshold:
2: 0.30
4: 0.15
8: 0.10
num_samples: 30
needles: [2, 4, 8]
# max_prompt_tokens: 32768 # Optional; defaults to server max_model_len - max_tokens - 256
max_tokens: 2048
concurrency: 8
server_args: "--max-model-len 32768 --reasoning-parser qwen3"
```
## Notes
- Samples stream from three parquet shards (`{N}needle/{N}needle_0.parquet`); only the first few row groups are fetched, not the full 1.4 GB repo.
- `max_prompt_tokens` defaults to `max_model_len - max_tokens - 256`, i.e. fills whatever context the server advertises. Set `--max-model-len` on the server to control the smoke-test context length; override `--max-prompt-tokens` on the client to cap below that.
- Sample length is pre-filtered by `n_chars × 4 ≤ max_prompt_tokens`, then verified via the server's `/tokenize` endpoint under the actual chat template.
- Reasoning models: start the server with `--reasoning-parser <name>` (e.g. `qwen3`, `deepseek_r1`) so `<think>` goes to `message.reasoning_content` and doesn't contaminate the scored answer.
+2
View File
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
+7
View File
@@ -0,0 +1,7 @@
model_name: "Qwen/Qwen3.5-4B"
needles: [2, 4, 8]
match_ratio_threshold:
2: 0.99
4: 0.84
8: 0.76
server_args: "--max-model-len 128K --reasoning-parser qwen3"
@@ -0,0 +1 @@
Qwen3.5-4B.yaml
+54
View File
@@ -0,0 +1,54 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from pathlib import Path
def pytest_addoption(parser):
"""Add custom command line options."""
parser.addoption(
"--config-list-file",
default="configs/models-small.txt",
help="File containing list of config files to test",
)
def pytest_generate_tests(metafunc):
"""Generate test parameters from config files."""
if "config_filename" in metafunc.fixturenames:
config_list_file = metafunc.config.getoption("--config-list-file")
config_list_path = Path(config_list_file)
if not config_list_path.is_absolute():
test_dir_path = Path(__file__).parent / config_list_file
if test_dir_path.exists():
config_list_path = test_dir_path
else:
config_list_path = Path.cwd() / config_list_file
print(f"Looking for config list at: {config_list_path}")
config_files = []
if config_list_path.exists():
config_dir = config_list_path.parent
with open(config_list_path) as f:
for line in f:
line = line.strip()
if line and not line.startswith("#"):
config_path = config_dir / line
if config_path.exists():
config_files.append(config_path)
print(f" ✓ Found: {config_path}")
else:
print(f" ✗ Missing: {config_path}")
else:
print(f"Config list file not found: {config_list_path}")
if config_files:
metafunc.parametrize(
"config_filename",
config_files,
ids=[config_file.stem for config_file in config_files],
)
else:
print("No config files found, test will be skipped")
+333
View File
@@ -0,0 +1,333 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""MRCR long-context evaluation for vLLM's OpenAI-compatible server.
Streams samples from `openai/mrcr` on HuggingFace, sends chat completions to
the server, and scores each response with a prefix-gated SequenceMatcher ratio
against the reference answer.
"""
import argparse
import asyncio
import json
import time
from difflib import SequenceMatcher
import aiohttp
import numpy as np
import requests
from tqdm.asyncio import tqdm
DATASET_REPO = "openai/mrcr"
NEEDLE_SHARDS = {
2: "2needle/2needle_0.parquet",
4: "4needle/4needle_0.parquet",
8: "8needle/8needle_0.parquet",
}
# Reserve headroom for chat-template tokens on top of the messages.
PROMPT_SAFETY_BUFFER = 256
# Pre-filter heuristic before the authoritative /tokenize check.
CHARS_PER_TOKEN = 4
# Skip chain-of-thought on reasoning models; ignored by non-reasoning templates.
DEFAULT_EXTRA_BODY: dict = {"chat_template_kwargs": {"enable_thinking": False}}
def discover_server_model(base_url: str) -> tuple[str, int | None]:
"""Return (model_id, max_model_len) from /v1/models."""
resp = requests.get(f"{base_url}/v1/models", timeout=30)
resp.raise_for_status()
data = resp.json().get("data", [])
if not data:
raise RuntimeError(f"No models advertised at {base_url}/v1/models")
entry = data[0]
return entry["id"], entry.get("max_model_len")
def count_chat_tokens(base_url: str, model: str, messages: list[dict]) -> int:
"""Return the chat-template-rendered token count via /tokenize."""
resp = requests.post(
f"{base_url}/tokenize",
json={"model": model, "messages": messages, "add_generation_prompt": True},
timeout=120,
)
resp.raise_for_status()
return int(resp.json()["count"])
def _load_mrcr_samples(
needles: list[int],
max_prompt_tokens: int,
num_samples: int,
seed: int,
base_url: str,
model_name: str,
) -> list[dict]:
"""Stream MRCR samples balanced across needle buckets, token-verified."""
try:
from datasets import load_dataset
except ImportError as e:
raise ImportError(
"MRCR eval requires `datasets`. Install with: uv pip install datasets"
) from e
max_chars = max_prompt_tokens * CHARS_PER_TOKEN
per_bucket = num_samples // len(needles)
leftover = num_samples - per_bucket * len(needles)
samples: list[dict] = []
for idx, n in enumerate(needles):
if n not in NEEDLE_SHARDS:
raise ValueError(f"Unsupported needle count {n}")
target = per_bucket + (1 if idx < leftover else 0)
if target == 0:
continue
ds = load_dataset(
DATASET_REPO,
data_files=NEEDLE_SHARDS[n],
split="train",
streaming=True,
).shuffle(seed=seed + n, buffer_size=16)
taken = 0
for row in ds:
if int(row.get("n_chars", 0)) > max_chars:
continue
prompt = row["prompt"]
messages = json.loads(prompt) if isinstance(prompt, str) else list(prompt)
n_tokens = count_chat_tokens(base_url, model_name, messages)
if n_tokens > max_prompt_tokens:
continue
samples.append(
{
"messages": messages,
"answer": row["answer"],
"random_string_to_prepend": row["random_string_to_prepend"],
"n_needles": int(row["n_needles"]),
"n_tokens": n_tokens,
}
)
taken += 1
if taken >= target:
break
if taken < target:
print(f"Warning: only {taken}/{target} samples for n_needles={n}")
if not samples:
raise RuntimeError("No MRCR samples fit; loosen max_prompt_tokens.")
return samples
def score_mrcr(response: str, answer: str, random_prefix: str) -> float:
"""Prefix-gated SequenceMatcher ratio; 0 if the prefix is missing."""
if not response.startswith(random_prefix):
return 0.0
stripped = response[len(random_prefix) :]
return SequenceMatcher(a=answer, b=stripped, autojunk=False).ratio()
async def _call_chat(
session: aiohttp.ClientSession,
url: str,
model: str,
messages: list[dict],
max_tokens: int,
temperature: float,
seed: int | None,
extra_body: dict,
) -> tuple[str, int]:
data = {
"model": model,
"messages": messages,
"temperature": temperature,
"max_tokens": max_tokens,
**extra_body,
}
if seed is not None:
data["seed"] = seed
try:
async with session.post(f"{url}/v1/chat/completions", json=data) as resp:
resp.raise_for_status()
result = await resp.json()
text = result["choices"][0]["message"]["content"] or ""
return text, result.get("usage", {}).get("completion_tokens", 0)
except Exception as e:
print(f"chat request failed: {e}")
return "", 0
def evaluate_mrcr(
model_name: str | None = None,
num_samples: int = 40,
needles: list[int] | None = None,
max_prompt_tokens: int | None = None,
max_tokens: int = 2048,
host: str = "http://127.0.0.1",
port: int = 8000,
temperature: float = 0.0,
seed: int | None = 42,
concurrency: int = 8,
extra_body: dict | None = None,
) -> dict:
"""Run MRCR against a vLLM server; auto-discovers model and context."""
needles = needles or [2, 4, 8]
extra_body = DEFAULT_EXTRA_BODY if extra_body is None else extra_body
base_url = f"{host}:{port}"
discovered_model, server_max_len = discover_server_model(base_url)
if model_name is None:
model_name = discovered_model
if max_prompt_tokens is None:
if server_max_len is None:
raise RuntimeError(
"Server did not advertise max_model_len; pass --max-prompt-tokens."
)
max_prompt_tokens = max(512, server_max_len - max_tokens - PROMPT_SAFETY_BUFFER)
print(
f"Model: {model_name} | max_prompt_tokens={max_prompt_tokens} "
f"(server max_model_len={server_max_len}, max_tokens={max_tokens})"
)
samples = _load_mrcr_samples(
needles=needles,
max_prompt_tokens=max_prompt_tokens,
num_samples=num_samples,
seed=seed or 0,
base_url=base_url,
model_name=model_name,
)
tok_counts = [s["n_tokens"] for s in samples]
print(
f"Loaded {len(samples)} samples (needles={needles}, "
f"tokens={min(tok_counts)}-{max(tok_counts)})"
)
async def run():
sem = asyncio.Semaphore(concurrency)
responses = [""] * len(samples)
out_tokens = [0] * len(samples)
async def one(session, i):
async with sem:
text, toks = await _call_chat(
session=session,
url=base_url,
model=model_name,
messages=samples[i]["messages"],
max_tokens=max_tokens,
temperature=temperature,
seed=seed,
extra_body=extra_body,
)
responses[i] = text
out_tokens[i] = toks
timeout = aiohttp.ClientTimeout(total=1800)
async with aiohttp.ClientSession(timeout=timeout) as session:
await tqdm.gather(
*[one(session, i) for i in range(len(samples))], desc="MRCR"
)
return responses, out_tokens
tic = time.perf_counter()
responses, out_tokens = asyncio.run(run())
latency = time.perf_counter() - tic
scores = np.array(
[
score_mrcr(r, s["answer"], s["random_string_to_prepend"])
for r, s in zip(responses, samples)
]
)
prefix_hits = np.array(
[
r.startswith(s["random_string_to_prepend"])
for r, s in zip(responses, samples)
]
)
per_needle = {
f"match_ratio_n{n}": float(
scores[np.array([s["n_needles"] == n for s in samples])].mean()
)
for n in needles
if any(s["n_needles"] == n for s in samples)
}
total_out = int(sum(out_tokens))
return {
"model": model_name,
"match_ratio": float(scores.mean()),
"prefix_hit_rate": float(prefix_hits.mean()),
"per_needle": per_needle,
"num_samples": len(samples),
"latency": latency,
"total_output_tokens": total_out,
"tokens_per_second": total_out / latency if latency > 0 else 0.0,
"max_tokens": max_tokens,
"needles": needles,
"max_prompt_tokens": max_prompt_tokens,
}
def main() -> None:
p = argparse.ArgumentParser(description="MRCR evaluation for vLLM serve")
p.add_argument("--model", default=None, help="Default: discovered from /v1/models")
p.add_argument("--num-samples", type=int, default=40)
p.add_argument(
"--needles", type=int, nargs="+", default=[2, 4, 8], choices=[2, 4, 8]
)
p.add_argument(
"--max-prompt-tokens",
type=int,
default=None,
help="Default: server max_model_len - max_tokens - buffer",
)
p.add_argument("--max-tokens", type=int, default=2048)
p.add_argument("--host", default="http://127.0.0.1")
p.add_argument("--port", type=int, default=8000)
p.add_argument("--temperature", type=float, default=0.0)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--concurrency", type=int, default=8)
p.add_argument(
"--extra-body",
default=None,
help="JSON merged into each request. "
"Pass '{}' to disable the default enable_thinking=false.",
)
p.add_argument("--save-results", default=None)
args = p.parse_args()
extra_body = json.loads(args.extra_body) if args.extra_body else None
result = evaluate_mrcr(
model_name=args.model,
num_samples=args.num_samples,
needles=args.needles,
max_prompt_tokens=args.max_prompt_tokens,
max_tokens=args.max_tokens,
host=args.host,
port=args.port,
temperature=args.temperature,
seed=args.seed,
concurrency=args.concurrency,
extra_body=extra_body,
)
print("\nResults:")
print(f" match_ratio: {result['match_ratio']:.4f}")
print(f" prefix_hit_rate: {result['prefix_hit_rate']:.4f}")
for k, v in result["per_needle"].items():
print(f" {k}: {v:.4f}")
print(f" samples: {result['num_samples']}")
print(f" latency: {result['latency']:.1f}s")
print(f" output tok/s: {result['tokens_per_second']:.1f}")
if args.save_results:
with open(args.save_results, "w") as f:
json.dump(result, f, indent=2)
if __name__ == "__main__":
main()
+83
View File
@@ -0,0 +1,83 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
MRCR long-context accuracy test.
Usage:
pytest -s -v tests/evals/mrcr/test_mrcr_correctness.py \
--config-list-file=configs/models-small.txt
"""
import shlex
import yaml
from tests.utils import RemoteOpenAIServer
from .mrcr_eval import evaluate_mrcr
def _split_host_port(url: str, default_port: int = 8000) -> tuple[str, int]:
if "://" in url:
url = url.split("://", 1)[1]
host_port = url.split("/", 1)[0]
if ":" in host_port:
host, p = host_port.split(":", 1)
return f"http://{host}", int(p)
return f"http://{host_port}", default_port
def test_mrcr_correctness(config_filename):
cfg = yaml.safe_load(config_filename.read_text(encoding="utf-8"))
server_args = shlex.split(cfg.get("server_args", ""))
server_args += ["--trust-remote-code", "--disable-uvicorn-access-log"]
print(
f"MRCR eval for {cfg['model_name']} (threshold {cfg['match_ratio_threshold']})"
)
with RemoteOpenAIServer(
cfg["model_name"],
server_args,
env_dict=cfg.get("env"),
max_wait_seconds=cfg.get("startup_max_wait_seconds", 600),
) as server:
host, port = _split_host_port(server.url_for("v1"))
results = evaluate_mrcr(
model_name=cfg.get("model_name"),
num_samples=cfg.get("num_samples", 40),
needles=cfg.get("needles", [2, 4, 8]),
max_prompt_tokens=cfg.get("max_prompt_tokens"),
max_tokens=cfg.get("max_tokens", 2048),
host=host,
port=port,
concurrency=cfg.get("concurrency", 8),
extra_body=cfg.get("extra_body"),
)
threshold = cfg["match_ratio_threshold"]
tol = cfg.get("tolerance", 0.05)
print(f" match_ratio: {results['match_ratio']:.4f}")
print(f" prefix_hit_rate: {results['prefix_hit_rate']:.4f}")
for k, v in results["per_needle"].items():
print(f" {k}: {v:.4f}")
failures: list[str] = []
if isinstance(threshold, dict):
for n, expected in threshold.items():
key = f"match_ratio_n{int(n)}"
measured = results["per_needle"].get(key)
if measured is None:
failures.append(f"{key}: no samples collected")
elif measured < expected - tol:
failures.append(f"{key}: {measured:.4f} < {expected:.4f} - {tol:.4f}")
else:
measured = results["match_ratio"]
if measured < threshold - tol:
failures.append(
f"match_ratio: {measured:.4f} < {threshold:.4f} - {tol:.4f}"
)
assert not failures, "MRCR thresholds failed: " + "; ".join(failures)
+91
View File
@@ -0,0 +1,91 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
from torch import Tensor
from torch.fx.experimental.proxy_tensor import make_fx
import vllm.ir.op
from vllm.ir.op import IrOp, IrOpInplaceOverload
@vllm.ir.register_op(allow_inplace=True)
def _custom_mm2(x: Tensor, w: Tensor) -> Tensor:
return x @ w
@_custom_mm2.register_impl("regular")
def _custom_mm2_regular(x: Tensor, w: Tensor) -> Tensor:
return x @ w + 1
@_custom_mm2.register_impl("inplace", inplace=True)
def _custom_mm2_inplace(x: Tensor, w: Tensor) -> Tensor:
x.copy_(x @ w + 2)
return x
class TestInplaceOp:
def test_registration(self):
# Test that the inplace op is registered correctly.
assert "_custom_mm2" in IrOp.registry
assert IrOp.registry["_custom_mm2"] is _custom_mm2
assert _custom_mm2.torch_op is torch.ops.vllm_ir._custom_mm2.default
assert isinstance(_custom_mm2.maybe_inplace, IrOpInplaceOverload)
assert (
_custom_mm2.maybe_inplace.torch_op
is torch.ops.vllm_ir._custom_mm2.maybe_inplace
)
def test_inplace_dispatching(self):
# check that the correct implementation is dispatched based on priority,
# and inplace semantics hold
w = torch.randn(3, 3)
x = torch.randn(2, 3)
x1 = x.clone()
with _custom_mm2.set_priority(["regular"]):
result_regular = _custom_mm2.maybe_inplace(x, w)
# check that the regular op does not modify x
torch.testing.assert_close(x, x1, atol=0, rtol=0)
with _custom_mm2.set_priority(["inplace"]):
result_inplace: Tensor = _custom_mm2.maybe_inplace(x, w)
# check that the inplace op returns x directly
assert result_inplace.data_ptr() == x.data_ptr()
torch.testing.assert_close(result_inplace, x1 @ w + 2)
torch.testing.assert_close(result_regular, x1 @ w + 1)
def test_default_dispatching(self):
# check that the correct implementation is dispatched,
# and ops do not modify inputs when using the default overload
w = torch.randn(3, 3)
x = torch.randn(2, 3)
x1 = x.clone()
with _custom_mm2.set_priority(["regular"]):
result_regular = _custom_mm2(x, w)
with _custom_mm2.set_priority(["inplace"]):
result_inplace = _custom_mm2(x, w)
# check that x was not modified by either impl
torch.testing.assert_close(x, x1, atol=0, rtol=0)
torch.testing.assert_close(result_inplace, x1 @ w + 2)
torch.testing.assert_close(result_regular, x1 @ w + 1)
def test_trace(self):
# Test that the inplace op can be used in a graph.
def func(x: Tensor, y: Tensor) -> Tensor:
return _custom_mm2.maybe_inplace(x, y)
x = torch.randn(2, 3)
y = torch.randn(3, 4)
graph = make_fx(func)(x, y)
assert any(
node.target == torch.ops.vllm_ir._custom_mm2.maybe_inplace
for node in graph.graph.nodes
)
+27 -16
View File
@@ -21,7 +21,7 @@ class CustomError(Exception):
pass
@vllm.ir.register_op
@vllm.ir.register_op(allow_inplace=True)
def _custom_add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
return x + y
@@ -129,11 +129,15 @@ class TestIrOpCustomAdd:
@pytest.mark.parametrize("enable_torch_wrap", [True, False])
@pytest.mark.parametrize("symbolic_trace", [True, False])
@pytest.mark.parametrize("overload", ["default", "maybe_inplace"])
def test_trace_sees_single_custom_op(
self, symbolic_trace: bool, enable_torch_wrap: bool
self, symbolic_trace: bool, enable_torch_wrap: bool, overload: str
):
op_fn = _custom_add if overload == "default" else _custom_add.maybe_inplace
torch_op = getattr(torch.ops.vllm_ir._custom_add, overload)
def fn(x, y):
return _custom_add(x, y)
return op_fn(x, y)
def find_fn(target: Any, gm: fx.GraphModule):
return gm.graph.find_nodes(op="call_function", target=target)
@@ -155,7 +159,7 @@ class TestIrOpCustomAdd:
torch.testing.assert_close(out_fx, out_eager)
# check that IR nodes only appear if enable_torch_wrap=True
ir_nodes = find_fn(torch.ops.vllm_ir._custom_add.default, gm)
ir_nodes = find_fn(torch_op, gm)
if enable_torch_wrap:
assert len(ir_nodes) == 1, gm.code
else:
@@ -167,7 +171,7 @@ class TestIrOpCustomAdd:
else:
gm = make_fx(fn)(torch.randn(2, 2), torch.randn(2, 2))
ir_nodes = find_fn(torch.ops.vllm_ir._custom_add.default, gm)
ir_nodes = find_fn(torch_op, gm)
assert len(ir_nodes) == 1, gm.code
@@ -176,9 +180,12 @@ def impl_a(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
return x + y + 10
@_custom_add.register_impl("impl_b")
@_custom_add.register_impl("impl_b", inplace=True)
def impl_b(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
return x + y + 20
"""Computes x+y+20"""
x.add_(y)
x.add_(20)
return x
@_custom_add.register_impl("impl_even", supports_args=lambda x, y: x.size(1) % 2 == 0)
@@ -243,19 +250,23 @@ class TestIrOpImplDispatch:
# Restored to empty
assert _custom_add.get_priority() == []
def test_dispatch_priority_order(self):
@pytest.mark.parametrize("overload", ["default", "maybe_inplace"])
def test_dispatch_priority_order(self, overload: str):
op_fn = _custom_add if overload == "default" else _custom_add.maybe_inplace
torch_op = getattr(torch.ops.vllm_ir._custom_add, overload)
x = torch.tensor(1, dtype=torch.int32)
y = torch.tensor(2, dtype=torch.int32)
with _custom_add.set_priority(["impl_b", "impl_a"]):
assert _custom_add.dispatch(x, y) is impl_b
out1 = _custom_add(x, y)
out2 = torch.ops.vllm_ir._custom_add(x, y)
out1 = op_fn(x.clone(), y)
out2 = torch_op(x.clone(), y)
with _custom_add.set_priority(["impl_a"]):
assert _custom_add.dispatch(x, y) is impl_a
out3 = _custom_add(x, y)
out4 = torch.ops.vllm_ir._custom_add(x, y)
out3 = op_fn(x.clone(), y)
out4 = torch_op(x.clone(), y)
# impl_b
assert out1.item() == 1 + 2 + 20
@@ -265,18 +276,18 @@ class TestIrOpImplDispatch:
assert out4.item() == 1 + 2 + 10
def test_unsupported_impl_filtered(self):
@_custom_add.register_impl("unsupported", supported=False)
def impl_bad(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
@_custom_add.register_impl("impl_unsupported", supported=False)
def impl_unsupported(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
return x + y + 999
x = torch.tensor(1, dtype=torch.int32)
y = torch.tensor(2, dtype=torch.int32)
with _custom_add.set_priority(["unsupported", "impl_a"]):
with _custom_add.set_priority(["impl_unsupported", "impl_a"]):
assert _custom_add.get_priority() == ["impl_a"]
out = _custom_add(x, y)
# impl_bad skipped → impl_a
# impl_unsupported skipped → impl_a
assert out.item() == 1 + 2 + 10
def test_supports_args_runtime_dispatch_and_warning(
@@ -5,12 +5,17 @@ import pytest
import torch
from tests.kernels.quantization.nvfp4_utils import (
dequant_nvfp4_kv_cache,
dequantize_nvfp4_to_dtype,
get_nvfp4_global_scale,
)
from vllm.platforms import current_platform
from vllm.utils.math_utils import round_up
from vllm.utils.torch_utils import set_random_seed
from vllm.utils.torch_utils import (
nvfp4_kv_cache_full_dim,
nvfp4_kv_cache_split_views,
set_random_seed,
)
if not current_platform.is_device_capability_family(100):
pytest.skip(
@@ -33,6 +38,117 @@ def to_float8(x, dtype=torch.float8_e4m3fn):
return x_scl_sat.to(dtype), scale.float().reciprocal()
def build_paged_kv_metadata(
seq_lens: torch.Tensor,
block_tables: torch.Tensor,
block_size: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Build paged-KV indptr/indices/last_page_lens from seq_lens + block_tables."""
kv_indptr = [0]
kv_indices = []
kv_last_page_lens = []
for i in range(len(seq_lens)):
sl = int(seq_lens[i])
assert sl > 0
nb = (sl + block_size - 1) // block_size
kv_indices.extend(block_tables[i, :nb].tolist())
kv_indptr.append(kv_indptr[-1] + nb)
kv_last_page_lens.append(sl % block_size or block_size)
return (
torch.tensor(kv_indptr, dtype=torch.int32),
torch.tensor(kv_indices, dtype=torch.int32),
torch.tensor(kv_last_page_lens, dtype=torch.int32),
)
def make_nvfp4_kv_cache(
kv_bf16_hnd: torch.Tensor, block_size: int, head_size: int
) -> tuple:
"""Quantize bf16 KV cache to nvfp4 via reshape_and_cache_flash.
Returns (k_data, v_data), (k_scales, v_scales), kv_scale, ref_kv_bf16.
"""
num_blocks, _, num_kv_heads, _, _ = kv_bf16_hnd.shape
kv_scale_val = (kv_bf16_hnd.abs().amax() / 448.0).item()
kv_scale_tensor = torch.tensor(
kv_scale_val, dtype=torch.float32, device=kv_bf16_hnd.device
)
# Allocate in HND physical order, permute to NHD logical order.
# hnd_order swaps dims 2↔3; it is its own inverse.
full_dim = nvfp4_kv_cache_full_dim(head_size)
hnd_order = (0, 1, 3, 2, 4)
kv_cache = torch.zeros(
(num_blocks, 2, num_kv_heads, block_size, full_dim),
dtype=torch.uint8,
device=kv_bf16_hnd.device,
).permute(*hnd_order)
# Flatten NHD [N, T, H, D] → token tensors [N*T, H, D] for the kernel.
num_tokens = num_blocks * block_size
k_tokens = (
kv_bf16_hnd[:, 0]
.permute(0, 2, 1, 3)
.reshape(num_tokens, num_kv_heads, head_size)
)
v_tokens = (
kv_bf16_hnd[:, 1]
.permute(0, 2, 1, 3)
.reshape(num_tokens, num_kv_heads, head_size)
)
slot_mapping = torch.arange(num_tokens, dtype=torch.long, device=kv_bf16_hnd.device)
# reshape_and_cache_flash: kernel receives kv_cache[:, 0] and [:, 1]
# (full K/V buffers containing both data and scale).
torch.ops._C_cache_ops.reshape_and_cache_flash(
k_tokens,
v_tokens,
kv_cache[:, 0],
kv_cache[:, 1],
slot_mapping,
"nvfp4",
kv_scale_tensor,
kv_scale_tensor,
)
# Split in HND order for trtllm kernel (expects HND numTokensPerPage).
kv_cache_hnd = kv_cache.permute(*hnd_order)
(k_data, v_data), (k_scales, v_scales) = nvfp4_kv_cache_split_views(kv_cache_hnd)
# Dequantize for the FA2 reference baseline.
ref_k = dequant_nvfp4_kv_cache(
k_data, k_scales, kv_scale_val, head_size, block_size
).to(torch.bfloat16)
ref_v = dequant_nvfp4_kv_cache(
v_data, v_scales, kv_scale_val, head_size, block_size
).to(torch.bfloat16)
ref_kv_bf16 = torch.stack([ref_k, ref_v], dim=1) # [N, 2, H, T, D]
return (k_data, v_data), (k_scales, v_scales), kv_scale_val, ref_kv_bf16
def make_quantized_kv_cache(
kv_cache: torch.Tensor,
kv_quant_dtype: torch.dtype,
block_size: int,
head_size: int,
) -> tuple:
"""Quantize kv_cache based on dtype. Returns (kv_cache, kv_cache_sf,
kv_scale, ref_kv_cache, is_nvfp4_kv)."""
is_nvfp4_kv = kv_quant_dtype == FP4_DTYPE
if is_nvfp4_kv:
data, scales, kv_scale, ref = make_nvfp4_kv_cache(
kv_cache, block_size, head_size
)
return data, scales, kv_scale, ref, True
elif kv_quant_dtype == FP8_DTYPE:
kv_fp8, kv_scale = to_float8(kv_cache)
ref = kv_fp8.to(kv_cache.dtype) * kv_scale
return kv_fp8, None, kv_scale, ref, False
else:
return kv_cache, None, 1.0, kv_cache, False
DTYPE = [torch.bfloat16]
QUANT_DTYPES = [
# (q_quant_dtype, kv_quant_dtype, o_quant_dtype)
@@ -41,6 +157,7 @@ QUANT_DTYPES = [
(FP8_DTYPE, FP8_DTYPE, None),
(FP8_DTYPE, FP8_DTYPE, FP8_DTYPE),
(FP8_DTYPE, FP8_DTYPE, FP4_DTYPE),
(FP8_DTYPE, FP4_DTYPE, FP8_DTYPE), # nvfp4 KV cache
]
BATCH_SIZE = [4, 12]
MAX_SEQ_LENS = [(1024, 4096)]
@@ -127,35 +244,19 @@ def test_flashinfer_trtllm_decode_with_baseline(
max_seq_len = torch.max(seq_lens).item()
kv_cache = torch.randn(kv_cache_shape, dtype=dtype)
if kv_quant_dtype == FP8_DTYPE:
kv_cache, kv_scale = to_float8(kv_cache)
ref_kv_cache = kv_cache.to(dtype) * kv_scale
else:
kv_scale = 1.0
ref_kv_cache = kv_cache
kv_cache, kv_cache_sf, kv_scale, ref_kv_cache, is_nvfp4_kv = (
make_quantized_kv_cache(kv_cache, kv_quant_dtype, block_size, head_size)
)
k_scale = v_scale = kv_scale
max_num_blocks_per_seq = (max_seq_len + block_size - 1) // block_size
block_tables = torch.randint(
0, NUM_BLOCKS, (batch_size, max_num_blocks_per_seq), dtype=torch.int32
)
kv_indptr = [0]
kv_indices = []
kv_last_page_lens = []
for i in range(batch_size):
seq_len = seq_lens[i]
assert seq_len > 0
num_blocks = (seq_len + block_size - 1) // block_size
kv_indices.extend(block_tables[i, :num_blocks])
kv_indptr.append(kv_indptr[-1] + num_blocks)
kv_last_page_len = seq_len % block_size
if kv_last_page_len == 0:
kv_last_page_len = block_size
kv_last_page_lens.append(kv_last_page_len)
kv_indptr = torch.tensor(kv_indptr, dtype=torch.int32)
kv_indices = torch.tensor(kv_indices, dtype=torch.int32)
kv_last_page_lens = torch.tensor(kv_last_page_lens, dtype=torch.int32)
kv_indptr, kv_indices, kv_last_page_lens = build_paged_kv_metadata(
seq_lens, block_tables, block_size
)
workspace_buffer = torch.zeros(128 * 1024 * 1024, dtype=torch.int8)
# Baseline Decode
@@ -225,6 +326,7 @@ def test_flashinfer_trtllm_decode_with_baseline(
sinks=sinks,
o_sf_scale=o_sf_scale_float,
out=output_trtllm,
kv_cache_sf=kv_cache_sf,
)
if o_quant_dtype == FP8_DTYPE:
output_trtllm = output_trtllm.to(dtype) * o_scale
@@ -237,7 +339,9 @@ def test_flashinfer_trtllm_decode_with_baseline(
)
output_trtllm = output_trtllm.reshape(-1, query.shape[1], query.shape[2])
if q_quant_dtype == FP8_DTYPE and o_quant_dtype == FP4_DTYPE:
if is_nvfp4_kv:
rtol, atol = 1.0, 1.0 # nvfp4 has higher quantization error
elif q_quant_dtype == FP8_DTYPE and o_quant_dtype == FP4_DTYPE:
rtol, atol = 7e-2, 9e-2
elif q_quant_dtype == FP8_DTYPE and o_quant_dtype == FP8_DTYPE:
rtol, atol = 3e-2, 4e-2
@@ -287,7 +391,12 @@ def test_flashinfer_trtllm_prefill_with_baseline(
kv_quant_dtype = kv_quant_dtype or dtype
o_quant_dtype = o_quant_dtype or dtype
if q_quant_dtype != kv_quant_dtype:
# FP8 Q + nvfp4 KV is the required combination for the nvfp4 KV path.
# All other mixed Q/KV dtype combinations are unsupported.
is_nvfp4_kv = kv_quant_dtype == FP4_DTYPE
if q_quant_dtype != kv_quant_dtype and not (
q_quant_dtype == FP8_DTYPE and is_nvfp4_kv
):
pytest.skip("Skipped mixed QKV dtypes for prefill")
max_q_len, max_kv_len = max_seq_lens
@@ -329,35 +438,19 @@ def test_flashinfer_trtllm_prefill_with_baseline(
max_seq_len = torch.max(seq_lens).item()
kv_cache = torch.randn(kv_cache_shape, dtype=dtype)
if kv_quant_dtype == FP8_DTYPE:
kv_cache, kv_scale = to_float8(kv_cache)
ref_kv_cache = kv_cache.to(dtype) * kv_scale
else:
kv_scale = 1.0
ref_kv_cache = kv_cache
kv_cache, kv_cache_sf, kv_scale, ref_kv_cache, is_nvfp4_kv = (
make_quantized_kv_cache(kv_cache, kv_quant_dtype, block_size, head_size)
)
k_scale = v_scale = kv_scale
max_num_blocks_per_seq = (max_seq_len + block_size - 1) // block_size
block_tables = torch.randint(
0, NUM_BLOCKS, (batch_size, max_num_blocks_per_seq), dtype=torch.int32
)
kv_indptr = [0]
kv_indices = []
kv_last_page_lens = []
for i in range(batch_size):
seq_len = seq_lens[i]
assert seq_len > 0
num_blocks = (seq_len + block_size - 1) // block_size
kv_indices.extend(block_tables[i, :num_blocks])
kv_indptr.append(kv_indptr[-1] + num_blocks)
kv_last_page_len = seq_len % block_size
if kv_last_page_len == 0:
kv_last_page_len = block_size
kv_last_page_lens.append(kv_last_page_len)
kv_indptr = torch.tensor(kv_indptr, dtype=torch.int32)
kv_indices = torch.tensor(kv_indices, dtype=torch.int32)
kv_last_page_lens = torch.tensor(kv_last_page_lens, dtype=torch.int32)
kv_indptr, kv_indices, kv_last_page_lens = build_paged_kv_metadata(
seq_lens, block_tables, block_size
)
workspace_buffer = torch.zeros(128 * 1024 * 1024, dtype=torch.int8)
# Baseline Prefill
@@ -431,6 +524,7 @@ def test_flashinfer_trtllm_prefill_with_baseline(
sinks=sinks,
o_sf_scale=o_sf_scale_float,
out=output_trtllm,
kv_cache_sf=kv_cache_sf,
)
if o_quant_dtype == FP8_DTYPE:
output_trtllm = output_trtllm.to(dtype) * o_scale
@@ -443,7 +537,9 @@ def test_flashinfer_trtllm_prefill_with_baseline(
)
output_trtllm = output_trtllm.reshape(-1, query.shape[1], query.shape[2])
if q_quant_dtype == FP8_DTYPE and o_quant_dtype == FP4_DTYPE:
if is_nvfp4_kv:
rtol, atol = 1.0, 1.5 # nvfp4 has higher quantization error
elif q_quant_dtype == FP8_DTYPE and o_quant_dtype == FP4_DTYPE:
rtol, atol = 3e-1, 4e-1
elif q_quant_dtype == FP8_DTYPE and o_quant_dtype == FP8_DTYPE:
rtol, atol = 4e-2, 6e-2
+205 -1
View File
@@ -28,7 +28,9 @@ def test_rms_norm_registration():
"native": True,
"vllm_c": current_platform.is_cuda_alike(),
"aiter": current_platform.is_rocm(),
"oink": False,
"oink": current_platform.has_device_capability(100)
and hasattr(torch.ops, "oink")
and hasattr(torch.ops.oink, "rmsnorm"),
"xpu_kernels": current_platform.is_xpu(),
}
@@ -67,6 +69,14 @@ class TestRMSNorm:
out2 = rms_norm_native(x * 2.0, weight, epsilon=epsilon)
torch.testing.assert_close(out2, out, rtol=get_default_rtol(out), atol=1e-3)
# Mean square should be approximately 1 (ignoring epsilon and weight scaling)
combined_norm = out.float() / weight.float()
variance = combined_norm.pow(2).mean(dim=-1)
# After RMS normalization, variance should be close to 1
torch.testing.assert_close(
variance, torch.ones_like(variance), rtol=1e-2, atol=1e-2
)
# Check behavior with and without weight
weight1 = torch.ones_like(weight)
out3 = rms_norm_native(x, weight1, epsilon=epsilon)
@@ -129,3 +139,197 @@ def test_aiter_rejects_unsupported_dtypes():
num_tokens=8, hidden_size=4096, dtype=dtype, epsilon=1e-5
)
assert not impl.supports_args(*args), f"aiter should reject dtype={dtype}"
fused_add_rms_norm_native = ir.ops.fused_add_rms_norm.impls["native"].impl_fn
@pytest.mark.skipif(
not current_platform.is_cuda_alike() and not current_platform.is_xpu(),
reason="Currently only kernels on CUDA, ROCm and XPU",
)
def test_fused_add_rms_norm_registration():
expected = {
"native": True,
"vllm_c": current_platform.is_cuda_alike(),
"aiter": current_platform.is_rocm(),
"oink": current_platform.has_device_capability(100)
and hasattr(torch.ops, "oink")
and hasattr(torch.ops.oink, "fused_add_rms_norm"),
"xpu_kernels": current_platform.is_xpu(),
}
actual = {
provider: impl.supported
for provider, impl in ir.ops.fused_add_rms_norm.impls.items()
}
assert actual == expected
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32])
@pytest.mark.parametrize("n_tokens", NUM_TOKENS)
@pytest.mark.parametrize("hidden_size", COMMON_HIDDEN_SIZES)
@pytest.mark.parametrize("epsilon", [1e-6, 1e-5])
@pytest.mark.skipif(
not current_platform.is_cuda_alike() and not current_platform.is_xpu(),
reason="Currently only kernels on CUDA, ROCm and XPU",
)
class TestFusedAddRMSNorm:
@classmethod
def setup_class(cls, **kwargs):
torch.set_default_device(current_platform.device_type)
def test_native_semantics(self, dtype, n_tokens, hidden_size, epsilon):
x, x_residual, weight, eps = ir.ops.fused_add_rms_norm.generate_inputs(
num_tokens=4, hidden_size=8, dtype=dtype, epsilon=epsilon
)
out, residual_out = fused_add_rms_norm_native(x, x_residual, weight, eps)
# Check shape, dtype, device
assert out.shape == x.shape
assert out.dtype == x.dtype
assert out.device == x.device
assert residual_out.shape == x_residual.shape
assert residual_out.dtype == x_residual.dtype
assert residual_out.device == x_residual.device
# Check that residual_out = x + x_residual
expected_residual = (x.float() + x_residual.float()).to(dtype)
torch.testing.assert_close(
residual_out, expected_residual, rtol=1e-3, atol=1e-3
)
# Verify that the output is RMS normalized version of (x + x_residual)
expected_out = rms_norm_native(expected_residual, weight, epsilon)
assert_close(
ir.ops.fused_add_rms_norm,
(out, residual_out),
(expected_out, expected_residual),
)
# Check the scaling property of rms norm
out1, _ = fused_add_rms_norm_native(
x, torch.zeros_like(x), weight, epsilon=epsilon
)
out2, _ = fused_add_rms_norm_native(
x * 2.0, torch.zeros_like(x), weight, epsilon=epsilon
)
torch.testing.assert_close(out2, out1, rtol=get_default_rtol(out), atol=1e-3)
# Check behavior with and without weight
weight1 = torch.ones_like(weight)
out3, _ = fused_add_rms_norm_native(x, x_residual, weight1, eps)
out4, _ = fused_add_rms_norm_native(x, x_residual, None, eps)
torch.testing.assert_close(out3, out4)
@pytest.mark.parametrize("provider", supported_providers(ir.ops.fused_add_rms_norm))
def test_impls(self, dtype, n_tokens, hidden_size, epsilon, provider):
impl = ir.ops.fused_add_rms_norm.impls[provider]
x, x_residual, weight, eps = ir.ops.fused_add_rms_norm.generate_inputs(
num_tokens=n_tokens, hidden_size=hidden_size, dtype=dtype, epsilon=epsilon
)
args = (x, x_residual, weight, eps, None)
if not impl.supports_args(*args):
pytest.skip(f"{provider} does not support args")
ref_output, ref_residual = fused_add_rms_norm_native(*clone_args(args))
output, residual = impl.impl_fn(*clone_args(args))
assert_close(ir.ops.fused_add_rms_norm, output, ref_output)
assert_close(ir.ops.fused_add_rms_norm, residual, ref_residual)
# check that dispatched call matches direct call
with ir.ops.fused_add_rms_norm.set_priority([provider, "native"]):
out_dispatched, residual_dispatched = ir.ops.fused_add_rms_norm(*args[:4])
out_direct, residual_direct = impl.impl_fn(*clone_args(args))
torch.testing.assert_close(out_dispatched, out_direct, rtol=0.0, atol=0.0)
torch.testing.assert_close(
residual_dispatched, residual_direct, rtol=0.0, atol=0.0
)
# none of these support variance_size override
assert not impl.supports_args(x, x_residual, weight, epsilon, 4)
assert not impl.supports_args(x, x_residual, weight, epsilon, variance_size=4)
# test weight=None behavior
out_no_weight, residual_no_weight = impl.impl_fn(
x.clone(), x_residual.clone(), None, epsilon
)
out_unit_weight, residual_unit_weight = impl.impl_fn(
x.clone(), x_residual.clone(), torch.ones_like(weight), epsilon
)
assert_close(ir.ops.fused_add_rms_norm, out_no_weight, out_unit_weight)
assert_close(
ir.ops.fused_add_rms_norm, residual_no_weight, residual_unit_weight
)
@pytest.mark.parametrize("provider", ["vllm_c"])
def test_inplace_semantics(self, dtype, n_tokens, hidden_size, epsilon, provider):
"""Test that inplace implementations reuse inputs,
for maybe_inplace overload but not for default overload."""
impl = ir.ops.fused_add_rms_norm.impls[provider]
if not impl.supported:
pytest.skip(f"{provider} impl not supported on this platform")
x, x_residual, weight, eps = ir.ops.fused_add_rms_norm.generate_inputs(
num_tokens=n_tokens, hidden_size=hidden_size, dtype=dtype, epsilon=epsilon
)
# Test default overload - should NOT modify inputs even with inplace impl
x_default = x.clone()
x_residual_default = x_residual.clone()
x_default_ptr = x_default.data_ptr()
x_residual_default_ptr = x_residual_default.data_ptr()
with ir.ops.fused_add_rms_norm.set_priority([provider, "native"]):
out_default, residual_default = ir.ops.fused_add_rms_norm(
x_default, x_residual_default, weight, eps
)
# Default should NOT be inplace (even with inplace implementation)
assert out_default.data_ptr() != x_default_ptr
assert residual_default.data_ptr() != x_residual_default_ptr
torch.testing.assert_close(x, x_default, rtol=0.0, atol=0.0)
torch.testing.assert_close(x_residual, x_residual_default, rtol=0.0, atol=0.0)
# Test maybe_inplace overload - should modify inputs with inplace impl
x_inplace = x.clone()
x_residual_inplace = x_residual.clone()
x_inplace_ptr = x_inplace.data_ptr()
x_residual_inplace_ptr = x_residual_inplace.data_ptr()
with ir.ops.fused_add_rms_norm.set_priority([provider, "native"]):
out_inplace, residual_inplace = ir.ops.fused_add_rms_norm.maybe_inplace(
x_inplace, x_residual_inplace, weight, eps
)
# maybe_inplace should be inplace
assert out_inplace.data_ptr() == x_inplace_ptr
assert residual_inplace.data_ptr() == x_residual_inplace_ptr
# Both should produce same results
torch.testing.assert_close(out_default, out_inplace, atol=0.0, rtol=0.0)
torch.testing.assert_close(
residual_default, residual_inplace, atol=0.0, rtol=0.0
)
@pytest.mark.parametrize("provider", supported_providers(ir.ops.fused_add_rms_norm))
def test_torch_opcheck(self, dtype, n_tokens, hidden_size, epsilon, provider):
args = ir.ops.fused_add_rms_norm.generate_inputs(
num_tokens=n_tokens, hidden_size=hidden_size, dtype=dtype, epsilon=epsilon
)
args = args + (None,) # Add variance_size parameter
# When checking the torch op, we have to set priority and use dispatch
with ir.ops.fused_add_rms_norm.set_priority([provider, "native"]):
torch.library.opcheck(torch.ops.vllm_ir.fused_add_rms_norm.default, args)
# Only test maybe_inplace with non-inplace implementations
# Inplace implementations return aliases of inputs which is not allowed.
# We break this invariant, but we also convert maybe_inplace to the default
# overload during compilation, so maybe_inplace never reaches Inductor.
if not ir.ops.fused_add_rms_norm.impls[provider].inplace:
torch.library.opcheck(
torch.ops.vllm_ir.fused_add_rms_norm.maybe_inplace, args
)
+207
View File
@@ -0,0 +1,207 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Tests for the FlashInfer TRTLLM NvFP4 MoE backend
(`TrtLlmNvFp4ExpertsModular`).
Covers the activations the wrapper claims to support SiLU, RELU^2 (non-gated),
and GELU including a Gemma4-shaped case (128 experts, top-k 8,
intermediate_size 704) that exercises the non-256-aligned padding path.
"""
import pytest
import torch
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
from tests.kernels.moe.utils import make_test_quant_config
from tests.kernels.quantization.nvfp4_utils import (
FLOAT4_E2M1_MAX,
FLOAT8_E4M3_MAX,
dequantize_nvfp4_to_dtype,
)
from tests.kernels.utils import torch_moe
from vllm import _custom_ops as ops
from vllm.config import ParallelConfig, VllmConfig, set_current_vllm_config
from vllm.model_executor.layers.fused_moe import fused_topk
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.all2all_utils import (
maybe_make_prepare_finalize,
)
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEConfig,
FusedMoEParallelConfig,
RoutingMethodType,
)
from vllm.model_executor.layers.fused_moe.experts.trtllm_nvfp4_moe import (
TrtLlmNvFp4ExpertsModular,
)
from vllm.platforms import current_platform
from vllm.utils.flashinfer import has_flashinfer_trtllm_fused_moe
from vllm.utils.math_utils import next_power_of_2
from vllm.utils.torch_utils import set_random_seed
if pytest and (
not has_flashinfer_trtllm_fused_moe()
or not current_platform.has_device_capability(100)
):
pytest.skip(
"Requires flashinfer TRTLLM fused MoE and NvFP4 (SM100)",
allow_module_level=True,
)
# (m, n, k) = (tokens, intermediate_size_per_partition, hidden_dim).
# The (64, 704, 4096) row matches Gemma4's MoE shape and exercises the
# non-256-aligned intermediate (padded inside the wrapper).
MNK_FACTORS = [
(2, 1024, 1024),
(64, 2048, 1536),
(64, 704, 4096),
]
@pytest.mark.parametrize("m,n,k", MNK_FACTORS)
@pytest.mark.parametrize("e", [128])
@pytest.mark.parametrize("topk", [8])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize(
"activation",
[MoEActivation.SILU, MoEActivation.RELU2_NO_MUL, MoEActivation.GELU],
)
@torch.inference_mode()
def test_trtllm_fp4_moe_no_graph(
m: int,
n: int,
k: int,
e: int,
topk: int,
dtype: torch.dtype,
activation: MoEActivation,
workspace_init,
):
# FlashInfer's trtllm_batched_gemm_runner has no precompiled tile
# config for non-gated RELU^2 at non-256-aligned intermediate_size
# (e.g. Gemma4's 704). Other activations (SiLU/GELU) work at the
# same shape. Tracked upstream in FlashInfer; unrelated to this
# PR's GELU enablement (Gemma4 uses GeGLU, not non-gated RELU^2).
if activation == MoEActivation.RELU2_NO_MUL and (m, n, k) == (64, 704, 4096):
pytest.skip(
"FlashInfer trtllm_batched_gemm_runner: no valid tile config "
"for non-gated RELU^2 at intermediate_size=704 "
"(getValidConfigIndices throws). Tracked upstream."
)
set_random_seed(7)
with set_current_vllm_config(
VllmConfig(parallel_config=ParallelConfig(pipeline_parallel_size=1))
):
a = torch.randn((m, k), device="cuda", dtype=dtype) / 10
quant_blocksize = 16
is_gated_act = activation.is_gated
w1_q, w2_q, quant_config = make_test_quant_config(
e,
n,
k,
in_dtype=dtype,
quant_dtype="nvfp4",
block_shape=None,
per_act_token_quant=False,
make_gate=is_gated_act,
# The TRT-LLM FP4 MoE kernel rejects swizzled (padded) activation
# scales — its numel-based vec_size check requires numel == M*K/16.
# Match what oracle/nvfp4.py does for this backend.
is_nvfp4_scale_swizzled=False,
)
score = torch.randn((m, e), device="cuda", dtype=dtype)
topk_weights, topk_ids, _ = fused_topk(a, score, topk, renormalize=False)
moe_config = FusedMoEConfig(
num_experts=e,
experts_per_token=topk,
hidden_dim=k,
intermediate_size_per_partition=n,
num_local_experts=e,
num_logical_experts=e,
activation=activation,
device="cuda",
moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
in_dtype=dtype,
is_act_and_mul=is_gated_act,
routing_method=RoutingMethodType.TopK,
max_num_tokens=next_power_of_2(m),
)
trtllm_experts = mk.FusedMoEKernel(
maybe_make_prepare_finalize(
moe=moe_config,
quant_config=quant_config,
allow_new_interface=True,
use_monolithic=False,
),
TrtLlmNvFp4ExpertsModular(moe_config=moe_config, quant_config=quant_config),
inplace=False,
)
trtllm_output = trtllm_experts.apply(
hidden_states=a,
w1=w1_q,
w2=w2_q,
topk_weights=topk_weights,
topk_ids=topk_ids,
activation=activation,
global_num_experts=e,
expert_map=None,
apply_router_weight_on_input=False,
)
# Reference: round-trip activations and weights through FP4
# quant/dequant so the comparison isolates kernel/activation behavior
# from quantization error.
a_global_scale = ((FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX) / a.abs().max()).to(
torch.float32
)
a_fp4, a_scale_interleaved = ops.scaled_fp4_quant(a, a_global_scale)
a_in_dtype = dequantize_nvfp4_to_dtype(
a_fp4,
a_scale_interleaved,
a_global_scale,
dtype=a.dtype,
device=a.device,
block_size=quant_blocksize,
)
w1_d = torch.empty(
(e, (2 if is_gated_act else 1) * n, k), device="cuda", dtype=dtype
)
w2_d = torch.empty((e, k, n), device="cuda", dtype=dtype)
for idx in range(e):
w1_d[idx] = dequantize_nvfp4_to_dtype(
w1_q[idx],
quant_config.w1_scale[idx],
(1 / quant_config.g1_alphas[idx]),
dtype=dtype,
device=w1_q.device,
block_size=quant_blocksize,
)
w2_d[idx] = dequantize_nvfp4_to_dtype(
w2_q[idx],
quant_config.w2_scale[idx],
(1 / quant_config.g2_alphas[idx]),
dtype=dtype,
device=w2_q.device,
block_size=quant_blocksize,
)
torch_output = torch_moe(
a_in_dtype, w1_d, w2_d, score, topk, activation=activation
)
torch.testing.assert_close(torch_output, trtllm_output, atol=2e-1, rtol=2e-1)
if __name__ == "__main__":
test_trtllm_fp4_moe_no_graph(
64, 704, 4096, 128, 8, torch.bfloat16, MoEActivation.GELU, None
)
+2
View File
@@ -402,6 +402,7 @@ def make_test_quant_config(
per_act_token_quant: bool = False,
block_shape: list[int] | None = None,
make_gate: bool = True,
is_nvfp4_scale_swizzled: bool = True,
) -> tuple[torch.Tensor, torch.Tensor, FusedMoEQuantConfig]:
(_, w1, w1_s, w1_gs), (_, w2, w2_s, w2_gs) = make_test_weights(
e,
@@ -442,6 +443,7 @@ def make_test_quant_config(
# TODO: make sure this is handled properly
g1_alphas=(1 / w1_gs) if w1_gs is not None else None,
g2_alphas=(1 / w2_gs) if w2_gs is not None else None,
is_nvfp4_scale_swizzled=is_nvfp4_scale_swizzled,
),
)
@@ -0,0 +1,308 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import huggingface_hub
import pytest
import torch
from safetensors import safe_open
from vllm.model_executor.layers.quantization.utils import (
nvfp4_emulation_utils,
)
from vllm.model_executor.layers.quantization.utils.nvfp4_emulation_utils import (
dequantize_to_dtype,
ref_nvfp4_quant_dequant,
)
from vllm.platforms import current_platform
from vllm.triton_utils import triton
@pytest.mark.skipif(
not current_platform.is_cuda_alike(),
reason="Triton NVFP4 kernel requires CUDA.",
)
def test_triton_dequantize_nvfp4(monkeypatch) -> None:
"""Test the Triton dequantization kernel against the CPU reference
using real NVFP4 weights from a checkpoint.
Tests both 2D (attention projection) and 3D (stacked MoE experts).
"""
checkpoint_path = huggingface_hub.snapshot_download(
"nvidia/Qwen3-30B-A3B-NVFP4",
allow_patterns=["model-00001-of-00004.safetensors"],
)
shard_path = f"{checkpoint_path}/model-00001-of-00004.safetensors"
block_size = 16
with safe_open(shard_path, framework="pt", device="cpu") as f:
all_keys = list(f.keys())
# 2D case: attention projection
tensor_fp4_2d = f.get_tensor("model.layers.9.self_attn.k_proj.weight")
tensor_sf_2d = f.get_tensor("model.layers.9.self_attn.k_proj.weight_scale")
global_scale_2d = f.get_tensor("model.layers.9.self_attn.k_proj.weight_scale_2")
# 3D case: stack ALL experts for layer 9 up_proj
expert_prefix = "model.layers.9.mlp.experts."
expert_indices = sorted(
int(key.split(".")[5])
for key in all_keys
if key.startswith(expert_prefix) and key.endswith(".up_proj.weight")
)
assert len(expert_indices) > 0
all_fp4 = []
all_sf = []
all_global_scale = []
for index in expert_indices:
name = f"{expert_prefix}{index}.up_proj"
all_fp4.append(f.get_tensor(f"{name}.weight"))
all_sf.append(f.get_tensor(f"{name}.weight_scale"))
all_global_scale.append(f.get_tensor(f"{name}.weight_scale_2"))
tensor_fp4_3d = torch.stack(all_fp4)
tensor_sf_3d = torch.stack(all_sf)
global_scale_3d = torch.stack(all_global_scale)
test_cases = [
("2D base", tensor_fp4_2d, tensor_sf_2d, global_scale_2d),
(
"2D 2x rows",
tensor_fp4_2d.repeat(2, 1),
tensor_sf_2d.repeat(2, 1),
global_scale_2d,
),
(
"2D 4x rows",
tensor_fp4_2d.repeat(4, 1),
tensor_sf_2d.repeat(4, 1),
global_scale_2d,
),
(
"2D 2x cols",
tensor_fp4_2d.repeat(1, 2),
tensor_sf_2d.repeat(1, 2),
global_scale_2d,
),
("3D base", tensor_fp4_3d, tensor_sf_3d, global_scale_3d),
(
"3D 2x experts",
tensor_fp4_3d.repeat(2, 1, 1),
tensor_sf_3d.repeat(2, 1, 1),
global_scale_3d.repeat(2),
),
(
"3D 2x rows",
tensor_fp4_3d.repeat(1, 2, 1),
tensor_sf_3d.repeat(1, 2, 1),
global_scale_3d,
),
(
"3D 2x cols",
tensor_fp4_3d.repeat(1, 1, 2),
tensor_sf_3d.repeat(1, 1, 2),
global_scale_3d,
),
]
quantiles = [0.5, 0.001, 0.999]
# Move the E2M1 lookup table to CUDA ahead of time, as would normally
# happen during model loading (process_weights_after_loading). Both the
# Triton and PyTorch reference paths run on CUDA.
nvfp4_emulation_utils.kE2M1ToFloat_handle.val = (
nvfp4_emulation_utils.kE2M1ToFloat_handle.val.cuda()
)
for label, tensor_fp4, tensor_sf, global_scale in test_cases:
fp4_cuda = tensor_fp4.cuda()
sf_cuda = tensor_sf.cuda()
gs_cuda = global_scale.cuda()
# Triton path
triton_result = dequantize_to_dtype(
fp4_cuda,
sf_cuda,
gs_cuda,
torch.bfloat16,
block_size,
swizzle=False,
)
# Reference path (PyTorch ops on CUDA, Triton dispatch disabled)
with monkeypatch.context() as m:
m.setattr(
nvfp4_emulation_utils.current_platform,
"is_cuda_alike",
lambda: False,
)
reference = dequantize_to_dtype(
fp4_cuda,
sf_cuda,
gs_cuda,
torch.bfloat16,
block_size,
swizzle=False,
)
torch.testing.assert_close(triton_result, reference, atol=0, rtol=0)
# Benchmark
shape = list(tensor_fp4.shape)
def _triton_bench(
fp4_cuda=fp4_cuda,
scale_cuda=sf_cuda,
global_scale_cuda=gs_cuda,
block_size=block_size,
):
return dequantize_to_dtype(
fp4_cuda,
scale_cuda,
global_scale_cuda,
torch.bfloat16,
block_size,
swizzle=False,
)
triton_ms, triton_min, triton_max = triton.testing.do_bench(
_triton_bench, quantiles=quantiles
)
def _reference_bench(
fp4_cuda=fp4_cuda,
scale_cuda=sf_cuda,
global_scale_cuda=gs_cuda,
block_size=block_size,
):
with monkeypatch.context() as m2:
m2.setattr(
nvfp4_emulation_utils.current_platform,
"is_cuda_alike",
lambda: False,
)
dequantize_to_dtype(
fp4_cuda,
scale_cuda,
global_scale_cuda,
torch.bfloat16,
block_size,
swizzle=False,
)
ref_ms, ref_min, ref_max = triton.testing.do_bench(
_reference_bench, quantiles=quantiles
)
speedup = ref_ms / triton_ms if triton_ms > 0 else float("inf")
print(f" dequantize {label} {shape}:")
print(
f" triton: median={triton_ms:.3f}ms, "
f"min={triton_min:.3f}ms, max={triton_max:.3f}ms"
)
print(
f" reference: median={ref_ms:.3f}ms, "
f"min={ref_min:.3f}ms, max={ref_max:.3f}ms"
)
print(f" speedup: {speedup:.2f}x")
@pytest.mark.skipif(
not current_platform.is_cuda_alike(),
reason="Triton NVFP4 kernel requires CUDA.",
)
@pytest.mark.parametrize(
"m, k",
[
(1, 16),
(1, 4096),
(2, 4096),
(4, 4096),
(8, 4096),
(16, 4096),
(24, 4096),
(32, 4096),
(1, 8192),
(2, 8192),
(4, 8192),
(8, 8192),
(16, 8192),
(24, 8192),
(32, 8192),
(1, 32),
(2, 48),
(7, 64),
(16, 128),
(33, 160),
(128, 256),
(256, 512),
(1024, 1024),
(5120, 2048),
(2048, 4096),
(4096, 7168),
(8192, 8192),
(128, 16384),
],
)
@pytest.mark.parametrize("global_scale_value", [0.5, 1.0, 0.001])
def test_triton_nvfp4_quant_dequant(
monkeypatch, m: int, k: int, global_scale_value: float
) -> None:
"""Test the Triton quant-dequant kernel against the CPU reference."""
block_size = 16
x = torch.randn(m, k, dtype=torch.bfloat16, device="cuda")
global_scale = torch.tensor(global_scale_value, dtype=torch.float32, device="cuda")
# Triton path
triton_result = ref_nvfp4_quant_dequant(x, global_scale, block_size)
# CPU reference path
with monkeypatch.context() as mp:
mp.setattr(
nvfp4_emulation_utils.current_platform,
"is_cuda_alike",
lambda: False,
)
reference = ref_nvfp4_quant_dequant(x.cpu(), global_scale.cpu(), block_size)
torch.testing.assert_close(triton_result.cpu(), reference, atol=0, rtol=0)
# Benchmark (both paths on CUDA tensors for fair comparison)
quantiles = [0.5, 0.001, 0.999]
def _triton_bench(
input_tensor=x, input_global_scale=global_scale, input_block_size=block_size
):
return ref_nvfp4_quant_dequant(
input_tensor, input_global_scale, input_block_size
)
triton_ms, triton_min, triton_max = triton.testing.do_bench(
_triton_bench, quantiles=quantiles
)
def _reference_bench(
input_tensor=x, input_global_scale=global_scale, input_block_size=block_size
):
with monkeypatch.context() as mp2:
mp2.setattr(
nvfp4_emulation_utils.current_platform,
"is_cuda_alike",
lambda: False,
)
ref_nvfp4_quant_dequant(input_tensor, input_global_scale, input_block_size)
ref_ms, ref_min, ref_max = triton.testing.do_bench(
_reference_bench, quantiles=quantiles
)
speedup = ref_ms / triton_ms if triton_ms > 0 else float("inf")
print(f" quant_dequant [{m}x{k}] gs={global_scale_value}:")
print(
f" triton: median={triton_ms:.3f}ms, "
f"min={triton_min:.3f}ms, max={triton_max:.3f}ms"
)
print(
f" reference: median={ref_ms:.3f}ms, "
f"min={ref_min:.3f}ms, max={ref_max:.3f}ms"
)
print(f" speedup: {speedup:.2f}x")
@@ -73,11 +73,6 @@ def test_per_token_group_quant_fp8(
# Larger shapes with padding
(127, 7168, 128),
(253, 640, 128),
# Non-power-of-2 group size
(4, 768, 96), # 768/96=8 groups, no padding
(3, 768, 96), # 768/96=8 groups, MN padding
(4, 480, 96), # 480/96=5 groups, K padding
(1, 480, 96), # both MN and K padding
],
)
@pytest.mark.parametrize("poisoned_scales", [False, True])
@@ -161,6 +156,188 @@ def test_per_token_group_quant_fp8_packed(
)
@pytest.mark.skipif(
not current_platform.is_cuda(), reason="DeepGEMM not available on this platform"
)
def test_per_token_group_quant_fp8_packed_all_zero():
"""All-zero input must produce well-defined UE8M0 scale bytes via the eps
floor in the kernel's UE8M0 path. Locks down the all-zero behavior before
optimization.
The CUDA kernel computes:
y_s = eps / fp8_max
y_s = exp2(ceil(log2(fmax(y_s, 1e-10))))
For all-zero input, eps/fp8_max < 1e-10, so the inner fmax clamps back to
1e-10, giving exp2(ceil(log2(1e-10))) = exp2(-33) => UE8M0 byte 0x5E (94).
"""
device = "cuda"
num_tokens, hidden_dim, group_size = 4, 7168, 128
x = torch.zeros((num_tokens, hidden_dim), device=device, dtype=torch.bfloat16)
out_q, out_s_packed = fp8_utils.per_token_group_quant_fp8_packed_for_deepgemm(
x,
group_size=group_size,
use_ue8m0=True,
)
# Quantized values must be all zero.
assert torch.equal(
out_q.view(torch.uint8),
torch.zeros_like(out_q, dtype=torch.uint8),
), "All-zero input should produce all-zero FP8 output"
# UE8M0 byte produced by the kernel for all-zero input.
# The kernel's inner fmax(y_s, 1e-10) clamps eps/fp8_max back to 1e-10.
# 1e-10 as float32 has biased exponent 0x5D and a non-zero mantissa, so
# the kernel's bit-twiddle (exp_bits + (mant_bits != 0)) rounds up to
# 0x5E. This matches exp2(ceil(log2(1e-10))) = exp2(-33).
expected_exp_byte = 0x5E
mn = num_tokens
groups_per_row = hidden_dim // group_size
k_num_packed = (groups_per_row + 3) // 4
tma_aligned_mn = ((mn + 3) // 4) * 4
num_scale_elems = mn + (k_num_packed - 1) * tma_aligned_mn
# All valid scale slots must contain the expected packed value.
# Padding slots must be zero.
actual = torch.as_strided(out_s_packed, (num_scale_elems,), (1,)).cpu()
expected = torch.zeros(num_scale_elems, dtype=torch.int32, device="cpu")
for row in range(mn):
for g in range(groups_per_row):
pack_col = g // 4
pos = g % 4
idx = pack_col * tma_aligned_mn + row
expected[idx] |= expected_exp_byte << (pos * 8)
assert torch.equal(actual, expected), "All-zero scale bytes mismatch"
@pytest.mark.skipif(
not current_platform.is_cuda(), reason="DeepGEMM not available on this platform"
)
def test_per_token_group_quant_fp8_packed_mantissa_rounds_up():
"""Inputs whose absmax/max_8bit produces a non-power-of-2 force the
mantissa-rounding-up branch (exp_byte += 1). Locks down this behavior
before optimization."""
device = "cuda"
num_tokens, hidden_dim, group_size = 4, 7168, 128
# Build a tensor whose per-group absmax = 1.5 * fp8_max * 2^k for various k.
# fp8_max = torch.finfo(torch.float8_e4m3fn).max = 448.0.
# Then absmax/fp8_max = 1.5 * 2^k -> non-zero mantissa, triggers ceil
# rounding to 2^(k+1). Use k=0 for simplicity; the bf16 representation of
# 1.5*448=672.0 is exact.
x = torch.full(
(num_tokens, hidden_dim),
672.0,
device=device,
dtype=torch.bfloat16,
)
out_q, out_s_packed = fp8_utils.per_token_group_quant_fp8_packed_for_deepgemm(
x,
group_size=group_size,
use_ue8m0=True,
)
with patch("vllm.platforms.current_platform.is_cuda", return_value=False):
ref_q, ref_s = fp8_utils.per_token_group_quant_fp8(
x,
group_size,
use_ue8m0=True,
)
assert torch.equal(out_q, ref_q), "Quantized output mismatch"
mn = num_tokens
groups_per_row = hidden_dim // group_size
k_num_packed = (groups_per_row + 3) // 4
tma_aligned_mn = ((mn + 3) // 4) * 4
num_scale_elems = mn + (k_num_packed - 1) * tma_aligned_mn
ref_s_flat = ref_s.reshape(mn, groups_per_row)
ref_exponents = (ref_s_flat.view(torch.int32) >> 23) & 0xFF
expected = torch.zeros(num_scale_elems, dtype=torch.int32, device="cpu")
for row in range(mn):
for g in range(groups_per_row):
pack_col = g // 4
pos = g % 4
idx = pack_col * tma_aligned_mn + row
expected[idx] |= int(ref_exponents[row, g].item()) << (pos * 8)
actual = torch.as_strided(out_s_packed, (num_scale_elems,), (1,)).cpu()
assert torch.equal(actual, expected), "Scale bytes mismatch"
@pytest.mark.parametrize(
"num_tokens,hidden_dim",
[
(1, 7168), # mn padded 1 -> 4
(2, 7168), # mn padded 2 -> 4
(3, 7168), # mn padded 3 -> 4
(5, 7168), # mn padded 5 -> 8
(127, 7168), # mn padded 127 -> 128
(253, 640), # both mn and groups padded
(1, 384), # extreme: 1 group, 1 mn row -> both axes padded
],
)
@pytest.mark.skipif(
not current_platform.is_cuda(), reason="DeepGEMM not available on this platform"
)
def test_per_token_group_quant_fp8_packed_zero_fills_padded_output_q(
num_tokens, hidden_dim
):
"""When output_q is allocated with shape (tma_aligned_mn, k) instead of
(mn, k), the kernel must overwrite the padded mn rows with zeros so
callers can use ``torch.empty`` instead of ``torch.zeros``."""
device = "cuda"
group_size = 128
torch.manual_seed(42)
x = torch.randn((num_tokens, hidden_dim), device=device, dtype=torch.bfloat16) * 8
mn = num_tokens
groups_per_row = hidden_dim // group_size
k_num_packed = (groups_per_row + 3) // 4
tma_aligned_mn = ((mn + 3) // 4) * 4
fp8_dtype = torch.float8_e4m3fn
finfo = torch.finfo(fp8_dtype)
# Allocate output_q with the padded mn extent and pre-fill with 0xFF
# so the kernel cannot rely on a clean buffer.
out_q = torch.empty((tma_aligned_mn, hidden_dim), device=device, dtype=fp8_dtype)
out_q.view(torch.uint8).fill_(0xFF)
out_s_packed = torch.empty_strided(
(mn, k_num_packed),
(1, tma_aligned_mn),
device=device,
dtype=torch.int32,
)
torch.ops._C.per_token_group_fp8_quant_packed(
x, out_q, out_s_packed, group_size, 1e-10, finfo.min, finfo.max
)
# Live rows must match the Triton reference.
with patch("vllm.platforms.current_platform.is_cuda", return_value=False):
ref_q, _ = fp8_utils.per_token_group_quant_fp8(x, group_size, use_ue8m0=True)
assert torch.equal(out_q[:mn], ref_q), "Live region mismatch"
# Padded rows must be all-zero; without this, downstream TMA loads would
# see uninitialised data.
if tma_aligned_mn > mn:
padded_bytes = out_q[mn:tma_aligned_mn].view(torch.uint8)
assert padded_bytes.eq(0).all(), (
f"Padded rows [{mn}, {tma_aligned_mn}) not zeroed; "
f"{padded_bytes.ne(0).sum().item()} non-zero bytes"
)
@pytest.mark.parametrize("shape", [(32, 128), (64, 256), (16, 512)])
@pytest.mark.parametrize("group_size", [64, 128])
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
@@ -59,6 +59,34 @@ def test_reload_lifecycle():
assert tensor.__dict__ == materialized_tensor.__dict__
def test_materialize_layer_preserves_non_meta_tensors():
"""Ensure that materialize_layer does not overwrite non meta tensors."""
layer = torch.nn.Linear(2, 3, bias=True)
# Create a non meta bias tensor and meta weight, which can happen with FP8
bias_values = torch.ones(3)
layer.bias.data.copy_(bias_values)
layer.weight = torch.nn.Parameter(layer.weight.data.to("meta"))
assert layer.weight.is_meta
assert not layer.bias.is_meta
# materialize the layer weights after the bias is initialized
info = LayerReloadingInfo(
restore_metadata=({}, {}),
restore_device=torch.device("cpu"),
)
materialize_layer(layer, info)
# Ensure the weight materialized off meta
assert not layer.weight.is_meta
assert layer.weight.device.type == "cpu"
# Ensure that the bias is (still) not meta and values are unchanged
assert not layer.bias.is_meta
assert torch.equal(layer.bias.data, bias_values)
def test_model_cleanup(dist_init, default_vllm_config):
layer = QKVParallelLinear(2, 3, 4)
assert layer.weight.weight_loader.__self__ is layer
@@ -23,11 +23,7 @@ from vllm.model_executor.layers.fused_moe.router.fused_topk_router import (
vllm_topk_sigmoid,
vllm_topk_softmax,
)
from vllm.model_executor.layers.layernorm import (
RMSNorm,
dispatch_rocm_rmsnorm_func,
fused_add_rms_norm,
)
from vllm.model_executor.layers.layernorm import RMSNorm
from vllm.platforms import current_platform
RMS_NORM_SUPPORTED_DTYPES = [torch.float16, torch.bfloat16]
@@ -153,26 +149,3 @@ def test_topk_sigmoid_dispatch(use_rocm_aiter: bool):
assert topk_func == rocm_aiter_ops.topk_sigmoid
else:
assert topk_func == vllm_topk_sigmoid
@pytest.mark.parametrize("add_residual", [False])
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
@pytest.mark.parametrize("use_rocm_aiter", [True, False])
@pytest.mark.skipif(
not current_platform.is_rocm(), reason="AITER is a feature exclusive for ROCm"
)
def test_rms_norm_dispatch(
add_residual: bool, dtype: torch.dtype, use_rocm_aiter: bool
):
rms_norm_func = dispatch_rocm_rmsnorm_func(dtype, use_rocm_aiter)
should_use_rocm_aiter = (
current_platform.is_rocm()
and use_rocm_aiter
and dtype in RMS_NORM_SUPPORTED_DTYPES
)
if should_use_rocm_aiter:
assert rms_norm_func == rocm_aiter_ops.rms_norm2d_with_add
else:
assert rms_norm_func == fused_add_rms_norm
+78 -41
View File
@@ -1,60 +1,97 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import multiprocessing
import types
import pytest
import torch
from vllm.platforms import current_platform
def _load_oink_ops_module():
# Import the module normally (vllm is installed as an editable package in CI).
from vllm import _oink_ops
def _test_oink_availability_impl(
device_capability: tuple[int, int],
has_rmsnorm: bool,
has_fused_add_rms_norm: bool,
expected_available: bool,
expected_fused: bool,
) -> None:
"""Test OINK support detection with mocked state."""
import torch
return _oink_ops
from vllm import platforms
# Mock device capability (class method, override on class)
dc = platforms.interface.DeviceCapability(*device_capability)
platforms.current_platform.__class__.get_device_capability = lambda device_id=0: dc
# Mock oink ops
oink_ops = types.SimpleNamespace()
if has_rmsnorm:
oink_ops.rmsnorm = lambda x, w, eps: x
if has_fused_add_rms_norm:
oink_ops.fused_add_rms_norm = lambda x, residual, w, eps: None
torch.ops.oink = oink_ops
# Now import vllm modules with mocks in place (fresh import with mocked platform)
import vllm.kernels.oink_ops # noqa: F401
from vllm.ir.ops import fused_add_rms_norm, rms_norm
# Verify support checks
assert rms_norm.impls["oink"].supported is expected_available
assert fused_add_rms_norm.impls["oink"].supported is expected_fused
def test_oink_availability_checks(monkeypatch: pytest.MonkeyPatch):
_oink_ops = _load_oink_ops_module()
@pytest.mark.parametrize(
"device_capability,has_rmsnorm,has_fused_add_rms_norm,expected_available,expected_fused",
[
# Case 1: < SM100, ops not supported
((9, 0), True, False, False, False),
# Case 2: CUDA available and SM100, rmsnorm op registered
((10, 0), True, False, True, False),
# Case 3: SM100 with both rmsnorm and fused_add_rms_norm
((10, 0), True, True, True, True),
],
)
@pytest.mark.skipif(not current_platform.is_cuda(), reason="Only test on CUDA")
def test_oink_availability_checks(
device_capability: tuple[int, int],
has_rmsnorm: bool,
has_fused_add_rms_norm: bool,
expected_available: bool,
expected_fused: bool,
):
"""Test OINK support detection with clean import state for each parameter set."""
# Ensure the ops namespace exists and is mutable for tests.
monkeypatch.setattr(
torch.ops,
"oink",
types.SimpleNamespace(rmsnorm=lambda x, w, eps: x),
raising=False,
)
# Case 1: CUDA not available.
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
assert _oink_ops.is_oink_available_for_device(0) is False
# Case 2: CUDA available but < SM100.
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda idx: (9, 0))
assert _oink_ops.is_oink_available_for_device(0) is False
# Case 3: CUDA available and SM100, rmsnorm op registered.
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda idx: (10, 0))
assert _oink_ops.is_oink_available_for_device(0) is True
# fused op presence probe
assert _oink_ops.has_fused_add_rms_norm() is False
monkeypatch.setattr(
torch.ops,
"oink",
types.SimpleNamespace(
rmsnorm=lambda x, w, eps: x,
fused_add_rms_norm=lambda x, residual, w, eps: None,
# Use spawn to run function in fresh process with clean imports
# TODO migrate to spawn utility:
# https://github.com/vllm-project/vllm/issues/41415
ctx = multiprocessing.get_context("spawn")
process = ctx.Process(
target=_test_oink_availability_impl,
args=(
device_capability,
has_rmsnorm,
has_fused_add_rms_norm,
expected_available,
expected_fused,
),
raising=False,
)
assert _oink_ops.has_fused_add_rms_norm() is True
process.start()
process.join()
if process.exitcode != 0:
raise AssertionError(
f"Subprocess test failed with exit code {process.exitcode}"
)
def test_can_view_as_2d_stride_guard():
# Import the helper from the layernorm module.
from vllm.model_executor.layers.layernorm import _can_view_as_2d
# No global import
import torch
# Import the helper from the kernels module.
from vllm.kernels.oink_ops import _can_view_as_2d
x = torch.zeros((2, 3, 4))
assert _can_view_as_2d(x) is True
@@ -805,6 +805,32 @@ VLM_TEST_SETTINGS = {
max_num_seqs=2,
patch_hf_runner=model_utils.molmo_patch_hf_runner,
),
"moondream3": VLMTestInfo(
models=["moondream/moondream3-preview"],
test_type=VLMTestType.IMAGE,
prompt_formatter=identity,
img_idx_to_prompt=lambda idx: "<|endoftext|><image>",
# Common-image coverage here targets query/caption. The native
# detect/point skills are not exposed by vLLM.
single_image_prompts=IMAGE_ASSETS.prompts(
{
"stop_sign": "<vlm_image><|md_reserved_0|>query<|md_reserved_1|>What is this sign?<|md_reserved_2|>", # noqa: E501
"cherry_blossom": (
"<vlm_image><|md_reserved_0|>query<|md_reserved_1|>What season is shown?<|md_reserved_2|>" # noqa: E501
),
}
),
max_model_len=4096,
max_num_seqs=2,
dtype="bfloat16",
hf_processor=model_utils.moondream3_processor,
patch_hf_runner=model_utils.moondream3_patch_hf_runner,
# Single size factor to avoid GPU OOM when running multiple test
# cases sequentially (9B MoE model uses ~18 GiB per instance).
image_size_factors=[(1.0,)],
# Moondream3 is 9B params with MoE, needs significant GPU memory
marks=[large_gpu_mark(min_gb=48)],
),
"ovis1_6-gemma2": VLMTestInfo(
models=["AIDC-AI/Ovis1.6-Gemma2-9B"],
test_type=(VLMTestType.IMAGE, VLMTestType.MULTI_IMAGE),
@@ -0,0 +1,176 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Generation tests for Moondream3 query and caption support."""
import pytest
from tests.models.registry import HF_EXAMPLE_MODELS
from vllm.platforms import current_platform
from ....conftest import IMAGE_ASSETS, ImageTestAssets
from ....utils import large_gpu_mark, multi_gpu_test
MOONDREAM3_MODEL_ID = "moondream/moondream3-preview"
MOONDREAM3_TOKENIZER = "moondream/starmie-v1"
HF_IMAGE_PROMPTS = IMAGE_ASSETS.prompts(
{
"stop_sign": "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>What color is the stop sign?<|md_reserved_2|>", # noqa: E501
"cherry_blossom": "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>What color are the flowers?<|md_reserved_2|>", # noqa: E501
}
)
def make_query_prompt(question: str) -> str:
"""Create a direct-answer query prompt for Moondream3."""
return (
"<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>"
f"{question}<|md_reserved_2|>"
)
def make_caption_prompt(length: str = "normal") -> str:
"""Create a caption prompt for Moondream3."""
return (
"<|endoftext|><image><|md_reserved_0|>"
f"describe<|md_reserved_1|>{length}<|md_reserved_2|>"
)
@multi_gpu_test(num_gpus=2)
@large_gpu_mark(min_gb=80)
def test_tensor_parallel(image_assets: ImageTestAssets):
import gc
from vllm import LLM, SamplingParams
from vllm.distributed.parallel_state import destroy_model_parallel
destroy_model_parallel()
gc.collect()
current_platform.empty_cache()
llm = LLM(
model=MOONDREAM3_MODEL_ID,
tokenizer=MOONDREAM3_TOKENIZER,
trust_remote_code=True,
dtype="bfloat16",
tensor_parallel_size=2,
max_model_len=1024,
enforce_eager=True,
limit_mm_per_prompt={"image": 1},
gpu_memory_utilization=0.45,
)
image = image_assets[0].pil_image
prompt = make_query_prompt("What color is the stop sign?")
try:
outputs = llm.generate(
{"prompt": prompt, "multi_modal_data": {"image": image}},
SamplingParams(max_tokens=20, temperature=0),
)
assert len(outputs) > 0
assert outputs[0].outputs[0].text is not None
finally:
del llm
gc.collect()
current_platform.empty_cache()
@pytest.fixture(scope="module")
def llm():
model_info = HF_EXAMPLE_MODELS.get_hf_info("Moondream3ForCausalLM")
model_info.check_transformers_version(on_fail="skip")
from vllm import LLM
try:
return LLM(
model=MOONDREAM3_MODEL_ID,
tokenizer=MOONDREAM3_TOKENIZER,
trust_remote_code=True,
dtype="bfloat16",
max_model_len=2048,
enforce_eager=True,
limit_mm_per_prompt={"image": 1},
gpu_memory_utilization=0.45,
)
except Exception as exc:
pytest.skip(f"Failed to load {MOONDREAM3_MODEL_ID}: {exc}")
@large_gpu_mark(min_gb=48)
def test_model_loading(llm):
assert llm is not None
@large_gpu_mark(min_gb=48)
def test_query_skill(llm, image_assets: ImageTestAssets):
from vllm import SamplingParams
image = image_assets[0].pil_image
prompt = make_query_prompt("What color is the stop sign?")
outputs = llm.generate(
{"prompt": prompt, "multi_modal_data": {"image": image}},
SamplingParams(max_tokens=50, temperature=0),
)
output_text = outputs[0].outputs[0].text
assert output_text is not None
assert len(output_text) > 0
@large_gpu_mark(min_gb=48)
def test_caption_skill(llm, image_assets: ImageTestAssets):
from vllm import SamplingParams
image = image_assets[1].pil_image
prompt = make_caption_prompt()
outputs = llm.generate(
{"prompt": prompt, "multi_modal_data": {"image": image}},
SamplingParams(max_tokens=100, temperature=0),
)
output_text = outputs[0].outputs[0].text
assert output_text is not None
assert len(output_text) > 0
@large_gpu_mark(min_gb=48)
def test_batched_inference(llm, image_assets: ImageTestAssets):
from vllm import SamplingParams
images = [asset.pil_image for asset in image_assets]
prompts = [
{"prompt": prompt, "multi_modal_data": {"image": img}}
for img, prompt in zip(images, HF_IMAGE_PROMPTS)
]
outputs = llm.generate(prompts, SamplingParams(max_tokens=50, temperature=0))
assert len(outputs) == len(images)
for output in outputs:
assert output.outputs[0].text is not None
assert len(output.outputs[0].text) > 0
@pytest.mark.parametrize("asset_name", ["stop_sign", "cherry_blossom"])
@large_gpu_mark(min_gb=48)
def test_image_assets(llm, image_assets: ImageTestAssets, asset_name: str):
from vllm import SamplingParams
asset_idx = 0 if asset_name == "stop_sign" else 1
image = image_assets[asset_idx].pil_image
prompt = HF_IMAGE_PROMPTS[asset_idx]
outputs = llm.generate(
{"prompt": prompt, "multi_modal_data": {"image": image}},
SamplingParams(max_tokens=50, temperature=0),
)
output_text = outputs[0].outputs[0].text
assert output_text is not None
assert len(output_text) > 0
@@ -3,6 +3,7 @@
import pytest
from vllm.assets.image import ImageAsset
from vllm.multimodal.video import sample_frames_from_video
from ....conftest import VIDEO_ASSETS
@@ -11,6 +12,7 @@ models = ["Qwen/Qwen2.5-VL-3B-Instruct"]
target_dtype = "bfloat16"
VIDEO_PLACEHOLDER = "<|vision_start|><|video_pad|><|vision_end|>"
IMAGE_PLACEHOLDER = "<|vision_start|><|image_pad|><|vision_end|>"
def qwen2_5_vl_chat_template(*query):
@@ -28,6 +30,25 @@ VIDEO_PROMPTS = VIDEO_ASSETS.prompts(
)
WINDOW_ATTN_IMAGE_PROMPT = qwen2_5_vl_chat_template(
IMAGE_PLACEHOLDER,
"Describe the image.",
)
def _window_attention_regression_image():
# image from regression issue: https://github.com/vllm-project/vllm/issues/15122
image = ImageAsset("hato").pil_image
return image.resize((image.width // 2, image.height // 2))
def _encoder_cudagraph_config(*, max_vision_items: int) -> dict:
return {
"cudagraph_mm_encoder": True,
"encoder_cudagraph_max_vision_items_per_batch": max_vision_items,
}
@pytest.mark.core_model
@pytest.mark.parametrize("model", models)
@pytest.mark.parametrize("video_pruning_rate", [0.0, 0.75])
@@ -146,3 +167,77 @@ def test_qwen2_5_vl_evs_batched_videos(
# Ensure the output is a string
assert isinstance(output_text, str)
@pytest.mark.core_model
@pytest.mark.parametrize("model", models)
@pytest.mark.parametrize("dtype", [target_dtype])
@pytest.mark.parametrize("max_tokens", [128])
@pytest.mark.parametrize("use_bytecode_hook", [True, False])
def test_qwen2_5_vl_window_attention_image(
vllm_runner,
model,
dtype: str,
max_tokens: int,
use_bytecode_hook: bool,
monkeypatch,
) -> None:
"""Regression test for Qwen2.5 window-attention image path."""
monkeypatch.setenv("VLLM_USE_BYTECODE_HOOK", "1" if use_bytecode_hook else "0")
prompt = [WINDOW_ATTN_IMAGE_PROMPT]
images = [[_window_attention_regression_image()]]
with vllm_runner(
model,
runner="generate",
max_model_len=4096,
dtype=dtype,
limit_mm_per_prompt={"image": 1},
compilation_config=_encoder_cudagraph_config(max_vision_items=1),
) as vllm_model:
outputs = vllm_model.generate_greedy(prompt, max_tokens, images=images)
assert len(outputs) == 1
output_ids, output_text = outputs[0]
assert len(output_ids) > 0
assert len(output_text) > 0
assert isinstance(output_text, str)
@pytest.mark.core_model
@pytest.mark.parametrize("model", models)
@pytest.mark.parametrize("dtype", [target_dtype])
@pytest.mark.parametrize("max_tokens", [128])
@pytest.mark.parametrize("use_bytecode_hook", [True, False])
def test_qwen2_5_vl_window_attention_image_batch(
vllm_runner,
model,
dtype: str,
max_tokens: int,
use_bytecode_hook: bool,
monkeypatch,
) -> None:
"""Regression test window-attention with a small image batch."""
monkeypatch.setenv("VLLM_USE_BYTECODE_HOOK", "1" if use_bytecode_hook else "0")
image = _window_attention_regression_image()
prompts = [WINDOW_ATTN_IMAGE_PROMPT, WINDOW_ATTN_IMAGE_PROMPT]
images = [[image], [image]]
with vllm_runner(
model,
runner="generate",
max_model_len=4096,
max_num_seqs=2,
dtype=dtype,
limit_mm_per_prompt={"image": 1},
compilation_config=_encoder_cudagraph_config(max_vision_items=2),
) as vllm_model:
outputs = vllm_model.generate_greedy(prompts, max_tokens, images=images)
assert len(outputs) == 2
for output_ids, output_text in outputs:
assert len(output_ids) > 0
assert len(output_text) > 0
assert isinstance(output_text, str)
@@ -54,7 +54,18 @@ MODEL_CONFIGS: dict[str, VitCudagraphTestConfig] = {
needs_video_metadata=True,
marks=[pytest.mark.core_model],
),
# TODO: Add more models below.
"qwen2_5_vl": VitCudagraphTestConfig(
model="Qwen/Qwen2.5-VL-3B-Instruct",
image_prompt=qwen_vl_chat_template(
"<|vision_start|><|image_pad|><|vision_end|>What is in this image?"
),
video_prompt=qwen_vl_chat_template(
"<|vision_start|><|video_pad|><|vision_end|>"
"Describe this video in one sentence."
),
needs_video_metadata=False,
marks=[pytest.mark.core_model],
),
}
@@ -38,6 +38,7 @@ def run_test(
limit_mm_per_prompt: dict[str, int],
vllm_runner_kwargs: dict[str, Any] | None,
hf_model_kwargs: dict[str, Any] | None,
hf_processor: Callable[[str], Any] | None,
patch_hf_runner: Callable[[HfRunner], HfRunner] | None,
runner: RunnerOption = "auto",
distributed_executor_backend: str | None = None,
@@ -116,8 +117,18 @@ def run_test(
)
vllm_outputs_per_mm.append(vllm_output)
hf_runner_kwargs: dict[str, Any] = {}
if model_info.tokenizer:
hf_runner_kwargs["tokenizer_name"] = model_info.tokenizer
if hf_processor is not None:
hf_runner_kwargs["processor"] = hf_processor(model)
hf_model = hf_runner(
model, dtype=dtype, auto_cls=auto_cls, model_kwargs=hf_model_kwargs
model,
dtype=dtype,
auto_cls=auto_cls,
model_kwargs=hf_model_kwargs,
**hf_runner_kwargs,
)
# Some models need to patch things like the model processor, e.g., internvl
@@ -1336,3 +1336,221 @@ def voxtral_patch_hf_runner(hf_model: "HfRunner") -> "HfRunner":
hf_model.get_inputs = patched_get_inputs # type: ignore[method-assign, assignment]
hf_model.model.generate = patched_generate # type: ignore[method-assign]
return hf_model
def moondream3_processor(model: str):
from vllm.transformers_utils.processors.moondream3 import Moondream3Processor
return Moondream3Processor.from_pretrained(model, trust_remote_code=True)
def moondream3_patch_hf_runner(hf_model: HfRunner) -> HfRunner:
"""Patch HfRunner for Moondream3."""
moondream_processor = hf_model.processor
def processor(*args, text="", images=None, **kwargs):
if images is None:
return moondream_processor(text=text, **kwargs)
images_list = [images] if isinstance(images, Image) else images
return moondream_processor(images=images_list, text=text, **kwargs)
hf_model.processor = processor
# Expose the LM head for logprob extraction.
hf_model.model.get_output_embeddings = lambda: hf_model.model.model.text.lm_head
native_model = hf_model.model.model # MoondreamModel instance
from torch.nn import functional as F
from vllm.model_executor.models.moondream3 import reconstruct_from_crops
# Resolve the placeholder tokens from the tokenizer instead of hard-coding.
image_placeholder_ids = moondream_processor.tokenizer.encode(
"<image>", add_special_tokens=False
)
def _normalize_tiling(tilings):
"""Extract (h, w) tuple from various tiling container formats."""
tiling = tilings
if isinstance(tiling, torch.Tensor):
tiling = tuple(tiling.squeeze().tolist())
elif isinstance(tiling, (list, tuple)):
t0 = tiling[0]
if isinstance(t0, torch.Tensor):
tiling = tuple(t0.tolist())
elif isinstance(t0, (list, tuple)):
tiling = tuple(t0)
return tiling
def _encode_vision(pixel_values, tilings):
"""Run preprocessed crops through vision encoder + projection."""
device = native_model.device
dtype = native_model.vision.pos_emb.dtype
config = native_model.config
pv = pixel_values
while pv.dim() > 4:
pv = pv.squeeze(0)
pv = pv.to(device=device, dtype=dtype)
features = native_model._vis_enc(pv)
grid_size = config.vision.crop_size // config.vision.enc_patch_size
global_feat = features[0]
if features.shape[0] > 1 and tilings is not None:
tiling = _normalize_tiling(tilings)
local = features[1:].view(-1, grid_size, grid_size, config.vision.enc_dim)
reconstructed = reconstruct_from_crops(
local,
tiling,
config.vision.overlap_margin,
patch_size=1,
)
else:
reconstructed = global_feat.view(
grid_size, grid_size, config.vision.enc_dim
)
return native_model._vis_proj(global_feat, reconstructed)
def _find_subsequence(seq, subseq):
"""Find start index of subseq in seq, or None."""
n = len(subseq)
for i in range(len(seq) - n + 1):
if seq[i : i + n] == subseq:
return i
return None
def _generate(
self,
input_ids=None,
pixel_values=None,
tilings=None,
attention_mask=None,
**kwargs,
):
max_new_tokens = kwargs.get("max_new_tokens", 128)
return_dict = kwargs.get("return_dict_in_generate", False)
output_hs = kwargs.get("output_hidden_states", False)
if pixel_values is None:
sequences = input_ids
if return_dict:
return types.SimpleNamespace(
sequences=sequences,
hidden_states=() if output_hs else None,
)
return sequences
# Processor may return lists; extract the single element.
if isinstance(pixel_values, (list, tuple)):
pixel_values = pixel_values[0]
if (
isinstance(tilings, (list, tuple))
and tilings
and not isinstance(tilings[0], int)
):
tilings = tilings[0]
hf_model.model._setup_caches()
native_model.use_flex_decoding = False
device = native_model.device
config = native_model.config
with torch.inference_mode():
for block in native_model.text.blocks:
block.kv_cache.k_cache.zero_()
block.kv_cache.v_cache.zero_()
img_emb = _encode_vision(pixel_values, tilings)
bos_emb = F.embedding(
torch.tensor([[config.tokenizer.bos_id]], device=device),
native_model.text.wte,
)
img_input = torch.cat([bos_emb, img_emb.unsqueeze(0)], dim=1)
prefix_len = img_input.size(1)
mask = native_model.attn_mask[:, :, :prefix_len, :]
pos_ids = torch.arange(prefix_len, dtype=torch.long, device=device)
native_model._prefill(img_input, mask, pos_ids, None)
ids = input_ids.squeeze(0).tolist()
img_start = _find_subsequence(ids, image_placeholder_ids)
if img_start is None:
sequences = input_ids
if return_dict:
return types.SimpleNamespace(
sequences=sequences,
hidden_states=() if output_hs else None,
)
return sequences
prompt_tokens = ids[img_start + len(image_placeholder_ids) :]
if not prompt_tokens:
sequences = input_ids
if return_dict:
return types.SimpleNamespace(
sequences=sequences,
hidden_states=() if output_hs else None,
)
return sequences
prompt_tensor = torch.tensor([prompt_tokens], device=device)
prompt_emb = F.embedding(prompt_tensor, native_model.text.wte)
prompt_len = prompt_emb.size(1)
mask = native_model.attn_mask[:, :, prefix_len : prefix_len + prompt_len, :]
pos_ids = torch.arange(
prefix_len,
prefix_len + prompt_len,
dtype=torch.long,
device=device,
)
hidden = native_model._prefill(prompt_emb, mask, pos_ids, None)
pos = prefix_len + prompt_len
hidden_last = native_model.text.post_ln(hidden[:, -1:, :])
logits = native_model.text.lm_head(hidden_last.squeeze(1))
generated = []
all_hidden_states = []
# Record the hidden state that predicted each generated token.
prev_hs = hidden_last
for _ in range(max_new_tokens):
next_token = logits.argmax(dim=-1).item()
if next_token == 0:
break
generated.append(next_token)
if output_hs:
all_hidden_states.append((prev_hs,))
next_emb = F.embedding(
torch.tensor([[next_token]], device=device),
native_model.text.wte,
)
mask = native_model.attn_mask[:, :, pos : pos + 1, :]
pos_ids_step = torch.tensor([pos], dtype=torch.long, device=device)
hidden = native_model._prefill(next_emb, mask, pos_ids_step, None)
hidden_last = native_model.text.post_ln(hidden[:, -1:, :])
prev_hs = hidden_last
logits = native_model.text.lm_head(hidden_last.squeeze(1))
pos += 1
result_ids = ids + generated
sequences = torch.tensor([result_ids], device=device)
if return_dict:
return types.SimpleNamespace(
sequences=sequences,
hidden_states=tuple(all_hidden_states) if output_hs else None,
)
return sequences
hf_model.model.generate = types.MethodType(_generate, hf_model.model)
return hf_model
@@ -133,6 +133,7 @@ class VLMTestInfo(NamedTuple):
# Exposed options for HF runner
hf_model_kwargs: dict[str, Any] | None = None
hf_processor: Callable[[str], Any] | None = None
# Indicates we should explicitly pass the EOS from the tokenizer
use_tokenizer_eos: bool = False
auto_cls: type[_BaseAutoModelClass] = AutoModelForCausalLM
@@ -196,6 +197,7 @@ class VLMTestInfo(NamedTuple):
"comparator": self.comparator,
"get_stop_token_ids": self.get_stop_token_ids,
"hf_model_kwargs": self.hf_model_kwargs,
"hf_processor": self.hf_processor,
"stop_str": self.stop_str,
"patch_hf_runner": self.patch_hf_runner,
}
@@ -12,6 +12,60 @@ from ...utils import build_model_context
GEMMA4_MODEL_ID = "google/gemma-4-E2B-it"
@pytest.mark.parametrize(
"image_width,image_height,max_soft_tokens",
[
# Production repro: a 3x900 image (extreme aspect ratio) made the
# prompt-side estimator return 289 while the HF Gemma 4 image
# processor's vision tower output capped at 280, producing the
# "Attempted to assign 280 multimodal tokens to 289 placeholders"
# mismatch that crashed EngineCore.
(900, 3, 280),
(3, 900, 280),
# Same pathology should hold for the video-frame budget (70 tokens).
(900, 3, 70),
# And for any other supported budget.
(4000, 2, 1120),
],
)
@pytest.mark.parametrize("model_id", [GEMMA4_MODEL_ID])
def test_compute_num_soft_tokens_does_not_exceed_max_soft_tokens(
model_id: str,
image_width: int,
image_height: int,
max_soft_tokens: int,
):
"""Regression for the Gemma 3/4 multimodal crash.
`_compute_num_soft_tokens` must never return a value larger than
`max_soft_tokens`. The HF Gemma 4 image processor clamps its vision
tower output to that value; if the prompt-side estimator returns more,
the prompt has more `image` placeholder tokens than the encoder will
fill, and `_merge_multimodal_embeddings` raises `ValueError` deep in
the model forward.
"""
ctx = build_model_context(
model_id,
mm_processor_kwargs={"do_pan_and_scan": True},
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
num_soft_tokens = processor.info._compute_num_soft_tokens(
image_width=image_width,
image_height=image_height,
max_soft_tokens=max_soft_tokens,
)
assert num_soft_tokens <= max_soft_tokens, (
f"_compute_num_soft_tokens returned {num_soft_tokens} for "
f"image_width={image_width}, image_height={image_height}, "
f"max_soft_tokens={max_soft_tokens} — exceeds the cap that the HF "
f"image processor enforces on its vision tower output. This is "
f"the placeholder/encoder count mismatch that crashes EngineCore."
)
@pytest.mark.parametrize("model_id", [GEMMA4_MODEL_ID])
def test_limit_mm_per_prompt(
image_assets: ImageTestAssets,
@@ -0,0 +1,553 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for Moondream3 multimodal processing.
Includes:
- Processor creation and application tests
- Image tokenization and placeholder expansion tests
- Tiling and cropping logic tests (CPU-based)
- Pixel normalization tests
"""
import numpy as np
import pytest
import torch
from vllm.multimodal import MULTIMODAL_REGISTRY
from ....conftest import ImageTestAssets
from ...utils import build_model_context
MOONDREAM3_MODEL_ID = "moondream/moondream3-preview"
# Expected multimodal prefix: BOS + 729 image tokens.
EXPECTED_IMAGE_TOKENS = 730
# Vision encoder constants
CROP_SIZE = 378
PATCH_SIZE = 14
MAX_CROPS = 12
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_processor_creation(model_id: str):
"""Test that Moondream3 processor can be created."""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
assert processor is not None
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_processor_apply(
image_assets: ImageTestAssets,
model_id: str,
):
"""Test that Moondream3 processor can process inputs.
NOTE: The prompt includes the leading BOS token because Moondream3
pre-fills BOS and image embeddings together.
"""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
prompt = "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>What is this?<|md_reserved_2|>" # noqa: E501
mm_data = {"image": [image_assets[0].pil_image]}
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs={},
)
assert "prompt_token_ids" in processed_inputs
image_placeholders = processed_inputs["mm_placeholders"]["image"]
assert len(image_placeholders) == 1
assert image_placeholders[0].length == EXPECTED_IMAGE_TOKENS
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_processor_pixel_values(
image_assets: ImageTestAssets,
model_id: str,
):
"""Test that pixel values are correctly produced."""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
prompt = "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>What is this?<|md_reserved_2|>" # noqa: E501
mm_data = {"image": [image_assets[0].pil_image]}
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs={},
)
# Check mm_kwargs contains pixel_values
mm_kwargs = processed_inputs.get("mm_kwargs")
assert mm_kwargs is not None
mm_data_result = mm_kwargs.get_data()
assert "pixel_values" in mm_data_result
# Verify pixel_values shape
pixel_values = mm_data_result["pixel_values"]
assert pixel_values.dim() == 5 # [batch, num_crops, C, H, W]
assert pixel_values.shape[2] == 3 # RGB channels
assert pixel_values.shape[3] == 378 # crop height
assert pixel_values.shape[4] == 378 # crop width
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_processor_image_token_expansion(
image_assets: ImageTestAssets,
model_id: str,
):
"""Test that <image> placeholder is expanded to correct number of tokens."""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
prompt = "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>Describe.<|md_reserved_2|>" # noqa: E501
mm_data = {"image": [image_assets[0].pil_image]}
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) == 1
assert image_placeholders[0].length == EXPECTED_IMAGE_TOKENS
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_multi_crop_tiling(
model_id: str,
):
"""Test that large images produce correct multi-crop tiling."""
from PIL import Image
from vllm.transformers_utils.processors.moondream3 import Moondream3Processor
processor = Moondream3Processor.from_pretrained(model_id, trust_remote_code=True)
# Create a large image that requires multiple crops
large_image = Image.new("RGB", (1000, 1000), color="blue")
pixel_values, tiling = processor.preprocess_image(large_image)
# Large images should produce more than 1x1 tiling
assert tiling[0] >= 1 and tiling[1] >= 1
# Check that we have global crop + local crops
expected_crops = tiling[0] * tiling[1] + 1
assert pixel_values.shape[0] == expected_crops
@pytest.mark.parametrize(
"image_size",
[
(500, 500),
(800, 600),
(1920, 1080),
],
)
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_tiling_various_sizes(
image_size: tuple[int, int],
model_id: str,
):
"""Test tiling with various image sizes."""
from PIL import Image
from vllm.transformers_utils.processors.moondream3 import Moondream3Processor
processor = Moondream3Processor.from_pretrained(model_id, trust_remote_code=True)
width, height = image_size
image = Image.new("RGB", (width, height), color="red")
pixel_values, tiling = processor.preprocess_image(image)
# Basic shape checks
assert pixel_values.dim() == 4 # [num_crops, C, H, W]
assert pixel_values.shape[1] == 3 # RGB
assert pixel_values.shape[2] == 378 # crop height
assert pixel_values.shape[3] == 378 # crop width
# Tiling should respect max_crops (12)
assert tiling[0] * tiling[1] <= 12
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_pixel_normalization(
model_id: str,
):
"""Test that pixel values are normalized to [-1, 1] range."""
from PIL import Image
from vllm.transformers_utils.processors.moondream3 import Moondream3Processor
processor = Moondream3Processor.from_pretrained(model_id, trust_remote_code=True)
# Create test image
image = Image.new("RGB", (378, 378), color="green")
pixel_values, _ = processor.preprocess_image(image)
# Normalization: (x - 0.5) / 0.5 = 2*x - 1
# For input [0, 1], output should be [-1, 1]
assert pixel_values.min() >= -1.0
assert pixel_values.max() <= 1.0
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_chat_template_with_image(
image_assets: ImageTestAssets,
model_id: str,
):
"""Test that chat template correctly formats BOS + image + prompt."""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
tokenizer = ctx.tokenizer
# Use the chat template format
prompt = "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>What is this?<|md_reserved_2|>" # noqa: E501
mm_data = {"image": [image_assets[0].pil_image]}
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs={},
)
token_ids = processed_inputs["prompt_token_ids"]
# BOS token (<|endoftext|>) should be token ID 0
bos_token_id = tokenizer.encode("<|endoftext|>", add_special_tokens=False)[0]
assert bos_token_id == 0
# First token should be BOS
assert token_ids[0] == bos_token_id
@pytest.mark.parametrize(
"content",
[
pytest.param(
[
{
"type": "image_url",
"image_url": {"url": "https://example.invalid/image.png"},
},
{"type": "text", "text": "What is in this image?"},
],
id="image-first",
),
pytest.param(
[
{"type": "text", "text": "What is in this image?"},
{
"type": "image_url",
"image_url": {"url": "https://example.invalid/image.png"},
},
],
id="text-first",
),
],
)
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_chat_template_content_list_uses_moondream_image_prefix(
image_assets: ImageTestAssets,
content: list[dict[str, object]],
model_id: str,
):
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
hf_processor = processor.info.get_hf_processor()
prompt = hf_processor.tokenizer.apply_chat_template(
[{"role": "user", "content": content}],
chat_template=hf_processor.chat_template,
tokenize=False,
)
expected_prompt = (
"<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>"
"What is in this image?<|md_reserved_2|>"
)
assert prompt == expected_prompt
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data({"image": [image_assets[0].pil_image]}),
hf_processor_mm_kwargs={},
)
image_placeholders = processed_inputs["mm_placeholders"]["image"]
assert len(image_placeholders) == 1
assert image_placeholders[0].length == EXPECTED_IMAGE_TOKENS
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_bos_token_always_first(
image_assets: ImageTestAssets,
model_id: str,
):
"""Test that BOS token (ID 0) is always at position 0."""
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
# Start with BOS token explicitly
prompt = "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>Describe this image.<|md_reserved_2|>" # noqa: E501
mm_data = {"image": [image_assets[0].pil_image]}
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data(mm_data),
hf_processor_mm_kwargs={},
)
token_ids = processed_inputs["prompt_token_ids"]
# Token ID 0 (<|endoftext|>) should be the first token
assert token_ids[0] == 0, (
f"Expected BOS token (0) at position 0, got {token_ids[0]}"
)
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_processor_with_small_image(
model_id: str,
):
"""Test processor with image smaller than crop size."""
from PIL import Image
from vllm.transformers_utils.processors.moondream3 import Moondream3Processor
processor = Moondream3Processor.from_pretrained(model_id, trust_remote_code=True)
# Small image (smaller than crop size)
small_image = Image.new("RGB", (100, 100), color="yellow")
pixel_values, tiling = processor.preprocess_image(small_image)
# Small images should use 1x1 tiling
assert tiling == (1, 1)
# Should have 2 crops (global + 1 local)
assert pixel_values.shape[0] == 2
@pytest.mark.parametrize(
"image_kind",
[
pytest.param("numpy_hwc", id="numpy-hwc"),
pytest.param("numpy_chw", id="numpy-chw"),
pytest.param("torch_chw", id="torch-chw"),
],
)
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_preprocess_image_accepts_non_pil_inputs(
image_assets: ImageTestAssets,
image_kind: str,
model_id: str,
):
from vllm.transformers_utils.processors.moondream3 import Moondream3Processor
processor = Moondream3Processor.from_pretrained(model_id, trust_remote_code=True)
pil_image = image_assets[0].pil_image.convert("RGB")
hwc_array = np.asarray(pil_image)
expected_pixel_values, expected_tiling = processor.preprocess_image(pil_image)
if image_kind == "numpy_hwc":
image = hwc_array
elif image_kind == "numpy_chw":
image = np.transpose(hwc_array, (2, 0, 1))
else:
image = torch.from_numpy(np.transpose(hwc_array, (2, 0, 1)).copy())
pixel_values, tiling = processor.preprocess_image(image)
assert tiling == expected_tiling
assert pixel_values.shape == expected_pixel_values.shape
assert pixel_values.dtype == torch.bfloat16
assert torch.equal(pixel_values, expected_pixel_values)
@pytest.mark.parametrize("image_kind", ["numpy_chw", "torch_chw"])
@pytest.mark.parametrize("model_id", [MOONDREAM3_MODEL_ID])
def test_processor_apply_accepts_non_pil_image_inputs(
image_assets: ImageTestAssets,
image_kind: str,
model_id: str,
):
ctx = build_model_context(
model_id,
limit_mm_per_prompt={"image": 1},
)
processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
prompt = "<|endoftext|><image><|md_reserved_0|>query<|md_reserved_1|>What is this?<|md_reserved_2|>" # noqa: E501
hwc_array = np.asarray(image_assets[0].pil_image.convert("RGB"))
chw_array = np.transpose(hwc_array, (2, 0, 1)).copy()
image = chw_array if image_kind == "numpy_chw" else torch.from_numpy(chw_array)
processed_inputs = processor(
prompt,
mm_items=processor.info.parse_mm_data({"image": [image]}),
hf_processor_mm_kwargs={},
)
image_placeholders = processed_inputs["mm_placeholders"]["image"]
assert len(image_placeholders) == 1
assert image_placeholders[0].length == EXPECTED_IMAGE_TOKENS
mm_kwargs = processed_inputs["mm_kwargs"].get_data()
assert mm_kwargs["pixel_values"].shape[2:] == (3, 378, 378)
class TestMoondream3TilingLogic:
"""CPU-based tests for Moondream3 tiling selection logic.
These tests validate the select_tiling() function which determines
how images are divided into crops for the vision encoder.
"""
def test_small_image_no_tiling(self):
"""Small images should use 1x1 tiling."""
from vllm.transformers_utils.processors.moondream3 import select_tiling
tiling = select_tiling(
height=300, width=300, crop_size=CROP_SIZE, max_crops=MAX_CROPS
)
assert tiling == (1, 1)
def test_exact_crop_size(self):
"""Image exactly at crop size should use 1x1."""
from vllm.transformers_utils.processors.moondream3 import select_tiling
tiling = select_tiling(
height=CROP_SIZE, width=CROP_SIZE, crop_size=CROP_SIZE, max_crops=MAX_CROPS
)
assert tiling == (1, 1)
def test_large_square_image(self):
"""Large square image should use multiple tiles."""
from vllm.transformers_utils.processors.moondream3 import select_tiling
tiling = select_tiling(
height=800, width=800, crop_size=CROP_SIZE, max_crops=MAX_CROPS
)
h_tiles, w_tiles = tiling
assert h_tiles >= 2
assert w_tiles >= 2
assert h_tiles * w_tiles <= MAX_CROPS
def test_wide_image(self):
"""Wide image should have more width tiles."""
from vllm.transformers_utils.processors.moondream3 import select_tiling
tiling = select_tiling(
height=400, width=1200, crop_size=CROP_SIZE, max_crops=MAX_CROPS
)
h_tiles, w_tiles = tiling
assert w_tiles >= h_tiles
def test_tall_image(self):
"""Tall image should have more height tiles."""
from vllm.transformers_utils.processors.moondream3 import select_tiling
tiling = select_tiling(
height=1200, width=400, crop_size=CROP_SIZE, max_crops=MAX_CROPS
)
h_tiles, w_tiles = tiling
assert h_tiles >= w_tiles
def test_respects_max_crops(self):
"""Tiling should not exceed max_crops."""
from vllm.transformers_utils.processors.moondream3 import select_tiling
tiling = select_tiling(
height=2000, width=2000, crop_size=CROP_SIZE, max_crops=4
)
h_tiles, w_tiles = tiling
assert h_tiles * w_tiles <= 4
class TestMoondream3VisionShapes:
"""CPU-based tests for vision encoder expected shapes.
These tests verify the mathematical relationships between
crop size, patch size, and token counts.
"""
def test_expected_patch_count(self):
"""Test 378/14 = 27 patches per side, 729 total."""
patches_per_side = CROP_SIZE // PATCH_SIZE
total_patches = patches_per_side**2
assert patches_per_side == 27
assert total_patches == EXPECTED_IMAGE_TOKENS - 1
def test_patch_embedding_input_dim(self):
"""Test patch embedding input dimension."""
channels = 3
input_dim = PATCH_SIZE * PATCH_SIZE * channels
assert input_dim == 14 * 14 * 3
assert input_dim == 588
class TestMoondream3TauAttention:
"""CPU-based tests for tau attention scaling components.
These tests validate the tau attention formula used in Moondream3:
- Token-based: tok_q = tanh(gelu(qkv) @ tau_wq.T)
- Position-based: tau_pos = 1 + (sigmoid(alpha * log(pos+1)) - 0.5)
"""
def test_tau_position_range(self):
"""Test tau position scaling produces values in valid range."""
num_heads = 32
seq_len = 100
tau_alpha = torch.randn(num_heads)
positions = torch.arange(seq_len)
pos_float = (positions.float() + 1.0).clamp(min=1e-6)
pos_log = pos_float.log()
tau_pos = 1.0 + (torch.sigmoid(tau_alpha[:, None] * pos_log[None, :]) - 0.5)
assert tau_pos.shape == (num_heads, seq_len)
# tau_pos should be between 0.5 and 1.5
assert tau_pos.min() >= 0.5
assert tau_pos.max() <= 1.5
def test_tau_token_output_range(self):
"""Test tau token scaling output is bounded by tanh."""
import torch.nn.functional as F
seq_len = 100
qkv_dim = 6144 # 2048 * 3
num_heads = 32
qkv = torch.randn(seq_len, qkv_dim)
tau_wq = torch.randn(num_heads, qkv_dim)
tok_feat = F.gelu(qkv)
tok_q = torch.tanh(tok_feat @ tau_wq.t())
assert tok_q.shape == (seq_len, num_heads)
# tanh output is bounded by [-1, 1]
assert tok_q.min() >= -1.0
assert tok_q.max() <= 1.0
@@ -0,0 +1,114 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
from vllm.model_executor.models.nano_nemotron_vl import NemotronH_Nano_VL_V2
class _TextOnlyMultiModalConfig:
def get_limit_per_prompt(self, modality: str) -> int:
return 0
class _ImageOnlyMultiModalConfig:
def get_limit_per_prompt(self, modality: str) -> int:
return 1 if modality == "image" else 0
class _ModelConfig:
multimodal_config = _TextOnlyMultiModalConfig()
class _ImageOnlyModelConfig:
multimodal_config = _ImageOnlyMultiModalConfig()
class _LanguageModel:
def __init__(self) -> None:
self.loaded_weights: list[tuple[str, object]] = []
def load_weights(self, weights):
self.loaded_weights = list(weights)
class _MissingMultiModalModule:
def named_parameters(self):
raise AssertionError("multimodal weights should not be inspected")
def load_weights(self, weights):
raise AssertionError("multimodal weights should not be loaded")
class _AdapterModule:
def named_parameters(self):
return []
class _VisionModel:
def __init__(self) -> None:
self.loaded_weights: list[tuple[str, object]] = []
def load_weights(self, weights):
self.loaded_weights = list(weights)
def test_nano_nemotron_vl_skips_multimodal_weights_in_text_only_mode():
model = object.__new__(NemotronH_Nano_VL_V2)
language_model = _LanguageModel()
object.__setattr__(model, "model_config", _ModelConfig())
object.__setattr__(model, "language_model", language_model)
object.__setattr__(model, "mlp1", _AdapterModule())
object.__setattr__(model, "vision_model", _MissingMultiModalModule())
object.__setattr__(model, "sound_encoder", None)
language_weight = object()
model.load_weights(
[
("language_model.layers.0.weight", language_weight),
("mlp1.0.weight", object()),
("vision_model.radio_model.encoder.weight", object()),
("sound_encoder.encoder.weight", object()),
]
)
assert language_model.loaded_weights == [("layers.0.weight", language_weight)]
def test_nano_nemotron_vl_loads_vision_weights_without_sound_encoder():
model = object.__new__(NemotronH_Nano_VL_V2)
language_model = _LanguageModel()
vision_model = _VisionModel()
object.__setattr__(model, "model_config", _ImageOnlyModelConfig())
object.__setattr__(model, "language_model", language_model)
object.__setattr__(model, "mlp1", _AdapterModule())
object.__setattr__(model, "vision_model", vision_model)
object.__setattr__(model, "sound_encoder", None)
language_weight = object()
vision_weight = object()
model.load_weights(
[
("language_model.layers.0.weight", language_weight),
("vision_model.radio_model.encoder.weight", vision_weight),
]
)
assert language_model.loaded_weights == [("layers.0.weight", language_weight)]
assert vision_model.loaded_weights == [
("radio_model.encoder.weight", vision_weight)
]
def test_nano_nemotron_vl_requires_sound_encoder_for_sound_weights():
model = object.__new__(NemotronH_Nano_VL_V2)
language_model = _LanguageModel()
vision_model = _VisionModel()
object.__setattr__(model, "model_config", _ImageOnlyModelConfig())
object.__setattr__(model, "language_model", language_model)
object.__setattr__(model, "mlp1", _AdapterModule())
object.__setattr__(model, "vision_model", vision_model)
object.__setattr__(model, "sound_encoder", None)
with pytest.raises(AssertionError):
model.load_weights([("sound_encoder.encoder.weight", object())])
+20 -48
View File
@@ -946,13 +946,6 @@ _MULTIMODAL_EXAMPLE_MODELS = {
"HCXVisionForCausalLM": _HfExamplesInfo(
"naver-hyperclovax/HyperCLOVAX-SEED-Vision-Instruct-3B",
trust_remote_code=True,
max_transformers_version="4.57",
transformers_version_reason={
"vllm": (
"Custom config cannot be loaded with Transformers "
"v5 because `text_config` is not always set"
)
},
),
"HCXVisionV2ForCausalLM": _HfExamplesInfo(
"naver-hyperclovax/HyperCLOVAX-SEED-Think-32B",
@@ -1122,6 +1115,16 @@ _MULTIMODAL_EXAMPLE_MODELS = {
extras={"olmo": "allenai/Molmo-7B-O-0924"},
trust_remote_code=True,
),
"Moondream3ForCausalLM": _HfExamplesInfo(
"moondream/moondream3-preview",
tokenizer="moondream/starmie-v1",
trust_remote_code=True,
),
"HfMoondream": _HfExamplesInfo(
"moondream/moondream3-preview",
tokenizer="moondream/starmie-v1",
trust_remote_code=True,
),
"Molmo2ForConditionalGeneration": _HfExamplesInfo(
"allenai/Molmo2-8B",
extras={"olmo": "allenai/Molmo2-O-7B"},
@@ -1138,30 +1141,17 @@ _MULTIMODAL_EXAMPLE_MODELS = {
"NemotronH_Nano_VL_V2": _HfExamplesInfo(
"nvidia/NVIDIA-Nemotron-Nano-12B-v2-VL-BF16",
max_model_len=4096,
# NemotronH layers are constructed via `hybrid_override_pattern`:
# NemotronH layers are constructed via `hybrid_override_pattern`
use_original_num_layers=True,
hf_overrides={
"vision_config": PretrainedConfig(
args={
"min_num_patches": 1, # Trigger image dynamic res
"max_num_patches": 12,
"model": "vit_huge_patch16_224",
},
# Trigger conv3d:
video_temporal_patch_size=2,
),
"text_config": {
"num_hidden_layers": 2,
"hybrid_override_pattern": "M*",
},
"text_config": {"num_hidden_layers": 2, "hybrid_override_pattern": "M*"},
},
trust_remote_code=True,
),
# NemotronH_Nano_Omni_Reasoning_V3 is an alias for NemotronH_Nano_VL_V2
# Use the same registry test as NemotronH_Nano_VL_V2 above
"NemotronH_Nano_Omni_Reasoning_V3": _HfExamplesInfo(
"nvidia/NVIDIA-Nemotron-Nano-12B-v2-VL-BF16",
"nvidia/Nemotron-3-Nano-Omni-30B-A3B-Reasoning-BF16",
max_model_len=4096,
# NemotronH layers are constructed via `hybrid_override_pattern`
use_original_num_layers=True,
hf_overrides={
"vision_config": PretrainedConfig(
@@ -1171,35 +1161,17 @@ _MULTIMODAL_EXAMPLE_MODELS = {
"model": "vit_huge_patch16_224",
},
video_temporal_patch_size=2,
# TODO(nhaber): This is `true` in the official `config.json`,
# but this causes a processor exception in the tests due to a known bug
# with mixed-resolution video when `true`. To be resolved.
video_maintain_aspect_ratio=False,
),
"text_config": {
"num_hidden_layers": 2,
"hybrid_override_pattern": "M*",
},
"text_config": {"num_hidden_layers": 2, "hybrid_override_pattern": "M*"},
},
trust_remote_code=True,
),
# NemotronH_Super_Omni_Reasoning_V3 is an alias for NemotronH_Nano_VL_V2 as well
# Use the same registry test as NemotronH_Nano_VL_V2 above
"NemotronH_Super_Omni_Reasoning_V3": _HfExamplesInfo(
"nvidia/NVIDIA-Nemotron-Nano-12B-v2-VL-BF16",
max_model_len=4096,
use_original_num_layers=True,
hf_overrides={
"vision_config": PretrainedConfig(
args={
"min_num_patches": 1,
"max_num_patches": 12,
"model": "vit_huge_patch16_224",
},
video_temporal_patch_size=2,
),
"text_config": {
"num_hidden_layers": 2,
"hybrid_override_pattern": "M*",
},
},
trust_remote_code=True,
"nvidia/Nemotron-3-Nano-Omni-30B-A3B-Reasoning-BF16", is_available_online=False
),
"OpenCUAForConditionalGeneration": _HfExamplesInfo(
"xlangai/OpenCUA-7B",
+14
View File
@@ -523,6 +523,17 @@ def dummy_hf_overrides(
text_config.update(update_dict)
# Update n_layers and moe configs for Moondream3 model
if model_arch in ("Moondream3ForCausalLM", "HfMoondream"):
text_config.update(
{
"n_layers": num_hidden_layers,
"moe_num_experts": num_experts,
"moe_experts_per_token": 2,
"moe_start_layer": num_hidden_layers,
}
)
if hasattr(hf_config, "vision_config"):
hf_config.vision_config.update(
{
@@ -531,6 +542,9 @@ def dummy_hf_overrides(
}
)
if model_arch in ("Moondream3ForCausalLM", "HfMoondream"):
hf_config.vision_config.update({"enc_n_layers": 1})
# e.g.: ibm-granite/granite-speech-3.3-2b
if hasattr(hf_config, "encoder_config"):
hf_config.encoder_config.update(
+1
View File
@@ -70,4 +70,5 @@ def test_cpu_offload_compressed_tensors(monkeypatch):
["--enforce_eager"],
["--enforce_eager", "--cpu-offload-gb", "1"],
max_wait_seconds=480,
include_seeded_sampling=False,
)
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from unittest.mock import MagicMock
import pytest
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
@@ -12,6 +14,20 @@ from vllm.tokenizers import get_tokenizer
REASONING_MODEL_NAME = "moonshotai/Kimi-K2.5"
@pytest.fixture
def mock_kimi_k2_tokenizer():
tokenizer = MagicMock()
tokenizer.get_vocab.return_value = {
"<think>": 100,
"</think>": 101,
"<|tool_calls_section_begin|>": 200,
"<|tool_calls_section_end|>": 201,
"<|tool_call_begin|>": 202,
"<|tool_call_end|>": 203,
}
return tokenizer
@pytest.fixture(scope="module")
def kimi_k2_tokenizer():
return get_tokenizer(tokenizer_name=REASONING_MODEL_NAME, trust_remote_code=True)
@@ -153,3 +169,50 @@ def test_streaming_tool_section_ends_reasoning(kimi_k2_tokenizer):
)
assert isinstance(result, DeltaMessage)
assert result.content == "<|tool_calls_section_begin|>"
def test_streaming_end_token_id_buffered(mock_kimi_k2_tokenizer):
"""When stop sequences buffer text, </think> ID arrives before its text.
The token ID is present in delta_token_ids but the actual string is not
yet in delta_text (still buffered). The parser must return None to wait
for the next delta, instead of calling find() which returns -1 and
silently corrupting the text split.
"""
parser = KimiK2ReasoningParser(mock_kimi_k2_tokenizer)
think_id = parser._start_token_id
end_think_id = parser._end_token_id
# Simulate: </think> ID arrived but text not yet flushed.
# Two token IDs in delta to bypass the single-special-token guard.
result = parser.extract_reasoning_streaming(
previous_text="some reasoning",
current_text="some reasoning extra",
delta_text="extra", # </think> text not yet flushed
previous_token_ids=[think_id],
current_token_ids=[think_id, end_think_id, 999],
delta_token_ids=[end_think_id, 999],
)
assert result is None
def test_streaming_tool_section_id_buffered(mock_kimi_k2_tokenizer):
"""When stop sequences buffer text, tool section start ID arrives before its text.
Same buffering scenario as above but for <|tool_calls_section_begin|>.
Without the guard, find() returns -1 and delta_text[:tool_index] silently
drops the last character of reasoning.
"""
parser = KimiK2ReasoningParser(mock_kimi_k2_tokenizer)
think_id = parser._start_token_id
tool_begin_id = parser._tool_section_start_token_id
result = parser.extract_reasoning_streaming(
previous_text="some reasoning",
current_text="some reasoning extra",
delta_text="extra", # tool section text not yet flushed
previous_token_ids=[think_id],
current_token_ids=[think_id, tool_begin_id, 999],
delta_token_ids=[tool_begin_id, 999],
)
assert result is None
@@ -0,0 +1,576 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Offline unit tests for `prompt_embeds` chat-completion content parts."""
from __future__ import annotations
import inspect
import io
from typing import Final
from unittest import mock
import pybase64 as base64
import pytest
import regex as re
import torch
from transformers import AutoTokenizer
from vllm.entrypoints.chat_utils import (
_ENABLE_PROMPT_EMBEDS_ERROR,
_PROMPT_EMBEDS_MISSING_DATA_ERROR,
_RESERVED_PLACEHOLDER_IN_TEXT_ERROR,
MM_PARSER_MAP,
MODALITY_PLACEHOLDERS_MAP,
PROMPT_EMBEDS_PLACEHOLDER_TOKEN,
parse_chat_messages,
parse_chat_messages_async,
)
from vllm.renderers.hf import (
_PROMPT_EMBEDS_PLACEHOLDER_SPAN_MISMATCH_ERROR,
_build_mixed_prompt_embeds,
_build_prompt_embeds_positions,
_build_prompt_embeds_updates,
_ensure_prompt_embeds_placeholder_token,
_expand_prompt_embeds_placeholders,
)
# Cover distinct tokenizer families:
# GPT2TokenizerFast (BPE, OpenAI-style)
# Qwen2TokenizerFast (SentencePiece BPE variant)
# BertTokenizerFast (WordPiece)
TOKENIZER_IDS: Final[list[str]] = [
"gpt2",
"Qwen/Qwen2.5-1.5B-Instruct",
"bert-base-uncased",
]
@pytest.fixture(params=TOKENIZER_IDS, ids=TOKENIZER_IDS)
def tokenizer(request):
"""A fresh tokenizer instance per tokenizer family."""
return AutoTokenizer.from_pretrained(request.param)
# Minimal chat template that works with any tokenizer. Iterates
# `message.content` as either a string or a list of dicts (openai format).
_SIMPLE_CHAT_TEMPLATE: Final[str] = (
"{% for m in messages %}"
"{% if m['content'] is string %}{{m['content']}}"
"{% else %}{% for p in m['content'] %}{{p['text']}}{% endfor %}"
"{% endif %}\n{% endfor %}"
)
async def _maybe_await(fn, *args, **kwargs):
"""Call *fn* and `await` the result if it's a coroutine."""
result = fn(*args, **kwargs)
if inspect.iscoroutine(result):
result = await result
return result
# Parametrize over sync / async parse paths so every end-to-end test
# exercises both.
_PARSE_FUNCTIONS = [parse_chat_messages, parse_chat_messages_async]
@pytest.fixture(params=_PARSE_FUNCTIONS, ids=["sync", "async"])
def parse_fn(request):
"""Either the sync or async `parse_chat_messages` callable."""
return request.param
def _encode_tensor(t: torch.Tensor) -> str:
buf = io.BytesIO()
torch.save(t, buf)
return base64.b64encode(buf.getvalue()).decode("utf-8")
_MOCK_HIDDEN_SIZE: Final[int] = 8
_MOCK_DTYPE: Final[torch.dtype] = torch.float32
def _make_mock_model_config(*, enable_prompt_embeds: bool = True) -> mock.MagicMock:
mc = mock.MagicMock()
mc.enable_prompt_embeds = enable_prompt_embeds
mc.multimodal_config = None
mc.allowed_local_media_path = None
mc.allowed_media_domains = None
# Test text-only code path in `MultiModalItemTracker.resolve_items`.
mc.is_multimodal_model = False
# `safe_load_prompt_embeds` pins each tensor to the model's hidden_size
# and dtype, so the mock must return concrete values.
mc.get_hidden_size.return_value = _MOCK_HIDDEN_SIZE
mc.dtype = _MOCK_DTYPE
return mc
def test_prompt_embeds_keys_registered():
assert "prompt_embeds" in MODALITY_PLACEHOLDERS_MAP
assert MODALITY_PLACEHOLDERS_MAP["prompt_embeds"] == "<##PROMPT_EMBEDS##>"
assert "prompt_embeds" in MM_PARSER_MAP
def test_ensure_placeholder_token_is_single_token_and_idempotent(tokenizer):
"""Ensure the placeholder token is a single token and that multiple calls to
"ensure" are idempotent, across all tokenizer families."""
tid1 = _ensure_prompt_embeds_placeholder_token(tokenizer)
tid2 = _ensure_prompt_embeds_placeholder_token(tokenizer)
assert tid1 == tid2
ids = tokenizer.encode(PROMPT_EMBEDS_PLACEHOLDER_TOKEN, add_special_tokens=False)
assert ids == [tid1]
# Repeating it in a string N times must produce exactly that many tokens.
N = 5
ids_rep = tokenizer.encode(
PROMPT_EMBEDS_PLACEHOLDER_TOKEN * N, add_special_tokens=False
)
assert ids_rep == [tid1] * N
def test_parse_chat_messages_openai_format():
NUM_TOKENS = 3
t = torch.randn(NUM_TOKENS, _MOCK_HIDDEN_SIZE, dtype=_MOCK_DTYPE)
b64 = _encode_tensor(t)
mc = _make_mock_model_config()
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Hello "},
{"type": "prompt_embeds", "data": b64},
{"type": "text", "text": " world"},
],
}
]
conv, mm_data, _ = parse_chat_messages(
messages,
mc,
content_format="openai",
)
# The middle content part is rewritten to a single placeholder-token
# sentinel.
texts = [p["text"] for p in conv[0]["content"]]
assert texts == [
"Hello ",
PROMPT_EMBEDS_PLACEHOLDER_TOKEN,
" world",
]
assert mm_data is not None and "prompt_embeds" in mm_data
assert torch.equal(mm_data["prompt_embeds"][0], t)
# Each layout entry is one content part:
# ("text", "A") -> {"type": "text", "text": "A"}
# ("embed", N) -> {"type": "prompt_embeds", "data": <base64 of (N, H) tensor>}
@pytest.mark.parametrize(
"layout",
[
# Case: Single embed only.
[("embed", 2)],
# Case: Embed at the start of the message.
[("embed", 3), ("text", "B")],
# Case: Embed at the end of the message.
[("text", "A"), ("embed", 1)],
# Case: Embed sandwiched between text spans.
[("text", "A"), ("embed", 2), ("text", "B")],
# Case: Multiple embeds with text in between.
[("text", "A"), ("embed", 2), ("text", "B"), ("embed", 3)],
# Case: Adjacent embeds with no separating text.
[("embed", 1), ("embed", 2)],
# Case: Multiple text spans before a trailing embed.
[("text", "A"), ("text", "B"), ("embed", 1)],
# Case: Long-ish run mixing both kinds.
[
("text", "head"),
("embed", 4),
("text", "mid"),
("embed", 1),
("embed", 2),
("text", "tail"),
],
],
ids=[
"single-embed",
"embed-then-text",
"text-then-embed",
"text-embed-text",
"text-embed-text-embed",
"adjacent-embeds",
"text-text-embed",
"long-mixed-run",
],
)
@pytest.mark.parametrize(
"interleave_mm_strings",
# `None`: text-only path where `multimodal_config` is absent.
# `False`: non-interleave multimodal path (the common default).
# `True`: sentinel-substitution interleave path.
# All three must preserve the request ordering of prompt_embeds
# relative to surrounding text because prompt_embeds are spliced at the
# token offset during rendering.
[None, False, True],
ids=["text-only", "interleave-off", "interleave-on"],
)
def test_parse_chat_messages_string_format_preserves_position(
layout, interleave_mm_strings
):
mc = _make_mock_model_config()
if interleave_mm_strings is not None:
mm_cfg = mock.MagicMock()
mm_cfg.interleave_mm_strings = interleave_mm_strings
mc.multimodal_config = mm_cfg
content: list[dict] = []
expected_parts: list[str] = []
expected_embeds: list[torch.Tensor] = []
for kind, value in layout:
if kind == "text":
content.append({"type": "text", "text": value})
expected_parts.append(value)
else: # prompt embeds
num_tokens = value
t = torch.randn(num_tokens, _MOCK_HIDDEN_SIZE, dtype=_MOCK_DTYPE)
expected_embeds.append(t)
content.append({"type": "prompt_embeds", "data": _encode_tensor(t)})
# Parser emits ONE sentinel per part.
expected_parts.append(PROMPT_EMBEDS_PLACEHOLDER_TOKEN)
messages = [{"role": "user", "content": content}]
conv, mm_data, _ = parse_chat_messages(
messages,
mc,
content_format="string",
)
assert conv[0]["content"] == "\n".join(expected_parts)
assert mm_data is not None and "prompt_embeds" in mm_data
assert len(mm_data["prompt_embeds"]) == len(expected_embeds)
for got, want in zip(mm_data["prompt_embeds"], expected_embeds, strict=True):
assert torch.equal(got, want)
def test_parse_chat_messages_requires_flag():
t = torch.randn(2, 4)
b64 = _encode_tensor(t)
mc = _make_mock_model_config(enable_prompt_embeds=False)
messages = [
{
"role": "user",
"content": [{"type": "prompt_embeds", "data": b64}],
}
]
with pytest.raises(ValueError, match=_ENABLE_PROMPT_EMBEDS_ERROR):
parse_chat_messages(
messages,
mc,
content_format="openai",
)
def test_parse_chat_messages_rejects_missing_data():
# `data` is marked `Required` on `ChatCompletionContentPartPromptEmbedsParam`;
# malformed requests without `data` must surface a clear validation error
# rather than being silently dropped.
mc = _make_mock_model_config()
messages = [
{
"role": "user",
"content": [{"type": "prompt_embeds"}], # no `data`
}
]
with pytest.raises(ValueError, match=_PROMPT_EMBEDS_MISSING_DATA_ERROR):
parse_chat_messages(
messages,
mc,
content_format="openai",
)
# Reserved placeholder guard: when `enable_prompt_embeds=True` the tokenizer is
# mutated to make `<prompt_embeds>` a single unsplittable token. Any user text
# containing that literal sequence would tokenize to the same sentinel ID and
# be mistaken for a splice point, so we reject it at parse time.
_PLACEHOLDER_ERROR_PATTERN: Final[str] = re.sub(
r"\\{[^}]*\\}", ".*", re.escape(_RESERVED_PLACEHOLDER_IN_TEXT_ERROR)
)
@pytest.mark.parametrize(
"content",
[
# Case: Top-level string content (wrapped as a single text part).
f"hello {PROMPT_EMBEDS_PLACEHOLDER_TOKEN} world",
# Case: List with a typed text part containing the placeholder.
[{"type": "text", "text": f"leading {PROMPT_EMBEDS_PLACEHOLDER_TOKEN}"}],
# Case: List with a plain-string part (no wrapping dict).
[f"raw string {PROMPT_EMBEDS_PLACEHOLDER_TOKEN}"],
],
ids=["top-level-string", "typed-text-part", "plain-string-part"],
)
def test_parse_chat_messages_rejects_placeholder_in_user_text(content):
mc = _make_mock_model_config() # enable_prompt_embeds=True by default
messages = [{"role": "user", "content": content}]
with pytest.raises(ValueError, match=_PLACEHOLDER_ERROR_PATTERN):
parse_chat_messages(messages, mc, content_format="openai")
def test_parse_chat_messages_allows_placeholder_in_text_when_feature_disabled():
# When `enable_prompt_embeds=False` the tokenizer is never mutated, so the
# literal `<prompt_embeds>` is just ordinary text and must pass through.
mc = _make_mock_model_config(enable_prompt_embeds=False)
messages = [
{
"role": "user",
"content": f"benign mention of {PROMPT_EMBEDS_PLACEHOLDER_TOKEN} here",
}
]
conv, mm_data, _ = parse_chat_messages(messages, mc, content_format="openai")
assert mm_data is None or "prompt_embeds" not in mm_data
# Text reaches the rendered conversation unchanged.
texts = [p["text"] for p in conv[0]["content"]]
assert PROMPT_EMBEDS_PLACEHOLDER_TOKEN in "".join(texts)
# Token-stream spec: ints are regular token IDs, tuples `(N,)` expand to
# a placeholder span of length N (creates corresponding `(N, H)` tensor).
# `expected` lists the `(start_idx, length)` pairs that
# `_build_prompt_embeds_positions` should return.
@pytest.mark.parametrize(
"stream, expected",
[
# Case: Single run in the middle.
([10, 20, (3,), 30], [(2, 3)]),
# Case: Single run at the start.
([(2,), 10, 20], [(0, 2)]),
# Case: Single run at the end.
([10, 20, (4,)], [(2, 4)]),
# Case: Two runs with tokens between.
([1, (2,), 2, 3, (3,), 4], [(1, 2), (5, 3)]),
# Case: Adjacent runs (no separating tokens).
([(1,), (2,)], [(0, 1), (1, 2)]),
# Case: Three runs.
([5, (2,), 6, (1,), 7, (3,), 8], [(1, 2), (4, 1), (6, 3)]),
],
ids=[
"single-middle",
"single-start",
"single-end",
"two-runs-separated",
"two-runs-adjacent",
"three-runs",
],
)
def test_build_positions(tokenizer, stream, expected):
H = 4
tid = _ensure_prompt_embeds_placeholder_token(tokenizer)
tensors: list[torch.Tensor] = []
token_ids: list[int] = []
for item in stream:
if isinstance(item, tuple):
length = item[0]
tensors.append(torch.randn(length, H))
token_ids.extend([tid] * length)
else:
token_ids.append(item)
mm_updates = _build_prompt_embeds_updates(tensors, tid)
positions = _build_prompt_embeds_positions(token_ids, len(tensors), mm_updates)
assert positions == expected
def test_build_positions_length_mismatch(tokenizer):
N1, H1 = 2, 4
N2, H2 = 3, 4
tid = _ensure_prompt_embeds_placeholder_token(tokenizer)
# 2 tensors expected but only a single placeholder run in the token
# stream (simulating dropping the second one).
tensors = [torch.randn(N1, H1), torch.randn(N2, H2)]
token_ids = [1, tid, tid, 2, 3]
mm_updates = _build_prompt_embeds_updates(tensors, tid)
# The error constant is a `str.format` template, escape it and turn
# the `{field}` placeholders into `.*` so it matches any substitution.
pattern = re.sub(
r"\\{[^}]*\\}", ".*", re.escape(_PROMPT_EMBEDS_PLACEHOLDER_SPAN_MISMATCH_ERROR)
)
with pytest.raises(ValueError, match=pattern):
_build_prompt_embeds_positions(token_ids, len(tensors), mm_updates)
# ints = regular token IDs (any value)
# (N,) = embed span of length N
@pytest.mark.parametrize(
"stream",
[
[10, 20, (3,), 30],
[(2,), 10, 20],
[10, 20, (4,)],
[1, (2,), 2, 3, (3,), 4],
[(1,), (2,)],
[5, (2,), 6, (1,), 7, (3,), 8],
],
ids=[
"single-middle",
"single-start",
"single-end",
"two-spans-separated",
"two-spans-adjacent",
"three-spans",
],
)
def test_build_mixed_prompt_embeds(stream):
H = 8
_PLACEHOLDER = 0 # sentinel for embed positions in token_ids
tensors: list[torch.Tensor] = []
token_ids: list[int] = []
positions: list[tuple[int, int]] = []
cursor = 0
for item in stream:
if isinstance(item, tuple):
length = item[0]
tensors.append(torch.randn(length, H))
positions.append((cursor, length))
token_ids.extend([_PLACEHOLDER] * length)
cursor += length
else:
token_ids.append(item)
cursor += 1
embeds, mask = _build_mixed_prompt_embeds(token_ids, tensors, positions)
assert embeds.shape == (len(token_ids), H)
assert len(mask) == len(token_ids)
# Mask: False exactly at embed positions, True everywhere else.
expected_mask = torch.ones(len(token_ids), dtype=torch.bool)
for start, length in positions:
expected_mask[start : start + length] = False
assert mask == expected_mask.tolist()
# Embed rows match input tensors at the right positions.
for tensor, (start, length) in zip(tensors, positions):
assert torch.equal(embeds[start : start + length], tensor)
# Non-embed positions remain zero-filled.
assert torch.all(embeds[expected_mask] == 0)
# End-to-end tests: each runs both sync and async parse paths via the
# `parse_fn` fixture.
@pytest.mark.asyncio
@pytest.mark.parametrize("role", ["user", "system"])
async def test_end_to_end_expand_and_build(tokenizer, parse_fn, role):
"""Full renderer pipeline: parse -> chat template -> expand -> locate
-> build mixed prompt, across tokenizers, roles, and sync/async."""
tokenizer.chat_template = _SIMPLE_CHAT_TEMPLATE
tid = _ensure_prompt_embeds_placeholder_token(tokenizer)
LEN_A, LEN_B = 3, 2
t_a = torch.randn(LEN_A, _MOCK_HIDDEN_SIZE, dtype=_MOCK_DTYPE)
t_b = torch.randn(LEN_B, _MOCK_HIDDEN_SIZE, dtype=_MOCK_DTYPE)
NUM_TENSORS = 2
mc = _make_mock_model_config()
messages = [
{
"role": role,
"content": [
{"type": "text", "text": "Hello "},
{"type": "prompt_embeds", "data": _encode_tensor(t_a)},
{"type": "text", "text": " world "},
{"type": "prompt_embeds", "data": _encode_tensor(t_b)},
{"type": "text", "text": "!"},
],
}
]
conv, mm_data, _ = await _maybe_await(
parse_fn, messages, mc, content_format="openai"
)
tensors = list(mm_data["prompt_embeds"])
assert len(tensors) == NUM_TENSORS
# Tokenize: each prompt_embeds part becomes 1 placeholder token.
# `return_dict=False` to get a flat `list[int]` on transformers v5
# (where the default flipped to True and yields a `BatchEncoding` dict).
token_ids = tokenizer.apply_chat_template(conv, tokenize=True, return_dict=False)
assert sum(t == tid for t in token_ids) == NUM_TENSORS
# Expand, locate, and build.
mm_updates = _build_prompt_embeds_updates(tensors, tid)
expanded = _expand_prompt_embeds_placeholders(token_ids, mm_updates)
assert len(expanded) == len(token_ids) + LEN_A + LEN_B - NUM_TENSORS
positions = _build_prompt_embeds_positions(expanded, len(tensors), mm_updates)
assert positions[0][1] == LEN_A
assert positions[1][1] == LEN_B
embeds, mask = _build_mixed_prompt_embeds(expanded, tensors, positions)
assert embeds.shape == (len(expanded), _MOCK_HIDDEN_SIZE)
assert mask.count(False) == LEN_A + LEN_B
assert torch.equal(embeds[positions[0][0] : positions[0][0] + LEN_A], t_a)
assert torch.equal(embeds[positions[1][0] : positions[1][0] + LEN_B], t_b)
@pytest.mark.asyncio
async def test_end_to_end_multi_message_conversation(tokenizer, parse_fn):
"""Full pipeline with prompt_embeds spread across system + user messages,
verifying ordering and positioning in the final token stream."""
tokenizer.chat_template = _SIMPLE_CHAT_TEMPLATE
tid = _ensure_prompt_embeds_placeholder_token(tokenizer)
LEN_SYS, LEN_USR = 4, 3
t_sys = torch.randn(LEN_SYS, _MOCK_HIDDEN_SIZE, dtype=_MOCK_DTYPE)
t_usr = torch.randn(LEN_USR, _MOCK_HIDDEN_SIZE, dtype=_MOCK_DTYPE)
NUM_TENSORS = 2 # t_sys and t_usr.
mc = _make_mock_model_config()
messages = [
{
"role": "system",
"content": [
{"type": "text", "text": "You are helpful."},
{"type": "prompt_embeds", "data": _encode_tensor(t_sys)},
],
},
{
"role": "user",
"content": [
{"type": "prompt_embeds", "data": _encode_tensor(t_usr)},
{"type": "text", "text": "Summarize."},
],
},
]
conv, mm_data, _ = await _maybe_await(
parse_fn, messages, mc, content_format="openai"
)
tensors = list(mm_data["prompt_embeds"])
assert len(tensors) == NUM_TENSORS
# Tokenize, expand, locate, and build.
# `return_dict=False` to get a flat `list[int]` on transformers v5
# (where the default flipped to True and yields a `BatchEncoding` dict).
token_ids = tokenizer.apply_chat_template(conv, tokenize=True, return_dict=False)
mm_updates = _build_prompt_embeds_updates(tensors, tid)
expanded = _expand_prompt_embeds_placeholders(token_ids, mm_updates)
positions = _build_prompt_embeds_positions(expanded, len(tensors), mm_updates)
assert positions[0][1] == LEN_SYS
assert positions[1][1] == LEN_USR
# System embed must appear before user embed in the token stream.
assert positions[0][0] < positions[1][0]
embeds, mask = _build_mixed_prompt_embeds(expanded, tensors, positions)
assert embeds.shape == (len(expanded), _MOCK_HIDDEN_SIZE)
assert mask.count(False) == LEN_SYS + LEN_USR
assert torch.equal(embeds[positions[0][0] : positions[0][0] + LEN_SYS], t_sys)
assert torch.equal(embeds[positions[1][0] : positions[1][0] + LEN_USR], t_usr)
+17 -8
View File
@@ -39,6 +39,11 @@ class MockModelConfig:
is_encoder_decoder: bool = False
is_multimodal_model: bool = False
renderer_num_workers: int = 1
hidden_size: int = 768
dtype: torch.dtype = torch.float32
def get_hidden_size(self) -> int:
return self.hidden_size
@dataclass
@@ -384,12 +389,13 @@ class TestRenderEmbedPrompt:
assert torch.equal(results[0]["prompt_embeds"], tensor_input)
def test_multiple_prompt_embeds(self):
renderer = _build_renderer(MockModelConfig())
hidden_size = 512
renderer = _build_renderer(MockModelConfig(hidden_size=hidden_size))
# Create multiple test tensors
tensor_inputs = [
torch.randn(8, 512, dtype=torch.float32),
torch.randn(12, 512, dtype=torch.float32),
torch.randn(8, hidden_size, dtype=torch.float32),
torch.randn(12, hidden_size, dtype=torch.float32),
]
prompts = renderer.render_prompts(
@@ -432,13 +438,15 @@ class TestRenderEmbedPrompt:
assert torch.equal(results[0]["prompt_embeds"], expected)
def test_prompt_embed_different_dtypes(self):
renderer = _build_renderer(MockModelConfig())
hidden_size = 256
# Test different supported dtypes
dtypes = [torch.float32, torch.float16, torch.bfloat16]
for dtype in dtypes:
tensor_input = torch.randn(5, 256, dtype=dtype)
renderer = _build_renderer(
MockModelConfig(hidden_size=hidden_size, dtype=dtype)
)
tensor_input = torch.randn(5, hidden_size, dtype=dtype)
prompts = renderer.render_prompts(
_preprocess_prompt(
@@ -474,10 +482,11 @@ class TestRenderEmbedPrompt:
assert results[0]["prompt_embeds"].shape == (10, 768)
def test_both_prompts_and_embeds(self):
renderer = _build_renderer(MockModelConfig())
hidden_size = 256
renderer = _build_renderer(MockModelConfig(hidden_size=hidden_size))
text_input = "Hello world"
tensor_input = torch.randn(5, 256, dtype=torch.float32)
tensor_input = torch.randn(5, hidden_size, dtype=torch.float32)
prompts = renderer.render_prompts(
_preprocess_prompt(
@@ -12,6 +12,7 @@ import pybase64 as base64
import pytest
import torch
from vllm.exceptions import VLLMValidationError
from vllm.multimodal.media import AudioEmbeddingMediaIO, ImageEmbeddingMediaIO
from vllm.renderers.embed_utils import safe_load_prompt_embeds
@@ -53,8 +54,14 @@ def _create_malicious_sparse_tensor() -> torch.Tensor:
values = torch.tensor([1.0])
shape = (3, 3)
# Create sparse tensor (this will be invalid)
sparse_tensor = torch.sparse_coo_tensor(indices, values, shape, dtype=torch.float32)
# Create sparse tensor (this will be invalid). Pass `check_invariants=False`
# explicitly so this fixture is robust to process-wide invariant-check state
# left enabled by other tests (the global flag isn't thread-local, and
# concurrent users of the `check_sparse_tensor_invariants` context manager
# can leak the "enabled" state across tests).
sparse_tensor = torch.sparse_coo_tensor(
indices, values, shape, dtype=torch.float32, check_invariants=False
)
return sparse_tensor
@@ -117,7 +124,7 @@ class TestPromptEmbedsValidation:
shape = (10, 10)
malicious_tensor = torch.sparse_coo_tensor(
indices, values, shape, dtype=torch.float32
indices, values, shape, dtype=torch.float32, check_invariants=False
)
encoded = _encode_tensor(malicious_tensor)
@@ -132,13 +139,69 @@ class TestPromptEmbedsValidation:
shape = (10, 10)
malicious_tensor = torch.sparse_coo_tensor(
indices, values, shape, dtype=torch.float32
indices, values, shape, dtype=torch.float32, check_invariants=False
)
encoded = _encode_tensor(malicious_tensor)
with pytest.raises((RuntimeError, ValueError)):
safe_load_prompt_embeds(model_config, encoded)
def test_hidden_size_mismatch_rejected(self, model_config):
"""Tensors whose trailing dim doesn't match the model's hidden_size
must be rejected at parse time."""
# opt-125m has hidden_size=768, passing 512 triggers the check.
wrong_hidden = torch.randn(10, 512, dtype=torch.float32)
encoded = _encode_tensor(wrong_hidden)
with pytest.raises(VLLMValidationError, match="hidden_size"):
safe_load_prompt_embeds(model_config, encoded)
def test_float_dtype_mismatch_cast_to_model_dtype(self, model_config):
"""Tensors whose dtype doesn't match the model's dtype but are still
floating-point are cast, since API clients generally can't know the
server's `--dtype` setting ahead of time."""
# Fixture pins model dtype to float32, upload a bfloat16 tensor.
mismatched_float = torch.randn(10, 768, dtype=torch.bfloat16)
encoded = _encode_tensor(mismatched_float)
result = safe_load_prompt_embeds(model_config, encoded)
assert result.dtype == torch.float32
assert result.shape == mismatched_float.shape
def test_non_float_dtype_rejected(self, model_config):
"""Non-floating-point dtypes cannot be safely cast for embeddings
(e.g. integer tensors almost certainly indicate caller confusion),
so they are rejected at parse time."""
non_float = torch.randint(0, 100, (10, 768), dtype=torch.int32)
encoded = _encode_tensor(non_float)
with pytest.raises(VLLMValidationError, match="floating-point"):
safe_load_prompt_embeds(model_config, encoded)
def test_non_2d_tensor_rejected(self, model_config):
"""Tensors that aren't 2D (even after squeezing a leading dim)
must be rejected with a clear error."""
# A 1D tensor cannot be interpreted as (num_tokens, hidden_size).
bad = torch.randn(768, dtype=torch.float32)
encoded = _encode_tensor(bad)
with pytest.raises(VLLMValidationError, match="2D tensor"):
safe_load_prompt_embeds(model_config, encoded)
def test_non_tensor_payload_rejected(self, model_config):
"""Deserializing to a non-Tensor object must raise a clear error
instead of propagating an AssertionError."""
# `torch.save` will serialize a plain dict; `weights_only=True` allows
# loading built-in containers, so this exercises the isinstance check.
buffer = io.BytesIO()
torch.save({"not": "a tensor"}, buffer)
buffer.seek(0)
encoded = base64.b64encode(buffer.read())
with pytest.raises(VLLMValidationError, match="torch.Tensor"):
safe_load_prompt_embeds(model_config, encoded)
class TestImageEmbedsValidation:
"""Test sparse tensor validation in image embeddings (Chat API)."""
+37 -3
View File
@@ -1295,11 +1295,14 @@ def test_ir_op_priority_default():
# Assert default is applied to ops
priority_config = IrOpPriorityConfig.with_default(["vllm_c", "native"])
assert priority_config.rms_norm == ["vllm_c", "native"]
assert priority_config.fused_add_rms_norm == ["vllm_c", "native"]
# Assert single ops override the default
assert IrOpPriorityConfig.with_default(
["vllm_c", "native"], rms_norm=["oink", "native"]
) == IrOpPriorityConfig(rms_norm=["oink", "native"])
priority_config = IrOpPriorityConfig.with_default(
["native"], rms_norm=["oink", "native"]
)
assert priority_config.rms_norm == ["oink", "native"]
assert priority_config.fused_add_rms_norm == ["native"]
def test_ir_op_priority_str():
@@ -1318,3 +1321,34 @@ def test_ir_op_priority_str():
with pytest.raises(pydantic.ValidationError):
# must be list of only strings
priority_config = IrOpPriorityConfig(rms_norm=["vllm_c", 4, "native"])
def test_ir_op_priority_ctx():
"""Test that the priority-setting context sets priority correctly."""
from vllm import ir
from vllm.config.kernel import IrOpPriorityConfig
priority = IrOpPriorityConfig.with_default(["native"], rms_norm=["vllm_c"])
priority2 = IrOpPriorityConfig.with_default(
["native"], fused_add_rms_norm=["vllm_c"]
)
with priority.set_priority():
assert ir.ops.rms_norm.get_priority() == ["vllm_c", "native"]
assert ir.ops.fused_add_rms_norm.get_priority() == ["native"]
with priority2.set_priority():
assert ir.ops.rms_norm.get_priority() == ["native"]
assert ir.ops.fused_add_rms_norm.get_priority() == ["vllm_c", "native"]
# context restored
assert ir.ops.rms_norm.get_priority() == ["vllm_c", "native"]
assert ir.ops.fused_add_rms_norm.get_priority() == ["native"]
with pytest.raises(ValueError), priority2.set_priority():
assert ir.ops.rms_norm.get_priority() == ["native"]
assert ir.ops.fused_add_rms_norm.get_priority() == ["vllm_c", "native"]
raise ValueError
# context restored even after exception
assert ir.ops.rms_norm.get_priority() == ["vllm_c", "native"]
assert ir.ops.fused_add_rms_norm.get_priority() == ["native"]
+1
View File
@@ -40,6 +40,7 @@ def _model_config():
multimodal_config=None,
allowed_local_media_path="",
allowed_media_domains=None,
enable_prompt_embeds=False,
)
@@ -6,6 +6,15 @@
import json
from unittest.mock import MagicMock
import pytest
from xgrammar import StructuralTag
from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionNamedFunction,
ChatCompletionNamedToolChoiceParam,
ChatCompletionRequest,
ChatCompletionToolsParam,
)
from vllm.tool_parsers import ToolParserManager
from vllm.tool_parsers.deepseekv4_tool_parser import DeepSeekV4ToolParser
@@ -20,6 +29,43 @@ PARAM_START = '<DSMLparameter name="'
PARAM_END = "</DSMLparameter>"
@pytest.fixture
def sample_tools() -> list[ChatCompletionToolsParam]:
return [
ChatCompletionToolsParam(
type="function",
function={
"name": "get_current_weather",
"description": "Get the current weather",
"parameters": {
"type": "object",
"properties": {
"city": {"type": "string", "description": "The city name"},
"state": {"type": "string", "description": "The state code"},
"unit": {"type": "string", "enum": ["fahrenheit", "celsius"]},
},
"required": ["city", "state"],
},
},
),
ChatCompletionToolsParam(
type="function",
function={
"name": "calculate_area",
"description": "Calculate area of a shape",
"parameters": {
"type": "object",
"properties": {
"shape": {"type": "string"},
"dimensions": {"type": "object"},
"precision": {"type": "integer"},
},
},
},
),
]
def make_parser(tools=None) -> DeepSeekV4ToolParser:
return DeepSeekV4ToolParser(MOCK_TOKENIZER, tools=tools)
@@ -121,3 +167,39 @@ def test_streaming_extracts_complete_invokes():
]
assert names == ["search"]
assert json.loads(reconstruct_args(deltas)) == {"query": "deepseek v4"}
def test_get_vllm_registry_structural_tag_returns_structural_tag(
sample_tools: list[ChatCompletionToolsParam],
) -> None:
parser = make_parser()
req = ChatCompletionRequest(
messages=[],
model="m",
tools=sample_tools,
tool_choice="auto",
)
tag = parser.get_structural_tag(req)
assert isinstance(tag, StructuralTag)
req = ChatCompletionRequest(
messages=[],
model="m",
tools=sample_tools,
tool_choice="required",
)
tag = parser.get_structural_tag(req)
assert isinstance(tag, StructuralTag)
if sample_tools:
tool = sample_tools[0]
req = ChatCompletionRequest(
messages=[],
model="m",
tools=sample_tools,
)
req.tool_choice = ChatCompletionNamedToolChoiceParam(
function=ChatCompletionNamedFunction(name=tool.function.name)
)
tag = parser.get_structural_tag(req)
assert isinstance(tag, StructuralTag)
@@ -6,8 +6,11 @@ from collections.abc import Generator
import pytest
from openai.types.responses.function_tool import FunctionTool
from xgrammar import StructuralTag
from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionNamedFunction,
ChatCompletionNamedToolChoiceParam,
ChatCompletionRequest,
ChatCompletionToolsParam,
)
@@ -108,6 +111,27 @@ def sample_tools(request):
]
def _as_chat_completion_tools(
tools: list[ChatCompletionToolsParam | FunctionTool],
) -> list[ChatCompletionToolsParam]:
normalized: list[ChatCompletionToolsParam] = []
for tool in tools:
if isinstance(tool, ChatCompletionToolsParam):
normalized.append(tool)
else:
normalized.append(
ChatCompletionToolsParam(
type="function",
function={
"name": tool.name,
"description": tool.description,
"parameters": tool.parameters,
},
)
)
return normalized
def assert_tool_calls(
actual_tool_calls: list[ToolCall], expected_tool_calls: list[ToolCall]
):
@@ -1146,3 +1170,88 @@ def test_no_double_serialization_string_args(qwen3_tool_parser):
args = json.loads(raw_arguments)
assert args["message"] == "hello world"
assert '\\"hello world\\"' not in raw_arguments
def test_get_vllm_registry_structural_tag_returns_structural_tag(
qwen3_tool_parser: Qwen3CoderToolParser,
sample_tools: list[ChatCompletionToolsParam],
) -> None:
request_tools = _as_chat_completion_tools(sample_tools)
req = ChatCompletionRequest(
messages=[],
model="m",
tools=request_tools,
tool_choice="auto",
)
tag = qwen3_tool_parser.get_structural_tag(req)
assert isinstance(tag, StructuralTag)
req = ChatCompletionRequest(
messages=[],
model="m",
tools=request_tools,
tool_choice="required",
)
tag = qwen3_tool_parser.get_structural_tag(req)
assert isinstance(tag, StructuralTag)
if request_tools:
tool = request_tools[0]
req = ChatCompletionRequest(
messages=[],
model="m",
tools=request_tools,
)
req.tool_choice = ChatCompletionNamedToolChoiceParam(
function=ChatCompletionNamedFunction(name=tool.function.name)
)
tag = qwen3_tool_parser.get_structural_tag(req)
assert isinstance(tag, StructuralTag)
@pytest.mark.parametrize("include_reasoning", [True, False])
def test_adjust_request_auto_uses_vllm_registry_structural_tag(
monkeypatch: pytest.MonkeyPatch,
qwen3_tool_parser: Qwen3CoderToolParser,
sample_tools: list[ChatCompletionToolsParam],
include_reasoning: bool,
) -> None:
monkeypatch.setattr(
"vllm.tool_parsers.abstract_tool_parser.VLLM_ENFORCE_STRICT_TOOL_CALLING",
True,
)
request_tools = _as_chat_completion_tools(sample_tools)
req = ChatCompletionRequest(
messages=[],
model="m",
tools=request_tools,
tool_choice="auto",
include_reasoning=include_reasoning,
)
out = qwen3_tool_parser.adjust_request(req)
assert out.structured_outputs is not None
assert out.structured_outputs.structural_tag is not None
assert isinstance(out.structured_outputs.structural_tag, str)
loaded = json.loads(out.structured_outputs.structural_tag)
assert isinstance(loaded, dict)
def test_adjust_request_required_prefers_structural_tag(
monkeypatch: pytest.MonkeyPatch,
qwen3_tool_parser: Qwen3CoderToolParser,
sample_tools: list[ChatCompletionToolsParam],
) -> None:
monkeypatch.setattr(
"vllm.tool_parsers.abstract_tool_parser.VLLM_ENFORCE_STRICT_TOOL_CALLING",
True,
)
request_tools = _as_chat_completion_tools(sample_tools)
req = ChatCompletionRequest(
messages=[],
model="m",
tools=request_tools,
tool_choice="required",
)
out = qwen3_tool_parser.adjust_request(req)
assert out.structured_outputs is not None
assert out.structured_outputs.structural_tag is not None
+126 -28
View File
@@ -2,6 +2,7 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import asyncio
import atexit
import contextlib
import copy
import functools
@@ -134,6 +135,11 @@ class RemoteVLLMServer:
"""
DUMMY_API_KEY = "token-abc123" # vLLM's OpenAI server does not need API key
_active_servers: set["RemoteVLLMServer"] = set()
_active_servers_lock = threading.RLock()
_cleanup_hooks_registered = False
_signal_hooks_registered = False
_previous_signal_handlers: dict[int, Any] = {}
proc: subprocess.Popen
def _create_cli_subcommand(self):
@@ -209,6 +215,7 @@ class RemoteVLLMServer:
)
self._pre_download_model(model, args)
self._shutdown_complete = False
# Record GPU memory before server start so we know what
# "released" looks like.
@@ -221,6 +228,7 @@ class RemoteVLLMServer:
)
self._start_server(model, vllm_serve_args, env_dict)
self._register_active_server()
max_wait_seconds = max_wait_seconds or 480
try:
self._wait_for_server(url=self.url_for("health"), timeout=max_wait_seconds)
@@ -246,8 +254,70 @@ class RemoteVLLMServer:
(when the server fails to start). Must be safe to call even if
the process is already dead.
"""
self._terminate_process_tree()
self._wait_for_gpu_memory_release()
if self._shutdown_complete:
return
self._shutdown_complete = True
try:
self._terminate_process_tree()
self._wait_for_gpu_memory_release()
finally:
self._unregister_active_server()
@classmethod
def _ensure_cleanup_hooks_registered(cls) -> None:
"""Register process-exit cleanup for detached server subprocesses."""
root_cls = RemoteVLLMServer
with root_cls._active_servers_lock:
if not root_cls._cleanup_hooks_registered:
atexit.register(root_cls._shutdown_active_servers)
root_cls._cleanup_hooks_registered = True
if (
threading.current_thread() is threading.main_thread()
and not root_cls._signal_hooks_registered
):
for signum in (signal.SIGTERM, signal.SIGINT):
root_cls._previous_signal_handlers[signum] = signal.getsignal(
signum
)
signal.signal(signum, root_cls._handle_parent_signal)
root_cls._signal_hooks_registered = True
def _register_active_server(self) -> None:
"""Track this server so parent-process exits still clean it up."""
RemoteVLLMServer._ensure_cleanup_hooks_registered()
with RemoteVLLMServer._active_servers_lock:
RemoteVLLMServer._active_servers.add(self)
def _unregister_active_server(self) -> None:
with RemoteVLLMServer._active_servers_lock:
RemoteVLLMServer._active_servers.discard(self)
@classmethod
def _shutdown_active_servers(cls) -> None:
"""Best-effort shutdown for all live RemoteVLLMServer instances."""
with cls._active_servers_lock:
servers = list(cls._active_servers)
for server in servers:
with contextlib.suppress(Exception):
server._shutdown()
@classmethod
def _handle_parent_signal(cls, signum, frame) -> None:
"""Clean up detached servers before letting the signal terminate pytest."""
cls._shutdown_active_servers()
previous_handler = cls._previous_signal_handlers.get(signum, signal.SIG_DFL)
if callable(previous_handler):
previous_handler(signum, frame)
elif previous_handler == signal.SIG_IGN:
return
elif signum == signal.SIGINT:
raise KeyboardInterrupt
else:
raise SystemExit(128 + signum)
def _terminate_process_tree(self) -> None:
"""Kill the server process tree without waiting for GPU memory release.
@@ -315,6 +385,9 @@ class RemoteVLLMServer:
if not servers:
return
for server in servers:
server._shutdown_complete = True
threads = [
threading.Thread(
target=s._terminate_process_tree,
@@ -339,7 +412,11 @@ class RemoteVLLMServer:
else s._pre_server_gpu_memory
),
)
earliest._wait_for_gpu_memory_release()
try:
earliest._wait_for_gpu_memory_release()
finally:
for server in servers:
server._unregister_active_server()
def _kill_process_group_survivors(
self, pgid: int | None, timeout: float = 15.0
@@ -705,6 +782,7 @@ def _test_completion(
model: str,
prompt: str,
token_ids: list[int],
include_seeded_sampling: bool = True,
):
results = []
@@ -739,33 +817,40 @@ def _test_completion(
}
)
# test seeded random sampling
completion = client.completions.create(
model=model, prompt=prompt, max_tokens=5, seed=33, temperature=1.0
)
if include_seeded_sampling:
# test seeded random sampling
completion = client.completions.create(
model=model, prompt=prompt, max_tokens=5, seed=33, temperature=1.0
)
results.append(
{
"test": "seeded_sampling",
"text": completion.choices[0].text,
"finish_reason": completion.choices[0].finish_reason,
"usage": completion.usage,
}
)
results.append(
{
"test": "seeded_sampling",
"text": completion.choices[0].text,
"finish_reason": completion.choices[0].finish_reason,
"usage": completion.usage,
}
)
# test seeded random sampling with multiple prompts
completion = client.completions.create(
model=model, prompt=[prompt, prompt], max_tokens=5, seed=33, temperature=1.0
)
# test seeded random sampling with multiple prompts
completion = client.completions.create(
model=model,
prompt=[prompt, prompt],
max_tokens=5,
seed=33,
temperature=1.0,
)
results.append(
{
"test": "seeded_sampling",
"text": [choice.text for choice in completion.choices],
"finish_reason": [choice.finish_reason for choice in completion.choices],
"usage": completion.usage,
}
)
results.append(
{
"test": "seeded_sampling",
"text": [choice.text for choice in completion.choices],
"finish_reason": [
choice.finish_reason for choice in completion.choices
],
"usage": completion.usage,
}
)
# test simple list
batch = client.completions.create(
@@ -960,6 +1045,7 @@ def compare_two_settings(
*,
method: str = "generate",
max_wait_seconds: float | None = None,
include_seeded_sampling: bool = True,
) -> None:
"""
Launch API server with two different sets of arguments/environments
@@ -971,6 +1057,8 @@ def compare_two_settings(
arg2: The second set of arguments to pass to the API server.
env1: The first set of environment variables to pass to the API server.
env2: The second set of environment variables to pass to the API server.
include_seeded_sampling: Whether to include temperature=1.0 seeded
sampling checks in the default generate comparison.
"""
compare_all_settings(
@@ -979,6 +1067,7 @@ def compare_two_settings(
[env1, env2],
method=method,
max_wait_seconds=max_wait_seconds,
include_seeded_sampling=include_seeded_sampling,
)
@@ -989,6 +1078,7 @@ def compare_all_settings(
*,
method: str = "generate",
max_wait_seconds: float | None = None,
include_seeded_sampling: bool = True,
) -> None:
"""
Launch API server with several different sets of arguments/environments
@@ -997,6 +1087,8 @@ def compare_all_settings(
model: The model to test.
all_args: A list of argument lists to pass to the API server.
all_envs: A list of environment dictionaries to pass to the API server.
include_seeded_sampling: Whether to include temperature=1.0 seeded
sampling checks in the default generate comparison.
"""
trust_remote_code = False
@@ -1057,7 +1149,13 @@ def compare_all_settings(
)
if method == "generate":
results += _test_completion(client, model, prompt, token_ids)
results += _test_completion(
client,
model,
prompt,
token_ids,
include_seeded_sampling=include_seeded_sampling,
)
elif method == "generate_close":
results += _test_completion_close(client, model, prompt)
elif method == "generate_chat":
@@ -0,0 +1,162 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for canonicalize_singleton_dim_strides.
Background
----------
When num_kv_heads_per_rank == 1 (e.g. Qwen3.5-397B with TP=8 1 KV head
per rank), PyTorch's is_contiguous() returns True for *any* stride on the
size-1 dimension. The KV cache allocator can therefore produce a tensor
where that singleton dim has stride = 1 element (2 bytes for bf16) instead
of the canonical product-of-remaining-dims value.
CUDA TMA (used by FlashInfer XQA SM90 and Flash-Attention 3/4 on H100+)
requires all non-outermost strides to be multiples of 16 bytes. A 2-byte
stride triggers cudaErrorIllegalInstruction.
canonicalize_singleton_dim_strides() patches degenerate strides on all
size-1 dimensions via torch.as_strided zero-copy.
The degenerate stride manifests at different positions in different backends:
- FlashInfer: stride(-3) after kv_cache.permute() shape [..., 1, B, D]
- FlashAttention: stride(-2) after kv_cache.unbind(0) shape [N, B, 1, D]
"""
import torch
from vllm.utils.torch_utils import canonicalize_singleton_dim_strides
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _inject_degenerate_stride(t: torch.Tensor, dim: int) -> torch.Tensor:
"""Return a view of t with a degenerate (stride=1) on a size-1 dim."""
assert t.shape[dim] == 1, f"dim {dim} must have size 1"
strides = list(t.stride())
strides[dim] = 1 # inject the bug
return t.as_strided(t.shape, strides)
# ---------------------------------------------------------------------------
# Tests: canonicalize_singleton_dim_strides
# ---------------------------------------------------------------------------
class TestCanonicalizeSingletonDimStrides:
def test_flashinfer_layout_dim_neg3(self):
"""FlashInfer path: degenerate stride at dim -3 (num_kv_heads)."""
# Shape after permute: [num_blocks, 2, num_kv_heads, block_size, head_size]
num_blocks, block_size, head_size = 64, 16, 128
t = torch.zeros(num_blocks, 2, 1, block_size, head_size, dtype=torch.bfloat16)
t_deg = _inject_degenerate_stride(t, dim=-3)
assert t_deg.stride(-3) == 1 # confirm degenerate
assert t_deg.is_contiguous() # PyTorch doesn't notice
fixed = canonicalize_singleton_dim_strides(t_deg)
assert fixed.stride(-3) == block_size * head_size # canonical = 2048
assert fixed.stride(-2) == head_size # inner dims unchanged
assert fixed.stride(-1) == 1
def test_flash_attn_layout_dim_neg2(self):
"""FlashAttention path: degenerate stride at dim -2 (num_kv_heads)."""
# Shape after unbind(0): [num_blocks, block_size, num_kv_heads, head_size]
num_blocks, block_size, head_size = 64, 16, 128
t = torch.zeros(num_blocks, block_size, 1, head_size, dtype=torch.bfloat16)
t_deg = _inject_degenerate_stride(t, dim=-2)
assert t_deg.stride(-2) == 1
assert t_deg.is_contiguous()
fixed = canonicalize_singleton_dim_strides(t_deg)
assert fixed.stride(-2) == head_size # canonical = 128
assert fixed.stride(-1) == 1
def test_canonical_strides_returned_as_is(self):
"""No degenerate strides → same object returned (no copy, no new view)."""
t = torch.zeros(64, 2, 1, 16, 128, dtype=torch.bfloat16)
result = canonicalize_singleton_dim_strides(t)
assert result is t
def test_multi_kv_heads_unchanged(self):
"""num_kv_heads > 1 → strides are already canonical → unchanged."""
t = torch.zeros(16, 2, 4, 16, 128, dtype=torch.bfloat16)
original_strides = t.stride()
result = canonicalize_singleton_dim_strides(t)
assert result.stride() == original_strides
def test_data_pointer_preserved(self):
"""Fix is zero-copy: same underlying storage."""
t = torch.zeros(8, 2, 1, 16, 128, dtype=torch.bfloat16)
t_deg = _inject_degenerate_stride(t, dim=-3)
fixed = canonicalize_singleton_dim_strides(t_deg)
assert fixed.data_ptr() == t_deg.data_ptr()
assert fixed.storage_offset() == t_deg.storage_offset()
def test_multiple_singleton_dims(self):
"""All size-1 dims with degenerate strides are fixed."""
# Shape: [1, 1, 8, 32] — two size-1 dims
t = torch.zeros(1, 1, 8, 32, dtype=torch.float16)
# Both size-1 dims get degenerate strides
t_deg = t.as_strided(t.shape, (1, 1, 32, 1)) # both leading dims = 1
fixed = canonicalize_singleton_dim_strides(t_deg)
assert fixed.stride(0) == 1 * 8 * 32 # canonical: 256
assert fixed.stride(1) == 1 * 8 * 32 # canonical: 256 (same since size-1)
assert fixed.stride(2) == 32
assert fixed.stride(3) == 1
def test_various_shapes_flashinfer(self):
"""Correctness across different block_size / head_size for FlashInfer layout."""
for block_size, head_size in [(16, 64), (16, 128), (32, 128), (16, 256)]:
t = torch.zeros(8, 2, 1, block_size, head_size, dtype=torch.bfloat16)
t_deg = _inject_degenerate_stride(t, dim=-3)
fixed = canonicalize_singleton_dim_strides(t_deg)
assert fixed.stride(-3) == block_size * head_size, (
f"Failed for block_size={block_size}, head_size={head_size}: "
f"got stride(-3)={fixed.stride(-3)}"
)
def test_various_shapes_flash_attn(self):
"""Correctness across different shapes for FlashAttention layout."""
for block_size, head_size in [(16, 64), (16, 128), (32, 128)]:
t = torch.zeros(8, block_size, 1, head_size, dtype=torch.bfloat16)
t_deg = _inject_degenerate_stride(t, dim=-2)
fixed = canonicalize_singleton_dim_strides(t_deg)
assert fixed.stride(-2) == head_size, (
f"Failed for block_size={block_size}, head_size={head_size}: "
f"got stride(-2)={fixed.stride(-2)}"
)
def test_tma_alignment_satisfied_after_fix_bf16(self):
"""After fix, all strides meet 16-byte TMA alignment for bf16."""
t = torch.zeros(64, 2, 1, 16, 128, dtype=torch.bfloat16)
t_deg = _inject_degenerate_stride(t, dim=-3)
fixed = canonicalize_singleton_dim_strides(t_deg)
element_size = fixed.element_size() # 2 bytes for bf16
for i, s in enumerate(fixed.stride()):
assert (s * element_size) % 16 == 0 or i == len(fixed.stride()) - 1, (
f"dim {i} stride {s} * {element_size} bytes not 16-byte aligned"
)
def test_non_contiguous_outer_dims_preserved(self):
"""Outer (non-size-1) non-contiguous strides are left unchanged."""
# Simulate cross-layer unified allocation: num_blocks stride is non-canonical
# but the inner dims should be fixed.
base = torch.zeros(200, 2, 1, 16, 128, dtype=torch.bfloat16)
# Slice every 2nd block → non-canonical outer stride
t_sliced = base[::2] # shape [100, 2, 1, 16, 128], stride[0] = 2*canonical
t_deg = _inject_degenerate_stride(t_sliced, dim=-3)
fixed = canonicalize_singleton_dim_strides(t_deg)
# Outer stride should be unchanged (not a size-1 dim)
assert fixed.stride(0) == t_sliced.stride(0)
# Inner degenerate stride should be fixed
assert fixed.stride(-3) == 16 * 128

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