Compare commits

...
Author SHA1 Message Date
yewentao256 77230471c0 remove torch 2.9, 2.10 workround
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-04-30 21:10:50 +00:00
Stefano CastagnettaandGitHub efb4cdf2b8 [CI/Build] Skip Prithvi/Terratorch model-registry tests when terratorch is missing (#41389)
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
2026-04-30 12:47:55 -07:00
Stefano CastagnettaGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
92a7c121b6 [CI] Add MTP coverage: Qwen3.5 correctness + no-sync spec decode (#40472)
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-30 12:24:09 -07:00
Jee Jee LiandGitHub 307b17ce33 [DSV4] Avoid redundant dtype conversion. (#41374)
Signed-off-by: Jee Jee Li <pandaleefree@gmail.com>
2026-04-30 09:57:27 -07:00
3ca6ca210f xpu docker: pin oneAPI to 2025.3 and avoid unintended 2026 upgrade (#41380)
Signed-off-by: wendyliu235 <wenjun.liu@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-04-30 16:02:23 +00:00
Stefano CastagnettaandGitHub 10558f5f46 [CI/Build] Skip terratorch + torchgeo while PyPI has lightning quarantined (#41377)
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
2026-04-30 07:59:07 -07:00
121dbe7a22 [ROCm] ROCm DeepEP API updated to latest (#39721)
Signed-off-by: Tej Kiran <vpolamre@amd.com>
Signed-off-by: tej <37236721+itej89@users.noreply.github.com>
Co-authored-by: TJian <tunjian.tan@embeddedllm.com>
Co-authored-by: HAIAI <39548240+HAIAI@users.noreply.github.com>
2026-04-30 07:46:59 -07:00
Matthew BonanniandGitHub f03d82efdd [UX][Bugfix] Fix OOM by setting PyTorch max_split_size_mb during model loading (#41268)
Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
2026-04-30 07:46:54 -07:00
a7fb008510 [EPLB] Optimize memory overhead in Nixl communicator (#40013)
Signed-off-by: ilmarkov <markovilya197@gmail.com>
Signed-off-by: Markov Ilya <markovilya19@gmail.com>
Co-authored-by: Markov Ilya <markovilya19@gmail.com>
Co-authored-by: Nicolò Lucchesi <nlucches@redhat.com>
2026-04-30 07:46:49 -07:00
Harry MellorandGitHub ff449b6426 Stop mergify labelling from skipping pre-commit (#41362)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-04-30 05:48:38 -07:00
3527229517 [Doc] Fix RTD build: pytorch.org/docs/stable/objects.inv returns 404 (#41353)
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-04-30 05:06:44 -07:00
Xiaoshuang WangGitHubCyrus Leungmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
b55b26520c [MoE] Make MoERunnerInterface a PluggableLayer for OOT support (#35178)
Signed-off-by: wxsIcey <1790571317@qq.com>
Signed-off-by: Icey <1790571317@qq.com>
Co-authored-by: Cyrus Leung <tlleungac@connect.ust.hk>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-30 03:31:08 -07:00
snadampalandGitHub 3179e53135 [P/D] Prefill compute optimizations with bi-directional KV cache transfers between P and D nodes (#32553)
Signed-off-by: Sunita Nadampalli <nadampal@amazon.com>
2026-04-30 10:14:20 +00:00
Nicolò LucchesiandGitHub efdc95674d [KVConnector] MultiConnector SupportsHMA (#39571)
Signed-off-by: NickLucche <nlucches@redhat.com>
2026-04-30 02:10:50 -07:00
Chenxi QianGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
54146a9bf9 [Bugfix] correct h matrix layout in chunk_kda output kernel (#40956)
Signed-off-by: ChenxiQian <chenxi.qian.cq@outlook.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-30 16:22:41 +08:00
ca97f7b9bb Fix Gemma4 MoE expert weight remapping (#41206)
Signed-off-by: sunghoon.baek <sunghoon.baek@connectfy.cloud>
Co-authored-by: sunghoon.baek <sunghoon.baek@connectfy.cloud>
Co-authored-by: OpenAI Codex <codex@openai.com>
2026-04-30 00:12:42 -07:00
Ekagra RanjanandGitHub a04e0cf3b8 Fix Cohere ASR after HF upgrade (#40582)
Signed-off-by: Ekagra Ranjan <3116519+ekagra-ranjan@users.noreply.github.com>
2026-04-29 23:39:04 -07:00
cb1b02d0e8 [Frontend] Add VLLM_SKIP_MODEL_NAME_VALIDATION environment variable (#34676)
Signed-off-by: Dhruv Singal <dhruvsingalabc@gmail.com>
Signed-off-by: Dhruv Singal <dsingal@Dhruvs-MacBook-Pro.local>
Signed-off-by: Your Name <you@example.com>
Signed-off-by: vLLM Assistant <assistant@vllm.ai>
Signed-off-by: Simon Mo <simon.mo@hey.com>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
Co-authored-by: Dhruv Singal <dsingal@Dhruvs-MacBook-Pro.local>
Co-authored-by: Your Name <you@example.com>
Co-authored-by: OpenCode <noreply@openai.com>
Co-authored-by: Simon Mo <simon.mo@hey.com>
2026-04-29 23:19:09 -07:00
a749a33d8d [Bugfix] Fix persistent_topk cooperative deadlock at TopK=1024 (#41189)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-29 21:03:45 -07:00
c42981d034 [Refactor][kv_offload] KV Offloading maintainability improvements (#40538)
Signed-off-by: Martin Hickey <martin.hickey@ie.ibm.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-04-30 05:55:31 +03:00
Wei ZhaoandGitHub 0ff1bf9bb1 [Bugfix] Fix failure to allocate KV blocks error (#41282)
Signed-off-by: wzhao18 <wzhao18.sz@gmail.com>
2026-04-29 18:44:07 -07:00
0ab67c0222 [CI] Add key field to all test_areas pipeline steps (#41201)
Signed-off-by: khluu <khluu000@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-04-29 16:59:16 -07:00
Rohan PotdarandGitHub 3795d7acf4 [ROCm][Bugfix][GPTOSS]: fix input_ids and expert_map args for quark w4a8 gptoss (#41165)
Signed-off-by: Rohan138 <rohanpotdar138@gmail.com>
2026-04-29 16:39:01 -07:00
Nick HillandGitHub 18599bfdf2 [Ci][BugFix] Fix slow DP tests due to bad teardown logic (#41166)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-04-29 19:31:00 -04:00
Thien TranandGitHub 296741d025 [DSv4] Use cvt PTX for FP32->FP4 conversion (#41015)
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
2026-04-29 16:16:40 -07:00
a966aaed30 [Bugfix][MLA] Size arange_buffer to max_num_batched_tokens to prevent CUDA IMA (#39277)
Signed-off-by: UranusSeven <109661872+UranusSeven@users.noreply.github.com>
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
Co-authored-by: Woosuk Kwon <woosuk@inferact.ai>
2026-04-29 16:14:50 -07:00
Hemanth AcharyaandGitHub 6841f5dc77 [ROCm] Add env flags to disable dynamic MXFP4 quant and enable AITER tuned GEMMs for Attention Projection Layers (#39987)
Signed-off-by: Hemanth Acharya <heachary@amd.com>
2026-04-29 16:07:46 -07:00
roikoren755andGitHub c2fb013312 [Bugfix][Compile] Fix gc.collect/empty_cache patch arity in CUDAGraphWrapper (#41235)
Signed-off-by: Roi Koren <roik@nvidia.com>
2026-04-29 21:59:18 +00:00
ccfb620c62 Create tests/distributed/test_mnnvl_alltoall.py (#35241)
Signed-off-by: Rishi Puri <riship@nvidia.com>
Signed-off-by: Claude <claude@anthropic.com>
Signed-off-by: Stefano Castagnetta <scastagnetta@nvidia.com>
Co-authored-by: Claude <claude@anthropic.com>
Co-authored-by: Stefano Castagnetta <scastagnetta@nvidia.com>
2026-04-29 21:56:56 +00:00
0335316a9b [BUG] Two phase pause to prevent deadlock (#39366)
Signed-off-by: ahao-anyscale <ahao@anyscale.com>
Signed-off-by: Aaron Hao <ahao@anyscale.com>
Co-authored-by: Junjie Zhang <junj.jay.zhang@gmail.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
2026-04-29 17:51:03 -04:00
Rohan PotdarandGitHub 944e138bcf [ROCm][Bugfix]: W4A4 MOE using emulation instead of AITER on MXFP4-supported hardware (#41175)
Signed-off-by: Rohan138 <rohanpotdar138@gmail.com>
2026-04-29 16:39:03 -05:00
Luochao WangGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
b58669cb42 [Perf][Spec Decode] Avoid per-step numpy allocation in prepare_next_t… (#41043)
Signed-off-by: wangluochao902 <wangluochao902@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-29 14:20:13 -07:00
Isotr0pyandGitHub 1628239eb2 [Multimodal][Render] Skip mm processor initialization and warmup for text-only mode (#41246)
Signed-off-by: Isotr0py <mozf@mail2.sysu.edu.cn>
2026-04-29 14:16:19 -07:00
yzong-rhandGitHub 93da1fe97a [CI] Add temperature to bfcl eval, default greedy (#41059)
Signed-off-by: Yifan Zong <yzong@redhat.com>
2026-04-29 14:01:57 -07:00
Andrew BarnesandGitHub 169988a3c0 [ROCm] Use quant_dtype in per_token_quant instead of hardcoded FP8 (#39121)
Signed-off-by: Bortlesboat <bortstheboat@gmail.com>
2026-04-29 20:46:01 +00:00
ChaunceyGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
faab189554 [Feature]: IndexCache support for DSA models (#37735)
Signed-off-by: chaunceyjiang <chaunceyjiang@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-29 15:15:35 -04:00
Laith SakkaandGitHub 6f20f81cbf Replace shape_invariants with simpler apprach in dynamic_arg_dims utilizing shape_id property. (#36194)
Signed-off-by: Laith Sakka <lsakka@meta.com>
2026-04-29 18:32:15 +00:00
136 changed files with 5012 additions and 2306 deletions
@@ -28,6 +28,7 @@
# BFCL_MAX_MODEL_LEN - Max model length (default: 4096)
# BFCL_PORT - Server port (default: 8000)
# BFCL_REASONING_PARSER - Reasoning parser name (default: disabled)
# BFCL_TEMPERATURE - Temperature (default: 0.0)
# BFCL_EXTRA_ARGS - Additional vLLM server args
set -euo pipefail
@@ -43,6 +44,7 @@ TP_SIZE="${BFCL_TP_SIZE:-1}"
MAX_MODEL_LEN="${BFCL_MAX_MODEL_LEN:-4096}"
PORT="${BFCL_PORT:-8000}"
REASONING_PARSER="${BFCL_REASONING_PARSER:-}"
TEMPERATURE="${BFCL_TEMPERATURE:-0.0}"
EXTRA_ARGS="${BFCL_EXTRA_ARGS:-}"
# Set up output directory
@@ -139,7 +141,7 @@ echo "vLLM server is ready. (started in ${SECONDS_WAITED}s)"
# be patched in-process so BFCL knows to use the OpenAI-compatible handler
# against our local vLLM server.
bfcl_exit_code=0
python3 - "$MODEL" "$TEST_CATEGORY" "$NUM_THREADS" "$PORT" "$API_TYPE" "$OUTPUT_DIR" << 'PYEOF' || bfcl_exit_code=$?
python3 - "$MODEL" "$TEST_CATEGORY" "$NUM_THREADS" "$PORT" "$API_TYPE" "$TEMPERATURE" "$OUTPUT_DIR" << 'PYEOF' || bfcl_exit_code=$?
import os
import sys
@@ -148,7 +150,8 @@ test_category = sys.argv[2]
num_threads = int(sys.argv[3])
port = sys.argv[4]
api_type = sys.argv[5]
output_dir = sys.argv[6] if len(sys.argv) > 6 and sys.argv[6] else os.getcwd()
temperature = float(sys.argv[6])
output_dir = sys.argv[7] if len(sys.argv) > 7 and sys.argv[7] else os.getcwd()
os.environ["OPENAI_BASE_URL"] = f"http://localhost:{port}/v1"
os.environ["OPENAI_API_KEY"] = "dummy"
@@ -204,6 +207,7 @@ gen_kwargs["model"] = [model]
gen_kwargs["test_category"] = [c.strip() for c in test_category.split(",")]
gen_kwargs["skip_server_setup"] = True
gen_kwargs["num_threads"] = num_threads
gen_kwargs["temperature"] = temperature
generate(**gen_kwargs)
# ---- evaluate ----
+2
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: V1 attention (H100)
key: v1-attention-h100
timeout_in_minutes: 30
device: h100
source_file_dependencies:
@@ -14,6 +15,7 @@ steps:
- pytest -v -s v1/attention
- label: V1 attention (B200)
key: v1-attention-b200
timeout_in_minutes: 30
device: b200
source_file_dependencies:
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Basic Correctness
key: basic-correctness
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
+2
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Benchmarks CLI Test
key: benchmarks-cli-test
timeout_in_minutes: 20
device: h200_18gb
source_file_dependencies:
@@ -12,6 +13,7 @@ steps:
- pytest -v -s benchmarks/
- label: Attention Benchmarks Smoke Test (B200)
key: attention-benchmarks-smoke-test-b200
device: b200
num_gpus: 2
optional: true
+13
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Sequence Parallel Correctness Tests (2 GPUs)
key: sequence-parallel-correctness-tests-2-gpus
timeout_in_minutes: 50
working_dir: "/vllm-workspace/"
num_devices: 2
@@ -17,6 +18,7 @@ steps:
- pytest -v -s tests/compile/correctness_e2e/test_sequence_parallel.py
- label: Sequence Parallel Correctness Tests (2xH100)
key: sequence-parallel-correctness-tests-2xh100
timeout_in_minutes: 50
working_dir: "/vllm-workspace/"
device: h100
@@ -27,6 +29,7 @@ steps:
- pytest -v -s tests/compile/correctness_e2e/test_sequence_parallel.py
- label: AsyncTP Correctness Tests (2xH100)
key: asynctp-correctness-tests-2xh100
timeout_in_minutes: 50
working_dir: "/vllm-workspace/"
device: h100
@@ -37,6 +40,7 @@ steps:
- pytest -v -s tests/compile/correctness_e2e/test_async_tp.py
- label: AsyncTP Correctness Tests (B200)
key: asynctp-correctness-tests-b200
timeout_in_minutes: 50
working_dir: "/vllm-workspace/"
device: b200
@@ -47,6 +51,7 @@ steps:
- pytest -v -s tests/compile/correctness_e2e/test_async_tp.py
- label: Distributed Compile Unit Tests (2xH100)
key: distributed-compile-unit-tests-2xh100
timeout_in_minutes: 20
working_dir: "/vllm-workspace/"
device: h100
@@ -60,6 +65,7 @@ steps:
- pytest -s -v tests/compile/passes/distributed
- label: Fusion and Compile Unit Tests (2xB200)
key: fusion-and-compile-unit-tests-2xb200
timeout_in_minutes: 20
working_dir: "/vllm-workspace/"
device: b200
@@ -89,6 +95,7 @@ steps:
- pytest -v -s tests/compile/fullgraph/test_full_graph.py::test_fp8_kv_scale_compile
- label: Fusion E2E Quick (H100)
key: fusion-e2e-quick-h100
timeout_in_minutes: 15
working_dir: "/vllm-workspace/"
device: h100
@@ -107,6 +114,7 @@ steps:
- pytest -v -s tests/compile/fusions_e2e/test_tp1_quant.py -k "inductor_partition and not +rms_norm and +quant_fp8 and (qwen3 or deepseek)"
- label: Fusion E2E Config Sweep (H100)
key: fusion-e2e-config-sweep-h100
timeout_in_minutes: 30
working_dir: "/vllm-workspace/"
device: h100
@@ -126,6 +134,7 @@ steps:
- pytest -v -s tests/compile/fusions_e2e/test_tp1_quant.py -k "llama-3"
- label: Fusion E2E Config Sweep (B200)
key: fusion-e2e-config-sweep-b200
timeout_in_minutes: 30
working_dir: "/vllm-workspace/"
device: b200
@@ -139,6 +148,7 @@ steps:
- pytest -v -s tests/compile/fusions_e2e/test_tp1_quant.py -k "inductor_partition and (FLASHINFER and not +rms_norm and (not +quant_fp8 or +quant_fp8 and (qwen3 or deepseek)) or llama-3)"
- label: Fusion E2E TP2 Quick (H100)
key: fusion-e2e-tp2-quick-h100
timeout_in_minutes: 20
working_dir: "/vllm-workspace/"
device: h100
@@ -156,6 +166,7 @@ steps:
- pytest -v -s tests/compile/fusions_e2e/test_tp2_async_tp.py -k "inductor_partition and not +rms_norm and (not +quant_fp8 or +quant_fp8 and (qwen3 or deepseek))"
- label: Fusion E2E TP2 AR-RMS Config Sweep (H100)
key: fusion-e2e-tp2-ar-rms-config-sweep-h100
timeout_in_minutes: 40
working_dir: "/vllm-workspace/"
device: h100
@@ -175,6 +186,7 @@ steps:
- pytest -v -s tests/compile/fusions_e2e/test_tp2_ar_rms.py -k "llama-3"
- label: Fusion E2E TP2 AsyncTP Config Sweep (H100)
key: fusion-e2e-tp2-asynctp-config-sweep-h100
timeout_in_minutes: 40
working_dir: "/vllm-workspace/"
device: h100
@@ -194,6 +206,7 @@ steps:
- pytest -v -s tests/compile/fusions_e2e/test_tp2_async_tp.py -k "llama-3"
- label: Fusion E2E TP2 (B200)
key: fusion-e2e-tp2-b200
timeout_in_minutes: 20
working_dir: "/vllm-workspace/"
device: b200
+2
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Platform Tests (CUDA)
key: platform-tests-cuda
timeout_in_minutes: 15
device: h200_18gb
source_file_dependencies:
@@ -13,6 +14,7 @@ steps:
- pytest -v -s cuda/test_platform_no_cuda_init.py
- label: Cudagraph
key: cudagraph
timeout_in_minutes: 20
source_file_dependencies:
- tests/v1/cudagraph
+8
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Distributed NixlConnector PD accuracy (4 GPUs)
key: distributed-nixlconnector-pd-accuracy-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -13,6 +14,7 @@ steps:
- uv pip install --system -r /vllm-workspace/requirements/kv_connectors.txt
- bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
- label: Distributed FlashInfer NixlConnector PD accuracy (4 GPUs)
key: distributed-flashinfer-nixlconnector-pd-accuracy-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -24,6 +26,7 @@ steps:
- FLASHINFER=1 bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
- label: DP EP Distributed NixlConnector PD accuracy tests (4 GPUs)
key: dp-ep-distributed-nixlconnector-pd-accuracy-tests-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -35,6 +38,7 @@ steps:
- DP_EP=1 bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
- label: CrossLayer KV layout Distributed NixlConnector PD accuracy tests (4 GPUs)
key: crosslayer-kv-layout-distributed-nixlconnector-pd-accuracy-tests-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -46,6 +50,7 @@ steps:
- CROSS_LAYERS_BLOCKS=True bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
- label: Hybrid SSM NixlConnector PD accuracy tests (4 GPUs)
key: hybrid-ssm-nixlconnector-pd-accuracy-tests-4-gpus
timeout_in_minutes: 20
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -57,6 +62,7 @@ steps:
- HYBRID_SSM=1 bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
- label: MultiConnector (Nixl+Offloading) PD accuracy (2 GPUs)
key: multiconnector-nixl-offloading-pd-accuracy-2-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -71,6 +77,7 @@ steps:
- bash v1/kv_connector/nixl_integration/run_multi_connector_accuracy_test.sh
- label: NixlConnector PD + Spec Decode acceptance (2 GPUs)
key: nixlconnector-pd-spec-decode-acceptance-2-gpus
timeout_in_minutes: 30
device: a100
working_dir: "/vllm-workspace/tests"
@@ -84,6 +91,7 @@ steps:
- bash v1/kv_connector/nixl_integration/spec_decode_acceptance_test.sh
- label: MultiConnector (Nixl+Offloading) PD edge cases (2 GPUs)
key: multiconnector-nixl-offloading-pd-edge-cases-2-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 2
+19
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Distributed Comm Ops
key: distributed-comm-ops
timeout_in_minutes: 20
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -16,6 +17,7 @@ steps:
- pytest -v -s distributed/test_shm_storage.py
- label: Distributed DP Tests (2 GPUs)
key: distributed-dp-tests-2-gpus
timeout_in_minutes: 20
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -37,6 +39,7 @@ steps:
- DP_SIZE=2 pytest -v -s entrypoints/openai/test_multi_api_servers.py
- label: Distributed Compile + RPC Tests (2 GPUs)
key: distributed-compile-rpc-tests-2-gpus
timeout_in_minutes: 20
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -59,6 +62,7 @@ steps:
- pytest -v -s ./compile/test_wrapper.py
- label: Distributed Torchrun + Shutdown Tests (2 GPUs)
key: distributed-torchrun-shutdown-tests-2-gpus
timeout_in_minutes: 20
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -81,6 +85,7 @@ steps:
- pytest -v -s v1/worker/test_worker_memory_snapshot.py
- label: Distributed Torchrun + Examples (4 GPUs)
key: distributed-torchrun-examples-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace"
num_devices: 4
@@ -112,6 +117,7 @@ steps:
- VLLM_ALLOW_INSECURE_SERIALIZATION=1 python3 examples/rl/rlhf_ipc.py
- label: Distributed DP Tests (4 GPUs)
key: distributed-dp-tests-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -132,6 +138,7 @@ steps:
- pytest -v -s distributed/test_utils.py
- label: Distributed Compile + Comm (4 GPUs)
key: distributed-compile-comm-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -153,6 +160,7 @@ steps:
- pytest -v -s distributed/test_multiproc_executor.py::test_multiproc_executor_multi_node
- label: Distributed Tests (8 GPUs)(H100)
key: distributed-tests-8-gpus-h100
timeout_in_minutes: 10
device: h100
num_devices: 8
@@ -164,6 +172,8 @@ steps:
- vllm/v1/engine/llm_engine.py
- vllm/v1/executor/uniproc_executor.py
- vllm/v1/worker/gpu_worker.py
- tests/distributed/test_mnnvl_alltoall.py
commands:
# https://github.com/NVIDIA/nccl/issues/1838
- export NCCL_CUMEM_HOST_ENABLE=0
@@ -171,6 +181,7 @@ steps:
- torchrun --nproc-per-node=8 ../examples/features/torchrun/torchrun_dp_example_offline.py --tp-size=2 --pp-size=1 --dp-size=4 --enable-ep
- label: Distributed Tests (4 GPUs)(A100)
key: distributed-tests-4-gpus-a100
device: a100
optional: true
num_devices: 4
@@ -185,6 +196,7 @@ steps:
- pytest -v -s -x lora/test_mixtral.py
- label: Distributed Tests (2 GPUs)(H100)
key: distributed-tests-2-gpus-h100
timeout_in_minutes: 15
device: h100
optional: true
@@ -199,6 +211,7 @@ steps:
- pytest -v -s tests/distributed/test_packed_tensor.py
- label: Distributed Tests (2 GPUs)(B200)
key: distributed-tests-2-gpus-b200
device: b200
optional: true
working_dir: "/vllm-workspace/"
@@ -207,8 +220,12 @@ steps:
- pytest -v -s tests/distributed/test_context_parallel.py
- pytest -v -s tests/distributed/test_nccl_symm_mem_allreduce.py
- pytest -v -s tests/v1/distributed/test_dbo.py
- pytest -v -s tests/distributed/test_mnnvl_alltoall.py
- label: 2 Node Test (4 GPUs)
key: 2-node-test-4-gpus
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -226,6 +243,7 @@ steps:
- ./.buildkite/scripts/run-multi-node-test.sh /vllm-workspace/tests 2 2 $IMAGE_TAG "VLLM_TEST_SAME_HOST=0 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_same_node.py | grep 'Same node test passed' && NUM_NODES=2 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_node_count.py | grep 'Node count test passed' && python3 ../examples/features/data_parallel/data_parallel_offline.py -dp=2 -tp=1 --dp-num-nodes=2 --dp-node-rank=0 --dp-master-addr=192.168.10.10 --dp-master-port=12345 --enforce-eager --trust-remote-code && VLLM_MULTI_NODE=1 pytest -v -s distributed/test_multi_node_assignment.py && VLLM_MULTI_NODE=1 pytest -v -s distributed/test_pipeline_parallel.py" "VLLM_TEST_SAME_HOST=0 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_same_node.py | grep 'Same node test passed' && NUM_NODES=2 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_node_count.py | grep 'Node count test passed' && python3 ../examples/features/data_parallel/data_parallel_offline.py -dp=2 -tp=1 --dp-num-nodes=2 --dp-node-rank=1 --dp-master-addr=192.168.10.10 --dp-master-port=12345 --enforce-eager --trust-remote-code"
- label: Pipeline + Context Parallelism (4 GPUs)
key: pipeline-context-parallelism-4-gpus
timeout_in_minutes: 60
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -240,6 +258,7 @@ steps:
- pytest -v -s distributed/test_pipeline_parallel.py
- label: RayExecutorV2 (4 GPUs)
key: rayexecutorv2-4-gpus
timeout_in_minutes: 60
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: DeepSeek V2-Lite Accuracy
key: deepseek-v2-lite-accuracy
timeout_in_minutes: 60
device: h100
optional: true
@@ -12,6 +13,7 @@ steps:
- bash .buildkite/scripts/scheduled_integration_test/deepseek_v2_lite_ep_eplb.sh 0.25 200 8010
- label: Qwen3-30B-A3B-FP8-block Accuracy
key: qwen3-30b-a3b-fp8-block-accuracy
timeout_in_minutes: 60
device: h100
optional: true
@@ -21,6 +23,7 @@ steps:
- bash .buildkite/scripts/scheduled_integration_test/qwen30b_a3b_fp8_block_ep_eplb.sh 0.8 200 8020
- label: Qwen3-30B-A3B-FP8-block Accuracy (B200)
key: qwen3-30b-a3b-fp8-block-accuracy-b200
timeout_in_minutes: 60
device: b200
optional: true
@@ -30,6 +33,7 @@ steps:
- bash .buildkite/scripts/scheduled_integration_test/qwen30b_a3b_fp8_block_ep_eplb.sh 0.8 200 8020 2 1
- label: Qwen3-30B-A3B-FP8 DP4 Async EPLB Accuracy
key: qwen3-30b-a3b-fp8-dp4-async-eplb-accuracy
timeout_in_minutes: 60
device: h100
optional: true
@@ -39,6 +43,7 @@ steps:
- bash .buildkite/scripts/scheduled_integration_test/qwen30b_a3b_fp8_dp4_async_eplb.sh 0.8 200 8050
- label: DeepSeek V2-Lite Prefetch Offload Accuracy (H100)
key: deepseek-v2-lite-prefetch-offload-accuracy-h100
timeout_in_minutes: 60
device: h100
optional: true
+7
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Engine
key: engine
timeout_in_minutes: 15
device: h200_18gb
source_file_dependencies:
@@ -16,6 +17,7 @@ steps:
- pytest -v -s engine test_sequence.py test_config.py test_logger.py test_vllm_port.py
- label: Engine (1 GPU)
key: engine-1-gpu
timeout_in_minutes: 30
source_file_dependencies:
- vllm/v1/engine/
@@ -25,6 +27,7 @@ steps:
- pytest -v -s v1/engine --ignore v1/engine/test_preprocess_error_handling.py
- label: e2e Scheduling (1 GPU)
key: e2e-scheduling-1-gpu
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -34,6 +37,7 @@ steps:
- pytest -v -s v1/e2e/general/test_async_scheduling.py
- label: e2e Core (1 GPU)
key: e2e-core-1-gpu
timeout_in_minutes: 30
source_file_dependencies:
- vllm/v1/
@@ -42,6 +46,7 @@ steps:
- pytest -v -s v1/e2e/general --ignore v1/e2e/general/test_async_scheduling.py
- label: V1 e2e (2 GPUs)
key: v1-e2e-2-gpus
timeout_in_minutes: 60 # TODO: Fix timeout after we have more confidence in the test stability
optional: true
num_devices: 2
@@ -58,6 +63,7 @@ steps:
- image-build-amd
- label: V1 e2e (4 GPUs)
key: v1-e2e-4-gpus
timeout_in_minutes: 60 # TODO: Fix timeout after we have more confidence in the test stability
optional: true
num_devices: 4
@@ -74,6 +80,7 @@ steps:
- image-build-amd
- label: V1 e2e (4xH100)
key: v1-e2e-4xh100
timeout_in_minutes: 60
device: h100
num_devices: 4
+10 -1
View File
@@ -2,7 +2,8 @@ group: Entrypoints
depends_on:
- image-build
steps:
- label: Entrypoints Unit Tests
- label: Entrypoints Unit Tests
key: entrypoints-unit-tests
timeout_in_minutes: 10
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -13,6 +14,7 @@ steps:
- pytest -v -s entrypoints/ --ignore=entrypoints/llm --ignore=entrypoints/rpc --ignore=entrypoints/sleep --ignore=entrypoints/serve/instrumentator --ignore=entrypoints/openai --ignore=entrypoints/offline_mode --ignore=entrypoints/test_chat_utils.py --ignore=entrypoints/pooling
- label: Entrypoints Integration (LLM)
key: entrypoints-integration-llm
timeout_in_minutes: 40
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -26,6 +28,7 @@ steps:
- pytest -v -s entrypoints/offline_mode # Needs to avoid interference with other tests
- label: Entrypoints Integration (API Server openai - Part 1)
key: entrypoints-integration-api-server-openai-part-1
timeout_in_minutes: 50
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -43,6 +46,7 @@ steps:
- label: Entrypoints Integration (API Server openai - Part 2)
key: entrypoints-integration-api-server-openai-part-2
timeout_in_minutes: 50
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -60,6 +64,7 @@ steps:
- image-build-amd
- label: Entrypoints Integration (API Server openai - Part 3)
key: entrypoints-integration-api-server-openai-part-3
timeout_in_minutes: 50
device: h200_18gb
working_dir: "/vllm-workspace/tests"
@@ -72,6 +77,7 @@ steps:
- pytest -v -s entrypoints/openai --ignore=entrypoints/openai/chat_completion --ignore=entrypoints/openai/completion --ignore=entrypoints/openai/speech_to_text/ --ignore=entrypoints/openai/correctness/ --ignore=entrypoints/openai/tool_parsers/ --ignore=entrypoints/openai/responses --ignore=entrypoints/openai/test_multi_api_servers.py
- label: Entrypoints Integration (API Server 2)
key: entrypoints-integration-api-server-2
timeout_in_minutes: 130
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -86,6 +92,7 @@ steps:
- pytest -v -s tool_use
- label: Entrypoints Integration (Pooling)
key: entrypoints-integration-pooling
timeout_in_minutes: 50
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -96,6 +103,7 @@ steps:
- pytest -v -s entrypoints/pooling
- label: Entrypoints Integration (Responses API)
key: entrypoints-integration-responses-api
timeout_in_minutes: 50
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -105,6 +113,7 @@ steps:
- pytest -v -s entrypoints/openai/responses
- label: OpenAI API Correctness
key: openai-api-correctness
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: EPLB Algorithm
key: eplb-algorithm
timeout_in_minutes: 15
device: h200_18gb
working_dir: "/vllm-workspace/tests"
@@ -15,6 +16,7 @@ steps:
- pytest -v -s distributed/test_eplb_utils.py
- label: EPLB Execution # 17min
key: eplb-execution
timeout_in_minutes: 27
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -26,6 +28,7 @@ steps:
- pytest -v -s distributed/test_eplb_spec_decode.py
- label: Elastic EP Scaling Test
key: elastic-ep-scaling-test
timeout_in_minutes: 20
device: h100
working_dir: "/vllm-workspace/tests"
+25
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: vLLM IR Tests
key: vllm-ir-tests
timeout_in_minutes: 10
device: h200_18gb
working_dir: "/vllm-workspace/"
@@ -14,6 +15,7 @@ steps:
- pytest -v -s tests/kernels/ir
- label: Kernels Core Operation Test
key: kernels-core-operation-test
timeout_in_minutes: 75
source_file_dependencies:
- csrc/
@@ -23,6 +25,7 @@ steps:
- pytest -v -s kernels/core --ignore=kernels/core/test_minimax_reduce_rms.py kernels/test_concat_mla_q.py
- label: Kernels MiniMax Reduce RMS Test (2 GPUs)
key: kernels-minimax-reduce-rms-test-2-gpus
timeout_in_minutes: 15
num_devices: 2
device: h100
@@ -36,6 +39,7 @@ steps:
- pytest -v -s kernels/core/test_minimax_reduce_rms.py
- label: Kernels Attention Test %N
key: kernels-attention-test
timeout_in_minutes: 35
source_file_dependencies:
- csrc/attention/
@@ -49,6 +53,7 @@ steps:
parallelism: 2
- label: Kernels Quantization Test %N
key: kernels-quantization-test
timeout_in_minutes: 90
source_file_dependencies:
- csrc/quantization/
@@ -59,6 +64,7 @@ steps:
parallelism: 2
- label: Kernels MoE Test %N
key: kernels-moe-test
timeout_in_minutes: 25
source_file_dependencies:
- csrc/quantization/cutlass_w8a8/moe/
@@ -74,6 +80,7 @@ steps:
parallelism: 5
- label: Kernels Mamba Test
key: kernels-mamba-test
timeout_in_minutes: 45
source_file_dependencies:
- csrc/mamba/
@@ -82,7 +89,18 @@ steps:
commands:
- pytest -v -s kernels/mamba
- label: Kernels KDA Test
timeout_in_minutes: 20
source_file_dependencies:
- vllm/model_executor/layers/fla/ops/kda.py
- vllm/model_executor/layers/fla/ops/chunk_delta_h.py
- vllm/model_executor/layers/fla/ops/l2norm.py
- tests/kernels/test_kda.py
commands:
- pytest -v -s kernels/test_kda.py
- label: Kernels DeepGEMM Test (H100)
key: kernels-deepgemm-test-h100
timeout_in_minutes: 45
device: h100
num_devices: 1
@@ -104,6 +122,7 @@ steps:
- pytest -v -s quantization/test_cutlass_w4a16.py
- label: Kernels (B200)
key: kernels-b200
timeout_in_minutes: 30
working_dir: "/vllm-workspace/"
device: b200
@@ -152,6 +171,7 @@ steps:
- pytest -v -s tests/models/quantization/test_nvfp4.py
- label: Kernels Helion Test
key: kernels-helion-test
timeout_in_minutes: 30
device: h100
source_file_dependencies:
@@ -163,6 +183,7 @@ steps:
- label: Kernels FP8 MoE Test (1 H100)
key: kernels-fp8-moe-test-1-h100
timeout_in_minutes: 90
device: h100
num_devices: 1
@@ -179,6 +200,7 @@ steps:
- pytest -v -s kernels/moe/test_triton_moe_ptpc_fp8.py
- label: Kernels FP8 MoE Test (2 H100s)
key: kernels-fp8-moe-test-2-h100s
timeout_in_minutes: 90
device: h100
num_devices: 2
@@ -188,6 +210,7 @@ steps:
- pytest -v -s kernels/moe/test_deepep_moe.py
- label: Kernels Fp4 MoE Test (B200)
key: kernels-fp4-moe-test-b200
timeout_in_minutes: 60
device: b200
num_devices: 1
@@ -200,6 +223,7 @@ steps:
- label: Kernels FusedMoE Layer Test (2 H100s)
key: kernels-fusedmoe-layer-test-2-h100s
timeout_in_minutes: 90
device: h100
num_devices: 2
@@ -216,6 +240,7 @@ steps:
- label: Kernels FusedMoE Layer Test (2 B200s)
key: kernels-fusedmoe-layer-test-2-b200s
timeout_in_minutes: 90
device: b200
num_devices: 2
+11
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: LM Eval Small Models
key: lm-eval-small-models
timeout_in_minutes: 75
source_file_dependencies:
- csrc/
@@ -24,6 +25,7 @@ steps:
# - pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large.txt --tp-size=4
- label: LM Eval Large Models (4 GPUs)(H100)
key: lm-eval-large-models-4-gpus-h100
device: h100
optional: true
num_devices: 4
@@ -36,6 +38,7 @@ steps:
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large-hopper.txt --tp-size=4
- label: LM Eval Small Models (B200)
key: lm-eval-small-models-b200
timeout_in_minutes: 120
device: b200
optional: true
@@ -46,6 +49,7 @@ steps:
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-blackwell.txt
- label: LM Eval Qwen3.5 Models (B200)
key: lm-eval-qwen3-5-models-b200
timeout_in_minutes: 120
device: b200
optional: true
@@ -62,6 +66,7 @@ steps:
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-qwen35-blackwell.txt
- label: LM Eval Large Models (H200)
key: lm-eval-large-models-h200
timeout_in_minutes: 60
device: h200
optional: true
@@ -70,6 +75,7 @@ steps:
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-h200.txt
- label: MoE Refactor Integration Test (H100 - TEMPORARY)
key: moe-refactor-integration-test-h100-temporary
device: h100
optional: true
num_devices: 2
@@ -77,6 +83,7 @@ steps:
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=evals/gsm8k/configs/moe-refactor/config-h100.txt
- label: MoE Refactor Integration Test (B200 - TEMPORARY)
key: moe-refactor-integration-test-b200-temporary
device: b200
optional: true
num_devices: 2
@@ -84,6 +91,7 @@ steps:
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=evals/gsm8k/configs/moe-refactor/config-b200.txt
- label: MoE Refactor Integration Test (B200 DP - TEMPORARY)
key: moe-refactor-integration-test-b200-dp-temporary
device: b200
optional: true
num_devices: 2
@@ -92,6 +100,7 @@ steps:
- label: LM Eval TurboQuant KV Cache
key: lm-eval-turboquant-kv-cache
timeout_in_minutes: 75
source_file_dependencies:
- vllm/model_executor/layers/quantization/turboquant/
@@ -102,6 +111,7 @@ steps:
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=evals/gsm8k/configs/models-turboquant.txt
- label: GPQA Eval (GPT-OSS) (H100)
key: gpqa-eval-gpt-oss-h100
timeout_in_minutes: 120
device: h100
optional: true
@@ -115,6 +125,7 @@ steps:
- pytest -s -v evals/gpt_oss/test_gpqa_correctness.py --config-list-file=configs/models-h100.txt
- label: GPQA Eval (GPT-OSS) (B200)
key: gpqa-eval-gpt-oss-b200
timeout_in_minutes: 120
device: b200
optional: true
+2
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: LoRA %N
key: lora
timeout_in_minutes: 30
source_file_dependencies:
- vllm/lora
@@ -13,6 +14,7 @@ steps:
- label: LoRA TP (Distributed)
key: lora-tp-distributed
timeout_in_minutes: 30
num_devices: 4
source_file_dependencies:
+14 -1
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: V1 Spec Decode
key: v1-spec-decode
timeout_in_minutes: 30
source_file_dependencies:
- vllm/
@@ -18,6 +19,7 @@ steps:
- image-build-amd
- label: V1 Sample + Logits
key: v1-sample-logits
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -41,6 +43,7 @@ steps:
- image-build-amd
- label: V1 Core + KV + Metrics
key: v1-core-kv-metrics
timeout_in_minutes: 30
source_file_dependencies:
- vllm/
@@ -71,6 +74,7 @@ steps:
- image-build-amd
- label: V1 Others (CPU)
key: v1-others-cpu
depends_on:
- image-build-cpu
source_file_dependencies:
@@ -86,6 +90,7 @@ steps:
- pytest -v -s -m 'cpu_test' v1/metrics
- label: Regression
key: regression
timeout_in_minutes: 20
device: h200_18gb
source_file_dependencies:
@@ -97,6 +102,7 @@ steps:
working_dir: "/vllm-workspace/tests" # optional
- label: Examples
key: examples
timeout_in_minutes: 45
working_dir: "/vllm-workspace/examples"
source_file_dependencies:
@@ -128,6 +134,7 @@ steps:
- python3 features/speculative_decoding/spec_decode_offline.py --test --method eagle3 --num_spec_tokens 3 --dataset-name hf --dataset-path philschmid/mt-bench --num-prompts 80 --temp 0 --top-p 1.0 --top-k -1 --tp 1 --enable-chunked-prefill --max-model-len 1536
- label: Metrics, Tracing (2 GPUs)
key: metrics-tracing-2-gpus
timeout_in_minutes: 20
num_devices: 2
source_file_dependencies:
@@ -142,6 +149,7 @@ steps:
- pytest -v -s v1/tracing
- label: Python-only Installation
key: python-only-installation
depends_on: ~
timeout_in_minutes: 20
source_file_dependencies:
@@ -151,6 +159,7 @@ steps:
- bash standalone_tests/python_only_compile.sh
- label: Async Engine, Inputs, Utils, Worker
key: async-engine-inputs-utils-worker
timeout_in_minutes: 50
source_file_dependencies:
- vllm/
@@ -163,7 +172,8 @@ steps:
- pytest -v -s utils_
- label: Async Engine, Inputs, Utils, Worker, Config (CPU)
depends_on:
key: async-engine-inputs-utils-worker-config-cpu
depends_on:
- image-build-cpu
timeout_in_minutes: 30
source_file_dependencies:
@@ -196,6 +206,7 @@ steps:
- pytest -v -s config
- label: Batch Invariance (H100)
key: batch-invariance-h100
timeout_in_minutes: 30
device: h100
source_file_dependencies:
@@ -211,6 +222,7 @@ steps:
- VLLM_TEST_MODEL=Qwen/Qwen3-30B-A3B-Thinking-2507-FP8 pytest -v -s v1/determinism/test_batch_invariance.py::test_v1_generation_is_deterministic_across_batch_sizes_with_needle[FLASH_ATTN]
- label: Batch Invariance (B200)
key: batch-invariance-b200
timeout_in_minutes: 30
device: b200
source_file_dependencies:
@@ -227,6 +239,7 @@ steps:
- pytest -v -s v1/determinism/test_nvfp4_batch_invariant.py
- label: Acceptance Length Test (Large Models) # optional
key: acceptance-length-test-large-models
timeout_in_minutes: 25
gpu: h100
optional: true
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Model Executor
key: model-executor
timeout_in_minutes: 35
source_file_dependencies:
- vllm/engine/arg_utils.py
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Model Runner V2 Core Tests
key: model-runner-v2-core-tests
timeout_in_minutes: 45
source_file_dependencies:
- vllm/v1/worker/gpu/
@@ -25,6 +26,7 @@ steps:
- pytest -v -s entrypoints/llm/test_struct_output_generate.py -k "xgrammar and not speculative_config6 and not speculative_config7 and not speculative_config8 and not speculative_config0"
- label: Model Runner V2 Examples
key: model-runner-v2-examples
timeout_in_minutes: 45
working_dir: "/vllm-workspace/examples"
source_file_dependencies:
@@ -60,6 +62,7 @@ steps:
- python3 features/speculative_decoding/spec_decode_offline.py --test --method eagle3 --num_spec_tokens 3 --dataset-name hf --dataset-path philschmid/mt-bench --num-prompts 80 --temp 0 --top-p 1.0 --top-k -1 --tp 1 --enable-chunked-prefill --max-model-len 1536
- label: Model Runner V2 Distributed (2 GPUs)
key: model-runner-v2-distributed-2-gpus
timeout_in_minutes: 45
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -80,6 +83,7 @@ steps:
- TP_SIZE=1 DP_SIZE=2 pytest -v -s v1/distributed/test_eagle_dp.py
- label: Model Runner V2 Pipeline Parallelism (4 GPUs)
key: model-runner-v2-pipeline-parallelism-4-gpus
timeout_in_minutes: 60
working_dir: "/vllm-workspace/tests"
num_devices: 4
@@ -95,6 +99,7 @@ steps:
- pytest -v -s distributed/test_pp_cudagraph.py -k "not ray"
- label: Model Runner V2 Spec Decode
key: model-runner-v2-spec-decode
timeout_in_minutes: 30
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
+6
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Basic Models Tests (Initialization)
key: basic-models-tests-initialization
timeout_in_minutes: 45
torch_nightly: true
source_file_dependencies:
@@ -16,6 +17,7 @@ steps:
torch_nightly: {}
- label: Basic Models Tests (Extra Initialization) %N
key: basic-models-tests-extra-initialization
timeout_in_minutes: 45
source_file_dependencies:
- vllm/model_executor/models/
@@ -31,6 +33,7 @@ steps:
torch_nightly: {}
- label: Basic Models Tests (Other)
key: basic-models-tests-other
timeout_in_minutes: 45
source_file_dependencies:
- vllm/
@@ -47,6 +50,7 @@ steps:
- label: Basic Models Test (Other CPU) # 5min
key: basic-models-test-other-cpu
depends_on:
- image-build-cpu
timeout_in_minutes: 10
@@ -59,6 +63,7 @@ steps:
- pytest -v -s models/test_utils.py models/test_vision.py
- label: Transformers Nightly Models
key: transformers-nightly-models
working_dir: "/vllm-workspace/"
optional: true
soft_fail: true
@@ -74,6 +79,7 @@ steps:
- VLLM_WORKER_MULTIPROC_METHOD=spawn python3 examples/generate/multimodal/audio_language_offline.py --model-type whisper
- label: Transformers Backward Compatibility Models Test
key: transformers-backward-compatibility-models-test
working_dir: "/vllm-workspace/"
optional: true
soft_fail: true
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Distributed Model Tests (2 GPUs)
key: distributed-model-tests-2-gpus
timeout_in_minutes: 50
working_dir: "/vllm-workspace/tests"
num_devices: 2
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Language Models Tests (Standard)
key: language-models-tests-standard
timeout_in_minutes: 25
source_file_dependencies:
- vllm/
@@ -15,6 +16,7 @@ steps:
torch_nightly: {}
- label: Language Models Tests (Extra Standard) %N
key: language-models-tests-extra-standard
timeout_in_minutes: 45
source_file_dependencies:
- vllm/model_executor/models/
@@ -31,6 +33,7 @@ steps:
torch_nightly: {}
- label: Language Models Tests (Hybrid) %N
key: language-models-tests-hybrid
timeout_in_minutes: 75
source_file_dependencies:
- vllm/
@@ -47,6 +50,7 @@ steps:
torch_nightly: {}
- label: Language Models Test (Extended Generation) # 80min
key: language-models-test-extended-generation
timeout_in_minutes: 110
optional: true
source_file_dependencies:
@@ -69,6 +73,7 @@ steps:
- pytest -v -s models/language/generation -m '(not core_model) and (not hybrid_model)'
- label: Language Models Test (PPL)
key: language-models-test-ppl
timeout_in_minutes: 110
device: h200_18gb
optional: true
@@ -79,6 +84,7 @@ steps:
- pytest -v -s models/language/generation_ppl_test
- label: Language Models Test (Extended Pooling) # 36min
key: language-models-test-extended-pooling
timeout_in_minutes: 50
optional: true
source_file_dependencies:
@@ -93,6 +99,7 @@ steps:
- image-build-amd
- label: Language Models Test (MTEB)
key: language-models-test-mteb
timeout_in_minutes: 110
device: h200_18gb
optional: true
+12 -1
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: "Multi-Modal Models (Standard) 1: qwen2"
key: multi-modal-models-standard-1-qwen2
timeout_in_minutes: 45
device: h200_18gb
source_file_dependencies:
@@ -19,6 +20,7 @@ steps:
- image-build-amd
- label: "Multi-Modal Models (Standard) 2: qwen3 + gemma"
key: multi-modal-models-standard-2-qwen3-gemma
timeout_in_minutes: 45
device: h200_18gb
source_file_dependencies:
@@ -36,6 +38,7 @@ steps:
- image-build-amd
- label: "Multi-Modal Models (Standard) 3: llava + qwen2_vl"
key: multi-modal-models-standard-3-llava-qwen2-vl
timeout_in_minutes: 45
source_file_dependencies:
- vllm/
@@ -51,6 +54,7 @@ steps:
- image-build-amd
- label: "Multi-Modal Models (Standard) 4: other + whisper"
key: multi-modal-models-standard-4-other-whisper
timeout_in_minutes: 45
source_file_dependencies:
- vllm/
@@ -67,7 +71,8 @@ steps:
- image-build-amd
- label: Multi-Modal Processor (CPU)
depends_on:
key: multi-modal-processor-cpu
depends_on:
- image-build-cpu
timeout_in_minutes: 60
source_file_dependencies:
@@ -80,6 +85,7 @@ steps:
- pytest -v -s models/multimodal/processing --ignore models/multimodal/processing/test_tensor_schema.py
- label: Multi-Modal Processor # 44min
key: multi-modal-processor
timeout_in_minutes: 60
device: h200_18gb
source_file_dependencies:
@@ -91,6 +97,7 @@ steps:
- pytest -v -s models/multimodal/processing/test_tensor_schema.py
- label: Multi-Modal Accuracy Eval (Small Models) # 50min
key: multi-modal-accuracy-eval-small-models
timeout_in_minutes: 70
working_dir: "/vllm-workspace/.buildkite/lm-eval-harness"
source_file_dependencies:
@@ -101,6 +108,7 @@ steps:
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-mm-small.txt --tp-size=1
- label: Multi-Modal Models (Extended Generation 1)
key: multi-modal-models-extended-generation-1
optional: true
source_file_dependencies:
- vllm/
@@ -117,6 +125,7 @@ steps:
- image-build-amd
- label: Multi-Modal Models (Extended Generation 2)
key: multi-modal-models-extended-generation-2
optional: true
source_file_dependencies:
- vllm/
@@ -126,6 +135,7 @@ steps:
- pytest -v -s models/multimodal/generation/test_common.py -m 'split(group=0) and not core_model'
- label: Multi-Modal Models (Extended Generation 3)
key: multi-modal-models-extended-generation-3
optional: true
source_file_dependencies:
- vllm/
@@ -135,6 +145,7 @@ steps:
- pytest -v -s models/multimodal/generation/test_common.py -m 'split(group=1) and not core_model'
- label: Multi-Modal Models (Extended Pooling)
key: multi-modal-models-extended-pooling
optional: true
device: h200_18gb
source_file_dependencies:
+1
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Plugin Tests (2 GPUs)
key: plugin-tests-2-gpus
timeout_in_minutes: 60
working_dir: "/vllm-workspace/tests"
num_devices: 2
+6
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: PyTorch Compilation Unit Tests
key: pytorch-compilation-unit-tests
timeout_in_minutes: 10
source_file_dependencies:
- vllm/
@@ -18,6 +19,7 @@ steps:
- "find compile/ -maxdepth 1 -name 'test_*.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'"
- label: PyTorch Compilation Unit Tests (H100)
key: pytorch-compilation-unit-tests-h100
timeout_in_minutes: 30
device: h100
num_devices: 1
@@ -28,6 +30,7 @@ steps:
- "find compile/h100/ -name 'test_*.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'"
- label: PyTorch Compilation Passes Unit Tests
key: pytorch-compilation-passes-unit-tests
timeout_in_minutes: 20
source_file_dependencies:
- vllm/
@@ -36,6 +39,7 @@ steps:
- pytest -s -v compile/passes --ignore compile/passes/distributed
- label: PyTorch Fullgraph Smoke Test
key: pytorch-fullgraph-smoke-test
timeout_in_minutes: 35
source_file_dependencies:
- vllm/
@@ -48,6 +52,7 @@ steps:
- "find compile/fullgraph/ -name 'test_*.py' -not -name 'test_full_graph.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'"
- label: PyTorch Fullgraph
key: pytorch-fullgraph
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -58,6 +63,7 @@ steps:
- pytest -v -s compile/fullgraph/test_full_graph.py -k 'not test_fp8_kv_scale_compile'
- label: Pytorch Nightly Dependency Override Check # 2min
key: pytorch-nightly-dependency-override-check
# if this test fails, it means the nightly torch version is not compatible with some
# of the dependencies. Please check the error message and add the package to whitelist
# in /vllm/tools/pre_commit/generate_nightly_torch_test.py
+3
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Quantization
key: quantization
timeout_in_minutes: 90
source_file_dependencies:
- csrc/
@@ -21,6 +22,7 @@ steps:
- VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py
- label: Quantized MoE Test (B200)
key: quantized-moe-test-b200
timeout_in_minutes: 60
working_dir: "/vllm-workspace/"
device: b200
@@ -38,6 +40,7 @@ steps:
- pytest -s -v tests/quantization/test_blackwell_moe.py
- label: Quantized Models Test
key: quantized-models-test
timeout_in_minutes: 60
source_file_dependencies:
- vllm/model_executor/layers/quantization
+1
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Ray Dependency Compatibility Check
key: ray-dependency-compatibility-check
# Informational only — does not block the pipeline.
# If this fails, it means the PR introduces a dependency that
# conflicts with Ray's dependency constraints.
+1
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Samplers Test
key: samplers-test
timeout_in_minutes: 75
source_file_dependencies:
- vllm/model_executor/layers
+21
View File
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Spec Decode Eagle
key: spec-decode-eagle
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -13,6 +14,7 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "eagle_correctness"
- label: Spec Decode Eagle Nightly B200
key: spec-decode-eagle-nightly-b200
timeout_in_minutes: 30
device: b200
optional: true
@@ -24,6 +26,7 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "eagle_correctness"
- label: Spec Decode Speculators + MTP
key: spec-decode-speculators-mtp
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -35,6 +38,7 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "speculators or mtp_correctness"
- label: Spec Decode Speculators + MTP Nightly B200
key: spec-decode-speculators-mtp-nightly-b200
timeout_in_minutes: 30
device: b200
optional: true
@@ -47,6 +51,7 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "speculators or mtp_correctness"
- label: Spec Decode Ngram + Suffix
key: spec-decode-ngram-suffix
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -57,6 +62,7 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "ngram or suffix"
- label: Spec Decode Draft Model
key: spec-decode-draft-model
timeout_in_minutes: 30
device: h200_18gb
source_file_dependencies:
@@ -67,6 +73,7 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "draft_model or no_sync or batch_inference"
- label: Spec Decode Draft Model Nightly B200
key: spec-decode-draft-model-nightly-b200
timeout_in_minutes: 30
device: b200
optional: true
@@ -78,6 +85,7 @@ steps:
- pytest -v -s v1/e2e/spec_decode -k "draft_model or no_sync or batch_inference"
- label: DFlash Speculators Correctness
key: dflash-speculators-correctness
timeout_in_minutes: 30
device: h100
optional: true
@@ -89,3 +97,16 @@ steps:
commands:
- export VLLM_ALLOW_INSECURE_SERIALIZATION=1
- pytest -v -s v1/spec_decode/test_speculators_dflash.py -m slow_test
- label: Spec Decode MTP hybrid (B200)
timeout_in_minutes: 30
device: b200
optional: true
source_file_dependencies:
- vllm/v1/spec_decode/
- vllm/v1/worker/gpu/spec_decode/
- vllm/model_executor/models/qwen3_5.py
- vllm/model_executor/models/qwen3_5_mtp.py
- tests/v1/e2e/spec_decode/
commands:
- pytest -v -s v1/e2e/spec_decode -k "qwen3_5-hybrid"
@@ -3,6 +3,7 @@ depends_on:
- image-build
steps:
- label: Weight Loading Multiple GPU # 33min
key: weight-loading-multiple-gpu
timeout_in_minutes: 45
working_dir: "/vllm-workspace/tests"
num_devices: 2
+2 -11
View File
@@ -16,11 +16,7 @@ permissions:
jobs:
pre-run-check:
if: >-
github.event_name == 'pull_request' &&
(github.event.action != 'labeled' ||
github.event.label.name == 'ready' ||
github.event.label.name == 'verified')
if: github.event_name == 'pull_request'
runs-on: ubuntu-latest
steps:
- name: Check PR label and author merge count
@@ -49,12 +45,7 @@ jobs:
pre-commit:
needs: pre-run-check
if: >-
always() &&
(github.event.action != 'labeled' ||
github.event.label.name == 'ready' ||
github.event.label.name == 'verified') &&
(needs.pre-run-check.result == 'success' || needs.pre-run-check.result == 'skipped')
if: always() && (needs.pre-run-check.result == 'success' || needs.pre-run-check.result == 'skipped')
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@8e8c483db84b4bee98b60c0593521ed34d9990e8 # v6.0.1
@@ -217,6 +217,7 @@ async def send_request(
min_tokens: int | None = None,
max_tokens: int | None = None,
timeout_sec: int = 120,
conversation_id: str | None = None,
) -> ServerResponse:
payload = {
"model": model,
@@ -225,6 +226,9 @@ async def send_request(
"temperature": 0.0,
}
if conversation_id is not None:
payload["conversation_id"] = conversation_id
if stream:
payload["stream"] = True
payload["stream_options"] = {"include_usage": False}
@@ -419,6 +423,7 @@ async def send_turn(
min_tokens,
max_tokens,
req_args.timeout_sec,
conversation_id=conv_id,
)
if response.valid is False:
+59 -4
View File
@@ -82,18 +82,73 @@ void launch_persistent_topk(const torch::Tensor& logits,
size_t smem_size = P::kFixedSmemLarge + chunk_size * sizeof(uint32_t);
if (smem_size < P::kSmemMedium) smem_size = P::kSmemMedium;
// Query occupancy for the instantiation that will actually launch;
// overestimating it deadlocks the cooperative barrier.
int occupancy = 1;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occupancy, P::persistent_topk_kernel<TopK, 4>, P::kThreadsPerBlock,
smem_size);
cudaError_t occ_err = cudaSuccess;
if (vec_size == 4) {
occ_err = cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occupancy, P::persistent_topk_kernel<TopK, 4>, P::kThreadsPerBlock,
smem_size);
} else if (vec_size == 2) {
occ_err = cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occupancy, P::persistent_topk_kernel<TopK, 2>, P::kThreadsPerBlock,
smem_size);
} else {
occ_err = cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occupancy, P::persistent_topk_kernel<TopK, 1>, P::kThreadsPerBlock,
smem_size);
}
TORCH_CHECK(occ_err == cudaSuccess,
"persistent_topk occupancy query failed: ",
cudaGetErrorString(occ_err));
if (occupancy < 1) occupancy = 1;
uint32_t max_resident_ctas = static_cast<uint32_t>(num_sms) * occupancy;
// The cooperative spin-wait barrier only runs when at least one row hits
// the radix path (seq_len > RADIX_THRESHOLD). Below that, non-CTA-0 CTAs
// early-exit, so oversubscription can't deadlock and headroom is wasted.
const bool needs_cooperative =
static_cast<uint32_t>(max_seq_len) > P::RADIX_THRESHOLD;
const uint32_t hw_resident_cap =
static_cast<uint32_t>(num_sms) * static_cast<uint32_t>(occupancy);
uint32_t max_resident_ctas = hw_resident_cap;
if (needs_cooperative) {
// Reserve one CTA per SM when occupancy allows; fall back to a single
// CTA when occupancy == 1 (the most deadlock-prone case — any straggler
// kernel that takes the only slot on one SM hangs the barrier). Never
// drop below one full group's worth.
uint32_t headroom = (occupancy > 1) ? static_cast<uint32_t>(num_sms) : 1u;
if (max_resident_ctas >= headroom + ctas_per_group) {
max_resident_ctas -= headroom;
}
}
uint32_t num_groups = std::min(max_resident_ctas / ctas_per_group,
static_cast<uint32_t>(num_rows));
if (num_groups == 0) num_groups = 1;
uint32_t total_ctas = num_groups * ctas_per_group;
// If the cooperative launch wouldn't fit, fall back to FilteredTopK
// instead of deadlocking. Only relevant when needs_cooperative.
if (needs_cooperative && total_ctas > hw_resident_cap) {
TORCH_CHECK(max_smem_per_block >= 128 * 1024,
"persistent_topk would oversubscribe and the FilteredTopK "
"fallback requires >=128KB smem per block (have ",
max_smem_per_block, "). total_ctas=", total_ctas,
" > num_sms*occupancy=", hw_resident_cap, " (TopK=", TopK,
", vec_size=", vec_size, ", ctas_per_group=", ctas_per_group,
", smem=", smem_size, ").");
cudaError_t status =
vllm::FilteredTopKRaggedTransform<float, int32_t, TopK>(
logits.data_ptr<float>(), output.data_ptr<int32_t>(),
lengths.data_ptr<int32_t>(), static_cast<uint32_t>(num_rows),
static_cast<uint32_t>(TopK), static_cast<uint32_t>(stride),
stream);
TORCH_CHECK(status == cudaSuccess,
"FilteredTopK fallback failed: ", cudaGetErrorString(status));
return;
}
size_t state_bytes = num_groups * sizeof(P::RadixRowState);
TORCH_CHECK(workspace.size(0) >= static_cast<int64_t>(state_bytes),
"workspace too small, need ", state_bytes, " bytes");
+15 -8
View File
@@ -200,9 +200,9 @@ RUN cd /opt/rixl && \
# DeepEP build stage
FROM base AS build_deep
ARG ROCSHMEM_BRANCH="ba0bf0f3"
ARG ROCSHMEM_BRANCH="f0acb0c6"
ARG ROCSHMEM_REPO="https://github.com/ROCm/rocm-systems.git"
ARG DEEPEP_BRANCH="5d90af8b"
ARG DEEPEP_BRANCH="a9ea9774"
ARG DEEPEP_REPO="https://github.com/ROCm/DeepEP.git"
ARG DEEPEP_NIC="cx7"
ARG DEEPEP_ROCM_ARCH="gfx942;gfx950"
@@ -213,18 +213,15 @@ RUN git clone ${ROCSHMEM_REPO} \
&& git checkout ${ROCSHMEM_BRANCH} \
&& mkdir -p projects/rocshmem/build \
&& cd projects/rocshmem/build \
&& bash ../scripts/build_configs/all_backends \
-DCMAKE_INSTALL_PREFIX="${ROCSHMEM_DIR}" \
-DROCM_PATH=/opt/rocm \
-DGPU_TARGETS="${DEEPEP_ROCM_ARCH}" \
-DUSE_EXTERNAL_MPI=OFF
&& INSTALL_PREFIX=${ROCSHMEM_DIR} \
../scripts/build_configs/all_backends -DUSE_EXTERNAL_MPI=OFF
# Build DeepEP wheel.
# DeepEP looks for rocshmem at ROCSHMEM_DIR.
RUN git clone ${DEEPEP_REPO} \
&& cd DeepEP \
&& git checkout ${DEEPEP_BRANCH} \
&& python3 setup.py --variant rocm --nic ${DEEPEP_NIC} bdist_wheel --dist-dir=/app/deep_install
&& python3 setup.py --variant rocm --rocm-explicit-ctx --nic ${DEEPEP_NIC} bdist_wheel --dist-dir=/app/deep_install
# MoRI runtime dependencies live in Dockerfile.rocm so NIC backend changes do
# not force users to rebuild the long-lived Dockerfile.rocm_base image.
@@ -388,6 +385,16 @@ RUN --mount=type=bind,from=export_vllm,src=/,target=/install \
# above are not available once that RUN step completes.
COPY --from=export_vllm /*.whl /opt/vllm-wheels/
# Update rdma-core to support latest rocshmem
ARG DEEPEP_NIC
RUN if [ "${DEEPEP_NIC}" = "cx7" ] || [ "${DEEPEP_NIC}" = "io" ]; then \
git clone --branch v62.0 --depth 1 https://github.com/linux-rdma/rdma-core.git /tmp/rdma-core && \
cd /tmp/rdma-core && \
mkdir -p build && cd build && \
cmake -GNinja -DCMAKE_INSTALL_PREFIX=/usr -DNO_MAN_PAGES=1 .. && \
ninja && ninja install && ldconfig && rm -rf /tmp/rdma-core; \
fi
# Install RIXL wheel
RUN --mount=type=bind,from=build_rixl,src=/app/install,target=/rixl_install \
uv pip install --system /rixl_install/*.whl
+14 -5
View File
@@ -5,9 +5,6 @@ WORKDIR /workspace/
ARG PYTHON_VERSION=3.12
ARG PIP_EXTRA_INDEX_URL="https://download.pytorch.org/whl/xpu"
RUN wget -O- https://apt.repos.intel.com/intel-gpg-keys/GPG-PUB-KEY-INTEL-SW-PRODUCTS.PUB | gpg --dearmor | tee /usr/share/keyrings/oneapi-archive-keyring.gpg > /dev/null && \
echo "deb [signed-by=/usr/share/keyrings/oneapi-archive-keyring.gpg] https://apt.repos.intel.com/oneapi all main" | tee /etc/apt/sources.list.d/oneAPI.list
RUN apt clean && apt-get update -y && \
apt-get install -y --no-install-recommends --fix-missing \
curl \
@@ -26,8 +23,20 @@ RUN apt clean && apt-get update -y && \
python3.12-dev \
python3-pip
RUN apt update && apt upgrade -y && \
apt install -y intel-oneapi-compiler-dpcpp-cpp-2025.3
# Add oneAPI repo, pin oneAPI to 2025.3, then install pinned packages in one layer.
RUN wget -O- https://apt.repos.intel.com/intel-gpg-keys/GPG-PUB-KEY-INTEL-SW-PRODUCTS.PUB | gpg --dearmor | tee /usr/share/keyrings/oneapi-archive-keyring.gpg > /dev/null && \
echo "deb [signed-by=/usr/share/keyrings/oneapi-archive-keyring.gpg] https://apt.repos.intel.com/oneapi all main" | tee /etc/apt/sources.list.d/oneAPI.list && \
printf '%s\n' \
'Package: intel-oneapi-* intel-deep-learning-essentials* intel-pti*' \
'Pin: version 2025.3*' \
'Pin-Priority: 1001' \
> /etc/apt/preferences.d/oneapi-2025.3.pref && \
apt-get update -y && \
apt-get install -y --no-install-recommends \
intel-oneapi-compiler-dpcpp-cpp-2025.3 \
intel-oneapi-mkl-devel-2025.3 \
intel-oneapi-dnnl-devel-2025.3 && \
rm -rf /var/lib/apt/lists/*
# Install UMD
RUN mkdir neo && \
+54
View File
@@ -0,0 +1,54 @@
# IndexCache
IndexCache reduces redundant top-k computation in DeepSeek-V3.2 (DSA) models by caching and reusing top-k indices across layers.
## Background
DeepSeek-V3.2 uses a DeepSeek Sparse Attention (DSA) mechanism where top-k token selection is computed per layer. For deep models with many layers, this computation can be expensive. IndexCache allows skipping redundant top-k computations by reusing indices from previous layers.
See: [IndexCache Paper](https://arxiv.org/abs/2603.12201)
## Usage
### CLI
```bash
vllm serve deepseek-ai/DeepSeek-V3.2 \
--hf-overrides '{"use_index_cache": true, "index_topk_freq": 4}' ...
```
### Configuration Reference
| Parameter | Type | Default | Description |
|----------------------|------|---------|--------------------------------------------------------------------------------------------------------------------------------------------------|
| `use_index_cache` | bool | false | Enable IndexCache. Must be set to true to use this feature |
| `index_topk_freq` | int | 1 | Frequency (in layers) at which top-k is computed. 1 = compute on every layer (disabled), 4 = compute on 1/4 of layers |
| `index_topk_pattern` | str | null | Per-layer F/S pattern. Overrides index_topk_freq if set. Each character maps to one DSA layer: F = Full, S = Shared |
### Configuration Examples
**Using `index_topk_freq`** (compute every N layers):
```bash
vllm serve deepseek-ai/DeepSeek-V3.2 \
--hf-overrides '{"use_index_cache": true, "index_topk_freq": 4}' ...
```
**Using `index_topk_pattern`** (explicit per-layer control):
```bash
# custom pattern for 61 layers: F = compute, S = reuse
vllm serve deepseek-ai/DeepSeek-V3.2 \
--hf-overrides '{"use_index_cache": true, "index_topk_pattern": "FFSFSSSFSSFFFSSSFFFSFSSSSSSFFSFFSFFSSFFFFFFSFFFFFSFFSSSSSSFSF"}'
```
## How It Works
1. When IndexCache is enabled, layers marked with `"F"` (Full) calculate and store top-k indices
2. Subsequent layers marked with `"S"` (Shared) receive the cached indices from the previous layer instead of recomputing
3. The cached indices are passed through the layer stack, reducing total computation
## Requirements
- DeepSeek-V3.2 or compatible DSA model
- `use_index_cache: true` via `--hf-overrides`
+13 -1
View File
@@ -38,8 +38,20 @@ class MockCustomOp:
return decorator
class MockPluggableLayer:
@staticmethod
def register(name):
def decorator(cls):
return cls
return decorator
mock_if_no_torch("vllm._C", MagicMock())
mock_if_no_torch("vllm.model_executor.custom_op", MagicMock(CustomOp=MockCustomOp))
mock_if_no_torch(
"vllm.model_executor.custom_op",
MagicMock(CustomOp=MockCustomOp, PluggableLayer=MockPluggableLayer),
)
mock_if_no_torch(
"vllm.utils.torch_utils", MagicMock(direct_register_custom_op=lambda *a, **k: None)
)
@@ -0,0 +1,562 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Disaggregated Prefill/Decode Proxy with Bidirectional KV Transfer
This proxy sits between clients and a vLLM Prefill/Decode (P/D) deployment,
routing multi-turn chat requests so that each turn reuses KV cache blocks
from the previous turn's Decode node via bidirectional KV transfer.
Architecture:
Client ──► Proxy ──► Prefill (P) ──► Decode (D)
│ │ │
│ kv_transfer_params flow: │
│ D finish ──► proxy caches │
│ next turn ──► proxy sends │
│ cached D blocks to P ──► │
│ P reads D blocks (bidir) │
│ P sends its blocks to D │
Per-request flow:
1. Client sends chat/completions request to proxy.
2. Proxy looks up cached D block info from the previous turn
(keyed by conversation_id).
3. If cache hit, proxy attaches D's block info to the request
so P can read D's KV blocks instead of recomputing.
4. Proxy sends request to P (max_tokens=1, non-streaming).
5. P returns kv_transfer_params with its own block info.
6. Proxy forwards request + P's block info to D (streaming).
7. D streams the response. The final chunk includes D's
kv_transfer_params, which the proxy caches for the next turn.
8. Proxy returns D's response to the client.
Conversation isolation:
Each request must include a ``conversation_id`` field (top-level in
the JSON body) to scope the KV cache across turns. Without it, the
proxy cannot link turns and falls back to no-cache behavior.
Usage:
python disagg_proxy_multiturn.py \\
--host 0.0.0.0 --port 8000 \\
--prefiller-host 10.0.0.1 --prefiller-port 8100 \\
--decoder-host 10.0.0.2 --decoder-port 8200
Dependencies:
pip install fastapi uvicorn httpx
"""
from __future__ import annotations
import argparse
import itertools
import json
import logging
import os
import time
import uuid
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from typing import Any
import httpx
from fastapi import FastAPI, Request
from fastapi.responses import JSONResponse, StreamingResponse
# Logging
logging.basicConfig(
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
level=logging.INFO,
)
logger = logging.getLogger("disagg_proxy")
# Data structures
@dataclass
class CachedKVEntry:
"""KV transfer parameters cached from D's response for one turn."""
kv_transfer_params: dict[str, Any]
timestamp: float = field(default_factory=time.time)
class ConversationKVCache:
"""Per-conversation KV block cache.
Each conversation is identified by a ``conversation_id`` supplied by
the client. After D finishes a turn, its ``kv_transfer_params`` are
stored here. On the next turn, the proxy retrieves them so P can
read D's blocks via bidirectional KV transfer.
"""
def __init__(self, ttl_seconds: float = 600.0) -> None:
self._store: dict[str, CachedKVEntry] = {}
self._ttl = ttl_seconds
def get(self, conversation_id: str) -> dict[str, Any] | None:
"""Retrieve and consume cached KV params for a conversation.
Returns a *copy* of the kv_transfer_params dict, or None.
The entry is removed after retrieval (single-use).
"""
entry = self._store.pop(conversation_id, None)
if entry is None:
return None
age = time.time() - entry.timestamp
if age > self._ttl:
logger.info(
"conv=%s: stale cache entry (age=%.1fs > ttl=%.1fs), discarding",
conversation_id,
age,
self._ttl,
)
return None
logger.info(
"conv=%s: cache HIT (age=%.1fs)",
conversation_id,
age,
)
return dict(entry.kv_transfer_params)
def put(self, conversation_id: str, kv_params: dict[str, Any]) -> None:
"""Store D's kv_transfer_params for a conversation."""
self._store[conversation_id] = CachedKVEntry(
kv_transfer_params=dict(kv_params), # defensive copy
)
logger.info(
"conv=%s: cached D blocks (remote_request_id=%s, blocks=%d)",
conversation_id,
kv_params.get("remote_request_id", "?"),
len(kv_params.get("remote_block_ids", [[]])[0])
if kv_params.get("remote_block_ids")
else 0,
)
def evict_stale(self) -> int:
"""Remove entries older than TTL. Returns count of evicted entries."""
now = time.time()
stale = [
cid
for cid, entry in self._store.items()
if now - entry.timestamp > self._ttl
]
for cid in stale:
del self._store[cid]
return len(stale)
@property
def size(self) -> int:
return len(self._store)
# Global state
kv_cache = ConversationKVCache(
ttl_seconds=450.0
) # Must be < VLLM_NIXL_ABORT_REQUEST_TIMEOUT (480s)
# Service client helpers
@dataclass
class ServiceClient:
"""Wrapper around an httpx.AsyncClient for a P or D instance."""
client: httpx.AsyncClient
host: str
port: int
id: int
def _make_headers(request_id: str) -> dict[str, str]:
"""Build HTTP headers for upstream requests."""
headers = {"X-Request-Id": request_id}
api_key = os.environ.get("OPENAI_API_KEY")
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
return headers
async def _send_to_prefill(
client: ServiceClient,
endpoint: str,
req_data: dict[str, Any],
request_id: str,
) -> dict[str, Any]:
"""Send a non-streaming prefill request (max_tokens=1).
Returns the JSON response from P, which includes kv_transfer_params.
"""
payload = req_data.copy()
payload["stream"] = False
payload["max_tokens"] = 1
payload.pop("max_completion_tokens", None)
payload.pop("min_tokens", None)
payload.pop("stream_options", None)
resp = await client.client.post(
endpoint,
json=payload,
headers=_make_headers(request_id),
)
resp.raise_for_status()
return resp.json()
async def _stream_from_decode(
client: ServiceClient,
endpoint: str,
req_data: dict[str, Any],
request_id: str,
conversation_id: str,
) -> tuple[str, str | None, dict[str, Any] | None, str, str | None, int | None]:
"""Stream response from D, capturing text and kv_transfer_params.
Returns (collected_text, finish_reason, kv_params, response_id, created).
Also stores kv_params in the conversation cache.
"""
payload = req_data.copy()
payload["stream"] = True
collected_text = ""
finish_reason: str | None = None
response_id: str | None = None
model_name: str | None = None
created: int | None = None
captured_kv: dict[str, Any] | None = None
async with client.client.stream(
"POST",
endpoint,
json=payload,
headers=_make_headers(request_id),
) as resp:
resp.raise_for_status()
async for line in resp.aiter_lines():
if not line or not line.startswith("data: "):
continue
if line == "data: [DONE]":
break
try:
chunk = json.loads(line[6:])
except json.JSONDecodeError:
continue
if response_id is None:
response_id = chunk.get("id")
model_name = chunk.get("model")
created = chunk.get("created")
for choice in chunk.get("choices", []):
collected_text += choice.get("text", "")
delta = choice.get("delta", {})
collected_text += delta.get("content", "")
if choice.get("finish_reason"):
finish_reason = choice["finish_reason"]
kv_params = chunk.get("kv_transfer_params")
if kv_params:
kv_params["remote_host"] = client.host
captured_kv = kv_params
if conversation_id:
kv_cache.put(conversation_id, kv_params)
return (
collected_text,
finish_reason,
captured_kv,
response_id or request_id,
model_name,
created,
)
async def _stream_from_decode_sse(
client: ServiceClient,
endpoint: str,
req_data: dict[str, Any],
request_id: str,
conversation_id: str,
):
"""Yield SSE chunks from D to the client, capturing kv_transfer_params."""
payload = req_data.copy()
payload["stream"] = True
async with client.client.stream(
"POST",
endpoint,
json=payload,
headers=_make_headers(request_id),
) as resp:
resp.raise_for_status()
async for line in resp.aiter_lines():
if not line:
yield "\n"
continue
if line.startswith("data: ") and line != "data: [DONE]":
try:
chunk = json.loads(line[6:])
kv_params = chunk.get("kv_transfer_params")
if kv_params and conversation_id:
kv_params["remote_host"] = client.host
kv_cache.put(conversation_id, kv_params)
except json.JSONDecodeError:
pass
yield line + "\n"
# FastAPI application
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Initialize HTTP clients for P and D instances."""
app.state.prefill_clients: list[ServiceClient] = []
app.state.decode_clients: list[ServiceClient] = []
for i, (host, port) in enumerate(global_args.prefiller_instances):
app.state.prefill_clients.append(
ServiceClient(
client=httpx.AsyncClient(
timeout=None,
base_url=f"http://{host}:{port}/v1",
),
host=host,
port=port,
id=i,
)
)
for i, (host, port) in enumerate(global_args.decoder_instances):
app.state.decode_clients.append(
ServiceClient(
client=httpx.AsyncClient(
timeout=None,
base_url=f"http://{host}:{port}/v1",
),
host=host,
port=port,
id=i,
)
)
app.state.prefill_iter = itertools.cycle(range(len(app.state.prefill_clients)))
app.state.decode_iter = itertools.cycle(range(len(app.state.decode_clients)))
logger.info(
"Ready: %d prefill, %d decode instances",
len(app.state.prefill_clients),
len(app.state.decode_clients),
)
yield
for sc in app.state.prefill_clients + app.state.decode_clients:
await sc.client.aclose()
app = FastAPI(title="Disaggregated P/D Proxy (Multi-turn)", lifespan=lifespan)
def _next_client(app_state, role: str) -> ServiceClient:
if role == "prefill":
return app_state.prefill_clients[next(app_state.prefill_iter)]
return app_state.decode_clients[next(app_state.decode_iter)]
# Request handler
async def _handle_request(api_path: str, request: Request):
"""Core request handler for both /v1/chat/completions and /v1/completions."""
req_data = await request.json()
request_id = str(uuid.uuid4())
conversation_id: str = req_data.pop("conversation_id", "")
client_wants_stream = req_data.get("stream", False)
if not conversation_id:
logger.warning(
"[%s] No conversation_id provided — KV cache reuse disabled "
"for this request. Add a 'conversation_id' field to enable "
"cross-turn KV sharing.",
request_id,
)
# Step 1: Look up cached D blocks from the previous turn
cached_kv = kv_cache.get(conversation_id) if conversation_id else None
if cached_kv:
# Tell P to read D's blocks (bidirectional transfer)
cached_kv["do_remote_decode"] = True
cached_kv["do_remote_prefill"] = False
req_data["kv_transfer_params"] = cached_kv
logger.info(
"[%s] conv=%s: sending D's cached blocks to P (remote_request_id=%s)",
request_id,
conversation_id,
cached_kv.get("remote_request_id"),
)
else:
# No cached blocks — P recomputes from scratch
req_data["kv_transfer_params"] = {
"do_remote_decode": True,
"do_remote_prefill": False,
"remote_engine_id": None,
"remote_block_ids": None,
"remote_host": None,
"remote_port": None,
}
logger.info("[%s] conv=%s: cache MISS", request_id, conversation_id)
# Step 2: Send to Prefill node (non-streaming, max_tokens=1)
prefill_client = _next_client(request.app.state, "prefill")
t0 = time.time()
prefill_resp = await _send_to_prefill(
prefill_client,
api_path,
req_data,
request_id,
)
logger.info(
"[%s] Prefill done in %.0fms",
request_id,
(time.time() - t0) * 1000,
)
# Attach P's kv_transfer_params for D to read P's blocks
p_kv_params = prefill_resp.get("kv_transfer_params", {})
if p_kv_params:
p_kv_params["remote_host"] = prefill_client.host
req_data["kv_transfer_params"] = p_kv_params
# Step 3: Stream from Decode node, capturing kv_transfer_params
decode_client = _next_client(request.app.state, "decode")
if client_wants_stream:
return StreamingResponse(
_stream_from_decode_sse(
decode_client,
api_path,
req_data,
request_id,
conversation_id,
),
media_type="text/event-stream",
)
text, finish_reason, _, resp_id, model, created = await _stream_from_decode(
decode_client,
api_path,
req_data,
request_id,
conversation_id,
)
# Build OpenAI-compatible response
is_chat = "messages" in req_data
if is_chat:
body = {
"id": resp_id,
"object": "chat.completion",
"created": created or int(time.time()),
"model": model or req_data.get("model", ""),
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": text},
"finish_reason": finish_reason,
}
],
"usage": None,
}
else:
body = {
"id": resp_id,
"object": "text_completion",
"created": created or int(time.time()),
"model": model or req_data.get("model", ""),
"choices": [
{
"index": 0,
"text": text,
"logprobs": None,
"finish_reason": finish_reason,
}
],
"usage": None,
}
return JSONResponse(content=body)
# Routes
@app.post("/v1/chat/completions")
async def chat_completions(request: Request):
return await _handle_request("/chat/completions", request)
@app.post("/v1/completions")
async def completions(request: Request):
return await _handle_request("/completions", request)
@app.get("/health")
async def health():
evicted = kv_cache.evict_stale()
return {
"status": "ok",
"cached_conversations": kv_cache.size,
"evicted_stale": evicted,
}
# CLI
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(
description="Disaggregated P/D proxy with bidirectional KV transfer",
)
p.add_argument("--host", default="0.0.0.0")
p.add_argument("--port", type=int, default=8000)
p.add_argument(
"--prefiller-host",
"--prefiller-hosts",
dest="prefiller_hosts",
nargs="+",
default=["localhost"],
)
p.add_argument(
"--prefiller-port",
"--prefiller-ports",
dest="prefiller_ports",
type=int,
nargs="+",
default=[8100],
)
p.add_argument(
"--decoder-host",
"--decoder-hosts",
dest="decoder_hosts",
nargs="+",
default=["localhost"],
)
p.add_argument(
"--decoder-port",
"--decoder-ports",
dest="decoder_ports",
type=int,
nargs="+",
default=[8200],
)
args = p.parse_args()
if len(args.prefiller_hosts) != len(args.prefiller_ports):
p.error("Number of prefiller hosts must match ports")
if len(args.decoder_hosts) != len(args.decoder_ports):
p.error("Number of decoder hosts must match ports")
args.prefiller_instances = list(zip(args.prefiller_hosts, args.prefiller_ports))
args.decoder_instances = list(zip(args.decoder_hosts, args.decoder_ports))
return args
if __name__ == "__main__":
global global_args
global_args = parse_args()
import uvicorn
uvicorn.run(app, host=global_args.host, port=global_args.port)
+2 -1
View File
@@ -105,7 +105,8 @@ plugins:
- https://docs.aiohttp.org/en/stable/objects.inv
- https://pillow.readthedocs.io/en/stable/objects.inv
- https://numpy.org/doc/stable/objects.inv
- https://pytorch.org/docs/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
- redirects:
redirect_maps:
features/spec_decode/README.md: features/speculative_decoding/README.md
+5 -1
View File
@@ -61,7 +61,11 @@ fastsafetensors>=0.2.2 # 0.2.2 contains important fixes for multi-GPU mem usage
instanttensor>=0.1.5
pydantic>=2.12 # 2.11 leads to error on python 3.13
decord==0.6.0; platform_machine == "x86_64"
terratorch >= 1.2.2 # Required for Prithvi tests
# terratorch is temporarily disabled while PyPI has the `lightning` package
# in `quarantined` status (every published terratorch version transitively
# requires `lightning`, so the resolver fails with "no versions of lightning").
# Re-enable once PyPI lifts the quarantine. Tracked in #41376.
# terratorch >= 1.2.2 # Required for Prithvi tests
imagehash # Required for Prithvi tests
segmentation-models-pytorch > 0.4.0 # Required for Prithvi tests
+11 -274
View File
@@ -1,15 +1,9 @@
# This file was autogenerated by uv via the following command:
# uv pip compile requirements/test/cuda.in -c requirements/cuda.txt -o requirements/test/cuda.txt --index-strategy unsafe-best-match --torch-backend cu130 --python-platform x86_64-manylinux_2_28 --python-version 3.12
absl-py==2.1.0
# via
# rouge-score
# tensorboard
# via rouge-score
accelerate==1.13.0
# via peft
aenum==3.1.16
# via lightly
affine==2.4.0
# via rasterio
aiohappyeyeballs==2.6.1
# via aiohttp
aiohttp==3.13.3
@@ -25,22 +19,14 @@ aiohttp-cors==0.8.1
# via ray
aiosignal==1.4.0
# via aiohttp
albucore==0.0.16
# via terratorch
albumentations==1.4.6
# via
# -r requirements/test/cuda.in
# terratorch
# via -r requirements/test/cuda.in
alembic==1.16.4
# via optuna
annotated-doc==0.0.4
# via fastapi
annotated-types==0.7.0
# via pydantic
antlr4-python3-runtime==4.9.3
# via
# hydra-core
# omegaconf
anyio==4.6.2.post1
# via
# httpx
@@ -54,12 +40,10 @@ arrow==1.3.0
attrs==24.2.0
# via
# aiohttp
# fiona
# hypothesis
# jsonlines
# jsonschema
# pytest-subtests
# rasterio
# referencing
audioread==3.0.1
# via librosa
@@ -78,9 +62,7 @@ backoff==2.2.1
# -r requirements/test/cuda.in
# schemathesis
bitsandbytes==0.49.2
# via
# -r requirements/test/cuda.in
# lightning
# via -r requirements/test/cuda.in
black==24.10.0
# via datamodel-code-generator
blobfile==3.0.0
@@ -103,15 +85,9 @@ cachetools==5.5.2
# via google-auth
certifi==2024.8.30
# via
# fiona
# httpcore
# httpx
# lightly
# pyogrio
# pyproj
# rasterio
# requests
# sentry-sdk
cffi==2.0.0
# via
# cryptography
@@ -125,25 +101,12 @@ chz==0.3.0
click==8.1.7
# via
# black
# click-plugins
# cligj
# fiona
# jiwer
# nltk
# rasterio
# ray
# schemathesis
# typer
# uvicorn
# wandb
click-plugins==1.1.1.2
# via
# fiona
# rasterio
cligj==0.7.2
# via
# fiona
# rasterio
colorama==0.4.6
# via
# perceptron
@@ -191,8 +154,6 @@ decorator==5.1.1
# via librosa
decord==0.6.0
# via -r requirements/test/cuda.in
diffusers==0.36.0
# via terratorch
dill==0.3.8
# via
# datasets
@@ -207,14 +168,10 @@ docker==7.1.0
# via gpt-oss
docopt==0.6.2
# via num2words
docstring-parser==0.17.0
# via jsonargparse
einops==0.8.1
# via
# -r requirements/test/cuda.in
# encodec
# terratorch
# torchgeo
# vector-quantize-pytorch
# vocos
einx==0.3.0
@@ -244,13 +201,10 @@ filelock==3.16.1
# -c requirements/common.txt
# blobfile
# datasets
# diffusers
# huggingface-hub
# ray
# torch
# virtualenv
fiona==1.10.1
# via torchgeo
fonttools==4.55.0
# via matplotlib
fqdn==1.5.1
@@ -267,9 +221,6 @@ fsspec==2024.12.0
# evaluate
# fastparquet
# huggingface-hub
# lightning
# pytorch-lightning
# tacoreader
# torch
ftfy==6.3.1
# via open-clip-torch
@@ -277,12 +228,6 @@ genai-perf==0.0.16
# via -r requirements/test/cuda.in
genson==1.3.0
# via datamodel-code-generator
geopandas==1.0.1
# via terratorch
gitdb==4.0.12
# via gitpython
gitpython==3.1.44
# via wandb
google-api-core==2.24.2
# via
# google-cloud-core
@@ -317,7 +262,6 @@ grpcio==1.78.0
# -r requirements/test/cuda.in
# grpcio-reflection
# ray
# tensorboard
grpcio-reflection==1.78.0
# via -r requirements/test/cuda.in
h11==0.14.0
@@ -326,8 +270,6 @@ h11==0.14.0
# uvicorn
h2==4.3.0
# via httpx
h5py==3.13.0
# via terratorch
harfile==0.3.0
# via schemathesis
hf-xet==1.4.3
@@ -343,7 +285,6 @@ httpcore==1.0.6
httpx==0.27.2
# via
# -r requirements/test/cuda.in
# diffusers
# huggingface-hub
# perceptron
# schemathesis
@@ -351,23 +292,17 @@ huggingface-hub==1.10.2
# via
# accelerate
# datasets
# diffusers
# evaluate
# open-clip-torch
# peft
# segmentation-models-pytorch
# sentence-transformers
# terratorch
# timm
# tokenizers
# transformers
# vocos
humanize==4.11.0
# via runai-model-streamer
hydra-core==1.3.2
# via
# lightly
# lightning
hyperframe==6.1.0
# via h2
hypothesis==6.131.0
@@ -392,11 +327,7 @@ imagehash==4.3.2
imageio==2.37.0
# via scikit-image
importlib-metadata==8.7.0
# via
# diffusers
# opentelemetry-api
importlib-resources==6.5.2
# via typeshed-client
# via opentelemetry-api
inflect==5.6.2
# via datamodel-code-generator
iniconfig==2.0.0
@@ -426,14 +357,8 @@ joblib==1.4.2
# librosa
# nltk
# scikit-learn
jsonargparse==4.46.0
# via
# lightning
# terratorch
jsonlines==4.0.0
# via lm-eval
jsonnet==0.21.0
# via jsonargparse
jsonpointer==3.0.0
# via jsonschema
jsonschema==4.23.0
@@ -452,10 +377,6 @@ kaleido==0.2.1
# via genai-perf
kiwisolver==1.4.7
# via matplotlib
kornia==0.8.1
# via torchgeo
kornia-rs==0.1.9
# via kornia
lazy-loader==0.4
# via
# librosa
@@ -464,21 +385,6 @@ libnacl==2.1.0
# via tensorizer
librosa==0.10.2.post1
# via -r requirements/test/cuda.in
lightly==1.5.22
# via
# terratorch
# torchgeo
lightly-utils==0.0.2
# via lightly
lightning==2.6.1
# via
# terratorch
# torchgeo
lightning-utilities==0.14.3
# via
# lightning
# pytorch-lightning
# torchmetrics
llvmlite==0.47.0
# via numba
lm-eval==0.4.11
@@ -490,8 +396,6 @@ lxml==5.3.0
# sacrebleu
mako==1.3.10
# via alembic
markdown==3.8.2
# via tensorboard
markdown-it-py==3.0.0
# via rich
markupsafe==3.0.1
@@ -500,11 +404,7 @@ markupsafe==3.0.1
# mako
# werkzeug
matplotlib==3.9.2
# via
# -r requirements/test/cuda.in
# lightning
# pycocotools
# torchgeo
# via -r requirements/test/cuda.in
mbstrdecoder==1.1.3
# via
# dataproperty
@@ -559,7 +459,6 @@ numpy==2.2.6
# via
# -r requirements/test/cuda.in
# accelerate
# albucore
# albumentations
# bitsandbytes
# bm25s
@@ -567,19 +466,14 @@ numpy==2.2.6
# cupy-cuda12x
# datasets
# decord
# diffusers
# einx
# encodec
# evaluate
# fastparquet
# genai-perf
# geopandas
# h5py
# imagehash
# imageio
# librosa
# lightly
# lightly-utils
# lm-eval
# matplotlib
# mistral-common
@@ -591,11 +485,7 @@ numpy==2.2.6
# patsy
# peft
# perceptron
# pycocotools
# pyogrio
# pywavelets
# rasterio
# rioxarray
# rouge-score
# runai-model-streamer
# sacrebleu
@@ -603,21 +493,14 @@ numpy==2.2.6
# scikit-learn
# scipy
# segmentation-models-pytorch
# shapely
# soxr
# statsmodels
# tensorboard
# tensorboardx
# tensorizer
# terratorch
# tifffile
# torchgeo
# torchmetrics
# torchvision
# transformers
# tritonclient
# vocos
# xarray
nvidia-cublas==13.1.0.3
# via
# cuda-toolkit
@@ -657,10 +540,6 @@ nvidia-nvshmem-cu13==3.4.5
# via torch
nvidia-nvtx==13.0.85
# via cuda-toolkit
omegaconf==2.3.0
# via
# hydra-core
# lightning
open-clip-torch==2.32.0
# via -r requirements/test/cuda.in
openai-harmony==0.0.4
@@ -675,7 +554,6 @@ opencv-python-headless==4.13.0.90
# via
# -c requirements/common.txt
# -r requirements/test/cuda.in
# albucore
# albumentations
# mistral-common
openpyxl==3.1.5
@@ -710,44 +588,27 @@ packaging==24.2
# datasets
# evaluate
# fastparquet
# geopandas
# huggingface-hub
# hydra-core
# kornia
# lazy-loader
# lightning
# lightning-utilities
# matplotlib
# optuna
# peft
# plotly
# pooch
# pyogrio
# pytest
# pytest-rerunfailures
# pytorch-lightning
# ray
# rioxarray
# scikit-image
# statsmodels
# tensorboard
# tensorboardx
# torchmetrics
# transformers
# typepy
# wandb
# xarray
pandas==2.2.3
# via
# datasets
# evaluate
# fastparquet
# genai-perf
# geopandas
# statsmodels
# tacoreader
# torchgeo
# xarray
pathspec==0.12.1
# via black
pathvalidate==3.2.1
@@ -762,25 +623,20 @@ perf-analyzer==0.1.0
# via genai-perf
pillow==10.4.0
# via
# diffusers
# genai-perf
# imagehash
# imageio
# lightly-utils
# matplotlib
# mistral-common
# perceptron
# scikit-image
# segmentation-models-pytorch
# tensorboard
# torchgeo
# torchvision
platformdirs==4.3.6
# via
# black
# pooch
# virtualenv
# wandb
plotly==5.24.1
# via
# -r requirements/test/cuda.in
@@ -817,10 +673,7 @@ protobuf==6.33.6
# opentelemetry-proto
# proto-plus
# ray
# tensorboard
# tensorboardx
# tensorizer
# wandb
psutil==6.1.0
# via
# accelerate
@@ -834,16 +687,12 @@ pyarrow==23.0.0
# via
# datasets
# genai-perf
# tacoreader
# terratorch
pyasn1==0.6.1
# via
# pyasn1-modules
# rsa
pyasn1-modules==0.4.2
# via google-auth
pycocotools==2.0.8
# via terratorch
pycountry==24.6.1
# via pydantic-extra-types
pycparser==2.22
@@ -858,13 +707,11 @@ pydantic==2.12.0
# datamodel-code-generator
# fastapi
# gpt-oss
# lightly
# mistral-common
# mteb
# openai-harmony
# pydantic-extra-types
# ray
# wandb
pydantic-core==2.41.1
# via pydantic
pydantic-extra-types==2.10.5
@@ -873,17 +720,8 @@ pygments==2.18.0
# via rich
pyjwt==2.11.0
# via msal
pyogrio==0.11.0
# via geopandas
pyparsing==3.2.0
# via
# matplotlib
# rasterio
pyproj==3.7.1
# via
# geopandas
# rioxarray
# torchgeo
# via matplotlib
pyrate-limiter==3.7.0
# via schemathesis
pystemmer==3.0.0
@@ -920,22 +758,15 @@ pytest-subtests==0.14.1
# via schemathesis
pytest-timeout==2.3.1
# via -r requirements/test/cuda.in
python-box==7.3.2
# via terratorch
python-dateutil==2.9.0.post0
# via
# arrow
# botocore
# lightly
# matplotlib
# pandas
# typepy
python-rapidjson==1.20
# via tritonclient
pytorch-lightning==2.5.2
# via
# lightly
# lightning
pytrec-eval-terrier==0.5.7
# via mteb
pytz==2024.2
@@ -952,26 +783,16 @@ pyyaml==6.0.2
# datasets
# genai-perf
# huggingface-hub
# jsonargparse
# lightning
# omegaconf
# optuna
# peft
# pytorch-lightning
# ray
# responses
# schemathesis
# timm
# transformers
# vocos
# wandb
rapidfuzz==3.12.1
# via jiwer
rasterio==1.4.3
# via
# rioxarray
# terratorch
# torchgeo
ray==2.48.0
# via -r requirements/test/cuda.in
redis==5.2.0
@@ -982,7 +803,6 @@ referencing==0.35.1
# jsonschema-specifications
regex==2026.2.28
# via
# diffusers
# nltk
# open-clip-torch
# sacrebleu
@@ -994,13 +814,11 @@ requests==2.32.3
# azure-core
# buildkite-test-collector
# datasets
# diffusers
# docker
# evaluate
# google-api-core
# google-cloud-storage
# gpt-oss
# lightly
# lm-eval
# mistral-common
# msal
@@ -1010,9 +828,7 @@ requests==2.32.3
# responses
# schemathesis
# starlette-testclient
# tacoreader
# tiktoken
# wandb
responses==0.25.3
# via genai-perf
rfc3339-validator==0.1.4
@@ -1022,13 +838,9 @@ rfc3987==1.3.8
rich==13.9.4
# via
# genai-perf
# lightning
# mteb
# perceptron
# terratorch
# typer
rioxarray==0.19.0
# via terratorch
rouge-score==0.1.2
# via lm-eval
rpds-py==0.20.1
@@ -1037,8 +849,6 @@ rpds-py==0.20.1
# referencing
rsa==4.9.1
# via google-auth
rtree==1.4.0
# via torchgeo
runai-model-streamer==0.15.7
# via -r requirements/test/cuda.in
runai-model-streamer-azure==0.15.7
@@ -1054,7 +864,6 @@ sacrebleu==2.4.3
safetensors==0.4.5
# via
# accelerate
# diffusers
# open-clip-torch
# peft
# segmentation-models-pytorch
@@ -1063,9 +872,7 @@ safetensors==0.4.5
schemathesis==3.39.15
# via -r requirements/test/cuda.in
scikit-image==0.25.2
# via
# albumentations
# terratorch
# via albumentations
scikit-learn==1.5.2
# via
# albumentations
@@ -1073,7 +880,6 @@ scikit-learn==1.5.2
# lm-eval
# mteb
# sentence-transformers
# terratorch
scipy==1.13.1
# via
# albumentations
@@ -1087,27 +893,16 @@ scipy==1.13.1
# statsmodels
# vocos
segmentation-models-pytorch==0.5.0
# via
# -r requirements/test/cuda.in
# terratorch
# torchgeo
# via -r requirements/test/cuda.in
sentence-transformers==5.2.0
# via
# -r requirements/test/cuda.in
# mteb
sentry-sdk==2.52.0
# via wandb
setuptools==77.0.3
# via
# -c requirements/common.txt
# lightning-utilities
# pytablewriter
# tensorboard
# torch
shapely==2.1.1
# via
# geopandas
# torchgeo
shellingham==1.5.4
# via
# perceptron
@@ -1116,15 +911,12 @@ six==1.16.0
# via
# -c requirements/common.txt
# junit-xml
# lightly
# opencensus
# python-dateutil
# rfc3339-validator
# rouge-score
smart-open==7.1.0
# via ray
smmap==5.0.2
# via gitdb
sniffio==1.3.1
# via
# anyio
@@ -1166,8 +958,6 @@ tabledata==1.3.3
# via pytablewriter
tabulate==0.9.0
# via sacrebleu
tacoreader==0.5.6
# via terratorch
tblib==3.1.0
# via -r requirements/test/cuda.in
tcolorpy==0.1.6
@@ -1177,26 +967,14 @@ tenacity==9.1.2
# gpt-oss
# lm-eval
# plotly
tensorboard==2.20.0
# via terratorch
tensorboard-data-server==0.7.2
# via tensorboard
tensorboardx==2.6.4
# via lightning
tensorizer==2.10.1
# via -r requirements/test/cuda.in
termcolor==3.1.0
# via
# gpt-oss
# terratorch
terratorch==1.2.2
# via -r requirements/test/cuda.in
# via gpt-oss
threadpoolctl==3.5.0
# via scikit-learn
tifffile==2025.3.30
# via
# scikit-image
# terratorch
# via scikit-image
tiktoken==0.12.0
# via
# -c requirements/common.txt
@@ -1208,8 +986,6 @@ timm==1.0.17
# -r requirements/test/cuda.in
# open-clip-torch
# segmentation-models-pytorch
# terratorch
# torchgeo
tokenizers==0.22.2
# via
# -c requirements/common.txt
@@ -1227,21 +1003,14 @@ torch==2.11.0+cu130
# bitsandbytes
# encodec
# instanttensor
# kornia
# lightly
# lightning
# mteb
# open-clip-torch
# peft
# pytorch-lightning
# runai-model-streamer
# segmentation-models-pytorch
# sentence-transformers
# tensorizer
# terratorch
# timm
# torchgeo
# torchmetrics
# torchvision
# vector-quantize-pytorch
# vocos
@@ -1251,31 +1020,18 @@ torchaudio==2.11.0+cu130
# -r requirements/test/cuda.in
# encodec
# vocos
torchgeo==0.7.0
# via terratorch
torchmetrics==1.7.4
# via
# lightning
# pytorch-lightning
# terratorch
# torchgeo
torchvision==0.26.0+cu130
# via
# -c requirements/cuda.txt
# -r requirements/test/cuda.in
# lightly
# open-clip-torch
# segmentation-models-pytorch
# terratorch
# timm
# torchgeo
tqdm==4.67.3
# via
# datasets
# evaluate
# huggingface-hub
# lightly
# lightning
# lm-eval
# mteb
# nltk
@@ -1283,11 +1039,8 @@ tqdm==4.67.3
# optuna
# peft
# pqdm
# pytorch-lightning
# segmentation-models-pytorch
# sentence-transformers
# tacoreader
# terratorch
# transformers
transformers==5.5.3
# via
@@ -1316,8 +1069,6 @@ typer==0.15.2
# transformers
types-python-dateutil==2.9.0.20241206
# via arrow
typeshed-client==2.8.2
# via jsonargparse
typing-extensions==4.15.0
# via
# -c requirements/common.txt
@@ -1332,8 +1083,6 @@ typing-extensions==4.15.0
# grpcio
# huggingface-hub
# librosa
# lightning
# lightning-utilities
# lm-eval
# mistral-common
# mteb
@@ -1344,16 +1093,12 @@ typing-extensions==4.15.0
# pydantic
# pydantic-core
# pydantic-extra-types
# pytorch-lightning
# sentence-transformers
# sqlalchemy
# starlette
# torch
# torchgeo
# typer
# typeshed-client
# typing-inspection
# wandb
typing-inspection==0.4.2
# via pydantic
tzdata==2024.2
@@ -1365,10 +1110,8 @@ urllib3==2.2.3
# blobfile
# botocore
# docker
# lightly
# requests
# responses
# sentry-sdk
# tritonclient
uvicorn==0.35.0
# via gpt-oss
@@ -1378,22 +1121,16 @@ virtualenv==20.31.2
# via ray
vocos==0.1.0
# via -r requirements/test/cuda.in
wandb==0.24.2
# via terratorch
wcwidth==0.2.13
# via ftfy
webcolors==24.11.1
# via jsonschema
werkzeug==3.1.3
# via
# schemathesis
# tensorboard
# via schemathesis
word2number==1.1
# via lm-eval
wrapt==1.17.2
# via smart-open
xarray==2025.7.1
# via rioxarray
xxhash==3.5.0
# via
# datasets
+8 -2
View File
@@ -61,7 +61,11 @@ pydantic>=2.12 # 2.11 leads to error on python 3.13
decord==0.6.0
# Prithvi tests
terratorch>=1.2.2
# terratorch is temporarily disabled while PyPI has the `lightning` package
# in `quarantined` status (every published terratorch version transitively
# requires `lightning`, so the resolver fails with "no versions of lightning").
# Re-enable once PyPI lifts the quarantine. Tracked in #41376.
# terratorch>=1.2.2
imagehash # Required for Prithvi tests
segmentation-models-pytorch>0.4.0 # Required for Prithvi tests
@@ -79,5 +83,7 @@ plotly # required for perf comparison html report
# ROCm-specific extras (not in CUDA cuda.in)
rapidfuzz
torchgeo==0.7.0
# torchgeo also pulled in `lightning` transitively; disabled for the same
# quarantine reason as terratorch above. Restore once the quarantine clears.
# torchgeo==0.7.0
multiprocess==0.70.16
+13 -265
View File
@@ -1,15 +1,9 @@
# This file was autogenerated by uv via the following command:
# uv pip compile requirements/test/rocm.in -c requirements/rocm.txt -o requirements/test/rocm.txt --index-strategy unsafe-best-match --python-platform x86_64-manylinux_2_28 --python-version 3.12 --no-emit-package torch --no-emit-package torchvision --no-emit-package torchaudio --no-emit-package triton --no-emit-package cuda-bindings --no-emit-package cuda-pathfinder --no-emit-package cuda-toolkit --no-emit-package cupy-cuda12x --no-emit-package nvidia-cublas --no-emit-package nvidia-cuda-cupti --no-emit-package nvidia-cuda-nvrtc --no-emit-package nvidia-cuda-runtime --no-emit-package nvidia-cudnn --no-emit-package nvidia-cufft --no-emit-package nvidia-cufile --no-emit-package nvidia-curand --no-emit-package nvidia-cusolver --no-emit-package nvidia-cusparse --no-emit-package nvidia-cusparselt --no-emit-package nvidia-nccl --no-emit-package nvidia-nvjitlink --no-emit-package nvidia-nvshmem --no-emit-package nvidia-nvtx --no-emit-package nvidia-cublas-cu12 --no-emit-package nvidia-cuda-cupti-cu12 --no-emit-package nvidia-cuda-nvrtc-cu12 --no-emit-package nvidia-cuda-runtime-cu12 --no-emit-package nvidia-cudnn-cu12 --no-emit-package nvidia-cufft-cu12 --no-emit-package nvidia-cufile-cu12 --no-emit-package nvidia-curand-cu12 --no-emit-package nvidia-cusolver-cu12 --no-emit-package nvidia-cusparse-cu12 --no-emit-package nvidia-cusparselt-cu12 --no-emit-package nvidia-nccl-cu12 --no-emit-package nvidia-nvjitlink-cu12 --no-emit-package nvidia-nvshmem-cu12 --no-emit-package nvidia-nvtx-cu12 --no-emit-package nvidia-cublas-cu13 --no-emit-package nvidia-cuda-cupti-cu13 --no-emit-package nvidia-cuda-nvrtc-cu13 --no-emit-package nvidia-cuda-runtime-cu13 --no-emit-package nvidia-cudnn-cu13 --no-emit-package nvidia-cufft-cu13 --no-emit-package nvidia-cufile-cu13 --no-emit-package nvidia-curand-cu13 --no-emit-package nvidia-cusolver-cu13 --no-emit-package nvidia-cusparse-cu13 --no-emit-package nvidia-cusparselt-cu13 --no-emit-package nvidia-nccl-cu13 --no-emit-package nvidia-nvjitlink-cu13 --no-emit-package nvidia-nvshmem-cu13 --no-emit-package nvidia-nvtx-cu13
absl-py==2.4.0
# via
# rouge-score
# tensorboard
# via rouge-score
accelerate==1.13.0
# via peft
aenum==3.1.17
# via lightly
affine==2.4.0
# via rasterio
aiohappyeyeballs==2.6.1
# via aiohttp
aiohttp==3.13.3
@@ -25,12 +19,8 @@ aiohttp-cors==0.8.1
# via ray
aiosignal==1.4.0
# via aiohttp
albucore==0.1.2
# via terratorch
albumentations==1.4.6
# via
# -r requirements/test/rocm.in
# terratorch
# via -r requirements/test/rocm.in
alembic==1.18.4
# via optuna
annotated-doc==0.0.4
@@ -43,10 +33,6 @@ anthropic==0.93.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
antlr4-python3-runtime==4.9.3
# via
# hydra-core
# omegaconf
anyio==4.13.0
# via
# anthropic
@@ -67,11 +53,9 @@ astor==0.8.1
attrs==26.1.0
# via
# aiohttp
# fiona
# jsonlines
# jsonschema
# pytest-subtests
# rasterio
# referencing
audioread==3.0.1
# via librosa
@@ -90,9 +74,7 @@ backoff==2.2.1
# -r requirements/test/rocm.in
# schemathesis
bitsandbytes==0.49.2
# via
# -r requirements/test/rocm.in
# lightning
# via -r requirements/test/rocm.in
black==26.3.1
# via datamodel-code-generator
blake3==1.0.8
@@ -119,13 +101,8 @@ cbor2==5.9.0
# via -r requirements/test/../common.txt
certifi==2026.2.25
# via
# fiona
# httpcore
# httpx
# lightly
# pyogrio
# pyproj
# rasterio
# requests
# sentry-sdk
cffi==1.17.1
@@ -143,24 +120,13 @@ chz==0.4.0
click==8.3.1
# via
# black
# click-plugins
# cligj
# fiona
# jiwer
# nltk
# rasterio
# ray
# rich-toolkit
# schemathesis
# typer
# uvicorn
# wandb
click-plugins==1.1.1.2
# via fiona
cligj==0.7.2
# via
# fiona
# rasterio
cloudpickle==3.1.2
# via -r requirements/test/../common.txt
colorama==0.4.6
@@ -211,8 +177,6 @@ depyf==0.20.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
diffusers==0.37.0
# via terratorch
dill==0.3.8
# via
# datasets
@@ -237,16 +201,12 @@ docker==7.1.0
docopt==0.6.2
# via num2words
docstring-parser==0.17.0
# via
# anthropic
# jsonargparse
# via anthropic
einops==0.8.2
# via
# -r requirements/test/../common.txt
# -r requirements/test/rocm.in
# encodec
# terratorch
# torchgeo
# vector-quantize-pytorch
# vocos
einx==0.4.2
@@ -283,14 +243,11 @@ filelock==3.25.2
# -r requirements/test/../common.txt
# blobfile
# datasets
# diffusers
# huggingface-hub
# python-discovery
# ray
# torch
# virtualenv
fiona==1.10.1
# via torchgeo
fonttools==4.62.1
# via matplotlib
fqdn==1.5.1
@@ -307,9 +264,6 @@ fsspec==2025.3.0
# evaluate
# fastparquet
# huggingface-hub
# lightning
# pytorch-lightning
# tacoreader
# torch
ftfy==6.3.1
# via open-clip-torch
@@ -317,16 +271,10 @@ genai-perf==0.0.16
# via -r requirements/test/rocm.in
genson==1.3.0
# via datamodel-code-generator
geopandas==1.1.3
# via terratorch
gguf==0.18.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
gitdb==4.0.12
# via gitpython
gitpython==3.1.46
# via wandb
google-api-core==2.30.0
# via
# google-cloud-core
@@ -366,7 +314,6 @@ grpcio==1.78.0
# grpcio-reflection
# opentelemetry-exporter-otlp-proto-grpc
# ray
# tensorboard
grpcio-reflection==1.78.0
# via
# -c requirements/rocm.txt
@@ -377,8 +324,6 @@ h11==0.16.0
# uvicorn
h2==4.3.0
# via httpx
h5py==3.16.0
# via terratorch
harfile==0.4.0
# via schemathesis
hf-xet==1.4.3
@@ -397,7 +342,6 @@ httpx==0.27.2
# via
# -r requirements/test/rocm.in
# anthropic
# diffusers
# fastapi
# fastapi-cloud-cli
# huggingface-hub
@@ -412,23 +356,17 @@ huggingface-hub==1.10.2
# via
# accelerate
# datasets
# diffusers
# evaluate
# open-clip-torch
# peft
# segmentation-models-pytorch
# sentence-transformers
# terratorch
# timm
# tokenizers
# transformers
# vocos
humanize==4.15.0
# via runai-model-streamer
hydra-core==1.3.2
# via
# lightly
# lightning
hyperframe==6.1.0
# via h2
hypothesis==6.151.9
@@ -455,11 +393,7 @@ imagehash==4.3.2
imageio==2.37.3
# via scikit-image
importlib-metadata==8.7.1
# via
# diffusers
# opentelemetry-api
importlib-resources==6.5.2
# via typeshed-client
# via opentelemetry-api
inflect==7.5.0
# via datamodel-code-generator
iniconfig==2.3.0
@@ -497,14 +431,8 @@ joblib==1.5.3
# librosa
# nltk
# scikit-learn
jsonargparse==4.47.0
# via
# lightning
# terratorch
jsonlines==4.0.0
# via lm-eval
jsonnet==0.21.0
# via jsonargparse
jsonpointer==3.1.0
# via jsonschema
jsonschema==4.26.0
@@ -524,10 +452,6 @@ kaleido==1.0.0
# via genai-perf
kiwisolver==1.5.0
# via matplotlib
kornia==0.8.2
# via torchgeo
kornia-rs==0.1.10
# via kornia
lark==1.2.2
# via
# -c requirements/common.txt
@@ -540,21 +464,6 @@ libnacl==2.1.0
# via tensorizer
librosa==0.10.2.post1
# via -r requirements/test/rocm.in
lightly==1.5.22
# via
# terratorch
# torchgeo
lightly-utils==0.0.2
# via lightly
lightning==2.6.1
# via
# terratorch
# torchgeo
lightning-utilities==0.15.3
# via
# lightning
# pytorch-lightning
# torchmetrics
llguidance==1.3.0
# via
# -c requirements/common.txt
@@ -580,8 +489,6 @@ lxml==6.0.2
# sacrebleu
mako==1.3.10
# via alembic
markdown==3.10.2
# via tensorboard
markdown-it-py==4.0.0
# via rich
markupsafe==3.0.3
@@ -590,10 +497,7 @@ markupsafe==3.0.3
# mako
# werkzeug
matplotlib==3.10.8
# via
# -r requirements/test/rocm.in
# lightning
# torchgeo
# via -r requirements/test/rocm.in
mbstrdecoder==1.1.4
# via
# dataproperty
@@ -660,14 +564,11 @@ numba==0.65.0
# -c requirements/rocm.txt
# -r requirements/test/rocm.in
# librosa
numkong==7.1.1
# via albucore
numpy==2.2.6
# via
# -r requirements/test/../common.txt
# -r requirements/test/rocm.in
# accelerate
# albucore
# albumentations
# bitsandbytes
# bm25s
@@ -675,20 +576,15 @@ numpy==2.2.6
# cupy-cuda12x
# datasets
# decord
# diffusers
# einx
# encodec
# evaluate
# fastparquet
# genai-perf
# geopandas
# gguf
# h5py
# imagehash
# imageio
# librosa
# lightly
# lightly-utils
# lm-eval
# matplotlib
# mistral-common
@@ -700,12 +596,8 @@ numpy==2.2.6
# patsy
# peft
# perceptron
# pycocotools
# pyogrio
# pytrec-eval-terrier
# pywavelets
# rasterio
# rioxarray
# rouge-score
# runai-model-streamer
# sacrebleu
@@ -714,27 +606,16 @@ numpy==2.2.6
# scipy
# segmentation-models-pytorch
# sentence-transformers
# shapely
# soundfile
# soxr
# statsmodels
# tensorboard
# tensorboardx
# tensorizer
# terratorch
# tifffile
# torchgeo
# torchmetrics
# torchvision
# transformers
# tritonclient
# vocos
# xarray
# xgrammar
omegaconf==2.3.0
# via
# hydra-core
# lightning
open-clip-torch==2.32.0
# via -r requirements/test/rocm.in
openai==2.31.0
@@ -824,46 +705,29 @@ packaging==26.0
# datasets
# evaluate
# fastparquet
# geopandas
# huggingface-hub
# hydra-core
# kaleido
# kornia
# lazy-loader
# lightning
# lightning-utilities
# lm-format-enforcer
# matplotlib
# optuna
# peft
# plotly
# pooch
# pyogrio
# pytest
# pytest-rerunfailures
# pytorch-lightning
# ray
# rioxarray
# scikit-image
# statsmodels
# tensorboard
# tensorboardx
# torchmetrics
# transformers
# typepy
# wandb
# xarray
pandas==3.0.1
# via
# datasets
# evaluate
# fastparquet
# genai-perf
# geopandas
# statsmodels
# tacoreader
# torchgeo
# xarray
partial-json-parser==0.2.1.1.post7
# via -r requirements/test/../common.txt
pathspec==1.0.4
@@ -881,18 +745,14 @@ perf-analyzer==0.1.0
pillow==12.1.1
# via
# -r requirements/test/../common.txt
# diffusers
# genai-perf
# imagehash
# imageio
# lightly-utils
# matplotlib
# mistral-common
# perceptron
# scikit-image
# segmentation-models-pytorch
# tensorboard
# torchgeo
# torchvision
platformdirs==4.3.6
# via
@@ -900,7 +760,6 @@ platformdirs==4.3.6
# pooch
# python-discovery
# virtualenv
# wandb
plotly==6.6.0
# via
# -r requirements/test/rocm.in
@@ -946,10 +805,7 @@ protobuf==6.33.6
# opentelemetry-proto
# proto-plus
# ray
# tensorboard
# tensorboardx
# tensorizer
# wandb
psutil==7.2.2
# via
# -r requirements/test/../common.txt
@@ -966,16 +822,12 @@ pyarrow==23.0.1
# via
# datasets
# genai-perf
# tacoreader
# terratorch
pyasn1==0.6.3
# via pyasn1-modules
pyasn1-modules==0.4.2
# via google-auth
pybase64==1.4.3
# via -r requirements/test/../common.txt
pycocotools==2.0.11
# via terratorch
pycountry==26.2.16
# via pydantic-extra-types
pycparser==3.0
@@ -994,7 +846,6 @@ pydantic==2.12.5
# fastapi
# fastapi-cloud-cli
# gpt-oss
# lightly
# lm-format-enforcer
# mcp
# mistral-common
@@ -1005,7 +856,6 @@ pydantic==2.12.5
# pydantic-extra-types
# pydantic-settings
# ray
# wandb
# xgrammar
pydantic-core==2.41.5
# via pydantic
@@ -1023,17 +873,8 @@ pyjwt==2.12.1
# via
# mcp
# msal
pyogrio==0.12.1
# via geopandas
pyparsing==3.3.2
# via
# matplotlib
# rasterio
pyproj==3.7.2
# via
# geopandas
# rioxarray
# torchgeo
# via matplotlib
pyrate-limiter==3.9.0
# via schemathesis
pystemmer==3.0.0
@@ -1070,13 +911,10 @@ pytest-subtests==0.14.2
# via schemathesis
pytest-timeout==2.3.1
# via -r requirements/test/rocm.in
python-box==7.4.1
# via terratorch
python-dateutil==2.9.0.post0
# via
# arrow
# botocore
# lightly
# matplotlib
# pandas
# typepy
@@ -1096,10 +934,6 @@ python-rapidjson==1.23
# via tritonclient
pytokens==0.4.1
# via black
pytorch-lightning==2.6.1
# via
# lightly
# lightning
pytrec-eval-terrier==0.5.10
# via mteb
pytz==2026.1.post1
@@ -1116,13 +950,9 @@ pyyaml==6.0.3
# genai-perf
# gguf
# huggingface-hub
# jsonargparse
# lightning
# lm-format-enforcer
# omegaconf
# optuna
# peft
# pytorch-lightning
# ray
# responses
# schemathesis
@@ -1130,7 +960,6 @@ pyyaml==6.0.3
# transformers
# uvicorn
# vocos
# wandb
pyzmq==27.1.0
# via
# -c requirements/common.txt
@@ -1139,11 +968,6 @@ rapidfuzz==3.12.1
# via
# -r requirements/test/rocm.in
# jiwer
rasterio==1.5.0
# via
# rioxarray
# terratorch
# torchgeo
ray==2.54.0
# via -r requirements/test/rocm.in
redis==7.3.0
@@ -1155,7 +979,6 @@ referencing==0.37.0
regex==2026.2.28
# via
# -r requirements/test/../common.txt
# diffusers
# nltk
# open-clip-torch
# sacrebleu
@@ -1168,14 +991,12 @@ requests==2.32.5
# azure-core
# buildkite-test-collector
# datasets
# diffusers
# docker
# evaluate
# gguf
# google-api-core
# google-cloud-storage
# gpt-oss
# lightly
# lm-eval
# mistral-common
# msal
@@ -1186,9 +1007,7 @@ requests==2.32.5
# responses
# schemathesis
# starlette-testclient
# tacoreader
# tiktoken
# wandb
responses==0.26.0
# via genai-perf
rfc3339-validator==0.1.4
@@ -1198,11 +1017,9 @@ rfc3987==1.3.8
rich==14.3.3
# via
# genai-perf
# lightning
# mteb
# perceptron
# rich-toolkit
# terratorch
# typer
rich-toolkit==0.19.7
# via
@@ -1210,16 +1027,12 @@ rich-toolkit==0.19.7
# fastapi-cloud-cli
rignore==0.7.6
# via fastapi-cloud-cli
rioxarray==0.22.0
# via terratorch
rouge-score==0.1.2
# via lm-eval
rpds-py==0.30.0
# via
# jsonschema
# referencing
rtree==1.4.1
# via torchgeo
runai-model-streamer==0.15.7
# via
# -c requirements/rocm.txt
@@ -1237,7 +1050,6 @@ sacrebleu==2.6.0
safetensors==0.7.0
# via
# accelerate
# diffusers
# open-clip-torch
# peft
# segmentation-models-pytorch
@@ -1246,9 +1058,7 @@ safetensors==0.7.0
schemathesis==3.39.15
# via -r requirements/test/rocm.in
scikit-image==0.26.0
# via
# albumentations
# terratorch
# via albumentations
scikit-learn==1.8.0
# via
# albumentations
@@ -1256,7 +1066,6 @@ scikit-learn==1.8.0
# lm-eval
# mteb
# sentence-transformers
# terratorch
scipy==1.17.1
# via
# albumentations
@@ -1271,10 +1080,7 @@ scipy==1.17.1
# statsmodels
# vocos
segmentation-models-pytorch==0.5.0
# via
# -r requirements/test/rocm.in
# terratorch
# torchgeo
# via -r requirements/test/rocm.in
sentence-transformers==5.3.0
# via
# -r requirements/test/rocm.in
@@ -1282,9 +1088,7 @@ sentence-transformers==5.3.0
sentencepiece==0.2.1
# via -r requirements/test/../common.txt
sentry-sdk==2.55.0
# via
# fastapi-cloud-cli
# wandb
# via fastapi-cloud-cli
setproctitle==1.3.7
# via -r requirements/test/../common.txt
setuptools==79.0.1
@@ -1294,12 +1098,7 @@ setuptools==79.0.1
# -r requirements/test/../common.txt
# model-hosting-container-standards
# pytablewriter
# tensorboard
# torch
shapely==2.1.2
# via
# geopandas
# torchgeo
shellingham==1.5.4
# via
# perceptron
@@ -1311,15 +1110,12 @@ six==1.17.0
# -c requirements/common.txt
# -r requirements/test/../common.txt
# junit-xml
# lightly
# opencensus
# python-dateutil
# rfc3339-validator
# rouge-score
smart-open==7.5.1
# via ray
smmap==5.0.3
# via gitdb
sniffio==1.3.1
# via
# anthropic
@@ -1358,8 +1154,6 @@ starlette-testclient==0.4.1
# via schemathesis
statsmodels==0.14.6
# via genai-perf
stringzilla==4.6.0
# via albucore
structlog==25.5.0
# via gpt-oss
supervisor==4.3.0
@@ -1372,8 +1166,6 @@ tabledata==1.3.4
# via pytablewriter
tabulate==0.10.0
# via sacrebleu
tacoreader==0.5.6
# via terratorch
tblib==3.1.0
# via -r requirements/test/rocm.in
tcolorpy==0.1.7
@@ -1382,28 +1174,16 @@ tenacity==9.1.4
# via
# gpt-oss
# lm-eval
tensorboard==2.20.0
# via terratorch
tensorboard-data-server==0.7.2
# via tensorboard
tensorboardx==2.6.4
# via lightning
tensorizer==2.10.1
# via
# -c requirements/rocm.txt
# -r requirements/test/rocm.in
termcolor==3.3.0
# via
# gpt-oss
# terratorch
terratorch==1.2.2
# via -r requirements/test/rocm.in
# via gpt-oss
threadpoolctl==3.6.0
# via scikit-learn
tifffile==2026.3.3
# via
# scikit-image
# terratorch
# via scikit-image
tiktoken==0.12.0
# via
# -c requirements/common.txt
@@ -1417,8 +1197,6 @@ timm==1.0.17
# -r requirements/test/rocm.in
# open-clip-torch
# segmentation-models-pytorch
# terratorch
# torchgeo
tokenizers==0.22.2
# via
# -c requirements/common.txt
@@ -1429,16 +1207,6 @@ tomli==2.4.0
# via schemathesis
tomli-w==1.2.0
# via schemathesis
torchgeo==0.7.0
# via
# -r requirements/test/rocm.in
# terratorch
torchmetrics==1.9.0
# via
# lightning
# pytorch-lightning
# terratorch
# torchgeo
tqdm==4.67.3
# via
# -r requirements/test/../common.txt
@@ -1446,8 +1214,6 @@ tqdm==4.67.3
# evaluate
# gguf
# huggingface-hub
# lightly
# lightning
# lm-eval
# mteb
# nltk
@@ -1456,11 +1222,8 @@ tqdm==4.67.3
# optuna
# peft
# pqdm
# pytorch-lightning
# segmentation-models-pytorch
# sentence-transformers
# tacoreader
# terratorch
# transformers
transformers==5.5.3
# via
@@ -1492,8 +1255,6 @@ typer==0.24.1
# huggingface-hub
# perceptron
# transformers
typeshed-client==2.9.0
# via jsonargparse
typing-extensions==4.15.0
# via
# -c requirements/common.txt
@@ -1511,8 +1272,6 @@ typing-extensions==4.15.0
# grpcio
# huggingface-hub
# librosa
# lightning
# lightning-utilities
# lm-eval
# mcp
# mistral-common
@@ -1527,18 +1286,14 @@ typing-extensions==4.15.0
# pydantic
# pydantic-core
# pydantic-extra-types
# pytorch-lightning
# referencing
# rich-toolkit
# sentence-transformers
# sqlalchemy
# starlette
# torch
# torchgeo
# typeguard
# typeshed-client
# typing-inspection
# wandb
# xgrammar
typing-inspection==0.4.2
# via
@@ -1555,7 +1310,6 @@ urllib3==2.6.3
# blobfile
# botocore
# docker
# lightly
# requests
# responses
# sentry-sdk
@@ -1575,8 +1329,6 @@ virtualenv==21.2.0
# via ray
vocos==0.1.0
# via -r requirements/test/rocm.in
wandb==0.25.1
# via terratorch
watchfiles==1.1.1
# via
# -r requirements/test/../common.txt
@@ -1588,15 +1340,11 @@ webcolors==25.10.0
websockets==16.0
# via uvicorn
werkzeug==3.1.6
# via
# schemathesis
# tensorboard
# via schemathesis
word2number==1.1
# via lm-eval
wrapt==2.1.2
# via smart-open
xarray==2026.2.0
# via rioxarray
xgrammar==0.1.33
# via
# -c requirements/common.txt
+2 -4
View File
@@ -150,13 +150,11 @@ def test_full_graph(
if is_torch_equal_or_newer("2.9.0.dev")
]
+ [
# Test get_raw_stream patch with compile_sizes
# This tests that TorchInductor autotune works correctly with get_raw_stream
# patch in torch 2.9 and without patch in torch 2.10+
# Cover compile_sizes autotune path.
(
CompilationConfig(
mode=CompilationMode.VLLM_COMPILE,
compile_sizes=[1, 2], # Triggers autotune which uses get_raw_stream
compile_sizes=[1, 2], # Triggers the autotune path.
cudagraph_mode=CUDAGraphMode.NONE,
),
"facebook/opt-125m",
-24
View File
@@ -24,7 +24,6 @@ from vllm.engine.arg_utils import EngineArgs
from vllm.platforms import current_platform
from vllm.utils.torch_utils import (
_is_torch_equal_or_newer,
is_torch_equal,
)
from vllm.v1.cudagraph_dispatcher import CudagraphDispatcher
@@ -43,29 +42,6 @@ def test_version():
assert not _is_torch_equal_or_newer("2.7.1", "2.8.0.dev")
def test_get_raw_stream_patch():
"""Test that get_raw_stream patch is applied only for torch 2.9.0 or 2.9.1."""
import builtins
# Check if get_raw_stream exists in builtins
has_patch = hasattr(builtins, "get_raw_stream")
# Import torch to get actual version
is_torch_2_9 = is_torch_equal("2.9.0") or is_torch_equal("2.9.1")
if is_torch_2_9:
# For torch 2.9.x, the patch should be applied
assert has_patch, "get_raw_stream should be patched for torch 2.9.x"
# Verify it's callable (it should be the _cuda_getCurrentRawStream function)
get_raw_stream = builtins.get_raw_stream # type: ignore[attr-defined]
assert callable(get_raw_stream)
# Verify it's the correct function from torch._C
from torch._C import _cuda_getCurrentRawStream
assert get_raw_stream is _cuda_getCurrentRawStream
def test_copy_pass():
vllm_config = VllmConfig()
inductor_pass = FixFunctionalizationPass(vllm_config)
@@ -55,12 +55,10 @@ def test_dynamic_shapes_compilation(
evaluate_guards,
):
"""Test that all dynamic shapes types compile successfully"""
if use_bytecode_hook and shapes_type == DynamicShapesType.UNBACKED:
pytest.skip("UNBACKED dynamic shapes require VLLM_USE_BYTECODE_HOOK=0")
if evaluate_guards and shapes_type == DynamicShapesType.UNBACKED:
pytest.skip("unbacked dynamic shapes do not add guards")
# TODO is this still a requirement?
if evaluate_guards and use_aot_compile:
pytest.skip("evaluate_guards requires use_aot_compile=0")
+774
View File
@@ -0,0 +1,774 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Tests for MNNVL AllToAll operations.
Requires: docker run ... --cap-add=SYS_PTRACE ...
Run: pytest tests/distributed/test_mnnvl_alltoall.py -v
"""
import os
import traceback
import pytest
import torch
import torch.multiprocessing as mp
from vllm.distributed import get_ep_group
from vllm.utils.flashinfer import (
has_flashinfer_nvlink_one_sided,
has_flashinfer_nvlink_two_sided,
)
from vllm.utils.network_utils import get_open_port
from ..utils import init_test_distributed_environment
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _has_sys_ptrace() -> bool:
"""Check for SYS_PTRACE capability (bit 19 in CapEff)."""
try:
with open("/proc/self/status") as f:
for line in f:
if line.startswith("CapEff:"):
return bool(int(line.split()[1], 16) & (1 << 19))
except Exception:
pass
return False
def _spawn_workers(worker_fn, world_size, *, dp_size=None):
"""Spawn one process per GPU, run worker_fn, assert all succeed.
Uses an mp.Queue to propagate worker tracebacks back to the parent
so pytest shows the actual failure, not just an exit code.
"""
if mp.get_start_method(allow_none=True) is None:
mp.set_start_method("spawn")
port = str(get_open_port())
# Allocate a second port for DP master when dp_size is set, so the
# distributed init port and DP port can't collide even under xdist.
dp_port = str(get_open_port()) if dp_size is not None else None
err_queue: mp.Queue = mp.Queue()
procs = []
for rank in range(world_size):
p = mp.Process(
target=_run_worker,
args=(rank, world_size, port, worker_fn, dp_size, dp_port, err_queue),
)
p.start()
procs.append(p)
for p in procs:
p.join()
# Collect any errors from workers before asserting.
errors = []
while not err_queue.empty():
errors.append(err_queue.get_nowait())
err_queue.close()
err_queue.join_thread()
if errors:
pytest.fail("Worker(s) failed:\n" + "\n---\n".join(errors))
def _run_worker(rank, world_size, port, worker_fn, dp_size, dp_port, err_queue):
"""Per-process setup: device, distributed env, then call worker_fn.
Args:
dp_size: If set, initialize with tp=1 and data_parallel_size=dp_size.
Otherwise use tp=world_size (default for EP-based tests).
dp_port: Separate port for the DP master (only used when dp_size is set).
err_queue: Queue for propagating tracebacks to the parent process.
"""
try:
os.environ.pop("CUDA_VISIBLE_DEVICES", None)
torch.accelerator.set_device_index(rank)
if dp_size is not None:
_init_dp_environment(world_size, rank, port, dp_size, dp_port)
else:
init_test_distributed_environment(world_size, 1, rank, port)
worker_fn(rank, world_size)
torch.distributed.barrier()
except Exception:
err_queue.put(f"[Rank {rank}]\n{traceback.format_exc()}")
# Don't re-raise: the parent reads errors from err_queue.
# A non-zero exit from the re-raise would be redundant.
import sys
sys.exit(1)
def _init_dp_environment(world_size, rank, port, dp_size, dp_port):
"""Initialize distributed env with data parallelism.
Sets up tp=1, pp=1, dp=dp_size. Each process is one DP rank
with local rank 0 within its (trivial) tp*pp group.
Args:
port: Port for torch.distributed init.
dp_port: Separate port for the DP master group init.
"""
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.config.parallel import ParallelConfig
from vllm.distributed.parallel_state import (
ensure_model_parallel_initialized,
init_distributed_environment,
)
vllm_config = VllmConfig()
vllm_config.parallel_config = ParallelConfig(
data_parallel_size=dp_size,
data_parallel_rank=rank,
# Pre-populate port list so __post_init__ doesn't auto-generate
# random ports. All DP ranks must agree on the same port.
_data_parallel_master_port_list=[int(dp_port)],
)
with set_current_vllm_config(vllm_config):
# rank=0 here because each DP rank has a single (tp=1,pp=1) process,
# so the local rank within the tp*pp group is always 0.
# init_distributed_environment will offset by data_parallel_rank.
init_distributed_environment(
world_size=1, # tp * pp = 1
rank=0,
distributed_init_method=f"tcp://localhost:{port}",
local_rank=rank,
)
ensure_model_parallel_initialized(1, 1)
def _make_forward_context(rank, world_size, num_tokens_per_rank):
"""Create a forward context with mock DP metadata for AgRs tests.
Returns a context manager suitable for ``with`` statements.
The real DPMetadata (with sp_local_sizes etc.) is created internally
by set_forward_context from num_tokens_across_dp; the attn_metadata
placeholder just satisfies the "attn_metadata is not None" guard.
"""
from vllm.config.parallel import ParallelConfig
from vllm.config.vllm import VllmConfig
from vllm.forward_context import set_forward_context
class _AttnMeta:
"""Minimal placeholder so set_forward_context's
``attn_metadata is not None`` guard (forward_context.py:334)
is satisfied. The real DPMetadata is built from num_tokens_across_dp."""
dp_metadata = None
vllm_config = VllmConfig()
vllm_config.parallel_config = ParallelConfig(
data_parallel_size=world_size,
is_moe_model=True,
data_parallel_rank=rank,
)
return set_forward_context(
_AttnMeta(),
vllm_config,
num_tokens=num_tokens_per_rank,
num_tokens_across_dp=torch.tensor(
[num_tokens_per_rank] * world_size, dtype=torch.int
),
)
# ---------------------------------------------------------------------------
# Skip conditions
# ---------------------------------------------------------------------------
requires_multi_gpu = pytest.mark.skipif(
torch.accelerator.device_count() < 2, reason="Need >= 2 GPUs"
)
requires_two_sided = pytest.mark.skipif(
not has_flashinfer_nvlink_two_sided(),
reason="FlashInfer NVLink two-sided not available",
)
requires_one_sided = pytest.mark.skipif(
not has_flashinfer_nvlink_one_sided(),
reason="FlashInfer NVLink one-sided not available",
)
requires_ptrace = pytest.mark.skipif(
not _has_sys_ptrace(),
reason="SYS_PTRACE required (docker run --cap-add=SYS_PTRACE)",
)
# NOTE: No module-level pytestmark here. The FlashInfer lifecycle tests have
# their own @requires_two_sided / @requires_one_sided decorators, and
# test_args_dispatch_combine uses only standard torch.distributed ops and
# should run even when FlashInfer NVLink backends are not installed.
# ---------------------------------------------------------------------------
# Test 1: Two-sided manager lifecycle (init, cleanup, reinit, ensure_init)
# ---------------------------------------------------------------------------
#
# Tests FlashInferNVLinkTwoSidedManager which wraps FlashInfer's MnnvlMoe.
# initialize() allocates MNNVL shared workspaces via MnnvlMoe.get_moe_workspaces,
# which uses pidfd_getfd() to share memory file descriptors across processes —
# hence the SYS_PTRACE requirement.
#
# Uses EP group (get_ep_group) because the two-sided manager is constructed
# with an EP-scoped communicator in production. With tp=world_size the EP
# group spans all ranks, giving us a multi-rank group for testing.
# ---------------------------------------------------------------------------
def _two_sided_lifecycle_worker(rank, world_size):
from vllm.distributed.device_communicators.all2all import (
FlashInferNVLinkTwoSidedManager,
)
cpu_group = get_ep_group().cpu_group
num_gpus = torch.accelerator.device_count()
manager = FlashInferNVLinkTwoSidedManager(cpu_group)
# Not initialized yet
assert not manager.initialized
assert manager.rank == rank
assert manager.world_size == world_size
# Initialize
manager.initialize(world_size=world_size, rank=rank, gpus_per_node=num_gpus)
assert manager.initialized
assert manager.workspace_tensor is not None
assert manager.prepare_workspace_tensor is not None
assert manager.mapping is not None
torch.distributed.barrier()
# Cleanup
manager.cleanup()
assert not manager.initialized
assert manager.workspace_tensor is None
assert manager.prepare_workspace_tensor is None
torch.distributed.barrier()
# Reinitialize
manager.initialize(world_size=world_size, rank=rank, gpus_per_node=num_gpus)
assert manager.initialized
torch.distributed.barrier()
# ensure_alltoall_workspace_initialized is idempotent when already init'd
assert manager.ensure_alltoall_workspace_initialized()
assert manager.initialized
manager.cleanup()
assert not manager.initialized
@requires_multi_gpu
@requires_two_sided
@requires_ptrace
@pytest.mark.parametrize("world_size", [2])
def test_two_sided_manager_lifecycle(world_size):
"""Test init, cleanup, reinit, and ensure_initialized idempotency."""
_spawn_workers(_two_sided_lifecycle_worker, world_size)
# ---------------------------------------------------------------------------
# Test 2: One-sided manager lifecycle (init, cleanup, reinit)
# ---------------------------------------------------------------------------
#
# Tests FlashInferNVLinkOneSidedManager which wraps FlashInfer's MoeAlltoAll.
# initialize() creates MoeAlltoAll with an MnnvlConfig, which allocates MNNVL
# shared workspaces — same cross-process memory sharing as two-sided, hence
# the SYS_PTRACE requirement.
#
# Uses DP group (get_dp_group) because the one-sided manager's initialize()
# internally calls get_dp_group() to set up the MnnvlConfig communicator.
# We therefore need a real DP group with world_size > 1, which requires
# dp_size=world_size via _init_dp_environment.
# ---------------------------------------------------------------------------
def _one_sided_lifecycle_worker(rank, world_size):
from vllm.distributed.device_communicators.all2all import (
FlashInferNVLinkOneSidedManager,
)
from vllm.distributed.parallel_state import get_dp_group
cpu_group = get_dp_group().cpu_group
manager = FlashInferNVLinkOneSidedManager(cpu_group)
assert not manager.initialized
assert manager.rank == rank
assert manager.world_size == world_size
init_kwargs = dict(
max_num_tokens=1024,
top_k=2,
num_experts=world_size * 8,
hidden_size=4096,
)
# Initialize
manager.initialize(**init_kwargs)
assert manager.initialized
assert manager.moe_alltoall is not None
assert manager.mapping is not None
torch.distributed.barrier()
# Cleanup
manager.cleanup()
assert not manager.initialized
assert manager.moe_alltoall is None
torch.distributed.barrier()
# Reinitialize with different token count
manager.initialize(**{**init_kwargs, "max_num_tokens": 2048})
assert manager.initialized
torch.distributed.barrier()
manager.cleanup()
@requires_multi_gpu
@requires_one_sided
@requires_ptrace
@pytest.mark.parametrize("world_size", [2])
def test_one_sided_manager_lifecycle(world_size):
"""Test init, cleanup, and reinit with different params."""
_spawn_workers(
_one_sided_lifecycle_worker,
world_size,
dp_size=world_size,
)
# ---------------------------------------------------------------------------
# Test 3: AgRs dispatch/combine with value validation
# ---------------------------------------------------------------------------
#
# Tests AgRsAll2AllManager which uses only standard torch.distributed
# all_gatherv / reduce_scatterv — no FlashInfer or MNNVL dependency.
# This test validates the reference all-to-all implementation that other
# backends are compared against.
# ---------------------------------------------------------------------------
def _args_dispatch_combine_worker(rank, world_size):
from vllm.distributed.device_communicators.all2all import AgRsAll2AllManager
from vllm.forward_context import get_forward_context
cpu_group = get_ep_group().cpu_group
device = torch.device(f"cuda:{rank}")
hidden_size = 64
tokens_per_rank = 16
experts_per_token = 2
num_experts = world_size * 4
total_tokens = world_size * tokens_per_rank
# Deterministic per-rank data: rank r has value (r + 1)
hidden = torch.full(
(tokens_per_rank, hidden_size),
float(rank + 1),
device=device,
dtype=torch.float32,
)
router = torch.full(
(tokens_per_rank, num_experts),
float(rank + 1) * 10,
device=device,
dtype=torch.float32,
)
weights = torch.full(
(tokens_per_rank, experts_per_token),
float(rank + 1) * 100,
device=device,
dtype=torch.float32,
)
ids = torch.full(
(tokens_per_rank, experts_per_token),
rank,
device=device,
dtype=torch.long,
)
with _make_forward_context(rank, world_size, tokens_per_rank):
manager = AgRsAll2AllManager(cpu_group)
dp_metadata = get_forward_context().dp_metadata
with dp_metadata.sp_local_sizes(sequence_parallel_size=1):
# -- dispatch_router_logits --
d_hidden, d_router = manager.dispatch_router_logits(
hidden.clone(),
router.clone(),
is_sequence_parallel=True,
)
assert d_hidden.shape == (total_tokens, hidden_size)
assert d_router.shape == (total_tokens, num_experts)
for r in range(world_size):
s = r * tokens_per_rank
e = (r + 1) * tokens_per_rank
torch.testing.assert_close(
d_hidden[s:e],
torch.full_like(d_hidden[s:e], float(r + 1)),
)
torch.testing.assert_close(
d_router[s:e],
torch.full_like(d_router[s:e], float(r + 1) * 10),
)
# -- dispatch --
d_hidden2, d_weights, d_ids = manager.dispatch(
hidden.clone(),
weights.clone(),
ids.clone(),
is_sequence_parallel=True,
)
assert d_hidden2.shape == (total_tokens, hidden_size)
assert d_weights.shape == (total_tokens, experts_per_token)
assert d_ids.shape == (total_tokens, experts_per_token)
for r in range(world_size):
s = r * tokens_per_rank
e = (r + 1) * tokens_per_rank
torch.testing.assert_close(
d_weights[s:e],
torch.full_like(d_weights[s:e], float(r + 1) * 100),
)
assert (d_ids[s:e] == r).all()
# -- combine (reduce-scatter) --
# Each token i has value i in all columns; after reduce-scatter
# each rank gets its slice, summed across ranks.
expert_out = (
torch.arange(total_tokens, device=device, dtype=torch.float32)
.unsqueeze(1)
.expand(total_tokens, hidden_size)
.contiguous()
)
combined = manager.combine(expert_out, is_sequence_parallel=True)
assert combined.shape == (tokens_per_rank, hidden_size)
for i in range(tokens_per_rank):
expected_val = float(rank * tokens_per_rank + i) * world_size
torch.testing.assert_close(
combined[i],
torch.full_like(combined[i], expected_val),
)
torch.distributed.barrier()
@requires_multi_gpu
@pytest.mark.parametrize("world_size", [2])
def test_args_dispatch_combine(world_size):
"""Validate dispatch gathers all-rank data and combine reduces correctly."""
_spawn_workers(_args_dispatch_combine_worker, world_size)
# ---------------------------------------------------------------------------
# Test 4: FlashInfer two-sided dispatch/combine data communication
# ---------------------------------------------------------------------------
#
# Tests actual data flow through the FlashInfer NVLink two-sided backend
# by calling flashinfer_alltoall_dispatch (with defer_input_quant=True to
# skip quantization) and flashinfer_alltoall_combine, then verifying exact
# round-trip values. Dispatch sends each token once per distinct expert
# rank, and combine performs an unweighted sum, so:
# dispatch(hidden) → identity → combine = hidden * num_distinct_ranks(i)
# ---------------------------------------------------------------------------
def _two_sided_data_worker(rank, world_size):
from vllm.distributed.device_communicators.all2all import (
FlashInferNVLinkTwoSidedManager,
)
from vllm.distributed.parallel_state import get_dp_group
from vllm.forward_context import get_forward_context
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEQuantConfig,
FusedMoEQuantDesc,
)
from vllm.model_executor.layers.fused_moe.prepare_finalize.flashinfer_nvlink_two_sided import ( # noqa: E501
flashinfer_alltoall_combine,
flashinfer_alltoall_dispatch,
)
# Use DP group because MnnvlMoe workspace allocation calls get_dp_group()
# internally and requires dp_size == ep_size.
cpu_group = get_dp_group().cpu_group
device = torch.device(f"cuda:{rank}")
num_gpus = torch.accelerator.device_count()
hidden_size = 128
tokens_per_rank = 32
experts_per_token = 2
num_experts = world_size * 4
# Initialize the FlashInfer two-sided manager
manager = FlashInferNVLinkTwoSidedManager(cpu_group)
manager.initialize(world_size=world_size, rank=rank, gpus_per_node=num_gpus)
assert manager.initialized
torch.distributed.barrier()
# Create deterministic per-rank test data
torch.manual_seed(rank + 42)
hidden = torch.randn(
tokens_per_rank,
hidden_size,
device=device,
dtype=torch.bfloat16,
)
# Assign each token to experts spread across ranks so tokens move between GPUs
topk_ids = torch.randint(
0,
num_experts,
(tokens_per_rank, experts_per_token),
device=device,
dtype=torch.int32,
)
topk_weights = torch.rand(
tokens_per_rank,
experts_per_token,
device=device,
dtype=torch.float32,
)
# Unquantized config: quant_dtype=None means moe_kernel_quantize_input is a no-op
no_quant = FusedMoEQuantDesc()
quant_config = FusedMoEQuantConfig(
_a1=no_quant,
_a2=no_quant,
_w1=no_quant,
_w2=no_quant,
)
assert quant_config.quant_dtype is None # sanity: no quantization
with _make_forward_context(rank, world_size, tokens_per_rank):
dp_metadata = get_forward_context().dp_metadata
with dp_metadata.sp_local_sizes(sequence_parallel_size=1):
local_sizes = dp_metadata.get_chunk_sizes_across_dp_rank()
# --- FlashInfer two-sided dispatch ---
alltoall_info, fi_topk_ids, fi_topk_weights, fi_hidden, fi_scale = (
flashinfer_alltoall_dispatch(
manager,
local_sizes,
hidden.clone(),
None, # no global scale
topk_ids.clone(),
topk_weights.clone(),
experts_per_token,
num_experts,
quant_config,
defer_input_quant=True,
)
)
assert fi_scale is None # deferred quant: no scale produced
assert fi_hidden is not None
assert fi_hidden.shape[1] == hidden_size
assert fi_hidden.numel() > 0
# --- Round-trip exact verification ---
# The all-to-all sends each token once per *distinct* expert
# rank. Combine performs an unweighted sum of the per-rank
# contributions. With identity expert (feeding dispatched
# hidden straight back):
# result[i] = hidden[i] * num_distinct_expert_ranks(i)
combined = flashinfer_alltoall_combine(
manager,
fi_hidden,
top_k=experts_per_token,
token_count=tokens_per_rank,
alltoall_info=alltoall_info,
)
assert combined.shape == (tokens_per_rank, hidden_size)
experts_per_rank = num_experts // world_size
expert_ranks = topk_ids // experts_per_rank # (tokens, top_k)
num_distinct = torch.tensor(
[len(set(row.tolist())) for row in expert_ranks],
device=device,
dtype=torch.float32,
).unsqueeze(1) # (tokens, 1)
expected = (hidden.float() * num_distinct).to(hidden.dtype)
torch.testing.assert_close(combined, expected)
# --- Linearity check with scaled expert output ---
# Scaling the expert output by a constant should scale the
# combined result by the same constant.
scale = 3.0
combined_scaled = flashinfer_alltoall_combine(
manager,
fi_hidden * scale,
top_k=experts_per_token,
token_count=tokens_per_rank,
alltoall_info=alltoall_info,
)
expected_scaled = (hidden.float() * num_distinct * scale).to(hidden.dtype)
torch.testing.assert_close(combined_scaled, expected_scaled)
torch.distributed.barrier()
manager.cleanup()
@requires_multi_gpu
@requires_two_sided
@requires_ptrace
@pytest.mark.parametrize("world_size", [2])
def test_two_sided_dispatch_combine(world_size):
"""Test FlashInfer two-sided dispatch/combine with exact value verification."""
_spawn_workers(_two_sided_data_worker, world_size, dp_size=world_size)
# ---------------------------------------------------------------------------
# Test 5: FlashInfer one-sided dispatch/combine data communication
# ---------------------------------------------------------------------------
#
# Tests actual data flow through the FlashInfer NVLink one-sided backend
# by calling MoeAlltoAll.dispatch() and MoeAlltoAll.combine() directly
# with synthetic payloads, then verifying shapes and round-trip consistency.
# ---------------------------------------------------------------------------
def _one_sided_data_worker(rank, world_size):
from vllm.distributed.device_communicators.all2all import (
FlashInferNVLinkOneSidedManager,
)
from vllm.distributed.parallel_state import get_dp_group
from vllm.forward_context import get_forward_context
cpu_group = get_dp_group().cpu_group
device = torch.device(f"cuda:{rank}")
hidden_size = 256
tokens_per_rank = 32
experts_per_token = 2
num_experts = world_size * 8
# Initialize the one-sided manager
manager = FlashInferNVLinkOneSidedManager(cpu_group)
manager.initialize(
max_num_tokens=tokens_per_rank,
top_k=experts_per_token,
num_experts=num_experts,
hidden_size=hidden_size,
)
assert manager.initialized
assert manager.moe_alltoall is not None
with _make_forward_context(rank, world_size, tokens_per_rank):
dp_metadata = get_forward_context().dp_metadata
with dp_metadata.sp_local_sizes(sequence_parallel_size=1):
local_sizes = dp_metadata.get_chunk_sizes_across_dp_rank()
runtime_max_tokens = max(local_sizes)
# Create test data with raw tensors matching the nvfp4 payload
# sizes the workspace was allocated for:
# a1q: (tokens, hidden_size // 2) — nvfp4 hidden states
# a1q_scale: (tokens, hidden_size // 16) — fp8 scaling factors
torch.manual_seed(rank + 42)
a1q = torch.randint(
0,
256,
(tokens_per_rank, hidden_size // 2),
device=device,
dtype=torch.uint8,
)
a1q_scale = torch.randint(
0,
256,
(tokens_per_rank, hidden_size // 16),
device=device,
dtype=torch.uint8,
)
topk_ids = torch.randint(
0,
num_experts,
(tokens_per_rank, experts_per_token),
device=device,
dtype=torch.int32,
)
topk_weights = torch.rand(
tokens_per_rank,
experts_per_token,
device=device,
dtype=torch.float32,
)
# --- One-sided dispatch ---
payloads = [a1q, a1q_scale, topk_ids, topk_weights]
recv_payloads = manager.moe_alltoall.dispatch(
token_selected_experts=topk_ids,
input_payloads=payloads,
runtime_max_tokens_per_rank=runtime_max_tokens,
)
assert len(recv_payloads) == 4
recv_a1q, recv_scale, recv_ids, recv_weights = recv_payloads
assert recv_a1q.numel() > 0
assert recv_ids.numel() > 0
# --- Round-trip exact verification ---
# The dispatch routes each token once per *distinct* expert
# rank. Combine performs an unweighted sum of per-rank
# contributions. With constant expert output (all 1s):
# result[i] = 1.0 * num_distinct_expert_ranks(i)
expert_output = torch.ones(
world_size,
runtime_max_tokens,
hidden_size,
device=device,
dtype=torch.bfloat16,
)
combined = manager.moe_alltoall.combine(
payload=expert_output,
runtime_max_tokens_per_rank=runtime_max_tokens,
)
assert combined.shape == (tokens_per_rank, hidden_size)
experts_per_rank = num_experts // world_size
expert_ranks = topk_ids // experts_per_rank # (tokens, top_k)
num_distinct = torch.tensor(
[len(set(row.tolist())) for row in expert_ranks],
device=device,
dtype=torch.bfloat16,
).unsqueeze(1) # (tokens, 1)
expected = num_distinct.expand_as(combined)
torch.testing.assert_close(combined, expected)
# --- Linearity check with scaled expert output ---
# Scaling the expert output by a constant should scale the
# combined result by the same constant.
# Re-dispatch to reset internal state (one-sided requires a
# fresh dispatch before each combine).
manager.moe_alltoall.dispatch(
token_selected_experts=topk_ids,
input_payloads=payloads,
runtime_max_tokens_per_rank=runtime_max_tokens,
)
scale = 3.0
combined_scaled = manager.moe_alltoall.combine(
payload=expert_output * scale,
runtime_max_tokens_per_rank=runtime_max_tokens,
)
expected_scaled = (expected * scale).to(torch.bfloat16)
torch.testing.assert_close(combined_scaled, expected_scaled)
torch.distributed.barrier()
manager.cleanup()
@requires_multi_gpu
@requires_one_sided
@requires_ptrace
@pytest.mark.parametrize("world_size", [2])
def test_one_sided_dispatch_combine(world_size):
"""Test FlashInfer one-sided dispatch/combine with actual data flow."""
_spawn_workers(_one_sided_data_worker, world_size, dp_size=world_size)
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import importlib.util
import numpy as np
import pybase64 as base64
import pytest
@@ -10,7 +12,16 @@ import torch
from tests.utils import RemoteOpenAIServer
from vllm.utils.serial_utils import tensor2base64
# Prithvi requires terratorch, which is temporarily unavailable while PyPI has
# `lightning` quarantined (#41376). Skip just the Prithvi case; leave the
# Qwen3-VL case in the same file untouched.
_TERRATORCH_AVAILABLE = importlib.util.find_spec("terratorch") is not None
@pytest.mark.skipif(
not _TERRATORCH_AVAILABLE,
reason="terratorch unavailable while PyPI has `lightning` quarantined; see #41376",
)
@pytest.mark.parametrize(
"model_name", ["ibm-nasa-geospatial/Prithvi-EO-2.0-300M-TL-Sen1Floods11"]
)
@@ -167,9 +167,8 @@ def run_evaluation(
"model_config",
[
("openai/whisper-large-v3", 12.744980),
# TODO (ekagra): turn on after asr release
# CohereASR is used to test the variable encoder length code paths
# ("CohereLabs/cohere-transcribe-03-2026", 11.92),
("CohereLabs/cohere-transcribe-03-2026", 11.92),
],
)
# Original dataset is 20GB+ in size, hence we use a pre-filtered slice.
@@ -6,7 +6,9 @@ from unittest.mock import MagicMock
import pytest
import vllm.envs as envs
from vllm.entrypoints.openai.engine.serving import GenerationError, OpenAIServing
from vllm.envs import disable_envs_cache
@pytest.mark.asyncio
@@ -60,3 +62,35 @@ async def test_convert_generation_error_to_streaming_response():
assert isinstance(error_json, str)
assert "Internal server error" in error_json
assert "InternalServerError" in error_json
def test_is_model_supported_skip_name_validation_env(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""When VLLM_SKIP_MODEL_NAME_VALIDATION is set, accept any model id."""
disable_envs_cache()
monkeypatch.delenv("VLLM_SKIP_MODEL_NAME_VALIDATION", raising=False)
mock_engine = MagicMock()
mock_engine.model_config = MagicMock()
mock_engine.model_config.max_model_len = 100
mock_models = MagicMock()
mock_models.is_base_model.return_value = False
serving = OpenAIServing(
engine_client=mock_engine,
models=mock_models,
request_logger=None,
)
assert serving._is_model_supported("not-a-registered-model") is False
monkeypatch.setenv("VLLM_SKIP_MODEL_NAME_VALIDATION", "1")
disable_envs_cache()
assert envs.VLLM_SKIP_MODEL_NAME_VALIDATION is True
assert serving._is_model_supported("not-a-registered-model") is True
monkeypatch.setenv("VLLM_SKIP_MODEL_NAME_VALIDATION", "true")
disable_envs_cache()
assert envs.VLLM_SKIP_MODEL_NAME_VALIDATION is True
assert serving._is_model_supported("another-alias") is True
+228 -4
View File
@@ -3,12 +3,11 @@
"""
Round-trip tests for compressor → FP8 quant + KV cache insert → gather + dequant.
Two paths tested:
Four test functions cover five paths:
A) DeepseekV4 Attention: head_dim=512 (448 FP8 nope + 64 bf16 rope), quant_block=64
B) Indexer: head_dim=128 (all FP8), quant_block=128
These serve as golden references for validating the future fused
compressor+quant+cache kernel.
C) DeepseekV4 Attention magnitude range: correctness across small/large values
D) Indexer fused Triton kernel: compress+norm+rope+quant+insert
"""
import math
@@ -21,6 +20,12 @@ from vllm.v1.attention.ops.deepseek_v4_ops import (
dequantize_and_gather_k_cache,
quantize_and_insert_k_cache,
)
from vllm.v1.attention.ops.deepseek_v4_ops.fused_compress_quant_cache import (
_fused_kv_compress_norm_rope_insert_indexer_attn,
_fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn,
)
from .test_fused_indexer_q_rope_quant import quantize_to_mxfp4
def _ue8m0_reference(x: torch.Tensor, block_size: int, fp8_max: float):
@@ -309,3 +314,222 @@ def test_deepseek_v4_quant_magnitude_range():
f"Token {t}: rel_err={rel_err:.4f}, abs_diff={abs_diff:.6f}, "
f"magnitude={magnitude:.4f}"
)
# ── Test D: Indexer fused K-cache insert (Triton kernels) ────────────────────
#
# Both kernels share the same Triton signature; use_fp4 selects between them.
# Full pipeline: state-cache gather → softmax-weighted compress → RMSNorm →
# GPT-J RoPE → quant (MXFP4 or FP8) → paged cache insert.
def _reference_kv_compress_norm_rope(
state_cache: torch.Tensor,
block_table: torch.Tensor,
positions: torch.Tensor,
rms_weight: torch.Tensor,
cos_sin_cache: torch.Tensor,
compress_ratio: int = 1,
overlap: int = 0,
use_fp4: bool = False,
rms_eps: float = 1e-6,
fp8_max: float = 448.0,
):
"""Compress → RMSNorm → GPT-J RoPE → quantize.
Gathers (1+overlap)*compress_ratio state entries per output token, applies
per-element softmax over the scores, and computes the weighted kv sum.
Returns (quantized_values, scale) matching the kernel's output layout.
"""
device = state_cache.device
head_dim = rms_weight.shape[0]
rope_dim = cos_sin_cache.shape[-1]
state_block_size = state_cache.shape[1]
state_width = state_cache.shape[-1] // 2
nope_dim = head_dim - rope_dim
total = (1 + overlap) * compress_ratio
results = []
for pos in positions.tolist():
src = torch.arange(pos - total + 1, pos + 1, dtype=torch.int64, device=device)
valid = src >= 0
idx = src.clamp(min=0)
pages = block_table[0, idx // state_block_size]
offsets = idx % state_block_size
raw = state_cache[pages, offsets].float() # [total, state_dim]
# Group 0 (tokens 0..cr-1): kv[:H], score[SW:SW+H]
# Group 1 (tokens cr..2cr-1): kv[H:2H], score[SW+H:SW+2H]
if overlap:
sw = state_width
g0_kv = raw[:compress_ratio, :head_dim]
g1_kv = raw[compress_ratio:, head_dim : 2 * head_dim]
g0_scores = raw[:compress_ratio, sw : sw + head_dim]
g1_scores = raw[compress_ratio:, sw + head_dim : sw + 2 * head_dim]
kv = torch.cat([g0_kv, g1_kv])
scores = torch.cat([g0_scores, g1_scores])
else:
kv = raw[:, :head_dim]
scores = raw[:, state_width : state_width + head_dim]
scores[~valid] = float("-inf")
kv[~valid] = 0.0
weights = torch.softmax(scores, dim=0)
compressed = (kv * weights).sum(dim=0) # [H]
var = (compressed * compressed).mean()
normed = compressed * torch.rsqrt(var + rms_eps) * rms_weight.float()
compressed_pos = (pos // compress_ratio) * compress_ratio
cos, sin = cos_sin_cache[compressed_pos].float().chunk(2)
nope, rope = normed.split([nope_dim, rope_dim])
rope = torch.stack(
[rope[0::2] * cos - rope[1::2] * sin, rope[1::2] * cos + rope[0::2] * sin],
dim=-1,
).reshape(rope_dim)
results.append(torch.cat([nope, rope]).to(state_cache.dtype))
result = torch.stack(results)
if use_fp4:
return quantize_to_mxfp4(result)
else:
pairs = [
_ue8m0_reference(result[t], head_dim, fp8_max) for t in range(len(result))
]
quants, scales = zip(*pairs)
return torch.stack(quants), torch.cat(scales)
@pytest.mark.parametrize("num_tokens", [1, 7, 32])
@pytest.mark.parametrize("kv_block_size", [16, 32])
@pytest.mark.parametrize("use_fp4", [False, True])
def test_fused_kv_insert_indexer(num_tokens: int, kv_block_size: int, use_fp4: bool):
"""Fused K compress+norm+rope+quant+insert for the indexer KV cache."""
HEAD_DIM = 128
ROPE_DIM = 64
BLOCK_SIZE = 16
RMS_EPS = 1e-6
FP8_MAX = 448.0
device = "cuda"
torch.manual_seed(42)
compress_ratio = 4
if use_fp4:
TOKEN_STRIDE = HEAD_DIM // 2 # packed nibbles: 64 bytes
SCALE_DIM = HEAD_DIM // 32 # ue8m0 bytes: 4
QUANT_BLOCK = 32
kernel = _fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn
else:
TOKEN_STRIDE = HEAD_DIM # FP8 bytes: 128
SCALE_DIM = 4 # 1 float32: 4 bytes
QUANT_BLOCK = HEAD_DIM
kernel = _fused_kv_compress_norm_rope_insert_indexer_attn
# overlap=1 whenever compress_ratio==4, matching DeepseekCompressor logic.
overlap = 1 if compress_ratio == 4 else 0
coff = 1 + overlap # multiplier for state_dim per entry
num_pages = (compress_ratio * num_tokens - 1) // BLOCK_SIZE + 2
state_cache = torch.randn(
num_pages,
BLOCK_SIZE,
2 * coff * HEAD_DIM, # kv_state + score_state, each coff*HEAD_DIM wide
dtype=torch.bfloat16,
device=device,
)
block_table = torch.arange(num_pages, dtype=torch.int32, device=device).unsqueeze(0)
token_to_req = torch.zeros(num_tokens, dtype=torch.int32, device=device)
slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)
positions = torch.arange(
compress_ratio - 1,
compress_ratio * num_tokens,
compress_ratio,
dtype=torch.int64,
device=device,
)
rms_weight = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device=device)
cos_sin_cache = torch.randn(compress_ratio * num_tokens, ROPE_DIM, device=device)
kv_n_blocks = (num_tokens + kv_block_size - 1) // kv_block_size + 1
kv_cache = torch.zeros(
kv_n_blocks,
kv_block_size * (TOKEN_STRIDE + SCALE_DIM),
dtype=torch.uint8,
device=device,
)
kernel[(num_tokens,)](
state_cache,
state_cache.stride(0),
state_cache.stride(1),
token_to_req,
positions,
slot_mapping,
block_table,
block_table.stride(0),
BLOCK_SIZE,
rms_weight,
RMS_EPS,
cos_sin_cache,
cos_sin_cache.stride(0),
kv_cache,
slot_mapping,
kv_block_size,
HEAD_SIZE=HEAD_DIM,
TRITON_BLOCK_SIZE=HEAD_DIM,
STATE_WIDTH=coff * HEAD_DIM,
COMPRESS_RATIO=compress_ratio,
OVERLAP=overlap,
ROPE_HEAD_DIM=ROPE_DIM,
FP8_MAX=FP8_MAX,
QUANT_BLOCK=QUANT_BLOCK,
TOKEN_STRIDE=TOKEN_STRIDE,
SCALE_DIM=SCALE_DIM,
KV_BLOCK_STRIDE=kv_cache.stride(0),
num_warps=1,
)
k_quant, scale = _reference_kv_compress_norm_rope(
state_cache,
block_table,
positions,
rms_weight,
cos_sin_cache,
compress_ratio,
overlap,
use_fp4,
rms_eps=RMS_EPS,
fp8_max=FP8_MAX,
)
if use_fp4:
for i in range(num_tokens):
blk, pos = i // kv_block_size, i % kv_block_size
val_off = pos * TOKEN_STRIDE
fp4_actual = kv_cache[blk, val_off : val_off + TOKEN_STRIDE]
assert torch.equal(k_quant[i], fp4_actual), (
f"token {i}: packed nibbles differ, "
f"{(k_quant[i] != fp4_actual).sum()} "
f"/ {TOKEN_STRIDE}"
)
scale_off = kv_block_size * TOKEN_STRIDE + pos * SCALE_DIM
scale_actual = kv_cache[blk, scale_off : scale_off + SCALE_DIM]
assert torch.equal(scale_actual, scale[i]), (
f"token {i}: ue8m0 {scale_actual.tolist()} != {scale[i].tolist()}"
)
else:
k_quant = k_quant.view(torch.uint8)
for i in range(num_tokens):
blk, pos = i // kv_block_size, i % kv_block_size
val_off = pos * TOKEN_STRIDE
assert torch.equal(
k_quant[i], kv_cache[blk, val_off : val_off + TOKEN_STRIDE]
), f"token {i}: FP8 bytes differ"
scale_off = kv_block_size * TOKEN_STRIDE + pos * SCALE_DIM
actual_scale = kv_cache[blk, scale_off : scale_off + SCALE_DIM].view(
torch.float32
)
assert torch.equal(actual_scale, scale[i : i + 1]), (
f"token {i}: scale {actual_scale.item()} != {scale[i].item()}"
)
@@ -30,6 +30,56 @@ N_HEAD = 64
MAX_POS = 4096
def quantize_to_mxfp4(
x: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Reference MXFP4 quantization.
Args:
x: [..., head_dim] where head_dim is divisible by 32
Returns:
packed: [..., head_dim//2] uint8 2 E2M1 nibbles/byte, low nibble = even index
scales: [..., head_dim//32] uint8 1 ue8m0 byte
"""
MXFP4_BLOCK_SIZE = 32
orig_shape = x.shape
head_dim = orig_shape[-1]
n_blocks = head_dim // MXFP4_BLOCK_SIZE
x_f32 = x.float().reshape(-1, n_blocks, MXFP4_BLOCK_SIZE)
# Per-block ue8m0 scale: 2^ceil(log2(amax / 6.0)), stored as byte = exp + 127
# 6 * 2^-126 is from https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/inference/kernel.py#L163
amax = x_f32.abs().amax(dim=-1, keepdim=True).clamp(min=6 * (2**-126))
log2_ratio = (amax * (1.0 / 6.0)).log2().ceil().clamp(-127.0, 127.0)
scale = log2_ratio.exp2()
ue8m0 = (log2_ratio + 127.0).to(torch.uint8) # [*, n_blocks]
# E2M1 round-to-nearest-even: midpoints round to the even code.
# E2M1 values: [0.00, 0.50, 1.00, 1.50, 2.00, 3.00, 4.00, 6.00]
# boundaries: [ 0.25, 0.75, 1.25, 1.75, 2.50, 3.50, 5.00]
x_scaled = (x_f32 / scale).clamp(-6.0, 6.0)
abs_x = x_scaled.abs()
code = torch.zeros_like(abs_x, dtype=torch.int32)
code = torch.where(abs_x > 0.25, 1, code)
code = torch.where(abs_x >= 0.75, 2, code)
code = torch.where(abs_x > 1.25, 3, code)
code = torch.where(abs_x >= 1.75, 4, code)
code = torch.where(abs_x > 2.5, 5, code)
code = torch.where(abs_x >= 3.5, 6, code)
code = torch.where(abs_x > 5.0, 7, code)
sign = ((x_scaled.view(torch.int32) >> 31) & 1).to(torch.uint8)
nibble = code.to(torch.uint8) | (sign << 3)
# Pack: even-index element → low nibble, odd-index → high nibble
nibble_flat = nibble.reshape(-1, head_dim)
packed = (nibble_flat[:, 0::2] | (nibble_flat[:, 1::2] << 4)).contiguous()
packed = packed.reshape(*orig_shape[:-1], head_dim // 2)
scales = ue8m0.view(*orig_shape[:-1], n_blocks)
return packed, scales
def _reference(
positions: torch.Tensor,
q: torch.Tensor,
@@ -37,6 +87,7 @@ def _reference(
weights: torch.Tensor,
softmax_scale: float,
head_scale: float,
use_fp4: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
q_rot = q.clone()
ops.rotary_embedding(
@@ -49,22 +100,33 @@ def _reference(
HEAD_DIM - ROPE_DIM, # rope_dim_offset → rotate the tail
False,
)
q_fp8, q_scale = per_token_group_quant_fp8(
q_rot.view(-1, HEAD_DIM).contiguous(),
HEAD_DIM,
use_ue8m0=True,
)
q_fp8 = q_fp8.view(-1, N_HEAD, HEAD_DIM)
q_scale = q_scale.view(-1, N_HEAD)
weights_out = weights.to(torch.float32) * q_scale * softmax_scale * head_scale
return q_fp8, weights_out
if use_fp4:
q_packed, ue8m0 = quantize_to_mxfp4(q_rot.view(-1, N_HEAD, HEAD_DIM))
# Pack 4 ue8m0 bytes into 1 int32
q_scale = ue8m0.view(torch.int32).squeeze(-1)
# FP4 path: q_scale stays separate (cannot be folded into a per-token scalar)
weights_out = weights.to(torch.float32) * softmax_scale * head_scale
return (q_packed, q_scale), weights_out
else:
q_fp8, q_scale = per_token_group_quant_fp8(
q_rot.view(-1, HEAD_DIM).contiguous(),
HEAD_DIM,
use_ue8m0=True,
)
q_fp8 = q_fp8.view(-1, N_HEAD, HEAD_DIM)
q_scale = q_scale.view(-1, N_HEAD)
weights_out = weights.to(torch.float32) * q_scale * softmax_scale * head_scale
return q_fp8, weights_out
@pytest.mark.parametrize("num_tokens", [1, 7, 32, 257])
@pytest.mark.parametrize("cache_dtype", [torch.float32, torch.bfloat16])
@pytest.mark.parametrize("use_fp4", [False, True])
@torch.inference_mode()
def test_fused_indexer_q_rope_quant_matches_unfused(num_tokens, cache_dtype):
def test_fused_indexer_q_rope_quant_matches_unfused(num_tokens, cache_dtype, use_fp4):
device = "cuda"
torch.manual_seed(0)
@@ -77,21 +139,32 @@ def test_fused_indexer_q_rope_quant_matches_unfused(num_tokens, cache_dtype):
softmax_scale = HEAD_DIM**-0.5
head_scale = N_HEAD**-0.5
q_fp8_ref, weights_ref = _reference(
positions, q, cos_sin_cache, weights, softmax_scale, head_scale
q_quant_ref, weights_ref = _reference(
positions, q, cos_sin_cache, weights, softmax_scale, head_scale, use_fp4
)
q_fp8_fused, weights_fused = fused_indexer_q_rope_quant(
positions, q.clone(), cos_sin_cache, weights, softmax_scale, head_scale
q_quant_fused, weights_fused = fused_indexer_q_rope_quant(
positions, q.clone(), cos_sin_cache, weights, softmax_scale, head_scale, use_fp4
)
if use_fp4:
q_quant_ref, q_scale_ref = q_quant_ref
q_quant_fused, q_scale_fused = q_quant_fused
assert torch.equal(q_scale_ref, q_scale_fused), (
f"q_scale mismatch: "
f"{(q_scale_ref != q_scale_fused).sum().item()} "
f"/ {q_scale_ref.numel()} bytes differ"
)
# fp8 tensors aren't directly comparable via torch.equal — reinterpret as int8.
ref_bits = q_fp8_ref.view(torch.int8)
fused_bits = q_fp8_fused.view(torch.int8)
ref_bits = q_quant_ref.view(torch.int8)
fused_bits = q_quant_fused.view(torch.int8)
assert torch.equal(ref_bits, fused_bits), (
f"q_fp8 mismatch: "
f"q_quant_fused mismatch: "
f"{(ref_bits != fused_bits).sum().item()} / {ref_bits.numel()} bytes differ"
)
assert weights_fused.dtype == torch.float32
assert torch.equal(weights_ref, weights_fused), (
f"weights mismatch: max abs diff "
f"{(weights_ref - weights_fused).abs().max().item()}"
+157
View File
@@ -0,0 +1,157 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Precision tests for vllm's chunk_kda Triton operator.
Compares chunk_kda against a naive recurrent reference (float32).
Uses torch.rand for q/k/v to match FLA's test pattern.
"""
import pytest
import torch
import torch.nn.functional as F
from vllm.model_executor.layers.fla.ops.kda import chunk_kda
from vllm.model_executor.layers.fla.ops.l2norm import l2norm_fwd
DEVICE = "cuda"
def naive_recurrent_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Naive recurrent KDA reference, ported from FLA's naive.py."""
dtype = v.dtype
B, T, H, K = q.shape
V = v.shape[-1]
if scale is None:
scale = K**-0.5
q, k, v, g, beta = (x.to(torch.float) for x in [q, k, v, g, beta])
q = q * scale
S = k.new_zeros(B, H, K, V).to(q)
if initial_state is not None:
S += initial_state
o = torch.zeros_like(v)
for i in range(T):
q_i, k_i, v_i, g_i, b_i = q[:, i], k[:, i], v[:, i], g[:, i], beta[:, i]
S = S * g_i[..., None].exp()
S = S + torch.einsum(
"bhk,bhv->bhkv",
b_i[..., None] * k_i,
v_i - (k_i[..., None] * S).sum(-2),
)
o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S)
if not output_final_state:
S = None
return o.to(dtype), S
def assert_close(
name: str,
ref: torch.Tensor,
tri: torch.Tensor,
ratio: float,
err_atol: float = 1e-6,
):
"""RMSE-based relative error comparison."""
abs_err = (ref.detach() - tri.detach()).flatten().abs().max().item()
rmse_diff = (ref.detach() - tri.detach()).flatten().square().mean().sqrt().item()
rmse_base = ref.detach().flatten().square().mean().sqrt().item()
rel_err = rmse_diff / (rmse_base + 1e-8)
print(f"{name:>4} | abs={abs_err:.6f} | rmse={rel_err:.6f} | thr={ratio}")
if abs_err <= err_atol:
return
assert not torch.isnan(ref).any(), f"{name}: NaN detected in ref"
assert not torch.isnan(tri).any(), f"{name}: NaN detected in tri"
assert rel_err < ratio, (
f"{name}: max abs err {abs_err:.6f}, rmse ratio {rel_err:.6f} >= {ratio}"
)
@pytest.mark.parametrize(
("H", "D", "cu_seqlens", "dtype"),
[
pytest.param(
*test,
id="H{}-D{}-cu{}-{}".format(*test),
)
for test in [
(32, 128, [0, 64], torch.float16),
(32, 128, [0, 1024], torch.float16),
(32, 128, [0, 15], torch.float16),
(32, 128, [0, 256, 512, 768, 1024], torch.float16),
(32, 128, [0, 15, 100, 300, 1200], torch.float16),
(64, 128, [0, 256, 500, 1000], torch.float16),
(32, 128, [0, 8192], torch.float16),
(32, 128, [0, 256, 500, 1000], torch.bfloat16),
]
],
)
@torch.inference_mode()
def test_chunk_kda(
H: int,
D: int,
cu_seqlens: list[int],
dtype: torch.dtype,
):
T = cu_seqlens[-1]
torch.manual_seed(42)
B = 1
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
N = len(cu_seqlens) - 1
q = torch.rand(B, T, H, D, dtype=dtype, device=DEVICE)
k = torch.rand(B, T, H, D, dtype=dtype, device=DEVICE)
v = torch.rand(B, T, H, D, dtype=dtype, device=DEVICE)
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)).to(
dtype
)
beta = torch.rand(B, T, H, dtype=dtype, device=DEVICE).sigmoid()
h0 = torch.randn(N, H, D, D, dtype=torch.float32, device=DEVICE)
# Naive reference with l2norm_fwd (same kernel as chunk_kda)
ref_outputs = []
ref_states = []
for i in range(N):
s, e = cu_seqlens[i], cu_seqlens[i + 1]
q_i = l2norm_fwd(q[:, s:e].contiguous())
k_i = l2norm_fwd(k[:, s:e].contiguous())
o_i, ht_i = naive_recurrent_kda(
q_i,
k_i,
v[:, s:e],
g[:, s:e],
beta[:, s:e],
initial_state=h0[i],
output_final_state=True,
)
ref_outputs.append(o_i)
ref_states.append(ht_i)
ref_o = torch.cat(ref_outputs, dim=1)
ref_ht = torch.cat(ref_states, dim=0)
# h0 transposed to (V, K) layout for the kernel; naive uses (K, V)
tri_o, tri_ht = chunk_kda(
q=q.clone(),
k=k.clone(),
v=v.clone(),
g=g.clone(),
beta=beta.clone(),
initial_state=h0.transpose(-1, -2).contiguous().clone(),
output_final_state=True,
cu_seqlens=cu_seqlens_t,
use_qk_l2norm_in_kernel=True,
)
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
assert not torch.isnan(tri_ht).any(), "Triton output ht contains NaN"
assert_close("o", ref_o, tri_o, 0.005)
assert_close("ht", ref_ht, tri_ht.transpose(-1, -2).contiguous(), 0.005)
@@ -1,11 +1,18 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import importlib.util
import pytest
import torch
from ....conftest import VllmRunner
pytestmark = pytest.mark.skipif(
importlib.util.find_spec("terratorch") is None,
reason="terratorch unavailable while PyPI has `lightning` quarantined; see #41376",
)
def _run_test(
vllm_runner: type[VllmRunner],
@@ -311,6 +311,9 @@ def _test_processing_correctness(
baseline_processor,
cached_processor,
batch_idx,
hit_rate,
num_batches,
simplify_rate,
)
@@ -320,6 +323,9 @@ def _test_processing_correctness_one(
baseline_processor: BaseMultiModalProcessor,
cached_processor: BaseMultiModalProcessor,
batch_idx: int,
hit_rate: float,
num_batches: int,
simplify_rate: float,
):
model_type = model_config.hf_config.model_type
@@ -343,7 +349,11 @@ def _test_processing_correctness_one(
baseline_tokenized_result,
cached_tokenized_result,
ignore_mm_keys=ignore_mm_keys,
msg=f"Failed ({batch_idx=}, {token_prompt=}, {mm_data=})",
msg=(
f"Failed ({batch_idx=}, {hit_rate=}, "
f"{num_batches=}, {simplify_rate=}, "
f"{text_prompt=}, {token_prompt=}, {mm_data=})"
),
)
if text_prompt is not None:
@@ -362,21 +372,33 @@ def _test_processing_correctness_one(
baseline_text_result,
cached_text_result,
ignore_mm_keys=ignore_mm_keys,
msg=f"Failed ({batch_idx=}, {text_prompt=}, {mm_data=})",
msg=(
f"Failed ({batch_idx=}, {hit_rate=}, "
f"{num_batches=}, {simplify_rate=}, "
f"{text_prompt=}, {token_prompt=}, {mm_data=})"
),
)
_assert_inputs_equal(
baseline_text_result,
baseline_tokenized_result,
ignore_mm_keys=ignore_mm_keys,
msg=f"Failed ({batch_idx=}, {text_prompt=}, {token_prompt=}, {mm_data=})",
msg=(
f"Failed ({batch_idx=}, {hit_rate=}, "
f"{num_batches=}, {simplify_rate=}, "
f"{text_prompt=}, {token_prompt=}, {mm_data=})"
),
)
_assert_inputs_equal(
cached_text_result,
cached_tokenized_result,
ignore_mm_keys=ignore_mm_keys,
msg=f"Failed ({batch_idx=}, {text_prompt=}, {token_prompt=}, {mm_data=})",
msg=(
f"Failed ({batch_idx=}, {hit_rate=}, "
f"{num_batches=}, {simplify_rate=}, "
f"{text_prompt=}, {token_prompt=}, {mm_data=})"
),
)
@@ -408,6 +430,8 @@ def test_processing_correctness(
"correctness test as is. Let's revisit adapting this "
"test once more realtime models exist."
)
if model_id == "CohereLabs/cohere-transcribe-03-2026":
pytest.skip("Fix later")
_test_processing_correctness(
model_id,
+1 -3
View File
@@ -1396,9 +1396,7 @@ _MULTIMODAL_EXAMPLE_MODELS = {
),
# [Encoder-decoder]
"CohereAsrForConditionalGeneration": _HfExamplesInfo(
"CohereLabs/cohere-transcribe-03-2026",
trust_remote_code=True,
is_available_online=False, # TODO (ekagra): revert after asr release
"CohereLabs/cohere-transcribe-03-2026", trust_remote_code=True
),
"NemotronParseForConditionalGeneration": _HfExamplesInfo(
"nvidia/NVIDIA-Nemotron-Parse-v1.1", trust_remote_code=True
+10
View File
@@ -109,6 +109,16 @@ def can_initialize(
"which is not configured in test environment"
)
if model_arch in ("PrithviGeoSpatialMAE", "Terratorch"):
import importlib.util
if importlib.util.find_spec("terratorch") is None:
pytest.skip(
"terratorch is not installed; "
"temporarily skipped while PyPI has `lightning` quarantined "
"(see #41376)"
)
if model_arch in ["DeepseekV32ForCausalLM", "GlmMoeDsaForCausalLM"]:
from vllm.platforms import current_platform
+11
View File
@@ -36,6 +36,17 @@ def test_registry_imports(model_arch):
check_max_version=False,
check_version_reason="vllm",
)
if model_arch in ("PrithviGeoSpatialMAE", "Terratorch"):
import importlib.util
if importlib.util.find_spec("terratorch") is None:
pytest.skip(
"terratorch is not installed; "
"temporarily skipped while PyPI has `lightning` quarantined "
"(see #41376)"
)
# Ensure all model classes can be imported successfully
model_cls = ModelRegistry._try_load_model_cls(model_arch)
assert model_cls is not None
+7
View File
@@ -1,12 +1,19 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import importlib.util
import pytest
import torch
from tests.conftest import VllmRunner
from tests.utils import create_new_process_for_each_test
pytestmark = pytest.mark.skipif(
importlib.util.find_spec("terratorch") is None,
reason="terratorch unavailable while PyPI has `lightning` quarantined; see #41376",
)
@create_new_process_for_each_test() # Hangs otherwise
@pytest.mark.parametrize(
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import importlib.util
import io
import imagehash
@@ -11,6 +12,11 @@ from PIL import Image
from tests.utils import RemoteOpenAIServer
from vllm.entrypoints.pooling.pooling.protocol import IOProcessorResponse
pytestmark = pytest.mark.skipif(
importlib.util.find_spec("terratorch") is None,
reason="terratorch unavailable while PyPI has `lightning` quarantined; see #41376",
)
models_config = {
"ibm-nasa-geospatial/Prithvi-EO-2.0-300M-TL-Sen1Floods11": {
"image_url": "https://huggingface.co/christian-pinto/Prithvi-EO-2.0-300M-TL-VLLM/resolve/main/valencia_example_2024-10-26.tiff", # noqa: E501
@@ -1,63 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for QuarkConfig.maybe_update_config.
Fetches real HF configs (metadata only, no model weights) to verify
that dynamic_mxfp4_quant is only enabled for DeepSeek-V3-family models.
Run: pytest tests/quantization/test_quark_maybe_update_config.py -v
"""
import pytest
from transformers import AutoConfig
from vllm.model_executor.layers.quantization.quark.quark import QuarkConfig
def _make_quark_config() -> QuarkConfig:
"""Create a minimal QuarkConfig for testing."""
return QuarkConfig(quant_config={}, kv_cache_group=[], pack_method="reorder")
# ---------------------------------------------------------------------------
# Non-deepseek models must not flip dynamic_mxfp4_quant
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"model_name",
["amd/MiniMax-M2.1-MXFP4"],
)
def test_non_deepseek_model_stays_false(model_name: str):
"""Non-deepseek_v3 models must not enable dynamic_mxfp4_quant."""
hf_config = AutoConfig.from_pretrained(model_name, trust_remote_code=True)
qcfg = _make_quark_config()
qcfg.maybe_update_config(model_name, hf_config=hf_config)
assert qcfg.dynamic_mxfp4_quant is False
# ---------------------------------------------------------------------------
# DeepSeek-V3 family + fp4 must enable dynamic_mxfp4_quant
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"model_name",
["amd/DeepSeek-R1-MXFP4-ASQ"],
)
def test_deepseek_family_fp4_enables_flag(model_name: str):
hf_config = AutoConfig.from_pretrained(model_name, trust_remote_code=True)
qcfg = _make_quark_config()
qcfg.maybe_update_config(model_name, hf_config=hf_config)
assert qcfg.dynamic_mxfp4_quant is True
# ---------------------------------------------------------------------------
# Missing hf_config → warn and stay False
# ---------------------------------------------------------------------------
def test_missing_hf_config_stays_false():
qcfg = _make_quark_config()
qcfg.maybe_update_config("some/model")
assert qcfg.dynamic_mxfp4_quant is False
-103
View File
@@ -1,103 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for the FxGraphCachePickler.dumps ValueError patch in env_override.py.
Validates that _apply_fxgraphcache_pickle_patch correctly wraps a pickler's
dumps method to convert ValueError into a bypass exception, without affecting
other exception types or normal return values.
"""
import pytest
from vllm.env_override import _apply_fxgraphcache_pickle_patch
class _BypassStub(Exception):
"""Stand-in for BypassFxGraphCache in unit tests."""
class TestApplyFxgraphcachePicklePatch:
def test_valueerror_converted_to_bypass(self):
class Pickler:
def dumps(self, obj):
raise ValueError("can't serialize blocked layout")
_apply_fxgraphcache_pickle_patch(Pickler, _BypassStub)
with pytest.raises(_BypassStub, match="Failed to pickle cache key"):
Pickler().dumps(object())
def test_original_valueerror_chained(self):
class Pickler:
def dumps(self, obj):
raise ValueError("bad tensor layout")
_apply_fxgraphcache_pickle_patch(Pickler, _BypassStub)
with pytest.raises(_BypassStub) as exc_info:
Pickler().dumps(object())
cause = exc_info.value.__cause__
assert isinstance(cause, ValueError)
assert str(cause) == "bad tensor layout"
def test_non_valueerror_propagates(self):
class Pickler:
def dumps(self, obj):
raise TypeError("unexpected type")
_apply_fxgraphcache_pickle_patch(Pickler, _BypassStub)
with pytest.raises(TypeError, match="unexpected type"):
Pickler().dumps(object())
def test_normal_return_preserved(self):
sentinel = b"serialized-graph-key"
class Pickler:
def dumps(self, obj):
return sentinel
_apply_fxgraphcache_pickle_patch(Pickler, _BypassStub)
assert Pickler().dumps(object()) is sentinel
def test_idempotent(self):
class Pickler:
def dumps(self, obj):
return b"ok"
_apply_fxgraphcache_pickle_patch(Pickler, _BypassStub)
first_dumps = Pickler.dumps
_apply_fxgraphcache_pickle_patch(Pickler, _BypassStub)
assert Pickler.dumps is first_dumps
def test_sentinel_attribute_set(self):
class Pickler:
def dumps(self, obj):
return b"ok"
assert not hasattr(Pickler.dumps, "_vllm_patched")
assert not getattr(Pickler, "_vllm_fxgraph_dumps_patched", False)
_apply_fxgraphcache_pickle_patch(Pickler, _BypassStub)
assert Pickler.dumps._vllm_patched is True # type: ignore[attr-defined]
assert Pickler._vllm_fxgraph_dumps_patched is True # type: ignore[attr-defined]
def test_patch_applied_in_current_environment():
"""Integration: verify patch state matches current torch version."""
from torch._inductor.codecache import FxGraphCachePickler
from vllm.utils.torch_utils import is_torch_equal_or_newer
should_be_patched = is_torch_equal_or_newer(
"2.10.0"
) and not is_torch_equal_or_newer("2.11.0")
assert getattr(FxGraphCachePickler, "_vllm_fxgraph_dumps_patched", False) == (
should_be_patched
)
assert hasattr(FxGraphCachePickler.dumps, "_vllm_patched") == should_be_patched
+55 -6
View File
@@ -14,9 +14,10 @@ import signal
import subprocess
import sys
import tempfile
import threading
import time
import warnings
from collections.abc import Callable, Iterable
from collections.abc import Callable, Iterable, Sequence
from contextlib import ExitStack, contextmanager, suppress
from multiprocessing import Process
from pathlib import Path
@@ -62,8 +63,6 @@ from vllm.utils.torch_utils import (
FP8_DTYPE = current_platform.fp8_dtype()
if current_platform.is_rocm():
import threading
from amdsmi import (
amdsmi_get_gpu_vram_usage,
amdsmi_get_processor_handles,
@@ -247,6 +246,16 @@ 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()
def _terminate_process_tree(self) -> None:
"""Kill the server process tree without waiting for GPU memory release.
Split out from ``_shutdown`` so that ``shutdown_many`` can run this
phase in parallel for sibling servers and then wait for GPU memory
release once at the end.
"""
pid = self.proc.pid
# Get the process group ID. Because we used
@@ -288,9 +297,49 @@ class RemoteVLLMServer:
# prevent VRAM from being reclaimed by the driver.
self._kill_process_group_survivors(pgid)
# Wait for GPU memory to actually be freed, not just
# "stabilized at whatever level it's at".
self._wait_for_gpu_memory_release()
@classmethod
def shutdown_many(cls, servers: Sequence["RemoteVLLMServer"]) -> None:
"""Shut down multiple sibling servers and wait for GPU memory once.
Test fixtures that hold several ``RemoteVLLMServer`` instances at
once must NOT shut them down by calling each server's ``__exit__``
sequentially: every server measures total GPU memory across all
visible devices in ``_wait_for_gpu_memory_release``, so the first
server's wait blocks the full timeout because later sibling
servers are still holding GPU memory.
Instead, this method terminates every server's process tree in
parallel, then runs the GPU-memory-release wait once against the
earliest recorded baseline (memory before any server started).
"""
if not servers:
return
threads = [
threading.Thread(
target=s._terminate_process_tree,
name=f"shutdown-{s.proc.pid}",
daemon=True,
)
for s in servers
]
for t in threads:
t.start()
for t in threads:
t.join()
# Use the smallest pre-server baseline so the wait targets memory
# usage before *any* of these sibling servers started, not after
# earlier siblings had already allocated.
earliest = min(
servers,
key=lambda s: (
float("inf")
if s._pre_server_gpu_memory is None
else s._pre_server_gpu_memory
),
)
earliest._wait_for_gpu_memory_release()
def _kill_process_group_survivors(
self, pgid: int | None, timeout: float = 15.0
@@ -432,3 +432,52 @@ def test_chunked_local_attention_get_num_blocks_to_allocate():
)
== 15
)
def test_predictor_matches_allocator_blocks_calculation_with_admission_cap():
"""In forward steps, `get_num_blocks_to_allocate` must return exactly what
`allocate_new_blocks` will pull; otherwise `block_pool.get_new_blocks`
raises `ValueError: Cannot get N free blocks from the pool`.
"""
block_size = 2
sliding_window = 8 # 4-block live window
cap = sliding_window // block_size
spec = SlidingWindowSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
sliding_window=sliding_window,
)
block_pool = BlockPool(
num_gpu_blocks=100, enable_caching=True, hash_block_size=block_size
)
manager = SlidingWindowManager(
spec,
block_pool=block_pool,
enable_caching=False,
kv_cache_group_id=0,
max_admission_blocks_per_request=cap,
)
request_id = "req"
total_computed = 0
# Walk through request forward steps. Check num_blocks returned by
# `get_num_blocks_to_allocate` matches what `allocate_new_blocks` pulls
for num_tokens in (4, 8, 12, 16):
predicted = manager.get_num_blocks_to_allocate(
request_id=request_id,
num_tokens=num_tokens,
new_computed_blocks=[],
total_computed_tokens=total_computed,
num_tokens_main_model=num_tokens,
)
new_blocks = manager.allocate_new_blocks(
request_id, num_tokens=num_tokens, num_tokens_main_model=num_tokens
)
assert predicted == len(new_blocks), (
f"num_tokens={num_tokens}: predictor returned {predicted} "
f"but allocator pulled {len(new_blocks)}"
)
total_computed = num_tokens
+122 -7
View File
@@ -365,10 +365,13 @@ async def test_dp_pause_keep_race_staggered_engines():
async def staggered_pause_keep(method: str, *args) -> Any:
if method != "pause_scheduler" or not args or args[0] != "keep":
return await original_call_utility(method, *args)
# Send pause(keep) to engine 0 first
await client._call_utility_async(
method, *args, engine=client.core_engines[0]
# Fire pause(keep) to engine 0 (don't await — with DP
# two-phase pause, consensus requires all ranks).
pause_0 = asyncio.create_task(
client._call_utility_async(method, *args, engine=client.core_engines[0])
)
# Let the event loop send the message to engine 0.
await asyncio.sleep(0.5)
# In the middle: send two requests (race window)
sp = SamplingParams(max_tokens=5, ignore_eos=True)
@@ -384,11 +387,13 @@ async def test_dp_pause_keep_race_staggered_engines():
t2 = asyncio.create_task(consume_gen("race-2"))
mid_pause_tasks.extend([t1, t2])
await asyncio.sleep(3)
# Then send pause(keep) to engine 1
result = await client._call_utility_async(
method, *args, engine=client.core_engines[1]
# Fire pause(keep) to engine 1, then await both so
# consensus can be reached.
pause_1 = asyncio.create_task(
client._call_utility_async(method, *args, engine=client.core_engines[1])
)
return result
results = await asyncio.gather(pause_0, pause_1)
return results[0]
client.call_utility_async = staggered_pause_keep
@@ -398,3 +403,113 @@ async def test_dp_pause_keep_race_staggered_engines():
assert not await engine.is_paused()
# Let the two requests we sent mid-pause complete
await asyncio.gather(*mid_pause_tasks)
@pytest.mark.asyncio
async def test_dp_pause_barrier_request_deadlock():
"""
Test that start_dp_wave is ignored while paused.
Sequence:
1. Pause all engines (PAUSED_ALL).
2. Send barrier to engine 0 only blocks in dist.barrier(dp_group).
3. Send a request routed to engine 1.
4. Wait for any (buggy) START_DP_WAVE propagation.
5. Send barrier to engine 1 completes in fixed code, deadlocks
in buggy code because engine 1 is stuck in EP all-to-all.
"""
if DP_SIZE != 2:
pytest.skip("requires DP_SIZE=2")
with ExitStack() as after:
engine_args = _get_dp_pause_engine_args(expert_parallel=True)
engine = AsyncLLM.from_engine_args(engine_args)
after.callback(engine.shutdown)
client = engine.engine_core
# Cache get_supported_tasks so that generate() won't need to
# send a utility call to all engines (which would hang once
# engine 0 is blocked in the barrier).
await engine.get_supported_tasks()
# Pause all engines normally — no staggering.
await engine.pause_generation(mode="keep")
assert await engine.is_paused()
original_call_utility = client.call_utility_async
mid_barrier_tasks: list[asyncio.Task] = []
async def staggered_barrier(method: str, *args) -> Any:
if method != "barrier":
return await original_call_utility(method, *args)
# Send barrier to engine 0 only — it blocks in
# dist.barrier(dp_group) waiting for engine 1.
barrier_0 = asyncio.create_task(
client._call_utility_async(method, *args, engine=client.core_engines[0])
)
await asyncio.sleep(1)
# While engine 0 is blocked, send a request routed
# specifically to engine 1.
sp = SamplingParams(max_tokens=5, ignore_eos=True)
engine_1 = client.core_engines[1]
original_get_engine = client.get_core_engine_for_request
def route_to_engine_1(req):
client.reqs_in_flight[req.request_id] = engine_1
return engine_1
client.get_core_engine_for_request = route_to_engine_1
async def consume_gen(req_id: str) -> None:
async for _ in engine.generate(
request_id=req_id,
prompt=DP_PAUSE_PROMPT,
sampling_params=sp,
):
pass
t1 = asyncio.create_task(consume_gen("race-1"))
mid_barrier_tasks.append(t1)
# Yield so generate() preprocessing completes and
# add_request_async is called (which, in buggy code,
# would send FIRST_REQ and wake engine 1).
for _ in range(200):
await asyncio.sleep(0)
client.get_core_engine_for_request = original_get_engine
# Wait for any START_DP_WAVE to propagate and for
# engine 1 to potentially enter execute_dummy_batch.
await asyncio.sleep(5)
# Now send barrier to engine 1. In buggy code engine 1
# is stuck in execute_dummy_batch (EP all-to-all) while
# engine 0 is stuck in dist.barrier(dp_group) — deadlock.
result = await client._call_utility_async(
method, *args, engine=client.core_engines[1]
)
await barrier_0
return result
client.call_utility_async = staggered_barrier
# Drive the staggered barrier. Old code deadlocks here.
try:
await asyncio.wait_for(client.call_utility_async("barrier"), timeout=30)
except asyncio.TimeoutError:
for t in mid_barrier_tasks:
t.cancel()
pytest.fail(
"Staggered barrier deadlocked — FIRST_REQ sent while "
"paused caused collective-op mismatch between engines"
)
await engine.resume_generation()
assert not await engine.is_paused()
# Let the two requests we sent mid-barrier complete.
await asyncio.gather(*mid_barrier_tasks)
+6 -5
View File
@@ -111,11 +111,12 @@ class ExternalLBServerManager:
def __exit__(self, exc_type, exc_val, exc_tb):
"""Stop all server instances."""
while self.servers:
try:
self.servers.pop()[0].__exit__(exc_type, exc_val, exc_tb)
except Exception as e:
print(f"Error stopping server: {e}")
servers = [s for s, _ in self.servers]
self.servers.clear()
try:
RemoteOpenAIServer.shutdown_many(servers)
except Exception as e:
print(f"Error stopping servers: {e}")
@pytest.fixture(scope="module")
+6 -5
View File
@@ -134,11 +134,12 @@ class HybridLBServerManager:
def __exit__(self, exc_type, exc_val, exc_tb):
"""Stop all server instances."""
while self.servers:
try:
self.servers.pop()[0].__exit__(exc_type, exc_val, exc_tb)
except Exception as e:
print(f"Error stopping server: {e}")
servers = [s for s, _ in self.servers]
self.servers.clear()
try:
RemoteOpenAIServer.shutdown_many(servers)
except Exception as e:
print(f"Error stopping servers: {e}")
@pytest.fixture(scope="module")
+14 -14
View File
@@ -228,13 +228,13 @@ class MultinodeInternalLBServerManager:
def __exit__(self, exc_type, exc_val, exc_tb):
"""Stop all server instances."""
while self.servers:
if server := self.servers.pop():
try:
server[0].__exit__(exc_type, exc_val, exc_tb)
except Exception as e:
print(f"Error stopping server: {e}")
traceback.print_exc()
servers = [entry[0] for entry in self.servers if entry is not None]
self.servers.clear()
try:
RemoteOpenAIServer.shutdown_many(servers)
except Exception as e:
print(f"Error stopping servers: {e}")
traceback.print_exc()
class APIOnlyServerManager:
@@ -370,13 +370,13 @@ class APIOnlyServerManager:
def __exit__(self, exc_type, exc_val, exc_tb):
"""Stop both server instances."""
while self.servers:
if server := self.servers.pop():
try:
server[0].__exit__(exc_type, exc_val, exc_tb)
except Exception as e:
print(f"Error stopping server: {e}")
traceback.print_exc()
servers = [entry[0] for entry in self.servers if entry is not None]
self.servers.clear()
try:
RemoteOpenAIServer.shutdown_many(servers)
except Exception as e:
print(f"Error stopping servers: {e}")
traceback.print_exc()
@pytest.fixture(scope="module")
@@ -82,6 +82,13 @@ SPEC_DECODE_CONFIGS = [
2,
id="eagle-mla-deepseek",
),
pytest.param(
"Qwen/Qwen3.5-0.8B-Base",
"Qwen/Qwen3.5-0.8B-Base",
"mtp",
1,
id="mtp-qwen3_5-hybrid",
),
]
@@ -104,6 +111,14 @@ def test_no_sync_with_spec_decode(
from vllm import LLM, SamplingParams
from vllm.distributed import cleanup_dist_env_and_memory
# Qwen3.5 is a VLM; without this, profile_run runs the ViT warmup
# and peaks well above the 18GB MIG slice used by one of the CI lanes.
# This test only exercises text generation, so the vision tower is
# never needed.
extra_kwargs: dict = {}
if "Qwen3.5" in model:
extra_kwargs["limit_mm_per_prompt"] = {"image": 0, "video": 0}
llm = LLM(
model=model,
max_model_len=256,
@@ -114,6 +129,7 @@ def test_no_sync_with_spec_decode(
},
enforce_eager=True,
async_scheduling=True,
**extra_kwargs,
)
# Assert async scheduling is actually active before running inference.
+24 -1
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import os
import random
from collections.abc import Iterable
from dataclasses import dataclass
@@ -723,8 +724,13 @@ def test_eagle_correctness_heavy(
[
(("mtp", "XiaomiMiMo/MiMo-7B-Base", 1), False, 0.5), # ref: 65%-70%
(("mtp", "ZixiQi/DeepSeek-V3-4layers-MTP-FP8", 1), False, 0.0), # dummy model
(
("mtp", "Qwen/Qwen3.5-0.8B-Base", 1),
False,
0.20,
), # hybrid + MTP, ref: ~34%-35%
],
ids=["mimo", "deepseek"],
ids=["mimo", "deepseek", "qwen3_5-hybrid"],
)
@single_gpu_only
@large_gpu_mark(min_gb=20)
@@ -750,13 +756,29 @@ def test_mtp_correctness(
method, model_name, tp_size = model_setup
_skip_if_insufficient_gpus_for_tp(tp_size)
if "Qwen3.5" in model_name and os.environ.get("VLLM_USE_V2_MODEL_RUNNER"):
pytest.skip(
"Model Runner V2 does not yet support hybrid models "
"(Qwen3.5 mixes Mamba-style GDN with attention layers)."
)
attn_backend = "TRITON_ATTN" if current_platform.is_rocm() else "auto"
# Qwen3.5 is a VLM; without this, profile_run runs the ViT warmup
# and peaks well above the 18GB MIG slice used by one of the CI
# lanes. This test only exercises text generation, so the vision
# tower is never needed.
extra_kwargs: dict[str, Any] = {}
if "Qwen3.5" in model_name:
extra_kwargs["limit_mm_per_prompt"] = {"image": 0, "video": 0}
ref_llm = LLM(
model=model_name,
max_model_len=2048,
tensor_parallel_size=tp_size,
trust_remote_code=True,
attention_backend=attn_backend,
**extra_kwargs,
)
ref_outputs = ref_llm.chat(test_prompts, sampling_config)
evaluate_llm_for_gsm8k(
@@ -777,6 +799,7 @@ def test_mtp_correctness(
},
max_model_len=2048,
attention_backend=attn_backend,
**extra_kwargs,
)
# MTP supports async scheduling; assert it is active by default.
assert spec_llm.llm_engine.vllm_config.scheduler_config.async_scheduling
@@ -15,7 +15,7 @@ from vllm.distributed.kv_transfer.kv_connector.v1.offloading.scheduler import (
OffloadingConnectorScheduler,
)
from vllm.v1.core.kv_cache_utils import BlockHash
from vllm.v1.kv_offload.abstract import (
from vllm.v1.kv_offload.base import (
OffloadingEvent,
OffloadingManager,
ReqContext,
@@ -20,7 +20,7 @@ from vllm.v1.kv_cache_interface import (
MLAAttentionSpec,
UniformTypeKVCacheSpecs,
)
from vllm.v1.kv_offload.spec import (
from vllm.v1.kv_offload.base import (
CanonicalKVCacheRef,
CanonicalKVCaches,
OffloadingSpec,
@@ -37,15 +37,15 @@ from vllm.v1.kv_cache_interface import (
KVCacheConfig,
KVCacheGroupSpec,
)
from vllm.v1.kv_offload.abstract import (
from vllm.v1.kv_offload.base import (
GPULoadStoreSpec,
LoadStoreSpec,
OffloadingManager,
OffloadingSpec,
OffloadKey,
PrepareStoreOutput,
make_offload_key,
)
from vllm.v1.kv_offload.mediums import GPULoadStoreSpec
from vllm.v1.kv_offload.spec import OffloadingSpec
from vllm.v1.kv_offload.worker.worker import (
OffloadingHandler,
TransferResult,
@@ -0,0 +1,915 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for bi-directional KV cache transfer between P and D nodes.
Tests cover the new behaviors added by the bi-directional KV transfer PR:
1. P-node scheduler lifecycle: P pulls KV from D using remote_block_ids,
eliminating redundant prefill computation in multi-turn conversations.
2. P-node metadata: NixlConnectorMetadata correctly populates recv metadata
when P pulls KV from D (do_remote_decode=True + remote_block_ids).
3. P-node worker: start_load_kv processes reqs_to_recv for KV pull from D.
4. D-node request_finished: returns kv_transfer_params with remote_block_ids
and remote_num_tokens so P can pull KV in future turns.
5. Edge cases:
- No double read after reschedule (_remote_blocks_processed flag)
- remote_num_tokens bounded by block capacity (num_computed_tokens)
- kv_recompute_threshold skips small transfers
- P-node holds blocks for D after finishing
- Cache MISS first turn falls back to local prefill
- Partial remote coverage: P pulls partial, computes the rest
- _remote_blocks_processed flag persists across reschedules
P-node flags: do_remote_prefill=False (prefill locally),
do_remote_decode=True (don't decode locally, send KV to D).
P pulls KV from D when remote_block_ids is not None and
external tokens > 0.
"""
import copy
import time
from unittest.mock import patch
import pytest
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.connector import (
NixlConnector,
NixlConnectorMetadata,
)
from vllm.forward_context import ForwardContext
from vllm.v1.outputs import (
EMPTY_MODEL_RUNNER_OUTPUT,
KVConnectorOutput,
)
from vllm.v1.request import RequestStatus
from .test_nixl_connector import FakeNixlConnectorWorker, FakeNixlWrapper
from .utils import (
assert_scheduler_empty,
create_model_runner_output,
create_request,
create_scheduler,
create_vllm_config,
make_kv_cache_config,
)
pytestmark = pytest.mark.cpu_test
# Common extra config for all bi-directional KV transfer tests.
BIDIR_KV_EXTRA_CONFIG = {"bidirectional_kv_xfer": True, "kv_recompute_threshold": 0}
# Helpers
def _make_p_node_turn2_request(
request_id, block_size, num_tokens, num_remote_blocks=3, remote_num_tokens=None
):
"""Create a P-node Turn 2 request with remote_block_ids from D."""
request = create_request(
request_id=request_id,
block_size=block_size,
num_tokens=num_tokens,
do_remote_decode=True,
)
if remote_num_tokens is None:
remote_num_tokens = num_remote_blocks * block_size
request.kv_transfer_params["remote_block_ids"] = [list(range(num_remote_blocks))]
request.kv_transfer_params["remote_num_tokens"] = remote_num_tokens
request.kv_transfer_params["remote_engine_id"] = "decode-engine"
request.kv_transfer_params["remote_request_id"] = f"decode-{request_id}"
request.kv_transfer_params["remote_host"] = "decode-host"
request.kv_transfer_params["remote_port"] = 5678
return request
def _make_connector_with_fake_worker(
hand_shake_latency=0, cycles_before_done=0, do_handshake=True
):
"""Create a NixlConnector with FakeNixlConnectorWorker."""
vllm_config = create_vllm_config()
kv_cache_config = make_kv_cache_config(block_size=16, num_blocks=2)
connector = NixlConnector(vllm_config, KVConnectorRole.WORKER, kv_cache_config)
connector.connector_worker = FakeNixlConnectorWorker(
vllm_config,
connector.engine_id,
hand_shake_latency=hand_shake_latency,
kv_cache_config=kv_cache_config,
)
worker = connector.connector_worker
assert isinstance(worker.nixl_wrapper, FakeNixlWrapper)
worker.nixl_wrapper.set_cycles_before_xfer_done(cycles_before_done)
worker.kv_cache_layout = "HND"
if do_handshake:
remote_agents = worker._nixl_handshake(
host="localhost",
port=1234,
remote_tp_size=1,
expected_engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
)
worker._remote_agents[FakeNixlConnectorWorker.REMOTE_ENGINE_ID] = remote_agents
return connector, worker
def _make_p_node_recv_metadata(request_id, local_blocks, remote_blocks):
"""Build NixlConnectorMetadata for P-node pulling KV from D."""
meta = NixlConnectorMetadata()
meta.add_new_req_to_recv(
request_id=request_id,
local_block_ids=(local_blocks,),
kv_transfer_params={
"do_remote_prefill": False,
"do_remote_decode": True,
"remote_block_ids": (remote_blocks,),
"remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
"remote_request_id": f"decode-{request_id}",
"remote_host": "localhost",
"remote_port": 1234,
"remote_tp_size": 1,
},
)
return meta
def _do_load_kv(connector, metadata):
"""Bind metadata and call start_load_kv."""
connector.bind_connector_metadata(metadata)
ctx = ForwardContext(no_compile_layers={}, attn_metadata={}, slot_mapping={})
connector.start_load_kv(ctx)
# 1. P-node scheduler lifecycle tests
def test_multiturn_lifecycle():
"""Full two-turn lifecycle on the P node:
Turn 1: P prefills locally (do_remote_prefill=False), sends KV to D
(do_remote_decode=True). Finishes LENGTH_CAPPED with remote_block_ids.
Turn 2: P receives remote_block_ids from D. P pulls KV from D because
remote_block_ids is not None and external tokens > 0. Computes only
new tokens, finishes LENGTH_CAPPED."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
t1 = create_request(
request_id=100, block_size=BS, num_tokens=int(BS * 2.5), do_remote_decode=True
)
scheduler.add_request(t1)
t1_id = t1.request_id
so = scheduler.schedule()
mro = create_model_runner_output(reqs=[t1])
eco = scheduler.update_from_output(so, mro)
assert t1.status == RequestStatus.FINISHED_LENGTH_CAPPED
kv = eco[0].outputs[0].kv_transfer_params
assert kv and sum(len(g) for g in kv["remote_block_ids"]) > 0
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
t2 = _make_p_node_turn2_request(200, BS, int(BS * 2.5))
scheduler.add_request(t2)
t2_id = t2.request_id
so = scheduler.schedule()
assert t2.status == RequestStatus.WAITING_FOR_REMOTE_KVS
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_recving={t2_id})
scheduler.update_from_output(so, mro)
so = scheduler.schedule()
mro = create_model_runner_output(reqs=[t2])
scheduler.update_from_output(so, mro)
assert t2.status == RequestStatus.FINISHED_LENGTH_CAPPED
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={t1_id, t2_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_first_turn_no_remote_blocks():
"""First turn: P has no remote_block_ids from D yet.
Standard local prefill, returns kv_transfer_params for future turns."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=3, block_size=BS, num_tokens=int(BS * 2.5), do_remote_decode=True
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
assert req.status != RequestStatus.WAITING_FOR_REMOTE_KVS
mro = create_model_runner_output(reqs=[req])
eco = scheduler.update_from_output(so, mro)
assert req.status == RequestStatus.FINISHED_LENGTH_CAPPED
assert eco[0].outputs[0].kv_transfer_params is not None
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_abort_p_side_during_send():
"""P-side do_remote_decode=True: blocks held until finished_sending."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=42, block_size=BS, num_tokens=int(BS * 2.5), do_remote_decode=True
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
assert req_id in scheduler.requests
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
assert req_id in scheduler.requests
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_abort_p_side_non_length_capped():
"""P-side abort with non-LENGTH_CAPPED → immediate block free."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=44, block_size=BS, num_tokens=int(BS * 2.5), do_remote_decode=True
)
req.sampling_params.max_tokens = 100
req.max_tokens = 100
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
scheduler.finish_requests([req_id], RequestStatus.FINISHED_ABORTED)
conn = scheduler.connector.connector_scheduler
assert req_id in conn._reqs_not_processed
assert req_id not in scheduler.requests
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
assert_scheduler_empty(scheduler)
def test_remote_blocks_exceed_prompt_tokens():
"""D provides more remote tokens than P's prompt needs.
P caps external tokens to prompt length."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
NUM_TOKENS = int(BS * 2.5)
req = _make_p_node_turn2_request(
300, BS, NUM_TOKENS, num_remote_blocks=5, remote_num_tokens=5 * BS
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
assert req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
assert req.num_computed_tokens == NUM_TOKENS
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_recving={req_id})
scheduler.update_from_output(so, mro)
so = scheduler.schedule()
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
assert req.status == RequestStatus.FINISHED_LENGTH_CAPPED
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_p_node_pulls_partial_last_block_from_d():
"""D sends remote_block_ids with partially filled last block.
remote_num_tokens < len(remote_block_ids) * block_size.
P pulls only remote_num_tokens worth of external tokens."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
num_remote_blocks = 3
remote_num_tokens = int(BS * 2.5)
assert remote_num_tokens < num_remote_blocks * BS
NUM_TOKENS = int(BS * 3.5)
req = _make_p_node_turn2_request(
400,
BS,
NUM_TOKENS,
num_remote_blocks=num_remote_blocks,
remote_num_tokens=remote_num_tokens,
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
assert req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_recving={req_id})
scheduler.update_from_output(so, mro)
so = scheduler.schedule()
assert len(scheduler.running) == 1
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
assert req.status == RequestStatus.FINISHED_LENGTH_CAPPED
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
# 2. P-node metadata tests
def test_add_new_req_to_recv_populates_remote_meta():
"""add_new_req_to_recv correctly populates RemoteMeta for P-node
bi-directional KV pull from D."""
meta = NixlConnectorMetadata()
kv_params = {
"remote_block_ids": [[0, 1, 2]],
"remote_engine_id": "decode-engine",
"remote_request_id": "decode-req-123",
"remote_host": "decode-host",
"remote_port": 5678,
}
local_block_ids = ([10, 11, 12],)
meta.add_new_req_to_recv(
request_id="test-req",
local_block_ids=local_block_ids,
kv_transfer_params=kv_params,
)
assert "test-req" in meta.reqs_to_recv
rm = meta.reqs_to_recv["test-req"]
assert rm.remote is not None
assert rm.remote.block_ids == kv_params["remote_block_ids"]
assert rm.remote.engine_id == "decode-engine"
assert rm.remote.request_id == "decode-req-123"
assert rm.remote.host == "decode-host"
assert rm.remote.port == 5678
assert rm.local_block_ids == local_block_ids
def test_build_connector_meta_recv_entries():
"""P-node scheduler: do_remote_decode=True + remote_block_ids →
_reqs_need_recv populated, build_connector_meta produces reqs_to_recv."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = _make_p_node_turn2_request(1, BS, int(BS * 2.5))
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
assert req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
meta = so.kv_connector_metadata
assert isinstance(meta, NixlConnectorMetadata)
assert req_id in meta.reqs_to_recv
rm = meta.reqs_to_recv[req_id]
assert rm.remote is not None
assert rm.remote.engine_id == "decode-engine"
def test_build_connector_meta_clears_reqs_need_recv():
"""After build_connector_meta, _reqs_need_recv is cleared."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = _make_p_node_turn2_request(2, BS, int(BS * 2.5))
scheduler.add_request(req)
conn = scheduler.connector.connector_scheduler
scheduler.schedule()
assert len(conn._reqs_need_recv) == 0
def test_build_connector_meta_multiple_requests():
"""Multiple P-node requests all included in reqs_to_recv."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
reqs = [_make_p_node_turn2_request(10 + i, BS, int(BS * 2.5)) for i in range(3)]
for r in reqs:
scheduler.add_request(r)
so = scheduler.schedule()
meta = so.kv_connector_metadata
assert isinstance(meta, NixlConnectorMetadata)
assert len(meta.reqs_to_recv) == 3
for r in reqs:
assert r.request_id in meta.reqs_to_recv
# 3. P-node worker tests (FakeNixlWrapper)
@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker.NixlWrapper",
FakeNixlWrapper,
)
def test_p_node_pull_kv_from_d(dist_init):
"""P node pulls KV from D via start_load_kv with reqs_to_recv."""
connector, worker = _make_connector_with_fake_worker()
meta = _make_p_node_recv_metadata("req-p1", [10, 11, 12], [20, 21, 22])
_do_load_kv(connector, meta)
assert "req-p1" in worker._recving_metadata
_, done_recving = connector.get_finished(finished_req_ids=set())
assert "req-p1" in done_recving
@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker.NixlWrapper",
FakeNixlWrapper,
)
def test_p_node_pull_then_send_kv(dist_init):
"""Full P-node bi-directional: pull KV from D → prefill →
send KV back to D via notification."""
connector, worker = _make_connector_with_fake_worker()
meta = _make_p_node_recv_metadata("req-p2", [10, 11], [20, 21])
_do_load_kv(connector, meta)
_, done_recving = connector.get_finished(finished_req_ids=set())
assert "req-p2" in done_recving
worker._reqs_to_send["req-p2"] = time.perf_counter() + 60
worker._reqs_to_process.add("req-p2")
notif = f"req-p2:{worker.world_size}".encode()
orig = worker.nixl_wrapper.get_new_notifs
worker.nixl_wrapper.get_new_notifs = lambda: {"agent": [notif]}
done_sending, _ = connector.get_finished(finished_req_ids=set())
assert "req-p2" in done_sending
worker.nixl_wrapper.get_new_notifs = orig
@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker.NixlWrapper",
FakeNixlWrapper,
)
def test_p_node_deferred_pull_on_no_handshake(dist_init):
"""P defers KV pull when no prior handshake exists."""
connector, worker = _make_connector_with_fake_worker(
hand_shake_latency=0, do_handshake=False
)
meta = _make_p_node_recv_metadata("req-p3", [10, 11], [20, 21])
_do_load_kv(connector, meta)
assert "req-p3" in worker._recving_metadata
timeout = 3.0
start = time.perf_counter()
while time.perf_counter() - start < timeout:
connector.bind_connector_metadata(NixlConnectorMetadata())
ctx = ForwardContext(no_compile_layers={}, attn_metadata={}, slot_mapping={})
connector.start_load_kv(ctx)
_, done = connector.get_finished(finished_req_ids=set())
if "req-p3" in done:
return
time.sleep(0.2)
raise AssertionError("Transfer did not complete after async handshake")
# 4. D-node request_finished returns kv_transfer_params (new behavior)
def test_d_node_request_finished_returns_kv_params():
"""D-node request_finished returns kv_transfer_params with
do_remote_decode=True, remote_block_ids, remote_num_tokens
for P to pull. These params go directly to P node."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=1, block_size=BS, num_tokens=int(BS * 2.5), do_remote_prefill=True
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
scheduler.update_from_output(
so, create_model_runner_output(reqs=[], finished_recving={req_id})
)
so = scheduler.schedule()
eco = scheduler.update_from_output(
so, create_model_runner_output(reqs=[req], use_eos=True)
)
assert req.status == RequestStatus.FINISHED_STOPPED
kv = eco[0].outputs[0].kv_transfer_params
assert kv is not None
assert kv["do_remote_decode"] is True
assert kv["do_remote_prefill"] is False
assert "remote_block_ids" in kv
assert "remote_num_tokens" in kv
assert kv["remote_num_tokens"] > 0
def test_d_node_request_finished_delays_block_free():
"""D-node holds blocks (delay_free=True) until P reads them."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=2, block_size=BS, num_tokens=int(BS * 2.5), do_remote_prefill=True
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
scheduler.update_from_output(
so, create_model_runner_output(reqs=[], finished_recving={req_id})
)
so = scheduler.schedule()
scheduler.update_from_output(
so, create_model_runner_output(reqs=[req], use_eos=True)
)
assert req_id in scheduler.requests
conn = scheduler.connector.connector_scheduler
assert req_id in conn._reqs_need_send
def test_d_node_request_finished_remote_num_tokens():
"""D-node kv_transfer_params includes correct remote_num_tokens."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=3, block_size=BS, num_tokens=int(BS * 2.5), do_remote_prefill=True
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
scheduler.update_from_output(
so, create_model_runner_output(reqs=[], finished_recving={req_id})
)
so = scheduler.schedule()
eco = scheduler.update_from_output(
so, create_model_runner_output(reqs=[req], use_eos=True)
)
kv = eco[0].outputs[0].kv_transfer_params
assert kv["remote_num_tokens"] > 0
assert sum(len(g) for g in kv["remote_block_ids"]) > 0
def test_d_node_partial_last_block_remote_num_tokens():
"""D-node: remote_num_tokens < len(remote_block_ids) * block_size
when last block is partially filled."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=5, block_size=BS, num_tokens=int(BS * 2.5), do_remote_prefill=True
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
scheduler.update_from_output(
so, create_model_runner_output(reqs=[], finished_recving={req_id})
)
so = scheduler.schedule()
eco = scheduler.update_from_output(
so, create_model_runner_output(reqs=[req], use_eos=True)
)
kv = eco[0].outputs[0].kv_transfer_params
total_blocks = sum(len(g) for g in kv["remote_block_ids"])
assert total_blocks == 3
assert kv["remote_num_tokens"] < total_blocks * BS
assert kv["remote_num_tokens"] > 0
# 5. Edge case tests
def test_no_double_read_blocks_after_reschedule():
"""Edge case 1: update_state_after_alloc called twice for the same
bidirectional request (once on initial schedule, once after
WAITING_FOR_REMOTE_KVS reschedule). The _remote_blocks_processed
flag must prevent the request from being added to _reqs_need_recv
twice, which would cause P to read D's blocks twice."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = _make_p_node_turn2_request(500, BS, int(BS * 2.5))
scheduler.add_request(req)
req_id = req.request_id
conn = scheduler.connector.connector_scheduler
# First schedule: request enters WAITING_FOR_REMOTE_KVS,
# _reqs_need_recv populated then cleared by build_connector_meta.
so = scheduler.schedule()
assert req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
meta = so.kv_connector_metadata
assert isinstance(meta, NixlConnectorMetadata)
assert req_id in meta.reqs_to_recv
# _reqs_need_recv should be cleared after build_connector_meta
assert len(conn._reqs_need_recv) == 0
# Simulate recv completion
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_recving={req_id})
scheduler.update_from_output(so, mro)
# Second schedule after recv: update_state_after_alloc called again.
# The _remote_blocks_processed flag should prevent re-entry.
so = scheduler.schedule()
meta2 = so.kv_connector_metadata
assert isinstance(meta2, NixlConnectorMetadata)
# Must NOT be in reqs_to_recv again
assert req_id not in meta2.reqs_to_recv
# Clean up
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
assert req.status == RequestStatus.FINISHED_LENGTH_CAPPED
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_remote_num_tokens_bounded_by_blocks():
"""Edge case 2: D-node request_finished must return
remote_num_tokens <= len(remote_block_ids) * block_size.
request.num_tokens includes the last sampled token which has no KV
in the cache, so remote_num_tokens must use num_computed_tokens."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=501,
block_size=BS,
num_tokens=int(BS * 2.5),
do_remote_prefill=True,
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
scheduler.update_from_output(
so, create_model_runner_output(reqs=[], finished_recving={req_id})
)
so = scheduler.schedule()
eco = scheduler.update_from_output(
so, create_model_runner_output(reqs=[req], use_eos=True)
)
kv = eco[0].outputs[0].kv_transfer_params
assert kv is not None
total_blocks = sum(len(g) for g in kv["remote_block_ids"])
max_tokens_in_blocks = total_blocks * BS
assert kv["remote_num_tokens"] <= max_tokens_in_blocks, (
f"remote_num_tokens ({kv['remote_num_tokens']}) exceeds "
f"block capacity ({max_tokens_in_blocks})"
)
assert kv["remote_num_tokens"] > 0
def test_kv_recompute_threshold_skips_small_transfer():
"""Edge case 3: When remote tokens are below kv_recompute_threshold,
P should skip the remote pull and compute locally instead of
entering WAITING_FOR_REMOTE_KVS."""
threshold = 256
vllm_config = create_vllm_config(
kv_connector_extra_config={
"bidirectional_kv_xfer": True,
"kv_recompute_threshold": threshold,
},
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
# Create request where remote tokens (48) < threshold (256)
req = _make_p_node_turn2_request(
502,
BS,
int(BS * 2.5),
num_remote_blocks=3,
remote_num_tokens=3 * BS,
)
scheduler.add_request(req)
so = scheduler.schedule()
# Should NOT enter WAITING_FOR_REMOTE_KVS — threshold not met
assert req.status != RequestStatus.WAITING_FOR_REMOTE_KVS
assert req.status == RequestStatus.RUNNING
# Clean up
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
assert req.status == RequestStatus.FINISHED_LENGTH_CAPPED
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req.request_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_p_node_finished_holds_blocks_for_d():
"""Edge case 4: P-node finishes with FINISHED_LENGTH_CAPPED and
do_remote_decode=True. P must hold blocks (delay_free=True) and
return kv_transfer_params with do_remote_prefill=True so D can
read P's blocks."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=503,
block_size=BS,
num_tokens=int(BS * 2.5),
do_remote_decode=True,
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
mro = create_model_runner_output(reqs=[req])
eco = scheduler.update_from_output(so, mro)
assert req.status == RequestStatus.FINISHED_LENGTH_CAPPED
kv = eco[0].outputs[0].kv_transfer_params
assert kv is not None
# P-node finished: should tell D to pull (do_remote_prefill=True)
assert kv["do_remote_prefill"] is True
assert kv["do_remote_decode"] is False
assert "remote_block_ids" in kv
assert sum(len(g) for g in kv["remote_block_ids"]) > 0
# Blocks should be held (request still tracked)
assert req_id in scheduler.requests
# Clean up: simulate D reading and notifying
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_cache_miss_first_turn_no_remote_pull():
"""Edge case 5: First turn with do_remote_decode=True but no
remote_block_ids (cache MISS). P should prefill locally with
num_external_tokens=0 and not enter WAITING_FOR_REMOTE_KVS."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = create_request(
request_id=504,
block_size=BS,
num_tokens=int(BS * 2.5),
do_remote_decode=True,
)
# No remote_block_ids set — this is a cache MISS
assert req.kv_transfer_params.get("remote_block_ids") is None
scheduler.add_request(req)
so = scheduler.schedule()
# Should NOT wait for remote KVs
assert req.status != RequestStatus.WAITING_FOR_REMOTE_KVS
assert req.status == RequestStatus.RUNNING
# Clean up
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req.request_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_partial_remote_tokens_less_than_prompt():
"""Edge case 6: D's remote_num_tokens covers only part of P's
prompt. P should pull remote_num_tokens worth of external tokens
and compute the rest locally."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
NUM_TOKENS = int(BS * 4.5) # 72 tokens
# D provides only 2 blocks (32 tokens) out of 72
req = _make_p_node_turn2_request(
505,
BS,
NUM_TOKENS,
num_remote_blocks=2,
remote_num_tokens=2 * BS,
)
scheduler.add_request(req)
req_id = req.request_id
so = scheduler.schedule()
assert req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
# num_computed_tokens should reflect the external tokens pulled
# (capped to remote_num_tokens, not full prompt)
assert req.num_computed_tokens < NUM_TOKENS
# Complete the transfer and finish
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_recving={req_id})
scheduler.update_from_output(so, mro)
so = scheduler.schedule()
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
assert req.status == RequestStatus.FINISHED_LENGTH_CAPPED
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
def test_remote_blocks_processed_flag_persists():
"""Edge case 7: After recv completes and request is rescheduled,
the _remote_blocks_processed flag in kv_transfer_params prevents
the bidirectional path from re-entering _reqs_need_recv."""
vllm_config = create_vllm_config(
kv_connector_extra_config=BIDIR_KV_EXTRA_CONFIG,
)
scheduler = create_scheduler(vllm_config)
BS = vllm_config.cache_config.block_size
req = _make_p_node_turn2_request(506, BS, int(BS * 2.5))
scheduler.add_request(req)
req_id = req.request_id
conn = scheduler.connector.connector_scheduler
# First schedule → WAITING_FOR_REMOTE_KVS
so = scheduler.schedule()
assert req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
# Recv completes
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_recving={req_id})
scheduler.update_from_output(so, mro)
# Verify the flag is set
assert req.kv_transfer_params.get("_remote_blocks_processed") is True
# Next schedule: update_state_after_alloc is called again.
# _reqs_need_recv must NOT contain this request.
so = scheduler.schedule()
assert req_id not in conn._reqs_need_recv
meta = so.kv_connector_metadata
assert isinstance(meta, NixlConnectorMetadata)
assert req_id not in meta.reqs_to_recv
# Clean up
mro = create_model_runner_output(reqs=[req])
scheduler.update_from_output(so, mro)
so = scheduler.schedule()
scheduler.update_from_output(so, EMPTY_MODEL_RUNNER_OUTPUT)
so = scheduler.schedule()
mro = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
mro.kv_connector_output = KVConnectorOutput(finished_sending={req_id})
scheduler.update_from_output(so, mro)
assert_scheduler_empty(scheduler)
@@ -8,13 +8,18 @@ from typing import Any
from unittest.mock import MagicMock
import pytest
import torch
from tests.v1.kv_connector.unit.utils import create_vllm_config
from vllm import LLM, SamplingParams
from vllm.config import KVTransferConfig
from vllm.distributed.kv_transfer.kv_connector.factory import KVConnectorFactory
from vllm.distributed.kv_transfer.kv_connector.v1 import KVConnectorRole
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorBase_V1
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorBase_V1,
SupportsHMA,
supports_hma,
)
from vllm.distributed.kv_transfer.kv_connector.v1.metrics import KVConnectorStats
from vllm.distributed.kv_transfer.kv_connector.v1.multi_connector import (
MultiConnector,
@@ -83,8 +88,43 @@ class MockConnector(KVConnectorBase_V1):
pass
# Register the mock connector
class MockHMAConnector(KVConnectorBase_V1, SupportsHMA):
"""Mock connector that supports HMA for testing."""
def __new__(cls, *args, **kwargs):
mock = MagicMock(spec_set=cls)
return mock
def start_load_kv(self, forward_context, **kwargs):
pass
def wait_for_layer_load(self, layer_name):
pass
def save_kv_layer(self, layer_name, kv_layer, attn_metadata, **kwargs):
pass
def wait_for_save(self):
pass
def build_connector_meta(self, scheduler_output):
return None
def get_num_new_matched_tokens(self, request, num_computed_tokens):
return (0, False)
def update_state_after_alloc(self, request, blocks, num_tokens) -> None:
pass
def request_finished_all_groups(self, request, block_ids):
return (False, None)
# Register mock connectors
KVConnectorFactory.register_connector("MockConnector", __name__, MockConnector.__name__)
KVConnectorFactory.register_connector(
"MockHMAConnector", __name__, MockHMAConnector.__name__
)
@pytest.fixture
@@ -920,3 +960,133 @@ def test_multi_connector_worker_metadata(mc):
mc.update_connector_output(kv_connector_output)
assert_update_connector_output_called(mc)
assert kv_connector_output.kv_connector_worker_meta == mc_worker_meta_01a_01b
def _make_multi_connector(connector_names: list[str]) -> MultiConnector:
"""Build a MultiConnector wrapping the given registered connectors."""
vllm_config = create_vllm_config()
connectors = [
{
"kv_connector": name,
"kv_role": "kv_both",
"kv_connector_module_path": "tests.v1.kv_connector.unit.test_multi_connector", # noqa: E501
}
for name in connector_names
]
vllm_config.kv_transfer_config = KVTransferConfig(
kv_connector="MultiConnector",
kv_role="kv_both",
kv_connector_extra_config={"connectors": connectors},
)
kv_cache_config = KVCacheConfig(
num_blocks=0,
kv_cache_tensors=[],
kv_cache_groups=[],
)
return MultiConnector(
vllm_config=vllm_config,
role=KVConnectorRole.WORKER,
kv_cache_config=kv_cache_config,
)
def test_multi_connector_hma_opt_in():
"""
MultiConnector currently assumes HMA is opt-in: it needs
--no-disable-hybrid-kv-cache-manager to be enabled.
At runtime, _all_support_hma is True only when every sub-connector
implements SupportsHMA. Test all combinations of HMA / non-HMA
sub-connectors.
"""
assert supports_hma(MultiConnector)
# -- All non-HMA connectors => _all_support_hma is False --
mc_none = _make_multi_connector(["MockConnector", "MockConnector"])
assert not supports_hma(mc_none._connectors[0])
assert not supports_hma(mc_none._connectors[1])
assert mc_none._all_support_hma is False
# -- All HMA connectors => _all_support_hma is True --
mc_all = _make_multi_connector(["MockHMAConnector", "MockHMAConnector"])
assert supports_hma(mc_all._connectors[0])
assert supports_hma(mc_all._connectors[1])
assert mc_all._all_support_hma is True
# -- Mixed: first HMA, second non-HMA => _all_support_hma is False --
mc_mixed1 = _make_multi_connector(["MockHMAConnector", "MockConnector"])
assert supports_hma(mc_mixed1._connectors[0])
assert not supports_hma(mc_mixed1._connectors[1])
assert mc_mixed1._all_support_hma is False
# -- Mixed: first non-HMA, second HMA => _all_support_hma is False --
mc_mixed2 = _make_multi_connector(["MockConnector", "MockHMAConnector"])
assert not supports_hma(mc_mixed2._connectors[0])
assert supports_hma(mc_mixed2._connectors[1])
assert mc_mixed2._all_support_hma is False
@pytest.mark.skipif(
not torch.cuda.is_available(), reason="Requires GPU to instantiate LLM"
)
def test_multi_connector_mixed_hma_disables_hybrid_kv_cache(monkeypatch):
"""
When MultiConnector wraps a mix of HMA (NixlConnector) and non-HMA
(MockConnector) sub-connectors, verify that:
1. The scheduler's MultiConnector has _all_support_hma == False.
2. vLLM auto-disables the hybrid KV cache manager (no preference expressed by user)
"""
from unittest.mock import patch
from tests.v1.kv_connector.unit.test_nixl_connector import FakeNixlWrapper
monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
kv_transfer_config = KVTransferConfig(
kv_connector="MultiConnector",
kv_role="kv_both",
kv_connector_extra_config={
"connectors": [
{
"kv_connector": "NixlConnector",
"kv_role": "kv_both",
},
{
"kv_connector": "MockConnector",
"kv_role": "kv_both",
"kv_connector_module_path": (
"tests.v1.kv_connector.unit.test_multi_connector"
),
},
],
},
)
with patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker.NixlWrapper",
FakeNixlWrapper,
):
llm = LLM(
model="Qwen/Qwen3-0.6B",
enforce_eager=True,
gpu_memory_utilization=0.3,
max_model_len=128,
max_num_seqs=1,
max_num_batched_tokens=128,
kv_transfer_config=kv_transfer_config,
)
try:
# HMA should be auto-disabled when user has not expressed a preference.
assert (
llm.llm_engine.vllm_config.scheduler_config.disable_hybrid_kv_cache_manager
is True
)
# The scheduler-side MultiConnector should detect the mixed
# HMA support among its sub-connectors.
scheduler = llm.llm_engine.engine_core.engine_core.scheduler
mc = scheduler.connector
assert isinstance(mc, MultiConnector)
assert mc._all_support_hma is False
finally:
llm.llm_engine.engine_core.shutdown()
@@ -9,14 +9,15 @@ import torch
from vllm.platforms import current_platform
from vllm.utils.torch_utils import set_random_seed
from vllm.v1.kv_offload.cpu.shared_offload_region import SharedOffloadRegion
from vllm.v1.kv_offload.mediums import CPULoadStoreSpec, GPULoadStoreSpec
from vllm.v1.kv_offload.spec import (
from vllm.v1.kv_offload.base import (
CanonicalKVCacheRef,
CanonicalKVCaches,
CanonicalKVCacheTensor,
GPULoadStoreSpec,
)
from vllm.v1.kv_offload.worker.cpu_gpu import CpuGpuOffloadingHandlers
from vllm.v1.kv_offload.cpu.common import CPULoadStoreSpec
from vllm.v1.kv_offload.cpu.gpu_worker import CpuGpuOffloadingHandlers
from vllm.v1.kv_offload.cpu.shared_offload_region import SharedOffloadRegion
NUM_GPU_BLOCKS = [64]
NUM_CPU_BLOCKS = [256]
@@ -6,7 +6,7 @@ from dataclasses import dataclass
import numpy as np
import pytest
from vllm.v1.kv_offload.abstract import (
from vllm.v1.kv_offload.base import (
LoadStoreSpec,
OffloadingEvent,
OffloadKey,
@@ -14,9 +14,9 @@ from vllm.v1.kv_offload.abstract import (
ReqContext,
make_offload_key,
)
from vllm.v1.kv_offload.cpu.common import CPULoadStoreSpec
from vllm.v1.kv_offload.cpu.manager import CPUOffloadingManager
from vllm.v1.kv_offload.cpu.policies.arc import ARCCachePolicy
from vllm.v1.kv_offload.mediums import CPULoadStoreSpec
from vllm.v1.kv_offload.reuse_manager import FilterReusedOffloadingManager
+1 -1
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from vllm.v1.kv_offload.abstract import LoadStoreSpec
from vllm.v1.kv_offload.base import LoadStoreSpec
from vllm.v1.kv_offload.worker.worker import (
OffloadingHandler,
OffloadingWorker,
+9 -2
View File
@@ -778,7 +778,7 @@ def _rocm_aiter_per_token_quant_impl(
assert quant_dtype in [torch.int8, FP8_DTYPE]
out_shape = x.shape
out = torch.empty(x.shape, dtype=FP8_DTYPE, device=x.device)
out = torch.empty(x.shape, dtype=quant_dtype, device=x.device)
if scale is None:
scale = torch.empty((*out_shape[:-1], 1), dtype=torch.float32, device=x.device)
dynamic_per_token_scaled_quant(
@@ -798,7 +798,7 @@ def _rocm_aiter_per_token_quant_fake(
) -> tuple[torch.Tensor, torch.Tensor]:
out_shape = x.shape
return (
torch.empty(x.shape, dtype=FP8_DTYPE, device=x.device),
torch.empty(x.shape, dtype=quant_dtype, device=x.device),
torch.empty((*out_shape[:-1], 1), dtype=torch.float32, device=x.device),
)
@@ -1355,6 +1355,13 @@ class rocm_aiter_ops:
def is_triton_gemm_enabled(cls) -> bool:
return cls._AITER_ENABLED and cls._TRITON_UNQUANT_GEMM
@classmethod
@if_aiter_supported
def is_tgemm_enabled(cls) -> bool:
from vllm.platforms.rocm import on_gfx950
return cls.is_linear_enabled() and on_gfx950()
@staticmethod
@if_aiter_supported
def register_ops_once() -> None:
-43
View File
@@ -203,47 +203,6 @@ def is_compile_cache_enabled(
)
def _patch_standalone_compile_atomic_save() -> None:
"""Backport of pytorch/pytorch#162432 for torch < 2.10.0.
Patches CompiledArtifact.save() to use write_atomic for binary format,
preventing corrupt cache files when multiple processes compile
concurrently.
"""
from torch._inductor.codecache import write_atomic
from torch._inductor.standalone_compile import CompiledArtifact as cls
if getattr(cls.save, "_vllm_patched", False):
return
original_save = cls.save
def _save(
self: Any, *, path: str, format: Literal["binary", "unpacked"] = "binary"
) -> None:
if format != "binary":
return original_save(self, path=path, format=format)
from torch._dynamo.utils import dynamo_timed
from torch._inductor.codecache import torch_key
from torch.utils._appending_byte_serializer import BytesWriter
with dynamo_timed("CompiledArtifact.save"):
assert self._artifacts is not None
artifact_bytes, cache_info = self._artifacts
assert len(cache_info.aot_autograd_artifacts) == 1, cache_info
key = cache_info.aot_autograd_artifacts[0]
assert not os.path.isdir(path)
writer = BytesWriter()
writer.write_bytes(torch_key())
writer.write_str(key)
writer.write_bytes(artifact_bytes)
write_atomic(path, writer.to_bytes())
_save._vllm_patched = True # type: ignore[attr-defined]
cls.save = _save # type: ignore[assignment]
logger.debug("Patched %s.save for atomic writes (torch < 2.10)", cls.__name__)
class InductorStandaloneAdaptor(CompilerInterface):
"""
The adaptor for the Inductor compiler.
@@ -257,8 +216,6 @@ class InductorStandaloneAdaptor(CompilerInterface):
name = "inductor_standalone"
def __init__(self, save_format: Literal["binary", "unpacked"]) -> None:
if not is_torch_equal_or_newer("2.10.0"):
_patch_standalone_compile_atomic_save()
self.save_format = save_format
def compute_hash(self, vllm_config: VllmConfig) -> str:
+7 -2
View File
@@ -290,9 +290,14 @@ class CUDAGraphWrapper:
# across layers will make the cudagraph capture very slow.
# therefore, we only run gc for the first graph,
# and disable gc for the rest of the graphs.
stack.enter_context(patch("gc.collect", lambda: None))
stack.enter_context(
patch("torch.accelerator.empty_cache", lambda: None)
patch("gc.collect", lambda *args, **kwargs: None)
)
stack.enter_context(
patch(
"torch.accelerator.empty_cache",
lambda *args, **kwargs: None,
)
)
if self.graph_pool is not None:
+73 -33
View File
@@ -32,6 +32,9 @@ from vllm.utils.torch_utils import is_torch_equal_or_newer
from .monitor import monitor_profiling_run, monitor_torch_compile
# shape_id parameter was added to mark_unbacked in PyTorch 2.11.0
_SUPPORTS_SHAPE_ID = is_torch_equal_or_newer("2.11.0")
if TYPE_CHECKING:
# Only added on nightly/2.10 so wrap
try:
@@ -89,7 +92,7 @@ def support_torch_compile(
@overload
def support_torch_compile(
*,
dynamic_arg_dims: dict[str, int | list[int]] | None,
dynamic_arg_dims: dict[str, int | list[int] | dict[int, str]] | None,
) -> Callable[[type[_T]], type[_T]]: ...
@@ -103,7 +106,7 @@ def support_torch_compile(
@overload
def support_torch_compile(
*,
dynamic_arg_dims: dict[str, int | list[int]] | None,
dynamic_arg_dims: dict[str, int | list[int] | dict[int, str]] | None,
mark_unbacked_dims: dict[str, int | list[int]] | None,
) -> Callable[[type[_T]], type[_T]]: ...
@@ -115,11 +118,10 @@ def support_torch_compile(cls: type[_T]) -> type[_T]: ...
def support_torch_compile(
cls: type[_T] | None = None,
*,
dynamic_arg_dims: dict[str, int | list[int]] | None = None,
dynamic_arg_dims: dict[str, int | list[int] | dict[int, str]] | None = None,
mark_unbacked_dims: dict[str, int | list[int]] | None = None,
enable_if: Callable[[VllmConfig], bool] | None = None,
is_encoder: bool = False,
shape_invariants: Callable[..., None] = lambda *args, **kwargs: None,
) -> Callable[[type[_T]], type[_T]] | type[_T]:
"""
A decorator to add support for compiling the forward method of a class.
@@ -141,8 +143,12 @@ def support_torch_compile(
```
`dynamic_arg_dims` is a dictionary that maps argument names to the dynamic
dimensions of the argument. The dynamic dimensions can be either a single
integer or a list of integers.
dimensions of the argument. The value can be:
- int: a single dimension index (e.g., 0)
- list[int]: multiple dimension indices (e.g., [0, 1])
- dict[int, str]: dimension to shape_id mapping for shape relations
(e.g., {0: "b"}). Dimensions with the same shape_id share the same
unbacked symbol.
if `dynamic_arg_dims` is `None`, it is inferred from the type annotation
of the `forward` method, based on the following default rules:
@@ -189,7 +195,7 @@ def support_torch_compile(
torch._check(input_ids.size()[0] == inputs_embeds.size()[0])
This enforces constraints on the symbolic shapes without hardcoding
specific values. It is needed for some models to avoid data dependent
errors.
errors and maximize perf when unbacked shapes are used.
"""
def cls_decorator_helper(cls: type[_T]) -> type[_T]:
@@ -229,13 +235,13 @@ def support_torch_compile(
raise ValueError(
f"Argument {k} not found in the forward method of {cls}"
)
return _support_torch_compile(
cls,
inferred_dynamic_arg_dims,
mark_unbacked_dims,
enable_if,
is_encoder,
shape_invariants,
)
if cls is not None:
@@ -324,15 +330,13 @@ def _try_load_aot_compiled_fn(
def _support_torch_compile(
cls: type[_T],
dynamic_arg_dims: dict[str, int | list[int]],
dynamic_arg_dims: dict[str, int | list[int] | dict[int, str]],
mark_unbacked_dims: dict[str, int | list[int]] | None = None,
enable_if: Callable[[VllmConfig], bool] | None = None,
is_encoder: bool = False,
shape_invariants: Callable[..., None] = lambda *args, **kwargs: None,
) -> type[_T]:
"""
A decorator to add support for compiling the forward method of a class.
"""
"""Internal implementation of support_torch_compile decorator."""
if TorchCompileWithNoGuardsWrapper in cls.__bases__:
# support decorating multiple times
return cls
@@ -392,7 +396,8 @@ def _support_torch_compile(
if self.do_not_compile:
return
self._check_shape_invariants = shape_invariants
self._dynamic_arg_dims = dynamic_arg_dims
self.was_aot_compile_fn_loaded_from_disk = False
compilation_counter.num_models_seen += 1
self.compiled = False
@@ -409,48 +414,83 @@ def _support_torch_compile(
def _mark_dynamic_inputs(
mod: type[_T], ds_type: DynamicShapesType, *args: Any, **kwargs: Any
) -> None:
def mark_dynamic(arg: torch.Tensor, dims: list[int]) -> None:
def mark_dynamic(
arg: torch.Tensor, dim_shape_pairs: list[tuple[int, str | None]]
) -> None:
if ds_type == DynamicShapesType.UNBACKED:
if is_torch_equal_or_newer("2.10.0"):
for dim in dims:
torch._dynamo.decorators.mark_unbacked(
arg, dim, hint_override=arg.size()[dim]
)
for dim, shape_id in dim_shape_pairs:
if shape_id is not None:
if not _SUPPORTS_SHAPE_ID:
raise RuntimeError(
f"shape_id='{shape_id}' requires PyTorch >= 2.11.0"
)
torch._dynamo.decorators.mark_unbacked(
arg,
dim,
hint_override=arg.size()[dim],
shape_id=shape_id,
)
else:
torch._dynamo.decorators.mark_unbacked(
arg,
dim,
hint_override=arg.size()[dim],
)
else:
# For older versions, we can't use hint_override or shape_id
dims = [dim for dim, _ in dim_shape_pairs]
torch._dynamo.decorators.mark_unbacked(arg, dims)
else:
dims = [dim for dim, _ in dim_shape_pairs]
torch._dynamo.mark_dynamic(arg, dims)
sig = inspect.signature(mod.__class__.forward) # type: ignore[attr-defined]
bound_args = sig.bind(mod, *args, **kwargs)
bound_args.apply_defaults()
for k, dims in dynamic_arg_dims.items():
# Normalize dynamic_arg_dims to dict[str, dict[int, str | None]]
normalized_dims: dict[str, dict[int, str | None]] = {}
for k, v in dynamic_arg_dims.items():
if isinstance(v, dict):
normalized_dims[k] = {dim: shape_id for dim, shape_id in v.items()}
elif isinstance(v, int):
normalized_dims[k] = {v: None}
else:
normalized_dims[k] = {d: None for d in v}
for k, dim_to_shape_id in normalized_dims.items():
arg = bound_args.arguments.get(k)
if arg is not None:
dims = [dims] if isinstance(dims, int) else dims
dims = list(dim_to_shape_id.keys())
if isinstance(arg, torch.Tensor):
# In case dims is specified with negative indexing
dims = [arg.ndim + dim if dim < 0 else dim for dim in dims]
mark_dynamic(arg, dims)
dim_shape_pairs = [
(arg.ndim + d if d < 0 else d, dim_to_shape_id.get(d))
for d in dims
]
mark_dynamic(arg, dim_shape_pairs)
elif isinstance(arg, IntermediateTensors):
for tensor in arg.tensors.values():
# In case dims is specified with negative indexing
dims = [tensor.ndim + dim if dim < 0 else dim for dim in dims]
mark_dynamic(tensor, dims)
dim_shape_pairs = [
(tensor.ndim + d if d < 0 else d, dim_to_shape_id.get(d))
for d in dims
]
mark_dynamic(tensor, dim_shape_pairs)
else:
raise ValueError(
"Unsupported dynamic dimensions"
f" {dims} for argument {k} with type {type(arg)}."
f"Unsupported dynamic dimensions {dims} "
f"for argument {k} with type {type(arg)}."
)
if mark_unbacked_dims:
for k, dims in mark_unbacked_dims.items():
for k, dims_val in mark_unbacked_dims.items():
arg = bound_args.arguments.get(k)
if arg is not None:
dims = [dims] if isinstance(dims, int) else dims
dims = [dims_val] if isinstance(dims_val, int) else list(dims_val)
if isinstance(arg, torch.Tensor):
# In case dims is specified with negative indexing
dims = [arg.ndim + dim if dim < 0 else dim for dim in dims]
dims = [arg.ndim + d if d < 0 else d for d in dims]
if is_torch_equal_or_newer("2.10.0"):
for dim in dims:
torch._dynamo.decorators.mark_unbacked(
+3 -19
View File
@@ -53,12 +53,6 @@ class TorchCompileWithNoGuardsWrapper:
since we drop all guards.
"""
def check_invariants_and_forward(self, *args: Any, **kwargs: Any) -> Any:
assert hasattr(self, "_check_shape_invariants")
self._check_shape_invariants(*args, **kwargs)
return self.forward(*args, **kwargs)
def _call_with_optional_nvtx_range(
self, callable_fn: Callable[P, R], *args: P.args, **kwargs: P.kwargs
) -> Any:
@@ -115,6 +109,9 @@ class TorchCompileWithNoGuardsWrapper:
"compilation_config.dynamic_shapes_config.evaluate_guards "
"requires VLLM_USE_BYTECODE_HOOK=0. "
)
assert ds_type != DynamicShapesType.UNBACKED, (
"UNBACKED dynamic shapes do not add guards"
)
options["guard_filter_fn"] = lambda x: [
entry.guard_type == "SHAPE_ENV" for entry in x
@@ -130,19 +127,6 @@ class TorchCompileWithNoGuardsWrapper:
compiled_ptr: Any = self.forward
# Validate that unbacked dynamic shapes require VLLM_USE_BYTECODE_HOOK=False
if ds_type == DynamicShapesType.UNBACKED:
# reason is that bytecode does torch._dynamo.eval_frame.
# remove_from_cache(self.original_code_object()) to force a new
# re-compilation. And if we use
# compiled_ptr = self.check_invariants_and_forward
# it will reset all entries.
assert not envs.VLLM_USE_BYTECODE_HOOK, (
"UNBACKED dynamic shapes requires VLLM_USE_BYTECODE_HOOK=0. "
)
assert not self.evaluate_guards, "UNBACKED dynamic shapes do not add guards"
compiled_ptr = self.check_invariants_and_forward
# Apply the constrain_to_fx_strides patch before first compilation.
# This covers STOCK_TORCH_COMPILE and DYNAMO_ONCE paths. The VLLM
# compile paths call this from their own compile() methods too.
+27
View File
@@ -663,6 +663,33 @@ class ParallelConfig:
aggregated_has_unfinished = bool(tensor.item())
return aggregated_has_unfinished
@staticmethod
def sync_dp_state(
dp_group: ProcessGroup, has_unfinished: bool, pending_pause: bool
) -> tuple[bool, bool]:
"""Combined all-reduce for DP state synchronization.
Uses a single SUM all-reduce on a 2-element tensor:
[0] = 1 if this rank has unfinished work, else 0.
SUM > 0 logical OR across ranks any rank has work.
[1] = 1 if this rank has a pending pause request, else 0.
SUM == dp_size all ranks reached pause consensus.
has_unfinished_global is true if any rank has unfinished work,
or if some ranks are waiting for a pause consensus.
Returns:
(has_unfinished_global, pause_consensus)
"""
tensor = torch.tensor(
[int(has_unfinished), int(pending_pause)], dtype=torch.int32, device="cpu"
)
torch.distributed.all_reduce(tensor, op=ReduceOp.SUM, group=dp_group)
dp_size = dp_group.size()
pause_count = tensor[1].item()
has_unfinished_global = tensor[0].item() > 0 or pause_count % dp_size != 0
return has_unfinished_global, pause_count == dp_size
@staticmethod
def sync_kv_cache_memory_size(dp_group: ProcessGroup, kv_cache_memory: int) -> int:
if kv_cache_memory == -1:
@@ -10,7 +10,6 @@ import vllm.envs as envs
from vllm.distributed import get_dp_group, get_ep_group
from vllm.forward_context import get_forward_context
from vllm.logger import init_logger
from vllm.platforms import current_platform
from vllm.utils.flashinfer import (
has_flashinfer_nvlink_one_sided,
has_flashinfer_nvlink_two_sided,
@@ -225,11 +224,8 @@ class DeepEPHTAll2AllManager(DeepEPAll2AllManagerBase):
num_rdma_bytes=num_rdma_bytes,
low_latency_mode=False,
num_qps_per_rank=num_qps_per_rank,
explicitly_destroy=True,
)
if not current_platform.is_rocm():
kwargs.update(
explicitly_destroy=True,
)
return kwargs
def get_handle(self, kwargs):
@@ -303,13 +299,10 @@ class DeepEPLLAll2AllManager(DeepEPAll2AllManagerBase):
num_rdma_bytes=num_rdma_bytes,
low_latency_mode=True,
num_qps_per_rank=num_qps_per_rank,
allow_nvlink_for_low_latency_mode=True,
allow_mnnvl=envs.VLLM_DEEPEP_LOW_LATENCY_USE_MNNVL,
explicitly_destroy=True,
)
if not current_platform.is_rocm():
kwargs.update(
allow_nvlink_for_low_latency_mode=True,
allow_mnnvl=envs.VLLM_DEEPEP_LOW_LATENCY_USE_MNNVL,
explicitly_destroy=True,
)
return kwargs
def get_handle(self, kwargs):
+204 -199
View File
@@ -11,6 +11,7 @@ from abc import ABC, abstractmethod
from collections.abc import Sequence
from datetime import timedelta
import numpy as np
import torch
from torch.distributed import (
P2POp,
@@ -47,15 +48,25 @@ class EplbCommunicator(ABC):
"""Abstract EPLB communicator for expert weight transfers."""
@abstractmethod
def add_send(self, tensor: torch.Tensor, dst_rank: int) -> None:
def add_send(
self,
tensors: list[torch.Tensor],
dst_rank: int,
expert_id: int,
) -> None:
pass
@abstractmethod
def add_recv(self, tensor: torch.Tensor, src_rank: int) -> None:
def add_recv(
self,
tensors: list[torch.Tensor],
src_rank: int,
expert_id: int,
) -> None:
pass
@abstractmethod
def execute(self) -> None:
def execute(self, old_indices: np.ndarray | None = None) -> None:
pass
@property
@@ -85,27 +96,39 @@ class TorchDistNcclEplbCommunicator(EplbCommunicator):
self._p2p_ops: list[P2POp] = []
self._log_initialized()
def add_send(self, tensor: torch.Tensor, dst_rank: int) -> None:
self._p2p_ops.append(
P2POp(
torch.distributed.isend,
tensor,
dst_rank,
self._ep_group,
def add_send(
self,
tensors: list[torch.Tensor],
dst_rank: int,
expert_id: int, # unused by this backend
) -> None:
for tensor in tensors:
self._p2p_ops.append(
P2POp(
torch.distributed.isend,
tensor,
dst_rank,
self._ep_group,
)
)
)
def add_recv(self, tensor: torch.Tensor, src_rank: int) -> None:
self._p2p_ops.append(
P2POp(
torch.distributed.irecv,
tensor,
src_rank,
self._ep_group,
def add_recv(
self,
tensors: list[torch.Tensor],
src_rank: int,
expert_id: int, # unused by this backend
) -> None:
for tensor in tensors:
self._p2p_ops.append(
P2POp(
torch.distributed.irecv,
tensor,
src_rank,
self._ep_group,
)
)
)
def execute(self) -> None:
def execute(self, old_indices: np.ndarray | None = None) -> None:
if not self._p2p_ops:
return
try:
@@ -130,13 +153,25 @@ class TorchDistGlooStagedEplbCommunicator(EplbCommunicator):
self._ops: list[tuple[str, torch.Tensor, int]] = []
self._log_initialized()
def add_send(self, tensor: torch.Tensor, dst_rank: int) -> None:
self._ops.append(("send", tensor, dst_rank))
def add_send(
self,
tensors: list[torch.Tensor],
dst_rank: int,
expert_id: int, # unused by this backend
) -> None:
for tensor in tensors:
self._ops.append(("send", tensor, dst_rank))
def add_recv(self, tensor: torch.Tensor, src_rank: int) -> None:
self._ops.append(("recv", tensor, src_rank))
def add_recv(
self,
tensors: list[torch.Tensor],
src_rank: int,
expert_id: int, # unused by this backend
) -> None:
for tensor in tensors:
self._ops.append(("recv", tensor, src_rank))
def execute(self) -> None:
def execute(self, old_indices: np.ndarray | None = None) -> None:
if not self._ops:
return
@@ -207,17 +242,17 @@ class NixlEplbCommunicator(EplbCommunicator):
self._cuda_stream = cuda_stream
self._world_size = cpu_group.size()
self._rank = cpu_group.rank()
self._send_tensors: dict[torch.dtype, list[list[torch.Tensor]]] = {}
self._recv_tensors: dict[torch.dtype, list[list[torch.Tensor]]] = {}
self._dtypes: list[torch.dtype] = []
# expert_id -> weight tensors to pack into the send buffer.
self._expert_send_map: dict[int, list[torch.Tensor]] = {}
# src_rank -> expert_id -> weight tensors to unpack after transfer.
self._recv_map: dict[int, dict[int, list[torch.Tensor]]] = {}
self._num_local_experts: int = expert_weights[0].shape[0]
self._device = expert_weights[0].device
for tensor in expert_weights:
assert tensor.device == self._device, (
"All local EPLB tensors are expected to be on the same device: "
f"expected={self._device}, got={tensor.device}"
)
if tensor.dtype not in self._dtypes:
self._dtypes.append(tensor.dtype)
config = (
nixl_agent_config(capture_telemetry=False)
@@ -228,13 +263,12 @@ class NixlEplbCommunicator(EplbCommunicator):
self._nixl_memory_type = "VRAM"
self._registered_desc: object | None = None
self._remote_agents: dict[int, str] = {}
self._remote_send_meta: dict[int, tuple[int, int, int]] = {}
self._remote_send_meta: dict[int, tuple[int, int]] = {}
self._send_buffer: torch.Tensor = torch.empty(0)
self._recv_buffer: torch.Tensor = torch.empty(0)
self._peer_partition_bytes: int = 0
self._dtype_max_bytes: dict[torch.dtype, int] = {}
self._expert_bytes: int = 0
self._cuda_device_id = int(self._device.index or 0)
self._xfer_cache: dict[tuple[int, int, int], tuple[int, int, int]] = {}
self._init_step("buffers", self._init_registered_buffers, expert_weights)
self._init_step("agents", self._init_remote_agents)
self._init_step("send meta", self._exchange_remote_send_meta)
@@ -258,34 +292,33 @@ class NixlEplbCommunicator(EplbCommunicator):
uid = uuid.uuid4().hex[:8]
return f"eplb-{self._rank}{pp_suffix}-{uid}"
def _get_peer_buckets(
def add_send(
self,
bucket_map: dict[torch.dtype, list[list[torch.Tensor]]],
dtype: torch.dtype,
) -> list[list[torch.Tensor]]:
peer_buckets = bucket_map.get(dtype)
if peer_buckets is None:
peer_buckets = [[] for _ in range(self._world_size)]
bucket_map[dtype] = peer_buckets
return peer_buckets
def add_send(self, tensor: torch.Tensor, dst_rank: int) -> None:
tensors: list[torch.Tensor],
dst_rank: int,
expert_id: int,
) -> None:
assert dst_rank != self._rank, (
"EPLB communicator should not enqueue same-rank sends: "
f"rank={self._rank}, dst_rank={dst_rank}"
)
self._get_peer_buckets(self._send_tensors, tensor.dtype)[dst_rank].append(
tensor
)
# An expert sent to multiple peers is packed only once; skip duplicates.
if expert_id not in self._expert_send_map:
self._expert_send_map[expert_id] = tensors
def add_recv(self, tensor: torch.Tensor, src_rank: int) -> None:
def add_recv(
self,
tensors: list[torch.Tensor],
src_rank: int,
expert_id: int,
) -> None:
assert src_rank != self._rank, (
"EPLB communicator should not enqueue same-rank recvs: "
f"rank={self._rank}, src_rank={src_rank}"
)
self._get_peer_buckets(self._recv_tensors, tensor.dtype)[src_rank].append(
tensor
)
recv_experts = self._recv_map.setdefault(src_rank, {})
if expert_id not in recv_experts:
recv_experts[expert_id] = tensors
def _init_remote_agents(self) -> None:
local_metadata = self._nixl_wrapper.get_agent_metadata()
@@ -303,30 +336,18 @@ class NixlEplbCommunicator(EplbCommunicator):
)
def _init_registered_buffers(self, expert_weights: Sequence[torch.Tensor]) -> None:
total_max_bytes = 0
for dtype in self._dtypes:
max_numel = max(
sum(t.numel() for t in expert_weights if t.dtype == dtype), 1
)
max_bytes = max_numel * dtype.itemsize
self._dtype_max_bytes[dtype] = max_bytes
total_max_bytes += max_bytes
self._peer_partition_bytes = total_max_bytes
# The send buffer needs world_size partitions because remote peers
# READ from fixed offsets (rank * partition_bytes).
# This allocates world_size * partition_bytes
# which can cause OOM on large models.
# TODO(ilmarkov): shrink to const * partition_bytes and execute
# communication in multiple steps dealing with the worst case.
send_total_bytes = self._peer_partition_bytes * self._world_size
total_bytes = max(sum(t.nbytes for t in expert_weights), 1)
assert total_bytes % self._num_local_experts == 0, (
f"Number of bytes in moe layer {total_bytes} is not divisible "
f"by number of local experts {self._num_local_experts}"
)
self._expert_bytes = total_bytes // self._num_local_experts
self._send_buffer = torch.empty(
send_total_bytes, device=self._device, dtype=torch.uint8
total_bytes, device=self._device, dtype=torch.uint8
)
self._recv_buffer = torch.empty(
self._peer_partition_bytes, device=self._device, dtype=torch.uint8
total_bytes, device=self._device, dtype=torch.uint8
)
descs = self._nixl_wrapper.get_reg_descs([self._send_buffer, self._recv_buffer])
@@ -336,12 +357,11 @@ class NixlEplbCommunicator(EplbCommunicator):
def _exchange_remote_send_meta(self) -> None:
"""Exchange send-buffer metadata so each rank can build dynamic
descriptors at execute time."""
local_meta: tuple[int, int, int] = (
local_meta: tuple[int, int] = (
self._send_buffer.data_ptr(),
self._peer_partition_bytes,
self._cuda_device_id,
)
gathered_meta: list[tuple[int, int, int] | None] = [None] * self._world_size
gathered_meta: list[tuple[int, int] | None] = [None] * self._world_size
torch.distributed.all_gather_object(
gathered_meta, local_meta, group=self._cpu_group
)
@@ -353,14 +373,11 @@ class NixlEplbCommunicator(EplbCommunicator):
@staticmethod
def _pack_send_buffer(
peer_tensors: list[torch.Tensor],
in_tensors: list[torch.Tensor],
send_buffer: torch.Tensor,
byte_offset: int,
) -> int:
"""
Returns the byte offset after the last written byte.
"""
for tensor in peer_tensors:
) -> None:
for tensor in in_tensors:
raw = tensor.reshape(-1).view(torch.uint8)
if raw.numel() == 0:
continue
@@ -368,18 +385,14 @@ class NixlEplbCommunicator(EplbCommunicator):
raw, non_blocking=True
)
byte_offset += raw.numel()
return byte_offset
@staticmethod
def _unpack_recv_buffer(
recv_buffer: torch.Tensor,
peer_tensors: list[torch.Tensor],
out_tensors: list[torch.Tensor],
byte_offset: int,
) -> int:
"""
Returns the byte offset after the last read byte.
"""
for tensor in peer_tensors:
) -> None:
for tensor in out_tensors:
num_bytes = tensor.numel() * tensor.element_size()
if num_bytes == 0:
continue
@@ -388,19 +401,6 @@ class NixlEplbCommunicator(EplbCommunicator):
non_blocking=True,
)
byte_offset += num_bytes
return byte_offset
def _release_all_cached_handles(self) -> None:
"""Best-effort release of every cached dlist and xfer handle."""
for local_dlist, remote_dlist, xfer in self._xfer_cache.values():
for release_fn, handle in (
(self._nixl_wrapper.release_xfer_handle, xfer),
(self._nixl_wrapper.release_dlist_handle, local_dlist),
(self._nixl_wrapper.release_dlist_handle, remote_dlist),
):
with contextlib.suppress(Exception):
release_fn(handle)
self._xfer_cache.clear()
def _wait_for_all_transfers(self, handles: list[int]) -> None:
pending = set(handles)
@@ -418,82 +418,68 @@ class NixlEplbCommunicator(EplbCommunicator):
if pending:
time.sleep(0.0005)
def _get_or_create_xfer(self, src: int, total_bytes: int, recv_offset: int) -> int:
"""Return a cached xfer handle or create and cache a new one."""
key = (src, total_bytes, recv_offset)
cached = self._xfer_cache.get(key)
if cached is not None:
return cached[2]
def _create_peer_xfer(
self,
src: int,
local_descs: list[tuple[int, int, int]],
remote_descs: list[tuple[int, int, int]],
) -> tuple[int, int, int]:
"""Create a batched xfer for multiple descriptors from one peer.
recv_base = self._recv_buffer.data_ptr()
Each element in *local_descs* / *remote_descs* is an
``(address, size, device_id)`` tuple.
Returns ``(local_dlist, remote_dlist, xfer_handle)``.
"""
local_desc = self._nixl_wrapper.get_xfer_descs(
[
(
recv_base + recv_offset,
total_bytes,
self._cuda_device_id,
)
],
self._nixl_memory_type,
local_descs, self._nixl_memory_type
)
local_handle = self._nixl_wrapper.prep_xfer_dlist(
"NIXL_INIT_AGENT",
local_desc,
)
remote_base, remote_part_bytes, remote_dev = self._remote_send_meta[src]
agent_name = self._remote_agents[src]
remote_desc = self._nixl_wrapper.get_xfer_descs(
[
(
remote_base + self._rank * remote_part_bytes,
total_bytes,
remote_dev,
)
],
self._nixl_memory_type,
remote_descs, self._nixl_memory_type
)
remote_handle = self._nixl_wrapper.prep_xfer_dlist(
agent_name,
self._remote_agents[src],
remote_desc,
)
indices = list(range(len(local_descs)))
xfer_handle = self._nixl_wrapper.make_prepped_xfer(
"READ",
local_handle,
[0],
indices,
remote_handle,
[0],
indices,
)
self._xfer_cache[key] = (local_handle, remote_handle, xfer_handle)
return xfer_handle
return (local_handle, remote_handle, xfer_handle)
def execute(self) -> None:
xfer_handles: list[int] = []
def execute(self, old_indices: np.ndarray | None = None) -> None:
assert old_indices is not None, (
"NixlEplbCommunicator.execute requires old_indices"
)
xfer_entries: list[tuple[int, int, int]] = []
try:
# Phase 1: pack send buffers.
n = self._num_local_experts
rank_experts = old_indices[: self._world_size * n].reshape(
self._world_size, n
)
# Build expert_id -> send slot mapping per rank.
expert_to_send_slot: list[dict[int, int]] = [
{int(eid): i for i, eid in enumerate(row) if eid != -1}
for row in rank_experts
]
# Phase 1: pack each expert at its slot offset in the send buffer.
with torch.cuda.stream(self._cuda_stream):
for dst in range(self._world_size):
byte_offset = dst * self._peer_partition_bytes
for dtype in self._dtypes:
peer_tensors = self._send_tensors.get(
dtype, [[] for _ in range(self._world_size)]
)[dst]
actual_bytes = sum(
t.numel() * t.element_size() for t in peer_tensors
)
if actual_bytes > self._dtype_max_bytes[dtype]:
raise RuntimeError(
"NIXL EPLB send overflow for dtype "
f"{dtype}: peer={dst}, "
f"required={actual_bytes}, "
f"capacity={self._dtype_max_bytes[dtype]}"
)
byte_offset = self._pack_send_buffer(
peer_tensors,
self._send_buffer,
byte_offset,
)
for expert_id, tensors in self._expert_send_map.items():
slot = expert_to_send_slot[self._rank][expert_id]
byte_offset = slot * self._expert_bytes
self._pack_send_buffer(tensors, self._send_buffer, byte_offset)
# Ensure all packed data is visible in device memory before pulls.
if self._cuda_stream is not None:
@@ -508,58 +494,65 @@ class NixlEplbCommunicator(EplbCommunicator):
timeout=timedelta(minutes=5),
)
# Phase 2: look up or create descriptors and issue all READs.
# Data from all peers is packed sequentially into the single
# partition-sized recv buffer at running offsets.
recv_offsets: dict[int, int] = {}
# Phase 2: issue one batched READ per peer.
recv_offsets: dict[tuple[int, int], int] = {}
recv_offset = 0
recv_base = self._recv_buffer.data_ptr()
for src in range(self._world_size):
if src == self._rank:
continue
actual_total_bytes = 0
for dtype in self._dtypes:
peer_tensors = self._recv_tensors.get(
dtype, [[] for _ in range(self._world_size)]
)[src]
actual_total_bytes += sum(
t.numel() * t.element_size() for t in peer_tensors
)
if actual_total_bytes == 0:
recv_experts = self._recv_map.get(src)
if not recv_experts:
continue
recv_offsets[src] = recv_offset
xfer_handle = self._get_or_create_xfer(
src, actual_total_bytes, recv_offset
expert_ids = list(recv_experts.keys())
remote_base, remote_dev = self._remote_send_meta[src]
local_descs: list[tuple[int, int, int]] = []
remote_descs: list[tuple[int, int, int]] = []
for expert_id in expert_ids:
slot = expert_to_send_slot[src][expert_id]
remote_off = slot * self._expert_bytes
recv_offsets[(src, expert_id)] = recv_offset
local_descs.append(
(
recv_base + recv_offset,
self._expert_bytes,
self._cuda_device_id,
)
)
remote_descs.append(
(remote_base + remote_off, self._expert_bytes, remote_dev)
)
recv_offset += self._expert_bytes
assert recv_offset <= self._recv_buffer.nbytes
local_h, remote_h, xfer_h = self._create_peer_xfer(
src, local_descs, remote_descs
)
self._nixl_wrapper.transfer(xfer_handle)
xfer_handles.append(xfer_handle)
recv_offset += actual_total_bytes
self._nixl_wrapper.transfer(xfer_h)
xfer_entries.append((local_h, remote_h, xfer_h))
# Phase 3: single wait for all in-flight transfers, then unpack.
self._wait_for_all_transfers(xfer_handles)
# Phase 3: wait for all in-flight transfers, then unpack.
self._wait_for_all_transfers([x[2] for x in xfer_entries])
with torch.cuda.stream(self._cuda_stream):
for src, offset in recv_offsets.items():
byte_offset = offset
for dtype in self._dtypes:
peer_tensors = self._recv_tensors.get(
dtype, [[] for _ in range(self._world_size)]
)[src]
byte_offset = self._unpack_recv_buffer(
self._recv_buffer,
peer_tensors,
byte_offset,
)
except Exception:
self._release_all_cached_handles()
raise
for (src, expert_id), offset in recv_offsets.items():
self._unpack_recv_buffer(
self._recv_buffer,
self._recv_map[src][expert_id],
offset,
)
finally:
self._send_tensors.clear()
self._recv_tensors.clear()
for local_h, remote_h, xfer_h in xfer_entries:
with contextlib.suppress(Exception):
self._nixl_wrapper.release_xfer_handle(xfer_h)
with contextlib.suppress(Exception):
self._nixl_wrapper.release_dlist_handle(local_h)
with contextlib.suppress(Exception):
self._nixl_wrapper.release_dlist_handle(remote_h)
self._expert_send_map.clear()
self._recv_map.clear()
def __del__(self) -> None:
try:
self._release_all_cached_handles()
if self._registered_desc is not None:
self._nixl_wrapper.deregister_memory(self._registered_desc)
self._registered_desc = None
@@ -588,15 +581,27 @@ class PyNcclEplbCommunicator(EplbCommunicator):
self._pynccl_comm.group_start()
self._group_started = True
def add_send(self, tensor: torch.Tensor, dst_rank: int) -> None:
def add_send(
self,
tensors: list[torch.Tensor],
dst_rank: int,
expert_id: int, # unused by this backend
) -> None:
self._ensure_group_started()
self._pynccl_comm.send(tensor, dst_rank, stream=self._cuda_stream)
for tensor in tensors:
self._pynccl_comm.send(tensor, dst_rank, stream=self._cuda_stream)
def add_recv(self, tensor: torch.Tensor, src_rank: int) -> None:
def add_recv(
self,
tensors: list[torch.Tensor],
src_rank: int,
expert_id: int, # unused by this backend
) -> None:
self._ensure_group_started()
self._pynccl_comm.recv(tensor, src_rank, stream=self._cuda_stream)
for tensor in tensors:
self._pynccl_comm.recv(tensor, src_rank, stream=self._cuda_stream)
def execute(self) -> None:
def execute(self, old_indices: np.ndarray | None = None) -> None:
if self._group_started:
self._pynccl_comm.group_end()
self._group_started = False
+8 -5
View File
@@ -294,9 +294,9 @@ def move_to_buffer(
recver_pos = remainder_start + sender_pos
if recver_pos < len(ranks_to_recv):
recv_ranks.append(ranks_to_recv[recver_pos])
expert_tensors = [w[src] for w in expert_weights]
for dst in recv_ranks:
for w in expert_weights:
communicator.add_send(w[src], dst)
communicator.add_send(expert_tensors, dst, expert_id=int(expert))
# 3. Post recvs
if recv_count > 0:
@@ -325,11 +325,14 @@ def move_to_buffer(
src = ranks_to_send[recver_pos // num_dst_per_sender]
else:
src = ranks_to_send[recver_pos - remainder_start]
for b in expert_weights_buffers:
communicator.add_recv(b[dst], src)
communicator.add_recv(
[b[dst] for b in expert_weights_buffers],
src,
expert_id=int(expert),
)
# 4. Execute the P2P operations. The real communication happens here.
communicator.execute()
communicator.execute(old_indices=old_indices)
# wait for the communication to finish
return TransferMetadata(
is_unchanged=is_unchanged,
@@ -1,9 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import copy
from collections.abc import Iterable
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, cast
import torch
@@ -18,6 +18,8 @@ from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorMetadata,
KVConnectorRole,
KVConnectorWorkerMetadata,
SupportsHMA,
supports_hma,
)
from vllm.distributed.kv_transfer.kv_connector.v1.metrics import (
KVConnectorPromMetrics,
@@ -123,7 +125,7 @@ class MultiKVConnectorPromMetrics(KVConnectorPromMetrics):
self._prom_metrics[connector_id].observe(stats_data["data"], engine_idx)
class MultiConnector(KVConnectorBase_V1):
class MultiConnector(KVConnectorBase_V1, SupportsHMA):
"""
A wrapper for using multiple KVConnectors at the same time.
@@ -166,6 +168,12 @@ class MultiConnector(KVConnectorBase_V1):
self._connectors.append(connector_cls(temp_config, role, kv_cache_config))
self._ktc_kv_transfer_config.append(temp_config.kv_transfer_config)
self._all_support_hma = all(supports_hma(c) for c in self._connectors)
assert (
vllm_config.scheduler_config.disable_hybrid_kv_cache_manager
or self._all_support_hma
), "HMA should not be enabled unless all sub-connectors support it"
# A mapping from request id to the index of the connector chosen to
# load the request from (if any).
self._requests_to_connector: dict[str, int] = {}
@@ -436,15 +444,17 @@ class MultiConnector(KVConnectorBase_V1):
for c in self._connectors:
c.set_xfer_handshake_metadata(metadata)
def request_finished(
def _aggregate_request_finished(
self,
request: "Request",
blocks: list[int],
per_connector_fn: Callable[
[KVConnectorBase_V1], tuple[bool, dict[str, Any] | None]
],
) -> tuple[bool, dict[str, Any] | None]:
async_saves = 0
kv_txfer_params = None
for c in self._connectors:
async_save, txfer_params = c.request_finished(request, blocks)
async_save, txfer_params = per_connector_fn(c)
if async_save:
async_saves += 1
if txfer_params is not None:
@@ -458,11 +468,39 @@ class MultiConnector(KVConnectorBase_V1):
if async_saves > 1:
self._extra_async_saves[request.request_id] = async_saves - 1
# Clean up other state for this request.
self._requests_to_connector.pop(request.request_id, None)
return async_saves > 0, kv_txfer_params
def request_finished(
self,
request: "Request",
blocks: list[int],
) -> tuple[bool, dict[str, Any] | None]:
return self._aggregate_request_finished(
request,
lambda c: c.request_finished(request, blocks),
)
def request_finished_all_groups(
self,
request: "Request",
block_ids: tuple[list[int], ...],
) -> tuple[bool, dict[str, Any] | None]:
if not self._all_support_hma:
assert len(block_ids) == 1, (
"HMA with multiple kv_cache_groups requires all "
"sub-connectors to support HMA"
)
return self.request_finished(request, block_ids[0])
return self._aggregate_request_finished(
request,
lambda c: cast(SupportsHMA, c).request_finished_all_groups(
request, block_ids
),
)
def take_events(self) -> Iterable["KVCacheEvent"]:
for c in self._connectors:
yield from c.take_events()
@@ -119,6 +119,30 @@ class NixlConnectorScheduler:
for n_tokens, block_size in sw_sizes_tokens
]
# Threshold to decide whether to compute kv cache locally
# or pull from a remote node: minimum number of remote
# tokens to amortize the xfer latencies
self.kv_recompute_threshold: int = int(
vllm_config.kv_transfer_config.get_from_extra_config(
"kv_recompute_threshold", 64
)
)
# Bi-directional KV transfer feature supports KV block
# transfers from D node to P node
self.is_bidirectional_kv_xfer_enabled = (
vllm_config.kv_transfer_config.get_from_extra_config(
"bidirectional_kv_xfer", False
)
)
if self.is_bidirectional_kv_xfer_enabled and self.kv_recompute_threshold > 0:
logger.info(
"Bidirectional KV transfer is enabled and the kv "
"recompute threshold is set to %d tokens",
self.kv_recompute_threshold,
)
def shutdown(self):
self._stop_event.set()
if self._nixl_handshake_listener_t is not None:
@@ -298,6 +322,44 @@ class NixlConnectorScheduler:
if params is not None and params.get("do_remote_decode") and self._has_mamba:
self._truncate_mamba_request_for_prefill(request)
if (
params is not None
and params.get("do_remote_decode")
and params.get("remote_block_ids")
and all(
p in params
for p in (
"remote_engine_id",
"remote_request_id",
"remote_host",
"remote_port",
)
)
):
# Decode node has kv blocks for part of prefill request, so, provide them
# as an external token count to scheduler.
# The tokens will be loaded if not already present
# in the prefill node local cache
remote_num_tokens = params.get("remote_num_tokens") or 0
count = (
min(remote_num_tokens, request.num_prompt_tokens) - num_computed_tokens
)
if count > 0:
# Check kv_recompute_threshold: skip pull if
# remote tokens are below the threshold.
if (
self.kv_recompute_threshold > 0
and count < self.kv_recompute_threshold
):
logger.debug(
"Skipping remote pull for %s: %d remote tokens < threshold %d",
request.request_id,
count,
self.kv_recompute_threshold,
)
return 0, False
return count, True
# No remote prefill for this request.
return 0, False
@@ -315,13 +377,19 @@ class NixlConnectorScheduler:
if not params:
return
if params.get("do_remote_decode"):
if params.get("do_remote_decode") or (
params.get("do_remote_prefill") and self.is_bidirectional_kv_xfer_enabled
):
self._reqs_in_batch.add(request.request_id)
if self.use_host_buffer and params.get("do_remote_decode"):
# NOTE: when accelerator is not directly supported by Nixl,
# prefilled blocks need to be saved to host memory before transfer.
self._reqs_need_save[request.request_id] = request
elif params.get("do_remote_prefill"):
elif params.get("do_remote_prefill") or (
params.get("do_remote_decode")
and self.is_bidirectional_kv_xfer_enabled
and not params.get("_remote_blocks_processed")
):
if params.get("remote_block_ids"):
if all(
p in params
@@ -333,8 +401,8 @@ class NixlConnectorScheduler:
)
):
# If remote_blocks and num_external_tokens = 0, we have
# a full prefix cache hit on the D worker. We need to call
# send_notif in _read_blocks to free the memory on the P.
# a full prefix cache hit on the local node. We need to call
# send_notif in _read_blocks to free the memory on the remote node.
unhashed_local_block_ids: BlockIds = (
blocks.get_unhashed_block_ids_all_groups()
@@ -362,6 +430,7 @@ class NixlConnectorScheduler:
assert num_external_tokens == 0
# Only trigger 1 KV transfer per request.
params["do_remote_prefill"] = False
params["_remote_blocks_processed"] = True
def _build_save_meta(
self,
@@ -450,6 +519,9 @@ class NixlConnectorScheduler:
if not params:
return False, None
is_p_node = bool(params.get("do_remote_decode"))
is_d_node = not is_p_node
if params.get("do_remote_prefill"):
# If do_remote_prefill is still True when the request is finished,
# update_state_after_alloc must not have been called (the request
@@ -461,9 +533,13 @@ class NixlConnectorScheduler:
params["do_remote_prefill"] = False
return False, None
if not params.get("do_remote_decode"):
if is_d_node and not self.is_bidirectional_kv_xfer_enabled:
return False, None
if request.status != RequestStatus.FINISHED_LENGTH_CAPPED:
if request.status not in (
RequestStatus.FINISHED_LENGTH_CAPPED,
RequestStatus.FINISHED_STOPPED,
):
# Also include the case of a P/D Prefill request with immediate
# block free (eg abort). Stop tracking this request.
self._reqs_not_processed.add(request.request_id)
@@ -474,7 +550,7 @@ class NixlConnectorScheduler:
# TODO: check whether block_ids actually ever be 0. If not we could
# remove the conditional below
delay_free_blocks = any(len(group) > 0 for group in block_ids)
remote_num_tokens = 0
if delay_free_blocks:
# Prefill request on remote. It will be read from D upon completion
logger.debug(
@@ -492,13 +568,16 @@ class NixlConnectorScheduler:
# Here we "unpad" blocks to send the actual remote blocks to be read.
block_ids = self.get_sw_clipped_blocks(block_ids)
remote_num_tokens = request.num_computed_tokens
return delay_free_blocks, dict(
do_remote_prefill=True,
do_remote_decode=False,
do_remote_prefill=is_p_node,
do_remote_decode=is_d_node,
remote_block_ids=block_ids,
remote_engine_id=self.engine_id,
remote_request_id=request.request_id,
remote_host=self.side_channel_host,
remote_port=self.side_channel_port,
tp_size=self.vllm_config.parallel_config.tensor_parallel_size,
remote_num_tokens=remote_num_tokens,
)
@@ -18,15 +18,15 @@ from vllm.logger import init_logger
from vllm.utils.math_utils import cdiv
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_offload.abstract import (
from vllm.v1.kv_offload.base import (
GPULoadStoreSpec,
OffloadingManager,
OffloadingSpec,
OffloadKey,
ReqContext,
get_offload_block_hash,
make_offload_key,
)
from vllm.v1.kv_offload.mediums import GPULoadStoreSpec
from vllm.v1.kv_offload.spec import OffloadingSpec
from vllm.v1.outputs import KVConnectorOutput
from vllm.v1.request import Request
@@ -25,7 +25,7 @@ from vllm.v1.kv_cache_interface import (
MambaSpec,
UniformTypeKVCacheSpecs,
)
from vllm.v1.kv_offload.spec import (
from vllm.v1.kv_offload.base import (
CanonicalKVCacheRef,
CanonicalKVCaches,
CanonicalKVCacheTensor,
@@ -755,6 +755,8 @@ class OpenAIServing:
def _is_model_supported(self, model_name: str | None) -> bool:
if not model_name:
return True
if envs.VLLM_SKIP_MODEL_NAME_VALIDATION:
return True
return self.models.is_base_model(model_name)
+1 -428
View File
@@ -87,7 +87,7 @@ _maybe_set_cuda_compatibility_path()
import torch
from vllm.logger import init_logger
from vllm.utils.torch_utils import is_torch_equal, is_torch_equal_or_newer
from vllm.utils.torch_utils import is_torch_equal_or_newer
logger = init_logger(__name__)
@@ -112,384 +112,6 @@ os.environ["TORCHINDUCTOR_COMPILE_THREADS"] = "1"
# in the environment.
os.environ.setdefault("TRITON_CACHE_AUTOTUNING", "1")
# ===================================================
# torch 2.9 Inductor PythonWrapperCodegen monkeypatch
# ===================================================
# This change monkeypatches memory_plan_reuse in pytorch 2.9.0 to work around
# a test failure for test_multi_graph_piecewise_compile_outputs_equal.
# For more context, see https://github.com/pytorch/pytorch/pull/165514.
def memory_plan_reuse_patched(self):
import torch._inductor.ir as ir
from torch._inductor.codegen.wrapper import (
EnterSubgraphLine,
ExitSubgraphLine,
MemoryPlanningLine,
MemoryPlanningState,
SubgraphPythonWrapperCodegen,
)
from torch._inductor.virtualized import V
def get_output_names(graph_outputs) -> list[str]:
import itertools
names = []
shape_counter = itertools.count(0)
none_counter = itertools.count(0)
for node in graph_outputs:
if isinstance(node, ir.NoneAsConstantBuffer):
names.append(f"{V.graph.name}_none{next(none_counter)}")
elif isinstance(node, ir.ShapeAsConstantBuffer):
names.append(f"{V.graph.name}_shape{next(shape_counter)}")
else:
names.append(node.get_name())
return names
if (
isinstance(V.graph.wrapper_code, SubgraphPythonWrapperCodegen)
and V.graph.wrapper_code.partition_signatures is not None
):
out_names = get_output_names(
V.graph.wrapper_code.partition_signatures.output_nodes
)
else:
out_names = V.graph.get_output_names()
while (
self.lines
and isinstance(self.lines[-1], MemoryPlanningLine)
and self.lines[-1].node.name not in out_names # type: ignore[attr-defined]
):
# these lines will be pointless
self.lines.pop()
# codegen allocations in two passes
planning_states = [MemoryPlanningState()]
past_planning_states = []
for i in range(len(self.lines)):
line = self.lines[i]
if isinstance(line, MemoryPlanningLine):
self.lines[i] = line.plan(planning_states[-1])
elif isinstance(line, EnterSubgraphLine):
planning_states.append(MemoryPlanningState())
elif isinstance(line, ExitSubgraphLine):
past_planning_states.append(planning_states.pop())
past_planning_states.append(planning_states.pop())
assert len(planning_states) == 0
# ===================================================
# torch 2.9 Inductor get_graph_partition_signature monkeypatch
# ===================================================
# This change monkeypatches get_graph_partition_signature in pytorch 2.9.0 to
# fix inductor partition + attention-nvfp4 quant fusion, tested in
# `tests/compile/test_fusion_attn.py::test_attn_quant`.
# For more context, see https://github.com/pytorch/pytorch/pull/165815.
def get_graph_partition_signature_patched(
self, partitions, skip_cudagraphs: list[bool]
):
"""
Gets signature for each graph partition, including input nodes, output nodes, and
whether deallocating an input within graph partition.
"""
from torch._inductor import dependencies
from torch._inductor.ir import GraphPartitionSignature, MutationOutput, NoneLayout
from torch._inductor.virtualized import V
from torch.utils._ordered_set import OrderedSet
signatures = []
unmet_output_names = OrderedSet(V.graph.get_output_names())
name_to_node = self.get_name_to_nodes()
def is_none_layout(buf_name: str) -> bool:
"""
Checks if buf_name is NoneLayout. Buffers with NoneLayout is not allocated
so graph partition should not take it as inputs or outputs.
"""
buf = self.name_to_buf.get(buf_name, None)
if buf is None:
return False
if isinstance(buf.node.layout, NoneLayout):
if isinstance(buf.node, MutationOutput) and (
real_name := self.mutation_real_name.get(buf_name, None)
):
return is_none_layout(real_name)
return True
return False
for partition, skip_cudagraph in zip(
reversed(partitions), reversed(skip_cudagraphs)
):
output_names: OrderedSet[str] = OrderedSet()
for node in partition:
output_names.update(node.outputs_by_name.keys())
returned_output_names = output_names.intersection(unmet_output_names)
# all reads/writes are partition inputs except those generated
# within the partition and tensor constants
read_writes = dependencies.ReadWrites.merge_list(
[node.read_writes for node in partition]
)
# WeakDep is fake dependency on unused buffer. It should not appear
# in partition_input_names for inputs that are actually read or written.
partition_input_names = (
OrderedSet(
[
x.name
for x in read_writes.reads | read_writes.writes
if not is_none_layout(x.name)
]
)
- output_names
)
partition_input_names = OrderedSet(
self.mutation_real_name.get(name, name) for name in partition_input_names
)
buffer_names_to_free: OrderedSet[str] = OrderedSet()
for node in partition:
buffer_names_to_free.update(node.last_usage)
# buffer_names_to_free may contain buffers allocated in previous
# graph partitions. These buffers should also be a partition
# input.
extra_input_names = [
name
for name in (buffer_names_to_free - output_names)
if name in name_to_node
]
partition_input_names.update(extra_input_names)
input_nodes = {
name: name_to_node[name]
for name in partition_input_names
if name in name_to_node
}
input_deallocation = {
name: name in buffer_names_to_free
for name in partition_input_names
if name in name_to_node
}
# if an input tensor is not freed in the partition function, it should
# also be returned as an output. This brings benefits to cudagraph
# since the returned output tensor is a cudagraph managed tensor with
# a static tensor address.
extra_output_names = [
name
for name in partition_input_names
if name in name_to_node and name not in buffer_names_to_free
]
returned_output_names.update(extra_output_names)
returned_output_names = OrderedSet(
self.mutation_real_name.get(name, name) for name in returned_output_names
)
output_nodes = [
name_to_node[name]
for name in returned_output_names
if not is_none_layout(name)
]
constant_names = [
name for name in partition_input_names if name in V.graph.constants
]
symbol_inputs = self.get_graph_partition_symbol_inputs(partition, input_nodes)
partition_signature = GraphPartitionSignature(
symbol_inputs,
input_nodes,
output_nodes,
input_deallocation,
skip_cudagraph,
constant_names,
)
signatures.append(partition_signature)
unmet_output_names = partition_input_names.union(
unmet_output_names - returned_output_names
)
return signatures[::-1]
# ========================================
# torch 2.9 Inductor Scheduler monkeypatch
# ========================================
# This change monkeypatches a function in Inductor to work around the following
# bug: https://github.com/vllm-project/vllm/issues/26678
#
# The bug occurs when `use_inductor_graph_partition` is turned on and there
# exists operators inside of `splitting_ops` that have an in-place mutation. In
# vllm, this specifically occurs on the operator
# vllm.unified_attention_with_output. In this case, inductor does not populate
# the inductor IR's `origin_node` field, causing an assertion error when trying
# to access the node's `origin_node` field.
#
# So, we will monkeypatch torch._inductor.scheduler.Scheduler.should_partition
# so that it does not access the inductor IR node's `origin_node` field and just
# returns True if a node is registered as having a custom partition function.
# This is ok for now since vllm's implementation of the custom partition
# functions just return True.
# ========================================
def should_partition_patched(self, node, should_log: bool = False) -> bool:
# This is a patched version of
# torch._inductor.scheduler.Scheduler.should_partition that modifies
# the following piece of code so that we always return True:
# https://github.com/pytorch/pytorch/blob/ecb53078faf86ca1b33277df33b82985675bb011/torch/_inductor/scheduler.py#L4712-L4724
"""Return True if we should partition the inductor graph on this node"""
import torch._inductor.ir as ir
from torch._inductor.scheduler import (
BaseSchedulerNode,
FusedSchedulerNode,
)
from torch._inductor.utils import (
_unstable_customized_partition_wrapper,
is_cudagraph_unsafe_op,
maybe_log_cudagraph_partition,
)
# Allow users to manually specify if a node should be partitioned
# Can only do this for FallbackKernels
ir_node = node.node
if isinstance(ir_node, torch._inductor.ir.FallbackKernel) and (
op := ir_node.op_overload
):
op_overload_packet_name = op.name()
op_overload_name = (
f"{op_overload_packet_name}.{op._overloadname}"
if isinstance(op, torch._ops.OpOverload)
else op_overload_packet_name
)
if (
op_overload_packet_name
in torch._inductor.config.custom_should_partition_ops
or op_overload_name in torch._inductor.config.custom_should_partition_ops
):
assert isinstance(op, torch._ops.OpOverload)
return True
# When not using cudagraphs, keep all kernels in the `call` function
# instead of graph partition functions, since graph partition only brings
# benefit to cudagraph
if (
not torch._inductor.config.triton.cudagraphs
and _unstable_customized_partition_wrapper.wrapper is None
):
return True
# avoid duplicating logs when should_partition is called multiple times
# on the same node
def noop_log(msg: str, node: BaseSchedulerNode | None) -> None:
return
log_partition_reason = maybe_log_cudagraph_partition if should_log else noop_log
if isinstance(node, FusedSchedulerNode):
return any(self.should_partition(snode) for snode in node.snodes)
assert node.node is not None
if not node.is_gpu():
log_partition_reason("non gpu ops", node=node)
return True
if isinstance(node.node, ir.DeviceCopy):
log_partition_reason("DeviceCopy ops", node=node)
return True
if isinstance(node.node, ir.Conditional):
log_partition_reason("Conditional ops", node=node)
return True
if getattr(node.node, "unbacked_bindings", None):
log_partition_reason("unbacked binding ops", node=node)
return True
if is_cudagraph_unsafe_op(node.node):
log_partition_reason("CUDAGraph-unsafe custom ops", node=node)
return True
return False
def _update_scheduler_patched(self) -> None:
# Copied from torch._inductor.graph.GrahLowering._update_scheduler. Patches
# this method so that we can patch Scheduler.should_partition with the
# function above
"""
(Re)initializes the scheduler member. When initializing the scheduler, no CUBIN
files should be generated (to avoid biasing any benchmarks and pessimizing
fusion decisions).
"""
import torch._inductor.config as config
from torch._inductor.scheduler import Scheduler
Scheduler.should_partition = should_partition_patched
Scheduler.get_graph_partition_signature = get_graph_partition_signature_patched
with config.patch("triton.store_cubin", False):
self.scheduler = Scheduler(self.operations)
# ===================================================
# torch 2.9 Inductor get_raw_stream workaround
# ===================================================
# Workaround for TorchInductor autotune using get_raw_stream() without defining it.
# This occurs when compile_sizes > 1 in compilation_config.
# For more context, see https://github.com/vllm-project/vllm/issues/30905.
def _patch_get_raw_stream_if_needed():
"""Workaround for TorchInductor autotune get_raw_stream() bug."""
from vllm.utils.torch_utils import is_torch_equal
# Only apply the patch for torch 2.9.0 or 2.9.1
if is_torch_equal("2.9.0") or is_torch_equal("2.9.1"):
import builtins
# Check if CUDA functionality is available without initializing CUDA
# _cuda_getCurrentRawStream only exists in CUDA builds of PyTorch
if hasattr(torch._C, "_cuda_getCurrentRawStream"):
from torch._C import _cuda_getCurrentRawStream as _get_raw_stream
builtins.get_raw_stream = _get_raw_stream # type: ignore[attr-defined]
_patch_get_raw_stream_if_needed()
if is_torch_equal("2.9.0"):
from torch._inductor.codegen.wrapper import PythonWrapperCodegen
from torch._inductor.graph import GraphLowering
from torch.utils._config_module import _Config, _ConfigEntry
# `custom_should_partition_ops` is a new config after 2.9.0. So this would
# not overwrite any user configs.
torch._inductor.config._config["custom_should_partition_ops"] = _ConfigEntry(
_Config(default=[])
)
PythonWrapperCodegen.memory_plan_reuse = memory_plan_reuse_patched
GraphLowering._update_scheduler = _update_scheduler_patched
# ===================================================
# torch <2.12 GraphCaptureOutput.get_runtime_env monkeypatch
# ===================================================
@@ -586,55 +208,6 @@ if is_torch_equal_or_newer("2.10.0") and not is_torch_equal_or_newer("2.12.0.dev
GraphCaptureOutput.get_runtime_env = _patched_get_runtime_env
# ===================================================
# torch 2.10 FxGraphCachePickler.dumps ValueError fix
# ===================================================
# PyTorch 2.10's FxGraphCachePickler.dumps() doesn't catch ValueError,
# causing torch.compile cache failures when tensors with non-standard
# layouts (e.g. blocked-layout prepacked weights) are serialized.
# PyTorch mainline fixed this in pytorch/pytorch#176557 (merged 2026-03-04).
# This is a thin backport for 2.10 users; remove once 2.10 is dropped.
def _apply_fxgraphcache_pickle_patch(pickler_cls, bypass_cls):
"""Wrap pickler_cls.dumps to convert ValueError into bypass_cls.
Idempotent: sets `_vllm_fxgraph_dumps_patched` on the class after the
first apply to prevent re-application. The wrapper function is also
marked with `_vllm_patched` as an additional safeguard.
"""
if getattr(pickler_cls, "_vllm_fxgraph_dumps_patched", False):
return
original_dumps = pickler_cls.dumps
if hasattr(original_dumps, "_vllm_patched"):
return
def patched_dumps(self, obj):
try:
return original_dumps(self, obj)
except ValueError as e:
raise bypass_cls("Failed to pickle cache key") from e
patched_dumps._vllm_patched = True # type: ignore[attr-defined]
pickler_cls.dumps = patched_dumps
pickler_cls._vllm_fxgraph_dumps_patched = True # type: ignore[attr-defined]
def _patch_fxgraphcache_pickle_if_needed():
"""Apply FxGraphCachePickler.dumps ValueError backport when on torch 2.10.x."""
from vllm.utils.torch_utils import is_torch_equal_or_newer
if not is_torch_equal_or_newer("2.10.0") or is_torch_equal_or_newer("2.11.0"):
return
from torch._inductor.codecache import BypassFxGraphCache, FxGraphCachePickler
_apply_fxgraphcache_pickle_patch(FxGraphCachePickler, BypassFxGraphCache)
_patch_fxgraphcache_pickle_if_needed()
# ===================================================
# torch 2.11 Inductor cpp codegen indirect_assert scalar-mask fix
# ===================================================
+14
View File
@@ -255,6 +255,10 @@ if TYPE_CHECKING:
VLLM_LORA_DISABLE_PDL: bool = False
VLLM_ENABLE_CUDA_COMPATIBILITY: bool = False
VLLM_CUDA_COMPATIBILITY_PATH: str | None = None
VLLM_SKIP_MODEL_NAME_VALIDATION: bool = False
"""If set, vLLM will skip model name validation in API requests.
This allows any model name to be accepted in the 'model' field of requests,
making the server model-name agnostic. Useful for proxy/gateway scenarios."""
VLLM_ELASTIC_EP_SCALE_UP_LAUNCH: bool = False
VLLM_ELASTIC_EP_DRAIN_REQUESTS: bool = False
VLLM_MEMORY_PROFILER_ESTIMATE_CUDAGRAPHS: bool = True
@@ -999,6 +1003,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
# use aiter linear op if aiter ops are enabled
# The following list of related ops
# - scaled_mm (per-tensor / rowwise)
# - use aiter tuned gemms for unquantized gemms
"VLLM_ROCM_USE_AITER_LINEAR": lambda: (
os.getenv("VLLM_ROCM_USE_AITER_LINEAR", "True").lower() in ("true", "1")
),
@@ -1710,6 +1715,14 @@ environment_variables: dict[str, Callable[[], Any]] = {
"VLLM_CUDA_COMPATIBILITY_PATH": lambda: os.environ.get(
"VLLM_CUDA_COMPATIBILITY_PATH", None
),
# Skip model name validation in OpenAI API requests.
# When set to 1, any model name will be accepted in the 'model' field
# of API requests. This is useful for proxy/gateway scenarios where
# the actual model is served but different names may be used in requests.
"VLLM_SKIP_MODEL_NAME_VALIDATION": lambda: (
os.getenv("VLLM_SKIP_MODEL_NAME_VALIDATION", "0").strip().lower()
in ("1", "true")
),
# Whether it is a scale up launch engine for elastic EP,
# Should only be set by EngineCoreClient.
"VLLM_ELASTIC_EP_SCALE_UP_LAUNCH": lambda: bool(
@@ -1885,6 +1898,7 @@ def compile_factors() -> dict[str, object]:
"VLLM_TEST_FORCE_LOAD_FORMAT",
"VLLM_ENABLE_CUDA_COMPATIBILITY",
"VLLM_CUDA_COMPATIBILITY_PATH",
"VLLM_SKIP_MODEL_NAME_VALIDATION",
"LOCAL_RANK",
"CUDA_VISIBLE_DEVICES",
"NO_COLOR",
+6 -7
View File
@@ -1076,10 +1076,10 @@ def chunk_gla_fwd_kernel_o(
)
p_h = tl.make_block_ptr(
h + (i_tg * H + i_h) * K * V,
(K, V),
(V, 1),
(i_k * BK, i_v * BV),
(BK, BV),
(V, K),
(K, 1),
(i_v * BV, i_k * BK),
(BV, BK),
(1, 0),
)
@@ -1090,12 +1090,11 @@ def chunk_gla_fwd_kernel_o(
b_g = tl.load(p_g, boundary_check=(0, 1))
# [BT, BK]
b_qg = (b_q * exp(b_g)).to(b_q.dtype)
# [BK, BV]
# [BV, BK]
b_h = tl.load(p_h, boundary_check=(0, 1))
# works but dkw, owing to divine benevolence
# [BT, BV]
if i_k >= 0:
b_o += tl.dot(b_qg, b_h.to(b_qg.dtype))
b_o += tl.dot(b_qg, tl.trans(b_h).to(b_qg.dtype))
p_v = tl.make_block_ptr(
v + (bos * H + i_h) * V,
(T, V),
+11 -12
View File
@@ -1487,15 +1487,19 @@ class FusedMoE(PluggableLayer):
"w2_input_scale",
}
# Parameters of non-expert submodules that live inside runner (MoERunner).
# These must be excluded from EPLB weight rearrangement.
NON_EXPERT_PREFIXES = (
"runner._shared_experts.",
"runner.gate.",
"runner.routed_input_transform.",
"runner.routed_output_transform.",
)
assert all(
weight.is_contiguous()
for name, weight in weights
if not (
name.startswith("_shared_experts.")
or name.startswith("_gate.")
or name.startswith("_routed_input_transform.")
or name.startswith("_routed_output_transform.")
)
if not name.startswith(NON_EXPERT_PREFIXES)
and name not in NON_EXPERT_WEIGHTS
)
@@ -1504,12 +1508,7 @@ class FusedMoE(PluggableLayer):
for name, weight in weights
if name not in NON_EXPERT_WEIGHTS
and weight.shape != torch.Size([])
and not name.startswith("_shared_experts.")
# exclude parameters from non-expert submodules,
# e.g. gate/shared/transforms.
and not name.startswith("_gate.")
and not name.startswith("_routed_input_transform.")
and not name.startswith("_routed_output_transform.")
and not name.startswith(NON_EXPERT_PREFIXES)
]
def set_eplb_state(
@@ -16,7 +16,6 @@ from vllm.model_executor.layers.fused_moe.utils import (
moe_kernel_quantize_input,
normalize_batched_scales_shape,
)
from vllm.platforms import current_platform
from vllm.v1.worker.ubatching import (
dbo_current_ubatch_id,
dbo_enabled,
@@ -290,46 +289,29 @@ class DeepEPLLPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
# Dispatch
dispatch_topk_ids = self._map_global_to_physical_ids(topk_ids)
if current_platform.is_rocm():
(
expert_x,
expert_num_tokens,
handle,
_,
hook,
) = self.buffer.low_latency_dispatch(
a1,
dispatch_topk_ids,
self.max_tokens_per_rank,
num_experts,
use_fp8=self.use_fp8_dispatch,
async_finish=False,
return_recv_hook=True,
)
else:
(
expert_x,
expert_num_tokens,
handle,
_,
hook,
) = self.buffer.low_latency_dispatch(
a1,
dispatch_topk_ids,
self.max_tokens_per_rank,
num_experts,
use_fp8=self.use_fp8_dispatch,
round_scale=self.use_ue8m0_dispatch,
use_ue8m0=self.use_ue8m0_dispatch,
**(dict(use_nvfp4=True) if use_nvfp4 else dict()),
**(
dict(x_global_scale=qc_a1_gscale_or_scale)
if qc_a1_gscale_or_scale is not None and nvfp4_dispatch
else dict()
),
async_finish=False,
return_recv_hook=True,
)
(
expert_x,
expert_num_tokens,
handle,
_,
hook,
) = self.buffer.low_latency_dispatch(
a1,
dispatch_topk_ids,
self.max_tokens_per_rank,
num_experts,
use_fp8=self.use_fp8_dispatch,
round_scale=self.use_ue8m0_dispatch,
use_ue8m0=self.use_ue8m0_dispatch,
**(dict(use_nvfp4=True) if use_nvfp4 else dict()),
**(
dict(x_global_scale=qc_a1_gscale_or_scale)
if qc_a1_gscale_or_scale is not None and nvfp4_dispatch
else dict()
),
async_finish=False,
return_recv_hook=True,
)
self.handles[a2a_idx] = handle
return (
@@ -252,7 +252,8 @@ def rocm_aiter_fused_experts(
else:
quant_method = QuantMethod.NO.value
# mxfp4: both w4a4 (quark) and w4a16 (oracle CK) use BLOCK_1X32
# mxfp4 i.e. w4a4, w4a16 uses BLOCK_1X32
# mxfp6 and mxfp8 are unsupported in AITER currently and use emulation instead
if quant_config.use_mxfp4_w4a4 or quant_config.use_mxfp4_w4a16:
quant_method = QuantMethod.BLOCK_1X32.value
# w8a8 block-scaled
@@ -220,7 +220,7 @@ class MoERunner(MoERunnerInterface):
self.routed_output_transform = routed_output_transform
self.routed_scaling_factor = routed_scaling_factor
self.gate = gate
self.quant_method = quant_method
self._quant_method = quant_method
self.enable_dbo = enable_dbo
self._shared_experts: SharedExperts | None = None
@@ -263,7 +263,7 @@ class MoERunner(MoERunnerInterface):
def _replace_quant_method(self, quant_method: FusedMoEMethodBase):
if self._shared_experts is not None:
self._shared_experts._quant_method = quant_method
self.quant_method = quant_method
self._quant_method = quant_method
def is_internal_router(self) -> bool:
return self.gate is not None
@@ -330,8 +330,8 @@ class MoERunner(MoERunnerInterface):
@property
def _fused_output_is_reduced(self) -> bool:
return (
self.quant_method.moe_kernel is not None
and self.quant_method.moe_kernel.output_is_reduced()
self._quant_method.moe_kernel is not None
and self._quant_method.moe_kernel.output_is_reduced()
)
def _maybe_reduce_shared_expert_output(
@@ -407,7 +407,7 @@ class MoERunner(MoERunnerInterface):
)
transformed_hidden_dim = hidden_states.shape[-1]
if (
not self.quant_method.skip_forward_padding
not self._quant_method.skip_forward_padding
and self.moe_config.hidden_dim != transformed_hidden_dim
):
hidden_states = F.pad(
@@ -451,8 +451,8 @@ class MoERunner(MoERunnerInterface):
shared_experts_input, SharedExpertsOrder.NO_OVERLAP
)
if self.quant_method.is_monolithic:
fused_out = self.quant_method.apply_monolithic(
if self._quant_method.is_monolithic:
fused_out = self._quant_method.apply_monolithic(
layer=layer,
x=hidden_states,
router_logits=router_logits,
@@ -467,7 +467,7 @@ class MoERunner(MoERunnerInterface):
# Passing shared_experts_input in case SharedExpertsOrder is
# MK_INTERNAL_OVERLAPPED.
fused_out = self.quant_method.apply(
fused_out = self._quant_method.apply(
layer=layer,
x=hidden_states,
topk_weights=topk_weights,
@@ -618,7 +618,7 @@ class MoERunner(MoERunnerInterface):
@property
def do_naive_dispatch_combine(self) -> bool:
return (
self.moe_config.dp_size > 1 and not self.quant_method.supports_internal_mk
self.moe_config.dp_size > 1 and not self._quant_method.supports_internal_mk
)
def _maybe_dispatch(

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