forked from Karylab-cklius/vllm
Compare commits
154
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
78d2334aab | ||
|
|
b53b1c7ffe | ||
|
|
79ca54d221 | ||
|
|
09f3cd5c10 | ||
|
|
ea6078fe6a | ||
|
|
a0df04e477 | ||
|
|
e2352c2974 | ||
|
|
25faa1f4cc | ||
|
|
4583630b56 | ||
|
|
21da47dabe | ||
|
|
6c379b9e54 | ||
|
|
5099474633 | ||
|
|
058cc0a8b6 | ||
|
|
837db7605e | ||
|
|
bf2a393034 | ||
|
|
d682968aa9 | ||
|
|
021cdf72bc | ||
|
|
4cb5e746b6 | ||
|
|
22cc891108 | ||
|
|
afdcbd5d39 | ||
|
|
8d4f54966c | ||
|
|
351c72d6e5 | ||
|
|
7299e6509e | ||
|
|
08985351f3 | ||
|
|
5fd3b276f8 | ||
|
|
1e9f04da14 | ||
|
|
702214146c | ||
|
|
a331589394 | ||
|
|
e945169207 | ||
|
|
554352a311 | ||
|
|
421c1ec448 | ||
|
|
b4c80ec0fd | ||
|
|
f428718ffe | ||
|
|
4403af8fb5 | ||
|
|
d57888efa4 | ||
|
|
ed938ad7db | ||
|
|
731fb3323d | ||
|
|
8dd8b6ed78 | ||
|
|
e1a5fc406b | ||
|
|
b4092176b9 | ||
|
|
ebbb2d55ac | ||
|
|
2959a9273a | ||
|
|
1797576237 | ||
|
|
0d339cf135 | ||
|
|
5fd21eb0b2 | ||
|
|
9d4b87f4f0 | ||
|
|
58b2e89642 | ||
|
|
2659f60a1a | ||
|
|
091386a99b | ||
|
|
d112eb1ac7 | ||
|
|
2a47a9ff0f | ||
|
|
9c7c74bf10 | ||
|
|
5e27b2baf4 | ||
|
|
eb0fdeb1e8 | ||
|
|
46f74e144b | ||
|
|
0a7bacdcac | ||
|
|
8b2b566ea7 | ||
|
|
0b131b16c9 | ||
|
|
bcb518ad7a | ||
|
|
06e1e0885c | ||
|
|
1a59078c87 | ||
|
|
fa85ead2f3 | ||
|
|
e28e8c8782 | ||
|
|
ee0fd6984a | ||
|
|
d537122398 | ||
|
|
3d20275bb4 | ||
|
|
f694d43b33 | ||
|
|
3c6084bb0d | ||
|
|
68ff30d40e | ||
|
|
6d8fff5698 | ||
|
|
e2c58570ea | ||
|
|
43fa24e832 | ||
|
|
93bbe94d3a | ||
|
|
17bc144556 | ||
|
|
295232a26a | ||
|
|
56e4345226 | ||
|
|
e9993a52aa | ||
|
|
a46abb7ae6 | ||
|
|
4c62663315 | ||
|
|
d78650cf97 | ||
|
|
5bdc01bcc3 | ||
|
|
20a5f8b43b | ||
|
|
7b5d60cc37 | ||
|
|
14b438a98b | ||
|
|
2785a5e0e6 | ||
|
|
efd15e192a | ||
|
|
556b063e45 | ||
|
|
aa0ac8a661 | ||
|
|
b831374cf1 | ||
|
|
71bc19dbdd | ||
|
|
ef2c40dc00 | ||
|
|
4bf699d310 | ||
|
|
520828789c | ||
|
|
9d4dc4ca2f | ||
|
|
b9684d99e9 | ||
|
|
4fadf9c92c | ||
|
|
d8d95998dc | ||
|
|
475a6ad18a | ||
|
|
f2beaa80c8 | ||
|
|
8e27a9c215 | ||
|
|
7d567172fc | ||
|
|
f00e163f35 | ||
|
|
44b2512767 | ||
|
|
188c68798e | ||
|
|
c45f681932 | ||
|
|
89e8645a9e | ||
|
|
88a9cdd439 | ||
|
|
6f612fbedf | ||
|
|
506ec6d656 | ||
|
|
a52205bccf | ||
|
|
3d34f8cbdc | ||
|
|
eb04c769d3 | ||
|
|
ce3ef17bec | ||
|
|
bf5149b516 | ||
|
|
cca3365b73 | ||
|
|
040df8f2ea | ||
|
|
ced32bb474 | ||
|
|
c5e5c33fcd | ||
|
|
a8c86eeb16 | ||
|
|
7e179e4bc0 | ||
|
|
405c7cf283 | ||
|
|
3f53e2138f | ||
|
|
d53f4593ce | ||
|
|
ad32608e24 | ||
|
|
b2cfae777d | ||
|
|
3f1ff1ff14 | ||
|
|
c69c73418a | ||
|
|
ebf3a6d705 | ||
|
|
c4fd9794e9 | ||
|
|
7ad894c86a | ||
|
|
a7fdfeef72 | ||
|
|
8bf374955f | ||
|
|
9096659edb | ||
|
|
81d8f4ebac | ||
|
|
a9a8a32dcd | ||
|
|
9d808e2309 | ||
|
|
f3858d5422 | ||
|
|
259ff891be | ||
|
|
6607a80dab | ||
|
|
b8bd773fe4 | ||
|
|
2addbb9cc9 | ||
|
|
e3cfea2e1b | ||
|
|
f99260d2aa | ||
|
|
3f65e21e32 | ||
|
|
b00e76ff72 | ||
|
|
f4359a70f9 | ||
|
|
3afe659b6b | ||
|
|
16e91176cf | ||
|
|
ab8b0fe338 | ||
|
|
d467a2a7f2 | ||
|
|
76a373eff4 | ||
|
|
25ee659db0 | ||
|
|
eacff17c8d | ||
|
|
7b375c8502 |
@@ -0,0 +1,68 @@
|
||||
group: Intel
|
||||
steps:
|
||||
- label: ":docker: Build XPU image"
|
||||
soft_fail: true
|
||||
optional: true
|
||||
depends_on: []
|
||||
key: image-build-xpu
|
||||
commands:
|
||||
- bash -lc '.buildkite/image_build/image_build_xpu.sh "public.ecr.aws/q9t5s3a7" "vllm-ci-test-repo" "$BUILDKITE_COMMIT"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: -1 # Agent was lost
|
||||
limit: 2
|
||||
- exit_status: -10 # Agent was lost
|
||||
limit: 2
|
||||
- label: "XPU example Test"
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
timeout_in_minutes: 30
|
||||
optional: true
|
||||
device: intel_gpu
|
||||
no_plugin: true
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
source_file_dependencies:
|
||||
- .buildkite/hardware_tests/intel_xpu_ci/test-intel.yaml
|
||||
- .buildkite/scripts/hardware_ci/run-intel-ci-test.sh
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'bash .buildkite/scripts/hardware_ci/run-intel-ci-test.sh example'
|
||||
- label: "XPU V1 test"
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
timeout_in_minutes: 30
|
||||
optional: true
|
||||
device: intel_gpu
|
||||
no_plugin: true
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
source_file_dependencies:
|
||||
- .buildkite/hardware_tests/intel_xpu_ci/test-intel.yaml
|
||||
- .buildkite/scripts/hardware_ci/run-intel-ci-test.sh
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'bash .buildkite/scripts/hardware_ci/run-intel-ci-test.sh v1'
|
||||
- label: "XPU server test"
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
timeout_in_minutes: 30
|
||||
optional: true
|
||||
device: intel_gpu
|
||||
no_plugin: true
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
source_file_dependencies:
|
||||
- .buildkite/hardware_tests/intel_xpu_ci/test-intel.yaml
|
||||
- .buildkite/scripts/hardware_ci/run-intel-ci-test.sh
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'bash .buildkite/scripts/hardware_ci/run-intel-ci-test.sh server'
|
||||
@@ -57,13 +57,16 @@ steps:
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
|
||||
'pip install lm_eval[api]>=0.4.12 &&
|
||||
export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
|
||||
cd tests &&
|
||||
pytest -v -s v1/logits_processors --ignore=v1/logits_processors/test_custom_online.py --ignore=v1/logits_processors/test_custom_offline.py &&
|
||||
pytest -v -s v1/test_oracle.py &&
|
||||
pytest -v -s v1/test_request.py &&
|
||||
pytest -v -s v1/test_outputs.py &&
|
||||
pytest -v -s v1/sample/test_topk_topp_sampler.py'
|
||||
pytest -v -s v1/sample/test_topk_topp_sampler.py &&
|
||||
pytest -v -s v1/sample/test_logprobs.py &&
|
||||
pytest -v -s v1/sample/test_logprobs_e2e.py'
|
||||
|
||||
- label: XPU CPU Offload
|
||||
timeout_in_minutes: 60
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
group: Model Runner V2 Intel
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
steps:
|
||||
- label: Model Runner V2 Core Tests (Intel)
|
||||
timeout_in_minutes: 45
|
||||
device: intel_gpu
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/v1/worker/gpu/
|
||||
- vllm/v1/worker/gpu_worker.py
|
||||
- vllm/v1/core/sched/
|
||||
- vllm/v1/attention/
|
||||
- tests/v1/engine/test_llm_engine.py
|
||||
- tests/v1/e2e/
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
|
||||
cd tests &&
|
||||
pytest -v -s v1/engine/test_llm_engine.py -k "not test_engine_metrics" &&
|
||||
ENFORCE_EAGER=1 pytest -v -s v1/e2e/general/test_async_scheduling.py -k "not ngram" &&
|
||||
pytest -v -s v1/e2e/general/test_min_tokens.py'
|
||||
|
||||
- label: Model Runner V2 Examples (Intel)
|
||||
timeout_in_minutes: 45
|
||||
device: intel_gpu
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/v1/worker/gpu/
|
||||
- vllm/v1/core/sched/
|
||||
- vllm/v1/worker/gpu_worker.py
|
||||
- examples/basic/offline_inference/
|
||||
- examples/generate/multimodal/
|
||||
- examples/features/
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
|
||||
cd examples &&
|
||||
python3 basic/offline_inference/chat.py &&
|
||||
python3 basic/offline_inference/generate.py --model facebook/opt-125m &&
|
||||
python3 generate/multimodal/vision_language_offline.py --seed 0 &&
|
||||
python3 features/automatic_prefix_caching/prefix_caching_offline.py'
|
||||
@@ -60,6 +60,7 @@ steps:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh &&
|
||||
pytest -v -s v1/core --ignore=v1/core/test_reset_prefix_cache_e2e.py --ignore=v1/core/test_scheduler_e2e.py &&
|
||||
pytest -v -s v1/engine --ignore=v1/engine/test_output_processor.py &&
|
||||
pytest -v -s v1/sample --ignore=v1/sample/test_logprobs.py --ignore=v1/sample/test_logprobs_e2e.py -k "not test_topk_only and not test_topp_only and not test_topk_and_topp" &&
|
||||
@@ -87,3 +88,20 @@ steps:
|
||||
cd tests &&
|
||||
pytest -v -s entrypoints/multimodal/openai/chat_completion/test_audio_in_video.py &&
|
||||
pytest -v -s benchmarks/test_serve_cli.py'
|
||||
- label: "XPU quantization test"
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
timeout_in_minutes: 30
|
||||
device: intel_gpu
|
||||
no_plugin: true
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- .buildkite/intel_jobs/test-intel.yaml
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
pytest -v -s quantization/test_auto_round.py'
|
||||
@@ -0,0 +1,51 @@
|
||||
#!/bin/bash
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
test_suite="${1:-}"
|
||||
|
||||
if [[ -z "${test_suite}" ]]; then
|
||||
echo "Usage: $0 <example|v1|server>" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
case "${test_suite}" in
|
||||
example)
|
||||
pip install tblib==3.1.0
|
||||
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 -O3 -cc.cudagraph_mode=NONE
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager -tp 2 --distributed-executor-backend mp
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager --attention-backend=TRITON_ATTN
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager --quantization fp8
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager --kv-cache-dtype fp8
|
||||
python3 examples/basic/offline_inference/generate.py --model nvidia/Llama-3.1-8B-Instruct-FP8 --block-size 64 --enforce-eager --quantization modelopt --kv-cache-dtype fp8 --attention-backend TRITON_ATTN --max-model-len 4096
|
||||
python3 examples/basic/offline_inference/generate.py --model superjob/Qwen3-4B-Instruct-2507-GPTQ-Int4 --block-size 64 --enforce-eager --max-model-len 8192
|
||||
python3 examples/basic/offline_inference/generate.py --model ibm-research/PowerMoE-3b --block-size 64 --enforce-eager -tp 2
|
||||
python3 examples/basic/offline_inference/generate.py --model ibm-research/PowerMoE-3b --block-size 64 --enforce-eager -tp 2 --enable-expert-parallel
|
||||
python3 examples/basic/offline_inference/generate.py --model superjob/Qwen3-4B-Instruct-2507-GPTQ-Int4 --max-model-len 8192
|
||||
;;
|
||||
v1)
|
||||
cd tests
|
||||
|
||||
pytest -v -s v1/core --ignore=v1/core/test_reset_prefix_cache_e2e.py --ignore=v1/core/test_scheduler_e2e.py
|
||||
pytest -v -s v1/engine --ignore=v1/engine/test_output_processor.py
|
||||
pytest -v -s v1/sample --ignore=v1/sample/test_logprobs.py --ignore=v1/sample/test_logprobs_e2e.py -k "not test_topk_only and not test_topp_only and not test_topk_and_topp"
|
||||
pytest -v -s v1/worker --ignore=v1/worker/test_gpu_model_runner.py --ignore=v1/worker/test_worker_memory_snapshot.py
|
||||
pytest -v -s v1/structured_output
|
||||
pytest -v -s v1/test_serial_utils.py
|
||||
pytest -v -s v1/spec_decode --ignore=v1/spec_decode/test_max_len.py --ignore=v1/spec_decode/test_speculators_eagle3.py --ignore=v1/spec_decode/test_acceptance_length.py
|
||||
pytest -v -s v1/kv_connector/unit --ignore=v1/kv_connector/unit/test_multi_connector.py --ignore=v1/kv_connector/unit/test_example_connector.py --ignore=v1/kv_connector/unit/test_lmcache_integration.py --ignore=v1/kv_connector/unit/test_hf3fs_client.py --ignore=v1/kv_connector/unit/test_hf3fs_connector.py --ignore=v1/kv_connector/unit/test_hf3fs_metadata_server.py --ignore=v1/kv_connector/unit/test_offloading_connector.py
|
||||
;;
|
||||
server)
|
||||
pip install av
|
||||
cd tests
|
||||
|
||||
pytest -v -s entrypoints/multimodal/openai/chat_completion/test_audio_in_video.py
|
||||
pytest -v -s benchmarks/test_serve_cli.py
|
||||
;;
|
||||
*)
|
||||
echo "Unknown Intel test suite: ${test_suite}" >&2
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
@@ -243,8 +243,10 @@ container_name="xpu_${BUILDKITE_COMMIT}_$(tr -dc A-Za-z0-9 < /dev/urandom | head
|
||||
|
||||
# ---- Command source selection ----
|
||||
commands=""
|
||||
commands_source=""
|
||||
if [[ -n "${VLLM_TEST_COMMANDS:-}" ]]; then
|
||||
commands="${VLLM_TEST_COMMANDS}"
|
||||
commands_source="env"
|
||||
echo "Commands sourced from VLLM_TEST_COMMANDS (quoting preserved)"
|
||||
elif [[ $# -gt 0 ]]; then
|
||||
all_yaml=true
|
||||
@@ -303,8 +305,12 @@ if [[ -z "$commands" ]]; then
|
||||
fi
|
||||
|
||||
echo "Raw commands: $commands"
|
||||
commands=$(re_quote_pytest_markers "$commands")
|
||||
echo "After re-quoting: $commands"
|
||||
if [[ "$commands_source" != "env" ]]; then
|
||||
commands=$(re_quote_pytest_markers "$commands")
|
||||
echo "After re-quoting: $commands"
|
||||
else
|
||||
echo "Skipping re-quoting for VLLM_TEST_COMMANDS input"
|
||||
fi
|
||||
commands=$(apply_intel_test_overrides "$commands")
|
||||
echo "Final commands: $commands"
|
||||
|
||||
|
||||
+25
-24
@@ -415,22 +415,6 @@ steps:
|
||||
commands:
|
||||
- pytest -v -s kernels/mamba
|
||||
|
||||
#----------------------------------------------------------- mi250 · lora ------------------------------------------------------------#
|
||||
|
||||
- label: LoRA %N # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
|
||||
agent_pool: mi250_1
|
||||
parallelism: 4
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/lora
|
||||
- tests/lora
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s lora --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --ignore=lora/test_chatglm3_tp.py --ignore=lora/test_llama_tp.py --ignore=lora/test_qwen3_with_multi_loras.py --ignore=lora/test_olmoe_tp.py --ignore=lora/test_deepseekv2_tp.py --ignore=lora/test_gptoss_tp.py --ignore=lora/test_qwen3moe_tp.py --ignore=lora/test_qwen35_densemodel_lora.py
|
||||
|
||||
#------------------------------------------------------ mi250 · models / basic -------------------------------------------------------#
|
||||
|
||||
- label: Basic Models Test (Other CPU) # TBD
|
||||
@@ -608,6 +592,11 @@ steps:
|
||||
- pytest -v -s plugins_tests/test_bge_m3_sparse_io_processor_plugins.py
|
||||
- pip uninstall bge_m3_sparse_plugin -y
|
||||
# END: `bge_m3_sparse io_processor` test
|
||||
# BEGIN: `colbert_query io_processor` test
|
||||
- pip install -e ./plugins/colbert_query_plugin
|
||||
- pytest -v -s plugins_tests/test_colbert_query_io_processor_plugins.py
|
||||
- pip uninstall colbert_query_plugin -y
|
||||
# END: `colbert_query io_processor` test
|
||||
# BEGIN: `stat_logger` plugins test
|
||||
- pip install -e ./plugins/vllm_add_dummy_stat_logger
|
||||
- pytest -v -s plugins_tests/test_stats_logger_plugins.py
|
||||
@@ -1694,6 +1683,20 @@ steps:
|
||||
|
||||
#----------------------------------------------------------- mi300 · lora ------------------------------------------------------------#
|
||||
|
||||
- label: LoRA %N # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
agent_pool: mi300_1
|
||||
parallelism: 4
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/lora
|
||||
- tests/lora
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s lora --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --ignore=lora/test_chatglm3_tp.py --ignore=lora/test_llama_tp.py --ignore=lora/test_qwen3_with_multi_loras.py --ignore=lora/test_olmoe_tp.py --ignore=lora/test_deepseekv2_tp.py --ignore=lora/test_gptoss_tp.py --ignore=lora/test_qwen3moe_tp.py --ignore=lora/test_qwen35_densemodel_lora.py
|
||||
|
||||
- label: LoRA TP (Distributed) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
@@ -1778,10 +1781,9 @@ steps:
|
||||
- tests/models/multimodal/generation
|
||||
- tests/models/multimodal/test_mapping.py
|
||||
commands:
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/AndreasKaratzas/mamba@rocm-7.0-v2.3.0'
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
- pytest -v -s models/language/generation -m hybrid_model --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB
|
||||
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal/generation -m 'not core_model' --ignore models/multimodal/generation/test_common.py
|
||||
- pytest -v -s models/multimodal/test_mapping.py
|
||||
|
||||
- label: Multi-Modal Models (Extended Generation 2) # TBD
|
||||
timeout_in_minutes: 180
|
||||
@@ -1793,9 +1795,8 @@ steps:
|
||||
- vllm/
|
||||
- tests/models/multimodal/generation
|
||||
commands:
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/AndreasKaratzas/mamba@rocm-7.0-v2.3.0'
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
- pytest -v -s models/language/generation -m '(not core_model) and (not hybrid_model)'
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal/generation/test_common.py -m 'split(group=0) and not core_model'
|
||||
|
||||
|
||||
- label: Multi-Modal Models (Extended Generation 3) # TBD
|
||||
@@ -2754,7 +2755,7 @@ steps:
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- python3 benchmarks/attention_benchmarks/benchmark.py --backends ROCM_ATTN ROCM_AITER_FA ROCM_AITER_UNIFIED_ATTN --batch-specs "8q1s1k" --repeats 1 --warmup-iters 1
|
||||
- python3 benchmarks/attention_benchmarks/benchmark.py --backends ROCM_ATTN ROCM_AITER_FA ROCM_AITER_UNIFIED_ATTN --batch-specs "8q1s1k"
|
||||
|
||||
#-------------------------------------------------------- mi355 · distributed --------------------------------------------------------#
|
||||
|
||||
|
||||
@@ -23,4 +23,4 @@ steps:
|
||||
- benchmarks/attention_benchmarks/
|
||||
- vllm/v1/attention/
|
||||
commands:
|
||||
- python3 benchmarks/attention_benchmarks/benchmark.py --backends flash flashinfer --batch-specs "8q1s1k" --repeats 1 --warmup-iters 1
|
||||
- python3 benchmarks/attention_benchmarks/benchmark.py --backends flash flashinfer --batch-specs "8q1s1k"
|
||||
|
||||
@@ -12,6 +12,17 @@ steps:
|
||||
commands:
|
||||
- pytest -v -s lora --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --ignore=lora/test_chatglm3_tp.py --ignore=lora/test_llama_tp.py --ignore=lora/test_qwen3_with_multi_loras.py --ignore=lora/test_olmoe_tp.py --ignore=lora/test_deepseekv2_tp.py --ignore=lora/test_gptoss_tp.py --ignore=lora/test_qwen3moe_tp.py --ignore=lora/test_qwen35_densemodel_lora.py
|
||||
parallelism: 4
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
timeout_in_minutes: 60
|
||||
source_file_dependencies:
|
||||
- vllm/lora
|
||||
- tests/lora
|
||||
- vllm/platforms/rocm.py
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
|
||||
- label: LoRA TP (Distributed)
|
||||
|
||||
@@ -21,6 +21,12 @@ steps:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
# TODO: create another `optional` test group for slow tests
|
||||
- pytest -v -s -m 'not slow_test' v1/spec_decode
|
||||
mirror:
|
||||
amd:
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 65
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: V1 Sample + Logits
|
||||
key: v1-sample-logits
|
||||
|
||||
@@ -30,7 +30,6 @@ steps:
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal/generation/test_common.py -m core_model -k "qwen3 or gemma"
|
||||
- pytest -v -s models/multimodal/generation/test_qwen2_5_vl.py -m core_model
|
||||
- pytest -v -s models/multimodal/generation/test_vit_cudagraph.py -m core_model
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
@@ -63,9 +62,16 @@ steps:
|
||||
- tests/models/multimodal
|
||||
commands:
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/generation/test_vit_cudagraph.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal/generation/test_vit_cudagraph.py -m core_model
|
||||
- pytest models/multimodal/generation/test_memory_leak.py -m core_model
|
||||
- cd .. && VLLM_WORKER_MULTIPROC_METHOD=spawn pytest -v -s tests/models/multimodal/generation/test_whisper.py -m core_model # Otherwise, mp_method="spawn" doesn't work
|
||||
mirror:
|
||||
amd:
|
||||
soft_fail: true
|
||||
device: mi325_1
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Multi-Modal Processor (CPU)
|
||||
key: multi-modal-processor-cpu
|
||||
|
||||
@@ -27,6 +27,10 @@ steps:
|
||||
- pip install -e ./plugins/bge_m3_sparse_plugin
|
||||
- pytest -v -s plugins_tests/test_bge_m3_sparse_io_processor_plugins.py
|
||||
- pip uninstall bge_m3_sparse_plugin -y
|
||||
# test colbert_query io_processor plugin
|
||||
- pip install -e ./plugins/colbert_query_plugin
|
||||
- pytest -v -s plugins_tests/test_colbert_query_io_processor_plugins.py
|
||||
- pip uninstall colbert_query_plugin -y
|
||||
# end io_processor plugins test
|
||||
# begin stat_logger plugins test
|
||||
- pip install -e ./plugins/vllm_add_dummy_stat_logger
|
||||
|
||||
@@ -99,9 +99,13 @@ steps:
|
||||
- vllm/v1/engine/
|
||||
- vllm/v1/worker/
|
||||
- tests/utils.py
|
||||
- tests/v1/distributed/test_external_lb_dp.py
|
||||
- tests/v1/distributed/test_hybrid_lb_dp.py
|
||||
- tests/v1/distributed/test_internal_lb_dp.py
|
||||
commands:
|
||||
- export VLLM_USE_RUST_FRONTEND=1
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- export NCCL_CUMEM_HOST_ENABLE=0
|
||||
- TP_SIZE=1 DP_SIZE=4 pytest -v -s v1/distributed/test_internal_lb_dp.py -k "not 4 and not server_info"
|
||||
- TP_SIZE=1 DP_SIZE=2 pytest -v -s v1/distributed/test_external_lb_dp.py -k "not 4 and not server_info"
|
||||
- TP_SIZE=1 DP_SIZE=4 pytest -v -s v1/distributed/test_hybrid_lb_dp.py -k "not 4 and not server_info"
|
||||
|
||||
@@ -120,16 +120,6 @@
|
||||
/vllm/model_executor/models/transformers @hmellor
|
||||
/tests/models/test_transformers.py @hmellor
|
||||
|
||||
# Observability
|
||||
/vllm/config/observability.py @markmc
|
||||
/vllm/v1/metrics @markmc
|
||||
/tests/v1/metrics @markmc
|
||||
/vllm/tracing.py @markmc
|
||||
/tests/v1/tracing/test_tracing.py @markmc
|
||||
/vllm/config/kv_events.py @markmc
|
||||
/vllm/distributed/kv_events.py @markmc
|
||||
/tests/distributed/test_events.py @markmc
|
||||
|
||||
# Docs
|
||||
/docs/mkdocs @hmellor
|
||||
/docs/**/*.yml @hmellor
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# Custom self-hosted runner labels (e.g. the autoscaling vllm-runners pool) so
|
||||
# actionlint doesn't flag them as unknown in `runs-on`.
|
||||
self-hosted-runner:
|
||||
labels:
|
||||
- vllm-runners
|
||||
@@ -388,9 +388,13 @@ pull_request_rules:
|
||||
- or:
|
||||
- files~=^tests/tool_use/
|
||||
- files~=^tests/tool_parsers/
|
||||
- files~=^tests/parser/
|
||||
- files~=^tests/reasoning/
|
||||
- files~=^tests/entrypoints/openai/.*tool.*
|
||||
- files~=^tests/entrypoints/anthropic/.*tool.*
|
||||
- files~=^vllm/tool_parsers/
|
||||
- files~=^vllm/parser/
|
||||
- files~=^vllm/reasoning/
|
||||
- files=docs/features/tool_calling.md
|
||||
- files~=^examples/tool_calling/
|
||||
actions:
|
||||
|
||||
@@ -46,12 +46,16 @@ jobs:
|
||||
pre-commit:
|
||||
needs: pre-run-check
|
||||
if: always() && (needs.pre-run-check.result == 'success' || needs.pre-run-check.result == 'skipped')
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: [self-hosted, linux, x64, vllm-runners]
|
||||
steps:
|
||||
- uses: actions/checkout@8e8c483db84b4bee98b60c0593521ed34d9990e8 # v6.0.1
|
||||
- uses: actions/setup-python@83679a892e2d95755f2dac6acb0bfd1e9ac5d548 # v6.1.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
# Provide shellcheck on PATH so tools/pre_commit/shellcheck.sh skips its
|
||||
# wget + tar -xJ self-download, which the self-hosted runner image lacks
|
||||
# (no wget/xz). Pinned to shellcheck 0.10.0 to match the script's "stable".
|
||||
- run: python -m pip install shellcheck-py==0.10.0.1
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/markdownlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
|
||||
@@ -363,35 +363,6 @@ if(VLLM_GPU_LANG STREQUAL "CUDA")
|
||||
SRCS "${VLLM_EXT_SRC}"
|
||||
CUDA_ARCHS "${CUDA_ARCHS}")
|
||||
|
||||
# Expert-specialization MXFP8 blockscaled grouped kernels (SM100+).
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(ES_MXFP8_GROUPED_MM_ARCHS "10.0f;11.0f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(ES_MXFP8_GROUPED_MM_ARCHS "10.0a;10.1a;10.3a" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND ES_MXFP8_GROUPED_MM_ARCHS)
|
||||
set(ES_MXFP8_GROUPED_MM_SRCS
|
||||
"csrc/libtorch_stable/moe/mxfp8_moe/cutlass_mxfp8_grouped_mm.cu"
|
||||
"csrc/libtorch_stable/moe/mxfp8_moe/mxfp8_experts_quant.cu")
|
||||
set_gencode_flags_for_srcs(
|
||||
SRCS "${ES_MXFP8_GROUPED_MM_SRCS}"
|
||||
CUDA_ARCHS "${ES_MXFP8_GROUPED_MM_ARCHS}")
|
||||
list(APPEND VLLM_STABLE_EXT_SRC "${ES_MXFP8_GROUPED_MM_SRCS}")
|
||||
list(APPEND VLLM_GPU_FLAGS "-DENABLE_ES_MXFP8_GROUPED_MM_SM100=1")
|
||||
message(STATUS "Building ES MXFP8 grouped kernels for archs: ${ES_MXFP8_GROUPED_MM_ARCHS}")
|
||||
else()
|
||||
if (NOT ${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8
|
||||
AND ES_MXFP8_GROUPED_MM_ARCHS)
|
||||
message(STATUS "Not building ES MXFP8 grouped kernels as CUDA Compiler version is "
|
||||
"not >= 12.8.")
|
||||
else()
|
||||
message(STATUS "Not building ES MXFP8 grouped kernels as no compatible archs found "
|
||||
"in CUDA target architectures.")
|
||||
endif()
|
||||
endif()
|
||||
|
||||
|
||||
|
||||
# if CUDA endif
|
||||
endif()
|
||||
|
||||
|
||||
@@ -53,6 +53,16 @@ from common import (
|
||||
from vllm.v1.worker.workspace import init_workspace_manager
|
||||
|
||||
|
||||
def _str2bool(v) -> bool:
|
||||
if isinstance(v, bool):
|
||||
return v
|
||||
if v.lower() in ("true", "1", "yes", "t"):
|
||||
return True
|
||||
if v.lower() in ("false", "0", "no", "f"):
|
||||
return False
|
||||
raise argparse.ArgumentTypeError(f"expected a boolean, got {v!r}")
|
||||
|
||||
|
||||
def run_standard_attention_benchmark(config: BenchmarkConfig) -> BenchmarkResult:
|
||||
"""Run standard attention benchmark (Flash/Triton/FlashInfer)."""
|
||||
from runner import run_attention_benchmark
|
||||
@@ -485,6 +495,20 @@ def main():
|
||||
help="Prefill backends to compare (fa2, fa3, fa4). "
|
||||
"Uses the first decode backend for impl construction.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fp8-output-scale",
|
||||
type=float,
|
||||
help="Static per-tensor scale enabling the MLA prefill FP8-output "
|
||||
"comparison on FA4 (fused write vs standalone post-quant).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fuse-quant-op",
|
||||
nargs="+",
|
||||
type=_str2bool,
|
||||
help="FP8-output write path(s) to run: false = bf16 attention + "
|
||||
"standalone static-FP8 quant, true = FA4 writes FP8 directly. "
|
||||
"Default: both.",
|
||||
)
|
||||
|
||||
# Batch specifications
|
||||
parser.add_argument(
|
||||
@@ -618,6 +642,12 @@ def main():
|
||||
# Prefill backends (e.g., ["fa3", "fa4"])
|
||||
args.prefill_backends = yaml_config.get("prefill_backends", None)
|
||||
|
||||
# FP8 output benchmark knobs; CLI wins.
|
||||
if args.fp8_output_scale is None:
|
||||
args.fp8_output_scale = yaml_config.get("fp8_output_scale", None)
|
||||
if args.fuse_quant_op is None:
|
||||
args.fuse_quant_op = yaml_config.get("fuse_quant_op", None)
|
||||
|
||||
# Check for special modes
|
||||
args.mode = yaml_config.get("mode", None)
|
||||
|
||||
@@ -787,8 +817,59 @@ def main():
|
||||
"skipped (timings are placeholder zeros).[/]"
|
||||
)
|
||||
|
||||
# FA4 fused FP8 output vs standalone post-quant, on the same fa4 kernel:
|
||||
# the delta is the post-quant kernel the fused path removes.
|
||||
fp8_output_scale = getattr(args, "fp8_output_scale", None)
|
||||
if fp8_output_scale is not None:
|
||||
decode_backend = backends[0]
|
||||
fuse_variants = args.fuse_quant_op or [False, True]
|
||||
label_of = {False: "post_quant", True: "fused"}
|
||||
console.print(
|
||||
f"[yellow]FP8 output comparison @ scale={fp8_output_scale} "
|
||||
f"(prefill=fa4, decode impl={decode_backend})[/]"
|
||||
)
|
||||
fp8_results = []
|
||||
total = len(fuse_variants) * len(args.batch_specs)
|
||||
with tqdm(total=total, desc="FP8 output benchmarking") as pbar:
|
||||
for spec in args.batch_specs:
|
||||
for fuse in fuse_variants:
|
||||
config = BenchmarkConfig(
|
||||
backend=decode_backend,
|
||||
batch_spec=spec,
|
||||
num_layers=args.num_layers,
|
||||
head_dim=args.head_dim,
|
||||
num_q_heads=args.num_q_heads,
|
||||
num_kv_heads=args.num_kv_heads,
|
||||
block_size=args.block_size,
|
||||
device=args.device,
|
||||
repeats=args.repeats,
|
||||
warmup_iters=args.warmup_iters,
|
||||
profile_memory=args.profile_memory,
|
||||
kv_cache_dtype=args.kv_cache_dtype,
|
||||
use_cuda_graphs=args.cuda_graphs,
|
||||
prefill_backend="fa4",
|
||||
)
|
||||
result = run_benchmark(
|
||||
config, output_scale=fp8_output_scale, fuse_quant_op=fuse
|
||||
)
|
||||
label = label_of[fuse]
|
||||
labeled_config = replace(result.config, backend=label)
|
||||
result = replace(result, config=labeled_config)
|
||||
fp8_results.append(result)
|
||||
|
||||
if not result.success:
|
||||
console.print(f"[red]Error {label} {spec}: {result.error}[/]")
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
console.print("\n[bold green]FP8 Output Results:[/]")
|
||||
formatter = ResultsFormatter(console)
|
||||
labels = [label_of[f] for f in fuse_variants]
|
||||
formatter.print_table(fp8_results, labels, compare_to_fastest=True)
|
||||
all_results = fp8_results
|
||||
|
||||
# Handle special mode: decode_vs_prefill comparison
|
||||
if hasattr(args, "mode") and args.mode == "decode_vs_prefill":
|
||||
elif hasattr(args, "mode") and args.mode == "decode_vs_prefill":
|
||||
console.print("[yellow]Mode: Decode vs Prefill pipeline comparison[/]")
|
||||
console.print(
|
||||
"[dim]For each query length, testing both decode and prefill pipelines[/]"
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# MLA prefill FP8-output microbenchmark (FA4).
|
||||
# Compares the fused FP8 write against bf16 attention + a standalone static-FP8
|
||||
# quant; the delta is the post-quant kernel the fused path removes.
|
||||
# DeepSeek-Coder-V2-Lite dims; FA4 needs SM100/110.
|
||||
#
|
||||
# Usage:
|
||||
# python benchmark.py --config configs/mla_fa4_fp8_output.yaml
|
||||
|
||||
description: "MLA prefill FA4 fused-FP8 output vs post-quant"
|
||||
|
||||
model:
|
||||
name: "deepseek-v2-lite"
|
||||
num_layers: 27
|
||||
num_q_heads: 16
|
||||
num_kv_heads: 1
|
||||
head_dim: 576
|
||||
kv_lora_rank: 512
|
||||
qk_nope_head_dim: 128
|
||||
qk_rope_head_dim: 64
|
||||
v_head_dim: 128
|
||||
block_size: 128
|
||||
|
||||
# Pure prefill (q_len == kv_len) so every token goes through forward_mha.
|
||||
batch_specs:
|
||||
- "q512"
|
||||
- "q1k"
|
||||
- "q2k"
|
||||
- "q4k"
|
||||
- "q8k"
|
||||
- "2q4k"
|
||||
- "4q4k"
|
||||
- "8q4k"
|
||||
|
||||
# Only used to construct the MLA impl; the pure-prefill specs skip decode.
|
||||
decode_backends:
|
||||
- CUTLASS_MLA
|
||||
|
||||
# Sweep the two FP8 write paths (prefill backend is fixed to fa4).
|
||||
fp8_output_scale: 0.1
|
||||
fuse_quant_op: [false, true]
|
||||
|
||||
device: "cuda:0"
|
||||
repeats: 50
|
||||
warmup_iters: 10
|
||||
@@ -708,6 +708,8 @@ def _run_single_benchmark(
|
||||
device: torch.device,
|
||||
indexer=None,
|
||||
kv_cache_dtype: str | None = None,
|
||||
output_scale: float | None = None,
|
||||
fuse_quant_op: bool = False,
|
||||
) -> BenchmarkResult:
|
||||
"""
|
||||
Run a single benchmark iteration.
|
||||
@@ -721,6 +723,11 @@ def _run_single_benchmark(
|
||||
mla_dims: MLA dimension configuration
|
||||
device: Target device
|
||||
indexer: Optional MockIndexer for sparse backends
|
||||
output_scale: Static per-tensor FP8 scale for prefill output. None
|
||||
keeps the plain bf16 output (no quantization).
|
||||
fuse_quant_op: With output_scale set, True lets the prefill kernel write
|
||||
FP8 directly; False runs bf16 attention then a standalone static-FP8
|
||||
quant. The delta isolates the saved post-quant kernel.
|
||||
|
||||
Returns:
|
||||
BenchmarkResult with timing statistics
|
||||
@@ -824,23 +831,55 @@ def _run_single_benchmark(
|
||||
num_prefill, mla_dims, query_fmt, device, torch.bfloat16
|
||||
)
|
||||
|
||||
# Prefill FP8 output: fused (kernel writes e4m3) vs separate post-quant.
|
||||
prefill_fp8_output = None
|
||||
prefill_output_scale = None
|
||||
prefill_quant_op = None
|
||||
if has_prefill and output_scale is not None:
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
prefill_output_scale = torch.tensor(
|
||||
[output_scale], device=device, dtype=torch.float32
|
||||
)
|
||||
if fuse_quant_op:
|
||||
prefill_fp8_output = torch.empty_like(
|
||||
prefill_inputs["output"], dtype=current_platform.fp8_dtype()
|
||||
)
|
||||
else:
|
||||
from vllm.model_executor.layers.quantization.input_quant_fp8 import (
|
||||
QuantFP8,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
GroupShape,
|
||||
)
|
||||
|
||||
prefill_quant_op = QuantFP8(static=True, group_shape=GroupShape.PER_TENSOR)
|
||||
|
||||
fused_output = output_scale is not None and fuse_quant_op
|
||||
|
||||
# Build forward function (runs a single decode/prefill pass)
|
||||
def forward_fn():
|
||||
results = []
|
||||
if has_decode:
|
||||
results.append(impl.forward_mqa(decode_inputs, kv_cache, metadata, layer))
|
||||
if has_prefill:
|
||||
results.append(
|
||||
impl.forward_mha(
|
||||
prefill_inputs["q"],
|
||||
prefill_inputs["k_c_normed"],
|
||||
prefill_inputs["k_pe"],
|
||||
kv_cache,
|
||||
metadata,
|
||||
prefill_inputs["k_scale"],
|
||||
prefill_inputs["output"],
|
||||
)
|
||||
out = impl.forward_mha(
|
||||
prefill_inputs["q"],
|
||||
prefill_inputs["k_c_normed"],
|
||||
prefill_inputs["k_pe"],
|
||||
kv_cache,
|
||||
metadata,
|
||||
prefill_inputs["k_scale"],
|
||||
prefill_fp8_output if fused_output else prefill_inputs["output"],
|
||||
prefill_output_scale if fused_output else None,
|
||||
)
|
||||
if fused_output:
|
||||
out = prefill_fp8_output
|
||||
elif prefill_quant_op is not None:
|
||||
out, _ = prefill_quant_op(
|
||||
prefill_inputs["output"], prefill_output_scale
|
||||
)
|
||||
results.append(out)
|
||||
return results[0] if len(results) == 1 else tuple(results)
|
||||
|
||||
def benchmark_fn():
|
||||
@@ -881,6 +920,8 @@ def _run_mla_benchmark_batched(
|
||||
configs_with_params: list[tuple], # [(config, threshold, num_splits), ...]
|
||||
index_topk: int = 2048,
|
||||
prefill_backend: str | None = None,
|
||||
output_scale: float | None = None,
|
||||
fuse_quant_op: bool = False,
|
||||
) -> list[BenchmarkResult]:
|
||||
"""
|
||||
Unified batched MLA benchmark runner for all backends.
|
||||
@@ -1020,6 +1061,8 @@ def _run_mla_benchmark_batched(
|
||||
device,
|
||||
indexer=indexer,
|
||||
kv_cache_dtype=kv_cache_dtype,
|
||||
output_scale=output_scale,
|
||||
fuse_quant_op=fuse_quant_op,
|
||||
)
|
||||
results.append(result)
|
||||
|
||||
@@ -1047,6 +1090,8 @@ def run_mla_benchmark(
|
||||
num_kv_splits: int | None = None,
|
||||
index_topk: int = 2048,
|
||||
prefill_backend: str | None = None,
|
||||
output_scale: float | None = None,
|
||||
fuse_quant_op: bool = False,
|
||||
) -> BenchmarkResult | list[BenchmarkResult]:
|
||||
"""
|
||||
Unified MLA benchmark runner for all backends.
|
||||
@@ -1066,6 +1111,9 @@ def run_mla_benchmark(
|
||||
index_topk: Topk value for sparse MLA backends (default 2048)
|
||||
prefill_backend: Prefill backend name (e.g., "fa3", "fa4").
|
||||
When set, forces the specified FlashAttention version for prefill.
|
||||
output_scale: Static per-tensor FP8 scale for prefill output (None = bf16).
|
||||
fuse_quant_op: With output_scale set, fuse the FP8 write into the prefill
|
||||
kernel vs a standalone post-quant kernel. See _run_single_benchmark.
|
||||
|
||||
Returns:
|
||||
BenchmarkResult (single mode) or list of BenchmarkResult (batched mode)
|
||||
@@ -1090,7 +1138,12 @@ def run_mla_benchmark(
|
||||
|
||||
# Use unified batched execution
|
||||
results = _run_mla_benchmark_batched(
|
||||
backend, configs_with_params, index_topk, prefill_backend=prefill_backend
|
||||
backend,
|
||||
configs_with_params,
|
||||
index_topk,
|
||||
prefill_backend=prefill_backend,
|
||||
output_scale=output_scale,
|
||||
fuse_quant_op=fuse_quant_op,
|
||||
)
|
||||
|
||||
# Return single result or list based on input
|
||||
|
||||
@@ -0,0 +1,358 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Benchmark and regression-test pinned (page-locked) CPU memory for vLLM.
|
||||
|
||||
Verifies that enabling pinned memory does not regress throughput or latency
|
||||
compared to unpinned memory. Each condition runs in an isolated ``spawn``
|
||||
subprocess so both start from a cold CUDA context, giving an unbiased
|
||||
comparison.
|
||||
|
||||
Usage
|
||||
-----
|
||||
Run all tests with the default model::
|
||||
|
||||
python benchmarks/benchmark_pin_memory.py -v
|
||||
|
||||
Override the model and optional max-model-len::
|
||||
|
||||
python benchmarks/benchmark_pin_memory.py --model unsloth/Qwen3-1.7B -v
|
||||
python benchmarks/benchmark_pin_memory.py --model unsloth/Qwen3-1.7B \
|
||||
--max-model-len 8192 -v
|
||||
|
||||
Run only throughput or latency tests::
|
||||
|
||||
python benchmarks/benchmark_pin_memory.py -v -k test_throughput
|
||||
python benchmarks/benchmark_pin_memory.py -v -k test_latency
|
||||
|
||||
Run only the v1 or v2 runner variant::
|
||||
|
||||
python benchmarks/benchmark_pin_memory.py -v -k v1
|
||||
python benchmarks/benchmark_pin_memory.py -v -k v2
|
||||
|
||||
Note: on WSL2, v1 runner tests are skipped because pin memory is not available
|
||||
for the v1 runner without cpu_offload_gb. Run on other platforms to exercise v1.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import multiprocessing
|
||||
import sys
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
|
||||
# Allow up to 2% degradation. Both benchmark runs start from an identical
|
||||
# cold CUDA context (separate spawn subprocesses), so the measured difference
|
||||
# reflects the genuine pin_memory overhead rather than cold/warm ordering bias.
|
||||
_THROUGHPUT_TOLERANCE = 0.98
|
||||
_THROUGHPUT_NUM_REQUESTS = 200
|
||||
_THROUGHPUT_INPUT_LEN = 128
|
||||
_THROUGHPUT_OUTPUT_LEN = 512
|
||||
_THROUGHPUT_MAX_NUM_SEQS = 128
|
||||
|
||||
# Latency benchmark constants — match latency.py defaults.
|
||||
_LATENCY_TOLERANCE = 1.02 # Allow up to 2% latency regression.
|
||||
_LATENCY_BATCH_SIZE = 64
|
||||
_LATENCY_INPUT_LEN = 32
|
||||
_LATENCY_OUTPUT_LEN = 128
|
||||
_LATENCY_WARMUP_ITERS = 5
|
||||
_LATENCY_BENCH_ITERS = 15
|
||||
|
||||
_DEFAULT_MODEL = "unsloth/Qwen3-1.7B"
|
||||
_DEFAULT_MAX_MODEL_LEN = 16384
|
||||
|
||||
|
||||
def _benchmark_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(add_help=False)
|
||||
parser.add_argument("--model", default=_DEFAULT_MODEL)
|
||||
parser.add_argument("--max-model-len", type=int, default=_DEFAULT_MAX_MODEL_LEN)
|
||||
args, _ = parser.parse_known_args()
|
||||
return args
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model() -> str:
|
||||
return _benchmark_args().model
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def max_model_len() -> int:
|
||||
return _benchmark_args().max_model_len
|
||||
|
||||
|
||||
def _skip_if_pin_memory_not_available(engine_args_kwargs: dict) -> None:
|
||||
"""Skip the current pytest test if pin_memory is unavailable for this config."""
|
||||
import vllm.utils.platform_utils as pu
|
||||
from vllm.config import set_current_vllm_config
|
||||
from vllm.engine.arg_utils import EngineArgs
|
||||
|
||||
vllm_config = EngineArgs(**engine_args_kwargs).create_engine_config()
|
||||
with set_current_vllm_config(vllm_config):
|
||||
pu.is_pin_memory_available.cache_clear()
|
||||
if not pu.is_pin_memory_available():
|
||||
import os
|
||||
|
||||
runner = "v2" if os.environ.get("VLLM_USE_V2_MODEL_RUNNER") == "1" else "v1"
|
||||
model = engine_args_kwargs.get("model", "unknown")
|
||||
print(
|
||||
f"\033[33mSKIP: pin_memory not available for "
|
||||
f"{runner} runner, model={model}\033[0m"
|
||||
)
|
||||
pytest.skip("pin_memory not available for this configuration")
|
||||
|
||||
|
||||
def _throughput_worker(
|
||||
pin: bool,
|
||||
engine_args_kwargs: dict,
|
||||
q: "multiprocessing.Queue[float]",
|
||||
v2_mode: bool = False,
|
||||
) -> None:
|
||||
"""Run throughput benchmark in a fresh spawn subprocess.
|
||||
|
||||
Delegates to vllm/benchmarks/throughput.py main() using the random dataset,
|
||||
so the methodology matches the official benchmark. Results are written to a
|
||||
temp JSON file and forwarded through the queue as tokens/s.
|
||||
|
||||
v2_mode: when True, monkeypatches is_uva_available() to always return True
|
||||
so the v2 model runner's UVA buffers remain functional even when pin=False.
|
||||
This isolates the non-UVA pin_memory paths in v2.
|
||||
"""
|
||||
import vllm.utils.platform_utils as pu
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
pu.is_pin_memory_available.cache_clear()
|
||||
pu.is_uva_available.cache_clear()
|
||||
type(current_platform).is_pin_memory_available = classmethod(lambda cls: pin)
|
||||
if v2_mode:
|
||||
pu.is_uva_available = lambda: True
|
||||
|
||||
from vllm.benchmarks.throughput import add_cli_args
|
||||
from vllm.benchmarks.throughput import main as throughput_main
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
add_cli_args(parser)
|
||||
args = parser.parse_args([])
|
||||
|
||||
for key, val in engine_args_kwargs.items():
|
||||
setattr(args, key, val)
|
||||
args.max_num_seqs = _THROUGHPUT_MAX_NUM_SEQS
|
||||
args.dataset_name = "random"
|
||||
args.input_len = _THROUGHPUT_INPUT_LEN
|
||||
args.output_len = _THROUGHPUT_OUTPUT_LEN
|
||||
# Nullify defaults that conflict with explicit input/output_len.
|
||||
args.random_input_len = None
|
||||
args.random_output_len = None
|
||||
args.random_prefix_len = None
|
||||
args.num_prompts = _THROUGHPUT_NUM_REQUESTS
|
||||
args.seed = 0
|
||||
args.disable_detokenize = True
|
||||
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
|
||||
tmp_path = f.name
|
||||
args.output_json = tmp_path
|
||||
|
||||
throughput_main(args)
|
||||
|
||||
with open(tmp_path) as f:
|
||||
results = json.load(f)
|
||||
q.put(results["tokens_per_second"])
|
||||
|
||||
|
||||
def _run_throughput_benchmark(
|
||||
pin: bool,
|
||||
engine_args_kwargs: dict,
|
||||
v2_mode: bool = False,
|
||||
) -> float:
|
||||
ctx = multiprocessing.get_context("spawn")
|
||||
q = ctx.Queue()
|
||||
p = ctx.Process(
|
||||
target=_throughput_worker,
|
||||
args=(pin, engine_args_kwargs, q, v2_mode),
|
||||
)
|
||||
p.start()
|
||||
p.join()
|
||||
if p.exitcode != 0:
|
||||
raise RuntimeError(
|
||||
f"Throughput benchmark subprocess (pin={pin}) exited with code {p.exitcode}"
|
||||
)
|
||||
return q.get()
|
||||
|
||||
|
||||
def _latency_worker(
|
||||
pin: bool,
|
||||
engine_args_kwargs: dict,
|
||||
q: "multiprocessing.Queue[dict]",
|
||||
v2_mode: bool = False,
|
||||
) -> None:
|
||||
"""Run latency benchmark in a fresh spawn subprocess.
|
||||
|
||||
Follows latency.py methodology: fixed batch of dummy token IDs, warmup
|
||||
iterations to reach steady state, then timed iterations reduced to avg
|
||||
and percentiles. Results are written to a temp JSON file by latency_main
|
||||
and forwarded through the queue.
|
||||
"""
|
||||
import vllm.utils.platform_utils as pu
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
pu.is_pin_memory_available.cache_clear()
|
||||
pu.is_uva_available.cache_clear()
|
||||
type(current_platform).is_pin_memory_available = classmethod(lambda cls: pin)
|
||||
if v2_mode:
|
||||
pu.is_uva_available = lambda: True
|
||||
|
||||
from vllm.benchmarks.latency import add_cli_args
|
||||
from vllm.benchmarks.latency import main as latency_main
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
add_cli_args(parser)
|
||||
args = parser.parse_args([])
|
||||
|
||||
for key, val in engine_args_kwargs.items():
|
||||
setattr(args, key, val)
|
||||
args.input_len = _LATENCY_INPUT_LEN
|
||||
args.output_len = _LATENCY_OUTPUT_LEN
|
||||
args.batch_size = _LATENCY_BATCH_SIZE
|
||||
args.num_iters_warmup = _LATENCY_WARMUP_ITERS
|
||||
args.num_iters = _LATENCY_BENCH_ITERS
|
||||
args.profile = False
|
||||
args.disable_detokenize = True
|
||||
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
|
||||
tmp_path = f.name
|
||||
args.output_json = tmp_path
|
||||
|
||||
latency_main(args)
|
||||
|
||||
with open(tmp_path) as f:
|
||||
results = json.load(f)
|
||||
q.put(results)
|
||||
|
||||
|
||||
def _run_latency_benchmark(
|
||||
pin: bool,
|
||||
engine_args_kwargs: dict,
|
||||
v2_mode: bool = False,
|
||||
) -> dict:
|
||||
ctx = multiprocessing.get_context("spawn")
|
||||
q = ctx.Queue()
|
||||
p = ctx.Process(
|
||||
target=_latency_worker,
|
||||
args=(pin, engine_args_kwargs, q, v2_mode),
|
||||
)
|
||||
p.start()
|
||||
p.join()
|
||||
if p.exitcode != 0:
|
||||
raise RuntimeError(
|
||||
f"Latency benchmark subprocess (pin={pin}) exited with code {p.exitcode}"
|
||||
)
|
||||
return q.get()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"test_v2_runner",
|
||||
[
|
||||
pytest.param(False, id="v1"),
|
||||
pytest.param(True, id="v2"),
|
||||
],
|
||||
)
|
||||
class TestPinnedMemory:
|
||||
"""Verify pinned memory yields >= throughput vs unpinned via real vLLM inference."""
|
||||
|
||||
def test_throughput(self, monkeypatch, test_v2_runner, model, max_model_len):
|
||||
"""Benchmark throughput with pin_memory forced on then off.
|
||||
|
||||
Delegates to vllm/benchmarks/throughput.py main() with the random
|
||||
dataset. Each condition runs in an isolated spawn subprocess so both
|
||||
start from a cold CUDA context, giving an unbiased comparison.
|
||||
"""
|
||||
monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
|
||||
monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "1" if test_v2_runner else "0")
|
||||
|
||||
engine_args_kwargs = dict(
|
||||
model=model,
|
||||
gpu_memory_utilization=0.88,
|
||||
max_model_len=max_model_len,
|
||||
enable_prefix_caching=False,
|
||||
)
|
||||
|
||||
_skip_if_pin_memory_not_available(engine_args_kwargs)
|
||||
|
||||
unpinned_tps = _run_throughput_benchmark(
|
||||
False, engine_args_kwargs, v2_mode=test_v2_runner
|
||||
)
|
||||
pinned_tps = _run_throughput_benchmark(
|
||||
True, engine_args_kwargs, v2_mode=test_v2_runner
|
||||
)
|
||||
|
||||
pct_diff = (pinned_tps - unpinned_tps) / unpinned_tps * 100
|
||||
runner = "v2" if test_v2_runner else "v1"
|
||||
print(
|
||||
f"\n=== Throughput results ({runner} runner, {model}) ==="
|
||||
f"\npin_memory=True: {pinned_tps:.1f} tok/s"
|
||||
f"\npin_memory=False: {unpinned_tps:.1f} tok/s"
|
||||
f"\nDifference: {pct_diff:+.1f}% (pinned vs unpinned)"
|
||||
)
|
||||
|
||||
assert pinned_tps >= unpinned_tps * _THROUGHPUT_TOLERANCE, (
|
||||
f"Pinned throughput ({pinned_tps:.1f} tok/s) fell more than "
|
||||
f"{(1.0 - _THROUGHPUT_TOLERANCE) * 100:.1f}% below "
|
||||
f"unpinned ({unpinned_tps:.1f} tok/s)."
|
||||
)
|
||||
|
||||
def test_latency(self, monkeypatch, test_v2_runner, model, max_model_len):
|
||||
"""Benchmark per-batch latency with pin_memory forced on then off.
|
||||
|
||||
Follows vllm/benchmarks/latency.py: fixed dummy-token batch, warmup
|
||||
iterations to reach steady state, then timed iterations reduced to avg
|
||||
and percentiles. Subprocesses run serially so each gets a cold CUDA
|
||||
context without GPU memory pressure from the other run.
|
||||
"""
|
||||
monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
|
||||
monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "1" if test_v2_runner else "0")
|
||||
|
||||
engine_args_kwargs = dict(
|
||||
model=model,
|
||||
gpu_memory_utilization=0.88,
|
||||
max_model_len=max_model_len,
|
||||
enable_prefix_caching=False,
|
||||
)
|
||||
|
||||
_skip_if_pin_memory_not_available(engine_args_kwargs)
|
||||
|
||||
unpinned = _run_latency_benchmark(
|
||||
False, engine_args_kwargs, v2_mode=test_v2_runner
|
||||
)
|
||||
pinned = _run_latency_benchmark(
|
||||
True, engine_args_kwargs, v2_mode=test_v2_runner
|
||||
)
|
||||
|
||||
pct_diff = (
|
||||
(pinned["avg_latency"] - unpinned["avg_latency"])
|
||||
/ unpinned["avg_latency"]
|
||||
* 100
|
||||
)
|
||||
runner = "v2" if test_v2_runner else "v1"
|
||||
print(
|
||||
f"\n=== Latency results ({runner} runner, {model}) ==="
|
||||
f"\npin_memory=True: avg={pinned['avg_latency']:.3f}s"
|
||||
f" p50={pinned['percentiles']['50']:.3f}s"
|
||||
f" p99={pinned['percentiles']['99']:.3f}s"
|
||||
f"\npin_memory=False: avg={unpinned['avg_latency']:.3f}s"
|
||||
f" p50={unpinned['percentiles']['50']:.3f}s"
|
||||
f" p99={unpinned['percentiles']['99']:.3f}s"
|
||||
f"\nDifference: {pct_diff:+.1f}% (pinned vs unpinned)"
|
||||
)
|
||||
|
||||
assert pinned["avg_latency"] <= unpinned["avg_latency"] * _LATENCY_TOLERANCE, (
|
||||
f"Pinned avg latency ({pinned['avg_latency']:.3f}s) exceeded "
|
||||
f"unpinned ({unpinned['avg_latency']:.3f}s) by more than "
|
||||
f"{(_LATENCY_TOLERANCE - 1.0) * 100:.1f}%."
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
_parser = argparse.ArgumentParser(add_help=False)
|
||||
_parser.add_argument("--model", default=_DEFAULT_MODEL)
|
||||
_parser.add_argument("--max-model-len", type=int, default=_DEFAULT_MAX_MODEL_LEN)
|
||||
_, _remaining = _parser.parse_known_args()
|
||||
sys.exit(pytest.main([__file__] + _remaining))
|
||||
@@ -32,19 +32,17 @@ message(STATUS "fmha_sm100 is available at ${fmha_sm100_SOURCE_DIR}")
|
||||
|
||||
add_custom_target(fmha_sm100)
|
||||
|
||||
set(FMHA_SM100_PY_ROOT "${fmha_sm100_SOURCE_DIR}/python/fmha_sm100")
|
||||
|
||||
install(FILES
|
||||
"${fmha_sm100_SOURCE_DIR}/python/fmha_sm100/__init__.py"
|
||||
"${fmha_sm100_SOURCE_DIR}/python/fmha_sm100/sparse.py"
|
||||
"${FMHA_SM100_PY_ROOT}/__init__.py"
|
||||
"${FMHA_SM100_PY_ROOT}/sparse.py"
|
||||
DESTINATION vllm/third_party/fmha_sm100
|
||||
COMPONENT fmha_sm100)
|
||||
|
||||
install(DIRECTORY "${fmha_sm100_SOURCE_DIR}/python/fmha_sm100/cute/"
|
||||
install(DIRECTORY "${FMHA_SM100_PY_ROOT}/cute/"
|
||||
DESTINATION vllm/third_party/fmha_sm100/cute
|
||||
COMPONENT fmha_sm100
|
||||
FILES_MATCHING
|
||||
REGEX "/__pycache__(/.*)?$" EXCLUDE
|
||||
REGEX ".*\\.pyc$" EXCLUDE
|
||||
PATTERN "example.py" EXCLUDE
|
||||
PATTERN "test_*.py" EXCLUDE
|
||||
PATTERN "*.py"
|
||||
PATTERN "build_k2q_csr.cu")
|
||||
PATTERN "__pycache__" EXCLUDE
|
||||
PATTERN "*.pyc" EXCLUDE
|
||||
PATTERN ".git*" EXCLUDE)
|
||||
|
||||
+25
-25
@@ -11,13 +11,25 @@ static inline cpu_attention::Fp8KVCacheDataType parse_fp8_kv_dtype(
|
||||
return cpu_attention::Fp8KVCacheDataType::kAuto;
|
||||
}
|
||||
|
||||
bool cpu_attn_has_isa(const std::string& isa) {
|
||||
if (isa == "rvv") {
|
||||
#if defined(__riscv) && defined(__riscv_v_min_vlen) && __riscv_v_min_vlen == 128
|
||||
return true;
|
||||
#else
|
||||
return false;
|
||||
#endif
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
torch::Tensor get_scheduler_metadata(
|
||||
const int64_t num_req, const int64_t num_heads_q,
|
||||
const int64_t num_heads_kv, const int64_t head_dim,
|
||||
const torch::Tensor& seq_lens, at::ScalarType dtype,
|
||||
const torch::Tensor& query_start_loc, const bool casual,
|
||||
const torch::Tensor& query_start_loc, const bool causal,
|
||||
const int64_t window_size, const std::string& isa_hint,
|
||||
const bool enable_kv_split) {
|
||||
const bool enable_kv_split,
|
||||
const std::optional<torch::Tensor>& dynamic_causal) {
|
||||
cpu_attention::ISA isa;
|
||||
if (isa_hint == "amx") {
|
||||
isa = cpu_attention::ISA::AMX;
|
||||
@@ -44,24 +56,13 @@ torch::Tensor get_scheduler_metadata(
|
||||
input.head_dim = head_dim;
|
||||
input.query_start_loc = query_start_loc.data_ptr<int32_t>();
|
||||
input.seq_lens = seq_lens.data_ptr<int32_t>();
|
||||
if (window_size != -1) {
|
||||
input.left_sliding_window_size = window_size - 1;
|
||||
if (casual) {
|
||||
input.right_sliding_window_size = 0;
|
||||
} else {
|
||||
input.right_sliding_window_size = window_size - 1;
|
||||
}
|
||||
} else {
|
||||
input.left_sliding_window_size = -1;
|
||||
if (casual) {
|
||||
input.right_sliding_window_size = 0;
|
||||
} else {
|
||||
input.right_sliding_window_size = -1;
|
||||
}
|
||||
}
|
||||
input.casual = casual;
|
||||
|
||||
input.sliding_window_size = window_size;
|
||||
input.causal = causal;
|
||||
input.isa = isa;
|
||||
input.enable_kv_split = enable_kv_split;
|
||||
input.dynamic_causal =
|
||||
dynamic_causal.has_value() ? dynamic_causal->data_ptr<bool>() : nullptr;
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(dtype, "get_scheduler_metadata", [&]() {
|
||||
CPU_ATTN_DISPATCH(head_dim, isa, 0, [&]() {
|
||||
@@ -175,10 +176,11 @@ void cpu_attention_with_kv_cache(
|
||||
const torch::Tensor& seq_lens, // [num_tokens]
|
||||
const double scale, const bool causal,
|
||||
const std::optional<torch::Tensor>& alibi_slopes, // [num_heads]
|
||||
const int64_t sliding_window_left, const int64_t sliding_window_right,
|
||||
const int64_t sliding_window,
|
||||
const torch::Tensor& block_table, // [num_tokens, max_block_num]
|
||||
const double softcap, const torch::Tensor& scheduler_metadata,
|
||||
const std::optional<torch::Tensor>& s_aux, // [num_heads]
|
||||
const std::optional<torch::Tensor>& s_aux, // [num_heads]
|
||||
const std::optional<torch::Tensor>& dynamic_causal, // [num_reqs]
|
||||
const double k_scale = 1.0, const double v_scale = 1.0,
|
||||
const std::string& kv_cache_dtype = "auto") {
|
||||
TORCH_CHECK_EQ(query.dim(), 3);
|
||||
@@ -220,13 +222,11 @@ void cpu_attention_with_kv_cache(
|
||||
input.alibi_slopes =
|
||||
alibi_slopes.has_value() ? alibi_slopes->data_ptr<float>() : nullptr;
|
||||
input.s_aux = s_aux.has_value() ? s_aux->data_ptr<c10::BFloat16>() : nullptr;
|
||||
input.dynamic_causal =
|
||||
dynamic_causal.has_value() ? dynamic_causal->data_ptr<bool>() : nullptr;
|
||||
input.scale = scale;
|
||||
input.causal = causal;
|
||||
input.sliding_window_left = sliding_window_left;
|
||||
input.sliding_window_right = sliding_window_right;
|
||||
if (input.causal) {
|
||||
input.sliding_window_right = 0;
|
||||
}
|
||||
input.sliding_window_size = sliding_window;
|
||||
input.softcap = static_cast<float>(softcap);
|
||||
|
||||
if (is_fp8) {
|
||||
|
||||
+62
-27
@@ -388,13 +388,13 @@ class AttentionScheduler {
|
||||
int32_t head_dim;
|
||||
int32_t* query_start_loc;
|
||||
int32_t* seq_lens;
|
||||
int32_t left_sliding_window_size;
|
||||
int32_t right_sliding_window_size;
|
||||
bool casual;
|
||||
int32_t sliding_window_size;
|
||||
bool causal;
|
||||
cpu_attention::ISA isa;
|
||||
int32_t max_num_q_per_iter; // max Q head num can be hold in registers
|
||||
int32_t kv_block_alignment; // context length alignment requirement
|
||||
bool enable_kv_split;
|
||||
bool* dynamic_causal;
|
||||
};
|
||||
|
||||
static constexpr int32_t MaxQTileIterNum = 128;
|
||||
@@ -403,7 +403,8 @@ class AttentionScheduler {
|
||||
: available_cache_size_(cpu_utils::get_available_l2_size()) {}
|
||||
|
||||
torch::Tensor schedule(const ScheduleInput& input) const {
|
||||
const bool casual = input.casual;
|
||||
const bool causal = input.causal;
|
||||
const bool is_dynamic_causal = input.dynamic_causal != nullptr;
|
||||
const int32_t thread_num = omp_get_max_threads();
|
||||
const int64_t cache_size = cpu_utils::get_available_l2_size();
|
||||
const int32_t max_num_q_per_iter = input.max_num_q_per_iter;
|
||||
@@ -434,8 +435,7 @@ class AttentionScheduler {
|
||||
const int32_t default_tile_token_num = default_tile_size / q_head_per_kv;
|
||||
const int32_t split_kv_q_token_num_threshold =
|
||||
input.enable_kv_split ? 1 : 0;
|
||||
const int32_t left_sliding_window_size = input.left_sliding_window_size;
|
||||
const int32_t right_sliding_window_size = input.right_sliding_window_size;
|
||||
const int32_t sliding_window_size = input.sliding_window_size;
|
||||
TORCH_CHECK_LE(split_kv_q_token_num_threshold * q_head_per_kv, 16);
|
||||
|
||||
// get total kv len
|
||||
@@ -444,7 +444,9 @@ class AttentionScheduler {
|
||||
const int32_t seq_len = input.seq_lens[req_id];
|
||||
const int32_t q_token_num =
|
||||
input.query_start_loc[req_id + 1] - input.query_start_loc[req_id];
|
||||
const int32_t q_start_pos = (casual ? (seq_len - q_token_num) : 0);
|
||||
const bool req_causal =
|
||||
is_dynamic_causal ? input.dynamic_causal[req_id] : causal;
|
||||
const int32_t q_start_pos = seq_len - q_token_num;
|
||||
const int32_t kv_start_pos = 0;
|
||||
const int32_t kv_end_pos = seq_len;
|
||||
|
||||
@@ -456,7 +458,7 @@ class AttentionScheduler {
|
||||
const int32_t q_tile_pos_right = q_tile_pos_left + q_tile_token_num;
|
||||
const auto [kv_tile_pos_left, kv_tile_pos_right] = calcu_kv_tile_pos(
|
||||
kv_start_pos, kv_end_pos, q_tile_pos_left, q_tile_pos_right,
|
||||
left_sliding_window_size, right_sliding_window_size);
|
||||
sliding_window_size, req_causal);
|
||||
const auto [aligned_kv_tile_pos_left, aligned_kv_tile_pos_right] =
|
||||
align_kv_tile_pos(kv_tile_pos_left, kv_tile_pos_right,
|
||||
kv_len_alignment);
|
||||
@@ -484,7 +486,9 @@ class AttentionScheduler {
|
||||
const int32_t seq_len = input.seq_lens[req_id];
|
||||
const int32_t q_token_num =
|
||||
input.query_start_loc[req_id + 1] - input.query_start_loc[req_id];
|
||||
const int32_t q_start_pos = (casual ? (seq_len - q_token_num) : 0);
|
||||
const bool req_causal =
|
||||
is_dynamic_causal ? input.dynamic_causal[req_id] : causal;
|
||||
const int32_t q_start_pos = seq_len - q_token_num;
|
||||
const int32_t kv_start_pos = 0;
|
||||
const int32_t kv_end_pos = seq_len;
|
||||
int32_t local_split_id = 0;
|
||||
@@ -498,7 +502,7 @@ class AttentionScheduler {
|
||||
const int32_t q_tile_pos_right = q_tile_pos_left + q_tile_token_num;
|
||||
const auto [kv_tile_pos_left, kv_tile_pos_right] = calcu_kv_tile_pos(
|
||||
kv_start_pos, kv_end_pos, q_tile_pos_left, q_tile_pos_right,
|
||||
left_sliding_window_size, right_sliding_window_size);
|
||||
sliding_window_size, req_causal);
|
||||
const auto [aligned_kv_tile_pos_left, aligned_kv_tile_pos_right] =
|
||||
align_kv_tile_pos(kv_tile_pos_left, kv_tile_pos_right,
|
||||
kv_len_alignment);
|
||||
@@ -708,15 +712,41 @@ class AttentionScheduler {
|
||||
return metadata_tensor;
|
||||
}
|
||||
|
||||
FORCE_INLINE static std::pair<int32_t, int32_t> calcu_sliding_window_size(
|
||||
int32_t window_size, bool causal) {
|
||||
int32_t left_sliding_window_size, right_sliding_window_size;
|
||||
if (window_size != -1) {
|
||||
left_sliding_window_size = window_size - 1;
|
||||
if (causal) {
|
||||
right_sliding_window_size = 0;
|
||||
} else {
|
||||
right_sliding_window_size = window_size - 1;
|
||||
}
|
||||
} else {
|
||||
left_sliding_window_size = -1;
|
||||
if (causal) {
|
||||
right_sliding_window_size = 0;
|
||||
} else {
|
||||
right_sliding_window_size = -1;
|
||||
}
|
||||
}
|
||||
|
||||
return {left_sliding_window_size, right_sliding_window_size};
|
||||
}
|
||||
|
||||
FORCE_INLINE static std::pair<int32_t, int32_t> calcu_kv_tile_pos(
|
||||
int32_t kv_left_pos, int32_t kv_right_pos, int32_t q_left_pos,
|
||||
int32_t q_right_pos, int32_t sliding_window_left,
|
||||
int32_t sliding_window_right) {
|
||||
if (sliding_window_left != -1) {
|
||||
kv_left_pos = std::max(kv_left_pos, q_left_pos - sliding_window_left);
|
||||
int32_t q_right_pos, int32_t window_size, bool causal) {
|
||||
auto [left_sliding_window_size, right_sliding_window_size] =
|
||||
calcu_sliding_window_size(window_size, causal);
|
||||
|
||||
if (left_sliding_window_size != -1) {
|
||||
kv_left_pos =
|
||||
std::max(kv_left_pos, q_left_pos - left_sliding_window_size);
|
||||
}
|
||||
if (sliding_window_right != -1) {
|
||||
kv_right_pos = std::min(kv_right_pos, q_right_pos + sliding_window_right);
|
||||
if (right_sliding_window_size != -1) {
|
||||
kv_right_pos =
|
||||
std::min(kv_right_pos, q_right_pos + right_sliding_window_size);
|
||||
}
|
||||
return {kv_left_pos, kv_right_pos};
|
||||
}
|
||||
@@ -805,10 +835,10 @@ struct AttentionInput {
|
||||
int32_t* block_table;
|
||||
float* alibi_slopes;
|
||||
c10::BFloat16* s_aux;
|
||||
bool* dynamic_causal;
|
||||
float scale;
|
||||
bool causal;
|
||||
int32_t sliding_window_left;
|
||||
int32_t sliding_window_right;
|
||||
int32_t sliding_window_size;
|
||||
float softcap;
|
||||
// FP8 KV cache scales (used by FP8 attention implementations)
|
||||
float k_scale_fp8 = 1.0f;
|
||||
@@ -1442,15 +1472,16 @@ class AttentionMainLoop {
|
||||
const int64_t q_head_num_stride = input->query_num_heads_stride;
|
||||
const int64_t kv_cache_head_num_stride = input->cache_num_kv_heads_stride;
|
||||
const int64_t kv_cache_block_num_stride = input->cache_num_blocks_stride;
|
||||
const int32_t sliding_window_left = input->sliding_window_left;
|
||||
const int32_t sliding_window_right = input->sliding_window_right;
|
||||
const int32_t sliding_window_size = input->sliding_window_size;
|
||||
const int32_t block_size = input->block_size;
|
||||
const float scale = input->scale;
|
||||
const float softcap_scale = input->softcap;
|
||||
const float* alibi_slopes = input->alibi_slopes;
|
||||
const c10::BFloat16* s_aux = input->s_aux;
|
||||
const bool* dynamic_causal = input->dynamic_causal;
|
||||
const bool is_dynamic_causal = dynamic_causal != nullptr;
|
||||
|
||||
const bool casual = input->causal;
|
||||
const bool causal = input->causal;
|
||||
int32_t* const block_table = input->block_table;
|
||||
const int64_t block_table_stride = input->blt_num_tokens_stride;
|
||||
|
||||
@@ -1533,6 +1564,11 @@ class AttentionMainLoop {
|
||||
&curr_workitem_groups[workitem_group_idx];
|
||||
|
||||
const int32_t current_group_idx = current_workitem_group->req_id;
|
||||
const int32_t current_group_causal =
|
||||
is_dynamic_causal ? dynamic_causal[current_group_idx] : causal;
|
||||
auto [sliding_window_left, sliding_window_right] =
|
||||
AttentionScheduler::calcu_sliding_window_size(
|
||||
sliding_window_size, current_group_causal);
|
||||
const int32_t kv_start_pos =
|
||||
current_workitem_group->kv_split_pos_start;
|
||||
const int32_t kv_end_pos = current_workitem_group->kv_split_pos_end;
|
||||
@@ -1560,8 +1596,7 @@ class AttentionMainLoop {
|
||||
const int32_t q_end = input->query_start_loc[current_group_idx + 1];
|
||||
const int32_t q_start = input->query_start_loc[current_group_idx];
|
||||
const int32_t seq_len = input->seq_lens[current_group_idx];
|
||||
const int32_t q_start_pos =
|
||||
(casual ? seq_len - (q_end - q_start) : 0);
|
||||
const int32_t q_start_pos = seq_len - (q_end - q_start);
|
||||
const int32_t block_num = (seq_len + block_size - 1) / block_size;
|
||||
// Only apply sink for the first KV split
|
||||
bool use_sink = (s_aux != nullptr &&
|
||||
@@ -1611,8 +1646,8 @@ class AttentionMainLoop {
|
||||
const auto [kv_tile_start_pos, kv_tile_end_pos] =
|
||||
AttentionScheduler::calcu_kv_tile_pos(
|
||||
kv_start_pos, kv_end_pos, q_tile_start_pos,
|
||||
q_tile_end_pos, sliding_window_left,
|
||||
sliding_window_right);
|
||||
q_tile_end_pos, sliding_window_size,
|
||||
current_group_causal);
|
||||
const auto [rounded_kv_tile_start_pos, rounded_kv_tile_end_pos] =
|
||||
AttentionScheduler::align_kv_tile_pos(
|
||||
kv_tile_start_pos, kv_tile_end_pos, blocksize_alignment);
|
||||
@@ -1725,8 +1760,8 @@ class AttentionMainLoop {
|
||||
actual_kv_tile_pos_right] =
|
||||
AttentionScheduler::calcu_kv_tile_pos(
|
||||
kv_tile_pos_left, kv_tile_pos_right, q_tile_pos_left,
|
||||
q_tile_pos_right, sliding_window_left,
|
||||
sliding_window_right);
|
||||
q_tile_pos_right, sliding_window_size,
|
||||
current_group_causal);
|
||||
const int32_t q_iter_idx =
|
||||
q_head_tile_token_offset / curr_max_q_token_num_per_iter;
|
||||
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
#include <sleef.h>
|
||||
|
||||
#include "cpu/cpu_types.hpp"
|
||||
#include "cpu/utils.hpp"
|
||||
#include "cpu/micro_gemm/cpu_micro_gemm_vec.hpp"
|
||||
@@ -163,7 +165,6 @@ void gelu_tanh_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
vec_op::FP32Vec16 w1_vec(0.7978845608028654);
|
||||
vec_op::FP32Vec16 w2_vec(0.5);
|
||||
vec_op::FP32Vec16 w3_vec(0.044715);
|
||||
alignas(64) float temp[16];
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < dim; n += 16) {
|
||||
@@ -171,12 +172,9 @@ void gelu_tanh_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
vec_op::FP32Vec16 up_vec(up + n);
|
||||
auto gate_pow3_vec = gate_vec * gate_vec * gate_vec;
|
||||
auto inner_vec = w1_vec * (gate_vec + w3_vec * gate_pow3_vec);
|
||||
|
||||
inner_vec.save(temp);
|
||||
for (int32_t i = 0; i < 16; ++i) {
|
||||
temp[i] = std::tanh(temp[i]);
|
||||
}
|
||||
vec_op::FP32Vec16 tanh_vec(temp);
|
||||
// Note: can't use fast_exp form because diffusiongemma will generate
|
||||
// wrong results
|
||||
vec_op::FP32Vec16 tanh_vec(Sleef_tanhf16_u10(inner_vec.reg));
|
||||
auto gelu_tanh = gate_vec * w2_vec * (one_vec + tanh_vec);
|
||||
auto gated_output_fp32 = up_vec * gelu_tanh;
|
||||
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
|
||||
|
||||
+41
-22
@@ -4,8 +4,9 @@ namespace {
|
||||
template <typename scalar_t>
|
||||
void rms_norm_impl(scalar_t* __restrict__ out,
|
||||
const scalar_t* __restrict__ input,
|
||||
const scalar_t* __restrict__ weight, const float epsilon,
|
||||
const int num_tokens, const int hidden_size) {
|
||||
const scalar_t* __restrict__ weight, const bool has_weight,
|
||||
const float epsilon, const int num_tokens,
|
||||
const int hidden_size) {
|
||||
using scalar_vec_t = vec_op::vec_t<scalar_t>;
|
||||
constexpr int VEC_ELEM_NUM = scalar_vec_t::get_elem_num();
|
||||
TORCH_CHECK(hidden_size % VEC_ELEM_NUM == 0);
|
||||
@@ -27,12 +28,15 @@ void rms_norm_impl(scalar_t* __restrict__ out,
|
||||
|
||||
for (int j = 0; j < hidden_size; j += VEC_ELEM_NUM) {
|
||||
scalar_vec_t x(input_p + j);
|
||||
scalar_vec_t w(weight + j);
|
||||
|
||||
vec_op::FP32Vec8 fp32_x(x);
|
||||
vec_op::FP32Vec8 fp32_w(w);
|
||||
|
||||
vec_op::FP32Vec8 fp32_out = fp32_x * fp32_s_variance * fp32_w;
|
||||
vec_op::FP32Vec8 fp32_out;
|
||||
if (has_weight) {
|
||||
scalar_vec_t w(weight + j);
|
||||
vec_op::FP32Vec8 fp32_w(w);
|
||||
fp32_out = fp32_x * fp32_s_variance * fp32_w;
|
||||
} else {
|
||||
fp32_out = fp32_x * fp32_s_variance;
|
||||
}
|
||||
|
||||
scalar_vec_t out(fp32_out);
|
||||
out.save(output_p + j);
|
||||
@@ -44,8 +48,8 @@ template <typename scalar_t>
|
||||
void fused_add_rms_norm_impl(scalar_t* __restrict__ input,
|
||||
scalar_t* __restrict__ residual,
|
||||
const scalar_t* __restrict__ weight,
|
||||
const float epsilon, const int num_tokens,
|
||||
const int hidden_size) {
|
||||
const bool has_weight, const float epsilon,
|
||||
const int num_tokens, const int hidden_size) {
|
||||
using scalar_vec_t = vec_op::vec_t<scalar_t>;
|
||||
constexpr int VEC_ELEM_NUM = scalar_vec_t::get_elem_num();
|
||||
TORCH_CHECK(hidden_size % VEC_ELEM_NUM == 0);
|
||||
@@ -72,13 +76,18 @@ void fused_add_rms_norm_impl(scalar_t* __restrict__ input,
|
||||
vec_op::FP32Vec8 fp32_s_variance(s_variance);
|
||||
|
||||
for (int j = 0; j < hidden_size; j += VEC_ELEM_NUM) {
|
||||
scalar_vec_t w(weight + j);
|
||||
scalar_vec_t res(residual_p + j);
|
||||
|
||||
vec_op::FP32Vec8 fp32_w(w);
|
||||
vec_op::FP32Vec8 fp32_res(res);
|
||||
|
||||
vec_op::FP32Vec8 fp32_out = fp32_res * fp32_s_variance * fp32_w;
|
||||
vec_op::FP32Vec8 fp32_out;
|
||||
if (has_weight) {
|
||||
scalar_vec_t w(weight + j);
|
||||
scalar_vec_t res(residual_p + j);
|
||||
vec_op::FP32Vec8 fp32_w(w);
|
||||
vec_op::FP32Vec8 fp32_res(res);
|
||||
fp32_out = fp32_res * fp32_s_variance * fp32_w;
|
||||
} else {
|
||||
scalar_vec_t res(residual_p + j);
|
||||
vec_op::FP32Vec8 fp32_res(res);
|
||||
fp32_out = fp32_res * fp32_s_variance;
|
||||
}
|
||||
|
||||
scalar_vec_t out(fp32_out);
|
||||
out.save(input_p + j);
|
||||
@@ -87,31 +96,41 @@ void fused_add_rms_norm_impl(scalar_t* __restrict__ input,
|
||||
}
|
||||
} // namespace
|
||||
|
||||
void rms_norm(torch::Tensor& out, torch::Tensor& input, torch::Tensor& weight,
|
||||
double epsilon) {
|
||||
void rms_norm(torch::Tensor& out, torch::Tensor& input,
|
||||
std::optional<torch::Tensor> weight, double epsilon) {
|
||||
int hidden_size = input.size(-1);
|
||||
int num_tokens = input.numel() / hidden_size;
|
||||
const bool has_weight = weight.has_value();
|
||||
if (has_weight) {
|
||||
TORCH_CHECK(weight->is_contiguous());
|
||||
}
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(input.scalar_type(), "rms_norm_impl", [&] {
|
||||
CPU_KERNEL_GUARD_IN(rms_norm_impl)
|
||||
rms_norm_impl(out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(),
|
||||
weight.data_ptr<scalar_t>(), epsilon, num_tokens,
|
||||
hidden_size);
|
||||
has_weight ? weight->data_ptr<scalar_t>() : nullptr,
|
||||
has_weight, epsilon, num_tokens, hidden_size);
|
||||
CPU_KERNEL_GUARD_OUT(rms_norm_impl)
|
||||
});
|
||||
}
|
||||
|
||||
void fused_add_rms_norm(torch::Tensor& input, torch::Tensor& residual,
|
||||
torch::Tensor& weight, double epsilon) {
|
||||
std::optional<torch::Tensor> weight, double epsilon) {
|
||||
int hidden_size = input.size(-1);
|
||||
int num_tokens = input.numel() / hidden_size;
|
||||
const bool has_weight = weight.has_value();
|
||||
if (has_weight) {
|
||||
TORCH_CHECK(weight->scalar_type() == input.scalar_type());
|
||||
TORCH_CHECK(weight->is_contiguous());
|
||||
}
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(
|
||||
input.scalar_type(), "fused_add_rms_norm_impl", [&] {
|
||||
CPU_KERNEL_GUARD_IN(fused_add_rms_norm_impl)
|
||||
fused_add_rms_norm_impl(
|
||||
input.data_ptr<scalar_t>(), residual.data_ptr<scalar_t>(),
|
||||
weight.data_ptr<scalar_t>(), epsilon, num_tokens, hidden_size);
|
||||
has_weight ? weight->data_ptr<scalar_t>() : nullptr, has_weight,
|
||||
epsilon, num_tokens, hidden_size);
|
||||
CPU_KERNEL_GUARD_OUT(fused_add_rms_norm_impl)
|
||||
});
|
||||
}
|
||||
|
||||
+15
-10
@@ -146,13 +146,16 @@ at::Tensor causal_conv1d_update_cpu(
|
||||
void activation_lut_bf16(torch::Tensor& out, torch::Tensor& input,
|
||||
const std::string& activation);
|
||||
|
||||
bool cpu_attn_has_isa(const std::string& isa);
|
||||
|
||||
torch::Tensor get_scheduler_metadata(
|
||||
const int64_t num_req, const int64_t num_heads_q,
|
||||
const int64_t num_heads_kv, const int64_t head_dim,
|
||||
const torch::Tensor& seq_lens, at::ScalarType dtype,
|
||||
const torch::Tensor& query_start_loc, const bool casual,
|
||||
const int64_t window_size, const std::string& isa_hint,
|
||||
const bool enable_kv_split);
|
||||
const bool enable_kv_split,
|
||||
const std::optional<torch::Tensor>& dynamic_causal);
|
||||
|
||||
void cpu_attn_reshape_and_cache(const torch::Tensor& key,
|
||||
const torch::Tensor& value,
|
||||
@@ -169,10 +172,10 @@ void cpu_attention_with_kv_cache(
|
||||
const torch::Tensor& query_start_loc, const torch::Tensor& seq_lens,
|
||||
const double scale, const bool causal,
|
||||
const std::optional<torch::Tensor>& alibi_slopes,
|
||||
const int64_t sliding_window_left, const int64_t sliding_window_right,
|
||||
const torch::Tensor& block_table, const double softcap,
|
||||
const torch::Tensor& scheduler_metadata,
|
||||
const std::optional<torch::Tensor>& s_aux, const double k_scale,
|
||||
const int64_t sliding_window_left, const torch::Tensor& block_table,
|
||||
const double softcap, const torch::Tensor& scheduler_metadata,
|
||||
const std::optional<torch::Tensor>& s_aux,
|
||||
const std::optional<torch::Tensor>& dynamic_causal, const double k_scale,
|
||||
const double v_scale, const std::string& kv_cache_dtype);
|
||||
|
||||
// Note: just for avoiding importing errors
|
||||
@@ -309,13 +312,13 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
// Layernorm
|
||||
// Apply Root Mean Square (RMS) Normalization to the input tensor.
|
||||
ops.def(
|
||||
"rms_norm(Tensor! out, Tensor input, Tensor weight, float epsilon) -> "
|
||||
"rms_norm(Tensor! out, Tensor input, Tensor? weight, float epsilon) -> "
|
||||
"()");
|
||||
ops.impl("rms_norm", torch::kCPU, &rms_norm);
|
||||
|
||||
// In-place fused Add and RMS Normalization.
|
||||
ops.def(
|
||||
"fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor weight, "
|
||||
"fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor? weight, "
|
||||
"float epsilon) -> ()");
|
||||
ops.impl("fused_add_rms_norm", torch::kCPU, &fused_add_rms_norm);
|
||||
|
||||
@@ -496,11 +499,12 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
ops.impl("fused_gdn_gating_cpu", torch::kCPU, &fused_gdn_gating_cpu);
|
||||
|
||||
// CPU attention kernels
|
||||
ops.def("cpu_attn_has_isa(str isa) -> bool", &cpu_attn_has_isa);
|
||||
ops.def(
|
||||
"get_scheduler_metadata(int num_req, int num_heads_q, int num_heads_kv, "
|
||||
"int head_dim, Tensor seq_lens, ScalarType dtype, Tensor "
|
||||
"query_start_loc, bool casual, int window_size, str isa_hint, bool "
|
||||
"enable_kv_split) -> Tensor",
|
||||
"enable_kv_split, Tensor? dynamic_causal) -> Tensor",
|
||||
&get_scheduler_metadata);
|
||||
ops.def(
|
||||
"cpu_attn_reshape_and_cache(Tensor key, Tensor value, Tensor(a2!) "
|
||||
@@ -512,8 +516,9 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"cpu_attention_with_kv_cache(Tensor query, Tensor key_cache, Tensor "
|
||||
"value_cache, Tensor(a3!) output, Tensor query_start_loc, Tensor "
|
||||
"seq_lens, float scale, bool causal, Tensor? alibi_slopes, SymInt "
|
||||
"sliding_window_left, SymInt sliding_window_right, Tensor block_table, "
|
||||
"float softcap, Tensor scheduler_metadata, Tensor? s_aux, "
|
||||
"sliding_window_size, Tensor block_table, "
|
||||
"float softcap, Tensor scheduler_metadata, Tensor? s_aux, Tensor? "
|
||||
"dynamic_causal, "
|
||||
"float k_scale=1.0, float v_scale=1.0, str kv_cache_dtype=\"auto\") -> "
|
||||
"()",
|
||||
&cpu_attention_with_kv_cache);
|
||||
|
||||
@@ -18,7 +18,7 @@
|
||||
* ROPE_DIM = 64 (RoPE applied to dims [NOPE_DIM, HEAD_DIM))
|
||||
* NOPE_DIM = 448
|
||||
* QUANT_BLOCK = 64 (UE8M0 FP8 quant block)
|
||||
* FP8_MAX = 448.0f
|
||||
* FP8_MAX = 224.0f on ROCm FNUZ / 448.0f on OCP
|
||||
* is_neox=false (GPT-J interleaved pairs)
|
||||
* cos_sin_cache layout [max_pos, rope_dim] = cos || sin (cos first, sin
|
||||
* second along last dim; each half is rope_dim/2 = 32 values)
|
||||
@@ -61,10 +61,11 @@
|
||||
#ifdef USE_ROCM
|
||||
// ROCm-compatible FP8 conversion helpers
|
||||
__device__ __forceinline__ uint8_t rocm_cvt_float_to_fp8_e4m3(float val) {
|
||||
#if defined(HIP_FP8_TYPE_OCP)
|
||||
__hip_fp8_e4m3 fp8_val(val);
|
||||
#else
|
||||
// gfx942 uses FNUZ FP8; other ROCm targets use OCP E4M3.
|
||||
#if defined(__gfx942__)
|
||||
__hip_fp8_e4m3_fnuz fp8_val(val);
|
||||
#else
|
||||
__hip_fp8_e4m3 fp8_val(val);
|
||||
#endif
|
||||
return reinterpret_cast<uint8_t&>(fp8_val);
|
||||
}
|
||||
@@ -90,7 +91,13 @@ constexpr int kQuantBlock = 64;
|
||||
constexpr int kNumQuantBlocks = kNopeDim / kQuantBlock; // 7
|
||||
constexpr int kScaleBytesPerToken = kNumQuantBlocks + 1; // 8 (7 real + 1 pad)
|
||||
constexpr int kTokenDataBytes = kNopeDim + kRopeDim * 2; // 448 + 128 = 576
|
||||
// FNUZ on gfx942 / OCP elsewhere. FNUZ uses 224.0 (not the dtype's raw
|
||||
// 240.0) to match the rest of vLLM's FNUZ pipeline.
|
||||
#if defined(USE_ROCM) && defined(__gfx942__)
|
||||
constexpr float kFp8Max = 224.0f;
|
||||
#else
|
||||
constexpr float kFp8Max = 448.0f;
|
||||
#endif
|
||||
|
||||
#ifndef USE_ROCM
|
||||
// When num_tokens is less than this threshold,
|
||||
|
||||
@@ -58,8 +58,15 @@
|
||||
|
||||
#include "../cuda_compat.h"
|
||||
#include "../type_convert.cuh"
|
||||
#include "../attention/dtype_fp8.cuh"
|
||||
#include "dispatch_utils.h"
|
||||
|
||||
#ifdef USE_ROCM
|
||||
#include "../quantization/w8a8/fp8/amd/quant_utils.cuh"
|
||||
#else
|
||||
#include "../quantization/w8a8/fp8/nvidia/quant_utils.cuh"
|
||||
#endif
|
||||
|
||||
#ifndef FINAL_MASK
|
||||
#ifdef USE_ROCM
|
||||
#define FINAL_MASK 0xffffffffffffffffULL
|
||||
@@ -186,6 +193,21 @@ __device__ __forceinline__ void storeElems(
|
||||
*reinterpret_cast<uint2*>(dst) = v;
|
||||
}
|
||||
|
||||
template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
|
||||
__device__ __forceinline__ void storeCacheElems(
|
||||
cache_t* __restrict__ dst, float const (&elems)[kElemsPerLane]) {
|
||||
if constexpr (kv_dt == Fp8KVCacheDataType::kAuto) {
|
||||
// kAuto means unquantized KV cache here: cache_t == scalar_t, so store the
|
||||
// model dtype directly. FP8 cache dtypes use the conversion path below.
|
||||
storeElems<scalar_t>(reinterpret_cast<scalar_t*>(dst), elems);
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kElemsPerLane; i++) {
|
||||
dst[i] = fp8::scaled_convert<cache_t, float, kv_dt>(elems[i], 1.0f);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ────────────────────────────────────────────────────────────────────────────
|
||||
// Kernel
|
||||
// ────────────────────────────────────────────────────────────────────────────
|
||||
@@ -202,7 +224,8 @@ __device__ __forceinline__ void storeElems(
|
||||
// V : nkv only if kInsertKV (V-cache insert; no warps in dense)
|
||||
// IQ: niq only if kIsSparse (norm+RoPE)
|
||||
// IK: 1 only if kIsSparse (norm+RoPE; +index-cache insert)
|
||||
template <typename scalar_t, bool kIsSparse, bool kInsertKV>
|
||||
template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt,
|
||||
bool kIsSparse, bool kInsertKV>
|
||||
__global__ void fusedMiniMaxM3QNormRopeKVInsertKernel(
|
||||
scalar_t* __restrict__ qkv, // [N, qkv_row] in/out (packs index if sparse)
|
||||
scalar_t* __restrict__ q_out, // [N, nq*128] contiguous, or nullptr
|
||||
@@ -215,7 +238,7 @@ __global__ void fusedMiniMaxM3QNormRopeKVInsertKernel(
|
||||
int64_t const* __restrict__ positions, // [N] i64
|
||||
int64_t const* __restrict__ slot_mapping, // main K/V slots or nullptr
|
||||
int64_t const* __restrict__ index_slot_mapping, // index K slots/nullptr
|
||||
scalar_t* __restrict__ kv_cache, // [nb,2,bs,nkv,128] or nullptr
|
||||
cache_t* __restrict__ kv_cache, // [nb,2,bs,nkv,128] or nullptr
|
||||
scalar_t* __restrict__ index_cache, // [nb*bs, 128] or nullptr
|
||||
float const eps, int const rotary_dim, int const num_tokens, int const nq,
|
||||
int const nkv, int const niq, int const block_size,
|
||||
@@ -355,7 +378,8 @@ __global__ void fusedMiniMaxM3QNormRopeKVInsertKernel(
|
||||
int const kv = isK ? 0 : 1;
|
||||
int64_t const off =
|
||||
b * kv_s_block + kv * kv_s_kv + t * kv_s_token + head * kv_s_head;
|
||||
storeElems<scalar_t>(kv_cache + off + dim_base, elems);
|
||||
storeCacheElems<scalar_t, cache_t, kv_dt>(kv_cache + off + dim_base,
|
||||
elems);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -373,13 +397,13 @@ __global__ void fusedMiniMaxM3QNormRopeKVInsertKernel(
|
||||
// ────────────────────────────────────────────────────────────────────────────
|
||||
// Launch wrapper
|
||||
// ────────────────────────────────────────────────────────────────────────────
|
||||
template <typename scalar_t>
|
||||
template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
|
||||
void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out,
|
||||
scalar_t const* q_norm_w, scalar_t const* k_norm_w,
|
||||
scalar_t const* iq_norm_w, scalar_t const* ik_norm_w,
|
||||
scalar_t const* cos_sin_cache,
|
||||
int64_t const* positions, int64_t const* slot_mapping,
|
||||
int64_t const* index_slot_mapping, scalar_t* kv_cache,
|
||||
int64_t const* index_slot_mapping, cache_t* kv_cache,
|
||||
scalar_t* index_cache, float const eps,
|
||||
int const rotary_dim, int const num_tokens,
|
||||
int const nq, int const nkv, int const niq,
|
||||
@@ -419,7 +443,8 @@ void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out,
|
||||
#define LAUNCH(IS_SPARSE, INSERT) \
|
||||
cudaLaunchKernelEx( \
|
||||
&config, \
|
||||
fusedMiniMaxM3QNormRopeKVInsertKernel<scalar_t, IS_SPARSE, INSERT>, \
|
||||
fusedMiniMaxM3QNormRopeKVInsertKernel<scalar_t, cache_t, kv_dt, \
|
||||
IS_SPARSE, INSERT>, \
|
||||
qkv, q_out, index_q_out, q_norm_w, k_norm_w, iq_norm_w, ik_norm_w, \
|
||||
cos_sin_cache, positions, slot_mapping, index_slot_mapping, kv_cache, \
|
||||
index_cache, eps, rotary_dim, num_tokens, nq, nkv, niq, block_size, \
|
||||
@@ -428,7 +453,8 @@ void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out,
|
||||
// ROCm: standard kernel launch syntax (no PDL/stream serialization).
|
||||
// clang-format off
|
||||
#define LAUNCH(IS_SPARSE, INSERT) \
|
||||
fusedMiniMaxM3QNormRopeKVInsertKernel<scalar_t, IS_SPARSE, INSERT> \
|
||||
fusedMiniMaxM3QNormRopeKVInsertKernel<scalar_t, cache_t, kv_dt, \
|
||||
IS_SPARSE, INSERT> \
|
||||
<<<grid, kBlockSize, 0, stream>>>( \
|
||||
qkv, q_out, index_q_out, q_norm_w, k_norm_w, iq_norm_w, \
|
||||
ik_norm_w, cos_sin_cache, positions, slot_mapping, \
|
||||
@@ -455,6 +481,33 @@ void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out,
|
||||
} // namespace minimax_m3_fused_ops
|
||||
} // namespace vllm
|
||||
|
||||
#define CALL_FUSED_MINIMAX_M3(_RAW_T, CACHE_T, KV_DTYPE) \
|
||||
vllm::minimax_m3_fused_ops::launchFusedMiniMaxM3<st, CACHE_T, KV_DTYPE>( \
|
||||
reinterpret_cast<st*>(qkv.data_ptr()), \
|
||||
q_out.has_value() ? reinterpret_cast<st*>(q_out->data_ptr()) : nullptr, \
|
||||
index_q_out.has_value() ? reinterpret_cast<st*>(index_q_out->data_ptr()) \
|
||||
: nullptr, \
|
||||
reinterpret_cast<st const*>(q_norm_weight.data_ptr()), \
|
||||
reinterpret_cast<st const*>(k_norm_weight.data_ptr()), \
|
||||
has_index ? reinterpret_cast<st const*>(index_q_norm_weight->data_ptr()) \
|
||||
: nullptr, \
|
||||
has_index ? reinterpret_cast<st const*>(index_k_norm_weight->data_ptr()) \
|
||||
: nullptr, \
|
||||
reinterpret_cast<st const*>(cos_sin_cache.data_ptr()), \
|
||||
reinterpret_cast<int64_t const*>(positions.data_ptr()), \
|
||||
insert_kv ? reinterpret_cast<int64_t const*>(slot_mapping->data_ptr()) \
|
||||
: nullptr, \
|
||||
insert_kv ? reinterpret_cast<int64_t const*>( \
|
||||
effective_index_slot_mapping->data_ptr()) \
|
||||
: nullptr, \
|
||||
insert_kv ? reinterpret_cast<CACHE_T*>(kv_cache->data_ptr()) : nullptr, \
|
||||
(insert_kv && has_index) \
|
||||
? reinterpret_cast<st*>(index_cache->data_ptr()) \
|
||||
: nullptr, \
|
||||
static_cast<float>(eps), static_cast<int>(rotary_dim), num_tokens, nq, \
|
||||
nkv, niq, static_cast<int>(block_size), kv_s_block, kv_s_kv, kv_s_token, \
|
||||
kv_s_head, has_index, insert_kv, stream)
|
||||
|
||||
// ────────────────────────────────────────────────────────────────────────────
|
||||
// Torch op wrapper
|
||||
// ────────────────────────────────────────────────────────────────────────────
|
||||
@@ -475,9 +528,14 @@ void fused_minimax_m3_qknorm_rope_kv_insert(
|
||||
int64_t block_size,
|
||||
std::optional<torch::stable::Tensor> q_out, // [N, nq*128] contiguous
|
||||
std::optional<torch::stable::Tensor>
|
||||
index_q_out) { // [N, niq*128] contiguous
|
||||
index_q_out, // [N, niq*128] contiguous
|
||||
const std::string& kv_cache_dtype) {
|
||||
STD_TORCH_CHECK(qkv.is_cuda() && qkv.is_contiguous(),
|
||||
"qkv must be contiguous CUDA");
|
||||
STD_TORCH_CHECK(
|
||||
qkv.scalar_type() == torch::headeronly::ScalarType::Half ||
|
||||
qkv.scalar_type() == torch::headeronly::ScalarType::BFloat16,
|
||||
"qkv must be float16 or bfloat16");
|
||||
STD_TORCH_CHECK(
|
||||
positions.is_cuda() &&
|
||||
positions.scalar_type() == torch::headeronly::ScalarType::Long,
|
||||
@@ -510,6 +568,8 @@ void fused_minimax_m3_qknorm_rope_kv_insert(
|
||||
// (1 head)]) right after [q|k|v] in the same row; the dense layer does not.
|
||||
bool const has_index = niq > 0;
|
||||
bool const insert_kv = kv_cache.has_value();
|
||||
vllm::Fp8KVCacheDataType const kv_dt =
|
||||
vllm::get_fp8_kv_cache_data_type(kv_cache_dtype);
|
||||
int const kHeadDim = vllm::minimax_m3_fused_ops::kHeadDim;
|
||||
int const expected_row =
|
||||
(nq + 2 * nkv + (has_index ? niq + 1 : 0)) * kHeadDim;
|
||||
@@ -552,8 +612,14 @@ void fused_minimax_m3_qknorm_rope_kv_insert(
|
||||
torch::headeronly::ScalarType::Long &&
|
||||
index_slot_mapping->numel() == slot_mapping->numel()),
|
||||
"index_slot_mapping must be int64 CUDA with slot_mapping length");
|
||||
STD_TORCH_CHECK(kv_cache->scalar_type() == qkv.scalar_type(),
|
||||
"kv_cache dtype must match qkv (bf16 cache only)");
|
||||
if (kv_dt == vllm::Fp8KVCacheDataType::kAuto) {
|
||||
STD_TORCH_CHECK(kv_cache->scalar_type() == qkv.scalar_type(),
|
||||
"auto kv_cache dtype must match qkv");
|
||||
} else {
|
||||
STD_TORCH_CHECK(
|
||||
kv_cache->scalar_type() == torch::headeronly::ScalarType::Byte,
|
||||
"fp8 kv_cache must use uint8 storage");
|
||||
}
|
||||
STD_TORCH_CHECK(index_cache.has_value() &&
|
||||
index_cache->scalar_type() == qkv.scalar_type(),
|
||||
"insert mode requires matching index_cache");
|
||||
@@ -601,35 +667,9 @@ void fused_minimax_m3_qknorm_rope_kv_insert(
|
||||
VLLM_STABLE_DISPATCH_HALF_TYPES(
|
||||
qkv.scalar_type(), "fused_minimax_m3_qknorm_rope_kv_insert", [&] {
|
||||
using st = scalar_t;
|
||||
vllm::minimax_m3_fused_ops::launchFusedMiniMaxM3<st>(
|
||||
reinterpret_cast<st*>(qkv.data_ptr()),
|
||||
q_out.has_value() ? reinterpret_cast<st*>(q_out->data_ptr())
|
||||
: nullptr,
|
||||
index_q_out.has_value()
|
||||
? reinterpret_cast<st*>(index_q_out->data_ptr())
|
||||
: nullptr,
|
||||
reinterpret_cast<st const*>(q_norm_weight.data_ptr()),
|
||||
reinterpret_cast<st const*>(k_norm_weight.data_ptr()),
|
||||
has_index
|
||||
? reinterpret_cast<st const*>(index_q_norm_weight->data_ptr())
|
||||
: nullptr,
|
||||
has_index
|
||||
? reinterpret_cast<st const*>(index_k_norm_weight->data_ptr())
|
||||
: nullptr,
|
||||
reinterpret_cast<st const*>(cos_sin_cache.data_ptr()),
|
||||
reinterpret_cast<int64_t const*>(positions.data_ptr()),
|
||||
insert_kv
|
||||
? reinterpret_cast<int64_t const*>(slot_mapping->data_ptr())
|
||||
: nullptr,
|
||||
insert_kv ? reinterpret_cast<int64_t const*>(
|
||||
effective_index_slot_mapping->data_ptr())
|
||||
: nullptr,
|
||||
insert_kv ? reinterpret_cast<st*>(kv_cache->data_ptr()) : nullptr,
|
||||
(insert_kv && has_index)
|
||||
? reinterpret_cast<st*>(index_cache->data_ptr())
|
||||
: nullptr,
|
||||
static_cast<float>(eps), static_cast<int>(rotary_dim), num_tokens,
|
||||
nq, nkv, niq, static_cast<int>(block_size), kv_s_block, kv_s_kv,
|
||||
kv_s_token, kv_s_head, has_index, insert_kv, stream);
|
||||
DISPATCH_BY_KV_CACHE_DTYPE(qkv.scalar_type(), kv_cache_dtype,
|
||||
CALL_FUSED_MINIMAX_M3);
|
||||
});
|
||||
}
|
||||
|
||||
#undef CALL_FUSED_MINIMAX_M3
|
||||
|
||||
@@ -11,7 +11,7 @@
|
||||
namespace vllm {
|
||||
|
||||
// TODO(woosuk): Further optimize this kernel.
|
||||
template <typename scalar_t, int VEC_SIZE, int NUM_DIMS>
|
||||
template <typename scalar_t, int VEC_SIZE, int NUM_DIMS, bool HasWeight>
|
||||
__global__ void rms_norm_kernel(
|
||||
scalar_t* __restrict__ out, // [..., hidden_size]
|
||||
const scalar_t* __restrict__ input, // [..., hidden_size]
|
||||
@@ -20,7 +20,7 @@ __global__ void rms_norm_kernel(
|
||||
const int64_t input_stride_d4, // input.stride(-4)
|
||||
const int64_t input_shape_d2, // input.size(-2)
|
||||
const int64_t input_shape_d3, // input.size(-3)
|
||||
const scalar_t* __restrict__ weight, // [hidden_size]
|
||||
const scalar_t* __restrict__ weight, // [hidden_size], null if !HasWeight
|
||||
const float epsilon, const int num_tokens, const int hidden_size) {
|
||||
__shared__ float s_variance;
|
||||
float variance = 0.0f;
|
||||
@@ -74,11 +74,19 @@ __global__ void rms_norm_kernel(
|
||||
for (int i = threadIdx.x; i < hidden_size / VEC_SIZE; i += blockDim.x) {
|
||||
vec_n_t<scalar_t, VEC_SIZE> dst;
|
||||
vec_n_t<scalar_t, VEC_SIZE> src1 = v_in[i];
|
||||
vec_n_t<scalar_t, VEC_SIZE> src2 = v_w[i];
|
||||
vec_n_t<scalar_t, VEC_SIZE> src2;
|
||||
if constexpr (HasWeight) {
|
||||
src2 = v_w[i];
|
||||
}
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC_SIZE; j++) {
|
||||
float x = static_cast<float>(src1.val[j]);
|
||||
dst.val[j] = static_cast<scalar_t>(x * s_variance) * src2.val[j];
|
||||
scalar_t normalized = static_cast<scalar_t>(x * s_variance);
|
||||
if constexpr (HasWeight) {
|
||||
dst.val[j] = normalized * src2.val[j];
|
||||
} else {
|
||||
dst.val[j] = normalized;
|
||||
}
|
||||
}
|
||||
v_out[i] = dst;
|
||||
}
|
||||
@@ -88,13 +96,13 @@ __global__ void rms_norm_kernel(
|
||||
Additional optimizations we can make in this case are
|
||||
packed and vectorized operations, which help with the
|
||||
memory latency bottleneck. */
|
||||
template <typename scalar_t, int width>
|
||||
template <typename scalar_t, int width, bool HasWeight>
|
||||
__global__ std::enable_if_t<(width > 0) && _typeConvert<scalar_t>::exists>
|
||||
fused_add_rms_norm_kernel(
|
||||
scalar_t* __restrict__ input, // [..., hidden_size]
|
||||
const int64_t input_stride,
|
||||
scalar_t* __restrict__ residual, // [..., hidden_size]
|
||||
const scalar_t* __restrict__ weight, // [hidden_size]
|
||||
const scalar_t* __restrict__ weight, // [hidden_size], null if !HasWeight
|
||||
const float epsilon, const int num_tokens, const int hidden_size) {
|
||||
// Sanity checks on our vector struct and type-punned pointer arithmetic
|
||||
static_assert(std::is_pod_v<_f16Vec<scalar_t, width>>);
|
||||
@@ -136,13 +144,21 @@ fused_add_rms_norm_kernel(
|
||||
int id = blockIdx.x * vec_hidden_size + idx;
|
||||
int64_t strided_id = blockIdx.x * vec_input_stride + idx;
|
||||
_f16Vec<scalar_t, width> res = residual_v[id];
|
||||
_f16Vec<scalar_t, width> w = weight_v[idx];
|
||||
_f16Vec<scalar_t, width> out;
|
||||
using Converter = _typeConvert<scalar_t>;
|
||||
if constexpr (HasWeight) {
|
||||
_f16Vec<scalar_t, width> w = weight_v[idx];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < width; ++j) {
|
||||
float x = Converter::convert(res.data[j]);
|
||||
out.data[j] = Converter::convert(x * s_variance) * w.data[j];
|
||||
for (int j = 0; j < width; ++j) {
|
||||
float x = Converter::convert(res.data[j]);
|
||||
out.data[j] = Converter::convert(x * s_variance) * w.data[j];
|
||||
}
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (int j = 0; j < width; ++j) {
|
||||
float x = Converter::convert(res.data[j]);
|
||||
out.data[j] = Converter::convert(x * s_variance);
|
||||
}
|
||||
}
|
||||
input_v[strided_id] = out;
|
||||
}
|
||||
@@ -151,13 +167,13 @@ fused_add_rms_norm_kernel(
|
||||
/* Generic fused_add_rms_norm_kernel
|
||||
The width field is not used here but necessary for other specializations.
|
||||
*/
|
||||
template <typename scalar_t, int width>
|
||||
template <typename scalar_t, int width, bool HasWeight>
|
||||
__global__ std::enable_if_t<(width == 0) || !_typeConvert<scalar_t>::exists>
|
||||
fused_add_rms_norm_kernel(
|
||||
scalar_t* __restrict__ input, // [..., hidden_size]
|
||||
const int64_t input_stride,
|
||||
scalar_t* __restrict__ residual, // [..., hidden_size]
|
||||
const scalar_t* __restrict__ weight, // [hidden_size]
|
||||
const scalar_t* __restrict__ weight, // [hidden_size], null if !HasWeight
|
||||
const float epsilon, const int num_tokens, const int hidden_size) {
|
||||
__shared__ float s_variance;
|
||||
float variance = 0.0f;
|
||||
@@ -181,23 +197,29 @@ fused_add_rms_norm_kernel(
|
||||
|
||||
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
|
||||
float x = (float)residual[blockIdx.x * hidden_size + idx];
|
||||
input[blockIdx.x * input_stride + idx] =
|
||||
(scalar_t)(x * s_variance) * weight[idx];
|
||||
if constexpr (HasWeight) {
|
||||
input[blockIdx.x * input_stride + idx] =
|
||||
(scalar_t)(x * s_variance) * weight[idx];
|
||||
} else {
|
||||
input[blockIdx.x * input_stride + idx] = (scalar_t)(x * s_variance);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace vllm
|
||||
|
||||
void rms_norm(torch::stable::Tensor& out, // [..., hidden_size]
|
||||
torch::stable::Tensor& input, // [..., hidden_size]
|
||||
torch::stable::Tensor& weight, // [hidden_size]
|
||||
void rms_norm(torch::stable::Tensor& out, // [..., hidden_size]
|
||||
torch::stable::Tensor& input, // [..., hidden_size]
|
||||
std::optional<torch::stable::Tensor> weight, // [hidden_size]
|
||||
double epsilon) {
|
||||
STD_TORCH_CHECK(out.is_contiguous());
|
||||
if (input.stride(-1) != 1) {
|
||||
input = torch::stable::contiguous(input);
|
||||
}
|
||||
STD_TORCH_CHECK(input.stride(-1) == 1);
|
||||
STD_TORCH_CHECK(weight.is_contiguous());
|
||||
if (weight.has_value()) {
|
||||
STD_TORCH_CHECK(weight->is_contiguous());
|
||||
}
|
||||
|
||||
int hidden_size = input.size(-1);
|
||||
|
||||
@@ -215,46 +237,69 @@ void rms_norm(torch::stable::Tensor& out, // [..., hidden_size]
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
input.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream();
|
||||
const bool has_weight = weight.has_value();
|
||||
VLLM_STABLE_DISPATCH_RANK234(num_dims, [&] {
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
|
||||
input.scalar_type(), "rms_norm_kernel", [&] {
|
||||
const scalar_t* weight_ptr =
|
||||
has_weight ? weight->const_data_ptr<scalar_t>() : nullptr;
|
||||
const int calculated_vec_size =
|
||||
std::gcd(16 / sizeof(scalar_t), hidden_size);
|
||||
const int block_size =
|
||||
std::min(hidden_size / calculated_vec_size, max_block_size);
|
||||
dim3 block(block_size);
|
||||
VLLM_STABLE_DISPATCH_VEC_SIZE(calculated_vec_size, [&] {
|
||||
vllm::rms_norm_kernel<scalar_t, vec_size, tensor_rank>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
out.mutable_data_ptr<scalar_t>(),
|
||||
input.const_data_ptr<scalar_t>(), input_stride_d2,
|
||||
input_stride_d3, input_stride_d4, input_shape_d2,
|
||||
input_shape_d3, weight.const_data_ptr<scalar_t>(), epsilon,
|
||||
num_tokens, hidden_size);
|
||||
if (has_weight) {
|
||||
vllm::rms_norm_kernel<scalar_t, vec_size, tensor_rank, true>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
out.mutable_data_ptr<scalar_t>(),
|
||||
input.const_data_ptr<scalar_t>(), input_stride_d2,
|
||||
input_stride_d3, input_stride_d4, input_shape_d2,
|
||||
input_shape_d3, weight_ptr, epsilon, num_tokens,
|
||||
hidden_size);
|
||||
} else {
|
||||
vllm::rms_norm_kernel<scalar_t, vec_size, tensor_rank, false>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
out.mutable_data_ptr<scalar_t>(),
|
||||
input.const_data_ptr<scalar_t>(), input_stride_d2,
|
||||
input_stride_d3, input_stride_d4, input_shape_d2,
|
||||
input_shape_d3, weight_ptr, epsilon, num_tokens,
|
||||
hidden_size);
|
||||
}
|
||||
});
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
#define LAUNCH_FUSED_ADD_RMS_NORM(width) \
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES( \
|
||||
input.scalar_type(), "fused_add_rms_norm_kernel", [&] { \
|
||||
vllm::fused_add_rms_norm_kernel<scalar_t, width> \
|
||||
<<<grid, block, 0, stream>>>( \
|
||||
input.mutable_data_ptr<scalar_t>(), input_stride, \
|
||||
residual.mutable_data_ptr<scalar_t>(), \
|
||||
weight.const_data_ptr<scalar_t>(), epsilon, num_tokens, \
|
||||
hidden_size); \
|
||||
#define LAUNCH_FUSED_ADD_RMS_NORM(width, has_weight) \
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES( \
|
||||
input.scalar_type(), "fused_add_rms_norm_kernel", [&] { \
|
||||
if (has_weight) { \
|
||||
vllm::fused_add_rms_norm_kernel<scalar_t, width, true> \
|
||||
<<<grid, block, 0, stream>>>( \
|
||||
input.mutable_data_ptr<scalar_t>(), input_stride, \
|
||||
residual.mutable_data_ptr<scalar_t>(), \
|
||||
weight->const_data_ptr<scalar_t>(), epsilon, num_tokens, \
|
||||
hidden_size); \
|
||||
} else { \
|
||||
vllm::fused_add_rms_norm_kernel<scalar_t, width, false> \
|
||||
<<<grid, block, 0, stream>>>( \
|
||||
input.mutable_data_ptr<scalar_t>(), input_stride, \
|
||||
residual.mutable_data_ptr<scalar_t>(), nullptr, epsilon, \
|
||||
num_tokens, hidden_size); \
|
||||
} \
|
||||
});
|
||||
|
||||
void fused_add_rms_norm(torch::stable::Tensor& input, // [..., hidden_size]
|
||||
torch::stable::Tensor& residual, // [..., hidden_size]
|
||||
torch::stable::Tensor& weight, // [hidden_size]
|
||||
std::optional<torch::stable::Tensor> weight,
|
||||
double epsilon) {
|
||||
STD_TORCH_CHECK(weight.scalar_type() == input.scalar_type());
|
||||
STD_TORCH_CHECK(input.scalar_type() == residual.scalar_type());
|
||||
STD_TORCH_CHECK(residual.is_contiguous());
|
||||
STD_TORCH_CHECK(weight.is_contiguous());
|
||||
if (weight.has_value()) {
|
||||
STD_TORCH_CHECK(weight->scalar_type() == input.scalar_type());
|
||||
STD_TORCH_CHECK(weight->is_contiguous());
|
||||
}
|
||||
int hidden_size = input.size(-1);
|
||||
int64_t input_stride = input.stride(-2);
|
||||
int num_tokens = input.numel() / hidden_size;
|
||||
@@ -269,30 +314,33 @@ void fused_add_rms_norm(torch::stable::Tensor& input, // [..., hidden_size]
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
input.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream();
|
||||
/*If the tensor types are FP16/BF16, try to use the optimized kernel
|
||||
with packed + vectorized ops.
|
||||
Max optimization is achieved with a width-8 vector of FP16/BF16s
|
||||
since we can load at most 128 bits at once in a global memory op.
|
||||
However, this requires each tensor's data to be aligned to 16
|
||||
bytes.
|
||||
*/
|
||||
constexpr int vector_width = 8;
|
||||
constexpr int req_alignment_bytes = vector_width * 2;
|
||||
auto inp_ptr = reinterpret_cast<std::uintptr_t>(input.data_ptr());
|
||||
auto res_ptr = reinterpret_cast<std::uintptr_t>(residual.data_ptr());
|
||||
auto wt_ptr = reinterpret_cast<std::uintptr_t>(weight.data_ptr());
|
||||
constexpr int vector_width = 8;
|
||||
constexpr int req_alignment_bytes =
|
||||
vector_width * 2; // vector_width * sizeof(bfloat16 or float16) (float32
|
||||
// falls back to non-vectorized version anyway)
|
||||
bool ptrs_are_aligned = inp_ptr % req_alignment_bytes == 0 &&
|
||||
res_ptr % req_alignment_bytes == 0 &&
|
||||
wt_ptr % req_alignment_bytes == 0;
|
||||
bool offsets_are_multiple_of_vector_width =
|
||||
hidden_size % vector_width == 0 && input_stride % vector_width == 0;
|
||||
bool batch_invariant_launch = vllm::vllm_is_batch_invariant();
|
||||
if (ptrs_are_aligned && offsets_are_multiple_of_vector_width &&
|
||||
!batch_invariant_launch) {
|
||||
LAUNCH_FUSED_ADD_RMS_NORM(8);
|
||||
const bool has_weight = weight.has_value();
|
||||
if (has_weight) {
|
||||
auto wt_ptr = reinterpret_cast<std::uintptr_t>(weight->data_ptr());
|
||||
bool ptrs_are_aligned = inp_ptr % req_alignment_bytes == 0 &&
|
||||
res_ptr % req_alignment_bytes == 0 &&
|
||||
wt_ptr % req_alignment_bytes == 0;
|
||||
if (ptrs_are_aligned && offsets_are_multiple_of_vector_width &&
|
||||
!batch_invariant_launch) {
|
||||
LAUNCH_FUSED_ADD_RMS_NORM(8, true);
|
||||
} else {
|
||||
LAUNCH_FUSED_ADD_RMS_NORM(0, true);
|
||||
}
|
||||
} else {
|
||||
LAUNCH_FUSED_ADD_RMS_NORM(0);
|
||||
bool ptrs_are_aligned = inp_ptr % req_alignment_bytes == 0 &&
|
||||
res_ptr % req_alignment_bytes == 0;
|
||||
if (ptrs_are_aligned && offsets_are_multiple_of_vector_width &&
|
||||
!batch_invariant_launch) {
|
||||
LAUNCH_FUSED_ADD_RMS_NORM(8, false);
|
||||
} else {
|
||||
LAUNCH_FUSED_ADD_RMS_NORM(0, false);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,69 +0,0 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
// Adapted from SGLang:
|
||||
// https://github.com/sgl-project/sglang/blob/ded068a76e00878881d52d5bfb791e0f60d7311b/sgl-kernel/csrc/expert_specialization/es_sm100_mxfp8_blockscaled.cu
|
||||
|
||||
#include <torch/csrc/stable/library.h>
|
||||
#include <torch/csrc/stable/tensor.h>
|
||||
#include "libtorch_stable/torch_utils.h"
|
||||
|
||||
#include "cutlass_mxfp8_grouped_mm_launcher.cuh"
|
||||
|
||||
void cutlass_mxfp8_grouped_mm(const torch::stable::Tensor& a,
|
||||
const torch::stable::Tensor& b,
|
||||
const torch::stable::Tensor& sfa,
|
||||
const torch::stable::Tensor& sfb,
|
||||
torch::stable::Tensor& d,
|
||||
const torch::stable::Tensor& problem_sizes,
|
||||
const torch::stable::Tensor& expert_offsets,
|
||||
const torch::stable::Tensor& blockscale_offsets) {
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
STD_TORCH_CHECK(problem_sizes.dim() == 2, "problem_sizes must be 2D tensor");
|
||||
STD_TORCH_CHECK(problem_sizes.size(1) == 3,
|
||||
"problem_sizes must have shape (num_experts, 3)");
|
||||
STD_TORCH_CHECK(
|
||||
problem_sizes.size(0) == expert_offsets.size(0),
|
||||
"Number of experts in problem_sizes must match expert_offsets");
|
||||
STD_TORCH_CHECK(
|
||||
problem_sizes.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"problem_sizes must be int32");
|
||||
STD_TORCH_CHECK(
|
||||
expert_offsets.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"expert_offsets must be int32");
|
||||
STD_TORCH_CHECK(
|
||||
blockscale_offsets.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"blockscale_offsets must be int32");
|
||||
STD_TORCH_CHECK(a.dim() == 2,
|
||||
"a must be a 2D tensor of shape (num_tokens, k)");
|
||||
STD_TORCH_CHECK(b.dim() == 3,
|
||||
"b must be a 3D tensor of shape (num_experts, k, n)");
|
||||
STD_TORCH_CHECK(a.size(1) == b.size(1) && a.size(1) % 128 == 0,
|
||||
"k should align 128");
|
||||
STD_TORCH_CHECK(b.size(2) % 128 == 0, "n should align 128");
|
||||
STD_TORCH_CHECK(a.stride(1) == 1, "a must be row major");
|
||||
STD_TORCH_CHECK(b.stride(1) == 1, "b must be column major");
|
||||
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
a.get_device_index());
|
||||
auto stream = get_current_cuda_stream(a.get_device_index());
|
||||
if (d.scalar_type() == torch::headeronly::ScalarType::BFloat16) {
|
||||
expert_specialization::cutlass_mxfp8_grouped_mm_dispatch_out_dtype<
|
||||
cutlass::bfloat16_t>(a, b, sfa, sfb, d, problem_sizes, expert_offsets,
|
||||
blockscale_offsets, stream);
|
||||
} else if (d.scalar_type() == torch::headeronly::ScalarType::Half) {
|
||||
expert_specialization::cutlass_mxfp8_grouped_mm_dispatch_out_dtype<
|
||||
cutlass::half_t>(a, b, sfa, sfb, d, problem_sizes, expert_offsets,
|
||||
blockscale_offsets, stream);
|
||||
} else {
|
||||
STD_TORCH_CHECK(false, "dtype must be kFloat16 or kBFloat16");
|
||||
}
|
||||
#else
|
||||
STD_TORCH_CHECK(false,
|
||||
"No implemented cutlass_mxfp8_grouped_mm for "
|
||||
"current device");
|
||||
#endif
|
||||
}
|
||||
|
||||
STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, m) {
|
||||
m.impl("cutlass_mxfp8_grouped_mm", TORCH_BOX(&cutlass_mxfp8_grouped_mm));
|
||||
}
|
||||
@@ -1,141 +0,0 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
// Adapted from SGLang:
|
||||
// https://github.com/sgl-project/sglang/blob/ded068a76e00878881d52d5bfb791e0f60d7311b/sgl-kernel/csrc/expert_specialization/es_sm100_mxfp8_blockscaled_functor.cuh
|
||||
|
||||
#pragma once
|
||||
#include <cuda.h>
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/util/packed_stride.hpp"
|
||||
#include "cutlass_mxfp8_grouped_mm_traits.cuh"
|
||||
|
||||
namespace expert_specialization {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <typename GemmTraits>
|
||||
struct CutlassMxfp8GroupedMmOffsetFunctor {
|
||||
using Gemm = typename GemmTraits::Gemm;
|
||||
using ElementA = typename Gemm::ElementA;
|
||||
using ElementB = typename Gemm::ElementB;
|
||||
using ElementSF = typename GemmTraits::ElementSF;
|
||||
using ElementD = typename GemmTraits::ElementOutput;
|
||||
// Input
|
||||
int* expert_offsets{nullptr};
|
||||
int* blockscale_offsets{nullptr};
|
||||
// Output
|
||||
ElementA* a_base{nullptr};
|
||||
ElementB* b_base{nullptr};
|
||||
ElementSF* sfa_base{nullptr};
|
||||
ElementSF* sfb_base{nullptr};
|
||||
ElementD* d_base{nullptr};
|
||||
ElementA** a_offsets{nullptr};
|
||||
ElementB** b_offsets{nullptr};
|
||||
ElementSF** sfa_offsets{nullptr};
|
||||
ElementSF** sfb_offsets{nullptr};
|
||||
ElementD** d_offsets{nullptr};
|
||||
|
||||
CutlassMxfp8GroupedMmOffsetFunctor() = default;
|
||||
CutlassMxfp8GroupedMmOffsetFunctor(
|
||||
int* _expert_offsets, int* _blockscale_offsets, ElementA* _a_base,
|
||||
ElementB* _b_base, ElementSF* _sfa_base, ElementSF* _sfb_base,
|
||||
ElementD* _d_base, ElementA** _a_offsets, ElementB** _b_offsets,
|
||||
ElementSF** _sfa_offsets, ElementSF** _sfb_offsets, ElementD** _d_offsets)
|
||||
: expert_offsets{_expert_offsets},
|
||||
blockscale_offsets{_blockscale_offsets},
|
||||
a_base(_a_base),
|
||||
b_base(_b_base),
|
||||
sfa_base(_sfa_base),
|
||||
sfb_base(_sfb_base),
|
||||
d_base(_d_base),
|
||||
a_offsets(_a_offsets),
|
||||
b_offsets(_b_offsets),
|
||||
sfa_offsets(_sfa_offsets),
|
||||
sfb_offsets(_sfb_offsets),
|
||||
d_offsets(_d_offsets) {}
|
||||
|
||||
void CUTE_DEVICE operator()(int64_t expert_id, int m, int n, int k) {
|
||||
int64_t expert_offset = static_cast<int64_t>(expert_offsets[expert_id]);
|
||||
int64_t blockscale_offset =
|
||||
static_cast<int64_t>(blockscale_offsets[expert_id]);
|
||||
int64_t a_stride = expert_offset * k;
|
||||
int64_t b_stride = expert_id * k * n;
|
||||
int64_t d_stride = expert_offset * n;
|
||||
int64_t sfa_stride = blockscale_offset * (k / 32);
|
||||
int64_t sfb_stride = expert_id * n * (k / 32);
|
||||
|
||||
a_offsets[expert_id] = a_base + a_stride;
|
||||
b_offsets[expert_id] = b_base + b_stride;
|
||||
sfa_offsets[expert_id] = sfa_base + sfa_stride;
|
||||
sfb_offsets[expert_id] = sfb_base + sfb_stride;
|
||||
d_offsets[expert_id] = d_base + d_stride;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename GemmTraits>
|
||||
struct CutlassMxfp8GroupedMmLayoutFunctor {
|
||||
using Sm1xxBlkScaledConfig = typename GemmTraits::Sm1xxBlkScaledConfig;
|
||||
using LayoutSFA = typename GemmTraits::LayoutSFA;
|
||||
using LayoutSFB = typename GemmTraits::LayoutSFB;
|
||||
LayoutSFA* layout_sfa_base{nullptr};
|
||||
LayoutSFB* layout_sfb_base{nullptr};
|
||||
|
||||
CutlassMxfp8GroupedMmLayoutFunctor() = default;
|
||||
CutlassMxfp8GroupedMmLayoutFunctor(LayoutSFA* _layout_sfa_base,
|
||||
LayoutSFB* _layout_sfb_base)
|
||||
: layout_sfa_base(_layout_sfa_base), layout_sfb_base(_layout_sfb_base) {}
|
||||
|
||||
void CUTE_DEVICE operator()(int64_t expert_id, int m, int n, int k) {
|
||||
LayoutSFA* layout_sfa_ptr = layout_sfa_base + expert_id;
|
||||
LayoutSFB* layout_sfb_ptr = layout_sfb_base + expert_id;
|
||||
*layout_sfa_ptr = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(
|
||||
cute::make_shape(m, n, k, 1));
|
||||
*layout_sfb_ptr = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(
|
||||
cute::make_shape(m, n, k, 1));
|
||||
}
|
||||
};
|
||||
|
||||
template <typename GemmTraits>
|
||||
struct CutlassMxfp8GroupedMmStrideFunctor {
|
||||
using StrideA = typename GemmTraits::StrideA;
|
||||
using StrideB = typename GemmTraits::StrideB;
|
||||
using StrideD = typename GemmTraits::StrideD;
|
||||
StrideA* stride_A_base{nullptr};
|
||||
StrideB* stride_B_base{nullptr};
|
||||
StrideD* stride_D_base{nullptr};
|
||||
|
||||
CutlassMxfp8GroupedMmStrideFunctor() = default;
|
||||
CutlassMxfp8GroupedMmStrideFunctor(StrideA* _stride_A_base,
|
||||
StrideB* _stride_B_base,
|
||||
StrideD* _stride_D_base)
|
||||
: stride_A_base(_stride_A_base),
|
||||
stride_B_base(_stride_B_base),
|
||||
stride_D_base(_stride_D_base) {}
|
||||
|
||||
void CUTE_DEVICE operator()(int64_t expert_id, int m, int n, int k) {
|
||||
StrideA* stride_A = stride_A_base + expert_id;
|
||||
StrideB* stride_B = stride_B_base + expert_id;
|
||||
StrideD* stride_D = stride_D_base + expert_id;
|
||||
*stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, 1});
|
||||
*stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n, k, 1});
|
||||
*stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n, 1});
|
||||
}
|
||||
};
|
||||
|
||||
template <typename OffsetFunctor, typename LayoutFunctor,
|
||||
typename StrideFunctor>
|
||||
__global__ void cutlassMxfp8GroupedMmPreComputeKernel(
|
||||
int* problem_sizes, OffsetFunctor offset_functor,
|
||||
LayoutFunctor layout_functor, StrideFunctor stride_functor) {
|
||||
int64_t expert_id = static_cast<int64_t>(threadIdx.x);
|
||||
int m = problem_sizes[expert_id * 3 + 0];
|
||||
int n = problem_sizes[expert_id * 3 + 1];
|
||||
int k = problem_sizes[expert_id * 3 + 2];
|
||||
|
||||
offset_functor(expert_id, m, n, k);
|
||||
layout_functor(expert_id, m, n, k);
|
||||
stride_functor(expert_id, m, n, k);
|
||||
}
|
||||
|
||||
} // namespace expert_specialization
|
||||
@@ -1,198 +0,0 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
// Adapted from SGLang:
|
||||
// https://github.com/sgl-project/sglang/blob/ded068a76e00878881d52d5bfb791e0f60d7311b/sgl-kernel/csrc/expert_specialization/es_sm100_mxfp8_blockscaled_launcher.cuh
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <torch/csrc/stable/tensor.h>
|
||||
#include <torch/headeronly/util/Exception.h>
|
||||
|
||||
#include <cassert>
|
||||
#include <iostream>
|
||||
#include <string>
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass_mxfp8_grouped_mm_functor.cuh"
|
||||
#include "cutlass_mxfp8_grouped_mm_traits.cuh"
|
||||
#include "libtorch_stable/torch_utils.h"
|
||||
|
||||
namespace expert_specialization {
|
||||
|
||||
template <typename GemmTraits>
|
||||
void cutlass_mxfp8_grouped_mm_pre_compute(
|
||||
torch::stable::Tensor& a_ptrs, torch::stable::Tensor& b_ptrs,
|
||||
torch::stable::Tensor& sfa_ptrs, torch::stable::Tensor& sfb_ptrs,
|
||||
torch::stable::Tensor& d_ptrs, torch::stable::Tensor& stride_a,
|
||||
torch::stable::Tensor& stride_b, torch::stable::Tensor& stride_d,
|
||||
torch::stable::Tensor& layout_sfa, torch::stable::Tensor& layout_sfb,
|
||||
const torch::stable::Tensor& a, const torch::stable::Tensor& b,
|
||||
const torch::stable::Tensor& sfa, const torch::stable::Tensor& sfb,
|
||||
const torch::stable::Tensor& d, const torch::stable::Tensor& problem_sizes,
|
||||
const torch::stable::Tensor& expert_offsets,
|
||||
const torch::stable::Tensor& blockscale_offsets, cudaStream_t stream) {
|
||||
using OffsetFunctor = CutlassMxfp8GroupedMmOffsetFunctor<GemmTraits>;
|
||||
using ElementA = typename OffsetFunctor::ElementA;
|
||||
using ElementB = typename OffsetFunctor::ElementB;
|
||||
using ElementSF = typename OffsetFunctor::ElementSF;
|
||||
using ElementD = typename OffsetFunctor::ElementD;
|
||||
|
||||
using LayoutFunctor = CutlassMxfp8GroupedMmLayoutFunctor<GemmTraits>;
|
||||
using LayoutSFA = typename LayoutFunctor::LayoutSFA;
|
||||
using LayoutSFB = typename LayoutFunctor::LayoutSFB;
|
||||
|
||||
using StrideFunctor = CutlassMxfp8GroupedMmStrideFunctor<GemmTraits>;
|
||||
using StrideA = typename StrideFunctor::StrideA;
|
||||
using StrideB = typename StrideFunctor::StrideB;
|
||||
using StrideD = typename StrideFunctor::StrideD;
|
||||
|
||||
int num_experts = static_cast<int>(expert_offsets.size(0));
|
||||
STD_TORCH_CHECK(num_experts <= 1024,
|
||||
"Number of experts cannot exceed 1024, the maximum number of "
|
||||
"threads per block.");
|
||||
|
||||
OffsetFunctor offset_functor(
|
||||
reinterpret_cast<int*>(expert_offsets.data_ptr()),
|
||||
reinterpret_cast<int*>(blockscale_offsets.data_ptr()),
|
||||
reinterpret_cast<ElementA*>(a.data_ptr()),
|
||||
reinterpret_cast<ElementB*>(b.data_ptr()),
|
||||
reinterpret_cast<ElementSF*>(sfa.data_ptr()),
|
||||
reinterpret_cast<ElementSF*>(sfb.data_ptr()),
|
||||
reinterpret_cast<ElementD*>(d.data_ptr()),
|
||||
reinterpret_cast<ElementA**>(a_ptrs.data_ptr()),
|
||||
reinterpret_cast<ElementB**>(b_ptrs.data_ptr()),
|
||||
reinterpret_cast<ElementSF**>(sfa_ptrs.data_ptr()),
|
||||
reinterpret_cast<ElementSF**>(sfb_ptrs.data_ptr()),
|
||||
reinterpret_cast<ElementD**>(d_ptrs.data_ptr()));
|
||||
LayoutFunctor layout_functor(
|
||||
reinterpret_cast<LayoutSFA*>(layout_sfa.data_ptr()),
|
||||
reinterpret_cast<LayoutSFB*>(layout_sfb.data_ptr()));
|
||||
StrideFunctor stride_functor(reinterpret_cast<StrideA*>(stride_a.data_ptr()),
|
||||
reinterpret_cast<StrideB*>(stride_b.data_ptr()),
|
||||
reinterpret_cast<StrideD*>(stride_d.data_ptr()));
|
||||
cutlassMxfp8GroupedMmPreComputeKernel<<<1, num_experts, 0, stream>>>(
|
||||
static_cast<int*>(problem_sizes.data_ptr()), offset_functor,
|
||||
layout_functor, stride_functor);
|
||||
}
|
||||
|
||||
template <typename GemmTraits>
|
||||
void cutlass_mxfp8_grouped_mm(const torch::stable::Tensor& a_ptrs,
|
||||
const torch::stable::Tensor& b_ptrs,
|
||||
const torch::stable::Tensor& sfa_ptrs,
|
||||
const torch::stable::Tensor& sfb_ptrs,
|
||||
const torch::stable::Tensor& d_ptrs,
|
||||
const torch::stable::Tensor& stride_a,
|
||||
const torch::stable::Tensor& stride_b,
|
||||
const torch::stable::Tensor& stride_d,
|
||||
const torch::stable::Tensor& layout_sfa,
|
||||
const torch::stable::Tensor& layout_sfb,
|
||||
const torch::stable::Tensor& problem_sizes,
|
||||
cudaStream_t stream) {
|
||||
using Gemm = typename GemmTraits::Gemm;
|
||||
using ElementA = typename Gemm::ElementA;
|
||||
using ElementB = typename Gemm::ElementB;
|
||||
using ElementSF = typename GemmTraits::ElementSF;
|
||||
using ElementD = typename GemmTraits::ElementOutput;
|
||||
using StrideA = typename GemmTraits::StrideA;
|
||||
using StrideB = typename GemmTraits::StrideB;
|
||||
using StrideD = typename GemmTraits::StrideD;
|
||||
using LayoutSFA = typename GemmTraits::LayoutSFA;
|
||||
using LayoutSFB = typename GemmTraits::LayoutSFB;
|
||||
using UnderlyingProblemShape =
|
||||
typename GemmTraits::ProblemShape::UnderlyingProblemShape;
|
||||
|
||||
cutlass::KernelHardwareInfo hw_info;
|
||||
hw_info.device_id = d_ptrs.get_device_index();
|
||||
hw_info.sm_count = get_device_prop()->multiProcessorCount;
|
||||
hw_info.cluster_shape = GemmTraits::MMAConfig::preferred_cluster;
|
||||
hw_info.cluster_shape_fallback = GemmTraits::MMAConfig::fallback_cluster;
|
||||
|
||||
int num_experts = static_cast<int>(problem_sizes.size(0));
|
||||
|
||||
UnderlyingProblemShape* underlying_problem_shape =
|
||||
reinterpret_cast<UnderlyingProblemShape*>(problem_sizes.data_ptr());
|
||||
|
||||
typename Gemm::Arguments arguments = {
|
||||
cutlass::gemm::GemmUniversalMode::kGrouped,
|
||||
{num_experts, underlying_problem_shape, nullptr},
|
||||
{reinterpret_cast<const ElementA**>(a_ptrs.data_ptr()),
|
||||
reinterpret_cast<StrideA*>(stride_a.data_ptr()),
|
||||
reinterpret_cast<const ElementB**>(b_ptrs.data_ptr()),
|
||||
reinterpret_cast<StrideB*>(stride_b.data_ptr()),
|
||||
reinterpret_cast<const ElementSF**>(sfa_ptrs.data_ptr()),
|
||||
reinterpret_cast<LayoutSFA*>(layout_sfa.data_ptr()),
|
||||
reinterpret_cast<const ElementSF**>(sfb_ptrs.data_ptr()),
|
||||
reinterpret_cast<LayoutSFB*>(layout_sfb.data_ptr())},
|
||||
{{},
|
||||
nullptr,
|
||||
nullptr,
|
||||
reinterpret_cast<ElementD**>(d_ptrs.data_ptr()),
|
||||
reinterpret_cast<StrideD*>(stride_d.data_ptr())},
|
||||
hw_info,
|
||||
{} // Scheduler
|
||||
};
|
||||
|
||||
Gemm gemm;
|
||||
|
||||
auto can_implement_status = gemm.can_implement(arguments);
|
||||
STD_TORCH_CHECK(can_implement_status == cutlass::Status::kSuccess,
|
||||
"Failed to implement GEMM");
|
||||
|
||||
size_t workspace_size = gemm.get_workspace_size(arguments);
|
||||
torch::stable::Tensor workspace = torch::stable::empty(
|
||||
{static_cast<int64_t>(workspace_size)},
|
||||
torch::headeronly::ScalarType::Byte, std::nullopt, d_ptrs.device());
|
||||
|
||||
auto status = gemm.initialize(arguments, workspace.data_ptr(), stream);
|
||||
STD_TORCH_CHECK(status == cutlass::Status::kSuccess,
|
||||
"Failed to initialize GEMM");
|
||||
|
||||
status = gemm.run(stream, nullptr, true); // Enable PDL
|
||||
STD_TORCH_CHECK(status == cutlass::Status::kSuccess, "Failed to run GEMM");
|
||||
}
|
||||
|
||||
template <typename OutType>
|
||||
void cutlass_mxfp8_grouped_mm_dispatch_out_dtype(
|
||||
const torch::stable::Tensor& a, const torch::stable::Tensor& b,
|
||||
const torch::stable::Tensor& sfa, const torch::stable::Tensor& sfb,
|
||||
torch::stable::Tensor& d, const torch::stable::Tensor& problem_sizes,
|
||||
const torch::stable::Tensor& expert_offsets,
|
||||
const torch::stable::Tensor& blockscale_offsets, cudaStream_t stream) {
|
||||
int num_experts = static_cast<int>(problem_sizes.size(0));
|
||||
auto device = a.device();
|
||||
|
||||
torch::stable::Tensor a_ptrs = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
torch::stable::Tensor b_ptrs = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
torch::stable::Tensor sfa_ptrs = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
torch::stable::Tensor sfb_ptrs = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
torch::stable::Tensor d_ptrs = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
|
||||
torch::stable::Tensor stride_a = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
torch::stable::Tensor stride_b = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
torch::stable::Tensor stride_d = torch::stable::empty(
|
||||
num_experts, torch::headeronly::ScalarType::Long, std::nullopt, device);
|
||||
torch::stable::Tensor layout_sfa =
|
||||
torch::stable::empty({num_experts, 5}, torch::headeronly::ScalarType::Int,
|
||||
std::nullopt, device);
|
||||
torch::stable::Tensor layout_sfb =
|
||||
torch::stable::empty({num_experts, 5}, torch::headeronly::ScalarType::Int,
|
||||
std::nullopt, device);
|
||||
|
||||
using GemmTraits = CutlassMxfp8GroupedMmGemmTraits<MMA1SMConfig, OutType>;
|
||||
cutlass_mxfp8_grouped_mm_pre_compute<GemmTraits>(
|
||||
a_ptrs, b_ptrs, sfa_ptrs, sfb_ptrs, d_ptrs, stride_a, stride_b, stride_d,
|
||||
layout_sfa, layout_sfb, a, b, sfa, sfb, d, problem_sizes, expert_offsets,
|
||||
blockscale_offsets, stream);
|
||||
cutlass_mxfp8_grouped_mm<GemmTraits>(
|
||||
a_ptrs, b_ptrs, sfa_ptrs, sfb_ptrs, d_ptrs, stride_a, stride_b, stride_d,
|
||||
layout_sfa, layout_sfb, problem_sizes, stream);
|
||||
}
|
||||
|
||||
} // namespace expert_specialization
|
||||
@@ -1,127 +0,0 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
// Adapted from SGLang:
|
||||
// https://github.com/sgl-project/sglang/blob/ded068a76e00878881d52d5bfb791e0f60d7311b/sgl-kernel/csrc/expert_specialization/es_sm100_mxfp8_blockscaled_traits.cuh
|
||||
|
||||
#pragma once
|
||||
|
||||
// Misc
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/arch/arch.h"
|
||||
#include "cutlass/arch/mma.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/detail/sm100_blockscaled_layout.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
#include "cutlass/gemm/group_array_problem_shape.hpp"
|
||||
#include "cutlass/layout/layout.h"
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
#include "cutlass/numeric_size.h"
|
||||
|
||||
// Collective Builder
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/fusion/sm90_callbacks_tma_warpspecialized.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
|
||||
// Integration
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
|
||||
namespace expert_specialization {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// Different configs for 1SM and 2SM MMA kernel
|
||||
struct MMA1SMConfig {
|
||||
using MmaTileShape = Shape<_128, _128, _128>;
|
||||
using KernelSchedule =
|
||||
cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
|
||||
const static dim3 preferred_cluster;
|
||||
const static dim3 fallback_cluster;
|
||||
};
|
||||
const dim3 MMA1SMConfig::preferred_cluster(1, 4, 1);
|
||||
const dim3 MMA1SMConfig::fallback_cluster(1, 2, 1);
|
||||
|
||||
template <typename _MMAConfig, typename OutputDtype>
|
||||
struct CutlassMxfp8GroupedMmGemmTraits {
|
||||
using MMAConfig = _MMAConfig;
|
||||
using ElementInput = cutlass::float_e4m3_t;
|
||||
using ElementOutput = OutputDtype;
|
||||
using ProblemShape = cutlass::gemm::GroupProblemShape<Shape<int, int, int>>;
|
||||
|
||||
// A matrix configuration
|
||||
using ElementA = cutlass::mx_float8_t<ElementInput>;
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
constexpr static int AlignmentA = 32;
|
||||
|
||||
// B matrix configuration
|
||||
using ElementB = cutlass::mx_float8_t<ElementInput>;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
constexpr static int AlignmentB = 32;
|
||||
|
||||
// C/D matrix configuration
|
||||
using ElementC = void;
|
||||
using ElementD = ElementOutput;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
constexpr static int AlignmentC = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr static int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
using ElementAccumulator = float;
|
||||
|
||||
static constexpr auto RoundStyle = cutlass::FloatRoundStyle::round_to_nearest;
|
||||
using CustomEVTIdentity = // acc
|
||||
cutlass::epilogue::fusion::Sm90EVT<
|
||||
cutlass::epilogue::fusion::Sm90Compute<
|
||||
cutlass::epilogue::thread::Identity, ElementD, ElementAccumulator,
|
||||
RoundStyle>,
|
||||
cutlass::epilogue::fusion::Sm90AccFetch>;
|
||||
|
||||
// Core kernel configurations
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
|
||||
using StageCountType = cutlass::gemm::collective::StageCountAuto;
|
||||
|
||||
// Runtime Cluster Shape
|
||||
using ClusterShape = Shape<int32_t, int32_t, _1>;
|
||||
|
||||
// Define Epilogue
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag, OperatorClass, typename MMAConfig::MmaTileShape,
|
||||
ClusterShape, Shape<_64, _64>, ElementAccumulator, ElementAccumulator,
|
||||
ElementC, LayoutC*, AlignmentC, ElementD, LayoutD*, AlignmentD,
|
||||
typename MMAConfig::EpilogueSchedule,
|
||||
CustomEVTIdentity>::CollectiveOp;
|
||||
|
||||
// Define Mainloop
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag, OperatorClass, ElementA, LayoutA*, AlignmentA, ElementB,
|
||||
LayoutB*, AlignmentB, ElementAccumulator,
|
||||
typename MMAConfig::MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(
|
||||
sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
typename MMAConfig::KernelSchedule>::CollectiveOp;
|
||||
|
||||
// Define GemmKernel
|
||||
using GemmKernel =
|
||||
cutlass::gemm::kernel::GemmUniversal<ProblemShape, CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
|
||||
using ElementSF = typename Gemm::GemmKernel::ElementSF;
|
||||
using StrideA = typename Gemm::GemmKernel::InternalStrideA;
|
||||
using StrideB = typename Gemm::GemmKernel::InternalStrideB;
|
||||
using StrideC = typename Gemm::GemmKernel::InternalStrideC;
|
||||
using StrideD = typename Gemm::GemmKernel::InternalStrideD;
|
||||
using LayoutSFA =
|
||||
typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
|
||||
using LayoutSFB =
|
||||
typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
|
||||
using Sm1xxBlkScaledConfig =
|
||||
typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
|
||||
};
|
||||
|
||||
} // namespace expert_specialization
|
||||
@@ -1,66 +0,0 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
// Adapted from SGLang:
|
||||
// https://github.com/sgl-project/sglang/blob/ded068a76e00878881d52d5bfb791e0f60d7311b/sgl-kernel/csrc/expert_specialization/es_sm100_mxfp8_blockscaled_group_quant.cu
|
||||
|
||||
#include <torch/csrc/stable/library.h>
|
||||
#include <torch/csrc/stable/tensor.h>
|
||||
#include "libtorch_stable/torch_utils.h"
|
||||
|
||||
#include "mxfp8_experts_quant.cuh"
|
||||
|
||||
void mxfp8_experts_quant(const torch::stable::Tensor& input,
|
||||
const torch::stable::Tensor& problem_sizes,
|
||||
const torch::stable::Tensor& expert_offsets,
|
||||
const torch::stable::Tensor& blockscale_offsets,
|
||||
torch::stable::Tensor& quant_output,
|
||||
torch::stable::Tensor& scale_factor) {
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
STD_TORCH_CHECK(input.dim() == 2, "input must be 2D tensor");
|
||||
STD_TORCH_CHECK(input.size(1) % 128 == 0, "k must align to 128");
|
||||
STD_TORCH_CHECK(input.stride(1) == 1, "input must be row major");
|
||||
STD_TORCH_CHECK(problem_sizes.dim() == 2, "problem_sizes must be 2D tensor");
|
||||
STD_TORCH_CHECK(
|
||||
problem_sizes.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"problem_sizes must be int32");
|
||||
STD_TORCH_CHECK(
|
||||
expert_offsets.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"expert_offsets must be int32");
|
||||
STD_TORCH_CHECK(
|
||||
blockscale_offsets.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"blockscale_offsets must be int32");
|
||||
|
||||
auto groups = problem_sizes.size(0);
|
||||
STD_TORCH_CHECK(
|
||||
expert_offsets.dim() == 1 && expert_offsets.size(0) == groups,
|
||||
"expert_offsets must be 1D and have size equal to the number of groups");
|
||||
STD_TORCH_CHECK(
|
||||
blockscale_offsets.dim() == 1 && blockscale_offsets.size(0) == groups,
|
||||
"blockscale_offsets must be 1D and have size equal to the number of "
|
||||
"groups");
|
||||
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
input.get_device_index());
|
||||
if (input.scalar_type() == torch::headeronly::ScalarType::BFloat16) {
|
||||
expert_specialization::launch_mxfp8_experts_quant<__nv_bfloat16>(
|
||||
input, problem_sizes, expert_offsets, blockscale_offsets, quant_output,
|
||||
scale_factor);
|
||||
} else if (input.scalar_type() == torch::headeronly::ScalarType::Half) {
|
||||
expert_specialization::launch_mxfp8_experts_quant<__half>(
|
||||
input, problem_sizes, expert_offsets, blockscale_offsets, quant_output,
|
||||
scale_factor);
|
||||
} else {
|
||||
STD_TORCH_CHECK(false, "dtype must be kFloat16 or kBFloat16");
|
||||
}
|
||||
#else
|
||||
STD_TORCH_CHECK(false,
|
||||
"No implemented mxfp8_experts_quant for "
|
||||
"current device");
|
||||
#endif
|
||||
}
|
||||
|
||||
// Registered here (not torch_bindings.cpp) because ENABLE_ES_MXFP8_GROUPED_MM
|
||||
// is applied only under COMPILE_LANGUAGE:CUDA.
|
||||
STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, m) {
|
||||
m.impl("mxfp8_experts_quant", TORCH_BOX(&mxfp8_experts_quant));
|
||||
}
|
||||
@@ -1,416 +0,0 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
// Adapted from SGLang:
|
||||
// https://github.com/sgl-project/sglang/blob/ded068a76e00878881d52d5bfb791e0f60d7311b/sgl-kernel/csrc/expert_specialization/es_sm100_mxfp8_blockscaled_group_quant.cuh
|
||||
|
||||
#pragma once
|
||||
#include <cuda.h>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
|
||||
#include <torch/csrc/inductor/aoti_torch/c/shim.h>
|
||||
#include <torch/csrc/stable/macros.h>
|
||||
#include <torch/csrc/stable/tensor.h>
|
||||
#include <torch/headeronly/util/Exception.h>
|
||||
|
||||
#include <cuda/ptx>
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
#include "libtorch_stable/torch_utils.h"
|
||||
|
||||
namespace expert_specialization {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
constexpr uint32_t THREAD_BLOCK_SIZE = 128;
|
||||
constexpr uint32_t WARP_SIZE = 32;
|
||||
constexpr int BLOCK_M = 128;
|
||||
constexpr int BLOCK_K = 128;
|
||||
using ThrLayout = Layout<Shape<_16, _8>, Stride<_8, _1>>;
|
||||
using ValLayout = Layout<Shape<_1, _16>>;
|
||||
using SfR2SThrLayout = Layout<Shape<_16, _4>, Stride<_4, _1>>;
|
||||
using SfR2SValLayout = Layout<Shape<_1, _1>>;
|
||||
using ScaleFactorTileLayout =
|
||||
Layout<Shape<Shape<_32, _4>, _4>, Stride<Stride<_16, _4>, _1>>;
|
||||
|
||||
// Fast reciprocal.
|
||||
inline __device__ float reciprocal_approximate_ftz(float a) {
|
||||
float b;
|
||||
asm volatile("rcp.approx.ftz.f32 %0, %1;\n" : "=f"(b) : "f"(a));
|
||||
return b;
|
||||
}
|
||||
|
||||
// Some code references TRT-LLM:
|
||||
// https://github.com/NVIDIA/TensorRT-LLM/blob/main/cpp/tensorrt_llm/kernels/quantization.cuh
|
||||
template <typename FragmentS, typename FragmentD>
|
||||
__inline__ __device__ uint8_t cvt_warp_fp16_to_mxfp8(FragmentS& fragment_s,
|
||||
FragmentD& fragment_d) {
|
||||
using FragmentSLayout = typename FragmentS::layout_type;
|
||||
using FragmentDLayout = typename FragmentD::layout_type;
|
||||
FragmentSLayout fragment_s_layout;
|
||||
FragmentDLayout fragment_d_layout;
|
||||
static_assert(is_static<FragmentSLayout>::value &&
|
||||
size(fragment_s_layout) == 16);
|
||||
static_assert(is_static<FragmentDLayout>::value &&
|
||||
size(fragment_d_layout) == 16);
|
||||
|
||||
constexpr int eles_per_thr = 16;
|
||||
using ValType = typename FragmentS::element_type;
|
||||
using VecType = std::conditional_t<std::is_same_v<ValType, __nv_bfloat16>,
|
||||
__nv_bfloat162, __half2>;
|
||||
VecType vec[8];
|
||||
// Assign vals
|
||||
vec[0].x = fragment_s(Int<0>{});
|
||||
vec[0].y = fragment_s(Int<1>{});
|
||||
vec[1].x = fragment_s(Int<2>{});
|
||||
vec[1].y = fragment_s(Int<3>{});
|
||||
vec[2].x = fragment_s(Int<4>{});
|
||||
vec[2].y = fragment_s(Int<5>{});
|
||||
vec[3].x = fragment_s(Int<6>{});
|
||||
vec[3].y = fragment_s(Int<7>{});
|
||||
vec[4].x = fragment_s(Int<8>{});
|
||||
vec[4].y = fragment_s(Int<9>{});
|
||||
vec[5].x = fragment_s(Int<10>{});
|
||||
vec[5].y = fragment_s(Int<11>{});
|
||||
vec[6].x = fragment_s(Int<12>{});
|
||||
vec[6].y = fragment_s(Int<13>{});
|
||||
vec[7].x = fragment_s(Int<14>{});
|
||||
vec[7].y = fragment_s(Int<15>{});
|
||||
|
||||
auto local_max = __habs2(vec[0]);
|
||||
for (int i = 1; i < eles_per_thr / 2; i++) {
|
||||
local_max = __hmax2(__habs2(vec[i]), local_max);
|
||||
}
|
||||
local_max = __hmax2(__shfl_xor_sync(uint32_t(-1), local_max, 1), local_max);
|
||||
|
||||
// Get the final absolute maximum values.
|
||||
float block_max(0.0f);
|
||||
if constexpr (std::is_same_v<ValType, __nv_bfloat16>) {
|
||||
block_max = __bfloat162float(__hmax(local_max.x, local_max.y));
|
||||
} else {
|
||||
block_max = __half2float(__hmax(local_max.x, local_max.y));
|
||||
}
|
||||
// Get the SF (max value of the vector / max value of mxfp8).
|
||||
float sf_val = block_max * reciprocal_approximate_ftz(448.0f);
|
||||
// 8 bits representation of the SF.
|
||||
uint8_t fp8_sf_val;
|
||||
|
||||
__nv_fp8_e8m0 tmp_sf_val;
|
||||
tmp_sf_val.__x =
|
||||
__nv_cvt_float_to_e8m0(sf_val, __NV_SATFINITE, cudaRoundPosInf);
|
||||
sf_val = static_cast<float>(tmp_sf_val);
|
||||
fp8_sf_val = tmp_sf_val.__x;
|
||||
// Get the output scale (reciprocal of the SFValue).
|
||||
float output_scale =
|
||||
block_max != 0.f ? reciprocal_approximate_ftz(sf_val) : 0.0f;
|
||||
|
||||
// Convert the input to float.
|
||||
float2 fp2_vals[eles_per_thr / 2];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < eles_per_thr / 2; i++) {
|
||||
if constexpr (std::is_same_v<ValType, __half>) {
|
||||
fp2_vals[i] = __half22float2(vec[i]);
|
||||
} else {
|
||||
fp2_vals[i] = __bfloat1622float2(vec[i]);
|
||||
}
|
||||
fp2_vals[i].x *= output_scale;
|
||||
fp2_vals[i].y *= output_scale;
|
||||
}
|
||||
union {
|
||||
uint8_t bytes[16];
|
||||
__nv_fp8x2_e4m3 elts[8];
|
||||
} u;
|
||||
u.elts[0] = __nv_fp8x2_e4m3(fp2_vals[0]);
|
||||
u.elts[1] = __nv_fp8x2_e4m3(fp2_vals[1]);
|
||||
u.elts[2] = __nv_fp8x2_e4m3(fp2_vals[2]);
|
||||
u.elts[3] = __nv_fp8x2_e4m3(fp2_vals[3]);
|
||||
u.elts[4] = __nv_fp8x2_e4m3(fp2_vals[4]);
|
||||
u.elts[5] = __nv_fp8x2_e4m3(fp2_vals[5]);
|
||||
u.elts[6] = __nv_fp8x2_e4m3(fp2_vals[6]);
|
||||
u.elts[7] = __nv_fp8x2_e4m3(fp2_vals[7]);
|
||||
fragment_d(Int<0>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[0]);
|
||||
fragment_d(Int<1>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[1]);
|
||||
fragment_d(Int<2>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[2]);
|
||||
fragment_d(Int<3>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[3]);
|
||||
fragment_d(Int<4>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[4]);
|
||||
fragment_d(Int<5>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[5]);
|
||||
fragment_d(Int<6>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[6]);
|
||||
fragment_d(Int<7>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[7]);
|
||||
fragment_d(Int<8>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[8]);
|
||||
fragment_d(Int<9>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[9]);
|
||||
fragment_d(Int<10>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[10]);
|
||||
fragment_d(Int<11>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[11]);
|
||||
fragment_d(Int<12>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[12]);
|
||||
fragment_d(Int<13>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[13]);
|
||||
fragment_d(Int<14>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[14]);
|
||||
fragment_d(Int<15>{}) = cutlass::float_e4m3_t::bitcast(u.bytes[15]);
|
||||
return fp8_sf_val;
|
||||
}
|
||||
|
||||
template <typename TensorS, typename TensorP, typename TensorD,
|
||||
typename TensorSharedSF, typename TensorSF, typename TiledCopyG2R,
|
||||
typename TiledCopyR2G, typename TiledCopyR2S>
|
||||
__inline__ __device__ void mxfp8_experts_quant_tile(
|
||||
TensorS& tensor_s, TensorP& tensor_p, TensorD& tensor_d,
|
||||
TensorSharedSF& tensor_shared_sf, TensorSF& tensor_sf, int m,
|
||||
TiledCopyG2R& tiled_copy_g2r, TiledCopyR2G& tiled_copy_r2g,
|
||||
TiledCopyR2S& tiled_copy_r2s) {
|
||||
static_assert(size(get<0>(typename TensorS::layout_type{})) == 128 &&
|
||||
size(get<1>(typename TensorS::layout_type{})) == 128 &&
|
||||
stride(get<1>(typename TensorS::layout_type{})) == 1);
|
||||
static_assert(size(get<0>(typename TensorD::layout_type{})) == 128 &&
|
||||
size(get<1>(typename TensorD::layout_type{})) == 128 &&
|
||||
stride(get<1>(typename TensorD::layout_type{})) == 1);
|
||||
static_assert(size(get<0>(typename TensorP::layout_type{})) == 128 &&
|
||||
size(get<1>(typename TensorP::layout_type{})) == 128);
|
||||
static_assert(size(get<0>(typename TensorSharedSF::layout_type{})) == 128 &&
|
||||
size(get<1>(typename TensorSharedSF::layout_type{})) == 4);
|
||||
static_assert(size(get<0>(typename TensorSF::layout_type{})) == 128 &&
|
||||
size(get<1>(typename TensorSF::layout_type{})) == 4);
|
||||
|
||||
using Tiler_MN = typename TiledCopyG2R::Tiler_MN;
|
||||
auto tiler_mn = Tiler_MN{};
|
||||
static_assert(size<0>(tiler_mn) == 16 && size<1>(tiler_mn) == 128);
|
||||
|
||||
auto tiled_tensor_s = tiled_divide(tensor_s, tiler_mn);
|
||||
auto tiled_tensor_p = tiled_divide(tensor_p, tiler_mn);
|
||||
auto tiled_tensor_d = tiled_divide(tensor_d, tiler_mn);
|
||||
static_assert(size<2>(tiled_tensor_s) == 1);
|
||||
static_assert(size<2>(tiled_tensor_p) == 1);
|
||||
static_assert(size<2>(tiled_tensor_d) == 1);
|
||||
auto squeeze_tiled_tensor_s = take<0, 2>(tiled_tensor_s);
|
||||
auto squeeze_tiled_tensor_p = take<0, 2>(tiled_tensor_p);
|
||||
auto squeeze_tiled_tensor_d = take<0, 2>(tiled_tensor_d);
|
||||
|
||||
using SF_Tiler_MN = typename TiledCopyR2S::Tiler_MN;
|
||||
auto sf_tiler_mn = SF_Tiler_MN{};
|
||||
static_assert(size<0>(sf_tiler_mn) == 16 && size<1>(sf_tiler_mn) == 4);
|
||||
|
||||
auto tiled_tensor_sf = tiled_divide(tensor_sf, sf_tiler_mn);
|
||||
auto tiled_tensor_shared_sf = tiled_divide(tensor_shared_sf, sf_tiler_mn);
|
||||
auto squeeze_tiled_tensor_sf = take<0, 2>(tiled_tensor_sf);
|
||||
auto squeeze_tiled_tensor_shared_sf = take<0, 2>(tiled_tensor_shared_sf);
|
||||
|
||||
constexpr int tile_loop_count = size<1>(tiled_tensor_s);
|
||||
constexpr int rows_in_tile = 16;
|
||||
// We don't need to clear shared memory
|
||||
// clear(squeeze_tiled_tensor_shared_sf);
|
||||
#pragma unroll 4
|
||||
for (int t = 0; t < tile_loop_count; t++) {
|
||||
if (t * rows_in_tile >= m) {
|
||||
break;
|
||||
}
|
||||
auto current_copy_tile_s = tensor<0>(squeeze_tiled_tensor_s(_, t));
|
||||
auto current_copy_tile_p = tensor<0>(squeeze_tiled_tensor_p(_, t));
|
||||
auto current_copy_tile_d = tensor<0>(squeeze_tiled_tensor_d(_, t));
|
||||
auto current_copy_tile_sf = tensor<0>(squeeze_tiled_tensor_sf(_, t));
|
||||
auto current_copy_tile_shared_sf =
|
||||
tensor<0>(squeeze_tiled_tensor_shared_sf(_, t));
|
||||
|
||||
// Global to Register copy
|
||||
auto thr_copy_g2r = tiled_copy_g2r.get_thread_slice(threadIdx.x);
|
||||
auto thr_tile_g2r_s = thr_copy_g2r.partition_S(current_copy_tile_s);
|
||||
auto thr_tile_g2r_p = thr_copy_g2r.partition_S(current_copy_tile_p);
|
||||
auto input_fragment = make_fragment_like(thr_tile_g2r_s);
|
||||
|
||||
// Register to Global copy
|
||||
auto thr_copy_r2g = tiled_copy_r2g.get_thread_slice(threadIdx.x);
|
||||
auto thr_tile_r2g_d = thr_copy_r2g.partition_D(current_copy_tile_d);
|
||||
auto thr_tile_r2g_p = thr_copy_r2g.partition_D(current_copy_tile_p);
|
||||
auto output_fragment = make_fragment_like(thr_tile_r2g_d);
|
||||
|
||||
// Register to Shared copy
|
||||
auto thr_copy_r2s = tiled_copy_r2s.get_thread_slice(threadIdx.x / 2);
|
||||
auto thr_tile_r2s_shared_sf =
|
||||
thr_copy_r2s.partition_D(current_copy_tile_shared_sf);
|
||||
auto shared_sf_fragment = make_fragment_like(thr_tile_r2s_shared_sf);
|
||||
|
||||
// CopyG2R & convert & CopyR2G
|
||||
copy_if(tiled_copy_g2r, thr_tile_g2r_p, thr_tile_g2r_s, input_fragment);
|
||||
uint8_t fp8_sf_val =
|
||||
cvt_warp_fp16_to_mxfp8(input_fragment, output_fragment);
|
||||
copy_if(tiled_copy_r2g, thr_tile_r2g_p, output_fragment, thr_tile_r2g_d);
|
||||
shared_sf_fragment[0] = fp8_sf_val;
|
||||
|
||||
// Before first copy r2s, clear shared memory and wait previous group
|
||||
if (t == 0 && threadIdx.x == 0) {
|
||||
// Wait for the group to have completed reading from shared memory.
|
||||
cuda::ptx::cp_async_bulk_wait_group_read(cuda::ptx::n32_t<0>());
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (threadIdx.x % 2 == 0) {
|
||||
copy(tiled_copy_r2s, shared_sf_fragment, thr_tile_r2s_shared_sf);
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// Wait for shared memory writes to be visible to TMA engine.
|
||||
cuda::ptx::fence_proxy_async(cuda::ptx::space_shared); // b)
|
||||
__syncthreads();
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
cuda::ptx::cp_async_bulk(cuda::ptx::space_global, cuda::ptx::space_shared,
|
||||
squeeze_tiled_tensor_sf.data().get(),
|
||||
squeeze_tiled_tensor_shared_sf.data().get(), 512);
|
||||
// Wait for TMA transfer to have finished reading shared memory.
|
||||
// Create a "bulk async-group" out of the previous bulk copy operation.
|
||||
cuda::ptx::cp_async_bulk_commit_group();
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
template <typename T_IN, typename TiledCopyG2R, typename TiledCopyR2G,
|
||||
typename TiledCopyR2S>
|
||||
__global__ void mxfp8_experts_quant_kernel(
|
||||
const T_IN* input, const int* problem_sizes, const int* expert_offsets,
|
||||
const int* blockscale_offsets, cutlass::float_e4m3_t* quant_output,
|
||||
uint8_t* scale_factor, int groups, TiledCopyG2R tiled_copy_g2r,
|
||||
TiledCopyR2G tiled_copy_r2g, TiledCopyR2S tiled_copy_r2s) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000
|
||||
__shared__ __align__(512) uint8_t shared_memory[512];
|
||||
ScaleFactorTileLayout scale_factor_tile_layout{};
|
||||
auto scale_factor_shared =
|
||||
make_tensor(make_smem_ptr(shared_memory),
|
||||
scale_factor_tile_layout); // ((_32,_4), _4):((_16,_4), _1)
|
||||
// TODO: Transform Groupwise Schedule into a more efficient Schedule
|
||||
for (int g = 0; g < groups; g++) {
|
||||
int m = problem_sizes[g * 3 + 0];
|
||||
int k = problem_sizes[g * 3 + 2];
|
||||
int64_t expert_offset = static_cast<int64_t>(expert_offsets[g]);
|
||||
int64_t blockscale_offset = static_cast<int64_t>(blockscale_offsets[g]);
|
||||
|
||||
auto input_tensor = make_tensor(
|
||||
make_gmem_ptr(input + expert_offset * k),
|
||||
make_layout(make_shape(m, k),
|
||||
LayoutRight{})); // (M, K):(K, 1) half_t/bfloat16_t
|
||||
|
||||
auto quant_output_tensor = make_tensor(
|
||||
make_gmem_ptr(quant_output + expert_offset * k),
|
||||
make_layout(make_shape(m, k),
|
||||
LayoutRight{})); // (M, K):(K, 1) cutlass::float_e4m3_t
|
||||
|
||||
auto scale_factor_shape = make_shape(ceil_div(m, 128) * 128, k / 32);
|
||||
auto scale_factor_layout = tile_to_shape(scale_factor_tile_layout,
|
||||
scale_factor_shape, LayoutRight{});
|
||||
// layout<0>(layout<0>(scale_factor_layout)) (_32,_4):(_16,_4) -- static
|
||||
// layout<1>(layout<0>(scale_factor_layout)) M_align_128 / 128 -- dynamic
|
||||
// shape dynamic stride layout<0>(layout<1>(scale_factor_layout)) _4:_1 --
|
||||
// static layout<1>(layout<1>(scale_factor_layout)) (K / 32) / 4 : _512 --
|
||||
// dynamic shape static stride
|
||||
|
||||
// Reshape to zipped layout for 1D indexing
|
||||
auto zipped_scale_factor_layout = make_layout(
|
||||
make_layout(layout<0>(layout<0>(scale_factor_layout)),
|
||||
layout<0>(layout<1>(scale_factor_layout))),
|
||||
make_layout(
|
||||
layout<1>(layout<0>(scale_factor_layout)),
|
||||
layout<1>(layout<1>(
|
||||
scale_factor_layout)))); // (((_32,_4),_4),(M_align_128 /
|
||||
// 128,(K / 32) /
|
||||
// 4)):(((_16,_4),_1),(?,_512))
|
||||
|
||||
auto scale_factor_tensor =
|
||||
make_tensor(make_gmem_ptr(scale_factor + blockscale_offset * (k / 32)),
|
||||
zipped_scale_factor_layout);
|
||||
|
||||
// Used for cases where M is not divisible by 128 (most scenarios).
|
||||
auto input_shape = shape(input_tensor); // (M, K):(K, 1)
|
||||
auto identity_tensor = make_identity_tensor(input_shape);
|
||||
auto predict_tensor = cute::lazy::transform(
|
||||
identity_tensor, [&](auto c) { return elem_less(c, input_shape); });
|
||||
|
||||
// (_128, _128)
|
||||
auto tiler = make_shape(Int<BLOCK_M>{}, Int<BLOCK_K>{});
|
||||
|
||||
auto tiled_input_tensor = zipped_divide(
|
||||
input_tensor, tiler); // ((128, 128), (cdiv(M, 128), cdiv(K, 128)))
|
||||
auto tiled_quant_output_tensor =
|
||||
zipped_divide(quant_output_tensor,
|
||||
tiler); // ((128, 128), (cdiv(M, 128), cdiv(K, 128)))
|
||||
auto tiled_predict_tensor = zipped_divide(
|
||||
predict_tensor, tiler); // ((128, 128), (cdiv(M, 128), cdiv(K, 128)))
|
||||
|
||||
auto total_tiles =
|
||||
size<1>(tiled_input_tensor); // cdiv(M, 128) * cdiv(K, 128)
|
||||
decltype(total_tiles) blk_offset = blockIdx.x;
|
||||
while (blk_offset < total_tiles) {
|
||||
auto current_input_tile = tensor<0>(tiled_input_tensor(_, blk_offset));
|
||||
auto current_quant_output_tile =
|
||||
tensor<0>(tiled_quant_output_tensor(_, blk_offset));
|
||||
auto current_predict_tile =
|
||||
tensor<0>(tiled_predict_tensor(_, blk_offset));
|
||||
auto current_scale_factor_tile =
|
||||
tensor<0>(scale_factor_tensor(_, blk_offset));
|
||||
|
||||
mxfp8_experts_quant_tile<
|
||||
decltype(current_input_tile), decltype(current_predict_tile),
|
||||
decltype(current_quant_output_tile), decltype(scale_factor_shared),
|
||||
decltype(current_scale_factor_tile), TiledCopyG2R, TiledCopyR2G,
|
||||
TiledCopyR2S>(current_input_tile, current_predict_tile,
|
||||
current_quant_output_tile, scale_factor_shared,
|
||||
current_scale_factor_tile, m, tiled_copy_g2r,
|
||||
tiled_copy_r2g, tiled_copy_r2s);
|
||||
blk_offset += gridDim.x;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
template <typename T_IN>
|
||||
void launch_mxfp8_experts_quant(const torch::stable::Tensor& input,
|
||||
const torch::stable::Tensor& problem_sizes,
|
||||
const torch::stable::Tensor& expert_offsets,
|
||||
const torch::stable::Tensor& blockscale_offsets,
|
||||
torch::stable::Tensor& quant_output,
|
||||
torch::stable::Tensor& scale_factor) {
|
||||
ThrLayout thr_layout{};
|
||||
ValLayout val_layout{};
|
||||
SfR2SThrLayout r2s_thr_layout{};
|
||||
SfR2SValLayout r2s_val_layout{};
|
||||
|
||||
using CopyOpG2R =
|
||||
UniversalCopy<cutlass::AlignedArray<T_IN, size(val_layout)>>;
|
||||
using CopyAtomG2R = cute::Copy_Atom<CopyOpG2R, T_IN>;
|
||||
auto tiled_copy_g2r = cute::make_tiled_copy(
|
||||
CopyAtomG2R{}, thr_layout, val_layout); // Tiler_MN: (16, 128)
|
||||
|
||||
using CopyOpR2G = UniversalCopy<
|
||||
cutlass::AlignedArray<cutlass::float_e4m3_t, size(val_layout)>>;
|
||||
using CopyAtomR2G = cute::Copy_Atom<CopyOpR2G, cutlass::float_e4m3_t>;
|
||||
auto tiled_copy_r2g = cute::make_tiled_copy(
|
||||
CopyAtomR2G{}, thr_layout, val_layout); // Tiler_MN: (16, 128)
|
||||
|
||||
using CopyOpR2S =
|
||||
UniversalCopy<cutlass::AlignedArray<uint8_t, size(r2s_val_layout)>>;
|
||||
using CopyAtomR2S = cute::Copy_Atom<CopyOpR2S, uint8_t>;
|
||||
auto tiled_copy_r2s = cute::make_tiled_copy(
|
||||
CopyAtomR2S{}, r2s_thr_layout, r2s_val_layout); // Tiler_MN: (16, 4)
|
||||
|
||||
int max_active_blocks_per_sm = -1;
|
||||
STD_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
|
||||
&max_active_blocks_per_sm,
|
||||
mxfp8_experts_quant_kernel<T_IN, decltype(tiled_copy_g2r),
|
||||
decltype(tiled_copy_r2g),
|
||||
decltype(tiled_copy_r2s)>,
|
||||
THREAD_BLOCK_SIZE, 0));
|
||||
|
||||
dim3 grid(get_device_prop()->multiProcessorCount * max_active_blocks_per_sm,
|
||||
1, 1);
|
||||
dim3 block(THREAD_BLOCK_SIZE, 1, 1);
|
||||
int num_experts = static_cast<int>(problem_sizes.size(0));
|
||||
auto stream = get_current_cuda_stream(input.get_device_index());
|
||||
mxfp8_experts_quant_kernel<T_IN, decltype(tiled_copy_g2r),
|
||||
decltype(tiled_copy_r2g), decltype(tiled_copy_r2s)>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<const T_IN*>(input.data_ptr()),
|
||||
reinterpret_cast<const int*>(problem_sizes.data_ptr()),
|
||||
reinterpret_cast<const int*>(expert_offsets.data_ptr()),
|
||||
reinterpret_cast<const int*>(blockscale_offsets.data_ptr()),
|
||||
reinterpret_cast<cutlass::float_e4m3_t*>(quant_output.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(scale_factor.data_ptr()), num_experts,
|
||||
tiled_copy_g2r, tiled_copy_r2g, tiled_copy_r2s);
|
||||
}
|
||||
|
||||
} // namespace expert_specialization
|
||||
@@ -3,6 +3,9 @@
|
||||
#include <torch/csrc/stable/library.h>
|
||||
#include <torch/csrc/stable/tensor.h>
|
||||
|
||||
#include <optional>
|
||||
#include <string>
|
||||
|
||||
void per_token_group_quant_fp8(const torch::stable::Tensor& input,
|
||||
torch::stable::Tensor& output_q,
|
||||
torch::stable::Tensor& output_s,
|
||||
@@ -185,11 +188,12 @@ torch::stable::Tensor hadacore_transform(torch::stable::Tensor& x,
|
||||
|
||||
// Layernorm kernels (shared CUDA/ROCm)
|
||||
void rms_norm(torch::stable::Tensor& out, torch::stable::Tensor& input,
|
||||
torch::stable::Tensor& weight, double epsilon);
|
||||
std::optional<torch::stable::Tensor> weight, double epsilon);
|
||||
|
||||
void fused_add_rms_norm(torch::stable::Tensor& input,
|
||||
torch::stable::Tensor& residual,
|
||||
torch::stable::Tensor& weight, double epsilon);
|
||||
std::optional<torch::stable::Tensor> weight,
|
||||
double epsilon);
|
||||
|
||||
// Layernorm-quant kernels (shared CUDA/ROCm)
|
||||
void rms_norm_static_fp8_quant(torch::stable::Tensor& out,
|
||||
@@ -297,7 +301,8 @@ void fused_minimax_m3_qknorm_rope_kv_insert(
|
||||
std::optional<torch::stable::Tensor> kv_cache,
|
||||
std::optional<torch::stable::Tensor> index_cache, int64_t block_size,
|
||||
std::optional<torch::stable::Tensor> q_out,
|
||||
std::optional<torch::stable::Tensor> index_q_out);
|
||||
std::optional<torch::stable::Tensor> index_q_out,
|
||||
const std::string& kv_cache_dtype);
|
||||
|
||||
// Sampler kernels (shared CUDA/ROCm)
|
||||
void apply_repetition_penalties_(
|
||||
|
||||
@@ -24,13 +24,21 @@ __device__ inline void vectorize_with_alignment(
|
||||
ScaOp&& scalar_op) { // InT -> OutT
|
||||
static_assert(VEC_SIZE > 0 && (VEC_SIZE & (VEC_SIZE - 1)) == 0,
|
||||
"VEC_SIZE must be a positive power-of-two");
|
||||
constexpr int WIDTH = VEC_SIZE * sizeof(InT); // eg: 64 B
|
||||
constexpr int WIDTH = VEC_SIZE * sizeof(InT); // eg: 16 B
|
||||
constexpr int OUT_WIDTH = VEC_SIZE * sizeof(OutT); // eg: 16 B
|
||||
uintptr_t addr = reinterpret_cast<uintptr_t>(in);
|
||||
uintptr_t out_addr = reinterpret_cast<uintptr_t>(out);
|
||||
|
||||
// fast path when the whole region is already aligned
|
||||
// Note: currently the output is guaranteed to be same as the input, so we
|
||||
// don't check it here, comments here just for future reference.
|
||||
bool can_vec = ((addr & (WIDTH - 1)) == 0) && ((len & (VEC_SIZE - 1)) == 0);
|
||||
// fast path when input and output are both fully aligned. The vector
|
||||
// load/store below go through vec_n_t<T, VEC_SIZE>, declared
|
||||
// __align__(VEC_SIZE * sizeof(T)), so each side must be aligned to its
|
||||
// own vector width. out is NOT generally co-aligned with in: e.g.
|
||||
// reshape_and_cache_flash writes KV-cache rows whose byte offset is a
|
||||
// multiple of head_size, which for head sizes that are not a multiple
|
||||
// of VEC_SIZE puts some rows off the vector-width boundary.
|
||||
bool can_vec = ((addr & (WIDTH - 1)) == 0) &&
|
||||
((out_addr & (OUT_WIDTH - 1)) == 0) &&
|
||||
((len & (VEC_SIZE - 1)) == 0);
|
||||
if (can_vec) {
|
||||
int num_vec = len / VEC_SIZE;
|
||||
|
||||
@@ -55,6 +63,16 @@ __device__ inline void vectorize_with_alignment(
|
||||
prefix_elems /= sizeof(InT);
|
||||
prefix_elems = min(prefix_elems, len); // 0 ≤ prefix < 16
|
||||
|
||||
// the prefix below aligns in; if that does not also align out (their
|
||||
// addresses differ modulo the vector width), vectorizing is impossible
|
||||
// and the whole copy must stay scalar.
|
||||
if (((out_addr + prefix_elems * sizeof(OutT)) & (OUT_WIDTH - 1)) != 0) {
|
||||
for (int i = tid; i < len; i += stride) {
|
||||
scalar_op(out[i], in[i]);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
// 1. prefill the when it is unsafe to vectorize
|
||||
for (int i = tid; i < prefix_elems; i += stride) {
|
||||
scalar_op(out[i], in[i]);
|
||||
|
||||
@@ -304,9 +304,17 @@ __global__ void per_token_group_quant_8bit_packed_register_kernel(
|
||||
const int sf_k_idx = blockIdx.x * kGroupsPerBlockX + sf_k_local;
|
||||
const int mn_idx = blockIdx.y * kRowsPerBlock + row_local;
|
||||
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
asm volatile("griddepcontrol.wait;");
|
||||
#endif
|
||||
|
||||
if (mn_idx >= tma_aligned_mn) {
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
asm volatile("griddepcontrol.launch_dependents;");
|
||||
#endif
|
||||
return;
|
||||
}
|
||||
|
||||
const bool is_valid_group = (mn_idx < mn) && (sf_k_idx < groups_per_row);
|
||||
|
||||
// Load 16 input elements (32 B) into registers as two adjacent uint4
|
||||
@@ -417,6 +425,10 @@ __global__ void per_token_group_quant_8bit_packed_register_kernel(
|
||||
static_cast<int64_t>(mn_idx) * groups_per_row * GROUP_SIZE +
|
||||
sf_k_idx * GROUP_SIZE + lane_id * VEC_SIZE;
|
||||
*reinterpret_cast<uint4*>(group_output) = packed_out;
|
||||
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
asm volatile("griddepcontrol.launch_dependents;");
|
||||
#endif
|
||||
}
|
||||
|
||||
// Public entry point: register-resident packed quant kernel.
|
||||
@@ -495,23 +507,54 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
|
||||
|
||||
auto dst_type = output_q.scalar_type();
|
||||
|
||||
#define LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, KX, RY) \
|
||||
do { \
|
||||
dim3 grid(static_cast<unsigned int>(blocks_x), \
|
||||
static_cast<unsigned int>(blocks_y)); \
|
||||
dim3 block(num_threads); \
|
||||
per_token_group_quant_8bit_packed_register_kernel<T, DST_DTYPE, 128, KX, \
|
||||
RY> \
|
||||
<<<grid, block, 0, stream>>>( \
|
||||
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
|
||||
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
|
||||
static_cast<int>(padded_groups_per_row), \
|
||||
static_cast<int>(groups_per_row), static_cast<int>(mn), \
|
||||
static_cast<int>(output_q_mn_extent), \
|
||||
static_cast<int>(tma_aligned_mn), num_scale_elems, \
|
||||
static_cast<float>(eps), static_cast<float>(min_8bit), \
|
||||
static_cast<float>(max_8bit)); \
|
||||
} while (0)
|
||||
// PDL (Programmatic Dependent Launch) is NVIDIA-only; ROCm/HIP has no
|
||||
// equivalent launch attribute, so fall back to a classic launch there.
|
||||
#ifndef USE_ROCM
|
||||
#define LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, KX, RY) \
|
||||
do { \
|
||||
cudaLaunchConfig_t config = {}; \
|
||||
config.gridDim = dim3(static_cast<unsigned int>(blocks_x), \
|
||||
static_cast<unsigned int>(blocks_y)); \
|
||||
config.blockDim = dim3(num_threads); \
|
||||
config.dynamicSmemBytes = 0; \
|
||||
config.stream = stream; \
|
||||
cudaLaunchAttribute attrs[1]; \
|
||||
attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; \
|
||||
attrs[0].val.programmaticStreamSerializationAllowed = 1; \
|
||||
config.numAttrs = 1; \
|
||||
config.attrs = attrs; \
|
||||
cudaLaunchKernelEx( \
|
||||
&config, \
|
||||
per_token_group_quant_8bit_packed_register_kernel<T, DST_DTYPE, 128, \
|
||||
KX, RY>, \
|
||||
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
|
||||
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
|
||||
static_cast<int>(padded_groups_per_row), \
|
||||
static_cast<int>(groups_per_row), static_cast<int>(mn), \
|
||||
static_cast<int>(output_q_mn_extent), \
|
||||
static_cast<int>(tma_aligned_mn), num_scale_elems, \
|
||||
static_cast<float>(eps), static_cast<float>(min_8bit), \
|
||||
static_cast<float>(max_8bit)); \
|
||||
} while (0)
|
||||
#else
|
||||
#define LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, KX, RY) \
|
||||
do { \
|
||||
dim3 grid(static_cast<unsigned int>(blocks_x), \
|
||||
static_cast<unsigned int>(blocks_y)); \
|
||||
dim3 block(num_threads); \
|
||||
per_token_group_quant_8bit_packed_register_kernel<T, DST_DTYPE, 128, KX, \
|
||||
RY> \
|
||||
<<<grid, block, 0, stream>>>( \
|
||||
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
|
||||
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
|
||||
static_cast<int>(padded_groups_per_row), \
|
||||
static_cast<int>(groups_per_row), static_cast<int>(mn), \
|
||||
static_cast<int>(output_q_mn_extent), \
|
||||
static_cast<int>(tma_aligned_mn), num_scale_elems, \
|
||||
static_cast<float>(eps), static_cast<float>(min_8bit), \
|
||||
static_cast<float>(max_8bit)); \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#define LAUNCH_REG_KERNEL(T, DST_DTYPE) \
|
||||
do { \
|
||||
|
||||
@@ -308,22 +308,6 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
|
||||
"awq_dequantize(Tensor _kernel, Tensor _scaling_factors, "
|
||||
"Tensor _zeros, SymInt split_k_iters, int thx, int thy) -> Tensor");
|
||||
|
||||
// Expert-specialization mxfp8 blockscaled grouped quantization (SM100+).
|
||||
ops.def(
|
||||
"mxfp8_experts_quant("
|
||||
" Tensor input, Tensor problem_sizes, Tensor expert_offsets,"
|
||||
" Tensor blockscale_offsets, Tensor! quant_output, Tensor! scale_factor)"
|
||||
" -> ()");
|
||||
// conditionally compiled so impl registration is in source file
|
||||
|
||||
// Expert-specialization mxfp8 blockscaled grouped GEMM (SM100+).
|
||||
ops.def(
|
||||
"cutlass_mxfp8_grouped_mm("
|
||||
" Tensor a, Tensor b, Tensor sfa, Tensor sfb, Tensor! out,"
|
||||
" Tensor problem_sizes, Tensor expert_offsets, Tensor blockscale_offsets)"
|
||||
" -> ()");
|
||||
// conditionally compiled so impl registration is in source file
|
||||
|
||||
// DeepSeek V3 fused A GEMM (SM 9.0+, bf16 only, 1-16 tokens).
|
||||
// conditionally compiled so impl registration is in source file
|
||||
ops.def(
|
||||
@@ -369,12 +353,13 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
|
||||
|
||||
// Apply Root Mean Square (RMS) Normalization to the input tensor.
|
||||
ops.def(
|
||||
"rms_norm(Tensor! result, Tensor input, Tensor weight, float epsilon) -> "
|
||||
"rms_norm(Tensor! result, Tensor input, Tensor? weight, float epsilon) "
|
||||
"-> "
|
||||
"()");
|
||||
|
||||
// In-place fused Add and RMS Normalization.
|
||||
ops.def(
|
||||
"fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor weight, "
|
||||
"fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor? weight, "
|
||||
"float epsilon) -> ()");
|
||||
|
||||
// Layernorm-quant
|
||||
@@ -471,7 +456,8 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
|
||||
"int num_index_heads, "
|
||||
"Tensor? slot_mapping, Tensor? index_slot_mapping, "
|
||||
"Tensor!? kv_cache, Tensor!? index_cache, "
|
||||
"int block_size, Tensor!? q_out, Tensor!? index_q_out) -> ()");
|
||||
"int block_size, Tensor!? q_out, Tensor!? index_q_out, "
|
||||
"str kv_cache_dtype) -> ()");
|
||||
|
||||
// Apply repetition penalties to logits in-place.
|
||||
ops.def(
|
||||
|
||||
+3
-3
@@ -34,11 +34,11 @@ torch::Tensor weak_ref_tensor(torch::Tensor& tensor) {
|
||||
// rms_norm and fused_add_rms_norm declarations also exist in
|
||||
// csrc/libtorch_stable/ops.h (torch::stable ABI for CUDA). They remain here
|
||||
// because the CPU build still uses these torch::Tensor declarations.
|
||||
void rms_norm(torch::Tensor& out, torch::Tensor& input, torch::Tensor& weight,
|
||||
double epsilon);
|
||||
void rms_norm(torch::Tensor& out, torch::Tensor& input,
|
||||
std::optional<torch::Tensor> weight, double epsilon);
|
||||
|
||||
void fused_add_rms_norm(torch::Tensor& input, torch::Tensor& residual,
|
||||
torch::Tensor& weight, double epsilon);
|
||||
std::optional<torch::Tensor> weight, double epsilon);
|
||||
|
||||
// rotary_embedding also exist in csrc/libtorch_stable/ops.h (torch::stable
|
||||
// ABI for CUDA). It remains here because the CPU build still uses these
|
||||
|
||||
@@ -131,8 +131,8 @@ CMD ["/bin/bash"]
|
||||
# never included in the final runtime image (mirrors ROCm's build_rixl stage).
|
||||
FROM vllm-base AS ucx-nixl-build
|
||||
|
||||
ARG UCX_VERSION=e5d98879705239d254ede40b4a52891850cb5349
|
||||
ARG NIXL_VERSION=0.7.0
|
||||
ARG UCX_VERSION=v1.21.0-rc2
|
||||
ARG NIXL_VERSION=0.10.1
|
||||
|
||||
# Build-time only: compiler, autotools, and verbs dev headers
|
||||
RUN apt-get update -y && apt-get install -y --no-install-recommends \
|
||||
@@ -167,8 +167,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
|
||||
FROM vllm-base AS vllm-openai
|
||||
|
||||
ARG UCX_VERSION=e5d98879705239d254ede40b4a52891850cb5349
|
||||
ARG NIXL_VERSION=0.7.0
|
||||
ARG NIXL_VERSION=0.10.1
|
||||
|
||||
# Copy compiled UCX runtime libraries and the pre-built NIXL wheel.
|
||||
# No compiler or autotools are installed in this stage.
|
||||
@@ -192,7 +191,8 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
ibverbs-providers \
|
||||
librdmacm1t64 \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& uv pip install --no-deps /tmp/nixl_wheels/nixl-*.whl \
|
||||
&& uv pip install --no-deps /tmp/nixl_wheels/nixl*.whl \
|
||||
&& uv pip install nixl==${NIXL_VERSION} \
|
||||
&& rm -rf /tmp/nixl_wheels
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
|
||||
@@ -214,9 +214,9 @@ hardware and configuration.
|
||||
| Backend | Description | Dtypes | Compute Cap. | Notes |
|
||||
| ------- | ----------- | ------ | ------------ | ----- |
|
||||
| `FLASH_ATTN`‡ | FlashAttention varlen (FA2/FA3/FA4) | fp16, bf16 | Any | FA4 on SM100+, FA3 on SM90, FA2 otherwise |
|
||||
| `TRTLLM_RAGGED` | TensorRT-LLM ragged attention | fp16, bf16 | 10.x | DeepSeek R1 dims only |
|
||||
| `FLASHINFER` | FlashInfer CUTLASS backend | fp16, bf16 | 10.x | DeepSeek R1 dims only |
|
||||
| `TOKENSPEED_MLA` | | fp16, bf16 | 10.x | DeepSeek R1 dims only |
|
||||
| `TRTLLM_RAGGED` | TensorRT-LLM ragged attention | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) or (qk_nope_head_dim=192, qk_rope_head_dim=64, v_head_dim=256) only |
|
||||
| `FLASHINFER` | FlashInfer CUTLASS backend | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) only |
|
||||
| `TOKENSPEED_MLA` | | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) only |
|
||||
|
||||
> **‡** Automatic selection tries FlashAttention first. On Blackwell
|
||||
> (SM100), the fallback order is TRT-LLM Ragged, FlashInfer, then
|
||||
@@ -252,6 +252,6 @@ default on NVIDIA is `FLASHMLA_SPARSE_DSV4`.
|
||||
|
||||
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
|
||||
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ |
|
||||
| `FLASHINFER_MLA_SPARSE_DSV4` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
|
||||
| `FLASHMLA_SPARSE_DSV4` | bf16 | `auto`, `bfloat16`, `fp8_ds_mla`, `fp8` | 256 | 512 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x-10.x |
|
||||
| `FLASHINFER_MLA_SPARSE_DSV4` | bf16 | `auto`, `bfloat16`, `fp8` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
|
||||
| `FLASHMLA_SPARSE_DSV4` | bf16 | `auto`, `fp8_ds_mla`, `fp8` | 256 | 512 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x-10.x |
|
||||
| `ROCM_FLASHMLA_SPARSE_DSV4` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
The [CUDA Graphs](cuda_graphs.md) infrastructure in vLLM primarily targets the **decoder** (language model) forward pass. vLLM also supports capturing the **encoder** (vision transformer) forward pass as CUDA Graphs, independently from the decoder. This is based on <https://github.com/vllm-project/vllm/pull/35963>.
|
||||
|
||||
For two-tower vision encoders (e.g., DeepSeek-OCR's SAM + CLIP with dynamic tiling), a **dual-path graph** mode captures two independent sets of CUDA graphs — one for the global image path and one for the local patch path — enabling independent budget selection and partial eager fallback per path. This is based on <https://github.com/vllm-project/vllm/pull/43586>.
|
||||
|
||||
!!! note
|
||||
Encoder CUDA Graphs are orthogonal to decoder CUDA Graphs — both can be enabled simultaneously. Encoder graphs capture the vision encoder execution (e.g., ViT in Qwen3-VL), while decoder graphs capture the language model execution as described in the [CUDA Graphs design document](cuda_graphs.md).
|
||||
|
||||
@@ -11,6 +13,8 @@ Vision encoder inference incurs CUDA kernel launch overhead on the host side. Th
|
||||
|
||||
Encoder CUDA Graphs eliminate this overhead by pre-capturing the full encoder forward pass at multiple token budget levels during model initialization, then replaying the appropriate graph at runtime.
|
||||
|
||||
For two-tower vision encoders such as DeepSeek-OCR (SAM + CLIP with dynamic tiling), the global image path and local patch path have independent token profiles (272 tokens per global image vs. 100 tokens per local patch). Capturing a single monolithic graph for both paths would significantly reduce packing efficiency. The dual-path graph mode captures each path as a separate set of budgets, allowing the manager to pack and replay each path independently.
|
||||
|
||||
## Design
|
||||
|
||||
The encoder CUDA Graph system uses a **budget-based capture/replay** strategy, managed by [EncoderCudaGraphManager][vllm.v1.worker.encoder_cudagraph.EncoderCudaGraphManager]. The system contains the following core components:
|
||||
@@ -37,10 +41,14 @@ class BudgetGraphMetadata:
|
||||
|
||||
Budgets are auto-generated as power-of-2 levels from a model-provided range via `get_encoder_cudagraph_budget_range()`, with the maximum budget always included even if it does not fall on a power-of-2 boundary. Budgets can also be explicitly specified by the user via `encoder_cudagraph_token_budgets` in `CompilationConfig`.
|
||||
|
||||
When `EncoderCudaGraphConfig.enable_dual_path_graph` is `True`, the manager generates two independent budget lists — `global_token_budgets` (multiples of `global_token_per_image`) and `local_token_budgets` (multiples of `local_token_per_patch`) — and stores captured graphs under `budget_graphs["global"]` and `budget_graphs["local"]` respectively.
|
||||
|
||||
### Greedy bin-packing at runtime
|
||||
|
||||
When a batch of images arrives, the manager sorts images by output token count (smallest first) and greedily packs as many images as possible into each sub-batch while staying within the **largest** token budget and the maximum batch size. Once a sub-batch is finalized (the next image would overflow either constraint), the manager finds the **smallest** budget that fits the sub-batch's total tokens and replays the corresponding CUDA Graph. This repeats until the batch is exhausted. Images that exceed all budgets fall back to eager execution.
|
||||
|
||||
For dual-path models, the manager routes to `_execute_local_dual_path()`, which constrains both global and local token budgets simultaneously during packing (see [Dual-Path graph capture](#dual-path-graph-capture)).
|
||||
|
||||
For each graph replay:
|
||||
|
||||
1. Call `prepare_encoder_cudagraph_replay_buffers()` to compute buffer values (including `pixel_values` and precomputed metadata) from actual batch inputs.
|
||||
@@ -48,6 +56,42 @@ For each graph replay:
|
||||
3. Replay the CUDA Graph.
|
||||
4. Clone outputs from `output_buffer` (cloning is necessary since the buffer is reused across replays).
|
||||
|
||||
### Dual-Path graph capture
|
||||
|
||||
For two-tower vision encoders (e.g., DeepSeek-OCR), the `EncoderCudaGraphConfig` sets `enable_dual_path_graph=True` and provides `global_token_per_image` / `local_token_per_patch`. The manager captures two independent sets of CUDA graphs — one for the **global** image path and one for the **local** patch path — stored under `budget_graphs["global"]` and `budget_graphs["local"]` respectively.
|
||||
|
||||
**Budget generation.** Two separate budget lists are generated:
|
||||
|
||||
* `global_token_budgets` — power-of-2 multiples of `global_token_per_image` (e.g., `[272, 544, 1088, 2176, 4352, 8704, 13824]` for DeepSeek-OCR).
|
||||
* `local_token_budgets` — power-of-2 multiples of `local_token_per_patch` (e.g., `[0, 100, 200, 400, 800, 1600, 3200, 6400, 12800]` for DeepSeek-OCR). A budget of `0` is always included to handle images with no local patches (images ≤ 640×640 that produce only global features).
|
||||
|
||||
Both lists are capped at the same `max_budget`.
|
||||
|
||||
**Dual-path greedy packing.** Each `EncoderItemSpec` provides both `global_output_tokens` (constant per image) and `local_output_tokens` (proportional to the patch count). The dual-path packing algorithm constrains both budgets simultaneously:
|
||||
|
||||
* Sort images by total output tokens (global + local), smallest first.
|
||||
* Greedily pack images: an image is added to the current sub-batch only if both the accumulated global tokens ≤ `max_global_budget` **and** the accumulated local tokens ≤ `max_local_budget`, with the image count ≤ `max_batch_size`.
|
||||
* Once either constraint would overflow, finalize the sub-batch and find the smallest fitting budget **independently** for each path.
|
||||
* Repeat until all images are packed.
|
||||
|
||||
**Partial graph fallback.** After packing, each sub-batch falls into one of four execution scenarios:
|
||||
|
||||
| Global budget | Local budget | Execution |
|
||||
| :---: | :---: | --- |
|
||||
| Found | Found | Both paths use CUDA graph replay |
|
||||
| Found | `None` | Global graph replay + local path skipped (no patches) |
|
||||
| `None` | Found | Global eager fallback + local graph replay |
|
||||
| `None` | `None` | Both paths fall back to eager execution |
|
||||
|
||||
Note that the `0`-budget graph is never actually replayed for local — it signals that local patch processing should be skipped entirely.
|
||||
|
||||
**Buffer keys per path.** Global and local paths use different buffer keys. For DeepSeek-OCR, the global path uses `pixel_values` (full images, shape `[B, 3, 1280, 1280]`) while the local path uses `images_crop` (patches, shape `[P, 3, 1024, 1024]`). The manager iterates over each captured graph's own `input_buffers.keys()` rather than a shared `buffer_keys` list, so both paths can use different buffers.
|
||||
|
||||
**Post-processing.** The `postprocess_encoder_output` method receives a `local_output` parameter (a tensor or `None`) containing the local-path encoder output. The model is responsible for assembling global and local features into the final per-image embedding. For DeepSeek-OCR, this means reshaping the global output into `[B, 272, n_embed]`, the local output into `[P, 100, n_embed]`, assembling patch grids with newline tokens, and concatenating `[patches_grid, global, view_separator]` for each image.
|
||||
|
||||
!!! note
|
||||
The dual-path design enables partial CUDA graph coverage — one path can hit while the other falls back to eager. This avoids wasted compute on zero-padded patch buffers for untiled images and avoids graph invalidation caused by variable `crop_shape` per image.
|
||||
|
||||
### Data-parallel support
|
||||
|
||||
When `mm_encoder_tp_mode="data"`, the manager distributes images across TP ranks using load-balanced assignment via `get_load_balance_assignment`, executes locally on each rank, then gathers results back in the original order via `tensor_model_parallel_all_gather`.
|
||||
@@ -67,29 +111,31 @@ Models opt-in to encoder CUDA Graphs by implementing the [SupportsEncoderCudaGra
|
||||
|
||||
* `get_encoder_cudagraph_config()` — returns static configuration (supported modalities, buffer keys, output hidden size, padding logics, max frames per video).
|
||||
* `get_encoder_cudagraph_budget_range(vllm_config)` — returns `(min_budget, max_budget)` for auto-inference of token budgets.
|
||||
* `get_encoder_cudagraph_item_specs(mm_kwargs)` — returns `list[EncoderItemSpec]` describing each item with its input size and output token count. Replaces the former three separate methods (`get_num_items`, `get_per_item_output_tokens`, `get_per_item_input_sizes`).
|
||||
* `get_encoder_cudagraph_item_specs(mm_kwargs)` — returns `list[EncoderItemSpec]` describing each item with its input size, total output token count (`output_tokens`), and optionally per-path token counts (`global_output_tokens`, `local_output_tokens`) for dual-path models.
|
||||
* `select_encoder_cudagraph_items(mm_kwargs, indices)` — extracts a sub-batch of items by index, used during greedy packing and DP sharding.
|
||||
* `prepare_encoder_cudagraph_capture_inputs(...)` — creates dummy inputs for graph capture. Returns `EncoderCudaGraphCaptureInputs` with a single `values: dict[str, torch.Tensor]` that contains all buffers to be recorded into the graph.
|
||||
* `prepare_encoder_cudagraph_replay_buffers(mm_kwargs, max_batch_size, max_frames_per_batch)` — computes buffer values from actual batch inputs. Returns `EncoderCudaGraphReplayBuffers` with a `values` dict whose keys match `buffer_keys` in the config.
|
||||
* `encoder_cudagraph_forward(inputs: dict[str, torch.Tensor])` — forward pass accepting only fixed-shaped input tensors (the captured `values` dict). Called during both capture and replay. The `pixel_values` tensor is included in `inputs` alongside metadata buffers.
|
||||
* `encoder_eager_forward(mm_kwargs)` — fallback eager forward when no graph fits.
|
||||
* `postprocess_encoder_output(...)` — post-process encoder output, delegates to `scatter_output_slices` by default.
|
||||
* `prepare_encoder_cudagraph_capture_inputs(..., path="default")` — creates dummy inputs for graph capture. The `path` parameter (`"global"` or `"local"`) tells the model which path to generate dummy inputs for. Returns `EncoderCudaGraphCaptureInputs` with a single `values: dict[str, torch.Tensor]` that contains all buffers to be recorded into the graph.
|
||||
* `prepare_encoder_cudagraph_replay_buffers(mm_kwargs, max_batch_size, max_frames_per_batch, path="default")` — computes buffer values from actual batch inputs. The `path` parameter selects which modality keys to extract from `mm_kwargs`. Returns `EncoderCudaGraphReplayBuffers` with a `values` dict whose keys match the captured graph's `input_buffers.keys()`.
|
||||
* `encoder_cudagraph_forward(inputs: dict[str, torch.Tensor], path="default")` — forward pass accepting only fixed-shaped input tensors (the captured `values` dict). Called during both capture and replay. The `path` parameter dispatches to the correct encoder sub-module (e.g., global vs. local path for DeepSeek-OCR).
|
||||
* `encoder_eager_forward(mm_kwargs, path="default")` — fallback eager forward when no graph fits. When `path` is `"global"` or `"local"`, runs only that encoder path without graph capture.
|
||||
* `postprocess_encoder_output(..., local_output=None)` — post-process encoder output. The `local_output` parameter receives the local-path encoder output tensor (or `None`), enabling dual-path models to assemble global and local features into the final per-image embedding.
|
||||
|
||||
!!! note
|
||||
The `SupportsEncoderCudaGraph` protocol is designed to be model-agnostic. New vision encoder models can opt-in by implementing the protocol methods without modifying the manager.
|
||||
|
||||
**Supported models:**
|
||||
|
||||
| Architecture | Models | CG for Image | CG for Video |
|
||||
| ------------ | ------ | ------------ | ------------ |
|
||||
| `Llama4ForConditionalGeneration` | `Llama 4` | ✅︎ | - |
|
||||
| `InternVLChatModel` | `InternVL3.5`, `InternVL3`, `InternVL2.5`, `InternVL2` | ✅︎ | ✅︎ |
|
||||
| `Qwen2VLForConditionalGeneration` | `Qwen2-VL` | ✅︎ | ✅︎ |
|
||||
| `Qwen2_5_VLForConditionalGeneration` | `Qwen2.5-VL` | ✅︎ | ✅︎ |
|
||||
| `Qwen3VLForConditionalGeneration` | `Qwen3-VL` | ✅︎ | ✅︎ |
|
||||
| `Qwen3_5ForConditionalGeneration` | `Qwen3.5` | ✅︎ | ✅︎ |
|
||||
| `Step3VLForConditionalGeneration` | `Step3-VL` | ✅︎ | ❌︎ |
|
||||
| `Glm4vForConditionalGeneration` | `GLM-4.1V, GLM-4.6V-Flash` | ✅︎ | ✅︎ |
|
||||
| Architecture | Models | CG for Image | CG for Video | Dual-Path Graph |
|
||||
| ------------ | ------ | ------------ | ------------ | --------------- |
|
||||
| `DeepseekOCRForCausalLM` | `DeepSeek-OCR` | ✅︎ | ❌︎ | ✅︎ |
|
||||
| `Glm4vForConditionalGeneration` | `GLM-4.1V, GLM-4.6V-Flash` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `InternVLChatModel` | `InternVL3.5`, `InternVL3`, `InternVL2.5`, `InternVL2` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `KimiVLForConditionalGeneration` | `Kimi-VL` | ✅︎ | ❌︎ | ❌︎ |
|
||||
| `Llama4ForConditionalGeneration` | `Llama 4` | ✅︎ | ❌︎ | ❌︎ |
|
||||
| `Qwen2VLForConditionalGeneration` | `Qwen2-VL` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `Qwen2_5_VLForConditionalGeneration` | `Qwen2.5-VL` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `Qwen3VLForConditionalGeneration` | `Qwen3-VL` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `Qwen3_5ForConditionalGeneration` | `Qwen3.5` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `Step3VLForConditionalGeneration` | `Step3-VL` | ✅︎ | ❌︎ | ❌︎ |
|
||||
|
||||
!!! note
|
||||
Encoder CUDA Graphs have currently been tested with `--mm-encoder-attn-backend=FLASH_ATTN` and `--mm-encoder-attn-backend=FLASHINFER` on Blackwell GPUs.
|
||||
@@ -104,6 +150,8 @@ Three fields in `CompilationConfig` control encoder CUDA Graphs:
|
||||
* `encoder_cudagraph_max_vision_items_per_batch` (`int`, default `0`) — maximum number of images/videos per batch during capture. If 0 (default), auto-inferred as `max_budget // min_budget`.
|
||||
* `encoder_cudagraph_max_frames_per_batch` (`int`, default `None`) — maximum number of video frames per batch during capture. If `None` (default), auto-inferred as `encoder_cudagraph_max_vision_items_per_batch * max_frames_per_video` (`max_frames_per_video` is a model-specific value from `EncoderCudaGraphConfig`, computed by `get_max_frames_per_video()` on the model). If we limit the video count per prompt to `0`, it will also be set to `0` (i.e., fall back to image-only mode).
|
||||
|
||||
Dual-path mode is configured at the model level via `EncoderCudaGraphConfig` fields (`enable_dual_path_graph`, `global_token_per_image`, `local_token_per_patch`) — no additional user configuration is required. The manager automatically generates separate budget lists and routes to dual-path execution when the model opts in.
|
||||
|
||||
## Usage guide
|
||||
|
||||
### Image inference
|
||||
|
||||
@@ -127,6 +127,29 @@ PYTHONHASHSEED=0 vllm serve ...
|
||||
- FS thread counts: tune `n_read_threads` and `n_write_threads` to the parallelism your storage can sustain. Reads are latency-sensitive on the prefill path, so prefer more read threads when prefill hit rates are high.
|
||||
- Sharing `root_dir` across runs: runs with the same model, `block_size`, parallelism layout, and dtype share files under the same `<digest>` subdirectory. Changing any of these produces a new subdirectory; old ones are orphaned but harmless. Delete them to reclaim disk.
|
||||
|
||||
## Per-Request Selective Offload
|
||||
|
||||
Individual requests can cap how many of their tokens are eligible for offload by setting `max_offload_tokens` in the request's `kv_transfer_params`. Only the first `max_offload_tokens` tokens of the request are offloaded; blocks beyond that point are skipped on the store path. This is useful when a known prefix (e.g., a system prompt or shared context) is worth caching but later request-specific tokens are not.
|
||||
|
||||
| Key | Type | Notes |
|
||||
| --- | --- | --- |
|
||||
| `max_offload_tokens` | non-negative `int` | Upper bound on tokens to offload for this request. `0` disables offload for the request entirely; omit the key (or set to `None`) for no cap. Non-`int`, negative, or `bool` values are rejected with a warning and treated as no cap. |
|
||||
|
||||
!!! note
|
||||
`max_offload_tokens` is experimental and subject to change.
|
||||
|
||||
Example (OpenAI-compatible completions request):
|
||||
|
||||
```json
|
||||
{
|
||||
"model": "<model>",
|
||||
"prompt": "...",
|
||||
"kv_transfer_params": {
|
||||
"max_offload_tokens": 1024
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Further Reading
|
||||
|
||||
- [vLLM blog: KV Offloading Connector](https://vllm.ai/blog/2026-01-08-kv-offloading-connector) — motivation, architecture (DMA-based async transfer), and benchmarks (TTFT and throughput).
|
||||
|
||||
@@ -205,6 +205,7 @@ the vLLM JSON config.
|
||||
- `load_async` (bool): Enable asynchronous loading for better compute-I/O overlap. Default: `true`.
|
||||
- `enable_cross_layers_blocks` (bool): Enable cross-layer block packing for reduced store operations. Default: `false`.
|
||||
- `lookup_rpc_port` (int): Custom port for the ZMQ lookup RPC socket. Default: `0`.
|
||||
- `cache_prefix` (str): Namespace prepended to every store key. Lets separate deployments share one Mooncake master without polluting each other — instances configured with different prefixes never see each other's cached blocks, even for identical prompts. All instances that should share a prefix cache must use the same value. Default: `""` (no prefix; keys are byte-identical to the unprefixed format).
|
||||
|
||||
## Notes
|
||||
|
||||
|
||||
@@ -49,7 +49,7 @@ To run an AWQ model with vLLM, you can use [TheBloke/Llama-2-7b-Chat-AWQ](https:
|
||||
```bash
|
||||
python examples/deployment/llm_engine_example.py \
|
||||
--model TheBloke/Llama-2-7b-Chat-AWQ \
|
||||
--quantization awq
|
||||
--quantization auto_awq
|
||||
```
|
||||
|
||||
AWQ models are also supported directly through the LLM entrypoint:
|
||||
@@ -70,7 +70,7 @@ AWQ models are also supported directly through the LLM entrypoint:
|
||||
sampling_params = SamplingParams(temperature=0.8, top_p=0.95)
|
||||
|
||||
# Create an LLM.
|
||||
llm = LLM(model="TheBloke/Llama-2-7b-Chat-AWQ", quantization="AWQ")
|
||||
llm = LLM(model="TheBloke/Llama-2-7b-Chat-AWQ", quantization="auto_awq")
|
||||
# Generate texts from the prompts. The output is a list of RequestOutput objects
|
||||
# that contain the prompt, generated text, and other information.
|
||||
outputs = llm.generate(prompts, sampling_params)
|
||||
|
||||
@@ -142,6 +142,10 @@ VLLM_USE_PRECOMPILED=1 VLLM_PRECOMPILED_WHEEL_VARIANT=cpu VLLM_TARGET_DEVICE=cpu
|
||||
=== "IBM Z (S390X)"
|
||||
--8<-- "docs/getting_started/installation/cpu.s390x.inc.md:build-image-from-source"
|
||||
|
||||
## AMD Zen optimizations {#amd-zen-optimizations}
|
||||
|
||||
--8<-- "docs/getting_started/installation/cpu.x86.inc.md:amd-zen-optimizations"
|
||||
|
||||
## Related runtime environment variables
|
||||
|
||||
- `VLLM_CPU_KVCACHE_SPACE`: specify the KV Cache size (e.g, `VLLM_CPU_KVCACHE_SPACE=40` means 40 GiB space for KV cache), larger setting will allow vLLM to run more requests in parallel. This parameter should be set based on the hardware configuration and memory management pattern of users. Default value is `0`.
|
||||
@@ -149,12 +153,14 @@ VLLM_USE_PRECOMPILED=1 VLLM_PRECOMPILED_WHEEL_VARIANT=cpu VLLM_TARGET_DEVICE=cpu
|
||||
- `VLLM_CPU_NUM_OF_RESERVED_CPU`: specify the number of CPU cores which are not dedicated to the OpenMP threads for each rank. The variable only takes effect when VLLM_CPU_OMP_THREADS_BIND is set to `auto`. Default value is `None`. If the value is not set and use `auto` thread binding, no CPU will be reserved for `world_size == 1`, 1 CPU per rank will be reserved for `world_size > 1`.
|
||||
- `CPU_VISIBLE_MEMORY_NODES`: specify visible NUMA memory nodes for vLLM CPU workers, similar to ```CUDA_VISIBLE_DEVICES```. The variable only takes effect when VLLM_CPU_OMP_THREADS_BIND is set to `auto`. The variable provides more control for the auto thread-binding feature, such as masking nodes and changing nodes binding sequence.
|
||||
- `VLLM_CPU_SGL_KERNEL` (x86 only, Experimental): whether to use small-batch optimized kernels for linear layer and MoE layer, especially for low-latency requirements like online serving. The kernels require AMX instruction set, BFloat16 weight type and weight shapes divisible by 32. Default is `0` (False).
|
||||
- `VLLM_ZENTORCH_WEIGHT_PREPACK` (AMD Zen only): when `ZenCpuPlatform` is active, eagerly prepack linear weights into ZenDNN's blocked layout at model load time, eliminating per-inference layout conversion overhead. Default is `1` (enabled). See [AMD Zen optimizations](#amd-zen-optimizations).
|
||||
|
||||
## FAQ
|
||||
|
||||
### Which `dtype` should be used?
|
||||
|
||||
- Currently, vLLM CPU uses model default settings as `dtype`. However, due to unstable float16 support in torch CPU, it is recommended to explicitly set `dtype=bfloat16` if there are any performance or accuracy problem.
|
||||
- On AMD Zen CPUs (`ZenCpuPlatform`), `float16` is **not** supported. Only `bfloat16` and `float32` are accepted; models declared with `float16` are auto-downcast to `bfloat16` at model load time. See [AMD Zen optimizations](#amd-zen-optimizations).
|
||||
|
||||
### How to launch a vLLM service on CPU?
|
||||
|
||||
@@ -227,6 +233,25 @@ By providing MODEL_FILTER and DTYPE_FILTER, only commands for related model ID a
|
||||
ON_CPU=1 SERVING_JSON=serving-tests-cpu-text.json DRY_RUN=1 MODEL_FILTER=meta-llama/Llama-3.1-8B-Instruct DTYPE_FILTER=bfloat16 bash .buildkite/performance-benchmarks/scripts/run-performance-benchmarks.sh
|
||||
```
|
||||
|
||||
### How do I enable AMD Zen optimizations? {#how-do-i-enable-amd-zen-optimizations}
|
||||
|
||||
On an AMD Zen 4 / Zen 5 CPU, install the CPU wheel with the `zen` extra so vLLM pulls the tested `zentorch` version for that release:
|
||||
|
||||
```bash
|
||||
export VLLM_VERSION=$(curl -s https://api.github.com/repos/vllm-project/vllm/releases/latest | jq -r .tag_name | sed 's/^v//')
|
||||
uv pip install "vllm[zen]" --extra-index-url https://wheels.vllm.ai/${VLLM_VERSION}/cpu --index-strategy first-index --torch-backend cpu
|
||||
```
|
||||
|
||||
vLLM auto-detects the platform and routes linear layers through ZenDNN-optimized kernels - no flag needed. To verify it is engaged, look for the platform-selection line in the server's startup logs:
|
||||
|
||||
```bash
|
||||
vllm serve Qwen/Qwen3-0.6B 2>&1 | grep "AMD Zen CPU detected with zentorch installed"
|
||||
```
|
||||
|
||||
For per-backend dispatch details (which kernel each linear layer was bound to), re-run with `VLLM_LOGGING_LEVEL=DEBUG` and grep for `CPU unquantized GEMM dispatch`.
|
||||
|
||||
See [AMD Zen optimizations](#amd-zen-optimizations) for detection rules, supported dtypes, and the `VLLM_ZENTORCH_WEIGHT_PREPACK` knob.
|
||||
|
||||
### How to decide `VLLM_CPU_OMP_THREADS_BIND`?
|
||||
|
||||
- Default `auto` thread-binding is recommended for most cases. Ideally, each OpenMP thread will be bound to a dedicated physical core respectively, threads of each rank will be bound to the same NUMA node respectively, and 1 CPU per rank will be reserved for other vLLM components when `world_size > 1`. If you have any performance problems or unexpected binding behaviours, please try to bind threads as following.
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
<!-- markdownlint-disable MD041 -->
|
||||
<!-- markdownlint-disable MD041 MD051 -->
|
||||
--8<-- [start:installation]
|
||||
|
||||
vLLM supports basic model inferencing and serving on x86 CPU platform, with data types FP32, FP16 and BF16.
|
||||
@@ -200,7 +200,19 @@ docker build -f docker/Dockerfile.cpu \
|
||||
--target vllm-openai .
|
||||
```
|
||||
|
||||
#### Launching the OpenAI server
|
||||
#### Building with AMD Zen optimizations
|
||||
|
||||
For AMD Zen 4 / Zen 5 hosts (`linux/amd64` only), use the `vllm-openai-zen` target. It extends the default `vllm-openai` image and adds `zentorch` via the `vllm[zen]` extra so `ZenCpuPlatform` auto-activates at runtime:
|
||||
|
||||
```bash
|
||||
docker build -f docker/Dockerfile.cpu \
|
||||
--tag vllm-cpu-zen-env \
|
||||
--target vllm-openai-zen .
|
||||
```
|
||||
|
||||
The resulting image accepts the same arguments and environment variables as `vllm-openai` (see [Launching the OpenAI server](#launching-the-openai-server) below); no extra flag is needed to engage Zen optimizations. See [AMD Zen optimizations](cpu.md#amd-zen-optimizations) for runtime behavior and the supported-dtype caveats.
|
||||
|
||||
#### Launching the OpenAI server {#launching-the-openai-server}
|
||||
|
||||
```bash
|
||||
docker run --rm \
|
||||
@@ -216,5 +228,36 @@ docker run --rm \
|
||||
```
|
||||
|
||||
--8<-- [end:build-image-from-source]
|
||||
--8<-- [start:amd-zen-optimizations]
|
||||
|
||||
On AMD Zen CPUs, vLLM auto-selects `ZenCpuPlatform` (a subclass of `CpuPlatform`) which dispatches linear layers through [`zentorch`](https://github.com/amd/ZenDNN-pytorch-plugin)'s ZenDNN-optimized kernels. See the FAQ entry [How do I enable AMD Zen optimizations?](#how-do-i-enable-amd-zen-optimizations) for the install command.
|
||||
|
||||
### Detection rules
|
||||
|
||||
`ZenCpuPlatform` is selected when **all** of the following hold:
|
||||
|
||||
- vLLM is built for CPU
|
||||
- `/proc/cpuinfo` reports `AuthenticAMD` and `avx512`
|
||||
- `import zentorch` succeeds
|
||||
|
||||
Otherwise, vLLM falls back to the default `CpuPlatform` (oneDNN / sgl-kernel paths).
|
||||
|
||||
### Supported dtypes
|
||||
|
||||
`float16` is **not** supported on `ZenCpuPlatform`. `ZenCpuPlatform.supported_dtypes` advertises only `bfloat16` and `float32`, so models declared with `torch_dtype=float16` are auto-downcast to `bfloat16` at load time with the standard `"Your device 'cpu' doesn't support torch.float16. Falling back to torch.bfloat16 for compatibility."` warning emitted from `vllm/config/model.py`.
|
||||
|
||||
### Environment variables
|
||||
|
||||
- `VLLM_ZENTORCH_WEIGHT_PREPACK` (default `1`): eagerly prepacks linear weights into ZenDNN's blocked layout at model load time, eliminating per-inference layout conversion overhead. Set to `0` to disable.
|
||||
|
||||
### Docker
|
||||
|
||||
The `vllm-openai-zen` Docker target (in `docker/Dockerfile.cpu`) extends the default `vllm-openai` image with `vllm[zen]`. Build it with `docker build -f docker/Dockerfile.cpu --target vllm-openai-zen .` — see [Building with AMD Zen optimizations](#building-with-amd-zen-optimizations) for the full command and run instructions.
|
||||
|
||||
### Reference
|
||||
|
||||
For the design rationale, see [RFC #35089: In-Tree AMD Zen CPU Backend via zentorch](https://github.com/vllm-project/vllm/issues/35089).
|
||||
|
||||
--8<-- [end:amd-zen-optimizations]
|
||||
--8<-- [start:extra-information]
|
||||
--8<-- [end:extra-information]
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
# CPU - Intel® Xeon®
|
||||
|
||||
!!! note "AMD Zen CPUs"
|
||||
On AMD Zen 4 / Zen 5 CPUs, AMD Zen optimizations are auto-enabled when the [`zentorch`](https://github.com/amd/ZenDNN-pytorch-plugin) package is installed. All models supported by vLLM on CPU are supported on AMD Zen as well; model compatibility does not change. This page reflects the current CPU reference validation matrix on Intel systems. See [AMD Zen optimizations](../../getting_started/installation/cpu.md#amd-zen-optimizations) for details.
|
||||
|
||||
## Validated Hardware
|
||||
|
||||
| Hardware |
|
||||
|
||||
@@ -374,7 +374,6 @@ th {
|
||||
| `BailingMoeForCausalLM` | Ling | `inclusionAI/Ling-lite-1.5`, `inclusionAI/Ling-plus`, etc. | ✅︎ | ✅︎ |
|
||||
| `BailingMoeV2ForCausalLM` | Ling | `inclusionAI/Ling-mini-2.0`, etc. | ✅︎ | ✅︎ |
|
||||
| `BailingMoeV2_5ForCausalLM` | Ling | `inclusionAI/Ling-2.5-1T`, `inclusionAI/Ring-2.5-1T` | | ✅︎ |
|
||||
| `BambaForCausalLM` | Bamba | `ibm-ai-platform/Bamba-9B-fp8`, `ibm-ai-platform/Bamba-9B` | ✅︎ | ✅︎ |
|
||||
| `BloomForCausalLM` | BLOOM, BLOOMZ, BLOOMChat | `bigscience/bloom`, `bigscience/bloomz`, etc. | | ✅︎ |
|
||||
| `ChatGLMModel`, `ChatGLMForConditionalGeneration` | ChatGLM | `zai-org/chatglm2-6b`, `zai-org/chatglm3-6b`, `thu-coai/ShieldLM-6B-chatglm3`, etc. | ✅︎ | ✅︎ |
|
||||
| `CohereForCausalLM`, `Cohere2ForCausalLM` | Command-R, Command-A | `CohereLabs/c4ai-command-r-v01`, `CohereLabs/c4ai-command-r7b-12-2024`, `CohereLabs/c4ai-command-a-03-2025`, `CohereLabs/command-a-reasoning-08-2025`, etc. | ✅︎ | ✅︎ |
|
||||
@@ -386,7 +385,6 @@ th {
|
||||
| `DeepseekV2ForCausalLM` | DeepSeek-V2 | `deepseek-ai/DeepSeek-V2`, `deepseek-ai/DeepSeek-V2-Chat`, etc. | ✅︎ | ✅︎ |
|
||||
| `DeepseekV3ForCausalLM` | DeepSeek-V3 | `deepseek-ai/DeepSeek-V3`, `deepseek-ai/DeepSeek-R1`, `deepseek-ai/DeepSeek-V3.1`, etc. | ✅︎ | ✅︎ |
|
||||
| `DeepseekV4ForCausalLM` | DeepSeek-V4 | `deepseek-ai/DeepSeek-V4-Flash`, `deepseek-ai/DeepSeek-V4-Pro`, etc. | | ✅︎ |
|
||||
| `Dots1ForCausalLM` | dots.llm1 | `rednote-hilab/dots.llm1.base`, `rednote-hilab/dots.llm1.inst`, etc. | | ✅︎ |
|
||||
| `DotsOCRForCausalLM` | dots_ocr | `rednote-hilab/dots.ocr` | ✅︎ | ✅︎ |
|
||||
| `Ernie4_5ForCausalLM` | Ernie4.5 | `baidu/ERNIE-4.5-0.3B-PT`, etc. | ✅︎ | ✅︎ |
|
||||
| `Ernie4_5_MoeForCausalLM` | Ernie4.5MoE | `baidu/ERNIE-4.5-21B-A3B-PT`, `baidu/ERNIE-4.5-300B-A47B-PT`, etc. | ✅︎ | ✅︎ |
|
||||
@@ -419,6 +417,7 @@ th {
|
||||
| `GritLM` | GritLM | `parasail-ai/GritLM-7B-vllm`. | ✅︎ | ✅︎ |
|
||||
| `Grok1ModelForCausalLM` | Grok1 | `hpcai-tech/grok-1`. | ✅︎ | ✅︎ |
|
||||
| `Grok1ForCausalLM` | Grok2 | `xai-org/grok-2` | ✅︎ | ✅︎ |
|
||||
| `HrmTextForCausalLM` | HRM-Text | `sapientinc/HRM-Text-1B`, etc. | | |
|
||||
| `HunYuanDenseV1ForCausalLM` | Hunyuan Dense | `tencent/Hunyuan-7B-Instruct` | ✅︎ | ✅︎ |
|
||||
| `HunYuanMoEV1ForCausalLM` | Hunyuan-A13B | `tencent/Hunyuan-A13B-Instruct`, `tencent/Hunyuan-A13B-Pretrain`, `tencent/Hunyuan-A13B-Instruct-FP8`, etc. | ✅︎ | ✅︎ |
|
||||
| `HYV3ForCausalLM` | HY3 | `tencent/Hy3-preview-Base`, `tencent/Hy3-preview` | ✅︎ | ✅︎ |
|
||||
|
||||
@@ -28,9 +28,9 @@ For more information on implementation, see [Low Level `layerwise` API](#low-lev
|
||||
Online quantization refers to when a user provides full precision weights and those weights are quantized on-the-fly as they are loaded into the model. The layerwise reloading system handles this by treating online quantization as a **processing** step, which is then handled in an online way both during first-time load and during reload. A typical online quantization method implementation should look like this:
|
||||
|
||||
```python
|
||||
class Fp8OnlineLinearMethod(Fp8LinearMethod):
|
||||
"""Online version of Fp8LinearMethod which loads a full precision checkpoint
|
||||
and quantizes weights during loading."""
|
||||
class Fp8PerTensorOnlineLinearMethod(LinearMethodBase):
|
||||
"""Online version of FP8 per-tensor quantization which loads a full
|
||||
precision checkpoint and quantizes weights during loading."""
|
||||
|
||||
uses_meta_device: bool = True
|
||||
|
||||
|
||||
@@ -2533,15 +2533,17 @@ MODELS_NEED_VIDEO_METADATA = [
|
||||
|
||||
MODELS_SUPPORT_VIT_CUDA_GRAPH = [
|
||||
"llama4",
|
||||
"internvl_chat",
|
||||
"qwen2_vl",
|
||||
"qwen2_5_vl",
|
||||
"qwen3_vl",
|
||||
"qwen3_vl_moe",
|
||||
"qwen2_vl",
|
||||
"kimi_vl",
|
||||
"qwen3_5",
|
||||
"qwen3_5_moe",
|
||||
"internvl_chat",
|
||||
"stepvl",
|
||||
"glm4_1v",
|
||||
"deepseek_ocr",
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -116,7 +116,9 @@
|
||||
}
|
||||
{%- endmacro -%}
|
||||
{%- macro format_argument(argument, escape_keys=True) -%}
|
||||
{%- if argument is string -%}
|
||||
{%- if argument is none -%}
|
||||
{{- 'null' -}}
|
||||
{%- elif argument is string -%}
|
||||
{{- '<|"|>' + argument + '<|"|>' -}}
|
||||
{%- elif argument is boolean -%}
|
||||
{{- 'true' if argument else 'false' -}}
|
||||
@@ -172,18 +174,21 @@
|
||||
{{- '<tool_response|>' -}}
|
||||
{%- endmacro -%}
|
||||
|
||||
{%- set ns = namespace(prev_message_type=None) -%}
|
||||
{#- ===== SETUP ===== -#}
|
||||
{%- set ns = namespace(prev_message_type=None, prev_non_tool_role=None) -%}
|
||||
{%- set loop_messages = messages -%}
|
||||
{%- set enable_thinking = enable_thinking | default(false) -%}
|
||||
{%- set preserve_thinking = preserve_thinking | default(false) -%}
|
||||
{{- bos_token -}}
|
||||
{#- Handle System/Tool Definitions Block -#}
|
||||
{%- if (enable_thinking is defined and enable_thinking) or tools or messages[0]['role'] in ['system', 'developer'] -%}
|
||||
{%- if enable_thinking or tools or (messages and messages[0]['role'] in ['system', 'developer']) -%}
|
||||
{{- '<|turn>system\n' -}}
|
||||
{#- Inject Thinking token at the very top of the FIRST system turn -#}
|
||||
{%- if enable_thinking is defined and enable_thinking -%}
|
||||
{%- if enable_thinking -%}
|
||||
{{- '<|think|>\n' -}}
|
||||
{%- set ns.prev_message_type = 'think' -%}
|
||||
{%- endif -%}
|
||||
{%- if messages[0]['role'] in ['system', 'developer'] -%}
|
||||
{%- if messages and messages[0]['role'] in ['system', 'developer'] -%}
|
||||
{%- if messages[0]['content'] is string -%}
|
||||
{{- messages[0]['content'] | trim -}}
|
||||
{%- elif messages[0]['content'] is sequence -%}
|
||||
@@ -217,31 +222,24 @@
|
||||
{%- if message['role'] != 'tool' -%}
|
||||
{%- set ns.prev_message_type = None -%}
|
||||
{%- set role = 'model' if message['role'] == 'assistant' else message['role'] -%}
|
||||
{#- Detect continuation: suppress duplicate <|turn>model when previous non-tool message was also assistant -#}
|
||||
{%- set prev_nt = namespace(role=None, found=false) -%}
|
||||
{%- if loop.index0 > 0 -%}
|
||||
{%- for j in range(loop.index0 - 1, -1, -1) -%}
|
||||
{%- if not prev_nt.found -%}
|
||||
{%- if loop_messages[j]['role'] != 'tool' -%}
|
||||
{%- set prev_nt.role = loop_messages[j]['role'] -%}
|
||||
{%- set prev_nt.found = true -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- endif -%}
|
||||
{%- set continue_same_model_turn = (role == 'model' and prev_nt.role == 'assistant') -%}
|
||||
{#- Detect continuation using tracked state — O(1) instead of O(n) backward scan -#}
|
||||
{%- set continue_same_model_turn = (role == 'model' and ns.prev_non_tool_role == 'assistant') -%}
|
||||
{%- if not continue_same_model_turn -%}
|
||||
{{- '<|turn>' + role + '\n' }}
|
||||
{%- if role == 'model' and not enable_thinking and not (message.get('reasoning') or message.get('reasoning_content')) -%}
|
||||
{{- '<|channel>thought\n<channel|>' -}}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
|
||||
{#- Render reasoning/reasoning_content as thinking channel -#}
|
||||
{%- set thinking_text = message.get('reasoning') or message.get('reasoning_content') -%}
|
||||
{%- if thinking_text and loop.index0 > ns_turn.last_user_idx and message.get('tool_calls') -%}
|
||||
{%- set thinking_gate = (loop.index0 > ns_turn.last_user_idx) or preserve_thinking -%}
|
||||
{%- if thinking_text and thinking_gate -%}
|
||||
{{- '<|channel>thought\n' + thinking_text + '\n<channel|>' -}}
|
||||
{%- endif -%}
|
||||
|
||||
{%- if message['tool_calls'] -%}
|
||||
{%- for tool_call in message['tool_calls'] -%}
|
||||
{%- if message.get('tool_calls') -%}
|
||||
{%- for tool_call in message.get('tool_calls') -%}
|
||||
{%- set function = tool_call['function'] -%}
|
||||
{{- '<|tool_call>call:' + function['name'] + '{' -}}
|
||||
{%- if function['arguments'] is mapping -%}
|
||||
@@ -251,8 +249,13 @@
|
||||
{%- set ns_args.found_first = true -%}
|
||||
{{- key -}}:{{- format_argument(value, escape_keys=False) -}}
|
||||
{%- endfor -%}
|
||||
{%- elif function['arguments'] is string -%}
|
||||
{{- function['arguments'] -}}
|
||||
{%- elif function['arguments'] is none -%}
|
||||
{%- else -%}
|
||||
{{- raise_exception(
|
||||
"chat_template: tool_calls[].function.arguments must be a "
|
||||
"JSON object (mapping), not a string. Deserialize arguments "
|
||||
"before passing to the template."
|
||||
) -}}
|
||||
{%- endif -%}
|
||||
{{- '}<tool_call|>' -}}
|
||||
{%- endfor -%}
|
||||
@@ -262,7 +265,7 @@
|
||||
{%- set ns_tr_out = namespace(flag=false) -%}
|
||||
{%- if message.get('tool_responses') -%}
|
||||
{#- Legacy: tool_responses embedded on the assistant message (Google/Gemma native) -#}
|
||||
{%- for tool_response in message['tool_responses'] -%}
|
||||
{%- for tool_response in message.get('tool_responses') -%}
|
||||
{{- format_tool_response_block(tool_response['name'] | default('unknown', true), tool_response['response']) -}}
|
||||
{%- set ns_tr_out.flag = true -%}
|
||||
{%- set ns.prev_message_type = 'tool_response' -%}
|
||||
@@ -277,8 +280,8 @@
|
||||
{%- else -%}
|
||||
{%- set follow = loop_messages[k] -%}
|
||||
{#- Resolve tool_call_id to function name -#}
|
||||
{%- set ns_tname = namespace(name=follow.get('name') | default('unknown', true)) -%}
|
||||
{%- for tc in message['tool_calls'] -%}
|
||||
{%- set ns_tname = namespace(name=follow.get('name') or 'unknown') -%}
|
||||
{%- for tc in message.get('tool_calls') -%}
|
||||
{%- if tc.get('id') == follow.get('tool_call_id') -%}
|
||||
{%- set ns_tname.name = tc['function']['name'] -%}
|
||||
{%- endif -%}
|
||||
@@ -296,9 +299,9 @@
|
||||
{%- endfor -%}
|
||||
{{- format_tool_response_block(ns_tname.name, ns_txt.s) -}}
|
||||
{%- for part in tool_body -%}
|
||||
{%- if part.get('type') == 'image' -%}
|
||||
{%- if part.get('type') in ['image', 'image_url'] -%}
|
||||
{{- '<|image|>' -}}
|
||||
{%- elif part.get('type') == 'audio' -%}
|
||||
{%- elif part.get('type') in ['audio', 'input_audio'] -%}
|
||||
{{- '<|audio|>' -}}
|
||||
{%- elif part.get('type') == 'video' -%}
|
||||
{{- '<|video|>' -}}
|
||||
@@ -314,29 +317,26 @@
|
||||
{%- endif -%}
|
||||
|
||||
{%- set captured_content -%}
|
||||
{%- if message['content'] is string -%}
|
||||
{%- if message.get('content') is string -%}
|
||||
{%- if role == 'model' -%}
|
||||
{{- strip_thinking(message['content']) -}}
|
||||
{%- else -%}
|
||||
{{- message['content'] | trim -}}
|
||||
{%- endif -%}
|
||||
{%- elif message['content'] is sequence -%}
|
||||
{%- elif message.get('content') is sequence -%}
|
||||
{%- for item in message['content'] -%}
|
||||
{%- if item['type'] == 'text' -%}
|
||||
{%- if item.get('type') == 'text' -%}
|
||||
{%- if role == 'model' -%}
|
||||
{{- strip_thinking(item['text']) -}}
|
||||
{%- else -%}
|
||||
{{- item['text'] | trim -}}
|
||||
{%- endif -%}
|
||||
{%- elif item['type'] == 'image' -%}
|
||||
{%- elif item.get('type') in ['image', 'image_url'] -%}
|
||||
{{- '<|image|>' -}}
|
||||
{%- set ns.prev_message_type = 'image' -%}
|
||||
{%- elif item['type'] == 'audio' -%}
|
||||
{%- elif item.get('type') in ['audio', 'input_audio'] -%}
|
||||
{{- '<|audio|>' -}}
|
||||
{%- set ns.prev_message_type = 'audio' -%}
|
||||
{%- elif item['type'] == 'video' -%}
|
||||
{%- elif item.get('type') == 'video' -%}
|
||||
{{- '<|video|>' -}}
|
||||
{%- set ns.prev_message_type = 'video' -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- endif -%}
|
||||
@@ -345,19 +345,43 @@
|
||||
{{- captured_content -}}
|
||||
{%- set has_content = captured_content | trim | length > 0 -%}
|
||||
|
||||
{#- Forward-scan: find next non-tool message role for continuation detection -#}
|
||||
{%- set next_nt = namespace(role=None, found=false) -%}
|
||||
{%- for j in range(loop.index0 + 1, loop_messages | length) -%}
|
||||
{%- if not next_nt.found -%}
|
||||
{%- if loop_messages[j]['role'] != 'tool' -%}
|
||||
{%- set next_nt.role = loop_messages[j]['role'] -%}
|
||||
{%- set next_nt.found = true -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
|
||||
{%- set continues_into_next = (
|
||||
role == 'model'
|
||||
and next_nt.role == 'assistant'
|
||||
and (not message.get('tool_calls') or ns_tr_out.flag)
|
||||
) -%}
|
||||
|
||||
{%- if ns.prev_message_type == 'tool_call' and not ns_tr_out.flag -%}
|
||||
{{- '<|tool_response>' -}}
|
||||
{%- elif continues_into_next -%}
|
||||
{{- '\n' -}}
|
||||
{%- elif not (ns_tr_out.flag and not has_content) -%}
|
||||
{{- '<turn|>\n' -}}
|
||||
{%- endif -%}
|
||||
|
||||
{#- Track previous non-tool role for next iteration (avoids O(n) backward scan) -#}
|
||||
{%- set ns.prev_non_tool_role = message['role'] -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
|
||||
{%- if add_generation_prompt -%}
|
||||
{%- if ns.prev_message_type != 'tool_response' and ns.prev_message_type != 'tool_call' -%}
|
||||
{{- '<|turn>model\n' -}}
|
||||
{%- if not enable_thinking | default(false) -%}
|
||||
{%- if not enable_thinking -%}
|
||||
{{- '<|channel>thought\n<channel|>' -}}
|
||||
{%- endif -%}
|
||||
{%- elif ns.prev_message_type == 'tool_response' and enable_thinking -%}
|
||||
{{- '<|channel>thought\n' -}}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
|
||||
@@ -2,5 +2,5 @@ lmcache >= 0.3.9
|
||||
# CuPy 14.1.0 imports pytest from cupy.testing._random. Use <14.1.0
|
||||
# until a fixed newer release is verified for runtime images.
|
||||
cupy-cuda13x < 14.1.0
|
||||
nixl >= 1.1.0 # Required for disaggregated prefill
|
||||
nixl == 1.2.0 # Required for disaggregated prefill
|
||||
mooncake-transfer-engine >= 0.3.8
|
||||
|
||||
@@ -11,7 +11,7 @@ numba == 0.65.0 # Required for N-gram speculative decoding
|
||||
datasets
|
||||
peft
|
||||
pytest-asyncio
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
packaging>=24.2
|
||||
setuptools>=77.0.3,<80.0.0
|
||||
setuptools-scm>=8
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# testing
|
||||
pytest
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
pytest-forked
|
||||
pytest-asyncio
|
||||
pytest-rerunfailures
|
||||
|
||||
@@ -966,7 +966,7 @@ tenacity==9.1.2
|
||||
# gpt-oss
|
||||
# lm-eval
|
||||
# plotly
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
# via -r requirements/test/cuda.in
|
||||
termcolor==3.1.0
|
||||
# via gpt-oss
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# testing
|
||||
pytest
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
pytest-forked
|
||||
pytest-asyncio
|
||||
pytest-rerunfailures
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
# testing
|
||||
pytest
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
pytest-forked
|
||||
pytest-asyncio
|
||||
pytest-rerunfailures
|
||||
|
||||
@@ -1182,7 +1182,7 @@ tenacity==9.1.4
|
||||
# via
|
||||
# gpt-oss
|
||||
# lm-eval
|
||||
tensorizer==2.10.1
|
||||
tensorizer==2.12.1
|
||||
# via
|
||||
# -c requirements/rocm.txt
|
||||
# -r requirements/test/rocm.in
|
||||
|
||||
@@ -13,7 +13,7 @@ pytest-shard
|
||||
absl-py
|
||||
accelerate
|
||||
arctic-inference
|
||||
lm_eval[api]
|
||||
lm_eval[api]>=0.4.12
|
||||
modelscope
|
||||
|
||||
# --- Audio Processing ---
|
||||
|
||||
@@ -33,7 +33,6 @@ arctic-inference==0.1.1
|
||||
attrs==26.1.0
|
||||
# via
|
||||
# aiohttp
|
||||
# jsonlines
|
||||
# jsonschema
|
||||
# referencing
|
||||
audioread==3.0.1
|
||||
@@ -225,8 +224,6 @@ joblib==1.5.3
|
||||
# librosa
|
||||
# nltk
|
||||
# scikit-learn
|
||||
jsonlines==4.0.0
|
||||
# via lm-eval
|
||||
jsonschema==4.26.0
|
||||
# via
|
||||
# -c requirements/common.txt
|
||||
@@ -247,7 +244,7 @@ librosa==0.10.2.post1
|
||||
# via -r requirements/test/xpu.in
|
||||
llvmlite==0.47.0
|
||||
# via numba
|
||||
lm-eval==0.4.11
|
||||
lm-eval==0.4.12
|
||||
# via -r requirements/test/xpu.in
|
||||
lxml==6.0.2
|
||||
# via
|
||||
@@ -734,5 +731,3 @@ xxhash==3.6.0
|
||||
# evaluate
|
||||
yarl==1.23.0
|
||||
# via aiohttp
|
||||
zstandard==0.25.0
|
||||
# via lm-eval
|
||||
|
||||
@@ -12,4 +12,4 @@ ray[data]
|
||||
setuptools==78.1.0
|
||||
setuptools-rust>=1.9.0
|
||||
nixl==0.3.0
|
||||
tpu-inference==0.21.0
|
||||
tpu-inference==0.22.1
|
||||
|
||||
@@ -16,5 +16,5 @@ torch==2.12.0
|
||||
torchaudio
|
||||
torchvision
|
||||
|
||||
auto_round_lib>=0.13.0
|
||||
auto_round_lib>=0.13.3
|
||||
vllm_xpu_kernels @ https://github.com/vllm-project/vllm-xpu-kernels/releases/download/v0.1.9.1/vllm_xpu_kernels-0.1.9.1-cp38-abi3-manylinux_2_28_x86_64.whl
|
||||
|
||||
Generated
+1
@@ -5880,6 +5880,7 @@ dependencies = [
|
||||
"expect-test",
|
||||
"futures",
|
||||
"http-body",
|
||||
"indexmap 2.13.0",
|
||||
"itertools 0.14.0",
|
||||
"libc",
|
||||
"llm-multimodal",
|
||||
|
||||
+1
-1
@@ -105,7 +105,7 @@ tonic-prost = "0.14.5"
|
||||
tonic-prost-build = "0.14.5"
|
||||
tool-parser = "1.2.0"
|
||||
tower = { version = "0.5.3", features = ["util"] }
|
||||
tower-http = { version = "0.6.8", features = ["trace"] }
|
||||
tower-http = { version = "0.6.8", features = ["cors", "trace"] }
|
||||
tracing = { version = "0.1.44", features = ["release_max_level_debug"] }
|
||||
tracing-futures = { version = "0.2.5", features = ["futures-03"] }
|
||||
tracing-subscriber = { version = "0.3.20", features = ["env-filter", "fmt"] }
|
||||
|
||||
@@ -145,8 +145,8 @@ impl ChatLlm {
|
||||
self.text.tokenizer_vocab_size()
|
||||
}
|
||||
|
||||
/// Model vocabulary size, else `None`.
|
||||
pub fn model_vocab_size(&self) -> Option<usize> {
|
||||
/// Model vocabulary size from the model config.
|
||||
pub fn model_vocab_size(&self) -> usize {
|
||||
self.text.model_vocab_size()
|
||||
}
|
||||
|
||||
|
||||
+65
-2
@@ -23,8 +23,8 @@ use vllm_engine_core_client::TransportMode;
|
||||
use vllm_managed_engine::ManagedEngineConfig;
|
||||
use vllm_managed_engine::cli::{ManagedEngineArgs, repartition_managed_engine_args};
|
||||
use vllm_server::{
|
||||
ApiServerOptions, ChatTemplateContentFormatOption, Config, CoordinatorMode, HttpListenerMode,
|
||||
ParserSelection, RendererSelection,
|
||||
ApiServerOptions, ChatTemplateContentFormatOption, Config, CoordinatorMode, CorsConfig,
|
||||
HttpListenerMode, ParserSelection, RendererSelection,
|
||||
};
|
||||
|
||||
use crate::cli::unsupported::UnsupportedArgs;
|
||||
@@ -84,6 +84,13 @@ pub enum Command {
|
||||
Serve(ServeArgs),
|
||||
}
|
||||
|
||||
/// A JSON-encoded list of strings, matching Python's `json.loads` CLI type for
|
||||
/// the CORS list arguments (e.g. `--allowed-origins '["*"]'`). Parsing the whole
|
||||
/// value as one item keeps clap from treating the field as a repeated flag.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct JsonStringList(pub Vec<String>);
|
||||
|
||||
/// Runtime arguments shared by the external-engine and managed-engine paths.
|
||||
#[serde_as]
|
||||
#[derive(Educe, Clone, Args, PartialEq, Eq, Deserialize)]
|
||||
@@ -127,6 +134,11 @@ pub struct SharedRuntimeArgs {
|
||||
/// `config.json`.
|
||||
#[arg(long)]
|
||||
pub max_model_len: Option<u32>,
|
||||
/// Maximum number of log probabilities to return when `logprobs` is
|
||||
/// specified in sampling parameters. `-1` means no cap.
|
||||
#[arg(long, value_parser = clap::value_parser!(i32).range(-1..), allow_negative_numbers = true)]
|
||||
#[serde(default)]
|
||||
pub max_logprobs: Option<i32>,
|
||||
/// TCP port for the gRPC Generate service. When not set, no gRPC server is
|
||||
/// started.
|
||||
#[arg(long)]
|
||||
@@ -215,6 +227,30 @@ pub struct SharedRuntimeArgs {
|
||||
#[serde(default)]
|
||||
pub served_model_name: Vec<String>,
|
||||
|
||||
/// CORS allowed origins as a JSON list. `["*"]` allows any origin.
|
||||
#[arg(long, value_parser = parse_json::<JsonStringList>, value_name = "JSON", default_value = r#"["*"]"#)]
|
||||
#[serde(default = "default_cors_wildcard")]
|
||||
pub allowed_origins: JsonStringList,
|
||||
|
||||
/// CORS allowed methods as a JSON list. `["*"]` allows the standard set.
|
||||
#[arg(long, value_parser = parse_json::<JsonStringList>, value_name = "JSON", default_value = r#"["*"]"#)]
|
||||
#[serde(default = "default_cors_wildcard")]
|
||||
pub allowed_methods: JsonStringList,
|
||||
|
||||
/// CORS allowed request headers as a JSON list. `["*"]` mirrors the request.
|
||||
#[arg(long, value_parser = parse_json::<JsonStringList>, value_name = "JSON", default_value = r#"["*"]"#)]
|
||||
#[serde(default = "default_cors_wildcard")]
|
||||
pub allowed_headers: JsonStringList,
|
||||
|
||||
/// Allow CORS credentials (cookies, authorization headers).
|
||||
#[arg(
|
||||
long,
|
||||
default_missing_value = "true",
|
||||
num_args = 0..=1
|
||||
)]
|
||||
#[serde(default)]
|
||||
pub allow_credentials: bool,
|
||||
|
||||
/// Unsupported Python vLLM frontend arguments recognized but not yet
|
||||
/// implemented in Rust.
|
||||
#[educe(Debug(ignore))]
|
||||
@@ -254,16 +290,19 @@ impl SharedRuntimeArgs {
|
||||
input_address: String,
|
||||
output_address: String,
|
||||
coordinator_address: Option<String>,
|
||||
engine_start_index: u32,
|
||||
engine_count: usize,
|
||||
) -> Config {
|
||||
let ready_timeout = self.ready_timeout();
|
||||
let shutdown_timeout = self.shutdown_timeout();
|
||||
let api_server_options = self.api_server_options();
|
||||
let cors = self.cors_config();
|
||||
|
||||
Config {
|
||||
transport_mode: TransportMode::Bootstrapped {
|
||||
input_address,
|
||||
output_address,
|
||||
engine_start_index,
|
||||
engine_count,
|
||||
ready_timeout,
|
||||
},
|
||||
@@ -281,7 +320,9 @@ impl SharedRuntimeArgs {
|
||||
chat_template: self.chat_template,
|
||||
default_chat_template_kwargs: self.default_chat_template_kwargs,
|
||||
chat_template_content_format: self.chat_template_content_format,
|
||||
max_logprobs: self.max_logprobs,
|
||||
api_server_options,
|
||||
cors,
|
||||
api_keys: self.api_key,
|
||||
disable_log_stats: self.disable_log_stats,
|
||||
grpc_port: self.grpc_port,
|
||||
@@ -303,6 +344,7 @@ impl SharedRuntimeArgs {
|
||||
let ready_timeout = self.ready_timeout();
|
||||
let shutdown_timeout = self.shutdown_timeout();
|
||||
let api_server_options = self.api_server_options();
|
||||
let cors = self.cors_config();
|
||||
|
||||
Config {
|
||||
transport_mode: TransportMode::HandshakeOwner {
|
||||
@@ -324,7 +366,9 @@ impl SharedRuntimeArgs {
|
||||
chat_template: self.chat_template,
|
||||
default_chat_template_kwargs: self.default_chat_template_kwargs,
|
||||
chat_template_content_format: self.chat_template_content_format,
|
||||
max_logprobs: self.max_logprobs,
|
||||
api_server_options,
|
||||
cors,
|
||||
api_keys: self.api_key,
|
||||
disable_log_stats: self.disable_log_stats,
|
||||
grpc_port: self.grpc_port,
|
||||
@@ -339,12 +383,25 @@ impl SharedRuntimeArgs {
|
||||
enable_request_id_headers: self.enable_request_id_headers,
|
||||
}
|
||||
}
|
||||
|
||||
fn cors_config(&self) -> CorsConfig {
|
||||
CorsConfig {
|
||||
allow_origins: self.allowed_origins.0.clone(),
|
||||
allow_methods: self.allowed_methods.0.clone(),
|
||||
allow_headers: self.allowed_headers.0.clone(),
|
||||
allow_credentials: self.allow_credentials,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn default_engine_ready_timeout_secs() -> u64 {
|
||||
600
|
||||
}
|
||||
|
||||
fn default_cors_wildcard() -> JsonStringList {
|
||||
JsonStringList(vec!["*".to_string()])
|
||||
}
|
||||
|
||||
fn parse_json<T: DeserializeOwned>(value: &str) -> Result<T, String> {
|
||||
serde_json::from_str(value).map_err(|e| format!("invalid JSON object: {}", e.as_report()))
|
||||
}
|
||||
@@ -380,6 +437,10 @@ pub struct FrontendArgs {
|
||||
/// `stats_update_address`.
|
||||
#[arg(long)]
|
||||
pub coordinator_address: Option<String>,
|
||||
/// First data-parallel engine rank expected to register with this
|
||||
/// bootstrapped frontend.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub engine_start_index: u32,
|
||||
/// Total number of data-parallel engines expected for this frontend.
|
||||
#[arg(long, default_value_t = 1)]
|
||||
pub engine_count: usize,
|
||||
@@ -397,6 +458,7 @@ impl FrontendArgs {
|
||||
self.input_address,
|
||||
self.output_address,
|
||||
self.coordinator_address,
|
||||
self.engine_start_index,
|
||||
self.engine_count,
|
||||
)
|
||||
}
|
||||
@@ -467,6 +529,7 @@ impl ServeArgs {
|
||||
self.managed_engine.clone().into_config(
|
||||
self.runtime.model.clone(),
|
||||
self.runtime.max_model_len,
|
||||
self.runtime.max_logprobs,
|
||||
self.runtime.language_model_only,
|
||||
self.runtime.disable_log_stats,
|
||||
self.runtime.shutdown_timeout,
|
||||
|
||||
+203
-14
@@ -38,6 +38,7 @@ fn serve_args_forward_python_flags_with_separator() {
|
||||
max_model_len: Some(
|
||||
512,
|
||||
),
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
chat_template: None,
|
||||
@@ -48,6 +49,22 @@ fn serve_args_forward_python_flags_with_separator() {
|
||||
enable_request_id_headers: false,
|
||||
disable_log_stats: false,
|
||||
served_model_name: [],
|
||||
allowed_origins: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allowed_methods: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allowed_headers: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allow_credentials: false,
|
||||
},
|
||||
managed_engine: ManagedEngineArgs {
|
||||
python: "../vllm/.venv/bin/python",
|
||||
@@ -134,6 +151,29 @@ fn serve_args_forward_disable_log_stats_to_managed_engine() {
|
||||
assert_eq!(config.python_args, vec!["--disable-log-stats"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_args_forward_max_logprobs_to_frontend_and_managed_engine() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--max-logprobs",
|
||||
"-1",
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Serve(args) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
assert_eq!(args.runtime.max_logprobs, Some(-1));
|
||||
|
||||
let frontend_config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
|
||||
assert_eq!(frontend_config.max_logprobs, Some(-1));
|
||||
|
||||
let engine_config = args.to_managed_engine_config(5555);
|
||||
assert_eq!(engine_config.python_args, vec!["--max-logprobs", "-1"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_args_auto_forward_python_multi_char_alias_without_separator() {
|
||||
let cli = Cli::try_parse_from(["vllm-rs", "serve", "Qwen/Qwen3-0.6B", "-tp", "2"]).unwrap();
|
||||
@@ -312,11 +352,17 @@ fn serve_args_reject_unknown_renderer_value() {
|
||||
|
||||
#[test]
|
||||
fn serve_args_reject_unsupported_flag_arg() {
|
||||
let error = Cli::try_parse_from(["vllm-rs", "serve", "Qwen/Qwen3-0.6B", "--allow-credentials"])
|
||||
.unwrap_err();
|
||||
let error = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--ssl-keyfile",
|
||||
"/tmp/key.pem",
|
||||
])
|
||||
.unwrap_err();
|
||||
|
||||
expect![[r#"
|
||||
error: invalid value 'true' for '--allow-credentials [<ALLOW_CREDENTIALS>]': argument is not implemented in Rust frontend yet
|
||||
error: invalid value '/tmp/key.pem' for '--ssl-keyfile <SSL_KEYFILE>': argument is not implemented in Rust frontend yet
|
||||
|
||||
Remove this unsupported argument to continue.
|
||||
|
||||
@@ -324,8 +370,7 @@ fn serve_args_reject_unsupported_flag_arg() {
|
||||
This may lead to unexpected behavior as the Rust frontend will completely ignore that argument.
|
||||
|
||||
For more information, try '--help'.
|
||||
"#]]
|
||||
.assert_eq(&error.to_string());
|
||||
"#]].assert_eq(&error.to_string());
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -379,6 +424,7 @@ fn frontend_args_accept_json() {
|
||||
coordinator_address: Some(
|
||||
"tcp://127.0.0.1:7000",
|
||||
),
|
||||
engine_start_index: 0,
|
||||
engine_count: 1,
|
||||
runtime: SharedRuntimeArgs {
|
||||
model: "Qwen/Qwen3-0.6B",
|
||||
@@ -388,6 +434,7 @@ fn frontend_args_accept_json() {
|
||||
renderer: Auto,
|
||||
language_model_only: false,
|
||||
max_model_len: None,
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
chat_template: None,
|
||||
@@ -398,6 +445,22 @@ fn frontend_args_accept_json() {
|
||||
enable_request_id_headers: false,
|
||||
disable_log_stats: false,
|
||||
served_model_name: [],
|
||||
allowed_origins: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allowed_methods: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allowed_headers: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allow_credentials: false,
|
||||
},
|
||||
},
|
||||
),
|
||||
@@ -431,6 +494,7 @@ fn frontend_args_json_applies_defaults() {
|
||||
assert_eq!(args.runtime.reasoning_parser, ParserSelection::Auto);
|
||||
assert_eq!(args.runtime.renderer, RendererSelection::Auto);
|
||||
assert_eq!(args.runtime.max_model_len, None);
|
||||
assert_eq!(args.runtime.max_logprobs, None);
|
||||
assert_eq!(args.runtime.shutdown_timeout, 0);
|
||||
}
|
||||
|
||||
@@ -446,7 +510,7 @@ fn frontend_args_json_accepts_supported_non_default_fields() {
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","engine_ready_timeout_secs":42,"tool_call_parser":"hermes","reasoning_parser":"qwen3_thinking","tokenizer_mode":"deepseek_v32","language_model_only":true,"max_model_len":8192,"shutdown_timeout":3}"#,
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","engine_ready_timeout_secs":42,"tool_call_parser":"hermes","reasoning_parser":"qwen3_thinking","tokenizer_mode":"deepseek_v32","language_model_only":true,"max_model_len":8192,"max_logprobs":-1,"shutdown_timeout":3}"#,
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
@@ -465,6 +529,7 @@ fn frontend_args_json_accepts_supported_non_default_fields() {
|
||||
assert_eq!(args.runtime.renderer, RendererSelection::DeepSeekV32);
|
||||
assert!(args.runtime.language_model_only);
|
||||
assert_eq!(args.runtime.max_model_len, Some(8192));
|
||||
assert_eq!(args.runtime.max_logprobs, Some(-1));
|
||||
assert_eq!(args.runtime.shutdown_timeout, 3);
|
||||
}
|
||||
|
||||
@@ -531,6 +596,71 @@ fn frontend_args_json_sets_prompt_tokens_details_flag() {
|
||||
assert!(args.runtime.enable_prompt_tokens_details);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_args_parse_cors_flags() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--allowed-origins",
|
||||
r#"["http://a.com","http://b.com"]"#,
|
||||
"--allowed-methods",
|
||||
r#"["GET","POST"]"#,
|
||||
"--allow-credentials",
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Serve(serve) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
assert_eq!(
|
||||
serve.runtime.allowed_origins.0,
|
||||
["http://a.com", "http://b.com"]
|
||||
);
|
||||
assert_eq!(serve.runtime.allowed_methods.0, ["GET", "POST"]);
|
||||
assert!(serve.runtime.allow_credentials);
|
||||
// Unspecified lists keep the permissive default.
|
||||
assert_eq!(serve.runtime.allowed_headers.0, ["*"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_args_cors_defaults_are_permissive() {
|
||||
let cli = Cli::try_parse_from(["vllm-rs", "serve", "Qwen/Qwen3-0.6B"]).unwrap();
|
||||
|
||||
let Command::Serve(serve) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
assert_eq!(serve.runtime.allowed_origins.0, ["*"]);
|
||||
assert_eq!(serve.runtime.allowed_methods.0, ["*"]);
|
||||
assert_eq!(serve.runtime.allowed_headers.0, ["*"]);
|
||||
assert!(!serve.runtime.allow_credentials);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn frontend_args_json_parses_cors_fields() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"frontend",
|
||||
"--listen-fd",
|
||||
"3",
|
||||
"--input-address",
|
||||
"ipc:///tmp/input.sock",
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","allowed_origins":["http://a.com"],"allow_credentials":true}"#,
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Frontend(args) = cli.command else {
|
||||
panic!("expected frontend args");
|
||||
};
|
||||
assert_eq!(args.runtime.allowed_origins.0, ["http://a.com"]);
|
||||
assert!(args.runtime.allow_credentials);
|
||||
// Unspecified lists fall back to the permissive default via serde.
|
||||
assert_eq!(args.runtime.allowed_methods.0, ["*"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn frontend_args_json_rejects_unsupported_fields() {
|
||||
let error = Cli::try_parse_from([
|
||||
@@ -543,14 +673,14 @@ fn frontend_args_json_rejects_unsupported_fields() {
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","allow_credentials":true}"#,
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","ssl_keyfile":"/tmp/key.pem"}"#,
|
||||
])
|
||||
.unwrap_err();
|
||||
|
||||
expect![[r#"
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","allow_credentials":true}' for '--args-json <JSON>':
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","ssl_keyfile":"/tmp/key.pem"}' for '--args-json <JSON>':
|
||||
The following arguments are not implemented in Rust frontend yet:
|
||||
- allow_credentials
|
||||
- ssl_keyfile
|
||||
|
||||
Remove these arguments to continue.
|
||||
|
||||
@@ -570,15 +700,15 @@ fn frontend_args_json_aggregates_multiple_unsupported_fields() {
|
||||
"--output-address",
|
||||
"ipc:///tmp/output.sock",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","allow_credentials":true,"ssl_keyfile":"/tmp/key.pem"}"#,
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","ssl_keyfile":"/tmp/key.pem"}"#,
|
||||
])
|
||||
.unwrap_err();
|
||||
|
||||
let actual = error.to_string().replace(": \n", ":\n");
|
||||
expect![[r#"
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","allow_credentials":true,"ssl_keyfile":"/tmp/key.pem"}' for '--args-json <JSON>':
|
||||
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","ssl_keyfile":"/tmp/key.pem"}' for '--args-json <JSON>':
|
||||
The following arguments are not implemented in Rust frontend yet:
|
||||
- allow_credentials
|
||||
- response_role
|
||||
- ssl_keyfile
|
||||
|
||||
Remove these arguments to continue.
|
||||
@@ -792,6 +922,7 @@ fn serve_args_accept_handshake_aliases() {
|
||||
renderer: Auto,
|
||||
language_model_only: false,
|
||||
max_model_len: None,
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
chat_template: None,
|
||||
@@ -802,6 +933,22 @@ fn serve_args_accept_handshake_aliases() {
|
||||
enable_request_id_headers: false,
|
||||
disable_log_stats: false,
|
||||
served_model_name: [],
|
||||
allowed_origins: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allowed_methods: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allowed_headers: JsonStringList(
|
||||
[
|
||||
"*",
|
||||
],
|
||||
),
|
||||
allow_credentials: false,
|
||||
},
|
||||
managed_engine: ManagedEngineArgs {
|
||||
python: "python3",
|
||||
@@ -917,11 +1064,24 @@ fn serve_frontend_config_uses_dp_address_as_advertised_host() {
|
||||
chat_template: None,
|
||||
default_chat_template_kwargs: None,
|
||||
chat_template_content_format: Auto,
|
||||
max_logprobs: None,
|
||||
api_server_options: ApiServerOptions {
|
||||
enable_log_requests: false,
|
||||
enable_prompt_tokens_details: false,
|
||||
enable_request_id_headers: false,
|
||||
},
|
||||
cors: CorsConfig {
|
||||
allow_origins: [
|
||||
"*",
|
||||
],
|
||||
allow_methods: [
|
||||
"*",
|
||||
],
|
||||
allow_headers: [
|
||||
"*",
|
||||
],
|
||||
allow_credentials: false,
|
||||
},
|
||||
api_keys: [],
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
@@ -985,11 +1145,24 @@ fn serve_frontend_config_keeps_tcp_transport_for_non_local_only_topology() {
|
||||
chat_template: None,
|
||||
default_chat_template_kwargs: None,
|
||||
chat_template_content_format: Auto,
|
||||
max_logprobs: None,
|
||||
api_server_options: ApiServerOptions {
|
||||
enable_log_requests: false,
|
||||
enable_prompt_tokens_details: false,
|
||||
enable_request_id_headers: false,
|
||||
},
|
||||
cors: CorsConfig {
|
||||
allow_origins: [
|
||||
"*",
|
||||
],
|
||||
allow_methods: [
|
||||
"*",
|
||||
],
|
||||
allow_headers: [
|
||||
"*",
|
||||
],
|
||||
allow_credentials: false,
|
||||
},
|
||||
api_keys: [],
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
@@ -1033,8 +1206,10 @@ fn frontend_config_uses_external_coordinator_when_coordinator_address_is_present
|
||||
"ipc:///tmp/output.sock",
|
||||
"--coordinator-address",
|
||||
"tcp://127.0.0.1:7000",
|
||||
"--engine-start-index",
|
||||
"3",
|
||||
"--engine-count",
|
||||
"2",
|
||||
"1",
|
||||
"--args-json",
|
||||
r#"{"model_tag":"Qwen/Qwen3-0.6B"}"#,
|
||||
])
|
||||
@@ -1050,7 +1225,8 @@ fn frontend_config_uses_external_coordinator_when_coordinator_address_is_present
|
||||
transport_mode: Bootstrapped {
|
||||
input_address: "ipc:///tmp/input.sock",
|
||||
output_address: "ipc:///tmp/output.sock",
|
||||
engine_count: 2,
|
||||
engine_start_index: 3,
|
||||
engine_count: 1,
|
||||
ready_timeout: 600s,
|
||||
},
|
||||
coordinator_mode: External {
|
||||
@@ -1068,11 +1244,24 @@ fn frontend_config_uses_external_coordinator_when_coordinator_address_is_present
|
||||
chat_template: None,
|
||||
default_chat_template_kwargs: None,
|
||||
chat_template_content_format: Auto,
|
||||
max_logprobs: None,
|
||||
api_server_options: ApiServerOptions {
|
||||
enable_log_requests: false,
|
||||
enable_prompt_tokens_details: false,
|
||||
enable_request_id_headers: false,
|
||||
},
|
||||
cors: CorsConfig {
|
||||
allow_origins: [
|
||||
"*",
|
||||
],
|
||||
allow_methods: [
|
||||
"*",
|
||||
],
|
||||
allow_headers: [
|
||||
"*",
|
||||
],
|
||||
allow_credentials: false,
|
||||
},
|
||||
api_keys: [],
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
|
||||
@@ -202,14 +202,6 @@ pub struct EngineUnsupportedArgs {
|
||||
#[arg(long)]
|
||||
pub tokenizer_revision: Option<Unsupported>,
|
||||
|
||||
/// Maximum number of log probabilities to return when `logprobs` is
|
||||
/// specified in `SamplingParams`. The default value comes the default for
|
||||
/// the OpenAI Chat Completions API. -1 means no cap, i.e. all
|
||||
/// (output_length * vocab_size) logprobs are allowed to be returned and
|
||||
/// it may cause OOM.
|
||||
#[arg(long)]
|
||||
pub max_logprobs: Option<Unsupported>,
|
||||
|
||||
/// Skip initialization of tokenizer and detokenizer. Expects valid
|
||||
/// `prompt_token_ids` and `None` for prompt from the input. The generated
|
||||
/// output will contain token ids.
|
||||
@@ -534,27 +526,6 @@ pub struct ServerUnsupportedArgs {
|
||||
#[arg(long)]
|
||||
pub disable_access_log_for_endpoints: Option<Noop>,
|
||||
|
||||
/// Allow credentials.
|
||||
#[arg(
|
||||
long,
|
||||
visible_alias = "no-allow-credentials",
|
||||
default_missing_value = "true",
|
||||
num_args = 0..=1
|
||||
)]
|
||||
pub allow_credentials: Option<Unsupported>,
|
||||
|
||||
/// Allowed origins.
|
||||
#[arg(long)]
|
||||
pub allowed_origins: Option<Unsupported>,
|
||||
|
||||
/// Allowed methods.
|
||||
#[arg(long)]
|
||||
pub allowed_methods: Option<Unsupported>,
|
||||
|
||||
/// Allowed headers.
|
||||
#[arg(long)]
|
||||
pub allowed_headers: Option<Unsupported>,
|
||||
|
||||
/// The file path to the SSL key file.
|
||||
#[arg(long)]
|
||||
pub ssl_keyfile: Option<Unsupported>,
|
||||
|
||||
@@ -56,6 +56,9 @@ pub enum TransportMode {
|
||||
/// Output PULL socket address that engines will connect to for
|
||||
/// responses.
|
||||
output_address: String,
|
||||
/// First data-parallel engine rank expected to register on this
|
||||
/// transport.
|
||||
engine_start_index: u32,
|
||||
/// Total number of engines expected to register on this transport.
|
||||
engine_count: usize,
|
||||
/// Maximum time to wait for all expected engines to register.
|
||||
@@ -246,6 +249,7 @@ impl EngineCoreClient {
|
||||
TransportMode::Bootstrapped {
|
||||
input_address,
|
||||
output_address,
|
||||
engine_start_index,
|
||||
engine_count,
|
||||
ready_timeout,
|
||||
} => {
|
||||
@@ -256,6 +260,7 @@ impl EngineCoreClient {
|
||||
transport::connect_bootstrapped(
|
||||
input_address,
|
||||
output_address,
|
||||
*engine_start_index,
|
||||
*engine_count,
|
||||
*ready_timeout,
|
||||
)
|
||||
@@ -409,6 +414,24 @@ impl EngineCoreClient {
|
||||
.expect("engine core client requires at least one engine")
|
||||
}
|
||||
|
||||
/// Return the world size (TP * PP) from the parallel config, if available.
|
||||
pub fn world_size(&self) -> u64 {
|
||||
self.engines
|
||||
.first()
|
||||
.expect("engine core client requires at least one engine")
|
||||
.ready_response
|
||||
.world_size
|
||||
}
|
||||
|
||||
/// Return the data parallel size from the parallel config, if available.
|
||||
pub fn data_parallel_size(&self) -> u64 {
|
||||
self.engines
|
||||
.first()
|
||||
.expect("engine core client requires at least one engine")
|
||||
.ready_response
|
||||
.data_parallel_size
|
||||
}
|
||||
|
||||
/// Get the model name associated with this client used for metrics
|
||||
/// labeling.
|
||||
pub fn model_name(&self) -> &str {
|
||||
|
||||
@@ -52,6 +52,8 @@ pub fn default_ready_response() -> EngineCoreReadyResponse {
|
||||
dp_stats_address: None,
|
||||
dtype: ModelDtype::Float32,
|
||||
vllm_version: "test-vllm-version".to_string(),
|
||||
world_size: 1,
|
||||
data_parallel_size: 1,
|
||||
kv_cache_size_tokens: None,
|
||||
kv_cache_max_concurrency: None,
|
||||
}
|
||||
|
||||
@@ -44,6 +44,10 @@ pub struct EngineCoreReadyResponse {
|
||||
pub dtype: ModelDtype,
|
||||
/// Python vLLM version reported by the engine process.
|
||||
pub vllm_version: String,
|
||||
/// World size (TP * PP) from the parallel config.
|
||||
pub world_size: u64,
|
||||
/// Data parallelism size from the parallel config.
|
||||
pub data_parallel_size: u64,
|
||||
/// Total KV cache capacity in tokens, if reported.
|
||||
pub kv_cache_size_tokens: Option<u64>,
|
||||
/// Maximum achievable request concurrency given the KV cache, if reported.
|
||||
|
||||
@@ -44,6 +44,14 @@ fn default_repetition_penalty() -> f32 {
|
||||
1.0
|
||||
}
|
||||
|
||||
fn default_temperature() -> f32 {
|
||||
1.0
|
||||
}
|
||||
|
||||
fn default_max_tokens() -> u32 {
|
||||
16
|
||||
}
|
||||
|
||||
mod classified_outputs;
|
||||
pub mod dtype;
|
||||
pub mod handshake;
|
||||
@@ -246,24 +254,28 @@ pub struct StructuredOutputsParams {
|
||||
///
|
||||
/// Original Python definition:
|
||||
/// <https://github.com/vllm-project/vllm/blob/f22d6e026798a74e6542a52ef776c054f2de572a/vllm/sampling_params.py#L155-L291>
|
||||
// Python's SamplingParams is `omit_defaults=True`, so msgpack drops
|
||||
// default-valued keys; default the whole struct. Per-field fns cover the
|
||||
// non-zero defaults.
|
||||
#[serde_with::skip_serializing_none]
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, DefaultFromSerde)]
|
||||
#[serde(default)]
|
||||
pub struct EngineCoreSamplingParams {
|
||||
/// Controls randomness. Lower values are more deterministic; zero means
|
||||
/// greedy sampling.
|
||||
#[serde(default = "default_temperature")]
|
||||
pub temperature: f32,
|
||||
/// Cumulative probability threshold for nucleus sampling.
|
||||
#[serde(default = "default_top_p")]
|
||||
pub top_p: f32,
|
||||
/// Maximum number of top tokens to consider. `0` means all tokens.
|
||||
#[serde(default)]
|
||||
pub top_k: u32,
|
||||
/// Random seed used by the sampler when present.
|
||||
pub seed: Option<i64>,
|
||||
/// Maximum number of tokens to generate per output sequence.
|
||||
#[serde(default = "default_max_tokens")]
|
||||
pub max_tokens: u32,
|
||||
/// Minimum number of tokens to generate before EOS or stop-token handling.
|
||||
#[serde(default)]
|
||||
pub min_tokens: u32,
|
||||
/// Number of log probabilities to return per generated token.
|
||||
///
|
||||
@@ -274,7 +286,6 @@ pub struct EngineCoreSamplingParams {
|
||||
/// `None` disables prompt logprobs. `-1` requests the full vocabulary.
|
||||
pub prompt_logprobs: Option<i32>,
|
||||
/// Minimum probability threshold for token sampling.
|
||||
#[serde(default)]
|
||||
pub min_p: f32,
|
||||
/// Frequency penalty applied by the sampler.
|
||||
pub frequency_penalty: f32,
|
||||
@@ -301,16 +312,13 @@ pub struct EngineCoreSamplingParams {
|
||||
pub all_stop_token_ids: BTreeSet<u32>,
|
||||
/// Logit biases to apply during sampling.
|
||||
/// Keys are token IDs
|
||||
#[serde(default)]
|
||||
pub logit_bias: Option<HashMap<u32, f32>>,
|
||||
/// Restrict output to these token IDs only.
|
||||
#[serde(default)]
|
||||
pub allowed_token_ids: Option<Vec<u32>>,
|
||||
/// Tokenized bad words to avoid during generation.
|
||||
#[serde(default, rename = "_bad_words_token_ids")]
|
||||
#[serde(rename = "_bad_words_token_ids")]
|
||||
pub bad_words_token_ids: Option<Vec<Vec<u32>>>,
|
||||
/// Parameters for configuring structured outputs (guided decoding).
|
||||
#[serde(default)]
|
||||
pub structured_outputs: Option<StructuredOutputsParams>,
|
||||
/// Specific token IDs for which log probabilities should be returned at
|
||||
/// each position.
|
||||
@@ -318,15 +326,12 @@ pub struct EngineCoreSamplingParams {
|
||||
/// When set, the engine returns logprobs for exactly these tokens in
|
||||
/// addition to the sampled/scored token. Mutually exclusive with the
|
||||
/// `logprobs` count field in practice.
|
||||
#[serde(default)]
|
||||
pub logprob_token_ids: Option<Vec<u32>>,
|
||||
/// If `Some(true)`, the request will not attempt to read from the prefix
|
||||
/// cache; newly computed blocks may still populate the cache. `None`
|
||||
/// defers to engine-core defaults.
|
||||
#[serde(default)]
|
||||
pub skip_reading_prefix_cache: Option<bool>,
|
||||
/// Additional request parameters for custom extensions (from `vllm_xargs`).
|
||||
#[serde(default)]
|
||||
pub extra_args: Option<HashMap<String, serde_json::Value>>,
|
||||
}
|
||||
|
||||
@@ -640,4 +645,58 @@ mod tests {
|
||||
let value = serde_json::to_value(params).unwrap();
|
||||
assert_eq!(value["_backend"], "guidance");
|
||||
}
|
||||
|
||||
/// A real `sampling_params` is a sparse `omit_defaults` map; absent fields
|
||||
/// must fall back to defaults. `python_compat` can't catch this since Rust
|
||||
/// encodes full maps (see `engine_core_request_serializes_as_full_array`).
|
||||
#[test]
|
||||
fn decodes_sampling_params_with_omitted_defaults() {
|
||||
let sampling_params = Value::Map(vec![
|
||||
(
|
||||
Value::from("stop_token_ids"),
|
||||
Value::Array(vec![Value::from(151643u32)]),
|
||||
),
|
||||
(Value::from("skip_reading_prefix_cache"), Value::from(false)),
|
||||
]);
|
||||
let request = Value::Array(vec![
|
||||
Value::from("req-omit-defaults"),
|
||||
Value::Array(vec![
|
||||
Value::from(1u32),
|
||||
Value::from(2u32),
|
||||
Value::from(3u32),
|
||||
]),
|
||||
Value::Nil,
|
||||
sampling_params,
|
||||
Value::Nil,
|
||||
Value::from(1.0f64),
|
||||
]);
|
||||
|
||||
let mut bytes = Vec::new();
|
||||
rmpv::encode::write_value(&mut bytes, &request).unwrap();
|
||||
|
||||
let decoded: EngineCoreRequest = decode_msgpack(&bytes)
|
||||
.expect("a real omit_defaults request must decode (regression: missing field)");
|
||||
|
||||
assert_eq!(decoded.request_id, "req-omit-defaults");
|
||||
let sampling = decoded.sampling_params.expect("sampling params present");
|
||||
|
||||
assert_eq!(sampling.stop_token_ids, vec![151643]);
|
||||
assert_eq!(sampling.skip_reading_prefix_cache, Some(false));
|
||||
|
||||
// Omitted fields -> Python defaults.
|
||||
assert_eq!(sampling.temperature, 1.0);
|
||||
assert_eq!(sampling.top_p, 1.0);
|
||||
assert_eq!(sampling.top_k, 0);
|
||||
assert_eq!(sampling.seed, None);
|
||||
assert_eq!(sampling.max_tokens, 16);
|
||||
assert_eq!(sampling.min_tokens, 0);
|
||||
assert_eq!(sampling.min_p, 0.0);
|
||||
assert_eq!(sampling.frequency_penalty, 0.0);
|
||||
assert_eq!(sampling.presence_penalty, 0.0);
|
||||
assert_eq!(sampling.repetition_penalty, 1.0);
|
||||
assert_eq!(sampling.logprobs, None);
|
||||
assert_eq!(sampling.prompt_logprobs, None);
|
||||
assert_eq!(sampling.eos_token_id, None);
|
||||
assert!(sampling.all_stop_token_ids.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@ use crate::mock_engine::{
|
||||
MockEngineConfig, MockEngineDataSockets, connect_to_bootstrapped_frontend, connect_to_frontend,
|
||||
default_ready_response,
|
||||
};
|
||||
use crate::protocol::handshake::HandshakeInitMessage;
|
||||
use crate::protocol::handshake::{EngineCoreReadyResponse, HandshakeInitMessage};
|
||||
|
||||
/// Per-test IPC endpoint namespace backed by a unique temporary directory.
|
||||
///
|
||||
@@ -62,6 +62,15 @@ fn test_mock_engine_config() -> MockEngineConfig {
|
||||
}
|
||||
}
|
||||
|
||||
fn test_mock_engine_config_with_ready(ready_response: EngineCoreReadyResponse) -> MockEngineConfig {
|
||||
MockEngineConfig {
|
||||
local: true,
|
||||
headless: true,
|
||||
ready_response,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// Complete the engine-core handshake and connect mock input/output sockets
|
||||
/// plus optional coordinator sockets.
|
||||
pub async fn setup_mock_engine_sockets(
|
||||
@@ -147,3 +156,49 @@ where
|
||||
});
|
||||
(shutdown_tx, engine_task)
|
||||
}
|
||||
|
||||
/// Like [`setup_mock_engine`] but uses a custom ready response for the
|
||||
/// handshake, allowing tests to control `world_size`, `data_parallel_size`,
|
||||
/// etc.
|
||||
async fn setup_mock_engine_with_ready(
|
||||
engine_handshake: String,
|
||||
engine_id: impl Into<EngineId>,
|
||||
ready_response: EngineCoreReadyResponse,
|
||||
) -> (DealerSocket, PushSocket) {
|
||||
let config = test_mock_engine_config_with_ready(ready_response);
|
||||
let MockEngineSockets { data_sockets, .. } =
|
||||
connect_to_frontend(engine_handshake, engine_id, config)
|
||||
.await
|
||||
.expect("connect mock engine with custom ready response");
|
||||
let MockEngineDataSockets { dealer, push } =
|
||||
data_sockets.into_iter().next().expect("mock engine data socket");
|
||||
(dealer, push)
|
||||
}
|
||||
|
||||
/// Like [`spawn_mock_engine_task`] but uses a custom ready response for the
|
||||
/// handshake, allowing tests to set `world_size` and `data_parallel_size` to
|
||||
/// non-default values.
|
||||
pub fn spawn_mock_engine_task_with_ready<F>(
|
||||
engine_handshake: String,
|
||||
engine_id: impl Into<EngineId>,
|
||||
ready_response: EngineCoreReadyResponse,
|
||||
run: F,
|
||||
) -> (oneshot::Sender<()>, tokio::task::JoinHandle<()>)
|
||||
where
|
||||
F: for<'a> FnOnce(
|
||||
&'a mut DealerSocket,
|
||||
&'a mut PushSocket,
|
||||
) -> Pin<Box<dyn Future<Output = ()> + Send + 'a>>
|
||||
+ Send
|
||||
+ 'static,
|
||||
{
|
||||
let (shutdown_tx, shutdown_rx) = oneshot::channel();
|
||||
let engine_id = engine_id.into();
|
||||
let engine_task = tokio::spawn(async move {
|
||||
let (mut dealer, mut push) =
|
||||
setup_mock_engine_with_ready(engine_handshake, engine_id, ready_response).await;
|
||||
run(&mut dealer, &mut push).await;
|
||||
let _ = shutdown_rx.await;
|
||||
});
|
||||
(shutdown_tx, engine_task)
|
||||
}
|
||||
|
||||
@@ -303,6 +303,7 @@ fn bootstrapped_test_config(
|
||||
transport_mode: TransportMode::Bootstrapped {
|
||||
input_address,
|
||||
output_address,
|
||||
engine_start_index: 0,
|
||||
engine_count,
|
||||
ready_timeout,
|
||||
},
|
||||
@@ -312,6 +313,34 @@ fn bootstrapped_test_config(
|
||||
}
|
||||
}
|
||||
|
||||
fn bootstrapped_test_config_with_start_index(
|
||||
input_address: String,
|
||||
output_address: String,
|
||||
engine_start_index: u32,
|
||||
engine_count: usize,
|
||||
ready_timeout: Duration,
|
||||
client_index: u32,
|
||||
coordinator_mode: Option<CoordinatorMode>,
|
||||
) -> EngineCoreClientConfig {
|
||||
let mut config = bootstrapped_test_config(
|
||||
input_address,
|
||||
output_address,
|
||||
engine_count,
|
||||
ready_timeout,
|
||||
client_index,
|
||||
coordinator_mode,
|
||||
);
|
||||
let TransportMode::Bootstrapped {
|
||||
engine_start_index: start,
|
||||
..
|
||||
} = &mut config.transport_mode
|
||||
else {
|
||||
unreachable!("bootstrapped_test_config returns bootstrapped transport")
|
||||
};
|
||||
*start = engine_start_index;
|
||||
config
|
||||
}
|
||||
|
||||
async fn recv_xpub_message(xpub: &mut XPubSocket) -> Vec<bytes::Bytes> {
|
||||
xpub.recv().await.unwrap().into_vec()
|
||||
}
|
||||
@@ -2438,6 +2467,7 @@ fn python_msgpack_fixtures_match_rust_encoding() {
|
||||
let stdout = String::from_utf8(output.stdout).unwrap();
|
||||
let mut lines = stdout.lines();
|
||||
let request_hex = lines.next().expect("missing request fixture line");
|
||||
let defaults_request_hex = lines.next().expect("missing defaults request fixture line");
|
||||
let multimodal_request_hex = lines.next().expect("missing multimodal request fixture line");
|
||||
let outputs_hex = lines.next().expect("missing outputs fixture line");
|
||||
let inline_logprobs_frames = lines.next().expect("missing inline logprobs fixture line");
|
||||
@@ -2455,6 +2485,42 @@ fn python_msgpack_fixtures_match_rust_encoding() {
|
||||
let expected_request = sample_request();
|
||||
assert_eq!(decoded_request, expected_request);
|
||||
|
||||
// All-default sampling params -> empty map; must decode to Python defaults.
|
||||
let defaults_request_bytes = hex::decode(defaults_request_hex).unwrap();
|
||||
let decoded_defaults: EngineCoreRequest =
|
||||
rmp_serde::from_slice(&defaults_request_bytes).unwrap();
|
||||
assert_eq!(decoded_defaults.request_id, "req-defaults");
|
||||
let sampling = decoded_defaults
|
||||
.sampling_params
|
||||
.expect("defaults request carries sampling params");
|
||||
assert_eq!(
|
||||
sampling,
|
||||
EngineCoreSamplingParams {
|
||||
temperature: 1.0,
|
||||
top_p: 1.0,
|
||||
top_k: 0,
|
||||
seed: None,
|
||||
max_tokens: 16,
|
||||
min_tokens: 0,
|
||||
logprobs: None,
|
||||
prompt_logprobs: None,
|
||||
min_p: 0.0,
|
||||
frequency_penalty: 0.0,
|
||||
presence_penalty: 0.0,
|
||||
repetition_penalty: 1.0,
|
||||
stop_token_ids: Vec::new(),
|
||||
eos_token_id: None,
|
||||
all_stop_token_ids: BTreeSet::new(),
|
||||
logit_bias: None,
|
||||
allowed_token_ids: None,
|
||||
bad_words_token_ids: None,
|
||||
structured_outputs: None,
|
||||
logprob_token_ids: None,
|
||||
skip_reading_prefix_cache: None,
|
||||
extra_args: None,
|
||||
},
|
||||
);
|
||||
|
||||
let decoded_multimodal_request: EngineCoreRequest =
|
||||
rmp_serde::from_slice(&multimodal_request_bytes).unwrap();
|
||||
assert_eq!(decoded_multimodal_request, sample_multimodal_request());
|
||||
@@ -2652,6 +2718,90 @@ async fn bootstrapped_connects_with_contiguous_engine_ids() {
|
||||
client.shutdown().await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn bootstrapped_connects_with_nonzero_engine_start_index() {
|
||||
init_tracing();
|
||||
let ipc = IpcNamespace::new().unwrap();
|
||||
let input_address = ipc.input_endpoint();
|
||||
let output_address = ipc.output_endpoint();
|
||||
|
||||
let client_task = tokio::spawn({
|
||||
let input_address = input_address.clone();
|
||||
let output_address = output_address.clone();
|
||||
async move {
|
||||
EngineCoreClient::connect(bootstrapped_test_config_with_start_index(
|
||||
input_address,
|
||||
output_address,
|
||||
3,
|
||||
1,
|
||||
Duration::from_secs(2),
|
||||
0,
|
||||
None,
|
||||
))
|
||||
.await
|
||||
.unwrap()
|
||||
}
|
||||
});
|
||||
|
||||
let (_dealer, _push) =
|
||||
setup_bootstrapped_mock_engine(input_address, output_address, &[0x03, 0x00]).await;
|
||||
let client = client_task.await.unwrap();
|
||||
|
||||
assert_eq!(client.engine_count(), 1);
|
||||
let engine_ids =
|
||||
client.engine_identities().into_iter().map(|id| id.to_vec()).collect::<Vec<_>>();
|
||||
assert_eq!(engine_ids, vec![vec![0x03, 0x00]]);
|
||||
|
||||
client.shutdown().await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn bootstrapped_rejects_unexpected_engine_id_for_start_index() {
|
||||
init_tracing();
|
||||
let ipc = IpcNamespace::new().unwrap();
|
||||
let input_address = ipc.input_endpoint();
|
||||
let output_address = ipc.output_endpoint();
|
||||
|
||||
let client_task = tokio::spawn({
|
||||
let input_address = input_address.clone();
|
||||
let output_address = output_address.clone();
|
||||
async move {
|
||||
EngineCoreClient::connect(bootstrapped_test_config_with_start_index(
|
||||
input_address,
|
||||
output_address,
|
||||
3,
|
||||
1,
|
||||
Duration::from_secs(2),
|
||||
0,
|
||||
None,
|
||||
))
|
||||
.await
|
||||
}
|
||||
});
|
||||
|
||||
let _ = crate::mock_engine::connect_to_bootstrapped_frontend(
|
||||
input_address,
|
||||
output_address,
|
||||
&[0x00, 0x00],
|
||||
crate::mock_engine::MockEngineConfig {
|
||||
local: true,
|
||||
headless: true,
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await;
|
||||
let error = match client_task.await.unwrap() {
|
||||
Ok(_) => panic!("bootstrapped connect should reject unexpected engine id"),
|
||||
Err(error) => error,
|
||||
};
|
||||
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("received input registration for unexpected engine id")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn bootstrapped_connect_times_out_without_registration() {
|
||||
init_tracing();
|
||||
|
||||
@@ -31,12 +31,13 @@ class FinishReason(IntEnum):
|
||||
REPETITION = 4
|
||||
|
||||
|
||||
class EngineCoreSamplingParams(msgspec.Struct, dict=True):
|
||||
# Mirror of real SamplingParams; omit_defaults makes fixtures match real maps.
|
||||
class EngineCoreSamplingParams(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
temperature: float = 1.0
|
||||
top_p: float = 1.0
|
||||
top_k: int = 0
|
||||
seed: int | None = None
|
||||
max_tokens: int = 65536
|
||||
max_tokens: int = 16
|
||||
min_tokens: int = 0
|
||||
min_p: float = 0.0
|
||||
frequency_penalty: float = 0.0
|
||||
@@ -135,6 +136,16 @@ request = EngineCoreRequest(
|
||||
client_index=0,
|
||||
)
|
||||
|
||||
# All defaults -> empty map. Regression guard for the sparse-map decode.
|
||||
defaults_request = EngineCoreRequest(
|
||||
request_id="req-defaults",
|
||||
prompt_token_ids=[5, 6, 7],
|
||||
mm_features=None,
|
||||
sampling_params=EngineCoreSamplingParams(),
|
||||
pooling_params=None,
|
||||
arrival_time=1.0,
|
||||
)
|
||||
|
||||
multimodal_tensor = np.array([[1.0, 2.0], [3.5, 4.25]], dtype=np.float32)
|
||||
multimodal_features = [
|
||||
{
|
||||
@@ -347,6 +358,8 @@ class EngineCoreReadyResponse:
|
||||
dp_stats_address: str | None
|
||||
dtype: str
|
||||
vllm_version: str
|
||||
world_size: int
|
||||
data_parallel_size: int
|
||||
kv_cache_size_tokens: int | None = None
|
||||
kv_cache_max_concurrency: float | None = None
|
||||
|
||||
@@ -358,9 +371,12 @@ ready_response = EngineCoreReadyResponse(
|
||||
dp_stats_address=None,
|
||||
dtype="float32",
|
||||
vllm_version="0.0.0",
|
||||
data_parallel_size=1,
|
||||
world_size=1,
|
||||
)
|
||||
|
||||
print(msgspec.msgpack.encode(request).hex())
|
||||
print(msgspec.msgpack.encode(defaults_request).hex())
|
||||
print(msgpack.packb(multimodal_request_wire, use_bin_type=True).hex())
|
||||
print(msgspec.msgpack.encode(outputs).hex())
|
||||
print(" ".join(frame.hex() for frame in encode_output_frames(inline_logprobs)))
|
||||
|
||||
@@ -327,6 +327,7 @@ pub async fn connect_handshake(
|
||||
pub async fn connect_bootstrapped(
|
||||
input_address: &str,
|
||||
output_address: &str,
|
||||
engine_start_index: u32,
|
||||
engine_count: usize,
|
||||
ready_timeout: Duration,
|
||||
) -> Result<ConnectedTransport> {
|
||||
@@ -342,8 +343,8 @@ pub async fn connect_bootstrapped(
|
||||
|
||||
let engines = wait_for_input_registrations(
|
||||
&mut input_socket,
|
||||
// TODO: follow start rank
|
||||
(0..engine_count).map(|index| EngineId::from((index as u16).to_le_bytes().to_vec())),
|
||||
(0..engine_count)
|
||||
.map(|offset| EngineId::from_engine_index(engine_start_index + offset as u32)),
|
||||
ready_timeout,
|
||||
)
|
||||
.await?;
|
||||
|
||||
@@ -71,6 +71,7 @@ impl ManagedEngineArgs {
|
||||
self,
|
||||
model: String,
|
||||
max_model_len: Option<u32>,
|
||||
max_logprobs: Option<i32>,
|
||||
language_model_only: bool,
|
||||
disable_log_stats: bool,
|
||||
shutdown_timeout: u64,
|
||||
@@ -82,6 +83,10 @@ impl ManagedEngineArgs {
|
||||
python_args.push("--max-model-len".to_string());
|
||||
python_args.push(max_model_len.to_string());
|
||||
}
|
||||
if let Some(max_logprobs) = max_logprobs {
|
||||
python_args.push("--max-logprobs".to_string());
|
||||
python_args.push(max_logprobs.to_string());
|
||||
}
|
||||
if language_model_only {
|
||||
python_args.push("--language-model-only".to_string());
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ axum.workspace = true
|
||||
educe.workspace = true
|
||||
futures.workspace = true
|
||||
http-body.workspace = true
|
||||
indexmap.workspace = true
|
||||
itertools.workspace = true
|
||||
libc.workspace = true
|
||||
llm-multimodal.workspace = true
|
||||
|
||||
@@ -14,8 +14,8 @@ use tokio_util::sync::CancellationToken;
|
||||
use tracing_subscriber::EnvFilter;
|
||||
use vllm_engine_core_client::TransportMode;
|
||||
use vllm_server::{
|
||||
ApiServerOptions, ChatTemplateContentFormatOption, Config, CoordinatorMode, HttpListenerMode,
|
||||
ParserSelection, RendererSelection, serve,
|
||||
ApiServerOptions, ChatTemplateContentFormatOption, Config, CoordinatorMode, CorsConfig,
|
||||
HttpListenerMode, ParserSelection, RendererSelection, serve,
|
||||
};
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
@@ -68,7 +68,9 @@ async fn main() -> Result<()> {
|
||||
chat_template: None,
|
||||
default_chat_template_kwargs: None,
|
||||
chat_template_content_format: ChatTemplateContentFormatOption::Auto,
|
||||
max_logprobs: None,
|
||||
api_server_options: ApiServerOptions::default(),
|
||||
cors: CorsConfig::default(),
|
||||
api_keys: Vec::new(),
|
||||
disable_log_stats: false,
|
||||
grpc_port: None,
|
||||
|
||||
@@ -2,7 +2,8 @@ use std::collections::HashMap;
|
||||
use std::fmt;
|
||||
use std::time::Duration;
|
||||
|
||||
use anyhow::Result;
|
||||
use anyhow::{Result, bail};
|
||||
use axum::http::{HeaderName, HeaderValue, Method};
|
||||
use educe::Educe;
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
@@ -45,6 +46,59 @@ pub struct ApiServerOptions {
|
||||
pub enable_request_id_headers: bool,
|
||||
}
|
||||
|
||||
/// CORS settings mirroring Python's `CORSMiddleware`; the default is permissive.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
|
||||
pub struct CorsConfig {
|
||||
/// Allowed origins. `["*"]` allows any origin.
|
||||
pub allow_origins: Vec<String>,
|
||||
/// Allowed methods. `["*"]` allows the standard method set.
|
||||
pub allow_methods: Vec<String>,
|
||||
/// Allowed request headers. `["*"]` mirrors the requested headers.
|
||||
pub allow_headers: Vec<String>,
|
||||
/// Whether to allow credentials (cookies, authorization headers).
|
||||
pub allow_credentials: bool,
|
||||
}
|
||||
|
||||
impl Default for CorsConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
allow_origins: vec!["*".to_string()],
|
||||
allow_methods: vec!["*".to_string()],
|
||||
allow_headers: vec!["*".to_string()],
|
||||
allow_credentials: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl CorsConfig {
|
||||
/// Validate that non-wildcard values parse into HTTP types, so the CORS
|
||||
/// layer can be built infallibly after startup validation has run.
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
for origin in &self.allow_origins {
|
||||
if origin != "*" {
|
||||
origin.parse::<HeaderValue>().map_err(|e| {
|
||||
anyhow::anyhow!("invalid --allowed-origins value {origin:?}: {e}")
|
||||
})?;
|
||||
}
|
||||
}
|
||||
for method in &self.allow_methods {
|
||||
if method != "*" {
|
||||
method.parse::<Method>().map_err(|e| {
|
||||
anyhow::anyhow!("invalid --allowed-methods value {method:?}: {e}")
|
||||
})?;
|
||||
}
|
||||
}
|
||||
for header in &self.allow_headers {
|
||||
if header != "*" {
|
||||
header.parse::<HeaderName>().map_err(|e| {
|
||||
anyhow::anyhow!("invalid --allowed-headers value {header:?}: {e}")
|
||||
})?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Normalized runtime configuration for the minimal OpenAI-compatible server.
|
||||
#[derive(Educe, Clone, PartialEq, Eq, Serialize)]
|
||||
#[educe(Debug)]
|
||||
@@ -77,8 +131,13 @@ pub struct Config {
|
||||
pub default_chat_template_kwargs: Option<HashMap<String, Value>>,
|
||||
/// How to serialize `message.content` for chat-template rendering.
|
||||
pub chat_template_content_format: ChatTemplateContentFormatOption,
|
||||
/// Optional maximum number of top log probabilities accepted by the
|
||||
/// frontend. `None` delegates to the text layer default.
|
||||
pub max_logprobs: Option<i32>,
|
||||
/// HTTP/API-server behavior switches.
|
||||
pub api_server_options: ApiServerOptions,
|
||||
/// CORS settings applied to every HTTP response.
|
||||
pub cors: CorsConfig,
|
||||
/// API keys accepted as bearer tokens for guarded routes.
|
||||
#[serde(skip_serializing)]
|
||||
#[educe(Debug(method(fmt_redacted_api_keys)))]
|
||||
@@ -98,6 +157,15 @@ impl Config {
|
||||
/// startup.
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
vllm_chat::validate_parser_overrides(&self.tool_call_parser, &self.reasoning_parser)?;
|
||||
self.cors.validate()?;
|
||||
if let Some(max_logprobs) = self.max_logprobs
|
||||
&& max_logprobs < -1
|
||||
{
|
||||
bail!(
|
||||
"max_logprobs must be non-negative or -1, got {}",
|
||||
max_logprobs
|
||||
);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -74,12 +74,11 @@ impl IntoResponse for ApiError {
|
||||
}
|
||||
}
|
||||
|
||||
/// Classify a text-pipeline submit failure: tokenized-prompt validation
|
||||
/// failures (the prompt is too long for the model, or empty after
|
||||
/// tokenization) are the client's fault and map to HTTP 400, mirroring the
|
||||
/// Python frontend. Everything else stays an internal 500.
|
||||
/// Classify a text-pipeline submit failure: request validation failures are
|
||||
/// the client's fault and map to HTTP 400, mirroring the Python frontend.
|
||||
/// Everything else stays an internal 500.
|
||||
pub fn text_submit_error(context: &'static str, error: vllm_text::Error) -> ApiError {
|
||||
if is_prompt_validation_error(&error) {
|
||||
if is_request_validation_error(&error) {
|
||||
return invalid_request!("{error}");
|
||||
}
|
||||
server_error!("{}: {}", context, error.to_report_string())
|
||||
@@ -90,18 +89,20 @@ pub fn text_submit_error(context: &'static str, error: vllm_text::Error) -> ApiE
|
||||
pub fn chat_submit_error(context: &'static str, error: vllm_chat::Error) -> ApiError {
|
||||
match &error {
|
||||
vllm_chat::Error::PromptTooLong { .. } => invalid_request!("{error}"),
|
||||
vllm_chat::Error::Text(text_error) if is_prompt_validation_error(text_error) => {
|
||||
vllm_chat::Error::Text(text_error) if is_request_validation_error(text_error) => {
|
||||
invalid_request!("{error}")
|
||||
}
|
||||
_ => server_error!("{}: {}", context, error.to_report_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_prompt_validation_error(error: &vllm_text::Error) -> bool {
|
||||
fn is_request_validation_error(error: &vllm_text::Error) -> bool {
|
||||
matches!(
|
||||
error,
|
||||
vllm_text::Error::PromptTooLong { .. }
|
||||
| vllm_text::Error::EmptyPromptTokenIds { .. }
|
||||
| vllm_text::Error::Logprobs(_)
|
||||
| vllm_text::Error::OutOfVocab(_)
|
||||
// An empty tokenized prompt detected later, at request prepare
|
||||
// time, surfaces through the transparent Llm wrapper.
|
||||
| vllm_text::Error::Llm(vllm_llm::Error::EmptyPromptTokenIds { .. })
|
||||
@@ -145,6 +146,44 @@ mod tests {
|
||||
assert_eq!(api_error.status_code(), StatusCode::BAD_REQUEST);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn logprobs_validation_maps_to_invalid_request() {
|
||||
let error = vllm_text::Error::Logprobs(vllm_text::LogprobsError::TooManyCount {
|
||||
parameter: "logprobs",
|
||||
requested: 1000,
|
||||
max_allowed: 20,
|
||||
});
|
||||
let api_error = text_submit_error("failed to submit completion request", error);
|
||||
assert_eq!(api_error.status_code(), StatusCode::BAD_REQUEST);
|
||||
let response = api_error.to_error_response();
|
||||
assert_eq!(response.error.error_type, "invalid_request_error");
|
||||
assert!(response.error.message.contains("logprobs"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_wrapped_logprobs_validation_maps_to_invalid_request() {
|
||||
let error = vllm_chat::Error::Text(vllm_text::Error::Logprobs(
|
||||
vllm_text::LogprobsError::TooManyCount {
|
||||
parameter: "prompt_logprobs",
|
||||
requested: 1000,
|
||||
max_allowed: 20,
|
||||
},
|
||||
));
|
||||
let api_error = chat_submit_error("failed to submit chat request", error);
|
||||
assert_eq!(api_error.status_code(), StatusCode::BAD_REQUEST);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn out_of_vocab_validation_maps_to_invalid_request() {
|
||||
let error = vllm_text::Error::OutOfVocab(vllm_text::OutOfVocabError {
|
||||
parameter: "logprob_token_ids",
|
||||
token_ids: vec![1000],
|
||||
vocab_size: 1000,
|
||||
});
|
||||
let api_error = text_submit_error("failed to submit completion request", error);
|
||||
assert_eq!(api_error.status_code(), StatusCode::BAD_REQUEST);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn other_submit_errors_stay_internal() {
|
||||
let error = vllm_text::Error::Tokenizer("backend exploded".to_string());
|
||||
|
||||
@@ -16,7 +16,7 @@ use std::sync::{Arc, OnceLock};
|
||||
use anyhow::{Context as _, Result};
|
||||
use axum::Router;
|
||||
use axum::serve::ListenerExt as _;
|
||||
pub use config::{ApiServerOptions, Config, CoordinatorMode, HttpListenerMode};
|
||||
pub use config::{ApiServerOptions, Config, CoordinatorMode, CorsConfig, HttpListenerMode};
|
||||
use tokio::net::TcpListener;
|
||||
use tokio::time::{Instant, sleep_until};
|
||||
use tokio_stream::wrappers::TcpListenerStream;
|
||||
@@ -90,7 +90,7 @@ async fn build_state(config: &Config) -> Result<Arc<AppState>> {
|
||||
.context("failed to connect to engine core")?;
|
||||
|
||||
let llm = Llm::new(client).with_log_stats(!config.disable_log_stats);
|
||||
let text = TextLlm::new(llm, text_backend);
|
||||
let text = TextLlm::new(llm, text_backend).with_max_logprobs(config.max_logprobs);
|
||||
|
||||
let chat = ChatLlm::new(text, chat_backend)
|
||||
.with_tool_call_parser(config.tool_call_parser.clone())
|
||||
@@ -98,9 +98,11 @@ async fn build_state(config: &Config) -> Result<Arc<AppState>> {
|
||||
|
||||
Ok(Arc::new(
|
||||
AppState::new(served_model_names, chat)
|
||||
.with_model_path(config.model.clone())
|
||||
.with_api_server_options(config.api_server_options)
|
||||
.with_server_info(ServerInfoSnapshot::from_config(config))
|
||||
.with_api_keys(config.api_keys.clone()),
|
||||
.with_api_keys(config.api_keys.clone())
|
||||
.with_cors(config.cors.clone()),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use indexmap::IndexMap;
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
use vllm_engine_core_client::EngineCoreClient;
|
||||
use vllm_engine_core_client::protocol::lora::LoraRequest;
|
||||
@@ -15,8 +15,8 @@ pub(crate) struct LoraModelResolution {
|
||||
|
||||
/// Runtime registry for dynamically loaded LoRA adapters.
|
||||
pub(crate) struct LoraManager {
|
||||
/// Dynamically loaded LoRA adapters keyed by public model name.
|
||||
requests: RwLock<BTreeMap<String, LoraRequest>>,
|
||||
/// Dynamically loaded LoRA adapters keyed by public model name, in load order.
|
||||
requests: RwLock<IndexMap<String, LoraRequest>>,
|
||||
/// Monotonic adapter id allocator. LoRA ids are one-indexed.
|
||||
id_counter: AtomicU64,
|
||||
/// Serialize dynamic LoRA registry updates around engine utility calls.
|
||||
@@ -51,18 +51,15 @@ pub(crate) enum UnloadLoraError {
|
||||
impl LoraManager {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
requests: RwLock::new(BTreeMap::new()),
|
||||
requests: RwLock::new(IndexMap::new()),
|
||||
id_counter: AtomicU64::new(0),
|
||||
update_lock: Mutex::new(()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Return base served model names plus dynamically loaded LoRA adapter
|
||||
/// names.
|
||||
pub async fn served_model_names(&self, base_model_names: &[String]) -> Vec<String> {
|
||||
let mut names = base_model_names.to_vec();
|
||||
names.extend(self.requests.read().await.keys().cloned());
|
||||
names
|
||||
/// Snapshot loaded LoRA adapters in load order.
|
||||
pub async fn served_lora_requests(&self) -> Vec<LoraRequest> {
|
||||
self.requests.read().await.values().cloned().collect()
|
||||
}
|
||||
|
||||
/// Resolve the requested model against one consistent LoRA registry
|
||||
@@ -163,6 +160,6 @@ impl LoraManager {
|
||||
});
|
||||
}
|
||||
|
||||
Ok(self.requests.write().await.remove(lora_name).unwrap_or(lora_request))
|
||||
Ok(self.requests.write().await.shift_remove(lora_name).unwrap_or(lora_request))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
//! CORS support mirroring Python's Starlette `CORSMiddleware`.
|
||||
//!
|
||||
//! Built on `tower_http::cors::CorsLayer`, configured to reproduce Starlette's
|
||||
//! `CORSMiddleware` behavior for the `--allowed-origins` / `--allowed-methods` /
|
||||
//! `--allowed-headers` / `--allow-credentials` settings. Two intentional
|
||||
//! behavioral differences remain, both invisible to real clients:
|
||||
//!
|
||||
//! - A rejected preflight returns `200` (empty) rather than Starlette's
|
||||
//! `400 "Disallowed CORS ..."`. The browser denies the request either way
|
||||
//! (the disallowed `Access-Control-Allow-*` headers are simply absent), and
|
||||
//! tower-http makes the preflight reject decision inside its short-circuit,
|
||||
//! so matching the `400` would mean re-implementing the layer.
|
||||
//! - A bare `OPTIONS` (no `Access-Control-Request-Method`) returns `200`
|
||||
//! rather than `405`. No real client sends one.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::Request;
|
||||
use axum::http::{HeaderName, HeaderValue, Method, header};
|
||||
use axum::middleware::Next;
|
||||
use axum::response::Response;
|
||||
use tower_http::cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer};
|
||||
|
||||
use crate::config::CorsConfig;
|
||||
|
||||
/// The method set that `"*"` expands to.
|
||||
const ALL_METHODS: [Method; 7] = [
|
||||
Method::DELETE,
|
||||
Method::GET,
|
||||
Method::HEAD,
|
||||
Method::OPTIONS,
|
||||
Method::PATCH,
|
||||
Method::POST,
|
||||
Method::PUT,
|
||||
];
|
||||
|
||||
/// Headers always treated as allowed (the CORS safelist).
|
||||
const SAFELISTED_HEADERS: [&str; 4] = [
|
||||
"accept",
|
||||
"accept-language",
|
||||
"content-language",
|
||||
"content-type",
|
||||
];
|
||||
|
||||
fn is_wildcard(values: &[String]) -> bool {
|
||||
values.iter().any(|value| value == "*")
|
||||
}
|
||||
|
||||
/// Build a `CorsLayer` from the resolved [`CorsConfig`].
|
||||
///
|
||||
/// Values are assumed valid: [`CorsConfig::validate`] runs at startup before
|
||||
/// the router is built.
|
||||
pub fn cors_layer(cfg: &CorsConfig) -> CorsLayer {
|
||||
let wildcard_origins = is_wildcard(&cfg.allow_origins);
|
||||
|
||||
let allow_origin = if wildcard_origins {
|
||||
if cfg.allow_credentials {
|
||||
// `*` with credentials is illegal, so reflect the request origin
|
||||
// instead; this also avoids tower-http's wildcard+credentials panic.
|
||||
AllowOrigin::mirror_request()
|
||||
} else {
|
||||
AllowOrigin::any()
|
||||
}
|
||||
} else {
|
||||
AllowOrigin::list(
|
||||
cfg.allow_origins
|
||||
.iter()
|
||||
.map(|origin| origin.parse::<HeaderValue>().expect("validated origin"))
|
||||
.collect::<Vec<_>>(),
|
||||
)
|
||||
};
|
||||
|
||||
// Expand `*` to an explicit list rather than `Any`, so we emit the method
|
||||
// names (not `*`) and never hit tower-http's `Any`+credentials panic.
|
||||
let allow_methods = if is_wildcard(&cfg.allow_methods) {
|
||||
AllowMethods::list(ALL_METHODS)
|
||||
} else {
|
||||
AllowMethods::list(
|
||||
cfg.allow_methods
|
||||
.iter()
|
||||
.map(|method| method.parse::<Method>().expect("validated method"))
|
||||
.collect::<Vec<_>>(),
|
||||
)
|
||||
};
|
||||
|
||||
let allow_headers = if is_wildcard(&cfg.allow_headers) {
|
||||
// `*` mirrors the requested headers.
|
||||
AllowHeaders::mirror_request()
|
||||
} else {
|
||||
// Union the safelisted headers, lowercased and sorted.
|
||||
let mut names: Vec<String> = SAFELISTED_HEADERS.iter().map(|s| s.to_string()).collect();
|
||||
names.extend(cfg.allow_headers.iter().map(|h| h.to_ascii_lowercase()));
|
||||
names.sort();
|
||||
names.dedup();
|
||||
AllowHeaders::list(
|
||||
names
|
||||
.iter()
|
||||
.map(|header| header.parse::<HeaderName>().expect("validated header"))
|
||||
.collect::<Vec<_>>(),
|
||||
)
|
||||
};
|
||||
|
||||
// Emit `Vary: Origin` only when the allow-origin is dynamic (explicit
|
||||
// origins, or credentials); the wildcard + no-credentials case emits no
|
||||
// `Vary` at all, and an empty list disables the header here.
|
||||
let vary: Vec<HeaderName> = if !wildcard_origins || cfg.allow_credentials {
|
||||
vec![header::ORIGIN]
|
||||
} else {
|
||||
vec![]
|
||||
};
|
||||
|
||||
CorsLayer::new()
|
||||
.allow_origin(allow_origin)
|
||||
.allow_methods(allow_methods)
|
||||
.allow_headers(allow_headers)
|
||||
.allow_credentials(cfg.allow_credentials)
|
||||
.max_age(Duration::from_secs(600))
|
||||
.vary(vary)
|
||||
}
|
||||
|
||||
/// Strip CORS response headers when the request carried no `Origin`.
|
||||
///
|
||||
/// A request without an `Origin` should carry no CORS headers, but tower-http
|
||||
/// emits `Vary` and `Access-Control-Allow-*` unconditionally. Removing them on
|
||||
/// no-`Origin` requests keeps non-CORS responses (e.g. `/health`, plain `curl`)
|
||||
/// clean.
|
||||
pub async fn strip_cors_on_no_origin(req: Request, next: Next) -> Response {
|
||||
let had_origin = req.headers().contains_key(header::ORIGIN);
|
||||
let mut response = next.run(req).await;
|
||||
if !had_origin {
|
||||
let headers = response.headers_mut();
|
||||
headers.remove(header::VARY);
|
||||
headers.remove(header::ACCESS_CONTROL_ALLOW_ORIGIN);
|
||||
headers.remove(header::ACCESS_CONTROL_ALLOW_CREDENTIALS);
|
||||
headers.remove(header::ACCESS_CONTROL_ALLOW_METHODS);
|
||||
headers.remove(header::ACCESS_CONTROL_ALLOW_HEADERS);
|
||||
headers.remove(header::ACCESS_CONTROL_MAX_AGE);
|
||||
headers.remove(header::ACCESS_CONTROL_EXPOSE_HEADERS);
|
||||
}
|
||||
response
|
||||
}
|
||||
@@ -1,9 +1,11 @@
|
||||
mod auth;
|
||||
mod cors;
|
||||
mod load;
|
||||
mod metrics;
|
||||
mod request_id;
|
||||
|
||||
pub use auth::authenticate_api_key;
|
||||
pub use cors::{cors_layer, strip_cors_on_no_origin};
|
||||
pub use load::track_server_load;
|
||||
pub use metrics::track_http_metrics;
|
||||
pub use request_id::set_request_id_header;
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
mod abort_requests;
|
||||
mod cache;
|
||||
mod collective_rpc;
|
||||
mod health;
|
||||
@@ -11,6 +12,7 @@ mod server_info;
|
||||
mod sleep;
|
||||
mod tokenize;
|
||||
mod version;
|
||||
mod world_size;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -91,6 +93,7 @@ fn build_router_with_options(
|
||||
.route("/reset_mm_cache", post(cache::reset_mm_cache))
|
||||
.route("/reset_encoder_cache", post(cache::reset_encoder_cache))
|
||||
.route("/collective_rpc", post(collective_rpc::collective_rpc))
|
||||
.route("/abort_requests", post(abort_requests::abort_requests))
|
||||
.route("/sleep", post(sleep::sleep))
|
||||
.route("/wake_up", post(sleep::wake_up))
|
||||
.route("/is_sleeping", get(sleep::is_sleeping))
|
||||
@@ -98,6 +101,7 @@ fn build_router_with_options(
|
||||
.route("/resume", post(pause::resume))
|
||||
.route("/is_paused", get(pause::is_paused))
|
||||
.route("/server_info", get(server_info::server_info))
|
||||
.route("/get_world_size", get(world_size::get_world_size))
|
||||
}
|
||||
|
||||
let enable_request_id_headers = state.api_server_options.enable_request_id_headers;
|
||||
@@ -108,7 +112,9 @@ fn build_router_with_options(
|
||||
state.clone(),
|
||||
middleware::track_server_load,
|
||||
))
|
||||
.layer(from_fn(middleware::track_http_metrics));
|
||||
.layer(from_fn(middleware::track_http_metrics))
|
||||
.layer(middleware::cors_layer(&state.cors))
|
||||
.layer(from_fn(middleware::strip_cors_on_no_origin));
|
||||
|
||||
if enable_api_key_auth {
|
||||
router = router.layer(from_fn_with_state(
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::Json;
|
||||
use axum::extract::State;
|
||||
use axum::extract::rejection::JsonRejection;
|
||||
use axum::http::StatusCode;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::error::ApiError;
|
||||
use crate::state::AppState;
|
||||
use crate::utils::utility_call_error;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(crate) struct AbortRequestsRequest {
|
||||
request_ids: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
pub async fn abort_requests(
|
||||
State(state): State<Arc<AppState>>,
|
||||
body: Result<Json<AbortRequestsRequest>, JsonRejection>,
|
||||
) -> Result<StatusCode, ApiError> {
|
||||
let Json(body) = body.map_err(|error| ApiError::json_parse_error(error.body_text()))?;
|
||||
let request_ids = body.request_ids.ok_or_else(|| {
|
||||
ApiError::invalid_request(
|
||||
"Missing 'request_ids' in request body".to_string(),
|
||||
Some("request_ids"),
|
||||
)
|
||||
})?;
|
||||
|
||||
state
|
||||
.chat
|
||||
.abort(&request_ids)
|
||||
.await
|
||||
.map_err(|error| utility_call_error("abort_requests", error))?;
|
||||
|
||||
Ok(StatusCode::OK)
|
||||
}
|
||||
@@ -53,13 +53,6 @@ pub async fn chat_completions(
|
||||
let request_context = resolve_request_context(&headers, body.request_id.as_deref());
|
||||
let lora_resolution = state.resolve_model_with_loras(Some(&body.model)).await;
|
||||
|
||||
if let Err(err) = validate::validate_token_id_ranges(
|
||||
&body,
|
||||
state.tokenizer_vocab_size(),
|
||||
state.model_vocab_size(),
|
||||
) {
|
||||
return err.into_response();
|
||||
}
|
||||
let prepared = match prepare_chat_request(body, &lora_resolution, request_context) {
|
||||
Ok(prepared) => prepared,
|
||||
Err(error) => return error.into_response(),
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
use super::types::ChatCompletionRequest;
|
||||
use crate::error::{ApiError, bail_invalid_request};
|
||||
use crate::routes::openai::utils::token_ids::{validate_allowed_token_ids, validate_logit_bias};
|
||||
use crate::routes::openai::utils::types::{ChatMessage, Tool, ToolChoice, ToolChoiceValue};
|
||||
|
||||
/// Enforce the minimal compatibility contract for the Rust OpenAI server.
|
||||
@@ -154,21 +153,6 @@ fn validate_function_tools(tools: &[Tool], param: &'static str) -> Result<(), Ap
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Reject out-of-vocab token ids, mirroring the Python input processor:
|
||||
/// `allowed_token_ids` against the tokenizer vocab, `logit_bias` keys against the
|
||||
/// model vocab (skipped when the model size is unknown).
|
||||
pub(super) fn validate_token_id_ranges(
|
||||
request: &ChatCompletionRequest,
|
||||
tokenizer_vocab_size: usize,
|
||||
model_vocab_size: Option<usize>,
|
||||
) -> Result<(), ApiError> {
|
||||
validate_allowed_token_ids(request.allowed_token_ids.as_deref(), tokenizer_vocab_size)?;
|
||||
validate_logit_bias(
|
||||
request.logit_bias.as_ref(),
|
||||
model_vocab_size.unwrap_or(usize::MAX),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::HashMap;
|
||||
@@ -176,7 +160,7 @@ mod tests {
|
||||
use serde_json::json;
|
||||
use vllm_chat::ReasoningEffort;
|
||||
|
||||
use super::{validate_request_compat, validate_token_id_ranges};
|
||||
use super::validate_request_compat;
|
||||
use crate::routes::openai::chat_completions::types::ChatCompletionRequest;
|
||||
use crate::routes::openai::utils::structured_outputs::ResponseFormat;
|
||||
use crate::routes::openai::utils::types::{
|
||||
@@ -188,32 +172,6 @@ mod tests {
|
||||
names.iter().map(|s| s.to_string()).collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_token_id_ranges_rejects_oob_and_accepts_in_vocab() {
|
||||
// allowed_token_ids are bounded by the tokenizer vocab
|
||||
let mut request = base_request();
|
||||
request.allowed_token_ids = Some(vec![5, 1_000_000]);
|
||||
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
|
||||
// logit_bias is bounded by the larger model vocab: an id between the two
|
||||
// vocabs is valid and must not be rejected (the parity regression we fix)
|
||||
let mut request = base_request();
|
||||
request.logit_bias = Some(HashMap::from([("150".to_string(), 1.0)]));
|
||||
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_ok());
|
||||
// logit_bias beyond the model vocab -> reject
|
||||
let mut request = base_request();
|
||||
request.logit_bias = Some(HashMap::from([("1000000".to_string(), 1.0)]));
|
||||
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
|
||||
// all in-vocab -> accept
|
||||
let mut request = base_request();
|
||||
request.allowed_token_ids = Some(vec![5, 50]);
|
||||
request.logit_bias = Some(HashMap::from([("50".to_string(), 1.0)]));
|
||||
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_ok());
|
||||
// unknown sizes -> skip
|
||||
let mut request = base_request();
|
||||
request.allowed_token_ids = Some(vec![1_000_000]);
|
||||
assert!(validate_token_id_ranges(&request, usize::MAX, None).is_ok());
|
||||
}
|
||||
|
||||
fn base_request() -> ChatCompletionRequest {
|
||||
ChatCompletionRequest {
|
||||
model: "Qwen/Qwen1.5-0.5B-Chat".to_string(),
|
||||
|
||||
@@ -2,6 +2,7 @@ mod convert;
|
||||
mod types;
|
||||
mod validate;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::convert::Infallible;
|
||||
use std::result::Result;
|
||||
use std::sync::Arc;
|
||||
@@ -16,12 +17,15 @@ use futures::{Stream, StreamExt as _, pin_mut};
|
||||
use thiserror_ext::AsReport as _;
|
||||
use tracing::{debug, error, info, trace};
|
||||
use tracing_futures::Instrument as _;
|
||||
use vllm_text::{DecodedTextEvent, FinishReason, TextOutputStream, TextOutputStreamExt as _};
|
||||
use vllm_text::{
|
||||
DecodedPromptLogprobs, DecodedTextEvent, FinishReason, TextOutputStream,
|
||||
TextOutputStreamExt as _,
|
||||
};
|
||||
|
||||
use self::convert::{ResponseOptions, prepare_completion_request};
|
||||
use super::utils::logprobs::{
|
||||
collected_logprobs_to_openai, decoded_logprobs_to_openai, decoded_prompt_logprobs_to_maps,
|
||||
text_len,
|
||||
decoded_prompt_logprobs_to_openai, text_len,
|
||||
};
|
||||
use super::utils::types::Usage;
|
||||
use crate::config::ApiServerOptions;
|
||||
@@ -47,13 +51,6 @@ pub async fn completions(
|
||||
let request_context = resolve_request_context(&headers, body.request_id.as_deref());
|
||||
let lora_resolution = state.resolve_model_with_loras(Some(&body.model)).await;
|
||||
|
||||
if let Err(err) = validate::validate_token_id_ranges(
|
||||
&body,
|
||||
state.tokenizer_vocab_size(),
|
||||
state.model_vocab_size(),
|
||||
) {
|
||||
return err.into_response();
|
||||
}
|
||||
let prepared = match prepare_completion_request(body, &lora_resolution, request_context) {
|
||||
Ok(prepared) => prepared,
|
||||
Err(error) => return error.into_response(),
|
||||
@@ -126,6 +123,7 @@ async fn collect_completion(
|
||||
include_usage: _,
|
||||
// Ignored: non-streaming responses are collected before usage is attached.
|
||||
include_continuous_usage: _,
|
||||
prompt_only,
|
||||
echo,
|
||||
requested_logprobs,
|
||||
include_prompt_logprobs,
|
||||
@@ -143,17 +141,17 @@ async fn collect_completion(
|
||||
.map(|sr| serde_json::to_value(sr).expect("StopReason must serialize to JSON"));
|
||||
|
||||
let prompt_char_count = echo.as_ref().map(|prompt| text_len(prompt)).unwrap_or_default();
|
||||
let prompt_logprobs = if include_prompt_logprobs {
|
||||
let prompt_logprobs = collected.prompt_logprobs.as_ref().ok_or_else(|| {
|
||||
server_error!(
|
||||
"completion response requested prompt_logprobs but generation returned none"
|
||||
)
|
||||
let logprobs = if requested_logprobs.is_some() && prompt_only {
|
||||
let prompt = echo.as_deref().ok_or_else(|| {
|
||||
server_error!("prompt-only completion response missing echoed prompt")
|
||||
})?;
|
||||
Some(prompt_logprobs)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let logprobs = if requested_logprobs.is_some() {
|
||||
Some(prompt_only_logprobs_to_openai(
|
||||
collected.prompt_logprobs.as_ref(),
|
||||
prompt,
|
||||
collected.prompt_token_ids.as_ref(),
|
||||
return_tokens_as_token_ids,
|
||||
)?)
|
||||
} else if requested_logprobs.is_some() {
|
||||
Some(collected_logprobs_to_openai(
|
||||
&collected,
|
||||
echo.is_some(),
|
||||
@@ -163,10 +161,18 @@ async fn collect_completion(
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let prompt_logprobs =
|
||||
prompt_logprobs.map(|lp| decoded_prompt_logprobs_to_maps(lp, return_tokens_as_token_ids));
|
||||
let prompt_logprobs = if include_prompt_logprobs {
|
||||
Some(prompt_logprobs_to_maps(
|
||||
collected.prompt_logprobs.as_ref(),
|
||||
collected.prompt_token_ids.as_ref(),
|
||||
return_tokens_as_token_ids,
|
||||
)?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let text = match &echo {
|
||||
None => collected.text,
|
||||
Some(prompt) if prompt_only => prompt.clone(),
|
||||
Some(prompt) => format!("{prompt}{}", collected.text),
|
||||
};
|
||||
let finish_reason = completion_finish_reason_to_openai(finish_reason)?.to_string();
|
||||
@@ -218,6 +224,7 @@ async fn completion_chunk_stream(
|
||||
ResponseOptions {
|
||||
include_usage,
|
||||
include_continuous_usage,
|
||||
prompt_only,
|
||||
echo,
|
||||
requested_logprobs,
|
||||
// Ignored: streaming prompt logprobs are rejected for Python parity.
|
||||
@@ -246,14 +253,30 @@ async fn completion_chunk_stream(
|
||||
while let Some(next) = stream.next().await {
|
||||
match next {
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids, ..
|
||||
prompt_token_ids,
|
||||
prompt_logprobs,
|
||||
}) => {
|
||||
debug!("completion stream started");
|
||||
continuous_usage.set_prompt_tokens(prompt_token_ids.len());
|
||||
if let Some(prompt) = echo.as_ref() {
|
||||
visible_text_len = text_len(prompt);
|
||||
let mut chunk =
|
||||
delta_chunk(&request_id, &response_model, created, prompt.clone(), None);
|
||||
let logprobs = if prompt_only && requested_logprobs.is_some() {
|
||||
Some(prompt_only_logprobs_to_openai(
|
||||
prompt_logprobs.as_ref(),
|
||||
prompt,
|
||||
prompt_token_ids.as_ref(),
|
||||
return_tokens_as_token_ids,
|
||||
)?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let mut chunk = delta_chunk(
|
||||
&request_id,
|
||||
&response_model,
|
||||
created,
|
||||
prompt.clone(),
|
||||
logprobs,
|
||||
);
|
||||
if return_token_ids && first_chunk {
|
||||
if let Some(choice) = chunk.choices.first_mut() {
|
||||
choice.prompt_token_ids = Some(prompt_token_ids.to_vec());
|
||||
@@ -278,6 +301,48 @@ async fn completion_chunk_stream(
|
||||
logprobs,
|
||||
finished,
|
||||
}) => {
|
||||
// Prompt-only streaming already emitted the echoed prompt in the Start chunk.
|
||||
// The one generated token is only used to drive the engine to a finished event,
|
||||
// so hide its delta and forward only the terminal finish/usage metadata.
|
||||
if prompt_only {
|
||||
if let Some(finished) = finished {
|
||||
if enable_log_requests {
|
||||
info!(
|
||||
stream = true,
|
||||
model = %response_model,
|
||||
prompt_tokens = finished.usage.prompt_token_count,
|
||||
output_tokens = finished.usage.output_token_count,
|
||||
finish_reason = finished.finish_reason.as_str(),
|
||||
"completion finished"
|
||||
);
|
||||
}
|
||||
continuous_usage.set_final_counts(
|
||||
finished.usage.prompt_token_count,
|
||||
finished.usage.output_token_count,
|
||||
);
|
||||
let final_chunk = final_chunk(
|
||||
&request_id,
|
||||
&response_model,
|
||||
created,
|
||||
finished.finish_reason,
|
||||
)?;
|
||||
yield_chunk!(final_chunk);
|
||||
|
||||
if include_usage {
|
||||
y.yield_ok(CompletionSseChunk::Usage(usage_chunk(
|
||||
&request_id,
|
||||
&response_model,
|
||||
created,
|
||||
Usage::from_token_usage(
|
||||
finished.usage,
|
||||
enable_prompt_tokens_details,
|
||||
),
|
||||
)))
|
||||
.await;
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
let delta_text_len = text_len(&delta);
|
||||
let logprobs = if requested_logprobs.is_some() {
|
||||
let decoded_logprobs = logprobs.as_ref().ok_or_else(|| {
|
||||
@@ -393,6 +458,57 @@ fn completion_finish_reason_to_openai(
|
||||
}
|
||||
}
|
||||
|
||||
fn prompt_only_logprobs_to_openai(
|
||||
prompt_logprobs: Option<&DecodedPromptLogprobs>,
|
||||
prompt: &str,
|
||||
prompt_token_ids: &[u32],
|
||||
return_tokens_as_token_ids: bool,
|
||||
) -> Result<LogProbs, ApiError> {
|
||||
if let Some(prompt_logprobs) = prompt_logprobs {
|
||||
return decoded_prompt_logprobs_to_openai(prompt_logprobs, 0, return_tokens_as_token_ids);
|
||||
}
|
||||
|
||||
if let [token_id] = prompt_token_ids {
|
||||
let token = if return_tokens_as_token_ids {
|
||||
format!("token_id:{token_id}")
|
||||
} else {
|
||||
prompt.to_string()
|
||||
};
|
||||
|
||||
return Ok(LogProbs {
|
||||
tokens: vec![token],
|
||||
token_logprobs: vec![None],
|
||||
top_logprobs: vec![None],
|
||||
text_offset: vec![0],
|
||||
});
|
||||
}
|
||||
|
||||
Err(server_error!(
|
||||
"prompt-only completion requested logprobs but generation returned none"
|
||||
))
|
||||
}
|
||||
|
||||
fn prompt_logprobs_to_maps(
|
||||
prompt_logprobs: Option<&DecodedPromptLogprobs>,
|
||||
prompt_token_ids: &[u32],
|
||||
return_tokens_as_token_ids: bool,
|
||||
) -> Result<Vec<Option<HashMap<String, f32>>>, ApiError> {
|
||||
if let Some(prompt_logprobs) = prompt_logprobs {
|
||||
return Ok(decoded_prompt_logprobs_to_maps(
|
||||
prompt_logprobs,
|
||||
return_tokens_as_token_ids,
|
||||
));
|
||||
}
|
||||
|
||||
if let [_token_id] = prompt_token_ids {
|
||||
return Ok(vec![None]);
|
||||
}
|
||||
|
||||
Err(server_error!(
|
||||
"completion response requested prompt_logprobs but generation returned none"
|
||||
))
|
||||
}
|
||||
|
||||
fn usage_chunk(
|
||||
request_id: &str,
|
||||
response_model: &str,
|
||||
@@ -456,8 +572,8 @@ mod tests {
|
||||
use futures::{StreamExt as _, stream};
|
||||
use itertools::Itertools as _;
|
||||
use vllm_text::{
|
||||
DecodedLogprobs, DecodedPositionLogprobs, DecodedTextEvent, DecodedTokenLogprob,
|
||||
FinishReason, Finished,
|
||||
DecodedLogprobs, DecodedPositionLogprobs, DecodedPromptLogprobs, DecodedTextEvent,
|
||||
DecodedTokenLogprob, FinishReason, Finished,
|
||||
};
|
||||
|
||||
use super::{
|
||||
@@ -620,4 +736,314 @@ mod tests {
|
||||
CompletionSseChunk::Chunk(_) => panic!("expected usage chunk"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn collect_completion_hides_internal_prompt_only_token() {
|
||||
let stream = stream::iter(vec![
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids: vec![1, 2].into(),
|
||||
prompt_logprobs: None,
|
||||
}),
|
||||
Ok(DecodedTextEvent::TextDelta {
|
||||
delta: " leaked".to_string(),
|
||||
token_ids: vec![3],
|
||||
logprobs: None,
|
||||
finished: Some(Finished {
|
||||
usage: vllm_llm::TokenUsage {
|
||||
prompt_token_count: 2,
|
||||
output_token_count: 1,
|
||||
cached_token_count: 0,
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
|
||||
let response = super::collect_completion(
|
||||
stream,
|
||||
"cmpl-1".to_string(),
|
||||
"model".to_string(),
|
||||
1,
|
||||
ApiServerOptions::default(),
|
||||
ResponseOptions {
|
||||
prompt_only: true,
|
||||
echo: Some("hello".to_string()),
|
||||
return_token_ids: true,
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("collect completion");
|
||||
|
||||
assert_eq!(response.choices[0].text, "hello");
|
||||
assert_eq!(response.choices[0].token_ids.as_deref(), Some(&[3][..]));
|
||||
assert_eq!(
|
||||
response.choices[0].prompt_token_ids.as_deref(),
|
||||
Some(&[1, 2][..])
|
||||
);
|
||||
let usage = response.usage.expect("usage");
|
||||
assert_eq!(usage.prompt_tokens, 2);
|
||||
assert_eq!(usage.completion_tokens, Some(1));
|
||||
assert_eq!(usage.total_tokens, 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn collect_completion_maps_prompt_logprobs_for_single_token_prompt() {
|
||||
let stream = stream::iter(vec![
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids: vec![9707].into(),
|
||||
prompt_logprobs: None,
|
||||
}),
|
||||
Ok(DecodedTextEvent::TextDelta {
|
||||
delta: " leaked".to_string(),
|
||||
token_ids: vec![3],
|
||||
logprobs: None,
|
||||
finished: Some(Finished {
|
||||
usage: vllm_llm::TokenUsage {
|
||||
prompt_token_count: 1,
|
||||
output_token_count: 1,
|
||||
cached_token_count: 0,
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
|
||||
let response = super::collect_completion(
|
||||
stream,
|
||||
"cmpl-1".to_string(),
|
||||
"model".to_string(),
|
||||
1,
|
||||
ApiServerOptions::default(),
|
||||
ResponseOptions {
|
||||
prompt_only: true,
|
||||
echo: Some("Hello".to_string()),
|
||||
requested_logprobs: Some(1),
|
||||
include_prompt_logprobs: true,
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("collect completion");
|
||||
|
||||
let choice = &response.choices[0];
|
||||
assert_eq!(choice.text, "Hello");
|
||||
assert_eq!(choice.prompt_logprobs, Some(vec![None]));
|
||||
let logprobs = choice.logprobs.as_ref().expect("logprobs");
|
||||
assert_eq!(logprobs.tokens, vec!["Hello".to_string()]);
|
||||
assert_eq!(logprobs.token_logprobs, vec![None]);
|
||||
assert_eq!(logprobs.top_logprobs, vec![None]);
|
||||
assert_eq!(logprobs.text_offset, vec![0]);
|
||||
let usage = response.usage.expect("usage");
|
||||
assert_eq!(usage.prompt_tokens, 1);
|
||||
assert_eq!(usage.completion_tokens, Some(1));
|
||||
assert_eq!(usage.total_tokens, 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn completion_chunk_stream_hides_internal_prompt_only_token() {
|
||||
let stream = stream::iter(vec![
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids: vec![1, 2].into(),
|
||||
prompt_logprobs: None,
|
||||
}),
|
||||
Ok(DecodedTextEvent::TextDelta {
|
||||
delta: " leaked".to_string(),
|
||||
token_ids: vec![3],
|
||||
logprobs: None,
|
||||
finished: Some(Finished {
|
||||
usage: vllm_llm::TokenUsage {
|
||||
prompt_token_count: 2,
|
||||
output_token_count: 1,
|
||||
cached_token_count: 0,
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
|
||||
let chunks = completion_chunk_stream(
|
||||
stream,
|
||||
"cmpl-1".to_string(),
|
||||
"model".to_string(),
|
||||
1,
|
||||
ApiServerOptions::default(),
|
||||
ResponseOptions {
|
||||
include_usage: true,
|
||||
prompt_only: true,
|
||||
echo: Some("hello".to_string()),
|
||||
return_token_ids: true,
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.collect::<Vec<_>>()
|
||||
.await;
|
||||
|
||||
let chunks: Vec<_> = chunks.into_iter().try_collect().expect("stream should succeed");
|
||||
assert_eq!(chunks.len(), 3);
|
||||
|
||||
match &chunks[0] {
|
||||
CompletionSseChunk::Chunk(chunk) => {
|
||||
assert_eq!(chunk.choices[0].text, "hello");
|
||||
assert_eq!(
|
||||
chunk.choices[0].prompt_token_ids.as_deref(),
|
||||
Some(&[1, 2][..])
|
||||
);
|
||||
}
|
||||
CompletionSseChunk::Usage(_) => panic!("expected prompt chunk"),
|
||||
}
|
||||
match &chunks[1] {
|
||||
CompletionSseChunk::Chunk(chunk) => {
|
||||
assert_eq!(chunk.choices[0].text, "");
|
||||
assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("length"));
|
||||
}
|
||||
CompletionSseChunk::Usage(_) => panic!("expected final chunk"),
|
||||
}
|
||||
match &chunks[2] {
|
||||
CompletionSseChunk::Usage(chunk) => {
|
||||
let usage = chunk.usage.as_ref().expect("usage");
|
||||
assert_eq!(usage.prompt_tokens, 2);
|
||||
assert_eq!(usage.completion_tokens, Some(1));
|
||||
assert_eq!(usage.total_tokens, 3);
|
||||
}
|
||||
CompletionSseChunk::Chunk(_) => panic!("expected usage chunk"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn completion_chunk_stream_maps_prompt_logprobs_for_single_token_prompt() {
|
||||
let stream = stream::iter(vec![
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids: vec![9707].into(),
|
||||
prompt_logprobs: None,
|
||||
}),
|
||||
Ok(DecodedTextEvent::TextDelta {
|
||||
delta: " leaked".to_string(),
|
||||
token_ids: vec![3],
|
||||
logprobs: None,
|
||||
finished: Some(Finished {
|
||||
usage: vllm_llm::TokenUsage {
|
||||
prompt_token_count: 1,
|
||||
output_token_count: 1,
|
||||
cached_token_count: 0,
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
|
||||
let chunks = completion_chunk_stream(
|
||||
stream,
|
||||
"cmpl-1".to_string(),
|
||||
"model".to_string(),
|
||||
1,
|
||||
ApiServerOptions::default(),
|
||||
ResponseOptions {
|
||||
prompt_only: true,
|
||||
echo: Some("Hello".to_string()),
|
||||
requested_logprobs: Some(1),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.collect::<Vec<_>>()
|
||||
.await;
|
||||
|
||||
let chunks: Vec<_> = chunks.into_iter().try_collect().expect("stream should succeed");
|
||||
assert_eq!(chunks.len(), 2);
|
||||
|
||||
match &chunks[0] {
|
||||
CompletionSseChunk::Chunk(chunk) => {
|
||||
assert_eq!(chunk.choices[0].text, "Hello");
|
||||
let logprobs = chunk.choices[0].logprobs.as_ref().expect("logprobs");
|
||||
assert_eq!(logprobs.tokens, vec!["Hello".to_string()]);
|
||||
assert_eq!(logprobs.token_logprobs, vec![None]);
|
||||
assert_eq!(logprobs.top_logprobs, vec![None]);
|
||||
assert_eq!(logprobs.text_offset, vec![0]);
|
||||
}
|
||||
CompletionSseChunk::Usage(_) => panic!("expected prompt chunk"),
|
||||
}
|
||||
match &chunks[1] {
|
||||
CompletionSseChunk::Chunk(chunk) => {
|
||||
assert_eq!(chunk.choices[0].text, "");
|
||||
assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("length"));
|
||||
}
|
||||
CompletionSseChunk::Usage(_) => panic!("expected final chunk"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn completion_chunk_stream_maps_prompt_only_logprobs() {
|
||||
let stream = stream::iter(vec![
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids: vec![1, 2].into(),
|
||||
prompt_logprobs: Some(DecodedPromptLogprobs {
|
||||
first_token_id: 1,
|
||||
first_token: "he".to_string(),
|
||||
scored_positions: vec![DecodedPositionLogprobs {
|
||||
entries: vec![DecodedTokenLogprob {
|
||||
token_id: 2,
|
||||
token: "llo".to_string(),
|
||||
logprob: -0.2,
|
||||
rank: 1,
|
||||
}],
|
||||
}],
|
||||
}),
|
||||
}),
|
||||
Ok(DecodedTextEvent::TextDelta {
|
||||
delta: " leaked".to_string(),
|
||||
token_ids: vec![3],
|
||||
logprobs: None,
|
||||
finished: Some(Finished {
|
||||
usage: vllm_llm::TokenUsage {
|
||||
prompt_token_count: 2,
|
||||
output_token_count: 1,
|
||||
cached_token_count: 0,
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
|
||||
let chunks = completion_chunk_stream(
|
||||
stream,
|
||||
"cmpl-1".to_string(),
|
||||
"model".to_string(),
|
||||
1,
|
||||
ApiServerOptions::default(),
|
||||
ResponseOptions {
|
||||
prompt_only: true,
|
||||
echo: Some("hello".to_string()),
|
||||
requested_logprobs: Some(1),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.collect::<Vec<_>>()
|
||||
.await;
|
||||
|
||||
let chunks: Vec<_> = chunks.into_iter().try_collect().expect("stream should succeed");
|
||||
assert_eq!(chunks.len(), 2);
|
||||
|
||||
match &chunks[0] {
|
||||
CompletionSseChunk::Chunk(chunk) => {
|
||||
assert_eq!(chunk.choices[0].text, "hello");
|
||||
let logprobs = chunk.choices[0].logprobs.as_ref().expect("logprobs");
|
||||
assert_eq!(logprobs.tokens, vec!["he".to_string(), "llo".to_string()]);
|
||||
assert_eq!(logprobs.token_logprobs, vec![None, Some(-0.2)]);
|
||||
assert_eq!(logprobs.text_offset, vec![0, 2]);
|
||||
}
|
||||
CompletionSseChunk::Usage(_) => panic!("expected prompt chunk"),
|
||||
}
|
||||
match &chunks[1] {
|
||||
CompletionSseChunk::Chunk(chunk) => {
|
||||
assert_eq!(chunk.choices[0].text, "");
|
||||
assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("length"));
|
||||
}
|
||||
CompletionSseChunk::Usage(_) => panic!("expected final chunk"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -27,6 +27,8 @@ pub(super) struct ResponseOptions {
|
||||
pub include_usage: bool,
|
||||
/// Whether every streamed chunk should carry cumulative usage.
|
||||
pub include_continuous_usage: bool,
|
||||
/// Whether the caller requested prompt-only echo via `max_tokens=0`.
|
||||
pub prompt_only: bool,
|
||||
/// Original text prompt that should be echoed back northbound when
|
||||
/// `echo=true`.
|
||||
pub echo: Option<String>,
|
||||
@@ -68,11 +70,13 @@ pub(super) fn prepare_completion_request(
|
||||
})?),
|
||||
None => None,
|
||||
};
|
||||
let prompt_logprobs = request.prompt_logprobs.or(if request.echo && !request.stream {
|
||||
logprobs
|
||||
} else {
|
||||
None
|
||||
});
|
||||
let prompt_only = request.echo && request.max_tokens == Some(0);
|
||||
let prompt_logprobs =
|
||||
request.prompt_logprobs.or(if request.echo && (!request.stream || prompt_only) {
|
||||
logprobs
|
||||
} else {
|
||||
None
|
||||
});
|
||||
let include_usage = (request.stream_options.as_ref())
|
||||
.and_then(|options| options.include_usage)
|
||||
.unwrap_or(false);
|
||||
@@ -83,6 +87,11 @@ pub(super) fn prepare_completion_request(
|
||||
.and_then(|options| options.continuous_usage_stats)
|
||||
.unwrap_or(false);
|
||||
let include_prompt_logprobs = prompt_logprobs.is_some();
|
||||
let max_tokens = if prompt_only {
|
||||
Some(1)
|
||||
} else {
|
||||
request.max_tokens
|
||||
};
|
||||
let echo = request.echo.then(|| request.prompt.as_text().cloned()).flatten();
|
||||
|
||||
let structured_outputs =
|
||||
@@ -97,7 +106,7 @@ pub(super) fn prepare_completion_request(
|
||||
top_p: request.top_p,
|
||||
top_k: request.top_k,
|
||||
seed: request.seed,
|
||||
max_tokens: request.max_tokens,
|
||||
max_tokens,
|
||||
min_tokens: request.min_tokens,
|
||||
logprobs,
|
||||
prompt_logprobs,
|
||||
@@ -138,6 +147,7 @@ pub(super) fn prepare_completion_request(
|
||||
options: ResponseOptions {
|
||||
include_usage,
|
||||
include_continuous_usage,
|
||||
prompt_only,
|
||||
echo,
|
||||
requested_logprobs: request.logprobs,
|
||||
include_prompt_logprobs,
|
||||
@@ -325,6 +335,57 @@ mod tests {
|
||||
|
||||
assert_eq!(prepared.options.echo, Some("hello".to_string()));
|
||||
assert_eq!(prepared.text_request.sampling_params.max_tokens, Some(7));
|
||||
assert!(!prepared.options.prompt_only);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_completion_request_lowers_prompt_only_echo_as_one_internal_token() {
|
||||
let request: CompletionRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"prompt": "hello",
|
||||
"stream": false,
|
||||
"echo": true,
|
||||
"max_tokens": 0
|
||||
}))
|
||||
.expect("parse request");
|
||||
|
||||
let prepared = prepare_completion_request(
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
assert!(prepared.options.prompt_only);
|
||||
assert_eq!(prepared.options.echo, Some("hello".to_string()));
|
||||
assert_eq!(prepared.text_request.sampling_params.max_tokens, Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepare_completion_request_enables_prompt_logprobs_for_stream_prompt_only_echo() {
|
||||
let request: CompletionRequest = serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
"prompt": "hello",
|
||||
"echo": true,
|
||||
"stream": true,
|
||||
"max_tokens": 0,
|
||||
"logprobs": 3
|
||||
}))
|
||||
.expect("parse request");
|
||||
|
||||
let prepared = prepare_completion_request(
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
assert!(prepared.options.prompt_only);
|
||||
assert_eq!(prepared.text_request.sampling_params.logprobs, Some(3));
|
||||
assert_eq!(
|
||||
prepared.text_request.sampling_params.prompt_logprobs,
|
||||
Some(3)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -2,9 +2,6 @@ use vllm_text::Prompt;
|
||||
|
||||
use super::types::CompletionRequest;
|
||||
use crate::error::{ApiError, bail_invalid_request};
|
||||
use crate::routes::openai::utils::token_ids::{
|
||||
validate_allowed_token_ids, validate_logit_bias, validate_prompt_token_ids,
|
||||
};
|
||||
|
||||
/// Enforce the minimal compatibility contract for the Rust OpenAI server.
|
||||
pub(super) fn validate_request_compat(
|
||||
@@ -29,8 +26,11 @@ pub(super) fn validate_request_compat(
|
||||
bail_invalid_request!(param = "n", "Only n=1 is supported.");
|
||||
}
|
||||
|
||||
if request.max_tokens == Some(0) {
|
||||
bail_invalid_request!(param = "max_tokens", "max_tokens must be greater than 0.");
|
||||
if request.max_tokens == Some(0) && !request.echo {
|
||||
bail_invalid_request!(
|
||||
param = "max_tokens",
|
||||
"max_tokens=0 is only supported when echo=true."
|
||||
);
|
||||
}
|
||||
|
||||
if request.echo && matches!(request.prompt, Prompt::TokenIds(_)) {
|
||||
@@ -98,63 +98,13 @@ pub(super) fn validate_request_compat(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Reject out-of-vocab token ids, mirroring the Python input processor. A token-id
|
||||
/// prompt may reference ids the engine embeds beyond either vocab alone (Qwen3
|
||||
/// extra LM tokens, multimodal placeholders), so it is bounded by the union of the
|
||||
/// tokenizer and model vocabularies; `allowed_token_ids` by the tokenizer vocab;
|
||||
/// `logit_bias` keys by the model vocab (skipped when the model size is unknown).
|
||||
pub(super) fn validate_token_id_ranges(
|
||||
request: &CompletionRequest,
|
||||
tokenizer_vocab_size: usize,
|
||||
model_vocab_size: Option<usize>,
|
||||
) -> Result<(), ApiError> {
|
||||
let prompt_bound = tokenizer_vocab_size.max(model_vocab_size.unwrap_or(0));
|
||||
validate_prompt_token_ids(&request.prompt, prompt_bound)?;
|
||||
validate_allowed_token_ids(request.allowed_token_ids.as_deref(), tokenizer_vocab_size)?;
|
||||
validate_logit_bias(
|
||||
request.logit_bias.as_ref(),
|
||||
model_vocab_size.unwrap_or(usize::MAX),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
use vllm_text::Prompt;
|
||||
|
||||
use super::{validate_request_compat, validate_token_id_ranges};
|
||||
use super::validate_request_compat;
|
||||
use crate::routes::openai::completions::types::CompletionRequest;
|
||||
|
||||
#[test]
|
||||
fn validate_token_id_ranges_rejects_oob_prompt_and_params() {
|
||||
// a token-id prompt below both vocabs is accepted (the engine can embed it)
|
||||
let mut request = base_request();
|
||||
request.prompt = Prompt::TokenIds(vec![5, 150]);
|
||||
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_ok());
|
||||
// an id at or above the union of the two vocabs is rejected
|
||||
let mut request = base_request();
|
||||
request.prompt = Prompt::TokenIds(vec![5, 200]);
|
||||
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
|
||||
// an id beyond the model vocab but within the (larger) tokenizer vocab is
|
||||
// accepted: the engine embeds added/placeholder ids above the model vocab,
|
||||
// matching the Python input processor's max(tokenizer, model) bound
|
||||
let mut request = base_request();
|
||||
request.prompt = Prompt::TokenIds(vec![150]);
|
||||
assert!(validate_token_id_ranges(&request, 200, Some(100)).is_ok());
|
||||
// falls back to the tokenizer vocab when the model size is unknown
|
||||
let mut request = base_request();
|
||||
request.prompt = Prompt::TokenIds(vec![150]);
|
||||
assert!(validate_token_id_ranges(&request, 100, None).is_err());
|
||||
// allowed_token_ids are bounded by the tokenizer vocab -> reject
|
||||
let mut request = base_request();
|
||||
request.allowed_token_ids = Some(vec![150]);
|
||||
assert!(validate_token_id_ranges(&request, 100, Some(200)).is_err());
|
||||
// unknown sizes -> skip
|
||||
let mut request = base_request();
|
||||
request.prompt = Prompt::TokenIds(vec![1_000_000]);
|
||||
assert!(validate_token_id_ranges(&request, usize::MAX, None).is_ok());
|
||||
}
|
||||
|
||||
fn base_request() -> CompletionRequest {
|
||||
serde_json::from_value(json!({
|
||||
"model": "Qwen/Qwen1.5-0.5B-Chat",
|
||||
@@ -219,4 +169,30 @@ mod tests {
|
||||
validate_request_compat(&request, &served_names(&["Qwen/Qwen1.5-0.5B-Chat"])).is_ok()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_request_compat_accepts_prompt_only_echo() {
|
||||
let request = CompletionRequest {
|
||||
stream: false,
|
||||
echo: true,
|
||||
max_tokens: Some(0),
|
||||
..base_request()
|
||||
};
|
||||
assert!(
|
||||
validate_request_compat(&request, &served_names(&["Qwen/Qwen1.5-0.5B-Chat"])).is_ok()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_request_compat_rejects_prompt_only_without_echo() {
|
||||
let request = CompletionRequest {
|
||||
stream: false,
|
||||
echo: false,
|
||||
max_tokens: Some(0),
|
||||
..base_request()
|
||||
};
|
||||
assert!(
|
||||
validate_request_compat(&request, &served_names(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use axum::Json;
|
||||
use axum::extract::State;
|
||||
@@ -6,19 +7,39 @@ use axum::extract::State;
|
||||
use crate::routes::openai::utils::types::{ListModelsResponse, ModelObject};
|
||||
use crate::state::AppState;
|
||||
|
||||
/// Return all configured served model names in OpenAI `list models` format.
|
||||
// Frontend marker; Python uses "vllm".
|
||||
const OWNED_BY: &str = "vllm-frontend-rs";
|
||||
|
||||
/// Base cards carry `max_model_len` and `root` = model path; LoRA cards carry
|
||||
/// `root` = adapter path and `parent` = base model. LoRA cards follow load order.
|
||||
pub async fn list_models(State(state): State<Arc<AppState>>) -> Json<ListModelsResponse> {
|
||||
let model_names = state.served_model_names_with_loras().await;
|
||||
let created = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_secs() as i64;
|
||||
let max_model_len = state.chat.engine_core_client().max_model_len();
|
||||
let model_path = state.model_path().map(str::to_string);
|
||||
|
||||
let base_cards = state.served_model_names().iter().map(|name| ModelObject {
|
||||
id: name.clone(),
|
||||
object: "model".to_string(),
|
||||
created,
|
||||
owned_by: OWNED_BY.to_string(),
|
||||
root: Some(model_path.clone().unwrap_or_else(|| name.clone())),
|
||||
parent: None,
|
||||
max_model_len: Some(max_model_len),
|
||||
});
|
||||
|
||||
let primary = state.primary_model_name().to_string();
|
||||
let lora_cards = state.served_lora_requests().await.into_iter().map(|lora| ModelObject {
|
||||
id: lora.lora_name,
|
||||
object: "model".to_string(),
|
||||
created,
|
||||
owned_by: OWNED_BY.to_string(),
|
||||
root: Some(lora.lora_path),
|
||||
parent: Some(lora.base_model_name.unwrap_or_else(|| primary.clone())),
|
||||
max_model_len: None,
|
||||
});
|
||||
|
||||
Json(ListModelsResponse {
|
||||
object: "list".to_string(),
|
||||
data: model_names
|
||||
.into_iter()
|
||||
.map(|name| ModelObject {
|
||||
id: name,
|
||||
object: "model".to_string(),
|
||||
created: 0,
|
||||
owned_by: "vllm-frontend-rs".to_string(),
|
||||
})
|
||||
.collect(),
|
||||
data: base_cards.chain(lora_cards).collect(),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
pub mod logprobs;
|
||||
pub mod structured_outputs;
|
||||
pub mod token_ids;
|
||||
pub mod types;
|
||||
pub mod usage;
|
||||
pub mod validated_json;
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use vllm_text::Prompt;
|
||||
|
||||
use crate::error::{ApiError, bail_invalid_request};
|
||||
|
||||
/// Reject token-id prompt entries at or above `bound` (the highest in-vocab id is
|
||||
/// `bound - 1`).
|
||||
pub(crate) fn validate_prompt_token_ids(prompt: &Prompt, bound: usize) -> Result<(), ApiError> {
|
||||
if let Prompt::TokenIds(ids) = prompt
|
||||
&& let Some(&bad) = ids.iter().find(|&&id| id as usize >= bound)
|
||||
{
|
||||
bail_invalid_request!(
|
||||
param = "prompt",
|
||||
"prompt contains out-of-vocab token id {bad}; vocabulary size is {bound}."
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Reject `allowed_token_ids` entries at or above `bound`.
|
||||
pub(crate) fn validate_allowed_token_ids(
|
||||
allowed_token_ids: Option<&[u32]>,
|
||||
bound: usize,
|
||||
) -> Result<(), ApiError> {
|
||||
if let Some(ids) = allowed_token_ids
|
||||
&& let Some(&bad) = ids.iter().find(|&&id| id as usize >= bound)
|
||||
{
|
||||
bail_invalid_request!(
|
||||
param = "allowed_token_ids",
|
||||
"allowed_token_ids contains out-of-vocab token id {bad}; vocabulary size is {bound}."
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Reject `logit_bias` keys at or above `bound`.
|
||||
pub(crate) fn validate_logit_bias(
|
||||
logit_bias: Option<&HashMap<String, f32>>,
|
||||
bound: usize,
|
||||
) -> Result<(), ApiError> {
|
||||
if let Some(bias) = logit_bias {
|
||||
for key in bias.keys() {
|
||||
if let Ok(id) = key.parse::<u32>()
|
||||
&& id as usize >= bound
|
||||
{
|
||||
bail_invalid_request!(
|
||||
param = "logit_bias",
|
||||
"logit_bias contains out-of-vocab token id {id}; vocabulary size is {bound}."
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -457,6 +457,12 @@ pub struct ModelObject {
|
||||
pub object: String,
|
||||
pub created: i64,
|
||||
pub owned_by: String,
|
||||
/// Backend model path (base cards) or adapter path (LoRA cards).
|
||||
pub root: Option<String>,
|
||||
/// Base model a LoRA adapter derives from; `null` for base models.
|
||||
pub parent: Option<String>,
|
||||
/// Maximum context length; `null` for LoRA adapter cards.
|
||||
pub max_model_len: Option<u32>,
|
||||
}
|
||||
|
||||
/// Response body for `GET /v1/models`.
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user