forked from Karylab-cklius/vllm
Compare commits
137
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e1a763558c | ||
|
|
84aeec9f22 | ||
|
|
82a770ddbd | ||
|
|
a09a9bace1 | ||
|
|
6c20d467a2 | ||
|
|
cb59d0a351 | ||
|
|
0ab1bded36 | ||
|
|
040cbf95cc | ||
|
|
5b3762a7f0 | ||
|
|
4d30c510ce | ||
|
|
6700813f86 | ||
|
|
eb44b3aaa4 | ||
|
|
7a98c7a392 | ||
|
|
0d9e60619b | ||
|
|
1134545b6f | ||
|
|
3e0c887511 | ||
|
|
adfbbc1005 | ||
|
|
adc98f04d0 | ||
|
|
8def3cdde2 | ||
|
|
616c9bd0f4 | ||
|
|
8688a06d67 | ||
|
|
f25953cc59 | ||
|
|
d9aa35161d | ||
|
|
6bcda970fd | ||
|
|
ea0e9c8f2e | ||
|
|
94ed0bf4e0 | ||
|
|
1940c8441e | ||
|
|
72d16aee15 | ||
|
|
e78a0c8e59 | ||
|
|
97a98006b0 | ||
|
|
0a684ab0c0 | ||
|
|
0d9210a502 | ||
|
|
1d874867ea | ||
|
|
2e2e626b40 | ||
|
|
af91f4b3e4 | ||
|
|
2396a61108 | ||
|
|
97a668152b | ||
|
|
58b2012aa2 | ||
|
|
b7c20d0cfa | ||
|
|
a2b1f9fc3b | ||
|
|
642076d26c | ||
|
|
5feb3950e5 | ||
|
|
4ec199b66a | ||
|
|
7ca017778f | ||
|
|
fbfe58133d | ||
|
|
9dd62d80ab | ||
|
|
f878367898 | ||
|
|
bd091079cb | ||
|
|
b23bd73f54 | ||
|
|
e2d7adeb64 | ||
|
|
15cb8e140d | ||
|
|
f007cceb42 | ||
|
|
0a5069e4e3 | ||
|
|
8ce53a616e | ||
|
|
ae10e855ab | ||
|
|
530ee36a0d | ||
|
|
d835ad572c | ||
|
|
47d0597ca2 | ||
|
|
818cf61e91 | ||
|
|
c01618fdc8 | ||
|
|
823eaf667d | ||
|
|
f1f1259692 | ||
|
|
df13b5aef5 | ||
|
|
4938d44a3b | ||
|
|
37bf988c2f | ||
|
|
9459fc6471 | ||
|
|
5245c80564 | ||
|
|
9bc266d923 | ||
|
|
5c9f6557d7 | ||
|
|
dcfebf93f4 | ||
|
|
752bd10647 | ||
|
|
2730b657c4 | ||
|
|
1dcbbd9cac | ||
|
|
ace9fda495 | ||
|
|
ef0aa7ca2f | ||
|
|
e6d1310b2a | ||
|
|
ac5f38a0f7 | ||
|
|
b6ff8a2f50 | ||
|
|
9243e0124e | ||
|
|
df362b2d6d | ||
|
|
7c2acd38b7 | ||
|
|
a287eb163f | ||
|
|
e94243893d | ||
|
|
29c0ec4d63 | ||
|
|
c7ce03bcbd | ||
|
|
c233d90aa8 | ||
|
|
d96aee0951 | ||
|
|
c71a583aa9 | ||
|
|
f12b80c6ef | ||
|
|
da64db78b9 | ||
|
|
425c4eafb0 | ||
|
|
02c01f442b | ||
|
|
fae543015c | ||
|
|
c9be3a8aa1 | ||
|
|
41ea2dd44a | ||
|
|
088c0be268 | ||
|
|
fcd2255d16 | ||
|
|
b5433b6f50 | ||
|
|
cc25f028b7 | ||
|
|
c4cd2bd544 | ||
|
|
5784507da4 | ||
|
|
bf578e1abd | ||
|
|
efed8a1e83 | ||
|
|
11d291511a | ||
|
|
877dae9c68 | ||
|
|
c4dd6d78fd | ||
|
|
ce2aecc4dc | ||
|
|
f38f3d11fb | ||
|
|
d4b4562917 | ||
|
|
7b3192523e | ||
|
|
4c6e2e4b30 | ||
|
|
8502958810 | ||
|
|
ce4bdcbda4 | ||
|
|
d5b1ec2684 | ||
|
|
867ff69733 | ||
|
|
109b736b86 | ||
|
|
69d4f5ef63 | ||
|
|
426d48bfa1 | ||
|
|
26c909ed74 | ||
|
|
fb1d8ccaf5 | ||
|
|
9354f22204 | ||
|
|
17fdd42100 | ||
|
|
472d330c21 | ||
|
|
3b6c96a101 | ||
|
|
4d4e04f452 | ||
|
|
67fe73b2b4 | ||
|
+1 |
ee8f36d0b3 | ||
|
+1 |
f3e9497e92 | ||
|
|
fe784ff22e | ||
|
|
b88abb5036 | ||
|
|
67f9046e4a | ||
|
|
f17be06fbe | ||
|
|
2cab53ddee | ||
|
|
ab0a20d151 | ||
|
|
4a394bfcda | ||
|
|
c95c663049 | ||
|
|
ab3c1aedf3 |
@@ -18,6 +18,8 @@ steps:
|
||||
- tests/kernels/quantization/test_cpu_fp8_scaled_mm.py
|
||||
- tests/kernels/mamba/cpu/test_cpu_gdn_ops.py
|
||||
- tests/kernels/mamba/test_cpu_short_conv.py
|
||||
- tests/kernels/mamba/test_causal_conv1d.py
|
||||
- tests/kernels/mamba/test_mamba_ssm.py
|
||||
commands:
|
||||
- |
|
||||
bash .buildkite/scripts/hardware_ci/run-cpu-test.sh 30m "
|
||||
@@ -28,7 +30,9 @@ steps:
|
||||
pytest -x -v -s tests/kernels/test_onednn.py
|
||||
pytest -x -v -s tests/kernels/test_awq_int4_to_int8.py
|
||||
pytest -x -v -s tests/kernels/quantization/test_cpu_fp8_scaled_mm.py
|
||||
pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py"
|
||||
pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_causal_conv1d.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_mamba_ssm.py"
|
||||
|
||||
# Note: SDE can't be downloaded from CI host because of AWS WAF
|
||||
# - label: CPU-Compatibility Tests
|
||||
|
||||
+397
-367
@@ -31,8 +31,46 @@ steps:
|
||||
- text: "What is the release version?"
|
||||
key: release-version
|
||||
|
||||
- group: "Build Python wheels"
|
||||
- group: "Build CUDA 13.0 Python wheels"
|
||||
key: "build-wheels"
|
||||
steps:
|
||||
- label: "Build wheel - aarch64 - CUDA 13.0"
|
||||
depends_on: ~
|
||||
id: build-wheel-arm64-cuda-13-0
|
||||
agents:
|
||||
queue: arm64_cpu_queue_release
|
||||
commands:
|
||||
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_AARCH64}\" --build-arg BUILD_OS=manylinux --build-arg BUILD_BASE_IMAGE=pytorch/manylinuxaarch64-builder:cuda13.0 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
|
||||
- "mkdir artifacts"
|
||||
- "docker run --rm -v $(pwd)/artifacts:/artifacts_host vllm-ci:build-image bash -c 'cp -r dist /artifacts_host && chmod -R a+rw /artifacts_host'"
|
||||
- "bash .buildkite/scripts/upload-nightly-wheels.sh"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)" release-wheels'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
- label: "Build wheel - x86_64 - CUDA 13.0"
|
||||
depends_on: ~
|
||||
id: build-wheel-x86-cuda-13-0
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_X86}\" --build-arg BUILD_OS=manylinux --build-arg BUILD_BASE_IMAGE=pytorch/manylinux2_28-builder:cuda13.0 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
|
||||
- "mkdir artifacts"
|
||||
- "docker run --rm -v $(pwd)/artifacts:/artifacts_host vllm-ci:build-image bash -c 'cp -r dist /artifacts_host && chmod -R a+rw /artifacts_host'"
|
||||
- "bash .buildkite/scripts/upload-nightly-wheels.sh"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)" release-wheels'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
- block: "Unblock to build additional Python wheels"
|
||||
depends_on: ~
|
||||
key: block-build-additional-wheels
|
||||
if: build.env("NIGHTLY") != "1"
|
||||
|
||||
- group: "Build additional Python wheels"
|
||||
key: "build-additional-wheels"
|
||||
depends_on: block-build-additional-wheels
|
||||
allow_dependency_failure: true
|
||||
steps:
|
||||
- label: "Build wheel - aarch64 - CUDA 12.9"
|
||||
depends_on: ~
|
||||
@@ -48,20 +86,6 @@ steps:
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
- label: "Build wheel - aarch64 - CUDA 13.0"
|
||||
depends_on: ~
|
||||
id: build-wheel-arm64-cuda-13-0
|
||||
agents:
|
||||
queue: arm64_cpu_queue_release
|
||||
commands:
|
||||
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_AARCH64}\" --build-arg BUILD_OS=manylinux --build-arg BUILD_BASE_IMAGE=pytorch/manylinuxaarch64-builder:cuda13.0 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
|
||||
- "mkdir artifacts"
|
||||
- "docker run --rm -v $(pwd)/artifacts:/artifacts_host vllm-ci:build-image bash -c 'cp -r dist /artifacts_host && chmod -R a+rw /artifacts_host'"
|
||||
- "bash .buildkite/scripts/upload-nightly-wheels.sh"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)" release-wheels'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
- label: "Build wheel - aarch64 - CPU"
|
||||
depends_on: ~
|
||||
id: build-wheel-arm64-cpu
|
||||
@@ -113,7 +137,7 @@ steps:
|
||||
- 'mv artifacts/reassembled/wheel "artifacts/dist/$$wheel_name"'
|
||||
- "aws sts get-caller-identity"
|
||||
- "VLLM_WHEEL_PLATFORM=macos bash .buildkite/scripts/upload-nightly-wheels.sh"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)"'
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)" release-wheels'
|
||||
plugins:
|
||||
- aws-assume-role-with-web-identity#v1.6.0:
|
||||
role-arn: arn:aws:iam::936637512419:role/vllm-release-macos-wheel-uploader
|
||||
@@ -133,20 +157,6 @@ steps:
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
- label: "Build wheel - x86_64 - CUDA 13.0"
|
||||
depends_on: ~
|
||||
id: build-wheel-x86-cuda-13-0
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg USE_SCCACHE=1 --build-arg GIT_REPO_CHECK=1 --build-arg CUDA_VERSION=13.0.2 --build-arg torch_cuda_arch_list=\"${CUDA_ARCH_X86}\" --build-arg BUILD_OS=manylinux --build-arg BUILD_BASE_IMAGE=pytorch/manylinux2_28-builder:cuda13.0 --tag vllm-ci:build-image --target build --progress plain -f docker/Dockerfile ."
|
||||
- "mkdir artifacts"
|
||||
- "docker run --rm -v $(pwd)/artifacts:/artifacts_host vllm-ci:build-image bash -c 'cp -r dist /artifacts_host && chmod -R a+rw /artifacts_host'"
|
||||
- "bash .buildkite/scripts/upload-nightly-wheels.sh"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)" release-wheels'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
- label: "Build wheel - x86_64 - CPU"
|
||||
depends_on: ~
|
||||
id: build-wheel-x86-cpu
|
||||
@@ -162,12 +172,26 @@ steps:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
- label: "Generate and upload wheel indices"
|
||||
key: generate-wheel-indices
|
||||
depends_on: "build-wheels"
|
||||
allow_dependency_failure: true
|
||||
if: build.env("NIGHTLY") != "1"
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "UPDATE_VERSION_INDEX=0 bash .buildkite/scripts/generate-and-upload-nightly-index.sh"
|
||||
|
||||
- label: "Regenerate indices with additional wheels"
|
||||
key: generate-additional-wheel-indices
|
||||
depends_on:
|
||||
- build-wheels
|
||||
- build-additional-wheels
|
||||
- generate-wheel-indices
|
||||
allow_dependency_failure: true
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/generate-and-upload-nightly-index.sh"
|
||||
- 'UPDATE_NIGHTLY_INDEX="$${NIGHTLY:-0}" bash .buildkite/scripts/generate-and-upload-nightly-index.sh'
|
||||
|
||||
- block: "Unblock to build release Docker images"
|
||||
depends_on: ~
|
||||
@@ -566,366 +590,370 @@ steps:
|
||||
#
|
||||
# =============================================================================
|
||||
|
||||
# ROCm Job 1: Build ROCm Base Wheels (with S3 caching)
|
||||
- label: ":rocm: Build ROCm Base Image & Wheels"
|
||||
id: build-rocm-base-wheels
|
||||
- group: "Build ROCm Wheel / Image "
|
||||
key: "build-rocm-wheel-image"
|
||||
depends_on: ~
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- |
|
||||
set -euo pipefail
|
||||
steps:
|
||||
# ROCm Job 1: Build ROCm Base Wheels (with S3 caching)
|
||||
- label: ":rocm: Build ROCm Base Image & Wheels"
|
||||
id: build-rocm-base-wheels
|
||||
depends_on: ~
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- |
|
||||
set -euo pipefail
|
||||
|
||||
# Generate cache key
|
||||
CACHE_KEY=$$(.buildkite/scripts/cache-rocm-base-wheels.sh key)
|
||||
ECR_CACHE_TAG="public.ecr.aws/q9t5s3a7/vllm-release-repo:$${CACHE_KEY}-rocm-base"
|
||||
# Generate cache key
|
||||
CACHE_KEY=$$(.buildkite/scripts/cache-rocm-base-wheels.sh key)
|
||||
ECR_CACHE_TAG="public.ecr.aws/q9t5s3a7/vllm-release-repo:$${CACHE_KEY}-rocm-base"
|
||||
|
||||
echo "========================================"
|
||||
echo "ROCm Base Build Configuration"
|
||||
echo "========================================"
|
||||
echo " CACHE_KEY: $${CACHE_KEY}"
|
||||
echo " ECR_CACHE_TAG: $${ECR_CACHE_TAG}"
|
||||
echo "========================================"
|
||||
|
||||
# Login to ECR
|
||||
aws ecr-public get-login-password --region us-east-1 | \
|
||||
docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7
|
||||
|
||||
IMAGE_EXISTS=false
|
||||
WHEELS_EXIST=false
|
||||
|
||||
# Check ECR for Docker image
|
||||
echo "========================================"
|
||||
echo "ROCm Base Build Configuration"
|
||||
echo "========================================"
|
||||
echo " CACHE_KEY: $${CACHE_KEY}"
|
||||
echo " ECR_CACHE_TAG: $${ECR_CACHE_TAG}"
|
||||
echo "========================================"
|
||||
|
||||
# Login to ECR
|
||||
aws ecr-public get-login-password --region us-east-1 | \
|
||||
docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7
|
||||
|
||||
IMAGE_EXISTS=false
|
||||
WHEELS_EXIST=false
|
||||
|
||||
# Check ECR for Docker image
|
||||
|
||||
if docker manifest inspect "$${ECR_CACHE_TAG}" > /dev/null 2>&1; then
|
||||
IMAGE_EXISTS=true
|
||||
echo "ECR image cache HIT"
|
||||
fi
|
||||
|
||||
# Check S3 for wheels
|
||||
WHEEL_CACHE_STATUS=$(.buildkite/scripts/cache-rocm-base-wheels.sh check)
|
||||
if [ "$${WHEEL_CACHE_STATUS}" = "hit" ]; then
|
||||
WHEELS_EXIST=true
|
||||
echo "S3 wheels cache HIT"
|
||||
fi
|
||||
if docker manifest inspect "$${ECR_CACHE_TAG}" > /dev/null 2>&1; then
|
||||
IMAGE_EXISTS=true
|
||||
echo "ECR image cache HIT"
|
||||
fi
|
||||
|
||||
# Check S3 for wheels
|
||||
WHEEL_CACHE_STATUS=$(.buildkite/scripts/cache-rocm-base-wheels.sh check)
|
||||
if [ "$${WHEEL_CACHE_STATUS}" = "hit" ]; then
|
||||
WHEELS_EXIST=true
|
||||
echo "S3 wheels cache HIT"
|
||||
fi
|
||||
|
||||
|
||||
# Scenario 1: Both cached (best case)
|
||||
if [ "$${IMAGE_EXISTS}" = "true" ] && [ "$${WHEELS_EXIST}" = "true" ]; then
|
||||
echo ""
|
||||
echo "FULL CACHE HIT - Reusing both image and wheels"
|
||||
echo ""
|
||||
|
||||
# Scenario 1: Both cached (best case)
|
||||
if [ "$${IMAGE_EXISTS}" = "true" ] && [ "$${WHEELS_EXIST}" = "true" ]; then
|
||||
echo ""
|
||||
echo "FULL CACHE HIT - Reusing both image and wheels"
|
||||
echo ""
|
||||
|
||||
# Download wheels
|
||||
.buildkite/scripts/cache-rocm-base-wheels.sh download
|
||||
|
||||
# Save ECR tag for downstream jobs
|
||||
buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}"
|
||||
|
||||
# Scenario 2: Full rebuild needed
|
||||
else
|
||||
echo ""
|
||||
echo " CACHE MISS - Building from scratch..."
|
||||
echo ""
|
||||
|
||||
# Build full base image and push to ECR
|
||||
DOCKER_BUILDKIT=1 docker buildx build \
|
||||
--file docker/Dockerfile.rocm_base \
|
||||
--tag "$${ECR_CACHE_TAG}" \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
--push \
|
||||
.
|
||||
|
||||
# Build wheel extraction stage
|
||||
DOCKER_BUILDKIT=1 docker buildx build \
|
||||
--file docker/Dockerfile.rocm_base \
|
||||
--tag rocm-base-debs:$${BUILDKITE_BUILD_NUMBER} \
|
||||
--target debs_wheel_release \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
--load \
|
||||
.
|
||||
|
||||
# Extract and upload wheels
|
||||
mkdir -p artifacts/rocm-base-wheels
|
||||
cid=$(docker create rocm-base-debs:$${BUILDKITE_BUILD_NUMBER})
|
||||
docker cp $${cid}:/app/debs/. artifacts/rocm-base-wheels/
|
||||
docker rm $${cid}
|
||||
|
||||
.buildkite/scripts/cache-rocm-base-wheels.sh upload
|
||||
# Download wheels
|
||||
.buildkite/scripts/cache-rocm-base-wheels.sh download
|
||||
|
||||
# Save ECR tag for downstream jobs
|
||||
buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}"
|
||||
|
||||
# Scenario 2: Full rebuild needed
|
||||
else
|
||||
echo ""
|
||||
echo " CACHE MISS - Building from scratch..."
|
||||
echo ""
|
||||
|
||||
# Build full base image and push to ECR
|
||||
DOCKER_BUILDKIT=1 docker buildx build \
|
||||
--file docker/Dockerfile.rocm_base \
|
||||
--tag "$${ECR_CACHE_TAG}" \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
--push \
|
||||
.
|
||||
|
||||
# Build wheel extraction stage
|
||||
DOCKER_BUILDKIT=1 docker buildx build \
|
||||
--file docker/Dockerfile.rocm_base \
|
||||
--tag rocm-base-debs:$${BUILDKITE_BUILD_NUMBER} \
|
||||
--target debs_wheel_release \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
--load \
|
||||
.
|
||||
|
||||
# Extract and upload wheels
|
||||
mkdir -p artifacts/rocm-base-wheels
|
||||
cid=$(docker create rocm-base-debs:$${BUILDKITE_BUILD_NUMBER})
|
||||
docker cp $${cid}:/app/debs/. artifacts/rocm-base-wheels/
|
||||
docker rm $${cid}
|
||||
|
||||
.buildkite/scripts/cache-rocm-base-wheels.sh upload
|
||||
|
||||
# Cache base docker image to ECR
|
||||
docker push "$${ECR_CACHE_TAG}"
|
||||
|
||||
buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}"
|
||||
|
||||
echo ""
|
||||
echo " Build complete - Image and wheels cached"
|
||||
fi
|
||||
# Cache base docker image to ECR
|
||||
docker push "$${ECR_CACHE_TAG}"
|
||||
|
||||
buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}"
|
||||
|
||||
echo ""
|
||||
echo " Build complete - Image and wheels cached"
|
||||
fi
|
||||
|
||||
artifact_paths:
|
||||
- "artifacts/rocm-base-wheels/*.whl"
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
artifact_paths:
|
||||
- "artifacts/rocm-base-wheels/*.whl"
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
|
||||
# ROCm Job 2: Build vLLM ROCm Wheel
|
||||
- label: ":python: Build vLLM ROCm Wheel - x86_64"
|
||||
id: build-rocm-vllm-wheel
|
||||
depends_on:
|
||||
- step: build-rocm-base-wheels
|
||||
allow_failure: false
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
timeout_in_minutes: 180
|
||||
commands:
|
||||
# Download artifacts and prepare Docker image
|
||||
- |
|
||||
set -euo pipefail
|
||||
# ROCm Job 2: Build vLLM ROCm Wheel
|
||||
- label: ":python: Build vLLM ROCm Wheel - x86_64"
|
||||
id: build-rocm-vllm-wheel
|
||||
depends_on:
|
||||
- step: build-rocm-base-wheels
|
||||
allow_failure: false
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
timeout_in_minutes: 180
|
||||
commands:
|
||||
# Download artifacts and prepare Docker image
|
||||
- |
|
||||
set -euo pipefail
|
||||
|
||||
# Ensure git tags are up-to-date (Buildkite's default fetch doesn't update tags)
|
||||
# This fixes version detection when tags are moved/force-pushed
|
||||
echo "Fetching latest tags from origin..."
|
||||
git fetch --tags --force origin
|
||||
|
||||
# Log tag information for debugging version detection
|
||||
echo "========================================"
|
||||
echo "Git Tag Verification"
|
||||
echo "========================================"
|
||||
echo "Current HEAD: $(git rev-parse HEAD)"
|
||||
echo "git describe --tags: $(git describe --tags 2>/dev/null || echo 'No tags found')"
|
||||
echo ""
|
||||
echo "Recent tags (pointing to commits near HEAD):"
|
||||
git tag -l --sort=-creatordate | head -5
|
||||
echo "setuptools_scm version detection:"
|
||||
pip install -q setuptools_scm 2>/dev/null || true
|
||||
python3 -c "import setuptools_scm; print(' Detected version:', setuptools_scm.get_version())" 2>/dev/null || echo " (setuptools_scm not available in this environment)"
|
||||
echo "========================================"
|
||||
# Ensure git tags are up-to-date (Buildkite's default fetch doesn't update tags)
|
||||
# This fixes version detection when tags are moved/force-pushed
|
||||
echo "Fetching latest tags from origin..."
|
||||
git fetch --tags --force origin
|
||||
|
||||
# Log tag information for debugging version detection
|
||||
echo "========================================"
|
||||
echo "Git Tag Verification"
|
||||
echo "========================================"
|
||||
echo "Current HEAD: $(git rev-parse HEAD)"
|
||||
echo "git describe --tags: $(git describe --tags 2>/dev/null || echo 'No tags found')"
|
||||
echo ""
|
||||
echo "Recent tags (pointing to commits near HEAD):"
|
||||
git tag -l --sort=-creatordate | head -5
|
||||
echo "setuptools_scm version detection:"
|
||||
pip install -q setuptools_scm 2>/dev/null || true
|
||||
python3 -c "import setuptools_scm; print(' Detected version:', setuptools_scm.get_version())" 2>/dev/null || echo " (setuptools_scm not available in this environment)"
|
||||
echo "========================================"
|
||||
|
||||
# Download wheel artifacts from current build
|
||||
echo "Downloading wheel artifacts from current build"
|
||||
buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" .
|
||||
# Download wheel artifacts from current build
|
||||
echo "Downloading wheel artifacts from current build"
|
||||
buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" .
|
||||
|
||||
# Get ECR image tag from metadata (set by build-rocm-base-wheels)
|
||||
ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')"
|
||||
if [ -z "$${ECR_IMAGE_TAG}" ]; then
|
||||
echo "ERROR: rocm-base-image-tag metadata not found"
|
||||
echo "This should have been set by the build-rocm-base-wheels job"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Login to ECR
|
||||
aws ecr-public get-login-password --region us-east-1 | \
|
||||
docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7
|
||||
|
||||
# Pull base Docker image from ECR
|
||||
docker pull "$${ECR_IMAGE_TAG}"
|
||||
|
||||
echo "Loaded base image: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Prepare base wheels for Docker build context
|
||||
mkdir -p docker/context/base-wheels
|
||||
touch docker/context/base-wheels/.keep
|
||||
cp artifacts/rocm-base-wheels/*.whl docker/context/base-wheels/
|
||||
echo "Base wheels for vLLM build:"
|
||||
ls -lh docker/context/base-wheels/
|
||||
# Get ECR image tag from metadata (set by build-rocm-base-wheels)
|
||||
ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')"
|
||||
if [ -z "$${ECR_IMAGE_TAG}" ]; then
|
||||
echo "ERROR: rocm-base-image-tag metadata not found"
|
||||
echo "This should have been set by the build-rocm-base-wheels job"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Login to ECR
|
||||
aws ecr-public get-login-password --region us-east-1 | \
|
||||
docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7
|
||||
|
||||
# Pull base Docker image from ECR
|
||||
docker pull "$${ECR_IMAGE_TAG}"
|
||||
|
||||
echo "Loaded base image: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Prepare base wheels for Docker build context
|
||||
mkdir -p docker/context/base-wheels
|
||||
touch docker/context/base-wheels/.keep
|
||||
cp artifacts/rocm-base-wheels/*.whl docker/context/base-wheels/
|
||||
echo "Base wheels for vLLM build:"
|
||||
ls -lh docker/context/base-wheels/
|
||||
|
||||
echo "========================================"
|
||||
echo "Building vLLM wheel with:"
|
||||
echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}"
|
||||
echo " BUILDKITE_BRANCH: $${BUILDKITE_BRANCH}"
|
||||
echo " BASE_IMAGE: $${ECR_IMAGE_TAG}"
|
||||
echo "========================================"
|
||||
echo "========================================"
|
||||
echo "Building vLLM wheel with:"
|
||||
echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}"
|
||||
echo " BUILDKITE_BRANCH: $${BUILDKITE_BRANCH}"
|
||||
echo " BASE_IMAGE: $${ECR_IMAGE_TAG}"
|
||||
echo "========================================"
|
||||
|
||||
# Build vLLM wheel using local checkout (REMOTE_VLLM=0)
|
||||
DOCKER_BUILDKIT=1 docker build \
|
||||
--file docker/Dockerfile.rocm \
|
||||
--target export_vllm_wheel_release \
|
||||
--output type=local,dest=rocm-dist \
|
||||
--build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \
|
||||
--build-arg REMOTE_VLLM=0 \
|
||||
--build-arg GIT_REPO_CHECK=1 \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
.
|
||||
echo "Built vLLM wheel:"
|
||||
ls -lh rocm-dist/*.whl
|
||||
# Copy wheel to artifacts directory
|
||||
mkdir -p artifacts/rocm-vllm-wheel
|
||||
cp rocm-dist/*.whl artifacts/rocm-vllm-wheel/
|
||||
echo "Final vLLM wheel:"
|
||||
ls -lh artifacts/rocm-vllm-wheel/
|
||||
artifact_paths:
|
||||
- "artifacts/rocm-vllm-wheel/*.whl"
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
# Build vLLM wheel using local checkout (REMOTE_VLLM=0)
|
||||
DOCKER_BUILDKIT=1 docker build \
|
||||
--file docker/Dockerfile.rocm \
|
||||
--target export_vllm_wheel_release \
|
||||
--output type=local,dest=rocm-dist \
|
||||
--build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \
|
||||
--build-arg REMOTE_VLLM=0 \
|
||||
--build-arg GIT_REPO_CHECK=1 \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
.
|
||||
echo "Built vLLM wheel:"
|
||||
ls -lh rocm-dist/*.whl
|
||||
# Copy wheel to artifacts directory
|
||||
mkdir -p artifacts/rocm-vllm-wheel
|
||||
cp rocm-dist/*.whl artifacts/rocm-vllm-wheel/
|
||||
echo "Final vLLM wheel:"
|
||||
ls -lh artifacts/rocm-vllm-wheel/
|
||||
artifact_paths:
|
||||
- "artifacts/rocm-vllm-wheel/*.whl"
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
|
||||
# ROCm Job 3: Upload Wheels to S3
|
||||
- label: ":s3: Upload ROCm Wheels to S3"
|
||||
id: upload-rocm-wheels
|
||||
depends_on:
|
||||
- step: build-rocm-vllm-wheel
|
||||
allow_failure: false
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
timeout_in_minutes: 60
|
||||
commands:
|
||||
# Download all wheel artifacts and run upload
|
||||
- |
|
||||
set -euo pipefail
|
||||
# ROCm Job 3: Upload Wheels to S3
|
||||
- label: ":s3: Upload ROCm Wheels to S3"
|
||||
id: upload-rocm-wheels
|
||||
depends_on:
|
||||
- step: build-rocm-vllm-wheel
|
||||
allow_failure: false
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
timeout_in_minutes: 60
|
||||
commands:
|
||||
# Download all wheel artifacts and run upload
|
||||
- |
|
||||
set -euo pipefail
|
||||
|
||||
# Download artifacts from current build
|
||||
echo "Downloading artifacts from current build"
|
||||
buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" .
|
||||
buildkite-agent artifact download "artifacts/rocm-vllm-wheel/*.whl" .
|
||||
# Download artifacts from current build
|
||||
echo "Downloading artifacts from current build"
|
||||
buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" .
|
||||
buildkite-agent artifact download "artifacts/rocm-vllm-wheel/*.whl" .
|
||||
|
||||
# Run upload script
|
||||
bash .buildkite/scripts/upload-rocm-wheels.sh
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
# # Run upload script
|
||||
bash .buildkite/scripts/upload-rocm-wheels.sh
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
|
||||
# ROCm Job 4: Annotate ROCm Wheel Release
|
||||
- label: ":memo: Annotate ROCm wheel release"
|
||||
id: annotate-rocm-release
|
||||
depends_on:
|
||||
- upload-rocm-wheels
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/annotate-rocm-release.sh"
|
||||
env:
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
# ROCm Job 4: Annotate ROCm Wheel Release
|
||||
- label: ":memo: Annotate ROCm wheel release"
|
||||
id: annotate-rocm-release
|
||||
depends_on:
|
||||
- upload-rocm-wheels
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/annotate-rocm-release.sh"
|
||||
env:
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
|
||||
# ROCm Job 5: Generate Root Index for ROCm Wheels (for release only)
|
||||
# This is the job to create https://wheels.vllm.ai/rocm/ index allowing
|
||||
# users to install with `uv pip install vllm --extra-index-url https://wheels.vllm.ai/rocm/`
|
||||
- block: "Generate Root Index for ROCm Wheels for Release"
|
||||
key: block-generate-root-index-rocm-wheels
|
||||
depends_on: upload-rocm-wheels
|
||||
# ROCm Job 5: Generate Root Index for ROCm Wheels (for release only)
|
||||
# This is the job to create https://wheels.vllm.ai/rocm/ index allowing
|
||||
# users to install with `uv pip install vllm --extra-index-url https://wheels.vllm.ai/rocm/`
|
||||
- block: "Generate Root Index for ROCm Wheels for Release"
|
||||
key: block-generate-root-index-rocm-wheels
|
||||
depends_on: upload-rocm-wheels
|
||||
|
||||
- label: ":package: Generate Root Index for ROCm Wheels for Release"
|
||||
depends_on: block-generate-root-index-rocm-wheels
|
||||
id: generate-root-index-rocm-wheels
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "bash tools/vllm-rocm/generate-rocm-wheels-root-index.sh"
|
||||
env:
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
VARIANT: "rocm723"
|
||||
- label: ":package: Generate Root Index for ROCm Wheels for Release"
|
||||
depends_on: block-generate-root-index-rocm-wheels
|
||||
id: generate-root-index-rocm-wheels
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
commands:
|
||||
- "bash tools/vllm-rocm/generate-rocm-wheels-root-index.sh"
|
||||
env:
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
VARIANT: "rocm723"
|
||||
|
||||
# ROCm Job 6: Build ROCm Release Docker Image
|
||||
- label: ":docker: Build release image - x86_64 - ROCm"
|
||||
id: build-rocm-release-image
|
||||
depends_on:
|
||||
- step: block-build-release-images
|
||||
allow_failure: true
|
||||
- step: build-rocm-base-wheels
|
||||
allow_failure: false
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
timeout_in_minutes: 60
|
||||
commands:
|
||||
- |
|
||||
set -euo pipefail
|
||||
|
||||
# Login to ECR
|
||||
aws ecr-public get-login-password --region us-east-1 | \
|
||||
docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7
|
||||
|
||||
# Get ECR image tag from metadata (set by build-rocm-base-wheels)
|
||||
ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')"
|
||||
if [ -z "$${ECR_IMAGE_TAG}" ]; then
|
||||
echo "ERROR: rocm-base-image-tag metadata not found"
|
||||
echo "This should have been set by the build-rocm-base-wheels job"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Pull base Docker image from ECR
|
||||
docker pull "$${ECR_IMAGE_TAG}"
|
||||
|
||||
echo "Loaded base image: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Pass the base image ECR tag to downstream steps (nightly publish)
|
||||
buildkite-agent meta-data set "rocm-base-ecr-tag" "$${ECR_IMAGE_TAG}"
|
||||
|
||||
echo "========================================"
|
||||
echo "Building vLLM ROCm release image with:"
|
||||
echo " BASE_IMAGE: $${ECR_IMAGE_TAG}"
|
||||
echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}"
|
||||
echo "========================================"
|
||||
|
||||
# Build vLLM ROCm release image using cached base
|
||||
DOCKER_BUILDKIT=1 docker build \
|
||||
--build-arg max_jobs=16 \
|
||||
--build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
--tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm \
|
||||
--target vllm-openai \
|
||||
--progress plain \
|
||||
-f docker/Dockerfile.rocm .
|
||||
|
||||
# Push to ECR
|
||||
docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm
|
||||
# ROCm Job 6: Build ROCm Release Docker Image
|
||||
- label: ":docker: Build release image - x86_64 - ROCm"
|
||||
id: build-rocm-release-image
|
||||
depends_on:
|
||||
- step: block-build-release-images
|
||||
allow_failure: true
|
||||
- step: build-rocm-base-wheels
|
||||
allow_failure: false
|
||||
agents:
|
||||
queue: cpu_queue_release
|
||||
timeout_in_minutes: 60
|
||||
commands:
|
||||
- |
|
||||
set -euo pipefail
|
||||
|
||||
# Login to ECR
|
||||
aws ecr-public get-login-password --region us-east-1 | \
|
||||
docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7
|
||||
|
||||
# Get ECR image tag from metadata (set by build-rocm-base-wheels)
|
||||
ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')"
|
||||
if [ -z "$${ECR_IMAGE_TAG}" ]; then
|
||||
echo "ERROR: rocm-base-image-tag metadata not found"
|
||||
echo "This should have been set by the build-rocm-base-wheels job"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Pull base Docker image from ECR
|
||||
docker pull "$${ECR_IMAGE_TAG}"
|
||||
|
||||
echo "Loaded base image: $${ECR_IMAGE_TAG}"
|
||||
|
||||
# Pass the base image ECR tag to downstream steps (nightly publish)
|
||||
buildkite-agent meta-data set "rocm-base-ecr-tag" "$${ECR_IMAGE_TAG}"
|
||||
|
||||
echo "========================================"
|
||||
echo "Building vLLM ROCm release image with:"
|
||||
echo " BASE_IMAGE: $${ECR_IMAGE_TAG}"
|
||||
echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}"
|
||||
echo "========================================"
|
||||
|
||||
# Build vLLM ROCm release image using cached base
|
||||
DOCKER_BUILDKIT=1 docker build \
|
||||
--build-arg max_jobs=16 \
|
||||
--build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \
|
||||
--build-arg USE_SCCACHE=1 \
|
||||
--build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \
|
||||
--build-arg SCCACHE_REGION_NAME=us-west-2 \
|
||||
--build-arg SCCACHE_S3_NO_CREDENTIALS=0 \
|
||||
--tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm \
|
||||
--target vllm-openai \
|
||||
--progress plain \
|
||||
-f docker/Dockerfile.rocm .
|
||||
|
||||
# Push to ECR
|
||||
docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm
|
||||
|
||||
echo ""
|
||||
echo " Successfully built and pushed ROCm release image"
|
||||
echo " Image: public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm"
|
||||
echo ""
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
echo ""
|
||||
echo " Successfully built and pushed ROCm release image"
|
||||
echo " Image: public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm"
|
||||
echo ""
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
S3_BUCKET: "vllm-wheels"
|
||||
|
||||
- label: "Publish nightly XPU image to DockerHub"
|
||||
depends_on:
|
||||
- create-manifest-xpu
|
||||
if: build.env("NIGHTLY") == "1"
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/xpu/push-nightly-builds-xpu.sh"
|
||||
- "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-xpu"
|
||||
plugins:
|
||||
- docker-login#v3.0.0:
|
||||
username: vllmbot
|
||||
password-env: DOCKERHUB_TOKEN
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
- label: "Publish nightly XPU image to DockerHub"
|
||||
depends_on:
|
||||
- create-manifest-xpu
|
||||
if: build.env("NIGHTLY") == "1"
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/xpu/push-nightly-builds-xpu.sh"
|
||||
- "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-xpu"
|
||||
plugins:
|
||||
- docker-login#v3.0.0:
|
||||
username: vllmbot
|
||||
password-env: DOCKERHUB_TOKEN
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
|
||||
- label: "Publish nightly ROCm image to DockerHub"
|
||||
depends_on:
|
||||
- build-rocm-release-image
|
||||
if: build.env("NIGHTLY") == "1"
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/push-nightly-builds-rocm.sh"
|
||||
# Clean up old nightly builds (keep only last 14)
|
||||
- "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-rocm"
|
||||
- "bash .buildkite/scripts/cleanup-nightly-builds.sh base-nightly- vllm/vllm-openai-rocm"
|
||||
plugins:
|
||||
- docker-login#v3.0.0:
|
||||
username: vllmbot
|
||||
password-env: DOCKERHUB_TOKEN
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
- label: "Publish nightly ROCm image to DockerHub"
|
||||
depends_on:
|
||||
- build-rocm-release-image
|
||||
if: build.env("NIGHTLY") == "1"
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/push-nightly-builds-rocm.sh"
|
||||
# Clean up old nightly builds (keep only last 14)
|
||||
- "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-rocm"
|
||||
- "bash .buildkite/scripts/cleanup-nightly-builds.sh base-nightly- vllm/vllm-openai-rocm"
|
||||
plugins:
|
||||
- docker-login#v3.0.0:
|
||||
username: vllmbot
|
||||
password-env: DOCKERHUB_TOKEN
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
|
||||
# =============================================================================
|
||||
# Publish to DockerHub and PyPI (at the end so all builds complete first)
|
||||
@@ -974,6 +1002,8 @@ steps:
|
||||
depends_on:
|
||||
- input-release-version
|
||||
- build-wheels
|
||||
- build-additional-wheels
|
||||
- generate-additional-wheel-indices
|
||||
|
||||
- label: "Upload release wheels to PyPI"
|
||||
depends_on:
|
||||
|
||||
@@ -45,8 +45,10 @@ $PYTHON .buildkite/scripts/generate-nightly-index.py --version "$SUBPATH" --curr
|
||||
echo "Uploading indices to $S3_COMMIT_PREFIX"
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "$S3_COMMIT_PREFIX"
|
||||
|
||||
# copy to /nightly/ only if it is on the main branch and not a PR
|
||||
if [[ "$BUILDKITE_BRANCH" == "main" && "$BUILDKITE_PULL_REQUEST" == "false" ]]; then
|
||||
# copy to /nightly/ only when enabled for a main branch build that is not a PR
|
||||
if [[ "${UPDATE_NIGHTLY_INDEX:-1}" == "1" && \
|
||||
"$BUILDKITE_BRANCH" == "main" && \
|
||||
"$BUILDKITE_PULL_REQUEST" == "false" ]]; then
|
||||
echo "Uploading indices to overwrite /nightly/"
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "s3://$BUCKET/nightly/"
|
||||
fi
|
||||
@@ -67,7 +69,7 @@ pure_version="${version%%+*}"
|
||||
echo "Pure version (without variant): $pure_version"
|
||||
|
||||
# re-generate and copy to /<pure_version>/ only if it does not have "dev" in the version
|
||||
if [[ "$version" != *"dev"* ]]; then
|
||||
if [[ "${UPDATE_VERSION_INDEX:-1}" == "1" && "$version" != *"dev"* ]]; then
|
||||
echo "Re-generating indices for /$pure_version/"
|
||||
rm -rf "${INDICES_OUTPUT_DIR:?}"
|
||||
mkdir -p "$INDICES_OUTPUT_DIR"
|
||||
|
||||
@@ -40,7 +40,9 @@ function cpu_tests() {
|
||||
pytest -x -v -s tests/kernels/moe/test_cpu_fused_moe.py
|
||||
pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py
|
||||
pytest -x -v -s tests/kernels/moe/test_cpu_int4_moe.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_cpu_short_conv.py"
|
||||
pytest -x -v -s tests/kernels/mamba/test_cpu_short_conv.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_causal_conv1d.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_mamba_ssm.py"
|
||||
|
||||
# skip tests requiring model downloads if HF_TOKEN is not set
|
||||
# due to rate-limits
|
||||
@@ -97,3 +99,4 @@ function cpu_tests() {
|
||||
# All of CPU tests are expected to be finished less than 40 mins.
|
||||
export -f cpu_tests
|
||||
timeout 2h bash -c cpu_tests
|
||||
|
||||
|
||||
@@ -29,6 +29,7 @@ PYO3_PYTHON_VERSION="${PYO3_PYTHON_VERSION:-3.12}"
|
||||
CARGO_SORT_VERSION_REQ="${CARGO_SORT_VERSION_REQ:-2}"
|
||||
CARGO_DENY_VERSION_REQ="${CARGO_DENY_VERSION_REQ:-0.20}"
|
||||
CARGO_NEXTEST_VERSION_REQ="${CARGO_NEXTEST_VERSION_REQ:-0.9}"
|
||||
CARGO_LLVM_COV_VERSION="${CARGO_LLVM_COV_VERSION:-0.8.7}"
|
||||
|
||||
log_section() {
|
||||
echo "--- $*"
|
||||
@@ -106,6 +107,18 @@ install_cargo_nextest() {
|
||||
"cargo-nextest@${CARGO_NEXTEST_VERSION_REQ}"
|
||||
}
|
||||
|
||||
install_cargo_llvm_cov() {
|
||||
log_section "Installing cargo-llvm-cov ${CARGO_LLVM_COV_VERSION}"
|
||||
local toolchain
|
||||
toolchain="$(rust_toolchain)"
|
||||
rustup component add --toolchain "$toolchain" llvm-tools-preview
|
||||
cargo binstall \
|
||||
--no-confirm \
|
||||
--force \
|
||||
--secure \
|
||||
"cargo-llvm-cov@${CARGO_LLVM_COV_VERSION}"
|
||||
}
|
||||
|
||||
install_uv() {
|
||||
log_section "Installing uv ${UV_VERSION}"
|
||||
curl -L --proto '=https' --tlsv1.2 -sSf \
|
||||
@@ -176,14 +189,41 @@ run_tests() {
|
||||
setup_pyo3_python
|
||||
install_cargo_binstall
|
||||
install_cargo_nextest
|
||||
install_cargo_llvm_cov
|
||||
|
||||
log_section "Running cargo nextest"
|
||||
cargo nextest run \
|
||||
log_section "Running cargo nextest with Rust coverage"
|
||||
mkdir -p artifacts
|
||||
export LLVM_PROFILE_FILE_NAME="vllm-rust-unit-%4m.profraw"
|
||||
cargo llvm-cov clean \
|
||||
--manifest-path rust/Cargo.toml \
|
||||
--profraw-only
|
||||
|
||||
set +e
|
||||
cargo llvm-cov nextest \
|
||||
--manifest-path rust/Cargo.toml \
|
||||
--workspace \
|
||||
--all-features \
|
||||
--locked \
|
||||
--no-fail-fast
|
||||
--no-fail-fast \
|
||||
--no-clean \
|
||||
--lcov \
|
||||
--output-path artifacts/rust-unit.lcov \
|
||||
--ignore-filename-regex='/\.cargo/(registry|git)/|/rustc/|/target/'
|
||||
local coverage_rc=$?
|
||||
|
||||
local upload_rc=0
|
||||
if [[ $coverage_rc -eq 0 ]]; then
|
||||
# shellcheck source=.buildkite/scripts/rust-coverage.sh
|
||||
source .buildkite/scripts/rust-coverage.sh
|
||||
rust_coverage_upload artifacts/rust-unit.lcov rust-unit
|
||||
upload_rc=$?
|
||||
fi
|
||||
set -e
|
||||
|
||||
if [[ $coverage_rc -ne 0 ]]; then
|
||||
return "$coverage_rc"
|
||||
fi
|
||||
return "$upload_rc"
|
||||
}
|
||||
|
||||
install_protoc
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
#!/bin/sh
|
||||
|
||||
RUST_CODECOV_VERSION="v11.3.1"
|
||||
RUST_CODECOV_SHA256="ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d"
|
||||
|
||||
rust_coverage_repo_root() {
|
||||
if [ -f /vllm-workspace/.buildkite/scripts/rust-coverage.sh ]; then
|
||||
printf '%s\n' /vllm-workspace
|
||||
elif [ -n "${BUILDKITE_BUILD_CHECKOUT_PATH:-}" ] \
|
||||
&& [ -d "$BUILDKITE_BUILD_CHECKOUT_PATH" ]; then
|
||||
printf '%s\n' "$BUILDKITE_BUILD_CHECKOUT_PATH"
|
||||
else
|
||||
git rev-parse --show-toplevel
|
||||
fi
|
||||
}
|
||||
|
||||
rust_coverage_start() {
|
||||
RUST_COVERAGE_FLAG=${1:?coverage flag is required}
|
||||
RUST_COVERAGE_DIR="/tmp/vllm-rust-coverage/${BUILDKITE_JOB_ID:-local}"
|
||||
export RUST_COVERAGE_FLAG RUST_COVERAGE_DIR
|
||||
mkdir -p "$RUST_COVERAGE_DIR"
|
||||
LLVM_PROFILE_FILE="$RUST_COVERAGE_DIR/rust-%4m.profraw"
|
||||
export LLVM_PROFILE_FILE
|
||||
trap rust_coverage_finalize 0
|
||||
}
|
||||
|
||||
rust_coverage_objects() {
|
||||
rust_cov_objects_manifest="$(dirname "$(command -v llvm-cov)")/../objects"
|
||||
python3 - "$rust_cov_objects_manifest" <<'PY'
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
for relative in Path(sys.argv[1]).read_text().splitlines():
|
||||
for entry in sys.path:
|
||||
path = Path(entry or ".").resolve() / relative
|
||||
if path.is_file():
|
||||
print(path)
|
||||
break
|
||||
else:
|
||||
raise RuntimeError(f"installed Rust coverage object was not found: {relative}")
|
||||
PY
|
||||
}
|
||||
|
||||
rust_coverage_collect() {
|
||||
rust_cov_collect_flag=${1:?coverage flag is required}
|
||||
rust_cov_collect_lcov="$RUST_COVERAGE_DIR/$rust_cov_collect_flag.lcov"
|
||||
|
||||
rust_cov_collect_objects=$(rust_coverage_objects) || return 1
|
||||
rust_cov_collect_primary=
|
||||
set --
|
||||
while IFS= read -r rust_cov_collect_object; do
|
||||
if [ -z "$rust_cov_collect_primary" ]; then
|
||||
rust_cov_collect_primary=$rust_cov_collect_object
|
||||
else
|
||||
set -- "$@" "--object=$rust_cov_collect_object"
|
||||
fi
|
||||
done <<EOF
|
||||
$rust_cov_collect_objects
|
||||
EOF
|
||||
|
||||
llvm-profdata merge \
|
||||
-sparse \
|
||||
"$RUST_COVERAGE_DIR"/*.profraw \
|
||||
-o "$RUST_COVERAGE_DIR/merged.profdata" || return 1
|
||||
llvm-cov export \
|
||||
"$rust_cov_collect_primary" \
|
||||
"$@" \
|
||||
--format=lcov \
|
||||
--instr-profile="$RUST_COVERAGE_DIR/merged.profdata" \
|
||||
--ignore-filename-regex='/\.cargo/(registry|git)/|/rustc/|/target/' \
|
||||
> "$rust_cov_collect_lcov" || return 1
|
||||
RUST_COVERAGE_LCOV=$rust_cov_collect_lcov
|
||||
export RUST_COVERAGE_LCOV
|
||||
}
|
||||
|
||||
rust_coverage_upload() {
|
||||
rust_cov_upload_lcov=${1:?LCOV path is required}
|
||||
rust_cov_upload_flag=${2:?coverage flag is required}
|
||||
rust_cov_upload_repo_root=$(rust_coverage_repo_root) || return 1
|
||||
|
||||
if [ "$(uname -m)" != "x86_64" ]; then
|
||||
echo "Rust coverage upload currently supports x86_64 CI agents" >&2
|
||||
return 1
|
||||
fi
|
||||
|
||||
rust_cov_upload_codecov_dir=$(mktemp -d /tmp/codecov-bin.XXXXXX) \
|
||||
|| return 1
|
||||
curl -fsSL \
|
||||
"https://github.com/codecov/codecov-cli/releases/download/${RUST_CODECOV_VERSION}/codecovcli_linux" \
|
||||
-o "$rust_cov_upload_codecov_dir/codecov" || return 1
|
||||
echo "$RUST_CODECOV_SHA256 $rust_cov_upload_codecov_dir/codecov" \
|
||||
| sha256sum -c - || return 1
|
||||
chmod +x "$rust_cov_upload_codecov_dir/codecov" || return 1
|
||||
|
||||
rust_cov_upload_slug="vllm-project/vllm"
|
||||
if [ -n "${BUILDKITE_PULL_REQUEST:-}" ] \
|
||||
&& [ "${BUILDKITE_PULL_REQUEST}" != "false" ] \
|
||||
&& [ -n "${BUILDKITE_PULL_REQUEST_REPO:-}" ]; then
|
||||
rust_cov_upload_slug=$(echo "$BUILDKITE_PULL_REQUEST_REPO" \
|
||||
| sed -E 's#(git@|https?://)([^/:]+)[:/]([^/]+/[^/.]+)(\.git)?$#\3#')
|
||||
case "$rust_cov_upload_slug" in
|
||||
*/*) ;;
|
||||
*) rust_cov_upload_slug="vllm-project/vllm" ;;
|
||||
esac
|
||||
fi
|
||||
|
||||
rust_cov_upload_branch=${BUILDKITE_BRANCH:?BUILDKITE_BRANCH is required}
|
||||
if [ -z "${CODECOV_TOKEN:-}" ]; then
|
||||
# Codecov accepts tokenless public uploads on unprotected branch names.
|
||||
# A colon-separated prefix keeps feature-branch and fork uploads from
|
||||
# requiring a repository secret.
|
||||
if [ -n "${BUILDKITE_PULL_REQUEST:-}" ] \
|
||||
&& [ "${BUILDKITE_PULL_REQUEST}" != "false" ]; then
|
||||
rust_cov_upload_branch="pr${BUILDKITE_PULL_REQUEST}:$rust_cov_upload_branch"
|
||||
else
|
||||
rust_cov_upload_branch="buildkite:$rust_cov_upload_branch"
|
||||
fi
|
||||
fi
|
||||
|
||||
set --
|
||||
set -- "$@" upload-process
|
||||
set -- "$@" --file "$rust_cov_upload_lcov"
|
||||
# LCOV paths are mapped server-side by codecov.yml. Skip the CLI's local
|
||||
# source-line fix scanning, which is unrelated to path mapping.
|
||||
set -- "$@" --disable-search --disable-file-fixes
|
||||
set -- "$@" --fail-on-error --git-service github
|
||||
set -- "$@" --build "${BUILDKITE_BUILD_NUMBER:?BUILDKITE_BUILD_NUMBER is required}"
|
||||
set -- "$@" --branch "$rust_cov_upload_branch"
|
||||
set -- "$@" --sha "${BUILDKITE_COMMIT:?BUILDKITE_COMMIT is required}"
|
||||
set -- "$@" --slug "$rust_cov_upload_slug"
|
||||
set -- "$@" --flag "$rust_cov_upload_flag"
|
||||
set -- "$@" --name "${rust_cov_upload_flag}-${BUILDKITE_JOB_ID:?BUILDKITE_JOB_ID is required}"
|
||||
set -- "$@" --dir "$rust_cov_upload_repo_root"
|
||||
set -- "$@" --network-root-folder "$rust_cov_upload_repo_root"
|
||||
if [ -n "${BUILDKITE_PULL_REQUEST:-}" ] \
|
||||
&& [ "${BUILDKITE_PULL_REQUEST}" != "false" ]; then
|
||||
set -- "$@" --pr "$BUILDKITE_PULL_REQUEST"
|
||||
fi
|
||||
|
||||
rust_cov_upload_log="$rust_cov_upload_codecov_dir/codecov.log"
|
||||
# E2E steps run from tests/, so execute from the repository root to resolve
|
||||
# codecov.yml and repository paths consistently.
|
||||
(
|
||||
cd "$rust_cov_upload_repo_root" || exit 1
|
||||
"$rust_cov_upload_codecov_dir/codecov" "$@"
|
||||
) >"$rust_cov_upload_log" 2>&1
|
||||
rust_cov_upload_rc=$?
|
||||
cat "$rust_cov_upload_log"
|
||||
# v11.3.1 can log API failures while returning zero even with
|
||||
# --fail-on-error. Preserve the strict CI contract explicitly.
|
||||
if grep -aEq 'error.* -- ' "$rust_cov_upload_log"; then
|
||||
echo "Codecov CLI reported an upload error" >&2
|
||||
rust_cov_upload_rc=1
|
||||
fi
|
||||
rm -rf "$rust_cov_upload_codecov_dir"
|
||||
return "$rust_cov_upload_rc"
|
||||
}
|
||||
|
||||
rust_coverage_finalize() {
|
||||
rust_cov_finalize_test_rc=$?
|
||||
trap - 0
|
||||
set +e
|
||||
|
||||
rust_coverage_collect "$RUST_COVERAGE_FLAG"
|
||||
rust_cov_finalize_collect_rc=$?
|
||||
|
||||
rust_cov_finalize_upload_rc=0
|
||||
if [ "$rust_cov_finalize_collect_rc" -eq 0 ]; then
|
||||
rust_coverage_upload "$RUST_COVERAGE_LCOV" "$RUST_COVERAGE_FLAG"
|
||||
rust_cov_finalize_upload_rc=$?
|
||||
fi
|
||||
|
||||
find "$RUST_COVERAGE_DIR" -type f -name '*.profraw' -delete
|
||||
|
||||
if [ "$rust_cov_finalize_test_rc" -ne 0 ]; then
|
||||
exit "$rust_cov_finalize_test_rc"
|
||||
fi
|
||||
if [ "$rust_cov_finalize_collect_rc" -ne 0 ]; then
|
||||
exit "$rust_cov_finalize_collect_rc"
|
||||
fi
|
||||
exit "$rust_cov_finalize_upload_rc"
|
||||
}
|
||||
@@ -113,8 +113,8 @@ $PYTHON .buildkite/scripts/generate-nightly-index.py \
|
||||
echo "Uploading indices to $S3_COMMIT_PREFIX"
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "$S3_COMMIT_PREFIX"
|
||||
|
||||
# Update rocm/nightly/ if on main branch and not a PR
|
||||
if [[ "$BUILDKITE_BRANCH" == "main" && "$BUILDKITE_PULL_REQUEST" == "false" ]] || [[ "$NIGHTLY" == "1" ]]; then
|
||||
# Only scheduled nightly builds should update the moving nightly index.
|
||||
if [[ "${NIGHTLY:-0}" == "1" ]]; then
|
||||
echo "Updating rocm/nightly/ index..."
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "s3://$BUCKET/rocm/nightly/"
|
||||
fi
|
||||
@@ -147,7 +147,7 @@ echo ""
|
||||
echo "Install command (by commit):"
|
||||
echo " pip install vllm --extra-index-url https://${BUCKET}.s3.amazonaws.com/$ROCM_SUBPATH/"
|
||||
echo ""
|
||||
if [[ "$BUILDKITE_BRANCH" == "main" ]] || [[ "$NIGHTLY" == "1" ]]; then
|
||||
if [[ "${NIGHTLY:-0}" == "1" ]]; then
|
||||
echo "Install command (nightly):"
|
||||
echo " pip install vllm --extra-index-url https://${BUCKET}.s3.amazonaws.com/rocm/nightly/"
|
||||
fi
|
||||
|
||||
@@ -18,6 +18,7 @@ steps:
|
||||
- pytest -v -s cuda/test_platform_no_cuda_init.py
|
||||
|
||||
- label: Cudagraph
|
||||
device: h200_35gb
|
||||
key: cudagraph
|
||||
timeout_in_minutes: 30
|
||||
source_file_dependencies:
|
||||
@@ -28,4 +29,4 @@ steps:
|
||||
commands:
|
||||
- pytest -v -s v1/cudagraph/test_cudagraph_dispatch.py
|
||||
- pytest -v -s v1/cudagraph/test_cudagraph_mode.py
|
||||
- pytest -v -s v1/cudagraph/test_breakable_cudagraph.py
|
||||
- pytest -v -s v1/cudagraph/test_breakable_cudagraph.py
|
||||
|
||||
@@ -3,6 +3,7 @@ depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Entrypoints Unit Tests
|
||||
device: h200_35gb
|
||||
key: entrypoints-unit-tests
|
||||
timeout_in_minutes: 25
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -15,6 +16,7 @@ steps:
|
||||
- pytest -v -s entrypoints/weight_transfer
|
||||
|
||||
- label: Entrypoints Integration (LLM)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-llm
|
||||
timeout_in_minutes: 60
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -55,6 +57,7 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (API Server OpenAI - Part 1)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-api-server-openai-part-1
|
||||
timeout_in_minutes: 45
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -73,6 +76,7 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (API Server OpenAI - Part 2)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-api-server-openai-part-2
|
||||
timeout_in_minutes: 45
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -92,6 +96,7 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (API Server Generate)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-api-server-generate
|
||||
timeout_in_minutes: 50
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -114,6 +119,7 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (Responses API)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-responses-api
|
||||
timeout_in_minutes: 50
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -148,6 +154,7 @@ steps:
|
||||
- pytest -v -s entrypoints/multimodal
|
||||
|
||||
- label: Entrypoints Integration (Pooling)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-pooling
|
||||
timeout_in_minutes: 50
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
|
||||
@@ -15,6 +15,7 @@ steps:
|
||||
- pytest -v -s tests/kernels/ir
|
||||
|
||||
- label: Kernels Core Operation Test
|
||||
device: h200_35gb
|
||||
key: kernels-core-operation-test
|
||||
timeout_in_minutes: 120
|
||||
source_file_dependencies:
|
||||
@@ -163,6 +164,7 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: Kernels Mamba Test
|
||||
device: h200_35gb
|
||||
key: kernels-mamba-test
|
||||
timeout_in_minutes: 40
|
||||
source_file_dependencies:
|
||||
@@ -235,6 +237,11 @@ steps:
|
||||
- vllm/model_executor/kernels/linear/cute_dsl/ll_bf16.py
|
||||
- vllm/model_executor/kernels/linear/cute_dsl/_ll_bf16_dotprod.py
|
||||
- vllm/model_executor/kernels/linear/cute_dsl/_ll_bf16_splitk.py
|
||||
- vllm/cute_utils/
|
||||
- vllm/model_executor/layers/mamba/ops/gdn_chunk_cutedsl/
|
||||
- vllm/model_executor/layers/fused_moe/router/bf16x3_router_gemm_cutedsl.py
|
||||
- tests/kernels/mamba/test_gdn_prefill_cutedsl.py
|
||||
- tests/kernels/test_bf16x3_router_gemm_cutedsl.py
|
||||
- tests/kernels/test_ll_bf16_gemm.py
|
||||
- tests/kernels/test_top_k_per_row.py
|
||||
commands:
|
||||
@@ -264,6 +271,8 @@ steps:
|
||||
- pytest -v -s tests/kernels/moe/test_flashinfer_moe.py
|
||||
- pytest -v -s tests/kernels/moe/test_trtllm_nvfp4_moe.py
|
||||
- pytest -v -s tests/kernels/moe/test_cutedsl_moe.py
|
||||
- pytest -v -s tests/kernels/mamba/test_gdn_prefill_cutedsl.py
|
||||
- pytest -v -s tests/kernels/test_bf16x3_router_gemm_cutedsl.py
|
||||
- pytest -v -s tests/kernels/test_ll_bf16_gemm.py
|
||||
# e2e
|
||||
- pytest -v -s tests/models/quantization/test_nvfp4.py
|
||||
|
||||
@@ -78,6 +78,28 @@ steps:
|
||||
commands:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-small-tp.txt
|
||||
|
||||
- label: LM Eval PCP (4xB200)
|
||||
key: lm-eval-pcp-4xb200
|
||||
timeout_in_minutes: 360
|
||||
device: b200-k8s
|
||||
num_devices: 4
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
- tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml
|
||||
- tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml
|
||||
- tests/evals/gsm8k/configs/models-pcp.txt
|
||||
- vllm/model_executor/layers/quantization
|
||||
- vllm/config/parallel.py
|
||||
- vllm/distributed/parallel_state.py
|
||||
- vllm/model_executor/layers/attention/mla_attention.py
|
||||
- vllm/model_executor/layers/attention/pcp.py
|
||||
- vllm/v1/worker/gpu/model_runner.py
|
||||
- vllm/v1/worker/gpu/pcp_manager.py
|
||||
autorun_on_main: true
|
||||
commands:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-pcp.txt
|
||||
|
||||
- label: LM Eval Large Models EP (2xB200)
|
||||
key: lm-eval-large-models-ep-2xb200
|
||||
timeout_in_minutes: 60
|
||||
|
||||
@@ -64,8 +64,9 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: V1 Core + KV + Metrics
|
||||
device: h200_35gb
|
||||
key: v1-core-kv-metrics
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 80
|
||||
source_file_dependencies:
|
||||
- vllm/config/
|
||||
- vllm/distributed/
|
||||
|
||||
@@ -3,13 +3,16 @@ depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Model Executor
|
||||
device: h200_35gb
|
||||
key: model-executor
|
||||
timeout_in_minutes: 45
|
||||
source_file_dependencies:
|
||||
- vllm/engine/arg_utils.py
|
||||
- vllm/config/model.py
|
||||
- vllm/model_executor
|
||||
- vllm/model_executor/warmup
|
||||
- tests/model_executor
|
||||
- tests/model_executor/test_jit_warmup.py
|
||||
- tests/entrypoints/openai/completion/test_tensorizer_entrypoint.py
|
||||
commands:
|
||||
- apt-get update && apt-get install -y curl libsodium23
|
||||
@@ -33,7 +36,9 @@ steps:
|
||||
- vllm/engine/arg_utils.py
|
||||
- vllm/config/model.py
|
||||
- vllm/model_executor
|
||||
- vllm/model_executor/warmup
|
||||
- tests/model_executor
|
||||
- tests/model_executor/test_jit_warmup.py
|
||||
- tests/entrypoints/openai/completion/test_tensorizer_entrypoint.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
|
||||
@@ -21,6 +21,7 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: Language Models Tests (Extra Standard) %N
|
||||
device: h200_35gb
|
||||
key: language-models-tests-extra-standard
|
||||
timeout_in_minutes: 40
|
||||
source_file_dependencies:
|
||||
@@ -51,8 +52,8 @@ steps:
|
||||
- tests/models/language/pooling/test_classification.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
|
||||
- label: Language Models Tests (Hybrid) %N
|
||||
device: h200_35gb
|
||||
key: language-models-tests-hybrid
|
||||
timeout_in_minutes: 65
|
||||
source_file_dependencies:
|
||||
@@ -63,8 +64,8 @@ steps:
|
||||
# Note: also needed to run plamo2 model in vLLM
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/state-spaces/mamba@v2.3.0'
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
# Shard hybrid language model tests
|
||||
- pytest -v -s models/language/generation -m hybrid_model --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB
|
||||
# Shard the hybrid language model tests that are numerically stable on Hopper.
|
||||
- pytest -v -s models/language/generation -m hybrid_model -k 'not granite-4.0-tiny-preview' --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB
|
||||
parallelism: 2
|
||||
mirror:
|
||||
amd:
|
||||
@@ -77,6 +78,20 @@ steps:
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
- pytest -v -s models/language/generation -m hybrid_model --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB
|
||||
|
||||
# Granite 4 hybrid generation is sensitive to hardware-specific Triton SSD
|
||||
# autotuning (https://github.com/vllm-project/vllm/issues/25194). Keep this one
|
||||
# correctness test on L4 until its H200 output matches the Transformers reference.
|
||||
- label: Language Models Tests (Granite L4 Compatibility)
|
||||
key: language-models-tests-granite-l4-compatibility
|
||||
timeout_in_minutes: 65
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/models/language/generation
|
||||
commands:
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/state-spaces/mamba@v2.3.0'
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
- pytest -v -s models/language/generation -m hybrid_model -k 'granite-4.0-tiny-preview'
|
||||
|
||||
- label: Language Models Test (Extended Generation) # 80min
|
||||
device: h200_35gb
|
||||
key: language-models-test-extended-generation
|
||||
|
||||
@@ -119,6 +119,7 @@ steps:
|
||||
- vllm/model_executor/model_loader/
|
||||
|
||||
- label: Multi-Modal Models (Extended Generation 1)
|
||||
device: h200_35gb
|
||||
key: multi-modal-models-extended-generation-1
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -116,8 +116,9 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: PyTorch Fullgraph Smoke Test
|
||||
device: h200_35gb
|
||||
key: pytorch-fullgraph-smoke-test
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 90
|
||||
source_file_dependencies:
|
||||
- vllm/__init__.py
|
||||
- vllm/_aiter_ops.py
|
||||
@@ -149,7 +150,42 @@ steps:
|
||||
# as it is a heavy test that is covered in other steps.
|
||||
# Use `find` to launch multiple instances of pytest so that
|
||||
# they do not suffer from https://github.com/vllm-project/vllm/issues/28965
|
||||
- "find compile/fullgraph/ -name 'test_*.py' -not -name 'test_full_graph.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'"
|
||||
- "find compile/fullgraph/ -name 'test_*.py' -not -name 'test_full_cudagraph.py' -not -name 'test_full_graph.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'"
|
||||
|
||||
# Hopper-only DeepSeek-V2-Lite cases in this file require two 29.3-GiB model
|
||||
# instances and cannot fit a 35GB MIG slice. L4 retains the original coverage:
|
||||
# those SM90 cases skip while the architecture-compatible cases still run.
|
||||
- label: PyTorch Fullgraph CUDAGraph (L4 Compatibility)
|
||||
key: pytorch-fullgraph-cudagraph-l4-compatibility
|
||||
timeout_in_minutes: 60
|
||||
source_file_dependencies:
|
||||
- vllm/__init__.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/_custom_ops.py
|
||||
- vllm/compilation/
|
||||
- vllm/config/
|
||||
- vllm/distributed/
|
||||
- vllm/engine/
|
||||
- vllm/env_override.py
|
||||
- vllm/envs.py
|
||||
- vllm/forward_context.py
|
||||
- vllm/inputs/
|
||||
- vllm/ir/
|
||||
- vllm/kernels/
|
||||
- vllm/logger.py
|
||||
- vllm/model_executor/
|
||||
- vllm/multimodal/
|
||||
- vllm/platforms/
|
||||
- vllm/plugins/
|
||||
- vllm/sampling_params.py
|
||||
- vllm/sequence.py
|
||||
- vllm/transformers_utils/
|
||||
- vllm/triton_utils/
|
||||
- vllm/utils/
|
||||
- vllm/v1/
|
||||
- tests/compile
|
||||
commands:
|
||||
- pytest -s -v compile/fullgraph/test_full_cudagraph.py
|
||||
|
||||
- label: PyTorch Fullgraph
|
||||
key: pytorch-fullgraph
|
||||
|
||||
@@ -3,8 +3,11 @@ depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Quantization
|
||||
device: h200_35gb
|
||||
key: quantization
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 75
|
||||
env:
|
||||
VLLM_USE_V2_MODEL_RUNNER: "0"
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
- vllm/model_executor/layers/quantization
|
||||
@@ -19,9 +22,13 @@ steps:
|
||||
# TODO(jerryzh168): resolve the above comment
|
||||
- uv pip install --system torchao==0.17.0 --index-url https://download.pytorch.org/whl/cu130
|
||||
- uv pip install --system conch-triton-kernels
|
||||
- VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py
|
||||
# The SM90-only checkpoint currently contains a removed weight_chan_scale
|
||||
# parameter. It was not exercised by the previous L4 job.
|
||||
- VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py -k 'not test_compressed_tensors_w4a8_fp8' --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
|
||||
parallelism: 8
|
||||
|
||||
- label: Quantized Fusions
|
||||
device: h200_35gb
|
||||
key: quantized-fusions
|
||||
timeout_in_minutes: 20
|
||||
source_file_dependencies:
|
||||
@@ -52,10 +59,14 @@ steps:
|
||||
- pytest -s -v tests/quantization/test_blackwell_moe.py
|
||||
|
||||
- label: Quantized Models Test
|
||||
device: h200_35gb
|
||||
key: quantized-models-test
|
||||
timeout_in_minutes: 50
|
||||
timeout_in_minutes: 65
|
||||
env:
|
||||
VLLM_USE_V2_MODEL_RUNNER: "0"
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/layers/quantization
|
||||
- tests/models/quantization
|
||||
commands:
|
||||
- pytest -v -s models/quantization
|
||||
- pytest -v -s models/quantization --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
|
||||
parallelism: 3
|
||||
|
||||
@@ -8,6 +8,11 @@ steps:
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- rust/
|
||||
- build_rust.sh
|
||||
- tools/build_rust.py
|
||||
- rust-toolchain.toml
|
||||
- .buildkite/scripts/rust-coverage.sh
|
||||
- codecov.yml
|
||||
- vllm/benchmarks/
|
||||
- vllm/entrypoints/openai/
|
||||
- vllm/entrypoints/serve/
|
||||
@@ -23,6 +28,7 @@ steps:
|
||||
- tests/entrypoints/openai/test_uds.py
|
||||
- tests/v1/sample/test_logprobs_e2e.py
|
||||
commands:
|
||||
- . /vllm-workspace/.buildkite/scripts/rust-coverage.sh && rust_coverage_start rust-e2e
|
||||
- export VLLM_USE_RUST_FRONTEND=1
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -v -s benchmarks/test_serve_cli.py -k "not insecure and not (test_bench_serve and not test_bench_serve_chat)"
|
||||
@@ -43,6 +49,11 @@ steps:
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- rust/
|
||||
- build_rust.sh
|
||||
- tools/build_rust.py
|
||||
- rust-toolchain.toml
|
||||
- .buildkite/scripts/rust-coverage.sh
|
||||
- codecov.yml
|
||||
- vllm/entrypoints/openai/
|
||||
- vllm/entrypoints/serve/
|
||||
- vllm/v1/engine/
|
||||
@@ -54,6 +65,7 @@ steps:
|
||||
# - tests/entrypoints/serve/dev/test_sleep.py
|
||||
- tests/entrypoints/serve/tokenize/test_tokenization.py
|
||||
commands:
|
||||
- . /vllm-workspace/.buildkite/scripts/rust-coverage.sh && rust_coverage_start rust-e2e
|
||||
- export VLLM_USE_RUST_FRONTEND=1
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- PYTHONPATH=/vllm-workspace pytest -v -s entrypoints/serve/dev/rpc/test_collective_rpc.py
|
||||
@@ -72,24 +84,37 @@ steps:
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- rust/
|
||||
- build_rust.sh
|
||||
- tools/build_rust.py
|
||||
- rust-toolchain.toml
|
||||
- .buildkite/scripts/rust-coverage.sh
|
||||
- codecov.yml
|
||||
- vllm/entrypoints/openai/
|
||||
- tests/utils.py
|
||||
- tests/entrypoints/openai/correctness/test_lmeval.py
|
||||
commands:
|
||||
- . /vllm-workspace/.buildkite/scripts/rust-coverage.sh && rust_coverage_start rust-e2e
|
||||
- export VLLM_USE_RUST_FRONTEND=1
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine
|
||||
|
||||
- label: Rust Frontend Tool Use
|
||||
device: h200_35gb
|
||||
timeout_in_minutes: 25
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- rust/
|
||||
- build_rust.sh
|
||||
- tools/build_rust.py
|
||||
- rust-toolchain.toml
|
||||
- .buildkite/scripts/rust-coverage.sh
|
||||
- codecov.yml
|
||||
- vllm/entrypoints/openai/
|
||||
- vllm/tool_parsers/
|
||||
- tests/utils.py
|
||||
- tests/tool_use/
|
||||
commands:
|
||||
- . /vllm-workspace/.buildkite/scripts/rust-coverage.sh && rust_coverage_start rust-e2e
|
||||
- export VLLM_USE_RUST_FRONTEND=1
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -v -s tool_use --ignore=tool_use/mistral --models llama3.2 -k "not test_response_format_with_tool_choice_required and not test_parallel_tool_calls_false and not test_tool_call_and_choice"
|
||||
@@ -100,6 +125,11 @@ steps:
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- rust/
|
||||
- build_rust.sh
|
||||
- tools/build_rust.py
|
||||
- rust-toolchain.toml
|
||||
- .buildkite/scripts/rust-coverage.sh
|
||||
- codecov.yml
|
||||
- vllm/distributed/
|
||||
- vllm/engine/
|
||||
- vllm/executor/
|
||||
@@ -110,6 +140,7 @@ steps:
|
||||
- tests/v1/distributed/test_hybrid_lb_dp.py
|
||||
- tests/v1/distributed/test_internal_lb_dp.py
|
||||
commands:
|
||||
- . /vllm-workspace/.buildkite/scripts/rust-coverage.sh && rust_coverage_start rust-e2e
|
||||
- export VLLM_USE_RUST_FRONTEND=1
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- export NCCL_CUMEM_HOST_ENABLE=0
|
||||
|
||||
@@ -26,5 +26,7 @@ steps:
|
||||
- rust-toolchain.toml
|
||||
- .buildkite/test_areas/rust_frontend_cargo.yaml
|
||||
- .buildkite/scripts/run-rust-frontend-cargo-ci.sh
|
||||
- .buildkite/scripts/rust-coverage.sh
|
||||
- codecov.yml
|
||||
commands:
|
||||
- .buildkite/scripts/run-rust-frontend-cargo-ci.sh test
|
||||
|
||||
@@ -47,6 +47,7 @@
|
||||
|
||||
# Rust Frontend
|
||||
/rust/ @BugenZhao @njhill
|
||||
/rust/src/bench @esmeetu
|
||||
/build_rust.sh @BugenZhao @njhill
|
||||
/rust-toolchain.toml @BugenZhao @njhill
|
||||
/.buildkite/test_areas/rust* @BugenZhao @njhill
|
||||
|
||||
@@ -257,3 +257,4 @@ vllm/grpc/vllm_engine_pb2.pyi
|
||||
|
||||
# Ignore generated cpu headers
|
||||
csrc/cpu/cpu_attn_dispatch_generated.h
|
||||
rust-coverage-tools/
|
||||
|
||||
@@ -30,7 +30,7 @@ repos:
|
||||
- id: markdownlint-cli2
|
||||
language_version: lts
|
||||
args: [--fix]
|
||||
exclude: ^CLAUDE\.md$
|
||||
exclude: (^|/)CLAUDE\.md$
|
||||
- repo: https://github.com/rhysd/actionlint
|
||||
rev: v1.7.7
|
||||
hooks:
|
||||
|
||||
@@ -48,7 +48,7 @@ vLLM is flexible and easy to use with:
|
||||
- Tool calling and reasoning parsers
|
||||
- OpenAI-compatible API server, plus Anthropic Messages API and gRPC support
|
||||
- Efficient multi-LoRA support for dense and MoE layers
|
||||
- Support for NVIDIA GPUs, AMD GPUs, and x86/ARM/PowerPC CPUs. Additionally, diverse hardware plugins such as Google TPUs, Intel Gaudi, IBM Spyre, Huawei Ascend, Rebellions NPU, Apple Silicon, MetaX GPU, and more.
|
||||
- Support for NVIDIA GPUs, AMD GPUs, Intel GPUs, and x86/ARM/PowerPC CPUs. Additionally, diverse hardware plugins such as Google TPUs, Intel Gaudi, IBM Spyre, Huawei Ascend, Rebellions NPU, Apple Silicon, MetaX GPU, and more.
|
||||
|
||||
vLLM seamlessly supports 200+ model architectures on Hugging Face, including:
|
||||
|
||||
|
||||
@@ -69,12 +69,11 @@ def make_inputs(total_tokens, num_reqs, block_size):
|
||||
# Output workspace
|
||||
dst = torch.zeros(total_tokens, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
|
||||
|
||||
seq_lens_t = torch.tensor(seq_lens, dtype=torch.int32, device="cuda")
|
||||
workspace_starts_t = torch.tensor(
|
||||
workspace_starts, dtype=torch.int32, device="cuda"
|
||||
)
|
||||
|
||||
return cache, dst, block_table, seq_lens_t, workspace_starts_t
|
||||
return cache, dst, block_table, workspace_starts_t
|
||||
|
||||
|
||||
def bench_scenario(label, num_reqs, total_tokens_list, save_path):
|
||||
@@ -94,7 +93,7 @@ def bench_scenario(label, num_reqs, total_tokens_list, save_path):
|
||||
)
|
||||
)
|
||||
def bench_fn(total_tokens, provider, num_reqs):
|
||||
cache, dst, block_table, seq_lens_t, ws_starts = make_inputs(
|
||||
cache, dst, block_table, ws_starts = make_inputs(
|
||||
total_tokens, num_reqs, BLOCK_SIZE
|
||||
)
|
||||
|
||||
@@ -102,7 +101,7 @@ def bench_scenario(label, num_reqs, total_tokens_list, save_path):
|
||||
|
||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
||||
lambda: ops.cp_gather_and_upconvert_fp8_kv_cache(
|
||||
cache, dst, block_table, seq_lens_t, ws_starts, num_reqs
|
||||
cache, dst, block_table, ws_starts, num_reqs
|
||||
),
|
||||
quantiles=quantiles,
|
||||
rep=500,
|
||||
|
||||
@@ -8,6 +8,8 @@
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ROOT="$(cd "$(dirname "$0")" && pwd)"
|
||||
CARGO_LLVM_COV_VERSION="0.8.7"
|
||||
COVERAGE_TOOLS_DIR="$REPO_ROOT/rust-coverage-tools"
|
||||
|
||||
# Read the required toolchain from rust-toolchain.toml.
|
||||
TOOLCHAIN=$(grep '^channel' "$REPO_ROOT/rust-toolchain.toml" | sed 's/.*= *"\(.*\)"/\1/')
|
||||
@@ -30,4 +32,39 @@ else
|
||||
PROFILE_ARG="--release"
|
||||
fi
|
||||
|
||||
rm -rf "$COVERAGE_TOOLS_DIR"
|
||||
mkdir -p "$COVERAGE_TOOLS_DIR/bin" "$COVERAGE_TOOLS_DIR/lib"
|
||||
|
||||
if [[ "${VLLM_RUST_COVERAGE:-0}" == "1" ]]; then
|
||||
# rustc wrapper flags are invisible to Cargo's normal fingerprinting.
|
||||
# Keep instrumented intermediates isolated when local builds switch modes.
|
||||
export CARGO_TARGET_DIR="$REPO_ROOT/rust/target/coverage"
|
||||
rustup component add --toolchain "$TOOLCHAIN" llvm-tools-preview
|
||||
cargo +"$TOOLCHAIN" install \
|
||||
--locked \
|
||||
--version "$CARGO_LLVM_COV_VERSION" \
|
||||
cargo-llvm-cov
|
||||
|
||||
eval "$(
|
||||
cargo +"$TOOLCHAIN" llvm-cov show-env \
|
||||
--manifest-path "$REPO_ROOT/rust/Cargo.toml" \
|
||||
--sh
|
||||
)"
|
||||
|
||||
# Build scripts and proc macros can run during compilation. Their profiles
|
||||
# are unrelated to runtime coverage and would otherwise pollute the tree.
|
||||
export LLVM_PROFILE_FILE=/dev/null
|
||||
export VLLM_RUST_COVERAGE_OBJECTS="$COVERAGE_TOOLS_DIR/objects"
|
||||
fi
|
||||
|
||||
python3 "$REPO_ROOT/tools/build_rust.py" "$PROFILE_ARG"
|
||||
|
||||
if [[ "${VLLM_RUST_COVERAGE:-0}" == "1" ]]; then
|
||||
LLVM_BIN_DIR="$(dirname "$(rustup run "$TOOLCHAIN" rustc \
|
||||
--print target-libdir)")/bin"
|
||||
|
||||
cp "$LLVM_BIN_DIR"/{llvm-cov,llvm-profdata} "$COVERAGE_TOOLS_DIR/bin/"
|
||||
chmod 0755 "$COVERAGE_TOOLS_DIR/bin/"*
|
||||
cp -L "$LLVM_BIN_DIR"/../lib/libLLVM.so* "$COVERAGE_TOOLS_DIR/lib/"
|
||||
chmod 0644 "$COVERAGE_TOOLS_DIR/lib/"*
|
||||
fi
|
||||
|
||||
@@ -430,6 +430,7 @@ set(VLLM_EXT_SRC
|
||||
"csrc/cpu/layernorm.cpp"
|
||||
"csrc/cpu/mla_decode.cpp"
|
||||
"csrc/cpu/pos_encoding.cpp"
|
||||
"csrc/cpu/mamba_cpu.cpp"
|
||||
"csrc/moe/dynamic_4bit_int_moe_cpu.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp")
|
||||
@@ -489,6 +490,7 @@ if (ENABLE_X86_ISA)
|
||||
"csrc/cpu/spec_decode_utils.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/dnnl_kernels.cpp"
|
||||
"csrc/cpu/mamba_cpu.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp"
|
||||
# TODO: Remove these files
|
||||
"csrc/cpu/activation.cpp"
|
||||
@@ -502,6 +504,7 @@ if (ENABLE_X86_ISA)
|
||||
"csrc/cpu/utils.cpp"
|
||||
"csrc/cpu/spec_decode_utils.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/mamba_cpu.cpp"
|
||||
"csrc/cpu/dnnl_kernels.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp"
|
||||
# TODO: Remove these files
|
||||
|
||||
@@ -22,7 +22,7 @@ if(QUTLASS_SRC_DIR)
|
||||
set(qutlass_BINARY_DIR "${CMAKE_BINARY_DIR}/qutlass-binary-dir-unused")
|
||||
else()
|
||||
set(_QUTLASS_UPSTREAM_REPO "https://github.com/IST-DASLab/qutlass.git")
|
||||
set(_QUTLASS_UPSTREAM_TAG "830d2c4537c7396e14a02a46fbddd18b5d107c65")
|
||||
set(_QUTLASS_UPSTREAM_TAG "e74319e3405ce6d71965732880f5dc1f52371f64")
|
||||
|
||||
set(_qutlass_fc_root "${FETCHCONTENT_BASE_DIR}")
|
||||
if(NOT _qutlass_fc_root)
|
||||
@@ -125,8 +125,6 @@ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND QUTLASS_ARCHS)
|
||||
CUDA_ARCHS "${QUTLASS_ARCHS}"
|
||||
)
|
||||
|
||||
# QuTLASS uses legacy ATen headers and cannot be built with TORCH_TARGET_VERSION.
|
||||
# Keep it as its own extension (registers torch.ops._qutlass_C).
|
||||
define_extension_target(
|
||||
_qutlass_C
|
||||
DESTINATION vllm
|
||||
@@ -139,9 +137,11 @@ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND QUTLASS_ARCHS)
|
||||
WITH_SOABI)
|
||||
|
||||
target_compile_definitions(_qutlass_C PRIVATE
|
||||
QUTLASS_DISABLE_PYBIND=1
|
||||
QUTLASS_MINIMAL_BUILD=1
|
||||
TARGET_CUDA_ARCH=${QUTLASS_TARGET_CC}
|
||||
CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1)
|
||||
CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1
|
||||
TORCH_TARGET_VERSION=0x020B000000000000ULL
|
||||
USE_CUDA)
|
||||
|
||||
set_property(SOURCE ${QUTLASS_SOURCES} APPEND PROPERTY COMPILE_OPTIONS
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:--expt-relaxed-constexpr --use_fast_math -O3>
|
||||
|
||||
@@ -14,7 +14,7 @@ else()
|
||||
FetchContent_Declare(
|
||||
tml_fa4
|
||||
GIT_REPOSITORY https://github.com/vllm-project/tml-fa4.git
|
||||
GIT_TAG 13374f0c855acc1add1bf30444bd67aebbc24a8e
|
||||
GIT_TAG b206834606ed5b5f21f8eed6b0683f528ea9cf7d
|
||||
GIT_PROGRESS TRUE
|
||||
CONFIGURE_COMMAND ""
|
||||
BUILD_COMMAND "")
|
||||
|
||||
@@ -39,7 +39,7 @@ else()
|
||||
FetchContent_Declare(
|
||||
vllm-flash-attn
|
||||
GIT_REPOSITORY https://github.com/vllm-project/flash-attention.git
|
||||
GIT_TAG caaa4eb59845388a20b1f435ecaafb4bd9517ad8
|
||||
GIT_TAG 168920233059c48de6199e2cda74003b2ce3d199
|
||||
GIT_PROGRESS TRUE
|
||||
# Don't share the vllm-flash-attn build between build types
|
||||
BINARY_DIR ${CMAKE_BINARY_DIR}/vllm-flash-attn
|
||||
|
||||
+13
@@ -10,3 +10,16 @@ fixes:
|
||||
- "/usr/local/lib/python3.*/site-packages/vllm/::vllm/"
|
||||
- "/usr/lib/python3.*/dist-packages/vllm/::vllm/"
|
||||
- "/usr/lib/python3.*/site-packages/vllm/::vllm/"
|
||||
# Map Rust sources built in the E2E image and on Buildkite agents.
|
||||
- "/workspace/rust/::rust/"
|
||||
- "/var/lib/buildkite-agent/.*/rust/::rust/"
|
||||
|
||||
flags:
|
||||
rust-unit:
|
||||
paths:
|
||||
- rust/
|
||||
carryforward: false
|
||||
rust-e2e:
|
||||
paths:
|
||||
- rust/
|
||||
carryforward: false
|
||||
|
||||
+1
-2
@@ -67,9 +67,8 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
torch::Tensor const& src_cache, // [NUM_BLOCKS, BLOCK_SIZE, 656]
|
||||
torch::Tensor const& dst, // [TOT_TOKENS, 576]
|
||||
torch::Tensor const& block_table, // [BATCH, BLOCK_INDICES]
|
||||
torch::Tensor const& seq_lens, // [BATCH]
|
||||
torch::Tensor const& workspace_starts, // [BATCH]
|
||||
int64_t batch_size);
|
||||
int64_t batch_size, std::optional<torch::Tensor> seq_starts = std::nullopt);
|
||||
|
||||
// Indexer K quantization and cache function
|
||||
void indexer_k_quant_and_cache(
|
||||
|
||||
@@ -102,7 +102,9 @@ class TileGemm82 {
|
||||
kv_cache_t* __restrict__ curr_b = b_tile;
|
||||
|
||||
for (int32_t k = 0; k < dynamic_k_size; ++k) {
|
||||
auto [fp32_b_0_reg, fp32_b_1_reg] = load_b_pair_vec(curr_b);
|
||||
auto fp32_b_regs = load_b_pair_vec(curr_b);
|
||||
auto fp32_b_0_reg = fp32_b_regs.first;
|
||||
auto fp32_b_1_reg = fp32_b_regs.second;
|
||||
|
||||
float* __restrict__ curr_m_a = curr_a;
|
||||
vec_op::unroll_loop<int32_t, M>([&](int32_t i) {
|
||||
|
||||
@@ -336,13 +336,14 @@ struct FP32Vec8 : public Vec<FP32Vec8> {
|
||||
reg.val[1] = fp16_to_fp32_bits(raw_lo);
|
||||
}
|
||||
float reduce_sum() const {
|
||||
AliasReg ar;
|
||||
ar.reg = reg;
|
||||
float result = 0;
|
||||
unroll_loop<int, VEC_ELEM_NUM>(
|
||||
[&result, &ar](int i) { result += ar.values[i]; });
|
||||
|
||||
return result;
|
||||
// VSX horizontal reduction: 3 vector ops instead of 8 scalar adds.
|
||||
// Step 1: pairwise sum of the two 4-wide halves
|
||||
__vector float s = vec_add(reg.val[0], reg.val[1]);
|
||||
// Step 2: rotate by 8 bytes (2 floats) and add
|
||||
s = vec_add(s, vec_sld(s, s, 8));
|
||||
// Step 3: rotate by 4 bytes (1 float) and add => all lanes hold total
|
||||
s = vec_add(s, vec_sld(s, s, 4));
|
||||
return vec_extract(s, 0);
|
||||
}
|
||||
FP32Vec8 exp() const {
|
||||
f32x4x2_t out;
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
//
|
||||
// CPU at::Tensor wrappers for Mamba decode-step kernels defined in
|
||||
// mamba_kernels.hpp.
|
||||
|
||||
#include "cpu/mamba_kernels.hpp"
|
||||
|
||||
#include <ATen/ATen.h>
|
||||
#include <torch/library.h>
|
||||
#include <c10/util/Optional.h>
|
||||
|
||||
#include "cpu_types.hpp"
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// causal_conv1d_update
|
||||
// ---------------------------------------------------------------------------
|
||||
at::Tensor causal_conv1d_update_cpu_impl(
|
||||
at::Tensor& x, at::Tensor& conv_state, const at::Tensor& weight,
|
||||
const c10::optional<at::Tensor>& bias,
|
||||
const c10::optional<std::string>& activation,
|
||||
const c10::optional<at::Tensor>& conv_state_indices,
|
||||
const c10::optional<at::Tensor>& query_start_loc, int64_t pad_slot_id) {
|
||||
bool do_silu = false;
|
||||
if (activation.has_value()) {
|
||||
const std::string& act = activation.value();
|
||||
do_silu = (act == "silu" || act == "swish");
|
||||
}
|
||||
|
||||
at::ScalarType dtype = x.scalar_type();
|
||||
|
||||
// Input x: contiguous in native dtype.
|
||||
at::Tensor x_c = x.is_contiguous() ? x : x.contiguous();
|
||||
|
||||
// conv_state: NEVER copy the full paged tensor just for layout reasons.
|
||||
// If the dtype matches we work directly on conv_state (contiguous or not)
|
||||
// by extracting strides and passing them to the kernel.
|
||||
// Only a dtype-conversion copy is made when types differ (rare for BF16).
|
||||
bool state_type_ok = (conv_state.scalar_type() == dtype);
|
||||
at::Tensor state_c = state_type_ok ? conv_state : conv_state.to(dtype);
|
||||
// state_c and conv_state may be non-contiguous — that is intentional.
|
||||
|
||||
// Weight: coerce to same dtype if needed (should match in practice)
|
||||
at::Tensor w_c =
|
||||
(weight.scalar_type() != dtype)
|
||||
? weight.to(dtype).contiguous()
|
||||
: (weight.is_contiguous() ? weight : weight.contiguous());
|
||||
|
||||
// Bias stays float32 (small scalar, used only for fp32 accumulation)
|
||||
at::Tensor bias_f32;
|
||||
if (bias.has_value() && bias.value().defined())
|
||||
bias_f32 = bias.value().to(at::kFloat).contiguous();
|
||||
|
||||
int64_t batch = x_c.size(0);
|
||||
int64_t dim = x_c.size(1);
|
||||
int64_t seqlen = (x_c.dim() == 3) ? x_c.size(2) : 1;
|
||||
int64_t width = w_c.size(1);
|
||||
int64_t state_len = state_c.size(2);
|
||||
|
||||
// Extract strides — works for contiguous AND non-contiguous (transposed)
|
||||
// state. stride(0): between cache slots (e.g. num_slots × dim × width-1 in
|
||||
// contiguous) stride(1): between conv channels (dim stride) stride(2):
|
||||
// between state elements (=1 when contiguous, =dim when transposed)
|
||||
int64_t stride_s_slot = state_c.stride(0);
|
||||
int64_t stride_s_dim = state_c.stride(1);
|
||||
int64_t stride_s_state = state_c.stride(2);
|
||||
|
||||
at::Tensor out = x_c.clone(); // native dtype, no float32 alloc
|
||||
|
||||
const int32_t* cache_idx_ptr = nullptr;
|
||||
at::Tensor cache_idx_int;
|
||||
if (conv_state_indices.has_value()) {
|
||||
cache_idx_int = conv_state_indices.value().to(at::kInt).contiguous();
|
||||
cache_idx_ptr = cache_idx_int.data_ptr<int32_t>();
|
||||
}
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(dtype, "causal_conv1d_update", [&] {
|
||||
mamba_cpu::causal_conv1d_update_kernel<scalar_t>(
|
||||
x_c.data_ptr<scalar_t>(), state_c.data_ptr<scalar_t>(), stride_s_slot,
|
||||
stride_s_dim, stride_s_state, w_c.data_ptr<scalar_t>(),
|
||||
bias_f32.defined() ? bias_f32.data_ptr<float>() : nullptr,
|
||||
out.data_ptr<scalar_t>(), cache_idx_ptr,
|
||||
static_cast<int32_t>(pad_slot_id), batch, dim, seqlen, width, state_len,
|
||||
do_silu);
|
||||
});
|
||||
|
||||
// Write back only when a type-conversion copy was made.
|
||||
// Layout-only non-contiguity is handled via strides above — no copy needed.
|
||||
if (!state_type_ok) conv_state.copy_(state_c);
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// selective_state_update
|
||||
// ---------------------------------------------------------------------------
|
||||
void selective_state_update_cpu_impl(
|
||||
at::Tensor& state, // (nstates, nheads, dim, dstate)
|
||||
const at::Tensor& x, // (N, nheads, dim)
|
||||
const at::Tensor& dt, const at::Tensor& A, const at::Tensor& B,
|
||||
const at::Tensor& C, const c10::optional<at::Tensor>& D,
|
||||
const c10::optional<at::Tensor>& z,
|
||||
const c10::optional<at::Tensor>& dt_bias, bool dt_softplus,
|
||||
const c10::optional<at::Tensor>& state_batch_indices,
|
||||
const c10::optional<at::Tensor>& dst_state_batch_indices,
|
||||
int64_t null_block_id, at::Tensor& out,
|
||||
const c10::optional<at::Tensor>& num_accepted_tokens,
|
||||
const c10::optional<at::Tensor>& cu_seqlens) {
|
||||
at::ScalarType state_type = state.scalar_type();
|
||||
at::ScalarType input_type = x.scalar_type();
|
||||
|
||||
// x, B, C must be contiguous and match input_type
|
||||
auto ensure_input = [input_type](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = (t.scalar_type() != input_type) ? t.to(input_type) : t;
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
at::Tensor x_in = ensure_input(x);
|
||||
at::Tensor B_in = ensure_input(B);
|
||||
at::Tensor C_in = ensure_input(C);
|
||||
at::Tensor z_in;
|
||||
if (z.has_value() && z.value().defined()) z_in = ensure_input(z.value());
|
||||
|
||||
// A, D, dt_bias are float32 model parameters that arrive here as expanded
|
||||
// tensors, e.g. A is (nheads, head_dim, dstate) with strides (1, 0, 0).
|
||||
// We need just the scalar value per head as a (nheads,) 1-D array so that
|
||||
// A_ptr[h] in the kernel correctly reads head h's value.
|
||||
//
|
||||
// Strategy: peel trailing expanded (stride=0) dims via .select(), which is
|
||||
// a zero-copy view. For A: (nheads, head_dim, dstate) strides (1,0,0)
|
||||
// → .select(2,0) → (nheads, head_dim) strides (1,0)
|
||||
// → .select(1,0) → (nheads,) stride (1,) ← contiguous, free.
|
||||
// No allocation, no type conversion (A is already float32).
|
||||
auto to_per_head_1d_f32 = [](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = t;
|
||||
// Peel trailing dimensions that are broadcast (stride=0 or size=1)
|
||||
while (r.dim() > 1) r = r.select(r.dim() - 1, 0);
|
||||
if (r.scalar_type() != at::kFloat) r = r.to(at::kFloat);
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
|
||||
at::Tensor A_f32 = to_per_head_1d_f32(A); // (nheads,) float32
|
||||
at::Tensor D_f32, dt_bias_f32;
|
||||
if (D.has_value() && D.value().defined())
|
||||
D_f32 = to_per_head_1d_f32(D.value());
|
||||
if (dt_bias.has_value() && dt_bias.value().defined())
|
||||
dt_bias_f32 = to_per_head_1d_f32(dt_bias.value());
|
||||
|
||||
// dt: reduce (N, nheads, head_dim) expanded tensor → (N, nheads) BEFORE
|
||||
// the type conversion so we convert head_dim x fewer elements.
|
||||
at::Tensor dt_f32;
|
||||
{
|
||||
// If dt was expanded to (N, nheads, head_dim) with stride-0 in dim 2,
|
||||
// take a zero-copy view of index 0 along that dim first.
|
||||
at::Tensor t2 = (dt.dim() == 3) ? dt.select(2, 0) : dt; // (N, nheads)
|
||||
at::Tensor t3 = (t2.scalar_type() != at::kFloat) ? t2.to(at::kFloat) : t2;
|
||||
dt_f32 = t3.is_contiguous() ? t3 : t3.contiguous();
|
||||
}
|
||||
|
||||
int64_t nheads = state.size(1);
|
||||
int64_t dim = state.size(2);
|
||||
int64_t dstate = state.size(3);
|
||||
int64_t N = (cu_seqlens.has_value() && cu_seqlens.value().defined())
|
||||
? cu_seqlens.value().size(0) - 1
|
||||
: x_in.size(0);
|
||||
int64_t ngroups = B_in.size(1);
|
||||
|
||||
// Strides
|
||||
int64_t stride_state_n = state.stride(0);
|
||||
int64_t stride_state_h = state.stride(1);
|
||||
int64_t stride_state_d = state.stride(2);
|
||||
int64_t stride_x_n = x_in.stride(0);
|
||||
int64_t stride_x_h = x_in.stride(1);
|
||||
int64_t stride_dt_n = dt_f32.stride(0); // dt is (N, nheads)
|
||||
int64_t stride_BC_n = B_in.stride(0);
|
||||
int64_t stride_BC_g = B_in.stride(1);
|
||||
int64_t stride_out_n = out.stride(0);
|
||||
int64_t stride_out_h = out.stride(1);
|
||||
|
||||
// Optional index pointers
|
||||
auto get_int32_ptr =
|
||||
[](const c10::optional<at::Tensor>& opt) -> const int32_t* {
|
||||
return (opt.has_value() && opt.value().defined())
|
||||
? opt.value().data_ptr<int32_t>()
|
||||
: nullptr;
|
||||
};
|
||||
const int32_t* sbi_ptr = get_int32_ptr(state_batch_indices);
|
||||
const int32_t* dsbi_ptr = get_int32_ptr(dst_state_batch_indices);
|
||||
const int32_t* nat_ptr = get_int32_ptr(num_accepted_tokens);
|
||||
const int32_t* csl_ptr = get_int32_ptr(cu_seqlens);
|
||||
|
||||
// Dispatch on (state_t, input_t, out_t): write directly into `out`
|
||||
// without any intermediate float32 buffer.
|
||||
VLLM_DISPATCH_FLOATING_TYPES(state_type, "ssu_state", [&] {
|
||||
using state_t = scalar_t;
|
||||
VLLM_DISPATCH_FLOATING_TYPES(input_type, "ssu_input", [&] {
|
||||
using input_t = scalar_t;
|
||||
VLLM_DISPATCH_FLOATING_TYPES(out.scalar_type(), "ssu_out", [&] {
|
||||
using out_t = scalar_t;
|
||||
mamba_cpu::selective_state_update_kernel<state_t, input_t, out_t>(
|
||||
state.data_ptr<state_t>(), stride_state_n, stride_state_h,
|
||||
stride_state_d, x_in.data_ptr<input_t>(), stride_x_n, stride_x_h,
|
||||
dt_f32.data_ptr<float>(), stride_dt_n, A_f32.data_ptr<float>(),
|
||||
B_in.data_ptr<input_t>(), C_in.data_ptr<input_t>(), stride_BC_n,
|
||||
stride_BC_g, D_f32.defined() ? D_f32.data_ptr<float>() : nullptr,
|
||||
z_in.defined() ? z_in.data_ptr<input_t>() : nullptr,
|
||||
dt_bias_f32.defined() ? dt_bias_f32.data_ptr<float>() : nullptr,
|
||||
out.data_ptr<out_t>(), stride_out_n, stride_out_h, sbi_ptr,
|
||||
dsbi_ptr, static_cast<int32_t>(null_block_id), nat_ptr, csl_ptr, N,
|
||||
nheads, ngroups, dim, dstate, dt_softplus);
|
||||
});
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// mamba_chunk_scan_fwd_cpu
|
||||
// ---------------------------------------------------------------------------
|
||||
void mamba_chunk_scan_fwd_cpu_impl(
|
||||
at::Tensor& out, // [seqlen, nheads, headdim] — pre-allocated by caller
|
||||
at::Tensor&
|
||||
final_states, // [batch, nheads, headdim, dstate] float32 contiguous
|
||||
const at::Tensor& x, // [seqlen, nheads, headdim]
|
||||
const at::Tensor&
|
||||
dt, // [seqlen, nheads] float32 (preprocessed: bias+softplus+clamp)
|
||||
const at::Tensor& A, // [nheads] float32
|
||||
const at::Tensor& B, // [seqlen, ngroups, dstate]
|
||||
const at::Tensor& C, // [seqlen, ngroups, dstate]
|
||||
const c10::optional<at::Tensor>& D, // [nheads] float32 (optional)
|
||||
const c10::optional<at::Tensor>& z, // [seqlen, nheads, headdim] (optional)
|
||||
const at::Tensor& cu_seqlens // [batch+1] int32
|
||||
) {
|
||||
const at::ScalarType input_type = x.scalar_type();
|
||||
|
||||
auto ensure_contig = [input_type](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = (t.scalar_type() != input_type) ? t.to(input_type) : t;
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
at::Tensor x_in = ensure_contig(x);
|
||||
at::Tensor B_in = ensure_contig(B);
|
||||
at::Tensor C_in = ensure_contig(C);
|
||||
at::Tensor z_in;
|
||||
if (z.has_value() && z.value().defined()) z_in = ensure_contig(z.value());
|
||||
|
||||
// A and D are float32 model parameters, potentially broadcast-expanded.
|
||||
// Strip trailing broadcast dims to get a contiguous (nheads,) array.
|
||||
auto to_per_head_f32 = [](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = t;
|
||||
while (r.dim() > 1) r = r.select(r.dim() - 1, 0);
|
||||
if (r.scalar_type() != at::kFloat) r = r.to(at::kFloat);
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
at::Tensor A_f32 = to_per_head_f32(A);
|
||||
at::Tensor D_f32;
|
||||
if (D.has_value() && D.value().defined()) D_f32 = to_per_head_f32(D.value());
|
||||
|
||||
// dt: [seqlen, nheads] float32 — caller has applied bias+softplus+clamp in
|
||||
// Python.
|
||||
at::Tensor dt_c = dt.is_contiguous() ? dt : dt.contiguous();
|
||||
if (dt_c.scalar_type() != at::kFloat) dt_c = dt_c.to(at::kFloat);
|
||||
|
||||
at::Tensor cu_int = cu_seqlens.to(at::kInt).contiguous();
|
||||
|
||||
const int64_t batch = final_states.size(0);
|
||||
const int64_t nheads = final_states.size(1);
|
||||
const int64_t headdim = final_states.size(2);
|
||||
const int64_t dstate = final_states.size(3);
|
||||
const int64_t ngroups = B_in.size(1);
|
||||
|
||||
TORCH_CHECK(final_states.is_contiguous(),
|
||||
"mamba_chunk_scan_fwd_cpu: final_states must be contiguous");
|
||||
TORCH_CHECK(out.is_contiguous(),
|
||||
"mamba_chunk_scan_fwd_cpu: out must be contiguous (writes via "
|
||||
"raw data_ptr)");
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(input_type, "mamba_chunk_scan_fwd_cpu", [&] {
|
||||
mamba_cpu::mamba_chunk_scan_fwd_kernel<scalar_t>(
|
||||
final_states.data_ptr<float>(), x_in.data_ptr<scalar_t>(),
|
||||
dt_c.data_ptr<float>(), A_f32.data_ptr<float>(),
|
||||
B_in.data_ptr<scalar_t>(), C_in.data_ptr<scalar_t>(),
|
||||
D_f32.defined() ? D_f32.data_ptr<float>() : nullptr,
|
||||
z_in.defined() ? z_in.data_ptr<scalar_t>() : nullptr,
|
||||
out.data_ptr<scalar_t>(), cu_int.data_ptr<int32_t>(), batch, nheads,
|
||||
ngroups, headdim, dstate);
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,382 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
//
|
||||
// Fused CPU vector kernels for Mamba decode-step hotspots:
|
||||
// - causal_conv1d_update (depthwise 1-D conv state roll + compute)
|
||||
// - selective_state_update (SSM recurrence, single-step)
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cpu_types.hpp"
|
||||
#include <cmath>
|
||||
#include <cstring>
|
||||
#include <cstdint>
|
||||
#include <algorithm>
|
||||
|
||||
namespace mamba_cpu {
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// causal_conv1d_update — templated for native BF16/FP32
|
||||
//
|
||||
// state_ptr may point to a NON-CONTIGUOUS paged KV cache tensor.
|
||||
// Explicit strides are passed so the kernel writes directly into the
|
||||
// correct memory locations without making a contiguous copy of the full
|
||||
// paged tensor (which was the source of the 34-41% direct_copy_kernel).
|
||||
//
|
||||
// stride_s_slot = state.stride(0) — between cache slots
|
||||
// stride_s_dim = state.stride(1) — between conv_dim channels
|
||||
// stride_s_state = state.stride(2) — between state elements
|
||||
//
|
||||
// When stride_s_state == 1 (contiguous), the memmove fast path is used.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename scalar_t>
|
||||
inline void causal_conv1d_update_kernel(
|
||||
const scalar_t* __restrict__ x_ptr, scalar_t* __restrict__ state_ptr,
|
||||
int64_t stride_s_slot, int64_t stride_s_dim, int64_t stride_s_state,
|
||||
const scalar_t* __restrict__ weight_ptr, const float* __restrict__ bias_ptr,
|
||||
scalar_t* __restrict__ out_ptr, const int32_t* __restrict__ cache_idxs,
|
||||
int32_t pad_slot_id, int64_t batch, int64_t dim, int64_t seqlen,
|
||||
int64_t width, int64_t state_len, bool do_silu) {
|
||||
#pragma omp parallel for
|
||||
for (int64_t b = 0; b < batch; ++b) {
|
||||
int64_t cache_idx = (cache_idxs != nullptr) ? cache_idxs[b] : b;
|
||||
if (cache_idx == pad_slot_id) continue;
|
||||
|
||||
for (int64_t t = 0; t < seqlen; ++t) {
|
||||
const scalar_t* x_b = x_ptr + (b * dim * seqlen + t);
|
||||
scalar_t* out_b = out_ptr + (b * dim * seqlen + t);
|
||||
// Base of this slot in the (possibly non-contiguous) paged state
|
||||
scalar_t* s_base = state_ptr + cache_idx * stride_s_slot;
|
||||
|
||||
for (int64_t d = 0; d < dim; ++d) {
|
||||
float x_val = static_cast<float>(x_b[d * seqlen]);
|
||||
scalar_t* sd = s_base + d * stride_s_dim; // start of this dim's state
|
||||
const scalar_t* w = weight_ptr + d * width;
|
||||
|
||||
// Accumulate in float32 for precision
|
||||
float acc = (bias_ptr != nullptr) ? bias_ptr[d] : 0.0f;
|
||||
for (int64_t k = 0; k < state_len; ++k) {
|
||||
acc += static_cast<float>(w[k]) *
|
||||
static_cast<float>(sd[k * stride_s_state]);
|
||||
}
|
||||
acc += static_cast<float>(w[state_len]) * x_val;
|
||||
|
||||
// Shift state left and append new input.
|
||||
// Use memmove when contiguous (stride==1); element loop otherwise.
|
||||
if (stride_s_state == 1) {
|
||||
if (state_len > 1)
|
||||
std::memmove(sd, sd + 1, (state_len - 1) * sizeof(scalar_t));
|
||||
if (state_len > 0) sd[state_len - 1] = static_cast<scalar_t>(x_val);
|
||||
} else {
|
||||
for (int64_t k = 0; k < state_len - 1; ++k)
|
||||
sd[k * stride_s_state] = sd[(k + 1) * stride_s_state];
|
||||
if (state_len > 0)
|
||||
sd[(state_len - 1) * stride_s_state] = static_cast<scalar_t>(x_val);
|
||||
}
|
||||
|
||||
if (do_silu) {
|
||||
float sigmoid = (acc >= 0) ? 1.0f / (1.0f + std::exp(-acc))
|
||||
: std::exp(acc) / (1.0f + std::exp(acc));
|
||||
acc *= sigmoid;
|
||||
}
|
||||
out_b[d * seqlen] = static_cast<scalar_t>(acc);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// selective_state_update
|
||||
//
|
||||
// Template parameters:
|
||||
// state_t - dtype of ssm_state cache (typically BFloat16)
|
||||
// input_t - dtype of x, B, C (typically BFloat16)
|
||||
// out_t - dtype of output tensor (typically BFloat16)
|
||||
// Write directly — no float32 intermediate buffer needed.
|
||||
//
|
||||
// A, D, dt_bias are accepted as const float* (they are always float32
|
||||
// model parameters in Mamba2). This eliminates the per-call float32→BF16
|
||||
// conversion and the .contiguous() materialisation of the broadcast-expand.
|
||||
//
|
||||
// dt is accepted as a (N, nheads) scalar-per-head tensor, not as the
|
||||
// (N, nheads, head_dim) expansion, so no .contiguous() copy is needed.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename state_t, typename input_t, typename out_t = float>
|
||||
inline void selective_state_update_kernel(
|
||||
state_t* __restrict__ state_ptr, int64_t stride_state_n,
|
||||
int64_t stride_state_h, int64_t stride_state_d,
|
||||
const input_t* __restrict__ x_ptr, int64_t stride_x_n, int64_t stride_x_h,
|
||||
// dt: (N, nheads) — scalar per head, NOT expanded to head_dim
|
||||
const float* __restrict__ dt_ptr, int64_t stride_dt_n,
|
||||
// A: (nheads,) float32 — scalar per head
|
||||
const float* __restrict__ A_ptr, const input_t* __restrict__ B_ptr,
|
||||
const input_t* __restrict__ C_ptr, int64_t stride_BC_n, int64_t stride_BC_g,
|
||||
// D: (nheads,) float32 — scalar per head (nullptr if not used)
|
||||
const float* __restrict__ D_ptr,
|
||||
// z: same shape as x (optional)
|
||||
const input_t* __restrict__ z_ptr,
|
||||
// dt_bias: (nheads,) float32 — scalar per head (nullptr if not used)
|
||||
const float* __restrict__ dt_bias_ptr, out_t* __restrict__ out_ptr,
|
||||
int64_t stride_out_n, int64_t stride_out_h,
|
||||
const int32_t* __restrict__ state_batch_indices,
|
||||
const int32_t* __restrict__ dst_state_batch_indices, int32_t null_block_id,
|
||||
const int32_t* __restrict__ num_accepted_tokens,
|
||||
const int32_t* __restrict__ cu_seqlens, int64_t N, int64_t nheads,
|
||||
int64_t ngroups, int64_t dim, int64_t dstate, bool dt_softplus) {
|
||||
using state_vec_t = vec_op::vec_t<state_t>;
|
||||
using input_vec_t = vec_op::vec_t<input_t>;
|
||||
constexpr int VEC_ELEM_NUM = 8;
|
||||
|
||||
int64_t nheads_per_group = nheads / ngroups;
|
||||
|
||||
for (int64_t seq_idx = 0; seq_idx < N; ++seq_idx) {
|
||||
int64_t bos, seq_len;
|
||||
if (cu_seqlens != nullptr) {
|
||||
bos = cu_seqlens[seq_idx];
|
||||
seq_len = cu_seqlens[seq_idx + 1] - bos;
|
||||
} else {
|
||||
bos = seq_idx;
|
||||
seq_len = 1;
|
||||
}
|
||||
|
||||
int64_t state_read_idx = (state_batch_indices != nullptr)
|
||||
? state_batch_indices[seq_idx]
|
||||
: seq_idx;
|
||||
if (state_read_idx == null_block_id) continue;
|
||||
|
||||
int64_t state_write_idx = (num_accepted_tokens == nullptr)
|
||||
? ((dst_state_batch_indices != nullptr)
|
||||
? dst_state_batch_indices[seq_idx]
|
||||
: state_read_idx)
|
||||
: -1;
|
||||
|
||||
state_t* s = state_ptr + state_read_idx * stride_state_n;
|
||||
|
||||
for (int64_t t = 0; t < seq_len; ++t) {
|
||||
int64_t token_idx = bos + t;
|
||||
const input_t* x_tok = x_ptr + token_idx * stride_x_n;
|
||||
// dt: (N, nheads) — one float per head per token
|
||||
const float* dt_tok = dt_ptr + token_idx * stride_dt_n;
|
||||
const input_t* B_tok = B_ptr + token_idx * stride_BC_n;
|
||||
const input_t* C_tok = C_ptr + token_idx * stride_BC_n;
|
||||
out_t* out_tok = out_ptr + token_idx * stride_out_n;
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t h = 0; h < nheads; ++h) {
|
||||
int64_t g = h / nheads_per_group;
|
||||
const input_t* x_h = x_tok + h * stride_x_h;
|
||||
const input_t* B_g = B_tok + g * stride_BC_g;
|
||||
const input_t* C_g = C_tok + g * stride_BC_g;
|
||||
out_t* out_h = out_tok + h * stride_out_h;
|
||||
state_t* s_h = s + h * stride_state_h;
|
||||
|
||||
// Read scalars-per-head (A, dt, dt_bias, D) — no per-dim indexing
|
||||
float dt_val = dt_tok[h];
|
||||
if (dt_bias_ptr != nullptr) dt_val += dt_bias_ptr[h];
|
||||
if (dt_softplus) {
|
||||
dt_val = (dt_val <= 20.0f) ? std::log1p(std::exp(dt_val)) : dt_val;
|
||||
}
|
||||
const float A_val = A_ptr[h]; // scalar: same for all dim, dstate
|
||||
const float D_val = (D_ptr != nullptr) ? D_ptr[h] : 0.0f;
|
||||
|
||||
const input_t* z_h =
|
||||
(z_ptr != nullptr) ? z_ptr + token_idx * stride_x_n + h * stride_x_h
|
||||
: nullptr;
|
||||
|
||||
vec_op::FP32Vec8 dt_vec(dt_val);
|
||||
// dA = exp(A * dt): A and dt are SCALARS per head, so compute once
|
||||
// and broadcast. This saves 7 redundant std::exp() calls that
|
||||
// FP32Vec8::exp() would otherwise make on the broadcast vector.
|
||||
const float dA_scalar = std::exp(A_val * dt_val);
|
||||
vec_op::FP32Vec8 dA(dA_scalar); // broadcast
|
||||
|
||||
for (int64_t d = 0; d < dim; ++d) {
|
||||
float x_val = static_cast<float>(x_h[d]);
|
||||
|
||||
vec_op::FP32Vec8 out_vec(0.0f);
|
||||
state_t* s_hd = s_h + d * stride_state_d;
|
||||
const input_t* B_g_base = B_g;
|
||||
const input_t* C_g_base = C_g;
|
||||
|
||||
vec_op::FP32Vec8 x_vec(x_val);
|
||||
// dBx = B * x * dt — same dA for all dstate (A is scalar)
|
||||
// s_new = s * dA + B * x * dt
|
||||
|
||||
int64_t n = 0;
|
||||
for (; n <= dstate - VEC_ELEM_NUM; n += VEC_ELEM_NUM) {
|
||||
vec_op::FP32Vec8 B_v((input_vec_t(B_g_base + n)));
|
||||
vec_op::FP32Vec8 C_v((input_vec_t(C_g_base + n)));
|
||||
vec_op::FP32Vec8 s_v((state_vec_t(s_hd + n)));
|
||||
|
||||
vec_op::FP32Vec8 dBx = B_v * x_vec * dt_vec;
|
||||
vec_op::FP32Vec8 s_new = s_v * dA + dBx;
|
||||
|
||||
state_vec_t(s_new).save(s_hd + n);
|
||||
out_vec = out_vec + s_new * C_v;
|
||||
}
|
||||
|
||||
float out_val = out_vec.reduce_sum();
|
||||
for (; n < dstate; ++n) {
|
||||
// Reuse dA_scalar computed once per head — no exp() re-call
|
||||
float dBx = static_cast<float>(B_g[n]) * x_val * dt_val;
|
||||
float s_new = static_cast<float>(s_hd[n]) * dA_scalar + dBx;
|
||||
s_hd[n] = static_cast<state_t>(s_new);
|
||||
out_val += s_new * static_cast<float>(C_g[n]);
|
||||
}
|
||||
|
||||
if (D_ptr != nullptr) out_val += x_val * D_val;
|
||||
if (z_h != nullptr) {
|
||||
float z_val = static_cast<float>(z_h[d]);
|
||||
float sigmoid = (z_val >= 0)
|
||||
? 1.0f / (1.0f + std::exp(-z_val))
|
||||
: std::exp(z_val) / (1.0f + std::exp(z_val));
|
||||
out_val *= z_val * sigmoid;
|
||||
}
|
||||
out_h[d] = static_cast<out_t>(out_val);
|
||||
}
|
||||
}
|
||||
|
||||
if (num_accepted_tokens != nullptr &&
|
||||
dst_state_batch_indices != nullptr) {
|
||||
int64_t token_dst_idx = dst_state_batch_indices[seq_idx * seq_len + t];
|
||||
if (token_dst_idx != null_block_id && token_dst_idx != state_read_idx) {
|
||||
state_t* dst_s = state_ptr + token_dst_idx * stride_state_n;
|
||||
std::memmove(dst_s, s, nheads * stride_state_h * sizeof(state_t));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (num_accepted_tokens == nullptr && state_write_idx != null_block_id &&
|
||||
state_write_idx != state_read_idx) {
|
||||
state_t* dst_s = state_ptr + state_write_idx * stride_state_n;
|
||||
std::memmove(dst_s, s, nheads * stride_state_h * sizeof(state_t));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// mamba_chunk_scan_fwd
|
||||
//
|
||||
// Prefill SSM recurrence for Mamba2 / SSD models.
|
||||
//
|
||||
// Key difference from selective_state_update_kernel (decode path):
|
||||
// - #pragma omp parallel for collapse(2) is OUTSIDE the time loop.
|
||||
// Each thread owns a (batch, head) slice and runs the entire token
|
||||
// sequence without any per-token OpenMP synchronisation overhead.
|
||||
// For seqlen=256, this eliminates 256 thread-barrier launches per batch.
|
||||
//
|
||||
// `dt` arrives already processed (float32, after bias + softplus + clamp)
|
||||
// to keep this kernel simple. Preprocessing is done in the Python wrapper.
|
||||
//
|
||||
// `states_ptr` points to the [batch, nheads, headdim, dstate] float32 output
|
||||
// tensor, pre-initialised by the caller (zero or from initial_states).
|
||||
// Each (b, h) slice is private to exactly one thread via collapse(2), so
|
||||
// there are no write conflicts.
|
||||
//
|
||||
// D is treated as a scalar per head ([nheads] float32).
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename input_t>
|
||||
inline void mamba_chunk_scan_fwd_kernel(
|
||||
float* __restrict__ states_ptr, // [batch, nheads, headdim, dstate] f32
|
||||
const input_t* __restrict__ x_ptr, // [seqlen, nheads, headdim]
|
||||
const float* __restrict__ dt_ptr, // [seqlen, nheads] f32 (preprocessed)
|
||||
const float* __restrict__ A_ptr, // [nheads] f32
|
||||
const input_t* __restrict__ B_ptr, // [seqlen, ngroups, dstate]
|
||||
const input_t* __restrict__ C_ptr, // [seqlen, ngroups, dstate]
|
||||
const float* __restrict__ D_ptr, // [nheads] f32 (nullable)
|
||||
const input_t* __restrict__ z_ptr, // [seqlen, nheads, headdim] (nullable)
|
||||
input_t* __restrict__ out_ptr, // [seqlen, nheads, headdim]
|
||||
const int32_t* __restrict__ cu_seqlens, // [batch+1] int32
|
||||
int64_t batch, int64_t nheads, int64_t ngroups, int64_t headdim,
|
||||
int64_t dstate) {
|
||||
using input_vec_t = vec_op::vec_t<input_t>;
|
||||
constexpr int VEC_ELEM_NUM = 8;
|
||||
|
||||
const int64_t nheads_per_group = nheads / ngroups;
|
||||
// states layout: [batch, nheads, headdim, dstate] contiguous (caller
|
||||
// guarantee)
|
||||
const int64_t stride_s_b = nheads * headdim * dstate;
|
||||
const int64_t stride_s_h = headdim * dstate;
|
||||
// stride_s_d = dstate, stride_s_n = 1
|
||||
|
||||
#pragma omp parallel for collapse(2) schedule(static)
|
||||
for (int64_t b = 0; b < batch; ++b) {
|
||||
for (int64_t h = 0; h < nheads; ++h) {
|
||||
const int64_t seq_start = cu_seqlens[b];
|
||||
const int64_t seq_end = cu_seqlens[b + 1];
|
||||
const int64_t g = h / nheads_per_group;
|
||||
|
||||
const float A_val = A_ptr[h];
|
||||
const float D_val = (D_ptr != nullptr) ? D_ptr[h] : 0.0f;
|
||||
|
||||
// Working state slice: states[b, h, :, :] — float32, headdim * dstate.
|
||||
// Fits in L1/L2 for typical dims (e.g. 64*128*4 = 32 KB).
|
||||
float* s_bh = states_ptr + b * stride_s_b + h * stride_s_h;
|
||||
|
||||
for (int64_t t = seq_start; t < seq_end; ++t) {
|
||||
const input_t* x_h = x_ptr + t * nheads * headdim + h * headdim;
|
||||
const float* dt_h = dt_ptr + t * nheads + h;
|
||||
const input_t* B_g = B_ptr + t * ngroups * dstate + g * dstate;
|
||||
const input_t* C_g = C_ptr + t * ngroups * dstate + g * dstate;
|
||||
const input_t* z_h = (z_ptr != nullptr)
|
||||
? z_ptr + t * nheads * headdim + h * headdim
|
||||
: nullptr;
|
||||
input_t* out_h = out_ptr + t * nheads * headdim + h * headdim;
|
||||
|
||||
const float dt_val = *dt_h;
|
||||
const float dA_val = std::exp(A_val * dt_val);
|
||||
const vec_op::FP32Vec8 dA_vec(dA_val); // broadcast scalar
|
||||
const vec_op::FP32Vec8 dt_vec(dt_val);
|
||||
|
||||
for (int64_t d = 0; d < headdim; ++d) {
|
||||
const float x_val = static_cast<float>(x_h[d]);
|
||||
float* s_bhd = s_bh + d * dstate; // [dstate] contiguous float32
|
||||
|
||||
// Vectorised SSM update + readout over dstate:
|
||||
// s_new = s * dA + x * dt * B
|
||||
// y += s_new * C
|
||||
int64_t n = 0;
|
||||
vec_op::FP32Vec8 y_vec(0.0f);
|
||||
const vec_op::FP32Vec8 x_vec(x_val);
|
||||
|
||||
for (; n <= dstate - VEC_ELEM_NUM; n += VEC_ELEM_NUM) {
|
||||
const vec_op::FP32Vec8 B_v((input_vec_t(B_g + n)));
|
||||
const vec_op::FP32Vec8 C_v((input_vec_t(C_g + n)));
|
||||
const vec_op::FP32Vec8 s_v(s_bhd + n);
|
||||
|
||||
const vec_op::FP32Vec8 s_new = s_v * dA_vec + x_vec * dt_vec * B_v;
|
||||
s_new.save(s_bhd + n);
|
||||
y_vec = y_vec + s_new * C_v;
|
||||
}
|
||||
|
||||
float y_val = y_vec.reduce_sum();
|
||||
|
||||
// Scalar tail for remaining dstate elements
|
||||
for (; n < dstate; ++n) {
|
||||
const float B_n = static_cast<float>(B_g[n]);
|
||||
const float C_n = static_cast<float>(C_g[n]);
|
||||
const float s_new = s_bhd[n] * dA_val + x_val * dt_val * B_n;
|
||||
s_bhd[n] = s_new;
|
||||
y_val += s_new * C_n;
|
||||
}
|
||||
|
||||
// D skip connection (scalar per head)
|
||||
if (D_ptr != nullptr) y_val += x_val * D_val;
|
||||
|
||||
// z gating: out = y * z * sigmoid(z) (SiLU)
|
||||
if (z_h != nullptr) {
|
||||
const float z_val = static_cast<float>(z_h[d]);
|
||||
const float sigmoid =
|
||||
(z_val >= 0.0f) ? 1.0f / (1.0f + std::exp(-z_val))
|
||||
: std::exp(z_val) / (1.0f + std::exp(z_val));
|
||||
y_val *= z_val * sigmoid;
|
||||
}
|
||||
|
||||
out_h[d] = static_cast<input_t>(y_val);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace mamba_cpu
|
||||
@@ -213,6 +213,32 @@ void compute_slot_mapping_kernel_impl(const torch::Tensor query_start_loc,
|
||||
torch::Tensor slot_mapping,
|
||||
const int64_t block_size);
|
||||
|
||||
at::Tensor causal_conv1d_update_cpu_impl(
|
||||
at::Tensor& x, at::Tensor& conv_state, const at::Tensor& weight,
|
||||
const c10::optional<at::Tensor>& bias,
|
||||
const c10::optional<std::string>& activation,
|
||||
const c10::optional<at::Tensor>& conv_state_indices,
|
||||
const c10::optional<at::Tensor>& query_start_loc, int64_t pad_slot_id);
|
||||
|
||||
void selective_state_update_cpu_impl(
|
||||
at::Tensor& state, const at::Tensor& x, const at::Tensor& dt,
|
||||
const at::Tensor& A, const at::Tensor& B, const at::Tensor& C,
|
||||
const c10::optional<at::Tensor>& D, const c10::optional<at::Tensor>& z,
|
||||
const c10::optional<at::Tensor>& dt_bias, bool dt_softplus,
|
||||
const c10::optional<at::Tensor>& state_batch_indices,
|
||||
const c10::optional<at::Tensor>& dst_state_batch_indices,
|
||||
int64_t null_block_id, at::Tensor& out,
|
||||
const c10::optional<at::Tensor>& num_accepted_tokens,
|
||||
const c10::optional<at::Tensor>& cu_seqlens);
|
||||
|
||||
void mamba_chunk_scan_fwd_cpu_impl(at::Tensor& out, at::Tensor& final_states,
|
||||
const at::Tensor& x, const at::Tensor& dt,
|
||||
const at::Tensor& A, const at::Tensor& B,
|
||||
const at::Tensor& C,
|
||||
const c10::optional<at::Tensor>& D,
|
||||
const c10::optional<at::Tensor>& z,
|
||||
const at::Tensor& cu_seqlens);
|
||||
|
||||
void init_cpu_memory_env(std::vector<int64_t> node_ids);
|
||||
|
||||
namespace cpu_utils {
|
||||
@@ -595,6 +621,30 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"block_size) -> ()",
|
||||
&compute_slot_mapping_kernel_impl);
|
||||
|
||||
// Mamba CPU kernels
|
||||
ops.def(
|
||||
"causal_conv1d_update_cpu_vec("
|
||||
"Tensor(a0!) x, Tensor(a1!) conv_state, Tensor weight, "
|
||||
"Tensor? bias, str? activation, Tensor? conv_state_indices, "
|
||||
"Tensor? query_start_loc, SymInt pad_slot_id) -> Tensor",
|
||||
&causal_conv1d_update_cpu_impl);
|
||||
|
||||
ops.def(
|
||||
"selective_state_update_cpu("
|
||||
"Tensor(a0!) state, Tensor x, Tensor dt, Tensor A, Tensor B, Tensor C, "
|
||||
"Tensor? D, Tensor? z, Tensor? dt_bias, bool dt_softplus, "
|
||||
"Tensor? state_batch_indices, Tensor? dst_state_batch_indices, "
|
||||
"SymInt null_block_id, Tensor(a13!) out, "
|
||||
"Tensor? num_accepted_tokens, Tensor? cu_seqlens) -> ()",
|
||||
&selective_state_update_cpu_impl);
|
||||
|
||||
ops.def(
|
||||
"mamba_chunk_scan_fwd_cpu("
|
||||
"Tensor(a0!) out, Tensor(a1!) final_states, "
|
||||
"Tensor x, Tensor dt, Tensor A, Tensor B, Tensor C, "
|
||||
"Tensor? D, Tensor? z, Tensor cu_seqlens) -> ()",
|
||||
&mamba_chunk_scan_fwd_cpu_impl);
|
||||
|
||||
ops.def("init_cpu_memory_env(SymInt[] node_ids) -> ()", &init_cpu_memory_env);
|
||||
|
||||
// Speculative decoding kernels
|
||||
|
||||
@@ -1174,7 +1174,8 @@ __global__ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
const int32_t num_reqs, const int32_t block_size,
|
||||
const int32_t total_tokens, const int64_t block_table_stride,
|
||||
const int64_t cache_block_stride, const int64_t cache_entry_stride,
|
||||
const int64_t dst_entry_stride) {
|
||||
const int64_t dst_entry_stride,
|
||||
const int32_t* __restrict__ seq_starts) { // Optional source offsets
|
||||
const int flat_warp_id = (blockIdx.x * blockDim.x + threadIdx.x) >> 5;
|
||||
if (flat_warp_id >= total_tokens) return;
|
||||
const int lane_id = threadIdx.x & 31;
|
||||
@@ -1192,7 +1193,8 @@ __global__ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
|
||||
// Compute physical token address via block table
|
||||
const int out_token_id = flat_warp_id;
|
||||
const int token_offset = out_token_id - workspace_starts[req_id];
|
||||
int token_offset = out_token_id - workspace_starts[req_id];
|
||||
if (seq_starts != nullptr) token_offset += seq_starts[req_id];
|
||||
const int cache_block_idx = token_offset / block_size;
|
||||
const int offset_in_block = token_offset % block_size;
|
||||
const int physical_block =
|
||||
@@ -1383,9 +1385,9 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
torch::stable::Tensor const& src_cache, // [NUM_BLOCKS, BLOCK_SIZE, 656]
|
||||
torch::stable::Tensor const& dst, // [TOT_TOKENS, 576]
|
||||
torch::stable::Tensor const& block_table, // [BATCH, BLOCK_INDICES]
|
||||
torch::stable::Tensor const& seq_lens, // [BATCH]
|
||||
torch::stable::Tensor const& workspace_starts, // [BATCH]
|
||||
int64_t batch_size) {
|
||||
int64_t batch_size,
|
||||
std::optional<torch::stable::Tensor> seq_starts = std::nullopt) {
|
||||
torch::stable::accelerator::DeviceGuard device_guard(
|
||||
src_cache.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream();
|
||||
@@ -1396,20 +1398,25 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
STD_TORCH_CHECK(
|
||||
block_table.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"block_table must be int32");
|
||||
STD_TORCH_CHECK(seq_lens.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"seq_lens must be int32");
|
||||
STD_TORCH_CHECK(
|
||||
workspace_starts.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"workspace_starts must be int32");
|
||||
if (seq_starts.has_value()) {
|
||||
STD_TORCH_CHECK(
|
||||
seq_starts.value().scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"seq_starts must be int32");
|
||||
}
|
||||
|
||||
STD_TORCH_CHECK(src_cache.device() == dst.device(),
|
||||
"src_cache and dst must be on the same device");
|
||||
STD_TORCH_CHECK(src_cache.device() == block_table.device(),
|
||||
"src_cache and block_table must be on the same device");
|
||||
STD_TORCH_CHECK(src_cache.device() == seq_lens.device(),
|
||||
"src_cache and seq_lens must be on the same device");
|
||||
STD_TORCH_CHECK(src_cache.device() == workspace_starts.device(),
|
||||
"src_cache and workspace_starts must be on the same device");
|
||||
if (seq_starts.has_value()) {
|
||||
STD_TORCH_CHECK(src_cache.device() == seq_starts.value().device(),
|
||||
"src_cache and seq_starts must be on the same device");
|
||||
}
|
||||
auto dtype = src_cache.scalar_type();
|
||||
STD_TORCH_CHECK(
|
||||
dtype == torch::headeronly::ScalarType::Byte || // uint8
|
||||
@@ -1438,6 +1445,9 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
constexpr int warps_per_block = 8;
|
||||
const int grid_size = (total_tokens + warps_per_block - 1) / warps_per_block;
|
||||
const int block_size_threads = warps_per_block * 32; // 256 threads
|
||||
const int32_t* seq_starts_ptr =
|
||||
seq_starts.has_value() ? seq_starts.value().const_data_ptr<int32_t>()
|
||||
: nullptr;
|
||||
|
||||
vllm::cp_gather_and_upconvert_fp8_kv_cache<<<grid_size, block_size_threads, 0,
|
||||
stream>>>(
|
||||
@@ -1446,7 +1456,7 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
workspace_starts.const_data_ptr<int32_t>(),
|
||||
static_cast<int32_t>(batch_size), block_size, total_tokens,
|
||||
block_table_stride, cache_block_stride, cache_entry_stride,
|
||||
dst_entry_stride);
|
||||
dst_entry_stride, seq_starts_ptr);
|
||||
}
|
||||
|
||||
// Macro to dispatch the kernel based on the data type.
|
||||
|
||||
@@ -71,6 +71,73 @@ __device__ __forceinline__ float toFloat(T value) {
|
||||
}
|
||||
}
|
||||
|
||||
#ifndef USE_ROCM
|
||||
// Adapted from:
|
||||
// https://github.com/sgl-project/sglang/blob/main/python/sglang/jit_kernel/csrc/deepseek_v4/hash_topk.cuh
|
||||
template <typename OutIndType, typename HashIndType>
|
||||
__launch_bounds__(128) __global__
|
||||
void dsv4HashTopkSoftplusSqrt(const float* input, float* output,
|
||||
OutIndType* indices, int num_rows,
|
||||
int num_experts, float routed_scaling_factor,
|
||||
const HashIndType* input_ids,
|
||||
const HashIndType* tid2eid) {
|
||||
const int warp = (blockIdx.x * blockDim.x + threadIdx.x) / 32;
|
||||
const int lane = threadIdx.x % 32;
|
||||
if (warp >= num_rows) return;
|
||||
const int64_t token_id = load_index_as_int64(input_ids, warp);
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
cudaGridDependencySynchronize();
|
||||
#endif
|
||||
int expert = 0;
|
||||
float weight = 0.f;
|
||||
if (lane < 6) {
|
||||
// only load and calculate for 6 experts
|
||||
expert = static_cast<int>(tid2eid[token_id * 6 + lane]);
|
||||
const float x = input[warp * num_experts + expert];
|
||||
weight = sqrtf(fmaxf(x, 0.f) + __logf(1.f + __expf(-fabsf(x))));
|
||||
}
|
||||
float weight_sum = weight;
|
||||
#pragma unroll
|
||||
for (int mask = 16; mask > 0; mask >>= 1) {
|
||||
// sum in warp
|
||||
weight_sum += VLLM_SHFL_XOR_SYNC(weight_sum, mask);
|
||||
}
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
#endif
|
||||
if (lane < 6) {
|
||||
const int offset = warp * 6 + lane;
|
||||
output[offset] =
|
||||
weight * routed_scaling_factor / (weight_sum > 0.f ? weight_sum : 1.f);
|
||||
indices[offset] = static_cast<OutIndType>(expert);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename OutIndType, typename HashIndType>
|
||||
void launchDsv4HashTopk(const float* input, float* output, OutIndType* indices,
|
||||
int num_rows, int num_experts,
|
||||
double routed_scaling_factor,
|
||||
const HashIndType* input_ids,
|
||||
const HashIndType* tid2eid, cudaStream_t stream) {
|
||||
if (num_rows == 0) return;
|
||||
auto* kernel = &dsv4HashTopkSoftplusSqrt<OutIndType, HashIndType>;
|
||||
cudaLaunchConfig_t config = {};
|
||||
config.gridDim = (num_rows + 3) / 4;
|
||||
config.blockDim = 128;
|
||||
config.stream = stream;
|
||||
cudaLaunchAttribute attr;
|
||||
attr.id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attr.val.programmaticStreamSerializationAllowed = 1;
|
||||
config.attrs = &attr;
|
||||
config.numAttrs = 1;
|
||||
const float scale = static_cast<float>(routed_scaling_factor);
|
||||
cudaLaunchKernelEx(&config, kernel, input, output, indices, num_rows,
|
||||
num_experts, scale, input_ids, tid2eid);
|
||||
}
|
||||
#endif
|
||||
|
||||
// ====================== TopK softplus_sqrt things
|
||||
// ===============================
|
||||
|
||||
@@ -556,6 +623,17 @@ void topkGatingSoftplusSqrtKernelLauncher(
|
||||
const float* correction_bias, const bool use_hash,
|
||||
const HashIndType* input_ids, const HashIndType* tid2eid,
|
||||
cudaStream_t stream) {
|
||||
#ifndef USE_ROCM
|
||||
if constexpr (std::is_same_v<InputType, float>) {
|
||||
if (use_hash && topk == 6 && renormalize &&
|
||||
(num_experts == 256 || num_experts == 384)) {
|
||||
launchDsv4HashTopk<IndType, HashIndType>(
|
||||
gating_output, topk_weights, topk_indices, num_tokens, num_experts,
|
||||
routed_scaling_factor, input_ids, tid2eid, stream);
|
||||
return;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
static constexpr int WARPS_PER_TB = 4;
|
||||
static constexpr int BYTES_PER_LDG_POWER_OF_2 = 16;
|
||||
// for bfloat16 dtype, we need 4 bytes loading to make sure num_experts
|
||||
|
||||
@@ -527,9 +527,9 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
// 656]
|
||||
torch::stable::Tensor const& dst, // [TOT_TOKENS, 576]
|
||||
torch::stable::Tensor const& block_table, // [BATCH, BLOCK_INDICES]
|
||||
torch::stable::Tensor const& seq_lens, // [BATCH]
|
||||
torch::stable::Tensor const& workspace_starts, // [BATCH]
|
||||
int64_t batch_size);
|
||||
int64_t batch_size,
|
||||
std::optional<torch::stable::Tensor> seq_starts = std::nullopt);
|
||||
|
||||
// Indexer K quantization and cache function
|
||||
void indexer_k_quant_and_cache(
|
||||
|
||||
@@ -39,11 +39,15 @@ __global__ void marlin_int4_fp8_preprocess_kernel_awq(
|
||||
// AWQ zeros: (size_k // group_size, size_n // 8)
|
||||
const int32_t* __restrict__ qzeros, int32_t size_n, int32_t size_k,
|
||||
int32_t group_size) {
|
||||
int32_t val =
|
||||
qweight[(blockIdx.x * 32 + threadIdx.x) * size_n / 8 + blockIdx.y];
|
||||
int32_t zero =
|
||||
qzeros[(blockIdx.x * 32 + threadIdx.x) / group_size * size_n / 8 +
|
||||
blockIdx.y];
|
||||
// Thread mapping: threadIdx.x -> column dim (coalesced read within a row),
|
||||
// blockIdx.x -> row dim. Adjacent threads read consecutive int32 in the
|
||||
// same row (stride 1) instead of striding across rows (stride size_n/8).
|
||||
int col = blockIdx.y * 32 + threadIdx.x;
|
||||
if (col >= size_n / 8) return;
|
||||
(void)size_k;
|
||||
|
||||
int32_t val = qweight[blockIdx.x * (size_n / 8) + col];
|
||||
int32_t zero = qzeros[blockIdx.x / group_size * (size_n / 8) + col];
|
||||
int32_t new_val = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -58,7 +62,7 @@ __global__ void marlin_int4_fp8_preprocess_kernel_awq(
|
||||
zero >>= 4;
|
||||
}
|
||||
|
||||
output[(blockIdx.x * 32 + threadIdx.x) * size_n / 8 + blockIdx.y] = new_val;
|
||||
output[blockIdx.x * (size_n / 8) + col] = new_val;
|
||||
}
|
||||
|
||||
torch::stable::Tensor marlin_int4_fp8_preprocess(
|
||||
@@ -102,7 +106,7 @@ torch::stable::Tensor marlin_int4_fp8_preprocess(
|
||||
"qweight.size(0) % qzeros.size(0) != 0");
|
||||
STD_TORCH_CHECK(group_size % 8 == 0, "group_size % 8 != 0");
|
||||
|
||||
dim3 blocks(size_k / 32, size_n / 8);
|
||||
dim3 blocks(size_k, (size_n / 8 + 31) / 32);
|
||||
marlin_int4_fp8_preprocess_kernel_awq<<<blocks, 32, 0, stream>>>(
|
||||
reinterpret_cast<const int32_t*>(qweight.const_data_ptr()),
|
||||
reinterpret_cast<int32_t*>(output.mutable_data_ptr()),
|
||||
|
||||
@@ -847,8 +847,8 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C_cache_ops, ops) {
|
||||
|
||||
ops.def(
|
||||
"cp_gather_and_upconvert_fp8_kv_cache(Tensor src_cache, Tensor! dst, "
|
||||
"Tensor block_table, Tensor seq_lens, Tensor workspace_starts, int "
|
||||
"batch_size) -> ()");
|
||||
"Tensor block_table, Tensor workspace_starts, int batch_size, Tensor? "
|
||||
"seq_starts) -> ()");
|
||||
|
||||
ops.def(
|
||||
"indexer_k_quant_and_cache(Tensor k, Tensor! kv_cache, Tensor "
|
||||
|
||||
@@ -294,6 +294,9 @@ FROM base AS rust-build
|
||||
ARG BUILD_OS
|
||||
ARG USE_SCCACHE
|
||||
ARG SCCACHE_ENDPOINT
|
||||
# Temporary default for the initial CI validation. Set this back to 0 when
|
||||
# ci-infra passes VLLM_RUST_COVERAGE=1 explicitly.
|
||||
ARG VLLM_RUST_COVERAGE=1
|
||||
|
||||
# Install native tools needed only for Rust/protoc builds.
|
||||
RUN if [ "${BUILD_OS}" = "manylinux" ]; then \
|
||||
@@ -902,6 +905,14 @@ COPY ./vllm/collect_env.py .
|
||||
# note that this uses vllm installed by `pip`
|
||||
FROM vllm-base AS test
|
||||
|
||||
COPY --from=rust-build \
|
||||
/workspace/rust-coverage-tools/ \
|
||||
/opt/vllm-rust-coverage/
|
||||
|
||||
ENV PATH=/opt/vllm-rust-coverage/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=/opt/vllm-rust-coverage/lib:${LD_LIBRARY_PATH}
|
||||
ENV LLVM_PROFILE_FILE=/dev/null
|
||||
|
||||
ADD . /vllm-workspace/
|
||||
|
||||
ARG PYTHON_VERSION
|
||||
|
||||
@@ -46,6 +46,9 @@
|
||||
"TORCH_CUDA_ARCH_LIST": {
|
||||
"default": "7.5 8.0 8.6 8.9 9.0 10.0 11.0 12.0"
|
||||
},
|
||||
"VLLM_RUST_COVERAGE": {
|
||||
"default": "1"
|
||||
},
|
||||
"MAX_JOBS": {
|
||||
"default": "2"
|
||||
},
|
||||
|
||||
@@ -5,7 +5,7 @@ vLLM uses the following environment variables to configure the system:
|
||||
!!! warning
|
||||
Please note that `VLLM_PORT` and `VLLM_HOST_IP` set the port and ip for vLLM's **internal usage**. It is not the port and ip for the API server. If you use `--host $VLLM_HOST_IP` and `--port $VLLM_PORT` to start the API server, it will not work.
|
||||
|
||||
All environment variables used by vLLM are prefixed with `VLLM_`. **Special care should be taken for Kubernetes users**: please do not name the service as `vllm`, otherwise environment variables set by Kubernetes might conflict with vLLM's environment variables, because [Kubernetes sets environment variables for each service with the capitalized service name as the prefix](https://kubernetes.io/docs/concepts/services-networking/service/#environment-variables).
|
||||
Most vLLM-specific environment variables are prefixed with `VLLM_` (a handful of standard names — for example `CUDA_VISIBLE_DEVICES`, `MAX_JOBS`, `S3_ACCESS_KEY_ID`/`S3_SECRET_ACCESS_KEY`/`S3_ENDPOINT_URL`, `DO_NOT_TRACK`, `NO_COLOR` — are also read directly when set). **Special care should be taken for Kubernetes users**: please do not name the service as `vllm`, otherwise environment variables set by Kubernetes might conflict with vLLM's environment variables, because [Kubernetes sets environment variables for each service with the capitalized service name as the prefix](https://kubernetes.io/docs/concepts/services-networking/service/#environment-variables).
|
||||
|
||||
```python
|
||||
--8<-- "vllm/envs.py:env-vars-definition"
|
||||
|
||||
@@ -6,7 +6,11 @@ vLLM maintains a per-commit wheel repository (commonly referred to as "nightly")
|
||||
|
||||
### Wheel Building
|
||||
|
||||
Wheels are built in the `Release` pipeline (`.buildkite/release-pipeline.yaml`) after a PR is merged into the main branch, with multiple variants:
|
||||
Wheels are built in the `Release` pipeline
|
||||
(`.buildkite/release-pipeline.yaml`) after a PR is merged into the main branch.
|
||||
Regular builds produce the CUDA 13.0 wheels for x86_64 and aarch64. Additional
|
||||
wheel variants and ROCm builds can be unblocked on demand and run automatically
|
||||
when `NIGHTLY=1`:
|
||||
|
||||
- **Backend variants**: `cpu` and `cuXXX` (e.g., `cu129`, `cu130`).
|
||||
- **Architecture variants**: `x86_64` and `aarch64`.
|
||||
|
||||
@@ -164,8 +164,8 @@ Priority is **1 = highest** (tried first).
|
||||
| `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`, `fp8_e5m2` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | 9.x |
|
||||
| `FLASH_ATTN` | FA4* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | ≥10.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 |
|
||||
@@ -205,7 +205,7 @@ 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=192, qk_rope_head_dim=64, v_head_dim=256) (FA2/FA3 only) |
|
||||
| `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 |
|
||||
|
||||
@@ -306,7 +306,7 @@ Supported quantization scheme/hardware combinations:
|
||||
|
||||
- Pass: [`vllm/compilation/passes/fusion/rms_quant_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/rms_quant_fusion.py)
|
||||
- ROCm AITER pass: [`vllm/compilation/passes/fusion/rocm_aiter_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/rocm_aiter_fusion.py)
|
||||
- CUDA/HIP kernels: [`csrc/layernorm_quant_kernels.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/layernorm_quant_kernels.cu)
|
||||
- CUDA/HIP kernels: [`csrc/libtorch_stable/layernorm_quant_kernels.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/libtorch_stable/layernorm_quant_kernels.cu)
|
||||
|
||||
### SiLU+Mul + Quantization (`fuse_act_quant`)
|
||||
|
||||
@@ -332,7 +332,7 @@ Supported quantization scheme/hardware combinations:
|
||||
- Pass: [`vllm/compilation/passes/fusion/act_quant_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/act_quant_fusion.py)
|
||||
- ROCm AITER pass: [`vllm/compilation/passes/fusion/rocm_aiter_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/rocm_aiter_fusion.py)
|
||||
- CUDA/HIP kernels: [`csrc/quantization/`](https://github.com/vllm-project/vllm/blob/main/csrc/quantization/)
|
||||
- Fused SiLU+Mul+BlockQuant kernel: [`csrc/quantization/fused_kernels/fused_silu_mul_block_quant.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/quantization/fused_kernels/fused_silu_mul_block_quant.cu)
|
||||
- Fused SiLU+Mul+BlockQuant kernel: [`csrc/libtorch_stable/quantization/fused_kernels/fused_silu_mul_block_quant.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/libtorch_stable/quantization/fused_kernels/fused_silu_mul_block_quant.cu)
|
||||
|
||||
### RMSNorm + Padding (`fuse_act_padding`)
|
||||
|
||||
|
||||
@@ -107,6 +107,7 @@ Batch invariance has been tested and verified on the following models:
|
||||
- **Llama 3**: Llama3.1 and 3.2 series, `meta-llama/Llama-3.2-3B-Instruct` for example
|
||||
- **GPT-OSS**: `openai/gpt-oss-20b`, `openai/gpt-oss-120b`
|
||||
- **Mistral**: `mistralai/Mistral-7B-v0.3`
|
||||
- **Phi series**: `microsoft/Phi-3.5-mini-instruct`
|
||||
|
||||
Other models may also work, but these have been explicitly validated. If you encounter issues with a specific model, please report them on the [GitHub issue tracker](https://github.com/vllm-project/vllm/issues/new/choose).
|
||||
|
||||
|
||||
@@ -68,13 +68,14 @@ vllm serve <model> \
|
||||
| --- | --- | --- | --- | --- |
|
||||
| `spec_name` | no | `CPUOffloadingSpec` | both | Set to `TieringOffloadingSpec` for multi-tier. |
|
||||
| `cpu_bytes_to_use` | yes | — | both | Total bytes of host memory reserved for the CPU tier across all workers (not per-worker). |
|
||||
| `block_size` | no | GPU block size | both | Offloaded block size in tokens; must be a multiple of the GPU block size. |
|
||||
| `block_size` | no | GPU block size | both | Offloaded block size in tokens; must be a multiple of the GPU block size. Mutually exclusive with `blocks_per_chunk`. |
|
||||
| `blocks_per_chunk` | no | `1` | both | Offloaded chunk size in GPU blocks; must be > 0. Alternative to `block_size` for models whose KV cache groups have different block sizes. |
|
||||
| `eviction_policy` | no | `lru` | both | Primary tier policy: `lru` or `arc`. |
|
||||
| `store_threshold` | no | `0` | single-tier | Min lookups before a block is offloaded. Values ≥ 2 are rejected by `TieringOffloadingSpec`. |
|
||||
| `max_tracker_size` | no | `64000` | single-tier | Max entries in the lookup tracker. |
|
||||
| `secondary_tiers` | no | `[]` | multi-tier | List of secondary tier configs (see below). |
|
||||
| `offload_prompt_only` | no | `true` | both | If `true`, only prompt (prefill) blocks are offloaded; decode blocks are skipped. |
|
||||
| `self_describing_kv_events` | no | `false` | single-tier | Opt-in. When `true` *and* KV cache events are enabled (`--kv-events-config` with `enable_kv_cache_events`), the connector emits self-describing block-granular `BlockStored`/`BlockRemoved` payloads (constituent block hashes, whole-chunk `token_ids`, per-block `block_size`, parent hash, LoRA + group/cache-spec metadata) instead of the placeholder fallback, so external KV-event consumers can index offloaded blocks. Inert unless events are enabled. Currently rejected by `TieringOffloadingSpec`. Full-attention groups only; sliding-window/SSM groups keep the placeholder fallback. In chunk mode (`block_size` > GPU block size), overlapping chunks re-announce shared per-block hashes, so consumers must reference-count (deduplicate) repeated store/remove announcements. |
|
||||
| `self_describing_kv_events` | no | `false` | both | Opt-in. When `true` *and* KV cache events are enabled (`--kv-events-config` with `enable_kv_cache_events`), the connector emits self-describing block-granular `BlockStored`/`BlockRemoved` payloads (constituent block hashes, whole-chunk `token_ids`, per-block `block_size`, parent hash, LoRA + group/cache-spec metadata) instead of the placeholder fallback, so external KV-event consumers can index offloaded blocks. Inert unless events are enabled. With `TieringOffloadingSpec`, a CPU promotion is self-describing when a local request observes its primary-tier `HIT` before event translation; otherwise its stored event may retain the placeholder, while a later `HIT` can backfill metadata for removal. Pending-removal/re-promotion races and externally initiated promotions may also produce placeholders, and consumers must ignore removals for unknown hashes. Full-attention groups only; sliding-window/SSM groups keep the placeholder fallback. In chunk mode (`block_size` > GPU block size, or `blocks_per_chunk` > 1), overlapping chunks re-announce shared per-block hashes, so consumers must reference-count (deduplicate) repeated store/remove announcements. |
|
||||
| `spec_module_path` | no | — | both | Python import path for a custom `OffloadingSpec` not in the built-in registry. Required only when `spec_name` is not built-in (advanced). |
|
||||
|
||||
## Secondary Tiers
|
||||
@@ -83,9 +84,11 @@ Each entry in `secondary_tiers` is a dict with a required `type` field plus tier
|
||||
|
||||
The filesystem and object-store tiers can publish hash-only `BlockStored` KV events for blocks they successfully store, tagged with a stable per-tier `medium` (`FS` for the filesystem tier, `OBJ` for the object-store tier). Set `enable_kv_events: true` in the tier's entry to opt in; events are published only when KV cache events are also enabled globally via `--kv-events-config`.
|
||||
|
||||
Set the optional `locality` tier field to `LOCAL` or `REMOTE` to describe the tier's storage location relative to the publishing vLLM instance. `LOCAL` marks storage local to that instance, while `REMOTE` marks storage that is not local to it. When the setting is omitted, locality is unspecified. vLLM does not infer it from the tier type, so an OBJ tier is not implicitly `REMOTE`. A KV event includes `locality` only when the tier explicitly configures it. This metadata describes the tier property without implying that a consumer can already route requests to its blocks.
|
||||
|
||||
### Filesystem (FS)
|
||||
|
||||
The filesystem tier (`type: "fs"`) writes blocks to a directory on local storage.
|
||||
The filesystem tier (`type: "fs"`) writes blocks to a filesystem directory.
|
||||
|
||||
| Key | Required | Default | Notes |
|
||||
| --- | --- | --- | --- |
|
||||
@@ -94,6 +97,7 @@ The filesystem tier (`type: "fs"`) writes blocks to a directory on local storage
|
||||
| `n_read_threads` | no | `16` | Read-priority I/O threads (load path). |
|
||||
| `n_write_threads` | no | `16` | Write-priority I/O threads (store path). |
|
||||
| `enable_kv_events` | no | `false` | Publish `BlockStored` KV events (medium `FS`) for successfully stored blocks. Requires KV cache events to be enabled globally. |
|
||||
| `locality` | no | unspecified | `LOCAL` or `REMOTE` relative to the publishing vLLM instance. Included in the tier's KV events only when explicitly configured. |
|
||||
|
||||
Each thread group prefers its own queue but pulls from the other when its primary queue is empty, so a write-heavy or read-heavy burst won't leave the off-priority queue waiting. Size the totals to your storage's effective concurrency.
|
||||
|
||||
@@ -134,6 +138,7 @@ The object-store tier (`type: "obj"`) offloads blocks to an S3-compatible object
|
||||
| `prefix` | no | `""` | Key prefix prepended to all object keys. |
|
||||
| `io_threads` | no | `4` | Number of NIXL OBJ backend I/O threads. |
|
||||
| `enable_kv_events` | no | `false` | Publish `BlockStored` KV events (medium `OBJ`) for successfully stored blocks. Requires KV cache events to be enabled globally. |
|
||||
| `locality` | no | unspecified | `LOCAL` or `REMOTE` relative to the publishing vLLM instance. Included in the tier's KV events only when explicitly configured; OBJ does not imply `REMOTE`. |
|
||||
|
||||
`store_config` fields:
|
||||
|
||||
@@ -175,7 +180,7 @@ Rather than embedding `host`/`port` in each `secondary_tiers` entry, set them on
|
||||
|
||||
- `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.
|
||||
- For single-tier (CPU-only) setups, set `cpu_bytes_to_use` larger than the aggregate GPU KV cache. Because offloading is immediate, a smaller CPU tier just mirrors what the GPU already holds and adds no hit rate.
|
||||
- `block_size`: larger offloaded blocks reduce per-block bookkeeping overhead but increase the granularity of lookups. Must be a multiple of the GPU block size.
|
||||
- `block_size` / `blocks_per_chunk`: larger offloaded chunks reduce per-block bookkeeping overhead but increase the granularity of lookups.
|
||||
- FS thread counts: tune `n_read_threads` and `n_write_threads` to the parallelism your storage can sustain. Reads are latency-sensitive on the prefill path, so prefer more read threads when prefill hit rates are high.
|
||||
- Sharing `root_dir` across runs: runs with the same model, `block_size`, parallelism layout, and dtype share files under the same `<digest>` subdirectory. Changing any of these produces a new subdirectory; old ones are orphaned but harmless. Delete them to reclaim disk.
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ Currently, there are no pre-built XPU wheels.
|
||||
|
||||
- First, install required [driver](https://dgpu-docs.intel.com/driver/installation.html#installing-gpu-drivers).
|
||||
- Second, install Python packages for vLLM XPU backend building (Intel OneAPI dependencies are installed automatically as part of `torch-xpu`, see [PyTorch XPU get started](https://docs.pytorch.org/docs/stable/notes/get_start_xpu.html)):
|
||||
- Start from vllm-xpu-kernels v0.1.10, we recommend user upgrade driver to [compute runtime 26.18](https://github.com/intel/compute-runtime/releases/tag/26.14.37833.4) release, to avoid potential compatibility issue.
|
||||
- Start from vllm-xpu-kernels v0.1.10, we recommend user upgrade driver to [compute runtime 26.18](https://github.com/intel/compute-runtime/releases/tag/26.18.38308.1) release, to avoid potential compatibility issue.
|
||||
|
||||
```bash
|
||||
git clone https://github.com/vllm-project/vllm.git
|
||||
@@ -58,7 +58,40 @@ VLLM_TARGET_DEVICE=xpu pip install --no-build-isolation -e . -v
|
||||
--8<-- [end:build-wheel-from-source]
|
||||
--8<-- [start:pre-built-images]
|
||||
|
||||
Currently, we release prebuilt XPU images at docker [hub](https://hub.docker.com/r/intel/vllm/tags) based on vLLM released version. For more information, please refer release [note](https://github.com/intel/ai-containers/blob/main/vllm).
|
||||
vLLM offers official Docker images for deployment.
|
||||
The images can be used to run OpenAI compatible server and are available on Docker Hub as [vllm/vllm-openai-xpu](https://hub.docker.com/r/vllm/vllm-openai-xpu/tags).
|
||||
|
||||
- `vllm/vllm-openai-xpu:latest` — stable release, available starting from v0.26.0
|
||||
- `vllm/vllm-openai-xpu:nightly` — preview build from the latest development branch, use this if you want the latest features and fixes
|
||||
|
||||
```bash
|
||||
docker run --rm \
|
||||
--network=host \
|
||||
--device /dev/dri:/dev/dri \
|
||||
-v /dev/dri/by-path:/dev/dri/by-path \
|
||||
-v ~/.cache/huggingface:/root/.cache/huggingface \
|
||||
--env "HF_TOKEN=$HF_TOKEN" \
|
||||
--ipc=host \
|
||||
--privileged \
|
||||
vllm/vllm-openai-xpu:<tag> \
|
||||
--model Qwen/Qwen3-0.6B
|
||||
```
|
||||
|
||||
To use the docker image as base for development, you can launch it in interactive session through overriding the entrypoint.
|
||||
|
||||
???+ console "Commands"
|
||||
```bash
|
||||
docker run --rm -it \
|
||||
--network=host \
|
||||
--device /dev/dri:/dev/dri \
|
||||
-v /dev/dri/by-path:/dev/dri/by-path \
|
||||
-v ~/.cache/huggingface:/root/.cache/huggingface \
|
||||
--env "HF_TOKEN=$HF_TOKEN" \
|
||||
--ipc=host \
|
||||
--privileged \
|
||||
--entrypoint /bin/bash \
|
||||
vllm/vllm-openai-xpu:<tag>
|
||||
```
|
||||
|
||||
--8<-- [end:pre-built-images]
|
||||
--8<-- [start:build-image-from-source]
|
||||
|
||||
@@ -65,6 +65,15 @@ This guide will help you quickly get started with vLLM to perform:
|
||||
!!! tip
|
||||
A nightly Docker image is also available as [vllm/vllm-openai-rocm:nightly](https://hub.docker.com/r/vllm/vllm-openai-rocm/tags) for testing the latest development builds.
|
||||
|
||||
=== "Intel GPU"
|
||||
|
||||
vLLM supports Intel GPUs through the XPU backend. Pre-built XPU wheels will be available soon.
|
||||
|
||||
Official Docker images for Intel GPUs are added to the vLLM release starting from v0.26.0. Nightly Docker image is also available as [vllm/vllm-openai-xpu:nightly](https://hub.docker.com/r/vllm/vllm-openai-xpu/tags).
|
||||
|
||||
!!! tip
|
||||
For more detailed instructions, including building from source and Docker image setup, please refer to the [GPU installation guide](installation/gpu.md) and select the "Intel XPU" tab.
|
||||
|
||||
=== "Google TPU"
|
||||
|
||||
To run vLLM on Google TPUs, you need to install the `vllm-tpu` package.
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
{% extends "base.html" %}
|
||||
|
||||
{% block announce %}
|
||||
<p>You are viewing the latest developer preview docs. <a href="https://docs.vllm.ai/en/stable/">Click here</a> to view docs for the latest stable release.</p>
|
||||
<p>You are viewing the latest developer preview docs. <a href="https://docs.vllm.ai/en/stable/{{ page.url }}">Click here</a> to view docs for the latest stable release.</p>
|
||||
{% endblock %}
|
||||
|
||||
@@ -31,10 +31,8 @@
|
||||
| THUDM/CodeGeex4-All-9B | CodeGeexForCausalLM | ✅ | | |
|
||||
| chuhac/TeleChat2-35B | LlamaForCausalLM (TeleChat2 based on Llama arch) | ✅ | | |
|
||||
| 01-ai/Yi1.5-34B-Chat | YiForCausalLM | ✅ | | |
|
||||
| THUDM/CodeGeex4-All-9B | CodeGeexForCausalLM | ✅ | | |
|
||||
| deepseek-ai/DeepSeek-Coder-33B-base | DeepSeekCoderForCausalLM | ✅ | | |
|
||||
| meta-llama/Llama-2-13b-chat-hf | LlamaForCausalLM | ✅ | | |
|
||||
| THUDM/CodeGeex4-All-9B | CodeGeexForCausalLM | ✅ | | |
|
||||
| Qwen/Qwen1.5-14B-Chat | QwenForCausalLM | ✅ | | |
|
||||
| Qwen/Qwen1.5-32B-Chat | QwenForCausalLM | ✅ | | |
|
||||
| RedHatAI/Meta-Llama-3.1-8B-Instruct-FP8-dynamic | LlamaForCausalLM | | ✅ | |
|
||||
|
||||
@@ -220,8 +220,6 @@ For multi-node deployment, add these EPLB flags to each node's command. We recom
|
||||
|
||||
- Use simulator flags `VLLM_MOE_ROUTING_SIMULATION_STRATEGY=uniform_random` and `VLLM_RANDOMIZE_DP_DUMMY_INPUTS=1` so token routing is balanced across EP ranks.
|
||||
|
||||
- Increasing `VLLM_MOE_DP_CHUNK_SIZE` may increase throughput by increasing the maximum batch size for inter-rank token transfers. This may cause DeepEP to throw `assert self.nvshmem_qp_depth >= (num_max_dispatch_tokens_per_rank + 1) * 2`, which can be fixed by increasing environment variable `NVSHMEM_QP_DEPTH`.
|
||||
|
||||
## Disaggregated Serving (Prefill/Decode Split)
|
||||
|
||||
For production deployments requiring strict SLA guarantees for time-to-first-token and inter-token latency, disaggregated serving allows independent scaling of prefill and decode operations.
|
||||
|
||||
@@ -137,8 +137,12 @@ For further details on renderer APIs, please refer to [this page](renderer.md).
|
||||
|
||||
### Derenderer APIs
|
||||
|
||||
- `/v1/completions/derender` - Derenderer completion requests
|
||||
- `/v1/chat/completions/derender` - Derenderer chat completion requests
|
||||
For further details on derenderer APIs, please refer to [this page](derenderer.md).
|
||||
|
||||
- [Chat Completions Derender API](derenderer.md) (`/v1/chat/completions/derender`)
|
||||
- Derender chat completion requests
|
||||
- [Completions Derender API](derenderer.md) (`/v1/completions/derender`)
|
||||
- Derender completion requests
|
||||
|
||||
## Tokenize APIs
|
||||
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
# Derenderer APIs
|
||||
|
||||
The derenderer API is the post processing counterpart to the [Renderer APIs](renderer.md). Where `/render` turns a request into token ID (preprocessing), `/derender` turns generated token IDs back into a fully formed OpenAI compatible response (detokenization, reasoning parsing, tool call parsing), all without a GPU.
|
||||
|
||||
This closes the loop for a token-in / token-out engine in disaggregated serving:
|
||||
|
||||
- **GPU less post processing**: Detokenization, reasoning parsing, and tool call parsing run on the same GPU less frontend that hosts `/render`
|
||||
- **Parser parity**: The derenderer reuses vLLM's tool and reasoning parsers, so a disaggregated deployment produces the same `content`/`reasoning`/ `tool_calls` split as a standard `vllm serve` server
|
||||
- **Non-streaming**: The endpoints expect a complete `GenerateResponse` with all token IDs present and perform one-shot parsing. Streaming derender would require a separate endpoint design and is not currently supported but is in the pipeline
|
||||
|
||||
Both endpoints are hosted by the GPU less rendering server started with [`vllm launch render`](../../cli/launch/render.md), alongside the `/render`
|
||||
endpoints.
|
||||
|
||||
## Pipeline
|
||||
|
||||
```text
|
||||
render generate derender
|
||||
request ───────────────▶ token_ids ─────────▶ token_ids ──────────▶ response
|
||||
(chat / (GPU less) (token-in / (GPU less) (OpenAI
|
||||
completion) │ token-out engine) ▲ compatible)
|
||||
└─────────────── request + prompt_tokens ──┘
|
||||
```
|
||||
|
||||
The derender step needs more than the engine's `token_ids`. It also consumes the original `chat_request`/`completion_request` and `prompt_tokens` carried over from the render step (see [Request format](#request-format)) so the tool and reasoning parsers have the context they need.
|
||||
|
||||
## API Reference
|
||||
|
||||
- Chat Completions Derender API (`/v1/chat/completions/derender`)
|
||||
- Post process a single `GenerateResponse` into a `ChatCompletionResponse`
|
||||
- Completions Derender API (`/v1/completions/derender`)
|
||||
- Post process a list of `GenerateResponse` objects (one per prompt) into a `CompletionResponse`
|
||||
|
||||
## Request format
|
||||
|
||||
Each request wraps the engine's `GenerateResponse`(s) together with the caller metadata needed to reconstruct the final response without a GPU.
|
||||
|
||||
`/v1/chat/completions/derender`:
|
||||
|
||||
??? code
|
||||
|
||||
```python
|
||||
--8<-- "vllm/entrypoints/scale_out/token_in_token_out/protocol.py:derender-chat-request"
|
||||
```
|
||||
|
||||
`/v1/completions/derender`:
|
||||
|
||||
??? code
|
||||
|
||||
```python
|
||||
--8<-- "vllm/entrypoints/scale_out/token_in_token_out/protocol.py:derender-completion-request"
|
||||
```
|
||||
|
||||
Oversized payloads are rejected with a `400` before any `tokenizer.decode()` or parser runs.
|
||||
|
||||
## Example
|
||||
|
||||
The example below drives the full `render → generate → derender` round trip for a chat request against a GPU less render server (`/render`, `/derender`) and a token-in / token-out engine (`/inference/v1/generate`).
|
||||
|
||||
```python
|
||||
import httpx
|
||||
|
||||
MODEL = "meta-llama/Llama-3.2-1B-Instruct"
|
||||
RENDER = "http://localhost:8100" # vllm launch render ...
|
||||
ENGINE = "http://localhost:8200" # token-in / token-out engine
|
||||
|
||||
chat_request = {
|
||||
"model": MODEL,
|
||||
"messages": [{"role": "user", "content": "What is 2+2?"}],
|
||||
"max_tokens": 32,
|
||||
}
|
||||
|
||||
with httpx.Client(timeout=60.0) as client:
|
||||
# 1. Render: request -> token IDs (GPU less)
|
||||
generate_request = client.post(
|
||||
f"{RENDER}/v1/chat/completions/render", json=chat_request
|
||||
).json()
|
||||
prompt_tokens = len(generate_request["token_ids"])
|
||||
|
||||
# 2. Generate: token IDs -> token IDs (token-in / token-out engine)
|
||||
generate_response = client.post(
|
||||
f"{ENGINE}/inference/v1/generate", json=generate_request
|
||||
).json()
|
||||
|
||||
# 3. Derender: token IDs -> ChatCompletionResponse (GPU less)
|
||||
response = client.post(
|
||||
f"{RENDER}/v1/chat/completions/derender",
|
||||
json={
|
||||
"model": MODEL,
|
||||
"generate_response": generate_response,
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"chat_request": chat_request,
|
||||
},
|
||||
).json()
|
||||
|
||||
print(response["choices"][0]["message"]["content"])
|
||||
```
|
||||
|
||||
Passing `chat_request` lets the derenderer run the configured tool and reasoning parsers. This means `response["choices"][0]["message"]` carries the same `content` / `reasoning` / `tool_calls` split a `vllm serve` server would produce. Omit `chat_request` for plain detokenization only.
|
||||
@@ -12,3 +12,5 @@ Our renderer API is designed to disaggregate the render phase(preprocessing) and
|
||||
- Render completion requests
|
||||
- [Chat Completions Render API](renderer.md) (`/v1/chat/completions/render`)
|
||||
- Render chat completions
|
||||
|
||||
For the post processing counterpart that turns generated token IDs back into OpenAI compatible responses, see the [Derenderer APIs](derenderer.md).
|
||||
|
||||
@@ -67,7 +67,7 @@ The Transcriptions API supports uploading audio files in various formats includi
|
||||
- `response_format`: Format of the response ("json", "text") (optional)
|
||||
- `temperature`: Sampling temperature between 0 and 1 (optional)
|
||||
|
||||
For the complete list of supported parameters including sampling parameters and vLLM extensions, see the [protocol definitions](https://github.com/vllm-project/vllm/blob/main/vllm/entrypoints/openai/protocol.py#L2182).
|
||||
For the complete list of supported parameters including sampling parameters and vLLM extensions, see the [protocol definitions](https://github.com/vllm-project/vllm/blob/main/vllm/entrypoints/speech_to_text/transcription/protocol.py).
|
||||
|
||||
**Response Format:**
|
||||
|
||||
|
||||
@@ -155,8 +155,10 @@ When `--api-key` is configured, the following `/v1` endpoints require Bearer tok
|
||||
- `/v1/chat/completions` - Chat completions
|
||||
- `/v1/chat/completions/batch` - Batch chat completions
|
||||
- `/v1/chat/completions/render` - Render chat completion requests
|
||||
- `/v1/chat/completions/derender` - Derender chat completion requests
|
||||
- `/v1/completions` - Text completions
|
||||
- `/v1/completions/render` - Render completion requests
|
||||
- `/v1/completions/derender` - Derender completion requests
|
||||
- `/v1/embeddings` - Generate embeddings
|
||||
- `/v1/audio/transcriptions` - Audio transcription
|
||||
- `/v1/audio/translations` - Audio translation
|
||||
|
||||
@@ -210,8 +210,31 @@ async def stream_decode_response(session, response, request_id):
|
||||
await session.close()
|
||||
|
||||
|
||||
def example_round_robin_dp_loader(request_number, dp_size):
|
||||
return request_nums % dp_size
|
||||
def flat_interleaved_dp_route(request_number, instances):
|
||||
"""Flat round-robin over the full (instance, dp_rank) slot space.
|
||||
|
||||
ONE counter over (n_instances * dp_size) slots, so instance-selection and
|
||||
DP-rank-selection are derived from the SAME index and can never alias. The
|
||||
previous scheme computed instance = req % n and rank = req % dp from the
|
||||
same counter with n | dp, which locked each instance to a stride-n subset
|
||||
of its ranks (e.g. 2 prefill instances -> 4 of 8 ranks each -> half the
|
||||
GPUs never receive a request, so the deployment falsely appears not to
|
||||
scale).
|
||||
|
||||
Interleaved order — inst0_r0, inst1_r0, inst0_r1, inst1_r1, ... — so
|
||||
consecutive requests alternate instances AND every rank gets walked.
|
||||
|
||||
Assumes homogeneous dp_size across a role's instances (true for the
|
||||
DP<->DP and DP<->TP deployments this proxy targets). Returns
|
||||
(instance_index, dp_rank); dp_rank is None when dp_size == 1 (e.g. a TP
|
||||
decode), which avoids forwarding an out-of-range data-parallel rank.
|
||||
"""
|
||||
n = len(instances)
|
||||
dp = instances[0]["dp_size"]
|
||||
slot = (request_number - 1) % (n * dp)
|
||||
inst_idx = slot % n
|
||||
dp_rank = (slot // n) if dp > 1 else None
|
||||
return inst_idx, dp_rank
|
||||
|
||||
|
||||
@app.route("/health", methods=["GET"])
|
||||
@@ -252,18 +275,21 @@ async def handle_request(api: str, request: Request):
|
||||
503,
|
||||
)
|
||||
)
|
||||
pid = request_nums % len(prefill_instances)
|
||||
did = request_nums % len(decode_instances)
|
||||
# Flat interleaved round-robin (see flat_interleaved_dp_route): ONE
|
||||
# counter over the full (instance, dp_rank) slot space per role, so
|
||||
# instance-selection and DP-rank-selection derive from the same index
|
||||
# and can never alias. The old scheme keyed both on request_nums with
|
||||
# n_instances | dp_size, stranding half the ranks (e.g. in 2P_DP8EP).
|
||||
pid, selected_prefill_dp_rank = flat_interleaved_dp_route(
|
||||
request_nums, prefill_instances
|
||||
)
|
||||
# Decode instance selection uses the same interleaved walk; in READ
|
||||
# mode the decode reads KV from selected_prefill_dp_rank, so the
|
||||
# decode's own dp_rank is not forwarded here.
|
||||
did, _ = flat_interleaved_dp_route(request_nums, decode_instances)
|
||||
prefill_instance_endpoint = prefill_instances[pid]
|
||||
decode_instance_endpoint = decode_instances[did]
|
||||
|
||||
selected_prefill_dp_rank = None
|
||||
if prefill_instance_endpoint["dp_size"] > 1:
|
||||
selected_prefill_dp_rank = example_round_robin_dp_loader(
|
||||
request_nums // len(prefill_instance_endpoint),
|
||||
prefill_instance_endpoint["dp_size"],
|
||||
)
|
||||
|
||||
# Embed both zmq_addresses in the request_id so the connector can parse
|
||||
# the peer's host/ports from it, similar to P2P-NCCL
|
||||
uid = str(uuid.uuid4()).replace("-", "")
|
||||
@@ -427,9 +453,33 @@ if __name__ == "__main__":
|
||||
args = parser.parse_args()
|
||||
|
||||
t = start_service_discovery("0.0.0.0", 36367)
|
||||
app.debug = True
|
||||
# High-concurrency hardening. Quart's app.run() uses a shallow listen
|
||||
# backlog (100) and, with app.debug=True, adds per-request overhead that
|
||||
# starves the single accept loop. Under a burst of ~512 simultaneous client
|
||||
# connections the backlog overflows and the kernel RSTs the excess, so
|
||||
# clients see "ClientOSError: [Errno 104] Connection reset by peer" before
|
||||
# any response (~16% request loss at c=512). Serve via hypercorn with debug
|
||||
# OFF and a deep backlog so the burst QUEUES (higher TTFT) instead of being
|
||||
# reset -> 100% request success.
|
||||
app.debug = False
|
||||
app.config["BODY_TIMEOUT"] = 360000
|
||||
app.config["RESPONSE_TIMEOUT"] = 360000
|
||||
|
||||
app.run(host="0.0.0.0", port=args.port)
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
from hypercorn.asyncio import serve as _hypercorn_serve
|
||||
from hypercorn.config import Config as _HypercornConfig
|
||||
|
||||
_hcfg = _HypercornConfig()
|
||||
_hcfg.bind = [f"0.0.0.0:{args.port}"]
|
||||
# Deep listen backlog so a wide connection burst queues, not RSTs. NOTE:
|
||||
# effective backlog is capped by the host's net.core.somaxconn (proxy runs
|
||||
# --network host); kernel 6.x defaults to 4096. Override via
|
||||
# PROXY_LISTEN_BACKLOG.
|
||||
_hcfg.backlog = int(os.environ.get("PROXY_LISTEN_BACKLOG", "4096"))
|
||||
# Long-lived SSE streams (8k1k decode ~5 min): never reap on keepalive.
|
||||
_hcfg.keep_alive_timeout = 360000.0
|
||||
|
||||
asyncio.run(_hypercorn_serve(app, _hcfg))
|
||||
t.join()
|
||||
|
||||
@@ -42,12 +42,16 @@ class BlockStored(KVCacheEvent):
|
||||
"""
|
||||
|
||||
group_idx: int | None = None
|
||||
kv_cache_spec_kind: str | None = None
|
||||
kv_cache_spec_sliding_window: int | None = None
|
||||
locality: str | None = None
|
||||
|
||||
|
||||
class BlockRemoved(KVCacheEvent):
|
||||
block_hashes: list[ExternalBlockHash]
|
||||
medium: str | None
|
||||
group_idx: int | None = None
|
||||
locality: str | None = None
|
||||
|
||||
|
||||
class AllBlocksCleared(KVCacheEvent):
|
||||
|
||||
@@ -17,6 +17,7 @@ from transformers import AutoProcessor, AutoTokenizer
|
||||
from vllm import LLM, EngineArgs, SamplingParams
|
||||
from vllm.lora.request import LoRARequest
|
||||
from vllm.multimodal.utils import fetch_image
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.argparse_utils import FlexibleArgumentParser
|
||||
|
||||
QUESTION = "What is the content of each image?"
|
||||
@@ -1443,6 +1444,8 @@ def run_generate(
|
||||
engine_args.seed = seed
|
||||
if tensor_parallel_size is not None:
|
||||
engine_args.tensor_parallel_size = tensor_parallel_size
|
||||
if current_platform.is_rocm():
|
||||
os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn"
|
||||
llm = LLM.from_engine_args(engine_args)
|
||||
|
||||
sampling_params = SamplingParams(
|
||||
@@ -1484,6 +1487,8 @@ def run_chat(
|
||||
engine_args.seed = seed
|
||||
if tensor_parallel_size is not None:
|
||||
engine_args.tensor_parallel_size = tensor_parallel_size
|
||||
if current_platform.is_rocm():
|
||||
os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn"
|
||||
llm = LLM.from_engine_args(engine_args)
|
||||
|
||||
sampling_params = (
|
||||
|
||||
@@ -21,6 +21,7 @@ from vllm.assets.image import ImageAsset
|
||||
from vllm.assets.video import VideoAsset
|
||||
from vllm.lora.request import LoRARequest
|
||||
from vllm.multimodal.image import convert_image_mode
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.argparse_utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
@@ -2646,6 +2647,8 @@ def main(args):
|
||||
if args.tensor_parallel_size is not None:
|
||||
engine_args.tensor_parallel_size = args.tensor_parallel_size
|
||||
engine_args = maybe_add_vit_cuda_graph_compilation_config(args, engine_args)
|
||||
if current_platform.is_rocm():
|
||||
os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn"
|
||||
llm = LLM.from_engine_args(engine_args)
|
||||
|
||||
# Don't want to check the flag multiple times, so just hijack `prompts`.
|
||||
|
||||
+2
-1
@@ -122,7 +122,8 @@ python = "./.venv"
|
||||
[tool.typos.files]
|
||||
# these files may be written in non english words
|
||||
extend-exclude = ["tests/models/fixtures/*", "tests/prompts/*", "tests/tokenizers_/*",
|
||||
"benchmarks/sonnet.txt", "tests/lora/data/*", "build/*",
|
||||
"benchmarks/sonnet.txt", "rust/src/bench/src/datasets/sonnet.txt",
|
||||
"tests/lora/data/*", "build/*",
|
||||
"examples/pooling/token_embed/*", "tests/models/language/pooling/*",
|
||||
"vllm/third_party/*", "vllm/entrypoints/serve/instrumentator/static/*",
|
||||
"tests/entrypoints/speech_to_text/transcription/test_transcription_validation.py",
|
||||
|
||||
@@ -379,7 +379,7 @@ inflect==5.6.2
|
||||
# via datamodel-code-generator
|
||||
iniconfig==2.0.0
|
||||
# via pytest
|
||||
instanttensor==0.1.5
|
||||
instanttensor==0.1.9
|
||||
# via -r requirements/test/cuda.in
|
||||
interegular==0.3.3
|
||||
# via lm-format-enforcer
|
||||
|
||||
@@ -58,7 +58,7 @@ arctic-inference == 0.1.1; platform_machine == "x86_64" # Required for suffix de
|
||||
numba == 0.65.0 # Required for N-gram speculative decoding
|
||||
runai-model-streamer[s3,gcs,azure]==0.15.7
|
||||
fastsafetensors>=0.3.2
|
||||
instanttensor>=0.1.5; platform_machine == "x86_64"
|
||||
instanttensor>=0.1.9; platform_machine == "x86_64"
|
||||
decord==0.6.0; platform_machine == "x86_64"
|
||||
# terratorch is temporarily disabled while PyPI has the `lightning` package
|
||||
# in `quarantined` status (every published terratorch version transitively
|
||||
|
||||
@@ -398,7 +398,7 @@ inflect==5.6.2
|
||||
# via datamodel-code-generator
|
||||
iniconfig==2.0.0
|
||||
# via pytest
|
||||
instanttensor==0.1.5
|
||||
instanttensor==0.1.9
|
||||
# via -r requirements/test/cuda.in
|
||||
interegular==0.3.3
|
||||
# via lm-format-enforcer
|
||||
|
||||
@@ -44,5 +44,5 @@ numba == 0.65.0 # Required for N-gram speculative decoding
|
||||
numpy
|
||||
runai-model-streamer[s3,gcs,azure]==0.15.7
|
||||
fastsafetensors>=0.3.2
|
||||
instanttensor>=0.1.5
|
||||
instanttensor>=0.1.9
|
||||
pydantic>=2.12 # 2.11 leads to error on python 3.13
|
||||
|
||||
@@ -54,7 +54,7 @@ arctic-inference==0.1.1 # Required for suffix decoding test
|
||||
numba==0.65.0 # Required for N-gram speculative decoding
|
||||
runai-model-streamer[s3,gcs,azure]==0.15.7
|
||||
fastsafetensors>=0.3.2
|
||||
instanttensor>=0.1.5
|
||||
instanttensor>=0.1.9
|
||||
decord==0.6.0
|
||||
|
||||
# Prithvi tests
|
||||
|
||||
@@ -391,7 +391,7 @@ inflect==7.5.0
|
||||
# via datamodel-code-generator
|
||||
iniconfig==2.3.0
|
||||
# via pytest
|
||||
instanttensor==0.1.6
|
||||
instanttensor==0.1.9
|
||||
# via -r requirements/test/rocm.in
|
||||
interegular==0.3.3
|
||||
# via lm-format-enforcer
|
||||
|
||||
@@ -18,4 +18,4 @@ torchvision
|
||||
torchcodec >= 0.14 # Required for the torchcodec video decoding backend
|
||||
|
||||
auto_round_lib==0.14.1
|
||||
vllm_xpu_kernels @ https://github.com/vllm-project/vllm-xpu-kernels/releases/download/v0.1.11/vllm_xpu_kernels-0.1.11-cp38-abi3-manylinux_2_28_x86_64.whl
|
||||
vllm_xpu_kernels @ https://github.com/vllm-project/vllm-xpu-kernels/releases/download/v0.1.11.1/vllm_xpu_kernels-0.1.11.1-cp38-abi3-manylinux_2_28_x86_64.whl
|
||||
|
||||
Generated
+82
-11
@@ -2585,6 +2585,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841"
|
||||
dependencies = [
|
||||
"autocfg",
|
||||
"libm",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -3305,6 +3306,16 @@ dependencies = [
|
||||
"getrandom 0.3.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rand_distr"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6a8615d50dcf34fa31f7ab52692afec947c4dd0ab803cc87cb3b0b4570ff7463"
|
||||
dependencies = [
|
||||
"num-traits",
|
||||
"rand 0.9.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rawpointer"
|
||||
version = "0.2.1"
|
||||
@@ -3549,6 +3560,15 @@ dependencies = [
|
||||
"rustc-hash 2.1.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rlimit"
|
||||
version = "0.11.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f35ee2729c56bb610f6dba436bf78135f728b7373bdffae2ec815b2d3eb98cc3"
|
||||
dependencies = [
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rmp"
|
||||
version = "0.8.15"
|
||||
@@ -4855,6 +4875,7 @@ dependencies = [
|
||||
"futures-core",
|
||||
"pin-project-lite",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -4916,9 +4937,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "tonic"
|
||||
version = "0.14.5"
|
||||
version = "0.14.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fec7c61a0695dc1887c1b53952990f3ad2e3a31453e1f49f10e75424943a93ec"
|
||||
checksum = "ac2a5518c70fa84342385732db33fb3f44bc4cc748936eb5833d2df34d6445ef"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"axum",
|
||||
@@ -4945,9 +4966,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "tonic-build"
|
||||
version = "0.14.5"
|
||||
version = "0.14.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1882ac3bf5ef12877d7ed57aad87e75154c11931c2ba7e6cde5e22d63522c734"
|
||||
checksum = "c68f61875ac5293cf72e6c8cf0158086428c82c37229e98c840878f1706b0322"
|
||||
dependencies = [
|
||||
"prettyplease",
|
||||
"proc-macro2",
|
||||
@@ -4956,10 +4977,23 @@ dependencies = [
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tonic-prost"
|
||||
version = "0.14.5"
|
||||
name = "tonic-health"
|
||||
version = "0.14.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a55376a0bbaa4975a3f10d009ad763d8f4108f067c7c2e74f3001fb49778d309"
|
||||
checksum = "fcfab99db777fba2802f0dfa861d1628d1ae916fb199d29819941f139ae85082"
|
||||
dependencies = [
|
||||
"prost",
|
||||
"tokio",
|
||||
"tokio-stream",
|
||||
"tonic",
|
||||
"tonic-prost",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tonic-prost"
|
||||
version = "0.14.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "50849f68853be452acf590cde0b146665b8d507b3b8af17261df47e02c209ea0"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"prost",
|
||||
@@ -4968,9 +5002,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "tonic-prost-build"
|
||||
version = "0.14.5"
|
||||
version = "0.14.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f3144df636917574672e93d0f56d7edec49f90305749c668df5101751bb8f95a"
|
||||
checksum = "654e5643eff75d7f8c99197ce1440ed19a3474eada74c12bbac488b2cafdae27"
|
||||
dependencies = [
|
||||
"prettyplease",
|
||||
"proc-macro2",
|
||||
@@ -5439,6 +5473,41 @@ version = "0.9.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a"
|
||||
|
||||
[[package]]
|
||||
name = "vllm-bench"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"chrono",
|
||||
"clap",
|
||||
"dirs",
|
||||
"futures",
|
||||
"hf-hub",
|
||||
"image",
|
||||
"indicatif",
|
||||
"mimalloc",
|
||||
"rand 0.9.2",
|
||||
"rand_distr",
|
||||
"rayon",
|
||||
"reqwest 0.12.28",
|
||||
"rlimit",
|
||||
"rustc-hash 1.1.0",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror 2.0.18",
|
||||
"thiserror-ext",
|
||||
"tiktoken-rs 0.9.1",
|
||||
"tokenizers",
|
||||
"tokio",
|
||||
"tokio-stream",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"url",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "vllm-chat"
|
||||
version = "0.1.0"
|
||||
@@ -5507,6 +5576,7 @@ dependencies = [
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"uuid",
|
||||
"vllm-bench",
|
||||
"vllm-chat",
|
||||
"vllm-engine-core-client",
|
||||
"vllm-managed-engine",
|
||||
@@ -5676,6 +5746,7 @@ dependencies = [
|
||||
"tokio-stream",
|
||||
"tokio-util",
|
||||
"tonic",
|
||||
"tonic-health",
|
||||
"tonic-prost",
|
||||
"tonic-prost-build",
|
||||
"tower",
|
||||
@@ -6267,9 +6338,9 @@ checksum = "9edde0db4769d2dc68579893f2306b26c6ecfbe0ef499b013d731b7b9247e0b9"
|
||||
|
||||
[[package]]
|
||||
name = "xgrammar-structural-tag"
|
||||
version = "0.1.0+xgrammar.0.2.2.4d145cc"
|
||||
version = "0.2.0+xgrammar.0.2.4.dd729e7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2436dea2393d55a3b188588aa300c5a8afe8f45a77da52c611fb4498a6c876e6"
|
||||
checksum = "d4d24c842efc3c24e9756aa426d530cbdac0980e49af223cb384e276e981ca0a"
|
||||
dependencies = [
|
||||
"auto_impl",
|
||||
"serde",
|
||||
|
||||
+16
-5
@@ -1,5 +1,6 @@
|
||||
[workspace]
|
||||
members = [
|
||||
"src/bench",
|
||||
"src/chat",
|
||||
"src/cmd",
|
||||
"src/engine-core-client",
|
||||
@@ -32,8 +33,10 @@ base64 = "0.22.1"
|
||||
bytemuck = { version = "1.25.0", features = ["extern_crate_alloc"] }
|
||||
byteorder = "1.5.0"
|
||||
bytes = "1.12.0"
|
||||
chrono = "0.4.42"
|
||||
clap = { version = "4.5.38", features = ["derive", "env"] }
|
||||
criterion = "0.5.1"
|
||||
dirs = "6.0.0"
|
||||
easy-ext = "1.0.3"
|
||||
educe = "0.6.0"
|
||||
enum-as-inner = "0.7.0"
|
||||
@@ -50,7 +53,9 @@ hyper-util = { version = "0.1.20", features = [
|
||||
"service",
|
||||
"tokio",
|
||||
] }
|
||||
image = { version = "0.25.9", default-features = false, features = ["jpeg"] }
|
||||
indexmap = "2.13.0"
|
||||
indicatif = "0.18.4"
|
||||
itertools = "0.14.0"
|
||||
libc = "0.2.177"
|
||||
llm-multimodal = { git = "https://github.com/smg-project/llm-multimodal", rev = "5390032d6dc8a3e6fdc83acd320260367eb4b9b5", default-features = false, features = ["native-tls"] }
|
||||
@@ -71,10 +76,13 @@ prost-types = "0.14.3"
|
||||
pyo3 = "0.28.3"
|
||||
pythonize = "0.28.0"
|
||||
rand = "0.9.2"
|
||||
rand_distr = "0.5.1"
|
||||
rayon = "1.11.0"
|
||||
reasoning-parser = "1.2.2"
|
||||
reqwest = { version = "0.12.8", default-features = false, features = ["native-tls"] }
|
||||
reqwest-0-13 = { package = "reqwest", version = "0.13.4", default-features = false, features = ["native-tls"] }
|
||||
riptoken = { version = "0.3.0", default-features = false }
|
||||
rlimit = "0.11.0"
|
||||
rmp-serde = "1.3.1"
|
||||
rmpv = { version = "1.3.1", features = ["with-serde"] }
|
||||
rustc-hash = "1.1.0"
|
||||
@@ -110,10 +118,11 @@ tokio = { version = "1.47.1", features = [
|
||||
tokio-openssl = "0.6"
|
||||
tokio-stream = "0.1"
|
||||
tokio-util = { version = "0.7.18", features = ["rt"] }
|
||||
tonic = "0.14.5"
|
||||
tonic-build = "0.14.5"
|
||||
tonic-prost = "0.14.5"
|
||||
tonic-prost-build = "0.14.5"
|
||||
tonic = "0.14.6"
|
||||
tonic-build = "0.14.6"
|
||||
tonic-health = "0.14.6"
|
||||
tonic-prost = "0.14.6"
|
||||
tonic-prost-build = "0.14.6"
|
||||
tool-parser = "1.2.0"
|
||||
tower = { version = "0.5.3", features = ["util"] }
|
||||
tower-http = { version = "0.6.8", features = ["cors", "trace"] }
|
||||
@@ -121,8 +130,10 @@ tracing = { version = "0.1.44", features = ["release_max_level_debug"] }
|
||||
tracing-futures = { version = "0.2.5", features = ["futures-03"] }
|
||||
tracing-subscriber = { version = "0.3.20", features = ["env-filter", "fmt"] }
|
||||
trait-set = "0.3.0"
|
||||
url = "2.5.7"
|
||||
uuid = { version = "1.22.0", features = ["v4"] }
|
||||
validator = { version = "0.20.0", features = ["derive"] }
|
||||
vllm-bench = { path = "src/bench" }
|
||||
vllm-chat = { path = "src/chat" }
|
||||
vllm-engine-core-client = { path = "src/engine-core-client" }
|
||||
vllm-llm = { path = "src/llm" }
|
||||
@@ -133,7 +144,7 @@ vllm-server = { path = "src/server" }
|
||||
vllm-text = { path = "src/text" }
|
||||
vllm-tokenizer = { path = "src/tokenizer" }
|
||||
winnow = { version = "1.0.2", features = ["simd"] }
|
||||
xgrammar-structural-tag = "0.1.0"
|
||||
xgrammar-structural-tag = "0.2.0"
|
||||
zeromq = { version = "0.6.0", default-features = false, features = [
|
||||
"tokio-runtime",
|
||||
"all-transport",
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
# AGENTS.md
|
||||
|
||||
## Project Overview
|
||||
|
||||
Rust rewrite of `vllm bench serve` — a high-performance benchmark client for vLLM serving endpoints. Standalone binary, no Python dependency at runtime.
|
||||
|
||||
Member crate `vllm-bench` of the `rust/` workspace. Uses workspace dependencies and lints; the workspace `[profile.release]` (thin LTO, `panic = "abort"`) applies. Note the workspace bans rustls/ring (`rust/deny.toml`) — all HTTP must stay on native-tls, which is why HF Hub downloads go through `src/hub.rs` (async hf-hub API bridged to sync) instead of hf-hub's ureq backend.
|
||||
|
||||
## Build & Test
|
||||
|
||||
Run from the `rust/` workspace root:
|
||||
|
||||
```bash
|
||||
# Build release binary (rust/target/release/vllm-bench)
|
||||
cargo build -p vllm-bench --release
|
||||
|
||||
# Run all tests
|
||||
cargo test -p vllm-bench
|
||||
|
||||
# Run ignored integration tests (requires network for tokenizer download)
|
||||
cargo test -p vllm-bench -- --ignored
|
||||
```
|
||||
|
||||
## Architecture
|
||||
|
||||
- `src/main.rs` — Entry point, mimalloc, tokio runtime, mode dispatch (compare/sweep/multi-run/multi-turn/single)
|
||||
- `src/cli.rs` — clap derive CLI args (~50+ flags)
|
||||
- `src/config.rs` — Validated config from CLI; `GoodputConfig`, `RampUpConfig`, sampling param merging
|
||||
- `src/error.rs` — `BenchError` enum (Http, Json, Tokenizer, Config, EndpointTimeout, Backend, Io)
|
||||
- `src/benchmark.rs` — Core benchmark orchestrator (spawn-per-request with tokio + Semaphore; fetches speculative decoding metrics from `/metrics`)
|
||||
- `src/multi_turn.rs` — Multi-turn conversation orchestrator (channel-based worker pool, sequential turns per conversation)
|
||||
- `src/sweep.rs` — Concurrency/rate parameter sweep (`--sweep-max-concurrency`, `--sweep-request-rate`)
|
||||
- `src/multi_run.rs` — N-run aggregation with mean/std/min/max/CV (`--num-runs`)
|
||||
- `src/compare.rs` — Side-by-side diff of two result JSON files (`--compare`)
|
||||
- `src/tokenizer.rs` — `TokenizerKind` enum: Local(HuggingFace), Tiktoken, OR Server-side `/tokenize`+`/detokenize` fallback
|
||||
- `src/tiktoken.rs` — Tiktoken BPE loader (`.tiktoken`/`.model` files; built-in encodings o200k_base/cl100k_base; pat_str extraction from Python source)
|
||||
- `src/hub.rs` — `HubRepo`: sync facade over hf-hub's async (reqwest/native-tls) API — per-download thread with its own runtime; the sync ureq backend is unusable here because it pulls rustls, which `rust/deny.toml` bans
|
||||
- `src/rate_control.rs` — Gamma/Poisson request scheduling + linear/exponential ramp-up
|
||||
- `src/ready_checker.rs` — Endpoint readiness with retry
|
||||
- `src/backends/` — Backend implementations (enum dispatch, not trait objects)
|
||||
- `mod.rs` — `Backend` enum, `RequestFuncInput`/`RequestFuncOutput` (includes `messages` field for multi-turn)
|
||||
- `streaming.rs` — SSE parser (`StreamedResponseHandler`) with speculative JSON parse for split TCP segments
|
||||
- `openai_completions.rs` — `/v1/completions` backend
|
||||
- `openai_chat.rs` — `/v1/chat/completions` backend (uses `input.messages` when set; zero-copy raw JSON payload for multimodal)
|
||||
- `pooling.rs` — Non-streaming pooling/embedding backends: `openai-embeddings`, `openai-embeddings-chat`, `vllm-pooling`, `vllm-rerank`
|
||||
- `src/datasets/random.rs` — Random dataset generation with rayon parallelism
|
||||
- `src/datasets/random_mm.rs` — Random multimodal dataset (synthetic JPEG images, bucket config sampling, pre-serialized JSON fragments); `--enable-multimodal-chat` pre-builds the chat `messages` array at dataset time (mirrors Python's `apply_multimodal_chat_transformation`)
|
||||
- `src/datasets/sharegpt.rs` — ShareGPT JSON loader + HuggingFace Hub auto-download with caching
|
||||
- `src/datasets/sonnet.rs` and `src/datasets/sonnet.txt` — Sonnet dataset (built-in Shakespeare sonnets via `include_str!("sonnet.txt")`; controllable token length + shared prefix; mirrors Python `SonnetDataset`)
|
||||
- `src/datasets/speed_bench.rs` — NVIDIA SPEED-Bench loader (HF datasets-server API, 6 configs, 11 categories, local cache)
|
||||
- `src/datasets/hf_dataset.rs` — Generic HuggingFace dataset loader (datasets-server API, column auto-detection)
|
||||
- `src/datasets/custom.rs` — Custom JSONL dataset (`{"prompt": ..., "output_tokens": ...}` per line; `--custom-output-len -1` uses per-line output_tokens; prompts always sent raw — no client-side chat template)
|
||||
- `src/datasets/prefix_repetition.rs` — Prefix repetition dataset (N shared prefixes × fresh random suffixes, standard prefix-cache stress; mirrors Python `PrefixRepetitionRandomDataset`)
|
||||
- `src/datasets/random_rerank.rs` — Random rerank dataset (one query + batched documents per request for `vllm-rerank`; `--no-reranker` for embedding-based scoring; mirrors Python `RandomDatasetForReranking`)
|
||||
- `src/datasets/multi_turn.rs` — Multi-turn synthetic generator + ShareGPT multi-turn loader (3-tier prefix sharing: global/conversation/unique-suffix; `per_turn_input_len`)
|
||||
- `src/metrics/mod.rs` — `BenchmarkMetrics` and `MultiTurnMetrics` structs
|
||||
- `src/metrics/calculator.rs` — TTFT/TPOT/ITL/E2EL/throughput stats, goodput SLO checking, peak concurrency, `calculate_multi_turn_metrics`
|
||||
- `src/metrics/steady_state.rs` — Steady-state window detection (in-flight concurrency plateau via two-pointer start/end merge) + plateau throughput/TTFT/TPOT; gated on `--max-concurrency` set + `--request-rate inf` (closed-loop)
|
||||
- `src/output/console.rs` — Terminal output matching Python format + multi-turn per-turn breakdown
|
||||
- `src/output/json.rs` — JSON result file (compatible with Python schema) + multi-turn JSON with `per_turn_metrics`
|
||||
|
||||
## Key Design Decisions
|
||||
|
||||
- **Enum dispatch** for backends (avoids async trait object issues with `dyn`)
|
||||
- **reqwest http1_only()** to match Python aiohttp behavior
|
||||
- **rayon** for parallel dataset generation (key perf win over Python)
|
||||
- **mimalloc** global allocator to reduce contention at 1400+ concurrency (page-agnostic; works on aarch64 64K-page kernels where jemalloc aborts with `LG_PAGE=12` builds)
|
||||
- **Arc\<str\> prompts** zero-copy sharing across tokio tasks (~3GB savings at 100k prompts with 8k-token inputs)
|
||||
- **Spawn-per-request** `tokio::spawn` + `Semaphore` (matches Python asyncio pattern)
|
||||
- **Speculative JSON parse** in SSE handler — detects complete JSON before `\n\n` arrives, improving TTFT/ITL accuracy when TCP segments split
|
||||
- **Tokenizer fallback chain**: Local HF → Tiktoken (`.tiktoken`/`.model` + built-in encodings) → Server-side `/tokenize`+`/detokenize`. Blocking HTTP in rayon threads for server fallback.
|
||||
- **hf-hub** for downloading tokenizers and datasets from HuggingFace Hub
|
||||
- **Pre-serialized mm fragments** (`Arc<str>`) for multimodal: image content stored as JSON strings, zero-copy concatenated into payload — avoids deep-cloning ~200KB+ base64 per request
|
||||
- **Steady-state metrics** (default-on in closed-loop): measure throughput/TTFT/TPOT only over the saturated plateau to cut run-to-run variance at high concurrency; `steady_state` is an `Option` in JSON (`#[serde(default)]` for backward compat), null when the scope gate fails or `--no-steady-state`
|
||||
- **`--prompt-token-ids`** (random dataset only): send token-ID arrays instead of text to skip server-side tokenization; also skips the token-length verification pass (counts exact by construction)
|
||||
- **`--random-range-ratio`** follows Python semantics: lengths sampled uniformly from `[len*(1-r), len*(1+r)]`, default `0.0` = fixed; accepts a float in `[0,1)` or `'{"input": r1, "output": r2}'`. (The pre-2026-07 Rust-only form `[len*r, len]` with default 1.0 is rejected with a migration hint.)
|
||||
- **`prompt_list`** (`Arc<[Arc<str>]>` on `SampleRequest`/`RequestFuncInput`): multiple inputs per request for pooling backends — embeddings batches (`--random-batch-size`) send `"input": [...]`, rerank sends `[0]` as query + `[1..]` as documents
|
||||
- JSON output schema must match Python `vllm bench serve` exactly
|
||||
|
||||
## Common Issues
|
||||
|
||||
- **localhost vs 127.0.0.1**: Some systems resolve `localhost` to IPv6 `::1` while vLLM listens on IPv4 only. Use `127.0.0.1` or the actual hostname.
|
||||
- **Models without tokenizer.json** (e.g., `nvidia/Kimi-K2.5-NVFP4`): Automatically falls back to server-side tokenization. Can also use `--tokenizer` to point to a model with `tokenizer.json`.
|
||||
- **usage.completion_tokens parsing**: vLLM sends final usage chunk with `"choices":[]` (empty array). The usage `if` must be separate from the choices `if` (not `else if`).
|
||||
|
||||
## Typical Usage
|
||||
|
||||
```bash
|
||||
# Embedding benchmark (openai-embeddings, 8 inputs batched per request)
|
||||
./target/release/vllm-bench \
|
||||
--backend openai-embeddings \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model BAAI/bge-large-en-v1.5 \
|
||||
--dataset-name random \
|
||||
--random-input-len 512 \
|
||||
--random-batch-size 8 \
|
||||
--num-prompts 1000 \
|
||||
--save-result
|
||||
|
||||
# vLLM rerank benchmark (one query + 8 documents per request)
|
||||
./target/release/vllm-bench \
|
||||
--backend vllm-rerank \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model BAAI/bge-reranker-v2-m3 \
|
||||
--dataset-name random-rerank \
|
||||
--random-input-len 512 \
|
||||
--random-batch-size 8 \
|
||||
--num-prompts 500 \
|
||||
--save-result
|
||||
|
||||
# Prefix-cache stress (10 shared prefixes, 256+256 tokens)
|
||||
./target/release/vllm-bench \
|
||||
--backend vllm \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model nvidia/Kimi-K2.5-NVFP4 \
|
||||
--dataset-name prefix_repetition \
|
||||
--prefix-repetition-prefix-len 256 \
|
||||
--prefix-repetition-suffix-len 256 \
|
||||
--prefix-repetition-num-prefixes 10 \
|
||||
--num-prompts 1000
|
||||
|
||||
# Custom JSONL workload ({"prompt": ..., "output_tokens": ...} per line)
|
||||
./target/release/vllm-bench \
|
||||
--backend openai-chat \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model nvidia/Kimi-K2.5-NVFP4 \
|
||||
--dataset-name custom \
|
||||
--dataset-path workload.jsonl \
|
||||
--custom-output-len -1 \
|
||||
--num-prompts 1000
|
||||
|
||||
# Random dataset
|
||||
./target/release/vllm-bench \
|
||||
--backend vllm \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model nvidia/Kimi-K2.5-NVFP4 \
|
||||
--dataset-name random \
|
||||
--random-input-len 8192 \
|
||||
--random-output-len 1024 \
|
||||
--ignore-eos \
|
||||
--num-prompts 4096 \
|
||||
--percentile-metrics "ttft,tpot,itl,e2el" \
|
||||
--save-result \
|
||||
--max-concurrency 1400
|
||||
|
||||
# Random multimodal dataset (VLM benchmark)
|
||||
./target/release/vllm-bench \
|
||||
--backend openai-chat \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model Qwen/Qwen2.5-VL-7B-Instruct \
|
||||
--dataset-name random-mm \
|
||||
--random-input-len 512 \
|
||||
--random-output-len 128 \
|
||||
--num-prompts 100 \
|
||||
--random-mm-base-items-per-request 1 \
|
||||
--random-mm-limit-mm-per-prompt '{"image": 1, "video": 0}' \
|
||||
--random-mm-bucket-config '{(1024, 800, 1): 1.0}'
|
||||
|
||||
# HuggingFace dataset (WildChat)
|
||||
./target/release/vllm-bench \
|
||||
--backend openai-chat \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model nvidia/Kimi-K2.5-NVFP4 \
|
||||
--dataset-name hf \
|
||||
--dataset-path allenai/WildChat-4.8M \
|
||||
--hf-split train \
|
||||
--num-prompts 1000 \
|
||||
--save-result
|
||||
|
||||
# HuggingFace dataset (LongBench with subset)
|
||||
./target/release/vllm-bench \
|
||||
--backend openai-chat \
|
||||
--base-url http://gb200-10:30000 \
|
||||
--model nvidia/Kimi-K2.5-NVFP4 \
|
||||
--dataset-name hf \
|
||||
--dataset-path THUDM/LongBench \
|
||||
--hf-subset narrativeqa \
|
||||
--hf-split test \
|
||||
--hf-output-len 512 \
|
||||
--num-prompts 200
|
||||
```
|
||||
@@ -0,0 +1 @@
|
||||
@AGENTS.md
|
||||
@@ -0,0 +1,40 @@
|
||||
[package]
|
||||
name = "vllm-bench"
|
||||
version.workspace = true
|
||||
edition.workspace = true
|
||||
description = "High-performance benchmark client for vLLM serving endpoints"
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
anyhow.workspace = true
|
||||
base64.workspace = true
|
||||
bytes.workspace = true
|
||||
chrono.workspace = true
|
||||
clap.workspace = true
|
||||
dirs.workspace = true
|
||||
futures.workspace = true
|
||||
hf-hub.workspace = true
|
||||
image.workspace = true
|
||||
indicatif.workspace = true
|
||||
mimalloc.workspace = true
|
||||
rand.workspace = true
|
||||
rand_distr.workspace = true
|
||||
rayon.workspace = true
|
||||
reqwest = { workspace = true, features = ["json", "stream", "http2"] }
|
||||
rlimit.workspace = true
|
||||
rustc-hash.workspace = true
|
||||
serde = { workspace = true, features = ["rc"] }
|
||||
serde_json = { workspace = true, features = ["raw_value"] }
|
||||
thiserror.workspace = true
|
||||
thiserror-ext.workspace = true
|
||||
tiktoken-rs.workspace = true
|
||||
tokenizers.workspace = true
|
||||
tokio.workspace = true
|
||||
tokio-stream.workspace = true
|
||||
tracing.workspace = true
|
||||
tracing-subscriber.workspace = true
|
||||
url.workspace = true
|
||||
uuid.workspace = true
|
||||
|
||||
[lints]
|
||||
workspace = true
|
||||
@@ -0,0 +1,810 @@
|
||||
# vllm-bench
|
||||
|
||||
High-performance Rust benchmark client for vLLM serving endpoints. A drop-in replacement for `vllm bench serve` with near-instant startup, parallel dataset generation, and a fraction of the memory overhead — and no Python at runtime.
|
||||
|
||||
```bash
|
||||
vllm-bench --backend vllm --base-url http://127.0.0.1:8000 \
|
||||
--model <model> --dataset-name random \
|
||||
--random-input-len 1024 --random-output-len 128 \
|
||||
--num-prompts 1000 --max-concurrency 200
|
||||
```
|
||||
|
||||
## Highlights
|
||||
|
||||
- **Fast** — ~7 ms startup, single ~7 MB static binary, no Python imports.
|
||||
- **Scales** — `Arc<str>` prompt sharing + mimalloc keep memory <100 MB at 1400+ concurrency.
|
||||
- **Many datasets** — `random`, `random-mm` (VLM), `sharegpt`, `sonnet`, `speed-bench`, and any HuggingFace dataset.
|
||||
- **Many backends** — completions, chat, embeddings, pooling, and rerank.
|
||||
- **Beyond a single run** — concurrency/rate **sweeps**, **multi-run** stats, **multi-turn** conversations, **LoRA** multi-adapter, and result **comparison**.
|
||||
- **Steady-state metrics** — throughput/latency measured over the saturated plateau, excluding ramp-up and drain.
|
||||
- **Parity** — JSON output schema and timing semantics match Python `vllm bench serve` exactly.
|
||||
|
||||
### Performance vs. Python
|
||||
|
||||
| Metric | Python | Rust |
|
||||
| -------- | -------- | ------ |
|
||||
| Startup time | Multi-second (import vllm + numpy + aiohttp) | ~7 ms |
|
||||
| 100k random prompts (input_len=8192) | Minutes | Seconds (rayon parallelism) |
|
||||
| Binary size | — | ~7 MB |
|
||||
| Peak memory at 1400 concurrency | High (GIL + per-object overhead) | <100 MB (`Arc<str>` prompt sharing) |
|
||||
|
||||
## Contents
|
||||
|
||||
- [Install](#install)
|
||||
- [Quick Start](#quick-start)
|
||||
- [Usage Examples](#usage-examples)
|
||||
- [Supported Backends](#supported-backends)
|
||||
- [Supported Datasets](#supported-datasets)
|
||||
- [Metrics](#metrics)
|
||||
- [CLI Reference](#cli-reference)
|
||||
- [Tokenizer Support](#tokenizer-support)
|
||||
- [Output Format](#output-format)
|
||||
- [Architecture](#architecture)
|
||||
- [Environment Variables](#environment-variables)
|
||||
|
||||
## Install
|
||||
|
||||
### Prebuilt binaries (Linux)
|
||||
|
||||
```bash
|
||||
curl -fsSL https://github.com/vllm-project/vllm-bench/releases/latest/download/vllm-bench-$(uname -m)-linux-musl -o vllm-bench && chmod +x vllm-bench
|
||||
```
|
||||
|
||||
### With Cargo
|
||||
|
||||
Install straight from the repository (builds from source; requires [Rust](https://rustup.rs/) stable and a C compiler for the native tokenizer dependency):
|
||||
|
||||
```bash
|
||||
cargo install --git https://github.com/vllm-project/vllm-bench vllm-bench
|
||||
```
|
||||
|
||||
The trailing `vllm-bench` selects the package — the repo also ships a `mock-llm-server` binary, so omitting it fails with `multiple packages with binaries found`. The binary is installed to `~/.cargo/bin/`.
|
||||
|
||||
### Build from source
|
||||
|
||||
Requires [Rust](https://rustup.rs/) (stable).
|
||||
|
||||
```bash
|
||||
git clone https://github.com/vllm-project/vllm-bench.git
|
||||
cd vllm-bench
|
||||
./install.sh # builds release and installs to ~/.local/bin
|
||||
# or: ./install.sh --to ~/bin
|
||||
```
|
||||
|
||||
## Quick Start
|
||||
|
||||
Point it at a running vLLM server and benchmark with synthetic prompts:
|
||||
|
||||
```bash
|
||||
vllm-bench \
|
||||
--backend vllm \
|
||||
--base-url http://127.0.0.1:8000 \
|
||||
--model <model-name> \
|
||||
--dataset-name random \
|
||||
--random-input-len 1024 \
|
||||
--random-output-len 128 \
|
||||
--num-prompts 1000 \
|
||||
--max-concurrency 200
|
||||
```
|
||||
|
||||
> **Tip:** prefer `127.0.0.1` over `localhost` — some systems resolve `localhost` to IPv6 `::1` while vLLM listens on IPv4 only.
|
||||
|
||||
Add `--save-result` to write a JSON file, or `--dry-run` to generate and inspect the dataset without sending any requests.
|
||||
|
||||
## Usage Examples
|
||||
|
||||
<details open>
|
||||
<summary><b>Generation (completions / chat)</b></summary>
|
||||
|
||||
```bash
|
||||
# Full production-style run with percentile metrics and result file
|
||||
vllm-bench \
|
||||
--backend vllm \
|
||||
--base-url http://127.0.0.1:8000 \
|
||||
--model nvidia/Kimi-K2.5-NVFP4 \
|
||||
--dataset-name random \
|
||||
--random-input-len 8192 \
|
||||
--random-output-len 1024 \
|
||||
--ignore-eos \
|
||||
--num-prompts 4096 \
|
||||
--percentile-metrics "ttft,tpot,itl,e2el" \
|
||||
--save-result \
|
||||
--max-concurrency 1400
|
||||
|
||||
# Send token IDs instead of text (pure vLLM: skips server-side tokenization,
|
||||
# exact token counts, faster). Random dataset only.
|
||||
vllm-bench \
|
||||
--backend vllm \
|
||||
--base-url http://127.0.0.1:8000 \
|
||||
--model <model-name> \
|
||||
--dataset-name random \
|
||||
--random-input-len 1024 \
|
||||
--prompt-token-ids \
|
||||
--num-prompts 1000
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Datasets (ShareGPT / Sonnet / HuggingFace / SPEED-Bench)</b></summary>
|
||||
|
||||
```bash
|
||||
# ShareGPT (auto-downloads from HuggingFace on first run, cached afterwards)
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name sharegpt --num-prompts 500 --save-result
|
||||
|
||||
# ShareGPT with an explicit local file
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name sharegpt --dataset-path /path/to/ShareGPT_V3.json \
|
||||
--num-prompts 500 --save-result
|
||||
|
||||
# Sonnet — built-in Shakespeare sonnets, no dataset file needed.
|
||||
# Generates prompts of a controllable token length with a shared prefix.
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name sonnet \
|
||||
--sonnet-input-len 550 --sonnet-output-len 150 --sonnet-prefix-len 200 \
|
||||
--num-prompts 500
|
||||
|
||||
# Any public HuggingFace dataset (auto-downloads, auto-detects columns)
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name hf --dataset-path allenai/WildChat-4.8M \
|
||||
--hf-split train --num-prompts 1000 --save-result
|
||||
|
||||
# HuggingFace dataset with subset + fixed output length (LongBench)
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name hf --dataset-path THUDM/LongBench \
|
||||
--hf-subset narrativeqa --hf-split test --hf-output-len 512 --num-prompts 200
|
||||
|
||||
# Gated HuggingFace dataset (requires HF_TOKEN)
|
||||
HF_TOKEN=hf_xxx vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name hf --dataset-path lmsys/lmsys-chat-1m \
|
||||
--hf-split train --hf-output-len 256 --num-prompts 1000
|
||||
|
||||
# SPEED-Bench for speculative decoding evaluation (auto-downloads, cached)
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name speed-bench --speed-bench-config qualitative \
|
||||
--num-prompts 200 --output-len 256 --save-result
|
||||
|
||||
# SPEED-Bench throughput split with entropy category filter + input truncation
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name speed-bench --speed-bench-config throughput_16k \
|
||||
--speed-bench-max-input-len 10240 --speed-bench-category low_entropy \
|
||||
--num-prompts 500 --output-len 256 --max-concurrency 200 --save-result
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Multimodal (VLM with synthetic images)</b></summary>
|
||||
|
||||
```bash
|
||||
# One synthetic image per request
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 \
|
||||
--model Qwen/Qwen2.5-VL-7B-Instruct \
|
||||
--dataset-name random-mm \
|
||||
--random-input-len 512 --random-output-len 128 --num-prompts 100 \
|
||||
--random-mm-base-items-per-request 1 \
|
||||
--random-mm-limit-mm-per-prompt '{"image": 1, "video": 0}' \
|
||||
--random-mm-bucket-config '{(1024, 800, 1): 1.0}'
|
||||
|
||||
# Multiple images per request, mixed resolutions
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 \
|
||||
--model Qwen/Qwen2.5-VL-7B-Instruct \
|
||||
--dataset-name random-mm \
|
||||
--random-input-len 256 --random-output-len 128 --num-prompts 50 \
|
||||
--random-mm-base-items-per-request 3 \
|
||||
--random-mm-limit-mm-per-prompt '{"image": 5, "video": 0}' \
|
||||
--random-mm-bucket-config '{(256,256,1): 0.5, (720,1280,1): 0.5}'
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Embedding / Pooling / Rerank</b></summary>
|
||||
|
||||
```bash
|
||||
# Text embedding
|
||||
vllm-bench \
|
||||
--backend openai-embeddings --base-url http://127.0.0.1:8000 \
|
||||
--model BAAI/bge-large-en-v1.5 \
|
||||
--dataset-name random --random-input-len 512 --num-prompts 1000 \
|
||||
--max-concurrency 200 --save-result
|
||||
|
||||
# Chat-format embedding (supports multimodal content)
|
||||
vllm-bench \
|
||||
--backend openai-embeddings-chat --base-url http://127.0.0.1:8000 \
|
||||
--model BAAI/bge-large-en-v1.5 \
|
||||
--dataset-name sharegpt --num-prompts 500 --save-result
|
||||
|
||||
# vLLM native pooling endpoint
|
||||
vllm-bench \
|
||||
--backend vllm-pooling --base-url http://127.0.0.1:8000 \
|
||||
--model BAAI/bge-large-en-v1.5 \
|
||||
--dataset-name random --random-input-len 256 --num-prompts 1000 --save-result
|
||||
|
||||
# Rerank (query from dataset, documents via --extra-body)
|
||||
vllm-bench \
|
||||
--backend vllm-rerank --base-url http://127.0.0.1:8000 \
|
||||
--model BAAI/bge-reranker-v2-m3 \
|
||||
--dataset-name sharegpt --num-prompts 500 \
|
||||
--extra-body '{"documents": ["document to rerank"]}' --save-result
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Rate control, ramp-up & goodput</b></summary>
|
||||
|
||||
```bash
|
||||
# Ramp from 10 → 100 RPS with goodput SLO tracking
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--num-prompts 2000 \
|
||||
--ramp-up-strategy linear --ramp-up-start-rps 10 --ramp-up-end-rps 100 \
|
||||
--goodput ttft:200 e2el:5000 \
|
||||
--save-result
|
||||
|
||||
# Fixed Poisson arrival rate at 50 RPS
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--num-prompts 2000 --request-rate 50 --burstiness 1.0
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Sweep — find the optimal concurrency / rate</b></summary>
|
||||
|
||||
```bash
|
||||
# Sweep over concurrency values
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--num-prompts 500 \
|
||||
--sweep-max-concurrency 1,10,50,100,200,500,1000
|
||||
|
||||
# Sweep over request rates
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--num-prompts 500 \
|
||||
--sweep-request-rate 1,10,50,100,inf
|
||||
|
||||
# Scale work with concurrency and reset the prefix cache between points
|
||||
# (--sweep-num-prompts-factor sets num_prompts = concurrency * factor;
|
||||
# --reset-prefix-cache requires VLLM_SERVER_DEV_MODE=1 on the server)
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--sweep-max-concurrency 1,10,50,100 \
|
||||
--sweep-num-prompts-factor 20 \
|
||||
--reset-prefix-cache
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Multi-run & comparison</b></summary>
|
||||
|
||||
```bash
|
||||
# Run 5 times, report mean/std/min/max with coefficient of variation
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--num-prompts 1000 --max-concurrency 200 --num-runs 5
|
||||
|
||||
# Compare two saved result files side-by-side (no server needed)
|
||||
vllm-bench --compare baseline.json optimized.json
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Multi-turn conversations</b></summary>
|
||||
|
||||
```bash
|
||||
# Synthetic multi-turn (controllable per-turn token lengths)
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name random --multi-turn --multi-turn-num-turns 5 \
|
||||
--random-input-len 512 --random-output-len 256 \
|
||||
--num-prompts 50 --multi-turn-concurrency 10 \
|
||||
--percentile-metrics "ttft,tpot,itl,e2el" --save-result
|
||||
|
||||
# Variable turn count per conversation + per-turn input length for turns 1+
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name random --multi-turn \
|
||||
--multi-turn-min-turns 2 --multi-turn-max-turns 8 \
|
||||
--random-input-len 2048 --per-turn-input-len 256 --random-output-len 128 \
|
||||
--num-prompts 100 --multi-turn-concurrency 20
|
||||
|
||||
# ShareGPT conversations (loads all turns, not just the first two)
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--dataset-name sharegpt --multi-turn \
|
||||
--num-prompts 50 --multi-turn-concurrency 10 --save-result
|
||||
|
||||
# Think time between turns
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--multi-turn --multi-turn-num-turns 3 --multi-turn-delay-ms 500 \
|
||||
--num-prompts 100 --multi-turn-concurrency 20
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>LoRA multi-adapter</b></summary>
|
||||
|
||||
```bash
|
||||
# Distribute requests across N adapters registered on the server.
|
||||
# --model stays the BASE model (tokenizer / readiness / /tokenize use it);
|
||||
# the per-request `model` field is rewritten to one of --lora-modules.
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 \
|
||||
--model Qwen/Qwen3-30B-A3B \
|
||||
--lora-modules sql-lora-1 sql-lora-2 sql-lora-3 sql-lora-4 \
|
||||
sql-lora-5 sql-lora-6 sql-lora-7 sql-lora-8 \
|
||||
--lora-assignment random \
|
||||
--dataset-name random --random-input-len 1024 --random-output-len 256 \
|
||||
--num-prompts 1000 --max-concurrency 64 --save-result
|
||||
|
||||
# Deterministic round-robin assignment (request i -> adapter[i % N])
|
||||
vllm-bench \
|
||||
--backend openai-chat --base-url http://127.0.0.1:8000 \
|
||||
--model Qwen/Qwen3-30B-A3B \
|
||||
--lora-modules sql-lora-1 sql-lora-2 sql-lora-3 sql-lora-4 \
|
||||
--lora-assignment round-robin \
|
||||
--dataset-name random --num-prompts 1000
|
||||
```
|
||||
|
||||
Server side — start vLLM with `--enable-lora` and one `name=path` pair per adapter:
|
||||
|
||||
```bash
|
||||
vllm serve <base-model> \
|
||||
--enable-lora --max-loras 8 --max-lora-rank 16 \
|
||||
--lora-modules \
|
||||
sql-lora-1=jeeejeee/qwen3-moe-text2sql-spider \
|
||||
sql-lora-2=jeeejeee/qwen3-moe-text2sql-spider \
|
||||
...
|
||||
```
|
||||
|
||||
Set `--max-loras` ≥ number of adapter names to keep them all resident (clean steady-state numbers), or lower to stress the LoRA swap path.
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Profiling & dry-run</b></summary>
|
||||
|
||||
```bash
|
||||
# Trigger vLLM server-side profiling (start before, stop after the benchmark)
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--num-prompts 100 --profile
|
||||
|
||||
# Defer profiling until the server batch is full, then capture for 10s
|
||||
vllm-bench \
|
||||
--backend vllm --base-url http://127.0.0.1:8000 --model <model-name> \
|
||||
--num-prompts 2000 --max-concurrency 256 \
|
||||
--profile --profile-batch-threshold 200 --profile-duration 10
|
||||
|
||||
# Dry run: generate dataset, print stats, send nothing
|
||||
vllm-bench \
|
||||
--model <model-name> --num-prompts 100000 --random-input-len 8192 --dry-run
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
## Supported Backends
|
||||
|
||||
### Generation
|
||||
|
||||
| Backend | API Endpoint | Description |
|
||||
| --------- | ------------- | ------------- |
|
||||
| `vllm` / `openai` | `/v1/completions` | OpenAI-compatible completions (streaming) |
|
||||
| `openai-chat` | `/v1/chat/completions` | OpenAI-compatible chat completions (streaming, multimodal) |
|
||||
|
||||
### Embedding / Pooling
|
||||
|
||||
| Backend | API Endpoint | Description |
|
||||
| --------- | ------------- | ------------- |
|
||||
| `openai-embeddings` | `/v1/embeddings` | Text embedding (accepts text or token IDs) |
|
||||
| `openai-embeddings-chat` | `/v1/embeddings` | Chat-format embedding (supports multimodal content) |
|
||||
| `vllm-pooling` | `/v1/pooling` | vLLM native pooling endpoint |
|
||||
| `vllm-rerank` | `/v1/rerank` | vLLM reranking (query from prompt, documents via `--extra-body`) |
|
||||
|
||||
Pooling backends are non-streaming and report E2EL (end-to-end latency) only. Use `--dataset-name sharegpt`, `sonnet`, or `hf` for text-based embedding/rerank benchmarks, or `random` for token-ID-based embedding benchmarks.
|
||||
|
||||
## Supported Datasets
|
||||
|
||||
| Dataset | Description |
|
||||
| --------- | ------------- |
|
||||
| `random` | Synthetic prompts with exact token-length matching (default) |
|
||||
| `random-mm` | Synthetic multimodal prompts with random JPEG images for VLM benchmarking (requires `openai-chat`) |
|
||||
| `sharegpt` | Real conversations from ShareGPT (auto-downloads from HuggingFace, or use `--dataset-path`) |
|
||||
| `sonnet` | Built-in Shakespeare sonnets; controllable token length + shared prefix, no dataset file needed |
|
||||
| `speed-bench` | NVIDIA SPEED-Bench for speculative decoding evaluation (auto-downloads, 11 categories) |
|
||||
| `hf` | Any HuggingFace dataset (auto-downloads via datasets-server API, auto-detects chat/text columns) |
|
||||
|
||||
## Metrics
|
||||
|
||||
### Generation backends
|
||||
|
||||
- **TTFT** (Time to First Token) — latency from request send to first token received
|
||||
- **TPOT** (Time per Output Token) — average time between output tokens
|
||||
- **ITL** (Inter-Token Latency) — per-token latency distribution
|
||||
- **E2EL** (End-to-End Latency) — total request latency
|
||||
- **Throughput** — requests/sec, output tokens/sec, peak output tokens/sec, total tokens/sec
|
||||
- **Concurrency** — peak concurrent requests
|
||||
- **Goodput** — requests/sec meeting all specified SLOs (with `--goodput`)
|
||||
|
||||
### Pooling / embedding backends
|
||||
|
||||
- **E2EL** — total request latency (mean, median, std, percentiles)
|
||||
- **Throughput** — requests/sec, input tokens/sec
|
||||
- **Concurrency** — peak concurrent requests
|
||||
|
||||
### Steady-state metrics
|
||||
|
||||
When `--max-concurrency` is set and `--request-rate` is `inf` (closed-loop mode), the benchmark automatically reports an additional **Steady-State Metrics** block. It measures throughput and latency only over the window during which in-flight concurrency stays at or above a fraction of `--max-concurrency`, excluding the ramp-up and drain phases. This sharply reduces run-to-run variance at very high concurrency.
|
||||
|
||||
The block reports request/input/output/total token throughput plus TTFT (mean, median, percentiles) and TPOT (mean, median, P90, P99) over the detected plateau, along with the window bounds and how many requests fell inside it. Tune it with `--steady-state-threshold` (default `0.95`) and `--steady-state-min-window`, or disable with `--no-steady-state`. The result JSON carries a `steady_state` object (null when not computed).
|
||||
|
||||
## CLI Reference
|
||||
|
||||
Run `vllm-bench --help` for the authoritative list. Grouped reference below.
|
||||
|
||||
<details>
|
||||
<summary><b>Server connection</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--backend` | `openai` | Backend type (`vllm`, `openai`, `openai-chat`, `openai-embeddings`, `openai-embeddings-chat`, `vllm-pooling`, `vllm-rerank`) |
|
||||
| `--base-url` | — | Server base URL (overrides `--host`/`--port`) |
|
||||
| `--host` | `127.0.0.1` | Server host |
|
||||
| `--port` | `8000` | Server port |
|
||||
| `--endpoint` | Auto | API endpoint path (auto-selected per backend) |
|
||||
| `--insecure` | `false` | Disable SSL certificate verification |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Model & tokenizer</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--model` | Auto-detect | Model name (fetched from `/v1/models` if omitted) |
|
||||
| `--served-model-name` | — | Model name used in API requests |
|
||||
| `--tokenizer` | Same as model | Tokenizer name or path (supports HF, tiktoken, server fallback) |
|
||||
| `--tokenizer-mode` | `auto` | Tokenizer mode (`auto`, `hf`, `slow`, `mistral`) |
|
||||
| `--trust-remote-code` | `false` | Trust remote code for tokenizer |
|
||||
| `--skip-tokenizer-init` | `false` | Skip tokenizer initialization |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Dataset</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--dataset-name` | `random` | Dataset type (`random`, `random-mm`, `sharegpt`, `sonnet`, `speed-bench`, `hf`) |
|
||||
| `--dataset-path` | — | Path to dataset file (optional for `sharegpt`/`sonnet`, which auto-source) |
|
||||
| `--num-prompts` | `1000` | Number of prompts to generate (conversations in multi-turn mode) |
|
||||
| `--max-model-len` | — | Filter out requests where `prompt_len + output_len` exceeds this context length |
|
||||
| `--input-len` | — | Override input length (general) |
|
||||
| `--output-len` | — | Override output length (general) |
|
||||
| `--no-oversample` | `false` | Don't oversample if dataset is smaller than `--num-prompts` |
|
||||
| `--disable-shuffle` | `false` | Don't shuffle the dataset |
|
||||
| `--seed` | `0` | Random seed for reproducibility |
|
||||
| **Random** | | |
|
||||
| `--random-input-len` | `1024` | Input token length |
|
||||
| `--random-output-len` | `128` | Output token length |
|
||||
| `--random-prefix-len` | `0` | Shared prefix length |
|
||||
| `--random-range-ratio` | `1.0` | Length jitter, range `(0, 1]`. Lengths sampled from `[ratio × target, target]`; `1.0` = fixed length |
|
||||
| `--prompt-token-ids` | `false` | Send prompts as token-ID arrays (skips server-side tokenization, exact counts). Random dataset only |
|
||||
| **Random multimodal** | | |
|
||||
| `--random-mm-base-items-per-request` | `1` | Base number of multimodal items (images) per request |
|
||||
| `--random-mm-num-mm-items-range-ratio` | `0.0` | Range ratio for varying item count per request |
|
||||
| `--random-mm-limit-mm-per-prompt` | `{"image": 255, "video": 1}` | Per-modality hard caps (JSON) |
|
||||
| `--random-mm-bucket-config` | `{(256,256,1): 0.5, (720,1280,1): 0.5}` | `(height,width,frames)` → probability (Python tuple syntax; frames=1 = image) |
|
||||
| **ShareGPT** | | |
|
||||
| `--sharegpt-output-len` | — | Override output length |
|
||||
| **Sonnet** | | |
|
||||
| `--sonnet-input-len` | `550` | Input tokens per request |
|
||||
| `--sonnet-output-len` | `150` | Output tokens per request |
|
||||
| `--sonnet-prefix-len` | `200` | Prefix tokens shared across requests |
|
||||
| **SPEED-Bench** | | |
|
||||
| `--speed-bench-config` | `qualitative` | Split (`qualitative`, `throughput_1k`/`2k`/`8k`/`16k`/`32k`) |
|
||||
| `--speed-bench-category` | — | Filter by category (`low_entropy`, `high_entropy`, `mixed_entropy`, `coding`, `math`, …) |
|
||||
| `--speed-bench-max-input-len` | — | Truncate prompts to at most N tokens |
|
||||
| **HuggingFace** | | |
|
||||
| `--hf-split` | Auto | Split (`train`, `test`, `validation`); auto-detected if omitted |
|
||||
| `--hf-subset` | — | Subset/config name (e.g. `narrativeqa` for LongBench) |
|
||||
| `--hf-output-len` | — | Fixed output length for all requests (overrides dataset-derived length) |
|
||||
| `--hf-text-column` | Auto | Column containing prompt text; auto-detected from common patterns |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Rate control</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--request-rate` | `inf` | Requests per second (`inf` = all at once) |
|
||||
| `--burstiness` | `1.0` | Burstiness factor (1.0 = Poisson, >1 = bursty) |
|
||||
| `--max-concurrency` | `num-prompts` | Maximum concurrent requests (semaphore) |
|
||||
| `--ramp-up-strategy` | — | Ramp-up mode (`linear` or `exponential`) |
|
||||
| `--ramp-up-start-rps` | — | Starting request rate for ramp-up |
|
||||
| `--ramp-up-end-rps` | — | Ending request rate for ramp-up |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Sampling parameters</b></summary>
|
||||
|
||||
| Flag | Description |
|
||||
| ------ | ------------- |
|
||||
| `--temperature` | Temperature (server default if omitted) |
|
||||
| `--top-p` | Top-p (nucleus) sampling |
|
||||
| `--top-k` | Top-k sampling |
|
||||
| `--min-p` | Min-p sampling |
|
||||
| `--frequency-penalty` | Frequency penalty |
|
||||
| `--presence-penalty` | Presence penalty |
|
||||
| `--repetition-penalty` | Repetition penalty |
|
||||
|
||||
Merged into the request body. Only effective with generation backends (`vllm`, `openai`, `openai-chat`); ignored by pooling/embedding backends.
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Output & results</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--save-result` | `false` | Save results to JSON file |
|
||||
| `--save-detailed` | `false` | Include per-request data in JSON (input/output lens, ITLs, texts) |
|
||||
| `--append-result` | `false` | Append to existing JSON file (JSONL format) |
|
||||
| `--result-dir` | — | Directory for result files |
|
||||
| `--result-filename` | Auto | Custom result filename |
|
||||
| `--percentile-metrics` | `ttft,tpot,itl,e2el` | Metrics for percentile reporting (pooling defaults to `e2el` only) |
|
||||
| `--metric-percentiles` | `99` | Percentile values to compute |
|
||||
| `--sweep-summary-percentiles` | — | Extra percentiles for sweep summary tables (auto-added to computed set) |
|
||||
| `--goodput` | — | SLO pairs for goodput (`ttft:100 tpot:50 e2el:500`, values in ms) |
|
||||
| `--disable-tqdm` | `false` | Disable progress bar |
|
||||
| `--label` | — | Label prefix for result files |
|
||||
| `--metadata` | — | Key-value metadata (`KEY=VALUE`, repeatable) |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Request options</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--ignore-eos` | `false` | Ignore EOS token (force full output length) |
|
||||
| `--logprobs` | — | Number of logprobs per token |
|
||||
| `--num-warmups` | `0` | Warmup requests before benchmarking |
|
||||
| `--ready-check-timeout-sec` | `0` | Endpoint readiness timeout (0 = skip) |
|
||||
| `--request-id-prefix` | Auto (UUID) | Prefix for request IDs |
|
||||
| `--header` | — | Extra headers (`KEY=VALUE`, repeatable) |
|
||||
| `--extra-body` | — | Extra JSON body parameters |
|
||||
| `--dry-run` | `false` | Generate dataset only, skip benchmark |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Steady-state metrics</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--steady-state-threshold` | `0.95` | Fraction of `--max-concurrency` at which the steady-state window opens, range (0, 1] |
|
||||
| `--steady-state-min-window` | Auto | Minimum window duration (s) below which a warning is attached. Default `max(10, 0.1 × run_duration)` |
|
||||
| `--no-steady-state` | `false` | Disable steady-state metrics computation |
|
||||
|
||||
Computed only when `--max-concurrency` is set and `--request-rate` is `inf`.
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Profiling</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--profile` | `false` | Trigger vLLM server-side profiling (`/start_profile` before, `/stop_profile` after) |
|
||||
| `--profile-batch-threshold` | — | Defer profiling until `/metrics` reports ≥ N running requests, then capture. Requires `--profile` |
|
||||
| `--profile-duration` | `5.0` | Seconds to capture once the batch threshold is reached. Requires `--profile-batch-threshold` |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Sweep mode</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--sweep-max-concurrency` | — | Comma-separated concurrency values to sweep (e.g. `1,10,50,100,500`) |
|
||||
| `--sweep-request-rate` | — | Comma-separated rate values to sweep, supports `inf` (e.g. `1,10,100,inf`) |
|
||||
| `--sweep-num-prompts-factor` | — | Set `num_prompts = concurrency × factor` per concurrency sweep point |
|
||||
| `--reset-prefix-cache` | `false` | Reset the server's prefix cache before each sweep iteration (requires `VLLM_SERVER_DEV_MODE=1`) |
|
||||
|
||||
Runs the benchmark once per value, then prints a summary table comparing throughput and latency across all sweep points and identifies the best-throughput configuration. Works in multi-turn mode too. `--sweep-summary-percentiles` appends extra TTFT/TPOT/E2EL columns to the summary, auto-adding any missing percentiles to the computed set so they also appear in result JSON.
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Multi-turn conversation benchmark</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--multi-turn` | `false` | Enable multi-turn conversation mode (requires `--backend openai-chat`) |
|
||||
| `--multi-turn-num-turns` | `3` | Turns per conversation (synthetic mode) |
|
||||
| `--multi-turn-min-turns` | `0` | Minimum turns per conversation (0 = use `--multi-turn-num-turns`) |
|
||||
| `--multi-turn-max-turns` | `0` | Maximum turns per conversation (0 = `--multi-turn-num-turns` synthetic / uncapped ShareGPT) |
|
||||
| `--multi-turn-concurrency` | — | Concurrent conversations (defaults to `--max-concurrency` or `--num-prompts`) |
|
||||
| `--multi-turn-delay-ms` | `0` | Delay between turns in ms (simulates user think time) |
|
||||
| `--per-turn-input-len` | `0` | Input token length for turns 1+ (0 = use `--random-input-len` for all turns) |
|
||||
| `--multi-turn-prefix-global-ratio` | `0.0` | Fraction of per-turn input shared across all conversations (random dataset only) |
|
||||
| `--multi-turn-prefix-conversation-ratio` | `0.0` | Fraction shared within each conversation (random dataset only) |
|
||||
|
||||
With `--multi-turn`, `--num-prompts` controls the number of **conversations**, not individual requests.
|
||||
|
||||
**How it works:**
|
||||
|
||||
- Turn 1: send `[user_1]`, get `assistant_1`
|
||||
- Turn 2: send `[user_1, assistant_1, user_2]`, get `assistant_2`
|
||||
- Turn N: send full history + `user_N` — measures growing-context performance
|
||||
|
||||
**Data sources:**
|
||||
|
||||
- `--dataset-name random` — synthetic conversations with controllable per-turn token lengths. Auto-sets `min_tokens` to enforce output length without `ignore_eos`.
|
||||
- `--dataset-name sharegpt` — loads all turns (not just the first two); filters for entries with ≥ 2 real turns.
|
||||
|
||||
**Prefix sharing** (random dataset): when `--multi-turn-prefix-global-ratio` or `--multi-turn-prefix-conversation-ratio` is > 0, each turn sends a fixed-length message (no history accumulation) composed of a global prefix + per-conversation prefix + unique suffix. The two ratios must sum to < 1.0.
|
||||
|
||||
**Router affinity:** every turn sends `X-Session-ID: {conversation_id}` for KV-cache reuse behind a vLLM router.
|
||||
|
||||
**Output:** overall metrics plus a per-turn breakdown (TTFT/TPOT/ITL/E2EL by turn index). Expect TTFT to climb across turns due to growing context. JSON includes a `per_turn_metrics` array.
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>LoRA multi-adapter</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--lora-modules` | — | Adapter names registered on the server (`vllm serve --lora-modules name=path`). Each request's `model` field is rewritten to one of these. Repeatable |
|
||||
| `--lora-assignment` | `random` | Distribution: `random` (uniform, seeded by `--seed`) or `round-robin` (deterministic `i % N`) |
|
||||
|
||||
`--model` must stay the **base** model — its tokenizer builds prompts, and `/v1/models`, `/tokenize`, ready check, and warmup all use it. Only the per-request `model` field in completions/chat payloads is rewritten to the assigned adapter (vLLM routes by name).
|
||||
|
||||
**Assignment scope:** per request in single-shot mode; **per conversation** (sticky across all turns) in multi-turn mode, to avoid breaking prefix-cache reuse mid-dialog.
|
||||
|
||||
**Reproducibility:** with `--lora-assignment random`, the same `--seed` + same `--lora-modules` list yields identical request-to-adapter mappings. Pooling/embedding backends are rejected — LoRA routing applies to generative paths only.
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Multi-run & comparison</b></summary>
|
||||
|
||||
| Flag | Default | Description |
|
||||
| ------ | --------- | ------------- |
|
||||
| `--num-runs` | `1` | Run benchmark N times; report mean/std/min/max with CV |
|
||||
| `--compare` | — | Compare two result JSON files side-by-side (skips benchmarking) |
|
||||
|
||||
`--num-runs` aggregates metrics across runs and reports the coefficient of variation (CV) for throughput stability. `--compare` reads two previously-saved result files and prints a diff with delta, % change, and improvement/regression markers.
|
||||
|
||||
</details>
|
||||
|
||||
## Tokenizer Support
|
||||
|
||||
Tokenizers are loaded with a three-tier fallback chain:
|
||||
|
||||
1. **Local HuggingFace** — `tokenizer.json` from a local path or the Hub (fastest)
|
||||
2. **Tiktoken** — `.tiktoken` / `.model` format for Kimi, Qwen, etc. (auto-extracts `pat_str` from Python source)
|
||||
3. **Server-side** — falls back to vLLM's `/tokenize` + `/detokenize` endpoints
|
||||
|
||||
For the `random` dataset, prompt token lengths are verified against the server on the first run and cached; subsequent runs with the same model+server skip verification. Verification is also skipped when `--prompt-token-ids` is set (token counts are exact by construction).
|
||||
|
||||
Models without `tokenizer.json` (e.g. `nvidia/Kimi-K2.5-NVFP4`) fall back to server-side tokenization automatically; you can also point `--tokenizer` at a model that ships `tokenizer.json`.
|
||||
|
||||
## Output Format
|
||||
|
||||
JSON output is compatible with the `vllm bench serve` Python schema. Result files are named:
|
||||
|
||||
```text
|
||||
{label}-{rate}qps-concurrency{max_concurrency}-{model}-{timestamp}.json
|
||||
```
|
||||
|
||||
Use `--append-result` to append multiple runs to the same file in JSONL format. `--save-detailed` adds per-request arrays (input/output lengths, ITLs, generated text).
|
||||
|
||||
## Architecture
|
||||
|
||||
<details>
|
||||
<summary><b>Source layout</b></summary>
|
||||
|
||||
```text
|
||||
src/
|
||||
├── main.rs # Entry point, mimalloc, tokio runtime, mode dispatch
|
||||
├── cli.rs # clap CLI argument definitions
|
||||
├── config.rs # Validated config, goodput/ramp-up parsing
|
||||
├── benchmark.rs # Core orchestrator (schedule, spawn, collect, verify, profile)
|
||||
├── multi_turn.rs # Multi-turn conversation orchestrator (channel workers)
|
||||
├── compare.rs # Result diff (--compare file_a.json file_b.json)
|
||||
├── sweep.rs # Parameter sweep (--sweep-max-concurrency, --sweep-request-rate)
|
||||
├── multi_run.rs # Multi-run statistics (--num-runs N)
|
||||
├── rate_control.rs # Gamma/Poisson scheduling + linear/exponential ramp-up
|
||||
├── ready_checker.rs # Endpoint readiness with retry
|
||||
├── tokenizer.rs # Tokenizer abstraction (HF, tiktoken, server)
|
||||
├── tiktoken.rs # Tiktoken BPE loader with pat_str extraction
|
||||
├── error.rs # Error types
|
||||
├── backends/
|
||||
│ ├── mod.rs # Backend enum dispatch, typed SSE structs
|
||||
│ ├── streaming.rs # SSE stream parser with speculative JSON parse
|
||||
│ ├── openai_completions.rs # /v1/completions backend
|
||||
│ ├── openai_chat.rs # /v1/chat/completions backend
|
||||
│ └── pooling.rs # Embedding/pooling/rerank backends (non-streaming)
|
||||
├── datasets/
|
||||
│ ├── mod.rs # SampleRequest, ConversationTurn, MultiTurnConversation types
|
||||
│ ├── random.rs # Random dataset with rayon parallelism
|
||||
│ ├── random_mm.rs # Random multimodal dataset (JPEG generation, bucket sampling)
|
||||
│ ├── multi_turn.rs # Multi-turn synthetic + ShareGPT conversation generators
|
||||
│ ├── sharegpt.rs # ShareGPT JSON dataset loader
|
||||
│ ├── sonnet.rs # Sonnet dataset (built-in Shakespeare sonnets)
|
||||
│ ├── speed_bench.rs # NVIDIA SPEED-Bench loader (auto-download + cache)
|
||||
│ └── hf_dataset.rs # Generic HuggingFace dataset (auto-download, column detection)
|
||||
├── metrics/
|
||||
│ ├── mod.rs # BenchmarkMetrics, MultiTurnMetrics structs
|
||||
│ ├── calculator.rs # Percentile/throughput/goodput/peak/multi-turn computation
|
||||
│ └── steady_state.rs # Steady-state window detection + plateau metrics
|
||||
└── output/
|
||||
├── mod.rs
|
||||
├── console.rs # Terminal output (matches Python format)
|
||||
└── json.rs # JSON result serialization (Python-compatible schema)
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
### Key design decisions
|
||||
|
||||
- **reqwest + tokio** — HTTP client with connection pooling, forced HTTP/1.1, TCP_NODELAY to match Python's aiohttp and avoid Nagle latency inflation on TTFT
|
||||
- **mimalloc** — global allocator to reduce contention under high concurrency (1400+ tasks); page-agnostic, runs on aarch64 4K- and 64K-page kernels
|
||||
- **`Arc<str>` prompts** — zero-copy prompt sharing across tokio tasks, eliminating ~3 GB peak memory at 100k requests with 8k-token prompts
|
||||
- **Spawn-per-request** — `tokio::spawn` per request with a `Semaphore` for concurrency control (matches Python's asyncio pattern)
|
||||
- **rayon** — parallel dataset generation across CPU cores (200–500× faster than Python for 100k+ prompts)
|
||||
- **Enum dispatch** — backend variants instead of trait objects (avoids async trait-object limitations)
|
||||
- **Typed SSE deserialization** — `CompletionChunk`/`ChatChunk` structs skip unused JSON fields (cheaper than `serde_json::Value`)
|
||||
- **Speculative JSON parse** — SSE handler uses `serde_json::value::RawValue` to detect complete JSON before `\n\n` arrives, improving TTFT/ITL accuracy when TCP segments split
|
||||
- **Connection error retry** — automatic retry with backoff on connection reset/timeout/refused (up to 3 attempts)
|
||||
- **Tokenizer verification cache** — server-side token-length verification is cached per model+server pair
|
||||
|
||||
### Behavioral parity with Python
|
||||
|
||||
The Rust implementation matches Python `vllm bench serve` in:
|
||||
|
||||
- SSE streaming protocol handling (including speculative parse for split TCP segments)
|
||||
- Timing semantics (monotonic `Instant` matching Python's `time.perf_counter()`)
|
||||
- Chat vs. completions differences (`max_completion_tokens` vs. `max_tokens`, Content-Type, timestamp placement)
|
||||
- JSON output schema (all fields, key naming, `request_rate` as the string `"inf"`)
|
||||
- Rate control (Gamma distribution, normalization, burstiness, linear/exponential ramp-up)
|
||||
- Metrics (TTFT/TPOT/ITL/E2EL percentiles, peak tokens/sec, peak concurrency, goodput)
|
||||
- Sampling parameters merged into the request body via `extra_body` (same precedence rules)
|
||||
|
||||
## Environment Variables
|
||||
|
||||
| Variable | Description |
|
||||
| ---------- | ------------- |
|
||||
| `OPENAI_API_KEY` | API key for authenticated endpoints (cached, not read per-request) |
|
||||
| `HF_TOKEN` | HuggingFace token for gated model tokenizers and gated datasets |
|
||||
| `TOKIO_WORKER_THREADS` | Override tokio worker thread count (default: physical cores) |
|
||||
|
||||
## License
|
||||
|
||||
Apache-2.0
|
||||
</content>
|
||||
</invoke>
|
||||
@@ -0,0 +1,213 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
pub mod openai_chat;
|
||||
pub mod openai_completions;
|
||||
pub mod pooling;
|
||||
pub mod streaming;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
// --- Typed SSE chunk structs for zero-alloc deserialization ---
|
||||
// Using typed deserialization avoids building a full serde_json::Value tree.
|
||||
// Only the fields we need are extracted; everything else is skipped by serde.
|
||||
|
||||
/// Completions API streaming chunk (minimal fields).
|
||||
#[derive(Deserialize)]
|
||||
pub struct CompletionChunk {
|
||||
#[serde(default)]
|
||||
pub choices: Vec<CompletionChoice>,
|
||||
pub usage: Option<ChunkUsage>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct CompletionChoice {
|
||||
pub text: Option<String>,
|
||||
}
|
||||
|
||||
/// Chat API streaming chunk (minimal fields).
|
||||
#[derive(Deserialize)]
|
||||
pub struct ChatChunk {
|
||||
#[serde(default)]
|
||||
pub choices: Vec<ChatChoice>,
|
||||
pub usage: Option<ChunkUsage>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct ChatChoice {
|
||||
pub delta: Option<ChatDelta>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct ChatDelta {
|
||||
pub content: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct ChunkUsage {
|
||||
pub completion_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
use crate::cli::BackendKind;
|
||||
use crate::error::Result;
|
||||
|
||||
/// Input for a single benchmark request.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RequestFuncInput {
|
||||
pub prompt: Arc<str>,
|
||||
pub api_url: String,
|
||||
pub prompt_len: usize,
|
||||
pub output_len: usize,
|
||||
pub model: String,
|
||||
pub model_name: Option<String>,
|
||||
pub logprobs: Option<usize>,
|
||||
pub extra_headers: Option<HashMap<String, String>>,
|
||||
pub extra_body: Option<serde_json::Value>,
|
||||
pub ignore_eos: bool,
|
||||
pub request_id: Option<String>,
|
||||
/// Pre-built messages array for multi-turn conversations.
|
||||
/// When set, the chat backend uses this instead of building from `prompt`.
|
||||
pub messages: Option<serde_json::Value>,
|
||||
/// Pre-computed token IDs for this prompt.
|
||||
/// When set, the completions backend sends these directly via `prompt_token_ids`
|
||||
/// instead of the text `prompt`, skipping server-side tokenization.
|
||||
pub prompt_token_ids: Option<Arc<[u32]>>,
|
||||
/// Multimodal content as pre-serialized JSON fragments.
|
||||
/// When set, the chat backend concatenates these directly into the payload bytes,
|
||||
/// avoiding any parsing or deep-cloning of base64 image data.
|
||||
pub multi_modal_content: Option<Arc<[Arc<str>]>>,
|
||||
/// Complete pre-serialized chat `messages` array (--enable-multimodal-chat).
|
||||
/// When set, the chat backend splices it verbatim into the payload bytes,
|
||||
/// taking precedence over `messages`, `prompt`, and `multi_modal_content`.
|
||||
pub chat_messages_json: Option<Arc<str>>,
|
||||
/// Multiple text inputs for one request (pooling backends only):
|
||||
/// embeddings batch (`"input": [...]`) or rerank query+documents.
|
||||
pub prompt_list: Option<Arc<[Arc<str>]>>,
|
||||
}
|
||||
|
||||
/// Output from a single benchmark request including timing metrics.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct RequestFuncOutput {
|
||||
pub generated_text: String,
|
||||
pub success: bool,
|
||||
pub latency: f64,
|
||||
pub output_tokens: usize,
|
||||
pub ttft: f64,
|
||||
pub itl: Vec<f64>,
|
||||
pub tpot: f64,
|
||||
pub prompt_len: usize,
|
||||
pub error: String,
|
||||
pub start_time: f64,
|
||||
}
|
||||
|
||||
impl Default for RequestFuncOutput {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
generated_text: String::new(),
|
||||
success: false,
|
||||
latency: 0.0,
|
||||
output_tokens: 0,
|
||||
ttft: 0.0,
|
||||
itl: Vec::new(),
|
||||
tpot: 0.0,
|
||||
prompt_len: 0,
|
||||
error: String::new(),
|
||||
start_time: 0.0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for RequestFuncInput {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
prompt: Arc::from(""),
|
||||
api_url: String::new(),
|
||||
prompt_len: 0,
|
||||
output_len: 0,
|
||||
model: String::new(),
|
||||
model_name: None,
|
||||
logprobs: None,
|
||||
extra_headers: None,
|
||||
extra_body: None,
|
||||
ignore_eos: false,
|
||||
request_id: None,
|
||||
messages: None,
|
||||
prompt_token_ids: None,
|
||||
multi_modal_content: None,
|
||||
chat_messages_json: None,
|
||||
prompt_list: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Enum dispatch for backend implementations (avoids async trait object issues).
|
||||
#[derive(Clone)]
|
||||
pub enum Backend {
|
||||
OpenAICompletions(openai_completions::OpenAICompletionsBackend),
|
||||
OpenAIChat(openai_chat::OpenAIChatBackend),
|
||||
Pooling(pooling::PoolingBackend),
|
||||
}
|
||||
|
||||
impl Backend {
|
||||
/// Send a single request and collect timing metrics.
|
||||
pub async fn send_request(
|
||||
&self,
|
||||
input: &RequestFuncInput,
|
||||
client: &reqwest::Client,
|
||||
) -> Result<RequestFuncOutput> {
|
||||
match self {
|
||||
Backend::OpenAICompletions(b) => b.send_request(input, client).await,
|
||||
Backend::OpenAIChat(b) => b.send_request(input, client).await,
|
||||
Backend::Pooling(b) => b.send_request(input, client).await,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Get a backend by kind.
|
||||
pub fn get_backend(kind: BackendKind) -> Result<Backend> {
|
||||
match kind {
|
||||
BackendKind::Vllm | BackendKind::Openai => Ok(Backend::OpenAICompletions(
|
||||
openai_completions::OpenAICompletionsBackend,
|
||||
)),
|
||||
BackendKind::OpenaiChat => Ok(Backend::OpenAIChat(openai_chat::OpenAIChatBackend)),
|
||||
kind if kind.is_pooling() => Ok(Backend::Pooling(pooling::PoolingBackend { kind })),
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Cached API key to avoid per-request env var syscall.
|
||||
static API_KEY: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
|
||||
|
||||
fn cached_api_key() -> &'static Option<String> {
|
||||
API_KEY.get_or_init(|| std::env::var("OPENAI_API_KEY").ok())
|
||||
}
|
||||
|
||||
/// Build common headers including auth and extras.
|
||||
pub fn build_headers(
|
||||
content_type: Option<&str>,
|
||||
extra_headers: &Option<HashMap<String, String>>,
|
||||
request_id: &Option<String>,
|
||||
) -> HashMap<String, String> {
|
||||
let mut headers = HashMap::new();
|
||||
|
||||
if let Some(ct) = content_type {
|
||||
headers.insert("Content-Type".to_string(), ct.to_string());
|
||||
}
|
||||
|
||||
if let Some(api_key) = cached_api_key() {
|
||||
headers.insert("Authorization".to_string(), format!("Bearer {api_key}"));
|
||||
}
|
||||
|
||||
if let Some(extra) = extra_headers {
|
||||
headers.extend(extra.iter().map(|(k, v)| (k.clone(), v.clone())));
|
||||
}
|
||||
|
||||
if let Some(rid) = request_id {
|
||||
headers.insert("x-request-id".to_string(), rid.clone());
|
||||
}
|
||||
|
||||
headers
|
||||
}
|
||||
@@ -0,0 +1,376 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::time::Instant;
|
||||
|
||||
use futures::StreamExt;
|
||||
|
||||
use super::streaming::{StreamedResponseHandler, trim_bytes};
|
||||
use super::{ChatChunk, RequestFuncInput, RequestFuncOutput, build_headers};
|
||||
use crate::error::Result;
|
||||
|
||||
/// Backend for OpenAI Chat Completions API (/v1/chat/completions).
|
||||
#[derive(Clone)]
|
||||
pub struct OpenAIChatBackend;
|
||||
|
||||
impl OpenAIChatBackend {
|
||||
pub async fn send_request(
|
||||
&self,
|
||||
input: &RequestFuncInput,
|
||||
client: &reqwest::Client,
|
||||
) -> Result<RequestFuncOutput> {
|
||||
// Content-Type is set below by `.json()` / `.header()`; keep it out of
|
||||
// headers_map to avoid a duplicate that strict gateways reject.
|
||||
let headers_map = build_headers(None, &input.extra_headers, &input.request_id);
|
||||
|
||||
let mut output = RequestFuncOutput {
|
||||
prompt_len: input.prompt_len,
|
||||
itl: Vec::with_capacity(input.output_len.max(1)),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let st = Instant::now();
|
||||
|
||||
let mut most_recent_timestamp = st;
|
||||
let mut generated_text = String::new();
|
||||
let mut first_token_received = false;
|
||||
|
||||
// Build request: use zero-copy raw JSON for multimodal, serde_json for text-only
|
||||
let mut request =
|
||||
if input.multi_modal_content.is_some() || input.chat_messages_json.is_some() {
|
||||
let payload_bytes = build_mm_payload(input);
|
||||
client
|
||||
.post(&input.api_url)
|
||||
.header("content-type", "application/json")
|
||||
.body(payload_bytes)
|
||||
} else {
|
||||
let payload = build_text_payload(input);
|
||||
client.post(&input.api_url).json(&payload)
|
||||
};
|
||||
for (k, v) in &headers_map {
|
||||
request = request.header(k, v);
|
||||
}
|
||||
|
||||
match request.send().await {
|
||||
Ok(response) => {
|
||||
if response.status().is_success() {
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let mut stream = response.bytes_stream();
|
||||
|
||||
while let Some(chunk_result) = stream.next().await {
|
||||
let chunk_bytes = match chunk_result {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
output.success = false;
|
||||
output.error = format!("Stream error: {e}");
|
||||
return Ok(output);
|
||||
}
|
||||
};
|
||||
|
||||
let trimmed_bytes = trim_bytes(&chunk_bytes);
|
||||
if trimmed_bytes.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let messages = handler.add_chunk(trimmed_bytes);
|
||||
for message in messages {
|
||||
// Skip SSE comments
|
||||
if message.starts_with(':') {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Handle multi-field SSE events (e.g., Dynamo sends
|
||||
// "event: message\ndata: {...}"). Extract the data: line.
|
||||
let raw = if message.contains('\n') {
|
||||
match message.lines().find(|l| l.starts_with("data: ")) {
|
||||
Some(l) => l,
|
||||
None => continue,
|
||||
}
|
||||
} else {
|
||||
message.as_str()
|
||||
};
|
||||
|
||||
let chunk = raw.strip_prefix("data: ").unwrap_or(raw);
|
||||
|
||||
if chunk == "[DONE]" {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Python chat backend: timestamp is captured for ALL
|
||||
// non-DONE messages, and most_recent_timestamp is updated
|
||||
// unconditionally (outside `if choices:`). This differs from
|
||||
// completions which only timestamps content chunks.
|
||||
let timestamp = Instant::now();
|
||||
|
||||
let data: ChatChunk = match serde_json::from_str(chunk) {
|
||||
Ok(d) => d,
|
||||
Err(_) => continue,
|
||||
};
|
||||
|
||||
if !data.choices.is_empty() {
|
||||
let content = data.choices[0]
|
||||
.delta
|
||||
.as_ref()
|
||||
.and_then(|d| d.content.as_deref())
|
||||
.unwrap_or("");
|
||||
|
||||
if !first_token_received {
|
||||
first_token_received = true;
|
||||
output.ttft = timestamp.duration_since(st).as_secs_f64();
|
||||
} else {
|
||||
output.itl.push(
|
||||
timestamp
|
||||
.duration_since(most_recent_timestamp)
|
||||
.as_secs_f64(),
|
||||
);
|
||||
}
|
||||
|
||||
generated_text.push_str(content);
|
||||
}
|
||||
// Separate `if` (not `else if`) — Dynamo may send
|
||||
// both choices and usage in the same chunk.
|
||||
if let Some(ref usage) = data.usage
|
||||
&& let Some(ct) = usage.completion_tokens
|
||||
{
|
||||
output.output_tokens = ct as usize;
|
||||
}
|
||||
|
||||
most_recent_timestamp = timestamp;
|
||||
}
|
||||
}
|
||||
|
||||
output.generated_text = generated_text;
|
||||
output.success = true;
|
||||
output.latency = most_recent_timestamp.duration_since(st).as_secs_f64();
|
||||
} else {
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
output.error = if body.is_empty() {
|
||||
format!("HTTP {status}")
|
||||
} else {
|
||||
format!("HTTP {status}: {body}")
|
||||
};
|
||||
output.success = false;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
output.success = false;
|
||||
output.error = format!("{e:#}");
|
||||
}
|
||||
}
|
||||
|
||||
Ok(output)
|
||||
}
|
||||
}
|
||||
|
||||
/// Build a JSON payload for text-only (non-multimodal) requests using serde_json.
|
||||
fn build_text_payload(input: &RequestFuncInput) -> serde_json::Value {
|
||||
let model = input.model_name.as_deref().unwrap_or(&input.model);
|
||||
|
||||
let messages = if let Some(ref msgs) = input.messages {
|
||||
msgs.clone()
|
||||
} else {
|
||||
let content = serde_json::json!([
|
||||
{"type": "text", "text": input.prompt}
|
||||
]);
|
||||
serde_json::json!([{"role": "user", "content": content}])
|
||||
};
|
||||
|
||||
let mut payload = serde_json::json!({
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"max_completion_tokens": input.output_len,
|
||||
"stream": true,
|
||||
"stream_options": {
|
||||
"include_usage": true,
|
||||
},
|
||||
});
|
||||
|
||||
if input.ignore_eos {
|
||||
payload["ignore_eos"] = serde_json::json!(true);
|
||||
}
|
||||
if let Some(serde_json::Value::Object(map)) = input.extra_body.as_ref() {
|
||||
for (k, v) in map {
|
||||
payload[k] = v.clone();
|
||||
}
|
||||
}
|
||||
|
||||
payload
|
||||
}
|
||||
|
||||
/// Build the JSON payload as raw bytes for multimodal requests.
|
||||
///
|
||||
/// This is the zero-copy fast path: pre-serialized mm content fragments
|
||||
/// (each ~200KB+ of base64 image data) are concatenated directly into the
|
||||
/// output buffer without being parsed, cloned, or re-serialized.
|
||||
///
|
||||
/// Saves ~200KB of allocation + copy per image per request compared to
|
||||
/// the serde_json::Value approach.
|
||||
fn build_mm_payload(input: &RequestFuncInput) -> Vec<u8> {
|
||||
let model = input.model_name.as_deref().unwrap_or(&input.model);
|
||||
|
||||
// Estimate total size: JSON overhead (~300 bytes) + prompt + mm fragments
|
||||
let mm_total: usize = input
|
||||
.multi_modal_content
|
||||
.as_ref()
|
||||
.map(|mm| mm.iter().map(|f| f.len() + 1).sum())
|
||||
.unwrap_or(0)
|
||||
+ input.chat_messages_json.as_ref().map_or(0, |m| m.len());
|
||||
let estimated = 512 + input.prompt.len() * 2 + mm_total;
|
||||
let mut json = String::with_capacity(estimated);
|
||||
|
||||
// {"model": <model>
|
||||
json.push_str(r#"{"model":"#);
|
||||
// serde_json::to_string on &str produces a JSON-escaped quoted string
|
||||
json.push_str(&serde_json::to_string(model).unwrap());
|
||||
|
||||
json.push_str(r#","messages":"#);
|
||||
if let Some(ref msgs) = input.chat_messages_json {
|
||||
// --enable-multimodal-chat: the dataset pre-built the full messages
|
||||
// array (text + mm parts); splice it verbatim.
|
||||
json.push_str(msgs);
|
||||
} else {
|
||||
let mm = input.multi_modal_content.as_ref().unwrap();
|
||||
|
||||
// [{"role":"user","content":[ <text part>
|
||||
json.push_str(r#"[{"role":"user","content":[{"type":"text","text":""#);
|
||||
// JSON-escape the prompt text (handles \n, \t, unicode, quotes)
|
||||
push_json_escaped_str(&mut json, &input.prompt);
|
||||
json.push_str(r#""}"#);
|
||||
|
||||
// ,<mm fragment 1>,<mm fragment 2>,...
|
||||
for fragment in mm.iter() {
|
||||
json.push(',');
|
||||
json.push_str(fragment);
|
||||
}
|
||||
|
||||
// Close content, message, messages
|
||||
json.push_str(r#"]}]"#);
|
||||
}
|
||||
|
||||
// ,"max_completion_tokens": N, "stream": true, ...
|
||||
json.push_str(r##","max_completion_tokens":"##);
|
||||
json.push_str(&input.output_len.to_string());
|
||||
json.push_str(r##","stream":true,"stream_options":{"include_usage":true}"##);
|
||||
|
||||
if input.ignore_eos {
|
||||
json.push_str(r#","ignore_eos":true"#);
|
||||
}
|
||||
|
||||
// Merge extra_body key-value pairs, skipping keys already set above
|
||||
if let Some(serde_json::Value::Object(map)) = input.extra_body.as_ref() {
|
||||
for (k, v) in map {
|
||||
match k.as_str() {
|
||||
"model"
|
||||
| "messages"
|
||||
| "max_completion_tokens"
|
||||
| "stream"
|
||||
| "stream_options"
|
||||
| "ignore_eos" => continue,
|
||||
_ => {
|
||||
json.push(',');
|
||||
json.push_str(&serde_json::to_string(k).unwrap());
|
||||
json.push(':');
|
||||
json.push_str(&serde_json::to_string(v).unwrap());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
json.push('}');
|
||||
json.into_bytes()
|
||||
}
|
||||
|
||||
/// Write a JSON-escaped string (without surrounding quotes) into the buffer.
|
||||
///
|
||||
/// Handles: `\n`, `\r`, `\t`, `\\`, `\"`, and control characters.
|
||||
/// This avoids the allocation of `serde_json::to_string` which produces
|
||||
/// a new String with surrounding quotes.
|
||||
fn push_json_escaped_str(buf: &mut String, s: &str) {
|
||||
use std::fmt::Write;
|
||||
for ch in s.chars() {
|
||||
match ch {
|
||||
'"' => buf.push_str(r#"\""#),
|
||||
'\\' => buf.push_str(r"\\"),
|
||||
'\n' => buf.push_str(r"\n"),
|
||||
'\r' => buf.push_str(r"\r"),
|
||||
'\t' => buf.push_str(r"\t"),
|
||||
c if c.is_control() => {
|
||||
// \uXXXX escape for control characters
|
||||
for unit in c.encode_utf16(&mut [0; 2]) {
|
||||
write!(buf, "\\u{unit:04x}").unwrap();
|
||||
}
|
||||
}
|
||||
c => buf.push(c),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn mm_input() -> RequestFuncInput {
|
||||
let frag: Arc<str> =
|
||||
Arc::from(r#"{"type":"image_url","image_url":{"url":"data:image/jpeg;base64,AAAA"}}"#);
|
||||
RequestFuncInput {
|
||||
prompt: Arc::from("hello \"world\"\nline2"),
|
||||
model: "test-model".to_string(),
|
||||
output_len: 128,
|
||||
multi_modal_content: Some(Arc::from(vec![frag])),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// Regression test: the assembled multimodal payload must be valid JSON
|
||||
/// (the text part once shipped without its opening quote — a raw-string
|
||||
/// delimiter eating the trailing `"` in `"text":"`).
|
||||
#[test]
|
||||
fn test_build_mm_payload_is_valid_json() {
|
||||
let payload = build_mm_payload(&mm_input());
|
||||
let v: serde_json::Value =
|
||||
serde_json::from_slice(&payload).expect("mm payload must be valid JSON");
|
||||
assert_eq!(v["model"], "test-model");
|
||||
assert_eq!(v["messages"][0]["role"], "user");
|
||||
let content = v["messages"][0]["content"].as_array().unwrap();
|
||||
assert_eq!(content[0]["type"], "text");
|
||||
assert_eq!(content[0]["text"], "hello \"world\"\nline2");
|
||||
assert_eq!(content[1]["type"], "image_url");
|
||||
assert_eq!(v["max_completion_tokens"], 128);
|
||||
assert_eq!(v["stream"], true);
|
||||
assert_eq!(v["stream_options"]["include_usage"], true);
|
||||
}
|
||||
|
||||
/// --enable-multimodal-chat (dataset pre-built messages) must produce a
|
||||
/// payload semantically identical to the fragment-assembly path.
|
||||
#[test]
|
||||
fn test_chat_messages_json_path_equivalent_to_fragment_path() {
|
||||
let base = mm_input();
|
||||
let fragment_payload = build_mm_payload(&base);
|
||||
|
||||
let mut chat = base.clone();
|
||||
let mm = chat.multi_modal_content.take().unwrap();
|
||||
let msgs = crate::datasets::random_mm::build_chat_messages_json(&chat.prompt, Some(&mm));
|
||||
chat.chat_messages_json = Some(Arc::from(msgs.as_str()));
|
||||
let chat_payload = build_mm_payload(&chat);
|
||||
|
||||
let a: serde_json::Value = serde_json::from_slice(&fragment_payload).unwrap();
|
||||
let b: serde_json::Value = serde_json::from_slice(&chat_payload).unwrap();
|
||||
assert_eq!(a, b);
|
||||
}
|
||||
|
||||
/// ignore_eos and extra_body must survive the raw-splice path.
|
||||
#[test]
|
||||
fn test_mm_payload_tail_fields() {
|
||||
let mut input = mm_input();
|
||||
input.ignore_eos = true;
|
||||
input.extra_body = Some(serde_json::json!({"temperature": 0.5, "stream": false}));
|
||||
let v: serde_json::Value = serde_json::from_slice(&build_mm_payload(&input)).unwrap();
|
||||
assert_eq!(v["ignore_eos"], true);
|
||||
assert_eq!(v["temperature"], 0.5);
|
||||
// keys already set above must not be overridden by extra_body
|
||||
assert_eq!(v["stream"], true);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::error::Error as StdError;
|
||||
use std::time::Instant;
|
||||
|
||||
use futures::StreamExt;
|
||||
|
||||
use super::streaming::{StreamedResponseHandler, trim_bytes};
|
||||
use super::{CompletionChunk, RequestFuncInput, RequestFuncOutput, build_headers};
|
||||
use crate::error::Result;
|
||||
|
||||
/// Backend for OpenAI-compatible Completions API (/v1/completions).
|
||||
/// Used by "vllm" and "openai" backends.
|
||||
#[derive(Clone)]
|
||||
pub struct OpenAICompletionsBackend;
|
||||
|
||||
impl OpenAICompletionsBackend {
|
||||
pub async fn send_request(
|
||||
&self,
|
||||
input: &RequestFuncInput,
|
||||
client: &reqwest::Client,
|
||||
) -> Result<RequestFuncOutput> {
|
||||
let model = input.model_name.as_deref().unwrap_or(&input.model);
|
||||
|
||||
// When prompt_token_ids are available, send them as the `prompt` value
|
||||
// (JSON array of integers). vLLM's completions API accepts both string
|
||||
// and token ID array as `prompt`, skipping server-side tokenization.
|
||||
let prompt_value = if let Some(ref token_ids) = input.prompt_token_ids {
|
||||
serde_json::json!(token_ids.as_ref())
|
||||
} else {
|
||||
serde_json::json!(input.prompt)
|
||||
};
|
||||
|
||||
let mut payload = serde_json::json!({
|
||||
"model": model,
|
||||
"prompt": prompt_value,
|
||||
"max_tokens": input.output_len,
|
||||
"stream": true,
|
||||
"stream_options": {
|
||||
"include_usage": true,
|
||||
},
|
||||
});
|
||||
|
||||
// Always include logprobs (null when not set) — matches Python which
|
||||
// sends logprobs=None explicitly rather than omitting the key.
|
||||
payload["logprobs"] = match input.logprobs {
|
||||
Some(n) => serde_json::json!(n),
|
||||
None => serde_json::Value::Null,
|
||||
};
|
||||
|
||||
// Apply ignore_eos and extra_body
|
||||
if input.ignore_eos {
|
||||
payload["ignore_eos"] = serde_json::json!(true);
|
||||
}
|
||||
if let Some(serde_json::Value::Object(map)) = input.extra_body.as_ref() {
|
||||
for (k, v) in map {
|
||||
payload[k] = v.clone();
|
||||
}
|
||||
}
|
||||
|
||||
let headers_map = build_headers(None, &input.extra_headers, &input.request_id);
|
||||
|
||||
let mut output = RequestFuncOutput {
|
||||
prompt_len: input.prompt_len,
|
||||
itl: Vec::with_capacity(input.output_len.max(1)),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let st = Instant::now();
|
||||
// start_time is overwritten by benchmark.rs with monotonic offset
|
||||
|
||||
let mut most_recent_timestamp = st;
|
||||
let mut generated_text = String::new();
|
||||
let mut first_chunk_received = false;
|
||||
|
||||
let mut request = client.post(&input.api_url).json(&payload);
|
||||
for (k, v) in &headers_map {
|
||||
request = request.header(k, v);
|
||||
}
|
||||
|
||||
match request.send().await {
|
||||
Ok(response) => {
|
||||
if response.status().is_success() {
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let mut stream = response.bytes_stream();
|
||||
|
||||
while let Some(chunk_result) = stream.next().await {
|
||||
let chunk_bytes = match chunk_result {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
output.success = false;
|
||||
output.error = format!("Stream error: {e}");
|
||||
return Ok(output);
|
||||
}
|
||||
};
|
||||
|
||||
let trimmed_bytes = trim_bytes(&chunk_bytes);
|
||||
if trimmed_bytes.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let messages = handler.add_chunk(trimmed_bytes);
|
||||
for message in messages {
|
||||
// Skip SSE comments
|
||||
if message.starts_with(':') {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Handle multi-field SSE events (e.g., Dynamo sends
|
||||
// "event: message\ndata: {...}"). Extract the data: line.
|
||||
let raw = if message.contains('\n') {
|
||||
match message.lines().find(|l| l.starts_with("data: ")) {
|
||||
Some(l) => l,
|
||||
None => continue,
|
||||
}
|
||||
} else {
|
||||
message.as_str()
|
||||
};
|
||||
|
||||
let chunk = raw.strip_prefix("data: ").unwrap_or(raw);
|
||||
|
||||
if chunk == "[DONE]" {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Typed deserialization — avoids allocating a full
|
||||
// serde_json::Value tree; only extracts needed fields.
|
||||
let data: CompletionChunk = match serde_json::from_str(chunk) {
|
||||
Ok(d) => d,
|
||||
Err(_) => continue,
|
||||
};
|
||||
|
||||
if !data.choices.is_empty() {
|
||||
let text = data.choices[0].text.as_deref().unwrap_or("");
|
||||
|
||||
let timestamp = Instant::now();
|
||||
|
||||
if !first_chunk_received {
|
||||
first_chunk_received = true;
|
||||
output.ttft = timestamp.duration_since(st).as_secs_f64();
|
||||
} else {
|
||||
output.itl.push(
|
||||
timestamp
|
||||
.duration_since(most_recent_timestamp)
|
||||
.as_secs_f64(),
|
||||
);
|
||||
}
|
||||
|
||||
most_recent_timestamp = timestamp;
|
||||
generated_text.push_str(text);
|
||||
}
|
||||
// Separate `if` (not `else if`) — Dynamo may send
|
||||
// both choices and usage in the same chunk.
|
||||
if let Some(ref usage) = data.usage
|
||||
&& let Some(ct) = usage.completion_tokens
|
||||
{
|
||||
output.output_tokens = ct as usize;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if first_chunk_received {
|
||||
output.success = true;
|
||||
} else {
|
||||
output.success = false;
|
||||
output.error = "Never received a valid chunk to calculate TTFT. \
|
||||
This response will be marked as failed!"
|
||||
.to_string();
|
||||
}
|
||||
output.generated_text = generated_text;
|
||||
output.latency = most_recent_timestamp.duration_since(st).as_secs_f64();
|
||||
} else {
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
output.error = if body.is_empty() {
|
||||
format!("HTTP {status}")
|
||||
} else {
|
||||
format!("HTTP {status}: {body}")
|
||||
};
|
||||
output.success = false;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
output.success = false;
|
||||
// Capture full error chain for debugging
|
||||
let mut error_msg = format!("{e}");
|
||||
let mut source = e.source();
|
||||
while let Some(cause) = source {
|
||||
error_msg.push_str(&format!("\n Caused by: {cause}"));
|
||||
source = cause.source();
|
||||
}
|
||||
output.error = error_msg;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(output)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,326 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
//! Pooling/embedding backends: non-streaming HTTP POST for embedding, pooling, and rerank
|
||||
//! endpoints.
|
||||
//!
|
||||
//! Supported variants:
|
||||
//! - `openai-embeddings`: Standard OpenAI `/v1/embeddings` with text input
|
||||
//! - `openai-embeddings-chat`: OpenAI `/v1/embeddings` with chat message format (supports
|
||||
//! multimodal)
|
||||
//! - `vllm-pooling`: vLLM `/v1/pooling` endpoint
|
||||
//! - `vllm-rerank`: vLLM `/v1/rerank` endpoint (query + documents)
|
||||
|
||||
use std::time::Instant;
|
||||
|
||||
use crate::backends::{RequestFuncInput, RequestFuncOutput, build_headers};
|
||||
use crate::cli::BackendKind;
|
||||
use crate::error::Result;
|
||||
|
||||
/// Response from embedding/pooling endpoints (minimal fields for usage extraction).
|
||||
#[derive(serde::Deserialize)]
|
||||
struct PoolingResponse {
|
||||
usage: Option<PoolingUsage>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct PoolingUsage {
|
||||
prompt_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PoolingBackend {
|
||||
pub kind: BackendKind,
|
||||
}
|
||||
|
||||
impl PoolingBackend {
|
||||
pub async fn send_request(
|
||||
&self,
|
||||
input: &RequestFuncInput,
|
||||
client: &reqwest::Client,
|
||||
) -> Result<RequestFuncOutput> {
|
||||
// Preserve client-side prompt_len as fallback if server doesn't report usage.
|
||||
let mut output = RequestFuncOutput {
|
||||
prompt_len: input.prompt_len,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let headers = build_headers(
|
||||
Some("application/json"),
|
||||
&input.extra_headers,
|
||||
&input.request_id,
|
||||
);
|
||||
|
||||
let payload = self.build_payload(input);
|
||||
|
||||
let mut request = client.post(&input.api_url);
|
||||
for (k, v) in &headers {
|
||||
request = request.header(k, v);
|
||||
}
|
||||
|
||||
let st = Instant::now();
|
||||
|
||||
let response = match request.json(&payload).send().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
output.error = format!("Request failed: {e}");
|
||||
return Ok(output);
|
||||
}
|
||||
};
|
||||
|
||||
if response.status().is_success() {
|
||||
let latency = st.elapsed().as_secs_f64();
|
||||
output.latency = latency;
|
||||
output.ttft = latency;
|
||||
output.success = true;
|
||||
|
||||
// Parse usage from response; keep client-side prompt_len as fallback.
|
||||
match response.json::<PoolingResponse>().await {
|
||||
Ok(data) => {
|
||||
if let Some(usage) = data.usage
|
||||
&& let Some(tokens) = usage.prompt_tokens
|
||||
{
|
||||
output.prompt_len = tokens as usize;
|
||||
}
|
||||
}
|
||||
Err(_) => {
|
||||
// Response parsed but no usage — keep client-side prompt_len
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
output.error = format!("HTTP {status}: {body}");
|
||||
}
|
||||
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn build_payload(&self, input: &RequestFuncInput) -> serde_json::Value {
|
||||
let model = input.model_name.as_deref().unwrap_or(&input.model);
|
||||
|
||||
// For "input" field (openai-embeddings, vllm-pooling): a batched request
|
||||
// (--random-batch-size) sends the text list; otherwise prefer prompt_token_ids
|
||||
// when available. The random dataset sets prompt="" and relies on token IDs;
|
||||
// the OpenAI embeddings API accepts both text strings and token ID arrays.
|
||||
// Note: embeddings-chat uses text in messages; vllm-rerank uses text as query.
|
||||
let input_value = if let Some(ref list) = input.prompt_list {
|
||||
serde_json::json!(list.iter().map(|s| s.as_ref()).collect::<Vec<&str>>())
|
||||
} else if let Some(ref token_ids) = input.prompt_token_ids {
|
||||
serde_json::json!(token_ids.as_ref())
|
||||
} else {
|
||||
serde_json::json!(input.prompt.as_ref())
|
||||
};
|
||||
|
||||
let is_vllm_backend = matches!(
|
||||
self.kind,
|
||||
BackendKind::VllmPooling | BackendKind::VllmRerank
|
||||
);
|
||||
|
||||
let mut payload = match self.kind {
|
||||
BackendKind::OpenaiEmbeddings => {
|
||||
let mut p = serde_json::json!({
|
||||
"model": model,
|
||||
"input": input_value,
|
||||
});
|
||||
// truncate_prompt_tokens is vLLM-specific; only include for vLLM backends
|
||||
// to avoid breaking standard OpenAI providers.
|
||||
if is_vllm_backend {
|
||||
p["truncate_prompt_tokens"] = serde_json::json!(-1);
|
||||
}
|
||||
p
|
||||
}
|
||||
BackendKind::OpenaiEmbeddingsChat => {
|
||||
// Chat format: uses text prompt in messages array (for multimodal support).
|
||||
// Python's _get_chat_content always returns a content array.
|
||||
// Use raw string concatenation for multimodal fragments (zero-copy,
|
||||
// avoids re-parsing ~200KB+ base64 per image).
|
||||
let content_json = build_chat_content_json(input);
|
||||
|
||||
let mut p = serde_json::json!({
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": content_json}],
|
||||
});
|
||||
if is_vllm_backend {
|
||||
p["truncate_prompt_tokens"] = serde_json::json!(-1);
|
||||
}
|
||||
p
|
||||
}
|
||||
BackendKind::VllmPooling => {
|
||||
serde_json::json!({
|
||||
"model": model,
|
||||
"input": input_value,
|
||||
"truncate_prompt_tokens": -1,
|
||||
})
|
||||
}
|
||||
BackendKind::VllmRerank => {
|
||||
// random-rerank dataset: prompt_list = [query, doc1, doc2, ...]
|
||||
// (mirrors Python async_request_vllm_rerank).
|
||||
if let Some(ref list) = input.prompt_list {
|
||||
if list.len() < 2 {
|
||||
tracing::warn!(
|
||||
backend = "vllm-rerank",
|
||||
inputs = list.len(),
|
||||
"rerank request has no documents"
|
||||
);
|
||||
}
|
||||
let query = list.first().map(|s| s.as_ref()).unwrap_or("");
|
||||
let documents: Vec<&str> = list.iter().skip(1).map(|s| s.as_ref()).collect();
|
||||
serde_json::json!({
|
||||
"model": model,
|
||||
"query": query,
|
||||
"documents": documents,
|
||||
"truncate_prompt_tokens": -1,
|
||||
})
|
||||
} else {
|
||||
// Legacy path: text prompt as query, documents via --extra-body.
|
||||
let query = input.prompt.as_ref();
|
||||
if query.is_empty() && input.prompt_token_ids.is_some() {
|
||||
tracing::warn!(
|
||||
backend = "vllm-rerank",
|
||||
dataset = "random",
|
||||
"rerank request has an empty query; use the random-rerank dataset"
|
||||
);
|
||||
}
|
||||
serde_json::json!({
|
||||
"model": model,
|
||||
"query": query,
|
||||
"truncate_prompt_tokens": -1,
|
||||
})
|
||||
}
|
||||
}
|
||||
_ => unreachable!("PoolingBackend with non-pooling kind"),
|
||||
};
|
||||
|
||||
// Merge extra_body fields into payload
|
||||
if let Some(ref extra) = input.extra_body
|
||||
&& let (Some(base), Some(extra_obj)) = (payload.as_object_mut(), extra.as_object())
|
||||
{
|
||||
for (k, v) in extra_obj {
|
||||
base.insert(k.clone(), v.clone());
|
||||
}
|
||||
}
|
||||
|
||||
payload
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the chat content JSON array for embeddings-chat.
|
||||
/// Uses raw string concatenation for multimodal fragments to avoid
|
||||
/// re-parsing large base64 image data (matching openai_chat.rs approach).
|
||||
fn build_chat_content_json(input: &RequestFuncInput) -> serde_json::Value {
|
||||
if input.multi_modal_content.is_none() {
|
||||
// Text-only: return content array with single text element
|
||||
return serde_json::json!([{
|
||||
"type": "text",
|
||||
"text": input.prompt.as_ref(),
|
||||
}]);
|
||||
}
|
||||
|
||||
// Multimodal: build JSON string manually for zero-copy fragment embedding
|
||||
let mm = input.multi_modal_content.as_ref().unwrap();
|
||||
let prompt = input.prompt.as_ref();
|
||||
|
||||
let mm_total: usize = mm.iter().map(|f| f.len() + 1).sum();
|
||||
let mut json = String::with_capacity(64 + prompt.len() * 2 + mm_total);
|
||||
|
||||
// [{"type":"text","text":"<prompt>"}
|
||||
json.push_str(r#"[{"type":"text","text":""#);
|
||||
push_json_escaped_str(&mut json, prompt);
|
||||
json.push_str(r#""}"#);
|
||||
|
||||
// ,<mm fragment 1>,<mm fragment 2>,...
|
||||
for fragment in mm.iter() {
|
||||
json.push(',');
|
||||
json.push_str(fragment);
|
||||
}
|
||||
|
||||
json.push(']');
|
||||
|
||||
// Parse the assembled string into a Value for embedding in the payload.
|
||||
// This parse is O(n) but operates on the pre-built string once, not per-fragment.
|
||||
serde_json::from_str(&json).unwrap_or_else(|_| {
|
||||
serde_json::json!([{
|
||||
"type": "text",
|
||||
"text": input.prompt.as_ref(),
|
||||
}])
|
||||
})
|
||||
}
|
||||
|
||||
/// Escape a string for safe JSON embedding (matching openai_chat.rs).
|
||||
fn push_json_escaped_str(buf: &mut String, s: &str) {
|
||||
use std::fmt::Write;
|
||||
for ch in s.chars() {
|
||||
match ch {
|
||||
'"' => buf.push_str(r#"\""#),
|
||||
'\\' => buf.push_str(r"\\"),
|
||||
'\n' => buf.push_str(r"\n"),
|
||||
'\r' => buf.push_str(r"\r"),
|
||||
'\t' => buf.push_str(r"\t"),
|
||||
c if c < '\x20' => {
|
||||
let _ = write!(buf, "\\u{:04x}", c as u32);
|
||||
}
|
||||
c => buf.push(c),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn list(items: &[&str]) -> Option<Arc<[Arc<str>]>> {
|
||||
Some(items.iter().map(|s| Arc::from(*s)).collect())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_embeddings_payload_batched_input() {
|
||||
let backend = PoolingBackend {
|
||||
kind: BackendKind::OpenaiEmbeddings,
|
||||
};
|
||||
let input = RequestFuncInput {
|
||||
model: "bge".to_string(),
|
||||
prompt_list: list(&["t1", "t2", "t3"]),
|
||||
..Default::default()
|
||||
};
|
||||
let payload = backend.build_payload(&input);
|
||||
assert_eq!(payload["input"], serde_json::json!(["t1", "t2", "t3"]));
|
||||
assert_eq!(payload["model"], "bge");
|
||||
// truncate_prompt_tokens is vLLM-specific and deliberately omitted for
|
||||
// the plain OpenAI embeddings backend.
|
||||
assert!(payload.get("truncate_prompt_tokens").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rerank_payload_query_and_documents() {
|
||||
let backend = PoolingBackend {
|
||||
kind: BackendKind::VllmRerank,
|
||||
};
|
||||
let input = RequestFuncInput {
|
||||
model: "reranker".to_string(),
|
||||
prompt_list: list(&["the query", "doc a", "doc b"]),
|
||||
..Default::default()
|
||||
};
|
||||
let payload = backend.build_payload(&input);
|
||||
assert_eq!(payload["query"], "the query");
|
||||
assert_eq!(payload["documents"], serde_json::json!(["doc a", "doc b"]));
|
||||
assert_eq!(payload["truncate_prompt_tokens"], -1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rerank_payload_legacy_single_prompt() {
|
||||
let backend = PoolingBackend {
|
||||
kind: BackendKind::VllmRerank,
|
||||
};
|
||||
let input = RequestFuncInput {
|
||||
model: "reranker".to_string(),
|
||||
prompt: Arc::from("query text"),
|
||||
..Default::default()
|
||||
};
|
||||
let payload = backend.build_payload(&input);
|
||||
assert_eq!(payload["query"], "query text");
|
||||
assert!(payload.get("documents").is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
/// SSE streaming response handler.
|
||||
///
|
||||
/// Accumulates incoming byte chunks and extracts complete SSE messages.
|
||||
/// Mirrors Python's `StreamedResponseHandler` from endpoint_request_func.py:22-60.
|
||||
pub struct StreamedResponseHandler {
|
||||
buffer: String,
|
||||
/// Reusable message buffer — avoids allocating a new Vec per `add_chunk` call.
|
||||
messages: Vec<String>,
|
||||
}
|
||||
|
||||
impl StreamedResponseHandler {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
buffer: String::with_capacity(4096),
|
||||
messages: Vec::with_capacity(4),
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a chunk of bytes and return any complete SSE messages.
|
||||
///
|
||||
/// The returned slice borrows from the handler and is valid until the next
|
||||
/// `add_chunk` call.
|
||||
pub fn add_chunk(&mut self, chunk_bytes: &[u8]) -> &[String] {
|
||||
self.messages.clear();
|
||||
|
||||
let chunk_str = String::from_utf8_lossy(chunk_bytes);
|
||||
self.buffer.push_str(&chunk_str);
|
||||
|
||||
// Split by double newlines (SSE message separator)
|
||||
while let Some(pos) = self.buffer.find("\n\n") {
|
||||
let message = self.buffer[..pos].trim().to_string();
|
||||
// Efficiently remove consumed bytes by shifting remaining data
|
||||
self.buffer.drain(..pos + 2);
|
||||
if !message.is_empty() {
|
||||
self.messages.push(message);
|
||||
}
|
||||
}
|
||||
|
||||
// Handle buffered data without trailing `\n\n`.
|
||||
// Matches Python's speculative json.loads() in StreamedResponseHandler.
|
||||
// This matters for TTFT/ITL accuracy: when a data message and its `\n\n`
|
||||
// arrive in separate TCP segments, we want to emit the message at the
|
||||
// first segment's arrival time, not the second.
|
||||
//
|
||||
// Also handles multi-field SSE events where the buffer may start with
|
||||
// "event: ...\ndata: ..." (Dynamo frontend).
|
||||
let data_start = if self.buffer.starts_with("data: ") {
|
||||
Some(0)
|
||||
} else {
|
||||
// Look for a "data: " line in multi-field events
|
||||
self.buffer.find("\ndata: ").map(|p| p + 1)
|
||||
};
|
||||
if let Some(offset) = data_start {
|
||||
let content = self.buffer[offset + 6..].trim();
|
||||
if content == "[DONE]"
|
||||
|| (!content.is_empty()
|
||||
&& serde_json::from_str::<&serde_json::value::RawValue>(content).is_ok())
|
||||
{
|
||||
self.messages.push(self.buffer.trim().to_string());
|
||||
self.buffer.clear();
|
||||
}
|
||||
}
|
||||
|
||||
&self.messages
|
||||
}
|
||||
}
|
||||
|
||||
/// Trim leading/trailing ASCII whitespace from a byte slice.
|
||||
pub fn trim_bytes(bytes: &[u8]) -> &[u8] {
|
||||
let start = bytes.iter().position(|b| !b.is_ascii_whitespace()).unwrap_or(bytes.len());
|
||||
let end = bytes
|
||||
.iter()
|
||||
.rposition(|b| !b.is_ascii_whitespace())
|
||||
.map(|p| p + 1)
|
||||
.unwrap_or(start);
|
||||
&bytes[start..end]
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_basic_sse() {
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let msgs =
|
||||
handler.add_chunk(b"data: {\"choices\":[{\"text\":\"hi\"}]}\n\ndata: [DONE]\n\n");
|
||||
assert_eq!(msgs.len(), 2);
|
||||
assert!(msgs[0].contains("choices"));
|
||||
assert!(msgs[1].contains("[DONE]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_split_chunks() {
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let msgs1 = handler.add_chunk(b"data: {\"cho");
|
||||
assert!(msgs1.is_empty());
|
||||
let msgs2 = handler.add_chunk(b"ices\":[{\"text\":\"a\"}]}\n\n");
|
||||
assert_eq!(msgs2.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_comment_lines() {
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let msgs = handler.add_chunk(b": ping\n\ndata: {\"test\":1}\n\n");
|
||||
assert_eq!(msgs.len(), 2);
|
||||
assert!(msgs[0].starts_with(":"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_done_without_newlines() {
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let msgs = handler.add_chunk(b"data: [DONE]");
|
||||
assert_eq!(msgs.len(), 1);
|
||||
assert!(msgs[0].contains("[DONE]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_incomplete_json_in_buffer() {
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let msgs = handler.add_chunk(b"data: {\"partial\":");
|
||||
assert!(msgs.is_empty());
|
||||
// Complete JSON without \n\n — speculative parse emits it
|
||||
let msgs2 = handler.add_chunk(b"true}");
|
||||
assert_eq!(msgs2.len(), 1);
|
||||
assert!(msgs2[0].contains("partial"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_multi_field_sse_event() {
|
||||
// Dynamo frontend sends "event: message\ndata: {...}\n\n"
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let msgs =
|
||||
handler.add_chunk(b"event: message\ndata: {\"choices\":[{\"text\":\"hi\"}]}\n\n");
|
||||
assert_eq!(msgs.len(), 1);
|
||||
assert!(msgs[0].contains("choices"));
|
||||
assert!(msgs[0].contains("event: message"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_multi_field_sse_speculative_parse() {
|
||||
// Multi-field event without trailing \n\n — speculative parse should emit it
|
||||
let mut handler = StreamedResponseHandler::new();
|
||||
let msgs = handler.add_chunk(b"event: message\ndata: {\"choices\":[{\"text\":\"hi\"}]}");
|
||||
assert_eq!(msgs.len(), 1);
|
||||
assert!(msgs[0].contains("choices"));
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,743 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::fmt;
|
||||
|
||||
/// Backend type for the benchmark endpoint.
|
||||
#[derive(clap::ValueEnum, Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum BackendKind {
|
||||
#[value(name = "vllm")]
|
||||
Vllm,
|
||||
#[value(name = "openai")]
|
||||
Openai,
|
||||
#[value(name = "openai-chat")]
|
||||
OpenaiChat,
|
||||
#[value(name = "openai-embeddings")]
|
||||
OpenaiEmbeddings,
|
||||
#[value(name = "openai-embeddings-chat")]
|
||||
OpenaiEmbeddingsChat,
|
||||
#[value(name = "vllm-pooling")]
|
||||
VllmPooling,
|
||||
#[value(name = "vllm-rerank")]
|
||||
VllmRerank,
|
||||
}
|
||||
|
||||
impl BackendKind {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Vllm => "vllm",
|
||||
Self::Openai => "openai",
|
||||
Self::OpenaiChat => "openai-chat",
|
||||
Self::OpenaiEmbeddings => "openai-embeddings",
|
||||
Self::OpenaiEmbeddingsChat => "openai-embeddings-chat",
|
||||
Self::VllmPooling => "vllm-pooling",
|
||||
Self::VllmRerank => "vllm-rerank",
|
||||
}
|
||||
}
|
||||
|
||||
/// Return true if the backend is compatible with OpenAI-style API and sampling parameters.
|
||||
pub fn is_openai_compatible(self) -> bool {
|
||||
match self {
|
||||
Self::Vllm | Self::Openai | Self::OpenaiChat => true,
|
||||
Self::OpenaiEmbeddings
|
||||
| Self::OpenaiEmbeddingsChat
|
||||
| Self::VllmPooling
|
||||
| Self::VllmRerank => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Return true if the backend is a pooling/embedding backend (non-generative).
|
||||
pub fn is_pooling(self) -> bool {
|
||||
matches!(
|
||||
self,
|
||||
Self::OpenaiEmbeddings
|
||||
| Self::OpenaiEmbeddingsChat
|
||||
| Self::VllmPooling
|
||||
| Self::VllmRerank
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for BackendKind {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
/// Dataset to benchmark with.
|
||||
#[derive(clap::ValueEnum, Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum DatasetName {
|
||||
#[value(name = "random")]
|
||||
Random,
|
||||
#[value(name = "random-mm")]
|
||||
RandomMm,
|
||||
#[value(name = "sharegpt")]
|
||||
ShareGpt,
|
||||
#[value(name = "sonnet")]
|
||||
Sonnet,
|
||||
#[value(name = "speed-bench", alias = "speed_bench")]
|
||||
SpeedBench,
|
||||
#[value(name = "hf")]
|
||||
Hf,
|
||||
#[value(name = "custom")]
|
||||
Custom,
|
||||
#[value(name = "prefix_repetition", alias = "prefix-repetition")]
|
||||
PrefixRepetition,
|
||||
#[value(name = "random-rerank")]
|
||||
RandomRerank,
|
||||
}
|
||||
|
||||
/// Ramp-up strategy for request rate.
|
||||
#[derive(clap::ValueEnum, Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum RampUpStrategy {
|
||||
#[value(name = "linear")]
|
||||
Linear,
|
||||
#[value(name = "exponential")]
|
||||
Exponential,
|
||||
}
|
||||
|
||||
/// Strategy for assigning LoRA modules to requests.
|
||||
#[derive(clap::ValueEnum, Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum LoraAssignment {
|
||||
#[value(name = "random")]
|
||||
Random,
|
||||
#[value(name = "round-robin")]
|
||||
RoundRobin,
|
||||
}
|
||||
|
||||
/// SPEED-Bench dataset split/config.
|
||||
#[derive(clap::ValueEnum, Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SpeedBenchConfig {
|
||||
#[value(name = "qualitative")]
|
||||
Qualitative,
|
||||
#[value(name = "throughput_1k")]
|
||||
Throughput1k,
|
||||
#[value(name = "throughput_2k")]
|
||||
Throughput2k,
|
||||
#[value(name = "throughput_8k")]
|
||||
Throughput8k,
|
||||
#[value(name = "throughput_16k")]
|
||||
Throughput16k,
|
||||
#[value(name = "throughput_32k")]
|
||||
Throughput32k,
|
||||
}
|
||||
|
||||
impl SpeedBenchConfig {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Qualitative => "qualitative",
|
||||
Self::Throughput1k => "throughput_1k",
|
||||
Self::Throughput2k => "throughput_2k",
|
||||
Self::Throughput8k => "throughput_8k",
|
||||
Self::Throughput16k => "throughput_16k",
|
||||
Self::Throughput32k => "throughput_32k",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for SpeedBenchConfig {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
/// High-performance benchmark client for vLLM serving endpoints.
|
||||
#[derive(clap::Args, Debug, Clone)]
|
||||
pub struct BenchServeArgs {
|
||||
/// The type of backend or endpoint to use for the benchmark.
|
||||
#[arg(long, default_value = "openai")]
|
||||
pub backend: BackendKind,
|
||||
|
||||
/// Server or API base url if not using http host and port.
|
||||
#[arg(long)]
|
||||
pub base_url: Option<String>,
|
||||
|
||||
/// Server host.
|
||||
#[arg(long, default_value = "127.0.0.1")]
|
||||
pub host: String,
|
||||
|
||||
/// Server port.
|
||||
#[arg(long, default_value_t = 8000)]
|
||||
pub port: u16,
|
||||
|
||||
/// API endpoint. Auto-selected based on --backend if not specified.
|
||||
#[arg(long)]
|
||||
pub endpoint: Option<String>,
|
||||
|
||||
/// Name of the model. If not specified, will fetch from server.
|
||||
#[arg(long)]
|
||||
pub model: Option<String>,
|
||||
|
||||
/// The model name used in the API (for --served-model-name).
|
||||
#[arg(long)]
|
||||
pub served_model_name: Option<String>,
|
||||
|
||||
/// Name or path of the tokenizer.
|
||||
#[arg(long)]
|
||||
pub tokenizer: Option<String>,
|
||||
|
||||
/// Tokenizer mode (auto, hf, slow, mistral).
|
||||
#[arg(long, default_value = "auto")]
|
||||
pub tokenizer_mode: String,
|
||||
|
||||
/// Skip initialization of tokenizer.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub skip_tokenizer_init: bool,
|
||||
|
||||
/// Trust remote code for tokenizer.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub trust_remote_code: bool,
|
||||
|
||||
/// Dataset name.
|
||||
#[arg(long, default_value = "random")]
|
||||
pub dataset_name: DatasetName,
|
||||
|
||||
/// General input length for datasets.
|
||||
#[arg(long)]
|
||||
pub input_len: Option<usize>,
|
||||
|
||||
/// General output length for datasets.
|
||||
#[arg(long)]
|
||||
pub output_len: Option<usize>,
|
||||
|
||||
/// Maximum model context length. Requests with prompt_len + output_len above this are filtered
|
||||
/// out.
|
||||
#[arg(long)]
|
||||
pub max_model_len: Option<usize>,
|
||||
|
||||
/// Random dataset input length.
|
||||
#[arg(long, default_value_t = 1024)]
|
||||
pub random_input_len: usize,
|
||||
|
||||
/// Random dataset output length.
|
||||
#[arg(long, default_value_t = 128)]
|
||||
pub random_output_len: usize,
|
||||
|
||||
/// Random dataset prefix length.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub random_prefix_len: usize,
|
||||
|
||||
/// Per-turn input length for turns 1+ in multi-turn mode.
|
||||
/// 0 = fallback to --random-input-len for all turns.
|
||||
/// Mirrors sglang bench_multiturn.py --sub-question-input-length.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub per_turn_input_len: usize,
|
||||
|
||||
/// Range ratio for sampling input/output lengths, matching Python
|
||||
/// `vllm bench serve`: lengths are drawn uniformly from
|
||||
/// [len*(1-r), len*(1+r)]. 0.0 (the default) = exact target lengths.
|
||||
/// Accepts a single float in [0, 1) or a JSON object
|
||||
/// '{"input": r1, "output": r2}' for independent control.
|
||||
/// NOTE: semantics changed — the old Rust-only form sampled [len*r, len]
|
||||
/// with default 1.0; old values like 1.0 are now rejected.
|
||||
#[arg(long, default_value = "0.0")]
|
||||
pub random_range_ratio: String,
|
||||
|
||||
/// Batch multiple generated inputs into one request (embeddings/pooling
|
||||
/// backends only). E.g. 8 sends "input": [t1..t8] per request. Mirrors
|
||||
/// Python --random-batch-size. Default 1 = no batching.
|
||||
#[arg(long, default_value_t = 1)]
|
||||
pub random_batch_size: usize,
|
||||
|
||||
/// random-rerank: the served model is NOT a reranker (embedding-based
|
||||
/// scoring). Changes query/document length accounting to mirror Python
|
||||
/// --no-reranker.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub no_reranker: bool,
|
||||
|
||||
/// Bimodal prefix-cache (random dataset): fraction of prompts that are "warm"
|
||||
/// and reuse a shared cached prefix. 0.0 = off (default). E.g. 0.8 = 80% warm
|
||||
/// (prefix-cache hit), 20% cold (full prefill). Requires --random-cache-ratio > 0
|
||||
/// and --prompt-token-ids. In this mode --random-input-len is the TOTAL length.
|
||||
#[arg(long, default_value_t = 0.0)]
|
||||
pub random_cache_hit_fraction: f64,
|
||||
|
||||
/// Bimodal prefix-cache (random dataset): fraction of each WARM prompt's length
|
||||
/// that is the shared cached prefix. 0.0 = off (default). E.g. 0.95 = 95% cached,
|
||||
/// 5% unique suffix. Used with --random-cache-hit-fraction.
|
||||
#[arg(long, default_value_t = 0.0)]
|
||||
pub random_cache_ratio: f64,
|
||||
|
||||
/// Send prompt as token ID arrays instead of text strings.
|
||||
/// By default, prompts are decoded to text for maximum
|
||||
/// compatibility. Enable this for pure vLLM deployments to skip server-side
|
||||
/// tokenization (faster, exact token counts).
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub prompt_token_ids: bool,
|
||||
|
||||
// --- Random multimodal dataset ---
|
||||
/// Base number of multimodal items (images/videos) per request.
|
||||
#[arg(long, default_value_t = 1)]
|
||||
pub random_mm_base_items_per_request: usize,
|
||||
|
||||
/// Range ratio for varying the number of multimodal items per request.
|
||||
/// Items sampled from [floor(n*(1-r)), ceil(n*(1+r))].
|
||||
#[arg(long, default_value_t = 0.0)]
|
||||
pub random_mm_num_mm_items_range_ratio: f64,
|
||||
|
||||
/// Per-modality hard caps as JSON, e.g. '{"image": 3, "video": 0}'.
|
||||
#[arg(long, default_value = "{\"image\": 255, \"video\": 1}")]
|
||||
pub random_mm_limit_mm_per_prompt: String,
|
||||
|
||||
/// Bucket config mapping (height,width,num_frames) to probability.
|
||||
/// Uses Python-style syntax: '{(256,256,1): 0.5, (720,1280,1): 0.5}'.
|
||||
/// num_frames=1 means image, num_frames>1 means video.
|
||||
#[arg(long, default_value = "{(256,256,1): 0.5, (720,1280,1): 0.5}")]
|
||||
pub random_mm_bucket_config: String,
|
||||
|
||||
/// Enable multimodal chat transformation for datasets that support it.
|
||||
/// The dataset pre-builds the OpenAI chat `messages` array (text part +
|
||||
/// multimodal items) at generation time, and the request sends it verbatim.
|
||||
/// Mirrors Python's --enable-multimodal-chat. Currently applies to random-mm.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub enable_multimodal_chat: bool,
|
||||
|
||||
// --- Custom dataset (JSONL) ---
|
||||
/// Output tokens per request for the custom dataset. Set to -1 to use the
|
||||
/// per-line "output_tokens" field from the JSONL file instead.
|
||||
#[arg(long, default_value_t = 256, allow_negative_numbers = true)]
|
||||
pub custom_output_len: i64,
|
||||
|
||||
/// Skip applying a chat template to custom dataset prompts.
|
||||
/// NOTE: the Rust client never renders chat templates client-side, so this
|
||||
/// is always effectively on; passing it silences the informational notice.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub skip_chat_template: bool,
|
||||
|
||||
// --- Prefix repetition dataset ---
|
||||
/// Shared-prefix token length for the prefix_repetition dataset.
|
||||
#[arg(long, default_value_t = 256)]
|
||||
pub prefix_repetition_prefix_len: usize,
|
||||
|
||||
/// Per-request random suffix token length for the prefix_repetition dataset.
|
||||
#[arg(long, default_value_t = 256)]
|
||||
pub prefix_repetition_suffix_len: usize,
|
||||
|
||||
/// Number of distinct shared prefixes for the prefix_repetition dataset.
|
||||
/// Requests are split evenly across prefixes (num-prompts / num-prefixes each).
|
||||
#[arg(long, default_value_t = 10)]
|
||||
pub prefix_repetition_num_prefixes: usize,
|
||||
|
||||
/// Output tokens per request for the prefix_repetition dataset.
|
||||
#[arg(long, default_value_t = 128)]
|
||||
pub prefix_repetition_output_len: usize,
|
||||
|
||||
/// Number of prompts to generate.
|
||||
#[arg(long, default_value_t = 1000)]
|
||||
pub num_prompts: usize,
|
||||
|
||||
/// Number of requests per second. Use "inf" for all at once.
|
||||
#[arg(long, default_value_t = f64::INFINITY)]
|
||||
pub request_rate: f64,
|
||||
|
||||
/// Burstiness factor of request generation.
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
pub burstiness: f64,
|
||||
|
||||
/// Maximum number of concurrent requests.
|
||||
#[arg(long)]
|
||||
pub max_concurrency: Option<usize>,
|
||||
|
||||
/// Fraction of --max-concurrency at which the steady-state window opens.
|
||||
/// Range: (0.0, 1.0]. Used only when --max-concurrency is set and
|
||||
/// --request-rate is inf.
|
||||
#[arg(long, default_value_t = 0.95)]
|
||||
pub steady_state_threshold: f64,
|
||||
|
||||
/// Minimum steady-state window duration in seconds. Below this, a warning
|
||||
/// is attached. If unset, computed as max(10.0, 0.1 * run_duration).
|
||||
#[arg(long)]
|
||||
pub steady_state_min_window: Option<f64>,
|
||||
|
||||
/// Disable steady-state metrics computation entirely.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub no_steady_state: bool,
|
||||
|
||||
/// Disable tqdm progress bar.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub disable_tqdm: bool,
|
||||
|
||||
/// Number of warmup requests.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub num_warmups: usize,
|
||||
|
||||
/// Use vLLM profiling. --profiler-config must be provided on the server.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub profile: bool,
|
||||
|
||||
/// Minimum server batch size (num_requests_running) before starting the
|
||||
/// profiler. When set, profiling is deferred until the /metrics endpoint
|
||||
/// reports at least this many running requests, then captures for
|
||||
/// --profile-duration seconds. Requires --profile.
|
||||
#[arg(long)]
|
||||
pub profile_batch_threshold: Option<usize>,
|
||||
|
||||
/// How many seconds to capture once the batch threshold is reached.
|
||||
/// Defaults to 5. Requires --profile and --profile-batch-threshold.
|
||||
#[arg(long, default_value_t = 5.0)]
|
||||
pub profile_duration: f64,
|
||||
|
||||
/// Save benchmark results to a JSON file.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub save_result: bool,
|
||||
|
||||
/// Save detailed per-request results.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub save_detailed: bool,
|
||||
|
||||
/// Directory to save benchmark JSON results.
|
||||
#[arg(long)]
|
||||
pub result_dir: Option<String>,
|
||||
|
||||
/// Filename to save benchmark JSON results.
|
||||
#[arg(long)]
|
||||
pub result_filename: Option<String>,
|
||||
|
||||
/// Random seed.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub seed: u64,
|
||||
|
||||
/// Set ignore_eos flag when sending the benchmark request.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub ignore_eos: bool,
|
||||
|
||||
/// Comma-separated list of metrics to report percentiles for.
|
||||
#[arg(long)]
|
||||
pub percentile_metrics: Option<String>,
|
||||
|
||||
/// Comma-separated list of percentiles for selected metrics.
|
||||
#[arg(long, default_value = "99")]
|
||||
pub metric_percentiles: String,
|
||||
|
||||
/// Comma-separated list of extra percentiles to show in sweep summaries.
|
||||
#[arg(long)]
|
||||
pub sweep_summary_percentiles: Option<String>,
|
||||
|
||||
/// The label (prefix) of the benchmark results.
|
||||
#[arg(long)]
|
||||
pub label: Option<String>,
|
||||
|
||||
/// Number of logprobs-per-token to compute.
|
||||
#[arg(long)]
|
||||
pub logprobs: Option<usize>,
|
||||
|
||||
/// Prefix for request IDs.
|
||||
#[arg(long)]
|
||||
pub request_id_prefix: Option<String>,
|
||||
|
||||
/// Maximum time to wait for endpoint readiness in seconds.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub ready_check_timeout_sec: u64,
|
||||
|
||||
/// Key-value pairs for extra headers (KEY=VALUE).
|
||||
#[arg(long = "header", num_args = 1..)]
|
||||
pub headers: Option<Vec<String>>,
|
||||
|
||||
/// JSON string for extra body parameters.
|
||||
#[arg(long)]
|
||||
pub extra_body: Option<String>,
|
||||
|
||||
/// Key-value pairs for metadata (KEY=VALUE).
|
||||
#[arg(long = "metadata", num_args = 1..)]
|
||||
pub metadata: Option<Vec<String>>,
|
||||
|
||||
/// Dry run: only generate dataset and print stats, don't benchmark.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub dry_run: bool,
|
||||
|
||||
// --- Sampling parameters ---
|
||||
/// Top-p sampling parameter. Only affects openai-compatible backends.
|
||||
#[arg(long)]
|
||||
pub top_p: Option<f64>,
|
||||
|
||||
/// Top-k sampling parameter. Only affects openai-compatible backends.
|
||||
#[arg(long)]
|
||||
pub top_k: Option<i64>,
|
||||
|
||||
/// Min-p sampling parameter. Only affects openai-compatible backends.
|
||||
#[arg(long)]
|
||||
pub min_p: Option<f64>,
|
||||
|
||||
/// Temperature sampling parameter. Only affects openai-compatible backends.
|
||||
#[arg(long)]
|
||||
pub temperature: Option<f64>,
|
||||
|
||||
/// Frequency penalty sampling parameter. Only affects openai-compatible backends.
|
||||
#[arg(long)]
|
||||
pub frequency_penalty: Option<f64>,
|
||||
|
||||
/// Presence penalty sampling parameter. Only affects openai-compatible backends.
|
||||
#[arg(long)]
|
||||
pub presence_penalty: Option<f64>,
|
||||
|
||||
/// Repetition penalty sampling parameter. Only affects openai-compatible backends.
|
||||
#[arg(long)]
|
||||
pub repetition_penalty: Option<f64>,
|
||||
|
||||
// --- SSL ---
|
||||
/// Disable SSL certificate verification.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub insecure: bool,
|
||||
|
||||
// --- Ramp-up ---
|
||||
/// Ramp-up strategy for request rate (linear or exponential).
|
||||
#[arg(long)]
|
||||
pub ramp_up_strategy: Option<RampUpStrategy>,
|
||||
|
||||
/// Starting request rate for ramp-up (RPS).
|
||||
#[arg(long)]
|
||||
pub ramp_up_start_rps: Option<f64>,
|
||||
|
||||
/// Ending request rate for ramp-up (RPS).
|
||||
#[arg(long)]
|
||||
pub ramp_up_end_rps: Option<f64>,
|
||||
|
||||
// --- Goodput ---
|
||||
/// Service level objectives for goodput as "KEY:VALUE" pairs (e.g. ttft:100 tpot:50 e2el:500).
|
||||
/// Values are in milliseconds.
|
||||
#[arg(long = "goodput", num_args = 1..)]
|
||||
pub goodput: Option<Vec<String>>,
|
||||
|
||||
// --- Result ---
|
||||
/// Append the benchmark result to the existing JSON file.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub append_result: bool,
|
||||
|
||||
// --- ShareGPT dataset ---
|
||||
/// Path to dataset file (required for sharegpt dataset).
|
||||
#[arg(long)]
|
||||
pub dataset_path: Option<String>,
|
||||
|
||||
/// Override output length for ShareGPT dataset.
|
||||
#[arg(long)]
|
||||
pub sharegpt_output_len: Option<usize>,
|
||||
|
||||
/// Do not oversample if dataset is smaller than num_prompts.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub no_oversample: bool,
|
||||
|
||||
/// Do not shuffle the dataset.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub disable_shuffle: bool,
|
||||
|
||||
// --- Sonnet dataset ---
|
||||
/// Number of input tokens per request (sonnet dataset).
|
||||
#[arg(long, default_value_t = crate::datasets::sonnet::DEFAULT_INPUT_LEN)]
|
||||
pub sonnet_input_len: usize,
|
||||
|
||||
/// Number of output tokens per request (sonnet dataset).
|
||||
#[arg(long, default_value_t = crate::datasets::sonnet::DEFAULT_OUTPUT_LEN)]
|
||||
pub sonnet_output_len: usize,
|
||||
|
||||
/// Number of prefix tokens shared across requests (sonnet dataset).
|
||||
#[arg(long, default_value_t = crate::datasets::sonnet::DEFAULT_PREFIX_LEN)]
|
||||
pub sonnet_prefix_len: usize,
|
||||
|
||||
/// SPEED-Bench config/split (qualitative, throughput_1k, throughput_2k, throughput_8k,
|
||||
/// throughput_16k, throughput_32k).
|
||||
#[arg(long, default_value = "qualitative")]
|
||||
pub speed_bench_config: SpeedBenchConfig,
|
||||
|
||||
/// Filter SPEED-Bench by category (e.g. low_entropy, high_entropy, coding, math).
|
||||
#[arg(long)]
|
||||
pub speed_bench_category: Option<String>,
|
||||
|
||||
/// Truncate SPEED-Bench prompts to at most this many tokens.
|
||||
/// Useful for creating custom input lengths from larger splits (e.g. --speed-bench-config
|
||||
/// throughput_16k --speed-bench-max-input-len 10240).
|
||||
#[arg(long)]
|
||||
pub speed_bench_max_input_len: Option<usize>,
|
||||
|
||||
// --- HuggingFace dataset ---
|
||||
/// HuggingFace dataset split (e.g. train, test, validation).
|
||||
#[arg(long)]
|
||||
pub hf_split: Option<String>,
|
||||
|
||||
/// HuggingFace dataset subset/config name.
|
||||
#[arg(long)]
|
||||
pub hf_subset: Option<String>,
|
||||
|
||||
/// Fixed output length for HF dataset requests (overrides dataset-derived length).
|
||||
#[arg(long)]
|
||||
pub hf_output_len: Option<usize>,
|
||||
|
||||
/// Column name containing the prompt text. Auto-detected if not specified.
|
||||
#[arg(long)]
|
||||
pub hf_text_column: Option<String>,
|
||||
|
||||
// --- Compare mode ---
|
||||
/// Compare two benchmark result JSON files (e.g. --compare a.json b.json).
|
||||
/// Prints side-by-side metrics with delta and % change. Skips benchmarking.
|
||||
#[arg(long = "compare", num_args = 2, value_names = ["FILE_A", "FILE_B"])]
|
||||
pub compare: Option<Vec<String>>,
|
||||
|
||||
// --- Sweep mode ---
|
||||
/// Sweep over max-concurrency values (comma-separated, e.g. --sweep-max-concurrency
|
||||
/// 1,10,50,100,500).
|
||||
#[arg(long)]
|
||||
pub sweep_max_concurrency: Option<String>,
|
||||
|
||||
/// When sweeping concurrency, set num_prompts = concurrency * this factor for each sweep
|
||||
/// point.
|
||||
#[arg(long)]
|
||||
pub sweep_num_prompts_factor: Option<usize>,
|
||||
|
||||
/// Sweep over request-rate values (comma-separated, supports "inf", e.g. --sweep-request-rate
|
||||
/// 1,10,100,inf).
|
||||
#[arg(long)]
|
||||
pub sweep_request_rate: Option<String>,
|
||||
|
||||
/// Reset the server's prefix cache before each sweep iteration.
|
||||
/// Requires VLLM_SERVER_DEV_MODE=1 on the vLLM server.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub reset_prefix_cache: bool,
|
||||
|
||||
// --- Multi-run ---
|
||||
/// Number of benchmark runs for statistical aggregation.
|
||||
#[arg(long, default_value_t = 1)]
|
||||
pub num_runs: usize,
|
||||
|
||||
// --- Multi-turn conversation benchmark ---
|
||||
/// Enable multi-turn conversation benchmark mode.
|
||||
#[arg(long, default_value_t = false)]
|
||||
pub multi_turn: bool,
|
||||
|
||||
/// Number of turns per conversation in synthetic multi-turn mode.
|
||||
#[arg(long, default_value_t = 3)]
|
||||
pub multi_turn_num_turns: usize,
|
||||
|
||||
/// Minimum turns per conversation. 0 = use --multi-turn-num-turns.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub multi_turn_min_turns: usize,
|
||||
|
||||
/// Maximum turns per conversation.
|
||||
/// For synthetic multi-turn, 0 = use --multi-turn-num-turns.
|
||||
/// For ShareGPT multi-turn, 0 = uncapped.
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub multi_turn_max_turns: usize,
|
||||
|
||||
/// Number of concurrent conversations (defaults to max-concurrency or num-prompts).
|
||||
#[arg(long)]
|
||||
pub multi_turn_concurrency: Option<usize>,
|
||||
|
||||
/// Delay between turns in milliseconds (simulates user think time).
|
||||
#[arg(long, default_value_t = 0)]
|
||||
pub multi_turn_delay_ms: u64,
|
||||
|
||||
/// Fraction of per-turn input tokens shared across ALL conversations (0.0–1.0).
|
||||
/// When > 0, enables prefix sharing mode: each turn sends a fixed-length message
|
||||
/// (no history accumulation). Only works with --dataset-name random.
|
||||
#[arg(long, default_value_t = 0.0)]
|
||||
pub multi_turn_prefix_global_ratio: f64,
|
||||
|
||||
/// Fraction of per-turn input tokens shared within each conversation (0.0–1.0).
|
||||
/// When > 0, enables prefix sharing mode: each turn sends a fixed-length message
|
||||
/// (no history accumulation). Only works with --dataset-name random.
|
||||
#[arg(long, default_value_t = 0.0)]
|
||||
pub multi_turn_prefix_conversation_ratio: f64,
|
||||
|
||||
// --- LoRA ---
|
||||
/// LoRA adapter names registered on the server (server-side
|
||||
/// `--lora-modules name=path`). Each request's `model` field is rewritten
|
||||
/// to one of these names; tokenizer and other endpoints keep using --model.
|
||||
/// In multi-turn mode, one adapter is assigned per conversation (sticky
|
||||
/// across turns).
|
||||
#[arg(long = "lora-modules", num_args = 1..)]
|
||||
pub lora_modules: Option<Vec<String>>,
|
||||
|
||||
/// Strategy for assigning LoRA adapters to requests.
|
||||
/// 'random' (default) picks uniformly at random; 'round-robin' cycles
|
||||
/// through `--lora-modules` deterministically (i % N).
|
||||
#[arg(long = "lora-assignment", default_value = "random")]
|
||||
pub lora_assignment: LoraAssignment,
|
||||
}
|
||||
|
||||
impl BenchServeArgs {
|
||||
/// Resolve the base URL from explicit --base-url or from --host/--port.
|
||||
pub fn resolve_base_url(&self) -> String {
|
||||
if let Some(ref base) = self.base_url {
|
||||
base.clone()
|
||||
} else {
|
||||
format!("http://{}:{}", self.host, self.port)
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve the API endpoint, auto-selecting based on backend if not explicit.
|
||||
pub fn resolve_endpoint(&self) -> String {
|
||||
if let Some(ref ep) = self.endpoint {
|
||||
return ep.clone();
|
||||
}
|
||||
match self.backend {
|
||||
BackendKind::OpenaiChat => "/v1/chat/completions".to_string(),
|
||||
BackendKind::Vllm | BackendKind::Openai => "/v1/completions".to_string(),
|
||||
BackendKind::OpenaiEmbeddings | BackendKind::OpenaiEmbeddingsChat => {
|
||||
"/v1/embeddings".to_string()
|
||||
}
|
||||
BackendKind::VllmPooling => "/v1/pooling".to_string(),
|
||||
BackendKind::VllmRerank => "/v1/rerank".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve the full API URL.
|
||||
pub fn resolve_api_url(&self) -> String {
|
||||
format!("{}{}", self.resolve_base_url(), self.resolve_endpoint())
|
||||
}
|
||||
|
||||
/// Parse extra headers from KEY=VALUE pairs.
|
||||
pub fn parse_headers(
|
||||
&self,
|
||||
) -> crate::error::Result<Option<std::collections::HashMap<String, String>>> {
|
||||
match &self.headers {
|
||||
None => Ok(None),
|
||||
Some(items) => {
|
||||
let mut map = std::collections::HashMap::new();
|
||||
for item in items {
|
||||
let (k, v) = item.split_once('=').ok_or_else(|| {
|
||||
crate::error::BenchError::Config(
|
||||
"Invalid header format. Use KEY=VALUE".into(),
|
||||
)
|
||||
})?;
|
||||
map.insert(k.trim().to_string(), v.trim().to_string());
|
||||
}
|
||||
Ok(Some(map))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse extra body JSON.
|
||||
pub fn parse_extra_body(&self) -> crate::error::Result<Option<serde_json::Value>> {
|
||||
match &self.extra_body {
|
||||
None => Ok(None),
|
||||
Some(s) => {
|
||||
let v: serde_json::Value = serde_json::from_str(s).map_err(|e| {
|
||||
crate::error::BenchError::Config(format!("Invalid --extra-body JSON: {e}"))
|
||||
})?;
|
||||
Ok(Some(v))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Generate the request ID prefix (auto-generate if not provided).
|
||||
pub fn get_request_id_prefix(&self) -> String {
|
||||
self.request_id_prefix
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("bench-{}-", &uuid::Uuid::new_v4().to_string()[..8]))
|
||||
}
|
||||
|
||||
/// Resolve input/output lengths, applying --input-len/--output-len overrides.
|
||||
pub fn resolved_random_input_len(&self) -> usize {
|
||||
self.input_len.unwrap_or(self.random_input_len)
|
||||
}
|
||||
|
||||
pub fn resolved_random_output_len(&self) -> usize {
|
||||
self.output_len.unwrap_or(self.random_output_len)
|
||||
}
|
||||
|
||||
pub fn resolved_per_turn_input_len(&self) -> usize {
|
||||
if self.per_turn_input_len > 0 {
|
||||
self.per_turn_input_len
|
||||
} else {
|
||||
self.resolved_random_input_len()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,302 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use crate::error::{BenchError, Result};
|
||||
|
||||
/// Metric definition for comparison: name, JSON key, and whether lower is better.
|
||||
struct MetricDef {
|
||||
label: &'static str,
|
||||
key: &'static str,
|
||||
lower_is_better: bool,
|
||||
}
|
||||
|
||||
const METRICS: &[MetricDef] = &[
|
||||
MetricDef {
|
||||
label: "Request throughput (req/s)",
|
||||
key: "request_throughput",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Output throughput (tok/s)",
|
||||
key: "output_throughput",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Total token throughput (tok/s)",
|
||||
key: "total_token_throughput",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Peak output tokens/s",
|
||||
key: "max_output_tokens_per_s",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Peak concurrent requests",
|
||||
key: "max_concurrent_requests",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Mean TTFT (ms)",
|
||||
key: "mean_ttft_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Median TTFT (ms)",
|
||||
key: "median_ttft_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "P99 TTFT (ms)",
|
||||
key: "p99_ttft_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Mean TPOT (ms)",
|
||||
key: "mean_tpot_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Median TPOT (ms)",
|
||||
key: "median_tpot_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "P99 TPOT (ms)",
|
||||
key: "p99_tpot_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Mean ITL (ms)",
|
||||
key: "mean_itl_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Median ITL (ms)",
|
||||
key: "median_itl_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "P99 ITL (ms)",
|
||||
key: "p99_itl_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Mean E2EL (ms)",
|
||||
key: "mean_e2el_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Median E2EL (ms)",
|
||||
key: "median_e2el_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "P99 E2EL (ms)",
|
||||
key: "p99_e2el_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Completed requests",
|
||||
key: "completed",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Failed requests",
|
||||
key: "failed",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "Duration (s)",
|
||||
key: "duration",
|
||||
lower_is_better: true,
|
||||
},
|
||||
];
|
||||
|
||||
const STEADY_STATE_METRICS: &[MetricDef] = &[
|
||||
MetricDef {
|
||||
label: "SS Request throughput (req/s)",
|
||||
key: "request_throughput",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS Output throughput (tok/s)",
|
||||
key: "output_throughput",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS Input throughput (tok/s)",
|
||||
key: "input_throughput",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS Total token throughput (tok/s)",
|
||||
key: "total_token_throughput",
|
||||
lower_is_better: false,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS Mean TTFT (ms)",
|
||||
key: "mean_ttft_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS Median TTFT (ms)",
|
||||
key: "median_ttft_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS Mean TPOT (ms)",
|
||||
key: "mean_tpot_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS Median TPOT (ms)",
|
||||
key: "median_tpot_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS P90 TPOT (ms)",
|
||||
key: "p90_tpot_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
MetricDef {
|
||||
label: "SS P99 TPOT (ms)",
|
||||
key: "p99_tpot_ms",
|
||||
lower_is_better: true,
|
||||
},
|
||||
];
|
||||
|
||||
/// Compare two benchmark result JSON files and print a side-by-side table.
|
||||
pub fn compare_results(file_a: &str, file_b: &str) -> Result<()> {
|
||||
let json_a = load_result_json(file_a)?;
|
||||
let json_b = load_result_json(file_b)?;
|
||||
|
||||
// Print header with file context
|
||||
let model_a = json_a.get("model_id").and_then(|v| v.as_str()).unwrap_or("?");
|
||||
let model_b = json_b.get("model_id").and_then(|v| v.as_str()).unwrap_or("?");
|
||||
let date_a = json_a.get("date").and_then(|v| v.as_str()).unwrap_or("?");
|
||||
let date_b = json_b.get("date").and_then(|v| v.as_str()).unwrap_or("?");
|
||||
|
||||
println!("{:=^90}", " Benchmark Comparison ");
|
||||
println!(" A: {} (model: {}, date: {})", file_a, model_a, date_a);
|
||||
println!(" B: {} (model: {}, date: {})", file_b, model_b, date_b);
|
||||
println!();
|
||||
|
||||
// Print comparison table
|
||||
println!(
|
||||
"{:<35} {:>12} {:>12} {:>10} {:>8}",
|
||||
"Metric", "A", "B", "Delta", "Change"
|
||||
);
|
||||
println!("{:-<35} {:->12} {:->12} {:->10} {:->8}", "", "", "", "", "");
|
||||
|
||||
for metric in METRICS {
|
||||
let val_a = get_f64(&json_a, metric.key);
|
||||
let val_b = get_f64(&json_b, metric.key);
|
||||
|
||||
match (val_a, val_b) {
|
||||
(Some(a), Some(b)) => print_diff_row(metric, a, b),
|
||||
_ => {
|
||||
// One or both values missing — skip
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Steady-state section — both sides must have the block; otherwise render N/A.
|
||||
let ss_a = json_a.get("steady_state");
|
||||
let ss_b = json_b.get("steady_state");
|
||||
let both_present =
|
||||
matches!(ss_a, Some(v) if !v.is_null()) && matches!(ss_b, Some(v) if !v.is_null());
|
||||
|
||||
println!();
|
||||
println!("{:=^70}", " Steady-State Comparison ");
|
||||
if !both_present {
|
||||
println!("N/A — one or both runs have no steady-state window");
|
||||
} else {
|
||||
let ss_a = ss_a.unwrap();
|
||||
let ss_b = ss_b.unwrap();
|
||||
for m in STEADY_STATE_METRICS {
|
||||
let a = ss_a.get(m.key).and_then(|v| v.as_f64());
|
||||
let b = ss_b.get(m.key).and_then(|v| v.as_f64());
|
||||
match (a, b) {
|
||||
(Some(a), Some(b)) => print_diff_row(m, a, b),
|
||||
_ => println!("{:<35} N/A", m.label),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
println!("{:=<90}", "");
|
||||
println!();
|
||||
println!("Legend: + = improvement, - = regression (relative to A → B)");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn print_diff_row(metric: &MetricDef, a: f64, b: f64) {
|
||||
let delta = b - a;
|
||||
let pct = if a.abs() > 1e-10 {
|
||||
(delta / a) * 100.0
|
||||
} else if b.abs() > 1e-10 {
|
||||
f64::INFINITY
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
// Determine if change is good/bad/neutral
|
||||
let marker = if delta.abs() < 1e-10 {
|
||||
" "
|
||||
} else if metric.lower_is_better {
|
||||
if delta < 0.0 { "+" } else { "-" }
|
||||
} else if delta > 0.0 {
|
||||
"+"
|
||||
} else {
|
||||
"-"
|
||||
};
|
||||
|
||||
let delta_str = format_delta(delta);
|
||||
let pct_str = if pct.is_infinite() {
|
||||
"inf%".to_string()
|
||||
} else {
|
||||
format!("{:+.1}%", pct)
|
||||
};
|
||||
|
||||
println!(
|
||||
"{:<35} {:>12} {:>12} {:>10} {:>7}{}",
|
||||
metric.label,
|
||||
format_value(a),
|
||||
format_value(b),
|
||||
delta_str,
|
||||
pct_str,
|
||||
marker,
|
||||
);
|
||||
}
|
||||
|
||||
fn load_result_json(path: &str) -> Result<serde_json::Value> {
|
||||
let content = std::fs::read_to_string(path)
|
||||
.map_err(|e| BenchError::Config(format!("Cannot read result file '{path}': {e}")))?;
|
||||
|
||||
// Support JSONL: take the last line (most recent run)
|
||||
let json_str = content.lines().rfind(|l| !l.trim().is_empty()).unwrap_or(&content);
|
||||
|
||||
serde_json::from_str(json_str)
|
||||
.map_err(|e| BenchError::Config(format!("Cannot parse JSON from '{path}': {e}")))
|
||||
}
|
||||
|
||||
fn get_f64(json: &serde_json::Value, key: &str) -> Option<f64> {
|
||||
json.get(key).and_then(|v| v.as_f64())
|
||||
}
|
||||
|
||||
fn format_value(v: f64) -> String {
|
||||
if v == v.floor() && v.abs() < 1e12 {
|
||||
format!("{}", v as i64)
|
||||
} else {
|
||||
format!("{:.2}", v)
|
||||
}
|
||||
}
|
||||
|
||||
fn format_delta(d: f64) -> String {
|
||||
if d == d.floor() && d.abs() < 1e12 {
|
||||
format!("{:+}", d as i64)
|
||||
} else {
|
||||
format!("{:+.2}", d)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,183 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
//! Custom dataset: JSONL file with one request per line.
|
||||
//!
|
||||
//! ```jsonl
|
||||
//! {"prompt": "What is the capital of India?", "output_tokens": 10}
|
||||
//! {"prompt": "What is the capital of Iran?", "output_tokens": 1520}
|
||||
//! ```
|
||||
//!
|
||||
//! Mirrors Python's `CustomDataset`. `output_tokens` is optional unless
|
||||
//! `--custom-output-len -1` is passed. Unlike Python, prompts are always sent
|
||||
//! raw (no client-side chat template; see `--skip-chat-template`).
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rand::SeedableRng;
|
||||
use rand::rngs::StdRng;
|
||||
use rand::seq::SliceRandom;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::SampleRequest;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CustomLine {
|
||||
prompt: String,
|
||||
output_tokens: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// Load the custom JSONL dataset.
|
||||
///
|
||||
/// `output_len < 0` means "use the per-line output_tokens field" (Python's
|
||||
/// `--custom-output-len -1`); otherwise `output_len` applies to every request.
|
||||
pub fn load_custom_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
path: &str,
|
||||
num_requests: usize,
|
||||
output_len: i64,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
no_oversample: bool,
|
||||
disable_shuffle: bool,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
let content = std::fs::read_to_string(path)
|
||||
.map_err(|e| BenchError::Config(format!("Failed to read custom dataset '{path}': {e}")))?;
|
||||
|
||||
let mut lines: Vec<CustomLine> = Vec::new();
|
||||
for (lineno, line) in content.lines().enumerate() {
|
||||
let line = line.trim();
|
||||
if line.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let parsed: CustomLine = serde_json::from_str(line).map_err(|e| {
|
||||
BenchError::Config(format!(
|
||||
"Invalid JSONL at {path}:{}: {e} (each line must be an object \
|
||||
with a 'prompt' field)",
|
||||
lineno + 1
|
||||
))
|
||||
})?;
|
||||
lines.push(parsed);
|
||||
}
|
||||
if lines.is_empty() {
|
||||
return Err(BenchError::Config(format!(
|
||||
"Custom dataset '{path}' contains no entries"
|
||||
)));
|
||||
}
|
||||
|
||||
// Python shuffles the loaded data (seeded) before taking num_requests.
|
||||
if !disable_shuffle {
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
lines.shuffle(&mut rng);
|
||||
}
|
||||
|
||||
let mut requests: Vec<SampleRequest> = Vec::with_capacity(num_requests.min(lines.len()));
|
||||
for (i, item) in lines.iter().enumerate() {
|
||||
if requests.len() >= num_requests {
|
||||
break;
|
||||
}
|
||||
|
||||
let expected_output_len = if output_len < 0 {
|
||||
let raw = item.output_tokens.as_ref().ok_or_else(|| {
|
||||
BenchError::Config(
|
||||
"custom dataset: --custom-output-len -1 requires an \
|
||||
'output_tokens' field on every line"
|
||||
.into(),
|
||||
)
|
||||
})?;
|
||||
raw.as_i64().filter(|v| *v > 0).ok_or_else(|| {
|
||||
BenchError::Config(format!(
|
||||
"custom dataset: invalid 'output_tokens' value {raw}: \
|
||||
must be a positive integer"
|
||||
))
|
||||
})? as usize
|
||||
} else {
|
||||
output_len as usize
|
||||
};
|
||||
|
||||
let prompt_len = tokenizer.encode(&item.prompt, true)?.len();
|
||||
requests.push(SampleRequest {
|
||||
prompt: Arc::from(item.prompt.as_str()),
|
||||
prompt_len,
|
||||
expected_output_len,
|
||||
request_id: Some(format!("{request_id_prefix}{i}")),
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
|
||||
super::oversample_requests(
|
||||
&mut requests,
|
||||
num_requests,
|
||||
request_id_prefix,
|
||||
no_oversample,
|
||||
);
|
||||
Ok(requests)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn write_temp_jsonl(name: &str, content: &str) -> String {
|
||||
let path = std::env::temp_dir().join(format!("vllm-bench-custom-{name}.jsonl"));
|
||||
std::fs::write(&path, content).unwrap();
|
||||
path.to_string_lossy().into_owned()
|
||||
}
|
||||
|
||||
/// gpt2 via built-in tiktoken encoding — loads without network access.
|
||||
fn test_tokenizer() -> TokenizerKind {
|
||||
TokenizerKind::Tiktoken(
|
||||
crate::tiktoken::load_builtin_tiktoken("gpt2")
|
||||
.expect("gpt2 built-in tiktoken should always load without network"),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_custom_dataset_basic() {
|
||||
let path = write_temp_jsonl(
|
||||
"basic",
|
||||
r#"{"prompt": "hello world", "output_tokens": 10}
|
||||
{"prompt": "foo bar baz", "output_tokens": 20}
|
||||
"#,
|
||||
);
|
||||
let reqs = load_custom_dataset(&test_tokenizer(), &path, 2, 256, 0, "t-", true, true)
|
||||
.expect("load should succeed");
|
||||
assert_eq!(reqs.len(), 2);
|
||||
// Fixed output_len (256) wins over per-line output_tokens by default
|
||||
assert!(reqs.iter().all(|r| r.expected_output_len == 256));
|
||||
assert_eq!(&*reqs[0].prompt, "hello world");
|
||||
assert!(reqs[0].prompt_len > 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_custom_dataset_per_line_output_tokens() {
|
||||
let path = write_temp_jsonl(
|
||||
"perline",
|
||||
r#"{"prompt": "hello", "output_tokens": 10}
|
||||
{"prompt": "world", "output_tokens": 20}
|
||||
"#,
|
||||
);
|
||||
let reqs = load_custom_dataset(&test_tokenizer(), &path, 2, -1, 0, "t-", true, true)
|
||||
.expect("load should succeed");
|
||||
assert_eq!(reqs[0].expected_output_len, 10);
|
||||
assert_eq!(reqs[1].expected_output_len, 20);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_custom_dataset_missing_output_tokens_errors() {
|
||||
let path = write_temp_jsonl("missing", r#"{"prompt": "hello"}"#);
|
||||
let err = load_custom_dataset(&test_tokenizer(), &path, 1, -1, 0, "t-", true, true)
|
||||
.expect_err("should fail without output_tokens");
|
||||
assert!(err.to_string().contains("output_tokens"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_custom_dataset_missing_prompt_errors() {
|
||||
let path = write_temp_jsonl("noprompt", r#"{"text": "hello"}"#);
|
||||
assert!(
|
||||
load_custom_dataset(&test_tokenizer(), &path, 1, 256, 0, "t-", true, true).is_err()
|
||||
);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,210 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
pub mod custom;
|
||||
pub mod hf_dataset;
|
||||
pub mod multi_turn;
|
||||
pub mod prefix_repetition;
|
||||
mod progress;
|
||||
pub mod random;
|
||||
pub mod random_mm;
|
||||
pub mod random_rerank;
|
||||
pub mod sharegpt;
|
||||
pub mod sonnet;
|
||||
pub mod speed_bench;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
/// Represents a single inference request for benchmarking.
|
||||
/// Matches Python's SampleRequest dataclass from datasets.py:71-82.
|
||||
///
|
||||
/// `prompt` uses `Arc<str>` to avoid expensive String clones when distributing
|
||||
/// requests across tokio tasks. At 100k prompts with 8k tokens each, this saves
|
||||
/// ~3GB of peak memory vs cloning String per task.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SampleRequest {
|
||||
pub prompt: Arc<str>,
|
||||
pub prompt_len: usize,
|
||||
pub expected_output_len: usize,
|
||||
pub request_id: Option<String>,
|
||||
/// Pre-computed token IDs for this prompt.
|
||||
/// When set, the completions backend sends these directly via `prompt_token_ids`
|
||||
/// instead of the text `prompt`, avoiding server-side re-tokenization.
|
||||
pub prompt_token_ids: Option<Arc<[u32]>>,
|
||||
/// Multimodal content items as pre-serialized JSON fragments.
|
||||
/// Each `Arc<str>` is a complete JSON object string, e.g.
|
||||
/// `{"type":"image_url","image_url":{"url":"data:image/jpeg;base64,..."}}`
|
||||
///
|
||||
/// Pre-serialized to avoid:
|
||||
/// 1. `serde_json::Value` tree overhead (3 Maps + keys per image)
|
||||
/// 2. Deep-cloning ~200KB+ base64 data when building request payloads
|
||||
///
|
||||
/// Double-`Arc` for zero-cost sharing: outer Arc for the slice, inner Arc for each fragment.
|
||||
pub multi_modal_content: Option<Arc<[Arc<str>]>>,
|
||||
/// Pre-serialized OpenAI chat `messages` array as a complete JSON string,
|
||||
/// e.g. `[{"role":"user","content":[{"type":"text","text":"..."},{"type":"image_url",...}]}]`.
|
||||
///
|
||||
/// Set by datasets when `--enable-multimodal-chat` is on (mirrors Python's
|
||||
/// `apply_multimodal_chat_transformation`: the dataset builds the chat messages
|
||||
/// and the backend sends them verbatim). When set, `multi_modal_content` is None
|
||||
/// and the mm items are embedded here instead. `prompt` still holds the text part
|
||||
/// for token accounting and /tokenize verification.
|
||||
pub chat_messages_json: Option<Arc<str>>,
|
||||
/// Multiple text inputs for one request (pooling backends only).
|
||||
/// Embeddings send it as `"input": [t1, t2, ...]` (--random-batch-size);
|
||||
/// rerank sends `[0]` as the query and `[1..]` as documents (random-rerank).
|
||||
/// Mirrors Python's list-valued `SampleRequest.prompt`.
|
||||
pub prompt_list: Option<Arc<[Arc<str>]>>,
|
||||
}
|
||||
|
||||
impl Default for SampleRequest {
|
||||
/// Empty request; struct-update base so dataset builders only spell out the
|
||||
/// fields they set (new optional fields then don't touch every call site).
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
prompt: Arc::from(""),
|
||||
prompt_len: 0,
|
||||
expected_output_len: 0,
|
||||
request_id: None,
|
||||
prompt_token_ids: None,
|
||||
multi_modal_content: None,
|
||||
chat_messages_json: None,
|
||||
prompt_list: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Oversample `requests` up to `num_requests` by cloning random entries
|
||||
/// (seeded by list length for determinism), renumbering their request ids.
|
||||
/// No-op when enough samples exist, `no_oversample` is set, or the list is empty.
|
||||
/// Mirrors Python `BenchmarkDataset.maybe_oversample_requests`.
|
||||
pub fn oversample_requests(
|
||||
requests: &mut Vec<SampleRequest>,
|
||||
num_requests: usize,
|
||||
request_id_prefix: &str,
|
||||
no_oversample: bool,
|
||||
) {
|
||||
use rand::rngs::StdRng;
|
||||
use rand::{Rng, SeedableRng};
|
||||
|
||||
if requests.len() >= num_requests || requests.is_empty() {
|
||||
return;
|
||||
}
|
||||
if no_oversample {
|
||||
tracing::info!(
|
||||
samples = requests.len(),
|
||||
requested = num_requests,
|
||||
"skipping dataset oversampling"
|
||||
);
|
||||
return;
|
||||
}
|
||||
let original_len = requests.len();
|
||||
let mut rng = StdRng::seed_from_u64(original_len as u64);
|
||||
for i in 0..(num_requests - original_len) {
|
||||
let mut req = requests[rng.random_range(0..original_len)].clone();
|
||||
req.request_id = Some(format!("{request_id_prefix}{}", original_len + i));
|
||||
requests.push(req);
|
||||
}
|
||||
tracing::info!(
|
||||
original_samples = original_len,
|
||||
samples = requests.len(),
|
||||
"oversampled dataset"
|
||||
);
|
||||
}
|
||||
|
||||
/// Group already-generated single-input requests into batched requests of
|
||||
/// `batch_size` inputs each (embeddings/pooling only). Mirrors Python
|
||||
/// `RandomDataset.sample` batching: prompt becomes a list, prompt_len is the
|
||||
/// sum over the batch, request ids are renumbered per batch.
|
||||
/// `batch_size <= 1` returns the input unchanged.
|
||||
pub fn batch_requests(
|
||||
requests: Vec<SampleRequest>,
|
||||
batch_size: usize,
|
||||
request_id_prefix: &str,
|
||||
) -> Vec<SampleRequest> {
|
||||
if batch_size <= 1 {
|
||||
return requests;
|
||||
}
|
||||
requests
|
||||
.chunks(batch_size)
|
||||
.enumerate()
|
||||
.map(|(batch_idx, batch)| SampleRequest {
|
||||
prompt_list: Some(batch.iter().map(|r| r.prompt.clone()).collect()),
|
||||
prompt_len: batch.iter().map(|r| r.prompt_len).sum(),
|
||||
expected_output_len: 0,
|
||||
request_id: Some(format!("{request_id_prefix}{batch_idx}")),
|
||||
..Default::default()
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn req(prompt: &str, len: usize) -> SampleRequest {
|
||||
SampleRequest {
|
||||
prompt: Arc::from(prompt),
|
||||
prompt_len: len,
|
||||
expected_output_len: 128,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_batch_requests_groups_and_sums() {
|
||||
let reqs = vec![
|
||||
req("a", 10),
|
||||
req("b", 20),
|
||||
req("c", 30),
|
||||
req("d", 40),
|
||||
req("e", 50),
|
||||
];
|
||||
let batched = batch_requests(reqs, 2, "t-");
|
||||
assert_eq!(batched.len(), 3); // 2 + 2 + 1
|
||||
let first = batched[0].prompt_list.as_ref().unwrap();
|
||||
assert_eq!(first.len(), 2);
|
||||
assert_eq!(&*first[0], "a");
|
||||
assert_eq!(batched[0].prompt_len, 30);
|
||||
assert_eq!(batched[0].expected_output_len, 0);
|
||||
assert_eq!(batched[0].request_id.as_deref(), Some("t-0"));
|
||||
assert_eq!(batched[2].prompt_list.as_ref().unwrap().len(), 1);
|
||||
assert_eq!(batched[2].prompt_len, 50);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_batch_requests_size_one_is_identity() {
|
||||
let reqs = vec![req("a", 10), req("b", 20)];
|
||||
let out = batch_requests(reqs, 1, "t-");
|
||||
assert_eq!(out.len(), 2);
|
||||
assert!(out[0].prompt_list.is_none());
|
||||
assert_eq!(&*out[0].prompt, "a");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_oversample_requests() {
|
||||
let mut reqs = vec![req("a", 10), req("b", 20)];
|
||||
oversample_requests(&mut reqs, 5, "t-", false);
|
||||
assert_eq!(reqs.len(), 5);
|
||||
assert_eq!(reqs[4].request_id.as_deref(), Some("t-4"));
|
||||
|
||||
let mut reqs = vec![req("a", 10)];
|
||||
oversample_requests(&mut reqs, 5, "t-", true); // no_oversample
|
||||
assert_eq!(reqs.len(), 1);
|
||||
}
|
||||
}
|
||||
|
||||
/// A single turn in a multi-turn conversation.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ConversationTurn {
|
||||
pub user_message: Arc<str>,
|
||||
pub user_message_len: usize,
|
||||
pub expected_output_len: usize,
|
||||
}
|
||||
|
||||
/// A complete multi-turn conversation with all turns pre-generated.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MultiTurnConversation {
|
||||
pub conversation_id: String,
|
||||
pub turns: Vec<ConversationTurn>,
|
||||
}
|
||||
@@ -0,0 +1,788 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rand::rngs::StdRng;
|
||||
use rand::seq::SliceRandom;
|
||||
use rand::{Rng, SeedableRng};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use super::{ConversationTurn, MultiTurnConversation};
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
/// Configuration for generating random multi-turn conversations.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MultiTurnRandomConfig {
|
||||
pub num_conversations: usize,
|
||||
pub min_turns: usize,
|
||||
pub max_turns: usize,
|
||||
/// Shared prefix length prepended to the conversation.
|
||||
///
|
||||
/// In normal accumulated-history mode this is added to turn 0, so all
|
||||
/// later turns inherit it through history. In no-history prefix-sharing
|
||||
/// mode it is added to every independent turn.
|
||||
pub prefix_len: usize,
|
||||
/// Input length for turn 0.
|
||||
pub input_len: usize,
|
||||
/// Input length for turns 1+. 0 = fallback to input_len.
|
||||
pub per_turn_input_len: usize,
|
||||
pub output_len: usize,
|
||||
pub seed: u64,
|
||||
pub request_id_prefix: String,
|
||||
pub prefix_sharing_config: Option<PrefixSharingConfig>,
|
||||
}
|
||||
|
||||
/// Configuration for 3-tier prefix sharing in multi-turn user messages.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PrefixSharingConfig {
|
||||
/// Fraction of per-turn input tokens shared across ALL conversations.
|
||||
pub global_ratio: f64,
|
||||
/// Fraction of per-turn input tokens shared within each conversation.
|
||||
pub conversation_ratio: f64,
|
||||
}
|
||||
|
||||
/// Generate a deterministic token sequence from allowed tokens using offset+modulo.
|
||||
fn make_token_seq(allowed_tokens: &[u32], offset: usize, len: usize) -> Vec<u32> {
|
||||
let at_len = allowed_tokens.len();
|
||||
(0..len).map(|i| allowed_tokens[(offset + i) % at_len]).collect()
|
||||
}
|
||||
|
||||
/// Generate synthetic multi-turn conversations with random user messages.
|
||||
///
|
||||
/// Each conversation has `num_turns` turns, each with a random user prompt
|
||||
/// of `input_len` tokens and `output_len` expected output tokens.
|
||||
pub fn generate_multi_turn_random(
|
||||
tokenizer: &TokenizerKind,
|
||||
cfg: &MultiTurnRandomConfig,
|
||||
) -> Result<Vec<MultiTurnConversation>> {
|
||||
let num_conversations = cfg.num_conversations;
|
||||
let min_turns = cfg.min_turns;
|
||||
let max_turns = cfg.max_turns;
|
||||
let prefix_len = cfg.prefix_len;
|
||||
let input_len = cfg.input_len;
|
||||
let output_len = cfg.output_len;
|
||||
let seed = cfg.seed;
|
||||
let request_id_prefix = &cfg.request_id_prefix;
|
||||
let allowed_tokens = tokenizer.get_allowed_tokens();
|
||||
if allowed_tokens.is_empty() {
|
||||
return Err(BenchError::Tokenizer("No allowed tokens found".into()));
|
||||
}
|
||||
|
||||
let vocab_size = tokenizer.vocab_size() as usize;
|
||||
let num_special = tokenizer.num_special_tokens_to_add();
|
||||
let real_input_len = input_len.saturating_sub(num_special);
|
||||
let real_per_turn_len = if cfg.per_turn_input_len > 0 {
|
||||
cfg.per_turn_input_len.saturating_sub(num_special)
|
||||
} else {
|
||||
real_input_len
|
||||
};
|
||||
|
||||
if real_input_len < 1 {
|
||||
return Err(BenchError::Config(format!(
|
||||
"--random-input-len too small: with {num_special} special tokens, \
|
||||
effective input length is {real_input_len}"
|
||||
)));
|
||||
}
|
||||
if real_per_turn_len < 1 {
|
||||
return Err(BenchError::Config(format!(
|
||||
"--per-turn-input-len too small: with {num_special} special tokens, \
|
||||
effective per-turn input length is {real_per_turn_len}"
|
||||
)));
|
||||
}
|
||||
|
||||
// Prefix sharing mode: generate 3-tier prefixed messages
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
if let Some(ref ps_cfg) = cfg.prefix_sharing_config {
|
||||
return generate_prefix_sharing_conversations(
|
||||
tokenizer,
|
||||
cfg,
|
||||
ps_cfg,
|
||||
&allowed_tokens,
|
||||
&mut rng,
|
||||
);
|
||||
}
|
||||
let shared_prefix_text =
|
||||
generate_shared_prefix_text(tokenizer, &allowed_tokens, prefix_len, seed)?;
|
||||
|
||||
// Pre-generate per-conversation turn counts and per-turn offsets deterministically.
|
||||
// Turn counts are drawn first so the RNG sequence is stable regardless of vocab_size.
|
||||
let conv_turn_counts: Vec<usize> = (0..num_conversations)
|
||||
.map(|_| {
|
||||
if min_turns == max_turns {
|
||||
min_turns
|
||||
} else {
|
||||
rng.random_range(min_turns..=max_turns)
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let offsets: Vec<Vec<usize>> = conv_turn_counts
|
||||
.iter()
|
||||
.map(|&n| (0..n).map(|_| rng.random_range(0..vocab_size)).collect())
|
||||
.collect();
|
||||
|
||||
// Parallel generation across conversations
|
||||
offsets
|
||||
.par_iter()
|
||||
.enumerate()
|
||||
.map(|(conv_idx, conv_offsets)| {
|
||||
let mut turns = Vec::with_capacity(conv_offsets.len());
|
||||
for (turn_idx, &offset) in conv_offsets.iter().enumerate() {
|
||||
let target_len = if turn_idx == 0 {
|
||||
real_input_len
|
||||
} else {
|
||||
real_per_turn_len
|
||||
};
|
||||
// Use max_turns stride to keep offsets unique across variable-length convs
|
||||
let inner_seq = make_token_seq(
|
||||
&allowed_tokens,
|
||||
offset + conv_idx * max_turns + turn_idx,
|
||||
target_len,
|
||||
);
|
||||
|
||||
let (prompt, adjusted) =
|
||||
gen_prompt_to_target_len(tokenizer, &inner_seq, target_len)?;
|
||||
let (prompt, token_len) = if turn_idx == 0 && !shared_prefix_text.is_empty() {
|
||||
let combined = format!("{}{}", &*shared_prefix_text, prompt);
|
||||
let token_len = tokenizer.encode(&combined, false)?.len();
|
||||
(combined, token_len)
|
||||
} else {
|
||||
(prompt, adjusted.len())
|
||||
};
|
||||
|
||||
turns.push(ConversationTurn {
|
||||
user_message: Arc::from(prompt),
|
||||
user_message_len: token_len,
|
||||
expected_output_len: output_len,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(MultiTurnConversation {
|
||||
conversation_id: format!("{request_id_prefix}conv-{conv_idx}"),
|
||||
turns,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Generate conversations with 3-tier prefix sharing.
|
||||
///
|
||||
/// Each turn's user message = [global_prefix][conversation_prefix][unique_suffix].
|
||||
/// No history accumulation — each turn sends only its own fixed-length message.
|
||||
fn generate_prefix_sharing_conversations(
|
||||
tokenizer: &TokenizerKind,
|
||||
cfg: &MultiTurnRandomConfig,
|
||||
ps_cfg: &PrefixSharingConfig,
|
||||
allowed_tokens: &[u32],
|
||||
rng: &mut StdRng,
|
||||
) -> Result<Vec<MultiTurnConversation>> {
|
||||
let num_conversations = cfg.num_conversations;
|
||||
let min_turns = cfg.min_turns;
|
||||
let max_turns = cfg.max_turns;
|
||||
let prefix_len = cfg.prefix_len;
|
||||
let output_len = cfg.output_len;
|
||||
let request_id_prefix = &cfg.request_id_prefix;
|
||||
|
||||
let num_special = tokenizer.num_special_tokens_to_add();
|
||||
let real_input_len = cfg.input_len.saturating_sub(num_special);
|
||||
let real_per_turn_len = if cfg.per_turn_input_len > 0 {
|
||||
cfg.per_turn_input_len.saturating_sub(num_special)
|
||||
} else {
|
||||
real_input_len
|
||||
};
|
||||
|
||||
// Compute segment lengths from turn-0 (real_input_len) so the shared prefix
|
||||
// bytes stay byte-identical across all turns regardless of per_turn_input_len.
|
||||
let global_len = (real_input_len as f64 * ps_cfg.global_ratio).floor() as usize;
|
||||
let conv_len = (real_input_len as f64 * ps_cfg.conversation_ratio).floor() as usize;
|
||||
let unique_len = real_input_len.saturating_sub(global_len + conv_len);
|
||||
|
||||
// Validate that turns 1+ still have room for a non-empty unique suffix
|
||||
if real_per_turn_len <= global_len + conv_len {
|
||||
return Err(BenchError::Config(format!(
|
||||
"--per-turn-input-len ({real_per_turn_len} after special tokens) is too small: \
|
||||
global_len={global_len} + conv_len={conv_len} already fills the budget. \
|
||||
Increase --per-turn-input-len or reduce prefix ratios."
|
||||
)));
|
||||
}
|
||||
|
||||
let at_len = allowed_tokens.len();
|
||||
let shared_prefix_text =
|
||||
generate_shared_prefix_text(tokenizer, allowed_tokens, prefix_len, cfg.seed)?;
|
||||
|
||||
// Generate global prefix text once
|
||||
let global_text: Arc<str> = if global_len > 0 {
|
||||
let offset: usize = rng.random_range(0..at_len);
|
||||
let seq = make_token_seq(allowed_tokens, offset, global_len);
|
||||
let (text, _) = gen_prompt_to_target_len(tokenizer, &seq, global_len)?;
|
||||
Arc::from(text)
|
||||
} else {
|
||||
Arc::from("")
|
||||
};
|
||||
|
||||
// Generate per-conversation prefix texts
|
||||
let conv_texts: Vec<Arc<str>> = if conv_len > 0 {
|
||||
let mut texts = Vec::with_capacity(num_conversations);
|
||||
for conv_idx in 0..num_conversations {
|
||||
let offset: usize = rng.random_range(0..at_len);
|
||||
let seq = make_token_seq(allowed_tokens, offset + conv_idx, conv_len);
|
||||
let (text, _) = gen_prompt_to_target_len(tokenizer, &seq, conv_len)?;
|
||||
texts.push(Arc::from(text));
|
||||
}
|
||||
texts
|
||||
} else {
|
||||
vec![Arc::from(""); num_conversations]
|
||||
};
|
||||
|
||||
// Pre-generate per-conversation turn counts and unique offsets deterministically.
|
||||
let vocab_size = tokenizer.vocab_size() as usize;
|
||||
let conv_turn_counts: Vec<usize> = (0..num_conversations)
|
||||
.map(|_| {
|
||||
if min_turns == max_turns {
|
||||
min_turns
|
||||
} else {
|
||||
rng.random_range(min_turns..=max_turns)
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let unique_offsets: Vec<Vec<usize>> = conv_turn_counts
|
||||
.iter()
|
||||
.map(|&n| (0..n).map(|_| rng.random_range(0..vocab_size)).collect())
|
||||
.collect();
|
||||
|
||||
// Parallel generation across conversations
|
||||
unique_offsets
|
||||
.par_iter()
|
||||
.enumerate()
|
||||
.map(|(conv_idx, conv_offsets)| {
|
||||
let mut turns = Vec::with_capacity(conv_offsets.len());
|
||||
for (turn_idx, &offset) in conv_offsets.iter().enumerate() {
|
||||
// Turn 0 uses unique_len derived from real_input_len;
|
||||
// turns 1+ use per-turn unique_len (prefix bytes stay identical).
|
||||
let turn_unique_len = if turn_idx == 0 {
|
||||
unique_len
|
||||
} else {
|
||||
real_per_turn_len.saturating_sub(global_len + conv_len)
|
||||
};
|
||||
|
||||
// Generate unique suffix
|
||||
let unique_text = if turn_unique_len > 0 {
|
||||
let seq = make_token_seq(
|
||||
allowed_tokens,
|
||||
offset + conv_idx * max_turns + turn_idx,
|
||||
turn_unique_len,
|
||||
);
|
||||
let (text, _) = gen_prompt_to_target_len(tokenizer, &seq, turn_unique_len)?;
|
||||
text
|
||||
} else {
|
||||
String::new()
|
||||
};
|
||||
|
||||
// Concatenate: optional random prefix + global + conversation + unique.
|
||||
// Prefix-sharing mode sends each turn independently, so the random
|
||||
// prefix must be included on every turn to be present in every request.
|
||||
let combined = format!(
|
||||
"{}{}{}{}",
|
||||
&*shared_prefix_text, &*global_text, &*conv_texts[conv_idx], unique_text
|
||||
);
|
||||
// Re-encode to get actual token count (BPE boundary effects)
|
||||
let token_len = tokenizer.encode(&combined, false)?.len();
|
||||
|
||||
turns.push(ConversationTurn {
|
||||
user_message: Arc::from(combined),
|
||||
user_message_len: token_len,
|
||||
expected_output_len: output_len,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(MultiTurnConversation {
|
||||
conversation_id: format!("{request_id_prefix}conv-{conv_idx}"),
|
||||
turns,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn generate_shared_prefix_text(
|
||||
tokenizer: &TokenizerKind,
|
||||
allowed_tokens: &[u32],
|
||||
prefix_len: usize,
|
||||
seed: u64,
|
||||
) -> Result<Arc<str>> {
|
||||
if prefix_len == 0 {
|
||||
return Ok(Arc::from(""));
|
||||
}
|
||||
|
||||
let mut rng = StdRng::seed_from_u64(seed.wrapping_add(0xDEAD));
|
||||
let tokens: Vec<u32> = (0..prefix_len)
|
||||
.map(|_| allowed_tokens[rng.random_range(0..allowed_tokens.len())])
|
||||
.collect();
|
||||
let (text, _) = gen_prompt_to_target_len(tokenizer, &tokens, prefix_len)?;
|
||||
Ok(Arc::from(text))
|
||||
}
|
||||
|
||||
/// Load multi-turn conversations from a ShareGPT dataset.
|
||||
///
|
||||
/// Walks ALL turns in each entry (not just first 2). Filters entries
|
||||
/// with at least 4 messages (2 user + 2 assistant = 2 real turns).
|
||||
pub fn load_sharegpt_multi_turn(
|
||||
tokenizer: &TokenizerKind,
|
||||
dataset_path: &str,
|
||||
num_conversations: usize,
|
||||
output_len_override: Option<usize>,
|
||||
max_turns: Option<usize>,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
) -> Result<Vec<MultiTurnConversation>> {
|
||||
let content = std::fs::read_to_string(dataset_path).map_err(|e| {
|
||||
BenchError::Config(format!(
|
||||
"Failed to read ShareGPT file '{dataset_path}': {e}"
|
||||
))
|
||||
})?;
|
||||
|
||||
let data: serde_json::Value = serde_json::from_str(&content)
|
||||
.map_err(|e| BenchError::Config(format!("Invalid JSON in ShareGPT file: {e}")))?;
|
||||
|
||||
let entries = data
|
||||
.as_array()
|
||||
.ok_or_else(|| BenchError::Config("ShareGPT file must contain a JSON array".into()))?;
|
||||
|
||||
// Filter entries with at least 4 messages (2 turns: user+assistant+user+assistant)
|
||||
let mut filtered: Vec<&serde_json::Value> = entries
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
entry
|
||||
.get("conversations")
|
||||
.and_then(|c| c.as_array())
|
||||
.map(|a| a.len() >= 4)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.collect();
|
||||
|
||||
if filtered.is_empty() {
|
||||
return Err(BenchError::Config(
|
||||
"No valid multi-turn entries in ShareGPT file (need at least 4 messages per entry)"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
|
||||
// Shuffle
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
filtered.shuffle(&mut rng);
|
||||
|
||||
let mut conversations = Vec::new();
|
||||
|
||||
for entry in &filtered {
|
||||
if conversations.len() >= num_conversations {
|
||||
break;
|
||||
}
|
||||
|
||||
let msgs = entry["conversations"].as_array().unwrap();
|
||||
let mut turns = Vec::new();
|
||||
|
||||
// Walk alternating human/gpt pairs, stopping early once max_turns reached
|
||||
// to avoid tokenizing turns that would be discarded by truncate().
|
||||
let mut i = 0;
|
||||
while i + 1 < msgs.len() {
|
||||
if let Some(m) = max_turns
|
||||
&& turns.len() >= m
|
||||
{
|
||||
break;
|
||||
}
|
||||
let from = msgs[i].get("from").and_then(|f| f.as_str()).unwrap_or("");
|
||||
let user_text = msgs[i].get("value").and_then(|v| v.as_str()).unwrap_or("");
|
||||
let assistant_text = msgs[i + 1].get("value").and_then(|v| v.as_str()).unwrap_or("");
|
||||
|
||||
// Expect human then gpt
|
||||
if from != "human" || user_text.is_empty() {
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
|
||||
let user_ids = tokenizer.encode(user_text, false)?;
|
||||
let user_len = user_ids.len();
|
||||
|
||||
let expected_output_len = if let Some(override_len) = output_len_override {
|
||||
override_len
|
||||
} else {
|
||||
let assistant_ids = tokenizer.encode(assistant_text, false)?;
|
||||
assistant_ids.len().max(1)
|
||||
};
|
||||
|
||||
turns.push(ConversationTurn {
|
||||
user_message: Arc::from(user_text),
|
||||
user_message_len: user_len,
|
||||
expected_output_len,
|
||||
});
|
||||
|
||||
i += 2;
|
||||
}
|
||||
|
||||
if turns.len() >= 2 {
|
||||
let conv_idx = conversations.len();
|
||||
conversations.push(MultiTurnConversation {
|
||||
conversation_id: format!("{request_id_prefix}conv-{conv_idx}"),
|
||||
turns,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if conversations.is_empty() {
|
||||
return Err(BenchError::Config(
|
||||
"No valid multi-turn conversations after filtering ShareGPT dataset.".into(),
|
||||
));
|
||||
}
|
||||
|
||||
// Oversample if needed
|
||||
if conversations.len() < num_conversations {
|
||||
let original_len = conversations.len();
|
||||
let needed = num_conversations - original_len;
|
||||
for i in 0..needed {
|
||||
let mut conv = conversations[rng.random_range(0..original_len)].clone();
|
||||
conv.conversation_id = format!("{request_id_prefix}conv-{}", original_len + i);
|
||||
conversations.push(conv);
|
||||
}
|
||||
tracing::info!(
|
||||
original_conversations = original_len,
|
||||
conversations = conversations.len(),
|
||||
"oversampled multi-turn conversations"
|
||||
);
|
||||
}
|
||||
|
||||
Ok(conversations)
|
||||
}
|
||||
|
||||
/// Ensure decoded-then-encoded prompt length matches the target.
|
||||
fn gen_prompt_to_target_len(
|
||||
tokenizer: &TokenizerKind,
|
||||
token_sequence: &[u32],
|
||||
target_len: usize,
|
||||
) -> Result<(String, Vec<u32>)> {
|
||||
let max_retry = 20;
|
||||
let mut tokens = token_sequence.to_vec();
|
||||
|
||||
for retry in 0..=max_retry {
|
||||
let prompt = tokenizer.decode(&tokens, true)?;
|
||||
tokens = tokenizer.encode(&prompt, false)?;
|
||||
|
||||
if retry >= max_retry {
|
||||
// BPE tokenizers can oscillate by ±1 on certain boundaries.
|
||||
// For benchmark random content, accept close-enough and truncate/pad.
|
||||
if tokens.len() > target_len {
|
||||
tokens.truncate(target_len);
|
||||
}
|
||||
// If still short by 1-2 tokens, accept as-is — negligible for benchmarks.
|
||||
// Re-decode after truncation to ensure prompt string matches token vector.
|
||||
let prompt = tokenizer.decode(&tokens, true)?;
|
||||
return Ok((prompt, tokens));
|
||||
}
|
||||
|
||||
if tokens.len() == target_len {
|
||||
return Ok((prompt, tokens));
|
||||
} else if tokens.len() < target_len {
|
||||
let allowed = tokenizer.get_allowed_tokens();
|
||||
let needed = target_len - tokens.len();
|
||||
if allowed.is_empty() {
|
||||
let vocab_size = tokenizer.vocab_size() as usize;
|
||||
for j in 0..needed {
|
||||
tokens.push(((tokens.len() + j) % vocab_size) as u32);
|
||||
}
|
||||
} else {
|
||||
for j in 0..needed {
|
||||
tokens.push(allowed[(tokens.len() + j) % allowed.len()]);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
tokens.truncate(target_len);
|
||||
}
|
||||
}
|
||||
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn common_prefix_bytes(strings: &[&str]) -> usize {
|
||||
if strings.is_empty() {
|
||||
return 0;
|
||||
}
|
||||
let first = strings[0].as_bytes();
|
||||
let mut len = first.len();
|
||||
for s in &strings[1..] {
|
||||
let b = s.as_bytes();
|
||||
len = len.min(b.len());
|
||||
for i in 0..len {
|
||||
if first[i] != b[i] {
|
||||
len = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
len
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_prefix_sharing_structure() {
|
||||
let tok = crate::tokenizer::load_tokenizer("nvidia/Kimi-K2.5-NVFP4", false, None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let cfg = MultiTurnRandomConfig {
|
||||
num_conversations: 5,
|
||||
min_turns: 3,
|
||||
max_turns: 3,
|
||||
prefix_len: 0,
|
||||
input_len: 1000,
|
||||
per_turn_input_len: 0,
|
||||
output_len: 100,
|
||||
seed: 42,
|
||||
request_id_prefix: "test-".to_string(),
|
||||
prefix_sharing_config: Some(PrefixSharingConfig {
|
||||
global_ratio: 0.1,
|
||||
conversation_ratio: 0.8,
|
||||
}),
|
||||
};
|
||||
|
||||
let conversations = generate_multi_turn_random(&tok, &cfg).unwrap();
|
||||
assert_eq!(conversations.len(), 5);
|
||||
|
||||
let messages: Vec<Vec<&str>> = conversations
|
||||
.iter()
|
||||
.map(|c| c.turns.iter().map(|t| &*t.user_message).collect())
|
||||
.collect();
|
||||
|
||||
// 1. Global prefix: all messages share a common prefix
|
||||
let all_msgs: Vec<&str> = messages.iter().flat_map(|v| v.iter().copied()).collect();
|
||||
let global_prefix = common_prefix_bytes(&all_msgs);
|
||||
println!("Global prefix bytes: {global_prefix}");
|
||||
assert!(global_prefix > 0, "Global prefix must be non-empty");
|
||||
|
||||
// 2. Conversation prefix: turns within same conversation share more
|
||||
for (i, conv_msgs) in messages.iter().enumerate() {
|
||||
let conv_prefix = common_prefix_bytes(conv_msgs);
|
||||
println!("Conv {i} prefix bytes: {conv_prefix} (global: {global_prefix})");
|
||||
assert!(
|
||||
conv_prefix > global_prefix,
|
||||
"Conv prefix ({conv_prefix}) must exceed global prefix ({global_prefix})"
|
||||
);
|
||||
}
|
||||
|
||||
// 3. Different conversations diverge after global prefix
|
||||
let cross = common_prefix_bytes(&[messages[0][0], messages[1][0]]);
|
||||
let within = common_prefix_bytes(&messages[0]);
|
||||
println!("Cross-conv prefix: {cross}, within-conv prefix: {within}");
|
||||
assert!(
|
||||
cross < within,
|
||||
"Cross-conv ({cross}) must be < within-conv ({within})"
|
||||
);
|
||||
|
||||
// 4. Turns within same conversation are not identical (unique suffix)
|
||||
for (i, conv_msgs) in messages.iter().enumerate() {
|
||||
for a in 0..conv_msgs.len() {
|
||||
for b in (a + 1)..conv_msgs.len() {
|
||||
assert_ne!(
|
||||
conv_msgs[a], conv_msgs[b],
|
||||
"Conv {i} turn {a} and {b} must differ"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 5. Token lengths approximately match target
|
||||
for (i, conv) in conversations.iter().enumerate() {
|
||||
for (j, turn) in conv.turns.iter().enumerate() {
|
||||
let diff = (turn.user_message_len as i64 - 1000).abs();
|
||||
println!(
|
||||
"Conv {i} turn {j}: {} tokens (diff {diff})",
|
||||
turn.user_message_len
|
||||
);
|
||||
assert!(
|
||||
diff <= 10,
|
||||
"Token len {} too far from 1000",
|
||||
turn.user_message_len
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
println!("All prefix sharing checks passed!");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_per_turn_input_len_default_mode() {
|
||||
let tok = crate::tokenizer::load_tokenizer("nvidia/Kimi-K2.5-NVFP4", false, None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let cfg = MultiTurnRandomConfig {
|
||||
num_conversations: 4,
|
||||
min_turns: 3,
|
||||
max_turns: 3,
|
||||
prefix_len: 0,
|
||||
input_len: 512,
|
||||
per_turn_input_len: 128,
|
||||
output_len: 64,
|
||||
seed: 1,
|
||||
request_id_prefix: "test-".to_string(),
|
||||
prefix_sharing_config: None,
|
||||
};
|
||||
|
||||
let conversations = generate_multi_turn_random(&tok, &cfg).unwrap();
|
||||
assert_eq!(conversations.len(), 4);
|
||||
|
||||
for (i, conv) in conversations.iter().enumerate() {
|
||||
assert_eq!(conv.turns.len(), 3);
|
||||
for (j, turn) in conv.turns.iter().enumerate() {
|
||||
let expected = if j == 0 { 512usize } else { 128usize };
|
||||
let diff = (turn.user_message_len as i64 - expected as i64).abs();
|
||||
println!(
|
||||
"Conv {i} turn {j}: {} tokens (expected ~{expected}, diff {diff})",
|
||||
turn.user_message_len
|
||||
);
|
||||
assert!(
|
||||
diff <= 5,
|
||||
"Conv {i} turn {j}: token len {} too far from {expected}",
|
||||
turn.user_message_len
|
||||
);
|
||||
}
|
||||
}
|
||||
println!("per_turn_input_len default-mode checks passed!");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_variable_turns_range() {
|
||||
let tok = crate::tokenizer::load_tokenizer("nvidia/Kimi-K2.5-NVFP4", false, None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let cfg = MultiTurnRandomConfig {
|
||||
num_conversations: 50,
|
||||
min_turns: 2,
|
||||
max_turns: 5,
|
||||
prefix_len: 0,
|
||||
input_len: 256,
|
||||
per_turn_input_len: 0,
|
||||
output_len: 32,
|
||||
seed: 7,
|
||||
request_id_prefix: "test-".to_string(),
|
||||
prefix_sharing_config: None,
|
||||
};
|
||||
|
||||
let conversations = generate_multi_turn_random(&tok, &cfg).unwrap();
|
||||
assert_eq!(conversations.len(), 50);
|
||||
|
||||
let mut distinct_counts = std::collections::HashSet::new();
|
||||
for conv in &conversations {
|
||||
let n = conv.turns.len();
|
||||
assert!((2..=5).contains(&n), "turn count {n} out of [2,5]");
|
||||
distinct_counts.insert(n);
|
||||
}
|
||||
assert!(
|
||||
distinct_counts.len() >= 2,
|
||||
"expected at least 2 distinct turn counts, got {distinct_counts:?}"
|
||||
);
|
||||
println!("variable_turns_range checks passed! counts: {distinct_counts:?}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_variable_turns_fixed() {
|
||||
let tok = crate::tokenizer::load_tokenizer("nvidia/Kimi-K2.5-NVFP4", false, None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let cfg = MultiTurnRandomConfig {
|
||||
num_conversations: 10,
|
||||
min_turns: 4,
|
||||
max_turns: 4,
|
||||
prefix_len: 0,
|
||||
input_len: 256,
|
||||
per_turn_input_len: 0,
|
||||
output_len: 32,
|
||||
seed: 42,
|
||||
request_id_prefix: "test-".to_string(),
|
||||
prefix_sharing_config: None,
|
||||
};
|
||||
|
||||
let conversations = generate_multi_turn_random(&tok, &cfg).unwrap();
|
||||
for conv in &conversations {
|
||||
assert_eq!(conv.turns.len(), 4, "expected exactly 4 turns");
|
||||
}
|
||||
println!("variable_turns_fixed checks passed!");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_per_turn_input_len_prefix_sharing() {
|
||||
let tok = crate::tokenizer::load_tokenizer("nvidia/Kimi-K2.5-NVFP4", false, None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Turn 0 input_len=1000, turns 1+ per_turn_input_len=600
|
||||
// global_len ≈ 100 (10%), conv_len ≈ 800 (80%), unique ≈ 100
|
||||
// per-turn unique ≈ 600 - 900 = negative → would error; use smaller ratios
|
||||
// global=0.05 (50), conv=0.5 (500), unique_t0=450, unique_t1=600-550=50
|
||||
let cfg = MultiTurnRandomConfig {
|
||||
num_conversations: 4,
|
||||
min_turns: 3,
|
||||
max_turns: 3,
|
||||
prefix_len: 0,
|
||||
input_len: 1000,
|
||||
per_turn_input_len: 600,
|
||||
output_len: 64,
|
||||
seed: 3,
|
||||
request_id_prefix: "test-".to_string(),
|
||||
prefix_sharing_config: Some(PrefixSharingConfig {
|
||||
global_ratio: 0.05,
|
||||
conversation_ratio: 0.50,
|
||||
}),
|
||||
};
|
||||
|
||||
let conversations = generate_multi_turn_random(&tok, &cfg).unwrap();
|
||||
assert_eq!(conversations.len(), 4);
|
||||
|
||||
let messages: Vec<Vec<&str>> = conversations
|
||||
.iter()
|
||||
.map(|c| c.turns.iter().map(|t| &*t.user_message).collect())
|
||||
.collect();
|
||||
|
||||
// Global prefix bytes shared across all turns of all conversations
|
||||
let all_msgs: Vec<&str> = messages.iter().flat_map(|v| v.iter().copied()).collect();
|
||||
let global_prefix = common_prefix_bytes(&all_msgs);
|
||||
assert!(global_prefix > 0, "Global prefix must be non-empty");
|
||||
|
||||
// Within each conversation, prefix grows (conv prefix longer than global)
|
||||
for (i, conv_msgs) in messages.iter().enumerate() {
|
||||
let conv_prefix = common_prefix_bytes(conv_msgs);
|
||||
assert!(
|
||||
conv_prefix > global_prefix,
|
||||
"Conv {i}: conv_prefix ({conv_prefix}) must exceed global ({global_prefix})"
|
||||
);
|
||||
}
|
||||
|
||||
// Turn 0 length ≈ 1000, turns 1+ ≈ 600
|
||||
for (i, conv) in conversations.iter().enumerate() {
|
||||
for (j, turn) in conv.turns.iter().enumerate() {
|
||||
let expected = if j == 0 { 1000usize } else { 600usize };
|
||||
let diff = (turn.user_message_len as i64 - expected as i64).abs();
|
||||
println!(
|
||||
"Conv {i} turn {j}: {} tokens (expected ~{expected})",
|
||||
turn.user_message_len
|
||||
);
|
||||
assert!(
|
||||
diff <= 10,
|
||||
"Conv {i} turn {j}: token len {} too far from {expected}",
|
||||
turn.user_message_len
|
||||
);
|
||||
}
|
||||
}
|
||||
println!("per_turn_input_len prefix-sharing checks passed!");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
//! Prefix repetition dataset: N distinct shared prefixes, each reused by
|
||||
//! `num_prompts / num_prefixes` requests with a fresh random suffix.
|
||||
//! The standard prefix-cache stress workload; mirrors Python's
|
||||
//! `PrefixRepetitionRandomDataset`.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rand::rngs::StdRng;
|
||||
use rand::seq::SliceRandom;
|
||||
use rand::{Rng, SeedableRng};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use super::SampleRequest;
|
||||
use super::random::gen_prompt_decode_to_target_len;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
/// Generate the prefix repetition dataset.
|
||||
///
|
||||
/// Like Python, `num_requests % num_prefixes` remainder requests are dropped:
|
||||
/// the total is `(num_requests / num_prefixes) * num_prefixes`.
|
||||
pub fn generate_prefix_repetition_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
num_requests: usize,
|
||||
prefix_len: usize,
|
||||
suffix_len: usize,
|
||||
num_prefixes: usize,
|
||||
output_len: usize,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
disable_shuffle: bool,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
let prompts_per_prefix = num_requests / num_prefixes;
|
||||
if prompts_per_prefix == 0 {
|
||||
return Err(BenchError::Config(format!(
|
||||
"num_prompts ({num_requests}) must be >= num_prefixes ({num_prefixes})"
|
||||
)));
|
||||
}
|
||||
let total = prompts_per_prefix * num_prefixes;
|
||||
if total != num_requests {
|
||||
tracing::info!(
|
||||
requested = num_requests,
|
||||
generated = total,
|
||||
prefixes = num_prefixes,
|
||||
prompts_per_prefix,
|
||||
dropped = num_requests - total,
|
||||
"adjusted prefix-repetition request count"
|
||||
);
|
||||
}
|
||||
|
||||
let allowed_tokens = tokenizer.get_allowed_tokens();
|
||||
if allowed_tokens.is_empty() {
|
||||
return Err(BenchError::Tokenizer("No allowed tokens found".into()));
|
||||
}
|
||||
let allowed_ref = &allowed_tokens;
|
||||
|
||||
// Exact-length random token block: decode -> re-encode -> converge to target.
|
||||
let gen_block = |target_len: usize, item_seed: u64| -> Result<Vec<u32>> {
|
||||
let mut rng = StdRng::seed_from_u64(item_seed);
|
||||
let tokens: Vec<u32> = (0..target_len)
|
||||
.map(|_| allowed_ref[rng.random_range(0..allowed_ref.len())])
|
||||
.collect();
|
||||
let (_, adjusted) =
|
||||
gen_prompt_decode_to_target_len(tokenizer, &tokens, target_len, false, allowed_ref)?;
|
||||
Ok(adjusted)
|
||||
};
|
||||
|
||||
// Generate the shared prefixes (one per group), then suffixes in parallel.
|
||||
let prefixes: Vec<Vec<u32>> = (0..num_prefixes)
|
||||
.map(|p| gen_block(prefix_len, seed.wrapping_add(0xF1F0).wrapping_add(p as u64)))
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
|
||||
let rid_prefix = request_id_prefix.to_string();
|
||||
let mut requests: Vec<SampleRequest> = (0..total)
|
||||
.into_par_iter()
|
||||
.map(|i| {
|
||||
let prefix_tokens = &prefixes[i / prompts_per_prefix];
|
||||
let suffix_tokens = gen_block(suffix_len, seed.wrapping_add(0xBEEF + i as u64))?;
|
||||
|
||||
let mut combined = Vec::with_capacity(prefix_tokens.len() + suffix_tokens.len());
|
||||
combined.extend_from_slice(prefix_tokens);
|
||||
combined.extend_from_slice(&suffix_tokens);
|
||||
let prompt = tokenizer.decode(&combined, true)?;
|
||||
|
||||
Ok(SampleRequest {
|
||||
prompt: Arc::from(prompt),
|
||||
prompt_len: combined.len(),
|
||||
expected_output_len: output_len,
|
||||
request_id: Some(format!("{rid_prefix}{i}")),
|
||||
..Default::default()
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
|
||||
// Interleave prefixes (Python shuffles too) so one prefix group isn't sent
|
||||
// as a contiguous burst.
|
||||
if !disable_shuffle {
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
requests.shuffle(&mut rng);
|
||||
}
|
||||
|
||||
Ok(requests)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// gpt2 via built-in tiktoken encoding — loads without network access.
|
||||
fn test_tokenizer() -> TokenizerKind {
|
||||
TokenizerKind::Tiktoken(
|
||||
crate::tiktoken::load_builtin_tiktoken("gpt2")
|
||||
.expect("gpt2 built-in tiktoken should always load without network"),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_prefix_repetition_structure() {
|
||||
let tok = test_tokenizer();
|
||||
// 7 requests / 3 prefixes -> 2 per prefix, 6 total (remainder dropped like Python)
|
||||
let reqs = generate_prefix_repetition_dataset(&tok, 7, 32, 16, 3, 64, 0, "t-", true)
|
||||
.expect("generation should succeed");
|
||||
assert_eq!(reqs.len(), 6);
|
||||
assert!(reqs.iter().all(|r| r.expected_output_len == 64));
|
||||
// Exact-length blocks: prompt_len == prefix + suffix
|
||||
assert!(
|
||||
reqs.iter().all(|r| r.prompt_len == 32 + 16),
|
||||
"lens: {:?}",
|
||||
reqs.iter().map(|r| r.prompt_len).collect::<Vec<_>>()
|
||||
);
|
||||
// Consecutive pairs (shuffle disabled) share a common prefix; requests
|
||||
// from different groups don't.
|
||||
let common = |a: &str, b: &str| -> usize {
|
||||
a.bytes().zip(b.bytes()).take_while(|(x, y)| x == y).count()
|
||||
};
|
||||
let same_group = common(&reqs[0].prompt, &reqs[1].prompt);
|
||||
let diff_group = common(&reqs[0].prompt, &reqs[2].prompt);
|
||||
assert!(
|
||||
same_group > diff_group,
|
||||
"same-group shared prefix ({same_group}) should exceed cross-group ({diff_group})"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_prefix_repetition_too_few_requests_errors() {
|
||||
let tok = test_tokenizer();
|
||||
assert!(generate_prefix_repetition_dataset(&tok, 2, 32, 16, 3, 64, 0, "t-", true).is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use indicatif::{ProgressBar, ProgressStyle};
|
||||
|
||||
const REPORT_INTERVAL: Duration = Duration::from_secs(10);
|
||||
|
||||
/// Reports row download progress to an interactive progress bar, or through
|
||||
/// periodic tracing events when the progress bar is hidden on a non-TTY.
|
||||
pub(super) struct RowDownloadReporter {
|
||||
progress: ProgressBar,
|
||||
next_report: Instant,
|
||||
}
|
||||
|
||||
impl RowDownloadReporter {
|
||||
/// Creates a reporter that emits non-TTY updates every 10 seconds.
|
||||
pub fn new() -> Self {
|
||||
let progress = ProgressBar::new(0);
|
||||
progress.set_style(
|
||||
ProgressStyle::with_template(
|
||||
"{spinner:.green} Fetching rows [{bar:30.cyan/blue}] {pos}/{len}",
|
||||
)
|
||||
.unwrap()
|
||||
.progress_chars("#>-"),
|
||||
);
|
||||
Self {
|
||||
progress,
|
||||
next_report: Instant::now() + REPORT_INTERVAL,
|
||||
}
|
||||
}
|
||||
|
||||
/// Updates the current row count and reports progress when due.
|
||||
pub fn update(&mut self, rows: usize, total: u64) {
|
||||
let rows = rows as u64;
|
||||
let total = total.max(rows);
|
||||
self.progress.set_length(total);
|
||||
self.progress.set_position(rows);
|
||||
|
||||
if self.should_report(Instant::now()) {
|
||||
tracing::info!(rows, total, "fetching dataset rows");
|
||||
}
|
||||
}
|
||||
|
||||
/// Clears the interactive progress bar after the download completes.
|
||||
pub fn finish(self) {
|
||||
self.progress.finish_and_clear();
|
||||
}
|
||||
|
||||
fn should_report(&mut self, now: Instant) -> bool {
|
||||
if !self.progress.is_hidden() || now < self.next_report {
|
||||
return false;
|
||||
}
|
||||
self.next_report = now + REPORT_INTERVAL;
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn hidden_reporter_uses_ten_second_deadline() {
|
||||
let start = Instant::now();
|
||||
let mut reporter = RowDownloadReporter {
|
||||
progress: ProgressBar::hidden(),
|
||||
next_report: start + REPORT_INTERVAL,
|
||||
};
|
||||
|
||||
assert!(!reporter.should_report(start + Duration::from_secs(9)));
|
||||
assert!(reporter.should_report(start + Duration::from_secs(10)));
|
||||
assert!(!reporter.should_report(start + Duration::from_secs(19)));
|
||||
assert!(reporter.should_report(start + Duration::from_secs(20)));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,497 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rayon::prelude::*;
|
||||
|
||||
use super::SampleRequest;
|
||||
use crate::config::RangeRatio;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
/// Generate random dataset with rayon parallelism.
|
||||
///
|
||||
/// This is the key performance win — Python does sequential tokenizer calls
|
||||
/// while Rust parallelizes across CPU cores with native tokenizer speed.
|
||||
///
|
||||
/// Mirrors Python's RandomDataset.sample() from datasets.py:470-560.
|
||||
pub fn generate_random_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
num_requests: usize,
|
||||
input_len: usize,
|
||||
output_len: usize,
|
||||
prefix_len: usize,
|
||||
range_ratio: RangeRatio,
|
||||
cache_hit_fraction: f64,
|
||||
cache_ratio: f64,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
use_token_ids: bool,
|
||||
batch_size: usize,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
let vocab_size = tokenizer.vocab_size();
|
||||
let allowed_tokens = tokenizer.get_allowed_tokens();
|
||||
if allowed_tokens.is_empty() {
|
||||
return Err(BenchError::Tokenizer("No allowed tokens found".into()));
|
||||
}
|
||||
|
||||
if batch_size > 1 && use_token_ids {
|
||||
return Err(BenchError::Config(
|
||||
"--random-batch-size > 1 is not supported with --prompt-token-ids".into(),
|
||||
));
|
||||
}
|
||||
|
||||
let num_special = tokenizer.num_special_tokens_to_add();
|
||||
let real_input_len = input_len.saturating_sub(num_special);
|
||||
|
||||
// Python semantics: sample uniformly from [len*(1-r), len*(1+r)].
|
||||
let (input_low, input_high) = range_ratio.input_bounds(real_input_len);
|
||||
let (output_low, output_high) = range_ratio.output_bounds(output_len);
|
||||
if !range_ratio.is_fixed() {
|
||||
tracing::info!(
|
||||
input_low,
|
||||
input_high,
|
||||
output_low,
|
||||
output_high,
|
||||
"sampling random request lengths"
|
||||
);
|
||||
}
|
||||
|
||||
// Bimodal prefix-cache mode: a fraction of prompts (warm) reuse a shared cached
|
||||
// prefix covering `cache_ratio` of their length; the rest (cold) are fully unique.
|
||||
// Models e.g. "80% of prompts have 95% of input cached" with
|
||||
// --random-cache-hit-fraction 0.8 --random-cache-ratio 0.95. In this mode
|
||||
// --random-input-len is the TOTAL prompt length L (the cached prefix is part of L),
|
||||
// and --random-prefix-len is ignored.
|
||||
let bimodal = cache_hit_fraction > 0.0 && cache_ratio > 0.0;
|
||||
if bimodal {
|
||||
if !use_token_ids {
|
||||
return Err(BenchError::Config(
|
||||
"bimodal prefix-cache (--random-cache-hit-fraction) requires --prompt-token-ids \
|
||||
so warm prompts send identical token IDs and actually hit the prefix cache"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
if cache_hit_fraction > 1.0 || cache_ratio > 1.0 {
|
||||
return Err(BenchError::Config(
|
||||
"--random-cache-hit-fraction and --random-cache-ratio must be in [0, 1]".into(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
// Length of the shared cached base prefix.
|
||||
let base_len = if bimodal {
|
||||
((input_high as f64) * cache_ratio).ceil() as usize
|
||||
} else {
|
||||
prefix_len
|
||||
};
|
||||
|
||||
// Validate (non-bimodal keeps the original check)
|
||||
if !bimodal {
|
||||
let min_total = prefix_len + input_low;
|
||||
if min_total < 1 {
|
||||
return Err(BenchError::Config(format!(
|
||||
"--random-input-len too small: with {num_special} special tokens and \
|
||||
range_ratio={:?}, minimum total input is {min_total}",
|
||||
range_ratio
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
// Generate the shared base prefix once (sequential, only happens once).
|
||||
let prefix_token_ids = if base_len > 0 {
|
||||
generate_prefix(tokenizer, &allowed_tokens, base_len, seed)?
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
// Pre-generate per-request sampling params using deterministic RNG
|
||||
use rand::rngs::StdRng;
|
||||
use rand::{Rng, SeedableRng};
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
|
||||
struct RequestParams {
|
||||
cached_len: usize, // tokens taken from the shared base (cache-hittable)
|
||||
suffix_len: usize, // unique tokens appended after the cached prefix
|
||||
output_len: usize,
|
||||
offset: usize,
|
||||
}
|
||||
|
||||
let params: Vec<RequestParams> = (0..num_requests)
|
||||
.map(|_| {
|
||||
let ol = if output_low == output_high {
|
||||
output_low
|
||||
} else {
|
||||
rng.random_range(output_low..=output_high)
|
||||
};
|
||||
let off = rng.random_range(0..vocab_size as usize);
|
||||
if bimodal {
|
||||
// Total length L from the input distribution; prefix is part of L.
|
||||
let l = if input_low == input_high {
|
||||
input_low
|
||||
} else {
|
||||
rng.random_range(input_low..=input_high)
|
||||
};
|
||||
let warm = rng.random::<f64>() < cache_hit_fraction;
|
||||
let cached = if warm {
|
||||
(((l as f64) * cache_ratio).round() as usize).min(base_len).min(l)
|
||||
} else {
|
||||
0
|
||||
};
|
||||
RequestParams {
|
||||
cached_len: cached,
|
||||
suffix_len: l - cached,
|
||||
output_len: ol,
|
||||
offset: off,
|
||||
}
|
||||
} else {
|
||||
// Original behavior: full shared prefix + variable unique input.
|
||||
let il = if input_low == input_high {
|
||||
input_low
|
||||
} else {
|
||||
rng.random_range(input_low..=input_high)
|
||||
};
|
||||
RequestParams {
|
||||
cached_len: prefix_len,
|
||||
suffix_len: il,
|
||||
output_len: ol,
|
||||
offset: off,
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Phase 1: Generate all token sequences (parallel, fast — just array ops)
|
||||
let prefix_ref = &prefix_token_ids;
|
||||
let allowed_ref = &allowed_tokens;
|
||||
let rid_prefix = request_id_prefix.to_string();
|
||||
|
||||
let token_sequences: Vec<Vec<u32>> = params
|
||||
.par_iter()
|
||||
.enumerate()
|
||||
.map(|(i, p)| {
|
||||
let at_len = allowed_ref.len();
|
||||
let mut seq = Vec::with_capacity(p.cached_len + p.suffix_len);
|
||||
seq.extend_from_slice(&prefix_ref[..p.cached_len]);
|
||||
for j in 0..p.suffix_len {
|
||||
seq.push(allowed_ref[(p.offset + i + j) % at_len]);
|
||||
}
|
||||
seq
|
||||
})
|
||||
.collect();
|
||||
|
||||
let target_lens: Vec<usize> = params.iter().map(|p| p.cached_len + p.suffix_len).collect();
|
||||
|
||||
if use_token_ids {
|
||||
// Fast path: store token IDs directly. The completions backend sends
|
||||
// them as `"prompt": [id1, id2, ...]`, bypassing both client-side decode
|
||||
// and server-side tokenization. Token counts are exact by construction.
|
||||
let result: Vec<SampleRequest> = token_sequences
|
||||
.into_par_iter()
|
||||
.enumerate()
|
||||
.map(|(i, tokens)| SampleRequest {
|
||||
prompt: Arc::from(""),
|
||||
prompt_len: target_lens[i],
|
||||
expected_output_len: params[i].output_len,
|
||||
request_id: Some(format!("{rid_prefix}{i}")),
|
||||
prompt_token_ids: Some(Arc::from(tokens)),
|
||||
..Default::default()
|
||||
})
|
||||
.collect();
|
||||
Ok(result)
|
||||
} else {
|
||||
// Default path: decode tokens to text, re-encode,
|
||||
// truncate to target length, decode again. Sends text prompts for maximum
|
||||
let result: Vec<SampleRequest> = token_sequences
|
||||
.into_par_iter()
|
||||
.enumerate()
|
||||
.map(|(i, tokens)| {
|
||||
let target = target_lens[i];
|
||||
// decode → encode → truncate → decode
|
||||
let prompt_text = tokenizer.decode(&tokens, true)?;
|
||||
let mut re_encoded = tokenizer.encode(&prompt_text, false)?;
|
||||
re_encoded.truncate(target);
|
||||
let prompt = tokenizer.decode(&re_encoded, true)?;
|
||||
let prompt_len = re_encoded.len();
|
||||
Ok(SampleRequest {
|
||||
prompt: Arc::from(prompt),
|
||||
prompt_len,
|
||||
expected_output_len: params[i].output_len,
|
||||
request_id: Some(format!("{rid_prefix}{i}")),
|
||||
..Default::default()
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
Ok(super::batch_requests(result, batch_size, request_id_prefix))
|
||||
}
|
||||
}
|
||||
|
||||
fn generate_prefix(
|
||||
tokenizer: &TokenizerKind,
|
||||
allowed_tokens: &[u32],
|
||||
prefix_len: usize,
|
||||
seed: u64,
|
||||
) -> Result<Vec<u32>> {
|
||||
use rand::rngs::StdRng;
|
||||
use rand::{Rng, SeedableRng};
|
||||
|
||||
let mut rng = StdRng::seed_from_u64(seed.wrapping_add(0xDEAD));
|
||||
let tokens: Vec<u32> = (0..prefix_len)
|
||||
.map(|_| allowed_tokens[rng.random_range(0..allowed_tokens.len())])
|
||||
.collect();
|
||||
|
||||
let (_, adjusted) =
|
||||
gen_prompt_decode_to_target_len(tokenizer, &tokens, prefix_len, false, allowed_tokens)?;
|
||||
Ok(adjusted)
|
||||
}
|
||||
|
||||
/// Ensure decoded-then-encoded prompt length matches the target.
|
||||
///
|
||||
/// Mirrors Python's `gen_prompt_decode_to_target_len` from datasets.py:381-435.
|
||||
pub(crate) fn gen_prompt_decode_to_target_len(
|
||||
tokenizer: &TokenizerKind,
|
||||
token_sequence: &[u32],
|
||||
target_len: usize,
|
||||
add_special_tokens: bool,
|
||||
allowed_tokens: &[u32],
|
||||
) -> Result<(String, Vec<u32>)> {
|
||||
let max_retry = 20;
|
||||
let mut tokens = token_sequence.to_vec();
|
||||
|
||||
for retry in 0..=max_retry {
|
||||
let prompt = tokenizer.decode(&tokens, true)?;
|
||||
tokens = tokenizer.encode(&prompt, add_special_tokens)?;
|
||||
|
||||
if retry >= max_retry {
|
||||
if tokens.len() != target_len {
|
||||
return Err(BenchError::Tokenizer(format!(
|
||||
"Token length mismatch after {max_retry} retries: \
|
||||
target={target_len}, actual={}. \
|
||||
encode/decode roundtrip cannot converge.",
|
||||
tokens.len()
|
||||
)));
|
||||
}
|
||||
return Ok((prompt, tokens));
|
||||
}
|
||||
|
||||
if tokens.len() == target_len {
|
||||
return Ok((prompt, tokens));
|
||||
} else if tokens.len() < target_len {
|
||||
// Pad with tokens from the allowed set (UTF-8-safe for tiktoken)
|
||||
let needed = target_len - tokens.len();
|
||||
if allowed_tokens.is_empty() {
|
||||
let vocab_size = tokenizer.vocab_size() as usize;
|
||||
for j in 0..needed {
|
||||
tokens.push(((tokens.len() + j) % vocab_size) as u32);
|
||||
}
|
||||
} else {
|
||||
for j in 0..needed {
|
||||
tokens.push(allowed_tokens[(tokens.len() + j) % allowed_tokens.len()]);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Truncate
|
||||
tokens.truncate(target_len);
|
||||
}
|
||||
}
|
||||
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::tokenizer;
|
||||
|
||||
// Integration test requires a tokenizer, so only run with --ignored
|
||||
#[test]
|
||||
#[ignore]
|
||||
fn test_generate_random_dataset_token_ids() {
|
||||
let tokenizer =
|
||||
TokenizerKind::Tiktoken(crate::tiktoken::load_builtin_tiktoken("gpt2").unwrap());
|
||||
let requests = generate_random_dataset(
|
||||
&tokenizer,
|
||||
10, // num_requests
|
||||
128, // input_len
|
||||
32, // output_len
|
||||
0, // prefix_len
|
||||
RangeRatio {
|
||||
input: 0.0,
|
||||
output: 0.0,
|
||||
}, // range_ratio (0.0 = fixed length)
|
||||
0.0, // cache_hit_fraction (0 = bimodal off)
|
||||
0.0, // cache_ratio
|
||||
42, // seed
|
||||
"test-",
|
||||
true, // use_token_ids
|
||||
1, // batch_size
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(requests.len(), 10);
|
||||
for req in &requests {
|
||||
assert!(req.prompt_token_ids.is_some());
|
||||
assert_eq!(req.prompt_token_ids.as_ref().unwrap().len(), req.prompt_len);
|
||||
assert!(req.prompt_len > 0);
|
||||
assert_eq!(req.expected_output_len, 32);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore]
|
||||
fn test_generate_random_dataset_text() {
|
||||
let tokenizer =
|
||||
TokenizerKind::Tiktoken(crate::tiktoken::load_builtin_tiktoken("gpt2").unwrap());
|
||||
let requests = generate_random_dataset(
|
||||
&tokenizer,
|
||||
10, // num_requests
|
||||
128, // input_len
|
||||
32, // output_len
|
||||
0, // prefix_len
|
||||
RangeRatio {
|
||||
input: 0.0,
|
||||
output: 0.0,
|
||||
}, // range_ratio (0.0 = fixed length)
|
||||
0.0, // cache_hit_fraction (0 = bimodal off)
|
||||
0.0, // cache_ratio
|
||||
42, // seed
|
||||
"test-",
|
||||
false, // use_token_ids = false → text prompts
|
||||
1, // batch_size
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(requests.len(), 10);
|
||||
for req in &requests {
|
||||
assert!(req.prompt_token_ids.is_none());
|
||||
assert!(!req.prompt.is_empty());
|
||||
assert!(req.prompt_len > 0);
|
||||
assert!(req.prompt_len <= 128);
|
||||
assert_eq!(req.expected_output_len, 32);
|
||||
}
|
||||
}
|
||||
|
||||
/// Test that generated prompts have EXACT target token length (token ID mode).
|
||||
#[test]
|
||||
#[ignore]
|
||||
fn test_token_length_exact_local() {
|
||||
let tokenizer =
|
||||
TokenizerKind::Tiktoken(crate::tiktoken::load_builtin_tiktoken("gpt2").unwrap());
|
||||
let target_len = 512;
|
||||
let requests = generate_random_dataset(
|
||||
&tokenizer,
|
||||
50,
|
||||
target_len,
|
||||
64,
|
||||
0,
|
||||
RangeRatio {
|
||||
input: 0.0,
|
||||
output: 0.0,
|
||||
},
|
||||
0.0,
|
||||
0.0,
|
||||
123,
|
||||
"len-test-",
|
||||
true,
|
||||
1,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
for (i, req) in requests.iter().enumerate() {
|
||||
let token_ids = req.prompt_token_ids.as_ref().expect("should have token IDs");
|
||||
assert_eq!(
|
||||
token_ids.len(),
|
||||
target_len,
|
||||
"Request {i}: expected {target_len} token IDs, got {}",
|
||||
token_ids.len()
|
||||
);
|
||||
assert_eq!(req.prompt_len, target_len);
|
||||
}
|
||||
}
|
||||
|
||||
/// Test that tiktoken tokenizer produces exact target token lengths (token ID mode).
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_token_length_exact_tiktoken() {
|
||||
// Use Qwen2.5 which has a tiktoken-format tokenizer
|
||||
let tokenizer = tokenizer::load_tokenizer("Qwen/Qwen2.5-0.5B", false, None).await;
|
||||
let tokenizer = match tokenizer {
|
||||
Ok(t) => t,
|
||||
Err(e) => {
|
||||
eprintln!("Skipping tiktoken test (tokenizer unavailable): {e}");
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// Verify it's actually a tiktoken tokenizer or local — either way test convergence
|
||||
let target_len = 256;
|
||||
let requests = generate_random_dataset(
|
||||
&tokenizer,
|
||||
20,
|
||||
target_len,
|
||||
32,
|
||||
0,
|
||||
RangeRatio {
|
||||
input: 0.0,
|
||||
output: 0.0,
|
||||
},
|
||||
0.0,
|
||||
0.0,
|
||||
42,
|
||||
"tiktoken-test-",
|
||||
true,
|
||||
1,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
for (i, req) in requests.iter().enumerate() {
|
||||
let token_ids = req.prompt_token_ids.as_ref().expect("should have token IDs");
|
||||
assert_eq!(
|
||||
token_ids.len(),
|
||||
target_len,
|
||||
"Request {i}: expected {target_len} token IDs, got {}",
|
||||
token_ids.len()
|
||||
);
|
||||
assert_eq!(req.prompt_len, target_len);
|
||||
}
|
||||
}
|
||||
|
||||
/// Test encode/decode roundtrip stability for tiktoken.
|
||||
/// After one decode→encode cycle with UTF-8-safe tokens, length must not drift.
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_tiktoken_roundtrip_stability() {
|
||||
let tokenizer = tokenizer::load_tokenizer("Qwen/Qwen2.5-0.5B", false, None).await;
|
||||
let tokenizer = match tokenizer {
|
||||
Ok(t) => t,
|
||||
Err(e) => {
|
||||
eprintln!("Skipping roundtrip test (tokenizer unavailable): {e}");
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
let allowed = tokenizer.get_allowed_tokens();
|
||||
assert!(!allowed.is_empty(), "allowed tokens should not be empty");
|
||||
|
||||
// Build a sequence from allowed tokens only
|
||||
use rand::rngs::StdRng;
|
||||
use rand::{Rng, SeedableRng};
|
||||
let mut rng = StdRng::seed_from_u64(99);
|
||||
let seq: Vec<u32> = (0..512).map(|_| allowed[rng.random_range(0..allowed.len())]).collect();
|
||||
|
||||
let decoded = tokenizer.decode(&seq, true).unwrap();
|
||||
let re_encoded = tokenizer.encode(&decoded, false).unwrap();
|
||||
let re_decoded = tokenizer.decode(&re_encoded, true).unwrap();
|
||||
let re_re_encoded = tokenizer.encode(&re_decoded, false).unwrap();
|
||||
|
||||
// After first cycle, length should stabilize
|
||||
assert_eq!(
|
||||
re_encoded.len(),
|
||||
re_re_encoded.len(),
|
||||
"Roundtrip should stabilize: first re-encode={}, second re-encode={}",
|
||||
re_encoded.len(),
|
||||
re_re_encoded.len()
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,657 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::io::Cursor;
|
||||
use std::sync::Arc;
|
||||
|
||||
use base64::Engine as _;
|
||||
use rand::rngs::StdRng;
|
||||
use rand::{Rng, SeedableRng};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use super::SampleRequest;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
/// A bucket key: (height, width, num_frames). num_frames=1 means image, >1 means video.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MmBucketKey {
|
||||
pub height: u32,
|
||||
pub width: u32,
|
||||
pub num_frames: u32,
|
||||
}
|
||||
|
||||
/// Per-modality hard caps.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MmLimitPerPrompt {
|
||||
pub image: usize,
|
||||
pub video: usize,
|
||||
}
|
||||
|
||||
impl Default for MmLimitPerPrompt {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
image: 255,
|
||||
video: 1,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse the limit-mm-per-prompt JSON string, e.g. `{"image": 3, "video": 0}`.
|
||||
pub fn parse_limit_mm_per_prompt(s: &str) -> Result<MmLimitPerPrompt> {
|
||||
let v: serde_json::Value = serde_json::from_str(s)
|
||||
.map_err(|e| BenchError::Config(format!("Invalid --random-mm-limit-mm-per-prompt: {e}")))?;
|
||||
let obj = v.as_object().ok_or_else(|| {
|
||||
BenchError::Config("--random-mm-limit-mm-per-prompt must be a JSON object".into())
|
||||
})?;
|
||||
let image = obj.get("image").and_then(|v| v.as_u64()).unwrap_or(255) as usize;
|
||||
let video = obj.get("video").and_then(|v| v.as_u64()).unwrap_or(1) as usize;
|
||||
Ok(MmLimitPerPrompt { image, video })
|
||||
}
|
||||
|
||||
/// Parse the bucket config string in Python-style syntax.
|
||||
///
|
||||
/// Accepts: `{(256,256,1): 0.5, (720,1280,1): 0.5}`
|
||||
/// Each key is `(height, width, num_frames)` and value is the probability weight.
|
||||
pub fn parse_bucket_config(s: &str) -> Result<Vec<(MmBucketKey, f64)>> {
|
||||
let trimmed = s.trim();
|
||||
let inner = trimmed
|
||||
.strip_prefix('{')
|
||||
.and_then(|s| s.strip_suffix('}'))
|
||||
.ok_or_else(|| BenchError::Config("Bucket config must be wrapped in {}".into()))?;
|
||||
|
||||
let mut buckets = Vec::new();
|
||||
let chars: Vec<char> = inner.chars().collect();
|
||||
let len = chars.len();
|
||||
let mut i = 0;
|
||||
|
||||
while i < len {
|
||||
// Skip whitespace and commas
|
||||
while i < len && (chars[i].is_whitespace() || chars[i] == ',') {
|
||||
i += 1;
|
||||
}
|
||||
if i >= len {
|
||||
break;
|
||||
}
|
||||
|
||||
// Expect '('
|
||||
if chars[i] != '(' {
|
||||
return Err(BenchError::Config(format!(
|
||||
"Expected '(' in bucket config at position {i}"
|
||||
)));
|
||||
}
|
||||
i += 1;
|
||||
|
||||
// Read until ')'
|
||||
let tuple_start = i;
|
||||
while i < len && chars[i] != ')' {
|
||||
i += 1;
|
||||
}
|
||||
if i >= len {
|
||||
return Err(BenchError::Config("Unclosed '(' in bucket config".into()));
|
||||
}
|
||||
let tuple_str: String = chars[tuple_start..i].iter().collect();
|
||||
i += 1; // skip ')'
|
||||
|
||||
// Skip whitespace, then expect ':'
|
||||
while i < len && chars[i].is_whitespace() {
|
||||
i += 1;
|
||||
}
|
||||
if i >= len || chars[i] != ':' {
|
||||
return Err(BenchError::Config(
|
||||
"Expected ':' after tuple in bucket config".into(),
|
||||
));
|
||||
}
|
||||
i += 1;
|
||||
|
||||
// Skip whitespace
|
||||
while i < len && chars[i].is_whitespace() {
|
||||
i += 1;
|
||||
}
|
||||
|
||||
// Read the probability value until ',' or end
|
||||
let val_start = i;
|
||||
while i < len && chars[i] != ',' {
|
||||
i += 1;
|
||||
}
|
||||
let val_str: String = chars[val_start..i].iter().collect();
|
||||
|
||||
// Parse tuple
|
||||
let parts: Vec<&str> = tuple_str.split(',').collect();
|
||||
if parts.len() != 3 {
|
||||
return Err(BenchError::Config(format!(
|
||||
"Bucket key must have 3 values (height,width,num_frames), got: ({tuple_str})"
|
||||
)));
|
||||
}
|
||||
|
||||
let height: u32 = parts[0].trim().parse().map_err(|_| {
|
||||
BenchError::Config(format!(
|
||||
"Invalid height in bucket config: '{}'",
|
||||
parts[0].trim()
|
||||
))
|
||||
})?;
|
||||
let width: u32 = parts[1].trim().parse().map_err(|_| {
|
||||
BenchError::Config(format!(
|
||||
"Invalid width in bucket config: '{}'",
|
||||
parts[1].trim()
|
||||
))
|
||||
})?;
|
||||
let num_frames: u32 = parts[2].trim().parse().map_err(|_| {
|
||||
BenchError::Config(format!(
|
||||
"Invalid num_frames in bucket config: '{}'",
|
||||
parts[2].trim()
|
||||
))
|
||||
})?;
|
||||
let prob: f64 = val_str.trim().parse().map_err(|_| {
|
||||
BenchError::Config(format!(
|
||||
"Invalid probability in bucket config: '{}'",
|
||||
val_str.trim()
|
||||
))
|
||||
})?;
|
||||
|
||||
if prob < 0.0 {
|
||||
return Err(BenchError::Config(format!(
|
||||
"Bucket probability must be non-negative, got: {prob}"
|
||||
)));
|
||||
}
|
||||
|
||||
buckets.push((
|
||||
MmBucketKey {
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
},
|
||||
prob,
|
||||
));
|
||||
}
|
||||
|
||||
if buckets.is_empty() {
|
||||
return Err(BenchError::Config(
|
||||
"Bucket config must have at least one entry".into(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(buckets)
|
||||
}
|
||||
|
||||
/// JSON fragment prefix/suffix for image content blocks.
|
||||
const IMG_JSON_PREFIX: &str = r#"{"type":"image_url","image_url":{"url":"data:image/jpeg;base64,"#;
|
||||
const IMG_JSON_SUFFIX: &str = r#""}}"#;
|
||||
|
||||
/// Generate a synthetic random JPEG image and return it as a pre-serialized JSON fragment.
|
||||
///
|
||||
/// Builds the complete JSON string in a single allocation:
|
||||
/// `{"type":"image_url","image_url":{"url":"data:image/jpeg;base64,<b64>"}}`
|
||||
///
|
||||
/// The base64 data is written directly into the final string — no intermediate
|
||||
/// String or format!() copy.
|
||||
fn generate_random_image(width: u32, height: u32, rng: &mut StdRng) -> Result<Arc<str>> {
|
||||
let pixel_count = (width as usize) * (height as usize) * 3;
|
||||
let mut pixels = vec![0u8; pixel_count];
|
||||
rng.fill(pixels.as_mut_slice());
|
||||
|
||||
let img = image::RgbImage::from_raw(width, height, pixels)
|
||||
.ok_or_else(|| BenchError::Config("Failed to create image from random pixels".into()))?;
|
||||
|
||||
// Pre-allocate JPEG buffer (random pixels compress poorly, estimate ~60% of raw)
|
||||
let estimated_jpeg = pixel_count * 3 / 5;
|
||||
let mut buf = Cursor::new(Vec::with_capacity(estimated_jpeg));
|
||||
img.write_to(&mut buf, image::ImageFormat::Jpeg)
|
||||
.map_err(|e| BenchError::Config(format!("Failed to encode JPEG: {e}")))?;
|
||||
|
||||
let jpeg_bytes = buf.into_inner();
|
||||
|
||||
// Pre-compute exact output size: prefix + base64_len + suffix
|
||||
let b64_len = jpeg_bytes.len().div_ceil(3) * 4;
|
||||
let total_len = IMG_JSON_PREFIX.len() + b64_len + IMG_JSON_SUFFIX.len();
|
||||
|
||||
// Single allocation: write base64 directly into the JSON fragment string
|
||||
let mut json_fragment = String::with_capacity(total_len);
|
||||
json_fragment.push_str(IMG_JSON_PREFIX);
|
||||
base64::engine::general_purpose::STANDARD.encode_string(&jpeg_bytes, &mut json_fragment);
|
||||
json_fragment.push_str(IMG_JSON_SUFFIX);
|
||||
|
||||
Ok(Arc::from(json_fragment))
|
||||
}
|
||||
|
||||
/// Sample multimodal items for a single request.
|
||||
///
|
||||
/// Returns a list of (height, width, num_frames) tuples.
|
||||
fn sample_mm_items(
|
||||
rng: &mut StdRng,
|
||||
min_items: usize,
|
||||
max_items: usize,
|
||||
buckets: &[(MmBucketKey, f64)],
|
||||
limit: &MmLimitPerPrompt,
|
||||
) -> Vec<MmBucketKey> {
|
||||
let num_items = if min_items == max_items {
|
||||
min_items
|
||||
} else {
|
||||
rng.random_range(min_items..=max_items)
|
||||
};
|
||||
|
||||
// Filter to non-zero probability buckets
|
||||
let active_buckets: Vec<&(MmBucketKey, f64)> =
|
||||
buckets.iter().filter(|(_, p)| *p > 0.0).collect();
|
||||
if active_buckets.is_empty() || num_items == 0 {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let total_weight: f64 = active_buckets.iter().map(|(_, p)| p).sum();
|
||||
if total_weight <= 0.0 {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let mut result = Vec::with_capacity(num_items);
|
||||
let mut image_count = 0usize;
|
||||
let mut video_count = 0usize;
|
||||
|
||||
for _ in 0..num_items {
|
||||
// Build normalized weights considering remaining capacity
|
||||
let mut weights: Vec<f64> = Vec::with_capacity(active_buckets.len());
|
||||
for (key, prob) in &active_buckets {
|
||||
let is_video = key.num_frames > 1;
|
||||
let at_limit = if is_video {
|
||||
video_count >= limit.video
|
||||
} else {
|
||||
image_count >= limit.image
|
||||
};
|
||||
weights.push(if at_limit { 0.0 } else { *prob });
|
||||
}
|
||||
|
||||
let w_total: f64 = weights.iter().sum();
|
||||
if w_total <= 0.0 {
|
||||
break; // All modalities at limit
|
||||
}
|
||||
|
||||
// Weighted random selection (strict `<` to avoid selecting zero-weight buckets)
|
||||
let r = rng.random::<f64>() * w_total;
|
||||
let mut cumulative = 0.0;
|
||||
// Default to last non-zero-weight bucket (floating-point accumulation fallback)
|
||||
let mut selected_idx = weights.iter().rposition(|w| *w > 0.0).unwrap_or(0);
|
||||
for (i, w) in weights.iter().enumerate() {
|
||||
cumulative += w;
|
||||
if r < cumulative {
|
||||
selected_idx = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
let (key, _) = &active_buckets[selected_idx];
|
||||
if key.num_frames > 1 {
|
||||
video_count += 1;
|
||||
} else {
|
||||
image_count += 1;
|
||||
}
|
||||
result.push(key.clone());
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// Generate random multimodal dataset.
|
||||
///
|
||||
/// Mirrors Python's RandomMultiModalDataset.sample() from datasets.py.
|
||||
/// Generates text prompts with exact token lengths and random images/videos.
|
||||
pub fn generate_random_mm_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
num_requests: usize,
|
||||
input_len: usize,
|
||||
output_len: usize,
|
||||
prefix_len: usize,
|
||||
range_ratio: crate::config::RangeRatio,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
base_items_per_request: usize,
|
||||
num_mm_items_range_ratio: f64,
|
||||
limit: &MmLimitPerPrompt,
|
||||
buckets: &[(MmBucketKey, f64)],
|
||||
enable_multimodal_chat: bool,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
if !(0.0..=1.0).contains(&num_mm_items_range_ratio) {
|
||||
return Err(BenchError::Config(
|
||||
"num_mm_items_range_ratio must be in [0, 1]".into(),
|
||||
));
|
||||
}
|
||||
|
||||
// Check for video buckets with non-zero probability
|
||||
for (key, prob) in buckets {
|
||||
if key.num_frames > 1 && *prob > 0.0 {
|
||||
return Err(BenchError::Config(
|
||||
"Video generation (num_frames > 1) is not yet supported in Rust. \
|
||||
Set video bucket probabilities to 0.0."
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
// Compute item count bounds
|
||||
let n = base_items_per_request as f64;
|
||||
let r = num_mm_items_range_ratio;
|
||||
let min_items = (n * (1.0 - r)).floor().max(0.0) as usize;
|
||||
let max_items = (n * (1.0 + r)).ceil() as usize;
|
||||
// Clamp to total modality limit
|
||||
let total_limit = limit.image + limit.video;
|
||||
let max_items = max_items.min(total_limit);
|
||||
let min_items = min_items.min(max_items);
|
||||
|
||||
let vocab_size = tokenizer.vocab_size();
|
||||
let allowed_tokens = tokenizer.get_allowed_tokens();
|
||||
if allowed_tokens.is_empty() {
|
||||
return Err(BenchError::Tokenizer("No allowed tokens found".into()));
|
||||
}
|
||||
|
||||
let num_special = tokenizer.num_special_tokens_to_add();
|
||||
let real_input_len = input_len.saturating_sub(num_special);
|
||||
|
||||
// Python semantics: sample uniformly from [len*(1-r), len*(1+r)].
|
||||
let (input_low, input_high) = range_ratio.input_bounds(real_input_len);
|
||||
let (output_low, output_high) = range_ratio.output_bounds(output_len);
|
||||
|
||||
// Pre-generate per-request params
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
|
||||
struct RequestParams {
|
||||
input_len: usize,
|
||||
output_len: usize,
|
||||
offset: usize,
|
||||
}
|
||||
|
||||
let params: Vec<RequestParams> = (0..num_requests)
|
||||
.map(|_| {
|
||||
let il = if input_low == input_high {
|
||||
input_low
|
||||
} else {
|
||||
rng.random_range(input_low..=input_high)
|
||||
};
|
||||
let ol = if output_low == output_high {
|
||||
output_low
|
||||
} else {
|
||||
rng.random_range(output_low..=output_high)
|
||||
};
|
||||
let off = rng.random_range(0..vocab_size as usize);
|
||||
RequestParams {
|
||||
input_len: il,
|
||||
output_len: ol,
|
||||
offset: off,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Pre-generate multimodal item configs per request
|
||||
let mm_configs: Vec<Vec<MmBucketKey>> = (0..num_requests)
|
||||
.map(|_| sample_mm_items(&mut rng, min_items, max_items, buckets, limit))
|
||||
.collect();
|
||||
|
||||
// Generate text prompts (need text for chat backend, not just token IDs)
|
||||
let prefix_token_ids = if prefix_len > 0 {
|
||||
generate_prefix(tokenizer, &allowed_tokens, prefix_len, seed)?
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
// Generate token sequences
|
||||
let prefix_ref = &prefix_token_ids;
|
||||
let allowed_ref = &allowed_tokens;
|
||||
|
||||
let token_sequences: Vec<Vec<u32>> = params
|
||||
.par_iter()
|
||||
.enumerate()
|
||||
.map(|(i, p)| {
|
||||
let at_len = allowed_ref.len();
|
||||
let mut seq = Vec::with_capacity(prefix_ref.len() + p.input_len);
|
||||
seq.extend_from_slice(prefix_ref);
|
||||
for j in 0..p.input_len {
|
||||
seq.push(allowed_ref[(p.offset + i + j) % at_len]);
|
||||
}
|
||||
seq
|
||||
})
|
||||
.collect();
|
||||
|
||||
let target_lens: Vec<usize> = params.iter().map(|p| prefix_len + p.input_len).collect();
|
||||
|
||||
// Decode tokens to text (chat backend needs text prompts for multimodal)
|
||||
let prompts: Result<Vec<String>> = token_sequences
|
||||
.into_par_iter()
|
||||
.enumerate()
|
||||
.map(|(i, tokens)| {
|
||||
let (text, _adjusted) = super::random::gen_prompt_decode_to_target_len(
|
||||
tokenizer,
|
||||
&tokens,
|
||||
target_lens[i],
|
||||
false,
|
||||
allowed_ref,
|
||||
)?;
|
||||
Ok(text)
|
||||
})
|
||||
.collect();
|
||||
let prompts = prompts?;
|
||||
|
||||
// Generate images for each request (parallel per request)
|
||||
// Each request gets its own RNG seeded deterministically.
|
||||
let rid_prefix = request_id_prefix.to_string();
|
||||
let result: Vec<SampleRequest> = prompts
|
||||
.into_par_iter()
|
||||
.enumerate()
|
||||
.map(|(i, prompt)| {
|
||||
let mut item_rng =
|
||||
StdRng::seed_from_u64(seed.wrapping_add(i as u64).wrapping_add(0xBEEF));
|
||||
let mm_items: Vec<Arc<str>> = mm_configs[i]
|
||||
.iter()
|
||||
.map(|key| {
|
||||
generate_random_image(key.width, key.height, &mut item_rng)
|
||||
.expect("Image generation should not fail")
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mm_content: Option<Arc<[Arc<str>]>> = if mm_items.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(Arc::from(mm_items))
|
||||
};
|
||||
|
||||
// --enable-multimodal-chat: pre-build the full chat `messages` array
|
||||
// (text part + mm items) at dataset time, mirroring Python's
|
||||
// apply_multimodal_chat_transformation. mm content moves inside the
|
||||
// messages string; the backend splices it verbatim.
|
||||
let (mm_content, chat_messages_json) = if enable_multimodal_chat {
|
||||
let msgs = build_chat_messages_json(&prompt, mm_content.as_deref());
|
||||
(None, Some(Arc::from(msgs.as_str())))
|
||||
} else {
|
||||
(mm_content, None)
|
||||
};
|
||||
|
||||
SampleRequest {
|
||||
prompt: Arc::from(prompt.as_str()),
|
||||
prompt_len: target_lens[i],
|
||||
expected_output_len: params[i].output_len,
|
||||
request_id: Some(format!("{rid_prefix}{i}")),
|
||||
multi_modal_content: mm_content,
|
||||
chat_messages_json,
|
||||
..Default::default()
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// Pre-serialize the OpenAI chat `messages` array for --enable-multimodal-chat.
|
||||
///
|
||||
/// Produces `[{"role":"user","content":[{"type":"text","text":"..."},<frag>,...]}]`
|
||||
/// by concatenating the JSON-escaped prompt with the pre-serialized mm fragments,
|
||||
/// so the ~200KB+ base64 image data is never parsed or re-serialized.
|
||||
pub(crate) fn build_chat_messages_json(prompt: &str, mm_items: Option<&[Arc<str>]>) -> String {
|
||||
let mm_total: usize =
|
||||
mm_items.map(|items| items.iter().map(|f| f.len() + 1).sum()).unwrap_or(0);
|
||||
let mut msgs = String::with_capacity(64 + prompt.len() * 2 + mm_total);
|
||||
msgs.push_str(r#"[{"role":"user","content":[{"type":"text","text":"#);
|
||||
// serde_json::to_string on &str produces a JSON-escaped quoted string
|
||||
msgs.push_str(&serde_json::to_string(prompt).unwrap());
|
||||
msgs.push('}');
|
||||
for fragment in mm_items.unwrap_or(&[]) {
|
||||
msgs.push(',');
|
||||
msgs.push_str(fragment);
|
||||
}
|
||||
msgs.push_str("]}]");
|
||||
msgs
|
||||
}
|
||||
|
||||
fn generate_prefix(
|
||||
tokenizer: &TokenizerKind,
|
||||
allowed_tokens: &[u32],
|
||||
prefix_len: usize,
|
||||
seed: u64,
|
||||
) -> Result<Vec<u32>> {
|
||||
let mut rng = StdRng::seed_from_u64(seed.wrapping_add(0xDEAD));
|
||||
let tokens: Vec<u32> = (0..prefix_len)
|
||||
.map(|_| allowed_tokens[rng.random_range(0..allowed_tokens.len())])
|
||||
.collect();
|
||||
|
||||
let (_, adjusted) = super::random::gen_prompt_decode_to_target_len(
|
||||
tokenizer,
|
||||
&tokens,
|
||||
prefix_len,
|
||||
false,
|
||||
allowed_tokens,
|
||||
)?;
|
||||
Ok(adjusted)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_parse_bucket_config_basic() {
|
||||
let input = "{(256,256,1): 0.5, (720,1280,1): 0.5}";
|
||||
let buckets = parse_bucket_config(input).unwrap();
|
||||
assert_eq!(buckets.len(), 2);
|
||||
assert_eq!(buckets[0].0.height, 256);
|
||||
assert_eq!(buckets[0].0.width, 256);
|
||||
assert_eq!(buckets[0].0.num_frames, 1);
|
||||
assert!((buckets[0].1 - 0.5).abs() < 1e-10);
|
||||
assert_eq!(buckets[1].0.height, 720);
|
||||
assert_eq!(buckets[1].0.width, 1280);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_bucket_config_single() {
|
||||
let input = "{(1024, 800, 1): 1.0}";
|
||||
let buckets = parse_bucket_config(input).unwrap();
|
||||
assert_eq!(buckets.len(), 1);
|
||||
assert_eq!(buckets[0].0.height, 1024);
|
||||
assert_eq!(buckets[0].0.width, 800);
|
||||
assert_eq!(buckets[0].0.num_frames, 1);
|
||||
assert!((buckets[0].1 - 1.0).abs() < 1e-10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_bucket_config_with_video() {
|
||||
let input = "{(256,256,1): 0.4, (720,1280,1): 0.4, (720,1280,16): 0.2}";
|
||||
let buckets = parse_bucket_config(input).unwrap();
|
||||
assert_eq!(buckets.len(), 3);
|
||||
assert_eq!(buckets[2].0.num_frames, 16);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_chat_messages_json_valid_and_ordered() {
|
||||
let frag: Arc<str> =
|
||||
Arc::from(r#"{"type":"image_url","image_url":{"url":"data:image/jpeg;base64,AAAA"}}"#);
|
||||
let msgs = build_chat_messages_json("hi \"there\"\nline2", Some(&[frag]));
|
||||
let v: serde_json::Value = serde_json::from_str(&msgs).expect("must be valid JSON");
|
||||
assert_eq!(v.as_array().unwrap().len(), 1);
|
||||
assert_eq!(v[0]["role"], "user");
|
||||
let content = v[0]["content"].as_array().unwrap();
|
||||
assert_eq!(content.len(), 2);
|
||||
assert_eq!(content[0]["type"], "text");
|
||||
assert_eq!(content[0]["text"], "hi \"there\"\nline2");
|
||||
assert_eq!(content[1]["type"], "image_url");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_chat_messages_json_text_only() {
|
||||
let msgs = build_chat_messages_json("plain", None);
|
||||
let v: serde_json::Value = serde_json::from_str(&msgs).unwrap();
|
||||
let content = v[0]["content"].as_array().unwrap();
|
||||
assert_eq!(content.len(), 1);
|
||||
assert_eq!(content[0]["text"], "plain");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_limit_mm_per_prompt() {
|
||||
let input = r#"{"image": 3, "video": 0}"#;
|
||||
let limit = parse_limit_mm_per_prompt(input).unwrap();
|
||||
assert_eq!(limit.image, 3);
|
||||
assert_eq!(limit.video, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_limit_mm_per_prompt_defaults() {
|
||||
let input = r#"{}"#;
|
||||
let limit = parse_limit_mm_per_prompt(input).unwrap();
|
||||
assert_eq!(limit.image, 255);
|
||||
assert_eq!(limit.video, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sample_mm_items_basic() {
|
||||
let mut rng = StdRng::seed_from_u64(42);
|
||||
let buckets = vec![
|
||||
(
|
||||
MmBucketKey {
|
||||
height: 256,
|
||||
width: 256,
|
||||
num_frames: 1,
|
||||
},
|
||||
0.5,
|
||||
),
|
||||
(
|
||||
MmBucketKey {
|
||||
height: 720,
|
||||
width: 1280,
|
||||
num_frames: 1,
|
||||
},
|
||||
0.5,
|
||||
),
|
||||
];
|
||||
let limit = MmLimitPerPrompt { image: 5, video: 0 };
|
||||
let items = sample_mm_items(&mut rng, 2, 3, &buckets, &limit);
|
||||
assert!(items.len() >= 2 && items.len() <= 3);
|
||||
for item in &items {
|
||||
assert_eq!(item.num_frames, 1);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sample_mm_items_respects_limit() {
|
||||
let mut rng = StdRng::seed_from_u64(42);
|
||||
let buckets = vec![(
|
||||
MmBucketKey {
|
||||
height: 256,
|
||||
width: 256,
|
||||
num_frames: 1,
|
||||
},
|
||||
1.0,
|
||||
)];
|
||||
let limit = MmLimitPerPrompt { image: 2, video: 0 };
|
||||
let items = sample_mm_items(&mut rng, 5, 5, &buckets, &limit);
|
||||
// Should be capped at 2 due to image limit
|
||||
assert_eq!(items.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_generate_random_image() {
|
||||
let mut rng = StdRng::seed_from_u64(42);
|
||||
let result = generate_random_image(64, 64, &mut rng).unwrap();
|
||||
// Result is a pre-serialized JSON fragment
|
||||
assert!(
|
||||
result
|
||||
.starts_with(r#"{"type":"image_url","image_url":{"url":"data:image/jpeg;base64,"#)
|
||||
);
|
||||
assert!(result.ends_with(r#""}}"#));
|
||||
// Verify it's valid JSON
|
||||
let parsed: serde_json::Value = serde_json::from_str(&result).unwrap();
|
||||
assert_eq!(parsed["type"], "image_url");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
//! Random dataset specialized for scoring/rerank benchmarks: each request is
|
||||
//! one query plus a batch of documents. Mirrors Python's
|
||||
//! `RandomDatasetForReranking`.
|
||||
//!
|
||||
//! With `is_reranker` (default): the query and each document share the
|
||||
//! request's token budget (`query + sep + doc ~= input_len`), and every
|
||||
//! batched request counts the query once per document pair.
|
||||
//! With `--no-reranker` (embedding-based scoring): the query is just another
|
||||
//! embedding input occupying the first batch slot.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rand::rngs::StdRng;
|
||||
use rand::{Rng, SeedableRng};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use super::SampleRequest;
|
||||
use crate::config::RangeRatio;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
pub fn generate_random_rerank_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
num_requests: usize,
|
||||
input_len: usize,
|
||||
range_ratio: RangeRatio,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
batch_size: usize,
|
||||
is_reranker: bool,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
let allowed_tokens = tokenizer.get_allowed_tokens();
|
||||
if allowed_tokens.is_empty() {
|
||||
return Err(BenchError::Tokenizer("No allowed tokens found".into()));
|
||||
}
|
||||
let allowed_ref = &allowed_tokens;
|
||||
|
||||
let num_special = tokenizer.num_special_tokens_to_add();
|
||||
let real_input_len = input_len.saturating_sub(num_special);
|
||||
|
||||
let n_sep_tokens = usize::from(is_reranker);
|
||||
let query_len_param = if is_reranker {
|
||||
(real_input_len / 2).saturating_sub(n_sep_tokens)
|
||||
} else {
|
||||
real_input_len
|
||||
};
|
||||
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
let sample = |rng: &mut StdRng, (low, high): (usize, usize)| -> usize {
|
||||
if low == high {
|
||||
low
|
||||
} else {
|
||||
rng.random_range(low..=high)
|
||||
}
|
||||
};
|
||||
|
||||
// One query length for the whole run, like Python.
|
||||
let query_len = sample(&mut rng, range_ratio.input_bounds(query_len_param));
|
||||
|
||||
// --no-reranker folds the query into the first batch slot.
|
||||
let (num_docs, docs_per_batch, doc_len_param) = if is_reranker {
|
||||
let doc_len = real_input_len.saturating_sub(query_len).saturating_sub(n_sep_tokens);
|
||||
(num_requests, batch_size, doc_len)
|
||||
} else {
|
||||
(num_requests - 1, batch_size - 1, real_input_len)
|
||||
};
|
||||
if doc_len_param == 0 {
|
||||
return Err(BenchError::Config(format!(
|
||||
"random-rerank: --random-input-len {input_len} leaves no budget for documents \
|
||||
(query_len={query_len})"
|
||||
)));
|
||||
}
|
||||
|
||||
// Pre-sample per-document lengths and offsets deterministically.
|
||||
let doc_bounds = range_ratio.input_bounds(doc_len_param);
|
||||
let doc_params: Vec<(usize, usize)> = (0..num_docs)
|
||||
.map(|_| {
|
||||
(
|
||||
sample(&mut rng, doc_bounds),
|
||||
rng.random_range(0..allowed_ref.len()),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
let query_offset = rng.random_range(0..allowed_ref.len());
|
||||
|
||||
// Exact-length text: token sequence -> decode -> re-encode -> truncate -> decode.
|
||||
let gen_text = |target: usize, offset: usize, index: usize| -> Result<(Arc<str>, usize)> {
|
||||
let at_len = allowed_ref.len();
|
||||
let tokens: Vec<u32> =
|
||||
(0..target).map(|j| allowed_ref[(offset + index + j) % at_len]).collect();
|
||||
let text = tokenizer.decode(&tokens, true)?;
|
||||
let mut re_encoded = tokenizer.encode(&text, false)?;
|
||||
re_encoded.truncate(target);
|
||||
let final_text = tokenizer.decode(&re_encoded, true)?;
|
||||
Ok((Arc::from(final_text), re_encoded.len()))
|
||||
};
|
||||
|
||||
let (query_prompt, query_input_len) = gen_text(query_len, query_offset, 0)?;
|
||||
|
||||
let docs: Vec<(Arc<str>, usize)> = doc_params
|
||||
.par_iter()
|
||||
.enumerate()
|
||||
.map(|(i, (len, offset))| gen_text(*len, *offset, i + 1))
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
|
||||
// Batch documents; every request is [query, doc1, doc2, ...].
|
||||
let rid_prefix = request_id_prefix.to_string();
|
||||
let requests = docs
|
||||
.chunks(docs_per_batch)
|
||||
.enumerate()
|
||||
.map(|(batch_idx, batch)| {
|
||||
let query_contrib = if is_reranker {
|
||||
(query_input_len + n_sep_tokens) * batch.len()
|
||||
} else {
|
||||
query_input_len
|
||||
};
|
||||
let mut prompt_list: Vec<Arc<str>> = Vec::with_capacity(batch.len() + 1);
|
||||
prompt_list.push(query_prompt.clone());
|
||||
prompt_list.extend(batch.iter().map(|(text, _)| text.clone()));
|
||||
SampleRequest {
|
||||
prompt_list: Some(Arc::from(prompt_list)),
|
||||
prompt_len: query_contrib + batch.iter().map(|(_, len)| len).sum::<usize>(),
|
||||
expected_output_len: 0,
|
||||
request_id: Some(format!("{rid_prefix}{batch_idx}")),
|
||||
..Default::default()
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(requests)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// gpt2 via built-in tiktoken encoding — loads without network access.
|
||||
fn test_tokenizer() -> TokenizerKind {
|
||||
TokenizerKind::Tiktoken(
|
||||
crate::tiktoken::load_builtin_tiktoken("gpt2")
|
||||
.expect("gpt2 built-in tiktoken should always load without network"),
|
||||
)
|
||||
}
|
||||
|
||||
fn fixed_ratio() -> RangeRatio {
|
||||
RangeRatio::parse("0.0").unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_random_rerank_reranker_mode() {
|
||||
let tok = test_tokenizer();
|
||||
// 6 docs in batches of 3 -> 2 requests of [query, d1, d2, d3]
|
||||
let reqs = generate_random_rerank_dataset(&tok, 6, 128, fixed_ratio(), 0, "t-", 3, true)
|
||||
.expect("generation should succeed");
|
||||
assert_eq!(reqs.len(), 2);
|
||||
for r in &reqs {
|
||||
let list = r.prompt_list.as_ref().expect("prompt_list must be set");
|
||||
assert_eq!(list.len(), 4);
|
||||
assert_eq!(r.expected_output_len, 0);
|
||||
assert!(r.prompt_len > 0);
|
||||
}
|
||||
// Same query shared across requests
|
||||
assert_eq!(
|
||||
reqs[0].prompt_list.as_ref().unwrap()[0],
|
||||
reqs[1].prompt_list.as_ref().unwrap()[0]
|
||||
);
|
||||
// Reranker budget: query+sep+doc pairs stay near input_len per pair
|
||||
// (query ~63, doc ~64 for input_len=128, gpt2 has no special tokens)
|
||||
let list = reqs[0].prompt_list.as_ref().unwrap();
|
||||
let query_tokens = tok.encode(&list[0], false).unwrap().len();
|
||||
assert!(query_tokens <= 64, "query too long: {query_tokens}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_random_rerank_no_reranker_mode() {
|
||||
let tok = test_tokenizer();
|
||||
// no-reranker: query occupies first slot; 5 non-query docs in batches of 2
|
||||
let reqs = generate_random_rerank_dataset(&tok, 6, 64, fixed_ratio(), 0, "t-", 3, false)
|
||||
.expect("generation should succeed");
|
||||
// 6-1=5 docs, batches of 3-1=2 -> 3 requests
|
||||
assert_eq!(reqs.len(), 3);
|
||||
assert_eq!(reqs[0].prompt_list.as_ref().unwrap().len(), 3);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,216 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rand::rngs::StdRng;
|
||||
use rand::seq::SliceRandom;
|
||||
use rand::{Rng, SeedableRng};
|
||||
|
||||
use super::SampleRequest;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
/// Default validation bounds matching Python's is_valid_sequence() defaults.
|
||||
const MIN_LEN: usize = 4;
|
||||
const MAX_PROMPT_LEN: usize = 1024;
|
||||
const MAX_TOTAL_LEN: usize = 2048;
|
||||
|
||||
/// Default HuggingFace dataset repo and filename for ShareGPT.
|
||||
const DEFAULT_SHAREGPT_REPO: &str = "anon8231489123/ShareGPT_Vicuna_unfiltered";
|
||||
const DEFAULT_SHAREGPT_FILE: &str = "ShareGPT_V3_unfiltered_cleaned_split.json";
|
||||
|
||||
/// Download the default ShareGPT dataset from HuggingFace Hub.
|
||||
/// Uses hf-hub's built-in cache — subsequent calls return the cached path instantly.
|
||||
pub async fn download_sharegpt_dataset() -> Result<String> {
|
||||
tracing::info!(
|
||||
repository = DEFAULT_SHAREGPT_REPO,
|
||||
file = DEFAULT_SHAREGPT_FILE,
|
||||
"downloading ShareGPT dataset"
|
||||
);
|
||||
let repo = crate::hub::HubRepo::dataset(DEFAULT_SHAREGPT_REPO.to_string())
|
||||
.map_err(BenchError::Config)?;
|
||||
let path = repo.get(DEFAULT_SHAREGPT_FILE).await.map_err(|e| {
|
||||
BenchError::Config(format!(
|
||||
"Failed to download ShareGPT dataset from '{DEFAULT_SHAREGPT_REPO}': {e}"
|
||||
))
|
||||
})?;
|
||||
let path_str = path.to_string_lossy().to_string();
|
||||
tracing::info!(dataset = "sharegpt", path = %path_str, "dataset is ready");
|
||||
Ok(path_str)
|
||||
}
|
||||
|
||||
/// Load and sample from a ShareGPT-format JSON dataset.
|
||||
///
|
||||
/// Mirrors Python's ShareGPTDataset from datasets.py:1230-1313.
|
||||
pub fn load_sharegpt_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
dataset_path: &str,
|
||||
num_requests: usize,
|
||||
output_len_override: Option<usize>,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
no_oversample: bool,
|
||||
disable_shuffle: bool,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
// Load JSON file
|
||||
let content = std::fs::read_to_string(dataset_path).map_err(|e| {
|
||||
BenchError::Config(format!(
|
||||
"Failed to read ShareGPT file '{dataset_path}': {e}"
|
||||
))
|
||||
})?;
|
||||
|
||||
let data: serde_json::Value = serde_json::from_str(&content)
|
||||
.map_err(|e| BenchError::Config(format!("Invalid JSON in ShareGPT file: {e}")))?;
|
||||
|
||||
let entries = data
|
||||
.as_array()
|
||||
.ok_or_else(|| BenchError::Config("ShareGPT file must contain a JSON array".into()))?;
|
||||
|
||||
// Filter entries with at least 2 conversation turns
|
||||
let mut filtered: Vec<&serde_json::Value> = entries
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
entry
|
||||
.get("conversations")
|
||||
.and_then(|c| c.as_array())
|
||||
.map(|a| a.len() >= 2)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.collect();
|
||||
|
||||
if filtered.is_empty() {
|
||||
return Err(BenchError::Config(
|
||||
"No valid entries in ShareGPT file (need at least 2 conversation turns)".into(),
|
||||
));
|
||||
}
|
||||
|
||||
// Shuffle (unless disabled)
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
if !disable_shuffle {
|
||||
filtered.shuffle(&mut rng);
|
||||
}
|
||||
|
||||
// Sample requests
|
||||
let mut samples = Vec::new();
|
||||
let mut ind = 0;
|
||||
|
||||
for entry in &filtered {
|
||||
if samples.len() >= num_requests {
|
||||
break;
|
||||
}
|
||||
|
||||
let conversations = entry["conversations"].as_array().unwrap();
|
||||
let prompt = conversations[0]["value"].as_str().unwrap_or("");
|
||||
let completion = conversations[1]["value"].as_str().unwrap_or("");
|
||||
|
||||
if prompt.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Tokenize prompt and completion
|
||||
let prompt_ids = tokenizer.encode(prompt, false)?;
|
||||
let prompt_len = prompt_ids.len();
|
||||
|
||||
let new_output_len = if let Some(override_len) = output_len_override {
|
||||
override_len
|
||||
} else {
|
||||
let completion_ids = tokenizer.encode(completion, false)?;
|
||||
completion_ids.len()
|
||||
};
|
||||
|
||||
// Validate sequence lengths (matching Python's is_valid_sequence)
|
||||
let skip_min_output = output_len_override.is_some();
|
||||
if !is_valid_sequence(prompt_len, new_output_len, skip_min_output) {
|
||||
continue;
|
||||
}
|
||||
|
||||
samples.push(SampleRequest {
|
||||
prompt: Arc::from(prompt),
|
||||
prompt_len,
|
||||
expected_output_len: new_output_len,
|
||||
request_id: Some(format!("{request_id_prefix}{ind}")),
|
||||
..Default::default()
|
||||
});
|
||||
ind += 1;
|
||||
}
|
||||
|
||||
// Oversample if dataset is smaller than requested
|
||||
if samples.len() < num_requests {
|
||||
if no_oversample {
|
||||
tracing::info!(
|
||||
dataset = "sharegpt",
|
||||
samples = samples.len(),
|
||||
requested = num_requests,
|
||||
"skipping dataset oversampling"
|
||||
);
|
||||
} else if !samples.is_empty() {
|
||||
let needed = num_requests - samples.len();
|
||||
let original_len = samples.len();
|
||||
for i in 0..needed {
|
||||
let mut req = samples[rng.random_range(0..original_len)].clone();
|
||||
req.request_id = Some(format!("{request_id_prefix}{}", original_len + i));
|
||||
samples.push(req);
|
||||
}
|
||||
tracing::info!(
|
||||
dataset = "sharegpt",
|
||||
original_samples = original_len,
|
||||
samples = samples.len(),
|
||||
"oversampled dataset"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if samples.is_empty() {
|
||||
return Err(BenchError::Config(
|
||||
"No valid samples after filtering ShareGPT dataset. \
|
||||
Try relaxing constraints or using a larger dataset."
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(samples)
|
||||
}
|
||||
|
||||
/// Validate a sequence based on prompt and output lengths.
|
||||
/// Mirrors Python's is_valid_sequence() from datasets.py:260-284.
|
||||
fn is_valid_sequence(
|
||||
prompt_len: usize,
|
||||
output_len: usize,
|
||||
skip_min_output_len_check: bool,
|
||||
) -> bool {
|
||||
if prompt_len < MIN_LEN {
|
||||
return false;
|
||||
}
|
||||
if !skip_min_output_len_check && output_len < MIN_LEN {
|
||||
return false;
|
||||
}
|
||||
if prompt_len > MAX_PROMPT_LEN {
|
||||
return false;
|
||||
}
|
||||
if prompt_len + output_len > MAX_TOTAL_LEN {
|
||||
return false;
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_is_valid_sequence() {
|
||||
// Valid
|
||||
assert!(is_valid_sequence(100, 50, false));
|
||||
// Prompt too short
|
||||
assert!(!is_valid_sequence(3, 50, false));
|
||||
// Output too short
|
||||
assert!(!is_valid_sequence(100, 3, false));
|
||||
// Output too short but skip check
|
||||
assert!(is_valid_sequence(100, 1, true));
|
||||
// Prompt too long
|
||||
assert!(!is_valid_sequence(1025, 50, false));
|
||||
// Combined too long
|
||||
assert!(!is_valid_sequence(1024, 1025, false));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rand::SeedableRng;
|
||||
use rand::rngs::StdRng;
|
||||
use rand::seq::IndexedRandom;
|
||||
|
||||
use super::SampleRequest;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
/// Default values mirror Python's SonnetDataset defaults (datasets.py).
|
||||
pub const DEFAULT_PREFIX_LEN: usize = 200;
|
||||
pub const DEFAULT_INPUT_LEN: usize = 550;
|
||||
pub const DEFAULT_OUTPUT_LEN: usize = 150;
|
||||
|
||||
const BASE_PROMPT: &str = "Pick as many lines as you can from these poem lines:\n";
|
||||
|
||||
/// Shakespeare's sonnets, public domain. Bundled so `--dataset-name sonnet` works
|
||||
/// out of the box without `--dataset-path`. Source:
|
||||
/// https://raw.githubusercontent.com/vllm-project/vllm/main/benchmarks/sonnet.txt
|
||||
const BUILTIN_SONNET: &str = include_str!("sonnet.txt");
|
||||
|
||||
/// Load the sonnet dataset and generate `num_requests` prompts targeting `input_len`
|
||||
/// total prompt tokens.
|
||||
///
|
||||
/// Mirrors Python's `SonnetDataset.sample()` from `vllm/benchmarks/datasets/datasets.py`.
|
||||
/// The Rust port skips `apply_chat_template` (no Jinja runtime here): `base_offset`
|
||||
/// and `prompt_len` are computed from the raw text. The resulting prompts are slightly
|
||||
/// shorter than Python's chat-template-formatted version (off by the chat scaffolding
|
||||
/// tokens, typically <20).
|
||||
pub fn load_sonnet_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
dataset_path: Option<&str>,
|
||||
num_requests: usize,
|
||||
input_len: usize,
|
||||
output_len: usize,
|
||||
prefix_len: usize,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
let content = match dataset_path {
|
||||
Some(path) => std::fs::read_to_string(path)
|
||||
.map_err(|e| BenchError::Config(format!("Failed to read sonnet file '{path}': {e}")))?,
|
||||
None => BUILTIN_SONNET.to_string(),
|
||||
};
|
||||
|
||||
// Match Python's f.readlines(): keep trailing newlines so that joining lines
|
||||
// reconstructs the original text without inserting extra separators.
|
||||
let lines: Vec<String> = content.split_inclusive('\n').map(|s| s.to_string()).collect();
|
||||
|
||||
if lines.is_empty() {
|
||||
let src = dataset_path.unwrap_or("<built-in>");
|
||||
return Err(BenchError::Config(format!("Sonnet file '{src}' is empty")));
|
||||
}
|
||||
|
||||
// Average tokens per line (used to estimate how many lines to draw).
|
||||
let mut total_tokens: usize = 0;
|
||||
for line in &lines {
|
||||
let ids = tokenizer.encode(line, false)?;
|
||||
total_tokens += ids.len();
|
||||
}
|
||||
let avg_len = total_tokens as f64 / lines.len() as f64;
|
||||
if avg_len <= 0.0 {
|
||||
return Err(BenchError::Config(
|
||||
"Sonnet lines tokenized to zero tokens on average".into(),
|
||||
));
|
||||
}
|
||||
|
||||
let base_ids = tokenizer.encode(BASE_PROMPT, false)?;
|
||||
let base_offset = base_ids.len();
|
||||
if input_len <= base_offset {
|
||||
return Err(BenchError::Config(format!(
|
||||
"--sonnet-input-len ({input_len}) must be larger than the base prompt length ({base_offset})"
|
||||
)));
|
||||
}
|
||||
|
||||
let num_input_lines = ((input_len - base_offset) as f64 / avg_len).round() as i64;
|
||||
let num_prefix_lines =
|
||||
(((prefix_len as i64 - base_offset as i64) as f64) / avg_len).round().max(0.0) as i64;
|
||||
let num_input_lines = num_input_lines.max(0) as usize;
|
||||
let num_prefix_lines = (num_prefix_lines as usize).min(lines.len());
|
||||
let num_input_lines = num_input_lines.max(num_prefix_lines);
|
||||
|
||||
let prefix_lines: &[String] = &lines[..num_prefix_lines];
|
||||
let extras_per_request = num_input_lines - num_prefix_lines;
|
||||
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
let mut samples = Vec::with_capacity(num_requests);
|
||||
let mut ind = 0usize;
|
||||
let mut attempts = 0usize;
|
||||
let max_attempts = num_requests.saturating_mul(20).max(1000);
|
||||
|
||||
while samples.len() < num_requests {
|
||||
if attempts >= max_attempts {
|
||||
return Err(BenchError::Config(format!(
|
||||
"Could not assemble {num_requests} sonnet prompts under input_len={input_len} \
|
||||
after {attempts} attempts. Try increasing --sonnet-input-len."
|
||||
)));
|
||||
}
|
||||
attempts += 1;
|
||||
|
||||
let mut prompt = String::with_capacity(BASE_PROMPT.len() + 256 * num_input_lines);
|
||||
prompt.push_str(BASE_PROMPT);
|
||||
for line in prefix_lines {
|
||||
prompt.push_str(line);
|
||||
}
|
||||
for _ in 0..extras_per_request {
|
||||
// random.choices with replacement — duplicates are allowed.
|
||||
let line = lines.choose(&mut rng).unwrap();
|
||||
prompt.push_str(line);
|
||||
}
|
||||
|
||||
let prompt_ids = tokenizer.encode(&prompt, false)?;
|
||||
let prompt_len = prompt_ids.len();
|
||||
if prompt_len <= input_len {
|
||||
samples.push(SampleRequest {
|
||||
prompt: Arc::from(prompt),
|
||||
prompt_len,
|
||||
expected_output_len: output_len,
|
||||
request_id: Some(format!("{request_id_prefix}{ind}")),
|
||||
..Default::default()
|
||||
});
|
||||
ind += 1;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(samples)
|
||||
}
|
||||
@@ -0,0 +1,518 @@
|
||||
FROM fairest creatures we desire increase,
|
||||
That thereby beauty's rose might never die,
|
||||
But as the riper should by time decease,
|
||||
His tender heir might bear his memory:
|
||||
But thou, contracted to thine own bright eyes,
|
||||
Feed'st thy light'st flame with self-substantial fuel,
|
||||
Making a famine where abundance lies,
|
||||
Thyself thy foe, to thy sweet self too cruel.
|
||||
Thou that art now the world's fresh ornament
|
||||
And only herald to the gaudy spring,
|
||||
Within thine own bud buriest thy content
|
||||
And, tender churl, makest waste in niggarding.
|
||||
Pity the world, or else this glutton be,
|
||||
To eat the world's due, by the grave and thee.
|
||||
When forty winters shall beseige thy brow,
|
||||
And dig deep trenches in thy beauty's field,
|
||||
Thy youth's proud livery, so gazed on now,
|
||||
Will be a tatter'd weed, of small worth held:
|
||||
Then being ask'd where all thy beauty lies,
|
||||
Where all the treasure of thy lusty days,
|
||||
To say, within thine own deep-sunken eyes,
|
||||
Were an all-eating shame and thriftless praise.
|
||||
How much more praise deserved thy beauty's use,
|
||||
If thou couldst answer 'This fair child of mine
|
||||
Shall sum my count and make my old excuse,'
|
||||
Proving his beauty by succession thine!
|
||||
This were to be new made when thou art old,
|
||||
And see thy blood warm when thou feel'st it cold.
|
||||
Look in thy glass, and tell the face thou viewest
|
||||
Now is the time that face should form another;
|
||||
Whose fresh repair if now thou not renewest,
|
||||
Thou dost beguile the world, unbless some mother.
|
||||
For where is she so fair whose unear'd womb
|
||||
Disdains the tillage of thy husbandry?
|
||||
Or who is he so fond will be the tomb
|
||||
Of his self-love, to stop posterity?
|
||||
Thou art thy mother's glass, and she in thee
|
||||
Calls back the lovely April of her prime:
|
||||
So thou through windows of thine age shall see
|
||||
Despite of wrinkles this thy golden time.
|
||||
But if thou live, remember'd not to be,
|
||||
Die single, and thine image dies with thee.
|
||||
Unthrifty loveliness, why dost thou spend
|
||||
Upon thyself thy beauty's legacy?
|
||||
Nature's bequest gives nothing but doth lend,
|
||||
And being frank she lends to those are free.
|
||||
Then, beauteous niggard, why dost thou abuse
|
||||
The bounteous largess given thee to give?
|
||||
Profitless usurer, why dost thou use
|
||||
So great a sum of sums, yet canst not live?
|
||||
For having traffic with thyself alone,
|
||||
Thou of thyself thy sweet self dost deceive.
|
||||
Then how, when nature calls thee to be gone,
|
||||
What acceptable audit canst thou leave?
|
||||
Thy unused beauty must be tomb'd with thee,
|
||||
Which, used, lives th' executor to be.
|
||||
Those hours, that with gentle work did frame
|
||||
The lovely gaze where every eye doth dwell,
|
||||
Will play the tyrants to the very same
|
||||
And that unfair which fairly doth excel:
|
||||
For never-resting time leads summer on
|
||||
To hideous winter and confounds him there;
|
||||
Sap cheque'd with frost and lusty leaves quite gone,
|
||||
Beauty o'ersnow'd and bareness every where:
|
||||
Then, were not summer's distillation left,
|
||||
A liquid prisoner pent in walls of glass,
|
||||
Beauty's effect with beauty were bereft,
|
||||
Nor it nor no remembrance what it was:
|
||||
But flowers distill'd though they with winter meet,
|
||||
Leese but their show; their substance still lives sweet.
|
||||
Then let not winter's ragged hand deface
|
||||
In thee thy summer, ere thou be distill'd:
|
||||
Make sweet some vial; treasure thou some place
|
||||
With beauty's treasure, ere it be self-kill'd.
|
||||
That use is not forbidden usury,
|
||||
Which happies those that pay the willing loan;
|
||||
That's for thyself to breed another thee,
|
||||
Or ten times happier, be it ten for one;
|
||||
Ten times thyself were happier than thou art,
|
||||
If ten of thine ten times refigured thee:
|
||||
Then what could death do, if thou shouldst depart,
|
||||
Leaving thee living in posterity?
|
||||
Be not self-will'd, for thou art much too fair
|
||||
To be death's conquest and make worms thine heir.
|
||||
Lo! in the orient when the gracious light
|
||||
Lifts up his burning head, each under eye
|
||||
Doth homage to his new-appearing sight,
|
||||
Serving with looks his sacred majesty;
|
||||
And having climb'd the steep-up heavenly hill,
|
||||
Resembling strong youth in his middle age,
|
||||
yet mortal looks adore his beauty still,
|
||||
Attending on his golden pilgrimage;
|
||||
But when from highmost pitch, with weary car,
|
||||
Like feeble age, he reeleth from the day,
|
||||
The eyes, 'fore duteous, now converted are
|
||||
From his low tract and look another way:
|
||||
So thou, thyself out-going in thy noon,
|
||||
Unlook'd on diest, unless thou get a son.
|
||||
Music to hear, why hear'st thou music sadly?
|
||||
Sweets with sweets war not, joy delights in joy.
|
||||
Why lovest thou that which thou receivest not gladly,
|
||||
Or else receivest with pleasure thine annoy?
|
||||
If the true concord of well-tuned sounds,
|
||||
By unions married, do offend thine ear,
|
||||
They do but sweetly chide thee, who confounds
|
||||
In singleness the parts that thou shouldst bear.
|
||||
Mark how one string, sweet husband to another,
|
||||
Strikes each in each by mutual ordering,
|
||||
Resembling sire and child and happy mother
|
||||
Who all in one, one pleasing note do sing:
|
||||
Whose speechless song, being many, seeming one,
|
||||
Sings this to thee: 'thou single wilt prove none.'
|
||||
Is it for fear to wet a widow's eye
|
||||
That thou consumest thyself in single life?
|
||||
Ah! if thou issueless shalt hap to die.
|
||||
The world will wail thee, like a makeless wife;
|
||||
The world will be thy widow and still weep
|
||||
That thou no form of thee hast left behind,
|
||||
When every private widow well may keep
|
||||
By children's eyes her husband's shape in mind.
|
||||
Look, what an unthrift in the world doth spend
|
||||
Shifts but his place, for still the world enjoys it;
|
||||
But beauty's waste hath in the world an end,
|
||||
And kept unused, the user so destroys it.
|
||||
No love toward others in that bosom sits
|
||||
That on himself such murderous shame commits.
|
||||
For shame! deny that thou bear'st love to any,
|
||||
Who for thyself art so unprovident.
|
||||
Grant, if thou wilt, thou art beloved of many,
|
||||
But that thou none lovest is most evident;
|
||||
For thou art so possess'd with murderous hate
|
||||
That 'gainst thyself thou stick'st not to conspire.
|
||||
Seeking that beauteous roof to ruinate
|
||||
Which to repair should be thy chief desire.
|
||||
O, change thy thought, that I may change my mind!
|
||||
Shall hate be fairer lodged than gentle love?
|
||||
Be, as thy presence is, gracious and kind,
|
||||
Or to thyself at least kind-hearted prove:
|
||||
Make thee another self, for love of me,
|
||||
That beauty still may live in thine or thee.
|
||||
As fast as thou shalt wane, so fast thou growest
|
||||
In one of thine, from that which thou departest;
|
||||
And that fresh blood which youngly thou bestowest
|
||||
Thou mayst call thine when thou from youth convertest.
|
||||
Herein lives wisdom, beauty and increase:
|
||||
Without this, folly, age and cold decay:
|
||||
If all were minded so, the times should cease
|
||||
And threescore year would make the world away.
|
||||
Let those whom Nature hath not made for store,
|
||||
Harsh featureless and rude, barrenly perish:
|
||||
Look, whom she best endow'd she gave the more;
|
||||
Which bounteous gift thou shouldst in bounty cherish:
|
||||
She carved thee for her seal, and meant thereby
|
||||
Thou shouldst print more, not let that copy die.
|
||||
When I do count the clock that tells the time,
|
||||
And see the brave day sunk in hideous night;
|
||||
When I behold the violet past prime,
|
||||
And sable curls all silver'd o'er with white;
|
||||
When lofty trees I see barren of leaves
|
||||
Which erst from heat did canopy the herd,
|
||||
And summer's green all girded up in sheaves
|
||||
Borne on the bier with white and bristly beard,
|
||||
Then of thy beauty do I question make,
|
||||
That thou among the wastes of time must go,
|
||||
Since sweets and beauties do themselves forsake
|
||||
And die as fast as they see others grow;
|
||||
And nothing 'gainst Time's scythe can make defence
|
||||
Save breed, to brave him when he takes thee hence.
|
||||
O, that you were yourself! but, love, you are
|
||||
No longer yours than you yourself here live:
|
||||
Against this coming end you should prepare,
|
||||
And your sweet semblance to some other give.
|
||||
So should that beauty which you hold in lease
|
||||
Find no determination: then you were
|
||||
Yourself again after yourself's decease,
|
||||
When your sweet issue your sweet form should bear.
|
||||
Who lets so fair a house fall to decay,
|
||||
Which husbandry in honour might uphold
|
||||
Against the stormy gusts of winter's day
|
||||
And barren rage of death's eternal cold?
|
||||
O, none but unthrifts! Dear my love, you know
|
||||
You had a father: let your son say so.
|
||||
Not from the stars do I my judgment pluck;
|
||||
And yet methinks I have astronomy,
|
||||
But not to tell of good or evil luck,
|
||||
Of plagues, of dearths, or seasons' quality;
|
||||
Nor can I fortune to brief minutes tell,
|
||||
Pointing to each his thunder, rain and wind,
|
||||
Or say with princes if it shall go well,
|
||||
By oft predict that I in heaven find:
|
||||
But from thine eyes my knowledge I derive,
|
||||
And, constant stars, in them I read such art
|
||||
As truth and beauty shall together thrive,
|
||||
If from thyself to store thou wouldst convert;
|
||||
Or else of thee this I prognosticate:
|
||||
Thy end is truth's and beauty's doom and date.
|
||||
When I consider every thing that grows
|
||||
Holds in perfection but a little moment,
|
||||
That this huge stage presenteth nought but shows
|
||||
Whereon the stars in secret influence comment;
|
||||
When I perceive that men as plants increase,
|
||||
Cheered and cheque'd even by the self-same sky,
|
||||
Vaunt in their youthful sap, at height decrease,
|
||||
And wear their brave state out of memory;
|
||||
Then the conceit of this inconstant stay
|
||||
Sets you most rich in youth before my sight,
|
||||
Where wasteful Time debateth with Decay,
|
||||
To change your day of youth to sullied night;
|
||||
And all in war with Time for love of you,
|
||||
As he takes from you, I engraft you new.
|
||||
But wherefore do not you a mightier way
|
||||
Make war upon this bloody tyrant, Time?
|
||||
And fortify yourself in your decay
|
||||
With means more blessed than my barren rhyme?
|
||||
Now stand you on the top of happy hours,
|
||||
And many maiden gardens yet unset
|
||||
With virtuous wish would bear your living flowers,
|
||||
Much liker than your painted counterfeit:
|
||||
So should the lines of life that life repair,
|
||||
Which this, Time's pencil, or my pupil pen,
|
||||
Neither in inward worth nor outward fair,
|
||||
Can make you live yourself in eyes of men.
|
||||
To give away yourself keeps yourself still,
|
||||
And you must live, drawn by your own sweet skill.
|
||||
Who will believe my verse in time to come,
|
||||
If it were fill'd with your most high deserts?
|
||||
Though yet, heaven knows, it is but as a tomb
|
||||
Which hides your life and shows not half your parts.
|
||||
If I could write the beauty of your eyes
|
||||
And in fresh numbers number all your graces,
|
||||
The age to come would say 'This poet lies:
|
||||
Such heavenly touches ne'er touch'd earthly faces.'
|
||||
So should my papers yellow'd with their age
|
||||
Be scorn'd like old men of less truth than tongue,
|
||||
And your true rights be term'd a poet's rage
|
||||
And stretched metre of an antique song:
|
||||
But were some child of yours alive that time,
|
||||
You should live twice; in it and in my rhyme.
|
||||
Shall I compare thee to a summer's day?
|
||||
Thou art more lovely and more temperate:
|
||||
Rough winds do shake the darling buds of May,
|
||||
And summer's lease hath all too short a date:
|
||||
Sometime too hot the eye of heaven shines,
|
||||
And often is his gold complexion dimm'd;
|
||||
And every fair from fair sometime declines,
|
||||
By chance or nature's changing course untrimm'd;
|
||||
But thy eternal summer shall not fade
|
||||
Nor lose possession of that fair thou owest;
|
||||
Nor shall Death brag thou wander'st in his shade,
|
||||
When in eternal lines to time thou growest:
|
||||
So long as men can breathe or eyes can see,
|
||||
So long lives this and this gives life to thee.
|
||||
Devouring Time, blunt thou the lion's paws,
|
||||
And make the earth devour her own sweet brood;
|
||||
Pluck the keen teeth from the fierce tiger's jaws,
|
||||
And burn the long-lived phoenix in her blood;
|
||||
Make glad and sorry seasons as thou fleets,
|
||||
And do whate'er thou wilt, swift-footed Time,
|
||||
To the wide world and all her fading sweets;
|
||||
But I forbid thee one most heinous crime:
|
||||
O, carve not with thy hours my love's fair brow,
|
||||
Nor draw no lines there with thine antique pen;
|
||||
Him in thy course untainted do allow
|
||||
For beauty's pattern to succeeding men.
|
||||
Yet, do thy worst, old Time: despite thy wrong,
|
||||
My love shall in my verse ever live young.
|
||||
A woman's face with Nature's own hand painted
|
||||
Hast thou, the master-mistress of my passion;
|
||||
A woman's gentle heart, but not acquainted
|
||||
With shifting change, as is false women's fashion;
|
||||
An eye more bright than theirs, less false in rolling,
|
||||
Gilding the object whereupon it gazeth;
|
||||
A man in hue, all 'hues' in his controlling,
|
||||
Much steals men's eyes and women's souls amazeth.
|
||||
And for a woman wert thou first created;
|
||||
Till Nature, as she wrought thee, fell a-doting,
|
||||
And by addition me of thee defeated,
|
||||
By adding one thing to my purpose nothing.
|
||||
But since she prick'd thee out for women's pleasure,
|
||||
Mine be thy love and thy love's use their treasure.
|
||||
So is it not with me as with that Muse
|
||||
Stirr'd by a painted beauty to his verse,
|
||||
Who heaven itself for ornament doth use
|
||||
And every fair with his fair doth rehearse
|
||||
Making a couplement of proud compare,
|
||||
With sun and moon, with earth and sea's rich gems,
|
||||
With April's first-born flowers, and all things rare
|
||||
That heaven's air in this huge rondure hems.
|
||||
O' let me, true in love, but truly write,
|
||||
And then believe me, my love is as fair
|
||||
As any mother's child, though not so bright
|
||||
As those gold candles fix'd in heaven's air:
|
||||
Let them say more than like of hearsay well;
|
||||
I will not praise that purpose not to sell.
|
||||
My glass shall not persuade me I am old,
|
||||
So long as youth and thou are of one date;
|
||||
But when in thee time's furrows I behold,
|
||||
Then look I death my days should expiate.
|
||||
For all that beauty that doth cover thee
|
||||
Is but the seemly raiment of my heart,
|
||||
Which in thy breast doth live, as thine in me:
|
||||
How can I then be elder than thou art?
|
||||
O, therefore, love, be of thyself so wary
|
||||
As I, not for myself, but for thee will;
|
||||
Bearing thy heart, which I will keep so chary
|
||||
As tender nurse her babe from faring ill.
|
||||
Presume not on thy heart when mine is slain;
|
||||
Thou gavest me thine, not to give back again.
|
||||
As an unperfect actor on the stage
|
||||
Who with his fear is put besides his part,
|
||||
Or some fierce thing replete with too much rage,
|
||||
Whose strength's abundance weakens his own heart.
|
||||
So I, for fear of trust, forget to say
|
||||
The perfect ceremony of love's rite,
|
||||
And in mine own love's strength seem to decay,
|
||||
O'ercharged with burden of mine own love's might.
|
||||
O, let my books be then the eloquence
|
||||
And dumb presagers of my speaking breast,
|
||||
Who plead for love and look for recompense
|
||||
More than that tongue that more hath more express'd.
|
||||
O, learn to read what silent love hath writ:
|
||||
To hear with eyes belongs to love's fine wit.
|
||||
Mine eye hath play'd the painter and hath stell'd
|
||||
Thy beauty's form in table of my heart;
|
||||
My body is the frame wherein 'tis held,
|
||||
And perspective it is the painter's art.
|
||||
For through the painter must you see his skill,
|
||||
To find where your true image pictured lies;
|
||||
Which in my bosom's shop is hanging still,
|
||||
That hath his windows glazed with thine eyes.
|
||||
Now see what good turns eyes for eyes have done:
|
||||
Mine eyes have drawn thy shape, and thine for me
|
||||
Are windows to my breast, where-through the sun
|
||||
Delights to peep, to gaze therein on thee;
|
||||
Yet eyes this cunning want to grace their art;
|
||||
They draw but what they see, know not the heart.
|
||||
Let those who are in favour with their stars
|
||||
Of public honour and proud titles boast,
|
||||
Whilst I, whom fortune of such triumph bars,
|
||||
Unlook'd for joy in that I honour most.
|
||||
Great princes' favourites their fair leaves spread
|
||||
But as the marigold at the sun's eye,
|
||||
And in themselves their pride lies buried,
|
||||
For at a frown they in their glory die.
|
||||
The painful warrior famoused for fight,
|
||||
After a thousand victories once foil'd,
|
||||
Is from the book of honour razed quite,
|
||||
And all the rest forgot for which he toil'd:
|
||||
Then happy I, that love and am beloved
|
||||
Where I may not remove nor be removed.
|
||||
Lord of my love, to whom in vassalage
|
||||
Thy merit hath my duty strongly knit,
|
||||
To thee I send this written embassage,
|
||||
To witness duty, not to show my wit:
|
||||
Duty so great, which wit so poor as mine
|
||||
May make seem bare, in wanting words to show it,
|
||||
But that I hope some good conceit of thine
|
||||
In thy soul's thought, all naked, will bestow it;
|
||||
Till whatsoever star that guides my moving
|
||||
Points on me graciously with fair aspect
|
||||
And puts apparel on my tatter'd loving,
|
||||
To show me worthy of thy sweet respect:
|
||||
Then may I dare to boast how I do love thee;
|
||||
Till then not show my head where thou mayst prove me.
|
||||
Weary with toil, I haste me to my bed,
|
||||
The dear repose for limbs with travel tired;
|
||||
But then begins a journey in my head,
|
||||
To work my mind, when body's work's expired:
|
||||
For then my thoughts, from far where I abide,
|
||||
Intend a zealous pilgrimage to thee,
|
||||
And keep my drooping eyelids open wide,
|
||||
Looking on darkness which the blind do see
|
||||
Save that my soul's imaginary sight
|
||||
Presents thy shadow to my sightless view,
|
||||
Which, like a jewel hung in ghastly night,
|
||||
Makes black night beauteous and her old face new.
|
||||
Lo! thus, by day my limbs, by night my mind,
|
||||
For thee and for myself no quiet find.
|
||||
How can I then return in happy plight,
|
||||
That am debarr'd the benefit of rest?
|
||||
When day's oppression is not eased by night,
|
||||
But day by night, and night by day, oppress'd?
|
||||
And each, though enemies to either's reign,
|
||||
Do in consent shake hands to torture me;
|
||||
The one by toil, the other to complain
|
||||
How far I toil, still farther off from thee.
|
||||
I tell the day, to please them thou art bright
|
||||
And dost him grace when clouds do blot the heaven:
|
||||
So flatter I the swart-complexion'd night,
|
||||
When sparkling stars twire not thou gild'st the even.
|
||||
But day doth daily draw my sorrows longer
|
||||
And night doth nightly make grief's strength seem stronger.
|
||||
When, in disgrace with fortune and men's eyes,
|
||||
I all alone beweep my outcast state
|
||||
And trouble deal heaven with my bootless cries
|
||||
And look upon myself and curse my fate,
|
||||
Wishing me like to one more rich in hope,
|
||||
Featured like him, like him with friends possess'd,
|
||||
Desiring this man's art and that man's scope,
|
||||
With what I most enjoy contented least;
|
||||
Yet in these thoughts myself almost despising,
|
||||
Haply I think on thee, and then my state,
|
||||
Like to the lark at break of day arising
|
||||
From sullen earth, sings hymns at heaven's gate;
|
||||
For thy sweet love remember'd such wealth brings
|
||||
That then I scorn to change my state with kings.
|
||||
When to the sessions of sweet silent thought
|
||||
I summon up remembrance of things past,
|
||||
I sigh the lack of many a thing I sought,
|
||||
And with old woes new wail my dear time's waste:
|
||||
Then can I drown an eye, unused to flow,
|
||||
For precious friends hid in death's dateless night,
|
||||
And weep afresh love's long since cancell'd woe,
|
||||
And moan the expense of many a vanish'd sight:
|
||||
Then can I grieve at grievances foregone,
|
||||
And heavily from woe to woe tell o'er
|
||||
The sad account of fore-bemoaned moan,
|
||||
Which I new pay as if not paid before.
|
||||
But if the while I think on thee, dear friend,
|
||||
All losses are restored and sorrows end.
|
||||
Thy bosom is endeared with all hearts,
|
||||
Which I by lacking have supposed dead,
|
||||
And there reigns love and all love's loving parts,
|
||||
And all those friends which I thought buried.
|
||||
How many a holy and obsequious tear
|
||||
Hath dear religious love stol'n from mine eye
|
||||
As interest of the dead, which now appear
|
||||
But things removed that hidden in thee lie!
|
||||
Thou art the grave where buried love doth live,
|
||||
Hung with the trophies of my lovers gone,
|
||||
Who all their parts of me to thee did give;
|
||||
That due of many now is thine alone:
|
||||
Their images I loved I view in thee,
|
||||
And thou, all they, hast all the all of me.
|
||||
If thou survive my well-contented day,
|
||||
When that churl Death my bones with dust shall cover,
|
||||
And shalt by fortune once more re-survey
|
||||
These poor rude lines of thy deceased lover,
|
||||
Compare them with the bettering of the time,
|
||||
And though they be outstripp'd by every pen,
|
||||
Reserve them for my love, not for their rhyme,
|
||||
Exceeded by the height of happier men.
|
||||
O, then vouchsafe me but this loving thought:
|
||||
'Had my friend's Muse grown with this growing age,
|
||||
A dearer birth than this his love had brought,
|
||||
To march in ranks of better equipage:
|
||||
But since he died and poets better prove,
|
||||
Theirs for their style I'll read, his for his love.'
|
||||
Full many a glorious morning have I seen
|
||||
Flatter the mountain-tops with sovereign eye,
|
||||
Kissing with golden face the meadows green,
|
||||
Gilding pale streams with heavenly alchemy;
|
||||
Anon permit the basest clouds to ride
|
||||
With ugly rack on his celestial face,
|
||||
And from the forlorn world his visage hide,
|
||||
Stealing unseen to west with this disgrace:
|
||||
Even so my sun one early morn did shine
|
||||
With all triumphant splendor on my brow;
|
||||
But out, alack! he was but one hour mine;
|
||||
The region cloud hath mask'd him from me now.
|
||||
Yet him for this my love no whit disdaineth;
|
||||
Suns of the world may stain when heaven's sun staineth.
|
||||
Why didst thou promise such a beauteous day,
|
||||
And make me travel forth without my cloak,
|
||||
To let base clouds o'ertake me in my way,
|
||||
Hiding thy bravery in their rotten smoke?
|
||||
'Tis not enough that through the cloud thou break,
|
||||
To dry the rain on my storm-beaten face,
|
||||
For no man well of such a salve can speak
|
||||
That heals the wound and cures not the disgrace:
|
||||
Nor can thy shame give physic to my grief;
|
||||
Though thou repent, yet I have still the loss:
|
||||
The offender's sorrow lends but weak relief
|
||||
To him that bears the strong offence's cross.
|
||||
Ah! but those tears are pearl which thy love sheds,
|
||||
And they are rich and ransom all ill deeds.
|
||||
No more be grieved at that which thou hast done:
|
||||
Roses have thorns, and silver fountains mud;
|
||||
Clouds and eclipses stain both moon and sun,
|
||||
And loathsome canker lives in sweetest bud.
|
||||
All men make faults, and even I in this,
|
||||
Authorizing thy trespass with compare,
|
||||
Myself corrupting, salving thy amiss,
|
||||
Excusing thy sins more than thy sins are;
|
||||
For to thy sensual fault I bring in sense--
|
||||
Thy adverse party is thy advocate--
|
||||
And 'gainst myself a lawful plea commence:
|
||||
Such civil war is in my love and hate
|
||||
That I an accessary needs must be
|
||||
To that sweet thief which sourly robs from me.
|
||||
Let me confess that we two must be twain,
|
||||
Although our undivided loves are one:
|
||||
So shall those blots that do with me remain
|
||||
Without thy help by me be borne alone.
|
||||
In our two loves there is but one respect,
|
||||
Though in our lives a separable spite,
|
||||
Which though it alter not love's sole effect,
|
||||
Yet doth it steal sweet hours from love's delight.
|
||||
I may not evermore acknowledge thee,
|
||||
Lest my bewailed guilt should do thee shame,
|
||||
Nor thou with public kindness honour me,
|
||||
Unless thou take that honour from thy name:
|
||||
But do not so; I love thee in such sort
|
||||
As, thou being mine, mine is thy good report.
|
||||
As a decrepit father takes delight
|
||||
To see his active child do deeds of youth,
|
||||
So I, made lame by fortune's dearest spite,
|
||||
Take all my comfort of thy worth and truth.
|
||||
For whether beauty, birth, or wealth, or wit,
|
||||
Or any of these all, or all, or more,
|
||||
Entitled in thy parts do crowned sit,
|
||||
I make my love engrafted to this store:
|
||||
So then I am not lame, poor, nor despised,
|
||||
Whilst that this shadow doth such substance give
|
||||
That I in thy abundance am sufficed
|
||||
And by a part of all thy glory live.
|
||||
Look, what is best, that best I wish in thee:
|
||||
This wish I have; then ten times happy me!
|
||||
@@ -0,0 +1,310 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use rand::rngs::StdRng;
|
||||
use rand::seq::SliceRandom;
|
||||
use rand::{Rng, SeedableRng};
|
||||
|
||||
use super::SampleRequest;
|
||||
use super::progress::RowDownloadReporter;
|
||||
use crate::cli::SpeedBenchConfig;
|
||||
use crate::error::{BenchError, Result};
|
||||
use crate::tokenizer::TokenizerKind;
|
||||
|
||||
/// Marker text for masked entries that need external fetch.
|
||||
const MASKED_PREFIX: &str = "FULL BENCHMARK DATA SHOULD BE FETCHED";
|
||||
|
||||
/// Cache directory for downloaded SPEED-Bench datasets.
|
||||
fn cache_dir() -> std::path::PathBuf {
|
||||
dirs::cache_dir()
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("/tmp"))
|
||||
.join("vllm-bench")
|
||||
.join("datasets")
|
||||
}
|
||||
|
||||
/// Download SPEED-Bench dataset from HuggingFace datasets-server API.
|
||||
/// Results are cached as JSON locally for subsequent runs.
|
||||
pub async fn download_speed_bench(config: SpeedBenchConfig) -> Result<String> {
|
||||
let config_name = config.as_str();
|
||||
|
||||
let dir = cache_dir();
|
||||
std::fs::create_dir_all(&dir)?;
|
||||
let cache_path = dir.join(format!("speed-bench-{config_name}.json"));
|
||||
|
||||
// Return cached file if it exists
|
||||
if cache_path.exists() {
|
||||
let path_str = cache_path.to_string_lossy().to_string();
|
||||
tracing::info!(config = config_name, path = %path_str, "using cached SPEED-Bench dataset");
|
||||
return Ok(path_str);
|
||||
}
|
||||
|
||||
tracing::info!(config = config_name, "downloading SPEED-Bench dataset");
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(120))
|
||||
.build()
|
||||
.map_err(|e| BenchError::Config(format!("Failed to build HTTP client: {e}")))?;
|
||||
|
||||
let mut all_rows: Vec<serde_json::Value> = Vec::new();
|
||||
let mut offset = 0usize;
|
||||
let page_size = 100usize;
|
||||
let mut progress = RowDownloadReporter::new();
|
||||
|
||||
loop {
|
||||
let url = format!(
|
||||
"https://datasets-server.huggingface.co/rows\
|
||||
?dataset=nvidia/SPEED-Bench\
|
||||
&config={config_name}\
|
||||
&split=test\
|
||||
&offset={offset}\
|
||||
&length={page_size}"
|
||||
);
|
||||
|
||||
// Retry on transient errors (502, 503, timeouts)
|
||||
let max_retries = 3;
|
||||
let mut data: Option<serde_json::Value> = None;
|
||||
for attempt in 0..=max_retries {
|
||||
let resp = match client.get(&url).send().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
if attempt < max_retries {
|
||||
tokio::time::sleep(std::time::Duration::from_secs(
|
||||
2 * (attempt as u64 + 1),
|
||||
))
|
||||
.await;
|
||||
continue;
|
||||
}
|
||||
return Err(BenchError::Config(format!(
|
||||
"SPEED-Bench download failed after {max_retries} retries: {e}"
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
if resp.status().is_server_error() && attempt < max_retries {
|
||||
tokio::time::sleep(std::time::Duration::from_secs(2 * (attempt as u64 + 1))).await;
|
||||
continue;
|
||||
}
|
||||
|
||||
if !resp.status().is_success() {
|
||||
return Err(BenchError::Config(format!(
|
||||
"SPEED-Bench API returned HTTP {}",
|
||||
resp.status()
|
||||
)));
|
||||
}
|
||||
|
||||
data = Some(resp.json().await.map_err(|e| {
|
||||
BenchError::Config(format!("Failed to parse SPEED-Bench API response: {e}"))
|
||||
})?);
|
||||
break;
|
||||
}
|
||||
|
||||
let data = data.unwrap();
|
||||
|
||||
let rows = data["rows"]
|
||||
.as_array()
|
||||
.ok_or_else(|| BenchError::Config("No 'rows' in API response".into()))?;
|
||||
|
||||
if rows.is_empty() {
|
||||
break;
|
||||
}
|
||||
|
||||
for row in rows {
|
||||
if let Some(row_data) = row.get("row") {
|
||||
all_rows.push(row_data.clone());
|
||||
}
|
||||
}
|
||||
|
||||
let fetched = rows.len();
|
||||
offset += fetched;
|
||||
|
||||
let total = data["num_rows_total"].as_u64().unwrap_or(0);
|
||||
progress.update(offset, total);
|
||||
|
||||
if fetched < page_size {
|
||||
break;
|
||||
}
|
||||
}
|
||||
progress.finish();
|
||||
|
||||
if all_rows.is_empty() {
|
||||
return Err(BenchError::Config(
|
||||
"SPEED-Bench download returned no rows".into(),
|
||||
));
|
||||
}
|
||||
|
||||
// Save to cache
|
||||
let json_str = serde_json::to_string(&all_rows)?;
|
||||
std::fs::write(&cache_path, &json_str)?;
|
||||
|
||||
let path_str = cache_path.to_string_lossy().to_string();
|
||||
tracing::info!(
|
||||
config = config_name,
|
||||
rows = all_rows.len(),
|
||||
path = %path_str,
|
||||
"saved SPEED-Bench dataset"
|
||||
);
|
||||
Ok(path_str)
|
||||
}
|
||||
|
||||
/// Load SPEED-Bench dataset and convert to SampleRequests.
|
||||
///
|
||||
/// Filters out masked entries and optionally filters by category.
|
||||
/// Requires an output length override since SPEED-Bench has no reference outputs.
|
||||
pub fn load_speed_bench_dataset(
|
||||
tokenizer: &TokenizerKind,
|
||||
dataset_path: &str,
|
||||
num_requests: usize,
|
||||
output_len: usize,
|
||||
seed: u64,
|
||||
request_id_prefix: &str,
|
||||
category_filter: Option<&str>,
|
||||
no_oversample: bool,
|
||||
disable_shuffle: bool,
|
||||
max_input_len: Option<usize>,
|
||||
) -> Result<Vec<SampleRequest>> {
|
||||
let content = std::fs::read_to_string(dataset_path).map_err(|e| {
|
||||
BenchError::Config(format!(
|
||||
"Failed to read SPEED-Bench file '{dataset_path}': {e}"
|
||||
))
|
||||
})?;
|
||||
|
||||
let entries: Vec<serde_json::Value> = serde_json::from_str(&content)
|
||||
.map_err(|e| BenchError::Config(format!("Invalid JSON in SPEED-Bench file: {e}")))?;
|
||||
|
||||
// Filter entries
|
||||
let mut filtered: Vec<&serde_json::Value> = entries
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
// Must have turns array with at least one non-empty entry
|
||||
let turns = match entry.get("turns").and_then(|t| t.as_array()) {
|
||||
Some(t) if !t.is_empty() => t,
|
||||
_ => return false,
|
||||
};
|
||||
|
||||
// Skip masked entries
|
||||
let first_turn = turns[0].as_str().unwrap_or("");
|
||||
if first_turn.starts_with(MASKED_PREFIX) || first_turn.is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Single-turn only: skip multi-turn entries
|
||||
let is_multiturn = entry.get("multiturn").and_then(|m| m.as_bool()).unwrap_or(false);
|
||||
if is_multiturn {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Category filter
|
||||
if let Some(cat) = category_filter {
|
||||
let entry_cat = entry.get("category").and_then(|c| c.as_str()).unwrap_or("");
|
||||
if entry_cat != cat {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
true
|
||||
})
|
||||
.collect();
|
||||
|
||||
if filtered.is_empty() {
|
||||
let cat_msg = category_filter.map(|c| format!(" with category '{c}'")).unwrap_or_default();
|
||||
return Err(BenchError::Config(format!(
|
||||
"No valid single-turn entries in SPEED-Bench{cat_msg}. \
|
||||
Try a different --speed-bench-config or remove --speed-bench-category filter."
|
||||
)));
|
||||
}
|
||||
|
||||
// Shuffle
|
||||
let mut rng = StdRng::seed_from_u64(seed);
|
||||
if !disable_shuffle {
|
||||
filtered.shuffle(&mut rng);
|
||||
}
|
||||
|
||||
// Build SampleRequests
|
||||
let mut samples = Vec::new();
|
||||
let mut idx = 0;
|
||||
|
||||
for entry in &filtered {
|
||||
if samples.len() >= num_requests {
|
||||
break;
|
||||
}
|
||||
|
||||
let turns = entry["turns"].as_array().unwrap();
|
||||
let prompt = turns[0].as_str().unwrap_or("");
|
||||
|
||||
// Tokenize to get prompt length
|
||||
let prompt_ids = tokenizer.encode(prompt, false)?;
|
||||
let prompt_len = prompt_ids.len();
|
||||
|
||||
if prompt_len < 4 {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Truncate if max_input_len is set
|
||||
let (final_prompt, final_len) = if let Some(max_len) = max_input_len {
|
||||
if prompt_len > max_len {
|
||||
let truncated_ids = &prompt_ids[..max_len];
|
||||
let truncated_text = tokenizer.decode(truncated_ids, true)?;
|
||||
(Arc::from(truncated_text.as_str()), max_len)
|
||||
} else {
|
||||
(Arc::from(prompt), prompt_len)
|
||||
}
|
||||
} else {
|
||||
(Arc::from(prompt), prompt_len)
|
||||
};
|
||||
|
||||
samples.push(SampleRequest {
|
||||
prompt: final_prompt,
|
||||
prompt_len: final_len,
|
||||
expected_output_len: output_len,
|
||||
request_id: Some(format!("{request_id_prefix}{idx}")),
|
||||
..Default::default()
|
||||
});
|
||||
idx += 1;
|
||||
}
|
||||
|
||||
// Oversample if needed
|
||||
if samples.len() < num_requests {
|
||||
if no_oversample {
|
||||
tracing::info!(
|
||||
dataset = "speed-bench",
|
||||
samples = samples.len(),
|
||||
requested = num_requests,
|
||||
"skipping dataset oversampling"
|
||||
);
|
||||
} else if !samples.is_empty() {
|
||||
let original_len = samples.len();
|
||||
let needed = num_requests - original_len;
|
||||
for i in 0..needed {
|
||||
let mut req = samples[rng.random_range(0..original_len)].clone();
|
||||
req.request_id = Some(format!("{request_id_prefix}{}", original_len + i));
|
||||
samples.push(req);
|
||||
}
|
||||
tracing::info!(
|
||||
dataset = "speed-bench",
|
||||
original_samples = original_len,
|
||||
samples = samples.len(),
|
||||
"oversampled dataset"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if samples.is_empty() {
|
||||
return Err(BenchError::Config(
|
||||
"No valid samples after filtering SPEED-Bench dataset.".into(),
|
||||
));
|
||||
}
|
||||
|
||||
let mut cat_counts: std::collections::HashMap<&str, usize> = std::collections::HashMap::new();
|
||||
for entry in &filtered[..filtered.len().min(samples.len())] {
|
||||
let cat = entry.get("category").and_then(|c| c.as_str()).unwrap_or("unknown");
|
||||
*cat_counts.entry(cat).or_insert(0) += 1;
|
||||
}
|
||||
let mut cats: Vec<_> = cat_counts.into_iter().collect();
|
||||
cats.sort_by_key(|b| std::cmp::Reverse(b.1));
|
||||
let cat_str: Vec<String> = cats.iter().map(|(k, v)| format!("{k}:{v}")).collect();
|
||||
tracing::info!(categories = %cat_str.join(", "), "computed SPEED-Bench category distribution");
|
||||
|
||||
Ok(samples)
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Error, Debug)]
|
||||
pub enum BenchError {
|
||||
#[error("HTTP request failed: {0}")]
|
||||
Http(#[from] reqwest::Error),
|
||||
|
||||
#[error("JSON error: {0}")]
|
||||
Json(#[from] serde_json::Error),
|
||||
|
||||
#[error("Tokenizer error: {0}")]
|
||||
Tokenizer(String),
|
||||
|
||||
/// The server's /tokenize//detokenize endpoint is not usable (4xx status:
|
||||
/// not exposed, or rejected by a gateway such as LLM-d/EPP that returns
|
||||
/// 400 instead of 404). Callers treat this as "skip verification", unlike
|
||||
/// `Tokenizer` errors which are genuine failures.
|
||||
#[error("tokenize endpoint unavailable: {0}")]
|
||||
TokenizeUnavailable(String),
|
||||
|
||||
#[error("Configuration error: {0}")]
|
||||
Config(String),
|
||||
|
||||
#[error("Endpoint not ready after {0}s: {1}")]
|
||||
EndpointTimeout(u64, String),
|
||||
|
||||
#[error("Backend error: {0}")]
|
||||
Backend(String),
|
||||
|
||||
#[error("IO error: {0}")]
|
||||
Io(#[from] std::io::Error),
|
||||
}
|
||||
|
||||
pub type Result<T> = std::result::Result<T, BenchError>;
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user