forked from Karylab-cklius/vllm
Compare commits
22
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
135453b715 | ||
|
|
a707288c1e | ||
|
|
638f8fa979 | ||
|
|
cbaa80fede | ||
|
|
84a1066ccc | ||
|
|
d801ae8c26 | ||
|
|
65df49eba3 | ||
|
|
2a2ac21d3d | ||
|
|
c6fc95806b | ||
|
|
581b5e9afc | ||
|
|
5536fc0c01 | ||
|
|
7f95e66a11 | ||
|
|
b1687527b8 | ||
|
|
171019ab19 | ||
|
|
879a8c3180 | ||
|
|
1b57eb41f2 | ||
|
|
21943d4c25 | ||
|
|
f396bee56f | ||
|
|
215e2f7990 | ||
|
|
e175192d33 | ||
|
|
a54f0d1049 | ||
|
|
48698b1b9b |
@@ -8,6 +8,7 @@ run_all_patterns:
|
||||
- "CMakeLists.txt"
|
||||
- "requirements/common.txt"
|
||||
- "requirements/cuda.txt"
|
||||
- "requirements/kv_connectors.txt"
|
||||
- "requirements/build/cuda.txt"
|
||||
- "requirements/test/cuda.txt"
|
||||
- "setup.py"
|
||||
|
||||
@@ -28,6 +28,7 @@ steps:
|
||||
- "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)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -41,6 +42,7 @@ steps:
|
||||
- "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)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -54,6 +56,7 @@ steps:
|
||||
- "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)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -67,6 +70,7 @@ steps:
|
||||
- "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)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -80,6 +84,7 @@ steps:
|
||||
- "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)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -93,6 +98,7 @@ steps:
|
||||
- "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)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -138,6 +144,7 @@ steps:
|
||||
# re-tag to default image tag and push, just in case arm64 build fails
|
||||
- "docker tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m) public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT"
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)"'
|
||||
|
||||
- label: "Build release image - aarch64 - CUDA 13.0"
|
||||
depends_on: ~
|
||||
@@ -160,6 +167,7 @@ steps:
|
||||
--progress plain \
|
||||
-f docker/Dockerfile .
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)"'
|
||||
|
||||
- label: "Build release image - x86_64 - CUDA 12.9"
|
||||
depends_on: ~
|
||||
@@ -184,6 +192,7 @@ steps:
|
||||
# re-tag to default image tag and push, just in case arm64 build fails
|
||||
- "docker tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129"
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129"'
|
||||
|
||||
- label: "Build release image - aarch64 - CUDA 12.9"
|
||||
depends_on: ~
|
||||
@@ -205,6 +214,7 @@ steps:
|
||||
--progress plain \
|
||||
-f docker/Dockerfile .
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129"'
|
||||
|
||||
- label: "Build release image - x86_64 - CUDA 13.0 - Ubuntu 24.04"
|
||||
depends_on: ~
|
||||
@@ -231,6 +241,7 @@ steps:
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-ubuntu2404"
|
||||
- "docker tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-ubuntu2404 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-ubuntu2404"
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-ubuntu2404"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-ubuntu2404"'
|
||||
|
||||
- label: "Build release image - aarch64 - CUDA 13.0 - Ubuntu 24.04"
|
||||
depends_on: ~
|
||||
@@ -255,6 +266,7 @@ steps:
|
||||
--progress plain \
|
||||
-f docker/Dockerfile .
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-ubuntu2404"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-ubuntu2404"'
|
||||
|
||||
- label: "Build release image - x86_64 - CUDA 12.9 - Ubuntu 24.04"
|
||||
depends_on: ~
|
||||
@@ -280,6 +292,7 @@ steps:
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129-ubuntu2404"
|
||||
- "docker tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129-ubuntu2404 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129-ubuntu2404"
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129-ubuntu2404"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129-ubuntu2404"'
|
||||
|
||||
- label: "Build release image - aarch64 - CUDA 12.9 - Ubuntu 24.04"
|
||||
depends_on: ~
|
||||
@@ -303,6 +316,7 @@ steps:
|
||||
--progress plain \
|
||||
-f docker/Dockerfile .
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129-ubuntu2404"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-$(uname -m)-cu129-ubuntu2404"'
|
||||
|
||||
- block: "Build release image for x86_64 CPU"
|
||||
key: block-cpu-release-image-build
|
||||
@@ -320,6 +334,7 @@ steps:
|
||||
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg GIT_REPO_CHECK=1 --build-arg VLLM_CPU_X86=true --tag public.ecr.aws/q9t5s3a7/vllm-cpu-release-repo:$(buildkite-agent meta-data get release-version) --tag public.ecr.aws/q9t5s3a7/vllm-cpu-release-repo:latest --progress plain --target vllm-openai -f docker/Dockerfile.cpu ."
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-cpu-release-repo:latest"
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-cpu-release-repo:$(buildkite-agent meta-data get release-version)"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-cpu-release-repo:$(buildkite-agent meta-data get release-version)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -339,6 +354,7 @@ steps:
|
||||
- "DOCKER_BUILDKIT=1 docker build --build-arg max_jobs=16 --build-arg GIT_REPO_CHECK=1 --tag public.ecr.aws/q9t5s3a7/vllm-arm64-cpu-release-repo:$(buildkite-agent meta-data get release-version) --tag public.ecr.aws/q9t5s3a7/vllm-arm64-cpu-release-repo:latest --progress plain --target vllm-openai -f docker/Dockerfile.cpu ."
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-arm64-cpu-release-repo:latest"
|
||||
- "docker push public.ecr.aws/q9t5s3a7/vllm-arm64-cpu-release-repo:$(buildkite-agent meta-data get release-version)"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "public.ecr.aws/q9t5s3a7/vllm-arm64-cpu-release-repo:$(buildkite-agent meta-data get release-version)"'
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
|
||||
@@ -356,15 +372,7 @@ steps:
|
||||
- "aws ecr-public get-login-password --region us-east-1 | docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7"
|
||||
- "docker manifest create public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-x86_64 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-aarch64 --amend"
|
||||
- "docker manifest push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT"
|
||||
|
||||
- label: "Annotate release workflow - CUDA 13.0"
|
||||
depends_on:
|
||||
- create-multi-arch-manifest
|
||||
id: annotate-release-workflow
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/annotate-release.sh"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "Manifest: CUDA 13.0" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT"'
|
||||
|
||||
- label: "Create multi-arch manifest - CUDA 12.9"
|
||||
depends_on:
|
||||
@@ -377,6 +385,7 @@ steps:
|
||||
- "aws ecr-public get-login-password --region us-east-1 | docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7"
|
||||
- "docker manifest create public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-x86_64-cu129 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-aarch64-cu129 --amend"
|
||||
- "docker manifest push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "Manifest: CUDA 12.9" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129"'
|
||||
|
||||
- label: "Create multi-arch manifest - CUDA 13.0 - Ubuntu 24.04"
|
||||
depends_on:
|
||||
@@ -389,6 +398,7 @@ steps:
|
||||
- "aws ecr-public get-login-password --region us-east-1 | docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7"
|
||||
- "docker manifest create public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-ubuntu2404 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-x86_64-ubuntu2404 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-aarch64-ubuntu2404 --amend"
|
||||
- "docker manifest push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-ubuntu2404"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "Manifest: CUDA 13.0 Ubuntu 24.04" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-ubuntu2404"'
|
||||
|
||||
- label: "Create multi-arch manifest - CUDA 12.9 - Ubuntu 24.04"
|
||||
depends_on:
|
||||
@@ -401,6 +411,7 @@ steps:
|
||||
- "aws ecr-public get-login-password --region us-east-1 | docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7"
|
||||
- "docker manifest create public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129-ubuntu2404 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-x86_64-cu129-ubuntu2404 public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-aarch64-cu129-ubuntu2404 --amend"
|
||||
- "docker manifest push public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129-ubuntu2404"
|
||||
- 'bash .buildkite/scripts/annotate-build-artifact.sh "Manifest: CUDA 12.9 Ubuntu 24.04" "public.ecr.aws/q9t5s3a7/vllm-release-repo:$BUILDKITE_COMMIT-cu129-ubuntu2404"'
|
||||
|
||||
- label: "Publish nightly multi-arch image to DockerHub"
|
||||
depends_on:
|
||||
@@ -438,59 +449,6 @@ steps:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
|
||||
- block: "Publish release images to DockerHub"
|
||||
key: block-publish-release-images
|
||||
depends_on:
|
||||
- create-multi-arch-manifest
|
||||
- create-multi-arch-manifest-cuda-12-9
|
||||
- create-multi-arch-manifest-ubuntu2404
|
||||
- create-multi-arch-manifest-cuda-12-9-ubuntu2404
|
||||
- build-rocm-release-image
|
||||
- input-release-version
|
||||
# Wait for CPU builds if their block steps were unblocked, so publish
|
||||
# doesn't race the in-progress CPU build. allow_failure lets publish
|
||||
# proceed when the operator legitimately leaves the CPU block steps
|
||||
# unblocked or the CPU build fails.
|
||||
- step: build-cpu-release-image-x86
|
||||
allow_failure: true
|
||||
- step: build-cpu-release-image-arm64
|
||||
allow_failure: true
|
||||
if: build.env("NIGHTLY") != "1"
|
||||
|
||||
- label: "Publish release images to DockerHub"
|
||||
depends_on:
|
||||
- block-publish-release-images
|
||||
key: publish-release-images-dockerhub
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/publish-release-images.sh"
|
||||
plugins:
|
||||
- docker-login#v3.0.0:
|
||||
username: vllmbot
|
||||
password-env: DOCKERHUB_TOKEN
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
|
||||
- group: "Publish wheels"
|
||||
key: "publish-wheels"
|
||||
steps:
|
||||
- block: "Confirm update release wheels to PyPI (experimental, use with caution)?"
|
||||
key: block-upload-release-wheels
|
||||
depends_on:
|
||||
- input-release-version
|
||||
- build-wheels
|
||||
|
||||
- label: "Upload release wheels to PyPI"
|
||||
depends_on:
|
||||
- block-upload-release-wheels
|
||||
id: upload-release-wheels
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/upload-release-wheels-pypi.sh"
|
||||
|
||||
# =============================================================================
|
||||
# ROCm Release Pipeline (x86_64 only)
|
||||
# =============================================================================
|
||||
@@ -604,7 +562,7 @@ steps:
|
||||
echo ""
|
||||
echo " Build complete - Image and wheels cached"
|
||||
fi
|
||||
|
||||
|
||||
artifact_paths:
|
||||
- "artifacts/rocm-base-wheels/*.whl"
|
||||
env:
|
||||
@@ -820,7 +778,7 @@ steps:
|
||||
|
||||
# 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"
|
||||
@@ -847,3 +805,60 @@ steps:
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
|
||||
# =============================================================================
|
||||
# Publish to DockerHub and PyPI (at the end so all builds complete first)
|
||||
# =============================================================================
|
||||
|
||||
- block: "Publish release images to DockerHub"
|
||||
key: block-publish-release-images
|
||||
depends_on:
|
||||
- create-multi-arch-manifest
|
||||
- create-multi-arch-manifest-cuda-12-9
|
||||
- create-multi-arch-manifest-ubuntu2404
|
||||
- create-multi-arch-manifest-cuda-12-9-ubuntu2404
|
||||
- build-rocm-release-image
|
||||
- input-release-version
|
||||
# Wait for CPU builds if their block steps were unblocked, so publish
|
||||
# doesn't race the in-progress CPU build. allow_failure lets publish
|
||||
# proceed when the operator legitimately leaves the CPU block steps
|
||||
# unblocked or the CPU build fails.
|
||||
- step: build-cpu-release-image-x86
|
||||
allow_failure: true
|
||||
- step: build-cpu-release-image-arm64
|
||||
allow_failure: true
|
||||
if: build.env("NIGHTLY") != "1"
|
||||
|
||||
- label: "Publish release images to DockerHub"
|
||||
depends_on:
|
||||
- block-publish-release-images
|
||||
key: publish-release-images-dockerhub
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/publish-release-images.sh"
|
||||
plugins:
|
||||
- docker-login#v3.0.0:
|
||||
username: vllmbot
|
||||
password-env: DOCKERHUB_TOKEN
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
DOCKERHUB_USERNAME: "vllmbot"
|
||||
|
||||
- group: "Publish wheels"
|
||||
key: "publish-wheels"
|
||||
steps:
|
||||
- block: "Confirm update release wheels to PyPI (experimental, use with caution)?"
|
||||
key: block-upload-release-wheels
|
||||
depends_on:
|
||||
- input-release-version
|
||||
- build-wheels
|
||||
|
||||
- label: "Upload release wheels to PyPI"
|
||||
depends_on:
|
||||
- block-upload-release-wheels
|
||||
id: upload-release-wheels
|
||||
agents:
|
||||
queue: small_cpu_queue_release
|
||||
commands:
|
||||
- "bash .buildkite/scripts/upload-release-wheels-pypi.sh"
|
||||
|
||||
+9
@@ -0,0 +1,9 @@
|
||||
#!/bin/bash
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# Append a build artifact line to the Buildkite annotation.
|
||||
# Usage: annotate-build-artifact.sh <label> <value>
|
||||
set -e
|
||||
echo "- **${1}**: \`${2}\`" | \
|
||||
buildkite-agent annotate --append --style 'info' --context 'release-artifacts'
|
||||
@@ -1,27 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
set -ex
|
||||
|
||||
# Get release version, default to 1.0.0.dev for nightly/per-commit builds
|
||||
RELEASE_VERSION=$(buildkite-agent meta-data get release-version 2>/dev/null | sed 's/^v//')
|
||||
if [ -z "${RELEASE_VERSION}" ]; then
|
||||
RELEASE_VERSION="1.0.0.dev"
|
||||
fi
|
||||
|
||||
buildkite-agent annotate --style 'info' --context 'release-workflow' << EOF
|
||||
To download the wheel (by commit):
|
||||
\`\`\`
|
||||
aws s3 cp s3://vllm-wheels/${BUILDKITE_COMMIT}/vllm-${RELEASE_VERSION}-cp38-abi3-manylinux_2_35_x86_64.whl .
|
||||
aws s3 cp s3://vllm-wheels/${BUILDKITE_COMMIT}/vllm-${RELEASE_VERSION}-cp38-abi3-manylinux_2_35_aarch64.whl .
|
||||
|
||||
(Optional) For CUDA 12.9:
|
||||
aws s3 cp s3://vllm-wheels/${BUILDKITE_COMMIT}/vllm-${RELEASE_VERSION}+cu129-cp38-abi3-manylinux_2_31_x86_64.whl .
|
||||
aws s3 cp s3://vllm-wheels/${BUILDKITE_COMMIT}/vllm-${RELEASE_VERSION}+cu129-cp38-abi3-manylinux_2_31_aarch64.whl .
|
||||
|
||||
(Optional) For CPU:
|
||||
aws s3 cp s3://vllm-wheels/${BUILDKITE_COMMIT}/vllm-${RELEASE_VERSION}+cpu-cp38-abi3-manylinux_2_35_x86_64.whl .
|
||||
aws s3 cp s3://vllm-wheels/${BUILDKITE_COMMIT}/vllm-${RELEASE_VERSION}+cpu-cp38-abi3-manylinux_2_35_aarch64.whl .
|
||||
\`\`\`
|
||||
|
||||
Docker images are published automatically by the "Publish release images to DockerHub" pipeline step.
|
||||
EOF
|
||||
@@ -105,7 +105,11 @@ steps:
|
||||
device: h100
|
||||
num_devices: 1
|
||||
source_file_dependencies:
|
||||
- cmake/external_projects/deepgemm.cmake
|
||||
- tools/install_deepgemm.sh
|
||||
- tools/build_deepgemm_C.py
|
||||
- tools/setup_deepgemm_pythons.sh
|
||||
- tools/check_wheel_deepgemm.py
|
||||
- vllm/utils/deep_gemm.py
|
||||
- vllm/model_executor/layers/fused_moe
|
||||
- vllm/model_executor/layers/quantization
|
||||
@@ -115,6 +119,7 @@ steps:
|
||||
- tests/kernels/attention/test_deepgemm_attention.py
|
||||
- tests/quantization/test_cutlass_w4a16.py
|
||||
commands:
|
||||
- python3 ../tools/check_wheel_deepgemm.py
|
||||
- pytest -v -s kernels/quantization/test_block_fp8.py
|
||||
- pytest -v -s kernels/moe/test_deepgemm.py
|
||||
- pytest -v -s kernels/moe/test_batched_deepgemm.py
|
||||
|
||||
@@ -9,6 +9,9 @@ PATH=${cuda_home}/bin:$PATH
|
||||
LD_LIBRARY_PATH=${cuda_home}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install requirements
|
||||
if [ "$(echo $2 | cut -d. -f1)" = "12" ]; then
|
||||
sed -i 's/^nvidia-cutlass-dsl\[cu13\]>=/nvidia-cutlass-dsl>=/' requirements/cuda.txt
|
||||
fi
|
||||
$python_executable -m pip install -r requirements/build/cuda.txt -r requirements/cuda.txt
|
||||
|
||||
# Limit the number of parallel jobs to avoid OOM
|
||||
|
||||
@@ -53,48 +53,67 @@ cuda_archs_loose_intersection(DEEPGEMM_ARCHS
|
||||
if(DEEPGEMM_ARCHS)
|
||||
message(STATUS "DeepGEMM CUDA architectures: ${DEEPGEMM_ARCHS}")
|
||||
|
||||
find_package(CUDAToolkit REQUIRED)
|
||||
# Build _C once per interpreter in DEEPGEMM_PYTHON_INTERPRETERS (":"-
|
||||
# separated paths) so the wheel imports cleanly on every supported Python.
|
||||
# Unset → fall back to the build interpreter (editable / source builds).
|
||||
# The compile is delegated to tools/build_deepgemm_C.py and always runs
|
||||
# against the build interpreter's torch — target Pythons don't need torch.
|
||||
# Note: empty-but-set env vars are still DEFINED in cmake; treat empty as
|
||||
# unset so an empty interpreter list falls back to the build interpreter
|
||||
# rather than silently skipping the per-Python build.
|
||||
if(NOT "$ENV{DEEPGEMM_PYTHON_INTERPRETERS}" STREQUAL "")
|
||||
string(REPLACE ":" ";" _dg_pythons "$ENV{DEEPGEMM_PYTHON_INTERPRETERS}")
|
||||
else()
|
||||
set(_dg_pythons "${Python_EXECUTABLE}")
|
||||
endif()
|
||||
message(STATUS "DeepGEMM _C will be built for: ${_dg_pythons}")
|
||||
|
||||
#
|
||||
# Build the _C pybind11 extension from DeepGEMM's C++ source.
|
||||
# This is a CXX-only module — CUDA kernels are JIT-compiled at runtime.
|
||||
#
|
||||
Python_add_library(_deep_gemm_C MODULE WITH_SOABI
|
||||
"${deepgemm_SOURCE_DIR}/csrc/python_api.cpp")
|
||||
# Header set fed to add_custom_command's DEPENDS so a header-only edit
|
||||
# (in upstream DeepGEMM or its vendored cutlass/fmt) re-triggers the
|
||||
# rebuild. add_custom_command does no implicit header scanning, unlike
|
||||
# add_library.
|
||||
file(GLOB_RECURSE _dg_headers
|
||||
"${deepgemm_SOURCE_DIR}/csrc/*.h"
|
||||
"${deepgemm_SOURCE_DIR}/csrc/*.hpp"
|
||||
"${deepgemm_SOURCE_DIR}/deep_gemm/include/*.h"
|
||||
"${deepgemm_SOURCE_DIR}/deep_gemm/include/*.hpp"
|
||||
"${deepgemm_SOURCE_DIR}/deep_gemm/include/*.cuh")
|
||||
|
||||
# The pybind11 module name must be _C to match DeepGEMM's Python imports.
|
||||
set_target_properties(_deep_gemm_C PROPERTIES OUTPUT_NAME "_C")
|
||||
|
||||
target_compile_definitions(_deep_gemm_C PRIVATE
|
||||
"-DTORCH_EXTENSION_NAME=_C")
|
||||
|
||||
target_include_directories(_deep_gemm_C PRIVATE
|
||||
"${deepgemm_SOURCE_DIR}/csrc"
|
||||
"${deepgemm_SOURCE_DIR}/deep_gemm/include"
|
||||
"${deepgemm_SOURCE_DIR}/third-party/cutlass/include"
|
||||
"${deepgemm_SOURCE_DIR}/third-party/cutlass/tools/util/include"
|
||||
"${deepgemm_SOURCE_DIR}/third-party/fmt/include")
|
||||
|
||||
target_compile_options(_deep_gemm_C PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-O3>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-Wno-psabi>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-Wno-deprecated-declarations>)
|
||||
|
||||
# torch_python is required because DeepGEMM uses pybind11 type casters
|
||||
# for at::Tensor (via PYBIND11_MODULE), unlike vLLM's own extensions which
|
||||
# use torch::Library custom ops.
|
||||
find_library(TORCH_PYTHON_LIBRARY torch_python
|
||||
PATHS "${TORCH_INSTALL_PREFIX}/lib"
|
||||
REQUIRED)
|
||||
|
||||
target_link_libraries(_deep_gemm_C PRIVATE
|
||||
torch ${TORCH_LIBRARIES} "${TORCH_PYTHON_LIBRARY}"
|
||||
CUDA::cudart CUDA::nvrtc)
|
||||
|
||||
# Install the shared library into the vendored package directory
|
||||
install(TARGETS _deep_gemm_C
|
||||
LIBRARY DESTINATION vllm/third_party/deep_gemm
|
||||
COMPONENT _deep_gemm_C)
|
||||
set(_dg_markers)
|
||||
set(_dg_seen_soabis)
|
||||
foreach(_pybin IN LISTS _dg_pythons)
|
||||
execute_process(
|
||||
COMMAND "${_pybin}" -c
|
||||
"import sysconfig; print(sysconfig.get_config_var('SOABI'))"
|
||||
OUTPUT_VARIABLE _dg_soabi
|
||||
OUTPUT_STRIP_TRAILING_WHITESPACE
|
||||
COMMAND_ERROR_IS_FATAL ANY)
|
||||
# Dedup so duplicate paths (or two paths resolving to the same CPython)
|
||||
# don't register conflicting build rules.
|
||||
if(_dg_soabi IN_LIST _dg_seen_soabis)
|
||||
continue()
|
||||
endif()
|
||||
list(APPEND _dg_seen_soabis "${_dg_soabi}")
|
||||
set(_dg_dir "${CMAKE_CURRENT_BINARY_DIR}/deepgemm_C_${_dg_soabi}")
|
||||
set(_dg_marker "${_dg_dir}/.built")
|
||||
add_custom_command(
|
||||
OUTPUT "${_dg_marker}"
|
||||
COMMAND "${Python_EXECUTABLE}"
|
||||
"${CMAKE_SOURCE_DIR}/tools/build_deepgemm_C.py"
|
||||
"${deepgemm_SOURCE_DIR}" "${_dg_dir}" "${_pybin}"
|
||||
COMMAND "${CMAKE_COMMAND}" -E touch "${_dg_marker}"
|
||||
DEPENDS "${CMAKE_SOURCE_DIR}/tools/build_deepgemm_C.py"
|
||||
"${deepgemm_SOURCE_DIR}/csrc/python_api.cpp"
|
||||
${_dg_headers}
|
||||
COMMENT "Building DeepGEMM _C for ${_pybin}"
|
||||
VERBATIM)
|
||||
list(APPEND _dg_markers "${_dg_marker}")
|
||||
install(DIRECTORY "${_dg_dir}/"
|
||||
DESTINATION vllm/third_party/deep_gemm
|
||||
COMPONENT _deep_gemm_C
|
||||
FILES_MATCHING PATTERN "_C.cpython-*.so")
|
||||
endforeach()
|
||||
add_custom_target(_deep_gemm_C ALL DEPENDS ${_dg_markers})
|
||||
|
||||
#
|
||||
# Vendor DeepGEMM Python package files
|
||||
|
||||
@@ -156,6 +156,17 @@ inline int GetGroupsPerBlock(int64_t num_groups) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
// Largest divisor of padded_groups_per_row that is <= 16. ry = 16 / kx.
|
||||
inline int GetGroupsPerBlockX(int64_t padded_groups_per_row) {
|
||||
if (padded_groups_per_row % 16 == 0) {
|
||||
return 16;
|
||||
}
|
||||
if (padded_groups_per_row % 8 == 0) {
|
||||
return 8;
|
||||
}
|
||||
return 4;
|
||||
}
|
||||
|
||||
void per_token_group_quant_8bit(const torch::stable::Tensor& input,
|
||||
torch::stable::Tensor& output_q,
|
||||
torch::stable::Tensor& output_s,
|
||||
@@ -247,11 +258,11 @@ void per_token_group_quant_8bit(const torch::stable::Tensor& input,
|
||||
//
|
||||
// Constraints: GROUP_SIZE % (THREADS_PER_GROUP * VEC_SIZE) == 0; for
|
||||
// THREADS_PER_GROUP=8 and bf16/fp16 (VEC_SIZE=16), this means GROUP_SIZE=128.
|
||||
template <typename T, typename DST_DTYPE, int GROUP_SIZE>
|
||||
template <typename T, typename DST_DTYPE, int GROUP_SIZE, int kGroupsPerBlockX,
|
||||
int kRowsPerBlock>
|
||||
__global__ void per_token_group_quant_8bit_packed_register_kernel(
|
||||
const T* __restrict__ input, void* __restrict__ output_q,
|
||||
unsigned int* __restrict__ output_s_packed, const int64_t num_groups_padded,
|
||||
const int groups_per_block, const int padded_groups_per_row,
|
||||
unsigned int* __restrict__ output_s_packed, const int padded_groups_per_row,
|
||||
const int groups_per_row, const int mn, const int output_q_mn_extent,
|
||||
const int tma_aligned_mn, const int64_t num_scale_elems, const float eps,
|
||||
const float min_8bit, const float max_8bit) {
|
||||
@@ -260,27 +271,25 @@ __global__ void per_token_group_quant_8bit_packed_register_kernel(
|
||||
constexpr int VEC_SIZE = 32 / sizeof(T); // 16 for bf16/fp16
|
||||
static_assert(GROUP_SIZE == THREADS_PER_GROUP * VEC_SIZE,
|
||||
"GROUP_SIZE must equal THREADS_PER_GROUP * VEC_SIZE");
|
||||
// Each group's 8 threads must live in a single warp octet so the
|
||||
// 0xffu << (threadIdx.x & 24u) shuffle mask selects exactly the lanes
|
||||
// that share a group. Requires 32 % THREADS_PER_GROUP == 0 and the host
|
||||
// to launch num_threads as a multiple of THREADS_PER_GROUP (which it does
|
||||
// via num_threads = groups_per_block * THREADS_PER_GROUP).
|
||||
static_assert(32 % THREADS_PER_GROUP == 0,
|
||||
"THREADS_PER_GROUP must divide warp size for the shuffle "
|
||||
"mask to be valid");
|
||||
static_assert(
|
||||
kGroupsPerBlockX > 0 && (kGroupsPerBlockX & (kGroupsPerBlockX - 1)) == 0,
|
||||
"kGroupsPerBlockX must be a positive power of 2");
|
||||
static_assert(kRowsPerBlock > 0, "kRowsPerBlock must be positive");
|
||||
|
||||
const int local_group_id = threadIdx.x / THREADS_PER_GROUP;
|
||||
const int lane_id = threadIdx.x % THREADS_PER_GROUP;
|
||||
|
||||
const int64_t block_group_id = blockIdx.x * groups_per_block;
|
||||
const int64_t global_group_id = block_group_id + local_group_id;
|
||||
if (global_group_id >= num_groups_padded) {
|
||||
const int sf_k_local = local_group_id % kGroupsPerBlockX;
|
||||
const int row_local = local_group_id / kGroupsPerBlockX;
|
||||
const int sf_k_idx = blockIdx.x * kGroupsPerBlockX + sf_k_local;
|
||||
const int mn_idx = blockIdx.y * kRowsPerBlock + row_local;
|
||||
|
||||
if (mn_idx >= tma_aligned_mn) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int sf_k_idx =
|
||||
static_cast<int>(global_group_id % padded_groups_per_row);
|
||||
const int mn_idx = static_cast<int>(global_group_id / padded_groups_per_row);
|
||||
const bool is_valid_group = (mn_idx < mn) && (sf_k_idx < groups_per_row);
|
||||
|
||||
// Load 16 input elements (32 B) into registers as two adjacent uint4
|
||||
@@ -443,34 +452,53 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
|
||||
|
||||
constexpr int THREADS_PER_GROUP = 8;
|
||||
const int64_t padded_groups_per_row = k_num_packed_sfk * 4;
|
||||
const int64_t num_groups_padded = tma_aligned_mn * padded_groups_per_row;
|
||||
const int64_t num_scale_elems = mn + (k_num_packed_sfk - 1) * tma_aligned_mn;
|
||||
const int groups_per_block = GetGroupsPerBlock(num_groups_padded);
|
||||
|
||||
STD_TORCH_CHECK(padded_groups_per_row % 4 == 0,
|
||||
"padded_groups_per_row=", padded_groups_per_row,
|
||||
" is not a multiple of 4.");
|
||||
const int kx = GetGroupsPerBlockX(padded_groups_per_row);
|
||||
const int ry = 16 / kx;
|
||||
const int64_t blocks_x = padded_groups_per_row / kx;
|
||||
const int64_t blocks_y = (tma_aligned_mn + ry - 1) / ry;
|
||||
const int num_threads = (kx * ry) * THREADS_PER_GROUP;
|
||||
// CUDA caps grid.x and grid.y at 2^31 - 1; guard against pathological inputs.
|
||||
STD_TORCH_CHECK(blocks_x <= static_cast<int64_t>(INT32_MAX) &&
|
||||
blocks_y <= static_cast<int64_t>(INT32_MAX),
|
||||
"per_token_group_quant_8bit_packed grid too large: (",
|
||||
blocks_x, ", ", blocks_y, ").");
|
||||
|
||||
auto dst_type = output_q.scalar_type();
|
||||
const int64_t num_blocks = num_groups_padded / groups_per_block;
|
||||
const int num_threads = groups_per_block * THREADS_PER_GROUP;
|
||||
// CUDA caps grid.x at 2^31 - 1; this fits any realistic shape but guard
|
||||
// against pathological inputs.
|
||||
STD_TORCH_CHECK(num_blocks <= static_cast<int64_t>(INT32_MAX),
|
||||
"per_token_group_quant_8bit_packed grid too large: ",
|
||||
num_blocks, " blocks (max ", INT32_MAX, ").");
|
||||
|
||||
#define LAUNCH_REG_KERNEL(T, DST_DTYPE) \
|
||||
do { \
|
||||
dim3 grid(static_cast<unsigned int>(num_blocks)); \
|
||||
dim3 block(num_threads); \
|
||||
per_token_group_quant_8bit_packed_register_kernel<T, DST_DTYPE, 128> \
|
||||
<<<grid, block, 0, stream>>>( \
|
||||
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
|
||||
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
|
||||
num_groups_padded, groups_per_block, \
|
||||
static_cast<int>(padded_groups_per_row), \
|
||||
static_cast<int>(groups_per_row), static_cast<int>(mn), \
|
||||
static_cast<int>(output_q_mn_extent), \
|
||||
static_cast<int>(tma_aligned_mn), num_scale_elems, \
|
||||
static_cast<float>(eps), static_cast<float>(min_8bit), \
|
||||
static_cast<float>(max_8bit)); \
|
||||
#define LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, KX, RY) \
|
||||
do { \
|
||||
dim3 grid(static_cast<unsigned int>(blocks_x), \
|
||||
static_cast<unsigned int>(blocks_y)); \
|
||||
dim3 block(num_threads); \
|
||||
per_token_group_quant_8bit_packed_register_kernel<T, DST_DTYPE, 128, KX, \
|
||||
RY> \
|
||||
<<<grid, block, 0, stream>>>( \
|
||||
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
|
||||
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
|
||||
static_cast<int>(padded_groups_per_row), \
|
||||
static_cast<int>(groups_per_row), static_cast<int>(mn), \
|
||||
static_cast<int>(output_q_mn_extent), \
|
||||
static_cast<int>(tma_aligned_mn), num_scale_elems, \
|
||||
static_cast<float>(eps), static_cast<float>(min_8bit), \
|
||||
static_cast<float>(max_8bit)); \
|
||||
} while (0)
|
||||
|
||||
#define LAUNCH_REG_KERNEL(T, DST_DTYPE) \
|
||||
do { \
|
||||
if (kx == 16) { \
|
||||
LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, 16, 1); \
|
||||
} else if (kx == 8) { \
|
||||
LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, 8, 2); \
|
||||
} else if (kx == 4) { \
|
||||
LAUNCH_REG_KERNEL_INST(T, DST_DTYPE, 4, 4); \
|
||||
} else { \
|
||||
STD_TORCH_CHECK(false, "Unsupported kx value ", kx); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
VLLM_STABLE_DISPATCH_HALF_TYPES(
|
||||
@@ -488,6 +516,7 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
|
||||
}));
|
||||
|
||||
#undef LAUNCH_REG_KERNEL
|
||||
#undef LAUNCH_REG_KERNEL_INST
|
||||
}
|
||||
|
||||
void per_token_group_quant_fp8(const torch::stable::Tensor& input,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#include <cstdint>
|
||||
|
||||
/* Adapted from ./csrc/quantization/gguf/mmq.cuh
|
||||
based on ./vllm/model_executor/layers/fused_moe/fused_moe.py */
|
||||
based on ./vllm/model_executor/layers/fused_moe/experts/triton_moe.py */
|
||||
template <typename scalar_t, int qk, int qr, int qi, bool need_sum,
|
||||
typename block_q_t, int mmq_x, int mmq_y, int nwarps,
|
||||
allocate_tiles_cuda_t allocate_tiles, load_tiles_cuda_t load_tiles,
|
||||
|
||||
+18
-1
@@ -199,7 +199,10 @@ COPY requirements/cuda.txt requirements/cuda.txt
|
||||
COPY use_existing_torch.py use_existing_torch.py
|
||||
COPY pyproject.toml pyproject.toml
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
if [ "${PYTORCH_NIGHTLY}" = "1" ]; then \
|
||||
if [ "$(echo $CUDA_VERSION | cut -d. -f1)" = "12" ]; then \
|
||||
sed -i 's/^nvidia-cutlass-dsl\[cu13\]>=/nvidia-cutlass-dsl>=/' requirements/cuda.txt; \
|
||||
fi \
|
||||
&& if [ "${PYTORCH_NIGHTLY}" = "1" ]; then \
|
||||
echo "Installing torch nightly..." \
|
||||
&& uv pip install --python /opt/venv/bin/python3 torch torchaudio torchvision --pre \
|
||||
--index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/nightly/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.') \
|
||||
@@ -301,6 +304,15 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
python3 use_existing_torch.py --prefix; \
|
||||
fi
|
||||
|
||||
# Provision a bare interpreter for each CPython covered by `requires-python`
|
||||
# so DeepGEMM `_C` is built once per Python and bundled side-by-side in the
|
||||
# wheel; cmake reads DEEPGEMM_PYTHON_INTERPRETERS in deepgemm.cmake's
|
||||
# foreach loop. The matrix is derived from pyproject.toml.
|
||||
COPY tools/setup_deepgemm_pythons.sh tools/build_deepgemm_C.py tools/
|
||||
ENV DEEPGEMM_VENV_PREFIX=/opt/dgenv
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
tools/setup_deepgemm_pythons.sh > /tmp/dg_pythons.txt
|
||||
|
||||
# Build the vLLM wheel
|
||||
# if USE_SCCACHE is set, use sccache to speed up compilation
|
||||
# AWS credentials mounted at ~/.aws/credentials for sccache S3 auth (optional)
|
||||
@@ -328,6 +340,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
&& export VLLM_PRECOMPILED_WHEEL_COMMIT="${VLLM_MERGE_BASE_COMMIT}" \
|
||||
&& export VLLM_MAIN_CUDA_VERSION="${VLLM_MAIN_CUDA_VERSION}" \
|
||||
&& export VLLM_DOCKER_BUILD_CONTEXT=1 \
|
||||
&& export DEEPGEMM_PYTHON_INTERPRETERS=$(cat /tmp/dg_pythons.txt) \
|
||||
&& sccache --show-stats \
|
||||
&& python3 setup.py bdist_wheel --dist-dir=dist --py-limited-api=cp38 \
|
||||
&& sccache --show-stats; \
|
||||
@@ -345,6 +358,7 @@ RUN --mount=type=cache,target=/root/.cache/ccache \
|
||||
export VLLM_USE_PRECOMPILED="${VLLM_USE_PRECOMPILED}" && \
|
||||
export VLLM_PRECOMPILED_WHEEL_COMMIT="${VLLM_MERGE_BASE_COMMIT}" && \
|
||||
export VLLM_DOCKER_BUILD_CONTEXT=1 && \
|
||||
export DEEPGEMM_PYTHON_INTERPRETERS=$(cat /tmp/dg_pythons.txt) && \
|
||||
python3 setup.py bdist_wheel --dist-dir=dist --py-limited-api=cp38; \
|
||||
fi
|
||||
|
||||
@@ -616,6 +630,9 @@ ARG PYTORCH_CUDA_INDEX_BASE_URL
|
||||
COPY requirements/common.txt /tmp/common.txt
|
||||
COPY requirements/cuda.txt /tmp/requirements-cuda.txt
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
if [ "$(echo $CUDA_VERSION | cut -d. -f1)" = "12" ]; then \
|
||||
sed -i 's/^nvidia-cutlass-dsl\[cu13\]>=/nvidia-cutlass-dsl>=/' /tmp/requirements-cuda.txt; \
|
||||
fi && \
|
||||
uv pip install --system -r /tmp/requirements-cuda.txt \
|
||||
--extra-index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.') && \
|
||||
rm /tmp/requirements-cuda.txt /tmp/common.txt
|
||||
|
||||
@@ -142,7 +142,7 @@ We use "mamba-like" to refer to layers that possess a state that is updated in-p
|
||||
For implementing new custom mamba-like layers, one should inherit from `MambaBase` and implement the methods `get_state_dtype`, `get_state_shape` to calculate the data types and state shapes at runtime, as well as `mamba_type` and `get_attn_backend`.
|
||||
It is also necessary to implement the "attention meta-data" class which handles the meta-data that is common across all layers.
|
||||
Please see [`LinearAttentionMetadata`](../../../vllm/v1/attention/backends/linear_attn.py) or [`ShortConvAttentionMetadata`](../../../vllm/v1/attention/backends/short_conv_attn.py) for examples of this.
|
||||
It is also worth noting that we should update `MAMBA_TYPE_TO_BACKEND_MAP` and `MambaAttentionBackendEnum` in [`registry.py`](../../../vllm/v1/attention/backends/registry.py) when adding a new mamba backend.
|
||||
It is also worth noting that we should update `MambaAttentionBackendEnum` in [`registry.py`](../../../vllm/v1/attention/backends/registry.py) when adding a new mamba backend.
|
||||
Finally, if one wants to support torch compile and CUDA graphs, it necessary to wrap the call to the mamba-like layer inside a custom op and register it.
|
||||
Please see the calls to `direct_register_custom_op` in [vllm/model_executor/models/minimax_text_01.py](../../../vllm/model_executor/models/minimax_text_01.py) or [vllm/model_executor/layers/mamba/short_conv.py](../../../vllm/model_executor/layers/mamba/short_conv.py) for examples of this.
|
||||
The new custom op should then be added to the list `_attention_ops` in [vllm/config/compilation.py](../../../vllm/config/compilation.py) to ensure that piecewise CUDA graphs works as intended.
|
||||
|
||||
@@ -138,7 +138,7 @@ For example:
|
||||
|
||||
--8<-- "vllm/model_executor/models/transformers/moe.py:transformers_fused_moe"
|
||||
|
||||
--8<-- "vllm/model_executor/layers/fused_moe/fused_moe.py:grouped_topk"
|
||||
--8<-- "vllm/model_executor/layers/fused_moe/router/grouped_topk_router.py:grouped_topk"
|
||||
```
|
||||
|
||||
**9. Norm:**
|
||||
|
||||
@@ -80,14 +80,14 @@ To be used with a particular `FusedMoEPrepareAndFinalizeModular` subclass, MoE k
|
||||
|
||||
| Kernel | Input act. format | Quant. types | Quant. format | Activation function | Apply Weight On Input | Modular | Source |
|
||||
| ------ | ----------------- | ------------ | ------------- | ------------------- | --------------------- | ------- | ------ |
|
||||
| triton | standard | all<sup>1</sup> | G,A,T | silu, gelu,</br>swigluoai,</br>silu_no_mul,</br>gelu_no_mul | Y | Y | [`fused_experts`][vllm.model_executor.layers.fused_moe.fused_moe.fused_experts],</br>[`TritonExperts`][vllm.model_executor.layers.fused_moe.fused_moe.TritonExperts] |
|
||||
| triton | standard | all<sup>1</sup> | G,A,T | silu, gelu,</br>swigluoai,</br>silu_no_mul,</br>gelu_no_mul | Y | Y | [`fused_experts`][vllm.model_executor.layers.fused_moe.fused_moe.fused_experts],</br>[`TritonExperts`][vllm.model_executor.layers.fused_moe.experts.triton_moe.TritonExperts] |
|
||||
| triton (batched) | batched | all<sup>1</sup> | G,A,T | silu, gelu | <sup>6</sup> | Y | [`BatchedTritonExperts`][vllm.model_executor.layers.fused_moe.fused_batched_moe.BatchedTritonExperts] |
|
||||
| deep gemm | standard,</br>batched | fp8 | G(128),A,T | silu, gelu | <sup>6</sup> | Y | </br>[`DeepGemmExperts`][vllm.model_executor.layers.fused_moe.experts.deep_gemm_moe.DeepGemmExperts],</br>[`BatchedDeepGemmExperts`][vllm.model_executor.layers.fused_moe.experts.batched_deep_gemm_moe.BatchedDeepGemmExperts] |
|
||||
| cutlass_fp4 | standard,</br>batched | nvfp4 | A,T | silu | Y | Y | [`CutlassExpertsFp4`][vllm.model_executor.layers.fused_moe.experts.cutlass_moe.CutlassExpertsFp4] |
|
||||
| cutlass_fp8 | standard,</br>batched | fp8 | A,T | silu, gelu | Y | Y | [`CutlassExpertsFp8`][vllm.model_executor.layers.fused_moe.experts.cutlass_moe.CutlassExpertsFp8],</br>[`CutlasBatchedExpertsFp8`][vllm.model_executor.layers.fused_moe.experts.cutlass_moe.CutlassBatchedExpertsFp8] |
|
||||
| flashinfer | standard | nvfp4,</br>fp8 | T | <sup>5</sup> | N | Y | [`FlashInferExperts`][vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe.FlashInferExperts] |
|
||||
| flashinfer | standard | nvfp4,</br>fp8 | T | <sup>5</sup> | N | Y | [`FlashInferExperts`][vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe.FlashInferExperts] |
|
||||
| gpt oss triton | standard | N/A | N/A | <sup>5</sup> | Y | Y | [`triton_kernel_fused_experts`][vllm.model_executor.layers.fused_moe.experts.gpt_oss_triton_kernels_moe.triton_kernel_fused_experts],</br>[`OAITritonExperts`][vllm.model_executor.layers.fused_moe.experts.gpt_oss_triton_kernels_moe.OAITritonExperts] |
|
||||
| marlin | standard,</br>batched | <sup>3</sup> / N/A | <sup>3</sup> / N/A | silu,</br>swigluoai | Y | Y | [`fused_marlin_moe`][vllm.model_executor.layers.fused_moe.fused_marlin_moe.fused_marlin_moe],</br>[`MarlinExperts`][vllm.model_executor.layers.fused_moe.fused_marlin_moe.MarlinExperts],</br>[`BatchedMarlinExperts`][vllm.model_executor.layers.fused_moe.fused_marlin_moe.BatchedMarlinExperts] |
|
||||
| marlin | standard,</br>batched | <sup>3</sup> / N/A | <sup>3</sup> / N/A | silu,</br>swigluoai | Y | Y | [`fused_marlin_moe`][vllm.model_executor.layers.fused_moe.experts.marlin_moe.fused_marlin_moe],</br>[`MarlinExperts`][vllm.model_executor.layers.fused_moe.experts.marlin_moe.MarlinExperts],</br>[`BatchedMarlinExperts`][vllm.model_executor.layers.fused_moe.experts.marlin_moe.BatchedMarlinExperts] |
|
||||
| trtllm | standard | mxfp4,</br>nvfp4 | G(16),G(32) | <sup>5</sup> | N | Y | [`TrtLlmMxfp4ExpertsMonolithic`][vllm.model_executor.layers.fused_moe.experts.trtllm_mxfp4_moe.TrtLlmMxfp4ExpertsMonolithic],</br>[`TrtLlmMxfp4ExpertsModular`][vllm.model_executor.layers.fused_moe.experts.trtllm_mxfp4_moe.TrtLlmMxfp4ExpertsModular],</br>[`TrtLlmNvFp4ExpertsMonolithic`][vllm.model_executor.layers.fused_moe.experts.trtllm_nvfp4_moe.TrtLlmNvFp4ExpertsMonolithic],</br>[`TrtLlmNvfp4ExpertsModular`][vllm.model_executor.layers.fused_moe.experts.trtllm_nvfp4_moe.TrtLlmNvFp4ExpertsModular] |
|
||||
| rocm aiter moe | standard | mxfp4,</br>fp8 | G(32),G(128),A,T | silu, gelu,</br>swigluoai | Y | N | `rocm_aiter_fused_experts`,</br>`AiterExperts` |
|
||||
| cpu_fused_moe | standard | N/A | N/A | silu | N | N | [`CPUFusedMOE`][vllm.model_executor.layers.fused_moe.cpu_fused_moe.CPUFusedMOE] |
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
# MooncakeStoreConnector Usage Guide
|
||||
|
||||
MooncakeStoreConnector is a KV cache connector that uses [MooncakeDistributedStore](https://github.com/kvcache-ai/Mooncake) as a shared KV cache pool. Unlike `MooncakeConnector` which does direct point-to-point KV transfer between prefiller and decoder, MooncakeStoreConnector enables KV cache offloading to an external distributed store, supporting:
|
||||
|
||||
- **CPU offloading**: Extend effective KV cache capacity by offloading to CPU memory via Mooncake's transfer engine.
|
||||
- **Prefix caching across instances**: Hash-based deduplication allows multiple vLLM instances to share cached KV blocks through the store.
|
||||
- **Single-node and multi-node deployment**: Works both as a standalone KV cache extension and in disaggregated prefill-decode setups.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
### Install Mooncake
|
||||
|
||||
Install mooncake through pip:
|
||||
|
||||
```bash
|
||||
uv pip install mooncake-transfer-engine
|
||||
```
|
||||
|
||||
Refer to the [Mooncake official repository](https://github.com/kvcache-ai/Mooncake) for more installation instructions and building from source.
|
||||
|
||||
### Start the Mooncake Master Server
|
||||
|
||||
The Mooncake master manages metadata and coordinates the distributed store. Start it before launching vLLM:
|
||||
|
||||
```bash
|
||||
mooncake_master --port 50051
|
||||
```
|
||||
|
||||
Default ports:
|
||||
|
||||
- RPC: 50051
|
||||
|
||||
Multiple vLLM instances can share the same master server.
|
||||
|
||||
### Configure Mooncake
|
||||
|
||||
Create a JSON configuration file (e.g., `mooncake_config.json`):
|
||||
|
||||
```json
|
||||
{
|
||||
"metadata_server": "P2PHANDSHAKE",
|
||||
"master_server_address": "127.0.0.1:50051",
|
||||
"global_segment_size": "80GB",
|
||||
"local_buffer_size": "4GB",
|
||||
"protocol": "rdma",
|
||||
"device_name": ""
|
||||
}
|
||||
```
|
||||
|
||||
- `protocol`: Use `"rdma"` for best performance. `"tcp"` works as a fallback.
|
||||
- `global_segment_size`: CPU memory contributed to the distributed pool (per GPU).
|
||||
- `local_buffer_size`: Private buffer for this node's own operations (per GPU).
|
||||
|
||||
Set the config path via environment variable:
|
||||
|
||||
```bash
|
||||
export MOONCAKE_CONFIG_PATH=/path/to/mooncake_config.json
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
### Single-Node KV Cache Offloading
|
||||
|
||||
Use MooncakeStoreConnector to offload KV cache to CPU memory, extending the effective cache size:
|
||||
|
||||
```bash
|
||||
MOONCAKE_CONFIG_PATH=mooncake_config.json \
|
||||
vllm serve meta-llama/Llama-3.1-8B-Instruct \
|
||||
--kv-transfer-config '{"kv_connector":"MooncakeStoreConnector","kv_role":"kv_both"}'
|
||||
```
|
||||
|
||||
### Disaggregated Prefill-Decode (XpYd)
|
||||
|
||||
In disaggregated prefill-decode mode, use `MultiConnector` to combine `MooncakeConnector` (point-to-point KV transfer) with `MooncakeStoreConnector` (shared KV cache pool). This enables both direct P2P transfer between prefiller and decoder, and cross-instance prefix cache sharing via the distributed store.
|
||||
**Prefiller Node:**
|
||||
|
||||
```bash
|
||||
MOONCAKE_CONFIG_PATH=mooncake_config.json \
|
||||
VLLM_MOONCAKE_BOOTSTRAP_PORT=50052 \
|
||||
vllm serve meta-llama/Llama-3.1-8B-Instruct \
|
||||
--port 8100 \
|
||||
--kv-transfer-config '{
|
||||
"kv_connector": "MultiConnector",
|
||||
"kv_role": "kv_producer",
|
||||
"kv_connector_extra_config": {
|
||||
"connectors": [
|
||||
{
|
||||
"kv_connector": "MooncakeConnector",
|
||||
"kv_role": "kv_producer"
|
||||
},
|
||||
{
|
||||
"kv_connector": "MooncakeStoreConnector",
|
||||
"kv_role": "kv_producer"
|
||||
}
|
||||
]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**Decoder Node:**
|
||||
|
||||
```bash
|
||||
MOONCAKE_CONFIG_PATH=mooncake_config.json \
|
||||
VLLM_MOONCAKE_BOOTSTRAP_PORT=50053 \
|
||||
vllm serve meta-llama/Llama-3.1-8B-Instruct \
|
||||
--port 8200 \
|
||||
--kv-transfer-config '{
|
||||
"kv_connector": "MultiConnector",
|
||||
"kv_role": "kv_consumer",
|
||||
"kv_connector_extra_config": {
|
||||
"connectors": [
|
||||
{
|
||||
"kv_connector": "MooncakeConnector",
|
||||
"kv_role": "kv_consumer"
|
||||
},
|
||||
{
|
||||
"kv_connector": "MooncakeStoreConnector",
|
||||
"kv_role": "kv_consumer"
|
||||
}
|
||||
]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**Proxy:**
|
||||
|
||||
A disaggregation proxy is required to route requests between prefiller and decoder nodes. The proxy assigns `do_remote_prefill=True` / `do_remote_decode=True` to coordinate P2P transfer via `MooncakeConnector`. Refer to the [MooncakeConnector usage guide](mooncake_connector_usage.md) for proxy setup details.
|
||||
|
||||
## Environment Variables
|
||||
|
||||
| Variable | Description | Default |
|
||||
| --- | --- | --- |
|
||||
| `MOONCAKE_CONFIG_PATH` | Path to Mooncake JSON config file | (required) |
|
||||
| `VLLM_MOONCAKE_BOOTSTRAP_PORT` | Bootstrap port for MooncakeConnector P2P transfer (disagg mode only) | 8998 |
|
||||
|
||||
## KV Transfer Config
|
||||
|
||||
### KV Role Options
|
||||
|
||||
- **kv_producer**: For prefiller instances that store KV caches to the pool.
|
||||
- **kv_consumer**: For decoder instances that load KV caches from the pool.
|
||||
- **kv_both**: The instance both stores and loads KV caches. Use this for single-node CPU offloading.
|
||||
|
||||
### kv_connector_extra_config
|
||||
|
||||
- `load_async` (bool): Enable asynchronous loading for better compute-I/O overlap. Default: `true`.
|
||||
- `enable_cross_layers_blocks` (bool): Enable cross-layer block packing for reduced store operations. Default: `false`.
|
||||
- `discard_partial_chunks` (bool): Discard partial block chunks during store. Default: `true`.
|
||||
- `lookup_rpc_port` (int): Custom port for the ZMQ lookup RPC socket. Default: `0`.
|
||||
|
||||
## Notes
|
||||
|
||||
### Cross-DP Prefix Cache Hits
|
||||
|
||||
When running with data parallelism, set a fixed `PYTHONHASHSEED` so that block hashes are consistent across DP ranks:
|
||||
|
||||
```bash
|
||||
PYTHONHASHSEED=0 vllm serve ...
|
||||
```
|
||||
|
||||
Without this, identical prompts may produce different block hashes on different DP ranks, preventing cross-instance prefix cache hits.
|
||||
@@ -385,7 +385,7 @@ th {
|
||||
| `DeepseekForCausalLM` | DeepSeek | `deepseek-ai/deepseek-llm-67b-base`, `deepseek-ai/deepseek-llm-7b-chat`, etc. | ✅︎ | ✅︎ |
|
||||
| `DeepseekV2ForCausalLM` | DeepSeek-V2 | `deepseek-ai/DeepSeek-V2`, `deepseek-ai/DeepSeek-V2-Chat`, etc. | ✅︎ | ✅︎ |
|
||||
| `DeepseekV3ForCausalLM` | DeepSeek-V3 | `deepseek-ai/DeepSeek-V3`, `deepseek-ai/DeepSeek-R1`, `deepseek-ai/DeepSeek-V3.1`, etc. | ✅︎ | ✅︎ |
|
||||
| `DeepseekV4ForCausalLM` | DeepSeek-V4 | `deepseek-ai/DeepSeek-V4-Flash`, `deepseek-ai/DeepSeek-V4-Pro`, etc. | | |
|
||||
| `DeepseekV4ForCausalLM` | DeepSeek-V4 | `deepseek-ai/DeepSeek-V4-Flash`, `deepseek-ai/DeepSeek-V4-Pro`, etc. | | ✅︎ |
|
||||
| `Dots1ForCausalLM` | dots.llm1 | `rednote-hilab/dots.llm1.base`, `rednote-hilab/dots.llm1.inst`, etc. | | ✅︎ |
|
||||
| `DotsOCRForCausalLM` | dots_ocr | `rednote-hilab/dots.ocr` | ✅︎ | ✅︎ |
|
||||
| `Ernie4_5ForCausalLM` | Ernie4.5 | `baidu/ERNIE-4.5-0.3B-PT`, etc. | ✅︎ | ✅︎ |
|
||||
|
||||
@@ -263,7 +263,7 @@
|
||||
{%- if message.get('tool_responses') -%}
|
||||
{#- Legacy: tool_responses embedded on the assistant message (Google/Gemma native) -#}
|
||||
{%- for tool_response in message['tool_responses'] -%}
|
||||
{{- format_tool_response_block(tool_response['name'] | default('unknown'), tool_response['response']) -}}
|
||||
{{- format_tool_response_block(tool_response['name'] | default('unknown', true), tool_response['response']) -}}
|
||||
{%- set ns_tr_out.flag = true -%}
|
||||
{%- set ns.prev_message_type = 'tool_response' -%}
|
||||
{%- endfor -%}
|
||||
@@ -277,7 +277,7 @@
|
||||
{%- else -%}
|
||||
{%- set follow = loop_messages[k] -%}
|
||||
{#- Resolve tool_call_id to function name -#}
|
||||
{%- set ns_tname = namespace(name=follow.get('name') | default('unknown')) -%}
|
||||
{%- set ns_tname = namespace(name=follow.get('name') | default('unknown', true)) -%}
|
||||
{%- for tc in message['tool_calls'] -%}
|
||||
{%- if tc.get('id') == follow.get('tool_call_id') -%}
|
||||
{%- set ns_tname.name = tc['function']['name'] -%}
|
||||
|
||||
@@ -21,5 +21,5 @@ nvidia-cudnn-frontend>=1.13.0,<1.19.0
|
||||
fastsafetensors >= 0.2.2
|
||||
|
||||
# QuACK and Cutlass DSL for FA4 (cute-DSL implementation)
|
||||
nvidia-cutlass-dsl>=4.4.2
|
||||
nvidia-cutlass-dsl[cu13]>=4.4.2
|
||||
quack-kernels>=0.3.3
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
lmcache >= 0.3.9
|
||||
nixl[cu13] >= 0.7.1, <= 0.10.1 # Required for disaggregated prefill
|
||||
nixl-cu12 >= 0.7.1, <= 0.10.1
|
||||
nixl-cu13 >= 0.7.1, <= 0.10.1
|
||||
nixl >= 1.1.0 # Required for disaggregated prefill
|
||||
mooncake-transfer-engine >= 0.3.8
|
||||
|
||||
@@ -969,6 +969,9 @@ def get_requirements() -> list[str]:
|
||||
# vllm-flash-attn is built only for CUDA 12.x.
|
||||
# Skip for other versions.
|
||||
continue
|
||||
if "nvidia-cutlass-dsl[cu13]" in req and cuda_major == "12":
|
||||
# [cu13] extra is the default; strip it on CUDA 12 builds.
|
||||
req = req.replace("nvidia-cutlass-dsl[cu13]", "nvidia-cutlass-dsl")
|
||||
modified_requirements.append(req)
|
||||
requirements = modified_requirements
|
||||
elif _is_hip():
|
||||
|
||||
@@ -13,6 +13,7 @@ from vllm.model_executor.layers.mamba.ops.ssu_dispatch import (
|
||||
selective_state_update,
|
||||
)
|
||||
from vllm.utils.torch_utils import set_random_seed
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
from vllm.v1.kv_cache_interface import (
|
||||
KVCacheConfig,
|
||||
KVCacheGroupSpec,
|
||||
@@ -27,7 +28,9 @@ except ImportError:
|
||||
HAS_FLASHINFER = False
|
||||
|
||||
|
||||
def _kv_cache_config_with_ssu(mamba_type: str = "mamba2") -> KVCacheConfig:
|
||||
def _kv_cache_config_with_ssu(
|
||||
mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2,
|
||||
) -> KVCacheConfig:
|
||||
spec = MambaSpec(
|
||||
block_size=16,
|
||||
shapes=((16, 64),),
|
||||
@@ -77,7 +80,12 @@ def test_uninitialized_backend_raises():
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"mamba_type", ["linear_attention", "gdn_attention", "short_conv"]
|
||||
"mamba_type",
|
||||
[
|
||||
MambaAttentionBackendEnum.LINEAR,
|
||||
MambaAttentionBackendEnum.GDN_ATTN,
|
||||
MambaAttentionBackendEnum.SHORT_CONV,
|
||||
],
|
||||
)
|
||||
def test_init_is_noop_for_non_ssu_mamba_type(mamba_type):
|
||||
import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod
|
||||
|
||||
@@ -237,7 +237,7 @@ if has_mori():
|
||||
)
|
||||
|
||||
if has_flashinfer_cutlass_fused_moe() and current_platform.has_device_capability(100):
|
||||
from vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import (
|
||||
FlashInferExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.prepare_finalize.flashinfer_nvlink_two_sided import ( # noqa: E501
|
||||
@@ -298,7 +298,7 @@ if has_flashinfer_cutlass_fused_moe() and current_platform.has_device_capability
|
||||
)
|
||||
|
||||
if has_aiter():
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
AiterExperts,
|
||||
)
|
||||
|
||||
|
||||
@@ -18,12 +18,12 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
RoutingMethodType,
|
||||
fp8_w8a8_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import (
|
||||
FlashInferExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.trtllm_fp8_moe import (
|
||||
TrtLlmFp8ExpertsMonolithic,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe import (
|
||||
FlashInferExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import fused_experts
|
||||
from vllm.model_executor.layers.quantization.utils.flashinfer_utils import (
|
||||
rotate_weights_for_fi_trtllm_fp8_per_tensor_moe,
|
||||
|
||||
@@ -22,7 +22,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEParallelConfig,
|
||||
RoutingMethodType,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import (
|
||||
FlashInferExperts,
|
||||
is_valid_flashinfer_cutlass_fused_moe,
|
||||
)
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
fused_marlin_moe,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.router.grouped_topk_router import (
|
||||
|
||||
@@ -32,7 +32,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
int4_w4a16_moe_quant_config,
|
||||
int8_w8a16_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
batched_fused_marlin_moe,
|
||||
fused_marlin_moe,
|
||||
)
|
||||
|
||||
@@ -20,7 +20,7 @@ if not current_platform.is_rocm():
|
||||
pytest.skip("This test can only run on ROCm.", allow_module_level=True)
|
||||
|
||||
# this import statement is needed to ensure the ops are registered
|
||||
import vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe # noqa: F401
|
||||
import vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe # noqa: F401
|
||||
|
||||
# need to import once to ensure the ops are registered
|
||||
# Check if aiter package is installed
|
||||
|
||||
@@ -15,7 +15,7 @@ from vllm.model_executor.layers.fused_moe.activation import MoEActivation
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FUSED_MOE_UNQUANTIZED_CONFIG,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
# Test parameters
|
||||
@@ -151,7 +151,7 @@ def test_triton_experts_no_mul_activation(
|
||||
@torch.inference_mode()
|
||||
def test_workspace_shapes_no_mul_vs_gated():
|
||||
"""Test that workspace shapes differ correctly between gated and non-gated."""
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
|
||||
|
||||
M, N, K, topk = 64, 256, 128, 2
|
||||
|
||||
@@ -192,7 +192,7 @@ def test_workspace_shapes_no_mul_vs_gated():
|
||||
@torch.inference_mode()
|
||||
def test_adjust_n_for_activation():
|
||||
"""Test the adjust_N_for_activation method."""
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
|
||||
|
||||
experts = TritonExperts(
|
||||
moe_config=make_dummy_moe_config(),
|
||||
|
||||
@@ -158,7 +158,7 @@ def test_select_cuda_flashinfer_trtllm_backend(mock_is_supported_trtllm, monkeyp
|
||||
return_value=(False, None),
|
||||
)
|
||||
@patch(
|
||||
"vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe.FlashInferExperts.is_supported_config",
|
||||
"vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe.FlashInferExperts.is_supported_config",
|
||||
return_value=(True, None),
|
||||
)
|
||||
@pytest.mark.skipif(
|
||||
|
||||
@@ -17,12 +17,14 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEQuantConfig,
|
||||
RoutingMethodType,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import (
|
||||
TritonExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_batched_moe import (
|
||||
BatchedTritonExperts,
|
||||
NaiveBatchedExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import (
|
||||
TritonExperts,
|
||||
fused_experts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.modular_kernel import FusedMoEKernel
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import vllm.model_executor.layers.mhc as mhc_ops # noqa: F401
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.torch_utils import set_random_seed
|
||||
|
||||
DEVICE = current_platform.device_type
|
||||
|
||||
|
||||
def sinkhorn_normalize_ref(x: torch.Tensor, repeat: int, eps: float) -> torch.Tensor:
|
||||
x = x.softmax(-1) + eps
|
||||
x = x / (x.sum(-2, keepdim=True) + eps)
|
||||
for _ in range(repeat - 1):
|
||||
x = x / (x.sum(-1, keepdim=True) + eps)
|
||||
x = x / (x.sum(-2, keepdim=True) + eps)
|
||||
return x
|
||||
|
||||
|
||||
def mhc_pre_ref(
|
||||
residual: torch.Tensor,
|
||||
fn: torch.Tensor,
|
||||
hc_scale: torch.Tensor,
|
||||
hc_base: torch.Tensor,
|
||||
rms_eps: float,
|
||||
hc_pre_eps: float,
|
||||
hc_sinkhorn_eps: float,
|
||||
hc_post_mult_value: float,
|
||||
sinkhorn_repeat: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""mHC pre reference kernel from tilelang repo: https://github.com/tile-ai/tilelang/blob/d135bd1cd2d2eee74fbb41dd0a0831a427194c86/examples/deepseek_mhc/example_mhc_pre.py#L303"""
|
||||
hc_mult = residual.shape[-2]
|
||||
|
||||
residual_flat = residual.flatten(-2, -1).float()
|
||||
sqrsum = residual_flat.square().sum(-1)
|
||||
mixes = (
|
||||
residual_flat @ fn.T * (sqrsum.unsqueeze(-1) / fn.shape[-1] + rms_eps).rsqrt()
|
||||
)
|
||||
|
||||
hc_scale = torch.cat(
|
||||
[
|
||||
hc_scale[0].expand(hc_mult),
|
||||
hc_scale[1].expand(hc_mult),
|
||||
hc_scale[2].expand(hc_mult * hc_mult),
|
||||
],
|
||||
)
|
||||
mixes = mixes * hc_scale + hc_base
|
||||
|
||||
pre_mix = mixes[:, :hc_mult].sigmoid().unsqueeze(-1) + hc_pre_eps
|
||||
post_mix = (
|
||||
mixes[:, hc_mult : 2 * hc_mult].sigmoid() * hc_post_mult_value
|
||||
).unsqueeze(-1)
|
||||
res_mix = mixes[:, 2 * hc_mult :].view(-1, hc_mult, hc_mult)
|
||||
|
||||
res_mix = sinkhorn_normalize_ref(
|
||||
res_mix, repeat=sinkhorn_repeat, eps=hc_sinkhorn_eps
|
||||
)
|
||||
|
||||
layer_input = (residual * pre_mix).sum(-2).bfloat16()
|
||||
|
||||
return post_mix, res_mix, layer_input
|
||||
|
||||
|
||||
def mhc_post_ref(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
post_layer_mix: torch.Tensor,
|
||||
comb_res_mix: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""mHC post reference kernel from tilelang repo: https://github.com/tile-ai/tilelang/blob/d135bd1cd2d2eee74fbb41dd0a0831a427194c86/examples/deepseek_mhc/example_mhc_post.py#L68"""
|
||||
term2 = torch.bmm(comb_res_mix.mT, residual.float())
|
||||
return (x.float().unsqueeze(-2) * post_layer_mix + term2).bfloat16()
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not current_platform.is_cuda(),
|
||||
reason="CUDA required",
|
||||
)
|
||||
@pytest.mark.parametrize("num_tokens", [1, 4, 8, 128])
|
||||
@pytest.mark.parametrize("hidden_size", [4096, 7168])
|
||||
@pytest.mark.parametrize("hc_mult", [4])
|
||||
def test_mhc_fused_post_pre(num_tokens, hidden_size, hc_mult):
|
||||
torch.set_default_device(DEVICE)
|
||||
set_random_seed(0)
|
||||
|
||||
x = torch.randn((num_tokens, hidden_size), dtype=torch.bfloat16)
|
||||
residual = torch.randn((num_tokens, hc_mult, hidden_size), dtype=torch.bfloat16)
|
||||
post_layer_mix = torch.randn((num_tokens, hc_mult, 1), dtype=torch.float32)
|
||||
comb_res_mix = torch.randn((num_tokens, hc_mult, hc_mult), dtype=torch.float32)
|
||||
|
||||
hc_mult2 = hc_mult * hc_mult
|
||||
hc_mult3 = hc_mult * 2 + hc_mult2
|
||||
fn = (
|
||||
torch.randn((hc_mult3, hc_mult, hidden_size), dtype=torch.float)
|
||||
* 1e-4
|
||||
* (1 + torch.arange(hc_mult).mul(0.01).view(1, -1, 1))
|
||||
).flatten(1, 2)
|
||||
hc_scale = torch.randn((3,), dtype=torch.float) * 0.1
|
||||
hc_base = torch.randn((hc_mult3,), dtype=torch.float) * 0.1
|
||||
|
||||
hc_sinkhorn_eps = hc_pre_eps = rms_eps = 1e-6
|
||||
sinkhorn_repeat = 20
|
||||
hc_post_alpha = 1.0
|
||||
|
||||
def run_ref():
|
||||
residual_ref = mhc_post_ref(x, residual, post_layer_mix, comb_res_mix)
|
||||
post_mix_ref, res_mix_ref, layer_input_ref = mhc_pre_ref(
|
||||
residual_ref,
|
||||
fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
rms_eps,
|
||||
hc_pre_eps,
|
||||
hc_sinkhorn_eps,
|
||||
hc_post_alpha,
|
||||
sinkhorn_repeat,
|
||||
)
|
||||
return residual_ref, post_mix_ref, res_mix_ref, layer_input_ref
|
||||
|
||||
residual_ref, post_mix_ref, res_mix_ref, layer_input_ref = run_ref()
|
||||
|
||||
residual, post_mix, res_mix, x = torch.ops.vllm.mhc_fused_post_pre(
|
||||
x,
|
||||
residual,
|
||||
post_layer_mix,
|
||||
comb_res_mix,
|
||||
fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
rms_eps,
|
||||
hc_pre_eps,
|
||||
hc_sinkhorn_eps,
|
||||
hc_post_alpha,
|
||||
sinkhorn_repeat,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(residual, residual_ref, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(post_mix, post_mix_ref, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(res_mix, res_mix_ref, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(x, layer_input_ref, atol=1e-2, rtol=1e-2)
|
||||
@@ -0,0 +1,56 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.model_executor.models.molmo2 import build_flat_image_bool_length
|
||||
|
||||
|
||||
def test_build_flat_image_bool_length_matches_molmoweb_processor_tokens():
|
||||
hf_config = SimpleNamespace(
|
||||
image_patch_id=151938,
|
||||
low_res_image_start_token_id=151940,
|
||||
image_start_token_id=151936,
|
||||
image_col_id=151939,
|
||||
image_end_token_id=151937,
|
||||
)
|
||||
image_grids = torch.tensor([[14, 14, 14, 23]], dtype=torch.long)
|
||||
|
||||
image_tokens, num_image_tokens = build_flat_image_bool_length(
|
||||
image_grids,
|
||||
hf_config,
|
||||
image_use_col_tokens=True,
|
||||
use_single_crop_col_tokens=None,
|
||||
use_single_crop_start_token=False,
|
||||
)
|
||||
|
||||
assert num_image_tokens.tolist() == [550]
|
||||
assert len(image_tokens) == 550
|
||||
assert image_tokens[0].item() == hf_config.image_start_token_id
|
||||
assert (image_tokens == hf_config.image_col_id).sum().item() == 28
|
||||
|
||||
|
||||
def test_build_flat_image_bool_length_respects_disabled_col_tokens():
|
||||
hf_config = SimpleNamespace(
|
||||
image_patch_id=151938,
|
||||
low_res_image_start_token_id=151940,
|
||||
image_start_token_id=151936,
|
||||
image_col_id=151939,
|
||||
image_end_token_id=151937,
|
||||
)
|
||||
image_grids = torch.tensor([[2, 3, 5, 7]], dtype=torch.long)
|
||||
|
||||
image_tokens, num_image_tokens = build_flat_image_bool_length(
|
||||
image_grids,
|
||||
hf_config,
|
||||
image_use_col_tokens=False,
|
||||
use_single_crop_col_tokens=False,
|
||||
use_single_crop_start_token=True,
|
||||
)
|
||||
|
||||
assert num_image_tokens.tolist() == [45]
|
||||
assert len(image_tokens) == 45
|
||||
assert image_tokens[0].item() == hf_config.low_res_image_start_token_id
|
||||
assert (image_tokens == hf_config.image_col_id).sum().item() == 0
|
||||
@@ -13,6 +13,7 @@ from vllm.model_executor.models.minimax_text_01 import MiniMaxText01LinearAttent
|
||||
from vllm.v1.attention.backends.linear_attn import LinearAttentionBackend
|
||||
from vllm.v1.attention.backends.mamba1_attn import Mamba1AttentionBackend
|
||||
from vllm.v1.attention.backends.mamba2_attn import Mamba2AttentionBackend
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
from vllm.v1.attention.backends.short_conv_attn import ShortConvAttentionBackend
|
||||
|
||||
|
||||
@@ -32,7 +33,7 @@ from vllm.v1.attention.backends.short_conv_attn import ShortConvAttentionBackend
|
||||
use_rms_norm=True,
|
||||
),
|
||||
Mamba1AttentionBackend,
|
||||
"mamba1",
|
||||
MambaAttentionBackendEnum.MAMBA1,
|
||||
),
|
||||
(
|
||||
MambaMixer2,
|
||||
@@ -48,7 +49,7 @@ from vllm.v1.attention.backends.short_conv_attn import ShortConvAttentionBackend
|
||||
head_dim=32,
|
||||
),
|
||||
Mamba2AttentionBackend,
|
||||
"mamba2",
|
||||
MambaAttentionBackendEnum.MAMBA2,
|
||||
),
|
||||
(
|
||||
MiniMaxText01LinearAttention,
|
||||
@@ -64,7 +65,7 @@ from vllm.v1.attention.backends.short_conv_attn import ShortConvAttentionBackend
|
||||
linear_layer_idx=0,
|
||||
),
|
||||
LinearAttentionBackend,
|
||||
"linear_attention",
|
||||
MambaAttentionBackendEnum.LINEAR,
|
||||
),
|
||||
(
|
||||
ShortConv,
|
||||
@@ -74,7 +75,7 @@ from vllm.v1.attention.backends.short_conv_attn import ShortConvAttentionBackend
|
||||
layer_idx=0,
|
||||
),
|
||||
ShortConvAttentionBackend,
|
||||
"short_conv",
|
||||
MambaAttentionBackendEnum.SHORT_CONV,
|
||||
),
|
||||
],
|
||||
)
|
||||
@@ -97,10 +98,14 @@ def test_mamba_layers_get_attn_backend(
|
||||
@pytest.mark.parametrize(
|
||||
"layer_class,expected_backend,expected_mamba_type",
|
||||
[
|
||||
(MambaMixer, Mamba1AttentionBackend, "mamba1"),
|
||||
(MambaMixer2, Mamba2AttentionBackend, "mamba2"),
|
||||
(MiniMaxText01LinearAttention, LinearAttentionBackend, "linear_attention"),
|
||||
(ShortConv, ShortConvAttentionBackend, "short_conv"),
|
||||
(MambaMixer, Mamba1AttentionBackend, MambaAttentionBackendEnum.MAMBA1),
|
||||
(MambaMixer2, Mamba2AttentionBackend, MambaAttentionBackendEnum.MAMBA2),
|
||||
(
|
||||
MiniMaxText01LinearAttention,
|
||||
LinearAttentionBackend,
|
||||
MambaAttentionBackendEnum.LINEAR,
|
||||
),
|
||||
(ShortConv, ShortConvAttentionBackend, MambaAttentionBackendEnum.SHORT_CONV),
|
||||
],
|
||||
)
|
||||
def test_mamba_layers_have_unified_interface(
|
||||
|
||||
@@ -0,0 +1,258 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from vllm.config import set_current_vllm_config
|
||||
from vllm.distributed.kv_events import BlockStored
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
|
||||
KVConnectorRole,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import (
|
||||
connector,
|
||||
worker,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501
|
||||
MooncakeStoreConnectorMetadata,
|
||||
)
|
||||
from vllm.v1.outputs import KVConnectorOutput
|
||||
|
||||
from .utils import create_vllm_config
|
||||
|
||||
|
||||
def _make_vllm_config():
|
||||
return create_vllm_config(
|
||||
kv_connector="MooncakeStoreConnector",
|
||||
kv_role="kv_both",
|
||||
)
|
||||
|
||||
|
||||
def _make_block_stored() -> BlockStored:
|
||||
return BlockStored(
|
||||
block_hashes=[b"hash"],
|
||||
parent_block_hash=None,
|
||||
token_ids=[1, 2, 3],
|
||||
block_size=16,
|
||||
lora_id=None,
|
||||
medium="cpu",
|
||||
lora_name=None,
|
||||
)
|
||||
|
||||
|
||||
def test_scheduler_role_initializes_store_scheduler_only():
|
||||
vllm_config = _make_vllm_config()
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreScheduler"
|
||||
) as mock_scheduler,
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreWorker"
|
||||
) as mock_worker,
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.SCHEDULER)
|
||||
|
||||
mock_scheduler.assert_called_once_with(vllm_config)
|
||||
mock_worker.assert_not_called()
|
||||
assert conn.connector_scheduler is mock_scheduler.return_value
|
||||
assert conn.connector_worker is None
|
||||
|
||||
|
||||
def test_worker_role_initializes_store_worker_on_rank0():
|
||||
vllm_config = _make_vllm_config()
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreScheduler"
|
||||
) as mock_scheduler,
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreWorker"
|
||||
) as mock_worker,
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
|
||||
mock_scheduler.assert_not_called()
|
||||
mock_worker.assert_called_once_with(vllm_config)
|
||||
assert conn.connector_scheduler is None
|
||||
assert conn.connector_worker is mock_worker.return_value
|
||||
|
||||
|
||||
def test_worker_role_initializes_on_nonzero_rank():
|
||||
vllm_config = _make_vllm_config()
|
||||
vllm_config.parallel_config.rank = 1
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreWorker"
|
||||
) as mock_worker,
|
||||
):
|
||||
connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
|
||||
mock_worker.assert_called_once_with(vllm_config)
|
||||
|
||||
|
||||
def test_lookup_rpc_path_uses_data_parallel_index_in_dense_dp():
|
||||
vllm_config = _make_vllm_config()
|
||||
vllm_config.parallel_config.data_parallel_rank = 0
|
||||
vllm_config.parallel_config.data_parallel_index = 3
|
||||
|
||||
path = worker.get_zmq_rpc_path_lookup(vllm_config)
|
||||
|
||||
assert path.endswith("_dp_rank3")
|
||||
|
||||
|
||||
def test_lookup_rpc_path_uses_local_rank_when_local_engines_only():
|
||||
vllm_config = _make_vllm_config()
|
||||
vllm_config.parallel_config.data_parallel_index = 7
|
||||
vllm_config.parallel_config.data_parallel_rank_local = 1
|
||||
vllm_config.parallel_config.data_parallel_hybrid_lb = True
|
||||
|
||||
path = worker.get_zmq_rpc_path_lookup(vllm_config)
|
||||
|
||||
assert path.endswith("_dp_rank1")
|
||||
|
||||
|
||||
def test_worker_methods_delegate_to_store_worker():
|
||||
vllm_config = _make_vllm_config()
|
||||
kv_caches = {"layer0": MagicMock()}
|
||||
metadata = MooncakeStoreConnectorMetadata(set(), set())
|
||||
finished_req_ids = {"req-1"}
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreWorker"
|
||||
) as mock_worker_cls,
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
|
||||
worker_inst = mock_worker_cls.return_value
|
||||
worker_inst.get_finished.return_value = ({"req-1"}, {"req-2"})
|
||||
conn.bind_connector_metadata(metadata)
|
||||
|
||||
conn.register_kv_caches(kv_caches)
|
||||
result = conn.get_finished(finished_req_ids)
|
||||
|
||||
worker_inst.register_kv_caches.assert_called_once_with(kv_caches)
|
||||
worker_inst.get_finished.assert_called_once_with(finished_req_ids, metadata)
|
||||
assert result == ({"req-1"}, {"req-2"})
|
||||
|
||||
|
||||
def test_get_kv_connector_kv_cache_events_returns_none_when_empty():
|
||||
vllm_config = _make_vllm_config()
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreWorker"
|
||||
) as mock_worker_cls,
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
|
||||
mock_worker_cls.return_value.get_kv_events.return_value = []
|
||||
assert conn.get_kv_connector_kv_cache_events() is None
|
||||
|
||||
|
||||
def test_get_kv_connector_kv_cache_events_wraps_worker_events():
|
||||
vllm_config = _make_vllm_config()
|
||||
event = _make_block_stored()
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreWorker"
|
||||
) as mock_worker_cls,
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
|
||||
mock_worker_cls.return_value.get_kv_events.return_value = [event]
|
||||
kv_events = conn.get_kv_connector_kv_cache_events()
|
||||
|
||||
assert isinstance(kv_events, connector.MooncakeStoreKVEvents)
|
||||
assert kv_events.get_number_of_workers() == 1
|
||||
assert kv_events.get_all_events() == [event]
|
||||
|
||||
|
||||
def test_prefer_cross_layer_blocks_from_config():
|
||||
# Default: disabled
|
||||
vllm_config = _make_vllm_config()
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreScheduler"
|
||||
),
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.SCHEDULER)
|
||||
assert conn.prefer_cross_layer_blocks is False
|
||||
|
||||
# Enabled via config
|
||||
vllm_config_enabled = create_vllm_config(
|
||||
kv_connector="MooncakeStoreConnector",
|
||||
kv_role="kv_both",
|
||||
kv_connector_extra_config={"enable_cross_layers_blocks": "true"},
|
||||
)
|
||||
with (
|
||||
set_current_vllm_config(vllm_config_enabled),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreScheduler"
|
||||
),
|
||||
):
|
||||
conn_enabled = connector.MooncakeStoreConnector(
|
||||
vllm_config_enabled, KVConnectorRole.SCHEDULER
|
||||
)
|
||||
assert conn_enabled.prefer_cross_layer_blocks is True
|
||||
|
||||
|
||||
def test_register_cross_layers_kv_cache_delegates_to_worker():
|
||||
vllm_config = _make_vllm_config()
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreWorker"
|
||||
) as mock_worker_cls,
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
|
||||
fake_tensor = MagicMock()
|
||||
fake_backend = MagicMock()
|
||||
conn.register_cross_layers_kv_cache(fake_tensor, fake_backend)
|
||||
|
||||
worker_inst = mock_worker_cls.return_value
|
||||
worker_inst.register_cross_layers_kv_caches.assert_called_once_with(fake_tensor)
|
||||
|
||||
|
||||
def test_update_connector_output_and_take_events():
|
||||
vllm_config = _make_vllm_config()
|
||||
event = _make_block_stored()
|
||||
|
||||
with (
|
||||
set_current_vllm_config(vllm_config),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"connector.MooncakeStoreScheduler"
|
||||
),
|
||||
):
|
||||
conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.SCHEDULER)
|
||||
|
||||
kv_events = connector.MooncakeStoreKVEvents(num_workers=1)
|
||||
kv_events.add_events([event])
|
||||
conn.update_connector_output(KVConnectorOutput(kv_cache_events=kv_events))
|
||||
|
||||
assert conn._kv_cache_events is kv_events
|
||||
assert list(conn.take_events()) == [event]
|
||||
assert conn._kv_cache_events is None
|
||||
@@ -0,0 +1,300 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import threading
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import (
|
||||
worker,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501
|
||||
ChunkedTokenDatabase,
|
||||
KeyMetadata,
|
||||
ReqMeta,
|
||||
)
|
||||
|
||||
|
||||
def _make_store_sending_thread(
|
||||
store: MagicMock,
|
||||
) -> worker.KVCacheStoreSendingThread:
|
||||
token_database = ChunkedTokenDatabase(
|
||||
KeyMetadata("test-model", 0, 0, 0, 0), block_size=16
|
||||
)
|
||||
token_database.set_kv_caches_base_addr([0x1000])
|
||||
token_database.set_block_len([256])
|
||||
thread = worker.KVCacheStoreSendingThread(
|
||||
store=store,
|
||||
token_database=token_database,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
put_step=1,
|
||||
kv_role="kv_producer",
|
||||
ready_event=threading.Event(),
|
||||
)
|
||||
thread.request_queue.task_done = MagicMock()
|
||||
return thread
|
||||
|
||||
|
||||
def _make_store_req(req_id: str, block_hashes: list[bytes]) -> ReqMeta:
|
||||
return ReqMeta(
|
||||
req_id=req_id,
|
||||
token_len_chunk=32,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=block_hashes,
|
||||
can_save=True,
|
||||
original_block_size=16,
|
||||
)
|
||||
|
||||
|
||||
def test_store_sending_thread_skips_request_during_cpu_pressure():
|
||||
store = MagicMock()
|
||||
store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
|
||||
store.batch_put_from_multi_buffers.side_effect = [
|
||||
[-200, -200],
|
||||
[256, 256],
|
||||
[256, 256],
|
||||
]
|
||||
thread = _make_store_sending_thread(store)
|
||||
|
||||
thread.add_stored_request("req-a")
|
||||
thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))
|
||||
|
||||
assert thread._store_pressure_active is True
|
||||
assert "req-a" in thread._skip_store_requests
|
||||
assert store.batch_put_from_multi_buffers.call_count == 1
|
||||
|
||||
thread.add_stored_request("req-a")
|
||||
thread._handle_request(_make_store_req("req-a", [b"a2", b"a3"]))
|
||||
|
||||
assert store.batch_put_from_multi_buffers.call_count == 1
|
||||
|
||||
thread.add_stored_request("req-b")
|
||||
thread._handle_request(_make_store_req("req-b", [b"b0", b"b1"]))
|
||||
|
||||
assert thread._store_pressure_active is False
|
||||
assert "req-a" not in thread._skip_store_requests
|
||||
assert store.batch_put_from_multi_buffers.call_count == 2
|
||||
|
||||
thread.add_stored_request("req-a")
|
||||
thread._handle_request(_make_store_req("req-a", [b"a4", b"a5"]))
|
||||
|
||||
assert store.batch_put_from_multi_buffers.call_count == 3
|
||||
|
||||
|
||||
def test_store_sending_thread_only_skips_on_no_available_handle():
|
||||
store = MagicMock()
|
||||
store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
|
||||
store.batch_put_from_multi_buffers.side_effect = [
|
||||
[-500, -500],
|
||||
[256, 256],
|
||||
]
|
||||
thread = _make_store_sending_thread(store)
|
||||
|
||||
thread.add_stored_request("req-a")
|
||||
thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))
|
||||
|
||||
assert thread._store_pressure_active is False
|
||||
assert "req-a" not in thread._skip_store_requests
|
||||
assert store.batch_put_from_multi_buffers.call_count == 1
|
||||
|
||||
thread.add_stored_request("req-a")
|
||||
thread._handle_request(_make_store_req("req-a", [b"a2", b"a3"]))
|
||||
|
||||
assert store.batch_put_from_multi_buffers.call_count == 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers for register_kv_caches tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _auto_set_ready_event(*args, **kwargs):
|
||||
"""Side effect for mocked thread constructors that auto-sets ready_event."""
|
||||
for arg in args:
|
||||
if isinstance(arg, threading.Event):
|
||||
arg.set()
|
||||
for val in kwargs.values():
|
||||
if isinstance(val, threading.Event):
|
||||
val.set()
|
||||
return MagicMock()
|
||||
|
||||
|
||||
def _make_bare_worker(
|
||||
*,
|
||||
num_gpu_blocks: int = 10,
|
||||
block_size: int = 16,
|
||||
kv_role: str = "kv_both",
|
||||
) -> worker.MooncakeStoreWorker:
|
||||
"""Construct a MooncakeStoreWorker via __new__, bypassing __init__.
|
||||
|
||||
Sets only the attributes that register_kv_caches() reads so we can
|
||||
test the stride-based layout detection without a real
|
||||
MooncakeDistributedStore.
|
||||
"""
|
||||
w = object.__new__(worker.MooncakeStoreWorker)
|
||||
w.cache_config = MagicMock()
|
||||
w.cache_config.num_gpu_blocks = num_gpu_blocks
|
||||
w.store = MagicMock()
|
||||
w.store.register_buffer.return_value = 0
|
||||
w.use_mla = False
|
||||
w.token_database = ChunkedTokenDatabase(
|
||||
KeyMetadata("test-model", 0, 0, 0, 0), block_size=block_size
|
||||
)
|
||||
w.kv_role = kv_role
|
||||
w.block_size = block_size
|
||||
w.tp_rank = 0
|
||||
w.put_step = 1
|
||||
w.enable_kv_events = False
|
||||
w.kv_send_thread = None
|
||||
w.kv_recv_thread = None
|
||||
return w
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# register_kv_caches tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_register_kv_caches_blocks_first_single_segment():
|
||||
"""Blocks-first layout (FlashInfer/MLA): one segment per layer."""
|
||||
num_blocks = 10
|
||||
page_size_elements = 64 # elements per block
|
||||
w = _make_bare_worker(num_gpu_blocks=num_blocks)
|
||||
|
||||
# Shape: (num_blocks, page_size_elements) — blocks outermost, no outer_dims
|
||||
tensor = torch.zeros(num_blocks, page_size_elements, dtype=torch.float16)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreSendingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreRecvingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
):
|
||||
w.register_kv_caches({"layer0": tensor})
|
||||
|
||||
assert len(w.kv_caches_base_addr) == 1
|
||||
assert w.kv_caches_base_addr[0] == tensor.untyped_storage().data_ptr()
|
||||
|
||||
expected_block_len = tensor.untyped_storage().nbytes() // num_blocks
|
||||
assert len(w.block_len) == 1
|
||||
assert w.block_len[0] == expected_block_len
|
||||
|
||||
w.store.register_buffer.assert_called_once_with(
|
||||
tensor.untyped_storage().data_ptr(),
|
||||
tensor.untyped_storage().nbytes(),
|
||||
)
|
||||
|
||||
|
||||
def test_register_kv_caches_kv_first_two_segments():
|
||||
"""K/V-first layout (FlashAttn): two segments (K, V) per layer."""
|
||||
num_blocks = 10
|
||||
block_size_tokens = 16
|
||||
num_kv_heads = 4
|
||||
head_size = 8
|
||||
|
||||
w = _make_bare_worker(num_gpu_blocks=num_blocks)
|
||||
|
||||
# Shape: (2, num_blocks, block_size, num_kv_heads, head_size) — K/V outermost
|
||||
tensor = torch.zeros(
|
||||
2,
|
||||
num_blocks,
|
||||
block_size_tokens,
|
||||
num_kv_heads,
|
||||
head_size,
|
||||
dtype=torch.float16,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreSendingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreRecvingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
):
|
||||
w.register_kv_caches({"layer0": tensor})
|
||||
|
||||
# K/V-first: dim 0 has stride > page_size, so 2 segments
|
||||
assert len(w.kv_caches_base_addr) == 2
|
||||
assert len(w.block_len) == 2
|
||||
|
||||
el = tensor.element_size()
|
||||
seg_stride = tensor.stride(0) * el # stride of the K/V dim in bytes
|
||||
base = tensor.untyped_storage().data_ptr()
|
||||
assert w.kv_caches_base_addr[0] == base
|
||||
assert w.kv_caches_base_addr[1] == base + seg_stride
|
||||
assert w.block_len[0] == seg_stride // num_blocks
|
||||
assert w.block_len[1] == seg_stride // num_blocks
|
||||
|
||||
|
||||
def test_register_kv_caches_cross_layer_single_segment():
|
||||
"""Cross-layer tensor: single segment with block_len = page_size * num_layers."""
|
||||
num_blocks = 10
|
||||
num_layers = 4
|
||||
per_layer_page_elements = 64 # elements per layer per block
|
||||
|
||||
w = _make_bare_worker(num_gpu_blocks=num_blocks)
|
||||
|
||||
# Cross-layer blocks-first tensor: all layers packed into a single
|
||||
# contiguous block. Shape (num_blocks, num_layers * per_layer_page)
|
||||
# mimics the physical layout after stride reordering.
|
||||
total_page_elements = num_layers * per_layer_page_elements
|
||||
tensor = torch.zeros(num_blocks, total_page_elements, dtype=torch.float16)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreSendingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreRecvingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
):
|
||||
# Use the cross-layer wrapper key, same as register_cross_layers_kv_caches
|
||||
w.register_kv_caches({"__cross_layer__": tensor})
|
||||
|
||||
assert len(w.kv_caches_base_addr) == 1
|
||||
assert w.kv_caches_base_addr[0] == tensor.untyped_storage().data_ptr()
|
||||
|
||||
expected_block_len = tensor.untyped_storage().nbytes() // num_blocks
|
||||
# block_len should be per_layer_page_size * num_layers
|
||||
assert (
|
||||
expected_block_len
|
||||
== num_layers * per_layer_page_elements * tensor.element_size()
|
||||
)
|
||||
assert len(w.block_len) == 1
|
||||
assert w.block_len[0] == expected_block_len
|
||||
|
||||
# Also verify via register_cross_layers_kv_caches wrapper
|
||||
w2 = _make_bare_worker(num_gpu_blocks=num_blocks)
|
||||
with (
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreSendingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
|
||||
"worker.KVCacheStoreRecvingThread",
|
||||
side_effect=_auto_set_ready_event,
|
||||
),
|
||||
):
|
||||
w2.register_cross_layers_kv_caches(tensor)
|
||||
|
||||
assert w2.kv_caches_base_addr == w.kv_caches_base_addr
|
||||
assert w2.block_len == w.block_len
|
||||
@@ -117,10 +117,10 @@ def test_already_stored_block_not_evicted_during_prepare_store(eviction_policy):
|
||||
|
||||
# store [1, 2] and complete
|
||||
manager.prepare_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
manager.complete_store(to_keys([1, 2]))
|
||||
manager.complete_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
|
||||
# touch [1] to make block 2 the LRU candidate
|
||||
manager.touch(to_keys([1]))
|
||||
manager.touch(to_keys([1]), _EMPTY_REQ_CTX)
|
||||
|
||||
# prepare_store([2, 3, 4, 5]):
|
||||
# - block 2 is already stored -> filtered out of keys_to_store
|
||||
@@ -137,7 +137,7 @@ def test_already_stored_block_not_evicted_during_prepare_store(eviction_policy):
|
||||
)
|
||||
|
||||
# complete_store must not silently drop block 2
|
||||
manager.complete_store(to_keys([2, 3, 4, 5]))
|
||||
manager.complete_store(to_keys([2, 3, 4, 5]), _EMPTY_REQ_CTX)
|
||||
|
||||
# block 2 must still be present in the cache
|
||||
assert manager.lookup(to_key(2), _EMPTY_REQ_CTX) is True
|
||||
@@ -171,7 +171,7 @@ def test_cpu_manager():
|
||||
assert list(cpu_manager.take_events()) == []
|
||||
|
||||
# complete store [1, 2]
|
||||
cpu_manager.complete_store(to_keys([1, 2]))
|
||||
cpu_manager.complete_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
verify_events(cpu_manager.take_events(), expected_stores=({1, 2},))
|
||||
|
||||
# lookup [1, 2]
|
||||
@@ -199,7 +199,7 @@ def test_cpu_manager():
|
||||
assert cpu_manager.prepare_store(to_keys([1, 6]), _EMPTY_REQ_CTX) is None
|
||||
|
||||
# complete store [2, 3, 4, 5]
|
||||
cpu_manager.complete_store(to_keys([2, 3, 4, 5]))
|
||||
cpu_manager.complete_store(to_keys([2, 3, 4, 5]), _EMPTY_REQ_CTX)
|
||||
|
||||
# lookup (now that we have [2, 3, 4, 5])
|
||||
assert cpu_manager.lookup(to_key(1), _EMPTY_REQ_CTX) is False
|
||||
@@ -217,7 +217,7 @@ def test_cpu_manager():
|
||||
assert cpu_manager.prepare_store(to_keys([6, 7, 8]), _EMPTY_REQ_CTX) is None
|
||||
|
||||
# complete load [2, 3]
|
||||
cpu_manager.complete_load(to_keys([2, 3]))
|
||||
cpu_manager.complete_load(to_keys([2, 3]), _EMPTY_REQ_CTX)
|
||||
|
||||
# prepare store [6, 7, 8] -> evicts [2, 3, 4] (oldest)
|
||||
prepare_store_output = cpu_manager.prepare_store(to_keys([6, 7, 8]), _EMPTY_REQ_CTX)
|
||||
@@ -231,10 +231,10 @@ def test_cpu_manager():
|
||||
)
|
||||
|
||||
# complete store [6, 7, 8]
|
||||
cpu_manager.complete_store(to_keys([6, 7, 8]))
|
||||
cpu_manager.complete_store(to_keys([6, 7, 8]), _EMPTY_REQ_CTX)
|
||||
|
||||
# touch [5, 6, 7] (move to end of LRU order)
|
||||
cpu_manager.touch(to_keys([5, 6, 7]))
|
||||
cpu_manager.touch(to_keys([5, 6, 7]), _EMPTY_REQ_CTX)
|
||||
|
||||
# prepare store [7, 9] -> evicts [8] (oldest following previous touch)
|
||||
prepare_store_output = cpu_manager.prepare_store(to_keys([9]), _EMPTY_REQ_CTX)
|
||||
@@ -248,7 +248,7 @@ def test_cpu_manager():
|
||||
)
|
||||
|
||||
# complete store [7, 9] with failure
|
||||
cpu_manager.complete_store(to_keys([7, 9]), success=False)
|
||||
cpu_manager.complete_store(to_keys([7, 9]), _EMPTY_REQ_CTX, success=False)
|
||||
|
||||
# assert [7] is still stored, but [9] is not
|
||||
assert cpu_manager.lookup(to_key(7), _EMPTY_REQ_CTX) is True
|
||||
@@ -304,7 +304,7 @@ class TestARCPolicy:
|
||||
assert list(cpu_manager.take_events()) == []
|
||||
|
||||
# complete store [1, 2]
|
||||
cpu_manager.complete_store(to_keys([1, 2]))
|
||||
cpu_manager.complete_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
verify_events(cpu_manager.take_events(), expected_stores=({1, 2},))
|
||||
|
||||
# lookup [1, 2]
|
||||
@@ -325,14 +325,14 @@ class TestARCPolicy:
|
||||
|
||||
# store and complete block 1
|
||||
cpu_manager.prepare_store(to_keys([1]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([1]))
|
||||
cpu_manager.complete_store(to_keys([1]), _EMPTY_REQ_CTX)
|
||||
|
||||
# block 1 starts in T1 (recent)
|
||||
assert to_keys([1])[0] in arc_policy.t1
|
||||
assert to_keys([1])[0] not in arc_policy.t2
|
||||
|
||||
# touch block 1 (simulate second access)
|
||||
cpu_manager.touch(to_keys([1]))
|
||||
cpu_manager.touch(to_keys([1]), _EMPTY_REQ_CTX)
|
||||
|
||||
# block 1 should now be in T2 (frequent)
|
||||
assert to_keys([1])[0] not in arc_policy.t1
|
||||
@@ -357,7 +357,7 @@ class TestARCPolicy:
|
||||
evicted_keys=[],
|
||||
),
|
||||
)
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]))
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]), _EMPTY_REQ_CTX)
|
||||
|
||||
# prepare load [2, 3] (increases ref_cnt)
|
||||
prepare_load_output = cpu_manager.prepare_load(to_keys([2, 3]), _EMPTY_REQ_CTX)
|
||||
@@ -368,7 +368,7 @@ class TestARCPolicy:
|
||||
assert cpu_manager.prepare_store(to_keys([5, 6, 7]), _EMPTY_REQ_CTX) is None
|
||||
|
||||
# complete load [2, 3]
|
||||
cpu_manager.complete_load(to_keys([2, 3]))
|
||||
cpu_manager.complete_load(to_keys([2, 3]), _EMPTY_REQ_CTX)
|
||||
|
||||
# now prepare store [5, 6, 7] should succeed
|
||||
# ARC will evict blocks one at a time from T1 as needed
|
||||
@@ -389,20 +389,20 @@ class TestARCPolicy:
|
||||
|
||||
# store blocks 1, 2 (fills cache)
|
||||
cpu_manager.prepare_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([1, 2]))
|
||||
cpu_manager.complete_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
|
||||
initial_target = arc_policy.target_t1_size
|
||||
|
||||
# store block 3, evicting block 1 (moves to B1 ghost list)
|
||||
cpu_manager.prepare_store(to_keys([3]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([3]))
|
||||
cpu_manager.complete_store(to_keys([3]), _EMPTY_REQ_CTX)
|
||||
|
||||
# block 1 should be in B1 (ghost list)
|
||||
assert to_keys([1])[0] in arc_policy.b1
|
||||
|
||||
# touch block 1 (cache miss, but in B1)
|
||||
# this should increase target_t1_size (favor recency)
|
||||
cpu_manager.touch(to_keys([1]))
|
||||
cpu_manager.touch(to_keys([1]), _EMPTY_REQ_CTX)
|
||||
|
||||
# target should have increased
|
||||
assert arc_policy.target_t1_size > initial_target
|
||||
@@ -416,10 +416,10 @@ class TestARCPolicy:
|
||||
|
||||
# store blocks 1, 2, 3, 4
|
||||
cpu_manager.prepare_store(to_keys([1, 2, 3, 4]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]))
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]), _EMPTY_REQ_CTX)
|
||||
|
||||
# promote blocks 3, 4 to T2 by touching them
|
||||
cpu_manager.touch(to_keys([3, 4]))
|
||||
cpu_manager.touch(to_keys([3, 4]), _EMPTY_REQ_CTX)
|
||||
|
||||
# now: T1 = {1, 2}, T2 = {3, 4}
|
||||
assert len(arc_policy.t1) == 2
|
||||
@@ -434,7 +434,7 @@ class TestARCPolicy:
|
||||
assert output is not None
|
||||
assert to_keys([1]) == output.evicted_keys
|
||||
|
||||
cpu_manager.complete_store(to_keys([5]))
|
||||
cpu_manager.complete_store(to_keys([5]), _EMPTY_REQ_CTX)
|
||||
|
||||
# block 1 should be in B1 (ghost list)
|
||||
assert to_keys([1])[0] in arc_policy.b1
|
||||
@@ -450,12 +450,12 @@ class TestARCPolicy:
|
||||
|
||||
# fill cache with blocks 1, 2
|
||||
cpu_manager.prepare_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([1, 2]))
|
||||
cpu_manager.complete_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
|
||||
# store many blocks to fill ghost lists
|
||||
for i in range(3, 20):
|
||||
cpu_manager.prepare_store(to_keys([i]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([i]))
|
||||
cpu_manager.complete_store(to_keys([i]), _EMPTY_REQ_CTX)
|
||||
|
||||
# ghost lists should not exceed cache_capacity
|
||||
assert len(arc_policy.b1) <= arc_policy.cache_capacity
|
||||
@@ -470,14 +470,14 @@ class TestARCPolicy:
|
||||
|
||||
# store blocks 1, 2, 3, 4
|
||||
cpu_manager.prepare_store(to_keys([1, 2, 3, 4]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]))
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]), _EMPTY_REQ_CTX)
|
||||
|
||||
# promote 3, 4 to T2
|
||||
cpu_manager.touch(to_keys([3, 4]))
|
||||
cpu_manager.touch(to_keys([3, 4]), _EMPTY_REQ_CTX)
|
||||
|
||||
# T1 = {1, 2}, T2 = {3, 4}
|
||||
# touch [1, 3, 4] - should promote 1 to T2, and move 3,4 to end of T2
|
||||
cpu_manager.touch(to_keys([1, 3, 4]))
|
||||
cpu_manager.touch(to_keys([1, 3, 4]), _EMPTY_REQ_CTX)
|
||||
|
||||
# T1 = {2}, T2 = {1, 3, 4} (in that order, with 4 most recent)
|
||||
assert len(arc_policy.t1) == 1
|
||||
@@ -503,7 +503,7 @@ class TestARCPolicy:
|
||||
|
||||
# store blocks 1, 2, 3, 4
|
||||
cpu_manager.prepare_store(to_keys([1, 2, 3, 4]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]))
|
||||
cpu_manager.complete_store(to_keys([1, 2, 3, 4]), _EMPTY_REQ_CTX)
|
||||
|
||||
# prepare store block 5 (will evict block 1)
|
||||
prepare_store_output = cpu_manager.prepare_store(to_keys([5]), _EMPTY_REQ_CTX)
|
||||
@@ -511,7 +511,7 @@ class TestARCPolicy:
|
||||
assert len(prepare_store_output.evicted_keys) == 1
|
||||
|
||||
# complete store with failure
|
||||
cpu_manager.complete_store(to_keys([5]), success=False)
|
||||
cpu_manager.complete_store(to_keys([5]), _EMPTY_REQ_CTX, success=False)
|
||||
|
||||
# block 5 should not be in cache
|
||||
assert cpu_manager.lookup(to_key(5), _EMPTY_REQ_CTX) is False
|
||||
@@ -532,7 +532,7 @@ class TestARCPolicy:
|
||||
|
||||
# store [1, 2]
|
||||
cpu_manager.prepare_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
cpu_manager.complete_store(to_keys([1, 2]))
|
||||
cpu_manager.complete_store(to_keys([1, 2]), _EMPTY_REQ_CTX)
|
||||
|
||||
# store [3, 4, 5] -> evicts [1]
|
||||
prepare_store_output = cpu_manager.prepare_store(
|
||||
@@ -540,10 +540,10 @@ class TestARCPolicy:
|
||||
)
|
||||
assert prepare_store_output is not None
|
||||
assert len(prepare_store_output.evicted_keys) == 1
|
||||
cpu_manager.complete_store(to_keys([3, 4, 5]))
|
||||
cpu_manager.complete_store(to_keys([3, 4, 5]), _EMPTY_REQ_CTX)
|
||||
|
||||
# promote some blocks to T2
|
||||
cpu_manager.touch(to_keys([2, 3]))
|
||||
cpu_manager.touch(to_keys([2, 3]), _EMPTY_REQ_CTX)
|
||||
|
||||
# T1 has {4, 5}, T2 has {2, 3}
|
||||
assert len(arc_policy.t1) == 2
|
||||
@@ -552,7 +552,7 @@ class TestARCPolicy:
|
||||
# store [6] -> should evict from T1 (4 is oldest in T1)
|
||||
prepare_store_output = cpu_manager.prepare_store(to_keys([6]), _EMPTY_REQ_CTX)
|
||||
assert prepare_store_output is not None
|
||||
cpu_manager.complete_store(to_keys([6]))
|
||||
cpu_manager.complete_store(to_keys([6]), _EMPTY_REQ_CTX)
|
||||
|
||||
# verify blocks 2, 3 (in T2) are still present
|
||||
assert cpu_manager.lookup(to_key(2), _EMPTY_REQ_CTX) is True
|
||||
@@ -609,4 +609,4 @@ def test_filter_reused_manager():
|
||||
assert prepare_store_output is not None
|
||||
assert prepare_store_output.keys_to_store == []
|
||||
|
||||
manager.complete_store(to_keys([1]))
|
||||
manager.complete_store(to_keys([1]), _EMPTY_REQ_CTX)
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Build DeepGEMM's `_C` pybind11 extension for a target Python.
|
||||
|
||||
Driven from `cmake/external_projects/deepgemm.cmake`. The driver is the
|
||||
build interpreter (which has torch); the *target* Python is only used for
|
||||
its header path and SOABI. This avoids needing torch installed in N venvs
|
||||
to produce N matching `.so` files.
|
||||
|
||||
Usage: python build_deepgemm_C.py <DEEPGEMM_SRC_DIR> <OUTPUT_DIR> <TARGET_PY>
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from torch.utils import cpp_extension
|
||||
|
||||
if len(sys.argv) != 4:
|
||||
sys.exit(f"usage: {sys.argv[0]} <SRC> <OUT> <TARGET_PY>")
|
||||
|
||||
src = Path(sys.argv[1]).resolve()
|
||||
out = Path(sys.argv[2]).resolve()
|
||||
target_py = sys.argv[3]
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
info = json.loads(
|
||||
subprocess.check_output(
|
||||
[
|
||||
target_py,
|
||||
"-c",
|
||||
"import sysconfig, json; "
|
||||
"print(json.dumps({k: sysconfig.get_config_var(k) "
|
||||
"for k in ('EXT_SUFFIX', 'INCLUDEPY')}))",
|
||||
]
|
||||
).decode()
|
||||
)
|
||||
|
||||
cuda_home = cpp_extension.CUDA_HOME
|
||||
if cuda_home is None:
|
||||
sys.exit("CUDA_HOME not found; cannot build DeepGEMM _C")
|
||||
# CCCL lives outside the standard CUDAToolkit search, mirroring DeepGEMM's
|
||||
# own setup.py.
|
||||
includes = [
|
||||
info["INCLUDEPY"],
|
||||
f"{cuda_home}/include",
|
||||
f"{cuda_home}/include/cccl",
|
||||
str(src / "csrc"),
|
||||
str(src / "deep_gemm/include"),
|
||||
str(src / "third-party/cutlass/include"),
|
||||
str(src / "third-party/cutlass/tools/util/include"),
|
||||
str(src / "third-party/fmt/include"),
|
||||
*cpp_extension.include_paths(device_type="cuda"),
|
||||
]
|
||||
|
||||
cmd = [
|
||||
os.environ.get("CXX", "g++"),
|
||||
"-shared",
|
||||
"-fPIC",
|
||||
"-std=c++20",
|
||||
"-O3",
|
||||
"-g0",
|
||||
"-Wno-psabi",
|
||||
"-Wno-deprecated-declarations",
|
||||
"-DTORCH_API_INCLUDE_EXTENSION_H",
|
||||
"-DTORCH_EXTENSION_NAME=_C",
|
||||
f"-D_GLIBCXX_USE_CXX11_ABI={int(torch.compiled_with_cxx11_abi())}",
|
||||
*(f"-I{p}" for p in includes),
|
||||
str(src / "csrc/python_api.cpp"),
|
||||
*(f"-L{p}" for p in cpp_extension.library_paths(device_type="cuda")),
|
||||
f"-L{cuda_home}/lib64",
|
||||
"-ltorch",
|
||||
"-ltorch_python",
|
||||
"-ltorch_cpu",
|
||||
"-ltorch_cuda",
|
||||
"-lc10",
|
||||
"-lc10_cuda",
|
||||
"-lcudart",
|
||||
"-lnvrtc",
|
||||
"-o",
|
||||
str(out / f"_C{info['EXT_SUFFIX']}"),
|
||||
]
|
||||
print("[build_deepgemm_C] " + " ".join(cmd), flush=True)
|
||||
subprocess.check_call(cmd)
|
||||
@@ -0,0 +1,41 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
"""Assert the installed vLLM has a `_C.cpython-X.Y-*.so` for every CPython
|
||||
covered by `requires-python`. Fails closed if a Python's `.so` is missing
|
||||
from the wheel — i.e. the regression that surfaced in #41476/#41512.
|
||||
|
||||
Run from a CI test job after vLLM is installed, e.g. the H100 deepgemm
|
||||
kernel tests in .buildkite/test_areas/kernels.yaml.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import regex as re
|
||||
import tomllib
|
||||
|
||||
SO_RE = re.compile(r"^_C\.cpython-(\d)(\d+)-")
|
||||
|
||||
|
||||
def required_pythons() -> list[str]:
|
||||
pyproject = Path(__file__).resolve().parent.parent / "pyproject.toml"
|
||||
spec = tomllib.loads(pyproject.read_text())["project"]["requires-python"]
|
||||
m = re.match(r">=3\.(\d+),<3\.(\d+)", spec)
|
||||
if not m:
|
||||
sys.exit(f"unexpected requires-python format: {spec!r}")
|
||||
return [f"3.{v}" for v in range(int(m[1]), int(m[2]))]
|
||||
|
||||
|
||||
spec = importlib.util.find_spec("vllm.third_party.deep_gemm")
|
||||
if spec is None or spec.origin is None:
|
||||
sys.exit("vllm.third_party.deep_gemm not importable; is vllm installed?")
|
||||
pkg_dir = Path(spec.origin).parent
|
||||
|
||||
found = {f"{m[1]}.{m[2]}" for f in os.listdir(pkg_dir) if (m := SO_RE.match(f))}
|
||||
required = required_pythons()
|
||||
missing = [v for v in required if v not in found]
|
||||
print(f"deepgemm _C: found {sorted(found)}, required {required}, missing {missing}")
|
||||
sys.exit(1 if missing else 0)
|
||||
Executable
+49
@@ -0,0 +1,49 @@
|
||||
#!/usr/bin/env bash
|
||||
# Provision bare Python interpreters for the DeepGEMM `_C` per-Python build
|
||||
# and print a colon-separated list of their paths to stdout.
|
||||
#
|
||||
# Each target Python only needs a working interpreter — torch is not
|
||||
# installed since `tools/build_deepgemm_C.py` runs from the build interpreter.
|
||||
# uv re-uses any matching system Python and downloads a managed build
|
||||
# otherwise.
|
||||
#
|
||||
# Usage:
|
||||
# export DEEPGEMM_PYTHON_INTERPRETERS=$(tools/setup_deepgemm_pythons.sh)
|
||||
# python setup.py bdist_wheel --dist-dir=dist --py-limited-api=cp38
|
||||
#
|
||||
# With no args, expands to every CPython covered by `requires-python` in
|
||||
# pyproject.toml. Pass explicit versions (e.g. `3.10 3.11`) to override.
|
||||
#
|
||||
# Skip this script if you don't have uv: set DEEPGEMM_PYTHON_INTERPRETERS
|
||||
# directly to existing interpreter paths. Editable / single-Python builds
|
||||
# don't need the env var at all (cmake falls back to the build interpreter).
|
||||
#
|
||||
# Optional: DEEPGEMM_VENV_PREFIX (default: /tmp/dgenv).
|
||||
set -euo pipefail
|
||||
|
||||
if [ "$#" -eq 0 ]; then
|
||||
# Derive the matrix from `requires-python = ">=3.X,<3.Y"` in pyproject.toml.
|
||||
pyproject="$(dirname "$0")/../pyproject.toml"
|
||||
spec=$(grep -E '^requires-python' "$pyproject" \
|
||||
| grep -oE '>=3\.[0-9]+,<3\.[0-9]+')
|
||||
lo=${spec#>=3.}; lo=${lo%%,*}
|
||||
hi=${spec##*<3.}
|
||||
set -- $(seq "$lo" $((hi - 1)) | sed 's/^/3./')
|
||||
fi
|
||||
|
||||
prefix="${DEEPGEMM_VENV_PREFIX:-/tmp/dgenv}"
|
||||
mkdir -p "$prefix"
|
||||
|
||||
paths=""
|
||||
for V in "$@"; do
|
||||
venv="$prefix/$V"
|
||||
# Force a managed (uv-downloaded) Python so dev headers are bundled.
|
||||
# System Pythons on the build base may lack headers (manylinux's
|
||||
# /opt/python/cpXY-cpXY are off PATH; an apt-installed python3.X often
|
||||
# has no -dev), and the per-Python build needs Python.h.
|
||||
[ -x "$venv/bin/python" ] || \
|
||||
uv venv --python "$V" "$venv" --python-preference only-managed --seed \
|
||||
>/dev/null
|
||||
paths="$paths:$venv/bin/python"
|
||||
done
|
||||
echo "${paths#:}"
|
||||
@@ -9,6 +9,9 @@ from vllm.config.utils import config
|
||||
from vllm.logger import init_logger
|
||||
from vllm.utils.hashing import safe_hash
|
||||
|
||||
DEFAULT_SAFETENSORS_PREFETCH_NUM_THREADS = 8
|
||||
DEFAULT_SAFETENSORS_PREFETCH_BLOCK_SIZE = 16 * 1024 * 1024
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.model_executor.model_loader import LoadFormats
|
||||
from vllm.model_executor.model_loader.tensorizer import TensorizerConfig
|
||||
@@ -79,6 +82,15 @@ class LoadConfig:
|
||||
was quantized using torchao and saved using safetensors.
|
||||
Needs `torchao >= 0.14.0`.
|
||||
"""
|
||||
safetensors_prefetch_num_threads: int = Field(
|
||||
default=DEFAULT_SAFETENSORS_PREFETCH_NUM_THREADS, ge=1
|
||||
)
|
||||
"""Number of worker threads used to prefetch safetensors checkpoint files
|
||||
into the OS page cache when safetensors prefetching is enabled."""
|
||||
safetensors_prefetch_block_size: int = Field(
|
||||
default=DEFAULT_SAFETENSORS_PREFETCH_BLOCK_SIZE, ge=1
|
||||
)
|
||||
"""Read size in bytes for each safetensors checkpoint file prefetch."""
|
||||
model_loader_extra_config: dict | TensorizerConfig = Field(default_factory=dict)
|
||||
"""Extra config for model loader. This will be passed to the model loader
|
||||
corresponding to the chosen load_format."""
|
||||
|
||||
@@ -197,6 +197,11 @@ KVConnectorFactory.register_connector(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_connector",
|
||||
"MooncakeConnector",
|
||||
)
|
||||
KVConnectorFactory.register_connector(
|
||||
"MooncakeStoreConnector",
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.connector",
|
||||
"MooncakeStoreConnector",
|
||||
)
|
||||
KVConnectorFactory.register_connector(
|
||||
"FlexKVConnectorV1",
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.flexkv_connector",
|
||||
|
||||
@@ -8,6 +8,7 @@ import uvicorn
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from pydantic import BaseModel
|
||||
|
||||
from vllm.config import ParallelConfig
|
||||
from vllm.distributed.kv_transfer.kv_connector.utils import EngineId
|
||||
from vllm.logger import init_logger
|
||||
|
||||
@@ -16,6 +17,15 @@ WorkerAddr = str
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def get_mooncake_dp_engine_index(parallel_config: ParallelConfig) -> int:
|
||||
"""Return the per-engine DP index used for Mooncake side channels."""
|
||||
if parallel_config.local_engines_only:
|
||||
assert parallel_config.data_parallel_rank_local is not None
|
||||
return parallel_config.data_parallel_rank_local
|
||||
|
||||
return parallel_config.data_parallel_index
|
||||
|
||||
|
||||
class RegisterWorkerPayload(BaseModel):
|
||||
engine_id: EngineId
|
||||
dp_rank: int
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
@@ -0,0 +1,229 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# Adapted from vllm-project/vllm-ascend
|
||||
# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/).
|
||||
"""MooncakeStoreConnector - KV cache connector using MooncakeDistributedStore.
|
||||
|
||||
Unlike MooncakeConnector which does direct P2P transfer, this connector
|
||||
uses MooncakeDistributedStore as a shared KV cache pool. Both producer
|
||||
and consumer instances read/write KV to/from the store independently,
|
||||
enabling prefix caching via hash-based deduplication.
|
||||
"""
|
||||
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.distributed.kv_events import (
|
||||
KVCacheEvent,
|
||||
KVConnectorKVEvents,
|
||||
KVEventAggregator,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
|
||||
KVConnectorBase_V1,
|
||||
KVConnectorMetadata,
|
||||
KVConnectorRole,
|
||||
)
|
||||
from vllm.forward_context import ForwardContext
|
||||
from vllm.logger import init_logger
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
|
||||
from vllm.v1.core.sched.output import SchedulerOutput
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig
|
||||
from vllm.v1.outputs import KVConnectorOutput
|
||||
from vllm.v1.request import Request
|
||||
|
||||
from .data import MooncakeStoreConnectorMetadata
|
||||
from .scheduler import MooncakeStoreScheduler
|
||||
from .worker import MooncakeStoreWorker
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MooncakeStoreKVEvents(KVConnectorKVEvents):
|
||||
"""KV event aggregation for MooncakeStoreConnector."""
|
||||
|
||||
def __init__(self, num_workers: int) -> None:
|
||||
self._aggregator = KVEventAggregator(num_workers)
|
||||
|
||||
def add_events(self, events: list[KVCacheEvent]) -> None:
|
||||
self._aggregator.add_events(events)
|
||||
|
||||
def aggregate(self) -> "MooncakeStoreKVEvents":
|
||||
common_events = self._aggregator.get_common_events()
|
||||
self._aggregator.clear_events()
|
||||
self._aggregator.add_events(common_events)
|
||||
self._aggregator.reset_workers()
|
||||
return self
|
||||
|
||||
def increment_workers(self, count: int = 1) -> None:
|
||||
self._aggregator.increment_workers(count)
|
||||
|
||||
def get_all_events(self) -> list[KVCacheEvent]:
|
||||
return self._aggregator.get_all_events()
|
||||
|
||||
def get_number_of_workers(self) -> int:
|
||||
return self._aggregator.get_number_of_workers()
|
||||
|
||||
def clear_events(self) -> None:
|
||||
self._aggregator.clear_events()
|
||||
self._aggregator.reset_workers()
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<MooncakeStoreKVEvents events={self.get_all_events()}>"
|
||||
|
||||
|
||||
class MooncakeStoreConnector(KVConnectorBase_V1):
|
||||
"""KV connector using MooncakeDistributedStore as shared KV pool."""
|
||||
|
||||
@property
|
||||
def prefer_cross_layer_blocks(self) -> bool:
|
||||
extra_config = self._kv_transfer_config.kv_connector_extra_config
|
||||
return (
|
||||
str(extra_config.get("enable_cross_layers_blocks", "False")).lower()
|
||||
== "true"
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vllm_config: VllmConfig,
|
||||
role: KVConnectorRole,
|
||||
kv_cache_config: KVCacheConfig | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
vllm_config=vllm_config,
|
||||
role=role,
|
||||
kv_cache_config=kv_cache_config, # type: ignore[arg-type]
|
||||
)
|
||||
assert vllm_config.kv_transfer_config is not None
|
||||
self.kv_role = vllm_config.kv_transfer_config.kv_role
|
||||
self._kv_cache_events: MooncakeStoreKVEvents | None = None
|
||||
|
||||
self.connector_scheduler: MooncakeStoreScheduler | None = None
|
||||
self.connector_worker: MooncakeStoreWorker | None = None
|
||||
|
||||
if role == KVConnectorRole.SCHEDULER:
|
||||
self.connector_scheduler = MooncakeStoreScheduler(vllm_config)
|
||||
else:
|
||||
self.connector_worker = MooncakeStoreWorker(vllm_config)
|
||||
|
||||
# ============================================================
|
||||
# Scheduler-side methods
|
||||
# ============================================================
|
||||
|
||||
def get_num_new_matched_tokens(
|
||||
self,
|
||||
request: Request,
|
||||
num_computed_tokens: int,
|
||||
) -> tuple[int, bool]:
|
||||
assert self.connector_scheduler is not None
|
||||
return self.connector_scheduler.get_num_new_matched_tokens(
|
||||
request, num_computed_tokens
|
||||
)
|
||||
|
||||
def update_state_after_alloc(
|
||||
self,
|
||||
request: Request,
|
||||
blocks: KVCacheBlocks,
|
||||
num_external_tokens: int,
|
||||
):
|
||||
assert self.connector_scheduler is not None
|
||||
return self.connector_scheduler.update_state_after_alloc(
|
||||
request, blocks, num_external_tokens
|
||||
)
|
||||
|
||||
def build_connector_meta(
|
||||
self,
|
||||
scheduler_output: SchedulerOutput,
|
||||
) -> KVConnectorMetadata:
|
||||
assert self.connector_scheduler is not None
|
||||
return self.connector_scheduler.build_connector_meta(scheduler_output)
|
||||
|
||||
def request_finished(
|
||||
self,
|
||||
request: Request,
|
||||
block_ids: list[int],
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
assert self.connector_scheduler is not None
|
||||
return self.connector_scheduler.request_finished(request, block_ids)
|
||||
|
||||
def update_connector_output(self, connector_output: KVConnectorOutput):
|
||||
kv_cache_events = connector_output.kv_cache_events
|
||||
if not kv_cache_events or not isinstance(
|
||||
kv_cache_events, MooncakeStoreKVEvents
|
||||
):
|
||||
return
|
||||
|
||||
if self._kv_cache_events is None:
|
||||
self._kv_cache_events = kv_cache_events
|
||||
else:
|
||||
self._kv_cache_events.add_events(kv_cache_events.get_all_events())
|
||||
self._kv_cache_events.increment_workers(
|
||||
kv_cache_events.get_number_of_workers()
|
||||
)
|
||||
|
||||
def take_events(self) -> Iterable[KVCacheEvent]:
|
||||
if self._kv_cache_events is not None:
|
||||
self._kv_cache_events.aggregate()
|
||||
yield from self._kv_cache_events.get_all_events()
|
||||
self._kv_cache_events.clear_events()
|
||||
self._kv_cache_events = None
|
||||
|
||||
# ============================================================
|
||||
# Worker-side methods
|
||||
# ============================================================
|
||||
|
||||
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
|
||||
assert self.connector_worker is not None
|
||||
self.connector_worker.register_kv_caches(kv_caches)
|
||||
|
||||
def register_cross_layers_kv_cache(
|
||||
self, kv_cache: torch.Tensor, attn_backend: type
|
||||
):
|
||||
assert self.connector_worker is not None
|
||||
self.connector_worker.register_cross_layers_kv_caches(kv_cache)
|
||||
|
||||
def start_load_kv(self, forward_context: ForwardContext, **kwargs: Any) -> None:
|
||||
# No-op: loads are issued in get_finished() for compute overlap.
|
||||
pass
|
||||
|
||||
def wait_for_layer_load(self, layer_name: str) -> None:
|
||||
# No layerwise support - no-op
|
||||
return
|
||||
|
||||
def save_kv_layer(
|
||||
self,
|
||||
layer_name: str,
|
||||
kv_layer: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
# No layerwise support - no-op
|
||||
return
|
||||
|
||||
def wait_for_save(self):
|
||||
# No-op: stores are issued in get_finished() for compute overlap.
|
||||
pass
|
||||
|
||||
def get_finished(
|
||||
self, finished_req_ids: set[str]
|
||||
) -> tuple[set[str] | None, set[str] | None]:
|
||||
assert self.connector_worker is not None
|
||||
metadata = self._get_connector_metadata()
|
||||
assert isinstance(metadata, MooncakeStoreConnectorMetadata)
|
||||
return self.connector_worker.get_finished(finished_req_ids, metadata)
|
||||
|
||||
def get_kv_connector_kv_cache_events(
|
||||
self,
|
||||
) -> MooncakeStoreKVEvents | None:
|
||||
assert self.connector_worker is not None
|
||||
events = self.connector_worker.get_kv_events()
|
||||
if not events:
|
||||
return None
|
||||
|
||||
kv_events = MooncakeStoreKVEvents(num_workers=1)
|
||||
kv_events.add_events(events)
|
||||
return kv_events
|
||||
@@ -0,0 +1,276 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# Adapted from vllm-project/vllm-ascend
|
||||
# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/).
|
||||
"""Data classes for MooncakeStoreConnector."""
|
||||
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
|
||||
KVConnectorMetadata,
|
||||
)
|
||||
from vllm.logger import init_logger
|
||||
from vllm.utils.math_utils import cdiv
|
||||
from vllm.v1.core.kv_cache_utils import BlockHash
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class KeyMetadata:
|
||||
"""Metadata for constructing pool keys."""
|
||||
|
||||
model_name: str
|
||||
tp_rank: int
|
||||
pcp_rank: int
|
||||
dcp_rank: int
|
||||
pp_rank: int
|
||||
|
||||
|
||||
@dataclass(order=True)
|
||||
class PoolKey:
|
||||
"""Key for addressing KV cache blocks in the distributed store."""
|
||||
|
||||
key_metadata: KeyMetadata
|
||||
chunk_hash: str
|
||||
|
||||
def __hash__(self):
|
||||
return hash(
|
||||
(
|
||||
self.key_metadata.model_name,
|
||||
self.key_metadata.tp_rank,
|
||||
self.key_metadata.pcp_rank,
|
||||
self.key_metadata.dcp_rank,
|
||||
self.key_metadata.pp_rank,
|
||||
self.chunk_hash,
|
||||
)
|
||||
)
|
||||
|
||||
def to_string(self) -> str:
|
||||
return (
|
||||
f"{self.key_metadata.model_name}"
|
||||
f"@tp_rank:{self.key_metadata.tp_rank}"
|
||||
f"@pcp{self.key_metadata.pcp_rank}"
|
||||
f"@dcp{self.key_metadata.dcp_rank}"
|
||||
f"@pp_rank:{self.key_metadata.pp_rank}"
|
||||
f"@{self.chunk_hash}"
|
||||
)
|
||||
|
||||
|
||||
class ChunkedTokenDatabase:
|
||||
"""Maps token positions to store keys and GPU memory addresses."""
|
||||
|
||||
def __init__(self, metadata: KeyMetadata, block_size: int):
|
||||
self.metadata = metadata
|
||||
self.block_size = block_size
|
||||
self.kv_caches_base_addr: list[int] = []
|
||||
self.block_len: list[int] = []
|
||||
|
||||
def _make_key_by_hash(self, chunk_hash: str) -> PoolKey:
|
||||
return PoolKey(self.metadata, chunk_hash)
|
||||
|
||||
def set_kv_caches_base_addr(self, kv_caches_base_addr: list[int]):
|
||||
self.kv_caches_base_addr = kv_caches_base_addr
|
||||
|
||||
def set_block_len(self, block_len: list[int]):
|
||||
for length in block_len:
|
||||
if length % self.block_size != 0:
|
||||
raise ValueError(f"block_len {length} % {self.block_size} != 0")
|
||||
self.block_len = block_len
|
||||
|
||||
def prepare_value(
|
||||
self, start: int, end: int, block_ids: list[int]
|
||||
) -> tuple[list[int], list[int], int]:
|
||||
"""Compute memory addresses and sizes for a token range.
|
||||
|
||||
Returns:
|
||||
(addr_list, size_list, block_id)
|
||||
"""
|
||||
addr_list = []
|
||||
size_list = []
|
||||
block_id = block_ids[start // self.block_size]
|
||||
length = len(self.block_len)
|
||||
for index, base_addr in enumerate(self.kv_caches_base_addr):
|
||||
addr = base_addr + block_id * self.block_len[index % length]
|
||||
size = self.block_len[index % length] // self.block_size * (end - start)
|
||||
addr_list.append(addr)
|
||||
size_list.append(size)
|
||||
return addr_list, size_list, block_id
|
||||
|
||||
def process_tokens(
|
||||
self,
|
||||
token_len: int,
|
||||
block_hashes: list[BlockHash] | list[str],
|
||||
mask_num: int = 0,
|
||||
) -> Iterable[tuple[int, int, PoolKey]]:
|
||||
"""Process tokens and yield (start_idx, end_idx, pool_key) tuples.
|
||||
|
||||
Args:
|
||||
token_len: Total number of tokens.
|
||||
block_hashes: Block hashes for each block.
|
||||
mask_num: Number of tokens to skip from the beginning.
|
||||
"""
|
||||
if not block_hashes:
|
||||
return
|
||||
if not isinstance(block_hashes[0], str):
|
||||
block_hashes = [
|
||||
h.hex() # type: ignore[union-attr]
|
||||
for h in block_hashes
|
||||
]
|
||||
for chunk_id, hash_val in enumerate(block_hashes):
|
||||
start_idx = chunk_id * self.block_size
|
||||
if start_idx >= token_len:
|
||||
break
|
||||
end_idx = min(start_idx + self.block_size, token_len)
|
||||
if start_idx < mask_num:
|
||||
continue
|
||||
else:
|
||||
yield (
|
||||
start_idx,
|
||||
end_idx,
|
||||
self._make_key_by_hash(
|
||||
hash_val # type: ignore[arg-type]
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoadSpec:
|
||||
"""Specification for loading KV cache from external store."""
|
||||
|
||||
vllm_cached_tokens: int
|
||||
kvpool_cached_tokens: int
|
||||
can_load: bool
|
||||
token_len: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestTracker:
|
||||
"""Tracks per-request state across scheduler ticks."""
|
||||
|
||||
req_id: str
|
||||
token_len: int
|
||||
allocated_block_ids: list[int]
|
||||
num_saved_tokens: int = 0
|
||||
token_ids: list[int] | None = None
|
||||
# Snapshot of the prefill range length at tracker creation time.
|
||||
# For a fresh request this is len(prompt). For a resumed-from-preemption
|
||||
# request it includes previously-generated tokens, which are re-prefilled.
|
||||
prefill_end_tokens: int = 0
|
||||
|
||||
def update(
|
||||
self,
|
||||
new_block_ids: tuple[list[int], ...] | list[int],
|
||||
) -> None:
|
||||
if len(new_block_ids) == 0:
|
||||
new_block_ids = []
|
||||
elif isinstance(new_block_ids, tuple):
|
||||
new_block_ids = new_block_ids[0]
|
||||
elif isinstance(new_block_ids, list):
|
||||
pass
|
||||
else:
|
||||
raise ValueError(f"Unsupported new_block_ids type {type(new_block_ids)}")
|
||||
self.allocated_block_ids.extend(new_block_ids)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReqMeta:
|
||||
"""Per-request metadata for store put/get operations."""
|
||||
|
||||
req_id: str
|
||||
token_len_chunk: int
|
||||
block_ids: list[int]
|
||||
block_hashes: list[BlockHash]
|
||||
|
||||
can_save: bool | None = None
|
||||
load_spec: LoadSpec | None = None
|
||||
is_last_chunk: bool | None = None
|
||||
current_event: torch.cuda.Event | None = None
|
||||
|
||||
token_ids: list[int] | None = None
|
||||
original_block_size: int | None = None
|
||||
|
||||
@staticmethod
|
||||
def from_request_tracker(
|
||||
tracker: RequestTracker,
|
||||
block_size: int,
|
||||
load_spec: LoadSpec | None = None,
|
||||
skip_save: bool | None = False,
|
||||
block_hashes: list[BlockHash] | None = None,
|
||||
is_last_chunk: bool | None = None,
|
||||
discard_partial_chunks: bool = True,
|
||||
original_block_size: int | None = None,
|
||||
) -> "ReqMeta | None":
|
||||
"""Create ReqMeta from a RequestTracker."""
|
||||
if block_hashes is None:
|
||||
block_hashes = []
|
||||
input_token_len = tracker.token_len
|
||||
|
||||
chunk_boundary = (
|
||||
cdiv(tracker.num_saved_tokens + 1, block_size) * block_size
|
||||
if discard_partial_chunks
|
||||
else 0
|
||||
)
|
||||
num_tokens_to_save = (
|
||||
(input_token_len // block_size * block_size)
|
||||
if discard_partial_chunks
|
||||
else input_token_len
|
||||
)
|
||||
|
||||
skip_save = skip_save or num_tokens_to_save < chunk_boundary
|
||||
if skip_save and load_spec is None:
|
||||
return None
|
||||
|
||||
if not skip_save:
|
||||
tracker.num_saved_tokens = num_tokens_to_save
|
||||
|
||||
token_ids = None
|
||||
if tracker.token_ids:
|
||||
token_ids = tracker.token_ids
|
||||
|
||||
if load_spec is not None and load_spec.can_load:
|
||||
logger.debug(
|
||||
"Scheduled to load %d tokens for request %s",
|
||||
load_spec.kvpool_cached_tokens,
|
||||
tracker.req_id,
|
||||
)
|
||||
else:
|
||||
load_spec = None
|
||||
|
||||
logger.debug(
|
||||
"request:%s, meta save spec:%s, meta load spec:%s",
|
||||
tracker.req_id,
|
||||
not skip_save,
|
||||
load_spec,
|
||||
)
|
||||
return ReqMeta(
|
||||
req_id=tracker.req_id,
|
||||
token_len_chunk=num_tokens_to_save,
|
||||
block_ids=tracker.allocated_block_ids,
|
||||
can_save=not skip_save,
|
||||
load_spec=load_spec,
|
||||
block_hashes=block_hashes,
|
||||
is_last_chunk=is_last_chunk,
|
||||
token_ids=token_ids,
|
||||
original_block_size=original_block_size,
|
||||
)
|
||||
|
||||
|
||||
class MooncakeStoreConnectorMetadata(KVConnectorMetadata):
|
||||
"""Metadata passed from scheduler to worker."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
unfinished_request_ids: set[str],
|
||||
preempted_req_ids: set[str],
|
||||
):
|
||||
self.requests: list[ReqMeta] = []
|
||||
self.unfinished_request_ids = unfinished_request_ids
|
||||
self.preempted_req_ids = preempted_req_ids
|
||||
|
||||
def add_request(self, req_meta: ReqMeta) -> None:
|
||||
self.requests.append(req_meta)
|
||||
@@ -0,0 +1,380 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# Adapted from vllm-project/vllm-ascend
|
||||
# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/).
|
||||
"""Scheduler-side logic for MooncakeStoreConnector."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
|
||||
KVConnectorMetadata,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501
|
||||
LoadSpec,
|
||||
MooncakeStoreConnectorMetadata,
|
||||
ReqMeta,
|
||||
RequestTracker,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.worker import ( # noqa: E501
|
||||
LookupKeyClient,
|
||||
)
|
||||
from vllm.logger import init_logger
|
||||
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
|
||||
from vllm.v1.core.sched.output import NewRequestData, SchedulerOutput
|
||||
from vllm.v1.request import Request
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _new_req_prefill_tokens(request: NewRequestData) -> list[int]:
|
||||
"""Tokens this prefill will compute KV for.
|
||||
|
||||
Under the v2 model runner, resumed-from-preemption requests appear in
|
||||
``scheduled_new_reqs`` with ``prefill_token_ids`` set to the request's full
|
||||
token list (prompt + previously-generated). For all other cases this falls
|
||||
back to the original prompt.
|
||||
"""
|
||||
if request.prefill_token_ids is not None:
|
||||
return request.prefill_token_ids
|
||||
assert request.prompt_token_ids is not None
|
||||
return request.prompt_token_ids
|
||||
|
||||
|
||||
class MooncakeStoreScheduler:
|
||||
"""Scheduler-side component for MooncakeStoreConnector."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig):
|
||||
assert vllm_config.kv_transfer_config is not None
|
||||
self.kv_role = vllm_config.kv_transfer_config.kv_role
|
||||
self.load_async = vllm_config.kv_transfer_config.kv_connector_extra_config.get(
|
||||
"load_async", True
|
||||
)
|
||||
self.client = LookupKeyClient(vllm_config)
|
||||
|
||||
self.pcp_size = vllm_config.parallel_config.prefill_context_parallel_size
|
||||
self.dcp_size = vllm_config.parallel_config.decode_context_parallel_size
|
||||
self.original_block_size = vllm_config.cache_config.block_size
|
||||
self._block_size = vllm_config.cache_config.block_size
|
||||
if self.pcp_size > 1:
|
||||
self._block_size *= self.pcp_size
|
||||
if self.dcp_size > 1:
|
||||
self._block_size *= self.dcp_size
|
||||
|
||||
self._discard_partial_chunks = (
|
||||
vllm_config.kv_transfer_config.get_from_extra_config(
|
||||
"discard_partial_chunks", True
|
||||
)
|
||||
)
|
||||
|
||||
# Per-request state
|
||||
self.load_specs: dict[str, LoadSpec] = {} # to be loaded
|
||||
self._request_trackers: dict[str, RequestTracker] = {} # scheduled new requests
|
||||
self._preempted_req_ids: set[str] = set() # preempted requests
|
||||
self._unfinished_requests: dict[str, tuple[Request, list[int]]] = {}
|
||||
self._unfinished_request_ids: set[str] = set()
|
||||
|
||||
def get_num_new_matched_tokens(
|
||||
self,
|
||||
request: Request,
|
||||
num_computed_tokens: int,
|
||||
) -> tuple[int, bool]:
|
||||
"""Check for external KV cache hit."""
|
||||
# Look up against the full prefill range, not just the prompt.
|
||||
if self._discard_partial_chunks:
|
||||
token_len = request.num_tokens // self._block_size * self._block_size
|
||||
else:
|
||||
token_len = request.num_tokens
|
||||
|
||||
if token_len < self._block_size:
|
||||
return 0, False
|
||||
|
||||
num_external_hit_tokens = self.client.lookup(token_len, request.block_hashes)
|
||||
|
||||
if num_external_hit_tokens == request.num_tokens:
|
||||
num_external_hit_tokens -= 1
|
||||
|
||||
if num_external_hit_tokens < num_computed_tokens:
|
||||
need_to_allocate = 0
|
||||
else:
|
||||
need_to_allocate = num_external_hit_tokens - num_computed_tokens
|
||||
|
||||
logger.debug(
|
||||
"Reqid: %s, Total tokens %d, kvpool hit tokens: %d, need to load: %d",
|
||||
request.request_id,
|
||||
request.num_tokens,
|
||||
num_external_hit_tokens,
|
||||
need_to_allocate,
|
||||
)
|
||||
|
||||
if need_to_allocate <= 0:
|
||||
return 0, False
|
||||
|
||||
self.load_specs[request.request_id] = LoadSpec(
|
||||
vllm_cached_tokens=num_computed_tokens,
|
||||
kvpool_cached_tokens=num_external_hit_tokens,
|
||||
can_load=False,
|
||||
)
|
||||
|
||||
return need_to_allocate, self.load_async
|
||||
|
||||
def update_state_after_alloc(
|
||||
self,
|
||||
request: Request,
|
||||
blocks: KVCacheBlocks,
|
||||
num_external_tokens: int,
|
||||
):
|
||||
"""Update state after block allocation."""
|
||||
local_block_ids: list[int] = []
|
||||
if num_external_tokens > 0:
|
||||
local_block_ids = blocks.get_block_ids()[0]
|
||||
|
||||
self._unfinished_requests[request.request_id] = (request, local_block_ids)
|
||||
self._unfinished_request_ids.add(request.request_id)
|
||||
|
||||
if request.request_id not in self.load_specs:
|
||||
return
|
||||
|
||||
if num_external_tokens == 0:
|
||||
self.load_specs[request.request_id].can_load = False
|
||||
return
|
||||
|
||||
assert (
|
||||
num_external_tokens > 0
|
||||
and num_external_tokens
|
||||
== self.load_specs[request.request_id].kvpool_cached_tokens
|
||||
- self.load_specs[request.request_id].vllm_cached_tokens
|
||||
), (
|
||||
f"Mismatch in number of tokens: {num_external_tokens} vs "
|
||||
f"{self.load_specs[request.request_id].kvpool_cached_tokens} - "
|
||||
f"{self.load_specs[request.request_id].vllm_cached_tokens}"
|
||||
f" for request {request.request_id}"
|
||||
)
|
||||
|
||||
self.load_specs[request.request_id].can_load = True
|
||||
|
||||
def build_connector_meta(
|
||||
self, scheduler_output: SchedulerOutput
|
||||
) -> KVConnectorMetadata:
|
||||
"""Build connector metadata for this scheduler step."""
|
||||
force_skip_save = self.kv_role == "kv_consumer"
|
||||
|
||||
for finished_req_id in scheduler_output.finished_req_ids:
|
||||
self.load_specs.pop(finished_req_id, None)
|
||||
self._request_trackers.pop(finished_req_id, None)
|
||||
self._unfinished_requests.pop(finished_req_id, None)
|
||||
self._unfinished_request_ids.discard(finished_req_id)
|
||||
self._preempted_req_ids.discard(finished_req_id)
|
||||
|
||||
preempted_ids = scheduler_output.preempted_req_ids or set()
|
||||
self._preempted_req_ids.update(preempted_ids)
|
||||
for req_id in preempted_ids:
|
||||
self._request_trackers.pop(req_id, None)
|
||||
self._unfinished_requests.pop(req_id, None)
|
||||
|
||||
meta = MooncakeStoreConnectorMetadata(
|
||||
self._unfinished_request_ids,
|
||||
preempted_ids,
|
||||
)
|
||||
|
||||
# Handle new requests
|
||||
for request in scheduler_output.scheduled_new_reqs:
|
||||
load_spec = self.load_specs.pop(request.req_id, None)
|
||||
num_tokens_to_compute = (
|
||||
request.num_computed_tokens
|
||||
+ scheduler_output.num_scheduled_tokens[request.req_id]
|
||||
)
|
||||
assert request.req_id in self._unfinished_requests
|
||||
request_tuple = self._unfinished_requests.get(request.req_id)
|
||||
request_real = request_tuple[0] # type: ignore[index]
|
||||
|
||||
if not isinstance(request.block_ids[0], list):
|
||||
unfolded_block_ids = request.block_ids.copy()
|
||||
else:
|
||||
# TODO: support HMA
|
||||
unfolded_block_ids = request.block_ids[0].copy()
|
||||
|
||||
prefill_tokens = _new_req_prefill_tokens(request)
|
||||
request_tracker = RequestTracker(
|
||||
req_id=request.req_id,
|
||||
token_len=num_tokens_to_compute,
|
||||
allocated_block_ids=unfolded_block_ids,
|
||||
num_saved_tokens=0,
|
||||
token_ids=prefill_tokens[:num_tokens_to_compute],
|
||||
prefill_end_tokens=len(prefill_tokens),
|
||||
)
|
||||
self._request_trackers[request.req_id] = request_tracker
|
||||
|
||||
last_chunk_tokens_num = (
|
||||
(len(prefill_tokens) // self._block_size * self._block_size)
|
||||
if self._discard_partial_chunks
|
||||
else len(prefill_tokens)
|
||||
)
|
||||
|
||||
req_meta = ReqMeta.from_request_tracker(
|
||||
request_tracker,
|
||||
self._block_size,
|
||||
load_spec=load_spec,
|
||||
skip_save=force_skip_save,
|
||||
block_hashes=request_real.block_hashes,
|
||||
is_last_chunk=(request_tracker.token_len >= last_chunk_tokens_num),
|
||||
discard_partial_chunks=self._discard_partial_chunks,
|
||||
original_block_size=self.original_block_size,
|
||||
)
|
||||
if req_meta is not None:
|
||||
meta.add_request(req_meta)
|
||||
|
||||
# Handle cached (running, or MRV1 resumed-from-preemption) requests
|
||||
cached_reqs = scheduler_output.scheduled_cached_reqs
|
||||
if not force_skip_save:
|
||||
for i, req_id in enumerate(cached_reqs.req_ids):
|
||||
new_block_ids = cached_reqs.new_block_ids[i]
|
||||
if not new_block_ids:
|
||||
continue
|
||||
|
||||
req_meta = None
|
||||
if req_id in self._preempted_req_ids:
|
||||
# Resumed after preemption
|
||||
if isinstance(new_block_ids, tuple):
|
||||
block_ids_list = new_block_ids[0].copy()
|
||||
else:
|
||||
block_ids_list = new_block_ids.copy()
|
||||
self._preempted_req_ids.discard(req_id)
|
||||
load_spec = self.load_specs.pop(req_id, None)
|
||||
request_tuple = self._unfinished_requests.get(req_id)
|
||||
request_real = request_tuple[0] # type: ignore[index]
|
||||
num_tokens_to_compute = (
|
||||
request_real.num_computed_tokens
|
||||
+ scheduler_output.num_scheduled_tokens[req_id]
|
||||
)
|
||||
# On resume, the request re-prefills prompt + previously
|
||||
# generated tokens (all_token_ids).
|
||||
prefill_tokens = list(request_real.all_token_ids)
|
||||
request_tracker = RequestTracker(
|
||||
req_id=req_id,
|
||||
token_len=num_tokens_to_compute,
|
||||
allocated_block_ids=block_ids_list,
|
||||
num_saved_tokens=0,
|
||||
token_ids=prefill_tokens[:num_tokens_to_compute].copy(),
|
||||
prefill_end_tokens=len(prefill_tokens),
|
||||
)
|
||||
self._request_trackers[req_id] = request_tracker
|
||||
|
||||
last_chunk_tokens_num = (
|
||||
(len(prefill_tokens) // self._block_size * self._block_size)
|
||||
if self._discard_partial_chunks
|
||||
else len(prefill_tokens)
|
||||
)
|
||||
req_meta = ReqMeta.from_request_tracker(
|
||||
request_tracker,
|
||||
self._block_size,
|
||||
load_spec=load_spec,
|
||||
skip_save=force_skip_save,
|
||||
block_hashes=request_real.block_hashes,
|
||||
is_last_chunk=(
|
||||
request_tracker.token_len >= last_chunk_tokens_num
|
||||
),
|
||||
discard_partial_chunks=self._discard_partial_chunks,
|
||||
original_block_size=self.original_block_size,
|
||||
)
|
||||
else:
|
||||
# Decode/chunked request
|
||||
request_tracker = self._request_trackers[req_id]
|
||||
num_new_tokens = scheduler_output.num_scheduled_tokens[req_id]
|
||||
req_tuple = self._unfinished_requests.get(req_id)
|
||||
if req_tuple:
|
||||
unfinished_req = req_tuple[0]
|
||||
num_current_tokens = request_tracker.token_len
|
||||
new_token_ids = unfinished_req.all_token_ids[
|
||||
num_current_tokens : num_current_tokens + num_new_tokens
|
||||
]
|
||||
request_tracker.token_len += len(new_token_ids)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Request {req_id} is not in _unfinished_requests"
|
||||
)
|
||||
num_computed_token = cached_reqs.num_computed_tokens[i]
|
||||
# Use the tracker's snapshot of the prefill range so resumed
|
||||
# requests keep saving past the original prompt boundary.
|
||||
prefill_end = request_tracker.prefill_end_tokens
|
||||
if num_computed_token >= prefill_end:
|
||||
continue
|
||||
request_tracker.update(new_block_ids)
|
||||
|
||||
last_chunk_tokens_num = (
|
||||
(prefill_end // self._block_size * self._block_size)
|
||||
if self._discard_partial_chunks
|
||||
else prefill_end
|
||||
)
|
||||
req_meta = ReqMeta.from_request_tracker(
|
||||
request_tracker,
|
||||
self._block_size,
|
||||
load_spec=None,
|
||||
skip_save=force_skip_save,
|
||||
block_hashes=unfinished_req.block_hashes,
|
||||
is_last_chunk=(
|
||||
request_tracker.token_len >= last_chunk_tokens_num
|
||||
),
|
||||
discard_partial_chunks=self._discard_partial_chunks,
|
||||
original_block_size=self.original_block_size,
|
||||
)
|
||||
|
||||
if req_meta is not None:
|
||||
meta.add_request(req_meta)
|
||||
|
||||
# Handle requests with pending load specs not yet scheduled
|
||||
request_ids = [req.req_id for req in scheduler_output.scheduled_new_reqs]
|
||||
for request_id, (
|
||||
unfinished_req,
|
||||
block_ids,
|
||||
) in self._unfinished_requests.items():
|
||||
if request_id not in request_ids and request_id not in cached_reqs.req_ids:
|
||||
load_spec = self.load_specs.pop(request_id, None)
|
||||
if not load_spec:
|
||||
continue
|
||||
num_tokens_to_compute = load_spec.kvpool_cached_tokens
|
||||
if (num_tokens_to_compute % self._block_size != 0) and (
|
||||
num_tokens_to_compute == unfinished_req.num_tokens - 1
|
||||
):
|
||||
num_tokens_to_compute = num_tokens_to_compute + 1
|
||||
request_tracker = RequestTracker(
|
||||
req_id=request_id,
|
||||
token_len=num_tokens_to_compute,
|
||||
allocated_block_ids=block_ids,
|
||||
num_saved_tokens=0,
|
||||
)
|
||||
self._request_trackers[request_id] = request_tracker
|
||||
req_meta = ReqMeta.from_request_tracker(
|
||||
request_tracker,
|
||||
self._block_size,
|
||||
load_spec=load_spec,
|
||||
skip_save=None,
|
||||
block_hashes=unfinished_req.block_hashes,
|
||||
discard_partial_chunks=self._discard_partial_chunks,
|
||||
)
|
||||
if req_meta is not None:
|
||||
meta.add_request(req_meta)
|
||||
|
||||
return meta
|
||||
|
||||
def request_finished(
|
||||
self,
|
||||
request: Request,
|
||||
block_ids: list[int],
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
"""Determine whether to delay freeing blocks for async save."""
|
||||
if self.kv_role == "kv_consumer":
|
||||
return False, None
|
||||
tracker = self._request_trackers.get(request.request_id)
|
||||
assert tracker is not None
|
||||
if tracker.num_saved_tokens <= 0:
|
||||
return False, None
|
||||
delay_free_blocks = len(block_ids) > 0
|
||||
if delay_free_blocks:
|
||||
logger.debug(
|
||||
"Delaying free of %d blocks for request %s",
|
||||
len(block_ids),
|
||||
request.request_id,
|
||||
)
|
||||
return delay_free_blocks, None
|
||||
@@ -0,0 +1,979 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# The transfer-thread scaffolding (KVTransferThread, KVCacheStoreSendingThread,
|
||||
# KVCacheStoreRecvingThread) is adapted from vllm-project/vllm-ascend
|
||||
# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/).
|
||||
"""Worker-side logic for MooncakeStoreConnector.
|
||||
|
||||
Includes the store worker, transfer threads, lookup server,
|
||||
and MooncakeDistributedStore integration.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import regex as re
|
||||
import torch
|
||||
import zmq
|
||||
|
||||
import vllm.envs as envs
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.distributed import (
|
||||
get_dcp_group,
|
||||
get_pcp_group,
|
||||
get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
)
|
||||
from vllm.distributed.kv_events import BlockStored
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import (
|
||||
get_mooncake_dp_engine_index,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501
|
||||
ChunkedTokenDatabase,
|
||||
KeyMetadata,
|
||||
MooncakeStoreConnectorMetadata,
|
||||
ReqMeta,
|
||||
)
|
||||
from vllm.logger import init_logger
|
||||
from vllm.utils.network_utils import get_ip, make_zmq_socket
|
||||
from vllm.v1.core.kv_cache_utils import BlockHash, maybe_convert_block_hash
|
||||
from vllm.v1.serial_utils import MsgpackDecoder, MsgpackEncoder
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
DEFAULT_GLOBAL_SEGMENT_SIZE = 4 * 1024 * 1024 * 1024 # 4 GiB
|
||||
DEFAULT_LOCAL_BUFFER_SIZE = 4 * 1024 * 1024 * 1024 # 4 GiB
|
||||
MOONCAKE_NO_AVAILABLE_HANDLE = -200
|
||||
|
||||
|
||||
@dataclass
|
||||
class MooncakeStoreConfig:
|
||||
"""Configuration for MooncakeDistributedStore."""
|
||||
|
||||
metadata_server: str
|
||||
global_segment_size: int
|
||||
local_buffer_size: int
|
||||
protocol: str
|
||||
device_name: str
|
||||
master_server_address: str
|
||||
|
||||
@staticmethod
|
||||
def from_file(file_path: str) -> "MooncakeStoreConfig":
|
||||
with open(file_path) as file:
|
||||
config = json.load(file)
|
||||
return MooncakeStoreConfig(
|
||||
metadata_server=config.get("metadata_server", ""),
|
||||
global_segment_size=_parse_size(
|
||||
config.get("global_segment_size", DEFAULT_GLOBAL_SEGMENT_SIZE)
|
||||
),
|
||||
local_buffer_size=_parse_size(
|
||||
config.get("local_buffer_size", DEFAULT_LOCAL_BUFFER_SIZE)
|
||||
),
|
||||
protocol=config.get("protocol", "rdma"),
|
||||
device_name=config.get("device_name", ""),
|
||||
master_server_address=config.get("master_server_address", ""),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def load_from_env() -> "MooncakeStoreConfig":
|
||||
config_path = os.getenv("MOONCAKE_CONFIG_PATH")
|
||||
if not config_path:
|
||||
raise ValueError(
|
||||
"The environment variable 'MOONCAKE_CONFIG_PATH' is not set."
|
||||
)
|
||||
return MooncakeStoreConfig.from_file(config_path)
|
||||
|
||||
|
||||
def _parse_size(value: Any) -> int:
|
||||
"""Parse storage size strings with units: GB, MB, KB, B."""
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
if not isinstance(value, str):
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError) as e:
|
||||
raise TypeError(f"Unsupported type for size: {type(value)}") from e
|
||||
|
||||
cleaned = value.strip().lower()
|
||||
if not cleaned:
|
||||
raise ValueError("Size cannot be empty.")
|
||||
|
||||
unit_multipliers = {
|
||||
"gb": 1024**3,
|
||||
"mb": 1024**2,
|
||||
"kb": 1024,
|
||||
"b": 1,
|
||||
}
|
||||
match = re.match(r"^\s*([\d.]+)\s*(gb|mb|kb|b)?\s*$", cleaned)
|
||||
if not match:
|
||||
raise ValueError(f"Invalid format: '{value}'")
|
||||
|
||||
number_str = match.group(1)
|
||||
unit = match.group(2) or "b"
|
||||
multiplier = unit_multipliers[unit]
|
||||
|
||||
try:
|
||||
numeric_value = float(number_str)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"Invalid numeric value '{number_str}' in: '{value}'") from exc
|
||||
return int(numeric_value * multiplier)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Transfer Threads
|
||||
# ============================================================
|
||||
|
||||
|
||||
class KVTransferThread(threading.Thread):
|
||||
"""Base class for async KV cache transfer threads."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store: Any,
|
||||
token_database: ChunkedTokenDatabase,
|
||||
block_size: int,
|
||||
tp_rank: int,
|
||||
ready_event: threading.Event,
|
||||
name: str,
|
||||
):
|
||||
super().__init__(daemon=True, name=name)
|
||||
self.store = store
|
||||
self.ready_event = ready_event
|
||||
self.block_size = block_size
|
||||
self.tp_rank = tp_rank
|
||||
self.token_database = token_database
|
||||
self.done_task_lock = threading.Lock()
|
||||
self.request_queue: queue.Queue[Any] = queue.Queue()
|
||||
self.finished_requests: set[str] = set()
|
||||
self.kv_event_lock = threading.Lock()
|
||||
self.kv_events: list[BlockStored] = []
|
||||
|
||||
def add_request(self, request: ReqMeta) -> None:
|
||||
self.request_queue.put(request)
|
||||
|
||||
def get_and_clear_finished_requests(self) -> set[str]:
|
||||
with self.done_task_lock:
|
||||
finished = self.finished_requests.copy()
|
||||
self.finished_requests.clear()
|
||||
return finished
|
||||
|
||||
def set_finished_request(self, req_id: str):
|
||||
with self.done_task_lock:
|
||||
self.finished_requests.add(req_id)
|
||||
|
||||
def run(self):
|
||||
self.ready_event.set()
|
||||
while True:
|
||||
try:
|
||||
request_data = self.request_queue.get()
|
||||
if request_data is None:
|
||||
logger.warning("Received a None request!")
|
||||
self.request_queue.task_done()
|
||||
continue
|
||||
self._handle_request(request_data)
|
||||
except Exception as e:
|
||||
logger.error("Error in %s: %s", self.name, e)
|
||||
|
||||
def _handle_request(self, req_meta: Any):
|
||||
pass
|
||||
|
||||
def update_kv_event(self, events: list[BlockStored]):
|
||||
with self.kv_event_lock:
|
||||
self.kv_events.extend(events)
|
||||
|
||||
def get_kv_events(self) -> list[BlockStored]:
|
||||
with self.kv_event_lock:
|
||||
events = self.kv_events.copy()
|
||||
self.kv_events.clear()
|
||||
return events
|
||||
|
||||
|
||||
class KVCacheStoreSendingThread(KVTransferThread):
|
||||
"""Background thread for storing KV cache blocks to the store."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store: Any,
|
||||
token_database: ChunkedTokenDatabase,
|
||||
block_size: int,
|
||||
tp_rank: int,
|
||||
put_step: int,
|
||||
kv_role: str,
|
||||
ready_event: threading.Event,
|
||||
enable_kv_event: bool = False,
|
||||
):
|
||||
super().__init__(
|
||||
store,
|
||||
token_database,
|
||||
block_size,
|
||||
tp_rank,
|
||||
ready_event,
|
||||
name="KVCacheStoreSendingThread",
|
||||
)
|
||||
self.put_step = put_step
|
||||
self.kv_role = kv_role
|
||||
self.stored_requests: defaultdict[str, int] = defaultdict(int)
|
||||
self.enable_kv_event = enable_kv_event
|
||||
|
||||
# Pause store requests when CPU offloading is under pressure.
|
||||
self._store_pressure_active = False
|
||||
self._skip_store_requests: set[str] = set()
|
||||
|
||||
def add_stored_request(self, req_id: str):
|
||||
with self.done_task_lock:
|
||||
self.stored_requests[req_id] += 1
|
||||
|
||||
def dec_stored_request(self, req_id: str):
|
||||
with self.done_task_lock:
|
||||
if req_id in self.stored_requests:
|
||||
self.stored_requests[req_id] -= 1
|
||||
|
||||
def delete_finished_stored_request(self, req_id: str):
|
||||
with self.done_task_lock:
|
||||
if req_id in self.stored_requests:
|
||||
del self.stored_requests[req_id]
|
||||
self._skip_store_requests.discard(req_id)
|
||||
|
||||
def _should_skip_request(self, req_id: str) -> bool:
|
||||
with self.done_task_lock:
|
||||
return self._store_pressure_active and req_id in self._skip_store_requests
|
||||
|
||||
def _mark_request_skipped_for_pressure(self, req_id: str) -> bool:
|
||||
with self.done_task_lock:
|
||||
already_skipped = req_id in self._skip_store_requests
|
||||
self._store_pressure_active = True
|
||||
self._skip_store_requests.add(req_id)
|
||||
return already_skipped
|
||||
|
||||
def _clear_store_pressure(self) -> bool:
|
||||
with self.done_task_lock:
|
||||
if not self._store_pressure_active and not self._skip_store_requests:
|
||||
return False
|
||||
self._store_pressure_active = False
|
||||
self._skip_store_requests.clear()
|
||||
return True
|
||||
|
||||
def _handle_request(self, req_meta: ReqMeta):
|
||||
token_len = req_meta.token_len_chunk
|
||||
block_ids = req_meta.block_ids
|
||||
req_id = req_meta.req_id
|
||||
current_event = req_meta.current_event
|
||||
|
||||
if req_id not in self.stored_requests:
|
||||
self.request_queue.task_done()
|
||||
return
|
||||
if self._should_skip_request(req_id):
|
||||
logger.debug(
|
||||
"Skipping Mooncake store for request %s while CPU offloading "
|
||||
"is under pressure",
|
||||
req_id,
|
||||
)
|
||||
self.dec_stored_request(req_id)
|
||||
self.request_queue.task_done()
|
||||
return
|
||||
|
||||
starts = []
|
||||
ends = []
|
||||
keys = []
|
||||
block_hashes: list[BlockHash] = []
|
||||
for index, (start, end, key) in enumerate(
|
||||
self.token_database.process_tokens(token_len, req_meta.block_hashes)
|
||||
):
|
||||
starts.append(start)
|
||||
ends.append(end)
|
||||
keys.append(key.to_string())
|
||||
block_hashes.append(req_meta.block_hashes[index])
|
||||
|
||||
# Apply put_step striding for TP
|
||||
starts = starts[self.tp_rank % self.put_step :: self.put_step]
|
||||
ends = ends[self.tp_rank % self.put_step :: self.put_step]
|
||||
keys = keys[self.tp_rank % self.put_step :: self.put_step]
|
||||
block_hashes = block_hashes[self.tp_rank % self.put_step :: self.put_step]
|
||||
|
||||
if not keys:
|
||||
self.dec_stored_request(req_id)
|
||||
return
|
||||
|
||||
# Check which blocks already exist (dedup)
|
||||
exists_states = self.store.batch_is_exist(keys)
|
||||
missing_indices = [i for i, exists in enumerate(exists_states) if exists != 1]
|
||||
|
||||
if not missing_indices:
|
||||
self.dec_stored_request(req_id)
|
||||
return
|
||||
|
||||
starts = [starts[i] for i in missing_indices]
|
||||
ends = [ends[i] for i in missing_indices]
|
||||
keys = [keys[i] for i in missing_indices]
|
||||
block_hashes = [block_hashes[i] for i in missing_indices]
|
||||
|
||||
logger.debug(
|
||||
"Storing KV cache for %d out of %d blocks "
|
||||
"(missing_count=%d) for request %s",
|
||||
len(keys),
|
||||
token_len // self.block_size,
|
||||
len(missing_indices),
|
||||
req_id,
|
||||
)
|
||||
|
||||
addrs = []
|
||||
sizes = []
|
||||
stored_events: list[BlockStored] = []
|
||||
prev_key = None
|
||||
new_block_hashes = [maybe_convert_block_hash(bh) for bh in block_hashes]
|
||||
|
||||
for index, start in enumerate(starts):
|
||||
addr, size, _ = self.token_database.prepare_value(
|
||||
start, ends[index], block_ids
|
||||
)
|
||||
addrs.append(addr)
|
||||
sizes.append(size)
|
||||
|
||||
if self.enable_kv_event:
|
||||
token_ids = (
|
||||
req_meta.token_ids[start : ends[index]]
|
||||
if req_meta.token_ids is not None
|
||||
else None
|
||||
)
|
||||
stored_event = BlockStored(
|
||||
block_hashes=[new_block_hashes[index]],
|
||||
parent_block_hash=prev_key,
|
||||
token_ids=token_ids,
|
||||
block_size=req_meta.original_block_size,
|
||||
lora_id=None,
|
||||
medium="cpu",
|
||||
lora_name=None,
|
||||
)
|
||||
stored_events.append(stored_event)
|
||||
prev_key = new_block_hashes[index]
|
||||
|
||||
if current_event is not None:
|
||||
current_event.synchronize()
|
||||
|
||||
try:
|
||||
res = self.store.batch_put_from_multi_buffers(keys, addrs, sizes)
|
||||
failed = [i for i, v in enumerate(res) if v < 0]
|
||||
if failed:
|
||||
# Compute total bytes attempted for this batch
|
||||
total_bytes = sum(sum(s) if isinstance(s, list) else s for s in sizes)
|
||||
failed_codes = set(res[i] for i in failed)
|
||||
logger.warning(
|
||||
"batch_put failed: %d/%d keys failed "
|
||||
"(codes=%s, batch_bytes=%d, num_keys=%d), "
|
||||
"first_key=%s",
|
||||
len(failed),
|
||||
len(keys),
|
||||
failed_codes,
|
||||
total_bytes,
|
||||
len(keys),
|
||||
keys[0] if keys else "N/A",
|
||||
)
|
||||
if (
|
||||
MOONCAKE_NO_AVAILABLE_HANDLE in failed_codes
|
||||
and not self._mark_request_skipped_for_pressure(req_id)
|
||||
):
|
||||
logger.warning(
|
||||
"Detected Mooncake CPU offloading pressure "
|
||||
"(NO_AVAILABLE_HANDLE); skipping future store "
|
||||
"batches for request %s until a later store "
|
||||
"batch succeeds",
|
||||
req_id,
|
||||
)
|
||||
elif self._clear_store_pressure():
|
||||
logger.info(
|
||||
"Mooncake CPU offloading pressure cleared after a "
|
||||
"successful store batch"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("Failed to put key %s, error: %s", keys, e)
|
||||
|
||||
if self.enable_kv_event and stored_events:
|
||||
self.update_kv_event(stored_events)
|
||||
|
||||
self.dec_stored_request(req_id)
|
||||
self.request_queue.task_done()
|
||||
|
||||
|
||||
class KVCacheStoreRecvingThread(KVTransferThread):
|
||||
"""Background thread for loading KV cache blocks from the store."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store: Any,
|
||||
token_database: ChunkedTokenDatabase,
|
||||
block_size: int,
|
||||
tp_rank: int,
|
||||
ready_event: threading.Event,
|
||||
):
|
||||
super().__init__(
|
||||
store,
|
||||
token_database,
|
||||
block_size,
|
||||
tp_rank,
|
||||
ready_event,
|
||||
name="KVCacheStoreRecvingThread",
|
||||
)
|
||||
|
||||
def _handle_request(self, req_meta: ReqMeta):
|
||||
token_len = req_meta.load_spec.token_len # type: ignore[union-attr]
|
||||
req_id = req_meta.req_id
|
||||
mask_num = (
|
||||
req_meta.load_spec.vllm_cached_tokens # type: ignore[union-attr]
|
||||
// self.block_size
|
||||
* self.block_size
|
||||
)
|
||||
|
||||
addr_list = []
|
||||
size_list = []
|
||||
key_list = []
|
||||
for start, end, key in self.token_database.process_tokens(
|
||||
token_len, req_meta.block_hashes, mask_num
|
||||
):
|
||||
addr, size, _ = self.token_database.prepare_value(
|
||||
start, end, req_meta.block_ids
|
||||
)
|
||||
key_list.append(key.to_string())
|
||||
addr_list.append(addr)
|
||||
size_list.append(size)
|
||||
|
||||
# Rotate lists by tp_rank for load balancing
|
||||
key_list_c = (
|
||||
key_list[self.tp_rank % len(key_list) :]
|
||||
+ key_list[: self.tp_rank % len(key_list)]
|
||||
)
|
||||
addr_list_c = (
|
||||
addr_list[self.tp_rank % len(addr_list) :]
|
||||
+ addr_list[: self.tp_rank % len(addr_list)]
|
||||
)
|
||||
size_list_c = (
|
||||
size_list[self.tp_rank % len(size_list) :]
|
||||
+ size_list[: self.tp_rank % len(size_list)]
|
||||
)
|
||||
|
||||
try:
|
||||
res = self.store.batch_get_into_multi_buffers(
|
||||
key_list_c, addr_list_c, size_list_c
|
||||
)
|
||||
failed = [
|
||||
(key, value)
|
||||
for key, value in zip(key_list_c, res, strict=True)
|
||||
if value < 0
|
||||
]
|
||||
if failed:
|
||||
logger.warning(
|
||||
"Failed to get %d Mooncake keys (batch_keys=%d, first_failures=%s)",
|
||||
len(failed),
|
||||
len(key_list_c),
|
||||
failed[:3],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to get Mooncake batch %s, error: %s",
|
||||
key_list_c[:3],
|
||||
e,
|
||||
)
|
||||
|
||||
self.set_finished_request(req_id)
|
||||
self.request_queue.task_done()
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Store Worker
|
||||
# ============================================================
|
||||
|
||||
|
||||
class MooncakeStoreWorker:
|
||||
"""Worker-side component for MooncakeStoreConnector."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig):
|
||||
try:
|
||||
from mooncake.store import MooncakeDistributedStore # type: ignore
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"Please install mooncake by following the instructions at "
|
||||
"https://github.com/kvcache-ai/Mooncake/blob/main/doc/"
|
||||
"en/build.md to run vLLM with MooncakeStoreConnector."
|
||||
) from e
|
||||
|
||||
model_config = vllm_config.model_config
|
||||
parallel_config = vllm_config.parallel_config
|
||||
|
||||
self.dp_rank = get_mooncake_dp_engine_index(parallel_config)
|
||||
self.tp_rank = get_tensor_model_parallel_rank()
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.pp_size = parallel_config.pipeline_parallel_size
|
||||
self.pp_rank = (parallel_config.rank // self.tp_size) % self.pp_size
|
||||
|
||||
self.pcp_size = get_pcp_group().world_size
|
||||
self.pcp_rank = get_pcp_group().rank_in_group if self.pcp_size > 1 else 0
|
||||
self.dcp_size = get_dcp_group().world_size
|
||||
self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_size > 1 else 0
|
||||
|
||||
assert vllm_config.kv_transfer_config is not None
|
||||
self.kv_role = vllm_config.kv_transfer_config.kv_role
|
||||
self.load_async = vllm_config.kv_transfer_config.kv_connector_extra_config.get(
|
||||
"load_async", True
|
||||
)
|
||||
self.cache_config = vllm_config.cache_config
|
||||
self.original_block_size = self.cache_config.block_size
|
||||
self.block_size = self.cache_config.block_size
|
||||
if self.pcp_size > 1:
|
||||
self.block_size *= self.pcp_size
|
||||
if self.dcp_size > 1:
|
||||
self.block_size *= self.dcp_size
|
||||
self.num_layers = model_config.get_num_layers(parallel_config)
|
||||
|
||||
self.use_mla = False
|
||||
if (
|
||||
hasattr(model_config, "use_mla")
|
||||
and isinstance(model_config.use_mla, bool)
|
||||
and model_config.use_mla
|
||||
):
|
||||
self.use_mla = True
|
||||
|
||||
if self.use_mla:
|
||||
self.num_kv_head = 1
|
||||
else:
|
||||
self.num_kv_head = model_config.get_total_num_kv_heads()
|
||||
|
||||
if self.num_kv_head < self.tp_size:
|
||||
self.put_step = self.tp_size // self.num_kv_head
|
||||
self.head_or_tp_rank = self.tp_rank // self.put_step
|
||||
else:
|
||||
self.head_or_tp_rank = self.tp_rank
|
||||
self.put_step = 1
|
||||
|
||||
self.metadata = KeyMetadata(
|
||||
model_name=model_config.model.rstrip("/").split("/")[-1],
|
||||
tp_rank=self.head_or_tp_rank,
|
||||
pcp_rank=self.pcp_rank,
|
||||
dcp_rank=self.dcp_rank,
|
||||
pp_rank=self.pp_rank,
|
||||
)
|
||||
|
||||
self.token_database = ChunkedTokenDatabase(self.metadata, self.block_size)
|
||||
|
||||
# Initialize MooncakeDistributedStore with its own TransferEngine
|
||||
store_config = MooncakeStoreConfig.load_from_env()
|
||||
self.store = MooncakeDistributedStore()
|
||||
|
||||
local_seg = get_ip()
|
||||
config_dict = {
|
||||
"local_hostname": local_seg,
|
||||
"metadata_server": store_config.metadata_server,
|
||||
"global_segment_size": str(store_config.global_segment_size),
|
||||
"local_buffer_size": str(store_config.local_buffer_size),
|
||||
"protocol": store_config.protocol,
|
||||
"rdma_devices": store_config.device_name,
|
||||
"master_server_addr": store_config.master_server_address,
|
||||
}
|
||||
ret = self.store.setup(config_dict)
|
||||
if ret != 0:
|
||||
msg = "Initialize MooncakeDistributedStore failed."
|
||||
logger.error(msg)
|
||||
raise RuntimeError(msg)
|
||||
|
||||
kv_event_config = vllm_config.kv_events_config
|
||||
self.enable_kv_events = False
|
||||
if kv_event_config and kv_event_config.enable_kv_cache_events:
|
||||
self.enable_kv_events = True
|
||||
|
||||
self.kv_send_thread: KVCacheStoreSendingThread | None = None
|
||||
self.kv_recv_thread: KVCacheStoreRecvingThread | None = None
|
||||
self.finished_store_req: set[str] = set()
|
||||
|
||||
# Start lookup server on rank 0 for scheduler-side prefix queries
|
||||
self.lookup_server: LookupKeyServer | None = None
|
||||
if vllm_config.parallel_config.rank == 0:
|
||||
self.lookup_server = LookupKeyServer(self, vllm_config)
|
||||
|
||||
def register_cross_layers_kv_caches(self, kv_cache: torch.Tensor) -> None:
|
||||
"""Register a cross-layers KV cache tensor.
|
||||
|
||||
Wraps the unified tensor in a single-entry dict so that the
|
||||
existing stride-based logic in register_kv_caches() produces
|
||||
the correct single-segment result (block_len = page_size * num_layers).
|
||||
"""
|
||||
self.register_kv_caches({"__cross_layer__": kv_cache})
|
||||
|
||||
def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
|
||||
"""Register KV cache tensors and start transfer threads."""
|
||||
# TODO(yifan): we haven't supported HMA yet.
|
||||
first_kv_cache = next(iter(kv_caches.values()))
|
||||
|
||||
# num_blocks from cache_config is authoritative (set after
|
||||
# profiling, before KV cache allocation).
|
||||
assert self.cache_config.num_gpu_blocks is not None
|
||||
self.num_blocks = self.cache_config.num_gpu_blocks
|
||||
|
||||
# Detect the KV cache memory layout using the stride-based
|
||||
# approach from simple_kv_offload/worker.py.
|
||||
#
|
||||
# The physical layout varies across attention backends:
|
||||
# FlashAttn/ROCm : (2, num_blocks, ...) → K/V outermost
|
||||
# FlashInfer/MLA : (num_blocks, ...) → blocks outermost
|
||||
#
|
||||
# We derive page_size_bytes = storage.nbytes() // num_blocks,
|
||||
# then classify dims: any dim whose byte-stride exceeds
|
||||
# page_size_bytes must be an outer segment dim (e.g. the K/V
|
||||
# dim of size 2). For those backends we register each segment
|
||||
# (K, V) as a separate base-address so that the per-block
|
||||
# offset arithmetic in prepare_value() stays correct.
|
||||
storage = first_kv_cache.untyped_storage()
|
||||
el = first_kv_cache.element_size()
|
||||
page_size_bytes = storage.nbytes() // self.num_blocks
|
||||
outer_dims = [
|
||||
d
|
||||
for d in range(first_kv_cache.ndim)
|
||||
if first_kv_cache.stride(d) * el > page_size_bytes
|
||||
]
|
||||
|
||||
# Register buffers with the store (deduplicate shared storages)
|
||||
# and record per-segment base addresses for every layer.
|
||||
seen_ptrs: set[int] = set()
|
||||
self.kv_caches_base_addr: list[int] = []
|
||||
self.block_len: list[int] = []
|
||||
|
||||
for cache in kv_caches.values():
|
||||
cache_storage = cache.untyped_storage()
|
||||
base_addr = cache_storage.data_ptr()
|
||||
region_len = cache_storage.nbytes()
|
||||
|
||||
if base_addr not in seen_ptrs:
|
||||
seen_ptrs.add(base_addr)
|
||||
ret = self.store.register_buffer(base_addr, region_len)
|
||||
if ret != 0:
|
||||
logger.error(
|
||||
"register_buffer failed for addr %#x len %d: %d",
|
||||
base_addr,
|
||||
region_len,
|
||||
ret,
|
||||
)
|
||||
|
||||
if not outer_dims:
|
||||
# Blocks-first layout (FlashInfer / MLA): one segment.
|
||||
self.kv_caches_base_addr.append(base_addr)
|
||||
self.block_len.append(page_size_bytes)
|
||||
else:
|
||||
# K/V-first layout (FlashAttn / ROCm): split segments.
|
||||
seg_stride = cache.stride(outer_dims[0]) * el
|
||||
for idx in range(cache.shape[outer_dims[0]]):
|
||||
self.kv_caches_base_addr.append(base_addr + idx * seg_stride)
|
||||
self.block_len.append(seg_stride // self.num_blocks)
|
||||
|
||||
logger.info(
|
||||
"Registering KV_Caches. use_mla: %s, shape %s, "
|
||||
"num_blocks: %d, block_len: %s, "
|
||||
"per_key_bytes: %d, "
|
||||
"num_segments: %d",
|
||||
self.use_mla,
|
||||
first_kv_cache.shape,
|
||||
self.num_blocks,
|
||||
list(set(self.block_len)),
|
||||
sum(self.block_len),
|
||||
len(self.kv_caches_base_addr),
|
||||
)
|
||||
|
||||
self.token_database.set_kv_caches_base_addr(self.kv_caches_base_addr)
|
||||
self.token_database.set_block_len(self.block_len)
|
||||
|
||||
# Start transfer threads
|
||||
if self.kv_role in ["kv_producer", "kv_both"]:
|
||||
ready_event_sending = threading.Event()
|
||||
self.kv_send_thread = KVCacheStoreSendingThread(
|
||||
self.store,
|
||||
self.token_database,
|
||||
self.block_size,
|
||||
self.tp_rank,
|
||||
self.put_step,
|
||||
self.kv_role,
|
||||
ready_event_sending,
|
||||
self.enable_kv_events,
|
||||
)
|
||||
self.kv_send_thread.start()
|
||||
|
||||
ready_event_recving = threading.Event()
|
||||
self.kv_recv_thread = KVCacheStoreRecvingThread(
|
||||
self.store,
|
||||
self.token_database,
|
||||
self.block_size,
|
||||
self.tp_rank,
|
||||
ready_event_recving,
|
||||
)
|
||||
self.kv_recv_thread.start()
|
||||
ready_event_recving.wait()
|
||||
|
||||
def start_load_kv(
|
||||
self,
|
||||
metadata: MooncakeStoreConnectorMetadata,
|
||||
):
|
||||
"""No-op: loads are issued in get_finished() for overlap."""
|
||||
pass
|
||||
|
||||
def wait_for_save(
|
||||
self,
|
||||
metadata: MooncakeStoreConnectorMetadata,
|
||||
):
|
||||
"""No-op: stores are issued in get_finished() for overlap."""
|
||||
pass
|
||||
|
||||
def get_finished(
|
||||
self,
|
||||
finished_req_ids: set[str],
|
||||
meta: MooncakeStoreConnectorMetadata,
|
||||
) -> tuple[set[str], set[str]]:
|
||||
"""Issue all I/O and get completed send/recv request IDs.
|
||||
|
||||
All load and store I/O requests are issued here (after model
|
||||
compute is launched on the compute stream) for better
|
||||
compute-I/O overlap.
|
||||
"""
|
||||
# Issue async loads
|
||||
for request in meta.requests:
|
||||
load_spec = request.load_spec
|
||||
if load_spec is None or not load_spec.can_load:
|
||||
continue
|
||||
|
||||
token_len = request.token_len_chunk
|
||||
if (load_spec.kvpool_cached_tokens % self.block_size != 0) and (
|
||||
load_spec.kvpool_cached_tokens == token_len - 1
|
||||
):
|
||||
token_len = load_spec.kvpool_cached_tokens + 1
|
||||
else:
|
||||
token_len = load_spec.kvpool_cached_tokens
|
||||
load_spec.token_len = token_len
|
||||
|
||||
assert self.kv_recv_thread is not None
|
||||
self.kv_recv_thread.add_request(request)
|
||||
|
||||
assert self.load_async, "load_async must be True for better performance."
|
||||
# Issue stores with CUDA event synchronization
|
||||
if self.kv_role in ["kv_producer", "kv_both"]:
|
||||
current_event = None
|
||||
for request in meta.requests:
|
||||
if request.can_save:
|
||||
current_event = torch.cuda.Event()
|
||||
current_event.record()
|
||||
break
|
||||
|
||||
for request in meta.requests:
|
||||
if not request.can_save:
|
||||
continue
|
||||
request.current_event = current_event
|
||||
assert self.kv_send_thread is not None
|
||||
self.kv_send_thread.add_stored_request(request.req_id)
|
||||
self.kv_send_thread.add_request(request)
|
||||
|
||||
# Check completion of previously queued transfers
|
||||
done_sending = (
|
||||
self._get_and_clear_finished_sending(finished_req_ids, meta)
|
||||
if self.kv_role in ["kv_producer", "kv_both"]
|
||||
else set()
|
||||
)
|
||||
|
||||
done_recving = (
|
||||
self.kv_recv_thread.get_and_clear_finished_requests()
|
||||
if self.load_async and self.kv_recv_thread is not None
|
||||
else set()
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"Completed send: %d, recv: %d, tp_rank: %d",
|
||||
len(done_sending),
|
||||
len(done_recving),
|
||||
self.tp_rank,
|
||||
)
|
||||
return done_sending, done_recving
|
||||
|
||||
def _get_and_clear_finished_sending(
|
||||
self,
|
||||
finished_req_ids: set[str],
|
||||
meta: MooncakeStoreConnectorMetadata,
|
||||
) -> set[str]:
|
||||
assert self.kv_send_thread is not None
|
||||
finished_sending: set[str] = set()
|
||||
|
||||
for req_id in meta.preempted_req_ids:
|
||||
self.kv_send_thread.delete_finished_stored_request(req_id)
|
||||
|
||||
for req_id in self.kv_send_thread.stored_requests.copy():
|
||||
if (
|
||||
self.kv_send_thread.stored_requests[req_id] == 0
|
||||
and req_id in self.finished_store_req
|
||||
):
|
||||
self.finished_store_req.remove(req_id)
|
||||
finished_sending.add(req_id)
|
||||
self.kv_send_thread.delete_finished_stored_request(req_id)
|
||||
|
||||
for req_id in finished_req_ids:
|
||||
req_remain_jobs = self.kv_send_thread.stored_requests.get(req_id)
|
||||
if req_remain_jobs == 0:
|
||||
finished_sending.add(req_id)
|
||||
self.kv_send_thread.delete_finished_stored_request(req_id)
|
||||
elif req_remain_jobs is not None:
|
||||
self.finished_store_req.add(req_id)
|
||||
|
||||
return finished_sending
|
||||
|
||||
def lookup(
|
||||
self,
|
||||
token_len: int,
|
||||
block_hashes: list[BlockHash],
|
||||
) -> int:
|
||||
"""Check how many prefix tokens exist in the store.
|
||||
|
||||
Checks across all TP ranks and PP ranks.
|
||||
"""
|
||||
end = 0
|
||||
keys: list[str] = []
|
||||
try:
|
||||
starts: list[int] = []
|
||||
for start, end, key in self.token_database.process_tokens(
|
||||
token_len, block_hashes
|
||||
):
|
||||
keys.append(key.to_string())
|
||||
starts.append(start)
|
||||
|
||||
# Expand keys for all TP ranks
|
||||
multi_tp_keys = keys[:]
|
||||
for i in range(1, min(self.tp_size, self.num_kv_head)):
|
||||
for item in keys:
|
||||
new_str = item.replace("@tp_rank:0", f"@tp_rank:{i}", 1)
|
||||
multi_tp_keys.append(new_str)
|
||||
|
||||
# Expand keys for all PP ranks
|
||||
pp_base_keys = multi_tp_keys.copy()
|
||||
for i in range(1, self.pp_size):
|
||||
for item in pp_base_keys:
|
||||
new_str = item.replace("@pp_rank:0", f"@pp_rank:{i}", 1)
|
||||
multi_tp_keys.append(new_str)
|
||||
|
||||
res = self.store.batch_is_exist(multi_tp_keys)
|
||||
|
||||
num_block = len(keys)
|
||||
multi_tp_values = [
|
||||
res[i * num_block : (i + 1) * num_block]
|
||||
for i in range(min(self.tp_size, self.num_kv_head) * self.pp_size)
|
||||
]
|
||||
index = self._find_min_first_non_one_index(multi_tp_values)
|
||||
if index != -1:
|
||||
return starts[index]
|
||||
except Exception as e:
|
||||
logger.error("Remote connection failed in lookup: %s", e)
|
||||
return 0
|
||||
return end
|
||||
|
||||
@staticmethod
|
||||
def _find_min_first_non_one_index(
|
||||
arr: list[list[int]],
|
||||
) -> int:
|
||||
try:
|
||||
return min(idx for row in arr for idx, val in enumerate(row) if val != 1)
|
||||
except ValueError:
|
||||
return -1
|
||||
|
||||
def get_kv_events(self) -> list[BlockStored]:
|
||||
if self.enable_kv_events and self.kv_send_thread is not None:
|
||||
return self.kv_send_thread.get_kv_events()
|
||||
return []
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Lookup Key Server
|
||||
# ============================================================
|
||||
|
||||
|
||||
class LookupKeyServer:
|
||||
"""ZMQ server on worker rank 0 for handling prefix lookup queries."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store_worker: MooncakeStoreWorker,
|
||||
vllm_config: VllmConfig,
|
||||
):
|
||||
self.decoder = MsgpackDecoder()
|
||||
self.ctx = zmq.Context() # type: ignore[attr-defined]
|
||||
socket_path = get_zmq_rpc_path_lookup(vllm_config)
|
||||
self._ipc_path = socket_path.removeprefix("ipc://")
|
||||
if os.path.exists(self._ipc_path):
|
||||
os.unlink(self._ipc_path)
|
||||
self.socket = make_zmq_socket(
|
||||
self.ctx,
|
||||
socket_path,
|
||||
zmq.REP, # type: ignore[attr-defined]
|
||||
bind=True,
|
||||
)
|
||||
|
||||
self.store_worker = store_worker
|
||||
self.running = True
|
||||
|
||||
def process_request():
|
||||
while self.running:
|
||||
all_frames = self.socket.recv_multipart(copy=False)
|
||||
token_len = int.from_bytes(all_frames[0], byteorder="big")
|
||||
hash_frames = all_frames[1:]
|
||||
hashes_str = self.decoder.decode(hash_frames)
|
||||
result = self.store_worker.lookup(token_len, hashes_str)
|
||||
response = result.to_bytes(4, "big")
|
||||
self.socket.send(response)
|
||||
|
||||
self.thread = threading.Thread(target=process_request, daemon=True)
|
||||
self.thread.start()
|
||||
|
||||
def close(self):
|
||||
self.socket.close(linger=0)
|
||||
if os.path.exists(self._ipc_path):
|
||||
os.unlink(self._ipc_path)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Lookup Key Client
|
||||
# ============================================================
|
||||
|
||||
|
||||
class LookupKeyClient:
|
||||
"""ZMQ client for querying prefix cache hits from worker."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig):
|
||||
self.encoder = MsgpackEncoder()
|
||||
self.ctx = zmq.Context() # type: ignore[attr-defined]
|
||||
socket_path = get_zmq_rpc_path_lookup(vllm_config)
|
||||
self.socket = make_zmq_socket(
|
||||
self.ctx,
|
||||
socket_path,
|
||||
zmq.REQ, # type: ignore[attr-defined]
|
||||
bind=False,
|
||||
)
|
||||
|
||||
def lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
|
||||
hash_strs = [h.hex() for h in block_hashes]
|
||||
hash_frames = self.encoder.encode(hash_strs)
|
||||
token_len_bytes = token_len.to_bytes(4, byteorder="big")
|
||||
all_frames = [token_len_bytes] + list(hash_frames)
|
||||
self.socket.send_multipart(all_frames, copy=False)
|
||||
resp = self.socket.recv()
|
||||
result = int.from_bytes(resp, "big")
|
||||
return result
|
||||
|
||||
def close(self):
|
||||
self.socket.close(linger=0)
|
||||
|
||||
|
||||
def get_zmq_rpc_path_lookup(vllm_config: VllmConfig) -> str:
|
||||
"""Construct IPC path for ZMQ lookup socket."""
|
||||
dp_rank = get_mooncake_dp_engine_index(vllm_config.parallel_config)
|
||||
base_url = envs.VLLM_RPC_BASE_PATH
|
||||
rpc_port = 0
|
||||
assert vllm_config.kv_transfer_config is not None
|
||||
extra_config = vllm_config.kv_transfer_config.kv_connector_extra_config
|
||||
if "lookup_rpc_port" in extra_config:
|
||||
rpc_port = extra_config["lookup_rpc_port"]
|
||||
uid = os.getuid()
|
||||
logger.debug("Base URL: %s, RPC Port: %s, UID: %s", base_url, rpc_port, uid)
|
||||
return f"ipc://{base_url}/lookup_rpc_port_{rpc_port}_uid{uid}_dp_rank{dp_rank}"
|
||||
@@ -291,7 +291,7 @@ class OffloadingConnectorScheduler:
|
||||
self.config.kv_group_configs, req_status.group_states
|
||||
):
|
||||
if group_config.sliding_window_size_in_blocks is None:
|
||||
self.manager.touch(group_state.offload_keys)
|
||||
self.manager.touch(group_state.offload_keys, req_status.req_context)
|
||||
else:
|
||||
# we aim to keep just blocks that are necessary to hit
|
||||
# the original request (+ decoded blocks)
|
||||
@@ -300,7 +300,10 @@ class OffloadingConnectorScheduler:
|
||||
group_state.num_hit_blocks
|
||||
- group_config.sliding_window_size_in_blocks,
|
||||
)
|
||||
self.manager.touch(group_state.offload_keys[blocks_to_skip:])
|
||||
self.manager.touch(
|
||||
group_state.offload_keys[blocks_to_skip:],
|
||||
req_status.req_context,
|
||||
)
|
||||
|
||||
def _lookup(self, req_status: RequestOffloadState) -> int | None:
|
||||
"""
|
||||
@@ -802,14 +805,13 @@ class OffloadingConnectorScheduler:
|
||||
continue
|
||||
assert job_status.pending_count == 0
|
||||
|
||||
req_status = self._req_status[job_status.req_id]
|
||||
if job_status.is_store:
|
||||
self.manager.complete_store(job_status.keys)
|
||||
self.manager.complete_store(job_status.keys, req_status.req_context)
|
||||
else:
|
||||
self.manager.complete_load(job_status.keys)
|
||||
self.manager.complete_load(job_status.keys, req_status.req_context)
|
||||
if self._blocks_being_loaded:
|
||||
self._blocks_being_loaded.difference_update(job_status.keys)
|
||||
|
||||
req_status = self._req_status[job_status.req_id]
|
||||
if self._block_id_to_pending_jobs:
|
||||
# Sliding window blocks are tracked from store creation
|
||||
# and must be cleaned up unconditionally.
|
||||
|
||||
@@ -13,6 +13,7 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
|
||||
from vllm.model_executor.layers.mamba.mamba_utils import is_conv_state_dim_first
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
from vllm.v1.kv_cache_interface import MambaSpec
|
||||
|
||||
|
||||
@@ -103,7 +104,7 @@ def derive_mamba_conv_split(
|
||||
MambaConvSplitInfo with per-rank x_local, b_local, conv_rows,
|
||||
conv_dtype_size, and ssm_sizes (conv_state_bytes, ssm_state_bytes).
|
||||
"""
|
||||
if mamba_spec.mamba_type != "mamba2":
|
||||
if mamba_spec.mamba_type != MambaAttentionBackendEnum.MAMBA2:
|
||||
raise NotImplementedError(
|
||||
f"3-read conv transfer only supports Mamba2 models, "
|
||||
f"got mamba_type={mamba_spec.mamba_type!r}. "
|
||||
|
||||
@@ -355,7 +355,11 @@ def _compute_kwargs(cls: ConfigType) -> dict[str, dict[str, Any]]:
|
||||
if name == "max_model_len":
|
||||
kwargs[name]["type"] = human_readable_int_or_auto
|
||||
kwargs[name]["help"] += f"\n\n{human_readable_int_or_auto.__doc__}"
|
||||
elif name in ("max_num_batched_tokens", "kv_cache_memory_bytes"):
|
||||
elif name in (
|
||||
"max_num_batched_tokens",
|
||||
"kv_cache_memory_bytes",
|
||||
"safetensors_prefetch_block_size",
|
||||
):
|
||||
kwargs[name]["type"] = human_readable_int
|
||||
kwargs[name]["help"] += f"\n\n{human_readable_int.__doc__}"
|
||||
else:
|
||||
@@ -424,6 +428,8 @@ class EngineArgs:
|
||||
allowed_media_domains: list[str] | None = ModelConfig.allowed_media_domains
|
||||
download_dir: str | None = LoadConfig.download_dir
|
||||
safetensors_load_strategy: str | None = LoadConfig.safetensors_load_strategy
|
||||
safetensors_prefetch_num_threads: int = LoadConfig.safetensors_prefetch_num_threads
|
||||
safetensors_prefetch_block_size: int = LoadConfig.safetensors_prefetch_block_size
|
||||
load_format: str | LoadFormats = LoadConfig.load_format
|
||||
config_format: str = ModelConfig.config_format
|
||||
dtype: ModelDType = ModelConfig.dtype
|
||||
@@ -844,6 +850,14 @@ class EngineArgs:
|
||||
load_group.add_argument(
|
||||
"--safetensors-load-strategy", **load_kwargs["safetensors_load_strategy"]
|
||||
)
|
||||
load_group.add_argument(
|
||||
"--safetensors-prefetch-num-threads",
|
||||
**load_kwargs["safetensors_prefetch_num_threads"],
|
||||
)
|
||||
load_group.add_argument(
|
||||
"--safetensors-prefetch-block-size",
|
||||
**load_kwargs["safetensors_prefetch_block_size"],
|
||||
)
|
||||
load_group.add_argument(
|
||||
"--model-loader-extra-config", **load_kwargs["model_loader_extra_config"]
|
||||
)
|
||||
@@ -1584,6 +1598,8 @@ class EngineArgs:
|
||||
load_format=self.load_format,
|
||||
download_dir=self.download_dir,
|
||||
safetensors_load_strategy=self.safetensors_load_strategy,
|
||||
safetensors_prefetch_num_threads=self.safetensors_prefetch_num_threads,
|
||||
safetensors_prefetch_block_size=self.safetensors_prefetch_block_size,
|
||||
model_loader_extra_config=self.model_loader_extra_config,
|
||||
ignore_patterns=self.ignore_patterns,
|
||||
use_tqdm_on_load=self.use_tqdm_on_load,
|
||||
|
||||
@@ -111,6 +111,9 @@ class ChatCompletionResponse(OpenAIBaseModel):
|
||||
# vLLM-specific fields that are not in OpenAI spec
|
||||
prompt_logprobs: list[dict[int, Logprob] | None] | None = None
|
||||
prompt_token_ids: list[int] | None = None
|
||||
# Rendered prompt text from chat templating (only set when
|
||||
# ``return_prompt_text=True`` on the request).
|
||||
prompt_text: str | None = None
|
||||
kv_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="KVTransfer parameters."
|
||||
)
|
||||
@@ -138,6 +141,9 @@ class ChatCompletionStreamResponse(OpenAIBaseModel):
|
||||
system_fingerprint: str | None = None
|
||||
# not part of the OpenAI spec but for tracing the tokens
|
||||
prompt_token_ids: list[int] | None = None
|
||||
# Rendered prompt text from chat templating (only set when
|
||||
# ``return_prompt_text=True`` on the request); only sent on the first chunk.
|
||||
prompt_text: str | None = None
|
||||
|
||||
|
||||
class ChatCompletionToolsParam(OpenAIBaseModel):
|
||||
@@ -352,6 +358,15 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
"need to map generated text back to input tokens."
|
||||
),
|
||||
)
|
||||
return_prompt_text: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"If true, the response will include ``prompt_text`` containing the "
|
||||
"prompt string produced by chat templating. In streaming mode it "
|
||||
"is sent only on the first chunk. This is useful for inspecting "
|
||||
"exactly what was fed into the model."
|
||||
),
|
||||
)
|
||||
|
||||
cache_salt: str | None = Field(
|
||||
default=None,
|
||||
|
||||
@@ -508,6 +508,9 @@ class OpenAIServingChat(OpenAIServing):
|
||||
# the role
|
||||
role = self.get_chat_request_role(request)
|
||||
|
||||
# ``res.prompt`` is the rendered chat-templated prompt
|
||||
prompt_text = res.prompt if request.return_prompt_text else None
|
||||
|
||||
# NOTE num_choices defaults to 1 so this usually executes
|
||||
# once per request
|
||||
for i in range(num_choices):
|
||||
@@ -533,6 +536,7 @@ class OpenAIServingChat(OpenAIServing):
|
||||
if request.return_token_ids
|
||||
else None
|
||||
),
|
||||
prompt_text=prompt_text,
|
||||
)
|
||||
|
||||
# if continuous usage stats are requested, add it
|
||||
@@ -1371,6 +1375,9 @@ class OpenAIServingChat(OpenAIServing):
|
||||
if final_res.prompt_routed_experts is not None:
|
||||
prompt_routed_experts = final_res.prompt_routed_experts.tolist()
|
||||
|
||||
# ``final_res.prompt`` is the rendered chat-templated prompt text
|
||||
prompt_text = final_res.prompt if request.return_prompt_text else None
|
||||
|
||||
response = ChatCompletionResponse(
|
||||
id=request_id,
|
||||
created=created_time,
|
||||
@@ -1382,6 +1389,7 @@ class OpenAIServingChat(OpenAIServing):
|
||||
prompt_token_ids=(
|
||||
final_res.prompt_token_ids if request.return_token_ids else None
|
||||
),
|
||||
prompt_text=prompt_text,
|
||||
kv_transfer_params=final_res.kv_transfer_params,
|
||||
prompt_routed_experts=prompt_routed_experts,
|
||||
)
|
||||
|
||||
@@ -195,9 +195,9 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
|
||||
# There are two LoRA layers
|
||||
# the output_sizes in MergedColumnParallelLinear is not sharded by tp
|
||||
# we need to divide it by the tp_size to get correct slices size
|
||||
output_sizes = self.base_layer.output_sizes
|
||||
self.output_sizes = self.base_layer.output_sizes
|
||||
self.output_slices = tuple(
|
||||
divide(output_size, self.tp_size) for output_size in output_sizes
|
||||
divide(output_size, self.tp_size) for output_size in self.output_sizes
|
||||
)
|
||||
self.n_slices = len(self.output_slices)
|
||||
self.output_ids = (self.tp_rank,) * self.n_slices
|
||||
@@ -261,6 +261,42 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
|
||||
]
|
||||
return sliced_lora_b
|
||||
|
||||
def expand_packed_lora(
|
||||
self,
|
||||
lora_a: list[torch.Tensor],
|
||||
lora_b: list[torch.Tensor],
|
||||
) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
||||
"""
|
||||
Expand packed adapter groups when they don't match n_slices.
|
||||
E.g. in_proj_qkv (covers Q+K+V) + in_proj_z
|
||||
"""
|
||||
expanded_a: list[torch.Tensor] = []
|
||||
expanded_b: list[torch.Tensor] = []
|
||||
start_idx = 0
|
||||
for a_i, b_i in zip(lora_a, lora_b):
|
||||
# Determine which output slices this b_i covers.
|
||||
b_rows, cu_rows, covered = b_i.shape[0], 0, 0
|
||||
for i in range(start_idx, self.n_slices):
|
||||
cu_rows += self.output_sizes[i]
|
||||
if cu_rows == b_rows:
|
||||
covered = i - start_idx + 1
|
||||
break
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Cannot determine how to split lora_b with {b_rows} rows "
|
||||
f"into {self.n_slices} slices with output sizes "
|
||||
f"{self.output_sizes} starting from index {start_idx}."
|
||||
)
|
||||
# Split b_i into per-slice tensors and replicate a_i for each.
|
||||
start = 0
|
||||
for j in range(covered):
|
||||
size = self.output_sizes[start_idx + j]
|
||||
expanded_b.append(b_i[start : start + size, :])
|
||||
expanded_a.append(a_i)
|
||||
start += size
|
||||
start_idx += covered
|
||||
return expanded_a, expanded_b
|
||||
|
||||
def set_lora(
|
||||
self,
|
||||
index: int,
|
||||
@@ -269,6 +305,12 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
|
||||
):
|
||||
self.reset_lora(index)
|
||||
|
||||
# Expand packed adapter groups when they don't match n_slices.
|
||||
# E.g. in_proj_qkv (covers Q+K+V) + in_proj_z as 2 groups for a
|
||||
# 4-slice layer: split b_qkv by output_sizes and replicate a_qkv.
|
||||
if isinstance(lora_b, list) and len(lora_b) != self.n_slices:
|
||||
lora_a, lora_b = self.expand_packed_lora(lora_a, lora_b)
|
||||
|
||||
if self.tp_size > 1:
|
||||
lora_a = self.slice_lora_a(lora_a)
|
||||
lora_b = self.slice_lora_b(lora_b)
|
||||
@@ -497,18 +539,14 @@ class MergedColumnParallelLinearWithShardedLoRA(MergedColumnParallelLinearWithLo
|
||||
def slice_lora_a(
|
||||
self, lora_a: list[torch.Tensor | None]
|
||||
) -> list[torch.Tensor | None]:
|
||||
# NOTE: lora_a contains 2 subloras, and each sublora could be None.
|
||||
output_shard_size = self.lora_a_stacked[0].shape[2]
|
||||
output_start_idx = self.tp_rank * output_shard_size
|
||||
lora_a = [
|
||||
lora_a[0][output_start_idx : output_start_idx + output_shard_size, :]
|
||||
if lora_a[0] is not None
|
||||
else None,
|
||||
lora_a[1][output_start_idx : output_start_idx + output_shard_size, :]
|
||||
if lora_a[1] is not None
|
||||
else None,
|
||||
return [
|
||||
lora_a_i[output_start_idx : output_start_idx + output_shard_size, :]
|
||||
if (lora_a_i := lora_a[i]) is not None
|
||||
else None
|
||||
for i in range(len(lora_a))
|
||||
]
|
||||
return lora_a
|
||||
|
||||
def apply(self, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
return _mcp_apply(x, bias, self)
|
||||
|
||||
@@ -563,11 +563,16 @@ class LoRAModelManager:
|
||||
else:
|
||||
parts = module_name.split(".")
|
||||
replacements = self.packed_modules_mapping[parts[-1]]
|
||||
n_slices = getattr(module, "n_slices", len(replacements))
|
||||
if module.__class__.__name__ == "FusedMoEWithLoRA":
|
||||
replacements = replacements[
|
||||
: len(module.lora_a_stacked) // self.lora_slots
|
||||
]
|
||||
subloras: list[LoRALayerWeights | None] = []
|
||||
# HACK: overrides replacements for qkvz = qkv + z case.
|
||||
# Any better methods to handle this case?
|
||||
if n_slices != len(replacements):
|
||||
replacements = [f"slice_{i}" for i in range(n_slices)]
|
||||
for i, r in enumerate(replacements):
|
||||
lora = LoRALayerWeights.create_dummy_lora_weights(
|
||||
module_name + "." + r,
|
||||
|
||||
@@ -85,6 +85,13 @@ if HAS_TRITON:
|
||||
from vllm.model_executor.layers.fused_moe.experts.deep_gemm_moe import (
|
||||
DeepGemmExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
AiterExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import (
|
||||
TritonExperts,
|
||||
TritonWNA16Experts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.xpu_moe import (
|
||||
XPUExperts,
|
||||
XPUExpertsFp8,
|
||||
@@ -94,14 +101,9 @@ if HAS_TRITON:
|
||||
BatchedTritonExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import (
|
||||
TritonExperts,
|
||||
TritonWNA16Experts,
|
||||
fused_experts,
|
||||
get_config_file_name,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
AiterExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.router.fused_topk_router import (
|
||||
fused_topk,
|
||||
)
|
||||
|
||||
+5
-2
@@ -29,6 +29,7 @@ from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
||||
from vllm.model_executor.layers.fused_moe.utils import (
|
||||
_resize_cache,
|
||||
disable_inplace,
|
||||
swiglu_limit_func,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.marlin_utils import (
|
||||
get_marlin_input_dtype,
|
||||
@@ -50,8 +51,6 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.scalar_type import ScalarType, scalar_types
|
||||
|
||||
from .utils import swiglu_limit_func
|
||||
|
||||
|
||||
def _fused_marlin_moe(
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -414,6 +413,7 @@ def batched_fused_marlin_moe(
|
||||
is_k_full: bool = True,
|
||||
output: torch.Tensor | None = None,
|
||||
inplace: bool = False,
|
||||
clamp_limit: float | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
This function massages the inputs so the batched hidden_states can be
|
||||
@@ -536,6 +536,7 @@ def batched_fused_marlin_moe(
|
||||
intermediate_cache2=intermediate_cache2,
|
||||
output=output.view(-1, K) if output is not None else output,
|
||||
is_k_full=is_k_full,
|
||||
clamp_limit=clamp_limit,
|
||||
)
|
||||
|
||||
output = output.view(B, BATCH_TOKENS_MAX, K)
|
||||
@@ -769,6 +770,7 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase):
|
||||
sort_indices2=self.w2_g_idx_sort_indices,
|
||||
is_k_full=self.is_k_full,
|
||||
input_dtype=self.input_dtype,
|
||||
clamp_limit=self.gemm1_clamp_limit,
|
||||
)
|
||||
return
|
||||
|
||||
@@ -971,4 +973,5 @@ class BatchedMarlinExperts(MarlinExpertsBase):
|
||||
sort_indices1=self.w13_g_idx_sort_indices,
|
||||
sort_indices2=self.w2_g_idx_sort_indices,
|
||||
is_k_full=self.is_k_full,
|
||||
clamp_limit=self.gemm1_clamp_limit,
|
||||
)
|
||||
@@ -20,7 +20,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.utils import moe_kernel_quantize_input
|
||||
from vllm.model_executor.layers.quantization.utils.nvfp4_emulation_utils import (
|
||||
dequantize_to_dtype,
|
||||
|
||||
@@ -20,7 +20,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.utils import moe_kernel_quantize_input
|
||||
from vllm.model_executor.layers.quantization.utils.mxfp4_utils import dequant_mxfp4
|
||||
from vllm.model_executor.layers.quantization.utils.mxfp6_utils import dequant_mxfp6
|
||||
|
||||
@@ -0,0 +1,522 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Triton-based MoE expert implementations."""
|
||||
|
||||
import torch
|
||||
|
||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||
from vllm import _custom_ops as ops
|
||||
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEParallelConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import (
|
||||
_prepare_expert_assignment,
|
||||
invoke_fused_moe_triton_kernel,
|
||||
invoke_fused_moe_wna16_triton_kernel,
|
||||
try_get_optimal_moe_config,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.lora_experts_mixin import (
|
||||
LoRAExpertsMixin,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.moe_align_block_size import (
|
||||
moe_align_block_size,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
||||
TopKWeightAndReduceNoOP,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.utils import (
|
||||
_resize_cache,
|
||||
moe_kernel_quantize_input,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
QuantKey,
|
||||
kFp8Dynamic128Sym,
|
||||
kFp8DynamicTensorSym,
|
||||
kFp8DynamicTokenSym,
|
||||
kFp8Static128BlockSym,
|
||||
kFp8StaticChannelSym,
|
||||
kFp8StaticTensorSym,
|
||||
kInt8DynamicTokenSym,
|
||||
kInt8StaticChannelSym,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.triton_utils import tl
|
||||
|
||||
|
||||
class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
|
||||
"""Triton-based fused MoE expert implementation."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
moe_config: FusedMoEConfig,
|
||||
quant_config: FusedMoEQuantConfig,
|
||||
):
|
||||
# Whether quantized MOE runs natively, or through
|
||||
# higher-precision + activation QDQ.
|
||||
self.quantization_emulation = False
|
||||
super().__init__(moe_config, quant_config)
|
||||
|
||||
@staticmethod
|
||||
def activation_format() -> mk.FusedMoEActivationFormat:
|
||||
return mk.FusedMoEActivationFormat.Standard
|
||||
|
||||
@staticmethod
|
||||
def _supports_current_device() -> bool:
|
||||
return current_platform.is_cuda_alike() or current_platform.is_xpu()
|
||||
|
||||
@staticmethod
|
||||
def _supports_no_act_and_mul() -> bool:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _supports_quant_scheme(
|
||||
weight_key: QuantKey | None,
|
||||
activation_key: QuantKey | None,
|
||||
) -> bool:
|
||||
# INT8 requires at least 7.5 (Turing).
|
||||
device_supports_int8 = (
|
||||
current_platform.is_cuda()
|
||||
and current_platform.has_device_capability((7, 5))
|
||||
)
|
||||
|
||||
supported: list[tuple[QuantKey | None, QuantKey | None]] = [(None, None)]
|
||||
if device_supports_int8:
|
||||
supported.append((kInt8StaticChannelSym, kInt8DynamicTokenSym))
|
||||
if current_platform.supports_fp8():
|
||||
supported += [
|
||||
(kFp8Static128BlockSym, kFp8Dynamic128Sym),
|
||||
(kFp8StaticChannelSym, kFp8DynamicTokenSym),
|
||||
(kFp8StaticTensorSym, kFp8DynamicTokenSym),
|
||||
(kFp8StaticTensorSym, kFp8StaticTensorSym),
|
||||
(kFp8StaticTensorSym, kFp8DynamicTensorSym),
|
||||
]
|
||||
return (weight_key, activation_key) in supported
|
||||
|
||||
@staticmethod
|
||||
def _supports_activation(activation: MoEActivation) -> bool:
|
||||
return activation in [
|
||||
MoEActivation.SILU,
|
||||
MoEActivation.GELU,
|
||||
MoEActivation.GELU_TANH,
|
||||
MoEActivation.SWIGLUOAI,
|
||||
MoEActivation.SWIGLUSTEP,
|
||||
MoEActivation.SILU_NO_MUL,
|
||||
MoEActivation.GELU_NO_MUL,
|
||||
MoEActivation.GELU_TANH_NO_MUL,
|
||||
MoEActivation.RELU2_NO_MUL,
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool:
|
||||
return not (
|
||||
moe_parallel_config.use_fi_nvl_two_sided_kernels
|
||||
or moe_parallel_config.use_fi_nvl_one_sided_kernels
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_batch_invariance():
|
||||
return True
|
||||
|
||||
def supports_expert_map(self) -> bool:
|
||||
return True
|
||||
|
||||
def finalize_weight_and_reduce_impl(self) -> mk.TopKWeightAndReduce:
|
||||
return TopKWeightAndReduceNoOP()
|
||||
|
||||
def workspace_shapes(
|
||||
self,
|
||||
M: int,
|
||||
N: int,
|
||||
K: int,
|
||||
topk: int,
|
||||
global_num_experts: int,
|
||||
local_num_experts: int,
|
||||
expert_tokens_meta: mk.ExpertTokensMetadata | None,
|
||||
activation: MoEActivation,
|
||||
) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
|
||||
activation_out_dim = self.adjust_N_for_activation(N, activation)
|
||||
workspace1 = (M, topk, max(activation_out_dim, K))
|
||||
workspace2 = (M, topk, max(N, K))
|
||||
output = (M, K)
|
||||
return (workspace1, workspace2, output)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
w1: torch.Tensor,
|
||||
w2: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
activation: MoEActivation,
|
||||
global_num_experts: int,
|
||||
expert_map: torch.Tensor | None,
|
||||
a1q_scale: torch.Tensor | None,
|
||||
a2_scale: torch.Tensor | None,
|
||||
workspace13: torch.Tensor,
|
||||
workspace2: torch.Tensor,
|
||||
expert_tokens_meta: mk.ExpertTokensMetadata | None,
|
||||
apply_router_weight_on_input: bool,
|
||||
):
|
||||
# Check constraints.
|
||||
if self.quant_config.use_int4_w4a16:
|
||||
assert hidden_states.size(-1) // 2 == w1.size(2), "Hidden size mismatch"
|
||||
else:
|
||||
assert hidden_states.size(-1) == w1.size(2), (
|
||||
f"Hidden size mismatch {hidden_states.size(-1)} != {w1.size(2)}"
|
||||
)
|
||||
|
||||
assert hidden_states.is_contiguous(), "Hidden_states must be contiguous"
|
||||
assert hidden_states.dim() == 2
|
||||
assert w1.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert w2.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert hidden_states.dtype in [
|
||||
torch.float32,
|
||||
torch.float16,
|
||||
torch.bfloat16,
|
||||
torch.float8_e4m3fn,
|
||||
torch.float8_e4m3fnuz,
|
||||
]
|
||||
|
||||
E, num_tokens, N, K, top_k_num = self.moe_problem_size(
|
||||
hidden_states, w1, w2, topk_ids
|
||||
)
|
||||
|
||||
if global_num_experts == -1:
|
||||
global_num_experts = E
|
||||
|
||||
config = try_get_optimal_moe_config(
|
||||
w1.size(),
|
||||
w2.size(),
|
||||
top_k_num,
|
||||
self.quant_config.config_name(hidden_states.dtype),
|
||||
num_tokens,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
if hidden_states.dtype == torch.bfloat16:
|
||||
compute_type = tl.bfloat16
|
||||
elif hidden_states.dtype == torch.float16:
|
||||
compute_type = tl.float16
|
||||
elif hidden_states.dtype == torch.float32:
|
||||
compute_type = tl.float32
|
||||
elif (
|
||||
hidden_states.dtype == torch.float8_e4m3fn
|
||||
or hidden_states.dtype == torch.float8_e4m3fnuz
|
||||
):
|
||||
compute_type = tl.bfloat16
|
||||
else:
|
||||
raise ValueError(f"Unsupported compute_type: {hidden_states.dtype}")
|
||||
|
||||
# Note that the output tensor might be in workspace1
|
||||
intermediate_cache1 = _resize_cache(workspace2, (num_tokens, top_k_num, N))
|
||||
cache2_dim = self.adjust_N_for_activation(N, activation)
|
||||
intermediate_cache2 = _resize_cache(
|
||||
workspace13, (num_tokens * top_k_num, cache2_dim)
|
||||
)
|
||||
intermediate_cache3 = _resize_cache(workspace2, (num_tokens, top_k_num, K))
|
||||
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = (
|
||||
_prepare_expert_assignment(
|
||||
topk_ids,
|
||||
config,
|
||||
num_tokens,
|
||||
top_k_num,
|
||||
global_num_experts,
|
||||
expert_map,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
)
|
||||
|
||||
invoke_fused_moe_triton_kernel(
|
||||
hidden_states,
|
||||
w1,
|
||||
intermediate_cache1,
|
||||
a1q_scale,
|
||||
self.w1_scale,
|
||||
None, # topk_weights
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
False, # mul_routed_weights
|
||||
top_k_num,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
|
||||
use_int8_w8a8=self.quant_config.use_int8_w8a8,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
per_channel_quant=self.per_act_token_quant,
|
||||
block_shape=self.block_shape,
|
||||
B_bias=self.w1_bias,
|
||||
)
|
||||
|
||||
# LoRA w13: applied to intermediate_cache1 before activation, using
|
||||
# hidden_states as the lora_a input. moe_lora_align_block_size is
|
||||
# called once here and results reused for the w2 LoRA below.
|
||||
sorted_token_ids_lora = None
|
||||
expert_ids_lora = None
|
||||
num_tokens_post_padded_lora = None
|
||||
token_lora_mapping = None
|
||||
lora_context = self._lora_context
|
||||
if lora_context is not None:
|
||||
(
|
||||
sorted_token_ids_lora,
|
||||
expert_ids_lora,
|
||||
num_tokens_post_padded_lora,
|
||||
token_lora_mapping,
|
||||
) = self.apply_w13_lora(
|
||||
lora_context,
|
||||
y=intermediate_cache1,
|
||||
x=hidden_states,
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
expert_map=expert_map,
|
||||
w1=w1,
|
||||
w2=w2,
|
||||
num_tokens=num_tokens,
|
||||
top_k_num=top_k_num,
|
||||
)
|
||||
|
||||
self.activation(
|
||||
activation, intermediate_cache2, intermediate_cache1.view(-1, N)
|
||||
)
|
||||
|
||||
a2q_scale: torch.Tensor | None = None
|
||||
|
||||
qintermediate_cache2, a2q_scale = moe_kernel_quantize_input(
|
||||
intermediate_cache2,
|
||||
a2_scale,
|
||||
self.quant_dtype,
|
||||
self.per_act_token_quant,
|
||||
self.block_shape,
|
||||
quantization_emulation=self.quantization_emulation,
|
||||
)
|
||||
|
||||
invoke_fused_moe_triton_kernel(
|
||||
qintermediate_cache2,
|
||||
w2,
|
||||
intermediate_cache3,
|
||||
a2q_scale,
|
||||
self.w2_scale,
|
||||
topk_weights,
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
not apply_router_weight_on_input,
|
||||
1,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
|
||||
use_int8_w8a8=self.quant_config.use_int8_w8a8,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
per_channel_quant=self.per_act_token_quant,
|
||||
block_shape=self.block_shape,
|
||||
B_bias=self.w2_bias,
|
||||
)
|
||||
|
||||
# LoRA w2: applied to intermediate_cache3 before moe_sum, using the
|
||||
# unquantized intermediate_cache2 as the lora_a input. Reuses the
|
||||
# sorted_token_ids_lora computed above.
|
||||
if lora_context is not None:
|
||||
self.apply_w2_lora(
|
||||
lora_context,
|
||||
y=intermediate_cache3,
|
||||
x=intermediate_cache2,
|
||||
topk_weights=topk_weights,
|
||||
sorted_token_ids_lora=sorted_token_ids_lora,
|
||||
expert_ids_lora=expert_ids_lora,
|
||||
num_tokens_post_padded_lora=num_tokens_post_padded_lora,
|
||||
token_lora_mapping=token_lora_mapping,
|
||||
num_tokens=num_tokens,
|
||||
w1=w1,
|
||||
w2=w2,
|
||||
top_k_num=top_k_num,
|
||||
)
|
||||
|
||||
# separate function is required for MoE + LoRA
|
||||
self.moe_sum(intermediate_cache3, output)
|
||||
|
||||
def moe_sum(self, input: torch.Tensor, output: torch.Tensor) -> None:
|
||||
ops.moe_sum(input, output)
|
||||
|
||||
|
||||
class TritonWNA16Experts(TritonExperts):
|
||||
@staticmethod
|
||||
def _supports_current_device() -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_no_act_and_mul() -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_quant_scheme(
|
||||
weight_key: QuantKey | None,
|
||||
activation_key: QuantKey | None,
|
||||
) -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_activation(activation: MoEActivation) -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
w1: torch.Tensor,
|
||||
w2: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
activation: MoEActivation,
|
||||
global_num_experts: int,
|
||||
expert_map: torch.Tensor | None,
|
||||
a1q_scale: torch.Tensor | None,
|
||||
a2_scale: torch.Tensor | None,
|
||||
workspace13: torch.Tensor,
|
||||
workspace2: torch.Tensor,
|
||||
expert_tokens_meta: mk.ExpertTokensMetadata | None,
|
||||
apply_router_weight_on_input: bool,
|
||||
):
|
||||
# Check constraints.
|
||||
if self.quant_config.use_int4_w4a16:
|
||||
assert hidden_states.size(-1) // 2 == w1.size(2), "Hidden size mismatch"
|
||||
else:
|
||||
assert hidden_states.size(-1) == w1.size(2), (
|
||||
f"Hidden size mismatch {hidden_states.size(-1)} != {w1.size(2)}"
|
||||
)
|
||||
|
||||
assert hidden_states.is_contiguous(), "Hidden_states must be contiguous"
|
||||
assert hidden_states.dim() == 2
|
||||
assert w1.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert w2.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert hidden_states.dtype in [
|
||||
torch.float32,
|
||||
torch.float16,
|
||||
torch.bfloat16,
|
||||
torch.float8_e4m3fn,
|
||||
torch.float8_e4m3fnuz,
|
||||
]
|
||||
|
||||
E, num_tokens, N, K, top_k_num = self.moe_problem_size(
|
||||
hidden_states, w1, w2, topk_ids
|
||||
)
|
||||
|
||||
if global_num_experts == -1:
|
||||
global_num_experts = E
|
||||
|
||||
config = try_get_optimal_moe_config(
|
||||
w1.size(),
|
||||
w2.size(),
|
||||
top_k_num,
|
||||
self.quant_config.config_name(hidden_states.dtype),
|
||||
num_tokens,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
if hidden_states.dtype == torch.bfloat16:
|
||||
compute_type = tl.bfloat16
|
||||
elif hidden_states.dtype == torch.float16:
|
||||
compute_type = tl.float16
|
||||
elif hidden_states.dtype == torch.float32:
|
||||
compute_type = tl.float32
|
||||
elif (
|
||||
hidden_states.dtype == torch.float8_e4m3fn
|
||||
or hidden_states.dtype == torch.float8_e4m3fnuz
|
||||
):
|
||||
compute_type = tl.bfloat16
|
||||
else:
|
||||
raise ValueError(f"Unsupported compute_type: {hidden_states.dtype}")
|
||||
|
||||
# Note that the output tensor might be in workspace1
|
||||
intermediate_cache1 = _resize_cache(workspace2, (num_tokens, top_k_num, N))
|
||||
activation_out_dim = self.adjust_N_for_activation(N, activation)
|
||||
intermediate_cache2 = _resize_cache(
|
||||
workspace13, (num_tokens * top_k_num, activation_out_dim)
|
||||
)
|
||||
intermediate_cache3 = _resize_cache(workspace2, (num_tokens, top_k_num, K))
|
||||
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
||||
topk_ids, config["BLOCK_SIZE_M"], global_num_experts, expert_map
|
||||
)
|
||||
|
||||
invoke_fused_moe_wna16_triton_kernel(
|
||||
hidden_states,
|
||||
w1,
|
||||
intermediate_cache1,
|
||||
self.w1_scale,
|
||||
self.quant_config.w1_zp,
|
||||
None, # topk_weights
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
False, # mul_routed_weights
|
||||
top_k_num,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
self.activation(
|
||||
activation, intermediate_cache2, intermediate_cache1.view(-1, N)
|
||||
)
|
||||
|
||||
a2q_scale: torch.Tensor | None = None
|
||||
|
||||
qintermediate_cache2, a2q_scale = moe_kernel_quantize_input(
|
||||
intermediate_cache2,
|
||||
a2_scale,
|
||||
self.quant_dtype,
|
||||
self.per_act_token_quant,
|
||||
self.block_shape,
|
||||
)
|
||||
|
||||
invoke_fused_moe_wna16_triton_kernel(
|
||||
qintermediate_cache2,
|
||||
w2,
|
||||
intermediate_cache3,
|
||||
self.w2_scale,
|
||||
self.quant_config.w2_zp,
|
||||
topk_weights,
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
not apply_router_weight_on_input,
|
||||
1,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
# separate function is required for MoE + LoRA
|
||||
self.moe_sum(intermediate_cache3, output)
|
||||
@@ -20,34 +20,16 @@ from vllm.model_executor.layers.fused_moe.activation import (
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FUSED_MOE_UNQUANTIZED_CONFIG,
|
||||
FusedMoEConfig,
|
||||
FusedMoEParallelConfig,
|
||||
FusedMoEQuantConfig,
|
||||
_get_config_dtype_str,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.lora_experts_mixin import LoRAExpertsMixin
|
||||
from vllm.model_executor.layers.fused_moe.moe_align_block_size import (
|
||||
moe_align_block_size,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
||||
TopKWeightAndReduceNoOP,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.utils import (
|
||||
_resize_cache,
|
||||
disable_inplace,
|
||||
moe_kernel_quantize_input,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
QuantKey,
|
||||
kFp8Dynamic128Sym,
|
||||
kFp8DynamicTensorSym,
|
||||
kFp8DynamicTokenSym,
|
||||
kFp8Static128BlockSym,
|
||||
kFp8StaticChannelSym,
|
||||
kFp8StaticTensorSym,
|
||||
kInt8DynamicTokenSym,
|
||||
kInt8StaticChannelSym,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.triton_utils import tl, triton
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
@@ -1885,479 +1867,3 @@ def fused_experts_impl(
|
||||
)
|
||||
|
||||
return out_hidden_states
|
||||
|
||||
|
||||
class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
|
||||
"""Triton-based fused MoE expert implementation."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
moe_config: FusedMoEConfig,
|
||||
quant_config: FusedMoEQuantConfig,
|
||||
):
|
||||
# Whether quantized MOE runs natively, or through
|
||||
# higher-precision + activation QDQ.
|
||||
self.quantization_emulation = False
|
||||
super().__init__(moe_config, quant_config)
|
||||
|
||||
@staticmethod
|
||||
def activation_format() -> mk.FusedMoEActivationFormat:
|
||||
return mk.FusedMoEActivationFormat.Standard
|
||||
|
||||
@staticmethod
|
||||
def _supports_current_device() -> bool:
|
||||
return current_platform.is_cuda_alike() or current_platform.is_xpu()
|
||||
|
||||
@staticmethod
|
||||
def _supports_no_act_and_mul() -> bool:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _supports_quant_scheme(
|
||||
weight_key: QuantKey | None,
|
||||
activation_key: QuantKey | None,
|
||||
) -> bool:
|
||||
# INT8 requires at least 7.5 (Turing).
|
||||
device_supports_int8 = (
|
||||
current_platform.is_cuda()
|
||||
and current_platform.has_device_capability((7, 5))
|
||||
)
|
||||
|
||||
supported: list[tuple[QuantKey | None, QuantKey | None]] = [(None, None)]
|
||||
if device_supports_int8:
|
||||
supported.append((kInt8StaticChannelSym, kInt8DynamicTokenSym))
|
||||
if current_platform.supports_fp8():
|
||||
supported += [
|
||||
(kFp8Static128BlockSym, kFp8Dynamic128Sym),
|
||||
(kFp8StaticChannelSym, kFp8DynamicTokenSym),
|
||||
(kFp8StaticTensorSym, kFp8DynamicTokenSym),
|
||||
(kFp8StaticTensorSym, kFp8StaticTensorSym),
|
||||
(kFp8StaticTensorSym, kFp8DynamicTensorSym),
|
||||
]
|
||||
return (weight_key, activation_key) in supported
|
||||
|
||||
@staticmethod
|
||||
def _supports_activation(activation: MoEActivation) -> bool:
|
||||
return activation in [
|
||||
MoEActivation.SILU,
|
||||
MoEActivation.GELU,
|
||||
MoEActivation.GELU_TANH,
|
||||
MoEActivation.SWIGLUOAI,
|
||||
MoEActivation.SWIGLUSTEP,
|
||||
MoEActivation.SILU_NO_MUL,
|
||||
MoEActivation.GELU_NO_MUL,
|
||||
MoEActivation.GELU_TANH_NO_MUL,
|
||||
MoEActivation.RELU2_NO_MUL,
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool:
|
||||
return not (
|
||||
moe_parallel_config.use_fi_nvl_two_sided_kernels
|
||||
or moe_parallel_config.use_fi_nvl_one_sided_kernels
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_batch_invariance():
|
||||
return True
|
||||
|
||||
def supports_expert_map(self) -> bool:
|
||||
return True
|
||||
|
||||
def finalize_weight_and_reduce_impl(self) -> mk.TopKWeightAndReduce:
|
||||
return TopKWeightAndReduceNoOP()
|
||||
|
||||
def workspace_shapes(
|
||||
self,
|
||||
M: int,
|
||||
N: int,
|
||||
K: int,
|
||||
topk: int,
|
||||
global_num_experts: int,
|
||||
local_num_experts: int,
|
||||
expert_tokens_meta: mk.ExpertTokensMetadata | None,
|
||||
activation: MoEActivation,
|
||||
) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
|
||||
activation_out_dim = self.adjust_N_for_activation(N, activation)
|
||||
workspace1 = (M, topk, max(activation_out_dim, K))
|
||||
workspace2 = (M, topk, max(N, K))
|
||||
output = (M, K)
|
||||
return (workspace1, workspace2, output)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
w1: torch.Tensor,
|
||||
w2: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
activation: MoEActivation,
|
||||
global_num_experts: int,
|
||||
expert_map: torch.Tensor | None,
|
||||
a1q_scale: torch.Tensor | None,
|
||||
a2_scale: torch.Tensor | None,
|
||||
workspace13: torch.Tensor,
|
||||
workspace2: torch.Tensor,
|
||||
expert_tokens_meta: mk.ExpertTokensMetadata | None,
|
||||
apply_router_weight_on_input: bool,
|
||||
):
|
||||
# Check constraints.
|
||||
if self.quant_config.use_int4_w4a16:
|
||||
assert hidden_states.size(-1) // 2 == w1.size(2), "Hidden size mismatch"
|
||||
else:
|
||||
assert hidden_states.size(-1) == w1.size(2), (
|
||||
f"Hidden size mismatch {hidden_states.size(-1)} != {w1.size(2)}"
|
||||
)
|
||||
|
||||
assert hidden_states.is_contiguous(), "Hidden_states must be contiguous"
|
||||
assert hidden_states.dim() == 2
|
||||
assert w1.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert w2.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert hidden_states.dtype in [
|
||||
torch.float32,
|
||||
torch.float16,
|
||||
torch.bfloat16,
|
||||
torch.float8_e4m3fn,
|
||||
torch.float8_e4m3fnuz,
|
||||
]
|
||||
|
||||
E, num_tokens, N, K, top_k_num = self.moe_problem_size(
|
||||
hidden_states, w1, w2, topk_ids
|
||||
)
|
||||
|
||||
if global_num_experts == -1:
|
||||
global_num_experts = E
|
||||
|
||||
config = try_get_optimal_moe_config(
|
||||
w1.size(),
|
||||
w2.size(),
|
||||
top_k_num,
|
||||
self.quant_config.config_name(hidden_states.dtype),
|
||||
num_tokens,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
if hidden_states.dtype == torch.bfloat16:
|
||||
compute_type = tl.bfloat16
|
||||
elif hidden_states.dtype == torch.float16:
|
||||
compute_type = tl.float16
|
||||
elif hidden_states.dtype == torch.float32:
|
||||
compute_type = tl.float32
|
||||
elif (
|
||||
hidden_states.dtype == torch.float8_e4m3fn
|
||||
or hidden_states.dtype == torch.float8_e4m3fnuz
|
||||
):
|
||||
compute_type = tl.bfloat16
|
||||
else:
|
||||
raise ValueError(f"Unsupported compute_type: {hidden_states.dtype}")
|
||||
|
||||
# Note that the output tensor might be in workspace1
|
||||
intermediate_cache1 = _resize_cache(workspace2, (num_tokens, top_k_num, N))
|
||||
cache2_dim = self.adjust_N_for_activation(N, activation)
|
||||
intermediate_cache2 = _resize_cache(
|
||||
workspace13, (num_tokens * top_k_num, cache2_dim)
|
||||
)
|
||||
intermediate_cache3 = _resize_cache(workspace2, (num_tokens, top_k_num, K))
|
||||
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = (
|
||||
_prepare_expert_assignment(
|
||||
topk_ids,
|
||||
config,
|
||||
num_tokens,
|
||||
top_k_num,
|
||||
global_num_experts,
|
||||
expert_map,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
)
|
||||
|
||||
invoke_fused_moe_triton_kernel(
|
||||
hidden_states,
|
||||
w1,
|
||||
intermediate_cache1,
|
||||
a1q_scale,
|
||||
self.w1_scale,
|
||||
None, # topk_weights
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
False, # mul_routed_weights
|
||||
top_k_num,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
|
||||
use_int8_w8a8=self.quant_config.use_int8_w8a8,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
per_channel_quant=self.per_act_token_quant,
|
||||
block_shape=self.block_shape,
|
||||
B_bias=self.w1_bias,
|
||||
)
|
||||
|
||||
# LoRA w13: applied to intermediate_cache1 before activation, using
|
||||
# hidden_states as the lora_a input. moe_lora_align_block_size is
|
||||
# called once here and results reused for the w2 LoRA below.
|
||||
sorted_token_ids_lora = None
|
||||
expert_ids_lora = None
|
||||
num_tokens_post_padded_lora = None
|
||||
token_lora_mapping = None
|
||||
lora_context = self._lora_context
|
||||
if lora_context is not None:
|
||||
(
|
||||
sorted_token_ids_lora,
|
||||
expert_ids_lora,
|
||||
num_tokens_post_padded_lora,
|
||||
token_lora_mapping,
|
||||
) = self.apply_w13_lora(
|
||||
lora_context,
|
||||
y=intermediate_cache1,
|
||||
x=hidden_states,
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
expert_map=expert_map,
|
||||
w1=w1,
|
||||
w2=w2,
|
||||
num_tokens=num_tokens,
|
||||
top_k_num=top_k_num,
|
||||
)
|
||||
|
||||
self.activation(
|
||||
activation, intermediate_cache2, intermediate_cache1.view(-1, N)
|
||||
)
|
||||
|
||||
a2q_scale: torch.Tensor | None = None
|
||||
|
||||
qintermediate_cache2, a2q_scale = moe_kernel_quantize_input(
|
||||
intermediate_cache2,
|
||||
a2_scale,
|
||||
self.quant_dtype,
|
||||
self.per_act_token_quant,
|
||||
self.block_shape,
|
||||
quantization_emulation=self.quantization_emulation,
|
||||
)
|
||||
|
||||
invoke_fused_moe_triton_kernel(
|
||||
qintermediate_cache2,
|
||||
w2,
|
||||
intermediate_cache3,
|
||||
a2q_scale,
|
||||
self.w2_scale,
|
||||
topk_weights,
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
not apply_router_weight_on_input,
|
||||
1,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
|
||||
use_int8_w8a8=self.quant_config.use_int8_w8a8,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
per_channel_quant=self.per_act_token_quant,
|
||||
block_shape=self.block_shape,
|
||||
B_bias=self.w2_bias,
|
||||
)
|
||||
|
||||
# LoRA w2: applied to intermediate_cache3 before moe_sum, using the
|
||||
# unquantized intermediate_cache2 as the lora_a input. Reuses the
|
||||
# sorted_token_ids_lora computed above.
|
||||
if lora_context is not None:
|
||||
self.apply_w2_lora(
|
||||
lora_context,
|
||||
y=intermediate_cache3,
|
||||
x=intermediate_cache2,
|
||||
topk_weights=topk_weights,
|
||||
sorted_token_ids_lora=sorted_token_ids_lora,
|
||||
expert_ids_lora=expert_ids_lora,
|
||||
num_tokens_post_padded_lora=num_tokens_post_padded_lora,
|
||||
token_lora_mapping=token_lora_mapping,
|
||||
num_tokens=num_tokens,
|
||||
w1=w1,
|
||||
w2=w2,
|
||||
top_k_num=top_k_num,
|
||||
)
|
||||
|
||||
# separate function is required for MoE + LoRA
|
||||
self.moe_sum(intermediate_cache3, output)
|
||||
|
||||
def moe_sum(self, input: torch.Tensor, output: torch.Tensor) -> None:
|
||||
ops.moe_sum(input, output)
|
||||
|
||||
|
||||
class TritonWNA16Experts(TritonExperts):
|
||||
@staticmethod
|
||||
def _supports_current_device() -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_no_act_and_mul() -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_quant_scheme(
|
||||
weight_key: QuantKey | None,
|
||||
activation_key: QuantKey | None,
|
||||
) -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_activation(activation: MoEActivation) -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool:
|
||||
raise NotImplementedError(
|
||||
"TritonWNA16Experts is not yet used by an Oracle. "
|
||||
"This method should not be called."
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
w1: torch.Tensor,
|
||||
w2: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
activation: MoEActivation,
|
||||
global_num_experts: int,
|
||||
expert_map: torch.Tensor | None,
|
||||
a1q_scale: torch.Tensor | None,
|
||||
a2_scale: torch.Tensor | None,
|
||||
workspace13: torch.Tensor,
|
||||
workspace2: torch.Tensor,
|
||||
expert_tokens_meta: mk.ExpertTokensMetadata | None,
|
||||
apply_router_weight_on_input: bool,
|
||||
):
|
||||
# Check constraints.
|
||||
if self.quant_config.use_int4_w4a16:
|
||||
assert hidden_states.size(-1) // 2 == w1.size(2), "Hidden size mismatch"
|
||||
else:
|
||||
assert hidden_states.size(-1) == w1.size(2), (
|
||||
f"Hidden size mismatch {hidden_states.size(-1)} != {w1.size(2)}"
|
||||
)
|
||||
|
||||
assert hidden_states.is_contiguous(), "Hidden_states must be contiguous"
|
||||
assert hidden_states.dim() == 2
|
||||
assert w1.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert w2.stride(-1) == 1, "Stride of last dimension must be 1"
|
||||
assert hidden_states.dtype in [
|
||||
torch.float32,
|
||||
torch.float16,
|
||||
torch.bfloat16,
|
||||
torch.float8_e4m3fn,
|
||||
torch.float8_e4m3fnuz,
|
||||
]
|
||||
|
||||
E, num_tokens, N, K, top_k_num = self.moe_problem_size(
|
||||
hidden_states, w1, w2, topk_ids
|
||||
)
|
||||
|
||||
if global_num_experts == -1:
|
||||
global_num_experts = E
|
||||
|
||||
config = try_get_optimal_moe_config(
|
||||
w1.size(),
|
||||
w2.size(),
|
||||
top_k_num,
|
||||
self.quant_config.config_name(hidden_states.dtype),
|
||||
num_tokens,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
if hidden_states.dtype == torch.bfloat16:
|
||||
compute_type = tl.bfloat16
|
||||
elif hidden_states.dtype == torch.float16:
|
||||
compute_type = tl.float16
|
||||
elif hidden_states.dtype == torch.float32:
|
||||
compute_type = tl.float32
|
||||
elif (
|
||||
hidden_states.dtype == torch.float8_e4m3fn
|
||||
or hidden_states.dtype == torch.float8_e4m3fnuz
|
||||
):
|
||||
compute_type = tl.bfloat16
|
||||
else:
|
||||
raise ValueError(f"Unsupported compute_type: {hidden_states.dtype}")
|
||||
|
||||
# Note that the output tensor might be in workspace1
|
||||
intermediate_cache1 = _resize_cache(workspace2, (num_tokens, top_k_num, N))
|
||||
activation_out_dim = self.adjust_N_for_activation(N, activation)
|
||||
intermediate_cache2 = _resize_cache(
|
||||
workspace13, (num_tokens * top_k_num, activation_out_dim)
|
||||
)
|
||||
intermediate_cache3 = _resize_cache(workspace2, (num_tokens, top_k_num, K))
|
||||
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
||||
topk_ids, config["BLOCK_SIZE_M"], global_num_experts, expert_map
|
||||
)
|
||||
|
||||
invoke_fused_moe_wna16_triton_kernel(
|
||||
hidden_states,
|
||||
w1,
|
||||
intermediate_cache1,
|
||||
self.w1_scale,
|
||||
self.quant_config.w1_zp,
|
||||
None, # topk_weights
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
False, # mul_routed_weights
|
||||
top_k_num,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
self.activation(
|
||||
activation, intermediate_cache2, intermediate_cache1.view(-1, N)
|
||||
)
|
||||
|
||||
a2q_scale: torch.Tensor | None = None
|
||||
|
||||
qintermediate_cache2, a2q_scale = moe_kernel_quantize_input(
|
||||
intermediate_cache2,
|
||||
a2_scale,
|
||||
self.quant_dtype,
|
||||
self.per_act_token_quant,
|
||||
self.block_shape,
|
||||
)
|
||||
|
||||
invoke_fused_moe_wna16_triton_kernel(
|
||||
qintermediate_cache2,
|
||||
w2,
|
||||
intermediate_cache3,
|
||||
self.w2_scale,
|
||||
self.quant_config.w2_zp,
|
||||
topk_weights,
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
not apply_router_weight_on_input,
|
||||
1,
|
||||
config,
|
||||
compute_type=compute_type,
|
||||
use_int8_w8a16=self.quant_config.use_int8_w8a16,
|
||||
use_int4_w4a16=self.quant_config.use_int4_w4a16,
|
||||
block_shape=self.block_shape,
|
||||
)
|
||||
|
||||
# separate function is required for MoE + LoRA
|
||||
self.moe_sum(intermediate_cache3, output)
|
||||
|
||||
@@ -26,15 +26,15 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEQuantConfig,
|
||||
RoutingMethodType,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
init_aiter_topK_meta_data,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe_method_base import (
|
||||
FusedMoEMethodBase,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe_modular_method import (
|
||||
FusedMoEModularMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
init_aiter_topK_meta_data,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.router.router_factory import (
|
||||
create_fused_moe_router,
|
||||
)
|
||||
|
||||
@@ -123,7 +123,7 @@ def backend_to_kernel_cls(
|
||||
return [TrtLlmFp8ExpertsMonolithic, TrtLlmFp8ExpertsModular]
|
||||
|
||||
elif backend == Fp8MoeBackend.FLASHINFER_CUTLASS:
|
||||
from vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import ( # noqa: E501
|
||||
FlashInferExperts,
|
||||
)
|
||||
|
||||
@@ -144,14 +144,14 @@ def backend_to_kernel_cls(
|
||||
return [BatchedDeepGemmExperts]
|
||||
|
||||
elif backend == Fp8MoeBackend.MARLIN:
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
MarlinExperts,
|
||||
)
|
||||
|
||||
return [MarlinExperts]
|
||||
|
||||
elif backend == Fp8MoeBackend.TRITON:
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import (
|
||||
TritonExperts,
|
||||
)
|
||||
|
||||
@@ -165,7 +165,7 @@ def backend_to_kernel_cls(
|
||||
return [BatchedTritonExperts]
|
||||
|
||||
elif backend == Fp8MoeBackend.AITER:
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
AiterExperts,
|
||||
)
|
||||
|
||||
|
||||
@@ -46,7 +46,7 @@ def backend_to_kernel_cls(
|
||||
backend: Int8MoeBackend,
|
||||
) -> list[type[mk.FusedMoEExperts]]:
|
||||
if backend == Int8MoeBackend.TRITON:
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import (
|
||||
TritonExperts,
|
||||
)
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
BatchedMarlinExperts,
|
||||
MarlinExperts,
|
||||
)
|
||||
@@ -42,14 +42,14 @@ def backend_to_kernel_cls(
|
||||
) -> list[type[mk.FusedMoEExperts]]:
|
||||
"""Return the experts class for the given backend, or None for NONE."""
|
||||
if backend == WNA16MoEBackend.MARLIN:
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
MarlinExperts,
|
||||
)
|
||||
|
||||
return [MarlinExperts]
|
||||
|
||||
elif backend == WNA16MoEBackend.BATCHED_MARLIN:
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
BatchedMarlinExperts,
|
||||
)
|
||||
|
||||
|
||||
@@ -124,7 +124,7 @@ def backend_to_kernel_cls(
|
||||
Mxfp4MoeBackend.FLASHINFER_CUTLASS_MXFP4_BF16,
|
||||
Mxfp4MoeBackend.FLASHINFER_CUTLASS_MXFP4_MXFP8,
|
||||
):
|
||||
from vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import ( # noqa: E501
|
||||
FlashInferExperts,
|
||||
)
|
||||
|
||||
@@ -160,21 +160,21 @@ def backend_to_kernel_cls(
|
||||
]
|
||||
|
||||
elif backend == Mxfp4MoeBackend.MARLIN:
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
MarlinExperts,
|
||||
)
|
||||
|
||||
return [MarlinExperts]
|
||||
|
||||
elif backend == Mxfp4MoeBackend.BATCHED_MARLIN:
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
BatchedMarlinExperts,
|
||||
)
|
||||
|
||||
return [BatchedMarlinExperts]
|
||||
|
||||
elif backend == Mxfp4MoeBackend.AITER_MXFP4_BF16:
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
AiterExperts,
|
||||
)
|
||||
|
||||
|
||||
@@ -89,7 +89,7 @@ def backend_to_kernel_cls(
|
||||
]
|
||||
|
||||
elif backend == NvFp4MoeBackend.FLASHINFER_CUTLASS:
|
||||
from vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import ( # noqa: E501
|
||||
FlashInferExperts,
|
||||
)
|
||||
|
||||
@@ -117,7 +117,7 @@ def backend_to_kernel_cls(
|
||||
return [CutlassExpertsFp4]
|
||||
|
||||
elif backend == NvFp4MoeBackend.MARLIN:
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
MarlinExperts,
|
||||
)
|
||||
|
||||
|
||||
@@ -95,21 +95,23 @@ def backend_to_kernel_cls(
|
||||
return TrtLlmBf16Experts
|
||||
|
||||
elif backend == UnquantizedMoeBackend.FLASHINFER_CUTLASS:
|
||||
from vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import ( # noqa: E501
|
||||
FlashInferExperts,
|
||||
)
|
||||
|
||||
return FlashInferExperts
|
||||
|
||||
elif backend == UnquantizedMoeBackend.AITER:
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
AiterExperts,
|
||||
)
|
||||
|
||||
return AiterExperts
|
||||
|
||||
elif backend == UnquantizedMoeBackend.TRITON:
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import (
|
||||
TritonExperts,
|
||||
)
|
||||
|
||||
return TritonExperts
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
RoutingMethodType,
|
||||
get_routing_method_type,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
rocm_aiter_grouped_topk,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.router.base_router import BaseRouter
|
||||
|
||||
@@ -11,8 +11,8 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.cutlass_moe import CutlassExpertsFp8
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.fallback import FallbackExperts
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
|
||||
|
||||
@@ -14,8 +14,8 @@ from vllm.model_executor.layers.fused_moe.experts.deep_gemm_moe import (
|
||||
_valid_deep_gemm,
|
||||
_valid_deep_gemm_shape,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts
|
||||
from vllm.model_executor.layers.fused_moe.fallback import FallbackExperts
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import TritonExperts
|
||||
from vllm.utils.deep_gemm import (
|
||||
is_deep_gemm_e8m0_used,
|
||||
)
|
||||
|
||||
@@ -17,6 +17,7 @@ from vllm.model_executor.model_loader.weight_utils import sharded_weight_loader
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
from .fla.ops.kda import (
|
||||
FusedRMSNormGated,
|
||||
@@ -84,8 +85,8 @@ direct_register_custom_op(
|
||||
|
||||
class KimiDeltaAttention(nn.Module, MambaBase):
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "gdn_attention"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.GDN_ATTN
|
||||
|
||||
def get_state_dtype(
|
||||
self,
|
||||
|
||||
@@ -8,6 +8,7 @@ import torch
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
|
||||
from vllm.v1.attention.backend import AttentionBackend
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
from vllm.v1.attention.selector import get_mamba_attn_backend
|
||||
from vllm.v1.kv_cache_interface import KVCacheSpec, MambaSpec
|
||||
|
||||
@@ -33,7 +34,7 @@ class MambaBase(AttentionLayerBase):
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def mamba_type(self) -> str:
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
|
||||
@@ -64,6 +64,7 @@ from vllm.utils.torch_utils import (
|
||||
direct_register_custom_op,
|
||||
)
|
||||
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
# Optional ROCm AITER Triton kernels for the GDN decode fast-path.
|
||||
# Availability is checked centrally via rocm_aiter_ops; the actual function
|
||||
@@ -237,8 +238,8 @@ class ChunkGatedDeltaRule(CustomOp):
|
||||
@PluggableLayer.register("gated_delta_net_attention")
|
||||
class GatedDeltaNetAttention(PluggableLayer, MambaBase):
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "gdn_attention"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.GDN_ATTN
|
||||
|
||||
def get_state_dtype(self) -> tuple[torch.dtype, torch.dtype]:
|
||||
return MambaStateDtypeCalculator.gated_delta_net_state_dtype(
|
||||
@@ -263,7 +264,6 @@ class GatedDeltaNetAttention(PluggableLayer, MambaBase):
|
||||
config: Qwen3NextConfig,
|
||||
vllm_config: VllmConfig,
|
||||
prefix: str = "",
|
||||
create_in_proj_qkvz: bool = True,
|
||||
gqa_interleaved_layout=False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -323,32 +323,14 @@ class GatedDeltaNetAttention(PluggableLayer, MambaBase):
|
||||
# we need to create qkvz_proj adaptively here.
|
||||
# When create_in_proj_qkvz is False (e.g. LoRA enabled in Qwen3.5),
|
||||
# in_proj_qkv and in_proj_z are created separately instead.
|
||||
self.has_lora_projections = not create_in_proj_qkvz
|
||||
if create_in_proj_qkvz:
|
||||
self.in_proj_qkvz = self.create_qkvz_proj(
|
||||
hidden_size=self.hidden_size,
|
||||
key_dim=self.key_dim,
|
||||
value_dim=self.value_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.in_proj_qkvz",
|
||||
)
|
||||
else:
|
||||
# LoRA case (Qwen3.5 only): keep q/k/v and z as separate modules
|
||||
# so that LoRA adapters can be applied independently.
|
||||
self.in_proj_qkv = MergedColumnParallelLinear(
|
||||
input_size=self.hidden_size,
|
||||
output_sizes=[self.key_dim, self.key_dim, self.value_dim],
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.in_proj_qkv",
|
||||
)
|
||||
self.in_proj_z = ColumnParallelLinear(
|
||||
input_size=self.hidden_size,
|
||||
output_size=self.value_dim,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.in_proj_z",
|
||||
)
|
||||
self.in_proj_qkvz = self.create_qkvz_proj(
|
||||
hidden_size=self.hidden_size,
|
||||
key_dim=self.key_dim,
|
||||
value_dim=self.value_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.in_proj_qkvz",
|
||||
)
|
||||
|
||||
# ba_proj doesn't support blockwise fp8 quantization.
|
||||
# Qwen3-Next and Qwen3.5 have different in_proj_ba checkpoint
|
||||
# layouts, so we use a factory method to create the projection.
|
||||
@@ -707,7 +689,7 @@ class GatedDeltaNetAttention(PluggableLayer, MambaBase):
|
||||
):
|
||||
"""ROCm forward using AITER Triton fused projection+attention when
|
||||
available, otherwise falling back to the generic CUDA path."""
|
||||
if not self.has_lora_projections and GDN_AITER_TRITON_AVAILABLE:
|
||||
if GDN_AITER_TRITON_AVAILABLE:
|
||||
num_tokens = hidden_states.size(0)
|
||||
projected_states_qkvz, _ = self.in_proj_qkvz(hidden_states)
|
||||
projected_states_ba, _ = self.in_proj_ba(hidden_states)
|
||||
@@ -752,37 +734,27 @@ class GatedDeltaNetAttention(PluggableLayer, MambaBase):
|
||||
# ============================================================
|
||||
# Part 1: Input Projection
|
||||
# ============================================================
|
||||
if self.has_lora_projections:
|
||||
# LoRA path (Qwen3.5 only): separate in_proj_qkv and in_proj_z
|
||||
mixed_qkv, _ = self.in_proj_qkv(hidden_states)
|
||||
ba, _ = self.in_proj_ba(hidden_states)
|
||||
z, _ = self.in_proj_z(hidden_states)
|
||||
mixed_qkvz, _ = self.in_proj_qkvz(hidden_states)
|
||||
ba, _ = self.in_proj_ba(hidden_states)
|
||||
|
||||
if self.gqa_interleaved_layout:
|
||||
# Qwen3-Next: unpack the interleaved GQA layout
|
||||
query, key, value, z, b, a = self.fix_query_key_value_ordering(
|
||||
mixed_qkvz, ba
|
||||
)
|
||||
query, key, value = map(
|
||||
lambda x: rearrange(x, "l p d -> l (p d)"), (query, key, value)
|
||||
)
|
||||
mixed_qkv = torch.cat((query, key, value), dim=-1)
|
||||
else:
|
||||
# Qwen3.5: weights are already in [q, k, v, z] and [b, a] order
|
||||
qkv_size = (self.key_dim * 2 + self.value_dim) // self.tp_size
|
||||
z_size = self.value_dim // self.tp_size
|
||||
mixed_qkv, z = mixed_qkvz.split([qkv_size, z_size], dim=-1)
|
||||
z = z.reshape(z.size(0), -1, self.head_v_dim)
|
||||
b, a = ba.chunk(2, dim=-1)
|
||||
b = b.contiguous()
|
||||
a = a.contiguous()
|
||||
else:
|
||||
mixed_qkvz, _ = self.in_proj_qkvz(hidden_states)
|
||||
ba, _ = self.in_proj_ba(hidden_states)
|
||||
|
||||
if self.gqa_interleaved_layout:
|
||||
# Qwen3-Next: unpack the interleaved GQA layout
|
||||
query, key, value, z, b, a = self.fix_query_key_value_ordering(
|
||||
mixed_qkvz, ba
|
||||
)
|
||||
query, key, value = map(
|
||||
lambda x: rearrange(x, "l p d -> l (p d)"), (query, key, value)
|
||||
)
|
||||
mixed_qkv = torch.cat((query, key, value), dim=-1)
|
||||
else:
|
||||
# Qwen3.5: weights are already in [q, k, v, z] and [b, a] order
|
||||
qkv_size = (self.key_dim * 2 + self.value_dim) // self.tp_size
|
||||
z_size = self.value_dim // self.tp_size
|
||||
mixed_qkv, z = mixed_qkvz.split([qkv_size, z_size], dim=-1)
|
||||
z = z.reshape(z.size(0), -1, self.head_v_dim)
|
||||
b, a = ba.chunk(2, dim=-1)
|
||||
b = b.contiguous()
|
||||
a = a.contiguous()
|
||||
|
||||
# ============================================================
|
||||
# Part 2: Core Attention (Custom Op)
|
||||
@@ -822,8 +794,6 @@ class GatedDeltaNetAttention(PluggableLayer, MambaBase):
|
||||
"""
|
||||
num_tokens = hidden_states.size(0)
|
||||
|
||||
assert not self.has_lora_projections, "lora isn't supported on XPU."
|
||||
|
||||
# ============================================================
|
||||
# Part 1: Input Projection
|
||||
# ============================================================
|
||||
|
||||
@@ -32,6 +32,7 @@ from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.linear_attn import LinearAttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
|
||||
@CustomOp.register("minimax_text01_rmsnorm_tp")
|
||||
@@ -246,8 +247,8 @@ class MiniMaxText01LinearKernel:
|
||||
|
||||
class MiniMaxText01LinearAttention(nn.Module, MambaBase):
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "linear_attention"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.LINEAR
|
||||
|
||||
def get_state_dtype(self) -> tuple[torch.dtype]:
|
||||
assert self.model_config is not None
|
||||
|
||||
@@ -42,6 +42,7 @@ from vllm.utils.torch_utils import (
|
||||
)
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.mamba1_attn import Mamba1AttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
|
||||
# Adapted from transformers.models.mamba.modeling_mamba.MambaMixer
|
||||
@@ -476,8 +477,8 @@ class MambaMixer(MambaBase, PluggableLayer):
|
||||
)
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "mamba1"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.MAMBA1
|
||||
|
||||
def _time_proj_bias(self) -> torch.Tensor | None:
|
||||
if hasattr(self.dt_proj, "bias") and self.dt_proj.bias is not None:
|
||||
|
||||
@@ -52,6 +52,7 @@ from vllm.utils.torch_utils import (
|
||||
)
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.mamba2_attn import Mamba2AttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
# Added by the IBM Team, 2024
|
||||
|
||||
@@ -935,8 +936,8 @@ class MambaMixer2(MambaBase, PluggableLayer):
|
||||
)
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "mamba2"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.MAMBA2
|
||||
|
||||
|
||||
def mamba_mixer2(
|
||||
|
||||
@@ -37,7 +37,7 @@ def _causal_conv1d_fwd_kernel( # continuous batching
|
||||
num_cache_lines: tl.constexpr, # added to support vLLM larger cache lines
|
||||
# Strides
|
||||
stride_x_dim: tl.constexpr, # stride to get to next feature-value,
|
||||
stride_x_token: tl.constexpr, # stride to get to next token (same feature-index, same sequence-index)
|
||||
stride_x_token: tl.int64, # stride to get to next token (same feature-index, same sequence-index)
|
||||
stride_w_dim: tl.constexpr, # stride to get to next dim-axis value
|
||||
stride_w_width: tl.constexpr, # stride to get to next width-axis value
|
||||
stride_istate_seq: tl.constexpr,
|
||||
@@ -45,7 +45,7 @@ def _causal_conv1d_fwd_kernel( # continuous batching
|
||||
stride_istate_token: tl.constexpr,
|
||||
stride_cache_indices: tl.constexpr,
|
||||
stride_o_dim: tl.constexpr,
|
||||
stride_o_token: tl.constexpr,
|
||||
stride_o_token: tl.int64,
|
||||
stride_block_m: tl.constexpr, # Stride block to align divided by BLOCK_M
|
||||
# others
|
||||
pad_slot_id: tl.constexpr,
|
||||
@@ -769,7 +769,7 @@ def _causal_conv1d_update_kernel(
|
||||
# Strides
|
||||
stride_x_seq: tl.constexpr,
|
||||
stride_x_dim: tl.constexpr,
|
||||
stride_x_token: tl.constexpr,
|
||||
stride_x_token: tl.int64,
|
||||
stride_w_dim: tl.constexpr,
|
||||
stride_w_width: tl.constexpr,
|
||||
stride_conv_state_seq: tl.constexpr,
|
||||
@@ -778,7 +778,7 @@ def _causal_conv1d_update_kernel(
|
||||
stride_state_indices: tl.constexpr,
|
||||
stride_o_seq: tl.constexpr,
|
||||
stride_o_dim: tl.constexpr,
|
||||
stride_o_token: tl.constexpr,
|
||||
stride_o_token: tl.int64,
|
||||
# others
|
||||
null_block_id: tl.constexpr,
|
||||
# Meta-parameters
|
||||
|
||||
@@ -14,6 +14,7 @@ import torch
|
||||
|
||||
from vllm.config.mamba import MambaBackendEnum, MambaConfig
|
||||
from vllm.logger import init_logger
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig, MambaSpec
|
||||
|
||||
@@ -200,7 +201,8 @@ def initialize_mamba_ssu_backend(
|
||||
"""
|
||||
if not any(
|
||||
isinstance(g.kv_cache_spec, MambaSpec)
|
||||
and g.kv_cache_spec.mamba_type in ("mamba1", "mamba2")
|
||||
and g.kv_cache_spec.mamba_type
|
||||
in (MambaAttentionBackendEnum.MAMBA1, MambaAttentionBackendEnum.MAMBA2)
|
||||
for g in kv_cache_config.kv_cache_groups
|
||||
):
|
||||
return
|
||||
|
||||
@@ -25,6 +25,7 @@ from vllm.model_executor.layers.mamba.ops.causal_conv1d import (
|
||||
)
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
from vllm.v1.attention.backends.short_conv_attn import ShortConvAttentionMetadata
|
||||
|
||||
|
||||
@@ -223,8 +224,8 @@ class ShortConv(MambaBase, CustomOp):
|
||||
)
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "short_conv"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.SHORT_CONV
|
||||
|
||||
|
||||
def short_conv(
|
||||
|
||||
@@ -441,6 +441,131 @@ def mhc_post_tilelang(
|
||||
T.pdl_trigger()
|
||||
|
||||
|
||||
@tilelang.jit(
|
||||
pass_configs={
|
||||
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
|
||||
tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
|
||||
tilelang.PassConfigKey.TL_PTXAS_REGISTER_USAGE_LEVEL: 10,
|
||||
},
|
||||
)
|
||||
def mhc_fused_tilelang(
|
||||
comb_mix,
|
||||
residual_in,
|
||||
post_mix,
|
||||
x_in,
|
||||
weight_t,
|
||||
yp_out,
|
||||
rp_out,
|
||||
residual_out,
|
||||
hc: int,
|
||||
hidden: int,
|
||||
n_out: int,
|
||||
n_thr: int = 256,
|
||||
h_blk: int = 256,
|
||||
tile_n: int = 1,
|
||||
split_k: int = 1,
|
||||
) -> tilelang.JITKernel:
|
||||
"""Fused mhc post-mapping + pre-norm GEMM FMA"""
|
||||
m = T.dynamic("num_tokens")
|
||||
split_k = T.dynamic("split_k")
|
||||
h = hidden
|
||||
h_blk = math.gcd(hidden, h_blk)
|
||||
h_per_split = h // split_k
|
||||
n_tiles = n_out // tile_n
|
||||
|
||||
comb_mix: T.Tensor((m, hc, hc), T.float32) # type: ignore[no-redef, valid-type]
|
||||
residual_in: T.Tensor((m, hc, h), T.bfloat16) # type: ignore[no-redef, valid-type]
|
||||
post_mix: T.Tensor((m, hc), T.float32) # type: ignore[no-redef, valid-type]
|
||||
x_in: T.Tensor((m, h), T.bfloat16) # type: ignore[no-redef, valid-type]
|
||||
weight_t: T.Tensor((n_out, hc, h), T.float32) # type: ignore[no-redef, valid-type]
|
||||
yp_out: T.Tensor((split_k, m, n_out), T.float32) # type: ignore[no-redef, valid-type]
|
||||
rp_out: T.Tensor((split_k, m), T.float32) # type: ignore[no-redef, valid-type]
|
||||
residual_out: T.Tensor((m, hc, h), T.bfloat16) # type: ignore[no-redef, valid-type]
|
||||
|
||||
h_iters = h_per_split // n_thr
|
||||
num_warps = n_thr // 32
|
||||
|
||||
with T.Kernel(m, n_tiles, split_k, threads=n_thr) as (i_n, i_nt, i_ks):
|
||||
tid = T.get_thread_binding()
|
||||
warp_id = T.get_warp_idx()
|
||||
lane = T.get_lane_idx()
|
||||
|
||||
s_warp = T.alloc_shared((num_warps, tile_n + 1), T.float32)
|
||||
s_post = T.alloc_shared((hc,), T.float32)
|
||||
s_comb = T.alloc_shared((hc, hc), T.float32)
|
||||
|
||||
pm = T.alloc_local((hc,), T.float32)
|
||||
cm = T.alloc_local((hc, hc), T.float32)
|
||||
acc = T.alloc_local((tile_n,), T.float32)
|
||||
sqr = T.alloc_local((1,), T.float32)
|
||||
new_r = T.alloc_local((hc,), T.float32)
|
||||
|
||||
T.clear(acc)
|
||||
T.clear(sqr)
|
||||
h_split_start = i_ks * h_per_split
|
||||
|
||||
T.pdl_sync()
|
||||
|
||||
T.copy(post_mix[i_n, 0], s_post)
|
||||
T.copy(comb_mix[i_n, 0, 0], s_comb)
|
||||
|
||||
for j in T.unroll(hc):
|
||||
pm[j] = s_post[j]
|
||||
for j in T.unroll(hc):
|
||||
for k in T.unroll(hc):
|
||||
cm[k, j] = s_comb[k, j]
|
||||
|
||||
# Each thread owns h_iters elements of the k-split's h slice.
|
||||
for it in T.serial(h_iters):
|
||||
h_idx = h_split_start + it * n_thr + tid
|
||||
|
||||
# Compute new residual from layer output and past residual
|
||||
for j in T.unroll(hc):
|
||||
new_r[j] = pm[j] * x_in[i_n, h_idx]
|
||||
for k in T.unroll(hc):
|
||||
new_r[j] += cm[k, j] * residual_in[i_n, k, h_idx]
|
||||
|
||||
# populate residual_out and compute sqr sum
|
||||
if i_nt == 0:
|
||||
for j in T.unroll(hc):
|
||||
residual_out[i_n, j, h_idx] = new_r[j]
|
||||
sqr[0] += new_r[j] * new_r[j]
|
||||
|
||||
# Per-thread FMA into acc[n]
|
||||
for n in T.unroll(tile_n):
|
||||
for j in T.unroll(hc):
|
||||
acc[n] += weight_t[i_nt * tile_n + n, j, h_idx] * new_r[j]
|
||||
|
||||
for n in T.unroll(tile_n):
|
||||
acc[n] = T.warp_reduce_sum(acc[n])
|
||||
if i_nt == 0:
|
||||
sqr[0] = T.warp_reduce_sum(sqr[0])
|
||||
|
||||
# Cross-warp reduce via shared mem
|
||||
if lane == 0:
|
||||
for n in T.unroll(tile_n):
|
||||
s_warp[warp_id, n] = acc[n]
|
||||
if i_nt == 0:
|
||||
s_warp[warp_id, tile_n] = sqr[0]
|
||||
T.sync_threads()
|
||||
|
||||
# Warp 0 does the final cross-warp sum and writes outputs
|
||||
if warp_id == 0:
|
||||
if lane < tile_n:
|
||||
v = T.alloc_var(T.float32, init=0.0)
|
||||
for w in T.unroll(num_warps):
|
||||
v += s_warp[w, lane]
|
||||
yp_out[i_ks, i_n, i_nt * tile_n + lane] = v
|
||||
|
||||
if i_nt == 0 and lane == 0:
|
||||
v2 = T.alloc_var(T.float32, init=0.0)
|
||||
for w in T.unroll(num_warps):
|
||||
v2 += s_warp[w, tile_n]
|
||||
rp_out[i_ks, i_n] = v2
|
||||
|
||||
T.pdl_trigger()
|
||||
|
||||
|
||||
def mhc_post(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
@@ -468,6 +593,218 @@ def mhc_post(
|
||||
return out
|
||||
|
||||
|
||||
def mhc_fused_post_pre(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
post_layer_mix: torch.Tensor,
|
||||
comb_res_mix: torch.Tensor,
|
||||
fn: torch.Tensor,
|
||||
hc_scale: torch.Tensor,
|
||||
hc_base: torch.Tensor,
|
||||
rms_eps: float,
|
||||
hc_pre_eps: float,
|
||||
hc_sinkhorn_eps: float,
|
||||
hc_post_mult_value: float,
|
||||
sinkhorn_repeat: int,
|
||||
n_splits: int = 1,
|
||||
tile_n: int = 1,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Run one MHC post block followed by the next MHC pre block.
|
||||
|
||||
Returns:
|
||||
residual_cur: post-mapped residual, shape (..., hc_mult, hidden_size)
|
||||
post_mix_cur: shape (..., hc_mult, 1)
|
||||
comb_mix_cur: shape (..., hc_mult, hc_mult)
|
||||
layer_input_cur: shape (..., hidden_size)
|
||||
"""
|
||||
|
||||
assert residual.dtype == torch.bfloat16
|
||||
assert x.dtype == torch.bfloat16
|
||||
assert post_layer_mix.dtype == torch.float32
|
||||
assert comb_res_mix.dtype == torch.float32
|
||||
assert fn.dtype == torch.float32
|
||||
assert hc_scale.dtype == torch.float32
|
||||
assert hc_base.dtype == torch.float32
|
||||
|
||||
hc_mult = residual.shape[-2]
|
||||
hidden_size = residual.shape[-1]
|
||||
hc_mult2 = hc_mult * hc_mult
|
||||
hc_mult3 = hc_mult * 2 + hc_mult2
|
||||
hc_hidden_size = hc_mult * hidden_size
|
||||
outer_shape = residual.shape[:-2]
|
||||
|
||||
assert x.shape == (*outer_shape, hidden_size)
|
||||
assert post_layer_mix.shape in (
|
||||
(*outer_shape, hc_mult, 1),
|
||||
(*outer_shape, hc_mult),
|
||||
)
|
||||
assert comb_res_mix.shape == (*outer_shape, hc_mult, hc_mult)
|
||||
assert fn.shape == (hc_mult3, hc_hidden_size)
|
||||
assert hc_scale.shape == (3,)
|
||||
assert hc_base.shape == (hc_mult3,)
|
||||
|
||||
assert n_splits in (1, 2, 4, 8)
|
||||
assert hidden_size % n_splits == 0
|
||||
|
||||
residual_flat = residual.view(-1, hc_mult, hidden_size)
|
||||
num_tokens = residual_flat.shape[0]
|
||||
x_flat = x.view(num_tokens, hidden_size)
|
||||
post_layer_mix_flat = post_layer_mix.view(num_tokens, hc_mult)
|
||||
comb_res_mix_flat = comb_res_mix.view(num_tokens, hc_mult, hc_mult)
|
||||
|
||||
fma_token_threshold = 16
|
||||
if num_tokens <= fma_token_threshold:
|
||||
# TODO(gnovack): investigate autotuning these heuristics
|
||||
tile_n = 2 if num_tokens < 8 else 3
|
||||
n_splits = 8 if (num_tokens < 8 and hidden_size <= 4096) else 4
|
||||
else:
|
||||
# these number are from deepgemm kernel impl
|
||||
block_k = 64
|
||||
block_m = 64
|
||||
n_splits = compute_num_split(block_k, hc_hidden_size, cdiv(num_tokens, block_m))
|
||||
|
||||
gemm_out_mul = torch.empty(
|
||||
n_splits,
|
||||
num_tokens,
|
||||
hc_mult3,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
gemm_out_sqrsum = torch.empty(
|
||||
n_splits,
|
||||
num_tokens,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
residual_cur = torch.empty_like(residual_flat)
|
||||
post_mix_cur = torch.empty(
|
||||
num_tokens,
|
||||
hc_mult,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
comb_mix_cur = torch.empty(
|
||||
num_tokens,
|
||||
hc_mult2,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
layer_input_cur = torch.empty(
|
||||
num_tokens,
|
||||
hidden_size,
|
||||
dtype=torch.bfloat16,
|
||||
device=residual.device,
|
||||
)
|
||||
|
||||
if num_tokens <= fma_token_threshold:
|
||||
mhc_fused_tilelang(
|
||||
comb_res_mix_flat,
|
||||
residual_flat,
|
||||
post_layer_mix_flat,
|
||||
x_flat,
|
||||
fn.view(hc_mult3, hc_mult, hidden_size),
|
||||
gemm_out_mul,
|
||||
gemm_out_sqrsum,
|
||||
residual_cur,
|
||||
hc_mult,
|
||||
hidden_size,
|
||||
hc_mult3,
|
||||
tile_n=tile_n,
|
||||
n_splits=n_splits,
|
||||
)
|
||||
else:
|
||||
mhc_post_tilelang(
|
||||
comb_res_mix_flat,
|
||||
residual_flat,
|
||||
post_layer_mix_flat,
|
||||
x_flat,
|
||||
residual_cur,
|
||||
residual.shape[-2],
|
||||
residual.shape[-1],
|
||||
)
|
||||
|
||||
from vllm.utils.deep_gemm import tf32_hc_prenorm_gemm
|
||||
|
||||
tf32_hc_prenorm_gemm(
|
||||
residual_cur.view(num_tokens, hc_mult * hidden_size),
|
||||
fn,
|
||||
gemm_out_mul,
|
||||
gemm_out_sqrsum,
|
||||
n_splits,
|
||||
)
|
||||
|
||||
mhc_pre_big_fuse_tilelang(
|
||||
gemm_out_mul,
|
||||
gemm_out_sqrsum,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
residual_cur,
|
||||
post_mix_cur,
|
||||
comb_mix_cur,
|
||||
layer_input_cur,
|
||||
hidden_size,
|
||||
rms_eps,
|
||||
hc_pre_eps,
|
||||
hc_sinkhorn_eps,
|
||||
hc_post_mult_value,
|
||||
sinkhorn_repeat,
|
||||
n_splits,
|
||||
hc_mult,
|
||||
)
|
||||
|
||||
return (
|
||||
residual_cur.view(*outer_shape, hc_mult, hidden_size),
|
||||
post_mix_cur.view(*outer_shape, hc_mult, 1),
|
||||
comb_mix_cur.view(*outer_shape, hc_mult, hc_mult),
|
||||
layer_input_cur.view(*outer_shape, hidden_size),
|
||||
)
|
||||
|
||||
|
||||
def _mhc_fused_post_pre_fake(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
post_layer_mix: torch.Tensor,
|
||||
comb_res_mix: torch.Tensor,
|
||||
fn: torch.Tensor,
|
||||
hc_scale: torch.Tensor,
|
||||
hc_base: torch.Tensor,
|
||||
rms_eps: float,
|
||||
hc_pre_eps: float,
|
||||
hc_sinkhorn_eps: float,
|
||||
hc_post_mult_value: float,
|
||||
sinkhorn_repeat: int,
|
||||
n_splits: int = 1,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
hc_mult = residual.shape[-2]
|
||||
hidden_size = residual.shape[-1]
|
||||
outer_shape = residual.shape[:-2]
|
||||
|
||||
residual_cur = torch.empty_like(residual)
|
||||
post_mix_cur = torch.empty(
|
||||
*outer_shape,
|
||||
hc_mult,
|
||||
1,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
comb_mix_cur = torch.empty(
|
||||
*outer_shape,
|
||||
hc_mult,
|
||||
hc_mult,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
layer_input_cur = torch.empty(
|
||||
*outer_shape,
|
||||
hidden_size,
|
||||
dtype=torch.bfloat16,
|
||||
device=residual.device,
|
||||
)
|
||||
|
||||
return residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur
|
||||
|
||||
|
||||
def _mhc_post_fake(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
@@ -489,6 +826,12 @@ direct_register_custom_op(
|
||||
mutates_args=[],
|
||||
fake_impl=_mhc_post_fake,
|
||||
)
|
||||
direct_register_custom_op(
|
||||
op_name="mhc_fused_post_pre",
|
||||
op_func=mhc_fused_post_pre,
|
||||
mutates_args=[],
|
||||
fake_impl=_mhc_fused_post_pre_fake,
|
||||
)
|
||||
|
||||
|
||||
@tilelang.jit(
|
||||
|
||||
@@ -20,7 +20,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import fused_marlin_moe
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import fused_marlin_moe
|
||||
from vllm.model_executor.layers.fused_moe.layer import (
|
||||
FusedMoE,
|
||||
FusedMoEMethodBase,
|
||||
@@ -764,7 +764,7 @@ class AWQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
)
|
||||
|
||||
from vllm.model_executor.layers.fused_moe import modular_kernel as mk
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
BatchedMarlinExperts,
|
||||
MarlinExperts,
|
||||
)
|
||||
|
||||
+1
-1
@@ -17,7 +17,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.cutlass_moe import (
|
||||
CutlassExpertsMxfp4,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
MarlinExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
|
||||
|
||||
+1
-1
@@ -20,7 +20,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEQuantConfig,
|
||||
int4_w4a16_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
BatchedMarlinExperts,
|
||||
MarlinExperts,
|
||||
fused_marlin_moe,
|
||||
|
||||
@@ -26,7 +26,7 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
mxfp4_w4a16_moe_quant_config,
|
||||
ocp_mx_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import fused_marlin_moe
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import fused_marlin_moe
|
||||
from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
|
||||
TRITON_BACKENDS,
|
||||
Mxfp4MoeBackend,
|
||||
@@ -444,7 +444,7 @@ class QuarkW8A8Fp8MoEMethod(QuarkMoEMethod):
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
if self.rocm_aiter_moe_enabled:
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
rocm_aiter_fused_experts,
|
||||
)
|
||||
|
||||
@@ -909,7 +909,7 @@ class QuarkW4A8Fp8MoEMethod(QuarkMoEMethod):
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
rocm_aiter_fused_experts,
|
||||
)
|
||||
|
||||
@@ -1436,7 +1436,7 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
|
||||
|
||||
# AITER path
|
||||
# TODO: Refactor this to use modular MOE kernel as well.
|
||||
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
|
||||
rocm_aiter_fused_experts,
|
||||
)
|
||||
|
||||
|
||||
@@ -256,6 +256,12 @@ class DefaultModelLoader(BaseModelLoader):
|
||||
self.load_config.use_tqdm_on_load,
|
||||
self.load_config.safetensors_load_strategy,
|
||||
local_expert_ids=self.local_expert_ids,
|
||||
safetensors_prefetch_num_threads=(
|
||||
self.load_config.safetensors_prefetch_num_threads
|
||||
),
|
||||
safetensors_prefetch_block_size=(
|
||||
self.load_config.safetensors_prefetch_block_size
|
||||
),
|
||||
)
|
||||
else:
|
||||
if extra_config.get("enable_multithread_load"):
|
||||
|
||||
@@ -30,7 +30,11 @@ from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from vllm import envs
|
||||
from vllm.config import ModelConfig
|
||||
from vllm.config.load import LoadConfig
|
||||
from vllm.config.load import (
|
||||
DEFAULT_SAFETENSORS_PREFETCH_BLOCK_SIZE,
|
||||
DEFAULT_SAFETENSORS_PREFETCH_NUM_THREADS,
|
||||
LoadConfig,
|
||||
)
|
||||
from vllm.distributed import get_tensor_model_parallel_rank, get_world_group
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.quantization import (
|
||||
@@ -810,40 +814,57 @@ def _get_fs_type(files: list[str]) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
def _prefetch_checkpoint(file_path: str) -> None:
|
||||
def _prefetch_checkpoint(
|
||||
file_path: str,
|
||||
block_size: int = DEFAULT_SAFETENSORS_PREFETCH_BLOCK_SIZE,
|
||||
) -> None:
|
||||
"""Prefetch a checkpoint file into the OS page cache.
|
||||
|
||||
Reads the file in 16MB blocks so the kernel caches its pages before
|
||||
workers load the same file.
|
||||
Reads the file in blocks so the kernel caches its pages before workers load
|
||||
the same file.
|
||||
"""
|
||||
block_size = 16 * 1024 * 1024 # 16MB
|
||||
if block_size < 1:
|
||||
raise ValueError("safetensors prefetch block size must be >= 1")
|
||||
|
||||
with open(file_path, "rb") as f:
|
||||
while f.read(block_size):
|
||||
pass
|
||||
|
||||
|
||||
def _prefetch_all_checkpoints(sorted_files: list[str]) -> None:
|
||||
def _prefetch_all_checkpoints(
|
||||
sorted_files: list[str],
|
||||
num_prefetch_threads: int = DEFAULT_SAFETENSORS_PREFETCH_NUM_THREADS,
|
||||
block_size: int = DEFAULT_SAFETENSORS_PREFETCH_BLOCK_SIZE,
|
||||
) -> None:
|
||||
"""Start prefetching checkpoint files into page cache in a background thread."""
|
||||
if num_prefetch_threads < 1:
|
||||
raise ValueError("safetensors prefetch num threads must be >= 1")
|
||||
if block_size < 1:
|
||||
raise ValueError("safetensors prefetch block size must be >= 1")
|
||||
|
||||
if torch.distributed.is_initialized():
|
||||
rank = torch.distributed.get_rank()
|
||||
world_size = torch.distributed.get_world_size()
|
||||
else:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
num_prefetch_threads = 8
|
||||
paths_to_prefetch = sorted_files[rank::world_size]
|
||||
total_for_rank = len(paths_to_prefetch)
|
||||
|
||||
async def _prefetch_all() -> None:
|
||||
semaphore = asyncio.Semaphore(num_prefetch_threads)
|
||||
loop = asyncio.get_running_loop()
|
||||
completed = 0
|
||||
next_log_pct = 10
|
||||
|
||||
async def prefetch_one(path: str) -> None:
|
||||
async def prefetch_one(
|
||||
path: str,
|
||||
executor: concurrent.futures.ThreadPoolExecutor,
|
||||
) -> None:
|
||||
nonlocal completed, next_log_pct
|
||||
try:
|
||||
async with semaphore:
|
||||
await asyncio.to_thread(_prefetch_checkpoint, path)
|
||||
await loop.run_in_executor(
|
||||
executor, _prefetch_checkpoint, path, block_size
|
||||
)
|
||||
completed += 1
|
||||
if total_for_rank > 0 and next_log_pct <= 100:
|
||||
pct = 100 * completed / total_for_rank
|
||||
@@ -860,7 +881,12 @@ def _prefetch_all_checkpoints(sorted_files: list[str]) -> None:
|
||||
"Failed to prefetch checkpoint file %r.", path, exc_info=True
|
||||
)
|
||||
|
||||
await asyncio.gather(*(prefetch_one(p) for p in paths_to_prefetch))
|
||||
with concurrent.futures.ThreadPoolExecutor(
|
||||
max_workers=num_prefetch_threads
|
||||
) as executor:
|
||||
await asyncio.gather(
|
||||
*(prefetch_one(p, executor) for p in paths_to_prefetch)
|
||||
)
|
||||
|
||||
def _run_prefetch() -> None:
|
||||
start = time.perf_counter()
|
||||
@@ -871,7 +897,12 @@ def _prefetch_all_checkpoints(sorted_files: list[str]) -> None:
|
||||
elapsed,
|
||||
)
|
||||
|
||||
logger.info("Prefetching checkpoint files into page cache started (in background)")
|
||||
logger.info(
|
||||
"Prefetching checkpoint files into page cache started "
|
||||
"(in background, num_threads=%d, block_size=%d bytes)",
|
||||
num_prefetch_threads,
|
||||
block_size,
|
||||
)
|
||||
threading.Thread(target=_run_prefetch, daemon=True).start()
|
||||
|
||||
|
||||
@@ -880,6 +911,9 @@ def safetensors_weights_iterator(
|
||||
use_tqdm_on_load: bool,
|
||||
safetensors_load_strategy: str | None = None,
|
||||
local_expert_ids: set[int] | None = None,
|
||||
*,
|
||||
safetensors_prefetch_num_threads: int = DEFAULT_SAFETENSORS_PREFETCH_NUM_THREADS,
|
||||
safetensors_prefetch_block_size: int = DEFAULT_SAFETENSORS_PREFETCH_BLOCK_SIZE,
|
||||
) -> Generator[tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model safetensor files.
|
||||
|
||||
@@ -951,7 +985,11 @@ def safetensors_weights_iterator(
|
||||
)
|
||||
|
||||
if should_prefetch:
|
||||
_prefetch_all_checkpoints(sorted_files)
|
||||
_prefetch_all_checkpoints(
|
||||
sorted_files,
|
||||
num_prefetch_threads=safetensors_prefetch_num_threads,
|
||||
block_size=safetensors_prefetch_block_size,
|
||||
)
|
||||
|
||||
leftover_state_dict: dict[str, torch.Tensor] = {}
|
||||
for st_file in tqdm(
|
||||
|
||||
@@ -64,6 +64,7 @@ from vllm.model_executor.models.bailing_moe import BailingMLP
|
||||
from vllm.sequence import IntermediateTensors
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.linear_attn import LinearAttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
from .interfaces import HasInnerState, IsHybrid, SupportsPP
|
||||
from .utils import (
|
||||
@@ -444,8 +445,8 @@ class BailingMoELinearAttention(PluggableLayer, MambaBase):
|
||||
# --8<-- [end:bailing_moe_linear_attention]
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "linear_attention"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.LINEAR
|
||||
|
||||
def get_state_shape(self) -> tuple[tuple[int, ...], ...]:
|
||||
"""Return state shape for linear attention cache.
|
||||
|
||||
@@ -12,6 +12,7 @@ from vllm.compilation.decorators import support_torch_compile
|
||||
from vllm.config import VllmConfig, get_current_vllm_config
|
||||
from vllm.distributed import (
|
||||
get_ep_group,
|
||||
get_pp_group,
|
||||
get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
)
|
||||
@@ -49,6 +50,7 @@ from vllm.model_executor.layers.vocab_parallel_embedding import (
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
from vllm.model_executor.model_loader.weight_utils import default_weight_loader
|
||||
from vllm.model_executor.models.interfaces import SupportsPP
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.sequence import IntermediateTensors
|
||||
@@ -57,8 +59,10 @@ from vllm.utils.torch_utils import direct_register_custom_op
|
||||
|
||||
from .utils import (
|
||||
AutoWeightsLoader,
|
||||
PPMissingLayer,
|
||||
WeightsMapper,
|
||||
extract_layer_index,
|
||||
is_pp_missing_parameter,
|
||||
make_layers,
|
||||
maybe_prefix,
|
||||
)
|
||||
@@ -1199,23 +1203,53 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
x: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
input_ids: torch.Tensor | None,
|
||||
post_mix: torch.Tensor | None,
|
||||
res_mix: torch.Tensor | None,
|
||||
residual: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
residual = x
|
||||
x, post, comb = self.hc_pre(
|
||||
x, self.hc_attn_fn, self.hc_attn_scale, self.hc_attn_base
|
||||
)
|
||||
if residual is None:
|
||||
# Run standalone hc_pre on first layer
|
||||
residual = x
|
||||
x, post_mix, res_mix = self.hc_pre(
|
||||
x, self.hc_attn_fn, self.hc_attn_scale, self.hc_attn_base
|
||||
)
|
||||
else:
|
||||
residual, post_mix, res_mix, x = torch.ops.vllm.mhc_fused_post_pre(
|
||||
x,
|
||||
residual,
|
||||
post_mix,
|
||||
res_mix,
|
||||
self.hc_attn_fn,
|
||||
self.hc_attn_scale,
|
||||
self.hc_attn_base,
|
||||
self.rms_norm_eps,
|
||||
self.hc_eps,
|
||||
self.hc_eps,
|
||||
self.hc_post_alpha,
|
||||
self.hc_sinkhorn_iters,
|
||||
)
|
||||
|
||||
x = self.attn_norm(x)
|
||||
x = self.attn(positions, x, None)
|
||||
x = self.hc_post(x, residual, post, comb)
|
||||
|
||||
residual = x
|
||||
x, post, comb = self.hc_pre(
|
||||
x, self.hc_ffn_fn, self.hc_ffn_scale, self.hc_ffn_base
|
||||
residual, post_mix, res_mix, x = torch.ops.vllm.mhc_fused_post_pre(
|
||||
x,
|
||||
residual,
|
||||
post_mix,
|
||||
res_mix,
|
||||
self.hc_ffn_fn,
|
||||
self.hc_ffn_scale,
|
||||
self.hc_ffn_base,
|
||||
self.rms_norm_eps,
|
||||
self.hc_eps,
|
||||
self.hc_eps,
|
||||
self.hc_post_alpha,
|
||||
self.hc_sinkhorn_iters,
|
||||
)
|
||||
|
||||
x = self.ffn_norm(x)
|
||||
x = self.ffn(x, input_ids)
|
||||
x = self.hc_post(x, residual, post, comb)
|
||||
return x
|
||||
return x, residual, post_mix, res_mix
|
||||
|
||||
|
||||
@support_torch_compile
|
||||
@@ -1261,12 +1295,15 @@ class DeepseekV4Model(nn.Module):
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.embed_tokens",
|
||||
)
|
||||
if get_pp_group().is_first_rank:
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.embed_tokens",
|
||||
)
|
||||
else:
|
||||
self.embed_tokens = PPMissingLayer()
|
||||
|
||||
self.start_layer, self.end_layer, self.layers = make_layers(
|
||||
config.num_hidden_layers,
|
||||
@@ -1279,7 +1316,10 @@ class DeepseekV4Model(nn.Module):
|
||||
prefix=f"{prefix}.layers",
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, self.rms_norm_eps)
|
||||
if get_pp_group().is_last_rank:
|
||||
self.norm = RMSNorm(config.hidden_size, self.rms_norm_eps)
|
||||
else:
|
||||
self.norm = PPMissingLayer()
|
||||
|
||||
self.hc_head_fn = nn.Parameter(
|
||||
torch.empty(
|
||||
@@ -1304,16 +1344,42 @@ class DeepseekV4Model(nn.Module):
|
||||
# Pre-hc_head residual stream buffer for the MTP draft. Stable
|
||||
# address (outside the cudagraph pool) so the copy_ in forward()
|
||||
# refreshes it correctly across captured shapes.
|
||||
self._mtp_hidden_buffer = torch.empty(
|
||||
vllm_config.scheduler_config.max_num_batched_tokens,
|
||||
self.hc_dim,
|
||||
dtype=vllm_config.model_config.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
# refreshes it correctly across captured shapes. Only allocated on
|
||||
# the last PP rank — that's where MTP target hidden states are
|
||||
# produced.
|
||||
if get_pp_group().is_last_rank:
|
||||
self._mtp_hidden_buffer = torch.empty(
|
||||
vllm_config.scheduler_config.max_num_batched_tokens,
|
||||
self.hc_dim,
|
||||
dtype=vllm_config.model_config.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
else:
|
||||
self._mtp_hidden_buffer = None
|
||||
|
||||
def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.embed_tokens(input_ids)
|
||||
|
||||
def make_empty_intermediate_tensors(
|
||||
self,
|
||||
batch_size: int,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> IntermediateTensors:
|
||||
# PP intermediate tensors carry the multi-stream hidden_states
|
||||
# of shape (num_tokens, hc_mult, hidden_size) — V4 expands the
|
||||
# token embedding to hc_mult streams before the first decoder
|
||||
# layer and keeps that shape until hc_head() collapses it.
|
||||
return IntermediateTensors(
|
||||
{
|
||||
"hidden_states": torch.zeros(
|
||||
(batch_size, self.hc_mult, self.config.hidden_size),
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
@@ -1321,16 +1387,34 @@ class DeepseekV4Model(nn.Module):
|
||||
intermediate_tensors: IntermediateTensors | None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
) -> torch.Tensor | IntermediateTensors:
|
||||
hidden_states = self.embed_input_ids(input_ids)
|
||||
hidden_states = hidden_states.unsqueeze(-2).repeat(1, self.hc_mult, 1)
|
||||
if get_pp_group().is_first_rank:
|
||||
if inputs_embeds is not None:
|
||||
hidden_states = inputs_embeds
|
||||
else:
|
||||
hidden_states = self.embed_input_ids(input_ids)
|
||||
hidden_states = hidden_states.unsqueeze(-2).repeat(1, self.hc_mult, 1)
|
||||
else:
|
||||
assert intermediate_tensors is not None
|
||||
hidden_states = intermediate_tensors["hidden_states"]
|
||||
|
||||
if self.use_mega_moe:
|
||||
input_ids = input_ids.to(torch.int64)
|
||||
|
||||
residual, post_mix, res_mix = None, None, None
|
||||
for layer in islice(self.layers, self.start_layer, self.end_layer):
|
||||
hidden_states = layer(
|
||||
hidden_states, residual, post_mix, res_mix = layer(
|
||||
hidden_states,
|
||||
positions,
|
||||
input_ids,
|
||||
post_mix,
|
||||
res_mix,
|
||||
residual,
|
||||
)
|
||||
else:
|
||||
hidden_states = layer.hc_post(hidden_states, residual, post_mix, res_mix)
|
||||
|
||||
if not get_pp_group().is_last_rank:
|
||||
return IntermediateTensors({"hidden_states": hidden_states})
|
||||
|
||||
# Stash pre-hc_head residual for the MTP draft (captured copy_).
|
||||
num_tokens = hidden_states.shape[0]
|
||||
@@ -1380,6 +1464,8 @@ class DeepseekV4Model(nn.Module):
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
if is_pp_missing_parameter(name, self):
|
||||
break
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
@@ -1401,6 +1487,8 @@ class DeepseekV4Model(nn.Module):
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name_mapped = name.replace(weight_name, param_name)
|
||||
if is_pp_missing_parameter(name_mapped, self):
|
||||
continue
|
||||
param = params_dict[name_mapped]
|
||||
# We should ask the weight loader to return success or not
|
||||
# here since otherwise we may skip experts with other
|
||||
@@ -1422,12 +1510,16 @@ class DeepseekV4Model(nn.Module):
|
||||
loaded_params.add(name_mapped)
|
||||
continue
|
||||
elif "attn_sink" in name:
|
||||
if is_pp_missing_parameter(name, self):
|
||||
continue
|
||||
narrow_weight = loaded_weight[head_rank_start:head_rank_end]
|
||||
n = narrow_weight.shape[0]
|
||||
params_dict[name][:n].copy_(narrow_weight)
|
||||
loaded_params.add(name)
|
||||
continue
|
||||
else:
|
||||
if is_pp_missing_parameter(name, self):
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(
|
||||
param, "weight_loader", default_weight_loader
|
||||
@@ -1525,7 +1617,7 @@ def _make_deepseek_v4_weights_mapper(expert_dtype: str) -> WeightsMapper:
|
||||
)
|
||||
|
||||
|
||||
class DeepseekV4ForCausalLM(nn.Module):
|
||||
class DeepseekV4ForCausalLM(nn.Module, SupportsPP):
|
||||
model_cls = DeepseekV4Model
|
||||
|
||||
# Default mapper assumes the original FP4-expert checkpoint layout.
|
||||
@@ -1544,12 +1636,18 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
self.model = self.model_cls(
|
||||
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
|
||||
)
|
||||
self.lm_head = ParallelLMHead(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
prefix=maybe_prefix(prefix, "lm_head"),
|
||||
)
|
||||
if get_pp_group().is_last_rank:
|
||||
self.lm_head = ParallelLMHead(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
prefix=maybe_prefix(prefix, "lm_head"),
|
||||
)
|
||||
else:
|
||||
self.lm_head = PPMissingLayer()
|
||||
self.logits_processor = LogitsProcessor(config.vocab_size)
|
||||
self.make_empty_intermediate_tensors = (
|
||||
self.model.make_empty_intermediate_tensors
|
||||
)
|
||||
|
||||
def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.model.embed_input_ids(input_ids)
|
||||
|
||||
@@ -1338,6 +1338,9 @@ def exif_transpose(
|
||||
def build_flat_image_bool_length(
|
||||
image_grids: torch.LongTensor,
|
||||
hf_config: PretrainedConfig,
|
||||
image_use_col_tokens: bool = True,
|
||||
use_single_crop_col_tokens: bool | None = None,
|
||||
use_single_crop_start_token: bool = True,
|
||||
) -> tuple[torch.LongTensor, torch.LongTensor]:
|
||||
image_patch_id = hf_config.image_patch_id
|
||||
low_res_image_start_id = hf_config.low_res_image_start_token_id
|
||||
@@ -1353,7 +1356,17 @@ def build_flat_image_bool_length(
|
||||
h = image_grids[:, 2]
|
||||
w = image_grids[:, 3]
|
||||
|
||||
lengths = resized_h * resized_w + h * (w + 1) + 4 # [B]
|
||||
low_res_use_col_tokens = (
|
||||
image_use_col_tokens
|
||||
if use_single_crop_col_tokens is None
|
||||
else use_single_crop_col_tokens
|
||||
)
|
||||
low_res_extra = int(low_res_use_col_tokens)
|
||||
high_res_extra = int(image_use_col_tokens)
|
||||
|
||||
lengths = (
|
||||
resized_h * (resized_w + low_res_extra) + h * (w + high_res_extra) + 4
|
||||
) # [B]
|
||||
total_len = int(lengths.sum().item())
|
||||
|
||||
flat = torch.empty(total_len, dtype=torch.long, device=device)
|
||||
@@ -1363,16 +1376,24 @@ def build_flat_image_bool_length(
|
||||
resized_h_i, resized_w_i, h_i, w_i = image_grids[i].tolist()
|
||||
L_i = int(lengths[i].item())
|
||||
|
||||
num_low_res_patches = resized_h_i * resized_w_i
|
||||
|
||||
idx = offset
|
||||
|
||||
flat[idx] = low_res_image_start_id
|
||||
flat[idx] = (
|
||||
low_res_image_start_id if use_single_crop_start_token else image_start_id
|
||||
)
|
||||
idx += 1
|
||||
|
||||
if num_low_res_patches > 0:
|
||||
flat[idx : idx + num_low_res_patches] = image_patch_id
|
||||
idx += num_low_res_patches
|
||||
low_res_block_len = resized_w_i + low_res_extra
|
||||
if low_res_block_len > 0 and resized_h_i > 0:
|
||||
line = torch.empty(low_res_block_len, dtype=torch.long, device=device)
|
||||
if resized_w_i > 0:
|
||||
line[:resized_w_i] = image_patch_id
|
||||
if low_res_use_col_tokens:
|
||||
line[resized_w_i] = image_col_id
|
||||
|
||||
block = line.repeat(resized_h_i)
|
||||
flat[idx : idx + resized_h_i * low_res_block_len] = block
|
||||
idx += resized_h_i * low_res_block_len
|
||||
|
||||
flat[idx] = image_end_id
|
||||
idx += 1
|
||||
@@ -1380,12 +1401,13 @@ def build_flat_image_bool_length(
|
||||
flat[idx] = image_start_id
|
||||
idx += 1
|
||||
|
||||
block_len = w_i + 1
|
||||
block_len = w_i + high_res_extra
|
||||
if block_len > 0 and h_i > 0:
|
||||
line = torch.empty(block_len, dtype=torch.long, device=device)
|
||||
if w_i > 0:
|
||||
line[:w_i] = image_patch_id
|
||||
line[w_i] = image_col_id
|
||||
if image_use_col_tokens:
|
||||
line[w_i] = image_col_id
|
||||
|
||||
block = line.repeat(h_i)
|
||||
flat[idx : idx + h_i * block_len] = block
|
||||
@@ -2108,7 +2130,13 @@ class Molmo2MultiModalProcessor(BaseMultiModalProcessor[Molmo2ProcessingInfo]):
|
||||
(
|
||||
processed_outputs["image_tokens"],
|
||||
processed_outputs["num_image_tokens"],
|
||||
) = build_flat_image_bool_length(image_grids, hf_config)
|
||||
) = build_flat_image_bool_length(
|
||||
image_grids,
|
||||
hf_config,
|
||||
image_use_col_tokens=hf_processor.image_use_col_tokens,
|
||||
use_single_crop_col_tokens=hf_processor.use_single_crop_col_tokens,
|
||||
use_single_crop_start_token=hf_processor.use_single_crop_start_token,
|
||||
)
|
||||
|
||||
return BatchFeature({**processed_outputs, **all_video_outputs})
|
||||
|
||||
|
||||
@@ -92,6 +92,7 @@ from vllm.triton_utils.allocation import set_triton_allocator
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
from .interfaces import HasInnerState, IsHybrid, SupportsLoRA, SupportsPP
|
||||
from .utils import (
|
||||
@@ -136,8 +137,8 @@ class OlmoHybridGatedDeltaNet(nn.Module, MambaBase):
|
||||
"""
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "gdn_attention"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.GDN_ATTN
|
||||
|
||||
def get_state_dtype(self) -> tuple[torch.dtype, torch.dtype]:
|
||||
return MambaStateDtypeCalculator.gated_delta_net_state_dtype(
|
||||
|
||||
@@ -72,6 +72,7 @@ from vllm.sequence import IntermediateTensors
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
from vllm.v1.attention.backend import AttentionMetadata
|
||||
from vllm.v1.attention.backends.mamba2_attn import Mamba2AttentionMetadata
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
# Only used for type hinting.
|
||||
if TYPE_CHECKING:
|
||||
@@ -478,8 +479,8 @@ class Plamo2MambaMixer(MambaBase, PluggableLayer):
|
||||
)
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
return "mamba2"
|
||||
def mamba_type(self) -> MambaAttentionBackendEnum:
|
||||
return MambaAttentionBackendEnum.MAMBA2
|
||||
|
||||
|
||||
def plamo2_mamba_mixer(
|
||||
|
||||
@@ -138,7 +138,6 @@ class Qwen3_5DecoderLayer(Qwen3NextDecoderLayer):
|
||||
vllm_config=vllm_config,
|
||||
prefix=f"{prefix}.linear_attn",
|
||||
gqa_interleaved_layout=False,
|
||||
create_in_proj_qkvz=vllm_config.lora_config is None,
|
||||
)
|
||||
elif self.layer_type == "full_attention":
|
||||
self.self_attn = Qwen3NextAttention(
|
||||
@@ -217,7 +216,6 @@ class Qwen3_5Model(Qwen3NextModel):
|
||||
self.num_redundant_experts = eplb_config.num_redundant_experts
|
||||
|
||||
self.config = config
|
||||
self.enable_lora = vllm_config.lora_config is not None
|
||||
|
||||
self.vocab_size = config.vocab_size
|
||||
|
||||
@@ -276,6 +274,9 @@ class Qwen3_5Model(Qwen3NextModel):
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
# GDN
|
||||
("in_proj_qkvz", "in_proj_qkv", (0, 1, 2)),
|
||||
("in_proj_qkvz", "in_proj_z", 3),
|
||||
# self attention
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
@@ -287,21 +288,6 @@ class Qwen3_5Model(Qwen3NextModel):
|
||||
("in_proj_ba", "in_proj_a", 1),
|
||||
]
|
||||
|
||||
if self.enable_lora:
|
||||
stacked_params_mapping.extend(
|
||||
[
|
||||
("in_proj_qkv", "in_proj_qkv", (0, 1, 2)),
|
||||
("in_proj_z", "in_proj_z", 0),
|
||||
]
|
||||
)
|
||||
else:
|
||||
stacked_params_mapping.extend(
|
||||
[
|
||||
("in_proj_qkvz", "in_proj_qkv", (0, 1, 2)),
|
||||
("in_proj_qkvz", "in_proj_z", 3),
|
||||
]
|
||||
)
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
expert_params_mapping = self.get_expert_mapping()
|
||||
@@ -352,10 +338,7 @@ class Qwen3_5Model(Qwen3NextModel):
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
if param_name == "in_proj_z" and self.enable_lora:
|
||||
weight_loader(param, loaded_weight)
|
||||
else:
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
is_expert_weight = False
|
||||
@@ -485,15 +468,6 @@ class Qwen3_5ForCausalLMBase(
|
||||
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
|
||||
)
|
||||
|
||||
# When LoRA is enabled, GDN uses separate in_proj_qkv and in_proj_z
|
||||
# instead of merged in_proj_qkvz; pack mapping must match.
|
||||
if vllm_config.lora_config:
|
||||
base = getattr(Qwen3_5ForCausalLMBase, "packed_modules_mapping", {})
|
||||
self.packed_modules_mapping = {k: list(v) for k, v in base.items()}
|
||||
self.packed_modules_mapping.pop("in_proj_qkvz", None)
|
||||
self.packed_modules_mapping["in_proj_qkv"] = ["in_proj_qkv"]
|
||||
self.packed_modules_mapping["in_proj_z"] = ["in_proj_z"]
|
||||
|
||||
if get_pp_group().is_last_rank:
|
||||
if config.tie_word_embeddings:
|
||||
self.lm_head = self.model.embed_tokens
|
||||
@@ -586,7 +560,6 @@ class Qwen3_5ForConditionalGeneration(Qwen3VLForConditionalGeneration, IsHybrid)
|
||||
def __init__(self, *, vllm_config: VllmConfig, prefix: str = "model"):
|
||||
# protocols have not __init__ method, so we need to use nn.Module.__init__
|
||||
nn.Module.__init__(self)
|
||||
self.update_packed_mapping(enable_lora=vllm_config.lora_config is not None)
|
||||
config: Qwen3_5Config = vllm_config.model_config.hf_config
|
||||
quant_config = vllm_config.quant_config
|
||||
multimodal_config = vllm_config.model_config.multimodal_config
|
||||
@@ -614,17 +587,6 @@ class Qwen3_5ForConditionalGeneration(Qwen3VLForConditionalGeneration, IsHybrid)
|
||||
self.language_model.make_empty_intermediate_tensors
|
||||
)
|
||||
|
||||
def update_packed_mapping(self, enable_lora: bool):
|
||||
# When LoRA is enabled, GDN uses separate in_proj_qkv and in_proj_z
|
||||
if enable_lora:
|
||||
base = getattr(
|
||||
Qwen3_5ForConditionalGeneration, "packed_modules_mapping", {}
|
||||
)
|
||||
self.packed_modules_mapping = {k: list(v) for k, v in base.items()}
|
||||
self.packed_modules_mapping.pop("in_proj_qkvz", None)
|
||||
self.packed_modules_mapping["in_proj_qkv"] = ["in_proj_qkv"]
|
||||
self.packed_modules_mapping["in_proj_z"] = ["in_proj_z"]
|
||||
|
||||
def embed_input_ids(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
@@ -811,7 +773,6 @@ class Qwen3_5MoeForConditionalGeneration(
|
||||
def __init__(self, *, vllm_config: VllmConfig, prefix: str = "model"):
|
||||
# protocols have not __init__ method, so we need to use nn.Module.__init__
|
||||
nn.Module.__init__(self)
|
||||
self.update_packed_mapping(enable_lora=vllm_config.lora_config is not None)
|
||||
config: Qwen3_5MoeConfig = vllm_config.model_config.hf_config
|
||||
quant_config = vllm_config.quant_config
|
||||
multimodal_config = vllm_config.model_config.multimodal_config
|
||||
|
||||
@@ -204,8 +204,16 @@ def _rejection_greedy_sample_kernel_impl(
|
||||
bonus_token_ids,
|
||||
is_greedy,
|
||||
max_spec_len,
|
||||
uniform_probs=None,
|
||||
synthetic_conditional_rates=None,
|
||||
SYNTHETIC_MODE=False,
|
||||
):
|
||||
# C++ kernel expects int64 for all integer tensors.
|
||||
# Note: uniform_probs, synthetic_conditional_rates, and SYNTHETIC_MODE are
|
||||
# passed by the rejection sampler for synthetic mode support, but are not
|
||||
# yet implemented in the C++ CPU kernel. We accept them here to maintain
|
||||
# compatibility with the kernel calling convention.
|
||||
assert not SYNTHETIC_MODE, "Synthetic acceptance not supported with CPU sampling"
|
||||
orig_dtype = output_token_ids.dtype
|
||||
output_token_ids_i64 = _ensure_int64(output_token_ids)
|
||||
torch.ops._C.rejection_greedy_sample_kernel_impl(
|
||||
@@ -233,11 +241,18 @@ def _rejection_random_sample_kernel_impl(
|
||||
is_greedy,
|
||||
max_spec_len,
|
||||
vocab_size,
|
||||
synthetic_conditional_rates=None,
|
||||
NO_DRAFT_PROBS=False,
|
||||
SYNTHETIC_MODE=False,
|
||||
):
|
||||
# C++ kernel expects int64 for all integer tensors and float32 for probs.
|
||||
# uniform_probs is intentionally float64 in Python to avoid exact-zero
|
||||
# samples; cast to float32 here for C++ compatibility.
|
||||
# Note: synthetic_conditional_rates and SYNTHETIC_MODE are passed by the
|
||||
# rejection sampler for synthetic mode support, but are not yet implemented
|
||||
# in the C++ CPU kernel. We accept them here to maintain compatibility with
|
||||
# the kernel calling convention.
|
||||
assert not SYNTHETIC_MODE, "Synthetic acceptance not supported with CPU sampling"
|
||||
orig_dtype = output_token_ids.dtype
|
||||
output_token_ids_i64 = _ensure_int64(output_token_ids)
|
||||
torch.ops._C.rejection_random_sample_kernel_impl(
|
||||
|
||||
@@ -193,16 +193,6 @@ class MambaAttentionBackendEnum(Enum, metaclass=_AttentionBackendEnumMeta):
|
||||
_MAMBA_ATTN_OVERRIDES.pop(self, None)
|
||||
|
||||
|
||||
MAMBA_TYPE_TO_BACKEND_MAP = {
|
||||
"mamba1": MambaAttentionBackendEnum.MAMBA1.name,
|
||||
"mamba2": MambaAttentionBackendEnum.MAMBA2.name,
|
||||
"short_conv": MambaAttentionBackendEnum.SHORT_CONV.name,
|
||||
"linear_attention": MambaAttentionBackendEnum.LINEAR.name,
|
||||
"gdn_attention": MambaAttentionBackendEnum.GDN_ATTN.name,
|
||||
"custom": MambaAttentionBackendEnum.CUSTOM.name,
|
||||
}
|
||||
|
||||
|
||||
_ATTN_OVERRIDES: dict[AttentionBackendEnum, str] = {}
|
||||
_MAMBA_ATTN_OVERRIDES: dict[MambaAttentionBackendEnum, str] = {}
|
||||
|
||||
|
||||
@@ -459,25 +459,29 @@ def _decode_grouped_att_m_fwd(
|
||||
):
|
||||
# with is_mla there is only a single c_kv in smem.
|
||||
# could increase BLOCK or num_stages.
|
||||
BLOCK = 32
|
||||
Lk = k_buffer.shape[-1]
|
||||
Lv = v_buffer.shape[-1]
|
||||
|
||||
# [TODO] work around shmem limit on MI3xx
|
||||
if is_hip_ and Lk >= 576:
|
||||
BLOCK = 16
|
||||
|
||||
if Lk == 576:
|
||||
BLOCK_DMODEL = 512
|
||||
BLOCK_DPE = 64
|
||||
elif Lk == 288:
|
||||
BLOCK_DMODEL = 256
|
||||
BLOCK_DPE = 32
|
||||
# Align tile dimensions with latent rank for MLA to avoid shape mismatch.
|
||||
if is_mla:
|
||||
if not is_hip_ and Lk == 576:
|
||||
BLOCK_DMODEL = 512
|
||||
BLOCK_DPE = 64
|
||||
elif not is_hip_ and Lk == 288:
|
||||
BLOCK_DMODEL = 256
|
||||
BLOCK_DPE = 32
|
||||
else:
|
||||
BLOCK_DMODEL = triton.next_power_of_2(Lv)
|
||||
BLOCK_DPE = triton.next_power_of_2(Lk - Lv) if Lk > Lv else 0
|
||||
else:
|
||||
BLOCK_DMODEL = triton.next_power_of_2(Lk)
|
||||
BLOCK_DPE = 0
|
||||
BLOCK_DV = triton.next_power_of_2(Lv)
|
||||
|
||||
BLOCK = 32
|
||||
if is_hip_:
|
||||
BLOCK = 16
|
||||
|
||||
batch, head_num = q.shape[0], q.shape[1]
|
||||
kv_group_num = q.shape[1] // k_buffer.shape[-2]
|
||||
|
||||
@@ -496,6 +500,11 @@ def _decode_grouped_att_m_fwd(
|
||||
# https://github.com/triton-lang/triton/blob/main/third_party/amd/backend/compiler.py
|
||||
extra_kargs = {"waves_per_eu": 1, "matrix_instr_nonkdim": 16, "kpack": 2}
|
||||
num_stages = 1
|
||||
elif not is_hip_ and BLOCK_DMODEL >= 1024:
|
||||
# Avoid shared memory overflow on NVIDIA when BLOCK_DMODEL is large
|
||||
# like non-MLA D_QK=576, BLOCK_DMODEL=1024, BLOCK_H=16
|
||||
# exceeds 101376 bytes limit
|
||||
num_stages = 1
|
||||
|
||||
_fwd_grouped_kernel_stage1[grid](
|
||||
q,
|
||||
|
||||
@@ -12,7 +12,6 @@ from vllm.logger import init_logger
|
||||
from vllm.utils.import_utils import resolve_obj_by_qualname
|
||||
from vllm.v1.attention.backend import AttentionBackend, AttentionType
|
||||
from vllm.v1.attention.backends.registry import (
|
||||
MAMBA_TYPE_TO_BACKEND_MAP,
|
||||
MambaAttentionBackendEnum,
|
||||
)
|
||||
|
||||
@@ -138,7 +137,7 @@ def _cached_get_attn_backend(
|
||||
|
||||
|
||||
def get_mamba_attn_backend(
|
||||
mamba_type: str,
|
||||
mamba_type: MambaAttentionBackendEnum,
|
||||
) -> type[AttentionBackend]:
|
||||
"""Select which mamba attention backend to use and lazily import it."""
|
||||
return _cached_get_mamba_attn_backend(mamba_type)
|
||||
@@ -146,21 +145,11 @@ def get_mamba_attn_backend(
|
||||
|
||||
@cache
|
||||
def _cached_get_mamba_attn_backend(
|
||||
mamba_type: str,
|
||||
mamba_type: MambaAttentionBackendEnum,
|
||||
) -> type[AttentionBackend]:
|
||||
assert mamba_type and isinstance(mamba_type, str)
|
||||
assert mamba_type and isinstance(mamba_type, MambaAttentionBackendEnum)
|
||||
|
||||
selected_backend = None
|
||||
try:
|
||||
backend_name = MAMBA_TYPE_TO_BACKEND_MAP[mamba_type]
|
||||
selected_backend = MambaAttentionBackendEnum[backend_name]
|
||||
except KeyError as e:
|
||||
raise ValueError(
|
||||
f"Invalid mamba attention backend type: '{mamba_type}'. Valid "
|
||||
f"types are: {list(MAMBA_TYPE_TO_BACKEND_MAP.keys())}"
|
||||
) from e
|
||||
|
||||
mamba_attn_backend = selected_backend.get_class()
|
||||
mamba_attn_backend = mamba_type.get_class()
|
||||
if envs.VLLM_BATCH_INVARIANT and not mamba_attn_backend.supports_batch_invariance():
|
||||
raise RuntimeError(
|
||||
"VLLM batch_invariant mode is not supported for "
|
||||
|
||||
@@ -16,6 +16,7 @@ from typing_extensions import Self
|
||||
from vllm.logger import init_logger
|
||||
from vllm.utils.math_utils import cdiv, round_up
|
||||
from vllm.utils.torch_utils import get_dtype_size, nvfp4_kv_cache_full_dim
|
||||
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.config import VllmConfig
|
||||
@@ -421,6 +422,17 @@ class SlidingWindowSpec(AttentionSpec):
|
||||
|
||||
@property
|
||||
def real_page_size_bytes(self) -> int:
|
||||
# Mirror ``FullAttentionSpec.real_page_size_bytes`` for NVFP4 KV cache.
|
||||
if self.kv_quant_mode.is_nvfp4:
|
||||
last_dim = nvfp4_kv_cache_full_dim(
|
||||
self.head_size
|
||||
) + nvfp4_kv_cache_full_dim(self.head_size_v)
|
||||
return (
|
||||
self.block_size
|
||||
* self.num_kv_heads
|
||||
* last_dim
|
||||
* get_dtype_size(self.dtype)
|
||||
)
|
||||
return (
|
||||
self.block_size
|
||||
* self.num_kv_heads
|
||||
@@ -532,7 +544,7 @@ class MambaSpec(KVCacheSpec):
|
||||
shapes: tuple[tuple[int, ...], ...]
|
||||
dtypes: tuple[torch.dtype]
|
||||
page_size_padded: int | None = None
|
||||
mamba_type: str = "mamba2"
|
||||
mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2
|
||||
mamba_cache_mode: str = "none"
|
||||
num_speculative_blocks: int = 0
|
||||
|
||||
|
||||
@@ -147,22 +147,24 @@ class OffloadingManager(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def touch(self, keys: Collection[OffloadKey]):
|
||||
def touch(self, keys: Collection[OffloadKey], req_context: ReqContext):
|
||||
"""
|
||||
Mark the given blocks as recently used.
|
||||
This could in practice mean moving them to the end of an LRU list.
|
||||
|
||||
Args:
|
||||
keys: the keys identifying the blocks.
|
||||
req_context: per-request context (e.g. kv_transfer_params).
|
||||
"""
|
||||
return
|
||||
|
||||
def complete_load(self, keys: Collection[OffloadKey]):
|
||||
def complete_load(self, keys: Collection[OffloadKey], req_context: ReqContext):
|
||||
"""
|
||||
Marks previous blocks that were prepared to load as done loading.
|
||||
|
||||
Args:
|
||||
keys: the keys identifying the blocks.
|
||||
req_context: per-request context (e.g. kv_transfer_params).
|
||||
"""
|
||||
return
|
||||
|
||||
@@ -189,7 +191,12 @@ class OffloadingManager(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def complete_store(self, keys: Collection[OffloadKey], success: bool = True):
|
||||
def complete_store(
|
||||
self,
|
||||
keys: Collection[OffloadKey],
|
||||
req_context: ReqContext,
|
||||
success: bool = True,
|
||||
):
|
||||
"""
|
||||
Marks blocks which were previously prepared to be stored, as stored.
|
||||
Following this call, the blocks become loadable.
|
||||
@@ -198,6 +205,7 @@ class OffloadingManager(ABC):
|
||||
|
||||
Args:
|
||||
keys: the keys identifying the blocks.
|
||||
req_context: per-request context (e.g. kv_transfer_params).
|
||||
success: whether the blocks were stored successfully.
|
||||
"""
|
||||
return
|
||||
|
||||
@@ -106,10 +106,12 @@ class CPUOffloadingManager(OffloadingManager):
|
||||
blocks.append(block)
|
||||
return self._get_load_store_spec(keys, blocks)
|
||||
|
||||
def touch(self, keys: Collection[OffloadKey]) -> None:
|
||||
def touch(self, keys: Collection[OffloadKey], req_context: ReqContext) -> None:
|
||||
self._policy.touch(keys)
|
||||
|
||||
def complete_load(self, keys: Collection[OffloadKey]) -> None:
|
||||
def complete_load(
|
||||
self, keys: Collection[OffloadKey], req_context: ReqContext
|
||||
) -> None:
|
||||
for key in keys:
|
||||
block = self._policy.get(key)
|
||||
assert block is not None, f"Block {key!r} not found"
|
||||
@@ -172,7 +174,10 @@ class CPUOffloadingManager(OffloadingManager):
|
||||
)
|
||||
|
||||
def complete_store(
|
||||
self, keys: Collection[OffloadKey], success: bool = True
|
||||
self,
|
||||
keys: Collection[OffloadKey],
|
||||
req_context: ReqContext,
|
||||
success: bool = True,
|
||||
) -> None:
|
||||
stored_keys: list[OffloadKey] = []
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user