Compare commits
396
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
+11 |
f68f4fddea | ||
|
|
0ba2aa35a8 | ||
|
|
aaaeda98dc | ||
|
|
d9cd774198 | ||
|
|
70052fb924 | ||
|
|
318b527cc2 | ||
|
|
6a1acac3fe | ||
|
|
213f681f81 | ||
|
|
33c4f3551c | ||
|
|
caa9cad31e | ||
|
|
7513d071bd | ||
|
|
89f6aa3a9e | ||
|
|
84d26b9ee3 | ||
|
|
9e6746b3c7 | ||
|
|
972848f276 | ||
|
|
2279575cd9 | ||
|
|
5d8e90a966 | ||
|
|
e222c33f2f | ||
|
|
c064fa52b6 | ||
|
|
7e51939e25 | ||
|
|
9863102ed9 | ||
|
|
8c13ee5735 | ||
|
|
41798069f3 | ||
|
|
866fea2b99 | ||
|
|
d02df748bf | ||
|
|
453f01783d | ||
|
|
7b40fb9645 | ||
|
|
8eac21a602 | ||
|
|
a454a1dd25 | ||
|
|
833483f357 | ||
|
|
163ecba377 | ||
|
|
589a5b884b | ||
|
|
5c5434e2d8 | ||
|
|
dd72658e7d | ||
|
|
0d77325b10 | ||
|
|
2ac125123a | ||
|
|
7bdf8cc37c | ||
|
|
bf27e34ebb | ||
|
|
275556c35c | ||
|
|
d65acd83d8 | ||
|
|
80c9d5d5e0 | ||
|
|
da54a5bf05 | ||
|
|
1479bd9e9d | ||
|
|
0231dd5467 | ||
|
|
2659467497 | ||
|
|
a49d37c6b9 | ||
|
|
4501a6d56b | ||
|
|
e18f0037a5 | ||
|
|
b354734d17 | ||
|
|
b91a40e729 | ||
|
|
75ccdf3145 | ||
|
|
c6fe94b4d5 | ||
|
|
46f01a50ac | ||
|
|
f00efc5265 | ||
|
|
b0cb1da1bd | ||
|
|
0e36e3bbd1 | ||
|
|
494845e79f | ||
|
|
0416dab275 | ||
|
|
c8db00b16c | ||
|
|
80c7683923 | ||
|
|
638d6e9757 | ||
|
|
1ad84fea86 | ||
|
|
12213c6795 | ||
|
|
10c75477b0 | ||
|
|
ac36a7a1e7 | ||
|
|
521aa80f71 | ||
|
|
a76df87db8 | ||
|
|
a4904ba903 | ||
|
|
f83de6d44c | ||
|
|
239fc73553 | ||
|
|
76bf55240c | ||
|
|
9a698f3255 | ||
|
|
4080263bb2 | ||
|
|
fc5fda105f | ||
|
|
b07ec92faa | ||
|
|
27ffbfde8d | ||
|
|
229e01e9e1 | ||
|
|
191146dba5 | ||
|
|
f3a920a076 | ||
|
|
149daf0d72 | ||
|
|
917fdb5bf7 | ||
|
|
4b594b4aa1 | ||
|
|
7d10a4cfce | ||
|
|
910cc8543a | ||
|
|
431934522b | ||
|
|
61a09532f2 | ||
|
|
3de4b2bf3c | ||
|
|
b44311b6ef | ||
|
|
b0d7875180 | ||
|
|
53c2f20dd9 | ||
|
|
37e370fe93 | ||
|
|
2dc5a72e7e | ||
|
|
c79ff5f918 | ||
|
|
1a659a0c37 | ||
|
|
0f6cf7f628 | ||
|
|
c79ad3ae21 | ||
|
|
61c9ef986a | ||
|
|
d6dbdb9b0d | ||
|
|
06da482fb4 | ||
|
|
2f75e7f712 | ||
|
|
7c21548ce3 | ||
|
|
9df2f91232 | ||
|
|
387189c429 | ||
|
|
75576c63be | ||
|
|
16aca639b7 | ||
|
|
6049424b7e | ||
|
|
060b5f61dc | ||
|
|
ec59c1579f | ||
|
|
1750e443f2 | ||
|
|
ba18929079 | ||
|
|
0500ca6a58 | ||
|
|
a1c15bcb0f | ||
|
|
4809de7317 | ||
|
|
05781e21dd | ||
|
|
85f638a2b8 | ||
|
|
08e5067561 | ||
|
|
a7d00ec051 | ||
|
|
b8fb56d970 | ||
|
|
96a739289e | ||
|
|
60d443f738 | ||
|
|
1dca300653 | ||
|
|
fca252d59e | ||
|
|
33178f9006 | ||
|
|
b2b8f679d0 | ||
|
|
de6ec294ef | ||
|
|
61e10f0116 | ||
|
|
6e96891ba0 | ||
|
|
47f1b47a73 | ||
|
|
5aab491bc9 | ||
|
|
5812e1a66b | ||
|
|
8950394e0a | ||
|
|
7bb49be4d1 | ||
|
|
c67650f04b | ||
|
|
f890e1dbe2 | ||
|
|
040cbf95cc | ||
|
|
5b3762a7f0 | ||
|
|
4d30c510ce | ||
|
|
6700813f86 | ||
|
|
eb44b3aaa4 | ||
|
|
7a98c7a392 | ||
|
|
0d9e60619b | ||
|
|
1134545b6f | ||
|
|
3e0c887511 | ||
|
|
adfbbc1005 | ||
|
|
adc98f04d0 | ||
|
|
8def3cdde2 | ||
|
|
616c9bd0f4 | ||
|
|
8688a06d67 | ||
|
|
f25953cc59 | ||
|
|
d9aa35161d | ||
|
|
6bcda970fd | ||
|
|
ea0e9c8f2e | ||
|
|
94ed0bf4e0 | ||
|
|
1940c8441e | ||
|
|
72d16aee15 | ||
|
|
e78a0c8e59 | ||
|
|
97a98006b0 | ||
|
|
0a684ab0c0 | ||
|
|
0d9210a502 | ||
|
|
1d874867ea | ||
|
|
2e2e626b40 | ||
|
|
af91f4b3e4 | ||
|
|
2396a61108 | ||
|
|
97a668152b | ||
|
|
58b2012aa2 | ||
|
|
b7c20d0cfa | ||
|
|
a2b1f9fc3b | ||
|
|
642076d26c | ||
|
|
5feb3950e5 | ||
|
|
4ec199b66a | ||
|
|
7ca017778f | ||
|
|
fbfe58133d | ||
|
|
9dd62d80ab | ||
|
|
f878367898 | ||
|
|
bd091079cb | ||
|
|
b23bd73f54 | ||
|
|
e2d7adeb64 | ||
|
|
15cb8e140d | ||
|
|
f007cceb42 | ||
|
|
0a5069e4e3 | ||
|
|
8ce53a616e | ||
|
|
ae10e855ab | ||
|
|
530ee36a0d | ||
|
|
d835ad572c | ||
|
|
47d0597ca2 | ||
|
|
818cf61e91 | ||
|
|
c01618fdc8 | ||
|
|
823eaf667d | ||
|
|
f1f1259692 | ||
|
|
df13b5aef5 | ||
|
|
4938d44a3b | ||
|
|
37bf988c2f | ||
|
|
9459fc6471 | ||
|
|
5245c80564 | ||
|
|
9bc266d923 | ||
|
|
5c9f6557d7 | ||
|
|
dcfebf93f4 | ||
|
|
752bd10647 | ||
|
|
2730b657c4 | ||
|
|
1dcbbd9cac | ||
|
|
ace9fda495 | ||
|
|
ef0aa7ca2f | ||
|
|
e6d1310b2a | ||
|
|
ac5f38a0f7 | ||
|
|
b6ff8a2f50 | ||
|
|
9243e0124e | ||
|
|
df362b2d6d | ||
|
|
7c2acd38b7 | ||
|
|
a287eb163f | ||
|
|
e94243893d | ||
|
|
29c0ec4d63 | ||
|
|
c7ce03bcbd | ||
|
|
c233d90aa8 | ||
|
|
d96aee0951 | ||
|
|
c71a583aa9 | ||
|
|
f12b80c6ef | ||
|
|
da64db78b9 | ||
|
|
425c4eafb0 | ||
|
|
02c01f442b | ||
|
|
fae543015c | ||
|
|
c9be3a8aa1 | ||
|
|
41ea2dd44a | ||
|
|
088c0be268 | ||
|
|
fcd2255d16 | ||
|
|
b5433b6f50 | ||
|
|
cc25f028b7 | ||
|
|
c4cd2bd544 | ||
|
|
5784507da4 | ||
|
|
bf578e1abd | ||
|
|
efed8a1e83 | ||
|
|
11d291511a | ||
|
|
877dae9c68 | ||
|
|
c4dd6d78fd | ||
|
|
ce2aecc4dc | ||
|
|
f38f3d11fb | ||
|
|
d4b4562917 | ||
|
|
7b3192523e | ||
|
|
4c6e2e4b30 | ||
|
|
8502958810 | ||
|
|
ce4bdcbda4 | ||
|
|
d5b1ec2684 | ||
|
|
867ff69733 | ||
|
|
109b736b86 | ||
|
|
69d4f5ef63 | ||
|
|
426d48bfa1 | ||
|
|
26c909ed74 | ||
|
|
fb1d8ccaf5 | ||
|
|
9354f22204 | ||
|
|
17fdd42100 | ||
|
|
472d330c21 | ||
|
|
3b6c96a101 | ||
|
|
4d4e04f452 | ||
|
|
67fe73b2b4 | ||
|
+1 |
ee8f36d0b3 | ||
|
+1 |
f3e9497e92 | ||
|
|
fe784ff22e | ||
|
|
b88abb5036 | ||
|
|
67f9046e4a | ||
|
|
f17be06fbe | ||
|
|
2cab53ddee | ||
|
|
ab0a20d151 | ||
|
|
4a394bfcda | ||
|
|
c95c663049 | ||
|
|
ab3c1aedf3 | ||
|
+1 |
fb5ec0dc9e | ||
|
|
971dac2caa | ||
|
|
efa2e424f6 | ||
|
|
02bf9c7907 | ||
|
|
f61163e6c7 | ||
|
|
626c90b2d5 | ||
|
+1 |
251f7e478e | ||
|
|
ce65385618 | ||
|
|
7d56fe2adc | ||
|
|
75bdad40b5 | ||
|
|
d08eebad16 | ||
|
|
3e90d015ba | ||
|
|
7cd1d57b74 | ||
|
|
b8168e33e0 | ||
|
|
d803b44dbe | ||
|
|
530852f959 | ||
|
|
a317bc5739 | ||
|
|
a9531edfa6 | ||
|
|
8c3393f373 | ||
|
|
9f8cbfd8eb | ||
|
|
ea1d65fe6d | ||
|
|
f44f3d6f79 | ||
|
|
cc706b05a5 | ||
|
|
85e296950c | ||
|
|
dc9f845ddc | ||
|
|
12f2c515a7 | ||
|
+1 |
6570c9800c | ||
|
|
8bfd683901 | ||
|
|
7dc2698632 | ||
|
|
59b964f37d | ||
|
|
6a9f24aa8c | ||
|
|
ba47bb5be1 | ||
|
|
df8a0900df | ||
|
|
2db39c7049 | ||
|
|
3935829f89 | ||
|
|
7746961277 | ||
|
|
5de1add806 | ||
|
|
915dffaa5f | ||
|
|
81e13a0591 | ||
|
|
f95e3f0edb | ||
|
|
5a65ba5f17 | ||
|
|
9d1c695be5 | ||
|
|
3c1bc1fc0d | ||
|
|
3a5e88e629 | ||
|
|
0becb7486b | ||
|
|
2dab187f75 | ||
|
|
015b0320de | ||
|
|
4238b011a7 | ||
|
|
eb33ff34dd | ||
|
|
2bd8957627 | ||
|
|
3034c8d389 | ||
|
|
ecf4aa5ce2 | ||
|
|
49e777cf08 | ||
|
|
b7950e798f | ||
|
|
de100ffb62 | ||
|
|
43cd340247 | ||
|
|
1d99f0f421 | ||
|
|
0885b51981 | ||
|
|
6036bf110a | ||
|
|
2fa63e0fff | ||
|
|
61141ed265 | ||
|
|
05eed72aec | ||
|
|
5810e884f1 | ||
|
|
615834ee58 | ||
|
|
5811ed6a05 | ||
|
|
1b30ae4ca4 | ||
|
|
4e04bcbce6 | ||
|
|
66b6c684ab | ||
|
|
c0302d9497 | ||
|
|
313fae3e89 | ||
|
|
7aab6e2684 | ||
|
|
9dd2e72828 | ||
|
|
d119beb1b9 | ||
|
|
12a8057bfe | ||
|
|
e281ac663a | ||
|
|
adce068118 | ||
|
|
b6770d7b54 | ||
|
|
3b39fd284a | ||
|
|
6472131298 | ||
|
|
37aa52821d | ||
|
|
96d2ceda4b | ||
|
|
fdf2cf66d3 | ||
|
|
9b2be4e9a5 | ||
|
|
3ad85e0de4 | ||
|
|
4f7fffb92f | ||
|
|
6e073440b1 | ||
|
|
f7aadae5e5 | ||
|
|
442c421e79 | ||
|
|
0bd6b85a1f | ||
|
|
3ca242d1b6 | ||
|
|
7e950521b3 | ||
|
|
0f0f28b537 | ||
|
|
520a20ba4e | ||
|
|
9182e86971 | ||
|
|
313d01f507 | ||
|
|
05d4f8bba3 | ||
|
|
0b54201a04 | ||
|
|
32e632dfeb | ||
|
|
7ffb98e248 | ||
|
|
cdaa40d2a8 | ||
|
|
ca3618bc69 | ||
|
|
b2f7d2560a | ||
|
|
1ff9429655 | ||
|
|
af453e5647 | ||
|
|
32aef44388 | ||
|
|
7a74a9662b | ||
|
|
b6754f536e | ||
|
|
793cf79c89 | ||
|
|
50ac1c7bab | ||
|
|
f04d3f640e | ||
|
|
0a9396a25e | ||
|
|
038ec293b1 | ||
|
|
894ebb27f5 | ||
|
|
c9a788eedc | ||
|
|
0762f2afeb | ||
|
|
31be872f55 | ||
|
|
94c0ef3001 | ||
|
|
af1f036a70 | ||
|
|
95aab66e95 | ||
|
|
dcf4072da9 | ||
|
|
382bbd5144 | ||
|
|
b50ef9c6ed | ||
|
|
9e289c553c | ||
|
|
c4f5cd60da | ||
|
|
0b0ef8d7eb | ||
|
|
21472f32ea | ||
|
|
fec64fea75 | ||
|
|
8b8af2caf7 | ||
|
|
7738ef35b8 | ||
|
|
9a21f0d1a3 | ||
|
|
8ac8375270 | ||
|
|
7dc447dda7 |
@@ -18,6 +18,8 @@ steps:
|
||||
TERM: "xterm-256color"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1 # Transient Docker/BuildKit failure
|
||||
limit: 1
|
||||
- exit_status: -1 # Agent was lost
|
||||
limit: 1
|
||||
- exit_status: -10 # Agent was lost
|
||||
@@ -46,6 +48,8 @@ steps:
|
||||
VLLM_BRANCH: "$BUILDKITE_COMMIT"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1 # Transient Docker/BuildKit failure
|
||||
limit: 1
|
||||
- exit_status: -1 # Agent was lost
|
||||
limit: 1
|
||||
- exit_status: -10 # Agent was lost
|
||||
@@ -72,6 +76,8 @@ steps:
|
||||
VLLM_BRANCH: "$BUILDKITE_COMMIT"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1 # Transient Docker/BuildKit failure
|
||||
limit: 1
|
||||
- exit_status: -1 # Agent was lost
|
||||
limit: 1
|
||||
- exit_status: -10 # Agent was lost
|
||||
|
||||
@@ -18,6 +18,8 @@ steps:
|
||||
- tests/kernels/quantization/test_cpu_fp8_scaled_mm.py
|
||||
- tests/kernels/mamba/cpu/test_cpu_gdn_ops.py
|
||||
- tests/kernels/mamba/test_cpu_short_conv.py
|
||||
- tests/kernels/mamba/test_causal_conv1d.py
|
||||
- tests/kernels/mamba/test_mamba_ssm.py
|
||||
commands:
|
||||
- |
|
||||
bash .buildkite/scripts/hardware_ci/run-cpu-test.sh 30m "
|
||||
@@ -28,7 +30,9 @@ steps:
|
||||
pytest -x -v -s tests/kernels/test_onednn.py
|
||||
pytest -x -v -s tests/kernels/test_awq_int4_to_int8.py
|
||||
pytest -x -v -s tests/kernels/quantization/test_cpu_fp8_scaled_mm.py
|
||||
pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py"
|
||||
pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_causal_conv1d.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_mamba_ssm.py"
|
||||
|
||||
# Note: SDE can't be downloaded from CI host because of AWS WAF
|
||||
# - label: CPU-Compatibility Tests
|
||||
|
||||
@@ -99,6 +99,7 @@ steps:
|
||||
pytest -v -s v1/worker --ignore=v1/worker/test_gpu_model_runner.py --ignore=v1/worker/test_worker_memory_snapshot.py &&
|
||||
pytest -v -s v1/structured_output &&
|
||||
pytest -v -s v1/test_serial_utils.py &&
|
||||
pytest -v -s v1/e2e/general/test_correctness_sliding_window.py --deselect="tests/v1/e2e/general/test_correctness_sliding_window.py::test_sliding_window_retrieval[True-1-5-google/gemma-3-1b-it]" &&
|
||||
pytest -v -s v1/spec_decode --ignore=v1/spec_decode/test_max_len.py --ignore=v1/spec_decode/test_speculators_eagle3.py --ignore=v1/spec_decode/test_acceptance_length.py --ignore=v1/spec_decode/test_speculators_correctness.py &&
|
||||
pytest -v -s v1/kv_connector/unit --ignore=v1/kv_connector/unit/test_multi_connector.py --ignore=v1/kv_connector/unit/test_example_connector.py --ignore=v1/kv_connector/unit/test_lmcache_integration.py --ignore=v1/kv_connector/unit/test_hf3fs_client.py --ignore=v1/kv_connector/unit/test_hf3fs_connector.py --ignore=v1/kv_connector/unit/test_hf3fs_metadata_server.py --ignore=v1/kv_connector/unit/test_offloading_connector.py'
|
||||
- label: "XPU server test"
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# For hf script, without -t option (tensor parallel size).
|
||||
# bash .buildkite/lm-eval-harness/run-lm-eval-mmlupro-vllm-baseline.sh -m meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8 -l 250 -t 8 -f 5
|
||||
model_name: "meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8"
|
||||
rocm_safetensors_load_strategy: lazy
|
||||
required_gpu_arch:
|
||||
- gfx942
|
||||
- gfx950
|
||||
|
||||
@@ -72,6 +72,11 @@ def launch_lm_eval(eval_config, tp_size):
|
||||
if moe_backend is not None:
|
||||
model_args += f"moe_backend={moe_backend},"
|
||||
|
||||
if current_platform.is_rocm():
|
||||
rocm_load_strategy = eval_config.get("rocm_safetensors_load_strategy")
|
||||
if rocm_load_strategy is not None:
|
||||
model_args += f"safetensors_load_strategy={rocm_load_strategy},"
|
||||
|
||||
env_vars = eval_config.get("env_vars", None)
|
||||
with scoped_env_vars(env_vars):
|
||||
results = lm_eval.simple_evaluate(
|
||||
|
||||
+472
-399
File diff suppressed because it is too large
Load Diff
@@ -3,7 +3,8 @@
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# Append a build artifact line to the Buildkite annotation.
|
||||
# Usage: annotate-build-artifact.sh <label> <value>
|
||||
# Usage: annotate-build-artifact.sh <label> <value> <context>
|
||||
set -e
|
||||
echo "- **${1}**: \`${2}\`" | \
|
||||
buildkite-agent annotate --append --style 'info' --context 'release-artifacts'
|
||||
buildkite-agent annotate --append --style 'info' \
|
||||
--context "${3:?context is required}"
|
||||
|
||||
Executable
+32
@@ -0,0 +1,32 @@
|
||||
#!/usr/bin/env bash
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
#
|
||||
# Build the macOS arm64 CPU wheel natively on a macOS agent (the `macmini`
|
||||
# queue) into artifacts/dist/ for upload-nightly-wheels.sh.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# The Rust frontend build needs protoc.
|
||||
if ! command -v protoc >/dev/null 2>&1; then
|
||||
brew install protobuf
|
||||
fi
|
||||
|
||||
# upload-nightly-wheels.sh expects exactly one wheel.
|
||||
rm -rf artifacts/dist
|
||||
mkdir -p artifacts/dist
|
||||
|
||||
export VLLM_TARGET_DEVICE=cpu
|
||||
export VLLM_REQUIRE_RUST_FRONTEND=1
|
||||
export MACOSX_DEPLOYMENT_TARGET=11.0
|
||||
# uv's CPython is universal2; force an arm64-only build and tag so the wheel
|
||||
# isn't mislabelled universal2 and installed on Intel Macs where import fails.
|
||||
export ARCHFLAGS="-arch arm64"
|
||||
export _PYTHON_HOST_PLATFORM="macosx-11.0-arm64"
|
||||
export CMAKE_BUILD_PARALLEL_LEVEL="${CMAKE_BUILD_PARALLEL_LEVEL:-4}"
|
||||
|
||||
uv venv --python 3.12
|
||||
uv pip install -r requirements/build/cpu.txt --index-strategy unsafe-best-match
|
||||
uv build --wheel --no-build-isolation -o artifacts/dist
|
||||
|
||||
ls -l artifacts/dist/*.whl
|
||||
@@ -17,7 +17,7 @@ DEFAULT_REPO_SLUG="vllm-project/vllm"
|
||||
DEFAULT_CI_HCL_SOURCE="docker/ci-rocm.hcl"
|
||||
DEFAULT_CI_BASE_CONTENT_FILES="requirements/common.txt requirements/rocm.txt requirements/test/rocm.txt docker/Dockerfile.rocm_base docker/ci-rocm.hcl docker/docker-bake-rocm.hcl tools/install_torchcodec_rocm.sh tools/install_protoc.sh rust-toolchain.toml tests/vllm_test_utils .buildkite/scripts/ci-bake-rocm.sh .buildkite/scripts/rocm/build-ci-base.sh"
|
||||
DEFAULT_CI_BASE_DOCKERFILE="docker/Dockerfile.rocm"
|
||||
DEFAULT_CI_BASE_DOCKERFILE_STAGES="base rust_toolchain_input_0 rust_toolchain_input_1 rust-toolchain-input rust-toolchain build_rixl build_rocshmem build_deepep mori_base ci_base"
|
||||
DEFAULT_CI_BASE_DOCKERFILE_STAGES="base rust_toolchain_input_0 rust_toolchain_input_1 rust-toolchain-input rust-toolchain build_nixl build_rocshmem build_deepep mori_base ci_base"
|
||||
DEFAULT_CI_BASE_METADATA_VERSION="1"
|
||||
IMAGE_EXISTED_BEFORE_BUILD=0
|
||||
|
||||
@@ -285,7 +285,7 @@ get_content_arg_names() {
|
||||
fi | awk 'NF && !seen[$0]++'
|
||||
}
|
||||
|
||||
compute_ci_base_content_hash() {
|
||||
compute_ci_base_content_hash_once() {
|
||||
local -a content_paths=()
|
||||
local -a content_args=()
|
||||
local dockerfile="${CI_BASE_DOCKERFILE:-}"
|
||||
@@ -301,7 +301,8 @@ compute_ci_base_content_hash() {
|
||||
if [[ -n "${dockerfile}" ]]; then
|
||||
printf 'dockerfile:%s\n' "${dockerfile}"
|
||||
printf 'resolved-build-args:\n'
|
||||
hash_dockerfile_arg_values "${dockerfile}" "${content_args[@]}"
|
||||
hash_dockerfile_arg_values "${dockerfile}" "${content_args[@]}" \
|
||||
|| return 1
|
||||
if [[ -n "${stages}" ]]; then
|
||||
printf 'dockerfile-stages:%s\n' "${stages}"
|
||||
if [[ -f "${dockerfile}" ]]; then
|
||||
@@ -314,6 +315,53 @@ compute_ci_base_content_hash() {
|
||||
} | sha256sum | cut -d' ' -f1
|
||||
}
|
||||
|
||||
compute_ci_base_content_hash() {
|
||||
local attempts="${CI_BASE_HASH_ATTEMPTS:-3}"
|
||||
local delay_secs="${CI_BASE_HASH_RETRY_DELAY:-5}"
|
||||
local attempt=0
|
||||
local hash=""
|
||||
local failed=0
|
||||
local -a hashes=()
|
||||
|
||||
if [[ ! "${attempts}" =~ ^[1-9][0-9]*$ ]]; then
|
||||
echo "Invalid CI_BASE_HASH_ATTEMPTS: ${attempts}" >&2
|
||||
return 1
|
||||
fi
|
||||
if [[ ! "${delay_secs}" =~ ^[0-9]+$ ]]; then
|
||||
echo "Invalid CI_BASE_HASH_RETRY_DELAY: ${delay_secs}" >&2
|
||||
return 1
|
||||
fi
|
||||
|
||||
for ((attempt = 1; attempt <= attempts; attempt++)); do
|
||||
if ! hash=$(compute_ci_base_content_hash_once); then
|
||||
echo "ci_base content hash calculation ${attempt}/${attempts} failed" >&2
|
||||
failed=1
|
||||
else
|
||||
hashes+=("${hash}")
|
||||
echo "ci_base content hash calculation ${attempt}/${attempts}: ${hash}" >&2
|
||||
fi
|
||||
|
||||
if ((attempt < attempts)); then
|
||||
sleep "${delay_secs}"
|
||||
fi
|
||||
done
|
||||
|
||||
if ((failed)) || ((${#hashes[@]} != attempts)); then
|
||||
echo "Could not calculate a reliable ci_base content hash" >&2
|
||||
return 1
|
||||
fi
|
||||
|
||||
for hash in "${hashes[@]:1}"; do
|
||||
if [[ "${hash}" != "${hashes[0]}" ]]; then
|
||||
echo "ci_base content hash changed between calculations" >&2
|
||||
printf ' observed: %s\n' "${hashes[@]}" >&2
|
||||
return 1
|
||||
fi
|
||||
done
|
||||
|
||||
printf '%s\n' "${hashes[0]}"
|
||||
}
|
||||
|
||||
extract_dockerfile_arg_default() {
|
||||
local dockerfile="$1"
|
||||
local arg_name="$2"
|
||||
@@ -366,7 +414,11 @@ hash_dockerfile_arg_values() {
|
||||
printf 'arg:%s=%s\n' "${arg_name}" "${arg_value:-<empty>}"
|
||||
if [[ "${arg_name}" == "BASE_IMAGE" && -n "${arg_value}" ]]; then
|
||||
digest=$(resolve_image_digest "${arg_value}")
|
||||
printf 'arg:%s.digest=%s\n' "${arg_name}" "${digest:-unknown}"
|
||||
if [[ -z "${digest}" ]]; then
|
||||
echo "Failed to resolve digest for BASE_IMAGE=${arg_value}" >&2
|
||||
return 1
|
||||
fi
|
||||
printf 'arg:%s.digest=%s\n' "${arg_name}" "${digest}"
|
||||
fi
|
||||
done
|
||||
}
|
||||
@@ -1107,8 +1159,8 @@ ci_base_metadata_pairs() {
|
||||
metadata_pair "vllm.rocm.nic_backend" "$(resolve_dockerfile_arg_value "${dockerfile}" "NIC_BACKEND")"
|
||||
metadata_pair "vllm.rocm.ainic_version" "$(resolve_dockerfile_arg_value "${dockerfile}" "AINIC_VERSION")"
|
||||
metadata_pair "vllm.rocm.ubuntu_codename" "$(resolve_dockerfile_arg_value "${dockerfile}" "UBUNTU_CODENAME")"
|
||||
metadata_pair "vllm.rocm.rixl_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "RIXL_REPO")"
|
||||
metadata_pair "vllm.rocm.rixl_commit" "${RIXL_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "RIXL_BRANCH")}"
|
||||
metadata_pair "vllm.rocm.nixl_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "NIXL_REPO")"
|
||||
metadata_pair "vllm.rocm.nixl_commit" "${NIXL_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "NIXL_BRANCH")}"
|
||||
metadata_pair "vllm.rocm.ucx_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "UCX_REPO")"
|
||||
metadata_pair "vllm.rocm.ucx_commit" "${UCX_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "UCX_BRANCH")}"
|
||||
metadata_pair "vllm.rocm.rocshmem_repo" "$(resolve_dockerfile_arg_value "${dockerfile}" "ROCSHMEM_REPO")"
|
||||
@@ -1117,7 +1169,7 @@ ci_base_metadata_pairs() {
|
||||
metadata_pair "vllm.rocm.deepep_commit" "${DEEPEP_BRANCH:-$(resolve_dockerfile_arg_value "${dockerfile}" "DEEPEP_BRANCH")}"
|
||||
metadata_pair "vllm.rocm.deepep_nic" "$(resolve_dockerfile_arg_value "${dockerfile}" "DEEPEP_NIC")"
|
||||
metadata_pair "vllm.rocm.deepep_rocm_arch" "$(resolve_dockerfile_arg_value "${dockerfile}" "DEEPEP_ROCM_ARCH")"
|
||||
metadata_pair "vllm.rocm.rixl_cache_key" "${RIXL_CACHE_KEY:-}"
|
||||
metadata_pair "vllm.rocm.nixl_cache_key" "${NIXL_CACHE_KEY:-}"
|
||||
metadata_pair "vllm.rocm.rocshmem_cache_key" "${ROCSHMEM_CACHE_KEY:-}"
|
||||
metadata_pair "vllm.rocm.deepep_cache_key" "${DEEPEP_CACHE_KEY:-}"
|
||||
|
||||
@@ -1634,7 +1686,7 @@ extract_dependency_pins() {
|
||||
return 0
|
||||
fi
|
||||
|
||||
for var in RIXL_BRANCH UCX_BRANCH ROCSHMEM_BRANCH DEEPEP_BRANCH; do
|
||||
for var in NIXL_BRANCH UCX_BRANCH ROCSHMEM_BRANCH DEEPEP_BRANCH; do
|
||||
if [[ -n "${!var:-}" ]]; then
|
||||
echo "Using provided ${var}: ${!var}"
|
||||
continue
|
||||
@@ -1654,30 +1706,30 @@ extract_dependency_pins() {
|
||||
compute_dependency_cache_keys() {
|
||||
local bake_dir=""
|
||||
local dockerfile_rocm=""
|
||||
local rixl_branch=""
|
||||
local nixl_branch=""
|
||||
local ucx_branch=""
|
||||
local rocshmem_branch=""
|
||||
local deepep_branch=""
|
||||
local rixl_material=""
|
||||
local nixl_material=""
|
||||
local rocshmem_material=""
|
||||
local deepep_material=""
|
||||
|
||||
bake_dir=$(dirname "${VLLM_BAKE_FILE}")
|
||||
dockerfile_rocm="${bake_dir}/Dockerfile.rocm"
|
||||
rixl_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "RIXL_BRANCH")
|
||||
nixl_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "NIXL_BRANCH")
|
||||
ucx_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "UCX_BRANCH")
|
||||
rocshmem_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "ROCSHMEM_BRANCH")
|
||||
deepep_branch=$(resolve_dockerfile_arg_value "${dockerfile_rocm}" "DEEPEP_BRANCH")
|
||||
|
||||
if [[ -n "${rixl_branch}" && -n "${ucx_branch}" ]]; then
|
||||
rixl_material=$(compose_stage_cache_material "${dockerfile_rocm}" "base build_rixl")
|
||||
RIXL_CACHE_KEY=$(
|
||||
if [[ -n "${nixl_branch}" && -n "${ucx_branch}" ]]; then
|
||||
nixl_material=$(compose_stage_cache_material "${dockerfile_rocm}" "base build_nixl")
|
||||
NIXL_CACHE_KEY=$(
|
||||
compose_dependency_cache_key \
|
||||
"${rixl_branch}-ucx-${ucx_branch}" \
|
||||
"${rixl_material}"
|
||||
"${nixl_branch}-ucx-${ucx_branch}" \
|
||||
"${nixl_material}"
|
||||
)
|
||||
export RIXL_CACHE_KEY
|
||||
echo "RIXL dependency cache key: ${RIXL_CACHE_KEY}"
|
||||
export NIXL_CACHE_KEY
|
||||
echo "NIXL dependency cache key: ${NIXL_CACHE_KEY}"
|
||||
fi
|
||||
|
||||
if [[ -n "${rocshmem_branch}" ]]; then
|
||||
@@ -1728,11 +1780,11 @@ dependency_cache_ref_for_target() {
|
||||
local cache_repo="${DOCKERHUB_CACHE_REPO:-rocm/vllm-ci-cache}"
|
||||
|
||||
case "${target}" in
|
||||
rixl-rocm-ci)
|
||||
if [[ -n "${RIXL_CACHE_KEY:-}" ]]; then
|
||||
printf '%s\n' "${cache_repo}:rixl-rocm-${RIXL_CACHE_KEY}"
|
||||
elif [[ -n "${RIXL_BRANCH:-}" ]]; then
|
||||
printf '%s\n' "${cache_repo}:rixl-rocm-${RIXL_BRANCH}-ucx-${UCX_BRANCH:-}"
|
||||
nixl-rocm-ci)
|
||||
if [[ -n "${NIXL_CACHE_KEY:-}" ]]; then
|
||||
printf '%s\n' "${cache_repo}:nixl-rocm-${NIXL_CACHE_KEY}"
|
||||
elif [[ -n "${NIXL_BRANCH:-}" ]]; then
|
||||
printf '%s\n' "${cache_repo}:nixl-rocm-${NIXL_BRANCH}-ucx-${UCX_BRANCH:-}"
|
||||
fi
|
||||
;;
|
||||
rocshmem-rocm-ci)
|
||||
@@ -1763,7 +1815,7 @@ add_dependency_cache_target() {
|
||||
|
||||
resolve_ci_base_dependency_targets() {
|
||||
local mode="${ROCM_DEP_CACHE_EXPORT_MODE:-missing}"
|
||||
local rixl_ref=""
|
||||
local nixl_ref=""
|
||||
local rocshmem_ref=""
|
||||
local deepep_ref=""
|
||||
|
||||
@@ -1772,7 +1824,7 @@ resolve_ci_base_dependency_targets() {
|
||||
case "${mode}" in
|
||||
always)
|
||||
echo "ROCM_DEP_CACHE_EXPORT_MODE=always; exporting all dependency caches serially"
|
||||
for target in rixl-rocm-ci rocshmem-rocm-ci deepep-rocm-ci; do
|
||||
for target in nixl-rocm-ci rocshmem-rocm-ci deepep-rocm-ci; do
|
||||
if [[ -n "$(dependency_cache_ref_for_target "${target}")" ]]; then
|
||||
add_dependency_cache_target "${target}"
|
||||
fi
|
||||
@@ -1792,13 +1844,13 @@ resolve_ci_base_dependency_targets() {
|
||||
;;
|
||||
esac
|
||||
|
||||
if [[ "${mode}" != "always" && -n "${RIXL_CACHE_KEY:-}" ]]; then
|
||||
rixl_ref=$(dependency_cache_ref_for_target "rixl-rocm-ci")
|
||||
if dependency_cache_ref_exists "${rixl_ref}"; then
|
||||
echo "RIXL dependency cache exists: ${rixl_ref}"
|
||||
if [[ "${mode}" != "always" && -n "${NIXL_CACHE_KEY:-}" ]]; then
|
||||
nixl_ref=$(dependency_cache_ref_for_target "nixl-rocm-ci")
|
||||
if dependency_cache_ref_exists "${nixl_ref}"; then
|
||||
echo "NIXL dependency cache exists: ${nixl_ref}"
|
||||
else
|
||||
echo "RIXL dependency cache missing; will seed: ${rixl_ref}"
|
||||
add_dependency_cache_target "rixl-rocm-ci"
|
||||
echo "NIXL dependency cache missing; will seed: ${nixl_ref}"
|
||||
add_dependency_cache_target "nixl-rocm-ci"
|
||||
fi
|
||||
fi
|
||||
|
||||
@@ -1898,8 +1950,8 @@ confirm_remote_image_push() {
|
||||
fi
|
||||
|
||||
if [[ -z "${remote_revision}" \
|
||||
&& ${IMAGE_EXISTED_BEFORE_BUILD} -eq 0 \
|
||||
&& image_tag_is_commit_scoped ]]; then
|
||||
&& ${IMAGE_EXISTED_BEFORE_BUILD} -eq 0 ]] \
|
||||
&& image_tag_is_commit_scoped; then
|
||||
echo "Remote image exists under a commit-scoped tag; accepting push despite missing revision label."
|
||||
return 0
|
||||
fi
|
||||
@@ -2029,36 +2081,57 @@ upload_wheel_artifacts_if_present() {
|
||||
local wheel_dir="./wheel-export"
|
||||
local artifact_dir="artifacts/vllm-rocm-install"
|
||||
local archive_name="vllm-rocm-install.tar.gz"
|
||||
local metadata_dir="${wheel_dir}/.vllm-ci-artifact"
|
||||
local native_base_image=""
|
||||
local whl=""
|
||||
local whl_name=""
|
||||
local -a wheels=()
|
||||
|
||||
if ! should_upload_wheel_artifacts; then
|
||||
return 0
|
||||
fi
|
||||
|
||||
if [[ ! -d "${wheel_dir}" ]] || ! ls "${wheel_dir}"/*.whl >/dev/null 2>&1; then
|
||||
echo "No ROCm wheel artifacts found in ${wheel_dir}"
|
||||
return 0
|
||||
if [[ -d "${wheel_dir}" ]]; then
|
||||
mapfile -t wheels < <(find "${wheel_dir}" -maxdepth 1 -type f -name '*.whl' -print)
|
||||
fi
|
||||
if [[ ${#wheels[@]} -ne 1 ]]; then
|
||||
echo "Expected exactly one ROCm wheel in ${wheel_dir}; found ${#wheels[@]}" >&2
|
||||
return 1
|
||||
fi
|
||||
whl="${wheels[0]}"
|
||||
whl_name=$(basename "${whl}")
|
||||
native_base_image="${CI_BASE_IMAGE_TAG_COMMIT_REF:-${CI_BASE_IMAGE:-}}"
|
||||
if [[ -z "${native_base_image}" ]]; then
|
||||
echo "Native ROCm artifact requires a ci_base image reference" >&2
|
||||
return 1
|
||||
fi
|
||||
|
||||
echo "--- :package: Uploading ROCm vLLM install artifact"
|
||||
mkdir -p "${artifact_dir}"
|
||||
rm -rf "${artifact_dir}" "${metadata_dir}"
|
||||
mkdir -p "${artifact_dir}" "${metadata_dir}"
|
||||
|
||||
printf '%s\n' "${BUILDKITE_COMMIT:-local}" > "${metadata_dir}/commit.txt"
|
||||
printf '%s\n' "${native_base_image}" > "${metadata_dir}/native-base-image.txt"
|
||||
printf '%s\n' "${CI_BASE_IMAGE:-}" > "${metadata_dir}/ci-base-image.txt"
|
||||
printf '%s\n' "${IMAGE_TAG:-}" > "${metadata_dir}/fallback-image.txt"
|
||||
printf '%s\n' "${whl_name}" > "${metadata_dir}/wheel-filename.txt"
|
||||
|
||||
tar -C "${wheel_dir}" -czf "${artifact_dir}/${archive_name}" .
|
||||
(
|
||||
cd "${artifact_dir}"
|
||||
sha256sum "${archive_name}" > "${archive_name}.sha256"
|
||||
)
|
||||
echo "Created ${archive_name}: $(du -sh "${artifact_dir}/${archive_name}" | cut -f1)"
|
||||
printf '%s\n' "${CI_BASE_IMAGE:-}" > "${artifact_dir}/ci-base-image.txt"
|
||||
printf '%s\n' "${IMAGE_TAG:-}" > "${artifact_dir}/fallback-image.txt"
|
||||
|
||||
for whl in "${wheel_dir}"/*.whl; do
|
||||
[[ -f "${whl}" ]] || continue
|
||||
whl_name=$(basename "${whl}")
|
||||
cp "${whl}" "${artifact_dir}/${whl_name}"
|
||||
echo "Copied ${whl_name}: $(du -sh "${artifact_dir}/${whl_name}" | cut -f1)"
|
||||
done
|
||||
cp "${metadata_dir}"/*.txt "${artifact_dir}/"
|
||||
cp "${whl}" "${artifact_dir}/${whl_name}"
|
||||
echo "Copied ${whl_name}: $(du -sh "${artifact_dir}/${whl_name}" | cut -f1)"
|
||||
|
||||
if command -v buildkite-agent >/dev/null 2>&1; then
|
||||
buildkite-agent artifact upload "${artifact_dir}/*"
|
||||
buildkite-agent artifact upload "${artifact_dir}/*" || return 1
|
||||
echo "ROCm vLLM install artifacts uploaded to ${artifact_dir}/"
|
||||
elif [[ "${BUILDKITE:-false}" == "true" ]]; then
|
||||
echo "buildkite-agent not found; cannot upload required ROCm artifacts" >&2
|
||||
return 1
|
||||
else
|
||||
echo "Not in Buildkite, skipping artifact upload"
|
||||
fi
|
||||
@@ -2090,6 +2163,11 @@ main() {
|
||||
echo "BAKE_PRINT_ONLY=1 set; skipping build"
|
||||
return 0
|
||||
fi
|
||||
if should_upload_wheel_artifacts; then
|
||||
# wheel-export is an output directory, not a BuildKit cache. Starting
|
||||
# clean prevents a failed/retried export from packaging a stale wheel.
|
||||
rm -rf ./wheel-export
|
||||
fi
|
||||
seed_dependency_caches_if_needed
|
||||
run_bake
|
||||
upload_wheel_artifacts_if_present
|
||||
|
||||
@@ -45,8 +45,10 @@ $PYTHON .buildkite/scripts/generate-nightly-index.py --version "$SUBPATH" --curr
|
||||
echo "Uploading indices to $S3_COMMIT_PREFIX"
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "$S3_COMMIT_PREFIX"
|
||||
|
||||
# copy to /nightly/ only if it is on the main branch and not a PR
|
||||
if [[ "$BUILDKITE_BRANCH" == "main" && "$BUILDKITE_PULL_REQUEST" == "false" ]]; then
|
||||
# copy to /nightly/ only when enabled for a main branch build that is not a PR
|
||||
if [[ "${UPDATE_NIGHTLY_INDEX:-1}" == "1" && \
|
||||
"$BUILDKITE_BRANCH" == "main" && \
|
||||
"$BUILDKITE_PULL_REQUEST" == "false" ]]; then
|
||||
echo "Uploading indices to overwrite /nightly/"
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "s3://$BUCKET/nightly/"
|
||||
fi
|
||||
@@ -67,7 +69,7 @@ pure_version="${version%%+*}"
|
||||
echo "Pure version (without variant): $pure_version"
|
||||
|
||||
# re-generate and copy to /<pure_version>/ only if it does not have "dev" in the version
|
||||
if [[ "$version" != *"dev"* ]]; then
|
||||
if [[ "${UPDATE_VERSION_INDEX:-1}" == "1" && "$version" != *"dev"* ]]; then
|
||||
echo "Re-generating indices for /$pure_version/"
|
||||
rm -rf "${INDICES_OUTPUT_DIR:?}"
|
||||
mkdir -p "$INDICES_OUTPUT_DIR"
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#!/bin/bash
|
||||
|
||||
# This script runs tests inside the corresponding ROCm docker container.
|
||||
# It handles both single-node and multi-node test configurations.
|
||||
# This script runs ROCm tests either directly in a native CI pod or inside the
|
||||
# corresponding Docker container. Multi-node tests continue to use Docker.
|
||||
#
|
||||
# Multi-node detection: Instead of matching on fragile group names, we detect
|
||||
# multi-node jobs structurally by looking for the bracket command syntax
|
||||
@@ -34,10 +34,27 @@ set -o pipefail
|
||||
: "${CLICOLOR_FORCE:=1}"
|
||||
: "${PY_COLORS:=1}"
|
||||
: "${ROCM_DOCKER_TTY:=1}"
|
||||
: "${PYTHONFAULTHANDLER:=1}"
|
||||
: "${PYTEST_TIMEOUT:=2400}"
|
||||
if [[ " ${PYTEST_ADDOPTS:-} " != *" --color"* ]]; then
|
||||
PYTEST_ADDOPTS="${PYTEST_ADDOPTS:+${PYTEST_ADDOPTS} }--color=yes"
|
||||
fi
|
||||
export BUILDKIT_PROGRESS TERM FORCE_COLOR CLICOLOR_FORCE PY_COLORS PYTEST_ADDOPTS ROCM_DOCKER_TTY
|
||||
if [[ " ${PYTEST_ADDOPTS:-} " != *" --durations="* ]]; then
|
||||
PYTEST_ADDOPTS="${PYTEST_ADDOPTS:+${PYTEST_ADDOPTS} }--durations=25"
|
||||
fi
|
||||
if [[ " ${PYTEST_ADDOPTS:-} " != *" --durations-min="* ]]; then
|
||||
PYTEST_ADDOPTS="${PYTEST_ADDOPTS:+${PYTEST_ADDOPTS} }--durations-min=1.0"
|
||||
fi
|
||||
# Dump stacks after 25 minutes, then stop an individual test after 40 minutes.
|
||||
if [[ " ${PYTEST_ADDOPTS:-} " != *" faulthandler_timeout="* ]]; then
|
||||
PYTEST_ADDOPTS="${PYTEST_ADDOPTS:+${PYTEST_ADDOPTS} }-o faulthandler_timeout=1500"
|
||||
fi
|
||||
if [[ " ${PYTEST_ADDOPTS:-} " != *" --timeout-method="* &&
|
||||
" ${PYTEST_ADDOPTS:-} " != *" --timeout-method "* ]]; then
|
||||
PYTEST_ADDOPTS="${PYTEST_ADDOPTS:+${PYTEST_ADDOPTS} }--timeout-method=thread"
|
||||
fi
|
||||
export BUILDKIT_PROGRESS TERM FORCE_COLOR CLICOLOR_FORCE PY_COLORS PYTEST_ADDOPTS PYTEST_TIMEOUT ROCM_DOCKER_TTY
|
||||
export PYTHONFAULTHANDLER
|
||||
|
||||
# Export Python path for commands that run directly on the host. Containerized
|
||||
# tests set this to /vllm-workspace below so spawned Python processes do not
|
||||
@@ -53,6 +70,28 @@ report_docker_usage() {
|
||||
docker system df || true
|
||||
}
|
||||
|
||||
clear_ci_orchestration_env() {
|
||||
unset -v \
|
||||
VLLM_TEST_GROUP_NAME \
|
||||
VLLM_CI_REQUIRE_PERSISTENT_HF_CACHE \
|
||||
VLLM_CI_ARTIFACT_STEP \
|
||||
VLLM_TEST_CACHE \
|
||||
VLLM_CI_EXECUTION_MODE \
|
||||
VLLM_CI_WORKSPACE \
|
||||
VLLM_CI_REQUIRE_WORKSPACE_MOUNT \
|
||||
VLLM_TEST_COMMANDS \
|
||||
VLLM_CI_BRANCH \
|
||||
VLLM_CI_BASE_IMAGE \
|
||||
VLLM_CI_FALLBACK_IMAGE \
|
||||
VLLM_CI_DOCKER_DISABLED \
|
||||
VLLM_CI_ARTIFACT_GLOB \
|
||||
VLLM_CI_ARTIFACT_CHECKSUM_GLOB \
|
||||
VLLM_CI_EXPECTED_GPU_COUNT \
|
||||
VLLM_CI_USE_ARTIFACTS \
|
||||
VLLM_CI_RESULTS_ROOT \
|
||||
VLLM_ALLOW_DEPRECATED_BEAM_SEARCH
|
||||
}
|
||||
|
||||
cleanup_network() {
|
||||
local max_nodes=${NUM_NODES:-2}
|
||||
for node in $(seq 0 $((max_nodes - 1))); do
|
||||
@@ -145,7 +184,11 @@ prepare_artifact_image() {
|
||||
fi
|
||||
|
||||
cp "${wheel_dir}"/*.whl "${context_dir}/wheels/" || return 1
|
||||
tar -C "${wheel_dir}" --exclude='*.whl' -cf - . \
|
||||
tar -C "${wheel_dir}" \
|
||||
--exclude='*.whl' \
|
||||
--exclude='.vllm-ci-artifact' \
|
||||
--exclude='./.vllm-ci-artifact' \
|
||||
-cf - . \
|
||||
| tar -C "${workspace_dir}" -xf - || return 1
|
||||
cat > "${context_dir}/Dockerfile" <<'EOF'
|
||||
ARG BASE_IMAGE
|
||||
@@ -168,6 +211,259 @@ EOF
|
||||
return 0
|
||||
}
|
||||
|
||||
is_native_runtime() {
|
||||
[[ "${AMD_CI_RUNTIME:-}" == "native" || "${NATIVE_CI:-}" == "true" ]]
|
||||
}
|
||||
|
||||
validate_native_workspace() {
|
||||
local workspace_dir="${VLLM_CI_WORKSPACE:-/vllm-workspace}"
|
||||
local workspace_real=""
|
||||
local checkout_real=""
|
||||
local workspace_mount=""
|
||||
|
||||
mkdir -p "${workspace_dir}" || return 1
|
||||
workspace_real=$(readlink -m "${workspace_dir}") || return 1
|
||||
if [[ -n "${BUILDKITE_BUILD_CHECKOUT_PATH:-}" ]]; then
|
||||
checkout_real=$(readlink -m "${BUILDKITE_BUILD_CHECKOUT_PATH}") || return 1
|
||||
if [[ "${checkout_real}" == "${workspace_real}" \
|
||||
|| "${checkout_real}" == "${workspace_real}/"* \
|
||||
|| "${workspace_real}" == "${checkout_real}/"* ]]; then
|
||||
echo "Refusing to replace ${workspace_real}; it overlaps the Buildkite checkout ${checkout_real}" >&2
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
if [[ "${VLLM_CI_REQUIRE_WORKSPACE_MOUNT:-1}" == "1" ]]; then
|
||||
if ! command -v findmnt >/dev/null 2>&1; then
|
||||
echo "findmnt is required to verify the native workspace mount" >&2
|
||||
return 1
|
||||
fi
|
||||
workspace_mount=$(findmnt -n -T "${workspace_real}" -o TARGET 2>/dev/null || true)
|
||||
if [[ "$(readlink -m "${workspace_mount:-/}")" != "${workspace_real}" ]]; then
|
||||
echo "Native CI requires a dedicated volume mounted at ${workspace_real}" >&2
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
}
|
||||
|
||||
prepare_native_workspace() {
|
||||
if [[ "${VLLM_CI_USE_ARTIFACTS:-0}" != "1" ]]; then
|
||||
echo "Native CI requires VLLM_CI_USE_ARTIFACTS=1"
|
||||
return 1
|
||||
fi
|
||||
if ! command -v buildkite-agent >/dev/null 2>&1; then
|
||||
echo "buildkite-agent not found; cannot download ROCm wheel artifact"
|
||||
return 1
|
||||
fi
|
||||
validate_native_workspace || return 1
|
||||
|
||||
local artifact_glob="${VLLM_CI_ARTIFACT_GLOB:-artifacts/vllm-rocm-install/vllm-rocm-install.tar.gz}"
|
||||
local artifact_checksum_glob="${VLLM_CI_ARTIFACT_CHECKSUM_GLOB:-${artifact_glob}.sha256}"
|
||||
local artifact_step="${VLLM_CI_ARTIFACT_STEP:-image-build-amd}"
|
||||
local archive=""
|
||||
local checksum=""
|
||||
local download_dir=""
|
||||
local metadata_dir=""
|
||||
local recorded_base=""
|
||||
local recorded_commit=""
|
||||
local recorded_wheel=""
|
||||
local workspace_dir="${VLLM_CI_WORKSPACE:-/vllm-workspace}"
|
||||
local wheel_dir=""
|
||||
local attempt=0
|
||||
local attempt_dir=""
|
||||
local -a archives=()
|
||||
local -a checksums=()
|
||||
local -a wheels=()
|
||||
|
||||
artifact_work_dir=$(mktemp -d -t vllm-rocm-artifact.XXXXXX) || return 1
|
||||
wheel_dir="${artifact_work_dir}/wheels"
|
||||
mkdir -p "${wheel_dir}" || return 1
|
||||
|
||||
echo "--- Downloading ROCm wheel artifact from ${artifact_step} (native in-pod)"
|
||||
for attempt in 1 2 3; do
|
||||
attempt_dir="${artifact_work_dir}/download-${attempt}"
|
||||
rm -rf "${attempt_dir}" || return 1
|
||||
mkdir -p "${attempt_dir}" || return 1
|
||||
if buildkite-agent artifact download \
|
||||
"${artifact_glob}" "${attempt_dir}" --step "${artifact_step}" \
|
||||
&& buildkite-agent artifact download \
|
||||
"${artifact_checksum_glob}" "${attempt_dir}" --step "${artifact_step}"; then
|
||||
download_dir="${attempt_dir}"
|
||||
break
|
||||
fi
|
||||
echo "Artifact download attempt ${attempt}/3 failed"
|
||||
if [[ "${attempt}" -lt 3 ]]; then
|
||||
sleep $((attempt * 2))
|
||||
fi
|
||||
done
|
||||
if [[ -z "${download_dir}" ]]; then
|
||||
echo "Failed to download ${artifact_glob} and ${artifact_checksum_glob} from ${artifact_step}"
|
||||
return 1
|
||||
fi
|
||||
|
||||
mapfile -t archives < <(
|
||||
find "${download_dir}" -name "vllm-rocm-install.tar.gz" -type f -print
|
||||
)
|
||||
mapfile -t checksums < <(
|
||||
find "${download_dir}" -name "vllm-rocm-install.tar.gz.sha256" -type f -print
|
||||
)
|
||||
if [[ ${#archives[@]} -ne 1 || ${#checksums[@]} -ne 1 ]]; then
|
||||
echo "Expected exactly one ROCm archive and checksum; found ${#archives[@]} archive(s) and ${#checksums[@]} checksum(s)" >&2
|
||||
return 1
|
||||
fi
|
||||
archive="${archives[0]}"
|
||||
checksum="${checksums[0]}"
|
||||
if [[ "$(dirname "${archive}")" != "$(dirname "${checksum}")" ]]; then
|
||||
echo "ROCm archive and checksum were downloaded to different directories" >&2
|
||||
return 1
|
||||
fi
|
||||
(
|
||||
cd "$(dirname "${archive}")"
|
||||
sha256sum -c "$(basename "${checksum}")"
|
||||
) || return 1
|
||||
|
||||
tar --no-same-owner -xzf "${archive}" -C "${wheel_dir}" || return 1
|
||||
mapfile -t wheels < <(
|
||||
find "${wheel_dir}" -maxdepth 1 -type f -name '*.whl' -print
|
||||
)
|
||||
if [[ ${#wheels[@]} -ne 1 ]]; then
|
||||
echo "ROCm artifact must contain exactly one top-level wheel; found ${#wheels[@]}" >&2
|
||||
return 1
|
||||
fi
|
||||
metadata_dir="${wheel_dir}/.vllm-ci-artifact"
|
||||
for metadata_file in commit.txt native-base-image.txt wheel-filename.txt; do
|
||||
if [[ ! -s "${metadata_dir}/${metadata_file}" ]]; then
|
||||
echo "ROCm artifact metadata is missing ${metadata_file}" >&2
|
||||
return 1
|
||||
fi
|
||||
done
|
||||
for metadata_file in ci-base-image.txt fallback-image.txt; do
|
||||
if [[ ! -f "${metadata_dir}/${metadata_file}" ]]; then
|
||||
echo "ROCm artifact metadata is missing ${metadata_file}" >&2
|
||||
return 1
|
||||
fi
|
||||
done
|
||||
|
||||
recorded_commit=$(tr -d '\r\n' < "${metadata_dir}/commit.txt")
|
||||
recorded_base=$(tr -d '\r\n' < "${metadata_dir}/native-base-image.txt")
|
||||
recorded_wheel=$(tr -d '\r\n' < "${metadata_dir}/wheel-filename.txt")
|
||||
if [[ -z "${BUILDKITE_COMMIT:-}" || "${recorded_commit}" != "${BUILDKITE_COMMIT}" ]]; then
|
||||
echo "ROCm artifact commit ${recorded_commit} does not match ${BUILDKITE_COMMIT:-unset}" >&2
|
||||
return 1
|
||||
fi
|
||||
if [[ -z "${VLLM_CI_BASE_IMAGE:-}" || "${recorded_base}" != "${VLLM_CI_BASE_IMAGE}" ]]; then
|
||||
echo "ROCm artifact base ${recorded_base} does not match ${VLLM_CI_BASE_IMAGE:-unset}" >&2
|
||||
return 1
|
||||
fi
|
||||
if [[ "${recorded_wheel}" != "$(basename "${wheels[0]}")" ]]; then
|
||||
echo "ROCm artifact wheel manifest ${recorded_wheel} does not match $(basename "${wheels[0]}")" >&2
|
||||
return 1
|
||||
fi
|
||||
for required_dir in tests .buildkite requirements; do
|
||||
if [[ ! -d "${wheel_dir}/${required_dir}" ]]; then
|
||||
echo "ROCm wheel artifact did not contain ${required_dir}/" >&2
|
||||
return 1
|
||||
fi
|
||||
done
|
||||
|
||||
echo "--- Installing ROCm wheel into pod environment"
|
||||
python3 -m pip install --no-deps --force-reinstall "${wheels[0]}" || return 1
|
||||
|
||||
echo "--- Preparing ${workspace_dir} from artifact"
|
||||
find "${workspace_dir}" -mindepth 1 -maxdepth 1 -exec rm -rf -- {} + || return 1
|
||||
tar -C "${wheel_dir}" \
|
||||
--exclude='*.whl' \
|
||||
--exclude='.vllm-ci-artifact' \
|
||||
--exclude='./.vllm-ci-artifact' \
|
||||
-cf - . | tar --no-same-owner -C "${workspace_dir}" -xf - || return 1
|
||||
if [[ ! -d "${workspace_dir}/tests" ]]; then
|
||||
echo "Failed to stage the native test workspace" >&2
|
||||
return 1
|
||||
fi
|
||||
|
||||
return 0
|
||||
}
|
||||
|
||||
initialize_native_environment() {
|
||||
local job_id="${BUILDKITE_JOB_ID:-${BUILDKITE_PARALLEL_JOB:-local}}"
|
||||
local job_id_suffix=""
|
||||
local native_root=""
|
||||
local hf_mount=""
|
||||
|
||||
if [[ "$(id -u)" -ne 0 ]]; then
|
||||
echo "Native ROCm CI currently requires the ci_base container to run as root" >&2
|
||||
return 1
|
||||
fi
|
||||
|
||||
job_id="${job_id//[^A-Za-z0-9_.-]/_}"
|
||||
job_id_suffix="${job_id##*-}"
|
||||
job_id_suffix="${job_id_suffix:0:12}"
|
||||
native_root="/tmp/vllm-native-${job_id}"
|
||||
TMPDIR="/tmp/vllm-${job_id_suffix}/tmp"
|
||||
VLLM_RPC_BASE_PATH="/tmp"
|
||||
TORCHINDUCTOR_CACHE_DIR="${native_root}/cache/torchinductor"
|
||||
TRITON_CACHE_DIR="${native_root}/cache/triton"
|
||||
VLLM_CACHE_ROOT="${native_root}/cache/vllm"
|
||||
XDG_CACHE_HOME="${native_root}/cache/xdg"
|
||||
: "${HF_HOME:=/home/buildkite-agent/huggingface}"
|
||||
: "${HF_HUB_DOWNLOAD_TIMEOUT:=300}"
|
||||
: "${HF_HUB_ETAG_TIMEOUT:=60}"
|
||||
export TMPDIR VLLM_RPC_BASE_PATH
|
||||
export TORCHINDUCTOR_CACHE_DIR TRITON_CACHE_DIR VLLM_CACHE_ROOT XDG_CACHE_HOME
|
||||
export HF_HOME HF_HUB_DOWNLOAD_TIMEOUT HF_HUB_ETAG_TIMEOUT
|
||||
export PYTORCH_ROCM_ARCH=""
|
||||
|
||||
mkdir -p "${TMPDIR}" \
|
||||
"${TORCHINDUCTOR_CACHE_DIR}" \
|
||||
"${TRITON_CACHE_DIR}" \
|
||||
"${VLLM_CACHE_ROOT}" \
|
||||
"${XDG_CACHE_HOME}" \
|
||||
"${HF_HOME}" || return 1
|
||||
|
||||
echo "Native compile caches: VLLM_CACHE_ROOT=${VLLM_CACHE_ROOT} TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR}"
|
||||
|
||||
if [[ "${VLLM_CI_REQUIRE_PERSISTENT_HF_CACHE:-0}" == "1" ]]; then
|
||||
if ! command -v findmnt >/dev/null 2>&1; then
|
||||
echo "findmnt is required to verify the native Hugging Face cache mount" >&2
|
||||
return 1
|
||||
fi
|
||||
hf_mount=$(findmnt -n -T "${HF_HOME}" -o TARGET 2>/dev/null || true)
|
||||
if [[ -z "${hf_mount}" || "${hf_mount}" == "/" ]]; then
|
||||
echo "Native CI requires a persistent volume mounted at or above ${HF_HOME}" >&2
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
}
|
||||
|
||||
run_native_preflight() {
|
||||
local expected_gpus="${VLLM_CI_EXPECTED_GPU_COUNT:-1}"
|
||||
|
||||
if [[ ! "${expected_gpus}" =~ ^[0-9]+$ ]]; then
|
||||
echo "Invalid VLLM_CI_EXPECTED_GPU_COUNT=${expected_gpus}" >&2
|
||||
return 1
|
||||
fi
|
||||
|
||||
python3 -c "import encodings, importlib.metadata as im, importlib.util as iu; [im.version(d) for d in ('transformers', 'torch', 'ray', 'sympy', 'markupsafe', 'vllm')]; missing=[m for m in ('torch.utils.model_zoo', 'transformers.models.nomic_bert', 'ray.dag', 'sympy.physics', 'markupsafe._speedups') if iu.find_spec(m) is None]; assert not missing, missing" || return 1
|
||||
|
||||
if [[ "${expected_gpus}" == "0" ]]; then
|
||||
echo "Native CPU-only AMD job: skipping ROCm device validation"
|
||||
return 0
|
||||
fi
|
||||
|
||||
echo "--- ROCm info"
|
||||
rocminfo || return 1
|
||||
VLLM_CI_EXPECTED_GPU_COUNT="${expected_gpus}" python3 - <<'PY'
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
expected = int(os.environ["VLLM_CI_EXPECTED_GPU_COUNT"])
|
||||
assert torch.version.hip, "PyTorch is not a ROCm build"
|
||||
assert torch.cuda.is_available(), "ROCm GPU is not available to PyTorch"
|
||||
actual = torch.cuda.device_count()
|
||||
assert actual == expected, f"Expected {expected} ROCm GPU(s), found {actual}"
|
||||
PY
|
||||
}
|
||||
|
||||
is_multi_node() {
|
||||
local cmds="$1"
|
||||
# Primary signal: NUM_NODES environment variable set by the pipeline
|
||||
@@ -350,7 +646,58 @@ re_quote_pytest_markers() {
|
||||
# Main
|
||||
###############################################################################
|
||||
|
||||
# --- GPU initialization ---
|
||||
if is_native_runtime; then
|
||||
echo "--- Native in-pod ROCm CI (AMD_CI_RUNTIME=${AMD_CI_RUNTIME:-unset}, NATIVE_CI=${NATIVE_CI:-unset})"
|
||||
artifact_work_dir=""
|
||||
|
||||
cleanup_native_workspace() {
|
||||
if [[ -n "${artifact_work_dir}" ]]; then
|
||||
rm -rf "${artifact_work_dir}"
|
||||
fi
|
||||
}
|
||||
trap cleanup_native_workspace EXIT
|
||||
|
||||
if [[ -n "${VLLM_TEST_COMMANDS:-}" ]]; then
|
||||
commands="${VLLM_TEST_COMMANDS}"
|
||||
commands_source="env"
|
||||
else
|
||||
commands="$*"
|
||||
commands_source="argv"
|
||||
if [[ -z "$commands" ]]; then
|
||||
echo "Error: No test commands provided for native CI." >&2
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
if [[ "$commands_source" == "argv" ]]; then
|
||||
commands=$(re_quote_pytest_markers "$commands")
|
||||
fi
|
||||
|
||||
if is_multi_node "$commands"; then
|
||||
echo "Native CI does not support multi-node jobs yet."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if ! initialize_native_environment; then
|
||||
echo "Failed to initialize the native test environment"
|
||||
exit 1
|
||||
fi
|
||||
if ! prepare_native_workspace; then
|
||||
echo "Failed to prepare native test workspace"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
export PYTHONPATH="${VLLM_CI_WORKSPACE:-/vllm-workspace}"
|
||||
|
||||
echo "Native test commands: $commands"
|
||||
run_native_preflight || exit 1
|
||||
# Keep AMD CI orchestration variables out of vLLM's runtime environment.
|
||||
clear_ci_orchestration_env
|
||||
/bin/bash -o pipefail -c "${commands}"
|
||||
handle_pytest_exit "$?"
|
||||
fi
|
||||
|
||||
# --- GPU initialization for legacy Docker execution ---
|
||||
echo "--- ROCm info"
|
||||
rocminfo
|
||||
|
||||
@@ -452,25 +799,27 @@ fi
|
||||
|
||||
echo "Final commands: $commands"
|
||||
|
||||
# The ROCm test image often ships /vllm-workspace without .git (artifact tarball unpack).
|
||||
# tests/standalone_tests/python_only_compile.sh uses merge-base(HEAD, origin/main) for
|
||||
# wheels.vllm.ai; compute on the agent (full git checkout) and pass into the container.
|
||||
vllm_standalone_merge_base=""
|
||||
checkout="${BUILDKITE_BUILD_CHECKOUT_PATH:-}"
|
||||
if [[ -z "${checkout}" || ! -d "${checkout}" ]]; then
|
||||
checkout="."
|
||||
standalone_merge_base_env=()
|
||||
if [[ "$commands" == *python_only_compile.sh* ]]; then
|
||||
# The ROCm test image often ships /vllm-workspace without .git. Resolve the
|
||||
# wheels.vllm.ai commit from the agent checkout for this test only.
|
||||
vllm_standalone_merge_base=""
|
||||
checkout="${BUILDKITE_BUILD_CHECKOUT_PATH:-}"
|
||||
if [[ -z "${checkout}" || ! -d "${checkout}" ]]; then
|
||||
checkout="."
|
||||
fi
|
||||
# Pass safe.directory per-command because Buildkite uses mixed user IDs.
|
||||
if git -c "safe.directory=${checkout}" -C "${checkout}" rev-parse --is-inside-work-tree >/dev/null 2>&1; then
|
||||
vllm_standalone_merge_base="$(
|
||||
git -c "safe.directory=${checkout}" -C "${checkout}" merge-base HEAD origin/main 2>/dev/null || true
|
||||
)"
|
||||
fi
|
||||
if [[ -z "${vllm_standalone_merge_base}" ]]; then
|
||||
vllm_standalone_merge_base="${BUILDKITE_COMMIT:-}"
|
||||
fi
|
||||
echo "INFO: passing CI_STANDALONE_MERGE_BASE into container: ${vllm_standalone_merge_base}"
|
||||
standalone_merge_base_env=(-e "CI_STANDALONE_MERGE_BASE=${vllm_standalone_merge_base}")
|
||||
fi
|
||||
# Pass safe.directory per-command (-c) because buildkite runs will always fail
|
||||
# the next check on git 2.35.2+ due to mixed uses of root and buildkite-agent/uids.
|
||||
if git -c "safe.directory=${checkout}" -C "${checkout}" rev-parse --is-inside-work-tree >/dev/null 2>&1; then
|
||||
vllm_standalone_merge_base="$(
|
||||
git -c "safe.directory=${checkout}" -C "${checkout}" merge-base HEAD origin/main 2>/dev/null || true
|
||||
)"
|
||||
fi
|
||||
if [[ -z "${vllm_standalone_merge_base}" ]]; then
|
||||
vllm_standalone_merge_base="${BUILDKITE_COMMIT:-}"
|
||||
fi
|
||||
echo "INFO: passing VLLM_STANDALONE_MERGE_BASE into container: ${vllm_standalone_merge_base}"
|
||||
|
||||
MYPYTHONPATH="/vllm-workspace"
|
||||
|
||||
@@ -501,6 +850,7 @@ else
|
||||
fi
|
||||
|
||||
# --- Route: multi-node vs single-node ---
|
||||
clear_ci_orchestration_env
|
||||
if is_multi_node "$commands"; then
|
||||
echo "--- Multi-node job detected"
|
||||
export DCKR_VER=$(docker --version | sed 's/Docker version \(.*\), build .*/\1/')
|
||||
@@ -589,7 +939,9 @@ else
|
||||
-e FORCE_COLOR \
|
||||
-e CLICOLOR_FORCE \
|
||||
-e PY_COLORS \
|
||||
-e PYTHONFAULTHANDLER \
|
||||
-e PYTEST_ADDOPTS \
|
||||
-e PYTEST_TIMEOUT \
|
||||
-v "${HF_CACHE}:${HF_MOUNT}" \
|
||||
-e "HF_HOME=${HF_MOUNT}" \
|
||||
-e "PYTHONPATH=${MYPYTHONPATH}" \
|
||||
@@ -599,7 +951,7 @@ else
|
||||
-e "VLLM_CACHE_ROOT=${CONTAINER_CACHE_ROOT}/vllm" \
|
||||
-e "XDG_CACHE_HOME=${CONTAINER_CACHE_ROOT}/xdg" \
|
||||
-e "PYTORCH_ROCM_ARCH=" \
|
||||
-e "VLLM_STANDALONE_MERGE_BASE=${vllm_standalone_merge_base}" \
|
||||
"${standalone_merge_base_env[@]}" \
|
||||
--name "${container_name}" \
|
||||
"${image_name}" \
|
||||
/bin/bash -c "${CONTAINER_PREFLIGHT} && ${commands}"
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
#!/bin/bash
|
||||
set -euox pipefail
|
||||
|
||||
export VLLM_CPU_KVCACHE_SPACE=1
|
||||
export VLLM_CPU_KVCACHE_SPACE=1
|
||||
export VLLM_CPU_CI_ENV=1
|
||||
# Reduce sub-processes for acceleration
|
||||
export TORCH_COMPILE_DISABLE=1
|
||||
# Skip torch.compile via vLLM's --enforce-eager flag (passed below) instead of
|
||||
# TORCH_COMPILE_DISABLE=1, which torch 2.12 no longer treats as a silent no-op
|
||||
# when callers specify fullgraph=True.
|
||||
export VLLM_ENABLE_V1_MULTIPROCESSING=0
|
||||
|
||||
SDE_ARCHIVE="sde-external-10.7.0-2026-02-18-lin.tar.xz"
|
||||
@@ -49,15 +50,15 @@ wait_for_pid_and_check_log() {
|
||||
}
|
||||
|
||||
# Test Sky Lake (AVX512F)
|
||||
./sde/sde64 -skl -- python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --dtype bfloat16 > test_0.log 2>&1 &
|
||||
./sde/sde64 -skl -- python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --dtype bfloat16 --enforce-eager > test_0.log 2>&1 &
|
||||
PID_TEST_0=$!
|
||||
|
||||
# Test Cascade Lake (AVX512F + VNNI)
|
||||
./sde/sde64 -clx -- python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --dtype bfloat16 > test_1.log 2>&1 &
|
||||
./sde/sde64 -clx -- python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --dtype bfloat16 --enforce-eager > test_1.log 2>&1 &
|
||||
PID_TEST_1=$!
|
||||
|
||||
# Test Cooper Lake (AVX512F + VNNI + BF16)
|
||||
./sde/sde64 -cpx -- python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --dtype bfloat16 > test_2.log 2>&1 &
|
||||
./sde/sde64 -cpx -- python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --dtype bfloat16 --enforce-eager > test_2.log 2>&1 &
|
||||
PID_TEST_2=$!
|
||||
|
||||
wait_for_pid_and_check_log $PID_TEST_0 test_0.log
|
||||
|
||||
@@ -40,7 +40,9 @@ function cpu_tests() {
|
||||
pytest -x -v -s tests/kernels/moe/test_cpu_fused_moe.py
|
||||
pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py
|
||||
pytest -x -v -s tests/kernels/moe/test_cpu_int4_moe.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_cpu_short_conv.py"
|
||||
pytest -x -v -s tests/kernels/mamba/test_cpu_short_conv.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_causal_conv1d.py
|
||||
pytest -x -v -s tests/kernels/mamba/test_mamba_ssm.py"
|
||||
|
||||
# skip tests requiring model downloads if HF_TOKEN is not set
|
||||
# due to rate-limits
|
||||
@@ -97,3 +99,4 @@ function cpu_tests() {
|
||||
# All of CPU tests are expected to be finished less than 40 mins.
|
||||
export -f cpu_tests
|
||||
timeout 2h bash -c cpu_tests
|
||||
|
||||
|
||||
@@ -35,6 +35,7 @@ case "${test_suite}" in
|
||||
pytest -v -s v1/worker --ignore=v1/worker/test_gpu_model_runner.py --ignore=v1/worker/test_worker_memory_snapshot.py
|
||||
pytest -v -s v1/structured_output
|
||||
pytest -v -s v1/test_serial_utils.py
|
||||
pytest -v -s v1/e2e/general/test_correctness_sliding_window.py --deselect="tests/v1/e2e/general/test_correctness_sliding_window.py::test_sliding_window_retrieval[True-1-5-google/gemma-3-1b-it]"
|
||||
pytest -v -s v1/spec_decode --ignore=v1/spec_decode/test_max_len.py --ignore=v1/spec_decode/test_speculators_eagle3.py --ignore=v1/spec_decode/test_acceptance_length.py --ignore=v1/spec_decode/test_speculators_correctness.py
|
||||
pytest -v -s v1/kv_connector/unit --ignore=v1/kv_connector/unit/test_multi_connector.py --ignore=v1/kv_connector/unit/test_example_connector.py --ignore=v1/kv_connector/unit/test_lmcache_integration.py --ignore=v1/kv_connector/unit/test_hf3fs_client.py --ignore=v1/kv_connector/unit/test_hf3fs_connector.py --ignore=v1/kv_connector/unit/test_hf3fs_metadata_server.py --ignore=v1/kv_connector/unit/test_offloading_connector.py
|
||||
;;
|
||||
|
||||
@@ -13,6 +13,18 @@ metadata_get() {
|
||||
fi
|
||||
}
|
||||
|
||||
use_ci_base_if_present() {
|
||||
local ci_base_image=""
|
||||
|
||||
ci_base_image="$(metadata_get rocm-ci-base-image)"
|
||||
if [[ -z "${ci_base_image}" ]]; then
|
||||
return 1
|
||||
fi
|
||||
|
||||
export CI_BASE_IMAGE="${ci_base_image}"
|
||||
echo "Using ROCm ci_base image selected by the preceding build step: ${CI_BASE_IMAGE}"
|
||||
}
|
||||
|
||||
use_refreshed_base_if_present() {
|
||||
local base_refreshed=""
|
||||
|
||||
@@ -22,15 +34,12 @@ use_refreshed_base_if_present() {
|
||||
fi
|
||||
|
||||
export BASE_IMAGE
|
||||
export CI_BASE_IMAGE
|
||||
export IMAGE_TAG_LATEST
|
||||
|
||||
BASE_IMAGE="$(metadata_get rocm-base-image)"
|
||||
CI_BASE_IMAGE="$(metadata_get rocm-ci-base-image)"
|
||||
IMAGE_TAG_LATEST="$(metadata_get rocm-ci-image-descriptive)"
|
||||
|
||||
echo "Using refreshed ROCm base image for test image: ${BASE_IMAGE}"
|
||||
echo "Using refreshed ROCm ci_base image for test image: ${CI_BASE_IMAGE}"
|
||||
if [[ -n "${IMAGE_TAG_LATEST}" ]]; then
|
||||
echo "Also tagging full ROCm CI image as: ${IMAGE_TAG_LATEST}"
|
||||
fi
|
||||
@@ -41,6 +50,8 @@ use_refreshed_base_if_present() {
|
||||
main() {
|
||||
local base_refreshed=0
|
||||
|
||||
use_ci_base_if_present || true
|
||||
|
||||
if use_refreshed_base_if_present; then
|
||||
base_refreshed=1
|
||||
fi
|
||||
|
||||
@@ -6,8 +6,14 @@ set -ex
|
||||
# manylinux platform tag with auditwheel.
|
||||
# Index generation is handled separately by generate-and-upload-nightly-index.sh.
|
||||
|
||||
# shellcheck source=lib/manylinux.sh
|
||||
source .buildkite/scripts/lib/manylinux.sh
|
||||
# auditwheel is Linux-only; macOS wheels already carry a valid tag, so skip the
|
||||
# manylinux retag for them.
|
||||
WHEEL_PLATFORM="${VLLM_WHEEL_PLATFORM:-linux}"
|
||||
|
||||
if [[ "$WHEEL_PLATFORM" == "linux" ]]; then
|
||||
# shellcheck source=lib/manylinux.sh
|
||||
source .buildkite/scripts/lib/manylinux.sh
|
||||
fi
|
||||
|
||||
BUCKET="vllm-wheels"
|
||||
SUBPATH=$BUILDKITE_COMMIT
|
||||
@@ -27,8 +33,10 @@ wheel="${wheel_files[0]}"
|
||||
|
||||
# ========= detect manylinux tag and rename ==========
|
||||
|
||||
wheel="$(apply_manylinux_tag "$wheel")"
|
||||
echo "Renamed wheel to: $wheel"
|
||||
if [[ "$WHEEL_PLATFORM" == "linux" ]]; then
|
||||
wheel="$(apply_manylinux_tag "$wheel")"
|
||||
echo "Renamed wheel to: $wheel"
|
||||
fi
|
||||
|
||||
# Extract the version from the wheel
|
||||
version=$(unzip -p "$wheel" '**/METADATA' | grep '^Version: ' | cut -d' ' -f2)
|
||||
|
||||
@@ -113,8 +113,8 @@ $PYTHON .buildkite/scripts/generate-nightly-index.py \
|
||||
echo "Uploading indices to $S3_COMMIT_PREFIX"
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "$S3_COMMIT_PREFIX"
|
||||
|
||||
# Update rocm/nightly/ if on main branch and not a PR
|
||||
if [[ "$BUILDKITE_BRANCH" == "main" && "$BUILDKITE_PULL_REQUEST" == "false" ]] || [[ "$NIGHTLY" == "1" ]]; then
|
||||
# Only scheduled nightly builds should update the moving nightly index.
|
||||
if [[ "${NIGHTLY:-0}" == "1" ]]; then
|
||||
echo "Updating rocm/nightly/ index..."
|
||||
aws s3 cp --recursive "$INDICES_OUTPUT_DIR/" "s3://$BUCKET/rocm/nightly/"
|
||||
fi
|
||||
@@ -147,7 +147,7 @@ echo ""
|
||||
echo "Install command (by commit):"
|
||||
echo " pip install vllm --extra-index-url https://${BUCKET}.s3.amazonaws.com/$ROCM_SUBPATH/"
|
||||
echo ""
|
||||
if [[ "$BUILDKITE_BRANCH" == "main" ]] || [[ "$NIGHTLY" == "1" ]]; then
|
||||
if [[ "${NIGHTLY:-0}" == "1" ]]; then
|
||||
echo "Install command (nightly):"
|
||||
echo " pip install vllm --extra-index-url https://${BUCKET}.s3.amazonaws.com/rocm/nightly/"
|
||||
fi
|
||||
|
||||
+427
-239
File diff suppressed because it is too large
Load Diff
@@ -16,8 +16,9 @@ steps:
|
||||
parallelism: 2
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 95
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 125
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -4,7 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Basic Correctness
|
||||
key: basic-correctness
|
||||
timeout_in_minutes: 45
|
||||
timeout_in_minutes: 68
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -18,7 +18,8 @@ steps:
|
||||
- pytest -v -s basic_correctness/test_cpu_offload.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 70
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 60
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
@@ -4,7 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Benchmarks CLI Test
|
||||
key: benchmarks-cli-test
|
||||
timeout_in_minutes: 30
|
||||
timeout_in_minutes: 45
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -13,7 +13,9 @@ steps:
|
||||
- pytest -v -s benchmarks/
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 40
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ steps:
|
||||
- pytest -v -s cuda/test_platform_no_cuda_init.py
|
||||
|
||||
- label: Cudagraph
|
||||
device: h200_35gb
|
||||
key: cudagraph
|
||||
timeout_in_minutes: 30
|
||||
source_file_dependencies:
|
||||
@@ -25,7 +26,10 @@ steps:
|
||||
- vllm/v1/cudagraph_dispatcher.py
|
||||
- vllm/config/compilation.py
|
||||
- vllm/compilation
|
||||
- vllm/v1/worker/encoder_cudagraph.py
|
||||
- vllm/v1/worker/encoder_cudagraph_defs.py
|
||||
commands:
|
||||
- pytest -v -s v1/cudagraph/test_cudagraph_dispatch.py
|
||||
- pytest -v -s v1/cudagraph/test_cudagraph_mode.py
|
||||
- pytest -v -s v1/cudagraph/test_breakable_cudagraph.py
|
||||
- pytest -v -s v1/cudagraph/test_breakable_cudagraph.py
|
||||
- pytest -v -s v1/cudagraph/test_encoder_cudagraph.py
|
||||
|
||||
@@ -15,8 +15,9 @@ steps:
|
||||
- bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_4
|
||||
timeout_in_minutes: 85
|
||||
timeout_in_minutes: 60
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -65,8 +66,9 @@ steps:
|
||||
- DP_EP=1 bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_4
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 40
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -90,8 +92,9 @@ steps:
|
||||
- CROSS_LAYERS_BLOCKS=True bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_4
|
||||
timeout_in_minutes: 85
|
||||
timeout_in_minutes: 60
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -115,8 +118,9 @@ steps:
|
||||
- HYBRID_SSM=1 bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_4
|
||||
timeout_in_minutes: 80
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -171,8 +175,9 @@ steps:
|
||||
- bash v1/kv_connector/nixl_integration/config_sweep_spec_decode_test.sh
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_2
|
||||
timeout_in_minutes: 70
|
||||
timeout_in_minutes: 45
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -182,7 +187,7 @@ steps:
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- uv pip install --system -r /vllm-workspace/requirements/kv_connectors_rocm.txt
|
||||
- ATTENTION_BACKEND=TRITON_ATTN bash v1/kv_connector/nixl_integration/config_sweep_spec_decode_test.sh
|
||||
- KV_CACHE_MEMORY_BYTES=8G ATTENTION_BACKEND=TRITON_ATTN bash v1/kv_connector/nixl_integration/config_sweep_spec_decode_test.sh
|
||||
|
||||
- label: MultiConnector (Nixl+Offloading) PD edge cases (2 GPUs)
|
||||
key: multiconnector-nixl-offloading-pd-edge-cases-2-gpus
|
||||
@@ -198,3 +203,25 @@ steps:
|
||||
commands:
|
||||
- bash /vllm-workspace/.buildkite/scripts/install-kv-connectors.sh
|
||||
- bash v1/kv_connector/nixl_integration/run_multi_connector_edge_case_test.sh
|
||||
|
||||
# P TP 4 - D DPEP 4 test case for DSv4-Flash
|
||||
- label: DSv4-Flash Disaggregated DP EP
|
||||
key: dsv4-flash-disaggregated
|
||||
timeout_in_minutes: 60
|
||||
device: h200
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
num_devices: 8
|
||||
env:
|
||||
ENABLE_HMA_FLAG: "1"
|
||||
DP_EP: "1"
|
||||
GPU_MEMORY_UTILIZATION: "0.85"
|
||||
PREFILLER_TP_SIZE: "4"
|
||||
DECODER_TP_SIZE: "4"
|
||||
PREFILL_BLOCK_SIZE: "256"
|
||||
DECODE_BLOCK_SIZE: "256"
|
||||
MODEL_NAMES: "deepseek-ai/DeepSeek-V4-Flash"
|
||||
VLLM_SERVE_EXTRA_ARGS: "--trust-remote-code,--kv-cache-dtype,fp8"
|
||||
commands:
|
||||
- bash /vllm-workspace/.buildkite/scripts/install-kv-connectors.sh
|
||||
- bash v1/kv_connector/nixl_integration/run_accuracy_test.sh
|
||||
|
||||
@@ -39,7 +39,9 @@ steps:
|
||||
- DP_SIZE=2 pytest -v -s entrypoints/openai/test_multi_api_servers.py
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_2
|
||||
timeout_in_minutes: 45
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -28,8 +28,9 @@ steps:
|
||||
- pytest -v -s engine test_sequence.py test_config.py test_logger.py test_vllm_port.py test_jit_monitor.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 50
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 40
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
@@ -44,14 +45,14 @@ steps:
|
||||
- pytest -v -s v1/engine --ignore v1/engine/test_preprocess_error_handling.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 55
|
||||
device: mi250_1
|
||||
timeout_in_minutes: 45
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: e2e Scheduling (1 GPU)
|
||||
key: e2e-scheduling-1-gpu
|
||||
timeout_in_minutes: 35
|
||||
timeout_in_minutes: 53
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/v1/
|
||||
@@ -60,8 +61,8 @@ steps:
|
||||
- pytest -v -s v1/e2e/general/test_async_scheduling.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 70
|
||||
device: mi250_1
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
@@ -76,8 +77,8 @@ steps:
|
||||
- pytest -v -s v1/e2e/general --ignore v1/e2e/general/test_async_scheduling.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 60
|
||||
device: mi250_1
|
||||
timeout_in_minutes: 50
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -114,7 +115,9 @@ steps:
|
||||
- pytest -v -s v1/e2e/spec_decode/test_spec_decode.py -k "tensor_parallelism"
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_2
|
||||
timeout_in_minutes: 30
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Entrypoints Unit Tests
|
||||
device: h200_35gb
|
||||
key: entrypoints-unit-tests
|
||||
timeout_in_minutes: 25
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -15,6 +16,7 @@ steps:
|
||||
- pytest -v -s entrypoints/weight_transfer
|
||||
|
||||
- label: Entrypoints Integration (LLM)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-llm
|
||||
timeout_in_minutes: 60
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -28,16 +30,16 @@ steps:
|
||||
- pytest -v -s entrypoints/llm/offline_mode # Needs to avoid interference with other tests
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
# TODO(akaratza): Test after Torch >= 2.12 bump
|
||||
soft_fail: true
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (API Server)
|
||||
key: entrypoints-integration-api-server
|
||||
device: h200_35gb
|
||||
timeout_in_minutes: 50
|
||||
timeout_in_minutes: 75
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -50,13 +52,16 @@ steps:
|
||||
- pytest -v -s entrypoints/scale_out
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 65
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (API Server OpenAI - Part 1)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-api-server-openai-part-1
|
||||
timeout_in_minutes: 45
|
||||
timeout_in_minutes: 68
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -67,14 +72,16 @@ steps:
|
||||
- pytest -v -s entrypoints/openai --ignore=entrypoints/openai/completion --ignore=entrypoints/openai/chat_completion --ignore=entrypoints/openai/responses --ignore=entrypoints/openai/correctness
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 65
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (API Server OpenAI - Part 2)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-api-server-openai-part-2
|
||||
timeout_in_minutes: 45
|
||||
timeout_in_minutes: 83
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -86,12 +93,14 @@ steps:
|
||||
- pytest -v -s entrypoints/openai/completion --ignore=entrypoints/openai/completion/test_tensorizer_entrypoint.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 80
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 70
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (API Server Generate)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-api-server-generate
|
||||
timeout_in_minutes: 50
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -108,12 +117,14 @@ steps:
|
||||
- pytest -v -s entrypoints/anthropic
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 65
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Entrypoints Integration (Responses API)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-responses-api
|
||||
timeout_in_minutes: 50
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -148,8 +159,9 @@ steps:
|
||||
- pytest -v -s entrypoints/multimodal
|
||||
|
||||
- label: Entrypoints Integration (Pooling)
|
||||
device: h200_35gb
|
||||
key: entrypoints-integration-pooling
|
||||
timeout_in_minutes: 50
|
||||
timeout_in_minutes: 75
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -169,7 +181,9 @@ steps:
|
||||
- pytest -s entrypoints/openai/correctness/
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 30
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -16,7 +16,9 @@ steps:
|
||||
- pytest -v -s distributed/test_eplb_utils.py
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 30
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -15,6 +15,7 @@ steps:
|
||||
- pytest -v -s tests/kernels/ir
|
||||
|
||||
- label: Kernels Core Operation Test
|
||||
device: h200_35gb
|
||||
key: kernels-core-operation-test
|
||||
timeout_in_minutes: 120
|
||||
source_file_dependencies:
|
||||
@@ -79,7 +80,8 @@ steps:
|
||||
parallelism: 2
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 90
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
@@ -117,7 +119,9 @@ steps:
|
||||
parallelism: 2
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 120
|
||||
source_file_dependencies:
|
||||
- csrc/quantization/
|
||||
- vllm/model_executor/layers/quantization
|
||||
@@ -147,8 +151,9 @@ steps:
|
||||
parallelism: 5
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 65
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 55
|
||||
source_file_dependencies:
|
||||
- csrc/quantization/cutlass_w8a8/moe/
|
||||
- csrc/moe/
|
||||
@@ -163,6 +168,7 @@ steps:
|
||||
- image-build-amd
|
||||
|
||||
- label: Kernels Mamba Test
|
||||
device: h200_35gb
|
||||
key: kernels-mamba-test
|
||||
timeout_in_minutes: 40
|
||||
source_file_dependencies:
|
||||
@@ -176,9 +182,9 @@ steps:
|
||||
timeout_in_minutes: 25
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/layers/fla/ops/kda.py
|
||||
- vllm/model_executor/layers/fla/ops/chunk_delta_h.py
|
||||
- vllm/model_executor/layers/fla/ops/l2norm.py
|
||||
- vllm/third_party/flash_linear_attention/ops/kda.py
|
||||
- vllm/third_party/flash_linear_attention/ops/chunk_delta_h.py
|
||||
- vllm/third_party/flash_linear_attention/ops/l2norm.py
|
||||
- tests/kernels/test_kda.py
|
||||
commands:
|
||||
- pytest -v -s kernels/test_kda.py
|
||||
@@ -232,6 +238,15 @@ steps:
|
||||
- vllm/v1/attention/backends/mla/flashinfer_mla.py
|
||||
- vllm/v1/attention/selector.py
|
||||
- vllm/platforms/cuda.py
|
||||
- vllm/model_executor/kernels/linear/cute_dsl/ll_bf16.py
|
||||
- vllm/model_executor/kernels/linear/cute_dsl/_ll_bf16_dotprod.py
|
||||
- vllm/model_executor/kernels/linear/cute_dsl/_ll_bf16_splitk.py
|
||||
- vllm/cute_utils/
|
||||
- vllm/model_executor/layers/mamba/ops/gdn_chunk_cutedsl/
|
||||
- vllm/model_executor/layers/fused_moe/router/bf16x3_router_gemm_cutedsl.py
|
||||
- tests/kernels/mamba/test_gdn_prefill_cutedsl.py
|
||||
- tests/kernels/test_bf16x3_router_gemm_cutedsl.py
|
||||
- tests/kernels/test_ll_bf16_gemm.py
|
||||
- tests/kernels/test_top_k_per_row.py
|
||||
commands:
|
||||
- nvidia-smi
|
||||
@@ -260,6 +275,9 @@ steps:
|
||||
- pytest -v -s tests/kernels/moe/test_flashinfer_moe.py
|
||||
- pytest -v -s tests/kernels/moe/test_trtllm_nvfp4_moe.py
|
||||
- pytest -v -s tests/kernels/moe/test_cutedsl_moe.py
|
||||
- pytest -v -s tests/kernels/mamba/test_gdn_prefill_cutedsl.py
|
||||
- pytest -v -s tests/kernels/test_bf16x3_router_gemm_cutedsl.py
|
||||
- pytest -v -s tests/kernels/test_ll_bf16_gemm.py
|
||||
# e2e
|
||||
- pytest -v -s tests/models/quantization/test_nvfp4.py
|
||||
|
||||
|
||||
@@ -14,8 +14,9 @@ steps:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-small.txt
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 55
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 45
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -78,6 +79,28 @@ steps:
|
||||
commands:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-small-tp.txt
|
||||
|
||||
- label: LM Eval PCP (4xB200)
|
||||
key: lm-eval-pcp-4xb200
|
||||
timeout_in_minutes: 360
|
||||
device: b200-k8s
|
||||
num_devices: 4
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
- tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml
|
||||
- tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml
|
||||
- tests/evals/gsm8k/configs/models-pcp.txt
|
||||
- vllm/model_executor/layers/quantization
|
||||
- vllm/config/parallel.py
|
||||
- vllm/distributed/parallel_state.py
|
||||
- vllm/model_executor/layers/attention/mla_attention.py
|
||||
- vllm/model_executor/layers/attention/pcp.py
|
||||
- vllm/v1/worker/gpu/model_runner.py
|
||||
- vllm/v1/worker/gpu/pcp_manager.py
|
||||
autorun_on_main: true
|
||||
commands:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-pcp.txt
|
||||
|
||||
- label: LM Eval Large Models EP (2xB200)
|
||||
key: lm-eval-large-models-ep-2xb200
|
||||
timeout_in_minutes: 60
|
||||
@@ -103,7 +126,7 @@ steps:
|
||||
- vllm/transformers_utils/configs/qwen3_5_moe.py
|
||||
- vllm/model_executor/models/qwen3_next.py
|
||||
- vllm/model_executor/models/qwen3_next_mtp.py
|
||||
- vllm/model_executor/layers/fla/ops/
|
||||
- vllm/third_party/flash_linear_attention/ops/
|
||||
commands:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-qwen35-blackwell.txt
|
||||
|
||||
@@ -117,8 +140,9 @@ steps:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-h200.txt
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_8
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 40
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
commands:
|
||||
|
||||
@@ -14,9 +14,10 @@ steps:
|
||||
parallelism: 4
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
timeout_in_minutes: 65
|
||||
timeout_in_minutes: 85
|
||||
source_file_dependencies:
|
||||
- vllm/lora
|
||||
- tests/lora
|
||||
@@ -46,4 +47,4 @@ steps:
|
||||
- pytest -v -s -x lora/test_qwen3_with_multi_loras.py
|
||||
- pytest -v -s -x lora/test_olmoe_tp.py
|
||||
- pytest -v -s -x lora/test_gptoss_tp.py
|
||||
- pytest -v -s -x lora/test_qwen35_densemodel_lora.py
|
||||
- pytest -v -s -x lora/test_qwen35_densemodel_lora.py
|
||||
|
||||
@@ -23,14 +23,15 @@ steps:
|
||||
- pytest -v -s -m 'not slow_test' v1/spec_decode
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 75
|
||||
timeout_in_minutes: 50
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: V1 Sample + Logits
|
||||
key: v1-sample-logits
|
||||
timeout_in_minutes: 45
|
||||
timeout_in_minutes: 83
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/config/
|
||||
@@ -58,13 +59,16 @@ steps:
|
||||
- pytest -v -s v1/test_outputs.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 70
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: V1 Core + KV + Metrics
|
||||
device: h200_35gb
|
||||
key: v1-core-kv-metrics
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 80
|
||||
source_file_dependencies:
|
||||
- vllm/config/
|
||||
- vllm/distributed/
|
||||
@@ -88,6 +92,7 @@ steps:
|
||||
- tests/v1/kv_offload
|
||||
- tests/v1/simple_kv_offload
|
||||
- tests/v1/worker
|
||||
- tests/v1/streaming_input
|
||||
- tests/v1/kv_connector/unit
|
||||
- tests/v1/ec_connector/unit
|
||||
- tests/v1/metrics
|
||||
@@ -101,6 +106,7 @@ steps:
|
||||
- pytest -v -s v1/kv_offload
|
||||
- pytest -v -s v1/simple_kv_offload
|
||||
- pytest -v -s v1/worker
|
||||
- pytest -v -s v1/streaming_input
|
||||
- pytest -v -s -m 'not cpu_test' v1/kv_connector/unit
|
||||
- pytest -v -s -m 'not cpu_test' v1/ec_connector/unit
|
||||
- pytest -v -s -m 'not cpu_test' v1/metrics
|
||||
@@ -109,8 +115,9 @@ steps:
|
||||
- pytest -v -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 75
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 65
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
@@ -141,6 +148,7 @@ steps:
|
||||
- pytest -v -s -m 'cpu_test' v1/core
|
||||
- pytest -v -s v1/structured_output
|
||||
- pytest -v -s v1/test_serial_utils.py
|
||||
- pytest -v -s v1/cudagraph/test_cudagraph_manager.py
|
||||
- pytest -v -s -m 'cpu_test' v1/kv_connector/unit
|
||||
- pytest -v -s -m 'cpu_test' v1/metrics
|
||||
|
||||
@@ -204,7 +212,7 @@ steps:
|
||||
- vllm/multimodal
|
||||
- examples/
|
||||
commands:
|
||||
- pip install tensorizer # for tensorizer test
|
||||
- pip install --no-deps tensorizer # for tensorizer test
|
||||
# for basic
|
||||
- python3 basic/offline_inference/chat.py
|
||||
- python3 basic/offline_inference/generate.py --model facebook/opt-125m
|
||||
@@ -228,7 +236,9 @@ steps:
|
||||
- python3 features/speculative_decoding/spec_decode_offline.py --test --method eagle3 --num_spec_tokens 3 --dataset-name hf --dataset-path philschmid/mt-bench --num-prompts 80 --temp 0 --top-p 1.0 --top-k -1 --tp 1 --enable-chunked-prefill --max-model-len 1536
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 75
|
||||
source_file_dependencies:
|
||||
- vllm/entrypoints
|
||||
- vllm/multimodal
|
||||
@@ -264,10 +274,11 @@ steps:
|
||||
- pytest -v -s v1/tracing
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_2
|
||||
dind: false
|
||||
device: mi300_2
|
||||
timeout_in_minutes: 30
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
optional: true
|
||||
|
||||
- label: Python-only Installation
|
||||
key: python-only-installation
|
||||
@@ -282,8 +293,8 @@ steps:
|
||||
- bash standalone_tests/python_only_compile.sh
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 45
|
||||
device: mi250_1
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -3,13 +3,16 @@ depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Model Executor
|
||||
device: h200_35gb
|
||||
key: model-executor
|
||||
timeout_in_minutes: 45
|
||||
timeout_in_minutes: 60
|
||||
source_file_dependencies:
|
||||
- vllm/engine/arg_utils.py
|
||||
- vllm/config/model.py
|
||||
- vllm/model_executor
|
||||
- vllm/model_executor/warmup
|
||||
- tests/model_executor
|
||||
- tests/model_executor/test_jit_warmup.py
|
||||
- tests/entrypoints/openai/completion/test_tensorizer_entrypoint.py
|
||||
commands:
|
||||
- apt-get update && apt-get install -y curl libsodium23
|
||||
@@ -25,14 +28,18 @@ steps:
|
||||
- pytest -v -s entrypoints/openai/completion/test_tensorizer_entrypoint.py --timeout=900 --timeout-method=thread
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 60
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
- vllm/engine/arg_utils.py
|
||||
- vllm/config/model.py
|
||||
- vllm/model_executor
|
||||
- vllm/model_executor/warmup
|
||||
- tests/model_executor
|
||||
- tests/model_executor/test_jit_warmup.py
|
||||
- tests/entrypoints/openai/completion/test_tensorizer_entrypoint.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
|
||||
@@ -41,7 +41,7 @@ steps:
|
||||
commands:
|
||||
- set -x
|
||||
- export VLLM_USE_V2_MODEL_RUNNER=1
|
||||
- pip install tensorizer # for tensorizer test
|
||||
- pip install --no-deps tensorizer # for tensorizer test
|
||||
- python3 basic/offline_inference/chat.py # for basic
|
||||
- python3 basic/offline_inference/generate.py --model facebook/opt-125m
|
||||
#- python3 basic/offline_inference/generate.py --model meta-llama/Llama-2-13b-chat-hf --cpu-offload-gb 10 # TODO
|
||||
|
||||
@@ -42,10 +42,25 @@ steps:
|
||||
- pytest -v -s models/test_terratorch.py models/transformers/test_backend.py models/test_registry.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 50
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Inkling Unit Tests (B200)
|
||||
key: inkling-unit-tests-b200
|
||||
timeout_in_minutes: 40
|
||||
device: b200-k8s
|
||||
source_file_dependencies:
|
||||
- vllm/models/inkling/
|
||||
- vllm/cute_utils/
|
||||
- cmake/external_projects/tml_fa4.cmake
|
||||
- tests/models/inkling/
|
||||
commands:
|
||||
# FA4 kernel tests require SM100; the suite skips them elsewhere.
|
||||
- pytest -v -s models/inkling
|
||||
|
||||
- label: Basic Models Test (Other CPU) # 5min
|
||||
key: basic-models-test-other-cpu
|
||||
depends_on:
|
||||
|
||||
@@ -15,11 +15,14 @@ steps:
|
||||
- pytest -v -s models/language -m 'core_model and (not slow_test)'
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 45
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Language Models Tests (Extra Standard) %N
|
||||
device: h200_35gb
|
||||
key: language-models-tests-extra-standard
|
||||
timeout_in_minutes: 40
|
||||
source_file_dependencies:
|
||||
@@ -35,7 +38,9 @@ steps:
|
||||
parallelism: 2
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 40
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -49,8 +54,8 @@ steps:
|
||||
- tests/models/language/pooling/test_classification.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
|
||||
- label: Language Models Tests (Hybrid) %N
|
||||
device: h200_35gb
|
||||
key: language-models-tests-hybrid
|
||||
timeout_in_minutes: 65
|
||||
source_file_dependencies:
|
||||
@@ -58,16 +63,16 @@ steps:
|
||||
- tests/models/language/generation
|
||||
commands:
|
||||
# Install fast path packages for testing against transformers
|
||||
# Note: also needed to run plamo2 model in vLLM
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/state-spaces/mamba@v2.3.0'
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
# Shard hybrid language model tests
|
||||
- pytest -v -s models/language/generation -m hybrid_model --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB
|
||||
# Shard the hybrid language model tests that are numerically stable on Hopper.
|
||||
- pytest -v -s models/language/generation -m hybrid_model -k 'not granite-4.0-tiny-preview' --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB
|
||||
parallelism: 2
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 70
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 60
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
commands:
|
||||
@@ -75,6 +80,20 @@ steps:
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
- pytest -v -s models/language/generation -m hybrid_model --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB
|
||||
|
||||
# Granite 4 hybrid generation is sensitive to hardware-specific Triton SSD
|
||||
# autotuning (https://github.com/vllm-project/vllm/issues/25194). Keep this one
|
||||
# correctness test on L4 until its H200 output matches the Transformers reference.
|
||||
- label: Language Models Tests (Granite L4 Compatibility)
|
||||
key: language-models-tests-granite-l4-compatibility
|
||||
timeout_in_minutes: 65
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/models/language/generation
|
||||
commands:
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/state-spaces/mamba@v2.3.0'
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
- pytest -v -s models/language/generation -m hybrid_model -k 'granite-4.0-tiny-preview'
|
||||
|
||||
- label: Language Models Test (Extended Generation) # 80min
|
||||
device: h200_35gb
|
||||
key: language-models-test-extended-generation
|
||||
@@ -85,7 +104,6 @@ steps:
|
||||
- tests/models/language/generation
|
||||
commands:
|
||||
# Install fast path packages for testing against transformers
|
||||
# Note: also needed to run plamo2 model in vLLM
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/state-spaces/mamba@v2.3.0'
|
||||
- uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0'
|
||||
- pytest -v -s models/language/generation -m '(not core_model) and (not hybrid_model)'
|
||||
@@ -101,10 +119,10 @@ steps:
|
||||
commands:
|
||||
- pytest -v -s models/language/generation_ppl_test
|
||||
|
||||
- label: Language Models Test (Extended Pooling) # 36min
|
||||
- label: Language Models Test (Extended Pooling)
|
||||
device: h200_35gb
|
||||
key: language-models-test-extended-pooling
|
||||
timeout_in_minutes: 70
|
||||
timeout_in_minutes: 120
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -113,14 +131,15 @@ steps:
|
||||
- pytest -v -s models/language/pooling -m 'not core_model'
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 100
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 95
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: Language Models Test (MTEB)
|
||||
key: language-models-test-mteb
|
||||
timeout_in_minutes: 45
|
||||
timeout_in_minutes: 68
|
||||
device: h200_18gb
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -4,7 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: "Multi-Modal Models (Standard) 1: qwen2"
|
||||
key: multi-modal-models-standard-1-qwen2
|
||||
timeout_in_minutes: 45
|
||||
timeout_in_minutes: 68
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -14,13 +14,15 @@ steps:
|
||||
- pytest -v -s models/multimodal/generation/test_ultravox.py -m core_model
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 65
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: "Multi-Modal Models (Standard) 2: qwen3 + gemma"
|
||||
key: multi-modal-models-standard-2-qwen3-gemma
|
||||
timeout_in_minutes: 50
|
||||
timeout_in_minutes: 75
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -31,7 +33,9 @@ steps:
|
||||
- pytest -v -s models/multimodal/generation/test_qwen2_5_vl.py -m core_model
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
@@ -47,14 +51,15 @@ steps:
|
||||
- pytest -v -s models/multimodal/generation/test_qwen2_vl.py -m core_model
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
device: mi250_1
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: "Multi-Modal Models (Standard) 4: other + whisper"
|
||||
device: h200_35gb
|
||||
key: multi-modal-models-standard-4-other-whisper
|
||||
timeout_in_minutes: 50
|
||||
timeout_in_minutes: 75
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/models/multimodal
|
||||
@@ -65,7 +70,9 @@ steps:
|
||||
- cd .. && VLLM_WORKER_MULTIPROC_METHOD=spawn pytest -v -s tests/models/multimodal/generation/test_whisper.py -m core_model # Otherwise, mp_method="spawn" doesn't work
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 50
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
@@ -85,7 +92,7 @@ steps:
|
||||
|
||||
- label: Multi-Modal Processor # 44min
|
||||
key: multi-modal-processor
|
||||
timeout_in_minutes: 65
|
||||
timeout_in_minutes: 98
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -107,7 +114,9 @@ steps:
|
||||
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-mm-small.txt --tp-size=1
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 35
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -118,6 +127,7 @@ steps:
|
||||
- vllm/model_executor/model_loader/
|
||||
|
||||
- label: Multi-Modal Models (Extended Generation 1)
|
||||
device: h200_35gb
|
||||
key: multi-modal-models-extended-generation-1
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
@@ -129,7 +139,9 @@ steps:
|
||||
- pytest -v -s models/multimodal/test_mapping.py
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 90
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
@@ -164,8 +176,9 @@ steps:
|
||||
- pytest -v -s models/multimodal/pooling -m 'not core_model'
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 75
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 60
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -5,7 +5,7 @@ steps:
|
||||
- label: PyTorch Compilation Unit Tests
|
||||
device: h200_35gb
|
||||
key: pytorch-compilation-unit-tests
|
||||
timeout_in_minutes: 90
|
||||
timeout_in_minutes: 150
|
||||
source_file_dependencies:
|
||||
- vllm/__init__.py
|
||||
- vllm/_aiter_ops.py
|
||||
@@ -107,16 +107,11 @@ steps:
|
||||
- tests/compile/passes
|
||||
commands:
|
||||
- pytest -s -v compile/passes --ignore compile/passes/distributed
|
||||
mirror:
|
||||
amd:
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 65
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
|
||||
- label: PyTorch Fullgraph Smoke Test
|
||||
device: h200_35gb
|
||||
key: pytorch-fullgraph-smoke-test
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 90
|
||||
source_file_dependencies:
|
||||
- vllm/__init__.py
|
||||
- vllm/_aiter_ops.py
|
||||
@@ -148,7 +143,42 @@ steps:
|
||||
# as it is a heavy test that is covered in other steps.
|
||||
# Use `find` to launch multiple instances of pytest so that
|
||||
# they do not suffer from https://github.com/vllm-project/vllm/issues/28965
|
||||
- "find compile/fullgraph/ -name 'test_*.py' -not -name 'test_full_graph.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'"
|
||||
- "find compile/fullgraph/ -name 'test_*.py' -not -name 'test_full_cudagraph.py' -not -name 'test_full_graph.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'"
|
||||
|
||||
# Hopper-only DeepSeek-V2-Lite cases in this file require two 29.3-GiB model
|
||||
# instances and cannot fit a 35GB MIG slice. L4 retains the original coverage:
|
||||
# those SM90 cases skip while the architecture-compatible cases still run.
|
||||
- label: PyTorch Fullgraph CUDAGraph (L4 Compatibility)
|
||||
key: pytorch-fullgraph-cudagraph-l4-compatibility
|
||||
timeout_in_minutes: 60
|
||||
source_file_dependencies:
|
||||
- vllm/__init__.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/_custom_ops.py
|
||||
- vllm/compilation/
|
||||
- vllm/config/
|
||||
- vllm/distributed/
|
||||
- vllm/engine/
|
||||
- vllm/env_override.py
|
||||
- vllm/envs.py
|
||||
- vllm/forward_context.py
|
||||
- vllm/inputs/
|
||||
- vllm/ir/
|
||||
- vllm/kernels/
|
||||
- vllm/logger.py
|
||||
- vllm/model_executor/
|
||||
- vllm/multimodal/
|
||||
- vllm/platforms/
|
||||
- vllm/plugins/
|
||||
- vllm/sampling_params.py
|
||||
- vllm/sequence.py
|
||||
- vllm/transformers_utils/
|
||||
- vllm/triton_utils/
|
||||
- vllm/utils/
|
||||
- vllm/v1/
|
||||
- tests/compile
|
||||
commands:
|
||||
- pytest -s -v compile/fullgraph/test_full_cudagraph.py
|
||||
|
||||
- label: PyTorch Fullgraph
|
||||
key: pytorch-fullgraph
|
||||
@@ -197,7 +227,9 @@ steps:
|
||||
- bash standalone_tests/pytorch_nightly_dependency.sh
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 30
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -3,8 +3,11 @@ depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Quantization
|
||||
device: h200_35gb
|
||||
key: quantization
|
||||
timeout_in_minutes: 60
|
||||
timeout_in_minutes: 75
|
||||
env:
|
||||
VLLM_USE_V2_MODEL_RUNNER: "0"
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
- vllm/model_executor/layers/quantization
|
||||
@@ -19,9 +22,12 @@ steps:
|
||||
# TODO(jerryzh168): resolve the above comment
|
||||
- uv pip install --system torchao==0.17.0 --index-url https://download.pytorch.org/whl/cu130
|
||||
- uv pip install --system conch-triton-kernels
|
||||
- VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py
|
||||
# The SM90-only checkpoint currently contains a removed weight_chan_scale
|
||||
# parameter. It was not exercised by the previous L4 job.
|
||||
- VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py -k 'not test_compressed_tensors_w4a8_fp8'
|
||||
|
||||
- label: Quantized Fusions
|
||||
device: h200_35gb
|
||||
key: quantized-fusions
|
||||
timeout_in_minutes: 20
|
||||
source_file_dependencies:
|
||||
@@ -52,8 +58,11 @@ steps:
|
||||
- pytest -s -v tests/quantization/test_blackwell_moe.py
|
||||
|
||||
- label: Quantized Models Test
|
||||
device: h200_35gb
|
||||
key: quantized-models-test
|
||||
timeout_in_minutes: 50
|
||||
timeout_in_minutes: 65
|
||||
env:
|
||||
VLLM_USE_V2_MODEL_RUNNER: "0"
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/layers/quantization
|
||||
- tests/models/quantization
|
||||
|
||||
@@ -81,6 +81,7 @@ steps:
|
||||
- pytest -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine
|
||||
|
||||
- label: Rust Frontend Tool Use
|
||||
device: h200_35gb
|
||||
timeout_in_minutes: 25
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -19,8 +19,18 @@ steps:
|
||||
- VLLM_USE_FLASHINFER_SAMPLER=1 pytest -v -s samplers
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
device: mi250_1
|
||||
timeout_in_minutes: 40
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
- vllm/model_executor/layers
|
||||
- vllm/sampling_metadata.py
|
||||
- vllm/v1/sample/
|
||||
- vllm/entrypoints/generate/beam_search/
|
||||
- tests/samplers
|
||||
- tests/conftest.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s samplers
|
||||
|
||||
@@ -14,8 +14,9 @@ steps:
|
||||
- pytest -v -s v1/e2e/spec_decode -k "eagle_correctness"
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 60
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -53,8 +54,9 @@ steps:
|
||||
- pytest -v -s v1/e2e/spec_decode -k "speculators or mtp_correctness"
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 65
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 75
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -92,10 +94,9 @@ steps:
|
||||
- pytest -v -s v1/e2e/spec_decode -k "ngram or suffix"
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
timeout_in_minutes: 55
|
||||
# TODO(akaratza): Test after Torch >= 2.12 bump
|
||||
soft_fail: true
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 35
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
source_file_dependencies:
|
||||
@@ -119,7 +120,8 @@ steps:
|
||||
- pytest -v -s v1/e2e/spec_decode -k "draft_model or no_sync or batch_inference"
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
dind: false
|
||||
device: mi300_1
|
||||
timeout_in_minutes: 55
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
@@ -170,3 +172,19 @@ steps:
|
||||
- tests/v1/e2e/spec_decode/
|
||||
commands:
|
||||
- pytest -v -s v1/e2e/spec_decode -k "qwen3_5-hybrid"
|
||||
|
||||
- label: Spec Decode DeepSeek MTP Parallel Load (B200)
|
||||
key: spec-decode-deepseek-mtp-parallel-load-b200
|
||||
timeout_in_minutes: 30
|
||||
device: b200-k8s
|
||||
optional: true
|
||||
num_devices: 2
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/llm_base_proposer.py
|
||||
- vllm/v1/spec_decode/eagle.py
|
||||
- vllm/v1/worker/gpu/spec_decode/eagle/
|
||||
- vllm/model_executor/models/deepseek_mtp.py
|
||||
- vllm/model_executor/models/deepseek_v2.py
|
||||
- tests/v1/e2e/spec_decode/test_mtp_parallel_load.py
|
||||
commands:
|
||||
- pytest -v -s v1/e2e/spec_decode/test_mtp_parallel_load.py
|
||||
|
||||
@@ -15,7 +15,9 @@ steps:
|
||||
- bash weight_loading/run_model_weight_loading_test.sh -c weight_loading/models.txt
|
||||
mirror:
|
||||
amd:
|
||||
dind: false
|
||||
device: mi300_2
|
||||
timeout_in_minutes: 35
|
||||
depends_on:
|
||||
- image-build-amd
|
||||
commands:
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
dist
|
||||
vllm/*.so
|
||||
vllm/vllm-rs
|
||||
.git
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
|
||||
+2
-1
@@ -47,6 +47,7 @@
|
||||
|
||||
# Rust Frontend
|
||||
/rust/ @BugenZhao @njhill
|
||||
/rust/src/bench @esmeetu
|
||||
/build_rust.sh @BugenZhao @njhill
|
||||
/rust-toolchain.toml @BugenZhao @njhill
|
||||
/.buildkite/test_areas/rust* @BugenZhao @njhill
|
||||
@@ -172,7 +173,7 @@ mkdocs.yaml @hmellor
|
||||
# Kernels
|
||||
/vllm/v1/attention/ops/chunked_prefill_paged_decode.py @tdoublep
|
||||
/vllm/v1/attention/ops/triton_unified_attention.py @tdoublep
|
||||
/vllm/model_executor/layers/fla @ZJY0516 @vadiklyutiy
|
||||
/vllm/third_party/flash_linear_attention @ZJY0516 @vadiklyutiy
|
||||
|
||||
# ROCm related: specify owner with write access to notify AMD folks for careful code review
|
||||
/vllm/**/*rocm* @tjtanaa @dllehr-amd
|
||||
|
||||
@@ -181,6 +181,18 @@ pull_request_rules:
|
||||
add:
|
||||
- performance
|
||||
|
||||
- name: label-quantization
|
||||
description: Automatically apply quantization label
|
||||
conditions:
|
||||
- label != stale
|
||||
- or:
|
||||
- files~=^vllm/model_executor/layers/quantization/
|
||||
- title~=(?i)quant
|
||||
actions:
|
||||
label:
|
||||
add:
|
||||
- quantization
|
||||
|
||||
- name: label-qwen
|
||||
description: Automatically apply qwen label
|
||||
conditions:
|
||||
|
||||
@@ -130,6 +130,47 @@ jobs:
|
||||
},
|
||||
],
|
||||
},
|
||||
quantization: {
|
||||
keywords: [
|
||||
{
|
||||
term: "quantization",
|
||||
searchIn: "both"
|
||||
},
|
||||
{
|
||||
term: "quantized",
|
||||
searchIn: "both"
|
||||
},
|
||||
],
|
||||
},
|
||||
"intel-gpu": {
|
||||
// Keyword search - matches whole words only (with word boundaries)
|
||||
keywords: [
|
||||
{
|
||||
term: "B50",
|
||||
searchIn: "both"
|
||||
},
|
||||
{
|
||||
term: "B60",
|
||||
searchIn: "both"
|
||||
},
|
||||
{
|
||||
term: "B70",
|
||||
searchIn: "both"
|
||||
},
|
||||
{
|
||||
term: "intel gpu",
|
||||
searchIn: "both"
|
||||
},
|
||||
{
|
||||
term: "Arc GPU",
|
||||
searchIn: "both"
|
||||
},
|
||||
{
|
||||
term: "BMG",
|
||||
searchIn: "both"
|
||||
},
|
||||
],
|
||||
},
|
||||
// Add more label configurations here as needed
|
||||
// example: {
|
||||
// keywords: [...],
|
||||
@@ -323,7 +364,7 @@ jobs:
|
||||
// {users} will be replaced with @mentions
|
||||
const ccConfig = {
|
||||
rocm: {
|
||||
users: ['hongxiayang', 'tjtanaa', 'vllmellm'],
|
||||
users: ['hongxiayang', 'tjtanaa', 'vllmellm', 'giuseppegrossi'],
|
||||
message: 'CC {users} for ROCm-related issue',
|
||||
},
|
||||
mistral: {
|
||||
@@ -491,4 +532,4 @@ jobs:
|
||||
issue_number: context.issue.number,
|
||||
body: message,
|
||||
});
|
||||
core.notice(`Requested missing ROCm info from @${author}: ${missing.map(m => m.name).join(', ')}`);
|
||||
core.notice(`Requested missing ROCm info from @${author}: ${missing.map(m => m.name).join(', ')}`);
|
||||
|
||||
@@ -18,6 +18,9 @@ vllm/third_party/deep_gemm/
|
||||
# fmha_sm100 vendored package built from source
|
||||
vllm/third_party/fmha_sm100/
|
||||
|
||||
# tml-fa4 vendored package built from source
|
||||
vllm/third_party/tml_fa4/
|
||||
|
||||
# triton jit
|
||||
.triton
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ default_install_hook_types:
|
||||
default_stages:
|
||||
- pre-commit # Run locally
|
||||
- manual # Run in CI
|
||||
exclude: 'vllm/third_party/.*'
|
||||
exclude: 'vllm/third_party/.*|vllm/models/kimi_k3/nvidia/ops/third_party/.*|vllm/models/kimi_k3/amd/ops/third_party/.*'
|
||||
repos:
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.14.0
|
||||
@@ -30,7 +30,7 @@ repos:
|
||||
- id: markdownlint-cli2
|
||||
language_version: lts
|
||||
args: [--fix]
|
||||
exclude: ^CLAUDE\.md$
|
||||
exclude: (^|/)CLAUDE\.md$
|
||||
- repo: https://github.com/rhysd/actionlint
|
||||
rev: v1.7.7
|
||||
hooks:
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
collect_env.py
|
||||
vllm/model_executor/layers/fla/ops/*.py
|
||||
+64
-10
@@ -68,8 +68,8 @@ endif()
|
||||
# requirements.txt files and should be kept consistent. The ROCm torch
|
||||
# versions are derived from docker/Dockerfile.rocm
|
||||
#
|
||||
set(TORCH_SUPPORTED_VERSION_CUDA "2.11.0")
|
||||
set(TORCH_SUPPORTED_VERSION_ROCM "2.11.0")
|
||||
set(TORCH_SUPPORTED_VERSION_CUDA "2.13.0")
|
||||
set(TORCH_SUPPORTED_VERSION_ROCM "2.13.0")
|
||||
# TORCH_NIGHTLY=1 builds run against unpinned nightly wheels, so the supported-
|
||||
# version check would always warn. Only treat it as a nightly build when the
|
||||
# value is exactly "1" (the bootstrap exports TORCH_NIGHTLY=0 by default, which
|
||||
@@ -114,6 +114,11 @@ find_package(Torch REQUIRED)
|
||||
# Supported NVIDIA architectures.
|
||||
# This check must happen after find_package(Torch) because that's when CMAKE_CUDA_COMPILER_VERSION gets defined
|
||||
if(DEFINED CMAKE_CUDA_COMPILER_VERSION AND
|
||||
CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 13.4)
|
||||
# Rubin (10.7) can run SM100 family code, but CUDA 13.4 also supports
|
||||
# targeting it directly.
|
||||
set(CUDA_SUPPORTED_ARCHS "7.5;8.0;8.6;8.7;8.9;9.0;10.0;10.7;11.0;12.0")
|
||||
elseif(DEFINED CMAKE_CUDA_COMPILER_VERSION AND
|
||||
CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 13.0)
|
||||
# starting from CUDA 12.9 and Blackwell (10.0), we use family-specific targets (10.0f, 12.0f, etc)
|
||||
# to support the whole generation without specifying all sub-architectures
|
||||
@@ -411,8 +416,11 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
"csrc/libtorch_stable/mamba/selective_scan_fwd.cu"
|
||||
"csrc/libtorch_stable/cache_kernels.cu"
|
||||
"csrc/libtorch_stable/cache_kernels_fused.cu"
|
||||
"csrc/libtorch_stable/custom_all_gather_reduce_scatter.cu"
|
||||
"csrc/libtorch_stable/custom_all_gather_reduce_scatter_ops.cpp"
|
||||
"csrc/libtorch_stable/custom_all_reduce.cu"
|
||||
"csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu")
|
||||
"csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu"
|
||||
"csrc/libtorch_stable/fused_kimi_k3_mla_key_concat_kv_cache_kernel.cu")
|
||||
|
||||
if(VLLM_GPU_LANG STREQUAL "CUDA" AND
|
||||
DEFINED CMAKE_CUDA_COMPILER_VERSION AND
|
||||
@@ -420,7 +428,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(COOPERATIVE_TOPK_ARCHS
|
||||
"9.0a;10.0f;10.1f;10.3f;11.0f;12.0f;12.1f" "${CUDA_ARCHS}")
|
||||
"9.0a;10.0f;10.1f;10.3f;10.7f;11.0f;12.0f;12.1f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(COOPERATIVE_TOPK_ARCHS
|
||||
"9.0a;10.0a;10.1a;10.3a;12.0a;12.1a" "${CUDA_ARCHS}")
|
||||
@@ -695,7 +703,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
|
||||
# DeepSeek V3 fused A GEMM kernel (requires SM 9.0+, Hopper and later)
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(DSV3_FUSED_A_GEMM_ARCHS "9.0a;10.0f;11.0f;12.0f" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(DSV3_FUSED_A_GEMM_ARCHS "9.0a;10.0f;10.7f;11.0f;12.0f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(DSV3_FUSED_A_GEMM_ARCHS "9.0a;10.0a;10.1a;10.3a;12.0a;12.1a" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
@@ -815,7 +823,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
# The cutlass_scaled_mm kernels for Blackwell SM100 (c3x, i.e. CUTLASS 3.x)
|
||||
# require CUDA 12.8 or later
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(SCALED_MM_ARCHS "10.0f;11.0f" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(SCALED_MM_ARCHS "10.0f;10.7f;11.0f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(SCALED_MM_ARCHS "10.0a;10.1a;10.3a" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
@@ -899,7 +907,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
endif()
|
||||
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(SCALED_MM_ARCHS "10.0f;11.0f" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(SCALED_MM_ARCHS "10.0f;10.7f;11.0f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(SCALED_MM_ARCHS "10.0a;10.1a;10.3a" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
@@ -924,7 +932,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
|
||||
# moe_data.cu is used by all CUTLASS MoE kernels.
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(CUTLASS_MOE_DATA_ARCHS "9.0a;10.0f;11.0f;12.0f" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(CUTLASS_MOE_DATA_ARCHS "9.0a;10.0f;10.7f;11.0f;12.0f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(CUTLASS_MOE_DATA_ARCHS "9.0a;10.0a;10.1a;10.3a;12.0a;12.1a" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
@@ -981,7 +989,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
# SM10x/11x FP4 kernels. MXFP4 experts quantization is currently compiled
|
||||
# only in this block; SM12x has separate NVFP4 matmul/MoE kernels above.
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(FP4_SM100_ARCHS "10.0f;11.0f" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(FP4_SM100_ARCHS "10.0f;10.7f;11.0f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(FP4_SM100_ARCHS "10.0a;10.1a;10.3a" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
@@ -1047,7 +1055,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
# Runtime dispatch is gated in
|
||||
# vllm/v1/attention/backends/mla/cutlass_mla.py.
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(MLA_ARCHS "10.0f;11.0f" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(MLA_ARCHS "10.0f;10.7f;11.0f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(MLA_ARCHS "10.0a;10.1a;10.3a" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
@@ -1069,6 +1077,41 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
set(MLA_ARCHS)
|
||||
endif()
|
||||
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(FUSED_KDA_DECODE_ARCHS
|
||||
"9.0a;10.0f;12.0f" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
if(FUSED_KDA_DECODE_ARCHS)
|
||||
set(FUSED_KDA_DECODE_SRC
|
||||
"csrc/libtorch_stable/kimi_k3/fused_kda_decode_kernel.cu")
|
||||
set_gencode_flags_for_srcs(
|
||||
SRCS "${FUSED_KDA_DECODE_SRC}"
|
||||
CUDA_ARCHS "${FUSED_KDA_DECODE_ARCHS}")
|
||||
set_property(SOURCE ${FUSED_KDA_DECODE_SRC} APPEND PROPERTY
|
||||
COMPILE_OPTIONS "$<$<COMPILE_LANGUAGE:CUDA>:--use_fast_math>")
|
||||
list(APPEND VLLM_STABLE_EXT_SRC "${FUSED_KDA_DECODE_SRC}")
|
||||
message(STATUS
|
||||
"Building fused KDA decode for archs: ${FUSED_KDA_DECODE_ARCHS}")
|
||||
endif()
|
||||
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(KIMI_K3_ATTN_RES_ARCHS
|
||||
"10.0f" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
if(KIMI_K3_ATTN_RES_ARCHS)
|
||||
set(KIMI_K3_ATTN_RES_SRC
|
||||
"csrc/libtorch_stable/kimi_k3/attn_res_kernel.cu")
|
||||
set_gencode_flags_for_srcs(
|
||||
SRCS "${KIMI_K3_ATTN_RES_SRC}"
|
||||
CUDA_ARCHS "${KIMI_K3_ATTN_RES_ARCHS}")
|
||||
set_property(SOURCE ${KIMI_K3_ATTN_RES_SRC} APPEND PROPERTY
|
||||
COMPILE_OPTIONS
|
||||
"$<$<COMPILE_LANGUAGE:CUDA>:--expt-relaxed-constexpr;--expt-extended-lambda;--use_fast_math>")
|
||||
list(APPEND VLLM_STABLE_EXT_SRC "${KIMI_K3_ATTN_RES_SRC}")
|
||||
message(STATUS
|
||||
"Building Kimi K3 AttnRes for archs: ${KIMI_K3_ATTN_RES_ARCHS}")
|
||||
endif()
|
||||
|
||||
# Hadacore kernels
|
||||
cuda_archs_loose_intersection(HADACORE_ARCHS "8.0+PTX;9.0+PTX" "${CUDA_ARCHS}")
|
||||
if(HADACORE_ARCHS)
|
||||
@@ -1110,6 +1153,14 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
|
||||
target_compile_definitions(_C_stable_libtorch PRIVATE
|
||||
VLLM_ENABLE_COOPERATIVE_TOPK=1)
|
||||
endif()
|
||||
if(FUSED_KDA_DECODE_ARCHS)
|
||||
target_compile_definitions(_C_stable_libtorch PRIVATE
|
||||
VLLM_ENABLE_FUSED_KDA_DECODE=1)
|
||||
endif()
|
||||
if(KIMI_K3_ATTN_RES_ARCHS)
|
||||
target_compile_definitions(_C_stable_libtorch PRIVATE
|
||||
VLLM_ENABLE_KIMI_K3_ATTN_RES=1)
|
||||
endif()
|
||||
# Needed by CUTLASS kernels
|
||||
target_compile_definitions(_C_stable_libtorch PRIVATE
|
||||
CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1)
|
||||
@@ -1364,6 +1415,7 @@ if(VLLM_GPU_LANG STREQUAL "HIP")
|
||||
set(VLLM_ROCM_EXT_SRC
|
||||
"csrc/rocm/torch_bindings.cpp"
|
||||
"csrc/rocm/skinny_gemms.cu"
|
||||
"csrc/rocm/skinny_gemms_int4.cu"
|
||||
"csrc/rocm/attention.cu")
|
||||
|
||||
set(VLLM_ROCM_HAS_GFX1100 OFF)
|
||||
@@ -1406,7 +1458,9 @@ if (VLLM_GPU_LANG STREQUAL "CUDA")
|
||||
include(cmake/external_projects/deepgemm.cmake)
|
||||
include(cmake/external_projects/fmha_sm100.cmake)
|
||||
include(cmake/external_projects/flashmla.cmake)
|
||||
include(cmake/external_projects/flashkda.cmake)
|
||||
include(cmake/external_projects/qutlass.cmake)
|
||||
include(cmake/external_projects/tml_fa4.cmake)
|
||||
|
||||
# vllm-flash-attn should be last as it overwrites some CMake functions
|
||||
include(cmake/external_projects/vllm_flash_attn.cmake)
|
||||
|
||||
@@ -48,7 +48,7 @@ vLLM is flexible and easy to use with:
|
||||
- Tool calling and reasoning parsers
|
||||
- OpenAI-compatible API server, plus Anthropic Messages API and gRPC support
|
||||
- Efficient multi-LoRA support for dense and MoE layers
|
||||
- Support for NVIDIA GPUs, AMD GPUs, and x86/ARM/PowerPC CPUs. Additionally, diverse hardware plugins such as Google TPUs, Intel Gaudi, IBM Spyre, Huawei Ascend, Rebellions NPU, Apple Silicon, MetaX GPU, and more.
|
||||
- Support for NVIDIA GPUs, AMD GPUs, Intel GPUs, and x86/ARM/PowerPC CPUs. Additionally, diverse hardware plugins such as Google TPUs, Intel Gaudi, IBM Spyre, Huawei Ascend, Rebellions NPU, Apple Silicon, MetaX GPU, and more.
|
||||
|
||||
vLLM seamlessly supports 200+ model architectures on Hugging Face, including:
|
||||
|
||||
|
||||
@@ -75,7 +75,11 @@ def run_mla_benchmark(config: BenchmarkConfig, **kwargs) -> BenchmarkResult:
|
||||
from mla_runner import run_mla_benchmark as run_mla
|
||||
|
||||
return run_mla(
|
||||
config.backend, config, prefill_backend=config.prefill_backend, **kwargs
|
||||
config.backend,
|
||||
config,
|
||||
prefill_backend=config.prefill_backend,
|
||||
sparse_mla_force_mqa=config.sparse_mla_force_mqa,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@@ -592,6 +596,30 @@ def main():
|
||||
default="profile",
|
||||
help="Output file name for ncu profile (default: 'profile').",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--torch-profile",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Collect a PyTorch profiler Chrome trace for each benchmark run.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--torch-profile-dir",
|
||||
default=None,
|
||||
help="Directory for PyTorch profiler traces.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--torch-profile-iters",
|
||||
type=int,
|
||||
default=3,
|
||||
help="Number of forward passes to record per PyTorch profiler trace.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sparse-mla-mha-variants",
|
||||
nargs="+",
|
||||
default=None,
|
||||
choices=["dense_mha", "mqa"],
|
||||
help="Sparse MLA variants to run in mha_vs_mqa mode. Defaults to both.",
|
||||
)
|
||||
|
||||
# Parameter sweep (use YAML config for advanced sweeps)
|
||||
parser.add_argument(
|
||||
@@ -641,6 +669,7 @@ def main():
|
||||
|
||||
# Prefill backends (e.g., ["fa3", "fa4"])
|
||||
args.prefill_backends = yaml_config.get("prefill_backends", None)
|
||||
args.prefill_backend = yaml_config.get("prefill_backend", None)
|
||||
|
||||
# FP8 output benchmark knobs; CLI wins.
|
||||
if args.fp8_output_scale is None:
|
||||
@@ -683,6 +712,9 @@ def main():
|
||||
args.num_q_heads = model.get("num_q_heads", args.num_q_heads)
|
||||
args.num_kv_heads = model.get("num_kv_heads", args.num_kv_heads)
|
||||
args.block_size = model.get("block_size", args.block_size)
|
||||
args.max_model_len = model.get(
|
||||
"max_model_len", getattr(args, "max_model_len", None)
|
||||
)
|
||||
# MLA-specific dimensions
|
||||
args.kv_lora_rank = model.get("kv_lora_rank", args.kv_lora_rank)
|
||||
args.qk_nope_head_dim = model.get("qk_nope_head_dim", args.qk_nope_head_dim)
|
||||
@@ -701,6 +733,21 @@ def main():
|
||||
args.cuda_graphs = yaml_config["cuda_graphs"]
|
||||
if "ncu_profile" in yaml_config:
|
||||
args.ncu_profile = yaml_config["ncu_profile"]
|
||||
if "torch_profile" in yaml_config:
|
||||
args.torch_profile = yaml_config["torch_profile"]
|
||||
if "torch_profile_dir" in yaml_config:
|
||||
args.torch_profile_dir = yaml_config["torch_profile_dir"]
|
||||
if "torch_profile_iters" in yaml_config:
|
||||
args.torch_profile_iters = yaml_config["torch_profile_iters"]
|
||||
args.sparse_mla_topk_pattern = yaml_config.get(
|
||||
"sparse_mla_topk_pattern", "random"
|
||||
)
|
||||
args.sparse_mla_dense_mha_max_seq_len = yaml_config.get(
|
||||
"sparse_mla_dense_mha_max_seq_len", None
|
||||
)
|
||||
args.sparse_mla_mha_variants = yaml_config.get(
|
||||
"sparse_mla_mha_variants", args.sparse_mla_mha_variants
|
||||
)
|
||||
|
||||
# Parameter sweep configuration
|
||||
if "parameter_sweep" in yaml_config:
|
||||
@@ -842,8 +889,6 @@ def main():
|
||||
num_kv_heads=args.num_kv_heads,
|
||||
block_size=args.block_size,
|
||||
device=args.device,
|
||||
repeats=args.repeats,
|
||||
warmup_iters=args.warmup_iters,
|
||||
profile_memory=args.profile_memory,
|
||||
kv_cache_dtype=args.kv_cache_dtype,
|
||||
use_cuda_graphs=args.cuda_graphs,
|
||||
@@ -1063,6 +1108,133 @@ def main():
|
||||
f"\n [yellow]Prefill always faster for batch_size={bs}[/]"
|
||||
)
|
||||
|
||||
# Handle MHA vs MQA comparison mode for sparse MLA
|
||||
elif hasattr(args, "mode") and args.mode == "mha_vs_mqa":
|
||||
console.print("[yellow]Mode: MHA vs MQA comparison for sparse MLA[/]")
|
||||
|
||||
sparse_mla_topk_pattern = getattr(args, "sparse_mla_topk_pattern", "random")
|
||||
dense_mha_max_seq_len = getattr(args, "sparse_mla_dense_mha_max_seq_len", None)
|
||||
prefill_backend = getattr(args, "prefill_backend", None)
|
||||
if prefill_backend:
|
||||
console.print(f"Prefill backend: {prefill_backend}")
|
||||
available_variants = [
|
||||
("dense_mha", False, "dense"),
|
||||
("mqa", True, "auto"),
|
||||
]
|
||||
requested_variants = getattr(args, "sparse_mla_mha_variants", None)
|
||||
if requested_variants is not None:
|
||||
valid_variants = {label for label, _, _ in available_variants}
|
||||
invalid_variants = sorted(set(requested_variants) - valid_variants)
|
||||
if invalid_variants:
|
||||
raise ValueError(
|
||||
"Invalid sparse_mla_mha_variants entries: "
|
||||
f"{invalid_variants}. Valid variants are: "
|
||||
f"{sorted(valid_variants)}"
|
||||
)
|
||||
requested_variant_set = set(requested_variants)
|
||||
variants = [
|
||||
variant
|
||||
for variant in available_variants
|
||||
if variant[0] in requested_variant_set
|
||||
]
|
||||
else:
|
||||
variants = available_variants
|
||||
formatter = ResultsFormatter(console)
|
||||
total = 0
|
||||
for spec in args.batch_specs:
|
||||
q_len = max(request.q_len for request in parse_batch_spec(spec))
|
||||
for variant_label, _, _ in variants:
|
||||
if (
|
||||
variant_label == "dense_mha"
|
||||
and dense_mha_max_seq_len is not None
|
||||
and q_len > dense_mha_max_seq_len
|
||||
):
|
||||
continue
|
||||
total += len(backends)
|
||||
|
||||
with tqdm(total=total, desc="Benchmarking") as pbar:
|
||||
for spec in args.batch_specs:
|
||||
q_len = max(request.q_len for request in parse_batch_spec(spec))
|
||||
for backend in backends:
|
||||
for variant_label, force_mqa, mha_mode in variants:
|
||||
if (
|
||||
variant_label == "dense_mha"
|
||||
and dense_mha_max_seq_len is not None
|
||||
and q_len > dense_mha_max_seq_len
|
||||
):
|
||||
continue
|
||||
config = BenchmarkConfig(
|
||||
backend=f"{backend}_{variant_label}",
|
||||
batch_spec=spec,
|
||||
num_layers=args.num_layers,
|
||||
head_dim=args.head_dim,
|
||||
num_q_heads=args.num_q_heads,
|
||||
num_kv_heads=args.num_kv_heads,
|
||||
block_size=args.block_size,
|
||||
device=args.device,
|
||||
max_model_len=getattr(args, "max_model_len", None),
|
||||
kv_cache_dtype=args.kv_cache_dtype,
|
||||
profile_memory=args.profile_memory,
|
||||
use_cuda_graphs=args.cuda_graphs,
|
||||
ncu_profile=args.ncu_profile,
|
||||
torch_profile=args.torch_profile,
|
||||
torch_profile_dir=args.torch_profile_dir,
|
||||
torch_profile_iters=args.torch_profile_iters,
|
||||
warmup_ms=args.warmup_ms,
|
||||
kv_lora_rank=getattr(args, "kv_lora_rank", None),
|
||||
qk_nope_head_dim=getattr(args, "qk_nope_head_dim", None),
|
||||
qk_rope_head_dim=getattr(args, "qk_rope_head_dim", None),
|
||||
v_head_dim=getattr(args, "v_head_dim", None),
|
||||
sparse_mla_force_mqa=force_mqa,
|
||||
sparse_mla_mha_mode=mha_mode,
|
||||
sparse_mla_dense_mha_max_seq_len=dense_mha_max_seq_len,
|
||||
sparse_mla_topk_pattern=sparse_mla_topk_pattern,
|
||||
prefill_backend=prefill_backend,
|
||||
)
|
||||
|
||||
# run_mla_benchmark needs the real backend name
|
||||
from mla_runner import run_mla_benchmark as run_mla
|
||||
|
||||
run_label = f"{backend}_{variant_label} {spec}"
|
||||
pbar.set_postfix_str(run_label)
|
||||
|
||||
try:
|
||||
result = run_mla(
|
||||
backend,
|
||||
config,
|
||||
prefill_backend=prefill_backend,
|
||||
sparse_mla_force_mqa=force_mqa,
|
||||
)
|
||||
except Exception as e:
|
||||
result = BenchmarkResult(
|
||||
config=config,
|
||||
mean_time=float("inf"),
|
||||
median_time=float("inf"),
|
||||
std_time=0,
|
||||
min_time=float("inf"),
|
||||
max_time=float("inf"),
|
||||
error=str(e),
|
||||
)
|
||||
|
||||
all_results.append(result)
|
||||
if args.output_csv:
|
||||
formatter.save_csv(all_results, args.output_csv)
|
||||
if args.output_json:
|
||||
formatter.save_json(all_results, args.output_json)
|
||||
|
||||
if not result.success:
|
||||
console.print(
|
||||
f"[red]Error {backend}_{variant_label} "
|
||||
f"{spec}: {result.error}[/]"
|
||||
)
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
# Display results with variant labels as separate "backends"
|
||||
console.print("\n[bold green]MHA vs MQA Results:[/]")
|
||||
variant_backends = [f"{b}_{v}" for b in backends for v, _, _ in variants]
|
||||
formatter.print_table(all_results, variant_backends)
|
||||
|
||||
# Handle model parameter sweep mode
|
||||
elif hasattr(args, "model_parameter_sweep") and args.model_parameter_sweep:
|
||||
# Model parameter sweep
|
||||
|
||||
@@ -4,8 +4,10 @@
|
||||
"""Common utilities for attention benchmarking."""
|
||||
|
||||
import csv
|
||||
import gc
|
||||
import json
|
||||
import math
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import asdict, dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
@@ -44,10 +46,13 @@ def run_do_bench(
|
||||
kwargs: dict[str, Any] = {"return_mode": "all"}
|
||||
if use_cuda_graphs:
|
||||
result = triton.testing.do_bench_cudagraph(benchmark_fn, **kwargs)
|
||||
gc.collect()
|
||||
torch.accelerator.empty_cache()
|
||||
else:
|
||||
if warmup_ms is not None:
|
||||
kwargs["warmup"] = warmup_ms
|
||||
result = triton.testing.do_bench(benchmark_fn, **kwargs)
|
||||
torch.accelerator.synchronize()
|
||||
return result
|
||||
|
||||
|
||||
@@ -91,42 +96,6 @@ except ImportError:
|
||||
AttentionLayerBase = object # Fallback
|
||||
|
||||
|
||||
class MockKVBProj:
|
||||
"""Mock KV projection layer for MLA prefill mode.
|
||||
|
||||
Mimics ColumnParallelLinear behavior for kv_b_proj in MLA backends.
|
||||
Projects kv_c_normed to [qk_nope_head_dim + v_head_dim] per head.
|
||||
"""
|
||||
|
||||
def __init__(self, num_heads: int, qk_nope_head_dim: int, v_head_dim: int):
|
||||
self.num_heads = num_heads
|
||||
self.qk_nope_head_dim = qk_nope_head_dim
|
||||
self.v_head_dim = v_head_dim
|
||||
self.out_dim = qk_nope_head_dim + v_head_dim
|
||||
self.weight = torch.empty(0, dtype=torch.bfloat16)
|
||||
|
||||
def __call__(self, x: torch.Tensor) -> tuple[torch.Tensor]:
|
||||
"""
|
||||
Project kv_c_normed to output space.
|
||||
|
||||
Args:
|
||||
x: Input tensor [num_tokens, kv_lora_rank]
|
||||
|
||||
Returns:
|
||||
Tuple containing output tensor
|
||||
[num_tokens, num_heads, qk_nope_head_dim + v_head_dim]
|
||||
"""
|
||||
num_tokens = x.shape[0]
|
||||
result = torch.randn(
|
||||
num_tokens,
|
||||
self.num_heads,
|
||||
self.out_dim,
|
||||
device=x.device,
|
||||
dtype=x.dtype,
|
||||
)
|
||||
return (result,) # Return as tuple to match ColumnParallelLinear API
|
||||
|
||||
|
||||
class MockIndexer:
|
||||
"""Mock Indexer for sparse MLA backends.
|
||||
|
||||
@@ -158,6 +127,60 @@ class MockIndexer:
|
||||
)
|
||||
self.topk_indices_buffer[:num_tokens] = indices
|
||||
|
||||
def fill_indices(
|
||||
self,
|
||||
num_tokens: int,
|
||||
max_kv_len: int,
|
||||
pattern: str = "random",
|
||||
requests: Sequence[Any] | None = None,
|
||||
):
|
||||
if pattern == "random":
|
||||
self.fill_random_indices(num_tokens, max_kv_len)
|
||||
return
|
||||
if pattern == "prefix":
|
||||
indices = torch.arange(
|
||||
self.topk_tokens,
|
||||
dtype=torch.int32,
|
||||
device=self.topk_indices_buffer.device,
|
||||
)
|
||||
indices = (indices % max_kv_len).expand(num_tokens, -1)
|
||||
self.topk_indices_buffer[:num_tokens] = indices
|
||||
return
|
||||
if pattern == "sliding_window":
|
||||
if requests is None:
|
||||
start = max(max_kv_len - self.topk_tokens, 0)
|
||||
indices = torch.arange(
|
||||
start,
|
||||
start + self.topk_tokens,
|
||||
dtype=torch.int32,
|
||||
device=self.topk_indices_buffer.device,
|
||||
)
|
||||
indices = indices.clamp(max=max_kv_len - 1).expand(num_tokens, -1)
|
||||
self.topk_indices_buffer[:num_tokens] = indices
|
||||
return
|
||||
|
||||
rows = []
|
||||
offsets = torch.arange(
|
||||
self.topk_tokens,
|
||||
dtype=torch.int32,
|
||||
device=self.topk_indices_buffer.device,
|
||||
) - (self.topk_tokens - 1)
|
||||
for request in requests:
|
||||
q_len = request.q_len
|
||||
kv_len = request.kv_len
|
||||
context_len = kv_len - q_len
|
||||
positions = torch.arange(
|
||||
context_len,
|
||||
kv_len,
|
||||
dtype=torch.int32,
|
||||
device=self.topk_indices_buffer.device,
|
||||
)
|
||||
row_indices = positions[:, None] + offsets[None, :]
|
||||
rows.append(row_indices.clamp(min=0, max=kv_len - 1))
|
||||
self.topk_indices_buffer[:num_tokens] = torch.cat(rows, dim=0)
|
||||
return
|
||||
raise ValueError(f"Unknown sparse MLA topk pattern: {pattern}")
|
||||
|
||||
|
||||
class MockLayer(AttentionLayerBase):
|
||||
"""Mock attention layer with scale parameters and impl.
|
||||
@@ -252,10 +275,14 @@ class BenchmarkConfig:
|
||||
num_kv_heads: int
|
||||
block_size: int
|
||||
device: str
|
||||
max_model_len: int | None = None
|
||||
dtype: torch.dtype = torch.float16
|
||||
profile_memory: bool = False
|
||||
use_cuda_graphs: bool = False
|
||||
use_cuda_graphs: bool = True
|
||||
ncu_profile: bool = False
|
||||
torch_profile: bool = False
|
||||
torch_profile_dir: str | None = None
|
||||
torch_profile_iters: int = 3
|
||||
warmup_ms: int | None = None
|
||||
|
||||
# "auto" or "fp8"
|
||||
@@ -271,6 +298,10 @@ class BenchmarkConfig:
|
||||
# Backend-specific tuning
|
||||
num_kv_splits: int | None = None # CUTLASS MLA
|
||||
reorder_batch_threshold: int | None = None # FlashAttn MLA, FlashMLA
|
||||
sparse_mla_force_mqa: bool = False # Force MQA path for sparse MLA
|
||||
sparse_mla_mha_mode: str = "auto" # "auto" or "dense"
|
||||
sparse_mla_dense_mha_max_seq_len: int | None = None
|
||||
sparse_mla_topk_pattern: str = "random" # "random", "prefix", "sliding_window"
|
||||
num_splits: int | None = None # FlashAttention split-K (0=auto, 1=disabled)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,474 @@
|
||||
# Sparse MLA benchmark: forward_mha vs forward_mqa
|
||||
#
|
||||
# Usage:
|
||||
# python benchmark.py --config configs/mla_sparse_mha_vs_mqa.yaml
|
||||
#
|
||||
# Heatmap grid:
|
||||
# - batch_size: 1, 2, 4, 8, 16, 32
|
||||
# - seq_len: 32, 64, 128, 256, 512, 1024, 2048
|
||||
# - q_len: powers of two through seq_len
|
||||
#
|
||||
# Specs with q_len < seq_len include context; the q_len == seq_len diagonal
|
||||
# covers pure prefill.
|
||||
# The model shape below is the DP case. For the TP8 run, manually change
|
||||
# model.num_q_heads from 128 to 16 before rerunning this benchmark.
|
||||
|
||||
mode: mha_vs_mqa
|
||||
|
||||
model:
|
||||
name: "deepseek-v3"
|
||||
num_layers: 60
|
||||
num_q_heads: 128
|
||||
num_kv_heads: 1
|
||||
head_dim: 576
|
||||
kv_lora_rank: 512
|
||||
qk_nope_head_dim: 128
|
||||
qk_rope_head_dim: 64
|
||||
v_head_dim: 128
|
||||
block_size: 128
|
||||
max_model_len: 2048
|
||||
|
||||
batch_specs:
|
||||
# Batch size 1
|
||||
# seq_len = 32
|
||||
- "1q1s32"
|
||||
- "1q2s32"
|
||||
- "1q4s32"
|
||||
- "1q8s32"
|
||||
- "1q16s32"
|
||||
- "1q32"
|
||||
# seq_len = 64
|
||||
- "1q1s64"
|
||||
- "1q2s64"
|
||||
- "1q4s64"
|
||||
- "1q8s64"
|
||||
- "1q16s64"
|
||||
- "1q32s64"
|
||||
- "1q64"
|
||||
# seq_len = 128
|
||||
- "1q1s128"
|
||||
- "1q2s128"
|
||||
- "1q4s128"
|
||||
- "1q8s128"
|
||||
- "1q16s128"
|
||||
- "1q32s128"
|
||||
- "1q64s128"
|
||||
- "1q128"
|
||||
# seq_len = 256
|
||||
- "1q1s256"
|
||||
- "1q2s256"
|
||||
- "1q4s256"
|
||||
- "1q8s256"
|
||||
- "1q16s256"
|
||||
- "1q32s256"
|
||||
- "1q64s256"
|
||||
- "1q128s256"
|
||||
- "1q256"
|
||||
# seq_len = 512
|
||||
- "1q1s512"
|
||||
- "1q2s512"
|
||||
- "1q4s512"
|
||||
- "1q8s512"
|
||||
- "1q16s512"
|
||||
- "1q32s512"
|
||||
- "1q64s512"
|
||||
- "1q128s512"
|
||||
- "1q256s512"
|
||||
- "1q512"
|
||||
# seq_len = 1024
|
||||
- "1q1s1024"
|
||||
- "1q2s1024"
|
||||
- "1q4s1024"
|
||||
- "1q8s1024"
|
||||
- "1q16s1024"
|
||||
- "1q32s1024"
|
||||
- "1q64s1024"
|
||||
- "1q128s1024"
|
||||
- "1q256s1024"
|
||||
- "1q512s1024"
|
||||
- "1q1024"
|
||||
# seq_len = 2048
|
||||
- "1q1s2048"
|
||||
- "1q2s2048"
|
||||
- "1q4s2048"
|
||||
- "1q8s2048"
|
||||
- "1q16s2048"
|
||||
- "1q32s2048"
|
||||
- "1q64s2048"
|
||||
- "1q128s2048"
|
||||
- "1q256s2048"
|
||||
- "1q512s2048"
|
||||
- "1q1024s2048"
|
||||
- "1q2048"
|
||||
|
||||
# Batch size 2
|
||||
# seq_len = 32
|
||||
- "2q1s32"
|
||||
- "2q2s32"
|
||||
- "2q4s32"
|
||||
- "2q8s32"
|
||||
- "2q16s32"
|
||||
- "2q32"
|
||||
# seq_len = 64
|
||||
- "2q1s64"
|
||||
- "2q2s64"
|
||||
- "2q4s64"
|
||||
- "2q8s64"
|
||||
- "2q16s64"
|
||||
- "2q32s64"
|
||||
- "2q64"
|
||||
# seq_len = 128
|
||||
- "2q1s128"
|
||||
- "2q2s128"
|
||||
- "2q4s128"
|
||||
- "2q8s128"
|
||||
- "2q16s128"
|
||||
- "2q32s128"
|
||||
- "2q64s128"
|
||||
- "2q128"
|
||||
# seq_len = 256
|
||||
- "2q1s256"
|
||||
- "2q2s256"
|
||||
- "2q4s256"
|
||||
- "2q8s256"
|
||||
- "2q16s256"
|
||||
- "2q32s256"
|
||||
- "2q64s256"
|
||||
- "2q128s256"
|
||||
- "2q256"
|
||||
# seq_len = 512
|
||||
- "2q1s512"
|
||||
- "2q2s512"
|
||||
- "2q4s512"
|
||||
- "2q8s512"
|
||||
- "2q16s512"
|
||||
- "2q32s512"
|
||||
- "2q64s512"
|
||||
- "2q128s512"
|
||||
- "2q256s512"
|
||||
- "2q512"
|
||||
# seq_len = 1024
|
||||
- "2q1s1024"
|
||||
- "2q2s1024"
|
||||
- "2q4s1024"
|
||||
- "2q8s1024"
|
||||
- "2q16s1024"
|
||||
- "2q32s1024"
|
||||
- "2q64s1024"
|
||||
- "2q128s1024"
|
||||
- "2q256s1024"
|
||||
- "2q512s1024"
|
||||
- "2q1024"
|
||||
# seq_len = 2048
|
||||
- "2q1s2048"
|
||||
- "2q2s2048"
|
||||
- "2q4s2048"
|
||||
- "2q8s2048"
|
||||
- "2q16s2048"
|
||||
- "2q32s2048"
|
||||
- "2q64s2048"
|
||||
- "2q128s2048"
|
||||
- "2q256s2048"
|
||||
- "2q512s2048"
|
||||
- "2q1024s2048"
|
||||
- "2q2048"
|
||||
|
||||
# Batch size 4
|
||||
# seq_len = 32
|
||||
- "4q1s32"
|
||||
- "4q2s32"
|
||||
- "4q4s32"
|
||||
- "4q8s32"
|
||||
- "4q16s32"
|
||||
- "4q32"
|
||||
# seq_len = 64
|
||||
- "4q1s64"
|
||||
- "4q2s64"
|
||||
- "4q4s64"
|
||||
- "4q8s64"
|
||||
- "4q16s64"
|
||||
- "4q32s64"
|
||||
- "4q64"
|
||||
# seq_len = 128
|
||||
- "4q1s128"
|
||||
- "4q2s128"
|
||||
- "4q4s128"
|
||||
- "4q8s128"
|
||||
- "4q16s128"
|
||||
- "4q32s128"
|
||||
- "4q64s128"
|
||||
- "4q128"
|
||||
# seq_len = 256
|
||||
- "4q1s256"
|
||||
- "4q2s256"
|
||||
- "4q4s256"
|
||||
- "4q8s256"
|
||||
- "4q16s256"
|
||||
- "4q32s256"
|
||||
- "4q64s256"
|
||||
- "4q128s256"
|
||||
- "4q256"
|
||||
# seq_len = 512
|
||||
- "4q1s512"
|
||||
- "4q2s512"
|
||||
- "4q4s512"
|
||||
- "4q8s512"
|
||||
- "4q16s512"
|
||||
- "4q32s512"
|
||||
- "4q64s512"
|
||||
- "4q128s512"
|
||||
- "4q256s512"
|
||||
- "4q512"
|
||||
# seq_len = 1024
|
||||
- "4q1s1024"
|
||||
- "4q2s1024"
|
||||
- "4q4s1024"
|
||||
- "4q8s1024"
|
||||
- "4q16s1024"
|
||||
- "4q32s1024"
|
||||
- "4q64s1024"
|
||||
- "4q128s1024"
|
||||
- "4q256s1024"
|
||||
- "4q512s1024"
|
||||
- "4q1024"
|
||||
# seq_len = 2048
|
||||
- "4q1s2048"
|
||||
- "4q2s2048"
|
||||
- "4q4s2048"
|
||||
- "4q8s2048"
|
||||
- "4q16s2048"
|
||||
- "4q32s2048"
|
||||
- "4q64s2048"
|
||||
- "4q128s2048"
|
||||
- "4q256s2048"
|
||||
- "4q512s2048"
|
||||
- "4q1024s2048"
|
||||
- "4q2048"
|
||||
|
||||
# Batch size 8
|
||||
# seq_len = 32
|
||||
- "8q1s32"
|
||||
- "8q2s32"
|
||||
- "8q4s32"
|
||||
- "8q8s32"
|
||||
- "8q16s32"
|
||||
- "8q32"
|
||||
# seq_len = 64
|
||||
- "8q1s64"
|
||||
- "8q2s64"
|
||||
- "8q4s64"
|
||||
- "8q8s64"
|
||||
- "8q16s64"
|
||||
- "8q32s64"
|
||||
- "8q64"
|
||||
# seq_len = 128
|
||||
- "8q1s128"
|
||||
- "8q2s128"
|
||||
- "8q4s128"
|
||||
- "8q8s128"
|
||||
- "8q16s128"
|
||||
- "8q32s128"
|
||||
- "8q64s128"
|
||||
- "8q128"
|
||||
# seq_len = 256
|
||||
- "8q1s256"
|
||||
- "8q2s256"
|
||||
- "8q4s256"
|
||||
- "8q8s256"
|
||||
- "8q16s256"
|
||||
- "8q32s256"
|
||||
- "8q64s256"
|
||||
- "8q128s256"
|
||||
- "8q256"
|
||||
# seq_len = 512
|
||||
- "8q1s512"
|
||||
- "8q2s512"
|
||||
- "8q4s512"
|
||||
- "8q8s512"
|
||||
- "8q16s512"
|
||||
- "8q32s512"
|
||||
- "8q64s512"
|
||||
- "8q128s512"
|
||||
- "8q256s512"
|
||||
- "8q512"
|
||||
# seq_len = 1024
|
||||
- "8q1s1024"
|
||||
- "8q2s1024"
|
||||
- "8q4s1024"
|
||||
- "8q8s1024"
|
||||
- "8q16s1024"
|
||||
- "8q32s1024"
|
||||
- "8q64s1024"
|
||||
- "8q128s1024"
|
||||
- "8q256s1024"
|
||||
- "8q512s1024"
|
||||
- "8q1024"
|
||||
# seq_len = 2048
|
||||
- "8q1s2048"
|
||||
- "8q2s2048"
|
||||
- "8q4s2048"
|
||||
- "8q8s2048"
|
||||
- "8q16s2048"
|
||||
- "8q32s2048"
|
||||
- "8q64s2048"
|
||||
- "8q128s2048"
|
||||
- "8q256s2048"
|
||||
- "8q512s2048"
|
||||
- "8q1024s2048"
|
||||
- "8q2048"
|
||||
|
||||
# Batch size 16
|
||||
# seq_len = 32
|
||||
- "16q1s32"
|
||||
- "16q2s32"
|
||||
- "16q4s32"
|
||||
- "16q8s32"
|
||||
- "16q16s32"
|
||||
- "16q32"
|
||||
# seq_len = 64
|
||||
- "16q1s64"
|
||||
- "16q2s64"
|
||||
- "16q4s64"
|
||||
- "16q8s64"
|
||||
- "16q16s64"
|
||||
- "16q32s64"
|
||||
- "16q64"
|
||||
# seq_len = 128
|
||||
- "16q1s128"
|
||||
- "16q2s128"
|
||||
- "16q4s128"
|
||||
- "16q8s128"
|
||||
- "16q16s128"
|
||||
- "16q32s128"
|
||||
- "16q64s128"
|
||||
- "16q128"
|
||||
# seq_len = 256
|
||||
- "16q1s256"
|
||||
- "16q2s256"
|
||||
- "16q4s256"
|
||||
- "16q8s256"
|
||||
- "16q16s256"
|
||||
- "16q32s256"
|
||||
- "16q64s256"
|
||||
- "16q128s256"
|
||||
- "16q256"
|
||||
# seq_len = 512
|
||||
- "16q1s512"
|
||||
- "16q2s512"
|
||||
- "16q4s512"
|
||||
- "16q8s512"
|
||||
- "16q16s512"
|
||||
- "16q32s512"
|
||||
- "16q64s512"
|
||||
- "16q128s512"
|
||||
- "16q256s512"
|
||||
- "16q512"
|
||||
# seq_len = 1024
|
||||
- "16q1s1024"
|
||||
- "16q2s1024"
|
||||
- "16q4s1024"
|
||||
- "16q8s1024"
|
||||
- "16q16s1024"
|
||||
- "16q32s1024"
|
||||
- "16q64s1024"
|
||||
- "16q128s1024"
|
||||
- "16q256s1024"
|
||||
- "16q512s1024"
|
||||
- "16q1024"
|
||||
# seq_len = 2048
|
||||
- "16q1s2048"
|
||||
- "16q2s2048"
|
||||
- "16q4s2048"
|
||||
- "16q8s2048"
|
||||
- "16q16s2048"
|
||||
- "16q32s2048"
|
||||
- "16q64s2048"
|
||||
- "16q128s2048"
|
||||
- "16q256s2048"
|
||||
- "16q512s2048"
|
||||
- "16q1024s2048"
|
||||
- "16q2048"
|
||||
|
||||
# Batch size 32
|
||||
# seq_len = 32
|
||||
- "32q1s32"
|
||||
- "32q2s32"
|
||||
- "32q4s32"
|
||||
- "32q8s32"
|
||||
- "32q16s32"
|
||||
- "32q32"
|
||||
# seq_len = 64
|
||||
- "32q1s64"
|
||||
- "32q2s64"
|
||||
- "32q4s64"
|
||||
- "32q8s64"
|
||||
- "32q16s64"
|
||||
- "32q32s64"
|
||||
- "32q64"
|
||||
# seq_len = 128
|
||||
- "32q1s128"
|
||||
- "32q2s128"
|
||||
- "32q4s128"
|
||||
- "32q8s128"
|
||||
- "32q16s128"
|
||||
- "32q32s128"
|
||||
- "32q64s128"
|
||||
- "32q128"
|
||||
# seq_len = 256
|
||||
- "32q1s256"
|
||||
- "32q2s256"
|
||||
- "32q4s256"
|
||||
- "32q8s256"
|
||||
- "32q16s256"
|
||||
- "32q32s256"
|
||||
- "32q64s256"
|
||||
- "32q128s256"
|
||||
- "32q256"
|
||||
# seq_len = 512
|
||||
- "32q1s512"
|
||||
- "32q2s512"
|
||||
- "32q4s512"
|
||||
- "32q8s512"
|
||||
- "32q16s512"
|
||||
- "32q32s512"
|
||||
- "32q64s512"
|
||||
- "32q128s512"
|
||||
- "32q256s512"
|
||||
- "32q512"
|
||||
# seq_len = 1024
|
||||
- "32q1s1024"
|
||||
- "32q2s1024"
|
||||
- "32q4s1024"
|
||||
- "32q8s1024"
|
||||
- "32q16s1024"
|
||||
- "32q32s1024"
|
||||
- "32q64s1024"
|
||||
- "32q128s1024"
|
||||
- "32q256s1024"
|
||||
- "32q512s1024"
|
||||
- "32q1024"
|
||||
# seq_len = 2048
|
||||
- "32q1s2048"
|
||||
- "32q2s2048"
|
||||
- "32q4s2048"
|
||||
- "32q8s2048"
|
||||
- "32q16s2048"
|
||||
- "32q32s2048"
|
||||
- "32q64s2048"
|
||||
- "32q128s2048"
|
||||
- "32q256s2048"
|
||||
- "32q512s2048"
|
||||
- "32q1024s2048"
|
||||
- "32q2048"
|
||||
|
||||
backends:
|
||||
- FLASHMLA_SPARSE
|
||||
|
||||
device: "cuda:0"
|
||||
profile_memory: false
|
||||
sparse_mla_dense_mha_max_seq_len: 2048
|
||||
sparse_mla_topk_pattern: "random"
|
||||
|
||||
output:
|
||||
csv: "benchmark_output/mla_sparse_mha_vs_mqa.csv"
|
||||
json: "benchmark_output/mla_sparse_mha_vs_mqa.json"
|
||||
@@ -9,6 +9,8 @@ needing full VllmConfig integration.
|
||||
"""
|
||||
|
||||
import statistics
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -17,7 +19,6 @@ from common import (
|
||||
BenchmarkResult,
|
||||
MockHfConfig,
|
||||
MockIndexer,
|
||||
MockKVBProj,
|
||||
MockLayer,
|
||||
run_do_bench,
|
||||
run_ncu_profile,
|
||||
@@ -33,8 +34,59 @@ from vllm.config import (
|
||||
VllmConfig,
|
||||
set_current_vllm_config,
|
||||
)
|
||||
from vllm.model_executor.layers.linear import ColumnParallelLinear
|
||||
from vllm.v1.attention.backends.mla.prefill.registry import MLAPrefillBackendEnum
|
||||
|
||||
|
||||
def _safe_profile_name(value: str) -> str:
|
||||
return "".join(c if c.isalnum() or c in "._-" else "_" for c in value)
|
||||
|
||||
|
||||
def _create_kv_b_proj(
|
||||
mla_dims: dict,
|
||||
device: torch.device,
|
||||
):
|
||||
kv_b_proj = ColumnParallelLinear(
|
||||
mla_dims["kv_lora_rank"],
|
||||
mla_dims["num_q_heads"]
|
||||
* (mla_dims["qk_nope_head_dim"] + mla_dims["v_head_dim"]),
|
||||
bias=False,
|
||||
params_dtype=torch.bfloat16,
|
||||
quant_config=None,
|
||||
prefix="benchmark.kv_b_proj",
|
||||
).to(device)
|
||||
with torch.no_grad():
|
||||
kv_b_proj.weight.copy_(torch.randn_like(kv_b_proj.weight))
|
||||
return kv_b_proj
|
||||
|
||||
|
||||
def _ensure_single_rank_model_parallel() -> None:
|
||||
import torch.distributed as dist
|
||||
|
||||
from vllm.distributed import (
|
||||
ensure_model_parallel_initialized,
|
||||
init_distributed_environment,
|
||||
model_parallel_is_initialized,
|
||||
)
|
||||
|
||||
if not dist.is_available():
|
||||
return
|
||||
if not dist.is_initialized():
|
||||
with tempfile.NamedTemporaryFile(
|
||||
prefix="vllm_bench_dist_", delete=False
|
||||
) as init_file:
|
||||
distributed_init_method = f"file://{init_file.name}"
|
||||
init_distributed_environment(
|
||||
world_size=1,
|
||||
rank=0,
|
||||
distributed_init_method=distributed_init_method,
|
||||
local_rank=0,
|
||||
backend="nccl",
|
||||
)
|
||||
if not model_parallel_is_initialized():
|
||||
ensure_model_parallel_initialized(1, 1)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# VllmConfig Creation
|
||||
# ============================================================================
|
||||
@@ -66,10 +118,12 @@ def create_minimal_vllm_config(
|
||||
block_size: int = 128,
|
||||
max_num_seqs: int = 256,
|
||||
max_num_batched_tokens: int = 8192,
|
||||
max_model_len: int = 32768,
|
||||
mla_dims: dict | None = None,
|
||||
index_topk: int | None = None,
|
||||
prefill_backend: str | None = None,
|
||||
kv_cache_dtype: str = "auto",
|
||||
sparse_mla_force_mqa: bool = False,
|
||||
) -> VllmConfig:
|
||||
"""
|
||||
Create minimal VllmConfig for MLA benchmarks.
|
||||
@@ -86,6 +140,8 @@ def create_minimal_vllm_config(
|
||||
prefill_backend: Prefill backend name (e.g., "fa3", "fa4", "flashinfer",
|
||||
"trtllm"). Configures the attention config to force
|
||||
the specified prefill backend.
|
||||
sparse_mla_force_mqa: If True, forces all sparse MLA tokens through
|
||||
forward_mqa (even prefill tokens).
|
||||
|
||||
Returns:
|
||||
VllmConfig for benchmarking
|
||||
@@ -131,7 +187,7 @@ def create_minimal_vllm_config(
|
||||
trust_remote_code=True,
|
||||
dtype="bfloat16",
|
||||
seed=0,
|
||||
max_model_len=32768,
|
||||
max_model_len=max_model_len,
|
||||
quantization=None,
|
||||
enforce_eager=False,
|
||||
max_logprobs=20,
|
||||
@@ -163,7 +219,7 @@ def create_minimal_vllm_config(
|
||||
scheduler_config = SchedulerConfig(
|
||||
max_num_seqs=max_num_seqs,
|
||||
max_num_batched_tokens=max(max_num_batched_tokens, max_num_seqs),
|
||||
max_model_len=32768,
|
||||
max_model_len=max_model_len,
|
||||
is_encoder_decoder=False,
|
||||
enable_chunked_prefill=True,
|
||||
)
|
||||
@@ -192,6 +248,9 @@ def create_minimal_vllm_config(
|
||||
"flash_attn_version"
|
||||
]
|
||||
|
||||
if sparse_mla_force_mqa:
|
||||
vllm_config.attention_config.sparse_mla_force_mqa = True
|
||||
|
||||
return vllm_config
|
||||
|
||||
|
||||
@@ -548,12 +607,7 @@ def _create_backend_impl(
|
||||
# Calculate scale
|
||||
scale = 1.0 / np.sqrt(mla_dims["qk_nope_head_dim"] + mla_dims["qk_rope_head_dim"])
|
||||
|
||||
# Create mock kv_b_proj layer for prefill mode
|
||||
mock_kv_b_proj = MockKVBProj(
|
||||
num_heads=mla_dims["num_q_heads"],
|
||||
qk_nope_head_dim=mla_dims["qk_nope_head_dim"],
|
||||
v_head_dim=mla_dims["v_head_dim"],
|
||||
)
|
||||
kv_b_proj = _create_kv_b_proj(mla_dims, device)
|
||||
|
||||
# Create indexer for sparse backends
|
||||
indexer = None
|
||||
@@ -584,7 +638,7 @@ def _create_backend_impl(
|
||||
"qk_rope_head_dim": mla_dims["qk_rope_head_dim"],
|
||||
"qk_head_dim": mla_dims["qk_nope_head_dim"] + mla_dims["qk_rope_head_dim"],
|
||||
"v_head_dim": mla_dims["v_head_dim"],
|
||||
"kv_b_proj": mock_kv_b_proj,
|
||||
"kv_b_proj": kv_b_proj,
|
||||
}
|
||||
|
||||
# Add indexer for sparse backends
|
||||
@@ -785,14 +839,35 @@ def _run_single_benchmark(
|
||||
# Fill indexer with random indices for sparse backends
|
||||
is_sparse = backend_cfg.get("is_sparse", False)
|
||||
if is_sparse and indexer is not None:
|
||||
indexer.fill_random_indices(total_q, max_kv_len)
|
||||
indexer.fill_indices(
|
||||
total_q,
|
||||
max_kv_len,
|
||||
getattr(config, "sparse_mla_topk_pattern", "random"),
|
||||
)
|
||||
|
||||
# Determine which forward methods to use based on metadata.
|
||||
# Sparse MLA backends always use forward_mqa
|
||||
has_decode = is_sparse or getattr(metadata, "decode", None) is not None
|
||||
has_prefill = not is_sparse and getattr(metadata, "prefill", None) is not None
|
||||
# Non-sparse backends use .decode/.prefill sub-objects.
|
||||
# Sparse backends use num_decode_tokens/num_prefills directly.
|
||||
#
|
||||
# sparse_mla_force_mqa overrides: even for prefill metadata, use MQA.
|
||||
force_mqa = getattr(config, "sparse_mla_force_mqa", False)
|
||||
force_dense_mha = getattr(config, "sparse_mla_mha_mode", "auto") == "dense"
|
||||
if force_mqa:
|
||||
has_decode = True
|
||||
has_prefill = False
|
||||
elif is_sparse:
|
||||
has_decode = metadata.num_decode_tokens > 0
|
||||
has_prefill = metadata.num_prefills > 0
|
||||
else:
|
||||
has_decode = metadata.decode is not None
|
||||
has_prefill = metadata.prefill is not None
|
||||
if not has_decode and not has_prefill:
|
||||
raise RuntimeError("Metadata has neither decode nor prefill metadata")
|
||||
if is_sparse and force_dense_mha and not has_prefill:
|
||||
raise RuntimeError(
|
||||
"Sparse MLA dense_mha benchmark did not produce prefill metadata. "
|
||||
"Check reorder_batch_threshold/path forcing."
|
||||
)
|
||||
|
||||
num_decode = (
|
||||
metadata.num_decode_tokens
|
||||
@@ -871,7 +946,6 @@ def _run_single_benchmark(
|
||||
metadata,
|
||||
prefill_inputs["k_scale"],
|
||||
prefill_fp8_output if fused_output else prefill_inputs["output"],
|
||||
prefill_output_scale if fused_output else None,
|
||||
)
|
||||
if fused_output:
|
||||
out = prefill_fp8_output
|
||||
@@ -898,6 +972,48 @@ def _run_single_benchmark(
|
||||
throughput_tokens_per_sec=0.0,
|
||||
)
|
||||
|
||||
if config.torch_profile:
|
||||
profile_dir = Path(
|
||||
config.torch_profile_dir or "benchmark_outputs/torch_profiles"
|
||||
)
|
||||
profile_dir.mkdir(parents=True, exist_ok=True)
|
||||
trace_name = _safe_profile_name(f"{config.backend}_{config.batch_spec}")
|
||||
trace_path = profile_dir / f"{trace_name}.json"
|
||||
iters = max(config.torch_profile_iters, 1)
|
||||
|
||||
forward_fn()
|
||||
torch.accelerator.synchronize()
|
||||
with torch.profiler.profile(
|
||||
activities=[
|
||||
torch.profiler.ProfilerActivity.CPU,
|
||||
torch.profiler.ProfilerActivity.CUDA,
|
||||
],
|
||||
record_shapes=True,
|
||||
profile_memory=True,
|
||||
with_stack=False,
|
||||
) as prof:
|
||||
for _ in range(iters):
|
||||
forward_fn()
|
||||
torch.accelerator.synchronize()
|
||||
prof.step()
|
||||
prof.export_chrome_trace(str(trace_path))
|
||||
print(f"Saved PyTorch profiler trace to {trace_path}")
|
||||
print(
|
||||
prof.key_averages().table(
|
||||
sort_by="cuda_time_total",
|
||||
row_limit=25,
|
||||
)
|
||||
)
|
||||
return BenchmarkResult(
|
||||
config=config,
|
||||
mean_time=0.0,
|
||||
median_time=0.0,
|
||||
std_time=0.0,
|
||||
min_time=0.0,
|
||||
max_time=0.0,
|
||||
throughput_tokens_per_sec=0.0,
|
||||
)
|
||||
|
||||
all_ms = run_do_bench(benchmark_fn, config.use_cuda_graphs, config.warmup_ms)
|
||||
|
||||
# Convert ms to seconds per layer
|
||||
@@ -920,6 +1036,7 @@ def _run_mla_benchmark_batched(
|
||||
configs_with_params: list[tuple], # [(config, threshold, num_splits), ...]
|
||||
index_topk: int = 2048,
|
||||
prefill_backend: str | None = None,
|
||||
sparse_mla_force_mqa: bool = False,
|
||||
output_scale: float | None = None,
|
||||
fuse_quant_op: bool = False,
|
||||
) -> list[BenchmarkResult]:
|
||||
@@ -940,6 +1057,8 @@ def _run_mla_benchmark_batched(
|
||||
index_topk: Topk value for sparse MLA backends (default 2048)
|
||||
prefill_backend: Prefill backend name (e.g., "fa3", "fa4").
|
||||
When set, forces the specified FlashAttention version for prefill.
|
||||
sparse_mla_force_mqa: If True, forces all sparse MLA tokens through
|
||||
forward_mqa (even prefill tokens).
|
||||
|
||||
Returns:
|
||||
List of BenchmarkResult objects
|
||||
@@ -980,21 +1099,41 @@ def _run_mla_benchmark_batched(
|
||||
sum(r.q_len for r in parse_batch_spec(cfg.batch_spec))
|
||||
for cfg, *_ in configs_with_params
|
||||
)
|
||||
max_model_len = max(
|
||||
max_total_q,
|
||||
max(
|
||||
getattr(cfg, "max_model_len", None) or 32768
|
||||
for cfg, *_ in configs_with_params
|
||||
),
|
||||
)
|
||||
|
||||
# Create and set vLLM config for MLA (reused across all benchmarks)
|
||||
vllm_config = create_minimal_vllm_config(
|
||||
model_name="deepseek-v3", # Used only for model path
|
||||
block_size=block_size,
|
||||
max_num_batched_tokens=max_total_q,
|
||||
max_model_len=max_model_len,
|
||||
mla_dims=mla_dims, # Use custom dims from config or default
|
||||
index_topk=index_topk if is_sparse else None,
|
||||
prefill_backend=prefill_backend,
|
||||
kv_cache_dtype=kv_cache_dtype,
|
||||
sparse_mla_force_mqa=sparse_mla_force_mqa,
|
||||
)
|
||||
|
||||
results = []
|
||||
|
||||
# Initialize workspace manager (needed by metadata builders)
|
||||
from vllm.v1.worker.workspace import (
|
||||
init_workspace_manager,
|
||||
is_workspace_manager_initialized,
|
||||
)
|
||||
|
||||
if not is_workspace_manager_initialized():
|
||||
init_workspace_manager(device)
|
||||
|
||||
with set_current_vllm_config(vllm_config):
|
||||
_ensure_single_rank_model_parallel()
|
||||
|
||||
# Create backend impl, layer, builder, and indexer (reused across benchmarks)
|
||||
impl, layer, builder_instance, indexer = _create_backend_impl(
|
||||
backend_cfg,
|
||||
@@ -1040,9 +1179,20 @@ def _run_mla_benchmark_batched(
|
||||
for config, threshold, num_splits in configs_with_params:
|
||||
# Set threshold for this benchmark (FlashAttn/FlashMLA only)
|
||||
original_threshold = None
|
||||
if threshold is not None and builder_instance:
|
||||
effective_threshold = threshold
|
||||
force_dense_mha = (
|
||||
is_sparse
|
||||
and getattr(config, "sparse_mla_mha_mode", "auto") == "dense"
|
||||
and not getattr(config, "sparse_mla_force_mqa", False)
|
||||
)
|
||||
if force_dense_mha:
|
||||
# Sparse MLA normally treats q_len <= 1 as decode. Use an
|
||||
# impossible threshold so dense_mha benchmarks actually run
|
||||
# the prefill/MHA path, including q_len=1 short extends.
|
||||
effective_threshold = -1
|
||||
if effective_threshold is not None and builder_instance:
|
||||
original_threshold = builder_instance.reorder_batch_threshold
|
||||
builder_instance.reorder_batch_threshold = threshold
|
||||
builder_instance.reorder_batch_threshold = effective_threshold
|
||||
|
||||
# Set num_splits for CUTLASS
|
||||
original_num_splits = None
|
||||
@@ -1090,6 +1240,7 @@ def run_mla_benchmark(
|
||||
num_kv_splits: int | None = None,
|
||||
index_topk: int = 2048,
|
||||
prefill_backend: str | None = None,
|
||||
sparse_mla_force_mqa: bool = False,
|
||||
output_scale: float | None = None,
|
||||
fuse_quant_op: bool = False,
|
||||
) -> BenchmarkResult | list[BenchmarkResult]:
|
||||
@@ -1111,6 +1262,8 @@ def run_mla_benchmark(
|
||||
index_topk: Topk value for sparse MLA backends (default 2048)
|
||||
prefill_backend: Prefill backend name (e.g., "fa3", "fa4").
|
||||
When set, forces the specified FlashAttention version for prefill.
|
||||
sparse_mla_force_mqa: If True, forces all sparse MLA tokens through
|
||||
forward_mqa (even prefill tokens).
|
||||
output_scale: Static per-tensor FP8 scale for prefill output (None = bf16).
|
||||
fuse_quant_op: With output_scale set, fuse the FP8 write into the prefill
|
||||
kernel vs a standalone post-quant kernel. See _run_single_benchmark.
|
||||
@@ -1142,6 +1295,7 @@ def run_mla_benchmark(
|
||||
configs_with_params,
|
||||
index_topk,
|
||||
prefill_backend=prefill_backend,
|
||||
sparse_mla_force_mqa=sparse_mla_force_mqa,
|
||||
output_scale=output_scale,
|
||||
fuse_quant_op=fuse_quant_op,
|
||||
)
|
||||
|
||||
@@ -69,12 +69,11 @@ def make_inputs(total_tokens, num_reqs, block_size):
|
||||
# Output workspace
|
||||
dst = torch.zeros(total_tokens, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
|
||||
|
||||
seq_lens_t = torch.tensor(seq_lens, dtype=torch.int32, device="cuda")
|
||||
workspace_starts_t = torch.tensor(
|
||||
workspace_starts, dtype=torch.int32, device="cuda"
|
||||
)
|
||||
|
||||
return cache, dst, block_table, seq_lens_t, workspace_starts_t
|
||||
return cache, dst, block_table, workspace_starts_t
|
||||
|
||||
|
||||
def bench_scenario(label, num_reqs, total_tokens_list, save_path):
|
||||
@@ -94,7 +93,7 @@ def bench_scenario(label, num_reqs, total_tokens_list, save_path):
|
||||
)
|
||||
)
|
||||
def bench_fn(total_tokens, provider, num_reqs):
|
||||
cache, dst, block_table, seq_lens_t, ws_starts = make_inputs(
|
||||
cache, dst, block_table, ws_starts = make_inputs(
|
||||
total_tokens, num_reqs, BLOCK_SIZE
|
||||
)
|
||||
|
||||
@@ -102,7 +101,7 @@ def bench_scenario(label, num_reqs, total_tokens_list, save_path):
|
||||
|
||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
||||
lambda: ops.cp_gather_and_upconvert_fp8_kv_cache(
|
||||
cache, dst, block_table, seq_lens_t, ws_starts, num_reqs
|
||||
cache, dst, block_table, ws_starts, num_reqs
|
||||
),
|
||||
quantiles=quantiles,
|
||||
rep=500,
|
||||
|
||||
@@ -0,0 +1,367 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Benchmark the Kimi-K3 latent MoE addmm against CuTe residual GEMM.
|
||||
|
||||
The benchmark covers ``BF16[M, 3584] @ BF16[7168, 3584].T + BF16[M, 7168]``
|
||||
with FP32 accumulation and BF16 output. Both backends execute through CUDA
|
||||
Graph replay. Weights and residuals rotate across buffers exceeding L2 so the
|
||||
comparison models the full latent MoE projection-and-add path.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
import importlib.util
|
||||
import json
|
||||
import math
|
||||
import statistics
|
||||
from collections.abc import Callable, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import cutlass
|
||||
import cutlass.cute as cute
|
||||
import torch
|
||||
from cuda.bindings import driver as cuda
|
||||
from cuda.bindings.driver import CUstream
|
||||
from quack.compile_utils import make_fake_tensor
|
||||
|
||||
N = 7168
|
||||
K = 3584
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True, slots=True)
|
||||
class Config:
|
||||
block_size: int
|
||||
outputs_per_block: int
|
||||
k_unroll: int
|
||||
vector_width: int = 8
|
||||
|
||||
|
||||
def parse_config(value: str) -> Config:
|
||||
try:
|
||||
parts = [int(part) for part in value.split(",")]
|
||||
except ValueError as error:
|
||||
raise argparse.ArgumentTypeError(
|
||||
"config must be BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH]"
|
||||
) from error
|
||||
if len(parts) == 3:
|
||||
return Config(*parts)
|
||||
if len(parts) == 4:
|
||||
return Config(*parts)
|
||||
raise argparse.ArgumentTypeError(
|
||||
"config must be BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH]"
|
||||
)
|
||||
|
||||
|
||||
def production_residual_config(m: int) -> Config | None:
|
||||
"""The measured Latent-MoE residual config for M, from the K3 table."""
|
||||
from vllm.models.kimi_k3.nvidia.low_latency_gemm import KIMI_K3_PROJECTIONS
|
||||
|
||||
spec = KIMI_K3_PROJECTIONS.get((N, K))
|
||||
config = spec.residual_config(m) if spec is not None else None
|
||||
if config is None:
|
||||
return None
|
||||
return Config(
|
||||
config.block_size,
|
||||
config.outputs_per_block,
|
||||
config.k_unroll,
|
||||
config.vector_width,
|
||||
)
|
||||
|
||||
|
||||
def candidate_configs(mode: str, selected: Config | None, m: int) -> list[Config]:
|
||||
if mode == "selected":
|
||||
if selected is not None:
|
||||
return [selected]
|
||||
# No explicit --config: fall back to the production table for this M.
|
||||
config = production_residual_config(m)
|
||||
return [config] if config is not None else []
|
||||
if mode == "baseline":
|
||||
return [Config(224, 4, 2)]
|
||||
return [
|
||||
Config(block_size, outputs_per_block, k_unroll, vector_width)
|
||||
for vector_width in (4, 8)
|
||||
for block_size in (32, 64, 128, 224, 448)
|
||||
if block_size % 32 == 0 and K % (block_size * vector_width) == 0
|
||||
for outputs_per_block in (1, 2, 4, 7, 8)
|
||||
if N % outputs_per_block == 0
|
||||
for k_unroll in (1, 2, 4)
|
||||
]
|
||||
|
||||
|
||||
def load_kernel_class(path: Path):
|
||||
spec = importlib.util.spec_from_file_location("cute_skinny_device", path)
|
||||
if spec is None or spec.loader is None:
|
||||
raise RuntimeError(f"cannot load CuTe kernel from {path}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module.CuteSkinnyGemm
|
||||
|
||||
|
||||
def stream() -> CUstream:
|
||||
return CUstream(torch.cuda.current_stream().cuda_stream)
|
||||
|
||||
|
||||
def compile_kernel(kernel_class, m: int, config: Config, max_registers: int):
|
||||
element_type = cutlass.BFloat16
|
||||
n = cute.sym_int(divisibility=config.outputs_per_block)
|
||||
k = cute.sym_int(divisibility=config.block_size * config.vector_width)
|
||||
a = make_fake_tensor(element_type, (m, k), divisibility=config.vector_width)
|
||||
b = make_fake_tensor(element_type, (n, k), divisibility=config.vector_width)
|
||||
residual = make_fake_tensor(element_type, (m, n), divisibility=1)
|
||||
c = make_fake_tensor(element_type, (m, n), divisibility=1)
|
||||
kernel = kernel_class(
|
||||
element_type=element_type,
|
||||
num_rows=m,
|
||||
block_size=config.block_size,
|
||||
outputs_per_block=config.outputs_per_block,
|
||||
vector_width=config.vector_width,
|
||||
k_unroll=config.k_unroll,
|
||||
has_residual=True,
|
||||
use_pdl=True,
|
||||
)
|
||||
return cute.compile(
|
||||
kernel,
|
||||
a,
|
||||
b,
|
||||
residual,
|
||||
c,
|
||||
stream(),
|
||||
options=(
|
||||
"--enable-tvm-ffi --keep-cubin "
|
||||
f"--ptxas-options -maxrregcount={max_registers} "
|
||||
"--ptxas-options -lineinfo"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def resource_usage(compiled) -> dict[str, Any]:
|
||||
executor = getattr(compiled, "_default_executor", None)
|
||||
context = getattr(executor, "exec_context", None)
|
||||
functions = getattr(context, "kernel_functions", None)
|
||||
if not functions:
|
||||
return {"resource_metrics_available": False}
|
||||
|
||||
def attribute(name, function) -> int:
|
||||
error, value = cuda.cuFuncGetAttribute(name, function)
|
||||
if error != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError(f"cuFuncGetAttribute failed with {error}")
|
||||
return int(value)
|
||||
|
||||
registers = [
|
||||
attribute(cuda.CUfunction_attribute.CU_FUNC_ATTRIBUTE_NUM_REGS, function)
|
||||
for function in functions
|
||||
]
|
||||
local_bytes = [
|
||||
attribute(
|
||||
cuda.CUfunction_attribute.CU_FUNC_ATTRIBUTE_LOCAL_SIZE_BYTES,
|
||||
function,
|
||||
)
|
||||
for function in functions
|
||||
]
|
||||
return {
|
||||
"resource_metrics_available": True,
|
||||
"registers_per_thread": max(registers, default=0),
|
||||
"spill_bytes": max(local_bytes, default=0),
|
||||
}
|
||||
|
||||
|
||||
def rotating_buffer_count(m: int, multiplier: float, limit: int) -> int:
|
||||
properties = torch.cuda.get_device_properties(0)
|
||||
bytes_per_pair = (N * K + m * N) * 2
|
||||
target = math.ceil(multiplier * properties.L2_cache_size)
|
||||
return max(2, min(limit, math.ceil(target / bytes_per_pair)))
|
||||
|
||||
|
||||
def graph_samples(
|
||||
launch: Callable[[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], None],
|
||||
activation: torch.Tensor,
|
||||
weights: Sequence[torch.Tensor],
|
||||
residuals: Sequence[torch.Tensor],
|
||||
repeats: int,
|
||||
replays: int,
|
||||
) -> tuple[list[float], list[torch.Tensor]]:
|
||||
outputs = [torch.empty_like(residual) for residual in residuals]
|
||||
for weight, residual, output in zip(weights, residuals, outputs):
|
||||
launch(activation, weight, residual, output)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
for weight, residual, output in zip(weights, residuals, outputs):
|
||||
launch(activation, weight, residual, output)
|
||||
for _ in range(20):
|
||||
graph.replay()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
samples = []
|
||||
for _ in range(repeats):
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
for _ in range(replays):
|
||||
graph.replay()
|
||||
end.record()
|
||||
end.synchronize()
|
||||
samples.append(start.elapsed_time(end) * 1000.0 / (replays * len(weights)))
|
||||
return samples, outputs
|
||||
|
||||
|
||||
def summarize(samples: Sequence[float]) -> dict[str, Any]:
|
||||
ordered = sorted(samples)
|
||||
|
||||
def percentile(fraction: float) -> float:
|
||||
position = fraction * (len(ordered) - 1)
|
||||
lower = math.floor(position)
|
||||
upper = math.ceil(position)
|
||||
if lower == upper:
|
||||
return ordered[lower]
|
||||
weight = position - lower
|
||||
return ordered[lower] * (1.0 - weight) + ordered[upper] * weight
|
||||
|
||||
mean = statistics.mean(samples)
|
||||
return {
|
||||
"median_us": statistics.median(samples),
|
||||
"p10_us": percentile(0.1),
|
||||
"p90_us": percentile(0.9),
|
||||
"mean_us": mean,
|
||||
"cv_pct": statistics.pstdev(samples) / mean * 100.0,
|
||||
"samples_us": list(samples),
|
||||
}
|
||||
|
||||
|
||||
def correctness(
|
||||
output: torch.Tensor,
|
||||
activation: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
) -> dict[str, Any]:
|
||||
actual = output.float()
|
||||
reference = activation.float() @ weight.float().t() + residual.float()
|
||||
error = (actual - reference).abs()
|
||||
scaled_error = error / (reference.abs() + 1.0)
|
||||
cosine = torch.nn.functional.cosine_similarity(
|
||||
actual.flatten(), reference.flatten(), dim=0
|
||||
).item()
|
||||
return {
|
||||
"valid": cosine > 0.999,
|
||||
"cosine": cosine,
|
||||
"max_abs_error": error.max().item(),
|
||||
"max_scaled_error": scaled_error.max().item(),
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--kernel", type=Path, required=True)
|
||||
parser.add_argument("--output", type=Path, required=True)
|
||||
parser.add_argument(
|
||||
"--mode", choices=("baseline", "sweep", "selected"), default="baseline"
|
||||
)
|
||||
parser.add_argument("--config", type=parse_config)
|
||||
parser.add_argument("--m", type=int, action="append")
|
||||
parser.add_argument("--config-shard", type=int, default=0)
|
||||
parser.add_argument("--num-config-shards", type=int, default=1)
|
||||
parser.add_argument("--repeats", type=int, default=21)
|
||||
parser.add_argument("--replays", type=int, default=200)
|
||||
parser.add_argument("--cache-multiplier", type=float, default=3.0)
|
||||
parser.add_argument("--max-buffers", type=int, default=32)
|
||||
parser.add_argument("--max-registers", type=int, default=64)
|
||||
args = parser.parse_args()
|
||||
|
||||
token_counts = args.m or list(range(1, 17))
|
||||
if any(not 1 <= m <= 16 for m in token_counts):
|
||||
raise ValueError("expected 1 <= M <= 16")
|
||||
if not 0 <= args.config_shard < args.num_config_shards:
|
||||
raise ValueError("config shard must be in [0, num_config_shards)")
|
||||
torch.cuda.set_device(0)
|
||||
if torch.cuda.get_device_capability() != (10, 3):
|
||||
raise RuntimeError("this benchmark requires SM103")
|
||||
|
||||
kernel_class = load_kernel_class(args.kernel)
|
||||
properties = torch.cuda.get_device_properties(0)
|
||||
metadata = {
|
||||
"device": properties.name,
|
||||
"compute_capability": list(torch.cuda.get_device_capability()),
|
||||
"torch_version": torch.__version__,
|
||||
"cuda_version": torch.version.cuda,
|
||||
}
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
with args.output.open("w", encoding="utf-8") as output_file:
|
||||
for m in token_counts:
|
||||
configs = candidate_configs(args.mode, args.config, m)
|
||||
torch.manual_seed(20260722 + m)
|
||||
count = rotating_buffer_count(m, args.cache_multiplier, args.max_buffers)
|
||||
activation = torch.randn((m, K), device="cuda", dtype=torch.bfloat16)
|
||||
weights = [
|
||||
torch.randn((N, K), device="cuda", dtype=torch.bfloat16)
|
||||
for _ in range(count)
|
||||
]
|
||||
residuals = [
|
||||
torch.randn((m, N), device="cuda", dtype=torch.bfloat16)
|
||||
for _ in range(count)
|
||||
]
|
||||
candidates: list[tuple[str, Config | None]] = [("cublas_addmm", None)]
|
||||
candidates.extend(
|
||||
("cute_residual", config)
|
||||
for index, config in enumerate(configs)
|
||||
if index % args.num_config_shards == args.config_shard
|
||||
)
|
||||
for backend, config in candidates:
|
||||
row: dict[str, Any] = {
|
||||
"m": m,
|
||||
"n": N,
|
||||
"k": K,
|
||||
"backend": backend,
|
||||
"mode": args.mode,
|
||||
"config": dataclasses.asdict(config) if config else {},
|
||||
"num_buffers": count,
|
||||
"cache_multiplier": args.cache_multiplier,
|
||||
**metadata,
|
||||
}
|
||||
try:
|
||||
if backend == "cublas_addmm":
|
||||
launch = lambda a, b, residual, c: torch.addmm(
|
||||
residual, a, b.t(), out=c
|
||||
)
|
||||
else:
|
||||
if config is None:
|
||||
raise AssertionError("missing CuTe config")
|
||||
compiled = compile_kernel(
|
||||
kernel_class, m, config, args.max_registers
|
||||
)
|
||||
launch = lambda a, b, residual, c, fn=compiled: fn(
|
||||
a, b, residual, c, stream()
|
||||
)
|
||||
row.update(resource_usage(compiled))
|
||||
samples, outputs = graph_samples(
|
||||
launch,
|
||||
activation,
|
||||
weights,
|
||||
residuals,
|
||||
args.repeats,
|
||||
args.replays,
|
||||
)
|
||||
row.update(
|
||||
correctness(outputs[0], activation, weights[0], residuals[0])
|
||||
)
|
||||
row.update(summarize(samples))
|
||||
except Exception as error: # noqa: BLE001
|
||||
row.update(
|
||||
{
|
||||
"valid": False,
|
||||
"error": f"{type(error).__name__}: {error}",
|
||||
}
|
||||
)
|
||||
output_file.write(json.dumps(row, sort_keys=True) + "\n")
|
||||
output_file.flush()
|
||||
print(json.dumps(row, sort_keys=True), flush=True)
|
||||
|
||||
del activation, weights, residuals
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,806 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Benchmark the Kimi K3 latent-MoE tail and its up-projection kernels.
|
||||
|
||||
The ``up-projection`` subcommand isolates the TP-local dynamic and static-M
|
||||
skinny GEMMs. It rotates weights through a working set larger than L2 to model
|
||||
successive model layers.
|
||||
|
||||
The ``whole-tail`` subcommand measures the distributed operator. Its reference
|
||||
path includes two AllReduces, RMSNorm, the replicated up-projection, and the
|
||||
final add. CUDA-event samples report the slowest rank so cross-rank skew is
|
||||
included.
|
||||
|
||||
Examples:
|
||||
|
||||
.. code-block:: console
|
||||
|
||||
.venv/bin/python \
|
||||
benchmarks/kernels/benchmark_kimi_k3_latent_moe_tail.py up-projection
|
||||
|
||||
torchrun --nproc-per-node=8 \
|
||||
benchmarks/kernels/benchmark_kimi_k3_latent_moe_tail.py whole-tail
|
||||
|
||||
For multi-node runs, launch one ``torchrun`` agent per node and use a shared
|
||||
rendezvous endpoint.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import statistics
|
||||
from collections.abc import Callable, Sequence
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import cutlass
|
||||
import cutlass.utils as utils
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
from cuda.bindings import driver as cuda
|
||||
|
||||
from vllm.distributed import get_tp_group
|
||||
from vllm.distributed.parallel_state import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
set_custom_all_reduce,
|
||||
)
|
||||
from vllm.model_executor.warmup.cutedsl_warmup import cutedsl_warmup
|
||||
from vllm.models.kimi_k3.nvidia.ops import latent_moe_tail
|
||||
from vllm.models.kimi_k3.nvidia.ops.cute_dsl.latent_moe_tail import (
|
||||
fused_add_multicast_gemm,
|
||||
fused_add_multicast_skinny_gemm,
|
||||
)
|
||||
|
||||
HIDDEN_SIZE = 7168
|
||||
LATENT_SIZE = 3584
|
||||
RMS_EPS = 0.1
|
||||
MAX_NUM_TOKENS = 16
|
||||
MMA_TILER_MN = (64, 32)
|
||||
CLUSTER_SHAPE_MN = (1, 8)
|
||||
B_PRIME_STAGES = 2
|
||||
|
||||
|
||||
def parse_up_projection_config(
|
||||
value: str,
|
||||
) -> fused_add_multicast_skinny_gemm.SkinnyConfig:
|
||||
try:
|
||||
values = [int(part) for part in value.split(",")]
|
||||
except ValueError as error:
|
||||
raise argparse.ArgumentTypeError(
|
||||
"config must be BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH[,PREFETCH_B]]"
|
||||
) from error
|
||||
if len(values) in (3, 4):
|
||||
return fused_add_multicast_skinny_gemm.SkinnyConfig(*values)
|
||||
if len(values) == 5 and values[4] in (0, 1):
|
||||
return fused_add_multicast_skinny_gemm.SkinnyConfig(
|
||||
*values[:4],
|
||||
prefetch_b_before_pdl=bool(values[4]),
|
||||
)
|
||||
raise argparse.ArgumentTypeError(
|
||||
"config must be BLOCK,OUTPUTS,K_UNROLL"
|
||||
"[,VECTOR_WIDTH[,PREFETCH_B]], where PREFETCH_B is 0 or 1"
|
||||
)
|
||||
|
||||
|
||||
def parse_tail_skinny_config(
|
||||
value: str,
|
||||
) -> tuple[int, fused_add_multicast_skinny_gemm.SkinnyConfig]:
|
||||
try:
|
||||
values = [int(part) for part in value.split(",")]
|
||||
except ValueError as error:
|
||||
raise argparse.ArgumentTypeError(
|
||||
"config must be M,BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH[,PREFETCH_B]]"
|
||||
) from error
|
||||
if len(values) == 4:
|
||||
num_tokens, *config = values
|
||||
return num_tokens, fused_add_multicast_skinny_gemm.SkinnyConfig(*config)
|
||||
if len(values) == 5:
|
||||
num_tokens, *config = values
|
||||
return num_tokens, fused_add_multicast_skinny_gemm.SkinnyConfig(*config)
|
||||
if len(values) == 6 and values[5] in (0, 1):
|
||||
num_tokens, block, outputs, unroll, vector_width, prefetch = values
|
||||
return num_tokens, fused_add_multicast_skinny_gemm.SkinnyConfig(
|
||||
block,
|
||||
outputs,
|
||||
unroll,
|
||||
vector_width,
|
||||
bool(prefetch),
|
||||
)
|
||||
raise argparse.ArgumentTypeError(
|
||||
"config must be M,BLOCK,OUTPUTS,K_UNROLL"
|
||||
"[,VECTOR_WIDTH[,PREFETCH_B]], where PREFETCH_B is 0 or 1"
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
subparsers = parser.add_subparsers(dest="scope", required=True)
|
||||
|
||||
up_projection = subparsers.add_parser(
|
||||
"up-projection",
|
||||
help="Benchmark the isolated TP-local up-projection kernels.",
|
||||
)
|
||||
up_projection.add_argument(
|
||||
"--backend",
|
||||
choices=("dynamic", "skinny", "both"),
|
||||
default="both",
|
||||
)
|
||||
up_projection.add_argument("--tp-size", type=int, default=16)
|
||||
up_projection.add_argument(
|
||||
"--num-tokens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=[*range(1, 9), 16],
|
||||
)
|
||||
up_projection.add_argument(
|
||||
"--skinny-config",
|
||||
type=parse_up_projection_config,
|
||||
action="append",
|
||||
help="Benchmark a static-M config for every selected token count.",
|
||||
)
|
||||
up_projection.add_argument("--cache-multiplier", type=float, default=2.0)
|
||||
up_projection.add_argument("--max-weights", type=int, default=64)
|
||||
up_projection.add_argument("--warmup-replays", type=int, default=10)
|
||||
up_projection.add_argument("--samples", type=int, default=31)
|
||||
up_projection.add_argument("--output", type=Path)
|
||||
|
||||
whole_tail = subparsers.add_parser(
|
||||
"whole-tail",
|
||||
help="Benchmark the distributed latent-MoE tail operator.",
|
||||
)
|
||||
whole_tail.add_argument(
|
||||
"--backend",
|
||||
choices=("reference", "fused", "both"),
|
||||
default="both",
|
||||
)
|
||||
whole_tail.add_argument(
|
||||
"--num-tokens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=[1, 5, 8, 16],
|
||||
)
|
||||
whole_tail.add_argument("--warmup-replays", type=int, default=20)
|
||||
whole_tail.add_argument("--samples", type=int, default=51)
|
||||
whole_tail.add_argument(
|
||||
"--skinny-max-num-tokens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
help="Override the fused operator's static-M cutoff; use 0 for dynamic-only.",
|
||||
)
|
||||
whole_tail.add_argument(
|
||||
"--skinny-config",
|
||||
type=parse_tail_skinny_config,
|
||||
action="append",
|
||||
help="Override one static-M config for tuning.",
|
||||
)
|
||||
whole_tail.add_argument("--output", type=Path)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def percentile(samples: Sequence[float], fraction: float) -> float:
|
||||
ordered = sorted(samples)
|
||||
position = fraction * (len(ordered) - 1)
|
||||
lower = math.floor(position)
|
||||
upper = math.ceil(position)
|
||||
if lower == upper:
|
||||
return ordered[lower]
|
||||
upper_weight = position - lower
|
||||
return ordered[lower] * (1.0 - upper_weight) + ordered[upper] * upper_weight
|
||||
|
||||
|
||||
def summarize(samples_us: Sequence[float]) -> dict[str, Any]:
|
||||
mean_us = statistics.mean(samples_us)
|
||||
return {
|
||||
"median_us": statistics.median(samples_us),
|
||||
"p10_us": percentile(samples_us, 0.1),
|
||||
"p90_us": percentile(samples_us, 0.9),
|
||||
"mean_us": mean_us,
|
||||
"cv_pct": statistics.pstdev(samples_us) / mean_us * 100.0,
|
||||
"samples_us": list(samples_us),
|
||||
}
|
||||
|
||||
|
||||
def rotating_weight_count(
|
||||
shard_size: int,
|
||||
cache_multiplier: float,
|
||||
limit: int,
|
||||
) -> int:
|
||||
properties = torch.cuda.get_device_properties(
|
||||
torch.accelerator.current_device_index()
|
||||
)
|
||||
weight_bytes = shard_size * LATENT_SIZE * 2
|
||||
target_bytes = math.ceil(properties.L2_cache_size * cache_multiplier)
|
||||
return max(2, min(limit, math.ceil(target_bytes / weight_bytes)))
|
||||
|
||||
|
||||
def capture_up_projection_graph(
|
||||
launches: Sequence[Callable[[], None]],
|
||||
) -> torch.cuda.CUDAGraph:
|
||||
for launch in launches:
|
||||
launch()
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
for launch in launches:
|
||||
launch()
|
||||
torch.accelerator.synchronize()
|
||||
return graph
|
||||
|
||||
|
||||
def benchmark_up_projection_graph(
|
||||
graph: torch.cuda.CUDAGraph,
|
||||
*,
|
||||
operations_per_replay: int,
|
||||
warmup_replays: int,
|
||||
samples: int,
|
||||
) -> dict[str, Any]:
|
||||
for _ in range(warmup_replays):
|
||||
graph.replay()
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
samples_us = []
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
for _ in range(samples):
|
||||
start.record()
|
||||
graph.replay()
|
||||
end.record()
|
||||
end.synchronize()
|
||||
samples_us.append(start.elapsed_time(end) * 1000.0 / operations_per_replay)
|
||||
return summarize(samples_us)
|
||||
|
||||
|
||||
class DynamicKernel:
|
||||
def __init__(
|
||||
self,
|
||||
shard_size: int,
|
||||
mailbox: torch.Tensor,
|
||||
shared_shard: torch.Tensor,
|
||||
) -> None:
|
||||
self.shard_size = shard_size
|
||||
self.mailbox = mailbox
|
||||
self.mailbox_c = fused_add_multicast_gemm._as_cute(mailbox)
|
||||
compile_latent = torch.empty(
|
||||
(1, MAX_NUM_TOKENS, LATENT_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=mailbox.device,
|
||||
)
|
||||
compile_weight = torch.empty(
|
||||
(1, shard_size, LATENT_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=mailbox.device,
|
||||
)
|
||||
cluster_size = math.prod(CLUSTER_SHAPE_MN)
|
||||
max_active_clusters = utils.HardwareInfo().get_max_active_clusters(cluster_size)
|
||||
self.compiled = fused_add_multicast_gemm.compile_kernel(
|
||||
(MAX_NUM_TOKENS, shard_size, LATENT_SIZE, 1),
|
||||
fused_add_multicast_gemm._as_cute(
|
||||
compile_latent,
|
||||
dynamic_m=True,
|
||||
),
|
||||
fused_add_multicast_gemm._as_cute(compile_weight),
|
||||
self.mailbox_c,
|
||||
fused_add_multicast_gemm._as_cute(shared_shard),
|
||||
HIDDEN_SIZE,
|
||||
shard_size,
|
||||
MMA_TILER_MN,
|
||||
CLUSTER_SHAPE_MN,
|
||||
max_active_clusters,
|
||||
B_PRIME_STAGES,
|
||||
)
|
||||
|
||||
def launch(
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
shared_shard: torch.Tensor,
|
||||
) -> None:
|
||||
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
|
||||
self.compiled(
|
||||
fused_add_multicast_gemm._as_cute(
|
||||
latent.unsqueeze(0),
|
||||
dynamic_m=True,
|
||||
),
|
||||
fused_add_multicast_gemm._as_cute(weight.unsqueeze(0)),
|
||||
self.mailbox_c,
|
||||
fused_add_multicast_gemm._as_cute(shared_shard),
|
||||
cutlass.Int64(latent.shape[0]),
|
||||
cutlass.Int64(self.mailbox.data_ptr()),
|
||||
stream,
|
||||
)
|
||||
|
||||
|
||||
class SkinnyKernel:
|
||||
def __init__(
|
||||
self,
|
||||
num_tokens: int,
|
||||
shard_size: int,
|
||||
config: fused_add_multicast_skinny_gemm.SkinnyConfig,
|
||||
) -> None:
|
||||
self.compiled = fused_add_multicast_skinny_gemm.compile_kernel(
|
||||
num_rows=num_tokens,
|
||||
latent_dim=LATENT_SIZE,
|
||||
hidden_dim=HIDDEN_SIZE,
|
||||
shard_dim=shard_size,
|
||||
config=config,
|
||||
)
|
||||
|
||||
def launch(
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
shared_shard: torch.Tensor,
|
||||
mailbox: torch.Tensor,
|
||||
) -> None:
|
||||
self.compiled(
|
||||
fused_add_multicast_skinny_gemm._as_cute(latent),
|
||||
fused_add_multicast_skinny_gemm._as_cute(weight),
|
||||
fused_add_multicast_skinny_gemm._as_cute(shared_shard),
|
||||
cutlass.Int64(mailbox.data_ptr()),
|
||||
cuda.CUstream(torch.cuda.current_stream().cuda_stream),
|
||||
)
|
||||
|
||||
|
||||
def check_up_projection_output(
|
||||
actual: torch.Tensor,
|
||||
latent: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
shared_shard: torch.Tensor,
|
||||
) -> None:
|
||||
gemm = F.linear(latent.float(), weight.float()).to(torch.bfloat16)
|
||||
expected = (gemm.float() + shared_shard.float()).to(torch.bfloat16)
|
||||
torch.testing.assert_close(actual, expected, atol=8e-2, rtol=3e-2)
|
||||
|
||||
|
||||
def make_up_projection_launches(
|
||||
launch: Callable[[torch.Tensor, torch.Tensor, torch.Tensor], None],
|
||||
latent: torch.Tensor,
|
||||
weights: Sequence[torch.Tensor],
|
||||
shared_shard: torch.Tensor,
|
||||
) -> list[Callable[[], None]]:
|
||||
return [
|
||||
lambda weight=weight: launch(latent, weight, shared_shard) for weight in weights
|
||||
]
|
||||
|
||||
|
||||
def benchmark_up_projection(args: argparse.Namespace) -> None:
|
||||
if args.tp_size <= 0 or HIDDEN_SIZE % args.tp_size:
|
||||
raise ValueError("TP size must be positive and divide the hidden size")
|
||||
if any(not 1 <= num_tokens <= MAX_NUM_TOKENS for num_tokens in args.num_tokens):
|
||||
raise ValueError("--num-tokens values must be in [1, 16]")
|
||||
if args.cache_multiplier <= 0 or args.max_weights <= 0:
|
||||
raise ValueError("cache multiplier and max weights must be positive")
|
||||
if args.warmup_replays < 0 or args.samples <= 0:
|
||||
raise ValueError("warmup replays must be nonnegative and samples positive")
|
||||
|
||||
torch.accelerator.set_device_index(0)
|
||||
device = torch.device("cuda", 0)
|
||||
if torch.cuda.get_device_capability(device)[0] != 10:
|
||||
raise RuntimeError("Kimi K3 latent-MoE tail requires SM100")
|
||||
|
||||
shard_size = HIDDEN_SIZE // args.tp_size
|
||||
weight_count = rotating_weight_count(
|
||||
shard_size,
|
||||
args.cache_multiplier,
|
||||
args.max_weights,
|
||||
)
|
||||
torch.manual_seed(20260726)
|
||||
weights = [
|
||||
torch.randn(
|
||||
(shard_size, LATENT_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
/ LATENT_SIZE**0.5
|
||||
for _ in range(weight_count)
|
||||
]
|
||||
mailbox = torch.empty(
|
||||
(1, MAX_NUM_TOKENS, HIDDEN_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
shared = torch.randn(
|
||||
(MAX_NUM_TOKENS, HIDDEN_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
shared_shard = shared[:, :shard_size]
|
||||
use_dynamic = args.backend in ("dynamic", "both")
|
||||
use_skinny = args.backend in ("skinny", "both")
|
||||
dynamic_kernel = (
|
||||
DynamicKernel(shard_size, mailbox, shared_shard) if use_dynamic else None
|
||||
)
|
||||
|
||||
results = []
|
||||
for num_tokens in args.num_tokens:
|
||||
latent = torch.randn(
|
||||
(num_tokens, LATENT_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
result: dict[str, Any] = {"num_tokens": num_tokens}
|
||||
if dynamic_kernel is not None:
|
||||
launches = make_up_projection_launches(
|
||||
dynamic_kernel.launch,
|
||||
latent,
|
||||
weights,
|
||||
shared_shard,
|
||||
)
|
||||
graph = capture_up_projection_graph(launches)
|
||||
result["dynamic"] = benchmark_up_projection_graph(
|
||||
graph,
|
||||
operations_per_replay=len(launches),
|
||||
warmup_replays=args.warmup_replays,
|
||||
samples=args.samples,
|
||||
)
|
||||
check_up_projection_output(
|
||||
mailbox[0, :num_tokens, :shard_size],
|
||||
latent,
|
||||
weights[-1],
|
||||
shared_shard[:num_tokens],
|
||||
)
|
||||
if use_skinny:
|
||||
configs = args.skinny_config or [
|
||||
fused_add_multicast_skinny_gemm.config_for_m(
|
||||
num_tokens,
|
||||
shard_size,
|
||||
)
|
||||
]
|
||||
skinny_results = []
|
||||
for config in configs:
|
||||
skinny_kernel = SkinnyKernel(num_tokens, shard_size, config)
|
||||
|
||||
def launch_skinny(
|
||||
latent: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
shared_shard: torch.Tensor,
|
||||
*,
|
||||
skinny_kernel: SkinnyKernel = skinny_kernel,
|
||||
num_tokens: int = num_tokens,
|
||||
) -> None:
|
||||
skinny_kernel.launch(
|
||||
latent,
|
||||
weight,
|
||||
shared_shard[:num_tokens],
|
||||
mailbox,
|
||||
)
|
||||
|
||||
launches = make_up_projection_launches(
|
||||
launch_skinny,
|
||||
latent,
|
||||
weights,
|
||||
shared_shard,
|
||||
)
|
||||
graph = capture_up_projection_graph(launches)
|
||||
timing = benchmark_up_projection_graph(
|
||||
graph,
|
||||
operations_per_replay=len(launches),
|
||||
warmup_replays=args.warmup_replays,
|
||||
samples=args.samples,
|
||||
)
|
||||
check_up_projection_output(
|
||||
mailbox[0, :num_tokens, :shard_size],
|
||||
latent,
|
||||
weights[-1],
|
||||
shared_shard[:num_tokens],
|
||||
)
|
||||
skinny_results.append(
|
||||
{
|
||||
"config": asdict(config),
|
||||
**timing,
|
||||
}
|
||||
)
|
||||
result["skinny"] = skinny_results
|
||||
results.append(result)
|
||||
|
||||
properties = torch.cuda.get_device_properties(device)
|
||||
report = {
|
||||
"scope": "up-projection",
|
||||
"device": properties.name,
|
||||
"compute_capability": list(torch.cuda.get_device_capability(device)),
|
||||
"tp_size": args.tp_size,
|
||||
"shard_size": shard_size,
|
||||
"weight_count": weight_count,
|
||||
"cache_multiplier": args.cache_multiplier,
|
||||
"warmup_replays": args.warmup_replays,
|
||||
"samples": args.samples,
|
||||
"results": results,
|
||||
}
|
||||
rendered = json.dumps(report, indent=2)
|
||||
print(rendered, flush=True)
|
||||
if args.output is not None:
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(rendered + "\n", encoding="utf-8")
|
||||
|
||||
|
||||
def capture_tail_graph(
|
||||
operation: Callable[[], torch.Tensor],
|
||||
cpu_group: dist.ProcessGroup,
|
||||
) -> tuple[torch.cuda.CUDAGraph, torch.Tensor]:
|
||||
for _ in range(3):
|
||||
dist.barrier(group=cpu_group)
|
||||
output = operation()
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
dist.barrier(group=cpu_group)
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
output = operation()
|
||||
torch.accelerator.synchronize()
|
||||
return graph, output
|
||||
|
||||
|
||||
def benchmark_tail_graph(
|
||||
graph: torch.cuda.CUDAGraph,
|
||||
*,
|
||||
warmup_replays: int,
|
||||
samples: int,
|
||||
device_group: dist.ProcessGroup,
|
||||
cpu_group: dist.ProcessGroup,
|
||||
) -> dict[str, Any]:
|
||||
for _ in range(warmup_replays):
|
||||
graph.replay()
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
dist.barrier(group=cpu_group)
|
||||
starts = [torch.cuda.Event(enable_timing=True) for _ in range(samples + 1)]
|
||||
ends = [torch.cuda.Event(enable_timing=True) for _ in range(samples + 1)]
|
||||
for start, end in zip(starts, ends):
|
||||
start.record()
|
||||
graph.replay()
|
||||
end.record()
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
samples_us = torch.tensor(
|
||||
[start.elapsed_time(end) * 1000.0 for start, end in zip(starts, ends)],
|
||||
dtype=torch.float64,
|
||||
device=torch.accelerator.current_device_index(),
|
||||
)
|
||||
dist.all_reduce(samples_us, op=dist.ReduceOp.MAX, group=device_group)
|
||||
return summarize(samples_us[1:].tolist())
|
||||
|
||||
|
||||
def make_inputs(
|
||||
num_tokens: int,
|
||||
rank: int,
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
torch.manual_seed(20260726 + 100 * num_tokens + rank)
|
||||
routed = torch.randn(
|
||||
(num_tokens, LATENT_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
).mul_(0.01)
|
||||
shared = torch.randn(
|
||||
(num_tokens, HIDDEN_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
return routed, shared
|
||||
|
||||
|
||||
def make_reference(
|
||||
routed: torch.Tensor,
|
||||
shared: torch.Tensor,
|
||||
rms_weight: torch.Tensor,
|
||||
up_weight: torch.Tensor,
|
||||
device_group: dist.ProcessGroup,
|
||||
) -> Callable[[], torch.Tensor]:
|
||||
routed_workspace = torch.empty_like(routed)
|
||||
shared_workspace = torch.empty_like(shared)
|
||||
|
||||
def reference() -> torch.Tensor:
|
||||
routed_workspace.copy_(routed)
|
||||
dist.all_reduce(routed_workspace, group=device_group)
|
||||
normalized = F.rms_norm(
|
||||
routed_workspace,
|
||||
(LATENT_SIZE,),
|
||||
rms_weight,
|
||||
RMS_EPS,
|
||||
)
|
||||
projected = F.linear(normalized, up_weight)
|
||||
shared_workspace.copy_(shared)
|
||||
dist.all_reduce(shared_workspace, group=device_group)
|
||||
return projected.add(shared_workspace)
|
||||
|
||||
return reference
|
||||
|
||||
|
||||
def check_fused_output(
|
||||
fused_output: torch.Tensor,
|
||||
reference: Callable[[], torch.Tensor],
|
||||
cpu_group: dist.ProcessGroup,
|
||||
) -> None:
|
||||
dist.barrier(group=cpu_group)
|
||||
expected = reference()
|
||||
torch.testing.assert_close(fused_output, expected, atol=8e-2, rtol=3e-2)
|
||||
|
||||
|
||||
def benchmark_whole_tail(args: argparse.Namespace) -> None:
|
||||
if any(not 1 <= num_tokens <= 16 for num_tokens in args.num_tokens):
|
||||
raise ValueError("--num-tokens values must be in [1, 16]")
|
||||
if args.warmup_replays < 0 or args.samples <= 0:
|
||||
raise ValueError("warmup replays must be nonnegative and samples positive")
|
||||
if args.skinny_max_num_tokens is not None and any(
|
||||
not 0 <= cutoff <= 8 for cutoff in args.skinny_max_num_tokens
|
||||
):
|
||||
raise ValueError("--skinny-max-num-tokens must be in [0, 8]")
|
||||
skinny_configs = dict(args.skinny_config or ())
|
||||
if len(skinny_configs) != len(args.skinny_config or ()):
|
||||
raise ValueError("--skinny-config must not repeat an M value")
|
||||
if any(not 1 <= num_tokens <= 8 for num_tokens in skinny_configs):
|
||||
raise ValueError("--skinny-config M values must be in [1, 8]")
|
||||
if not {"RANK", "WORLD_SIZE", "LOCAL_RANK"} <= os.environ.keys():
|
||||
raise RuntimeError("launch this benchmark with torchrun")
|
||||
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
local_rank = int(os.environ["LOCAL_RANK"])
|
||||
device = torch.device("cuda", local_rank)
|
||||
torch.accelerator.set_device_index(device)
|
||||
init_distributed_environment()
|
||||
if world_size > 8:
|
||||
set_custom_all_reduce(False)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
device_group = get_tp_group().device_group
|
||||
cpu_group = dist.new_group(backend="gloo")
|
||||
|
||||
if torch.cuda.get_device_capability(device)[0] != 10:
|
||||
raise RuntimeError("Kimi K3 latent-MoE tail requires SM100")
|
||||
|
||||
torch.manual_seed(20260726)
|
||||
rms_weight = 1 + 0.1 * torch.randn(
|
||||
LATENT_SIZE,
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
up_weight = (
|
||||
torch.randn(
|
||||
(HIDDEN_SIZE, LATENT_SIZE),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
/ LATENT_SIZE**0.5
|
||||
)
|
||||
|
||||
use_reference = args.backend in ("reference", "both")
|
||||
use_fused = args.backend in ("fused", "both")
|
||||
fused_ops = []
|
||||
if use_fused:
|
||||
production_config_for_m = fused_add_multicast_skinny_gemm.config_for_m
|
||||
|
||||
def config_for_m(
|
||||
num_rows: int,
|
||||
shard_dim: int = 896,
|
||||
) -> fused_add_multicast_skinny_gemm.SkinnyConfig:
|
||||
config = skinny_configs.get(num_rows)
|
||||
if config is not None:
|
||||
return config
|
||||
return production_config_for_m(num_rows, shard_dim)
|
||||
|
||||
fused_add_multicast_skinny_gemm.config_for_m = config_for_m
|
||||
cutoffs = args.skinny_max_num_tokens or [latent_moe_tail._SKINNY_MAX_NUM_TOKENS]
|
||||
for cutoff in cutoffs:
|
||||
latent_moe_tail._SKINNY_MAX_NUM_TOKENS = cutoff
|
||||
latent_moe_tail.KimiK3LatentMoETailOp._instances.clear()
|
||||
fused_ops.append(
|
||||
(
|
||||
cutoff,
|
||||
latent_moe_tail.KimiK3LatentMoETailOp.initialize(
|
||||
hidden_size=HIDDEN_SIZE,
|
||||
latent_size=LATENT_SIZE,
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
rms_eps=RMS_EPS,
|
||||
),
|
||||
)
|
||||
)
|
||||
cutedsl_warmup()
|
||||
|
||||
results = []
|
||||
for num_tokens in args.num_tokens:
|
||||
routed, shared = make_inputs(num_tokens, rank, device)
|
||||
reference = make_reference(
|
||||
routed,
|
||||
shared,
|
||||
rms_weight,
|
||||
up_weight,
|
||||
device_group,
|
||||
)
|
||||
result: dict[str, Any] = {"num_tokens": num_tokens}
|
||||
if use_reference:
|
||||
reference_graph, _ = capture_tail_graph(reference, cpu_group)
|
||||
result["reference"] = benchmark_tail_graph(
|
||||
reference_graph,
|
||||
warmup_replays=args.warmup_replays,
|
||||
samples=args.samples,
|
||||
device_group=device_group,
|
||||
cpu_group=cpu_group,
|
||||
)
|
||||
for cutoff, fused_op in fused_ops:
|
||||
|
||||
def fused(
|
||||
routed: torch.Tensor = routed,
|
||||
shared: torch.Tensor = shared,
|
||||
fused_op: latent_moe_tail.KimiK3LatentMoETailOp = fused_op,
|
||||
) -> torch.Tensor:
|
||||
return fused_op(routed, shared, rms_weight, up_weight)
|
||||
|
||||
fused_graph, fused_output = capture_tail_graph(fused, cpu_group)
|
||||
fused_key = "fused" if len(fused_ops) == 1 else f"fused_skinny_max_{cutoff}"
|
||||
result[fused_key] = benchmark_tail_graph(
|
||||
fused_graph,
|
||||
warmup_replays=args.warmup_replays,
|
||||
samples=args.samples,
|
||||
device_group=device_group,
|
||||
cpu_group=cpu_group,
|
||||
)
|
||||
check_fused_output(fused_output, reference, cpu_group)
|
||||
if "reference" in result:
|
||||
speedup = (
|
||||
result["reference"]["median_us"] / result[fused_key]["median_us"]
|
||||
)
|
||||
if len(fused_ops) == 1:
|
||||
result["speedup"] = speedup
|
||||
else:
|
||||
result[f"{fused_key}_speedup"] = speedup
|
||||
results.append(result)
|
||||
|
||||
properties = torch.cuda.get_device_properties(device)
|
||||
report = {
|
||||
"scope": "whole-tail",
|
||||
"device": properties.name,
|
||||
"compute_capability": list(torch.cuda.get_device_capability(device)),
|
||||
"world_size": world_size,
|
||||
"torch_version": torch.__version__,
|
||||
"cuda_version": torch.version.cuda,
|
||||
"warmup_replays": args.warmup_replays,
|
||||
"samples": args.samples,
|
||||
"skinny_max_num_tokens": [cutoff for cutoff, _ in fused_ops],
|
||||
"skinny_configs": {
|
||||
str(num_tokens): asdict(config)
|
||||
for num_tokens, config in skinny_configs.items()
|
||||
},
|
||||
"timing_scope": {
|
||||
"reference": (
|
||||
"two input copies, two AllReduces, RMSNorm, full replicated "
|
||||
"up-projection GEMM, and final add"
|
||||
),
|
||||
"fused": (
|
||||
"routed AllReduce/RMSNorm plus shared ReduceScatter, sharded "
|
||||
"up-projection/multicast, and Lamport copy"
|
||||
),
|
||||
},
|
||||
"results": results,
|
||||
}
|
||||
if rank == 0:
|
||||
rendered = json.dumps(report, indent=2)
|
||||
print(rendered, flush=True)
|
||||
if args.output is not None:
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(rendered + "\n", encoding="utf-8")
|
||||
|
||||
dist.barrier(group=cpu_group)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.scope == "up-projection":
|
||||
benchmark_up_projection(args)
|
||||
return
|
||||
|
||||
from vllm.config import VllmConfig, set_current_vllm_config
|
||||
|
||||
with set_current_vllm_config(VllmConfig()):
|
||||
benchmark_whole_tail(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,239 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import statistics
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
import vllm._custom_ops as ops
|
||||
from vllm.distributed.device_communicators.custom_all_reduce import CustomAllreduce
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--tokens", type=int, nargs="+", default=[8, 32, 128, 1024])
|
||||
parser.add_argument("--hidden-size", type=int, default=7168)
|
||||
parser.add_argument("--graph-repeats", type=int, default=20)
|
||||
parser.add_argument("--warmup-replays", type=int, default=5)
|
||||
parser.add_argument("--samples", type=int, default=15)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def capture_graph(op: Callable[[], None], repeats: int) -> torch.cuda.CUDAGraph:
|
||||
stream = torch.cuda.Stream()
|
||||
stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.cuda.stream(stream):
|
||||
for _ in range(3):
|
||||
op()
|
||||
stream.synchronize()
|
||||
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph, stream=stream):
|
||||
for _ in range(repeats):
|
||||
op()
|
||||
torch.cuda.current_stream().wait_stream(stream)
|
||||
return graph
|
||||
|
||||
|
||||
def max_rank_graph_time(
|
||||
graph: torch.cuda.CUDAGraph,
|
||||
repeats: int,
|
||||
warmup_replays: int,
|
||||
samples: int,
|
||||
device_group: dist.ProcessGroup,
|
||||
cpu_group: dist.ProcessGroup,
|
||||
) -> float:
|
||||
for _ in range(warmup_replays):
|
||||
graph.replay()
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
timings = []
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
for _ in range(samples):
|
||||
dist.barrier(group=cpu_group)
|
||||
start.record()
|
||||
graph.replay()
|
||||
end.record()
|
||||
end.synchronize()
|
||||
elapsed = torch.tensor(
|
||||
start.elapsed_time(end) / repeats,
|
||||
dtype=torch.float64,
|
||||
device=torch.accelerator.current_device_index(),
|
||||
)
|
||||
dist.all_reduce(elapsed, op=dist.ReduceOp.MAX, group=device_group)
|
||||
timings.append(elapsed.item())
|
||||
return statistics.median(timings)
|
||||
|
||||
|
||||
def check_outputs(
|
||||
comm: CustomAllreduce,
|
||||
local: torch.Tensor,
|
||||
reduce_input: torch.Tensor,
|
||||
device_group: dist.ProcessGroup,
|
||||
) -> None:
|
||||
expected_gather = torch.empty(
|
||||
(local.shape[0] * dist.get_world_size(), local.shape[1]),
|
||||
dtype=local.dtype,
|
||||
device=local.device,
|
||||
)
|
||||
dist.all_gather_into_tensor(expected_gather, local, group=device_group)
|
||||
gathered = comm.custom_all_gather(local)
|
||||
assert gathered is not None
|
||||
torch.testing.assert_close(gathered, expected_gather)
|
||||
|
||||
expected_scatter = torch.empty_like(local)
|
||||
dist.reduce_scatter_tensor(
|
||||
expected_scatter,
|
||||
reduce_input.clone(),
|
||||
group=device_group,
|
||||
)
|
||||
scattered = comm.custom_reduce_scatter(reduce_input)
|
||||
assert scattered is not None
|
||||
torch.testing.assert_close(scattered, expected_scatter)
|
||||
|
||||
|
||||
def benchmark_shape(
|
||||
comm: CustomAllreduce,
|
||||
global_tokens: int,
|
||||
hidden_size: int,
|
||||
graph_repeats: int,
|
||||
warmup_replays: int,
|
||||
samples: int,
|
||||
device_group: dist.ProcessGroup,
|
||||
cpu_group: dist.ProcessGroup,
|
||||
) -> dict[str, float | int]:
|
||||
world_size = dist.get_world_size()
|
||||
rank = dist.get_rank()
|
||||
padded_tokens = (global_tokens + world_size - 1) // world_size * world_size
|
||||
local_tokens = padded_tokens // world_size
|
||||
local = torch.full(
|
||||
(local_tokens, hidden_size),
|
||||
rank + 1,
|
||||
dtype=torch.bfloat16,
|
||||
device=torch.accelerator.current_device_index(),
|
||||
)
|
||||
reduce_input = torch.full(
|
||||
(padded_tokens, hidden_size),
|
||||
rank + 1,
|
||||
dtype=torch.bfloat16,
|
||||
device=local.device,
|
||||
)
|
||||
check_outputs(comm, local, reduce_input, device_group)
|
||||
|
||||
custom_gather_out = torch.empty(
|
||||
(padded_tokens, hidden_size),
|
||||
dtype=local.dtype,
|
||||
device=local.device,
|
||||
)
|
||||
custom_scatter_out = torch.empty_like(local)
|
||||
nccl_gather_out = torch.empty_like(custom_gather_out)
|
||||
nccl_scatter_out = torch.empty_like(local)
|
||||
|
||||
def custom_ag() -> None:
|
||||
ops.mnnvl_lamport_all_gather(
|
||||
comm._ptr,
|
||||
local,
|
||||
custom_gather_out,
|
||||
comm.mnnvl_lamport_ag_local_ptr,
|
||||
comm.mnnvl_lamport_ag_multicast_ptr,
|
||||
comm.mnnvl_lamport_ag_epoch_ptr,
|
||||
comm.mnnvl_buffer_size,
|
||||
)
|
||||
|
||||
def custom_rs() -> None:
|
||||
ops.mnnvl_lamport_reduce_scatter(
|
||||
comm._ptr,
|
||||
reduce_input,
|
||||
custom_scatter_out,
|
||||
comm.mnnvl_lamport_rs_local_ptr,
|
||||
comm.mnnvl_lamport_rs_epoch_ptr,
|
||||
comm.mnnvl_buffer_size,
|
||||
)
|
||||
|
||||
def nccl_ag() -> None:
|
||||
dist.all_gather_into_tensor(nccl_gather_out, local, group=device_group)
|
||||
|
||||
def nccl_rs() -> None:
|
||||
dist.reduce_scatter_tensor(
|
||||
nccl_scatter_out,
|
||||
reduce_input,
|
||||
group=device_group,
|
||||
)
|
||||
|
||||
graphs = {
|
||||
"custom_ag_us": capture_graph(custom_ag, graph_repeats),
|
||||
"nccl_ag_us": capture_graph(nccl_ag, graph_repeats),
|
||||
"custom_rs_us": capture_graph(custom_rs, graph_repeats),
|
||||
"nccl_rs_us": capture_graph(nccl_rs, graph_repeats),
|
||||
}
|
||||
times = {
|
||||
name: max_rank_graph_time(
|
||||
graph,
|
||||
graph_repeats,
|
||||
warmup_replays,
|
||||
samples,
|
||||
device_group,
|
||||
cpu_group,
|
||||
)
|
||||
* 1000
|
||||
for name, graph in graphs.items()
|
||||
}
|
||||
torch.testing.assert_close(custom_gather_out, nccl_gather_out)
|
||||
torch.testing.assert_close(custom_scatter_out, nccl_scatter_out)
|
||||
return {
|
||||
"global_tokens": global_tokens,
|
||||
"padded_tokens": padded_tokens,
|
||||
"local_bytes": local.nbytes,
|
||||
"full_bytes": reduce_input.nbytes,
|
||||
**times,
|
||||
"ag_speedup": times["nccl_ag_us"] / times["custom_ag_us"],
|
||||
"rs_speedup": times["nccl_rs_us"] / times["custom_rs_us"],
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
local_rank = int(os.environ["LOCAL_RANK"])
|
||||
torch.accelerator.set_device_index(local_rank)
|
||||
dist.init_process_group("nccl")
|
||||
device_group = dist.group.WORLD
|
||||
cpu_group = dist.new_group(backend="gloo")
|
||||
|
||||
comm = CustomAllreduce(
|
||||
group=cpu_group,
|
||||
device=torch.device("cuda", local_rank),
|
||||
)
|
||||
assert not comm.disabled
|
||||
assert comm.world_size == 16
|
||||
assert comm.mnnvl_only
|
||||
assert comm.mnnvl_multicast_ptr
|
||||
|
||||
results = [
|
||||
benchmark_shape(
|
||||
comm,
|
||||
tokens,
|
||||
args.hidden_size,
|
||||
args.graph_repeats,
|
||||
args.warmup_replays,
|
||||
args.samples,
|
||||
device_group,
|
||||
cpu_group,
|
||||
)
|
||||
for tokens in args.tokens
|
||||
]
|
||||
if dist.get_rank() == 0:
|
||||
print(json.dumps(results, indent=2), flush=True)
|
||||
|
||||
comm.close()
|
||||
dist.destroy_process_group(cpu_group)
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,201 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Benchmark the RDNAHybridW4A16LinearKernel across decode and prefill shapes.
|
||||
|
||||
Usage:
|
||||
python benchmark_int4_gemm.py
|
||||
python benchmark_int4_gemm.py --models Qwen/Qwen3-4B
|
||||
python benchmark_int4_gemm.py --group-size 128
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import copy
|
||||
import itertools
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.triton_utils import triton
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Weight shapes: [K, N], TP_SPLIT_DIM
|
||||
# ---------------------------------------------------------------------------
|
||||
WEIGHT_SHAPES = {
|
||||
"Qwen/Qwen3-4B": [
|
||||
([2560, 3840], 1), # qkv_proj
|
||||
([2560, 2560], 0), # o_proj
|
||||
([2560, 19456], 1), # gate_up_proj
|
||||
([9728, 2560], 0), # down_proj
|
||||
],
|
||||
"Qwen/Qwen2.5-7B-Instruct": [
|
||||
([3584, 4608], 1),
|
||||
([3584, 3584], 0),
|
||||
([3584, 37888], 1),
|
||||
([18944, 3584], 0),
|
||||
],
|
||||
"trymirai/SmolLM2-1.7B-Instruct-AWQ": [
|
||||
([2048, 6144], 1), # qkv_proj
|
||||
([2048, 2048], 0), # o_proj
|
||||
([2048, 16384], 1), # gate_up_proj
|
||||
([8192, 2048], 0), # down_proj
|
||||
],
|
||||
"RedHatAI/Qwen3-8B-quantized.w4a16": [
|
||||
([4096, 6144], 1), # qkv_proj
|
||||
([4096, 4096], 0), # o_proj
|
||||
([4096, 24576], 1), # gate_up_proj
|
||||
([12288, 4096], 0), # down_proj
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Weight packing
|
||||
# ---------------------------------------------------------------------------
|
||||
def prepare_hybrid_weights(K, N, group_size, device="cuda"):
|
||||
"""Create random weights for benchmarking.
|
||||
|
||||
Returns (w_q_skinny, w_s_skinny, w_fp16, w_zp). The triton path derives
|
||||
its int32 view from w_q_skinny, so no separate int32 buffer is returned.
|
||||
"""
|
||||
num_groups = K // group_size
|
||||
|
||||
# Random packed weights — actual values don't matter for throughput
|
||||
w_q_skinny_i32 = torch.randint(
|
||||
0, 2**31, (N, K // 8), dtype=torch.int32, device=device
|
||||
)
|
||||
w_q_skinny = w_q_skinny_i32.view(torch.int8).contiguous()
|
||||
w_s_skinny = torch.randn(N, num_groups, dtype=torch.float16, device=device) * 0.01
|
||||
|
||||
# Raw per-group zero-points for asymmetric benchmarks
|
||||
w_zp = torch.randint(0, 16, (N, num_groups), dtype=torch.int32, device=device).to(
|
||||
torch.float16
|
||||
)
|
||||
|
||||
# FP16 baseline for F.linear
|
||||
w_fp16 = torch.randn(N, K, dtype=torch.float16, device=device) * 0.01
|
||||
|
||||
return w_q_skinny, w_s_skinny, w_fp16, w_zp
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Benchmark
|
||||
# ---------------------------------------------------------------------------
|
||||
PROVIDERS = ["torch-fp16", "hybrid-w4a16", "hybrid-w4a16-zp"]
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["batch_size"],
|
||||
x_vals=[1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096],
|
||||
x_log=False,
|
||||
line_arg="provider",
|
||||
line_vals=PROVIDERS,
|
||||
line_names=PROVIDERS,
|
||||
ylabel="TFLOP/s (larger is better)",
|
||||
plot_name="FP16 vs Hybrid W4A16",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(batch_size, provider, N, K, group_size, weights):
|
||||
M = batch_size
|
||||
device = "cuda"
|
||||
dtype = torch.float16
|
||||
a = torch.randn((M, K), device=device, dtype=dtype)
|
||||
|
||||
quantiles = [0.5, 0.2, 0.8]
|
||||
|
||||
if provider == "torch-fp16":
|
||||
w_fp16 = weights["w_fp16"]
|
||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
||||
lambda: torch.nn.functional.linear(a, w_fp16),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
elif provider in ("hybrid-w4a16", "hybrid-w4a16-zp"):
|
||||
from vllm.model_executor.kernels.linear.mixed_precision import (
|
||||
rdna_hybrid_w4a16 as _k,
|
||||
)
|
||||
|
||||
_rdna_hybrid_w4a16_apply_impl = _k._rdna_hybrid_w4a16_apply_impl
|
||||
from vllm.utils.platform_utils import num_compute_units
|
||||
|
||||
w = weights
|
||||
cu_count = num_compute_units()
|
||||
use_zp = provider == "hybrid-w4a16-zp"
|
||||
|
||||
def run():
|
||||
return _rdna_hybrid_w4a16_apply_impl(
|
||||
a,
|
||||
w["w_q_skinny"],
|
||||
w["w_s_skinny"],
|
||||
w["w_zp"] if use_zp else None,
|
||||
None, # bias
|
||||
cu_count,
|
||||
group_size,
|
||||
)
|
||||
|
||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
||||
run,
|
||||
quantiles=quantiles,
|
||||
)
|
||||
else:
|
||||
return 0.0, 0.0, 0.0
|
||||
|
||||
to_tflops = lambda t_ms: (2 * M * N * K) * 1e-12 / (t_ms * 1e-3)
|
||||
return to_tflops(ms), to_tflops(max_ms), to_tflops(min_ms)
|
||||
|
||||
|
||||
def prepare_shapes(args):
|
||||
KN_model_names = []
|
||||
for model, tp_size in itertools.product(args.models, args.tp_sizes):
|
||||
for KN, tp_dim in copy.deepcopy(WEIGHT_SHAPES[model]):
|
||||
KN[tp_dim] //= tp_size
|
||||
KN.append(model)
|
||||
KN_model_names.append(KN)
|
||||
return KN_model_names
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Benchmark RDNAHybridW4A16LinearKernel"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--models",
|
||||
nargs="+",
|
||||
type=str,
|
||||
default=["Qwen/Qwen3-4B"],
|
||||
choices=list(WEIGHT_SHAPES.keys()),
|
||||
)
|
||||
parser.add_argument("--tp-sizes", nargs="+", type=int, default=[1])
|
||||
parser.add_argument("--group-size", type=int, default=128)
|
||||
parser.add_argument("--save-path", type=str, default=None)
|
||||
args = parser.parse_args()
|
||||
|
||||
for K, N, model in prepare_shapes(args):
|
||||
group_size = args.group_size
|
||||
print(f"\n{'=' * 70}")
|
||||
print(f"{model}, N={N} K={K}, group_size={group_size}")
|
||||
print(f"{'=' * 70}")
|
||||
|
||||
w_q_skinny, w_s_skinny, w_fp16, w_zp = prepare_hybrid_weights(K, N, group_size)
|
||||
|
||||
weights = {
|
||||
"w_q_skinny": w_q_skinny,
|
||||
"w_s_skinny": w_s_skinny,
|
||||
"w_fp16": w_fp16,
|
||||
"w_zp": w_zp,
|
||||
}
|
||||
|
||||
save_path = args.save_path or f"bench_int4_res_n{N}_k{K}"
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
benchmark.run(
|
||||
print_data=True,
|
||||
show_plots=False,
|
||||
save_path=save_path,
|
||||
N=N,
|
||||
K=K,
|
||||
group_size=group_size,
|
||||
weights=weights,
|
||||
)
|
||||
|
||||
print("\nBenchmark finished!")
|
||||
@@ -0,0 +1,267 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""End-to-end autoregressive decode benchmark: ReplaySSM vs the standard SSM kernel.
|
||||
|
||||
Loads a hybrid Mamba2 model, replicates one prompt across the batch, and times a
|
||||
long greedy decode (CUDA graphs on) once with the standard kernel and once with
|
||||
ReplaySSM, then reports the per-step / throughput speedup. The two modes run in
|
||||
separate subprocesses so each gets a clean CUDA context.
|
||||
|
||||
The FlashInfer FP4-MoE autotuner is disabled by default (it is unstable under
|
||||
CUDA-graph capture on the pre-release Blackwell FP4 path); pass
|
||||
--no-disable-flashinfer-autotune for non-FP4 models.
|
||||
|
||||
Examples:
|
||||
python e2e_decode_speedup.py --model-id nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16
|
||||
python e2e_decode_speedup.py --dtype auto --buffer-len 16 \
|
||||
--model-id nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4 # B300 NVFP4
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
DEFAULT_PROMPT = "My cat wrote all this CUDA code for a new language model and"
|
||||
|
||||
MODE_LABEL = {"standard": "standard", "replayssm": "ReplaySSM"}
|
||||
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser(
|
||||
description="E2E decode speedup: ReplaySSM vs the standard SSM kernel."
|
||||
)
|
||||
p.add_argument("--model-id", default="nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16")
|
||||
p.add_argument("--prompt", default=DEFAULT_PROMPT)
|
||||
p.add_argument("--batch-size", type=int, default=256)
|
||||
p.add_argument("--num-steps", type=int, default=1000)
|
||||
p.add_argument("--warmup-steps", type=int, default=128)
|
||||
p.add_argument("--repeats", type=int, default=1)
|
||||
p.add_argument(
|
||||
"--buffer-len", type=int, default=16, help="ReplaySSM input-buffer length."
|
||||
)
|
||||
p.add_argument(
|
||||
"--dtype",
|
||||
default="bfloat16",
|
||||
choices=["bfloat16", "float16", "float32", "auto"],
|
||||
)
|
||||
p.add_argument("--gpu-memory-utilization", type=float, default=0.9)
|
||||
p.add_argument("--max-model-len", type=int, default=None)
|
||||
p.add_argument(
|
||||
"--disable-flashinfer-autotune",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Disable the FlashInfer FP4-MoE autotuner (default: on). "
|
||||
"It is unstable under CUDA-graph capture on the "
|
||||
"pre-release Blackwell FP4 path; pass "
|
||||
"--no-disable-flashinfer-autotune for non-FP4 models.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--mamba-ssm-cache-dtype",
|
||||
default="auto",
|
||||
choices=["auto", "float32", "float16", "bfloat16"],
|
||||
help="SSM state dtype (both modes). 'auto' = config-driven; "
|
||||
"'float32' = fp32 state, 'bfloat16' = s16 state.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--baseline-ssm-config",
|
||||
default="",
|
||||
help="Pin the STANDARD baseline's SSM launch config as "
|
||||
"'bsm,nw' via override_ssm_config (forces the in-process "
|
||||
"engine so the override reaches the kernel). Empty = off.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--worker",
|
||||
choices=["standard", "replayssm"],
|
||||
default=None,
|
||||
help=argparse.SUPPRESS,
|
||||
)
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def resolve_max_model_len(args) -> int:
|
||||
if args.max_model_len is not None:
|
||||
return args.max_model_len
|
||||
return args.num_steps + 256
|
||||
|
||||
|
||||
def run_worker(args):
|
||||
# override_ssm_config is a module global; it only reaches the model if the
|
||||
# engine runs in-process (default V1 spawns a separate EngineCore). Force it.
|
||||
if args.worker == "standard" and args.baseline_ssm_config:
|
||||
os.environ["VLLM_ENABLE_V1_MULTIPROCESSING"] = "0"
|
||||
|
||||
import torch
|
||||
|
||||
from vllm import LLM, SamplingParams
|
||||
|
||||
mode = args.worker
|
||||
max_model_len = resolve_max_model_len(args)
|
||||
|
||||
llm_kwargs = dict(
|
||||
model=args.model_id,
|
||||
tensor_parallel_size=1,
|
||||
dtype=args.dtype,
|
||||
max_model_len=max_model_len,
|
||||
trust_remote_code=True,
|
||||
enable_prefix_caching=False,
|
||||
enable_chunked_prefill=False,
|
||||
max_num_seqs=args.batch_size,
|
||||
max_num_batched_tokens=max(max_model_len, args.batch_size * 64),
|
||||
enforce_eager=False,
|
||||
disable_log_stats=True,
|
||||
gpu_memory_utilization=args.gpu_memory_utilization,
|
||||
# SSM state dtype (applies to both standard and ReplaySSM).
|
||||
mamba_ssm_cache_dtype=args.mamba_ssm_cache_dtype,
|
||||
)
|
||||
if args.disable_flashinfer_autotune:
|
||||
# FP4-MoE autotuner is unstable under CUDA-graph capture on Blackwell;
|
||||
# re-enable (--no-disable-flashinfer-autotune) only for non-FP4 models.
|
||||
llm_kwargs["kernel_config"] = {"enable_flashinfer_autotune": False}
|
||||
if mode == "replayssm":
|
||||
llm_kwargs.update(use_replayssm=True, replayssm_buffer_len=args.buffer_len)
|
||||
|
||||
_ssm_cm = None
|
||||
if mode == "standard" and args.baseline_ssm_config:
|
||||
from vllm.model_executor.layers.mamba.ops.mamba_ssm import override_ssm_config
|
||||
|
||||
_bsm, _nw = (int(x) for x in args.baseline_ssm_config.split(","))
|
||||
_ssm_cm = override_ssm_config((_bsm, _nw))
|
||||
_ssm_cm.__enter__() # active through LLM() graph capture + decode
|
||||
print(
|
||||
f"[{mode}] override_ssm_config -> (BLOCK_SIZE_M={_bsm}, num_warps={_nw})",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
llm = LLM(**llm_kwargs)
|
||||
prompts = [args.prompt] * args.batch_size
|
||||
|
||||
def timed_generate(n_tokens):
|
||||
sp = SamplingParams(
|
||||
n=1,
|
||||
temperature=0.0,
|
||||
ignore_eos=True,
|
||||
min_tokens=n_tokens,
|
||||
max_tokens=n_tokens,
|
||||
)
|
||||
if torch.accelerator.is_available():
|
||||
torch.accelerator.synchronize()
|
||||
t0 = time.perf_counter()
|
||||
outs = llm.generate(prompts, sp, use_tqdm=False)
|
||||
if torch.accelerator.is_available():
|
||||
torch.accelerator.synchronize()
|
||||
elapsed = time.perf_counter() - t0
|
||||
produced = min(len(o.outputs[0].token_ids) for o in outs)
|
||||
assert produced == n_tokens, f"expected {n_tokens} tokens, got {produced}"
|
||||
return elapsed
|
||||
|
||||
timed_generate(args.warmup_steps)
|
||||
|
||||
best = None
|
||||
for _ in range(args.repeats):
|
||||
elapsed = timed_generate(args.num_steps)
|
||||
tok_s = args.batch_size * args.num_steps / elapsed
|
||||
per_step_ms = elapsed / args.num_steps * 1e3
|
||||
print(
|
||||
f"[{mode}] {elapsed:.3f}s {tok_s:,.0f} tok/s {per_step_ms:.3f} ms/step",
|
||||
flush=True,
|
||||
)
|
||||
if best is None or elapsed < best["elapsed_s"]:
|
||||
best = {
|
||||
"mode": mode,
|
||||
"elapsed_s": elapsed,
|
||||
"tok_s": tok_s,
|
||||
"per_step_ms": per_step_ms,
|
||||
}
|
||||
|
||||
print("RESULT_JSON " + json.dumps(best), flush=True)
|
||||
if _ssm_cm is not None:
|
||||
_ssm_cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def run_one_mode(args, mode) -> dict:
|
||||
cmd = [
|
||||
sys.executable,
|
||||
__file__,
|
||||
"--worker",
|
||||
mode,
|
||||
"--model-id",
|
||||
args.model_id,
|
||||
"--prompt",
|
||||
args.prompt,
|
||||
"--batch-size",
|
||||
str(args.batch_size),
|
||||
"--num-steps",
|
||||
str(args.num_steps),
|
||||
"--warmup-steps",
|
||||
str(args.warmup_steps),
|
||||
"--repeats",
|
||||
str(args.repeats),
|
||||
"--buffer-len",
|
||||
str(args.buffer_len),
|
||||
"--dtype",
|
||||
args.dtype,
|
||||
"--gpu-memory-utilization",
|
||||
str(args.gpu_memory_utilization),
|
||||
"--mamba-ssm-cache-dtype",
|
||||
args.mamba_ssm_cache_dtype,
|
||||
"--baseline-ssm-config",
|
||||
args.baseline_ssm_config,
|
||||
]
|
||||
cmd.append(
|
||||
"--disable-flashinfer-autotune"
|
||||
if args.disable_flashinfer_autotune
|
||||
else "--no-disable-flashinfer-autotune"
|
||||
)
|
||||
if args.max_model_len is not None:
|
||||
cmd += ["--max-model-len", str(args.max_model_len)]
|
||||
|
||||
result = None
|
||||
proc = subprocess.Popen(
|
||||
cmd, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True, bufsize=1
|
||||
)
|
||||
for line in proc.stdout:
|
||||
sys.stdout.write(line)
|
||||
sys.stdout.flush()
|
||||
if line.startswith("RESULT_JSON "):
|
||||
result = json.loads(line[len("RESULT_JSON ") :])
|
||||
proc.wait()
|
||||
if proc.returncode != 0:
|
||||
raise RuntimeError(f"mode '{mode}' worker exited with {proc.returncode}")
|
||||
if result is None:
|
||||
raise RuntimeError(f"mode '{mode}' produced no RESULT_JSON line")
|
||||
return result
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
if args.worker is not None:
|
||||
run_worker(args)
|
||||
return
|
||||
|
||||
print(
|
||||
f"model={args.model_id} batch_size={args.batch_size} "
|
||||
f"steps={args.num_steps} buffer_len={args.buffer_len} dtype={args.dtype}"
|
||||
)
|
||||
|
||||
std = run_one_mode(args, "standard")
|
||||
fla = run_one_mode(args, "replayssm")
|
||||
speedup = std["per_step_ms"] / fla["per_step_ms"]
|
||||
|
||||
print()
|
||||
header = f"{'mode':<10}{'ms/step':>12}{'tok/s':>16}{'wall (s)':>12}"
|
||||
print(header)
|
||||
print("-" * len(header))
|
||||
for r in (std, fla):
|
||||
print(
|
||||
f"{MODE_LABEL[r['mode']]:<10}{r['per_step_ms']:>12.3f}"
|
||||
f"{r['tok_s']:>16,.0f}{r['elapsed_s']:>12.3f}"
|
||||
)
|
||||
print("-" * len(header))
|
||||
print(f"speedup (standard / ReplaySSM, per step): {speedup:.3f}x")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -430,6 +430,7 @@ set(VLLM_EXT_SRC
|
||||
"csrc/cpu/layernorm.cpp"
|
||||
"csrc/cpu/mla_decode.cpp"
|
||||
"csrc/cpu/pos_encoding.cpp"
|
||||
"csrc/cpu/mamba_cpu.cpp"
|
||||
"csrc/moe/dynamic_4bit_int_moe_cpu.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp")
|
||||
@@ -489,6 +490,7 @@ if (ENABLE_X86_ISA)
|
||||
"csrc/cpu/spec_decode_utils.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/dnnl_kernels.cpp"
|
||||
"csrc/cpu/mamba_cpu.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp"
|
||||
# TODO: Remove these files
|
||||
"csrc/cpu/activation.cpp"
|
||||
@@ -502,6 +504,7 @@ if (ENABLE_X86_ISA)
|
||||
"csrc/cpu/utils.cpp"
|
||||
"csrc/cpu/spec_decode_utils.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/mamba_cpu.cpp"
|
||||
"csrc/cpu/dnnl_kernels.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp"
|
||||
# TODO: Remove these files
|
||||
|
||||
@@ -28,9 +28,9 @@ if(DEEPGEMM_SRC_DIR)
|
||||
message(STATUS "DeepGEMM using local DEEPGEMM_SRC_DIR: ${deepgemm_SOURCE_DIR}")
|
||||
else()
|
||||
# Keep in sync with tools/install_deepgemm.sh
|
||||
set(_DEEPGEMM_UPSTREAM_REPO "https://github.com/deepseek-ai/DeepGEMM.git")
|
||||
set(_DEEPGEMM_UPSTREAM_REPO "git@github.com:Inferact/DeepGEMM.git")
|
||||
# NOTE: This is currently targeting nv-dev branch due to sm120 support
|
||||
set(_DEEPGEMM_UPSTREAM_TAG "a6b593d2826719dcf4892609af7b84ee23aaf32a")
|
||||
set(_DEEPGEMM_UPSTREAM_TAG "f5a76426fa084087169693fd0cd815223576d6e9")
|
||||
|
||||
set(_deepgemm_fc_root "${FETCHCONTENT_BASE_DIR}")
|
||||
if(NOT _deepgemm_fc_root)
|
||||
@@ -68,6 +68,9 @@ endif()
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8)
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.9)
|
||||
list(APPEND DEEPGEMM_SUPPORT_ARCHS "10.0f")
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.4)
|
||||
list(APPEND DEEPGEMM_SUPPORT_ARCHS "10.7f")
|
||||
endif()
|
||||
else()
|
||||
list(APPEND DEEPGEMM_SUPPORT_ARCHS "10.0a")
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
include(FetchContent)
|
||||
|
||||
if(DEFINED ENV{FLASH_KDA_SRC_DIR})
|
||||
set(FLASH_KDA_SRC_DIR $ENV{FLASH_KDA_SRC_DIR})
|
||||
endif()
|
||||
|
||||
if(FLASH_KDA_SRC_DIR)
|
||||
FetchContent_Declare(
|
||||
flashkda
|
||||
SOURCE_DIR ${FLASH_KDA_SRC_DIR}
|
||||
)
|
||||
else()
|
||||
FetchContent_Declare(
|
||||
flashkda
|
||||
GIT_REPOSITORY git@github.com:Inferact/FlashKDA.git
|
||||
GIT_TAG a3e42bbbece3bb38f7c426b880315294a336e82f
|
||||
GIT_PROGRESS TRUE
|
||||
GIT_SUBMODULES cutlass
|
||||
)
|
||||
endif()
|
||||
|
||||
FetchContent_MakeAvailable(flashkda)
|
||||
message(STATUS "FlashKDA is available at ${flashkda_SOURCE_DIR}")
|
||||
|
||||
set(FLASH_KDA_SUPPORT_ARCHS)
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.0)
|
||||
list(APPEND FLASH_KDA_SUPPORT_ARCHS "9.0a")
|
||||
endif()
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
list(APPEND FLASH_KDA_SUPPORT_ARCHS "10.0f" "12.0f")
|
||||
elseif(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.9)
|
||||
list(APPEND FLASH_KDA_SUPPORT_ARCHS "10.0a" "10.3a" "12.0a")
|
||||
endif()
|
||||
|
||||
cuda_archs_loose_intersection(
|
||||
FLASH_KDA_ARCHS "${FLASH_KDA_SUPPORT_ARCHS}" "${CUDA_ARCHS}")
|
||||
|
||||
if(FLASH_KDA_ARCHS)
|
||||
message(STATUS "FlashKDA CUDA architectures: ${FLASH_KDA_ARCHS}")
|
||||
|
||||
set(FLASH_KDA_SOURCES
|
||||
csrc/flashkda_registration.cpp
|
||||
${flashkda_SOURCE_DIR}/csrc/flash_kda.cpp
|
||||
${flashkda_SOURCE_DIR}/csrc/smxx/fwd_launch.cu)
|
||||
set(FLASH_KDA_INCLUDES
|
||||
${flashkda_SOURCE_DIR}/csrc
|
||||
${flashkda_SOURCE_DIR}/cutlass/include
|
||||
${flashkda_SOURCE_DIR}/cutlass/examples/common
|
||||
${flashkda_SOURCE_DIR}/cutlass/tools/util/include)
|
||||
|
||||
set_gencode_flags_for_srcs(
|
||||
SRCS "${FLASH_KDA_SOURCES}"
|
||||
CUDA_ARCHS "${FLASH_KDA_ARCHS}")
|
||||
|
||||
define_extension_target(
|
||||
_flashkda_C
|
||||
DESTINATION vllm
|
||||
LANGUAGE ${VLLM_GPU_LANG}
|
||||
SOURCES ${FLASH_KDA_SOURCES}
|
||||
COMPILE_FLAGS ${VLLM_GPU_FLAGS}
|
||||
ARCHITECTURES ${VLLM_GPU_ARCHES}
|
||||
INCLUDE_DIRECTORIES ${FLASH_KDA_INCLUDES}
|
||||
USE_SABI 3
|
||||
WITH_SOABI)
|
||||
|
||||
target_compile_options(_flashkda_C PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:-UPy_LIMITED_API --expt-relaxed-constexpr --expt-extended-lambda --use_fast_math -O3>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-UPy_LIMITED_API>)
|
||||
else()
|
||||
message(STATUS
|
||||
"FlashKDA will not compile: CUDA >=12.0 and a supported architecture "
|
||||
"(SM90, SM10x, or SM12x) are required")
|
||||
add_custom_target(_flashkda_C)
|
||||
endif()
|
||||
@@ -19,7 +19,7 @@ else()
|
||||
FetchContent_Declare(
|
||||
flashmla
|
||||
GIT_REPOSITORY https://github.com/vllm-project/FlashMLA
|
||||
GIT_TAG b70aff3d110a2b1a037e62eac295166b5143643a
|
||||
GIT_TAG a8f794d1251cbfd88a5011445dd5582289c727e4
|
||||
GIT_PROGRESS TRUE
|
||||
CONFIGURE_COMMAND ""
|
||||
BUILD_COMMAND ""
|
||||
@@ -35,7 +35,7 @@ set(FLASHMLA_VENDOR_DIR "${CMAKE_SOURCE_DIR}/vllm/third_party/flashmla")
|
||||
file(MAKE_DIRECTORY "${FLASHMLA_VENDOR_DIR}")
|
||||
file(READ "${flashmla_SOURCE_DIR}/flash_mla/flash_mla_interface.py"
|
||||
FLASHMLA_INTERFACE_CONTENT)
|
||||
string(REPLACE "import flash_mla.cuda as flash_mla_cuda"
|
||||
string(REPLACE "flash_mla_cuda = torch.ops._flashmla_C"
|
||||
"import vllm._flashmla_C\nflash_mla_cuda = torch.ops._flashmla_C"
|
||||
FLASHMLA_INTERFACE_CONTENT
|
||||
"${FLASHMLA_INTERFACE_CONTENT}")
|
||||
@@ -60,6 +60,9 @@ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.9)
|
||||
# CUDA 12.9 has introduced "Family-Specific Architecture Features"
|
||||
# this supports all compute_10x family
|
||||
list(APPEND SUPPORT_ARCHS "10.0f")
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.4)
|
||||
list(APPEND SUPPORT_ARCHS "10.7f")
|
||||
endif()
|
||||
elseif(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8)
|
||||
list(APPEND SUPPORT_ARCHS "10.0a")
|
||||
endif()
|
||||
@@ -72,7 +75,7 @@ if(FLASH_MLA_ARCHS)
|
||||
list(APPEND VLLM_FLASHMLA_GPU_FLAGS "--expt-relaxed-constexpr" "--expt-extended-lambda" "--use_fast_math")
|
||||
|
||||
set(FlashMLA_SOURCES
|
||||
${flashmla_SOURCE_DIR}/csrc/torch_api.cpp
|
||||
${flashmla_SOURCE_DIR}/csrc/api/api.cpp
|
||||
|
||||
# Misc kernels for decoding
|
||||
${flashmla_SOURCE_DIR}/csrc/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu
|
||||
@@ -128,6 +131,7 @@ if(FLASH_MLA_ARCHS)
|
||||
|
||||
set(FlashMLA_Extension_INCLUDES
|
||||
${flashmla_SOURCE_DIR}/csrc
|
||||
${flashmla_SOURCE_DIR}/csrc/kerutils/include
|
||||
${flashmla_SOURCE_DIR}/csrc/extension/sm90/dense_fp8/
|
||||
${flashmla_SOURCE_DIR}/csrc/cutlass/include
|
||||
${flashmla_SOURCE_DIR}/csrc/cutlass/tools/util/include
|
||||
@@ -152,15 +156,18 @@ if(FLASH_MLA_ARCHS)
|
||||
USE_SABI 3
|
||||
WITH_SOABI)
|
||||
|
||||
# Keep Stable ABI for the module, but *not* for CUDA/C++ files.
|
||||
# This prevents Py_LIMITED_API from affecting nvcc and C++ compiles.
|
||||
# Also enable C++20 for the FlashMLA sources (required for std::span, requires, etc.)
|
||||
# Enable C++20 for the FlashMLA sources (required for std::span, requires, etc.)
|
||||
target_compile_options(_flashmla_C PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:-UPy_LIMITED_API>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-UPy_LIMITED_API>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-std=c++20>
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:-std=c++20>)
|
||||
|
||||
# _flashmla_C is now ABI-stable torch 2.11+
|
||||
target_compile_definitions(_flashmla_C PRIVATE
|
||||
TORCH_TARGET_VERSION=0x020B000000000000ULL)
|
||||
if(VLLM_GPU_LANG STREQUAL "CUDA")
|
||||
target_compile_definitions(_flashmla_C PRIVATE USE_CUDA)
|
||||
endif()
|
||||
|
||||
define_extension_target(
|
||||
_flashmla_extension_C
|
||||
DESTINATION vllm
|
||||
@@ -172,15 +179,15 @@ if(FLASH_MLA_ARCHS)
|
||||
USE_SABI 3
|
||||
WITH_SOABI)
|
||||
|
||||
# Keep Stable ABI for the module, but *not* for CUDA/C++ files.
|
||||
# This prevents Py_LIMITED_API from affecting nvcc and C++ compiles.
|
||||
target_compile_options(_flashmla_extension_C PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:-UPy_LIMITED_API>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-UPy_LIMITED_API>)
|
||||
# _flashmla_extension_C is now ABI-stable w/ torch 2.11+
|
||||
target_compile_definitions(_flashmla_extension_C PRIVATE
|
||||
TORCH_TARGET_VERSION=0x020B000000000000ULL)
|
||||
if(VLLM_GPU_LANG STREQUAL "CUDA")
|
||||
target_compile_definitions(_flashmla_extension_C PRIVATE USE_CUDA)
|
||||
endif()
|
||||
else()
|
||||
message(STATUS "FlashMLA will not compile: unsupported CUDA architecture ${CUDA_ARCHS}")
|
||||
# Create empty targets for setup.py on unsupported systems
|
||||
add_custom_target(_flashmla_C)
|
||||
add_custom_target(_flashmla_extension_C)
|
||||
endif()
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ else()
|
||||
FetchContent_Declare(
|
||||
fmha_sm100
|
||||
GIT_REPOSITORY https://github.com/vllm-project/MSA.git
|
||||
GIT_TAG 2e63ec37a0fc29bc20f39cd1a52e0f5affc33a73
|
||||
GIT_TAG 890aaa1a37a598ad17ccff0827fea21540d381fa
|
||||
GIT_PROGRESS TRUE
|
||||
CONFIGURE_COMMAND ""
|
||||
BUILD_COMMAND ""
|
||||
|
||||
@@ -22,7 +22,7 @@ if(QUTLASS_SRC_DIR)
|
||||
set(qutlass_BINARY_DIR "${CMAKE_BINARY_DIR}/qutlass-binary-dir-unused")
|
||||
else()
|
||||
set(_QUTLASS_UPSTREAM_REPO "https://github.com/IST-DASLab/qutlass.git")
|
||||
set(_QUTLASS_UPSTREAM_TAG "830d2c4537c7396e14a02a46fbddd18b5d107c65")
|
||||
set(_QUTLASS_UPSTREAM_TAG "e74319e3405ce6d71965732880f5dc1f52371f64")
|
||||
|
||||
set(_qutlass_fc_root "${FETCHCONTENT_BASE_DIR}")
|
||||
if(NOT _qutlass_fc_root)
|
||||
@@ -55,7 +55,11 @@ message(STATUS "[QUTLASS] QuTLASS is available at ${qutlass_SOURCE_DIR}")
|
||||
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(QUTLASS_SM120_ARCHS "12.0f" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(QUTLASS_SM100_ARCHS "10.0f" "${CUDA_ARCHS}")
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.4)
|
||||
cuda_archs_loose_intersection(QUTLASS_SM100_ARCHS "10.0f;10.7f" "${CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(QUTLASS_SM100_ARCHS "10.0f" "${CUDA_ARCHS}")
|
||||
endif()
|
||||
else()
|
||||
cuda_archs_loose_intersection(QUTLASS_SM120_ARCHS "12.0a;12.1a" "${CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(QUTLASS_SM100_ARCHS "10.0a;10.3a" "${CUDA_ARCHS}")
|
||||
@@ -125,8 +129,6 @@ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND QUTLASS_ARCHS)
|
||||
CUDA_ARCHS "${QUTLASS_ARCHS}"
|
||||
)
|
||||
|
||||
# QuTLASS uses legacy ATen headers and cannot be built with TORCH_TARGET_VERSION.
|
||||
# Keep it as its own extension (registers torch.ops._qutlass_C).
|
||||
define_extension_target(
|
||||
_qutlass_C
|
||||
DESTINATION vllm
|
||||
@@ -139,9 +141,11 @@ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8 AND QUTLASS_ARCHS)
|
||||
WITH_SOABI)
|
||||
|
||||
target_compile_definitions(_qutlass_C PRIVATE
|
||||
QUTLASS_DISABLE_PYBIND=1
|
||||
QUTLASS_MINIMAL_BUILD=1
|
||||
TARGET_CUDA_ARCH=${QUTLASS_TARGET_CC}
|
||||
CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1)
|
||||
CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1
|
||||
TORCH_TARGET_VERSION=0x020B000000000000ULL
|
||||
USE_CUDA)
|
||||
|
||||
set_property(SOURCE ${QUTLASS_SOURCES} APPEND PROPERTY COMPILE_OPTIONS
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:--expt-relaxed-constexpr --use_fast_math -O3>
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
include(FetchContent)
|
||||
|
||||
if(DEFINED ENV{TML_FA4_SRC_DIR})
|
||||
set(TML_FA4_SRC_DIR $ENV{TML_FA4_SRC_DIR})
|
||||
endif()
|
||||
|
||||
if(TML_FA4_SRC_DIR)
|
||||
FetchContent_Declare(
|
||||
tml_fa4
|
||||
SOURCE_DIR ${TML_FA4_SRC_DIR}
|
||||
CONFIGURE_COMMAND ""
|
||||
BUILD_COMMAND "")
|
||||
else()
|
||||
FetchContent_Declare(
|
||||
tml_fa4
|
||||
GIT_REPOSITORY https://github.com/vllm-project/tml-fa4.git
|
||||
GIT_TAG b206834606ed5b5f21f8eed6b0683f528ea9cf7d
|
||||
GIT_PROGRESS TRUE
|
||||
CONFIGURE_COMMAND ""
|
||||
BUILD_COMMAND "")
|
||||
endif()
|
||||
|
||||
FetchContent_GetProperties(tml_fa4)
|
||||
if(NOT tml_fa4_POPULATED)
|
||||
FetchContent_Populate(tml_fa4)
|
||||
endif()
|
||||
message(STATUS "tml-fa4 is available at ${tml_fa4_SOURCE_DIR}")
|
||||
|
||||
add_custom_target(tml_fa4)
|
||||
|
||||
# Install into a private namespace so this implementation cannot shadow the
|
||||
# flash_attn package used by vLLM's standard attention backends.
|
||||
install(CODE "
|
||||
file(GLOB_RECURSE TML_FA4_PY_FILES
|
||||
\"${tml_fa4_SOURCE_DIR}/flash_attn/cute/*.py\")
|
||||
foreach(SRC_FILE \${TML_FA4_PY_FILES})
|
||||
file(RELATIVE_PATH REL_PATH
|
||||
\"${tml_fa4_SOURCE_DIR}/flash_attn/cute\" \${SRC_FILE})
|
||||
set(DST_FILE
|
||||
\"\${CMAKE_INSTALL_PREFIX}/vllm/third_party/tml_fa4/\${REL_PATH}\")
|
||||
get_filename_component(DST_DIR \${DST_FILE} DIRECTORY)
|
||||
file(MAKE_DIRECTORY \${DST_DIR})
|
||||
file(READ \${SRC_FILE} FILE_CONTENTS)
|
||||
string(REPLACE
|
||||
\"flash_attn.cute\"
|
||||
\"vllm.third_party.tml_fa4\"
|
||||
FILE_CONTENTS \"\${FILE_CONTENTS}\")
|
||||
file(WRITE \${DST_FILE} \"\${FILE_CONTENTS}\")
|
||||
endforeach()
|
||||
" COMPONENT tml_fa4)
|
||||
@@ -39,7 +39,7 @@ else()
|
||||
FetchContent_Declare(
|
||||
vllm-flash-attn
|
||||
GIT_REPOSITORY https://github.com/vllm-project/flash-attention.git
|
||||
GIT_TAG bb9a72e7dde0dc614ffc663e052cd6a19ce73a42
|
||||
GIT_TAG ed4b7342bc8f0489dd9b649d5288867e35fc6a32
|
||||
GIT_PROGRESS TRUE
|
||||
# Don't share the vllm-flash-attn build between build types
|
||||
BINARY_DIR ${CMAKE_BINARY_DIR}/vllm-flash-attn
|
||||
|
||||
+15
-5
@@ -396,14 +396,24 @@ function(cuda_archs_loose_intersection OUT_CUDA_ARCHS SRC_CUDA_ARCHS TGT_CUDA_AR
|
||||
# match — e.g. SRC="12.0f" matches TGT="12.1a" since SM121 is in the SM12x
|
||||
# family. The output uses TGT's value to preserve the user's compilation flags.
|
||||
set(_CUDA_ARCHS)
|
||||
# Resolve exact base matches before family fallbacks so a generic entry such
|
||||
# as 10.0f cannot consume a 10.7 target that has a 10.7f source entry.
|
||||
foreach(_arch ${_SRC_CUDA_ARCHS})
|
||||
if(_arch MATCHES "[af]$")
|
||||
string(REGEX REPLACE "[af]$" "" _base "${_arch}")
|
||||
if("${_base}" IN_LIST _TGT_CUDA_ARCHS)
|
||||
list(REMOVE_ITEM _SRC_CUDA_ARCHS "${_arch}")
|
||||
list(REMOVE_ITEM _TGT_CUDA_ARCHS "${_base}")
|
||||
list(APPEND _CUDA_ARCHS "${_arch}")
|
||||
endif()
|
||||
endif()
|
||||
endforeach()
|
||||
|
||||
foreach(_arch ${_SRC_CUDA_ARCHS})
|
||||
if(_arch MATCHES "[af]$")
|
||||
list(REMOVE_ITEM _SRC_CUDA_ARCHS "${_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _base "${_arch}")
|
||||
if ("${_base}" IN_LIST TGT_CUDA_ARCHS)
|
||||
list(REMOVE_ITEM _TGT_CUDA_ARCHS "${_base}")
|
||||
list(APPEND _CUDA_ARCHS "${_arch}")
|
||||
elseif("${_base}a" IN_LIST _TGT_CUDA_ARCHS)
|
||||
if("${_base}a" IN_LIST _TGT_CUDA_ARCHS)
|
||||
list(REMOVE_ITEM _TGT_CUDA_ARCHS "${_base}a")
|
||||
list(APPEND _CUDA_ARCHS "${_base}a")
|
||||
elseif("${_base}f" IN_LIST _TGT_CUDA_ARCHS)
|
||||
@@ -487,7 +497,7 @@ endfunction()
|
||||
|
||||
function(cuda_archs_sm90plus OUT_CUDA_ARCHS TGT_CUDA_ARCHS)
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
|
||||
cuda_archs_loose_intersection(_archs "9.0a;10.0f;11.0f;12.0f" "${TGT_CUDA_ARCHS}")
|
||||
cuda_archs_loose_intersection(_archs "9.0a;10.0f;10.7f;11.0f;12.0f" "${TGT_CUDA_ARCHS}")
|
||||
else()
|
||||
cuda_archs_loose_intersection(_archs "9.0a;10.0a;10.1a;10.3a;12.0a;12.1a" "${TGT_CUDA_ARCHS}")
|
||||
endif()
|
||||
|
||||
+1
-2
@@ -67,9 +67,8 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
torch::Tensor const& src_cache, // [NUM_BLOCKS, BLOCK_SIZE, 656]
|
||||
torch::Tensor const& dst, // [TOT_TOKENS, 576]
|
||||
torch::Tensor const& block_table, // [BATCH, BLOCK_INDICES]
|
||||
torch::Tensor const& seq_lens, // [BATCH]
|
||||
torch::Tensor const& workspace_starts, // [BATCH]
|
||||
int64_t batch_size);
|
||||
int64_t batch_size, std::optional<torch::Tensor> seq_starts = std::nullopt);
|
||||
|
||||
// Indexer K quantization and cache function
|
||||
void indexer_k_quant_and_cache(
|
||||
|
||||
@@ -102,7 +102,9 @@ class TileGemm82 {
|
||||
kv_cache_t* __restrict__ curr_b = b_tile;
|
||||
|
||||
for (int32_t k = 0; k < dynamic_k_size; ++k) {
|
||||
auto [fp32_b_0_reg, fp32_b_1_reg] = load_b_pair_vec(curr_b);
|
||||
auto fp32_b_regs = load_b_pair_vec(curr_b);
|
||||
auto fp32_b_0_reg = fp32_b_regs.first;
|
||||
auto fp32_b_1_reg = fp32_b_regs.second;
|
||||
|
||||
float* __restrict__ curr_m_a = curr_a;
|
||||
vec_op::unroll_loop<int32_t, M>([&](int32_t i) {
|
||||
|
||||
@@ -336,13 +336,14 @@ struct FP32Vec8 : public Vec<FP32Vec8> {
|
||||
reg.val[1] = fp16_to_fp32_bits(raw_lo);
|
||||
}
|
||||
float reduce_sum() const {
|
||||
AliasReg ar;
|
||||
ar.reg = reg;
|
||||
float result = 0;
|
||||
unroll_loop<int, VEC_ELEM_NUM>(
|
||||
[&result, &ar](int i) { result += ar.values[i]; });
|
||||
|
||||
return result;
|
||||
// VSX horizontal reduction: 3 vector ops instead of 8 scalar adds.
|
||||
// Step 1: pairwise sum of the two 4-wide halves
|
||||
__vector float s = vec_add(reg.val[0], reg.val[1]);
|
||||
// Step 2: rotate by 8 bytes (2 floats) and add
|
||||
s = vec_add(s, vec_sld(s, s, 8));
|
||||
// Step 3: rotate by 4 bytes (1 float) and add => all lanes hold total
|
||||
s = vec_add(s, vec_sld(s, s, 4));
|
||||
return vec_extract(s, 0);
|
||||
}
|
||||
FP32Vec8 exp() const {
|
||||
f32x4x2_t out;
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
//
|
||||
// CPU at::Tensor wrappers for Mamba decode-step kernels defined in
|
||||
// mamba_kernels.hpp.
|
||||
|
||||
#include "cpu/mamba_kernels.hpp"
|
||||
|
||||
#include <ATen/ATen.h>
|
||||
#include <torch/library.h>
|
||||
#include <c10/util/Optional.h>
|
||||
|
||||
#include "cpu_types.hpp"
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// causal_conv1d_update
|
||||
// ---------------------------------------------------------------------------
|
||||
at::Tensor causal_conv1d_update_cpu_impl(
|
||||
at::Tensor& x, at::Tensor& conv_state, const at::Tensor& weight,
|
||||
const c10::optional<at::Tensor>& bias,
|
||||
const c10::optional<std::string>& activation,
|
||||
const c10::optional<at::Tensor>& conv_state_indices,
|
||||
const c10::optional<at::Tensor>& query_start_loc, int64_t pad_slot_id) {
|
||||
bool do_silu = false;
|
||||
if (activation.has_value()) {
|
||||
const std::string& act = activation.value();
|
||||
do_silu = (act == "silu" || act == "swish");
|
||||
}
|
||||
|
||||
at::ScalarType dtype = x.scalar_type();
|
||||
|
||||
// Input x: contiguous in native dtype.
|
||||
at::Tensor x_c = x.is_contiguous() ? x : x.contiguous();
|
||||
|
||||
// conv_state: NEVER copy the full paged tensor just for layout reasons.
|
||||
// If the dtype matches we work directly on conv_state (contiguous or not)
|
||||
// by extracting strides and passing them to the kernel.
|
||||
// Only a dtype-conversion copy is made when types differ (rare for BF16).
|
||||
bool state_type_ok = (conv_state.scalar_type() == dtype);
|
||||
at::Tensor state_c = state_type_ok ? conv_state : conv_state.to(dtype);
|
||||
// state_c and conv_state may be non-contiguous — that is intentional.
|
||||
|
||||
// Weight: coerce to same dtype if needed (should match in practice)
|
||||
at::Tensor w_c =
|
||||
(weight.scalar_type() != dtype)
|
||||
? weight.to(dtype).contiguous()
|
||||
: (weight.is_contiguous() ? weight : weight.contiguous());
|
||||
|
||||
// Bias stays float32 (small scalar, used only for fp32 accumulation)
|
||||
at::Tensor bias_f32;
|
||||
if (bias.has_value() && bias.value().defined())
|
||||
bias_f32 = bias.value().to(at::kFloat).contiguous();
|
||||
|
||||
int64_t batch = x_c.size(0);
|
||||
int64_t dim = x_c.size(1);
|
||||
int64_t seqlen = (x_c.dim() == 3) ? x_c.size(2) : 1;
|
||||
int64_t width = w_c.size(1);
|
||||
int64_t state_len = state_c.size(2);
|
||||
|
||||
// Extract strides — works for contiguous AND non-contiguous (transposed)
|
||||
// state. stride(0): between cache slots (e.g. num_slots × dim × width-1 in
|
||||
// contiguous) stride(1): between conv channels (dim stride) stride(2):
|
||||
// between state elements (=1 when contiguous, =dim when transposed)
|
||||
int64_t stride_s_slot = state_c.stride(0);
|
||||
int64_t stride_s_dim = state_c.stride(1);
|
||||
int64_t stride_s_state = state_c.stride(2);
|
||||
|
||||
at::Tensor out = x_c.clone(); // native dtype, no float32 alloc
|
||||
|
||||
const int32_t* cache_idx_ptr = nullptr;
|
||||
at::Tensor cache_idx_int;
|
||||
if (conv_state_indices.has_value()) {
|
||||
cache_idx_int = conv_state_indices.value().to(at::kInt).contiguous();
|
||||
cache_idx_ptr = cache_idx_int.data_ptr<int32_t>();
|
||||
}
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(dtype, "causal_conv1d_update", [&] {
|
||||
mamba_cpu::causal_conv1d_update_kernel<scalar_t>(
|
||||
x_c.data_ptr<scalar_t>(), state_c.data_ptr<scalar_t>(), stride_s_slot,
|
||||
stride_s_dim, stride_s_state, w_c.data_ptr<scalar_t>(),
|
||||
bias_f32.defined() ? bias_f32.data_ptr<float>() : nullptr,
|
||||
out.data_ptr<scalar_t>(), cache_idx_ptr,
|
||||
static_cast<int32_t>(pad_slot_id), batch, dim, seqlen, width, state_len,
|
||||
do_silu);
|
||||
});
|
||||
|
||||
// Write back only when a type-conversion copy was made.
|
||||
// Layout-only non-contiguity is handled via strides above — no copy needed.
|
||||
if (!state_type_ok) conv_state.copy_(state_c);
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// selective_state_update
|
||||
// ---------------------------------------------------------------------------
|
||||
void selective_state_update_cpu_impl(
|
||||
at::Tensor& state, // (nstates, nheads, dim, dstate)
|
||||
const at::Tensor& x, // (N, nheads, dim)
|
||||
const at::Tensor& dt, const at::Tensor& A, const at::Tensor& B,
|
||||
const at::Tensor& C, const c10::optional<at::Tensor>& D,
|
||||
const c10::optional<at::Tensor>& z,
|
||||
const c10::optional<at::Tensor>& dt_bias, bool dt_softplus,
|
||||
const c10::optional<at::Tensor>& state_batch_indices,
|
||||
const c10::optional<at::Tensor>& dst_state_batch_indices,
|
||||
int64_t null_block_id, at::Tensor& out,
|
||||
const c10::optional<at::Tensor>& num_accepted_tokens,
|
||||
const c10::optional<at::Tensor>& cu_seqlens) {
|
||||
at::ScalarType state_type = state.scalar_type();
|
||||
at::ScalarType input_type = x.scalar_type();
|
||||
|
||||
// x, B, C must be contiguous and match input_type
|
||||
auto ensure_input = [input_type](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = (t.scalar_type() != input_type) ? t.to(input_type) : t;
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
at::Tensor x_in = ensure_input(x);
|
||||
at::Tensor B_in = ensure_input(B);
|
||||
at::Tensor C_in = ensure_input(C);
|
||||
at::Tensor z_in;
|
||||
if (z.has_value() && z.value().defined()) z_in = ensure_input(z.value());
|
||||
|
||||
// A, D, dt_bias are float32 model parameters that arrive here as expanded
|
||||
// tensors, e.g. A is (nheads, head_dim, dstate) with strides (1, 0, 0).
|
||||
// We need just the scalar value per head as a (nheads,) 1-D array so that
|
||||
// A_ptr[h] in the kernel correctly reads head h's value.
|
||||
//
|
||||
// Strategy: peel trailing expanded (stride=0) dims via .select(), which is
|
||||
// a zero-copy view. For A: (nheads, head_dim, dstate) strides (1,0,0)
|
||||
// → .select(2,0) → (nheads, head_dim) strides (1,0)
|
||||
// → .select(1,0) → (nheads,) stride (1,) ← contiguous, free.
|
||||
// No allocation, no type conversion (A is already float32).
|
||||
auto to_per_head_1d_f32 = [](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = t;
|
||||
// Peel trailing dimensions that are broadcast (stride=0 or size=1)
|
||||
while (r.dim() > 1) r = r.select(r.dim() - 1, 0);
|
||||
if (r.scalar_type() != at::kFloat) r = r.to(at::kFloat);
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
|
||||
at::Tensor A_f32 = to_per_head_1d_f32(A); // (nheads,) float32
|
||||
at::Tensor D_f32, dt_bias_f32;
|
||||
if (D.has_value() && D.value().defined())
|
||||
D_f32 = to_per_head_1d_f32(D.value());
|
||||
if (dt_bias.has_value() && dt_bias.value().defined())
|
||||
dt_bias_f32 = to_per_head_1d_f32(dt_bias.value());
|
||||
|
||||
// dt: reduce (N, nheads, head_dim) expanded tensor → (N, nheads) BEFORE
|
||||
// the type conversion so we convert head_dim x fewer elements.
|
||||
at::Tensor dt_f32;
|
||||
{
|
||||
// If dt was expanded to (N, nheads, head_dim) with stride-0 in dim 2,
|
||||
// take a zero-copy view of index 0 along that dim first.
|
||||
at::Tensor t2 = (dt.dim() == 3) ? dt.select(2, 0) : dt; // (N, nheads)
|
||||
at::Tensor t3 = (t2.scalar_type() != at::kFloat) ? t2.to(at::kFloat) : t2;
|
||||
dt_f32 = t3.is_contiguous() ? t3 : t3.contiguous();
|
||||
}
|
||||
|
||||
int64_t nheads = state.size(1);
|
||||
int64_t dim = state.size(2);
|
||||
int64_t dstate = state.size(3);
|
||||
int64_t N = (cu_seqlens.has_value() && cu_seqlens.value().defined())
|
||||
? cu_seqlens.value().size(0) - 1
|
||||
: x_in.size(0);
|
||||
int64_t ngroups = B_in.size(1);
|
||||
|
||||
// Strides
|
||||
int64_t stride_state_n = state.stride(0);
|
||||
int64_t stride_state_h = state.stride(1);
|
||||
int64_t stride_state_d = state.stride(2);
|
||||
int64_t stride_x_n = x_in.stride(0);
|
||||
int64_t stride_x_h = x_in.stride(1);
|
||||
int64_t stride_dt_n = dt_f32.stride(0); // dt is (N, nheads)
|
||||
int64_t stride_BC_n = B_in.stride(0);
|
||||
int64_t stride_BC_g = B_in.stride(1);
|
||||
int64_t stride_out_n = out.stride(0);
|
||||
int64_t stride_out_h = out.stride(1);
|
||||
|
||||
// Optional index pointers
|
||||
auto get_int32_ptr =
|
||||
[](const c10::optional<at::Tensor>& opt) -> const int32_t* {
|
||||
return (opt.has_value() && opt.value().defined())
|
||||
? opt.value().data_ptr<int32_t>()
|
||||
: nullptr;
|
||||
};
|
||||
const int32_t* sbi_ptr = get_int32_ptr(state_batch_indices);
|
||||
const int32_t* dsbi_ptr = get_int32_ptr(dst_state_batch_indices);
|
||||
const int32_t* nat_ptr = get_int32_ptr(num_accepted_tokens);
|
||||
const int32_t* csl_ptr = get_int32_ptr(cu_seqlens);
|
||||
|
||||
// Dispatch on (state_t, input_t, out_t): write directly into `out`
|
||||
// without any intermediate float32 buffer.
|
||||
VLLM_DISPATCH_FLOATING_TYPES(state_type, "ssu_state", [&] {
|
||||
using state_t = scalar_t;
|
||||
VLLM_DISPATCH_FLOATING_TYPES(input_type, "ssu_input", [&] {
|
||||
using input_t = scalar_t;
|
||||
VLLM_DISPATCH_FLOATING_TYPES(out.scalar_type(), "ssu_out", [&] {
|
||||
using out_t = scalar_t;
|
||||
mamba_cpu::selective_state_update_kernel<state_t, input_t, out_t>(
|
||||
state.data_ptr<state_t>(), stride_state_n, stride_state_h,
|
||||
stride_state_d, x_in.data_ptr<input_t>(), stride_x_n, stride_x_h,
|
||||
dt_f32.data_ptr<float>(), stride_dt_n, A_f32.data_ptr<float>(),
|
||||
B_in.data_ptr<input_t>(), C_in.data_ptr<input_t>(), stride_BC_n,
|
||||
stride_BC_g, D_f32.defined() ? D_f32.data_ptr<float>() : nullptr,
|
||||
z_in.defined() ? z_in.data_ptr<input_t>() : nullptr,
|
||||
dt_bias_f32.defined() ? dt_bias_f32.data_ptr<float>() : nullptr,
|
||||
out.data_ptr<out_t>(), stride_out_n, stride_out_h, sbi_ptr,
|
||||
dsbi_ptr, static_cast<int32_t>(null_block_id), nat_ptr, csl_ptr, N,
|
||||
nheads, ngroups, dim, dstate, dt_softplus);
|
||||
});
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// mamba_chunk_scan_fwd_cpu
|
||||
// ---------------------------------------------------------------------------
|
||||
void mamba_chunk_scan_fwd_cpu_impl(
|
||||
at::Tensor& out, // [seqlen, nheads, headdim] — pre-allocated by caller
|
||||
at::Tensor&
|
||||
final_states, // [batch, nheads, headdim, dstate] float32 contiguous
|
||||
const at::Tensor& x, // [seqlen, nheads, headdim]
|
||||
const at::Tensor&
|
||||
dt, // [seqlen, nheads] float32 (preprocessed: bias+softplus+clamp)
|
||||
const at::Tensor& A, // [nheads] float32
|
||||
const at::Tensor& B, // [seqlen, ngroups, dstate]
|
||||
const at::Tensor& C, // [seqlen, ngroups, dstate]
|
||||
const c10::optional<at::Tensor>& D, // [nheads] float32 (optional)
|
||||
const c10::optional<at::Tensor>& z, // [seqlen, nheads, headdim] (optional)
|
||||
const at::Tensor& cu_seqlens // [batch+1] int32
|
||||
) {
|
||||
const at::ScalarType input_type = x.scalar_type();
|
||||
|
||||
auto ensure_contig = [input_type](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = (t.scalar_type() != input_type) ? t.to(input_type) : t;
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
at::Tensor x_in = ensure_contig(x);
|
||||
at::Tensor B_in = ensure_contig(B);
|
||||
at::Tensor C_in = ensure_contig(C);
|
||||
at::Tensor z_in;
|
||||
if (z.has_value() && z.value().defined()) z_in = ensure_contig(z.value());
|
||||
|
||||
// A and D are float32 model parameters, potentially broadcast-expanded.
|
||||
// Strip trailing broadcast dims to get a contiguous (nheads,) array.
|
||||
auto to_per_head_f32 = [](const at::Tensor& t) -> at::Tensor {
|
||||
at::Tensor r = t;
|
||||
while (r.dim() > 1) r = r.select(r.dim() - 1, 0);
|
||||
if (r.scalar_type() != at::kFloat) r = r.to(at::kFloat);
|
||||
return r.is_contiguous() ? r : r.contiguous();
|
||||
};
|
||||
at::Tensor A_f32 = to_per_head_f32(A);
|
||||
at::Tensor D_f32;
|
||||
if (D.has_value() && D.value().defined()) D_f32 = to_per_head_f32(D.value());
|
||||
|
||||
// dt: [seqlen, nheads] float32 — caller has applied bias+softplus+clamp in
|
||||
// Python.
|
||||
at::Tensor dt_c = dt.is_contiguous() ? dt : dt.contiguous();
|
||||
if (dt_c.scalar_type() != at::kFloat) dt_c = dt_c.to(at::kFloat);
|
||||
|
||||
at::Tensor cu_int = cu_seqlens.to(at::kInt).contiguous();
|
||||
|
||||
const int64_t batch = final_states.size(0);
|
||||
const int64_t nheads = final_states.size(1);
|
||||
const int64_t headdim = final_states.size(2);
|
||||
const int64_t dstate = final_states.size(3);
|
||||
const int64_t ngroups = B_in.size(1);
|
||||
|
||||
TORCH_CHECK(final_states.is_contiguous(),
|
||||
"mamba_chunk_scan_fwd_cpu: final_states must be contiguous");
|
||||
TORCH_CHECK(out.is_contiguous(),
|
||||
"mamba_chunk_scan_fwd_cpu: out must be contiguous (writes via "
|
||||
"raw data_ptr)");
|
||||
|
||||
VLLM_DISPATCH_FLOATING_TYPES(input_type, "mamba_chunk_scan_fwd_cpu", [&] {
|
||||
mamba_cpu::mamba_chunk_scan_fwd_kernel<scalar_t>(
|
||||
final_states.data_ptr<float>(), x_in.data_ptr<scalar_t>(),
|
||||
dt_c.data_ptr<float>(), A_f32.data_ptr<float>(),
|
||||
B_in.data_ptr<scalar_t>(), C_in.data_ptr<scalar_t>(),
|
||||
D_f32.defined() ? D_f32.data_ptr<float>() : nullptr,
|
||||
z_in.defined() ? z_in.data_ptr<scalar_t>() : nullptr,
|
||||
out.data_ptr<scalar_t>(), cu_int.data_ptr<int32_t>(), batch, nheads,
|
||||
ngroups, headdim, dstate);
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,382 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
//
|
||||
// Fused CPU vector kernels for Mamba decode-step hotspots:
|
||||
// - causal_conv1d_update (depthwise 1-D conv state roll + compute)
|
||||
// - selective_state_update (SSM recurrence, single-step)
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cpu_types.hpp"
|
||||
#include <cmath>
|
||||
#include <cstring>
|
||||
#include <cstdint>
|
||||
#include <algorithm>
|
||||
|
||||
namespace mamba_cpu {
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// causal_conv1d_update — templated for native BF16/FP32
|
||||
//
|
||||
// state_ptr may point to a NON-CONTIGUOUS paged KV cache tensor.
|
||||
// Explicit strides are passed so the kernel writes directly into the
|
||||
// correct memory locations without making a contiguous copy of the full
|
||||
// paged tensor (which was the source of the 34-41% direct_copy_kernel).
|
||||
//
|
||||
// stride_s_slot = state.stride(0) — between cache slots
|
||||
// stride_s_dim = state.stride(1) — between conv_dim channels
|
||||
// stride_s_state = state.stride(2) — between state elements
|
||||
//
|
||||
// When stride_s_state == 1 (contiguous), the memmove fast path is used.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename scalar_t>
|
||||
inline void causal_conv1d_update_kernel(
|
||||
const scalar_t* __restrict__ x_ptr, scalar_t* __restrict__ state_ptr,
|
||||
int64_t stride_s_slot, int64_t stride_s_dim, int64_t stride_s_state,
|
||||
const scalar_t* __restrict__ weight_ptr, const float* __restrict__ bias_ptr,
|
||||
scalar_t* __restrict__ out_ptr, const int32_t* __restrict__ cache_idxs,
|
||||
int32_t pad_slot_id, int64_t batch, int64_t dim, int64_t seqlen,
|
||||
int64_t width, int64_t state_len, bool do_silu) {
|
||||
#pragma omp parallel for
|
||||
for (int64_t b = 0; b < batch; ++b) {
|
||||
int64_t cache_idx = (cache_idxs != nullptr) ? cache_idxs[b] : b;
|
||||
if (cache_idx == pad_slot_id) continue;
|
||||
|
||||
for (int64_t t = 0; t < seqlen; ++t) {
|
||||
const scalar_t* x_b = x_ptr + (b * dim * seqlen + t);
|
||||
scalar_t* out_b = out_ptr + (b * dim * seqlen + t);
|
||||
// Base of this slot in the (possibly non-contiguous) paged state
|
||||
scalar_t* s_base = state_ptr + cache_idx * stride_s_slot;
|
||||
|
||||
for (int64_t d = 0; d < dim; ++d) {
|
||||
float x_val = static_cast<float>(x_b[d * seqlen]);
|
||||
scalar_t* sd = s_base + d * stride_s_dim; // start of this dim's state
|
||||
const scalar_t* w = weight_ptr + d * width;
|
||||
|
||||
// Accumulate in float32 for precision
|
||||
float acc = (bias_ptr != nullptr) ? bias_ptr[d] : 0.0f;
|
||||
for (int64_t k = 0; k < state_len; ++k) {
|
||||
acc += static_cast<float>(w[k]) *
|
||||
static_cast<float>(sd[k * stride_s_state]);
|
||||
}
|
||||
acc += static_cast<float>(w[state_len]) * x_val;
|
||||
|
||||
// Shift state left and append new input.
|
||||
// Use memmove when contiguous (stride==1); element loop otherwise.
|
||||
if (stride_s_state == 1) {
|
||||
if (state_len > 1)
|
||||
std::memmove(sd, sd + 1, (state_len - 1) * sizeof(scalar_t));
|
||||
if (state_len > 0) sd[state_len - 1] = static_cast<scalar_t>(x_val);
|
||||
} else {
|
||||
for (int64_t k = 0; k < state_len - 1; ++k)
|
||||
sd[k * stride_s_state] = sd[(k + 1) * stride_s_state];
|
||||
if (state_len > 0)
|
||||
sd[(state_len - 1) * stride_s_state] = static_cast<scalar_t>(x_val);
|
||||
}
|
||||
|
||||
if (do_silu) {
|
||||
float sigmoid = (acc >= 0) ? 1.0f / (1.0f + std::exp(-acc))
|
||||
: std::exp(acc) / (1.0f + std::exp(acc));
|
||||
acc *= sigmoid;
|
||||
}
|
||||
out_b[d * seqlen] = static_cast<scalar_t>(acc);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// selective_state_update
|
||||
//
|
||||
// Template parameters:
|
||||
// state_t - dtype of ssm_state cache (typically BFloat16)
|
||||
// input_t - dtype of x, B, C (typically BFloat16)
|
||||
// out_t - dtype of output tensor (typically BFloat16)
|
||||
// Write directly — no float32 intermediate buffer needed.
|
||||
//
|
||||
// A, D, dt_bias are accepted as const float* (they are always float32
|
||||
// model parameters in Mamba2). This eliminates the per-call float32→BF16
|
||||
// conversion and the .contiguous() materialisation of the broadcast-expand.
|
||||
//
|
||||
// dt is accepted as a (N, nheads) scalar-per-head tensor, not as the
|
||||
// (N, nheads, head_dim) expansion, so no .contiguous() copy is needed.
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename state_t, typename input_t, typename out_t = float>
|
||||
inline void selective_state_update_kernel(
|
||||
state_t* __restrict__ state_ptr, int64_t stride_state_n,
|
||||
int64_t stride_state_h, int64_t stride_state_d,
|
||||
const input_t* __restrict__ x_ptr, int64_t stride_x_n, int64_t stride_x_h,
|
||||
// dt: (N, nheads) — scalar per head, NOT expanded to head_dim
|
||||
const float* __restrict__ dt_ptr, int64_t stride_dt_n,
|
||||
// A: (nheads,) float32 — scalar per head
|
||||
const float* __restrict__ A_ptr, const input_t* __restrict__ B_ptr,
|
||||
const input_t* __restrict__ C_ptr, int64_t stride_BC_n, int64_t stride_BC_g,
|
||||
// D: (nheads,) float32 — scalar per head (nullptr if not used)
|
||||
const float* __restrict__ D_ptr,
|
||||
// z: same shape as x (optional)
|
||||
const input_t* __restrict__ z_ptr,
|
||||
// dt_bias: (nheads,) float32 — scalar per head (nullptr if not used)
|
||||
const float* __restrict__ dt_bias_ptr, out_t* __restrict__ out_ptr,
|
||||
int64_t stride_out_n, int64_t stride_out_h,
|
||||
const int32_t* __restrict__ state_batch_indices,
|
||||
const int32_t* __restrict__ dst_state_batch_indices, int32_t null_block_id,
|
||||
const int32_t* __restrict__ num_accepted_tokens,
|
||||
const int32_t* __restrict__ cu_seqlens, int64_t N, int64_t nheads,
|
||||
int64_t ngroups, int64_t dim, int64_t dstate, bool dt_softplus) {
|
||||
using state_vec_t = vec_op::vec_t<state_t>;
|
||||
using input_vec_t = vec_op::vec_t<input_t>;
|
||||
constexpr int VEC_ELEM_NUM = 8;
|
||||
|
||||
int64_t nheads_per_group = nheads / ngroups;
|
||||
|
||||
for (int64_t seq_idx = 0; seq_idx < N; ++seq_idx) {
|
||||
int64_t bos, seq_len;
|
||||
if (cu_seqlens != nullptr) {
|
||||
bos = cu_seqlens[seq_idx];
|
||||
seq_len = cu_seqlens[seq_idx + 1] - bos;
|
||||
} else {
|
||||
bos = seq_idx;
|
||||
seq_len = 1;
|
||||
}
|
||||
|
||||
int64_t state_read_idx = (state_batch_indices != nullptr)
|
||||
? state_batch_indices[seq_idx]
|
||||
: seq_idx;
|
||||
if (state_read_idx == null_block_id) continue;
|
||||
|
||||
int64_t state_write_idx = (num_accepted_tokens == nullptr)
|
||||
? ((dst_state_batch_indices != nullptr)
|
||||
? dst_state_batch_indices[seq_idx]
|
||||
: state_read_idx)
|
||||
: -1;
|
||||
|
||||
state_t* s = state_ptr + state_read_idx * stride_state_n;
|
||||
|
||||
for (int64_t t = 0; t < seq_len; ++t) {
|
||||
int64_t token_idx = bos + t;
|
||||
const input_t* x_tok = x_ptr + token_idx * stride_x_n;
|
||||
// dt: (N, nheads) — one float per head per token
|
||||
const float* dt_tok = dt_ptr + token_idx * stride_dt_n;
|
||||
const input_t* B_tok = B_ptr + token_idx * stride_BC_n;
|
||||
const input_t* C_tok = C_ptr + token_idx * stride_BC_n;
|
||||
out_t* out_tok = out_ptr + token_idx * stride_out_n;
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t h = 0; h < nheads; ++h) {
|
||||
int64_t g = h / nheads_per_group;
|
||||
const input_t* x_h = x_tok + h * stride_x_h;
|
||||
const input_t* B_g = B_tok + g * stride_BC_g;
|
||||
const input_t* C_g = C_tok + g * stride_BC_g;
|
||||
out_t* out_h = out_tok + h * stride_out_h;
|
||||
state_t* s_h = s + h * stride_state_h;
|
||||
|
||||
// Read scalars-per-head (A, dt, dt_bias, D) — no per-dim indexing
|
||||
float dt_val = dt_tok[h];
|
||||
if (dt_bias_ptr != nullptr) dt_val += dt_bias_ptr[h];
|
||||
if (dt_softplus) {
|
||||
dt_val = (dt_val <= 20.0f) ? std::log1p(std::exp(dt_val)) : dt_val;
|
||||
}
|
||||
const float A_val = A_ptr[h]; // scalar: same for all dim, dstate
|
||||
const float D_val = (D_ptr != nullptr) ? D_ptr[h] : 0.0f;
|
||||
|
||||
const input_t* z_h =
|
||||
(z_ptr != nullptr) ? z_ptr + token_idx * stride_x_n + h * stride_x_h
|
||||
: nullptr;
|
||||
|
||||
vec_op::FP32Vec8 dt_vec(dt_val);
|
||||
// dA = exp(A * dt): A and dt are SCALARS per head, so compute once
|
||||
// and broadcast. This saves 7 redundant std::exp() calls that
|
||||
// FP32Vec8::exp() would otherwise make on the broadcast vector.
|
||||
const float dA_scalar = std::exp(A_val * dt_val);
|
||||
vec_op::FP32Vec8 dA(dA_scalar); // broadcast
|
||||
|
||||
for (int64_t d = 0; d < dim; ++d) {
|
||||
float x_val = static_cast<float>(x_h[d]);
|
||||
|
||||
vec_op::FP32Vec8 out_vec(0.0f);
|
||||
state_t* s_hd = s_h + d * stride_state_d;
|
||||
const input_t* B_g_base = B_g;
|
||||
const input_t* C_g_base = C_g;
|
||||
|
||||
vec_op::FP32Vec8 x_vec(x_val);
|
||||
// dBx = B * x * dt — same dA for all dstate (A is scalar)
|
||||
// s_new = s * dA + B * x * dt
|
||||
|
||||
int64_t n = 0;
|
||||
for (; n <= dstate - VEC_ELEM_NUM; n += VEC_ELEM_NUM) {
|
||||
vec_op::FP32Vec8 B_v((input_vec_t(B_g_base + n)));
|
||||
vec_op::FP32Vec8 C_v((input_vec_t(C_g_base + n)));
|
||||
vec_op::FP32Vec8 s_v((state_vec_t(s_hd + n)));
|
||||
|
||||
vec_op::FP32Vec8 dBx = B_v * x_vec * dt_vec;
|
||||
vec_op::FP32Vec8 s_new = s_v * dA + dBx;
|
||||
|
||||
state_vec_t(s_new).save(s_hd + n);
|
||||
out_vec = out_vec + s_new * C_v;
|
||||
}
|
||||
|
||||
float out_val = out_vec.reduce_sum();
|
||||
for (; n < dstate; ++n) {
|
||||
// Reuse dA_scalar computed once per head — no exp() re-call
|
||||
float dBx = static_cast<float>(B_g[n]) * x_val * dt_val;
|
||||
float s_new = static_cast<float>(s_hd[n]) * dA_scalar + dBx;
|
||||
s_hd[n] = static_cast<state_t>(s_new);
|
||||
out_val += s_new * static_cast<float>(C_g[n]);
|
||||
}
|
||||
|
||||
if (D_ptr != nullptr) out_val += x_val * D_val;
|
||||
if (z_h != nullptr) {
|
||||
float z_val = static_cast<float>(z_h[d]);
|
||||
float sigmoid = (z_val >= 0)
|
||||
? 1.0f / (1.0f + std::exp(-z_val))
|
||||
: std::exp(z_val) / (1.0f + std::exp(z_val));
|
||||
out_val *= z_val * sigmoid;
|
||||
}
|
||||
out_h[d] = static_cast<out_t>(out_val);
|
||||
}
|
||||
}
|
||||
|
||||
if (num_accepted_tokens != nullptr &&
|
||||
dst_state_batch_indices != nullptr) {
|
||||
int64_t token_dst_idx = dst_state_batch_indices[seq_idx * seq_len + t];
|
||||
if (token_dst_idx != null_block_id && token_dst_idx != state_read_idx) {
|
||||
state_t* dst_s = state_ptr + token_dst_idx * stride_state_n;
|
||||
std::memmove(dst_s, s, nheads * stride_state_h * sizeof(state_t));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (num_accepted_tokens == nullptr && state_write_idx != null_block_id &&
|
||||
state_write_idx != state_read_idx) {
|
||||
state_t* dst_s = state_ptr + state_write_idx * stride_state_n;
|
||||
std::memmove(dst_s, s, nheads * stride_state_h * sizeof(state_t));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// mamba_chunk_scan_fwd
|
||||
//
|
||||
// Prefill SSM recurrence for Mamba2 / SSD models.
|
||||
//
|
||||
// Key difference from selective_state_update_kernel (decode path):
|
||||
// - #pragma omp parallel for collapse(2) is OUTSIDE the time loop.
|
||||
// Each thread owns a (batch, head) slice and runs the entire token
|
||||
// sequence without any per-token OpenMP synchronisation overhead.
|
||||
// For seqlen=256, this eliminates 256 thread-barrier launches per batch.
|
||||
//
|
||||
// `dt` arrives already processed (float32, after bias + softplus + clamp)
|
||||
// to keep this kernel simple. Preprocessing is done in the Python wrapper.
|
||||
//
|
||||
// `states_ptr` points to the [batch, nheads, headdim, dstate] float32 output
|
||||
// tensor, pre-initialised by the caller (zero or from initial_states).
|
||||
// Each (b, h) slice is private to exactly one thread via collapse(2), so
|
||||
// there are no write conflicts.
|
||||
//
|
||||
// D is treated as a scalar per head ([nheads] float32).
|
||||
// ---------------------------------------------------------------------------
|
||||
template <typename input_t>
|
||||
inline void mamba_chunk_scan_fwd_kernel(
|
||||
float* __restrict__ states_ptr, // [batch, nheads, headdim, dstate] f32
|
||||
const input_t* __restrict__ x_ptr, // [seqlen, nheads, headdim]
|
||||
const float* __restrict__ dt_ptr, // [seqlen, nheads] f32 (preprocessed)
|
||||
const float* __restrict__ A_ptr, // [nheads] f32
|
||||
const input_t* __restrict__ B_ptr, // [seqlen, ngroups, dstate]
|
||||
const input_t* __restrict__ C_ptr, // [seqlen, ngroups, dstate]
|
||||
const float* __restrict__ D_ptr, // [nheads] f32 (nullable)
|
||||
const input_t* __restrict__ z_ptr, // [seqlen, nheads, headdim] (nullable)
|
||||
input_t* __restrict__ out_ptr, // [seqlen, nheads, headdim]
|
||||
const int32_t* __restrict__ cu_seqlens, // [batch+1] int32
|
||||
int64_t batch, int64_t nheads, int64_t ngroups, int64_t headdim,
|
||||
int64_t dstate) {
|
||||
using input_vec_t = vec_op::vec_t<input_t>;
|
||||
constexpr int VEC_ELEM_NUM = 8;
|
||||
|
||||
const int64_t nheads_per_group = nheads / ngroups;
|
||||
// states layout: [batch, nheads, headdim, dstate] contiguous (caller
|
||||
// guarantee)
|
||||
const int64_t stride_s_b = nheads * headdim * dstate;
|
||||
const int64_t stride_s_h = headdim * dstate;
|
||||
// stride_s_d = dstate, stride_s_n = 1
|
||||
|
||||
#pragma omp parallel for collapse(2) schedule(static)
|
||||
for (int64_t b = 0; b < batch; ++b) {
|
||||
for (int64_t h = 0; h < nheads; ++h) {
|
||||
const int64_t seq_start = cu_seqlens[b];
|
||||
const int64_t seq_end = cu_seqlens[b + 1];
|
||||
const int64_t g = h / nheads_per_group;
|
||||
|
||||
const float A_val = A_ptr[h];
|
||||
const float D_val = (D_ptr != nullptr) ? D_ptr[h] : 0.0f;
|
||||
|
||||
// Working state slice: states[b, h, :, :] — float32, headdim * dstate.
|
||||
// Fits in L1/L2 for typical dims (e.g. 64*128*4 = 32 KB).
|
||||
float* s_bh = states_ptr + b * stride_s_b + h * stride_s_h;
|
||||
|
||||
for (int64_t t = seq_start; t < seq_end; ++t) {
|
||||
const input_t* x_h = x_ptr + t * nheads * headdim + h * headdim;
|
||||
const float* dt_h = dt_ptr + t * nheads + h;
|
||||
const input_t* B_g = B_ptr + t * ngroups * dstate + g * dstate;
|
||||
const input_t* C_g = C_ptr + t * ngroups * dstate + g * dstate;
|
||||
const input_t* z_h = (z_ptr != nullptr)
|
||||
? z_ptr + t * nheads * headdim + h * headdim
|
||||
: nullptr;
|
||||
input_t* out_h = out_ptr + t * nheads * headdim + h * headdim;
|
||||
|
||||
const float dt_val = *dt_h;
|
||||
const float dA_val = std::exp(A_val * dt_val);
|
||||
const vec_op::FP32Vec8 dA_vec(dA_val); // broadcast scalar
|
||||
const vec_op::FP32Vec8 dt_vec(dt_val);
|
||||
|
||||
for (int64_t d = 0; d < headdim; ++d) {
|
||||
const float x_val = static_cast<float>(x_h[d]);
|
||||
float* s_bhd = s_bh + d * dstate; // [dstate] contiguous float32
|
||||
|
||||
// Vectorised SSM update + readout over dstate:
|
||||
// s_new = s * dA + x * dt * B
|
||||
// y += s_new * C
|
||||
int64_t n = 0;
|
||||
vec_op::FP32Vec8 y_vec(0.0f);
|
||||
const vec_op::FP32Vec8 x_vec(x_val);
|
||||
|
||||
for (; n <= dstate - VEC_ELEM_NUM; n += VEC_ELEM_NUM) {
|
||||
const vec_op::FP32Vec8 B_v((input_vec_t(B_g + n)));
|
||||
const vec_op::FP32Vec8 C_v((input_vec_t(C_g + n)));
|
||||
const vec_op::FP32Vec8 s_v(s_bhd + n);
|
||||
|
||||
const vec_op::FP32Vec8 s_new = s_v * dA_vec + x_vec * dt_vec * B_v;
|
||||
s_new.save(s_bhd + n);
|
||||
y_vec = y_vec + s_new * C_v;
|
||||
}
|
||||
|
||||
float y_val = y_vec.reduce_sum();
|
||||
|
||||
// Scalar tail for remaining dstate elements
|
||||
for (; n < dstate; ++n) {
|
||||
const float B_n = static_cast<float>(B_g[n]);
|
||||
const float C_n = static_cast<float>(C_g[n]);
|
||||
const float s_new = s_bhd[n] * dA_val + x_val * dt_val * B_n;
|
||||
s_bhd[n] = s_new;
|
||||
y_val += s_new * C_n;
|
||||
}
|
||||
|
||||
// D skip connection (scalar per head)
|
||||
if (D_ptr != nullptr) y_val += x_val * D_val;
|
||||
|
||||
// z gating: out = y * z * sigmoid(z) (SiLU)
|
||||
if (z_h != nullptr) {
|
||||
const float z_val = static_cast<float>(z_h[d]);
|
||||
const float sigmoid =
|
||||
(z_val >= 0.0f) ? 1.0f / (1.0f + std::exp(-z_val))
|
||||
: std::exp(z_val) / (1.0f + std::exp(z_val));
|
||||
y_val *= z_val * sigmoid;
|
||||
}
|
||||
|
||||
out_h[d] = static_cast<input_t>(y_val);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace mamba_cpu
|
||||
@@ -213,6 +213,32 @@ void compute_slot_mapping_kernel_impl(const torch::Tensor query_start_loc,
|
||||
torch::Tensor slot_mapping,
|
||||
const int64_t block_size);
|
||||
|
||||
at::Tensor causal_conv1d_update_cpu_impl(
|
||||
at::Tensor& x, at::Tensor& conv_state, const at::Tensor& weight,
|
||||
const c10::optional<at::Tensor>& bias,
|
||||
const c10::optional<std::string>& activation,
|
||||
const c10::optional<at::Tensor>& conv_state_indices,
|
||||
const c10::optional<at::Tensor>& query_start_loc, int64_t pad_slot_id);
|
||||
|
||||
void selective_state_update_cpu_impl(
|
||||
at::Tensor& state, const at::Tensor& x, const at::Tensor& dt,
|
||||
const at::Tensor& A, const at::Tensor& B, const at::Tensor& C,
|
||||
const c10::optional<at::Tensor>& D, const c10::optional<at::Tensor>& z,
|
||||
const c10::optional<at::Tensor>& dt_bias, bool dt_softplus,
|
||||
const c10::optional<at::Tensor>& state_batch_indices,
|
||||
const c10::optional<at::Tensor>& dst_state_batch_indices,
|
||||
int64_t null_block_id, at::Tensor& out,
|
||||
const c10::optional<at::Tensor>& num_accepted_tokens,
|
||||
const c10::optional<at::Tensor>& cu_seqlens);
|
||||
|
||||
void mamba_chunk_scan_fwd_cpu_impl(at::Tensor& out, at::Tensor& final_states,
|
||||
const at::Tensor& x, const at::Tensor& dt,
|
||||
const at::Tensor& A, const at::Tensor& B,
|
||||
const at::Tensor& C,
|
||||
const c10::optional<at::Tensor>& D,
|
||||
const c10::optional<at::Tensor>& z,
|
||||
const at::Tensor& cu_seqlens);
|
||||
|
||||
void init_cpu_memory_env(std::vector<int64_t> node_ids);
|
||||
|
||||
namespace cpu_utils {
|
||||
@@ -570,7 +596,8 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
#endif
|
||||
|
||||
// fused moe
|
||||
#if defined(__AVX512F__) || (defined(ARM_BF16_SUPPORT))
|
||||
#if defined(__AVX512F__) || \
|
||||
(defined(__aarch64__) && !defined(__APPLE__) && defined(ARM_BF16_SUPPORT))
|
||||
ops.def(
|
||||
"prepack_moe_weight(Tensor weight, Tensor(a1!) packed_weight, str isa) "
|
||||
"-> ()");
|
||||
@@ -581,7 +608,7 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"bool skip_weighted, "
|
||||
"str act, str isa) -> ()");
|
||||
ops.impl("cpu_fused_moe", torch::kCPU, &cpu_fused_moe);
|
||||
#endif // #if defined(__AVX512F__) || (defined(ARM_BF16_SUPPORT))
|
||||
#endif
|
||||
ops.def(
|
||||
"mla_decode_kvcache("
|
||||
" Tensor! out, Tensor query, Tensor kv_cache,"
|
||||
@@ -594,6 +621,30 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"block_size) -> ()",
|
||||
&compute_slot_mapping_kernel_impl);
|
||||
|
||||
// Mamba CPU kernels
|
||||
ops.def(
|
||||
"causal_conv1d_update_cpu_vec("
|
||||
"Tensor(a0!) x, Tensor(a1!) conv_state, Tensor weight, "
|
||||
"Tensor? bias, str? activation, Tensor? conv_state_indices, "
|
||||
"Tensor? query_start_loc, SymInt pad_slot_id) -> Tensor",
|
||||
&causal_conv1d_update_cpu_impl);
|
||||
|
||||
ops.def(
|
||||
"selective_state_update_cpu("
|
||||
"Tensor(a0!) state, Tensor x, Tensor dt, Tensor A, Tensor B, Tensor C, "
|
||||
"Tensor? D, Tensor? z, Tensor? dt_bias, bool dt_softplus, "
|
||||
"Tensor? state_batch_indices, Tensor? dst_state_batch_indices, "
|
||||
"SymInt null_block_id, Tensor(a13!) out, "
|
||||
"Tensor? num_accepted_tokens, Tensor? cu_seqlens) -> ()",
|
||||
&selective_state_update_cpu_impl);
|
||||
|
||||
ops.def(
|
||||
"mamba_chunk_scan_fwd_cpu("
|
||||
"Tensor(a0!) out, Tensor(a1!) final_states, "
|
||||
"Tensor x, Tensor dt, Tensor A, Tensor B, Tensor C, "
|
||||
"Tensor? D, Tensor? z, Tensor cu_seqlens) -> ()",
|
||||
&mamba_chunk_scan_fwd_cpu_impl);
|
||||
|
||||
ops.def("init_cpu_memory_env(SymInt[] node_ids) -> ()", &init_cpu_memory_env);
|
||||
|
||||
// Speculative decoding kernels
|
||||
|
||||
@@ -0,0 +1,326 @@
|
||||
#pragma once
|
||||
|
||||
#include "custom_collective_common.cuh"
|
||||
|
||||
namespace vllm {
|
||||
|
||||
constexpr int kMnnvlLamportAgThreads = 128;
|
||||
constexpr int kMnnvlLamportRsThreads = 256;
|
||||
constexpr int kMnnvlLamportConcurrentPollMaxPacks = 8192;
|
||||
|
||||
using CopyPack = array_t<uint64_t, 2>;
|
||||
|
||||
template <int ngpus>
|
||||
__global__ void __launch_bounds__(512, 1)
|
||||
cross_device_all_gather(RankData* _dp, RankSignals sg, Signal* self_sg,
|
||||
CopyPack* __restrict__ result, int rank,
|
||||
int size_per_rank) {
|
||||
auto dp = *_dp;
|
||||
int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
int stride = gridDim.x * blockDim.x;
|
||||
barrier_at_start<ngpus>(sg, self_sg, rank);
|
||||
#pragma unroll
|
||||
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
|
||||
auto src = reinterpret_cast<const CopyPack*>(dp.ptrs[src_rank]);
|
||||
auto dst = result + src_rank * size_per_rank;
|
||||
for (int idx = tid; idx < size_per_rank; idx += stride) {
|
||||
dst[idx] = src[idx];
|
||||
}
|
||||
}
|
||||
barrier_at_end<ngpus, true>(sg, self_sg, rank);
|
||||
}
|
||||
|
||||
template <typename T, int ngpus>
|
||||
__global__ void __launch_bounds__(512, 1)
|
||||
cross_device_reduce_scatter(RankData* _dp, RankSignals sg, Signal* self_sg,
|
||||
T* __restrict__ result, int rank,
|
||||
int size_per_rank) {
|
||||
using P = typename packed_t<T>::P;
|
||||
using A = typename packed_t<T>::A;
|
||||
auto dp = *_dp;
|
||||
auto offset = rank * size_per_rank;
|
||||
barrier_at_start<ngpus>(sg, self_sg, rank);
|
||||
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < size_per_rank;
|
||||
idx += gridDim.x * blockDim.x) {
|
||||
reinterpret_cast<P*>(result)[idx] =
|
||||
packed_reduce<P, ngpus, A>((const P**)&dp.ptrs[0], offset + idx);
|
||||
}
|
||||
barrier_at_end<ngpus, true>(sg, self_sg, rank);
|
||||
}
|
||||
|
||||
template <typename P>
|
||||
union LamportPack {
|
||||
P packed;
|
||||
uint32_t words[sizeof(P) / sizeof(uint32_t)];
|
||||
};
|
||||
|
||||
template <typename P>
|
||||
DINLINE LamportPack<P> load_lamport_pack(const P* ptr) {
|
||||
static_assert(sizeof(P) == 16);
|
||||
LamportPack<P> value;
|
||||
#if !defined(USE_ROCM)
|
||||
asm volatile("ld.volatile.global.v4.u32 {%0, %1, %2, %3}, [%4];"
|
||||
: "=r"(value.words[0]), "=r"(value.words[1]),
|
||||
"=r"(value.words[2]), "=r"(value.words[3])
|
||||
: "l"(ptr)
|
||||
: "memory");
|
||||
#else
|
||||
const volatile uint32_t* src = reinterpret_cast<const volatile uint32_t*>(ptr);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
|
||||
value.words[i] = src[i];
|
||||
}
|
||||
#endif
|
||||
return value;
|
||||
}
|
||||
|
||||
template <typename P>
|
||||
DINLINE bool is_lamport_dirty(const LamportPack<P>& value) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
|
||||
if (value.words[i] == 0x80000000U) return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
template <typename P>
|
||||
DINLINE P lamport_sentinel() {
|
||||
LamportPack<P> value;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
|
||||
value.words[i] = 0x80000000U;
|
||||
}
|
||||
return value.packed;
|
||||
}
|
||||
|
||||
template <typename P>
|
||||
DINLINE P sanitize_lamport_payload(P packed) {
|
||||
LamportPack<P> value{.packed = packed};
|
||||
#pragma unroll
|
||||
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
|
||||
if (value.words[i] == 0x80000000U) value.words[i] = 0;
|
||||
}
|
||||
return value.packed;
|
||||
}
|
||||
|
||||
template <typename P>
|
||||
DINLINE P wait_lamport_payload(const P* ptr) {
|
||||
auto value = load_lamport_pack(ptr);
|
||||
while (is_lamport_dirty(value)) value = load_lamport_pack(ptr);
|
||||
return value.packed;
|
||||
}
|
||||
|
||||
template <typename P, int ngpus>
|
||||
DINLINE void wait_lamport_payloads(const P* base, int rank, int rank_stride,
|
||||
P local_value, P (&values)[ngpus]) {
|
||||
bool ready[ngpus];
|
||||
#pragma unroll
|
||||
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
|
||||
ready[src_rank] = src_rank == rank;
|
||||
if (src_rank == rank) values[src_rank] = local_value;
|
||||
}
|
||||
|
||||
int remaining = ngpus - 1;
|
||||
while (remaining != 0) {
|
||||
#pragma unroll
|
||||
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
|
||||
if (!ready[src_rank]) {
|
||||
auto value = load_lamport_pack(base + src_rank * rank_stride);
|
||||
if (!is_lamport_dirty(value)) {
|
||||
values[src_rank] = value.packed;
|
||||
ready[src_rank] = true;
|
||||
--remaining;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename P, typename A, int ngpus>
|
||||
DINLINE P reduce_lamport_payloads(const P* current_local, const P* packed_input,
|
||||
int rank, int size_per_rank, int idx) {
|
||||
P source_zero =
|
||||
rank == 0 ? packed_input[idx] : wait_lamport_payload(current_local + idx);
|
||||
A tmp = upcast(source_zero);
|
||||
#pragma unroll
|
||||
for (int src_rank = 1; src_rank < ngpus; ++src_rank) {
|
||||
P value = src_rank == rank
|
||||
? packed_input[rank * size_per_rank + idx]
|
||||
: wait_lamport_payload(current_local +
|
||||
src_rank * size_per_rank + idx);
|
||||
packed_assign_add(tmp, upcast(value));
|
||||
}
|
||||
return sanitize_lamport_payload(downcast<P>(tmp));
|
||||
}
|
||||
|
||||
DINLINE void lamport_cta_arrive(uint32_t* counter) {
|
||||
#if !defined(USE_ROCM)
|
||||
if (threadIdx.x < 32) {
|
||||
asm volatile("barrier.cta.sync 1, %0;" : : "r"(blockDim.x) : "memory");
|
||||
if (threadIdx.x == 0) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000
|
||||
asm volatile("red.async.release.global.gpu.add.u32 [%0], 1;"
|
||||
:
|
||||
: "l"(counter)
|
||||
: "memory");
|
||||
#elif defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
|
||||
asm volatile("red.release.global.gpu.add.u32 [%0], 1;"
|
||||
:
|
||||
: "l"(counter)
|
||||
: "memory");
|
||||
#else
|
||||
atomicAdd(counter, 1);
|
||||
#endif
|
||||
}
|
||||
} else {
|
||||
asm volatile("barrier.cta.arrive 1, %0;" : : "r"(blockDim.x) : "memory");
|
||||
}
|
||||
#else
|
||||
__syncthreads();
|
||||
if (threadIdx.x == 0) atomicAdd(counter, 1);
|
||||
#endif
|
||||
}
|
||||
|
||||
template <typename T, int ngpus>
|
||||
__global__ void __launch_bounds__(kMnnvlLamportAgThreads, 1)
|
||||
mnnvl_lamport_all_gather(RankData* _dp, const T* __restrict__ input,
|
||||
T* __restrict__ result,
|
||||
T* __restrict__ multicast_buffer,
|
||||
uint32_t* __restrict__ epochs, int rank,
|
||||
int size_per_rank, int stage_size) {
|
||||
using P = typename packed_t<T>::P;
|
||||
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
|
||||
(__CUDA_ARCH__ >= 900)
|
||||
cudaGridDependencySynchronize();
|
||||
#endif
|
||||
auto dp = *_dp;
|
||||
int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
int stride = gridDim.x * blockDim.x;
|
||||
uint32_t epoch = epochs[0];
|
||||
int current_stage = epoch % 3;
|
||||
int dirty_stage = (epoch + 1) % 3;
|
||||
int dirty_size = epochs[2 + dirty_stage];
|
||||
auto local_buffer = reinterpret_cast<P*>(const_cast<void*>(dp.ptrs[rank]));
|
||||
auto current_local = local_buffer + current_stage * stage_size;
|
||||
auto dirty_local = local_buffer + dirty_stage * stage_size;
|
||||
auto current_multicast =
|
||||
reinterpret_cast<P*>(multicast_buffer) + current_stage * stage_size;
|
||||
auto packed_input = reinterpret_cast<const P*>(input);
|
||||
auto packed_result = reinterpret_cast<P*>(result);
|
||||
|
||||
int total_size = size_per_rank * ngpus;
|
||||
P local_value;
|
||||
if (tid < size_per_rank) {
|
||||
local_value = packed_input[tid];
|
||||
current_multicast[rank * size_per_rank + tid] =
|
||||
sanitize_lamport_payload(local_value);
|
||||
}
|
||||
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
|
||||
(__CUDA_ARCH__ >= 900)
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
#endif
|
||||
|
||||
lamport_cta_arrive(&epochs[1]);
|
||||
|
||||
for (int idx = tid; idx < dirty_size; idx += stride) {
|
||||
dirty_local[idx] = lamport_sentinel<P>();
|
||||
}
|
||||
|
||||
if (tid < size_per_rank) {
|
||||
#pragma unroll
|
||||
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
|
||||
int output_idx = src_rank * size_per_rank + tid;
|
||||
P value = src_rank == rank
|
||||
? local_value
|
||||
: wait_lamport_payload(current_local + output_idx);
|
||||
packed_result[output_idx] = value;
|
||||
}
|
||||
}
|
||||
|
||||
if (tid == 0) {
|
||||
while (*reinterpret_cast<volatile uint32_t*>(&epochs[1]) < gridDim.x);
|
||||
epochs[2 + current_stage] = total_size;
|
||||
epochs[0] = epoch + 1;
|
||||
epochs[1] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, int ngpus>
|
||||
__global__ void __launch_bounds__(kMnnvlLamportRsThreads, 1)
|
||||
mnnvl_lamport_reduce_scatter_kernel(RankData* _dp,
|
||||
const T* __restrict__ input,
|
||||
T* __restrict__ result,
|
||||
uint32_t* __restrict__ epochs, int rank,
|
||||
int size_per_rank, int stage_size) {
|
||||
using P = typename packed_t<T>::P;
|
||||
using A = typename packed_t<T>::A;
|
||||
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
|
||||
(__CUDA_ARCH__ >= 900)
|
||||
cudaGridDependencySynchronize();
|
||||
#endif
|
||||
auto dp = *_dp;
|
||||
int dst_rank = blockIdx.x % ngpus;
|
||||
int tile = blockIdx.x / ngpus;
|
||||
int idx = tile * blockDim.x + threadIdx.x;
|
||||
int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
int stride = gridDim.x * blockDim.x;
|
||||
uint32_t epoch = epochs[0];
|
||||
int current_stage = epoch % 3;
|
||||
int dirty_stage = (epoch + 1) % 3;
|
||||
int dirty_size = epochs[2 + dirty_stage];
|
||||
auto local_buffer = reinterpret_cast<P*>(const_cast<void*>(dp.ptrs[rank]));
|
||||
auto current_local = local_buffer + current_stage * stage_size;
|
||||
auto dirty_local = local_buffer + dirty_stage * stage_size;
|
||||
auto packed_input = reinterpret_cast<const P*>(input);
|
||||
|
||||
if (idx < size_per_rank && dst_rank != rank) {
|
||||
auto dst = reinterpret_cast<P*>(const_cast<void*>(dp.ptrs[dst_rank])) +
|
||||
current_stage * stage_size + rank * size_per_rank;
|
||||
auto src = packed_input + dst_rank * size_per_rank;
|
||||
dst[idx] = sanitize_lamport_payload(src[idx]);
|
||||
}
|
||||
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
|
||||
(__CUDA_ARCH__ >= 900)
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
#endif
|
||||
|
||||
lamport_cta_arrive(&epochs[1]);
|
||||
|
||||
for (int idx = tid; idx < dirty_size; idx += stride) {
|
||||
dirty_local[idx] = lamport_sentinel<P>();
|
||||
}
|
||||
|
||||
if (idx < size_per_rank && dst_rank == rank) {
|
||||
if constexpr (ngpus == 4) {
|
||||
if (size_per_rank > kMnnvlLamportConcurrentPollMaxPacks) {
|
||||
reinterpret_cast<P*>(result)[idx] =
|
||||
reduce_lamport_payloads<P, A, ngpus>(current_local, packed_input,
|
||||
rank, size_per_rank, idx);
|
||||
} else {
|
||||
P values[ngpus];
|
||||
wait_lamport_payloads<P, ngpus>(
|
||||
current_local + idx, rank, size_per_rank,
|
||||
packed_input[rank * size_per_rank + idx], values);
|
||||
A tmp = upcast(values[0]);
|
||||
#pragma unroll
|
||||
for (int src_rank = 1; src_rank < ngpus; ++src_rank) {
|
||||
packed_assign_add(tmp, upcast(values[src_rank]));
|
||||
}
|
||||
reinterpret_cast<P*>(result)[idx] =
|
||||
sanitize_lamport_payload(downcast<P>(tmp));
|
||||
}
|
||||
} else {
|
||||
reinterpret_cast<P*>(result)[idx] = reduce_lamport_payloads<P, A, ngpus>(
|
||||
current_local, packed_input, rank, size_per_rank, idx);
|
||||
}
|
||||
}
|
||||
|
||||
if (tid == 0) {
|
||||
while (*reinterpret_cast<volatile uint32_t*>(&epochs[1]) < gridDim.x);
|
||||
epochs[2 + current_stage] = size_per_rank * ngpus;
|
||||
epochs[0] = epoch + 1;
|
||||
epochs[1] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace vllm
|
||||
+20
-296
@@ -1,299 +1,8 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#if defined(USE_ROCM)
|
||||
typedef __hip_bfloat16 nv_bfloat16;
|
||||
#endif
|
||||
|
||||
#include <iostream>
|
||||
#include <array>
|
||||
#include <limits>
|
||||
#include <map>
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
#include <cstdlib>
|
||||
#include <cstring>
|
||||
#include "custom_collective_common.cuh"
|
||||
|
||||
namespace vllm {
|
||||
#define CUDACHECK(cmd) \
|
||||
do { \
|
||||
cudaError_t e = cmd; \
|
||||
if (e != cudaSuccess) { \
|
||||
printf("Failed: Cuda error %s:%d '%s'\n", __FILE__, __LINE__, \
|
||||
cudaGetErrorString(e)); \
|
||||
exit(EXIT_FAILURE); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
// Maximal number of blocks in allreduce kernel.
|
||||
constexpr int kMaxBlocks = 36;
|
||||
|
||||
// Default number of blocks in allreduce kernel.
|
||||
#ifndef USE_ROCM
|
||||
const int defaultBlockLimit = 36;
|
||||
CUpointer_attribute rangeStartAddrAttr = CU_POINTER_ATTRIBUTE_RANGE_START_ADDR;
|
||||
#else
|
||||
const int defaultBlockLimit = 16;
|
||||
hipPointer_attribute rangeStartAddrAttr =
|
||||
HIP_POINTER_ATTRIBUTE_RANGE_START_ADDR;
|
||||
#endif
|
||||
|
||||
// Counter may overflow, but it's fine since unsigned int overflow is
|
||||
// well-defined behavior.
|
||||
using FlagType = uint32_t;
|
||||
|
||||
// Two sets of peer counters are needed for two syncs: starting and ending an
|
||||
// operation. The reason is that it's possible for peer GPU block to arrive at
|
||||
// the second sync point while the current GPU block haven't passed the first
|
||||
// sync point. Thus, peer GPU may write counter+1 while current GPU is busy
|
||||
// waiting for counter. We use alternating counter array to avoid this
|
||||
// possibility.
|
||||
struct Signal {
|
||||
alignas(128) FlagType start[kMaxBlocks][8];
|
||||
alignas(128) FlagType end[kMaxBlocks][8];
|
||||
alignas(128) FlagType _flag[kMaxBlocks]; // incremental flags for each rank
|
||||
};
|
||||
|
||||
struct __align__(16) RankData {
|
||||
const void* ptrs[8];
|
||||
};
|
||||
|
||||
struct __align__(16) RankSignals {
|
||||
Signal* signals[8];
|
||||
};
|
||||
|
||||
// like std::array, but aligned
|
||||
template <typename T, int sz>
|
||||
struct __align__(alignof(T) * sz) array_t {
|
||||
T data[sz];
|
||||
using type = T;
|
||||
static constexpr int size = sz;
|
||||
};
|
||||
|
||||
// use packed type to maximize memory efficiency
|
||||
// goal: generate ld.128 and st.128 instructions
|
||||
template <typename T>
|
||||
struct packed_t {
|
||||
// the (P)acked type for load/store
|
||||
using P = array_t<T, 16 / sizeof(T)>;
|
||||
// the (A)ccumulator type for reduction
|
||||
using A = array_t<float, 16 / sizeof(T)>;
|
||||
};
|
||||
|
||||
#define DINLINE __device__ __forceinline__
|
||||
|
||||
// scalar cast functions
|
||||
DINLINE float upcast_s(half val) { return __half2float(val); }
|
||||
|
||||
template <typename T>
|
||||
DINLINE T downcast_s(float val);
|
||||
template <>
|
||||
DINLINE half downcast_s(float val) {
|
||||
return __float2half(val);
|
||||
}
|
||||
|
||||
// scalar add functions
|
||||
// for some reason when compiling with Pytorch, the + operator for half and
|
||||
// bfloat is disabled so we call the intrinsics directly
|
||||
DINLINE half& assign_add(half& a, half b) {
|
||||
a = __hadd(a, b);
|
||||
return a;
|
||||
}
|
||||
DINLINE float& assign_add(float& a, float b) { return a += b; }
|
||||
|
||||
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
|
||||
DINLINE float upcast_s(nv_bfloat16 val) { return __bfloat162float(val); }
|
||||
template <>
|
||||
DINLINE nv_bfloat16 downcast_s(float val) {
|
||||
return __float2bfloat16(val);
|
||||
}
|
||||
DINLINE nv_bfloat16& assign_add(nv_bfloat16& a, nv_bfloat16 b) {
|
||||
a = __hadd(a, b);
|
||||
return a;
|
||||
}
|
||||
#endif
|
||||
|
||||
template <typename T, int N>
|
||||
DINLINE array_t<T, N>& packed_assign_add(array_t<T, N>& a, array_t<T, N> b) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < N; i++) {
|
||||
assign_add(a.data[i], b.data[i]);
|
||||
}
|
||||
return a;
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
DINLINE array_t<float, N> upcast(array_t<T, N> val) {
|
||||
if constexpr (std::is_same<T, float>::value) {
|
||||
return val;
|
||||
} else {
|
||||
array_t<float, N> out;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < N; i++) {
|
||||
out.data[i] = upcast_s(val.data[i]);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename O>
|
||||
DINLINE O downcast(array_t<float, O::size> val) {
|
||||
if constexpr (std::is_same<typename O::type, float>::value) {
|
||||
return val;
|
||||
} else {
|
||||
O out;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < O::size; i++) {
|
||||
out.data[i] = downcast_s<typename O::type>(val.data[i]);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
}
|
||||
|
||||
#if !defined(USE_ROCM)
|
||||
|
||||
static DINLINE void st_flag_release(FlagType* flag_addr, FlagType flag) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
|
||||
asm volatile("st.release.sys.global.u32 [%1], %0;" ::"r"(flag),
|
||||
"l"(flag_addr));
|
||||
#else
|
||||
asm volatile("membar.sys; st.volatile.global.u32 [%1], %0;" ::"r"(flag),
|
||||
"l"(flag_addr));
|
||||
#endif
|
||||
}
|
||||
|
||||
static DINLINE FlagType ld_flag_acquire(FlagType* flag_addr) {
|
||||
FlagType flag;
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
|
||||
asm volatile("ld.acquire.sys.global.u32 %0, [%1];"
|
||||
: "=r"(flag)
|
||||
: "l"(flag_addr));
|
||||
#else
|
||||
asm volatile("ld.volatile.global.u32 %0, [%1]; membar.gl;"
|
||||
: "=r"(flag)
|
||||
: "l"(flag_addr));
|
||||
#endif
|
||||
return flag;
|
||||
}
|
||||
|
||||
static DINLINE void st_flag_volatile(FlagType* flag_addr, FlagType flag) {
|
||||
asm volatile("st.volatile.global.u32 [%1], %0;" ::"r"(flag), "l"(flag_addr));
|
||||
}
|
||||
|
||||
static DINLINE FlagType ld_flag_volatile(FlagType* flag_addr) {
|
||||
FlagType flag;
|
||||
asm volatile("ld.volatile.global.u32 %0, [%1];"
|
||||
: "=r"(flag)
|
||||
: "l"(flag_addr));
|
||||
return flag;
|
||||
}
|
||||
|
||||
// This function is meant to be used as the first synchronization in the all
|
||||
// reduce kernel. Thus, it doesn't need to make any visibility guarantees for
|
||||
// prior memory accesses. Note: volatile writes will not be reordered against
|
||||
// other volatile writes.
|
||||
template <int ngpus>
|
||||
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
|
||||
int rank) {
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
auto peer_counter_ptr = &sg.signals[threadIdx.x]->start[blockIdx.x][rank];
|
||||
auto self_counter_ptr = &self_sg->start[blockIdx.x][threadIdx.x];
|
||||
// Write the expected counter value to peer and wait for correct value
|
||||
// from peer.
|
||||
st_flag_volatile(peer_counter_ptr, flag);
|
||||
while (ld_flag_volatile(self_counter_ptr) != flag);
|
||||
}
|
||||
__syncthreads();
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
// This function is meant to be used as the second or the final
|
||||
// synchronization barrier in the all reduce kernel. If it's the final
|
||||
// synchronization barrier, we don't need to make any visibility guarantees
|
||||
// for prior memory accesses.
|
||||
template <int ngpus, bool final_sync = false>
|
||||
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
|
||||
__syncthreads();
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
auto peer_counter_ptr = &sg.signals[threadIdx.x]->end[blockIdx.x][rank];
|
||||
auto self_counter_ptr = &self_sg->end[blockIdx.x][threadIdx.x];
|
||||
// Write the expected counter value to peer and wait for correct value from
|
||||
// peer.
|
||||
if constexpr (!final_sync) {
|
||||
st_flag_release(peer_counter_ptr, flag);
|
||||
while (ld_flag_acquire(self_counter_ptr) != flag);
|
||||
} else {
|
||||
st_flag_volatile(peer_counter_ptr, flag);
|
||||
while (ld_flag_volatile(self_counter_ptr) != flag);
|
||||
}
|
||||
}
|
||||
if constexpr (!final_sync) __syncthreads();
|
||||
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
#else
|
||||
|
||||
template <int ngpus>
|
||||
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
|
||||
int rank) {
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
// simultaneously write to the corresponding flag of all ranks.
|
||||
// Latency = 1 p2p write
|
||||
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->start[blockIdx.x][rank],
|
||||
flag, __ATOMIC_RELAXED, __MEMORY_SCOPE_SYSTEM);
|
||||
// wait until we got true from all ranks
|
||||
while (__scoped_atomic_load_n(&self_sg->start[blockIdx.x][threadIdx.x],
|
||||
__ATOMIC_RELAXED,
|
||||
__MEMORY_SCOPE_DEVICE) < flag);
|
||||
}
|
||||
__syncthreads();
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
template <int ngpus, bool final_sync = false>
|
||||
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
|
||||
__syncthreads();
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
// simultaneously write to the corresponding flag of all ranks.
|
||||
// Latency = 1 p2p write
|
||||
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->end[blockIdx.x][rank],
|
||||
flag,
|
||||
final_sync ? __ATOMIC_RELAXED : __ATOMIC_RELEASE,
|
||||
__MEMORY_SCOPE_SYSTEM);
|
||||
// wait until we got true from all ranks
|
||||
while (
|
||||
__scoped_atomic_load_n(&self_sg->end[blockIdx.x][threadIdx.x],
|
||||
final_sync ? __ATOMIC_RELAXED : __ATOMIC_ACQUIRE,
|
||||
__MEMORY_SCOPE_DEVICE) < flag);
|
||||
}
|
||||
if constexpr (!final_sync) __syncthreads();
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
#endif
|
||||
|
||||
template <typename P, int ngpus, typename A>
|
||||
DINLINE P packed_reduce(const P* ptrs[], int idx) {
|
||||
A tmp = upcast(ptrs[0][idx]);
|
||||
#pragma unroll
|
||||
for (int i = 1; i < ngpus; i++) {
|
||||
packed_assign_add(tmp, upcast(ptrs[i][idx]));
|
||||
}
|
||||
return downcast<P>(tmp);
|
||||
}
|
||||
|
||||
template <typename T, int ngpus>
|
||||
__global__ void __launch_bounds__(512, 1)
|
||||
@@ -616,6 +325,21 @@ class CustomAllreduce {
|
||||
#undef KL
|
||||
}
|
||||
|
||||
void allgather(cudaStream_t stream, void* input, void* output, int size_bytes,
|
||||
int threads = 512, int block_limit = defaultBlockLimit);
|
||||
template <typename T>
|
||||
void mnnvl_lamport_allgather(cudaStream_t stream, T* input, T* output,
|
||||
void* local_buffer, void* multicast_buffer,
|
||||
uint32_t* epochs, int size_bytes,
|
||||
int stage_size_bytes);
|
||||
template <typename T>
|
||||
void reduce_scatter(cudaStream_t stream, T* input, T* output, int size,
|
||||
int threads = 512, int block_limit = defaultBlockLimit);
|
||||
template <typename T>
|
||||
void mnnvl_lamport_reduce_scatter(cudaStream_t stream, T* input, T* output,
|
||||
void* local_buffer, uint32_t* epochs,
|
||||
int size, int stage_size_bytes);
|
||||
|
||||
~CustomAllreduce() {
|
||||
for (auto [_, ptr] : ipc_handles_) {
|
||||
CUDACHECK(cudaIpcCloseMemHandle(ptr));
|
||||
@@ -625,8 +349,8 @@ class CustomAllreduce {
|
||||
|
||||
/**
|
||||
* To inspect PTX/SASS, copy paste this header file to compiler explorer and
|
||||
add a template instantiation:
|
||||
* add a template instantiation:
|
||||
* template void vllm::CustomAllreduce::allreduce<half>(cudaStream_t, half *,
|
||||
half *, int, int, int);
|
||||
*/
|
||||
} // namespace vllm
|
||||
* half *, int, int, int);
|
||||
*/
|
||||
} // namespace vllm
|
||||
|
||||
@@ -0,0 +1,332 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#if defined(USE_ROCM)
|
||||
typedef __hip_bfloat16 nv_bfloat16;
|
||||
#endif
|
||||
|
||||
#include <iostream>
|
||||
#include <array>
|
||||
#include <limits>
|
||||
#include <map>
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
#include <cstdlib>
|
||||
#include <cstring>
|
||||
|
||||
namespace vllm {
|
||||
constexpr int kMaxCustomCollectiveRanks = 16;
|
||||
|
||||
#define CUDACHECK(cmd) \
|
||||
do { \
|
||||
cudaError_t e = cmd; \
|
||||
if (e != cudaSuccess) { \
|
||||
printf("Failed: Cuda error %s:%d '%s'\n", __FILE__, __LINE__, \
|
||||
cudaGetErrorString(e)); \
|
||||
exit(EXIT_FAILURE); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
// Maximal number of blocks in allreduce kernel.
|
||||
constexpr int kMaxBlocks = 36;
|
||||
|
||||
// Default number of blocks in allreduce kernel.
|
||||
#ifndef USE_ROCM
|
||||
inline constexpr int defaultBlockLimit = 36;
|
||||
inline CUpointer_attribute rangeStartAddrAttr =
|
||||
CU_POINTER_ATTRIBUTE_RANGE_START_ADDR;
|
||||
#else
|
||||
inline constexpr int defaultBlockLimit = 16;
|
||||
inline hipPointer_attribute rangeStartAddrAttr =
|
||||
HIP_POINTER_ATTRIBUTE_RANGE_START_ADDR;
|
||||
#endif
|
||||
|
||||
// Counter may overflow, but it's fine since unsigned int overflow is
|
||||
// well-defined behavior.
|
||||
using FlagType = uint32_t;
|
||||
|
||||
// Two sets of peer counters are needed for two syncs: starting and ending an
|
||||
// operation. The reason is that it's possible for peer GPU block to arrive at
|
||||
// the second sync point while the current GPU block haven't passed the first
|
||||
// sync point. Thus, peer GPU may write counter+1 while current GPU is busy
|
||||
// waiting for counter. We use alternating counter array to avoid this
|
||||
// possibility.
|
||||
struct Signal {
|
||||
alignas(128) FlagType start[kMaxBlocks][kMaxCustomCollectiveRanks];
|
||||
alignas(128) FlagType end[kMaxBlocks][kMaxCustomCollectiveRanks];
|
||||
alignas(128) FlagType _flag[kMaxBlocks]; // incremental flags for each rank
|
||||
};
|
||||
|
||||
struct __align__(16) RankData {
|
||||
const void* ptrs[kMaxCustomCollectiveRanks];
|
||||
};
|
||||
|
||||
struct __align__(16) RankSignals {
|
||||
Signal* signals[kMaxCustomCollectiveRanks];
|
||||
};
|
||||
|
||||
// like std::array, but aligned
|
||||
template <typename T, int sz>
|
||||
struct __align__(alignof(T) * sz) array_t {
|
||||
T data[sz];
|
||||
using type = T;
|
||||
static constexpr int size = sz;
|
||||
};
|
||||
|
||||
// use packed type to maximize memory efficiency
|
||||
// goal: generate ld.128 and st.128 instructions
|
||||
template <typename T>
|
||||
struct packed_t {
|
||||
// the (P)acked type for load/store
|
||||
using P = array_t<T, 16 / sizeof(T)>;
|
||||
// the (A)ccumulator type for reduction
|
||||
using A = array_t<float, 16 / sizeof(T)>;
|
||||
};
|
||||
|
||||
#define DINLINE __device__ __forceinline__
|
||||
|
||||
// scalar cast functions
|
||||
DINLINE float upcast_s(half val) { return __half2float(val); }
|
||||
|
||||
template <typename T>
|
||||
DINLINE T downcast_s(float val);
|
||||
template <>
|
||||
DINLINE half downcast_s(float val) {
|
||||
return __float2half(val);
|
||||
}
|
||||
|
||||
// scalar add functions
|
||||
// for some reason when compiling with Pytorch, the + operator for half and
|
||||
// bfloat is disabled so we call the intrinsics directly
|
||||
DINLINE half& assign_add(half& a, half b) {
|
||||
a = __hadd(a, b);
|
||||
return a;
|
||||
}
|
||||
DINLINE float& assign_add(float& a, float b) { return a += b; }
|
||||
|
||||
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
|
||||
DINLINE float upcast_s(nv_bfloat16 val) { return __bfloat162float(val); }
|
||||
template <>
|
||||
DINLINE nv_bfloat16 downcast_s(float val) {
|
||||
return __float2bfloat16(val);
|
||||
}
|
||||
DINLINE nv_bfloat16& assign_add(nv_bfloat16& a, nv_bfloat16 b) {
|
||||
a = __hadd(a, b);
|
||||
return a;
|
||||
}
|
||||
#endif
|
||||
|
||||
template <typename T, int N>
|
||||
DINLINE array_t<T, N>& packed_assign_add(array_t<T, N>& a, array_t<T, N> b) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < N; i++) {
|
||||
assign_add(a.data[i], b.data[i]);
|
||||
}
|
||||
return a;
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
DINLINE array_t<float, N> upcast(array_t<T, N> val) {
|
||||
if constexpr (std::is_same<T, float>::value) {
|
||||
return val;
|
||||
} else {
|
||||
array_t<float, N> out;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < N; i++) {
|
||||
out.data[i] = upcast_s(val.data[i]);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename O>
|
||||
DINLINE O downcast(array_t<float, O::size> val) {
|
||||
if constexpr (std::is_same<typename O::type, float>::value) {
|
||||
return val;
|
||||
} else {
|
||||
O out;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < O::size; i++) {
|
||||
out.data[i] = downcast_s<typename O::type>(val.data[i]);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
}
|
||||
|
||||
#if !defined(USE_ROCM)
|
||||
|
||||
static DINLINE void st_flag_release(FlagType* flag_addr, FlagType flag) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
|
||||
asm volatile("st.release.sys.global.u32 [%1], %0;" ::"r"(flag),
|
||||
"l"(flag_addr));
|
||||
#else
|
||||
asm volatile("membar.sys; st.volatile.global.u32 [%1], %0;" ::"r"(flag),
|
||||
"l"(flag_addr));
|
||||
#endif
|
||||
}
|
||||
|
||||
static DINLINE FlagType ld_flag_acquire(FlagType* flag_addr) {
|
||||
FlagType flag;
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
|
||||
asm volatile("ld.acquire.sys.global.u32 %0, [%1];"
|
||||
: "=r"(flag)
|
||||
: "l"(flag_addr));
|
||||
#else
|
||||
asm volatile("ld.volatile.global.u32 %0, [%1]; membar.gl;"
|
||||
: "=r"(flag)
|
||||
: "l"(flag_addr));
|
||||
#endif
|
||||
return flag;
|
||||
}
|
||||
|
||||
static DINLINE void st_flag_volatile(FlagType* flag_addr, FlagType flag) {
|
||||
asm volatile("st.volatile.global.u32 [%1], %0;" ::"r"(flag), "l"(flag_addr));
|
||||
}
|
||||
|
||||
static DINLINE FlagType ld_flag_volatile(FlagType* flag_addr) {
|
||||
FlagType flag;
|
||||
asm volatile("ld.volatile.global.u32 %0, [%1];"
|
||||
: "=r"(flag)
|
||||
: "l"(flag_addr));
|
||||
return flag;
|
||||
}
|
||||
|
||||
// This function is meant to be used as the first synchronization in the all
|
||||
// reduce kernel. Thus, it doesn't need to make any visibility guarantees for
|
||||
// prior memory accesses. Note: volatile writes will not be reordered against
|
||||
// other volatile writes.
|
||||
template <int ngpus>
|
||||
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
|
||||
int rank) {
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
auto peer_counter_ptr = &sg.signals[threadIdx.x]->start[blockIdx.x][rank];
|
||||
auto self_counter_ptr = &self_sg->start[blockIdx.x][threadIdx.x];
|
||||
// Write the expected counter value to peer and wait for correct value
|
||||
// from peer.
|
||||
st_flag_volatile(peer_counter_ptr, flag);
|
||||
while (ld_flag_volatile(self_counter_ptr) != flag);
|
||||
}
|
||||
__syncthreads();
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
template <int ngpus>
|
||||
DINLINE void barrier_at_start_release(const RankSignals& sg, Signal* self_sg,
|
||||
int rank) {
|
||||
__syncthreads();
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
auto peer_counter_ptr = &sg.signals[threadIdx.x]->start[blockIdx.x][rank];
|
||||
auto self_counter_ptr = &self_sg->start[blockIdx.x][threadIdx.x];
|
||||
st_flag_release(peer_counter_ptr, flag);
|
||||
while (ld_flag_acquire(self_counter_ptr) != flag);
|
||||
}
|
||||
__syncthreads();
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
// This function is meant to be used as the second or the final
|
||||
// synchronization barrier in the all reduce kernel. If it's the final
|
||||
// synchronization barrier, we don't need to make any visibility guarantees
|
||||
// for prior memory accesses.
|
||||
template <int ngpus, bool final_sync = false>
|
||||
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
|
||||
__syncthreads();
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
auto peer_counter_ptr = &sg.signals[threadIdx.x]->end[blockIdx.x][rank];
|
||||
auto self_counter_ptr = &self_sg->end[blockIdx.x][threadIdx.x];
|
||||
// Write the expected counter value to peer and wait for correct value from
|
||||
// peer.
|
||||
if constexpr (!final_sync) {
|
||||
st_flag_release(peer_counter_ptr, flag);
|
||||
while (ld_flag_acquire(self_counter_ptr) != flag);
|
||||
} else {
|
||||
st_flag_volatile(peer_counter_ptr, flag);
|
||||
while (ld_flag_volatile(self_counter_ptr) != flag);
|
||||
}
|
||||
}
|
||||
if constexpr (!final_sync) __syncthreads();
|
||||
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
#else
|
||||
|
||||
template <int ngpus>
|
||||
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
|
||||
int rank) {
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
// simultaneously write to the corresponding flag of all ranks.
|
||||
// Latency = 1 p2p write
|
||||
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->start[blockIdx.x][rank],
|
||||
flag, __ATOMIC_RELAXED, __MEMORY_SCOPE_SYSTEM);
|
||||
// wait until we got true from all ranks
|
||||
while (__scoped_atomic_load_n(&self_sg->start[blockIdx.x][threadIdx.x],
|
||||
__ATOMIC_RELAXED,
|
||||
__MEMORY_SCOPE_DEVICE) < flag);
|
||||
}
|
||||
__syncthreads();
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
template <int ngpus>
|
||||
DINLINE void barrier_at_start_release(const RankSignals& sg, Signal* self_sg,
|
||||
int rank) {
|
||||
__syncthreads();
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->start[blockIdx.x][rank],
|
||||
flag, __ATOMIC_RELEASE, __MEMORY_SCOPE_SYSTEM);
|
||||
while (__scoped_atomic_load_n(&self_sg->start[blockIdx.x][threadIdx.x],
|
||||
__ATOMIC_ACQUIRE,
|
||||
__MEMORY_SCOPE_DEVICE) < flag);
|
||||
}
|
||||
__syncthreads();
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
template <int ngpus, bool final_sync = false>
|
||||
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
|
||||
__syncthreads();
|
||||
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
|
||||
if (threadIdx.x < ngpus) {
|
||||
// simultaneously write to the corresponding flag of all ranks.
|
||||
// Latency = 1 p2p write
|
||||
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->end[blockIdx.x][rank],
|
||||
flag,
|
||||
final_sync ? __ATOMIC_RELAXED : __ATOMIC_RELEASE,
|
||||
__MEMORY_SCOPE_SYSTEM);
|
||||
// wait until we got true from all ranks
|
||||
while (
|
||||
__scoped_atomic_load_n(&self_sg->end[blockIdx.x][threadIdx.x],
|
||||
final_sync ? __ATOMIC_RELAXED : __ATOMIC_ACQUIRE,
|
||||
__MEMORY_SCOPE_DEVICE) < flag);
|
||||
}
|
||||
if constexpr (!final_sync) __syncthreads();
|
||||
// use one thread to update flag
|
||||
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
|
||||
}
|
||||
|
||||
#endif
|
||||
|
||||
template <typename P, int ngpus, typename A>
|
||||
DINLINE P packed_reduce(const P* ptrs[], int idx) {
|
||||
A tmp = upcast(ptrs[0][idx]);
|
||||
#pragma unroll
|
||||
for (int i = 1; i < ngpus; i++) {
|
||||
packed_assign_add(tmp, upcast(ptrs[i][idx]));
|
||||
}
|
||||
return downcast<P>(tmp);
|
||||
}
|
||||
|
||||
} // namespace vllm
|
||||
@@ -0,0 +1,17 @@
|
||||
#include "core/registration.h"
|
||||
#include "flash_kda.h"
|
||||
|
||||
TORCH_LIBRARY(_flashkda_C, m) {
|
||||
m.def("get_workspace_size(int T_total, int H, int N=1) -> int",
|
||||
&get_workspace_size);
|
||||
m.def(
|
||||
"fwd(Tensor q, Tensor k, Tensor v, Tensor g, Tensor beta, float scale, "
|
||||
"Tensor(a!) out, Tensor workspace, Tensor A_log, Tensor dt_bias, "
|
||||
"float lower_bound, "
|
||||
"Tensor? initial_state=None, Tensor(b!)? final_state=None, "
|
||||
"Tensor? cu_seqlens=None) -> ()");
|
||||
}
|
||||
|
||||
TORCH_LIBRARY_IMPL(_flashkda_C, CUDA, m) { m.impl("fwd", &fwd); }
|
||||
|
||||
REGISTER_EXTENSION(_flashkda_C)
|
||||
@@ -464,6 +464,66 @@ __global__ void swigluoai_and_mul_kernel(
|
||||
}
|
||||
}
|
||||
|
||||
// SITU (Kimi SituGLU) gated activation. Non-interleaved layout:
|
||||
// input = [gate(d), up(d)] per token.
|
||||
// gate_out = beta * tanh(gate / beta) * sigmoid(gate)
|
||||
// up_out = (linear_beta > 0) ? linear_beta * tanh(up / linear_beta) : up
|
||||
// out = gate_out * up_out
|
||||
// Compute is done in fp32 and written straight to `out` -- no intermediate
|
||||
// tensors and no full-tensor fp32 upcast (the pure-torch forward_native
|
||||
// allocated ~8 fp32 temporaries per call, which blows up MoE profiling).
|
||||
template <typename scalar_t>
|
||||
__global__ void situ_and_mul_kernel(
|
||||
scalar_t* __restrict__ out, // [..., d]
|
||||
const scalar_t* __restrict__ input, // [..., 2, d]
|
||||
const int d, const float beta, const float linear_beta) {
|
||||
const int64_t row = blockIdx.x;
|
||||
const scalar_t* gate_ptr = input + row * 2 * d;
|
||||
const scalar_t* up_ptr = gate_ptr + d;
|
||||
scalar_t* out_ptr = out + row * d;
|
||||
const bool clamp_up = linear_beta > 0.0f;
|
||||
const float inv_beta = 1.0f / beta;
|
||||
const float inv_linear_beta = clamp_up ? 1.0f / linear_beta : 0.0f;
|
||||
for (int64_t idx = threadIdx.x; idx < d; idx += blockDim.x) {
|
||||
const float g = (float)VLLM_LDG(&gate_ptr[idx]);
|
||||
const float u = (float)VLLM_LDG(&up_ptr[idx]);
|
||||
const float gate_out = beta * tanhf(g * inv_beta) / (1.0f + expf(-g));
|
||||
const float up_out =
|
||||
clamp_up ? linear_beta * tanhf(u * inv_linear_beta) : u;
|
||||
out_ptr[idx] = (scalar_t)(gate_out * up_out);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
__global__ void masked_situ_and_mul_kernel(
|
||||
scalar_t* __restrict__ out, const scalar_t* __restrict__ input,
|
||||
const int* __restrict__ expert_num_tokens, const int max_num_tokens,
|
||||
const int d, const float beta, const float linear_beta) {
|
||||
const int expert = blockIdx.y;
|
||||
const int num_tokens = expert_num_tokens[expert];
|
||||
const int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
if (idx >= d || num_tokens == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const bool clamp_up = linear_beta > 0.0f;
|
||||
const float inv_beta = 1.0f / beta;
|
||||
const float inv_linear_beta = clamp_up ? 1.0f / linear_beta : 0.0f;
|
||||
const int64_t expert_row = static_cast<int64_t>(expert) * max_num_tokens;
|
||||
for (int token = 0; token < num_tokens; ++token) {
|
||||
const int64_t row = expert_row + token;
|
||||
const scalar_t* gate_ptr = input + row * 2 * d;
|
||||
const scalar_t* up_ptr = gate_ptr + d;
|
||||
scalar_t* out_ptr = out + row * d;
|
||||
const float g = (float)VLLM_LDG(&gate_ptr[idx]);
|
||||
const float u = (float)VLLM_LDG(&up_ptr[idx]);
|
||||
const float gate_out = beta * tanhf(g * inv_beta) / (1.0f + expf(-g));
|
||||
const float up_out =
|
||||
clamp_up ? linear_beta * tanhf(u * inv_linear_beta) : u;
|
||||
out_ptr[idx] = (scalar_t)(gate_out * up_out);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace vllm
|
||||
|
||||
#define LAUNCH_ACTIVATION_GATE_KERNEL_WITH_PARAM(KERNEL, PACKED_KERNEL, PARAM) \
|
||||
@@ -553,6 +613,54 @@ void swigluoai_and_mul(torch::stable::Tensor& out, // [..., d]
|
||||
double alpha, double limit) {
|
||||
LAUNCH_SIGLUOAI_AND_MUL(vllm::swigluoai_and_mul, alpha, limit);
|
||||
}
|
||||
|
||||
// Kimi SITU gated activation. `linear_beta <= 0` means "unset" (up passed
|
||||
// through), matching SituAndMul(linear_beta=None) on the Python side.
|
||||
void situ_and_mul(torch::stable::Tensor& out, // [..., d]
|
||||
torch::stable::Tensor& input, // [..., 2 * d]
|
||||
double beta, double linear_beta) {
|
||||
int d = input.size(-1) / 2;
|
||||
int64_t num_tokens = input.numel() / input.size(-1);
|
||||
if (num_tokens == 0) {
|
||||
return;
|
||||
}
|
||||
dim3 grid(num_tokens);
|
||||
dim3 block(std::min(d, 1024));
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
input.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream();
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
|
||||
input.scalar_type(), "situ_and_mul_kernel", [&] {
|
||||
vllm::situ_and_mul_kernel<scalar_t><<<grid, block, 0, stream>>>(
|
||||
out.mutable_data_ptr<scalar_t>(), input.const_data_ptr<scalar_t>(),
|
||||
d, (float)beta, (float)linear_beta);
|
||||
});
|
||||
}
|
||||
|
||||
void masked_situ_and_mul(torch::stable::Tensor& out, // [E, T, d]
|
||||
torch::stable::Tensor& input, // [E, T, 2 * d]
|
||||
const torch::stable::Tensor& expert_num_tokens,
|
||||
double beta, double linear_beta) {
|
||||
int num_experts = input.size(0);
|
||||
int max_num_tokens = input.size(1);
|
||||
int d = input.size(2) / 2;
|
||||
if (num_experts == 0 || max_num_tokens == 0) {
|
||||
return;
|
||||
}
|
||||
constexpr int block_size = 256;
|
||||
dim3 grid((d + block_size - 1) / block_size, num_experts);
|
||||
dim3 block(block_size);
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
input.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream();
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
|
||||
input.scalar_type(), "masked_situ_and_mul_kernel", [&] {
|
||||
vllm::masked_situ_and_mul_kernel<scalar_t><<<grid, block, 0, stream>>>(
|
||||
out.mutable_data_ptr<scalar_t>(), input.const_data_ptr<scalar_t>(),
|
||||
expert_num_tokens.const_data_ptr<int>(), max_num_tokens, d,
|
||||
(float)beta, (float)linear_beta);
|
||||
});
|
||||
}
|
||||
namespace vllm {
|
||||
|
||||
// Element-wise activation kernel template.
|
||||
|
||||
@@ -21,7 +21,10 @@ __global__ void merge_attn_states_kernel(
|
||||
const float* prefix_lse, const scalar_t* suffix_output,
|
||||
const float* suffix_lse, const uint num_tokens, const uint num_heads,
|
||||
const uint head_size, const uint prefix_head_stride,
|
||||
const uint output_head_stride, const uint prefix_num_tokens,
|
||||
const uint output_head_stride, const uint prefix_lse_head_stride,
|
||||
const uint prefix_lse_token_stride, const uint suffix_lse_head_stride,
|
||||
const uint suffix_lse_token_stride, const uint output_lse_head_stride,
|
||||
const uint output_lse_token_stride, const uint prefix_num_tokens,
|
||||
const float* output_scale) {
|
||||
// Inputs always load 128-bit packs (pack_size elements of scalar_t).
|
||||
// Outputs store pack_size elements of output_t, which is smaller for FP8.
|
||||
@@ -84,15 +87,19 @@ __global__ void merge_attn_states_kernel(
|
||||
}
|
||||
}
|
||||
if (output_lse != nullptr && pack_idx == 0) {
|
||||
float s_lse = suffix_lse[head_idx * num_tokens + token_idx];
|
||||
output_lse[head_idx * num_tokens + token_idx] = s_lse;
|
||||
float s_lse = suffix_lse[head_idx * suffix_lse_head_stride +
|
||||
token_idx * suffix_lse_token_stride];
|
||||
output_lse[head_idx * output_lse_head_stride +
|
||||
token_idx * output_lse_token_stride] = s_lse;
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
// For tokens within prefix range, merge prefix and suffix
|
||||
float p_lse = prefix_lse[head_idx * num_tokens + token_idx];
|
||||
float s_lse = suffix_lse[head_idx * num_tokens + token_idx];
|
||||
float p_lse = prefix_lse[head_idx * prefix_lse_head_stride +
|
||||
token_idx * prefix_lse_token_stride];
|
||||
float s_lse = suffix_lse[head_idx * suffix_lse_head_stride +
|
||||
token_idx * suffix_lse_token_stride];
|
||||
p_lse = std::isinf(p_lse) ? -std::numeric_limits<float>::infinity() : p_lse;
|
||||
s_lse = std::isinf(s_lse) ? -std::numeric_limits<float>::infinity() : s_lse;
|
||||
|
||||
@@ -132,7 +139,8 @@ __global__ void merge_attn_states_kernel(
|
||||
}
|
||||
// We only need to write to output_lse once per head.
|
||||
if (output_lse != nullptr && pack_idx == 0) {
|
||||
output_lse[head_idx * num_tokens + token_idx] = max_lse;
|
||||
output_lse[head_idx * output_lse_head_stride +
|
||||
token_idx * output_lse_token_stride] = max_lse;
|
||||
}
|
||||
return;
|
||||
}
|
||||
@@ -187,7 +195,8 @@ __global__ void merge_attn_states_kernel(
|
||||
// We only need to write to output_lse once per head.
|
||||
if (output_lse != nullptr && pack_idx == 0) {
|
||||
float out_lse = logf(out_se) + max_lse;
|
||||
output_lse[head_idx * num_tokens + token_idx] = out_lse;
|
||||
output_lse[head_idx * output_lse_head_stride +
|
||||
token_idx * output_lse_token_stride] = out_lse;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -221,6 +230,9 @@ __global__ void merge_attn_states_kernel(
|
||||
reinterpret_cast<scalar_t*>(suffix_output.data_ptr()), \
|
||||
reinterpret_cast<float*>(suffix_lse.data_ptr()), num_tokens, \
|
||||
num_heads, head_size, prefix_head_stride, output_head_stride, \
|
||||
prefix_lse_head_stride, prefix_lse_token_stride, \
|
||||
suffix_lse_head_stride, suffix_lse_token_stride, \
|
||||
output_lse_head_stride, output_lse_token_stride, \
|
||||
prefix_num_tokens, output_scale_ptr); \
|
||||
}
|
||||
|
||||
@@ -259,6 +271,19 @@ void merge_attn_states_launcher(
|
||||
const uint head_size = output.size(2);
|
||||
const uint prefix_head_stride = prefix_output.stride(1);
|
||||
const uint output_head_stride = output.stride(1);
|
||||
// lse tensors are [NUM_HEADS, NUM_TOKENS] but may be non-contiguous views
|
||||
// (e.g. a transpose of a backend's [NUM_TOKENS, NUM_HEADS] output), so index
|
||||
// them by their actual strides rather than assuming a contiguous layout.
|
||||
const uint prefix_lse_head_stride = prefix_lse.stride(0);
|
||||
const uint prefix_lse_token_stride = prefix_lse.stride(1);
|
||||
const uint suffix_lse_head_stride = suffix_lse.stride(0);
|
||||
const uint suffix_lse_token_stride = suffix_lse.stride(1);
|
||||
uint output_lse_head_stride = 0;
|
||||
uint output_lse_token_stride = 0;
|
||||
if (output_lse.has_value()) {
|
||||
output_lse_head_stride = output_lse.value().stride(0);
|
||||
output_lse_token_stride = output_lse.value().stride(1);
|
||||
}
|
||||
// Thread mapping is based on input BF16 pack_size
|
||||
const uint pack_size = 16 / sizeof(scalar_t);
|
||||
STD_TORCH_CHECK(head_size % pack_size == 0,
|
||||
|
||||
@@ -443,6 +443,55 @@ __global__ void concat_and_cache_mla_kernel(
|
||||
copy(k_pe, kv_cache, k_pe_stride, block_stride, pe_dim, kv_lora_rank);
|
||||
}
|
||||
|
||||
// Grouped variant of concat_and_cache_mla: inserts the context K/V for every
|
||||
// draft layer in a single launch. Grid is (num_tokens, num_layers); each layer
|
||||
// reads its own cache base pointer from kv_cache_ptrs (same pointer-array
|
||||
// pattern as copy_blocks_kernel). bf16 only, so it is a raw 16-bit copy with no
|
||||
// scaling or quantization; scalar_t is uint16_t for portability.
|
||||
template <typename scalar_t>
|
||||
__global__ void concat_and_cache_mla_grouped_kernel(
|
||||
const scalar_t* __restrict__ kv_c, // [num_layers, num_tokens,
|
||||
// kv_lora_rank]
|
||||
const scalar_t* __restrict__ k_pe, // [num_layers, num_tokens, pe_dim]
|
||||
const int64_t* __restrict__ kv_cache_ptrs, // [num_layers]
|
||||
const int64_t* __restrict__ slot_mapping, // [num_layers, num_tokens]
|
||||
const int64_t kv_c_layer_stride, const int64_t kv_c_token_stride,
|
||||
const int64_t k_pe_layer_stride, const int64_t k_pe_token_stride,
|
||||
const int64_t slot_layer_stride, const int64_t block_stride,
|
||||
const int64_t entry_stride, const int kv_lora_rank, const int pe_dim,
|
||||
const int block_size) {
|
||||
const int64_t token_idx = blockIdx.x;
|
||||
const int64_t layer_idx = blockIdx.y;
|
||||
const int64_t slot_idx =
|
||||
slot_mapping[layer_idx * slot_layer_stride + token_idx];
|
||||
// NOTE: slot_idx can be -1 if the token is padded
|
||||
if (slot_idx < 0) {
|
||||
return;
|
||||
}
|
||||
const int64_t block_idx = slot_idx / block_size;
|
||||
const int64_t block_offset = slot_idx % block_size;
|
||||
|
||||
scalar_t* __restrict__ kv_cache =
|
||||
reinterpret_cast<scalar_t*>(kv_cache_ptrs[layer_idx]);
|
||||
const scalar_t* __restrict__ kv_c_layer =
|
||||
kv_c + layer_idx * kv_c_layer_stride;
|
||||
const scalar_t* __restrict__ k_pe_layer =
|
||||
k_pe + layer_idx * k_pe_layer_stride;
|
||||
|
||||
auto copy = [&](const scalar_t* __restrict__ src, int64_t src_token_stride,
|
||||
int size, int offset) {
|
||||
for (int i = threadIdx.x; i < size; i += blockDim.x) {
|
||||
const int64_t src_idx = token_idx * src_token_stride + i;
|
||||
const int64_t dst_idx =
|
||||
block_idx * block_stride + block_offset * entry_stride + i + offset;
|
||||
kv_cache[dst_idx] = src[src_idx];
|
||||
}
|
||||
};
|
||||
|
||||
copy(kv_c_layer, kv_c_token_stride, kv_lora_rank, 0);
|
||||
copy(k_pe_layer, k_pe_token_stride, pe_dim, kv_lora_rank);
|
||||
}
|
||||
|
||||
template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
|
||||
__global__ void concat_and_cache_ds_mla_kernel(
|
||||
const scalar_t* __restrict__ kv_c, // [num_tokens, kv_lora_rank]
|
||||
@@ -902,6 +951,53 @@ void concat_and_cache_mla(
|
||||
}
|
||||
}
|
||||
|
||||
void concat_and_cache_mla_grouped(
|
||||
torch::stable::Tensor& kv_c, // [num_layers, num_tokens, kv_lora_rank]
|
||||
torch::stable::Tensor& k_pe, // [num_layers, num_tokens, pe_dim]
|
||||
torch::stable::Tensor& kv_cache_ptrs, // [num_layers] int64, on device
|
||||
torch::stable::Tensor& slot_mapping, // [num_layers, num_tokens] int64
|
||||
int64_t block_size, int64_t block_stride, int64_t entry_stride) {
|
||||
int num_layers = kv_c.size(0);
|
||||
int num_tokens = kv_c.size(1);
|
||||
int kv_lora_rank = kv_c.size(2);
|
||||
int pe_dim = k_pe.size(2);
|
||||
|
||||
STD_TORCH_CHECK(
|
||||
kv_c.scalar_type() == torch::headeronly::ScalarType::BFloat16 &&
|
||||
k_pe.scalar_type() == torch::headeronly::ScalarType::BFloat16,
|
||||
"concat_and_cache_mla_grouped only supports a bf16 KV cache; got kv_c=",
|
||||
kv_c.scalar_type(), ", k_pe=", k_pe.scalar_type());
|
||||
STD_TORCH_CHECK(
|
||||
kv_cache_ptrs.scalar_type() == torch::headeronly::ScalarType::Long,
|
||||
"kv_cache_ptrs must be int64");
|
||||
|
||||
if (num_tokens == 0 || num_layers == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int64_t kv_c_layer_stride = kv_c.stride(0);
|
||||
const int64_t kv_c_token_stride = kv_c.stride(1);
|
||||
const int64_t k_pe_layer_stride = k_pe.stride(0);
|
||||
const int64_t k_pe_token_stride = k_pe.stride(1);
|
||||
const int64_t slot_layer_stride = slot_mapping.stride(0);
|
||||
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
kv_c.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream();
|
||||
|
||||
dim3 grid(num_tokens, num_layers);
|
||||
dim3 block(std::min(kv_lora_rank, 512));
|
||||
vllm::concat_and_cache_mla_grouped_kernel<uint16_t>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<const uint16_t*>(kv_c.data_ptr()),
|
||||
reinterpret_cast<const uint16_t*>(k_pe.data_ptr()),
|
||||
kv_cache_ptrs.const_data_ptr<int64_t>(),
|
||||
slot_mapping.const_data_ptr<int64_t>(), kv_c_layer_stride,
|
||||
kv_c_token_stride, k_pe_layer_stride, k_pe_token_stride,
|
||||
slot_layer_stride, block_stride, entry_stride, kv_lora_rank, pe_dim,
|
||||
block_size);
|
||||
}
|
||||
|
||||
namespace vllm {
|
||||
|
||||
template <typename Tout, typename Tin, Fp8KVCacheDataType kv_dt>
|
||||
@@ -1025,6 +1121,9 @@ __global__ void gather_and_maybe_dequant_cache(
|
||||
batch_offset += offset;
|
||||
int32_t block_table_id = batch_offset / block_size;
|
||||
int32_t slot_id = batch_offset % block_size;
|
||||
// seq_starts may push the block index past the end of the batch's block
|
||||
// table row.
|
||||
if (block_table_id >= block_table_stride) continue;
|
||||
int32_t block_table_offset = batch_id * block_table_stride + block_table_id;
|
||||
int32_t block_id = block_table[block_table_offset];
|
||||
int64_t cache_offset =
|
||||
@@ -1174,7 +1273,8 @@ __global__ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
const int32_t num_reqs, const int32_t block_size,
|
||||
const int32_t total_tokens, const int64_t block_table_stride,
|
||||
const int64_t cache_block_stride, const int64_t cache_entry_stride,
|
||||
const int64_t dst_entry_stride) {
|
||||
const int64_t dst_entry_stride,
|
||||
const int32_t* __restrict__ seq_starts) { // Optional source offsets
|
||||
const int flat_warp_id = (blockIdx.x * blockDim.x + threadIdx.x) >> 5;
|
||||
if (flat_warp_id >= total_tokens) return;
|
||||
const int lane_id = threadIdx.x & 31;
|
||||
@@ -1192,7 +1292,8 @@ __global__ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
|
||||
// Compute physical token address via block table
|
||||
const int out_token_id = flat_warp_id;
|
||||
const int token_offset = out_token_id - workspace_starts[req_id];
|
||||
int token_offset = out_token_id - workspace_starts[req_id];
|
||||
if (seq_starts != nullptr) token_offset += seq_starts[req_id];
|
||||
const int cache_block_idx = token_offset / block_size;
|
||||
const int offset_in_block = token_offset % block_size;
|
||||
const int physical_block =
|
||||
@@ -1383,9 +1484,9 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
torch::stable::Tensor const& src_cache, // [NUM_BLOCKS, BLOCK_SIZE, 656]
|
||||
torch::stable::Tensor const& dst, // [TOT_TOKENS, 576]
|
||||
torch::stable::Tensor const& block_table, // [BATCH, BLOCK_INDICES]
|
||||
torch::stable::Tensor const& seq_lens, // [BATCH]
|
||||
torch::stable::Tensor const& workspace_starts, // [BATCH]
|
||||
int64_t batch_size) {
|
||||
int64_t batch_size,
|
||||
std::optional<torch::stable::Tensor> seq_starts = std::nullopt) {
|
||||
torch::stable::accelerator::DeviceGuard device_guard(
|
||||
src_cache.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream();
|
||||
@@ -1396,20 +1497,25 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
STD_TORCH_CHECK(
|
||||
block_table.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"block_table must be int32");
|
||||
STD_TORCH_CHECK(seq_lens.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"seq_lens must be int32");
|
||||
STD_TORCH_CHECK(
|
||||
workspace_starts.scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"workspace_starts must be int32");
|
||||
if (seq_starts.has_value()) {
|
||||
STD_TORCH_CHECK(
|
||||
seq_starts.value().scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"seq_starts must be int32");
|
||||
}
|
||||
|
||||
STD_TORCH_CHECK(src_cache.device() == dst.device(),
|
||||
"src_cache and dst must be on the same device");
|
||||
STD_TORCH_CHECK(src_cache.device() == block_table.device(),
|
||||
"src_cache and block_table must be on the same device");
|
||||
STD_TORCH_CHECK(src_cache.device() == seq_lens.device(),
|
||||
"src_cache and seq_lens must be on the same device");
|
||||
STD_TORCH_CHECK(src_cache.device() == workspace_starts.device(),
|
||||
"src_cache and workspace_starts must be on the same device");
|
||||
if (seq_starts.has_value()) {
|
||||
STD_TORCH_CHECK(src_cache.device() == seq_starts.value().device(),
|
||||
"src_cache and seq_starts must be on the same device");
|
||||
}
|
||||
auto dtype = src_cache.scalar_type();
|
||||
STD_TORCH_CHECK(
|
||||
dtype == torch::headeronly::ScalarType::Byte || // uint8
|
||||
@@ -1438,6 +1544,9 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
constexpr int warps_per_block = 8;
|
||||
const int grid_size = (total_tokens + warps_per_block - 1) / warps_per_block;
|
||||
const int block_size_threads = warps_per_block * 32; // 256 threads
|
||||
const int32_t* seq_starts_ptr =
|
||||
seq_starts.has_value() ? seq_starts.value().const_data_ptr<int32_t>()
|
||||
: nullptr;
|
||||
|
||||
vllm::cp_gather_and_upconvert_fp8_kv_cache<<<grid_size, block_size_threads, 0,
|
||||
stream>>>(
|
||||
@@ -1446,7 +1555,7 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
workspace_starts.const_data_ptr<int32_t>(),
|
||||
static_cast<int32_t>(batch_size), block_size, total_tokens,
|
||||
block_table_stride, cache_block_stride, cache_entry_stride,
|
||||
dst_entry_stride);
|
||||
dst_entry_stride, seq_starts_ptr);
|
||||
}
|
||||
|
||||
// Macro to dispatch the kernel based on the data type.
|
||||
|
||||
@@ -0,0 +1,362 @@
|
||||
#include "torch_utils.h"
|
||||
|
||||
#include <torch/csrc/stable/macros.h>
|
||||
#include <torch/csrc/stable/accelerator.h>
|
||||
#include <torch/csrc/stable/tensor.h>
|
||||
#include <torch/headeronly/core/ScalarType.h>
|
||||
|
||||
#include "custom_all_reduce.cuh"
|
||||
#include "custom_all_gather_reduce_scatter.cuh"
|
||||
|
||||
namespace vllm {
|
||||
|
||||
void CustomAllreduce::allgather(cudaStream_t stream, void* input, void* output,
|
||||
int size_bytes, int threads, int block_limit) {
|
||||
if (size_bytes % sizeof(CopyPack) != 0)
|
||||
throw std::runtime_error(
|
||||
"custom allgather requires input byte size to be a multiple of " +
|
||||
std::to_string(sizeof(CopyPack)));
|
||||
|
||||
auto ptrs = buffers_.at(input);
|
||||
int size_per_rank = size_bytes / sizeof(CopyPack);
|
||||
int total_size = size_per_rank * world_size_;
|
||||
int blocks = std::min(block_limit, (total_size + threads - 1) / threads);
|
||||
|
||||
#define AG_CASE(ngpus) \
|
||||
case ngpus: \
|
||||
cross_device_all_gather<ngpus><<<blocks, threads, 0, stream>>>( \
|
||||
ptrs, sg_, self_sg_, reinterpret_cast<CopyPack*>(output), rank_, \
|
||||
size_per_rank); \
|
||||
break;
|
||||
|
||||
switch (world_size_) {
|
||||
AG_CASE(2)
|
||||
AG_CASE(4)
|
||||
AG_CASE(6)
|
||||
AG_CASE(8)
|
||||
default:
|
||||
throw std::runtime_error(
|
||||
"custom allgather only supports num gpus in (2,4,6,8)");
|
||||
}
|
||||
#undef AG_CASE
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void CustomAllreduce::mnnvl_lamport_allgather(cudaStream_t stream, T* input,
|
||||
T* output, void* local_buffer,
|
||||
void* multicast_buffer,
|
||||
uint32_t* epochs, int size_bytes,
|
||||
int stage_size_bytes) {
|
||||
if (size_bytes % sizeof(typename packed_t<T>::P) != 0 ||
|
||||
stage_size_bytes % sizeof(typename packed_t<T>::P) != 0)
|
||||
throw std::runtime_error(
|
||||
"MNNVL Lamport allgather requires 16-byte aligned sizes");
|
||||
|
||||
auto ptrs = buffers_.at(local_buffer);
|
||||
int size_per_rank = size_bytes / sizeof(typename packed_t<T>::P);
|
||||
int stage_size = stage_size_bytes / sizeof(typename packed_t<T>::P);
|
||||
int blocks =
|
||||
(size_per_rank + kMnnvlLamportAgThreads - 1) / kMnnvlLamportAgThreads;
|
||||
|
||||
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000
|
||||
cudaLaunchAttribute attributes[1]{};
|
||||
attributes[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attributes[0].val.programmaticStreamSerializationAllowed = 1;
|
||||
cudaLaunchConfig_t config{.gridDim = dim3(blocks),
|
||||
.blockDim = dim3(kMnnvlLamportAgThreads),
|
||||
.dynamicSmemBytes = 0,
|
||||
.stream = stream,
|
||||
.attrs = attributes,
|
||||
.numAttrs = 1};
|
||||
#define MNNVL_LAMPORT_AG_LAUNCH(ngpus) \
|
||||
CUDACHECK(cudaLaunchKernelEx(&config, &mnnvl_lamport_all_gather<T, ngpus>, \
|
||||
ptrs, input, output, \
|
||||
reinterpret_cast<T*>(multicast_buffer), \
|
||||
epochs, rank_, size_per_rank, stage_size))
|
||||
#else
|
||||
#define MNNVL_LAMPORT_AG_LAUNCH(ngpus) \
|
||||
mnnvl_lamport_all_gather<T, ngpus> \
|
||||
<<<blocks, kMnnvlLamportAgThreads, 0, stream>>>( \
|
||||
ptrs, input, output, reinterpret_cast<T*>(multicast_buffer), \
|
||||
epochs, rank_, size_per_rank, stage_size)
|
||||
#endif
|
||||
|
||||
#define MNNVL_LAMPORT_AG_CASE(ngpus) \
|
||||
case ngpus: \
|
||||
MNNVL_LAMPORT_AG_LAUNCH(ngpus); \
|
||||
break;
|
||||
|
||||
switch (world_size_) {
|
||||
MNNVL_LAMPORT_AG_CASE(2)
|
||||
MNNVL_LAMPORT_AG_CASE(4)
|
||||
MNNVL_LAMPORT_AG_CASE(6)
|
||||
MNNVL_LAMPORT_AG_CASE(8)
|
||||
MNNVL_LAMPORT_AG_CASE(16)
|
||||
default:
|
||||
throw std::runtime_error(
|
||||
"MNNVL Lamport allgather only supports num gpus in (2,4,6,8,16)");
|
||||
}
|
||||
#undef MNNVL_LAMPORT_AG_CASE
|
||||
#undef MNNVL_LAMPORT_AG_LAUNCH
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void CustomAllreduce::reduce_scatter(cudaStream_t stream, T* input, T* output,
|
||||
int size, int threads, int block_limit) {
|
||||
auto packed_size = packed_t<T>::P::size;
|
||||
if (size % (packed_size * world_size_) != 0)
|
||||
throw std::runtime_error(
|
||||
"custom reduce-scatter requires each output shard byte size to be "
|
||||
"a multiple of 16");
|
||||
|
||||
auto ptrs = buffers_.at(input);
|
||||
int size_per_rank = size / packed_size / world_size_;
|
||||
int blocks = std::min(block_limit, (size_per_rank + threads - 1) / threads);
|
||||
|
||||
#define RS_CASE(ngpus) \
|
||||
case ngpus: \
|
||||
cross_device_reduce_scatter<T, ngpus><<<blocks, threads, 0, stream>>>( \
|
||||
ptrs, sg_, self_sg_, output, rank_, size_per_rank); \
|
||||
break;
|
||||
|
||||
switch (world_size_) {
|
||||
RS_CASE(2)
|
||||
RS_CASE(4)
|
||||
RS_CASE(6)
|
||||
RS_CASE(8)
|
||||
default:
|
||||
throw std::runtime_error(
|
||||
"custom reduce-scatter only supports num gpus in (2,4,6,8)");
|
||||
}
|
||||
#undef RS_CASE
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void CustomAllreduce::mnnvl_lamport_reduce_scatter(cudaStream_t stream,
|
||||
T* input, T* output,
|
||||
void* local_buffer,
|
||||
uint32_t* epochs, int size,
|
||||
int stage_size_bytes) {
|
||||
auto packed_size = packed_t<T>::P::size;
|
||||
if (size % (packed_size * world_size_) != 0 ||
|
||||
stage_size_bytes % sizeof(typename packed_t<T>::P) != 0)
|
||||
throw std::runtime_error(
|
||||
"MNNVL Lamport reduce-scatter requires 16-byte aligned sizes");
|
||||
|
||||
auto ptrs = buffers_.at(local_buffer);
|
||||
int size_per_rank = size / packed_size / world_size_;
|
||||
int stage_size = stage_size_bytes / sizeof(typename packed_t<T>::P);
|
||||
int blocks_per_rank =
|
||||
(size_per_rank + kMnnvlLamportRsThreads - 1) / kMnnvlLamportRsThreads;
|
||||
int blocks = blocks_per_rank * world_size_;
|
||||
|
||||
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000
|
||||
cudaLaunchAttribute attributes[1]{};
|
||||
attributes[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attributes[0].val.programmaticStreamSerializationAllowed = 1;
|
||||
cudaLaunchConfig_t config{.gridDim = dim3(blocks),
|
||||
.blockDim = dim3(kMnnvlLamportRsThreads),
|
||||
.dynamicSmemBytes = 0,
|
||||
.stream = stream,
|
||||
.attrs = attributes,
|
||||
.numAttrs = 1};
|
||||
#define MNNVL_LAMPORT_RS_LAUNCH(ngpus) \
|
||||
CUDACHECK(cudaLaunchKernelEx( \
|
||||
&config, &mnnvl_lamport_reduce_scatter_kernel<T, ngpus>, ptrs, input, \
|
||||
output, epochs, rank_, size_per_rank, stage_size))
|
||||
#else
|
||||
#define MNNVL_LAMPORT_RS_LAUNCH(ngpus) \
|
||||
mnnvl_lamport_reduce_scatter_kernel<T, ngpus> \
|
||||
<<<blocks, kMnnvlLamportRsThreads, 0, stream>>>( \
|
||||
ptrs, input, output, epochs, rank_, size_per_rank, stage_size)
|
||||
#endif
|
||||
|
||||
#define MNNVL_LAMPORT_RS_CASE(ngpus) \
|
||||
case ngpus: \
|
||||
MNNVL_LAMPORT_RS_LAUNCH(ngpus); \
|
||||
break;
|
||||
|
||||
switch (world_size_) {
|
||||
MNNVL_LAMPORT_RS_CASE(2)
|
||||
MNNVL_LAMPORT_RS_CASE(4)
|
||||
MNNVL_LAMPORT_RS_CASE(6)
|
||||
MNNVL_LAMPORT_RS_CASE(8)
|
||||
MNNVL_LAMPORT_RS_CASE(16)
|
||||
default:
|
||||
throw std::runtime_error(
|
||||
"MNNVL Lamport reduce-scatter only supports num gpus in "
|
||||
"(2,4,6,8,16)");
|
||||
}
|
||||
#undef MNNVL_LAMPORT_RS_CASE
|
||||
#undef MNNVL_LAMPORT_RS_LAUNCH
|
||||
}
|
||||
|
||||
} // namespace vllm
|
||||
|
||||
using fptr_t = int64_t;
|
||||
static_assert(sizeof(void*) == sizeof(fptr_t));
|
||||
|
||||
bool _is_weak_contiguous(torch::stable::Tensor& t);
|
||||
|
||||
void custom_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out, fptr_t _reg_buffer,
|
||||
int64_t reg_buffer_sz_bytes) {
|
||||
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
inp.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
|
||||
|
||||
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
|
||||
STD_TORCH_CHECK((inp.numel() * fa->world_size_) == (out.numel()));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(out));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(inp));
|
||||
auto input_size = inp.numel() * inp.element_size();
|
||||
auto reg_buffer = reinterpret_cast<void*>(_reg_buffer);
|
||||
STD_TORCH_CHECK(reg_buffer != nullptr);
|
||||
STD_TORCH_CHECK((input_size) <= (reg_buffer_sz_bytes));
|
||||
STD_CUDA_CHECK(cudaMemcpyAsync(reg_buffer, inp.const_data_ptr(), input_size,
|
||||
cudaMemcpyDeviceToDevice, stream));
|
||||
fa->allgather(stream, reg_buffer, out.mutable_data_ptr(), input_size);
|
||||
}
|
||||
|
||||
void mnnvl_lamport_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out, fptr_t _local_buffer,
|
||||
fptr_t _multicast_buffer, fptr_t _epoch_buffer,
|
||||
int64_t stage_sz_bytes) {
|
||||
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
inp.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
|
||||
|
||||
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
|
||||
STD_TORCH_CHECK((inp.numel() * fa->world_size_) == (out.numel()));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(out));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(inp));
|
||||
auto input_size = inp.numel() * inp.element_size();
|
||||
STD_TORCH_CHECK((input_size * fa->world_size_) <= stage_sz_bytes);
|
||||
auto local_buffer = reinterpret_cast<void*>(_local_buffer);
|
||||
auto multicast_buffer = reinterpret_cast<void*>(_multicast_buffer);
|
||||
auto epochs = reinterpret_cast<uint32_t*>(_epoch_buffer);
|
||||
switch (out.scalar_type()) {
|
||||
case torch::headeronly::ScalarType::Float: {
|
||||
fa->mnnvl_lamport_allgather<float>(
|
||||
stream, reinterpret_cast<float*>(inp.mutable_data_ptr()),
|
||||
reinterpret_cast<float*>(out.mutable_data_ptr()), local_buffer,
|
||||
multicast_buffer, epochs, input_size, stage_sz_bytes);
|
||||
break;
|
||||
}
|
||||
case torch::headeronly::ScalarType::Half: {
|
||||
fa->mnnvl_lamport_allgather<half>(
|
||||
stream, reinterpret_cast<half*>(inp.mutable_data_ptr()),
|
||||
reinterpret_cast<half*>(out.mutable_data_ptr()), local_buffer,
|
||||
multicast_buffer, epochs, input_size, stage_sz_bytes);
|
||||
break;
|
||||
}
|
||||
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
|
||||
case torch::headeronly::ScalarType::BFloat16: {
|
||||
fa->mnnvl_lamport_allgather<nv_bfloat16>(
|
||||
stream, reinterpret_cast<nv_bfloat16*>(inp.mutable_data_ptr()),
|
||||
reinterpret_cast<nv_bfloat16*>(out.mutable_data_ptr()), local_buffer,
|
||||
multicast_buffer, epochs, input_size, stage_sz_bytes);
|
||||
break;
|
||||
}
|
||||
#endif
|
||||
default:
|
||||
throw std::runtime_error(
|
||||
"MNNVL Lamport allgather only supports float32, float16 and "
|
||||
"bfloat16");
|
||||
}
|
||||
}
|
||||
|
||||
void custom_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out, fptr_t _reg_buffer,
|
||||
int64_t reg_buffer_sz_bytes) {
|
||||
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
inp.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
|
||||
|
||||
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
|
||||
STD_TORCH_CHECK((out.numel() * fa->world_size_) == (inp.numel()));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(out));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(inp));
|
||||
auto input_size = inp.numel() * inp.element_size();
|
||||
auto reg_buffer = reinterpret_cast<void*>(_reg_buffer);
|
||||
STD_TORCH_CHECK(reg_buffer != nullptr);
|
||||
STD_TORCH_CHECK((input_size) <= (reg_buffer_sz_bytes));
|
||||
STD_CUDA_CHECK(cudaMemcpyAsync(reg_buffer, inp.const_data_ptr(), input_size,
|
||||
cudaMemcpyDeviceToDevice, stream));
|
||||
switch (out.scalar_type()) {
|
||||
case torch::headeronly::ScalarType::Float: {
|
||||
fa->reduce_scatter<float>(
|
||||
stream, reinterpret_cast<float*>(reg_buffer),
|
||||
reinterpret_cast<float*>(out.mutable_data_ptr()), inp.numel());
|
||||
break;
|
||||
}
|
||||
case torch::headeronly::ScalarType::Half: {
|
||||
fa->reduce_scatter<half>(stream, reinterpret_cast<half*>(reg_buffer),
|
||||
reinterpret_cast<half*>(out.mutable_data_ptr()),
|
||||
inp.numel());
|
||||
break;
|
||||
}
|
||||
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
|
||||
case torch::headeronly::ScalarType::BFloat16: {
|
||||
fa->reduce_scatter<nv_bfloat16>(
|
||||
stream, reinterpret_cast<nv_bfloat16*>(reg_buffer),
|
||||
reinterpret_cast<nv_bfloat16*>(out.mutable_data_ptr()), inp.numel());
|
||||
break;
|
||||
}
|
||||
#endif
|
||||
default:
|
||||
throw std::runtime_error(
|
||||
"custom reduce-scatter only supports float32, float16 and bfloat16");
|
||||
}
|
||||
}
|
||||
|
||||
void mnnvl_lamport_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out,
|
||||
fptr_t _local_buffer, fptr_t _epoch_buffer,
|
||||
int64_t stage_sz_bytes) {
|
||||
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
|
||||
const torch::stable::accelerator::DeviceGuard device_guard(
|
||||
inp.get_device_index());
|
||||
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
|
||||
|
||||
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
|
||||
STD_TORCH_CHECK((out.numel() * fa->world_size_) == (inp.numel()));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(out));
|
||||
STD_TORCH_CHECK(_is_weak_contiguous(inp));
|
||||
auto input_size = inp.numel() * inp.element_size();
|
||||
STD_TORCH_CHECK(input_size <= stage_sz_bytes);
|
||||
auto local_buffer = reinterpret_cast<void*>(_local_buffer);
|
||||
auto epochs = reinterpret_cast<uint32_t*>(_epoch_buffer);
|
||||
switch (out.scalar_type()) {
|
||||
case torch::headeronly::ScalarType::Float: {
|
||||
fa->mnnvl_lamport_reduce_scatter<float>(
|
||||
stream, reinterpret_cast<float*>(inp.mutable_data_ptr()),
|
||||
reinterpret_cast<float*>(out.mutable_data_ptr()), local_buffer,
|
||||
epochs, inp.numel(), stage_sz_bytes);
|
||||
break;
|
||||
}
|
||||
case torch::headeronly::ScalarType::Half: {
|
||||
fa->mnnvl_lamport_reduce_scatter<half>(
|
||||
stream, reinterpret_cast<half*>(inp.mutable_data_ptr()),
|
||||
reinterpret_cast<half*>(out.mutable_data_ptr()), local_buffer, epochs,
|
||||
inp.numel(), stage_sz_bytes);
|
||||
break;
|
||||
}
|
||||
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
|
||||
case torch::headeronly::ScalarType::BFloat16: {
|
||||
fa->mnnvl_lamport_reduce_scatter<nv_bfloat16>(
|
||||
stream, reinterpret_cast<nv_bfloat16*>(inp.mutable_data_ptr()),
|
||||
reinterpret_cast<nv_bfloat16*>(out.mutable_data_ptr()), local_buffer,
|
||||
epochs, inp.numel(), stage_sz_bytes);
|
||||
break;
|
||||
}
|
||||
#endif
|
||||
default:
|
||||
throw std::runtime_error(
|
||||
"MNNVL Lamport reduce-scatter only supports float32, float16 and "
|
||||
"bfloat16");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
#include "ops.h"
|
||||
#include "core/registration.h"
|
||||
|
||||
#include <torch/csrc/stable/library.h>
|
||||
|
||||
STABLE_TORCH_LIBRARY_FRAGMENT(_C_custom_ar, custom_ag_rs) {
|
||||
custom_ag_rs.def(
|
||||
"custom_all_gather(int fa, Tensor inp, Tensor! out, int reg_buffer, "
|
||||
"int reg_buffer_sz_bytes) -> ()");
|
||||
custom_ag_rs.def(
|
||||
"mnnvl_lamport_all_gather(int fa, Tensor inp, Tensor! out, int "
|
||||
"local_buffer, int multicast_buffer, int epoch_buffer, int "
|
||||
"stage_sz_bytes) -> ()");
|
||||
custom_ag_rs.def(
|
||||
"custom_reduce_scatter(int fa, Tensor inp, Tensor! out, int reg_buffer, "
|
||||
"int reg_buffer_sz_bytes) -> ()");
|
||||
custom_ag_rs.def(
|
||||
"mnnvl_lamport_reduce_scatter(int fa, Tensor inp, Tensor! out, int "
|
||||
"local_buffer, int epoch_buffer, int stage_sz_bytes) -> ()");
|
||||
}
|
||||
|
||||
STABLE_TORCH_LIBRARY_IMPL(_C_custom_ar, CUDA, custom_ag_rs) {
|
||||
custom_ag_rs.impl("custom_all_gather", TORCH_BOX(&custom_all_gather));
|
||||
custom_ag_rs.impl("mnnvl_lamport_all_gather",
|
||||
TORCH_BOX(&mnnvl_lamport_all_gather));
|
||||
custom_ag_rs.impl("custom_reduce_scatter", TORCH_BOX(&custom_reduce_scatter));
|
||||
custom_ag_rs.impl("mnnvl_lamport_reduce_scatter",
|
||||
TORCH_BOX(&mnnvl_lamport_reduce_scatter));
|
||||
}
|
||||
@@ -18,14 +18,14 @@ fptr_t init_custom_ar(const std::vector<fptr_t>& fake_ipc_ptrs,
|
||||
torch::stable::Tensor& rank_data, int64_t rank,
|
||||
bool fully_connected) {
|
||||
int world_size = fake_ipc_ptrs.size();
|
||||
if (world_size > 8)
|
||||
throw std::invalid_argument("world size > 8 is not supported");
|
||||
if (world_size > vllm::kMaxCustomCollectiveRanks)
|
||||
throw std::invalid_argument("world size > 16 is not supported");
|
||||
if (world_size % 2 != 0)
|
||||
throw std::invalid_argument("Odd num gpus is not supported for now");
|
||||
if (rank < 0 || rank >= world_size)
|
||||
throw std::invalid_argument("invalid rank passed in");
|
||||
|
||||
vllm::Signal* ipc_ptrs[8];
|
||||
vllm::Signal* ipc_ptrs[vllm::kMaxCustomCollectiveRanks];
|
||||
for (int i = 0; i < world_size; i++) {
|
||||
ipc_ptrs[i] = reinterpret_cast<vllm::Signal*>(fake_ipc_ptrs[i]);
|
||||
}
|
||||
@@ -124,7 +124,7 @@ int64_t meta_size() { return sizeof(vllm::Signal); }
|
||||
void register_buffer(fptr_t _fa, const std::vector<fptr_t>& fake_ipc_ptrs) {
|
||||
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
|
||||
STD_TORCH_CHECK(fake_ipc_ptrs.size() == fa->world_size_);
|
||||
void* ipc_ptrs[8];
|
||||
void* ipc_ptrs[vllm::kMaxCustomCollectiveRanks];
|
||||
for (int i = 0; i < fake_ipc_ptrs.size(); i++) {
|
||||
ipc_ptrs[i] = reinterpret_cast<void*>(fake_ipc_ptrs[i]);
|
||||
}
|
||||
|
||||
@@ -21,6 +21,21 @@
|
||||
#define VLLM_STABLE_DISPATCH_FP8_CASE(enum_type, ...) \
|
||||
THO_PRIVATE_CASE_TYPE_USING_HINT(enum_type, fp8_t, __VA_ARGS__)
|
||||
|
||||
// Same idea, for dispatching on an int32/int64 index tensor (e.g. topk_ids)
|
||||
// nested inside a value-type dispatch. Named 'idx_t' instead of 'scalar_t'.
|
||||
#define VLLM_STABLE_DISPATCH_IDX_CASE(enum_type, ...) \
|
||||
THO_PRIVATE_CASE_TYPE_USING_HINT(enum_type, idx_t, __VA_ARGS__)
|
||||
|
||||
#define VLLM_STABLE_DISPATCH_CASE_IDX_TYPES(...) \
|
||||
VLLM_STABLE_DISPATCH_IDX_CASE(torch::headeronly::ScalarType::Int, \
|
||||
__VA_ARGS__) \
|
||||
VLLM_STABLE_DISPATCH_IDX_CASE(torch::headeronly::ScalarType::Long, \
|
||||
__VA_ARGS__)
|
||||
|
||||
#define VLLM_STABLE_DISPATCH_IDX_TYPES(TYPE, NAME, ...) \
|
||||
THO_DISPATCH_SWITCH(TYPE, NAME, \
|
||||
VLLM_STABLE_DISPATCH_CASE_IDX_TYPES(__VA_ARGS__))
|
||||
|
||||
#define VLLM_STABLE_DISPATCH_CASE_FLOATING_TYPES(...) \
|
||||
THO_DISPATCH_CASE(torch::headeronly::ScalarType::Float, __VA_ARGS__) \
|
||||
THO_DISPATCH_CASE(torch::headeronly::ScalarType::Half, __VA_ARGS__) \
|
||||
|
||||
@@ -647,17 +647,17 @@ __global__ __launch_bounds__(256, 1) void fused_a_gemm_kernel(
|
||||
#endif
|
||||
}
|
||||
|
||||
template <typename T, int kHdIn, int kHdOut, int kTileN>
|
||||
template <typename T, int kHdIn, int kHdOut, int kTileN, int kTileK = 256>
|
||||
void invokeFusedAGemm(T* output, T const* mat_a, T const* mat_b, int num_tokens,
|
||||
cudaStream_t const stream) {
|
||||
constexpr int gemm_m = kHdOut; // 2112
|
||||
int const gemm_n = num_tokens; // 1-16
|
||||
constexpr int gemm_k = kHdIn; // 7168
|
||||
cudaStream_t const stream, bool enable_pdl) {
|
||||
constexpr int gemm_m = kHdOut;
|
||||
int const gemm_n = num_tokens;
|
||||
constexpr int gemm_k = kHdIn;
|
||||
constexpr int batch_size = 1;
|
||||
std::swap(mat_a, mat_b);
|
||||
constexpr int tile_m = 16;
|
||||
constexpr int tile_n = kTileN; // 8 or 16
|
||||
constexpr int tile_k = std::max(256, 1024 / tile_n); // 256
|
||||
constexpr int tile_n = kTileN;
|
||||
constexpr int tile_k = kTileK;
|
||||
constexpr int max_stage_cnt =
|
||||
1024 * 192 / ((tile_m + tile_n) * tile_k * sizeof(bf16_t));
|
||||
constexpr int k_iter_cnt = gemm_k / tile_k;
|
||||
@@ -679,7 +679,8 @@ void invokeFusedAGemm(T* output, T const* mat_a, T const* mat_b, int num_tokens,
|
||||
config.stream = stream;
|
||||
cudaLaunchAttribute attrs[1];
|
||||
attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attrs[0].val.programmaticStreamSerializationAllowed = getEnvEnablePDL();
|
||||
attrs[0].val.programmaticStreamSerializationAllowed =
|
||||
enable_pdl || getEnvEnablePDL();
|
||||
config.numAttrs = 1;
|
||||
config.attrs = attrs;
|
||||
if (smem_bytes >= (48 * 1024)) {
|
||||
@@ -694,36 +695,48 @@ void invokeFusedAGemm(T* output, T const* mat_a, T const* mat_b, int num_tokens,
|
||||
output, mat_a, mat_b, gemm_n);
|
||||
}
|
||||
|
||||
template void invokeFusedAGemm<__nv_bfloat16, 7168, 2112, 8>(
|
||||
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, int num_tokens,
|
||||
cudaStream_t);
|
||||
|
||||
template void invokeFusedAGemm<__nv_bfloat16, 7168, 2112, 16>(
|
||||
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, int num_tokens,
|
||||
cudaStream_t);
|
||||
template <typename T, int kHdIn, int kHdOut, int kTileK = 256>
|
||||
void invokeFusedAGemmForTokens(T* output, T const* mat_a, T const* mat_b,
|
||||
int num_tokens, cudaStream_t const stream,
|
||||
bool enable_pdl) {
|
||||
if (num_tokens <= 8) {
|
||||
invokeFusedAGemm<T, kHdIn, kHdOut, 8, kTileK>(
|
||||
output, mat_a, mat_b, num_tokens, stream, enable_pdl);
|
||||
} else {
|
||||
invokeFusedAGemm<T, kHdIn, kHdOut, 16, kTileK>(
|
||||
output, mat_a, mat_b, num_tokens, stream, enable_pdl);
|
||||
}
|
||||
}
|
||||
|
||||
void dsv3_fused_a_gemm(torch::stable::Tensor& output,
|
||||
torch::stable::Tensor const& mat_a,
|
||||
torch::stable::Tensor const& mat_b) {
|
||||
torch::stable::Tensor const& mat_b, bool enable_pdl) {
|
||||
STD_TORCH_CHECK(mat_a.dim() == 2 && mat_b.dim() == 2 && output.dim() == 2);
|
||||
int const num_tokens = mat_a.size(0);
|
||||
int const hd_in = mat_a.size(1);
|
||||
int const hd_out = mat_b.size(1);
|
||||
|
||||
constexpr int kHdIn = 7168;
|
||||
constexpr int kHdOut = 2112;
|
||||
STD_TORCH_CHECK(num_tokens >= 1 && num_tokens <= 16,
|
||||
"required 1 <= mat_a.shape[0] <= 16");
|
||||
STD_TORCH_CHECK(hd_in == kHdIn, "required mat_a.shape[1] == 7168");
|
||||
STD_TORCH_CHECK(hd_out == kHdOut, "required mat_b.shape[1] == 2112");
|
||||
STD_TORCH_CHECK(output.size(0) == num_tokens,
|
||||
"required output.shape[0] == mat_a.shape[0]");
|
||||
STD_TORCH_CHECK(output.size(1) == hd_out,
|
||||
"required output.shape[1] == mat_b.shape[1]");
|
||||
STD_TORCH_CHECK(mat_b.size(0) == hd_in,
|
||||
"required mat_b.shape[0] == mat_a.shape[1]");
|
||||
|
||||
STD_TORCH_CHECK(mat_a.stride(1) == 1, "mat_a must be a row major tensor");
|
||||
STD_TORCH_CHECK(output.stride(1) == 1, "output must be a row major tensor");
|
||||
STD_TORCH_CHECK(mat_b.stride(0) == 1, "mat_b must be a column major tensor");
|
||||
STD_TORCH_CHECK(mat_a.get_device_index() == mat_b.get_device_index() &&
|
||||
mat_a.get_device_index() == output.get_device_index(),
|
||||
"mat_a, mat_b, and output must be on the same device");
|
||||
|
||||
// The kernels index global memory with raw pointers and packed strides, so
|
||||
// reject any padded or transposed view rather than reading out of bounds.
|
||||
STD_TORCH_CHECK(mat_a.stride(0) == hd_in && mat_a.stride(1) == 1,
|
||||
"mat_a must be a packed row-major [num_tokens, hd_in] tensor");
|
||||
STD_TORCH_CHECK(output.stride(0) == hd_out && output.stride(1) == 1,
|
||||
"output must be a packed row-major [num_tokens, hd_out] tensor");
|
||||
STD_TORCH_CHECK(mat_b.stride(0) == 1 && mat_b.stride(1) == hd_in,
|
||||
"mat_b must be a packed column-major [hd_in, hd_out] tensor");
|
||||
|
||||
STD_TORCH_CHECK(
|
||||
mat_a.scalar_type() == torch::headeronly::ScalarType::BFloat16 &&
|
||||
@@ -738,19 +751,86 @@ void dsv3_fused_a_gemm(torch::stable::Tensor& output,
|
||||
STD_TORCH_CHECK(getSMVersion() >= 90, "required CUDA ARCH >= SM_90");
|
||||
|
||||
auto stream = get_current_cuda_stream(mat_a.get_device_index());
|
||||
if (num_tokens <= 8) {
|
||||
invokeFusedAGemm<__nv_bfloat16, kHdIn, kHdOut, 8>(
|
||||
reinterpret_cast<__nv_bfloat16*>(output.mutable_data_ptr()),
|
||||
reinterpret_cast<__nv_bfloat16 const*>(mat_a.data_ptr()),
|
||||
reinterpret_cast<__nv_bfloat16 const*>(mat_b.data_ptr()), num_tokens,
|
||||
stream);
|
||||
} else {
|
||||
invokeFusedAGemm<__nv_bfloat16, kHdIn, kHdOut, 16>(
|
||||
reinterpret_cast<__nv_bfloat16*>(output.mutable_data_ptr()),
|
||||
reinterpret_cast<__nv_bfloat16 const*>(mat_a.data_ptr()),
|
||||
reinterpret_cast<__nv_bfloat16 const*>(mat_b.data_ptr()), num_tokens,
|
||||
stream);
|
||||
auto* output_ptr =
|
||||
reinterpret_cast<__nv_bfloat16*>(output.mutable_data_ptr());
|
||||
auto const* mat_a_ptr =
|
||||
reinterpret_cast<__nv_bfloat16 const*>(mat_a.data_ptr());
|
||||
auto const* mat_b_ptr =
|
||||
reinterpret_cast<__nv_bfloat16 const*>(mat_b.data_ptr());
|
||||
|
||||
#define DISPATCH_DSV3_SHAPE(HD_IN, HD_OUT) \
|
||||
if (hd_in == HD_IN && hd_out == HD_OUT) { \
|
||||
invokeFusedAGemmForTokens<__nv_bfloat16, HD_IN, HD_OUT>( \
|
||||
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, \
|
||||
enable_pdl); \
|
||||
return; \
|
||||
}
|
||||
|
||||
// Shapes the Kimi-K3 selector routes to dsv3_fused_a (see the dsv3 winners
|
||||
// in KIMI_K3_PROJECTIONS) plus the DeepSeek V2/V3 QKV A-projection.
|
||||
DISPATCH_DSV3_SHAPE(7168, 1536)
|
||||
DISPATCH_DSV3_SHAPE(7168, 2112)
|
||||
DISPATCH_DSV3_SHAPE(1536, 2304)
|
||||
DISPATCH_DSV3_SHAPE(1536, 4608)
|
||||
DISPATCH_DSV3_SHAPE(7168, 3584)
|
||||
DISPATCH_DSV3_SHAPE(768, 7168)
|
||||
// TP16 dsv3 winners, as (hd_in=K, hd_out=N). TP16 dense down_proj is absent
|
||||
// because hd_in=2112 is not a multiple of any supported tile_k.
|
||||
DISPATCH_DSV3_SHAPE(1536, 1152)
|
||||
DISPATCH_DSV3_SHAPE(7168, 768)
|
||||
DISPATCH_DSV3_SHAPE(7168, 3216)
|
||||
DISPATCH_DSV3_SHAPE(7168, 4224)
|
||||
|
||||
#ifdef VLLM_K3_BENCH_SHAPES
|
||||
// The selector routes these shapes to CuTe or the default GEMM, so they are
|
||||
// never reached in production. They are compiled only for offline
|
||||
// DSV3-vs-CuTe benchmarking.
|
||||
DISPATCH_DSV3_SHAPE(7168, 6288)
|
||||
DISPATCH_DSV3_SHAPE(1536, 7168)
|
||||
DISPATCH_DSV3_SHAPE(3584, 7168)
|
||||
DISPATCH_DSV3_SHAPE(7168, 8448)
|
||||
DISPATCH_DSV3_SHAPE(7168, 20480)
|
||||
DISPATCH_DSV3_SHAPE(7168, 3072)
|
||||
DISPATCH_DSV3_SHAPE(7168, 12448)
|
||||
DISPATCH_DSV3_SHAPE(3072, 7168)
|
||||
DISPATCH_DSV3_SHAPE(8448, 7168)
|
||||
DISPATCH_DSV3_SHAPE(7168, 16896)
|
||||
DISPATCH_DSV3_SHAPE(7168, 40960)
|
||||
#endif
|
||||
|
||||
#undef DISPATCH_DSV3_SHAPE
|
||||
|
||||
if (hd_in == 128 && hd_out == 1536) {
|
||||
invokeFusedAGemmForTokens<__nv_bfloat16, 128, 1536, 128>(
|
||||
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
|
||||
return;
|
||||
}
|
||||
if (hd_in == 128 && hd_out == 3072) {
|
||||
invokeFusedAGemmForTokens<__nv_bfloat16, 128, 3072, 128>(
|
||||
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
|
||||
return;
|
||||
}
|
||||
// TP16 KDA f_b_proj and shared_expert down_proj. Neither hd_in is a multiple
|
||||
// of 256, so both need the 128 tile_k.
|
||||
if (hd_in == 128 && hd_out == 768) {
|
||||
invokeFusedAGemmForTokens<__nv_bfloat16, 128, 768, 128>(
|
||||
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
|
||||
return;
|
||||
}
|
||||
if (hd_in == 384 && hd_out == 7168) {
|
||||
invokeFusedAGemmForTokens<__nv_bfloat16, 384, 7168, 128>(
|
||||
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
|
||||
return;
|
||||
}
|
||||
#ifdef VLLM_K3_BENCH_SHAPES
|
||||
if (hd_in == 4224 && hd_out == 7168) {
|
||||
invokeFusedAGemmForTokens<__nv_bfloat16, 4224, 7168, 128>(
|
||||
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
|
||||
STD_TORCH_CHECK(false, "unsupported DSV3 fused-A GEMM shape");
|
||||
}
|
||||
|
||||
STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, m) {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,954 @@
|
||||
/*
|
||||
* Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
|
||||
*/
|
||||
|
||||
// Production AttnRes forward for Blackwell (SM100).
|
||||
//
|
||||
// Warp-specialized online softmax + residual + RMSNorm:
|
||||
// - 1 producer warp issues cp.async.bulk row loads into shared memory.
|
||||
// - 8 consumer warps compute reductions and output.
|
||||
// - Q=res_weight*rms_weight remains in registers across persistent tokens.
|
||||
// - V rows are converted once and cached as FP32 in TMEM between passes.
|
||||
//
|
||||
// Integration contract: Kimi K3 H=7168, 1<=num_blocks<=8, and token-major
|
||||
// block residual storage.
|
||||
|
||||
#include "../torch_utils.h"
|
||||
|
||||
#include <cfloat>
|
||||
#include <cstdint>
|
||||
#include <cstdio>
|
||||
#include <cuda_runtime.h>
|
||||
#include <type_traits>
|
||||
|
||||
using bf16_t = __nv_bfloat16;
|
||||
|
||||
namespace sm100 {
|
||||
namespace fwd_prod_v2 {
|
||||
|
||||
constexpr int K_TILE = 1024;
|
||||
constexpr int N_CHUNK_DEFAULT = 4;
|
||||
constexpr int CHUNK_DEPTH = 2;
|
||||
constexpr int BLK = 288; // 1 producer warp + 8 consumer warps
|
||||
constexpr int CONSUMER_THREADS = BLK - 32; // 256
|
||||
constexpr int CONSUMER_WARPS = CONSUMER_THREADS / 32;
|
||||
constexpr int CONSUMER_GROUPS = 2; // two 128-thread consumer groups
|
||||
constexpr int CONSUMER_THREADS_PER_GROUP = CONSUMER_THREADS / CONSUMER_GROUPS;
|
||||
constexpr int FIRST_USER_NAMED_BARRIER = 8;
|
||||
|
||||
__device__ __forceinline__ const bf16_t* residual_addr(
|
||||
const bf16_t* block_res, const bf16_t* layer_res, int source, int N,
|
||||
int token, int block_stride_m, int block_stride_r, int H) {
|
||||
if (source < N - 1) {
|
||||
return block_res + static_cast<long long>(token) * block_stride_m +
|
||||
source * block_stride_r;
|
||||
}
|
||||
return layer_res + static_cast<long long>(token) * H;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ uint32_t elect_one_sync() {
|
||||
uint32_t pred = 0;
|
||||
uint32_t laneid = 0;
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .b32 %%rx;\n"
|
||||
".reg .pred %%px;\n"
|
||||
" elect.sync %%rx|%%px, %2;\n"
|
||||
"@%%px mov.s32 %1, 1;\n"
|
||||
" mov.s32 %0, %%rx;\n"
|
||||
"}\n"
|
||||
: "+r"(laneid), "+r"(pred)
|
||||
: "r"(0xffffffff));
|
||||
return pred;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void mbarrier_init(uint64_t& barrier,
|
||||
int thread_count) {
|
||||
uint32_t const barrier_addr =
|
||||
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
|
||||
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n" ::"r"(barrier_addr),
|
||||
"r"(thread_count));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void mbarrier_expect_tx(uint64_t& barrier,
|
||||
uint32_t bytes) {
|
||||
uint32_t const barrier_addr =
|
||||
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
|
||||
asm volatile("mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;\n" ::"r"(
|
||||
barrier_addr),
|
||||
"r"(bytes));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void mbarrier_wait(uint64_t& barrier, int phase) {
|
||||
uint32_t const barrier_addr =
|
||||
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .pred p;\n"
|
||||
"WAIT:\n"
|
||||
"mbarrier.try_wait.parity.shared::cta.b64 p, [%0], %1;\n"
|
||||
"@p bra DONE;\n"
|
||||
"bra WAIT;\n"
|
||||
"DONE:\n"
|
||||
"}\n" ::"r"(barrier_addr),
|
||||
"r"(phase));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void mbarrier_arrive(uint64_t& barrier) {
|
||||
uint32_t const barrier_addr =
|
||||
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .b64 state;\n"
|
||||
"mbarrier.arrive.shared::cta.b64 state, [%0];\n"
|
||||
"}\n" ::"r"(barrier_addr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void fence_mbarrier_init() {
|
||||
asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void named_barrier_sync(uint32_t num_threads,
|
||||
uint32_t user_barrier_id) {
|
||||
asm volatile(
|
||||
"bar.sync %0, %1;" ::"r"(user_barrier_id + FIRST_USER_NAMED_BARRIER),
|
||||
"r"(num_threads)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tmem_allocate(int num_columns, uint32_t* dst) {
|
||||
uint32_t const dst_addr =
|
||||
static_cast<uint32_t>(__cvta_generic_to_shared(dst));
|
||||
asm volatile(
|
||||
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" ::"r"(
|
||||
dst_addr),
|
||||
"r"(num_columns));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tmem_free(uint32_t tmem_ptr, int num_columns) {
|
||||
asm volatile(
|
||||
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" ::"r"(tmem_ptr),
|
||||
"r"(num_columns));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tmem_release_allocation_lock() {
|
||||
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tmem_store_wait() {
|
||||
asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory");
|
||||
}
|
||||
|
||||
template <int N, typename T>
|
||||
__device__ __forceinline__ void tmem_load(uint32_t src_addr, T* dst) {
|
||||
uint32_t* values = reinterpret_cast<uint32_t*>(dst);
|
||||
if constexpr (N == 8) {
|
||||
asm volatile(
|
||||
"tcgen05.ld.sync.aligned.32x32b.x8.b32"
|
||||
"{%0, %1, %2, %3, %4, %5, %6, %7}, [%8];\n"
|
||||
: "=r"(values[0]), "=r"(values[1]), "=r"(values[2]), "=r"(values[3]),
|
||||
"=r"(values[4]), "=r"(values[5]), "=r"(values[6]), "=r"(values[7])
|
||||
: "r"(src_addr));
|
||||
} else {
|
||||
static_assert(N == 4, "AttnRes TMEM helpers support x4 and x8");
|
||||
asm volatile(
|
||||
"tcgen05.ld.sync.aligned.32x32b.x4.b32"
|
||||
"{%0, %1, %2, %3}, [%4];\n"
|
||||
: "=r"(values[0]), "=r"(values[1]), "=r"(values[2]), "=r"(values[3])
|
||||
: "r"(src_addr));
|
||||
}
|
||||
}
|
||||
|
||||
template <int N, typename T>
|
||||
__device__ __forceinline__ void tmem_store(uint32_t dst_addr, T* src) {
|
||||
uint32_t* values = reinterpret_cast<uint32_t*>(src);
|
||||
if constexpr (N == 8) {
|
||||
asm volatile(
|
||||
"tcgen05.st.sync.aligned.32x32b.x8.b32"
|
||||
"[%8], {%0, %1, %2, %3, %4, %5, %6, %7};\n" ::"r"(values[0]),
|
||||
"r"(values[1]), "r"(values[2]), "r"(values[3]), "r"(values[4]),
|
||||
"r"(values[5]), "r"(values[6]), "r"(values[7]), "r"(dst_addr));
|
||||
} else {
|
||||
static_assert(N == 4, "AttnRes TMEM helpers support x4 and x8");
|
||||
asm volatile(
|
||||
"tcgen05.st.sync.aligned.32x32b.x4.b32"
|
||||
"[%4], {%0, %1, %2, %3};\n" ::"r"(values[0]),
|
||||
"r"(values[1]), "r"(values[2]), "r"(values[3]), "r"(dst_addr));
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 float2_add(const float2& a, const float2& b) {
|
||||
float2 result;
|
||||
asm volatile("add.rn.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t&>(result))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
return result;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 float2_mul(const float2& a, const float2& b) {
|
||||
float2 result;
|
||||
asm volatile("mul.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t&>(result))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
return result;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 float2_fma(const float2& a, const float2& b,
|
||||
const float2& c) {
|
||||
float2 result;
|
||||
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t&>(result))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(c)));
|
||||
return result;
|
||||
}
|
||||
|
||||
template <int NC>
|
||||
struct FwdSmemPlan {
|
||||
alignas(16) uint64_t bar_ready[CHUNK_DEPTH];
|
||||
alignas(16) uint64_t bar_consumed[CHUNK_DEPTH];
|
||||
alignas(16) uint64_t bar_output_norm_ready;
|
||||
alignas(16) float2 ws_stats[CONSUMER_WARPS][NC];
|
||||
uint32_t tmem_base;
|
||||
};
|
||||
|
||||
__device__ __forceinline__ void cp_async_bulk(void* smem_dst,
|
||||
const void* gmem_src, int bytes,
|
||||
uint64_t& mbar) {
|
||||
uint32_t const s = static_cast<uint32_t>(__cvta_generic_to_shared(smem_dst));
|
||||
uint32_t const m = static_cast<uint32_t>(__cvta_generic_to_shared(&mbar));
|
||||
asm volatile(
|
||||
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%0], "
|
||||
"[%1], %2, [%3];\n" ::"r"(s),
|
||||
"l"(gmem_src), "r"(bytes), "r"(m)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
template <int H, int NC = N_CHUNK_DEFAULT, bool RELEASE_TMEM = false,
|
||||
bool HAS_DELTA = false, bool HAS_OUTPUT_NORM = false,
|
||||
bool OUTPUT_NORM_IN_SMEM = false>
|
||||
__global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
|
||||
const bf16_t* __restrict__ block_res, bf16_t* __restrict__ layer_res,
|
||||
const bf16_t* __restrict__ delta, const bf16_t* __restrict__ res_w,
|
||||
const bf16_t* __restrict__ rms_w, bf16_t* __restrict__ output, int N, int T,
|
||||
int B, int block_stride_m, int block_stride_r, float rms_eps,
|
||||
const bf16_t* __restrict__ output_norm_weight, float output_norm_eps) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000 && __CUDA_ARCH__ < 1100
|
||||
constexpr float LOG2_E = 1.4426950408889634f;
|
||||
constexpr int N_CHUNK = NC;
|
||||
// The two-source specialization only consumes half of the TMEM columns.
|
||||
constexpr int TMEM_COLS_ALLOC = NC == 2 ? 128 : 256;
|
||||
constexpr int NUM_BUFS = CHUNK_DEPTH * NC;
|
||||
constexpr int NHT = H / K_TILE;
|
||||
constexpr int SLICES_PER_GROUP =
|
||||
(NHT + CONSUMER_GROUPS - 1) / CONSUMER_GROUPS;
|
||||
constexpr int VEC = 8;
|
||||
constexpr int ACC_PER_THREAD = H == 7168 ? 28 : SLICES_PER_GROUP * VEC;
|
||||
constexpr int TMEM_V_COLS_PER_GROUP = SLICES_PER_GROUP * N_CHUNK * VEC;
|
||||
constexpr int TMEM_V_COLS_TOTAL = CONSUMER_GROUPS * TMEM_V_COLS_PER_GROUP;
|
||||
static_assert(TMEM_V_COLS_TOTAL <= TMEM_COLS_ALLOC);
|
||||
static_assert(H >= 4096 && H <= 8192);
|
||||
static_assert(H % K_TILE == 0);
|
||||
|
||||
const int tid = threadIdx.x;
|
||||
const int wid = tid >> 5;
|
||||
const int lane = tid & 31;
|
||||
const int TB = T * B;
|
||||
const int num_ctas = gridDim.x;
|
||||
const int num_chunks = (N + N_CHUNK - 1) / N_CHUNK;
|
||||
|
||||
const int comp_wid = wid - 1;
|
||||
const int comp_tid = tid - 32;
|
||||
const int group = (comp_wid >= 4) ? 1 : 0;
|
||||
const int ct_in_group =
|
||||
(comp_tid >= 0) ? (comp_tid & (CONSUMER_THREADS_PER_GROUP - 1)) : -1;
|
||||
const int k_local = ct_in_group * VEC;
|
||||
|
||||
constexpr size_t V_BYTES = (size_t)NUM_BUFS * H * sizeof(bf16_t);
|
||||
constexpr size_t DELTA_BYTES =
|
||||
HAS_DELTA ? (size_t)CHUNK_DEPTH * H * sizeof(bf16_t) : 0;
|
||||
constexpr size_t OUTPUT_NORM_BYTES =
|
||||
OUTPUT_NORM_IN_SMEM ? (size_t)H * sizeof(bf16_t) : 0;
|
||||
extern __shared__ __align__(16) char smem_raw[];
|
||||
bf16_t* v_bufs = reinterpret_cast<bf16_t*>(smem_raw); // [NUM_BUFS][H]
|
||||
bf16_t* delta_bufs = reinterpret_cast<bf16_t*>(smem_raw + V_BYTES);
|
||||
bf16_t* output_norm_buf =
|
||||
reinterpret_cast<bf16_t*>(smem_raw + V_BYTES + DELTA_BYTES);
|
||||
FwdSmemPlan<NC>& plan = *reinterpret_cast<FwdSmemPlan<NC>*>(
|
||||
smem_raw + V_BYTES + DELTA_BYTES + OUTPUT_NORM_BYTES);
|
||||
|
||||
auto slot_of = [](long long gci, int n) {
|
||||
return (int)(gci % CHUNK_DEPTH) * N_CHUNK + n;
|
||||
};
|
||||
auto phase_of = [](long long gci) { return (int)((gci / CHUNK_DEPTH) & 1); };
|
||||
auto buf_ptr = [&](int slot) -> bf16_t* { return v_bufs + slot * H; };
|
||||
auto delta_buf_ptr = [&](int chunk_slot) -> bf16_t* {
|
||||
return delta_bufs + chunk_slot * H;
|
||||
};
|
||||
|
||||
if (wid == 0 && elect_one_sync()) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CHUNK_DEPTH; i++) {
|
||||
mbarrier_init(plan.bar_ready[i], 1);
|
||||
mbarrier_init(plan.bar_consumed[i], CONSUMER_WARPS);
|
||||
}
|
||||
if constexpr (OUTPUT_NORM_IN_SMEM) {
|
||||
mbarrier_init(plan.bar_output_norm_ready, 1);
|
||||
}
|
||||
fence_mbarrier_init();
|
||||
}
|
||||
|
||||
// gdc wait BEFORE tmem alloc
|
||||
cudaGridDependencySynchronize();
|
||||
|
||||
if (wid == 1) {
|
||||
tmem_allocate(TMEM_COLS_ALLOC, &plan.tmem_base);
|
||||
if constexpr (RELEASE_TMEM) {
|
||||
tmem_release_allocation_lock();
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if constexpr (OUTPUT_NORM_IN_SMEM) {
|
||||
if (wid == 0 && elect_one_sync()) {
|
||||
mbarrier_expect_tx(plan.bar_output_norm_ready, H * (int)sizeof(bf16_t));
|
||||
cp_async_bulk(output_norm_buf, output_norm_weight, H * sizeof(bf16_t),
|
||||
plan.bar_output_norm_ready);
|
||||
}
|
||||
}
|
||||
|
||||
const uint32_t my_v_tmem =
|
||||
comp_tid >= 0 ? plan.tmem_base + group * TMEM_V_COLS_PER_GROUP : 0;
|
||||
float q_cache[ACC_PER_THREAD];
|
||||
if (comp_tid >= 0) {
|
||||
#pragma unroll
|
||||
for (int si = 0; si < SLICES_PER_GROUP; si++) {
|
||||
if constexpr (H == 7168) {
|
||||
if (si == SLICES_PER_GROUP - 1) {
|
||||
int h_base = 6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 4; j++) {
|
||||
int h = h_base + j;
|
||||
q_cache[si * VEC + j] =
|
||||
__bfloat162float(rms_w[h]) * __bfloat162float(res_w[h]);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
int dt = si * CONSUMER_GROUPS + group;
|
||||
if (dt >= NHT) continue;
|
||||
int h_base = dt * K_TILE + k_local;
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC; j++) {
|
||||
int h = h_base + j;
|
||||
q_cache[si * VEC + j] =
|
||||
__bfloat162float(rms_w[h]) * __bfloat162float(res_w[h]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (wid == 0) {
|
||||
if (elect_one_sync()) {
|
||||
long long gci = 0;
|
||||
for (int tb = blockIdx.x; tb < TB; tb += num_ctas) {
|
||||
const int t = tb / B;
|
||||
for (int ci = 0; ci < num_chunks; ci++, gci++) {
|
||||
int ns = ci * N_CHUNK;
|
||||
int an = min(N_CHUNK, N - ns);
|
||||
int chunk_slot = (int)(gci % CHUNK_DEPTH);
|
||||
int pc = phase_of(gci);
|
||||
mbarrier_wait(plan.bar_consumed[chunk_slot], pc ^ 1);
|
||||
int transaction_bytes = an * H * (int)sizeof(bf16_t);
|
||||
if constexpr (HAS_DELTA) {
|
||||
int prefix_n = N - 1 - ns;
|
||||
if (prefix_n >= 0 && prefix_n < an) {
|
||||
transaction_bytes += H * (int)sizeof(bf16_t);
|
||||
}
|
||||
}
|
||||
mbarrier_expect_tx(plan.bar_ready[chunk_slot], transaction_bytes);
|
||||
#pragma unroll
|
||||
for (int n = 0; n < N_CHUNK; n++) {
|
||||
if (n >= an) continue;
|
||||
int slot = slot_of(gci, n);
|
||||
const bf16_t* src =
|
||||
residual_addr(block_res, layer_res, ns + n, N, t,
|
||||
block_stride_m, block_stride_r, H);
|
||||
cp_async_bulk(buf_ptr(slot), src, H * sizeof(bf16_t),
|
||||
plan.bar_ready[chunk_slot]);
|
||||
}
|
||||
if constexpr (HAS_DELTA) {
|
||||
int prefix_n = N - 1 - ns;
|
||||
if (prefix_n >= 0 && prefix_n < an) {
|
||||
cp_async_bulk(delta_buf_ptr(chunk_slot),
|
||||
delta + (long long)tb * H, H * sizeof(bf16_t),
|
||||
plan.bar_ready[chunk_slot]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
float acc32[ACC_PER_THREAD] = {};
|
||||
float eps_cache;
|
||||
asm volatile("mov.b32 %0, %1;" : "=f"(eps_cache) : "f"(rms_eps));
|
||||
|
||||
long long gci = 0;
|
||||
for (int tb = blockIdx.x; tb < TB; tb += num_ctas) {
|
||||
float m_running = -FLT_MAX;
|
||||
float s_running = 0.f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < ACC_PER_THREAD; i++) {
|
||||
acc32[i] = 0.f;
|
||||
}
|
||||
|
||||
for (int ci = 0; ci < num_chunks; ci++, gci++) {
|
||||
int ns = ci * N_CHUNK;
|
||||
int an = min(N_CHUNK, N - ns);
|
||||
int chunk_slot = (int)(gci % CHUNK_DEPTH);
|
||||
int pr = phase_of(gci);
|
||||
mbarrier_wait(plan.bar_ready[chunk_slot], pr);
|
||||
|
||||
float2 sq_local[N_CHUNK] = {};
|
||||
float2 dot_local[N_CHUNK] = {};
|
||||
|
||||
auto pass_A_body = [&](auto AN_TOK) {
|
||||
constexpr int AN = decltype(AN_TOK)::value;
|
||||
#pragma unroll
|
||||
for (int si = 0; si < SLICES_PER_GROUP; si++) {
|
||||
if constexpr (H == 7168) {
|
||||
if (si == SLICES_PER_GROUP - 1) {
|
||||
int h_base =
|
||||
6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
|
||||
const float* qv = &q_cache[si * VEC];
|
||||
#pragma unroll
|
||||
for (int n = 0; n < AN; n++) {
|
||||
int slot = slot_of(gci, n);
|
||||
int2 vp =
|
||||
*reinterpret_cast<const int2*>(buf_ptr(slot) + h_base);
|
||||
auto* v2 = reinterpret_cast<__nv_bfloat162*>(&vp);
|
||||
if constexpr (HAS_DELTA) {
|
||||
int prefix_n = N - 1 - ns;
|
||||
if (n == prefix_n) {
|
||||
const bf16_t* delta_ptr =
|
||||
delta_buf_ptr(chunk_slot) + h_base;
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 2; j++) {
|
||||
auto delta2 = *reinterpret_cast<const __nv_bfloat162*>(
|
||||
delta_ptr + 2 * j);
|
||||
v2[j] = __hadd2(v2[j], delta2);
|
||||
}
|
||||
*reinterpret_cast<int2*>(layer_res + (long long)tb * H +
|
||||
h_base) = vp;
|
||||
}
|
||||
}
|
||||
float2 f[2] = {__bfloat1622float2(v2[0]),
|
||||
__bfloat1622float2(v2[1])};
|
||||
tmem_store<4>(my_v_tmem + (si * N_CHUNK + n) * VEC, f);
|
||||
sq_local[n] = float2_fma(f[0], f[0], sq_local[n]);
|
||||
sq_local[n] = float2_fma(f[1], f[1], sq_local[n]);
|
||||
dot_local[n] =
|
||||
float2_fma(f[0], make_float2(qv[0], qv[1]), dot_local[n]);
|
||||
dot_local[n] =
|
||||
float2_fma(f[1], make_float2(qv[2], qv[3]), dot_local[n]);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
int dt = si * CONSUMER_GROUPS + group;
|
||||
if (dt >= NHT) continue;
|
||||
int h_base = dt * K_TILE + k_local;
|
||||
const float* qv = &q_cache[si * VEC];
|
||||
|
||||
#pragma unroll
|
||||
for (int n = 0; n < AN; n++) {
|
||||
int slot = slot_of(gci, n);
|
||||
int4 vp = *reinterpret_cast<const int4*>(buf_ptr(slot) + h_base);
|
||||
auto* v2 = reinterpret_cast<__nv_bfloat162*>(&vp);
|
||||
if constexpr (HAS_DELTA) {
|
||||
int prefix_n = N - 1 - ns;
|
||||
if (n == prefix_n) {
|
||||
const bf16_t* delta_ptr = delta_buf_ptr(chunk_slot) + h_base;
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC / 2; j++) {
|
||||
auto delta2 = *reinterpret_cast<const __nv_bfloat162*>(
|
||||
delta_ptr + 2 * j);
|
||||
v2[j] = __hadd2(v2[j], delta2);
|
||||
}
|
||||
*reinterpret_cast<int4*>(layer_res + (long long)tb * H +
|
||||
h_base) = vp;
|
||||
}
|
||||
}
|
||||
float2 f[4] = {
|
||||
__bfloat1622float2(v2[0]), __bfloat1622float2(v2[1]),
|
||||
__bfloat1622float2(v2[2]), __bfloat1622float2(v2[3])};
|
||||
tmem_store<VEC>(my_v_tmem + (si * N_CHUNK + n) * VEC, f);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC / 2; j++) {
|
||||
sq_local[n] = float2_fma(f[j], f[j], sq_local[n]);
|
||||
dot_local[n] = float2_fma(
|
||||
f[j], make_float2(qv[2 * j], qv[2 * j + 1]), dot_local[n]);
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
if constexpr (NC == 4) {
|
||||
switch (an) {
|
||||
case 4:
|
||||
pass_A_body(std::integral_constant<int, 4>{});
|
||||
break;
|
||||
case 3:
|
||||
pass_A_body(std::integral_constant<int, 3>{});
|
||||
break;
|
||||
case 2:
|
||||
pass_A_body(std::integral_constant<int, 2>{});
|
||||
break;
|
||||
case 1:
|
||||
pass_A_body(std::integral_constant<int, 1>{});
|
||||
break;
|
||||
default:
|
||||
__builtin_unreachable();
|
||||
}
|
||||
} else if constexpr (NC == 3) {
|
||||
switch (an) {
|
||||
case 3:
|
||||
pass_A_body(std::integral_constant<int, 3>{});
|
||||
break;
|
||||
case 2:
|
||||
pass_A_body(std::integral_constant<int, 2>{});
|
||||
break;
|
||||
case 1:
|
||||
pass_A_body(std::integral_constant<int, 1>{});
|
||||
break;
|
||||
default:
|
||||
__builtin_unreachable();
|
||||
}
|
||||
} else {
|
||||
static_assert(NC == 2);
|
||||
switch (an) {
|
||||
case 2:
|
||||
pass_A_body(std::integral_constant<int, 2>{});
|
||||
break;
|
||||
case 1:
|
||||
pass_A_body(std::integral_constant<int, 1>{});
|
||||
break;
|
||||
default:
|
||||
__builtin_unreachable();
|
||||
}
|
||||
}
|
||||
if (lane == 0) {
|
||||
mbarrier_arrive(plan.bar_consumed[chunk_slot]);
|
||||
}
|
||||
tmem_store_wait();
|
||||
|
||||
float2 reduce_pair[N_CHUNK];
|
||||
#pragma unroll
|
||||
for (int n = 0; n < N_CHUNK; n++) {
|
||||
reduce_pair[n] = make_float2(sq_local[n].x + sq_local[n].y,
|
||||
dot_local[n].x + dot_local[n].y);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int offset = 16; offset > 0; offset >>= 1) {
|
||||
#pragma unroll
|
||||
for (int n = 0; n < N_CHUNK; n++) {
|
||||
uint64_t packed = reinterpret_cast<uint64_t&>(reduce_pair[n]);
|
||||
packed = __shfl_xor_sync(0xffffffff, packed, offset);
|
||||
float2 other = reinterpret_cast<float2&>(packed);
|
||||
reduce_pair[n] = float2_add(reduce_pair[n], other);
|
||||
}
|
||||
}
|
||||
if (lane == 0) {
|
||||
#pragma unroll
|
||||
for (int n = 0; n < N_CHUNK; n++) {
|
||||
plan.ws_stats[comp_wid][n] = reduce_pair[n];
|
||||
}
|
||||
}
|
||||
named_barrier_sync(CONSUMER_THREADS, 0);
|
||||
|
||||
float local_rsig = 0.f;
|
||||
float local_logit = 0.f;
|
||||
int stat_n = lane / CONSUMER_WARPS;
|
||||
int stat_w = lane % CONSUMER_WARPS;
|
||||
float2 totals = {};
|
||||
if (stat_n < N_CHUNK) {
|
||||
totals = plan.ws_stats[stat_w][stat_n];
|
||||
}
|
||||
#pragma unroll
|
||||
for (int offset = CONSUMER_WARPS / 2; offset > 0; offset >>= 1) {
|
||||
totals.x +=
|
||||
__shfl_down_sync(0xffffffff, totals.x, offset, CONSUMER_WARPS);
|
||||
totals.y +=
|
||||
__shfl_down_sync(0xffffffff, totals.y, offset, CONSUMER_WARPS);
|
||||
}
|
||||
if (stat_n < N_CHUNK && stat_w == 0) {
|
||||
local_rsig = rsqrtf(totals.x / H + eps_cache);
|
||||
local_logit = totals.y * local_rsig;
|
||||
}
|
||||
float logit_n[N_CHUNK];
|
||||
#pragma unroll
|
||||
for (int n = 0; n < N_CHUNK; n++) {
|
||||
logit_n[n] = __shfl_sync(0xffffffff, local_logit, n * CONSUMER_WARPS);
|
||||
}
|
||||
|
||||
float m_chunk = -FLT_MAX;
|
||||
#pragma unroll
|
||||
for (int n = 0; n < N_CHUNK; n++) {
|
||||
if (n < an) m_chunk = fmaxf(m_chunk, logit_n[n]);
|
||||
}
|
||||
float m_new = fmaxf(m_running, m_chunk);
|
||||
float corr = exp2f((m_running - m_new) * LOG2_E);
|
||||
float w_n[N_CHUNK] = {};
|
||||
float w_sum = 0.f;
|
||||
#pragma unroll
|
||||
for (int n = 0; n < N_CHUNK; n++) {
|
||||
if (n < an) {
|
||||
w_n[n] = exp2f((logit_n[n] - m_new) * LOG2_E);
|
||||
w_sum += w_n[n];
|
||||
}
|
||||
}
|
||||
|
||||
auto pass_B_body = [&](auto AN_TOK) {
|
||||
constexpr int AN = decltype(AN_TOK)::value;
|
||||
#pragma unroll
|
||||
for (int si = 0; si < SLICES_PER_GROUP; si++) {
|
||||
if constexpr (H == 7168) {
|
||||
if (si == SLICES_PER_GROUP - 1) {
|
||||
float2 corr2 = make_float2(corr, corr);
|
||||
float2 a[2];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 2; j++) {
|
||||
float2 old = make_float2(acc32[si * VEC + 2 * j],
|
||||
acc32[si * VEC + 2 * j + 1]);
|
||||
a[j] = float2_mul(old, corr2);
|
||||
}
|
||||
float2 f_cache[AN][2];
|
||||
#pragma unroll
|
||||
for (int n = 0; n < AN; n++) {
|
||||
tmem_load<4>(my_v_tmem + (si * N_CHUNK + n) * VEC,
|
||||
f_cache[n]);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int n = 0; n < AN; n++) {
|
||||
float2 wn = make_float2(w_n[n], w_n[n]);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 2; j++) {
|
||||
a[j] = float2_fma(wn, f_cache[n][j], a[j]);
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 2; j++) {
|
||||
acc32[si * VEC + 2 * j] = a[j].x;
|
||||
acc32[si * VEC + 2 * j + 1] = a[j].y;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
int dt = si * CONSUMER_GROUPS + group;
|
||||
if (dt >= NHT) continue;
|
||||
float2 corr2 = make_float2(corr, corr);
|
||||
float2 a[VEC / 2];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC / 2; j++) {
|
||||
float2 old = make_float2(acc32[si * VEC + 2 * j],
|
||||
acc32[si * VEC + 2 * j + 1]);
|
||||
a[j] = float2_mul(old, corr2);
|
||||
}
|
||||
float2 f_cache[AN][VEC / 2];
|
||||
#pragma unroll
|
||||
for (int n = 0; n < AN; n++) {
|
||||
tmem_load<VEC>(my_v_tmem + (si * N_CHUNK + n) * VEC, f_cache[n]);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int n = 0; n < AN; n++) {
|
||||
float2 wn = make_float2(w_n[n], w_n[n]);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC / 2; j++) {
|
||||
a[j] = float2_fma(wn, f_cache[n][j], a[j]);
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC / 2; j++) {
|
||||
acc32[si * VEC + 2 * j] = a[j].x;
|
||||
acc32[si * VEC + 2 * j + 1] = a[j].y;
|
||||
}
|
||||
}
|
||||
};
|
||||
if constexpr (NC == 4) {
|
||||
switch (an) {
|
||||
case 4:
|
||||
pass_B_body(std::integral_constant<int, 4>{});
|
||||
break;
|
||||
case 3:
|
||||
pass_B_body(std::integral_constant<int, 3>{});
|
||||
break;
|
||||
case 2:
|
||||
pass_B_body(std::integral_constant<int, 2>{});
|
||||
break;
|
||||
case 1:
|
||||
pass_B_body(std::integral_constant<int, 1>{});
|
||||
break;
|
||||
default:
|
||||
__builtin_unreachable();
|
||||
}
|
||||
} else if constexpr (NC == 3) {
|
||||
switch (an) {
|
||||
case 3:
|
||||
pass_B_body(std::integral_constant<int, 3>{});
|
||||
break;
|
||||
case 2:
|
||||
pass_B_body(std::integral_constant<int, 2>{});
|
||||
break;
|
||||
case 1:
|
||||
pass_B_body(std::integral_constant<int, 1>{});
|
||||
break;
|
||||
default:
|
||||
__builtin_unreachable();
|
||||
}
|
||||
} else {
|
||||
static_assert(NC == 2);
|
||||
switch (an) {
|
||||
case 2:
|
||||
pass_B_body(std::integral_constant<int, 2>{});
|
||||
break;
|
||||
case 1:
|
||||
pass_B_body(std::integral_constant<int, 1>{});
|
||||
break;
|
||||
default:
|
||||
__builtin_unreachable();
|
||||
}
|
||||
}
|
||||
|
||||
s_running = s_running * corr + w_sum;
|
||||
m_running = m_new;
|
||||
}
|
||||
|
||||
float inv_s = 1.f / s_running;
|
||||
bf16_t* out_ptr = output + (long long)tb * H;
|
||||
float2 output_sq_pair = {};
|
||||
// When output RMSNorm is fused, the softmax denominator cancels:
|
||||
// (acc / s) * rsqrt(mean((acc / s)^2) + eps)
|
||||
// = acc * rsqrt(mean(acc^2) + eps * s^2).
|
||||
#pragma unroll
|
||||
for (int si = 0; si < SLICES_PER_GROUP; si++) {
|
||||
if constexpr (H == 7168) {
|
||||
if (si == SLICES_PER_GROUP - 1) {
|
||||
int h_base = 6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
|
||||
uint2 packed;
|
||||
auto* ov2 = reinterpret_cast<__nv_bfloat162*>(&packed);
|
||||
float2 inv2 = make_float2(inv_s, inv_s);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 2; j++) {
|
||||
float2 old = make_float2(acc32[si * VEC + 2 * j],
|
||||
acc32[si * VEC + 2 * j + 1]);
|
||||
if constexpr (HAS_OUTPUT_NORM) {
|
||||
output_sq_pair = float2_fma(old, old, output_sq_pair);
|
||||
} else {
|
||||
float2 mixed = float2_mul(old, inv2);
|
||||
ov2[j] = __float22bfloat162_rn(mixed);
|
||||
}
|
||||
}
|
||||
if constexpr (!HAS_OUTPUT_NORM) {
|
||||
*reinterpret_cast<uint2*>(out_ptr + h_base) = packed;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
int dt = si * CONSUMER_GROUPS + group;
|
||||
if (dt >= NHT) continue;
|
||||
int h_base = dt * K_TILE + k_local;
|
||||
uint4 packed;
|
||||
auto* ov2 = reinterpret_cast<__nv_bfloat162*>(&packed);
|
||||
float2 inv2 = make_float2(inv_s, inv_s);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC / 2; j++) {
|
||||
float2 old =
|
||||
make_float2(acc32[si * VEC + 2 * j], acc32[si * VEC + 2 * j + 1]);
|
||||
if constexpr (HAS_OUTPUT_NORM) {
|
||||
output_sq_pair = float2_fma(old, old, output_sq_pair);
|
||||
} else {
|
||||
float2 mixed = float2_mul(old, inv2);
|
||||
ov2[j] = __float22bfloat162_rn(mixed);
|
||||
}
|
||||
}
|
||||
if constexpr (!HAS_OUTPUT_NORM) {
|
||||
*reinterpret_cast<uint4*>(out_ptr + h_base) = packed;
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (HAS_OUTPUT_NORM) {
|
||||
if constexpr (OUTPUT_NORM_IN_SMEM) {
|
||||
// The immutable weight copy is acquired once, at its first use.
|
||||
if (tb == blockIdx.x) {
|
||||
mbarrier_wait(plan.bar_output_norm_ready, 0);
|
||||
}
|
||||
}
|
||||
float output_sq = output_sq_pair.x + output_sq_pair.y;
|
||||
#pragma unroll
|
||||
for (int offset = 16; offset > 0; offset >>= 1) {
|
||||
output_sq += __shfl_xor_sync(0xffffffff, output_sq, offset);
|
||||
}
|
||||
if (lane == 0) {
|
||||
plan.ws_stats[comp_wid][0] = make_float2(output_sq, 0.f);
|
||||
}
|
||||
named_barrier_sync(CONSUMER_THREADS, 0);
|
||||
float total_sq = lane < CONSUMER_WARPS ? plan.ws_stats[lane][0].x : 0.f;
|
||||
#pragma unroll
|
||||
for (int offset = CONSUMER_WARPS / 2; offset > 0; offset >>= 1) {
|
||||
total_sq +=
|
||||
__shfl_down_sync(0xffffffff, total_sq, offset, CONSUMER_WARPS);
|
||||
}
|
||||
if (lane == 0) {
|
||||
total_sq =
|
||||
rsqrtf(total_sq / H + output_norm_eps * s_running * s_running);
|
||||
}
|
||||
float output_rsigma = __shfl_sync(0xffffffff, total_sq, 0);
|
||||
#pragma unroll
|
||||
for (int si = 0; si < SLICES_PER_GROUP; si++) {
|
||||
if constexpr (H == 7168) {
|
||||
if (si == SLICES_PER_GROUP - 1) {
|
||||
int h_base = 6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
|
||||
uint2 packed;
|
||||
auto* values = reinterpret_cast<bf16_t*>(&packed);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 4; j++) {
|
||||
const bf16_t* weight_ptr =
|
||||
OUTPUT_NORM_IN_SMEM ? output_norm_buf : output_norm_weight;
|
||||
float weight = __bfloat162float(weight_ptr[h_base + j]);
|
||||
values[j] = __float2bfloat16(acc32[si * VEC + j] *
|
||||
output_rsigma * weight);
|
||||
}
|
||||
*reinterpret_cast<uint2*>(out_ptr + h_base) = packed;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
int dt = si * CONSUMER_GROUPS + group;
|
||||
if (dt >= NHT) continue;
|
||||
int h_base = dt * K_TILE + k_local;
|
||||
uint4 packed;
|
||||
auto* values = reinterpret_cast<bf16_t*>(&packed);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC; j++) {
|
||||
const bf16_t* weight_ptr =
|
||||
OUTPUT_NORM_IN_SMEM ? output_norm_buf : output_norm_weight;
|
||||
float weight = __bfloat162float(weight_ptr[h_base + j]);
|
||||
values[j] =
|
||||
__float2bfloat16(acc32[si * VEC + j] * output_rsigma * weight);
|
||||
}
|
||||
*reinterpret_cast<uint4*>(out_ptr + h_base) = packed;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
__syncthreads();
|
||||
if (wid == 1) {
|
||||
tmem_free(plan.tmem_base, TMEM_COLS_ALLOC);
|
||||
}
|
||||
#else
|
||||
if (threadIdx.x == 0) {
|
||||
printf("attn_res_fwd_online_v2_kernel requires sm_10x\n");
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
template <int H, int NC = N_CHUNK_DEFAULT, bool RELEASE_TMEM = false,
|
||||
bool HAS_DELTA = false, bool HAS_OUTPUT_NORM = false,
|
||||
bool OUTPUT_NORM_IN_SMEM = false>
|
||||
static void launch_fwd(const bf16_t* block_residual, bf16_t* layer_residual,
|
||||
const bf16_t* delta, const bf16_t* res_weight,
|
||||
const bf16_t* rms_weight, bf16_t* output, int N, int T,
|
||||
int B, float rms_eps, int num_sm, cudaStream_t stream,
|
||||
const bf16_t* output_norm_weight = nullptr,
|
||||
float output_norm_eps = 0.f, int block_stride_m = 0,
|
||||
int block_stride_r = 0) {
|
||||
constexpr size_t smem_size =
|
||||
((size_t)CHUNK_DEPTH * (NC + (HAS_DELTA ? 1 : 0)) * H * sizeof(bf16_t) +
|
||||
(OUTPUT_NORM_IN_SMEM ? (size_t)H * sizeof(bf16_t) : 0) +
|
||||
sizeof(FwdSmemPlan<NC>) + 15) &
|
||||
~size_t(15);
|
||||
auto kernel =
|
||||
&attn_res_fwd_online_v2_kernel<H, NC, RELEASE_TMEM, HAS_DELTA,
|
||||
HAS_OUTPUT_NORM, OUTPUT_NORM_IN_SMEM>;
|
||||
static bool attrs_set = false;
|
||||
if (!attrs_set) {
|
||||
if (smem_size > 48 * 1024) {
|
||||
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
}
|
||||
attrs_set = true;
|
||||
}
|
||||
int grid = RELEASE_TMEM ? num_sm * 2 : num_sm;
|
||||
cudaLaunchConfig_t config{};
|
||||
config.gridDim = grid;
|
||||
config.blockDim = BLK;
|
||||
config.dynamicSmemBytes = smem_size;
|
||||
config.stream = stream;
|
||||
cudaLaunchAttribute attrs[1];
|
||||
attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attrs[0].val.programmaticStreamSerializationAllowed = 1;
|
||||
config.attrs = attrs;
|
||||
config.numAttrs = 1;
|
||||
cudaLaunchKernelEx(&config, kernel, block_residual, layer_residual, delta,
|
||||
res_weight, rms_weight, output, N, T, B, block_stride_m,
|
||||
block_stride_r, rms_eps, output_norm_weight,
|
||||
output_norm_eps);
|
||||
}
|
||||
|
||||
} // namespace fwd_prod_v2
|
||||
} // namespace sm100
|
||||
|
||||
void kimi_k3_attn_res(torch::stable::Tensor& prefix,
|
||||
torch::stable::Tensor const& delta,
|
||||
torch::stable::Tensor const& blocks,
|
||||
torch::stable::Tensor const& norm_weight,
|
||||
torch::stable::Tensor const& qk_weight,
|
||||
torch::stable::Tensor const& output_norm_weight,
|
||||
torch::stable::Tensor& output, int64_t num_blocks,
|
||||
double eps, double output_norm_eps) {
|
||||
int const num_tokens = static_cast<int>(prefix.size(0));
|
||||
int const device = prefix.get_device_index();
|
||||
torch::stable::accelerator::DeviceGuard const device_guard(device);
|
||||
cudaDeviceProp const* properties = get_device_prop();
|
||||
STD_TORCH_CHECK(properties->major == 10,
|
||||
"Kimi K3 AttnRes requires the SM100 family");
|
||||
|
||||
using namespace sm100::fwd_prod_v2;
|
||||
// Two-source chunks and two resident CTAs are beneficial once setup is
|
||||
// amortized by the long, full eight-block prefill workload.
|
||||
if (num_blocks == 8 && num_tokens >= 4096) {
|
||||
launch_fwd<7168, 2, true, true, true, true>(
|
||||
static_cast<bf16_t const*>(blocks.data_ptr()),
|
||||
static_cast<bf16_t*>(prefix.data_ptr()),
|
||||
static_cast<bf16_t const*>(delta.data_ptr()),
|
||||
static_cast<bf16_t const*>(qk_weight.data_ptr()),
|
||||
static_cast<bf16_t const*>(norm_weight.data_ptr()),
|
||||
static_cast<bf16_t*>(output.data_ptr()),
|
||||
static_cast<int>(num_blocks) + 1, num_tokens, 1,
|
||||
static_cast<float>(eps), properties->multiProcessorCount,
|
||||
get_current_cuda_stream(device),
|
||||
static_cast<bf16_t const*>(output_norm_weight.data_ptr()),
|
||||
static_cast<float>(output_norm_eps), static_cast<int>(blocks.stride(0)),
|
||||
static_cast<int>(blocks.stride(1)));
|
||||
} else {
|
||||
launch_fwd<7168, 4, false, true, true, true>(
|
||||
static_cast<bf16_t const*>(blocks.data_ptr()),
|
||||
static_cast<bf16_t*>(prefix.data_ptr()),
|
||||
static_cast<bf16_t const*>(delta.data_ptr()),
|
||||
static_cast<bf16_t const*>(qk_weight.data_ptr()),
|
||||
static_cast<bf16_t const*>(norm_weight.data_ptr()),
|
||||
static_cast<bf16_t*>(output.data_ptr()),
|
||||
static_cast<int>(num_blocks) + 1, num_tokens, 1,
|
||||
static_cast<float>(eps), properties->multiProcessorCount,
|
||||
get_current_cuda_stream(device),
|
||||
static_cast<bf16_t const*>(output_norm_weight.data_ptr()),
|
||||
static_cast<float>(output_norm_eps), static_cast<int>(blocks.stride(0)),
|
||||
static_cast<int>(blocks.stride(1)));
|
||||
}
|
||||
cudaError_t const error = cudaGetLastError();
|
||||
STD_TORCH_CHECK(
|
||||
error == cudaSuccess,
|
||||
"Kimi K3 AttnRes kernel launch failed: ", cudaGetErrorString(error));
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -25,6 +25,7 @@
|
||||
#include "libtorch_stable/torch_utils.h"
|
||||
|
||||
#include <cmath>
|
||||
#include <cstdint>
|
||||
#include <tuple>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
@@ -448,7 +449,8 @@ enum ScoringFunc {
|
||||
SCORING_SIGMOID = 1 // apply sigmoid
|
||||
};
|
||||
|
||||
// Efficient sigmoid approximation from TensorRT-LLM
|
||||
// Adapted from
|
||||
// https://github.com/NVIDIA/TensorRT-LLM/blob/v1.3.0rc2/cpp/tensorrt_llm/kernels/noAuxTcKernels.cu
|
||||
__device__ inline float sigmoid_accurate(float x) {
|
||||
return 0.5f * tanhf(0.5f * x) + 0.5f;
|
||||
}
|
||||
@@ -890,6 +892,434 @@ __global__ void grouped_topk_fused_small_expert_count_kernel(
|
||||
#endif
|
||||
}
|
||||
|
||||
// Adapted from
|
||||
// https://github.com/flashinfer-ai/flashinfer/blob/06400d062a2d51564bbe781f6f811d0b75ca593e/include/flashinfer/trtllm/fused_moe/RoutingKernelTopK.cuh
|
||||
namespace single_group_topk {
|
||||
namespace detail {
|
||||
|
||||
static constexpr int BlockDim = 256;
|
||||
static constexpr uint32_t FullWarpMask = 0xffffffffU;
|
||||
static constexpr float InvalidScore = -INFINITY;
|
||||
|
||||
// TopK-only tuning: use wider workers and keep these tiers on the block path.
|
||||
template <int MaxNumExperts, int MaxNumTopExperts>
|
||||
static constexpr bool UseTunedBlockPath =
|
||||
MaxNumTopExperts == 16 && (MaxNumExperts == 896 || MaxNumExperts == 1024);
|
||||
|
||||
template <typename T, typename BiasT, ScoringFunc SF>
|
||||
__device__ __forceinline__ void preprocess_score(T input, BiasT correction_bias,
|
||||
float& unbiased_score,
|
||||
float& selection_score) {
|
||||
unbiased_score = 0.0F;
|
||||
selection_score = InvalidScore;
|
||||
float const input_float = cuda_cast<float, T>(input);
|
||||
float const bias = cuda_cast<float, BiasT>(correction_bias);
|
||||
if (!is_finite(input_float) || !is_finite(bias)) {
|
||||
return;
|
||||
}
|
||||
|
||||
float const unbiased = apply_scoring<SF>(input_float);
|
||||
float const biased = unbiased + bias;
|
||||
if constexpr (SF == SCORING_NONE) {
|
||||
if (!is_finite(biased)) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
unbiased_score = unbiased;
|
||||
selection_score = biased == 0.0F ? 0.0F : biased;
|
||||
}
|
||||
|
||||
template <typename IdxT>
|
||||
__device__ __forceinline__ void write_outputs(
|
||||
cg::thread_block_tile<WARP_SIZE> const& warp, float lane_selection_score,
|
||||
float lane_unbiased, int32_t lane_expert, int32_t lane, int32_t token,
|
||||
int32_t topk, float* topk_values, IdxT* topk_indices, bool renormalize,
|
||||
float routed_scaling_factor) {
|
||||
bool const finite_selection =
|
||||
lane < topk && lane_selection_score != InvalidScore;
|
||||
lane_unbiased = finite_selection ? lane_unbiased : 0.0F;
|
||||
unsigned const finite_mask = __ballot_sync(FullWarpMask, finite_selection);
|
||||
float const sum = cg::reduce(warp, lane_unbiased, cg::plus<float>{});
|
||||
|
||||
if (lane < topk) {
|
||||
float output = 0.0F;
|
||||
if (finite_mask == 0) {
|
||||
if (renormalize) {
|
||||
output = 1.0F / static_cast<float>(topk);
|
||||
}
|
||||
} else if (finite_selection) {
|
||||
float scale = routed_scaling_factor;
|
||||
if (renormalize) {
|
||||
scale /= sum + 1e-20F;
|
||||
}
|
||||
output = lane_unbiased * scale;
|
||||
}
|
||||
|
||||
int64_t const output_index = int64_t{token} * topk + lane;
|
||||
topk_values[output_index] = output;
|
||||
topk_indices[output_index] = static_cast<IdxT>(lane_expert);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
|
||||
int MaxNumExperts, int MaxNumTopExperts>
|
||||
__global__ void __launch_bounds__(BlockDim)
|
||||
single_group_topk_block_kernel(T const* scores, float* topk_values,
|
||||
IdxT* topk_indices, BiasT const* bias,
|
||||
int64_t num_experts, int64_t topk,
|
||||
bool renormalize,
|
||||
float routed_scaling_factor,
|
||||
bool enable_pdl) {
|
||||
static constexpr int NumChunks = (MaxNumExperts + WARP_SIZE - 1) / WARP_SIZE;
|
||||
static constexpr int WorkerValuesPerLane =
|
||||
UseTunedBlockPath<MaxNumExperts, MaxNumTopExperts> ? 8 : 4;
|
||||
static constexpr int ExpertsPerWorkerWarp = WorkerValuesPerLane * WARP_SIZE;
|
||||
using LaneOwnedRange =
|
||||
reduce_topk::HighExpertLaneOwnedTopKRange<MaxNumExperts,
|
||||
MaxNumTopExperts>;
|
||||
static constexpr int NumWorkerWarps =
|
||||
(MaxNumExperts + ExpertsPerWorkerWarp - 1) / ExpertsPerWorkerWarp;
|
||||
static constexpr int NumIntermediate = NumWorkerWarps * MaxNumTopExperts;
|
||||
static constexpr int MergeValuesPerLane =
|
||||
(NumIntermediate + WARP_SIZE - 1) / WARP_SIZE;
|
||||
static constexpr bool LaneOwnedResourcesFit =
|
||||
NumWorkerWarps <= BlockDim / WARP_SIZE && MergeValuesPerLane <= 64;
|
||||
static constexpr bool UseHierarchicalLaneTopK =
|
||||
LaneOwnedRange::kEnabled && LaneOwnedResourcesFit;
|
||||
|
||||
static_assert(NumChunks <= 64);
|
||||
static_assert(MaxNumTopExperts <= WARP_SIZE);
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
if (enable_pdl) {
|
||||
cudaGridDependencySynchronize();
|
||||
}
|
||||
#endif
|
||||
|
||||
__shared__ float __attribute((aligned(128))) biased_scores[MaxNumExperts];
|
||||
__shared__ float __attribute((aligned(128))) unbiased_scores[MaxNumExperts];
|
||||
|
||||
int32_t const token = static_cast<int32_t>(blockIdx.x);
|
||||
int32_t const lane = static_cast<int32_t>(threadIdx.x) % WARP_SIZE;
|
||||
int32_t const warp_id = static_cast<int32_t>(threadIdx.x) / WARP_SIZE;
|
||||
int32_t const num_experts_i32 = static_cast<int32_t>(num_experts);
|
||||
int32_t const topk_i32 = static_cast<int32_t>(topk);
|
||||
T const* token_scores = scores + int64_t{token} * num_experts;
|
||||
|
||||
for (int32_t expert = static_cast<int32_t>(threadIdx.x);
|
||||
expert < num_experts_i32; expert += BlockDim) {
|
||||
preprocess_score<T, BiasT, SF>(token_scores[expert], bias[expert],
|
||||
unbiased_scores[expert],
|
||||
biased_scores[expert]);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
auto warp = cg::tiled_partition<WARP_SIZE>(cg::this_thread_block());
|
||||
|
||||
if constexpr (UseHierarchicalLaneTopK) {
|
||||
__shared__ float
|
||||
__attribute((aligned(128))) intermediate_scores[NumIntermediate];
|
||||
__shared__ int32_t
|
||||
__attribute((aligned(128))) intermediate_indices[NumIntermediate];
|
||||
|
||||
if (warp_id < NumWorkerWarps) {
|
||||
float local_scores[WorkerValuesPerLane];
|
||||
int32_t local_indices[WorkerValuesPerLane];
|
||||
#pragma unroll
|
||||
for (int index = 0; index < WorkerValuesPerLane; ++index) {
|
||||
int32_t const expert =
|
||||
warp_id * ExpertsPerWorkerWarp + index * WARP_SIZE + lane;
|
||||
local_scores[index] =
|
||||
expert < num_experts_i32 ? biased_scores[expert] : InvalidScore;
|
||||
local_indices[index] = expert;
|
||||
}
|
||||
|
||||
float lane_score;
|
||||
int32_t lane_expert;
|
||||
reduce_topk::reduceTopKForLane<MaxNumTopExperts>(
|
||||
warp, lane_score, lane_expert, local_scores, local_indices,
|
||||
InvalidScore, lane);
|
||||
if (lane < MaxNumTopExperts) {
|
||||
int32_t const intermediate = warp_id * MaxNumTopExperts + lane;
|
||||
bool const active = lane < topk_i32;
|
||||
intermediate_scores[intermediate] = active ? lane_score : InvalidScore;
|
||||
intermediate_indices[intermediate] =
|
||||
active ? lane_expert : MaxNumExperts;
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (warp_id != 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
float merge_scores[MergeValuesPerLane];
|
||||
int32_t merge_indices[MergeValuesPerLane];
|
||||
#pragma unroll
|
||||
for (int index = 0; index < MergeValuesPerLane; ++index) {
|
||||
int32_t const intermediate = index * WARP_SIZE + lane;
|
||||
bool const active = intermediate < NumIntermediate;
|
||||
merge_scores[index] =
|
||||
active ? intermediate_scores[intermediate] : InvalidScore;
|
||||
merge_indices[index] =
|
||||
active ? intermediate_indices[intermediate] : MaxNumExperts;
|
||||
}
|
||||
|
||||
float lane_score;
|
||||
int32_t lane_expert;
|
||||
reduce_topk::reduceTopKForLane<MaxNumTopExperts>(
|
||||
warp, lane_score, lane_expert, merge_scores, merge_indices,
|
||||
InvalidScore, lane);
|
||||
float const lane_unbiased =
|
||||
lane < topk_i32 && lane_expert >= 0 && lane_expert < num_experts_i32
|
||||
? unbiased_scores[lane_expert]
|
||||
: 0.0F;
|
||||
write_outputs(warp, lane_score, lane_unbiased, lane_expert, lane, token,
|
||||
topk_i32, topk_values, topk_indices, renormalize,
|
||||
routed_scaling_factor);
|
||||
} else {
|
||||
if (warp_id != 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
float local_scores[NumChunks];
|
||||
int32_t local_indices[NumChunks];
|
||||
#pragma unroll
|
||||
for (int index = 0; index < NumChunks; ++index) {
|
||||
int32_t const expert = index * WARP_SIZE + lane;
|
||||
local_scores[index] =
|
||||
expert < num_experts_i32 ? biased_scores[expert] : InvalidScore;
|
||||
local_indices[index] = expert;
|
||||
}
|
||||
|
||||
float top_scores[MaxNumTopExperts];
|
||||
int32_t top_experts[MaxNumTopExperts];
|
||||
reduce_topk::reduceTopK(warp, top_scores, top_experts, local_scores,
|
||||
local_indices, InvalidScore, topk_i32);
|
||||
float const lane_score = lane < topk_i32 ? top_scores[lane] : InvalidScore;
|
||||
int32_t const lane_expert = lane < topk_i32 ? top_experts[lane] : -1;
|
||||
float const lane_unbiased =
|
||||
lane < topk_i32 && lane_expert >= 0 && lane_expert < num_experts_i32
|
||||
? unbiased_scores[lane_expert]
|
||||
: 0.0F;
|
||||
write_outputs(warp, lane_score, lane_unbiased, lane_expert, lane, token,
|
||||
topk_i32, topk_values, topk_indices, renormalize,
|
||||
routed_scaling_factor);
|
||||
}
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
if (enable_pdl) {
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
template <int MaxNumExperts>
|
||||
struct WarpTopKLaunchConfig {
|
||||
static constexpr int DefaultBlockDim =
|
||||
MaxNumExperts <= 1024 ? MaxNumExperts : 1024;
|
||||
static constexpr int BlockDim = DefaultBlockDim > 256 ? 256 : DefaultBlockDim;
|
||||
static constexpr int NumWarps = BlockDim / WARP_SIZE;
|
||||
static constexpr int MaxBlockScale =
|
||||
(DefaultBlockDim + BlockDim - 1) / BlockDim;
|
||||
static constexpr int MaxBlocks = 1024 * MaxBlockScale;
|
||||
|
||||
static_assert(BlockDim % WARP_SIZE == 0);
|
||||
|
||||
static uint32_t grid_dim(int64_t num_tokens) {
|
||||
int64_t const token_blocks = (num_tokens + NumWarps - 1) / NumWarps;
|
||||
int64_t const selected =
|
||||
token_blocks < MaxBlocks ? token_blocks : MaxBlocks;
|
||||
return static_cast<uint32_t>(selected > 0 ? selected : 1);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
|
||||
int MaxNumExperts, int MaxNumTopExperts>
|
||||
__global__ void __launch_bounds__(WarpTopKLaunchConfig<MaxNumExperts>::BlockDim)
|
||||
single_group_topk_warp_kernel(T const* scores, float* topk_values,
|
||||
IdxT* topk_indices, BiasT const* bias,
|
||||
int64_t num_tokens, int64_t num_experts,
|
||||
int64_t topk, bool renormalize,
|
||||
float routed_scaling_factor,
|
||||
bool enable_pdl) {
|
||||
static constexpr int NumChunks = (MaxNumExperts + WARP_SIZE - 1) / WARP_SIZE;
|
||||
static constexpr int WarpBlockDim =
|
||||
WarpTopKLaunchConfig<MaxNumExperts>::BlockDim;
|
||||
using LaneOwnedRange =
|
||||
reduce_topk::HighExpertLaneOwnedTopKRange<MaxNumExperts,
|
||||
MaxNumTopExperts>;
|
||||
|
||||
static_assert(NumChunks <= 64);
|
||||
static_assert(MaxNumTopExperts <= WARP_SIZE);
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
if (enable_pdl) {
|
||||
cudaGridDependencySynchronize();
|
||||
}
|
||||
#endif
|
||||
|
||||
int32_t const lane = static_cast<int32_t>(threadIdx.x) % WARP_SIZE;
|
||||
int32_t const warp_id = static_cast<int32_t>(threadIdx.x) / WARP_SIZE;
|
||||
int32_t const global_warp =
|
||||
static_cast<int32_t>(blockIdx.x) * WarpBlockDim / WARP_SIZE + warp_id;
|
||||
int32_t const global_warp_stride =
|
||||
static_cast<int32_t>(gridDim.x) * WarpBlockDim / WARP_SIZE;
|
||||
int32_t const num_experts_i32 = static_cast<int32_t>(num_experts);
|
||||
int32_t const topk_i32 = static_cast<int32_t>(topk);
|
||||
auto warp = cg::tiled_partition<WARP_SIZE>(cg::this_thread_block());
|
||||
|
||||
for (int32_t token = global_warp; token < num_tokens;
|
||||
token += global_warp_stride) {
|
||||
T const* token_scores = scores + int64_t{token} * num_experts;
|
||||
float local_scores[NumChunks];
|
||||
int32_t local_indices[NumChunks];
|
||||
#pragma unroll
|
||||
for (int index = 0; index < NumChunks; ++index) {
|
||||
int32_t const expert = index * WARP_SIZE + lane;
|
||||
float unbiased;
|
||||
float selection;
|
||||
if (expert < num_experts_i32) {
|
||||
preprocess_score<T, BiasT, SF>(token_scores[expert], bias[expert],
|
||||
unbiased, selection);
|
||||
} else {
|
||||
selection = InvalidScore;
|
||||
}
|
||||
local_scores[index] = selection;
|
||||
local_indices[index] = expert;
|
||||
}
|
||||
|
||||
float lane_score;
|
||||
int32_t lane_expert;
|
||||
if constexpr (LaneOwnedRange::kEnabled) {
|
||||
reduce_topk::reduceTopKForLane<MaxNumTopExperts>(
|
||||
warp, lane_score, lane_expert, local_scores, local_indices,
|
||||
InvalidScore, lane);
|
||||
} else {
|
||||
float top_scores[MaxNumTopExperts];
|
||||
int32_t top_experts[MaxNumTopExperts];
|
||||
reduce_topk::reduceTopK(warp, top_scores, top_experts, local_scores,
|
||||
local_indices, InvalidScore, topk_i32);
|
||||
lane_score = lane < topk_i32 ? top_scores[lane] : InvalidScore;
|
||||
lane_expert = lane < topk_i32 ? top_experts[lane] : -1;
|
||||
}
|
||||
|
||||
float lane_unbiased = 0.0F;
|
||||
if (lane < topk_i32 && lane_expert >= 0 && lane_expert < num_experts_i32) {
|
||||
lane_unbiased = lane_score - cuda_cast<float, BiasT>(bias[lane_expert]);
|
||||
}
|
||||
write_outputs(warp, lane_score, lane_unbiased, lane_expert, lane, token,
|
||||
topk_i32, topk_values, topk_indices, renormalize,
|
||||
routed_scaling_factor);
|
||||
}
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
if (enable_pdl) {
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
template <int Experts, int TopK>
|
||||
struct Tier {
|
||||
static constexpr int kExperts = Experts;
|
||||
static constexpr int kTopK = TopK;
|
||||
};
|
||||
|
||||
template <typename... Tiers>
|
||||
struct TierList {};
|
||||
|
||||
using SigmoidBiasTiers =
|
||||
TierList<Tier<128, 8>, Tier<256, 8>, Tier<384, 8>, Tier<512, 8>,
|
||||
Tier<512, 22>, Tier<768, 16>, Tier<896, 16>, Tier<1024, 16>>;
|
||||
|
||||
using PrecomputedSoftmaxBiasTiers =
|
||||
TierList<Tier<128, 4>, Tier<128, 8>, Tier<160, 8>, Tier<256, 8>,
|
||||
Tier<256, 16>, Tier<512, 8>, Tier<512, 16>, Tier<512, 22>,
|
||||
Tier<512, 32>, Tier<576, 8>, Tier<768, 16>, Tier<896, 16>,
|
||||
Tier<1024, 16>>;
|
||||
|
||||
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
|
||||
int MaxNumExperts, int MaxNumTopExperts>
|
||||
void launch(T* scores, float* topk_values, IdxT* topk_indices,
|
||||
BiasT const* bias, int64_t num_tokens, int64_t num_experts,
|
||||
int64_t topk, bool renormalize, double routed_scaling_factor,
|
||||
bool enable_pdl, cudaLaunchConfig_t& config) {
|
||||
config.dynamicSmemBytes = 0;
|
||||
bool const use_block_kernel =
|
||||
UseTunedBlockPath<MaxNumExperts, MaxNumTopExperts> ||
|
||||
MaxNumExperts > 1024 || num_experts >= 1024 ||
|
||||
(num_experts >= 256 && num_tokens <= 1024);
|
||||
if (use_block_kernel) {
|
||||
config.gridDim = static_cast<uint32_t>(num_tokens);
|
||||
config.blockDim = BlockDim;
|
||||
cudaLaunchKernelEx(
|
||||
&config,
|
||||
&single_group_topk_block_kernel<T, BiasT, IdxT, SF, MaxNumExperts,
|
||||
MaxNumTopExperts>,
|
||||
scores, topk_values, topk_indices, bias, num_experts, topk, renormalize,
|
||||
static_cast<float>(routed_scaling_factor), enable_pdl);
|
||||
} else {
|
||||
using WarpConfig = WarpTopKLaunchConfig<MaxNumExperts>;
|
||||
config.gridDim = WarpConfig::grid_dim(num_tokens);
|
||||
config.blockDim = WarpConfig::BlockDim;
|
||||
cudaLaunchKernelEx(
|
||||
&config,
|
||||
&single_group_topk_warp_kernel<T, BiasT, IdxT, SF, MaxNumExperts,
|
||||
MaxNumTopExperts>,
|
||||
scores, topk_values, topk_indices, bias, num_tokens, num_experts, topk,
|
||||
renormalize, static_cast<float>(routed_scaling_factor), enable_pdl);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF>
|
||||
bool dispatch(TierList<>*, T*, float*, IdxT*, BiasT const*, int64_t, int64_t,
|
||||
int64_t, bool, double, bool, cudaLaunchConfig_t&) {
|
||||
return false;
|
||||
}
|
||||
|
||||
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
|
||||
typename First, typename... Rest>
|
||||
bool dispatch(TierList<First, Rest...>*, T* scores, float* topk_values,
|
||||
IdxT* topk_indices, BiasT const* bias, int64_t num_tokens,
|
||||
int64_t num_experts, int64_t topk, bool renormalize,
|
||||
double routed_scaling_factor, bool enable_pdl,
|
||||
cudaLaunchConfig_t& config) {
|
||||
if (num_experts <= First::kExperts && topk <= First::kTopK) {
|
||||
launch<T, BiasT, IdxT, SF, First::kExperts, First::kTopK>(
|
||||
scores, topk_values, topk_indices, bias, num_tokens, num_experts, topk,
|
||||
renormalize, routed_scaling_factor, enable_pdl, config);
|
||||
return true;
|
||||
}
|
||||
return dispatch<T, BiasT, IdxT, SF>(
|
||||
static_cast<TierList<Rest...>*>(nullptr), scores, topk_values,
|
||||
topk_indices, bias, num_tokens, num_experts, topk, renormalize,
|
||||
routed_scaling_factor, enable_pdl, config);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF>
|
||||
bool invoke(T* scores, float* topk_values, IdxT* topk_indices,
|
||||
BiasT const* bias, int64_t num_tokens, int64_t num_experts,
|
||||
int64_t topk, bool renormalize, double routed_scaling_factor,
|
||||
bool enable_pdl, cudaLaunchConfig_t& config) {
|
||||
static_assert(SF == SCORING_NONE || SF == SCORING_SIGMOID);
|
||||
if constexpr (SF == SCORING_SIGMOID) {
|
||||
return detail::dispatch<T, BiasT, IdxT, SF>(
|
||||
static_cast<detail::SigmoidBiasTiers*>(nullptr), scores, topk_values,
|
||||
topk_indices, bias, num_tokens, num_experts, topk, renormalize,
|
||||
routed_scaling_factor, enable_pdl, config);
|
||||
} else {
|
||||
return detail::dispatch<T, BiasT, IdxT, SF>(
|
||||
static_cast<detail::PrecomputedSoftmaxBiasTiers*>(nullptr), scores,
|
||||
topk_values, topk_indices, bias, num_tokens, num_experts, topk,
|
||||
renormalize, routed_scaling_factor, enable_pdl, config);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace single_group_topk
|
||||
|
||||
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF>
|
||||
void invokeNoAuxTc(T* scores, float* topk_values, IdxT* topk_indices,
|
||||
BiasT const* bias, int64_t const num_tokens,
|
||||
@@ -905,6 +1335,12 @@ void invokeNoAuxTc(T* scores, float* topk_values, IdxT* topk_indices,
|
||||
attrs[0].val.programmaticStreamSerializationAllowed = enable_pdl;
|
||||
config.numAttrs = 1;
|
||||
config.attrs = attrs;
|
||||
if (n_group == 1 && topk_group == 1 &&
|
||||
single_group_topk::invoke<T, BiasT, IdxT, SF>(
|
||||
scores, topk_values, topk_indices, bias, num_tokens, num_experts,
|
||||
topk, renormalize, routed_scaling_factor, enable_pdl, config)) {
|
||||
return;
|
||||
}
|
||||
|
||||
// Check if we can use the optimized
|
||||
// grouped_topk_fused_small_expert_count_kernel
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
/*
|
||||
* Adapted from
|
||||
* https://github.com/NVIDIA/TensorRT-LLM/blob/v1.3.0rc2/cpp/tensorrt_llm/kernels/moeTopKFuncs.cuh
|
||||
* https://github.com/flashinfer-ai/flashinfer/blob/06400d062a2d51564bbe781f6f811d0b75ca593e/include/flashinfer/trtllm/fused_moe/RoutingKernelTopK.cuh
|
||||
* Copyright (c) 2026, The vLLM team.
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION. All rights
|
||||
* reserved. SPDX-License-Identifier: Apache-2.0
|
||||
@@ -23,6 +24,9 @@
|
||||
#include <cooperative_groups/reduce.h>
|
||||
#include <cub/cub.cuh>
|
||||
|
||||
#include <cstdint>
|
||||
#include <type_traits>
|
||||
|
||||
namespace vllm {
|
||||
namespace moe {
|
||||
namespace reduce_topk {
|
||||
@@ -38,11 +42,10 @@ struct TopKRedType {
|
||||
"Top K reduction only implemented for int, float, float16 and bfloat16");
|
||||
|
||||
using TypeCmp = std::conditional_t<sizeof(T) == 4, uint64_t, uint32_t>;
|
||||
using IdxT = std::conditional_t<sizeof(T) == 4, int32_t, int16_t>;
|
||||
|
||||
static constexpr int kMoveBits = (sizeof(T) == 4) ? 32 : 16;
|
||||
static constexpr int kMaxIdx = 65535;
|
||||
TypeCmp compValIdx;
|
||||
TypeCmp compVal;
|
||||
|
||||
static __host__ __device__ inline TypeCmp makeCmpVal(T val, int32_t idx = 0) {
|
||||
auto valueBits = cub::Traits<T>::TwiddleIn(
|
||||
@@ -69,69 +72,175 @@ struct TopKRedType {
|
||||
__host__ __device__ TopKRedType() = default;
|
||||
|
||||
__host__ __device__ TopKRedType(T val, int32_t idx)
|
||||
: compValIdx(makeCmpVal(val, idx)) {}
|
||||
: compVal(makeCmpVal(val, idx)) {}
|
||||
|
||||
__host__ __device__ operator TypeCmp() const noexcept { return compValIdx; }
|
||||
__host__ __device__ operator TypeCmp() const noexcept { return compVal; }
|
||||
|
||||
__device__ inline TypeCmp reduce(
|
||||
cg::thread_block_tile<kWARP_SIZE> const& warp) {
|
||||
return cg::reduce(warp, compValIdx, cg::greater<TypeCmp>{});
|
||||
#ifdef __CUDA_ARCH__
|
||||
static constexpr bool kHAS_FAST_REDUX = (__CUDA_ARCH__ / 100) >= 10;
|
||||
#else
|
||||
static constexpr bool kHAS_FAST_REDUX = false;
|
||||
#endif
|
||||
if constexpr (!kHAS_FAST_REDUX) {
|
||||
return cg::reduce(warp, compVal, cg::greater<TypeCmp>{});
|
||||
} else if constexpr (sizeof(TypeCmp) == 8) {
|
||||
uint32_t hi = static_cast<uint32_t>(compVal >> 32);
|
||||
uint32_t lo = static_cast<uint32_t>(compVal & 0xffffffffu);
|
||||
uint32_t maxHi;
|
||||
asm volatile("redux.sync.max.u32 %0, %1, 0xffffffff;\n"
|
||||
: "=r"(maxHi)
|
||||
: "r"(hi));
|
||||
uint32_t loContrib = hi == maxHi ? lo : 0u;
|
||||
uint32_t maxLo;
|
||||
asm volatile("redux.sync.max.u32 %0, %1, 0xffffffff;\n"
|
||||
: "=r"(maxLo)
|
||||
: "r"(loContrib));
|
||||
return (static_cast<TypeCmp>(maxHi) << 32) | static_cast<TypeCmp>(maxLo);
|
||||
} else {
|
||||
TypeCmp result;
|
||||
asm volatile("redux.sync.max.u32 %0, %1, 0xffffffff;\n"
|
||||
: "=r"(result)
|
||||
: "r"(compVal));
|
||||
return result;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <int K_, bool Enable_>
|
||||
struct TopKIdx {
|
||||
// by default, empty
|
||||
template <int N>
|
||||
struct IsPowerOf2 {
|
||||
static constexpr bool value = N > 0 && (N & (N - 1)) == 0;
|
||||
};
|
||||
|
||||
template <int K_>
|
||||
struct TopKIdx<K_, true> {
|
||||
static constexpr int K = K_;
|
||||
int32_t val[K];
|
||||
template <int N>
|
||||
struct NextPow2 {
|
||||
private:
|
||||
static constexpr unsigned u = static_cast<unsigned>(N - 1);
|
||||
static constexpr unsigned s1 = u | (u >> 1);
|
||||
static constexpr unsigned s2 = s1 | (s1 >> 2);
|
||||
static constexpr unsigned s3 = s2 | (s2 >> 4);
|
||||
static constexpr unsigned s4 = s3 | (s3 >> 8);
|
||||
static constexpr unsigned s5 = s4 | (s4 >> 16);
|
||||
|
||||
public:
|
||||
static constexpr int value = N <= 1 ? 1 : static_cast<int>(s5 + 1);
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#define TOPK_SWAP(I, J) \
|
||||
{ \
|
||||
auto pairMin = min(topK[I].compValIdx, topK[J].compValIdx); \
|
||||
auto pairMax = max(topK[I].compValIdx, topK[J].compValIdx); \
|
||||
topK[I].compValIdx = pairMax; \
|
||||
topK[J].compValIdx = pairMin; \
|
||||
template <int A, int B, int Size, typename T>
|
||||
__device__ __forceinline__ void topkCompareSwap(T* a) {
|
||||
if constexpr (A < Size && B < Size) {
|
||||
if (a[A] < a[B]) {
|
||||
T tmp = a[A];
|
||||
a[A] = a[B];
|
||||
a[B] = tmp;
|
||||
}
|
||||
} else {
|
||||
(void)a;
|
||||
}
|
||||
}
|
||||
|
||||
template <int I, int End, int Step, int PairStride, int Size, typename T>
|
||||
__device__ __forceinline__ void topkMergePairs(T* a) {
|
||||
if constexpr (I + Step < End) {
|
||||
topkCompareSwap<I, I + Step, Size, T>(a);
|
||||
topkMergePairs<I + PairStride, End, Step, PairStride, Size, T>(a);
|
||||
} else {
|
||||
(void)a;
|
||||
}
|
||||
}
|
||||
|
||||
template <int Lo, int N, int R, int Size, typename T>
|
||||
__device__ __forceinline__ void topkOEM(T* a) {
|
||||
constexpr int M = R * 2;
|
||||
if constexpr (M < N) {
|
||||
topkOEM<Lo, N, M, Size, T>(a);
|
||||
topkOEM<Lo + R, N - R, M, Size, T>(a);
|
||||
topkMergePairs<Lo + R, Lo + N, R, M, Size, T>(a);
|
||||
} else if constexpr (R < N) {
|
||||
topkCompareSwap<Lo, Lo + R, Size, T>(a);
|
||||
} else {
|
||||
(void)a;
|
||||
}
|
||||
}
|
||||
|
||||
template <int Lo, int N, int Size, typename T>
|
||||
__device__ __forceinline__ void topkSortBatcher(T* a) {
|
||||
if constexpr (N > 1) {
|
||||
constexpr int Half = N / 2;
|
||||
topkSortBatcher<Lo, Half, Size, T>(a);
|
||||
topkSortBatcher<Lo + Half, N - Half, Size, T>(a);
|
||||
topkOEM<Lo, N, 1, Size, T>(a);
|
||||
} else {
|
||||
(void)a;
|
||||
}
|
||||
}
|
||||
|
||||
template <int N, typename RedType>
|
||||
struct Sort;
|
||||
struct Sort {
|
||||
static_assert(N > 0 && N <= 64, "Sort only supports N in range [1, 64]");
|
||||
|
||||
static __device__ void run(RedType* topK) {
|
||||
if constexpr (IsPowerOf2<N>::value) {
|
||||
#pragma unroll
|
||||
for (int k = 2; k <= N; k *= 2) {
|
||||
#pragma unroll
|
||||
for (int j = k / 2; j > 0; j /= 2) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < N; ++i) {
|
||||
int ixj = i ^ j;
|
||||
if (ixj > i) {
|
||||
if ((i & k) == 0) {
|
||||
if (topK[i].compVal < topK[ixj].compVal) {
|
||||
auto tmp = topK[i].compVal;
|
||||
topK[i].compVal = topK[ixj].compVal;
|
||||
topK[ixj].compVal = tmp;
|
||||
}
|
||||
} else {
|
||||
if (topK[i].compVal > topK[ixj].compVal) {
|
||||
auto tmp = topK[i].compVal;
|
||||
topK[i].compVal = topK[ixj].compVal;
|
||||
topK[ixj].compVal = tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
constexpr int P = NextPow2<N>::value;
|
||||
topkSortBatcher<0, P, N, RedType>(topK);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<1, RedType> {
|
||||
static __device__ void run(RedType* topK) {}
|
||||
static __device__ void run(RedType*) {}
|
||||
};
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<2, RedType> {
|
||||
static __device__ void run(RedType* topK) { TOPK_SWAP(0, 1); }
|
||||
static __device__ void run(RedType* topK) { topkCompareSwap<0, 1, 2>(topK); }
|
||||
};
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<3, RedType> {
|
||||
static __device__ void run(RedType* topK) {
|
||||
TOPK_SWAP(0, 1);
|
||||
TOPK_SWAP(1, 2);
|
||||
TOPK_SWAP(0, 1);
|
||||
topkCompareSwap<0, 1, 3>(topK);
|
||||
topkCompareSwap<1, 2, 3>(topK);
|
||||
topkCompareSwap<0, 1, 3>(topK);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<4, RedType> {
|
||||
static __device__ void run(RedType* topK) {
|
||||
TOPK_SWAP(0, 2);
|
||||
TOPK_SWAP(1, 3);
|
||||
TOPK_SWAP(0, 1);
|
||||
TOPK_SWAP(2, 3);
|
||||
TOPK_SWAP(1, 2);
|
||||
topkCompareSwap<0, 2, 4>(topK);
|
||||
topkCompareSwap<1, 3, 4>(topK);
|
||||
topkCompareSwap<0, 1, 4>(topK);
|
||||
topkCompareSwap<2, 3, 4>(topK);
|
||||
topkCompareSwap<1, 2, 4>(topK);
|
||||
}
|
||||
};
|
||||
|
||||
@@ -147,110 +256,112 @@ __forceinline__ __device__ void reduceTopK(
|
||||
typename RedType::TypeCmp packedMax{};
|
||||
#pragma unroll
|
||||
for (int kk = 0; kk < actualK; ++kk) {
|
||||
topK =
|
||||
kk > 0 && packedMax == topK.compValIdx ? RedType{minValue, idx} : topK;
|
||||
// get the next largest value
|
||||
topK = kk > 0 && packedMax == topK.compVal ? RedType{minValue, idx} : topK;
|
||||
packedMax = topK.reduce(warp);
|
||||
RedType::unpack(out[kk], outIdx[kk], packedMax);
|
||||
}
|
||||
};
|
||||
|
||||
template <int K, typename Type, int N, bool IsSorted = false>
|
||||
__device__ void reduceTopKFunc(cg::thread_block_tile<kWARP_SIZE> const& warp,
|
||||
Type (&out)[K], int32_t (&outIdx)[K],
|
||||
Type (&value)[N], int32_t (&idx)[N],
|
||||
Type minValue, int actualK = K) {
|
||||
static_assert(K > 0, "Top K must have K > 0");
|
||||
static_assert(K < kWARP_SIZE, "Top K must have K < kWARP_SIZE");
|
||||
static_assert(N > 0, "Top K must have N > 0");
|
||||
static_assert(N < 5,
|
||||
"Only support candidates number less than or equal to 128");
|
||||
using RedType = TopKRedType<Type>;
|
||||
RedType topK[N];
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = RedType{value[nn], idx[nn]};
|
||||
}
|
||||
|
||||
if constexpr (!IsSorted) {
|
||||
Sort<N, RedType>::run(topK);
|
||||
}
|
||||
typename RedType::TypeCmp packedMax{};
|
||||
#pragma unroll
|
||||
for (int kk = 0; kk < actualK; ++kk) {
|
||||
bool update = kk > 0 && packedMax == topK[0].compValIdx;
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = update && nn == N - 1 ? RedType{minValue, idx[nn]}
|
||||
: update ? topK[nn + 1]
|
||||
: topK[nn];
|
||||
}
|
||||
// get the next largest value
|
||||
packedMax = topK[0].reduce(warp);
|
||||
RedType::unpack(out[kk], outIdx[kk], packedMax);
|
||||
}
|
||||
};
|
||||
|
||||
template <int K, typename Type, int N>
|
||||
__forceinline__ __device__ void reduceTopK(
|
||||
cg::thread_block_tile<kWARP_SIZE> const& warp, Type (&out)[K],
|
||||
int32_t (&outIdx)[K], Type (&value)[N], int32_t (&idx)[N],
|
||||
Type const minValue, int actualK = K) {
|
||||
static_assert(K > 0, "Top K must have K > 0");
|
||||
static_assert(K < kWARP_SIZE, "Top K must have K < kWARP_SIZE");
|
||||
static_assert(K <= kWARP_SIZE, "Top K must have K <= kWARP_SIZE");
|
||||
static_assert(N > 0, "Top K must have N > 0");
|
||||
static_assert(
|
||||
N <= 16,
|
||||
"Only support candidates number less than or equal to 16*32=512");
|
||||
static_assert(N <= 4 || N % 4 == 0,
|
||||
"Only support candidates number is a multiple of 4*32=128 or "
|
||||
"less than or equal to 4");
|
||||
static_assert(N <= 64,
|
||||
"Only support candidates number less than or equal to "
|
||||
"64*32=2048");
|
||||
using RedType = TopKRedType<Type>;
|
||||
RedType topK[N];
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = RedType{value[nn], idx[nn]};
|
||||
}
|
||||
|
||||
if constexpr (N <= 4) {
|
||||
reduceTopKFunc<K, Type, N>(warp, out, outIdx, value, idx, minValue,
|
||||
actualK);
|
||||
} else {
|
||||
constexpr int numLoops = N / 4;
|
||||
constexpr int numResults = (numLoops * K - 1) / kWARP_SIZE + 1;
|
||||
Sort<N, RedType>::run(topK);
|
||||
|
||||
Type topKBufferValue[numResults];
|
||||
int32_t topKBufferIdx[numResults];
|
||||
int32_t laneIdx = threadIdx.x % kWARP_SIZE;
|
||||
|
||||
for (int ii = 0; ii < numResults; ++ii) {
|
||||
topKBufferValue[ii] = minValue;
|
||||
topKBufferIdx[ii] = ii * kWARP_SIZE - 1;
|
||||
typename RedType::TypeCmp packedMax{};
|
||||
for (int kk = 0; kk < actualK; ++kk) {
|
||||
bool update = kk > 0 && packedMax == topK[0].compVal;
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = update && nn == N - 1 ? RedType{minValue, idx[nn]}
|
||||
: update ? topK[nn + 1]
|
||||
: topK[nn];
|
||||
}
|
||||
for (int loop = 0; loop < numLoops; ++loop) {
|
||||
int start = loop * 4;
|
||||
Type topKValue[K];
|
||||
int32_t topKIdx[K];
|
||||
Type inValue[4];
|
||||
int32_t inIdx[4];
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
inValue[i] = value[start + i];
|
||||
inIdx[i] = idx[start + i];
|
||||
}
|
||||
reduceTopKFunc<K, Type, 4>(warp, topKValue, topKIdx, inValue, inIdx,
|
||||
minValue, actualK);
|
||||
int inOffset = laneIdx % K;
|
||||
if (laneIdx >= loop * K && laneIdx < (loop + 1) * K) {
|
||||
topKBufferValue[0] = topKValue[inOffset];
|
||||
topKBufferIdx[0] = topKIdx[inOffset];
|
||||
}
|
||||
if (loop == numLoops - 1 && (laneIdx < (numLoops * K - kWARP_SIZE))) {
|
||||
topKBufferValue[1] = topKValue[inOffset];
|
||||
topKBufferIdx[1] = topKIdx[inOffset];
|
||||
}
|
||||
}
|
||||
|
||||
reduceTopKFunc<K, Type, numResults>(warp, out, outIdx, topKBufferValue,
|
||||
topKBufferIdx, minValue, actualK);
|
||||
packedMax = topK[0].reduce(warp);
|
||||
RedType::unpack(out[kk], outIdx[kk], packedMax);
|
||||
}
|
||||
};
|
||||
|
||||
#undef TOPK_SWAP
|
||||
template <int NumExperts, int NumTopExperts, int MinExperts, int MaxExperts,
|
||||
int MinTopExperts, int MaxTopExperts>
|
||||
struct LaneOwnedTopKRange {
|
||||
static_assert(MinExperts > 0 && MinExperts <= MaxExperts);
|
||||
static_assert(MinTopExperts > 0 && MinTopExperts <= MaxTopExperts);
|
||||
static constexpr bool kEnabled =
|
||||
NumExperts >= MinExperts && NumExperts <= MaxExperts &&
|
||||
NumTopExperts >= MinTopExperts && NumTopExperts <= MaxTopExperts;
|
||||
};
|
||||
|
||||
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_EXPERTS = 512;
|
||||
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_EXPERTS = 1024;
|
||||
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_TOP_EXPERTS = 9;
|
||||
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_TOP_EXPERTS = 16;
|
||||
|
||||
template <int NumExperts, int NumTopExperts>
|
||||
using HighExpertLaneOwnedTopKRange =
|
||||
LaneOwnedTopKRange<NumExperts, NumTopExperts,
|
||||
kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_EXPERTS,
|
||||
kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_EXPERTS,
|
||||
kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_TOP_EXPERTS,
|
||||
kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_TOP_EXPERTS>;
|
||||
|
||||
template <int K, typename Type, int N>
|
||||
__forceinline__ __device__ void reduceTopKForLane(
|
||||
cg::thread_block_tile<kWARP_SIZE> const& warp, Type& out, int32_t& outIdx,
|
||||
Type (&value)[N], int32_t (&idx)[N], Type const minValue, int32_t laneIdx) {
|
||||
static_assert(K > 0, "Top K must have K > 0");
|
||||
static_assert(K <= kWARP_SIZE, "Top K must have K <= kWARP_SIZE");
|
||||
static_assert(N > 0, "Top K must have N > 0");
|
||||
static_assert(N <= 64,
|
||||
"Only support candidates number less than or equal to "
|
||||
"64*32=2048");
|
||||
using RedType = TopKRedType<Type>;
|
||||
RedType topK[N];
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = RedType{value[nn], idx[nn]};
|
||||
}
|
||||
|
||||
Sort<N, RedType>::run(topK);
|
||||
|
||||
typename RedType::TypeCmp packedMax{};
|
||||
typename RedType::TypeCmp lanePacked{};
|
||||
#pragma unroll
|
||||
for (int kk = 0; kk < K; ++kk) {
|
||||
bool update = kk > 0 && packedMax == topK[0].compVal;
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = update && nn == N - 1 ? RedType{minValue, idx[nn]}
|
||||
: update ? topK[nn + 1]
|
||||
: topK[nn];
|
||||
}
|
||||
packedMax = topK[0].reduce(warp);
|
||||
if (laneIdx == kk) {
|
||||
lanePacked = packedMax;
|
||||
}
|
||||
}
|
||||
|
||||
if (laneIdx < K) {
|
||||
RedType::unpack(out, outIdx, lanePacked);
|
||||
} else {
|
||||
out = minValue;
|
||||
outIdx = -1;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace reduce_topk
|
||||
} // namespace moe
|
||||
|
||||
@@ -360,12 +360,24 @@ __global__ void count_and_sort_expert_tokens_kernel(
|
||||
template <typename scalar_t>
|
||||
constexpr int MOE_SUM_VEC = 16 / sizeof(scalar_t);
|
||||
|
||||
template <typename scalar_t, int TOPK>
|
||||
template <typename idx_t>
|
||||
__device__ __forceinline__ bool moe_sum_pad_aware_skip(
|
||||
const idx_t* __restrict__ topk_ids, const int32_t* __restrict__ expert_map,
|
||||
int64_t idx) {
|
||||
int64_t expert_id = static_cast<int64_t>(topk_ids[idx]);
|
||||
if (expert_id < 0) return true;
|
||||
if (expert_map != nullptr && expert_map[expert_id] < 0) return true;
|
||||
return false;
|
||||
}
|
||||
|
||||
template <typename scalar_t, typename idx_t, int TOPK, bool PAD_AWARE>
|
||||
__global__ void moe_sum_vec_kernel(
|
||||
scalar_t* __restrict__ out, // [num_tokens, d], contiguous
|
||||
const scalar_t* __restrict__ input, // [num_tokens, topk, d], d contiguous
|
||||
const int64_t num_tokens, const int d, const int64_t stride_token,
|
||||
const int64_t stride_topk) {
|
||||
const int64_t stride_topk, const idx_t* __restrict__ topk_ids,
|
||||
const int32_t* __restrict__ expert_map, const int64_t stride_tk_token,
|
||||
const int64_t stride_tk_k) {
|
||||
using vec_t = vllm::vec_n_t<scalar_t, MOE_SUM_VEC<scalar_t>>; // 16-byte pack
|
||||
constexpr int VEC = MOE_SUM_VEC<scalar_t>;
|
||||
const int64_t n_vec = d / VEC;
|
||||
@@ -375,6 +387,10 @@ __global__ void moe_sum_vec_kernel(
|
||||
const int64_t token = i / n_vec;
|
||||
const int64_t v = i % n_vec;
|
||||
const scalar_t* in_tok = input + token * stride_token + v * VEC;
|
||||
const idx_t* tk_tok = nullptr;
|
||||
if constexpr (PAD_AWARE) {
|
||||
tk_tok = topk_ids + token * stride_tk_token;
|
||||
}
|
||||
|
||||
float acc[VEC];
|
||||
#pragma unroll
|
||||
@@ -382,6 +398,11 @@ __global__ void moe_sum_vec_kernel(
|
||||
|
||||
#pragma unroll
|
||||
for (int k = 0; k < TOPK; ++k) {
|
||||
if constexpr (PAD_AWARE) {
|
||||
if (moe_sum_pad_aware_skip(tk_tok, expert_map, k * stride_tk_k)) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
vec_t packed = *reinterpret_cast<const vec_t*>(in_tok + k * stride_topk);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC; ++j) acc[j] += static_cast<float>(packed.val[j]);
|
||||
@@ -394,13 +415,16 @@ __global__ void moe_sum_vec_kernel(
|
||||
}
|
||||
}
|
||||
|
||||
// Runtime-topk variant of the above.
|
||||
template <typename scalar_t>
|
||||
// Runtime-topk variant of the above, for topk values outside the templated
|
||||
// set.
|
||||
template <typename scalar_t, typename idx_t, bool PAD_AWARE>
|
||||
__global__ void moe_sum_vec_dynamic_kernel(
|
||||
scalar_t* __restrict__ out, // [num_tokens, d], contiguous
|
||||
const scalar_t* __restrict__ input, // [num_tokens, topk, d], d contiguous
|
||||
const int64_t num_tokens, const int d, const int topk,
|
||||
const int64_t stride_token, const int64_t stride_topk) {
|
||||
const int64_t stride_token, const int64_t stride_topk,
|
||||
const idx_t* __restrict__ topk_ids, const int32_t* __restrict__ expert_map,
|
||||
const int64_t stride_tk_token, const int64_t stride_tk_k) {
|
||||
using vec_t = vllm::vec_n_t<scalar_t, MOE_SUM_VEC<scalar_t>>;
|
||||
constexpr int VEC = MOE_SUM_VEC<scalar_t>;
|
||||
const int64_t n_vec = d / VEC;
|
||||
@@ -410,12 +434,21 @@ __global__ void moe_sum_vec_dynamic_kernel(
|
||||
const int64_t token = i / n_vec;
|
||||
const int64_t v = i % n_vec;
|
||||
const scalar_t* in_tok = input + token * stride_token + v * VEC;
|
||||
const idx_t* tk_tok = nullptr;
|
||||
if constexpr (PAD_AWARE) {
|
||||
tk_tok = topk_ids + token * stride_tk_token;
|
||||
}
|
||||
|
||||
float acc[VEC];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC; ++j) acc[j] = 0.f;
|
||||
|
||||
for (int k = 0; k < topk; ++k) {
|
||||
if constexpr (PAD_AWARE) {
|
||||
if (moe_sum_pad_aware_skip(tk_tok, expert_map, k * stride_tk_k)) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
vec_t packed = *reinterpret_cast<const vec_t*>(in_tok + k * stride_topk);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < VEC; ++j) acc[j] += static_cast<float>(packed.val[j]);
|
||||
@@ -430,17 +463,28 @@ __global__ void moe_sum_vec_dynamic_kernel(
|
||||
|
||||
// Stride-aware scalar fallback: handles unaligned/non-vectorizable hidden dims
|
||||
// (including a non-contiguous hidden stride) via per-element strided reads.
|
||||
template <typename scalar_t>
|
||||
template <typename scalar_t, typename idx_t, bool PAD_AWARE>
|
||||
__global__ void moe_sum_scalar_kernel(
|
||||
scalar_t* __restrict__ out, // [num_tokens, d], contiguous
|
||||
const scalar_t* __restrict__ input, // [num_tokens, topk, d]
|
||||
const int d, const int topk, const int64_t stride_token,
|
||||
const int64_t stride_topk, const int64_t stride_hidden) {
|
||||
const int64_t stride_topk, const int64_t stride_hidden,
|
||||
const idx_t* __restrict__ topk_ids, const int32_t* __restrict__ expert_map,
|
||||
const int64_t stride_tk_token, const int64_t stride_tk_k) {
|
||||
const int64_t token_idx = blockIdx.x;
|
||||
const scalar_t* in_tok = input + token_idx * stride_token;
|
||||
const idx_t* tk_tok = nullptr;
|
||||
if constexpr (PAD_AWARE) {
|
||||
tk_tok = topk_ids + token_idx * stride_tk_token;
|
||||
}
|
||||
for (int64_t idx = threadIdx.x; idx < d; idx += blockDim.x) {
|
||||
float x = 0.f;
|
||||
for (int k = 0; k < topk; ++k) {
|
||||
if constexpr (PAD_AWARE) {
|
||||
if (moe_sum_pad_aware_skip(tk_tok, expert_map, k * stride_tk_k)) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
x += static_cast<float>(
|
||||
VLLM_LDG(&in_tok[k * stride_topk + idx * stride_hidden]));
|
||||
}
|
||||
@@ -711,8 +755,9 @@ void batched_moe_align_block_size(int64_t max_tokens_per_batch,
|
||||
}
|
||||
|
||||
void moe_sum(torch::stable::Tensor& input, // [num_tokens, topk, hidden_size]
|
||||
torch::stable::Tensor& output) // [num_tokens, hidden_size]
|
||||
{
|
||||
torch::stable::Tensor& output, // [num_tokens, hidden_size]
|
||||
std::optional<torch::stable::Tensor> topk_ids,
|
||||
std::optional<torch::stable::Tensor> expert_map) {
|
||||
// Output is dense and written in place, so it must be contiguous. The input
|
||||
// is read by its strides (no copy); only the hidden dim needs to be
|
||||
// contiguous to take the vectorized path.
|
||||
@@ -731,10 +776,102 @@ void moe_sum(torch::stable::Tensor& input, // [num_tokens, topk, hidden_size]
|
||||
const cudaStream_t stream =
|
||||
get_current_cuda_stream(output.get_device_index());
|
||||
|
||||
#define LAUNCH_MOE_SUM_VEC(TOPK) \
|
||||
vllm::moe::moe_sum_vec_kernel<scalar_t, TOPK> \
|
||||
<<<grid, dim3(block), 0, stream>>>( \
|
||||
out_ptr, in_ptr, num_tokens, hidden_size, stride_token, stride_topk)
|
||||
if (topk_ids.has_value()) {
|
||||
// Pad-aware reduce path
|
||||
const torch::stable::Tensor& tk = topk_ids.value();
|
||||
STD_TORCH_CHECK(tk.size(0) == num_tokens && tk.size(1) == topk,
|
||||
"moe_sum: topk_ids must have shape [num_tokens, topk]");
|
||||
const int64_t stride_tk_token = tk.stride(0);
|
||||
const int64_t stride_tk_k = tk.stride(1);
|
||||
|
||||
const int32_t* expert_map_ptr = nullptr;
|
||||
if (expert_map.has_value()) {
|
||||
STD_TORCH_CHECK(
|
||||
expert_map->scalar_type() == torch::headeronly::ScalarType::Int,
|
||||
"moe_sum: expert_map must be int32");
|
||||
expert_map_ptr =
|
||||
reinterpret_cast<const int32_t*>(expert_map->const_data_ptr());
|
||||
}
|
||||
|
||||
#define LAUNCH_MOE_SUM_PAD_AWARE_VEC(TOPK) \
|
||||
vllm::moe::moe_sum_vec_kernel<scalar_t, idx_t, TOPK, true> \
|
||||
<<<grid, dim3(block), 0, stream>>>( \
|
||||
out_ptr, in_ptr, num_tokens, hidden_size, stride_token, stride_topk, \
|
||||
topk_ids_ptr, expert_map_ptr, stride_tk_token, stride_tk_k)
|
||||
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
|
||||
input.scalar_type(), "moe_sum_pad_aware", [&] {
|
||||
constexpr int VEC = vllm::moe::MOE_SUM_VEC<scalar_t>;
|
||||
constexpr int WIDTH = VEC * sizeof(scalar_t);
|
||||
auto* out_ptr =
|
||||
reinterpret_cast<scalar_t*>(output.mutable_data_ptr());
|
||||
auto* in_ptr =
|
||||
reinterpret_cast<const scalar_t*>(input.const_data_ptr());
|
||||
|
||||
const bool can_vec =
|
||||
(stride_hidden == 1) && (hidden_size % VEC == 0) &&
|
||||
(stride_token % VEC == 0) && (stride_topk % VEC == 0) &&
|
||||
(reinterpret_cast<uintptr_t>(in_ptr) % WIDTH == 0) &&
|
||||
(reinterpret_cast<uintptr_t>(out_ptr) % WIDTH == 0);
|
||||
|
||||
VLLM_STABLE_DISPATCH_IDX_TYPES(
|
||||
tk.scalar_type(), "moe_sum_pad_aware_idx", [&] {
|
||||
auto* topk_ids_ptr =
|
||||
reinterpret_cast<const idx_t*>(tk.const_data_ptr());
|
||||
if (can_vec) {
|
||||
const int64_t n_vec = hidden_size / VEC;
|
||||
const int64_t total = num_tokens * n_vec;
|
||||
const int block = 256;
|
||||
const dim3 grid(
|
||||
std::min<int64_t>((total + block - 1) / block, 65535));
|
||||
switch (topk) {
|
||||
case 1:
|
||||
LAUNCH_MOE_SUM_PAD_AWARE_VEC(1);
|
||||
break;
|
||||
case 2:
|
||||
LAUNCH_MOE_SUM_PAD_AWARE_VEC(2);
|
||||
break;
|
||||
case 4:
|
||||
LAUNCH_MOE_SUM_PAD_AWARE_VEC(4);
|
||||
break;
|
||||
case 6:
|
||||
LAUNCH_MOE_SUM_PAD_AWARE_VEC(6);
|
||||
break;
|
||||
case 8:
|
||||
LAUNCH_MOE_SUM_PAD_AWARE_VEC(8);
|
||||
break;
|
||||
case 9:
|
||||
LAUNCH_MOE_SUM_PAD_AWARE_VEC(9);
|
||||
break;
|
||||
default:
|
||||
vllm::moe::moe_sum_vec_dynamic_kernel<scalar_t, idx_t,
|
||||
true>
|
||||
<<<grid, dim3(block), 0, stream>>>(
|
||||
out_ptr, in_ptr, num_tokens, hidden_size, topk,
|
||||
stride_token, stride_topk, topk_ids_ptr,
|
||||
expert_map_ptr, stride_tk_token, stride_tk_k);
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
dim3 grid(num_tokens);
|
||||
dim3 block(std::min(hidden_size, 1024));
|
||||
vllm::moe::moe_sum_scalar_kernel<scalar_t, idx_t, true>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
out_ptr, in_ptr, hidden_size, topk, stride_token,
|
||||
stride_topk, stride_hidden, topk_ids_ptr,
|
||||
expert_map_ptr, stride_tk_token, stride_tk_k);
|
||||
}
|
||||
});
|
||||
});
|
||||
#undef LAUNCH_MOE_SUM_PAD_AWARE_VEC
|
||||
return;
|
||||
}
|
||||
|
||||
#define LAUNCH_MOE_SUM_VEC(TOPK) \
|
||||
vllm::moe::moe_sum_vec_kernel<scalar_t, int32_t, TOPK, false> \
|
||||
<<<grid, dim3(block), 0, stream>>>(out_ptr, in_ptr, num_tokens, \
|
||||
hidden_size, stride_token, \
|
||||
stride_topk, nullptr, nullptr, 0, 0)
|
||||
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES(input.scalar_type(), "moe_sum", [&] {
|
||||
constexpr int VEC = vllm::moe::MOE_SUM_VEC<scalar_t>;
|
||||
@@ -774,18 +911,19 @@ void moe_sum(torch::stable::Tensor& input, // [num_tokens, topk, hidden_size]
|
||||
LAUNCH_MOE_SUM_VEC(9);
|
||||
break;
|
||||
default:
|
||||
vllm::moe::moe_sum_vec_dynamic_kernel<scalar_t>
|
||||
<<<grid, dim3(block), 0, stream>>>(out_ptr, in_ptr, num_tokens,
|
||||
hidden_size, topk,
|
||||
stride_token, stride_topk);
|
||||
vllm::moe::moe_sum_vec_dynamic_kernel<scalar_t, int32_t, false>
|
||||
<<<grid, dim3(block), 0, stream>>>(
|
||||
out_ptr, in_ptr, num_tokens, hidden_size, topk, stride_token,
|
||||
stride_topk, nullptr, nullptr, 0, 0);
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
dim3 grid(num_tokens);
|
||||
dim3 block(std::min(hidden_size, 1024));
|
||||
vllm::moe::moe_sum_scalar_kernel<scalar_t><<<grid, block, 0, stream>>>(
|
||||
out_ptr, in_ptr, hidden_size, topk, stride_token, stride_topk,
|
||||
stride_hidden);
|
||||
vllm::moe::moe_sum_scalar_kernel<scalar_t, int32_t, false>
|
||||
<<<grid, block, 0, stream>>>(out_ptr, in_ptr, hidden_size, topk,
|
||||
stride_token, stride_topk, stride_hidden,
|
||||
nullptr, nullptr, 0, 0);
|
||||
}
|
||||
});
|
||||
#undef LAUNCH_MOE_SUM_VEC
|
||||
@@ -948,4 +1086,4 @@ void moe_lora_align_block_size(
|
||||
has_expert_map);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,14 +9,16 @@ void topk_softmax(torch::stable::Tensor& topk_weights,
|
||||
torch::stable::Tensor& topk_indices,
|
||||
torch::stable::Tensor& token_expert_indices,
|
||||
torch::stable::Tensor& gating_output, bool renormalize,
|
||||
std::optional<torch::stable::Tensor> bias);
|
||||
std::optional<torch::stable::Tensor> bias,
|
||||
std::optional<torch::stable::Tensor> is_padding);
|
||||
|
||||
void topk_sigmoid(torch::stable::Tensor& topk_weights,
|
||||
torch::stable::Tensor& topk_indices,
|
||||
torch::stable::Tensor& token_expert_indices,
|
||||
torch::stable::Tensor& gating_output, bool renormalize,
|
||||
std::optional<torch::stable::Tensor> bias,
|
||||
double routed_scaling_factor);
|
||||
double routed_scaling_factor,
|
||||
std::optional<torch::stable::Tensor> is_padding);
|
||||
|
||||
void topk_softplus_sqrt(
|
||||
torch::stable::Tensor& topk_weights, torch::stable::Tensor& topk_indices,
|
||||
@@ -25,9 +27,12 @@ void topk_softplus_sqrt(
|
||||
double routed_scaling_factor,
|
||||
const std::optional<torch::stable::Tensor>& correction_bias,
|
||||
const std::optional<torch::stable::Tensor>& input_ids,
|
||||
const std::optional<torch::stable::Tensor>& tid2eid);
|
||||
const std::optional<torch::stable::Tensor>& tid2eid,
|
||||
const std::optional<torch::stable::Tensor>& is_padding);
|
||||
|
||||
void moe_sum(torch::stable::Tensor& input, torch::stable::Tensor& output);
|
||||
void moe_sum(torch::stable::Tensor& input, torch::stable::Tensor& output,
|
||||
std::optional<torch::stable::Tensor> topk_ids,
|
||||
std::optional<torch::stable::Tensor> expert_map);
|
||||
|
||||
void moe_align_block_size(
|
||||
torch::stable::Tensor topk_ids, int64_t num_experts, int64_t block_size,
|
||||
|
||||
@@ -174,7 +174,8 @@ __launch_bounds__(TPB) __global__ void moeTopK(
|
||||
const int end_expert,
|
||||
const bool renormalize,
|
||||
const float* bias,
|
||||
const double routed_scaling_factor)
|
||||
const double routed_scaling_factor,
|
||||
const bool* is_padding)
|
||||
{
|
||||
|
||||
using cub_kvp = cub::KeyValuePair<int, float>;
|
||||
@@ -228,12 +229,14 @@ __launch_bounds__(TPB) __global__ void moeTopK(
|
||||
const int expert = result_kvp.key;
|
||||
const bool node_uses_expert = expert >= start_expert && expert < end_expert;
|
||||
const bool should_process_row = row_is_active && node_uses_expert;
|
||||
const bool is_pad_row = is_padding != nullptr && is_padding[block_row];
|
||||
|
||||
const int idx = k * block_row + k_idx;
|
||||
// Return the unbiased scores for output weights
|
||||
output[idx] = inputs_after_softmax[thread_read_offset + expert];
|
||||
indices[idx] = should_process_row ? (expert - start_expert) : num_experts;
|
||||
assert(indices[idx] >= 0);
|
||||
indices[idx] = is_pad_row ? static_cast<IndType>(-1)
|
||||
: (should_process_row ? (expert - start_expert) : num_experts);
|
||||
assert(is_pad_row || indices[idx] >= 0);
|
||||
source_rows[idx] = k_idx * num_rows + block_row;
|
||||
if (renormalize) {
|
||||
selected_sum += inputs_after_softmax[thread_read_offset + expert];
|
||||
@@ -277,7 +280,7 @@ template <int VPT, int NUM_EXPERTS, int WARPS_PER_CTA, int BYTES_PER_LDG, int WA
|
||||
__launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
void topkGating(const InputType* input, const bool* finished, float* output, const int num_rows, IndType* indices,
|
||||
int* source_rows, const int k, const int start_expert, const int end_expert, const bool renormalize,
|
||||
const float* bias, const double routed_scaling_factor)
|
||||
const float* bias, const double routed_scaling_factor, const bool* is_padding)
|
||||
{
|
||||
static_assert(std::is_same_v<InputType, float> || std::is_same_v<InputType, __nv_bfloat16> ||
|
||||
std::is_same_v<InputType, __half>,
|
||||
@@ -545,12 +548,14 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
// Add a guard to ignore experts not included by this node
|
||||
const bool node_uses_expert = expert >= start_expert && expert < end_expert;
|
||||
const bool should_process_row = row_is_active && node_uses_expert;
|
||||
const bool is_pad_row = is_padding != nullptr && is_padding[thread_row];
|
||||
|
||||
// The lead thread from each sub-group will write out the final results to global memory. (This will be a
|
||||
// single) thread per row of the input/output matrices.
|
||||
const int idx = k * thread_row + k_idx;
|
||||
output[idx] = max_val;
|
||||
indices[idx] = should_process_row ? (expert - start_expert) : NUM_EXPERTS;
|
||||
indices[idx] = is_pad_row ? static_cast<IndType>(-1)
|
||||
: (should_process_row ? (expert - start_expert) : NUM_EXPERTS);
|
||||
source_rows[idx] = k_idx * num_rows + thread_row;
|
||||
if (renormalize) {
|
||||
selected_sum += max_val;
|
||||
@@ -605,7 +610,7 @@ struct TopkConstants
|
||||
template <int EXPERTS, int WARPS_PER_TB, int WARP_SIZE_PARAM, int MAX_BYTES_PER_LDG, typename IndType, typename InputType, ScoringFunc SF>
|
||||
void topkGatingLauncherHelper(const InputType* input, const bool* finished, float* output, IndType* indices,
|
||||
int* source_row, const int num_rows, const int k, const int start_expert, const int end_expert, const bool renormalize,
|
||||
const float* bias, const double routed_scaling_factor, cudaStream_t stream)
|
||||
const float* bias, const double routed_scaling_factor, cudaStream_t stream, const bool* is_padding)
|
||||
{
|
||||
static constexpr int BYTES_PER_LDG = MIN(MAX_BYTES_PER_LDG, sizeof(InputType) * EXPERTS);
|
||||
using Constants = detail::TopkConstants<EXPERTS, BYTES_PER_LDG, WARP_SIZE_PARAM, InputType>;
|
||||
@@ -616,7 +621,7 @@ void topkGatingLauncherHelper(const InputType* input, const bool* finished, floa
|
||||
|
||||
dim3 block_dim(WARP_SIZE_PARAM, WARPS_PER_TB);
|
||||
topkGating<VPT, EXPERTS, WARPS_PER_TB, BYTES_PER_LDG, WARP_SIZE_PARAM, IndType, InputType, SF><<<num_blocks, block_dim, 0, stream>>>(
|
||||
input, finished, output, num_rows, indices, source_row, k, start_expert, end_expert, renormalize, bias, routed_scaling_factor);
|
||||
input, finished, output, num_rows, indices, source_row, k, start_expert, end_expert, renormalize, bias, routed_scaling_factor, is_padding);
|
||||
}
|
||||
|
||||
#ifndef USE_ROCM
|
||||
@@ -627,7 +632,7 @@ void topkGatingLauncherHelper(const InputType* input, const bool* finished, floa
|
||||
IndType, InputType, SF>( \
|
||||
gating_output, nullptr, topk_weights, topk_indices, \
|
||||
token_expert_indices, num_tokens, topk, 0, num_experts, renormalize, \
|
||||
bias, routed_scaling_factor, stream);
|
||||
bias, routed_scaling_factor, stream, is_padding);
|
||||
#else
|
||||
#define LAUNCH_TOPK(NUM_EXPERTS, WARPS_PER_TB, MAX_BYTES) \
|
||||
if (WARP_SIZE == 64) { \
|
||||
@@ -635,13 +640,13 @@ void topkGatingLauncherHelper(const InputType* input, const bool* finished, floa
|
||||
IndType, InputType, SF>( \
|
||||
gating_output, nullptr, topk_weights, topk_indices, \
|
||||
token_expert_indices, num_tokens, topk, 0, num_experts, renormalize, \
|
||||
bias, routed_scaling_factor, stream); \
|
||||
bias, routed_scaling_factor, stream, is_padding); \
|
||||
} else if (WARP_SIZE == 32) { \
|
||||
topkGatingLauncherHelper<NUM_EXPERTS, WARPS_PER_TB, 32, MAX_BYTES, \
|
||||
IndType, InputType, SF>( \
|
||||
gating_output, nullptr, topk_weights, topk_indices, \
|
||||
token_expert_indices, num_tokens, topk, 0, num_experts, renormalize, \
|
||||
bias, routed_scaling_factor, stream); \
|
||||
bias, routed_scaling_factor, stream, is_padding); \
|
||||
} else { \
|
||||
assert(false && \
|
||||
"Unsupported warp size. Only 32 and 64 are supported for ROCm"); \
|
||||
@@ -661,7 +666,8 @@ void topkGatingKernelLauncher(
|
||||
const bool renormalize,
|
||||
const float* bias,
|
||||
const double routed_scaling_factor,
|
||||
cudaStream_t stream) {
|
||||
cudaStream_t stream,
|
||||
const bool* is_padding) {
|
||||
static constexpr int WARPS_PER_TB = 4;
|
||||
static constexpr int BYTES_PER_LDG_POWER_OF_2 = 16;
|
||||
#ifndef USE_ROCM
|
||||
@@ -736,7 +742,7 @@ void topkGatingKernelLauncher(
|
||||
}
|
||||
moeTopK<TPB><<<num_tokens, TPB, 0, stream>>>(
|
||||
workspace, nullptr, topk_weights, topk_indices, token_expert_indices,
|
||||
num_experts, topk, 0, num_experts, renormalize, bias, routed_scaling_factor);
|
||||
num_experts, topk, 0, num_experts, renormalize, bias, routed_scaling_factor, is_padding);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -755,7 +761,8 @@ void dispatch_topk_launch(
|
||||
int num_tokens, int num_experts, int topk, bool renormalize,
|
||||
std::optional<torch::stable::Tensor> bias,
|
||||
double routed_scaling_factor,
|
||||
cudaStream_t stream)
|
||||
cudaStream_t stream,
|
||||
std::optional<torch::stable::Tensor> is_padding)
|
||||
{
|
||||
const float* bias_ptr = nullptr;
|
||||
if (bias.has_value()) {
|
||||
@@ -769,6 +776,18 @@ void dispatch_topk_launch(
|
||||
bias_ptr = bias_tensor.const_data_ptr<float>();
|
||||
}
|
||||
|
||||
const bool* is_padding_ptr = nullptr;
|
||||
if (is_padding.has_value()) {
|
||||
const torch::stable::Tensor& is_padding_tensor = is_padding.value();
|
||||
STD_TORCH_CHECK(is_padding_tensor.scalar_type() == torch::headeronly::ScalarType::Bool,
|
||||
"is_padding tensor must be bool");
|
||||
STD_TORCH_CHECK(is_padding_tensor.dim() == 1, "is_padding tensor must be 1D");
|
||||
STD_TORCH_CHECK(is_padding_tensor.size(0) == num_tokens,
|
||||
"is_padding size mismatch, expected: ", num_tokens);
|
||||
STD_TORCH_CHECK(is_padding_tensor.is_contiguous(), "is_padding tensor must be contiguous");
|
||||
is_padding_ptr = is_padding_tensor.const_data_ptr<bool>();
|
||||
}
|
||||
|
||||
if (topk_indices.scalar_type() == torch::headeronly::ScalarType::Int) {
|
||||
vllm::moe::topkGatingKernelLauncher<int, ComputeType, SF>(
|
||||
reinterpret_cast<const ComputeType*>(gating_output.const_data_ptr()),
|
||||
@@ -777,7 +796,7 @@ void dispatch_topk_launch(
|
||||
token_expert_indices.mutable_data_ptr<int>(),
|
||||
softmax_workspace.mutable_data_ptr<float>(),
|
||||
num_tokens, num_experts, topk, renormalize,
|
||||
bias_ptr, routed_scaling_factor, stream);
|
||||
bias_ptr, routed_scaling_factor, stream, is_padding_ptr);
|
||||
} else if (topk_indices.scalar_type() == torch::headeronly::ScalarType::UInt32) {
|
||||
vllm::moe::topkGatingKernelLauncher<uint32_t, ComputeType, SF>(
|
||||
reinterpret_cast<const ComputeType*>(gating_output.const_data_ptr()),
|
||||
@@ -786,7 +805,7 @@ void dispatch_topk_launch(
|
||||
token_expert_indices.mutable_data_ptr<int>(),
|
||||
softmax_workspace.mutable_data_ptr<float>(),
|
||||
num_tokens, num_experts, topk, renormalize,
|
||||
bias_ptr, routed_scaling_factor, stream);
|
||||
bias_ptr, routed_scaling_factor, stream, is_padding_ptr);
|
||||
} else {
|
||||
STD_TORCH_CHECK(topk_indices.scalar_type() == torch::headeronly::ScalarType::Long);
|
||||
vllm::moe::topkGatingKernelLauncher<int64_t, ComputeType, SF>(
|
||||
@@ -796,7 +815,7 @@ void dispatch_topk_launch(
|
||||
token_expert_indices.mutable_data_ptr<int>(),
|
||||
softmax_workspace.mutable_data_ptr<float>(),
|
||||
num_tokens, num_experts, topk, renormalize,
|
||||
bias_ptr, routed_scaling_factor, stream);
|
||||
bias_ptr, routed_scaling_factor, stream, is_padding_ptr);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -806,7 +825,8 @@ void topk_softmax(
|
||||
torch::stable::Tensor& token_expert_indices, // [num_tokens, topk]
|
||||
torch::stable::Tensor& gating_output, // [num_tokens, num_experts]
|
||||
bool renormalize,
|
||||
std::optional<torch::stable::Tensor> bias)
|
||||
std::optional<torch::stable::Tensor> bias,
|
||||
std::optional<torch::stable::Tensor> is_padding)
|
||||
{
|
||||
const int num_experts = gating_output.size(-1);
|
||||
const auto num_tokens = gating_output.numel() / num_experts;
|
||||
@@ -825,15 +845,15 @@ void topk_softmax(
|
||||
if (gating_output.scalar_type() == torch::headeronly::ScalarType::Float) {
|
||||
dispatch_topk_launch<float, vllm::moe::SCORING_SOFTMAX>(gating_output, topk_weights, topk_indices,
|
||||
token_expert_indices, softmax_workspace, num_tokens, num_experts, topk, renormalize,
|
||||
bias, 1.0, stream);
|
||||
bias, 1.0, stream, is_padding);
|
||||
} else if (gating_output.scalar_type() == torch::headeronly::ScalarType::Half) {
|
||||
dispatch_topk_launch<__half, vllm::moe::SCORING_SOFTMAX>(gating_output, topk_weights, topk_indices,
|
||||
token_expert_indices, softmax_workspace, num_tokens, num_experts, topk, renormalize,
|
||||
bias, 1.0, stream);
|
||||
bias, 1.0, stream, is_padding);
|
||||
} else if (gating_output.scalar_type() == torch::headeronly::ScalarType::BFloat16) {
|
||||
dispatch_topk_launch<__nv_bfloat16, vllm::moe::SCORING_SOFTMAX>(gating_output, topk_weights, topk_indices,
|
||||
token_expert_indices, softmax_workspace, num_tokens, num_experts, topk, renormalize,
|
||||
bias, 1.0, stream);
|
||||
bias, 1.0, stream, is_padding);
|
||||
} else {
|
||||
STD_TORCH_CHECK(false, "Unsupported gating_output data type: ", gating_output.scalar_type());
|
||||
}
|
||||
@@ -846,7 +866,8 @@ void topk_sigmoid(
|
||||
torch::stable::Tensor& gating_output, // [num_tokens, num_experts]
|
||||
bool renormalize,
|
||||
std::optional<torch::stable::Tensor> bias,
|
||||
double routed_scaling_factor)
|
||||
double routed_scaling_factor,
|
||||
std::optional<torch::stable::Tensor> is_padding)
|
||||
{
|
||||
const int num_experts = gating_output.size(-1);
|
||||
const auto num_tokens = gating_output.numel() / num_experts;
|
||||
@@ -865,15 +886,15 @@ void topk_sigmoid(
|
||||
if (gating_output.scalar_type() == torch::headeronly::ScalarType::Float) {
|
||||
dispatch_topk_launch<float, vllm::moe::SCORING_SIGMOID>(gating_output, topk_weights, topk_indices,
|
||||
token_expert_indices, workspace, num_tokens, num_experts, topk, renormalize,
|
||||
bias, routed_scaling_factor, stream);
|
||||
bias, routed_scaling_factor, stream, is_padding);
|
||||
} else if (gating_output.scalar_type() == torch::headeronly::ScalarType::Half) {
|
||||
dispatch_topk_launch<__half, vllm::moe::SCORING_SIGMOID>(gating_output, topk_weights, topk_indices,
|
||||
token_expert_indices, workspace, num_tokens, num_experts, topk, renormalize,
|
||||
bias, routed_scaling_factor, stream);
|
||||
bias, routed_scaling_factor, stream, is_padding);
|
||||
} else if (gating_output.scalar_type() == torch::headeronly::ScalarType::BFloat16) {
|
||||
dispatch_topk_launch<__nv_bfloat16, vllm::moe::SCORING_SIGMOID>(gating_output, topk_weights, topk_indices,
|
||||
token_expert_indices, workspace, num_tokens, num_experts, topk, renormalize,
|
||||
bias, routed_scaling_factor, stream);
|
||||
bias, routed_scaling_factor, stream, is_padding);
|
||||
} else {
|
||||
STD_TORCH_CHECK(false, "Unsupported gating_output data type: ", gating_output.scalar_type());
|
||||
}
|
||||
|
||||
@@ -44,6 +44,12 @@ typedef __hip_bfloat162 __nv_bfloat162;
|
||||
namespace vllm {
|
||||
namespace moe {
|
||||
|
||||
template <typename HashIndType>
|
||||
__device__ __forceinline__ int64_t load_index_as_int64(const HashIndType* ptr,
|
||||
int64_t offset) {
|
||||
return static_cast<int64_t>(ptr[offset]);
|
||||
}
|
||||
|
||||
/// Aligned array type
|
||||
template <typename T,
|
||||
/// Number of elements in the array
|
||||
@@ -65,6 +71,80 @@ __device__ __forceinline__ float toFloat(T value) {
|
||||
}
|
||||
}
|
||||
|
||||
#ifndef USE_ROCM
|
||||
// Adapted from:
|
||||
// https://github.com/sgl-project/sglang/blob/main/python/sglang/jit_kernel/csrc/deepseek_v4/hash_topk.cuh
|
||||
template <typename OutIndType, typename HashIndType>
|
||||
__launch_bounds__(128) __global__
|
||||
void dsv4HashTopkSoftplusSqrt(const float* input, float* output,
|
||||
OutIndType* indices, int num_rows,
|
||||
int num_experts, float routed_scaling_factor,
|
||||
const HashIndType* input_ids,
|
||||
const HashIndType* tid2eid,
|
||||
const bool* is_padding) {
|
||||
const int warp = (blockIdx.x * blockDim.x + threadIdx.x) / 32;
|
||||
const int lane = threadIdx.x % 32;
|
||||
if (warp >= num_rows) return;
|
||||
const int64_t token_id = load_index_as_int64(input_ids, warp);
|
||||
const bool is_pad_row = is_padding != nullptr && is_padding[warp];
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
cudaGridDependencySynchronize();
|
||||
#endif
|
||||
int expert = 0;
|
||||
float weight = 0.f;
|
||||
if (lane < 6 && !is_pad_row) {
|
||||
// only load and calculate for 6 experts
|
||||
expert = static_cast<int>(tid2eid[token_id * 6 + lane]);
|
||||
const float x = input[warp * num_experts + expert];
|
||||
weight = sqrtf(fmaxf(x, 0.f) + __logf(1.f + __expf(-fabsf(x))));
|
||||
if (isnan(weight)) {
|
||||
weight = 0.f;
|
||||
}
|
||||
}
|
||||
float weight_sum = weight;
|
||||
#pragma unroll
|
||||
for (int mask = 16; mask > 0; mask >>= 1) {
|
||||
// sum in warp
|
||||
weight_sum += VLLM_SHFL_XOR_SYNC(weight_sum, mask);
|
||||
}
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
#endif
|
||||
if (lane < 6) {
|
||||
const int offset = warp * 6 + lane;
|
||||
output[offset] =
|
||||
weight * routed_scaling_factor / (weight_sum > 0.f ? weight_sum : 1.f);
|
||||
indices[offset] = !is_pad_row ? static_cast<OutIndType>(expert)
|
||||
: static_cast<OutIndType>(-1);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename OutIndType, typename HashIndType>
|
||||
void launchDsv4HashTopk(const float* input, float* output, OutIndType* indices,
|
||||
int num_rows, int num_experts,
|
||||
double routed_scaling_factor,
|
||||
const HashIndType* input_ids,
|
||||
const HashIndType* tid2eid, cudaStream_t stream,
|
||||
const bool* is_padding) {
|
||||
if (num_rows == 0) return;
|
||||
auto* kernel = &dsv4HashTopkSoftplusSqrt<OutIndType, HashIndType>;
|
||||
cudaLaunchConfig_t config = {};
|
||||
config.gridDim = (num_rows + 3) / 4;
|
||||
config.blockDim = 128;
|
||||
config.stream = stream;
|
||||
cudaLaunchAttribute attr;
|
||||
attr.id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attr.val.programmaticStreamSerializationAllowed = 1;
|
||||
config.attrs = &attr;
|
||||
config.numAttrs = 1;
|
||||
const float scale = static_cast<float>(routed_scaling_factor);
|
||||
cudaLaunchKernelEx(&config, kernel, input, output, indices, num_rows,
|
||||
num_experts, scale, input_ids, tid2eid, is_padding);
|
||||
}
|
||||
#endif
|
||||
|
||||
// ====================== TopK softplus_sqrt things
|
||||
// ===============================
|
||||
|
||||
@@ -86,14 +166,15 @@ __device__ __forceinline__ float toFloat(T value) {
|
||||
|
||||
template <int VPT, int NUM_EXPERTS, int WARPS_PER_CTA, int BYTES_PER_LDG,
|
||||
int WARP_SIZE_PARAM, bool USE_HASH, typename IndType,
|
||||
typename InputType = float>
|
||||
typename HashIndType, typename InputType = float>
|
||||
__launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
void topkGatingSoftplusSqrt(
|
||||
const InputType* input, const bool* finished, float* output,
|
||||
const int num_rows, IndType* indices, int* source_rows, const int k,
|
||||
const int start_expert, const int end_expert, const bool renormalize,
|
||||
double routed_scaling_factor, const float* correction_bias,
|
||||
const IndType* input_ids, const IndType* tid2eid) {
|
||||
const HashIndType* input_ids, const HashIndType* tid2eid,
|
||||
const bool* is_padding) {
|
||||
static_assert(std::is_same_v<InputType, float> ||
|
||||
std::is_same_v<InputType, __nv_bfloat16> ||
|
||||
std::is_same_v<InputType, __half>,
|
||||
@@ -158,6 +239,7 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
return;
|
||||
}
|
||||
const bool row_is_active = finished ? !finished[thread_row] : true;
|
||||
const bool is_pad_row = is_padding != nullptr && is_padding[thread_row];
|
||||
|
||||
// We finally start setting up the read pointers for each thread. First, each
|
||||
// thread jumps to the start of the row it will read.
|
||||
@@ -176,9 +258,12 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
cudaGridDependencySynchronize();
|
||||
#endif
|
||||
|
||||
// NOTE(zhuhaoran): dispatch different input types loading, BF16/FP16 convert
|
||||
// to float
|
||||
if constexpr (std::is_same_v<InputType, float>) {
|
||||
if (is_pad_row) {
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
row_chunk[ii] = 0.f;
|
||||
}
|
||||
} else if constexpr (std::is_same_v<InputType, float>) {
|
||||
using VecType = AlignedArray<float, ELTS_PER_LDG>;
|
||||
VecType* row_chunk_vec_ptr = reinterpret_cast<VecType*>(&row_chunk);
|
||||
const VecType* vec_thread_read_ptr =
|
||||
@@ -240,19 +325,30 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
|
||||
// Hash MoE path: indices are predetermined from lookup table
|
||||
if constexpr (USE_HASH) {
|
||||
const IndType token_id = input_ids[thread_row];
|
||||
const IndType* expert_indices_for_token = tid2eid + token_id * k;
|
||||
const int64_t token_id = load_index_as_int64(input_ids, thread_row);
|
||||
const int64_t token_expert_offset = token_id * static_cast<int64_t>(k);
|
||||
if (!is_pad_row) {
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
float val = row_chunk[ii];
|
||||
float val_b = val * beta;
|
||||
val = (val_b > threshold) ? val : (__logf(1.0f + __expf(val_b))) / beta;
|
||||
row_chunk[ii] = sqrtf(val);
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
float val = row_chunk[ii];
|
||||
float val_b = val * beta;
|
||||
val = (val_b > threshold) ? val : (__logf(1.0f + __expf(val_b))) / beta;
|
||||
val = sqrtf(val);
|
||||
|
||||
// Dummy/padding tokens can result in NaN values, so
|
||||
// clamp them to 0.0. Note: this clamp could likely be removed if
|
||||
// 'is_padding' is made mandatory
|
||||
if (isnan(val)) {
|
||||
val = 0.f;
|
||||
}
|
||||
row_chunk[ii] = val;
|
||||
}
|
||||
}
|
||||
float selected_sum = 0.f;
|
||||
#pragma unroll
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int expert = expert_indices_for_token[k_idx];
|
||||
const int expert = static_cast<int>(
|
||||
load_index_as_int64(tid2eid, token_expert_offset + k_idx));
|
||||
const int idx = k * thread_row + k_idx;
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
const int group_id = ii / ELTS_PER_LDG;
|
||||
@@ -261,7 +357,8 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
group_id * THREADS_PER_ROW * ELTS_PER_LDG +
|
||||
local_id;
|
||||
if (expert == expert_idx) {
|
||||
indices[idx] = expert;
|
||||
indices[idx] = !is_pad_row ? static_cast<IndType>(expert)
|
||||
: static_cast<IndType>(-1);
|
||||
selected_sum += row_chunk[ii];
|
||||
break;
|
||||
}
|
||||
@@ -285,7 +382,8 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
|
||||
#pragma unroll
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int expert = expert_indices_for_token[k_idx];
|
||||
const int expert = static_cast<int>(
|
||||
load_index_as_int64(tid2eid, token_expert_offset + k_idx));
|
||||
const int idx = k * thread_row + k_idx;
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
const int group_id = ii / ELTS_PER_LDG;
|
||||
@@ -304,23 +402,31 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
#endif
|
||||
return;
|
||||
} else {
|
||||
if (!is_pad_row) {
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
float val = row_chunk[ii];
|
||||
float val_b = val * beta;
|
||||
// Compute softplus: log(1 + exp(val)) with numerical stability
|
||||
// When val > threshold, softplus(x) ≈ x to avoid exp overflow
|
||||
val = (val_b > threshold) ? val : (__logf(1.0f + __expf(val_b))) / beta;
|
||||
val = sqrtf(val);
|
||||
if (correction_bias) {
|
||||
const int group_id = ii / ELTS_PER_LDG;
|
||||
const int local_id = ii % ELTS_PER_LDG;
|
||||
const int expert_idx = first_elt_read_by_thread +
|
||||
group_id * THREADS_PER_ROW * ELTS_PER_LDG +
|
||||
local_id;
|
||||
val = val + correction_bias[expert_idx];
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
float val = row_chunk[ii];
|
||||
float val_b = val * beta;
|
||||
// Compute softplus: log(1 + exp(val)) with numerical stability
|
||||
// When val > threshold, softplus(x) ≈ x to avoid exp overflow
|
||||
val = (val_b > threshold) ? val : (__logf(1.0f + __expf(val_b))) / beta;
|
||||
val = sqrtf(val);
|
||||
// Dummy/padding tokens can result in NaN values, so
|
||||
// clamp them to 0.0. Note: this clamp could likely be removed if
|
||||
// 'is_padding' is made mandatory
|
||||
if (isnan(val)) {
|
||||
val = 0.f;
|
||||
}
|
||||
if (correction_bias) {
|
||||
const int group_id = ii / ELTS_PER_LDG;
|
||||
const int local_id = ii % ELTS_PER_LDG;
|
||||
const int expert_idx = first_elt_read_by_thread +
|
||||
group_id * THREADS_PER_ROW * ELTS_PER_LDG +
|
||||
local_id;
|
||||
val = val + correction_bias[expert_idx];
|
||||
}
|
||||
row_chunk[ii] = val;
|
||||
}
|
||||
row_chunk[ii] = val;
|
||||
}
|
||||
|
||||
// Original TopK path: find top-k experts by score
|
||||
@@ -375,18 +481,19 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
|
||||
// Add a guard to ignore experts not included by this node
|
||||
const bool node_uses_expert =
|
||||
expert >= start_expert && expert < end_expert;
|
||||
const bool should_process_row = row_is_active && node_uses_expert;
|
||||
const bool should_process_row =
|
||||
row_is_active && node_uses_expert && !is_pad_row;
|
||||
|
||||
// The lead thread from each sub-group will write out the final results
|
||||
// to global memory. (This will be a single) thread per row of the
|
||||
// input/output matrices.
|
||||
const int idx = k * thread_row + k_idx;
|
||||
if (correction_bias != nullptr) {
|
||||
if (correction_bias != nullptr && should_process_row) {
|
||||
max_val -= correction_bias[expert];
|
||||
}
|
||||
output[idx] = max_val;
|
||||
indices[idx] =
|
||||
should_process_row ? (expert - start_expert) : NUM_EXPERTS;
|
||||
!is_pad_row ? expert - start_expert : static_cast<IndType>(-1);
|
||||
source_rows[idx] = k_idx * num_rows + thread_row;
|
||||
if (renormalize) {
|
||||
selected_sum += max_val;
|
||||
@@ -461,14 +568,15 @@ struct TopkConstants {
|
||||
}
|
||||
|
||||
template <int EXPERTS, int WARPS_PER_TB, int WARP_SIZE_PARAM,
|
||||
int MAX_BYTES_PER_LDG, typename IndType, typename InputType>
|
||||
int MAX_BYTES_PER_LDG, typename IndType, typename HashIndType,
|
||||
typename InputType>
|
||||
void topkGatingSoftplusSqrtLauncherHelper(
|
||||
const InputType* input, const bool* finished, float* output,
|
||||
IndType* indices, int* source_row, const int num_rows, const int k,
|
||||
const int start_expert, const int end_expert, const bool renormalize,
|
||||
double routed_scaling_factor, const float* correction_bias,
|
||||
const bool use_hash, const IndType* input_ids, const IndType* tid2eid,
|
||||
cudaStream_t stream) {
|
||||
const bool use_hash, const HashIndType* input_ids,
|
||||
const HashIndType* tid2eid, cudaStream_t stream, const bool* is_padding) {
|
||||
static constexpr int BYTES_PER_LDG =
|
||||
MIN(MAX_BYTES_PER_LDG, sizeof(InputType) * EXPERTS);
|
||||
using Constants =
|
||||
@@ -481,7 +589,8 @@ void topkGatingSoftplusSqrtLauncherHelper(
|
||||
DISPATCH_HASH(use_hash, USE_HASH, {
|
||||
auto* kernel =
|
||||
&topkGatingSoftplusSqrt<VPT, EXPERTS, WARPS_PER_TB, BYTES_PER_LDG,
|
||||
WARP_SIZE_PARAM, USE_HASH, IndType, InputType>;
|
||||
WARP_SIZE_PARAM, USE_HASH, IndType, HashIndType,
|
||||
InputType>;
|
||||
#ifndef USE_ROCM
|
||||
cudaLaunchConfig_t config = {};
|
||||
config.gridDim = num_blocks;
|
||||
@@ -496,12 +605,12 @@ void topkGatingSoftplusSqrtLauncherHelper(
|
||||
cudaLaunchKernelEx(&config, kernel, input, finished, output, num_rows,
|
||||
indices, source_row, k, start_expert, end_expert,
|
||||
renormalize, routed_scaling_factor, correction_bias,
|
||||
input_ids, tid2eid);
|
||||
input_ids, tid2eid, is_padding);
|
||||
#else
|
||||
kernel<<<num_blocks, block_dim, 0, stream>>>(
|
||||
input, finished, output, num_rows, indices, source_row, k, start_expert,
|
||||
end_expert, renormalize, routed_scaling_factor, correction_bias,
|
||||
input_ids, tid2eid);
|
||||
input_ids, tid2eid, is_padding);
|
||||
#endif
|
||||
})
|
||||
}
|
||||
@@ -515,7 +624,7 @@ void topkGatingSoftplusSqrtLauncherHelper(
|
||||
gating_output, nullptr, topk_weights, topk_indices, \
|
||||
token_expert_indices, num_tokens, topk, 0, num_experts, renormalize, \
|
||||
routed_scaling_factor, correction_bias, use_hash, input_ids, tid2eid, \
|
||||
stream);
|
||||
stream, is_padding);
|
||||
#else
|
||||
#define LAUNCH_SOFTPLUS_SQRT(NUM_EXPERTS, WARPS_PER_TB, MAX_BYTES) \
|
||||
if (WARP_SIZE == 64) { \
|
||||
@@ -524,27 +633,39 @@ void topkGatingSoftplusSqrtLauncherHelper(
|
||||
gating_output, nullptr, topk_weights, topk_indices, \
|
||||
token_expert_indices, num_tokens, topk, 0, num_experts, renormalize, \
|
||||
routed_scaling_factor, correction_bias, use_hash, input_ids, \
|
||||
tid2eid, stream); \
|
||||
tid2eid, stream, is_padding); \
|
||||
} else if (WARP_SIZE == 32) { \
|
||||
topkGatingSoftplusSqrtLauncherHelper<NUM_EXPERTS, WARPS_PER_TB, 32, \
|
||||
MAX_BYTES>( \
|
||||
gating_output, nullptr, topk_weights, topk_indices, \
|
||||
token_expert_indices, num_tokens, topk, 0, num_experts, renormalize, \
|
||||
routed_scaling_factor, correction_bias, use_hash, input_ids, \
|
||||
tid2eid, stream); \
|
||||
tid2eid, stream, is_padding); \
|
||||
} else { \
|
||||
assert(false && \
|
||||
"Unsupported warp size. Only 32 and 64 are supported for ROCm"); \
|
||||
}
|
||||
#endif
|
||||
|
||||
template <typename IndType, typename InputType>
|
||||
template <typename IndType, typename InputType, typename HashIndType = IndType>
|
||||
void topkGatingSoftplusSqrtKernelLauncher(
|
||||
const InputType* gating_output, float* topk_weights, IndType* topk_indices,
|
||||
int* token_expert_indices, const int num_tokens, const int num_experts,
|
||||
const int topk, const bool renormalize, double routed_scaling_factor,
|
||||
const float* correction_bias, const bool use_hash, const IndType* input_ids,
|
||||
const IndType* tid2eid, cudaStream_t stream) {
|
||||
const float* correction_bias, const bool use_hash,
|
||||
const HashIndType* input_ids, const HashIndType* tid2eid,
|
||||
cudaStream_t stream, const bool* is_padding) {
|
||||
#ifndef USE_ROCM
|
||||
if constexpr (std::is_same_v<InputType, float>) {
|
||||
if (use_hash && topk == 6 && renormalize &&
|
||||
(num_experts == 256 || num_experts == 384)) {
|
||||
launchDsv4HashTopk<IndType, HashIndType>(
|
||||
gating_output, topk_weights, topk_indices, num_tokens, num_experts,
|
||||
routed_scaling_factor, input_ids, tid2eid, stream, is_padding);
|
||||
return;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
static constexpr int WARPS_PER_TB = 4;
|
||||
static constexpr int BYTES_PER_LDG_POWER_OF_2 = 16;
|
||||
// for bfloat16 dtype, we need 4 bytes loading to make sure num_experts
|
||||
@@ -639,62 +760,77 @@ void dispatch_topk_softplus_sqrt_launch(
|
||||
int num_experts, int topk, bool renormalize, double routed_scaling_factor,
|
||||
const std::optional<torch::stable::Tensor>& correction_bias,
|
||||
const std::optional<torch::stable::Tensor>& input_ids,
|
||||
const std::optional<torch::stable::Tensor>& tid2eid, cudaStream_t stream) {
|
||||
const std::optional<torch::stable::Tensor>& tid2eid, cudaStream_t stream,
|
||||
const std::optional<torch::stable::Tensor>& is_padding) {
|
||||
const float* bias_ptr = nullptr;
|
||||
if (correction_bias.has_value()) {
|
||||
bias_ptr = correction_bias.value().const_data_ptr<float>();
|
||||
}
|
||||
bool use_hash = false;
|
||||
if (tid2eid.has_value()) {
|
||||
STD_TORCH_CHECK(input_ids.has_value(),
|
||||
"input_ids is required for hash MoE");
|
||||
use_hash = true;
|
||||
}
|
||||
if (topk_indices.scalar_type() == torch::headeronly::ScalarType::Int) {
|
||||
const int* input_ids_ptr = nullptr;
|
||||
const int* tid2eid_ptr = nullptr;
|
||||
if (tid2eid.has_value()) {
|
||||
input_ids_ptr = input_ids.value().const_data_ptr<int>();
|
||||
tid2eid_ptr = tid2eid.value().const_data_ptr<int>();
|
||||
|
||||
auto launch = [&](auto* topk_indices_ptr) {
|
||||
using OutIndType =
|
||||
typename std::remove_pointer<decltype(topk_indices_ptr)>::type;
|
||||
|
||||
const bool* is_padding_ptr = nullptr;
|
||||
if (is_padding.has_value()) {
|
||||
const torch::stable::Tensor& is_padding_tensor = is_padding.value();
|
||||
STD_TORCH_CHECK(is_padding_tensor.scalar_type() ==
|
||||
torch::headeronly::ScalarType::Bool,
|
||||
"is_padding tensor must be bool");
|
||||
STD_TORCH_CHECK(is_padding_tensor.dim() == 1,
|
||||
"is_padding tensor must be 1D");
|
||||
STD_TORCH_CHECK(is_padding_tensor.size(0) == num_tokens,
|
||||
"is_padding size mismatch, expected: ", num_tokens);
|
||||
STD_TORCH_CHECK(is_padding_tensor.is_contiguous(),
|
||||
"is_padding tensor must be contiguous");
|
||||
is_padding_ptr = is_padding_tensor.const_data_ptr<bool>();
|
||||
}
|
||||
|
||||
vllm::moe::topkGatingSoftplusSqrtKernelLauncher<int, ComputeType>(
|
||||
gating_output, topk_weights.mutable_data_ptr<float>(),
|
||||
topk_indices.mutable_data_ptr<int>(),
|
||||
token_expert_indices.mutable_data_ptr<int>(), num_tokens, num_experts,
|
||||
topk, renormalize, routed_scaling_factor, bias_ptr, use_hash,
|
||||
input_ids_ptr, tid2eid_ptr, stream);
|
||||
if (tid2eid.has_value()) {
|
||||
STD_TORCH_CHECK(input_ids.has_value(),
|
||||
"input_ids is required for hash MoE");
|
||||
STD_TORCH_CHECK(
|
||||
input_ids.value().scalar_type() == tid2eid.value().scalar_type(),
|
||||
"input_ids and tid2eid must have the same dtype");
|
||||
if (tid2eid.value().scalar_type() ==
|
||||
torch::headeronly::ScalarType::Long) {
|
||||
vllm::moe::topkGatingSoftplusSqrtKernelLauncher<OutIndType, ComputeType,
|
||||
int64_t>(
|
||||
gating_output, topk_weights.mutable_data_ptr<float>(),
|
||||
topk_indices_ptr, token_expert_indices.mutable_data_ptr<int>(),
|
||||
num_tokens, num_experts, topk, renormalize, routed_scaling_factor,
|
||||
bias_ptr, true, input_ids.value().const_data_ptr<int64_t>(),
|
||||
tid2eid.value().const_data_ptr<int64_t>(), stream, is_padding_ptr);
|
||||
} else {
|
||||
STD_TORCH_CHECK(tid2eid.value().scalar_type() ==
|
||||
torch::headeronly::ScalarType::Int);
|
||||
vllm::moe::topkGatingSoftplusSqrtKernelLauncher<OutIndType, ComputeType,
|
||||
int>(
|
||||
gating_output, topk_weights.mutable_data_ptr<float>(),
|
||||
topk_indices_ptr, token_expert_indices.mutable_data_ptr<int>(),
|
||||
num_tokens, num_experts, topk, renormalize, routed_scaling_factor,
|
||||
bias_ptr, true, input_ids.value().const_data_ptr<int>(),
|
||||
tid2eid.value().const_data_ptr<int>(), stream, is_padding_ptr);
|
||||
}
|
||||
} else {
|
||||
vllm::moe::topkGatingSoftplusSqrtKernelLauncher<OutIndType, ComputeType>(
|
||||
gating_output, topk_weights.mutable_data_ptr<float>(),
|
||||
topk_indices_ptr, token_expert_indices.mutable_data_ptr<int>(),
|
||||
num_tokens, num_experts, topk, renormalize, routed_scaling_factor,
|
||||
bias_ptr, false, static_cast<const OutIndType*>(nullptr),
|
||||
static_cast<const OutIndType*>(nullptr), stream, is_padding_ptr);
|
||||
}
|
||||
};
|
||||
|
||||
if (topk_indices.scalar_type() == torch::headeronly::ScalarType::Int) {
|
||||
launch(topk_indices.mutable_data_ptr<int>());
|
||||
} else if (topk_indices.scalar_type() ==
|
||||
torch::headeronly::ScalarType::UInt32) {
|
||||
const uint32_t* input_ids_ptr = nullptr;
|
||||
const uint32_t* tid2eid_ptr = nullptr;
|
||||
if (tid2eid.has_value()) {
|
||||
input_ids_ptr = input_ids.value().const_data_ptr<uint32_t>();
|
||||
tid2eid_ptr = tid2eid.value().const_data_ptr<uint32_t>();
|
||||
}
|
||||
vllm::moe::topkGatingSoftplusSqrtKernelLauncher<uint32_t, ComputeType>(
|
||||
gating_output, topk_weights.mutable_data_ptr<float>(),
|
||||
topk_indices.mutable_data_ptr<uint32_t>(),
|
||||
token_expert_indices.mutable_data_ptr<int>(), num_tokens, num_experts,
|
||||
topk, renormalize, routed_scaling_factor, bias_ptr, use_hash,
|
||||
input_ids_ptr, tid2eid_ptr, stream);
|
||||
launch(topk_indices.mutable_data_ptr<uint32_t>());
|
||||
} else {
|
||||
STD_TORCH_CHECK(topk_indices.scalar_type() ==
|
||||
torch::headeronly::ScalarType::Long);
|
||||
|
||||
const int64_t* input_ids_ptr = nullptr;
|
||||
const int64_t* tid2eid_ptr = nullptr;
|
||||
if (tid2eid.has_value()) {
|
||||
input_ids_ptr = input_ids.value().const_data_ptr<int64_t>();
|
||||
tid2eid_ptr = tid2eid.value().const_data_ptr<int64_t>();
|
||||
}
|
||||
|
||||
vllm::moe::topkGatingSoftplusSqrtKernelLauncher<int64_t, ComputeType>(
|
||||
gating_output, topk_weights.mutable_data_ptr<float>(),
|
||||
topk_indices.mutable_data_ptr<int64_t>(),
|
||||
token_expert_indices.mutable_data_ptr<int>(), num_tokens, num_experts,
|
||||
topk, renormalize, routed_scaling_factor, bias_ptr, use_hash,
|
||||
input_ids_ptr, tid2eid_ptr, stream);
|
||||
launch(topk_indices.mutable_data_ptr<int64_t>());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -706,7 +842,8 @@ void topk_softplus_sqrt(
|
||||
bool renormalize, double routed_scaling_factor,
|
||||
const std::optional<torch::stable::Tensor>& correction_bias,
|
||||
const std::optional<torch::stable::Tensor>& input_ids,
|
||||
const std::optional<torch::stable::Tensor>& tid2eid) {
|
||||
const std::optional<torch::stable::Tensor>& tid2eid,
|
||||
const std::optional<torch::stable::Tensor>& is_padding) {
|
||||
const int num_experts = gating_output.size(-1);
|
||||
const auto num_tokens = gating_output.numel() / num_experts;
|
||||
const int topk = topk_weights.size(-1);
|
||||
@@ -719,23 +856,24 @@ void topk_softplus_sqrt(
|
||||
dispatch_topk_softplus_sqrt_launch<float>(
|
||||
gating_output.const_data_ptr<float>(), topk_weights, topk_indices,
|
||||
token_expert_indices, num_tokens, num_experts, topk, renormalize,
|
||||
routed_scaling_factor, correction_bias, input_ids, tid2eid, stream);
|
||||
routed_scaling_factor, correction_bias, input_ids, tid2eid, stream,
|
||||
is_padding);
|
||||
} else if (gating_output.scalar_type() ==
|
||||
torch::headeronly::ScalarType::Half) {
|
||||
dispatch_topk_softplus_sqrt_launch<__half>(
|
||||
reinterpret_cast<const __half*>(gating_output.const_data_ptr()),
|
||||
topk_weights, topk_indices, token_expert_indices, num_tokens,
|
||||
num_experts, topk, renormalize, routed_scaling_factor, correction_bias,
|
||||
input_ids, tid2eid, stream);
|
||||
input_ids, tid2eid, stream, is_padding);
|
||||
} else if (gating_output.scalar_type() ==
|
||||
torch::headeronly::ScalarType::BFloat16) {
|
||||
dispatch_topk_softplus_sqrt_launch<__nv_bfloat16>(
|
||||
reinterpret_cast<const __nv_bfloat16*>(gating_output.const_data_ptr()),
|
||||
topk_weights, topk_indices, token_expert_indices, num_tokens,
|
||||
num_experts, topk, renormalize, routed_scaling_factor, correction_bias,
|
||||
input_ids, tid2eid, stream);
|
||||
input_ids, tid2eid, stream, is_padding);
|
||||
} else {
|
||||
STD_TORCH_CHECK(false, "Unsupported gating_output data type: ",
|
||||
gating_output.scalar_type());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,23 +8,28 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_moe_C, m) {
|
||||
m.def(
|
||||
"topk_softmax(Tensor! topk_weights, Tensor! topk_indices, Tensor! "
|
||||
"token_expert_indices, Tensor gating_output, bool renormalize, Tensor? "
|
||||
"bias) -> ()");
|
||||
"bias, Tensor? is_padding) -> ()");
|
||||
|
||||
// Apply topk sigmoid to the gating outputs.
|
||||
m.def(
|
||||
"topk_sigmoid(Tensor! topk_weights, Tensor! topk_indices, Tensor! "
|
||||
"token_expert_indices, Tensor gating_output, bool renormalize, "
|
||||
"Tensor? bias, float routed_scaling_factor) -> ()");
|
||||
"Tensor? bias, float routed_scaling_factor, Tensor? is_padding) -> ()");
|
||||
|
||||
m.def(
|
||||
"topk_softplus_sqrt(Tensor! topk_weights, Tensor! topk_indices, Tensor! "
|
||||
"token_expert_indices, Tensor gating_output, bool renormalize, float "
|
||||
"routed_scaling_factor, Tensor? "
|
||||
"bias, Tensor? input_ids, Tensor? tid2eid) -> ()");
|
||||
"bias, Tensor? input_ids, Tensor? tid2eid, Tensor? is_padding) -> ()");
|
||||
|
||||
// Calculate the result of moe by summing up the partial results
|
||||
// from all selected experts.
|
||||
m.def("moe_sum(Tensor input, Tensor! output) -> ()");
|
||||
// from all selected experts. topk_ids/expert_map are optional and, when
|
||||
// both given, enable pad-aware reduce that skips (token, expert)
|
||||
// slots that were never actually computed (unrouted, or routed to an
|
||||
// expert not owned by this rank under expert parallelism).
|
||||
m.def(
|
||||
"moe_sum(Tensor input, Tensor! output, Tensor? topk_ids=None, "
|
||||
"Tensor? expert_map=None) -> ()");
|
||||
|
||||
// Aligning the number of tokens to be processed by each expert such
|
||||
// that it is divisible by the block size.
|
||||
|
||||
+108
-2
@@ -276,6 +276,61 @@ void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert(
|
||||
torch::stable::Tensor const& cos_sin_cache, double eps,
|
||||
int64_t cache_block_size);
|
||||
|
||||
void fused_kimi_k3_mla_key_concat_kv_cache_insert(
|
||||
torch::stable::Tensor& q, torch::stable::Tensor const& k_nope,
|
||||
torch::stable::Tensor const& k_pe, torch::stable::Tensor const& kv_c_normed,
|
||||
torch::stable::Tensor& k_out, torch::stable::Tensor& k_cache,
|
||||
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
|
||||
std::optional<torch::stable::Tensor> position_ids,
|
||||
std::optional<torch::stable::Tensor> cos_sin_cache);
|
||||
|
||||
void fused_kimi_k3_mla_key_concat_ds_mla_insert(
|
||||
torch::stable::Tensor& q, torch::stable::Tensor const& k_nope,
|
||||
torch::stable::Tensor const& k_pe, torch::stable::Tensor const& kv_c_normed,
|
||||
torch::stable::Tensor& k_out, torch::stable::Tensor& k_cache,
|
||||
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
|
||||
std::optional<torch::stable::Tensor> position_ids,
|
||||
std::optional<torch::stable::Tensor> cos_sin_cache);
|
||||
|
||||
void fused_kimi_k3_mla_qkv_quant_kv_cache_fp8_insert(
|
||||
torch::stable::Tensor const& q, torch::stable::Tensor const& k_nope,
|
||||
torch::stable::Tensor const& k_pe, torch::stable::Tensor const& kv_c_normed,
|
||||
torch::stable::Tensor const& v, torch::stable::Tensor& q_fp8,
|
||||
torch::stable::Tensor& k_fp8, torch::stable::Tensor& v_fp8,
|
||||
torch::stable::Tensor& k_cache, torch::stable::Tensor const& slot_mapping,
|
||||
torch::stable::Tensor const& q_scale_inv,
|
||||
torch::stable::Tensor const& k_scale_inv,
|
||||
torch::stable::Tensor const& v_scale_inv,
|
||||
torch::stable::Tensor const& cache_scale_inv, int64_t cache_block_size,
|
||||
std::optional<torch::stable::Tensor> position_ids,
|
||||
std::optional<torch::stable::Tensor> cos_sin_cache);
|
||||
|
||||
void fused_kimi_k3_mla_decode_q_concat_kv_cache_insert(
|
||||
torch::stable::Tensor const& ql_nope, torch::stable::Tensor const& q_pe,
|
||||
torch::stable::Tensor const& kv_c_normed, torch::stable::Tensor const& k_pe,
|
||||
torch::stable::Tensor& mqa_q, torch::stable::Tensor& k_cache,
|
||||
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
|
||||
std::optional<torch::stable::Tensor> position_ids,
|
||||
std::optional<torch::stable::Tensor> cos_sin_cache);
|
||||
|
||||
void fused_kimi_k3_mla_decode_q_concat_kv_cache_fp8_insert(
|
||||
torch::stable::Tensor const& ql_nope, torch::stable::Tensor const& q_pe,
|
||||
torch::stable::Tensor const& kv_c_normed, torch::stable::Tensor const& k_pe,
|
||||
torch::stable::Tensor& mqa_q, torch::stable::Tensor& k_cache,
|
||||
torch::stable::Tensor const& slot_mapping,
|
||||
torch::stable::Tensor const& q_scale_inv,
|
||||
torch::stable::Tensor const& cache_scale_inv, int64_t cache_block_size,
|
||||
std::optional<torch::stable::Tensor> position_ids,
|
||||
std::optional<torch::stable::Tensor> cos_sin_cache);
|
||||
|
||||
void fused_kimi_k3_mla_decode_q_concat_ds_mla_insert(
|
||||
torch::stable::Tensor const& ql_nope, torch::stable::Tensor const& q_pe,
|
||||
torch::stable::Tensor const& kv_c_normed, torch::stable::Tensor const& k_pe,
|
||||
torch::stable::Tensor& mqa_q, torch::stable::Tensor& k_cache,
|
||||
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
|
||||
std::optional<torch::stable::Tensor> position_ids,
|
||||
std::optional<torch::stable::Tensor> cos_sin_cache);
|
||||
|
||||
void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert(
|
||||
torch::stable::Tensor const& q, torch::stable::Tensor const& kv,
|
||||
torch::stable::Tensor& q_fp8, torch::stable::Tensor& k_cache,
|
||||
@@ -315,6 +370,30 @@ void fused_minimax_m3_qknorm_rope_kv_insert(
|
||||
std::optional<torch::stable::Tensor> index_q_out,
|
||||
const std::string& kv_cache_dtype, bool skip_index_branch);
|
||||
|
||||
#ifdef VLLM_ENABLE_FUSED_KDA_DECODE
|
||||
void fused_kda_decode(
|
||||
torch::stable::Tensor const& x, torch::stable::Tensor const& weight,
|
||||
std::optional<torch::stable::Tensor> bias,
|
||||
torch::stable::Tensor& conv_state, torch::stable::Tensor const& raw_g,
|
||||
torch::stable::Tensor const& raw_beta, torch::stable::Tensor const& a_log,
|
||||
torch::stable::Tensor const& dt_bias,
|
||||
torch::stable::Tensor const& state_indices, torch::stable::Tensor& state,
|
||||
torch::stable::Tensor& out, std::optional<double> lower_bound,
|
||||
std::optional<torch::stable::Tensor> output_gate,
|
||||
std::optional<torch::stable::Tensor> norm_weight, double norm_eps);
|
||||
#endif
|
||||
|
||||
#ifdef VLLM_ENABLE_KIMI_K3_ATTN_RES
|
||||
void kimi_k3_attn_res(torch::stable::Tensor& prefix,
|
||||
torch::stable::Tensor const& delta,
|
||||
torch::stable::Tensor const& blocks,
|
||||
torch::stable::Tensor const& norm_weight,
|
||||
torch::stable::Tensor const& qk_weight,
|
||||
torch::stable::Tensor const& output_norm_weight,
|
||||
torch::stable::Tensor& output, int64_t num_blocks,
|
||||
double eps, double output_norm_eps);
|
||||
#endif
|
||||
|
||||
// Sampler kernels (shared CUDA/ROCm)
|
||||
void apply_repetition_penalties_(
|
||||
torch::stable::Tensor& logits, const torch::stable::Tensor& prompt_mask,
|
||||
@@ -372,6 +451,20 @@ fptr_t init_custom_ar(const std::vector<int64_t>& fake_ipc_ptrs,
|
||||
void all_reduce(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out, fptr_t reg_buffer,
|
||||
int64_t reg_buffer_sz_bytes);
|
||||
void custom_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out, fptr_t reg_buffer,
|
||||
int64_t reg_buffer_sz_bytes);
|
||||
void mnnvl_lamport_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out, fptr_t local_buffer,
|
||||
fptr_t multicast_buffer, fptr_t epoch_buffer,
|
||||
int64_t stage_sz_bytes);
|
||||
void custom_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out, fptr_t reg_buffer,
|
||||
int64_t reg_buffer_sz_bytes);
|
||||
void mnnvl_lamport_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
|
||||
torch::stable::Tensor& out,
|
||||
fptr_t local_buffer, fptr_t epoch_buffer,
|
||||
int64_t stage_sz_bytes);
|
||||
void dispose(fptr_t _fa);
|
||||
int64_t meta_size();
|
||||
void register_buffer(fptr_t _fa, const std::vector<int64_t>& fake_ipc_ptrs);
|
||||
@@ -410,6 +503,12 @@ void fatrelu_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input,
|
||||
double threshold);
|
||||
void swigluoai_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input,
|
||||
double alpha = 1.702, double limit = 7.0);
|
||||
void situ_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input,
|
||||
double beta = 1.0, double linear_beta = -1.0);
|
||||
void masked_situ_and_mul(torch::stable::Tensor& out,
|
||||
torch::stable::Tensor& input,
|
||||
const torch::stable::Tensor& expert_num_tokens,
|
||||
double beta = 1.0, double linear_beta = -1.0);
|
||||
void gelu_new(torch::stable::Tensor& out, torch::stable::Tensor& input);
|
||||
void gelu_fast(torch::stable::Tensor& out, torch::stable::Tensor& input);
|
||||
void gelu_quick(torch::stable::Tensor& out, torch::stable::Tensor& input);
|
||||
@@ -486,6 +585,13 @@ void concat_and_cache_mla(torch::stable::Tensor& kv_c,
|
||||
const std::string& kv_cache_dtype,
|
||||
torch::stable::Tensor& scale);
|
||||
|
||||
void concat_and_cache_mla_grouped(torch::stable::Tensor& kv_c,
|
||||
torch::stable::Tensor& k_pe,
|
||||
torch::stable::Tensor& kv_cache_ptrs,
|
||||
torch::stable::Tensor& slot_mapping,
|
||||
int64_t block_size, int64_t block_stride,
|
||||
int64_t entry_stride);
|
||||
|
||||
// NOTE: k_pe and kv_c order is flipped compared to concat_and_cache_mla
|
||||
void concat_and_cache_mla_rope_fused(
|
||||
torch::stable::Tensor& positions, torch::stable::Tensor& q_pe,
|
||||
@@ -527,9 +633,9 @@ void cp_gather_and_upconvert_fp8_kv_cache(
|
||||
// 656]
|
||||
torch::stable::Tensor const& dst, // [TOT_TOKENS, 576]
|
||||
torch::stable::Tensor const& block_table, // [BATCH, BLOCK_INDICES]
|
||||
torch::stable::Tensor const& seq_lens, // [BATCH]
|
||||
torch::stable::Tensor const& workspace_starts, // [BATCH]
|
||||
int64_t batch_size);
|
||||
int64_t batch_size,
|
||||
std::optional<torch::stable::Tensor> seq_starts = std::nullopt);
|
||||
|
||||
// Indexer K quantization and cache function
|
||||
void indexer_k_quant_and_cache(
|
||||
|
||||
@@ -39,11 +39,15 @@ __global__ void marlin_int4_fp8_preprocess_kernel_awq(
|
||||
// AWQ zeros: (size_k // group_size, size_n // 8)
|
||||
const int32_t* __restrict__ qzeros, int32_t size_n, int32_t size_k,
|
||||
int32_t group_size) {
|
||||
int32_t val =
|
||||
qweight[(blockIdx.x * 32 + threadIdx.x) * size_n / 8 + blockIdx.y];
|
||||
int32_t zero =
|
||||
qzeros[(blockIdx.x * 32 + threadIdx.x) / group_size * size_n / 8 +
|
||||
blockIdx.y];
|
||||
// Thread mapping: threadIdx.x -> column dim (coalesced read within a row),
|
||||
// blockIdx.x -> row dim. Adjacent threads read consecutive int32 in the
|
||||
// same row (stride 1) instead of striding across rows (stride size_n/8).
|
||||
int col = blockIdx.y * 32 + threadIdx.x;
|
||||
if (col >= size_n / 8) return;
|
||||
(void)size_k;
|
||||
|
||||
int32_t val = qweight[blockIdx.x * (size_n / 8) + col];
|
||||
int32_t zero = qzeros[blockIdx.x / group_size * (size_n / 8) + col];
|
||||
int32_t new_val = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -58,7 +62,7 @@ __global__ void marlin_int4_fp8_preprocess_kernel_awq(
|
||||
zero >>= 4;
|
||||
}
|
||||
|
||||
output[(blockIdx.x * 32 + threadIdx.x) * size_n / 8 + blockIdx.y] = new_val;
|
||||
output[blockIdx.x * (size_n / 8) + col] = new_val;
|
||||
}
|
||||
|
||||
torch::stable::Tensor marlin_int4_fp8_preprocess(
|
||||
@@ -102,7 +106,7 @@ torch::stable::Tensor marlin_int4_fp8_preprocess(
|
||||
"qweight.size(0) % qzeros.size(0) != 0");
|
||||
STD_TORCH_CHECK(group_size % 8 == 0, "group_size % 8 != 0");
|
||||
|
||||
dim3 blocks(size_k / 32, size_n / 8);
|
||||
dim3 blocks(size_k, (size_n / 8 + 31) / 32);
|
||||
marlin_int4_fp8_preprocess_kernel_awq<<<blocks, 32, 0, stream>>>(
|
||||
reinterpret_cast<const int32_t*>(qweight.const_data_ptr()),
|
||||
reinterpret_cast<int32_t*>(output.mutable_data_ptr()),
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user