forked from Karylab-cklius/vllm
Compare commits
308
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
aa0db604c1 | ||
|
|
15f1df36e2 | ||
|
|
b666400fcb | ||
|
|
4cce17a1a9 | ||
|
|
0a77b24eac | ||
|
|
c8f09e9cf2 | ||
|
|
5d9b6e0e06 | ||
|
|
cb95b2b98a | ||
|
|
d00bdaee51 | ||
|
|
4fcf47661a | ||
|
|
d626b371f6 | ||
|
|
e8ee5b83eb | ||
|
|
a1e5fe67b9 | ||
|
|
4d2e7ab5b1 | ||
|
|
40d45036cf | ||
|
|
ccf90ba784 | ||
|
|
6adacfcb65 | ||
|
|
14cb86c187 | ||
|
|
8213e8f880 | ||
|
|
3693f922ff | ||
|
|
5c18b961d6 | ||
|
|
f72b20976c | ||
|
|
610a3efcaf | ||
|
|
f414f90601 | ||
|
|
8625ec267b | ||
|
|
995e9a209e | ||
|
|
739e5945dc | ||
|
|
4d042ed85f | ||
|
|
10d9872d3a | ||
|
|
ccd0d1d906 | ||
|
|
d8ddb31644 | ||
|
|
1ce0318c68 | ||
|
|
8d825b87d6 | ||
|
|
1b19bd7589 | ||
|
|
200a727e94 | ||
|
|
edbc1abd1c | ||
|
|
0e39202ca9 | ||
|
|
9dd5ee0117 | ||
|
|
fa6ae31177 | ||
|
|
2a3c32ce67 | ||
|
|
4beeb0689c | ||
|
|
cae984060f | ||
|
|
715681c127 | ||
|
|
dc02271d76 | ||
|
|
4e4ad41d11 | ||
|
|
620e8924d9 | ||
|
|
f00c5539d7 | ||
|
|
21fab0a3db | ||
|
|
3244a2ebf2 | ||
|
|
72ff142c37 | ||
|
|
ee3c0c83db | ||
|
|
cc07dad789 | ||
|
|
17e787a779 | ||
|
|
639402f5a2 | ||
|
|
0f7be0f2f7 | ||
|
|
394ff86965 | ||
|
|
df1e30e74b | ||
|
|
bd8bd52308 | ||
|
|
59b2f7b640 | ||
|
|
92feb9991d | ||
|
|
d4cb783c10 | ||
|
|
eb92ba740a | ||
|
|
a3e750c0a5 | ||
|
|
da72daced2 | ||
|
|
8d0aabdde9 | ||
|
|
0f3ce4c74b | ||
|
|
af661a182d | ||
|
|
7f0b8f2020 | ||
|
|
11e2375fe2 | ||
|
|
fc645f1acc | ||
|
|
2d80cf9d6e | ||
|
|
e7cfd7c5b9 | ||
|
|
e816a8811f | ||
|
|
e281cb721c | ||
|
|
51cfc0e76c | ||
|
|
b87575d24b | ||
|
|
42c6bb4b75 | ||
|
|
ecd1ea1363 | ||
|
|
8f121f7879 | ||
|
|
cb5f7501cb | ||
|
|
8d0f908b98 | ||
|
|
c9dddc144b | ||
|
|
c1cc7344fb | ||
|
|
f976e3b98b | ||
|
|
d468322dc1 | ||
|
|
967146e7bd | ||
|
|
8e8a3becd1 | ||
|
|
1dfd64c1cc | ||
|
|
ad720aefe9 | ||
|
|
270e8a4102 | ||
|
|
f44afef6d6 | ||
|
|
447ce22212 | ||
|
|
65e4e46f66 | ||
|
|
49d20346e4 | ||
|
|
ef076c1b73 | ||
|
|
ec68d53b2b | ||
|
|
13e6b1b908 | ||
|
|
58c0a928c9 | ||
|
|
3dd60971de | ||
|
|
a5b17fba8f | ||
|
|
a7b308e60c | ||
|
|
c48b2b83bd | ||
|
|
68066a99d1 | ||
|
|
24151eb438 | ||
|
|
571e7d3cac | ||
|
|
b0cb81a05b | ||
|
|
e7a1387e73 | ||
|
|
f83de7196f | ||
|
|
3cd32300d6 | ||
|
|
445a2a4d1a | ||
|
|
55d037e2e5 | ||
|
|
ecbfbb8d61 | ||
|
|
e0613702ad | ||
|
|
9853a3c159 | ||
|
|
bb6047db13 | ||
|
|
467d3247c3 | ||
|
|
e5de19ff9a | ||
|
|
edee96519a | ||
|
|
adaabb8a55 | ||
|
|
f7cad67412 | ||
|
|
a8134aef4e | ||
|
|
2800706f06 | ||
|
|
0d310ffbeb | ||
|
|
d5f75fdf50 | ||
|
|
827268e98d | ||
|
|
56e19d7ee2 | ||
|
|
9036d4c464 | ||
|
|
a8c6ee9b78 | ||
|
|
3b1d9c3156 | ||
|
|
54d244f28f | ||
|
|
6c749399b7 | ||
|
|
91eea72330 | ||
|
|
df2503e125 | ||
|
|
c8d98f81f6 | ||
|
|
d87fb264df | ||
|
|
66c079ae83 | ||
|
|
b6c9be509e | ||
|
|
ed733802f0 | ||
|
|
8a34c5087a | ||
|
|
ed2f282bc8 | ||
|
|
9e78555743 | ||
|
|
e80e633927 | ||
|
|
490f17d0c7 | ||
|
|
ccf38056b1 | ||
|
|
2e98406048 | ||
|
|
ef5a226819 | ||
|
|
aec18492d0 | ||
|
|
2a49284c8a | ||
|
|
d37b378762 | ||
|
|
92fbec391b | ||
|
|
2f41d6c063 | ||
|
|
d9b481e248 | ||
|
|
c8661431e0 | ||
|
|
3aecdf08b4 | ||
|
|
eb4205fee5 | ||
|
|
bf0d29dddb | ||
|
|
83aea2147f | ||
|
|
2e9034c998 | ||
|
|
8332078cfd | ||
|
|
ba4a78eb5d | ||
|
|
fdcd95a1a3 | ||
|
|
4f1d426261 | ||
|
|
f3c7941ec8 | ||
|
|
88fa073594 | ||
|
|
a0dd7c27a5 | ||
|
|
3352bf8b03 | ||
|
|
7c94ae16c6 | ||
|
|
4e05add0af | ||
|
|
ad05edfbca | ||
|
|
2018137242 | ||
|
|
a776a48b1c | ||
|
|
8477fe427d | ||
|
|
a65a434cc3 | ||
|
|
c86cb2aeb8 | ||
|
|
e24e0a43a4 | ||
|
|
b55d830ec7 | ||
|
|
62c9357879 | ||
|
|
75e01a39a1 | ||
|
|
512c5eb455 | ||
|
|
13151a4df4 | ||
|
|
56c976c1b5 | ||
|
|
d74a306c4b | ||
|
|
0e9f0a516c | ||
|
|
8904fc4d19 | ||
|
|
1a2c17634e | ||
|
|
308cec5864 | ||
|
|
4e2ab1861d | ||
|
|
140cbb1186 | ||
|
|
6155bbd1dd | ||
|
|
78434b923c | ||
|
|
2488d1dca2 | ||
|
|
d734445fcd | ||
|
|
927975ead8 | ||
|
|
9ea7d670d8 | ||
|
|
7b80cd8ac3 | ||
|
|
2111997f96 | ||
|
|
5af684c319 | ||
|
|
d521dcdbcc | ||
|
|
5daf62271d | ||
|
|
ad3304425b | ||
|
|
70406eb1dc | ||
|
|
de10041d85 | ||
|
|
08bfedc152 | ||
|
|
0102bd2f4c | ||
|
|
83d09d36b5 | ||
|
|
92b9afeecd | ||
|
|
7310555482 | ||
|
|
96b5004b71 | ||
|
|
98e1a43af7 | ||
|
|
729eb59f60 | ||
|
|
6e1100889e | ||
|
|
edcc37a8ce | ||
|
|
79df4a794d | ||
|
|
7c139ab23f | ||
|
|
0be9516ea4 | ||
|
|
7b9de7c892 | ||
|
|
dd9342e6bc | ||
|
|
8060bb0333 | ||
|
|
da4c0e4db9 | ||
|
|
a9a0e0551f | ||
|
|
5c35517a3e | ||
|
|
a435e3108d | ||
|
|
2df2c85be4 | ||
|
|
62095e82c1 | ||
|
|
b2b2c5239e | ||
|
|
00d7b497b3 | ||
|
|
9c81f35b1a | ||
|
|
886ba99a1c | ||
|
|
f186cfe75e | ||
|
|
dfa5062a8f | ||
|
|
e8ebbdde83 | ||
|
|
94fbb09894 | ||
|
|
419e73cdfa | ||
|
|
f01482408c | ||
|
|
bfdc0a3a99 | ||
|
|
93bada494f | ||
|
|
608914de30 | ||
|
|
4ae218c122 | ||
|
|
f40d9879f2 | ||
|
|
47e605092b | ||
|
|
e69a265135 | ||
|
|
fef56c1855 | ||
|
|
c5e3454e5a | ||
|
|
f6983f01de | ||
|
|
3d1d72de29 | ||
|
|
16bfb9cdd4 | ||
|
|
780ba37458 | ||
|
|
9570654c6d | ||
|
|
334e81e90a | ||
|
|
430aacf912 | ||
|
|
d7ccecd2b7 | ||
|
|
d56e952239 | ||
|
|
56de443db1 | ||
|
|
4dd49b06f8 | ||
|
|
f53fa26e05 | ||
|
|
1af6f78ae5 | ||
|
|
228023b3a5 | ||
|
|
9a528260ef | ||
|
|
968ed02ace | ||
|
|
7d266abb22 | ||
|
|
156405d243 | ||
|
|
99e5539a67 | ||
|
|
a88ce94bbb | ||
|
|
2a36d8fb72 | ||
|
|
93726b2a1c | ||
|
|
8617f8676b | ||
|
|
06fd9ffcc4 | ||
|
|
cab4064cd5 | ||
|
|
062f1a2d70 | ||
|
|
81994e1d0e | ||
|
|
4b506ff90a | ||
|
|
5875bb2e9c | ||
|
|
f0d3ad9f3e | ||
|
|
121ea5a21f | ||
|
|
ab79863e6c | ||
|
|
5f1de2b14b | ||
|
|
a5a623d961 | ||
|
|
f8c3af2d85 | ||
|
|
50cd5674b3 | ||
|
|
7b1a7423be | ||
|
|
97f92c6b47 | ||
|
|
1fed50d74f | ||
|
|
f9bf662e5b | ||
|
|
14e2241f77 | ||
|
|
cdd23258cf | ||
|
|
cec6774e9b | ||
|
|
355be167e6 | ||
|
|
d67e21b26e | ||
|
|
6a2c13a6f0 | ||
|
|
b443e6702e | ||
|
|
34d73a3375 | ||
|
|
9d7beab915 | ||
|
|
f2ecfa9cd7 | ||
|
|
e4cdaf199d | ||
|
|
2b72935629 | ||
|
|
1903df8328 | ||
|
|
d872b0a082 | ||
|
|
84deceffb7 | ||
|
|
e269b614c0 | ||
|
|
24090c52f3 | ||
|
|
063fd29c98 | ||
|
|
156e12ba35 | ||
|
|
3e5c06dd7d | ||
|
|
cc08dad785 | ||
|
|
976293e374 | ||
|
|
6efd919548 | ||
|
|
2145abaade | ||
|
|
a17a1f12dc |
@@ -8,8 +8,8 @@ run_all_patterns:
|
||||
- "CMakeLists.txt"
|
||||
- "requirements/common.txt"
|
||||
- "requirements/cuda.txt"
|
||||
- "requirements/build.txt"
|
||||
- "requirements/test.txt"
|
||||
- "requirements/build/cuda.txt"
|
||||
- "requirements/test/cuda.txt"
|
||||
- "setup.py"
|
||||
- "csrc/"
|
||||
- "cmake/"
|
||||
|
||||
@@ -6,8 +6,8 @@ run_all_patterns:
|
||||
- "CMakeLists.txt"
|
||||
- "requirements/common.txt"
|
||||
- "requirements/xpu.txt"
|
||||
- "requirements/build.txt"
|
||||
- "requirements/test.txt"
|
||||
- "requirements/build/cuda.txt"
|
||||
- "requirements/test/cuda.txt"
|
||||
- "setup.py"
|
||||
- "csrc/"
|
||||
- "cmake/"
|
||||
|
||||
@@ -5,7 +5,6 @@ steps:
|
||||
depends_on: []
|
||||
device: amd_cpu
|
||||
no_plugin: true
|
||||
soft_fail: true
|
||||
commands:
|
||||
- >
|
||||
docker build
|
||||
|
||||
@@ -35,6 +35,7 @@ steps:
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager -tp 2 --distributed-executor-backend mp &&
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager --attention-backend=TRITON_ATTN &&
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager --quantization fp8 &&
|
||||
python3 examples/basic/offline_inference/generate.py --model facebook/opt-125m --block-size 64 --enforce-eager --kv-cache-dtype fp8 &&
|
||||
python3 examples/basic/offline_inference/generate.py --model superjob/Qwen3-4B-Instruct-2507-GPTQ-Int4 --block-size 64 --enforce-eager --max-model-len 8192 &&
|
||||
python3 examples/basic/offline_inference/generate.py --model ibm-research/PowerMoE-3b --block-size 64 --enforce-eager -tp 2 &&
|
||||
python3 examples/basic/offline_inference/generate.py --model ibm-research/PowerMoE-3b --block-size 64 --enforce-eager -tp 2 --enable-expert-parallel'
|
||||
@@ -61,4 +62,4 @@ steps:
|
||||
pytest -v -s v1/structured_output &&
|
||||
pytest -v -s v1/test_serial_utils.py &&
|
||||
pytest -v -s v1/spec_decode --ignore=v1/spec_decode/test_max_len.py --ignore=v1/spec_decode/test_tree_attention.py --ignore=v1/spec_decode/test_speculators_eagle3.py --ignore=v1/spec_decode/test_acceptance_length.py &&
|
||||
pytest -v -s v1/kv_connector/unit --ignore=v1/kv_connector/unit/test_multi_connector.py --ignore=v1/kv_connector/unit/test_nixl_connector.py --ignore=v1/kv_connector/unit/test_example_connector.py --ignore=v1/kv_connector/unit/test_lmcache_integration.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'
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
# 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"
|
||||
required_gpu_arch:
|
||||
- gfx942
|
||||
- gfx950
|
||||
tasks:
|
||||
- name: "mmlu_pro"
|
||||
metrics:
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
# For vllm script, with -t option (tensor parallel size)
|
||||
# bash .buildkite/lm-eval-harness/run-lm-eval-gsm-vllm-baseline.sh -m RedHatAI/Qwen2.5-VL-3B-Instruct-FP8-Dynamic -l 1319 -t 1
|
||||
model_name: "RedHatAI/Qwen2.5-VL-3B-Instruct-FP8-Dynamic"
|
||||
required_gpu_arch:
|
||||
- gfx942
|
||||
- gfx950
|
||||
tasks:
|
||||
- name: "gsm8k"
|
||||
metrics:
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
model_name: "Qwen/Qwen3-235B-A22B-Instruct-2507-FP8"
|
||||
required_gpu_arch:
|
||||
- gfx942
|
||||
- gfx950
|
||||
tasks:
|
||||
- name: "mmlu_pro"
|
||||
metrics:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
Qwen2.5-1.5B-Instruct.yaml
|
||||
Meta-Llama-3.2-1B-Instruct-INT8-compressed-tensors.yaml
|
||||
Meta-Llama-3-8B-Instruct-INT8-compressed-tensors-asym.yaml
|
||||
Meta-Llama-3-8B-Instruct-nonuniform-compressed-tensors.yaml
|
||||
Qwen2.5-VL-3B-Instruct-FP8-dynamic.yaml
|
||||
Qwen1.5-MoE-W4A16-compressed-tensors.yaml
|
||||
|
||||
@@ -13,6 +13,7 @@ import os
|
||||
from contextlib import contextmanager
|
||||
|
||||
import lm_eval
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
from vllm.platforms import current_platform
|
||||
@@ -89,9 +90,40 @@ def launch_lm_eval(eval_config, tp_size):
|
||||
return results
|
||||
|
||||
|
||||
def _check_rocm_gpu_arch_requirement(eval_config):
|
||||
"""Skip the test if the model requires a ROCm GPU arch not present.
|
||||
|
||||
Model YAML configs can specify::
|
||||
|
||||
required_gpu_arch:
|
||||
- gfx942
|
||||
- gfx950
|
||||
|
||||
The check only applies on ROCm. On other platforms (e.g. CUDA) the
|
||||
field is ignored so that shared config files work for both NVIDIA and
|
||||
AMD CI pipelines.
|
||||
"""
|
||||
required_archs = eval_config.get("required_gpu_arch")
|
||||
if not required_archs:
|
||||
return
|
||||
|
||||
if not current_platform.is_rocm():
|
||||
return
|
||||
|
||||
from vllm.platforms.rocm import _GCN_ARCH # noqa: E402
|
||||
|
||||
if not any(arch in _GCN_ARCH for arch in required_archs):
|
||||
pytest.skip(
|
||||
f"Model requires GPU arch {required_archs}, "
|
||||
f"but detected arch is '{_GCN_ARCH}'"
|
||||
)
|
||||
|
||||
|
||||
def test_lm_eval_correctness_param(config_filename, tp_size):
|
||||
eval_config = yaml.safe_load(config_filename.read_text(encoding="utf-8"))
|
||||
|
||||
_check_rocm_gpu_arch_requirement(eval_config)
|
||||
|
||||
results = launch_lm_eval(eval_config, tp_size)
|
||||
|
||||
rtol = eval_config.get("rtol", DEFAULT_RTOL)
|
||||
|
||||
@@ -19,7 +19,7 @@ has_new_python=$($PYTHON -c "print(1 if __import__('sys').version_info >= (3,12)
|
||||
if [[ "$has_new_python" -eq 0 ]]; then
|
||||
# use new python from docker
|
||||
docker pull python:3-slim
|
||||
PYTHON="docker run --rm -v $(pwd):/app -w /app python:3-slim python3"
|
||||
PYTHON="docker run --rm -u $(id -u):$(id -g) -v $(pwd):/app -w /app python:3-slim python3"
|
||||
fi
|
||||
|
||||
echo "Using python interpreter: $PYTHON"
|
||||
|
||||
@@ -35,23 +35,6 @@ export PYTHONPATH=".."
|
||||
# Helper Functions
|
||||
###############################################################################
|
||||
|
||||
wait_for_clean_gpus() {
|
||||
local timeout=${1:-300}
|
||||
local start=$SECONDS
|
||||
echo "--- Waiting for clean GPU state (timeout: ${timeout}s)"
|
||||
while true; do
|
||||
if grep -q clean /opt/amdgpu/etc/gpu_state; then
|
||||
echo "GPUs state is \"clean\""
|
||||
return
|
||||
fi
|
||||
if (( SECONDS - start >= timeout )); then
|
||||
echo "Error: GPUs did not reach clean state within ${timeout}s" >&2
|
||||
exit 1
|
||||
fi
|
||||
sleep 3
|
||||
done
|
||||
}
|
||||
|
||||
cleanup_docker() {
|
||||
# Get Docker's root directory
|
||||
docker_root=$(docker info -f '{{.DockerRootDir}}')
|
||||
@@ -365,19 +348,12 @@ apply_rocm_test_overrides() {
|
||||
###############################################################################
|
||||
|
||||
# --- GPU initialization ---
|
||||
echo "--- Confirming Clean Initial State"
|
||||
wait_for_clean_gpus
|
||||
|
||||
echo "--- ROCm info"
|
||||
rocminfo
|
||||
|
||||
# --- Docker housekeeping ---
|
||||
cleanup_docker
|
||||
|
||||
echo "--- Resetting GPUs"
|
||||
echo "reset" > /opt/amdgpu/etc/gpu_state
|
||||
wait_for_clean_gpus
|
||||
|
||||
# --- Pull test image ---
|
||||
echo "--- Pulling container"
|
||||
image_name="rocm/vllm-ci:${BUILDKITE_COMMIT}"
|
||||
|
||||
@@ -42,7 +42,7 @@ WORKDIR /workspace/vllm
|
||||
ENV no_proxy=localhost,127.0.0.1
|
||||
ENV PT_HPU_ENABLE_LAZY_COLLECTIVES=true
|
||||
|
||||
RUN bash -c 'pip install -r <(sed "/^torch/d" requirements/build.txt)'
|
||||
RUN bash -c 'pip install -r <(sed "/^torch/d" requirements/build/cuda.txt)'
|
||||
RUN VLLM_TARGET_DEVICE=empty pip install --no-build-isolation -e .
|
||||
RUN pip install git+https://github.com/vllm-project/vllm-gaudi.git
|
||||
|
||||
|
||||
@@ -50,6 +50,6 @@ docker run \
|
||||
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/spec_decode --ignore=v1/spec_decode/test_max_len.py --ignore=v1/spec_decode/test_tree_attention.py --ignore=v1/spec_decode/test_speculators_eagle3.py --ignore=v1/spec_decode/test_acceptance_length.py
|
||||
pytest -v -s v1/kv_connector/unit --ignore=v1/kv_connector/unit/test_multi_connector.py --ignore=v1/kv_connector/unit/test_nixl_connector.py --ignore=v1/kv_connector/unit/test_example_connector.py --ignore=v1/kv_connector/unit/test_lmcache_integration.py -k "not (test_register_kv_caches and FLASH_ATTN and True)"
|
||||
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
|
||||
pytest -v -s v1/test_serial_utils.py
|
||||
'
|
||||
|
||||
+40
-77
@@ -123,7 +123,7 @@ steps:
|
||||
soft_fail: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- requirements/nightly_torch_test.txt
|
||||
- requirements/test/nightly-torch.txt
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- bash standalone_tests/pytorch_nightly_dependency.sh
|
||||
@@ -532,28 +532,6 @@ steps:
|
||||
- pytest -v -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine
|
||||
|
||||
|
||||
- label: V1 Speculative Decoding (slow) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
|
||||
agent_pool: mi250_1
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/
|
||||
- vllm/model_executor/models/
|
||||
- vllm/v1/attention/
|
||||
- vllm/model_executor/layers/
|
||||
- tests/v1/spec_decode/
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_eagle.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_extract_hidden_states.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_max_len.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_mtp.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_ngram.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_speculators_eagle3.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_tree_attention.py
|
||||
|
||||
|
||||
- label: V1 attention (H100-MI250) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
|
||||
@@ -751,6 +729,7 @@ steps:
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
|
||||
agent_pool: mi250_1
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
@@ -1072,7 +1051,8 @@ steps:
|
||||
- tests/models/multimodal/test_mapping.py
|
||||
commands:
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal/generation/test_memory_leak.py -m core_model
|
||||
- cd .. && VLLM_WORKER_MULTIPROC_METHOD=spawn pytest -v -s tests/models/multimodal/generation/test_whisper.py -m core_model
|
||||
|
||||
|
||||
@@ -1877,28 +1857,6 @@ steps:
|
||||
- pytest -v -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine
|
||||
|
||||
|
||||
- label: V1 Speculative Decoding (slow) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi325]
|
||||
agent_pool: mi325_1
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/
|
||||
- vllm/model_executor/models/
|
||||
- vllm/v1/attention/
|
||||
- vllm/model_executor/layers/
|
||||
- tests/v1/spec_decode/
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_eagle.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_extract_hidden_states.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_max_len.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_mtp.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_ngram.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_speculators_eagle3.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_tree_attention.py
|
||||
|
||||
|
||||
- label: Acceptance Length Test (Large Models) # TBD
|
||||
timeout_in_minutes: 180
|
||||
@@ -1913,7 +1871,7 @@ steps:
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- export VLLM_ALLOW_INSECURE_SERIALIZATION=1
|
||||
- pytest -v -s v1/spec_decode/test_acceptance_length.py -m slow_test
|
||||
- pytest -v -s v1/spec_decode/test_acceptance_length.py
|
||||
|
||||
|
||||
- label: V1 attention (H100-MI325) # 14.5m
|
||||
@@ -2035,7 +1993,6 @@ steps:
|
||||
timeout_in_minutes: 38
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi325]
|
||||
agent_pool: mi325_1
|
||||
optional: true
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
@@ -2165,7 +2122,15 @@ steps:
|
||||
- vllm/platforms/rocm.py
|
||||
- tests/quantization
|
||||
commands:
|
||||
- uv pip install --system torchao==0.14.1
|
||||
|
||||
# temporary install here since we need nightly, will move to requirements/test.in
|
||||
# after torchao 0.12 release, and pin a working version of torchao nightly here
|
||||
|
||||
# since torchao nightly is only compatible with torch nightly currently
|
||||
# https://github.com/pytorch/ao/issues/2919, we'll have to skip new torchao tests for now
|
||||
# we can only upgrade after this is resolved
|
||||
# TODO(jerryzh168): resolve the above comment
|
||||
- uv pip install --system torchao==0.17.0
|
||||
- uv pip install --system conch-triton-kernels
|
||||
- VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py
|
||||
|
||||
@@ -2290,7 +2255,8 @@ steps:
|
||||
- tests/models/multimodal/generation
|
||||
commands:
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal/generation/test_memory_leak.py -m core_model
|
||||
- cd .. && VLLM_WORKER_MULTIPROC_METHOD=spawn pytest -v -s tests/models/multimodal/generation/test_whisper.py -m core_model
|
||||
|
||||
|
||||
@@ -2690,6 +2656,24 @@ steps:
|
||||
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-small.txt
|
||||
|
||||
|
||||
- label: LM Eval Small Models (MI325) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi325]
|
||||
agent_pool: mi325_1
|
||||
working_dir: "/vllm-workspace/.buildkite/lm-eval-harness"
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
- vllm/model_executor/layers/quantization
|
||||
- vllm/model_executor/models/
|
||||
- vllm/model_executor/model_loader/
|
||||
- vllm/v1/attention/backends/
|
||||
- vllm/v1/attention/selector.py
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-small-rocm.txt
|
||||
|
||||
|
||||
- label: LM Eval Small Models (B200-MI325) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi325]
|
||||
@@ -2906,10 +2890,10 @@ steps:
|
||||
- bash .buildkite/scripts/scheduled_integration_test/qwen3_next_mtp_async_eplb.sh 0.8 1319 8040
|
||||
|
||||
##### .buildkite/test_areas/compile.yaml #####
|
||||
# Slowly setting up the tests so that it is also easier for the
|
||||
# Slowly setting up the tests so that it is also easier for the
|
||||
# CI team to review and upstream to the pipelinev2.
|
||||
# The following tests are important for vLLM IR Ops refactoring,
|
||||
# which affects fusion passes on ROCm. So we have to
|
||||
# which affects fusion passes on ROCm. So we have to
|
||||
# enable them as as soon as possible.
|
||||
|
||||
## TODO: Enable the test in this group
|
||||
@@ -2988,7 +2972,7 @@ steps:
|
||||
|
||||
## There are no ops on ROCm for these tests.
|
||||
## The test still passes but the logs are not useful.
|
||||
## fused ops just call torch.ops.symm_mem which
|
||||
## fused ops just call torch.ops.symm_mem which
|
||||
## exists in ROCm even though they don't work
|
||||
# - label: AsyncTP Correctness Tests (2xH100-2xMI325)
|
||||
# - label: Fusion E2E TP2 Quick (H100-MI325)
|
||||
@@ -3160,28 +3144,6 @@ steps:
|
||||
- pytest -v -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine
|
||||
|
||||
|
||||
- label: V1 Speculative Decoding (slow) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx950nightly, amdmi355]
|
||||
agent_pool: mi355_1
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/
|
||||
- vllm/model_executor/models/
|
||||
- vllm/v1/attention/
|
||||
- vllm/model_executor/layers/
|
||||
- tests/v1/spec_decode/
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_eagle.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_extract_hidden_states.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_max_len.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_mtp.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_ngram.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_speculators_eagle3.py
|
||||
- pytest -v -s -m 'slow_test' v1/spec_decode/test_tree_attention.py
|
||||
|
||||
|
||||
- label: V1 attention (B200-MI355) # TBD
|
||||
timeout_in_minutes: 180
|
||||
mirror_hardwares: [amdexperimental, amdproduction, amdgfx950nightly, amdmi355]
|
||||
@@ -3320,7 +3282,7 @@ steps:
|
||||
- vllm/_aiter_ops.py
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- uv pip install --system torchao==0.14.1
|
||||
- uv pip install --system torchao==0.17.0
|
||||
- uv pip install --system conch-triton-kernels
|
||||
- VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py
|
||||
|
||||
@@ -3426,7 +3388,8 @@ steps:
|
||||
- tests/models/multimodal/generation
|
||||
commands:
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal/generation/test_memory_leak.py -m core_model
|
||||
- cd .. && VLLM_WORKER_MULTIPROC_METHOD=spawn pytest -v -s tests/models/multimodal/generation/test_whisper.py -m core_model
|
||||
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Basic Correctness
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/basic_correctness/test_basic_correctness
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Benchmarks CLI Test
|
||||
timeout_in_minutes: 20
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/benchmarks/
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Platform Tests (CUDA)
|
||||
timeout_in_minutes: 15
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/cuda
|
||||
|
||||
@@ -196,6 +196,7 @@ steps:
|
||||
- VLLM_ALLOW_INSECURE_SERIALIZATION=1 python3 examples/rl/rlhf_async_new_apis.py
|
||||
- VLLM_USE_DEEP_GEMM=1 VLLM_LOGGING_LEVEL=DEBUG python3 examples/offline_inference/data_parallel.py --model=Qwen/Qwen1.5-MoE-A2.7B -tp=1 -dp=2 --max-model-len=2048 --all2all-backend=deepep_high_throughput
|
||||
- pytest -v -s tests/v1/distributed/test_dbo.py
|
||||
- TP_SIZE=1 DP_SIZE=2 pytest -v -s tests/v1/distributed/test_eagle_dp.py
|
||||
|
||||
- label: Distributed Tests (2 GPUs)(B200)
|
||||
device: b200
|
||||
@@ -224,20 +225,6 @@ steps:
|
||||
commands:
|
||||
- ./.buildkite/scripts/run-multi-node-test.sh /vllm-workspace/tests 2 2 $IMAGE_TAG "VLLM_TEST_SAME_HOST=0 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_same_node.py | grep 'Same node test passed' && NUM_NODES=2 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_node_count.py | grep 'Node count test passed' && python3 ../examples/offline_inference/data_parallel.py -dp=2 -tp=1 --dp-num-nodes=2 --dp-node-rank=0 --dp-master-addr=192.168.10.10 --dp-master-port=12345 --enforce-eager --trust-remote-code && VLLM_MULTI_NODE=1 pytest -v -s distributed/test_multi_node_assignment.py && VLLM_MULTI_NODE=1 pytest -v -s distributed/test_pipeline_parallel.py" "VLLM_TEST_SAME_HOST=0 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_same_node.py | grep 'Same node test passed' && NUM_NODES=2 torchrun --nnodes 2 --nproc-per-node=2 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_node_count.py | grep 'Node count test passed' && python3 ../examples/offline_inference/data_parallel.py -dp=2 -tp=1 --dp-num-nodes=2 --dp-node-rank=1 --dp-master-addr=192.168.10.10 --dp-master-port=12345 --enforce-eager --trust-remote-code"
|
||||
|
||||
- label: MessageQueue TCP Multi-Node (2 GPUs)
|
||||
timeout_in_minutes: 10
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
num_devices: 1
|
||||
num_nodes: 2
|
||||
no_plugin: true
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
- vllm/distributed/device_communicators/shm_broadcast.py
|
||||
- vllm/distributed/parallel_state.py
|
||||
- tests/distributed/test_mq_tcp_multinode.py
|
||||
commands:
|
||||
- ./.buildkite/scripts/run-multi-node-test.sh /vllm-workspace/tests 2 1 $IMAGE_TAG "torchrun --nnodes 2 --nproc-per-node=1 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_mq_tcp_multinode.py" "torchrun --nnodes 2 --nproc-per-node=1 --rdzv_backend=c10d --rdzv_endpoint=192.168.10.10 distributed/test_mq_tcp_multinode.py"
|
||||
|
||||
- label: Distributed NixlConnector PD accuracy (4 GPUs)
|
||||
timeout_in_minutes: 30
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -282,6 +269,20 @@ steps:
|
||||
- uv pip install --system -r /vllm-workspace/requirements/kv_connectors.txt
|
||||
- HYBRID_SSM=1 bash v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh
|
||||
|
||||
- label: MultiConnector (Nixl+Offloading) PD accuracy (2 GPUs)
|
||||
timeout_in_minutes: 30
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
num_devices: 2
|
||||
source_file_dependencies:
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/nixl_connector.py
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/offloading/
|
||||
- tests/v1/kv_connector/nixl_integration/
|
||||
commands:
|
||||
- uv pip install --system -r /vllm-workspace/requirements/kv_connectors.txt
|
||||
- bash v1/kv_connector/nixl_integration/run_multi_connector_accuracy_test.sh
|
||||
|
||||
- label: NixlConnector PD + Spec Decode acceptance (2 GPUs)
|
||||
timeout_in_minutes: 30
|
||||
device: a100
|
||||
@@ -295,6 +296,20 @@ steps:
|
||||
- uv pip install --system -r /vllm-workspace/requirements/kv_connectors.txt
|
||||
- bash v1/kv_connector/nixl_integration/spec_decode_acceptance_test.sh
|
||||
|
||||
- label: MultiConnector (Nixl+Offloading) PD edge cases (2 GPUs)
|
||||
timeout_in_minutes: 30
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
num_devices: 2
|
||||
source_file_dependencies:
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/nixl_connector.py
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py
|
||||
- vllm/distributed/kv_transfer/kv_connector/v1/offloading/
|
||||
- tests/v1/kv_connector/nixl_integration/
|
||||
commands:
|
||||
- uv pip install --system -r /vllm-workspace/requirements/kv_connectors.txt
|
||||
- bash v1/kv_connector/nixl_integration/run_multi_connector_edge_case_test.sh
|
||||
|
||||
- label: Pipeline + Context Parallelism (4 GPUs)
|
||||
timeout_in_minutes: 60
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Engine
|
||||
timeout_in_minutes: 15
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/engine
|
||||
@@ -25,6 +26,7 @@ steps:
|
||||
|
||||
- label: e2e Scheduling (1 GPU)
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/v1/
|
||||
- tests/v1/e2e/general/
|
||||
|
||||
@@ -61,6 +61,7 @@ steps:
|
||||
|
||||
- label: Entrypoints Integration (API Server openai - Part 3)
|
||||
timeout_in_minutes: 50
|
||||
device: h200_18gb
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -105,6 +106,7 @@ steps:
|
||||
|
||||
- label: OpenAI API Correctness
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
- vllm/entrypoints/openai/
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: EPLB Algorithm
|
||||
timeout_in_minutes: 15
|
||||
device: h200_18gb
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
source_file_dependencies:
|
||||
- vllm/distributed/eplb
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: vLLM IR Tests
|
||||
timeout_in_minutes: 10
|
||||
device: h200_18gb
|
||||
working_dir: "/vllm-workspace/"
|
||||
source_file_dependencies:
|
||||
- vllm/ir
|
||||
@@ -17,10 +18,22 @@ steps:
|
||||
source_file_dependencies:
|
||||
- csrc/
|
||||
- tests/kernels/core
|
||||
- tests/kernels/test_top_k_per_row.py
|
||||
- tests/kernels/test_concat_mla_q.py
|
||||
commands:
|
||||
- pytest -v -s kernels/core kernels/test_top_k_per_row.py kernels/test_concat_mla_q.py
|
||||
- pytest -v -s kernels/core --ignore=kernels/core/test_minimax_reduce_rms.py kernels/test_concat_mla_q.py
|
||||
|
||||
- label: Kernels MiniMax Reduce RMS Test (2 GPUs)
|
||||
timeout_in_minutes: 15
|
||||
num_devices: 2
|
||||
device: h100
|
||||
source_file_dependencies:
|
||||
- csrc/minimax_reduce_rms_kernel.cu
|
||||
- csrc/minimax_reduce_rms_kernel.h
|
||||
- vllm/model_executor/layers/mamba/linear_attn.py
|
||||
- vllm/model_executor/layers/mamba/lamport_workspace.py
|
||||
- tests/kernels/core/test_minimax_reduce_rms.py
|
||||
commands:
|
||||
- pytest -v -s kernels/core/test_minimax_reduce_rms.py
|
||||
|
||||
- label: Kernels Attention Test %N
|
||||
timeout_in_minutes: 35
|
||||
@@ -106,6 +119,7 @@ steps:
|
||||
- vllm/v1/attention/backends/mla/flashinfer_mla.py
|
||||
- vllm/v1/attention/selector.py
|
||||
- vllm/platforms/cuda.py
|
||||
- tests/kernels/test_top_k_per_row.py
|
||||
commands:
|
||||
- nvidia-smi
|
||||
- python3 examples/basic/offline_inference/chat.py
|
||||
@@ -116,6 +130,7 @@ steps:
|
||||
- pytest -v -s tests/kernels/attention/test_flashinfer_trtllm_attention.py
|
||||
- pytest -v -s tests/kernels/attention/test_cutlass_mla_decode.py
|
||||
- pytest -v -s tests/kernels/attention/test_flashinfer_mla_decode.py
|
||||
- pytest -v -s tests/kernels/test_top_k_per_row.py
|
||||
# Quantization
|
||||
- pytest -v -s tests/kernels/quantization/test_cutlass_scaled_mm.py -k 'fp8'
|
||||
- pytest -v -s tests/kernels/quantization/test_nvfp4_quant.py
|
||||
@@ -179,3 +194,21 @@ steps:
|
||||
- pytest -v -s kernels/moe/test_flashinfer_moe.py
|
||||
- pytest -v -s kernels/moe/test_nvfp4_moe.py
|
||||
- pytest -v -s kernels/moe/test_ocp_mx_moe.py
|
||||
|
||||
|
||||
- label: Kernels FusedMoE Layer Test (2 H100s)
|
||||
timeout_in_minutes: 90
|
||||
device: h100
|
||||
num_devices: 2
|
||||
optional: true
|
||||
commands:
|
||||
- pytest -v -s kernels/moe/test_moe_layer.py
|
||||
|
||||
|
||||
- label: Kernels FusedMoE Layer Test (2 B200s)
|
||||
timeout_in_minutes: 90
|
||||
device: b200
|
||||
num_devices: 2
|
||||
optional: true
|
||||
commands:
|
||||
- pytest -v -s kernels/moe/test_moe_layer.py
|
||||
|
||||
@@ -19,6 +19,7 @@ steps:
|
||||
|
||||
- label: V1 Sample + Logits
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/v1/sample
|
||||
@@ -86,6 +87,7 @@ steps:
|
||||
|
||||
- label: Regression
|
||||
timeout_in_minutes: 20
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/test_regression
|
||||
@@ -174,6 +176,7 @@ steps:
|
||||
- tests/renderers
|
||||
- tests/standalone_tests/lazy_imports.py
|
||||
- tests/tokenizers_
|
||||
- tests/reasoning
|
||||
- tests/tool_parsers
|
||||
- tests/transformers_utils
|
||||
- tests/config
|
||||
@@ -187,6 +190,7 @@ steps:
|
||||
- pytest -v -s -m 'cpu_test' multimodal
|
||||
- pytest -v -s renderers
|
||||
- pytest -v -s tokenizers_
|
||||
- pytest -v -s reasoning --ignore=reasoning/test_seedoss_reasoning_parser.py --ignore=reasoning/test_glm4_moe_reasoning_parser.py --ignore=reasoning/test_gemma4_reasoning_parser.py
|
||||
- pytest -v -s tool_parsers
|
||||
- pytest -v -s transformers_utils
|
||||
- pytest -v -s config
|
||||
|
||||
@@ -78,7 +78,6 @@ steps:
|
||||
- TP_SIZE=1 DP_SIZE=2 pytest -v -s v1/distributed/test_async_llm_dp.py -k "not ray"
|
||||
- TP_SIZE=1 DP_SIZE=2 pytest -v -s v1/distributed/test_eagle_dp.py
|
||||
|
||||
# These require fix https://github.com/vllm-project/vllm/pull/36280
|
||||
- label: Model Runner V2 Pipeline Parallelism (4 GPUs)
|
||||
timeout_in_minutes: 60
|
||||
working_dir: "/vllm-workspace/tests"
|
||||
@@ -101,11 +100,13 @@ steps:
|
||||
- vllm/v1/worker/gpu/
|
||||
- vllm/v1/worker/gpu_worker.py
|
||||
- tests/v1/spec_decode/test_max_len.py
|
||||
- tests/v1/spec_decode/test_probabilistic_rejection_sampler_utils.py
|
||||
- tests/v1/spec_decode/test_synthetic_rejection_sampler_utils.py
|
||||
- tests/v1/e2e/spec_decode/test_spec_decode.py
|
||||
commands:
|
||||
- set -x
|
||||
- export VLLM_USE_V2_MODEL_RUNNER=1
|
||||
- pytest -v -s v1/spec_decode/test_max_len.py -k "eagle or mtp"
|
||||
- pytest -v -s v1/spec_decode/test_probabilistic_rejection_sampler_utils.py
|
||||
- pytest -v -s v1/spec_decode/test_synthetic_rejection_sampler_utils.py
|
||||
- pytest -v -s v1/e2e/spec_decode/test_spec_decode.py -k "eagle or mtp"
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Basic Models Tests (Initialization)
|
||||
timeout_in_minutes: 45
|
||||
device: h200_18gb
|
||||
torch_nightly: true
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
|
||||
@@ -38,7 +38,7 @@ steps:
|
||||
# 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.5.2'
|
||||
- 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
|
||||
parallelism: 2
|
||||
@@ -53,7 +53,7 @@ steps:
|
||||
# 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.5.2'
|
||||
- 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)'
|
||||
mirror:
|
||||
amd:
|
||||
@@ -67,6 +67,7 @@ steps:
|
||||
|
||||
- label: Language Models Test (PPL)
|
||||
timeout_in_minutes: 110
|
||||
device: h200_18gb
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
@@ -90,6 +91,7 @@ steps:
|
||||
|
||||
- label: Language Models Test (MTEB)
|
||||
timeout_in_minutes: 110
|
||||
device: h200_18gb
|
||||
optional: true
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: "Multi-Modal Models (Standard) 1: qwen2"
|
||||
timeout_in_minutes: 45
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/models/multimodal
|
||||
@@ -19,6 +20,7 @@ steps:
|
||||
|
||||
- label: "Multi-Modal Models (Standard) 2: qwen3 + gemma"
|
||||
timeout_in_minutes: 45
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/models/multimodal
|
||||
@@ -54,7 +56,8 @@ steps:
|
||||
- tests/models/multimodal
|
||||
commands:
|
||||
- pip install git+https://github.com/TIGER-AI-Lab/Mantis.git
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/processing
|
||||
- pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/processing
|
||||
- pytest models/multimodal/generation/test_memory_leak.py -m core_model
|
||||
- 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:
|
||||
@@ -77,6 +80,7 @@ steps:
|
||||
|
||||
- label: Multi-Modal Processor # 44min
|
||||
timeout_in_minutes: 60
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/models/multimodal
|
||||
@@ -131,6 +135,7 @@ steps:
|
||||
|
||||
- label: Multi-Modal Models (Extended Pooling)
|
||||
optional: true
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/models/multimodal/pooling
|
||||
|
||||
@@ -49,6 +49,7 @@ steps:
|
||||
|
||||
- label: PyTorch Fullgraph
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/compile
|
||||
@@ -60,8 +61,9 @@ steps:
|
||||
# if this test fails, it means the nightly torch version is not compatible with some
|
||||
# of the dependencies. Please check the error message and add the package to whitelist
|
||||
# in /vllm/tools/pre_commit/generate_nightly_torch_test.py
|
||||
device: h200_18gb
|
||||
soft_fail: true
|
||||
source_file_dependencies:
|
||||
- requirements/nightly_torch_test.txt
|
||||
- requirements/test/nightly-torch.txt
|
||||
commands:
|
||||
- bash standalone_tests/pytorch_nightly_dependency.sh
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
group: Quantization
|
||||
depends_on:
|
||||
depends_on:
|
||||
- image-build
|
||||
steps:
|
||||
- label: Quantization
|
||||
@@ -9,14 +9,14 @@ steps:
|
||||
- vllm/model_executor/layers/quantization
|
||||
- tests/quantization
|
||||
commands:
|
||||
# temporary install here since we need nightly, will move to requirements/test.in
|
||||
# temporary install here since we need nightly, will move to requirements/test/cuda.in
|
||||
# after torchao 0.12 release, and pin a working version of torchao nightly here
|
||||
|
||||
# since torchao nightly is only compatible with torch nightly currently
|
||||
# https://github.com/pytorch/ao/issues/2919, we'll have to skip new torchao tests for now
|
||||
# we can only upgrade after this is resolved
|
||||
# TODO(jerryzh168): resolve the above comment
|
||||
- uv pip install --system torchao==0.14.1 --index-url https://download.pytorch.org/whl/cu129
|
||||
- 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
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ steps:
|
||||
# If this fails, it means the PR introduces a dependency that
|
||||
# conflicts with Ray's dependency constraints.
|
||||
# See https://github.com/vllm-project/vllm/issues/33599
|
||||
device: h200_18gb
|
||||
soft_fail: true
|
||||
timeout_in_minutes: 10
|
||||
source_file_dependencies:
|
||||
|
||||
@@ -4,6 +4,7 @@ depends_on:
|
||||
steps:
|
||||
- label: Spec Decode Eagle
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/
|
||||
- vllm/v1/worker/gpu/spec_decode/
|
||||
@@ -13,6 +14,7 @@ steps:
|
||||
|
||||
- label: Spec Decode Speculators + MTP
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/
|
||||
- vllm/v1/worker/gpu/spec_decode/
|
||||
@@ -23,6 +25,7 @@ steps:
|
||||
|
||||
- label: Spec Decode Ngram + Suffix
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/
|
||||
- vllm/v1/worker/gpu/spec_decode/
|
||||
@@ -32,6 +35,7 @@ steps:
|
||||
|
||||
- label: Spec Decode Draft Model
|
||||
timeout_in_minutes: 30
|
||||
device: h200_18gb
|
||||
source_file_dependencies:
|
||||
- vllm/v1/spec_decode/
|
||||
- vllm/v1/worker/gpu/spec_decode/
|
||||
|
||||
+7
-7
@@ -3,7 +3,7 @@
|
||||
|
||||
# This lists cover the "core" components of vLLM that require careful review
|
||||
/vllm/compilation @zou3519 @youkaichao @ProExpertProg @BoyuanFeng @vadiklyutiy
|
||||
/vllm/distributed/kv_transfer @NickLucche @ApostaC @orozery
|
||||
/vllm/distributed/kv_transfer @NickLucche @ApostaC @orozery @xuechendi
|
||||
/vllm/lora @jeejeelee
|
||||
/vllm/model_executor/layers/attention @LucasWilkinson @MatthewBonanni
|
||||
/vllm/model_executor/layers/fused_moe @mgoin @pavanimajety
|
||||
@@ -120,16 +120,16 @@ mkdocs.yaml @hmellor
|
||||
/tools/pre_commit @hmellor
|
||||
|
||||
# CPU
|
||||
/vllm/v1/worker/cpu* @bigPYJ1151
|
||||
/vllm/v1/worker/cpu* @bigPYJ1151 @xuechendi
|
||||
/csrc/cpu @bigPYJ1151
|
||||
/vllm/platforms/cpu.py @bigPYJ1151
|
||||
/vllm/platforms/cpu.py @bigPYJ1151 @xuechendi
|
||||
/cmake/cpu_extension.cmake @bigPYJ1151
|
||||
/docker/Dockerfile.cpu @bigPYJ1151
|
||||
/docker/Dockerfile.cpu @bigPYJ1151 @xuechendi
|
||||
|
||||
# Intel GPU
|
||||
/vllm/v1/worker/xpu* @jikunshang
|
||||
/vllm/platforms/xpu.py @jikunshang
|
||||
/docker/Dockerfile.xpu @jikunshang
|
||||
/vllm/v1/worker/xpu* @jikunshang @xuechendi
|
||||
/vllm/platforms/xpu.py @jikunshang @xuechendi
|
||||
/docker/Dockerfile.xpu @jikunshang @xuechendi
|
||||
|
||||
# Nemotron-specific files
|
||||
/vllm/model_executor/models/*nemotron* @tomeras91
|
||||
|
||||
+29
-8
@@ -18,7 +18,7 @@ pull_request_rules:
|
||||
- name: comment-pre-commit-failure
|
||||
description: Comment on PR when pre-commit check fails
|
||||
conditions:
|
||||
- status-failure=pre-commit
|
||||
- check-failure=pre-commit
|
||||
- -closed
|
||||
- -draft
|
||||
actions:
|
||||
@@ -51,7 +51,7 @@ pull_request_rules:
|
||||
- name: comment-dco-failure
|
||||
description: Comment on PR when DCO check fails
|
||||
conditions:
|
||||
- status-failure=dco
|
||||
- check-failure=dco
|
||||
- -closed
|
||||
- -draft
|
||||
actions:
|
||||
@@ -83,8 +83,8 @@ pull_request_rules:
|
||||
- or:
|
||||
- files~=^examples/.*deepseek.*\.py
|
||||
- files~=^tests/.*deepseek.*\.py
|
||||
- files~=^vllm/entrypoints/openai/tool_parsers/.*deepseek.*\.py
|
||||
- files~=^vllm/model_executor/models/.*deepseek.*\.py
|
||||
- files~=^vllm/tool_parsers/.*deepseek.*\.py
|
||||
- files~=^vllm/reasoning/.*deepseek.*\.py
|
||||
- files~=^vllm/transformers_utils/.*deepseek.*\.py
|
||||
- title~=(?i)DeepSeek
|
||||
@@ -110,9 +110,10 @@ pull_request_rules:
|
||||
- or:
|
||||
- files~=^examples/.*llama.*\.py
|
||||
- files~=^tests/.*llama.*\.py
|
||||
- files~=^vllm/entrypoints/openai/tool_parsers/llama.*\.py
|
||||
- files~=^vllm/model_executor/models/.*llama.*\.py
|
||||
- files~=^vllm/transformers_utils/configs/.*llama.*\.py
|
||||
- files~=^vllm/reasoning/.*llama.*\.py
|
||||
- files~=^vllm/tool_parsers/.*llama.*\.py
|
||||
- files~=^vllm/transformers_utils/.*llama.*\.py
|
||||
- title~=(?i)llama
|
||||
actions:
|
||||
label:
|
||||
@@ -133,6 +134,23 @@ pull_request_rules:
|
||||
add:
|
||||
- multi-modality
|
||||
|
||||
- name: label-mistral
|
||||
description: Automatically apply mistral label
|
||||
conditions:
|
||||
- label != stale
|
||||
- or:
|
||||
- files~=^examples/.*mistral.*\.py
|
||||
- files~=^tests/.*mistral.*\.py
|
||||
- files~=^vllm/model_executor/models/.*mistral.*\.py
|
||||
- files~=^vllm/reasoning/.*mistral.*\.py
|
||||
- files~=^vllm/tool_parsers/.*mistral.*\.py
|
||||
- files~=^vllm/transformers_utils/.*mistral.*\.py
|
||||
- title~=(?i)Mistral
|
||||
actions:
|
||||
label:
|
||||
add:
|
||||
- mistral
|
||||
|
||||
- name: label-new-model
|
||||
description: Automatically apply new-model label
|
||||
conditions:
|
||||
@@ -167,7 +185,9 @@ pull_request_rules:
|
||||
- files~=^examples/.*qwen.*\.py
|
||||
- files~=^tests/.*qwen.*\.py
|
||||
- files~=^vllm/model_executor/models/.*qwen.*\.py
|
||||
- files~=^vllm/tool_parsers/.*qwen.*\.py
|
||||
- files~=^vllm/reasoning/.*qwen.*\.py
|
||||
- files~=^vllm/transformers_utils/.*qwen.*\.py
|
||||
- title~=(?i)Qwen
|
||||
actions:
|
||||
label:
|
||||
@@ -378,17 +398,18 @@ pull_request_rules:
|
||||
add:
|
||||
- tool-calling
|
||||
|
||||
- name: auto-rebase if approved, ready, and 40 commits behind main
|
||||
- name: auto-rebase to keep merge candidate within 1 day behind main
|
||||
conditions:
|
||||
- base = main
|
||||
- label=ready
|
||||
- "#approved-reviews-by >= 1"
|
||||
- "#commits-behind >= 40"
|
||||
- "#commits-behind >= 50"
|
||||
- "#check-failure = 0"
|
||||
- -closed
|
||||
- -draft
|
||||
- -conflict
|
||||
actions:
|
||||
rebase: {}
|
||||
update: {}
|
||||
|
||||
- name: ping author on conflicts and add 'needs-rebase' label
|
||||
conditions:
|
||||
|
||||
@@ -320,20 +320,25 @@ jobs:
|
||||
script: |
|
||||
// Configuration: Map labels to GitHub users to CC
|
||||
// You can add multiple users per label, and multiple label configurations
|
||||
// {users} will be replaced with @mentions
|
||||
const ccConfig = {
|
||||
rocm: {
|
||||
users: ['hongxiayang', 'tjtanaa', 'vllmellm'], // Add more users as needed: ['user1', 'user2', 'user3']
|
||||
message: 'CC {users} for ROCm-related issue' // {users} will be replaced with @mentions
|
||||
users: ['hongxiayang', 'tjtanaa', 'vllmellm'],
|
||||
message: 'CC {users} for ROCm-related issue',
|
||||
},
|
||||
mistral: {
|
||||
users: ['patrickvonplaten', 'juliendenize', 'andylolu2'],
|
||||
message: 'CC {users} for Mistral-related issue',
|
||||
},
|
||||
// Add more label -> user mappings here
|
||||
// Example:
|
||||
// cuda: {
|
||||
// users: ['user1', 'user2'],
|
||||
// message: 'CC {users} for CUDA-related issue'
|
||||
// message: 'CC {users} for CUDA-related issue',
|
||||
// },
|
||||
// performance: {
|
||||
// users: ['perfexpert'],
|
||||
// message: 'CC {users} for performance issue'
|
||||
// message: 'CC {users} for performance issue',
|
||||
// },
|
||||
};
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ jobs:
|
||||
|
||||
- name: Install dependencies and build vLLM
|
||||
run: |
|
||||
uv pip install -r requirements/cpu-build.txt --index-strategy unsafe-best-match
|
||||
uv pip install -r requirements/build/cpu.txt --index-strategy unsafe-best-match
|
||||
uv pip install -r requirements/cpu.txt --index-strategy unsafe-best-match
|
||||
uv pip install -e . --no-build-isolation
|
||||
env:
|
||||
|
||||
@@ -2,6 +2,7 @@ name: pre-commit
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
types: [opened, synchronize, reopened, labeled]
|
||||
push:
|
||||
branches: [main]
|
||||
|
||||
@@ -15,7 +16,11 @@ permissions:
|
||||
|
||||
jobs:
|
||||
pre-run-check:
|
||||
if: github.event_name == 'pull_request'
|
||||
if: >-
|
||||
github.event_name == 'pull_request' &&
|
||||
(github.event.action != 'labeled' ||
|
||||
github.event.label.name == 'ready' ||
|
||||
github.event.label.name == 'verified')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check PR label and author merge count
|
||||
@@ -44,7 +49,12 @@ jobs:
|
||||
|
||||
pre-commit:
|
||||
needs: pre-run-check
|
||||
if: always() && (needs.pre-run-check.result == 'success' || needs.pre-run-check.result == 'skipped')
|
||||
if: >-
|
||||
always() &&
|
||||
(github.event.action != 'labeled' ||
|
||||
github.event.label.name == 'ready' ||
|
||||
github.event.label.name == 'verified') &&
|
||||
(needs.pre-run-check.result == 'success' || needs.pre-run-check.result == 'skipped')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@8e8c483db84b4bee98b60c0593521ed34d9990e8 # v6.0.1
|
||||
|
||||
@@ -9,7 +9,7 @@ PATH=${cuda_home}/bin:$PATH
|
||||
LD_LIBRARY_PATH=${cuda_home}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install requirements
|
||||
$python_executable -m pip install -r requirements/build.txt -r requirements/cuda.txt
|
||||
$python_executable -m pip install -r requirements/build/cuda.txt -r requirements/cuda.txt
|
||||
|
||||
# Limit the number of parallel jobs to avoid OOM
|
||||
export MAX_JOBS=1
|
||||
|
||||
@@ -12,6 +12,9 @@ vllm/third_party/triton_kernels/*
|
||||
# FlashMLA interface copied from source
|
||||
vllm/third_party/flashmla/flash_mla_interface.py
|
||||
|
||||
# DeepGEMM vendored package built from source
|
||||
vllm/third_party/deep_gemm/
|
||||
|
||||
# triton jit
|
||||
.triton
|
||||
|
||||
@@ -26,6 +29,7 @@ __pycache__/
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
!requirements/build/
|
||||
cmake-build-*/
|
||||
CMakeUserPresets.json
|
||||
develop-eggs/
|
||||
|
||||
+65
-10
@@ -39,15 +39,24 @@ repos:
|
||||
rev: 0.11.1
|
||||
hooks:
|
||||
- id: pip-compile
|
||||
args: [requirements/test.in, -c, requirements/common.txt, -o, requirements/test.txt, --index-strategy, unsafe-best-match, --torch-backend, cu129, --python-platform, x86_64-manylinux_2_28, --python-version, "3.12"]
|
||||
files: ^requirements/test\.(in|txt)$
|
||||
args: [
|
||||
requirements/test/cuda.in,
|
||||
-c, requirements/cuda.txt,
|
||||
-o, requirements/test/cuda.txt,
|
||||
--index-strategy, unsafe-best-match,
|
||||
--torch-backend, cu130,
|
||||
--python-platform, x86_64-manylinux_2_28,
|
||||
--python-version, "3.12",
|
||||
]
|
||||
files: ^requirements/(common|cuda|test/cuda)\.(in|txt)$
|
||||
- id: pip-compile
|
||||
alias: pip-compile-rocm
|
||||
name: pip-compile-rocm
|
||||
args: [
|
||||
requirements/rocm-test.in, -o, requirements/rocm-test.txt,
|
||||
--index-strategy, unsafe-best-match,
|
||||
requirements/test/rocm.in,
|
||||
-c, requirements/rocm.txt,
|
||||
-o, requirements/test/rocm.txt,
|
||||
--index-strategy, unsafe-best-match,
|
||||
--python-platform, x86_64-manylinux_2_28,
|
||||
--python-version, "3.12",
|
||||
# Exclude torch and CUDA/NVIDIA packages
|
||||
@@ -59,30 +68,76 @@ repos:
|
||||
--no-emit-package, cuda-pathfinder,
|
||||
--no-emit-package, cuda-toolkit,
|
||||
--no-emit-package, cupy-cuda12x,
|
||||
# nvidia packages (unsuffixed / unified naming)
|
||||
--no-emit-package, nvidia-cublas,
|
||||
--no-emit-package, nvidia-cuda-cupti,
|
||||
--no-emit-package, nvidia-cuda-nvrtc,
|
||||
--no-emit-package, nvidia-cuda-runtime,
|
||||
--no-emit-package, nvidia-cudnn-cu13,
|
||||
--no-emit-package, nvidia-cudnn,
|
||||
--no-emit-package, nvidia-cufft,
|
||||
--no-emit-package, nvidia-cufile,
|
||||
--no-emit-package, nvidia-curand,
|
||||
--no-emit-package, nvidia-cusolver,
|
||||
--no-emit-package, nvidia-cusparse,
|
||||
--no-emit-package, nvidia-cusparselt,
|
||||
--no-emit-package, nvidia-nccl,
|
||||
--no-emit-package, nvidia-nvjitlink,
|
||||
--no-emit-package, nvidia-nvshmem,
|
||||
--no-emit-package, nvidia-nvtx,
|
||||
# nvidia cu12 packages
|
||||
--no-emit-package, nvidia-cublas-cu12,
|
||||
--no-emit-package, nvidia-cuda-cupti-cu12,
|
||||
--no-emit-package, nvidia-cuda-nvrtc-cu12,
|
||||
--no-emit-package, nvidia-cuda-runtime-cu12,
|
||||
--no-emit-package, nvidia-cudnn-cu12,
|
||||
--no-emit-package, nvidia-cufft-cu12,
|
||||
--no-emit-package, nvidia-cufile-cu12,
|
||||
--no-emit-package, nvidia-curand-cu12,
|
||||
--no-emit-package, nvidia-cusolver-cu12,
|
||||
--no-emit-package, nvidia-cusparse-cu12,
|
||||
--no-emit-package, nvidia-cusparselt-cu12,
|
||||
--no-emit-package, nvidia-nccl-cu12,
|
||||
--no-emit-package, nvidia-nvjitlink-cu12,
|
||||
--no-emit-package, nvidia-nvshmem-cu12,
|
||||
--no-emit-package, nvidia-nvtx-cu12,
|
||||
# nvidia cu13 packages
|
||||
--no-emit-package, nvidia-cublas-cu13,
|
||||
--no-emit-package, nvidia-cuda-cupti-cu13,
|
||||
--no-emit-package, nvidia-cuda-nvrtc-cu13,
|
||||
--no-emit-package, nvidia-cuda-runtime-cu13,
|
||||
--no-emit-package, nvidia-cudnn-cu13,
|
||||
--no-emit-package, nvidia-cufft-cu13,
|
||||
--no-emit-package, nvidia-cufile-cu13,
|
||||
--no-emit-package, nvidia-curand-cu13,
|
||||
--no-emit-package, nvidia-cusolver-cu13,
|
||||
--no-emit-package, nvidia-cusparse-cu13,
|
||||
--no-emit-package, nvidia-cusparselt-cu13,
|
||||
--no-emit-package, nvidia-nccl-cu13,
|
||||
--no-emit-package, nvidia-nvjitlink,
|
||||
--no-emit-package, nvidia-nvjitlink-cu13,
|
||||
--no-emit-package, nvidia-nvshmem-cu13,
|
||||
--no-emit-package, nvidia-nvtx,
|
||||
--no-emit-package, nvidia-nvtx-cu13,
|
||||
]
|
||||
files: ^requirements/rocm-test\.(in|txt)$
|
||||
files: ^requirements/(common|rocm|test/rocm)\.(in|txt)$
|
||||
- id: pip-compile
|
||||
alias: pip-compile-xpu
|
||||
name: pip-compile-xpu
|
||||
args: [
|
||||
requirements/test/xpu.in,
|
||||
-c, requirements/xpu.txt,
|
||||
-o, requirements/test/xpu.txt,
|
||||
--index-strategy, unsafe-best-match,
|
||||
--torch-backend, xpu,
|
||||
--python-platform, x86_64-manylinux_2_39,
|
||||
--python-version, "3.12",
|
||||
]
|
||||
files: ^requirements/(common|xpu|test/xpu)\.(in|txt)$
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: format-torch-nightly-test
|
||||
name: reformat nightly_torch_test.txt to be in sync with test.in
|
||||
name: reformat test/nightly-torch.txt to be in sync with test/cuda.in
|
||||
language: python
|
||||
entry: python tools/pre_commit/generate_nightly_torch_test.py
|
||||
files: ^requirements/test\.(in|txt)$
|
||||
files: ^requirements/test/cuda\.(in|txt)$
|
||||
- id: mypy-local
|
||||
name: Run mypy locally for lowest supported Python version
|
||||
entry: python tools/pre_commit/mypy.py 0 "3.10"
|
||||
|
||||
@@ -72,11 +72,11 @@ uv pip install -e . --torch-backend=auto
|
||||
|
||||
```bash
|
||||
# Install test dependencies.
|
||||
# requirements/test.txt is pinned to x86_64; on other platforms, use the
|
||||
# requirements/test/cuda.txt is pinned to x86_64; on other platforms, use the
|
||||
# unpinned source file instead:
|
||||
uv pip install -r requirements/test.in # resolves for current platform
|
||||
uv pip install -r requirements/test/cuda.in # resolves for current platform
|
||||
# Or on x86_64:
|
||||
uv pip install -r requirements/test.txt
|
||||
uv pip install -r requirements/test/cuda.txt
|
||||
|
||||
# Run a specific test file (use .venv/bin/python directly;
|
||||
# `source activate` does not persist in non-interactive shells):
|
||||
|
||||
+9
-6
@@ -56,8 +56,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.10.0")
|
||||
set(TORCH_SUPPORTED_VERSION_ROCM "2.10.0")
|
||||
set(TORCH_SUPPORTED_VERSION_CUDA "2.11.0")
|
||||
set(TORCH_SUPPORTED_VERSION_ROCM "2.11.0")
|
||||
|
||||
#
|
||||
# Try to find python package with an executable that exactly matches
|
||||
@@ -225,8 +225,8 @@ if(VLLM_GPU_LANG STREQUAL "HIP")
|
||||
# Certain HIP functions are marked as [[nodiscard]], yet vllm ignores the result which generates
|
||||
# a lot of warnings that always mask real issues. Suppressing until this is properly addressed.
|
||||
#
|
||||
set(CMAKE_${VLLM_GPU_LANG}_FLAGS "${CMAKE_${VLLM_GPU_LANG}_FLAGS} -Wno-unused-result")
|
||||
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-unused-result")
|
||||
set(CMAKE_${VLLM_GPU_LANG}_FLAGS "${CMAKE_${VLLM_GPU_LANG}_FLAGS} -Wno-unused-result -Wno-unused-value")
|
||||
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-unused-result -Wno-unused-value")
|
||||
endif()
|
||||
|
||||
#
|
||||
@@ -299,6 +299,7 @@ set(VLLM_EXT_SRC
|
||||
"csrc/quantization/w8a8/int8/scaled_quant.cu"
|
||||
"csrc/quantization/w8a8/fp8/common.cu"
|
||||
"csrc/quantization/fused_kernels/fused_layernorm_dynamic_per_token_quant.cu"
|
||||
"csrc/quantization/fused_kernels/fused_silu_mul_block_quant.cu"
|
||||
"csrc/quantization/gguf/gguf_kernel.cu"
|
||||
"csrc/quantization/activation_kernels.cu"
|
||||
"csrc/cuda_utils_kernels.cu"
|
||||
@@ -306,6 +307,8 @@ set(VLLM_EXT_SRC
|
||||
"csrc/torch_bindings.cpp")
|
||||
|
||||
if(VLLM_GPU_LANG STREQUAL "CUDA")
|
||||
list(APPEND VLLM_EXT_SRC "csrc/minimax_reduce_rms_kernel.cu")
|
||||
|
||||
SET(CUTLASS_ENABLE_HEADERS_ONLY ON CACHE BOOL "Enable only the header library")
|
||||
|
||||
# Set CUTLASS_REVISION. Used for FetchContent. Also fixes some bogus messages when building.
|
||||
@@ -340,8 +343,7 @@ if(VLLM_GPU_LANG STREQUAL "CUDA")
|
||||
|
||||
list(APPEND VLLM_EXT_SRC
|
||||
"csrc/quantization/awq/gemm_kernels.cu"
|
||||
"csrc/cutlass_extensions/common.cpp"
|
||||
"csrc/quantization/fused_kernels/fused_silu_mul_block_quant.cu")
|
||||
"csrc/cutlass_extensions/common.cpp")
|
||||
|
||||
set_gencode_flags_for_srcs(
|
||||
SRCS "${VLLM_EXT_SRC}"
|
||||
@@ -1222,6 +1224,7 @@ endif()
|
||||
|
||||
# For CUDA we also build and ship some external projects.
|
||||
if (VLLM_GPU_LANG STREQUAL "CUDA")
|
||||
include(cmake/external_projects/deepgemm.cmake)
|
||||
include(cmake/external_projects/flashmla.cmake)
|
||||
include(cmake/external_projects/qutlass.cmake)
|
||||
|
||||
|
||||
@@ -23,47 +23,54 @@ For events, please visit [vllm.ai/events](https://vllm.ai/events) to join us.
|
||||
|
||||
vLLM is a fast and easy-to-use library for LLM inference and serving.
|
||||
|
||||
Originally developed in the [Sky Computing Lab](https://sky.cs.berkeley.edu) at UC Berkeley, vLLM has evolved into a community-driven project with contributions from both academia and industry.
|
||||
Originally developed in the [Sky Computing Lab](https://sky.cs.berkeley.edu) at UC Berkeley, vLLM has grown into one of the most active open-source AI projects built and maintained by a diverse community of many dozens of academic institutions and companies from over 2000 contributors.
|
||||
|
||||
vLLM is fast with:
|
||||
|
||||
- State-of-the-art serving throughput
|
||||
- Efficient management of attention key and value memory with [**PagedAttention**](https://blog.vllm.ai/2023/06/20/vllm.html)
|
||||
- Continuous batching of incoming requests
|
||||
- Fast model execution with CUDA/HIP graph
|
||||
- Quantizations: [GPTQ](https://arxiv.org/abs/2210.17323), [AWQ](https://arxiv.org/abs/2306.00978), [AutoRound](https://arxiv.org/abs/2309.05516), INT4, INT8, and FP8
|
||||
- Optimized CUDA kernels, including integration with FlashAttention and FlashInfer
|
||||
- Speculative decoding
|
||||
- Chunked prefill
|
||||
- Continuous batching of incoming requests, chunked prefill, prefix caching
|
||||
- Fast and flexible model execution with piecewise and full CUDA/HIP graphs
|
||||
- Quantization: FP8, MXFP8/MXFP4, NVFP4, INT8, INT4, GPTQ/AWQ, GGUF, compressed-tensors, ModelOpt, TorchAO, and [more](https://docs.vllm.ai/en/latest/features/quantization/index.html)
|
||||
- Optimized attention kernels including FlashAttention, FlashInfer, TRTLLM-GEN, FlashMLA, and Triton
|
||||
- Optimized GEMM/MoE kernels for various precisions using CUTLASS, TRTLLM-GEN, CuTeDSL
|
||||
- Speculative decoding including n-gram, suffix, EAGLE, DFlash
|
||||
- Automatic kernel generation and graph-level transformations using torch.compile
|
||||
- Disaggregated prefill, decode, and encode
|
||||
|
||||
vLLM is flexible and easy to use with:
|
||||
|
||||
- Seamless integration with popular Hugging Face models
|
||||
- High-throughput serving with various decoding algorithms, including *parallel sampling*, *beam search*, and more
|
||||
- Tensor, pipeline, data and expert parallelism support for distributed inference
|
||||
- Tensor, pipeline, data, expert, and context parallelism for distributed inference
|
||||
- Streaming outputs
|
||||
- OpenAI-compatible API server
|
||||
- Support for NVIDIA GPUs, AMD CPUs and GPUs, Intel CPUs and GPUs, PowerPC CPUs, Arm CPUs, and TPU. Additionally, support for diverse hardware plugins such as Intel Gaudi, IBM Spyre and Huawei Ascend.
|
||||
- Prefix caching support
|
||||
- Multi-LoRA support
|
||||
- Generation of structured outputs using xgrammar or guidance
|
||||
- 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.
|
||||
|
||||
vLLM seamlessly supports most popular open-source models on HuggingFace, including:
|
||||
vLLM seamlessly supports 200+ model architectures on HuggingFace, including:
|
||||
|
||||
- Transformer-like LLMs (e.g., Llama)
|
||||
- Mixture-of-Expert LLMs (e.g., Mixtral, Deepseek-V2 and V3)
|
||||
- Embedding Models (e.g., E5-Mistral)
|
||||
- Multi-modal LLMs (e.g., LLaVA)
|
||||
- Decoder-only LLMs (e.g., Llama, Qwen, Gemma)
|
||||
- Mixture-of-Expert LLMs (e.g., Mixtral, DeepSeek-V3, Qwen-MoE, GPT-OSS)
|
||||
- Hybrid attention and state-space models (e.g., Mamba, Qwen3.5)
|
||||
- Multi-modal models (e.g., LLaVA, Qwen-VL, Pixtral)
|
||||
- Embedding and retrieval models (e.g., E5-Mistral, GTE, ColBERT)
|
||||
- Reward and classification models (e.g., Qwen-Math)
|
||||
|
||||
Find the full list of supported models [here](https://docs.vllm.ai/en/latest/models/supported_models.html).
|
||||
|
||||
## Getting Started
|
||||
|
||||
Install vLLM with `pip` or [from source](https://docs.vllm.ai/en/latest/getting_started/installation/gpu/index.html#build-wheel-from-source):
|
||||
Install vLLM with [`uv`](https://docs.astral.sh/uv/) (recommended) or `pip`:
|
||||
|
||||
```bash
|
||||
pip install vllm
|
||||
uv pip install vllm
|
||||
```
|
||||
|
||||
Or [build from source](https://docs.vllm.ai/en/latest/getting_started/installation/gpu/index.html#build-wheel-from-source) for development.
|
||||
|
||||
Visit our [documentation](https://docs.vllm.ai/en/latest/) to learn more.
|
||||
|
||||
- [Installation](https://docs.vllm.ai/en/latest/getting_started/installation.html)
|
||||
|
||||
@@ -9,11 +9,12 @@ os.environ["VLLM_USE_DEEP_GEMM"] = "0"
|
||||
import torch
|
||||
|
||||
from vllm.benchmarks.lib.utils import default_vllm_config
|
||||
from vllm.model_executor.layers.quantization.utils.fp8_utils import (
|
||||
W8A8BlockFp8LinearOp,
|
||||
from vllm.model_executor.kernels.linear import (
|
||||
init_fp8_linear_kernel,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
GroupShape,
|
||||
create_fp8_quant_key,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.w8a8_utils import (
|
||||
CUTLASS_BLOCK_FP8_SUPPORTED,
|
||||
@@ -70,11 +71,15 @@ def build_w8a8_block_fp8_runner(M, N, K, block_size, device, use_cutlass):
|
||||
weight_group_shape = GroupShape(block_n, block_k)
|
||||
act_quant_group_shape = GroupShape(1, block_k) # Per-token, per-group quantization
|
||||
|
||||
linear_op = W8A8BlockFp8LinearOp(
|
||||
weight_group_shape=weight_group_shape,
|
||||
act_quant_group_shape=act_quant_group_shape,
|
||||
cutlass_block_fp8_supported=use_cutlass,
|
||||
use_aiter_and_is_supported=False,
|
||||
linear_op = init_fp8_linear_kernel(
|
||||
weight_quant_key=create_fp8_quant_key(
|
||||
static=True, group_shape=weight_group_shape
|
||||
),
|
||||
activation_quant_key=create_fp8_quant_key(
|
||||
static=False, group_shape=act_quant_group_shape
|
||||
),
|
||||
out_dtype=torch.get_default_dtype(),
|
||||
module_name="build_w8a8_block_fp8_runner",
|
||||
)
|
||||
|
||||
def run():
|
||||
|
||||
@@ -9,6 +9,7 @@ from vllm.model_executor.layers.fused_moe.moe_align_block_size import (
|
||||
moe_align_block_size,
|
||||
)
|
||||
from vllm.triton_utils import triton
|
||||
from vllm.utils.torch_utils import set_random_seed
|
||||
|
||||
|
||||
def get_topk_ids(num_tokens: int, num_experts: int, topk: int) -> torch.Tensor:
|
||||
@@ -44,7 +45,7 @@ configs = list(
|
||||
def benchmark(num_tokens, num_experts, topk, ep_size, provider):
|
||||
"""Benchmark function for Triton."""
|
||||
block_size = 256
|
||||
torch.cuda.manual_seed_all(0)
|
||||
set_random_seed(0)
|
||||
topk_ids = get_topk_ids(num_tokens, num_experts, topk)
|
||||
|
||||
e_map = None
|
||||
|
||||
@@ -20,7 +20,7 @@ import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from vllm.model_executor.layers.fused_moe.batched_deep_gemm_moe import (
|
||||
from vllm.model_executor.layers.fused_moe.experts.batched_deep_gemm_moe import (
|
||||
persistent_masked_m_silu_mul_quant,
|
||||
)
|
||||
from vllm.triton_utils import tl, triton
|
||||
|
||||
@@ -16,6 +16,7 @@ from vllm.utils.deep_gemm import (
|
||||
fp8_gemm_nt,
|
||||
per_block_cast_to_fp8,
|
||||
)
|
||||
from vllm.utils.torch_utils import set_random_seed
|
||||
|
||||
|
||||
def benchmark_shape(
|
||||
@@ -235,9 +236,7 @@ def run_benchmarks(verbose: bool = False):
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
# Set seeds for reproducibility
|
||||
torch.manual_seed(42)
|
||||
torch.cuda.manual_seed(42)
|
||||
|
||||
set_random_seed(42)
|
||||
# Define benchmark shapes (m, n, k)
|
||||
shapes = [
|
||||
(8, 4096, 7168),
|
||||
|
||||
@@ -1439,6 +1439,12 @@ async def main() -> None:
|
||||
action="store_true",
|
||||
help="Export summary to Excel file (optional)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stats-json-output",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Export per-request stats (ttft_ms, tpot_ms, etc.) to a JSON file",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-v",
|
||||
"--verbose",
|
||||
@@ -1651,6 +1657,19 @@ async def main() -> None:
|
||||
warmup_runtime_sec=warmup_runtime_sec,
|
||||
)
|
||||
|
||||
if args.stats_json_output is not None:
|
||||
# Export per-request metrics as a JSON array for downstream analysis.
|
||||
stats_data = [s._asdict() for s in client_metrics]
|
||||
logger.info(
|
||||
f"{Color.GREEN}Writing per-request stats JSON: "
|
||||
f"{args.stats_json_output}{Color.RESET}"
|
||||
)
|
||||
os.makedirs(
|
||||
os.path.dirname(os.path.abspath(args.stats_json_output)), exist_ok=True
|
||||
)
|
||||
with open(args.stats_json_output, "w") as f:
|
||||
json.dump(stats_data, f, indent=2)
|
||||
|
||||
if args.output_file is not None:
|
||||
# Write a JSON file with the updated conversations
|
||||
# The "assistant" content will contain the answers from the tested LLM
|
||||
|
||||
@@ -349,6 +349,7 @@ endif()
|
||||
set(VLLM_EXT_SRC
|
||||
"csrc/cpu/activation.cpp"
|
||||
"csrc/cpu/utils.cpp"
|
||||
"csrc/cpu/spec_decode_utils.cpp"
|
||||
"csrc/cpu/layernorm.cpp"
|
||||
"csrc/cpu/mla_decode.cpp"
|
||||
"csrc/cpu/pos_encoding.cpp"
|
||||
@@ -383,6 +384,7 @@ if (ENABLE_X86_ISA)
|
||||
"csrc/cpu/cpu_wna16.cpp"
|
||||
"csrc/cpu/cpu_fused_moe.cpp"
|
||||
"csrc/cpu/utils.cpp"
|
||||
"csrc/cpu/spec_decode_utils.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/dnnl_kernels.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp"
|
||||
@@ -395,6 +397,7 @@ if (ENABLE_X86_ISA)
|
||||
|
||||
set(VLLM_EXT_SRC_AVX2
|
||||
"csrc/cpu/utils.cpp"
|
||||
"csrc/cpu/spec_decode_utils.cpp"
|
||||
"csrc/cpu/cpu_attn.cpp"
|
||||
"csrc/cpu/torch_bindings.cpp"
|
||||
# TODO: Remove these files
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
include(FetchContent)
|
||||
|
||||
# If DEEPGEMM_SRC_DIR is set, DeepGEMM is built from that directory
|
||||
# instead of downloading.
|
||||
# It can be set as an environment variable or passed as a cmake argument.
|
||||
# The environment variable takes precedence.
|
||||
if (DEFINED ENV{DEEPGEMM_SRC_DIR})
|
||||
set(DEEPGEMM_SRC_DIR $ENV{DEEPGEMM_SRC_DIR})
|
||||
endif()
|
||||
|
||||
if(DEEPGEMM_SRC_DIR)
|
||||
FetchContent_Declare(
|
||||
deepgemm
|
||||
SOURCE_DIR ${DEEPGEMM_SRC_DIR}
|
||||
CONFIGURE_COMMAND ""
|
||||
BUILD_COMMAND ""
|
||||
)
|
||||
else()
|
||||
# This ref should be kept in sync with tools/install_deepgemm.sh
|
||||
FetchContent_Declare(
|
||||
deepgemm
|
||||
GIT_REPOSITORY https://github.com/deepseek-ai/DeepGEMM.git
|
||||
GIT_TAG 477618cd51baffca09c4b0b87e97c03fe827ef03
|
||||
GIT_SUBMODULES "third-party/cutlass" "third-party/fmt"
|
||||
GIT_PROGRESS TRUE
|
||||
CONFIGURE_COMMAND ""
|
||||
BUILD_COMMAND ""
|
||||
)
|
||||
endif()
|
||||
|
||||
# Use FetchContent_Populate (not MakeAvailable) to avoid processing
|
||||
# DeepGEMM's own CMakeLists.txt which has incompatible find_package calls.
|
||||
FetchContent_GetProperties(deepgemm)
|
||||
if(NOT deepgemm_POPULATED)
|
||||
FetchContent_Populate(deepgemm)
|
||||
endif()
|
||||
message(STATUS "DeepGEMM is available at ${deepgemm_SOURCE_DIR}")
|
||||
|
||||
# DeepGEMM requires CUDA 12.3+ for SM90, 12.9+ for SM100
|
||||
set(DEEPGEMM_SUPPORT_ARCHS)
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.3)
|
||||
list(APPEND DEEPGEMM_SUPPORT_ARCHS "9.0a")
|
||||
endif()
|
||||
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.9)
|
||||
list(APPEND DEEPGEMM_SUPPORT_ARCHS "10.0f")
|
||||
elseif(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.8)
|
||||
list(APPEND DEEPGEMM_SUPPORT_ARCHS "10.0a")
|
||||
endif()
|
||||
|
||||
cuda_archs_loose_intersection(DEEPGEMM_ARCHS
|
||||
"${DEEPGEMM_SUPPORT_ARCHS}" "${CUDA_ARCHS}")
|
||||
|
||||
if(DEEPGEMM_ARCHS)
|
||||
message(STATUS "DeepGEMM CUDA architectures: ${DEEPGEMM_ARCHS}")
|
||||
|
||||
find_package(CUDAToolkit REQUIRED)
|
||||
|
||||
#
|
||||
# Build the _C pybind11 extension from DeepGEMM's C++ source.
|
||||
# This is a CXX-only module — CUDA kernels are JIT-compiled at runtime.
|
||||
#
|
||||
Python_add_library(_deep_gemm_C MODULE WITH_SOABI
|
||||
"${deepgemm_SOURCE_DIR}/csrc/python_api.cpp")
|
||||
|
||||
# The pybind11 module name must be _C to match DeepGEMM's Python imports.
|
||||
set_target_properties(_deep_gemm_C PROPERTIES OUTPUT_NAME "_C")
|
||||
|
||||
target_compile_definitions(_deep_gemm_C PRIVATE
|
||||
"-DTORCH_EXTENSION_NAME=_C")
|
||||
|
||||
target_include_directories(_deep_gemm_C PRIVATE
|
||||
"${deepgemm_SOURCE_DIR}/csrc"
|
||||
"${deepgemm_SOURCE_DIR}/deep_gemm/include"
|
||||
"${deepgemm_SOURCE_DIR}/third-party/cutlass/include"
|
||||
"${deepgemm_SOURCE_DIR}/third-party/cutlass/tools/util/include"
|
||||
"${deepgemm_SOURCE_DIR}/third-party/fmt/include")
|
||||
|
||||
target_compile_options(_deep_gemm_C PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-std=c++17>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-O3>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-Wno-psabi>
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-Wno-deprecated-declarations>)
|
||||
|
||||
# torch_python is required because DeepGEMM uses pybind11 type casters
|
||||
# for at::Tensor (via PYBIND11_MODULE), unlike vLLM's own extensions which
|
||||
# use torch::Library custom ops.
|
||||
find_library(TORCH_PYTHON_LIBRARY torch_python
|
||||
PATHS "${TORCH_INSTALL_PREFIX}/lib"
|
||||
REQUIRED)
|
||||
|
||||
target_link_libraries(_deep_gemm_C PRIVATE
|
||||
torch ${TORCH_LIBRARIES} "${TORCH_PYTHON_LIBRARY}"
|
||||
CUDA::cudart CUDA::nvrtc)
|
||||
|
||||
# Install the shared library into the vendored package directory
|
||||
install(TARGETS _deep_gemm_C
|
||||
LIBRARY DESTINATION vllm/third_party/deep_gemm
|
||||
COMPONENT _deep_gemm_C)
|
||||
|
||||
#
|
||||
# Vendor DeepGEMM Python package files
|
||||
#
|
||||
install(FILES
|
||||
"${deepgemm_SOURCE_DIR}/deep_gemm/__init__.py"
|
||||
DESTINATION vllm/third_party/deep_gemm
|
||||
COMPONENT _deep_gemm_C)
|
||||
|
||||
install(DIRECTORY "${deepgemm_SOURCE_DIR}/deep_gemm/utils/"
|
||||
DESTINATION vllm/third_party/deep_gemm/utils
|
||||
COMPONENT _deep_gemm_C
|
||||
FILES_MATCHING PATTERN "*.py")
|
||||
|
||||
install(DIRECTORY "${deepgemm_SOURCE_DIR}/deep_gemm/testing/"
|
||||
DESTINATION vllm/third_party/deep_gemm/testing
|
||||
COMPONENT _deep_gemm_C
|
||||
FILES_MATCHING PATTERN "*.py")
|
||||
|
||||
install(DIRECTORY "${deepgemm_SOURCE_DIR}/deep_gemm/legacy/"
|
||||
DESTINATION vllm/third_party/deep_gemm/legacy
|
||||
COMPONENT _deep_gemm_C
|
||||
FILES_MATCHING PATTERN "*.py")
|
||||
|
||||
# Generate envs.py (normally generated by DeepGEMM's setup.py build step)
|
||||
file(WRITE "${CMAKE_CURRENT_BINARY_DIR}/deep_gemm_envs.py"
|
||||
"# Pre-installed environment variables\npersistent_envs = dict()\n")
|
||||
install(FILES "${CMAKE_CURRENT_BINARY_DIR}/deep_gemm_envs.py"
|
||||
DESTINATION vllm/third_party/deep_gemm
|
||||
RENAME envs.py
|
||||
COMPONENT _deep_gemm_C)
|
||||
|
||||
#
|
||||
# Install include files needed for JIT compilation at runtime.
|
||||
# The JIT compiler finds these relative to the package directory.
|
||||
#
|
||||
|
||||
# DeepGEMM's own CUDA headers
|
||||
install(DIRECTORY "${deepgemm_SOURCE_DIR}/deep_gemm/include/"
|
||||
DESTINATION vllm/third_party/deep_gemm/include
|
||||
COMPONENT _deep_gemm_C)
|
||||
|
||||
# CUTLASS and CuTe headers (vendored for JIT, separate from vLLM's CUTLASS)
|
||||
install(DIRECTORY "${deepgemm_SOURCE_DIR}/third-party/cutlass/include/"
|
||||
DESTINATION vllm/third_party/deep_gemm/include
|
||||
COMPONENT _deep_gemm_C)
|
||||
|
||||
else()
|
||||
message(STATUS "DeepGEMM will not compile: "
|
||||
"unsupported CUDA architecture ${CUDA_ARCHS}")
|
||||
# Create empty target so setup.py doesn't fail on unsupported systems
|
||||
add_custom_target(_deep_gemm_C)
|
||||
endif()
|
||||
@@ -39,7 +39,7 @@ else()
|
||||
FetchContent_Declare(
|
||||
vllm-flash-attn
|
||||
GIT_REPOSITORY https://github.com/vllm-project/flash-attention.git
|
||||
GIT_TAG c0ec424fd8a546d0cbbf4bf050bbcfe837c55afb
|
||||
GIT_TAG f5bc33cfc02c744d24a2e9d50e6db656de40611c
|
||||
GIT_PROGRESS TRUE
|
||||
# Don't share the vllm-flash-attn build between build types
|
||||
BINARY_DIR ${CMAKE_BINARY_DIR}/vllm-flash-attn
|
||||
@@ -87,18 +87,30 @@ endforeach()
|
||||
#
|
||||
add_custom_target(_vllm_fa4_cutedsl_C)
|
||||
|
||||
# Copy flash_attn/cute directory (needed for FA4) and transform imports
|
||||
# The cute directory uses flash_attn.cute imports internally, which we replace
|
||||
# with vllm.vllm_flash_attn.cute to match our package structure.
|
||||
install(CODE "
|
||||
file(GLOB_RECURSE CUTE_PY_FILES \"${vllm-flash-attn_SOURCE_DIR}/flash_attn/cute/*.py\")
|
||||
foreach(SRC_FILE \${CUTE_PY_FILES})
|
||||
file(RELATIVE_PATH REL_PATH \"${vllm-flash-attn_SOURCE_DIR}/flash_attn/cute\" \${SRC_FILE})
|
||||
set(DST_FILE \"\${CMAKE_INSTALL_PREFIX}/vllm/vllm_flash_attn/cute/\${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.vllm_flash_attn.cute\" FILE_CONTENTS \"\${FILE_CONTENTS}\")
|
||||
file(WRITE \${DST_FILE} \"\${FILE_CONTENTS}\")
|
||||
endforeach()
|
||||
" COMPONENT _vllm_fa4_cutedsl_C)
|
||||
# Install flash_attn/cute directory (needed for FA4).
|
||||
# When using a local source dir (VLLM_FLASH_ATTN_SRC_DIR), create a symlink
|
||||
# so edits to cute-dsl Python files take effect immediately without rebuilding.
|
||||
# Otherwise, copy files and transform flash_attn.cute imports to
|
||||
# vllm.vllm_flash_attn.cute to match our package structure.
|
||||
if(VLLM_FLASH_ATTN_SRC_DIR)
|
||||
install(CODE "
|
||||
set(LINK_TARGET \"${vllm-flash-attn_SOURCE_DIR}/flash_attn/cute\")
|
||||
set(LINK_NAME \"\${CMAKE_INSTALL_PREFIX}/vllm/vllm_flash_attn/cute\")
|
||||
file(MAKE_DIRECTORY \"\${CMAKE_INSTALL_PREFIX}/vllm/vllm_flash_attn\")
|
||||
file(REMOVE_RECURSE \"\${LINK_NAME}\")
|
||||
file(CREATE_LINK \"\${LINK_TARGET}\" \"\${LINK_NAME}\" SYMBOLIC)
|
||||
" COMPONENT _vllm_fa4_cutedsl_C)
|
||||
else()
|
||||
install(CODE "
|
||||
file(GLOB_RECURSE CUTE_PY_FILES \"${vllm-flash-attn_SOURCE_DIR}/flash_attn/cute/*.py\")
|
||||
foreach(SRC_FILE \${CUTE_PY_FILES})
|
||||
file(RELATIVE_PATH REL_PATH \"${vllm-flash-attn_SOURCE_DIR}/flash_attn/cute\" \${SRC_FILE})
|
||||
set(DST_FILE \"\${CMAKE_INSTALL_PREFIX}/vllm/vllm_flash_attn/cute/\${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.vllm_flash_attn.cute\" FILE_CONTENTS \"\${FILE_CONTENTS}\")
|
||||
file(WRITE \${DST_FILE} \"\${FILE_CONTENTS}\")
|
||||
endforeach()
|
||||
" COMPONENT _vllm_fa4_cutedsl_C)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
/*
|
||||
* Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
namespace vllm {
|
||||
namespace cuda_async {
|
||||
|
||||
__device__ __forceinline__ void cp_async_shared_global_16_cg(
|
||||
void* smem_ptr, const void* glob_ptr) {
|
||||
#if defined(USE_ROCM)
|
||||
*reinterpret_cast<int4*>(smem_ptr) = *reinterpret_cast<const int4*>(glob_ptr);
|
||||
#elif defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
uint32_t smem = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
|
||||
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n"
|
||||
:
|
||||
: "r"(smem), "l"(glob_ptr));
|
||||
#elif defined(__CUDA_ARCH__)
|
||||
*reinterpret_cast<int4*>(smem_ptr) = *reinterpret_cast<const int4*>(glob_ptr);
|
||||
#else
|
||||
(void)smem_ptr;
|
||||
(void)glob_ptr;
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_shared_global_ca(void* smem_ptr,
|
||||
const void* glob_ptr,
|
||||
int size_bytes) {
|
||||
#if defined(USE_ROCM)
|
||||
if (size_bytes == 4) {
|
||||
*reinterpret_cast<uint32_t*>(smem_ptr) =
|
||||
*reinterpret_cast<const uint32_t*>(glob_ptr);
|
||||
} else if (size_bytes == 8) {
|
||||
*reinterpret_cast<uint64_t*>(smem_ptr) =
|
||||
*reinterpret_cast<const uint64_t*>(glob_ptr);
|
||||
} else {
|
||||
*reinterpret_cast<int4*>(smem_ptr) =
|
||||
*reinterpret_cast<const int4*>(glob_ptr);
|
||||
}
|
||||
#elif defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
uint32_t smem = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
|
||||
if (size_bytes == 4) {
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 4;\n"
|
||||
:
|
||||
: "r"(smem), "l"(glob_ptr));
|
||||
} else if (size_bytes == 8) {
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 8;\n"
|
||||
:
|
||||
: "r"(smem), "l"(glob_ptr));
|
||||
} else {
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;\n"
|
||||
:
|
||||
: "r"(smem), "l"(glob_ptr));
|
||||
}
|
||||
#elif defined(__CUDA_ARCH__)
|
||||
if (size_bytes == 4) {
|
||||
*reinterpret_cast<uint32_t*>(smem_ptr) =
|
||||
*reinterpret_cast<const uint32_t*>(glob_ptr);
|
||||
} else if (size_bytes == 8) {
|
||||
*reinterpret_cast<uint64_t*>(smem_ptr) =
|
||||
*reinterpret_cast<const uint64_t*>(glob_ptr);
|
||||
} else {
|
||||
*reinterpret_cast<int4*>(smem_ptr) =
|
||||
*reinterpret_cast<const int4*>(glob_ptr);
|
||||
}
|
||||
#else
|
||||
(void)smem_ptr;
|
||||
(void)glob_ptr;
|
||||
(void)size_bytes;
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_commit_group() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 && !defined(USE_ROCM)
|
||||
asm volatile("cp.async.commit_group;\n" ::);
|
||||
#endif
|
||||
}
|
||||
|
||||
template <int n>
|
||||
__device__ __forceinline__ void cp_async_wait_group() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 && !defined(USE_ROCM)
|
||||
asm volatile("cp.async.wait_group %0;\n" : : "n"(n));
|
||||
#endif
|
||||
}
|
||||
|
||||
} // namespace cuda_async
|
||||
} // namespace vllm
|
||||
@@ -17,6 +17,22 @@ enum class Fp8KVCacheDataType {
|
||||
kFp8E5M2 = 2,
|
||||
};
|
||||
|
||||
inline Fp8KVCacheDataType get_fp8_kv_cache_data_type(
|
||||
const std::string& dtype_str) {
|
||||
// dtype_str refers to CacheDType at vllm.config.cache.CacheDType
|
||||
if (dtype_str == "auto" || dtype_str == "float16" ||
|
||||
dtype_str == "bfloat16") {
|
||||
// unquantized kv cache
|
||||
return Fp8KVCacheDataType::kAuto;
|
||||
} else if (dtype_str == "fp8" || dtype_str == "fp8_ds_mla" ||
|
||||
dtype_str == "fp8_e4m3") {
|
||||
return Fp8KVCacheDataType::kFp8E4M3;
|
||||
} else if (dtype_str == "fp8_e5m2") {
|
||||
return Fp8KVCacheDataType::kFp8E5M2;
|
||||
}
|
||||
TORCH_CHECK(false, "Unsupported fp8 kv cache data type: ", dtype_str);
|
||||
}
|
||||
|
||||
// fp8 vector types for quantization of kv cache
|
||||
template <>
|
||||
struct Vec<uint8_t, 1> {
|
||||
|
||||
+45
-24
@@ -91,9 +91,9 @@ void swap_blocks_batch(const torch::Tensor& src_ptrs,
|
||||
|
||||
if (n == 0) return;
|
||||
|
||||
const int64_t* src_data = src_ptrs.data_ptr<int64_t>();
|
||||
const int64_t* dst_data = dst_ptrs.data_ptr<int64_t>();
|
||||
const int64_t* size_data = sizes.data_ptr<int64_t>();
|
||||
int64_t* src_data = src_ptrs.mutable_data_ptr<int64_t>();
|
||||
int64_t* dst_data = dst_ptrs.mutable_data_ptr<int64_t>();
|
||||
int64_t* size_data = sizes.mutable_data_ptr<int64_t>();
|
||||
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
@@ -104,28 +104,49 @@ void swap_blocks_batch(const torch::Tensor& src_ptrs,
|
||||
static_assert(sizeof(CUdeviceptr) == sizeof(int64_t));
|
||||
static_assert(sizeof(size_t) == sizeof(int64_t));
|
||||
#if !defined(USE_ROCM) && defined(CUDA_VERSION) && CUDA_VERSION >= 12080
|
||||
CUmemcpyAttributes attr = {};
|
||||
attr.srcAccessOrder = CU_MEMCPY_SRC_ACCESS_ORDER_STREAM;
|
||||
size_t attrs_idx = 0;
|
||||
size_t fail_idx = 0;
|
||||
CUresult result = cuMemcpyBatchAsync(
|
||||
reinterpret_cast<CUdeviceptr*>(const_cast<int64_t*>(dst_data)),
|
||||
reinterpret_cast<CUdeviceptr*>(const_cast<int64_t*>(src_data)),
|
||||
reinterpret_cast<size_t*>(const_cast<int64_t*>(size_data)),
|
||||
static_cast<size_t>(n), &attr, &attrs_idx, 1, &fail_idx,
|
||||
static_cast<CUstream>(stream));
|
||||
TORCH_CHECK(result == CUDA_SUCCESS, "cuMemcpyBatchAsync failed at index ",
|
||||
fail_idx, " with error ", result);
|
||||
#else
|
||||
// Fallback for CUDA < 12.8 and ROCm: individual async copies.
|
||||
// cudaMemcpyDefault lets the driver infer direction from pointer types.
|
||||
for (int64_t i = 0; i < n; i++) {
|
||||
cudaMemcpyAsync(reinterpret_cast<void*>(dst_data[i]),
|
||||
reinterpret_cast<void*>(src_data[i]),
|
||||
static_cast<size_t>(size_data[i]), cudaMemcpyDefault,
|
||||
stream);
|
||||
}
|
||||
// Resolve cuMemcpyBatchAsync at runtime via cuGetProcAddress so that
|
||||
// binaries compiled with CUDA 12.8+ still work on older drivers, and
|
||||
// we avoid the CUDA 13.0 header remapping (#define to _v2 signature).
|
||||
// The function pointer is cached after the first call.
|
||||
using BatchFn =
|
||||
CUresult (*)(CUdeviceptr*, CUdeviceptr*, size_t*, size_t,
|
||||
CUmemcpyAttributes*, size_t*, size_t, size_t*, CUstream);
|
||||
static BatchFn batch_fn = []() -> BatchFn {
|
||||
CUdriverProcAddressQueryResult sym_status;
|
||||
void* fn_ptr = nullptr;
|
||||
CUresult res = cuGetProcAddress("cuMemcpyBatchAsync", &fn_ptr, 12080,
|
||||
CU_GET_PROC_ADDRESS_DEFAULT, &sym_status);
|
||||
if (res != CUDA_SUCCESS || fn_ptr == nullptr) {
|
||||
return nullptr;
|
||||
}
|
||||
return reinterpret_cast<BatchFn>(fn_ptr);
|
||||
}();
|
||||
|
||||
if (batch_fn != nullptr) {
|
||||
CUmemcpyAttributes attr = {};
|
||||
attr.srcAccessOrder = CU_MEMCPY_SRC_ACCESS_ORDER_STREAM;
|
||||
size_t attrs_idx = 0;
|
||||
size_t fail_idx = 0;
|
||||
CUresult result = batch_fn(reinterpret_cast<CUdeviceptr*>(dst_data),
|
||||
reinterpret_cast<CUdeviceptr*>(src_data),
|
||||
reinterpret_cast<size_t*>(size_data),
|
||||
static_cast<size_t>(n), &attr, &attrs_idx, 1,
|
||||
&fail_idx, static_cast<CUstream>(stream));
|
||||
TORCH_CHECK(result == CUDA_SUCCESS, "cuMemcpyBatchAsync failed at index ",
|
||||
fail_idx, " with error ", result);
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
// Fallback for CUDA < 12.8, older drivers, and ROCm:
|
||||
// individual async copies.
|
||||
// cudaMemcpyDefault lets the driver infer direction from pointer types.
|
||||
for (int64_t i = 0; i < n; i++) {
|
||||
cudaMemcpyAsync(reinterpret_cast<void*>(dst_data[i]),
|
||||
reinterpret_cast<void*>(src_data[i]),
|
||||
static_cast<size_t>(size_data[i]), cudaMemcpyDefault,
|
||||
stream);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
namespace vllm {
|
||||
|
||||
@@ -53,7 +53,7 @@ class TileGemm82 {
|
||||
const int64_t ldb, const int64_t ldc,
|
||||
const int32_t block_size, const int32_t dynamic_k_size,
|
||||
const bool accum_c) {
|
||||
static_assert(0 < M <= 8);
|
||||
static_assert(0 < M && M <= 8);
|
||||
using load_vec_t = typename VecTypeTrait<kv_cache_t>::vec_t;
|
||||
|
||||
kv_cache_t* __restrict__ curr_b_0 = b_tile;
|
||||
|
||||
@@ -68,7 +68,7 @@ class TileGemm161 {
|
||||
const int64_t ldb, const int64_t ldc,
|
||||
const int32_t block_size, const int32_t dynamic_k_size,
|
||||
const bool accum_c) {
|
||||
static_assert(0 < M <= 16);
|
||||
static_assert(0 < M && M <= 16);
|
||||
using load_vec_t = typename VecTypeTrait<kv_cache_t>::vec_t;
|
||||
|
||||
kv_cache_t* __restrict__ curr_b_0 = b_tile;
|
||||
|
||||
@@ -39,7 +39,7 @@ class TileGemm82 {
|
||||
|
||||
template <int32_t M>
|
||||
static void gemm_micro(DEFINE_CPU_MICRO_GEMM_PARAMS) {
|
||||
static_assert(0 < M <= 8);
|
||||
static_assert(0 < M && M <= 8);
|
||||
using load_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
|
||||
scalar_t* __restrict__ curr_b_0 = b_ptr;
|
||||
|
||||
@@ -0,0 +1,409 @@
|
||||
#include "cpu_types.hpp"
|
||||
|
||||
#include <algorithm>
|
||||
|
||||
namespace cpu_utils {
|
||||
|
||||
void eagle_prepare_inputs_padded_kernel_impl(
|
||||
const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& valid_sampled_tokens_count,
|
||||
const torch::Tensor& query_start_loc_gpu,
|
||||
torch::Tensor& token_indices_to_sample,
|
||||
torch::Tensor& num_rejected_tokens_gpu, const int64_t num_reqs) {
|
||||
const int64_t* cu_draft_ptr = cu_num_draft_tokens.data_ptr<int64_t>();
|
||||
const int64_t* valid_count_ptr =
|
||||
valid_sampled_tokens_count.data_ptr<int64_t>();
|
||||
const int32_t* query_loc_ptr = query_start_loc_gpu.data_ptr<int32_t>();
|
||||
int32_t* indices_out_ptr = token_indices_to_sample.data_ptr<int32_t>();
|
||||
int64_t* rejected_out_ptr = num_rejected_tokens_gpu.data_ptr<int64_t>();
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < num_reqs; ++req_idx) {
|
||||
int64_t start_idx = req_idx == 0 ? 0 : cu_draft_ptr[req_idx - 1];
|
||||
int64_t num_draft_tokens = cu_draft_ptr[req_idx] - start_idx;
|
||||
int64_t num_valid_tokens = valid_count_ptr[req_idx];
|
||||
|
||||
int64_t num_rejected = 0;
|
||||
if (num_draft_tokens > 0) {
|
||||
num_rejected = num_draft_tokens + 1 - num_valid_tokens;
|
||||
}
|
||||
|
||||
int32_t q_last_tok_idx = query_loc_ptr[req_idx + 1] - 1;
|
||||
int32_t index_to_sample = q_last_tok_idx - num_rejected;
|
||||
|
||||
indices_out_ptr[req_idx] = index_to_sample;
|
||||
rejected_out_ptr[req_idx] = num_rejected;
|
||||
}
|
||||
}
|
||||
|
||||
void eagle_prepare_next_token_padded_kernel_impl(
|
||||
const torch::Tensor& sampled_token_ids,
|
||||
const torch::Tensor& discard_request_mask,
|
||||
const torch::Tensor& backup_next_token_ids, torch::Tensor& next_token_ids,
|
||||
torch::Tensor& valid_sampled_tokens_count, const int64_t vocab_size,
|
||||
const int64_t num_sampled_tokens_per_req, const int64_t num_reqs) {
|
||||
const int64_t* sampled_ids_ptr = sampled_token_ids.data_ptr<int64_t>();
|
||||
const bool* discard_mask_ptr = discard_request_mask.data_ptr<bool>();
|
||||
const int64_t* backup_ids_ptr = backup_next_token_ids.data_ptr<int64_t>();
|
||||
int64_t* next_ids_out_ptr = next_token_ids.data_ptr<int64_t>();
|
||||
int64_t* valid_count_out_ptr = valid_sampled_tokens_count.data_ptr<int64_t>();
|
||||
|
||||
const int64_t stride = sampled_token_ids.stride(0);
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < num_reqs; ++req_idx) {
|
||||
const int64_t* row_ptr = sampled_ids_ptr + req_idx * stride;
|
||||
int64_t valid_count = 0;
|
||||
int64_t last_valid_token = -1;
|
||||
|
||||
for (int64_t pos = 0; pos < num_sampled_tokens_per_req; ++pos) {
|
||||
int64_t token = row_ptr[pos];
|
||||
if (token != -1 && token < vocab_size) {
|
||||
valid_count++;
|
||||
last_valid_token = token;
|
||||
}
|
||||
}
|
||||
|
||||
bool discard = discard_mask_ptr[req_idx];
|
||||
if (discard) {
|
||||
next_ids_out_ptr[req_idx] = backup_ids_ptr[req_idx];
|
||||
valid_count_out_ptr[req_idx] = 0;
|
||||
} else {
|
||||
next_ids_out_ptr[req_idx] =
|
||||
(valid_count > 0) ? last_valid_token : backup_ids_ptr[req_idx];
|
||||
valid_count_out_ptr[req_idx] = valid_count;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void eagle_step_slot_mapping_metadata_kernel_impl(
|
||||
const torch::Tensor& positions, const torch::Tensor& block_table,
|
||||
torch::Tensor& seq_lens, torch::Tensor& out_clamped_positions,
|
||||
torch::Tensor& out_slot_mapping, const int64_t block_size,
|
||||
const int64_t max_model_len, const int64_t PAD_ID) {
|
||||
const int64_t batch_size = positions.size(0);
|
||||
const int64_t input_batch_size = out_slot_mapping.size(0);
|
||||
|
||||
const int64_t* pos_ptr = positions.data_ptr<int64_t>();
|
||||
const int32_t* bt_ptr = block_table.data_ptr<int32_t>();
|
||||
int32_t* seq_lens_ptr = seq_lens.data_ptr<int32_t>();
|
||||
int64_t* out_clamped_ptr = out_clamped_positions.data_ptr<int64_t>();
|
||||
int64_t* out_slot_ptr = out_slot_mapping.data_ptr<int64_t>();
|
||||
|
||||
const int64_t bt_stride = block_table.stride(0);
|
||||
const int64_t n_blocks_per_req = block_table.size(1);
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < input_batch_size; ++req_idx) {
|
||||
if (req_idx >= batch_size) {
|
||||
out_slot_ptr[req_idx] = PAD_ID;
|
||||
continue;
|
||||
}
|
||||
|
||||
int64_t position = pos_ptr[req_idx];
|
||||
int64_t new_position = position + 1;
|
||||
bool exceeds_max = new_position >= max_model_len;
|
||||
int64_t clamped_position = exceeds_max ? 0 : new_position;
|
||||
|
||||
out_clamped_ptr[req_idx] = clamped_position;
|
||||
|
||||
int64_t block_number = clamped_position / block_size;
|
||||
block_number = std::min(block_number, n_blocks_per_req - 1);
|
||||
int32_t block_id = bt_ptr[req_idx * bt_stride + block_number];
|
||||
int64_t slot_id = block_id * block_size + (clamped_position % block_size);
|
||||
out_slot_ptr[req_idx] = exceeds_max ? PAD_ID : slot_id;
|
||||
|
||||
int32_t seq_len = seq_lens_ptr[req_idx];
|
||||
int32_t new_seq_len = exceeds_max ? 1 : (seq_len + 1);
|
||||
new_seq_len = std::min(new_seq_len, static_cast<int32_t>(max_model_len));
|
||||
seq_lens_ptr[req_idx] = new_seq_len;
|
||||
}
|
||||
}
|
||||
|
||||
void copy_and_expand_eagle_inputs_kernel_impl(
|
||||
const torch::Tensor& target_token_ids,
|
||||
const torch::Tensor& target_positions, const torch::Tensor& next_token_ids,
|
||||
torch::Tensor& out_input_ids, torch::Tensor& out_positions,
|
||||
torch::Tensor& out_is_rejected_token_mask,
|
||||
torch::Tensor& out_is_masked_token_mask,
|
||||
torch::Tensor& out_new_token_indices,
|
||||
torch::Tensor& out_hidden_state_mapping,
|
||||
const torch::Tensor& query_start_loc, const torch::Tensor& query_end_loc,
|
||||
const int64_t padding_token_id, const int64_t parallel_drafting_token_id,
|
||||
const int64_t total_input_tokens,
|
||||
const int64_t num_padding_slots_per_request, const bool shift_input_ids) {
|
||||
const int64_t num_reqs = query_end_loc.size(0);
|
||||
|
||||
const int64_t* target_ids_ptr = target_token_ids.data_ptr<int64_t>();
|
||||
const int64_t* target_pos_ptr = target_positions.data_ptr<int64_t>();
|
||||
const int64_t* next_ids_ptr = next_token_ids.data_ptr<int64_t>();
|
||||
const int32_t* query_start_ptr = query_start_loc.data_ptr<int32_t>();
|
||||
const int32_t* query_end_ptr = query_end_loc.data_ptr<int32_t>();
|
||||
|
||||
int64_t* out_ids_ptr = out_input_ids.data_ptr<int64_t>();
|
||||
int64_t* out_pos_ptr = out_positions.data_ptr<int64_t>();
|
||||
bool* out_rej_mask_ptr = out_is_rejected_token_mask.data_ptr<bool>();
|
||||
bool* out_mask_ptr = out_is_masked_token_mask.data_ptr<bool>();
|
||||
int32_t* out_new_idx_ptr = out_new_token_indices.data_ptr<int32_t>();
|
||||
int32_t* out_hidden_map_ptr = out_hidden_state_mapping.data_ptr<int32_t>();
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < num_reqs; ++req_idx) {
|
||||
int32_t q_start = query_start_ptr[req_idx];
|
||||
int32_t next_q_start = query_start_ptr[req_idx + 1];
|
||||
int32_t q_end = query_end_ptr[req_idx];
|
||||
|
||||
int64_t num_valid_tokens =
|
||||
shift_input_ids ? (q_end - q_start) : (q_end - q_start + 1);
|
||||
int64_t input_offset = shift_input_ids ? 1 : 0;
|
||||
|
||||
int64_t out_start = q_start + req_idx * (num_padding_slots_per_request -
|
||||
(shift_input_ids ? 1 : 0));
|
||||
int64_t num_rejected = next_q_start - q_end - 1;
|
||||
int64_t total_output_tokens =
|
||||
num_valid_tokens + num_padding_slots_per_request + num_rejected;
|
||||
|
||||
int64_t start_pos = target_pos_ptr[q_start];
|
||||
int64_t bonus_token = next_ids_ptr[req_idx];
|
||||
|
||||
for (int64_t j = 0; j < total_output_tokens; ++j) {
|
||||
int64_t out_idx = out_start + j;
|
||||
bool is_valid = j < num_valid_tokens;
|
||||
bool is_bonus = j == num_valid_tokens;
|
||||
bool is_parallel = (j > num_valid_tokens) &&
|
||||
(j < num_valid_tokens + num_padding_slots_per_request);
|
||||
bool is_rejected = j >= num_valid_tokens + num_padding_slots_per_request;
|
||||
|
||||
int64_t in_idx =
|
||||
std::min(static_cast<int64_t>(q_start + input_offset + j),
|
||||
total_input_tokens - 1);
|
||||
|
||||
int64_t token_id = padding_token_id;
|
||||
if (is_valid)
|
||||
token_id = target_ids_ptr[in_idx];
|
||||
else if (is_bonus)
|
||||
token_id = bonus_token;
|
||||
else if (is_parallel)
|
||||
token_id = parallel_drafting_token_id;
|
||||
|
||||
out_ids_ptr[out_idx] = token_id;
|
||||
out_pos_ptr[out_idx] = is_rejected ? 0 : (start_pos + j);
|
||||
out_rej_mask_ptr[out_idx] = is_rejected;
|
||||
out_mask_ptr[out_idx] = is_parallel;
|
||||
|
||||
if (is_bonus || is_parallel) {
|
||||
int64_t new_token_local_idx = j - num_valid_tokens;
|
||||
int64_t new_token_out_idx =
|
||||
req_idx * num_padding_slots_per_request + new_token_local_idx;
|
||||
out_new_idx_ptr[new_token_out_idx] = out_idx;
|
||||
}
|
||||
}
|
||||
|
||||
if (shift_input_ids) {
|
||||
int64_t n_input = next_q_start - q_start;
|
||||
for (int64_t j = 0; j < n_input; ++j) {
|
||||
out_hidden_map_ptr[q_start + j] = out_start + j;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void rejection_greedy_sample_kernel_impl(
|
||||
torch::Tensor& output_token_ids, const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& draft_token_ids, const torch::Tensor& target_argmax,
|
||||
const torch::Tensor& bonus_token_ids,
|
||||
const std::optional<torch::Tensor>& is_greedy, const int64_t max_spec_len) {
|
||||
const int64_t batch_size = cu_num_draft_tokens.size(0);
|
||||
|
||||
int64_t* out_ptr = output_token_ids.data_ptr<int64_t>();
|
||||
const int64_t* cu_draft_ptr = cu_num_draft_tokens.data_ptr<int64_t>();
|
||||
const int64_t* draft_ids_ptr = draft_token_ids.data_ptr<int64_t>();
|
||||
const int64_t* target_argmax_ptr = target_argmax.data_ptr<int64_t>();
|
||||
const int64_t* bonus_ids_ptr = bonus_token_ids.data_ptr<int64_t>();
|
||||
const bool* greedy_ptr =
|
||||
is_greedy.has_value() ? is_greedy.value().data_ptr<bool>() : nullptr;
|
||||
|
||||
const int64_t out_stride = output_token_ids.stride(0);
|
||||
const int64_t bonus_stride = bonus_token_ids.stride(0);
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < batch_size; ++req_idx) {
|
||||
if (greedy_ptr && !greedy_ptr[req_idx]) continue;
|
||||
|
||||
int64_t start_idx = req_idx == 0 ? 0 : cu_draft_ptr[req_idx - 1];
|
||||
int64_t end_idx = cu_draft_ptr[req_idx];
|
||||
int64_t num_draft_tokens = end_idx - start_idx;
|
||||
|
||||
bool rejected = false;
|
||||
for (int64_t pos = 0; pos < num_draft_tokens; ++pos) {
|
||||
int64_t target_id = target_argmax_ptr[start_idx + pos];
|
||||
out_ptr[req_idx * out_stride + pos] = target_id;
|
||||
|
||||
if (draft_ids_ptr[start_idx + pos] != target_id) {
|
||||
rejected = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (!rejected) {
|
||||
out_ptr[req_idx * out_stride + num_draft_tokens] =
|
||||
bonus_ids_ptr[req_idx * bonus_stride];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void rejection_random_sample_kernel_impl(
|
||||
torch::Tensor& output_token_ids, const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& draft_token_ids,
|
||||
const std::optional<torch::Tensor>& draft_probs,
|
||||
const torch::Tensor& target_probs, const torch::Tensor& bonus_token_ids,
|
||||
const torch::Tensor& recovered_token_ids,
|
||||
const torch::Tensor& uniform_probs,
|
||||
const std::optional<torch::Tensor>& is_greedy, const int64_t max_spec_len,
|
||||
const int64_t vocab_size, const bool no_draft_probs) {
|
||||
const int64_t batch_size = cu_num_draft_tokens.size(0);
|
||||
|
||||
int64_t* out_ptr = output_token_ids.data_ptr<int64_t>();
|
||||
const int64_t* cu_draft_ptr = cu_num_draft_tokens.data_ptr<int64_t>();
|
||||
const int64_t* draft_ids_ptr = draft_token_ids.data_ptr<int64_t>();
|
||||
const float* draft_probs_ptr =
|
||||
no_draft_probs ? nullptr : draft_probs.value().data_ptr<float>();
|
||||
const float* target_probs_ptr = target_probs.data_ptr<float>();
|
||||
const int64_t* bonus_ids_ptr = bonus_token_ids.data_ptr<int64_t>();
|
||||
const int64_t* recovered_ids_ptr = recovered_token_ids.data_ptr<int64_t>();
|
||||
const float* uniform_probs_ptr = uniform_probs.data_ptr<float>();
|
||||
const bool* greedy_ptr =
|
||||
is_greedy.has_value() ? is_greedy.value().data_ptr<bool>() : nullptr;
|
||||
|
||||
const int64_t out_stride = output_token_ids.stride(0);
|
||||
const int64_t bonus_stride = bonus_token_ids.stride(0);
|
||||
const int64_t target_stride = target_probs.stride(0);
|
||||
const int64_t draft_probs_stride =
|
||||
no_draft_probs ? 0 : draft_probs.value().stride(0);
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < batch_size; ++req_idx) {
|
||||
if (greedy_ptr && greedy_ptr[req_idx]) continue;
|
||||
|
||||
int64_t start_idx = req_idx == 0 ? 0 : cu_draft_ptr[req_idx - 1];
|
||||
int64_t end_idx = cu_draft_ptr[req_idx];
|
||||
int64_t num_draft_tokens = end_idx - start_idx;
|
||||
|
||||
bool rejected = false;
|
||||
for (int64_t pos = 0; pos < num_draft_tokens; ++pos) {
|
||||
int64_t token_idx = start_idx + pos;
|
||||
int64_t draft_id = draft_ids_ptr[token_idx];
|
||||
|
||||
float p = target_probs_ptr[token_idx * target_stride + draft_id];
|
||||
float q =
|
||||
no_draft_probs
|
||||
? 1.0f
|
||||
: draft_probs_ptr[token_idx * draft_probs_stride + draft_id];
|
||||
float uniform_p = uniform_probs_ptr[token_idx];
|
||||
|
||||
float ratio = (q > 0.0f) ? (p / q) : 0.0f;
|
||||
|
||||
if (ratio >= uniform_p) {
|
||||
out_ptr[req_idx * out_stride + pos] = draft_id;
|
||||
} else {
|
||||
out_ptr[req_idx * out_stride + pos] = recovered_ids_ptr[token_idx];
|
||||
rejected = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (!rejected) {
|
||||
out_ptr[req_idx * out_stride + num_draft_tokens] =
|
||||
bonus_ids_ptr[req_idx * bonus_stride];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void expand_kernel_impl(torch::Tensor& output, const torch::Tensor& input,
|
||||
const torch::Tensor& cu_num_tokens,
|
||||
const int64_t replace_from, const int64_t replace_to) {
|
||||
const int64_t batch_size = cu_num_tokens.size(0);
|
||||
const int64_t* cu_tokens_ptr = cu_num_tokens.data_ptr<int64_t>();
|
||||
|
||||
int64_t* out_ptr = output.data_ptr<int64_t>();
|
||||
const int64_t* in_ptr = input.data_ptr<int64_t>();
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < batch_size; ++req_idx) {
|
||||
int64_t start_idx = req_idx == 0 ? 0 : cu_tokens_ptr[req_idx - 1];
|
||||
int64_t end_idx = cu_tokens_ptr[req_idx];
|
||||
int64_t val = in_ptr[req_idx];
|
||||
|
||||
if (val == replace_from) {
|
||||
val = replace_to;
|
||||
}
|
||||
|
||||
for (int64_t i = start_idx; i < end_idx; ++i) {
|
||||
out_ptr[i] = val;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void sample_recovered_tokens_kernel_impl(
|
||||
torch::Tensor& output_token_ids, const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& draft_token_ids,
|
||||
const std::optional<torch::Tensor>& draft_probs,
|
||||
const torch::Tensor& target_probs, const torch::Tensor& inv_q,
|
||||
const int64_t vocab_size, const bool no_draft_probs) {
|
||||
const int64_t batch_size = cu_num_draft_tokens.size(0);
|
||||
|
||||
int64_t* out_ptr = output_token_ids.data_ptr<int64_t>();
|
||||
const int64_t* cu_draft_ptr = cu_num_draft_tokens.data_ptr<int64_t>();
|
||||
const int64_t* draft_ids_ptr = draft_token_ids.data_ptr<int64_t>();
|
||||
const float* draft_probs_ptr =
|
||||
no_draft_probs ? nullptr : draft_probs.value().data_ptr<float>();
|
||||
const float* target_probs_ptr = target_probs.data_ptr<float>();
|
||||
const float* inv_q_ptr = inv_q.data_ptr<float>();
|
||||
|
||||
const int64_t target_stride = target_probs.stride(0);
|
||||
const int64_t draft_probs_stride =
|
||||
no_draft_probs ? 0 : draft_probs.value().stride(0);
|
||||
const int64_t inv_q_stride = inv_q.stride(0);
|
||||
|
||||
#pragma omp parallel for
|
||||
for (int64_t req_idx = 0; req_idx < batch_size; ++req_idx) {
|
||||
int64_t start_idx = req_idx == 0 ? 0 : cu_draft_ptr[req_idx - 1];
|
||||
int64_t end_idx = cu_draft_ptr[req_idx];
|
||||
int64_t num_draft_tokens = end_idx - start_idx;
|
||||
|
||||
const float* req_inv_q = inv_q_ptr + req_idx * inv_q_stride;
|
||||
|
||||
for (int64_t pos = 0; pos < num_draft_tokens; ++pos) {
|
||||
int64_t token_idx = start_idx + pos;
|
||||
int64_t draft_id = draft_ids_ptr[token_idx];
|
||||
|
||||
const float* token_target_probs =
|
||||
target_probs_ptr + token_idx * target_stride;
|
||||
const float* token_draft_probs =
|
||||
no_draft_probs ? nullptr
|
||||
: (draft_probs_ptr + token_idx * draft_probs_stride);
|
||||
|
||||
int64_t best_id = 0;
|
||||
float best_val = -1.0f;
|
||||
|
||||
for (int64_t v = 0; v < vocab_size; ++v) {
|
||||
float prob = token_target_probs[v];
|
||||
if (no_draft_probs) {
|
||||
if (v == draft_id) prob = 0.0f;
|
||||
} else {
|
||||
float diff = prob - token_draft_probs[v];
|
||||
prob = diff > 0.0f ? diff : 0.0f;
|
||||
}
|
||||
|
||||
float val = prob * req_inv_q[v];
|
||||
if (val > best_val) {
|
||||
best_val = val;
|
||||
best_id = v;
|
||||
}
|
||||
}
|
||||
out_ptr[token_idx] = best_id;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace cpu_utils
|
||||
@@ -138,6 +138,61 @@ void compute_slot_mapping_kernel_impl(const torch::Tensor query_start_loc,
|
||||
torch::Tensor slot_mapping,
|
||||
const int64_t block_size);
|
||||
|
||||
namespace cpu_utils {
|
||||
void eagle_prepare_inputs_padded_kernel_impl(
|
||||
const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& valid_sampled_tokens_count,
|
||||
const torch::Tensor& query_start_loc_gpu,
|
||||
torch::Tensor& token_indices_to_sample,
|
||||
torch::Tensor& num_rejected_tokens_gpu, const int64_t num_reqs);
|
||||
void eagle_prepare_next_token_padded_kernel_impl(
|
||||
const torch::Tensor& sampled_token_ids,
|
||||
const torch::Tensor& discard_request_mask,
|
||||
const torch::Tensor& backup_next_token_ids, torch::Tensor& next_token_ids,
|
||||
torch::Tensor& valid_sampled_tokens_count, const int64_t vocab_size,
|
||||
const int64_t num_sampled_tokens_per_req, const int64_t num_reqs);
|
||||
void eagle_step_slot_mapping_metadata_kernel_impl(
|
||||
const torch::Tensor& positions, const torch::Tensor& block_table,
|
||||
torch::Tensor& seq_lens, torch::Tensor& out_clamped_positions,
|
||||
torch::Tensor& out_slot_mapping, const int64_t block_size,
|
||||
const int64_t max_model_len, const int64_t PAD_ID);
|
||||
void copy_and_expand_eagle_inputs_kernel_impl(
|
||||
const torch::Tensor& target_token_ids,
|
||||
const torch::Tensor& target_positions, const torch::Tensor& next_token_ids,
|
||||
torch::Tensor& out_input_ids, torch::Tensor& out_positions,
|
||||
torch::Tensor& out_is_rejected_token_mask,
|
||||
torch::Tensor& out_is_masked_token_mask,
|
||||
torch::Tensor& out_new_token_indices,
|
||||
torch::Tensor& out_hidden_state_mapping,
|
||||
const torch::Tensor& query_start_loc, const torch::Tensor& query_end_loc,
|
||||
const int64_t padding_token_id, const int64_t parallel_drafting_token_id,
|
||||
const int64_t total_input_tokens,
|
||||
const int64_t num_padding_slots_per_request, const bool shift_input_ids);
|
||||
void rejection_greedy_sample_kernel_impl(
|
||||
torch::Tensor& output_token_ids, const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& draft_token_ids, const torch::Tensor& target_argmax,
|
||||
const torch::Tensor& bonus_token_ids,
|
||||
const std::optional<torch::Tensor>& is_greedy, const int64_t max_spec_len);
|
||||
void rejection_random_sample_kernel_impl(
|
||||
torch::Tensor& output_token_ids, const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& draft_token_ids,
|
||||
const std::optional<torch::Tensor>& draft_probs,
|
||||
const torch::Tensor& target_probs, const torch::Tensor& bonus_token_ids,
|
||||
const torch::Tensor& recovered_token_ids,
|
||||
const torch::Tensor& uniform_probs,
|
||||
const std::optional<torch::Tensor>& is_greedy, const int64_t max_spec_len,
|
||||
const int64_t vocab_size, const bool no_draft_probs);
|
||||
void expand_kernel_impl(torch::Tensor& output, const torch::Tensor& input,
|
||||
const torch::Tensor& cu_num_tokens,
|
||||
const int64_t replace_from, const int64_t replace_to);
|
||||
void sample_recovered_tokens_kernel_impl(
|
||||
torch::Tensor& output_token_ids, const torch::Tensor& cu_num_draft_tokens,
|
||||
const torch::Tensor& draft_token_ids,
|
||||
const std::optional<torch::Tensor>& draft_probs,
|
||||
const torch::Tensor& target_probs, const torch::Tensor& inv_q,
|
||||
const int64_t vocab_size, const bool no_draft_probs);
|
||||
} // namespace cpu_utils
|
||||
|
||||
TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
// vLLM custom ops
|
||||
|
||||
@@ -363,6 +418,70 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"positions, Tensor block_table, Tensor(a3!) slot_mapping, SymInt "
|
||||
"block_size) -> ()",
|
||||
&compute_slot_mapping_kernel_impl);
|
||||
|
||||
// Speculative decoding kernels
|
||||
ops.def(
|
||||
"eagle_prepare_inputs_padded_kernel_impl(Tensor cu_num_draft_tokens, "
|
||||
"Tensor valid_sampled_tokens_count, Tensor query_start_loc_gpu, "
|
||||
"Tensor(a3!) token_indices_to_sample, "
|
||||
"Tensor(a4!) num_rejected_tokens_gpu, "
|
||||
"SymInt num_reqs) -> ()",
|
||||
&cpu_utils::eagle_prepare_inputs_padded_kernel_impl);
|
||||
ops.def(
|
||||
"eagle_prepare_next_token_padded_kernel_impl("
|
||||
"Tensor sampled_token_ids, Tensor discard_request_mask, "
|
||||
"Tensor backup_next_token_ids, Tensor(a3!) next_token_ids, "
|
||||
"Tensor(a4!) valid_sampled_tokens_count, SymInt vocab_size, "
|
||||
"SymInt num_sampled_tokens_per_req, SymInt num_reqs) -> ()",
|
||||
&cpu_utils::eagle_prepare_next_token_padded_kernel_impl);
|
||||
ops.def(
|
||||
"eagle_step_slot_mapping_metadata_kernel_impl("
|
||||
"Tensor positions, Tensor block_table, Tensor(a2!) seq_lens, "
|
||||
"Tensor(a3!) out_clamped_positions, Tensor(a4!) out_slot_mapping, "
|
||||
"SymInt block_size, SymInt max_model_len, SymInt PAD_ID) -> ()",
|
||||
&cpu_utils::eagle_step_slot_mapping_metadata_kernel_impl);
|
||||
ops.def(
|
||||
"copy_and_expand_eagle_inputs_kernel_impl("
|
||||
"Tensor target_token_ids, Tensor target_positions, "
|
||||
"Tensor next_token_ids, Tensor(a3!) out_input_ids, "
|
||||
"Tensor(a4!) out_positions, "
|
||||
"Tensor(a5!) out_is_rejected_token_mask, "
|
||||
"Tensor(a6!) out_is_masked_token_mask, "
|
||||
"Tensor(a7!) out_new_token_indices, "
|
||||
"Tensor(a8!) out_hidden_state_mapping, "
|
||||
"Tensor query_start_loc, Tensor query_end_loc, "
|
||||
"SymInt padding_token_id, SymInt parallel_drafting_token_id, "
|
||||
"SymInt total_input_tokens, SymInt num_padding_slots_per_request, "
|
||||
"bool shift_input_ids) -> ()",
|
||||
&cpu_utils::copy_and_expand_eagle_inputs_kernel_impl);
|
||||
ops.def(
|
||||
"rejection_greedy_sample_kernel_impl("
|
||||
"Tensor(a0!) output_token_ids, Tensor cu_num_draft_tokens, "
|
||||
"Tensor draft_token_ids, Tensor target_argmax, "
|
||||
"Tensor bonus_token_ids, Tensor? is_greedy, "
|
||||
"SymInt max_spec_len) -> ()",
|
||||
&cpu_utils::rejection_greedy_sample_kernel_impl);
|
||||
ops.def(
|
||||
"rejection_random_sample_kernel_impl("
|
||||
"Tensor(a0!) output_token_ids, Tensor cu_num_draft_tokens, "
|
||||
"Tensor draft_token_ids, Tensor? draft_probs, "
|
||||
"Tensor target_probs, Tensor bonus_token_ids, "
|
||||
"Tensor recovered_token_ids, Tensor uniform_probs, "
|
||||
"Tensor? is_greedy, SymInt max_spec_len, SymInt vocab_size, "
|
||||
"bool no_draft_probs) -> ()",
|
||||
&cpu_utils::rejection_random_sample_kernel_impl);
|
||||
ops.def(
|
||||
"expand_kernel_impl(Tensor(a0!) output, Tensor input, "
|
||||
"Tensor cu_num_tokens, SymInt replace_from, "
|
||||
"SymInt replace_to) -> ()",
|
||||
&cpu_utils::expand_kernel_impl);
|
||||
ops.def(
|
||||
"sample_recovered_tokens_kernel_impl("
|
||||
"Tensor(a0!) output_token_ids, Tensor cu_num_draft_tokens, "
|
||||
"Tensor draft_token_ids, Tensor? draft_probs, "
|
||||
"Tensor target_probs, Tensor inv_q, SymInt vocab_size, "
|
||||
"bool no_draft_probs) -> ()",
|
||||
&cpu_utils::sample_recovered_tokens_kernel_impl);
|
||||
}
|
||||
|
||||
REGISTER_EXTENSION(TORCH_EXTENSION_NAME)
|
||||
|
||||
+2
-1
@@ -55,7 +55,8 @@ struct Counter {
|
||||
|
||||
inline int64_t get_available_l2_size() {
|
||||
static int64_t size = []() {
|
||||
const uint32_t l2_cache_size = at::cpu::L2_cache_size();
|
||||
auto caps = at::cpu::get_cpu_capabilities();
|
||||
const uint32_t l2_cache_size = caps.at("l2_cache_size").toInt();
|
||||
return l2_cache_size >> 1; // use 50% of L2 cache
|
||||
}();
|
||||
return size;
|
||||
|
||||
@@ -389,20 +389,28 @@ struct Sm90ColOrScalarBroadcastArray {
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
begin() {
|
||||
cute::Tensor pred = make_tensor<bool>(shape(tCgCol));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(pred); ++i) {
|
||||
pred(i) = get<0>(tCcCol(i)) < m;
|
||||
}
|
||||
|
||||
if (!params.col_broadcast) {
|
||||
fill(tCrCol, *(params.ptr_col_array[group]));
|
||||
return;
|
||||
}
|
||||
|
||||
// Filter so we don't issue redundant copies over stride-0 modes
|
||||
// (only works if 0-strides are in same location, which is by construction)
|
||||
copy_if(pred, filter(tCgCol), filter(tCrCol));
|
||||
// tCgCol has layout (CPY,CPY_M,CPY_N,EPI_M,EPI_N) where CPY_N and
|
||||
// EPI_N are stride-0 for the column broadcast. Slice those modes at
|
||||
// index 0 to avoid redundant copies AND ensure pred/data consistency
|
||||
static_assert(decltype(stride<2>(tCgCol))::value == 0, "Expected stride-0 CPY_N for col broadcast");
|
||||
static_assert(decltype(stride<4>(tCgCol))::value == 0, "Expected stride-0 EPI_N for col broadcast");
|
||||
|
||||
auto tCgCol_s = tCgCol(_,_,0,_,0); // (CPY,CPY_M,EPI_M)
|
||||
auto tCrCol_s = tCrCol(_,_,0,_,0); // (CPY,CPY_M,EPI_M)
|
||||
auto tCcCol_s = tCcCol(_,_,0,_,0); // (CPY,CPY_M,EPI_M)
|
||||
|
||||
cute::Tensor pred = make_tensor<bool>(shape(tCgCol_s));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(pred); ++i) {
|
||||
pred(i) = get<0>(tCcCol_s(i)) < m;
|
||||
}
|
||||
|
||||
copy_if(pred, tCgCol_s, tCrCol_s);
|
||||
}
|
||||
|
||||
template <typename ElementAccumulator, int FragmentSize>
|
||||
|
||||
@@ -382,20 +382,28 @@ struct Sm90ColOrScalarBroadcast {
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
begin() {
|
||||
cute::Tensor pred = make_tensor<bool>(shape(tCgCol));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(pred); ++i) {
|
||||
pred(i) = get<0>(tCcCol(i)) < m;
|
||||
}
|
||||
|
||||
if (!params.col_broadcast) {
|
||||
fill(tCrCol, *(params.ptr_col));
|
||||
return;
|
||||
}
|
||||
|
||||
// Filter so we don't issue redundant copies over stride-0 modes
|
||||
// (only works if 0-strides are in same location, which is by construction)
|
||||
copy_if(pred, filter(tCgCol), filter(tCrCol));
|
||||
// tCgCol has layout (CPY,CPY_M,CPY_N,EPI_M,EPI_N) where CPY_N and
|
||||
// EPI_N are stride-0 for the column broadcast. Slice those modes at
|
||||
// index 0 to avoid redundant copies AND ensure pred/data consistency
|
||||
static_assert(decltype(stride<2>(tCgCol))::value == 0, "Expected stride-0 CPY_N for col broadcast");
|
||||
static_assert(decltype(stride<4>(tCgCol))::value == 0, "Expected stride-0 EPI_N for col broadcast");
|
||||
|
||||
auto tCgCol_s = tCgCol(_,_,0,_,0); // (CPY,CPY_M,EPI_M)
|
||||
auto tCrCol_s = tCrCol(_,_,0,_,0); // (CPY,CPY_M,EPI_M)
|
||||
auto tCcCol_s = tCcCol(_,_,0,_,0); // (CPY,CPY_M,EPI_M)
|
||||
|
||||
cute::Tensor pred = make_tensor<bool>(shape(tCgCol_s));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(pred); ++i) {
|
||||
pred(i) = get<0>(tCcCol_s(i)) < m;
|
||||
}
|
||||
|
||||
copy_if(pred, tCgCol_s, tCrCol_s);
|
||||
}
|
||||
|
||||
template <typename ElementAccumulator, int FragmentSize>
|
||||
|
||||
@@ -19,8 +19,10 @@
|
||||
#include <type_traits>
|
||||
|
||||
#include <torch/cuda.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include "async_util.cuh"
|
||||
#include "cuda_compat.h"
|
||||
#include "dispatch_utils.h"
|
||||
#include "type_convert.cuh"
|
||||
@@ -86,6 +88,9 @@ inline __device__ __host__ T divUp(T m, T n) {
|
||||
} // namespace tensorrt_llm::common
|
||||
|
||||
namespace tensorrt_llm::kernels {
|
||||
|
||||
using namespace vllm::cuda_async;
|
||||
|
||||
// NOTE(zhuhaoran): This kernel is adapted from TensorRT-LLM implementation,
|
||||
// with added support for passing the cos_sin_cache as an input.
|
||||
// https://github.com/NVIDIA/TensorRT-LLM/blob/main/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.cu
|
||||
@@ -301,6 +306,237 @@ __global__ void fusedQKNormRopeKernel(
|
||||
#endif
|
||||
}
|
||||
|
||||
// Multi-token-head kernel: one warp processes HEADS_PER_WARP token-heads for
|
||||
// the same token, sharing cos/sin from shared memory via cp.async.
|
||||
// When HEADS_PER_WARP > 1 the warp reuses the loaded cos/sin across all heads,
|
||||
// hiding global-memory latency and improving occupancy for large batches.
|
||||
template <typename scalar_t_in, typename scalar_t_cache, int head_dim,
|
||||
bool interleave, int HEADS_PER_WARP>
|
||||
__global__ void fusedQKNormRopeKernelNTokenHeads(
|
||||
void* qkv_void, int const num_heads_q, int const num_heads_k,
|
||||
int const num_heads_v, float const eps, void const* q_weight_void,
|
||||
void const* k_weight_void, void const* cos_sin_cache_void,
|
||||
int64_t const* position_ids, int const num_tokens, int const rotary_dim) {
|
||||
#if (!defined(__CUDA_ARCH__) || __CUDA_ARCH__ < 800) && !defined(USE_ROCM)
|
||||
if constexpr ((std::is_same_v<scalar_t_in, c10::BFloat16>) ||
|
||||
std::is_same_v<scalar_t_cache, c10::BFloat16>) {
|
||||
return;
|
||||
} else {
|
||||
#endif
|
||||
|
||||
using Converter = vllm::_typeConvert<scalar_t_in>;
|
||||
static_assert(Converter::exists,
|
||||
"Input QKV data type is not supported for this CUDA "
|
||||
"architecture or toolkit version.");
|
||||
using T_in = typename Converter::hip_type;
|
||||
using T2_in = typename Converter::packed_hip_type;
|
||||
|
||||
using CacheConverter = vllm::_typeConvert<scalar_t_cache>;
|
||||
static_assert(CacheConverter::exists,
|
||||
"Cache data type is not supported for this CUDA architecture "
|
||||
"or toolkit version.");
|
||||
using T_cache = typename CacheConverter::hip_type;
|
||||
|
||||
extern __shared__ char smem_storage[];
|
||||
// Shared memory layout:
|
||||
// [0, cos_sin_bytes) : cos/sin for each warp (warpsPerBlock *
|
||||
// rotary_dim * sizeof(T_cache))
|
||||
// [cos_sin_bytes, ...) : QKV tiles
|
||||
// per warp (warpsPerBlock * HEADS_PER_WARP * 32 * elemSizeBytes)
|
||||
T_cache* const smem = reinterpret_cast<T_cache*>(smem_storage);
|
||||
|
||||
T_in* qkv = reinterpret_cast<T_in*>(qkv_void);
|
||||
T_in const* q_weight = reinterpret_cast<T_in const*>(q_weight_void);
|
||||
T_in const* k_weight = reinterpret_cast<T_in const*>(k_weight_void);
|
||||
T_cache const* cos_sin_cache =
|
||||
reinterpret_cast<T_cache const*>(cos_sin_cache_void);
|
||||
|
||||
int const warpsPerBlock = blockDim.x / 32;
|
||||
int const warpId = threadIdx.x / 32;
|
||||
int const laneId = threadIdx.x % 32;
|
||||
|
||||
int const total_qk_heads = num_heads_q + num_heads_k;
|
||||
int const num_heads = num_heads_q + num_heads_k + num_heads_v;
|
||||
int const head_chunks_per_token =
|
||||
(total_qk_heads + HEADS_PER_WARP - 1) / HEADS_PER_WARP;
|
||||
|
||||
int const warp_global = blockIdx.x * warpsPerBlock + warpId;
|
||||
int const tokenIdx = warp_global / head_chunks_per_token;
|
||||
int const headChunk = warp_global % head_chunks_per_token;
|
||||
int const first_head = headChunk * HEADS_PER_WARP;
|
||||
int const num_heads_this_warp =
|
||||
(first_head + HEADS_PER_WARP <= total_qk_heads)
|
||||
? HEADS_PER_WARP
|
||||
: (total_qk_heads - first_head);
|
||||
|
||||
if (tokenIdx >= num_tokens) return;
|
||||
|
||||
static_assert(head_dim % (32 * 2) == 0, "head_dim must be divisible by 64");
|
||||
constexpr int numElemsPerThread = head_dim / 32;
|
||||
constexpr int elemSizeBytes = numElemsPerThread * sizeof(__nv_bfloat16);
|
||||
static_assert(elemSizeBytes % 4 == 0,
|
||||
"elemSizeBytes must be a multiple of 4");
|
||||
constexpr int vecSize = elemSizeBytes / 4;
|
||||
using vec_T = typename tensorrt_llm::common::packed_as<uint, vecSize>::type;
|
||||
|
||||
int const cos_sin_bytes =
|
||||
warpsPerBlock * rotary_dim * static_cast<int>(sizeof(T_cache));
|
||||
int const qkv_tile_bytes = 32 * elemSizeBytes;
|
||||
char* const this_warp_head_smem =
|
||||
smem_storage + cos_sin_bytes +
|
||||
warpId * (HEADS_PER_WARP * qkv_tile_bytes);
|
||||
|
||||
// === Group 0: async load all heads' QKV into smem (issued first). ===
|
||||
for (int k = 0; k < num_heads_this_warp; ++k) {
|
||||
int const localHeadIdx = first_head + k;
|
||||
bool const isQ = localHeadIdx < num_heads_q;
|
||||
int const headIdx = isQ ? localHeadIdx : localHeadIdx - num_heads_q;
|
||||
int offWarp;
|
||||
if (isQ) {
|
||||
offWarp = tokenIdx * num_heads * head_dim + headIdx * head_dim;
|
||||
} else {
|
||||
offWarp = tokenIdx * num_heads * head_dim + num_heads_q * head_dim +
|
||||
headIdx * head_dim;
|
||||
}
|
||||
int const offThread = offWarp + laneId * numElemsPerThread;
|
||||
char* smem_dst =
|
||||
this_warp_head_smem + k * qkv_tile_bytes + laneId * elemSizeBytes;
|
||||
cp_async_shared_global_ca(smem_dst,
|
||||
reinterpret_cast<const char*>(&qkv[offThread]),
|
||||
elemSizeBytes);
|
||||
}
|
||||
cp_async_commit_group(); // commit group 0 (QKV)
|
||||
|
||||
// === Group 1: async load cos/sin into smem (issued second). ===
|
||||
int64_t const pos_id = position_ids[tokenIdx];
|
||||
T_cache const* const cache_ptr = cos_sin_cache + pos_id * rotary_dim;
|
||||
int const copy_bytes = rotary_dim * static_cast<int>(sizeof(T_cache));
|
||||
int const num_copies = (copy_bytes + 15) / 16;
|
||||
for (int copyId = laneId; copyId < num_copies; copyId += 32) {
|
||||
char* smem_ptr =
|
||||
reinterpret_cast<char*>(&smem[warpId * rotary_dim]) + copyId * 16;
|
||||
const char* glob_ptr =
|
||||
reinterpret_cast<const char*>(cache_ptr) + copyId * 16;
|
||||
cp_async_shared_global_16_cg(smem_ptr, glob_ptr);
|
||||
}
|
||||
cp_async_commit_group(); // commit group 1 (cos/sin)
|
||||
|
||||
// wait<1>: allow at most 1 pending group (group 1) → group 0 (QKV) is done.
|
||||
cp_async_wait_group<1>();
|
||||
|
||||
float elements[numElemsPerThread];
|
||||
float elements2[numElemsPerThread];
|
||||
int const rotary_lanes = rotary_dim / numElemsPerThread;
|
||||
int const embed_dim = rotary_dim / 2;
|
||||
T_cache const* const cos_smem = &smem[warpId * rotary_dim];
|
||||
T_cache const* const sin_smem = &smem[warpId * rotary_dim + embed_dim];
|
||||
|
||||
// Preload weights into registers once, reused across all heads.
|
||||
float q_w[numElemsPerThread];
|
||||
float k_w[numElemsPerThread];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < numElemsPerThread; i++) {
|
||||
int const dim = laneId * numElemsPerThread + i;
|
||||
q_w[i] = Converter::convert(q_weight[dim]);
|
||||
k_w[i] = Converter::convert(k_weight[dim]);
|
||||
}
|
||||
|
||||
for (int k = 0; k < num_heads_this_warp; ++k) {
|
||||
int const localHeadIdx = first_head + k;
|
||||
bool const isQ = localHeadIdx < num_heads_q;
|
||||
int const headIdx = isQ ? localHeadIdx : localHeadIdx - num_heads_q;
|
||||
|
||||
int offsetWarp;
|
||||
if (isQ) {
|
||||
offsetWarp = tokenIdx * num_heads * head_dim + headIdx * head_dim;
|
||||
} else {
|
||||
offsetWarp = tokenIdx * num_heads * head_dim + num_heads_q * head_dim +
|
||||
headIdx * head_dim;
|
||||
}
|
||||
int const offsetThread = offsetWarp + laneId * numElemsPerThread;
|
||||
|
||||
// === Part 1: QK Norm (read from smem; group 0 already done). ===
|
||||
float sumOfSquares = 0.0f;
|
||||
{
|
||||
char const* smem_src =
|
||||
this_warp_head_smem + k * qkv_tile_bytes + laneId * elemSizeBytes;
|
||||
vec_T vec = *reinterpret_cast<vec_T const*>(smem_src);
|
||||
constexpr int num_packed_elems = elemSizeBytes / sizeof(T2_in);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < num_packed_elems; i++) {
|
||||
T2_in packed_val = *(reinterpret_cast<T2_in*>(&vec) + i);
|
||||
float2 vals = Converter::convert(packed_val);
|
||||
sumOfSquares += vals.x * vals.x;
|
||||
sumOfSquares += vals.y * vals.y;
|
||||
elements[2 * i] = vals.x;
|
||||
elements[2 * i + 1] = vals.y;
|
||||
}
|
||||
}
|
||||
|
||||
sumOfSquares = tensorrt_llm::common::warpReduceSum(sumOfSquares);
|
||||
float rms_rcp = rsqrtf(sumOfSquares / static_cast<float>(head_dim) + eps);
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < numElemsPerThread; i++) {
|
||||
elements[i] *= rms_rcp * (isQ ? q_w[i] : k_w[i]);
|
||||
}
|
||||
|
||||
// On first head: wait for group 1 (cos/sin) before RoPE.
|
||||
if (k == 0) cp_async_wait_group<0>();
|
||||
|
||||
// === Part 2: RoPE using cos/sin from shared memory. ===
|
||||
if (laneId < rotary_lanes) {
|
||||
if constexpr (interleave) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < numElemsPerThread / 2; ++i) {
|
||||
int const idx0 = 2 * i;
|
||||
int const idx1 = 2 * i + 1;
|
||||
int const dim_idx = laneId * numElemsPerThread + idx0;
|
||||
float const val0 = elements[idx0];
|
||||
float const val1 = elements[idx1];
|
||||
int const half_dim = dim_idx / 2;
|
||||
float const cos_val = CacheConverter::convert(cos_smem[half_dim]);
|
||||
float const sin_val = CacheConverter::convert(sin_smem[half_dim]);
|
||||
elements[idx0] = val0 * cos_val - val1 * sin_val;
|
||||
elements[idx1] = val0 * sin_val + val1 * cos_val;
|
||||
}
|
||||
} else {
|
||||
__syncwarp();
|
||||
int const pairOffset = (rotary_dim / 2) / numElemsPerThread;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < numElemsPerThread; i++) {
|
||||
elements2[i] = __shfl_xor_sync(FINAL_MASK, elements[i], pairOffset);
|
||||
if (laneId < pairOffset) elements2[i] = -elements2[i];
|
||||
int dim_idx = laneId * numElemsPerThread + i;
|
||||
dim_idx = (dim_idx * 2) % rotary_dim;
|
||||
int const half_dim = dim_idx / 2;
|
||||
float const cos_val = CacheConverter::convert(cos_smem[half_dim]);
|
||||
float const sin_val = CacheConverter::convert(sin_smem[half_dim]);
|
||||
elements[i] = elements[i] * cos_val + elements2[i] * sin_val;
|
||||
}
|
||||
__syncwarp();
|
||||
}
|
||||
}
|
||||
|
||||
// Store.
|
||||
{
|
||||
vec_T vec;
|
||||
constexpr int num_packed_elems = elemSizeBytes / sizeof(T2_in);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < num_packed_elems; i++) {
|
||||
T2_in packed_val = Converter::convert(
|
||||
make_float2(elements[2 * i], elements[2 * i + 1]));
|
||||
*(reinterpret_cast<T2_in*>(&vec) + i) = packed_val;
|
||||
}
|
||||
*reinterpret_cast<vec_T*>(&qkv[offsetThread]) = vec;
|
||||
}
|
||||
}
|
||||
|
||||
#if (!defined(__CUDA_ARCH__) || __CUDA_ARCH__ < 800) && !defined(USE_ROCM)
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
// Borrowed from
|
||||
// https://github.com/flashinfer-ai/flashinfer/blob/8125d079a43e9a0ba463a4ed1b639cefd084cec9/include/flashinfer/pos_enc.cuh#L568
|
||||
#define DISPATCH_INTERLEAVE(interleave, INTERLEAVE, ...) \
|
||||
@@ -321,15 +557,12 @@ void launchFusedQKNormRope(void* qkv, int const num_tokens,
|
||||
void const* cos_sin_cache, bool const interleave,
|
||||
int64_t const* position_ids, cudaStream_t stream) {
|
||||
constexpr int blockSize = 256;
|
||||
|
||||
int const warpsPerBlock = blockSize / 32;
|
||||
int const totalQKHeads = num_heads_q + num_heads_k;
|
||||
int const totalWarps = num_tokens * totalQKHeads;
|
||||
|
||||
int const gridSize = common::divUp(totalWarps, warpsPerBlock);
|
||||
dim3 gridDim(gridSize);
|
||||
dim3 blockDim(blockSize);
|
||||
|
||||
switch (head_dim) {
|
||||
case 64:
|
||||
DISPATCH_INTERLEAVE(interleave, INTERLEAVE, {
|
||||
@@ -360,6 +593,118 @@ void launchFusedQKNormRope(void* qkv, int const num_tokens,
|
||||
"Unsupported head dimension for fusedQKNormRope: ", head_dim);
|
||||
}
|
||||
}
|
||||
|
||||
// Launch: one warp processes token_heads_per_warp token-heads (1, 2, 4, or 8).
|
||||
// When token_heads_per_warp == 1, delegates to the 1-head baseline above.
|
||||
template <typename scalar_t_in, typename scalar_t_cache>
|
||||
void launchFusedQKNormRopeNTokenHeads(
|
||||
void* qkv, int const num_tokens, int const num_heads_q,
|
||||
int const num_heads_k, int const num_heads_v, int const head_dim,
|
||||
int const rotary_dim, float const eps, void const* q_weight,
|
||||
void const* k_weight, void const* cos_sin_cache, bool const interleave,
|
||||
int64_t const* position_ids, int const token_heads_per_warp,
|
||||
cudaStream_t stream) {
|
||||
TORCH_CHECK(token_heads_per_warp == 1 || token_heads_per_warp == 2 ||
|
||||
token_heads_per_warp == 4 || token_heads_per_warp == 8,
|
||||
"token_heads_per_warp must be 1, 2, 4, or 8, got ",
|
||||
token_heads_per_warp);
|
||||
|
||||
// token_heads_per_warp == 1: delegate to the 1-head baseline kernel.
|
||||
if (token_heads_per_warp == 1) {
|
||||
launchFusedQKNormRope<scalar_t_in, scalar_t_cache>(
|
||||
qkv, num_tokens, num_heads_q, num_heads_k, num_heads_v, head_dim,
|
||||
rotary_dim, eps, q_weight, k_weight, cos_sin_cache, interleave,
|
||||
position_ids, stream);
|
||||
return;
|
||||
}
|
||||
|
||||
// NTokenHeads kernel uses cp.async to load cos/sin in 16-byte chunks.
|
||||
// If rotary_dim * sizeof(cache_dtype) is not a multiple of 16, the last
|
||||
// cp.async would write past the shared memory allocation.
|
||||
// Fall back to the base kernel instead of failing.
|
||||
{
|
||||
size_t const rotary_bytes =
|
||||
static_cast<size_t>(rotary_dim) *
|
||||
(std::is_same_v<scalar_t_cache, float> ? sizeof(float) : 2u);
|
||||
if (rotary_bytes % 16 != 0) {
|
||||
launchFusedQKNormRope<scalar_t_in, scalar_t_cache>(
|
||||
qkv, num_tokens, num_heads_q, num_heads_k, num_heads_v, head_dim,
|
||||
rotary_dim, eps, q_weight, k_weight, cos_sin_cache, interleave,
|
||||
position_ids, stream);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
constexpr int blockSize = 256;
|
||||
int const warpsPerBlock = blockSize / 32;
|
||||
int const totalQKHeads = num_heads_q + num_heads_k;
|
||||
// Grid: one warp per (token, head_chunk); same token → reuse cos/sin in smem.
|
||||
int const head_chunks_per_token =
|
||||
(totalQKHeads + token_heads_per_warp - 1) / token_heads_per_warp;
|
||||
int const total_warps = num_tokens * head_chunks_per_token;
|
||||
int const gridSize = common::divUp(total_warps, warpsPerBlock);
|
||||
dim3 gridDim(gridSize);
|
||||
dim3 blockDim(blockSize);
|
||||
// Cache element size: float=4, bfloat16=2 (host-safe; kernel uses same
|
||||
// layout).
|
||||
size_t const cache_elem_size =
|
||||
std::is_same_v<scalar_t_cache, float> ? sizeof(float) : 2u;
|
||||
// QKV smem: token_heads_per_warp tiles per warp, each tile 32*(head_dim/32*2)
|
||||
// = 2*head_dim bytes.
|
||||
size_t const qkv_smem_per_warp = static_cast<size_t>(token_heads_per_warp) *
|
||||
2u * static_cast<size_t>(head_dim);
|
||||
size_t const smem_bytes =
|
||||
warpsPerBlock * static_cast<size_t>(rotary_dim) * cache_elem_size +
|
||||
warpsPerBlock * qkv_smem_per_warp;
|
||||
|
||||
#define LAUNCH_N_TOKEN_HEADS(N) \
|
||||
do { \
|
||||
switch (head_dim) { \
|
||||
case 64: \
|
||||
DISPATCH_INTERLEAVE(interleave, INTERLEAVE, { \
|
||||
fusedQKNormRopeKernelNTokenHeads<scalar_t_in, scalar_t_cache, 64, \
|
||||
INTERLEAVE, (N)> \
|
||||
<<<gridDim, blockDim, smem_bytes, stream>>>( \
|
||||
qkv, num_heads_q, num_heads_k, num_heads_v, eps, q_weight, \
|
||||
k_weight, cos_sin_cache, position_ids, num_tokens, \
|
||||
rotary_dim); \
|
||||
}); \
|
||||
break; \
|
||||
case 128: \
|
||||
DISPATCH_INTERLEAVE(interleave, INTERLEAVE, { \
|
||||
fusedQKNormRopeKernelNTokenHeads<scalar_t_in, scalar_t_cache, 128, \
|
||||
INTERLEAVE, (N)> \
|
||||
<<<gridDim, blockDim, smem_bytes, stream>>>( \
|
||||
qkv, num_heads_q, num_heads_k, num_heads_v, eps, q_weight, \
|
||||
k_weight, cos_sin_cache, position_ids, num_tokens, \
|
||||
rotary_dim); \
|
||||
}); \
|
||||
break; \
|
||||
case 256: \
|
||||
DISPATCH_INTERLEAVE(interleave, INTERLEAVE, { \
|
||||
fusedQKNormRopeKernelNTokenHeads<scalar_t_in, scalar_t_cache, 256, \
|
||||
INTERLEAVE, (N)> \
|
||||
<<<gridDim, blockDim, smem_bytes, stream>>>( \
|
||||
qkv, num_heads_q, num_heads_k, num_heads_v, eps, q_weight, \
|
||||
k_weight, cos_sin_cache, position_ids, num_tokens, \
|
||||
rotary_dim); \
|
||||
}); \
|
||||
break; \
|
||||
default: \
|
||||
TORCH_CHECK(false, "Unsupported head dimension: ", head_dim); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
if (token_heads_per_warp == 2) {
|
||||
LAUNCH_N_TOKEN_HEADS(2);
|
||||
} else if (token_heads_per_warp == 4) {
|
||||
LAUNCH_N_TOKEN_HEADS(4);
|
||||
} else if (token_heads_per_warp == 8) {
|
||||
LAUNCH_N_TOKEN_HEADS(8);
|
||||
}
|
||||
#undef LAUNCH_N_TOKEN_HEADS
|
||||
}
|
||||
|
||||
} // namespace tensorrt_llm::kernels
|
||||
|
||||
void fused_qk_norm_rope(
|
||||
@@ -374,7 +719,8 @@ void fused_qk_norm_rope(
|
||||
torch::Tensor& k_weight, // RMSNorm weights for key [head_dim]
|
||||
torch::Tensor& cos_sin_cache, // Cos/sin cache [max_position, head_dim]
|
||||
bool is_neox, // Whether RoPE is applied in Neox style
|
||||
torch::Tensor& position_ids // Position IDs for RoPE [num_tokens]
|
||||
torch::Tensor& position_ids, // Position IDs for RoPE [num_tokens]
|
||||
int64_t forced_token_heads_per_warp // -1 = auto-select, >0 = forced value
|
||||
) {
|
||||
// Input validation
|
||||
CHECK_INPUT(qkv);
|
||||
@@ -414,15 +760,48 @@ void fused_qk_norm_rope(
|
||||
qkv.size(1) == total_heads * head_dim,
|
||||
"QKV tensor size must match total number of heads and head dimension");
|
||||
|
||||
auto stream = at::cuda::getCurrentCUDAStream(qkv.get_device());
|
||||
auto device_id = qkv.get_device();
|
||||
auto stream = at::cuda::getCurrentCUDAStream(device_id);
|
||||
|
||||
// Select token_heads_per_warp: forced value if >0, else auto-select.
|
||||
// Auto thresholds are calibrated on SM 9.0 (H100). On other architectures,
|
||||
// fall back to token_heads_per_warp=1 (base kernel) until profiled.
|
||||
int token_heads_per_warp;
|
||||
if (forced_token_heads_per_warp > 0) { // only support SM80+
|
||||
token_heads_per_warp = static_cast<int>(forced_token_heads_per_warp);
|
||||
} else {
|
||||
token_heads_per_warp = 1;
|
||||
auto* dev_prop = at::cuda::getDeviceProperties(device_id);
|
||||
int sm_version = dev_prop->major * 10 + dev_prop->minor;
|
||||
int64_t total_qk_units = num_tokens * (num_heads_q + num_heads_k);
|
||||
if (sm_version == 90) {
|
||||
if (head_dim >= 256) {
|
||||
if (total_qk_units < 4096LL) {
|
||||
token_heads_per_warp = 1;
|
||||
} else if (total_qk_units < 8192LL) {
|
||||
token_heads_per_warp = 2;
|
||||
} else {
|
||||
token_heads_per_warp = 4;
|
||||
}
|
||||
} else {
|
||||
if (total_qk_units < 10240LL) {
|
||||
token_heads_per_warp = 1;
|
||||
} else if (total_qk_units < 40960LL) {
|
||||
token_heads_per_warp = 4;
|
||||
} else {
|
||||
token_heads_per_warp = 8;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
VLLM_DISPATCH_HALF_TYPES(qkv.scalar_type(), "fused_qk_norm_rope_kernel", [&] {
|
||||
using qkv_scalar_t = scalar_t;
|
||||
VLLM_DISPATCH_FLOATING_TYPES(
|
||||
cos_sin_cache.scalar_type(), "fused_qk_norm_rope_kernel", [&] {
|
||||
using cache_scalar_t = scalar_t;
|
||||
tensorrt_llm::kernels::launchFusedQKNormRope<qkv_scalar_t,
|
||||
cache_scalar_t>(
|
||||
tensorrt_llm::kernels::launchFusedQKNormRopeNTokenHeads<
|
||||
qkv_scalar_t, cache_scalar_t>(
|
||||
qkv.data_ptr(), static_cast<int>(num_tokens),
|
||||
static_cast<int>(num_heads_q), static_cast<int>(num_heads_k),
|
||||
static_cast<int>(num_heads_v), static_cast<int>(head_dim),
|
||||
@@ -430,7 +809,7 @@ void fused_qk_norm_rope(
|
||||
q_weight.data_ptr(), k_weight.data_ptr(),
|
||||
cos_sin_cache.data_ptr(), !is_neox,
|
||||
reinterpret_cast<int64_t const*>(position_ids.data_ptr()),
|
||||
stream);
|
||||
token_heads_per_warp, stream);
|
||||
});
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -240,8 +240,9 @@ template <typename T, typename DST_DTYPE>
|
||||
__global__ void per_token_group_quant_8bit_packed_kernel(
|
||||
const T* __restrict__ input, void* __restrict__ output_q,
|
||||
unsigned int* __restrict__ output_s_packed, const int group_size,
|
||||
const int num_groups, const int groups_per_block, const int groups_per_row,
|
||||
const int mn, const int tma_aligned_mn, const float eps,
|
||||
const int num_groups_padded, const int groups_per_block,
|
||||
const int padded_groups_per_row, const int groups_per_row, const int mn,
|
||||
const int tma_aligned_mn, const int num_scale_elems, const float eps,
|
||||
const float min_8bit, const float max_8bit) {
|
||||
const int threads_per_group = 16;
|
||||
const int64_t local_group_id = threadIdx.x / threads_per_group;
|
||||
@@ -249,51 +250,62 @@ __global__ void per_token_group_quant_8bit_packed_kernel(
|
||||
|
||||
const int64_t block_group_id = blockIdx.x * groups_per_block;
|
||||
const int64_t global_group_id = block_group_id + local_group_id;
|
||||
if (global_group_id >= num_groups) {
|
||||
if (global_group_id >= num_groups_padded) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int64_t block_group_offset = global_group_id * group_size;
|
||||
// map flat group id to 2D indices (mn_idx, sf_k_idx)
|
||||
const int sf_k_idx =
|
||||
static_cast<int>(global_group_id % padded_groups_per_row);
|
||||
const int mn_idx = static_cast<int>(global_group_id / padded_groups_per_row);
|
||||
|
||||
const T* group_input = input + block_group_offset;
|
||||
DST_DTYPE* group_output =
|
||||
static_cast<DST_DTYPE*>(output_q) + block_group_offset;
|
||||
// whether it is a valid group (not padding)
|
||||
const bool is_valid_group = (mn_idx < mn) && (sf_k_idx < groups_per_row);
|
||||
|
||||
// shared memory to cache each group's data to avoid double DRAM reads.
|
||||
extern __shared__ __align__(16) char smem_raw[];
|
||||
T* smem = reinterpret_cast<T*>(smem_raw);
|
||||
T* smem_group = smem + local_group_id * group_size;
|
||||
const float y_s =
|
||||
ComputeGroupScale<T, true>(group_input, smem_group, group_size, lane_id,
|
||||
threads_per_group, eps, max_8bit);
|
||||
|
||||
// pack 4 scales into a uint32
|
||||
// compute scale for valid groups
|
||||
float y_s = 0.f;
|
||||
if (is_valid_group) {
|
||||
const T* group_input =
|
||||
input + static_cast<int64_t>(mn_idx) * groups_per_row * group_size +
|
||||
sf_k_idx * group_size;
|
||||
y_s = ComputeGroupScale<T, true>(group_input, smem_group, group_size,
|
||||
lane_id, threads_per_group, eps, max_8bit);
|
||||
}
|
||||
|
||||
// pack 4 scales into a uint32 exponent
|
||||
if (lane_id == 0) {
|
||||
// map flat group id to 2D indices (mn_idx, sf_k_idx)
|
||||
const int sf_k_idx = static_cast<int>(global_group_id % groups_per_row);
|
||||
const int mn_idx = static_cast<int>(global_group_id / groups_per_row);
|
||||
|
||||
if (mn_idx < mn) {
|
||||
// each uint32 in output_s_packed stores 4 packed scales
|
||||
const int sf_k_pack_idx = sf_k_idx / 4;
|
||||
const int pos = sf_k_idx % 4;
|
||||
// each uint32 in output_s_packed stores 4 packed scales
|
||||
const int sf_k_pack_idx = sf_k_idx / 4;
|
||||
const int pos = sf_k_idx % 4;
|
||||
const int out_idx = sf_k_pack_idx * tma_aligned_mn + mn_idx;
|
||||
|
||||
if (is_valid_group) {
|
||||
// reinterpret the UE8M0 scale y_s as IEEE bits, extract the 8-bit
|
||||
// exponent, and place it into the correct byte of the 32-bit word.
|
||||
const unsigned int bits = __float_as_uint(y_s);
|
||||
const unsigned int exponent = (bits >> 23u) & 0xffu;
|
||||
const unsigned int contrib = exponent << (pos * 8u);
|
||||
|
||||
const int out_idx = sf_k_pack_idx * tma_aligned_mn + mn_idx;
|
||||
// atomically OR 8-bit exponent into the packed scales buffer
|
||||
atomicOr(output_s_packed + out_idx, contrib);
|
||||
const uint8_t exponent = static_cast<uint8_t>((bits >> 23u) & 0xffu);
|
||||
reinterpret_cast<uint8_t*>(output_s_packed)[out_idx * 4 + pos] = exponent;
|
||||
} else if (out_idx < num_scale_elems) {
|
||||
// write zero for padding groups if within bounds of output_s_packed
|
||||
reinterpret_cast<uint8_t*>(output_s_packed)[out_idx * 4 + pos] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
|
||||
QuantizeGroup<T, DST_DTYPE>(smem_group, group_output, group_size, lane_id,
|
||||
threads_per_group, y_s, min_8bit, max_8bit);
|
||||
if (is_valid_group) {
|
||||
DST_DTYPE* group_output =
|
||||
static_cast<DST_DTYPE*>(output_q) +
|
||||
static_cast<int64_t>(mn_idx) * groups_per_row * group_size +
|
||||
sf_k_idx * group_size;
|
||||
QuantizeGroup<T, DST_DTYPE>(smem_group, group_output, group_size, lane_id,
|
||||
threads_per_group, y_s, min_8bit, max_8bit);
|
||||
}
|
||||
}
|
||||
|
||||
void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
|
||||
@@ -310,7 +322,6 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
|
||||
|
||||
const int64_t mn = input.numel() / k;
|
||||
const int64_t groups_per_row = k / group_size;
|
||||
const int64_t num_groups = mn * groups_per_row;
|
||||
|
||||
STD_TORCH_CHECK(output_s_packed.dim() == 2,
|
||||
"output_s_packed must be 2D, got dim=", output_s_packed.dim(),
|
||||
@@ -330,36 +341,46 @@ void per_token_group_quant_8bit_packed(const torch::stable::Tensor& input,
|
||||
"output_s_packed shape must be [", mn, ", ", k_num_packed_sfk,
|
||||
"], but got [", output_s_packed.size(0), ", ",
|
||||
output_s_packed.size(1), "].");
|
||||
// Verify column-major TMA-aligned layout
|
||||
STD_TORCH_CHECK(output_s_packed.stride(0) == 1 &&
|
||||
output_s_packed.stride(1) == tma_aligned_mn,
|
||||
"output_s_packed must have strides [1, ", tma_aligned_mn,
|
||||
"], but got [", output_s_packed.stride(0), ", ",
|
||||
output_s_packed.stride(1), "].");
|
||||
|
||||
cudaStream_t stream = get_current_cuda_stream();
|
||||
|
||||
constexpr int THREADS_PER_GROUP = 16;
|
||||
|
||||
const int groups_per_block = GetGroupsPerBlock(num_groups);
|
||||
// Expand the grid to cover MN and K padding so every byte in
|
||||
// output_s_packed is written (padding bytes get zeroed by the kernel).
|
||||
const int64_t padded_groups_per_row = k_num_packed_sfk * 4;
|
||||
const int64_t num_groups_padded = tma_aligned_mn * padded_groups_per_row;
|
||||
// Number of elements in output_s_packed.
|
||||
const int64_t num_scale_elems = mn + (k_num_packed_sfk - 1) * tma_aligned_mn;
|
||||
|
||||
const int groups_per_block = GetGroupsPerBlock(num_groups_padded);
|
||||
|
||||
auto dst_type = output_q.scalar_type();
|
||||
const int num_blocks = num_groups / groups_per_block;
|
||||
const int num_blocks = num_groups_padded / groups_per_block;
|
||||
const int num_threads = groups_per_block * THREADS_PER_GROUP;
|
||||
|
||||
// zero-initialize packed scales, since we use atomicOr to accumulate
|
||||
// exponents from different groups.
|
||||
torch::stable::zero_(output_s_packed);
|
||||
|
||||
#define LAUNCH_PACKED_KERNEL(T, DST_DTYPE) \
|
||||
do { \
|
||||
dim3 grid(num_blocks); \
|
||||
dim3 block(num_threads); \
|
||||
size_t smem_bytes = \
|
||||
static_cast<size_t>(groups_per_block) * group_size * sizeof(T); \
|
||||
per_token_group_quant_8bit_packed_kernel<T, DST_DTYPE> \
|
||||
<<<grid, block, smem_bytes, stream>>>( \
|
||||
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
|
||||
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
|
||||
static_cast<int>(group_size), static_cast<int>(num_groups), \
|
||||
groups_per_block, static_cast<int>(groups_per_row), \
|
||||
static_cast<int>(mn), static_cast<int>(tma_aligned_mn), \
|
||||
static_cast<float>(eps), static_cast<float>(min_8bit), \
|
||||
static_cast<float>(max_8bit)); \
|
||||
#define LAUNCH_PACKED_KERNEL(T, DST_DTYPE) \
|
||||
do { \
|
||||
dim3 grid(num_blocks); \
|
||||
dim3 block(num_threads); \
|
||||
size_t smem_bytes = \
|
||||
static_cast<size_t>(groups_per_block) * group_size * sizeof(T); \
|
||||
per_token_group_quant_8bit_packed_kernel<T, DST_DTYPE> \
|
||||
<<<grid, block, smem_bytes, stream>>>( \
|
||||
static_cast<const T*>(input.data_ptr()), output_q.data_ptr(), \
|
||||
reinterpret_cast<unsigned int*>(output_s_packed.data_ptr()), \
|
||||
static_cast<int>(group_size), static_cast<int>(num_groups_padded), \
|
||||
groups_per_block, static_cast<int>(padded_groups_per_row), \
|
||||
static_cast<int>(groups_per_row), static_cast<int>(mn), \
|
||||
static_cast<int>(tma_aligned_mn), \
|
||||
static_cast<int>(num_scale_elems), static_cast<float>(eps), \
|
||||
static_cast<float>(min_8bit), static_cast<float>(max_8bit)); \
|
||||
} while (0)
|
||||
|
||||
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
|
||||
|
||||
@@ -0,0 +1,879 @@
|
||||
|
||||
/*
|
||||
* Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#include <cooperative_groups.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include <torch/cuda.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include "cuda_compat.h"
|
||||
#include "cuda_utils.h"
|
||||
#include "core/registration.h"
|
||||
#include "minimax_reduce_rms_kernel.h"
|
||||
|
||||
#include <algorithm>
|
||||
|
||||
#define FINAL_MASK 0xffffffff
|
||||
#define MINIMAX_REDUCE_RMS_WARP_SIZE 32
|
||||
|
||||
namespace vllm {
|
||||
namespace tensorrt_llm {
|
||||
|
||||
template <int NRanks>
|
||||
struct LamportComm {
|
||||
__device__ __forceinline__ LamportComm(void** workspace, int rank) {
|
||||
counter_ptr = &reinterpret_cast<int*>(workspace[NRanks * 3])[0];
|
||||
flag_ptr = &reinterpret_cast<int*>(workspace[NRanks * 3])[2];
|
||||
clear_ptr = &reinterpret_cast<int64_t*>(workspace[NRanks * 3 + 1])[0];
|
||||
flag_value = *flag_ptr;
|
||||
auto comm_size = reinterpret_cast<int64_t*>(workspace[NRanks * 3 + 1])[1];
|
||||
clear_size = *clear_ptr;
|
||||
int data_offset = flag_value % 3;
|
||||
int clear_offset = (flag_value + 2) % 3;
|
||||
for (int r = 0; r < NRanks; ++r) {
|
||||
data_bufs[r] = reinterpret_cast<uint8_t*>(workspace[2 * NRanks + r]) +
|
||||
data_offset * comm_size;
|
||||
}
|
||||
clear_buf = reinterpret_cast<uint8_t*>(workspace[2 * NRanks + rank]) +
|
||||
clear_offset * comm_size;
|
||||
__syncthreads();
|
||||
if (threadIdx.x == 0) {
|
||||
atomicAdd(counter_ptr, 1);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void update(int64_t new_clear_size) {
|
||||
if (blockIdx.x == 0 && threadIdx.x == 0) {
|
||||
while (*reinterpret_cast<int volatile*>(counter_ptr) != gridDim.x) {
|
||||
}
|
||||
*flag_ptr = (flag_value + 1) % 3;
|
||||
*clear_ptr = new_clear_size;
|
||||
*counter_ptr = 0;
|
||||
}
|
||||
}
|
||||
|
||||
int* counter_ptr;
|
||||
int* flag_ptr;
|
||||
int64_t* clear_ptr;
|
||||
uint8_t* data_bufs[NRanks];
|
||||
uint8_t* clear_buf;
|
||||
int64_t clear_size;
|
||||
int flag_value;
|
||||
};
|
||||
|
||||
__device__ __forceinline__ bool is_neg_zero(float v) {
|
||||
return *reinterpret_cast<uint32_t*>(&v) == 0x80000000;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ bool is_neg_zero(float4 v) {
|
||||
return is_neg_zero(v.x) || is_neg_zero(v.y) || is_neg_zero(v.z) ||
|
||||
is_neg_zero(v.w);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float4 get_neg_zero() {
|
||||
float4 vec;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
reinterpret_cast<uint32_t*>(&vec)[i] = 0x80000000;
|
||||
}
|
||||
return vec;
|
||||
}
|
||||
|
||||
template <int Dim>
|
||||
__device__ __forceinline__ float rms_rsqrt(float& v, float eps) {
|
||||
constexpr float kInvDim = 1.0F / static_cast<float>(Dim);
|
||||
v = rsqrtf((v * kInvDim) + eps);
|
||||
return v;
|
||||
}
|
||||
|
||||
template <int Dim>
|
||||
__device__ __forceinline__ float4 rms_rsqrt(float4& v, float eps) {
|
||||
constexpr float kInvDim = 1.0F / static_cast<float>(Dim);
|
||||
v.x = rsqrtf((v.x * kInvDim) + eps);
|
||||
v.y = rsqrtf((v.y * kInvDim) + eps);
|
||||
v.z = rsqrtf((v.z * kInvDim) + eps);
|
||||
v.w = rsqrtf((v.w * kInvDim) + eps);
|
||||
return v;
|
||||
}
|
||||
__device__ __forceinline__ float4 ld_global_volatile(float4* addr) {
|
||||
float4 val;
|
||||
asm volatile("ld.volatile.global.v4.f32 {%0, %1, %2, %3}, [%4];"
|
||||
: "=f"(val.x), "=f"(val.y), "=f"(val.z), "=f"(val.w)
|
||||
: "l"(addr));
|
||||
return val;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float ld_global_volatile(float* addr) {
|
||||
float val;
|
||||
asm volatile("ld.volatile.global.f32 %0, [%1];" : "=f"(val) : "l"(addr));
|
||||
return val;
|
||||
}
|
||||
|
||||
// Used by the scalar (non-float4) kernel only
|
||||
template <typename T, int NUM>
|
||||
__inline__ __device__ T warpReduceSumV2(T* val) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < NUM; i++) {
|
||||
#pragma unroll
|
||||
for (int mask = 16; mask > 0; mask >>= 1)
|
||||
val[i] += __shfl_xor_sync(FINAL_MASK, val[i], mask, 32);
|
||||
}
|
||||
return (T)(0.0f);
|
||||
}
|
||||
|
||||
template <typename T, int NUM>
|
||||
__inline__ __device__ T blockReduceSumV2(T* val) {
|
||||
static __shared__ T shared[NUM][33];
|
||||
int lane = threadIdx.x & 0x1f;
|
||||
int wid = threadIdx.x >> 5;
|
||||
|
||||
warpReduceSumV2<T, NUM>(val);
|
||||
|
||||
if (lane == 0) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < NUM; i++) {
|
||||
shared[i][wid] = val[i];
|
||||
}
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
|
||||
bool is_mask = threadIdx.x < (blockDim.x / 32.f);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < NUM; i++) {
|
||||
val[i] = is_mask ? shared[i][lane] : (T)(0.0f);
|
||||
}
|
||||
warpReduceSumV2<T, NUM>(val);
|
||||
return (T)0.0f;
|
||||
}
|
||||
|
||||
// for float4 version
|
||||
template <uint32_t kNumThreads, typename T, int ArraySize = 4>
|
||||
__device__ __forceinline__ void local_warp_reduce_sum_array(
|
||||
T* value_ptr, uint32_t active_mask = 0xffffffffu) {
|
||||
static_assert(kNumThreads >= 1 &&
|
||||
kNumThreads <= MINIMAX_REDUCE_RMS_WARP_SIZE);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < ArraySize; ++i) {
|
||||
#pragma unroll
|
||||
for (int mask = kNumThreads / 2; mask > 0; mask >>= 1) {
|
||||
value_ptr[i] += __shfl_xor_sync(active_mask, value_ptr[i], mask,
|
||||
MINIMAX_REDUCE_RMS_WARP_SIZE);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
constexpr int next_pow2(int val) {
|
||||
int result = 1;
|
||||
while (result < val) {
|
||||
result <<= 1;
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <typename DType>
|
||||
class IndexHelper {
|
||||
public:
|
||||
__device__ __forceinline__ IndexHelper(MiniMaxReduceRMSParams const& params) {
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
namespace cg = cooperative_groups;
|
||||
cg::cluster_group cluster = cg::this_cluster();
|
||||
cg::grid_group grid = cg::this_grid();
|
||||
token_id = grid.cluster_rank();
|
||||
access_id_in_token = cluster.thread_rank();
|
||||
token_stride = grid.num_clusters();
|
||||
#else
|
||||
token_id = blockIdx.x;
|
||||
access_id_in_token = threadIdx.x;
|
||||
token_stride = gridDim.x;
|
||||
#endif
|
||||
access_id = token_id * params.hidden_dim / kElemsPerAccess<DType> +
|
||||
access_id_in_token;
|
||||
access_stride = token_stride * params.hidden_dim / kElemsPerAccess<DType>;
|
||||
tot_access = params.size_q / kElemsPerAccess<DType>;
|
||||
}
|
||||
|
||||
int token_id;
|
||||
int access_id_in_token;
|
||||
int token_stride;
|
||||
int access_id;
|
||||
int access_stride;
|
||||
int tot_access;
|
||||
};
|
||||
|
||||
/**
|
||||
* this kernel is used to for minimax attention module
|
||||
* input tensor [total_tokens, hidden_dim / tp_size], fp32
|
||||
* rms weight [hidden_dim / tp_size], bf16
|
||||
step 1: reduce from single rank to get the variance sum (reduce(input^2,
|
||||
dim=-1)) step 2: reduce from all ranks to get the variance sum
|
||||
(all_reduce(variance_sum)) step 3: calculate the rms norm (input *
|
||||
rsqrt(variance + eps)) in this case, max hidden_dim is 6144 (float data), for
|
||||
each token, we only need 6144 / 4 / tp_size = (1536 / tp_size) threads so we can
|
||||
assume cluster size is 1 (tp_size >= 2)
|
||||
*/
|
||||
template <typename DType, int NRanks>
|
||||
__global__ void __launch_bounds__(1024)
|
||||
minimax_reduce_rms_kernel_lamport(MiniMaxReduceRMSParams params) {
|
||||
IndexHelper<DType> index_helper(params);
|
||||
int token_id = index_helper.token_id;
|
||||
int access_id_in_token = index_helper.access_id_in_token;
|
||||
int token_stride = index_helper.token_stride;
|
||||
int access_id = index_helper.access_id;
|
||||
int access_stride = index_helper.access_stride;
|
||||
int tot_access = index_helper.tot_access;
|
||||
int tot_tokens = params.size_q / params.hidden_dim;
|
||||
float4 clear_vec = get_neg_zero();
|
||||
|
||||
LamportComm<NRanks> comm(params.workspace, params.rank);
|
||||
int clear_access = comm.clear_size / kElemsPerAccess<DType>;
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
asm volatile("griddepcontrol.wait;");
|
||||
#endif
|
||||
for (int idx = access_id; idx < tot_access;
|
||||
idx += access_stride, token_id += token_stride) {
|
||||
alignas(16) DType vals[kElemsPerAccess<DType>];
|
||||
float sum_variance = 0.F;
|
||||
*reinterpret_cast<float4*>(vals) =
|
||||
reinterpret_cast<float4*>(params.allreduce_in)[idx];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kElemsPerAccess<DType>; ++i) {
|
||||
sum_variance += static_cast<float>(vals[i]) * static_cast<float>(vals[i]);
|
||||
}
|
||||
blockReduceSumV2<float, 1>(&sum_variance);
|
||||
if (is_neg_zero(sum_variance)) {
|
||||
sum_variance = 0.F;
|
||||
}
|
||||
if (threadIdx.x == 0) {
|
||||
for (int r = 0; r < NRanks; ++r) {
|
||||
reinterpret_cast<float*>(
|
||||
comm.data_bufs[r])[(params.rank * tot_tokens) + token_id] =
|
||||
(sum_variance);
|
||||
}
|
||||
}
|
||||
|
||||
bool done = false;
|
||||
float vars_all_ranks[NRanks];
|
||||
while (!done) {
|
||||
done = true;
|
||||
#pragma unroll
|
||||
for (int r = 0; r < NRanks; ++r) {
|
||||
vars_all_ranks[r] = ld_global_volatile(&reinterpret_cast<float*>(
|
||||
comm.data_bufs[params.rank])[(r * tot_tokens) + token_id]);
|
||||
done &= !is_neg_zero(vars_all_ranks[r]);
|
||||
}
|
||||
}
|
||||
sum_variance = 0.F;
|
||||
#pragma unroll
|
||||
for (int r = 0; r < NRanks; ++r) {
|
||||
sum_variance += vars_all_ranks[r];
|
||||
}
|
||||
|
||||
DType norm_weight[kElemsPerAccess<DType>];
|
||||
*reinterpret_cast<typename ElemsPerAccess<DType>::vec_type*>(norm_weight) =
|
||||
reinterpret_cast<typename ElemsPerAccess<DType>::vec_type*>(
|
||||
params.rms_gamma)[access_id_in_token];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kElemsPerAccess<DType>; ++i) {
|
||||
vals[i] = static_cast<DType>(
|
||||
static_cast<float>(vals[i]) *
|
||||
rsqrtf(
|
||||
(sum_variance / static_cast<float>(params.hidden_dim) / NRanks) +
|
||||
params.rms_eps) *
|
||||
static_cast<float>(norm_weight[i]));
|
||||
}
|
||||
|
||||
reinterpret_cast<float4*>(params.rms_norm_out)[idx] =
|
||||
*reinterpret_cast<float4*>(vals);
|
||||
}
|
||||
for (int idx = access_id; idx < clear_access; idx += access_stride) {
|
||||
reinterpret_cast<float4*>(comm.clear_buf)[idx] = clear_vec;
|
||||
}
|
||||
comm.update(params.size_q * NRanks);
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
asm volatile("griddepcontrol.launch_dependents;");
|
||||
#endif
|
||||
}
|
||||
|
||||
/**
|
||||
* Float4 variant: process 4 rows at once, allreduce variance sums as float4 for
|
||||
* better memory coalescing. sum_variance is always float; applies to all DTypes
|
||||
* (half, bf16, float). When tot_tokens % 4 != 0, the last group pads rows with
|
||||
* zeros; padded rows are not written to rms_norm_out. IsQK: when true, process
|
||||
* Q+K in one loop with doubled comm buffer; when false, single-matrix (Q only).
|
||||
*/
|
||||
template <typename DType, int NRanks, int OriginQDim, int OriginKDim>
|
||||
__global__ void __launch_bounds__(1024)
|
||||
minimax_reduce_qk_rms_kernel_lamport_float4(MiniMaxReduceRMSParams params) {
|
||||
// Compile-time per-rank dimensions
|
||||
constexpr int RankQDim = OriginQDim / NRanks;
|
||||
constexpr int RankKDim = OriginKDim / NRanks;
|
||||
// Threads needed to cover one row of Q / K with float4 accesses
|
||||
constexpr int ThreadsPerRowQ = RankQDim / kElemsPerAccess<DType>;
|
||||
constexpr int ThreadsPerRowK = RankKDim / kElemsPerAccess<DType>;
|
||||
// Number of warps dedicated to Q / K
|
||||
constexpr int NumWarpQ = (ThreadsPerRowQ + MINIMAX_REDUCE_RMS_WARP_SIZE - 1) /
|
||||
MINIMAX_REDUCE_RMS_WARP_SIZE;
|
||||
constexpr int NumWarpK = (ThreadsPerRowK + MINIMAX_REDUCE_RMS_WARP_SIZE - 1) /
|
||||
MINIMAX_REDUCE_RMS_WARP_SIZE;
|
||||
|
||||
int tot_tokens = params.size_q / RankQDim;
|
||||
int tot_groups = (tot_tokens + 3) / 4; // ceiling; last group may be partial
|
||||
|
||||
// Memory strides for strided qkv tensors (elements -> float4-access units)
|
||||
int access_stride_q = (params.stride_q > 0 ? params.stride_q : RankQDim) /
|
||||
kElemsPerAccess<DType>;
|
||||
int access_stride_k = (params.stride_k > 0 ? params.stride_k : RankKDim) /
|
||||
kElemsPerAccess<DType>;
|
||||
// Output strides: default to contiguous (hidden_dim / hidden_dim_k)
|
||||
int access_stride_q_out =
|
||||
(params.stride_q_out > 0 ? params.stride_q_out : params.hidden_dim) /
|
||||
kElemsPerAccess<DType>;
|
||||
int access_stride_k_out =
|
||||
(params.stride_k_out > 0 ? params.stride_k_out : params.hidden_dim_k) /
|
||||
kElemsPerAccess<DType>;
|
||||
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
namespace cg = cooperative_groups;
|
||||
cg::cluster_group cluster = cg::this_cluster();
|
||||
cg::grid_group grid = cg::this_grid();
|
||||
int group_id = grid.cluster_rank();
|
||||
int access_id_in_token = cluster.thread_rank();
|
||||
int group_stride = grid.num_clusters();
|
||||
#else
|
||||
int group_id = blockIdx.x;
|
||||
int access_id_in_token = threadIdx.x;
|
||||
int group_stride = gridDim.x;
|
||||
#endif
|
||||
|
||||
bool is_q = (access_id_in_token < NumWarpQ * MINIMAX_REDUCE_RMS_WARP_SIZE);
|
||||
int k_thread_idx =
|
||||
access_id_in_token - (NumWarpQ * MINIMAX_REDUCE_RMS_WARP_SIZE);
|
||||
bool is_valid_q = (access_id_in_token < ThreadsPerRowQ);
|
||||
bool is_valid_k = (k_thread_idx >= 0 && k_thread_idx < ThreadsPerRowK);
|
||||
float4 clear_vec = get_neg_zero();
|
||||
|
||||
// Shared memory for two-level block reduction and scale broadcast
|
||||
__shared__ float block_reduce_sum[4][MINIMAX_REDUCE_RMS_WARP_SIZE + 1];
|
||||
__shared__ float global_scale_q[4];
|
||||
__shared__ float global_scale_k[4];
|
||||
|
||||
LamportComm<NRanks> comm(params.workspace, params.rank);
|
||||
|
||||
DType norm_weight[kElemsPerAccess<DType>]{};
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
asm volatile("griddepcontrol.wait;");
|
||||
#endif
|
||||
if (is_q) {
|
||||
if (is_valid_q) {
|
||||
*reinterpret_cast<typename ElemsPerAccess<DType>::vec_type*>(
|
||||
norm_weight) =
|
||||
reinterpret_cast<typename ElemsPerAccess<DType>::vec_type const*>(
|
||||
params.rms_gamma)[access_id_in_token];
|
||||
}
|
||||
} else {
|
||||
if (is_valid_k) {
|
||||
*reinterpret_cast<typename ElemsPerAccess<DType>::vec_type*>(
|
||||
norm_weight) =
|
||||
reinterpret_cast<typename ElemsPerAccess<DType>::vec_type const*>(
|
||||
params.rms_gamma_k)[k_thread_idx];
|
||||
}
|
||||
}
|
||||
|
||||
// Main loop: process one group of 4 tokens per iteration.
|
||||
for (int g = group_id; g < tot_groups; g += group_stride) {
|
||||
alignas(16) DType vals[4][kElemsPerAccess<DType>]{};
|
||||
float warp_sum_variance[4]{0.F, 0.F, 0.F, 0.F};
|
||||
|
||||
if (is_q) {
|
||||
#pragma unroll
|
||||
for (int row = 0; row < 4; ++row) {
|
||||
int token_r = g * 4 + row;
|
||||
if (token_r >= tot_tokens || !is_valid_q) {
|
||||
continue;
|
||||
}
|
||||
int idx_r = token_r * access_stride_q + access_id_in_token;
|
||||
*reinterpret_cast<float4*>(&vals[row][0]) =
|
||||
reinterpret_cast<float4 const*>(params.allreduce_in)[idx_r];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kElemsPerAccess<DType>; ++i) {
|
||||
float x = static_cast<float>(vals[row][i]);
|
||||
warp_sum_variance[row] += x * x;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (int row = 0; row < 4; ++row) {
|
||||
int token_r = g * 4 + row;
|
||||
if (token_r >= tot_tokens || !is_valid_k) {
|
||||
continue;
|
||||
}
|
||||
int idx_r = token_r * access_stride_k + k_thread_idx;
|
||||
*reinterpret_cast<float4*>(&vals[row][0]) =
|
||||
reinterpret_cast<float4 const*>(params.allreduce_in_k)[idx_r];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kElemsPerAccess<DType>; ++i) {
|
||||
float x = static_cast<float>(vals[row][i]);
|
||||
warp_sum_variance[row] += x * x;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
local_warp_reduce_sum_array<MINIMAX_REDUCE_RMS_WARP_SIZE, float, 4>(
|
||||
warp_sum_variance);
|
||||
// Warp lane 0 writes its warp's partial sum to shared memory
|
||||
int lane = threadIdx.x & (MINIMAX_REDUCE_RMS_WARP_SIZE - 1);
|
||||
if (lane == 0) {
|
||||
#pragma unroll
|
||||
for (int t = 0; t < 4; ++t) {
|
||||
block_reduce_sum[t][threadIdx.x / MINIMAX_REDUCE_RMS_WARP_SIZE] =
|
||||
warp_sum_variance[t];
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int tid = threadIdx.x;
|
||||
|
||||
if (tid < MINIMAX_REDUCE_RMS_WARP_SIZE) {
|
||||
constexpr int kNumWarpQPow2 =
|
||||
(next_pow2(NumWarpQ) > NRanks) ? next_pow2(NumWarpQ) : NRanks;
|
||||
float local_sum[4];
|
||||
#pragma unroll
|
||||
for (int t = 0; t < 4; ++t) {
|
||||
local_sum[t] = (tid < NumWarpQ) ? block_reduce_sum[t][tid] : 0.F;
|
||||
}
|
||||
// After this, all kNumWarpQPow2 lanes (including tid 0..NRanks-1) have
|
||||
// the total Q sum-of-squares for all 4 tokens.
|
||||
local_warp_reduce_sum_array<kNumWarpQPow2, float, 4>(local_sum);
|
||||
|
||||
if (tid < NRanks) {
|
||||
#pragma unroll
|
||||
for (int t = 0; t < 4; ++t) {
|
||||
if (is_neg_zero(local_sum[t])) {
|
||||
local_sum[t] = 0.F;
|
||||
}
|
||||
}
|
||||
// Parallel push: thread tid writes this rank's Q sum to rank tid's buf
|
||||
reinterpret_cast<float4*>(
|
||||
comm.data_bufs[tid])[(params.rank * tot_groups * 2) + (2 * g)] =
|
||||
*reinterpret_cast<float4*>(local_sum);
|
||||
|
||||
// Parallel pull: thread tid reads rank tid's contribution from
|
||||
// this rank's (params.rank's) buffer
|
||||
bool done = false;
|
||||
float4 var_all_ranks;
|
||||
while (!done) {
|
||||
done = true;
|
||||
var_all_ranks = ld_global_volatile(&reinterpret_cast<float4*>(
|
||||
comm.data_bufs[params.rank])[(tid * tot_groups * 2) + (2 * g)]);
|
||||
done &= !is_neg_zero(var_all_ranks);
|
||||
}
|
||||
|
||||
// Warp-level allreduce: each of the NRanks threads holds one rank's
|
||||
// partial sum; after this all NRanks threads have the global total.
|
||||
constexpr uint32_t kQActiveMask = (1u << NRanks) - 1u;
|
||||
local_warp_reduce_sum_array<NRanks, float, 4>(
|
||||
reinterpret_cast<float*>(&var_all_ranks), kQActiveMask);
|
||||
|
||||
// Thread 0 computes rsqrt with compile-time Dim and writes to smem
|
||||
if (tid == 0) {
|
||||
*reinterpret_cast<float4*>(global_scale_q) =
|
||||
rms_rsqrt<OriginQDim>(var_all_ranks, params.rms_eps);
|
||||
}
|
||||
}
|
||||
} else if (tid >= MINIMAX_REDUCE_RMS_WARP_SIZE * NumWarpQ &&
|
||||
tid < MINIMAX_REDUCE_RMS_WARP_SIZE * (NumWarpQ + 1)) {
|
||||
// --- K leader warp ---
|
||||
constexpr int kNumWarpKPow2 =
|
||||
(next_pow2(NumWarpK) > NRanks) ? next_pow2(NumWarpK) : NRanks;
|
||||
float local_sum[4];
|
||||
#pragma unroll
|
||||
for (int t = 0; t < 4; ++t) {
|
||||
local_sum[t] = (k_thread_idx < NumWarpK)
|
||||
? block_reduce_sum[t][NumWarpQ + k_thread_idx]
|
||||
: 0.F;
|
||||
}
|
||||
local_warp_reduce_sum_array<kNumWarpKPow2, float, 4>(local_sum);
|
||||
|
||||
if (k_thread_idx < NRanks) {
|
||||
#pragma unroll
|
||||
for (int t = 0; t < 4; ++t) {
|
||||
if (is_neg_zero(local_sum[t])) {
|
||||
local_sum[t] = 0.F;
|
||||
}
|
||||
}
|
||||
reinterpret_cast<float4*>(
|
||||
comm.data_bufs[k_thread_idx])[(params.rank * tot_groups * 2) +
|
||||
(2 * g + 1)] =
|
||||
*reinterpret_cast<float4*>(local_sum);
|
||||
|
||||
bool done = false;
|
||||
float4 var_all_ranks;
|
||||
while (!done) {
|
||||
done = true;
|
||||
var_all_ranks = ld_global_volatile(&reinterpret_cast<float4*>(
|
||||
comm.data_bufs[params.rank])[(k_thread_idx * tot_groups * 2) +
|
||||
(2 * g + 1)]);
|
||||
done &= !is_neg_zero(var_all_ranks);
|
||||
}
|
||||
|
||||
constexpr uint32_t kKActiveMask = (1u << NRanks) - 1u;
|
||||
local_warp_reduce_sum_array<NRanks, float, 4>(
|
||||
reinterpret_cast<float*>(&var_all_ranks), kKActiveMask);
|
||||
|
||||
if (k_thread_idx == 0) {
|
||||
*reinterpret_cast<float4*>(global_scale_k) =
|
||||
rms_rsqrt<OriginKDim>(var_all_ranks, params.rms_eps);
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (is_q) {
|
||||
#pragma unroll
|
||||
for (int t = 0; t < 4; ++t) {
|
||||
warp_sum_variance[t] = global_scale_q[t];
|
||||
}
|
||||
#pragma unroll
|
||||
for (int r = 0; r < 4; ++r) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kElemsPerAccess<DType>; ++i) {
|
||||
vals[r][i] = static_cast<DType>(static_cast<float>(vals[r][i]) *
|
||||
warp_sum_variance[r] *
|
||||
static_cast<float>(norm_weight[i]));
|
||||
}
|
||||
int token_r = g * 4 + r;
|
||||
if (token_r >= tot_tokens || !is_valid_q) {
|
||||
continue;
|
||||
}
|
||||
int idx_out = token_r * access_stride_q_out + access_id_in_token;
|
||||
reinterpret_cast<float4*>(params.rms_norm_out)[idx_out] =
|
||||
*reinterpret_cast<float4*>(&vals[r][0]);
|
||||
}
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (int t = 0; t < 4; ++t) {
|
||||
warp_sum_variance[t] = global_scale_k[t];
|
||||
}
|
||||
#pragma unroll
|
||||
for (int r = 0; r < 4; ++r) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kElemsPerAccess<DType>; ++i) {
|
||||
vals[r][i] = static_cast<DType>(static_cast<float>(vals[r][i]) *
|
||||
warp_sum_variance[r] *
|
||||
static_cast<float>(norm_weight[i]));
|
||||
}
|
||||
int token_r = g * 4 + r;
|
||||
if (token_r >= tot_tokens || !is_valid_k) {
|
||||
continue;
|
||||
}
|
||||
int idx_out = token_r * access_stride_k_out + k_thread_idx;
|
||||
reinterpret_cast<float4*>(params.rms_norm_out_k)[idx_out] =
|
||||
*reinterpret_cast<float4*>(&vals[r][0]);
|
||||
}
|
||||
}
|
||||
} // end group loop
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
|
||||
asm volatile("griddepcontrol.launch_dependents;");
|
||||
#endif
|
||||
|
||||
int clear_access = static_cast<int>(comm.clear_size / kElemsPerAccess<DType>);
|
||||
int clear_stride = group_stride * blockDim.x;
|
||||
for (int idx = group_id * blockDim.x + threadIdx.x; idx < clear_access;
|
||||
idx += clear_stride) {
|
||||
reinterpret_cast<float4*>(comm.clear_buf)[idx] = clear_vec;
|
||||
}
|
||||
|
||||
comm.update(static_cast<int64_t>(2) * tot_groups * kElemsPerAccess<DType> *
|
||||
NRanks);
|
||||
}
|
||||
|
||||
int get_sm_count() {
|
||||
static int sm_count = 0;
|
||||
if (sm_count == 0) {
|
||||
int device_id;
|
||||
CUDA_CHECK(cudaGetDevice(&device_id));
|
||||
cudaDeviceProp device_prop;
|
||||
cudaGetDeviceProperties(&device_prop, device_id);
|
||||
sm_count = device_prop.multiProcessorCount;
|
||||
}
|
||||
return sm_count;
|
||||
}
|
||||
|
||||
inline int getSMVersion(bool queryRealSmArch = false) {
|
||||
int device{-1};
|
||||
CUDA_CHECK(cudaGetDevice(&device));
|
||||
int sm_major = 0;
|
||||
int sm_minor = 0;
|
||||
CUDA_CHECK(cudaDeviceGetAttribute(&sm_major,
|
||||
cudaDevAttrComputeCapabilityMajor, device));
|
||||
CUDA_CHECK(cudaDeviceGetAttribute(&sm_minor,
|
||||
cudaDevAttrComputeCapabilityMinor, device));
|
||||
int sm = sm_major * 10 + sm_minor;
|
||||
if (sm == 121 && !queryRealSmArch) {
|
||||
return 120;
|
||||
}
|
||||
return sm;
|
||||
}
|
||||
|
||||
template <typename KernelFunc>
|
||||
int get_max_active_blocks(KernelFunc kernel, int block_size,
|
||||
int dynamic_smem = 0) {
|
||||
int max_active = 0;
|
||||
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
|
||||
&max_active, kernel, block_size, dynamic_smem));
|
||||
return std::max(max_active, 1);
|
||||
}
|
||||
|
||||
template <typename DType, int NRanks>
|
||||
void minimax_reduce_rms_kernel_launcher(MiniMaxReduceRMSParams const& params) {
|
||||
static int SM = getSMVersion();
|
||||
int token_num = params.size_q / params.hidden_dim;
|
||||
int sm_count = get_sm_count();
|
||||
int cluster_size = 1;
|
||||
int cluster_num = token_num;
|
||||
int threads_per_token = params.hidden_dim / kElemsPerAccess<DType>;
|
||||
int block_size = threads_per_token;
|
||||
|
||||
int max_blocks_per_sm = get_max_active_blocks(
|
||||
minimax_reduce_rms_kernel_lamport<DType, NRanks>, block_size);
|
||||
int max_grid = max_blocks_per_sm * sm_count;
|
||||
|
||||
int grid_size =
|
||||
(std::min(max_grid, cluster_num * cluster_size) / cluster_size) *
|
||||
cluster_size;
|
||||
|
||||
cudaLaunchConfig_t cfg;
|
||||
cfg.gridDim = grid_size;
|
||||
cfg.blockDim = block_size;
|
||||
cfg.dynamicSmemBytes = 0;
|
||||
cfg.stream = params.stream;
|
||||
|
||||
cudaLaunchAttribute attribute[2];
|
||||
attribute[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attribute[0].val.programmaticStreamSerializationAllowed = 1;
|
||||
attribute[1].id = cudaLaunchAttributeClusterDimension;
|
||||
attribute[1].val.clusterDim.x = cluster_size;
|
||||
attribute[1].val.clusterDim.y = 1;
|
||||
attribute[1].val.clusterDim.z = 1;
|
||||
cfg.attrs = attribute;
|
||||
cfg.numAttrs = SM >= 90 ? 2 : 0;
|
||||
|
||||
CUDA_CHECK(cudaLaunchKernelEx(
|
||||
&cfg, minimax_reduce_rms_kernel_lamport<DType, NRanks>, params));
|
||||
}
|
||||
|
||||
template <typename DType, int NRanks, int OriginQDim, int OriginKDim>
|
||||
void minimax_reduce_rms_kernel_launcher_float4(
|
||||
MiniMaxReduceRMSParams const& params) {
|
||||
TORCH_CHECK(params.size_q % params.hidden_dim == 0);
|
||||
TORCH_CHECK(params.hidden_dim % kElemsPerAccess<DType> == 0);
|
||||
if (params.stride_q > 0) {
|
||||
TORCH_CHECK(params.stride_q % kElemsPerAccess<DType> == 0);
|
||||
}
|
||||
TORCH_CHECK(params.allreduce_in_k != nullptr,
|
||||
"float4 QK kernel requires K input");
|
||||
TORCH_CHECK(params.hidden_dim >= params.hidden_dim_k);
|
||||
TORCH_CHECK(params.size_k % params.hidden_dim_k == 0);
|
||||
TORCH_CHECK(params.hidden_dim_k % kElemsPerAccess<DType> == 0);
|
||||
TORCH_CHECK(params.size_q / params.hidden_dim ==
|
||||
params.size_k / params.hidden_dim_k);
|
||||
if (params.stride_k > 0) {
|
||||
TORCH_CHECK(params.stride_k % kElemsPerAccess<DType> == 0);
|
||||
}
|
||||
|
||||
int token_num = params.size_q / params.hidden_dim;
|
||||
int tot_groups = (token_num + 3) / 4;
|
||||
if (tot_groups == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
static int SM = getSMVersion();
|
||||
int sm_count = get_sm_count();
|
||||
int cluster_size = 1;
|
||||
int cluster_num = tot_groups;
|
||||
|
||||
int access_per_row_q = params.hidden_dim / kElemsPerAccess<DType>;
|
||||
int access_per_row_k = params.hidden_dim_k / kElemsPerAccess<DType>;
|
||||
|
||||
// Round each section up to a warp boundary
|
||||
auto divUp = [](int a, int b) { return (a + b - 1) / b * b; };
|
||||
int block_size = divUp(access_per_row_q, MINIMAX_REDUCE_RMS_WARP_SIZE) +
|
||||
divUp(access_per_row_k, MINIMAX_REDUCE_RMS_WARP_SIZE);
|
||||
|
||||
auto kfn =
|
||||
minimax_reduce_qk_rms_kernel_lamport_float4<DType, NRanks, OriginQDim,
|
||||
OriginKDim>;
|
||||
|
||||
int max_blocks_per_sm = get_max_active_blocks(kfn, block_size);
|
||||
int max_grid = max_blocks_per_sm * sm_count;
|
||||
int grid_size =
|
||||
(std::min(max_grid, cluster_num * cluster_size) / cluster_size) *
|
||||
cluster_size;
|
||||
|
||||
cudaLaunchConfig_t cfg;
|
||||
cfg.gridDim = grid_size;
|
||||
cfg.blockDim = block_size;
|
||||
cfg.dynamicSmemBytes = 0;
|
||||
cfg.stream = params.stream;
|
||||
|
||||
cudaLaunchAttribute attribute[2];
|
||||
attribute[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
|
||||
attribute[0].val.programmaticStreamSerializationAllowed = 1;
|
||||
attribute[1].id = cudaLaunchAttributeClusterDimension;
|
||||
attribute[1].val.clusterDim.x = cluster_size;
|
||||
attribute[1].val.clusterDim.y = 1;
|
||||
attribute[1].val.clusterDim.z = 1;
|
||||
cfg.attrs = attribute;
|
||||
cfg.numAttrs = SM >= 90 ? 2 : 0;
|
||||
|
||||
CUDA_CHECK(cudaLaunchKernelEx(&cfg, kfn, params));
|
||||
}
|
||||
|
||||
template <int NRanks>
|
||||
void dispatch_dtype(MiniMaxReduceRMSParams const& params) {
|
||||
// Use the optimized QK float4 kernel when:
|
||||
// - K input is present, AND
|
||||
// - the full (NRanks * per-rank) dimensions match the MiniMax M2 shape.
|
||||
// Otherwise fall back to the scalar kernel.
|
||||
bool use_float4 = (params.allreduce_in_k != nullptr) &&
|
||||
(params.hidden_dim * params.nranks == 6144) &&
|
||||
(params.hidden_dim_k * params.nranks == 1024);
|
||||
|
||||
if (params.dtype == at::ScalarType::Half) {
|
||||
if (use_float4) {
|
||||
minimax_reduce_rms_kernel_launcher_float4<half, NRanks, 6144, 1024>(
|
||||
params);
|
||||
} else {
|
||||
minimax_reduce_rms_kernel_launcher<half, NRanks>(params);
|
||||
}
|
||||
} else if (params.dtype == at::ScalarType::BFloat16) {
|
||||
if (use_float4) {
|
||||
minimax_reduce_rms_kernel_launcher_float4<__nv_bfloat16, NRanks, 6144,
|
||||
1024>(params);
|
||||
} else {
|
||||
minimax_reduce_rms_kernel_launcher<__nv_bfloat16, NRanks>(params);
|
||||
}
|
||||
} else if (params.dtype == at::ScalarType::Float) {
|
||||
if (use_float4) {
|
||||
minimax_reduce_rms_kernel_launcher_float4<float, NRanks, 6144, 1024>(
|
||||
params);
|
||||
} else {
|
||||
minimax_reduce_rms_kernel_launcher<float, NRanks>(params);
|
||||
}
|
||||
} else {
|
||||
TORCH_CHECK(false, "Unsupported data type for minimax_reduce_rms_op");
|
||||
}
|
||||
}
|
||||
|
||||
void minimax_reduce_rms_op(MiniMaxReduceRMSParams const& params) {
|
||||
if (params.nranks == 2) {
|
||||
dispatch_dtype<2>(params);
|
||||
} else if (params.nranks == 4) {
|
||||
dispatch_dtype<4>(params);
|
||||
} else if (params.nranks == 8) {
|
||||
dispatch_dtype<8>(params);
|
||||
} else if (params.nranks == 16) {
|
||||
dispatch_dtype<16>(params);
|
||||
} else {
|
||||
TORCH_CHECK(false, "minimax_reduce_rms_op: unsupported ranks number!");
|
||||
}
|
||||
}
|
||||
} // namespace tensorrt_llm
|
||||
} // namespace vllm
|
||||
|
||||
torch::Tensor minimax_allreduce_rms(torch::Tensor const& input,
|
||||
torch::Tensor const& norm_weight,
|
||||
torch::Tensor workspace, int64_t const rank,
|
||||
int64_t const nranks, double const eps) {
|
||||
auto allreduce_params = vllm::tensorrt_llm::MiniMaxReduceRMSParams();
|
||||
|
||||
allreduce_params.nranks = static_cast<int>(nranks);
|
||||
allreduce_params.rank = static_cast<int>(rank);
|
||||
allreduce_params.dtype = input.scalar_type();
|
||||
allreduce_params.size_q = static_cast<int>(input.numel());
|
||||
allreduce_params.hidden_dim = static_cast<int>(input.size(-1));
|
||||
allreduce_params.stride_q = allreduce_params.hidden_dim;
|
||||
allreduce_params.workspace =
|
||||
reinterpret_cast<void**>(workspace.mutable_data_ptr());
|
||||
allreduce_params.allreduce_in = input.data_ptr();
|
||||
allreduce_params.rms_gamma = norm_weight.data_ptr();
|
||||
allreduce_params.rms_eps = static_cast<float>(eps);
|
||||
allreduce_params.stream = at::cuda::getCurrentCUDAStream(input.get_device());
|
||||
|
||||
torch::Tensor rms_norm_out = torch::empty_like(input);
|
||||
allreduce_params.rms_norm_out = rms_norm_out.mutable_data_ptr();
|
||||
|
||||
vllm::tensorrt_llm::minimax_reduce_rms_op(allreduce_params);
|
||||
|
||||
return rms_norm_out;
|
||||
}
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor> minimax_allreduce_rms_qk(
|
||||
torch::Tensor qkv, torch::Tensor const& norm_weight_q,
|
||||
torch::Tensor const& norm_weight_k, torch::Tensor workspace,
|
||||
int64_t const q_size, int64_t const kv_size, int64_t const rank,
|
||||
int64_t const nranks, double const eps) {
|
||||
TORCH_CHECK(qkv.dim() == 2, "minimax_allreduce_rms_qk: qkv must be 2D");
|
||||
TORCH_CHECK(qkv.is_contiguous(),
|
||||
"minimax_allreduce_rms_qk: qkv must be contiguous");
|
||||
int64_t qkv_dim = qkv.size(-1);
|
||||
TORCH_CHECK(qkv_dim == q_size + 2 * kv_size,
|
||||
"minimax_allreduce_rms_qk: qkv last dim must equal "
|
||||
"q_size + 2 * kv_size");
|
||||
TORCH_CHECK(rank < nranks,
|
||||
"minimax_allreduce_rms_qk: rank must be less than nranks");
|
||||
|
||||
int64_t num_tokens = qkv.size(0);
|
||||
int elem_bytes = qkv.element_size();
|
||||
|
||||
torch::Tensor q_out = torch::empty({num_tokens, q_size}, qkv.options());
|
||||
torch::Tensor k_out = torch::empty({num_tokens, kv_size}, qkv.options());
|
||||
|
||||
auto params = vllm::tensorrt_llm::MiniMaxReduceRMSParams();
|
||||
params.nranks = static_cast<int>(nranks);
|
||||
params.rank = static_cast<int>(rank);
|
||||
params.dtype = qkv.scalar_type();
|
||||
params.size_q = static_cast<int>(num_tokens * q_size);
|
||||
params.hidden_dim = static_cast<int>(q_size);
|
||||
params.size_k = static_cast<int>(num_tokens * kv_size);
|
||||
params.hidden_dim_k = static_cast<int>(kv_size);
|
||||
params.stride_q = static_cast<int>(qkv_dim);
|
||||
params.stride_k = static_cast<int>(qkv_dim);
|
||||
params.stride_q_out = 0; // q_out is contiguous; kernel uses hidden_dim
|
||||
params.stride_k_out = 0; // k_out is contiguous; kernel uses hidden_dim_k
|
||||
params.workspace = reinterpret_cast<void**>(workspace.mutable_data_ptr());
|
||||
|
||||
uint8_t* base = static_cast<uint8_t*>(qkv.data_ptr());
|
||||
params.allreduce_in = base;
|
||||
params.allreduce_in_k = base + q_size * elem_bytes;
|
||||
params.rms_gamma = norm_weight_q.data_ptr();
|
||||
params.rms_gamma_k = norm_weight_k.data_ptr();
|
||||
params.rms_eps = static_cast<float>(eps);
|
||||
params.stream = at::cuda::getCurrentCUDAStream(qkv.get_device());
|
||||
|
||||
params.rms_norm_out = q_out.mutable_data_ptr();
|
||||
params.rms_norm_out_k = k_out.mutable_data_ptr();
|
||||
|
||||
vllm::tensorrt_llm::minimax_reduce_rms_op(params);
|
||||
return {q_out, k_out};
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
/*
|
||||
* Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
|
||||
#include <torch/types.h>
|
||||
|
||||
namespace vllm {
|
||||
namespace tensorrt_llm {
|
||||
|
||||
template <typename DType>
|
||||
struct ElemsPerAccess;
|
||||
|
||||
template <>
|
||||
struct ElemsPerAccess<half> {
|
||||
static constexpr int value = 8;
|
||||
using vec_type = float4;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct ElemsPerAccess<nv_bfloat16> {
|
||||
static constexpr int value = 8;
|
||||
using vec_type = float4;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct ElemsPerAccess<float> {
|
||||
static constexpr int value = 4;
|
||||
using vec_type = float4;
|
||||
};
|
||||
|
||||
template <typename DType>
|
||||
static constexpr int kElemsPerAccess = ElemsPerAccess<DType>::value;
|
||||
|
||||
struct MiniMaxReduceRMSParams {
|
||||
int nranks{};
|
||||
int rank{};
|
||||
at::ScalarType dtype{at::ScalarType::Undefined};
|
||||
int size_q{};
|
||||
int hidden_dim{};
|
||||
int size_k{};
|
||||
int hidden_dim_k{};
|
||||
int stride_q{}; // row stride for q input (elements); when > hidden_dim,
|
||||
// q is part of a wider qkv tensor
|
||||
int stride_k{}; // row stride for k input (elements); when > hidden_dim_k,
|
||||
// k is part of a wider qkv tensor
|
||||
int stride_q_out{}; // row stride for q output (elements); 0 = contiguous
|
||||
int stride_k_out{}; // row stride for k output (elements); 0 = contiguous
|
||||
void** workspace{};
|
||||
void* allreduce_in{};
|
||||
void* rms_norm_out{};
|
||||
void* rms_gamma{};
|
||||
void* allreduce_in_k{};
|
||||
void* rms_norm_out_k{};
|
||||
void* rms_gamma_k{};
|
||||
float rms_eps{};
|
||||
cudaStream_t stream{};
|
||||
};
|
||||
|
||||
void minimax_reduce_rms_op(MiniMaxReduceRMSParams const& params);
|
||||
|
||||
} // namespace tensorrt_llm
|
||||
} // namespace vllm
|
||||
+18
-7
@@ -96,7 +96,8 @@ void fused_qk_norm_rope(torch::Tensor& qkv, int64_t num_heads_q,
|
||||
int64_t num_heads_k, int64_t num_heads_v,
|
||||
int64_t head_dim, double eps, torch::Tensor& q_weight,
|
||||
torch::Tensor& k_weight, torch::Tensor& cos_sin_cache,
|
||||
bool is_neox, torch::Tensor& position_ids);
|
||||
bool is_neox, torch::Tensor& position_ids,
|
||||
int64_t forced_token_heads_per_warp);
|
||||
|
||||
void apply_repetition_penalties_(torch::Tensor& logits,
|
||||
const torch::Tensor& prompt_mask,
|
||||
@@ -114,9 +115,9 @@ void top_k_per_row_decode(const torch::Tensor& logits, int64_t next_n,
|
||||
int64_t numRows, int64_t stride0, int64_t stride1,
|
||||
int64_t topK);
|
||||
|
||||
void large_context_topk(const torch::Tensor& score, torch::Tensor& indices,
|
||||
const torch::Tensor& lengths,
|
||||
std::optional<torch::Tensor> row_starts_opt);
|
||||
void persistent_topk(const torch::Tensor& logits, const torch::Tensor& lengths,
|
||||
torch::Tensor& output, torch::Tensor& workspace, int64_t k,
|
||||
int64_t max_seq_len);
|
||||
|
||||
void rms_norm_static_fp8_quant(torch::Tensor& out, torch::Tensor& input,
|
||||
torch::Tensor& weight, torch::Tensor& scale,
|
||||
@@ -143,13 +144,11 @@ void rms_norm_per_block_quant(torch::Tensor& out, torch::Tensor const& input,
|
||||
std::optional<torch::Tensor> residual,
|
||||
int64_t group_size, bool is_scale_transposed);
|
||||
|
||||
#ifndef USE_ROCM
|
||||
void silu_and_mul_per_block_quant(torch::Tensor& out,
|
||||
torch::Tensor const& input,
|
||||
torch::Tensor& scales, int64_t group_size,
|
||||
std::optional<torch::Tensor> scale_ub,
|
||||
bool is_scale_transposed);
|
||||
#endif
|
||||
|
||||
void rotary_embedding(torch::Tensor& positions, torch::Tensor& query,
|
||||
std::optional<torch::Tensor> key, int64_t head_size,
|
||||
@@ -310,4 +309,16 @@ int64_t qr_max_size();
|
||||
#ifndef USE_ROCM
|
||||
void dsv3_fused_a_gemm(torch::Tensor& output, torch::Tensor const& mat_a,
|
||||
torch::Tensor const& mat_b);
|
||||
#endif
|
||||
#endif
|
||||
|
||||
#ifndef USE_ROCM
|
||||
torch::Tensor minimax_allreduce_rms(torch::Tensor const& input,
|
||||
torch::Tensor const& norm_weight,
|
||||
torch::Tensor workspace, int64_t const rank,
|
||||
int64_t const nranks, double const eps);
|
||||
std::tuple<torch::Tensor, torch::Tensor> minimax_allreduce_rms_qk(
|
||||
torch::Tensor qkv, torch::Tensor const& norm_weight_q,
|
||||
torch::Tensor const& norm_weight_k, torch::Tensor workspace,
|
||||
int64_t const q_size, int64_t const kv_size, int64_t const rank,
|
||||
int64_t const nranks, double const eps);
|
||||
#endif
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -6,7 +6,7 @@
|
||||
|
||||
#include "libtorch_stable/quantization/vectorization.cuh"
|
||||
// TODO(luka/varun):refactor common.cuh to use this file instead
|
||||
#include "quantization/w8a8/fp8/common.cuh"
|
||||
#include "../w8a8/fp8/common.cuh"
|
||||
|
||||
namespace vllm {
|
||||
|
||||
|
||||
@@ -639,7 +639,9 @@ __inline__ __device__ Tout scaled_convert(const Tin& x, const float scale) {
|
||||
// function with template<typename scalar_t, typename cache_t,
|
||||
// Fp8KVCacheDataType kv_dt>.
|
||||
#define DISPATCH_BY_KV_CACHE_DTYPE(SRC_DTYPE, KV_DTYPE, FN) \
|
||||
if (KV_DTYPE == "auto") { \
|
||||
vllm::Fp8KVCacheDataType KV_CACHE_DTYPE = \
|
||||
vllm::get_fp8_kv_cache_data_type(KV_DTYPE); \
|
||||
if (KV_CACHE_DTYPE == vllm::Fp8KVCacheDataType::kAuto) { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, float, vllm::Fp8KVCacheDataType::kAuto); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
@@ -649,21 +651,18 @@ __inline__ __device__ Tout scaled_convert(const Tin& x, const float scale) {
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else { \
|
||||
if (KV_DTYPE == "fp8" || KV_DTYPE == "fp8_e4m3") { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
FN(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::BFloat16) { \
|
||||
FN(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else { \
|
||||
TORCH_CHECK(false, \
|
||||
"Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else if (KV_CACHE_DTYPE == vllm::Fp8KVCacheDataType::kFp8E4M3) { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
FN(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::BFloat16) { \
|
||||
FN(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported data type of kv cache: ", KV_DTYPE); \
|
||||
TORCH_CHECK(false, "Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported data type of kv cache: ", KV_DTYPE); \
|
||||
}
|
||||
|
||||
} // namespace fp8
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#pragma once
|
||||
|
||||
#include "libtorch_stable/quantization/vectorization.cuh"
|
||||
#include "quantization/utils.cuh"
|
||||
#include "../../utils.cuh"
|
||||
|
||||
#include <cmath>
|
||||
|
||||
|
||||
@@ -543,7 +543,9 @@ __inline__ __device__ Tout scaled_convert(const Tin& x, const float scale) {
|
||||
// function with template<typename scalar_t, typename cache_t,
|
||||
// Fp8KVCacheDataType kv_dt>.
|
||||
#define DISPATCH_BY_KV_CACHE_DTYPE(SRC_DTYPE, KV_DTYPE, FN) \
|
||||
if (KV_DTYPE == "auto") { \
|
||||
vllm::Fp8KVCacheDataType KV_CACHE_DTYPE = \
|
||||
vllm::get_fp8_kv_cache_data_type(KV_DTYPE); \
|
||||
if (KV_CACHE_DTYPE == vllm::Fp8KVCacheDataType::kAuto) { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, float, vllm::Fp8KVCacheDataType::kAuto); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
@@ -553,43 +555,28 @@ __inline__ __device__ Tout scaled_convert(const Tin& x, const float scale) {
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else { \
|
||||
if (KV_DTYPE == "fp8" || KV_DTYPE == "fp8_e4m3") { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
FN(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::BFloat16) { \
|
||||
FN(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else { \
|
||||
TORCH_CHECK(false, \
|
||||
"Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else if (KV_DTYPE == "fp8_e5m2") { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E5M2); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
FN(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E5M2); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::BFloat16) { \
|
||||
FN(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kFp8E5M2); \
|
||||
} else { \
|
||||
TORCH_CHECK(false, \
|
||||
"Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else if (KV_DTYPE == "fp8_ds_mla") { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
FN(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::BFloat16) { \
|
||||
FN(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else { \
|
||||
TORCH_CHECK(false, \
|
||||
"Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else if (KV_CACHE_DTYPE == vllm::Fp8KVCacheDataType::kFp8E4M3) { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
FN(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::BFloat16) { \
|
||||
FN(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3); \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported data type of kv cache: ", KV_DTYPE); \
|
||||
TORCH_CHECK(false, "Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else if (KV_CACHE_DTYPE == vllm::Fp8KVCacheDataType::kFp8E5M2) { \
|
||||
if (SRC_DTYPE == at::ScalarType::Float) { \
|
||||
FN(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E5M2); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::Half) { \
|
||||
FN(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E5M2); \
|
||||
} else if (SRC_DTYPE == at::ScalarType::BFloat16) { \
|
||||
FN(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kFp8E5M2); \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported input type of kv cache: ", SRC_DTYPE); \
|
||||
} \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported data type of kv cache: ", KV_DTYPE); \
|
||||
}
|
||||
|
||||
} // namespace fp8
|
||||
|
||||
+24
-9
@@ -564,8 +564,9 @@ template <int kNumThreadsPerBlock, bool useRadixSort,
|
||||
bool multipleBlocksPerRow = false, bool mergeBlocks = false>
|
||||
static __global__ __launch_bounds__(kNumThreadsPerBlock) void topKPerRowDecode(
|
||||
const float* logits, const int* seqLens, int* outIndices, int stride0,
|
||||
int stride1, const int topK, int next_n, float* outLogits = nullptr,
|
||||
const int numBlocksToMerge = 0, const int* indices = nullptr) {
|
||||
int stride1, const int topK, int next_n, int seqLensIs2D = 0,
|
||||
float* outLogits = nullptr, const int numBlocksToMerge = 0,
|
||||
const int* indices = nullptr) {
|
||||
// The number of bins in the histogram.
|
||||
static constexpr int kNumBins = 2048;
|
||||
|
||||
@@ -574,8 +575,16 @@ static __global__ __launch_bounds__(kNumThreadsPerBlock) void topKPerRowDecode(
|
||||
|
||||
// The range of logits within the row.
|
||||
int rowStart = 0;
|
||||
int seq_len = seqLens[rowIdx / next_n];
|
||||
int rowEnd = max(0, seq_len - next_n + (rowIdx % next_n) + 1);
|
||||
int batch_idx = rowIdx / next_n;
|
||||
int next_n_idx = rowIdx % next_n;
|
||||
// seqLensIs2D=0: 1D seqLens — all rows in a batch share the same seq_len;
|
||||
// kernel computes per-row effective length via offset.
|
||||
// seqLensIs2D=1: 2D seqLens — each logit row has its own pre-computed
|
||||
// effective length (flat index rowIdx = b*next_n + j maps
|
||||
// directly to seqLens[b, j] in C-contiguous layout).
|
||||
int seq_len = seqLensIs2D ? seqLens[rowIdx] : seqLens[batch_idx];
|
||||
int rowEnd =
|
||||
seqLensIs2D ? max(0, seq_len) : max(0, seq_len - next_n + next_n_idx + 1);
|
||||
|
||||
// Local pointers to this block
|
||||
if constexpr (!multipleBlocksPerRow && !mergeBlocks) {
|
||||
@@ -653,6 +662,11 @@ void top_k_per_row_decode(const torch::Tensor& logits, int64_t next_n,
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
const auto numColumns = logits.size(1);
|
||||
|
||||
// True if seqLens is 2D (B, next_n): each logit row has its own pre-computed
|
||||
// effective seq_len. False if seqLens is 1D (B,): all rows in a batch share
|
||||
// the same seq_len and the kernel computes the per-row offset itself.
|
||||
int seqLensIs2D = seqLens.dim() == 2 ? 1 : 0;
|
||||
|
||||
if (numColumns < kSortingAlgorithmThreshold) {
|
||||
// Use insertion sort
|
||||
vllm::topKPerRowDecode<kNumThreadsPerBlock, false>
|
||||
@@ -660,7 +674,7 @@ void top_k_per_row_decode(const torch::Tensor& logits, int64_t next_n,
|
||||
logits.data_ptr<float>(), seqLens.data_ptr<int>(),
|
||||
indices.data_ptr<int>(), static_cast<int>(stride0),
|
||||
static_cast<int>(stride1), static_cast<int>(topK),
|
||||
static_cast<int>(next_n));
|
||||
static_cast<int>(next_n), seqLensIs2D);
|
||||
} else if (numColumns < kSplitWorkThreshold) {
|
||||
// From this threshold, use radix sort instead
|
||||
vllm::topKPerRowDecode<kNumThreadsPerBlock, true>
|
||||
@@ -668,7 +682,7 @@ void top_k_per_row_decode(const torch::Tensor& logits, int64_t next_n,
|
||||
logits.data_ptr<float>(), seqLens.data_ptr<int>(),
|
||||
indices.data_ptr<int>(), static_cast<int>(stride0),
|
||||
static_cast<int>(stride1), static_cast<int>(topK),
|
||||
static_cast<int>(next_n));
|
||||
static_cast<int>(next_n), seqLensIs2D);
|
||||
} else {
|
||||
// Long sequences are run in two steps
|
||||
constexpr auto multipleBlocksPerRowConfig = 10;
|
||||
@@ -686,15 +700,16 @@ void top_k_per_row_decode(const torch::Tensor& logits, int64_t next_n,
|
||||
logits.data_ptr<float>(), seqLens.data_ptr<int>(),
|
||||
outIndicesAux.data_ptr<int>(), static_cast<int>(stride0),
|
||||
static_cast<int>(stride1), static_cast<int>(topK),
|
||||
static_cast<int>(next_n), outLogitsAux.data_ptr<float>());
|
||||
static_cast<int>(next_n), seqLensIs2D,
|
||||
outLogitsAux.data_ptr<float>());
|
||||
|
||||
constexpr int kNumThreadsPerBlockMerge = 1024;
|
||||
vllm::topKPerRowDecode<kNumThreadsPerBlockMerge, true, false, true>
|
||||
<<<numRows, kNumThreadsPerBlockMerge, topK * sizeof(int32_t), stream>>>(
|
||||
outLogitsAux.data_ptr<float>(), seqLens.data_ptr<int>(),
|
||||
indices.data_ptr<int>(), multipleBlocksPerRowConfig * topK, 1,
|
||||
static_cast<int>(topK), static_cast<int>(next_n), nullptr,
|
||||
multipleBlocksPerRowConfig, outIndicesAux.data_ptr<int>());
|
||||
static_cast<int>(topK), static_cast<int>(next_n), seqLensIs2D,
|
||||
nullptr, multipleBlocksPerRowConfig, outIndicesAux.data_ptr<int>());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,200 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Benchmark: Lamport all-gather vs NCCL."""
|
||||
|
||||
import ctypes
|
||||
import os
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
_cudart = ctypes.CDLL("libcudart.so")
|
||||
IPC = 64
|
||||
|
||||
|
||||
def _cc(r):
|
||||
if r:
|
||||
raise RuntimeError(f"err {r}")
|
||||
|
||||
|
||||
def ipc_buf(sz, rank, ws):
|
||||
p = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaMalloc(ctypes.byref(p), sz))
|
||||
_cc(_cudart.cudaMemset(p, 0, sz))
|
||||
_cc(_cudart.cudaDeviceSynchronize())
|
||||
h = (ctypes.c_byte * IPC)()
|
||||
_cc(_cudart.cudaIpcGetMemHandle(ctypes.byref(h), p))
|
||||
ah = [None] * ws
|
||||
dist.all_gather_object(ah, bytes(h))
|
||||
ptrs = []
|
||||
for i in range(ws):
|
||||
if i == rank:
|
||||
ptrs.append(p.value)
|
||||
else:
|
||||
hh = (ctypes.c_byte * IPC)(*ah[i])
|
||||
pp = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaIpcOpenMemHandle(ctypes.byref(pp), hh, ctypes.c_uint(1)))
|
||||
ptrs.append(pp.value)
|
||||
return ptrs
|
||||
|
||||
|
||||
def gpu_timer(fn, warmup=20, repeats=200):
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
s = torch.cuda.Event(enable_timing=True)
|
||||
e = torch.cuda.Event(enable_timing=True)
|
||||
s.record()
|
||||
for _ in range(repeats):
|
||||
fn()
|
||||
e.record()
|
||||
torch.cuda.synchronize()
|
||||
return s.elapsed_time(e) / repeats * 1000
|
||||
|
||||
|
||||
def gpu_timer_graph(fn, warmup=20, repeats=200):
|
||||
"""Time with CUDA graph to exclude CPU overhead."""
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
g = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(g):
|
||||
fn()
|
||||
for _ in range(5):
|
||||
g.replay()
|
||||
torch.cuda.synchronize()
|
||||
s = torch.cuda.Event(enable_timing=True)
|
||||
e = torch.cuda.Event(enable_timing=True)
|
||||
s.record()
|
||||
for _ in range(repeats):
|
||||
g.replay()
|
||||
e.record()
|
||||
torch.cuda.synchronize()
|
||||
return s.elapsed_time(e) / repeats * 1000
|
||||
|
||||
|
||||
def main():
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
ws = dist.get_world_size()
|
||||
torch.cuda.set_device(rank)
|
||||
dev = f"cuda:{rank}"
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
if rank == 0:
|
||||
from moe_allgather import _load_lib
|
||||
|
||||
lib = _load_lib()
|
||||
dist.barrier()
|
||||
if rank != 0:
|
||||
from moe_allgather import _load_lib
|
||||
|
||||
lib = _load_lib()
|
||||
dist.barrier()
|
||||
from moe_allgather import MoeAllGather
|
||||
|
||||
max_size = 8 * 1024 * 1024
|
||||
bp = ipc_buf(max_size, rank, ws)
|
||||
dist.barrier()
|
||||
|
||||
class FakeCA:
|
||||
pass
|
||||
|
||||
ca = FakeCA()
|
||||
ca.rank = rank
|
||||
ca.world_size = ws
|
||||
ca.device = torch.device(dev)
|
||||
ca.buffer_ptrs = bp
|
||||
ca.max_size = max_size
|
||||
ag = MoeAllGather(ca)
|
||||
dist.barrier()
|
||||
|
||||
configs = [
|
||||
("1tok", 1),
|
||||
("2tok", 2),
|
||||
("4tok", 4),
|
||||
("8tok", 8),
|
||||
("16tok", 16),
|
||||
("32tok", 32),
|
||||
("64tok", 64),
|
||||
("128tok", 128),
|
||||
("256tok", 256),
|
||||
]
|
||||
topk = 8
|
||||
hd = 3584
|
||||
sd = 448
|
||||
|
||||
if rank == 0:
|
||||
print(f"world_size={ws}, max_per_rank={ag.max_per_rank} bytes")
|
||||
print(f"{'config':<12} {'lamport_graph':>10} {'nccl_graph':>10} {'speedup':>8}")
|
||||
print("-" * 65)
|
||||
|
||||
for name, N in configs:
|
||||
# Check if data fits in buffer.
|
||||
cursor = 0
|
||||
per_tok = topk * 4 + topk * 4 + hd + sd
|
||||
cursor = N * per_tok
|
||||
cursor = (cursor + 15) & ~15
|
||||
if cursor > ag.max_per_rank:
|
||||
if rank == 0:
|
||||
print(f"{name:<12} {'skip (too large)':>40}")
|
||||
continue
|
||||
|
||||
ids = torch.randint(0, 256, (N, topk), dtype=torch.int32, device=dev)
|
||||
wt = torch.randn(N, topk, dtype=torch.float32, device=dev).abs()
|
||||
hs = torch.randint(0, 255, (N, hd), dtype=torch.uint8, device=dev)
|
||||
sc = torch.randint(0, 255, (N, sd), dtype=torch.uint8, device=dev)
|
||||
inputs = [ids, wt, hs, sc]
|
||||
|
||||
# Custom Lamport kernel.
|
||||
c_outs = [
|
||||
torch.empty(N * ws, *t.shape[1:], dtype=t.dtype, device=dev) for t in inputs
|
||||
]
|
||||
|
||||
def run_lamport():
|
||||
lib.moe_all_gather(
|
||||
ag._buf_ptrs_ptr,
|
||||
ag._counters_ptr,
|
||||
rank,
|
||||
ws,
|
||||
ag.seg_capacity,
|
||||
ag.rank_stride,
|
||||
inputs,
|
||||
c_outs,
|
||||
)
|
||||
|
||||
# Lamport with CUDA graph.
|
||||
try:
|
||||
lam_g_us = gpu_timer_graph(run_lamport)
|
||||
except Exception as ex:
|
||||
lam_g_us = float("nan")
|
||||
if rank == 0:
|
||||
print(f" [graph capture failed: {ex}]")
|
||||
|
||||
# NCCL 1×AG (concat into one tensor).
|
||||
cat_inp = torch.cat(
|
||||
[t.reshape(N, -1).contiguous().view(torch.uint8) for t in inputs],
|
||||
dim=1,
|
||||
).contiguous()
|
||||
cat_out = torch.empty(N * ws, cat_inp.shape[1], dtype=torch.uint8, device=dev)
|
||||
|
||||
def run_nccl():
|
||||
dist.all_gather_into_tensor(cat_out, cat_inp)
|
||||
|
||||
# NCCL with CUDA graph.
|
||||
try:
|
||||
nccl_g_us = gpu_timer_graph(run_nccl)
|
||||
except Exception:
|
||||
nccl_g_us = float("nan")
|
||||
|
||||
if rank == 0:
|
||||
speedup = nccl_g_us / lam_g_us if lam_g_us > 0 else float("nan")
|
||||
print(f"{name:<12} {lam_g_us:>9.1f}µ {nccl_g_us:>9.1f}µ {speedup:>7.2f}x")
|
||||
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,301 @@
|
||||
// Lamport-based MoE all-gather kernel for EP dispatch.
|
||||
//
|
||||
// Replaces the flag-barrier approach with a Lamport sentinel protocol
|
||||
// (inspired by FlashInfer's trtllm_allreduce_fusion).
|
||||
//
|
||||
// Key advantages over the flag-barrier approach:
|
||||
// - No explicit barriers (sentinels provide per-element synchronization).
|
||||
// - Push model: NVLink writes (fire-and-forget) instead of NVLink reads.
|
||||
// - Triple buffering: no end barrier needed.
|
||||
//
|
||||
// Gathers the MoE dispatch tensors from all EP ranks:
|
||||
// - topk_ids [N, topk] int32
|
||||
// - topk_weights [N, topk] float32 / bfloat16
|
||||
// - hidden_states [N, D_h] uint8 (NVFP4) / bfloat16
|
||||
// - quant_scales [N, D_s] (optional)
|
||||
//
|
||||
// Double-buffer layout in each rank's IPC buffer:
|
||||
// [Segment 0][Segment 1]
|
||||
// Each segment: [Rank 0 slot][Rank 1 slot]...[Rank N-1 slot]
|
||||
// Each rank slot: packed tensors at 16-byte aligned offsets.
|
||||
//
|
||||
// Sentinel: 0x80000000 (negative-zero in float32). The writer replaces
|
||||
// any data word matching the sentinel with 0 before pushing. The reader
|
||||
// spin-loads (volatile) until no sentinel words remain in the vector.
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <torch/extension.h>
|
||||
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <c10/cuda/CUDAStream.h>
|
||||
|
||||
#define DINLINE __device__ __forceinline__
|
||||
|
||||
constexpr uint32_t SENTINEL = 0x80000000u;
|
||||
constexpr int kMaxBlocks = 36;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Volatile 128-bit load/store and sentinel helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
static DINLINE int4 ld128v(const void* addr) {
|
||||
int4 v;
|
||||
asm volatile("ld.volatile.global.v4.b32 {%0,%1,%2,%3}, [%4];"
|
||||
: "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w)
|
||||
: "l"(addr));
|
||||
return v;
|
||||
}
|
||||
|
||||
static DINLINE bool has_sentinel(int4 v) {
|
||||
return reinterpret_cast<uint32_t&>(v.x) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.y) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.z) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.w) == SENTINEL;
|
||||
}
|
||||
|
||||
static DINLINE int4 remove_sentinel(int4 v) {
|
||||
if (reinterpret_cast<uint32_t&>(v.x) == SENTINEL) v.x = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.y) == SENTINEL) v.y = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.z) == SENTINEL) v.z = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.w) == SENTINEL) v.w = 0;
|
||||
return v;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Lamport all-gather kernel
|
||||
// ---------------------------------------------------------------------------
|
||||
//
|
||||
// Phase 1 — PUSH: each rank writes its packed data to ALL peers' current
|
||||
// segment via regular stores (NVLink push, fire-and-forget).
|
||||
// Phase 2 — CLEAR: each rank writes sentinels to the OLDEST segment of
|
||||
// its own buffer, preparing it for reuse.
|
||||
// Phase 3 — POLL + SCATTER: each rank volatile-loads from its own current
|
||||
// segment, spinning until sentinels disappear, then scatters
|
||||
// directly to per-tensor output arrays.
|
||||
// Phase 4 — ADVANCE: one thread advances the triple-buffer ring counter.
|
||||
|
||||
template <int ngpus, int nbufs>
|
||||
__global__ void __launch_bounds__(512, 1) moe_allgather_lamport_kernel(
|
||||
int64_t* buf_ptrs, // [ngpus] IPC buffer base addresses (device)
|
||||
int* counters, // [0] = unused, [1] = ring (0/1/2), [2] = prev total_sz
|
||||
int rank,
|
||||
int seg_capacity, // bytes per segment
|
||||
int rank_stride, // bytes per rank-slot within a segment
|
||||
int total_sz, // int4 units of actual packed data per rank
|
||||
// inputs (up to 4)
|
||||
const void* inp0, const void* inp1, const void* inp2, const void* inp3,
|
||||
int off0, int sz0, int off1, int sz1, int off2, int sz2, int off3, int sz3,
|
||||
// outputs (up to 4)
|
||||
void* out0, void* out1, void* out2, void* out3) {
|
||||
using V = int4;
|
||||
const int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
const int stride = gridDim.x * blockDim.x;
|
||||
|
||||
// Read segment index and previous clear size.
|
||||
const int seg = counters[1]; // 0 or 1
|
||||
const int prev_total_sz = counters[2]; // set by previous invocation
|
||||
const int cur_seg = seg;
|
||||
const int old_seg = 1 - seg;
|
||||
|
||||
char* bufs[ngpus];
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++)
|
||||
bufs[r] = reinterpret_cast<char*>(buf_ptrs[r]) + cur_seg * seg_capacity;
|
||||
|
||||
// Sentinel vector for clearing.
|
||||
V sent;
|
||||
sent.x = sent.y = sent.z = sent.w = static_cast<int>(SENTINEL);
|
||||
|
||||
// ---- Phase 1: PUSH local data to ALL peers ----
|
||||
// Write to peer_r's buffer at [rank * rank_stride + off_i].
|
||||
|
||||
#define PUSH(idx, inp_ptr, off_val, sz_val) \
|
||||
if constexpr (nbufs > (idx)) { \
|
||||
const V* src = reinterpret_cast<const V*>(inp_ptr); \
|
||||
for (int i = tid; i < (sz_val); i += stride) { \
|
||||
V val = remove_sentinel(src[i]); \
|
||||
_Pragma("unroll") for (int r = 0; r < ngpus; r++) { \
|
||||
reinterpret_cast<V*>(bufs[r] + rank * rank_stride + (off_val))[i] = \
|
||||
val; \
|
||||
} \
|
||||
} \
|
||||
}
|
||||
|
||||
PUSH(0, inp0, off0, sz0)
|
||||
PUSH(1, inp1, off1, sz1)
|
||||
PUSH(2, inp2, off2, sz2)
|
||||
PUSH(3, inp3, off3, sz3)
|
||||
#undef PUSH
|
||||
|
||||
// ---- Phase 2: CLEAR only the previously-written data in oldest segment ----
|
||||
// Only clear what the previous invocation actually wrote (per rank-slot).
|
||||
if (prev_total_sz > 0) {
|
||||
char* clr_base =
|
||||
reinterpret_cast<char*>(buf_ptrs[rank]) + old_seg * seg_capacity;
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++) {
|
||||
V* clr = reinterpret_cast<V*>(clr_base + r * rank_stride);
|
||||
for (int i = tid; i < prev_total_sz; i += stride) clr[i] = sent;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 3: POLL + SCATTER ----
|
||||
// Volatile-load from own buffer; spin until sentinel gone; scatter to output.
|
||||
char* my = bufs[rank];
|
||||
|
||||
#define POLL(idx, out_ptr, off_val, sz_val) \
|
||||
if constexpr (nbufs > (idx)) { \
|
||||
for (int i = tid; i < (sz_val); i += stride) { \
|
||||
_Pragma("unroll") for (int s = 0; s < ngpus; s++) { \
|
||||
V val; \
|
||||
do { \
|
||||
val = ld128v( \
|
||||
reinterpret_cast<V*>(my + s * rank_stride + (off_val)) + i); \
|
||||
} while (has_sentinel(val)); \
|
||||
reinterpret_cast<V*>(out_ptr)[s * (sz_val) + i] = val; \
|
||||
} \
|
||||
} \
|
||||
}
|
||||
|
||||
POLL(0, out0, off0, sz0)
|
||||
POLL(1, out1, off1, sz1)
|
||||
POLL(2, out2, off2, sz2)
|
||||
POLL(3, out3, off3, sz3)
|
||||
#undef POLL
|
||||
|
||||
// ---- Phase 4: ADVANCE ring counter + store clear size for next call ----
|
||||
// Stream serialization ensures the next kernel sees these updates.
|
||||
if (blockIdx.x == 0 && threadIdx.x == 0) {
|
||||
counters[1] = 1 - seg;
|
||||
counters[2] = total_sz;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Sentinel initialization kernel
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
__global__ void lamport_init_kernel(uint32_t* buf, int n) {
|
||||
int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
int stride = gridDim.x * blockDim.x;
|
||||
for (int i = tid; i < n; i += stride) buf[i] = SENTINEL;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Host launcher
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
struct TensorDesc {
|
||||
void* inp;
|
||||
int off;
|
||||
int sz;
|
||||
int64_t nbytes;
|
||||
};
|
||||
|
||||
static TensorDesc make_desc(torch::Tensor& inp, int64_t& cursor) {
|
||||
TORCH_CHECK(inp.is_contiguous(), "input must be contiguous");
|
||||
int64_t nbytes = inp.numel() * inp.element_size();
|
||||
TORCH_CHECK(nbytes % 16 == 0, "tensor byte size must be multiple of 16, got ",
|
||||
nbytes);
|
||||
cursor = (cursor + 15) & ~15;
|
||||
TensorDesc d;
|
||||
d.inp = inp.data_ptr();
|
||||
d.off = static_cast<int>(cursor);
|
||||
d.sz = static_cast<int>(nbytes / 16);
|
||||
d.nbytes = nbytes;
|
||||
cursor += nbytes;
|
||||
return d;
|
||||
}
|
||||
|
||||
void lamport_init(int64_t buf_ptr, int64_t nbytes) {
|
||||
auto stream = c10::cuda::getCurrentCUDAStream().stream();
|
||||
int n = static_cast<int>(nbytes / 4);
|
||||
lamport_init_kernel<<<256, 256, 0, stream>>>(
|
||||
reinterpret_cast<uint32_t*>(buf_ptr), n);
|
||||
}
|
||||
|
||||
void moe_all_gather(int64_t buf_ptrs_ptr, int64_t counters_ptr, int64_t rank,
|
||||
int64_t world_size, int64_t seg_capacity,
|
||||
int64_t rank_stride, std::vector<torch::Tensor>& inputs,
|
||||
std::vector<torch::Tensor>& outputs) {
|
||||
auto stream = c10::cuda::getCurrentCUDAStream().stream();
|
||||
int n = static_cast<int>(inputs.size());
|
||||
TORCH_CHECK(n >= 2 && n <= 4, "2-4 input tensors required");
|
||||
TORCH_CHECK(inputs.size() == outputs.size());
|
||||
|
||||
int64_t cursor = 0;
|
||||
TensorDesc descs[4] = {};
|
||||
for (int i = 0; i < n; i++) descs[i] = make_desc(inputs[i], cursor);
|
||||
TORCH_CHECK(cursor % 16 == 0);
|
||||
int total_sz = static_cast<int>(cursor / 16);
|
||||
TORCH_CHECK(cursor <= rank_stride, "packed data (", cursor,
|
||||
" bytes) exceeds rank_stride (", rank_stride, " bytes)");
|
||||
|
||||
int ws = static_cast<int>(world_size);
|
||||
for (int i = 0; i < n; i++) {
|
||||
TORCH_CHECK(outputs[i].is_contiguous());
|
||||
TORCH_CHECK(outputs[i].numel() == inputs[i].numel() * ws);
|
||||
}
|
||||
|
||||
int r = static_cast<int>(rank);
|
||||
int threads = 512;
|
||||
int blocks =
|
||||
std::max(1, std::min(kMaxBlocks, (total_sz + threads - 1) / threads));
|
||||
|
||||
void *inps[4] = {}, *outs[4] = {};
|
||||
int offs[4] = {}, szs[4] = {};
|
||||
for (int i = 0; i < n; i++) {
|
||||
inps[i] = descs[i].inp;
|
||||
offs[i] = descs[i].off;
|
||||
szs[i] = descs[i].sz;
|
||||
outs[i] = outputs[i].data_ptr();
|
||||
}
|
||||
|
||||
auto* bp = reinterpret_cast<int64_t*>(buf_ptrs_ptr);
|
||||
auto* ct = reinterpret_cast<int*>(counters_ptr);
|
||||
int sc = static_cast<int>(seg_capacity);
|
||||
int rs = static_cast<int>(rank_stride);
|
||||
|
||||
#define KL(ng, nb) \
|
||||
moe_allgather_lamport_kernel<ng, nb><<<blocks, threads, 0, stream>>>( \
|
||||
bp, ct, r, sc, rs, total_sz, inps[0], inps[1], inps[2], inps[3], \
|
||||
offs[0], szs[0], offs[1], szs[1], offs[2], szs[2], offs[3], szs[3], \
|
||||
outs[0], outs[1], outs[2], outs[3]);
|
||||
|
||||
#define GPU_CASE(ng) \
|
||||
case ng: \
|
||||
switch (n) { \
|
||||
case 2: \
|
||||
KL(ng, 2); \
|
||||
break; \
|
||||
case 3: \
|
||||
KL(ng, 3); \
|
||||
break; \
|
||||
case 4: \
|
||||
KL(ng, 4); \
|
||||
break; \
|
||||
} \
|
||||
break;
|
||||
|
||||
switch (ws) {
|
||||
GPU_CASE(2)
|
||||
GPU_CASE(4)
|
||||
GPU_CASE(6)
|
||||
GPU_CASE(8)
|
||||
default:
|
||||
TORCH_CHECK(false, "world_size must be 2, 4, 6, or 8");
|
||||
}
|
||||
#undef GPU_CASE
|
||||
#undef KL
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Python binding
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("moe_all_gather", &moe_all_gather, "Lamport MoE all-gather");
|
||||
m.def("lamport_init", &lamport_init,
|
||||
"Initialize Lamport buffer with sentinels");
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Lamport-based fused MoE all-gather for EP dispatch.
|
||||
|
||||
JIT-compiles the CUDA kernel on first use (cached afterwards).
|
||||
Uses a Lamport sentinel protocol (push writes + per-element sync)
|
||||
with triple-buffered IPC regions — no explicit barriers.
|
||||
|
||||
Usage:
|
||||
ag = MoeAllGather(custom_allreduce)
|
||||
ids_g, wt_g, hs_g, sc_g = ag.gather(topk_ids, topk_weights, hidden, scales)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
_lib = None
|
||||
|
||||
|
||||
def _load_lib():
|
||||
global _lib
|
||||
if _lib is not None:
|
||||
return _lib
|
||||
|
||||
from torch.utils.cpp_extension import load
|
||||
|
||||
src = str(Path(__file__).with_name("moe_allgather.cu"))
|
||||
_lib = load(
|
||||
name="moe_allgather_kernel",
|
||||
sources=[src],
|
||||
extra_cuda_cflags=["-O3", "--use_fast_math"],
|
||||
verbose=os.environ.get("MOE_AG_VERBOSE", "") == "1",
|
||||
)
|
||||
return _lib
|
||||
|
||||
|
||||
class MoeAllGather:
|
||||
"""Lamport-based MoE dispatch all-gather with triple buffering."""
|
||||
|
||||
def __init__(self, ca_comm):
|
||||
self.rank = ca_comm.rank
|
||||
self.world_size = ca_comm.world_size
|
||||
self.device = ca_comm.device
|
||||
self.buffer_ptrs = ca_comm.buffer_ptrs
|
||||
self.max_size = ca_comm.max_size
|
||||
|
||||
ws = self.world_size
|
||||
# Double-buffer layout: 2 segments, each with ws rank-slots.
|
||||
# Safe because kernels in the same stream are serialized, and the
|
||||
# Use the FIRST half of the IPC buffer (second half reserved for
|
||||
# MoeReduceScatter) to avoid overlapping writes.
|
||||
half_size = (self.max_size // 2) & ~15
|
||||
|
||||
# Double-buffer layout within our half: 2 segments, each ws rank-slots.
|
||||
# seg_capacity and rank_stride are 16-byte aligned.
|
||||
self.seg_capacity = (half_size // 2) & ~15
|
||||
self.rank_stride = (self.seg_capacity // ws) & ~15
|
||||
self.max_per_rank = self.rank_stride # max packed bytes per rank
|
||||
|
||||
# Buffer pointer array on device (no offset — first half).
|
||||
self._buf_ptrs = torch.zeros(
|
||||
8, dtype=torch.int64, device=f"cuda:{self.device.index}"
|
||||
)
|
||||
for i in range(ws):
|
||||
self._buf_ptrs[i] = self.buffer_ptrs[i]
|
||||
self._buf_ptrs_ptr = self._buf_ptrs.data_ptr()
|
||||
|
||||
# Counters on device: [0]=unused, [1]=seg (0/1), [2]=prev_total_sz.
|
||||
self._counters = torch.zeros(
|
||||
3, dtype=torch.int32, device=f"cuda:{self.device.index}"
|
||||
)
|
||||
self._counters_ptr = self._counters.data_ptr()
|
||||
|
||||
# Initialize our half with sentinel values.
|
||||
lib = _load_lib()
|
||||
lib.lamport_init(self.buffer_ptrs[self.rank], half_size)
|
||||
torch.accelerator.synchronize(self.device)
|
||||
|
||||
def gather(
|
||||
self,
|
||||
topk_ids: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
scales: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]:
|
||||
lib = _load_lib()
|
||||
ws = self.world_size
|
||||
|
||||
inputs = [topk_ids, topk_weights, hidden_states]
|
||||
if scales is not None:
|
||||
inputs.append(scales)
|
||||
|
||||
outputs = [
|
||||
torch.empty((t.shape[0] * ws, *t.shape[1:]), dtype=t.dtype, device=t.device)
|
||||
for t in inputs
|
||||
]
|
||||
|
||||
lib.moe_all_gather(
|
||||
self._buf_ptrs_ptr,
|
||||
self._counters_ptr,
|
||||
self.rank,
|
||||
self.world_size,
|
||||
self.seg_capacity,
|
||||
self.rank_stride,
|
||||
inputs,
|
||||
outputs,
|
||||
)
|
||||
|
||||
if scales is not None:
|
||||
return outputs[0], outputs[1], outputs[2], outputs[3]
|
||||
return outputs[0], outputs[1], outputs[2], None
|
||||
@@ -0,0 +1,261 @@
|
||||
// Lamport-based MoE reduce-scatter kernel for EP combine.
|
||||
//
|
||||
// JIT-compilable via torch.utils.cpp_extension — no vLLM build required.
|
||||
//
|
||||
// Reduce-scatters a bf16 tensor [N_total, D] across EP ranks. Each rank
|
||||
// contributes its partial MoE output; the kernel sums all contributions
|
||||
// and each rank receives its own slice of the result.
|
||||
//
|
||||
// Protocol (same as the all-gather variant):
|
||||
// 1. PUSH: write own data to all peers' Lamport buffers (NVLink push).
|
||||
// 2. CLEAR: write sentinels to old segment of own buffer.
|
||||
// 3. POLL + REDUCE: volatile-load all peers' data for own slice,
|
||||
// accumulate in fp32, convert back to bf16, store to output.
|
||||
// 4. ADVANCE: toggle double-buffer index.
|
||||
//
|
||||
// Sentinel: 0x80000000 (two bf16 negative-zeros packed in uint32).
|
||||
// For bf16 reduce, replacing -0 with +0 is lossless.
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <torch/extension.h>
|
||||
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <c10/cuda/CUDAStream.h>
|
||||
|
||||
#define DINLINE __device__ __forceinline__
|
||||
|
||||
constexpr uint32_t SENTINEL = 0x80000000u;
|
||||
constexpr int kMaxBlocks = 36;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Volatile 128-bit load and sentinel helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
static DINLINE int4 ld128v(const void* addr) {
|
||||
int4 v;
|
||||
asm volatile("ld.volatile.global.v4.b32 {%0,%1,%2,%3}, [%4];"
|
||||
: "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w)
|
||||
: "l"(addr));
|
||||
return v;
|
||||
}
|
||||
|
||||
static DINLINE bool has_sentinel(int4 v) {
|
||||
return reinterpret_cast<uint32_t&>(v.x) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.y) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.z) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.w) == SENTINEL;
|
||||
}
|
||||
|
||||
static DINLINE int4 remove_sentinel(int4 v) {
|
||||
if (reinterpret_cast<uint32_t&>(v.x) == SENTINEL) v.x = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.y) == SENTINEL) v.y = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.z) == SENTINEL) v.z = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.w) == SENTINEL) v.w = 0;
|
||||
return v;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// bf16 ↔ fp32 helpers for int4 (8 bf16 values = 16 bytes)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// Accumulate 8 bf16 values from an int4 into 8 fp32 accumulators.
|
||||
static DINLINE void accumulate_bf16(float* acc, int4 v) {
|
||||
const __nv_bfloat16* bp = reinterpret_cast<const __nv_bfloat16*>(&v);
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 8; k++) acc[k] += __bfloat162float(bp[k]);
|
||||
}
|
||||
|
||||
// Convert 8 fp32 accumulators to bf16 and pack into int4.
|
||||
static DINLINE int4 fp32_to_bf16_int4(const float* acc) {
|
||||
int4 out;
|
||||
__nv_bfloat16* bp = reinterpret_cast<__nv_bfloat16*>(&out);
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 8; k++) bp[k] = __float2bfloat16(acc[k]);
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Lamport reduce-scatter kernel
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <int ngpus>
|
||||
__global__ void __launch_bounds__(512, 1) moe_rs_lamport_kernel(
|
||||
int64_t* buf_ptrs, // [ngpus] IPC buffer base addresses (device)
|
||||
int* counters, // [0] = unused, [1] = seg (0/1), [2] = prev total_sz
|
||||
int rank,
|
||||
int seg_capacity, // bytes per segment
|
||||
int rank_stride, // bytes per rank-slot within a segment
|
||||
const void* input, // [N_total, D] bf16 — full input
|
||||
void* output, // [N_per_rank, D] bf16 — this rank's reduced slice
|
||||
int total_sz, // int4 units of full input per rank
|
||||
int slice_off, // int4 offset to this rank's slice within packed data
|
||||
int slice_sz) { // int4 units of this rank's slice
|
||||
using V = int4;
|
||||
const int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
const int stride = gridDim.x * blockDim.x;
|
||||
|
||||
// Read segment index and previous clear size.
|
||||
const int seg = counters[1];
|
||||
const int prev_total_sz = counters[2];
|
||||
const int cur_seg = seg;
|
||||
const int old_seg = 1 - seg;
|
||||
|
||||
char* bufs[ngpus];
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++)
|
||||
bufs[r] = reinterpret_cast<char*>(buf_ptrs[r]) + cur_seg * seg_capacity;
|
||||
|
||||
V sent;
|
||||
sent.x = sent.y = sent.z = sent.w = static_cast<int>(SENTINEL);
|
||||
|
||||
// ---- Phase 1: PUSH full input to ALL peers ----
|
||||
{
|
||||
const V* src = reinterpret_cast<const V*>(input);
|
||||
for (int i = tid; i < total_sz; i += stride) {
|
||||
V val = remove_sentinel(src[i]);
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++)
|
||||
reinterpret_cast<V*>(bufs[r] + rank * rank_stride)[i] = val;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 2: CLEAR old segment ----
|
||||
if (prev_total_sz > 0) {
|
||||
char* clr_base =
|
||||
reinterpret_cast<char*>(buf_ptrs[rank]) + old_seg * seg_capacity;
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++) {
|
||||
V* clr = reinterpret_cast<V*>(clr_base + r * rank_stride);
|
||||
for (int i = tid; i < prev_total_sz; i += stride) clr[i] = sent;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 3: POLL + REDUCE for own slice ----
|
||||
// Read all ranks' data at [slice_off, slice_off + slice_sz) from own buffer,
|
||||
// sum in fp32, store bf16 result.
|
||||
{
|
||||
char* my = bufs[rank];
|
||||
V* dst = reinterpret_cast<V*>(output);
|
||||
|
||||
for (int i = tid; i < slice_sz; i += stride) {
|
||||
float acc[8] = {0, 0, 0, 0, 0, 0, 0, 0};
|
||||
|
||||
#pragma unroll
|
||||
for (int s = 0; s < ngpus; s++) {
|
||||
V val;
|
||||
do {
|
||||
val = ld128v(reinterpret_cast<V*>(my + s * rank_stride) + slice_off +
|
||||
i);
|
||||
} while (has_sentinel(val));
|
||||
accumulate_bf16(acc, val);
|
||||
}
|
||||
|
||||
dst[i] = fp32_to_bf16_int4(acc);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 4: ADVANCE ----
|
||||
if (blockIdx.x == 0 && threadIdx.x == 0) {
|
||||
counters[1] = 1 - seg;
|
||||
counters[2] = total_sz;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Sentinel initialization
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
__global__ void lamport_init_kernel(uint32_t* buf, int n) {
|
||||
int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
int stride = gridDim.x * blockDim.x;
|
||||
for (int i = tid; i < n; i += stride) buf[i] = SENTINEL;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Host launcher
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
void lamport_init(int64_t buf_ptr, int64_t nbytes) {
|
||||
auto stream = c10::cuda::getCurrentCUDAStream().stream();
|
||||
int n = static_cast<int>(nbytes / 4);
|
||||
lamport_init_kernel<<<256, 256, 0, stream>>>(
|
||||
reinterpret_cast<uint32_t*>(buf_ptr), n);
|
||||
}
|
||||
|
||||
void moe_reduce_scatter(int64_t buf_ptrs_ptr, int64_t counters_ptr,
|
||||
int64_t rank, int64_t world_size, int64_t seg_capacity,
|
||||
int64_t rank_stride, torch::Tensor input,
|
||||
torch::Tensor output) {
|
||||
auto stream = c10::cuda::getCurrentCUDAStream().stream();
|
||||
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
|
||||
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
|
||||
TORCH_CHECK(input.scalar_type() == torch::kBFloat16,
|
||||
"input must be bf16, got ", input.scalar_type());
|
||||
TORCH_CHECK(output.scalar_type() == torch::kBFloat16, "output must be bf16");
|
||||
|
||||
int ws = static_cast<int>(world_size);
|
||||
int r = static_cast<int>(rank);
|
||||
|
||||
// Input: [N_total, D], Output: [N_per_rank, D]
|
||||
int64_t N_total = input.size(0);
|
||||
int64_t D = input.size(1);
|
||||
TORCH_CHECK(N_total % ws == 0, "N_total must be divisible by world_size");
|
||||
int64_t N_per_rank = N_total / ws;
|
||||
TORCH_CHECK(output.size(0) == N_per_rank);
|
||||
TORCH_CHECK(output.size(1) == D);
|
||||
|
||||
int64_t input_bytes = input.numel() * input.element_size();
|
||||
TORCH_CHECK(input_bytes % 16 == 0,
|
||||
"input byte size must be multiple of 16, got ", input_bytes);
|
||||
TORCH_CHECK(input_bytes <= rank_stride, "input (", input_bytes,
|
||||
" bytes) exceeds rank_stride (", rank_stride, " bytes)");
|
||||
|
||||
int total_sz = static_cast<int>(input_bytes / 16);
|
||||
int slice_sz = total_sz / ws;
|
||||
int slice_off = r * slice_sz;
|
||||
|
||||
int threads = 512;
|
||||
int blocks =
|
||||
std::max(1, std::min(kMaxBlocks, (total_sz + threads - 1) / threads));
|
||||
|
||||
auto* bp = reinterpret_cast<int64_t*>(buf_ptrs_ptr);
|
||||
auto* ct = reinterpret_cast<int*>(counters_ptr);
|
||||
int sc = static_cast<int>(seg_capacity);
|
||||
int rs = static_cast<int>(rank_stride);
|
||||
|
||||
#define LAUNCH(ng) \
|
||||
moe_rs_lamport_kernel<ng><<<blocks, threads, 0, stream>>>( \
|
||||
bp, ct, r, sc, rs, input.data_ptr(), output.data_ptr(), total_sz, \
|
||||
slice_off, slice_sz);
|
||||
|
||||
switch (ws) {
|
||||
case 2:
|
||||
LAUNCH(2);
|
||||
break;
|
||||
case 4:
|
||||
LAUNCH(4);
|
||||
break;
|
||||
case 6:
|
||||
LAUNCH(6);
|
||||
break;
|
||||
case 8:
|
||||
LAUNCH(8);
|
||||
break;
|
||||
default:
|
||||
TORCH_CHECK(false, "world_size must be 2, 4, 6, or 8");
|
||||
}
|
||||
#undef LAUNCH
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Python binding
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("moe_reduce_scatter", &moe_reduce_scatter,
|
||||
"Lamport MoE reduce-scatter");
|
||||
m.def("lamport_init", &lamport_init,
|
||||
"Initialize Lamport buffer with sentinels");
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Lamport-based MoE reduce-scatter for EP combine.
|
||||
|
||||
JIT-compiles the CUDA kernel on first use (cached afterwards).
|
||||
Uses the same Lamport sentinel protocol as the all-gather kernel.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
_lib = None
|
||||
|
||||
|
||||
def _load_lib():
|
||||
global _lib
|
||||
if _lib is not None:
|
||||
return _lib
|
||||
|
||||
from torch.utils.cpp_extension import load
|
||||
|
||||
src = str(Path(__file__).with_name("moe_reduce_scatter.cu"))
|
||||
_lib = load(
|
||||
name="moe_reduce_scatter_kernel",
|
||||
sources=[src],
|
||||
extra_cuda_cflags=["-O3", "--use_fast_math"],
|
||||
verbose=os.environ.get("MOE_RS_VERBOSE", "") == "1",
|
||||
)
|
||||
return _lib
|
||||
|
||||
|
||||
class MoeReduceScatter:
|
||||
"""Lamport-based MoE combine reduce-scatter with double buffering."""
|
||||
|
||||
def __init__(self, ca_comm):
|
||||
self.rank = ca_comm.rank
|
||||
self.world_size = ca_comm.world_size
|
||||
self.device = ca_comm.device
|
||||
self.buffer_ptrs = ca_comm.buffer_ptrs
|
||||
self.max_size = ca_comm.max_size
|
||||
|
||||
ws = self.world_size
|
||||
# Use the SECOND half of the IPC buffer (first half reserved for
|
||||
# MoeAllGather) to avoid overlapping writes.
|
||||
half_size = (self.max_size // 2) & ~15
|
||||
self.buffer_offset = half_size
|
||||
|
||||
# Double-buffer layout within our half: 2 segments, each ws rank-slots.
|
||||
self.seg_capacity = (half_size // 2) & ~15
|
||||
self.rank_stride = (self.seg_capacity // ws) & ~15
|
||||
self.max_per_rank = self.rank_stride
|
||||
|
||||
# Buffer pointer array on device — offset to our half.
|
||||
self._buf_ptrs = torch.zeros(
|
||||
8, dtype=torch.int64, device=f"cuda:{self.device.index}"
|
||||
)
|
||||
for i in range(ws):
|
||||
self._buf_ptrs[i] = self.buffer_ptrs[i] + self.buffer_offset
|
||||
self._buf_ptrs_ptr = self._buf_ptrs.data_ptr()
|
||||
|
||||
# Counters: [0]=unused, [1]=seg (0/1), [2]=prev_total_sz.
|
||||
self._counters = torch.zeros(
|
||||
3, dtype=torch.int32, device=f"cuda:{self.device.index}"
|
||||
)
|
||||
self._counters_ptr = self._counters.data_ptr()
|
||||
|
||||
# Initialize our half with sentinels.
|
||||
lib = _load_lib()
|
||||
lib.lamport_init(self.buffer_ptrs[self.rank] + self.buffer_offset, half_size)
|
||||
torch.accelerator.synchronize(self.device)
|
||||
|
||||
def reduce_scatter(
|
||||
self,
|
||||
input: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Reduce-scatter input [N_total, D] bf16 → output [N_per_rank, D] bf16."""
|
||||
lib = _load_lib()
|
||||
ws = self.world_size
|
||||
|
||||
assert input.dim() == 2
|
||||
N_total, D = input.shape
|
||||
assert N_total % ws == 0
|
||||
N_per_rank = N_total // ws
|
||||
|
||||
output = torch.empty((N_per_rank, D), dtype=input.dtype, device=input.device)
|
||||
|
||||
lib.moe_reduce_scatter(
|
||||
self._buf_ptrs_ptr,
|
||||
self._counters_ptr,
|
||||
self.rank,
|
||||
self.world_size,
|
||||
self.seg_capacity,
|
||||
self.rank_stride,
|
||||
input,
|
||||
output,
|
||||
)
|
||||
return output
|
||||
@@ -0,0 +1,308 @@
|
||||
// Lamport reduce-scatter fused with residual add + RMSNorm.
|
||||
//
|
||||
// Replaces three separate kernels (RS + residual_add + RMSNorm) with one:
|
||||
// 1. PUSH: write MoE output to all peers' Lamport buffers.
|
||||
// 2. CLEAR: write sentinels to old segment.
|
||||
// 3. POLL+REDUCE+FUSE (per-token):
|
||||
// a. Volatile-load from all peers, sum in fp32.
|
||||
// b. Add residual.
|
||||
// c. Compute RMSNorm (block reduction for variance).
|
||||
// d. Store normed output + updated residual.
|
||||
//
|
||||
// Saves: one kernel launch (~3-5µs) + one global memory round-trip per layer.
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <torch/extension.h>
|
||||
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <c10/cuda/CUDAStream.h>
|
||||
|
||||
#define DINLINE __device__ __forceinline__
|
||||
|
||||
constexpr uint32_t SENTINEL = 0x80000000u;
|
||||
constexpr int kMaxBlocks = 36;
|
||||
|
||||
// Each token has D=7168 bf16 values = 896 int4 vectors.
|
||||
// With 512 threads: ceil(896/512) = 2 int4 per thread = 16 fp32 values.
|
||||
constexpr int kMaxValsPerThread = 16;
|
||||
|
||||
static DINLINE int4 ld128v(const void* addr) {
|
||||
int4 v;
|
||||
asm volatile("ld.volatile.global.v4.b32 {%0,%1,%2,%3}, [%4];"
|
||||
: "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w)
|
||||
: "l"(addr));
|
||||
return v;
|
||||
}
|
||||
|
||||
static DINLINE bool has_sentinel(int4 v) {
|
||||
return reinterpret_cast<uint32_t&>(v.x) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.y) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.z) == SENTINEL |
|
||||
reinterpret_cast<uint32_t&>(v.w) == SENTINEL;
|
||||
}
|
||||
|
||||
static DINLINE int4 remove_sentinel(int4 v) {
|
||||
if (reinterpret_cast<uint32_t&>(v.x) == SENTINEL) v.x = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.y) == SENTINEL) v.y = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.z) == SENTINEL) v.z = 0;
|
||||
if (reinterpret_cast<uint32_t&>(v.w) == SENTINEL) v.w = 0;
|
||||
return v;
|
||||
}
|
||||
|
||||
// bf16 helpers
|
||||
static DINLINE void accumulate_bf16(float* acc, int4 v) {
|
||||
const __nv_bfloat16* bp = reinterpret_cast<const __nv_bfloat16*>(&v);
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 8; k++) acc[k] += __bfloat162float(bp[k]);
|
||||
}
|
||||
|
||||
static DINLINE void add_bf16_to_fp32(float* dst, int4 v) {
|
||||
const __nv_bfloat16* bp = reinterpret_cast<const __nv_bfloat16*>(&v);
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 8; k++) dst[k] += __bfloat162float(bp[k]);
|
||||
}
|
||||
|
||||
static DINLINE int4 fp32_to_bf16_int4(const float* vals) {
|
||||
int4 out;
|
||||
__nv_bfloat16* bp = reinterpret_cast<__nv_bfloat16*>(&out);
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 8; k++) bp[k] = __float2bfloat16(vals[k]);
|
||||
return out;
|
||||
}
|
||||
|
||||
// Block-level tree reduction in shared memory.
|
||||
static DINLINE float block_reduce_sum(float val, float* smem) {
|
||||
smem[threadIdx.x] = val;
|
||||
__syncthreads();
|
||||
for (int s = blockDim.x / 2; s > 0; s >>= 1) {
|
||||
if (threadIdx.x < s) smem[threadIdx.x] += smem[threadIdx.x + s];
|
||||
__syncthreads();
|
||||
}
|
||||
return smem[0];
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Fused reduce-scatter + residual + RMSNorm kernel
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <int ngpus>
|
||||
__global__ void __launch_bounds__(512, 1) moe_rs_fused_kernel(
|
||||
int64_t* buf_ptrs, int* counters, int rank, int seg_capacity,
|
||||
int rank_stride,
|
||||
const void* input, // [N_total, D] bf16 — MoE output
|
||||
const void* residual_in, // [N_per_rank, D] bf16 — skip connection
|
||||
const void* gamma, // [D] bf16 — RMSNorm weight
|
||||
void* normed_out, // [N_per_rank, D] bf16 — normed result
|
||||
void* residual_out, // [N_per_rank, D] bf16 — updated residual
|
||||
int total_sz, // int4 units of full input
|
||||
int slice_off, // int4 offset to this rank's slice
|
||||
int slice_sz, // int4 units of this rank's slice
|
||||
int D_int4, // int4 units per token (hidden_dim * 2 / 16)
|
||||
int N_per_rank, // tokens in this rank's slice
|
||||
float eps) { // RMSNorm epsilon
|
||||
using V = int4;
|
||||
const int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
const int stride = gridDim.x * blockDim.x;
|
||||
|
||||
const int seg = counters[1];
|
||||
const int prev_total_sz = counters[2];
|
||||
const int cur_seg = seg;
|
||||
const int old_seg = 1 - seg;
|
||||
|
||||
char* bufs[ngpus];
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++)
|
||||
bufs[r] = reinterpret_cast<char*>(buf_ptrs[r]) + cur_seg * seg_capacity;
|
||||
|
||||
V sent;
|
||||
sent.x = sent.y = sent.z = sent.w = static_cast<int>(SENTINEL);
|
||||
|
||||
// ---- Phase 1: PUSH full MoE output to ALL peers ----
|
||||
{
|
||||
const V* src = reinterpret_cast<const V*>(input);
|
||||
for (int i = tid; i < total_sz; i += stride) {
|
||||
V val = remove_sentinel(src[i]);
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++)
|
||||
reinterpret_cast<V*>(bufs[r] + rank * rank_stride)[i] = val;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 2: CLEAR old segment ----
|
||||
if (prev_total_sz > 0) {
|
||||
char* clr_base =
|
||||
reinterpret_cast<char*>(buf_ptrs[rank]) + old_seg * seg_capacity;
|
||||
#pragma unroll
|
||||
for (int r = 0; r < ngpus; r++) {
|
||||
V* clr = reinterpret_cast<V*>(clr_base + r * rank_stride);
|
||||
for (int i = tid; i < prev_total_sz; i += stride) clr[i] = sent;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 3: Fused POLL + REDUCE + RESIDUAL + RMSNORM ----
|
||||
// Each block handles one token. Only first N_per_rank blocks participate.
|
||||
if (blockIdx.x < N_per_rank) {
|
||||
const int token = blockIdx.x;
|
||||
const int token_off = slice_off + token * D_int4; // in the full buffer
|
||||
char* my = bufs[rank];
|
||||
|
||||
extern __shared__ float smem[];
|
||||
|
||||
// Register storage for intermediate fp32 values.
|
||||
float local_vals[kMaxValsPerThread];
|
||||
int n_vals = 0;
|
||||
float partial_sum_sq = 0.0f;
|
||||
|
||||
// Pass 1: poll + reduce + add residual + compute sum_sq
|
||||
for (int pos = threadIdx.x; pos < D_int4; pos += blockDim.x) {
|
||||
float acc[8] = {0, 0, 0, 0, 0, 0, 0, 0};
|
||||
|
||||
// Poll all ranks' data for this position.
|
||||
#pragma unroll
|
||||
for (int s = 0; s < ngpus; s++) {
|
||||
V val;
|
||||
do {
|
||||
val = ld128v(reinterpret_cast<V*>(my + s * rank_stride) + token_off +
|
||||
pos);
|
||||
} while (has_sentinel(val));
|
||||
accumulate_bf16(acc, val);
|
||||
}
|
||||
|
||||
// Add residual.
|
||||
V res = reinterpret_cast<const V*>(residual_in)[token * D_int4 + pos];
|
||||
add_bf16_to_fp32(acc, res);
|
||||
|
||||
// Store in registers and accumulate sum_sq.
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 8; k++) {
|
||||
local_vals[n_vals++] = acc[k];
|
||||
partial_sum_sq += acc[k] * acc[k];
|
||||
}
|
||||
}
|
||||
|
||||
// Block-level reduction: total sum of squares.
|
||||
float total_sum_sq = block_reduce_sum(partial_sum_sq, smem);
|
||||
float rms_scale = rsqrtf(total_sum_sq / (D_int4 * 8) + eps);
|
||||
|
||||
// Pass 2: apply RMSNorm, store outputs.
|
||||
n_vals = 0;
|
||||
const V* gamma_v = reinterpret_cast<const V*>(gamma);
|
||||
for (int pos = threadIdx.x; pos < D_int4; pos += blockDim.x) {
|
||||
V gv = gamma_v[pos];
|
||||
const __nv_bfloat16* gp = reinterpret_cast<const __nv_bfloat16*>(&gv);
|
||||
|
||||
// Build normed output and residual output.
|
||||
float normed_fp32[8], res_fp32[8];
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 8; k++) {
|
||||
float val = local_vals[n_vals++];
|
||||
res_fp32[k] = val; // residual_out
|
||||
normed_fp32[k] = val * rms_scale * __bfloat162float(gp[k]); // normed
|
||||
}
|
||||
|
||||
reinterpret_cast<V*>(normed_out)[token * D_int4 + pos] =
|
||||
fp32_to_bf16_int4(normed_fp32);
|
||||
reinterpret_cast<V*>(residual_out)[token * D_int4 + pos] =
|
||||
fp32_to_bf16_int4(res_fp32);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 4: ADVANCE ----
|
||||
if (blockIdx.x == 0 && threadIdx.x == 0) {
|
||||
counters[1] = 1 - seg;
|
||||
counters[2] = total_sz;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Sentinel init + host launcher
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
__global__ void lamport_init_kernel(uint32_t* buf, int n) {
|
||||
int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
int stride = gridDim.x * blockDim.x;
|
||||
for (int i = tid; i < n; i += stride) buf[i] = SENTINEL;
|
||||
}
|
||||
|
||||
void lamport_init(int64_t buf_ptr, int64_t nbytes) {
|
||||
auto stream = c10::cuda::getCurrentCUDAStream().stream();
|
||||
int n = static_cast<int>(nbytes / 4);
|
||||
lamport_init_kernel<<<256, 256, 0, stream>>>(
|
||||
reinterpret_cast<uint32_t*>(buf_ptr), n);
|
||||
}
|
||||
|
||||
void moe_rs_fused(int64_t buf_ptrs_ptr, int64_t counters_ptr, int64_t rank,
|
||||
int64_t world_size, int64_t seg_capacity, int64_t rank_stride,
|
||||
torch::Tensor input, // [N_total, D] bf16
|
||||
torch::Tensor residual_in, // [N_per_rank, D] bf16
|
||||
torch::Tensor gamma, // [D] bf16
|
||||
torch::Tensor normed_out, // [N_per_rank, D] bf16
|
||||
torch::Tensor residual_out, // [N_per_rank, D] bf16
|
||||
double eps) {
|
||||
auto stream = c10::cuda::getCurrentCUDAStream().stream();
|
||||
TORCH_CHECK(input.is_contiguous() && residual_in.is_contiguous());
|
||||
TORCH_CHECK(gamma.is_contiguous() && normed_out.is_contiguous());
|
||||
TORCH_CHECK(residual_out.is_contiguous());
|
||||
TORCH_CHECK(input.scalar_type() == torch::kBFloat16);
|
||||
|
||||
int ws = static_cast<int>(world_size);
|
||||
int r = static_cast<int>(rank);
|
||||
int64_t N_total = input.size(0);
|
||||
int64_t D = input.size(1);
|
||||
TORCH_CHECK(N_total % ws == 0);
|
||||
int N_per_rank = static_cast<int>(N_total / ws);
|
||||
|
||||
int64_t input_bytes = input.numel() * input.element_size();
|
||||
TORCH_CHECK(input_bytes % 16 == 0);
|
||||
TORCH_CHECK(input_bytes <= rank_stride);
|
||||
|
||||
int total_sz = static_cast<int>(input_bytes / 16);
|
||||
int D_int4 = static_cast<int>(D * 2 / 16); // bf16 elements → int4 units
|
||||
int slice_sz = total_sz / ws;
|
||||
int slice_off = r * slice_sz;
|
||||
|
||||
int threads = 512;
|
||||
// Need at least N_per_rank blocks for Phase 3 (one per token).
|
||||
int blocks = std::max(
|
||||
N_per_rank, std::min(kMaxBlocks, (total_sz + threads - 1) / threads));
|
||||
|
||||
auto* bp = reinterpret_cast<int64_t*>(buf_ptrs_ptr);
|
||||
auto* ct = reinterpret_cast<int*>(counters_ptr);
|
||||
int sc = static_cast<int>(seg_capacity);
|
||||
int rs = static_cast<int>(rank_stride);
|
||||
int smem = threads * sizeof(float);
|
||||
|
||||
#define LAUNCH(ng) \
|
||||
moe_rs_fused_kernel<ng><<<blocks, threads, smem, stream>>>( \
|
||||
bp, ct, r, sc, rs, input.data_ptr(), residual_in.data_ptr(), \
|
||||
gamma.data_ptr(), normed_out.data_ptr(), residual_out.data_ptr(), \
|
||||
total_sz, slice_off, slice_sz, D_int4, N_per_rank, \
|
||||
static_cast<float>(eps));
|
||||
|
||||
switch (ws) {
|
||||
case 2:
|
||||
LAUNCH(2);
|
||||
break;
|
||||
case 4:
|
||||
LAUNCH(4);
|
||||
break;
|
||||
case 6:
|
||||
LAUNCH(6);
|
||||
break;
|
||||
case 8:
|
||||
LAUNCH(8);
|
||||
break;
|
||||
default:
|
||||
TORCH_CHECK(false, "world_size must be 2, 4, 6, or 8");
|
||||
}
|
||||
#undef LAUNCH
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("moe_rs_fused", &moe_rs_fused,
|
||||
"Fused Lamport reduce-scatter + residual + RMSNorm");
|
||||
m.def("lamport_init", &lamport_init,
|
||||
"Initialize Lamport buffer with sentinels");
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
JIT wrapper for the fused reduce-scatter + residual + RMSNorm kernel.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
_lib = None
|
||||
|
||||
|
||||
def _load_lib():
|
||||
global _lib
|
||||
if _lib is not None:
|
||||
return _lib
|
||||
|
||||
from torch.utils.cpp_extension import load
|
||||
|
||||
src = str(Path(__file__).with_name("moe_rs_fused.cu"))
|
||||
_lib = load(
|
||||
name="moe_rs_fused_kernel",
|
||||
sources=[src],
|
||||
extra_cuda_cflags=["-O3", "--use_fast_math"],
|
||||
verbose=os.environ.get("MOE_RS_FUSED_VERBOSE", "") == "1",
|
||||
)
|
||||
return _lib
|
||||
@@ -0,0 +1,137 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Test for the Lamport-based MoE all-gather kernel."""
|
||||
|
||||
import ctypes
|
||||
import os
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
_cudart = ctypes.CDLL("libcudart.so")
|
||||
IPC = 64
|
||||
|
||||
|
||||
def _cc(r):
|
||||
if r:
|
||||
raise RuntimeError(f"CUDA err {r}")
|
||||
|
||||
|
||||
def ipc_buf(sz, rank, ws):
|
||||
p = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaMalloc(ctypes.byref(p), sz))
|
||||
_cc(_cudart.cudaMemset(p, 0, sz))
|
||||
_cc(_cudart.cudaDeviceSynchronize())
|
||||
h = (ctypes.c_byte * IPC)()
|
||||
_cc(_cudart.cudaIpcGetMemHandle(ctypes.byref(h), p))
|
||||
ah = [None] * ws
|
||||
dist.all_gather_object(ah, bytes(h))
|
||||
ptrs = []
|
||||
for i in range(ws):
|
||||
if i == rank:
|
||||
ptrs.append(p.value)
|
||||
else:
|
||||
hh = (ctypes.c_byte * IPC)(*ah[i])
|
||||
pp = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaIpcOpenMemHandle(ctypes.byref(pp), hh, ctypes.c_uint(1)))
|
||||
ptrs.append(pp.value)
|
||||
return ptrs
|
||||
|
||||
|
||||
def main():
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
ws = dist.get_world_size()
|
||||
torch.cuda.set_device(rank)
|
||||
dev = f"cuda:{rank}"
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
if rank == 0:
|
||||
from moe_allgather import MoeAllGather, _load_lib
|
||||
|
||||
_load_lib()
|
||||
dist.barrier()
|
||||
from moe_allgather import MoeAllGather, _load_lib
|
||||
|
||||
_load_lib()
|
||||
dist.barrier()
|
||||
|
||||
max_size = 8 * 1024 * 1024
|
||||
bp = ipc_buf(max_size, rank, ws)
|
||||
dist.barrier()
|
||||
|
||||
# Build a fake ca_comm-like object.
|
||||
class FakeCA:
|
||||
pass
|
||||
|
||||
ca = FakeCA()
|
||||
ca.rank = rank
|
||||
ca.world_size = ws
|
||||
ca.device = torch.device(dev)
|
||||
ca.buffer_ptrs = bp
|
||||
ca.max_size = max_size
|
||||
# meta_ptrs not needed for Lamport approach
|
||||
ag = MoeAllGather(ca)
|
||||
dist.barrier()
|
||||
|
||||
errors = 0
|
||||
|
||||
# Test with various token counts.
|
||||
for N in [1, 4, 16, 64]:
|
||||
topk = 8
|
||||
hd = 3584
|
||||
sd = 448
|
||||
ids = (
|
||||
torch.arange(N * topk, dtype=torch.int32, device=dev) + rank * 1000
|
||||
).reshape(N, topk)
|
||||
wt = torch.ones(N, topk, dtype=torch.float32, device=dev) * (rank + 1) * 0.1
|
||||
hs = torch.full((N, hd), rank + 1, dtype=torch.uint8, device=dev)
|
||||
sc = torch.full((N, sd), rank + 1, dtype=torch.uint8, device=dev)
|
||||
|
||||
ids_g, wt_g, hs_g, sc_g = ag.gather(ids, wt, hs, sc)
|
||||
|
||||
for src in range(ws):
|
||||
s, e = src * N, (src + 1) * N
|
||||
exp_ids = (
|
||||
torch.arange(N * topk, dtype=torch.int32, device=dev) + src * 1000
|
||||
).reshape(N, topk)
|
||||
if not torch.equal(ids_g[s:e], exp_ids):
|
||||
print(f"[{rank}] FAIL ids src={src} N={N}")
|
||||
errors += 1
|
||||
exp_wt = torch.full(
|
||||
(N, topk), (src + 1) * 0.1, dtype=torch.float32, device=dev
|
||||
)
|
||||
if not torch.allclose(wt_g[s:e], exp_wt):
|
||||
print(f"[{rank}] FAIL wt src={src} N={N}")
|
||||
errors += 1
|
||||
exp_hs = torch.full((N, hd), src + 1, dtype=torch.uint8, device=dev)
|
||||
if not torch.equal(hs_g[s:e], exp_hs):
|
||||
print(f"[{rank}] FAIL hs src={src} N={N}")
|
||||
errors += 1
|
||||
exp_sc = torch.full((N, sd), src + 1, dtype=torch.uint8, device=dev)
|
||||
if not torch.equal(sc_g[s:e], exp_sc):
|
||||
print(f"[{rank}] FAIL sc src={src} N={N}")
|
||||
errors += 1
|
||||
|
||||
# Without scales.
|
||||
ids_g2, wt_g2, hs_g2, _ = ag.gather(ids, wt, hs)
|
||||
for src in range(ws):
|
||||
s, e = src * N, (src + 1) * N
|
||||
exp_ids = (
|
||||
torch.arange(N * topk, dtype=torch.int32, device=dev) + src * 1000
|
||||
).reshape(N, topk)
|
||||
if not torch.equal(ids_g2[s:e], exp_ids):
|
||||
print(f"[{rank}] FAIL no-sc ids src={src} N={N}")
|
||||
errors += 1
|
||||
|
||||
dist.barrier()
|
||||
print(
|
||||
f"[rank {rank}] {'PASSED' if errors == 0 else f'FAILED ({errors})'} (ws={ws})"
|
||||
)
|
||||
dist.destroy_process_group()
|
||||
return errors
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,173 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Stress test: random data, check bitwise correctness against NCCL."""
|
||||
|
||||
import ctypes
|
||||
import os
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
_cudart = ctypes.CDLL("libcudart.so")
|
||||
IPC = 64
|
||||
|
||||
|
||||
def _cc(r):
|
||||
if r:
|
||||
raise RuntimeError(f"CUDA err {r}")
|
||||
|
||||
|
||||
def ipc_buf(sz, rank, ws):
|
||||
p = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaMalloc(ctypes.byref(p), sz))
|
||||
_cc(_cudart.cudaMemset(p, 0, sz))
|
||||
_cc(_cudart.cudaDeviceSynchronize())
|
||||
h = (ctypes.c_byte * IPC)()
|
||||
_cc(_cudart.cudaIpcGetMemHandle(ctypes.byref(h), p))
|
||||
ah = [None] * ws
|
||||
dist.all_gather_object(ah, bytes(h))
|
||||
ptrs = []
|
||||
for i in range(ws):
|
||||
if i == rank:
|
||||
ptrs.append(p.value)
|
||||
else:
|
||||
hh = (ctypes.c_byte * IPC)(*ah[i])
|
||||
pp = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaIpcOpenMemHandle(ctypes.byref(pp), hh, ctypes.c_uint(1)))
|
||||
ptrs.append(pp.value)
|
||||
return ptrs
|
||||
|
||||
|
||||
def main():
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
ws = dist.get_world_size()
|
||||
torch.cuda.set_device(rank)
|
||||
dev = f"cuda:{rank}"
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
if rank == 0:
|
||||
from moe_allgather import _load_lib
|
||||
|
||||
_load_lib()
|
||||
dist.barrier()
|
||||
from moe_allgather import MoeAllGather, _load_lib
|
||||
|
||||
_load_lib()
|
||||
dist.barrier()
|
||||
|
||||
max_size = 8 * 1024 * 1024
|
||||
bp = ipc_buf(max_size, rank, ws)
|
||||
dist.barrier()
|
||||
|
||||
class FakeCA:
|
||||
pass
|
||||
|
||||
ca = FakeCA()
|
||||
ca.rank = rank
|
||||
ca.world_size = ws
|
||||
ca.device = torch.device(dev)
|
||||
ca.buffer_ptrs = bp
|
||||
ca.max_size = max_size
|
||||
ag = MoeAllGather(ca)
|
||||
dist.barrier()
|
||||
|
||||
topk = 8
|
||||
hd = 3584
|
||||
sd = 448
|
||||
errors = 0
|
||||
total_checks = 0
|
||||
sentinel_collisions = 0
|
||||
|
||||
for trial in range(200):
|
||||
# All ranks must use the same N for NCCL reference.
|
||||
N_tensor = torch.randint(1, 65, (1,), device=dev)
|
||||
dist.broadcast(N_tensor, src=0)
|
||||
N = N_tensor.item()
|
||||
# Random data including possible sentinel values
|
||||
ids = torch.randint(0, 256, (N, topk), dtype=torch.int32, device=dev)
|
||||
wt = torch.randn(N, topk, dtype=torch.float32, device=dev)
|
||||
hs = torch.randint(0, 256, (N, hd), dtype=torch.uint8, device=dev)
|
||||
sc = torch.randint(0, 256, (N, sd), dtype=torch.uint8, device=dev)
|
||||
|
||||
# Count sentinel patterns in hidden_states (as uint32 view)
|
||||
hs_u32 = hs.view(torch.int32)
|
||||
sentinel_collisions += (hs_u32 == 0x80000000).sum().item()
|
||||
|
||||
# Custom kernel
|
||||
ids_g, wt_g, hs_g, sc_g = ag.gather(ids, wt, hs, sc)
|
||||
|
||||
# NCCL reference
|
||||
ids_ref = torch.empty(N * ws, topk, dtype=torch.int32, device=dev)
|
||||
wt_ref = torch.empty(N * ws, topk, dtype=torch.float32, device=dev)
|
||||
hs_ref = torch.empty(N * ws, hd, dtype=torch.uint8, device=dev)
|
||||
sc_ref = torch.empty(N * ws, sd, dtype=torch.uint8, device=dev)
|
||||
dist.all_gather_into_tensor(ids_ref, ids)
|
||||
dist.all_gather_into_tensor(wt_ref, wt)
|
||||
dist.all_gather_into_tensor(hs_ref, hs)
|
||||
dist.all_gather_into_tensor(sc_ref, sc)
|
||||
|
||||
# Compare
|
||||
if not torch.equal(ids_g, ids_ref):
|
||||
mismatches = (ids_g != ids_ref).sum().item()
|
||||
if trial < 5 or mismatches > 0:
|
||||
print(f"[{rank}] trial={trial} ids MISMATCH: {mismatches} elements")
|
||||
errors += 1
|
||||
if not torch.equal(wt_g, wt_ref):
|
||||
# Check for -0 vs +0 differences
|
||||
bit_diff = wt_g.view(torch.int32) != wt_ref.view(torch.int32)
|
||||
neg_zero_mask = wt_ref.view(torch.int32) == 0x80000000
|
||||
real_errors = bit_diff & ~neg_zero_mask
|
||||
if real_errors.any():
|
||||
print(
|
||||
f"[{rank}] trial={trial} wt MISMATCH (non-negzero): {real_errors.sum().item()}"
|
||||
)
|
||||
errors += 1
|
||||
if not torch.equal(hs_g, hs_ref):
|
||||
mismatches = (hs_g != hs_ref).sum().item()
|
||||
# Check if mismatches are due to sentinel collision
|
||||
hs_g_u32 = hs_g.view(torch.int32)
|
||||
hs_ref_u32 = hs_ref.view(torch.int32)
|
||||
diff_mask = hs_g_u32 != hs_ref_u32
|
||||
sentinel_mask = (hs_ref_u32 == 0x80000000) & diff_mask
|
||||
non_sentinel = diff_mask & ~sentinel_mask
|
||||
if non_sentinel.any():
|
||||
print(
|
||||
f"[{rank}] trial={trial} hs NON-SENTINEL MISMATCH: {non_sentinel.sum().item()}"
|
||||
)
|
||||
errors += 1
|
||||
elif sentinel_mask.any():
|
||||
if trial < 3:
|
||||
print(
|
||||
f"[{rank}] trial={trial} hs sentinel collision: "
|
||||
f"{sentinel_mask.sum().item()} words (expected rare)"
|
||||
)
|
||||
if not torch.equal(sc_g, sc_ref):
|
||||
mismatches = (sc_g != sc_ref).sum().item()
|
||||
sc_g_u32 = sc_g.view(torch.int32) if sc_g.numel() % 4 == 0 else None
|
||||
if sc_g_u32 is not None:
|
||||
sc_ref_u32 = sc_ref.view(torch.int32)
|
||||
diff_mask = sc_g_u32 != sc_ref_u32
|
||||
sentinel_mask = (sc_ref_u32 == 0x80000000) & diff_mask
|
||||
non_sentinel = diff_mask & ~sentinel_mask
|
||||
if non_sentinel.any():
|
||||
print(
|
||||
f"[{rank}] trial={trial} sc NON-SENTINEL MISMATCH: {non_sentinel.sum().item()}"
|
||||
)
|
||||
errors += 1
|
||||
|
||||
total_checks += 1
|
||||
|
||||
dist.barrier()
|
||||
print(
|
||||
f"[rank {rank}] {total_checks} trials, {errors} real errors, "
|
||||
f"{sentinel_collisions} sentinel patterns in hs data. "
|
||||
f"{'PASSED' if errors == 0 else 'FAILED'}"
|
||||
)
|
||||
dist.destroy_process_group()
|
||||
return errors
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,204 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Test + benchmark for Lamport MoE reduce-scatter."""
|
||||
|
||||
import ctypes
|
||||
import os
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
_cudart = ctypes.CDLL("libcudart.so")
|
||||
IPC = 64
|
||||
|
||||
|
||||
def _cc(r):
|
||||
if r:
|
||||
raise RuntimeError(f"CUDA err {r}")
|
||||
|
||||
|
||||
def ipc_buf(sz, rank, ws):
|
||||
p = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaMalloc(ctypes.byref(p), sz))
|
||||
_cc(_cudart.cudaMemset(p, 0, sz))
|
||||
_cc(_cudart.cudaDeviceSynchronize())
|
||||
h = (ctypes.c_byte * IPC)()
|
||||
_cc(_cudart.cudaIpcGetMemHandle(ctypes.byref(h), p))
|
||||
ah = [None] * ws
|
||||
dist.all_gather_object(ah, bytes(h))
|
||||
ptrs = []
|
||||
for i in range(ws):
|
||||
if i == rank:
|
||||
ptrs.append(p.value)
|
||||
else:
|
||||
hh = (ctypes.c_byte * IPC)(*ah[i])
|
||||
pp = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaIpcOpenMemHandle(ctypes.byref(pp), hh, ctypes.c_uint(1)))
|
||||
ptrs.append(pp.value)
|
||||
return ptrs
|
||||
|
||||
|
||||
def gpu_timer_graph(fn, warmup=20, repeats=200):
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
g = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(g):
|
||||
fn()
|
||||
for _ in range(5):
|
||||
g.replay()
|
||||
torch.cuda.synchronize()
|
||||
s = torch.cuda.Event(enable_timing=True)
|
||||
e = torch.cuda.Event(enable_timing=True)
|
||||
s.record()
|
||||
for _ in range(repeats):
|
||||
g.replay()
|
||||
e.record()
|
||||
torch.cuda.synchronize()
|
||||
return s.elapsed_time(e) / repeats * 1000
|
||||
|
||||
|
||||
def main():
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
ws = dist.get_world_size()
|
||||
torch.cuda.set_device(rank)
|
||||
dev = f"cuda:{rank}"
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
if rank == 0:
|
||||
from moe_reduce_scatter import _load_lib
|
||||
|
||||
_load_lib()
|
||||
dist.barrier()
|
||||
if rank != 0:
|
||||
from moe_reduce_scatter import _load_lib
|
||||
|
||||
_load_lib()
|
||||
dist.barrier()
|
||||
from moe_reduce_scatter import MoeReduceScatter
|
||||
|
||||
max_size = 8 * 1024 * 1024
|
||||
bp = ipc_buf(max_size, rank, ws)
|
||||
dist.barrier()
|
||||
|
||||
class FakeCA:
|
||||
pass
|
||||
|
||||
ca = FakeCA()
|
||||
ca.rank = rank
|
||||
ca.world_size = ws
|
||||
ca.device = torch.device(dev)
|
||||
ca.buffer_ptrs = bp
|
||||
ca.max_size = max_size
|
||||
rs = MoeReduceScatter(ca)
|
||||
dist.barrier()
|
||||
|
||||
D = 7168 # DeepSeek V3 hidden_dim
|
||||
errors = 0
|
||||
|
||||
# ---- Correctness tests ----
|
||||
for N_per_rank in [1, 4, 16]:
|
||||
N_total = N_per_rank * ws
|
||||
|
||||
# Each rank gets a deterministic input.
|
||||
torch.manual_seed(42)
|
||||
# All ranks create the SAME "ground truth" inputs for each rank.
|
||||
all_inputs = [
|
||||
torch.randn(N_total, D, dtype=torch.bfloat16, device=dev) for _ in range(ws)
|
||||
]
|
||||
# This rank's input is all_inputs[rank].
|
||||
my_input = all_inputs[rank]
|
||||
|
||||
# Custom reduce-scatter.
|
||||
custom_out = rs.reduce_scatter(my_input)
|
||||
|
||||
# NCCL reference: reduce_scatter_tensor.
|
||||
nccl_out = torch.empty(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
dist.reduce_scatter_tensor(nccl_out, my_input)
|
||||
|
||||
# bf16 summation order differs between our kernel and NCCL,
|
||||
# giving ~1-2 ULP differences. Use generous tolerance.
|
||||
max_diff = (custom_out.float() - nccl_out.float()).abs().max().item()
|
||||
if not torch.allclose(custom_out, nccl_out, atol=0.125, rtol=0.01):
|
||||
mismatches = (
|
||||
((custom_out.float() - nccl_out.float()).abs() > 0.125).sum().item()
|
||||
)
|
||||
print(
|
||||
f"[{rank}] N_per_rank={N_per_rank} MISMATCH: "
|
||||
f"max_diff={max_diff:.6f}, mismatches={mismatches}"
|
||||
)
|
||||
errors += 1
|
||||
else:
|
||||
if rank == 0:
|
||||
print(f" N_per_rank={N_per_rank}: PASS (max_diff={max_diff:.6f})")
|
||||
|
||||
# ---- Benchmark ----
|
||||
if rank == 0:
|
||||
print(f"\nworld_size={ws}, max_per_rank={rs.max_per_rank} bytes")
|
||||
print(
|
||||
f"{'config':<12} {'lamport':>10} {'lamp_graph':>10} "
|
||||
f"{'nccl':>10} {'nccl_graph':>10} {'speedup':>8}"
|
||||
)
|
||||
print("-" * 65)
|
||||
|
||||
configs = [
|
||||
("1tok", 1),
|
||||
("2tok", 2),
|
||||
("4tok", 4),
|
||||
("8tok", 8),
|
||||
("16tok", 16),
|
||||
("32tok", 32),
|
||||
("64tok", 64),
|
||||
("128tok", 128),
|
||||
("256tok", 256),
|
||||
]
|
||||
|
||||
for name, N_per_rank in configs:
|
||||
N_total = N_per_rank * ws
|
||||
input_bytes = N_total * D * 2 # bf16
|
||||
if input_bytes > rs.max_per_rank:
|
||||
if rank == 0:
|
||||
print(f"{name:<12} {'skip (too large)':>40}")
|
||||
continue
|
||||
|
||||
inp = torch.randn(N_total, D, dtype=torch.bfloat16, device=dev)
|
||||
c_out = torch.empty(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
n_out = torch.empty(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
|
||||
def run_lamport():
|
||||
rs.reduce_scatter(inp)
|
||||
|
||||
def run_nccl():
|
||||
dist.reduce_scatter_tensor(n_out, inp)
|
||||
|
||||
from bench_moe_allgather import gpu_timer
|
||||
|
||||
lam_us = gpu_timer(run_lamport)
|
||||
try:
|
||||
lam_g_us = gpu_timer_graph(run_lamport)
|
||||
except Exception:
|
||||
lam_g_us = float("nan")
|
||||
|
||||
nccl_us = gpu_timer(run_nccl)
|
||||
try:
|
||||
nccl_g_us = gpu_timer_graph(run_nccl)
|
||||
except Exception:
|
||||
nccl_g_us = float("nan")
|
||||
|
||||
if rank == 0:
|
||||
speedup = nccl_g_us / lam_g_us if lam_g_us > 0 else float("nan")
|
||||
print(
|
||||
f"{name:<12} {lam_us:>9.1f}µ {lam_g_us:>9.1f}µ "
|
||||
f"{nccl_us:>9.1f}µ {nccl_g_us:>9.1f}µ {speedup:>7.2f}x"
|
||||
)
|
||||
|
||||
dist.barrier()
|
||||
print(f"[rank {rank}] {'PASSED' if errors == 0 else f'FAILED ({errors})'}")
|
||||
dist.destroy_process_group()
|
||||
return errors
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,270 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Test + benchmark: fused RS + residual + RMSNorm vs separate kernels."""
|
||||
|
||||
import ctypes
|
||||
import os
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
_cudart = ctypes.CDLL("libcudart.so")
|
||||
IPC = 64
|
||||
|
||||
|
||||
def _cc(r):
|
||||
if r:
|
||||
raise RuntimeError(f"CUDA err {r}")
|
||||
|
||||
|
||||
def ipc_buf(sz, rank, ws):
|
||||
p = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaMalloc(ctypes.byref(p), sz))
|
||||
_cc(_cudart.cudaMemset(p, 0, sz))
|
||||
_cc(_cudart.cudaDeviceSynchronize())
|
||||
h = (ctypes.c_byte * IPC)()
|
||||
_cc(_cudart.cudaIpcGetMemHandle(ctypes.byref(h), p))
|
||||
ah = [None] * ws
|
||||
dist.all_gather_object(ah, bytes(h))
|
||||
ptrs = []
|
||||
for i in range(ws):
|
||||
if i == rank:
|
||||
ptrs.append(p.value)
|
||||
else:
|
||||
hh = (ctypes.c_byte * IPC)(*ah[i])
|
||||
pp = ctypes.c_void_p()
|
||||
_cc(_cudart.cudaIpcOpenMemHandle(ctypes.byref(pp), hh, ctypes.c_uint(1)))
|
||||
ptrs.append(pp.value)
|
||||
return ptrs
|
||||
|
||||
|
||||
def rms_norm_ref(x, gamma, eps):
|
||||
"""Reference RMSNorm in fp32."""
|
||||
xf = x.float()
|
||||
rms = torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + eps)
|
||||
return (xf * rms * gamma.float()).to(x.dtype)
|
||||
|
||||
|
||||
def gpu_timer(fn, warmup=20, repeats=200):
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
s = torch.cuda.Event(enable_timing=True)
|
||||
e = torch.cuda.Event(enable_timing=True)
|
||||
s.record()
|
||||
for _ in range(repeats):
|
||||
fn()
|
||||
e.record()
|
||||
torch.cuda.synchronize()
|
||||
return s.elapsed_time(e) / repeats * 1000
|
||||
|
||||
|
||||
def gpu_timer_graph(fn, warmup=20, repeats=200):
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
g = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(g):
|
||||
fn()
|
||||
for _ in range(5):
|
||||
g.replay()
|
||||
torch.cuda.synchronize()
|
||||
s = torch.cuda.Event(enable_timing=True)
|
||||
e = torch.cuda.Event(enable_timing=True)
|
||||
s.record()
|
||||
for _ in range(repeats):
|
||||
g.replay()
|
||||
e.record()
|
||||
torch.cuda.synchronize()
|
||||
return s.elapsed_time(e) / repeats * 1000
|
||||
|
||||
|
||||
def main():
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
ws = dist.get_world_size()
|
||||
torch.cuda.set_device(rank)
|
||||
dev = f"cuda:{rank}"
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
# Compile fused kernel
|
||||
if rank == 0:
|
||||
from torch.utils.cpp_extension import load
|
||||
|
||||
load(
|
||||
name="moe_rs_fused_kernel",
|
||||
sources=[os.path.join(os.path.dirname(__file__), "moe_rs_fused.cu")],
|
||||
extra_cuda_cflags=["-O3", "--use_fast_math"],
|
||||
verbose=False,
|
||||
)
|
||||
dist.barrier()
|
||||
from torch.utils.cpp_extension import load
|
||||
|
||||
fused_lib = load(
|
||||
name="moe_rs_fused_kernel",
|
||||
sources=[os.path.join(os.path.dirname(__file__), "moe_rs_fused.cu")],
|
||||
extra_cuda_cflags=["-O3", "--use_fast_math"],
|
||||
verbose=False,
|
||||
)
|
||||
|
||||
# Also compile separate RS kernel for comparison
|
||||
from moe_reduce_scatter import MoeReduceScatter
|
||||
|
||||
max_size = 8 * 1024 * 1024
|
||||
bp = ipc_buf(max_size, rank, ws)
|
||||
dist.barrier()
|
||||
|
||||
class FakeCA:
|
||||
pass
|
||||
|
||||
ca = FakeCA()
|
||||
ca.rank = rank
|
||||
ca.world_size = ws
|
||||
ca.device = torch.device(dev)
|
||||
ca.buffer_ptrs = bp
|
||||
ca.max_size = max_size
|
||||
|
||||
rs_separate = MoeReduceScatter(ca)
|
||||
|
||||
# Fused kernel setup (uses same buffer layout as MoeReduceScatter)
|
||||
half_size = (max_size // 2) & ~15
|
||||
buf_offset = half_size
|
||||
fused_seg_cap = (half_size // 2) & ~15
|
||||
fused_rank_stride = (fused_seg_cap // ws) & ~15
|
||||
|
||||
fused_buf_ptrs = torch.zeros(8, dtype=torch.int64, device=dev)
|
||||
for i in range(ws):
|
||||
fused_buf_ptrs[i] = bp[i] + buf_offset
|
||||
fused_counters = torch.zeros(3, dtype=torch.int32, device=dev)
|
||||
|
||||
# Init sentinels for fused kernel's buffer region
|
||||
fused_lib.lamport_init(bp[rank] + buf_offset, half_size)
|
||||
torch.cuda.synchronize()
|
||||
dist.barrier()
|
||||
|
||||
D = 7168
|
||||
eps = 1e-6
|
||||
gamma = torch.randn(D, dtype=torch.bfloat16, device=dev).abs() + 0.5
|
||||
errors = 0
|
||||
|
||||
# ---- Correctness ----
|
||||
for N_per_rank in [1, 4]:
|
||||
N_total = N_per_rank * ws
|
||||
torch.manual_seed(42 + rank)
|
||||
moe_out = torch.randn(N_total, D, dtype=torch.bfloat16, device=dev)
|
||||
residual = torch.randn(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
|
||||
# Reference: NCCL RS + add + norm
|
||||
rs_ref = torch.empty(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
dist.reduce_scatter_tensor(rs_ref, moe_out)
|
||||
ref_residual = residual + rs_ref
|
||||
ref_normed = rms_norm_ref(ref_residual, gamma, eps)
|
||||
|
||||
# Fused kernel
|
||||
normed_out = torch.empty(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
residual_out = torch.empty(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
fused_lib.moe_rs_fused(
|
||||
fused_buf_ptrs.data_ptr(),
|
||||
fused_counters.data_ptr(),
|
||||
rank,
|
||||
ws,
|
||||
fused_seg_cap,
|
||||
fused_rank_stride,
|
||||
moe_out,
|
||||
residual,
|
||||
gamma,
|
||||
normed_out,
|
||||
residual_out,
|
||||
eps,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Compare
|
||||
max_diff_res = (residual_out.float() - ref_residual.float()).abs().max().item()
|
||||
max_diff_norm = (normed_out.float() - ref_normed.float()).abs().max().item()
|
||||
|
||||
ok = max_diff_res < 0.125 and max_diff_norm < 0.125
|
||||
if rank == 0:
|
||||
print(
|
||||
f" N_per_rank={N_per_rank}: {'PASS' if ok else 'FAIL'} "
|
||||
f"(res_diff={max_diff_res:.4f}, norm_diff={max_diff_norm:.4f})"
|
||||
)
|
||||
if not ok:
|
||||
errors += 1
|
||||
|
||||
# ---- Benchmark ----
|
||||
if rank == 0:
|
||||
print(f"\nBenchmark: D={D}, world_size={ws}")
|
||||
print(
|
||||
f"{'config':<10} {'fused':>10} {'fused_g':>10} "
|
||||
f"{'RS+norm':>10} {'RS+norm_g':>10} {'speedup':>8}"
|
||||
)
|
||||
print("-" * 58)
|
||||
|
||||
for N_per_rank in [1, 2, 4, 8]:
|
||||
N_total = N_per_rank * ws
|
||||
input_bytes = N_total * D * 2
|
||||
if input_bytes > fused_rank_stride:
|
||||
if rank == 0:
|
||||
print(f"{N_per_rank}tok skip (too large)")
|
||||
continue
|
||||
|
||||
moe_out = torch.randn(N_total, D, dtype=torch.bfloat16, device=dev)
|
||||
residual = torch.randn(N_per_rank, D, dtype=torch.bfloat16, device=dev)
|
||||
normed_out = torch.empty_like(residual)
|
||||
residual_out = torch.empty_like(residual)
|
||||
rs_out = torch.empty_like(residual)
|
||||
|
||||
# Fused
|
||||
def run_fused():
|
||||
fused_lib.moe_rs_fused(
|
||||
fused_buf_ptrs.data_ptr(),
|
||||
fused_counters.data_ptr(),
|
||||
rank,
|
||||
ws,
|
||||
fused_seg_cap,
|
||||
fused_rank_stride,
|
||||
moe_out,
|
||||
residual,
|
||||
gamma,
|
||||
normed_out,
|
||||
residual_out,
|
||||
eps,
|
||||
)
|
||||
|
||||
# Separate: RS + add + norm (our Lamport RS + triton-like ops)
|
||||
def run_separate():
|
||||
rs_separate.reduce_scatter(moe_out)
|
||||
# Simulate add + RMSNorm (in practice this is a fused triton kernel)
|
||||
tmp = residual + rs_out
|
||||
torch.rsqrt(tmp.float().pow(2).mean(-1, keepdim=True) + eps)
|
||||
|
||||
fused_us = gpu_timer(run_fused)
|
||||
try:
|
||||
fused_g = gpu_timer_graph(run_fused)
|
||||
except Exception:
|
||||
fused_g = float("nan")
|
||||
|
||||
sep_us = gpu_timer(run_separate)
|
||||
try:
|
||||
sep_g = gpu_timer_graph(run_separate)
|
||||
except Exception:
|
||||
sep_g = float("nan")
|
||||
|
||||
if rank == 0:
|
||||
speedup = sep_g / fused_g if fused_g > 0 else float("nan")
|
||||
print(
|
||||
f"{N_per_rank}tok {fused_us:>9.1f}µ {fused_g:>9.1f}µ "
|
||||
f"{sep_us:>9.1f}µ {sep_g:>9.1f}µ {speedup:>7.2f}x"
|
||||
)
|
||||
|
||||
dist.barrier()
|
||||
print(f"[rank {rank}] {'PASSED' if errors == 0 else f'FAILED ({errors})'}")
|
||||
dist.destroy_process_group()
|
||||
return errors
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
+135
-352
@@ -1,373 +1,156 @@
|
||||
// Portions of this file are adapted from SGLang PR:
|
||||
// https://github.com/sgl-project/sglang/pull/11194
|
||||
// and
|
||||
// https://github.com/sgl-project/sglang/pull/17747
|
||||
// Persistent TopK kernel for DeepSeek V3 sparse attention indexer.
|
||||
// See persistent_topk.cuh for kernel implementation.
|
||||
|
||||
#include "cuda_compat.h"
|
||||
#include "dispatch_utils.h"
|
||||
|
||||
#include <torch/cuda.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <torch/all.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <algorithm>
|
||||
|
||||
#ifndef USE_ROCM
|
||||
#include <cub/cub.cuh>
|
||||
#else
|
||||
#include <hipcub/hipcub.hpp>
|
||||
#include "persistent_topk.cuh"
|
||||
#endif
|
||||
|
||||
namespace vllm {
|
||||
void persistent_topk(const torch::Tensor& logits, const torch::Tensor& lengths,
|
||||
torch::Tensor& output, torch::Tensor& workspace, int64_t k,
|
||||
int64_t max_seq_len) {
|
||||
#ifndef USE_ROCM
|
||||
TORCH_CHECK(logits.is_cuda(), "logits must be CUDA tensor");
|
||||
TORCH_CHECK(lengths.is_cuda(), "lengths must be CUDA tensor");
|
||||
TORCH_CHECK(output.is_cuda(), "output must be CUDA tensor");
|
||||
TORCH_CHECK(logits.dtype() == torch::kFloat32, "Only float32 supported");
|
||||
TORCH_CHECK(lengths.dtype() == torch::kInt32, "lengths must be int32");
|
||||
TORCH_CHECK(output.dtype() == torch::kInt32, "output must be int32");
|
||||
TORCH_CHECK(logits.dim() == 2, "logits must be 2D");
|
||||
TORCH_CHECK(lengths.dim() == 1 || lengths.dim() == 2,
|
||||
"lengths must be 1D or 2D");
|
||||
TORCH_CHECK(lengths.is_contiguous(), "lengths must be contiguous");
|
||||
TORCH_CHECK(output.dim() == 2, "output must be 2D");
|
||||
|
||||
constexpr int TopK = 2048; // DeepSeek V3 sparse attention top-k
|
||||
constexpr int kThreadsPerBlock = 1024; // Threads per block
|
||||
const int64_t num_rows = logits.size(0);
|
||||
const int64_t stride = logits.size(1);
|
||||
|
||||
// Shared memory budget
|
||||
#if defined(USE_ROCM)
|
||||
constexpr size_t kSmem = 48 * 1024; // ROCm default: 48KB
|
||||
#else
|
||||
// Reduced from 128KB to 32KB to improve occupancy.
|
||||
// Each radix pass needs at most ~TopK candidates in the threshold bin,
|
||||
// so 4K entries per round (2 rounds = 8K entries = 32KB) is sufficient.
|
||||
constexpr size_t kSmem = 8 * 1024 * sizeof(uint32_t); // 32KB (bytes)
|
||||
#endif
|
||||
TORCH_CHECK(lengths.numel() == num_rows, "lengths size mismatch");
|
||||
TORCH_CHECK(output.size(0) == num_rows && output.size(1) == k,
|
||||
"output size mismatch");
|
||||
namespace P = vllm::persistent;
|
||||
|
||||
struct FastTopKParams {
|
||||
const float* __restrict__ input; // [batch, seq_len] Logits
|
||||
const int32_t* __restrict__ row_starts; // [batch] Offset into each row
|
||||
// (optional)
|
||||
int32_t* __restrict__ indices; // [batch, TopK] Output top-k indices
|
||||
int32_t* __restrict__ lengths; // [batch] Sequence lengths per row
|
||||
int64_t input_stride; // Stride between rows
|
||||
};
|
||||
TORCH_CHECK(k == P::TopK, "k must be 2048");
|
||||
TORCH_CHECK(k <= stride, "k out of range");
|
||||
|
||||
__device__ __forceinline__ auto convert_to_uint32_v2(float x) -> uint32_t {
|
||||
uint32_t bits = __float_as_uint(x);
|
||||
return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u);
|
||||
}
|
||||
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
__device__ __forceinline__ auto convert_to_uint8(float x) -> uint8_t {
|
||||
__half h = __float2half_rn(x);
|
||||
uint16_t bits = __half_as_ushort(h);
|
||||
uint16_t key = (bits & 0x8000) ? static_cast<uint16_t>(~bits)
|
||||
: static_cast<uint16_t>(bits | 0x8000);
|
||||
return static_cast<uint8_t>(key >> 8);
|
||||
}
|
||||
|
||||
__device__ void naive_topk_cuda(const float* __restrict__ logits,
|
||||
int32_t* __restrict__ output_indices,
|
||||
int32_t seq_len) {
|
||||
const int thread_id = threadIdx.x;
|
||||
for (int i = thread_id; i < TopK; i += kThreadsPerBlock) {
|
||||
output_indices[i] = (i < seq_len) ? i : -1;
|
||||
}
|
||||
}
|
||||
|
||||
// Adapted from:
|
||||
// https://github.com/sgl-project/sglang/blob/v0.5.8/sgl-kernel/csrc/elementwise/topk.cu#L87
|
||||
// by: DarkSharpness
|
||||
// which at the same time is an optimized topk kernel copied from tilelang
|
||||
// kernel
|
||||
__device__ void fast_topk_cuda_tl(
|
||||
const float* __restrict__ logits, // Input logits [seq_len]
|
||||
int* __restrict__ output_indices, // Output top-k indices [TopK]
|
||||
int logits_offset, // Starting offset in logits array
|
||||
int seq_len) // Number of valid logits to process
|
||||
{
|
||||
constexpr int RADIX = 256;
|
||||
constexpr int MAX_BUFFERED_ITEMS = kSmem / (2 * sizeof(int));
|
||||
|
||||
alignas(128) __shared__ int shared_histogram[2][RADIX + 128];
|
||||
alignas(128) __shared__ int shared_output_count;
|
||||
alignas(128) __shared__ int shared_threshold_bin;
|
||||
alignas(128) __shared__ int shared_buffered_count[2];
|
||||
|
||||
extern __shared__ int buffered_indices[][MAX_BUFFERED_ITEMS];
|
||||
|
||||
const int thread_id = threadIdx.x;
|
||||
int remaining_k = TopK;
|
||||
|
||||
// Pass 0: Build coarse 8-bit histogram using FP16 high bits
|
||||
if (thread_id < RADIX + 1) {
|
||||
shared_histogram[0][thread_id] = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int idx = thread_id; idx < seq_len; idx += kThreadsPerBlock) {
|
||||
const auto bin = convert_to_uint8(logits[idx + logits_offset]);
|
||||
::atomicAdd(&shared_histogram[0][bin], 1);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Helper: Compute cumulative sum (suffix sum) over histogram using ping-pong
|
||||
// buffers
|
||||
auto compute_cumulative_sum = [&]() {
|
||||
static_assert(1 << 8 == RADIX,
|
||||
"Radix must be 256 for 8 unrolled iterations");
|
||||
#pragma unroll 8
|
||||
for (int i = 0; i < 8; ++i) {
|
||||
if (C10_LIKELY(thread_id < RADIX)) {
|
||||
const int stride = 1 << i;
|
||||
const int src_buffer = i & 1;
|
||||
const int dst_buffer = src_buffer ^ 1;
|
||||
|
||||
int value = shared_histogram[src_buffer][thread_id];
|
||||
if (thread_id < RADIX - stride) {
|
||||
value += shared_histogram[src_buffer][thread_id + stride];
|
||||
}
|
||||
shared_histogram[dst_buffer][thread_id] = value;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
};
|
||||
|
||||
compute_cumulative_sum();
|
||||
|
||||
// Find threshold bin where cumsum crosses remaining_k
|
||||
if (thread_id < RADIX && shared_histogram[0][thread_id] > remaining_k &&
|
||||
shared_histogram[0][thread_id + 1] <= remaining_k) {
|
||||
shared_threshold_bin = thread_id;
|
||||
shared_buffered_count[0] = 0;
|
||||
shared_output_count = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
const int threshold_bin = shared_threshold_bin;
|
||||
remaining_k -= shared_histogram[0][threshold_bin + 1];
|
||||
|
||||
// Early exit if threshold bin perfectly matches remaining_k
|
||||
if (remaining_k == 0) {
|
||||
for (int idx = thread_id; idx < seq_len; idx += kThreadsPerBlock) {
|
||||
const int bin = convert_to_uint8(logits[idx + logits_offset]);
|
||||
if (bin > threshold_bin) {
|
||||
const int output_pos = ::atomicAdd(&shared_output_count, 1);
|
||||
output_indices[output_pos] = idx;
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
return;
|
||||
static int num_sms = 0;
|
||||
static int max_smem_per_block = 0;
|
||||
if (num_sms == 0) {
|
||||
int device;
|
||||
cudaGetDevice(&device);
|
||||
cudaDeviceGetAttribute(&num_sms, cudaDevAttrMultiProcessorCount, device);
|
||||
cudaDeviceGetAttribute(&max_smem_per_block,
|
||||
cudaDevAttrMaxSharedMemoryPerBlockOptin, device);
|
||||
}
|
||||
|
||||
// Prepare for refinement passes: Process threshold bin
|
||||
__syncthreads();
|
||||
if (thread_id < RADIX + 1) {
|
||||
shared_histogram[0][thread_id] = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Scan all elements and:
|
||||
// 1. Write indices > threshold_bin to output
|
||||
// 2. Buffer indices == threshold_bin for refinement
|
||||
// 3. Build histogram for next refinement pass (fused optimization)
|
||||
for (int idx = thread_id; idx < seq_len; idx += kThreadsPerBlock) {
|
||||
const float logit_value = logits[idx + logits_offset];
|
||||
const int bin = convert_to_uint8(logit_value);
|
||||
|
||||
if (bin > threshold_bin) {
|
||||
// in top-k, write to output
|
||||
const int output_pos = ::atomicAdd(&shared_output_count, 1);
|
||||
output_indices[output_pos] = idx;
|
||||
} else if (bin == threshold_bin) {
|
||||
// Candidate for top-k, needs refinement
|
||||
const int buffer_pos = ::atomicAdd(&shared_buffered_count[0], 1);
|
||||
if (C10_LIKELY(buffer_pos < MAX_BUFFERED_ITEMS)) {
|
||||
buffered_indices[0][buffer_pos] = idx;
|
||||
// Fused: Build histogram for next pass
|
||||
const uint32_t fp32_bits = convert_to_uint32_v2(logit_value);
|
||||
const int next_bin = (fp32_bits >> 24) & 0xFF;
|
||||
::atomicAdd(&shared_histogram[0][next_bin], 1);
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// ============================================================================
|
||||
// Passes 1-4: Refine using 8-bit passes over FP32 bits
|
||||
// ============================================================================
|
||||
// FP32 bits [31:0] split into 4 bytes processed MSB-first:
|
||||
// Pass 1: bits [31:24], Pass 2: bits [23:16], Pass 3: bits [15:8], Pass 4:
|
||||
// bits [7:0]
|
||||
#pragma unroll 4
|
||||
for (int pass = 0; pass < 4; ++pass) {
|
||||
__shared__ int shared_final_k; // For final pass: remaining slots to fill
|
||||
const int src_buffer = pass % 2;
|
||||
const int dst_buffer = src_buffer ^ 1;
|
||||
|
||||
// Clamp buffered count to prevent overflow
|
||||
const int raw_buffered = shared_buffered_count[src_buffer];
|
||||
const int num_buffered =
|
||||
(raw_buffered < MAX_BUFFERED_ITEMS) ? raw_buffered : MAX_BUFFERED_ITEMS;
|
||||
|
||||
compute_cumulative_sum();
|
||||
|
||||
// Find threshold bin for this pass
|
||||
if (thread_id < RADIX && shared_histogram[0][thread_id] > remaining_k &&
|
||||
shared_histogram[0][thread_id + 1] <= remaining_k) {
|
||||
shared_threshold_bin = thread_id;
|
||||
shared_buffered_count[dst_buffer] = 0;
|
||||
shared_final_k = remaining_k - shared_histogram[0][thread_id + 1];
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
const int threshold_bin = shared_threshold_bin;
|
||||
remaining_k -= shared_histogram[0][threshold_bin + 1];
|
||||
|
||||
// Bit offset for this pass: 24, 16, 8, 0
|
||||
const int bit_offset = 24 - pass * 8;
|
||||
|
||||
// Early exit if threshold bin perfectly matches
|
||||
if (remaining_k == 0) {
|
||||
for (int i = thread_id; i < num_buffered; i += kThreadsPerBlock) {
|
||||
const int idx = buffered_indices[src_buffer][i];
|
||||
const uint32_t fp32_bits =
|
||||
convert_to_uint32_v2(logits[idx + logits_offset]);
|
||||
const int bin = (fp32_bits >> bit_offset) & 0xFF;
|
||||
if (bin > threshold_bin) {
|
||||
const int output_pos = ::atomicAdd(&shared_output_count, 1);
|
||||
output_indices[output_pos] = idx;
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
break;
|
||||
}
|
||||
|
||||
// Continue refinement
|
||||
__syncthreads();
|
||||
if (thread_id < RADIX + 1) {
|
||||
shared_histogram[0][thread_id] = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int i = thread_id; i < num_buffered; i += kThreadsPerBlock) {
|
||||
const int idx = buffered_indices[src_buffer][i];
|
||||
const float logit_value = logits[idx + logits_offset];
|
||||
const uint32_t fp32_bits = convert_to_uint32_v2(logit_value);
|
||||
const int bin = (fp32_bits >> bit_offset) & 0xFF;
|
||||
|
||||
if (bin > threshold_bin) {
|
||||
// Definitely in top-k
|
||||
const int output_pos = ::atomicAdd(&shared_output_count, 1);
|
||||
output_indices[output_pos] = idx;
|
||||
} else if (bin == threshold_bin) {
|
||||
if (pass == 3) {
|
||||
// Final pass (bits [7:0]): No more refinement possible
|
||||
// Fill remaining slots in reverse order to maintain descending order
|
||||
const int slot = ::atomicAdd(&shared_final_k, -1);
|
||||
if (slot > 0) {
|
||||
output_indices[TopK - slot] = idx;
|
||||
}
|
||||
} else {
|
||||
// Buffer for next pass and build next histogram
|
||||
const int buffer_pos =
|
||||
::atomicAdd(&shared_buffered_count[dst_buffer], 1);
|
||||
if (C10_LIKELY(buffer_pos < MAX_BUFFERED_ITEMS)) {
|
||||
buffered_indices[dst_buffer][buffer_pos] = idx;
|
||||
// Fused: Build histogram for next pass
|
||||
const int next_bit_offset = bit_offset - 8;
|
||||
const int next_bin = (fp32_bits >> next_bit_offset) & 0xFF;
|
||||
::atomicAdd(&shared_histogram[0][next_bin], 1);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
|
||||
__global__ __launch_bounds__(kThreadsPerBlock) void topk_kernel(
|
||||
const FastTopKParams params) {
|
||||
const auto& [input, row_starts, indices, lengths, input_stride] = params;
|
||||
const uint64_t batch_idx = blockIdx.x;
|
||||
const int logits_offset = row_starts == nullptr ? 0 : row_starts[batch_idx];
|
||||
const int seq_len = lengths[batch_idx];
|
||||
int* output_indices = indices + batch_idx * TopK;
|
||||
const float* logits = input + batch_idx * input_stride;
|
||||
|
||||
if (seq_len <= TopK) {
|
||||
// Shortcut: All elements are in top-k
|
||||
return naive_topk_cuda(logits, output_indices, seq_len);
|
||||
if (num_rows > 32 && max_smem_per_block >= 128 * 1024) {
|
||||
cudaError_t status = vllm::FilteredTopKRaggedTransform<float, int32_t>(
|
||||
logits.data_ptr<float>(), output.data_ptr<int32_t>(),
|
||||
lengths.data_ptr<int32_t>(), static_cast<uint32_t>(num_rows),
|
||||
static_cast<uint32_t>(k), static_cast<uint32_t>(stride), stream);
|
||||
TORCH_CHECK(status == cudaSuccess,
|
||||
"FilteredTopK failed: ", cudaGetErrorString(status));
|
||||
} else {
|
||||
return fast_topk_cuda_tl(logits, output_indices, logits_offset, seq_len);
|
||||
}
|
||||
}
|
||||
TORCH_CHECK(workspace.is_cuda(), "workspace must be CUDA tensor");
|
||||
TORCH_CHECK(workspace.dtype() == torch::kUInt8, "workspace must be uint8");
|
||||
|
||||
FastTopKParams get_params(
|
||||
const at::Tensor& score, const at::Tensor& lengths,
|
||||
std::optional<at::Tensor> row_starts_opt = std::nullopt,
|
||||
std::optional<at::Tensor> indices_opt = std::nullopt) {
|
||||
const int64_t batch_size = score.size(0);
|
||||
// Smem cap: smaller smem → more CTAs/group → more per-row parallelism for
|
||||
// large path. Empirically tuned.
|
||||
int effective_max_smem;
|
||||
if (num_rows <= 4) {
|
||||
effective_max_smem =
|
||||
std::min(max_smem_per_block, static_cast<int>(P::kSmemMedium));
|
||||
} else if (num_rows <= 8) {
|
||||
constexpr int kSmemCapMedium = 48 * 1024;
|
||||
effective_max_smem = std::min(max_smem_per_block, kSmemCapMedium);
|
||||
} else {
|
||||
effective_max_smem = max_smem_per_block;
|
||||
}
|
||||
|
||||
TORCH_CHECK(score.dim() == 2 && score.stride(1) == 1,
|
||||
"score must be 2D with contiguous rows");
|
||||
TORCH_CHECK(lengths.dim() == 1 && lengths.is_contiguous() &&
|
||||
lengths.size(0) == batch_size,
|
||||
"lengths must be 1D contiguous with size matching batch");
|
||||
size_t available_for_ordered =
|
||||
static_cast<size_t>(effective_max_smem) - P::kFixedSmemLarge;
|
||||
uint32_t max_chunk_elements =
|
||||
static_cast<uint32_t>(available_for_ordered / sizeof(uint32_t));
|
||||
|
||||
const int32_t* row_starts_ptr = nullptr;
|
||||
if (row_starts_opt.has_value()) {
|
||||
const auto& row_starts = *row_starts_opt;
|
||||
TORCH_CHECK(row_starts.dim() == 1 && row_starts.size(0) == batch_size,
|
||||
"row_starts must be 1D with size matching batch");
|
||||
row_starts_ptr = row_starts.data_ptr<int32_t>();
|
||||
uint32_t vec_size = 1;
|
||||
if (stride % 4 == 0)
|
||||
vec_size = 4;
|
||||
else if (stride % 2 == 0)
|
||||
vec_size = 2;
|
||||
|
||||
max_chunk_elements = (max_chunk_elements / vec_size) * vec_size;
|
||||
uint32_t min_chunk = vec_size * P::kThreadsPerBlock;
|
||||
if (max_chunk_elements < min_chunk) max_chunk_elements = min_chunk;
|
||||
|
||||
uint32_t ctas_per_group =
|
||||
(static_cast<uint32_t>(stride) + max_chunk_elements - 1) /
|
||||
max_chunk_elements;
|
||||
uint32_t chunk_size =
|
||||
(static_cast<uint32_t>(stride) + ctas_per_group - 1) / ctas_per_group;
|
||||
chunk_size = ((chunk_size + vec_size - 1) / vec_size) * vec_size;
|
||||
if (chunk_size > max_chunk_elements) chunk_size = max_chunk_elements;
|
||||
|
||||
size_t smem_size = P::kFixedSmemLarge + chunk_size * sizeof(uint32_t);
|
||||
if (smem_size < P::kSmemMedium) smem_size = P::kSmemMedium;
|
||||
|
||||
int occupancy = 1;
|
||||
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
|
||||
&occupancy, P::persistent_topk_kernel<4>, P::kThreadsPerBlock,
|
||||
smem_size);
|
||||
if (occupancy < 1) occupancy = 1;
|
||||
|
||||
uint32_t max_resident_ctas = static_cast<uint32_t>(num_sms) * occupancy;
|
||||
uint32_t num_groups = std::min(max_resident_ctas / ctas_per_group,
|
||||
static_cast<uint32_t>(num_rows));
|
||||
if (num_groups == 0) num_groups = 1;
|
||||
uint32_t total_ctas = num_groups * ctas_per_group;
|
||||
|
||||
size_t state_bytes = num_groups * sizeof(P::RadixRowState);
|
||||
TORCH_CHECK(workspace.size(0) >= static_cast<int64_t>(state_bytes),
|
||||
"workspace too small, need ", state_bytes, " bytes");
|
||||
|
||||
P::PersistentTopKParams params;
|
||||
params.input = logits.data_ptr<float>();
|
||||
params.output = output.data_ptr<int32_t>();
|
||||
params.lengths = lengths.data_ptr<int32_t>();
|
||||
params.num_rows = static_cast<uint32_t>(num_rows);
|
||||
params.stride = static_cast<uint32_t>(stride);
|
||||
params.chunk_size = chunk_size;
|
||||
params.row_states =
|
||||
reinterpret_cast<P::RadixRowState*>(workspace.data_ptr<uint8_t>());
|
||||
params.ctas_per_group = ctas_per_group;
|
||||
params.max_seq_len = static_cast<uint32_t>(max_seq_len);
|
||||
|
||||
#define LAUNCH_PERSISTENT(VS) \
|
||||
do { \
|
||||
auto kernel = &P::persistent_topk_kernel<VS>; \
|
||||
cudaError_t err = cudaFuncSetAttribute( \
|
||||
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size); \
|
||||
TORCH_CHECK(err == cudaSuccess, \
|
||||
"Failed to set smem: ", cudaGetErrorString(err)); \
|
||||
kernel<<<total_ctas, P::kThreadsPerBlock, smem_size, stream>>>(params); \
|
||||
} while (0)
|
||||
|
||||
if (vec_size == 4) {
|
||||
LAUNCH_PERSISTENT(4);
|
||||
} else if (vec_size == 2) {
|
||||
LAUNCH_PERSISTENT(2);
|
||||
} else {
|
||||
LAUNCH_PERSISTENT(1);
|
||||
}
|
||||
#undef LAUNCH_PERSISTENT
|
||||
}
|
||||
|
||||
int32_t* indices_ptr = nullptr;
|
||||
if (indices_opt.has_value()) {
|
||||
const auto& indices = *indices_opt;
|
||||
TORCH_CHECK(indices.dim() == 2 && indices.is_contiguous() &&
|
||||
indices.size(0) == batch_size && indices.size(1) == TopK,
|
||||
"indices must be 2D contiguous [batch, TopK]");
|
||||
indices_ptr = indices.data_ptr<int32_t>();
|
||||
}
|
||||
|
||||
return FastTopKParams{
|
||||
.input = score.data_ptr<float>(),
|
||||
.row_starts = row_starts_ptr,
|
||||
.indices = indices_ptr,
|
||||
.lengths = lengths.data_ptr<int32_t>(),
|
||||
.input_stride = score.stride(0),
|
||||
};
|
||||
}
|
||||
|
||||
template <auto* kernel_func, size_t smem_bytes>
|
||||
void setup_kernel_smem_once() {
|
||||
static const cudaError_t result = []() -> cudaError_t {
|
||||
#ifdef USE_ROCM
|
||||
auto func_ptr = reinterpret_cast<const void*>(kernel_func);
|
||||
cudaError_t err = cudaGetLastError();
|
||||
TORCH_CHECK(err == cudaSuccess,
|
||||
"persistent_topk failed: ", cudaGetErrorString(err));
|
||||
#else
|
||||
auto func_ptr = kernel_func;
|
||||
TORCH_CHECK(false, "persistent_topk is not supported on ROCm");
|
||||
#endif
|
||||
return cudaFuncSetAttribute(
|
||||
func_ptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes);
|
||||
}();
|
||||
|
||||
TORCH_CHECK(
|
||||
result == cudaSuccess,
|
||||
"Failed to set kernel shared memory limit: ", cudaGetErrorString(result));
|
||||
}
|
||||
|
||||
} // namespace vllm
|
||||
|
||||
void large_context_topk(
|
||||
const torch::Tensor& logits, torch::Tensor& indices,
|
||||
const torch::Tensor& seq_lens,
|
||||
std::optional<torch::Tensor> row_starts = std::nullopt) {
|
||||
TORCH_CHECK(logits.is_cuda(), "logits must be a CUDA tensor");
|
||||
TORCH_CHECK(indices.is_cuda(), "indices must be a CUDA tensor");
|
||||
TORCH_CHECK(seq_lens.is_cuda(), "seq_lens must be a CUDA tensor");
|
||||
if (row_starts.has_value()) {
|
||||
TORCH_CHECK(row_starts->is_cuda(), "row_starts must be a CUDA tensor");
|
||||
}
|
||||
|
||||
const auto params = vllm::get_params(logits, seq_lens, row_starts, indices);
|
||||
const int64_t batch_size = logits.size(0);
|
||||
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
const dim3 grid(static_cast<uint32_t>(batch_size));
|
||||
const dim3 block(vllm::kThreadsPerBlock);
|
||||
|
||||
vllm::setup_kernel_smem_once<vllm::topk_kernel, vllm::kSmem>();
|
||||
vllm::topk_kernel<<<grid, block, vllm::kSmem, stream>>>(params);
|
||||
|
||||
const cudaError_t result = cudaGetLastError();
|
||||
TORCH_CHECK(result == cudaSuccess,
|
||||
"large_context_topk kernel failed: ", cudaGetErrorString(result));
|
||||
}
|
||||
+40
-16
@@ -110,6 +110,18 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"silu_and_mul_quant(Tensor! result, Tensor input, Tensor scale) -> ()");
|
||||
ops.impl("silu_and_mul_quant", torch::kCUDA, &silu_and_mul_quant);
|
||||
|
||||
// Fused SiLU+Mul + per-block quantization
|
||||
ops.def(
|
||||
"silu_and_mul_per_block_quant("
|
||||
"Tensor! out, "
|
||||
"Tensor input, "
|
||||
"Tensor! scales, "
|
||||
"int group_size, "
|
||||
"Tensor? scale_ub=None, "
|
||||
"bool is_scale_transposed=False) -> ()");
|
||||
ops.impl("silu_and_mul_per_block_quant", torch::kCUDA,
|
||||
&silu_and_mul_per_block_quant);
|
||||
|
||||
ops.def("mul_and_silu(Tensor! out, Tensor input) -> ()");
|
||||
ops.impl("mul_and_silu", torch::kCUDA, &mul_and_silu);
|
||||
|
||||
@@ -161,7 +173,8 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"fused_qk_norm_rope(Tensor! qkv, int num_heads_q, "
|
||||
"int num_heads_k, int num_heads_v, int head_dim, float eps, "
|
||||
"Tensor q_weight, Tensor k_weight, Tensor cos_sin_cache, "
|
||||
"bool is_neox, Tensor position_ids) -> ()");
|
||||
"bool is_neox, Tensor position_ids, "
|
||||
"int forced_token_heads_per_warp=-1) -> ()");
|
||||
ops.impl("fused_qk_norm_rope", torch::kCUDA, &fused_qk_norm_rope);
|
||||
|
||||
// Apply repetition penalties to logits in-place
|
||||
@@ -185,10 +198,9 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
ops.impl("top_k_per_row_decode", torch::kCUDA, &top_k_per_row_decode);
|
||||
|
||||
ops.def(
|
||||
"large_context_topk(Tensor score, Tensor indices, Tensor lengths, "
|
||||
"Tensor? "
|
||||
"row_starts_opt) -> ()");
|
||||
ops.impl("large_context_topk", torch::kCUDA, &large_context_topk);
|
||||
"persistent_topk(Tensor logits, Tensor lengths, Tensor! output, "
|
||||
"Tensor workspace, int k, int max_seq_len) -> ()");
|
||||
ops.impl("persistent_topk", torch::kCUDA, &persistent_topk);
|
||||
|
||||
// Layernorm-quant
|
||||
// Apply Root Mean Square (RMS) Normalization to the input tensor.
|
||||
@@ -233,17 +245,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
|
||||
// Quantization ops
|
||||
#ifndef USE_ROCM
|
||||
// Fused SiLU+Mul + per-block quantization
|
||||
ops.def(
|
||||
"silu_and_mul_per_block_quant("
|
||||
"Tensor! out, "
|
||||
"Tensor input, "
|
||||
"Tensor! scales, "
|
||||
"int group_size, "
|
||||
"Tensor? scale_ub=None, "
|
||||
"bool is_scale_transposed=False) -> ()");
|
||||
ops.impl("silu_and_mul_per_block_quant", torch::kCUDA,
|
||||
&silu_and_mul_per_block_quant);
|
||||
// DeepSeek V3 fused A GEMM (SM 9.0+, bf16 only, 1-16 tokens).
|
||||
ops.def(
|
||||
"dsv3_fused_a_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()");
|
||||
@@ -496,6 +497,29 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
|
||||
"Tensor? b_qzeros, "
|
||||
"SymInt n, SymInt group_size, SymInt sm_count, SymInt sm_version, SymInt "
|
||||
"CUBLAS_M_THRESHOLD, bool has_zp, bool n32k16_reorder) -> Tensor");
|
||||
|
||||
ops.def(
|
||||
"minimax_allreduce_rms("
|
||||
"Tensor input,"
|
||||
"Tensor norm_weight,"
|
||||
"Tensor workspace,"
|
||||
"int rank,"
|
||||
"int nranks,"
|
||||
"float eps) -> Tensor");
|
||||
ops.impl("minimax_allreduce_rms", torch::kCUDA, &minimax_allreduce_rms);
|
||||
ops.def(
|
||||
"minimax_allreduce_rms_qk("
|
||||
"Tensor qkv,"
|
||||
"Tensor norm_weight_q,"
|
||||
"Tensor norm_weight_k,"
|
||||
"Tensor workspace,"
|
||||
"int q_size,"
|
||||
"int kv_size,"
|
||||
"int rank,"
|
||||
"int nranks,"
|
||||
"float eps) -> (Tensor, Tensor)");
|
||||
ops.impl("minimax_allreduce_rms_qk", torch::kCUDA, &minimax_allreduce_rms_qk);
|
||||
|
||||
// conditionally compiled so impl in source file
|
||||
#endif
|
||||
}
|
||||
|
||||
+28
-47
@@ -22,7 +22,7 @@
|
||||
# docker buildx bake -f docker/docker-bake.hcl -f docker/versions.json
|
||||
# =============================================================================
|
||||
|
||||
ARG CUDA_VERSION=12.9.1
|
||||
ARG CUDA_VERSION=13.0.0
|
||||
ARG PYTHON_VERSION=3.12
|
||||
ARG UBUNTU_VERSION=22.04
|
||||
|
||||
@@ -37,7 +37,7 @@ ARG UBUNTU_VERSION=22.04
|
||||
# compatibility with other Linux OSes. The main reason for this is that the
|
||||
# glibc version is baked into the distro, and binaries built with one glibc
|
||||
# version are not backwards compatible with OSes that use an earlier version.
|
||||
ARG BUILD_BASE_IMAGE=nvidia/cuda:${CUDA_VERSION}-devel-ubuntu20.04
|
||||
ARG BUILD_BASE_IMAGE=nvidia/cuda:${CUDA_VERSION}-devel-ubuntu22.04
|
||||
# Using cuda base image with minimal dependencies necessary for JIT compilation (FlashInfer, DeepGEMM, EP kernels)
|
||||
ARG FINAL_BASE_IMAGE=nvidia/cuda:${CUDA_VERSION}-base-ubuntu${UBUNTU_VERSION}
|
||||
|
||||
@@ -204,7 +204,7 @@ ARG PYTORCH_CUDA_INDEX_BASE_URL
|
||||
ARG PYTORCH_NIGHTLY
|
||||
|
||||
# Install build dependencies
|
||||
COPY requirements/build.txt requirements/build.txt
|
||||
COPY requirements/build/cuda.txt requirements/build/cuda.txt
|
||||
COPY use_existing_torch.py use_existing_torch.py
|
||||
COPY --from=base /workspace/torch_lib_versions.txt torch_lib_versions.txt
|
||||
|
||||
@@ -219,13 +219,13 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
if [ "${PYTORCH_NIGHTLY}" = "1" ]; then \
|
||||
echo "Installing build requirements without torch..." \
|
||||
&& python3 use_existing_torch.py --prefix \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build.txt \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build/cuda.txt \
|
||||
&& echo "Installing torch nightly..." \
|
||||
&& uv pip install --python /opt/venv/bin/python3 $(cat torch_lib_versions.txt | grep -i "^torch=" | xargs) --pre \
|
||||
--index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/nightly/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.'); \
|
||||
else \
|
||||
echo "Installing build requirements..." \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build.txt \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build/cuda.txt \
|
||||
--extra-index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.'); \
|
||||
fi
|
||||
|
||||
@@ -315,7 +315,7 @@ RUN --mount=type=cache,target=/root/.cache/ccache \
|
||||
#################### CSRC BUILD IMAGE ####################
|
||||
|
||||
#################### EXTENSIONS BUILD IMAGE ####################
|
||||
# Build DeepGEMM, DeepEP - runs in PARALLEL with csrc-build
|
||||
# Build DeepEP - runs in PARALLEL with csrc-build
|
||||
# This stage is independent and doesn't affect csrc cache
|
||||
FROM base AS extensions-build
|
||||
ARG CUDA_VERSION
|
||||
@@ -327,21 +327,6 @@ ENV UV_LINK_MODE=copy
|
||||
|
||||
WORKDIR /workspace
|
||||
|
||||
# Build DeepGEMM wheel
|
||||
# Default moved here from tools/install_deepgemm.sh for centralized version management
|
||||
ARG DEEPGEMM_GIT_REF=477618cd51baffca09c4b0b87e97c03fe827ef03
|
||||
COPY tools/install_deepgemm.sh /tmp/install_deepgemm.sh
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
mkdir -p /tmp/deepgemm/dist && \
|
||||
VLLM_DOCKER_BUILD_CONTEXT=1 TORCH_CUDA_ARCH_LIST="9.0a 10.0a" /tmp/install_deepgemm.sh \
|
||||
--cuda-version "${CUDA_VERSION}" \
|
||||
${DEEPGEMM_GIT_REF:+--ref "$DEEPGEMM_GIT_REF"} \
|
||||
--wheel-dir /tmp/deepgemm/dist || \
|
||||
echo "DeepGEMM build skipped (CUDA version requirement not met)"
|
||||
|
||||
# Ensure the wheel dir exists so COPY won't fail when DeepGEMM is skipped
|
||||
RUN mkdir -p /tmp/deepgemm/dist && touch /tmp/deepgemm/dist/.deepgemm_skipped
|
||||
|
||||
# Build DeepEP wheels
|
||||
COPY tools/ep_kernels/install_python_libraries.sh /tmp/install_python_libraries.sh
|
||||
# Defaults moved here from tools/ep_kernels/install_python_libraries.sh for centralized version management
|
||||
@@ -370,7 +355,7 @@ ARG PYTORCH_CUDA_INDEX_BASE_URL
|
||||
ARG PYTORCH_NIGHTLY
|
||||
|
||||
# Install build dependencies
|
||||
COPY requirements/build.txt requirements/build.txt
|
||||
COPY requirements/build/cuda.txt requirements/build/cuda.txt
|
||||
COPY use_existing_torch.py use_existing_torch.py
|
||||
COPY --from=base /workspace/torch_lib_versions.txt torch_lib_versions.txt
|
||||
|
||||
@@ -385,13 +370,13 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
if [ "${PYTORCH_NIGHTLY}" = "1" ]; then \
|
||||
echo "Installing build requirements without torch..." \
|
||||
&& python3 use_existing_torch.py --prefix \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build.txt \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build/cuda.txt \
|
||||
&& echo "Installing torch nightly..." \
|
||||
&& uv pip install --python /opt/venv/bin/python3 $(cat torch_lib_versions.txt | grep -i "^torch=" | xargs) --pre \
|
||||
--index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/nightly/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.'); \
|
||||
else \
|
||||
echo "Installing build requirements..." \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build.txt \
|
||||
&& uv pip install --python /opt/venv/bin/python3 -r requirements/build/cuda.txt \
|
||||
--extra-index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.'); \
|
||||
fi
|
||||
|
||||
@@ -426,7 +411,6 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
python3 setup.py bdist_wheel --dist-dir=dist --py-limited-api=cp38
|
||||
|
||||
# Copy extension wheels from extensions-build stage for later use
|
||||
COPY --from=extensions-build /tmp/deepgemm/dist /tmp/deepgemm/dist
|
||||
COPY --from=extensions-build /tmp/ep_kernels_workspace/dist /tmp/ep_kernels_workspace/dist
|
||||
|
||||
# Check the size of the wheel if RUN_WHEEL_CHECK is true
|
||||
@@ -466,8 +450,8 @@ ARG PYTORCH_NIGHTLY
|
||||
|
||||
# Install development dependencies
|
||||
COPY requirements/lint.txt requirements/lint.txt
|
||||
COPY requirements/test.in requirements/test.in
|
||||
COPY requirements/test.txt requirements/test.txt
|
||||
COPY requirements/test/cuda.in requirements/test/cuda.in
|
||||
COPY requirements/test/cuda.txt requirements/test/cuda.txt
|
||||
COPY requirements/dev.txt requirements/dev.txt
|
||||
COPY use_existing_torch.py use_existing_torch.py
|
||||
COPY --from=base /workspace/torch_lib_versions.txt torch_lib_versions.txt
|
||||
@@ -475,8 +459,8 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
if [ "${PYTORCH_NIGHTLY}" = "1" ]; then \
|
||||
echo "Installing dev requirements plus torch nightly..." \
|
||||
&& python3 use_existing_torch.py --prefix \
|
||||
&& cat torch_lib_versions.txt >> requirements/test.in \
|
||||
&& uv pip compile requirements/test.in -o requirements/test.txt --index-strategy unsafe-best-match \
|
||||
&& cat torch_lib_versions.txt >> requirements/test/cuda.in \
|
||||
&& uv pip compile requirements/test/cuda.in -o requirements/test/cuda.txt --index-strategy unsafe-best-match \
|
||||
--extra-index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/nightly/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.') \
|
||||
&& uv pip install --python /opt/venv/bin/python3 $(cat torch_lib_versions.txt | xargs) --pre \
|
||||
-r requirements/dev.txt \
|
||||
@@ -546,17 +530,23 @@ RUN apt-get update -y \
|
||||
# Install CUDA development tools for runtime JIT compilation
|
||||
# (FlashInfer, DeepGEMM, EP kernels all require compilation at runtime)
|
||||
RUN CUDA_VERSION_DASH=$(echo $CUDA_VERSION | cut -d. -f1,2 | tr '.' '-') && \
|
||||
CUDA_VERSION_SHORT=$(echo $CUDA_VERSION | cut -d. -f1,2) && \
|
||||
apt-get update -y && \
|
||||
apt-get install -y --no-install-recommends \
|
||||
apt-get install -y --no-install-recommends --allow-change-held-packages \
|
||||
cuda-nvcc-${CUDA_VERSION_DASH} \
|
||||
cuda-cudart-${CUDA_VERSION_DASH} \
|
||||
cuda-nvrtc-${CUDA_VERSION_DASH} \
|
||||
cuda-cuobjdump-${CUDA_VERSION_DASH} \
|
||||
libcurand-dev-${CUDA_VERSION_DASH} \
|
||||
libcublas-${CUDA_VERSION_DASH} \
|
||||
# Fixes nccl_allocator requiring nccl.h at runtime
|
||||
# https://github.com/vllm-project/vllm/blob/1336a1ea244fa8bfd7e72751cabbdb5b68a0c11a/vllm/distributed/device_communicators/pynccl_allocator.py#L22
|
||||
libnccl-dev && \
|
||||
# Required by fastsafetensors (fixes #20384)
|
||||
libnuma-dev && \
|
||||
# Fixes nccl_allocator requiring nccl.h at runtime
|
||||
# https://github.com/vllm-project/vllm/blob/1336a1ea244fa8bfd7e72751cabbdb5b68a0c11a/vllm/distributed/device_communicators/pynccl_allocator.py#L22
|
||||
# NCCL packages don't use the cuda-MAJOR-MINOR naming convention,
|
||||
# so we pin the version to match our CUDA version
|
||||
NCCL_VER=$(apt-cache madison libnccl-dev | grep "+cuda${CUDA_VERSION_SHORT}" | head -1 | awk -F'|' '{gsub(/^ +| +$/, "", $2); print $2}') && \
|
||||
apt-get install -y --no-install-recommends --allow-change-held-packages libnccl-dev=${NCCL_VER} libnccl2=${NCCL_VER} && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Install uv for faster pip installs
|
||||
@@ -689,15 +679,6 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
. /etc/environment && \
|
||||
uv pip list
|
||||
|
||||
# Install deepgemm wheel that has been built in the `build` stage
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
--mount=type=bind,from=build,source=/tmp/deepgemm/dist,target=/tmp/deepgemm/dist,ro \
|
||||
sh -c 'if ls /tmp/deepgemm/dist/*.whl >/dev/null 2>&1; then \
|
||||
uv pip install --system /tmp/deepgemm/dist/*.whl; \
|
||||
else \
|
||||
echo "No DeepGEMM wheels to install; skipping."; \
|
||||
fi'
|
||||
|
||||
# Pytorch now installs NVSHMEM, setting LD_LIBRARY_PATH
|
||||
ENV LD_LIBRARY_PATH=/usr/local/cuda/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
@@ -746,8 +727,8 @@ ARG PYTORCH_NIGHTLY
|
||||
|
||||
# Install development dependencies (for testing)
|
||||
COPY requirements/lint.txt requirements/lint.txt
|
||||
COPY requirements/test.in requirements/test.in
|
||||
COPY requirements/test.txt requirements/test.txt
|
||||
COPY requirements/test/cuda.in requirements/test/cuda.in
|
||||
COPY requirements/test/cuda.txt requirements/test/cuda.txt
|
||||
COPY requirements/dev.txt requirements/dev.txt
|
||||
COPY use_existing_torch.py use_existing_torch.py
|
||||
COPY --from=base /workspace/torch_lib_versions.txt torch_lib_versions.txt
|
||||
@@ -757,8 +738,8 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
if [ "${PYTORCH_NIGHTLY}" = "1" ]; then \
|
||||
echo "Installing dev requirements plus torch nightly..." \
|
||||
&& python3 use_existing_torch.py --prefix \
|
||||
&& cat torch_lib_versions.txt >> requirements/test.in \
|
||||
&& uv pip compile requirements/test.in -o requirements/test.txt --index-strategy unsafe-best-match \
|
||||
&& cat torch_lib_versions.txt >> requirements/test/cuda.in \
|
||||
&& uv pip compile requirements/test/cuda.in -o requirements/test/cuda.txt --index-strategy unsafe-best-match \
|
||||
--extra-index-url ${PYTORCH_CUDA_INDEX_BASE_URL}/nightly/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.') \
|
||||
&& uv pip install --system $(cat torch_lib_versions.txt | xargs) --pre \
|
||||
-r requirements/dev.txt \
|
||||
@@ -822,7 +803,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install --system -r /tmp/kv_connectors.txt --no-build || ( \
|
||||
# if the above fails, install from source
|
||||
apt-get update -y && \
|
||||
apt-get install -y --no-install-recommends ${BUILD_PKGS} && \
|
||||
apt-get install -y --no-install-recommends --allow-change-held-packages ${BUILD_PKGS} && \
|
||||
uv pip install --system -r /tmp/kv_connectors.txt --no-build-isolation && \
|
||||
apt-get purge -y ${BUILD_PKGS} && \
|
||||
# clean up -dev packages, keep runtime libraries
|
||||
|
||||
+14
-12
@@ -107,10 +107,10 @@ RUN if [ "$TARGETARCH" = "arm64" ] && [ "$VLLM_CPU_X86" != "0" ]; then \
|
||||
fi
|
||||
|
||||
# Copy build requirements
|
||||
COPY requirements/cpu-build.txt requirements/build.txt
|
||||
COPY requirements/build/cpu.txt requirements/build/cpu.txt
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install -r requirements/build.txt
|
||||
uv pip install -r requirements/build/cpu.txt
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -127,26 +127,28 @@ FROM base AS vllm-test-deps
|
||||
WORKDIR /vllm-workspace
|
||||
|
||||
# Copy test requirements
|
||||
COPY requirements/test.in requirements/cpu-test.in
|
||||
COPY requirements/test/cuda.in requirements/test/cpu.in
|
||||
|
||||
RUN \
|
||||
sed -i '/mamba_ssm/d' requirements/cpu-test.in && \
|
||||
sed -i '/mamba_ssm/d' requirements/test/cpu.in && \
|
||||
remove_packages_not_supported_on_aarch64() { \
|
||||
case "$(uname -m)" in \
|
||||
aarch64|arm64) \
|
||||
sed -i '/decord/d' requirements/cpu-test.in; \
|
||||
sed -i '/terratorch/d' requirements/cpu-test.in; \
|
||||
sed -i '/decord/d' requirements/test/cpu.in; \
|
||||
sed -i '/terratorch/d' requirements/test/cpu.in; \
|
||||
;; \
|
||||
esac; \
|
||||
}; \
|
||||
remove_packages_not_supported_on_aarch64 && \
|
||||
sed -i 's/^torch==.*/torch==2.10.0/g' requirements/cpu-test.in && \
|
||||
sed -i 's/torchaudio.*/torchaudio/g' requirements/cpu-test.in && \
|
||||
sed -i 's/torchvision.*/torchvision/g' requirements/cpu-test.in && \
|
||||
uv pip compile requirements/cpu-test.in -o requirements/cpu-test.txt --index-strategy unsafe-best-match --torch-backend cpu
|
||||
sed -i 's/^torch==.*/torch==2.11.0/g' requirements/test/cpu.in && \
|
||||
sed -i 's/torchaudio.*/torchaudio/g' requirements/test/cpu.in && \
|
||||
sed -i 's/torchvision.*/torchvision/g' requirements/test/cpu.in && \
|
||||
# Related issue: https://github.com/vllm-project/vllm/pull/38800#issuecomment-4228314305
|
||||
sed -i 's/^sentence-transformers.*/sentence-transformers==5.3.0/g' requirements/test/cpu.in && \
|
||||
uv pip compile requirements/test/cpu.in -o requirements/test/cpu.txt --index-strategy unsafe-best-match --torch-backend cpu
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install -r requirements/cpu-test.txt
|
||||
uv pip install -r requirements/test/cpu.txt
|
||||
|
||||
######################### DEV IMAGE #########################
|
||||
FROM vllm-build AS vllm-dev
|
||||
@@ -168,7 +170,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
--mount=type=bind,source=.git,target=.git \
|
||||
VLLM_TARGET_DEVICE=cpu python3 setup.py develop
|
||||
|
||||
COPY --from=vllm-test-deps /vllm-workspace/requirements/cpu-test.txt requirements/test.txt
|
||||
COPY --from=vllm-test-deps /vllm-workspace/requirements/test/cpu.txt requirements/test/cpu.txt
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install -r requirements/dev.txt && \
|
||||
|
||||
@@ -107,7 +107,7 @@ COPY . .
|
||||
RUN python3 use_existing_torch.py
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install --system -r requirements/build.txt
|
||||
uv pip install --system -r requirements/build/cuda.txt
|
||||
|
||||
ARG GIT_REPO_CHECK=0
|
||||
RUN --mount=type=bind,source=.git,target=.git \
|
||||
@@ -261,7 +261,7 @@ FROM vllm-base as test
|
||||
COPY tests/ tests/
|
||||
|
||||
# install build and runtime dependencies without stable torch version
|
||||
COPY requirements/nightly_torch_test.txt requirements/nightly_torch_test.txt
|
||||
COPY requirements/test/nightly-torch.txt requirements/test/nightly-torch.txt
|
||||
|
||||
# This timeout (in seconds) is necessary when installing some dependencies via uv since it's likely to time out
|
||||
# Reference: https://github.com/astral-sh/uv/pull/1694
|
||||
@@ -277,7 +277,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
ENV HF_HUB_ENABLE_HF_TRANSFER 1
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install --system -r requirements/nightly_torch_test.txt
|
||||
uv pip install --system -r requirements/test/nightly-torch.txt
|
||||
|
||||
# Logging to confirm the torch versions
|
||||
RUN pip freeze | grep -E 'torch|vllm|flashinfer'
|
||||
|
||||
@@ -251,7 +251,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
make -C /numactl install && \
|
||||
# sentencepiece.pc is in some pkgconfig inside uv cache
|
||||
export PKG_CONFIG_PATH=$(find / -type d -name "pkgconfig" 2>/dev/null | tr '\n' ':') && \
|
||||
nanobind_DIR=$(uv pip show nanobind | grep Location | sed 's/^Location: //;s/$/\/nanobind\/cmake/') && uv pip install -r /src/requirements/common.txt -r /src/requirements/cpu.txt -r /src/requirements/build.txt --no-build-isolation && \
|
||||
nanobind_DIR=$(uv pip show nanobind | grep Location | sed 's/^Location: //;s/$/\/nanobind\/cmake/') && uv pip install -r /src/requirements/common.txt -r /src/requirements/cpu.txt -r /src/requirements/build/cuda.txt --no-build-isolation && \
|
||||
cd /src/ && \
|
||||
uv build --wheel --out-dir /vllmwheel/ --no-build-isolation && \
|
||||
uv pip install /vllmwheel/*.whl
|
||||
|
||||
+22
-20
@@ -19,7 +19,8 @@ ENV PYTORCH_ROCM_ARCH=${ARG_PYTORCH_ROCM_ARCH:-${PYTORCH_ROCM_ARCH}}
|
||||
# Install some basic utilities
|
||||
RUN apt-get update -q -y && apt-get install -q -y \
|
||||
sqlite3 libsqlite3-dev libfmt-dev libmsgpack-dev libsuitesparse-dev \
|
||||
apt-transport-https ca-certificates wget curl
|
||||
apt-transport-https ca-certificates wget curl \
|
||||
libnuma-dev
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
# Remove sccache only if not using sccache (it exists in base image from Dockerfile.rocm_base)
|
||||
ARG USE_SCCACHE
|
||||
@@ -328,14 +329,14 @@ RUN --mount=type=bind,from=export_vllm,src=/,target=/install \
|
||||
--mount=type=cache,target=/root/.cache/uv \
|
||||
cd /install \
|
||||
&& uv pip install --system -r requirements/rocm.txt \
|
||||
&& uv pip install --system -r requirements/rocm-test.txt \
|
||||
&& uv pip install --system -r requirements/test/rocm.txt \
|
||||
&& pip uninstall -y vllm \
|
||||
&& uv pip install --system *.whl
|
||||
|
||||
# Verify that PyTorch is the ROCm build, not CUDA
|
||||
RUN python3 -c "import torch; assert torch.version.hip is not None, \
|
||||
f'Expected ROCm PyTorch but got CUDA (torch.version.cuda={torch.version.cuda}, torch.version.hip={torch.version.hip})'; \
|
||||
print(f'Verified: PyTorch {torch.__version__} with ROCm (HIP {torch.version.hip})')"
|
||||
# Persist the built wheel in the image so python_only_compile_rocm.sh can
|
||||
# reinstall it after removing compilers. The bind-mounted /install contents
|
||||
# above are not available once that RUN step completes.
|
||||
COPY --from=export_vllm /*.whl /opt/vllm-wheels/
|
||||
|
||||
# Install RIXL wheel
|
||||
RUN --mount=type=bind,from=build_rixl,src=/app/install,target=/rixl_install \
|
||||
@@ -390,20 +391,21 @@ ENV MIOPEN_DEBUG_CONV_GEMM=0
|
||||
RUN mkdir src && mv vllm src/vllm
|
||||
|
||||
# This is a workaround to ensure pytest exits with the correct status code in CI tests.
|
||||
RUN cat << 'EOF' > /vllm-workspace/conftest.py
|
||||
import os
|
||||
|
||||
_exit_code = 1
|
||||
|
||||
def pytest_sessionfinish(session, exitstatus):
|
||||
global _exit_code
|
||||
_exit_code = int(exitstatus)
|
||||
|
||||
def pytest_unconfigure(config):
|
||||
sys.stdout.flush()
|
||||
sys.stderr.flush()
|
||||
os._exit(_exit_code)
|
||||
EOF
|
||||
RUN printf '%s\n' \
|
||||
'import os' \
|
||||
'' \
|
||||
'_exit_code = 1' \
|
||||
'' \
|
||||
'def pytest_sessionfinish(session, exitstatus):' \
|
||||
' global _exit_code' \
|
||||
' _exit_code = int(exitstatus)' \
|
||||
'' \
|
||||
'def pytest_unconfigure(config):' \
|
||||
' import sys' \
|
||||
' sys.stdout.flush()' \
|
||||
' sys.stderr.flush()' \
|
||||
' os._exit(_exit_code)' \
|
||||
> /vllm-workspace/conftest.py
|
||||
|
||||
# -----------------------
|
||||
# Final vLLM image
|
||||
|
||||
@@ -9,7 +9,7 @@ ARG PYTORCH_AUDIO_BRANCH="v2.9.0"
|
||||
ARG PYTORCH_AUDIO_REPO="https://github.com/pytorch/audio.git"
|
||||
ARG FA_BRANCH="0e60e394"
|
||||
ARG FA_REPO="https://github.com/Dao-AILab/flash-attention.git"
|
||||
ARG AITER_BRANCH="v0.1.10.post2"
|
||||
ARG AITER_BRANCH="v0.1.10.post3"
|
||||
ARG AITER_REPO="https://github.com/ROCm/aiter.git"
|
||||
ARG MORI_BRANCH="2d02c6a9"
|
||||
ARG MORI_REPO="https://github.com/ROCm/mori.git"
|
||||
@@ -112,10 +112,14 @@ FROM base AS build_triton
|
||||
ARG TRITON_BRANCH
|
||||
ARG TRITON_REPO
|
||||
RUN git clone ${TRITON_REPO}
|
||||
# Cherry picking the following
|
||||
# https://github.com/triton-lang/triton/pull/8991
|
||||
# https://github.com/triton-lang/triton/pull/9541
|
||||
RUN cd triton \
|
||||
&& git checkout ${TRITON_BRANCH} \
|
||||
&& git config --global user.email "you@example.com" && git config --global user.name "Your Name" \
|
||||
&& git cherry-pick 555d04f \
|
||||
&& git cherry-pick dd998b6 \
|
||||
&& if [ ! -f setup.py ]; then cd python; fi \
|
||||
&& python3 setup.py bdist_wheel --dist-dir=dist \
|
||||
&& mkdir -p /app/install && cp dist/*.whl /app/install
|
||||
|
||||
@@ -93,13 +93,13 @@ RUN curl https://sh.rustup.rs -sSf | sh -s -- -y && \
|
||||
|
||||
FROM python-install AS torch-vision
|
||||
# Install torchvision
|
||||
ARG TORCH_VISION_VERSION=v0.25.0
|
||||
ARG TORCH_VISION_VERSION=v0.26.0
|
||||
WORKDIR /tmp
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
git clone https://github.com/pytorch/vision.git && \
|
||||
cd vision && \
|
||||
git checkout $TORCH_VISION_VERSION && \
|
||||
uv pip install torch==2.10.0 --index-url https://download.pytorch.org/whl/cpu && \
|
||||
uv pip install torch==2.11.0 --index-url https://download.pytorch.org/whl/cpu && \
|
||||
python setup.py bdist_wheel
|
||||
|
||||
FROM python-install AS hf-xet-builder
|
||||
@@ -253,7 +253,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
NUMBA_WHL_FILE=$(ls /tmp/numba-wheels/*.whl) && \
|
||||
OPENCV_WHL_FILE=$(ls /tmp/opencv-wheels/*.whl) && \
|
||||
OUTLINES_CORE_WHL_FILE=$(ls /tmp/outlines-core/dist/*.whl) && \
|
||||
uv pip install -v \
|
||||
uv pip install -v \
|
||||
$ARROW_WHL_FILE \
|
||||
$VISION_WHL_FILE \
|
||||
$HF_XET_WHL_FILE \
|
||||
@@ -262,7 +262,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
$OPENCV_WHL_FILE \
|
||||
$OUTLINES_CORE_WHL_FILE \
|
||||
--index-strategy unsafe-best-match \
|
||||
-r requirements/cpu-build.txt \
|
||||
-r requirements/build/cpu.txt \
|
||||
-r requirements/cpu.txt
|
||||
|
||||
|
||||
|
||||
@@ -76,20 +76,14 @@ ENV UV_LINK_MODE="copy"
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
--mount=type=bind,src=requirements/common.txt,target=/workspace/vllm/requirements/common.txt \
|
||||
--mount=type=bind,src=requirements/xpu.txt,target=/workspace/vllm/requirements/xpu.txt \
|
||||
--mount=type=bind,src=requirements/xpu-test.in,target=/workspace/vllm/requirements/xpu-test.in \
|
||||
--mount=type=bind,src=requirements/test/xpu.txt,target=/workspace/vllm/requirements/test/xpu.txt \
|
||||
uv pip install --upgrade pip && \
|
||||
uv pip install -r requirements/xpu.txt && \
|
||||
uv pip compile /workspace/vllm/requirements/xpu-test.in \
|
||||
-o /workspace/vllm/requirements/xpu-test.txt \
|
||||
-c /workspace/vllm/requirements/xpu.txt \
|
||||
--index-strategy unsafe-best-match \
|
||||
--extra-index-url ${PIP_EXTRA_INDEX_URL} \
|
||||
--python-version ${PYTHON_VERSION} && \
|
||||
uv pip install grpcio-tools protobuf nanobind && \
|
||||
source /opt/intel/oneapi/setvars.sh --force && \
|
||||
source /opt/intel/oneapi/ccl/2021.15/env/vars.sh --force && \
|
||||
export CMAKE_PREFIX_PATH="$(python3 -c 'import site; print(site.getsitepackages()[0])'):${CMAKE_PREFIX_PATH}" && \
|
||||
uv pip install --no-build-isolation -r /workspace/vllm/requirements/xpu-test.txt
|
||||
uv pip install --no-build-isolation -r /workspace/vllm/requirements/test/xpu.txt
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"_comment": "Auto-generated from Dockerfile ARGs. Do not edit manually. Run: python tools/generate_versions_json.py",
|
||||
"variable": {
|
||||
"CUDA_VERSION": {
|
||||
"default": "12.9.1"
|
||||
"default": "13.0.0"
|
||||
},
|
||||
"PYTHON_VERSION": {
|
||||
"default": "3.12"
|
||||
@@ -11,10 +11,10 @@
|
||||
"default": "22.04"
|
||||
},
|
||||
"BUILD_BASE_IMAGE": {
|
||||
"default": "nvidia/cuda:12.9.1-devel-ubuntu20.04"
|
||||
"default": "nvidia/cuda:13.0.0-devel-ubuntu22.04"
|
||||
},
|
||||
"FINAL_BASE_IMAGE": {
|
||||
"default": "nvidia/cuda:12.9.1-base-ubuntu22.04"
|
||||
"default": "nvidia/cuda:13.0.0-base-ubuntu22.04"
|
||||
},
|
||||
"GET_PIP_URL": {
|
||||
"default": "https://bootstrap.pypa.io/get-pip.py"
|
||||
@@ -52,9 +52,6 @@
|
||||
"vllm_target_device": {
|
||||
"default": "cuda"
|
||||
},
|
||||
"DEEPGEMM_GIT_REF": {
|
||||
"default": "477618cd51baffca09c4b0b87e97c03fe827ef03"
|
||||
},
|
||||
"DEEPEP_COMMIT_HASH": {
|
||||
"default": "73b6ea4"
|
||||
},
|
||||
|
||||
+27
-13
@@ -25,7 +25,7 @@ hide:
|
||||
|
||||
vLLM is a fast and easy-to-use library for LLM inference and serving.
|
||||
|
||||
Originally developed in the [Sky Computing Lab](https://sky.cs.berkeley.edu) at UC Berkeley, vLLM has evolved into a community-driven project with contributions from both academia and industry.
|
||||
Originally developed in the [Sky Computing Lab](https://sky.cs.berkeley.edu) at UC Berkeley, vLLM has grown into one of the most active open-source AI projects built and maintained by a diverse community of many dozens of academic institutions and companies from over 2000 contributors.
|
||||
|
||||
Where to get started with vLLM depends on the type of user. If you are looking to:
|
||||
|
||||
@@ -42,23 +42,37 @@ vLLM is fast with:
|
||||
|
||||
- State-of-the-art serving throughput
|
||||
- Efficient management of attention key and value memory with [**PagedAttention**](https://blog.vllm.ai/2023/06/20/vllm.html)
|
||||
- Continuous batching of incoming requests
|
||||
- Fast model execution with CUDA/HIP graph
|
||||
- Quantization: [GPTQ](https://arxiv.org/abs/2210.17323), [AWQ](https://arxiv.org/abs/2306.00978), INT4, INT8, and FP8
|
||||
- Optimized CUDA kernels, including integration with FlashAttention and FlashInfer.
|
||||
- Speculative decoding
|
||||
- Chunked prefill
|
||||
- Continuous batching of incoming requests, chunked prefill, prefix caching
|
||||
- Fast and flexible model execution with piecewise and full CUDA/HIP graphs
|
||||
- Quantization: FP8, MXFP8/MXFP4, NVFP4, INT8, INT4, GPTQ/AWQ, GGUF, compressed-tensors, ModelOpt, TorchAO, and [more](https://docs.vllm.ai/en/latest/features/quantization/index.html)
|
||||
- Optimized attention kernels including FlashAttention, FlashInfer, TRTLLM-GEN, FlashMLA, and Triton
|
||||
- Optimized GEMM/MoE kernels for various precisions using CUTLASS, TRTLLM-GEN, CuTeDSL
|
||||
- Speculative decoding including n-gram, suffix, EAGLE, DFlash
|
||||
- Automatic kernel generation and graph-level transformations using torch.compile
|
||||
- Disaggregated prefill, decode, and encode
|
||||
|
||||
vLLM is flexible and easy to use with:
|
||||
|
||||
- Seamless integration with popular HuggingFace models
|
||||
- Seamless integration with popular Hugging Face models
|
||||
- High-throughput serving with various decoding algorithms, including *parallel sampling*, *beam search*, and more
|
||||
- Tensor, pipeline, data and expert parallelism support for distributed inference
|
||||
- Tensor, pipeline, data, expert, and context parallelism for distributed inference
|
||||
- Streaming outputs
|
||||
- OpenAI-compatible API server
|
||||
- Support for NVIDIA GPUs, AMD CPUs and GPUs, Intel CPUs and GPUs, PowerPC CPUs, Arm CPUs, and TPU. Additionally, support for diverse hardware plugins such as Intel Gaudi, IBM Spyre and Huawei Ascend.
|
||||
- Prefix caching support
|
||||
- Multi-LoRA support
|
||||
- Generation of structured outputs using xgrammar or guidance
|
||||
- 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.
|
||||
|
||||
vLLM seamlessly supports 200+ model architectures on HuggingFace, including:
|
||||
|
||||
- Decoder-only LLMs (e.g., Llama, Qwen, Gemma)
|
||||
- Mixture-of-Expert LLMs (e.g., Mixtral, DeepSeek-V3, Qwen-MoE, GPT-OSS)
|
||||
- Hybrid attention and state-space models (e.g., Mamba, Qwen3.5)
|
||||
- Multi-modal models (e.g., LLaVA, Qwen-VL, Pixtral)
|
||||
- Embedding and retrieval models (e.g., E5-Mistral, GTE, ColBERT)
|
||||
- Reward and classification models (e.g., Qwen-Math)
|
||||
|
||||
Find the full list of supported models [here](./models/supported_models.md).
|
||||
|
||||
For more information, check out the following:
|
||||
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 325 KiB After Width: | Height: | Size: 315 KiB |
@@ -140,6 +140,80 @@ Data parallelism replicates the entire model across multiple GPU sets and proces
|
||||
Data parallelism can be combined with the other parallelism strategies and is set by `data_parallel_size=N`.
|
||||
Note that MoE layers will be sharded according to the product of the tensor parallel size and data parallel size.
|
||||
|
||||
### NUMA Binding for Multi-Socket GPU Nodes
|
||||
|
||||
On multi-socket GPU servers, GPU worker processes can lose performance if their
|
||||
CPU execution and memory allocation drift away from the NUMA node nearest to the
|
||||
GPU. vLLM can pin each worker with `numactl` before the Python subprocess starts,
|
||||
so the interpreter, imports, and early allocator state are created with the
|
||||
desired NUMA policy from the beginning.
|
||||
|
||||
Use `--numa-bind` to enable the feature. By default, vLLM auto-detects the
|
||||
GPU-to-NUMA mapping and uses `--cpunodebind=<node> --membind=<node>` for each
|
||||
worker. When you need a custom CPU policy, add `--numa-bind-cpus` and vLLM will
|
||||
switch to `--physcpubind=<cpu-list> --membind=<node>`.
|
||||
|
||||
These `--numa-bind*` options only apply to GPU execution processes. They do not
|
||||
configure the CPU backend's separate thread-affinity controls. Automatic
|
||||
GPU-to-NUMA detection is currently implemented for CUDA/NVML-based platforms;
|
||||
other GPU backends must provide explicit binding lists if they use these
|
||||
options.
|
||||
|
||||
`--numa-bind-nodes` takes one non-negative NUMA node index per visible GPU, in
|
||||
the same order as the GPU indices.
|
||||
`--numa-bind-cpus` takes one `numactl` CPU list per visible GPU, in the same
|
||||
order as the GPU indices. Each CPU list must use
|
||||
`numactl --physcpubind` syntax such as `0-3`, `0,2,4-7`, or `16-31,48-63`.
|
||||
|
||||
```bash
|
||||
# Auto-detect NUMA nodes for visible GPUs
|
||||
vllm serve meta-llama/Llama-3.1-8B-Instruct \
|
||||
--tensor-parallel-size 4 \
|
||||
--numa-bind
|
||||
|
||||
# Explicit NUMA-node mapping
|
||||
vllm serve meta-llama/Llama-3.1-8B-Instruct \
|
||||
--tensor-parallel-size 4 \
|
||||
--numa-bind \
|
||||
--numa-bind-nodes 0 0 1 1
|
||||
|
||||
# Explicit CPU pinning, useful for PCT or other high-frequency core layouts
|
||||
vllm serve meta-llama/Llama-3.1-8B-Instruct \
|
||||
--tensor-parallel-size 4 \
|
||||
--numa-bind \
|
||||
--numa-bind-nodes 0 0 1 1 \
|
||||
--numa-bind-cpus 0-3 4-7 48-51 52-55
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- CLI usage forces multiprocessing to use the `spawn` method automatically. If you enable NUMA binding through the Python API, also set `VLLM_WORKER_MULTIPROC_METHOD=spawn`.
|
||||
- Automatic detection relies on NVML and NUMA support from the host. If it cannot determine the mapping reliably, pass `--numa-bind-nodes` explicitly.
|
||||
- Explicit `--numa-bind-nodes` and `--numa-bind-cpus` values must be valid `numactl` inputs. vLLM does a small amount of validation, but the effective binding semantics are still determined by `numactl`.
|
||||
- The current implementation binds GPU execution processes such as `EngineCore` and multiprocessing workers. It does not apply NUMA binding to frontend API server processes or the DP coordinator.
|
||||
- In containerized environments, NUMA policy syscalls may require extra permissions, such as `--cap-add SYS_NICE` when running via `docker run`.
|
||||
|
||||
### CPU Backend Thread Affinity
|
||||
|
||||
The CPU backend uses a different mechanism from `--numa-bind`. CPU execution is
|
||||
configured through CPU-specific environment variables such as
|
||||
`VLLM_CPU_OMP_THREADS_BIND`, `VLLM_CPU_NUM_OF_RESERVED_CPU`, and
|
||||
`CPU_VISIBLE_MEMORY_NODES`, rather than the GPU-oriented `--numa-bind*` CLI
|
||||
options.
|
||||
|
||||
By default, `VLLM_CPU_OMP_THREADS_BIND=auto` derives OpenMP placement from the
|
||||
available CPU and NUMA topology for each CPU worker. To override the automatic
|
||||
policy, set `VLLM_CPU_OMP_THREADS_BIND` explicitly using the CPU list format
|
||||
documented for the CPU backend, or use `nobind` to disable this behavior.
|
||||
|
||||
For the current CPU backend setup and tuning guidance, see:
|
||||
|
||||
- [Related runtime environment variables](../getting_started/installation/cpu.md#related-runtime-environment-variables)
|
||||
- [How to decide `VLLM_CPU_OMP_THREADS_BIND`](../getting_started/installation/cpu.md#how-to-decide-vllm_cpu_omp_threads_bind)
|
||||
|
||||
The GPU-only `--numa-bind`, `--numa-bind-nodes`, and `--numa-bind-cpus` options
|
||||
do not configure CPU worker affinity.
|
||||
|
||||
### Batch-level DP for Multi-Modal Encoders
|
||||
|
||||
By default, TP is used to shard the weights of multi-modal encoders just like for language decoders,
|
||||
|
||||
@@ -49,10 +49,10 @@ If you are developing vLLM's Python and CUDA/C++ code, install Pytorch first:
|
||||
uv pip install torch torchvision torchaudio --extra-index-url https://download.pytorch.org/whl/cu129
|
||||
```
|
||||
|
||||
Then install the necessary build dependencies from `requirements/build.txt`, skipping `torch` as it was installed in the previous step:
|
||||
Then install the necessary build dependencies from `requirements/build/cuda.txt`, skipping `torch` as it was installed in the previous step:
|
||||
|
||||
```bash
|
||||
grep -v '^torch==' requirements/build.txt | uv pip install -r -
|
||||
grep -v '^torch==' requirements/build/cuda.txt | uv pip install -r -
|
||||
```
|
||||
|
||||
Finally install vLLM using:
|
||||
|
||||
@@ -16,10 +16,10 @@ Before setting up the incremental build:
|
||||
|
||||
2. **CUDA Toolkit:** Verify that the NVIDIA CUDA Toolkit is correctly installed and `nvcc` is accessible in your `PATH`. CMake relies on `nvcc` to compile CUDA code. You can typically find `nvcc` in `$CUDA_HOME/bin/nvcc` or by running `which nvcc`. If you encounter issues, refer to the [official CUDA Toolkit installation guides](https://developer.nvidia.com/cuda-toolkit-archive) and vLLM's main [GPU installation documentation](../getting_started/installation/gpu.md#troubleshooting) for troubleshooting. The `CMAKE_CUDA_COMPILER` variable in your `CMakeUserPresets.json` should also point to your `nvcc` binary.
|
||||
|
||||
3. **Build Tools:** It is highly recommended to install `ccache` for fast rebuilds by caching compilation results (e.g., `sudo apt install ccache` or `conda install ccache`). Also, ensure the core build dependencies like `cmake` and `ninja` are installed. These are installable through `requirements/build.txt` or your system's package manager.
|
||||
3. **Build Tools:** It is highly recommended to install `ccache` for fast rebuilds by caching compilation results (e.g., `sudo apt install ccache` or `conda install ccache`). Also, ensure the core build dependencies like `cmake` and `ninja` are installed. These are installable through `requirements/build/cuda.txt` or your system's package manager.
|
||||
|
||||
```console
|
||||
uv pip install -r requirements/build.txt --torch-backend=auto
|
||||
uv pip install -r requirements/build/cuda.txt --torch-backend=auto
|
||||
```
|
||||
|
||||
## Setting up the CMake Build Environment
|
||||
|
||||
@@ -57,9 +57,9 @@ Modular kernels are supported by the following `FusedMoEMethodBase` classes.
|
||||
|
||||
- [`ModelOptFp8MoEMethod`][vllm.model_executor.layers.quantization.modelopt.ModelOptFp8MoEMethod]
|
||||
- [`Fp8MoEMethod`][vllm.model_executor.layers.quantization.fp8.Fp8MoEMethod]
|
||||
- [`CompressedTensorsW4A4Nvfp4MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.CompressedTensorsW4A4Nvfp4MoEMethod]
|
||||
- [`CompressedTensorsW8A8Fp8MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.CompressedTensorsW8A8Fp8MoEMethod]
|
||||
- [`Mxfp4MoEMethod`][vllm.model_executor.layers.quantization.mxfp4.Mxfp4MoEMethod]
|
||||
- [`CompressedTensorsW4A4Nvfp4MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.compressed_tensors_moe_w4a4_nvfp4.CompressedTensorsW4A4Nvfp4MoEMethod]
|
||||
- [`CompressedTensorsW8A8Fp8MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.compressed_tensors_moe_w8a8_fp8.CompressedTensorsW8A8Fp8MoEMethod]
|
||||
- [`GptOssMxfp4MoEMethod`][vllm.model_executor.layers.quantization.mxfp4.GptOssMxfp4MoEMethod]
|
||||
- [`UnquantizedFusedMoEMethod`][vllm.model_executor.layers.fused_moe.layer.UnquantizedFusedMoEMethod]
|
||||
|
||||
## Fused Experts Kernels
|
||||
@@ -82,7 +82,7 @@ To be used with a particular `FusedMoEPrepareAndFinalizeModular` subclass, MoE k
|
||||
| ------ | ----------------- | ------------ | ------------- | ------------------- | --------------------- | ------- | ------ |
|
||||
| triton | standard | all<sup>1</sup> | G,A,T | silu, gelu,</br>swigluoai,</br>silu_no_mul,</br>gelu_no_mul | Y | Y | [`fused_experts`][vllm.model_executor.layers.fused_moe.fused_moe.fused_experts],</br>[`TritonExperts`][vllm.model_executor.layers.fused_moe.fused_moe.TritonExperts] |
|
||||
| triton (batched) | batched | all<sup>1</sup> | G,A,T | silu, gelu | <sup>6</sup> | Y | [`BatchedTritonExperts`][vllm.model_executor.layers.fused_moe.fused_batched_moe.BatchedTritonExperts] |
|
||||
| deep gemm | standard,</br>batched | fp8 | G(128),A,T | silu, gelu | <sup>6</sup> | Y | </br>[`DeepGemmExperts`][vllm.model_executor.layers.fused_moe.deep_gemm_moe.DeepGemmExperts],</br>[`BatchedDeepGemmExperts`][vllm.model_executor.layers.fused_moe.batched_deep_gemm_moe.BatchedDeepGemmExperts] |
|
||||
| deep gemm | standard,</br>batched | fp8 | G(128),A,T | silu, gelu | <sup>6</sup> | Y | </br>[`DeepGemmExperts`][vllm.model_executor.layers.fused_moe.experts.deep_gemm_moe.DeepGemmExperts],</br>[`BatchedDeepGemmExperts`][vllm.model_executor.layers.fused_moe.experts.batched_deep_gemm_moe.BatchedDeepGemmExperts] |
|
||||
| cutlass_fp4 | standard,</br>batched | nvfp4 | A,T | silu | Y | Y | [`CutlassExpertsFp4`][vllm.model_executor.layers.fused_moe.cutlass_moe.CutlassExpertsFp4] |
|
||||
| cutlass_fp8 | standard,</br>batched | fp8 | A,T | silu, gelu | Y | Y | [`CutlassExpertsFp8`][vllm.model_executor.layers.fused_moe.cutlass_moe.CutlassExpertsFp8],</br>[`CutlasBatchedExpertsFp8`][vllm.model_executor.layers.fused_moe.cutlass_moe.CutlassBatchedExpertsFp8] |
|
||||
| flashinfer | standard | nvfp4,</br>fp8 | T | <sup>5</sup> | N | Y | [`FlashInferExperts`][vllm.model_executor.layers.fused_moe.flashinfer_cutlass_moe.FlashInferExperts] |
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user