Compare commits

...
Author SHA1 Message Date
Jiangyun Zhuandkhluu 135453b715 [Bugfix] Install nvidia-cutlass-dsl[cu13] extra on CUDA 13 platforms (#42438)
Signed-off-by: zjy0516 <riverclouds.zhu@qq.com>
(cherry picked from commit 140dc2ec30)
2026-05-13 02:03:17 -07:00
sychen52andkhluu a707288c1e Patch SlidingWindowSpec.real_page_size_bytes for nvfp4 kv (#42464)
Signed-off-by: Shiyang Chen <shiychen@nvidia.com>
(cherry picked from commit a8c13d2837)
2026-05-13 02:03:07 -07:00
Alecandkhluu 638f8fa979 [PD] Bump NIXL connector dependency to 1.x (#42364)
Signed-off-by: Alec Flowers <aflowers@nvidia.com>
(cherry picked from commit 07534b8782)
2026-05-13 02:02:55 -07:00
Chao Leiandkhluu cbaa80fede [KV Transfer] Add MooncakeStoreConnector for KV cache offloading via Mooncake distributed store (#40900)
Signed-off-by: leichao.lc <leichao.lc@antgroup.com>
Signed-off-by: Yifan Qiao <yifanqiao@inferact.ai>
Co-authored-by: leichao.lc <leichao.lc@antgroup.com>
Co-authored-by: ivanium <yifanqiao@inferact.ai>
Co-authored-by: aoshen524 <aoshen@inferact.ai>
Co-authored-by: Dao007forever <daole@inferact.ai>
Co-authored-by: Teng Ma <sima.mt@alibaba-inc.com>
Co-authored-by: Pz1116 <zpbzpb123123@gmail.com>
Co-authored-by: foraxe <1055696449@qq.com>
Co-authored-by: Skywalker-EP <173423846@qq.com>
Co-authored-by: fems14 <1804143737@qq.com>
Co-authored-by: jianzs <zheng.shoujian@outlook.com>
Co-authored-by: baxingpiaochong <771405853@qq.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
(cherry picked from commit ebeb09d822)
2026-05-13 02:02:44 -07:00
Kevin H. Luu 84a1066ccc [CI] Inline build artifact annotations in release pipeline (#42357)
Signed-off-by: khluu <khluu000@gmail.com>
(cherry picked from commit 8c4fc4202a)
2026-05-13 02:02:30 -07:00
Michael Goinandkhluu d801ae8c26 [Build] Build bundled DeepGEMM _C per-Python so the wheel imports on every CPython (#41516)
Signed-off-by: mgoin <mgoin64@gmail.com>
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
(cherry picked from commit d077622d60)
2026-05-12 14:57:17 -07:00
Jiahan Chang (Cyrus)andkhluu 65df49eba3 [Perf] Use 2D-grid to eliminate divmod in W8W8 group quant (#42153)
Signed-off-by: jiahanc <173873397+jiahanc@users.noreply.github.com>
Co-authored-by: Yongye Zhu <zyy1102000@gmail.com>
(cherry picked from commit dd6b3a5ef5)
2026-05-12 14:57:06 -07:00
Kevin H. Luu 2a2ac21d3d [CI] Move DockerHub and PyPI publish steps to end of release pipeline (#42355)
Signed-off-by: khluu <khluu000@gmail.com>
(cherry picked from commit e1c8776e90)
2026-05-12 14:56:46 -07:00
Jee Jee Liandkhluu c6fc95806b [Bugfix] Fix DSV4 swiglu_limit on marlin backend (#42287)
Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
Co-authored-by: Yongye Zhu <zyy1102000@gmail.com>
(cherry picked from commit 53181384e0)
2026-05-12 14:56:29 -07:00
汪志鹏GitHubWang, Zhipeng | RASIA <zhipeng.wang@rakuten.com>CursorChauncey
581b5e9afc [Frontend] Return rendered prompt text in chat completion response (#42052)
Signed-off-by: Wang, Zhipeng | RASIA <zhipeng.wang@rakuten.com>
Co-authored-by: Wang, Zhipeng | RASIA <zhipeng.wang@rakuten.com>
Co-authored-by: Cursor <cursor@cursor.com>
Co-authored-by: Chauncey <chaunceyjiang@gmail.com>
2026-05-11 13:53:39 +08:00
wangxiyuanandGitHub 5536fc0c01 [Misc] Replace mamba_type string literals with MambaAttentionBackendEnum (#41188)
Signed-off-by: wangxiyuan <wangxiyuan1007@gmail.com>
2026-05-11 03:59:36 +00:00
vllmellmandGitHub 7f95e66a11 [ROCm][Bugfix]: dynamically align BLOCK_DMODEL with Lv in MLA decode kernel (#41119)
Signed-off-by: vllmellm <vllm.ellm@embeddedllm.com>
2026-05-11 11:14:19 +08:00
yzong-rhandGitHub b1687527b8 [Bugfix] Gemma 4 chat template crash with missing tool name and tool id (#42188)
Signed-off-by: Yifan <yzong@redhat.com>
2026-05-11 03:07:45 +00:00
171019ab19 add fused mhc_post_pre kernel (#41536)
Signed-off-by: george <george@inferact.ai>
Co-authored-by: george <george@inferact.ai>
2026-05-10 19:56:52 -07:00
Haoqing WangandGitHub 879a8c3180 Fix Molmo2 image token metadata (#42162)
Signed-off-by: Haoqi Wang <78337154+hqhq1025@users.noreply.github.com>
2026-05-11 01:19:21 +00:00
1b57eb41f2 [MoE] Move various experts classes to fused_moe/experts/ (#41979)
Signed-off-by: Jackmin801 <ongjackm@gmail.com>
Signed-off-by: Robert Shaw <robertgshaw2@gmail.com>
Signed-off-by: Jackmin801 <56836461+Jackmin801@users.noreply.github.com>
Signed-off-by: Bill Nell <bnell@redhat.com>
Co-authored-by: Jackmin801 <ongjackm@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Robert Shaw <robertgshaw2@gmail.com>
Co-authored-by: Robert Shaw <114415538+robertgshaw2-redhat@users.noreply.github.com>
Co-authored-by: Jackmin801 <56836461+Jackmin801@users.noreply.github.com>
2026-05-11 07:54:33 +08:00
Mohammad Miadh AngkadandGitHub 21943d4c25 [Performance] Make safetensors checkpoint prefetch settings configurable (#41499)
Signed-off-by: Mohammad Miadh Angkad <MAngkad.BSDSBA2027@aim.edu>
2026-05-10 15:55:15 +00:00
f396bee56f [DSV4] Add PP support for deepseek-v4 (#41694)
Signed-off-by: Isotr0py <mozf@mail2.sysu.edu.cn>
Co-authored-by: qizixi <22851944+zixi-qi@users.noreply.github.com>
2026-05-10 15:47:26 +00:00
VensenGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
215e2f7990 [Bugfix][Mamba] IMA in causal_conv1d kernel for long sequences (#41617)
Signed-off-by: vensen <vensenmu@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-05-10 12:38:28 +00:00
Ronen SchafferandGitHub e175192d33 [KV Offload] Pass ReqContext to touch(), complete_load(), and complete_store() (#41366)
Signed-off-by: Ronen Schaffer <ronen.schaffer@ibm.com>
2026-05-10 15:09:25 +03:00
a54f0d1049 [CPU] Fix spec decode kernel signatures for synthetic mode compatibility (#41932)
Signed-off-by: jmamou <jonathan.mamou@intel.com>
Signed-off-by: Jonathan Mamou <jonathan.mamou@intel.com>
Co-authored-by: Benjamin Chislett <chislett.ben@gmail.com>
2026-05-10 12:07:15 +00:00
Isotr0pyandGitHub 48698b1b9b [Bugfix] Fuse Qwen3.5 in_qkvz_proj forwarding with LoRA enabled (#37912)
Signed-off-by: Isotr0py <mozf@mail2.sysu.edu.cn>
Signed-off-by: Isotr0py <Isotr0py@outlook.com>
2026-05-10 10:59:02 +00:00
101 changed files with 4704 additions and 1022 deletions
+1
View File
@@ -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"
+79 -64
View File
@@ -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
View File
@@ -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'
-27
View File
@@ -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
+5
View File
@@ -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
+3
View File
@@ -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
+59 -40
View File
@@ -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 -1
View File
@@ -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
View File
@@ -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
+1 -1
View File
@@ -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.
+1 -1
View File
@@ -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:**
+3 -3
View File
@@ -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.
+1 -1
View File
@@ -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. | ✅︎ | ✅︎ |
+2 -2
View File
@@ -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'] -%}
+1 -1
View File
@@ -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 -3
View File
@@ -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
+3
View File
@@ -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():
+10 -2
View File
@@ -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,
)
+3 -3
View File
@@ -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,
+1 -1
View File
@@ -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 (
+1 -1
View File
@@ -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,
)
+1 -1
View File
@@ -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(
+3 -1
View File
@@ -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
+142
View File
@@ -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
+32 -32
View File
@@ -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)
+87
View File
@@ -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)
+41
View File
@@ -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)
+49
View File
@@ -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#:}"
+12
View File
@@ -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}. "
+17 -1
View File
@@ -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,
)
+49 -11
View File
@@ -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)
+5
View File
@@ -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,
)
@@ -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,
)
+3 -2
View File
@@ -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,
+2 -1
View File
@@ -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(
+343
View File
@@ -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,
)
@@ -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 (
@@ -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.
+130 -32
View File
@@ -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)
+38 -10
View File
@@ -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})
+3 -2
View File
@@ -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(
+3 -2
View File
@@ -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(
+4 -43
View File
@@ -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
+15
View File
@@ -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(
-10
View File
@@ -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,
+4 -15
View File
@@ -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 "
+13 -1
View File
@@ -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
+11 -3
View File
@@ -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
+8 -3
View File
@@ -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