Compare commits

...
Author SHA1 Message Date
Chinmay Kulkarni 4fda385093 [ZenCPU] Update device to "zen5"
Signed-off-by: Chinmay Kulkarni <Chinmay.Kulkarni@amd.com>
2026-05-05 06:03:42 -06:00
Chinmay Kulkarni ee5052de02 [ZenCPU] Added requirements files as dependencies for amd-zen-cpu-inference tests
Signed-off-by: Chinmay Kulkarni <Chinmay.Kulkarni@amd.com>
2026-05-05 06:03:33 -06:00
Chinmay Kulkarni ad387d78ca [ZenCPU] Add more source file dependencies to tests
Signed-off-by: Chinmay Kulkarni <Chinmay.Kulkarni@amd.com>
2026-05-05 06:03:32 -06:00
Chinmay Kulkarni 200cbdd308 [ZenCPU] Changes with respect to docker build and relevant cpu tests
Signed-off-by: Chinmay Kulkarni <Chinmay.Kulkarni@amd.com>
2026-05-05 06:03:31 -06:00
98661fe012 [Bugfix][KVConnector] Support DCP/PCP in OffloadingConnector (#41549)
Signed-off-by: Itay Etelis <itay.etelis@ibm.com>
Co-authored-by: Itay Etelis <itay.etelis@ibm.com>
2026-05-05 14:54:29 +03:00
Harry MellorandGitHub b0765bee17 Fix DeepSeek-OCR for Transformers v4 (#41460)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-05-05 11:11:21 +00:00
bairongzGitHubzhuangbaironggemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
0a201b60cf [Model] support Qianfan-OCR model (#40136)
Signed-off-by: bairongz <baiyuu.cs@gmail.com>
Signed-off-by: zhuangbairong <zhuangbairong@baidu.com>
Co-authored-by: zhuangbairong <zhuangbairong@baidu.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-05-05 10:51:25 +00:00
8b9ea2f881 [Feature] Add Triton kernel JIT compilation monitor for inference (#40137)
Signed-off-by: Artem Perevedentsev <aperevedents@nvidia.com>
Co-authored-by: Jiangyun Zhu <riverclouds.zhu@qq.com>
2026-05-05 14:08:57 +04:00
Kunshang JiandGitHub 2ceea42958 [XPU] use xpu topk topp sample kernel (#39285)
Signed-off-by: Kunshang Ji <kunshang.ji@intel.com>
2026-05-05 18:05:17 +08:00
bee126165f [P/D][Mooncake] Add KVConnectorStats for transfer observability (#40414)
Signed-off-by: Zhewen Li <zhewenli@inferact.ai>
Co-authored-by: Zhewen Li <zhewenli@inferact.ai>
2026-05-05 02:17:38 -07:00
BitTobyandGitHub 27cc676be3 [Model] Use AutoWeightsLoader for Plamo2 (#41699)
Signed-off-by: bittoby <218712309+bittoby@users.noreply.github.com>
2026-05-05 08:56:24 +00:00
4845aee6b7 [Benchmark] Add --trust-remote-code flag to multi-turn benchmark (#41661)
Signed-off-by: Dao Le <daole@inferact.ai>
Signed-off-by: Dao Le <Dao007forever@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-05-05 01:00:37 -07:00
BitTobyandGitHub 0c620d2e08 [Model] Use AutoWeightsLoader for CohereMoe (#41690)
Signed-off-by: bittoby <218712309+bittoby@users.noreply.github.com>
2026-05-05 04:44:15 +00:00
6bb924bbf3 [Model] Fix Gemma4 MoE activation mismatch (#41574)
Signed-off-by: Luciano Martins <lucianommartins@users.noreply.github.com>
Co-authored-by: Luciano Martins <lucianommartins@users.noreply.github.com>
2026-05-05 04:34:11 +00:00
czhu-cohereandGitHub eaec7be446 [BugFix] Preserve max_seq_len in ubatch metadata during CUDA graph capture (#40961)
Signed-off-by: root <conway.zhu@cohere.com>
Signed-off-by: <conway.zhu@cohere.com>
2026-05-05 04:29:34 +00:00
Jeffrey WangandGitHub f04fd1677b [Ray] Enable RayExecutorV2 by default (#41421)
Signed-off-by: Jeffrey Wang <jeffreywang@anyscale.com>
2026-05-05 04:27:34 +00:00
420b0a5c95 [Hardware][Power]Add Power VSX Attention Backend and fix l2 Cache Crash (#40451)
Signed-off-by: Akash Kaothalkar <akashkaothalkar@akashs-mbp.bl1-in.ibm.com>
Signed-off-by: Akash Kaothalkar <akash.kaothalkar@ibm.com>
Signed-off-by: Akash kaothalkar <akash.kaothalkar@ibm.com>
Co-authored-by: Akash Kaothalkar <akashkaothalkar@akashs-mbp.bl1-in.ibm.com>
Co-authored-by: Akash Kaothalkar <akash.kaothalkar@ibm.com>
Co-authored-by: Li, Jiang <jiang1.li@intel.com>
2026-05-04 20:51:09 -07:00
Bowen BaoandGitHub 1e9500410a [ROCm][Quantization][2/N] Refactor quark_moe w4a8 w/ oracle (#39136)
Signed-off-by: Bowen Bao <bowenbao@amd.com>
2026-05-04 19:50:38 -07:00
Nick HillandGitHub 416f9cdede [Perf][2/n] Eliminate GPU<->CPU syncs in pooling code (#41433)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-05-05 02:43:25 +00:00
685bf811d6 [XPU] enable is_act_and_mul for xpu (#37481)
Signed-off-by: Chendi Xue <chendi.xue@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-05-05 01:07:39 +00:00
Giancarlo DelfinandGitHub e1e4646b06 [Model Runner V2] Rebuild attn metadata between draft decode steps (#41162)
Signed-off-by: Giancarlo Delfin <gdelfin@inferact.ai>
2026-05-05 00:44:55 +00:00
4f2af1a7c0 [Feature] TurboQuant: support hybrid models and uniform quantization (#39931)
Signed-off-by: JartX <sagformas@epdcenter.es>
Signed-off-by: Jim Smith <jhsmith0@me.com>
Co-authored-by: Jim Smith <jhsmith0@me.com>
Co-authored-by: Sandermage <sandermage@users.noreply.github.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-05-04 20:14:01 -04:00
Wentao YeandGitHub 577b9623e6 [Bug] Fix status update address for non-MOE model within external dp mode (#40839)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-05-04 16:37:16 -07:00
Andreas KaratzasandGitHub 1cb0838721 [ROCm][CI] Fix MLA prefill scale for DeepSeek GSM8K (#41569)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-05-04 16:32:55 -07:00
71 changed files with 3483 additions and 715 deletions
+94
View File
@@ -22,6 +22,29 @@ steps:
pytest -x -v -s tests/kernels/test_onednn.py
pytest -x -v -s tests/kernels/test_awq_int4_to_int8.py"
- label: AMD-CPU-Kernel Tests
depends_on: []
soft_fail: false
device: zen5
no_plugin: true
source_file_dependencies:
- setup.py
- vllm/docker/Dockerfile.cpu
- vllm/requirements/cpu.txt
- vllm/requirements/build/cpu.txt
- csrc/cpu/
- cmake/cpu_extension.cmake
- CMakeLists.txt
- vllm/model_executor/layers/utils.py
- vllm/platforms/cpu.py
- vllm/platforms/zen_cpu.py
- vllm/platforms/__init__.py
- tests/model_executor/test_cpu_unquantized_gemm_dispatch.py
commands:
- |
bash .buildkite/scripts/hardware_ci/run-amd-cpu-test.sh 20m "
pytest -x -v -s tests/model_executor/test_cpu_unquantized_gemm_dispatch.py"
- label: CPU-Compatibility Tests
depends_on: []
device: intel_cpu
@@ -35,6 +58,26 @@ steps:
bash .buildkite/scripts/hardware_ci/run-cpu-test.sh 20m "
bash .buildkite/scripts/hardware_ci/run-cpu-compatibility-test.sh"
- label: AMD-CPU-Compatibility Tests
depends_on: []
soft_fail: false
device: zen5
no_plugin: true
source_file_dependencies:
- setup.py
- vllm/docker/Dockerfile.cpu
- vllm/requirements/cpu.txt
- vllm/requirements/build/cpu.txt
- vllm/platforms/cpu.py
- vllm/platforms/zen_cpu.py
- vllm/platforms/interface.py
- vllm/platforms/__init__.py
- tests/test_zen_cpu_platform_detection.py
commands:
- |
bash .buildkite/scripts/hardware_ci/run-amd-cpu-test.sh 20m "
pytest -x -v -s tests/test_zen_cpu_platform_detection.py"
- label: CPU-Language Generation and Pooling Model Tests
depends_on: []
device: intel_cpu
@@ -50,6 +93,32 @@ steps:
pytest -x -v -s tests/models/language/generation -m cpu_model
pytest -x -v -s tests/models/language/pooling -m cpu_model"
- label: AMD-CPU-Language Generation and Pooling Model Tests
depends_on: []
soft_fail: false
device: zen5
no_plugin: true
source_file_dependencies:
- setup.py
- vllm/docker/Dockerfile.cpu
- vllm/requirements/cpu.txt
- vllm/requirements/build/cpu.txt
- csrc/cpu/
- vllm/model_executor/layers/utils.py
- vllm/platforms/zen_cpu.py
- vllm/platforms/__init__.py
- setup.py
- vllm/platforms/cpu.py
- vllm/platforms/interface.py
- vllm/v1/worker/cpu_model_runner.py
- tests/models/language/generation/
- tests/models/language/pooling/
commands:
- |
bash .buildkite/scripts/hardware_ci/run-amd-cpu-test.sh 30m "
pytest -x -v -s tests/models/language/generation -m cpu_model
pytest -x -v -s tests/models/language/pooling -m cpu_model"
- label: CPU-Quantization Model Tests
depends_on: []
device: intel_cpu
@@ -98,6 +167,31 @@ steps:
bash .buildkite/scripts/hardware_ci/run-cpu-test.sh 10m "
bash .buildkite/scripts/hardware_ci/run-cpu-distributed-smoke-test.sh dp_tp"
- label: AMD-CPU-Distributed Tests
depends_on: []
soft_fail: false
device: zen5
no_plugin: true
source_file_dependencies:
- setup.py
- vllm/docker/Dockerfile.cpu
- vllm/requirements/cpu.txt
- vllm/requirements/build/cpu.txt
- csrc/cpu/shm.cpp
- vllm/v1/worker/cpu_worker.py
- vllm/v1/worker/gpu_worker.py
- vllm/v1/worker/cpu_model_runner.py
- vllm/v1/worker/gpu_model_runner.py
- vllm/platforms/cpu.py
- vllm/platforms/zen_cpu.py
- vllm/platforms/cpu.py
- vllm/distributed/parallel_state.py
- vllm/distributed/device_communicators/cpu_communicator.py
commands:
- |
bash .buildkite/scripts/hardware_ci/run-amd-cpu-test.sh 10m "
bash .buildkite/scripts/hardware_ci/run-cpu-distributed-smoke-test.sh"
- label: CPU-Multi-Modal Model Tests %N
depends_on: []
device: intel_cpu
+15
View File
@@ -27,6 +27,21 @@ steps:
- exit_status: -10 # Agent was lost
limit: 2
- label: ":docker: Build AMD CPU image"
soft_fail: true
key: image-build-amd-cpu
depends_on: []
commands:
- .buildkite/image_build/image_build_amd_cpu.sh $REGISTRY $REPO $BUILDKITE_COMMIT
env:
DOCKER_BUILDKIT: "1"
retry:
automatic:
- exit_status: -1 # Agent was lost
limit: 2
- exit_status: -10 # Agent was lost
limit: 2
- label: ":docker: Build HPU image"
soft_fail: true
depends_on: []
@@ -0,0 +1,34 @@
#!/bin/bash
set -e
if [[ $# -lt 3 ]]; then
echo "Usage: $0 <registry> <repo> <commit>"
exit 1
fi
REGISTRY=$1
REPO=$2
BUILDKITE_COMMIT=$3
# authenticate with AWS ECR
aws ecr-public get-login-password --region us-east-1 | docker login --username AWS --password-stdin "$REGISTRY"
# skip build if image already exists
if [[ -z $(docker manifest inspect "$REGISTRY"/"$REPO":"$BUILDKITE_COMMIT"-amd-cpu) ]]; then
echo "Image not found, proceeding with build..."
else
echo "Image found"
exit 0
fi
# build
docker build --file docker/Dockerfile.cpu \
--build-arg max_jobs=16 \
--build-arg buildkite_commit="$BUILDKITE_COMMIT" \
--build-arg VLLM_CPU_X86=true \
--tag "$REGISTRY"/"$REPO":"$BUILDKITE_COMMIT"-amd-cpu \
--target vllm-zen-test \
--progress plain .
# push
docker push "$REGISTRY"/"$REPO":"$BUILDKITE_COMMIT"-amd-cpu
@@ -0,0 +1,20 @@
#!/bin/bash
# This script build the CPU docker image and run the offline inference inside the container.
# It serves a sanity check for compilation and basic model usage.
set -euox pipefail
# allow to bind to different cores
CORE_RANGE=${CORE_RANGE:-48-95}
NUMA_NODE=${NUMA_NODE:-1}
IMAGE_NAME="amd-cpu-test-$NUMA_NODE"
TIMEOUT_VAL=$1
TEST_COMMAND=$2
# building the docker image
echo "--- :docker: Building Docker image"
docker build --progress plain --tag "$IMAGE_NAME" --target vllm-zen-test -f docker/Dockerfile.cpu .
# Run the image, setting --shm-size=4g for tensor parallel.
docker run --rm --cpuset-cpus="$CORE_RANGE" --cpuset-mems="$NUMA_NODE" -v ~/.cache/huggingface:/root/.cache/huggingface --privileged=true -e HF_TOKEN -e VLLM_CPU_KVCACHE_SPACE=16 -e VLLM_CPU_CI_ENV=1 -e VLLM_CPU_SIM_MULTI_NUMA=1 --shm-size=4g "$IMAGE_NAME" \
timeout "$TIMEOUT_VAL" bash -c "set -euox pipefail; echo \"--- Print packages\"; pip list; echo \"--- Running tests\"; ${TEST_COMMAND}"
+2 -1
View File
@@ -13,8 +13,9 @@ steps:
- tests/test_config
- tests/test_logger
- tests/test_vllm_port
- tests/test_jit_monitor.py
commands:
- pytest -v -s engine test_sequence.py test_config.py test_logger.py test_vllm_port.py
- pytest -v -s engine test_sequence.py test_config.py test_logger.py test_vllm_port.py test_jit_monitor.py
- label: Engine (1 GPU)
key: engine-1-gpu
@@ -1473,6 +1473,12 @@ async def main() -> None:
"(for example: --warmup-percentages=0%%,50%%)",
)
parser.add_argument(
"--trust-remote-code",
action="store_true",
help="Trust remote code when loading the tokenizer.",
)
args = parser.parse_args()
logger.info(args)
@@ -1515,7 +1521,9 @@ async def main() -> None:
np.random.seed(args.seed)
logger.info("Loading tokenizer")
tokenizer = AutoTokenizer.from_pretrained(args.model)
tokenizer = AutoTokenizer.from_pretrained(
args.model, trust_remote_code=args.trust_remote_code
)
await get_server_info(args.url)
+4
View File
@@ -29,6 +29,8 @@ torch::Tensor get_scheduler_metadata(
isa = cpu_attention::ISA::NEON;
} else if (isa_hint == "vxe") {
isa = cpu_attention::ISA::VXE;
} else if (isa_hint == "vsx") {
isa = cpu_attention::ISA::VSX;
} else {
TORCH_CHECK(false, "Unsupported CPU attention ISA hint: " + isa_hint);
}
@@ -129,6 +131,8 @@ void cpu_attn_reshape_and_cache(
return cpu_attention::ISA::NEON;
} else if (isa == "vxe") {
return cpu_attention::ISA::VXE;
} else if (isa == "vsx") {
return cpu_attention::ISA::VSX;
} else {
TORCH_CHECK(false, "Invalid ISA type: " + isa);
}
+4 -1
View File
@@ -12,7 +12,7 @@
#include "cpu/utils.hpp"
namespace cpu_attention {
enum class ISA { AMX, VEC, VEC16, NEON, VXE };
enum class ISA { AMX, VEC, VEC16, NEON, VXE, VSX };
// Mirrors csrc/attention/dtype_fp8.cuh Fp8KVCacheDataType exactly.
enum class Fp8KVCacheDataType {
@@ -164,6 +164,9 @@ struct AttentionMetadata {
case ISA::VXE:
ss << "VXE, ";
break;
case ISA::VSX:
ss << "VSX, ";
break;
}
ss << "workitem_group_num: " << workitem_group_num
<< ", reduction_item_num: " << reduction_item_num
+2 -2
View File
@@ -27,8 +27,8 @@ FORCE_INLINE std::pair<vec_op::FP32Vec16, vec_op::FP32Vec16> load_b_pair_vec(
return {vec_op::FP32Vec16(bf16_b_reg, 0), vec_op::FP32Vec16(bf16_b_reg, 1)};
} else {
using load_vec_t = typename VecTypeTrait<kv_cache_t>::vec_t;
return {vec_op::FP32Vec16(load_vec_t(ptr)),
vec_op::FP32Vec16(load_vec_t(ptr + 16))};
return std::make_pair(vec_op::FP32Vec16(load_vec_t(ptr)),
vec_op::FP32Vec16(load_vec_t(ptr + 16)));
}
}
+359
View File
@@ -0,0 +1,359 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#ifndef CPU_ATTN_VSX_HPP
#define CPU_ATTN_VSX_HPP
#include "cpu_attn_impl.hpp"
#include <altivec.h>
#include <type_traits>
namespace cpu_attention {
namespace {
// ppc64le Vector = 16 bytes (128 bits)
#define BLOCK_SIZE_ALIGNMENT 32
#define HEAD_SIZE_ALIGNMENT 32
#define MAX_Q_HEAD_NUM_PER_ITER 16
template <typename kv_cache_t>
FORCE_INLINE void load_row8_B_as_f32(const kv_cache_t* p, __vector float& b0,
__vector float& b1);
// [1] Float Specialization
template <>
FORCE_INLINE void load_row8_B_as_f32<float>(const float* p, __vector float& b0,
__vector float& b1) {
b0 = vec_xl(0, const_cast<float*>(p));
b1 = vec_xl(0, const_cast<float*>(p + 4));
}
// [2] BFloat16 Specialization (Little Endian ppc64le)
// On ppc64le (LE): BF16 bits should land in the HIGH 16 bits of each float32.
// Byte layout of float32 on LE: [byte0(LSB), byte1, byte2, byte3(MSB)]
// We need BF16 in bytes2-3 (high half) with bytes0-1 zeroed.
// vec_mergeh on LE interleaves elements 0..3: result_i = {a[i], b[i]}
// So vec_mergeh(zeros_u16, raw_u16) gives for each uint16 pair:
// uint16[2i] = zeros[i] -> low 16 bits of uint32 -> zeroed mantissa LSBs
// uint16[2i+1] = raw[i] -> high 16 bits of uint32 -> BF16 bits
// Cast to float32 gives exactly (bf16_bits << 16) per element.
template <>
FORCE_INLINE void load_row8_B_as_f32<c10::BFloat16>(const c10::BFloat16* p,
__vector float& b0,
__vector float& b1) {
__vector unsigned short raw = vec_xl(
0, reinterpret_cast<unsigned short*>(const_cast<c10::BFloat16*>(p)));
__vector unsigned short zeros = vec_splat_u16(0);
// LE: zeros in low 16 bits, raw in high 16 bits → bf16 << 16 == float32
b0 = (__vector float)vec_mergeh(zeros, raw);
b1 = (__vector float)vec_mergel(zeros, raw);
}
// Note: c10::Half (FP16) is not supported on PowerPC architecture
template <int32_t M, typename kv_cache_t>
FORCE_INLINE void gemm_micro_ppc64le_Mx8_Ku4(
const float* __restrict A, // [M x K]
const kv_cache_t* __restrict B, // [K x 8]
float* __restrict C, // [M x 8]
int64_t lda, int64_t ldb, int64_t ldc, int32_t K, bool accumulate) {
static_assert(1 <= M && M <= 8, "M must be in [1,8]");
#define ROWS_APPLY(OP) OP(0) OP(1) OP(2) OP(3) OP(4) OP(5) OP(6) OP(7)
#define IF_M(i) if constexpr (M > (i))
// 1. Define A pointers
#define DECL_A(i) const float* a##i = A + (i) * lda;
ROWS_APPLY(DECL_A)
#undef DECL_A
// 2. Define Accumulators (2 vectors covers 8 columns)
#define DECL_ACC(i) __vector float acc##i##_0, acc##i##_1;
ROWS_APPLY(DECL_ACC)
#undef DECL_ACC
// 3. Initialize Accumulators (Load C or Zero)
#define INIT_ACC(i) \
IF_M(i) { \
if (accumulate) { \
acc##i##_0 = vec_xl(0, const_cast<float*>(C + (i) * ldc + 0)); \
acc##i##_1 = vec_xl(0, const_cast<float*>(C + (i) * ldc + 4)); \
} else { \
acc##i##_0 = vec_splats(0.0f); \
acc##i##_1 = vec_splats(0.0f); \
} \
}
ROWS_APPLY(INIT_ACC)
#undef INIT_ACC
int32_t k = 0;
for (; k + 3 < K; k += 4) {
// Load 4 values of A for each Row M: A[k...k+3]
#define LOAD_A4(i) \
__vector float a##i##v; \
IF_M(i) a##i##v = vec_xl(0, const_cast<float*>(a##i + k));
ROWS_APPLY(LOAD_A4)
#undef LOAD_A4
// FMA for specific lane L of A
// ppc64le: vec_madd(b, vec_splat(a, lane), acc)
#define FMAS_LANE(i, aiv, L) \
IF_M(i) { \
__vector float a_broad = vec_splat(aiv, L); \
acc##i##_0 = vec_madd(b0, a_broad, acc##i##_0); \
acc##i##_1 = vec_madd(b1, a_broad, acc##i##_1); \
}
// Unroll K=0..3
{
__vector float b0, b1;
load_row8_B_as_f32<kv_cache_t>(B + (int64_t)(k + 0) * ldb, b0, b1);
#define STEP_K0(i) FMAS_LANE(i, a##i##v, 0)
ROWS_APPLY(STEP_K0)
#undef STEP_K0
}
{
__vector float b0, b1;
load_row8_B_as_f32<kv_cache_t>(B + (int64_t)(k + 1) * ldb, b0, b1);
#define STEP_K1(i) FMAS_LANE(i, a##i##v, 1)
ROWS_APPLY(STEP_K1)
#undef STEP_K1
}
{
__vector float b0, b1;
load_row8_B_as_f32<kv_cache_t>(B + (int64_t)(k + 2) * ldb, b0, b1);
#define STEP_K2(i) FMAS_LANE(i, a##i##v, 2)
ROWS_APPLY(STEP_K2)
#undef STEP_K2
}
{
__vector float b0, b1;
load_row8_B_as_f32<kv_cache_t>(B + (int64_t)(k + 3) * ldb, b0, b1);
#define STEP_K3(i) FMAS_LANE(i, a##i##v, 3)
ROWS_APPLY(STEP_K3)
#undef STEP_K3
}
#undef FMAS_LANE
}
for (; k < K; ++k) {
__vector float b0, b1;
load_row8_B_as_f32<kv_cache_t>(B + (int64_t)k * ldb, b0, b1);
#define TAIL_ROW(i) \
IF_M(i) { \
__vector float ai = vec_splats(*(a##i + k)); \
acc##i##_0 = vec_madd(b0, ai, acc##i##_0); \
acc##i##_1 = vec_madd(b1, ai, acc##i##_1); \
}
ROWS_APPLY(TAIL_ROW)
#undef TAIL_ROW
}
#define STORE_ROW(i) \
IF_M(i) { \
vec_xst(acc##i##_0, 0, C + (i) * ldc + 0); \
vec_xst(acc##i##_1, 0, C + (i) * ldc + 4); \
}
ROWS_APPLY(STORE_ROW)
#undef STORE_ROW
#undef ROWS_APPLY
#undef IF_M
}
template <int32_t N, typename kv_cache_t>
FORCE_INLINE void gemm_macro_ppc64le_Mx8_Ku4(const float* __restrict A,
const kv_cache_t* __restrict B,
float* __restrict C, int32_t M,
int32_t K, int64_t lda,
int64_t ldb, int64_t ldc,
bool accumulate) {
static_assert(N % 8 == 0, "N must be a multiple of 8");
for (int32_t m = 0; m < M;) {
int32_t mb = (M - m >= 8) ? 8 : (M - m >= 4) ? 4 : (M - m >= 2) ? 2 : 1;
const float* Ab = A + m * lda;
float* Cb = C + m * ldc;
for (int32_t n = 0; n < N; n += 8) {
const kv_cache_t* Bn = B + n;
float* Cn = Cb + n;
switch (mb) {
case 8:
gemm_micro_ppc64le_Mx8_Ku4<8, kv_cache_t>(Ab, Bn, Cn, lda, ldb, ldc,
K, accumulate);
break;
case 4:
gemm_micro_ppc64le_Mx8_Ku4<4, kv_cache_t>(Ab, Bn, Cn, lda, ldb, ldc,
K, accumulate);
break;
case 2:
gemm_micro_ppc64le_Mx8_Ku4<2, kv_cache_t>(Ab, Bn, Cn, lda, ldb, ldc,
K, accumulate);
break;
default:
gemm_micro_ppc64le_Mx8_Ku4<1, kv_cache_t>(Ab, Bn, Cn, lda, ldb, ldc,
K, accumulate);
break;
}
}
m += mb;
}
}
template <typename kv_cache_t>
class TileGemmPPC64 {
public:
template <AttentionGemmPhase phase, int32_t k_size>
FORCE_INLINE static void gemm(const int32_t m_size,
float* __restrict__ a_tile,
kv_cache_t* __restrict__ b_tile,
float* __restrict__ c_tile, const int64_t lda,
const int64_t ldb, const int64_t ldc,
const int32_t block_size,
const int32_t dynamic_k_size,
const bool accum_c) {
if constexpr (phase == AttentionGemmPhase::QK) {
gemm_macro_ppc64le_Mx8_Ku4<BLOCK_SIZE_ALIGNMENT, kv_cache_t>(
a_tile, b_tile, c_tile, m_size, k_size, lda, ldb, ldc, accum_c);
} else {
gemm_macro_ppc64le_Mx8_Ku4<HEAD_SIZE_ALIGNMENT, kv_cache_t>(
a_tile, b_tile, c_tile, m_size, dynamic_k_size, lda, ldb, ldc,
accum_c);
}
}
};
} // namespace
template <typename scalar_t, int64_t head_dim>
class AttentionImpl<ISA::VSX, scalar_t, head_dim> {
public:
using query_t = scalar_t;
using q_buffer_t = float;
using kv_cache_t = scalar_t;
using logits_buffer_t = float;
using partial_output_buffer_t = float;
using prob_buffer_t = float;
constexpr static int64_t BlockSizeAlignment = BLOCK_SIZE_ALIGNMENT;
constexpr static int64_t HeadDimAlignment = HEAD_SIZE_ALIGNMENT;
constexpr static int64_t MaxQHeadNumPerIteration = MAX_Q_HEAD_NUM_PER_ITER;
constexpr static int64_t HeadDim = head_dim;
constexpr static ISA ISAType = ISA::VSX;
constexpr static bool scale_on_logits =
false; // Scale is applied to Q during copy
public:
AttentionImpl() {}
template <template <typename tile_gemm_t> typename attention>
FORCE_INLINE void execute_attention(DEFINE_CPU_ATTENTION_PARAMS) {
attention<TileGemmPPC64<kv_cache_t>> attention_iteration;
attention_iteration(CPU_ATTENTION_PARAMS);
}
// Strides for Memory Layout
constexpr static int64_t k_cache_token_group_stride(
const int32_t block_size) {
return BlockSizeAlignment; // [head_dim, block_size] layout
}
constexpr static int64_t v_cache_token_group_stride(
const int32_t block_size) {
return head_dim * BlockSizeAlignment;
}
constexpr static int64_t v_cache_head_group_stride(const int32_t block_size) {
return HeadDimAlignment;
}
static void copy_q_heads_tile(scalar_t* __restrict__ src,
float* __restrict__ q_buffer,
const int32_t q_num,
const int32_t q_heads_per_kv,
const int64_t q_num_stride,
const int64_t q_head_stride, float scale) {
__vector float scale_vec = vec_splats(scale);
constexpr bool is_bf16 = std::is_same<scalar_t, c10::BFloat16>::value;
for (int32_t i = 0; i < q_num; ++i) {
for (int32_t h = 0; h < q_heads_per_kv; ++h) {
scalar_t* curr_src = src + i * q_num_stride + h * q_head_stride;
float* curr_dst =
q_buffer + i * q_heads_per_kv * head_dim + h * head_dim;
int32_t d = 0;
for (; d <= head_dim - 8; d += 8) {
__vector float v0, v1;
load_row8_B_as_f32<scalar_t>(curr_src + d, v0, v1);
v0 = vec_mul(v0, scale_vec);
v1 = vec_mul(v1, scale_vec);
vec_xst(v0, 0, curr_dst + d);
vec_xst(v1, 0, curr_dst + d + 4);
}
for (; d < head_dim; ++d) {
float val = static_cast<float>(curr_src[d]);
curr_dst[d] = val * scale;
}
}
}
}
static void reshape_and_cache(
const scalar_t* __restrict__ key, const scalar_t* __restrict__ value,
scalar_t* __restrict__ key_cache, scalar_t* __restrict__ value_cache,
const int64_t* __restrict__ slot_mapping, const int64_t token_num,
const int64_t key_token_num_stride, const int64_t value_token_num_stride,
const int64_t head_num, const int64_t key_head_num_stride,
const int64_t value_head_num_stride, const int64_t num_blocks,
const int64_t num_blocks_stride, const int64_t cache_head_num_stride,
const int64_t block_size, const int64_t block_size_stride,
const float k_inv = 0.0f, const float v_inv = 0.0f) {
// k_inv and v_inv are unused on VSX: FP8 KV cache is not supported on
// PowerPC. The parameters are present to match the common interface.
#pragma omp parallel for collapse(2)
for (int64_t token_idx = 0; token_idx < token_num; ++token_idx) {
for (int64_t head_idx = 0; head_idx < head_num; ++head_idx) {
const int64_t pos = slot_mapping[token_idx];
if (pos < 0) continue;
const int64_t block_idx = pos / block_size;
const int64_t block_offset = pos % block_size;
{
const scalar_t* key_src = key + token_idx * key_token_num_stride +
head_idx * key_head_num_stride;
scalar_t* key_dst = key_cache + block_idx * num_blocks_stride +
head_idx * cache_head_num_stride + block_offset;
for (int64_t i = 0, j = 0; i < head_dim; ++i, j += block_size) {
key_dst[j] = key_src[i];
}
}
{
const scalar_t* val_src = value + token_idx * value_token_num_stride +
head_idx * value_head_num_stride;
scalar_t* val_dst = value_cache + block_idx * num_blocks_stride +
head_idx * cache_head_num_stride +
block_offset * head_dim;
std::memcpy(val_dst, val_src, sizeof(scalar_t) * head_dim);
}
}
}
}
};
} // namespace cpu_attention
#undef BLOCK_SIZE_ALIGNMENT
#undef HEAD_SIZE_ALIGNMENT
#undef MAX_Q_HEAD_NUM_PER_ITER
#endif // CPU_ATTN_VSX_HPP
+4
View File
@@ -9,6 +9,10 @@
namespace vec_op {
// FP8 tag types for tag dispatch (see cpu_attn_vec.hpp)
struct fp8_e4m3_tag {};
struct fp8_e5m2_tag {};
// FIXME: FP16 is not fully supported in Torch-CPU
#define VLLM_DISPATCH_CASE_FLOATING_TYPES(...) \
AT_DISPATCH_CASE(at::ScalarType::Float, __VA_ARGS__) \
+13 -2
View File
@@ -20,6 +20,7 @@ ISA_TYPES = {
"VEC16": 2,
"NEON": 3,
"VXE": 4,
"VSX": 5,
}
# KV cache index: 0 = auto (same as scalar_t), 1 = fp8_e4m3, 2 = fp8_e5m2
@@ -37,7 +38,7 @@ KV_CACHE_CPP_TYPES = {
}
# ISAs supported for head_dims divisible by 32
ISA_FOR_32 = ["AMX", "NEON", "VEC", "VEC16", "VXE"]
ISA_FOR_32 = ["AMX", "NEON", "VEC", "VEC16", "VXE", "VSX"]
# ISAs supported for head_dims divisible by 16 only
ISA_FOR_16 = ["VEC16"]
@@ -148,6 +149,10 @@ def generate_header_file() -> str:
#include "cpu_attn_vxe.hpp"
#endif
#ifdef __powerpc__
#include "cpu_attn_vsx.hpp"
#endif
"""
header += generate_helper_function()
@@ -207,6 +212,11 @@ def generate_header_file() -> str:
["VXE", "VEC", "VEC16"],
fp8=False,
)
header += _macro_block(
"#elif defined(__powerpc__)",
["VSX", "VEC", "VEC16"],
fp8=False,
)
header += _macro_block(
"#elif defined(__AVX512F__)",
["VEC", "VEC16"],
@@ -223,7 +233,8 @@ def generate_header_file() -> str:
fp8=False,
)
header += (
"#endif /* CPU_CAPABILITY_AMXBF16 / __aarch64__ / __s390x__ */\n\n"
"#endif /* CPU_CAPABILITY_AMXBF16 / __aarch64__ / "
"__s390x__ / __powerpc__ */\n\n"
"#endif // CPU_ATTN_DISPATCH_GENERATED_H\n"
)
+1 -1
View File
@@ -54,7 +54,7 @@ struct Counter {
};
inline int64_t get_available_l2_size() {
#if defined(__s390x__)
#if defined(__s390x__) || defined(__powerpc__)
static int64_t size = []() {
uint32_t l2_cache_size = 0;
auto caps = at::cpu::get_cpu_capabilities();
+19
View File
@@ -250,3 +250,22 @@ RUN --mount=type=cache,target=/root/.cache/uv \
uv pip install "vllm[zen]"
ENTRYPOINT ["vllm", "serve"]
######################### ZEN CPU TEST IMAGE #########################
FROM vllm-openai-zen AS vllm-zen-test
COPY --from=vllm-test-deps /vllm-workspace/requirements/cpu-test.txt requirements/test.txt
RUN --mount=type=cache,target=/root/.cache/uv \
uv pip install -r requirements/test.txt
ADD ./tests/ ./tests/
ADD ./examples/ ./examples/
ADD ./benchmarks/ ./benchmarks/
ADD ./vllm/collect_env.py .
ADD ./.buildkite/ ./.buildkite/
RUN --mount=type=cache,target=/root/.cache/uv \
uv pip install -e tests/vllm_test_utils
ENTRYPOINT []
+1
View File
@@ -614,6 +614,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
| `Phi4MMForCausalLM` | Phi-4-multimodal | T + I<sup>+</sup> / T + A<sup>+</sup> / I<sup>+</sup> + A<sup>+</sup> | `microsoft/Phi-4-multimodal-instruct`, etc. | ✅︎ | ✅︎ |
| `Phi4ForCausalLMV` | Phi-4-reasoning-vision | T + I<sup>+</sup> | `microsoft/Phi-4-reasoning-vision-15B`, etc. | | ✅︎ |
| `PixtralForConditionalGeneration` | Ministral 3 (Mistral format), Mistral 3 (Mistral format), Mistral Large 3 (Mistral format), Pixtral (Mistral format) | T + I<sup>+</sup> | `mistralai/Ministral-3-3B-Instruct-2512`, `mistralai/Mistral-Small-3.1-24B-Instruct-2503`, `mistralai/Mistral-Large-3-675B-Instruct-2512` `mistralai/Pixtral-12B-2409` etc. | ✅︎ | ✅︎ |
| `QianfanOCRForConditionalGeneration` | QianfanOCR | T + I<sup>E+</sup> | `baidu/Qianfan-OCR`, etc. | ✅︎ | ✅︎ |
| `QwenVLForConditionalGeneration`<sup>^</sup> | Qwen-VL | T + I<sup>E+</sup> | `Qwen/Qwen-VL`, `Qwen/Qwen-VL-Chat`, etc. | ✅︎ | ✅︎ |
| `Qwen2AudioForConditionalGeneration` | Qwen2-Audio | T + A<sup>+</sup> | `Qwen/Qwen2-Audio-7B-Instruct` | | ✅︎ |
| `Qwen2VLForConditionalGeneration` | QVQ, Qwen2-VL | T + I<sup>E+</sup> + V<sup>E+</sup> | `Qwen/QVQ-72B-Preview`, `Qwen/Qwen2-VL-7B-Instruct`, `Qwen/Qwen2-VL-72B-Instruct`, etc. | ✅︎ | ✅︎ |
+1 -1
View File
@@ -98,7 +98,7 @@ For larger scale deployments especially, it can make sense to handle the orchest
In this case, it's more convenient to treat each DP rank like a separate vLLM deployment, with its own endpoint, and have an external router balance HTTP requests between them, making use of appropriate real-time telemetry from each server for routing decisions.
This can already be done trivially for non-MoE models, since each deployed server is fully independent. No data parallel CLI options need to be used for this.
This can already be done trivially for non-MoE models, since each deployed server is fully independent. In that case, launch independent vLLM instances without any `--data-parallel-*` arguments; external DP CLI options are only supported for MoE deployments.
We support an equivalent topology for MoE DP+EP which can be configured via the following CLI arguments.
+388 -8
View File
@@ -28,6 +28,25 @@ HOPPER_MXFP4_BF16_AVAILABLE = (
and has_flashinfer()
)
# ROCm platform and dependencies
ROCM_AVAILABLE = current_platform.is_rocm()
ROCM_TRITON_KERNELS_AVAILABLE = False
ROCM_AITER_AVAILABLE = False
ROCM_GFX950 = False
if ROCM_AVAILABLE:
from vllm._aiter_ops import rocm_aiter_ops
from vllm.platforms.rocm import on_gfx950
from vllm.utils.import_utils import has_triton_kernels
ROCM_TRITON_KERNELS_AVAILABLE = has_triton_kernels()
ROCM_GFX950 = on_gfx950()
ROCM_AITER_AVAILABLE = rocm_aiter_ops.is_enabled()
if ROCM_AITER_AVAILABLE:
from aiter.ops.triton.moe.quant_moe import upcast_from_mxfp
from aiter.ops.triton.quant import dynamic_mxfp4_quant
if TRTLLM_GEN_MXFP4_AVAILABLE:
from flashinfer import (
fp4_quantize,
@@ -111,6 +130,7 @@ def test_mxfp4_loading_and_execution_moe(vllm_runner, model_case: ModelCase):
def swiglu(x, alpha: float = 1.702, beta: float = 1.0, limit: float | None = None):
# Note we add an extra bias of 1 to the linear layer
# Uses chunked layout: first half is gate, second half is up
x_glu, x_linear = torch.chunk(x, 2, dim=-1)
if limit is not None:
x_glu = x_glu.clamp(max=limit)
@@ -119,6 +139,16 @@ def swiglu(x, alpha: float = 1.702, beta: float = 1.0, limit: float | None = Non
return out_glu * (x_linear + beta)
def swigluoai(x, alpha: float = 1.702, limit: float = 7.0):
# OAI swiglu uses interleaved layout: gate/up alternating
# See SwigluOAIAndMul in vllm/model_executor/layers/activation.py
gate, up = x[..., ::2], x[..., 1::2]
gate = gate.clamp(max=limit)
up = up.clamp(min=-limit, max=limit)
glu = gate * torch.sigmoid(gate * alpha)
return (up + 1) * glu
fp4_lookup_table = [0, 0.5, 1, 1.5, 2, 3, 4, 6, -0, -0.5, -1, -1.5, -2, -3, -4, -6]
@@ -168,8 +198,20 @@ def reference_moe(
beta,
limit,
act_type,
is_gated,
activation: str = "swiglu",
use_interleaved_layout: bool = False,
):
"""
Reference MoE implementation for accuracy testing.
Args:
activation: One of "swiglu", "silu", "relu2". Controls the activation
function used after the first MLP.
use_interleaved_layout: If True, uses interleaved gate/up layout
(gate=x[..., ::2], up=x[..., 1::2]) as used by SWIGLUOAI.
If False, uses chunked layout (gate, up = chunk(x, 2)) as used
by standard swiglu/silu.
"""
# renormalize routing
experts = torch.topk(roouting_logits, k=topk, dim=-1, sorted=True)
expert_weights = torch.nn.functional.softmax(experts.values, dim=1)
@@ -179,12 +221,21 @@ def reference_moe(
mlp1_weight = w13[expert_indices, ...]
mlp1_bias = bias13[expert_indices, ...]
t = torch.einsum("beck,bk->bec", mlp1_weight, t) + mlp1_bias
if is_gated:
t = swiglu(t, alpha=alpha, beta=beta, limit=limit)
else:
# Apply activation
if activation in ("swiglu", "silu"):
if use_interleaved_layout:
# SWIGLUOAI: interleaved gate/up layout
t = swigluoai(t, alpha=alpha, limit=limit)
else:
# Standard swiglu/silu: chunked layout
t = swiglu(t, alpha=alpha, beta=beta, limit=limit)
elif activation == "relu2":
# RELU2_NO_MUL: relu(x)^2
t = torch.relu(t)
t = t * t
else:
raise ValueError(f"Unknown activation: {activation}")
if act_type == "mxfp8":
t_quantized, t_scale = mxfp8_quantize(
@@ -585,7 +636,8 @@ def test_trtllm_gen_mxfp4_fused_moe(
beta,
limit,
act_type,
is_gated=True,
activation="swiglu",
use_interleaved_layout=False,
)
ref_result[start_idx:end_idx].copy_(chunk_result)
@@ -722,7 +774,8 @@ def test_flashinfer_cutlass_mxfp4_fused_moe(
beta,
limit,
"bf16",
is_gated=True,
activation="swiglu",
use_interleaved_layout=False,
)
from vllm.utils.flashinfer import flashinfer_cutlass_fused_moe
@@ -908,7 +961,8 @@ def test_flashinfer_cutlass_mxfp4_mxfp8_fused_moe(
beta,
limit,
"mxfp8",
is_gated=True,
activation="swiglu",
use_interleaved_layout=False,
)
# Prepare inputs for FlashInfer CUTLASS fused MoE
@@ -1080,7 +1134,8 @@ def test_trtllm_gen_mxfp8_block_scale_moe(
beta=0.0,
limit=None,
act_type="mxfp8",
is_gated=is_gated,
activation="swiglu" if is_gated else "relu2",
use_interleaved_layout=False,
)
# Shuffle weights/scales with the same indexed layout used by TRTLLM kernels.
@@ -1150,3 +1205,328 @@ def test_trtllm_gen_mxfp8_block_scale_moe(
# Block-scale MXFP8 kernels are approximate; require majority close.
check_accuracy(ref, out, atol=0.1, rtol=0.85, percent=0.8)
# -----------------------------------------------------------------------------
# ROCm Oracle-based kernel execution tests
# -----------------------------------------------------------------------------
# TODO: Further tighten the accuracy threshold.
# - More accurate ref moe to include activation quantization
# - Check aiter kernel accuracy. E.g., quant / dequant details.
ROCM_BACKEND_CONFIGS = {
"TRITON": {
"activation": "SWIGLUOAI",
"rtol": 0.3,
"percent": 0.95,
"requires_aiter": False,
"requires_gfx950": False,
},
"TRITON_UNFUSED": {
"activation": "SWIGLUOAI",
"rtol": 0.3,
"percent": 0.95,
"requires_aiter": False,
"requires_gfx950": False,
},
"AITER_MXFP4_BF16": {
"activation": "SILU",
"rtol": 1.0,
"percent": 0.7,
"requires_aiter": True,
"requires_gfx950": True,
},
"AITER_MXFP4_FP8": {
"activation": "SWIGLUOAI",
"rtol": 0.5,
"percent": 0.9,
"requires_aiter": True,
"requires_gfx950": True,
},
}
@pytest.mark.parametrize("backend_name", list(ROCM_BACKEND_CONFIGS.keys()))
@pytest.mark.parametrize("topk", [4])
@pytest.mark.parametrize("num_experts", [8])
@pytest.mark.parametrize("num_tokens,hidden_size,intermediate_size", [(16, 256, 256)])
@pytest.mark.skipif(
not ROCM_AVAILABLE,
reason="ROCm is required for this test",
)
@torch.inference_mode()
def test_rocm_mxfp4_moe_oracle(
backend_name: str,
topk: int,
num_experts: int,
num_tokens: int,
hidden_size: int,
intermediate_size: int,
):
"""
Test ROCm MXFP4 MoE using oracle functions.
This test validates that the oracle functions work end-to-end:
- select_mxfp4_moe_backend() selects a valid backend
- convert_to_mxfp4_moe_kernel_format() converts weights without error
- make_mxfp4_moe_quant_config() builds a valid quant config
- make_mxfp4_moe_kernel() creates a kernel that runs without error
- The kernel output is within accuracy tolerance of reference
"""
config = ROCM_BACKEND_CONFIGS[backend_name]
# Check platform requirements
if not ROCM_TRITON_KERNELS_AVAILABLE:
pytest.skip("triton_kernels required for quantization")
if config["requires_aiter"] and not ROCM_AITER_AVAILABLE:
pytest.skip(f"Backend {backend_name} requires AITER")
if config["requires_gfx950"] and not ROCM_GFX950:
pytest.skip(f"Backend {backend_name} requires GFX950")
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
Mxfp4MoeBackend,
backend_to_kernel_cls,
convert_to_mxfp4_moe_kernel_format,
make_mxfp4_moe_kernel,
make_mxfp4_moe_quant_config,
)
from vllm.v1.worker.workspace import init_workspace_manager
# Initialize workspace manager (needed for modular kernels)
init_workspace_manager(torch.accelerator.current_device_index())
# Map string to enum
backend = Mxfp4MoeBackend[backend_name]
# Get experts class from oracle
experts_cls_list = backend_to_kernel_cls(backend)
if experts_cls_list is None or len(experts_cls_list) == 0:
pytest.skip(f"Backend {backend_name} not available")
# Use first experts class
experts_cls = experts_cls_list[0]
torch.manual_seed(42)
dtype = torch.bfloat16
device = "cuda:0"
# Create MoE config with Renormalize routing (required by monolithic kernels)
from vllm.model_executor.layers.fused_moe import FusedMoEConfig
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEParallelConfig,
RoutingMethodType,
)
moe_config = FusedMoEConfig(
num_experts=num_experts,
experts_per_token=topk,
hidden_dim=hidden_size,
intermediate_size_per_partition=intermediate_size,
num_local_experts=num_experts,
num_logical_experts=num_experts,
moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
activation=MoEActivation[config["activation"]],
in_dtype=dtype,
device="cuda",
routing_method=RoutingMethodType.Renormalize,
)
# Create float weights in checkpoint format:
# w13: [num_experts, 2*intermediate_size, hidden_size]
# w2: [num_experts, hidden_size, intermediate_size]
w13_float = torch.randn(
num_experts, 2 * intermediate_size, hidden_size, dtype=dtype, device=device
)
w2_float = torch.randn(
num_experts, hidden_size, intermediate_size, dtype=dtype, device=device
)
# dynamic_mxfp4_quant expects 2D input, so reshape 3D weights
# w13: [E, 2*I, H] -> [E*2*I, H] -> quantize -> [E, 2*I, H//2]
# w2: [E, H, I] -> [E*H, I] -> quantize -> [E, H, I//2]
w13_2d = w13_float.reshape(-1, hidden_size)
w13_quant_2d, w13_scale_2d = dynamic_mxfp4_quant(w13_2d)
w13_quant = w13_quant_2d.reshape(num_experts, 2 * intermediate_size, -1)
w13_scale = w13_scale_2d.reshape(num_experts, 2 * intermediate_size, -1)
w2_2d = w2_float.reshape(-1, intermediate_size)
w2_quant_2d, w2_scale_2d = dynamic_mxfp4_quant(w2_2d)
w2_quant = w2_quant_2d.reshape(num_experts, hidden_size, -1)
w2_scale = w2_scale_2d.reshape(num_experts, hidden_size, -1)
w13_bias = torch.randn(
num_experts, 2 * intermediate_size, dtype=dtype, device=device
)
w2_bias = torch.randn(num_experts, hidden_size, dtype=dtype, device=device)
# Create static input scales for W4A8 backend (AITER_MXFP4_FP8)
w13_input_scale: torch.Tensor | None = None
w2_input_scale: torch.Tensor | None = None
if backend_name == "AITER_MXFP4_FP8":
# Static FP8 scales: one scale per expert
w13_input_scale = torch.ones(num_experts, dtype=torch.float32, device=device)
w2_input_scale = torch.ones(num_experts, dtype=torch.float32, device=device)
# Create mock layer for oracle functions
class MockLayer:
w13_weight: torch.Tensor
w2_weight: torch.Tensor
w13_weight_scale: torch.Tensor
w2_weight_scale: torch.Tensor
w13_input_scale: torch.Tensor | None
w2_input_scale: torch.Tensor | None
layer = MockLayer()
layer.w13_weight = w13_quant
layer.w2_weight = w2_quant
layer.w13_weight_scale = w13_scale
layer.w2_weight_scale = w2_scale
layer.w13_input_scale = w13_input_scale
layer.w2_input_scale = w2_input_scale
# Convert weights using oracle
w13_conv, w2_conv, w13_scale_conv, w2_scale_conv, w13_bias_conv, w2_bias_conv = (
convert_to_mxfp4_moe_kernel_format(
mxfp4_backend=backend,
layer=layer, # type: ignore[arg-type]
w13_weight=w13_quant,
w2_weight=w2_quant,
w13_weight_scale=w13_scale,
w2_weight_scale=w2_scale,
w13_bias=w13_bias,
w2_bias=w2_bias,
)
)
# Build quant config using oracle
quant_config = make_mxfp4_moe_quant_config(
mxfp4_backend=backend,
w1_scale=w13_scale_conv,
w2_scale=w2_scale_conv,
w1_bias=w13_bias_conv,
w2_bias=w2_bias_conv,
a1_scale=w13_input_scale,
a2_scale=w2_input_scale,
)
# Select activation based on backend
activation_name = str(config["activation"])
activation = MoEActivation[activation_name]
# Build kernel using oracle
assert quant_config is not None, "Failed to create quant config"
with set_current_vllm_config(VllmConfig()):
kernel = make_mxfp4_moe_kernel(
moe_quant_config=quant_config,
moe_config=moe_config,
mxfp4_backend=backend,
experts_cls=experts_cls,
routing_tables=None,
shared_experts=None,
)
# Create inputs
x = torch.randn(num_tokens, hidden_size, dtype=dtype, device=device)
router_logits = torch.randn(
num_tokens, num_experts, dtype=torch.float32, device=device
)
topk_weights, topk_ids = torch.topk(router_logits, k=topk, dim=-1, sorted=True)
topk_weights = torch.nn.functional.softmax(topk_weights, dim=-1)
# Run kernel - use appropriate method based on impl type
if kernel.is_monolithic:
# Monolithic impl uses router_logits
out = kernel.apply_monolithic(
hidden_states=x,
w1=w13_conv,
w2=w2_conv,
router_logits=router_logits,
activation=activation,
global_num_experts=num_experts,
expert_map=None,
apply_router_weight_on_input=False,
)
else:
# Modular impl uses topk_weights and topk_ids
out = kernel.apply(
hidden_states=x,
w1=w13_conv,
w2=w2_conv,
topk_weights=topk_weights,
topk_ids=topk_ids,
activation=activation,
global_num_experts=num_experts,
expert_map=None,
apply_router_weight_on_input=False,
)
# Verify output is valid (no NaN/Inf) and has expected shape
assert out.shape == (num_tokens, hidden_size), f"Unexpected shape: {out.shape}"
assert not torch.any(torch.isnan(out)), "Output contains NaN"
assert not torch.any(torch.isinf(out)), "Output contains Inf"
# Verify output has reasonable magnitude (not all zeros)
assert out.abs().max() > 0.01, "Output is effectively zero"
# Dequantize weights for reference computation
w13_dq = upcast_from_mxfp(
w13_quant.view(torch.uint8), w13_scale, torch.bfloat16, axis=-1
)
w2_dq = upcast_from_mxfp(
w2_quant.view(torch.uint8), w2_scale, torch.bfloat16, axis=-1
)
# Determine activation type and layout
# SWIGLUOAI uses interleaved layout (gate/up alternating)
# SILU uses chunked layout (first half gate, second half up)
use_interleaved = activation == MoEActivation.SWIGLUOAI
if activation in [MoEActivation.SWIGLUOAI, MoEActivation.SILU]:
act_name = "swiglu"
else:
act_name = "relu2"
ref = reference_moe(
router_logits,
topk,
num_experts,
x.to(torch.float32),
w13_dq.to(torch.float32),
w13_bias.to(torch.float32),
w2_dq.to(torch.float32),
w2_bias.to(torch.float32),
alpha=1.702 if activation == MoEActivation.SWIGLUOAI else 1.0,
beta=1.0 if activation == MoEActivation.SWIGLUOAI else 0.0,
limit=7.0 if activation == MoEActivation.SWIGLUOAI else None,
act_type="bf16",
activation=act_name,
use_interleaved_layout=use_interleaved,
)
# Compute and print accuracy statistics
diff = (ref.float() - out.float()).abs()
rel_diff = diff / (ref.float().abs() + 1e-6)
print(f"\n[{backend_name}] Accuracy statistics:")
print(
f" Reference: min={ref.min():.4f}, max={ref.max():.4f}, mean={ref.mean():.4f}"
)
print(
f" Output: min={out.min():.4f}, max={out.max():.4f}, mean={out.mean():.4f}"
)
print(
f" Abs diff: min={diff.min():.4f}, max={diff.max():.4f}, "
f"mean={diff.mean():.4f}"
)
print(
f" Rel diff: min={rel_diff.min():.4f}, max={rel_diff.max():.4f}, "
f"mean={rel_diff.mean():.4f}"
)
# Check what percentage of values are within various tolerances
for rtol in [0.1, 0.5, 1.0, 2.0]:
within_tol = (diff <= rtol * out.float().abs()).float().mean()
print(f" Within rtol={rtol}: {within_tol * 100:.1f}%")
# Check accuracy using per-backend thresholds
check_accuracy(ref, out, atol=0.1, rtol=config["rtol"], percent=config["percent"])
@@ -42,6 +42,11 @@ AITER_MODEL_LIST = [
pytest.mark.core_model,
pytest.mark.slow_test,
pytest.mark.cpu_model,
pytest.mark.skipif(
current_platform.is_zen_cpu(),
reason="bloom-560m ALiBi is currently not supported on\
AMD Zen CPUs due to lack of support for float16 compute.",
),
],
),
pytest.param(
+18 -2
View File
@@ -88,7 +88,15 @@ def load_reward_outputs(filename: "StrPath") -> list[list[float]]:
[
pytest.param(
"Qwen/Qwen2.5-Math-PRM-7B",
marks=[pytest.mark.core_model, pytest.mark.cpu_model],
marks=[
pytest.mark.core_model,
pytest.mark.cpu_model,
pytest.mark.skipif(
current_platform.is_zen_cpu(),
reason="Qwen2.5-Math-PRM-7B is currently not supported on\
AMD Zen CPUs due to lack of support for float16 compute.",
)
],
),
],
)
@@ -131,7 +139,15 @@ def test_prm_models(
[
pytest.param(
"Qwen/Qwen2.5-Math-PRM-7B",
marks=[pytest.mark.core_model, pytest.mark.cpu_model],
marks=[
pytest.mark.core_model,
pytest.mark.cpu_model,
pytest.mark.skipif(
current_platform.is_zen_cpu(),
reason="Qwen2.5-Math-PRM-7B is currently not supported on\
AMD Zen CPUs due to lack of support for float16 compute.",
)
],
),
],
)
@@ -928,6 +928,16 @@ VLM_TEST_SETTINGS = {
),
],
),
"qianfan_ocr": VLMTestInfo(
models=["baidu/Qianfan-OCR"],
test_type=(VLMTestType.IMAGE, VLMTestType.MULTI_IMAGE),
prompt_formatter=lambda img_prompt: f"<|im_start|>user\n{img_prompt}<|im_end|>\n<|im_start|>assistant\n", # noqa: E501
img_idx_to_prompt=lambda idx: "<image>",
max_model_len=4096,
use_tokenizer_eos=True,
auto_cls=AutoModelForImageTextToText,
hf_model_kwargs=model_utils.qianfan_ocr_hf_model_kwargs("baidu/Qianfan-OCR"),
),
"qwen_vl": VLMTestInfo(
models=["Qwen/Qwen-VL"],
test_type=(VLMTestType.IMAGE, VLMTestType.MULTI_IMAGE),
@@ -1554,3 +1554,94 @@ def moondream3_patch_hf_runner(hf_model: HfRunner) -> HfRunner:
hf_model.model.generate = types.MethodType(_generate, hf_model.model)
return hf_model
def qianfan_ocr_hf_model_kwargs(model_name: str) -> dict:
"""Return hf_model_kwargs with a patched config for QianfanOCR."""
from vllm.transformers_utils.configs.qianfan_ocr import QianfanOCRConfig
config = QianfanOCRConfig.from_pretrained(model_name)
vc = config.vision_config
if isinstance(vc.image_size, int):
vc.image_size = (vc.image_size, vc.image_size)
if isinstance(vc.patch_size, int):
vc.patch_size = (vc.patch_size, vc.patch_size)
return {"config": config}
def qianfan_ocr_patch_hf_runner(hf_model: HfRunner) -> HfRunner:
"""Patches an HfRunner instance to run QianfanOCR model inference.
QianfanOCR shares the same architecture as InternVLChatModel, so the
patching logic mirrors ``internvl_patch_hf_runner``. The only difference
is that we load the config via vllm's registered ``QianfanOCRConfig``
instead of relying on ``trust_remote_code``.
"""
class QianfanOCRProcessor:
def __init__(self, hf_runner: HfRunner):
self.tokenizer = hf_runner.tokenizer
from vllm.transformers_utils.configs.qianfan_ocr import QianfanOCRConfig
self.config = QianfanOCRConfig.from_pretrained(hf_runner.model_name)
self.vision_config = self.config.vision_config
self.use_thumbnail = self.config.use_thumbnail
self.min_num = self.config.min_dynamic_patch
self.max_num = self.config.max_dynamic_patch
self.image_size = self.vision_config.image_size
# Compute num_image_token from config instead of model attribute,
# since the transformers-native model doesn't expose it.
image_size = self.config.force_image_size or self.vision_config.image_size
patch_size = self.vision_config.patch_size
downsample_ratio = self.config.downsample_ratio
self.num_image_token = int(
(image_size // patch_size) ** 2 * (downsample_ratio**2)
)
def __call__(
self,
text: str,
images: PIL.Image.Image | list[PIL.Image.Image] = None,
**kwargs,
):
from vllm.transformers_utils.processors.internvl import (
image_to_pixel_values_internvl,
)
IMG_START = "<img>"
IMG_END = "</img>"
IMG_CONTEXT = "<IMG_CONTEXT>"
images = [images] if isinstance(images, PIL.Image.Image) else images
pixel_values_list = [
image_to_pixel_values_internvl(
image,
input_size=self.image_size,
min_num=self.min_num,
max_num=self.max_num,
use_thumbnail=self.use_thumbnail,
)
for image in images
]
num_patches_list = [pv.shape[0] for pv in pixel_values_list]
pixel_values = torch.cat(pixel_values_list, dim=0)
for num_patches in num_patches_list:
context_tokens = IMG_CONTEXT * self.num_image_token * num_patches
image_tokens = IMG_START + context_tokens + IMG_END
text = text.replace("<image>", image_tokens, 1)
prompt = self.tokenizer(text, return_tensors="pt")
prompt.update({"pixel_values": pixel_values})
return prompt
img_context_token_id = hf_model.tokenizer.convert_tokens_to_ids("<IMG_CONTEXT>")
hf_model.model.img_context_token_id = img_context_token_id
hf_model.processor = QianfanOCRProcessor(hf_model)
hf_model.model.get_output_embeddings = (
lambda: hf_model.model.language_model.get_output_embeddings()
)
hf_model.model.generate = types.MethodType(_internvl_generate, hf_model.model)
return hf_model
+4
View File
@@ -1264,6 +1264,10 @@ _MULTIMODAL_EXAMPLE_MODELS = {
},
tokenizer_mode="mistral",
),
"QianfanOCRForConditionalGeneration": _HfExamplesInfo(
"baidu/Qianfan-OCR",
min_transformers_version="5.6.0",
),
"QwenVLForConditionalGeneration": _HfExamplesInfo(
"Qwen/Qwen-VL",
extras={"chat": "Qwen/Qwen-VL-Chat"},
+82 -4
View File
@@ -182,22 +182,100 @@ class TestTurboQuantConfig:
# ---- Boundary skip layers ----
@staticmethod
def _dense_model_config(num_layers):
from types import SimpleNamespace
return SimpleNamespace(
is_hybrid=False,
hf_text_config=SimpleNamespace(num_hidden_layers=num_layers),
)
def test_boundary_skip_layers_basic(self):
layers = TurboQuantConfig.get_boundary_skip_layers(32)
mc = self._dense_model_config(32)
layers = TurboQuantConfig.get_boundary_skip_layers(mc)
assert layers == ["0", "1", "30", "31"]
def test_boundary_skip_layers_zero(self):
assert TurboQuantConfig.get_boundary_skip_layers(32, 0) == []
mc = self._dense_model_config(32)
assert TurboQuantConfig.get_boundary_skip_layers(mc, 0) == []
def test_boundary_skip_layers_small_model(self):
layers = TurboQuantConfig.get_boundary_skip_layers(4)
mc = self._dense_model_config(4)
layers = TurboQuantConfig.get_boundary_skip_layers(mc)
assert layers == ["0", "1", "2", "3"]
def test_boundary_skip_layers_cap_at_half(self):
layers = TurboQuantConfig.get_boundary_skip_layers(8, 10)
mc = self._dense_model_config(8)
layers = TurboQuantConfig.get_boundary_skip_layers(mc, 10)
assert len(layers) == 8
class TestHybridAttentionIndices:
"""Regression tests for boundary protection on hybrid models.
Hybrid models (attention + Mamba / linear-attention) identify KV-carrying
layers via layer_types / layers_block_type / attn_type_list. The helper
must return the *global* layer indices of the full-attention layers so
that kv_cache_dtype_skip_layers matches what extract_layer_index(prefix)
reports on the Attention layers at runtime.
"""
@staticmethod
def _fake_model_config(text_cfg=None, hf_cfg=None):
from types import SimpleNamespace
return SimpleNamespace(
hf_text_config=text_cfg if text_cfg is not None else SimpleNamespace(),
hf_config=hf_cfg if hf_cfg is not None else SimpleNamespace(),
)
def test_layer_types_full_attention(self):
from vllm.model_executor.layers.quantization.turboquant.config import (
_get_full_attention_layer_indices,
)
cfg = type("C", (), {})()
cfg.layer_types = [
"linear_attention",
"linear_attention",
"full_attention",
"linear_attention",
"full_attention",
"full_attention",
]
mc = self._fake_model_config(text_cfg=cfg)
assert _get_full_attention_layer_indices(mc) == [2, 4, 5]
def test_layers_block_type_jamba(self):
from vllm.model_executor.layers.quantization.turboquant.config import (
_get_full_attention_layer_indices,
)
cfg = type("C", (), {})()
cfg.layers_block_type = ["mamba", "attention", "mamba", "attention"]
mc = self._fake_model_config(text_cfg=cfg)
assert _get_full_attention_layer_indices(mc) == [1, 3]
def test_attn_type_list_minimax(self):
from vllm.model_executor.layers.quantization.turboquant.config import (
_get_full_attention_layer_indices,
)
hf = type("C", (), {})()
hf.attn_type_list = [0, 1, 0, 1, 1]
mc = self._fake_model_config(hf_cfg=hf)
assert _get_full_attention_layer_indices(mc) == [1, 3, 4]
def test_no_hybrid_hints_returns_empty(self):
from vllm.model_executor.layers.quantization.turboquant.config import (
_get_full_attention_layer_indices,
)
mc = self._fake_model_config()
assert _get_full_attention_layer_indices(mc) == []
# ============================================================================
# Centroids tests (CPU-only)
# ============================================================================
-2
View File
@@ -1215,8 +1215,6 @@ def test_scheduler_config_init():
("facebook/opt-125m", 1, False, False),
# Non-MoE model with DP>1 internal LB should need coordinator
("facebook/opt-125m", 2, False, True),
# Non-MoE model with DP>1 external LB should not need coordinator
("facebook/opt-125m", 2, True, False),
# MoE model with DP=1 should not need coordinator
("mistralai/Mixtral-8x7B-Instruct-v0.1", 1, False, False),
# MoE model with DP>1 internal LB should need both coordinator
+240
View File
@@ -0,0 +1,240 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import os
import sys
from types import SimpleNamespace
from unittest import mock
import pytest
from vllm.triton_utils import jit_monitor
@pytest.fixture(autouse=True)
def _reset_monitor():
"""Reset global monitor state between tests."""
jit_monitor._active = False
yield
jit_monitor._active = False
# ------------------------------------------------------------------
# Helpers — lightweight stand-ins for triton.knobs
# ------------------------------------------------------------------
def _make_fake_knobs(*, autotuning_print=False, jit_hook=None):
"""Build a minimal fake ``triton.knobs`` namespace."""
autotuning = SimpleNamespace(print=autotuning_print)
runtime = SimpleNamespace(jit_post_compile_hook=jit_hook)
return SimpleNamespace(autotuning=autotuning, runtime=runtime)
def _patch_triton_knobs(fake_knobs):
"""Context manager that makes ``from triton import knobs`` return *fake_knobs*."""
fake_triton = SimpleNamespace(knobs=fake_knobs)
return mock.patch.dict(sys.modules, {"triton": fake_triton})
# ------------------------------------------------------------------
# Unit tests (no GPU required, triton is mocked)
# ------------------------------------------------------------------
class TestActivateBasic:
def test_sets_active(self):
assert not jit_monitor.is_active()
with _patch_triton_knobs(_make_fake_knobs()):
jit_monitor.activate()
assert jit_monitor.is_active()
def test_idempotent(self):
fake = _make_fake_knobs()
with _patch_triton_knobs(fake):
jit_monitor.activate()
first_hook = fake.runtime.jit_post_compile_hook
jit_monitor.activate()
assert fake.runtime.jit_post_compile_hook is first_hook
def test_logs_info_on_activation(self):
with (
mock.patch.object(jit_monitor.logger, "info") as m,
_patch_triton_knobs(_make_fake_knobs()),
):
jit_monitor.activate()
m.assert_called_once()
assert "Kernel JIT monitor activated" in m.call_args[0][0]
class TestAutotuningPrint:
def test_enables_autotuning_print(self):
fake = _make_fake_knobs(autotuning_print=False)
with _patch_triton_knobs(fake):
jit_monitor.activate()
assert fake.autotuning.print is True
def test_respects_user_opt_out(self):
fake = _make_fake_knobs(autotuning_print=False)
with (
mock.patch.dict(os.environ, {"TRITON_PRINT_AUTOTUNING": "0"}),
_patch_triton_knobs(fake),
):
jit_monitor.activate()
assert fake.autotuning.print is False
def test_noop_when_user_already_enabled(self):
fake = _make_fake_knobs(autotuning_print=True)
with (
mock.patch.dict(os.environ, {"TRITON_PRINT_AUTOTUNING": "1"}),
_patch_triton_knobs(fake),
):
jit_monitor.activate()
assert fake.autotuning.print is True
class TestJitHook:
def test_hook_registered(self):
fake = _make_fake_knobs()
assert fake.runtime.jit_post_compile_hook is None
with _patch_triton_knobs(fake):
jit_monitor.activate()
assert fake.runtime.jit_post_compile_hook is not None
def test_hook_logs_warning(self):
fake = _make_fake_knobs()
with _patch_triton_knobs(fake):
jit_monitor.activate()
hook = fake.runtime.jit_post_compile_hook
mock_fn = SimpleNamespace(name="test_kernel")
with mock.patch.object(jit_monitor.logger, "warning") as m:
hook(
key="some_key",
repr="some_repr",
fn=mock_fn,
compile=lambda: None,
is_manual_warmup=False,
already_compiled=False,
)
m.assert_called_once()
msg = m.call_args[0][0] % m.call_args[0][1:]
assert "Triton kernel JIT compilation during inference" in msg
assert "test_kernel" in msg
def test_hook_chains_existing_hook(self):
existing = mock.MagicMock(return_value="existing_result")
fake = _make_fake_knobs(jit_hook=existing)
with _patch_triton_knobs(fake):
jit_monitor.activate()
hook = fake.runtime.jit_post_compile_hook
mock_fn = SimpleNamespace(name="chained_kernel")
kwargs = dict(
key="k",
repr="r",
fn=mock_fn,
compile=lambda: None,
is_manual_warmup=False,
already_compiled=False,
)
result = hook(**kwargs)
existing.assert_called_once()
assert result == "existing_result"
def test_hook_works_without_existing_hook(self):
fake = _make_fake_knobs(jit_hook=None)
with _patch_triton_knobs(fake):
jit_monitor.activate()
hook = fake.runtime.jit_post_compile_hook
mock_fn = SimpleNamespace(name="solo_kernel")
result = hook(
key="k",
repr="r",
fn=mock_fn,
compile=lambda: None,
is_manual_warmup=False,
already_compiled=False,
)
assert result is None
class TestNoTritonFallback:
def test_activate_without_triton(self):
with mock.patch.object(jit_monitor, "HAS_TRITON", False):
jit_monitor.activate()
assert jit_monitor.is_active()
# ------------------------------------------------------------------
# Integration tests (real Triton + GPU)
# ------------------------------------------------------------------
try:
import torch
_HAS_CUDA = torch.cuda.is_available()
except ImportError:
_HAS_CUDA = False
try:
import triton
import triton.language as tl
_HAS_TRITON = True
except ImportError:
_HAS_TRITON = False
_skip_no_gpu = pytest.mark.skipif(
not (_HAS_CUDA and _HAS_TRITON),
reason="Requires CUDA GPU and Triton",
)
if _HAS_TRITON:
@triton.jit
def _add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < n
x = tl.load(x_ptr + offs, mask=mask)
y = tl.load(y_ptr + offs, mask=mask)
tl.store(out_ptr + offs, x + y, mask=mask)
def _run_add_kernel(n: int, block: int = 256) -> None:
"""Launch ``_add_kernel`` with vectors of length *n*."""
x = torch.randn(n, device="cuda")
y = torch.randn(n, device="cuda")
out = torch.empty(n, device="cuda")
grid = ((n + block - 1) // block,)
_add_kernel[grid](x, y, out, n, BLOCK=block)
torch.accelerator.synchronize()
@_skip_no_gpu
class TestTritonJitHookIntegration:
"""End-to-end: real Triton kernel, real GPU, real hook."""
def test_no_warning_on_cached_shape(self):
_run_add_kernel(1024)
jit_monitor.activate()
with mock.patch.object(jit_monitor.logger, "warning") as w:
_run_add_kernel(1024)
w.assert_not_called()
def test_warning_on_new_constexpr(self):
_run_add_kernel(1024, block=256)
jit_monitor.activate()
with mock.patch.object(jit_monitor.logger, "warning") as w:
# Different BLOCK (a tl.constexpr) forces recompilation.
_run_add_kernel(1024, block=512)
w.assert_called()
msg = w.call_args[0][0] % w.call_args[0][1:]
assert "_add_kernel" in msg
+5 -3
View File
@@ -22,6 +22,7 @@ from vllm.config.vllm import set_current_vllm_config
from vllm.model_executor.layers.attention.mla_attention import (
QueryLenSupport,
_DecodeConcatQuantFP8,
get_mla_prefill_scale,
)
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.model_executor.layers.quantization.utils.quant_utils import GroupShape
@@ -785,7 +786,8 @@ def test_backend_correctness(
assert kv_lora_rank + qk_rope_head_dim == head_size, (
f"MLA dimensions don't match: {total_head_size} != {head_size}"
)
scale = 1.0 / (total_head_size**0.5)
decode_scale = 1.0 / (total_head_size**0.5)
prefill_scale = get_mla_prefill_scale(vllm_config.model_config)
# 2. Generate data and compute SDPA reference output for MLA
all_q_vllm, all_kv_c_vllm, all_k_pe_vllm = [], [], []
@@ -902,7 +904,7 @@ def test_backend_correctness(
v_sdpa_in = v_mqa.unsqueeze(0).transpose(1, 2)
sdpa_out_i_decode = torch.nn.functional.scaled_dot_product_attention(
q_sdpa_in, k_sdpa_in, v_sdpa_in, attn_mask=attn_mask, scale=scale
q_sdpa_in, k_sdpa_in, v_sdpa_in, attn_mask=attn_mask, scale=decode_scale
)
sdpa_out_i_decode = sdpa_out_i_decode.transpose(1, 2).squeeze(
0
@@ -938,7 +940,7 @@ def test_backend_correctness(
# Single attention call with custom mask
sdpa_out_i_prefill = torch.nn.functional.scaled_dot_product_attention(
q_sdpa_in, k_sdpa_in, v_sdpa_in, attn_mask=attn_mask, scale=scale
q_sdpa_in, k_sdpa_in, v_sdpa_in, attn_mask=attn_mask, scale=prefill_scale
)
sdpa_out_i_prefill = sdpa_out_i_prefill.transpose(1, 2).squeeze(0)
sdpa_out_i_prefill = sdpa_out_i_prefill.flatten(start_dim=-2)
@@ -2,12 +2,17 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for MLA prefill backend selector."""
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
import torch
from vllm.config import AttentionConfig, ModelConfig, VllmConfig
from vllm.model_executor.layers.attention.mla_attention import get_mla_prefill_scale
from vllm.model_executor.layers.rotary_embedding.deepseek_scaling_rope import (
yarn_get_mscale,
)
from vllm.platforms.interface import DeviceCapability
from vllm.v1.attention.backends.mla.prefill.registry import MLAPrefillBackendEnum
from vllm.v1.attention.backends.mla.prefill.selector import (
@@ -53,6 +58,62 @@ def _make_vllm_config(
return mock_vllm_config
class TestMLAPrefillScale:
"""Tests for the MLA prefill softmax scale."""
def test_uses_qk_head_dim_for_deepseek_v2_style_mla(self):
model_config = SimpleNamespace(
hf_text_config=SimpleNamespace(
q_lora_rank=None,
kv_lora_rank=512,
qk_nope_head_dim=128,
qk_rope_head_dim=64,
v_head_dim=128,
rope_parameters={"rope_type": "default"},
)
)
assert get_mla_prefill_scale(model_config) == pytest.approx(192**-0.5)
def test_applies_deepseek_yarn_mscale(self):
model_config = SimpleNamespace(
hf_text_config=SimpleNamespace(
q_lora_rank=None,
kv_lora_rank=512,
qk_nope_head_dim=128,
qk_rope_head_dim=64,
v_head_dim=128,
rope_parameters={
"rope_type": "yarn",
"factor": 40,
"mscale_all_dim": 0.707,
},
)
)
mscale = yarn_get_mscale(40, 0.707)
assert get_mla_prefill_scale(model_config) == pytest.approx(
192**-0.5 * mscale * mscale
)
def test_deepseek_v4_style_mla_does_not_apply_yarn_mscale(self):
model_config = SimpleNamespace(
hf_text_config=SimpleNamespace(
compress_ratios=[4],
q_lora_rank=1536,
head_dim=128,
qk_rope_head_dim=64,
rope_parameters={
"rope_type": "yarn",
"factor": 40,
"mscale_all_dim": 0.707,
},
)
)
assert get_mla_prefill_scale(model_config) == pytest.approx(128**-0.5)
class TestGetMLAPrefillBackend:
"""Tests for get_mla_prefill_backend (public API)."""
+1 -1
View File
@@ -14,7 +14,7 @@ import requests
from tests.utils import RemoteOpenAIServer
from vllm.platforms import current_platform
MODEL_NAME = "ibm-research/PowerMoE-3b"
MODEL_NAME = os.getenv("MODEL_NAME", "ibm-research/PowerMoE-3b")
# Number of data parallel ranks for external LB testing
DP_SIZE = int(os.getenv("DP_SIZE", "2"))
@@ -0,0 +1,281 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import threading
from unittest.mock import MagicMock
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_connector import (
MooncakeConnector,
MooncakeConnectorWorker,
SendBlockMeta,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.stats import (
MooncakeKVConnectorStats,
)
def test_is_empty_on_fresh_stats():
stats = MooncakeKVConnectorStats()
assert stats.is_empty()
assert stats.num_successful_transfers == 0
def test_record_transfer_and_reduce():
stats = MooncakeKVConnectorStats()
# 1 MB transfer in 1 ms -> 1000 MB/s throughput
stats.record_transfer(duration_s=0.001, total_bytes=1 * 2**20, num_descs=4)
# 2 MB transfer in 2 ms
stats.record_transfer(duration_s=0.002, total_bytes=2 * 2**20, num_descs=6)
assert not stats.is_empty()
assert stats.num_successful_transfers == 2
reduced = stats.reduce()
assert reduced["Num successful transfers"] == 2
# avg = (1 + 2) / 2 = 1.5 ms
assert reduced["Avg xfer time (ms)"] == 1.5
assert reduced["Avg MB per transfer"] == 1.5
# 3 MB total / 3 ms total = 1000 MB/s
assert reduced["Throughput (MB/s)"] == 1000.0
assert reduced["Avg number of descriptors"] == 5.0
assert reduced["Num failed transfers"] == 0
assert reduced["Num failed recvs"] == 0
assert reduced["Num KV expired reqs"] == 0
def test_record_failures_keeps_stats_non_empty():
stats = MooncakeKVConnectorStats()
stats.record_failed_transfer()
stats.record_failed_recv()
stats.record_kv_expired_req()
assert not stats.is_empty()
reduced = stats.reduce()
# No successful transfers -> latency/throughput all zero, but failure
# counters still surface.
assert reduced["Num successful transfers"] == 0
assert reduced["Num failed transfers"] == 1
assert reduced["Num failed recvs"] == 1
assert reduced["Num KV expired reqs"] == 1
def test_aggregate_sums_observations():
a = MooncakeKVConnectorStats()
b = MooncakeKVConnectorStats()
a.record_transfer(duration_s=0.001, total_bytes=1 * 2**20, num_descs=1)
b.record_transfer(duration_s=0.002, total_bytes=2 * 2**20, num_descs=2)
b.record_failed_transfer()
a.aggregate(b)
assert a.num_successful_transfers == 2
reduced = a.reduce()
assert reduced["Num successful transfers"] == 2
assert reduced["Num failed transfers"] == 1
def test_aggregate_with_empty_other_is_noop():
a = MooncakeKVConnectorStats()
a.record_transfer(duration_s=0.001, total_bytes=1, num_descs=1)
b = MooncakeKVConnectorStats()
a.aggregate(b)
assert a.num_successful_transfers == 1
def test_getstate_drops_lock_and_setstate_recreates_it():
# KVConnectorStats subclasses must be picklable (worker→scheduler IPC),
# but threading.Lock isn't — so __getstate__ strips it and __setstate__
# rebuilds a fresh per-process lock.
original = MooncakeKVConnectorStats()
original.record_transfer(duration_s=0.01, total_bytes=2048, num_descs=3)
state = original.__getstate__()
assert "_lock" not in state
rebuilt = MooncakeKVConnectorStats.__new__(MooncakeKVConnectorStats)
rebuilt.__setstate__(state)
assert rebuilt.data == original.data
# Lock works on the receiver side.
rebuilt.record_transfer(duration_s=0.02, total_bytes=4096, num_descs=5)
assert rebuilt.num_successful_transfers == 2
def test_concurrent_writers_keep_row_lengths_aligned():
# Multiple writers + a snapshot reader must never produce a snapshot
# with mismatched column lengths — reduce()'s
# len(descs) == num_successful_transfers assertion would fire.
stats = MooncakeKVConnectorStats()
stop = threading.Event()
writer_count = 4
snapshots: list[MooncakeKVConnectorStats] = []
def writer():
i = 0
while not stop.is_set():
stats.record_transfer(
duration_s=0.001 + i * 1e-9,
total_bytes=1024 + i,
num_descs=1 + (i % 8),
)
i += 1
def snapper():
while not stop.is_set():
snap = stats.clone_and_reset()
if not snap.is_empty():
# Force the same path the logger walks; reduce() will
# blow up on torn rows via its internal assert.
snap.reduce()
snapshots.append(snap)
threads = [threading.Thread(target=writer) for _ in range(writer_count)]
snapshotter = threading.Thread(target=snapper)
for t in threads:
t.start()
snapshotter.start()
# Short fixed window — long enough to interleave thousands of ops.
threading.Event().wait(0.2)
stop.set()
for t in threads:
t.join()
snapshotter.join()
# Final drain so we don't lose the in-flight tail.
final = stats.clone_and_reset()
if not final.is_empty():
final.reduce()
snapshots.append(final)
# Every snapshot's columns must have identical lengths (the invariant
# the lock protects), and the union must contain at least one row.
total_rows = 0
for snap in snapshots:
n = len(snap.data["transfer_duration"])
assert len(snap.data["bytes_transferred"]) == n
assert len(snap.data["num_descriptors"]) == n
total_rows += n
assert total_rows > 0
def test_clone_and_reset_hands_off_old_data():
stats = MooncakeKVConnectorStats()
stats.record_transfer(duration_s=0.001, total_bytes=1, num_descs=1)
stats.record_failed_recv()
snapshot = stats.clone_and_reset()
assert snapshot.num_successful_transfers == 1
assert not snapshot.is_empty()
# Original is now empty.
assert stats.is_empty()
assert stats.num_successful_transfers == 0
# Recording on the original does not mutate the snapshot.
stats.record_transfer(duration_s=0.005, total_bytes=2, num_descs=2)
assert snapshot.num_successful_transfers == 1
def test_build_kv_connector_stats_none_returns_empty_instance():
out = MooncakeConnector.build_kv_connector_stats()
assert isinstance(out, MooncakeKVConnectorStats)
assert out.is_empty()
def test_build_kv_connector_stats_with_data_round_trips():
original = MooncakeKVConnectorStats()
original.record_transfer(duration_s=0.01, total_bytes=1024, num_descs=3)
original.record_failed_transfer()
# Serialized form is the .data dict; build should reconstruct an instance
# that behaves the same.
rebuilt = MooncakeConnector.build_kv_connector_stats(data=original.data)
assert isinstance(rebuilt, MooncakeKVConnectorStats)
assert rebuilt.num_successful_transfers == 1
assert rebuilt.reduce()["Num failed transfers"] == 1
def _bare_worker() -> MooncakeConnectorWorker:
"""Construct a MooncakeConnectorWorker skipping __init__ (full init requires
a live TransferEngine). Only the attributes touched by the methods under
test are populated; role flags and async_zmq_ctx keep __del__'s shutdown
path a no-op."""
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.xfer_stats = MooncakeKVConnectorStats()
worker.engine = MagicMock()
worker.async_zmq_ctx = MagicMock()
worker.is_kv_consumer = True
worker.is_kv_producer = True
return worker
def test_send_blocks_records_success():
worker = _bare_worker()
worker.engine.batch_transfer_sync_write.return_value = 0
ret = worker._send_blocks(
"host:1234",
src_ptrs=[0x1000, 0x2000],
dst_ptrs=[0x3000, 0x4000],
lengths=[1024, 2048],
)
assert ret == 0
assert worker.xfer_stats.num_successful_transfers == 1
data = worker.xfer_stats.data
assert data["bytes_transferred"] == [1024 + 2048]
assert data["num_descriptors"] == [2]
assert data["num_failed_transfers"] == []
def test_send_blocks_records_failure():
worker = _bare_worker()
worker.engine.batch_transfer_sync_write.return_value = 1 # non-zero = fail
ret = worker._send_blocks("host:1234", [0x1000], [0x2000], [4096])
assert ret == 1
assert worker.xfer_stats.num_successful_transfers == 0
assert worker.xfer_stats.data["num_failed_transfers"] == [1]
def test_get_kv_connector_stats_returns_none_when_empty():
worker = _bare_worker()
assert worker.get_kv_connector_stats() is None
def test_get_kv_connector_stats_returns_and_resets():
worker = _bare_worker()
worker.engine.batch_transfer_sync_write.return_value = 0
worker._send_blocks("host:1234", [0x1000], [0x2000], [4096])
snapshot = worker.get_kv_connector_stats()
assert isinstance(snapshot, MooncakeKVConnectorStats)
assert snapshot.num_successful_transfers == 1
# Second call returns None because the worker's stats were reset.
assert worker.get_kv_connector_stats() is None
def test_expired_request_bumps_counter():
import asyncio
worker = _bare_worker()
worker.reqs_need_send = {
"tid1": SendBlockMeta(
p_req_id="req1",
transfer_id="tid1",
local_block_ids=[0, 1],
ready=asyncio.Event(),
expire_time=-1.0, # Already expired.
sending=0,
),
}
worker.finished_sending_reqs = set()
asyncio.run(worker.fetch_finished_sending_reqs())
assert worker.xfer_stats.data["num_kv_expired_reqs"] == [1]
# Expired transfer also cleaned out of reqs_need_send.
assert "tid1" not in worker.reqs_need_send
+35
View File
@@ -185,6 +185,35 @@ def _xpu_ops_deepseek_scaling_rope_fake(
return query, key
def _topk_topp_sample_impl(
random_sampled: torch.Tensor,
logits_to_return: torch.Tensor | None,
logits: torch.Tensor,
k: torch.Tensor | None,
p: torch.Tensor | None,
logprobs_mode: str,
seeds: torch.Tensor | None,
lambda_: float = 1.0,
) -> None:
torch.ops._xpu_C.topk_topp_sampler(
random_sampled, logits_to_return, logits, k, p, logprobs_mode, seeds, lambda_
)
return
def _topk_topp_sample_fake(
random_sampled: torch.Tensor,
logits_to_return: torch.Tensor | None,
logits: torch.Tensor,
k: torch.Tensor | None,
p: torch.Tensor | None,
logprobs_mode: str,
seeds: torch.Tensor | None,
lambda_: float = 1.0,
) -> None:
return
def _xpu_mxfp8_quantize_impl(
x: torch.Tensor, dtype: torch.dtype | None = None
) -> tuple[torch.Tensor, torch.Tensor]:
@@ -691,6 +720,12 @@ class xpu_ops:
fake_impl=_gdn_attention_core_xpu_fake,
)
direct_register_custom_op(
op_name="xpu_topk_topp_sampler",
op_func=_topk_topp_sample_impl,
fake_impl=_topk_topp_sample_fake,
)
_OPS_REGISTERED = True
+4 -2
View File
@@ -135,8 +135,10 @@ class ParallelConfig:
data_parallel_external_lb: bool = False
"""Whether to use "external" DP LB mode. Applies only to online serving
and when data_parallel_size > 0. This is useful for a "one-pod-per-rank"
wide-EP setup in Kubernetes. Set implicitly when --data-parallel-rank
is provided explicitly to vllm serve."""
wide-EP setup in Kubernetes. Supported only for MoE deployments; non-MoE
models should use independent vLLM instances without --data-parallel-*
arguments. Set implicitly when --data-parallel-rank is provided explicitly
to vllm serve."""
data_parallel_hybrid_lb: bool = False
"""Whether to use "hybrid" DP LB mode. Applies only to online serving
and when data_parallel_size > 0. Enables running an AsyncLLM
@@ -31,10 +31,14 @@ from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorRole,
SupportsHMA,
)
from vllm.distributed.kv_transfer.kv_connector.v1.metrics import KVConnectorStats
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import (
MooncakeBootstrapServer,
RegisterWorkerPayload,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.stats import (
MooncakeKVConnectorStats,
)
from vllm.distributed.parallel_state import (
get_pp_group,
get_tensor_model_parallel_rank,
@@ -457,6 +461,25 @@ class MooncakeConnector(KVConnectorBase_V1, SupportsHMA):
def wait_for_save(self):
pass
def get_kv_connector_stats(self) -> KVConnectorStats | None:
"""Return worker-local transfer stats since the last call.
Note the P/D asymmetry: because Mooncake is P-push (P calls
batch_transfer_sync_write), P records successful transfer latency,
bytes, and descriptor counts, while D only records failures
(recv/ZMQ errors). Aggregated NIXL-style dashboards will find
successful-transfer metrics on the P worker, not D.
"""
if self.connector_worker is None:
return None
return self.connector_worker.get_kv_connector_stats()
@classmethod
def build_kv_connector_stats(
cls, data: dict[str, Any] | None = None
) -> KVConnectorStats | None:
return MooncakeKVConnectorStats(data=data or {})
class MooncakeConnectorScheduler:
"""Implementation of Scheduler side methods"""
@@ -816,6 +839,8 @@ class MooncakeConnectorWorker:
self.finished_sending_reqs: set[ReqId] = set()
self.finished_recving_reqs: set[ReqId] = set()
self.xfer_stats = MooncakeKVConnectorStats()
self.block_size = vllm_config.cache_config.block_size
self.model_config = vllm_config.model_config
self.cache_config = vllm_config.cache_config
@@ -1340,11 +1365,23 @@ class MooncakeConnectorWorker:
ret_value = self.engine.batch_transfer_sync_write(
remote_session, src_ptrs, dst_ptrs, lengths
)
duration = time.perf_counter() - start_time
if ret_value == 0:
logger.debug(
"Sending to %s done, took %s",
self.xfer_stats.record_transfer(
duration_s=duration,
total_bytes=sum(lengths),
num_descs=len(src_ptrs),
)
logger.debug("Sending to %s done, took %s", remote_session, duration)
else:
self.xfer_stats.record_failed_transfer()
logger.warning(
"Sending to %s failed (ret=%s) after %s (%d descriptors, %d bytes)",
remote_session,
time.perf_counter() - start_time,
ret_value,
duration,
len(src_ptrs),
sum(lengths),
)
return ret_value
@@ -1445,6 +1482,7 @@ class MooncakeConnectorWorker:
send_meta.p_req_id,
envs.VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT,
)
self.xfer_stats.record_kv_expired_req()
finished_sending_reqs.add(send_meta.p_req_id)
expired_transfer_id.append(transfer_id)
@@ -1485,6 +1523,13 @@ class MooncakeConnectorWorker:
return finished_sending_reqs or None, finished_recving_reqs or None
def get_kv_connector_stats(self) -> KVConnectorStats | None:
"""Return transfer stats collected since the last call, or None
if nothing has been recorded in this interval."""
if self.xfer_stats.is_empty():
return None
return self.xfer_stats.clone_and_reset()
async def receive_kv_from_single_worker(
self,
worker_addr: str,
@@ -1531,6 +1576,7 @@ class MooncakeConnectorWorker:
req_ids,
response.err_msg,
)
self.xfer_stats.record_failed_recv()
return
self.process_pulling_result(response, pull_metas)
if response.status == MooncakeXferResponseStatus.FINISH:
@@ -1539,6 +1585,7 @@ class MooncakeConnectorWorker:
logger.debug("ZMQ context terminated, exiting Mooncake receiver thread.")
except Exception as e:
logger.error("MooncakeXferMetadata transfer failed for %s: %s", req_ids, e)
self.xfer_stats.record_failed_recv()
return
def process_pulling_result(
@@ -0,0 +1,146 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Stats container for the Mooncake connector."""
import threading
from dataclasses import dataclass
from typing import Any
import numpy as np
from vllm.distributed.kv_transfer.kv_connector.v1.metrics import (
KVConnectorStats,
)
# TODO(mooncake-stats): add MooncakePromMetrics (mirror NixlPromMetrics)
# and wire it via MooncakeConnector.build_prom_metrics in a follow-up PR.
@dataclass
class MooncakeKVConnectorStats(KVConnectorStats):
"""Container for Mooncake KV transfer performance metrics.
`_lock` serializes record_* against clone_and_reset so each row's
appends are atomic and column lengths stay aligned. Writers run on
the sender pool / receiver loop / sender loop; reader runs on the
main worker thread.
"""
def __post_init__(self):
self._lock = threading.Lock()
if not self.data:
self.reset()
# threading.Lock is not picklable; strip it from the wire form and
# rebuild a fresh per-process lock on the receiver side.
def __getstate__(self) -> dict[str, Any]:
state = self.__dict__.copy()
state.pop("_lock", None)
return state
def __setstate__(self, state: dict[str, Any]) -> None:
self.__dict__.update(state)
self._lock = threading.Lock()
def reset(self):
self.data: dict[str, list[float | int]] = {
"transfer_duration": [],
"bytes_transferred": [],
"num_descriptors": [],
"num_failed_transfers": [],
"num_failed_recvs": [],
"num_kv_expired_reqs": [],
}
def record_transfer(self, duration_s: float, total_bytes: int, num_descs: int):
with self._lock:
self.data["transfer_duration"].append(duration_s)
self.data["bytes_transferred"].append(total_bytes)
self.data["num_descriptors"].append(num_descs)
# Failure counters store a list of 1s so a future Prom counter can iterate
# with .inc(list_item), mirroring NIXL's NixlPromMetrics.observe.
def record_failed_transfer(self):
with self._lock:
self.data["num_failed_transfers"].append(1)
def record_failed_recv(self):
with self._lock:
self.data["num_failed_recvs"].append(1)
def record_kv_expired_req(self):
with self._lock:
self.data["num_kv_expired_reqs"].append(1)
def clone_and_reset(self) -> "MooncakeKVConnectorStats":
# Copy lists under the lock for length alignment; return a fresh
# instance so the snapshot has its own _lock.
with self._lock:
snapshot_data: dict[str, list[float | int]] = {
k: list(v) for k, v in self.data.items()
}
self.reset()
return MooncakeKVConnectorStats(data=snapshot_data)
def is_empty(self) -> bool:
return (
self.num_successful_transfers == 0
and len(self.data["num_failed_transfers"]) == 0
and len(self.data["num_failed_recvs"]) == 0
and len(self.data["num_kv_expired_reqs"]) == 0
)
def aggregate(self, other: KVConnectorStats) -> KVConnectorStats:
if not other.is_empty():
for k, v in other.data.items():
accumulator = self.data[k]
assert isinstance(accumulator, list)
accumulator.extend(v)
return self
def reduce(self) -> dict[str, int | float]:
num_failed_transfers = len(self.data["num_failed_transfers"])
num_failed_recvs = len(self.data["num_failed_recvs"])
num_kv_expired_reqs = len(self.data["num_kv_expired_reqs"])
if self.num_successful_transfers == 0:
return {
"Num successful transfers": 0,
"Avg xfer time (ms)": 0,
"P90 xfer time (ms)": 0,
"Avg MB per transfer": 0,
"Throughput (MB/s)": 0,
"Avg number of descriptors": 0,
"Num failed transfers": num_failed_transfers,
"Num failed recvs": num_failed_recvs,
"Num KV expired reqs": num_kv_expired_reqs,
}
xfer_time = np.asarray(self.data["transfer_duration"])
mb = np.asarray(self.data["bytes_transferred"]) / 2**20
descs = np.asarray(self.data["num_descriptors"], dtype=np.uint32)
n = len(descs)
assert n == self.num_successful_transfers
total_mb = mb.sum()
avg_mb = total_mb / n
total_time_seconds = xfer_time.sum()
throughput_mb_s = (
total_mb / total_time_seconds if total_time_seconds > 0 else 0.0
)
return {
"Num successful transfers": n,
"Avg xfer time (ms)": round(xfer_time.mean() * 1e3, 3),
"P90 xfer time (ms)": round(np.percentile(xfer_time, 90).item() * 1e3, 3),
"Avg MB per transfer": round(avg_mb, 3),
"Throughput (MB/s)": round(throughput_mb_s, 3),
"Avg number of descriptors": round(descs.mean(), 1),
"Num failed transfers": num_failed_transfers,
"Num failed recvs": num_failed_recvs,
"Num KV expired reqs": num_kv_expired_reqs,
}
@property
def num_successful_transfers(self) -> int:
return len(self.data["transfer_duration"])
+16 -18
View File
@@ -962,7 +962,9 @@ class EngineArgs:
"-dpn",
type=int,
help="Data parallel rank of this instance. "
"When set, enables external load balancer mode.",
"When set, enables external load balancer mode for MoE "
"data-parallel deployments. Unsupported for non-MoE models; "
"launch independent vLLM instances instead.",
)
parallel_group.add_argument(
"--data-parallel-start-rank",
@@ -1697,29 +1699,15 @@ class EngineArgs:
kv_offloading_backend=self.kv_offloading_backend,
)
# TurboQuant: auto-skip first/last 2 layers (boundary protection).
# These layers are most sensitive to quantization error.
# Users can add extra layers via --kv-cache-dtype-skip-layers.
if resolved_cache_dtype.startswith("turboquant_"):
if model_config.is_hybrid:
raise NotImplementedError(
"TurboQuant KV cache is not supported for hybrid "
"(attention + Mamba) models. Boundary layer protection "
"requires uniform attention layers."
)
from vllm.model_executor.layers.quantization.turboquant.config import (
TurboQuantConfig,
)
num_layers = model_config.hf_text_config.num_hidden_layers
boundary = TurboQuantConfig.get_boundary_skip_layers(num_layers)
boundary = TurboQuantConfig.get_boundary_skip_layers(model_config)
existing = set(cache_config.kv_cache_dtype_skip_layers)
merged = sorted(existing | set(boundary), key=lambda x: int(x))
cache_config.kv_cache_dtype_skip_layers = merged
logger.info(
"TQ: skipping layers %s for boundary protection (num_layers=%d)",
merged,
num_layers,
cache_config.kv_cache_dtype_skip_layers = sorted(
existing | set(boundary), key=int
)
ray_runtime_env = None
@@ -1793,6 +1781,16 @@ class EngineArgs:
data_parallel_external_lb = (
self.data_parallel_external_lb or self.data_parallel_rank is not None
)
if (
self.data_parallel_size > 1
and data_parallel_external_lb
and not model_config.is_moe
):
raise ValueError(
"Non-MoE models do not support external data parallel mode. "
"For external load balancing, launch independent vLLM "
"instances without --data-parallel-* arguments."
)
# Local DP rank = 1, use pure-external LB.
if data_parallel_external_lb:
assert self.data_parallel_rank is not None, (
+6 -2
View File
@@ -266,6 +266,7 @@ if TYPE_CHECKING:
VLLM_MEMORY_PROFILER_ESTIMATE_CUDAGRAPHS: bool = True
VLLM_NIXL_EP_MAX_NUM_RANKS: int = 32
VLLM_XPU_ENABLE_XPU_GRAPH: bool = False
VLLM_XPU_USE_SAMPLER_KERNEL: bool = True
VLLM_LORA_ENABLE_DUAL_STREAM: bool = False
@@ -782,9 +783,8 @@ environment_variables: dict[str, Callable[[], Any]] = {
),
# When True and distributed_executor_backend="ray", use RayExecutorV2
# (MQ-based) instead of RayDistributedExecutor (compiled-graph backend).
# TODO (jeffreywang): Enabled by default in vLLM 0.20.0.
"VLLM_USE_RAY_V2_EXECUTOR_BACKEND": lambda: bool(
int(os.getenv("VLLM_USE_RAY_V2_EXECUTOR_BACKEND", "0"))
int(os.getenv("VLLM_USE_RAY_V2_EXECUTOR_BACKEND", "1"))
),
# Use dedicated multiprocess context for workers.
# Both spawn and fork work
@@ -1776,6 +1776,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
"VLLM_XPU_ENABLE_XPU_GRAPH": lambda: bool(
int(os.getenv("VLLM_XPU_ENABLE_XPU_GRAPH", "0"))
),
# whether use xpu specific sample kernel
"VLLM_XPU_USE_SAMPLER_KERNEL": lambda: bool(
int(os.getenv("VLLM_XPU_USE_SAMPLER_KERNEL", "1"))
),
# Enable simple KV offload.
"VLLM_USE_SIMPLE_KV_OFFLOAD": lambda: bool(
int(os.getenv("VLLM_USE_SIMPLE_KV_OFFLOAD", "0"))
@@ -238,6 +238,9 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import (
kFp8StaticTensorSym,
kNvfp4Dynamic,
)
from vllm.model_executor.layers.rotary_embedding.deepseek_scaling_rope import (
yarn_get_mscale,
)
from vllm.platforms import current_platform
from vllm.utils.flashinfer import has_flashinfer
from vllm.utils.math_utils import cdiv, round_down
@@ -1327,6 +1330,35 @@ def get_mla_dims(model_config: ModelConfig) -> MLADims:
)
def get_mla_prefill_scale(model_config: ModelConfig) -> float:
hf_text_config = model_config.hf_text_config
mla_dims = get_mla_dims(model_config)
qk_head_dim = mla_dims.qk_nope_head_dim + mla_dims.qk_rope_head_dim
scale = qk_head_dim**-0.5
# Deepseek V4 disables YaRN mscale for attention; Deepseek V2/V3 applies
# the same mscale correction when constructing the MLA attention module.
if hasattr(hf_text_config, "compress_ratios"):
return scale
rope_parameters = getattr(hf_text_config, "rope_parameters", None)
if rope_parameters is None:
rope_parameters = getattr(hf_text_config, "rope_scaling", None)
if rope_parameters is None:
return scale
rope_type = rope_parameters.get("rope_type", rope_parameters.get("type"))
apply_yarn_scaling = rope_parameters.get("apply_yarn_scaling", True)
if rope_type != "default" and apply_yarn_scaling:
mscale_all_dim = rope_parameters.get("mscale_all_dim", False)
scaling_factor = rope_parameters["factor"]
mscale = yarn_get_mscale(float(scaling_factor), float(mscale_all_dim))
scale *= mscale * mscale
return scale
@functools.cache
def backend_supports_prefill_query_quantization() -> bool:
"""Check if the selected MLA prefill backend supports query quantization.
@@ -1527,7 +1559,7 @@ class MLACommonMetadataBuilder(AttentionMetadataBuilder[M]):
prefill_backend_cls = get_mla_prefill_backend(vllm_config)
self._prefill_backend = prefill_backend_cls(
num_heads=self.num_heads,
scale=self.model_config.get_head_size() ** -0.5,
scale=get_mla_prefill_scale(self.model_config),
kv_lora_rank=self.mla_dims.kv_lora_rank,
qk_nope_head_dim=self.mla_dims.qk_nope_head_dim,
qk_rope_head_dim=self.mla_dims.qk_rope_head_dim,
@@ -15,6 +15,7 @@ class MoEActivation(Enum):
# and produce output of shape [..., d]
SILU = "silu"
GELU = "gelu"
GELU_TANH = "gelu_tanh"
RELU2 = "relu2"
SWIGLUOAI = "swigluoai"
SWIGLUSTEP = "swiglustep"
@@ -24,6 +25,7 @@ class MoEActivation(Enum):
# NOTE: Non-gated activations require the "_no_mul" suffix to be present.
SILU_NO_MUL = "silu_no_mul"
GELU_NO_MUL = "gelu_no_mul"
GELU_TANH_NO_MUL = "gelu_tanh_no_mul"
RELU2_NO_MUL = "relu2_no_mul"
@property
@@ -53,6 +55,7 @@ class MoEActivation(Enum):
@classmethod
def from_str(cls, s: str) -> "MoEActivation":
"""Parse from string for backward compatibility."""
s = _STR_ALIASES.get(s, s)
for member in cls:
if member.value == s:
return member
@@ -61,20 +64,27 @@ class MoEActivation(Enum):
# Module-level lookup tables used by MoEActivation functions.
_STR_ALIASES: dict[str, str] = {
"gelu_pytorch_tanh": "gelu_tanh",
}
_CUSTOM_OP_NAMES: dict[MoEActivation, str] = {
MoEActivation.SILU: "silu_and_mul",
MoEActivation.GELU: "gelu_and_mul",
MoEActivation.GELU_TANH: "gelu_tanh_and_mul",
MoEActivation.SWIGLUOAI: "swigluoai_and_mul",
MoEActivation.SWIGLUSTEP: "swiglustep_and_mul",
MoEActivation.RELU2: "relu2",
MoEActivation.SILU_NO_MUL: "silu_and_mul",
MoEActivation.GELU_NO_MUL: "gelu_and_mul",
MoEActivation.GELU_TANH_NO_MUL: "gelu_tanh_and_mul",
MoEActivation.RELU2_NO_MUL: "relu2",
}
_WITHOUT_MUL: dict[MoEActivation, MoEActivation] = {
MoEActivation.SILU: MoEActivation.SILU_NO_MUL,
MoEActivation.GELU: MoEActivation.GELU_NO_MUL,
MoEActivation.GELU_TANH: MoEActivation.GELU_TANH_NO_MUL,
MoEActivation.RELU2: MoEActivation.RELU2_NO_MUL,
}
@@ -115,6 +125,8 @@ def apply_moe_activation(
torch.ops._C.silu_and_mul(output, input)
elif activation == MoEActivation.GELU:
torch.ops._C.gelu_and_mul(output, input)
elif activation == MoEActivation.GELU_TANH:
torch.ops._C.gelu_tanh_and_mul(output, input)
elif activation == MoEActivation.SWIGLUOAI:
torch.ops._C.swigluoai_and_mul(output, input)
elif activation == MoEActivation.SWIGLUSTEP:
@@ -127,6 +139,8 @@ def apply_moe_activation(
output.copy_(F.silu(input))
elif activation == MoEActivation.GELU_NO_MUL:
output.copy_(F.gelu(input))
elif activation == MoEActivation.GELU_TANH_NO_MUL:
output.copy_(F.gelu(input, approximate="tanh"))
elif activation == MoEActivation.RELU2_NO_MUL:
F.relu(input, inplace=True)
torch.square(input, out=output)
@@ -0,0 +1,292 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
from vllm._aiter_ops import rocm_aiter_ops
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEConfig,
FusedMoEParallelConfig,
FusedMoEQuantConfig,
RoutingMethodType,
)
from vllm.model_executor.layers.quantization.utils.quant_utils import (
QuantKey,
kFp8StaticTensorSym,
kMxfp4Static,
)
__all__ = [
"AiterW4A8ExpertsMonolithic",
"aiter_triton_kernel_w4a8_moe_forward",
]
def aiter_triton_kernel_w4a8_moe_forward(
hidden_states: torch.Tensor,
w1, # Tensor or triton_kernels.Tensor
w2, # Tensor or triton_kernels.Tensor
gating_output: torch.Tensor,
topk: int,
renormalize: bool,
activation: MoEActivation = MoEActivation.SWIGLUOAI,
quant_config: FusedMoEQuantConfig | None = None,
apply_router_weight_on_input: bool = False,
global_num_experts: int = -1,
expert_map: torch.Tensor | None = None,
unpadded_N_w1=None,
unpadded_K_w1=None,
unpadded_N_w2=None,
unpadded_K_w2=None,
):
assert (
quant_config is not None
and quant_config.use_mxfp4_w4a8
and rocm_aiter_ops.is_enabled()
)
from aiter.ops.triton.moe_routing.routing import routing as aiter_routing
routing_data, gather_idx, scatter_idx = aiter_routing(
gating_output, topk, sm_first=not renormalize
)
return triton_kernel_fused_mxfp4_w4a8_experts(
None,
hidden_states,
w1,
w2,
routing_data,
gather_idx,
scatter_idx,
activation=activation.value,
quant_config=quant_config,
apply_router_weight_on_input=apply_router_weight_on_input,
global_num_experts=global_num_experts,
expert_map=expert_map,
unpadded_N_w1=unpadded_N_w1,
unpadded_K_w1=unpadded_K_w1,
unpadded_N_w2=unpadded_N_w2,
unpadded_K_w2=unpadded_K_w2,
)
def triton_kernel_fused_mxfp4_w4a8_experts(
output_tensor: torch.Tensor,
hidden_states: torch.Tensor,
w1, # Tensor or triton_kernels.Tensor
w2, # Tensor or triton_kernels.Tensor
routing_data, # RoutingData
gather_indx, # GatherIndx
scatter_indx, # ScatterIndx
activation: str = "silu",
quant_config: FusedMoEQuantConfig | None = None,
swiglu_alpha: float = 1.702,
swiglu_limit: float = 7.0,
apply_router_weight_on_input: bool = False,
global_num_experts: int = -1,
expert_map: torch.Tensor | None = None,
a1q_scale: torch.Tensor | None = None,
unpadded_N_w1=None,
unpadded_K_w1=None,
unpadded_N_w2=None,
unpadded_K_w2=None,
) -> torch.Tensor:
assert quant_config is not None
# type check, uint8 means mxfp4
assert hidden_states.dtype == torch.bfloat16
assert quant_config.w1_bias is None or quant_config.w1_bias.dtype == torch.float32
assert quant_config.w2_bias is None or quant_config.w2_bias.dtype == torch.float32
# Shape check: weights are padded (e.g. hidden_size padded for
# GFX950 swizzle).
assert hidden_states.shape[-1] == w1.shape[-2]
assert w2.shape[-1] == w1.shape[1]
E, _, N = w1.shape
if global_num_experts == -1:
global_num_experts = E
gammas = routing_data.gate_scal if routing_data else None
from aiter.ops.triton.moe_op_gemm_a8w4 import moe_gemm_a8w4
from aiter.ops.triton.quant_moe import downcast_to_static_fp8
assert quant_config.w1_precision is not None, (
"w1_precision in quant config can't be None"
)
assert quant_config.w2_precision is not None, (
"w2_precision in quant config can't be None"
)
hidden_states = downcast_to_static_fp8(
hidden_states, quant_config.w1_precision.flex_ctx.lhs_data.scale
)
intermediate_cache1 = moe_gemm_a8w4(
hidden_states,
w1.storage.data,
None,
quant_config.w1_precision.weight_scale.storage.data,
quant_config.w1_precision.flex_ctx.lhs_data.scale,
quant_config.w2_precision.flex_ctx.lhs_data.scale,
quant_config.w1_bias,
routing_data,
gather_indx=gather_indx,
gammas=gammas if apply_router_weight_on_input else None,
swizzle_mx_scale="CDNA4_SCALE",
out_dtype=torch.float8_e4m3fn,
apply_swiglu=True,
alpha=swiglu_alpha,
limit=swiglu_limit,
unpadded_N=unpadded_N_w1,
unpadded_K=unpadded_K_w1,
)
intermediate_cache3 = moe_gemm_a8w4(
intermediate_cache1,
w2.storage.data,
None,
quant_config.w2_precision.weight_scale.storage.data,
quant_config.w2_precision.flex_ctx.lhs_data.scale,
None,
quant_config.w2_bias,
routing_data,
scatter_indx=scatter_indx,
gammas=None if apply_router_weight_on_input else gammas,
swizzle_mx_scale="CDNA4_SCALE",
unpadded_N=unpadded_N_w2,
unpadded_K=unpadded_K_w2,
)
return intermediate_cache3
class AiterW4A8ExpertsMonolithic(mk.FusedMoEExpertsMonolithic):
"""
Monolithic MXFP4 W4A8 expert using AITER triton kernels.
This backend uses:
- aiter.ops.triton.moe_routing.routing for routing
- aiter.ops.triton.moe_op_gemm_a8w4.moe_gemm_a8w4 for computation
Weight format: MXFP4 weights with GFX950 swizzle
Activation: Static FP8 quantization
"""
def __init__(
self,
moe_config: FusedMoEConfig,
quant_config: FusedMoEQuantConfig,
):
super().__init__(moe_config, quant_config)
self.topk = moe_config.experts_per_token
self.renormalize = moe_config.routing_method in (
RoutingMethodType.Renormalize,
RoutingMethodType.RenormalizeNaive,
)
@staticmethod
def activation_format() -> mk.FusedMoEActivationFormat:
return mk.FusedMoEActivationFormat.Standard
@staticmethod
def _supports_current_device() -> bool:
# Requires AITER and GFX950
if not rocm_aiter_ops.is_enabled():
return False
from vllm.platforms.rocm import on_gfx950
return on_gfx950()
@staticmethod
def _supports_no_act_and_mul() -> bool:
return False
@staticmethod
def _supports_quant_scheme(
weight_key: QuantKey | None,
activation_key: QuantKey | None,
) -> bool:
# W4A8: MXFP4 weights with static FP8 activations
SUPPORTED_W_A = [
(kMxfp4Static, kFp8StaticTensorSym),
]
return (weight_key, activation_key) in SUPPORTED_W_A
@staticmethod
def _supports_activation(activation: MoEActivation) -> bool:
# Only SILU activation (swiglu) is supported
return activation == MoEActivation.SWIGLUOAI
@staticmethod
def _supports_parallel_config(
moe_parallel_config: FusedMoEParallelConfig,
) -> bool:
return (
not moe_parallel_config.use_all2all_kernels
and not moe_parallel_config.enable_eplb
and moe_parallel_config.dp_size <= 1
)
@staticmethod
def _supports_routing_method(
routing_method: RoutingMethodType,
weight_key: QuantKey | None,
activation_key: QuantKey | None,
) -> bool:
return routing_method in [
RoutingMethodType.Renormalize,
RoutingMethodType.RenormalizeNaive,
]
@staticmethod
def _supports_router_logits_dtype(
router_logits_dtype: torch.dtype | None,
routing_method: RoutingMethodType,
) -> bool:
return True
def supports_expert_map(self) -> bool:
return False # Expert parallelism not yet supported
@property
def expects_unquantized_inputs(self) -> bool:
return True
def apply(
self,
hidden_states: torch.Tensor,
w1: torch.Tensor,
w2: torch.Tensor,
router_logits: torch.Tensor,
activation: MoEActivation,
global_num_experts: int,
expert_map: torch.Tensor | None,
a1q_scale: torch.Tensor | None,
apply_router_weight_on_input: bool,
# grouped topk + fused topk bias parameters
num_expert_group: int | None = None,
e_score_correction_bias: torch.Tensor | None = None,
routed_scaling_factor: float | None = None,
topk_group: int | None = None,
) -> torch.Tensor:
assert self.moe_config.intermediate_size_per_partition_unpadded is not None
assert self.moe_config.hidden_dim_unpadded is not None
return aiter_triton_kernel_w4a8_moe_forward(
hidden_states=hidden_states,
w1=w1,
w2=w2,
gating_output=router_logits,
topk=self.topk,
renormalize=self.renormalize,
global_num_experts=global_num_experts,
expert_map=expert_map,
quant_config=self.quant_config,
apply_router_weight_on_input=apply_router_weight_on_input,
unpadded_N_w1=self.moe_config.intermediate_size_per_partition_unpadded * 2,
unpadded_K_w1=self.moe_config.hidden_dim_unpadded,
unpadded_N_w2=self.moe_config.hidden_dim_unpadded,
unpadded_K_w2=self.moe_config.intermediate_size_per_partition_unpadded,
)
@@ -5,7 +5,6 @@ import torch
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
from vllm import _custom_ops as ops
from vllm._aiter_ops import rocm_aiter_ops
from vllm.logger import init_logger
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.config import (
@@ -286,35 +285,6 @@ def triton_kernel_moe_forward(
unpadded_N_w2=None,
unpadded_K_w2=None,
) -> torch.Tensor:
if (
quant_config is not None
and quant_config.use_mxfp4_w4a8
and rocm_aiter_ops.is_enabled()
):
from aiter.ops.triton.moe_routing.routing import routing as aiter_routing
routing_data, gather_idx, scatter_idx = aiter_routing(
gating_output, topk, sm_first=not renormalize
)
return triton_kernel_fused_mxfp4_w4a8_experts(
None,
hidden_states,
w1,
w2,
routing_data,
gather_idx,
scatter_idx,
activation=activation.value,
quant_config=quant_config,
apply_router_weight_on_input=apply_router_weight_on_input,
global_num_experts=global_num_experts,
expert_map=expert_map,
unpadded_N_w1=unpadded_N_w1,
unpadded_K_w1=unpadded_K_w1,
unpadded_N_w2=unpadded_N_w2,
unpadded_K_w2=unpadded_K_w2,
)
from triton_kernels.topk import topk as topk_fn
sm_first = not renormalize
@@ -471,99 +441,6 @@ def triton_kernel_fused_experts(
return output_tensor
# This is a triton implementation of the fused_experts function
def triton_kernel_fused_mxfp4_w4a8_experts(
output_tensor: torch.Tensor,
hidden_states: torch.Tensor,
w1, # Tensor or triton_kernels.Tensor
w2, # Tensor or triton_kernels.Tensor
routing_data, # RoutingData
gather_indx, # GatherIndx
scatter_indx, # ScatterIndx
activation: str = "silu",
quant_config: FusedMoEQuantConfig | None = None,
swiglu_alpha: float = 1.702,
swiglu_limit: float = 7.0,
apply_router_weight_on_input: bool = False,
global_num_experts: int = -1,
expert_map: torch.Tensor | None = None,
a1q_scale: torch.Tensor | None = None,
unpadded_N_w1=None,
unpadded_K_w1=None,
unpadded_N_w2=None,
unpadded_K_w2=None,
) -> torch.Tensor:
assert quant_config is not None
# type check, uint8 means mxfp4
assert hidden_states.dtype == torch.bfloat16
assert quant_config.w1_bias is None or quant_config.w1_bias.dtype == torch.float32
assert quant_config.w2_bias is None or quant_config.w2_bias.dtype == torch.float32
# Shape check: weights are padded (e.g. hidden_size padded for
# GFX950 swizzle).
assert hidden_states.shape[-1] == w1.shape[-2]
assert w2.shape[-1] == w1.shape[1]
E, _, N = w1.shape
if global_num_experts == -1:
global_num_experts = E
gammas = routing_data.gate_scal if routing_data else None
from aiter.ops.triton.moe_op_gemm_a8w4 import moe_gemm_a8w4
from aiter.ops.triton.quant_moe import downcast_to_static_fp8
assert quant_config.w1_precision is not None, (
"w1_precision in quant config can't be None"
)
assert quant_config.w2_precision is not None, (
"w2_precision in quant config can't be None"
)
hidden_states = downcast_to_static_fp8(
hidden_states, quant_config.w1_precision.flex_ctx.lhs_data.scale
)
intermediate_cache1 = moe_gemm_a8w4(
hidden_states,
w1.storage.data,
None,
quant_config.w1_precision.weight_scale.storage.data,
quant_config.w1_precision.flex_ctx.lhs_data.scale,
quant_config.w2_precision.flex_ctx.lhs_data.scale,
quant_config.w1_bias,
routing_data,
gather_indx=gather_indx,
gammas=gammas if apply_router_weight_on_input else None,
swizzle_mx_scale="CDNA4_SCALE",
out_dtype=torch.float8_e4m3fn,
apply_swiglu=True,
alpha=swiglu_alpha,
limit=swiglu_limit,
unpadded_N=unpadded_N_w1,
unpadded_K=unpadded_K_w1,
)
intermediate_cache3 = moe_gemm_a8w4(
intermediate_cache1,
w2.storage.data,
None,
quant_config.w2_precision.weight_scale.storage.data,
quant_config.w2_precision.flex_ctx.lhs_data.scale,
None,
quant_config.w2_bias,
routing_data,
scatter_indx=scatter_indx,
gammas=None if apply_router_weight_on_input else gammas,
swizzle_mx_scale="CDNA4_SCALE",
unpadded_N=unpadded_N_w2,
unpadded_K=unpadded_K_w2,
)
return intermediate_cache3
def make_routing_data(
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
@@ -62,7 +62,7 @@ class XPUExperts(mk.FusedMoEExpertsModular):
@staticmethod
def _supports_no_act_and_mul() -> bool:
return False
return True
@staticmethod
def _supports_activation(activation: MoEActivation) -> bool:
@@ -70,6 +70,7 @@ class XPUExperts(mk.FusedMoEExpertsModular):
MoEActivation.SILU,
MoEActivation.GELU,
MoEActivation.SWIGLUOAI,
MoEActivation.RELU2_NO_MUL,
]
@staticmethod
@@ -786,9 +786,11 @@ class BatchedTritonExperts(mk.FusedMoEExpertsModular):
return activation in [
MoEActivation.SILU,
MoEActivation.GELU,
MoEActivation.GELU_TANH,
MoEActivation.SWIGLUOAI,
MoEActivation.SILU_NO_MUL,
MoEActivation.GELU_NO_MUL,
MoEActivation.GELU_TANH_NO_MUL,
MoEActivation.RELU2_NO_MUL,
]
@@ -152,10 +152,12 @@ class HummingExpertsBase(mk.FusedMoEExpertsModular):
return activation in [
MoEActivation.SILU,
MoEActivation.GELU,
MoEActivation.GELU_TANH,
MoEActivation.SWIGLUOAI,
MoEActivation.SWIGLUSTEP,
MoEActivation.SILU_NO_MUL,
MoEActivation.GELU_NO_MUL,
MoEActivation.GELU_TANH_NO_MUL,
MoEActivation.RELU2_NO_MUL,
]
@@ -613,10 +613,12 @@ class MarlinExpertsBase(mk.FusedMoEExpertsModular):
return activation in [
MoEActivation.SILU,
MoEActivation.GELU,
MoEActivation.GELU_TANH,
MoEActivation.SWIGLUOAI,
MoEActivation.SWIGLUSTEP,
MoEActivation.SILU_NO_MUL,
MoEActivation.GELU_NO_MUL,
MoEActivation.GELU_TANH_NO_MUL,
MoEActivation.RELU2_NO_MUL,
]
@@ -1941,10 +1941,12 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
return activation in [
MoEActivation.SILU,
MoEActivation.GELU,
MoEActivation.GELU_TANH,
MoEActivation.SWIGLUOAI,
MoEActivation.SWIGLUSTEP,
MoEActivation.SILU_NO_MUL,
MoEActivation.GELU_NO_MUL,
MoEActivation.GELU_TANH_NO_MUL,
MoEActivation.RELU2_NO_MUL,
]
@@ -538,9 +538,11 @@ class FusedMoE(PluggableLayer):
# for heuristic purposes, so it must be initialized first.
self.quant_method: FusedMoEMethodBase = _get_quant_method()
if not self.moe_config.is_act_and_mul and not current_platform.is_cuda_alike():
if not self.moe_config.is_act_and_mul and not (
current_platform.is_cuda_alike() or current_platform.is_xpu()
):
raise NotImplementedError(
"is_act_and_mul=False is supported only for CUDA and ROCm for now"
"is_act_and_mul=False is supported only for CUDA and XPU for now"
)
if self.enable_eplb and not self.quant_method.supports_eplb:
@@ -19,6 +19,7 @@ from vllm.model_executor.layers.fused_moe.config import (
FusedMoEQuantConfig,
FusedMoEQuantDesc,
mxfp4_mxfp8_moe_quant_config,
mxfp4_w4a8_moe_quant_config,
mxfp4_w4a16_moe_quant_config,
ocp_mx_moe_quant_config,
)
@@ -26,9 +27,11 @@ from vllm.model_executor.layers.quantization.utils.mxfp4_utils import _swizzle_m
from vllm.model_executor.layers.quantization.utils.quant_utils import (
QuantKey,
kFp8Dynamic128Sym,
kFp8StaticTensorSym,
kMxfp4Static,
kMxfp8Dynamic,
)
from vllm.model_executor.layers.quantization.utils.w8a8_utils import all_close_1d
from vllm.platforms import current_platform
from vllm.utils.import_utils import has_triton_kernels
from vllm.utils.math_utils import round_up
@@ -59,8 +62,9 @@ class Mxfp4MoeBackend(Enum):
# Marlin
BATCHED_MARLIN = "BATCHED_MARLIN"
MARLIN = "MARLIN"
# ROCm AITER
AITER = "AITER"
# ROCm AITER backends
AITER_MXFP4_BF16 = "AITER_MXFP4_BF16" # W4A16: CK kernel
AITER_MXFP4_FP8 = "AITER_MXFP4_FP8" # W4A8: triton kernel
# Triton
TRITON = "TRITON"
TRITON_UNFUSED = "TRITON_UNFUSED"
@@ -72,6 +76,13 @@ class Mxfp4MoeBackend(Enum):
HUMMING = "HUMMING"
# AITER backends group
AITER_BACKENDS = (
Mxfp4MoeBackend.AITER_MXFP4_BF16,
Mxfp4MoeBackend.AITER_MXFP4_FP8,
)
# Backends that share the same TRTLLM weight format
TRTLLM_BACKENDS = (
Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_BF16,
@@ -159,13 +170,20 @@ def backend_to_kernel_cls(
return [BatchedMarlinExperts]
elif backend == Mxfp4MoeBackend.AITER:
elif backend == Mxfp4MoeBackend.AITER_MXFP4_BF16:
from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import (
AiterExperts,
)
return [AiterExperts]
elif backend == Mxfp4MoeBackend.AITER_MXFP4_FP8:
from vllm.model_executor.layers.fused_moe.experts.aiter_mxfp4_w4a8_moe import (
AiterW4A8ExpertsMonolithic,
)
return [AiterW4A8ExpertsMonolithic]
elif backend == Mxfp4MoeBackend.XPU:
from vllm.model_executor.layers.fused_moe.experts.xpu_moe import XPUExpertsMXFp4
@@ -194,7 +212,8 @@ def map_mxfp4_backend(runner_backend: MoEBackend) -> Mxfp4MoeBackend:
"triton_unfused": Mxfp4MoeBackend.TRITON_UNFUSED,
"humming": Mxfp4MoeBackend.HUMMING,
"marlin": Mxfp4MoeBackend.MARLIN,
"aiter": Mxfp4MoeBackend.AITER,
"aiter": Mxfp4MoeBackend.AITER_MXFP4_BF16,
"aiter_mxfp4_fp8": Mxfp4MoeBackend.AITER_MXFP4_FP8,
"xpu": Mxfp4MoeBackend.XPU,
"emulation": Mxfp4MoeBackend.EMULATION,
}
@@ -213,7 +232,8 @@ def _get_priority_backends_for_gpt_oss() -> list[Mxfp4MoeBackend]:
"""
_AVAILABLE_BACKENDS = [
Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_BF16,
Mxfp4MoeBackend.AITER,
Mxfp4MoeBackend.AITER_MXFP4_BF16,
Mxfp4MoeBackend.AITER_MXFP4_FP8,
Mxfp4MoeBackend.TRITON,
Mxfp4MoeBackend.FLASHINFER_CUTLASS_MXFP4_BF16,
# TRITON_UNFUSED has bug with MTP support
@@ -254,16 +274,28 @@ def _backend_activation_key(backend: Mxfp4MoeBackend) -> QuantKey | None:
Mxfp4MoeBackend.FLASHINFER_CUTLASS_MXFP4_MXFP8,
):
return kMxfp8Dynamic
return None
if backend == Mxfp4MoeBackend.AITER_MXFP4_FP8:
return kFp8StaticTensorSym
return None # BF16 activation
def select_gpt_oss_mxfp4_moe_backend(
def select_mxfp4_moe_backend(
config: FusedMoEConfig,
activation_key: QuantKey | None = None,
) -> tuple[Mxfp4MoeBackend, type[mk.FusedMoEExperts] | None]:
"""
Select the primary MXFP4 MoE backend.
Args:
config: MoE configuration
activation_key: Optional activation quantization key. If provided,
overrides the default activation key for backend selection.
Use kFp8StaticTensorSym for W4A8 scheme.
Note: Shape-specific fallbacks may still occur at runtime.
"""
# If activation_key is explicitly provided (e.g., W4A8), use it
requested_activation_key = activation_key
device_capability = current_platform.get_device_capability()
triton_kernels_supported = (
has_triton_kernels()
@@ -332,11 +364,17 @@ def select_gpt_oss_mxfp4_moe_backend(
and requested_backend == Mxfp4MoeBackend.MARLIN
):
requested_backend = Mxfp4MoeBackend.BATCHED_MARLIN
# Use requested_activation_key if provided, otherwise use backend default
act_key = (
requested_activation_key
if requested_activation_key is not None
else _backend_activation_key(requested_backend)
)
return _return_or_raise(
requested_backend,
config,
kMxfp4Static,
_backend_activation_key(requested_backend),
act_key,
activation_format,
)
@@ -408,10 +446,15 @@ def select_gpt_oss_mxfp4_moe_backend(
)
for backend in AVAILABLE_BACKENDS:
activation_key = _backend_activation_key(backend)
# Use requested_activation_key if provided, otherwise use backend default
act_key = (
requested_activation_key
if requested_activation_key is not None
else _backend_activation_key(backend)
)
for k_cls in backend_to_kernel_cls(backend):
supported, reason = k_cls.is_supported_config(
k_cls, config, kMxfp4Static, activation_key, activation_format
k_cls, config, kMxfp4Static, act_key, activation_format
)
if supported:
logger.info_once(_make_log_backend(backend))
@@ -438,7 +481,7 @@ def select_gpt_oss_mxfp4_moe_backend(
return Mxfp4MoeBackend.NONE, None
def select_mxfp4_moe_backend(
def select_deepseek_v4_mxfp4_moe_backend(
config: FusedMoEConfig,
) -> tuple[Mxfp4MoeBackend, type[mk.FusedMoEExperts] | None]:
"""
@@ -836,7 +879,7 @@ def convert_gpt_oss_weight_to_mxfp4_moe_kernel_format(
w2_bias,
)
elif mxfp4_backend == Mxfp4MoeBackend.AITER:
elif mxfp4_backend == Mxfp4MoeBackend.AITER_MXFP4_BF16:
from vllm._aiter_ops import rocm_aiter_ops
if w13_bias is not None:
@@ -898,6 +941,63 @@ def convert_gpt_oss_weight_to_mxfp4_moe_kernel_format(
w2_bias,
)
elif mxfp4_backend == Mxfp4MoeBackend.AITER_MXFP4_FP8:
# W4A8: MXFP4 weights + static FP8 activations (triton kernel)
from triton_kernels.matmul_ogs import FlexCtx, PrecisionConfig
from triton_kernels.numerics import InFlexData
if w13_bias is not None:
w13_bias = w13_bias.to(torch.float32)
if w2_bias is not None:
w2_bias = w2_bias.to(torch.float32)
# Process static FP8 input scales (reduce to scalar, warn if not uniform)
w13_input_scale = layer.w13_input_scale
w2_input_scale = layer.w2_input_scale
if w13_input_scale is None or w2_input_scale is None:
raise ValueError(
"W4A8 (AITER_MXFP4_FP8) requires static input scales, but found "
"w13_input_scale or w2_input_scale is None."
)
if not all_close_1d(w13_input_scale) or not all_close_1d(w2_input_scale):
logger.warning_once(
"Found input_scales that are not equal for "
"fp8 MoE layer. Using the maximum across experts "
"for each layer."
)
w13_input_scale = w13_input_scale.max().to(torch.float32)
w2_input_scale = w2_input_scale.max().to(torch.float32)
# Swizzle weights for GFX950
w13_weight, w13_flex, w13_scale = _swizzle_mxfp4(w13_weight, w13_weight_scale)
w2_weight, w2_flex, w2_scale = _swizzle_mxfp4(w2_weight, w2_weight_scale)
# Create InFlexData for activation scales
lhs_data13 = InFlexData(scale=w13_input_scale)
lhs_data2 = InFlexData(scale=w2_input_scale)
# Create PrecisionConfig with both weight and activation info
w13_precision_config = PrecisionConfig(
weight_scale=w13_scale,
flex_ctx=FlexCtx(rhs_data=w13_flex, lhs_data=lhs_data13),
)
w2_precision_config = PrecisionConfig(
weight_scale=w2_scale,
flex_ctx=FlexCtx(rhs_data=w2_flex, lhs_data=lhs_data2),
)
del layer.w13_weight
del layer.w2_weight
return (
w13_weight,
w2_weight,
w13_precision_config,
w2_precision_config,
w13_bias,
w2_bias,
)
elif mxfp4_backend in TRITON_BACKENDS:
from triton_kernels.matmul_ogs import FlexCtx, PrecisionConfig
@@ -1220,6 +1320,8 @@ def make_mxfp4_moe_quant_config(
swiglu_limit: float | None = None,
w1_bias: torch.Tensor | None = None,
w2_bias: torch.Tensor | None = None,
a1_scale: torch.Tensor | None = None,
a2_scale: torch.Tensor | None = None,
layer: torch.nn.Module | None = None,
) -> FusedMoEQuantConfig | None:
"""Create a FusedMoEQuantConfig for the given MXFP4 backend."""
@@ -1262,6 +1364,17 @@ def make_mxfp4_moe_quant_config(
gemm1_beta=gemm1_beta,
gemm1_clamp_limit=swiglu_limit,
)
elif mxfp4_backend == Mxfp4MoeBackend.AITER_MXFP4_FP8:
# W4A8: MXFP4 weights + static FP8 activations
return mxfp4_w4a8_moe_quant_config(
w1_scale=w1_scale,
w2_scale=w2_scale,
a1_scale=a1_scale,
a2_scale=a2_scale,
w1_bias=w1_bias,
w2_bias=w2_bias,
block_shape=None,
)
elif mxfp4_backend in (
Mxfp4MoeBackend.MARLIN,
Mxfp4MoeBackend.BATCHED_MARLIN,
@@ -1269,7 +1382,7 @@ def make_mxfp4_moe_quant_config(
Mxfp4MoeBackend.TRITON_UNFUSED,
Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_BF16,
Mxfp4MoeBackend.FLASHINFER_CUTLASS_MXFP4_BF16,
Mxfp4MoeBackend.AITER,
Mxfp4MoeBackend.AITER_MXFP4_BF16,
):
return mxfp4_w4a16_moe_quant_config(
w1_bias=w1_bias,
@@ -68,21 +68,23 @@ class MeanPool(SequencePoolingMethod):
"partial prefill not supported with MEAN pooling"
)
prompt_lens = pooling_cursor.prompt_lens_cpu.to(
hidden_states.device, dtype=torch.int64, non_blocking=True
)
num_seqs = prompt_lens.numel()
prompt_lens_cpu = pooling_cursor.prompt_lens_cpu
num_seqs = prompt_lens_cpu.numel()
hidden_size = hidden_states.shape[-1]
if num_seqs == 0:
# early return for empty batch
return hidden_states.new_empty((0, hidden_size), dtype=torch.float32)
# eg. [2, 1, 3] -> [0, 0, 1, 2, 2, 2]
# Build segment_ids on CPU so repeat_interleave doesn't need to sync
# GPU->CPU to learn its data-dependent output length, then upload
# non-blocking. eg. [2, 1, 3] -> [0, 0, 1, 2, 2, 2]
segment_ids = torch.repeat_interleave(
torch.arange(num_seqs, device=hidden_states.device, dtype=torch.long),
prompt_lens,
torch.arange(num_seqs, dtype=torch.long),
prompt_lens_cpu,
).to(hidden_states.device, non_blocking=True)
prompt_lens = prompt_lens_cpu.to(
hidden_states.device, dtype=torch.int64, non_blocking=True
)
segment_sums = torch.zeros(
(num_seqs, hidden_size),
+34 -4
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import dataclasses
from collections.abc import Mapping, Set
from itertools import groupby
@@ -80,9 +81,11 @@ class DispatchPooler(Pooler):
pooling_metadata: PoolingMetadata,
) -> PoolerOutput:
poolers_by_task = self.poolers_by_task
cursor = pooling_metadata.pooling_cursor
outputs = list[torch.Tensor | None]()
offset = 0
token_offset = 0
for task, group in groupby(pooling_metadata.tasks):
if not (pooler := poolers_by_task.get(task)):
raise ValueError(
@@ -91,10 +94,37 @@ class DispatchPooler(Pooler):
)
num_items = len(list(group))
group_output: PoolerOutput = pooler(
hidden_states,
pooling_metadata[offset : offset + num_items],
)
group_metadata = pooling_metadata[offset : offset + num_items]
if cursor is None:
group_hidden_states = hidden_states
else:
# Slice out this group's tokens so sub-poolers see only their
# portion of the batch. Token offset is computed from the CPU
# `num_scheduled_tokens_cpu` to avoid a GPU->CPU sync.
group_cursor = group_metadata.pooling_cursor
num_group_tokens = int(group_cursor.num_scheduled_tokens_cpu.sum())
group_hidden_states = hidden_states[
token_offset : token_offset + num_group_tokens
]
if token_offset:
# Shift first/last indices to be relative to the slice
# so seqwise poolers (which index `hidden_states` directly)
# remain correct.
pooling_cursor = dataclasses.replace(
group_cursor,
first_token_indices_gpu=(
group_cursor.first_token_indices_gpu - token_offset
),
last_token_indices_gpu=(
group_cursor.last_token_indices_gpu - token_offset
),
)
group_metadata = dataclasses.replace(
group_metadata, pooling_cursor=pooling_cursor
)
token_offset += num_group_tokens
group_output: PoolerOutput = pooler(group_hidden_states, group_metadata)
outputs.extend(group_output)
offset += num_items
@@ -47,17 +47,12 @@ class AllPool(TokenPoolingMethod):
pooling_metadata: PoolingMetadata,
) -> list[TokenPoolingMethodOutputItem]:
pooling_cursor = pooling_metadata.get_pooling_cursor()
split_sizes = pooling_cursor.num_scheduled_tokens_cpu.tolist()
if split_sizes:
# DispatchPooler passes the full hidden_states tensor.
# slice out the subgroup once, then split it by
# per-request token counts
group_start = int(pooling_cursor.first_token_indices_gpu[0].item())
group_end = int(pooling_cursor.last_token_indices_gpu[-1].item()) + 1
hidden_states_group = hidden_states[group_start:group_end]
hidden_states_lst = list(hidden_states_group.split(split_sizes))
else:
hidden_states_lst = []
# Use the already-CPU num_scheduled_tokens tensor so `.tolist()`
# doesn't trigger a GPU->CPU sync. torch.split produces the same
# consecutive slices as indexing with first/last per-sequence indices.
hidden_states_lst = list(
torch.split(hidden_states, pooling_cursor.num_scheduled_tokens_cpu.tolist())
)
if not self.enable_chunked_prefill:
return hidden_states_lst
@@ -95,12 +90,14 @@ class StepPool(AllPool):
pooling_metadata: PoolingMetadata,
) -> list[TokenPoolingMethodOutputItem]:
pooled_data_lst = super().forward(hidden_states, pooling_metadata)
prompt_token_ids = pooling_metadata.get_prompt_token_ids()
# Use the CPU copy of prompt_token_ids so the step_tag_id mask can be
# resolved to indices without a d2h sync from boolean indexing.
prompt_token_ids_cpu = pooling_metadata.get_prompt_token_ids_cpu()
pooling_params = pooling_metadata.pooling_params
pooled_data = list[torch.Tensor | None]()
for data, token_id, pooling_param in zip(
pooled_data_lst, prompt_token_ids, pooling_params
for data, token_id_cpu, pooling_param in zip(
pooled_data_lst, prompt_token_ids_cpu, pooling_params
):
# for unfinished chunked prefill
if data is None:
@@ -113,7 +110,9 @@ class StepPool(AllPool):
data = data[:, returned_token_ids]
if step_tag_id is not None:
data = data[token_id == step_tag_id]
idx_cpu = (token_id_cpu == step_tag_id).nonzero(as_tuple=True)[0]
idx = idx_cpu.to(data.device, non_blocking=True)
data = data[idx]
pooled_data.append(data)
@@ -24,7 +24,7 @@ from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
make_mxfp4_moe_kernel,
make_mxfp4_moe_quant_config,
mxfp4_round_up_hidden_size_and_intermediate_size,
select_gpt_oss_mxfp4_moe_backend,
select_deepseek_v4_mxfp4_moe_backend,
select_mxfp4_moe_backend,
)
from vllm.model_executor.layers.linear import LinearBase, UnquantizedLinearMethod
@@ -140,7 +140,7 @@ class GptOssMxfp4MoEMethod(FusedMoEMethodBase):
def __init__(self, moe: FusedMoEConfig):
super().__init__(moe)
self.weight_dtype = "gpt_oss_mxfp4"
self.mxfp4_backend, self.experts_cls = select_gpt_oss_mxfp4_moe_backend(moe)
self.mxfp4_backend, self.experts_cls = select_mxfp4_moe_backend(moe)
self.max_capture_size = (
get_current_vllm_config().compilation_config.max_cudagraph_capture_size
@@ -468,7 +468,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
def __init__(self, moe: FusedMoEConfig):
super().__init__(moe)
self.weight_dtype = "mxfp4"
self.mxfp4_backend, self.experts_cls = select_mxfp4_moe_backend(moe)
self.mxfp4_backend, self.experts_cls = select_deepseek_v4_mxfp4_moe_backend(moe)
self.max_capture_size = (
get_current_vllm_config().compilation_config.max_cudagraph_capture_size
@@ -35,19 +35,19 @@ from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
make_mxfp4_moe_kernel,
make_mxfp4_moe_quant_config,
mxfp4_round_up_hidden_size_and_intermediate_size,
select_gpt_oss_mxfp4_moe_backend,
select_mxfp4_moe_backend,
)
from vllm.model_executor.layers.quantization.utils.marlin_utils_fp8 import (
prepare_fp8_moe_layer_for_marlin,
)
from vllm.model_executor.layers.quantization.utils.mxfp4_utils import (
_swizzle_mxfp4,
)
from vllm.model_executor.layers.quantization.utils.ocp_mx_utils import (
OCP_MX_BLOCK_SIZE,
OCP_MX_Scheme,
)
from vllm.model_executor.layers.quantization.utils.quant_utils import GroupShape
from vllm.model_executor.layers.quantization.utils.quant_utils import (
GroupShape,
kFp8StaticTensorSym,
)
from vllm.model_executor.layers.quantization.utils.w8a8_utils import (
all_close_1d,
normalize_e4m3fn_to_e4m3fnuz,
@@ -62,7 +62,6 @@ logger = init_logger(__name__)
__all__ = [
"QuarkMoEMethod",
"QuarkOCP_MX_MoEMethod",
"QuarkOCP_MX_MoEMethod_OSS",
]
@@ -94,22 +93,9 @@ class QuarkMoEMethod(FusedMoEMethodBase):
elif quant_config._is_fp8_w8a8(weight_config, input_config):
return QuarkW8A8Fp8MoEMethod(weight_config, input_config, module.moe_config)
elif quant_config._is_w_ocp_mx_a_x(weight_config, input_config):
emulate = not current_platform.supports_mx() or not (
rocm_aiter_ops.is_fused_moe_enabled()
)
if (
input_config is not None
and input_config.get("dtype") == "fp8_e4m3"
and not input_config.get("is_dynamic")
and not emulate
):
return QuarkOCP_MX_MoEMethod_OSS(
weight_config, input_config, module.moe_config
)
else:
return QuarkOCP_MX_MoEMethod(
weight_config, input_config, module.moe_config
)
# All OCP MX schemes (W4A16, W4A8, etc.) handled by QuarkOCP_MX_MoEMethod
# Backend selection happens inside via oracle
return QuarkOCP_MX_MoEMethod(weight_config, input_config, module.moe_config)
elif quant_config._is_static_tensor_w8a8(
weight_config, input_config
) or quant_config._is_dynamic_per_token_w8a8(weight_config, input_config):
@@ -993,7 +979,7 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
self.experts_cls: type[mk.FusedMoEExperts] | None = None
self.moe_kernel: mk.FusedMoEKernel | None = None
# Used for triton kernel precision configs
# Used for triton kernel precision configs (W4A8, TRITON backends)
self.w13_precision_config = None
self.w2_precision_config = None
@@ -1002,6 +988,17 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
else:
self.static_input_scales = False
# Select backend based on OCP MX scheme
if self.ocp_mx_scheme == "w_mxfp4":
# W4A16: weight-only MXFP4
self.mxfp4_backend, self.experts_cls = select_mxfp4_moe_backend(moe)
elif self.ocp_mx_scheme == "w_mxfp4_a_fp8" and self.static_input_scales:
# W4A8: MXFP4 weights + static FP8 activations
self.mxfp4_backend, self.experts_cls = select_mxfp4_moe_backend(
moe, activation_key=kFp8StaticTensorSym
)
# Validation for unsupported schemes
if any(
self.ocp_mx_scheme.endswith(a_scheme)
for a_scheme in ["a_mxfp4", "a_mxfp6_e3m2", "a_mxfp6_e2m3"]
@@ -1026,7 +1023,7 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
)
# TODO: Remove once all OCP MX schemes use the kernel abstraction
_AITER_NATIVE_OCP_MX_SCHEMES = ("w_mxfp4", "w_mxfp4_a_mxfp4")
_AITER_NATIVE_OCP_MX_SCHEMES = ("w_mxfp4", "w_mxfp4_a_mxfp4", "w_mxfp4_a_fp8")
self.emulate = (
not current_platform.supports_mx()
or self.ocp_mx_scheme not in _AITER_NATIVE_OCP_MX_SCHEMES
@@ -1034,9 +1031,6 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
self.mxfp4_backend is Mxfp4MoeBackend.NONE or not self.use_rocm_aiter_moe
)
if self.ocp_mx_scheme == "w_mxfp4":
self.mxfp4_backend, self.experts_cls = select_gpt_oss_mxfp4_moe_backend(moe)
if self.emulate:
# We use the same code path between MXFP4/MXFP6 emulation.
self.mxfp4_backend = Mxfp4MoeBackend.EMULATION
@@ -1046,7 +1040,12 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
if self.mxfp4_backend != Mxfp4MoeBackend.NONE:
self.experts_cls = backend_to_kernel_cls(self.mxfp4_backend)[0]
if self.emulate:
# Log backend selection
if self.mxfp4_backend != Mxfp4MoeBackend.NONE:
logger.info_once(
f"Using {self.mxfp4_backend.value} backend for {self.ocp_mx_scheme}"
)
elif self.emulate:
logger.warning_once(
f"The current mode (supports_mx={current_platform.supports_mx()}, "
f"use_rocm_aiter_moe={self.use_rocm_aiter_moe}, "
@@ -1056,10 +1055,6 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
"QDQ (quantize and dequantize) will be used, with the linear "
"layers computed in high precision."
)
else:
logger.warning_once(
"The current mode supports native MoE MXFP4 computation"
)
def maybe_roundup_sizes(
self,
@@ -1204,6 +1199,11 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
layer.w2_input_scale = None
def process_weights_after_loading(self, layer):
# For MXFP4 schemes with native backend, use oracle
if self.mxfp4_backend != Mxfp4MoeBackend.NONE:
self._setup_kernel(layer)
return
if self.static_input_scales and self.input_dtype == "fp8":
# firstly, process activations if fp8 static input
if layer.w13_input_scale is None or layer.w2_input_scale is None:
@@ -1252,14 +1252,6 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
w2_input_scale, requires_grad=False
)
# For w_mxfp4, use oracle functions
if self.emulate or (
self.ocp_mx_scheme == "w_mxfp4"
and self.mxfp4_backend != Mxfp4MoeBackend.NONE
):
self._setup_kernel_via_oracle(layer)
return
# TODO(bowenbao): gradually migrate to oracles.
# Existing AITER path for w_mxfp4_a_mxfp4 and other schemes
from aiter.utility.fp4_utils import e8m0_shuffle
@@ -1298,46 +1290,48 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
torch.accelerator.empty_cache()
def _setup_kernel_via_oracle(self, layer: FusedMoE):
"""Setup kernel using oracle functions for w_mxfp4 scheme."""
w13 = layer.w13_weight
w2 = layer.w2_weight
w13_scale = layer.w13_weight_scale
w2_scale = layer.w2_weight_scale
def _setup_kernel(self, layer: FusedMoE):
"""Setup kernel using oracle functions for MXFP4 schemes (W4A16, W4A8)."""
w13_bias = getattr(layer, "w13_bias", None)
w2_bias = getattr(layer, "w2_bias", None)
# Convert weights to kernel format
# Convert weights to kernel format (handles all backend-specific logic)
w13, w2, w13_scale, w2_scale, w13_bias, w2_bias = (
convert_gpt_oss_weight_to_mxfp4_moe_kernel_format(
mxfp4_backend=self.mxfp4_backend,
layer=layer,
w13_weight=w13,
w2_weight=w2,
w13_weight_scale=w13_scale,
w2_weight_scale=w2_scale,
w13_weight=layer.w13_weight,
w2_weight=layer.w2_weight,
w13_weight_scale=layer.w13_weight_scale,
w2_weight_scale=layer.w2_weight_scale,
w13_bias=w13_bias,
w2_bias=w2_bias,
)
)
# For TRITON backends, weights are wrapped tensors from triton_kernels
# that don't support .detach(). Manually assign parameters.
if self.mxfp4_backend not in TRITON_BACKENDS:
replace_parameter(layer, "w13_weight", w13)
replace_parameter(layer, "w2_weight", w2)
replace_parameter(layer, "w13_weight_scale", w13_scale)
replace_parameter(layer, "w2_weight_scale", w2_scale)
else:
# Handle weight/scale assignment based on backend type
if self.mxfp4_backend in TRITON_BACKENDS or self.mxfp4_backend in (
Mxfp4MoeBackend.AITER_MXFP4_FP8,
):
# Triton-based backends: w13/w2 are triton_kernels.tensor.Tensor
# Store on layer for apply(), scales are PrecisionConfig
layer.w13_weight = w13
layer.w2_weight = w2
self.w13_precision_config = w13_scale
self.w2_precision_config = w2_scale
else:
# Standard backends: replace parameters
replace_parameter(layer, "w13_weight", w13)
replace_parameter(layer, "w2_weight", w2)
replace_parameter(layer, "w13_weight_scale", w13_scale)
replace_parameter(layer, "w2_weight_scale", w2_scale)
if w13_bias is not None and w2_bias is not None:
replace_parameter(layer, "w13_bias", w13_bias)
replace_parameter(layer, "w2_bias", w2_bias)
torch.accelerator.empty_cache()
# Build quant config and kernel
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
if self.moe_quant_config is not None and self.experts_cls is not None:
@@ -1353,22 +1347,26 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
def get_fused_moe_quant_config(
self, layer: torch.nn.Module
) -> FusedMoEQuantConfig | None:
# For w_mxfp4 with oracle backend, use oracle function
if self.ocp_mx_scheme == "w_mxfp4" and self.mxfp4_backend not in (
Mxfp4MoeBackend.NONE,
Mxfp4MoeBackend.EMULATION,
):
w1_scale = layer.w13_weight_scale
w2_scale = layer.w2_weight_scale
if self.mxfp4_backend in TRITON_BACKENDS:
# For oracle-based backends (W4A16, W4A8), use make_mxfp4_moe_quant_config
if self.mxfp4_backend not in (Mxfp4MoeBackend.NONE, Mxfp4MoeBackend.EMULATION):
# Determine scale source based on backend type
if self.mxfp4_backend in TRITON_BACKENDS or self.mxfp4_backend in (
Mxfp4MoeBackend.AITER_MXFP4_FP8,
):
w1_scale = self.w13_precision_config
w2_scale = self.w2_precision_config
else:
w1_scale = layer.w13_weight_scale
w2_scale = layer.w2_weight_scale
return make_mxfp4_moe_quant_config(
mxfp4_backend=self.mxfp4_backend,
w1_scale=w1_scale,
w2_scale=w2_scale,
w1_bias=getattr(layer, "w13_bias", None),
w2_bias=getattr(layer, "w2_bias", None),
a1_scale=getattr(layer, "w13_input_scale", None),
a2_scale=getattr(layer, "w2_input_scale", None),
)
# Emulation and other schemes
@@ -1421,7 +1419,7 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
topk_ids: torch.Tensor,
shared_experts_input: torch.Tensor | None,
) -> torch.Tensor:
# For oracle kernel or emulation kernel
# For oracle-based kernels (W4A16, W4A8) or emulation kernel
if self.moe_kernel is not None:
return self.moe_kernel.apply(
hidden_states=x,
@@ -1473,135 +1471,3 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod):
expert_map=layer.expert_map,
apply_router_weight_on_input=layer.apply_router_weight_on_input,
)
class QuarkOCP_MX_MoEMethod_OSS(QuarkOCP_MX_MoEMethod):
def __init__(
self,
weight_config: dict[str, Any],
input_config: dict[str, Any],
moe: FusedMoEConfig,
):
super().__init__(weight_config, input_config, moe)
def process_weights_after_loading(self, layer):
from triton_kernels.matmul_ogs import FlexCtx, PrecisionConfig
w13_bias = layer.w13_bias.to(torch.float32)
w2_bias = layer.w2_bias.to(torch.float32)
layer.w13_bias = torch.nn.Parameter(w13_bias, requires_grad=False)
layer.w2_bias = torch.nn.Parameter(w2_bias, requires_grad=False)
# FIXME warp need to be adjusted based on batch size
# only apply to batched mode
if self.moe.use_ep:
num_warps = 4 if self.moe.max_num_tokens <= 512 else 8
else:
num_warps = 8
w13_weight, w13_flex, w13_scale = _swizzle_mxfp4(
layer.w13_weight, layer.w13_weight_scale, num_warps
)
w2_weight, w2_flex, w2_scale = _swizzle_mxfp4(
layer.w2_weight, layer.w2_weight_scale, num_warps
)
self.w13_weight_triton_tensor = w13_weight
self.w2_weight_triton_tensor = w2_weight
# need to delete the original weights to save memory on single GPU
del layer.w13_weight
del layer.w2_weight
layer.w13_weight = None
layer.w2_weight = None
torch.accelerator.empty_cache()
if self.static_input_scales:
if layer.w13_input_scale is None or layer.w2_input_scale is None:
raise ValueError(
"QuantConfig has static quantization, but found "
"activation scales are None."
)
if not all_close_1d(layer.w13_input_scale) or not all_close_1d(
layer.w2_input_scale
):
logger.warning_once(
"Found input_scales that are not equal for "
"fp8 MoE layer. Using the maximum across experts "
"for each layer."
)
layer.w13_input_scale = torch.nn.Parameter(
layer.w13_input_scale.max().to(torch.float32), requires_grad=False
)
layer.w2_input_scale = torch.nn.Parameter(
layer.w2_input_scale.max().to(torch.float32), requires_grad=False
)
from triton_kernels.numerics import InFlexData
lhs_data13 = InFlexData(scale=layer.w13_input_scale)
lhs_data2 = InFlexData(scale=layer.w2_input_scale)
self.w13_precision_config = PrecisionConfig(
weight_scale=w13_scale,
flex_ctx=FlexCtx(rhs_data=w13_flex, lhs_data=lhs_data13),
)
self.w2_precision_config = PrecisionConfig(
weight_scale=w2_scale,
flex_ctx=FlexCtx(rhs_data=w2_flex, lhs_data=lhs_data2),
)
def get_fused_moe_quant_config(
self, layer: torch.nn.Module
) -> FusedMoEQuantConfig | None:
return mxfp4_w4a8_moe_quant_config(
w1_scale=self.w13_precision_config,
w2_scale=self.w2_precision_config,
a1_scale=layer.w13_input_scale,
a2_scale=layer.w2_input_scale,
w1_bias=layer.w13_bias,
w2_bias=layer.w2_bias,
block_shape=None,
)
@property
def is_monolithic(self) -> bool:
return True
def apply_monolithic(
self,
layer: FusedMoE,
x: torch.Tensor,
router_logits: torch.Tensor,
input_ids: torch.Tensor | None = None,
) -> torch.Tensor:
if layer.enable_eplb:
raise NotImplementedError(
f"EPLB not supported for {self.__class__.__name__} yet."
)
from vllm.model_executor.layers.fused_moe.experts.gpt_oss_triton_kernels_moe import ( # noqa: E501
triton_kernel_moe_forward,
)
assert self.moe.hidden_dim_unpadded is not None
assert self.moe.intermediate_size_per_partition_unpadded is not None
return triton_kernel_moe_forward(
hidden_states=x,
w1=self.w13_weight_triton_tensor,
w2=self.w2_weight_triton_tensor,
gating_output=router_logits,
topk=layer.top_k,
renormalize=layer.renormalize,
global_num_experts=layer.global_num_experts,
expert_map=layer.expert_map,
quant_config=self.moe_quant_config,
apply_router_weight_on_input=layer.apply_router_weight_on_input,
unpadded_N_w1=self.moe.intermediate_size_per_partition_unpadded * 2,
unpadded_K_w1=self.moe.hidden_dim_unpadded,
unpadded_N_w2=self.moe.hidden_dim_unpadded,
unpadded_K_w2=self.moe.intermediate_size_per_partition_unpadded,
)
@@ -2,8 +2,17 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""TurboQuant configuration."""
from __future__ import annotations
import logging
import math
from dataclasses import dataclass
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from vllm.config import ModelConfig
logger = logging.getLogger(__name__)
# Named TQ presets: each maps to frozen config parameters.
# key_quant_bits: 8 = FP8 keys, 3-4 = MSE (Lloyd-Max) quantized keys.
@@ -159,12 +168,34 @@ class TurboQuantConfig:
return s + (s % 2) # round up to even
@staticmethod
def get_boundary_skip_layers(num_layers: int, n: int = 2) -> list[str]:
"""Get layer indices to skip TQ compression (boundary protection).
def get_boundary_skip_layers(
model_config: ModelConfig,
n: int = 2,
) -> list[str]:
"""Layer indices to skip TQ compression (boundary protection).
Returns first N and last N layer indices as strings, suitable for
kv_cache_dtype_skip_layers.
For hybrid models (attention + Mamba/linear-attention), boundary
protection is disabled hybrids typically have only 8-12
full-attention layers and a hard n=2 on each side would cover
~40 % of them. The dense GSM8K baselines that motivate n=2
don't apply to hybrids.
For dense models, skips first N and last N attention layers.
Empirically required for aggressive presets (k3v4_nc, 3bit_nc)
without it GSM8K drops ~30 points on Qwen3-4B.
"""
if model_config.is_hybrid:
attn_indices = _get_full_attention_layer_indices(model_config)
if not attn_indices:
raise NotImplementedError(
"TurboQuant KV cache requires identifiable "
"full-attention layers, but none were found in "
"the hybrid model config."
)
logger.info("TQ hybrid: full-attention layers %s", attn_indices)
return []
num_layers = model_config.hf_text_config.num_hidden_layers
if n <= 0 or num_layers <= 0:
return []
n = min(n, num_layers // 2) # don't skip more than half
@@ -175,7 +206,7 @@ class TurboQuantConfig:
return [str(i) for i in indices]
@staticmethod
def from_cache_dtype(cache_dtype: str, head_dim: int) -> "TurboQuantConfig":
def from_cache_dtype(cache_dtype: str, head_dim: int) -> TurboQuantConfig:
"""Create config from a named preset.
Valid presets: turboquant_k8v4, turboquant_4bit_nc, etc.
@@ -193,3 +224,31 @@ class TurboQuantConfig:
value_quant_bits=preset["value_quant_bits"],
norm_correction=preset["norm_correction"],
)
def _get_full_attention_layer_indices(model_config: ModelConfig) -> list[int]:
"""Global indices of full-attention layers in a hybrid model.
Covers the conventions used across vLLM: ``layer_types`` (Qwen3.5/Next),
``layers_block_type`` (Jamba/Zamba2), ``attn_type_list`` (Minimax).
"""
text_cfg = model_config.hf_text_config
hf_cfg = model_config.hf_config
layer_types = getattr(text_cfg, "layer_types", None)
if layer_types is not None:
return [
i for i, t in enumerate(layer_types) if t in ("full_attention", "attention")
]
layers_block_type = getattr(text_cfg, "layers_block_type", None)
if layers_block_type is not None:
return [
i for i, t in enumerate(layers_block_type) if t in ("attention", "hybrid")
]
attn_type_list = getattr(hf_cfg, "attn_type_list", None)
if attn_type_list is not None:
return [i for i, t in enumerate(attn_type_list) if t == 1]
return []
+63 -59
View File
@@ -37,6 +37,7 @@ from vllm.sequence import IntermediateTensors
from .commandr import LayerNorm
from .interfaces import SupportsPP, SupportsQuant
from .utils import (
AutoWeightsLoader,
extract_layer_index,
is_pp_missing_parameter,
make_empty_intermediate_tensors_factory,
@@ -330,6 +331,7 @@ class CohereMoeModel(nn.Module):
quant_config = vllm_config.quant_config
self.config = config
self.quant_config = quant_config
self.vocab_size = config.vocab_size
self.org_vocab_size = config.vocab_size
self.embed_tokens = VocabParallelEmbedding(
@@ -378,63 +380,6 @@ class CohereMoeModel(nn.Module):
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
class CohereMoeForCausalLM(nn.Module, SupportsPP, SupportsQuant):
is_text_generation_model = True
packed_modules_mapping = {
"qkv_proj": [
"q_proj",
"k_proj",
"v_proj",
],
"gate_up_proj": [
"gate_proj",
"up_proj",
],
}
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
super().__init__()
config = vllm_config.model_config.hf_config
quant_config = vllm_config.quant_config
self.config = config
assert getattr(config, "tie_word_embeddings", True)
self.unpadded_vocab_size = config.vocab_size
self.quant_config = quant_config
self.logits_scale = config.logit_scale
self.logits_processor = LogitsProcessor(
self.unpadded_vocab_size, config.vocab_size, scale=self.logits_scale
)
self.model = CohereMoeModel(
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
)
self.make_empty_intermediate_tensors = (
self.model.make_empty_intermediate_tensors
)
def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.model.get_input_embeddings(input_ids)
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.model.get_input_embeddings(input_ids)
@torch.no_grad()
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
intermediate_tensors: IntermediateTensors | None = None,
inputs_embeds: torch.Tensor | None = None,
) -> torch.Tensor | IntermediateTensors:
return self.model(input_ids, positions, intermediate_tensors, inputs_embeds)
def compute_logits(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor | None:
return self.logits_processor(self.model.embed_tokens, hidden_states)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
stacked_params_mapping = [
("qkv_proj", "q_proj", "q"),
@@ -507,8 +452,6 @@ class CohereMoeForCausalLM(nn.Module, SupportsPP, SupportsQuant):
)
break
else:
if "lm_head.weight" in name:
continue
if (
name.endswith(".bias") or name.endswith("_bias")
) and name not in params_dict:
@@ -526,3 +469,64 @@ class CohereMoeForCausalLM(nn.Module, SupportsPP, SupportsQuant):
loaded_params.add(name)
return loaded_params
class CohereMoeForCausalLM(nn.Module, SupportsPP, SupportsQuant):
is_text_generation_model = True
packed_modules_mapping = {
"qkv_proj": [
"q_proj",
"k_proj",
"v_proj",
],
"gate_up_proj": [
"gate_proj",
"up_proj",
],
}
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
super().__init__()
config = vllm_config.model_config.hf_config
quant_config = vllm_config.quant_config
self.config = config
assert getattr(config, "tie_word_embeddings", True)
self.unpadded_vocab_size = config.vocab_size
self.quant_config = quant_config
self.logits_scale = config.logit_scale
self.logits_processor = LogitsProcessor(
self.unpadded_vocab_size, config.vocab_size, scale=self.logits_scale
)
self.model = CohereMoeModel(
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
)
self.make_empty_intermediate_tensors = (
self.model.make_empty_intermediate_tensors
)
def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.model.get_input_embeddings(input_ids)
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.model.get_input_embeddings(input_ids)
@torch.no_grad()
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
intermediate_tensors: IntermediateTensors | None = None,
inputs_embeds: torch.Tensor | None = None,
) -> torch.Tensor | IntermediateTensors:
return self.model(input_ids, positions, intermediate_tensors, inputs_embeds)
def compute_logits(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor | None:
return self.logits_processor(self.model.embed_tokens, hidden_states)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
loader = AutoWeightsLoader(self, skip_prefixes=["lm_head."])
return loader.load_weights(weights)
+1 -1
View File
@@ -360,7 +360,7 @@ class Gemma4MoE(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.experts",
custom_routing_function=routing_function,
activation="gelu",
activation="gelu_tanh",
)
def forward(self, x: torch.Tensor, router_logits: torch.Tensor) -> torch.Tensor:
+83 -85
View File
@@ -797,6 +797,83 @@ class Plamo2Model(torch.nn.Module):
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
# Update the weight names to be compatible with the vllm version
# of the model.
# Do not change the order of the replacements.
replacements = {
# Rename incompatible weight names.
".A_log": ".A",
".B_norm_weight": ".B_norm.weight",
".C_norm_weight": ".C_norm.weight",
".dt_norm_weight": ".dt_norm.weight",
".q_weight": ".q_norm.weight",
".k_weight": ".k_norm.weight",
}
# Apply replacements based on the defined mappings
for old, new in replacements.items():
if old in name:
name = name.replace(old, new)
# Reshape the in_proj weights to match the shape expected
# by MergedColumnParallelLinear.
# This works both for unquantized weights and
# for quantized weights.
# In the quantized case, the weights are already transposed.
# Also, in addition to the quantized weights,
# the zero points and scales have to be reshaped as well.
# Packing should not be affected by this.
if (
".mixer.in_proj.weight" in name
or "mixer.in_proj.qweight" in name
or "mixer.in_proj.scales" in name
or "mixer.in_proj.qzeros" in name
):
if "mixer.in_proj.weight" in name:
loaded_weight = loaded_weight.transpose(0, 1)
# for weight:
# loaded_weight.shape[0] == self.config.hidden_size
# for qweight:
# loaded_weight.shape[0] == self.config.hidden_size // param.pack_factor # noqa
# for scales and qzeros:
# loaded_weight.shape[0] == self.config.hidden_size // self.vllm_config.quant_config.group_size # noqa
loaded_weight = loaded_weight.reshape(
loaded_weight.shape[0], self.config.mamba_num_heads, -1
)
gate_weight, hidden_states_weight = loaded_weight.chunk(2, dim=-1)
gate_weight = gate_weight.reshape(loaded_weight.shape[0], -1)
hidden_states_weight = hidden_states_weight.reshape(
loaded_weight.shape[0], -1
)
loaded_weight = torch.cat([gate_weight, hidden_states_weight], dim=-1)
if "mixer.in_proj.weight" in name:
loaded_weight = loaded_weight.transpose(0, 1)
# Offset parameter with vllm's RMSNorm haven't been supported yet.
if ".pre_mixer_norm" in name:
loaded_weight += 1.0
elif ".post_mixer_norm" in name:
loaded_weight += 1.0 / 5
elif ".pre_mlp_norm" in name:
loaded_weight += 1.0
elif ".post_mlp_norm" in name:
loaded_weight += 1.0 / (5**1.5)
elif name == "norm.weight":
loaded_weight += 1.0
# Skip layers on other devices.
if is_pp_missing_parameter(name, self):
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
class Plamo2ForCausalLM(
torch.nn.Module, HasInnerState, SupportsLoRA, SupportsPP, IsHybrid
@@ -906,88 +983,9 @@ class Plamo2ForCausalLM(
logits = self.logits_processor(self.lm_head, hidden_states)
return logits
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
params_dict = dict(self.named_parameters())
for name, loaded_weight in weights:
# Both tie_word_embeddings=True and lm_head.weight in the safetensor
# at the same time causes dict key access error.
if name == "lm_head.weight" and self.config.tie_word_embeddings:
assert "lm_head.weight" not in params_dict
continue
# Same workaround as AutoWeightsLoader for GPTQModel
if any(
substr in name
for substr in AutoWeightsLoader.ROTARY_EMBEDS_UNUSED_WEIGHTS
):
continue
# Update the weight names to be compatible with the vllm version
# of the model.
# Do not change the order of the replacements.
replacements = {
# Rename incompatible weight names.
".A_log": ".A",
".B_norm_weight": ".B_norm.weight",
".C_norm_weight": ".C_norm.weight",
".dt_norm_weight": ".dt_norm.weight",
".q_weight": ".q_norm.weight",
".k_weight": ".k_norm.weight",
}
# Apply replacements based on the defined mappings
for old, new in replacements.items():
if old in name:
name = name.replace(old, new)
# Reshape the in_proj weights to match the shape expected
# by MergedColumnParallelLinear.
# This works both for unquantized weights and
# for quantized weights.
# In the quantized case, the weights are already transposed.
# Also, in addition to the quantized weights,
# the zero points and scales have to be reshaped as well.
# Packing should not be affected by this.
if (
".mixer.in_proj.weight" in name
or "mixer.in_proj.qweight" in name
or "mixer.in_proj.scales" in name
or "mixer.in_proj.qzeros" in name
):
if "mixer.in_proj.weight" in name:
loaded_weight = loaded_weight.transpose(0, 1)
# for weight:
# loaded_weight.shape[0] == self.config.hidden_size
# for qweight:
# loaded_weight.shape[0] == self.config.hidden_size // param.pack_factor # noqa
# for scales and qzeros:
# loaded_weight.shape[0] == self.config.hidden_size // self.vllm_config.quant_config.group_size # noqa
loaded_weight = loaded_weight.reshape(
loaded_weight.shape[0], self.config.mamba_num_heads, -1
)
gate_weight, hidden_states_weight = loaded_weight.chunk(2, dim=-1)
gate_weight = gate_weight.reshape(loaded_weight.shape[0], -1)
hidden_states_weight = hidden_states_weight.reshape(
loaded_weight.shape[0], -1
)
loaded_weight = torch.cat([gate_weight, hidden_states_weight], dim=-1)
if "mixer.in_proj.weight" in name:
loaded_weight = loaded_weight.transpose(0, 1)
# Offset parameter with vllm's RMSNorm haven't been supported yet.
if ".pre_mixer_norm" in name:
loaded_weight += 1.0
elif ".post_mixer_norm" in name:
loaded_weight += 1.0 / 5
elif ".pre_mlp_norm" in name:
loaded_weight += 1.0
elif ".post_mlp_norm" in name:
loaded_weight += 1.0 / (5**1.5)
elif "model.norm.weight" in name:
loaded_weight += 1.0
# Skip layers on other devices.
if is_pp_missing_parameter(name, self):
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
loader = AutoWeightsLoader(
self,
skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None),
)
return loader.load_weights(weights)
+92
View File
@@ -0,0 +1,92 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# QianfanOCR is built on InternVL with a Qwen3 language backbone.
# The model architecture and weights are fully compatible with InternVLChatModel,
# only the config model_type / architectures strings differ.
from transformers import PretrainedConfig
from vllm.model_executor.layers.quantization import QuantizationConfig
from vllm.model_executor.layers.quantization.fp8 import Fp8Config
from vllm.multimodal import MULTIMODAL_REGISTRY
from vllm.transformers_utils.processors.internvl import (
InternVLImageProcessor,
InternVLProcessor,
)
from .internvl import (
BaseInternVLDummyInputsBuilder,
BaseInternVLMultiModalProcessor,
BaseInternVLProcessingInfo,
InternVLChatModel,
)
class QianfanOCRProcessingInfo(BaseInternVLProcessingInfo):
"""Image-only ProcessingInfo for QianfanOCR (no video support)."""
def get_hf_processor(self, **kwargs: object) -> InternVLProcessor:
config = self.get_hf_config()
vision_config = config.vision_config
kwargs = self.ctx.get_merged_mm_kwargs(kwargs)
kwargs.setdefault("image_size", vision_config.image_size)
kwargs.setdefault("min_dynamic_patch", config.min_dynamic_patch)
kwargs.setdefault("max_dynamic_patch", config.max_dynamic_patch)
kwargs.setdefault("dynamic_image_size", config.dynamic_image_size)
kwargs.setdefault("use_thumbnail", config.use_thumbnail)
image_processor = InternVLImageProcessor(**kwargs)
image_size = image_processor.image_size
patch_size = vision_config.patch_size
downsample_ratio = config.downsample_ratio
image_seq_length = int((image_size // patch_size) ** 2 * (downsample_ratio**2))
return InternVLProcessor(
tokenizer=self.get_tokenizer(),
image_processor=image_processor,
video_processor=None,
image_seq_length=image_seq_length,
ctx_video_token=None,
)
@MULTIMODAL_REGISTRY.register_processor(
BaseInternVLMultiModalProcessor,
info=QianfanOCRProcessingInfo,
dummy_inputs=BaseInternVLDummyInputsBuilder,
)
class QianfanOCRForConditionalGeneration(InternVLChatModel):
"""QianfanOCR multimodal model.
Identical in structure to InternVLChatModel (InternViT vision encoder +
pixel-shuffle MLP connector + Qwen3 language model). This class exists
solely to register the ``QianfanOCRForConditionalGeneration`` architecture
name that appears in the model's config.json.
"""
def _patch_quant_config(
self, config: PretrainedConfig, quant_config: QuantizationConfig
) -> None:
super()._patch_quant_config(config, quant_config)
# ignore vit layers to preserve model performance
if isinstance(quant_config, Fp8Config):
_FP8_IGNORED_LAYERS = [
*(
layer
for i in range(config.vision_config.num_hidden_layers)
for layer in [
f"vision_model.encoder.layers.{i}.attn.qkv",
f"vision_model.encoder.layers.{i}.attn.proj",
f"vision_model.encoder.layers.{i}.mlp.fc1",
f"vision_model.encoder.layers.{i}.mlp.fc2",
]
),
"language_model.lm_head",
"mlp1.1",
"mlp1.3",
]
for layer in _FP8_IGNORED_LAYERS:
if layer not in quant_config.ignored_layers:
quant_config.ignored_layers.append(layer)
+4
View File
@@ -511,6 +511,10 @@ _MULTIMODAL_MODELS = {
"Phi4ForCausalLMV": ("phi4siglip", "Phi4ForCausalLMV"),
"Phi4MMForCausalLM": ("phi4mm", "Phi4MMForCausalLM"),
"PixtralForConditionalGeneration": ("pixtral", "PixtralForConditionalGeneration"),
"QianfanOCRForConditionalGeneration": (
"qianfan_ocr",
"QianfanOCRForConditionalGeneration",
),
"QwenVLForConditionalGeneration": ("qwen_vl", "QwenVLForConditionalGeneration"),
"Qwen2VLForConditionalGeneration": ("qwen2_vl", "Qwen2VLForConditionalGeneration"),
"Qwen2_5_VLForConditionalGeneration": (
+36
View File
@@ -545,6 +545,42 @@ class Platform:
dtype=kv_cache_dtype,
kv_quant_mode=kv_quant_mode,
).page_size_bytes
elif cache_config.cache_dtype.startswith("turboquant_"):
# TQ has a packed K|V layout; the standard FullAttentionSpec
# formula over-sizes it and trips unify_kv_cache_spec_page_size
# when all attention layers are TQ. With mixed skip+TQ the skip
# layers still use the standard layout — take max so mamba
# padding covers the largest actual page.
from vllm.model_executor.layers.quantization.turboquant.config import (
TurboQuantConfig,
)
from vllm.v1.kv_cache_interface import TQFullAttentionSpec
tq_cfg = TurboQuantConfig.from_cache_dtype(
cache_config.cache_dtype, model_config.get_head_size()
)
tq_page = TQFullAttentionSpec(
block_size=1,
num_kv_heads=model_config.get_num_kv_heads(parallel_config),
head_size=model_config.get_head_size(),
head_size_v=model_config.get_head_size(),
dtype=kv_cache_dtype,
kv_quant_mode=kv_quant_mode,
tq_slot_size=tq_cfg.slot_size_aligned,
).page_size_bytes
if cache_config.kv_cache_dtype_skip_layers:
skip_page = FullAttentionSpec(
block_size=1,
num_kv_heads=model_config.get_num_kv_heads(parallel_config),
head_size=model_config.get_head_size(),
dtype=model_config.dtype,
).page_size_bytes
# lcm, not max: skip_page is often not a multiple of
# tq_page, so max would leave per-layer page sizes
# un-unifiable downstream.
attn_page_size_1_token = lcm(tq_page, skip_page)
else:
attn_page_size_1_token = tq_page
else:
attn_page_size_1_token = FullAttentionSpec(
block_size=1,
+1
View File
@@ -125,6 +125,7 @@ _CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = LazyConfigDict(
step3_vl="Step3VLConfig",
step3_text="Step3TextConfig",
step3p5="Step3p5Config",
qianfan_ocr="QianfanOCRConfig",
qwen3_asr="Qwen3ASRConfig",
qwen3_next="Qwen3NextConfig",
qwen3_5="Qwen3_5Config",
@@ -70,6 +70,8 @@ _CLASS_TO_MODULE: dict[str, str] = {
"Step3VisionEncoderConfig": "vllm.transformers_utils.configs.step3_vl",
"Step3TextConfig": "vllm.transformers_utils.configs.step3_vl",
"Step3p5Config": "vllm.transformers_utils.configs.step3p5",
"QianfanOCRConfig": "vllm.transformers_utils.configs.qianfan_ocr",
"QianfanOCRVisionConfig": "vllm.transformers_utils.configs.qianfan_ocr",
"Qwen3ASRConfig": "vllm.transformers_utils.configs.qwen3_asr",
"Qwen3NextConfig": "vllm.transformers_utils.configs.qwen3_next",
"Qwen3_5Config": "vllm.transformers_utils.configs.qwen3_5",
@@ -135,6 +137,8 @@ __all__ = [
"Step3VisionEncoderConfig",
"Step3TextConfig",
"Step3p5Config",
"QianfanOCRConfig",
"QianfanOCRVisionConfig",
"Qwen3ASRConfig",
"Qwen3NextConfig",
"Qwen3_5Config",
@@ -101,7 +101,6 @@ else:
class DeepseekVLV2Config(PretrainedConfig):
model_type = "deepseek_vl_v2"
architectures: list[str] | None = None
tile_tag: str = "2D"
global_view_pos: str = "head"
@@ -114,17 +113,11 @@ class DeepseekVLV2Config(PretrainedConfig):
candidate_resolutions: tuple[tuple[int, int]] = ((384, 384),),
**kwargs,
):
if "architectures" not in kwargs:
kwargs["architectures"] = ["DeepseekVLV2ForCausalLM"]
architectures = kwargs.setdefault("architectures", ["DeepseekVLV2ForCausalLM"])
vision_config = kwargs.pop("vision_config", {})
self.vision_config = VisionEncoderConfig(**vision_config)
projector_config = kwargs.pop("projector_config", {})
self.projector_config = MlpProjectorConfig(**projector_config)
language_config = kwargs.pop("language_config", {})
self.text_config = DeepseekVLV2TextConfig(**language_config)
self.vision_config = VisionEncoderConfig(**kwargs.pop("vision_config", {}))
self.projector_config = MlpProjectorConfig(**kwargs.pop("projector_config", {}))
self.text_config = DeepseekVLV2TextConfig(**kwargs.pop("language_config", {}))
self.tile_tag = tile_tag
self.global_view_pos = global_view_pos
@@ -132,8 +125,8 @@ class DeepseekVLV2Config(PretrainedConfig):
self.vocab_size = self.text_config.vocab_size
# update model_type for OCR models
if "DeepseekOCRForCausalLM" in kwargs["architectures"]:
self.model_type = "deepseek_ocr"
elif "DeepseekOCR2ForCausalLM" in kwargs["architectures"]:
self.model_type = "deepseek_ocr2"
if "DeepseekOCRForCausalLM" in architectures:
kwargs["model_type"] = "deepseek_ocr"
elif "DeepseekOCR2ForCausalLM" in architectures:
kwargs["model_type"] = "deepseek_ocr2"
super().__init__(**kwargs)
@@ -0,0 +1,105 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from typing import Any
from transformers import PretrainedConfig
from transformers.models.auto import CONFIG_MAPPING
class QianfanOCRVisionConfig(PretrainedConfig):
model_type = "qianfan_ocr_vision"
def __init__(
self,
hidden_size: int = 1024,
intermediate_size: int = 4096,
num_hidden_layers: int = 24,
num_attention_heads: int = 16,
num_channels: int = 3,
image_size: int = 448,
patch_size: int = 14,
hidden_act: str = "gelu",
layer_norm_eps: float = 1e-6,
attention_dropout: float = 0.0,
drop_path_rate: float = 0.1,
qkv_bias: bool = True,
qk_normalization: bool = False,
norm_type: str = "layer_norm",
initializer_range: float = 0.02,
initializer_factor: float = 0.1,
use_mask_token: bool = False,
use_mean_pooling: bool = True,
**kwargs: Any,
):
super().__init__(**kwargs)
self.hidden_size = hidden_size
self.intermediate_size = intermediate_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.num_channels = num_channels
self.image_size = image_size
self.patch_size = patch_size
self.hidden_act = hidden_act
self.layer_norm_eps = layer_norm_eps
self.attention_dropout = attention_dropout
self.drop_path_rate = drop_path_rate
self.qkv_bias = qkv_bias
self.qk_normalization = qk_normalization
self.norm_type = norm_type
self.initializer_range = initializer_range
self.initializer_factor = initializer_factor
self.use_mask_token = use_mask_token
self.use_mean_pooling = use_mean_pooling
class QianfanOCRConfig(PretrainedConfig):
model_type = "qianfan_ocr"
def __init__(
self,
vision_config: dict | None = None,
text_config: dict | None = None,
downsample_ratio: float = 0.5,
dynamic_image_size: bool = True,
force_image_size: int = 448,
image_token_id: int = 151671,
max_dynamic_patch: int = 12,
min_dynamic_patch: int = 1,
pad2square: bool = False,
ps_version: str = "v2",
select_layer: int = -1,
template: str = "internvl2_5",
use_thumbnail: bool = True,
tie_word_embeddings: bool = False,
**kwargs: Any,
):
super().__init__(**kwargs)
if isinstance(vision_config, dict):
self.vision_config = QianfanOCRVisionConfig(**vision_config)
elif vision_config is None:
self.vision_config = QianfanOCRVisionConfig()
else:
self.vision_config = vision_config
if isinstance(text_config, dict):
model_type = text_config.get("model_type", "qwen3")
self.text_config = CONFIG_MAPPING[model_type](**text_config)
elif text_config is None:
self.text_config = CONFIG_MAPPING["qwen3"]()
else:
self.text_config = text_config
self.downsample_ratio = downsample_ratio
self.dynamic_image_size = dynamic_image_size
self.force_image_size = force_image_size
self.image_token_id = image_token_id
self.max_dynamic_patch = max_dynamic_patch
self.min_dynamic_patch = min_dynamic_patch
self.pad2square = pad2square
self.ps_version = ps_version
self.select_layer = select_layer
self.template = template
self.use_thumbnail = use_thumbnail
self.tie_word_embeddings = tie_word_embeddings
+113
View File
@@ -0,0 +1,113 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Monitor unexpected Triton kernel JIT compilation during inference.
After server warmup completes, any Triton JIT compilation or autotuning
event indicates a cache miss or unexpected input shape that causes a
latency spike. This module registers hooks in the Triton runtime to
detect and log such events so they can be investigated.
Currently monitors:
- Triton ``@triton.autotune`` cache misses (via ``knobs.autotuning.print``)
- Triton ``@triton.jit`` first-time compilations
(via ``knobs.runtime.jit_post_compile_hook``)
"""
import os
from vllm.logger import init_logger
from vllm.triton_utils.importing import HAS_TRITON
logger = init_logger(__name__)
_active: bool = False
def is_active() -> bool:
"""Return whether the JIT compilation monitor is currently active."""
return _active
def activate() -> None:
"""Enable JIT compilation monitoring after warmup.
Call once per worker process at the end of
:func:`compile_or_warm_up_model`. After activation every Triton
kernel compilation or autotuning benchmark that happens during
inference will be logged as a warning.
Safe to call multiple times subsequent calls are no-ops.
If the user has explicitly set ``TRITON_PRINT_AUTOTUNING=0`` in
their environment, autotuning printing is left disabled; the JIT
compilation hook is still registered regardless.
"""
global _active
if _active:
return
_active = True
_setup_triton_autotuning_print()
_setup_triton_jit_hook()
logger.info(
"Kernel JIT monitor activated — Triton JIT compilations "
"during inference will be logged as warnings."
)
# ------------------------------------------------------------------
# Triton autotuning print
# ------------------------------------------------------------------
def _setup_triton_autotuning_print() -> None:
"""Enable ``TRITON_PRINT_AUTOTUNING`` unless the user opted out."""
if not HAS_TRITON:
return
from triton import knobs # type: ignore[import-untyped]
user_val = os.environ.get("TRITON_PRINT_AUTOTUNING")
if user_val == "0":
logger.debug(
"TRITON_PRINT_AUTOTUNING=0 set by user — "
"autotuning messages will stay suppressed."
)
return
knobs.autotuning.print = True
# ------------------------------------------------------------------
# Triton JIT compilation hook
# ------------------------------------------------------------------
def _setup_triton_jit_hook() -> None:
"""Register a ``jit_post_compile_hook`` that warns on compilation."""
if not HAS_TRITON:
return
from triton import knobs # type: ignore[import-untyped]
existing_hook = knobs.runtime.jit_post_compile_hook
def _on_jit_compile(**kwargs):
# `jit_post_compile_hook` is Triton internal API and its
# signature has changed across releases (kwargs added/renamed).
# Accept **kwargs so an upstream change cannot crash this hook
# with TypeError, and forward the full kwarg set to any
# pre-existing hook unchanged.
fn = kwargs.get("fn")
fn_name = getattr(fn, "name", "<unknown>")
logger.warning_once(
"Triton kernel JIT compilation during inference: %s. "
"This causes a latency spike; consider extending warmup "
"to cover this shape/config.",
fn_name,
)
if existing_hook is not None:
return existing_hook(**kwargs)
return None
knobs.runtime.jit_post_compile_hook = _on_jit_compile
+12 -3
View File
@@ -29,7 +29,12 @@ from vllm.v1.kv_cache_interface import AttentionSpec, CrossAttentionSpec
logger = init_logger(__name__)
_CPU_ARCH_PREFER_MIXED_BATCH = (CpuArchEnum.X86, CpuArchEnum.ARM, CpuArchEnum.S390X)
_CPU_ARCH_PREFER_MIXED_BATCH = (
CpuArchEnum.X86,
CpuArchEnum.ARM,
CpuArchEnum.S390X,
CpuArchEnum.POWERPC,
)
class CPUAttentionBackend(AttentionBackend):
@@ -510,8 +515,10 @@ def _get_attn_isa(
)
return "vec16"
supports_amx = torch.cpu._is_amx_tile_supported()
supports_arm = current_platform.get_cpu_architecture() == CpuArchEnum.ARM
supports_vxe = current_platform.get_cpu_architecture() == CpuArchEnum.S390X
arch = current_platform.get_cpu_architecture()
supports_arm = arch == CpuArchEnum.ARM
supports_vxe = arch == CpuArchEnum.S390X
supports_vsx = arch == CpuArchEnum.POWERPC
supports_avx512 = torch.cpu._is_avx512_supported()
if fp8_kv and not supports_amx and not supports_avx512:
raise NotImplementedError(
@@ -525,6 +532,8 @@ def _get_attn_isa(
return "neon"
elif supports_vxe:
return "vxe"
elif supports_vsx:
return "vsx"
else:
return "vec"
else:
+10 -2
View File
@@ -331,12 +331,20 @@ class OffloadingSpec(ABC):
assert kv_transfer_config is not None
self.extra_config = kv_transfer_config.kv_connector_extra_config
parallel_config = vllm_config.parallel_config
context_parallel_factor = (
parallel_config.decode_context_parallel_size
* parallel_config.prefill_context_parallel_size
)
# block size used by vLLM for hashing request tokens for the sake
# of enabling prefix caching
self.hash_block_size = vllm_config.cache_config.block_size
self.hash_block_size = (
vllm_config.cache_config.block_size * context_parallel_factor
)
# gpu block size per group
self.gpu_block_size: tuple[int, ...] = tuple(
kv_cache_group.kv_cache_spec.block_size
kv_cache_group.kv_cache_spec.block_size * context_parallel_factor
for kv_cache_group in kv_cache_config.kv_cache_groups
)
+48
View File
@@ -82,6 +82,11 @@ class TopKTopPSampler(nn.Module):
self.forward = self.forward_native
else:
self.forward = self.forward_cpu
elif current_platform.is_xpu():
if envs.VLLM_XPU_USE_SAMPLER_KERNEL:
self.forward = self.forward_xpu
else:
self.forward = self.forward_native
elif (
logprobs_mode not in ("processed_logits", "processed_logprobs")
and rocm_aiter_ops.is_enabled()
@@ -243,6 +248,49 @@ class TopKTopPSampler(nn.Module):
return torch.multinomial(renorm_probs, num_samples=1).view(-1)
raise RuntimeError("aiter_sample was called with no active top-k or top-p.")
def forward_xpu(
self,
logits: torch.Tensor,
generators: dict[int, torch.Generator],
k: torch.Tensor | None,
p: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
if generators:
logger.warning_once(
"xpu kernel topk_topp_sampler does not support "
"per-request generators. Falling back to "
"PyTorch-native implementation."
)
return self.forward_native(logits, generators, k, p)
random_sampled = torch.empty(
logits.shape[0], dtype=torch.int64, device=logits.device
)
logits_to_return = None
if (
self.logprobs_mode == "processed_logits"
or self.logprobs_mode == "processed_logprobs"
):
logits_to_return = torch.empty_like(logits)
assert len(generators) != logits.shape[0], (
"xpu kernel topk_topp_sampler does not support batch-wise generators."
)
generator = torch.xpu.default_generators[logits.device.index]
state = generator.get_state()
seed, offset = state.view(torch.int64)
seeds = torch.tensor(
[seed, offset], dtype=torch.int64, device=torch.device("cpu")
)
# The XPU kernel expects k as int64 (Long), but the input batch
# stores top_k as int32. Cast here to avoid dtype mismatch.
if k is not None:
k = k.to(torch.int64)
torch.ops.vllm.xpu_topk_topp_sampler(
random_sampled, logits_to_return, logits, k, p, self.logprobs_mode, seeds
)
return random_sampled, logits_to_return
# Note: this is a workaround for
# https://github.com/pytorch/pytorch/pull/151218
+18 -4
View File
@@ -76,6 +76,8 @@ def gumbel_block_argmax(
pos_ptr,
processed_logits_ptr,
processed_logits_stride,
processed_logits_col_ptr,
vocab_size,
APPLY_TEMPERATURE: tl.constexpr,
):
req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx)
@@ -88,8 +90,15 @@ def gumbel_block_argmax(
if processed_logits_ptr is not None:
# Store the temperature-applied logits.
if processed_logits_col_ptr is not None:
col = tl.load(processed_logits_col_ptr)
else:
col = 0
tl.store(
processed_logits_ptr + req_state_idx * processed_logits_stride + block,
processed_logits_ptr
+ req_state_idx * processed_logits_stride
+ col * vocab_size
+ block,
logits,
mask=mask,
)
@@ -121,6 +130,7 @@ def _gumbel_sample_kernel(
local_max_stride,
processed_logits_ptr,
processed_logits_stride,
processed_logits_col_ptr,
logits_ptr,
logits_stride,
expanded_idx_mapping_ptr,
@@ -153,6 +163,8 @@ def _gumbel_sample_kernel(
pos_ptr,
processed_logits_ptr,
processed_logits_stride,
processed_logits_col_ptr,
vocab_size,
APPLY_TEMPERATURE=APPLY_TEMPERATURE,
)
token_id = block_idx * BLOCK_SIZE + idx
@@ -167,7 +179,8 @@ def gumbel_sample(
seed: torch.Tensor, # [max_num_reqs]
pos: torch.Tensor, # [num_tokens]
apply_temperature: bool,
processed_logits_out: torch.Tensor | None = None, # [num_reqs, vocab_size]
output_processed_logits: torch.Tensor | None = None,
output_processed_logits_col: torch.Tensor | None = None,
) -> torch.Tensor:
num_tokens, vocab_size = logits.shape
BLOCK_SIZE = 1024
@@ -179,8 +192,9 @@ def gumbel_sample(
local_argmax.stride(0),
local_max,
local_max.stride(0),
processed_logits_out,
processed_logits_out.stride(0) if processed_logits_out is not None else 0,
output_processed_logits,
output_processed_logits.stride(0) if output_processed_logits is not None else 0,
output_processed_logits_col,
logits,
logits.stride(0),
expanded_idx_mapping,
+175 -114
View File
@@ -89,9 +89,13 @@ class EagleSpeculator:
dtype=torch.int64,
device=device,
)
self.current_draft_step = torch.tensor(0, dtype=torch.int64, device=device)
self.last_token_indices = torch.zeros(
self.max_num_reqs, dtype=torch.int64, device=device
)
self.arange = torch.arange(
self.max_num_reqs + 1, dtype=torch.int32, device="cpu"
)
self.supports_mm_inputs = MULTIMODAL_REGISTRY.supports_multimodal_inputs(
self.draft_model_config
@@ -228,9 +232,10 @@ class EagleSpeculator:
logits: torch.Tensor,
idx_mapping: torch.Tensor,
pos: torch.Tensor,
step: int,
draft_step: torch.Tensor,
draft_logits: torch.Tensor | None,
) -> torch.Tensor:
if self.draft_logits is not None:
if draft_logits is not None:
# NOTE(woosuk): We must add 1 to the positions to match the Gumbel noise
# used for draft and target sampling.
return gumbel_sample(
@@ -240,7 +245,8 @@ class EagleSpeculator:
self.seeds,
pos + 1,
apply_temperature=True,
processed_logits_out=self.draft_logits[:, step],
output_processed_logits=draft_logits,
output_processed_logits_col=draft_step,
)
else:
return logits.argmax(dim=-1)
@@ -274,11 +280,63 @@ class EagleSpeculator:
logits,
idx_mapping,
pos,
step=0,
self.current_draft_step,
self.draft_logits,
)
self.hidden_states[:num_reqs] = hidden_states[last_token_indices]
self.input_buffers.positions[:num_reqs] = pos
def multi_step_decode(
self,
num_reqs: int,
skip_attn: bool,
batch_desc: BatchExecutionDescriptor,
num_tokens_across_dp: torch.Tensor | None,
) -> None:
positions = self.input_buffers.positions[:num_reqs]
query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1]
idx_mapping = self.idx_mapping[:num_reqs]
for step in range(1, self.num_speculative_steps):
attn_metadata = None
slot_mappings_by_layer = None
if not skip_attn:
# Build attention metadata and slot mappings for each draft
# decode step. It is necessary to rebuild the attention
# metadata even when replaying the FULL graph so that any
# attention metadata builder state is updated.
slot_mappings = self.block_tables.compute_slot_mappings(
idx_mapping,
query_start_loc,
positions,
batch_desc.num_tokens,
)
slot_mappings_by_layer = build_slot_mappings_by_layer(
slot_mappings, self.kv_cache_config
)
attn_metadata = self._build_draft_attn_metadata(
num_reqs=num_reqs,
num_reqs_padded=batch_desc.num_reqs or num_reqs,
num_tokens_padded=batch_desc.num_tokens,
)
# Update the current draft step.
self.current_draft_step.fill_(step)
# Generate draft tokens for the current step.
if batch_desc.cg_mode == CUDAGraphMode.FULL:
assert self.decode_cudagraph_manager is not None
self.decode_cudagraph_manager.run_fullgraph(batch_desc)
else:
self.generate_draft(
num_reqs,
batch_desc.num_tokens,
attn_metadata,
slot_mappings_by_layer,
num_tokens_across_dp=num_tokens_across_dp,
cudagraph_runtime_mode=batch_desc.cg_mode,
)
def generate_draft(
self,
num_reqs: int,
@@ -288,59 +346,52 @@ class EagleSpeculator:
num_tokens_across_dp: torch.Tensor | None,
cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
) -> None:
pos = self.input_buffers.positions[:num_reqs]
query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1]
idx_mapping = self.idx_mapping[:num_reqs]
for step in range(1, self.num_speculative_steps):
# Run the eagle model.
last_hidden_states, hidden_states = self.run_model(
num_tokens_padded,
attn_metadata,
slot_mappings,
num_tokens_across_dp,
cudagraph_runtime_mode,
)
last_hidden_states = last_hidden_states[:num_reqs]
hidden_states = hidden_states[:num_reqs]
logits = self.model.compute_logits(last_hidden_states)
positions = self.input_buffers.positions[:num_reqs]
# Run the eagle model forward pass.
last_hidden_states, hidden_states = self.run_model(
num_tokens_padded,
attn_metadata,
slot_mappings,
num_tokens_across_dp,
cudagraph_runtime_mode,
)
last_hidden_states = last_hidden_states[:num_reqs]
draft_tokens = self._sample_draft(
logits,
idx_mapping,
pos,
step=step,
)
self.draft_tokens[:num_reqs, step] = draft_tokens
# Sample the draft tokens.
logits = self.model.compute_logits(last_hidden_states)
draft_tokens = self._sample_draft(
logits,
idx_mapping,
positions,
self.current_draft_step,
self.draft_logits,
)
if step < self.num_speculative_steps - 1:
# Update the inputs for the next step.
update_eagle_inputs(
draft_tokens,
hidden_states,
self.input_buffers,
self.hidden_states,
self.max_model_len,
)
if attn_metadata is not None:
self.block_tables.compute_slot_mappings(
idx_mapping, query_start_loc, pos, num_tokens_padded
)
# Update the inputs for the next step.
update_eagle_draft_inputs(
draft_tokens,
self.current_draft_step,
hidden_states,
self.draft_tokens,
self.hidden_states,
self.input_buffers,
num_reqs,
self.max_model_len,
self.num_speculative_steps,
)
def _build_draft_attn_metadata(
self,
num_reqs: int,
num_reqs_padded: int,
num_tokens_padded: int,
max_query_len: int,
) -> dict[str, Any] | None:
if not self.draft_attn_layer_names:
return None
query_start_loc_cpu = (
torch.arange(num_reqs_padded + 1, dtype=torch.int32, device="cpu").clamp_(
max=num_reqs
)
* max_query_len
query_start_loc_cpu = torch.clamp(
self.arange[: num_reqs_padded + 1], max=num_reqs
)
block_tables = [
x[:num_reqs_padded] for x in self.block_tables.input_block_tables
@@ -354,7 +405,7 @@ class EagleSpeculator:
: num_reqs_padded + 1
],
query_start_loc_cpu=query_start_loc_cpu,
max_query_len=max_query_len,
max_query_len=1,
seq_lens=self.input_buffers.seq_lens[:num_reqs_padded],
max_seq_len=self.max_model_len,
block_tables=block_tables,
@@ -373,7 +424,7 @@ class EagleSpeculator:
self.last_token_indices.zero_()
# Capture the prefill routine (model forward + compute_logits +
# gumbel_sample).
# sample).
# For FULL graphs, the entire routine is recorded as one graph.
# For PIECEWISE, only the model's compiled regions are captured
# and the rest (compute_logits, gumbel_sample) runs eagerly.
@@ -387,10 +438,9 @@ class EagleSpeculator:
if self.num_speculative_steps == 1:
return
# Capture the decode draft generation loop (model forward +
# compute_logits + gumbel_sample + update_eagle_inputs, for
# each step). For FULL graphs, the entire multi-step loop is
# recorded as one graph.
# Capture the decode draft generation routine (model forward +
# compute_logits + sample + update_eagle_inputs) for a single
# step.
assert self.decode_cudagraph_manager is not None
self.decode_cudagraph_manager.capture(
self.generate_draft,
@@ -461,9 +511,10 @@ class EagleSpeculator:
# Get the input ids and last token indices for the speculator.
prepare_eagle_inputs(
self.last_token_indices,
self.current_draft_step,
self.input_buffers,
input_batch,
self.last_token_indices,
num_sampled,
num_rejected,
last_sampled,
@@ -473,12 +524,18 @@ class EagleSpeculator:
# When all requests are decoding (no true prefills), each has
# num_speculative_steps + 1 tokens, enabling FULL graph replay.
# Mixed or prefill-only batches fall back to PIECEWISE.
uniform_token_count = get_uniform_token_count(
num_reqs,
# Use the actual number of tokens without padding added by
# the target model during FULL cudagraph.
input_batch.num_tokens,
max_query_len,
)
prefill_batch_desc, num_tokens_across_dp = dispatch_cg_and_sync_dp(
self.prefill_cudagraph_manager,
num_reqs,
num_tokens,
get_uniform_token_count(num_reqs, num_tokens, max_query_len),
uniform_token_count,
dp_size=self.dp_size,
dp_rank=self.dp_rank,
need_eager=is_profile,
@@ -528,48 +585,21 @@ class EagleSpeculator:
need_eager=is_profile,
)
attn_metadata_updated = None
slot_mappings_updated = None
if not (dummy_run and skip_attn_for_dummy_run):
# Build attention metadata and slot mappings for the draft
# decode steps. It is necessary to rebuild the attention
# metadata even when replaying the FULL graph so that any
# attention metadata builder state is updated.
slot_mappings = self.block_tables.compute_slot_mappings(
self.idx_mapping[:num_reqs],
self.input_buffers.query_start_loc[: num_reqs + 1],
self.input_buffers.positions[:num_reqs],
decode_batch_desc.num_tokens,
)
slot_mappings_updated = build_slot_mappings_by_layer(
slot_mappings, self.kv_cache_config
)
attn_metadata_updated = self._build_draft_attn_metadata(
num_reqs=num_reqs,
num_reqs_padded=decode_batch_desc.num_reqs or num_reqs,
num_tokens_padded=decode_batch_desc.num_tokens,
max_query_len=1,
)
# Generate the remaining num_speculative_steps - 1 draft tokens.
self.multi_step_decode(
num_reqs,
dummy_run and skip_attn_for_dummy_run,
decode_batch_desc,
num_tokens_across_dp,
)
if decode_batch_desc.cg_mode == CUDAGraphMode.FULL:
# Replay the full graph for draft generation.
assert self.decode_cudagraph_manager is not None
self.decode_cudagraph_manager.run_fullgraph(decode_batch_desc)
else:
self.generate_draft(
num_reqs,
decode_batch_desc.num_tokens,
attn_metadata_updated,
slot_mappings_updated,
num_tokens_across_dp=num_tokens_across_dp,
cudagraph_runtime_mode=decode_batch_desc.cg_mode,
)
return self.draft_tokens[:num_reqs]
@triton.jit
def _prepare_eagle_inputs_kernel(
last_token_indices_ptr,
eagle_current_draft_step_ptr,
eagle_input_ids_ptr,
eagle_positions_ptr,
eagle_query_start_loc_ptr,
@@ -630,6 +660,8 @@ def _prepare_eagle_inputs_kernel(
# Copy sequence lengths.
tl.store(eagle_seq_lens_ptr + req_idx, seq_len)
if req_idx == (num_reqs - 1):
# Reset the current draft step to 0.
tl.store(eagle_current_draft_step_ptr, 0)
# Pad query_start_loc for CUDA graphs.
for i in range(num_reqs, max_num_reqs + 1, BLOCK_SIZE):
block = i + tl.arange(0, BLOCK_SIZE)
@@ -648,10 +680,11 @@ def _prepare_eagle_inputs_kernel(
def prepare_eagle_inputs(
input_buffers: InputBuffers,
input_batch: InputBatch,
# [num_reqs]
last_token_indices: torch.Tensor,
current_draft_step: torch.Tensor,
input_buffers: InputBuffers,
input_batch: InputBatch,
# [num_reqs]
num_sampled: torch.Tensor,
# [num_reqs]
@@ -665,6 +698,7 @@ def prepare_eagle_inputs(
num_reqs = input_batch.num_reqs
_prepare_eagle_inputs_kernel[(num_reqs,)](
last_token_indices,
current_draft_step,
input_buffers.input_ids,
input_buffers.positions,
input_buffers.query_start_loc,
@@ -685,7 +719,7 @@ def prepare_eagle_inputs(
@triton.jit
def _prepare_eagle_docode_kernel(
def _prepare_eagle_decode_kernel(
draft_tokens_ptr,
draft_tokens_stride,
target_seq_lens_ptr,
@@ -742,7 +776,7 @@ def prepare_eagle_decode(
max_num_reqs: int,
):
num_reqs = draft_tokens.shape[0]
_prepare_eagle_docode_kernel[(num_reqs + 1,)](
_prepare_eagle_decode_kernel[(num_reqs + 1,)](
draft_tokens,
draft_tokens.stride(0),
target_seq_lens,
@@ -758,36 +792,55 @@ def prepare_eagle_decode(
@triton.jit
def _update_eagle_inputs_kernel(
def _update_eagle_draft_inputs_kernel(
output_draft_tokens_ptr,
output_draft_tokens_stride,
next_input_hidden_states_ptr,
next_input_hidden_states_stride,
input_ids_ptr,
positions_ptr,
input_hidden_states_ptr,
input_hidden_states_stride,
seq_lens_ptr,
max_model_len,
draft_tokens_ptr,
output_hidden_states_ptr,
output_hidden_states_stride,
current_draft_step_ptr,
hidden_states_ptr,
hidden_states_stride,
hidden_size,
max_model_len,
num_speculative_steps,
BLOCK_SIZE: tl.constexpr,
):
req_idx = tl.program_id(0)
# Draft token -> Input ID.
# Write the sampled draft token into self.draft_tokens[req_idx, step].
draft_token = tl.load(draft_tokens_ptr + req_idx)
step = tl.load(current_draft_step_ptr)
tl.store(
output_draft_tokens_ptr + req_idx * output_draft_tokens_stride + step,
draft_token,
)
if step >= num_speculative_steps - 1:
# This is the final step. Skip updating draft forward inputs.
return
# Write the sampled draft token into the input ids tensor for the next
# forward pass.
tl.store(input_ids_ptr + req_idx, draft_token)
# Output hidden states -> Input hidden states.
# Copy hidden states into the input hidden states tensor for the next
# forward pass.
for i in range(0, hidden_size, BLOCK_SIZE):
block = i + tl.arange(0, BLOCK_SIZE)
mask = block < hidden_size
output_hidden_states = tl.load(
output_hidden_states_ptr + req_idx * output_hidden_states_stride + block,
hidden_states = tl.load(
hidden_states_ptr + req_idx * hidden_states_stride + block,
mask=mask,
)
tl.store(
input_hidden_states_ptr + req_idx * input_hidden_states_stride + block,
output_hidden_states,
next_input_hidden_states_ptr
+ req_idx * next_input_hidden_states_stride
+ block,
hidden_states,
mask=mask,
)
@@ -803,24 +856,32 @@ def _update_eagle_inputs_kernel(
tl.store(seq_lens_ptr + req_idx, seq_len)
def update_eagle_inputs(
def update_eagle_draft_inputs(
draft_tokens: torch.Tensor,
output_hidden_states: torch.Tensor,
input_buffers: InputBuffers,
current_draft_step: torch.Tensor,
hidden_states: torch.Tensor,
output_draft_tokens: torch.Tensor,
next_input_hidden_states: torch.Tensor,
input_buffers: InputBuffers,
num_reqs: int,
max_model_len: int,
num_speculative_steps: int,
):
num_reqs, hidden_size = output_hidden_states.shape
_update_eagle_inputs_kernel[(num_reqs,)](
_, hidden_size = hidden_states.shape
_update_eagle_draft_inputs_kernel[(num_reqs,)](
output_draft_tokens,
output_draft_tokens.stride(0),
next_input_hidden_states,
next_input_hidden_states.stride(0),
input_buffers.input_ids,
input_buffers.positions,
input_buffers.seq_lens,
draft_tokens,
current_draft_step,
hidden_states,
hidden_states.stride(0),
input_buffers.seq_lens,
max_model_len,
draft_tokens,
output_hidden_states,
output_hidden_states.stride(0),
hidden_size,
max_model_len,
num_speculative_steps,
BLOCK_SIZE=1024,
)
@@ -392,8 +392,10 @@ def _resample_kernel(
temp_ptr,
seed_ptr,
pos_ptr,
None,
0,
None, # processed_logits_ptr
0, # processed_logits_stride
None, # processed_logits_col_ptr
vocab_size,
APPLY_TEMPERATURE=False,
)
token_id = block_idx * BLOCK_SIZE + idx
+8
View File
@@ -711,6 +711,14 @@ class Worker(WorkerBase):
# the model initialization and profiling.
set_random_seed(self.model_config.seed)
# All warmup is done — start monitoring for unexpected JIT
# compilations that would cause latency spikes during inference.
from vllm.triton_utils.jit_monitor import (
activate as activate_triton_jit_monitor,
)
activate_triton_jit_monitor()
return CompilationTimes(
language_model=self.compilation_config.compilation_time,
encoder=self.compilation_config.encoder_compilation_time,
+3 -1
View File
@@ -214,7 +214,9 @@ def _make_metadata_with_slice(
seq_lens_cpu_upper_bound[-1] -= tokens_skipped
assert seq_lens_cpu_upper_bound is not None
max_seq_len = int(seq_lens_cpu_upper_bound.max())
# Preserve the max_seq_len override set during CUDA-graph capture so
# the attention backend selects the correct kernel for SWA layers.
max_seq_len = max(int(seq_lens_cpu_upper_bound.max()), attn_metadata.max_seq_len)
num_requests = request_slice.stop - request_slice.start
num_actual_tokens = token_slice.stop - token_slice.start