forked from Karylab-cklius/vllm
Compare commits
105
Commits
k3-release
...
v0.26.1rc0
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
53f6dd5c6f | ||
|
|
99115fcdcd | ||
|
|
1053e248f0 | ||
|
|
b5bcb3ce88 | ||
|
|
b2f9e4caa4 | ||
|
|
831d3848f1 | ||
|
|
fd10e8946d | ||
|
|
ed13deb376 | ||
|
|
99de48e98f | ||
|
|
bf2b45b5d6 | ||
|
|
8112b6c997 | ||
|
|
15d65f8669 | ||
|
|
e3c2fc3b3c | ||
|
|
2b465b2c42 | ||
|
|
3f47a8384d | ||
|
|
04502deca2 | ||
|
|
d2ca3002d9 | ||
|
|
27d7061ef6 | ||
|
|
ef9975d021 | ||
|
|
56c96b0d91 | ||
|
|
59a6b0411d | ||
|
|
dbccc5ae32 | ||
|
|
a89015c6df | ||
|
|
96fa3f42c9 | ||
|
|
81962bb699 | ||
|
|
92e8518d37 | ||
|
|
77cba0259f | ||
|
|
30fbd05537 | ||
|
|
0906123953 | ||
|
|
bc3629b1c4 | ||
|
|
312ea82e75 | ||
|
|
394beb633b | ||
|
|
7f599d7854 | ||
|
|
eb290ab673 | ||
|
|
cbc3a87200 | ||
|
|
afc94523c9 | ||
|
|
8061dc26bd | ||
|
|
fd9d2ede6f | ||
|
|
e09900436c | ||
|
|
5d07e268b1 | ||
|
|
d742856610 | ||
|
|
544cb724c8 | ||
|
|
c314af1abf | ||
|
|
f19ee27e39 | ||
|
|
5f89a03dcb | ||
|
|
53397fbfac | ||
|
|
49f31d7cee | ||
|
|
74d3b799e1 | ||
|
|
29fdeab254 | ||
|
|
ff6173997d | ||
|
|
8de50e46d4 | ||
|
|
da99ffcc13 | ||
|
|
8040ef2426 | ||
|
|
bf4f633b4c | ||
|
|
854c33f380 | ||
|
|
ac87549cbd | ||
|
|
439f336212 | ||
|
|
ffc4f08c8e | ||
|
|
50aa830482 | ||
|
|
f0553889c0 | ||
|
|
fdaa0d9e59 | ||
|
|
0934b26790 | ||
|
|
9e50e1037e | ||
|
|
b5b61c622c | ||
|
|
b68d7ef262 | ||
|
|
7154856f3d | ||
|
|
3f1d40960f | ||
|
|
0da6e7f3d6 | ||
|
|
5559679229 | ||
|
|
da3a252fd1 | ||
|
|
21fd9e85a0 | ||
|
|
8d28b48d01 | ||
|
|
0164022c90 | ||
|
|
30b0714031 | ||
|
|
2e860de498 | ||
|
|
7eca0e1a64 | ||
|
|
7a29a3c54c | ||
|
|
1240c74c0a | ||
|
|
48ebd6f2f1 | ||
|
|
b153ae6089 | ||
|
|
7a6a5b3667 | ||
|
|
0111002323 | ||
|
|
d30b1ecd1b | ||
|
|
dbd80cc031 | ||
|
|
70009fb934 | ||
|
|
ee1d996367 | ||
|
|
6b0103d1c9 | ||
|
|
9321aff536 | ||
|
|
26d725c334 | ||
|
|
7fe6d3c76b | ||
|
|
2e0da24150 | ||
|
+1 |
0b0bd2b5f6 | ||
|
|
33ef67e9fb | ||
|
|
3e74c60b9c | ||
|
|
d1a8ba63d9 | ||
|
|
1423569ff5 | ||
|
|
9a50464698 | ||
|
|
b9b6306ebe | ||
|
|
ca0defa343 | ||
|
|
0b1a8bb1f6 | ||
|
|
fe5145765f | ||
|
|
dbcc1cdd0a | ||
|
|
a82f1b388f | ||
|
|
190be7dad2 | ||
|
|
94682b79f4 |
@@ -0,0 +1,26 @@
|
||||
group: Benchmarks
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
steps:
|
||||
- label: Benchmarks CLI Test
|
||||
key: benchmarks-cli-test
|
||||
timeout_in_minutes: 40
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 1+
|
||||
mem: 16+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/benchmarks/
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
pytest -v -s benchmarks/'
|
||||
@@ -2,6 +2,44 @@ group: Engine Intel
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
steps:
|
||||
- label: Engine
|
||||
key: engine
|
||||
timeout_in_minutes: 40
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 1+
|
||||
mem: 16+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/compilation/
|
||||
- vllm/config/
|
||||
- vllm/engine/
|
||||
- vllm/entrypoints/logger.py
|
||||
- vllm/envs.py
|
||||
- vllm/logger.py
|
||||
- vllm/logging_utils/
|
||||
- vllm/platforms/
|
||||
- vllm/sequence.py
|
||||
- vllm/triton_utils/
|
||||
- vllm/utils/
|
||||
- tests/engine
|
||||
- tests/test_sequence
|
||||
- tests/test_config
|
||||
- tests/test_logger
|
||||
- tests/test_vllm_port
|
||||
- tests/test_jit_monitor.py
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
pytest -v -s engine/test_arg_utils.py test_sequence.py test_logger.py test_vllm_port.py test_jit_monitor.py'
|
||||
|
||||
- label: Engine (1 GPU)
|
||||
timeout_in_minutes: 30
|
||||
device: intel_gpu
|
||||
@@ -23,3 +61,41 @@ steps:
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
pytest -v -s v1/engine --ignore v1/engine/test_preprocess_error_handling.py'
|
||||
|
||||
- label: V1 e2e (2 GPUs)
|
||||
timeout_in_minutes: 30
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 2+
|
||||
mem: 16+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/compilation/
|
||||
- vllm/config/
|
||||
- vllm/distributed/
|
||||
- vllm/engine/
|
||||
- vllm/envs.py
|
||||
- vllm/forward_context.py
|
||||
- vllm/inputs/
|
||||
- vllm/logger.py
|
||||
- vllm/logging_utils/
|
||||
- vllm/model_executor/
|
||||
- vllm/multimodal/
|
||||
- vllm/platforms/
|
||||
- vllm/sampling_params.py
|
||||
- vllm/transformers_utils/
|
||||
- vllm/triton_utils/
|
||||
- vllm/utils/
|
||||
- vllm/v1/
|
||||
- tests/v1/e2e/spec_decode
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
pytest -v -s v1/e2e/spec_decode/test_spec_decode.py -k "tensor_parallelism"'
|
||||
|
||||
@@ -125,13 +125,13 @@ steps:
|
||||
pytest -v -s v1/kv_offload &&
|
||||
pytest -v -s v1/kv_connector/unit/test_offloading_connector.py'
|
||||
|
||||
- label: NixlConnector PD accuracy (2 GPUs)
|
||||
- label: NixlConnector PD accuracy (4 GPUs)
|
||||
timeout_in_minutes: 60
|
||||
num_devices: 2
|
||||
num_devices: 4
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 2+
|
||||
gpu: 4+
|
||||
mem: 16+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
@@ -148,7 +148,10 @@ steps:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh'
|
||||
bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh &&
|
||||
PREFILLER_TP_SIZE=2 DECODER_TP_SIZE=1 bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh &&
|
||||
PREFILLER_TP_SIZE=1 DECODER_TP_SIZE=2 bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh &&
|
||||
PREFILLER_TP_SIZE=2 DECODER_TP_SIZE=2 bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh'
|
||||
|
||||
- label: Regression
|
||||
key: regression
|
||||
@@ -259,3 +262,25 @@ steps:
|
||||
pytest -v -s detokenizer &&
|
||||
pytest -v -s -m "not cpu_test" ./multimodal &&
|
||||
pytest -v -s utils_ --ignore=utils_/test_mem_utils.py'
|
||||
|
||||
- label: Fusion Unit Tests
|
||||
timeout_in_minutes: 30
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 1+
|
||||
mem: 16+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/compilation/
|
||||
- tests/compile/passes/test_qk_norm_rope_fusion.py
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
pytest -v -s compile/passes/test_qk_norm_rope_fusion.py'
|
||||
@@ -0,0 +1,33 @@
|
||||
group: Model Executor Intel
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
steps:
|
||||
- label: Model Executor (Intel)
|
||||
key: model-executor-intel
|
||||
timeout_in_minutes: 45
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 1+
|
||||
mem: 24+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/engine/arg_utils.py
|
||||
- vllm/config/model.py
|
||||
- vllm/model_executor
|
||||
- tests/model_executor
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'apt-get update && apt-get install -y curl libsodium23 &&
|
||||
pip3 install tensorizer==2.10.1 &&
|
||||
pip3 install runai-model-streamer[s3,gcs,azure]\>=0.15.7 &&
|
||||
export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
|
||||
export PYTHONFAULTHANDLER=1 &&
|
||||
cd tests &&
|
||||
pytest -v -s model_executor -m "not slow_test" --ignore="model_executor/layers/test_rocm_unquantized_gemm.py" --deselect="tests/model_executor/model_loader/test_reload.py::test_kv_scale_reload"'
|
||||
@@ -8,7 +8,7 @@ steps:
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 2+
|
||||
mem: 16+
|
||||
mem: 24+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
@@ -28,7 +28,9 @@ steps:
|
||||
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
|
||||
cd tests &&
|
||||
pytest -v -s v1/engine/test_llm_engine.py -k "not test_engine_metrics" &&
|
||||
pytest -v -s v1/e2e/general/test_context_length.py &&
|
||||
ENFORCE_EAGER=1 pytest -v -s v1/e2e/general/test_async_scheduling.py -k "not ngram" &&
|
||||
pytest -v -s entrypoints/llm/test_struct_output_generate.py -k "xgrammar and not speculative_config6 and not speculative_config7 and not speculative_config8 and not speculative_config0" &&
|
||||
pytest -v -s v1/e2e/general/test_min_tokens.py'
|
||||
|
||||
- label: Model Runner V2 Examples (Intel)
|
||||
@@ -60,3 +62,55 @@ steps:
|
||||
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'
|
||||
|
||||
- label: Model Runner V2 Distributed (2 GPUs)
|
||||
timeout_in_minutes: 50
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 2+
|
||||
mem: 16+
|
||||
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
|
||||
- tests/basic_correctness/test_basic_correctness.py
|
||||
- tests/v1/distributed/test_async_llm_dp.py
|
||||
- tests/v1/distributed/test_eagle_dp.py
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
|
||||
cd tests &&
|
||||
TARGET_TEST_SUITE=L4 pytest -v -s basic_correctness/test_basic_correctness.py -m "distributed\(num_gpus=2\)" -k "not ray and not True"'
|
||||
|
||||
- label: Model Runner V2 Spec Decode
|
||||
timeout_in_minutes: 50
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 1+
|
||||
mem: 24+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/v1/worker/gpu/
|
||||
- vllm/v1/worker/gpu_worker.py
|
||||
- tests/v1/spec_decode/test_max_len.py
|
||||
- tests/v1/spec_decode/test_rejection_sampler_utils.py
|
||||
- tests/v1/e2e/spec_decode/test_spec_decode.py
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
|
||||
cd tests &&
|
||||
pytest -v -s v1/spec_decode/test_synthetic_rejection_sampler_utils.py'
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
group: Samplers Intel
|
||||
depends_on:
|
||||
- image-build-xpu
|
||||
steps:
|
||||
- label: Samplers Test (FlashInfer)
|
||||
key: samplers-test-flashinfer-intel
|
||||
timeout_in_minutes: 40
|
||||
device: intel_gpu
|
||||
agent_tags:
|
||||
label: production
|
||||
gpu: 1+
|
||||
mem: 24+
|
||||
no_plugin: true
|
||||
working_dir: "."
|
||||
env:
|
||||
REGISTRY: "public.ecr.aws/q9t5s3a7"
|
||||
REPO: "vllm-ci-test-repo"
|
||||
VLLM_TEST_DEVICE: "xpu"
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/layers
|
||||
- vllm/sampling_metadata.py
|
||||
- tests/samplers
|
||||
- tests/conftest.py
|
||||
- vllm/entrypoints/generate/beam_search
|
||||
commands:
|
||||
- >-
|
||||
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
|
||||
'cd tests &&
|
||||
VLLM_USE_FLASHINFER_SAMPLER=1 pytest -v -s samplers'
|
||||
@@ -7,6 +7,9 @@
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# The macmini queue uses persistent checkouts, so refresh tags for setuptools-scm.
|
||||
git fetch --tags --force origin
|
||||
|
||||
# The Rust frontend build needs protoc.
|
||||
if ! command -v protoc >/dev/null 2>&1; then
|
||||
brew install protobuf
|
||||
|
||||
@@ -387,6 +387,7 @@ initialize_native_environment() {
|
||||
local job_id="${BUILDKITE_JOB_ID:-${BUILDKITE_PARALLEL_JOB:-local}}"
|
||||
local job_id_suffix=""
|
||||
local native_root=""
|
||||
local hf_fstype=""
|
||||
local hf_mount=""
|
||||
|
||||
if [[ "$(id -u)" -ne 0 ]]; then
|
||||
@@ -405,11 +406,14 @@ initialize_native_environment() {
|
||||
VLLM_CACHE_ROOT="${native_root}/cache/vllm"
|
||||
XDG_CACHE_HOME="${native_root}/cache/xdg"
|
||||
: "${HF_HOME:=/home/buildkite-agent/huggingface}"
|
||||
# datasets uses POSIX locks that are unsupported by the shared HF NFS cache.
|
||||
# Keep processed datasets job-local while retaining the persistent Hub cache.
|
||||
HF_DATASETS_CACHE="${native_root}/cache/huggingface/datasets"
|
||||
: "${HF_HUB_DOWNLOAD_TIMEOUT:=300}"
|
||||
: "${HF_HUB_ETAG_TIMEOUT:=60}"
|
||||
export TMPDIR VLLM_RPC_BASE_PATH
|
||||
export TORCHINDUCTOR_CACHE_DIR TRITON_CACHE_DIR VLLM_CACHE_ROOT XDG_CACHE_HOME
|
||||
export HF_HOME HF_HUB_DOWNLOAD_TIMEOUT HF_HUB_ETAG_TIMEOUT
|
||||
export HF_HOME HF_DATASETS_CACHE HF_HUB_DOWNLOAD_TIMEOUT HF_HUB_ETAG_TIMEOUT
|
||||
export PYTORCH_ROCM_ARCH=""
|
||||
|
||||
mkdir -p "${TMPDIR}" \
|
||||
@@ -417,7 +421,8 @@ initialize_native_environment() {
|
||||
"${TRITON_CACHE_DIR}" \
|
||||
"${VLLM_CACHE_ROOT}" \
|
||||
"${XDG_CACHE_HOME}" \
|
||||
"${HF_HOME}" || return 1
|
||||
"${HF_HOME}" \
|
||||
"${HF_DATASETS_CACHE}" || return 1
|
||||
|
||||
echo "Native compile caches: VLLM_CACHE_ROOT=${VLLM_CACHE_ROOT} TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR}"
|
||||
|
||||
@@ -432,6 +437,18 @@ initialize_native_environment() {
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
|
||||
if command -v findmnt >/dev/null 2>&1; then
|
||||
hf_fstype=$(findmnt -n -T "${HF_HOME}" -o FSTYPE 2>/dev/null || true)
|
||||
fi
|
||||
if [[ "${hf_fstype}" == nfs || "${hf_fstype}" == nfs4 ]]; then
|
||||
# Keep hf-xet state local and avoid vectored writes on shared NFS.
|
||||
export HF_XET_CACHE="${native_root}/cache/hf-xet"
|
||||
export HF_XET_HIGH_PERFORMANCE=0
|
||||
export HF_XET_RECONSTRUCTION_USE_VECTORED_WRITE=0
|
||||
mkdir -p "${HF_XET_CACHE}" || return 1
|
||||
echo "Configured hf-xet for shared ${hf_fstype} cache at ${HF_HOME}"
|
||||
fi
|
||||
}
|
||||
|
||||
run_native_preflight() {
|
||||
|
||||
+25
-37
@@ -169,20 +169,6 @@ steps:
|
||||
- pip install helion==1.1.0
|
||||
- pytest -v -s kernels/helion/
|
||||
|
||||
- label: Kernels Mamba Test # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
|
||||
agent_pool: mi250_1
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- csrc/mamba/
|
||||
- tests/kernels/mamba
|
||||
- vllm/model_executor/layers/mamba/ops
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s kernels/mamba
|
||||
|
||||
#------------------------------------------------------ mi250 · models / basic -------------------------------------------------------#
|
||||
|
||||
- label: Basic Models Test (Other CPU) # TBD
|
||||
@@ -364,22 +350,6 @@ steps:
|
||||
commands:
|
||||
- pytest -v -s v1/e2e/spec_decode -k "speculators or mtp_correctness"
|
||||
|
||||
- label: V1 attention (H100-MI250) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
|
||||
agent_pool: mi250_1
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/config/attention.py
|
||||
- vllm/model_executor/layers/attention
|
||||
- vllm/v1/attention
|
||||
- tests/v1/attention
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/envs.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s v1/attention
|
||||
|
||||
- label: V1 others (CPU) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
|
||||
@@ -1577,11 +1547,12 @@ steps:
|
||||
commands:
|
||||
- pytest -v -s kernels/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
|
||||
|
||||
- label: Kernels Core Operation Test # TBD
|
||||
- label: Kernels Core Operation Test %N # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
dind: false
|
||||
agent_pool: mi300_1
|
||||
parallelism: 3
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
@@ -1592,7 +1563,7 @@ steps:
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s kernels/core --ignore=kernels/core/test_minimax_reduce_rms.py kernels/test_concat_mla_q.py kernels/test_top_k_per_row.py
|
||||
- pytest -v -s kernels/core --ignore=kernels/core/test_minimax_reduce_rms.py kernels/test_concat_mla_q.py kernels/test_top_k_per_row.py --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
|
||||
|
||||
- label: Kernels KDA Test # TBD
|
||||
timeout_in_minutes: 180
|
||||
@@ -1610,6 +1581,21 @@ steps:
|
||||
commands:
|
||||
- pytest -v -s kernels/test_kda.py
|
||||
|
||||
- label: Kernels Mamba Test # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
dind: false
|
||||
agent_pool: mi300_1
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- csrc/mamba/
|
||||
- tests/kernels/mamba
|
||||
- vllm/model_executor/layers/mamba/ops
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s kernels/mamba
|
||||
|
||||
- label: Kernels MoE Test %N # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
@@ -2161,7 +2147,7 @@ steps:
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
dind: false
|
||||
agent_pool: mi300_1
|
||||
parallelism: 4
|
||||
parallelism: 8
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/"
|
||||
source_file_dependencies:
|
||||
@@ -2694,11 +2680,12 @@ steps:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -v -s v1/kv_connector/extract_hidden_states_integration
|
||||
|
||||
- label: V1 attention (H100-MI300) # TBD
|
||||
- label: V1 attention (H100-MI300) %N # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
|
||||
dind: false
|
||||
agent_pool: mi300_1
|
||||
parallelism: 2
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
@@ -2710,7 +2697,7 @@ steps:
|
||||
- vllm/envs.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s v1/attention
|
||||
- pytest -v -s v1/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
|
||||
|
||||
- label: V1 Core + KV + Metrics # TBD
|
||||
timeout_in_minutes: 180
|
||||
@@ -3676,11 +3663,12 @@ steps:
|
||||
|
||||
#------------------------------------------------------------ mi355 · v1 -------------------------------------------------------------#
|
||||
|
||||
- label: V1 attention (B200-MI355) # TBD
|
||||
- label: V1 attention (B200-MI355) %N # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx950nightly, amdmi355]
|
||||
dind: false
|
||||
agent_pool: mi355_1
|
||||
parallelism: 2
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/config/attention.py
|
||||
@@ -3691,7 +3679,7 @@ steps:
|
||||
- vllm/envs.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s v1/attention
|
||||
- pytest -v -s v1/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
|
||||
|
||||
- label: V1 Core + KV + Metrics # TBD
|
||||
timeout_in_minutes: 180
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
group: Fault Tolerance
|
||||
depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Fault Tolerance E2E (2xH100)
|
||||
key: fault-tolerance-e2e-2xh100
|
||||
timeout_in_minutes: 35
|
||||
device: h100
|
||||
num_devices: 2
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/v1/fault_tolerance/
|
||||
- vllm/v1/worker/sentinel/
|
||||
- vllm/entrypoints/serve/fault_tolerance/
|
||||
- vllm/distributed/elastic_ep/
|
||||
- vllm/distributed/device_communicators/
|
||||
- vllm/v1/engine/
|
||||
- vllm/v1/worker/
|
||||
- tests/v1/fault_tolerance/
|
||||
- tests/v1/distributed/test_external_lb_dp.py
|
||||
commands:
|
||||
# Base image has no nixl; install it or has_nixl_ep() skips the tests.
|
||||
- bash /vllm-workspace/.buildkite/scripts/install-kv-connectors.sh
|
||||
# https://github.com/NVIDIA/nccl/issues/1838
|
||||
- export NCCL_CUMEM_HOST_ENABLE=0
|
||||
- pytest -v -s v1/fault_tolerance/test_fault_tolerance_e2e.py
|
||||
@@ -337,7 +337,7 @@ steps:
|
||||
|
||||
- label: LM Eval KV-Offload (2xH100)
|
||||
key: kv-offload-medium
|
||||
timeout_in_minutes: 30
|
||||
timeout_in_minutes: 45
|
||||
device: h100
|
||||
num_devices: 2
|
||||
source_file_dependencies:
|
||||
@@ -347,7 +347,7 @@ steps:
|
||||
- vllm/v1/simple_kv_offload/
|
||||
- tests/evals/gsm8k/test_gsm8k_offloading.py
|
||||
commands:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_offloading.py -k "qwen3.5-35b"
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_offloading.py -k "qwen3.5-35b or deepseek-v2-lite"
|
||||
|
||||
- label: LM Eval KV-Offload (4xH100)
|
||||
key: kv-offload-large
|
||||
|
||||
@@ -19,6 +19,7 @@ pull_request_rules:
|
||||
description: Comment on PR when pre-commit check fails
|
||||
conditions:
|
||||
- check-failure=pre-commit
|
||||
- -check-cancelled=pre-commit
|
||||
- -closed
|
||||
- -draft
|
||||
- or:
|
||||
@@ -232,6 +233,31 @@ pull_request_rules:
|
||||
add:
|
||||
- gpt-oss
|
||||
|
||||
- name: label-kimi
|
||||
description: Automatically apply kimi label
|
||||
conditions:
|
||||
- label != stale
|
||||
- or:
|
||||
- files~=(?i)kimi
|
||||
- files~=(?i)moonshot
|
||||
- title~=(?i)(?:kimi|moonshot)
|
||||
actions:
|
||||
label:
|
||||
add:
|
||||
- kimi
|
||||
|
||||
- name: label-k3
|
||||
description: Automatically apply k3 label (launch triage; retire after ramp-down)
|
||||
conditions:
|
||||
- label != stale
|
||||
- or:
|
||||
- files~=(?i)kimi[-_]?k3
|
||||
- title~=(?i)(?:kimi[-\s]?k3|\bk3\b)
|
||||
actions:
|
||||
label:
|
||||
add:
|
||||
- k3
|
||||
|
||||
- name: label-nvidia
|
||||
description: Automatically apply nvidia label
|
||||
conditions:
|
||||
|
||||
@@ -130,6 +130,25 @@ jobs:
|
||||
},
|
||||
],
|
||||
},
|
||||
kimi: {
|
||||
keywords: [
|
||||
{ term: "Kimi", searchIn: "both" },
|
||||
{ term: "Moonshot", searchIn: "both" },
|
||||
],
|
||||
substrings: [
|
||||
{ term: "moonshotai/", searchIn: "both" },
|
||||
{ term: "kimi", searchIn: "title" },
|
||||
],
|
||||
},
|
||||
k3: {
|
||||
keywords: [
|
||||
{ term: "Kimi K3", searchIn: "both" },
|
||||
{ term: "K3", searchIn: "title" },
|
||||
],
|
||||
substrings: [
|
||||
{ term: "moonshotai/kimi-k3", searchIn: "both" },
|
||||
],
|
||||
},
|
||||
quantization: {
|
||||
keywords: [
|
||||
{
|
||||
|
||||
@@ -173,9 +173,6 @@ venv.bak/
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
docs/argparse
|
||||
docs/examples/*
|
||||
!docs/examples/README.md
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
|
||||
@@ -3,6 +3,9 @@ MD007:
|
||||
MD013: false
|
||||
MD024:
|
||||
siblings_only: true
|
||||
MD025:
|
||||
# Allow front matter title to be different from the first heading in the document.
|
||||
front_matter_title: ""
|
||||
MD031:
|
||||
list_items: false
|
||||
MD033: false
|
||||
|
||||
@@ -260,10 +260,6 @@ repos:
|
||||
files: ^docker/(Dockerfile|versions\.json)$
|
||||
pass_filenames: false
|
||||
additional_dependencies: [dockerfile-parse]
|
||||
- id: attention-backend-docs
|
||||
name: Check attention backend documentation is up to date
|
||||
entry: python tools/pre_commit/generate_attention_backend_docs.py --check
|
||||
language: python
|
||||
- id: check-boolean-context-manager
|
||||
name: Check for boolean ops in with-statements
|
||||
entry: python tools/pre_commit/check_boolean_context_manager.py
|
||||
|
||||
@@ -1358,6 +1358,10 @@ def main():
|
||||
profile_memory=args.profile_memory,
|
||||
warmup_ms=args.warmup_ms,
|
||||
prefill_backend=pb,
|
||||
kv_lora_rank=args.kv_lora_rank,
|
||||
qk_nope_head_dim=args.qk_nope_head_dim,
|
||||
qk_rope_head_dim=args.qk_rope_head_dim,
|
||||
v_head_dim=args.v_head_dim,
|
||||
)
|
||||
|
||||
result = run_benchmark(config)
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import statistics
|
||||
|
||||
import torch
|
||||
from tabulate import tabulate
|
||||
|
||||
from vllm.models.inkling.nvidia.ops import qkvr_prep
|
||||
from vllm.utils.argparse_utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
def make_inputs(tokens: int, tp_size: int, is_local: bool):
|
||||
torch.manual_seed(0)
|
||||
num_q_heads = 64 // tp_size
|
||||
num_kv_heads = (16 if is_local else 8) // tp_size
|
||||
head_dim = 128
|
||||
d_rel = 16
|
||||
rel_extent = 512 if is_local else 1024
|
||||
page_size = 16
|
||||
num_blocks = (tokens + page_size - 1) // page_size
|
||||
q_width = num_q_heads * head_dim
|
||||
kv_width = num_kv_heads * head_dim
|
||||
r_width = num_q_heads * d_rel
|
||||
device = "cuda"
|
||||
|
||||
qkvr = torch.randn(
|
||||
tokens,
|
||||
q_width + 2 * kv_width + r_width,
|
||||
device=device,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
k_weight = torch.randn(kv_width, 4, device=device, dtype=torch.bfloat16)
|
||||
v_weight = torch.randn_like(k_weight)
|
||||
q_norm_weight = torch.randn(head_dim, device=device, dtype=torch.bfloat16)
|
||||
k_norm_weight = torch.randn_like(q_norm_weight)
|
||||
rel_proj = torch.randn(d_rel, rel_extent, device=device, dtype=torch.bfloat16)
|
||||
conv_cache = torch.zeros(
|
||||
num_blocks,
|
||||
num_kv_heads,
|
||||
page_size,
|
||||
2 * head_dim,
|
||||
device=device,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
key_cache = torch.empty(
|
||||
num_blocks,
|
||||
page_size,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
value_cache = torch.empty_like(key_cache)
|
||||
positions = torch.arange(tokens, device=device, dtype=torch.int64)
|
||||
block_table = torch.arange(num_blocks, device=device, dtype=torch.int32)[None]
|
||||
seq_idx = torch.zeros(tokens, device=device, dtype=torch.int32)
|
||||
slots = torch.arange(tokens, device=device, dtype=torch.int64)
|
||||
query_start = torch.zeros(tokens, device=device, dtype=torch.int32)
|
||||
log_scaling = None
|
||||
if not is_local:
|
||||
effective_n = (positions + 1).to(torch.float32)
|
||||
log_scaling = 1.0 + 0.1 * torch.log(torch.clamp(effective_n / 128000, min=1.0))
|
||||
return (
|
||||
qkvr,
|
||||
k_weight,
|
||||
v_weight,
|
||||
q_norm_weight,
|
||||
k_norm_weight,
|
||||
rel_proj,
|
||||
1e-6,
|
||||
num_q_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
d_rel,
|
||||
conv_cache,
|
||||
key_cache,
|
||||
value_cache,
|
||||
positions,
|
||||
block_table,
|
||||
seq_idx,
|
||||
slots,
|
||||
query_start,
|
||||
slots,
|
||||
0,
|
||||
head_dim,
|
||||
page_size,
|
||||
log_scaling,
|
||||
)
|
||||
|
||||
|
||||
def capture(implementation, inputs):
|
||||
outputs = []
|
||||
|
||||
def run():
|
||||
outputs[:] = implementation.fused_qkvr_prep(*inputs)
|
||||
|
||||
stream = torch.cuda.Stream()
|
||||
stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.cuda.stream(stream):
|
||||
for _ in range(3):
|
||||
run()
|
||||
torch.cuda.current_stream().wait_stream(stream)
|
||||
torch.accelerator.synchronize()
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
run()
|
||||
torch.accelerator.synchronize()
|
||||
return graph, outputs
|
||||
|
||||
|
||||
def time_graph(graph: torch.cuda.CUDAGraph, warmup: int, repeats: int) -> float:
|
||||
for _ in range(warmup):
|
||||
graph.replay()
|
||||
torch.accelerator.synchronize()
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
for _ in range(repeats):
|
||||
graph.replay()
|
||||
end.record()
|
||||
end.synchronize()
|
||||
return start.elapsed_time(end) * 1000 / repeats
|
||||
|
||||
|
||||
def benchmark(inputs, args) -> float:
|
||||
graph, _ = capture(qkvr_prep, inputs)
|
||||
return statistics.median(
|
||||
time_graph(graph, args.warmup, args.repeats) for _ in range(args.trials)
|
||||
)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main(args):
|
||||
rows = []
|
||||
for tp_size in args.tp_sizes:
|
||||
for tokens in args.tokens:
|
||||
for is_local in (True, False):
|
||||
triton_us = benchmark(make_inputs(tokens, tp_size, is_local), args)
|
||||
rows.append(
|
||||
[
|
||||
tp_size,
|
||||
tokens,
|
||||
"local" if is_local else "global",
|
||||
triton_us,
|
||||
]
|
||||
)
|
||||
|
||||
print("Inkling QKVR prep (CUDA graph, median latency)")
|
||||
print(
|
||||
tabulate(
|
||||
rows,
|
||||
headers=[
|
||||
"TP",
|
||||
"tokens",
|
||||
"scope",
|
||||
"Triton (us)",
|
||||
],
|
||||
floatfmt=("d", "d", "", ".2f"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = FlexibleArgumentParser()
|
||||
parser.add_argument(
|
||||
"--tokens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=[1 << power for power in range(15)],
|
||||
)
|
||||
parser.add_argument("--tp-sizes", type=int, nargs="+", default=[4, 8])
|
||||
parser.add_argument("--warmup", type=int, default=20)
|
||||
parser.add_argument("--repeats", type=int, default=200)
|
||||
parser.add_argument("--trials", type=int, default=5)
|
||||
main(parser.parse_args())
|
||||
@@ -154,7 +154,7 @@ def main(
|
||||
scale=scale,
|
||||
causal=True,
|
||||
alibi_slopes=None,
|
||||
sliding_window=window_size,
|
||||
sliding_window=window_size if sliding_window is not None else -1,
|
||||
block_table=block_tables,
|
||||
softcap=0,
|
||||
scheduler_metadata=metadata,
|
||||
|
||||
@@ -15,6 +15,7 @@ endif()
|
||||
#
|
||||
set(ENABLE_X86_ISA $ENV{VLLM_CPU_X86})
|
||||
set(ENABLE_ARM_BF16 $ENV{VLLM_CPU_ARM_BF16})
|
||||
set(ENABLE_ARM_I8MM $ENV{VLLM_CPU_ARM_I8MM})
|
||||
set(ENABLE_RVV_BF16 $ENV{VLLM_CPU_RVV_BF16})
|
||||
|
||||
include_directories("${CMAKE_SOURCE_DIR}/csrc")
|
||||
@@ -96,12 +97,14 @@ if (MACOSX_FOUND AND CMAKE_SYSTEM_PROCESSOR STREQUAL "arm64")
|
||||
set(ENABLE_NUMA OFF)
|
||||
check_sysctl(hw.optional.neon ASIMD_FOUND)
|
||||
check_sysctl(hw.optional.arm.FEAT_BF16 ARM_BF16_FOUND)
|
||||
check_sysctl(hw.optional.arm.FEAT_I8MM ARM_I8MM_FOUND)
|
||||
else()
|
||||
find_isa(${CPUINFO} "Power11" POWER11_FOUND)
|
||||
find_isa(${CPUINFO} "POWER10" POWER10_FOUND)
|
||||
find_isa(${CPUINFO} "POWER9" POWER9_FOUND)
|
||||
find_isa(${CPUINFO} "asimd" ASIMD_FOUND) # Check for ARM NEON support
|
||||
find_isa(${CPUINFO} "bf16" ARM_BF16_FOUND) # Check for ARM BF16 support
|
||||
find_isa(${CPUINFO} "i8mm" ARM_I8MM_FOUND) # Check for ARM I8MM support
|
||||
find_isa(${CPUINFO} "S390" S390_FOUND)
|
||||
find_isa(${CPUINFO} "zvfhmin" RVV_FP16_FOUND) # Check for RISC-V Vector FP16 support
|
||||
find_isa(${CPUINFO} "zvfbfmin" RVV_BF16_FOUND) # Check for RISC-V Vector BF16 support
|
||||
@@ -111,6 +114,11 @@ else()
|
||||
set(ARM_BF16_FOUND ON)
|
||||
message(STATUS "ARM BF16 support enabled via VLLM_CPU_ARM_BF16 environment variable")
|
||||
endif()
|
||||
if (ENABLE_ARM_I8MM)
|
||||
set(ARM_I8MM_FOUND ON)
|
||||
message(STATUS
|
||||
"ARM I8MM support enabled via VLLM_CPU_ARM_I8MM environment variable")
|
||||
endif()
|
||||
# Some kernels (e.g. Bianbu on Spacemit X100) do not report zvfbfmin
|
||||
# in /proc/cpuinfo despite hardware support. VLLM_CPU_RVV_BF16=1
|
||||
# overrides the detection result.
|
||||
@@ -166,6 +174,11 @@ elseif (ASIMD_FOUND)
|
||||
message(WARNING "BF16 functionality is not available")
|
||||
set(MARCH_FLAGS "-march=armv8.2-a+dotprod+fp16")
|
||||
endif()
|
||||
if(ARM_I8MM_FOUND)
|
||||
message(STATUS "I8MM extension detected")
|
||||
string(APPEND MARCH_FLAGS "+i8mm")
|
||||
add_compile_definitions(ARM_I8MM_SUPPORT)
|
||||
endif()
|
||||
list(APPEND CXX_COMPILE_FLAGS ${MARCH_FLAGS})
|
||||
elseif (S390_FOUND)
|
||||
message(STATUS "S390 detected")
|
||||
@@ -447,8 +460,13 @@ if (ASIMD_FOUND AND NOT APPLE_SILICON_FOUND)
|
||||
"csrc/cpu/shm.cpp"
|
||||
"csrc/cpu/activation_lut_bf16.cpp"
|
||||
"csrc/cpu/cpu_tanhf_neon.hpp"
|
||||
"csrc/cpu/cpu_fused_moe.cpp"
|
||||
${VLLM_EXT_SRC})
|
||||
if (ARM_BF16_FOUND)
|
||||
set(VLLM_EXT_SRC "csrc/cpu/cpu_fused_moe.cpp" ${VLLM_EXT_SRC})
|
||||
if (ARM_I8MM_FOUND)
|
||||
set(VLLM_EXT_SRC "csrc/cpu/cpu_fused_moe_int8.cpp" ${VLLM_EXT_SRC})
|
||||
endif()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if (POWER9_FOUND OR POWER10_FOUND OR POWER11_FOUND)
|
||||
|
||||
@@ -172,4 +172,15 @@
|
||||
|
||||
#endif // __riscv_v
|
||||
|
||||
// Power VSX
|
||||
#ifdef __powerpc__
|
||||
// FP32Vec16::exp() in cpu_types_vsx.hpp delegates to FP32Vec8::exp(), which
|
||||
// implements a vectorised 5-term minimax polynomial using VSX intrinsics.
|
||||
#define DEFINE_FAST_EXP \
|
||||
auto fast_exp = [&](const vec_op::FP32Vec16& vec) \
|
||||
__attribute__((always_inline)) { return vec.exp(); }; \
|
||||
auto fast_exp_f16 = fast_exp;
|
||||
|
||||
#endif // __powerpc__
|
||||
|
||||
#endif
|
||||
|
||||
+5
-187
@@ -1,5 +1,6 @@
|
||||
#include "cpu/cpu_types.hpp"
|
||||
#include "cpu/utils.hpp"
|
||||
#include "cpu/cpu_fused_moe_activations.hpp"
|
||||
#include "cpu/micro_gemm/cpu_micro_gemm_vec.hpp"
|
||||
#include "cpu/cpu_arch_macros.h"
|
||||
|
||||
@@ -43,193 +44,9 @@
|
||||
}()
|
||||
|
||||
namespace {
|
||||
enum class FusedMOEAct {
|
||||
SiluAndMul,
|
||||
SwigluOAIAndMul,
|
||||
GeluAndMul,
|
||||
GeluTanhAndMul,
|
||||
};
|
||||
|
||||
FusedMOEAct get_act_type(const std::string& act) {
|
||||
if (act == "silu") {
|
||||
return FusedMOEAct::SiluAndMul;
|
||||
} else if (act == "swigluoai") {
|
||||
return FusedMOEAct::SwigluOAIAndMul;
|
||||
} else if (act == "gelu") {
|
||||
return FusedMOEAct::GeluAndMul;
|
||||
} else if (act == "gelu_tanh") {
|
||||
return FusedMOEAct::GeluTanhAndMul;
|
||||
} else {
|
||||
TORCH_CHECK(false, "Invalid act type: " + act);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void swigluoai_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride,
|
||||
const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
#if !defined(__aarch64__)
|
||||
// For GPT-OSS interleaved gate-up weights
|
||||
alignas(64) static int32_t index[16] = {0, 2, 4, 6, 8, 10, 12, 14,
|
||||
16, 18, 20, 22, 24, 26, 28, 30};
|
||||
vec_op::INT32Vec16 index_vec(index);
|
||||
#endif
|
||||
vec_op::FP32Vec16 gate_up_max_vec(7.0);
|
||||
vec_op::FP32Vec16 up_min_vec(-7.0);
|
||||
vec_op::FP32Vec16 alpha_vec(1.702);
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
|
||||
DEFINE_FAST_EXP
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < n_size; n += 32) {
|
||||
// Note: AdvSIMD does not support gather loads
|
||||
#if defined(__aarch64__)
|
||||
vec_op::FP32Vec16 gate_vec(vec_op::uninit);
|
||||
vec_op::FP32Vec16 up_vec(vec_op::uninit);
|
||||
vec_op::FP32Vec16::load_even_odd(input + n, gate_vec, up_vec);
|
||||
#else
|
||||
vec_op::FP32Vec16 gate_vec(input + n, index_vec);
|
||||
vec_op::FP32Vec16 up_vec(input + n + 1, index_vec);
|
||||
#endif
|
||||
gate_vec = gate_vec.min(gate_up_max_vec);
|
||||
up_vec = up_vec.clamp(up_min_vec, gate_up_max_vec);
|
||||
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec * alpha_vec));
|
||||
auto glu = gate_vec * sigmoid_vec;
|
||||
auto gated_output_fp32 = (one_vec + up_vec) * glu;
|
||||
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
|
||||
gated_output.save(output + n / 2);
|
||||
}
|
||||
input += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void silu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride, const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
const int32_t dim = n_size / 2;
|
||||
float* __restrict__ gate = input;
|
||||
float* __restrict__ up = input + dim;
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
|
||||
DEFINE_FAST_EXP
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < dim; n += 16) {
|
||||
vec_op::FP32Vec16 gate_vec(gate + n);
|
||||
vec_op::FP32Vec16 up_vec(up + n);
|
||||
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec));
|
||||
auto silu = gate_vec * sigmoid_vec;
|
||||
auto gated_output_fp32 = up_vec * silu;
|
||||
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
|
||||
gated_output.save(output + n);
|
||||
}
|
||||
gate += input_stride;
|
||||
up += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void gelu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride, const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
const int32_t dim = n_size / 2;
|
||||
float* __restrict__ gate = input;
|
||||
float* __restrict__ up = input + dim;
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
vec_op::FP32Vec16 w1_vec(M_SQRT1_2);
|
||||
vec_op::FP32Vec16 w2_vec(0.5);
|
||||
alignas(64) float temp[16];
|
||||
|
||||
DEFINE_FAST_EXP
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < dim; n += 16) {
|
||||
vec_op::FP32Vec16 gate_vec(gate + n);
|
||||
vec_op::FP32Vec16 up_vec(up + n);
|
||||
auto er_input_vec = gate_vec * w1_vec;
|
||||
|
||||
er_input_vec.save(temp);
|
||||
for (int32_t i = 0; i < 16; ++i) {
|
||||
temp[i] = std::erf(temp[i]);
|
||||
}
|
||||
vec_op::FP32Vec16 er_vec(temp);
|
||||
auto gelu = gate_vec * w2_vec * (one_vec + er_vec);
|
||||
auto gated_output_fp32 = up_vec * gelu;
|
||||
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
|
||||
gated_output.save(output + n);
|
||||
}
|
||||
gate += input_stride;
|
||||
up += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void gelu_tanh_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride,
|
||||
const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
const int32_t dim = n_size / 2;
|
||||
float* __restrict__ gate = input;
|
||||
float* __restrict__ up = input + dim;
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
vec_op::FP32Vec16 w1_vec(0.7978845608028654);
|
||||
vec_op::FP32Vec16 w2_vec(0.5);
|
||||
vec_op::FP32Vec16 w3_vec(0.044715);
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < dim; n += 16) {
|
||||
vec_op::FP32Vec16 gate_vec(gate + n);
|
||||
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);
|
||||
// Note: can't use fast_exp form because diffusiongemma will generate
|
||||
// wrong results
|
||||
auto tanh_vec = inner_vec.tanh();
|
||||
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);
|
||||
gated_output.save(output + n);
|
||||
}
|
||||
gate += input_stride;
|
||||
up += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
FORCE_INLINE void apply_gated_act(const FusedMOEAct act,
|
||||
float* __restrict__ input,
|
||||
scalar_t* __restrict__ output,
|
||||
const int32_t m, const int32_t n,
|
||||
const int32_t input_stride,
|
||||
const int32_t output_stride) {
|
||||
switch (act) {
|
||||
case FusedMOEAct::SwigluOAIAndMul:
|
||||
swigluoai_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
case FusedMOEAct::SiluAndMul:
|
||||
silu_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
case FusedMOEAct::GeluAndMul:
|
||||
gelu_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
case FusedMOEAct::GeluTanhAndMul:
|
||||
gelu_tanh_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
default:
|
||||
TORCH_CHECK(false, "Unsupported act type.");
|
||||
}
|
||||
}
|
||||
using cpu_fused_moe_utils::apply_gated_act;
|
||||
using cpu_fused_moe_utils::FusedMOEAct;
|
||||
|
||||
template <typename scalar_t, typename gemm_t>
|
||||
void prepack_moe_weight_impl(scalar_t* __restrict__ weight_ptr,
|
||||
@@ -817,6 +634,7 @@ void fused_moe_impl(scalar_t* __restrict__ output, scalar_t* __restrict__ input,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
void prepack_moe_weight(
|
||||
@@ -864,7 +682,7 @@ void cpu_fused_moe(
|
||||
const int32_t input_size_2 = w2.size(2);
|
||||
const int32_t output_size_2 = w2.size(1);
|
||||
const int32_t topk_num = topk_id.size(1);
|
||||
const FusedMOEAct act_type = get_act_type(act);
|
||||
const FusedMOEAct act_type = cpu_fused_moe_utils::get_act_type(act);
|
||||
cpu_utils::ISA isa_type = cpu_utils::get_isa(isa);
|
||||
TORCH_CHECK(!skip_weighted || topk_num == 1,
|
||||
"skip_weighted is only supported for topk=1 on CPU");
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
#ifndef CPU_FUSED_MOE_ACTIVATIONS_HPP
|
||||
#define CPU_FUSED_MOE_ACTIVATIONS_HPP
|
||||
|
||||
#include <cmath>
|
||||
#include <cstdint>
|
||||
#include <string>
|
||||
|
||||
#include "cpu/cpu_arch_macros.h"
|
||||
#include "cpu/utils.hpp"
|
||||
|
||||
namespace cpu_fused_moe_utils {
|
||||
enum class FusedMOEAct {
|
||||
SiluAndMul,
|
||||
SwigluOAIAndMul,
|
||||
GeluAndMul,
|
||||
GeluTanhAndMul,
|
||||
};
|
||||
|
||||
inline FusedMOEAct get_act_type(const std::string& act) {
|
||||
if (act == "silu") {
|
||||
return FusedMOEAct::SiluAndMul;
|
||||
} else if (act == "swigluoai") {
|
||||
return FusedMOEAct::SwigluOAIAndMul;
|
||||
} else if (act == "gelu") {
|
||||
return FusedMOEAct::GeluAndMul;
|
||||
} else if (act == "gelu_tanh") {
|
||||
return FusedMOEAct::GeluTanhAndMul;
|
||||
} else {
|
||||
TORCH_CHECK(false, "Invalid act type: " + act);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void swigluoai_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride,
|
||||
const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
#if !defined(__aarch64__)
|
||||
// For GPT-OSS interleaved gate-up weights
|
||||
alignas(64) static int32_t index[16] = {0, 2, 4, 6, 8, 10, 12, 14,
|
||||
16, 18, 20, 22, 24, 26, 28, 30};
|
||||
vec_op::INT32Vec16 index_vec(index);
|
||||
#endif
|
||||
vec_op::FP32Vec16 gate_up_max_vec(7.0);
|
||||
vec_op::FP32Vec16 up_min_vec(-7.0);
|
||||
vec_op::FP32Vec16 alpha_vec(1.702);
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
|
||||
DEFINE_FAST_EXP
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < n_size; n += 32) {
|
||||
// Note: AdvSIMD does not support gather loads
|
||||
#if defined(__aarch64__)
|
||||
vec_op::FP32Vec16 gate_vec(vec_op::uninit);
|
||||
vec_op::FP32Vec16 up_vec(vec_op::uninit);
|
||||
vec_op::FP32Vec16::load_even_odd(input + n, gate_vec, up_vec);
|
||||
#else
|
||||
vec_op::FP32Vec16 gate_vec(input + n, index_vec);
|
||||
vec_op::FP32Vec16 up_vec(input + n + 1, index_vec);
|
||||
#endif
|
||||
gate_vec = gate_vec.min(gate_up_max_vec);
|
||||
up_vec = up_vec.clamp(up_min_vec, gate_up_max_vec);
|
||||
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec * alpha_vec));
|
||||
auto glu = gate_vec * sigmoid_vec;
|
||||
auto gated_output_fp32 = (one_vec + up_vec) * glu;
|
||||
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
|
||||
gated_output.save(output + n / 2);
|
||||
}
|
||||
input += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void silu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride, const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
const int32_t dim = n_size / 2;
|
||||
float* __restrict__ gate = input;
|
||||
float* __restrict__ up = input + dim;
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
|
||||
DEFINE_FAST_EXP
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < dim; n += 16) {
|
||||
vec_op::FP32Vec16 gate_vec(gate + n);
|
||||
vec_op::FP32Vec16 up_vec(up + n);
|
||||
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec));
|
||||
auto silu = gate_vec * sigmoid_vec;
|
||||
auto gated_output_fp32 = up_vec * silu;
|
||||
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
|
||||
gated_output.save(output + n);
|
||||
}
|
||||
gate += input_stride;
|
||||
up += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void gelu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride, const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
const int32_t dim = n_size / 2;
|
||||
float* __restrict__ gate = input;
|
||||
float* __restrict__ up = input + dim;
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
vec_op::FP32Vec16 w1_vec(M_SQRT1_2);
|
||||
vec_op::FP32Vec16 w2_vec(0.5);
|
||||
alignas(64) float temp[16];
|
||||
|
||||
DEFINE_FAST_EXP
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < dim; n += 16) {
|
||||
vec_op::FP32Vec16 gate_vec(gate + n);
|
||||
vec_op::FP32Vec16 up_vec(up + n);
|
||||
auto er_input_vec = gate_vec * w1_vec;
|
||||
|
||||
er_input_vec.save(temp);
|
||||
for (int32_t i = 0; i < 16; ++i) {
|
||||
temp[i] = std::erf(temp[i]);
|
||||
}
|
||||
vec_op::FP32Vec16 er_vec(temp);
|
||||
auto gelu = gate_vec * w2_vec * (one_vec + er_vec);
|
||||
auto gated_output_fp32 = up_vec * gelu;
|
||||
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
|
||||
gated_output.save(output + n);
|
||||
}
|
||||
gate += input_stride;
|
||||
up += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void gelu_tanh_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
|
||||
const int32_t m_size, const int32_t n_size,
|
||||
const int32_t input_stride,
|
||||
const int32_t output_stride) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
const int32_t dim = n_size / 2;
|
||||
float* __restrict__ gate = input;
|
||||
float* __restrict__ up = input + dim;
|
||||
vec_op::FP32Vec16 one_vec(1.0);
|
||||
vec_op::FP32Vec16 w1_vec(0.7978845608028654);
|
||||
vec_op::FP32Vec16 w2_vec(0.5);
|
||||
vec_op::FP32Vec16 w3_vec(0.044715);
|
||||
|
||||
for (int32_t m = 0; m < m_size; ++m) {
|
||||
for (int32_t n = 0; n < dim; n += 16) {
|
||||
vec_op::FP32Vec16 gate_vec(gate + n);
|
||||
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);
|
||||
// Note: can't use fast_exp form because diffusiongemma will generate
|
||||
// wrong results
|
||||
auto tanh_vec = inner_vec.tanh();
|
||||
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);
|
||||
gated_output.save(output + n);
|
||||
}
|
||||
gate += input_stride;
|
||||
up += input_stride;
|
||||
output += output_stride;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
FORCE_INLINE void apply_gated_act(const FusedMOEAct act,
|
||||
float* __restrict__ input,
|
||||
scalar_t* __restrict__ output,
|
||||
const int32_t m, const int32_t n,
|
||||
const int32_t input_stride,
|
||||
const int32_t output_stride) {
|
||||
switch (act) {
|
||||
case FusedMOEAct::SwigluOAIAndMul:
|
||||
swigluoai_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
case FusedMOEAct::SiluAndMul:
|
||||
silu_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
case FusedMOEAct::GeluAndMul:
|
||||
gelu_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
case FusedMOEAct::GeluTanhAndMul:
|
||||
gelu_tanh_and_mul(input, output, m, n, input_stride, output_stride);
|
||||
return;
|
||||
default:
|
||||
TORCH_CHECK(false, "Unsupported act type.");
|
||||
}
|
||||
}
|
||||
} // namespace cpu_fused_moe_utils
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,647 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
#include "cpu/cpu_arch_macros.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
|
||||
#include "cpu/cpu_fused_moe_activations.hpp"
|
||||
#include "cpu/cpu_types.hpp"
|
||||
#include "cpu/micro_gemm/cpu_micro_gemm_impl.hpp"
|
||||
#include "cpu/utils.hpp"
|
||||
|
||||
#if defined(ARM_I8MM_SUPPORT) && defined(ARM_BF16_SUPPORT)
|
||||
#include "cpu/micro_gemm/cpu_micro_gemm_int8_neon.hpp"
|
||||
#define NEON_DISPATCH(SCALAR_TYPE, ...) \
|
||||
case cpu_utils::ISA::NEON: { \
|
||||
using gemm_t = \
|
||||
cpu_micro_gemm::MicroGemmINT8<cpu_utils::ISA::NEON, SCALAR_TYPE>; \
|
||||
return __VA_ARGS__(); \
|
||||
}
|
||||
#else
|
||||
#define NEON_DISPATCH(SCALAR_TYPE, ...) case cpu_utils::ISA::NEON:
|
||||
#endif
|
||||
|
||||
#define CPU_INT8_ISA_DISPATCH_IMPL(ISA_TYPE, SCALAR_TYPE, ...) \
|
||||
[&] { \
|
||||
switch (ISA_TYPE) { \
|
||||
NEON_DISPATCH(SCALAR_TYPE, __VA_ARGS__) \
|
||||
default: { \
|
||||
TORCH_CHECK(false, "Invalid CPU ISA type."); \
|
||||
} \
|
||||
} \
|
||||
}()
|
||||
|
||||
namespace {
|
||||
using cpu_fused_moe_utils::apply_gated_act;
|
||||
using cpu_fused_moe_utils::FusedMOEAct;
|
||||
|
||||
template <typename gemm_t>
|
||||
void prepack_moe_weight_int8_impl(const int8_t* __restrict__ weight_ptr,
|
||||
int8_t* __restrict__ packed_weight_ptr,
|
||||
const int32_t expert_num,
|
||||
const int32_t output_size,
|
||||
const int32_t input_size,
|
||||
const int64_t expert_stride) {
|
||||
#pragma omp parallel for
|
||||
for (int32_t e_idx = 0; e_idx < expert_num; ++e_idx) {
|
||||
gemm_t::pack_weight(weight_ptr + expert_stride * e_idx,
|
||||
packed_weight_ptr + expert_stride * e_idx, output_size,
|
||||
input_size);
|
||||
}
|
||||
}
|
||||
|
||||
// INT8 MoE kernel, based on the original BF16 kernel in cpu_fused_moe.cpp
|
||||
template <typename scalar_t, typename gemm_t>
|
||||
void fused_moe_int8_impl(
|
||||
scalar_t* __restrict__ output, const scalar_t* __restrict__ input,
|
||||
const int8_t* __restrict__ w13, const int8_t* __restrict__ w2,
|
||||
const float* __restrict__ w13_scales, const float* __restrict__ w2_scales,
|
||||
scalar_t* __restrict__ w13_bias, scalar_t* __restrict__ w2_bias,
|
||||
const float* __restrict__ topk_weights, const int32_t* __restrict__ topk_id,
|
||||
const FusedMOEAct act_type, const int32_t token_num,
|
||||
const int32_t expert_num, const int32_t topk_num,
|
||||
const int32_t input_size_13, const int32_t output_size_13,
|
||||
const int32_t input_size_2, const int32_t output_size_2,
|
||||
const bool skip_weighted) {
|
||||
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
constexpr int32_t gemm_n_tile_size = gemm_t::NSize;
|
||||
constexpr int32_t gemm_m_tile_size = gemm_t::MaxMSize;
|
||||
constexpr int32_t min_w13_n_tile_size = 2 * gemm_n_tile_size;
|
||||
|
||||
TORCH_CHECK_EQ(input_size_13 % gemm_t::K, 0);
|
||||
TORCH_CHECK_EQ(input_size_2 % gemm_t::K, 0);
|
||||
TORCH_CHECK_EQ(output_size_13 % min_w13_n_tile_size, 0);
|
||||
TORCH_CHECK_EQ(output_size_2 % gemm_n_tile_size, 0);
|
||||
TORCH_CHECK_EQ(output_size_13 / 2, input_size_2);
|
||||
|
||||
const int32_t thread_num = cpu_utils::get_max_threads();
|
||||
const int32_t w13_input_buffer_size = cpu_utils::round_up<64>(
|
||||
gemm_m_tile_size * input_size_13 * sizeof(int8_t));
|
||||
const int32_t w2_input_buffer_size =
|
||||
cpu_utils::round_up<64>(gemm_m_tile_size * input_size_2 * sizeof(int8_t));
|
||||
|
||||
const int32_t w13_n_tile_size = [&]() {
|
||||
const int64_t cache_size = cpu_utils::get_available_l2_size();
|
||||
const int32_t n_size_cache_limit =
|
||||
(cache_size - w13_input_buffer_size) /
|
||||
(gemm_m_tile_size * sizeof(float) + input_size_13 * sizeof(int8_t));
|
||||
const int32_t n_size_thread_limit =
|
||||
output_size_13 / std::max(1, thread_num / topk_num);
|
||||
const int32_t n_size = cpu_utils::round_down<min_w13_n_tile_size>(
|
||||
std::min(n_size_cache_limit, n_size_thread_limit));
|
||||
return std::max(n_size, min_w13_n_tile_size);
|
||||
}();
|
||||
|
||||
const int32_t w2_n_tile_size = [&]() {
|
||||
const int64_t cache_size = cpu_utils::get_available_l2_size();
|
||||
const int32_t n_size_cache_limit =
|
||||
(cache_size - w2_input_buffer_size) / (input_size_2 * sizeof(int8_t));
|
||||
const int32_t n_size_thread_limit =
|
||||
output_size_2 / std::max(1, thread_num / topk_num);
|
||||
const int32_t n_size = cpu_utils::round_down<gemm_n_tile_size>(
|
||||
std::min(n_size_cache_limit, n_size_thread_limit));
|
||||
return std::max(n_size, gemm_n_tile_size);
|
||||
}();
|
||||
|
||||
int32_t common_buffer_offset = 0;
|
||||
const int32_t token_num_per_group_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset += cpu_utils::round_up<64>(expert_num * sizeof(int32_t));
|
||||
const int32_t cu_token_num_per_group_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset +=
|
||||
cpu_utils::round_up<64>((expert_num + 1) * sizeof(int32_t));
|
||||
const int32_t expanded_token_num = token_num * topk_num;
|
||||
const int32_t expand_token_id_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset +=
|
||||
cpu_utils::round_up<64>(expanded_token_num * sizeof(int32_t));
|
||||
const int32_t expand_token_id_index_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset +=
|
||||
cpu_utils::round_up<64>(expanded_token_num * sizeof(int32_t));
|
||||
const int32_t input_quant_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset +=
|
||||
cpu_utils::round_up<64>(token_num * input_size_13 * sizeof(int8_t));
|
||||
const int32_t input_scale_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset += cpu_utils::round_up<64>(token_num * sizeof(float));
|
||||
const int32_t w13_gemm_output_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset += cpu_utils::round_up<64>(
|
||||
expanded_token_num * input_size_2 * sizeof(scalar_t));
|
||||
const int32_t w13_output_scale_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset +=
|
||||
cpu_utils::round_up<64>(expanded_token_num * sizeof(float));
|
||||
const int32_t w2_gemm_output_buffer_offset = common_buffer_offset;
|
||||
common_buffer_offset += cpu_utils::round_up<64>(
|
||||
expanded_token_num * output_size_2 * sizeof(float));
|
||||
|
||||
int32_t gemm_thread_buffer_offset = 0;
|
||||
const int32_t gemm_input_buffer_offset = gemm_thread_buffer_offset;
|
||||
gemm_thread_buffer_offset +=
|
||||
std::max(w13_input_buffer_size, w2_input_buffer_size);
|
||||
const int32_t gemm_output_buffer_offset = gemm_thread_buffer_offset;
|
||||
gemm_thread_buffer_offset += cpu_utils::round_up<64>(
|
||||
gemm_m_tile_size * std::max(w13_n_tile_size, w2_n_tile_size) *
|
||||
sizeof(int32_t));
|
||||
|
||||
const int32_t ws_output_buffer_offset = 0;
|
||||
const int32_t ws_thread_buffer_size =
|
||||
cpu_utils::round_up<64>(output_size_2 * sizeof(float));
|
||||
const int32_t thread_buffer_size =
|
||||
std::max(gemm_thread_buffer_offset, ws_thread_buffer_size);
|
||||
const int32_t buffer_size =
|
||||
common_buffer_offset + thread_buffer_size * thread_num;
|
||||
cpu_utils::ScratchPadManager::get_scratchpad_manager()->realloc(buffer_size);
|
||||
uint8_t* common_buffer_start =
|
||||
cpu_utils::ScratchPadManager::get_scratchpad_manager()
|
||||
->get_data<uint8_t>();
|
||||
uint8_t* thread_buffer_start = common_buffer_start + common_buffer_offset;
|
||||
|
||||
int32_t* __restrict__ token_num_per_group_buffer = reinterpret_cast<int32_t*>(
|
||||
common_buffer_start + token_num_per_group_buffer_offset);
|
||||
int32_t* __restrict__ cu_token_num_per_group_buffer =
|
||||
reinterpret_cast<int32_t*>(common_buffer_start +
|
||||
cu_token_num_per_group_buffer_offset);
|
||||
int32_t* __restrict__ expand_token_id_buffer = reinterpret_cast<int32_t*>(
|
||||
common_buffer_start + expand_token_id_buffer_offset);
|
||||
int32_t* __restrict__ expand_token_id_index_buffer =
|
||||
reinterpret_cast<int32_t*>(common_buffer_start +
|
||||
expand_token_id_index_buffer_offset);
|
||||
int8_t* __restrict__ input_quant_buffer = reinterpret_cast<int8_t*>(
|
||||
common_buffer_start + input_quant_buffer_offset);
|
||||
float* __restrict__ input_scale_buffer =
|
||||
reinterpret_cast<float*>(common_buffer_start + input_scale_buffer_offset);
|
||||
|
||||
std::memset(token_num_per_group_buffer, 0, expert_num * sizeof(int32_t));
|
||||
for (int32_t i = 0; i < expanded_token_num; ++i) {
|
||||
++token_num_per_group_buffer[topk_id[i]];
|
||||
}
|
||||
|
||||
int32_t token_num_sum = 0;
|
||||
cu_token_num_per_group_buffer[0] = 0;
|
||||
int32_t* token_index_buffer = cu_token_num_per_group_buffer + 1;
|
||||
for (int32_t i = 0; i < expert_num; ++i) {
|
||||
token_index_buffer[i] = token_num_sum;
|
||||
token_num_sum += token_num_per_group_buffer[i];
|
||||
}
|
||||
|
||||
for (int32_t i = 0; i < token_num; ++i) {
|
||||
const int32_t* curr_topk_id = topk_id + i * topk_num;
|
||||
int32_t* curr_index_buffer = expand_token_id_index_buffer + i * topk_num;
|
||||
for (int32_t j = 0; j < topk_num; ++j) {
|
||||
const int32_t curr_expert_id = curr_topk_id[j];
|
||||
const int32_t curr_index = token_index_buffer[curr_expert_id]++;
|
||||
expand_token_id_buffer[curr_index] = i;
|
||||
curr_index_buffer[j] = curr_index;
|
||||
}
|
||||
}
|
||||
|
||||
// quantize inputs
|
||||
#pragma omp parallel for
|
||||
for (int32_t token_idx = 0; token_idx < token_num; ++token_idx) {
|
||||
gemm_t::quantize_row(input + token_idx * input_size_13,
|
||||
input_quant_buffer + token_idx * input_size_13,
|
||||
input_scale_buffer[token_idx], input_size_13);
|
||||
}
|
||||
|
||||
{
|
||||
alignas(64) cpu_utils::Counter counter;
|
||||
cpu_utils::Counter* counter_ptr = &counter;
|
||||
|
||||
// w13 GEMM + act
|
||||
#pragma omp parallel for schedule(static, 1)
|
||||
for (int32_t thread_id = 0; thread_id < thread_num; ++thread_id) {
|
||||
const int32_t task_num_per_expert =
|
||||
(output_size_13 + w13_n_tile_size - 1) / w13_n_tile_size;
|
||||
const int32_t task_num = task_num_per_expert * expert_num;
|
||||
uint8_t* __restrict__ thread_buffer =
|
||||
thread_buffer_start + thread_id * thread_buffer_size;
|
||||
int8_t* __restrict__ gemm_input_buffer =
|
||||
reinterpret_cast<int8_t*>(thread_buffer + gemm_input_buffer_offset);
|
||||
float* __restrict__ gemm_output_buffer =
|
||||
reinterpret_cast<float*>(thread_buffer + gemm_output_buffer_offset);
|
||||
auto* __restrict__ w13_gemm_output_buffer = reinterpret_cast<scalar_t*>(
|
||||
common_buffer_start + w13_gemm_output_buffer_offset);
|
||||
gemm_t gemm;
|
||||
|
||||
const int32_t w13_n_group_stride =
|
||||
gemm_t::WeightOCGroupSize * input_size_13;
|
||||
const int32_t w13_n_tile_stride = gemm_n_tile_size * input_size_13;
|
||||
|
||||
for (;;) {
|
||||
const int32_t task_id = counter_ptr->acquire_counter();
|
||||
if (task_id >= task_num) {
|
||||
break;
|
||||
}
|
||||
const int32_t curr_expert_id = task_id / task_num_per_expert;
|
||||
const int32_t curr_output_group_id = task_id % task_num_per_expert;
|
||||
const int32_t curr_token_num =
|
||||
token_num_per_group_buffer[curr_expert_id];
|
||||
if (curr_token_num == 0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int32_t actual_n_tile_size =
|
||||
std::min(w13_n_tile_size,
|
||||
output_size_13 - curr_output_group_id * w13_n_tile_size);
|
||||
const int32_t* __restrict__ curr_expand_token_id_buffer =
|
||||
expand_token_id_buffer +
|
||||
cu_token_num_per_group_buffer[curr_expert_id];
|
||||
scalar_t* __restrict__ curr_w13_gemm_output_buffer =
|
||||
w13_gemm_output_buffer +
|
||||
cu_token_num_per_group_buffer[curr_expert_id] * input_size_2 +
|
||||
curr_output_group_id * w13_n_tile_size / 2;
|
||||
|
||||
const int8_t* w13_weight_ptr_0 = nullptr;
|
||||
const int8_t* w13_weight_ptr_1 = nullptr;
|
||||
const float* w13_scale_ptr_0 = nullptr;
|
||||
const float* w13_scale_ptr_1 = nullptr;
|
||||
scalar_t* w13_bias_ptr_0 = nullptr;
|
||||
scalar_t* w13_bias_ptr_1 = nullptr;
|
||||
if (act_type == FusedMOEAct::SwigluOAIAndMul) {
|
||||
const int32_t output_offset = curr_output_group_id * w13_n_tile_size;
|
||||
w13_weight_ptr_0 = w13 +
|
||||
curr_expert_id * input_size_13 * output_size_13 +
|
||||
output_offset * input_size_13;
|
||||
w13_weight_ptr_1 =
|
||||
w13_weight_ptr_0 + actual_n_tile_size / 2 * input_size_13;
|
||||
w13_scale_ptr_0 =
|
||||
w13_scales + curr_expert_id * output_size_13 + output_offset;
|
||||
w13_scale_ptr_1 = w13_scale_ptr_0 + actual_n_tile_size / 2;
|
||||
if (w13_bias != nullptr) {
|
||||
w13_bias_ptr_0 =
|
||||
w13_bias + curr_expert_id * output_size_13 + output_offset;
|
||||
w13_bias_ptr_1 = w13_bias_ptr_0 + actual_n_tile_size / 2;
|
||||
}
|
||||
} else {
|
||||
const int32_t output_offset =
|
||||
curr_output_group_id * (w13_n_tile_size / 2);
|
||||
w13_weight_ptr_0 = w13 +
|
||||
curr_expert_id * input_size_13 * output_size_13 +
|
||||
output_offset * input_size_13;
|
||||
w13_weight_ptr_1 =
|
||||
w13_weight_ptr_0 + output_size_13 / 2 * input_size_13;
|
||||
w13_scale_ptr_0 =
|
||||
w13_scales + curr_expert_id * output_size_13 + output_offset;
|
||||
w13_scale_ptr_1 = w13_scale_ptr_0 + output_size_13 / 2;
|
||||
if (w13_bias != nullptr) {
|
||||
w13_bias_ptr_0 =
|
||||
w13_bias + curr_expert_id * output_size_13 + output_offset;
|
||||
w13_bias_ptr_1 = w13_bias_ptr_0 + output_size_13 / 2;
|
||||
}
|
||||
}
|
||||
|
||||
for (int32_t token_idx = 0; token_idx < curr_token_num;
|
||||
token_idx += gemm_m_tile_size) {
|
||||
const int32_t actual_token_num =
|
||||
std::min(gemm_m_tile_size, curr_token_num - token_idx);
|
||||
const int8_t* input_rows[gemm_m_tile_size];
|
||||
alignas(64) float input_scales[gemm_m_tile_size];
|
||||
// gather and pack
|
||||
for (int32_t i = 0; i < actual_token_num; ++i) {
|
||||
const int32_t curr_token_id = curr_expand_token_id_buffer[i];
|
||||
input_rows[i] = input_quant_buffer + curr_token_id * input_size_13;
|
||||
input_scales[i] = input_scale_buffer[curr_token_id];
|
||||
}
|
||||
gemm_t::pack_input_from_rows(input_rows, gemm_input_buffer,
|
||||
actual_token_num, input_size_13);
|
||||
curr_expand_token_id_buffer += actual_token_num;
|
||||
|
||||
const int8_t* w13_weight_ptr_0_iter = w13_weight_ptr_0;
|
||||
const int8_t* w13_weight_ptr_1_iter = w13_weight_ptr_1;
|
||||
const float* w13_scale_ptr_0_iter = w13_scale_ptr_0;
|
||||
const float* w13_scale_ptr_1_iter = w13_scale_ptr_1;
|
||||
scalar_t* w13_bias_ptr_0_iter = w13_bias_ptr_0;
|
||||
scalar_t* w13_bias_ptr_1_iter = w13_bias_ptr_1;
|
||||
float* w13_output_buffer_0_iter = gemm_output_buffer;
|
||||
float* w13_output_buffer_1_iter =
|
||||
gemm_output_buffer + actual_n_tile_size / 2;
|
||||
|
||||
for (int32_t i = 0; i < actual_n_tile_size;
|
||||
i += min_w13_n_tile_size) {
|
||||
auto* output_0_int32 =
|
||||
reinterpret_cast<int32_t*>(w13_output_buffer_0_iter);
|
||||
gemm.gemm(gemm_input_buffer, w13_weight_ptr_0_iter, output_0_int32,
|
||||
actual_token_num, input_size_13, w13_n_group_stride,
|
||||
actual_n_tile_size);
|
||||
gemm_t::dequantize_tile(output_0_int32, w13_output_buffer_0_iter,
|
||||
input_scales, w13_scale_ptr_0_iter,
|
||||
actual_token_num, gemm_n_tile_size,
|
||||
actual_n_tile_size);
|
||||
if (w13_bias != nullptr) {
|
||||
cpu_micro_gemm::add_bias_epilogue<gemm_n_tile_size>(
|
||||
w13_output_buffer_0_iter, w13_output_buffer_0_iter,
|
||||
w13_bias_ptr_0_iter, actual_token_num, actual_n_tile_size,
|
||||
actual_n_tile_size);
|
||||
w13_bias_ptr_0_iter += gemm_n_tile_size;
|
||||
}
|
||||
|
||||
auto* output_1_int32 =
|
||||
reinterpret_cast<int32_t*>(w13_output_buffer_1_iter);
|
||||
gemm.gemm(gemm_input_buffer, w13_weight_ptr_1_iter, output_1_int32,
|
||||
actual_token_num, input_size_13, w13_n_group_stride,
|
||||
actual_n_tile_size);
|
||||
gemm_t::dequantize_tile(output_1_int32, w13_output_buffer_1_iter,
|
||||
input_scales, w13_scale_ptr_1_iter,
|
||||
actual_token_num, gemm_n_tile_size,
|
||||
actual_n_tile_size);
|
||||
if (w13_bias != nullptr) {
|
||||
cpu_micro_gemm::add_bias_epilogue<gemm_n_tile_size>(
|
||||
w13_output_buffer_1_iter, w13_output_buffer_1_iter,
|
||||
w13_bias_ptr_1_iter, actual_token_num, actual_n_tile_size,
|
||||
actual_n_tile_size);
|
||||
w13_bias_ptr_1_iter += gemm_n_tile_size;
|
||||
}
|
||||
|
||||
w13_weight_ptr_0_iter += w13_n_tile_stride;
|
||||
w13_weight_ptr_1_iter += w13_n_tile_stride;
|
||||
w13_scale_ptr_0_iter += gemm_n_tile_size;
|
||||
w13_scale_ptr_1_iter += gemm_n_tile_size;
|
||||
w13_output_buffer_0_iter += gemm_n_tile_size;
|
||||
w13_output_buffer_1_iter += gemm_n_tile_size;
|
||||
}
|
||||
|
||||
apply_gated_act(act_type, gemm_output_buffer,
|
||||
curr_w13_gemm_output_buffer, actual_token_num,
|
||||
actual_n_tile_size, actual_n_tile_size, input_size_2);
|
||||
curr_w13_gemm_output_buffer += gemm_m_tile_size * input_size_2;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto* __restrict__ w13_gemm_output_buffer = reinterpret_cast<scalar_t*>(
|
||||
common_buffer_start + w13_gemm_output_buffer_offset);
|
||||
float* __restrict__ w13_output_scale_buffer = reinterpret_cast<float*>(
|
||||
common_buffer_start + w13_output_scale_buffer_offset);
|
||||
|
||||
// quantize w2 inputs - in place
|
||||
#pragma omp parallel for
|
||||
for (int32_t token_idx = 0; token_idx < expanded_token_num; ++token_idx) {
|
||||
scalar_t* input_row = w13_gemm_output_buffer + token_idx * input_size_2;
|
||||
int8_t* output_row = reinterpret_cast<int8_t*>(input_row);
|
||||
gemm_t::quantize_row(input_row, output_row,
|
||||
w13_output_scale_buffer[token_idx], input_size_2);
|
||||
}
|
||||
|
||||
{
|
||||
alignas(64) cpu_utils::Counter counter;
|
||||
cpu_utils::Counter* counter_ptr = &counter;
|
||||
|
||||
// w2 gemm
|
||||
#pragma omp parallel for schedule(static, 1)
|
||||
for (int32_t thread_id = 0; thread_id < thread_num; ++thread_id) {
|
||||
const int32_t task_num_per_expert =
|
||||
(output_size_2 + w2_n_tile_size - 1) / w2_n_tile_size;
|
||||
const int32_t task_num = task_num_per_expert * expert_num;
|
||||
uint8_t* __restrict__ thread_buffer =
|
||||
thread_buffer_start + thread_id * thread_buffer_size;
|
||||
int8_t* __restrict__ gemm_input_buffer =
|
||||
reinterpret_cast<int8_t*>(thread_buffer + gemm_input_buffer_offset);
|
||||
float* __restrict__ gemm_output_buffer =
|
||||
reinterpret_cast<float*>(thread_buffer + gemm_output_buffer_offset);
|
||||
float* __restrict__ w2_gemm_output_buffer = reinterpret_cast<float*>(
|
||||
common_buffer_start + w2_gemm_output_buffer_offset);
|
||||
gemm_t gemm;
|
||||
|
||||
const int32_t w2_n_group_stride =
|
||||
gemm_t::WeightOCGroupSize * input_size_2;
|
||||
const int32_t w2_n_tile_stride = gemm_n_tile_size * input_size_2;
|
||||
|
||||
for (;;) {
|
||||
const int32_t task_id = counter_ptr->acquire_counter();
|
||||
if (task_id >= task_num) {
|
||||
break;
|
||||
}
|
||||
const int32_t curr_expert_id = task_id / task_num_per_expert;
|
||||
const int32_t curr_output_group_id = task_id % task_num_per_expert;
|
||||
const int32_t curr_token_num =
|
||||
token_num_per_group_buffer[curr_expert_id];
|
||||
if (curr_token_num == 0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int32_t actual_n_tile_size =
|
||||
std::min(w2_n_tile_size,
|
||||
output_size_2 - curr_output_group_id * w2_n_tile_size);
|
||||
scalar_t* __restrict__ curr_w13_gemm_output_buffer =
|
||||
w13_gemm_output_buffer +
|
||||
cu_token_num_per_group_buffer[curr_expert_id] * input_size_2;
|
||||
float* __restrict__ curr_w13_output_scale_buffer =
|
||||
w13_output_scale_buffer +
|
||||
cu_token_num_per_group_buffer[curr_expert_id];
|
||||
float* __restrict__ curr_w2_gemm_output_buffer =
|
||||
w2_gemm_output_buffer +
|
||||
cu_token_num_per_group_buffer[curr_expert_id] * output_size_2 +
|
||||
curr_output_group_id * w2_n_tile_size;
|
||||
const int8_t* __restrict__ w2_weight_ptr =
|
||||
w2 + curr_expert_id * output_size_2 * input_size_2 +
|
||||
curr_output_group_id * w2_n_tile_size * input_size_2;
|
||||
const float* __restrict__ w2_scale_ptr =
|
||||
w2_scales + curr_expert_id * output_size_2 +
|
||||
curr_output_group_id * w2_n_tile_size;
|
||||
scalar_t* w2_bias_ptr = nullptr;
|
||||
if (w2_bias != nullptr) {
|
||||
w2_bias_ptr = w2_bias + curr_expert_id * output_size_2 +
|
||||
curr_output_group_id * w2_n_tile_size;
|
||||
}
|
||||
|
||||
for (int32_t token_idx = 0; token_idx < curr_token_num;
|
||||
token_idx += gemm_m_tile_size) {
|
||||
const int32_t actual_token_num =
|
||||
std::min(gemm_m_tile_size, curr_token_num - token_idx);
|
||||
const int8_t* input_rows[gemm_m_tile_size];
|
||||
alignas(64) float input_scales[gemm_m_tile_size];
|
||||
for (int32_t i = 0; i < actual_token_num; ++i) {
|
||||
input_rows[i] = reinterpret_cast<const int8_t*>(
|
||||
curr_w13_gemm_output_buffer + i * input_size_2);
|
||||
input_scales[i] = curr_w13_output_scale_buffer[i];
|
||||
}
|
||||
gemm_t::pack_input_from_rows(input_rows, gemm_input_buffer,
|
||||
actual_token_num, input_size_2);
|
||||
|
||||
const int8_t* w2_weight_ptr_iter = w2_weight_ptr;
|
||||
const float* w2_scale_ptr_iter = w2_scale_ptr;
|
||||
scalar_t* w2_bias_ptr_iter = w2_bias_ptr;
|
||||
float* curr_w2_gemm_output_buffer_iter = curr_w2_gemm_output_buffer;
|
||||
for (int32_t i = 0; i < actual_n_tile_size; i += gemm_n_tile_size) {
|
||||
auto* output_int32 = reinterpret_cast<int32_t*>(gemm_output_buffer);
|
||||
gemm.gemm(gemm_input_buffer, w2_weight_ptr_iter, output_int32,
|
||||
actual_token_num, input_size_2, w2_n_group_stride,
|
||||
gemm_n_tile_size);
|
||||
gemm_t::dequantize_tile(output_int32, gemm_output_buffer,
|
||||
input_scales, w2_scale_ptr_iter,
|
||||
actual_token_num, gemm_n_tile_size,
|
||||
gemm_n_tile_size);
|
||||
if (w2_bias != nullptr) {
|
||||
cpu_micro_gemm::add_bias_epilogue<gemm_n_tile_size>(
|
||||
gemm_output_buffer, gemm_output_buffer, w2_bias_ptr_iter,
|
||||
actual_token_num, gemm_n_tile_size, gemm_n_tile_size);
|
||||
w2_bias_ptr_iter += gemm_n_tile_size;
|
||||
}
|
||||
for (int32_t m_idx = 0; m_idx < actual_token_num; ++m_idx) {
|
||||
std::memcpy(
|
||||
curr_w2_gemm_output_buffer_iter + m_idx * output_size_2,
|
||||
gemm_output_buffer + m_idx * gemm_n_tile_size,
|
||||
gemm_n_tile_size * sizeof(float));
|
||||
}
|
||||
|
||||
w2_weight_ptr_iter += w2_n_tile_stride;
|
||||
w2_scale_ptr_iter += gemm_n_tile_size;
|
||||
curr_w2_gemm_output_buffer_iter += gemm_n_tile_size;
|
||||
}
|
||||
|
||||
curr_w13_gemm_output_buffer += gemm_m_tile_size * input_size_2;
|
||||
curr_w13_output_scale_buffer += gemm_m_tile_size;
|
||||
curr_w2_gemm_output_buffer += gemm_m_tile_size * output_size_2;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
{
|
||||
alignas(64) cpu_utils::Counter counter;
|
||||
cpu_utils::Counter* counter_ptr = &counter;
|
||||
|
||||
#pragma omp parallel for schedule(static, 1)
|
||||
for (int32_t thread_id = 0; thread_id < thread_num; ++thread_id) {
|
||||
uint8_t* __restrict__ thread_buffer =
|
||||
thread_buffer_start + thread_id * thread_buffer_size;
|
||||
float* __restrict__ ws_output_buffer =
|
||||
reinterpret_cast<float*>(thread_buffer + ws_output_buffer_offset);
|
||||
float* __restrict__ w2_gemm_output_buffer = reinterpret_cast<float*>(
|
||||
common_buffer_start + w2_gemm_output_buffer_offset);
|
||||
|
||||
for (;;) {
|
||||
const int32_t token_id = counter_ptr->acquire_counter();
|
||||
if (token_id >= token_num) {
|
||||
break;
|
||||
}
|
||||
int32_t* __restrict__ curr_expand_token_id_index_buffer =
|
||||
expand_token_id_index_buffer + token_id * topk_num;
|
||||
const float* __restrict__ curr_weight =
|
||||
topk_weights + token_id * topk_num;
|
||||
const float first_weight = skip_weighted ? 1.0f : curr_weight[0];
|
||||
scalar_t* __restrict__ curr_output_buffer =
|
||||
output + token_id * output_size_2;
|
||||
|
||||
if (topk_num > 1) {
|
||||
int32_t w2_output_idx = curr_expand_token_id_index_buffer[0];
|
||||
float* w2_output_iter =
|
||||
w2_gemm_output_buffer + w2_output_idx * output_size_2;
|
||||
float* ws_output_buffer_iter = ws_output_buffer;
|
||||
vec_op::FP32Vec16 weight_vec(first_weight);
|
||||
for (int32_t i = 0; i < output_size_2; i += 16) {
|
||||
vec_op::FP32Vec16 vec(w2_output_iter);
|
||||
(vec * weight_vec).save(ws_output_buffer_iter);
|
||||
w2_output_iter += 16;
|
||||
ws_output_buffer_iter += 16;
|
||||
}
|
||||
|
||||
for (int32_t idx = 1; idx < topk_num - 1; ++idx) {
|
||||
w2_output_idx = curr_expand_token_id_index_buffer[idx];
|
||||
w2_output_iter =
|
||||
w2_gemm_output_buffer + w2_output_idx * output_size_2;
|
||||
ws_output_buffer_iter = ws_output_buffer;
|
||||
weight_vec = vec_op::FP32Vec16(curr_weight[idx]);
|
||||
for (int32_t i = 0; i < output_size_2; i += 16) {
|
||||
vec_op::FP32Vec16 vec(w2_output_iter);
|
||||
vec_op::FP32Vec16 sum(ws_output_buffer_iter);
|
||||
(sum + vec * weight_vec).save(ws_output_buffer_iter);
|
||||
w2_output_iter += 16;
|
||||
ws_output_buffer_iter += 16;
|
||||
}
|
||||
}
|
||||
|
||||
const int32_t last_idx = topk_num - 1;
|
||||
w2_output_idx = curr_expand_token_id_index_buffer[last_idx];
|
||||
w2_output_iter =
|
||||
w2_gemm_output_buffer + w2_output_idx * output_size_2;
|
||||
ws_output_buffer_iter = ws_output_buffer;
|
||||
scalar_t* curr_output_buffer_iter = curr_output_buffer;
|
||||
weight_vec = vec_op::FP32Vec16(curr_weight[last_idx]);
|
||||
for (int32_t i = 0; i < output_size_2; i += 16) {
|
||||
vec_op::FP32Vec16 vec(w2_output_iter);
|
||||
vec_op::FP32Vec16 sum(ws_output_buffer_iter);
|
||||
scalar_vec_t(sum + vec * weight_vec).save(curr_output_buffer_iter);
|
||||
w2_output_iter += 16;
|
||||
ws_output_buffer_iter += 16;
|
||||
curr_output_buffer_iter += 16;
|
||||
}
|
||||
} else {
|
||||
const int32_t w2_output_idx = curr_expand_token_id_index_buffer[0];
|
||||
float* w2_output_iter =
|
||||
w2_gemm_output_buffer + w2_output_idx * output_size_2;
|
||||
scalar_t* curr_output_buffer_iter = curr_output_buffer;
|
||||
vec_op::FP32Vec16 weight_vec(first_weight);
|
||||
for (int32_t i = 0; i < output_size_2; i += 16) {
|
||||
vec_op::FP32Vec16 vec(w2_output_iter);
|
||||
scalar_vec_t(vec * weight_vec).save(curr_output_buffer_iter);
|
||||
w2_output_iter += 16;
|
||||
curr_output_buffer_iter += 16;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} // namespace
|
||||
|
||||
void prepack_moe_weight_int8(
|
||||
const torch::Tensor& weight, // [expert_num, output_size, input_size]
|
||||
torch::Tensor& packed_weight, const std::string& isa) {
|
||||
TORCH_CHECK(weight.is_contiguous());
|
||||
const int32_t expert_num = weight.size(0);
|
||||
const int32_t output_size = weight.size(1);
|
||||
const int32_t input_size = weight.size(2);
|
||||
const int64_t expert_stride = weight.stride(0);
|
||||
const cpu_utils::ISA isa_type = cpu_utils::get_isa(isa);
|
||||
TORCH_CHECK_EQ(output_size % 32, 0);
|
||||
|
||||
CPU_INT8_ISA_DISPATCH_IMPL(isa_type, c10::BFloat16, [&]() {
|
||||
TORCH_CHECK_EQ(input_size % gemm_t::K, 0);
|
||||
prepack_moe_weight_int8_impl<gemm_t>(
|
||||
weight.data_ptr<int8_t>(), packed_weight.data_ptr<int8_t>(), expert_num,
|
||||
output_size, input_size, expert_stride);
|
||||
});
|
||||
}
|
||||
|
||||
void cpu_fused_moe_int8(torch::Tensor& output, const torch::Tensor& input,
|
||||
const torch::Tensor& w13, const torch::Tensor& w2,
|
||||
const torch::Tensor& w13_scale,
|
||||
const torch::Tensor& w2_scale,
|
||||
const std::optional<torch::Tensor>& w13_bias,
|
||||
const std::optional<torch::Tensor>& w2_bias,
|
||||
const torch::Tensor& topk_weights,
|
||||
const torch::Tensor& topk_id, const bool skip_weighted,
|
||||
const std::string& act, const std::string& isa) {
|
||||
const int32_t token_num = input.size(0);
|
||||
const int32_t input_size_13 = input.size(1);
|
||||
const int64_t input_stride = input.stride(0);
|
||||
TORCH_CHECK_EQ(input_stride, input_size_13);
|
||||
const int32_t expert_num = w13.size(0);
|
||||
const int32_t output_size_13 = w13.size(1);
|
||||
const int32_t input_size_2 = w2.size(2);
|
||||
const int32_t output_size_2 = w2.size(1);
|
||||
const int32_t topk_num = topk_id.size(1);
|
||||
const FusedMOEAct act_type = cpu_fused_moe_utils::get_act_type(act);
|
||||
const cpu_utils::ISA isa_type = cpu_utils::get_isa(isa);
|
||||
TORCH_CHECK(!skip_weighted || topk_num == 1,
|
||||
"skip_weighted is only supported for topk=1 on CPU");
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(
|
||||
input.scalar_type(), "cpu_fused_moe_int8", [&]() {
|
||||
CPU_INT8_ISA_DISPATCH_IMPL(isa_type, scalar_t, [&]() {
|
||||
fused_moe_int8_impl<scalar_t, gemm_t>(
|
||||
output.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(),
|
||||
w13.data_ptr<int8_t>(), w2.data_ptr<int8_t>(),
|
||||
w13_scale.data_ptr<float>(), w2_scale.data_ptr<float>(),
|
||||
w13_bias.has_value() ? w13_bias->data_ptr<scalar_t>() : nullptr,
|
||||
w2_bias.has_value() ? w2_bias->data_ptr<scalar_t>() : nullptr,
|
||||
topk_weights.data_ptr<float>(), topk_id.data_ptr<int32_t>(),
|
||||
act_type, token_num, expert_num, topk_num, input_size_13,
|
||||
output_size_13, input_size_2, output_size_2, skip_weighted);
|
||||
});
|
||||
});
|
||||
}
|
||||
@@ -287,7 +287,7 @@ struct FP32Vec4 : public Vec<FP32Vec4> {
|
||||
|
||||
explicit FP32Vec4(__vector float data) : reg(data) {}
|
||||
|
||||
explicit FP32Vec4(const FP32Vec4& data) : reg(data.reg) {}
|
||||
FP32Vec4(const FP32Vec4& data) : reg(data.reg) {}
|
||||
};
|
||||
|
||||
struct FP32Vec8 : public Vec<FP32Vec8> {
|
||||
@@ -316,7 +316,7 @@ struct FP32Vec8 : public Vec<FP32Vec8> {
|
||||
|
||||
explicit FP32Vec8(f32x4x2_t data) : reg(data) {}
|
||||
|
||||
explicit FP32Vec8(const FP32Vec8& data) {
|
||||
FP32Vec8(const FP32Vec8& data) {
|
||||
reg.val[0] = data.reg.val[0];
|
||||
reg.val[1] = data.reg.val[1];
|
||||
}
|
||||
@@ -593,7 +593,7 @@ struct FP32Vec16 : public Vec<FP32Vec16> {
|
||||
explicit FP32Vec16(bool, const float* ptr) : FP32Vec16(ptr) {}
|
||||
explicit FP32Vec16(f32x4x4_t data) : reg(data) {}
|
||||
|
||||
explicit FP32Vec16(const FP32Vec16& data) {
|
||||
FP32Vec16(const FP32Vec16& data) {
|
||||
reg.val[0] = data.reg.val[0];
|
||||
reg.val[1] = data.reg.val[1];
|
||||
reg.val[2] = data.reg.val[2];
|
||||
@@ -747,6 +747,15 @@ struct FP32Vec16 : public Vec<FP32Vec16> {
|
||||
vec_abs(reg.val[2]), vec_abs(reg.val[3])}));
|
||||
}
|
||||
|
||||
FP32Vec16 exp() const {
|
||||
FP32Vec8 lo(f32x4x2_t{reg.val[0], reg.val[1]});
|
||||
FP32Vec8 hi(f32x4x2_t{reg.val[2], reg.val[3]});
|
||||
auto lo_e = lo.exp();
|
||||
auto hi_e = hi.exp();
|
||||
return FP32Vec16(f32x4x4_t{lo_e.reg.val[0], lo_e.reg.val[1],
|
||||
hi_e.reg.val[0], hi_e.reg.val[1]});
|
||||
}
|
||||
|
||||
float reduce_max() {
|
||||
__vector float max01 = vec_max(reg.val[0], reg.val[1]);
|
||||
__vector float max23 = vec_max(reg.val[2], reg.val[3]);
|
||||
|
||||
@@ -31,6 +31,9 @@ class MicroGemm {
|
||||
}
|
||||
};
|
||||
|
||||
template <cpu_utils::ISA isa, typename scalar_t>
|
||||
class MicroGemmINT8;
|
||||
|
||||
template <int32_t n_size, typename scalar_t>
|
||||
FORCE_INLINE void default_epilogue(float* __restrict__ c_ptr,
|
||||
scalar_t* __restrict__ d_ptr,
|
||||
|
||||
@@ -0,0 +1,424 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
#ifndef CPU_MICRO_GEMM_INT8_NEON_HPP
|
||||
#define CPU_MICRO_GEMM_INT8_NEON_HPP
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstdint>
|
||||
|
||||
#include "cpu/micro_gemm/cpu_micro_gemm_impl.hpp"
|
||||
|
||||
#include <arm_bf16.h>
|
||||
#include <arm_neon.h>
|
||||
#include <c10/util/BFloat16.h>
|
||||
#include <c10/util/Exception.h>
|
||||
#include <c10/util/Half.h>
|
||||
|
||||
namespace cpu_micro_gemm {
|
||||
|
||||
namespace neon_smmla {
|
||||
|
||||
constexpr int32_t K = 8;
|
||||
constexpr int32_t Cols = 2;
|
||||
constexpr int32_t TileSize = K * Cols;
|
||||
|
||||
FORCE_INLINE float32x4x2_t load_as_f32(const float* input) {
|
||||
float32x4x2_t result;
|
||||
result.val[0] = vld1q_f32(input);
|
||||
result.val[1] = vld1q_f32(input + 4);
|
||||
return result;
|
||||
}
|
||||
|
||||
FORCE_INLINE float32x4x2_t load_as_f32(const c10::Half* input) {
|
||||
const auto input_vec = vld1q_f16(reinterpret_cast<const float16_t*>(input));
|
||||
float32x4x2_t result;
|
||||
result.val[0] = vcvt_f32_f16(vget_low_f16(input_vec));
|
||||
result.val[1] = vcvt_f32_f16(vget_high_f16(input_vec));
|
||||
return result;
|
||||
}
|
||||
|
||||
FORCE_INLINE float32x4x2_t load_as_f32(const c10::BFloat16* input) {
|
||||
const auto input_vec = vld1q_bf16(reinterpret_cast<const bfloat16_t*>(input));
|
||||
float32x4x2_t result;
|
||||
result.val[0] = vcvt_f32_bf16(vget_low_bf16(input_vec));
|
||||
result.val[1] = vcvt_f32_bf16(vget_high_bf16(input_vec));
|
||||
return result;
|
||||
}
|
||||
|
||||
FORCE_INLINE void store_acc_rowpair(const int32x4_t acc01,
|
||||
const int32x4_t acc23,
|
||||
const int32x4_t acc45,
|
||||
const int32x4_t acc67,
|
||||
int32_t* __restrict__ c_ptr,
|
||||
const int64_t ldc, const int32_t m_rows) {
|
||||
if (m_rows == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
vst1q_s32(c_ptr, vcombine_s32(vget_low_s32(acc01), vget_low_s32(acc23)));
|
||||
vst1q_s32(c_ptr + 4, vcombine_s32(vget_low_s32(acc45), vget_low_s32(acc67)));
|
||||
|
||||
if (m_rows == 2) {
|
||||
vst1q_s32(c_ptr + ldc,
|
||||
vcombine_s32(vget_high_s32(acc01), vget_high_s32(acc23)));
|
||||
vst1q_s32(c_ptr + ldc + 4,
|
||||
vcombine_s32(vget_high_s32(acc45), vget_high_s32(acc67)));
|
||||
}
|
||||
}
|
||||
|
||||
FORCE_INLINE void gemm_micro_smmla_8x8_packed_a(
|
||||
const int8_t* __restrict__ a_packed, const int8_t* __restrict__ b_packed,
|
||||
int32_t* __restrict__ c_ptr, const int32_t m, const int32_t k_size,
|
||||
const int64_t ldc) {
|
||||
const int32x4_t zero = vdupq_n_s32(0);
|
||||
int32x4_t acc0101 = zero, acc0123 = zero, acc0145 = zero, acc0167 = zero;
|
||||
int32x4_t acc2301 = zero, acc2323 = zero, acc2345 = zero, acc2367 = zero;
|
||||
int32x4_t acc4501 = zero, acc4523 = zero, acc4545 = zero, acc4567 = zero;
|
||||
int32x4_t acc6701 = zero, acc6723 = zero, acc6745 = zero, acc6767 = zero;
|
||||
|
||||
const int8_t* __restrict__ a_tile = a_packed;
|
||||
const int8_t* __restrict__ b_tile = b_packed;
|
||||
|
||||
#pragma GCC unroll 8
|
||||
for (int32_t k_idx = 0; k_idx < k_size; k_idx += K) {
|
||||
const int8x16_t a_tile01 = vld1q_s8(a_tile);
|
||||
const int8x16_t a_tile23 = vld1q_s8(a_tile + TileSize);
|
||||
const int8x16_t a_tile45 = vld1q_s8(a_tile + 2 * TileSize);
|
||||
const int8x16_t a_tile67 = vld1q_s8(a_tile + 3 * TileSize);
|
||||
|
||||
const int8x16_t b_tile01 = vld1q_s8(b_tile);
|
||||
const int8x16_t b_tile23 = vld1q_s8(b_tile + TileSize);
|
||||
const int8x16_t b_tile45 = vld1q_s8(b_tile + 2 * TileSize);
|
||||
const int8x16_t b_tile67 = vld1q_s8(b_tile + 3 * TileSize);
|
||||
|
||||
acc0101 = vmmlaq_s32(acc0101, a_tile01, b_tile01);
|
||||
acc2301 = vmmlaq_s32(acc2301, a_tile23, b_tile01);
|
||||
acc4501 = vmmlaq_s32(acc4501, a_tile45, b_tile01);
|
||||
acc6701 = vmmlaq_s32(acc6701, a_tile67, b_tile01);
|
||||
|
||||
acc0123 = vmmlaq_s32(acc0123, a_tile01, b_tile23);
|
||||
acc2323 = vmmlaq_s32(acc2323, a_tile23, b_tile23);
|
||||
acc4523 = vmmlaq_s32(acc4523, a_tile45, b_tile23);
|
||||
acc6723 = vmmlaq_s32(acc6723, a_tile67, b_tile23);
|
||||
|
||||
acc0145 = vmmlaq_s32(acc0145, a_tile01, b_tile45);
|
||||
acc2345 = vmmlaq_s32(acc2345, a_tile23, b_tile45);
|
||||
acc4545 = vmmlaq_s32(acc4545, a_tile45, b_tile45);
|
||||
acc6745 = vmmlaq_s32(acc6745, a_tile67, b_tile45);
|
||||
|
||||
acc0167 = vmmlaq_s32(acc0167, a_tile01, b_tile67);
|
||||
acc2367 = vmmlaq_s32(acc2367, a_tile23, b_tile67);
|
||||
acc4567 = vmmlaq_s32(acc4567, a_tile45, b_tile67);
|
||||
acc6767 = vmmlaq_s32(acc6767, a_tile67, b_tile67);
|
||||
|
||||
a_tile += 4 * TileSize;
|
||||
b_tile += 4 * TileSize;
|
||||
}
|
||||
|
||||
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc,
|
||||
std::min(2, m));
|
||||
store_acc_rowpair(acc2301, acc2323, acc2345, acc2367, c_ptr + 2 * ldc, ldc,
|
||||
std::min(2, std::max(0, m - 2)));
|
||||
store_acc_rowpair(acc4501, acc4523, acc4545, acc4567, c_ptr + 4 * ldc, ldc,
|
||||
std::min(2, std::max(0, m - 4)));
|
||||
store_acc_rowpair(acc6701, acc6723, acc6745, acc6767, c_ptr + 6 * ldc, ldc,
|
||||
std::min(2, std::max(0, m - 6)));
|
||||
}
|
||||
|
||||
FORCE_INLINE void gemm_micro_smmla_4x16_packed_a(
|
||||
const int8_t* __restrict__ a_packed, const int8_t* __restrict__ b_packed,
|
||||
int32_t* __restrict__ c_ptr, const int32_t m, const int32_t k_size,
|
||||
const int64_t b_n_group_stride, const int64_t ldc) {
|
||||
const int32_t m_rows_01 = std::min(2, m);
|
||||
const int32_t m_rows_23 = std::min(2, std::max(0, m - 2));
|
||||
const int32x4_t zero = vdupq_n_s32(0);
|
||||
|
||||
int32x4_t acc0101 = zero, acc0123 = zero, acc0145 = zero, acc0167 = zero;
|
||||
int32x4_t acc2301 = zero, acc2323 = zero, acc2345 = zero, acc2367 = zero;
|
||||
int32x4_t acc0189 = zero, acc011011 = zero, acc011213 = zero,
|
||||
acc011415 = zero;
|
||||
int32x4_t acc2389 = zero, acc231011 = zero, acc231213 = zero,
|
||||
acc231415 = zero;
|
||||
|
||||
const int8_t* __restrict__ a_tile = a_packed;
|
||||
// note: b packs 8 panels contiguously, so we need 2 b_tile ptrs
|
||||
// for the 4x16 microkernel
|
||||
const int8_t* __restrict__ b_tile0 = b_packed;
|
||||
const int8_t* __restrict__ b_tile1 = b_packed + b_n_group_stride;
|
||||
|
||||
#pragma GCC unroll 8
|
||||
for (int32_t k_idx = 0; k_idx < k_size; k_idx += K) {
|
||||
const int8x16_t a_tile01 = vld1q_s8(a_tile);
|
||||
const int8x16_t a_tile23 = vld1q_s8(a_tile + TileSize);
|
||||
const int8x16_t b_tile01 = vld1q_s8(b_tile0);
|
||||
const int8x16_t b_tile23 = vld1q_s8(b_tile0 + TileSize);
|
||||
const int8x16_t b_tile45 = vld1q_s8(b_tile0 + 2 * TileSize);
|
||||
const int8x16_t b_tile67 = vld1q_s8(b_tile0 + 3 * TileSize);
|
||||
const int8x16_t b_tile89 = vld1q_s8(b_tile1);
|
||||
const int8x16_t b_tile1011 = vld1q_s8(b_tile1 + TileSize);
|
||||
const int8x16_t b_tile1213 = vld1q_s8(b_tile1 + 2 * TileSize);
|
||||
const int8x16_t b_tile1415 = vld1q_s8(b_tile1 + 3 * TileSize);
|
||||
|
||||
acc0101 = vmmlaq_s32(acc0101, a_tile01, b_tile01);
|
||||
acc2301 = vmmlaq_s32(acc2301, a_tile23, b_tile01);
|
||||
acc0123 = vmmlaq_s32(acc0123, a_tile01, b_tile23);
|
||||
acc2323 = vmmlaq_s32(acc2323, a_tile23, b_tile23);
|
||||
|
||||
acc0145 = vmmlaq_s32(acc0145, a_tile01, b_tile45);
|
||||
acc2345 = vmmlaq_s32(acc2345, a_tile23, b_tile45);
|
||||
acc0167 = vmmlaq_s32(acc0167, a_tile01, b_tile67);
|
||||
acc2367 = vmmlaq_s32(acc2367, a_tile23, b_tile67);
|
||||
|
||||
acc0189 = vmmlaq_s32(acc0189, a_tile01, b_tile89);
|
||||
acc2389 = vmmlaq_s32(acc2389, a_tile23, b_tile89);
|
||||
acc011011 = vmmlaq_s32(acc011011, a_tile01, b_tile1011);
|
||||
acc231011 = vmmlaq_s32(acc231011, a_tile23, b_tile1011);
|
||||
|
||||
acc011213 = vmmlaq_s32(acc011213, a_tile01, b_tile1213);
|
||||
acc231213 = vmmlaq_s32(acc231213, a_tile23, b_tile1213);
|
||||
acc011415 = vmmlaq_s32(acc011415, a_tile01, b_tile1415);
|
||||
acc231415 = vmmlaq_s32(acc231415, a_tile23, b_tile1415);
|
||||
|
||||
a_tile += 2 * TileSize;
|
||||
b_tile0 += 4 * TileSize;
|
||||
b_tile1 += 4 * TileSize;
|
||||
}
|
||||
|
||||
// rows 0-1, columns 0-7
|
||||
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc, m_rows_01);
|
||||
// rows 0-1, columns 8-15
|
||||
store_acc_rowpair(acc0189, acc011011, acc011213, acc011415, c_ptr + 8, ldc,
|
||||
m_rows_01);
|
||||
// rows 2-3, columns 0-7
|
||||
store_acc_rowpair(acc2301, acc2323, acc2345, acc2367, c_ptr + 2 * ldc, ldc,
|
||||
m_rows_23);
|
||||
// rows 2-3, columns 8-15
|
||||
store_acc_rowpair(acc2389, acc231011, acc231213, acc231415,
|
||||
c_ptr + 2 * ldc + 8, ldc, m_rows_23);
|
||||
}
|
||||
|
||||
} // namespace neon_smmla
|
||||
|
||||
template <typename scalar_t>
|
||||
class MicroGemmINT8<cpu_utils::ISA::NEON, scalar_t> {
|
||||
public:
|
||||
static constexpr int32_t K = neon_smmla::K;
|
||||
static constexpr int32_t Mr = 8;
|
||||
static constexpr int32_t Nr = 8;
|
||||
static constexpr int32_t NrGemv = 16;
|
||||
static constexpr int32_t MaxMSize = 8;
|
||||
static constexpr int32_t NSize = 32;
|
||||
static constexpr int32_t WeightOCGroupSize = Nr;
|
||||
static_assert(MaxMSize % Mr == 0);
|
||||
|
||||
static FORCE_INLINE void quantize_row(const scalar_t* input, int8_t* output,
|
||||
float& scale, const int32_t size) {
|
||||
TORCH_CHECK_EQ(size % K, 0);
|
||||
float32x4_t max_vec = vdupq_n_f32(0.0f);
|
||||
|
||||
for (int32_t i = 0; i < size; i += K) {
|
||||
const float32x4x2_t input_vec = neon_smmla::load_as_f32(input + i);
|
||||
max_vec = vmaxq_f32(max_vec, vabsq_f32(input_vec.val[0]));
|
||||
max_vec = vmaxq_f32(max_vec, vabsq_f32(input_vec.val[1]));
|
||||
}
|
||||
|
||||
const float abs_max = std::max(vmaxvq_f32(max_vec), 1.0e-7f);
|
||||
scale = abs_max / 127.0f;
|
||||
const float32x4_t inv_scale_vec = vdupq_n_f32(127.0f / abs_max);
|
||||
|
||||
for (int32_t i = 0; i < size; i += K) {
|
||||
const float32x4x2_t input_vec = neon_smmla::load_as_f32(input + i);
|
||||
const int32x4_t output_low =
|
||||
vcvtnq_s32_f32(vmulq_f32(input_vec.val[0], inv_scale_vec));
|
||||
const int32x4_t output_high =
|
||||
vcvtnq_s32_f32(vmulq_f32(input_vec.val[1], inv_scale_vec));
|
||||
const int16x8_t output_s16 =
|
||||
vcombine_s16(vqmovn_s32(output_low), vqmovn_s32(output_high));
|
||||
vst1_s8(output + i, vqmovn_s16(output_s16));
|
||||
}
|
||||
}
|
||||
|
||||
// with current code, fusing this into the gemm micro kernel didn't move the
|
||||
// needle
|
||||
static FORCE_INLINE void dequantize_tile(
|
||||
int32_t* input, float* output, const float* __restrict__ input_scales,
|
||||
const float* __restrict__ weight_scales, const int32_t m, const int32_t n,
|
||||
const int32_t stride) {
|
||||
TORCH_CHECK_EQ(n % 4, 0);
|
||||
for (int32_t m_idx = 0; m_idx < m; ++m_idx) {
|
||||
const float32x4_t input_scale_vec = vdupq_n_f32(input_scales[m_idx]);
|
||||
for (int32_t n_idx = 0; n_idx < n; n_idx += 4) {
|
||||
const int32x4_t input_vec = vld1q_s32(input + m_idx * stride + n_idx);
|
||||
const float32x4_t weight_scale_vec = vld1q_f32(weight_scales + n_idx);
|
||||
const float32x4_t output_vec =
|
||||
vmulq_f32(vcvtq_f32_s32(input_vec),
|
||||
vmulq_f32(input_scale_vec, weight_scale_vec));
|
||||
vst1q_f32(output + m_idx * stride + n_idx, output_vec);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// physical layout [
|
||||
// M / (8 or 4); Mr is 8 or 4
|
||||
// K / 8; K for smmla is 8
|
||||
// 4, ; 4 row-pairs for each 8 rows
|
||||
// 2, ; row-pair is 2 rows
|
||||
// 4 ; 4 elements per row
|
||||
// ]
|
||||
static void pack_input_from_rows(const int8_t* const* __restrict__ rows,
|
||||
int8_t* __restrict__ a_packed,
|
||||
const int32_t m, const int32_t k) {
|
||||
TORCH_CHECK(m > 0 && m <= MaxMSize);
|
||||
TORCH_CHECK(k % K == 0);
|
||||
const int8x8_t zero = vdup_n_s8(0);
|
||||
|
||||
for (int32_t row_base = 0; row_base < m; row_base += Mr) {
|
||||
const int32_t panel_m = std::min(Mr, m - row_base);
|
||||
const int8_t* const* panel_rows = rows + row_base;
|
||||
int8_t* __restrict__ out = a_packed + row_base * k;
|
||||
|
||||
// fast path for full 8-row panels (fast path for 4-row panels didn't move
|
||||
// the needle)
|
||||
if (panel_m == Mr) {
|
||||
const int8_t* __restrict__ row0 = panel_rows[0];
|
||||
const int8_t* __restrict__ row1 = panel_rows[1];
|
||||
const int8_t* __restrict__ row2 = panel_rows[2];
|
||||
const int8_t* __restrict__ row3 = panel_rows[3];
|
||||
const int8_t* __restrict__ row4 = panel_rows[4];
|
||||
const int8_t* __restrict__ row5 = panel_rows[5];
|
||||
const int8_t* __restrict__ row6 = panel_rows[6];
|
||||
const int8_t* __restrict__ row7 = panel_rows[7];
|
||||
int32_t k_idx = 0;
|
||||
for (; k_idx + 2 * K <= k; k_idx += 2 * K) {
|
||||
int8_t* __restrict__ block0 = out;
|
||||
int8_t* __restrict__ block1 = out + 4 * neon_smmla::TileSize;
|
||||
|
||||
int8x16_t a0 = vld1q_s8(row0 + k_idx);
|
||||
int8x16_t a1 = vld1q_s8(row1 + k_idx);
|
||||
vst1q_s8(block0, vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
|
||||
vst1q_s8(block1, vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
|
||||
|
||||
a0 = vld1q_s8(row2 + k_idx);
|
||||
a1 = vld1q_s8(row3 + k_idx);
|
||||
vst1q_s8(block0 + neon_smmla::TileSize,
|
||||
vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
|
||||
vst1q_s8(block1 + neon_smmla::TileSize,
|
||||
vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
|
||||
|
||||
a0 = vld1q_s8(row4 + k_idx);
|
||||
a1 = vld1q_s8(row5 + k_idx);
|
||||
vst1q_s8(block0 + 2 * neon_smmla::TileSize,
|
||||
vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
|
||||
vst1q_s8(block1 + 2 * neon_smmla::TileSize,
|
||||
vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
|
||||
|
||||
a0 = vld1q_s8(row6 + k_idx);
|
||||
a1 = vld1q_s8(row7 + k_idx);
|
||||
vst1q_s8(block0 + 3 * neon_smmla::TileSize,
|
||||
vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
|
||||
vst1q_s8(block1 + 3 * neon_smmla::TileSize,
|
||||
vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
|
||||
|
||||
out += 8 * neon_smmla::TileSize;
|
||||
}
|
||||
|
||||
for (; k_idx < k; k_idx += K) {
|
||||
int8x8_t a0 = vld1_s8(row0 + k_idx);
|
||||
int8x8_t a1 = vld1_s8(row1 + k_idx);
|
||||
vst1q_s8(out, vcombine_s8(a0, a1));
|
||||
|
||||
a0 = vld1_s8(row2 + k_idx);
|
||||
a1 = vld1_s8(row3 + k_idx);
|
||||
vst1q_s8(out + neon_smmla::TileSize, vcombine_s8(a0, a1));
|
||||
|
||||
a0 = vld1_s8(row4 + k_idx);
|
||||
a1 = vld1_s8(row5 + k_idx);
|
||||
vst1q_s8(out + 2 * neon_smmla::TileSize, vcombine_s8(a0, a1));
|
||||
|
||||
a0 = vld1_s8(row6 + k_idx);
|
||||
a1 = vld1_s8(row7 + k_idx);
|
||||
vst1q_s8(out + 3 * neon_smmla::TileSize, vcombine_s8(a0, a1));
|
||||
|
||||
out += 4 * neon_smmla::TileSize;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
const int32_t row_pairs = (panel_m <= 4) ? 2 : Mr / 2;
|
||||
for (int32_t k_idx = 0; k_idx < k; k_idx += K) {
|
||||
for (int32_t pair_idx = 0; pair_idx < row_pairs; ++pair_idx) {
|
||||
const int32_t row_idx = pair_idx * 2;
|
||||
const int8x8_t row0 =
|
||||
(row_idx < panel_m) ? vld1_s8(panel_rows[row_idx] + k_idx) : zero;
|
||||
const int8x8_t row1 = (row_idx + 1 < panel_m)
|
||||
? vld1_s8(panel_rows[row_idx + 1] + k_idx)
|
||||
: zero;
|
||||
vst1q_s8(out, vcombine_s8(row0, row1));
|
||||
out += neon_smmla::TileSize;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// physical layout [
|
||||
// N / 8; Nr is 8
|
||||
// K / 8; K for smmla is 8
|
||||
// 4, ; 4 col-pairs for each 8 cols
|
||||
// 2, ; col-pair is 2 cols
|
||||
// 4 ; 4 elements per col
|
||||
// ]
|
||||
static void pack_weight(const int8_t* __restrict__ weight,
|
||||
int8_t* __restrict__ packed_weight,
|
||||
const int32_t output_size, const int32_t input_size) {
|
||||
TORCH_CHECK(output_size % NSize == 0);
|
||||
TORCH_CHECK(input_size % K == 0);
|
||||
|
||||
for (int32_t o_idx = 0; o_idx < output_size; o_idx += Nr) {
|
||||
int8_t* __restrict__ dst = packed_weight + o_idx * input_size;
|
||||
for (int32_t k_idx = 0; k_idx < input_size; k_idx += K) {
|
||||
for (int32_t pair_idx = 0; pair_idx < Nr;
|
||||
pair_idx += neon_smmla::Cols) {
|
||||
const int8_t* __restrict__ row0 =
|
||||
weight + (o_idx + pair_idx) * input_size + k_idx;
|
||||
const int8_t* __restrict__ row1 = row0 + input_size;
|
||||
vst1q_s8(dst, vcombine_s8(vld1_s8(row0), vld1_s8(row1)));
|
||||
dst += neon_smmla::TileSize;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void gemm(const int8_t* __restrict__ a_packed,
|
||||
const int8_t* __restrict__ b_packed, int32_t* __restrict__ c,
|
||||
const int32_t m, const int32_t k, const int64_t b_n_group_stride,
|
||||
const int64_t ldc) const {
|
||||
TORCH_CHECK(m > 0 && m <= MaxMSize);
|
||||
TORCH_CHECK(k % K == 0);
|
||||
|
||||
for (int32_t n_idx = 0; n_idx < NSize; n_idx += NrGemv) {
|
||||
const int8_t* __restrict__ b_panel = b_packed + n_idx * k;
|
||||
|
||||
for (int32_t row_base = 0; row_base < m; row_base += Mr) {
|
||||
const int32_t panel_m = std::min(Mr, m - row_base);
|
||||
const int8_t* __restrict__ a_panel = a_packed + row_base * k;
|
||||
int32_t* __restrict__ c_panel = c + row_base * ldc + n_idx;
|
||||
|
||||
if (panel_m <= 4) {
|
||||
neon_smmla::gemm_micro_smmla_4x16_packed_a(
|
||||
a_panel, b_panel, c_panel, panel_m, k, b_n_group_stride, ldc);
|
||||
} else {
|
||||
neon_smmla::gemm_micro_smmla_8x8_packed_a(a_panel, b_panel, c_panel,
|
||||
panel_m, k, ldc);
|
||||
neon_smmla::gemm_micro_smmla_8x8_packed_a(
|
||||
a_panel, b_panel + b_n_group_stride, c_panel + Nr, panel_m, k,
|
||||
ldc);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cpu_micro_gemm
|
||||
|
||||
#endif
|
||||
@@ -1,3 +1,6 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
#ifndef CPU_MICRO_GEMM_NEON_HPP
|
||||
#define CPU_MICRO_GEMM_NEON_HPP
|
||||
|
||||
@@ -16,9 +19,6 @@ namespace {
|
||||
constexpr int32_t K = 4;
|
||||
constexpr int32_t Cols = 2;
|
||||
constexpr int32_t TileSize = K * Cols;
|
||||
constexpr int32_t Mr = 8;
|
||||
constexpr int32_t Nr = 8;
|
||||
constexpr int32_t Nr_gemv = 16;
|
||||
|
||||
// a = [a0, a1, a2, a3], b = [b0, b1, b2, b3] -> [a0, a1, b0, b1]
|
||||
FORCE_INLINE float32x4_t zip1_f32x4(const float32x4_t a, const float32x4_t b) {
|
||||
@@ -132,7 +132,7 @@ FORCE_INLINE void gemm_micro_bfmmla_8x8_packed_a(
|
||||
acc6767 = vbfmmlaq_f32(acc6767, a_tile67, b_tile67);
|
||||
|
||||
a_tile += 4 * TileSize;
|
||||
b_tile += Nr * K;
|
||||
b_tile += 4 * TileSize;
|
||||
}
|
||||
|
||||
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc,
|
||||
@@ -205,8 +205,8 @@ FORCE_INLINE void gemm_micro_bfmmla_4x16_packed_a(
|
||||
acc231415 = vbfmmlaq_f32(acc231415, a_tile23, b_tile1415);
|
||||
|
||||
a_tile += 2 * TileSize;
|
||||
b_tile0 += Nr * K;
|
||||
b_tile1 += Nr * K;
|
||||
b_tile0 += 4 * TileSize;
|
||||
b_tile1 += 4 * TileSize;
|
||||
}
|
||||
|
||||
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc, m_rows_01);
|
||||
@@ -223,6 +223,9 @@ FORCE_INLINE void gemm_micro_bfmmla_4x16_packed_a(
|
||||
template <typename scalar_t>
|
||||
class MicroGemm<cpu_utils::ISA::NEON, scalar_t> {
|
||||
public:
|
||||
static constexpr int32_t Mr = 8;
|
||||
static constexpr int32_t Nr = 8;
|
||||
static constexpr int32_t NrGemv = 16;
|
||||
static constexpr int32_t MaxMSize = 8;
|
||||
static constexpr int32_t NSize = 32;
|
||||
static constexpr int32_t WeightOCGroupSize = Nr;
|
||||
@@ -246,6 +249,9 @@ class MicroGemm<cpu_utils::ISA::NEON, c10::BFloat16> {
|
||||
public:
|
||||
using scalar_t = c10::BFloat16;
|
||||
|
||||
static constexpr int32_t Mr = 8;
|
||||
static constexpr int32_t Nr = 8;
|
||||
static constexpr int32_t NrGemv = 16;
|
||||
static constexpr int32_t MaxMSize = 8;
|
||||
static constexpr int32_t NSize = 32;
|
||||
static constexpr int32_t WeightOCGroupSize = Nr;
|
||||
@@ -253,7 +259,7 @@ class MicroGemm<cpu_utils::ISA::NEON, c10::BFloat16> {
|
||||
|
||||
public:
|
||||
// physical layout [
|
||||
// M / 8; Mr is 8
|
||||
// M / (8 or 4); Mr is 8 or 4
|
||||
// K / 4; K for bfmmla is 4
|
||||
// 4, ; 4 row-pairs for each 8 rows
|
||||
// 2, ; row-pair is 2 rows
|
||||
@@ -439,7 +445,7 @@ class MicroGemm<cpu_utils::ISA::NEON, c10::BFloat16> {
|
||||
(void)lda; // A is packed, so lda is not needed
|
||||
TORCH_CHECK_EQ(k % K, 0);
|
||||
|
||||
for (int32_t n_idx = 0; n_idx < NSize; n_idx += Nr_gemv) {
|
||||
for (int32_t n_idx = 0; n_idx < NSize; n_idx += NrGemv) {
|
||||
const bfloat16_t* __restrict__ b_panel =
|
||||
reinterpret_cast<const bfloat16_t*>(b_ptr) + n_idx * k;
|
||||
|
||||
|
||||
+173
-12
@@ -451,6 +451,90 @@ void causal_conv1d_update_kernel_impl(
|
||||
});
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void causal_conv1d_update_multi_kernel_impl(
|
||||
scalar_t* __restrict__ out,
|
||||
const scalar_t* __restrict__ input,
|
||||
scalar_t* __restrict__ conv_states,
|
||||
const scalar_t* __restrict__ weight,
|
||||
const scalar_t* __restrict__ bias,
|
||||
const int32_t* __restrict__ num_accepted_tokens,
|
||||
const int32_t* __restrict__ conv_indices,
|
||||
bool silu_activation,
|
||||
int64_t batch,
|
||||
int64_t dim,
|
||||
int64_t seqlen,
|
||||
int64_t width,
|
||||
int64_t state_len,
|
||||
int64_t conv_state_slot_stride) {
|
||||
constexpr int64_t BLOCK_N = block_size_n() * 2;
|
||||
const int64_t NB = div_up(dim, BLOCK_N);
|
||||
|
||||
AT_DISPATCH_BOOL2(bias != nullptr, has_bias, silu_activation, has_silu, [&] {
|
||||
at::parallel_for(0, batch * NB, 0, [&](int64_t begin, int64_t end) {
|
||||
int64_t bs{0}, nb{0};
|
||||
data_index_init(begin, bs, batch, nb, NB);
|
||||
|
||||
for (int64_t i = begin; i < end; ++i) {
|
||||
const int64_t nb_start = nb * BLOCK_N;
|
||||
const int64_t nb_size = std::min(dim - nb_start, BLOCK_N);
|
||||
const int32_t conv_state_index = conv_indices[bs];
|
||||
const int32_t history_offset = num_accepted_tokens[bs] - 1;
|
||||
|
||||
switch (width << 4 | nb_size >> 4) {
|
||||
case 0x42:
|
||||
tinygemm_kernel<scalar_t, 4, 32, has_bias, has_silu>::apply(
|
||||
input + bs * seqlen * dim + nb_start,
|
||||
weight + nb_start * width,
|
||||
out + bs * seqlen * dim + nb_start,
|
||||
has_bias ? bias + nb_start : nullptr,
|
||||
conv_states + conv_state_index * conv_state_slot_stride +
|
||||
history_offset * dim + nb_start,
|
||||
true,
|
||||
seqlen,
|
||||
dim,
|
||||
true);
|
||||
break;
|
||||
case 0x44:
|
||||
tinygemm_kernel<scalar_t, 4, 64, has_bias, has_silu>::apply(
|
||||
input + bs * seqlen * dim + nb_start,
|
||||
weight + nb_start * width,
|
||||
out + bs * seqlen * dim + nb_start,
|
||||
has_bias ? bias + nb_start : nullptr,
|
||||
conv_states + conv_state_index * conv_state_slot_stride +
|
||||
history_offset * dim + nb_start,
|
||||
true,
|
||||
seqlen,
|
||||
dim,
|
||||
true);
|
||||
break;
|
||||
default:
|
||||
TORCH_CHECK(false, "Unexpected block size, ", width, " x ", nb_size);
|
||||
}
|
||||
|
||||
data_index_step(bs, batch, nb, NB);
|
||||
}
|
||||
});
|
||||
});
|
||||
|
||||
at::parallel_for(0, batch, 0, [&](int64_t begin, int64_t end) {
|
||||
for (int64_t bs = begin; bs < end; ++bs) {
|
||||
const int32_t conv_state_index = conv_indices[bs];
|
||||
const int32_t num_accepted = num_accepted_tokens[bs];
|
||||
scalar_t* state = conv_states + conv_state_index * conv_state_slot_stride;
|
||||
|
||||
std::memmove(
|
||||
state,
|
||||
state + num_accepted * dim,
|
||||
(state_len - seqlen) * dim * sizeof(scalar_t));
|
||||
std::memcpy(
|
||||
state + (state_len - seqlen) * dim,
|
||||
input + bs * seqlen * dim,
|
||||
seqlen * dim * sizeof(scalar_t));
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
} // anonymous namespace
|
||||
|
||||
// from [dim, width] or [N, K]
|
||||
@@ -545,7 +629,7 @@ at::Tensor get_block_indices(const std::optional<at::Tensor>& offsets, int64_t n
|
||||
// query_start_loc: (batch + 1) int32
|
||||
// cache_indices: (batch) int32
|
||||
// has_initial_state: (batch) bool
|
||||
// conv_states: (..., dim, width - 1) itype
|
||||
// conv_states: (..., dim, state_len) itype, where state_len >= width - 1
|
||||
// activation: either None or "silu" or "swish"
|
||||
// pad_slot_id: int
|
||||
//
|
||||
@@ -586,11 +670,14 @@ at::Tensor causal_conv1d_fwd_cpu(
|
||||
CHECK_EQ(conv_states_val.scalar_type(), scalar_type);
|
||||
CHECK_GE(padded_batch, batch);
|
||||
CHECK_EQ(conv_states_val.size(1), dim);
|
||||
CHECK_EQ(conv_states_val.size(2), width - 1);
|
||||
const int64_t state_len = conv_states_val.size(2);
|
||||
CHECK_GE(state_len, width - 1);
|
||||
|
||||
// adjust `conv_states` to be contiguous on `dim`
|
||||
// should happen only once
|
||||
if (conv_states_val.stride(-2) != 1) {
|
||||
TORCH_CHECK(state_len == width - 1,
|
||||
"causal_conv1d_fwd_cpu: wide conv_states must be contiguous on dim.");
|
||||
auto conv_states_copy = conv_states_val.clone();
|
||||
conv_states_val.as_strided_({padded_batch, dim, width - 1}, {(width - 1) * dim, 1, dim});
|
||||
conv_states_val.copy_(conv_states_copy);
|
||||
@@ -651,14 +738,14 @@ at::Tensor causal_conv1d_fwd_cpu(
|
||||
|
||||
// API aligned with GPUs
|
||||
//
|
||||
// x: (batch, dim) or (batch, dim, seqlen)
|
||||
// x: (batch, dim) or (batch, seqlen, dim)
|
||||
// conv_state: (..., dim, state_len), where state_len >= width - 1
|
||||
// weight: (dim, width)
|
||||
// bias: (dim,)
|
||||
// cache_seqlens: (batch,), dtype int32.
|
||||
// num_accepted_tokens: (batch,), dtype int32.
|
||||
// conv_state_indices: (batch,), dtype int32
|
||||
// pad_slot_id: int
|
||||
// out: (batch, dim) or (batch, dim, seqlen)
|
||||
// out: (batch, dim) or (batch, seqlen, dim)
|
||||
//
|
||||
at::Tensor causal_conv1d_update_cpu(
|
||||
const at::Tensor& x,
|
||||
@@ -666,7 +753,7 @@ at::Tensor causal_conv1d_update_cpu(
|
||||
const at::Tensor& weight,
|
||||
const std::optional<at::Tensor>& bias,
|
||||
bool silu_activation,
|
||||
const std::optional<at::Tensor>& cache_seqlens,
|
||||
const std::optional<at::Tensor>& num_accepted_tokens,
|
||||
const std::optional<at::Tensor>& conv_state_indices,
|
||||
int64_t pad_slot_id,
|
||||
bool is_vnni) {
|
||||
@@ -674,13 +761,13 @@ at::Tensor causal_conv1d_update_cpu(
|
||||
CHECK_CONTIGUOUS(weight);
|
||||
auto packed_w = is_vnni ? weight : causal_conv1d_weight_pack(weight);
|
||||
|
||||
// TODO: add multi-token prediction support
|
||||
TORCH_CHECK(x.dim() == 2, "causal_conv1d_update_cpu: expect x to be 2D tensor.");
|
||||
TORCH_CHECK(!cache_seqlens.has_value(), "causal_conv1d_update_cpu: don't support cache_seqlens.");
|
||||
TORCH_CHECK(
|
||||
x.dim() == 2 || x.dim() == 3,
|
||||
"causal_conv1d_update_cpu: expect x to be 2D or 3D tensor.");
|
||||
|
||||
int64_t batch = x.size(0);
|
||||
int64_t dim = x.size(1);
|
||||
int64_t seqlen = 1;
|
||||
int64_t dim = x.dim() == 2 ? x.size(1) : x.size(2);
|
||||
int64_t seqlen = x.dim() == 2 ? 1 : x.size(1);
|
||||
int64_t width = weight.size(-1);
|
||||
|
||||
const auto scalar_type = x.scalar_type();
|
||||
@@ -690,10 +777,84 @@ at::Tensor causal_conv1d_update_cpu(
|
||||
|
||||
CHECK_EQ(conv_states.scalar_type(), scalar_type);
|
||||
CHECK_EQ(conv_states.size(1), dim);
|
||||
CHECK_EQ(conv_states.size(2), width - 1);
|
||||
const int64_t state_len = conv_states.size(2);
|
||||
CHECK_GE(state_len, width - 1);
|
||||
|
||||
if (x.dim() == 3) {
|
||||
TORCH_CHECK(
|
||||
num_accepted_tokens.has_value(),
|
||||
"causal_conv1d_update_cpu: num_accepted_tokens is required for 3D x.");
|
||||
TORCH_CHECK(
|
||||
conv_state_indices.has_value(),
|
||||
"causal_conv1d_update_cpu: conv_state_indices is required for 3D x.");
|
||||
CHECK_OPTIONAL_SHAPE_DTYPE(num_accepted_tokens, batch, at::kInt);
|
||||
TORCH_CHECK(
|
||||
width == 4,
|
||||
"causal_conv1d_update_cpu: support only width of 4 for 3D x.");
|
||||
TORCH_CHECK(
|
||||
seqlen > 0,
|
||||
"causal_conv1d_update_cpu: expect non-empty sequence for 3D x.");
|
||||
TORCH_CHECK(
|
||||
state_len >= seqlen,
|
||||
"causal_conv1d_update_cpu: state_len must be >= seqlen for 3D x.");
|
||||
TORCH_CHECK(
|
||||
conv_states.stride(-2) == 1 && conv_states.stride(-1) == dim,
|
||||
"causal_conv1d_update_cpu: 3D x requires SD conv_states layout.");
|
||||
|
||||
const int32_t* accepted_counts =
|
||||
num_accepted_tokens.value().data_ptr<int32_t>();
|
||||
const int32_t* indices = conv_state_indices.value().data_ptr<int32_t>();
|
||||
const int64_t num_slots = conv_states.size(0);
|
||||
for (int64_t bs = 0; bs < batch; ++bs) {
|
||||
const int32_t num_accepted = accepted_counts[bs];
|
||||
const int32_t conv_state_index = indices[bs];
|
||||
TORCH_CHECK(
|
||||
conv_state_index != pad_slot_id,
|
||||
"causal_conv1d_update_cpu: 3D x does not support pad slots.");
|
||||
TORCH_CHECK(
|
||||
conv_state_index >= 0 && conv_state_index < num_slots,
|
||||
"causal_conv1d_update_cpu: conv_state_indices out of range.");
|
||||
TORCH_CHECK(
|
||||
num_accepted >= 1 && num_accepted <= seqlen,
|
||||
"causal_conv1d_update_cpu: num_accepted_tokens must be in [1, "
|
||||
"seqlen].");
|
||||
TORCH_CHECK(
|
||||
num_accepted - 1 + width - 1 <= state_len,
|
||||
"causal_conv1d_update_cpu: history window exceeds conv_states.");
|
||||
}
|
||||
|
||||
int64_t conv_state_slot_stride = conv_states.stride(0);
|
||||
at::Tensor out = at::empty_like(x);
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(
|
||||
scalar_type, "causal_conv1d_update_multi_kernel_impl", [&] {
|
||||
causal_conv1d_update_multi_kernel_impl<scalar_t>(
|
||||
out.data_ptr<scalar_t>(),
|
||||
x.data_ptr<scalar_t>(),
|
||||
conv_states.data_ptr<scalar_t>(),
|
||||
packed_w.data_ptr<scalar_t>(),
|
||||
conditional_data_ptr<scalar_t>(bias),
|
||||
accepted_counts,
|
||||
indices,
|
||||
silu_activation,
|
||||
batch,
|
||||
dim,
|
||||
seqlen,
|
||||
width,
|
||||
state_len,
|
||||
conv_state_slot_stride);
|
||||
});
|
||||
return out;
|
||||
}
|
||||
|
||||
TORCH_CHECK(
|
||||
!num_accepted_tokens.has_value(),
|
||||
"causal_conv1d_update_cpu: num_accepted_tokens is only supported for 3D "
|
||||
"x.");
|
||||
|
||||
// adjust `conv_states` to be contiguous on `dim`
|
||||
if (conv_states.stride(-2) != 1) {
|
||||
TORCH_CHECK(state_len == width - 1,
|
||||
"causal_conv1d_update_cpu: wide conv_states must be contiguous on dim.");
|
||||
int64_t num_cache_lines = conv_states.size(0);
|
||||
auto conv_states_copy = conv_states.clone();
|
||||
conv_states.as_strided_({num_cache_lines, dim, width - 1}, {(width - 1) * dim, 1, dim});
|
||||
|
||||
@@ -147,7 +147,7 @@ at::Tensor causal_conv1d_fwd_cpu(
|
||||
at::Tensor causal_conv1d_update_cpu(
|
||||
const at::Tensor& x, const at::Tensor& conv_states,
|
||||
const at::Tensor& weight, const std::optional<at::Tensor>& bias,
|
||||
bool silu_activation, const std::optional<at::Tensor>& cache_seqlens,
|
||||
bool silu_activation, const std::optional<at::Tensor>& num_accepted_tokens,
|
||||
const std::optional<at::Tensor>& conv_state_indices, int64_t pad_slot_id,
|
||||
bool is_vnni);
|
||||
|
||||
@@ -207,6 +207,20 @@ void cpu_fused_moe(torch::Tensor& output, const torch::Tensor& input,
|
||||
const torch::Tensor& topk_id, const bool skip_weighted,
|
||||
const std::string& act, const std::string& isa);
|
||||
|
||||
void prepack_moe_weight_int8(const torch::Tensor& weight,
|
||||
torch::Tensor& packed_weight,
|
||||
const std::string& isa);
|
||||
|
||||
void cpu_fused_moe_int8(torch::Tensor& output, const torch::Tensor& input,
|
||||
const torch::Tensor& w13, const torch::Tensor& w2,
|
||||
const torch::Tensor& w13_scale,
|
||||
const torch::Tensor& w2_scale,
|
||||
const std::optional<torch::Tensor>& w13_bias,
|
||||
const std::optional<torch::Tensor>& w2_bias,
|
||||
const torch::Tensor& topk_weights,
|
||||
const torch::Tensor& topk_id, const bool skip_weighted,
|
||||
const std::string& act, const std::string& isa);
|
||||
|
||||
void compute_slot_mapping_kernel_impl(const torch::Tensor query_start_loc,
|
||||
const torch::Tensor positions,
|
||||
const torch::Tensor block_table,
|
||||
@@ -502,7 +516,8 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
ops.def(
|
||||
"causal_conv1d_update_cpu(Tensor x, Tensor(a!) conv_states, Tensor "
|
||||
"weight, Tensor? bias, bool silu_activation,"
|
||||
"Tensor? cache_seqlens, Tensor? conv_state_indices, int pad_slot_id, "
|
||||
"Tensor? num_accepted_tokens, Tensor? conv_state_indices, int "
|
||||
"pad_slot_id, "
|
||||
"bool is_vnni) -> Tensor");
|
||||
ops.impl("causal_conv1d_update_cpu", torch::kCPU, &causal_conv1d_update_cpu);
|
||||
#endif
|
||||
@@ -596,8 +611,7 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
#endif
|
||||
|
||||
// fused moe
|
||||
#if defined(__AVX512F__) || \
|
||||
(defined(__aarch64__) && !defined(__APPLE__) && defined(ARM_BF16_SUPPORT))
|
||||
#if defined(__AVX512F__) || (defined(ARM_BF16_SUPPORT) && !defined(__APPLE__))
|
||||
ops.def(
|
||||
"prepack_moe_weight(Tensor weight, Tensor(a1!) packed_weight, str isa) "
|
||||
"-> ()");
|
||||
@@ -608,7 +622,22 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"bool skip_weighted, "
|
||||
"str act, str isa) -> ()");
|
||||
ops.impl("cpu_fused_moe", torch::kCPU, &cpu_fused_moe);
|
||||
#endif
|
||||
#endif // #if defined(__AVX512F__) || (defined(ARM_BF16_SUPPORT) &&
|
||||
// !defined(__APPLE__))
|
||||
#if defined(ARM_I8MM_SUPPORT) && defined(ARM_BF16_SUPPORT) && \
|
||||
!defined(__APPLE__)
|
||||
ops.def(
|
||||
"prepack_moe_weight_int8(Tensor weight, Tensor(a1!) packed_weight, "
|
||||
"str isa) -> ()");
|
||||
ops.impl("prepack_moe_weight_int8", torch::kCPU, &prepack_moe_weight_int8);
|
||||
ops.def(
|
||||
"cpu_fused_moe_int8(Tensor(a0!) output, Tensor input, Tensor w13, "
|
||||
"Tensor w2, Tensor w13_scale, Tensor w2_scale, Tensor? w13_bias, "
|
||||
"Tensor? w2_bias, Tensor topk_weights, Tensor topk_id, bool "
|
||||
"skip_weighted, str act, str isa) -> ()");
|
||||
ops.impl("cpu_fused_moe_int8", torch::kCPU, &cpu_fused_moe_int8);
|
||||
#endif // #if defined(ARM_I8MM_SUPPORT) && defined(ARM_BF16_SUPPORT) &&
|
||||
// !defined(__APPLE__)
|
||||
ops.def(
|
||||
"mla_decode_kvcache("
|
||||
" Tensor! out, Tensor query, Tensor kv_cache,"
|
||||
|
||||
@@ -22,17 +22,25 @@ template <typename AllReduceKernel, typename T>
|
||||
__global__ __quickreduce_launch_bounds_two_shot__ static void
|
||||
allreduce_prototype_twoshot(T const* A, T* B, uint32_t N, uint32_t num_blocks,
|
||||
int rank, uint8_t** dbuffer_list,
|
||||
uint32_t data_offset, uint32_t flag_color,
|
||||
uint32_t data_offset, uint32_t* d_flag_counters,
|
||||
int64_t data_size_per_phase) {
|
||||
int block = blockIdx.x;
|
||||
int grid = gridDim.x;
|
||||
|
||||
// Load this block's counter from device memory and advance it on-device,
|
||||
// so the color keeps changing across graph replays instead of being frozen.
|
||||
uint32_t flag_color = d_flag_counters[blockIdx.x];
|
||||
|
||||
while (block < num_blocks) {
|
||||
AllReduceKernel::run(A, B, N, block, rank, dbuffer_list, data_offset,
|
||||
flag_color, data_size_per_phase);
|
||||
block += grid;
|
||||
flag_color++;
|
||||
}
|
||||
// All threads compute the same final value; one writer per block is enough.
|
||||
if (threadIdx.x == 0 && threadIdx.y == 0) {
|
||||
d_flag_counters[blockIdx.x] = flag_color;
|
||||
}
|
||||
}
|
||||
|
||||
#define TWOSHOT_DISPATCH(__codec) \
|
||||
@@ -42,21 +50,21 @@ allreduce_prototype_twoshot(T const* A, T* B, uint32_t N, uint32_t num_blocks,
|
||||
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
|
||||
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
|
||||
num_blocks, rank, dbuffer_list, data_offset, \
|
||||
flag_color, this->kMaxProblemSize); \
|
||||
d_flag_counters, this->kMaxProblemSize); \
|
||||
} else if (world_size == 4) { \
|
||||
using LineCodec = __codec<T, 4>; \
|
||||
using AllReduceKernel = AllReduceTwoshot<T, LineCodec, cast_bf2half>; \
|
||||
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
|
||||
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
|
||||
num_blocks, rank, dbuffer_list, data_offset, \
|
||||
flag_color, this->kMaxProblemSize); \
|
||||
d_flag_counters, this->kMaxProblemSize); \
|
||||
} else if (world_size == 8) { \
|
||||
using LineCodec = __codec<T, 8>; \
|
||||
using AllReduceKernel = AllReduceTwoshot<T, LineCodec, cast_bf2half>; \
|
||||
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
|
||||
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
|
||||
num_blocks, rank, dbuffer_list, data_offset, \
|
||||
flag_color, this->kMaxProblemSize); \
|
||||
d_flag_counters, this->kMaxProblemSize); \
|
||||
}
|
||||
|
||||
// INT3 only retains good performance on TP2 (world_size == 2). On TP4/TP8
|
||||
@@ -69,7 +77,7 @@ allreduce_prototype_twoshot(T const* A, T* B, uint32_t N, uint32_t num_blocks,
|
||||
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
|
||||
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
|
||||
num_blocks, rank, dbuffer_list, data_offset, \
|
||||
flag_color, this->kMaxProblemSize); \
|
||||
d_flag_counters, this->kMaxProblemSize); \
|
||||
} else { \
|
||||
throw std::runtime_error( \
|
||||
"INT3 quick all-reduce is only supported for world_size == 2 " \
|
||||
@@ -94,7 +102,7 @@ struct DeviceComms {
|
||||
static int constexpr kMaxWorldSize = 8;
|
||||
|
||||
bool initialized = false;
|
||||
uint32_t flag_color = 1;
|
||||
uint32_t* d_flag_counters = nullptr;
|
||||
int world_size;
|
||||
int rank;
|
||||
|
||||
@@ -128,6 +136,16 @@ struct DeviceComms {
|
||||
// Clear the flags buffer.
|
||||
HIP_CHECK(hipMemset(dbuffer, 0, flags_buffer_size));
|
||||
|
||||
// One flag-color counter per block, advanced by the kernel. Start at 1
|
||||
// to stay clear of the flags buffer we just zeroed.
|
||||
HIP_CHECK(hipMalloc(&d_flag_counters, kMaxNumBlocks * sizeof(uint32_t)));
|
||||
{
|
||||
std::vector<uint32_t> init_color(kMaxNumBlocks, 1u);
|
||||
HIP_CHECK(hipMemcpy(d_flag_counters, init_color.data(),
|
||||
kMaxNumBlocks * sizeof(uint32_t),
|
||||
hipMemcpyHostToDevice));
|
||||
}
|
||||
|
||||
// Device-side list of IPC buffers.
|
||||
buffer_list.resize(world_size);
|
||||
HIP_CHECK(hipMalloc(&dbuffer_list, world_size * sizeof(uint8_t*)));
|
||||
@@ -144,6 +162,12 @@ struct DeviceComms {
|
||||
hipIpcMemHandle_t const get_handle() { return buffer_ipc_handle; }
|
||||
|
||||
void destroy() {
|
||||
// Allocated before `initialized` flips true, so free it on its own guard
|
||||
// to avoid a leak if init fails partway through.
|
||||
if (d_flag_counters) {
|
||||
HIP_CHECK(hipFree(d_flag_counters));
|
||||
d_flag_counters = nullptr;
|
||||
}
|
||||
if (initialized) {
|
||||
for (int i = 0; i < world_size; i++) {
|
||||
if (i != rank) {
|
||||
@@ -211,8 +235,6 @@ struct DeviceComms {
|
||||
break;
|
||||
}
|
||||
HIP_CHECK(cudaGetLastError());
|
||||
// Rotate the flag color.
|
||||
flag_color += divceil(N, grid);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
+3
-1
@@ -56,7 +56,9 @@ nav:
|
||||
- API Reference:
|
||||
- api/README.md
|
||||
- api/vllm
|
||||
- CLI Reference: cli
|
||||
- CLI Reference:
|
||||
- cli/README.md
|
||||
- vllm: cli
|
||||
- Community:
|
||||
- community/*
|
||||
- Governance: governance
|
||||
|
||||
+7
-9
@@ -1,10 +1,8 @@
|
||||
nav:
|
||||
- README.md
|
||||
- serve.md
|
||||
- chat.md
|
||||
- complete.md
|
||||
- run-batch.md
|
||||
- vllm bench:
|
||||
- bench/**/*.md
|
||||
- vllm launch:
|
||||
- launch/**/*.md
|
||||
- "*.md"
|
||||
- bench:
|
||||
- bench/*.md
|
||||
- sweep:
|
||||
- bench/sweep/*.md
|
||||
- launch:
|
||||
- launch/*.md
|
||||
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm bench latency
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_latency.inc.md"
|
||||
@@ -1,55 +0,0 @@
|
||||
# vllm bench mm-processor
|
||||
|
||||
## Overview
|
||||
|
||||
`vllm bench mm-processor` profiles the multimodal input processor pipeline of
|
||||
vision-language models. It measures per-stage latency from the HuggingFace
|
||||
processor through to the encoder forward pass, helping you identify
|
||||
preprocessing bottlenecks and understand how different image resolutions or
|
||||
item counts affect end-to-end request time.
|
||||
|
||||
The benchmark supports two data sources: synthetic random multimodal inputs
|
||||
(`random-mm`) and HuggingFace datasets (`hf`). Warmup requests are run before
|
||||
measurement to ensure stable results.
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
vllm bench mm-processor \
|
||||
--model Qwen/Qwen2-VL-7B-Instruct \
|
||||
--dataset-name random-mm \
|
||||
--num-prompts 50 \
|
||||
--random-input-len 300 \
|
||||
--random-output-len 40 \
|
||||
--random-mm-base-items-per-request 2 \
|
||||
--random-mm-limit-mm-per-prompt '{"image": 3, "video": 0}' \
|
||||
--random-mm-bucket-config '{(256, 256, 1): 0.7, (720, 1280, 1): 0.3}'
|
||||
```
|
||||
|
||||
## Measured Stages
|
||||
|
||||
| Stage | Description |
|
||||
| ----- | ----------- |
|
||||
| `get_mm_hashes_secs` | Time spent hashing multimodal inputs |
|
||||
| `get_cache_missing_items_secs` | Time spent looking up the processor cache |
|
||||
| `apply_hf_processor_secs` | Time spent in the HuggingFace processor |
|
||||
| `merge_mm_kwargs_secs` | Time spent merging multimodal kwargs |
|
||||
| `apply_prompt_updates_secs` | Time spent updating prompt tokens |
|
||||
| `preprocessor_total_secs` | Total preprocessing time |
|
||||
| `encoder_forward_secs` | Time spent in the encoder model forward pass |
|
||||
| `num_encoder_calls` | Number of encoder invocations per request |
|
||||
|
||||
The benchmark also reports end-to-end latency (TTFT + decode time) per
|
||||
request. Use `--metric-percentiles` to select which percentiles to report
|
||||
(default: p99) and `--output-json` to save results.
|
||||
|
||||
For more examples (HF datasets, warmup, JSON output), see
|
||||
[Benchmarking CLI — Multimodal Processor Benchmark](../../benchmarking/cli.md#multimodal-processor-benchmark).
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_mm_processor.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm bench serve
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_serve.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm bench sweep plot
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_sweep_plot.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm bench sweep plot_pareto
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_sweep_plot_pareto.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm bench sweep serve
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_sweep_serve.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm bench sweep serve_workload
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_sweep_serve_workload.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm bench throughput
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/bench_throughput.inc.md"
|
||||
@@ -1,5 +0,0 @@
|
||||
# vllm chat
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/chat.inc.md"
|
||||
@@ -1,5 +0,0 @@
|
||||
# vllm complete
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/complete.inc.md"
|
||||
@@ -1,10 +0,0 @@
|
||||
<!-- markdownlint-disable MD041 -->
|
||||
When passing JSON CLI arguments, the following sets of arguments are equivalent:
|
||||
|
||||
- `--json-arg '{"key1": "value1", "key2": {"key3": "value2"}}'`
|
||||
- `--json-arg.key1 value1 --json-arg.key2.key3 value2`
|
||||
|
||||
Additionally, list elements can be passed individually using `+`:
|
||||
|
||||
- `--json-arg '{"key4": ["value3", "value4", "value5"]}'`
|
||||
- `--json-arg.key4+ value3 --json-arg.key4+='value4,value5'`
|
||||
@@ -1,22 +0,0 @@
|
||||
# vllm launch render
|
||||
|
||||
## Overview
|
||||
|
||||
`vllm launch render` starts a GPU-less rendering server for preprocessing and
|
||||
postprocessing only.
|
||||
|
||||
```bash
|
||||
vllm launch render meta-llama/Llama-3.2-1B-Instruct --port 8100
|
||||
```
|
||||
|
||||
This command reuses the standard serving parser, so model, frontend,
|
||||
networking, and related CLI options follow the same conventions as
|
||||
[`vllm serve`](../serve.md).
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/launch_render.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm run-batch
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/run-batch.inc.md"
|
||||
@@ -1,9 +0,0 @@
|
||||
# vllm serve
|
||||
|
||||
## JSON CLI Arguments
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## Arguments
|
||||
|
||||
--8<-- "docs/generated/argparse/serve.inc.md"
|
||||
@@ -11,12 +11,4 @@ Engine arguments control the behavior of the vLLM engine.
|
||||
|
||||
The engine argument classes, [EngineArgs][vllm.engine.arg_utils.EngineArgs] and [AsyncEngineArgs][vllm.engine.arg_utils.AsyncEngineArgs], are a combination of the configuration classes defined in [vllm.config][]. Therefore, if you are interested in developer documentation, we recommend looking at these configuration classes as they are the source of truth for types, defaults and docstrings.
|
||||
|
||||
--8<-- "docs/cli/json_tip.inc.md"
|
||||
|
||||
## `EngineArgs`
|
||||
|
||||
--8<-- "docs/generated/argparse/engine_args.inc.md"
|
||||
|
||||
## `AsyncEngineArgs`
|
||||
|
||||
--8<-- "docs/generated/argparse/async_engine_args.inc.md"
|
||||
--8<-- "gen:engine-args"
|
||||
|
||||
@@ -195,7 +195,7 @@ Provide a fast duration→token estimate to improve streaming usage statistics:
|
||||
The API server takes care of basic audio I/O and optional chunking before building prompts:
|
||||
|
||||
- Resampling: Input audio is resampled to `SpeechToTextConfig.sample_rate` using `AudioResampler`.
|
||||
- Chunking: If `SpeechToTextConfig.allow_audio_chunking` is True and the duration exceeds `max_audio_clip_s`, the server splits the audio into overlapping chunks and generates a prompt per chunk. Overlap is controlled by `overlap_chunk_second`.
|
||||
- Chunking: If `SpeechToTextConfig.allow_audio_chunking` is True and the duration exceeds `max_audio_clip_s`, the server splits the audio into chunks and generates a prompt per chunk. There is no overlap between chunks, overlap_chunk_second controls the size of the search window used to find the split point.
|
||||
- Energy-aware splitting: When `min_energy_split_window_size` is set, the server finds low-energy regions to minimize cutting within words.
|
||||
|
||||
Relevant server logic:
|
||||
|
||||
@@ -8,6 +8,26 @@ toc_depth: 2
|
||||
|
||||
--8<-- "docs/getting_started/installation/gpu.md:pre-built-images"
|
||||
|
||||
## Persist the compile cache across containers
|
||||
|
||||
Mounting the Hugging Face cache keeps model weights across containers, but each
|
||||
new container still starts with an empty `VLLM_CACHE_ROOT` (default
|
||||
`~/.cache/vllm`) and recompiles the model's `torch.compile` artifacts. Mount a
|
||||
named volume at that path to reuse the inductor, Triton, and AOT artifacts from
|
||||
the second container onward:
|
||||
|
||||
```bash
|
||||
docker run --rm --gpus all \
|
||||
-v ~/.cache/huggingface:/root/.cache/huggingface \
|
||||
-v vllm-cache:/root/.cache/vllm \
|
||||
-p 8000:8000 \
|
||||
vllm/vllm-openai:latest \
|
||||
meta-llama/Llama-3.1-8B-Instruct
|
||||
```
|
||||
|
||||
See [Faster Startup](../configuration/optimization.md#faster-startup) for the
|
||||
mechanism and for what invalidates the cache.
|
||||
|
||||
## Run as a non-root user
|
||||
|
||||
The CUDA `vllm/vllm-openai` image runs as root by default for backward
|
||||
|
||||
@@ -1,15 +1,9 @@
|
||||
# Attention Backend Feature Support
|
||||
|
||||
This document is auto-generated by `tools/pre_commit/generate_attention_backend_docs.py`.
|
||||
It shows the feature support for each registered attention backend
|
||||
based on the checks in `AttentionBackend.validate_configuration()`.
|
||||
|
||||
**Do not edit this file manually.** Run the following command to
|
||||
regenerate it:
|
||||
|
||||
```bash
|
||||
python tools/pre_commit/generate_attention_backend_docs.py
|
||||
```
|
||||
The priority and feature tables on this page are auto-generated from the
|
||||
attention backend registry by
|
||||
`docs/mkdocs/gen_files/generate_attention_backends.py`, based on the checks in
|
||||
`AttentionBackend.validate_configuration()`.
|
||||
|
||||
## Setting the Attention Backend
|
||||
|
||||
@@ -98,40 +92,11 @@ Priority is **1 = highest** (tried first).
|
||||
|
||||
### Standard Attention (MHA, MQA, GQA)
|
||||
|
||||
**Blackwell (SM 10.x):**
|
||||
|
||||
| Priority | Backend |
|
||||
| -------- | ------- |
|
||||
| 1 | `FLASHINFER` |
|
||||
| 2 | `FLASH_ATTN` |
|
||||
| 3 | `TRITON_ATTN` |
|
||||
| 4 | `FLEX_ATTENTION` |
|
||||
| 5 | `TURBOQUANT` |
|
||||
|
||||
**Ampere/Hopper (SM 8.x-9.x):**
|
||||
|
||||
| Priority | Backend |
|
||||
| -------- | ------- |
|
||||
| 1 | `FLASH_ATTN` |
|
||||
| 2 | `FLASHINFER` |
|
||||
| 3 | `TRITON_ATTN` |
|
||||
| 4 | `FLEX_ATTENTION` |
|
||||
| 5 | `TURBOQUANT` |
|
||||
--8<-- "gen:priority-standard"
|
||||
|
||||
### MLA Attention (DeepSeek-style)
|
||||
|
||||
**Blackwell (SM 10.x):**
|
||||
|
||||
| Priority | Backend |
|
||||
| -------- | ------- |
|
||||
| 1 | `FLASHINFER_MLA` |
|
||||
| 2 | `TOKENSPEED_MLA` |
|
||||
| 3 | `CUTLASS_MLA` |
|
||||
| 4 | `FLASH_ATTN_MLA` |
|
||||
| 5 | `FLASHMLA` |
|
||||
| 6 | `TRITON_MLA` |
|
||||
| 7 | `FLASHINFER_MLA_SPARSE`**\*** |
|
||||
| 8 | `FLASHMLA_SPARSE` |
|
||||
--8<-- "gen:priority-mla"
|
||||
|
||||
> **\*** For sparse MLA, FP8 KV cache always prefers `FLASHINFER_MLA_SPARSE`. With BF16 KV cache, `FLASHINFER_MLA_SPARSE` is preferred for low query-head counts (<= 16), while `FLASHMLA_SPARSE` is preferred otherwise.
|
||||
>
|
||||
@@ -157,24 +122,7 @@ Priority is **1 = highest** (tried first).
|
||||
|
||||
## Standard Attention (MHA, MQA, GQA) Backends
|
||||
|
||||
| Backend | Version | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | MM Prefix | DCP | Attention Types | Compute Cap. |
|
||||
| ------- | ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | --------- | --- | --------------- | ------------ |
|
||||
| `CPU_ATTN` | | fp16, bf16, fp32 | `auto`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 112, 128, 160, 192, 224, 256, 512 | ❌ | ✅ | ❌ | ❌ | All | N/A |
|
||||
| `FLASHINFER` | Native† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ❌ | ✅ | ❌ | ✅ | Decoder | 8.x-9.x |
|
||||
| `FLASHINFER` | XQA† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ❌ | ❌ | ❌ | ✅ | Decoder | 9.0 |
|
||||
| `FLASHINFER` | trtllm-gen† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `nvfp4` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ✅ | ✅ | ❌ | ✅ | Decoder | 10.x |
|
||||
| `FLASH_ATTN` | FA2* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ❌ | ✅ | All | ≥8.0 |
|
||||
| `FLASH_ATTN` | FA3* | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | 9.x |
|
||||
| `FLASH_ATTN` | FA4* | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | ≥10.0 |
|
||||
| `FLASH_ATTN_DIFFKV` | | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
|
||||
| `FLEX_ATTENTION` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder Only | Any |
|
||||
| `HPC_ATTN` | | fp16, bf16 | `auto`, `bfloat16`, `fp8_e4m3` | 64 | 128 | ❌ | ❌ | ❌ | ❌ | Decoder | ≥9.0 |
|
||||
| `ROCM_AITER_FA` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32 | 64, 128, 256 | ✅ | ✅ | ❌ | ❌ | Decoder | N/A |
|
||||
| `ROCM_AITER_UNIFIED_ATTN` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | Any | ✅ | ❌ | ✅ | ❌ | All | N/A |
|
||||
| `ROCM_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 128, 160, 192, 224, 256 | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder, Encoder Only | N/A |
|
||||
| `TRITON_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `int4_per_token_head`, `int8_per_token_head`, `fp8_per_token_head` | %16 | Any | ✅ | ✅ | ✅ | ❌ | All | Any |
|
||||
| `TRITON_ATTN_DIFFKV` | | fp16, bf16 | `auto`, `bfloat16` | Any | Any | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
|
||||
| `TURBOQUANT` | | fp16, bf16 | `turboquant_k8v4`, `turboquant_4bit_nc`, `turboquant_k3v4_nc`, `turboquant_3bit_nc` | 16, 32, 64, 128 | Any | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
|
||||
--8<-- "gen:table-standard"
|
||||
|
||||
> **†** FlashInfer Native is the regular FlashInfer path. XQA is the SM90 decode path exposed through FlashInfer's TRTLLM decode API. trtllm-gen is used on SM100 and supports sinks. Disable XQA/trtllm-gen via `--attention-config.use_trtllm_attention=0`.
|
||||
>
|
||||
@@ -188,9 +136,7 @@ automatic priority lists above. A lightning indexer scores KV blocks, the
|
||||
top-k blocks (plus fixed init/local blocks) are selected, and attention
|
||||
attends only to those blocks; index keys live in a separate side cache.
|
||||
|
||||
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | MM Prefix | DCP | Attention Types | Compute Cap. |
|
||||
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | --------- | --- | --------------- | ------------ |
|
||||
| `MINIMAX_M3_SPARSE` | bf16, fp16 | `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 128 | 128 | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
|
||||
--8<-- "gen:table-minimax"
|
||||
|
||||
## MLA (Multi-head Latent Attention) Backends
|
||||
|
||||
@@ -203,38 +149,20 @@ To explicitly select a prefill backend, use
|
||||
Otherwise, the prefill backend is selected automatically at runtime based on
|
||||
hardware and configuration.
|
||||
|
||||
| Backend | Description | Dtypes | Compute Cap. | Notes |
|
||||
| ------- | ----------- | ------ | ------------ | ----- |
|
||||
| `FLASH_ATTN`‡ | FlashAttention varlen (FA2/FA3/FA4) | fp16, bf16 | Any | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) (FA2/FA3/FA4) or (qk_nope_head_dim=64, qk_rope_head_dim=64, v_head_dim=128) (FA2/FA3/FA4) or (qk_nope_head_dim=192, qk_rope_head_dim=64, v_head_dim=256) (FA2/FA3 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 |
|
||||
--8<-- "gen:table-mla-prefill"
|
||||
|
||||
> **‡** Automatic selection tries FlashAttention first. On Blackwell
|
||||
> (SM100), the fallback order is TRT-LLM Ragged, FlashInfer, then
|
||||
> TokenSpeed MLA. On other GPUs, only FlashAttention is considered.
|
||||
> TokenSpeed MLA; for (qk_nope_head_dim=192, qk_rope_head_dim=64,
|
||||
> v_head_dim=256) TRT-LLM Ragged is tried before FlashAttention. On other
|
||||
> GPUs, only FlashAttention is considered.
|
||||
|
||||
### Decode Backends
|
||||
|
||||
MLA decode backends are selected using the standard
|
||||
`-ac.backend=<BACKEND>` argument (e.g., `FLASHMLA`, `TRITON_MLA`).
|
||||
|
||||
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
|
||||
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ |
|
||||
| `CUTLASS_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 128 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
|
||||
| `FLASHINFER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
|
||||
| `FLASHINFER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
|
||||
| `FLASHINFER_MLA_SPARSE_SM120` | bf16 | `auto`, `fp8`, `fp8_e4m3`, `fp8_ds_mla` | 64, 256 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | 12.x |
|
||||
| `FLASHMLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x-10.x |
|
||||
| `FLASHMLA_SPARSE` | bf16 | `auto`, `bfloat16`, `fp8_ds_mla` | 64 | 576 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x-10.x |
|
||||
| `FLASH_ATTN_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x |
|
||||
| `FLASH_ATTN_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | 64 | Any | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x |
|
||||
| `ROCM_AITER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %1 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
|
||||
| `ROCM_AITER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 1, 64 | Any | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | N/A |
|
||||
| `ROCM_AITER_TRITON_MLA` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
|
||||
| `TOKENSPEED_MLA` | fp16, bf16 | `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
|
||||
| `TRITON_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
|
||||
| `XPU_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | Any | 576 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | Any |
|
||||
--8<-- "gen:table-mla-decode"
|
||||
|
||||
### DeepSeek V4 Decode Backends
|
||||
|
||||
@@ -245,8 +173,4 @@ pipeline (compressor + SWA + indexer, 256-token blocks, head 512);
|
||||
default on NVIDIA is `FLASHINFER_MLA_SPARSE_DSV4` on SM12x and
|
||||
`FLASHMLA_SPARSE_DSV4` on other supported CUDA architectures.
|
||||
|
||||
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
|
||||
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ |
|
||||
| `FLASHINFER_MLA_SPARSE_DSV4` | bf16 | `auto`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_ds_mla` | 256 | 512 | ✅ | ❌ | ✅ | ❌ | ❌ | Decoder | 10.x, 12.x |
|
||||
| `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 |
|
||||
--8<-- "gen:table-mla-v4-decode"
|
||||
|
||||
@@ -129,6 +129,7 @@ Models opt-in to encoder CUDA Graphs by implementing the [SupportsEncoderCudaGra
|
||||
| `DeepseekOCRForCausalLM` | `DeepSeek-OCR` | ✅︎ | ❌︎ | ✅︎ |
|
||||
| `Gemma3ForConditionalGeneration` | `Gemma3` | ✅︎ | ❌︎ | ❌︎ |
|
||||
| `Glm4vForConditionalGeneration` | `GLM-4.1V, GLM-4.6V-Flash` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `Gemma4ForConditionalGeneration` | `Gemma-4` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `InternVLChatModel` | `InternVL3.5`, `InternVL3`, `InternVL2.5`, `InternVL2` | ✅︎ | ✅︎ | ❌︎ |
|
||||
| `KimiVLForConditionalGeneration` | `Kimi-VL` | ✅︎ | ❌︎ | ❌︎ |
|
||||
| `Llama4ForConditionalGeneration` | `Llama 4` | ✅︎ | ❌︎ | ❌︎ |
|
||||
|
||||
@@ -157,7 +157,7 @@ Object keys follow the same run-configuration digest scheme as the filesystem ti
|
||||
|
||||
The P2P tier (`type: "p2p"`) shares completed KV blocks between vLLM instances over RDMA via NIXL. Each instance binds a control socket on `host:port` and exchanges blocks directly with peers — no shared filesystem required.
|
||||
|
||||
PYTHONHASHSEED environment variable must be set to the same fixed value on all nodes.
|
||||
The `PYTHONHASHSEED` environment variable must be set to the same fixed value (e.g. `"0"`) on all nodes so that block content hashes match across instances (see [Cross-Process Sharing](#cross-process-sharing)). This is enforced: a P2P instance started without `PYTHONHASHSEED` set fails at startup, and each peer's value is verified during the connect handshake — a peer advertising a different `PYTHONHASHSEED` is rejected.
|
||||
|
||||
| Key | Required | Default | Notes |
|
||||
| --- | --- | --- | --- |
|
||||
@@ -176,6 +176,71 @@ Rather than embedding `host`/`port` in each `secondary_tiers` entry, set them on
|
||||
- `VLLM_P2P_SIDE_CHANNEL_HOST` (default `localhost`): address the P2P control socket binds to. It is used **verbatim** as both the bind address and the identity peers dial back — there is no auto-detection (this mirrors `VLLM_NIXL_SIDE_CHANNEL_HOST`). The default binds the loopback interface only, so peers on another host cannot reach it. **For any cross-host P2P deployment you must set this explicitly to the node's routable IP** (e.g. the pod IP) before launching `vllm serve` — otherwise remote peers will fail to connect. The NIXL agent name is a separate per-process identifier, so peers sharing a `host:port` never collide.
|
||||
- `VLLM_P2P_SIDE_CHANNEL_PORT` (default `5710`): base port for the P2P control socket. The port actually bound is `VLLM_P2P_SIDE_CHANNEL_PORT + data_parallel_index` — one socket per DP replica, matching NIXL (for DP=1 the offset is 0). The peer's port is passed as `remote_port` in `kv_transfer_params`; the router/EPP that selects the DP rank (e.g. via the `X-data-parallel-rank` header) computes `remote_port = base + rank`. The DP-index offset separates replicas *within* one deployment; two co-located *deployments* (a prefiller and a decoder on the same host) still need distinct base ports (e.g. decoder base `5711`) to avoid a bind collision.
|
||||
|
||||
#### Orchestration-Layer Protocol
|
||||
|
||||
The P2P tier does not decide *which* peer to pull from — that is the orchestration layer's job (the router/EPP and its scheduler). The orchestrator drives every transfer through a request's `kv_transfer_params` dict: it picks the request's role, allocates a unique transaction ID, and supplies the remote peer's address. All block lookup, hash matching, and NIXL transfer happen at the tier level below; the orchestrator only sets the correct role keys and enforces the allowed combinations.
|
||||
|
||||
Every vLLM instance is a symmetric **peer**. Per request it acts as a **consumer** (pulls KV blocks from a remote peer's CPU cache instead of computing locally) or a **producer** (serves blocks from its own CPU cache to remote consumers) — or both, on the same session, for different requests. Roles are chosen per request by the keys below; there are no fixed prefiller/decoder processes.
|
||||
|
||||
Three role keys are defined, each mapping to a sub-dict. All are optional; a request with none of them uses the tier only as a local CPU cache.
|
||||
|
||||
Each key names the **remote counterpart** this peer transfers with (not this
|
||||
peer's own role), so the name reads as "the remote ___ I transfer with".
|
||||
|
||||
| Key | Set on | Value fields | Meaning |
|
||||
| --- | --- | --- | --- |
|
||||
| `remote_decoder` | prefill producer request | `kv_request_id` | Peer computes KV and keeps it available in CPU cache for the remote decoder to pull. |
|
||||
| `remote_prefiller` | decode consumer request | `kv_request_id`, `remote_host`, `remote_port` | Peer pulls KV from the remote prefiller at the given address (classic P/D disaggregation). |
|
||||
| `remote_kv_source` | P2P consumer request | `kv_request_id`, `remote_host`, `remote_port` | Peer looks up and pulls whatever blocks the remote source currently holds in CPU cache. |
|
||||
|
||||
Field semantics:
|
||||
|
||||
- `kv_request_id` (str): unique transaction ID allocated by the orchestrator and pushed to every peer involved in the transfer; used to correlate the lookup, fetch, and transfer-done messages. The producer is implicit — it serves whatever block hashes it currently holds in its CPU cache for that ID.
|
||||
- `remote_host` (str): IP/hostname of the remote peer's control socket to query. Must be the peer's routable node IP (see [Environment Variables](#environment-variables)).
|
||||
- `remote_port` (int): the peer's bound control-socket port, i.e. `base + data_parallel_index` for the selected DP rank.
|
||||
|
||||
Allowed and forbidden combinations:
|
||||
|
||||
- **`remote_decoder` + `remote_kv_source`** is the only legal multi-key combination: a prefill producer may *also* act as a P2P consumer for the same request — skipping prefix prefill by pulling cached blocks from a source while still keeping its own computed blocks available for a downstream decoder.
|
||||
- Forbidden: `remote_prefiller` + `remote_decoder` (contradictory roles), `remote_prefiller` + `remote_kv_source` (two competing fetch sources), and all three together.
|
||||
|
||||
Minimal examples (values that would appear in the request's `kv_transfer_params`):
|
||||
|
||||
```python
|
||||
# Prefill producer — compute and keep KV for a remote decoder to pull
|
||||
kv_transfer_params = {"remote_decoder": {"kv_request_id": "<unique-transfer-id>"}}
|
||||
|
||||
# Decode consumer — pull KV from a specific prefiller (classic P/D)
|
||||
kv_transfer_params = {
|
||||
"remote_prefiller": {
|
||||
"kv_request_id": "<unique-transfer-id>",
|
||||
"remote_host": "<prefiller-node-ip>",
|
||||
"remote_port": 5710,
|
||||
}
|
||||
}
|
||||
|
||||
# P2P consumer — pull whatever the source already has cached
|
||||
kv_transfer_params = {
|
||||
"remote_kv_source": {
|
||||
"kv_request_id": "<unique-transfer-id>",
|
||||
"remote_host": "<source-node-ip>",
|
||||
"remote_port": 5710,
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Runtime handshake for a P2P (or P/D) pull, once the orchestrator has set the keys above:
|
||||
|
||||
1. Both peers already have listener threads on their control sockets (see [Environment Variables](#environment-variables)).
|
||||
2. **Lookup.** The consumer's tiering manager does per-block lookups; in P2P mode the tier returns `None` and registers the key. At `on_schedule_end` the consumer sends one **`LookupMsg`** (`kv_request_id` + block hashes) to the peer, per request, per step.
|
||||
3. The producer matches those hashes against its local CPU cache and replies with a **`LookupRespMsg`** carrying the hit block hashes.
|
||||
4. **Resolve.** Retried lookups now return hit / miss / in-flight. The consumer calls `submit_load` for hits only, allocating CPU slots only for hits.
|
||||
5. The consumer sends a **`FetchMsg`** (`kv_request_id`, block hashes, destination block indexes).
|
||||
6. The producer performs the **NIXL WRITE** transfer and sends **`TransferDone`** with a success status.
|
||||
7. On `get_finished`, hits are loaded into GPU as ordinary cache hits; misses are recomputed by the engine.
|
||||
|
||||
In classic **P/D mode** (`remote_prefiller` set, no `remote_kv_source`), the lookup phase (steps 2–4) is skipped: the decode consumer assumes the prefiller holds all of the request's blocks, so every block `lookup()` returns an immediate hit and the consumer jumps straight to the **`FetchMsg`** in step 5. The `LookupMsg`/`LookupRespMsg` round-trip only happens in P2P mode, where the consumer does not know in advance which blocks the peer has cached.
|
||||
|
||||
## Tuning Tips
|
||||
|
||||
- `cpu_bytes_to_use`: a bigger CPU tier means fewer trips to slower secondary tiers and a higher hit rate. The value is total across all workers, not per-worker. Leave headroom for the rest of the host workload.
|
||||
|
||||
@@ -818,16 +818,18 @@ Full example: [examples/generate/multimodal/openai_chat_completion_client_for_mu
|
||||
|
||||
#### Video Decoding Backend
|
||||
|
||||
vLLM decodes video bytes into frames using a selectable decoding backend. Three
|
||||
vLLM decodes video bytes into frames using a selectable decoding backend. Five
|
||||
backends are supported:
|
||||
|
||||
- `opencv` (default): OpenCV-based decoder.
|
||||
- `pyav`: PyAV decoder.
|
||||
- `torchcodec`: TorchCodec (PyTorch-native) decoder.
|
||||
- `pynvvideocodec`: NVIDIA NVDEC-based decoder.
|
||||
- `deepstream`: NVIDIA DeepStream NVDEC-based decoder.
|
||||
|
||||
All three backends are ultimately backed by FFmpeg. `torchcodec` lets
|
||||
you choose which FFmpeg version is used while `opencv` and `pyav` rely on
|
||||
whichever FFmpeg build they were linked against.
|
||||
The CPU backends are backed by FFmpeg. `torchcodec` lets you choose which FFmpeg
|
||||
version is used while `opencv` and `pyav` rely on whichever FFmpeg build they
|
||||
were linked against.
|
||||
|
||||
Select the backend by passing the `backend` parameter via `--media-io-kwargs`:
|
||||
|
||||
@@ -854,6 +856,21 @@ vllm serve Qwen/Qwen3-VL-30B-A3B-Instruct \
|
||||
--media-io-kwargs '{"video": {"backend": "torchcodec", "seek_mode": "approximate", "num_ffmpeg_threads": 4}}'
|
||||
```
|
||||
|
||||
**PyNvVideoCodec-specific parameters:**
|
||||
|
||||
- `hw_decoders`: Maximum number of concurrent hardware decoder slots retained
|
||||
by each API server process. It must be a positive integer and defaults to `2`,
|
||||
which is the recommended starting point for concurrent video workloads.
|
||||
Because vLLM reserves GPU memory for these slots at startup, this value cannot
|
||||
be overridden per request. Benchmark before increasing it because each
|
||||
additional slot increases the GPU memory reservation.
|
||||
|
||||
```bash
|
||||
# Example: explicitly use the recommended 2 hardware decoders
|
||||
vllm serve Qwen/Qwen3-VL-30B-A3B-Instruct \
|
||||
--media-io-kwargs '{"video": {"backend": "pynvvideocodec", "hw_decoders": 2}}'
|
||||
```
|
||||
|
||||
#### Video Frame Recovery
|
||||
|
||||
For improved robustness when processing potentially corrupted or truncated video files, vLLM supports optional frame recovery using a dynamic window forward-scan approach. When enabled, if a target frame fails to load during sequential reading, the next successfully grabbed frame (before the next target frame) will be used in its place.
|
||||
|
||||
@@ -19,6 +19,20 @@ following `quantization.quant_algo` values:
|
||||
- `NVFP4`: ModelOpt NVFP4 checkpoints (use `quantization="modelopt_fp4"`).
|
||||
- `MXFP8`: ModelOpt MXFP8 checkpoints (use `quantization="modelopt_mxfp8"`).
|
||||
|
||||
!!! note
|
||||
For NVFP4 checkpoints, vLLM selects a GEMM kernel automatically at load
|
||||
time from the backends available on the current platform (CUTLASS,
|
||||
FlashInfer, Marlin, and others). On GPUs without a supported native FP4
|
||||
GEMM kernel, vLLM falls back to weight-only (W4A16) execution via Marlin
|
||||
and logs a warning; this may reduce throughput for compute-heavy
|
||||
workloads. Use `--linear-backend` to override the automatic selection
|
||||
(this replaces the deprecated `VLLM_NVFP4_GEMM_BACKEND` environment
|
||||
variable). Values relevant to NVFP4 include `cutlass`,
|
||||
`flashinfer_cutlass`, `flashinfer_trtllm`, `flashinfer_cudnn`, and
|
||||
`marlin`; the full list is documented under `KernelConfig` on the
|
||||
[Engine Arguments](../../configuration/engine_args.md) page and shown by
|
||||
`vllm serve --help=KernelConfig`.
|
||||
|
||||
## Quantizing HuggingFace Models with PTQ
|
||||
|
||||
You can quantize HuggingFace models using the example scripts provided in the Model Optimizer repository. The primary script for LLM PTQ is typically found within the `examples/llm_ptq` directory.
|
||||
|
||||
+183
-48
@@ -2,6 +2,7 @@
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import importlib.metadata
|
||||
import importlib.util
|
||||
import inspect
|
||||
import logging
|
||||
import sys
|
||||
import textwrap
|
||||
@@ -10,17 +11,21 @@ from argparse import SUPPRESS, Action, HelpFormatter
|
||||
from collections.abc import Callable, Iterable
|
||||
from importlib.machinery import ModuleSpec
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
from typing import TYPE_CHECKING
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import mkdocs_gen_files
|
||||
import regex as re
|
||||
from pydantic_core import core_schema
|
||||
|
||||
logger = logging.getLogger("mkdocs")
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.parent
|
||||
ARGPARSE_DOC_DIR = ROOT_DIR / "docs/generated/argparse"
|
||||
|
||||
sys.path.insert(0, str(ROOT_DIR))
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
|
||||
from generated_content import fill_markers # noqa: E402
|
||||
|
||||
|
||||
def mock_if_no_torch(mock_module: str, mock: MagicMock):
|
||||
@@ -132,8 +137,8 @@ def auto_mock(module_name: str, attr: str, max_mocks: int = 100):
|
||||
|
||||
|
||||
bench_latency = auto_mock("vllm.benchmarks", "latency")
|
||||
bench_mm_processor = auto_mock("vllm.benchmarks", "mm_processor")
|
||||
bench_serve = auto_mock("vllm.benchmarks", "serve")
|
||||
bench_startup = auto_mock("vllm.benchmarks", "startup")
|
||||
bench_sweep_plot = auto_mock("vllm.benchmarks.sweep.plot", "SweepPlotArgs")
|
||||
bench_sweep_plot_pareto = auto_mock(
|
||||
"vllm.benchmarks.sweep.plot_pareto", "SweepPlotParetoArgs"
|
||||
@@ -142,12 +147,28 @@ bench_sweep_serve = auto_mock("vllm.benchmarks.sweep.serve", "SweepServeArgs")
|
||||
bench_sweep_serve_workload = auto_mock(
|
||||
"vllm.benchmarks.sweep.serve_workload", "SweepServeWorkloadArgs"
|
||||
)
|
||||
bench_sweep_startup = auto_mock("vllm.benchmarks.sweep.startup", "SweepStartupArgs")
|
||||
bench_throughput = auto_mock("vllm.benchmarks", "throughput")
|
||||
AsyncEngineArgs = auto_mock("vllm.engine.arg_utils", "AsyncEngineArgs")
|
||||
EngineArgs = auto_mock("vllm.engine.arg_utils", "EngineArgs")
|
||||
ChatCommand = auto_mock("vllm.entrypoints.cli.openai", "ChatCommand")
|
||||
CompleteCommand = auto_mock("vllm.entrypoints.cli.openai", "CompleteCommand")
|
||||
BenchmarkSubcommand = auto_mock(
|
||||
"vllm.entrypoints.cli.benchmark.main", "BenchmarkSubcommand"
|
||||
)
|
||||
import_bench_subcommands = auto_mock(
|
||||
"vllm.entrypoints.cli.benchmark.main", "_import_bench_subcommand_modules"
|
||||
)
|
||||
BenchmarkSubcommandBase = auto_mock(
|
||||
"vllm.entrypoints.cli.benchmark.base", "BenchmarkSubcommandBase"
|
||||
)
|
||||
BenchmarkMMProcessorSubcommand = auto_mock(
|
||||
"vllm.entrypoints.cli.benchmark.mm_processor", "BenchmarkMMProcessorSubcommand"
|
||||
)
|
||||
LaunchSubcommandBase = auto_mock("vllm.entrypoints.cli.launch", "LaunchSubcommandBase")
|
||||
launch_description = auto_mock("vllm.entrypoints.cli.launch", "DESCRIPTION")
|
||||
RenderSubcommand = auto_mock("vllm.entrypoints.cli.launch", "RenderSubcommand")
|
||||
sweep_subcommands = auto_mock("vllm.benchmarks.sweep.cli", "SUBCOMMANDS")
|
||||
openai_cli_args = auto_mock("vllm.entrypoints.openai", "cli_args")
|
||||
openai_run_batch = auto_mock("vllm.entrypoints.openai", "run_batch")
|
||||
|
||||
@@ -179,7 +200,7 @@ class MarkdownFormatter(HelpFormatter):
|
||||
|
||||
def add_text(self, text: str):
|
||||
if text:
|
||||
self._markdown_output.append(f"{text.strip()}\n\n")
|
||||
self._markdown_output.append(f"{inspect.cleandoc(text)}\n\n")
|
||||
|
||||
def add_usage(self, usage, actions, groups, prefix=None):
|
||||
pass
|
||||
@@ -241,49 +262,163 @@ def create_parser(add_cli_args, **kwargs) -> FlexibleArgumentParser:
|
||||
return _parser or parser
|
||||
|
||||
|
||||
def on_startup(command: Literal["build", "gh-deploy", "serve"], dirty: bool):
|
||||
logger.info("Generating argparse documentation")
|
||||
logger.debug("Root directory: %s", ROOT_DIR.resolve())
|
||||
logger.debug("Output directory: %s", ARGPARSE_DOC_DIR.resolve())
|
||||
|
||||
# Create the ARGPARSE_DOC_DIR if it doesn't exist
|
||||
if not ARGPARSE_DOC_DIR.exists():
|
||||
ARGPARSE_DOC_DIR.mkdir(parents=True)
|
||||
|
||||
# Create parsers to document
|
||||
parsers = {
|
||||
# Engine args
|
||||
"engine_args": create_parser(EngineArgs.add_cli_args),
|
||||
"async_engine_args": create_parser(
|
||||
AsyncEngineArgs.add_cli_args, async_args_only=True
|
||||
),
|
||||
# CLI
|
||||
"serve": create_parser(openai_cli_args.make_arg_parser),
|
||||
"chat": create_parser(ChatCommand.add_cli_args),
|
||||
"complete": create_parser(CompleteCommand.add_cli_args),
|
||||
"launch_render": create_parser(RenderSubcommand.add_cli_args),
|
||||
"run-batch": create_parser(openai_run_batch.make_arg_parser),
|
||||
# Benchmark CLI
|
||||
"bench_latency": create_parser(bench_latency.add_cli_args),
|
||||
"bench_mm_processor": create_parser(bench_mm_processor.add_cli_args),
|
||||
"bench_serve": create_parser(bench_serve.add_cli_args),
|
||||
"bench_sweep_plot": create_parser(bench_sweep_plot.add_cli_args),
|
||||
"bench_sweep_plot_pareto": create_parser(bench_sweep_plot_pareto.add_cli_args),
|
||||
"bench_sweep_serve": create_parser(bench_sweep_serve.add_cli_args),
|
||||
"bench_sweep_serve_workload": create_parser(
|
||||
bench_sweep_serve_workload.add_cli_args
|
||||
),
|
||||
"bench_throughput": create_parser(bench_throughput.add_cli_args),
|
||||
}
|
||||
|
||||
# Generate documentation for each parser
|
||||
for stem, parser in parsers.items():
|
||||
doc_path = ARGPARSE_DOC_DIR / f"{stem}.inc.md"
|
||||
# Specify encoding for building on Windows
|
||||
with open(doc_path, "w", encoding="utf-8") as f:
|
||||
f.write(super(type(parser), parser).format_help())
|
||||
logger.info("Argparse generated: %s", doc_path.relative_to(ROOT_DIR))
|
||||
def format_help(parser: FlexibleArgumentParser) -> str:
|
||||
"""Format a parser's help as markdown using `MarkdownFormatter`."""
|
||||
return super(type(parser), parser).format_help()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
on_startup("build", False)
|
||||
# Absolute docs URLs are kept in the help text because they are useful in the
|
||||
# terminal. Wrap them as markdown links so the `url_schemes` hook can rewrite
|
||||
# them into doc-relative links / cross-references at render time.
|
||||
_DOCS_URL = re.compile(r"https://docs\.vllm\.ai/en/[^/\s]+/[^\s)>]+")
|
||||
|
||||
|
||||
def linkify_docs_urls(text: str) -> str:
|
||||
"""Wrap bare docs.vllm.ai URLs in help text as markdown links."""
|
||||
return _DOCS_URL.sub(lambda m: f"[{m.group()}]({m.group()})", text)
|
||||
|
||||
|
||||
logger.info("Generating argparse documentation")
|
||||
logger.debug("Root directory: %s", ROOT_DIR.resolve())
|
||||
|
||||
# The JSON tip is always rendered immediately before generated argument content,
|
||||
# and the generator is its only consumer, so it lives here rather than in a
|
||||
# separate snippet file. (The runtime terminal equivalent is
|
||||
# `FlexibleArgumentParser._json_tip` in vllm/utils/argparse_utils.py.)
|
||||
JSON_TIP = """## JSON CLI Arguments
|
||||
|
||||
When passing JSON CLI arguments, the following sets of arguments are equivalent:
|
||||
|
||||
- `--json-arg '{"key1": "value1", "key2": {"key3": "value2"}}'`
|
||||
- `--json-arg.key1 value1 --json-arg.key2.key3 value2`
|
||||
|
||||
Additionally, list elements can be passed individually using `+`:
|
||||
|
||||
- `--json-arg '{"key4": ["value3", "value4", "value5"]}'`
|
||||
- `--json-arg.key4+ value3 --json-arg.key4+='value4,value5'`
|
||||
|
||||
"""
|
||||
|
||||
# Argument sections filled into `gen:` markers on handwritten pages
|
||||
engine_args = create_parser(EngineArgs.add_cli_args)
|
||||
async_engine_args = create_parser(AsyncEngineArgs.add_cli_args, async_args_only=True)
|
||||
fill_markers(
|
||||
"configuration/engine_args.md",
|
||||
{
|
||||
"engine-args": (
|
||||
f"{JSON_TIP}## `EngineArgs`\n\n"
|
||||
f"{linkify_docs_urls(format_help(engine_args))}"
|
||||
f"## `AsyncEngineArgs`\n\n"
|
||||
f"{linkify_docs_urls(format_help(async_engine_args))}"
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
# CLI reference pages generated entirely from their parser: page -> (parser, JSON tip)
|
||||
pages = {
|
||||
"cli/serve.md": (create_parser(openai_cli_args.make_arg_parser), True),
|
||||
"cli/chat.md": (create_parser(ChatCommand.add_cli_args), False),
|
||||
"cli/complete.md": (create_parser(CompleteCommand.add_cli_args), False),
|
||||
"cli/run-batch.md": (create_parser(openai_run_batch.make_arg_parser), True),
|
||||
"cli/launch/render.md": (create_parser(RenderSubcommand.add_cli_args), True),
|
||||
"cli/bench/latency.md": (create_parser(bench_latency.add_cli_args), True),
|
||||
# URL kept as `mm_processor` for back-compat; command name is `mm-processor`
|
||||
"cli/bench/mm_processor.md": (
|
||||
create_parser(BenchmarkMMProcessorSubcommand.add_cli_args),
|
||||
True,
|
||||
),
|
||||
"cli/bench/serve.md": (create_parser(bench_serve.add_cli_args), True),
|
||||
"cli/bench/startup.md": (create_parser(bench_startup.add_cli_args), True),
|
||||
"cli/bench/throughput.md": (create_parser(bench_throughput.add_cli_args), True),
|
||||
"cli/bench/sweep/plot.md": (create_parser(bench_sweep_plot.add_cli_args), True),
|
||||
"cli/bench/sweep/plot_pareto.md": (
|
||||
create_parser(bench_sweep_plot_pareto.add_cli_args),
|
||||
True,
|
||||
),
|
||||
"cli/bench/sweep/serve.md": (create_parser(bench_sweep_serve.add_cli_args), True),
|
||||
"cli/bench/sweep/serve_workload.md": (
|
||||
create_parser(bench_sweep_serve_workload.add_cli_args),
|
||||
True,
|
||||
),
|
||||
"cli/bench/sweep/startup.md": (
|
||||
create_parser(bench_sweep_startup.add_cli_args),
|
||||
True,
|
||||
),
|
||||
}
|
||||
|
||||
# Command name for pages whose file stem differs (URL kept for back-compat).
|
||||
COMMAND_NAMES = {"cli/bench/mm_processor.md": "mm-processor"}
|
||||
|
||||
for doc_path, (parser, json_tip) in pages.items():
|
||||
segments = Path(doc_path).relative_to("cli").with_suffix("").parts
|
||||
label = COMMAND_NAMES.get(doc_path, segments[-1])
|
||||
command = " ".join([*segments[:-1], label])
|
||||
# `title` frontmatter keeps the nav label to just this command's segment,
|
||||
# while the H1 stays the full `vllm ...` command for the page heading.
|
||||
content = f"---\ntitle: {label}\n---\n\n"
|
||||
content += f"# vllm {command}\n\n"
|
||||
if parser.description:
|
||||
content += f"## Overview\n\n{parser.description}\n\n"
|
||||
# Rendered above instead of at the top of the Arguments section
|
||||
parser.description = None
|
||||
if json_tip:
|
||||
content += JSON_TIP
|
||||
content += f"## Arguments\n\n{linkify_docs_urls(format_help(parser))}"
|
||||
with mkdocs_gen_files.open(doc_path, "w") as f:
|
||||
f.write(content)
|
||||
logger.debug("CLI reference generated: %s", doc_path)
|
||||
|
||||
logger.info("Total argparse docs generated: %d", len(pages) + 2)
|
||||
|
||||
|
||||
# --- Bare subcommand (group) pages -------------------------------------------
|
||||
# Mirror `vllm <group> --help`: an overview plus a table of child subcommands,
|
||||
# each linked to its reference page. Children are read from the CLI registries
|
||||
# so the listing can never drift from the actual subcommands. Each page is the
|
||||
# `README.md` of its command directory so it becomes that section's index and is
|
||||
# picked up by the existing nav globs.
|
||||
import_bench_subcommands() # populate BenchmarkSubcommandBase.__subclasses__()
|
||||
bench_subcommands = BenchmarkSubcommandBase.__subclasses__()
|
||||
bench_children = [(cmd.name, cmd.help) for cmd in bench_subcommands]
|
||||
|
||||
groups = {
|
||||
"cli/bench/README.md": (BenchmarkSubcommand.help, bench_children),
|
||||
"cli/launch/README.md": (
|
||||
launch_description,
|
||||
[(cmd.name, cmd.help) for cmd in LaunchSubcommandBase.__subclasses__()],
|
||||
),
|
||||
"cli/bench/sweep/README.md": (
|
||||
dict(bench_children).get("sweep"),
|
||||
[(args.parser_name, args.parser_help) for args, _ in sweep_subcommands],
|
||||
),
|
||||
}
|
||||
|
||||
# Doc paths that exist, so we only link a child that has a reference page.
|
||||
existing_pages = set(pages) | set(groups)
|
||||
|
||||
|
||||
def child_link(group_doc: str, name: str) -> str | None:
|
||||
group_dir = Path(group_doc).parent # cli/bench/README.md -> cli/bench
|
||||
for stem in (name, name.replace("-", "_")):
|
||||
# A leaf page (bench/latency.md) or a nested group index (sweep/README.md)
|
||||
for candidate in (group_dir / f"{stem}.md", group_dir / stem / "README.md"):
|
||||
if candidate.as_posix() in existing_pages:
|
||||
return candidate.relative_to(group_dir).as_posix()
|
||||
return None
|
||||
|
||||
|
||||
for doc_path, (overview, children) in groups.items():
|
||||
title = "vllm " + Path(doc_path).parent.relative_to("cli").as_posix()
|
||||
lines = [f"# {title.replace('/', ' ')}", ""]
|
||||
if overview:
|
||||
lines += ["## Overview", "", overview.strip(), ""]
|
||||
lines += ["## Subcommands", "", "| Command | Description |", "| --- | --- |"]
|
||||
for name, summary in children:
|
||||
link = child_link(doc_path, name)
|
||||
command = f"[`{name}`]({link})" if link else f"`{name}`"
|
||||
lines.append(f"| {command} | {(summary or '').strip()} |")
|
||||
with mkdocs_gen_files.open(doc_path, "w") as f:
|
||||
f.write("\n".join(lines) + "\n")
|
||||
logger.debug("CLI group reference generated: %s", doc_path)
|
||||
|
||||
logger.info("CLI group reference pages generated: %d", len(groups))
|
||||
+61
-404
@@ -9,33 +9,28 @@ based on the checks in AttentionBackend.validate_configuration().
|
||||
|
||||
This approach avoids requiring CUDA/ROCm/GPU libraries to be installed.
|
||||
|
||||
When used as a pre-commit hook, this script receives filenames as arguments
|
||||
and only runs the check if any of the relevant files were modified.
|
||||
It runs as an mkdocs-gen-files script, so the page is generated at docs build
|
||||
time rather than being committed to the repository.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import ast
|
||||
import fnmatch
|
||||
import logging
|
||||
import sys
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
|
||||
from generated_content import fill_markers # noqa: E402
|
||||
|
||||
logger = logging.getLogger("mkdocs")
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants and file paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
REPO_ROOT = Path(__file__).parent.parent.parent
|
||||
|
||||
RELEVANT_PATTERNS = [
|
||||
"vllm/v1/attention/backends/*.py",
|
||||
"vllm/v1/attention/backends/**/*.py",
|
||||
"vllm/models/minimax_m3/common/sparse_attention.py",
|
||||
"vllm/model_executor/layers/attention/mla_attention.py",
|
||||
"vllm/platforms/cuda.py",
|
||||
"tools/pre_commit/generate_attention_backend_docs.py",
|
||||
"docs/design/attention_backends.md",
|
||||
]
|
||||
REPO_ROOT = Path(__file__).parent.parent.parent.parent
|
||||
|
||||
BACKENDS_DIR = REPO_ROOT / "vllm" / "v1" / "attention" / "backends"
|
||||
REGISTRY_FILE = BACKENDS_DIR / "registry.py"
|
||||
@@ -55,19 +50,6 @@ BACKEND_KV_DTYPE_EXCLUDES: dict[str, set[str]] = {
|
||||
}
|
||||
|
||||
|
||||
def is_relevant_file(filepath: str) -> bool:
|
||||
"""Check if a file matches any of the relevant patterns."""
|
||||
path = Path(filepath)
|
||||
if path.is_absolute():
|
||||
try:
|
||||
path = path.relative_to(REPO_ROOT)
|
||||
except ValueError:
|
||||
return False
|
||||
path_str = str(path)
|
||||
|
||||
return any(fnmatch.fnmatch(path_str, pattern) for pattern in RELEVANT_PATTERNS)
|
||||
|
||||
|
||||
MLA_PREFILL_DIR = BACKENDS_DIR / "mla" / "prefill"
|
||||
MLA_PREFILL_REGISTRY_FILE = MLA_PREFILL_DIR / "registry.py"
|
||||
MLA_PREFILL_SELECTOR_FILE = MLA_PREFILL_DIR / "selector.py"
|
||||
@@ -960,7 +942,7 @@ def analyze_backend(backend_name: str, class_path: str) -> dict[str, Any] | None
|
||||
try:
|
||||
tree = ast.parse(file_path.read_text())
|
||||
except Exception as e:
|
||||
print(f" Warning: Could not parse {file_path}: {e}", file=sys.stderr)
|
||||
logger.warning("Could not parse %s: %s", file_path, e)
|
||||
return None
|
||||
|
||||
class_name = class_path.rsplit(".", 1)[1]
|
||||
@@ -1657,113 +1639,12 @@ def _render_table(
|
||||
return lines
|
||||
|
||||
|
||||
def generate_markdown_table(
|
||||
backends: list[dict[str, Any]], title: str, is_mla_table: bool = False
|
||||
) -> str:
|
||||
"""Generate a titled markdown table from backend info."""
|
||||
if not backends:
|
||||
return f"## {title}\n\nNo backends found.\n"
|
||||
has_versions = any(b.get("version") for b in backends)
|
||||
columns = _build_columns(is_mla_table, has_versions)
|
||||
lines = [f"## {title}", ""]
|
||||
lines.extend(_render_table(columns, backends))
|
||||
lines.append("")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Markdown section generators (usage, priority, legend, MLA)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def generate_usage_section() -> str:
|
||||
"""Generate the usage documentation section."""
|
||||
return """## Setting the Attention Backend
|
||||
|
||||
### Command Line
|
||||
|
||||
There are two ways to specify the backend from the command line:
|
||||
|
||||
**Option 1: Using `--attention-backend` (simple)**
|
||||
|
||||
```bash
|
||||
vllm serve <model> --attention-backend FLASH_ATTN
|
||||
```
|
||||
|
||||
**Option 2: Using `--attention-config.backend` / `-ac.backend` (structured config)**
|
||||
|
||||
```bash
|
||||
# Dot notation
|
||||
vllm serve <model> --attention-config.backend FLASH_ATTN
|
||||
vllm serve <model> -ac.backend FLASH_ATTN
|
||||
|
||||
# JSON format
|
||||
vllm serve <model> --attention-config '{"backend": "FLASH_ATTN"}'
|
||||
vllm serve <model> -ac '{"backend": "FLASH_ATTN"}'
|
||||
```
|
||||
|
||||
> **Note:** `--attention-backend` and `--attention-config.backend` are mutually
|
||||
> exclusive. Use one or the other, not both.
|
||||
|
||||
### Python API
|
||||
|
||||
Use `AttentionConfig` with the `LLM` class:
|
||||
|
||||
```python
|
||||
from vllm import LLM
|
||||
from vllm.config import AttentionConfig
|
||||
from vllm.v1.attention.backends.registry import AttentionBackendEnum
|
||||
|
||||
# Method 1: Using AttentionConfig with enum
|
||||
llm = LLM(
|
||||
model="Qwen/Qwen3-0.6B",
|
||||
attention_config=AttentionConfig(backend=AttentionBackendEnum.FLASH_ATTN),
|
||||
)
|
||||
|
||||
# Method 2: Using attention_backend parameter with string
|
||||
llm = LLM(
|
||||
model="Qwen/Qwen3-0.6B",
|
||||
attention_backend="FLASH_ATTN",
|
||||
)
|
||||
```
|
||||
|
||||
## Backend Selection Behavior
|
||||
|
||||
### Manual Selection
|
||||
|
||||
When you explicitly set a backend via `--attention-backend` or `AttentionConfig`:
|
||||
|
||||
1. The backend is **validated** against your configuration (model dtype, head
|
||||
size, compute capability, etc.)
|
||||
2. If the backend **doesn't support** your configuration, an error is raised
|
||||
with the specific reason
|
||||
3. If valid, the backend is used
|
||||
|
||||
Example error when selecting an incompatible backend:
|
||||
|
||||
```text
|
||||
ValueError: Selected backend FLASHMLA is not valid for this configuration.
|
||||
Reason: ['compute capability not supported']
|
||||
```
|
||||
|
||||
### Automatic Selection
|
||||
|
||||
When no backend is specified (the default):
|
||||
|
||||
1. vLLM iterates through backends in **priority order** (see tables below)
|
||||
2. Each backend is validated against your configuration
|
||||
3. The **first compatible backend** is selected
|
||||
4. If no backend is compatible, an error is raised listing all backends and
|
||||
their incompatibility reasons
|
||||
"""
|
||||
|
||||
|
||||
def _priority_table(
|
||||
title: str,
|
||||
backends: list[str],
|
||||
annotations: dict[str, str] | None = None,
|
||||
) -> list[str]:
|
||||
"""Generate a priority table for a list of backends."""
|
||||
"""Render a priority table for a list of backends."""
|
||||
|
||||
def _fmt(b: str) -> str:
|
||||
suffix = annotations.get(b, "") if annotations else ""
|
||||
@@ -1779,102 +1660,38 @@ def _priority_table(
|
||||
]
|
||||
|
||||
|
||||
def generate_priority_section(priorities: dict[str, list[str]]) -> str:
|
||||
"""Generate the priority ranking section."""
|
||||
lines = [
|
||||
"## Backend Priority (CUDA)",
|
||||
"",
|
||||
"When no backend is explicitly selected, vLLM chooses the first",
|
||||
"compatible backend from these priority-ordered lists.",
|
||||
"",
|
||||
"Priority is **1 = highest** (tried first).",
|
||||
"",
|
||||
"### Standard Attention (MHA, MQA, GQA)",
|
||||
"",
|
||||
]
|
||||
|
||||
sm100 = "Blackwell (SM 10.x)"
|
||||
ampere = "Ampere/Hopper (SM 8.x-9.x)"
|
||||
|
||||
if "standard_sm100" in priorities:
|
||||
lines.extend(_priority_table(sm100, priorities["standard_sm100"]))
|
||||
if "standard_default" in priorities:
|
||||
lines.extend(_priority_table(ampere, priorities["standard_default"]))
|
||||
|
||||
lines.extend(["### MLA Attention (DeepSeek-style)", ""])
|
||||
|
||||
mla_sm100_annotations = {
|
||||
"FLASHINFER_MLA_SPARSE": "**\\***",
|
||||
}
|
||||
if "mla_sm100" in priorities:
|
||||
lines.extend(
|
||||
_priority_table(sm100, priorities["mla_sm100"], mla_sm100_annotations)
|
||||
)
|
||||
if "mla_default" in priorities:
|
||||
lines.extend(_priority_table(ampere, priorities["mla_default"]))
|
||||
|
||||
if "mla_sm100" in priorities:
|
||||
lines.append(
|
||||
"> **\\*** For sparse MLA, FP8 KV cache always prefers "
|
||||
"`FLASHINFER_MLA_SPARSE`. With BF16 KV cache, `FLASHINFER_MLA_SPARSE` "
|
||||
"is preferred for low query-head counts (<= 16), while "
|
||||
"`FLASHMLA_SPARSE` is preferred otherwise."
|
||||
)
|
||||
lines.append(">")
|
||||
|
||||
lines.append(
|
||||
"> **Note:** ROCm and CPU platforms have their own selection logic. "
|
||||
"See the platform-specific documentation for details."
|
||||
)
|
||||
lines.append("")
|
||||
|
||||
return "\n".join(lines)
|
||||
_SM100 = "Blackwell (SM 10.x)"
|
||||
_AMPERE = "Ampere/Hopper (SM 8.x-9.x)"
|
||||
|
||||
|
||||
def generate_legend() -> str:
|
||||
"""Generate a legend explaining the table columns."""
|
||||
return """## Legend
|
||||
|
||||
| Column | Description |
|
||||
| ------ | ----------- |
|
||||
| **Dtypes** | Supported model data types (fp16, bf16, fp32) |
|
||||
| **KV Dtypes** | Supported KV cache data types (`auto`, `fp8`, `fp8_e4m3`, etc.) |
|
||||
| **Block Sizes** | Supported KV cache block sizes (%N means multiples of N) |
|
||||
| **Head Sizes** | Supported attention head sizes |
|
||||
| **Sink** | Attention sink support (for StreamingLLM) |
|
||||
| **Non-Causal** | Non-causal (bidirectional) attention support for decoder models |
|
||||
| **Sparse** | Sparse attention support (MLA only) |
|
||||
| **MM Prefix** | Multimodal prefix full attention support |
|
||||
| **DCP** | Decode Context Parallelism support (`--decode-context-parallel-size`) |
|
||||
| **Attention Types** | Supported attention patterns (Decoder, Encoder, Enc-Dec) |
|
||||
| **Compute Cap.** | Required CUDA compute capability (N/A for non-CUDA backends) |
|
||||
|
||||
**Symbols:** ✅ = Supported, ❌ = Not supported
|
||||
"""
|
||||
|
||||
|
||||
def generate_mla_section(
|
||||
prefill_backends: list[dict[str, Any]],
|
||||
decode_backends: list[dict[str, Any]],
|
||||
v4_decode_backends: list[dict[str, Any]] | None = None,
|
||||
def _priority_block(
|
||||
priorities: dict[str, list[str]],
|
||||
sm100_key: str,
|
||||
default_key: str,
|
||||
sm100_annotations: dict[str, str] | None = None,
|
||||
) -> str:
|
||||
"""Generate the complete MLA section with prefill and decode tables."""
|
||||
"""Render whichever priority tables exist for one attention category."""
|
||||
lines: list[str] = []
|
||||
if sm100_key in priorities:
|
||||
lines += _priority_table(_SM100, priorities[sm100_key], sm100_annotations)
|
||||
if default_key in priorities:
|
||||
lines += _priority_table(_AMPERE, priorities[default_key])
|
||||
return "\n".join(lines).strip()
|
||||
|
||||
|
||||
def _feature_table(backends: list[dict[str, Any]], is_mla: bool) -> str:
|
||||
"""Render a backend feature table (header, separator, one row per backend)."""
|
||||
has_versions = any(b.get("version") for b in backends)
|
||||
columns = _build_columns(is_mla, has_versions)
|
||||
return "\n".join(_render_table(columns, backends))
|
||||
|
||||
|
||||
def _mla_prefill_table(prefill_backends: list[dict[str, Any]]) -> str:
|
||||
"""Render the MLA prefill backend table."""
|
||||
lines = [
|
||||
"## MLA (Multi-head Latent Attention) Backends",
|
||||
"",
|
||||
"MLA uses separate backends for prefill and decode phases.",
|
||||
"",
|
||||
"### Prefill Backends",
|
||||
"",
|
||||
"To explicitly select a prefill backend, use",
|
||||
"`-ac.mla_prefill_backend=<BACKEND>` (e.g., `FLASH_ATTN`, `FLASHINFER`).",
|
||||
"Otherwise, the prefill backend is selected automatically at runtime based on",
|
||||
"hardware and configuration.",
|
||||
"",
|
||||
"| Backend | Description | Dtypes | Compute Cap. | Notes |",
|
||||
"| ------- | ----------- | ------ | ------------ | ----- |",
|
||||
]
|
||||
|
||||
for backend in prefill_backends:
|
||||
row = "| `{}`{} | {} | {} | {} | {} |".format(
|
||||
backend["name"],
|
||||
@@ -1885,87 +1702,21 @@ def generate_mla_section(
|
||||
backend.get("notes", ""),
|
||||
)
|
||||
lines.append(row.replace(" ", " "))
|
||||
|
||||
lines.extend(
|
||||
[
|
||||
"",
|
||||
"> **‡** Automatic selection tries FlashAttention first. On Blackwell",
|
||||
"> (SM100), the fallback order is TRT-LLM Ragged, FlashInfer, then",
|
||||
"> TokenSpeed MLA. On other GPUs, only FlashAttention is considered.",
|
||||
"",
|
||||
"### Decode Backends",
|
||||
"",
|
||||
"MLA decode backends are selected using the standard",
|
||||
"`-ac.backend=<BACKEND>` argument (e.g., `FLASHMLA`, `TRITON_MLA`).",
|
||||
"",
|
||||
]
|
||||
)
|
||||
|
||||
# Reuse data-driven table rendering for decode backends
|
||||
columns = _build_columns(is_mla=True, has_versions=False)
|
||||
lines.extend(_render_table(columns, decode_backends))
|
||||
|
||||
if v4_decode_backends:
|
||||
lines.extend(
|
||||
[
|
||||
"",
|
||||
"### DeepSeek V4 Decode Backends",
|
||||
"",
|
||||
"DeepSeek V4 sparse MLA uses its own decode backends, selected via",
|
||||
"`--attention-backend=<BACKEND>` (e.g., `FLASHMLA_SPARSE_DSV4`,",
|
||||
"`FLASHINFER_MLA_SPARSE_DSV4`). They share the V4 sparse-index",
|
||||
"pipeline (compressor + SWA + indexer, 256-token blocks, head 512);",
|
||||
"default on NVIDIA is `FLASHINFER_MLA_SPARSE_DSV4` on SM12x and",
|
||||
"`FLASHMLA_SPARSE_DSV4` on other supported CUDA architectures.",
|
||||
"",
|
||||
]
|
||||
)
|
||||
lines.extend(_render_table(columns, v4_decode_backends))
|
||||
|
||||
lines.append("")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def generate_minimax_section(backends: list[dict[str, Any]]) -> str:
|
||||
"""Generate the MiniMax M3 sparse attention section."""
|
||||
lines = [
|
||||
"## MiniMax M3 Sparse Attention Backends",
|
||||
"",
|
||||
'Block-sparse GQA backend used by MiniMax M3 sparse ("lightning indexer")',
|
||||
"layers. It is wired in directly by the model and is not part of the",
|
||||
"automatic priority lists above. A lightning indexer scores KV blocks, the",
|
||||
"top-k blocks (plus fixed init/local blocks) are selected, and attention",
|
||||
"attends only to those blocks; index keys live in a separate side cache.",
|
||||
"",
|
||||
]
|
||||
columns = _build_columns(is_mla=False, has_versions=False)
|
||||
lines.extend(_render_table(columns, backends))
|
||||
lines.append("")
|
||||
return "\n".join(lines)
|
||||
def build_blocks() -> dict[str, str]:
|
||||
"""Build the generated table blocks keyed by their `gen:` marker name.
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Top-level orchestration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def generate_docs() -> str:
|
||||
"""Generate the complete documentation."""
|
||||
Only the tables are generated here; the surrounding prose lives in the
|
||||
handwritten ``docs/design/attention_backends.md`` page.
|
||||
"""
|
||||
attention_backends_map = parse_registry()
|
||||
|
||||
# Parse priority lists from cuda.py
|
||||
priorities = parse_cuda_priority_lists()
|
||||
|
||||
# Parse FlashAttention FA2/FA3 feature differences
|
||||
fa_features = parse_flash_attn_features()
|
||||
|
||||
# Parse FlashInfer TRTLLM feature differences (native vs TRTLLM on Blackwell)
|
||||
fi_features = parse_flashinfer_trtllm_features()
|
||||
|
||||
# Parse MLA prefill backends
|
||||
mla_prefill_backends = parse_mla_prefill_backends()
|
||||
|
||||
# Collect backend info
|
||||
all_backends = []
|
||||
for backend_name, class_path in attention_backends_map.items():
|
||||
if backend_name in SKIP_BACKENDS:
|
||||
@@ -1973,17 +1724,14 @@ def generate_docs() -> str:
|
||||
info = analyze_backend(backend_name, class_path)
|
||||
if info:
|
||||
all_backends.append(info)
|
||||
|
||||
# Expand backends into version variants
|
||||
if fa_features:
|
||||
all_backends = _expand_flash_attn_variants(all_backends, fa_features)
|
||||
if fi_features:
|
||||
all_backends = _expand_flashinfer_variants(all_backends, fi_features)
|
||||
|
||||
# DeepSeek V4 (*_DSV4) decode backends and MiniMax M3 sparse backends each
|
||||
# get their own subsection rather than mixing into the main MLA / standard
|
||||
# tables (the ROCm V4 backend isn't flagged is_mla by the AST heuristic, so
|
||||
# filter purely on the name).
|
||||
# DeepSeek V4 (*_DSV4) and MiniMax M3 sparse backends get their own tables
|
||||
# rather than mixing into the main MLA / standard tables (the ROCm V4 backend
|
||||
# isn't flagged is_mla by the AST heuristic, so filter purely on the name).
|
||||
def _is_v4(b: dict[str, Any]) -> bool:
|
||||
return b["name"].endswith("_DSV4")
|
||||
|
||||
@@ -1999,112 +1747,21 @@ def generate_docs() -> str:
|
||||
if not b["is_mla"] and not _is_v4(b) and not _is_minimax(b)
|
||||
]
|
||||
|
||||
# Generate documentation
|
||||
script_path = "tools/pre_commit/generate_attention_backend_docs.py"
|
||||
doc_lines = [
|
||||
"# Attention Backend Feature Support",
|
||||
"",
|
||||
f"This document is auto-generated by `{script_path}`.",
|
||||
"It shows the feature support for each registered attention backend",
|
||||
"based on the checks in `AttentionBackend.validate_configuration()`.",
|
||||
"",
|
||||
"**Do not edit this file manually.** Run the following command to",
|
||||
"regenerate it:",
|
||||
"",
|
||||
"```bash",
|
||||
f"python {script_path}",
|
||||
"```",
|
||||
"",
|
||||
]
|
||||
|
||||
# Add usage documentation
|
||||
doc_lines.append(generate_usage_section())
|
||||
|
||||
# Add priority section
|
||||
doc_lines.append(generate_priority_section(priorities))
|
||||
|
||||
# Add legend and feature tables
|
||||
doc_lines.append(generate_legend())
|
||||
standard_title = "Standard Attention (MHA, MQA, GQA) Backends"
|
||||
doc_lines.append(
|
||||
generate_markdown_table(non_mla_backends, standard_title, is_mla_table=False)
|
||||
)
|
||||
# Add footnotes for version/variant distinctions (in table order)
|
||||
footnotes = []
|
||||
if fi_features:
|
||||
footnotes.append(
|
||||
"> **†** FlashInfer Native is the regular FlashInfer path. XQA is the "
|
||||
"SM90 decode path exposed through FlashInfer's TRTLLM decode API. "
|
||||
"trtllm-gen is used on SM100 and supports sinks. Disable XQA/trtllm-gen "
|
||||
"via `--attention-config.use_trtllm_attention=0`."
|
||||
)
|
||||
if fa_features:
|
||||
footnotes.append(
|
||||
"> **\\*** Specify the FlashAttention version via "
|
||||
"`--attention-config.flash_attn_version=2`, `3`, or `4`. "
|
||||
"Default is FA4 on SM100+ (Blackwell), FA3 on SM90 (Hopper), "
|
||||
"FA2 otherwise."
|
||||
)
|
||||
if footnotes:
|
||||
doc_lines.append("\n>\n".join(footnotes) + "\n")
|
||||
|
||||
# Add MiniMax M3 sparse section (separate category after standard GQA)
|
||||
if minimax_backends:
|
||||
doc_lines.append(generate_minimax_section(minimax_backends))
|
||||
|
||||
# Add MLA section with prefill and decode backends
|
||||
doc_lines.append(
|
||||
generate_mla_section(mla_prefill_backends, mla_backends, v4_decode_backends)
|
||||
)
|
||||
|
||||
return "\n".join(doc_lines)
|
||||
mla_sm100_annotations = {"FLASHINFER_MLA_SPARSE": "**\\***"}
|
||||
return {
|
||||
"priority-standard": _priority_block(
|
||||
priorities, "standard_sm100", "standard_default"
|
||||
),
|
||||
"priority-mla": _priority_block(
|
||||
priorities, "mla_sm100", "mla_default", mla_sm100_annotations
|
||||
),
|
||||
"table-standard": _feature_table(non_mla_backends, is_mla=False),
|
||||
"table-minimax": _feature_table(minimax_backends, is_mla=False),
|
||||
"table-mla-prefill": _mla_prefill_table(mla_prefill_backends),
|
||||
"table-mla-decode": _feature_table(mla_backends, is_mla=True),
|
||||
"table-mla-v4-decode": _feature_table(v4_decode_backends, is_mla=True),
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Generate attention backend documentation table"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
"-o",
|
||||
type=str,
|
||||
default=str(REPO_ROOT / "docs" / "design" / "attention_backends.md"),
|
||||
help="Output file path (default: docs/design/attention_backends.md)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--check",
|
||||
action="store_true",
|
||||
help="Check if the documentation is up to date (for pre-commit)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"files",
|
||||
nargs="*",
|
||||
help="Files to check (passed by pre-commit). If none are relevant, skip.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.files and not any(is_relevant_file(f) for f in args.files):
|
||||
sys.exit(0)
|
||||
|
||||
output_path = Path(args.output)
|
||||
new_content = generate_docs()
|
||||
|
||||
if args.check:
|
||||
needs_update = (
|
||||
not output_path.exists() or output_path.read_text() != new_content
|
||||
)
|
||||
if needs_update:
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
output_path.write_text(new_content)
|
||||
print(f"🔄 Regenerated: {output_path}")
|
||||
sys.exit(1)
|
||||
print(f"✅ Up to date: {output_path}")
|
||||
sys.exit(0)
|
||||
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
output_path.write_text(new_content)
|
||||
print(f"Generated: {output_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
logger.info("Generating attention backend documentation")
|
||||
fill_markers("design/attention_backends.md", build_blocks())
|
||||
+34
-41
@@ -5,16 +5,15 @@ import logging
|
||||
from dataclasses import dataclass
|
||||
from functools import cached_property
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
import mkdocs_awesome_nav.nav.directory as _nav_dir
|
||||
import mkdocs_gen_files
|
||||
import regex as re
|
||||
|
||||
logger = logging.getLogger("mkdocs")
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.parent
|
||||
ROOT_DIR_RELATIVE = "../../../../.."
|
||||
EXAMPLE_DIR = ROOT_DIR / "examples"
|
||||
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/examples"
|
||||
|
||||
|
||||
def title(text: str) -> str:
|
||||
@@ -197,44 +196,38 @@ class Example:
|
||||
return content
|
||||
|
||||
|
||||
def on_startup(command: Literal["build", "gh-deploy", "serve"], dirty: bool):
|
||||
# Monkey-patch dirname_to_title in awesome-nav so that sub-directory names are
|
||||
# title-cased (e.g. "Offline Inference" instead of "Offline inference").
|
||||
import mkdocs_awesome_nav.nav.directory as _nav_dir
|
||||
# Monkey-patch dirname_to_title in awesome-nav so that sub-directory names are
|
||||
# title-cased (e.g. "Offline Inference" instead of "Offline inference").
|
||||
_nav_dir.dirname_to_title = title
|
||||
logger.info("Generating example documentation")
|
||||
logger.debug("Root directory: %s", ROOT_DIR.resolve())
|
||||
logger.debug("Example directory: %s", EXAMPLE_DIR.resolve())
|
||||
|
||||
_nav_dir.dirname_to_title = title
|
||||
logger.info("Generating example documentation")
|
||||
logger.debug("Root directory: %s", ROOT_DIR.resolve())
|
||||
logger.debug("Example directory: %s", EXAMPLE_DIR.resolve())
|
||||
logger.debug("Example document directory: %s", EXAMPLE_DOC_DIR.resolve())
|
||||
categories = sorted(
|
||||
p for p in EXAMPLE_DIR.iterdir() if p.is_dir() and not p.name.startswith(".")
|
||||
)
|
||||
|
||||
# Create the EXAMPLE_DOC_DIR if it doesn't exist
|
||||
if not EXAMPLE_DOC_DIR.exists():
|
||||
EXAMPLE_DOC_DIR.mkdir(parents=True)
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
# Find categorised examples
|
||||
for category in categories:
|
||||
logger.info("Processing category: %s", category.stem)
|
||||
globs = [category.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path, category.stem))
|
||||
# Find examples in subdirectories
|
||||
globs = [category.glob(f"*/{pattern}") for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path.parent, category.stem))
|
||||
|
||||
categories = sorted(p for p in EXAMPLE_DIR.iterdir() if p.is_dir())
|
||||
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
# Find categorised examples
|
||||
for category in categories:
|
||||
logger.info("Processing category: %s", category.stem)
|
||||
globs = [category.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path, category.stem))
|
||||
# Find examples in subdirectories
|
||||
globs = [category.glob(f"*/{pattern}") for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path.parent, category.stem))
|
||||
|
||||
# Generate the example documentation
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
example_name = f"{example.path.stem}.md"
|
||||
doc_path = EXAMPLE_DOC_DIR / example.category / example_name
|
||||
if not doc_path.parent.exists():
|
||||
doc_path.parent.mkdir(parents=True)
|
||||
# Specify encoding for building on Windows
|
||||
with open(doc_path, "w+", encoding="utf-8") as f:
|
||||
f.write(example.generate())
|
||||
logger.debug("Example generated: %s", doc_path.relative_to(ROOT_DIR))
|
||||
logger.info("Total examples generated: %d", len(examples))
|
||||
# Generate the example documentation
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
doc_path = f"examples/{example.category}/{example.path.stem}.md"
|
||||
with mkdocs_gen_files.open(doc_path, "w") as f:
|
||||
f.write(example.generate())
|
||||
if example.main_file is not None:
|
||||
# Point the edit button at the example's source file
|
||||
edit_path = Path("..") / example.main_file.relative_to(ROOT_DIR)
|
||||
mkdocs_gen_files.set_edit_path(doc_path, str(edit_path))
|
||||
logger.debug("Example generated: %s", doc_path)
|
||||
logger.info("Total examples generated: %d", len(examples))
|
||||
@@ -2,27 +2,28 @@
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import ast
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
|
||||
from generated_content import fill_markers # noqa: E402
|
||||
|
||||
logger = logging.getLogger("mkdocs")
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.parent
|
||||
DOCS_DIR = ROOT_DIR / "docs"
|
||||
GENERATED_METRICS_DIR = DOCS_DIR / "generated" / "metrics"
|
||||
|
||||
# Files to scan for metric definitions - each will generate a separate table
|
||||
# Files to scan for metric definitions - each fills a `gen:` marker in
|
||||
# docs/usage/metrics.md with its table (the section heading and any preamble
|
||||
# live in the tracked page next to the marker).
|
||||
METRIC_SOURCE_FILES = [
|
||||
{"path": "vllm/v1/metrics/loggers.py", "output": "general.inc.md"},
|
||||
{
|
||||
"path": "vllm/v1/spec_decode/metrics.py",
|
||||
"output": "spec_decode.inc.md",
|
||||
},
|
||||
{"path": "vllm/v1/metrics/loggers.py", "key": "metrics-general"},
|
||||
{"path": "vllm/v1/spec_decode/metrics.py", "key": "metrics-spec-decode"},
|
||||
{
|
||||
"path": "vllm/distributed/kv_transfer/kv_connector/v1/nixl/stats.py",
|
||||
"output": "nixl_connector.inc.md",
|
||||
"key": "metrics-nixl",
|
||||
},
|
||||
{"path": "vllm/v1/metrics/perf.py", "output": "perf.inc.md"},
|
||||
{"path": "vllm/v1/metrics/perf.py", "key": "metrics-mfu"},
|
||||
]
|
||||
|
||||
|
||||
@@ -110,41 +111,27 @@ def generate_markdown_table(metrics: list[dict[str, str]]) -> str:
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
|
||||
def on_startup(command: Literal["build", "gh-deploy", "serve"], dirty: bool):
|
||||
"""Generate metrics documentation tables from source files."""
|
||||
logger.info("Generating metrics documentation")
|
||||
logger.info("Generating metrics documentation")
|
||||
|
||||
# Create generated directory if it doesn't exist
|
||||
GENERATED_METRICS_DIR.mkdir(parents=True, exist_ok=True)
|
||||
blocks = {}
|
||||
total_metrics = 0
|
||||
for source_config in METRIC_SOURCE_FILES:
|
||||
source_path = source_config["path"]
|
||||
|
||||
total_metrics = 0
|
||||
for source_config in METRIC_SOURCE_FILES:
|
||||
source_path = source_config["path"]
|
||||
output_file = source_config["output"]
|
||||
filepath = ROOT_DIR / source_path
|
||||
if not filepath.exists():
|
||||
raise FileNotFoundError(f"Metrics source file not found: {filepath}")
|
||||
|
||||
filepath = ROOT_DIR / source_path
|
||||
if not filepath.exists():
|
||||
raise FileNotFoundError(f"Metrics source file not found: {filepath}")
|
||||
logger.debug("Extracting metrics from: %s", source_path)
|
||||
metrics = extract_metrics_from_file(filepath)
|
||||
logger.debug("Found %d metrics in %s", len(metrics), source_path)
|
||||
|
||||
logger.debug("Extracting metrics from: %s", source_path)
|
||||
metrics = extract_metrics_from_file(filepath)
|
||||
logger.debug("Found %d metrics in %s", len(metrics), source_path)
|
||||
blocks[source_config["key"]] = generate_markdown_table(metrics).strip()
|
||||
total_metrics += len(metrics)
|
||||
|
||||
# Generate and write the markdown table for this source
|
||||
table_content = generate_markdown_table(metrics)
|
||||
output_path = GENERATED_METRICS_DIR / output_file
|
||||
with open(output_path, "w", encoding="utf-8") as f:
|
||||
f.write(table_content)
|
||||
|
||||
total_metrics += len(metrics)
|
||||
logger.info(
|
||||
"Generated metrics table: %s (%d metrics)",
|
||||
output_path.relative_to(ROOT_DIR),
|
||||
len(metrics),
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Total metrics generated: %d across %d files",
|
||||
total_metrics,
|
||||
len(METRIC_SOURCE_FILES),
|
||||
)
|
||||
fill_markers("usage/metrics.md", blocks)
|
||||
logger.info(
|
||||
"Total metrics generated: %d across %d files",
|
||||
total_metrics,
|
||||
len(METRIC_SOURCE_FILES),
|
||||
)
|
||||
@@ -0,0 +1,56 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Inline build-time generated content into existing docs pages.
|
||||
|
||||
Source pages mark where generated content goes with a snippet-style marker,
|
||||
`--8<-- "gen:<key>"`, so the insertion point is explicit and readable. The
|
||||
substitution happens here (at gen-files time, before mkdocs-gen-files shadows
|
||||
the page), not via pymdownx.snippets, so the content can be generated at build
|
||||
time without living in a real file on disk.
|
||||
|
||||
The `gen:` prefix keeps these markers distinct from real pymdownx.snippets
|
||||
includes, and `fill_markers` fails loudly if a marker is missing or left behind
|
||||
(pymdownx.snippets would otherwise silently drop an unsubstituted marker).
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import mkdocs_gen_files
|
||||
import regex as re
|
||||
|
||||
DOCS_DIR = Path(__file__).parent.parent.parent
|
||||
|
||||
_MARKER = '--8<-- "gen:{key}"'
|
||||
_ANY_MARKER = re.compile(r'--8<-- "gen:[^"]*"')
|
||||
|
||||
|
||||
def fill_markers(doc_path: str, blocks: dict[str, str]) -> None:
|
||||
"""Replace `--8<-- "gen:<key>"` markers in a docs page with generated content.
|
||||
|
||||
Args:
|
||||
doc_path: Docs-relative path of the source page to fill.
|
||||
blocks: Mapping of marker key to the markdown to insert in its place.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the source page does not exist.
|
||||
ValueError: If an expected marker is missing, or any `gen:` marker is
|
||||
left unsubstituted after filling.
|
||||
"""
|
||||
source = DOCS_DIR / doc_path
|
||||
if not source.exists():
|
||||
raise FileNotFoundError(f"Cannot fill markers in missing page: {doc_path}")
|
||||
|
||||
text = source.read_text()
|
||||
for key, content in blocks.items():
|
||||
marker = _MARKER.format(key=key)
|
||||
if marker not in text:
|
||||
raise ValueError(f"{doc_path}: missing marker {marker}")
|
||||
text = text.replace(marker, content)
|
||||
|
||||
if leftover := _ANY_MARKER.search(text):
|
||||
raise ValueError(f"{doc_path}: unsubstituted marker {leftover.group()}")
|
||||
|
||||
with mkdocs_gen_files.open(doc_path, "w") as f:
|
||||
f.write(text)
|
||||
# Keep the edit button pointing at the real source page
|
||||
mkdocs_gen_files.set_edit_path(doc_path, doc_path)
|
||||
@@ -19,6 +19,7 @@ The on_page_markdown hook passes the current page context to the preprocessor be
|
||||
each page is converted.
|
||||
"""
|
||||
|
||||
import posixpath
|
||||
from pathlib import Path
|
||||
|
||||
import regex as re
|
||||
@@ -38,18 +39,22 @@ TITLE = r"(?P<title>[^\[\]<>]+?)"
|
||||
REPO = r"(?P<repo>.+?/.+?)"
|
||||
TYPE = r"(?P<type>issues|pull|projects)"
|
||||
NUMBER = r"(?P<number>\d+)"
|
||||
VERSION = r"[^/\s]+"
|
||||
PATH = r"(?P<path>[^\s]+?)"
|
||||
FRAGMENT = r"(?P<fragment>#[^\s]+)?"
|
||||
URL = f"https://github.com/{REPO}/{TYPE}/{NUMBER}{FRAGMENT}"
|
||||
URL_GITHUB = f"https://github.com/{REPO}/{TYPE}/{NUMBER}{FRAGMENT}"
|
||||
RELATIVE = rf"(?!(https?|ftp)://|#){PATH}{FRAGMENT}"
|
||||
URL_DOCS = f"https://docs.vllm.ai/en/{VERSION}/{PATH}{FRAGMENT}"
|
||||
|
||||
# Common titles to use for GitHub links when none is provided in the link.
|
||||
TITLES = {"issues": "Issue ", "pull": "Pull Request ", "projects": "Project "}
|
||||
|
||||
# Regex to match GitHub issue, PR, and project links with optional titles.
|
||||
github_link = re.compile(rf"(\[{TITLE}\]\(|<){URL}(\)|>)")
|
||||
github_link = re.compile(rf"(\[{TITLE}\]\(|<){URL_GITHUB}(\)|>)")
|
||||
# Regex to match relative file links with optional titles.
|
||||
relative_link = re.compile(rf"\[{TITLE}\]\({RELATIVE}\)")
|
||||
# Regex to match absolute docs.vllm.ai links (should only exist in CLI).
|
||||
docs_link = re.compile(rf"\[{TITLE}\]\({URL_DOCS}\)")
|
||||
|
||||
|
||||
class UrlSchemesPreprocessor(Preprocessor):
|
||||
@@ -61,7 +66,8 @@ class UrlSchemesPreprocessor(Preprocessor):
|
||||
|
||||
def run(self, lines):
|
||||
page = self.ext.page
|
||||
if page is None or getattr(page.file, "abs_src_path", None) is None:
|
||||
files = self.ext.files
|
||||
if page is None:
|
||||
return lines
|
||||
|
||||
def replace_relative_link(match: re.Match) -> str:
|
||||
@@ -70,7 +76,7 @@ class UrlSchemesPreprocessor(Preprocessor):
|
||||
"""
|
||||
title = match.group("title")
|
||||
path = match.group("path")
|
||||
path = (Path(page.file.abs_src_path).parent / path).resolve()
|
||||
path = ((DOC_DIR / page.file.src_uri).parent / path).resolve()
|
||||
fragment = match.group("fragment") or ""
|
||||
|
||||
# Check if the path exists and is outside the docs dir
|
||||
@@ -105,9 +111,36 @@ class UrlSchemesPreprocessor(Preprocessor):
|
||||
url = f"https://github.com/{repo}/{type}/{number}{fragment}"
|
||||
return f"[{gh_icon} {title}]({url})"
|
||||
|
||||
def replace_docs_link(match: re.Match) -> str:
|
||||
"""Rewrite absolute docs.vllm.ai links as doc-relative links."""
|
||||
title = match.group("title")
|
||||
path = match.group("path").rstrip("/")
|
||||
fragment = match.group("fragment") or ""
|
||||
|
||||
# vllm.config.<Class> API reference -> mkdocstrings cross-reference
|
||||
if path == "api/vllm/config" and re.fullmatch(
|
||||
r"#vllm\.config\.\w+", fragment
|
||||
):
|
||||
ident = fragment[1:]
|
||||
return f"[`{ident}`][{ident}]"
|
||||
|
||||
# Other docs pages -> link relative to the current page, but only
|
||||
# when the target is a known docs page (real or generated); leave
|
||||
# unknown/external URLs untouched. This is correct even when the same
|
||||
# docstring is also rendered on its API reference page.
|
||||
src = f"{path.removesuffix('.html')}.md"
|
||||
if files.get_file_from_path(src) is None:
|
||||
return match.group(0)
|
||||
rel = posixpath.relpath(src, posixpath.dirname(page.file.src_uri))
|
||||
# Auto-wrapped bare URLs use the URL as their title; make it readable.
|
||||
if title.startswith("http"):
|
||||
title = path.removesuffix(".html")
|
||||
return f"[{title}]({rel}{fragment})"
|
||||
|
||||
markdown = "\n".join(lines)
|
||||
markdown = github_link.sub(replace_github_link, markdown)
|
||||
markdown = relative_link.sub(replace_relative_link, markdown)
|
||||
markdown = docs_link.sub(replace_docs_link, markdown)
|
||||
return markdown.split("\n")
|
||||
|
||||
|
||||
@@ -116,6 +149,7 @@ class UrlSchemesExtension(Extension):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.page = None
|
||||
self.files = None
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def extendMarkdown(self, md):
|
||||
@@ -138,4 +172,5 @@ def on_page_markdown(
|
||||
) -> str:
|
||||
"""Pass the current page context to the preprocessor."""
|
||||
_ext.page = page
|
||||
_ext.files = files
|
||||
return markdown
|
||||
|
||||
@@ -19,7 +19,7 @@ vLLM also supports model implementations that are available in Transformers. We
|
||||
|
||||
Currently, the Transformers modeling backend works for the following:
|
||||
|
||||
- Modalities: embedding models, language models and vision-language models*
|
||||
- Modalities: embedding models, language models, vision-language models* and audio-language models
|
||||
- Architectures: encoder-only, decoder-only, mixture-of-experts
|
||||
- Attention types: full attention and/or sliding attention
|
||||
|
||||
@@ -427,7 +427,6 @@ th {
|
||||
| `OlmoeForCausalLM` | OLMoE | `allenai/OLMoE-1B-7B-0924`, `allenai/OLMoE-1B-7B-0924-Instruct`, etc. | | ✅︎ |
|
||||
| `OPTForCausalLM` | OPT, OPT-IML | `facebook/opt-66b`, `facebook/opt-iml-max-30b`, etc. | ✅︎ | ✅︎ |
|
||||
| `OrionForCausalLM` | Orion | `OrionStarAI/Orion-14B-Base`, `OrionStarAI/Orion-14B-Chat`, etc. | | ✅︎ |
|
||||
| `OuroForCausalLM` | ouro | `ByteDance/Ouro-1.4B`, `ByteDance/Ouro-2.6B`, etc. | ✅︎ | |
|
||||
| `PanguEmbeddedForCausalLM` | openPangu-Embedded-7B | `FreedomIntelligence/openPangu-Embedded-7B-V1.1` | ✅︎ | ✅︎ |
|
||||
| `PanguProMoEV2ForCausalLM` | openpangu-pro-moe-v2 | | ✅︎ | ✅︎ |
|
||||
| `PanguUltraMoEForCausalLM` | openpangu-ultra-moe-718b-model | `FreedomIntelligence/openPangu-Ultra-MoE-718B-V1.1` | ✅︎ | ✅︎ |
|
||||
@@ -465,6 +464,7 @@ Some models are supported only via the [Transformers modeling backend](#transfor
|
||||
| `Olmo2ForCausalLM` | OLMo2 | `allenai/OLMo-2-0425-1B`, etc. | ✅︎ | ✅︎ |
|
||||
| `SmolLM3ForCausalLM` | SmolLM3 | `HuggingFaceTB/SmolLM3-3B` | ✅︎ | ✅︎ |
|
||||
| `Starcoder2ForCausalLM` | Starcoder2 | `bigcode/starcoder2-3b`, `bigcode/starcoder2-7b`, `bigcode/starcoder2-15b`, etc. | ✅︎ | ✅︎ |
|
||||
| `VaultGemmaForCausalLM` | VaultGemma | `google/vaultgemma-1b` | ✅︎ | ✅︎ |
|
||||
|
||||
!!! note
|
||||
Currently, the ROCm version of vLLM supports Mistral and Mixtral only for context lengths up to 4096.
|
||||
@@ -607,7 +607,8 @@ Some models are supported only via the [Transformers modeling backend](#transfor
|
||||
|
||||
| Architecture | Models | Inputs | Example HF Models | [LoRA](../features/lora.md) | [PP](../serving/parallelism_scaling.md) |
|
||||
| ------------ | ------ | ------ | ----------------- | --------------------------- | --------------------------------------- |
|
||||
| `Emu3ForConditionalGeneration` | Emu3 | T + I | `BAAI/Emu3-Chat-hf` | ✅︎ | ✅︎ |
|
||||
| `Emu3ForConditionalGeneration` | Emu3 | T + I<sup>+</sup> | `BAAI/Emu3-Chat-hf` | ✅︎ | ✅︎ |
|
||||
| `VibeVoiceAsrForConditionalGeneration` | VibeVoice-ASR | T + A<sup>+</sup> | `microsoft/VibeVoice-ASR-HF` | ✅︎ | ✅︎ |
|
||||
|
||||
<sup>^</sup> You need to set the architecture name via `--hf-overrides` to match the one in vLLM.</br>
|
||||
<sup>E</sup> Pre-computed embeddings can be inputted for this modality.</br>
|
||||
|
||||
@@ -28,7 +28,7 @@ DOCS_PATHS=(
|
||||
docs/ # Actual docs content
|
||||
examples/ # Examples are rendered in docs
|
||||
vllm/ # API & CLI reference
|
||||
requirements/test/cuda.txt # CLI reference (see docs/mkdocs/hooks/generate_argparse.py)
|
||||
requirements/test/cuda.txt # CLI reference (see docs/mkdocs/gen_files/generate_argparse.py)
|
||||
mkdocs.yaml # Affects build process
|
||||
.readthedocs.yaml # Affects build process
|
||||
requirements/docs.txt # Affects build process
|
||||
|
||||
@@ -35,21 +35,21 @@ The following metrics are exposed:
|
||||
|
||||
## General Metrics
|
||||
|
||||
--8<-- "docs/generated/metrics/general.inc.md"
|
||||
--8<-- "gen:metrics-general"
|
||||
|
||||
## Speculative Decoding Metrics
|
||||
|
||||
--8<-- "docs/generated/metrics/spec_decode.inc.md"
|
||||
--8<-- "gen:metrics-spec-decode"
|
||||
|
||||
## NIXL KV Connector Metrics
|
||||
|
||||
--8<-- "docs/generated/metrics/nixl_connector.inc.md"
|
||||
--8<-- "gen:metrics-nixl"
|
||||
|
||||
## Model Flops Utilization (MFU) Performance Metrics
|
||||
|
||||
These metrics are available via `--enable-mfu-metrics`:
|
||||
|
||||
--8<-- "docs/generated/metrics/perf.inc.md"
|
||||
--8<-- "gen:metrics-mfu"
|
||||
|
||||
## Deprecation Policy
|
||||
|
||||
|
||||
@@ -503,6 +503,45 @@ def run_gemma3n(questions: list[str], modality: str) -> ModelRequestData:
|
||||
)
|
||||
|
||||
|
||||
# Gemma 4
|
||||
def run_gemma4(questions: list[str], modality: str) -> ModelRequestData:
|
||||
assert modality in ("image", "video")
|
||||
model_name = "google/gemma-4-31B-it"
|
||||
|
||||
# NOTE: Gemma-4-31B is a large model. Users running into Out-Of-Memory (OOM)
|
||||
# errors might need to set `tensor_parallel_size` to > 1.
|
||||
engine_args = EngineArgs(
|
||||
model=model_name,
|
||||
max_model_len=4096,
|
||||
max_num_seqs=2,
|
||||
limit_mm_per_prompt={modality: 1},
|
||||
)
|
||||
|
||||
if modality == "image":
|
||||
prompts = [
|
||||
(
|
||||
"<bos><start_of_turn>user\n"
|
||||
f"<|image|>\n{question}<end_of_turn>\n"
|
||||
"<start_of_turn>model\n"
|
||||
)
|
||||
for question in questions
|
||||
]
|
||||
else: # video
|
||||
prompts = [
|
||||
(
|
||||
"<bos><start_of_turn>user\n"
|
||||
f"<|video|>\n{question}<end_of_turn>\n"
|
||||
"<start_of_turn>model\n"
|
||||
)
|
||||
for question in questions
|
||||
]
|
||||
|
||||
return ModelRequestData(
|
||||
engine_args=engine_args,
|
||||
prompts=prompts,
|
||||
)
|
||||
|
||||
|
||||
# GLM-4v
|
||||
def run_glm4v(questions: list[str], modality: str) -> ModelRequestData:
|
||||
assert modality == "image"
|
||||
@@ -2303,6 +2342,7 @@ model_example_map = {
|
||||
"exaone4_5": run_exaone4_5,
|
||||
"gemma3": run_gemma3,
|
||||
"gemma3n": run_gemma3n,
|
||||
"gemma4": run_gemma4,
|
||||
"glm4v": run_glm4v,
|
||||
"glm4_1v": run_glm4_1v,
|
||||
"glm4_5v": run_glm4_5v,
|
||||
@@ -2374,6 +2414,7 @@ MODELS_NEED_VIDEO_METADATA = [
|
||||
|
||||
MODELS_SUPPORT_VIT_CUDA_GRAPH = [
|
||||
"llama4",
|
||||
"gemma4",
|
||||
"qwen2_vl",
|
||||
"qwen2_5_vl",
|
||||
"qwen3_vl",
|
||||
|
||||
+7
-10
@@ -3,7 +3,6 @@ site_url: !ENV READTHEDOCS_CANONICAL_URL
|
||||
repo_url: https://github.com/vllm-project/vllm
|
||||
edit_uri: edit/main/docs/
|
||||
exclude_docs: |
|
||||
argparse
|
||||
*.inc.md
|
||||
*.template.md
|
||||
theme:
|
||||
@@ -50,24 +49,22 @@ theme:
|
||||
|
||||
hooks:
|
||||
- docs/mkdocs/hooks/remove_announcement.py
|
||||
- docs/mkdocs/hooks/generate_examples.py
|
||||
- docs/mkdocs/hooks/generate_argparse.py
|
||||
- docs/mkdocs/hooks/generate_metrics.py
|
||||
- docs/mkdocs/hooks/url_schemes.py
|
||||
- docs/mkdocs/hooks/autoref_code.py
|
||||
|
||||
plugins:
|
||||
- meta
|
||||
- search
|
||||
- gen-files:
|
||||
scripts:
|
||||
- docs/mkdocs/gen_files/generate_examples.py
|
||||
- docs/mkdocs/gen_files/generate_argparse.py
|
||||
- docs/mkdocs/gen_files/generate_metrics.py
|
||||
- docs/mkdocs/gen_files/generate_attention_backends.py
|
||||
- autorefs
|
||||
- awesome-nav
|
||||
- glightbox
|
||||
- git-revision-date-localized:
|
||||
# exclude autogenerated files
|
||||
exclude:
|
||||
- api/*
|
||||
- examples/*
|
||||
- generated/*
|
||||
- git-revision-date-localized
|
||||
- minify:
|
||||
minify_html: true
|
||||
minify_js: true
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
syntax = "proto3";
|
||||
package vllm;
|
||||
|
||||
service Control {
|
||||
rpc GetServerInfo (GetServerInfoRequest) returns (ServerInfo) {}
|
||||
rpc GetModelInfo (GetModelInfoRequest) returns (ModelInfo) {}
|
||||
rpc Abort (AbortRequest) returns (AbortResponse) {}
|
||||
}
|
||||
|
||||
message GetServerInfoRequest {}
|
||||
|
||||
message ServerInfo {
|
||||
string engine_version = 1;
|
||||
string api_version = 2;
|
||||
string instance_id = 3;
|
||||
ParallelismInfo parallelism = 4;
|
||||
uint32 max_model_len = 5;
|
||||
uint32 kv_block_size = 6;
|
||||
uint64 total_kv_blocks = 7;
|
||||
uint64 max_running_requests = 8;
|
||||
uint64 max_batched_tokens = 9;
|
||||
}
|
||||
|
||||
message ParallelismInfo {
|
||||
uint32 tensor_parallel_size = 1;
|
||||
uint32 pipeline_parallel_size = 2;
|
||||
uint32 data_parallel_size = 3;
|
||||
uint32 data_parallel_rank = 4;
|
||||
uint32 decode_context_parallel_size = 5;
|
||||
}
|
||||
|
||||
message GetModelInfoRequest {}
|
||||
|
||||
message ModelInfo {
|
||||
string model_id = 1;
|
||||
string served_model_name = 2;
|
||||
repeated string served_model_aliases = 3;
|
||||
|
||||
bool supports_text_input = 20;
|
||||
bool supports_token_ids_input = 21;
|
||||
bool supports_multimodal = 23;
|
||||
string reasoning_parser = 24;
|
||||
string tool_call_parser = 25;
|
||||
}
|
||||
|
||||
message AbortRequest {
|
||||
repeated string request_ids = 1;
|
||||
}
|
||||
|
||||
message AbortResponse {}
|
||||
@@ -7,17 +7,13 @@ package vllm;
|
||||
import "google/protobuf/struct.proto";
|
||||
|
||||
|
||||
service Generate {
|
||||
service Inference {
|
||||
// Generates text given a prompt
|
||||
rpc Generate (GenerateRequest) returns (GenerateResponse) {}
|
||||
// Generates text given a prompt, streaming the outputs
|
||||
rpc GenerateStream (GenerateRequest) returns (stream GenerateResponse) {}
|
||||
}
|
||||
|
||||
service Control {
|
||||
rpc Abort (AbortRequest) returns (AbortResponse) {}
|
||||
}
|
||||
|
||||
// ======================================================================================
|
||||
// Generate Request
|
||||
// ======================================================================================
|
||||
@@ -204,13 +200,3 @@ message CandidateTokenInfo {
|
||||
message TokenIds {
|
||||
repeated uint32 ids = 1;
|
||||
}
|
||||
|
||||
// ======================================================================================
|
||||
// Control
|
||||
// ======================================================================================
|
||||
|
||||
message AbortRequest {
|
||||
repeated string request_ids = 1;
|
||||
}
|
||||
|
||||
message AbortResponse {}
|
||||
@@ -255,6 +255,33 @@ impl ChatLlm {
|
||||
self.text.engine_core_client()
|
||||
}
|
||||
|
||||
/// Whether the loaded backend has a registered multimodal processor.
|
||||
pub fn supports_multimodal(&self) -> bool {
|
||||
self.processor.backend.multimodal_model_info().is_some()
|
||||
}
|
||||
|
||||
/// Effective tool-call parser name for this model, if parsing is enabled.
|
||||
pub fn tool_call_parser_name(&self) -> Option<&str> {
|
||||
match &self.tool_call_parser {
|
||||
ParserSelection::Auto => {
|
||||
ToolParserFactory::global().resolve_name_for_model(self.model_id())
|
||||
}
|
||||
ParserSelection::None => None,
|
||||
ParserSelection::Explicit(name) => Some(name),
|
||||
}
|
||||
}
|
||||
|
||||
/// Effective reasoning parser name for this model, if parsing is enabled.
|
||||
pub fn reasoning_parser_name(&self) -> Option<&str> {
|
||||
match &self.reasoning_parser {
|
||||
ParserSelection::Auto => {
|
||||
ReasoningParserFactory::global().resolve_name_for_model(self.model_id())
|
||||
}
|
||||
ParserSelection::None => None,
|
||||
ParserSelection::Explicit(name) => Some(name),
|
||||
}
|
||||
}
|
||||
|
||||
/// Render, tokenize, and submit one chat request.
|
||||
pub async fn chat(&self, request: ChatRequest) -> Result<ChatEventStream> {
|
||||
let (text_request, output_processor) = self
|
||||
|
||||
@@ -148,11 +148,6 @@ pub struct SharedRuntimeArgs {
|
||||
#[arg(long)]
|
||||
#[serde(default)]
|
||||
pub language_model_only: bool,
|
||||
/// Override the maximum model context length. When set, the frontend uses
|
||||
/// this value instead of the model's `max_position_embeddings` from
|
||||
/// `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)]
|
||||
@@ -664,7 +659,6 @@ impl ServeArgs {
|
||||
|
||||
self.managed_engine.clone().into_config(
|
||||
self.runtime.model.clone(),
|
||||
self.runtime.max_model_len,
|
||||
self.runtime.max_logprobs,
|
||||
profiler_config,
|
||||
reasoning_parser.as_deref(),
|
||||
|
||||
@@ -58,9 +58,6 @@ fn serve_args_forward_python_flags_with_separator() {
|
||||
reasoning_parser: Auto,
|
||||
renderer: Auto,
|
||||
language_model_only: false,
|
||||
max_model_len: Some(
|
||||
512,
|
||||
),
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
@@ -102,6 +99,9 @@ fn serve_args_forward_python_flags_with_separator() {
|
||||
handshake_port: None,
|
||||
data_parallel_size: 1,
|
||||
data_parallel_size_local: None,
|
||||
max_model_len: Some(
|
||||
"512",
|
||||
),
|
||||
python_args: [
|
||||
"--dtype",
|
||||
"float16",
|
||||
@@ -756,7 +756,6 @@ fn frontend_args_accept_json() {
|
||||
reasoning_parser: None,
|
||||
renderer: Auto,
|
||||
language_model_only: false,
|
||||
max_model_len: None,
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
@@ -823,11 +822,32 @@ fn frontend_args_json_applies_defaults() {
|
||||
assert_eq!(args.runtime.tool_call_parser, ParserSelection::None);
|
||||
assert_eq!(args.runtime.reasoning_parser, ParserSelection::None);
|
||||
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);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn frontend_args_json_ignores_engine_owned_max_model_len() {
|
||||
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","max_model_len":-1}"#,
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Frontend(args) = cli.command else {
|
||||
panic!("expected frontend args");
|
||||
};
|
||||
assert_eq!(args.runtime.model, "Qwen/Qwen3-0.6B");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn frontend_args_json_accepts_supported_non_default_fields() {
|
||||
let cli = Cli::try_parse_from([
|
||||
@@ -840,7 +860,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,"max_logprobs":-1,"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_logprobs":-1,"shutdown_timeout":3}"#,
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
@@ -858,11 +878,38 @@ 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);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_args_forward_auto_max_model_len_to_managed_engine() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"vllm-rs",
|
||||
"serve",
|
||||
"Qwen/Qwen3-0.6B",
|
||||
"--max-model-len",
|
||||
"auto",
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let Command::Serve(args) = cli.command else {
|
||||
panic!("expected serve args");
|
||||
};
|
||||
assert_eq!(args.managed_engine.max_model_len.as_deref(), Some("auto"));
|
||||
|
||||
let config = args.to_managed_engine_config(5555);
|
||||
expect![[r#"
|
||||
[
|
||||
"--max-model-len",
|
||||
"auto",
|
||||
"--reasoning-parser",
|
||||
"qwen3",
|
||||
]
|
||||
"#]]
|
||||
.assert_debug_eq(&config.python_args);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serve_args_accept_none_reasoning_parser() {
|
||||
let cli = Cli::try_parse_from([
|
||||
@@ -1278,7 +1325,6 @@ fn serve_args_accept_handshake_aliases() {
|
||||
reasoning_parser: Auto,
|
||||
renderer: Auto,
|
||||
language_model_only: false,
|
||||
max_model_len: None,
|
||||
max_logprobs: None,
|
||||
grpc_port: None,
|
||||
shutdown_timeout: 0,
|
||||
@@ -1322,6 +1368,7 @@ fn serve_args_accept_handshake_aliases() {
|
||||
),
|
||||
data_parallel_size: 4,
|
||||
data_parallel_size_local: None,
|
||||
max_model_len: None,
|
||||
python_args: [],
|
||||
},
|
||||
},
|
||||
|
||||
@@ -394,6 +394,17 @@ impl EngineCoreClient {
|
||||
self.engines.iter().map(|engine| &engine.ready_response).collect()
|
||||
}
|
||||
|
||||
/// Return the first engine's ready response.
|
||||
///
|
||||
/// Per-engine fields such as `data_parallel_rank` should be read through
|
||||
/// [`ready_responses`](Self::ready_responses).
|
||||
pub fn ready_response(&self) -> &EngineCoreReadyResponse {
|
||||
&self
|
||||
.engines
|
||||
.first()
|
||||
.expect("engine core client requires at least one engine")
|
||||
.ready_response
|
||||
}
|
||||
/// Return the engine-reported effective model dtype.
|
||||
pub fn model_dtype(&self) -> ModelDtype {
|
||||
self.engines
|
||||
|
||||
@@ -74,17 +74,13 @@ impl EngineRoutingState {
|
||||
///
|
||||
/// Scheduler stats can raise the load estimate above the frontend-local
|
||||
/// view, but they should not lower it below requests this frontend has
|
||||
/// already admitted. Waiting requests still get the same extra penalty
|
||||
/// as the original `waiting * 4 + running` score.
|
||||
/// already admitted.
|
||||
fn routing_score(&self) -> usize {
|
||||
const WAITING_WEIGHT: usize = 4;
|
||||
|
||||
let Some(stats) = self.last_scheduler_stats else {
|
||||
return self.inflight;
|
||||
};
|
||||
|
||||
let scheduler_total = stats.running + stats.waiting;
|
||||
self.inflight.max(scheduler_total) + stats.waiting * (WAITING_WEIGHT - 1)
|
||||
self.inflight.max(stats.running + stats.waiting)
|
||||
}
|
||||
|
||||
/// Replace the local routing view with a fresh real scheduler snapshot.
|
||||
@@ -750,7 +746,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routing_score_keeps_extra_waiting_penalty() {
|
||||
fn routing_score_counts_waiting_without_extra_penalty() {
|
||||
let state = EngineRoutingState {
|
||||
inflight: 1,
|
||||
last_scheduler_stats: Some(EngineLoadSnapshot {
|
||||
@@ -759,7 +755,7 @@ mod tests {
|
||||
}),
|
||||
};
|
||||
|
||||
assert_eq!(state.routing_score(), 14);
|
||||
assert_eq!(state.routing_score(), 5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -57,6 +57,13 @@ pub fn default_ready_response() -> EngineCoreReadyResponse {
|
||||
vllm_version: "test-vllm-version".to_string(),
|
||||
world_size: 1,
|
||||
data_parallel_size: 1,
|
||||
tensor_parallel_size: 1,
|
||||
pipeline_parallel_size: 1,
|
||||
decode_context_parallel_size: 1,
|
||||
data_parallel_rank: 0,
|
||||
max_num_seqs: 256,
|
||||
max_num_batched_tokens: 8192,
|
||||
instance_id: "test-instance".to_string(),
|
||||
kv_cache_size_tokens: None,
|
||||
kv_cache_max_concurrency: None,
|
||||
}
|
||||
|
||||
@@ -52,6 +52,21 @@ pub struct EngineCoreReadyResponse {
|
||||
pub world_size: u64,
|
||||
/// Data parallelism size from the parallel config.
|
||||
pub data_parallel_size: u64,
|
||||
// Required discovery metadata; EngineCore and client versions must match.
|
||||
/// Tensor-parallel size of this engine.
|
||||
pub tensor_parallel_size: u32,
|
||||
/// Pipeline-parallel size of this engine.
|
||||
pub pipeline_parallel_size: u32,
|
||||
/// Decode-context-parallel size of this engine.
|
||||
pub decode_context_parallel_size: u32,
|
||||
/// This engine's data-parallel rank.
|
||||
pub data_parallel_rank: u32,
|
||||
/// Scheduler cap on concurrently running sequences.
|
||||
pub max_num_seqs: u64,
|
||||
/// Scheduler cap on batched tokens per step.
|
||||
pub max_num_batched_tokens: u64,
|
||||
/// Unique identifier for this server instance.
|
||||
pub instance_id: String,
|
||||
/// 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.
|
||||
|
||||
@@ -363,6 +363,13 @@ class EngineCoreReadyResponse:
|
||||
vllm_version: str
|
||||
world_size: int
|
||||
data_parallel_size: int
|
||||
tensor_parallel_size: int
|
||||
pipeline_parallel_size: int
|
||||
decode_context_parallel_size: int
|
||||
data_parallel_rank: int
|
||||
max_num_seqs: int
|
||||
max_num_batched_tokens: int
|
||||
instance_id: str
|
||||
kv_cache_size_tokens: int | None = None
|
||||
kv_cache_max_concurrency: float | None = None
|
||||
|
||||
@@ -376,6 +383,13 @@ ready_response = EngineCoreReadyResponse(
|
||||
vllm_version="0.0.0",
|
||||
data_parallel_size=1,
|
||||
world_size=1,
|
||||
tensor_parallel_size=1,
|
||||
pipeline_parallel_size=1,
|
||||
decode_context_parallel_size=1,
|
||||
data_parallel_rank=0,
|
||||
max_num_seqs=256,
|
||||
max_num_batched_tokens=8192,
|
||||
instance_id="test-instance",
|
||||
)
|
||||
|
||||
print(msgspec.msgpack.encode(request).hex())
|
||||
|
||||
@@ -39,6 +39,12 @@ pub struct ManagedEngineArgs {
|
||||
/// Number of data parallel replicas to run on this node.
|
||||
#[arg(long)]
|
||||
pub data_parallel_size_local: Option<usize>,
|
||||
/// Maximum model context length forwarded to the managed Python engine.
|
||||
///
|
||||
/// Rust leaves validation to Python so values such as `auto` and
|
||||
/// human-readable integers retain their engine-owned semantics.
|
||||
#[arg(long)]
|
||||
pub max_model_len: Option<String>,
|
||||
|
||||
/// Additional arguments forwarded to `python -m vllm.entrypoints.cli.main
|
||||
/// serve ...`.
|
||||
@@ -78,7 +84,6 @@ impl ManagedEngineArgs {
|
||||
pub fn into_config(
|
||||
self,
|
||||
model: String,
|
||||
max_model_len: Option<u32>,
|
||||
max_logprobs: Option<i32>,
|
||||
profiler_config: Option<String>,
|
||||
reasoning_parser: Option<&str>,
|
||||
@@ -89,9 +94,9 @@ impl ManagedEngineArgs {
|
||||
) -> ManagedEngineConfig {
|
||||
let mut python_args = self.python_args;
|
||||
// Manually forward some args to the Python engine.
|
||||
if let Some(max_model_len) = max_model_len {
|
||||
if let Some(max_model_len) = self.max_model_len {
|
||||
python_args.push("--max-model-len".to_string());
|
||||
python_args.push(max_model_len.to_string());
|
||||
python_args.push(max_model_len);
|
||||
}
|
||||
if let Some(max_logprobs) = max_logprobs {
|
||||
python_args.push("--max-logprobs".to_string());
|
||||
|
||||
@@ -9,7 +9,13 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
.build_server(true)
|
||||
.build_client(true)
|
||||
.protoc_arg("--experimental_allow_proto3_optional") // be compatible with old compilers
|
||||
.compile_protos(&[format!("{proto_dir}/vllm_grpc.proto")], &[proto_dir])?;
|
||||
.compile_protos(
|
||||
&[
|
||||
format!("{proto_dir}/control.proto"),
|
||||
format!("{proto_dir}/inference.proto"),
|
||||
],
|
||||
&[proto_dir],
|
||||
)?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use thiserror_ext::AsReport as _;
|
||||
use tonic::{Request, Response, Status};
|
||||
use vllm_engine_core_client::protocol::handshake::EngineCoreReadyResponse;
|
||||
|
||||
use super::{ControlServer, pb};
|
||||
use crate::state::AppState;
|
||||
|
||||
pub(crate) type ControlGrpcService = ControlServer<ControlServiceImpl>;
|
||||
|
||||
/// gRPC control service backed by the shared application state.
|
||||
pub struct ControlServiceImpl {
|
||||
state: Arc<AppState>,
|
||||
}
|
||||
|
||||
impl ControlServiceImpl {
|
||||
pub fn new(state: Arc<AppState>) -> Self {
|
||||
Self { state }
|
||||
}
|
||||
|
||||
fn ready(&self) -> &EngineCoreReadyResponse {
|
||||
self.state.engine_core_client().ready_response()
|
||||
}
|
||||
|
||||
fn parallelism_info(&self) -> pb::ParallelismInfo {
|
||||
let ready = self.ready();
|
||||
pb::ParallelismInfo {
|
||||
tensor_parallel_size: ready.tensor_parallel_size,
|
||||
pipeline_parallel_size: ready.pipeline_parallel_size,
|
||||
data_parallel_size: ready.data_parallel_size.min(u64::from(u32::MAX)) as u32,
|
||||
data_parallel_rank: ready.data_parallel_rank,
|
||||
decode_context_parallel_size: ready.decode_context_parallel_size,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const GRPC_API_VERSION: &str = "vllm";
|
||||
|
||||
#[tonic::async_trait]
|
||||
impl pb::control_server::Control for ControlServiceImpl {
|
||||
async fn get_server_info(
|
||||
&self,
|
||||
_request: Request<pb::GetServerInfoRequest>,
|
||||
) -> Result<Response<pb::ServerInfo>, Status> {
|
||||
let ready = self.ready();
|
||||
Ok(Response::new(pb::ServerInfo {
|
||||
engine_version: ready.vllm_version.clone(),
|
||||
api_version: GRPC_API_VERSION.to_string(),
|
||||
instance_id: ready.instance_id.clone(),
|
||||
parallelism: Some(self.parallelism_info()),
|
||||
max_model_len: self.state.engine_core_client().max_model_len(),
|
||||
kv_block_size: ready.block_size.min(u64::from(u32::MAX)) as u32,
|
||||
total_kv_blocks: self.state.engine_core_client().total_num_gpu_blocks(),
|
||||
max_running_requests: ready.max_num_seqs,
|
||||
max_batched_tokens: ready.max_num_batched_tokens,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn get_model_info(
|
||||
&self,
|
||||
_request: Request<pb::GetModelInfoRequest>,
|
||||
) -> Result<Response<pb::ModelInfo>, Status> {
|
||||
let served = self.state.served_model_names();
|
||||
Ok(Response::new(pb::ModelInfo {
|
||||
model_id: self.state.chat.text().model_id().to_string(),
|
||||
served_model_name: self.state.primary_model_name().to_string(),
|
||||
served_model_aliases: served.iter().skip(1).cloned().collect(),
|
||||
// GenerateRequest accepts both prompt representations.
|
||||
supports_text_input: true,
|
||||
supports_token_ids_input: true,
|
||||
supports_multimodal: self.state.chat.supports_multimodal(),
|
||||
reasoning_parser: self
|
||||
.state
|
||||
.chat
|
||||
.reasoning_parser_name()
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
tool_call_parser: self
|
||||
.state
|
||||
.chat
|
||||
.tool_call_parser_name()
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
async fn abort(
|
||||
&self,
|
||||
request: Request<pb::AbortRequest>,
|
||||
) -> Result<Response<pb::AbortResponse>, Status> {
|
||||
let request_ids = request.into_inner().request_ids;
|
||||
if request_ids.is_empty() {
|
||||
return Ok(Response::new(pb::AbortResponse {}));
|
||||
}
|
||||
self.state
|
||||
.chat
|
||||
.abort(&request_ids)
|
||||
.await
|
||||
.map_err(|error| Status::internal(error.to_report_string()))?;
|
||||
Ok(Response::new(pb::AbortResponse {}))
|
||||
}
|
||||
}
|
||||
@@ -8,14 +8,14 @@ use tonic_health::ServingStatus;
|
||||
use tonic_health::server::HealthReporter;
|
||||
use tracing::{info, warn};
|
||||
|
||||
use super::{ControlGrpcService, GenerateGrpcService};
|
||||
use super::{ControlGrpcService, InferenceGrpcService};
|
||||
|
||||
pub(crate) async fn monitor_health(
|
||||
mut health_reporter: HealthReporter,
|
||||
mut engine_health: watch::Receiver<bool>,
|
||||
shutdown: CancellationToken,
|
||||
) {
|
||||
let generate_service = GenerateGrpcService::NAME;
|
||||
let inference_service = InferenceGrpcService::NAME;
|
||||
let control_service = ControlGrpcService::NAME;
|
||||
let status = ServingStatus::NotServing;
|
||||
let health_event_first = tokio::select! {
|
||||
@@ -45,7 +45,7 @@ pub(crate) async fn monitor_health(
|
||||
}
|
||||
};
|
||||
|
||||
health_reporter.set_not_serving::<GenerateGrpcService>().await;
|
||||
health_reporter.set_not_serving::<InferenceGrpcService>().await;
|
||||
health_reporter.set_not_serving::<ControlGrpcService>().await;
|
||||
// Both gRPC services use the same engine client, so overall server health
|
||||
// mirrors their shared engine health.
|
||||
@@ -59,7 +59,7 @@ pub(crate) async fn monitor_health(
|
||||
);
|
||||
}
|
||||
|
||||
health_reporter.clear_service_status(generate_service).await;
|
||||
health_reporter.clear_service_status(inference_service).await;
|
||||
health_reporter.clear_service_status(control_service).await;
|
||||
health_reporter.clear_service_status("").await;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
|
||||
use futures::{Stream, StreamExt as _};
|
||||
use thiserror_ext::AsReport as _;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use tonic::{Request, Response, Status};
|
||||
use tracing::info;
|
||||
use vllm_text::{DecodedTextEvent, TextOutputStreamExt as _};
|
||||
|
||||
use super::convert::{self, ResponseOpts};
|
||||
use super::{InferenceServer, pb};
|
||||
use crate::state::AppState;
|
||||
|
||||
pub(crate) type InferenceGrpcService = InferenceServer<InferenceServiceImpl>;
|
||||
|
||||
/// gRPC inference service backed by the shared application state.
|
||||
pub struct InferenceServiceImpl {
|
||||
state: Arc<AppState>,
|
||||
}
|
||||
|
||||
impl InferenceServiceImpl {
|
||||
pub fn new(state: Arc<AppState>) -> Self {
|
||||
Self { state }
|
||||
}
|
||||
}
|
||||
|
||||
#[tonic::async_trait]
|
||||
impl pb::inference_server::Inference for InferenceServiceImpl {
|
||||
type GenerateStreamStream =
|
||||
Pin<Box<dyn Stream<Item = Result<pb::GenerateResponse, Status>> + Send>>;
|
||||
|
||||
/// Unary generate: collect all output and return a single response.
|
||||
async fn generate(
|
||||
&self,
|
||||
request: Request<pb::GenerateRequest>,
|
||||
) -> Result<Response<pb::GenerateResponse>, Status> {
|
||||
let proto_req = request.into_inner();
|
||||
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
|
||||
let text_request =
|
||||
convert::to_text_request(proto_req, false, self.state.served_model_names())?;
|
||||
|
||||
let request_id = text_request.request_id.clone();
|
||||
info!(%request_id, "grpc generate (unary)");
|
||||
|
||||
let stream = self.state.chat.text().generate(text_request).await;
|
||||
let stream = stream.map_err(text_error_to_status)?;
|
||||
|
||||
let collected = stream.collect_output().await.map_err(text_error_to_status)?;
|
||||
|
||||
// Build the single aggregated response.
|
||||
let prompt_info = convert::to_prompt_info(
|
||||
&collected.prompt_token_ids,
|
||||
collected.prompt_logprobs.as_ref(),
|
||||
&response_opts,
|
||||
);
|
||||
|
||||
let finish_info = vllm_text::Finished {
|
||||
usage: collected.usage,
|
||||
finish_reason: collected.finish_reason,
|
||||
kv_transfer_params: collected.kv_transfer_params,
|
||||
ec_transfer_params: collected.ec_transfer_params,
|
||||
};
|
||||
|
||||
let outputs = convert::to_sequence_output(
|
||||
&collected.text,
|
||||
&collected.token_ids,
|
||||
collected.logprobs.as_ref(),
|
||||
Some(&finish_info),
|
||||
&response_opts,
|
||||
);
|
||||
|
||||
Ok(Response::new(pb::GenerateResponse {
|
||||
prompt_info: Some(prompt_info),
|
||||
outputs: Some(outputs),
|
||||
}))
|
||||
}
|
||||
|
||||
/// Streaming generate: yield incremental responses as tokens are produced.
|
||||
async fn generate_stream(
|
||||
&self,
|
||||
request: Request<pb::GenerateRequest>,
|
||||
) -> Result<Response<Self::GenerateStreamStream>, Status> {
|
||||
let proto_req = request.into_inner();
|
||||
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
|
||||
let text_request =
|
||||
convert::to_text_request(proto_req, true, self.state.served_model_names())?;
|
||||
|
||||
let request_id = text_request.request_id.clone();
|
||||
info!(%request_id, "grpc generate (stream)");
|
||||
|
||||
let stream = self.state.chat.text().generate(text_request).await;
|
||||
let stream = stream.map_err(text_error_to_status)?;
|
||||
|
||||
let (tx, rx) = mpsc::channel(32);
|
||||
|
||||
tokio::spawn(async move {
|
||||
futures::pin_mut!(stream);
|
||||
while let Some(event) = stream.next().await {
|
||||
let response = match event {
|
||||
Err(e) => Err(text_error_to_status(e)),
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids,
|
||||
prompt_logprobs,
|
||||
}) => {
|
||||
let prompt_info = convert::to_prompt_info(
|
||||
&prompt_token_ids,
|
||||
prompt_logprobs.as_ref(),
|
||||
&response_opts,
|
||||
);
|
||||
Ok(pb::GenerateResponse {
|
||||
prompt_info: Some(prompt_info),
|
||||
outputs: None,
|
||||
})
|
||||
}
|
||||
Ok(DecodedTextEvent::TextDelta {
|
||||
delta,
|
||||
token_ids,
|
||||
logprobs,
|
||||
finished,
|
||||
}) => Ok(pb::GenerateResponse {
|
||||
prompt_info: None,
|
||||
outputs: Some(convert::to_sequence_output(
|
||||
&delta,
|
||||
&token_ids,
|
||||
logprobs.as_ref(),
|
||||
finished.as_ref(),
|
||||
&response_opts,
|
||||
)),
|
||||
}),
|
||||
};
|
||||
|
||||
if tx.send(response).await.is_err() {
|
||||
// Client disconnected.
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let response_stream = ReceiverStream::new(rx);
|
||||
Ok(Response::new(Box::pin(response_stream)))
|
||||
}
|
||||
}
|
||||
|
||||
fn text_error_to_status(error: vllm_text::Error) -> Status {
|
||||
let message = error.to_report_string();
|
||||
if error.is_request_validation_error() {
|
||||
Status::invalid_argument(message)
|
||||
} else {
|
||||
Status::internal(message)
|
||||
}
|
||||
}
|
||||
@@ -1,203 +1,25 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
//! gRPC Generate service backed by the shared [`vllm_text::TextLlm`] facade.
|
||||
//! gRPC services backed by the shared application state.
|
||||
|
||||
mod control;
|
||||
mod convert;
|
||||
mod health;
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
|
||||
use futures::{Stream, StreamExt as _};
|
||||
use thiserror_ext::AsReport as _;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use tonic::{Request, Response, Status};
|
||||
use tracing::info;
|
||||
use vllm_text::{DecodedTextEvent, TextOutputStreamExt as _};
|
||||
|
||||
use self::convert::ResponseOpts;
|
||||
use crate::state::AppState;
|
||||
mod inference;
|
||||
|
||||
/// Generated protobuf/gRPC types for the `vllm` package.
|
||||
pub mod pb {
|
||||
tonic::include_proto!("vllm");
|
||||
}
|
||||
|
||||
pub(crate) use control::ControlGrpcService;
|
||||
pub use control::ControlServiceImpl;
|
||||
pub(crate) use health::monitor_health;
|
||||
pub(crate) use inference::InferenceGrpcService;
|
||||
pub use inference::InferenceServiceImpl;
|
||||
pub use pb::control_server::ControlServer;
|
||||
pub use pb::generate_server::GenerateServer;
|
||||
|
||||
pub(crate) type ControlGrpcService = ControlServer<ControlServiceImpl>;
|
||||
pub(crate) type GenerateGrpcService = GenerateServer<GenerateServiceImpl>;
|
||||
pub use pb::inference_server::InferenceServer;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
/// gRPC Generate service implementation backed by the shared application state.
|
||||
pub struct GenerateServiceImpl {
|
||||
state: Arc<AppState>,
|
||||
}
|
||||
|
||||
impl GenerateServiceImpl {
|
||||
pub fn new(state: Arc<AppState>) -> Self {
|
||||
Self { state }
|
||||
}
|
||||
}
|
||||
|
||||
/// gRPC control service backed by the shared application state.
|
||||
pub struct ControlServiceImpl {
|
||||
state: Arc<AppState>,
|
||||
}
|
||||
|
||||
impl ControlServiceImpl {
|
||||
pub fn new(state: Arc<AppState>) -> Self {
|
||||
Self { state }
|
||||
}
|
||||
}
|
||||
|
||||
#[tonic::async_trait]
|
||||
impl pb::control_server::Control for ControlServiceImpl {
|
||||
async fn abort(
|
||||
&self,
|
||||
request: Request<pb::AbortRequest>,
|
||||
) -> Result<Response<pb::AbortResponse>, Status> {
|
||||
let request_ids = request.into_inner().request_ids;
|
||||
if request_ids.is_empty() {
|
||||
return Ok(Response::new(pb::AbortResponse {}));
|
||||
}
|
||||
self.state
|
||||
.chat
|
||||
.abort(&request_ids)
|
||||
.await
|
||||
.map_err(|error| Status::internal(error.to_report_string()))?;
|
||||
Ok(Response::new(pb::AbortResponse {}))
|
||||
}
|
||||
}
|
||||
|
||||
#[tonic::async_trait]
|
||||
impl pb::generate_server::Generate for GenerateServiceImpl {
|
||||
type GenerateStreamStream =
|
||||
Pin<Box<dyn Stream<Item = Result<pb::GenerateResponse, Status>> + Send>>;
|
||||
|
||||
/// Unary generate: collect all output and return a single response.
|
||||
async fn generate(
|
||||
&self,
|
||||
request: Request<pb::GenerateRequest>,
|
||||
) -> Result<Response<pb::GenerateResponse>, Status> {
|
||||
let proto_req = request.into_inner();
|
||||
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
|
||||
let text_request =
|
||||
convert::to_text_request(proto_req, false, self.state.served_model_names())?;
|
||||
|
||||
let request_id = text_request.request_id.clone();
|
||||
info!(%request_id, "grpc generate (unary)");
|
||||
|
||||
let stream = self.state.chat.text().generate(text_request).await;
|
||||
let stream = stream.map_err(text_error_to_status)?;
|
||||
|
||||
let collected = stream.collect_output().await.map_err(text_error_to_status)?;
|
||||
|
||||
// Build the single aggregated response.
|
||||
let prompt_info = convert::to_prompt_info(
|
||||
&collected.prompt_token_ids,
|
||||
collected.prompt_logprobs.as_ref(),
|
||||
&response_opts,
|
||||
);
|
||||
|
||||
let finish_info = vllm_text::Finished {
|
||||
usage: collected.usage,
|
||||
finish_reason: collected.finish_reason,
|
||||
kv_transfer_params: collected.kv_transfer_params,
|
||||
ec_transfer_params: collected.ec_transfer_params,
|
||||
};
|
||||
|
||||
let outputs = convert::to_sequence_output(
|
||||
&collected.text,
|
||||
&collected.token_ids,
|
||||
collected.logprobs.as_ref(),
|
||||
Some(&finish_info),
|
||||
&response_opts,
|
||||
);
|
||||
|
||||
Ok(Response::new(pb::GenerateResponse {
|
||||
prompt_info: Some(prompt_info),
|
||||
outputs: Some(outputs),
|
||||
}))
|
||||
}
|
||||
|
||||
/// Streaming generate: yield incremental responses as tokens are produced.
|
||||
async fn generate_stream(
|
||||
&self,
|
||||
request: Request<pb::GenerateRequest>,
|
||||
) -> Result<Response<Self::GenerateStreamStream>, Status> {
|
||||
let proto_req = request.into_inner();
|
||||
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
|
||||
let text_request =
|
||||
convert::to_text_request(proto_req, true, self.state.served_model_names())?;
|
||||
|
||||
let request_id = text_request.request_id.clone();
|
||||
info!(%request_id, "grpc generate (stream)");
|
||||
|
||||
let stream = self.state.chat.text().generate(text_request).await;
|
||||
let stream = stream.map_err(text_error_to_status)?;
|
||||
|
||||
let (tx, rx) = mpsc::channel(32);
|
||||
|
||||
tokio::spawn(async move {
|
||||
futures::pin_mut!(stream);
|
||||
while let Some(event) = stream.next().await {
|
||||
let response = match event {
|
||||
Err(e) => Err(text_error_to_status(e)),
|
||||
Ok(DecodedTextEvent::Start {
|
||||
prompt_token_ids,
|
||||
prompt_logprobs,
|
||||
}) => {
|
||||
let prompt_info = convert::to_prompt_info(
|
||||
&prompt_token_ids,
|
||||
prompt_logprobs.as_ref(),
|
||||
&response_opts,
|
||||
);
|
||||
Ok(pb::GenerateResponse {
|
||||
prompt_info: Some(prompt_info),
|
||||
outputs: None,
|
||||
})
|
||||
}
|
||||
Ok(DecodedTextEvent::TextDelta {
|
||||
delta,
|
||||
token_ids,
|
||||
logprobs,
|
||||
finished,
|
||||
}) => Ok(pb::GenerateResponse {
|
||||
prompt_info: None,
|
||||
outputs: Some(convert::to_sequence_output(
|
||||
&delta,
|
||||
&token_ids,
|
||||
logprobs.as_ref(),
|
||||
finished.as_ref(),
|
||||
&response_opts,
|
||||
)),
|
||||
}),
|
||||
};
|
||||
|
||||
if tx.send(response).await.is_err() {
|
||||
// Client disconnected.
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let response_stream = ReceiverStream::new(rx);
|
||||
Ok(Response::new(Box::pin(response_stream)))
|
||||
}
|
||||
}
|
||||
|
||||
fn text_error_to_status(error: vllm_text::Error) -> Status {
|
||||
let message = error.to_report_string();
|
||||
if error.is_request_validation_error() {
|
||||
Status::invalid_argument(message)
|
||||
} else {
|
||||
Status::internal(message)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -25,12 +25,18 @@ use vllm_chat::{
|
||||
ChatBackend, ChatLlm, ChatRenderer, ChatRequest, ChatTextBackend, DefaultChatOutputProcessor,
|
||||
DynChatOutputProcessor, DynChatRenderer, NewChatOutputProcessorOptions, RenderedPrompt,
|
||||
};
|
||||
use vllm_engine_core_client::mock_engine::{
|
||||
DEFAULT_MOCK_BLOCK_SIZE, DEFAULT_MOCK_MAX_MODEL_LEN, DEFAULT_MOCK_NUM_GPU_BLOCKS,
|
||||
default_ready_response,
|
||||
};
|
||||
use vllm_engine_core_client::protocol::output::{
|
||||
EngineCoreFinishReason, EngineCoreOutput, EngineCoreOutputs, RequestBatchOutputs,
|
||||
};
|
||||
use vllm_engine_core_client::protocol::request::EngineCoreRequest;
|
||||
use vllm_engine_core_client::test_utils::{IpcNamespace, spawn_mock_engine_task};
|
||||
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig, EngineId};
|
||||
use vllm_engine_core_client::test_utils::{
|
||||
IpcNamespace, spawn_mock_engine_task, spawn_mock_engine_task_with_ready,
|
||||
};
|
||||
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig, EngineId, TransportMode};
|
||||
use vllm_llm::Llm;
|
||||
use vllm_text::tokenizer::DynTokenizer;
|
||||
use vllm_text::{Prompt, TextBackend};
|
||||
@@ -39,8 +45,8 @@ use zeromq::prelude::{SocketRecv, SocketSend};
|
||||
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
|
||||
|
||||
use super::pb::control_client::ControlClient;
|
||||
use super::pb::generate_client::GenerateClient;
|
||||
use super::{ControlServer, ControlServiceImpl, GenerateServer, GenerateServiceImpl, pb};
|
||||
use super::pb::inference_client::InferenceClient;
|
||||
use super::{ControlServer, ControlServiceImpl, InferenceServer, InferenceServiceImpl, pb};
|
||||
use crate::listener::{Listener, MaybeTlsListener};
|
||||
use crate::state::AppState;
|
||||
use crate::tls;
|
||||
@@ -202,7 +208,7 @@ async fn setup_grpc_service(
|
||||
engine_id: impl Into<EngineId>,
|
||||
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
|
||||
) -> (
|
||||
GenerateServer<GenerateServiceImpl>,
|
||||
InferenceServer<InferenceServiceImpl>,
|
||||
ControlServer<ControlServiceImpl>,
|
||||
tokio::sync::watch::Receiver<bool>,
|
||||
MockEngineTask,
|
||||
@@ -246,7 +252,7 @@ async fn setup_grpc_service(
|
||||
);
|
||||
let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat));
|
||||
(
|
||||
GenerateServer::new(GenerateServiceImpl::new(state.clone())),
|
||||
InferenceServer::new(InferenceServiceImpl::new(state.clone())),
|
||||
ControlServer::new(ControlServiceImpl::new(state)),
|
||||
engine_health,
|
||||
engine_task,
|
||||
@@ -259,30 +265,30 @@ async fn grpc_test_server(
|
||||
engine_id: impl Into<EngineId>,
|
||||
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
|
||||
) -> (
|
||||
GenerateClient<tonic::transport::Channel>,
|
||||
InferenceClient<tonic::transport::Channel>,
|
||||
tokio::task::JoinHandle<()>,
|
||||
MockEngineTask,
|
||||
) {
|
||||
let (generate_service, control_service, engine_health, engine_task) =
|
||||
let (inference_service, control_service, engine_health, engine_task) =
|
||||
setup_grpc_service(engine_id, output_specs).await;
|
||||
let (channel, server_task) = start_grpc_test_server(
|
||||
generate_service,
|
||||
inference_service,
|
||||
control_service,
|
||||
engine_health,
|
||||
tokio_util::sync::CancellationToken::new(),
|
||||
)
|
||||
.await;
|
||||
(GenerateClient::new(channel), server_task, engine_task)
|
||||
(InferenceClient::new(channel), server_task, engine_task)
|
||||
}
|
||||
|
||||
async fn start_grpc_test_server(
|
||||
generate_service: GenerateServer<GenerateServiceImpl>,
|
||||
inference_service: InferenceServer<InferenceServiceImpl>,
|
||||
control_service: ControlServer<ControlServiceImpl>,
|
||||
engine_health: tokio::sync::watch::Receiver<bool>,
|
||||
shutdown: tokio_util::sync::CancellationToken,
|
||||
) -> (Channel, tokio::task::JoinHandle<()>) {
|
||||
let (health_reporter, health_service) = health_reporter();
|
||||
health_reporter.set_serving::<GenerateServer<GenerateServiceImpl>>().await;
|
||||
health_reporter.set_serving::<InferenceServer<InferenceServiceImpl>>().await;
|
||||
health_reporter.set_serving::<ControlServer<ControlServiceImpl>>().await;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
|
||||
@@ -293,7 +299,7 @@ async fn start_grpc_test_server(
|
||||
let server = TonicServer::builder()
|
||||
.add_service(health_service)
|
||||
.add_service(control_service)
|
||||
.add_service(generate_service)
|
||||
.add_service(inference_service)
|
||||
.serve_with_incoming_shutdown(incoming, shutdown.clone().cancelled_owned());
|
||||
let health_monitor =
|
||||
super::monitor_health(health_reporter, engine_health, shutdown.clone());
|
||||
@@ -323,7 +329,7 @@ async fn grpc_tls_test_server(
|
||||
certs: &TestCerts,
|
||||
cert_reqs: i32,
|
||||
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
|
||||
let (generate_service, control_service, _engine_health, engine_task) =
|
||||
let (inference_service, control_service, _engine_health, engine_task) =
|
||||
setup_grpc_service(engine_id, output_specs).await;
|
||||
let context = tls::build_grpc_server_config(&server_tls(certs, cert_reqs))
|
||||
.expect("build grpc tls config");
|
||||
@@ -335,7 +341,7 @@ async fn grpc_tls_test_server(
|
||||
let incoming = MaybeTlsListener::tls(Listener::Tcp(listener), context);
|
||||
TonicServer::builder()
|
||||
.add_service(control_service)
|
||||
.add_service(generate_service)
|
||||
.add_service(inference_service)
|
||||
.serve_with_incoming(incoming)
|
||||
.await
|
||||
.expect("grpc tls server");
|
||||
@@ -351,7 +357,7 @@ async fn grpc_tls_client(
|
||||
certs: &TestCerts,
|
||||
addr: &str,
|
||||
identity: Option<&str>,
|
||||
) -> Result<GenerateClient<Channel>, tonic::transport::Error> {
|
||||
) -> Result<InferenceClient<Channel>, tonic::transport::Error> {
|
||||
let ca = certs.path("ca.pem");
|
||||
let identity = identity.map(|name| {
|
||||
(
|
||||
@@ -388,7 +394,7 @@ async fn grpc_tls_client(
|
||||
.expect("grpc endpoint")
|
||||
.connect_with_connector(connector)
|
||||
.await?;
|
||||
Ok(GenerateClient::new(channel))
|
||||
Ok(InferenceClient::new(channel))
|
||||
}
|
||||
|
||||
/// Complete a raw TLS handshake against the gRPC port (offering ALPN `h2`) for
|
||||
@@ -415,7 +421,7 @@ async fn grpc_server_with_keepalive(
|
||||
engine_id: impl Into<EngineId>,
|
||||
keepalive: Option<Duration>,
|
||||
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
|
||||
let (generate_service, control_service, _engine_health, engine_task) =
|
||||
let (inference_service, control_service, _engine_health, engine_task) =
|
||||
setup_grpc_service(engine_id, default_stream_output_specs()).await;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
|
||||
@@ -432,7 +438,7 @@ async fn grpc_server_with_keepalive(
|
||||
let incoming = MaybeTlsListener::plain(Listener::Tcp(listener));
|
||||
builder
|
||||
.add_service(control_service)
|
||||
.add_service(generate_service)
|
||||
.add_service(inference_service)
|
||||
.serve_with_incoming(incoming)
|
||||
.await
|
||||
.expect("grpc server");
|
||||
@@ -1083,20 +1089,20 @@ async fn grpc_without_keepalive_keeps_unresponsive_connection_open() {
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn control_abort_resolves_external_id_and_empty_is_noop() {
|
||||
let (generate_service, control_service, engine_health, engine_task) =
|
||||
let (inference_service, control_service, engine_health, engine_task) =
|
||||
setup_grpc_service(b"engine-grpc-abort-active", vec![(vec![b'h' as u32], None)]).await;
|
||||
let (channel, server_task) = start_grpc_test_server(
|
||||
generate_service,
|
||||
inference_service,
|
||||
control_service,
|
||||
engine_health,
|
||||
tokio_util::sync::CancellationToken::new(),
|
||||
)
|
||||
.await;
|
||||
let mut generate_client = GenerateClient::new(channel.clone());
|
||||
let mut inference_client = InferenceClient::new(channel.clone());
|
||||
let mut control_client = ControlClient::new(channel);
|
||||
let request_id = "test-abort-active";
|
||||
|
||||
let mut stream = generate_client
|
||||
let mut stream = inference_client
|
||||
.generate_stream(pb::GenerateRequest {
|
||||
request_id: request_id.to_string(),
|
||||
model: "test-model".to_string(),
|
||||
@@ -1171,14 +1177,127 @@ async fn control_abort_resolves_external_id_and_empty_is_noop() {
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn control_reports_server_and_model_info() {
|
||||
let (generate_service, control_service, engine_health, _engine_task) =
|
||||
setup_grpc_service(b"engine-grpc-info", default_stream_output_specs()).await;
|
||||
let (channel, server_task) = start_grpc_test_server(
|
||||
generate_service,
|
||||
control_service,
|
||||
engine_health,
|
||||
tokio_util::sync::CancellationToken::new(),
|
||||
)
|
||||
.await;
|
||||
let mut client = ControlClient::new(channel);
|
||||
|
||||
let server = client
|
||||
.get_server_info(pb::GetServerInfoRequest {})
|
||||
.await
|
||||
.expect("get server info")
|
||||
.into_inner();
|
||||
assert_eq!(server.engine_version, "test-vllm-version");
|
||||
assert_eq!(server.api_version, "vllm");
|
||||
assert_eq!(server.instance_id, "test-instance");
|
||||
assert_eq!(server.max_model_len, DEFAULT_MOCK_MAX_MODEL_LEN as u32);
|
||||
assert_eq!(server.kv_block_size, DEFAULT_MOCK_BLOCK_SIZE as u32);
|
||||
assert_eq!(server.total_kv_blocks, DEFAULT_MOCK_NUM_GPU_BLOCKS);
|
||||
assert_eq!(server.max_running_requests, 256);
|
||||
assert_eq!(server.max_batched_tokens, 8_192);
|
||||
let parallelism = server.parallelism.expect("parallelism metadata");
|
||||
assert_eq!(parallelism.tensor_parallel_size, 1);
|
||||
assert_eq!(parallelism.pipeline_parallel_size, 1);
|
||||
assert_eq!(parallelism.data_parallel_size, 1);
|
||||
assert_eq!(parallelism.data_parallel_rank, 0);
|
||||
assert_eq!(parallelism.decode_context_parallel_size, 1);
|
||||
|
||||
let model = client
|
||||
.get_model_info(pb::GetModelInfoRequest {})
|
||||
.await
|
||||
.expect("get model info")
|
||||
.into_inner();
|
||||
assert_eq!(model.model_id, "test-model");
|
||||
assert_eq!(model.served_model_name, "test-model");
|
||||
assert!(model.served_model_aliases.is_empty());
|
||||
assert!(model.supports_text_input);
|
||||
assert!(model.supports_token_ids_input);
|
||||
assert!(!model.supports_multimodal);
|
||||
assert!(model.reasoning_parser.is_empty());
|
||||
assert!(model.tool_call_parser.is_empty());
|
||||
|
||||
server_task.abort();
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn control_aggregates_multi_engine_capacity() {
|
||||
let ipc = IpcNamespace::new().expect("create ipc namespace");
|
||||
let handshake_address = ipc.handshake_endpoint();
|
||||
|
||||
let mut ready_0 = default_ready_response();
|
||||
ready_0.max_model_len = 8_192;
|
||||
ready_0.num_gpu_blocks = 10;
|
||||
ready_0.data_parallel_size = 2;
|
||||
|
||||
let mut ready_1 = default_ready_response();
|
||||
ready_1.max_model_len = 4_096;
|
||||
ready_1.num_gpu_blocks = 20;
|
||||
ready_1.data_parallel_size = 2;
|
||||
ready_1.data_parallel_rank = 1;
|
||||
|
||||
let engine_tasks = [ready_0, ready_1].map(|ready| {
|
||||
let engine_id = EngineId::from_engine_index(ready.data_parallel_rank);
|
||||
MockEngineTask::new(spawn_mock_engine_task_with_ready(
|
||||
handshake_address.clone(),
|
||||
engine_id,
|
||||
ready,
|
||||
|_, _| boxed_test_future(async {}),
|
||||
))
|
||||
});
|
||||
|
||||
let client = EngineCoreClient::connect(EngineCoreClientConfig {
|
||||
transport_mode: TransportMode::HandshakeOwner {
|
||||
handshake_address,
|
||||
advertised_host: "127.0.0.1".to_string(),
|
||||
engine_count: 2,
|
||||
ready_timeout: Duration::from_secs(2),
|
||||
local_input_address: Some(ipc.input_endpoint()),
|
||||
local_output_address: Some(ipc.output_endpoint()),
|
||||
},
|
||||
coordinator_mode: None,
|
||||
model_name: "test-model".to_string(),
|
||||
client_index: 0,
|
||||
})
|
||||
.await
|
||||
.expect("connect multi-engine client");
|
||||
let chat = ChatLlm::from_shared_backend(
|
||||
Llm::new(client),
|
||||
Arc::new(FakeTextBackend) as Arc<dyn ChatTextBackend>,
|
||||
);
|
||||
let service = ControlServiceImpl::new(Arc::new(AppState::new(
|
||||
vec!["test-model".to_string()],
|
||||
chat,
|
||||
)));
|
||||
|
||||
let server = pb::control_server::Control::get_server_info(
|
||||
&service,
|
||||
tonic::Request::new(pb::GetServerInfoRequest {}),
|
||||
)
|
||||
.await
|
||||
.expect("get server info")
|
||||
.into_inner();
|
||||
assert_eq!(server.max_model_len, 4_096);
|
||||
assert_eq!(server.total_kv_blocks, 30);
|
||||
|
||||
drop(engine_tasks);
|
||||
}
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() {
|
||||
let (generate_service, control_service, _connected_engine_health, _engine_task) =
|
||||
let (inference_service, control_service, _connected_engine_health, _engine_task) =
|
||||
setup_grpc_service(b"engine-grpc-health-failure", default_stream_output_specs()).await;
|
||||
let (engine_health_tx, engine_health) = tokio::sync::watch::channel(true);
|
||||
let (channel, server_task) = start_grpc_test_server(
|
||||
generate_service,
|
||||
inference_service,
|
||||
control_service,
|
||||
engine_health,
|
||||
tokio_util::sync::CancellationToken::new(),
|
||||
@@ -1187,7 +1306,7 @@ async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy()
|
||||
let mut health_client = HealthClient::new(channel);
|
||||
|
||||
let mut health_streams = Vec::new();
|
||||
for service in ["vllm.Generate", "vllm.Control", ""] {
|
||||
for service in ["vllm.Inference", "vllm.Control", ""] {
|
||||
let service_label = if service.is_empty() {
|
||||
"overall"
|
||||
} else {
|
||||
@@ -1242,14 +1361,14 @@ async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy()
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial]
|
||||
async fn grpc_health_watch_closes_on_graceful_shutdown() {
|
||||
let (generate_service, control_service, engine_health, _engine_task) = setup_grpc_service(
|
||||
let (inference_service, control_service, engine_health, _engine_task) = setup_grpc_service(
|
||||
b"engine-grpc-health-shutdown",
|
||||
default_stream_output_specs(),
|
||||
)
|
||||
.await;
|
||||
let shutdown = tokio_util::sync::CancellationToken::new();
|
||||
let (channel, server_task) = start_grpc_test_server(
|
||||
generate_service,
|
||||
inference_service,
|
||||
control_service,
|
||||
engine_health,
|
||||
shutdown.clone(),
|
||||
@@ -1258,43 +1377,43 @@ async fn grpc_health_watch_closes_on_graceful_shutdown() {
|
||||
let mut health_client = HealthClient::new(channel);
|
||||
let mut stream = health_client
|
||||
.watch(HealthCheckRequest {
|
||||
service: "vllm.Generate".to_string(),
|
||||
service: "vllm.Inference".to_string(),
|
||||
})
|
||||
.await
|
||||
.expect("start health watch for vllm.Generate")
|
||||
.expect("start health watch for vllm.Inference")
|
||||
.into_inner();
|
||||
|
||||
let initial = stream
|
||||
.message()
|
||||
.await
|
||||
.expect("read initial health status for vllm.Generate")
|
||||
.expect("read initial health status for vllm.Inference")
|
||||
.expect("health watch ended before its initial status");
|
||||
assert_eq!(
|
||||
initial.status,
|
||||
HealthServingStatus::Serving as i32,
|
||||
"unexpected initial health status for vllm.Generate"
|
||||
"unexpected initial health status for vllm.Inference"
|
||||
);
|
||||
|
||||
shutdown.cancel();
|
||||
|
||||
let update = tokio::time::timeout(Duration::from_secs(2), stream.message())
|
||||
.await
|
||||
.expect("timed out waiting for shutdown health update for vllm.Generate")
|
||||
.expect("failed to read shutdown health update for vllm.Generate")
|
||||
.expect("timed out waiting for shutdown health update for vllm.Inference")
|
||||
.expect("failed to read shutdown health update for vllm.Inference")
|
||||
.expect("health watch ended before its shutdown update");
|
||||
assert_eq!(
|
||||
update.status,
|
||||
HealthServingStatus::NotServing as i32,
|
||||
"unexpected shutdown health status for vllm.Generate"
|
||||
"unexpected shutdown health status for vllm.Inference"
|
||||
);
|
||||
|
||||
let stream_end = tokio::time::timeout(Duration::from_secs(2), stream.message())
|
||||
.await
|
||||
.expect("timed out waiting for vllm.Generate health watch to close")
|
||||
.expect("failed while closing vllm.Generate health watch");
|
||||
.expect("timed out waiting for vllm.Inference health watch to close")
|
||||
.expect("failed while closing vllm.Inference health watch");
|
||||
assert!(
|
||||
stream_end.is_none(),
|
||||
"vllm.Generate health watch remained open"
|
||||
"vllm.Inference health watch remained open"
|
||||
);
|
||||
|
||||
tokio::time::timeout(Duration::from_secs(2), server_task)
|
||||
|
||||
@@ -187,7 +187,7 @@ where
|
||||
let model = state.primary_model_name().to_owned();
|
||||
let app = extend_router(build_router(state.clone()));
|
||||
|
||||
// Optionally bind the gRPC Generate server on a separate port. Bind
|
||||
// Optionally bind the gRPC Inference server on a separate port. Bind
|
||||
// synchronously here so bind errors (port in use, permission denied, ...)
|
||||
// surface before serving rather than being deferred until shutdown.
|
||||
let grpc_setup = if let Some(grpc_port) = config.grpc_port {
|
||||
@@ -206,19 +206,19 @@ where
|
||||
.context("invalid gRPC TLS configuration")?;
|
||||
let (health_reporter, health_service) = health_reporter();
|
||||
let engine_health = state.engine_core_client().subscribe_health();
|
||||
health_reporter.set_serving::<grpc::GenerateGrpcService>().await;
|
||||
health_reporter.set_serving::<grpc::InferenceGrpcService>().await;
|
||||
health_reporter.set_serving::<grpc::ControlGrpcService>().await;
|
||||
let control_service =
|
||||
grpc::ControlGrpcService::new(grpc::ControlServiceImpl::new(state.clone()));
|
||||
let generate_service =
|
||||
grpc::GenerateGrpcService::new(grpc::GenerateServiceImpl::new(state.clone()));
|
||||
let inference_service =
|
||||
grpc::InferenceGrpcService::new(grpc::InferenceServiceImpl::new(state.clone()));
|
||||
let svc = TonicServer::builder()
|
||||
.http2_keepalive_interval(Some(GRPC_KEEPALIVE_INTERVAL))
|
||||
.http2_keepalive_timeout(Some(GRPC_KEEPALIVE_TIMEOUT))
|
||||
.layer(middleware::request_runtime_layer(state.clone()))
|
||||
.add_service(health_service)
|
||||
.add_service(control_service)
|
||||
.add_service(generate_service);
|
||||
.add_service(inference_service);
|
||||
info!(%addr, tls = grpc_tls.is_some(), "starting gRPC server");
|
||||
Some((grpc_listener, svc, grpc_tls, health_reporter, engine_health))
|
||||
} else {
|
||||
|
||||
@@ -30,8 +30,8 @@ const OFFLOADED_PATHS: &[&str] = &[
|
||||
"/detokenize",
|
||||
"/inference/v1/generate",
|
||||
// gRPC routes:
|
||||
"/vllm.Generate/Generate",
|
||||
"/vllm.Generate/GenerateStream",
|
||||
"/vllm.Inference/Generate",
|
||||
"/vllm.Inference/GenerateStream",
|
||||
];
|
||||
|
||||
/// Return a Tower layer that runs selected data-plane requests on the request runtime,
|
||||
@@ -124,8 +124,8 @@ mod tests {
|
||||
assert!(should_offload("/tokenize"));
|
||||
assert!(should_offload("/detokenize"));
|
||||
assert!(should_offload("/inference/v1/generate"));
|
||||
assert!(should_offload("/vllm.Generate/Generate"));
|
||||
assert!(should_offload("/vllm.Generate/GenerateStream"));
|
||||
assert!(should_offload("/vllm.Inference/Generate"));
|
||||
assert!(should_offload("/vllm.Inference/GenerateStream"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -83,9 +83,33 @@ def test_copy_pass():
|
||||
def test_custom_op():
|
||||
# proper syntax
|
||||
_ = CompilationConfig(custom_ops=["+quant_fp8", "-silu_and_mul"])
|
||||
_ = CompilationConfig(custom_ops=["none", "+rms_norm"])
|
||||
_ = CompilationConfig(custom_ops=["+rms_norm", "+rms_norm"])
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid syntax '"):
|
||||
_ = CompilationConfig(custom_ops=["quant_fp8"])
|
||||
for custom_ops in (["quant_fp8"], ["+"], ["-"]):
|
||||
with pytest.raises(ValueError, match="Invalid syntax '"):
|
||||
CompilationConfig(custom_ops=custom_ops)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("custom_ops", "config_kwargs", "match"),
|
||||
[
|
||||
(["all", "none"], {}, "can contain only one base mode"),
|
||||
(
|
||||
["none", "+rms_norm", "-rms_norm"],
|
||||
{},
|
||||
"cannot both enable and disable.*rms_norm",
|
||||
),
|
||||
(
|
||||
["-rotary_embedding"],
|
||||
{"pass_config": PassConfig(enable_qk_norm_rope_fusion=True)},
|
||||
"cannot both enable and disable.*rotary_embedding",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_reject_contradictory_custom_ops(custom_ops, config_kwargs, match):
|
||||
with pytest.raises(ValueError, match=match):
|
||||
CompilationConfig(custom_ops=custom_ops, **config_kwargs)
|
||||
|
||||
|
||||
# forked needed to workaround https://github.com/vllm-project/vllm/issues/21073
|
||||
|
||||
@@ -1,10 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from vllm.config.model import ModelConfig
|
||||
from vllm.config.multimodal import MultiModalConfig
|
||||
from vllm.transformers_utils.model_arch_config_convertor import (
|
||||
ModelArchConfigConvertorBase,
|
||||
)
|
||||
from vllm.v1.attention.backends.registry import AttentionBackendEnum
|
||||
|
||||
|
||||
@@ -68,3 +74,97 @@ def test_mm_encoder_attn_dtype_hash_updates(tmp_path):
|
||||
).compute_hash()
|
||||
assert base_hash != fp8_hash
|
||||
assert fp8_hash != fp8_static_hash
|
||||
|
||||
|
||||
def _make_mm_prefix_model_config(
|
||||
*,
|
||||
language_model_only: bool = False,
|
||||
) -> ModelConfig:
|
||||
model_config = MagicMock(spec=ModelConfig)
|
||||
model_config.multimodal_config = MultiModalConfig(
|
||||
language_model_only=language_model_only
|
||||
)
|
||||
# Bind real helper methods onto the mock.
|
||||
model_config._supports_multimodal_for_mm_prefix = (
|
||||
ModelConfig._supports_multimodal_for_mm_prefix.__get__(
|
||||
model_config, ModelConfig
|
||||
)
|
||||
)
|
||||
return model_config
|
||||
|
||||
|
||||
@pytest.mark.parametrize("supports_mm", [True, False])
|
||||
def test_supports_multimodal_for_mm_prefix_uses_registry(supports_mm: bool):
|
||||
model_config = _make_mm_prefix_model_config()
|
||||
|
||||
with patch(
|
||||
"vllm.multimodal.MULTIMODAL_REGISTRY.supports_multimodal_inputs",
|
||||
return_value=supports_mm,
|
||||
) as mocked:
|
||||
assert model_config._supports_multimodal_for_mm_prefix() is supports_mm
|
||||
mocked.assert_called_once_with(model_config)
|
||||
|
||||
# Sticky cache — registry must not be consulted again.
|
||||
with patch(
|
||||
"vllm.multimodal.MULTIMODAL_REGISTRY.supports_multimodal_inputs",
|
||||
side_effect=AssertionError("should use cache"),
|
||||
):
|
||||
assert model_config._supports_multimodal_for_mm_prefix() is supports_mm
|
||||
|
||||
|
||||
def test_supports_multimodal_for_mm_prefix_before_multimodal_config():
|
||||
model_config = _make_mm_prefix_model_config()
|
||||
model_config.multimodal_config = None
|
||||
|
||||
assert model_config._supports_multimodal_for_mm_prefix() is True
|
||||
assert not hasattr(model_config, "_supports_multimodal_inputs_cached")
|
||||
|
||||
|
||||
def test_language_model_only_disables_via_supports_multimodal_inputs():
|
||||
"""language_model_only zeros all limits, so registry reports text-only."""
|
||||
model_config = _make_mm_prefix_model_config(language_model_only=True)
|
||||
|
||||
with patch(
|
||||
"vllm.multimodal.MULTIMODAL_REGISTRY.supports_multimodal_inputs",
|
||||
return_value=False,
|
||||
):
|
||||
assert model_config._supports_multimodal_for_mm_prefix() is False
|
||||
|
||||
|
||||
def test_convertor_clears_mm_prefix_when_multimodal_disabled():
|
||||
hf_config = PretrainedConfig(
|
||||
model_type="gemma3",
|
||||
architectures=["Gemma3ForConditionalGeneration"],
|
||||
)
|
||||
hf_config.is_mm_prefix_lm = True
|
||||
convertor = ModelArchConfigConvertorBase(hf_config, hf_config)
|
||||
|
||||
assert convertor.is_mm_prefix_lm(supports_multimodal=True) is True
|
||||
assert convertor.is_mm_prefix_lm(supports_multimodal=False) is False
|
||||
|
||||
enabled = convertor.convert(supports_multimodal=True)
|
||||
disabled = convertor.convert(supports_multimodal=False)
|
||||
assert enabled.is_mm_prefix_lm is True
|
||||
assert disabled.is_mm_prefix_lm is False
|
||||
|
||||
|
||||
def test_sticky_cache_survives_text_subconfig_regeneration():
|
||||
"""with_hf_config deepcopies the cached decision onto text submodules."""
|
||||
model_config = _make_mm_prefix_model_config()
|
||||
with patch(
|
||||
"vllm.multimodal.MULTIMODAL_REGISTRY.supports_multimodal_inputs",
|
||||
return_value=False,
|
||||
):
|
||||
assert model_config._supports_multimodal_for_mm_prefix() is False
|
||||
|
||||
# Simulate deepcopy onto a Gemma4ForCausalLM-like config that would
|
||||
# otherwise fail registry lookup / return False incorrectly.
|
||||
text_config = _make_mm_prefix_model_config()
|
||||
text_config._supports_multimodal_inputs_cached = (
|
||||
model_config._supports_multimodal_inputs_cached
|
||||
)
|
||||
with patch(
|
||||
"vllm.multimodal.MULTIMODAL_REGISTRY.supports_multimodal_inputs",
|
||||
side_effect=AssertionError("must not re-query registry"),
|
||||
):
|
||||
assert text_config._supports_multimodal_for_mm_prefix() is False
|
||||
|
||||
+7
-11
@@ -31,7 +31,7 @@ import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from huggingface_hub import snapshot_download
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
from PIL import Image
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
@@ -1496,7 +1496,7 @@ _dummy_gemma2_embedding_path = os.path.join(temp_dir, "dummy_gemma2_embedding")
|
||||
def dummy_opt_path():
|
||||
json_path = os.path.join(_dummy_opt_path, "config.json")
|
||||
if not os.path.exists(_dummy_opt_path):
|
||||
snapshot_download(
|
||||
hf_api().snapshot_download(
|
||||
repo_id="facebook/opt-125m",
|
||||
local_dir=_dummy_opt_path,
|
||||
ignore_patterns=["*.bin", "*.bin.index.json", "*.pt", "*.h5", "*.msgpack"],
|
||||
@@ -1514,7 +1514,7 @@ def dummy_opt_path():
|
||||
def dummy_llava_path():
|
||||
json_path = os.path.join(_dummy_llava_path, "config.json")
|
||||
if not os.path.exists(_dummy_llava_path):
|
||||
snapshot_download(
|
||||
hf_api().snapshot_download(
|
||||
repo_id="llava-hf/llava-1.5-7b-hf",
|
||||
local_dir=_dummy_llava_path,
|
||||
ignore_patterns=[
|
||||
@@ -1539,7 +1539,7 @@ def dummy_llava_path():
|
||||
def dummy_gemma2_embedding_path():
|
||||
json_path = os.path.join(_dummy_gemma2_embedding_path, "config.json")
|
||||
if not os.path.exists(_dummy_gemma2_embedding_path):
|
||||
snapshot_download(
|
||||
hf_api().snapshot_download(
|
||||
repo_id="BAAI/bge-multilingual-gemma2",
|
||||
local_dir=_dummy_gemma2_embedding_path,
|
||||
ignore_patterns=[
|
||||
@@ -1729,13 +1729,9 @@ def disable_deepgemm_ue8m0(monkeypatch):
|
||||
|
||||
|
||||
def _should_clean_gpu_memory_between_tests() -> bool:
|
||||
setting = os.getenv("VLLM_TEST_CLEAN_GPU_MEMORY")
|
||||
if setting == "1":
|
||||
return True
|
||||
if setting == "0":
|
||||
return False
|
||||
# ROCm reclaims VRAM lazily; default to waiting between tests on ROCm CI.
|
||||
return current_platform.is_rocm()
|
||||
# This must stay opt-in: a function-scoped fixture cannot distinguish
|
||||
# stale VRAM from allocations owned by longer-lived module/session fixtures.
|
||||
return os.getenv("VLLM_TEST_CLEAN_GPU_MEMORY", "0") == "1"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
||||
@@ -7,6 +7,7 @@ Run `pytest tests/distributed/test_comm_ops.py`.
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
import ray
|
||||
@@ -19,6 +20,8 @@ from vllm.distributed import (
|
||||
tensor_model_parallel_all_reduce,
|
||||
tensor_model_parallel_reduce_scatter,
|
||||
)
|
||||
from vllm.distributed.device_communicators import flashinfer_all_reduce
|
||||
from vllm.distributed.device_communicators.cuda_communicator import CudaCommunicator
|
||||
from vllm.distributed.parallel_state import GroupCoordinator, TensorMetadata
|
||||
from vllm.v1.worker.gpu_worker import AsyncIntermediateTensors
|
||||
|
||||
@@ -278,6 +281,43 @@ def test_irecv_tensor_dict_send_allgather_postprocess_binds_keys(
|
||||
torch.testing.assert_close(td["b"], torch.ones(4, dtype=torch.int32))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("aliased", [False, True])
|
||||
def test_cuda_communicator_checkpoints_flashinfer_workspaces(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
aliased: bool,
|
||||
) -> None:
|
||||
group = object()
|
||||
normal_workspace = Mock()
|
||||
quant_workspace = normal_workspace if aliased else Mock()
|
||||
unique_workspaces = (
|
||||
[normal_workspace] if aliased else [normal_workspace, quant_workspace]
|
||||
)
|
||||
|
||||
monkeypatch.setattr(flashinfer_all_reduce, "_fi_ar_workspace", normal_workspace)
|
||||
monkeypatch.setattr(
|
||||
flashinfer_all_reduce, "_fi_ar_quant_workspace", quant_workspace
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
flashinfer_all_reduce,
|
||||
"_fi_ar_workspace_groups",
|
||||
{id(workspace): group for workspace in unique_workspaces},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
flashinfer_all_reduce, "TorchDistBackend", lambda group: group, raising=False
|
||||
)
|
||||
|
||||
communicator = CudaCommunicator.__new__(CudaCommunicator)
|
||||
communicator.cpu_group = group
|
||||
communicator.fi_ar_comm = None
|
||||
communicator.all2all_manager = None
|
||||
communicator.checkpoint_prepare()
|
||||
communicator.checkpoint_restore()
|
||||
|
||||
for workspace in unique_workspaces:
|
||||
workspace.checkpoint_prepare.assert_called_once_with()
|
||||
workspace.checkpoint_restore.assert_called_once_with(group)
|
||||
|
||||
|
||||
def test_async_intermediate_tensors_lazy_wait() -> None:
|
||||
work = _DummyWork()
|
||||
post_calls = {"n": 0}
|
||||
|
||||
@@ -15,14 +15,16 @@ from unittest.mock import patch
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
import vllm.envs as envs
|
||||
from vllm import LLM, SamplingParams
|
||||
from vllm.distributed import cleanup_dist_env_and_memory
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
from vllm.utils.network_utils import get_open_port
|
||||
|
||||
from ..utils import multi_gpu_test
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not current_platform.is_rocm(),
|
||||
reason="ROCm-only quick-reduce tests",
|
||||
@@ -151,6 +153,116 @@ def _run_two_gpu_quick_allreduce_test(
|
||||
)
|
||||
|
||||
|
||||
CUDAGRAPH_WORLD_SIZE = 2
|
||||
CUDAGRAPH_ROUNDS = 10
|
||||
CUDAGRAPH_NUM_ELEMENTS = 1 << 21 # 2M fp16 = 4 MB, above the quick-reduce thresholds
|
||||
|
||||
|
||||
def _quick_allreduce_cudagraph_worker(
|
||||
rank: int,
|
||||
world_size: int,
|
||||
port: int,
|
||||
quant_level: str,
|
||||
):
|
||||
# FP keeps the all-reduce bit-exact for small fp16 integers, so every
|
||||
# replayed round can be checked exactly.
|
||||
os.environ["VLLM_ROCM_QUICK_REDUCE_QUANTIZATION"] = quant_level
|
||||
os.environ["VLLM_ROCM_QUICK_REDUCE_CAST_BF16_TO_FP16"] = "0"
|
||||
_log(f"cudagraph worker start: rank={rank} quant={quant_level}")
|
||||
|
||||
device = torch.device(f"cuda:{rank}")
|
||||
torch.accelerator.set_device_index(device)
|
||||
dist.init_process_group(
|
||||
backend="gloo",
|
||||
init_method=f"tcp://127.0.0.1:{port}",
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
qar = None
|
||||
try:
|
||||
from vllm.distributed.device_communicators.quick_all_reduce import (
|
||||
QuickAllReduce,
|
||||
)
|
||||
|
||||
qar = QuickAllReduce(group=dist.GroupMember.WORLD, device=rank)
|
||||
assert not qar.disabled
|
||||
|
||||
N = CUDAGRAPH_NUM_ELEMENTS
|
||||
inp = torch.empty(N, dtype=torch.float16, device=device)
|
||||
out = torch.empty(N, dtype=torch.float16, device=device)
|
||||
assert qar.should_quick_allreduce(inp)
|
||||
|
||||
# Every rank contributes the same value v in a round, so the true
|
||||
# cross-rank all-reduce sum is simply world_size * v.
|
||||
def expected(v):
|
||||
return float(world_size * v)
|
||||
|
||||
if rank == 0:
|
||||
print(
|
||||
f"[repro] world_size={world_size} elems={N} regime={quant_level} fp16",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
# Warmup, then capture a graph with EXACTLY ONE quick-reduce (isolated
|
||||
# qr, so it is the sole writer of its flag slot -- the condition that
|
||||
# triggers the stale-flag bug).
|
||||
inp.fill_(1.0)
|
||||
qar.quick_all_reduce(inp, out=out)
|
||||
torch.accelerator.synchronize()
|
||||
dist.barrier()
|
||||
|
||||
g = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(g):
|
||||
qar.quick_all_reduce(inp, out=out)
|
||||
torch.accelerator.synchronize()
|
||||
dist.barrier()
|
||||
|
||||
for v in range(CUDAGRAPH_ROUNDS):
|
||||
inp.fill_(float(v)) # in-place: same value on every rank
|
||||
dist.barrier()
|
||||
g.replay()
|
||||
torch.accelerator.synchronize()
|
||||
dist.barrier()
|
||||
got = out.float()
|
||||
expect = expected(v)
|
||||
if rank == 0:
|
||||
print(f"round {v}: got={got[:10]}, expected={expect}", flush=True)
|
||||
mismatch = ~torch.isclose(got, torch.full_like(got, expect))
|
||||
num_mismatch = int(mismatch.sum().item())
|
||||
assert num_mismatch == 0, (
|
||||
f"rank={rank} round={v} expected={expect} "
|
||||
f"mismatched {num_mismatch}/{got.numel()} elements; "
|
||||
f"unique wrong values={torch.unique(got[mismatch])[:8].tolist()}"
|
||||
)
|
||||
_log(f"cudagraph worker complete: rank={rank} rounds={CUDAGRAPH_ROUNDS}")
|
||||
finally:
|
||||
if qar is not None:
|
||||
qar.close()
|
||||
if dist.is_initialized():
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
def _run_cudagraph_replay_test(*, world_size: int, quant_level: str):
|
||||
_log(f"launch {world_size}-GPU cudagraph replay case: quant={quant_level}")
|
||||
ctx = mp.get_context("spawn")
|
||||
port = get_open_port()
|
||||
procs = []
|
||||
|
||||
for rank in range(world_size):
|
||||
proc = ctx.Process(
|
||||
target=_quick_allreduce_cudagraph_worker,
|
||||
args=(rank, world_size, port, quant_level),
|
||||
)
|
||||
proc.start()
|
||||
procs.append(proc)
|
||||
|
||||
for proc in procs:
|
||||
proc.join(timeout=120)
|
||||
assert proc.exitcode == 0, f"worker exited with code {proc.exitcode}"
|
||||
_log(f"finished {world_size}-GPU cudagraph replay case: quant={quant_level}")
|
||||
|
||||
|
||||
MODEL_NAME = "Qwen/Qwen2.5-0.5B-Instruct"
|
||||
E2E_PREFILL_TOKENS = 1024
|
||||
E2E_MAX_MODEL_LEN = 1536
|
||||
@@ -210,11 +322,11 @@ def _log_prompt_summaries() -> None:
|
||||
@lru_cache(maxsize=1)
|
||||
def _get_model_path() -> str:
|
||||
try:
|
||||
path = snapshot_download(repo_id=MODEL_NAME, local_files_only=True)
|
||||
path = hf_api().snapshot_download(repo_id=MODEL_NAME, local_files_only=True)
|
||||
_log(f"using cached model snapshot: {path}")
|
||||
return path
|
||||
except Exception:
|
||||
path = snapshot_download(repo_id=MODEL_NAME)
|
||||
path = hf_api().snapshot_download(repo_id=MODEL_NAME)
|
||||
_log(f"downloaded model snapshot: {path}")
|
||||
return path
|
||||
|
||||
@@ -729,6 +841,18 @@ def test_quick_allreduce_two_gpu_correctness(quant_level):
|
||||
)
|
||||
|
||||
|
||||
@multi_gpu_test(num_gpus=CUDAGRAPH_WORLD_SIZE)
|
||||
def test_quick_allreduce_cudagraph_replay():
|
||||
# Regression test for the stale flag_color bug: a quick-reduce captured in a
|
||||
# CUDA graph must return correct results on every replay, not the data from
|
||||
# a previous round.
|
||||
_log("cudagraph replay case")
|
||||
_run_cudagraph_replay_test(
|
||||
world_size=CUDAGRAPH_WORLD_SIZE,
|
||||
quant_level="FP",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
current_platform.device_count() < WORLD_SIZE,
|
||||
reason="requires 2 ROCm GPUs",
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
import multiprocess as mp
|
||||
@@ -11,7 +12,12 @@ import numpy as np
|
||||
import pytest
|
||||
import torch.distributed as dist
|
||||
|
||||
from vllm.distributed.device_communicators.shm_broadcast import MessageQueue
|
||||
from vllm.distributed.device_communicators import shm_broadcast
|
||||
from vllm.distributed.device_communicators.shm_broadcast import (
|
||||
MessageQueue,
|
||||
ShmRingBuffer,
|
||||
check_shm_free_space,
|
||||
)
|
||||
from vllm.distributed.utils import StatelessProcessGroup
|
||||
from vllm.utils.network_utils import get_open_port
|
||||
from vllm.utils.system_utils import update_environment_variables
|
||||
@@ -522,3 +528,39 @@ def test_warning_logs(caplog_vllm):
|
||||
# Clean up when done
|
||||
writer.shutdown()
|
||||
reader.shutdown()
|
||||
|
||||
|
||||
def _fake_disk_usage(free_bytes: int):
|
||||
return SimpleNamespace(total=free_bytes, used=0, free=free_bytes)
|
||||
|
||||
|
||||
def test_check_shm_free_space_raises_when_insufficient(tmp_path):
|
||||
with (
|
||||
mock.patch.object(
|
||||
shm_broadcast.shutil, "disk_usage", return_value=_fake_disk_usage(32 << 20)
|
||||
),
|
||||
pytest.raises(RuntimeError, match="Insufficient space"),
|
||||
):
|
||||
check_shm_free_space(240 << 20, shm_path=str(tmp_path))
|
||||
|
||||
|
||||
def test_check_shm_free_space_passes_when_sufficient(tmp_path):
|
||||
with mock.patch.object(
|
||||
shm_broadcast.shutil, "disk_usage", return_value=_fake_disk_usage(512 << 20)
|
||||
):
|
||||
check_shm_free_space(240 << 20, shm_path=str(tmp_path))
|
||||
|
||||
|
||||
def test_check_shm_free_space_skipped_when_path_missing(tmp_path):
|
||||
check_shm_free_space(1 << 60, shm_path=str(tmp_path / "does-not-exist"))
|
||||
|
||||
|
||||
def test_shm_ring_buffer_creation_checks_free_space():
|
||||
with (
|
||||
mock.patch.object(
|
||||
shm_broadcast.shutil, "disk_usage", return_value=_fake_disk_usage(1 << 20)
|
||||
),
|
||||
mock.patch.object(shm_broadcast.os.path, "isdir", return_value=True),
|
||||
pytest.raises(RuntimeError, match="Insufficient space"),
|
||||
):
|
||||
ShmRingBuffer(n_reader=1, max_chunk_bytes=24 * 1024 * 1024, max_chunks=10)
|
||||
|
||||
@@ -190,30 +190,32 @@ number: "1" | "2"
|
||||
@pytest.fixture(scope="session")
|
||||
def qwen3_lora_files():
|
||||
"""Download Qwen3 LoRA files once per test session."""
|
||||
from huggingface_hub import snapshot_download
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
return snapshot_download(repo_id="charent/self_cognition_Alice")
|
||||
return hf_api().snapshot_download(repo_id="charent/self_cognition_Alice")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def qwen3_meowing_lora_files():
|
||||
"""Download Qwen3 LoRA files once per test session."""
|
||||
from huggingface_hub import snapshot_download
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
return snapshot_download(repo_id="Jackmin108/Qwen3-0.6B-Meow-LoRA")
|
||||
return hf_api().snapshot_download(repo_id="Jackmin108/Qwen3-0.6B-Meow-LoRA")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def qwen3_woofing_lora_files():
|
||||
"""Download Qwen3 LoRA files once per test session."""
|
||||
from huggingface_hub import snapshot_download
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
return snapshot_download(repo_id="Jackmin108/Qwen3-0.6B-Woof-LoRA")
|
||||
return hf_api().snapshot_download(repo_id="Jackmin108/Qwen3-0.6B-Woof-LoRA")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def opt125_lora_files() -> str:
|
||||
"""Download opt-125m LoRA files once per test session."""
|
||||
from huggingface_hub import snapshot_download
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
return snapshot_download(repo_id="peft-internal-testing/opt-125m-dummy-lora")
|
||||
return hf_api().snapshot_download(
|
||||
repo_id="peft-internal-testing/opt-125m-dummy-lora"
|
||||
)
|
||||
|
||||
+5
-3
@@ -12,10 +12,10 @@ import pytest_asyncio
|
||||
import safetensors
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from huggingface_hub import hf_hub_download
|
||||
from transformers import AutoConfig, AutoTokenizer
|
||||
|
||||
from tests.utils import RemoteOpenAIServer
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
from vllm.utils.serial_utils import tensor2base64
|
||||
from vllm.utils.torch_utils import is_torch_equal_or_newer
|
||||
|
||||
@@ -126,11 +126,13 @@ def qwen2audio_aligned_content_and_embeds_b64() -> tuple[str, str]:
|
||||
content = "Describe this audio."
|
||||
tokenizer = AutoTokenizer.from_pretrained(QWEN2AUDIO_MODEL, trust_remote_code=True)
|
||||
|
||||
index_path = hf_hub_download(QWEN2AUDIO_MODEL, "model.safetensors.index.json")
|
||||
index_path = hf_api().hf_hub_download(
|
||||
QWEN2AUDIO_MODEL, "model.safetensors.index.json"
|
||||
)
|
||||
with open(index_path) as f:
|
||||
weight_map = json.load(f)["weight_map"]
|
||||
embed_key = next(k for k in weight_map if k.endswith("embed_tokens.weight"))
|
||||
shard_path = hf_hub_download(QWEN2AUDIO_MODEL, weight_map[embed_key])
|
||||
shard_path = hf_api().hf_hub_download(QWEN2AUDIO_MODEL, weight_map[embed_key])
|
||||
with safetensors.safe_open(shard_path, framework="pt", device="cpu") as f:
|
||||
embed_weight = f.get_tensor(embed_key)
|
||||
embed_layer = nn.Embedding.from_pretrained(embed_weight.to(QWEN2AUDIO_DTYPE))
|
||||
|
||||
+3
-3
@@ -13,12 +13,12 @@ import pytest_asyncio
|
||||
import safetensors
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from huggingface_hub import hf_hub_download
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from tests.utils import RemoteOpenAIServer
|
||||
from vllm.assets.image import ImageAsset
|
||||
from vllm.multimodal.utils import encode_image_url
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
from vllm.utils.serial_utils import tensor2base64
|
||||
|
||||
MODEL_NAME = "Qwen/Qwen2-VL-2B-Instruct"
|
||||
@@ -84,11 +84,11 @@ def aligned_content_and_embeds_b64() -> tuple[str, str]:
|
||||
content = "Describe this image."
|
||||
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, trust_remote_code=True)
|
||||
|
||||
index_path = hf_hub_download(MODEL_NAME, "model.safetensors.index.json")
|
||||
index_path = hf_api().hf_hub_download(MODEL_NAME, "model.safetensors.index.json")
|
||||
with open(index_path) as f:
|
||||
weight_map = json.load(f)["weight_map"]
|
||||
embed_key = next(k for k in weight_map if k.endswith("embed_tokens.weight"))
|
||||
shard_path = hf_hub_download(MODEL_NAME, weight_map[embed_key])
|
||||
shard_path = hf_api().hf_hub_download(MODEL_NAME, weight_map[embed_key])
|
||||
with safetensors.safe_open(shard_path, framework="pt", device="cpu") as f:
|
||||
embed_weight = f.get_tensor(embed_key)
|
||||
embed_layer = nn.Embedding.from_pretrained(embed_weight.to(MODEL_DTYPE))
|
||||
|
||||
@@ -6,17 +6,19 @@ import os
|
||||
import openai # use the official client for correctness check
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
from tests.conftest import AudioTestAssets
|
||||
from tests.utils import RemoteOpenAIServer
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
# NOTE - the tests in this module are currently analogous to test_chat, but are
|
||||
# separated to avoid OOM killing due to module-scoped servers, since we
|
||||
# need a multimodal model for these tests.
|
||||
|
||||
# Contains a modality specific lora alongside the base model
|
||||
MULTIMODAL_MODEL_NAME = snapshot_download("microsoft/Phi-4-multimodal-instruct")
|
||||
MULTIMODAL_MODEL_NAME = hf_api().snapshot_download(
|
||||
"microsoft/Phi-4-multimodal-instruct"
|
||||
)
|
||||
AUDIO_LORA_PATH = os.path.join(MULTIMODAL_MODEL_NAME, "speech-lora")
|
||||
|
||||
ACTIVE_MM_LORA_RESPONSE = "Spoken text: The first words I spoke in the original chronograph, a little piece of practical poetry. Mary had a little lamb, it slept with quite a snow, and everywhere that Mary went, the lamb was sure to go." # noqa: E501
|
||||
|
||||
@@ -27,9 +27,9 @@ MODEL_NAME = "HuggingFaceH4/zephyr-7b-beta"
|
||||
@pytest.fixture(scope="module")
|
||||
def zephyr_lora_files():
|
||||
"""Download zephyr LoRA files once per test session."""
|
||||
from huggingface_hub import snapshot_download
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
return snapshot_download(repo_id="typeof/zephyr-7b-beta-lora")
|
||||
return hf_api().snapshot_download(repo_id="typeof/zephyr-7b-beta-lora")
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
|
||||
@@ -473,7 +473,7 @@ async def test_rerank_api_instruction_field(
|
||||
async def test_rerank_api_instruction_field_matches_chat_template_kwargs(
|
||||
server: tuple[RemoteOpenAIServer, str],
|
||||
):
|
||||
remote_server, _ = server
|
||||
remote_server, backend = server
|
||||
|
||||
doc_list = [
|
||||
document,
|
||||
@@ -514,4 +514,6 @@ async def test_rerank_api_instruction_field_matches_chat_template_kwargs(
|
||||
kwargs_scores = [
|
||||
r.relevance_score for r in sorted(kwargs_rerank.results, key=lambda x: x.index)
|
||||
]
|
||||
assert field_scores == pytest.approx(kwargs_scores)
|
||||
assert field_scores == pytest.approx(
|
||||
kwargs_scores, rel=get_tol(backend), abs=get_abs_tol(backend)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Unit tests for JinaRankingIOProcessor online request building."""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from vllm.entrypoints.pooling.base.io_processor import PoolingIOProcessor
|
||||
from vllm.entrypoints.pooling.scoring.io_processor import JinaRankingIOProcessor
|
||||
from vllm.entrypoints.pooling.scoring.protocol import RerankRequest
|
||||
from vllm.entrypoints.pooling.scoring.typing import ScoringData
|
||||
|
||||
pytestmark = pytest.mark.skip_global_cleanup
|
||||
|
||||
|
||||
def test_online_forwards_truncate_prompt_tokens_to_proxy(monkeypatch):
|
||||
"""The proxy request handed to the base factory must carry
|
||||
truncate_prompt_tokens/truncation_side from the real request.
|
||||
|
||||
JinaRankingIOProcessor swaps ctx.request for a proxy
|
||||
PoolingCompletionRequest before delegating to the base factory, which
|
||||
reads truncation off ctx.request. Dropping the fields on the proxy
|
||||
silently disables truncate_prompt_tokens for Jina rerank/score.
|
||||
"""
|
||||
proc = JinaRankingIOProcessor.__new__(JinaRankingIOProcessor)
|
||||
proc.valid_inputs_online = MagicMock(
|
||||
return_value=ScoringData(data_1=["query"], data_2=["doc"])
|
||||
)
|
||||
proc._get_token_limits = MagicMock(return_value=(0, 0))
|
||||
proc.ensure_str = MagicMock(side_effect=lambda data: list(data))
|
||||
proc.format_docs_prompts_func = MagicMock(return_value="formatted prompt")
|
||||
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _spy_base(self, ctx):
|
||||
captured["truncate_prompt_tokens"] = ctx.request.truncate_prompt_tokens
|
||||
captured["truncation_side"] = ctx.request.truncation_side
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(PoolingIOProcessor, "get_request_factory_online", _spy_base)
|
||||
|
||||
request = RerankRequest(
|
||||
model="m",
|
||||
query="query",
|
||||
documents=["doc"],
|
||||
truncate_prompt_tokens=512,
|
||||
truncation_side="left",
|
||||
)
|
||||
ctx = MagicMock()
|
||||
ctx.request = request
|
||||
ctx.prompt_extras = None
|
||||
|
||||
proc.get_request_factory_online(ctx)
|
||||
|
||||
assert captured["truncate_prompt_tokens"] == 512
|
||||
assert captured["truncation_side"] == "left"
|
||||
# The real request is restored after delegating.
|
||||
assert ctx.request is request
|
||||
@@ -6,7 +6,6 @@ from io import BytesIO
|
||||
import pybase64 as base64
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from huggingface_hub import hf_hub_download
|
||||
from PIL import Image
|
||||
from safetensors.torch import load_file
|
||||
from transformers import AutoModel, AutoTokenizer
|
||||
@@ -18,6 +17,7 @@ from vllm.entrypoints.chat_utils import (
|
||||
)
|
||||
from vllm.entrypoints.pooling.scoring.typing import ScoreMultiModalParam
|
||||
from vllm.entrypoints.pooling.scoring.utils import compute_maxsim_score
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
|
||||
class ColBERTScoringHfRunner(torch.nn.Module):
|
||||
@@ -38,7 +38,7 @@ class ColBERTScoringHfRunner(torch.nn.Module):
|
||||
).to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
path = hf_hub_download(model_name, filename="model.safetensors")
|
||||
path = hf_api().hf_hub_download(model_name, filename="model.safetensors")
|
||||
weights = load_file(path)
|
||||
|
||||
self.linear_weight = weights[linear_weights_key].to(self.device).float()
|
||||
|
||||
@@ -21,9 +21,9 @@ HEADER_SAGEMAKER_NEW_SESSION_ID = "X-Amzn-SageMaker-New-Session-Id"
|
||||
@pytest.fixture(scope="session")
|
||||
def smollm2_lora_files():
|
||||
"""Download LoRA files once per test session."""
|
||||
from huggingface_hub import snapshot_download
|
||||
from vllm.transformers_utils.repo_utils import hf_api
|
||||
|
||||
return snapshot_download(repo_id=LORA_ADAPTER_NAME_SMOLLM)
|
||||
return hf_api().snapshot_download(repo_id=LORA_ADAPTER_NAME_SMOLLM)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user