From d78650cf973b9a1d8c03160c691f98725f8edb16 Mon Sep 17 00:00:00 2001 From: Itay Alroy <75032521+itayalroy@users.noreply.github.com> Date: Wed, 17 Jun 2026 07:29:34 +0300 Subject: [PATCH 01/37] [CI][NIXL] Pin NIXL to 1.2.0 (#45843) Signed-off-by: Itay Alroy Signed-off-by: Itay Alroy <75032521+itayalroy@users.noreply.github.com> Co-authored-by: ovidiusm --- requirements/kv_connectors.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements/kv_connectors.txt b/requirements/kv_connectors.txt index 7a5b5f25c37..e0d494e9f21 100644 --- a/requirements/kv_connectors.txt +++ b/requirements/kv_connectors.txt @@ -2,5 +2,5 @@ lmcache >= 0.3.9 # CuPy 14.1.0 imports pytest from cupy.testing._random. Use <14.1.0 # until a fixed newer release is verified for runtime images. cupy-cuda13x < 14.1.0 -nixl >= 1.1.0 # Required for disaggregated prefill +nixl == 1.2.0 # Required for disaggregated prefill mooncake-transfer-engine >= 0.3.8 From 4c626633159887b0f2c962058c17c78f1434556d Mon Sep 17 00:00:00 2001 From: Thien Tran Date: Wed, 17 Jun 2026 12:38:03 +0800 Subject: [PATCH 02/37] [M3] Enable FP8 sparse GQA (#45744) Signed-off-by: Thien Tran Co-authored-by: Yongye Zhu --- cmake/external_projects/fmha_sm100.cmake | 18 ++- ...minimax_m3_qknorm_rope_kv_insert_kernel.cu | 120 ++++++++++++------ csrc/libtorch_stable/ops.h | 6 +- csrc/libtorch_stable/torch_bindings.cpp | 3 +- setup.py | 2 +- ..._fused_minimax_m3_qknorm_rope_kv_insert.py | 62 +++++++-- vllm/_custom_ops.py | 7 +- vllm/models/minimax_m3/amd/model.py | 70 +--------- .../minimax_m3/common/sparse_attention.py | 20 ++- vllm/models/minimax_m3/nvidia/model.py | 49 +------ 10 files changed, 179 insertions(+), 178 deletions(-) diff --git a/cmake/external_projects/fmha_sm100.cmake b/cmake/external_projects/fmha_sm100.cmake index 15610552f23..4a2414f5b83 100644 --- a/cmake/external_projects/fmha_sm100.cmake +++ b/cmake/external_projects/fmha_sm100.cmake @@ -32,19 +32,17 @@ message(STATUS "fmha_sm100 is available at ${fmha_sm100_SOURCE_DIR}") add_custom_target(fmha_sm100) +set(FMHA_SM100_PY_ROOT "${fmha_sm100_SOURCE_DIR}/python/fmha_sm100") + install(FILES - "${fmha_sm100_SOURCE_DIR}/python/fmha_sm100/__init__.py" - "${fmha_sm100_SOURCE_DIR}/python/fmha_sm100/sparse.py" + "${FMHA_SM100_PY_ROOT}/__init__.py" + "${FMHA_SM100_PY_ROOT}/sparse.py" DESTINATION vllm/third_party/fmha_sm100 COMPONENT fmha_sm100) -install(DIRECTORY "${fmha_sm100_SOURCE_DIR}/python/fmha_sm100/cute/" +install(DIRECTORY "${FMHA_SM100_PY_ROOT}/cute/" DESTINATION vllm/third_party/fmha_sm100/cute COMPONENT fmha_sm100 - FILES_MATCHING - REGEX "/__pycache__(/.*)?$" EXCLUDE - REGEX ".*\\.pyc$" EXCLUDE - PATTERN "example.py" EXCLUDE - PATTERN "test_*.py" EXCLUDE - PATTERN "*.py" - PATTERN "build_k2q_csr.cu") + PATTERN "__pycache__" EXCLUDE + PATTERN "*.pyc" EXCLUDE + PATTERN ".git*" EXCLUDE) diff --git a/csrc/libtorch_stable/fused_minimax_m3_qknorm_rope_kv_insert_kernel.cu b/csrc/libtorch_stable/fused_minimax_m3_qknorm_rope_kv_insert_kernel.cu index 5dd610f2878..06c8048cd90 100644 --- a/csrc/libtorch_stable/fused_minimax_m3_qknorm_rope_kv_insert_kernel.cu +++ b/csrc/libtorch_stable/fused_minimax_m3_qknorm_rope_kv_insert_kernel.cu @@ -58,8 +58,15 @@ #include "../cuda_compat.h" #include "../type_convert.cuh" +#include "../attention/dtype_fp8.cuh" #include "dispatch_utils.h" +#ifdef USE_ROCM + #include "../quantization/w8a8/fp8/amd/quant_utils.cuh" +#else + #include "../quantization/w8a8/fp8/nvidia/quant_utils.cuh" +#endif + #ifndef FINAL_MASK #ifdef USE_ROCM #define FINAL_MASK 0xffffffffffffffffULL @@ -186,6 +193,21 @@ __device__ __forceinline__ void storeElems( *reinterpret_cast(dst) = v; } +template +__device__ __forceinline__ void storeCacheElems( + cache_t* __restrict__ dst, float const (&elems)[kElemsPerLane]) { + if constexpr (kv_dt == Fp8KVCacheDataType::kAuto) { + // kAuto means unquantized KV cache here: cache_t == scalar_t, so store the + // model dtype directly. FP8 cache dtypes use the conversion path below. + storeElems(reinterpret_cast(dst), elems); + } else { +#pragma unroll + for (int i = 0; i < kElemsPerLane; i++) { + dst[i] = fp8::scaled_convert(elems[i], 1.0f); + } + } +} + // ──────────────────────────────────────────────────────────────────────────── // Kernel // ──────────────────────────────────────────────────────────────────────────── @@ -202,7 +224,8 @@ __device__ __forceinline__ void storeElems( // V : nkv only if kInsertKV (V-cache insert; no warps in dense) // IQ: niq only if kIsSparse (norm+RoPE) // IK: 1 only if kIsSparse (norm+RoPE; +index-cache insert) -template +template __global__ void fusedMiniMaxM3QNormRopeKVInsertKernel( scalar_t* __restrict__ qkv, // [N, qkv_row] in/out (packs index if sparse) scalar_t* __restrict__ q_out, // [N, nq*128] contiguous, or nullptr @@ -215,7 +238,7 @@ __global__ void fusedMiniMaxM3QNormRopeKVInsertKernel( int64_t const* __restrict__ positions, // [N] i64 int64_t const* __restrict__ slot_mapping, // main K/V slots or nullptr int64_t const* __restrict__ index_slot_mapping, // index K slots/nullptr - scalar_t* __restrict__ kv_cache, // [nb,2,bs,nkv,128] or nullptr + cache_t* __restrict__ kv_cache, // [nb,2,bs,nkv,128] or nullptr scalar_t* __restrict__ index_cache, // [nb*bs, 128] or nullptr float const eps, int const rotary_dim, int const num_tokens, int const nq, int const nkv, int const niq, int const block_size, @@ -355,7 +378,8 @@ __global__ void fusedMiniMaxM3QNormRopeKVInsertKernel( int const kv = isK ? 0 : 1; int64_t const off = b * kv_s_block + kv * kv_s_kv + t * kv_s_token + head * kv_s_head; - storeElems(kv_cache + off + dim_base, elems); + storeCacheElems(kv_cache + off + dim_base, + elems); } } } @@ -373,13 +397,13 @@ __global__ void fusedMiniMaxM3QNormRopeKVInsertKernel( // ──────────────────────────────────────────────────────────────────────────── // Launch wrapper // ──────────────────────────────────────────────────────────────────────────── -template +template void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out, scalar_t const* q_norm_w, scalar_t const* k_norm_w, scalar_t const* iq_norm_w, scalar_t const* ik_norm_w, scalar_t const* cos_sin_cache, int64_t const* positions, int64_t const* slot_mapping, - int64_t const* index_slot_mapping, scalar_t* kv_cache, + int64_t const* index_slot_mapping, cache_t* kv_cache, scalar_t* index_cache, float const eps, int const rotary_dim, int const num_tokens, int const nq, int const nkv, int const niq, @@ -419,7 +443,8 @@ void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out, #define LAUNCH(IS_SPARSE, INSERT) \ cudaLaunchKernelEx( \ &config, \ - fusedMiniMaxM3QNormRopeKVInsertKernel, \ + fusedMiniMaxM3QNormRopeKVInsertKernel, \ qkv, q_out, index_q_out, q_norm_w, k_norm_w, iq_norm_w, ik_norm_w, \ cos_sin_cache, positions, slot_mapping, index_slot_mapping, kv_cache, \ index_cache, eps, rotary_dim, num_tokens, nq, nkv, niq, block_size, \ @@ -428,7 +453,8 @@ void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out, // ROCm: standard kernel launch syntax (no PDL/stream serialization). // clang-format off #define LAUNCH(IS_SPARSE, INSERT) \ - fusedMiniMaxM3QNormRopeKVInsertKernel \ + fusedMiniMaxM3QNormRopeKVInsertKernel \ <<>>( \ qkv, q_out, index_q_out, q_norm_w, k_norm_w, iq_norm_w, \ ik_norm_w, cos_sin_cache, positions, slot_mapping, \ @@ -455,6 +481,33 @@ void launchFusedMiniMaxM3(scalar_t* qkv, scalar_t* q_out, scalar_t* index_q_out, } // namespace minimax_m3_fused_ops } // namespace vllm +#define CALL_FUSED_MINIMAX_M3(_RAW_T, CACHE_T, KV_DTYPE) \ + vllm::minimax_m3_fused_ops::launchFusedMiniMaxM3( \ + reinterpret_cast(qkv.data_ptr()), \ + q_out.has_value() ? reinterpret_cast(q_out->data_ptr()) : nullptr, \ + index_q_out.has_value() ? reinterpret_cast(index_q_out->data_ptr()) \ + : nullptr, \ + reinterpret_cast(q_norm_weight.data_ptr()), \ + reinterpret_cast(k_norm_weight.data_ptr()), \ + has_index ? reinterpret_cast(index_q_norm_weight->data_ptr()) \ + : nullptr, \ + has_index ? reinterpret_cast(index_k_norm_weight->data_ptr()) \ + : nullptr, \ + reinterpret_cast(cos_sin_cache.data_ptr()), \ + reinterpret_cast(positions.data_ptr()), \ + insert_kv ? reinterpret_cast(slot_mapping->data_ptr()) \ + : nullptr, \ + insert_kv ? reinterpret_cast( \ + effective_index_slot_mapping->data_ptr()) \ + : nullptr, \ + insert_kv ? reinterpret_cast(kv_cache->data_ptr()) : nullptr, \ + (insert_kv && has_index) \ + ? reinterpret_cast(index_cache->data_ptr()) \ + : nullptr, \ + static_cast(eps), static_cast(rotary_dim), num_tokens, nq, \ + nkv, niq, static_cast(block_size), kv_s_block, kv_s_kv, kv_s_token, \ + kv_s_head, has_index, insert_kv, stream) + // ──────────────────────────────────────────────────────────────────────────── // Torch op wrapper // ──────────────────────────────────────────────────────────────────────────── @@ -475,9 +528,14 @@ void fused_minimax_m3_qknorm_rope_kv_insert( int64_t block_size, std::optional q_out, // [N, nq*128] contiguous std::optional - index_q_out) { // [N, niq*128] contiguous + index_q_out, // [N, niq*128] contiguous + const std::string& kv_cache_dtype) { STD_TORCH_CHECK(qkv.is_cuda() && qkv.is_contiguous(), "qkv must be contiguous CUDA"); + STD_TORCH_CHECK( + qkv.scalar_type() == torch::headeronly::ScalarType::Half || + qkv.scalar_type() == torch::headeronly::ScalarType::BFloat16, + "qkv must be float16 or bfloat16"); STD_TORCH_CHECK( positions.is_cuda() && positions.scalar_type() == torch::headeronly::ScalarType::Long, @@ -510,6 +568,8 @@ void fused_minimax_m3_qknorm_rope_kv_insert( // (1 head)]) right after [q|k|v] in the same row; the dense layer does not. bool const has_index = niq > 0; bool const insert_kv = kv_cache.has_value(); + vllm::Fp8KVCacheDataType const kv_dt = + vllm::get_fp8_kv_cache_data_type(kv_cache_dtype); int const kHeadDim = vllm::minimax_m3_fused_ops::kHeadDim; int const expected_row = (nq + 2 * nkv + (has_index ? niq + 1 : 0)) * kHeadDim; @@ -552,8 +612,14 @@ void fused_minimax_m3_qknorm_rope_kv_insert( torch::headeronly::ScalarType::Long && index_slot_mapping->numel() == slot_mapping->numel()), "index_slot_mapping must be int64 CUDA with slot_mapping length"); - STD_TORCH_CHECK(kv_cache->scalar_type() == qkv.scalar_type(), - "kv_cache dtype must match qkv (bf16 cache only)"); + if (kv_dt == vllm::Fp8KVCacheDataType::kAuto) { + STD_TORCH_CHECK(kv_cache->scalar_type() == qkv.scalar_type(), + "auto kv_cache dtype must match qkv"); + } else { + STD_TORCH_CHECK( + kv_cache->scalar_type() == torch::headeronly::ScalarType::Byte, + "fp8 kv_cache must use uint8 storage"); + } STD_TORCH_CHECK(index_cache.has_value() && index_cache->scalar_type() == qkv.scalar_type(), "insert mode requires matching index_cache"); @@ -601,35 +667,9 @@ void fused_minimax_m3_qknorm_rope_kv_insert( VLLM_STABLE_DISPATCH_HALF_TYPES( qkv.scalar_type(), "fused_minimax_m3_qknorm_rope_kv_insert", [&] { using st = scalar_t; - vllm::minimax_m3_fused_ops::launchFusedMiniMaxM3( - reinterpret_cast(qkv.data_ptr()), - q_out.has_value() ? reinterpret_cast(q_out->data_ptr()) - : nullptr, - index_q_out.has_value() - ? reinterpret_cast(index_q_out->data_ptr()) - : nullptr, - reinterpret_cast(q_norm_weight.data_ptr()), - reinterpret_cast(k_norm_weight.data_ptr()), - has_index - ? reinterpret_cast(index_q_norm_weight->data_ptr()) - : nullptr, - has_index - ? reinterpret_cast(index_k_norm_weight->data_ptr()) - : nullptr, - reinterpret_cast(cos_sin_cache.data_ptr()), - reinterpret_cast(positions.data_ptr()), - insert_kv - ? reinterpret_cast(slot_mapping->data_ptr()) - : nullptr, - insert_kv ? reinterpret_cast( - effective_index_slot_mapping->data_ptr()) - : nullptr, - insert_kv ? reinterpret_cast(kv_cache->data_ptr()) : nullptr, - (insert_kv && has_index) - ? reinterpret_cast(index_cache->data_ptr()) - : nullptr, - static_cast(eps), static_cast(rotary_dim), num_tokens, - nq, nkv, niq, static_cast(block_size), kv_s_block, kv_s_kv, - kv_s_token, kv_s_head, has_index, insert_kv, stream); + DISPATCH_BY_KV_CACHE_DTYPE(qkv.scalar_type(), kv_cache_dtype, + CALL_FUSED_MINIMAX_M3); }); } + +#undef CALL_FUSED_MINIMAX_M3 diff --git a/csrc/libtorch_stable/ops.h b/csrc/libtorch_stable/ops.h index d5144d76818..49ce1744006 100644 --- a/csrc/libtorch_stable/ops.h +++ b/csrc/libtorch_stable/ops.h @@ -3,6 +3,9 @@ #include #include +#include +#include + void per_token_group_quant_fp8(const torch::stable::Tensor& input, torch::stable::Tensor& output_q, torch::stable::Tensor& output_s, @@ -297,7 +300,8 @@ void fused_minimax_m3_qknorm_rope_kv_insert( std::optional kv_cache, std::optional index_cache, int64_t block_size, std::optional q_out, - std::optional index_q_out); + std::optional index_q_out, + const std::string& kv_cache_dtype); // Sampler kernels (shared CUDA/ROCm) void apply_repetition_penalties_( diff --git a/csrc/libtorch_stable/torch_bindings.cpp b/csrc/libtorch_stable/torch_bindings.cpp index 7d9a39a7a4b..f360c83b8af 100644 --- a/csrc/libtorch_stable/torch_bindings.cpp +++ b/csrc/libtorch_stable/torch_bindings.cpp @@ -471,7 +471,8 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) { "int num_index_heads, " "Tensor? slot_mapping, Tensor? index_slot_mapping, " "Tensor!? kv_cache, Tensor!? index_cache, " - "int block_size, Tensor!? q_out, Tensor!? index_q_out) -> ()"); + "int block_size, Tensor!? q_out, Tensor!? index_q_out, " + "str kv_cache_dtype) -> ()"); // Apply repetition penalties to logits in-place. ops.def( diff --git a/setup.py b/setup.py index a0af43e2331..2aaa7dfc49c 100644 --- a/setup.py +++ b/setup.py @@ -1168,7 +1168,7 @@ package_data = { "third_party/deep_gemm/include/**/*.h", "third_party/deep_gemm/include/**/*.hpp", # fmha_sm100 sparse CuTe-DSL helper kernels (vendored via cmake) - "third_party/fmha_sm100/cute/src/sm100/build_k2q_csr/build_k2q_csr.cu", + "third_party/fmha_sm100/cute/**/*.cu", ] } diff --git a/tests/kernels/test_fused_minimax_m3_qknorm_rope_kv_insert.py b/tests/kernels/test_fused_minimax_m3_qknorm_rope_kv_insert.py index 3268d125bb2..96729614f82 100644 --- a/tests/kernels/test_fused_minimax_m3_qknorm_rope_kv_insert.py +++ b/tests/kernels/test_fused_minimax_m3_qknorm_rope_kv_insert.py @@ -106,7 +106,16 @@ def test_dense_norm_rope(num_tokens, num_heads, num_kv_heads): qkv_orig = qkv.clone() ops.fused_minimax_m3_qknorm_rope_kv_insert( - qkv, q_w, k_w, cos_sin, positions, num_heads, num_kv_heads, ROTARY_DIM, eps + qkv, + q_w, + k_w, + cos_sin, + positions, + num_heads, + num_kv_heads, + ROTARY_DIM, + eps, + kv_cache_dtype="auto", ) q_out, k_out, v_out = qkv.split([qsz, kvsz, kvsz], dim=-1) @@ -133,7 +142,8 @@ def test_dense_norm_rope(num_tokens, num_heads, num_kv_heads): @pytest.mark.parametrize("num_tokens", [1, 7, 64, 513]) @pytest.mark.parametrize("block_size", [16, 64]) -def test_sparse_full(num_tokens, block_size): +@pytest.mark.parametrize("kv_cache_dtype", ["auto", "fp8"]) +def test_sparse_full(num_tokens, block_size, kv_cache_dtype): torch.manual_seed(1) device, dtype, eps = "cuda", torch.bfloat16, 1e-6 base, max_pos = 5_000_000.0, 4096 @@ -158,8 +168,15 @@ def test_sparse_full(num_tokens, block_size): splits = [qsz, kvsz, kvsz, iqsz, iksz] num_blocks = (num_tokens + block_size - 1) // block_size + 1 + kv_cache_storage_dtype = torch.uint8 if kv_cache_dtype == "fp8" else dtype kv_cache = torch.zeros( - num_blocks, 2, block_size, num_kv_heads, HEAD_DIM, dtype=dtype, device=device + num_blocks, + 2, + block_size, + num_kv_heads, + HEAD_DIM, + dtype=kv_cache_storage_dtype, + device=device, ) index_cache = torch.zeros( num_blocks, block_size, HEAD_DIM, dtype=dtype, device=device @@ -195,11 +212,12 @@ def test_sparse_full(num_tokens, block_size): block_size, q_out, index_q, + kv_cache_dtype, ) # ── norm+rope parity. q/index_q land in their gather buffers; k/index_k are # rewritten in place inside qkv. ── - _, k_out, _, _, index_k = qkv.split(splits, dim=-1) + _, k_out, v_out, _, index_k = qkv.split(splits, dim=-1) q_in, k_in, v_in, iq_orig, ik_orig = qkv_orig.split(splits, dim=-1) q_ref = norm_rope_ref( q_in.view(num_tokens, num_heads, HEAD_DIM), q_w, positions, cos_sin, eps @@ -230,15 +248,33 @@ def test_sparse_full(num_tokens, block_size): # ── Cache inserts. ── # Main cache layout is [num_blocks, 2, block_size, num_kv_heads, head_dim] # (the K/V axis sits *before* block_size); index cache is [nb, bs, head_dim]. - idx_flat = index_cache.view(num_blocks * block_size, HEAD_DIM) k_ref_h = k_ref.view(num_tokens, num_kv_heads, HEAD_DIM) v_ref_h = v_in.view(num_tokens, num_kv_heads, HEAD_DIM) # v is raw (no norm/rope) - for t in range(num_tokens): - s = slot_mapping[t].item() - b, pos = s // block_size, s % block_size - torch.testing.assert_close( - kv_cache[b, 0, pos], k_ref_h[t], rtol=1e-2, atol=1e-2 + if kv_cache_dtype == "fp8": + expected_kv_cache = torch.zeros_like(kv_cache) + scale = torch.ones((), device=device) + ops.reshape_and_cache_flash( + k_out.view(num_tokens, num_kv_heads, HEAD_DIM), + v_out.view(num_tokens, num_kv_heads, HEAD_DIM), + expected_kv_cache[:, 0], + expected_kv_cache[:, 1], + slot_mapping, + kv_cache_dtype, + scale, + scale, ) - torch.testing.assert_close(kv_cache[b, 1, pos], v_ref_h[t], rtol=0, atol=0) - index_s = index_slot_mapping[t].item() - torch.testing.assert_close(idx_flat[index_s], ik_ref[t], rtol=1e-2, atol=1e-2) + torch.testing.assert_close(kv_cache, expected_kv_cache, rtol=0, atol=0) + else: + for t in range(num_tokens): + s = slot_mapping[t].item() + b, pos = s // block_size, s % block_size + torch.testing.assert_close( + kv_cache[b, 0, pos], k_ref_h[t], rtol=1e-2, atol=1e-2 + ) + torch.testing.assert_close(kv_cache[b, 1, pos], v_ref_h[t], rtol=0, atol=0) + + expected_index_cache = torch.zeros_like(index_cache).view(-1, HEAD_DIM) + expected_index_cache[index_slot_mapping] = index_k + torch.testing.assert_close( + index_cache.view(-1, HEAD_DIM), expected_index_cache, rtol=0, atol=0 + ) diff --git a/vllm/_custom_ops.py b/vllm/_custom_ops.py index 6f72a8a5156..5f52bc50eb2 100644 --- a/vllm/_custom_ops.py +++ b/vllm/_custom_ops.py @@ -2634,6 +2634,7 @@ def fused_minimax_m3_qknorm_rope_kv_insert( block_size: int = 0, q_out: torch.Tensor | None = None, index_q_out: torch.Tensor | None = None, + kv_cache_dtype: str = "auto", ) -> None: """Fused MiniMax-M3 attention pre-processing (in-place). @@ -2645,8 +2646,9 @@ def fused_minimax_m3_qknorm_rope_kv_insert( index_k]`` — the index branch is read straight out of ``qkv``. When ``kv_cache`` is given (sparse serving), also scatter-inserts the - normed/roped k & v into the paged bf16 KV cache by ``slot_mapping`` and the - index key into ``index_cache`` by ``index_slot_mapping``. If + normed/roped k & v into the paged KV cache by ``slot_mapping`` and the + index key into ``index_cache`` by ``index_slot_mapping``. ``kv_cache_dtype`` + selects the cache storage/conversion path. If ``index_slot_mapping`` is omitted, ``slot_mapping`` is used for both caches. If ``q_out`` / ``index_q_out`` (contiguous ``[N, nq*128]`` / ``[N, @@ -2675,6 +2677,7 @@ def fused_minimax_m3_qknorm_rope_kv_insert( block_size, q_out, index_q_out, + kv_cache_dtype, ) diff --git a/vllm/models/minimax_m3/amd/model.py b/vllm/models/minimax_m3/amd/model.py index b80d3b8b3b8..4bd92444728 100644 --- a/vllm/models/minimax_m3/amd/model.py +++ b/vllm/models/minimax_m3/amd/model.py @@ -426,6 +426,7 @@ class MiniMaxM3Attention(nn.Module): self.num_kv_heads, self.rotary_emb.rotary_dim, self.q_norm.variance_epsilon, + kv_cache_dtype="auto", ) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) attn_output = self.attn(q, k, v) @@ -533,11 +534,6 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): self.kv_cache_torch_dtype = kv_cache_dtype_str_to_dtype( self.kv_cache_dtype, vllm_config.model_config ) - # fp8 main-K/V cache: the fused qknorm+rope+kv-insert op is bf16-cache-only - # (asserts kv_cache dtype == qkv), so on the fp8 path we run it in - # norm+rope-only mode and write the cache via the fp8-capable - # reshape_and_cache_flash in _insert_kv. (index cache stays bf16.) - self._fp8_kv = "fp8" in self.kv_cache_dtype self.attn_backend = MiniMaxM3SparseBackend # Indexer and main attention are separate impls. On ROCm the SM100 gate @@ -592,37 +588,6 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): kv_quant_mode=get_kv_quant_mode(self.kv_cache_dtype), ) - def _insert_kv( - self, - key: torch.Tensor, - value: torch.Tensor, - index_key: torch.Tensor, - main_slot_mapping: torch.Tensor, - index_slot_mapping: torch.Tensor, - ) -> None: - """Write main K/V (fp8-quantizing) and index-K into their paged caches. - - Used only on the fp8-KV path: the fused #20 op is bf16-cache-only, so it - runs in norm+rope-only mode and the (already normed/roped) k/v/index_k are - written here via ``reshape_and_cache_flash`` (which honors kv_cache_dtype, - unit scale -- matching the fp8 read path added in #33). Mirrors the - pre-#20 unfused insert. The index cache stays bf16 (no quant). - """ - key_cache, value_cache = self.kv_cache.unbind(1) - scale = torch.ones((), device=key.device) - ops.reshape_and_cache_flash( - key.view(-1, self.num_kv_heads, self.head_dim), - value.view(-1, self.num_kv_heads, self.head_dim), - key_cache, - value_cache, - main_slot_mapping, - self.kv_cache_dtype, - scale, - scale, - ) - idx_cache = self.indexer.index_cache.kv_cache.view(-1, self.idx_head_dim) - idx_cache[index_slot_mapping] = index_key.to(idx_cache.dtype) - def forward( self, positions: torch.Tensor, @@ -636,11 +601,9 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): # of the single fused ``qkv`` tensor. Once the paged caches are bound the # kernel also inserts k/v and the index key into them (each with its own # slot_mapping); the memory-profiling run (caches unbound, no slot_mapping) - # short-circuits to zeros below. Replaces the - # q_norm/k_norm/rotary_emb/index_*_norm/index_rotary_emb/_insert_kv chain. - # (#20 fused_minimax_m3_qknorm_rope_kv_insert; HIP/CDNA path. The main and - # index slot mappings are read from the forward context's slot_mapping - # dict, matching the breakable-cudagraph path -- see nvidia/model.py.) + # short-circuits to zeros below. The main and index slot mappings are read + # from the forward context's slot_mapping dict, matching the + # breakable-cudagraph path -- see nvidia/model.py. cos_sin_cache = self.rotary_emb.cos_sin_cache rotary_dim = self.rotary_emb.rotary_dim eps = self.q_norm.variance_epsilon @@ -658,12 +621,6 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): index_slot_mapping = fwd_slot_mapping[self.indexer.index_cache.prefix] q = qkv.new_empty((num_tokens, self.q_size)) index_q = qkv.new_empty((num_tokens, self.index_q_size)) - # On the fp8-KV path the fused op cannot write the (fp8) cache, so pass - # kv_cache/index_cache = None -> insert_kv=False (norm+rope only): it still - # de-interleaves q/index_q and rewrites the normed/roped k & index_k in - # place in qkv, leaving v raw (correct -- v is never normed/roped). We then - # write the cache via _insert_kv below. - insert_via_fused = not self._fp8_kv ops.fused_minimax_m3_qknorm_rope_kv_insert( qkv, self.q_norm.weight, @@ -679,26 +636,13 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): self.num_idx_heads, main_slot_mapping, index_slot_mapping, - self.kv_cache if insert_via_fused else None, - self.indexer.index_cache.kv_cache if insert_via_fused else None, + self.kv_cache, + self.indexer.index_cache.kv_cache, self.kv_cache.size(2), # paged-cache block size q, index_q, + self.kv_cache_dtype, ) - if not insert_via_fused: - # Extract the normed/roped k, raw v, normed/roped index_k from qkv - # ([q | k | v | index_q | index_k], all head_dim=128) and fp8-insert. - kv = self.num_kv_heads * self.head_dim - # These are strided views into qkv (row stride = full qkv width), but - # their last dim is contiguous, so `_insert_kv`'s `.view(-1, nkv, - # head_dim)` works on them and `reshape_and_cache_flash` honors the - # input stride -- no `.contiguous()` needed (verified bit-identical; - # avoids a [N, kv] copy per step on the fp8-KV path). - k = qkv[:, self.q_size : self.q_size + kv] - v = qkv[:, self.q_size + kv : self.q_size + 2 * kv] - ik0 = self.q_size + 2 * kv + self.index_q_size - index_k = qkv[:, ik0 : ik0 + self.num_idx_heads * self.idx_head_dim] - self._insert_kv(k, v, index_k, main_slot_mapping, index_slot_mapping) output = torch.empty_like(q) attn_output = self._run_attention(q, index_q, output) diff --git a/vllm/models/minimax_m3/common/sparse_attention.py b/vllm/models/minimax_m3/common/sparse_attention.py index ab8f24186e6..b8d60e09e4b 100644 --- a/vllm/models/minimax_m3/common/sparse_attention.py +++ b/vllm/models/minimax_m3/common/sparse_attention.py @@ -22,6 +22,7 @@ import torch from vllm.config import VllmConfig from vllm.config.cache import CacheDType from vllm.forward_context import get_forward_context +from vllm.logger import init_logger from vllm.models.minimax_m3.common.ops.sparse_attn import ( SPARSE_BLOCK_SIZE, minimax_m3_sparse_attn, @@ -44,6 +45,8 @@ from vllm.v1.attention.backends.utils import ( ) from vllm.v1.kv_cache_interface import AttentionSpec, is_quantized_kv_cache +logger = init_logger(__name__) + class MiniMaxM3SparseBackend(AttentionBackend): """Block-sparse GQA backend for MiniMax M3 sparse attention layers.""" @@ -385,16 +388,25 @@ def select_main_impl_cls( ) -> type[MiniMaxM3SparseImpl]: """Pick the main attend impl off the main KV-cache dtype. - bf16 on Blackwell (SM100) uses the MSA attend; fp8 or non-Blackwell falls + Blackwell (SM100) uses the MSA attend for supported top-k block counts + when the KV cache is BF16 or FP8 E4M3; non-Blackwell and FP8 E5M2 fall back to Triton. The MSA module is imported lazily so AMD/non-SM100 never import fmha_sm100. """ - if ( + use_msa = ( current_platform.is_cuda() and current_platform.is_device_capability_family(100) and topk_blocks in (4, 8, 16, 32) - and not is_quantized_kv_cache(kv_cache_dtype) - ): + and kv_cache_dtype != "fp8_e5m2" + ) + selected = "MSA" if use_msa else "Triton" + logger.info_once( + "MiniMax M3 sparse attention selected %s (kv_cache_dtype=%s, topk_blocks=%s)", + selected, + kv_cache_dtype, + topk_blocks, + ) + if use_msa: from vllm.models.minimax_m3.nvidia.sparse_attention_msa import ( MiniMaxM3SparseMSAImpl, ) diff --git a/vllm/models/minimax_m3/nvidia/model.py b/vllm/models/minimax_m3/nvidia/model.py index e2cd62704fd..aaced78ed7c 100644 --- a/vllm/models/minimax_m3/nvidia/model.py +++ b/vllm/models/minimax_m3/nvidia/model.py @@ -66,10 +66,7 @@ from vllm.model_executor.models.utils import ( maybe_prefix, ) from vllm.model_executor.models.vision import run_dp_sharded_mrope_vision_model -from vllm.models.minimax_m3.common.indexer import ( - MiniMaxM3Indexer, - MiniMaxM3IndexerMetadata, -) +from vllm.models.minimax_m3.common.indexer import MiniMaxM3Indexer from vllm.models.minimax_m3.common.mm_preprocess import ( MiniMaxM3VLDummyInputsBuilder, MiniMaxM3VLMultiModalProcessor, @@ -78,7 +75,6 @@ from vllm.models.minimax_m3.common.mm_preprocess import ( from vllm.models.minimax_m3.common.sparse_attention import ( MiniMaxM3SparseBackend, MiniMaxM3SparseImpl, - MiniMaxM3SparseMetadata, select_main_impl_cls, ) from vllm.models.minimax_m3.common.vision_tower import MiniMaxVLVisionModel @@ -375,6 +371,7 @@ class MiniMaxM3Attention(nn.Module): self.num_kv_heads, self.rotary_emb.rotary_dim, self.q_norm.variance_epsilon, + kv_cache_dtype="auto", ) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) attn_output = self.attn(q, k, v) @@ -545,40 +542,6 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): kv_quant_mode=get_kv_quant_mode(self.kv_cache_dtype), ) - def _insert_kv( - self, key: torch.Tensor, value: torch.Tensor, index_key: torch.Tensor - ) -> None: - """Write main K/V and index-K into their paged caches. - - No-op during the profiling run, where caches are not yet bound and - ``attn_metadata`` is None. - """ - attn_metadata = get_forward_context().attn_metadata - if not isinstance(attn_metadata, dict): - return - main_meta = attn_metadata[self.layer_name] - index_meta = attn_metadata[self.indexer.index_cache.prefix] - assert isinstance(main_meta, MiniMaxM3SparseMetadata) - assert isinstance(index_meta, MiniMaxM3IndexerMetadata) - - # Identity scale: unused for the bf16 cache, required arg of the op. - key_cache, value_cache = self.kv_cache.unbind(1) - scale = torch.ones((), device=key.device) - ops.reshape_and_cache_flash( - key.view(-1, self.num_kv_heads, self.head_dim), - value.view(-1, self.num_kv_heads, self.head_dim), - key_cache, - value_cache, - main_meta.slot_mapping, - self.kv_cache_dtype, - scale, - scale, - ) - - # Index-key cache: single vector per token, scatter by slot. - idx_cache = self.indexer.index_cache.kv_cache.view(-1, self.idx_head_dim) - idx_cache[index_meta.slot_mapping] = index_key.to(idx_cache.dtype) - def forward( self, positions: torch.Tensor, @@ -592,10 +555,9 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): # of the single fused ``qkv`` tensor (the "5 results"). Once the paged # caches are bound the kernel also inserts k/v and the index key into # them; the initial memory-profiling run (caches unbound, no slot_mapping) - # short-circuits to zeros below. Replaces the - # q_norm/k_norm/rotary_emb/index_*_norm/index_rotary_emb/_insert_kv - # sequence. k/v and index_k are rewritten in place inside qkv (and - # scatter-inserted into the caches); q and index_q are de-interleaved + # short-circuits to zeros below. k/v and index_k are rewritten in place + # inside qkv (and scatter-inserted into the caches); q and index_q are + # de-interleaved # straight into the dedicated contiguous ``q``/``index_q`` buffers below. cos_sin_cache = self.rotary_emb.cos_sin_cache @@ -635,6 +597,7 @@ class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase): self.kv_cache.size(2), # paged-cache block size q, index_q, + self.kv_cache_dtype, ) output = torch.empty_like(q) From a46abb7ae68acc13a4fc5870db98619b3f97c6e0 Mon Sep 17 00:00:00 2001 From: Ting SUN Date: Wed, 17 Jun 2026 13:08:41 +0800 Subject: [PATCH 03/37] [Bugfix][Quantization] Reject unsupported compressed tensors KV cache schemes (#45312) Signed-off-by: Ting Sun Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> --- .../quantization/compressed_tensors/compressed_tensors.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors.py b/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors.py index 71fc809f4df..2231b2ca9af 100644 --- a/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors.py +++ b/vllm/model_executor/layers/quantization/compressed_tensors/compressed_tensors.py @@ -983,7 +983,7 @@ class CompressedTensorsKVCacheMethod(BaseKVCacheMethod): type_ = kv_cache_scheme.get("type") num_bits = kv_cache_scheme.get("num_bits") - if type_ != "float" and num_bits != 8: + if type_ != "float" or num_bits != 8: raise NotImplementedError( "Currently supported kv cache quantization is " "num_bits=8, type=float, however " From e9993a52aa1ad53598b43ca629a7a600cc14d001 Mon Sep 17 00:00:00 2001 From: Nick Hill Date: Tue, 16 Jun 2026 23:30:49 -0700 Subject: [PATCH 04/37] [BugFix][CI] Fix scheduler plugin test (#45897) Signed-off-by: Nick Hill --- tests/plugins_tests/test_scheduler_plugins.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/plugins_tests/test_scheduler_plugins.py b/tests/plugins_tests/test_scheduler_plugins.py index 45902cc874c..f416b888f51 100644 --- a/tests/plugins_tests/test_scheduler_plugins.py +++ b/tests/plugins_tests/test_scheduler_plugins.py @@ -10,7 +10,7 @@ from vllm.v1.engine.llm_engine import LLMEngine class DummyV1Scheduler(Scheduler): - def schedule(self): + def schedule(self, throttle_prefills: bool = False): raise Exception("Exception raised by DummyV1Scheduler") From 56e434522671890bdbe7362bcec46956521f951c Mon Sep 17 00:00:00 2001 From: Reid <61492567+reidliu41@users.noreply.github.com> Date: Wed, 17 Jun 2026 14:38:06 +0800 Subject: [PATCH 05/37] [Rust Frontend] Support prompt-only completions (#44938) Signed-off-by: reidliu41 --- .../server/src/routes/openai/completions.rs | 471 +++++++++++++++++- .../src/routes/openai/completions/convert.rs | 73 ++- .../src/routes/openai/completions/validate.rs | 33 +- 3 files changed, 550 insertions(+), 27 deletions(-) diff --git a/rust/src/server/src/routes/openai/completions.rs b/rust/src/server/src/routes/openai/completions.rs index de21dc3a1c3..95fb4db9a6f 100644 --- a/rust/src/server/src/routes/openai/completions.rs +++ b/rust/src/server/src/routes/openai/completions.rs @@ -2,6 +2,7 @@ mod convert; mod types; mod validate; +use std::collections::HashMap; use std::convert::Infallible; use std::result::Result; use std::sync::Arc; @@ -16,12 +17,15 @@ use futures::{Stream, StreamExt as _, pin_mut}; use thiserror_ext::AsReport as _; use tracing::{debug, error, info, trace}; use tracing_futures::Instrument as _; -use vllm_text::{DecodedTextEvent, FinishReason, TextOutputStream, TextOutputStreamExt as _}; +use vllm_text::{ + DecodedPromptLogprobs, DecodedTextEvent, FinishReason, TextOutputStream, + TextOutputStreamExt as _, +}; use self::convert::{ResponseOptions, prepare_completion_request}; use super::utils::logprobs::{ collected_logprobs_to_openai, decoded_logprobs_to_openai, decoded_prompt_logprobs_to_maps, - text_len, + decoded_prompt_logprobs_to_openai, text_len, }; use super::utils::types::Usage; use crate::config::ApiServerOptions; @@ -119,6 +123,7 @@ async fn collect_completion( include_usage: _, // Ignored: non-streaming responses are collected before usage is attached. include_continuous_usage: _, + prompt_only, echo, requested_logprobs, include_prompt_logprobs, @@ -136,17 +141,17 @@ async fn collect_completion( .map(|sr| serde_json::to_value(sr).expect("StopReason must serialize to JSON")); let prompt_char_count = echo.as_ref().map(|prompt| text_len(prompt)).unwrap_or_default(); - let prompt_logprobs = if include_prompt_logprobs { - let prompt_logprobs = collected.prompt_logprobs.as_ref().ok_or_else(|| { - server_error!( - "completion response requested prompt_logprobs but generation returned none" - ) + let logprobs = if requested_logprobs.is_some() && prompt_only { + let prompt = echo.as_deref().ok_or_else(|| { + server_error!("prompt-only completion response missing echoed prompt") })?; - Some(prompt_logprobs) - } else { - None - }; - let logprobs = if requested_logprobs.is_some() { + Some(prompt_only_logprobs_to_openai( + collected.prompt_logprobs.as_ref(), + prompt, + collected.prompt_token_ids.as_ref(), + return_tokens_as_token_ids, + )?) + } else if requested_logprobs.is_some() { Some(collected_logprobs_to_openai( &collected, echo.is_some(), @@ -156,10 +161,18 @@ async fn collect_completion( } else { None }; - let prompt_logprobs = - prompt_logprobs.map(|lp| decoded_prompt_logprobs_to_maps(lp, return_tokens_as_token_ids)); + let prompt_logprobs = if include_prompt_logprobs { + Some(prompt_logprobs_to_maps( + collected.prompt_logprobs.as_ref(), + collected.prompt_token_ids.as_ref(), + return_tokens_as_token_ids, + )?) + } else { + None + }; let text = match &echo { None => collected.text, + Some(prompt) if prompt_only => prompt.clone(), Some(prompt) => format!("{prompt}{}", collected.text), }; let finish_reason = completion_finish_reason_to_openai(finish_reason)?.to_string(); @@ -211,6 +224,7 @@ async fn completion_chunk_stream( ResponseOptions { include_usage, include_continuous_usage, + prompt_only, echo, requested_logprobs, // Ignored: streaming prompt logprobs are rejected for Python parity. @@ -239,14 +253,30 @@ async fn completion_chunk_stream( while let Some(next) = stream.next().await { match next { Ok(DecodedTextEvent::Start { - prompt_token_ids, .. + prompt_token_ids, + prompt_logprobs, }) => { debug!("completion stream started"); continuous_usage.set_prompt_tokens(prompt_token_ids.len()); if let Some(prompt) = echo.as_ref() { visible_text_len = text_len(prompt); - let mut chunk = - delta_chunk(&request_id, &response_model, created, prompt.clone(), None); + let logprobs = if prompt_only && requested_logprobs.is_some() { + Some(prompt_only_logprobs_to_openai( + prompt_logprobs.as_ref(), + prompt, + prompt_token_ids.as_ref(), + return_tokens_as_token_ids, + )?) + } else { + None + }; + let mut chunk = delta_chunk( + &request_id, + &response_model, + created, + prompt.clone(), + logprobs, + ); if return_token_ids && first_chunk { if let Some(choice) = chunk.choices.first_mut() { choice.prompt_token_ids = Some(prompt_token_ids.to_vec()); @@ -271,6 +301,48 @@ async fn completion_chunk_stream( logprobs, finished, }) => { + // Prompt-only streaming already emitted the echoed prompt in the Start chunk. + // The one generated token is only used to drive the engine to a finished event, + // so hide its delta and forward only the terminal finish/usage metadata. + if prompt_only { + if let Some(finished) = finished { + if enable_log_requests { + info!( + stream = true, + model = %response_model, + prompt_tokens = finished.usage.prompt_token_count, + output_tokens = finished.usage.output_token_count, + finish_reason = finished.finish_reason.as_str(), + "completion finished" + ); + } + continuous_usage.set_final_counts( + finished.usage.prompt_token_count, + finished.usage.output_token_count, + ); + let final_chunk = final_chunk( + &request_id, + &response_model, + created, + finished.finish_reason, + )?; + yield_chunk!(final_chunk); + + if include_usage { + y.yield_ok(CompletionSseChunk::Usage(usage_chunk( + &request_id, + &response_model, + created, + Usage::from_token_usage( + finished.usage, + enable_prompt_tokens_details, + ), + ))) + .await; + } + } + continue; + } let delta_text_len = text_len(&delta); let logprobs = if requested_logprobs.is_some() { let decoded_logprobs = logprobs.as_ref().ok_or_else(|| { @@ -386,6 +458,57 @@ fn completion_finish_reason_to_openai( } } +fn prompt_only_logprobs_to_openai( + prompt_logprobs: Option<&DecodedPromptLogprobs>, + prompt: &str, + prompt_token_ids: &[u32], + return_tokens_as_token_ids: bool, +) -> Result { + if let Some(prompt_logprobs) = prompt_logprobs { + return decoded_prompt_logprobs_to_openai(prompt_logprobs, 0, return_tokens_as_token_ids); + } + + if let [token_id] = prompt_token_ids { + let token = if return_tokens_as_token_ids { + format!("token_id:{token_id}") + } else { + prompt.to_string() + }; + + return Ok(LogProbs { + tokens: vec![token], + token_logprobs: vec![None], + top_logprobs: vec![None], + text_offset: vec![0], + }); + } + + Err(server_error!( + "prompt-only completion requested logprobs but generation returned none" + )) +} + +fn prompt_logprobs_to_maps( + prompt_logprobs: Option<&DecodedPromptLogprobs>, + prompt_token_ids: &[u32], + return_tokens_as_token_ids: bool, +) -> Result>>, ApiError> { + if let Some(prompt_logprobs) = prompt_logprobs { + return Ok(decoded_prompt_logprobs_to_maps( + prompt_logprobs, + return_tokens_as_token_ids, + )); + } + + if let [_token_id] = prompt_token_ids { + return Ok(vec![None]); + } + + Err(server_error!( + "completion response requested prompt_logprobs but generation returned none" + )) +} + fn usage_chunk( request_id: &str, response_model: &str, @@ -449,8 +572,8 @@ mod tests { use futures::{StreamExt as _, stream}; use itertools::Itertools as _; use vllm_text::{ - DecodedLogprobs, DecodedPositionLogprobs, DecodedTextEvent, DecodedTokenLogprob, - FinishReason, Finished, + DecodedLogprobs, DecodedPositionLogprobs, DecodedPromptLogprobs, DecodedTextEvent, + DecodedTokenLogprob, FinishReason, Finished, }; use super::{ @@ -613,4 +736,314 @@ mod tests { CompletionSseChunk::Chunk(_) => panic!("expected usage chunk"), } } + + #[tokio::test] + async fn collect_completion_hides_internal_prompt_only_token() { + let stream = stream::iter(vec![ + Ok(DecodedTextEvent::Start { + prompt_token_ids: vec![1, 2].into(), + prompt_logprobs: None, + }), + Ok(DecodedTextEvent::TextDelta { + delta: " leaked".to_string(), + token_ids: vec![3], + logprobs: None, + finished: Some(Finished { + usage: vllm_llm::TokenUsage { + prompt_token_count: 2, + output_token_count: 1, + cached_token_count: 0, + }, + finish_reason: FinishReason::Length, + kv_transfer_params: None, + }), + }), + ]); + + let response = super::collect_completion( + stream, + "cmpl-1".to_string(), + "model".to_string(), + 1, + ApiServerOptions::default(), + ResponseOptions { + prompt_only: true, + echo: Some("hello".to_string()), + return_token_ids: true, + ..Default::default() + }, + ) + .await + .expect("collect completion"); + + assert_eq!(response.choices[0].text, "hello"); + assert_eq!(response.choices[0].token_ids.as_deref(), Some(&[3][..])); + assert_eq!( + response.choices[0].prompt_token_ids.as_deref(), + Some(&[1, 2][..]) + ); + let usage = response.usage.expect("usage"); + assert_eq!(usage.prompt_tokens, 2); + assert_eq!(usage.completion_tokens, Some(1)); + assert_eq!(usage.total_tokens, 3); + } + + #[tokio::test] + async fn collect_completion_maps_prompt_logprobs_for_single_token_prompt() { + let stream = stream::iter(vec![ + Ok(DecodedTextEvent::Start { + prompt_token_ids: vec![9707].into(), + prompt_logprobs: None, + }), + Ok(DecodedTextEvent::TextDelta { + delta: " leaked".to_string(), + token_ids: vec![3], + logprobs: None, + finished: Some(Finished { + usage: vllm_llm::TokenUsage { + prompt_token_count: 1, + output_token_count: 1, + cached_token_count: 0, + }, + finish_reason: FinishReason::Length, + kv_transfer_params: None, + }), + }), + ]); + + let response = super::collect_completion( + stream, + "cmpl-1".to_string(), + "model".to_string(), + 1, + ApiServerOptions::default(), + ResponseOptions { + prompt_only: true, + echo: Some("Hello".to_string()), + requested_logprobs: Some(1), + include_prompt_logprobs: true, + ..Default::default() + }, + ) + .await + .expect("collect completion"); + + let choice = &response.choices[0]; + assert_eq!(choice.text, "Hello"); + assert_eq!(choice.prompt_logprobs, Some(vec![None])); + let logprobs = choice.logprobs.as_ref().expect("logprobs"); + assert_eq!(logprobs.tokens, vec!["Hello".to_string()]); + assert_eq!(logprobs.token_logprobs, vec![None]); + assert_eq!(logprobs.top_logprobs, vec![None]); + assert_eq!(logprobs.text_offset, vec![0]); + let usage = response.usage.expect("usage"); + assert_eq!(usage.prompt_tokens, 1); + assert_eq!(usage.completion_tokens, Some(1)); + assert_eq!(usage.total_tokens, 2); + } + + #[tokio::test] + async fn completion_chunk_stream_hides_internal_prompt_only_token() { + let stream = stream::iter(vec![ + Ok(DecodedTextEvent::Start { + prompt_token_ids: vec![1, 2].into(), + prompt_logprobs: None, + }), + Ok(DecodedTextEvent::TextDelta { + delta: " leaked".to_string(), + token_ids: vec![3], + logprobs: None, + finished: Some(Finished { + usage: vllm_llm::TokenUsage { + prompt_token_count: 2, + output_token_count: 1, + cached_token_count: 0, + }, + finish_reason: FinishReason::Length, + kv_transfer_params: None, + }), + }), + ]); + + let chunks = completion_chunk_stream( + stream, + "cmpl-1".to_string(), + "model".to_string(), + 1, + ApiServerOptions::default(), + ResponseOptions { + include_usage: true, + prompt_only: true, + echo: Some("hello".to_string()), + return_token_ids: true, + ..Default::default() + }, + ) + .collect::>() + .await; + + let chunks: Vec<_> = chunks.into_iter().try_collect().expect("stream should succeed"); + assert_eq!(chunks.len(), 3); + + match &chunks[0] { + CompletionSseChunk::Chunk(chunk) => { + assert_eq!(chunk.choices[0].text, "hello"); + assert_eq!( + chunk.choices[0].prompt_token_ids.as_deref(), + Some(&[1, 2][..]) + ); + } + CompletionSseChunk::Usage(_) => panic!("expected prompt chunk"), + } + match &chunks[1] { + CompletionSseChunk::Chunk(chunk) => { + assert_eq!(chunk.choices[0].text, ""); + assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("length")); + } + CompletionSseChunk::Usage(_) => panic!("expected final chunk"), + } + match &chunks[2] { + CompletionSseChunk::Usage(chunk) => { + let usage = chunk.usage.as_ref().expect("usage"); + assert_eq!(usage.prompt_tokens, 2); + assert_eq!(usage.completion_tokens, Some(1)); + assert_eq!(usage.total_tokens, 3); + } + CompletionSseChunk::Chunk(_) => panic!("expected usage chunk"), + } + } + + #[tokio::test] + async fn completion_chunk_stream_maps_prompt_logprobs_for_single_token_prompt() { + let stream = stream::iter(vec![ + Ok(DecodedTextEvent::Start { + prompt_token_ids: vec![9707].into(), + prompt_logprobs: None, + }), + Ok(DecodedTextEvent::TextDelta { + delta: " leaked".to_string(), + token_ids: vec![3], + logprobs: None, + finished: Some(Finished { + usage: vllm_llm::TokenUsage { + prompt_token_count: 1, + output_token_count: 1, + cached_token_count: 0, + }, + finish_reason: FinishReason::Length, + kv_transfer_params: None, + }), + }), + ]); + + let chunks = completion_chunk_stream( + stream, + "cmpl-1".to_string(), + "model".to_string(), + 1, + ApiServerOptions::default(), + ResponseOptions { + prompt_only: true, + echo: Some("Hello".to_string()), + requested_logprobs: Some(1), + ..Default::default() + }, + ) + .collect::>() + .await; + + let chunks: Vec<_> = chunks.into_iter().try_collect().expect("stream should succeed"); + assert_eq!(chunks.len(), 2); + + match &chunks[0] { + CompletionSseChunk::Chunk(chunk) => { + assert_eq!(chunk.choices[0].text, "Hello"); + let logprobs = chunk.choices[0].logprobs.as_ref().expect("logprobs"); + assert_eq!(logprobs.tokens, vec!["Hello".to_string()]); + assert_eq!(logprobs.token_logprobs, vec![None]); + assert_eq!(logprobs.top_logprobs, vec![None]); + assert_eq!(logprobs.text_offset, vec![0]); + } + CompletionSseChunk::Usage(_) => panic!("expected prompt chunk"), + } + match &chunks[1] { + CompletionSseChunk::Chunk(chunk) => { + assert_eq!(chunk.choices[0].text, ""); + assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("length")); + } + CompletionSseChunk::Usage(_) => panic!("expected final chunk"), + } + } + + #[tokio::test] + async fn completion_chunk_stream_maps_prompt_only_logprobs() { + let stream = stream::iter(vec![ + Ok(DecodedTextEvent::Start { + prompt_token_ids: vec![1, 2].into(), + prompt_logprobs: Some(DecodedPromptLogprobs { + first_token_id: 1, + first_token: "he".to_string(), + scored_positions: vec![DecodedPositionLogprobs { + entries: vec![DecodedTokenLogprob { + token_id: 2, + token: "llo".to_string(), + logprob: -0.2, + rank: 1, + }], + }], + }), + }), + Ok(DecodedTextEvent::TextDelta { + delta: " leaked".to_string(), + token_ids: vec![3], + logprobs: None, + finished: Some(Finished { + usage: vllm_llm::TokenUsage { + prompt_token_count: 2, + output_token_count: 1, + cached_token_count: 0, + }, + finish_reason: FinishReason::Length, + kv_transfer_params: None, + }), + }), + ]); + + let chunks = completion_chunk_stream( + stream, + "cmpl-1".to_string(), + "model".to_string(), + 1, + ApiServerOptions::default(), + ResponseOptions { + prompt_only: true, + echo: Some("hello".to_string()), + requested_logprobs: Some(1), + ..Default::default() + }, + ) + .collect::>() + .await; + + let chunks: Vec<_> = chunks.into_iter().try_collect().expect("stream should succeed"); + assert_eq!(chunks.len(), 2); + + match &chunks[0] { + CompletionSseChunk::Chunk(chunk) => { + assert_eq!(chunk.choices[0].text, "hello"); + let logprobs = chunk.choices[0].logprobs.as_ref().expect("logprobs"); + assert_eq!(logprobs.tokens, vec!["he".to_string(), "llo".to_string()]); + assert_eq!(logprobs.token_logprobs, vec![None, Some(-0.2)]); + assert_eq!(logprobs.text_offset, vec![0, 2]); + } + CompletionSseChunk::Usage(_) => panic!("expected prompt chunk"), + } + match &chunks[1] { + CompletionSseChunk::Chunk(chunk) => { + assert_eq!(chunk.choices[0].text, ""); + assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("length")); + } + CompletionSseChunk::Usage(_) => panic!("expected final chunk"), + } + } } diff --git a/rust/src/server/src/routes/openai/completions/convert.rs b/rust/src/server/src/routes/openai/completions/convert.rs index 2f6c760a990..9c306928590 100644 --- a/rust/src/server/src/routes/openai/completions/convert.rs +++ b/rust/src/server/src/routes/openai/completions/convert.rs @@ -27,6 +27,8 @@ pub(super) struct ResponseOptions { pub include_usage: bool, /// Whether every streamed chunk should carry cumulative usage. pub include_continuous_usage: bool, + /// Whether the caller requested prompt-only echo via `max_tokens=0`. + pub prompt_only: bool, /// Original text prompt that should be echoed back northbound when /// `echo=true`. pub echo: Option, @@ -68,11 +70,13 @@ pub(super) fn prepare_completion_request( })?), None => None, }; - let prompt_logprobs = request.prompt_logprobs.or(if request.echo && !request.stream { - logprobs - } else { - None - }); + let prompt_only = request.echo && request.max_tokens == Some(0); + let prompt_logprobs = + request.prompt_logprobs.or(if request.echo && (!request.stream || prompt_only) { + logprobs + } else { + None + }); let include_usage = (request.stream_options.as_ref()) .and_then(|options| options.include_usage) .unwrap_or(false); @@ -83,6 +87,11 @@ pub(super) fn prepare_completion_request( .and_then(|options| options.continuous_usage_stats) .unwrap_or(false); let include_prompt_logprobs = prompt_logprobs.is_some(); + let max_tokens = if prompt_only { + Some(1) + } else { + request.max_tokens + }; let echo = request.echo.then(|| request.prompt.as_text().cloned()).flatten(); let structured_outputs = @@ -97,7 +106,7 @@ pub(super) fn prepare_completion_request( top_p: request.top_p, top_k: request.top_k, seed: request.seed, - max_tokens: request.max_tokens, + max_tokens, min_tokens: request.min_tokens, logprobs, prompt_logprobs, @@ -138,6 +147,7 @@ pub(super) fn prepare_completion_request( options: ResponseOptions { include_usage, include_continuous_usage, + prompt_only, echo, requested_logprobs: request.logprobs, include_prompt_logprobs, @@ -325,6 +335,57 @@ mod tests { assert_eq!(prepared.options.echo, Some("hello".to_string())); assert_eq!(prepared.text_request.sampling_params.max_tokens, Some(7)); + assert!(!prepared.options.prompt_only); + } + + #[test] + fn prepare_completion_request_lowers_prompt_only_echo_as_one_internal_token() { + let request: CompletionRequest = serde_json::from_value(json!({ + "model": "Qwen/Qwen1.5-0.5B-Chat", + "prompt": "hello", + "stream": false, + "echo": true, + "max_tokens": 0 + })) + .expect("parse request"); + + let prepared = prepare_completion_request( + request, + &served(&["Qwen/Qwen1.5-0.5B-Chat"]), + ResolvedRequestContext::default(), + ) + .expect("prepare"); + + assert!(prepared.options.prompt_only); + assert_eq!(prepared.options.echo, Some("hello".to_string())); + assert_eq!(prepared.text_request.sampling_params.max_tokens, Some(1)); + } + + #[test] + fn prepare_completion_request_enables_prompt_logprobs_for_stream_prompt_only_echo() { + let request: CompletionRequest = serde_json::from_value(json!({ + "model": "Qwen/Qwen1.5-0.5B-Chat", + "prompt": "hello", + "echo": true, + "stream": true, + "max_tokens": 0, + "logprobs": 3 + })) + .expect("parse request"); + + let prepared = prepare_completion_request( + request, + &served(&["Qwen/Qwen1.5-0.5B-Chat"]), + ResolvedRequestContext::default(), + ) + .expect("prepare"); + + assert!(prepared.options.prompt_only); + assert_eq!(prepared.text_request.sampling_params.logprobs, Some(3)); + assert_eq!( + prepared.text_request.sampling_params.prompt_logprobs, + Some(3) + ); } #[test] diff --git a/rust/src/server/src/routes/openai/completions/validate.rs b/rust/src/server/src/routes/openai/completions/validate.rs index f19defe5e49..cbb040b90d0 100644 --- a/rust/src/server/src/routes/openai/completions/validate.rs +++ b/rust/src/server/src/routes/openai/completions/validate.rs @@ -26,8 +26,11 @@ pub(super) fn validate_request_compat( bail_invalid_request!(param = "n", "Only n=1 is supported."); } - if request.max_tokens == Some(0) { - bail_invalid_request!(param = "max_tokens", "max_tokens must be greater than 0."); + if request.max_tokens == Some(0) && !request.echo { + bail_invalid_request!( + param = "max_tokens", + "max_tokens=0 is only supported when echo=true." + ); } if request.echo && matches!(request.prompt, Prompt::TokenIds(_)) { @@ -166,4 +169,30 @@ mod tests { validate_request_compat(&request, &served_names(&["Qwen/Qwen1.5-0.5B-Chat"])).is_ok() ); } + + #[test] + fn validate_request_compat_accepts_prompt_only_echo() { + let request = CompletionRequest { + stream: false, + echo: true, + max_tokens: Some(0), + ..base_request() + }; + assert!( + validate_request_compat(&request, &served_names(&["Qwen/Qwen1.5-0.5B-Chat"])).is_ok() + ); + } + + #[test] + fn validate_request_compat_rejects_prompt_only_without_echo() { + let request = CompletionRequest { + stream: false, + echo: false, + max_tokens: Some(0), + ..base_request() + }; + assert!( + validate_request_compat(&request, &served_names(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err() + ); + } } From 295232a26a42837326531b48e891d54bab210dcb Mon Sep 17 00:00:00 2001 From: Sahil Singh Date: Wed, 17 Jun 2026 12:10:47 +0530 Subject: [PATCH 06/37] [Rust Frontend] Add /abort_requests endpoint (#44382) Signed-off-by: Sahil Singh --- rust/src/server/src/routes.rs | 2 + rust/src/server/src/routes/abort_requests.rs | 37 +++++++ rust/src/server/src/routes/tests.rs | 106 +++++++++++++++++++ 3 files changed, 145 insertions(+) create mode 100644 rust/src/server/src/routes/abort_requests.rs diff --git a/rust/src/server/src/routes.rs b/rust/src/server/src/routes.rs index 1cb83425ba1..3826ad40db7 100644 --- a/rust/src/server/src/routes.rs +++ b/rust/src/server/src/routes.rs @@ -1,3 +1,4 @@ +mod abort_requests; mod cache; mod collective_rpc; mod health; @@ -91,6 +92,7 @@ fn build_router_with_options( .route("/reset_mm_cache", post(cache::reset_mm_cache)) .route("/reset_encoder_cache", post(cache::reset_encoder_cache)) .route("/collective_rpc", post(collective_rpc::collective_rpc)) + .route("/abort_requests", post(abort_requests::abort_requests)) .route("/sleep", post(sleep::sleep)) .route("/wake_up", post(sleep::wake_up)) .route("/is_sleeping", get(sleep::is_sleeping)) diff --git a/rust/src/server/src/routes/abort_requests.rs b/rust/src/server/src/routes/abort_requests.rs new file mode 100644 index 00000000000..34fb041c800 --- /dev/null +++ b/rust/src/server/src/routes/abort_requests.rs @@ -0,0 +1,37 @@ +use std::sync::Arc; + +use axum::Json; +use axum::extract::State; +use axum::extract::rejection::JsonRejection; +use axum::http::StatusCode; +use serde::Deserialize; + +use crate::error::ApiError; +use crate::state::AppState; +use crate::utils::utility_call_error; + +#[derive(Debug, Deserialize)] +pub(crate) struct AbortRequestsRequest { + request_ids: Option>, +} + +pub async fn abort_requests( + State(state): State>, + body: Result, JsonRejection>, +) -> Result { + let Json(body) = body.map_err(|error| ApiError::json_parse_error(error.body_text()))?; + let request_ids = body.request_ids.ok_or_else(|| { + ApiError::invalid_request( + "Missing 'request_ids' in request body".to_string(), + Some("request_ids"), + ) + })?; + + state + .chat + .abort(&request_ids) + .await + .map_err(|error| utility_call_error("abort_requests", error))?; + + Ok(StatusCode::OK) +} diff --git a/rust/src/server/src/routes/tests.rs b/rust/src/server/src/routes/tests.rs index 9fb184162ce..5eb65a49853 100644 --- a/rust/src/server/src/routes/tests.rs +++ b/rust/src/server/src/routes/tests.rs @@ -4921,6 +4921,111 @@ async fn is_paused_route_returns_json_payload() { ); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn abort_requests_route_returns_ok_for_well_formed_body() { + let (app, engine_task) = + test_admin_app_with_engine_script(|_dealer, _push| boxed_test_future(async move {})).await; + + let response = app + .clone() + .call( + Request::builder() + .method("POST") + .uri("/abort_requests") + .header("content-type", "application/json") + .body(Body::from(r#"{"request_ids":["req-1","req-2"]}"#)) + .expect("build request"), + ) + .await + .expect("call app"); + + let status = response.status(); + let body = to_bytes(response.into_body(), usize::MAX).await.expect("read body"); + assert_eq!(status, StatusCode::OK, "{}", String::from_utf8_lossy(&body)); + assert!(body.is_empty()); + engine_task.abort_and_join().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn abort_requests_route_rejects_missing_request_ids() { + let (app, engine_task) = + test_admin_app_with_engine_script(|_dealer, _push| boxed_test_future(async move {})).await; + + let response = app + .clone() + .call( + Request::builder() + .method("POST") + .uri("/abort_requests") + .header("content-type", "application/json") + .body(Body::from(r#"{}"#)) + .expect("build request"), + ) + .await + .expect("call app"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let body = to_bytes(response.into_body(), usize::MAX).await.expect("read body"); + let json: serde_json::Value = serde_json::from_slice(&body).expect("decode json"); + assert_eq!(json["error"]["type"], "invalid_request_error"); + assert_eq!(json["error"]["param"], "request_ids"); + engine_task.abort_and_join().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn abort_requests_route_rejects_malformed_json() { + let (app, engine_task) = + test_admin_app_with_engine_script(|_dealer, _push| boxed_test_future(async move {})).await; + + let response = app + .clone() + .call( + Request::builder() + .method("POST") + .uri("/abort_requests") + .header("content-type", "application/json") + .body(Body::from(r#"{"request_ids": "#)) + .expect("build request"), + ) + .await + .expect("call app"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let body = to_bytes(response.into_body(), usize::MAX).await.expect("read body"); + let json: serde_json::Value = serde_json::from_slice(&body).expect("decode json"); + assert_eq!(json["error"]["type"], "invalid_request_error"); + engine_task.abort_and_join().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn abort_requests_route_accepts_empty_id_list() { + let (app, engine_task) = + test_admin_app_with_engine_script(|_dealer, _push| boxed_test_future(async move {})).await; + + let response = app + .clone() + .call( + Request::builder() + .method("POST") + .uri("/abort_requests") + .header("content-type", "application/json") + .body(Body::from(r#"{"request_ids":[]}"#)) + .expect("build request"), + ) + .await + .expect("call app"); + + let status = response.status(); + let body = to_bytes(response.into_body(), usize::MAX).await.expect("read body"); + assert_eq!(status, StatusCode::OK, "{}", String::from_utf8_lossy(&body)); + assert!(body.is_empty()); + engine_task.abort_and_join().await; +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn admin_routes_are_hidden_when_dev_mode_is_disabled() { @@ -4941,6 +5046,7 @@ async fn admin_routes_are_hidden_when_dev_mode_is_disabled() { ("POST", "/pause"), ("POST", "/resume"), ("POST", "/collective_rpc"), + ("POST", "/abort_requests"), ("POST", "/reset_prefix_cache"), ("POST", "/reset_mm_cache"), ("POST", "/reset_encoder_cache"), From 17bc1445562435b608041d434e9738440954159c Mon Sep 17 00:00:00 2001 From: Will Eaton Date: Wed, 17 Jun 2026 02:40:49 -0400 Subject: [PATCH 07/37] [Rust Frontend] Add serde defaults for omit_defaults fields in `EngineCoreSamplingParams` (#45848) Signed-off-by: Will Eaton --- .../engine-core-client/src/protocol/mod.rs | 81 ++++++++++++++++--- .../engine-core-client/src/tests/client.rs | 37 +++++++++ .../src/tests/python_compat.py | 16 +++- 3 files changed, 121 insertions(+), 13 deletions(-) diff --git a/rust/src/engine-core-client/src/protocol/mod.rs b/rust/src/engine-core-client/src/protocol/mod.rs index 6fae982d2ea..5e340b91176 100644 --- a/rust/src/engine-core-client/src/protocol/mod.rs +++ b/rust/src/engine-core-client/src/protocol/mod.rs @@ -44,6 +44,14 @@ fn default_repetition_penalty() -> f32 { 1.0 } +fn default_temperature() -> f32 { + 1.0 +} + +fn default_max_tokens() -> u32 { + 16 +} + mod classified_outputs; pub mod dtype; pub mod handshake; @@ -246,24 +254,28 @@ pub struct StructuredOutputsParams { /// /// Original Python definition: /// +// Python's SamplingParams is `omit_defaults=True`, so msgpack drops +// default-valued keys; default the whole struct. Per-field fns cover the +// non-zero defaults. #[serde_with::skip_serializing_none] -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, DefaultFromSerde)] +#[serde(default)] pub struct EngineCoreSamplingParams { /// Controls randomness. Lower values are more deterministic; zero means /// greedy sampling. + #[serde(default = "default_temperature")] pub temperature: f32, /// Cumulative probability threshold for nucleus sampling. #[serde(default = "default_top_p")] pub top_p: f32, /// Maximum number of top tokens to consider. `0` means all tokens. - #[serde(default)] pub top_k: u32, /// Random seed used by the sampler when present. pub seed: Option, /// Maximum number of tokens to generate per output sequence. + #[serde(default = "default_max_tokens")] pub max_tokens: u32, /// Minimum number of tokens to generate before EOS or stop-token handling. - #[serde(default)] pub min_tokens: u32, /// Number of log probabilities to return per generated token. /// @@ -274,7 +286,6 @@ pub struct EngineCoreSamplingParams { /// `None` disables prompt logprobs. `-1` requests the full vocabulary. pub prompt_logprobs: Option, /// Minimum probability threshold for token sampling. - #[serde(default)] pub min_p: f32, /// Frequency penalty applied by the sampler. pub frequency_penalty: f32, @@ -301,16 +312,13 @@ pub struct EngineCoreSamplingParams { pub all_stop_token_ids: BTreeSet, /// Logit biases to apply during sampling. /// Keys are token IDs - #[serde(default)] pub logit_bias: Option>, /// Restrict output to these token IDs only. - #[serde(default)] pub allowed_token_ids: Option>, /// Tokenized bad words to avoid during generation. - #[serde(default, rename = "_bad_words_token_ids")] + #[serde(rename = "_bad_words_token_ids")] pub bad_words_token_ids: Option>>, /// Parameters for configuring structured outputs (guided decoding). - #[serde(default)] pub structured_outputs: Option, /// Specific token IDs for which log probabilities should be returned at /// each position. @@ -318,15 +326,12 @@ pub struct EngineCoreSamplingParams { /// When set, the engine returns logprobs for exactly these tokens in /// addition to the sampled/scored token. Mutually exclusive with the /// `logprobs` count field in practice. - #[serde(default)] pub logprob_token_ids: Option>, /// If `Some(true)`, the request will not attempt to read from the prefix /// cache; newly computed blocks may still populate the cache. `None` /// defers to engine-core defaults. - #[serde(default)] pub skip_reading_prefix_cache: Option, /// Additional request parameters for custom extensions (from `vllm_xargs`). - #[serde(default)] pub extra_args: Option>, } @@ -640,4 +645,58 @@ mod tests { let value = serde_json::to_value(params).unwrap(); assert_eq!(value["_backend"], "guidance"); } + + /// A real `sampling_params` is a sparse `omit_defaults` map; absent fields + /// must fall back to defaults. `python_compat` can't catch this since Rust + /// encodes full maps (see `engine_core_request_serializes_as_full_array`). + #[test] + fn decodes_sampling_params_with_omitted_defaults() { + let sampling_params = Value::Map(vec![ + ( + Value::from("stop_token_ids"), + Value::Array(vec![Value::from(151643u32)]), + ), + (Value::from("skip_reading_prefix_cache"), Value::from(false)), + ]); + let request = Value::Array(vec![ + Value::from("req-omit-defaults"), + Value::Array(vec![ + Value::from(1u32), + Value::from(2u32), + Value::from(3u32), + ]), + Value::Nil, + sampling_params, + Value::Nil, + Value::from(1.0f64), + ]); + + let mut bytes = Vec::new(); + rmpv::encode::write_value(&mut bytes, &request).unwrap(); + + let decoded: EngineCoreRequest = decode_msgpack(&bytes) + .expect("a real omit_defaults request must decode (regression: missing field)"); + + assert_eq!(decoded.request_id, "req-omit-defaults"); + let sampling = decoded.sampling_params.expect("sampling params present"); + + assert_eq!(sampling.stop_token_ids, vec![151643]); + assert_eq!(sampling.skip_reading_prefix_cache, Some(false)); + + // Omitted fields -> Python defaults. + assert_eq!(sampling.temperature, 1.0); + assert_eq!(sampling.top_p, 1.0); + assert_eq!(sampling.top_k, 0); + assert_eq!(sampling.seed, None); + assert_eq!(sampling.max_tokens, 16); + assert_eq!(sampling.min_tokens, 0); + assert_eq!(sampling.min_p, 0.0); + assert_eq!(sampling.frequency_penalty, 0.0); + assert_eq!(sampling.presence_penalty, 0.0); + assert_eq!(sampling.repetition_penalty, 1.0); + assert_eq!(sampling.logprobs, None); + assert_eq!(sampling.prompt_logprobs, None); + assert_eq!(sampling.eos_token_id, None); + assert!(sampling.all_stop_token_ids.is_empty()); + } } diff --git a/rust/src/engine-core-client/src/tests/client.rs b/rust/src/engine-core-client/src/tests/client.rs index 83e6cb7ec22..7d24299deaf 100644 --- a/rust/src/engine-core-client/src/tests/client.rs +++ b/rust/src/engine-core-client/src/tests/client.rs @@ -2438,6 +2438,7 @@ fn python_msgpack_fixtures_match_rust_encoding() { let stdout = String::from_utf8(output.stdout).unwrap(); let mut lines = stdout.lines(); let request_hex = lines.next().expect("missing request fixture line"); + let defaults_request_hex = lines.next().expect("missing defaults request fixture line"); let multimodal_request_hex = lines.next().expect("missing multimodal request fixture line"); let outputs_hex = lines.next().expect("missing outputs fixture line"); let inline_logprobs_frames = lines.next().expect("missing inline logprobs fixture line"); @@ -2455,6 +2456,42 @@ fn python_msgpack_fixtures_match_rust_encoding() { let expected_request = sample_request(); assert_eq!(decoded_request, expected_request); + // All-default sampling params -> empty map; must decode to Python defaults. + let defaults_request_bytes = hex::decode(defaults_request_hex).unwrap(); + let decoded_defaults: EngineCoreRequest = + rmp_serde::from_slice(&defaults_request_bytes).unwrap(); + assert_eq!(decoded_defaults.request_id, "req-defaults"); + let sampling = decoded_defaults + .sampling_params + .expect("defaults request carries sampling params"); + assert_eq!( + sampling, + EngineCoreSamplingParams { + temperature: 1.0, + top_p: 1.0, + top_k: 0, + seed: None, + max_tokens: 16, + min_tokens: 0, + logprobs: None, + prompt_logprobs: None, + min_p: 0.0, + frequency_penalty: 0.0, + presence_penalty: 0.0, + repetition_penalty: 1.0, + stop_token_ids: Vec::new(), + eos_token_id: None, + all_stop_token_ids: BTreeSet::new(), + logit_bias: None, + allowed_token_ids: None, + bad_words_token_ids: None, + structured_outputs: None, + logprob_token_ids: None, + skip_reading_prefix_cache: None, + extra_args: None, + }, + ); + let decoded_multimodal_request: EngineCoreRequest = rmp_serde::from_slice(&multimodal_request_bytes).unwrap(); assert_eq!(decoded_multimodal_request, sample_multimodal_request()); diff --git a/rust/src/engine-core-client/src/tests/python_compat.py b/rust/src/engine-core-client/src/tests/python_compat.py index 89179b3fbfe..8398c874da0 100755 --- a/rust/src/engine-core-client/src/tests/python_compat.py +++ b/rust/src/engine-core-client/src/tests/python_compat.py @@ -31,12 +31,13 @@ class FinishReason(IntEnum): REPETITION = 4 -class EngineCoreSamplingParams(msgspec.Struct, dict=True): +# Mirror of real SamplingParams; omit_defaults makes fixtures match real maps. +class EngineCoreSamplingParams(msgspec.Struct, dict=True, omit_defaults=True): temperature: float = 1.0 top_p: float = 1.0 top_k: int = 0 seed: int | None = None - max_tokens: int = 65536 + max_tokens: int = 16 min_tokens: int = 0 min_p: float = 0.0 frequency_penalty: float = 0.0 @@ -135,6 +136,16 @@ request = EngineCoreRequest( client_index=0, ) +# All defaults -> empty map. Regression guard for the sparse-map decode. +defaults_request = EngineCoreRequest( + request_id="req-defaults", + prompt_token_ids=[5, 6, 7], + mm_features=None, + sampling_params=EngineCoreSamplingParams(), + pooling_params=None, + arrival_time=1.0, +) + multimodal_tensor = np.array([[1.0, 2.0], [3.5, 4.25]], dtype=np.float32) multimodal_features = [ { @@ -361,6 +372,7 @@ ready_response = EngineCoreReadyResponse( ) print(msgspec.msgpack.encode(request).hex()) +print(msgspec.msgpack.encode(defaults_request).hex()) print(msgpack.packb(multimodal_request_wire, use_bin_type=True).hex()) print(msgspec.msgpack.encode(outputs).hex()) print(" ".join(frame.hex() for frame in encode_output_frames(inline_logprobs))) From 93bbe94d3a3a49bc5a845edfa8ac75ada54614e1 Mon Sep 17 00:00:00 2001 From: arghyadeep sarkar Date: Wed, 17 Jun 2026 12:15:55 +0530 Subject: [PATCH 08/37] [Kernel] Add weightless RMSNorm CUDA kernels for has_weight=False (#41430) (#44109) Signed-off-by: hello-args --- csrc/cpu/layernorm.cpp | 63 ++++++--- csrc/cpu/torch_bindings.cpp | 4 +- csrc/libtorch_stable/layernorm_kernels.cu | 160 ++++++++++++++-------- csrc/libtorch_stable/ops.h | 5 +- csrc/libtorch_stable/torch_bindings.cpp | 5 +- csrc/ops.h | 6 +- tests/kernels/core/test_layernorm.py | 48 +++++++ vllm/_custom_ops.py | 10 +- vllm/kernels/vllm_c.py | 7 - vllm/kernels/xpu_ops.py | 20 ++- vllm/model_executor/layers/layernorm.py | 20 +-- 11 files changed, 231 insertions(+), 117 deletions(-) diff --git a/csrc/cpu/layernorm.cpp b/csrc/cpu/layernorm.cpp index a76ad08928a..704fb146338 100644 --- a/csrc/cpu/layernorm.cpp +++ b/csrc/cpu/layernorm.cpp @@ -4,8 +4,9 @@ namespace { template void rms_norm_impl(scalar_t* __restrict__ out, const scalar_t* __restrict__ input, - const scalar_t* __restrict__ weight, const float epsilon, - const int num_tokens, const int hidden_size) { + const scalar_t* __restrict__ weight, const bool has_weight, + const float epsilon, const int num_tokens, + const int hidden_size) { using scalar_vec_t = vec_op::vec_t; constexpr int VEC_ELEM_NUM = scalar_vec_t::get_elem_num(); TORCH_CHECK(hidden_size % VEC_ELEM_NUM == 0); @@ -27,12 +28,15 @@ void rms_norm_impl(scalar_t* __restrict__ out, for (int j = 0; j < hidden_size; j += VEC_ELEM_NUM) { scalar_vec_t x(input_p + j); - scalar_vec_t w(weight + j); - vec_op::FP32Vec8 fp32_x(x); - vec_op::FP32Vec8 fp32_w(w); - - vec_op::FP32Vec8 fp32_out = fp32_x * fp32_s_variance * fp32_w; + vec_op::FP32Vec8 fp32_out; + if (has_weight) { + scalar_vec_t w(weight + j); + vec_op::FP32Vec8 fp32_w(w); + fp32_out = fp32_x * fp32_s_variance * fp32_w; + } else { + fp32_out = fp32_x * fp32_s_variance; + } scalar_vec_t out(fp32_out); out.save(output_p + j); @@ -44,8 +48,8 @@ template void fused_add_rms_norm_impl(scalar_t* __restrict__ input, scalar_t* __restrict__ residual, const scalar_t* __restrict__ weight, - const float epsilon, const int num_tokens, - const int hidden_size) { + const bool has_weight, const float epsilon, + const int num_tokens, const int hidden_size) { using scalar_vec_t = vec_op::vec_t; constexpr int VEC_ELEM_NUM = scalar_vec_t::get_elem_num(); TORCH_CHECK(hidden_size % VEC_ELEM_NUM == 0); @@ -72,13 +76,18 @@ void fused_add_rms_norm_impl(scalar_t* __restrict__ input, vec_op::FP32Vec8 fp32_s_variance(s_variance); for (int j = 0; j < hidden_size; j += VEC_ELEM_NUM) { - scalar_vec_t w(weight + j); - scalar_vec_t res(residual_p + j); - - vec_op::FP32Vec8 fp32_w(w); - vec_op::FP32Vec8 fp32_res(res); - - vec_op::FP32Vec8 fp32_out = fp32_res * fp32_s_variance * fp32_w; + vec_op::FP32Vec8 fp32_out; + if (has_weight) { + scalar_vec_t w(weight + j); + scalar_vec_t res(residual_p + j); + vec_op::FP32Vec8 fp32_w(w); + vec_op::FP32Vec8 fp32_res(res); + fp32_out = fp32_res * fp32_s_variance * fp32_w; + } else { + scalar_vec_t res(residual_p + j); + vec_op::FP32Vec8 fp32_res(res); + fp32_out = fp32_res * fp32_s_variance; + } scalar_vec_t out(fp32_out); out.save(input_p + j); @@ -87,31 +96,41 @@ void fused_add_rms_norm_impl(scalar_t* __restrict__ input, } } // namespace -void rms_norm(torch::Tensor& out, torch::Tensor& input, torch::Tensor& weight, - double epsilon) { +void rms_norm(torch::Tensor& out, torch::Tensor& input, + std::optional weight, double epsilon) { int hidden_size = input.size(-1); int num_tokens = input.numel() / hidden_size; + const bool has_weight = weight.has_value(); + if (has_weight) { + TORCH_CHECK(weight->is_contiguous()); + } VLLM_DISPATCH_FLOATING_TYPES(input.scalar_type(), "rms_norm_impl", [&] { CPU_KERNEL_GUARD_IN(rms_norm_impl) rms_norm_impl(out.data_ptr(), input.data_ptr(), - weight.data_ptr(), epsilon, num_tokens, - hidden_size); + has_weight ? weight->data_ptr() : nullptr, + has_weight, epsilon, num_tokens, hidden_size); CPU_KERNEL_GUARD_OUT(rms_norm_impl) }); } void fused_add_rms_norm(torch::Tensor& input, torch::Tensor& residual, - torch::Tensor& weight, double epsilon) { + std::optional weight, double epsilon) { int hidden_size = input.size(-1); int num_tokens = input.numel() / hidden_size; + const bool has_weight = weight.has_value(); + if (has_weight) { + TORCH_CHECK(weight->scalar_type() == input.scalar_type()); + TORCH_CHECK(weight->is_contiguous()); + } VLLM_DISPATCH_FLOATING_TYPES( input.scalar_type(), "fused_add_rms_norm_impl", [&] { CPU_KERNEL_GUARD_IN(fused_add_rms_norm_impl) fused_add_rms_norm_impl( input.data_ptr(), residual.data_ptr(), - weight.data_ptr(), epsilon, num_tokens, hidden_size); + has_weight ? weight->data_ptr() : nullptr, has_weight, + epsilon, num_tokens, hidden_size); CPU_KERNEL_GUARD_OUT(fused_add_rms_norm_impl) }); } diff --git a/csrc/cpu/torch_bindings.cpp b/csrc/cpu/torch_bindings.cpp index b1a9342deec..2aad5e2387d 100644 --- a/csrc/cpu/torch_bindings.cpp +++ b/csrc/cpu/torch_bindings.cpp @@ -310,13 +310,13 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) { // Layernorm // Apply Root Mean Square (RMS) Normalization to the input tensor. ops.def( - "rms_norm(Tensor! out, Tensor input, Tensor weight, float epsilon) -> " + "rms_norm(Tensor! out, Tensor input, Tensor? weight, float epsilon) -> " "()"); ops.impl("rms_norm", torch::kCPU, &rms_norm); // In-place fused Add and RMS Normalization. ops.def( - "fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor weight, " + "fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor? weight, " "float epsilon) -> ()"); ops.impl("fused_add_rms_norm", torch::kCPU, &fused_add_rms_norm); diff --git a/csrc/libtorch_stable/layernorm_kernels.cu b/csrc/libtorch_stable/layernorm_kernels.cu index 37df6be329f..eb121b0b880 100644 --- a/csrc/libtorch_stable/layernorm_kernels.cu +++ b/csrc/libtorch_stable/layernorm_kernels.cu @@ -11,7 +11,7 @@ namespace vllm { // TODO(woosuk): Further optimize this kernel. -template +template __global__ void rms_norm_kernel( scalar_t* __restrict__ out, // [..., hidden_size] const scalar_t* __restrict__ input, // [..., hidden_size] @@ -20,7 +20,7 @@ __global__ void rms_norm_kernel( const int64_t input_stride_d4, // input.stride(-4) const int64_t input_shape_d2, // input.size(-2) const int64_t input_shape_d3, // input.size(-3) - const scalar_t* __restrict__ weight, // [hidden_size] + const scalar_t* __restrict__ weight, // [hidden_size], null if !HasWeight const float epsilon, const int num_tokens, const int hidden_size) { __shared__ float s_variance; float variance = 0.0f; @@ -74,11 +74,19 @@ __global__ void rms_norm_kernel( for (int i = threadIdx.x; i < hidden_size / VEC_SIZE; i += blockDim.x) { vec_n_t dst; vec_n_t src1 = v_in[i]; - vec_n_t src2 = v_w[i]; + vec_n_t src2; + if constexpr (HasWeight) { + src2 = v_w[i]; + } #pragma unroll for (int j = 0; j < VEC_SIZE; j++) { float x = static_cast(src1.val[j]); - dst.val[j] = static_cast(x * s_variance) * src2.val[j]; + scalar_t normalized = static_cast(x * s_variance); + if constexpr (HasWeight) { + dst.val[j] = normalized * src2.val[j]; + } else { + dst.val[j] = normalized; + } } v_out[i] = dst; } @@ -88,13 +96,13 @@ __global__ void rms_norm_kernel( Additional optimizations we can make in this case are packed and vectorized operations, which help with the memory latency bottleneck. */ -template +template __global__ std::enable_if_t<(width > 0) && _typeConvert::exists> fused_add_rms_norm_kernel( scalar_t* __restrict__ input, // [..., hidden_size] const int64_t input_stride, scalar_t* __restrict__ residual, // [..., hidden_size] - const scalar_t* __restrict__ weight, // [hidden_size] + const scalar_t* __restrict__ weight, // [hidden_size], null if !HasWeight const float epsilon, const int num_tokens, const int hidden_size) { // Sanity checks on our vector struct and type-punned pointer arithmetic static_assert(std::is_pod_v<_f16Vec>); @@ -136,13 +144,21 @@ fused_add_rms_norm_kernel( int id = blockIdx.x * vec_hidden_size + idx; int64_t strided_id = blockIdx.x * vec_input_stride + idx; _f16Vec res = residual_v[id]; - _f16Vec w = weight_v[idx]; _f16Vec out; using Converter = _typeConvert; + if constexpr (HasWeight) { + _f16Vec w = weight_v[idx]; #pragma unroll - for (int j = 0; j < width; ++j) { - float x = Converter::convert(res.data[j]); - out.data[j] = Converter::convert(x * s_variance) * w.data[j]; + for (int j = 0; j < width; ++j) { + float x = Converter::convert(res.data[j]); + out.data[j] = Converter::convert(x * s_variance) * w.data[j]; + } + } else { +#pragma unroll + for (int j = 0; j < width; ++j) { + float x = Converter::convert(res.data[j]); + out.data[j] = Converter::convert(x * s_variance); + } } input_v[strided_id] = out; } @@ -151,13 +167,13 @@ fused_add_rms_norm_kernel( /* Generic fused_add_rms_norm_kernel The width field is not used here but necessary for other specializations. */ -template +template __global__ std::enable_if_t<(width == 0) || !_typeConvert::exists> fused_add_rms_norm_kernel( scalar_t* __restrict__ input, // [..., hidden_size] const int64_t input_stride, scalar_t* __restrict__ residual, // [..., hidden_size] - const scalar_t* __restrict__ weight, // [hidden_size] + const scalar_t* __restrict__ weight, // [hidden_size], null if !HasWeight const float epsilon, const int num_tokens, const int hidden_size) { __shared__ float s_variance; float variance = 0.0f; @@ -181,23 +197,29 @@ fused_add_rms_norm_kernel( for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) { float x = (float)residual[blockIdx.x * hidden_size + idx]; - input[blockIdx.x * input_stride + idx] = - (scalar_t)(x * s_variance) * weight[idx]; + if constexpr (HasWeight) { + input[blockIdx.x * input_stride + idx] = + (scalar_t)(x * s_variance) * weight[idx]; + } else { + input[blockIdx.x * input_stride + idx] = (scalar_t)(x * s_variance); + } } } } // namespace vllm -void rms_norm(torch::stable::Tensor& out, // [..., hidden_size] - torch::stable::Tensor& input, // [..., hidden_size] - torch::stable::Tensor& weight, // [hidden_size] +void rms_norm(torch::stable::Tensor& out, // [..., hidden_size] + torch::stable::Tensor& input, // [..., hidden_size] + std::optional weight, // [hidden_size] double epsilon) { STD_TORCH_CHECK(out.is_contiguous()); if (input.stride(-1) != 1) { input = torch::stable::contiguous(input); } STD_TORCH_CHECK(input.stride(-1) == 1); - STD_TORCH_CHECK(weight.is_contiguous()); + if (weight.has_value()) { + STD_TORCH_CHECK(weight->is_contiguous()); + } int hidden_size = input.size(-1); @@ -215,46 +237,69 @@ void rms_norm(torch::stable::Tensor& out, // [..., hidden_size] const torch::stable::accelerator::DeviceGuard device_guard( input.get_device_index()); const cudaStream_t stream = get_current_cuda_stream(); + const bool has_weight = weight.has_value(); VLLM_STABLE_DISPATCH_RANK234(num_dims, [&] { VLLM_STABLE_DISPATCH_FLOATING_TYPES( input.scalar_type(), "rms_norm_kernel", [&] { + const scalar_t* weight_ptr = + has_weight ? weight->const_data_ptr() : nullptr; const int calculated_vec_size = std::gcd(16 / sizeof(scalar_t), hidden_size); const int block_size = std::min(hidden_size / calculated_vec_size, max_block_size); dim3 block(block_size); VLLM_STABLE_DISPATCH_VEC_SIZE(calculated_vec_size, [&] { - vllm::rms_norm_kernel - <<>>( - out.mutable_data_ptr(), - input.const_data_ptr(), input_stride_d2, - input_stride_d3, input_stride_d4, input_shape_d2, - input_shape_d3, weight.const_data_ptr(), epsilon, - num_tokens, hidden_size); + if (has_weight) { + vllm::rms_norm_kernel + <<>>( + out.mutable_data_ptr(), + input.const_data_ptr(), input_stride_d2, + input_stride_d3, input_stride_d4, input_shape_d2, + input_shape_d3, weight_ptr, epsilon, num_tokens, + hidden_size); + } else { + vllm::rms_norm_kernel + <<>>( + out.mutable_data_ptr(), + input.const_data_ptr(), input_stride_d2, + input_stride_d3, input_stride_d4, input_shape_d2, + input_shape_d3, weight_ptr, epsilon, num_tokens, + hidden_size); + } }); }); }); } -#define LAUNCH_FUSED_ADD_RMS_NORM(width) \ - VLLM_STABLE_DISPATCH_FLOATING_TYPES( \ - input.scalar_type(), "fused_add_rms_norm_kernel", [&] { \ - vllm::fused_add_rms_norm_kernel \ - <<>>( \ - input.mutable_data_ptr(), input_stride, \ - residual.mutable_data_ptr(), \ - weight.const_data_ptr(), epsilon, num_tokens, \ - hidden_size); \ +#define LAUNCH_FUSED_ADD_RMS_NORM(width, has_weight) \ + VLLM_STABLE_DISPATCH_FLOATING_TYPES( \ + input.scalar_type(), "fused_add_rms_norm_kernel", [&] { \ + if (has_weight) { \ + vllm::fused_add_rms_norm_kernel \ + <<>>( \ + input.mutable_data_ptr(), input_stride, \ + residual.mutable_data_ptr(), \ + weight->const_data_ptr(), epsilon, num_tokens, \ + hidden_size); \ + } else { \ + vllm::fused_add_rms_norm_kernel \ + <<>>( \ + input.mutable_data_ptr(), input_stride, \ + residual.mutable_data_ptr(), nullptr, epsilon, \ + num_tokens, hidden_size); \ + } \ }); void fused_add_rms_norm(torch::stable::Tensor& input, // [..., hidden_size] torch::stable::Tensor& residual, // [..., hidden_size] - torch::stable::Tensor& weight, // [hidden_size] + std::optional weight, double epsilon) { - STD_TORCH_CHECK(weight.scalar_type() == input.scalar_type()); STD_TORCH_CHECK(input.scalar_type() == residual.scalar_type()); STD_TORCH_CHECK(residual.is_contiguous()); - STD_TORCH_CHECK(weight.is_contiguous()); + if (weight.has_value()) { + STD_TORCH_CHECK(weight->scalar_type() == input.scalar_type()); + STD_TORCH_CHECK(weight->is_contiguous()); + } int hidden_size = input.size(-1); int64_t input_stride = input.stride(-2); int num_tokens = input.numel() / hidden_size; @@ -269,30 +314,33 @@ void fused_add_rms_norm(torch::stable::Tensor& input, // [..., hidden_size] const torch::stable::accelerator::DeviceGuard device_guard( input.get_device_index()); const cudaStream_t stream = get_current_cuda_stream(); - /*If the tensor types are FP16/BF16, try to use the optimized kernel - with packed + vectorized ops. - Max optimization is achieved with a width-8 vector of FP16/BF16s - since we can load at most 128 bits at once in a global memory op. - However, this requires each tensor's data to be aligned to 16 - bytes. - */ + constexpr int vector_width = 8; + constexpr int req_alignment_bytes = vector_width * 2; auto inp_ptr = reinterpret_cast(input.data_ptr()); auto res_ptr = reinterpret_cast(residual.data_ptr()); - auto wt_ptr = reinterpret_cast(weight.data_ptr()); - constexpr int vector_width = 8; - constexpr int req_alignment_bytes = - vector_width * 2; // vector_width * sizeof(bfloat16 or float16) (float32 - // falls back to non-vectorized version anyway) - bool ptrs_are_aligned = inp_ptr % req_alignment_bytes == 0 && - res_ptr % req_alignment_bytes == 0 && - wt_ptr % req_alignment_bytes == 0; bool offsets_are_multiple_of_vector_width = hidden_size % vector_width == 0 && input_stride % vector_width == 0; bool batch_invariant_launch = vllm::vllm_is_batch_invariant(); - if (ptrs_are_aligned && offsets_are_multiple_of_vector_width && - !batch_invariant_launch) { - LAUNCH_FUSED_ADD_RMS_NORM(8); + const bool has_weight = weight.has_value(); + if (has_weight) { + auto wt_ptr = reinterpret_cast(weight->data_ptr()); + bool ptrs_are_aligned = inp_ptr % req_alignment_bytes == 0 && + res_ptr % req_alignment_bytes == 0 && + wt_ptr % req_alignment_bytes == 0; + if (ptrs_are_aligned && offsets_are_multiple_of_vector_width && + !batch_invariant_launch) { + LAUNCH_FUSED_ADD_RMS_NORM(8, true); + } else { + LAUNCH_FUSED_ADD_RMS_NORM(0, true); + } } else { - LAUNCH_FUSED_ADD_RMS_NORM(0); + bool ptrs_are_aligned = inp_ptr % req_alignment_bytes == 0 && + res_ptr % req_alignment_bytes == 0; + if (ptrs_are_aligned && offsets_are_multiple_of_vector_width && + !batch_invariant_launch) { + LAUNCH_FUSED_ADD_RMS_NORM(8, false); + } else { + LAUNCH_FUSED_ADD_RMS_NORM(0, false); + } } } diff --git a/csrc/libtorch_stable/ops.h b/csrc/libtorch_stable/ops.h index 49ce1744006..9efc12e9f49 100644 --- a/csrc/libtorch_stable/ops.h +++ b/csrc/libtorch_stable/ops.h @@ -188,11 +188,12 @@ torch::stable::Tensor hadacore_transform(torch::stable::Tensor& x, // Layernorm kernels (shared CUDA/ROCm) void rms_norm(torch::stable::Tensor& out, torch::stable::Tensor& input, - torch::stable::Tensor& weight, double epsilon); + std::optional weight, double epsilon); void fused_add_rms_norm(torch::stable::Tensor& input, torch::stable::Tensor& residual, - torch::stable::Tensor& weight, double epsilon); + std::optional weight, + double epsilon); // Layernorm-quant kernels (shared CUDA/ROCm) void rms_norm_static_fp8_quant(torch::stable::Tensor& out, diff --git a/csrc/libtorch_stable/torch_bindings.cpp b/csrc/libtorch_stable/torch_bindings.cpp index f360c83b8af..0aabcc757dc 100644 --- a/csrc/libtorch_stable/torch_bindings.cpp +++ b/csrc/libtorch_stable/torch_bindings.cpp @@ -369,12 +369,13 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) { // Apply Root Mean Square (RMS) Normalization to the input tensor. ops.def( - "rms_norm(Tensor! result, Tensor input, Tensor weight, float epsilon) -> " + "rms_norm(Tensor! result, Tensor input, Tensor? weight, float epsilon) " + "-> " "()"); // In-place fused Add and RMS Normalization. ops.def( - "fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor weight, " + "fused_add_rms_norm(Tensor! input, Tensor! residual, Tensor? weight, " "float epsilon) -> ()"); // Layernorm-quant diff --git a/csrc/ops.h b/csrc/ops.h index b909c5711d4..ec3f5e187cc 100644 --- a/csrc/ops.h +++ b/csrc/ops.h @@ -34,11 +34,11 @@ torch::Tensor weak_ref_tensor(torch::Tensor& tensor) { // rms_norm and fused_add_rms_norm declarations also exist in // csrc/libtorch_stable/ops.h (torch::stable ABI for CUDA). They remain here // because the CPU build still uses these torch::Tensor declarations. -void rms_norm(torch::Tensor& out, torch::Tensor& input, torch::Tensor& weight, - double epsilon); +void rms_norm(torch::Tensor& out, torch::Tensor& input, + std::optional weight, double epsilon); void fused_add_rms_norm(torch::Tensor& input, torch::Tensor& residual, - torch::Tensor& weight, double epsilon); + std::optional weight, double epsilon); // rotary_embedding also exist in csrc/libtorch_stable/ops.h (torch::stable // ABI for CUDA). It remains here because the CPU build still uses these diff --git a/tests/kernels/core/test_layernorm.py b/tests/kernels/core/test_layernorm.py index c39d42c7593..fde09710b5d 100644 --- a/tests/kernels/core/test_layernorm.py +++ b/tests/kernels/core/test_layernorm.py @@ -6,6 +6,7 @@ import torch from tests.kernels.quant_utils import FP8_DTYPE from tests.kernels.utils import opcheck +from vllm import ir from vllm.model_executor.layers.layernorm import GemmaRMSNorm, RMSNorm from vllm.platforms import current_platform from vllm.utils.torch_utils import set_random_seed @@ -27,6 +28,10 @@ CUDA_DEVICES = [ ] +def _rms_norm_tolerance(dtype: torch.dtype) -> dict[str, float]: + return ir.ops.rms_norm.get_tolerance(dtype) + + @pytest.mark.parametrize("num_tokens", NUM_TOKENS) @pytest.mark.parametrize("hidden_size", HIDDEN_SIZES) @pytest.mark.parametrize("add_residual", ADD_RESIDUAL) @@ -81,6 +86,49 @@ def test_rms_norm( ) +@pytest.mark.parametrize("num_tokens", NUM_TOKENS) +@pytest.mark.parametrize("hidden_size", HIDDEN_SIZES) +@pytest.mark.parametrize("add_residual", ADD_RESIDUAL) +@pytest.mark.parametrize("dtype", DTYPES) +@pytest.mark.parametrize("seed", SEEDS) +@pytest.mark.parametrize("device", CUDA_DEVICES) +@torch.inference_mode() +def test_rms_norm_weightless( + default_vllm_config, + num_tokens: int, + hidden_size: int, + add_residual: bool, + dtype: torch.dtype, + seed: int, + device: str, +) -> None: + set_random_seed(seed) + torch.set_default_device(device) + layer = RMSNorm(hidden_size, has_weight=False).to(dtype=dtype) + x = torch.randn(num_tokens, hidden_size, dtype=dtype) + residual = torch.randn_like(x) if add_residual else None + + ref_out = layer.forward_native(x, residual) + out = layer(x, residual) + tol = _rms_norm_tolerance(dtype) + if add_residual: + torch.testing.assert_close(out[0], ref_out[0], **tol) + torch.testing.assert_close(out[1], ref_out[1], **tol) + else: + torch.testing.assert_close(out, ref_out, **tol) + + if residual is not None: + opcheck( + torch.ops._C.fused_add_rms_norm, + (x, residual, None, layer.variance_epsilon), + ) + else: + opcheck( + torch.ops._C.rms_norm, + (out, x, None, layer.variance_epsilon), + ) + + @pytest.mark.parametrize("num_tokens", NUM_TOKENS) @pytest.mark.parametrize("hidden_size", HIDDEN_SIZES) @pytest.mark.parametrize("add_residual", ADD_RESIDUAL) diff --git a/vllm/_custom_ops.py b/vllm/_custom_ops.py index 5f52bc50eb2..1b49c9159dc 100644 --- a/vllm/_custom_ops.py +++ b/vllm/_custom_ops.py @@ -315,13 +315,19 @@ def rotary_embedding( # layer norm ops def rms_norm( - out: torch.Tensor, input: torch.Tensor, weight: torch.Tensor, epsilon: float + out: torch.Tensor, + input: torch.Tensor, + weight: torch.Tensor | None, + epsilon: float, ) -> None: torch.ops._C.rms_norm(out, input, weight, epsilon) def fused_add_rms_norm( - input: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor, epsilon: float + input: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor | None, + epsilon: float, ) -> None: # Note: this func is batch invariant torch.ops._C.fused_add_rms_norm(input, residual, weight, epsilon) diff --git a/vllm/kernels/vllm_c.py b/vllm/kernels/vllm_c.py index 3b194b2ab93..6ae5d9939e3 100644 --- a/vllm/kernels/vllm_c.py +++ b/vllm/kernels/vllm_c.py @@ -25,9 +25,6 @@ rms_no_var_size = lambda x, weight, epsilon, variance_size=None: ( def rms_norm( x: Tensor, weight: Tensor | None, epsilon: float, variance_size: int | None = None ) -> Tensor: - if weight is None: - # Kernel requires weight tensor, pass ones - weight = torch.ones(x.shape[-1], device=x.device, dtype=x.dtype) assert variance_size is None # ROCm's vLLM C RMSNorm kernel operates on contiguous 2D tensors. # Higher-rank callers still normalize over the last dimension, so flatten @@ -64,10 +61,6 @@ def fused_add_rms_norm( epsilon: float, variance_size: int | None = None, ) -> tuple[Tensor, Tensor]: - if weight is None: - # Kernel requires weight tensor, pass ones - weight = torch.ones(x.shape[-1], device=x.device, dtype=x.dtype) - assert variance_size is None if IS_ROCM and (not x.is_contiguous() or not x_residual.is_contiguous()): output, residual = ir.ops.fused_add_rms_norm.impls["native"].impl_fn( diff --git a/vllm/kernels/xpu_ops.py b/vllm/kernels/xpu_ops.py index 5e7f90f7086..8a86b1226b4 100644 --- a/vllm/kernels/xpu_ops.py +++ b/vllm/kernels/xpu_ops.py @@ -29,10 +29,12 @@ rms_no_var = lambda x, weight, epsilon, variance_size=None: variance_size is Non def rms_norm( x: Tensor, weight: Tensor | None, epsilon: float, variance_size: int | None = None ) -> Tensor: - if weight is None: - # Kernel requires weight tensor, pass ones - weight = torch.ones(x.shape[-1], device=x.device, dtype=x.dtype) assert variance_size is None + if weight is None: + # Weightless _C ops are CUDA-only; native skips the multiply on XPU. + return ir.ops.rms_norm.impls["native"].impl_fn( + x, weight, epsilon, variance_size + ) output = torch.empty(x.shape, device=x.device, dtype=x.dtype) torch.ops._C.rms_norm(output, x, weight, epsilon) return output @@ -57,10 +59,14 @@ def fused_add_rms_norm( epsilon: float, variance_size: int | None = None, ) -> tuple[Tensor, Tensor]: - if weight is None: - # Kernel requires weight tensor, pass ones - weight = torch.ones(x.shape[-1], device=x.device, dtype=x.dtype) - assert variance_size is None + if weight is None: + # Weightless _C ops are CUDA-only; native skips the multiply on XPU. + output, residual = ir.ops.fused_add_rms_norm.impls["native"].impl_fn( + x, x_residual, weight, epsilon, variance_size + ) + x.copy_(output) + x_residual.copy_(residual) + return x, x_residual torch.ops._C.fused_add_rms_norm(x, x_residual, weight, epsilon) return x, x_residual diff --git a/vllm/model_executor/layers/layernorm.py b/vllm/model_executor/layers/layernorm.py index 13b0ae78131..8418245b825 100644 --- a/vllm/model_executor/layers/layernorm.py +++ b/vllm/model_executor/layers/layernorm.py @@ -9,7 +9,6 @@ import torch.nn.functional as F # Import kernels import vllm.kernels # noqa: F401 from vllm import envs, ir -from vllm.config import get_current_vllm_config from vllm.logger import init_logger from vllm.model_executor.custom_op import CustomOp from vllm.model_executor.layers.batch_invariant import rms_norm_batch_invariant @@ -65,19 +64,12 @@ class RMSNorm(CustomOp): if self.has_weight: self.weight = nn.Parameter(self.weight) - # Do not pass identity weight to native implementation (causes issue on TPU). - # Other implementations require weight to be passed even if all ones. - # Cheat and predict if native will be dispatched to: - # 1) if native is first in priority list - # 2) if variance_size_override is given (only supported by native impl) - # TODO(luka): address weight passing inconsistency: - # https://github.com/vllm-project/vllm/issues/39370 - priority = get_current_vllm_config().kernel_config.ir_op_priority - var_override = self.variance_size_override is not None - native_rms_norm = priority.rms_norm[0] == "native" or var_override - native_add_rms_norm = priority.fused_add_rms_norm[0] == "native" or var_override - self.pass_weight = self.has_weight or not native_rms_norm - self.pass_weight_add = self.has_weight or not native_add_rms_norm + # When has_weight=False, pass weight=None so implementations that + # support a weightless path can skip the per-channel multiply. + # Implementations that require weight (e.g. oink) fall back via IR + # op priority when weight=None is unsupported. + self.pass_weight = self.has_weight + self.pass_weight_add = self.has_weight def forward_native( self, From 43fa24e832971da51c97fc434eac2c2de870a446 Mon Sep 17 00:00:00 2001 From: Taneem Ibrahim Date: Wed, 17 Jun 2026 01:57:12 -0500 Subject: [PATCH 09/37] [Misc] Validate Cohere Embed Mixed Content Payloads (#45873) Signed-off-by: Taneem Ibrahim --- .../pooling/embed/test_io_processor.py | 42 +++++++++++++++++++ vllm/entrypoints/pooling/embed/protocol.py | 11 +++++ 2 files changed, 53 insertions(+) diff --git a/tests/entrypoints/pooling/embed/test_io_processor.py b/tests/entrypoints/pooling/embed/test_io_processor.py index f0dea740440..8f8f8faa8ad 100644 --- a/tests/entrypoints/pooling/embed/test_io_processor.py +++ b/tests/entrypoints/pooling/embed/test_io_processor.py @@ -150,6 +150,16 @@ class TestCohereEmbedRequestParsing: {"content": [{"type": "text", "text": "hello"}]}, ], }, + { + "model": "test", + "inputs": [ + { + "content": [ + {"type": "image_url", "image_url": {"url": "image-uri"}} + ] + }, + ], + }, ], ) def test_accepts_exactly_one_non_empty_input_field(self, request_body): @@ -157,6 +167,38 @@ class TestCohereEmbedRequestParsing: assert request.model == "test" + @pytest.mark.parametrize( + ("content", "error"), + [ + ( + {"type": "text"}, + "CohereEmbedContent with type='text' requires text", + ), + ( + {"type": "image_url"}, + "CohereEmbedContent with type='image_url' requires image_url.url", + ), + ( + {"type": "image_url", "image_url": {}}, + "CohereEmbedContent with type='image_url' requires image_url.url", + ), + ( + {"type": "image_url", "image_url": {"url": ""}}, + "CohereEmbedContent with type='image_url' requires image_url.url", + ), + ], + ) + def test_rejects_invalid_mixed_content_payloads(self, content, error): + with pytest.raises(ValidationError, match=error): + CohereEmbedRequest( + model="test", + inputs=[ + { + "content": [content], + }, + ], + ) + class TestResolveTruncation: """Unit tests for EmbedIOProcessor._resolve_cohere_truncation.""" diff --git a/vllm/entrypoints/pooling/embed/protocol.py b/vllm/entrypoints/pooling/embed/protocol.py index 8ec908f4511..2dcc848c8c7 100644 --- a/vllm/entrypoints/pooling/embed/protocol.py +++ b/vllm/entrypoints/pooling/embed/protocol.py @@ -207,6 +207,17 @@ class CohereEmbedContent(BaseModel): text: str | None = None image_url: dict[str, str] | None = None + @model_validator(mode="after") + def validate_content_payload(self): + if self.type == "text": + if self.text is None: + raise ValueError("CohereEmbedContent with type='text' requires text") + elif not self.image_url or not self.image_url.get("url"): + raise ValueError( + "CohereEmbedContent with type='image_url' requires image_url.url" + ) + return self + class CohereEmbedInput(BaseModel): content: list[CohereEmbedContent] From e2c58570eaae740fae1a5e50a01ceb18aa45c687 Mon Sep 17 00:00:00 2001 From: Bugen Zhao Date: Wed, 17 Jun 2026 15:32:40 +0800 Subject: [PATCH 10/37] [Rust Frontend] Support hybrid/external DP LB in Python supervised bootstrap (#45805) Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> Signed-off-by: Bugen Zhao --- .buildkite/test_areas/rust_frontend.yaml | 4 + rust/src/cmd/src/cli.rs | 7 ++ rust/src/cmd/src/cli/tests.rs | 8 +- rust/src/engine-core-client/src/client.rs | 5 + .../engine-core-client/src/tests/client.rs | 113 ++++++++++++++++++ rust/src/engine-core-client/src/transport.rs | 5 +- vllm/entrypoints/cli/serve.py | 9 +- vllm/v1/utils.py | 14 ++- 8 files changed, 159 insertions(+), 6 deletions(-) diff --git a/.buildkite/test_areas/rust_frontend.yaml b/.buildkite/test_areas/rust_frontend.yaml index 16d69f77345..f9abac2004e 100644 --- a/.buildkite/test_areas/rust_frontend.yaml +++ b/.buildkite/test_areas/rust_frontend.yaml @@ -99,9 +99,13 @@ steps: - vllm/v1/engine/ - vllm/v1/worker/ - tests/utils.py + - tests/v1/distributed/test_external_lb_dp.py + - tests/v1/distributed/test_hybrid_lb_dp.py - tests/v1/distributed/test_internal_lb_dp.py commands: - export VLLM_USE_RUST_FRONTEND=1 - export VLLM_WORKER_MULTIPROC_METHOD=spawn - export NCCL_CUMEM_HOST_ENABLE=0 - TP_SIZE=1 DP_SIZE=4 pytest -v -s v1/distributed/test_internal_lb_dp.py -k "not 4 and not server_info" + - TP_SIZE=1 DP_SIZE=2 pytest -v -s v1/distributed/test_external_lb_dp.py -k "not 4 and not server_info" + - TP_SIZE=1 DP_SIZE=4 pytest -v -s v1/distributed/test_hybrid_lb_dp.py -k "not 4 and not server_info" diff --git a/rust/src/cmd/src/cli.rs b/rust/src/cmd/src/cli.rs index 14291d7de05..003d96fa92b 100644 --- a/rust/src/cmd/src/cli.rs +++ b/rust/src/cmd/src/cli.rs @@ -290,6 +290,7 @@ impl SharedRuntimeArgs { input_address: String, output_address: String, coordinator_address: Option, + engine_start_index: u32, engine_count: usize, ) -> Config { let ready_timeout = self.ready_timeout(); @@ -301,6 +302,7 @@ impl SharedRuntimeArgs { transport_mode: TransportMode::Bootstrapped { input_address, output_address, + engine_start_index, engine_count, ready_timeout, }, @@ -435,6 +437,10 @@ pub struct FrontendArgs { /// `stats_update_address`. #[arg(long)] pub coordinator_address: Option, + /// First data-parallel engine rank expected to register with this + /// bootstrapped frontend. + #[arg(long, default_value_t = 0)] + pub engine_start_index: u32, /// Total number of data-parallel engines expected for this frontend. #[arg(long, default_value_t = 1)] pub engine_count: usize, @@ -452,6 +458,7 @@ impl FrontendArgs { self.input_address, self.output_address, self.coordinator_address, + self.engine_start_index, self.engine_count, ) } diff --git a/rust/src/cmd/src/cli/tests.rs b/rust/src/cmd/src/cli/tests.rs index 31ace0a77df..c57c23e017c 100644 --- a/rust/src/cmd/src/cli/tests.rs +++ b/rust/src/cmd/src/cli/tests.rs @@ -424,6 +424,7 @@ fn frontend_args_accept_json() { coordinator_address: Some( "tcp://127.0.0.1:7000", ), + engine_start_index: 0, engine_count: 1, runtime: SharedRuntimeArgs { model: "Qwen/Qwen3-0.6B", @@ -1205,8 +1206,10 @@ fn frontend_config_uses_external_coordinator_when_coordinator_address_is_present "ipc:///tmp/output.sock", "--coordinator-address", "tcp://127.0.0.1:7000", + "--engine-start-index", + "3", "--engine-count", - "2", + "1", "--args-json", r#"{"model_tag":"Qwen/Qwen3-0.6B"}"#, ]) @@ -1222,7 +1225,8 @@ fn frontend_config_uses_external_coordinator_when_coordinator_address_is_present transport_mode: Bootstrapped { input_address: "ipc:///tmp/input.sock", output_address: "ipc:///tmp/output.sock", - engine_count: 2, + engine_start_index: 3, + engine_count: 1, ready_timeout: 600s, }, coordinator_mode: External { diff --git a/rust/src/engine-core-client/src/client.rs b/rust/src/engine-core-client/src/client.rs index c646de567d0..b4357f77c7c 100644 --- a/rust/src/engine-core-client/src/client.rs +++ b/rust/src/engine-core-client/src/client.rs @@ -56,6 +56,9 @@ pub enum TransportMode { /// Output PULL socket address that engines will connect to for /// responses. output_address: String, + /// First data-parallel engine rank expected to register on this + /// transport. + engine_start_index: u32, /// Total number of engines expected to register on this transport. engine_count: usize, /// Maximum time to wait for all expected engines to register. @@ -246,6 +249,7 @@ impl EngineCoreClient { TransportMode::Bootstrapped { input_address, output_address, + engine_start_index, engine_count, ready_timeout, } => { @@ -256,6 +260,7 @@ impl EngineCoreClient { transport::connect_bootstrapped( input_address, output_address, + *engine_start_index, *engine_count, *ready_timeout, ) diff --git a/rust/src/engine-core-client/src/tests/client.rs b/rust/src/engine-core-client/src/tests/client.rs index 7d24299deaf..c00a4226854 100644 --- a/rust/src/engine-core-client/src/tests/client.rs +++ b/rust/src/engine-core-client/src/tests/client.rs @@ -303,6 +303,7 @@ fn bootstrapped_test_config( transport_mode: TransportMode::Bootstrapped { input_address, output_address, + engine_start_index: 0, engine_count, ready_timeout, }, @@ -312,6 +313,34 @@ fn bootstrapped_test_config( } } +fn bootstrapped_test_config_with_start_index( + input_address: String, + output_address: String, + engine_start_index: u32, + engine_count: usize, + ready_timeout: Duration, + client_index: u32, + coordinator_mode: Option, +) -> EngineCoreClientConfig { + let mut config = bootstrapped_test_config( + input_address, + output_address, + engine_count, + ready_timeout, + client_index, + coordinator_mode, + ); + let TransportMode::Bootstrapped { + engine_start_index: start, + .. + } = &mut config.transport_mode + else { + unreachable!("bootstrapped_test_config returns bootstrapped transport") + }; + *start = engine_start_index; + config +} + async fn recv_xpub_message(xpub: &mut XPubSocket) -> Vec { xpub.recv().await.unwrap().into_vec() } @@ -2689,6 +2718,90 @@ async fn bootstrapped_connects_with_contiguous_engine_ids() { client.shutdown().await.unwrap(); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn bootstrapped_connects_with_nonzero_engine_start_index() { + init_tracing(); + let ipc = IpcNamespace::new().unwrap(); + let input_address = ipc.input_endpoint(); + let output_address = ipc.output_endpoint(); + + let client_task = tokio::spawn({ + let input_address = input_address.clone(); + let output_address = output_address.clone(); + async move { + EngineCoreClient::connect(bootstrapped_test_config_with_start_index( + input_address, + output_address, + 3, + 1, + Duration::from_secs(2), + 0, + None, + )) + .await + .unwrap() + } + }); + + let (_dealer, _push) = + setup_bootstrapped_mock_engine(input_address, output_address, &[0x03, 0x00]).await; + let client = client_task.await.unwrap(); + + assert_eq!(client.engine_count(), 1); + let engine_ids = + client.engine_identities().into_iter().map(|id| id.to_vec()).collect::>(); + assert_eq!(engine_ids, vec![vec![0x03, 0x00]]); + + client.shutdown().await.unwrap(); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn bootstrapped_rejects_unexpected_engine_id_for_start_index() { + init_tracing(); + let ipc = IpcNamespace::new().unwrap(); + let input_address = ipc.input_endpoint(); + let output_address = ipc.output_endpoint(); + + let client_task = tokio::spawn({ + let input_address = input_address.clone(); + let output_address = output_address.clone(); + async move { + EngineCoreClient::connect(bootstrapped_test_config_with_start_index( + input_address, + output_address, + 3, + 1, + Duration::from_secs(2), + 0, + None, + )) + .await + } + }); + + let _ = crate::mock_engine::connect_to_bootstrapped_frontend( + input_address, + output_address, + &[0x00, 0x00], + crate::mock_engine::MockEngineConfig { + local: true, + headless: true, + ..Default::default() + }, + ) + .await; + let error = match client_task.await.unwrap() { + Ok(_) => panic!("bootstrapped connect should reject unexpected engine id"), + Err(error) => error, + }; + + assert!( + error + .to_string() + .contains("received input registration for unexpected engine id") + ); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn bootstrapped_connect_times_out_without_registration() { init_tracing(); diff --git a/rust/src/engine-core-client/src/transport.rs b/rust/src/engine-core-client/src/transport.rs index 360f94eda12..d0d9b4efe39 100644 --- a/rust/src/engine-core-client/src/transport.rs +++ b/rust/src/engine-core-client/src/transport.rs @@ -327,6 +327,7 @@ pub async fn connect_handshake( pub async fn connect_bootstrapped( input_address: &str, output_address: &str, + engine_start_index: u32, engine_count: usize, ready_timeout: Duration, ) -> Result { @@ -342,8 +343,8 @@ pub async fn connect_bootstrapped( let engines = wait_for_input_registrations( &mut input_socket, - // TODO: follow start rank - (0..engine_count).map(|index| EngineId::from((index as u16).to_le_bytes().to_vec())), + (0..engine_count) + .map(|offset| EngineId::from_engine_index(engine_start_index + offset as u32)), ready_timeout, ) .await?; diff --git a/vllm/entrypoints/cli/serve.py b/vllm/entrypoints/cli/serve.py index 415ab549cc7..8491e982165 100644 --- a/vllm/entrypoints/cli/serve.py +++ b/vllm/entrypoints/cli/serve.py @@ -328,6 +328,12 @@ def run_multi_api_server(args: argparse.Namespace): ) if rust_frontend_path: + if parallel_config.local_engines_only: + expected_engine_start_index = parallel_config.data_parallel_rank + expected_engine_count = parallel_config.data_parallel_size_local + else: + expected_engine_start_index = 0 + expected_engine_count = parallel_config.data_parallel_size # Start rust front-end process. api_server_manager = RustFrontendProcessManager( binary_path=rust_frontend_path, @@ -335,7 +341,8 @@ def run_multi_api_server(args: argparse.Namespace): args=args, input_address=addresses.inputs[0], output_address=addresses.outputs[0], - engine_count=parallel_config.data_parallel_size, + engine_start_index=expected_engine_start_index, + engine_count=expected_engine_count, stats_update_address=stats_update_address, ) else: diff --git a/vllm/v1/utils.py b/vllm/v1/utils.py index 9fd13a941e0..ba66358c66f 100644 --- a/vllm/v1/utils.py +++ b/vllm/v1/utils.py @@ -336,6 +336,7 @@ class RustFrontendProcessManager: args: argparse.Namespace, input_address: str, output_address: str, + engine_start_index: int, engine_count: int, stats_update_address: str | None = None, ): @@ -354,6 +355,8 @@ class RustFrontendProcessManager: input_address, "--output-address", output_address, + "--engine-start-index", + str(engine_start_index), "--engine-count", str(engine_count), ] @@ -362,7 +365,16 @@ class RustFrontendProcessManager: from vllm.entrypoints.serve.utils.api_utils import jsonify_non_default_args args_json = json.dumps( - jsonify_non_default_args(args, exclude={"api_server_count"}), + jsonify_non_default_args( + args, + exclude={ + "api_server_count", + # Python passes the bootstrapped engine range explicitly. + "data_parallel_rank", + "data_parallel_external_lb", + "data_parallel_hybrid_lb", + }, + ), sort_keys=True, ) cmd.extend(["--args-json", args_json]) From 6d8fff5698aea7225792b37c3b075cc64c4a0178 Mon Sep 17 00:00:00 2001 From: Itay Etelis <92247226+Etelis@users.noreply.github.com> Date: Wed, 17 Jun 2026 11:35:07 +0300 Subject: [PATCH 11/37] [KV Connector][Offloading] Avoid blocking the engine to flush offloads on idle (#45595) Signed-off-by: Itay Etelis Signed-off-by: Or Ozeri Signed-off-by: Itay Etelis Co-authored-by: Itay Etelis Co-authored-by: Or Ozeri Co-authored-by: Itay Etelis --- .../offloading_connector/test_scheduler.py | 85 +++---------------- .../kv_connector/v1/offloading/scheduler.py | 16 ++-- .../kv_connector/v1/offloading_connector.py | 4 + vllm/v1/kv_offload/base.py | 8 ++ vllm/v1/kv_offload/tiering/base.py | 8 ++ vllm/v1/kv_offload/tiering/manager.py | 9 ++ 6 files changed, 49 insertions(+), 81 deletions(-) diff --git a/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py b/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py index 8bd46184d64..973fcc63e31 100644 --- a/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py +++ b/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py @@ -262,12 +262,9 @@ def test_concurrent_lookups_of_the_same_prefix(request_runner, async_scheduling: runner.manager.prepare_store.side_effect = lambda keys, req_context: ( generate_store_output(keys) ) - # With sync scheduling, all-finished flush fires within this run. - # With async scheduling, the finish is delayed so flush fires later. runner.run( decoded_tokens=[EOS_TOKEN_ID], expected_stored=(0, 1, 2), - expected_flushed=(0, 1, 2) if not async_scheduling else (), ) # start a request to load the first block, but don't complete @@ -332,7 +329,6 @@ def test_abort_loading_requests(request_runner, async_scheduling: bool): runner.run( decoded_tokens=[EOS_TOKEN_ID], expected_stored=(0, 1, 2), - expected_flushed=(0, 1, 2) if not async_scheduling else (), ) # start a request to load the first block, but don't complete @@ -359,7 +355,6 @@ def test_abort_loading_requests(request_runner, async_scheduling: bool): runner.run( decoded_tokens=[], expected_loaded=(0, 1, 2), - expected_flushed=(0, 1, 2), ) # assert request is deleted @@ -774,7 +769,6 @@ def test_request_level_policy_stores_all_blocks(request_runner, async_scheduling runner.run( decoded_tokens=[EOS_TOKEN_ID], expected_stored=(0, 1, 2), - expected_flushed=(0, 1, 2) if not async_scheduling else (), ) # Reset GPU prefix cache so the next request must load from CPU. @@ -841,13 +835,8 @@ def test_fence_at_update_state_after_alloc(request_runner): runner.manager.prepare_store.side_effect = lambda keys, req_context: ( generate_store_output(keys) ) - runner.run( - decoded_tokens=[EOS_TOKEN_ID], - complete_transfers=False, - expected_stored=(0,), - expected_flushed=(0,), - ) - assert runner.connector_scheduler._block_id_to_pending_jobs == {} + runner.run(decoded_tokens=[EOS_TOKEN_ID], complete_transfers=False) + assert runner.connector_scheduler._block_id_to_pending_jobs runner.scheduler.reset_prefix_cache() runner.new_request(token_ids=[0] * 4) @@ -858,6 +847,8 @@ def test_fence_at_update_state_after_alloc(request_runner): runner.run( decoded_tokens=[], complete_transfers=False, + expected_stored=(0,), + expected_flushed=(0,), ) assert runner.connector_scheduler._block_id_to_pending_jobs == {} @@ -877,13 +868,8 @@ def test_fence_at_build_store_jobs(request_runner): runner.manager.prepare_store.side_effect = lambda keys, req_context: ( generate_store_output(keys) ) - runner.run( - decoded_tokens=[EOS_TOKEN_ID], - complete_transfers=False, - expected_stored=(0,), - expected_flushed=(0,), - ) - assert runner.connector_scheduler._block_id_to_pending_jobs == {} + runner.run(decoded_tokens=[EOS_TOKEN_ID], complete_transfers=False) + assert runner.connector_scheduler._block_id_to_pending_jobs runner.scheduler.reset_prefix_cache() runner.new_request(token_ids=[1] * 4) @@ -893,6 +879,8 @@ def test_fence_at_build_store_jobs(request_runner): ) runner.run( decoded_tokens=[EOS_TOKEN_ID], + expected_stored=(0,), + expected_flushed=(0,), ) assert runner.connector_scheduler._block_id_to_pending_jobs == {} @@ -966,10 +954,10 @@ def test_max_offload_tokens_validation(request_runner, async_scheduling: bool): lambda keys, req_context: generate_store_output(keys) ) - # With sync scheduling, the connector flushes completed stores when the - # request finishes; async scheduling defers the flush to the next step. - flushed_all = all_offsets if not async_scheduling else () - flushed_two = (0, 1, 2, 3, 4, 5) if not async_scheduling else () + # Pending offloads drain via non-blocking stepping, not a flush, so no + # blocks are flushed when the request finishes. + flushed_all: tuple[int, ...] = () + flushed_two: tuple[int, ...] = () # None -> no cap, all 9 offsets stored r = make_runner() @@ -1081,32 +1069,6 @@ def test_offload_prompt_only(request_runner, async_scheduling: bool): assert len(offered_keys) == num_prompt_blocks -def test_flush_all_jobs_when_no_requests_remain(request_runner): - """When all tracked requests are finished, build_connector_meta flushes - all pending jobs since there will be no future step to complete them.""" - block_size = 4 - block_size_factor = 1 - offloaded_block_size = block_size * block_size_factor - - runner = request_runner( - block_size=block_size, - num_gpu_blocks=100, - async_scheduling=False, - block_size_factor=block_size_factor, - ) - - runner.new_request(token_ids=[0] * offloaded_block_size) - runner.manager.prepare_store.side_effect = lambda keys, req_context: ( - generate_store_output(keys) - ) - runner.run( - decoded_tokens=[EOS_TOKEN_ID], - complete_transfers=False, - expected_stored=(0,), - expected_flushed=(0,), - ) - - @pytest.mark.parametrize("async_scheduling", [True, False]) def test_reset_cache(request_runner, async_scheduling: bool): """reset_cache flushes in-flight loads, calls manager.reset_cache(), resets @@ -1131,7 +1093,6 @@ def test_reset_cache(request_runner, async_scheduling: bool): runner.run( decoded_tokens=[EOS_TOKEN_ID], expected_stored=(0, 1, 2), - expected_flushed=(0, 1, 2) if not async_scheduling else (), ) # Reset GPU prefix cache then start a request that loads from CPU. @@ -1381,7 +1342,6 @@ def test_stale_sliding_window_block_after_prepare_store_failure( runner.run( decoded_tokens=[EOS_TOKEN_ID], expected_stored=(2, 3), - expected_flushed=(2, 3) if not async_scheduling else (), ) @@ -1409,7 +1369,6 @@ def test_skip_reading_prefix_cache(request_runner, async_scheduling: bool): runner.run( decoded_tokens=[EOS_TOKEN_ID], expected_stored=(0, 1, 2), - expected_flushed=(0, 1, 2) if not async_scheduling else (), ) # Reset GPU prefix cache so the next request cannot hit locally. @@ -1429,7 +1388,6 @@ def test_skip_reading_prefix_cache(request_runner, async_scheduling: bool): decoded_tokens=[EOS_TOKEN_ID], expected_loaded=(), # no CPU loads must happen expected_stored=(0, 1, 2), # tokens still offloaded to CPU - expected_flushed=(0, 1, 2) if not async_scheduling else (), ) # The external lookup must have been completely skipped. @@ -1936,15 +1894,6 @@ class TestEagle: (1, 0), (1, 1), ), - expected_flushed=( - (0, 0), - (0, 1), - (0, 2), - (1, 0), - (1, 1), - ) - if not async_scheduling - else (), ) @pytest.mark.parametrize("async_scheduling", [True, False]) @@ -1990,7 +1939,6 @@ class TestEagle: runner.run( decoded_tokens=[EOS_TOKEN_ID], expected_stored=((0, 0), (0, 1)), - expected_flushed=((0, 0), (0, 1)) if not async_scheduling else (), ) @pytest.mark.parametrize("async_scheduling", [True, False]) @@ -2087,15 +2035,6 @@ class TestEagle: (1, 0), (1, 1), ), - expected_flushed=( - (0, 0), - (0, 1), - (0, 2), - (1, 0), - (1, 1), - ) - if not async_scheduling - else (), ) runner.scheduler.reset_prefix_cache() diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py index 443d5b28d54..9c3cb7e5a5d 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py @@ -1016,14 +1016,6 @@ class OffloadingConnectorScheduler: for jid in self._block_id_to_pending_jobs[bid] ) - # If all tracked requests are finished, flush all pending jobs - # (both store and load) - there might not be a future scheduler - # step to trigger their completion. - if self._req_status and all( - rs.req.is_finished() for rs in self._req_status.values() - ): - self._current_batch_jobs_to_flush.update(self._jobs.keys()) - meta = OffloadingConnectorMetadata( load_jobs=self._current_batch_load_jobs, store_jobs=self._build_store_jobs(scheduler_output), @@ -1034,6 +1026,14 @@ class OffloadingConnectorScheduler: self._current_batch_allocated_block_ids = set() return meta + def has_pending_push_work(self) -> bool: + """Whether the engine must keep stepping. + + While True, build_connector_meta() and update_connector_output() + continue to be called even when no requests are scheduled. + """ + return bool(self._jobs) or self.manager.has_pending_work() + def update_connector_output(self, connector_output: KVConnectorOutput): """ Update KVConnector state from worker-side connectors output. diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py index 7d567ed4622..1c5986d5156 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py @@ -150,6 +150,10 @@ class OffloadingConnector(KVConnectorBase_V1, SupportsHMA): assert self.connector_scheduler is not None return self.connector_scheduler.build_connector_meta(scheduler_output) + def has_pending_push_work(self) -> bool: + assert self.connector_scheduler is not None + return self.connector_scheduler.has_pending_push_work() + def update_connector_output(self, connector_output: KVConnectorOutput): assert self.connector_scheduler is not None self.connector_scheduler.update_connector_output(connector_output) diff --git a/vllm/v1/kv_offload/base.py b/vllm/v1/kv_offload/base.py index 16f783190f9..15781bbc8a7 100644 --- a/vllm/v1/kv_offload/base.py +++ b/vllm/v1/kv_offload/base.py @@ -288,6 +288,14 @@ class OffloadingManager(ABC): """ return + def has_pending_work(self) -> bool: + """Whether this manager needs the engine to keep stepping. + + While True, on_schedule_end() and get_finished_jobs() continue + to be called even when no requests are scheduled. + """ + return False + def reset_cache(self) -> None: """Evict all tracked blocks and reset internal state.""" return diff --git a/vllm/v1/kv_offload/tiering/base.py b/vllm/v1/kv_offload/tiering/base.py index dd9178fc7c7..f9fbdf9495a 100644 --- a/vllm/v1/kv_offload/tiering/base.py +++ b/vllm/v1/kv_offload/tiering/base.py @@ -153,6 +153,14 @@ class SecondaryTierManager(ABC): """ pass + def has_pending_work(self) -> bool: + """Whether this tier needs the engine to keep stepping. + + While True, on_schedule_end() and get_finished_jobs() continue + to be called even when no requests are scheduled. + """ + return False + def touch(self, keys: Collection[OffloadKey], req_context: ReqContext): """ Mark blocks as recently used for eviction policy. diff --git a/vllm/v1/kv_offload/tiering/manager.py b/vllm/v1/kv_offload/tiering/manager.py index fbcccea1626..d13e1f1eea5 100644 --- a/vllm/v1/kv_offload/tiering/manager.py +++ b/vllm/v1/kv_offload/tiering/manager.py @@ -577,6 +577,15 @@ class TieringOffloadingManager(OffloadingManager): for tier in self.secondary_tiers: tier.on_schedule_end() + @override + def has_pending_work(self) -> bool: + # In-flight primary<->secondary transfers (pending promotions are + # translated to transfer jobs in on_schedule_end), plus any work the + # secondary tiers themselves still have outstanding. + return bool(self._transfer_jobs) or any( + tier.has_pending_work() for tier in self.secondary_tiers + ) + @override def take_events(self) -> Iterable[OffloadingEvent]: """Yield offloading events collected since the last call. From 68ff30d40e9f0ef1eaff33f99a155f719ccb7bb9 Mon Sep 17 00:00:00 2001 From: Joel Smith Date: Wed, 17 Jun 2026 09:35:27 +0100 Subject: [PATCH 12/37] [Bugfix] Fixes MiniCPM-O resampler device placement to avoid tensor device mismatch (#42332) Signed-off-by: j9smith --- vllm/model_executor/models/minicpmo.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/vllm/model_executor/models/minicpmo.py b/vllm/model_executor/models/minicpmo.py index a8786f677ba..bd8547420c6 100644 --- a/vllm/model_executor/models/minicpmo.py +++ b/vllm/model_executor/models/minicpmo.py @@ -719,7 +719,9 @@ class MiniCPMOBaseModel: def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self, skip_prefixes=["tts"]) - return loader.load_weights(weights) + loaded = loader.load_weights(weights) + self._ensure_resampler_device() + return loaded def subsequent_chunk_mask( self, From 3c6084bb0d5168be132df395d1859098807f0fec Mon Sep 17 00:00:00 2001 From: Nikhilesh Chhetri <106703537+nikhilesh-csa@users.noreply.github.com> Date: Wed, 17 Jun 2026 05:16:02 -0500 Subject: [PATCH 13/37] [Bugfix][Gemma4] Pre-initialise streaming reasoning state when prompt ends inside an open `<|channel>` (fixes #45834) (#45852) Signed-off-by: nikhilesh-csa --- .../engine/test_gemma4_streaming_reasoning.py | 208 ++++++++++++++++++ vllm/parser/abstract_parser.py | 7 + vllm/parser/engine/adapters.py | 3 + vllm/parser/engine/parser_engine.py | 12 + vllm/parser/gemma4.py | 28 +++ vllm/reasoning/abs_reasoning_parsers.py | 12 + 6 files changed, 270 insertions(+) diff --git a/tests/parser/engine/test_gemma4_streaming_reasoning.py b/tests/parser/engine/test_gemma4_streaming_reasoning.py index 05e2388ec2b..19e424701a9 100644 --- a/tests/parser/engine/test_gemma4_streaming_reasoning.py +++ b/tests/parser/engine/test_gemma4_streaming_reasoning.py @@ -25,12 +25,14 @@ CHANNEL_END_ID = 51 # TOOL_CALL_START_ID = 48 # <|tool_call> TOOL_CALL_END_ID = 49 # QUOTED_ID = 52 # <|"|> +NEW_TURN_ID = 53 # <|turn> SPECIAL_TOKEN_MAP = { CHANNEL_START_ID: "<|channel>", CHANNEL_END_ID: "", TOOL_CALL_START_ID: "<|tool_call>", TOOL_CALL_END_ID: "", QUOTED_ID: '<|"|>', + NEW_TURN_ID: "<|turn>", } SPECIAL_TEXT_TO_ID = {v: k for k, v in SPECIAL_TOKEN_MAP.items()} @@ -253,6 +255,212 @@ class TestGemma4StreamingReasoningThenToolCall: ) +# ── Prompt ends inside an open <|channel>thought\n block ───────────── + +_OPEN_REASONING_GEN_SEQUENCE: list[tuple[int, str]] = [ + (7001, "Sure"), + (7002, ","), + (7003, " the"), + (7004, " answer"), + (7005, " is"), + (7006, " 42"), + (CHANNEL_END_ID, ""), + (7007, "Hello"), + (7008, " world"), +] + + +class TestGemma4PromptOpenReasoning: + """When ``add_generation_prompt=True`` after a final tool response with + ``enable_thinking=True``, the Gemma4 chat template leaves the prompt + ending with ``<|channel>thought\\n`` — i.e. inside an open reasoning + channel. Tokens generated before ```` must be classified as + ``reasoning``, not visible ``content``. + + Regression test for vllm-project/vllm#45834. + """ + + @pytest.fixture + def open_reasoning_tokenizer(self): + return _make_tokenizer(_OPEN_REASONING_GEN_SEQUENCE) + + @pytest.fixture + def open_reasoning_parser(self, open_reasoning_tokenizer): + return Gemma4Parser(open_reasoning_tokenizer) + + @staticmethod + def _prompt_ids_open_channel() -> list[int]: + # Mimics a prompt that ends with ``...<|channel>thought\n``. The + # specific token ids for ``thought`` and ``\n`` are arbitrary — only + # the trailing ``<|channel>`` start token matters for detection. + return [CHANNEL_START_ID, 3000, 3001] + + def test_reasoning_not_leaked_into_content( + self, open_reasoning_parser, open_reasoning_tokenizer, request_obj + ): + results = _stream_tokens_batched( + open_reasoning_parser, + open_reasoning_tokenizer, + request_obj, + batch_size=1, + prompt_token_ids=self._prompt_ids_open_channel(), + ) + + reasoning, content, _ = _collect_fields(results) + + assert "Sure, the answer is 42" in reasoning, ( + f"Expected pre- tokens in reasoning, got " + f"reasoning={reasoning!r} content={content!r}" + ) + for leaked in ("Sure", "answer", "42"): + assert leaked not in content, ( + f"Reasoning text leaked into content: {content!r}" + ) + + def test_post_reasoning_text_in_content( + self, open_reasoning_parser, open_reasoning_tokenizer, request_obj + ): + results = _stream_tokens_batched( + open_reasoning_parser, + open_reasoning_tokenizer, + request_obj, + batch_size=1, + prompt_token_ids=self._prompt_ids_open_channel(), + ) + + _, content, _ = _collect_fields(results) + + assert "Hello world" in content, ( + f"Post- text missing from content: {content!r}" + ) + + def test_new_turn_prompt_unchanged(self, parser, mock_tokenizer, request_obj): + """When the prompt does NOT end in an open reasoning channel (e.g. a + new turn that ends with ``<|turn>model\\n``), behaviour must match + the existing flow — the model itself opens ``<|channel>``. + """ + results = _stream_tokens_batched( + parser, + mock_tokenizer, + request_obj, + batch_size=10, + # No <|channel> in the prompt tail. + prompt_token_ids=[9000, 9001], + ) + + reasoning, content, tool_calls = _collect_fields(results) + + assert "weather" in reasoning.lower(), ( + f"Expected reasoning about weather, got: {reasoning[:100]!r}" + ) + assert len(tool_calls) > 0, f"Tool calls missing — content={content!r}" + + +# ── Engine pre-initialised to REASONING + model still emits channel open ── + +_PRE_INIT_THOUGHT_GEN_SEQUENCE: list[tuple[int, str]] = [ + # Model naively emits the full reasoning opener even though the engine + # was pre-initialised to REASONING from the prompt. + (CHANNEL_START_ID, "<|channel>"), + (8000, "thought"), + (8001, "\n"), + (8002, "Reason"), + (8003, "ing"), + (8004, " body"), + (CHANNEL_END_ID, ""), + (8005, "Final"), + (8006, " content"), +] + + +class TestGemma4PreInitReasoningRobustness: + """Tests for the ``(REASONING, THINK_START)`` no-op transition and + cooperating ``thought\\n`` prefix stripping when the engine has been + pre-initialised to ``REASONING`` from the prompt. + + These cover the case the reviewer raised: prompt ends with + ``<|turn>model\\n`` (``is_reasoning_end`` returns ``False`` because + thinking is enabled, so the engine is pre-initialised), but the model + still emits its own ``<|channel>thought\\n…content``. The + ``thought\\n`` prefix must be stripped, the ``<|channel>`` must not + leak as text, and the post-```` text must appear as content. + """ + + @pytest.fixture + def pre_init_tokenizer(self): + return _make_tokenizer(_PRE_INIT_THOUGHT_GEN_SEQUENCE) + + @pytest.fixture + def pre_init_parser(self, pre_init_tokenizer): + return Gemma4Parser(pre_init_tokenizer) + + def test_redundant_channel_open_swallowed_after_new_turn( + self, pre_init_parser, pre_init_tokenizer, request_obj + ): + # Prompt ends with ``<|turn>model\n``-style sentinel. With + # ``enable_thinking=True`` (the default), ``is_reasoning_end`` + # returns ``False`` for a ``<|turn>`` tail, so the engine is + # pre-initialised to ``REASONING``. + results = _stream_tokens_batched( + pre_init_parser, + pre_init_tokenizer, + request_obj, + batch_size=1, + prompt_token_ids=[NEW_TURN_ID, 9100, 9101], + ) + + reasoning, content, _ = _collect_fields(results) + + # ``thought\n`` prefix must be stripped from reasoning even though + # the engine was pre-initialised to REASONING. + assert reasoning.startswith("Reason"), ( + f"thought\\n prefix leaked into reasoning: {reasoning!r}" + ) + assert "thought\n" not in reasoning, ( + f"thought\\n prefix leaked into reasoning: {reasoning!r}" + ) + assert "Reasoning body" in reasoning, f"Reasoning body missing: {reasoning!r}" + + # The redundant ``<|channel>`` opener must not appear as text. + assert "<|channel>" not in content, ( + f"<|channel> leaked into content: {content!r}" + ) + assert "<|channel>" not in reasoning, ( + f"<|channel> leaked into reasoning: {reasoning!r}" + ) + + # Post-```` text must appear as content. + assert "Final content" in content, ( + f"Post- text missing from content: {content!r}" + ) + + def test_redundant_channel_open_swallowed_after_open_channel_prompt( + self, pre_init_parser, pre_init_tokenizer, request_obj + ): + # Prompt already ends inside an open ``<|channel>`` block. Engine + # is pre-initialised to ``REASONING`` via the start-token check. + # Even if the model redundantly re-emits ``<|channel>thought\n``, + # the no-op transition + prefix stripping must keep output clean. + results = _stream_tokens_batched( + pre_init_parser, + pre_init_tokenizer, + request_obj, + batch_size=1, + prompt_token_ids=[CHANNEL_START_ID, 3000, 3001], + ) + + reasoning, content, _ = _collect_fields(results) + + assert "<|channel>" not in content, ( + f"<|channel> leaked into content: {content!r}" + ) + assert "thought\n" not in reasoning, ( + f"thought\\n prefix leaked into reasoning: {reasoning!r}" + ) + assert "Reasoning body" in reasoning + assert "Final content" in content + + # ── Second model output: two tool calls with holdback ──────────────── REASONING_TEXT_2 = ( diff --git a/vllm/parser/abstract_parser.py b/vllm/parser/abstract_parser.py index cf7dc2ec1fc..915d401f7bd 100644 --- a/vllm/parser/abstract_parser.py +++ b/vllm/parser/abstract_parser.py @@ -758,6 +758,13 @@ class DelegatingParser(Parser): prompt_token_ids ): state.reasoning_ended = True + else: + # Reasoning is still open at the end of the prompt; let the + # reasoning parser adjust its initial parsing state so the + # first generated tokens are classified correctly. + self._reasoning_parser.adjust_initial_state_from_prompt( + prompt_token_ids + ) current_text, current_token_ids = state.advance(delta_text, delta_token_ids) delta_message: DeltaMessage | None = None diff --git a/vllm/parser/engine/adapters.py b/vllm/parser/engine/adapters.py index 3efa918d1a4..2482dad437b 100644 --- a/vllm/parser/engine/adapters.py +++ b/vllm/parser/engine/adapters.py @@ -60,6 +60,9 @@ class ParserEngineReasoningAdapter(ReasoningParser): def is_reasoning_end(self, input_ids: Sequence[int]) -> bool: return self._parser_engine.is_reasoning_end(list(input_ids)) + def adjust_initial_state_from_prompt(self, prompt_token_ids: Sequence[int]) -> None: + self._parser_engine.adjust_initial_state_from_prompt(prompt_token_ids) + def extract_content_ids(self, input_ids: list[int]) -> list[int]: return self._parser_engine.extract_content_ids(input_ids) diff --git a/vllm/parser/engine/parser_engine.py b/vllm/parser/engine/parser_engine.py index 5aa853c6f69..237e2745632 100644 --- a/vllm/parser/engine/parser_engine.py +++ b/vllm/parser/engine/parser_engine.py @@ -101,6 +101,7 @@ class ParserEngine(Parser): self._reasoning_ended: bool = False self._streaming_initialized: bool = False + self._prompt_streaming_prepared: bool = False self._tool_slots: list[ToolCallSlot] = [] self._deferred_content: str = "" @@ -165,6 +166,10 @@ class ParserEngine(Parser): self._streaming_initialized = True self._reset(initial_state=initial_state) + def adjust_initial_state_from_prompt(self, prompt_token_ids: Sequence[int]) -> None: + """See :meth:`ReasoningParser.adjust_initial_state_from_prompt`.""" + return + def finish_streaming(self) -> DeltaMessage | None: events = self._engine.finish() return self._events_to_delta(events) if events else None @@ -176,6 +181,7 @@ class ParserEngine(Parser): self._deferred_content = "" self._deferred_reasoning = "" self._content_has_nonws = False + self._prompt_streaming_prepared = False def adjust_request( self, request: ChatCompletionRequest | ResponsesRequest @@ -364,6 +370,12 @@ class ParserEngine(Parser): *, finished: bool, ) -> DeltaMessage | None: + if not self._prompt_streaming_prepared and prompt_token_ids is not None: + # NOTE: call the hook BEFORE setting the flag, because the hook + # may invoke ``_reset`` (e.g. via ``initialize_streaming``) which + # clears ``_prompt_streaming_prepared``. + self.adjust_initial_state_from_prompt(prompt_token_ids) + self._prompt_streaming_prepared = True self._check_skip_tool_parsing(request) events = self._feed(delta_text, delta_token_ids) if finished: diff --git a/vllm/parser/gemma4.py b/vllm/parser/gemma4.py index cf0633c7ab0..5dd07e44e3e 100644 --- a/vllm/parser/gemma4.py +++ b/vllm/parser/gemma4.py @@ -353,6 +353,14 @@ def gemma4_config() -> ParserEngineConfig: ParserState.REASONING, (EventType.REASONING_START,), ), + # No-op: if we pre-initialised the engine to REASONING from the + # prompt (see ``adjust_initial_state_from_prompt``) but the model + # still emits its own ``<|channel>`` opener, swallow it instead + # of leaking it as TEXT_CHUNK. + (ParserState.REASONING, "THINK_START"): Transition( + ParserState.REASONING, + (), + ), (ParserState.REASONING, "THINK_END"): Transition( ParserState.CONTENT, (EventType.REASONING_END,), @@ -524,6 +532,26 @@ class Gemma4Parser(ParserEngine): return True return True + def adjust_initial_state_from_prompt(self, prompt_token_ids: Sequence[int]) -> None: + """Pre-initialise the engine to ``REASONING`` when the prompt does + not already end with reasoning concluded. + + This covers the post-tool-response continuation case where the chat + template leaves the prompt ending inside an open ``<|channel>`` + block (issue #45834). It is also safe in the common new-turn case + where the model itself emits ``<|channel>`` first: the no-op + ``(REASONING, THINK_START)`` transition swallows it, and the + ``thought\n`` prefix in the first reasoning chunk is stripped by + ``_events_to_delta`` as it already is in the default flow. + """ + if self.is_reasoning_end(list(prompt_token_ids)): + return + self._engine.reset(initial_state=ParserState.REASONING) + # Prevent a later default ``initialize_streaming()`` (e.g. from + # ``ParserEngineReasoningAdapter.extract_reasoning_streaming``) from + # clobbering this with ``CONTENT``. + self._streaming_initialized = True + def _events_to_delta( self, events: list[SemanticEvent], diff --git a/vllm/reasoning/abs_reasoning_parsers.py b/vllm/reasoning/abs_reasoning_parsers.py index 4e519f6aeb6..4e28e50702e 100644 --- a/vllm/reasoning/abs_reasoning_parsers.py +++ b/vllm/reasoning/abs_reasoning_parsers.py @@ -187,6 +187,18 @@ class ReasoningParser: """Adjust request parameters; override in subclasses as needed.""" return request + def adjust_initial_state_from_prompt(self, prompt_token_ids: Sequence[int]) -> None: + """Hook called once at the start of streaming with the prompt tokens. + + Gives parsers a chance to adjust their initial parsing state based on + the prompt — for example, when the chat template leaves the prompt + inside an open reasoning channel and the engine's default initial + state would otherwise misclassify the first generated tokens. + + Default is a no-op; override in subclasses as needed. + """ + return + def prepare_structured_tag( self, original_tag: str | None, From f694d43b33503f18dddd68857688c6826890a847 Mon Sep 17 00:00:00 2001 From: wentian-byte <3400259131@qq.com> Date: Wed, 17 Jun 2026 18:37:19 +0800 Subject: [PATCH 14/37] [Bugfix][test] Use Salesforce/wikitext for ppl tests (#45913) Co-authored-by: wentian-byte <192079369+wentian-byte@users.noreply.github.com> --- tests/models/language/generation_ppl_test/ppl_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/models/language/generation_ppl_test/ppl_utils.py b/tests/models/language/generation_ppl_test/ppl_utils.py index 59740505e82..2b5449bddcb 100644 --- a/tests/models/language/generation_ppl_test/ppl_utils.py +++ b/tests/models/language/generation_ppl_test/ppl_utils.py @@ -30,7 +30,7 @@ def wikitext_ppl_test( ): vllm_extra_kwargs = get_vllm_extra_kwargs(model_info, vllm_extra_kwargs) - dataset = load_dataset("wikitext", "wikitext-2-raw-v1", split="test") + dataset = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split="test") with vllm_runner( model_info.name, From 3d20275bb4d434f53055c3c0b645fd8bb072965e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Juan=20P=C3=A9rez=20de=20Algaba?= <124347725+jperezdealgaba@users.noreply.github.com> Date: Wed, 17 Jun 2026 13:07:13 +0200 Subject: [PATCH 15/37] fix(security): enforce audio decode duration limit in chat completions path (#45908) Signed-off-by: jperezde --- vllm/multimodal/media/audio.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/vllm/multimodal/media/audio.py b/vllm/multimodal/media/audio.py index 5e998be3fcb..c9e5753ca78 100644 --- a/vllm/multimodal/media/audio.py +++ b/vllm/multimodal/media/audio.py @@ -9,6 +9,7 @@ import numpy.typing as npt import pybase64 import torch +import vllm.envs as envs from vllm.logger import init_logger from vllm.multimodal.audio import resample_audio_pyav from vllm.utils.import_utils import PlaceholderModule @@ -235,7 +236,11 @@ class AudioMediaIO(MediaIO[tuple[npt.NDArray, float]]): self.kwargs = kwargs def load_bytes(self, data: bytes) -> tuple[npt.NDArray, float]: - return load_audio(BytesIO(data), sr=None) + return load_audio( + BytesIO(data), + sr=None, + max_duration_s=envs.VLLM_MAX_AUDIO_DECODE_DURATION_S, + ) def load_base64( self, @@ -245,7 +250,11 @@ class AudioMediaIO(MediaIO[tuple[npt.NDArray, float]]): return self.load_bytes(pybase64.b64decode(data)) def load_file(self, filepath: Path) -> tuple[npt.NDArray, float]: - return load_audio(filepath, sr=None) + return load_audio( + filepath, + sr=None, + max_duration_s=envs.VLLM_MAX_AUDIO_DECODE_DURATION_S, + ) def encode_base64( self, From d537122398df9f8720c307c8b2fa2a8d66f0ff78 Mon Sep 17 00:00:00 2001 From: vllmellm Date: Wed, 17 Jun 2026 19:41:29 +0800 Subject: [PATCH 16/37] [ROCm][Bugfix]: Fallback GFX942 sparse MLA ops to Triton (#45782) Signed-off-by: vllmellm --- .../v1/attention/ops/rocm_aiter_mla_sparse.py | 49 ++++++------------- 1 file changed, 15 insertions(+), 34 deletions(-) diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index c38a4780f78..51513a5a9f4 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -639,7 +639,6 @@ def rocm_aiter_sparse_attn_indexer( # careful! this will be None in dummy run attn_metadata = get_forward_context().attn_metadata fp8_dtype = current_platform.fp8_dtype() - from vllm import _custom_ops as ops from vllm.utils.torch_utils import _resolve_layer_name k_cache_prefix = _resolve_layer_name(k_cache_prefix) @@ -716,22 +715,13 @@ def rocm_aiter_sparse_attn_indexer( raise ValueError("k must be provided when skip_k_cache_insert is False") if not skip_k_cache_insert: - if _ON_GFX942: - ops.indexer_k_quant_and_cache( - k, - kv_cache, - slot_mapping, - quant_block_size, - scale_fmt, - ) - else: - indexer_k_quant_and_cache_triton( - k, - kv_cache, - slot_mapping, - quant_block_size, - scale_fmt, - ) + indexer_k_quant_and_cache_triton( + k, + kv_cache, + slot_mapping, + quant_block_size, + scale_fmt, + ) topk_indices_buffer[: hidden_states.shape[0]] = -1 if has_prefill: @@ -746,23 +736,14 @@ def rocm_aiter_sparse_attn_indexer( for chunk in prefill_metadata.chunks: k_fp8 = k_fp8_full[: chunk.total_seq_lens] k_scale = k_scale_full[: chunk.total_seq_lens] - if _ON_GFX942: - ops.cp_gather_indexer_k_quant_cache( - kv_cache, - k_fp8, - k_scale, - chunk.block_table, - chunk.cu_seq_lens, - ) - else: - cp_gather_indexer_k_quant_cache_triton( - kv_cache, - k_fp8, - k_scale, - chunk.block_table, - chunk.cu_seq_lens, - token_to_seq=chunk.token_to_seq, - ) + cp_gather_indexer_k_quant_cache_triton( + kv_cache, + k_fp8, + k_scale, + chunk.block_table, + chunk.cu_seq_lens, + token_to_seq=chunk.token_to_seq, + ) logits = rocm_fp8_mqa_logits( q_fp8[chunk.token_start : chunk.token_end], (k_fp8, k_scale.view(torch.float32)), From ee0fd6984ad6579204a388cd80021fea081c8869 Mon Sep 17 00:00:00 2001 From: Angelo Ruocco Date: Wed, 17 Jun 2026 13:58:00 +0200 Subject: [PATCH 17/37] docs, kv_offloading: add docs for selective offload (#45279) Signed-off-by: Angelo Ruocco --- docs/features/kv_offloading_usage.md | 23 +++++++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/docs/features/kv_offloading_usage.md b/docs/features/kv_offloading_usage.md index 39b826bfd56..93da2ed0361 100644 --- a/docs/features/kv_offloading_usage.md +++ b/docs/features/kv_offloading_usage.md @@ -127,6 +127,29 @@ PYTHONHASHSEED=0 vllm serve ... - FS thread counts: tune `n_read_threads` and `n_write_threads` to the parallelism your storage can sustain. Reads are latency-sensitive on the prefill path, so prefer more read threads when prefill hit rates are high. - Sharing `root_dir` across runs: runs with the same model, `block_size`, parallelism layout, and dtype share files under the same `` subdirectory. Changing any of these produces a new subdirectory; old ones are orphaned but harmless. Delete them to reclaim disk. +## Per-Request Selective Offload + +Individual requests can cap how many of their tokens are eligible for offload by setting `max_offload_tokens` in the request's `kv_transfer_params`. Only the first `max_offload_tokens` tokens of the request are offloaded; blocks beyond that point are skipped on the store path. This is useful when a known prefix (e.g., a system prompt or shared context) is worth caching but later request-specific tokens are not. + +| Key | Type | Notes | +| --- | --- | --- | +| `max_offload_tokens` | non-negative `int` | Upper bound on tokens to offload for this request. `0` disables offload for the request entirely; omit the key (or set to `None`) for no cap. Non-`int`, negative, or `bool` values are rejected with a warning and treated as no cap. | + +!!! note + `max_offload_tokens` is experimental and subject to change. + +Example (OpenAI-compatible completions request): + +```json +{ + "model": "", + "prompt": "...", + "kv_transfer_params": { + "max_offload_tokens": 1024 + } +} +``` + ## Further Reading - [vLLM blog: KV Offloading Connector](https://vllm.ai/blog/2026-01-08-kv-offloading-connector) — motivation, architecture (DMA-based async transfer), and benchmarks (TTFT and throughput). From e28e8c87820d620f3983ef6a6ef6fdd67deb7936 Mon Sep 17 00:00:00 2001 From: Hongxia Yang <62075498+hongxiayang@users.noreply.github.com> Date: Wed, 17 Jun 2026 08:02:40 -0400 Subject: [PATCH 18/37] [ROCm][Quant] Minimax-M3: Enable fp8_per_channel for bf16 weights on mi300x (#45854) Signed-off-by: Hongxia Yang Signed-off-by: tjtanaa Co-authored-by: tjtanaa --- vllm/model_executor/layers/fused_moe/config.py | 4 ++++ vllm/model_executor/layers/fused_moe/oracle/fp8.py | 2 ++ vllm/model_executor/layers/quantization/fp8.py | 2 ++ vllm/model_executor/layers/quantization/online/fp8.py | 2 ++ vllm/platforms/rocm.py | 1 + 5 files changed, 11 insertions(+) diff --git a/vllm/model_executor/layers/fused_moe/config.py b/vllm/model_executor/layers/fused_moe/config.py index 0755699d1a4..905a9bea3c5 100644 --- a/vllm/model_executor/layers/fused_moe/config.py +++ b/vllm/model_executor/layers/fused_moe/config.py @@ -603,6 +603,8 @@ def fp8_w8a8_moe_quant_config( a2_gscale: torch.Tensor | None = None, g1_alphas: torch.Tensor | None = None, g2_alphas: torch.Tensor | None = None, + gemm1_alpha: float | None = None, + gemm1_beta: float | None = None, gemm1_clamp_limit: float | None = None, ) -> FusedMoEQuantConfig: """ @@ -623,6 +625,8 @@ def fp8_w8a8_moe_quant_config( per_act_token_quant=per_act_token_quant, per_out_ch_quant=per_out_ch_quant, block_shape=block_shape, + gemm1_alpha=gemm1_alpha, + gemm1_beta=gemm1_beta, gemm1_clamp_limit=gemm1_clamp_limit, ) diff --git a/vllm/model_executor/layers/fused_moe/oracle/fp8.py b/vllm/model_executor/layers/fused_moe/oracle/fp8.py index acbf2cb46ad..1b5030b1909 100644 --- a/vllm/model_executor/layers/fused_moe/oracle/fp8.py +++ b/vllm/model_executor/layers/fused_moe/oracle/fp8.py @@ -568,6 +568,8 @@ def make_fp8_moe_quant_config( block_shape=block_shape, per_act_token_quant=per_act_token_quant, per_out_ch_quant=per_out_ch_quant, + gemm1_alpha=gemm1_alpha, + gemm1_beta=gemm1_beta, gemm1_clamp_limit=swiglu_limit, ) diff --git a/vllm/model_executor/layers/quantization/fp8.py b/vllm/model_executor/layers/quantization/fp8.py index 1c97e48771f..6fc9b0d8f65 100644 --- a/vllm/model_executor/layers/quantization/fp8.py +++ b/vllm/model_executor/layers/quantization/fp8.py @@ -787,6 +787,8 @@ class Fp8MoEMethod(FusedMoEMethodBase): a2_scale=a2_scale, block_shape=self.weight_block_size, swiglu_limit=getattr(layer, "swiglu_limit", None), + gemm1_alpha=getattr(layer, "swiglu_alpha", None), + gemm1_beta=getattr(layer, "swiglu_beta", None), ) # Inject biases into the quant config if the model has them diff --git a/vllm/model_executor/layers/quantization/online/fp8.py b/vllm/model_executor/layers/quantization/online/fp8.py index 10847851ccd..933fc7c9263 100644 --- a/vllm/model_executor/layers/quantization/online/fp8.py +++ b/vllm/model_executor/layers/quantization/online/fp8.py @@ -483,6 +483,8 @@ class _Fp8OnlineMoEBase(OnlineMoEMethodBase): per_act_token_quant=self.per_act_token_quant, per_out_ch_quant=self.per_out_ch_quant, swiglu_limit=getattr(layer, "swiglu_limit", None), + gemm1_alpha=getattr(layer, "swiglu_alpha", None), + gemm1_beta=getattr(layer, "swiglu_beta", None), ) diff --git a/vllm/platforms/rocm.py b/vllm/platforms/rocm.py index aaf1fdce36b..13695a142e8 100644 --- a/vllm/platforms/rocm.py +++ b/vllm/platforms/rocm.py @@ -459,6 +459,7 @@ class RocmPlatform(Platform): "modelopt_mixed", "fp8_per_tensor", "fp8_per_block", + "fp8_per_channel", "online", "gpt_oss_mxfp4", ] From fa85ead2f378621fc412ac3679eae88a5604ebe8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?O=C4=9Fuzhan=20KIR?= <86883236+oguzhankir@users.noreply.github.com> Date: Wed, 17 Jun 2026 15:14:01 +0300 Subject: [PATCH 19/37] [MM][Perf][CG] Support ViT full CUDA graph for Kimi-VL (#41992) Signed-off-by: oguz --- docs/design/cuda_graphs_multimodal.md | 1 + .../multimodal/vision_language_offline.py | 1 + .../generation/test_vit_cudagraph.py | 35 ++ vllm/model_executor/models/kimi_vl.py | 197 ++++++++++- vllm/model_executor/models/moonvit.py | 315 +++++++++++++++--- 5 files changed, 504 insertions(+), 45 deletions(-) diff --git a/docs/design/cuda_graphs_multimodal.md b/docs/design/cuda_graphs_multimodal.md index 379e5f16b52..1db82ffa688 100644 --- a/docs/design/cuda_graphs_multimodal.md +++ b/docs/design/cuda_graphs_multimodal.md @@ -129,6 +129,7 @@ Models opt-in to encoder CUDA Graphs by implementing the [SupportsEncoderCudaGra | `DeepseekOCRForCausalLM` | `DeepSeek-OCR` | ✅︎ | ❌︎ | ✅︎ | | `Glm4vForConditionalGeneration` | `GLM-4.1V, GLM-4.6V-Flash` | ✅︎ | ✅︎ | ❌︎ | | `InternVLChatModel` | `InternVL3.5`, `InternVL3`, `InternVL2.5`, `InternVL2` | ✅︎ | ✅︎ | ❌︎ | +| `KimiVLForConditionalGeneration` | `Kimi-VL` | ✅︎ | ❌︎ | ❌︎ | | `Llama4ForConditionalGeneration` | `Llama 4` | ✅︎ | ❌︎ | ❌︎ | | `Qwen2VLForConditionalGeneration` | `Qwen2-VL` | ✅︎ | ✅︎ | ❌︎ | | `Qwen2_5_VLForConditionalGeneration` | `Qwen2.5-VL` | ✅︎ | ✅︎ | ❌︎ | diff --git a/examples/generate/multimodal/vision_language_offline.py b/examples/generate/multimodal/vision_language_offline.py index 48521c52482..1b3741a3e42 100644 --- a/examples/generate/multimodal/vision_language_offline.py +++ b/examples/generate/multimodal/vision_language_offline.py @@ -2537,6 +2537,7 @@ MODELS_SUPPORT_VIT_CUDA_GRAPH = [ "qwen2_5_vl", "qwen3_vl", "qwen3_vl_moe", + "kimi_vl", "qwen3_5", "qwen3_5_moe", "internvl_chat", diff --git a/tests/models/multimodal/generation/test_vit_cudagraph.py b/tests/models/multimodal/generation/test_vit_cudagraph.py index 0496031988f..52b28ca8600 100644 --- a/tests/models/multimodal/generation/test_vit_cudagraph.py +++ b/tests/models/multimodal/generation/test_vit_cudagraph.py @@ -48,6 +48,13 @@ def internvl_chat_template(content: str) -> str: return f"<|im_start|>user\n{content}<|im_end|>\n<|im_start|>assistant\n" +def kimi_vl_chat_template(content: str) -> str: + return ( + f"<|im_user|>user<|im_middle|>{content}<|im_end|>" + "<|im_assistant|>assistant<|im_middle|>" + ) + + def step3_vl_chat_template(content: str) -> str: return ( "<|begin▁of▁sentence|> You are a helpful assistant.<|BOT|>user\n " @@ -100,6 +107,34 @@ MODEL_CONFIGS: dict[str, VitCudagraphTestConfig] = { needs_video_metadata=False, marks=[pytest.mark.core_model], ), + "kimi_vl": VitCudagraphTestConfig( + model="moonshotai/Kimi-VL-A3B-Instruct", + modalities=["image"], + image_prompt=kimi_vl_chat_template( + "<|media_start|>image<|media_content|><|media_pad|><|media_end|>" + "What is in this image?" + ), + needs_video_metadata=False, + # Single bucket sized to cover the test images' output tokens. + # The default auto-inferred range fans out into multiple power-of-2 + # buckets, each holding a full ViT capture pool. + compilation_config_overrides={ + "encoder_cudagraph_token_budgets": [1024], + }, + # Shrink to 1 text + 1 vision layer with random weights so the + # test runs on any CI GPU (incl. L4) and skips the multi-GiB + # weight download. The test only validates that encoder CG + # capture/replay functions correctly, not output quality. + vllm_runner_kwargs={ + "trust_remote_code": True, + "load_format": "dummy", + "hf_overrides": partial( + dummy_hf_overrides, + model_arch="KimiVLForConditionalGeneration", + ), + }, + marks=[pytest.mark.core_model], + ), "qwen3_vl": VitCudagraphTestConfig( model="Qwen/Qwen3-VL-2B-Instruct", image_prompt=qwen_vl_chat_template( diff --git a/vllm/model_executor/models/kimi_vl.py b/vllm/model_executor/models/kimi_vl.py index e3bc08c654d..2b08fc6c1fd 100644 --- a/vllm/model_executor/models/kimi_vl.py +++ b/vllm/model_executor/models/kimi_vl.py @@ -56,7 +56,11 @@ from vllm.config import VllmConfig from vllm.config.multimodal import BaseDummyOptions from vllm.inputs import MultiModalDataDict from vllm.model_executor.layers.linear import ReplicatedLinear -from vllm.model_executor.models.interfaces import SupportsMultiModal, SupportsPP +from vllm.model_executor.models.interfaces import ( + SupportsEncoderCudaGraph, + SupportsMultiModal, + SupportsPP, +) from vllm.model_executor.models.moonvit import MoonVitPretrainedModel from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import ( @@ -79,6 +83,7 @@ from vllm.multimodal.processing import ( from vllm.sequence import IntermediateTensors from vllm.transformers_utils.configs.kimi_vl import KimiVLConfig, MoonViTConfig from vllm.utils.tensor_schema import TensorSchema, TensorShape +from vllm.v1.worker.encoder_cudagraph_defs import EncoderCudaGraphReplayBuffers from .utils import AutoWeightsLoader, init_vllm_registered_model, maybe_prefix from .vision import is_vit_use_data_parallel, run_dp_sharded_mrope_vision_model @@ -287,7 +292,9 @@ class KimiVLMultiModalProcessor(BaseMultiModalProcessor[KimiVLProcessingInfo]): info=KimiVLProcessingInfo, dummy_inputs=KimiVLDummyInputsBuilder, ) -class KimiVLForConditionalGeneration(nn.Module, SupportsMultiModal, SupportsPP): +class KimiVLForConditionalGeneration( + nn.Module, SupportsMultiModal, SupportsEncoderCudaGraph, SupportsPP +): supports_encoder_tp_data = True @classmethod @@ -340,6 +347,192 @@ class KimiVLForConditionalGeneration(nn.Module, SupportsMultiModal, SupportsPP): self.media_placeholder: int = self.config.media_placeholder_token_id + self.model_config = model_config + + # -- SupportsEncoderCudaGraph protocol methods -- + + def get_encoder_cudagraph_config(self): + from vllm.v1.worker.encoder_cudagraph_defs import ( + EncoderCudaGraphConfig, + ) + + return EncoderCudaGraphConfig( + modalities=["image"], + buffer_keys=[ + "pixel_values", + "pos_embeds", + "rope_freqs_cis", + "cu_seqlens", + "max_seqlen", + "merge_gather_idx", + ], + out_hidden_size=self.hidden_size, + ) + + def get_encoder_cudagraph_budget_range( + self, + vllm_config, + ) -> tuple[int, int]: + # Min: estimated smallest possible encoder input. + # 224x224 image with patch_size=14 -> 16x16 patches, then merge + # kernel (2,2) -> 8x8 = 64 output tokens. + min_budget = 64 + max_budget = min( + vllm_config.scheduler_config.max_num_batched_tokens, + self.model_config.max_model_len, + ) + return (min_budget, max_budget) + + def _get_grid_hws( + self, + mm_kwargs: dict[str, Any], + ) -> list[tuple[int, int]]: + grid_hws = mm_kwargs["image_grid_hws"] + if not isinstance(grid_hws, list): + grid_hws = grid_hws.tolist() + return grid_hws + + def get_encoder_cudagraph_item_specs( + self, + mm_kwargs: dict[str, Any], + ): + from vllm.v1.worker.encoder_cudagraph_defs import EncoderItemSpec + + kh, kw = self.config.vision_config.merge_kernel_size + return [ + EncoderItemSpec( + input_size=h * w, + output_tokens=(h // kh) * (w // kw), + ) + for h, w in self._get_grid_hws(mm_kwargs) + ] + + def select_encoder_cudagraph_items( + self, + mm_kwargs: dict[str, Any], + indices: list[int], + ) -> dict[str, Any]: + grid_hws = self._get_grid_hws(mm_kwargs) + pixel_values = mm_kwargs["pixel_values"] + + if len(indices) == 0: + return { + "pixel_values": pixel_values[:0], + "image_grid_hws": pixel_values.new_zeros((0, 2), dtype=torch.long), + } + + patches_per_item = [h * w for h, w in grid_hws] + cum_patches = [0] + for p in patches_per_item: + cum_patches.append(cum_patches[-1] + p) + + selected_pv = torch.cat( + [pixel_values[cum_patches[i] : cum_patches[i + 1]] for i in indices] + ) + selected_grid = torch.tensor( + [grid_hws[i] for i in indices], + dtype=torch.long, + device=pixel_values.device, + ) + return { + "pixel_values": selected_pv, + "image_grid_hws": selected_grid, + } + + def prepare_encoder_cudagraph_capture_inputs( + self, + token_budget: int, + max_batch_size: int, + max_frames_per_batch: int, + device: torch.device, + dtype: torch.dtype, + path: str = "default", + ): + from vllm.v1.worker.encoder_cudagraph_defs import ( + EncoderCudaGraphCaptureInputs, + ) + + kh, kw = self.config.vision_config.merge_kernel_size + # Ceil so the buffer fits the worst case of one item using the full + # budget. Floor under-allocates when budget is not a multiple of + # max_batch_size. + per_mm_item_output = (token_budget + max_batch_size - 1) // max_batch_size + + # Shape the synthetic grid so neither dimension exceeds Rope2DPosEmb's + # precomputed range. Pack as wide a row as fits, then add rows. + rope = self.vision_tower.encoder.rope_2d + max_wo = rope.max_width // kw + wo = min(per_mm_item_output, max_wo) + ho = (per_mm_item_output + wo - 1) // wo + assert ho * kh <= rope.max_height, ( + f"per_mm_item_output={per_mm_item_output} exceeds RoPE grid capacity " + f"(max {(rope.max_height // kh) * (rope.max_width // kw)} tokens)" + ) + grid_hws_list = [(ho * kh, wo * kw) for _ in range(max_batch_size)] + + patch_size = self.config.vision_config.patch_size + if isinstance(patch_size, int): + patch_size = (patch_size, patch_size) + + total_patches = sum(h * w for h, w in grid_hws_list) + in_channels = 3 + dummy_pixel_values = torch.randn( + total_patches, + in_channels, + patch_size[0], + patch_size[1], + device=device, + dtype=dtype, + ) + + buffers = self.vision_tower.prepare_encoder_metadata( + grid_hws_list, + max_batch_size=max_batch_size, + max_seqlen_override=token_budget, + device=device, + ) + values = buffers | {"pixel_values": dummy_pixel_values} + + return EncoderCudaGraphCaptureInputs(values=values) + + def prepare_encoder_cudagraph_replay_buffers( + self, + mm_kwargs: dict[str, Any], + max_batch_size: int, + max_frames_per_batch: int, + path: str = "default", + ): + grid_hws_list = self._get_grid_hws(mm_kwargs) + buffers = self.vision_tower.prepare_encoder_metadata( + grid_hws_list, + max_batch_size=max_batch_size, + device=mm_kwargs["pixel_values"].device, + ) + values = buffers | {"pixel_values": mm_kwargs["pixel_values"]} + return EncoderCudaGraphReplayBuffers(values=values) + + def encoder_cudagraph_forward( + self, + values: dict[str, torch.Tensor], + path: str = "default", + ) -> torch.Tensor: + pixel_values = values.pop("pixel_values") + metadata = values + image_features = self.vision_tower( + pixel_values, grid_hw=None, encoder_metadata=metadata + ) + return self.multi_modal_projector(image_features) + + def encoder_eager_forward( + self, + mm_kwargs: dict[str, Any], + path: str = "default", + ) -> torch.Tensor: + pixel_values = mm_kwargs["pixel_values"] + image_grid_hws = mm_kwargs["image_grid_hws"] + image_features = self.vision_tower(pixel_values, image_grid_hws) + return self.multi_modal_projector(torch.cat(image_features)) + def _parse_and_validate_image_input( self, **kwargs: object ) -> KimiVLImageInputs | None: diff --git a/vllm/model_executor/models/moonvit.py b/vllm/model_executor/models/moonvit.py index 8c699865618..73e17cb9fb6 100644 --- a/vllm/model_executor/models/moonvit.py +++ b/vllm/model_executor/models/moonvit.py @@ -45,7 +45,9 @@ from collections.abc import Sequence from copy import deepcopy from functools import cached_property +from typing import Any +import numpy as np import torch import torch.nn as nn import torch.nn.functional as F @@ -110,23 +112,42 @@ class Learnable2DInterpPosEmb(nn.Module): def reset_parameters(self): nn.init.normal_(self.weight) - def forward(self, x: torch.Tensor, grid_hws: torch.Tensor) -> torch.Tensor: - pos_embs = [] - for shape in grid_hws.tolist(): - if shape == self.weight.shape[:-1]: + def get_pos_embeds( + self, + grid_hws_list: list[list[int]] | list[tuple[int, int]], + ) -> torch.Tensor: + """Build packed per-token positional embeddings for a list of grids. + + Returns a tensor of shape ``(sum(h * w), dim)`` formed by interpolating + the learned ``(height, width, dim)`` weight to each ``(h, w)`` grid and + concatenating the flattened results in the same order as + ``grid_hws_list``. Lives outside the captured CUDA graph so the + per-grid Python iteration is safe. + """ + weight_shape = list(self.weight.shape[:-1]) + pos_embs: list[torch.Tensor] = [] + for shape in grid_hws_list: + shape_list = [int(shape[0]), int(shape[1])] + if shape_list == weight_shape: pos_embs.append(self.weight.flatten(end_dim=1)) else: pos_embs.append( F.interpolate( self.weight.permute((2, 0, 1)).unsqueeze(0), - size=shape, + size=tuple(shape_list), mode=self.interpolation_mode, ) .squeeze(0) .permute((1, 2, 0)) .flatten(end_dim=1) ) - out = x + torch.cat(pos_embs) + if not pos_embs: + return self.weight.new_zeros((0, self.weight.shape[-1])) + return torch.cat(pos_embs) + + def forward(self, x: torch.Tensor, grid_hws: torch.Tensor) -> torch.Tensor: + pos_embs = self.get_pos_embeds(grid_hws.tolist()) + out = x + pos_embs return out @@ -158,19 +179,29 @@ class MoonVisionPatchEmbed(nn.Module): height=pos_emb_height, width=pos_emb_width, dim=out_dim ) - def forward(self, x: torch.Tensor, grid_hw: torch.Tensor) -> torch.Tensor: + def forward( + self, + x: torch.Tensor, + grid_hw: torch.Tensor | None = None, + *, + pos_embeds: torch.Tensor | None = None, + ) -> torch.Tensor: """ Args: x (L, Channels): input tensor grid_hw (N, 2): grid height and width + pos_embeds: precomputed positional embeddings of shape + ``(L, Cout)``. When provided, ``grid_hw`` is unused and the + CUDA-graph-incompatible interpolation in ``self.pos_emb`` is + skipped. Returns: (L, Cout) tensor """ x = self.proj(x).view(x.size(0), -1) - # apply positional embedding - x = self.pos_emb(x, grid_hw) - return x + if pos_embeds is not None: + return x + pos_embeds + return self.pos_emb(x, grid_hw) class Rope2DPosEmb(nn.Module): @@ -243,6 +274,35 @@ class Rope2DPosEmb(nn.Module): freqs_cis = freqs_cis.reshape(self.max_height, self.max_width, -1) return freqs_cis + def get_freqs_cis_by_seqlens_list( + self, + grid_hws_list: list[list[int]] | list[tuple[int, int]], + ) -> torch.Tensor: + """List-based variant of :meth:`get_freqs_cis_by_seqlens`. + + Accepts a Python list of ``(h, w)`` pairs so callers that already + operate outside the captured CUDA graph can avoid materializing a + tensor + ``.tolist()`` round-trip. + """ + assert all( + 1 <= h <= self.max_height and 1 <= w <= self.max_width + for h, w in grid_hws_list + ), ( + grid_hws_list, + self.max_height, + self.max_width, + ) + if not grid_hws_list: + return self.precomputed_freqs_cis.new_zeros((0, self.dim // 2)) + freqs_cis = torch.cat( + [ + self.precomputed_freqs_cis[:h, :w].reshape(-1, self.dim // 2) + for h, w in grid_hws_list + ], + dim=0, + ) + return freqs_cis + def get_freqs_cis_by_seqlens(self, grid_hws: torch.Tensor) -> torch.Tensor: """ Args: @@ -250,22 +310,7 @@ class Rope2DPosEmb(nn.Module): Returns: freqs_cis: tensor of shape (sum(t * height * width), dim//2) """ - shapes = grid_hws.tolist() - assert all( - 1 <= h <= self.max_height and 1 <= w <= self.max_width for h, w in shapes - ), ( - shapes, - self.max_height, - self.max_width, - ) - freqs_cis = torch.cat( - [ - self.precomputed_freqs_cis[:h, :w].reshape(-1, self.dim // 2) - for h, w in shapes - ], - dim=0, - ) - return freqs_cis + return self.get_freqs_cis_by_seqlens_list(grid_hws.tolist()) def get_freqs_cis_by_idx( self, pos_idx: torch.Tensor, pos_idx_mask: torch.Tensor @@ -392,11 +437,15 @@ class MoonVitEncoderLayer(nn.Module): x: torch.Tensor, cu_seqlens: torch.Tensor, rope_freqs_cis: torch.Tensor | None = None, + max_seqlen: torch.Tensor | None = None, ): """ Args: x (torch.Tensor): (seqlen, hidden_dim) cu_seqlens (torch.Tensor): + max_seqlen: Optional precomputed scalar tensor. When omitted it + is derived from ``cu_seqlens``, which produces a GPU scalar + that breaks CUDA graph capture. """ seq_length = x.size(0) xqkv, _ = self.wqkv(x) @@ -412,7 +461,8 @@ class MoonVitEncoderLayer(nn.Module): xq, xk = apply_rope(xq, xk, rope_freqs_cis) - max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max() + if max_seqlen is None: + max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max() attn_out = self.attn( xq.unsqueeze(0), xk.unsqueeze(0), @@ -433,10 +483,12 @@ class MoonVitEncoderLayer(nn.Module): hidden_states: torch.Tensor, cu_seqlens: torch.Tensor, rope_freqs_cis: torch.Tensor | None = None, + max_seqlen: torch.Tensor | None = None, ) -> torch.Tensor: """ Args: hidden_states: non-packed (B, N, D) or packed (L, D). if non-packed, seqlens should be None, if packed, seqlens should be set + max_seqlen: optional precomputed max-sequence-length scalar. Returns: output: same shape of input, non-packed (B, N, D) for non-packed input, (L, D) for packed input @@ -444,7 +496,10 @@ class MoonVitEncoderLayer(nn.Module): residual = hidden_states hidden_states = self.norm0(hidden_states) attn_out = self.attention_qkvpacked( - hidden_states, cu_seqlens, rope_freqs_cis=rope_freqs_cis + hidden_states, + cu_seqlens, + rope_freqs_cis=rope_freqs_cis, + max_seqlen=max_seqlen, ) hidden_states = residual + attn_out @@ -478,22 +533,39 @@ class MoonVitEncoder(nn.Module): ) self.final_layernorm = nn.LayerNorm(hidden_dim) - def forward( - self, hidden_states: torch.Tensor, grid_hw: torch.Tensor + def get_rope_freqs_cis( + self, + grid_hws_list: list[list[int]] | list[tuple[int, int]], ) -> torch.Tensor: - rope_freqs_cis = self.rope_2d.get_freqs_cis_by_seqlens(grid_hws=grid_hw) + return self.rope_2d.get_freqs_cis_by_seqlens_list(grid_hws_list) - lengths = torch.cat( - ( - torch.zeros(1, device=hidden_states.device, dtype=grid_hw.dtype), - (grid_hw[:, 0] * grid_hw[:, 1]).to(hidden_states.device), + def forward( + self, + hidden_states: torch.Tensor, + grid_hw: torch.Tensor | None = None, + *, + cu_seqlens: torch.Tensor | None = None, + rope_freqs_cis: torch.Tensor | None = None, + max_seqlen: torch.Tensor | None = None, + ) -> torch.Tensor: + if rope_freqs_cis is None: + rope_freqs_cis = self.rope_2d.get_freqs_cis_by_seqlens(grid_hws=grid_hw) + + if cu_seqlens is None: + lengths = torch.cat( + ( + torch.zeros(1, device=hidden_states.device, dtype=grid_hw.dtype), + (grid_hw[:, 0] * grid_hw[:, 1]).to(hidden_states.device), + ) ) - ) - cu_seqlens = lengths.cumsum(dim=0, dtype=torch.int32) + cu_seqlens = lengths.cumsum(dim=0, dtype=torch.int32) for _, block in enumerate(self.blocks): hidden_states = block( - hidden_states, cu_seqlens, rope_freqs_cis=rope_freqs_cis + hidden_states, + cu_seqlens, + rope_freqs_cis=rope_freqs_cis, + max_seqlen=max_seqlen, ) hidden_states = self.final_layernorm(hidden_states) @@ -530,6 +602,54 @@ def patch_merger( return outputs +def patch_merger_packed( + x: torch.Tensor, + gather_idx: torch.Tensor, + merge_kernel_size: tuple[int, int], +) -> torch.Tensor: + """CUDA-graph-safe equivalent of :func:`patch_merger`. + + Uses a precomputed index tensor to gather the per-token reshape + + permute that ``patch_merger`` does inside a Python loop. The output is + the concatenated 3D tensor ``(sum(new_h * new_w), kh * kw, d_model)``, + matching what ``torch.cat(patch_merger(...))`` would produce. + """ + kh, kw = merge_kernel_size + d_model = x.size(-1) + return x.index_select(0, gather_idx).view(-1, kh * kw, d_model) + + +def _build_merge_gather_idx( + grid_hws_list: list[list[int]] | list[tuple[int, int]], + merge_kernel_size: tuple[int, int], +) -> np.ndarray: + """Build the per-token gather indices used by :func:`patch_merger_packed`. + + For each item with grid (h, w) and merge kernel (kh, kw), the output + block at position (nh, nw) gathers the kh*kw input tokens at rows + (nh*kh + ih, nw*kw + iw) of that item, in (ih, iw) row-major order. + """ + kh, kw = merge_kernel_size + parts: list[np.ndarray] = [] + pre_sum = 0 + for h, w in grid_hws_list: + new_h, new_w = h // kh, w // kw + nh = np.arange(new_h, dtype=np.int64).reshape(new_h, 1, 1, 1) + nw = np.arange(new_w, dtype=np.int64).reshape(1, new_w, 1, 1) + ih = np.arange(kh, dtype=np.int64).reshape(1, 1, kh, 1) + iw = np.arange(kw, dtype=np.int64).reshape(1, 1, 1, kw) + # Linearized input row = (nh*kh + ih) * w + (nw*kw + iw), offset by + # the per-item base ``pre_sum``. Output is laid out as + # (new_h, new_w, kh, kw) which patch_merger flattens to + # (new_h*new_w, kh*kw). + idx = pre_sum + (nh * kh + ih) * w + (nw * kw + iw) + parts.append(idx.reshape(-1)) + pre_sum += h * w + if not parts: + return np.zeros(0, dtype=np.int64) + return np.concatenate(parts) + + class MoonVitPretrainedModel(PreTrainedModel): config_class = MoonViTConfig model_type = "moonvit" @@ -570,17 +690,126 @@ class MoonVitPretrainedModel(PreTrainedModel): prefix=f"{prefix}.encoder", ) + def prepare_encoder_metadata( + self, + grid_hws_list: list[list[int]] | list[tuple[int, int]], + *, + max_batch_size: int | None = None, + max_seqlen_override: int | None = None, + device: torch.device | None = None, + ) -> dict[str, Any]: + """Precompute every grid-dependent input the encoder needs. + + Used by the CUDA graph capture and replay paths to precompute + every grid-dependent input outside the captured graph, so per-grid + Python iteration and ``.tolist()`` round-trips are fine; the + values are then copied into fixed-shape buffers for replay. + + Args: + grid_hws_list: List of ``(h, w)`` patch-grid sizes per image. + max_batch_size: When set, ``cu_seqlens`` is right-padded with + its last value so the buffer covers up to this many + sequences. Required at CUDA graph capture/replay so the + buffer shape matches what was recorded; padding entries + are zero-length sequences and are ignored by varlen + attention. + max_seqlen_override: Override the per-replay max sequence + length scalar. At capture this must be a safe upper bound + (worst case: a single image consuming the full token + budget) because the value is baked into the captured + graph. + device: Device for the metadata tensors. Defaults to the + model's parameter device. + """ + if device is None: + device = next(self.parameters()).device + + # Normalize to a list of plain Python int pairs so the helpers + # below never need ``.tolist()`` on a tensor. + grid_pairs: list[tuple[int, int]] = [(int(h), int(w)) for h, w in grid_hws_list] + + metadata: dict[str, Any] = {} + + pos_embeds = self.patch_embed.pos_emb.get_pos_embeds(grid_pairs) + metadata["pos_embeds"] = pos_embeds.to(device=device) + + rope_freqs_cis = self.encoder.get_rope_freqs_cis(grid_pairs) + metadata["rope_freqs_cis"] = rope_freqs_cis.to(device=device) + + grid_arr = np.array(grid_pairs, dtype=np.int64) + seq_lens = (grid_arr[:, 0] * grid_arr[:, 1]).astype(np.int32) + cu_seqlens_np = np.concatenate( + [ + np.zeros(1, dtype=np.int32), + seq_lens.cumsum(dtype=np.int32), + ] + ) + + if max_batch_size is not None: + num_seqs = len(cu_seqlens_np) - 1 + if num_seqs < max_batch_size: + cu_seqlens_np = np.concatenate( + [ + cu_seqlens_np, + np.full( + max_batch_size - num_seqs, + cu_seqlens_np[-1], + dtype=np.int32, + ), + ] + ) + metadata["cu_seqlens"] = torch.from_numpy(cu_seqlens_np).to(device) + + if max_seqlen_override is not None: + max_seqlen_val = int(max_seqlen_override) + else: + max_seqlen_val = int(seq_lens.max()) if len(seq_lens) > 0 else 0 + # Keep on CPU: attention wrappers may call .item() on this scalar + # and we want that materialization to happen outside the captured + # graph (the value is constant per capture anyway). + metadata["max_seqlen"] = torch.tensor(max_seqlen_val, dtype=torch.int32) + + gather_idx_np = _build_merge_gather_idx(grid_pairs, self.merge_kernel_size) + metadata["merge_gather_idx"] = torch.from_numpy(gather_idx_np).to(device) + + return metadata + def forward( - self, pixel_values: torch.Tensor, grid_hw: torch.Tensor - ) -> torch.Tensor: + self, + pixel_values: torch.Tensor, + grid_hw: torch.Tensor, + *, + encoder_metadata: dict[str, Any] | None = None, + ) -> torch.Tensor | list[torch.Tensor]: """ Args: pixel_values (torch.Tensor): The input pixel values. grid_hw (torch.Tensor): The grid height and width. - - Returns: - torch.Tensor: The output tokens. + encoder_metadata: Optional precomputed metadata produced by + :meth:`prepare_encoder_metadata`. When provided every + ``.tolist()`` call in the forward path is skipped, the + returned tensor is the packed + ``(sum(new_h*new_w), kh*kw, hidden_size)`` form (suitable + for CUDA graph capture/replay), and ``grid_hw`` is unused. + When ``None`` the legacy path runs and returns a list of + per-image tensors. """ + if encoder_metadata is not None: + hidden_states = self.patch_embed( + pixel_values, pos_embeds=encoder_metadata["pos_embeds"] + ) + hidden_states = self.encoder( + hidden_states, + cu_seqlens=encoder_metadata["cu_seqlens"], + rope_freqs_cis=encoder_metadata["rope_freqs_cis"], + max_seqlen=encoder_metadata["max_seqlen"], + ) + return patch_merger_packed( + hidden_states, + encoder_metadata["merge_gather_idx"], + merge_kernel_size=self.merge_kernel_size, + ) + hidden_states = self.patch_embed(pixel_values, grid_hw) hidden_states = self.encoder(hidden_states, grid_hw) hidden_states = patch_merger( From 1a59078c873fb6dd14f23f17307a24977cdd9228 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Wed, 17 Jun 2026 20:23:44 +0800 Subject: [PATCH 20/37] [CI/Build] Avoid duplicate ViT CG test introduced by accident (#45654) Signed-off-by: Isotr0py --- .buildkite/test_areas/models_multimodal.yaml | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/.buildkite/test_areas/models_multimodal.yaml b/.buildkite/test_areas/models_multimodal.yaml index 48d24708358..a7358e8dbd6 100644 --- a/.buildkite/test_areas/models_multimodal.yaml +++ b/.buildkite/test_areas/models_multimodal.yaml @@ -30,7 +30,6 @@ steps: - pip install git+https://github.com/TIGER-AI-Lab/Mantis.git - pytest -v -s models/multimodal/generation/test_common.py -m core_model -k "qwen3 or gemma" - pytest -v -s models/multimodal/generation/test_qwen2_5_vl.py -m core_model - - pytest -v -s models/multimodal/generation/test_vit_cudagraph.py -m core_model mirror: amd: device: mi325_1 @@ -63,9 +62,16 @@ steps: - tests/models/multimodal commands: - pip install git+https://github.com/TIGER-AI-Lab/Mantis.git - - pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/processing + - pytest -v -s models/multimodal -m core_model --ignore models/multimodal/generation/test_common.py --ignore models/multimodal/generation/test_ultravox.py --ignore models/multimodal/generation/test_qwen2_5_vl.py --ignore models/multimodal/generation/test_qwen2_vl.py --ignore models/multimodal/generation/test_whisper.py --ignore models/multimodal/generation/test_memory_leak.py --ignore models/multimodal/generation/test_vit_cudagraph.py --ignore models/multimodal/processing + - pytest -v -s models/multimodal/generation/test_vit_cudagraph.py -m core_model - pytest models/multimodal/generation/test_memory_leak.py -m core_model - cd .. && VLLM_WORKER_MULTIPROC_METHOD=spawn pytest -v -s tests/models/multimodal/generation/test_whisper.py -m core_model # Otherwise, mp_method="spawn" doesn't work + mirror: + amd: + soft_fail: true + device: mi325_1 + depends_on: + - image-build-amd - label: Multi-Modal Processor (CPU) key: multi-modal-processor-cpu From 06e1e0885c30a96413850d7fd9a1e6f3cfde414b Mon Sep 17 00:00:00 2001 From: Chaojun Zhang Date: Wed, 17 Jun 2026 20:26:47 +0800 Subject: [PATCH 21/37] [XPU] Fix test_logprobs_e2e import error: pin lm-eval[api]>=0.4.12 (#44469) Signed-off-by: Chaojun Zhang --- .buildkite/intel_jobs/misc_intel.yaml | 6 ++++-- requirements/test/xpu.in | 2 +- requirements/test/xpu.txt | 7 +------ 3 files changed, 6 insertions(+), 9 deletions(-) diff --git a/.buildkite/intel_jobs/misc_intel.yaml b/.buildkite/intel_jobs/misc_intel.yaml index cdbbdc850be..d74494ed8b3 100644 --- a/.buildkite/intel_jobs/misc_intel.yaml +++ b/.buildkite/intel_jobs/misc_intel.yaml @@ -57,14 +57,16 @@ steps: commands: - >- bash .buildkite/scripts/hardware_ci/run-intel-test.sh - 'export VLLM_WORKER_MULTIPROC_METHOD=spawn && + 'pip install lm_eval[api]>=0.4.12 && + export VLLM_WORKER_MULTIPROC_METHOD=spawn && cd tests && pytest -v -s v1/logits_processors --ignore=v1/logits_processors/test_custom_online.py --ignore=v1/logits_processors/test_custom_offline.py && pytest -v -s v1/test_oracle.py && pytest -v -s v1/test_request.py && pytest -v -s v1/test_outputs.py && pytest -v -s v1/sample/test_topk_topp_sampler.py && - pytest -v -s v1/sample/test_logprobs.py' + pytest -v -s v1/sample/test_logprobs.py && + pytest -v -s v1/sample/test_logprobs_e2e.py' - label: XPU CPU Offload timeout_in_minutes: 60 diff --git a/requirements/test/xpu.in b/requirements/test/xpu.in index a828867845f..161e2c6871f 100644 --- a/requirements/test/xpu.in +++ b/requirements/test/xpu.in @@ -13,7 +13,7 @@ pytest-shard absl-py accelerate arctic-inference -lm_eval[api] +lm_eval[api]>=0.4.12 modelscope # --- Audio Processing --- diff --git a/requirements/test/xpu.txt b/requirements/test/xpu.txt index 820ce27bc3d..1b1f3c91c5e 100644 --- a/requirements/test/xpu.txt +++ b/requirements/test/xpu.txt @@ -33,7 +33,6 @@ arctic-inference==0.1.1 attrs==26.1.0 # via # aiohttp - # jsonlines # jsonschema # referencing audioread==3.0.1 @@ -225,8 +224,6 @@ joblib==1.5.3 # librosa # nltk # scikit-learn -jsonlines==4.0.0 - # via lm-eval jsonschema==4.26.0 # via # -c requirements/common.txt @@ -247,7 +244,7 @@ librosa==0.10.2.post1 # via -r requirements/test/xpu.in llvmlite==0.47.0 # via numba -lm-eval==0.4.11 +lm-eval==0.4.12 # via -r requirements/test/xpu.in lxml==6.0.2 # via @@ -734,5 +731,3 @@ xxhash==3.6.0 # evaluate yarl==1.23.0 # via aiohttp -zstandard==0.25.0 - # via lm-eval From bcb518ad7a2138ad51b5f9912c07c63af35cee86 Mon Sep 17 00:00:00 2001 From: Yi Liu Date: Wed, 17 Jun 2026 21:51:32 +0800 Subject: [PATCH 22/37] [quant][autoround]Refactor INC quantization into package with INCScheme orchestrator (#40601) Signed-off-by: yiliu30 Signed-off-by: Zhenzhong1 Signed-off-by: Zhenzhong Xu Co-authored-by: n1ck-guo Co-authored-by: Zhenzhong1 --- .buildkite/intel_jobs/test-intel.yaml | 17 + requirements/xpu.txt | 2 +- tests/quantization/test_auto_round.py | 758 ++++++++++++++++- .../model_executor/layers/quantization/inc.py | 794 ------------------ .../layers/quantization/inc/__init__.py | 6 + .../layers/quantization/inc/config_parser.py | 188 +++++ .../layers/quantization/inc/inc.py | 192 +++++ .../layers/quantization/inc/inc_linear.py | 47 ++ .../quantization/inc/schemes/__init__.py | 13 + .../quantization/inc/schemes/factory.py | 22 + .../quantization/inc/schemes/inc_scheme.py | 104 +++ .../inc/schemes/inc_wna16_linear.py | 399 +++++++++ .../inc/schemes/inc_wna16_scheme.py | 201 +++++ 13 files changed, 1940 insertions(+), 803 deletions(-) delete mode 100644 vllm/model_executor/layers/quantization/inc.py create mode 100644 vllm/model_executor/layers/quantization/inc/__init__.py create mode 100644 vllm/model_executor/layers/quantization/inc/config_parser.py create mode 100644 vllm/model_executor/layers/quantization/inc/inc.py create mode 100644 vllm/model_executor/layers/quantization/inc/inc_linear.py create mode 100644 vllm/model_executor/layers/quantization/inc/schemes/__init__.py create mode 100644 vllm/model_executor/layers/quantization/inc/schemes/factory.py create mode 100644 vllm/model_executor/layers/quantization/inc/schemes/inc_scheme.py create mode 100644 vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_linear.py create mode 100644 vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py diff --git a/.buildkite/intel_jobs/test-intel.yaml b/.buildkite/intel_jobs/test-intel.yaml index 4bf14f7064b..afeb11e06d5 100644 --- a/.buildkite/intel_jobs/test-intel.yaml +++ b/.buildkite/intel_jobs/test-intel.yaml @@ -87,3 +87,20 @@ steps: cd tests && pytest -v -s entrypoints/multimodal/openai/chat_completion/test_audio_in_video.py && pytest -v -s benchmarks/test_serve_cli.py' + - label: "XPU quantization test" + depends_on: + - image-build-xpu + timeout_in_minutes: 30 + device: intel_gpu + no_plugin: true + env: + REGISTRY: "public.ecr.aws/q9t5s3a7" + REPO: "vllm-ci-test-repo" + source_file_dependencies: + - vllm/ + - .buildkite/intel_jobs/test-intel.yaml + commands: + - >- + bash .buildkite/scripts/hardware_ci/run-intel-test.sh + 'cd tests && + pytest -v -s quantization/test_auto_round.py' \ No newline at end of file diff --git a/requirements/xpu.txt b/requirements/xpu.txt index 33a808866c4..f17e2281f7a 100644 --- a/requirements/xpu.txt +++ b/requirements/xpu.txt @@ -16,5 +16,5 @@ torch==2.12.0 torchaudio torchvision -auto_round_lib>=0.13.0 +auto_round_lib>=0.13.3 vllm_xpu_kernels @ https://github.com/vllm-project/vllm-xpu-kernels/releases/download/v0.1.9.1/vllm_xpu_kernels-0.1.9.1-cp38-abi3-manylinux_2_28_x86_64.whl diff --git a/tests/quantization/test_auto_round.py b/tests/quantization/test_auto_round.py index 9f5db821950..5cd599f7211 100644 --- a/tests/quantization/test_auto_round.py +++ b/tests/quantization/test_auto_round.py @@ -10,23 +10,765 @@ Run `pytest tests/quantization/test_auto_round.py`. import pytest +from vllm.model_executor.layers.fused_moe import RoutedExperts +from vllm.model_executor.layers.linear import LinearBase, UnquantizedLinearMethod +from vllm.model_executor.layers.quantization.auto_gptq import AutoGPTQConfig +from vllm.model_executor.layers.quantization.inc import INCConfig +from vllm.model_executor.layers.quantization.inc.config_parser import INCLayerConfig +from vllm.model_executor.layers.quantization.inc.inc_linear import INCLinearMethod +from vllm.model_executor.layers.quantization.inc.schemes import ( + INCWna16Scheme, + resolve_scheme, +) +from vllm.model_executor.layers.quantization.inc.schemes.inc_scheme import ( + INCLinearScheme, +) +from vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_linear import ( + INCARKLinearMethod, + INCWNA16LinearScheme, + INCXPULinearMethod, +) +from vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_scheme import ( + _resolve_awq_moe, + _resolve_gptq_moe, +) +from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead from vllm.platforms import current_platform MODELS = [ - "OPEA/Qwen2.5-0.5B-Instruct-int4-sym-inc", ##auto_round:auto_gptq - "Intel/Qwen2-0.5B-Instruct-int4-sym-AutoRound", ##auto_round:auto_awq + pytest.param( + "OPEA/Qwen2.5-0.5B-Instruct-int4-sym-inc", + id="auto_round:auto_gptq", + ), + pytest.param( + "Intel/Qwen2-0.5B-Instruct-int4-sym-AutoRound", + marks=pytest.mark.skipif( + not current_platform.is_cuda(), + reason="AWQ AutoRound model only supports CUDA backend for now.", + ), + id="auto_round:auto_awq", + ), ] @pytest.mark.skipif( - not current_platform.is_cpu() - and not current_platform.is_xpu() - and not current_platform.is_cuda(), - reason="only supports CPU/XPU/CUDA backend.", + not ( + current_platform.is_cpu() + or current_platform.is_xpu() + or current_platform.is_cuda() + ), + reason="Only supports CPU/XPU/CUDA backend.", ) @pytest.mark.parametrize("model", MODELS) -def test_auto_round(vllm_runner, model): +def test_auto_round_model(vllm_runner, model): with vllm_runner(model, enforce_eager=True) as llm: output = llm.generate_greedy(["The capital of France is"], max_tokens=8) + assert output - print(f"{output[0][1]}") + print(output[0][1]) + + +# --------------------------------------------------------------------------- +# Unit tests for INCConfig and related classes +# --------------------------------------------------------------------------- + + +class DummyLayer: + pass + + +class DummyFusedMoE: + pass + + +def make_config(**overrides) -> INCConfig: + kwargs = { + "weight_bits": 4, + "group_size": 128, + "sym": True, + "packing_format": "auto_round:auto_gptq", + "block_name_to_quantize": None, + "extra_config": None, + "data_type": "int", + "backend": "auto", + } + kwargs.update(overrides) + return INCConfig(**kwargs) + + +def make_layer_config(**overrides) -> INCLayerConfig: + kwargs = { + "bits": 4, + "group_size": 128, + "sym": True, + "packing_format": "auto_round:auto_gptq", + "backend": "auto", + "data_type": "int", + "quantized": True, + } + kwargs.update(overrides) + return INCLayerConfig(**kwargs) + + +def test_inc_config_parser_exact_match() -> None: + config = make_config( + extra_config={ + "layers.0.self_attn.q_proj": { + "bits": 8, + "group_size": 64, + "sym": False, + } + } + ) + + layer_config = config.config_parser.resolve( + DummyLayer(), "layers.0.self_attn.q_proj" + ) + + assert layer_config.bits == 8 + assert layer_config.group_size == 64 + assert layer_config.sym is False + assert layer_config.quantized is True + + +def test_inc_model_prefix_early_exit() -> None: + """extra_config keys with model. prefix trigger early unquantized return.""" + config = make_config( + extra_config={ + "model.layers.1.mlp.gate_proj": { + "bits": 16, + }, + } + ) + + # get_quant_method checks model. prefix for unquantized early-exit + result = config.get_quant_method(DummyLayer(), "layers.1.mlp.gate_proj") + assert isinstance(result, UnquantizedLinearMethod) + + +def test_inc_config_parser_regex_match() -> None: + config = make_config( + extra_config={ + r"layers\.\d+\.self_attn\.(q|k|v)_proj": { + "bits": 8, + "group_size": 64, + "sym": False, + } + } + ) + + layer_config = config.config_parser.resolve( + DummyLayer(), "layers.3.self_attn.q_proj" + ) + + assert layer_config.bits == 8 + assert layer_config.group_size == 64 + assert layer_config.sym is False + + +def test_inc_config_parser_invalid_regex_ignored() -> None: + config = make_config( + extra_config={ + "[invalid": { + "bits": 8, + "group_size": 64, + "sym": False, + } + } + ) + + layer_config = config.config_parser.resolve( + DummyLayer(), "layers.0.self_attn.q_proj" + ) + + assert layer_config.bits == 4 + assert layer_config.group_size == 128 + assert layer_config.sym is True + + +def test_inc_config_parser_block_name_to_quantize_marks_unquantized() -> None: + config = make_config(block_name_to_quantize=["layers.1"]) + + layer_config = config.config_parser.resolve( + DummyLayer(), "layers.0.self_attn.q_proj" + ) + + assert layer_config.bits == 16 + assert layer_config.group_size == -1 + assert layer_config.sym is True + assert layer_config.quantized is False + + +def test_inc_config_parser_parallel_lm_head_defaults_to_unquantized() -> None: + layer = object.__new__(ParallelLMHead) + config = make_config() + + layer_config = config.config_parser.resolve(layer, "lm_head") + + assert layer_config.quantized is False + assert layer_config.bits == 16 + + +def test_inc_config_parser_fused_moe_requires_consistent_configs() -> None: + config = make_config( + extra_config={ + "layers.0.block_sparse_moe.experts.0.w1": { + "bits": 4, + "group_size": 128, + "sym": True, + }, + "layers.0.block_sparse_moe.experts.0.w2": { + "bits": 8, + "group_size": 128, + "sym": True, + }, + } + ) + + with pytest.raises(ValueError, match="requires consistent quant config"): + config.config_parser.resolve(DummyFusedMoE(), "layers.0.block_sparse_moe") + + +def test_inc_config_parser_fused_module_requires_consistent_configs() -> None: + config = make_config( + extra_config={ + "layers.0.self_attn.q_proj": { + "bits": 4, + "group_size": 128, + "sym": True, + }, + "layers.0.self_attn.k_proj": { + "bits": 8, + "group_size": 128, + "sym": True, + }, + "layers.0.self_attn.v_proj": { + "bits": 4, + "group_size": 128, + "sym": True, + }, + } + ) + config.packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]} + + with pytest.raises(ValueError, match="requires consistent quant config"): + config.config_parser.resolve(DummyLayer(), "layers.0.self_attn.qkv_proj") + + +def test_inc_layer_config_mx_fp_helpers() -> None: + layer_config = INCLayerConfig( + bits=4, + group_size=32, + sym=True, + packing_format="", + backend="", + data_type="mx_fp", + quantized=True, + ) + + assert layer_config.is_mxfp4 is True + assert layer_config.is_mxfp8 is False + + +def test_inc_resolve_scheme_selects_wna16() -> None: + layer_config = INCLayerConfig( + bits=4, + group_size=128, + sym=True, + packing_format="auto_round:auto_gptq", + backend="auto", + data_type="int", + quantized=True, + ) + + scheme = resolve_scheme(layer_config) + + assert isinstance(scheme, INCWna16Scheme) + + +class DummyLinearScheme(INCLinearScheme): + def __init__(self) -> None: + self.calls: list[tuple] = [] + + @classmethod + def get_min_capability(cls) -> int: + return 0 + + def create_weights(self, *args, **kwargs) -> None: + self.calls.append(("create_weights", args, kwargs)) + + def process_weights_after_loading(self, layer) -> None: + self.calls.append(("process_weights_after_loading", layer)) + + def apply_weights(self, layer, x, bias=None): + self.calls.append(("apply_weights", layer, x, bias)) + return "applied" + + +def test_inc_linear_method_delegates() -> None: + scheme = DummyLinearScheme() + method = INCLinearMethod(scheme) + layer = DummyLayer() + + method.create_weights( + layer, + input_size_per_partition=1, + output_partition_sizes=[2], + input_size=1, + output_size=2, + params_dtype=None, + ) + method.process_weights_after_loading(layer) + result = method.apply(layer, "x", "b") + + assert result == "applied" + assert [call[0] for call in scheme.calls] == [ + "create_weights", + "process_weights_after_loading", + "apply_weights", + ] + + +def test_wna16_xpu_prefers_ark_when_available(monkeypatch) -> None: + class DummyQuantLinear: + pass + + monkeypatch.setattr(current_platform, "is_xpu", lambda: True) + monkeypatch.setattr(current_platform, "is_cpu", lambda: False) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_linear.get_ark_state", + lambda: (True, None, object(), DummyQuantLinear), + ) + + method = INCWna16Scheme().get_linear_method( + make_config(), + object(), + "layer", + make_layer_config(), + ) + + assert isinstance(method, INCLinearMethod) + assert isinstance(method.scheme, INCARKLinearMethod) + + +def test_wna16_xpu_falls_back_when_ark_unavailable(monkeypatch) -> None: + monkeypatch.setattr(current_platform, "is_xpu", lambda: True) + monkeypatch.setattr(current_platform, "is_cpu", lambda: False) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_linear.get_ark_state", + lambda: (False, "missing", None, None), + ) + + method = INCWna16Scheme().get_linear_method( + make_config(), + object(), + "layer", + make_layer_config(), + ) + + assert isinstance(method, INCLinearMethod) + assert isinstance(method.scheme, INCXPULinearMethod) + + +def test_wna16_cpu_gptq_prefers_ark_when_available(monkeypatch) -> None: + class DummyQuantLinear: + pass + + monkeypatch.setattr(current_platform, "is_xpu", lambda: False) + monkeypatch.setattr(current_platform, "is_cpu", lambda: True) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_linear.get_ark_state", + lambda: (True, None, object(), DummyQuantLinear), + ) + + method = INCWna16Scheme().get_linear_method( + make_config(), + object(), + "layer", + make_layer_config(), + ) + + assert isinstance(method, INCLinearMethod) + assert isinstance(method.scheme, INCARKLinearMethod) + + +def test_wna16_cpu_gptq_raises_when_ark_and_marlin_unavailable( + monkeypatch, +) -> None: + monkeypatch.setattr(current_platform, "is_xpu", lambda: False) + monkeypatch.setattr(current_platform, "is_cpu", lambda: True) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_linear.get_ark_state", + lambda: (False, "missing", None, None), + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_linear.check_marlin_supported", + lambda *args, **kwargs: False, + ) + + with pytest.raises(NotImplementedError, match="Only 4-bit and 8-bit symmetric"): + INCWna16Scheme().get_linear_method( + make_config(), + object(), + "layer", + make_layer_config(), + ) + + +def test_wna16_linear_gptq_uses_auto_gptq_when_supported(monkeypatch) -> None: + captured = {} + + class DummyMethod: + def __init__(self, cfg): + captured["cfg"] = cfg + + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.inc_wna16_linear." + "check_marlin_supported", + lambda *args, **kwargs: True, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.auto_gptq.AutoGPTQLinearMethod", + DummyMethod, + ) + + scheme = INCWNA16LinearScheme(make_layer_config()) + + assert isinstance(scheme.inner_method, DummyMethod) + assert isinstance(captured["cfg"], AutoGPTQConfig) + assert captured["cfg"].weight_bits == 4 + assert captured["cfg"].group_size == 128 + assert captured["cfg"].is_sym is True + + +def test_wna16_linear_gptq_unsupported_config_raises() -> None: + with pytest.raises(NotImplementedError, match="Only 4-bit and 8-bit symmetric"): + INCWNA16LinearScheme(make_layer_config(sym=False)) + + +def test_wna16_xpu_unsupported_config_still_raises(monkeypatch) -> None: + monkeypatch.setattr(current_platform, "is_xpu", lambda: True) + monkeypatch.setattr(current_platform, "is_cpu", lambda: False) + + with pytest.raises(NotImplementedError, match="unsupported config"): + INCWna16Scheme().get_linear_method( + make_config(sym=False), + object(), + "layer", + make_layer_config(sym=False), + ) + + +def test_inc_get_quant_method_unquantized_linear_returns_unquantized() -> None: + config = make_config(extra_config={"layer": {"bits": 16}}) + layer = object.__new__(LinearBase) + + method = config.get_quant_method(layer, "layer") + + assert isinstance(method, UnquantizedLinearMethod) + + +def test_inc_get_quant_method_unquantized_moe_returns_unquantized( + monkeypatch, +) -> None: + """Early-exit returns UnquantizedFusedMoEMethod for FusedMoE layers + when extra_config has bits >= 16.""" + config = make_config(extra_config={"layer": {"bits": 16}}) + layer = object.__new__(RoutedExperts) + layer.moe_config = None # UnquantizedFusedMoEMethod accepts moe_config + + class DummyUnquantizedFusedMoEMethod: + def __init__(self, moe_config) -> None: + self.moe_config = moe_config + + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.inc.UnquantizedFusedMoEMethod", + DummyUnquantizedFusedMoEMethod, + ) + + method = config.get_quant_method(layer, "layer") + + assert isinstance(method, DummyUnquantizedFusedMoEMethod) + assert method.moe_config is None + + +def test_inc_get_quant_method_linear_uses_resolved_scheme(monkeypatch) -> None: + config = make_config() + layer = object.__new__(LinearBase) + sentinel = object() + + class DummyScheme: + def get_linear_method(self, _config, _layer, _prefix, _layer_config): + return sentinel + + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.factory.resolve_scheme", + lambda _layer_config: DummyScheme(), + ) + + method = config.get_quant_method(layer, "layer") + + assert method is sentinel + + +def test_inc_get_quant_method_moe_uses_resolved_scheme(monkeypatch) -> None: + config = make_config() + layer = object.__new__(RoutedExperts) + sentinel = object() + + class DummyScheme: + def get_moe_method(self, _config, _layer, _prefix, _layer_config): + return sentinel + + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.inc.schemes.factory.resolve_scheme", + lambda _layer_config: DummyScheme(), + ) + + method = config.get_quant_method(layer, "layer") + + assert method is sentinel + + +def test_resolve_gptq_moe_falls_back_to_moe_wna16(monkeypatch) -> None: + captured = {} + + class DummyMoeConfig: + pass + + class DummyLayer: + moe_config = DummyMoeConfig() + + class DummyBuiltConfig: + pass + + built_config = DummyBuiltConfig() + + class DummyMethod: + def __init__(self, cfg, moe): + captured["cfg"] = cfg + captured["moe"] = moe + + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.utils.marlin_utils.check_marlin_supported", + lambda *args, **kwargs: False, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.moe_wna16.MoeWNA16Config.from_config", + lambda cfg: captured.update({"from_config": cfg}) or built_config, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.moe_wna16.MoeWNA16Method", + DummyMethod, + ) + + layer_config = INCLayerConfig( + bits=4, + group_size=128, + sym=True, + packing_format="auto_round:auto_gptq", + backend="auto", + data_type="int", + quantized=True, + ) + + _resolve_gptq_moe(DummyLayer(), layer_config) + + assert captured["from_config"] == { + "quant_method": "gptq", + "bits": 4, + "group_size": 128, + "sym": True, + "lm_head": False, + } + assert captured["cfg"] is built_config + assert captured["moe"] is DummyLayer.moe_config + + +def test_resolve_gptq_moe_uses_auto_gptq_when_supported(monkeypatch) -> None: + captured = {} + + class DummyMoeConfig: + pass + + class DummyLayer: + moe_config = DummyMoeConfig() + + class DummyMethod: + def __init__(self, cfg, moe): + captured["cfg"] = cfg + captured["moe"] = moe + + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.utils.marlin_utils.check_marlin_supported", + lambda *args, **kwargs: True, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.utils.marlin_utils." + "check_moe_marlin_supports_layer", + lambda *args, **kwargs: True, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.auto_gptq.AutoGPTQMoEMethod", + DummyMethod, + ) + + _resolve_gptq_moe(DummyLayer(), make_layer_config()) + + assert isinstance(captured["cfg"], AutoGPTQConfig) + assert captured["cfg"].weight_bits == 4 + assert captured["cfg"].group_size == 128 + assert captured["cfg"].is_sym is True + assert captured["moe"] is DummyLayer.moe_config + + +def test_resolve_awq_moe_uses_marlin_when_supported(monkeypatch) -> None: + captured = {} + + class DummyMoeConfig: + pass + + class DummyLayer: + moe_config = DummyMoeConfig() + + class DummyMethod: + def __init__(self, cfg, moe): + captured["cfg"] = cfg + captured["moe"] = moe + + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.utils.marlin_utils.check_marlin_supported", + lambda *args, **kwargs: True, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.utils.marlin_utils.check_moe_marlin_supports_layer", + lambda *args, **kwargs: True, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.awq_marlin.verify_marlin_supported", + lambda *args, **kwargs: None, + ) + monkeypatch.setattr( + "vllm.model_executor.layers.quantization.awq_marlin.AWQMarlinMoEMethod", + DummyMethod, + ) + + layer_config = INCLayerConfig( + bits=4, + group_size=128, + sym=False, + packing_format="auto_round:auto_awq", + backend="auto", + data_type="int", + quantized=True, + ) + + _resolve_awq_moe(DummyLayer(), layer_config) + + assert captured["cfg"].weight_bits == 4 + assert captured["cfg"].zero_point is True + assert captured["moe"] is DummyLayer.moe_config + + +# --------------------------------------------------------------------------- +# Tests for get_layer_config step 4 (fused QKV / packed_modules_mapping) +# --------------------------------------------------------------------------- + + +class TestGetLayerConfigFusedQKV: + """Tests for step-4 (fused QKV / packed_modules_mapping) logic. + + Focused on preventing false-positive substring matches. + """ + + def test_exact_fusion_key_match(self): + """A layer whose name contains 'qkv' maps to its extra_config entry.""" + config = make_config( + extra_config={ + "model.layers.0.self_attn.qkv_proj": {"bits": 8}, + } + ) + config.packed_modules_mapping = { + "qkv_proj": ["q_proj", "k_proj", "v_proj"], + } + bits, _, _ = config.get_layer_config( + DummyLayer(), "model.layers.0.self_attn.qkv_proj" + ) + assert bits == 8 + + def test_false_substring_match_does_not_override(self): + """Regression test for the false-substring-match bug. + + Scenario (Qwen3.6-35B-A3B VLM): + - packed_modules_mapping has "qkv" → ["qkv"] (from vision encoder). + - The GDN text-attention layer is named "in_proj_qkvz". + - "qkv" is a substring of "in_proj_qkvz", so old code would enter + step 4 and generate sub_name "in_proj_qkvz" (replacing "qkv" with + "qkv"). That name is NOT in extra_config, so get_config() falls + back to the global default (bits=4), even though correct is 16. + - Fix: skip the fusion key when none of the generated sub_names + actually exist in extra_config. + """ + config = make_config( + extra_config={ + "model.layers.0.in_proj_qkv": {"bits": 16}, + "model.layers.0.in_proj_z": {"bits": 16}, + } + ) + config.packed_modules_mapping = { + "qkv": ["qkv"], + } + bits, _, _ = config.get_layer_config( + DummyLayer(), "model.layers.0.in_proj_qkvz" + ) + # bits should be the global default (4) – no erroneous fusion match + assert bits == 4 + + def test_real_qkv_fusion_key_still_resolves(self): + """The true "qkv" fusion (vision encoder) still resolves correctly.""" + config = make_config( + extra_config={ + "vision_model.encoder.layers.0.self_attn.qkv": {"bits": 8}, + } + ) + config.packed_modules_mapping = { + "qkv": ["qkv"], + } + bits, _, _ = config.get_layer_config( + DummyLayer(), "vision_model.encoder.layers.0.self_attn.qkv" + ) + assert bits == 8 + + def test_mixed_fp16_and_int4_fused_layer(self): + """All sub-keys must agree; inconsistent configs raise ValueError.""" + config = make_config( + extra_config={ + "model.layers.0.self_attn.q_proj": {"bits": 16}, + "model.layers.0.self_attn.k_proj": {"bits": 4}, + "model.layers.0.self_attn.v_proj": {"bits": 4}, + } + ) + config.packed_modules_mapping = { + "qkv_proj": ["q_proj", "k_proj", "v_proj"], + } + with pytest.raises(ValueError, match="consistent quant config"): + config.get_layer_config(DummyLayer(), "model.layers.0.self_attn.qkv_proj") + + def test_fusion_triggered_by_regex_configured_sub_name(self): + """Fusion step 4 is still triggered when sub_names match via regex. + + Ensures the guard does not regress when extra_config uses regex + patterns instead of exact keys to configure sub-modules. + """ + config = make_config( + extra_config={ + r"model\.layers\.\d+\.self_attn\.(q|k|v)_proj": {"bits": 8}, + } + ) + config.packed_modules_mapping = { + "qkv_proj": ["q_proj", "k_proj", "v_proj"], + } + bits, _, _ = config.get_layer_config( + DummyLayer(), "model.layers.0.self_attn.qkv_proj" + ) + assert bits == 8 diff --git a/vllm/model_executor/layers/quantization/inc.py b/vllm/model_executor/layers/quantization/inc.py deleted file mode 100644 index 3a4d7d4039f..00000000000 --- a/vllm/model_executor/layers/quantization/inc.py +++ /dev/null @@ -1,794 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project - -from fractions import Fraction -from functools import lru_cache -from typing import TYPE_CHECKING, Any - -import regex as re -import torch -from torch.nn.parameter import Parameter - -from vllm.logger import init_logger -from vllm.model_executor.layers.fused_moe import RoutedExperts -from vllm.model_executor.layers.linear import ( - LinearBase, - LinearMethodBase, - UnquantizedLinearMethod, -) -from vllm.model_executor.layers.quantization import ( - QuantizationConfig, - QuantizationMethods, -) -from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead -from vllm.model_executor.parameter import ( - GroupQuantScaleParameter, - PackedvLLMParameter, - RowvLLMParameter, -) -from vllm.platforms import current_platform -from vllm.scalar_type import scalar_types - -if TYPE_CHECKING: - from vllm.model_executor.models.utils import WeightsMapper - -logger = init_logger(__name__) - - -class INCConfig(QuantizationConfig): - """Config class for Intel Neural Compressor (INC). - Repo: https://github.com/intel/neural-compressor - """ - - SUPPORTED_BITS = {2, 3, 4, 8} - SUPPORTED_DTYPES = {"int"} - SUPPORTED_FORMATS = {"auto_round:auto_gptq", "auto_round:auto_awq"} - SUPPORTED_BACKENDS = { - "auto", - "gptq", - "gptq:marlin", - "awq", - "awq:marlin", - "marlin", - } - - def __init__( - self, - weight_bits: int, - group_size: int, - sym: bool = True, - packing_format: str = "auto_round:auto_gptq", - block_name_to_quantize: str | list[str] | None = None, - extra_config: dict[str, Any] | None = None, - data_type: str = "int", - backend: str = "auto", - ) -> None: - super().__init__() - if weight_bits not in self.SUPPORTED_BITS: - raise ValueError( - f"Unsupported weight_bits: {weight_bits}, " - f"currently only support {self.SUPPORTED_BITS}." - ) - if data_type not in self.SUPPORTED_DTYPES: - raise ValueError( - f"Unsupported data_type: {data_type}," - f" currently only support {self.SUPPORTED_DTYPES}." - ) - if packing_format not in self.SUPPORTED_FORMATS: - raise ValueError( - f"Unsupported packing_format: {packing_format}, " - f"currently only support {self.SUPPORTED_FORMATS}." - ) - if backend not in self.SUPPORTED_BACKENDS: - raise ValueError( - f"Unsupported backend: {backend}, " - f"currently only support {self.SUPPORTED_BACKENDS}." - ) - - self.weight_bits = weight_bits - self.group_size = group_size - self.sym = sym - self.packing_format = packing_format - self.block_name_to_quantize = ( - block_name_to_quantize.split(",") - if isinstance(block_name_to_quantize, str) - else block_name_to_quantize - ) - self.extra_config = extra_config - self.data_type = data_type - self.backend = backend - self.pack_factor = Fraction(32, weight_bits) - - def __repr__(self) -> str: - return ( - f"INCConfig(weight_bits={self.weight_bits}, " - f"group_size={self.group_size}, sym={self.sym})" - ) - - @classmethod - def get_name(cls) -> QuantizationMethods: - return "inc" - - @classmethod - def get_supported_act_dtypes(cls) -> list[torch.dtype]: - return [torch.half, torch.bfloat16] - - @classmethod - def get_min_capability(cls) -> int: - return 60 - - @classmethod - def get_config_filenames(cls) -> list[str]: - return ["quantization_config.json"] - - @classmethod - def from_config(cls, config: dict[str, Any]) -> "INCConfig": - return cls( - weight_bits=cls.get_from_keys(config, ["bits"]), - group_size=cls.get_from_keys(config, ["group_size"]), - sym=cls.get_from_keys(config, ["sym"]), - packing_format=cls.get_from_keys_or( - config, ["packing_format"], "auto_round:auto_gptq" - ), - block_name_to_quantize=cls.get_from_keys_or( - config, ["block_name_to_quantize", "to_quant_block_names"], None - ), - extra_config=cls.get_from_keys_or(config, ["extra_config"], None), - data_type=cls.get_from_keys_or(config, ["data_type"], "int"), - backend=cls.get_from_keys_or(config, ["backend", "vllm_backend"], "auto"), - ) - - def get_layer_config(self, layer, layer_name: str): - def get_config(name: str, quantized: bool = True): - if not self.extra_config: - return ( - self.weight_bits if quantized else 16, - self.group_size if quantized else -1, - self.sym if quantized else True, - ) - - # exact match first - if name in self.extra_config: - cfg = self.extra_config[name] - return ( - cfg.get("bits", self.weight_bits if quantized else 16), - cfg.get("group_size", self.group_size if quantized else -1), - cfg.get("sym", self.sym if quantized else True), - ) - - REGEX_SPECIAL_CHARS = set(r"*+?^$()[]{}|\\") - for pattern, cfg in self.extra_config.items(): - if not isinstance(pattern, str) or not any( - c in REGEX_SPECIAL_CHARS for c in pattern - ): - continue - - try: - if re.search(re.compile(pattern), name) is not None: - return ( - cfg.get("bits", self.weight_bits if quantized else 16), - cfg.get("group_size", self.group_size if quantized else -1), - cfg.get("sym", self.sym if quantized else True), - ) - except re.error: - # Invalid regex, ignore. - continue - - return ( - self.weight_bits if quantized else 16, - self.group_size if quantized else -1, - self.sym if quantized else True, - ) - - # 1. Exact match from config - if self.extra_config and layer_name in self.extra_config: - return get_config(layer_name) - - # 2. Determine whether layer should be quantized - quantized = not isinstance(layer, ParallelLMHead) - if self.block_name_to_quantize: - quantized = any( - layer_name.startswith(name) for name in self.block_name_to_quantize - ) - - # 3. Handle fused MoE - if self.extra_config and "fusedmoe" in layer.__class__.__name__.lower(): - moe_configs = [ - get_config(name, quantized) - for name in self.extra_config - if name.startswith(layer_name) - ] - if moe_configs: - if len(set(moe_configs)) == 1: - return moe_configs[0] - raise ValueError( - f"Fused MoE layer '{layer_name}' requires " - f"consistent quant config for all sub-layers" - ) - - # 4. Handle fused QKV or other patterns - if self.extra_config: - for fusion_key, sub_keys in self.packed_modules_mapping.items(): - if fusion_key in layer_name and layer_name.count(fusion_key) == 1: - sub_names = [ - layer_name.replace(fusion_key, sub_key) for sub_key in sub_keys - ] - sub_configs = [get_config(name, quantized) for name in sub_names] - if len(set(sub_configs)) == 1: - return sub_configs[0] - raise ValueError( - f"Fused module '{layer_name}' requires " - f"consistent quant config for {sub_names}" - ) - - # 5. Fallback or try a regular expression match - return get_config(layer_name, quantized) - - def check_quantized(self, weight_bits: int) -> bool: - return weight_bits < 16 - - def apply_vllm_mapper(self, hf_to_vllm_mapper: "WeightsMapper"): - if self.block_name_to_quantize is not None: - self.block_name_to_quantize = hf_to_vllm_mapper.apply_list( - self.block_name_to_quantize - ) - if self.extra_config is not None: - self.extra_config = hf_to_vllm_mapper.apply_dict(self.extra_config) - - def apply_awq_quant_layer(self, layer, prefix: str, backend: str = "auto"): - from vllm.model_executor.layers.quantization.utils.marlin_utils import ( - check_marlin_supported, - check_moe_marlin_supports_layer, - ) - - weight_bits, group_size, sym = self.get_layer_config(layer, prefix) - if not self.check_quantized(weight_bits): - if isinstance(layer, (LinearBase, ParallelLMHead)): - return UnquantizedLinearMethod() - else: - return None - - logger.debug( - "[%s] Type: %s, Bits: %s, Group Size: %s, Sym: %s", - prefix, - layer.__class__.__name__, - weight_bits, - group_size, - sym, - ) - if backend == "auto" or "marlin" in backend: - AWQ_TYPE_MAP = { - 4: scalar_types.uint4, - 8: scalar_types.uint8, - } - use_marlin = (weight_bits in AWQ_TYPE_MAP) and check_marlin_supported( - AWQ_TYPE_MAP[weight_bits], group_size, not sym - ) - - if isinstance(layer, RoutedExperts): - use_marlin = use_marlin and check_moe_marlin_supports_layer( - layer, group_size - ) - - else: - use_marlin = False - if use_marlin: - from vllm.model_executor.layers.quantization.awq_marlin import ( - AWQMarlinConfig, - AWQMarlinLinearMethod, - AWQMarlinMoEMethod, - ) - - quant_args_marlin = AWQMarlinConfig( - weight_bits=weight_bits, - group_size=group_size, - zero_point=not sym, - lm_head_quantized=False, - full_config={}, - modules_to_not_convert=[], - ) - else: - from vllm.model_executor.layers.quantization.awq import ( - AWQConfig, - AWQLinearMethod, - ) - - quant_args = AWQConfig( - weight_bits=weight_bits, - group_size=group_size, - zero_point=not sym, - ) - - if isinstance(layer, RoutedExperts): - if use_marlin: - return AWQMarlinMoEMethod(quant_args_marlin, layer.moe_config) - from vllm.model_executor.layers.quantization.moe_wna16 import MoeWNA16Config - - config = { - "quant_method": "awq", - "bits": weight_bits, - "group_size": group_size, - "zero_point": not sym, - "lm_head": False, - } - return MoeWNA16Config.from_config(config).get_quant_method(layer, prefix) - - if isinstance(layer, (LinearBase, ParallelLMHead)): - if use_marlin: - return AWQMarlinLinearMethod(quant_args_marlin) - else: - return AWQLinearMethod(quant_args) - return None - - def apply_gptq_quant_layer(self, layer, prefix: str, backend: str = "auto"): - from vllm.model_executor.layers.quantization.utils.marlin_utils import ( - check_marlin_supported, - check_moe_marlin_supports_layer, - ) - - weight_bits, group_size, sym = self.get_layer_config(layer, prefix) - if not self.check_quantized(weight_bits): - if isinstance(layer, (LinearBase, ParallelLMHead)): - return UnquantizedLinearMethod() - else: - return None - - logger.debug( - "[%s] Type: %s, Bits: %s, Group Size: %s, Sym: %s", - prefix, - layer.__class__.__name__, - weight_bits, - group_size, - sym, - ) - if backend == "auto" or "marlin" in backend: - GPTQ_TYPE_MAP = { - (4, True): scalar_types.uint4b8, - (8, True): scalar_types.uint8b128, - } - use_marlin = (weight_bits, sym) in GPTQ_TYPE_MAP and check_marlin_supported( - GPTQ_TYPE_MAP[(weight_bits, sym)], group_size, has_zp=not sym - ) - if isinstance(layer, RoutedExperts): - use_marlin = use_marlin and check_moe_marlin_supports_layer( - layer, group_size - ) - else: - use_marlin = False - if use_marlin: - from vllm.model_executor.layers.quantization.auto_gptq import ( - AutoGPTQConfig, - AutoGPTQLinearMethod, - AutoGPTQMoEMethod, - ) - - quant_args_marlin = AutoGPTQConfig( - weight_bits=weight_bits, - group_size=group_size, - is_sym=sym, - lm_head_quantized=False, - desc_act=False, - dynamic={}, - full_config={}, - ) - - if isinstance(layer, RoutedExperts): - if use_marlin: - return AutoGPTQMoEMethod(quant_args_marlin, layer.moe_config) - else: - from vllm.model_executor.layers.quantization.moe_wna16 import ( - MoeWNA16Config, - ) - - config = { - "quant_method": "gptq", - "bits": weight_bits, - "group_size": group_size, - "sym": sym, - "lm_head": False, - } - return MoeWNA16Config.from_config(config).get_quant_method( - layer, prefix - ) - - if isinstance(layer, (LinearBase, ParallelLMHead)): - if use_marlin: - return AutoGPTQLinearMethod(quant_args_marlin) - else: - raise NotImplementedError( - f"INC quantization with bits={weight_bits}, sym={sym} " - "is not supported. Only 4-bit and 8-bit symmetric " - "quantization is supported with Marlin kernels." - ) - - return None - - def apply_xpu_w4a16_quant_layer(self, layer, prefix: str): - weight_bits, group_size, sym = self.get_layer_config(layer, prefix) - - if not self.check_quantized(weight_bits): - if isinstance(layer, (LinearBase, ParallelLMHead)): - return UnquantizedLinearMethod() - else: - return None - - if weight_bits != 4: - raise NotImplementedError( - f"INC on XPU only supports 4-bit quantization, " - f"got weight_bits={weight_bits}." - ) - if not sym: - raise NotImplementedError( - "INC W4A16 on XPU only supports symmetric quantization for now." - ) - - if isinstance(layer, (LinearBase, ParallelLMHead)): - is_ark_available, ark_error, _, _ = get_ark_state() - if is_ark_available: - return INCARKLinearMethod( - weight_bits=weight_bits, - group_size=group_size, - sym=sym, - ) - - logger.debug( - "ARK backend is unavailable for layer %s; " - "falling back to the default XPU INC path. Error: %s", - prefix, - ark_error or "unknown error", - ) - - return INCXPULinearMethod( - weight_bits=weight_bits, - group_size=group_size, - sym=sym, - ) - return None - - def apply_cpu_w4a16_quant_layer(self, layer, prefix: str): - weight_bits, group_size, sym = self.get_layer_config(layer, prefix) - if not self.check_quantized(weight_bits): - if isinstance(layer, (LinearBase, ParallelLMHead)): - return UnquantizedLinearMethod() - else: - return None - - if weight_bits != 4: - raise NotImplementedError( - f"INC on CPU only supports 4-bit quantization, " - f"got weight_bits={weight_bits}." - ) - if not sym: - raise NotImplementedError( - "INC W4A16 on CPU only supports symmetric quantization for now." - ) - if isinstance(layer, (LinearBase, ParallelLMHead)): - is_ark_available, ark_error, _, _ = get_ark_state() - if is_ark_available: - return INCARKLinearMethod( - weight_bits=weight_bits, - group_size=group_size, - sym=sym, - ) - - logger.debug( - "ARK backend is unavailable for layer %s; " - "falling back to the default CPU INC path. Error: %s", - prefix, - ark_error or "unknown error", - ) - - return self.apply_gptq_quant_layer(layer, prefix) - return None - - def get_quant_method(self, layer: torch.nn.Module, prefix: str): - if prefix and self.extra_config: - for layer_name in self.extra_config: - if ( - layer_name == prefix or layer_name == f"model.{prefix}" - ) and self.extra_config[layer_name].get("bits", 16) >= 16: - return UnquantizedLinearMethod() - - if current_platform.is_xpu(): - return self.apply_xpu_w4a16_quant_layer(layer, prefix) - is_gptq = "gptq" in self.packing_format or "gptq" in self.backend - if current_platform.is_cpu() and is_gptq: - return self.apply_cpu_w4a16_quant_layer(layer, prefix) - if is_gptq: - return self.apply_gptq_quant_layer(layer, prefix) - if "awq" in self.packing_format or "awq" in self.backend: - return self.apply_awq_quant_layer(layer, prefix) - - raise NotImplementedError( - f"Unsupported quantization configuration for layer '{prefix}'. " - f"Platform: CPU={current_platform.is_cpu()}. " - f"Platform: XPU={current_platform.is_xpu()}. " - f"Format: {self.packing_format}, Backend: {self.backend}." - ) - - @classmethod - def override_quantization_method( - cls, hf_quant_cfg, user_quant, hf_config=None - ) -> "QuantizationMethods | None": - """Override the `auto-round` method to `inc`.""" - is_auto_round_format = hf_quant_cfg.get("quant_method", None) == "auto-round" - if is_auto_round_format: - return cls.get_name() - return None - - -class INCXPULinearBase(LinearMethodBase): - def __init__(self, weight_bits: int, group_size: int, sym: bool): - self.weight_bits = weight_bits - self.group_size = group_size - self.sym = sym - self.pack_factor = 32 // weight_bits - - def _create_inc_weights( - self, - layer: torch.nn.Module, - input_size_per_partition: int, - output_partition_sizes: list[int], - params_dtype: torch.dtype, - weight_loader: Any, - group_size: int, - pack_factor: int, - ) -> None: - output_size_per_partition = sum(output_partition_sizes) - scales_and_zp_size = input_size_per_partition // group_size - - qweight = PackedvLLMParameter( - data=torch.empty( - input_size_per_partition // pack_factor, - output_size_per_partition, - dtype=torch.int32, - ), - input_dim=0, - output_dim=1, - packed_dim=0, - packed_factor=pack_factor, - weight_loader=weight_loader, - ) - - scales = GroupQuantScaleParameter( - data=torch.empty( - scales_and_zp_size, - output_size_per_partition, - dtype=params_dtype, - ), - input_dim=0, - output_dim=1, - weight_loader=weight_loader, - ) - - qzeros = PackedvLLMParameter( - data=torch.empty( - scales_and_zp_size, - output_size_per_partition // pack_factor, - dtype=torch.int32, - ), - input_dim=0, - output_dim=1, - packed_dim=1, - packed_factor=pack_factor, - weight_loader=weight_loader, - ) - - layer.register_parameter("qweight", qweight) - layer.register_parameter("scales", scales) - layer.register_parameter("qzeros", qzeros) - - g_idx = RowvLLMParameter( - data=torch.tensor( - [i // group_size for i in range(input_size_per_partition)], - dtype=torch.int32, - ), - input_dim=0, - weight_loader=weight_loader, - ) - layer.register_parameter("g_idx", g_idx) - - def create_weights( - self, - layer: torch.nn.Module, - input_size_per_partition: int, - output_partition_sizes: list[int], - input_size: int, - output_size: int, - params_dtype: torch.dtype, - **extra_weight_attrs, - ): - self._create_inc_weights( - layer=layer, - input_size_per_partition=input_size_per_partition, - output_partition_sizes=output_partition_sizes, - params_dtype=params_dtype, - weight_loader=extra_weight_attrs.get("weight_loader"), - group_size=self.group_size, - pack_factor=self.pack_factor, - ) - - -@lru_cache(maxsize=1) -def get_ark_state() -> tuple[bool, str | None, Any | None, Any | None]: - """Return ARK availability, error details, cached instance, and QuantLinear.""" - try: - import auto_round_kernel - from auto_round_kernel.qlinear import QuantLinear - - logger.info("Successfully imported auto_round_kernel.") - except ImportError as error: - return False, str(error), None, None - - ark_loader = getattr(auto_round_kernel, "_ark_instance", None) - if not callable(ark_loader): - return False, "auto_round_kernel does not expose _ark_instance().", None, None - - try: - ark_instance = ark_loader() - except Exception as error: - return False, str(error), None, None - - if ark_instance is None: - return False, "auto_round_kernel._ark_instance() returned None.", None, None - - return True, None, ark_instance, QuantLinear - - -class INCXPULinearMethod(INCXPULinearBase): - """XPU linear method for INC w4a16 GPTQ quantization (symmetric only). - - Repacks GPTQ weights from [in_packed, out] to oneDNN [out, in_packed] - layout and calls torch.ops._xpu_C.int4_gemm_w4a16. - - GPTQ format: qweight [in_packed, out] with sequential nibble order. - - Note: Asymmetric quantization (sym=false) is not for now. - - FIXME(yiliu30): Refine the implementation to reuse XPUwNa16LinearKernel. - """ - - def process_weights_after_loading(self, layer: torch.nn.Module) -> None: - """Repack GPTQ weights into kernel-ready NT layout.""" - device = layer.qweight.data.device - - # oneDNN int4 kernel requires strides[0]==1 ("NT format"), but GPTQ - # checkpoint is [K_packed, N] contiguous with strides (N, 1). - # Two transposes are needed — neither alone can achieve this: - # 1. .t().contiguous() → [N, K_packed] contiguous in memory - # 2. .t() → [K_packed, N] view with strides (1, K_packed) - # The result has the same logical shape but strides[0]==1 as required. - qweight_ct = layer.qweight.data.t().contiguous() - layer.qweight = Parameter(qweight_ct.t(), requires_grad=False) - - # Scales: [num_groups, out] — no change needed - layer.scales = Parameter(layer.scales.data, requires_grad=False) - - # Symmetric: GPTQ v1 stores qzeros=7, effective zp = 7+1 = 8 - # Kernel expects int8 scalar = 8 - layer.qzeros = Parameter( - torch.tensor([8], dtype=torch.int8, device=device), - requires_grad=False, - ) - - def apply( - self, - layer: torch.nn.Module, - x: torch.Tensor, - bias: torch.Tensor | None = None, - ) -> torch.Tensor: - # qweight is already in NT layout [K_packed, N] (strides (1, K_packed)) - # from process_weights_after_loading — pass directly to kernel. - out_shape = x.shape[:-1] + (layer.qweight.shape[1],) - reshaped_x = x.reshape(-1, x.shape[-1]) - out = torch.ops._xpu_C.int4_gemm_w4a16( - reshaped_x, - layer.qweight, - bias, - layer.scales, - layer.qzeros, - self.group_size, - None, # g_idx not needed: desc_act is always False for INC models - ) - return out.reshape(out_shape) - - -class INCARKLinearMethod(INCXPULinearBase): - """XPU & CPU w4a16 linear method for INC quantization utilizing the ARK backend. - - See: https://github.com/intel/auto-round/blob/main/auto_round_extension/ark/README.md - - Repacks GPTQ/INC weights into ARK's layout. - """ - - def __init__(self, weight_bits: int, group_size: int, sym: bool): - super().__init__(weight_bits=weight_bits, group_size=group_size, sym=sym) - - is_available, error_str, _, quant_linear_cls = get_ark_state() - if not is_available or quant_linear_cls is None: - reason = error_str or "unknown error" - raise ImportError(f"Failed to import auto_round_kernel. {reason}") - - self.QuantLinear = quant_linear_cls - - def create_weights( - self, - layer: torch.nn.Module, - input_size_per_partition: int, - output_partition_sizes: list[int], - input_size: int, - output_size: int, - params_dtype: torch.dtype, - **extra_weight_attrs, - ): - super().create_weights( - layer=layer, - input_size_per_partition=input_size_per_partition, - output_partition_sizes=output_partition_sizes, - input_size=input_size, - output_size=output_size, - params_dtype=params_dtype, - **extra_weight_attrs, - ) - layer.in_features = input_size_per_partition - layer.out_features = sum(output_partition_sizes) - layer.params_dtype = params_dtype - - def process_weights_after_loading(self, layer: torch.nn.Module) -> None: - if hasattr(layer, "input_size_per_partition"): - in_features = layer.input_size_per_partition - elif hasattr(layer, "input_size"): - in_features = layer.input_size - else: - raise AttributeError("Cannot determine in_features for layer.") - - if hasattr(layer, "output_partition_sizes"): - out_features = sum(layer.output_partition_sizes) - elif hasattr(layer, "output_size_per_partition"): - out_features = layer.output_size_per_partition - elif hasattr(layer, "output_size"): - out_features = layer.output_size - else: - out_features = layer.scales.shape[-1] - - ark_linear = self.QuantLinear( - bits=self.weight_bits, - group_size=self.group_size, - sym=self.sym, - in_features=in_features, - out_features=out_features, - bias=layer.bias is not None, - weight_dtype=layer.params_dtype, - ) - - ark_linear.to(layer.qweight.device) - - with torch.no_grad(): - ark_linear.qweight.copy_(layer.qweight.detach()) - - if hasattr(layer, "qzeros") and layer.qzeros is not None: - ark_linear.qzeros.copy_(layer.qzeros.detach()) - else: - ark_linear.qzeros = None - - ark_linear.scales.copy_(layer.scales.detach()) - - if hasattr(layer, "bias") and layer.bias is not None: - ark_linear.bias.copy_(layer.bias.detach()) - - ark_linear.post_init() - - layer.ark_linear = ark_linear - - del layer.qweight - if hasattr(layer, "qzeros"): - del layer.qzeros - del layer.scales - - def apply( - self, - layer: torch.nn.Module, - x: torch.Tensor, - bias: torch.Tensor | None = None, - ) -> torch.Tensor: - return layer.ark_linear.forward(x) diff --git a/vllm/model_executor/layers/quantization/inc/__init__.py b/vllm/model_executor/layers/quantization/inc/__init__.py new file mode 100644 index 00000000000..e7d9f4707f5 --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/__init__.py @@ -0,0 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from .inc import INCConfig + +__all__ = ["INCConfig"] diff --git a/vllm/model_executor/layers/quantization/inc/config_parser.py b/vllm/model_executor/layers/quantization/inc/config_parser.py new file mode 100644 index 00000000000..603b80b7cd0 --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/config_parser.py @@ -0,0 +1,188 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import regex as re + +from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead + +if TYPE_CHECKING: + import torch + + from .inc import INCConfig + + +@dataclass(frozen=True) +class INCLayerConfig: + bits: int + group_size: int + sym: bool + packing_format: str + backend: str + data_type: str + quantized: bool + + @property + def is_gptq(self) -> bool: + return "gptq" in self.packing_format or "gptq" in self.backend + + @property + def is_awq(self) -> bool: + return "awq" in self.packing_format or "awq" in self.backend + + @property + def is_wna16_int(self) -> bool: + return self.data_type == "int" and self.quantized + + @property + def is_mxfp4(self) -> bool: + return self.data_type == "mx_fp" and self.bits == 4 + + @property + def is_mxfp8(self) -> bool: + return self.data_type == "mx_fp" and self.bits == 8 + + +class INCConfigParser: + def __init__(self, config: "INCConfig") -> None: + self._config = config + + def resolve(self, layer: "torch.nn.Module", layer_name: str) -> INCLayerConfig: + bits, group_size, sym = self._resolve_raw(layer, layer_name) + return INCLayerConfig( + bits=bits, + group_size=group_size, + sym=sym, + packing_format=self._config.packing_format, + backend=self._config.backend, + data_type=self._config.data_type, + quantized=bits < 16, + ) + + def get_layer_config( + self, layer: "torch.nn.Module", layer_name: str + ) -> tuple[int, int, bool]: + layer_config = self.resolve(layer, layer_name) + return layer_config.bits, layer_config.group_size, layer_config.sym + + def _resolve_raw( + self, layer: "torch.nn.Module", layer_name: str + ) -> tuple[int, int, bool]: + REGEX_SPECIAL_CHARS = set(r"*+?^$()[]{}|\\") + + def is_explicitly_configured(name: str) -> bool: + """Return True if *name* has an explicit entry in extra_config, + either via exact key match or via a regex pattern key.""" + if not self._config.extra_config: + return False + if name in self._config.extra_config: + return True + for pattern in self._config.extra_config: + if not isinstance(pattern, str) or not any( + c in REGEX_SPECIAL_CHARS for c in pattern + ): + continue + try: + if re.search(re.compile(pattern), name) is not None: + return True + except re.error: + continue + return False + + def get_config(name: str, quantized: bool = True) -> tuple[int, int, bool]: + if not self._config.extra_config: + return ( + self._config.weight_bits if quantized else 16, + self._config.group_size if quantized else -1, + self._config.sym if quantized else True, + ) + + if name in self._config.extra_config: + cfg = self._config.extra_config[name] + return ( + cfg.get("bits", self._config.weight_bits if quantized else 16), + cfg.get( + "group_size", + self._config.group_size if quantized else -1, + ), + cfg.get("sym", self._config.sym if quantized else True), + ) + + regex_special_chars = set(r"*+?^$()[]{}|\\") + for pattern, cfg in self._config.extra_config.items(): + if not isinstance(pattern, str) or not any( + c in regex_special_chars for c in pattern + ): + continue + + try: + if re.search(re.compile(pattern), name) is not None: + return ( + cfg.get( + "bits", + self._config.weight_bits if quantized else 16, + ), + cfg.get( + "group_size", + self._config.group_size if quantized else -1, + ), + cfg.get("sym", self._config.sym if quantized else True), + ) + except re.error: + continue + + return ( + self._config.weight_bits if quantized else 16, + self._config.group_size if quantized else -1, + self._config.sym if quantized else True, + ) + + if self._config.extra_config and layer_name in self._config.extra_config: + return get_config(layer_name) + + quantized = not isinstance(layer, ParallelLMHead) + if self._config.block_name_to_quantize: + quantized = any( + layer_name.startswith(name) + for name in self._config.block_name_to_quantize + ) + + if self._config.extra_config and "fusedmoe" in layer.__class__.__name__.lower(): + moe_configs = [ + get_config(name, quantized) + for name in self._config.extra_config + if name.startswith(layer_name) + ] + if moe_configs: + if len(set(moe_configs)) == 1: + return moe_configs[0] + raise ValueError( + f"Fused MoE layer '{layer_name}' requires " + f"consistent quant config for all sub-layers" + ) + + if self._config.extra_config: + for fusion_key, sub_keys in self._config.packed_modules_mapping.items(): + if fusion_key in layer_name and layer_name.count(fusion_key) == 1: + sub_names = [ + layer_name.replace(fusion_key, sub_key) for sub_key in sub_keys + ] + # Only trigger if at least one sub_name is explicitly + # configured in extra_config (via exact match or regex). + # This prevents false matches when a short fusion_key + # (e.g. "qkv") is merely a substring of a longer layer + # name (e.g. "in_proj_qkvz") and none of the generated + # sub_names are actually configured. + if not any(is_explicitly_configured(n) for n in sub_names): + continue + sub_configs = [get_config(name, quantized) for name in sub_names] + if len(set(sub_configs)) == 1: + return sub_configs[0] + raise ValueError( + f"Fused module '{layer_name}' requires " + f"consistent quant config for {sub_names}" + ) + + return get_config(layer_name, quantized) diff --git a/vllm/model_executor/layers/quantization/inc/inc.py b/vllm/model_executor/layers/quantization/inc/inc.py new file mode 100644 index 00000000000..86fa7cefcfc --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/inc.py @@ -0,0 +1,192 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from fractions import Fraction +from typing import TYPE_CHECKING, Any + +import torch + +from vllm.logger import init_logger +from vllm.model_executor.layers.fused_moe import ( + RoutedExperts, + UnquantizedFusedMoEMethod, +) +from vllm.model_executor.layers.linear import ( + LinearBase, + UnquantizedLinearMethod, +) +from vllm.model_executor.layers.quantization import ( + QuantizationConfig, + QuantizationMethods, +) +from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead + +from .config_parser import INCConfigParser + +if TYPE_CHECKING: + from vllm.model_executor.models.utils import WeightsMapper + +logger = init_logger(__name__) + + +class INCConfig(QuantizationConfig): + """Config class for Intel Neural Compressor (INC). + Repo: https://github.com/intel/neural-compressor + """ + + SUPPORTED_BITS = {2, 3, 4, 8} + SUPPORTED_DTYPES = {"int"} + SUPPORTED_FORMATS = {"auto_round:auto_gptq", "auto_round:auto_awq"} + SUPPORTED_BACKENDS = { + "auto", + "gptq", + "gptq:marlin", + "awq", + "awq:marlin", + "marlin", + } + + def __init__( + self, + weight_bits: int, + group_size: int, + sym: bool = True, + packing_format: str = "auto_round:auto_gptq", + block_name_to_quantize: str | list[str] | None = None, + extra_config: dict[str, Any] | None = None, + data_type: str = "int", + backend: str = "auto", + ) -> None: + super().__init__() + if weight_bits not in self.SUPPORTED_BITS: + raise ValueError( + f"Unsupported weight_bits: {weight_bits}, " + f"currently only support {self.SUPPORTED_BITS}." + ) + if data_type not in self.SUPPORTED_DTYPES: + raise ValueError( + f"Unsupported data_type: {data_type}," + f" currently only support {self.SUPPORTED_DTYPES}." + ) + if packing_format not in self.SUPPORTED_FORMATS: + raise ValueError( + f"Unsupported packing_format: {packing_format}, " + f"currently only support {self.SUPPORTED_FORMATS}." + ) + if backend not in self.SUPPORTED_BACKENDS: + raise ValueError( + f"Unsupported backend: {backend}, " + f"currently only support {self.SUPPORTED_BACKENDS}." + ) + + self.weight_bits = weight_bits + self.group_size = group_size + self.sym = sym + self.packing_format = packing_format + self.block_name_to_quantize = ( + block_name_to_quantize.split(",") + if isinstance(block_name_to_quantize, str) + else block_name_to_quantize + ) + self.extra_config = extra_config + self.data_type = data_type + self.backend = backend + self.pack_factor = Fraction(32, weight_bits) + self.config_parser = INCConfigParser(self) + + def __repr__(self) -> str: + return ( + f"INCConfig(weight_bits={self.weight_bits}, " + f"group_size={self.group_size}, sym={self.sym})" + ) + + @classmethod + def get_name(cls) -> QuantizationMethods: + return "inc" + + @classmethod + def get_supported_act_dtypes(cls) -> list[torch.dtype]: + return [torch.half, torch.bfloat16] + + @classmethod + def get_min_capability(cls) -> int: + return 60 + + @classmethod + def get_config_filenames(cls) -> list[str]: + return ["quantization_config.json"] + + @classmethod + def from_config(cls, config: dict[str, Any]) -> "INCConfig": + return cls( + weight_bits=cls.get_from_keys(config, ["bits"]), + group_size=cls.get_from_keys(config, ["group_size"]), + sym=cls.get_from_keys(config, ["sym"]), + packing_format=cls.get_from_keys_or( + config, ["packing_format"], "auto_round:auto_gptq" + ), + block_name_to_quantize=cls.get_from_keys_or( + config, ["block_name_to_quantize", "to_quant_block_names"], None + ), + extra_config=cls.get_from_keys_or(config, ["extra_config"], None), + data_type=cls.get_from_keys_or(config, ["data_type"], "int"), + backend=cls.get_from_keys_or(config, ["backend", "vllm_backend"], "auto"), + ) + + def get_layer_config(self, layer, layer_name: str): + return self.config_parser.get_layer_config(layer, layer_name) + + def apply_vllm_mapper(self, hf_to_vllm_mapper: "WeightsMapper"): + if self.block_name_to_quantize is not None: + self.block_name_to_quantize = hf_to_vllm_mapper.apply_list( + self.block_name_to_quantize + ) + if self.extra_config is not None: + self.extra_config = hf_to_vllm_mapper.apply_dict(self.extra_config) + + def get_quant_method(self, layer: torch.nn.Module, prefix: str): + from .schemes.factory import resolve_scheme + + # Match original: check model.-prefixed names for unquantized layers + if prefix and self.extra_config: + for layer_name in self.extra_config: + if ( + layer_name == prefix or layer_name == f"model.{prefix}" + ) and self.extra_config[layer_name].get("bits", 16) >= 16: + if isinstance(layer, RoutedExperts): + return UnquantizedFusedMoEMethod(layer.moe_config) + return UnquantizedLinearMethod() + + layer_config = self.config_parser.resolve(layer, prefix) + if not layer_config.quantized: + if isinstance(layer, (LinearBase, ParallelLMHead)): + return UnquantizedLinearMethod() + if isinstance(layer, RoutedExperts): + return UnquantizedFusedMoEMethod(layer.moe_config) + return None + + logger.debug( + "[%s] Type: %s, Bits: %s, Group Size: %s, Sym: %s", + prefix, + layer.__class__.__name__, + layer_config.bits, + layer_config.group_size, + layer_config.sym, + ) + + scheme = resolve_scheme(layer_config) + if isinstance(layer, (LinearBase, ParallelLMHead)): + return scheme.get_linear_method(self, layer, prefix, layer_config) + if isinstance(layer, RoutedExperts): + return scheme.get_moe_method(self, layer, prefix, layer_config) + return None + + @classmethod + def override_quantization_method( + cls, hf_quant_cfg, user_quant, hf_config=None + ) -> "QuantizationMethods | None": + """Override the `auto-round` method to `inc`.""" + is_auto_round_format = hf_quant_cfg.get("quant_method", None) == "auto-round" + if is_auto_round_format: + return cls.get_name() + return None diff --git a/vllm/model_executor/layers/quantization/inc/inc_linear.py b/vllm/model_executor/layers/quantization/inc/inc_linear.py new file mode 100644 index 00000000000..9917a70194d --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/inc_linear.py @@ -0,0 +1,47 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from typing import TYPE_CHECKING + +import torch + +from vllm.model_executor.layers.linear import LinearMethodBase + +if TYPE_CHECKING: + from .schemes.inc_scheme import INCLinearScheme + + +class INCLinearMethod(LinearMethodBase): + def __init__(self, scheme: "INCLinearScheme") -> None: + self.scheme = scheme + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + return self.scheme.create_weights( + layer=layer, + input_size_per_partition=input_size_per_partition, + output_partition_sizes=output_partition_sizes, + input_size=input_size, + output_size=output_size, + params_dtype=params_dtype, + **extra_weight_attrs, + ) + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + return self.scheme.process_weights_after_loading(layer) + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + return self.scheme.apply_weights(layer, x, bias) diff --git a/vllm/model_executor/layers/quantization/inc/schemes/__init__.py b/vllm/model_executor/layers/quantization/inc/schemes/__init__.py new file mode 100644 index 00000000000..ea6c0a00d86 --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/schemes/__init__.py @@ -0,0 +1,13 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from .factory import resolve_scheme +from .inc_scheme import INCLinearScheme, INCScheme +from .inc_wna16_scheme import INCWna16Scheme + +__all__ = [ + "INCScheme", + "INCLinearScheme", + "INCWna16Scheme", + "resolve_scheme", +] diff --git a/vllm/model_executor/layers/quantization/inc/schemes/factory.py b/vllm/model_executor/layers/quantization/inc/schemes/factory.py new file mode 100644 index 00000000000..4ae85ed9a83 --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/schemes/factory.py @@ -0,0 +1,22 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from ..config_parser import INCLayerConfig + from .inc_scheme import INCScheme + + +def resolve_scheme(layer_config: "INCLayerConfig") -> "INCScheme": + from .inc_wna16_scheme import INCWna16Scheme + + scheme_list: list[type[INCScheme]] = [ + INCWna16Scheme, + ] + + for scheme_cls in scheme_list: + if scheme_cls.can_handle(layer_config): + return scheme_cls() + + raise NotImplementedError(f"No INC scheme found for layer config: {layer_config}") diff --git a/vllm/model_executor/layers/quantization/inc/schemes/inc_scheme.py b/vllm/model_executor/layers/quantization/inc/schemes/inc_scheme.py new file mode 100644 index 00000000000..b8bb263de05 --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/schemes/inc_scheme.py @@ -0,0 +1,104 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + import torch + + from vllm.model_executor.layers.fused_moe import FusedMoEMethodBase + from vllm.model_executor.layers.linear import LinearMethodBase + from vllm.model_executor.layers.quantization import QuantizationMethods + + from ..config_parser import INCLayerConfig + from ..inc import INCConfig + + +class INCScheme(ABC): + """One class per quant type. Single registration point for the factory. + + Each subclass defines: + - can_handle(): when does this scheme apply? + - get_linear_method(): required — how to quantize Linear layers + - get_moe_method(): optional — how to quantize MoE layers + - get_kvcache_method(): optional — how to quantize KV cache + + Schemes that don't support MoE/KVCache inherit the default raise. + """ + + @staticmethod + @abstractmethod + def can_handle(layer_config: "INCLayerConfig") -> bool: + raise NotImplementedError + + @abstractmethod + def get_linear_method( + self, + config: "INCConfig", + layer: "torch.nn.Module", + prefix: str, + layer_config: "INCLayerConfig", + ) -> "LinearMethodBase": + raise NotImplementedError + + def get_moe_method( + self, + config: "INCConfig", + layer: "torch.nn.Module", + prefix: str, + layer_config: "INCLayerConfig", + ) -> "FusedMoEMethodBase | None": + """Optional. Override if this scheme supports MoE. + Default raises NotImplementedError.""" + raise NotImplementedError( + f"{type(self).__name__} does not support MoE layers. " + f"Layer config: {layer_config}" + ) + + def get_kvcache_method( + self, + config: "INCConfig", + layer: "torch.nn.Module", + prefix: str, + layer_config: "INCLayerConfig", + ) -> "QuantizationMethods": + """Optional. Override if this scheme supports KV cache quantization. + Default raises NotImplementedError.""" + raise NotImplementedError( + f"{type(self).__name__} does not support KV cache quantization. " + f"Layer config: {layer_config}" + ) + + +class INCLinearScheme(ABC): + @classmethod + @abstractmethod + def get_min_capability(cls) -> int: + raise NotImplementedError + + @abstractmethod + def create_weights( + self, + layer: "torch.nn.Module", + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: "torch.dtype", + **extra_weight_attrs, + ) -> None: + raise NotImplementedError + + @abstractmethod + def process_weights_after_loading(self, layer: "torch.nn.Module") -> None: + raise NotImplementedError + + @abstractmethod + def apply_weights( + self, + layer: "torch.nn.Module", + x: "torch.Tensor", + bias: "torch.Tensor | None" = None, + ) -> "torch.Tensor": + raise NotImplementedError diff --git a/vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_linear.py b/vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_linear.py new file mode 100644 index 00000000000..e0ffc6ac287 --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_linear.py @@ -0,0 +1,399 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from functools import lru_cache +from typing import TYPE_CHECKING, Any + +import torch +from torch.nn.parameter import Parameter + +from vllm.logger import init_logger +from vllm.model_executor.layers.quantization.auto_gptq import AutoGPTQConfig +from vllm.model_executor.layers.quantization.awq import AWQConfig +from vllm.model_executor.layers.quantization.awq_marlin import AWQMarlinConfig +from vllm.model_executor.layers.quantization.utils.marlin_utils import ( + check_marlin_supported, +) +from vllm.model_executor.parameter import ( + GroupQuantScaleParameter, + PackedvLLMParameter, + RowvLLMParameter, +) +from vllm.scalar_type import scalar_types + +from .inc_scheme import INCLinearScheme + +logger = init_logger(__name__) + +if TYPE_CHECKING: + from ..config_parser import INCLayerConfig + + +@lru_cache(maxsize=1) +def get_ark_state() -> tuple[bool, str | None, Any | None, Any | None]: + """Return ARK availability, error details, cached module, and QuantLinear.""" + try: + import auto_round_kernel as ark + from auto_round_kernel.qlinear import QuantLinear + + logger.info("Successfully imported auto_round_kernel.") + except ImportError as error: + return False, str(error), None, None + + if getattr(ark, "cpu_lib", None) is None and getattr(ark, "xpu_lib", None) is None: + return ( + False, + "No ARK backend library is available.", + None, + None, + ) + logger.info("Successfully loaded auto_round_kernel backend library.") + + return True, None, ark, QuantLinear + + +class INCWNA16LinearScheme(INCLinearScheme): + def __init__(self, layer_config: "INCLayerConfig") -> None: + self.layer_config = layer_config + self.inner_method = self._build_inner_method() + + @classmethod + def get_min_capability(cls) -> int: + return 60 + + def _build_inner_method(self): + if self.layer_config.is_gptq: + return self._build_gptq_method() + if self.layer_config.is_awq: + return self._build_awq_method() + raise NotImplementedError( + f"WNA16 linear scheme does not support {self.layer_config}" + ) + + def _build_gptq_method(self): + gptq_type_map = { + (4, True): scalar_types.uint4b8, + (8, True): scalar_types.uint8b128, + } + use_marlin = ( + self.layer_config.backend == "auto" or "marlin" in self.layer_config.backend + ) and (self.layer_config.bits, self.layer_config.sym) in gptq_type_map + if use_marlin: + use_marlin = check_marlin_supported( + gptq_type_map[(self.layer_config.bits, self.layer_config.sym)], + self.layer_config.group_size, + has_zp=not self.layer_config.sym, + ) + + if use_marlin: + from vllm.model_executor.layers.quantization.auto_gptq import ( + AutoGPTQLinearMethod, + ) + + return AutoGPTQLinearMethod( + AutoGPTQConfig( + weight_bits=self.layer_config.bits, + group_size=self.layer_config.group_size, + desc_act=False, + is_sym=self.layer_config.sym, + lm_head_quantized=False, + dynamic={}, + full_config={}, + ) + ) + + raise NotImplementedError( + f"INC quantization with bits={self.layer_config.bits}, " + f"sym={self.layer_config.sym} is not supported. " + "Only 4-bit and 8-bit symmetric quantization is supported " + "with Marlin kernels." + ) + + def _build_awq_method(self): + awq_type_map = { + 4: scalar_types.uint4, + 8: scalar_types.uint8, + } + use_marlin = ( + self.layer_config.backend == "auto" or "marlin" in self.layer_config.backend + ) and self.layer_config.bits in awq_type_map + if use_marlin: + use_marlin = check_marlin_supported( + awq_type_map[self.layer_config.bits], + self.layer_config.group_size, + not self.layer_config.sym, + ) + + if use_marlin: + from vllm.model_executor.layers.quantization.awq_marlin import ( + AWQMarlinLinearMethod, + ) + + return AWQMarlinLinearMethod( + AWQMarlinConfig( + weight_bits=self.layer_config.bits, + group_size=self.layer_config.group_size, + zero_point=not self.layer_config.sym, + lm_head_quantized=False, + modules_to_not_convert=[], + full_config={}, + ) + ) + + from vllm.model_executor.layers.quantization.awq import AWQLinearMethod + + return AWQLinearMethod( + AWQConfig( + weight_bits=self.layer_config.bits, + group_size=self.layer_config.group_size, + zero_point=not self.layer_config.sym, + ) + ) + + def create_weights( + self, + layer: "torch.nn.Module", + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: "torch.dtype", + **extra_weight_attrs, + ) -> None: + return self.inner_method.create_weights( + layer=layer, + input_size_per_partition=input_size_per_partition, + output_partition_sizes=output_partition_sizes, + input_size=input_size, + output_size=output_size, + params_dtype=params_dtype, + **extra_weight_attrs, + ) + + def process_weights_after_loading(self, layer: "torch.nn.Module") -> None: + return self.inner_method.process_weights_after_loading(layer) + + def apply_weights( + self, + layer: "torch.nn.Module", + x: "torch.Tensor", + bias: "torch.Tensor | None" = None, + ) -> "torch.Tensor": + return self.inner_method.apply(layer, x, bias) + + +class INCXPULinearBase(INCLinearScheme): + def __init__(self, layer_config: "INCLayerConfig") -> None: + self.weight_bits = layer_config.bits + self.group_size = layer_config.group_size + self.sym = layer_config.sym + self.pack_factor = 32 // self.weight_bits + + @classmethod + def get_min_capability(cls) -> int: + return 0 + + def _create_inc_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + params_dtype: torch.dtype, + weight_loader: Any, + ) -> None: + output_size_per_partition = sum(output_partition_sizes) + scales_and_zp_size = input_size_per_partition // self.group_size + + qweight = PackedvLLMParameter( + data=torch.empty( + input_size_per_partition // self.pack_factor, + output_size_per_partition, + dtype=torch.int32, + ), + input_dim=0, + output_dim=1, + packed_dim=0, + packed_factor=self.pack_factor, + weight_loader=weight_loader, + ) + scales = GroupQuantScaleParameter( + data=torch.empty( + scales_and_zp_size, + output_size_per_partition, + dtype=params_dtype, + ), + input_dim=0, + output_dim=1, + weight_loader=weight_loader, + ) + qzeros = PackedvLLMParameter( + data=torch.empty( + scales_and_zp_size, + output_size_per_partition // self.pack_factor, + dtype=torch.int32, + ), + input_dim=0, + output_dim=1, + packed_dim=1, + packed_factor=self.pack_factor, + weight_loader=weight_loader, + ) + + layer.register_parameter("qweight", qweight) + layer.register_parameter("scales", scales) + layer.register_parameter("qzeros", qzeros) + + g_idx = RowvLLMParameter( + data=torch.tensor( + [i // self.group_size for i in range(input_size_per_partition)], + dtype=torch.int32, + ), + input_dim=0, + weight_loader=weight_loader, + ) + layer.register_parameter("g_idx", g_idx) + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ) -> None: + del input_size, output_size + self._create_inc_weights( + layer=layer, + input_size_per_partition=input_size_per_partition, + output_partition_sizes=output_partition_sizes, + params_dtype=params_dtype, + weight_loader=extra_weight_attrs.get("weight_loader"), + ) + + +class INCXPULinearMethod(INCXPULinearBase): + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + device = layer.qweight.data.device + + qweight_ct = layer.qweight.data.t().contiguous() + layer.qweight = Parameter(qweight_ct.t(), requires_grad=False) + layer.scales = Parameter(layer.scales.data, requires_grad=False) + layer.qzeros = Parameter( + torch.tensor([8], dtype=torch.int8, device=device), + requires_grad=False, + ) + + def apply_weights( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + out_shape = x.shape[:-1] + (layer.qweight.shape[1],) + reshaped_x = x.reshape(-1, x.shape[-1]) + out = torch.ops._xpu_C.int4_gemm_w4a16( + reshaped_x, + layer.qweight, + bias, + layer.scales, + layer.qzeros, + self.group_size, + None, + ) + return out.reshape(out_shape) + + +class INCARKLinearMethod(INCXPULinearBase): + def __init__(self, layer_config: "INCLayerConfig") -> None: + super().__init__(layer_config) + + is_available, error_str, _, quant_linear_cls = get_ark_state() + if not is_available or quant_linear_cls is None: + reason = error_str or "unknown error" + raise ImportError(f"Failed to import auto_round_kernel. {reason}") + + self.quant_linear_cls = quant_linear_cls + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ) -> None: + super().create_weights( + layer=layer, + input_size_per_partition=input_size_per_partition, + output_partition_sizes=output_partition_sizes, + input_size=input_size, + output_size=output_size, + params_dtype=params_dtype, + **extra_weight_attrs, + ) + layer.in_features = input_size_per_partition + layer.out_features = sum(output_partition_sizes) + layer.params_dtype = params_dtype + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + if hasattr(layer, "input_size_per_partition"): + in_features = layer.input_size_per_partition + elif hasattr(layer, "input_size"): + in_features = layer.input_size + else: + raise AttributeError("Cannot determine in_features for layer.") + + if hasattr(layer, "output_partition_sizes"): + out_features = sum(layer.output_partition_sizes) + elif hasattr(layer, "output_size_per_partition"): + out_features = layer.output_size_per_partition + elif hasattr(layer, "output_size"): + out_features = layer.output_size + else: + out_features = layer.scales.shape[-1] + + ark_linear = self.quant_linear_cls( + bits=self.weight_bits, + group_size=self.group_size, + sym=self.sym, + in_features=in_features, + out_features=out_features, + bias=layer.bias is not None, + weight_dtype=layer.params_dtype, + ) + ark_linear.to(layer.qweight.device) + + with torch.no_grad(): + ark_linear.qweight.copy_(layer.qweight.detach()) + if hasattr(layer, "qzeros") and layer.qzeros is not None: + ark_linear.qzeros.copy_(layer.qzeros.detach()) + else: + ark_linear.qzeros = None + ark_linear.scales.copy_(layer.scales.detach()) + if hasattr(layer, "bias") and layer.bias is not None: + ark_linear.bias.copy_(layer.bias.detach()) + + ark_linear.post_init() + layer.ark_linear = ark_linear + + del layer.qweight + if hasattr(layer, "qzeros"): + del layer.qzeros + del layer.scales + + def apply_weights( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + del bias + return layer.ark_linear.forward(x) + + +class INCXPUW4A16LinearScheme(INCXPULinearMethod): + pass diff --git a/vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py b/vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py new file mode 100644 index 00000000000..7b6c10de2a5 --- /dev/null +++ b/vllm/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py @@ -0,0 +1,201 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from typing import TYPE_CHECKING + +from vllm.logger import init_logger +from vllm.model_executor.layers.quantization.auto_gptq import AutoGPTQConfig +from vllm.model_executor.layers.quantization.awq_marlin import AWQMarlinConfig +from vllm.platforms import current_platform +from vllm.scalar_type import scalar_types + +from ..inc_linear import INCLinearMethod +from .inc_scheme import INCScheme + +if TYPE_CHECKING: + import torch + + from ..config_parser import INCLayerConfig + from ..inc import INCConfig + +logger = init_logger(__name__) + + +class INCWna16Scheme(INCScheme): + @staticmethod + def can_handle(layer_config: "INCLayerConfig") -> bool: + return layer_config.is_wna16_int + + def get_linear_method( + self, + config: "INCConfig", + layer: "torch.nn.Module", + prefix: str, + layer_config: "INCLayerConfig", + ): + del config, layer + if current_platform.is_xpu(): + if layer_config.bits == 4 and layer_config.sym: + from .inc_wna16_linear import ( + INCARKLinearMethod, + INCXPULinearMethod, + get_ark_state, + ) + + is_ark_available, ark_error, _, _ = get_ark_state() + if is_ark_available: + return INCLinearMethod(INCARKLinearMethod(layer_config)) + + logger.debug( + "ARK backend is unavailable for layer %s; " + "falling back to the default XPU INC path. Error: %s", + prefix, + ark_error or "unknown error", + ) + return INCLinearMethod(INCXPULinearMethod(layer_config)) + raise NotImplementedError(f"INC on XPU: unsupported config {layer_config}") + + if current_platform.is_cpu() and layer_config.is_gptq: + if layer_config.bits == 4 and layer_config.sym: + from .inc_wna16_linear import ( + INCARKLinearMethod, + INCWNA16LinearScheme, + get_ark_state, + ) + + is_ark_available, ark_error, _, _ = get_ark_state() + if is_ark_available: + return INCLinearMethod(INCARKLinearMethod(layer_config)) + + logger.debug( + "ARK backend is unavailable for layer %s; " + "falling back to the default CPU INC path. Error: %s", + prefix, + ark_error or "unknown error", + ) + return INCLinearMethod(INCWNA16LinearScheme(layer_config)) + raise NotImplementedError(f"INC on CPU: unsupported config {layer_config}") + + from .inc_wna16_linear import INCWNA16LinearScheme + + return INCLinearMethod(INCWNA16LinearScheme(layer_config)) + + def get_moe_method( + self, + config: "INCConfig", + layer: "torch.nn.Module", + prefix: str, + layer_config: "INCLayerConfig", + ): + del config, prefix + # XPU and CPU do not support MoE quantization yet + if current_platform.is_xpu() or current_platform.is_cpu(): + from vllm.model_executor.layers.fused_moe import ( + UnquantizedFusedMoEMethod, + ) + + return UnquantizedFusedMoEMethod(layer.moe_config) + if layer_config.is_gptq: + return _resolve_gptq_moe(layer, layer_config) + if layer_config.is_awq: + return _resolve_awq_moe(layer, layer_config) + raise NotImplementedError(f"WNA16 MoE does not support config {layer_config}") + + +def _resolve_gptq_moe(layer: "torch.nn.Module", layer_config: "INCLayerConfig"): + from vllm.model_executor.layers.quantization.auto_gptq import ( + AutoGPTQMoEMethod, + ) + from vllm.model_executor.layers.quantization.moe_wna16 import ( + MoeWNA16Config, + MoeWNA16Method, + ) + from vllm.model_executor.layers.quantization.utils.marlin_utils import ( + check_marlin_supported, + check_moe_marlin_supports_layer, + ) + + gptq_type_map = { + (4, True): scalar_types.uint4b8, + (8, True): scalar_types.uint8b128, + } + use_marlin = (layer_config.bits, layer_config.sym) in gptq_type_map + if use_marlin: + use_marlin = check_marlin_supported( + gptq_type_map[(layer_config.bits, layer_config.sym)], + layer_config.group_size, + has_zp=not layer_config.sym, + ) and check_moe_marlin_supports_layer(layer, layer_config.group_size) + + if use_marlin: + return AutoGPTQMoEMethod( + AutoGPTQConfig( + weight_bits=layer_config.bits, + group_size=layer_config.group_size, + desc_act=False, + is_sym=layer_config.sym, + lm_head_quantized=False, + dynamic={}, + full_config={}, + ), + layer.moe_config, + ) + + moe_config = MoeWNA16Config.from_config( + { + "quant_method": "gptq", + "bits": layer_config.bits, + "group_size": layer_config.group_size, + "sym": layer_config.sym, + "lm_head": False, + } + ) + return MoeWNA16Method(moe_config, layer.moe_config) + + +def _resolve_awq_moe(layer: "torch.nn.Module", layer_config: "INCLayerConfig"): + from vllm.model_executor.layers.quantization.awq_marlin import AWQMarlinMoEMethod + from vllm.model_executor.layers.quantization.moe_wna16 import ( + MoeWNA16Config, + MoeWNA16Method, + ) + from vllm.model_executor.layers.quantization.utils.marlin_utils import ( + check_marlin_supported, + check_moe_marlin_supports_layer, + ) + + awq_type_map = { + 4: scalar_types.uint4, + 8: scalar_types.uint8, + } + use_marlin = layer_config.bits in awq_type_map + if use_marlin: + use_marlin = check_marlin_supported( + awq_type_map[layer_config.bits], + layer_config.group_size, + not layer_config.sym, + ) and check_moe_marlin_supports_layer(layer, layer_config.group_size) + + if use_marlin: + return AWQMarlinMoEMethod( + AWQMarlinConfig( + weight_bits=layer_config.bits, + group_size=layer_config.group_size, + zero_point=not layer_config.sym, + lm_head_quantized=False, + modules_to_not_convert=[], + full_config={}, + ), + layer.moe_config, + ) + + moe_config = MoeWNA16Config.from_config( + { + "quant_method": "awq", + "bits": layer_config.bits, + "group_size": layer_config.group_size, + "zero_point": not layer_config.sym, + "lm_head": False, + } + ) + return MoeWNA16Method(moe_config, layer.moe_config) From 0b131b16c93308df90d013c5ac24c6fa25bd863b Mon Sep 17 00:00:00 2001 From: xaguilar-amd Date: Wed, 17 Jun 2026 16:05:34 +0200 Subject: [PATCH 23/37] [ROCm][AITER][Quark] Tag per-channel FP8 weights as PER_CHANNEL so AITER pre-shuffled GEMM is selected (#44626) Signed-off-by: Xavier Aguilar --- .../quantization/quark/schemes/quark_w8a8_fp8.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/vllm/model_executor/layers/quantization/quark/schemes/quark_w8a8_fp8.py b/vllm/model_executor/layers/quantization/quark/schemes/quark_w8a8_fp8.py index 6d94e26f960..280159700e6 100644 --- a/vllm/model_executor/layers/quantization/quark/schemes/quark_w8a8_fp8.py +++ b/vllm/model_executor/layers/quantization/quark/schemes/quark_w8a8_fp8.py @@ -16,8 +16,8 @@ from vllm.model_executor.layers.quantization.quark.schemes import QuarkScheme from vllm.model_executor.layers.quantization.utils.quant_utils import ( GroupShape, kFp8DynamicTokenSym, + kFp8StaticChannelSym, kFp8StaticTensorSym, - kFp8StaticTokenSym, ) from vllm.model_executor.layers.quantization.utils.w8a8_utils import ( normalize_e4m3fn_to_e4m3fnuz, @@ -49,13 +49,17 @@ class QuarkW8A8Fp8(QuarkScheme): per_token_activation = ( not self.is_static_input_scheme and self.input_qscheme == "per_channel" ) - per_token_weight = self.weight_qscheme == "per_channel" + per_channel_weight = self.weight_qscheme == "per_channel" self.activation_quant_key = ( kFp8DynamicTokenSym if per_token_activation else kFp8StaticTensorSym ) + # A per-output-channel weight scale is one fp32 value per weight row + # (length N). Tag it as ``GroupShape.PER_CHANNEL`` to match the + # canonical compressed-tensors CHANNEL strategy, so kernel selection + # (e.g. AITER's pre-shuffled FP8 GEMM) treats it uniformly. self.weight_quant_key = ( - kFp8StaticTokenSym if per_token_weight else kFp8StaticTensorSym + kFp8StaticChannelSym if per_channel_weight else kFp8StaticTensorSym ) self.out_dtype = torch.get_default_dtype() self.input_dtype = get_current_vllm_config().model_config.dtype From 8b2b566ea710e48149886aecf2e17a2c78d55c11 Mon Sep 17 00:00:00 2001 From: amirkl94 <203507526+amirkl94@users.noreply.github.com> Date: Wed, 17 Jun 2026 17:32:49 +0300 Subject: [PATCH 24/37] Feature: Enable Flashinfer non-gated MoE bf16 (#43853) Signed-off-by: Amir Klein <203507526+amirkl94@users.noreply.github.com> --- .../layers/fused_moe/experts/trtllm_bf16_moe.py | 13 ++++++++++--- .../layers/fused_moe/oracle/unquantized.py | 12 ++++++++++++ .../layers/quantization/utils/flashinfer_utils.py | 13 ++++++++----- 3 files changed, 30 insertions(+), 8 deletions(-) diff --git a/vllm/model_executor/layers/fused_moe/experts/trtllm_bf16_moe.py b/vllm/model_executor/layers/fused_moe/experts/trtllm_bf16_moe.py index 033b7dc7d93..550d6b5341d 100644 --- a/vllm/model_executor/layers/fused_moe/experts/trtllm_bf16_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/trtllm_bf16_moe.py @@ -11,6 +11,9 @@ from vllm.model_executor.layers.fused_moe.config import ( FusedMoEQuantConfig, RoutingMethodType, ) +from vllm.model_executor.layers.quantization.utils.flashinfer_utils import ( + activation_to_flashinfer_int, +) from vllm.model_executor.layers.quantization.utils.quant_utils import ( QuantKey, ) @@ -54,8 +57,8 @@ class TrtLlmBf16Experts(mk.FusedMoEExpertsMonolithic): @staticmethod def _supports_no_act_and_mul() -> bool: - """BF16 kernels do not support non-gated MoE""" - return False + """BF16 kernels support non-gated MoE via RELU2_NO_MUL.""" + return True @staticmethod def _supports_quant_scheme( @@ -67,7 +70,8 @@ class TrtLlmBf16Experts(mk.FusedMoEExpertsMonolithic): @staticmethod def _supports_activation(activation: MoEActivation) -> bool: - return activation in [MoEActivation.SILU] + """Supports SiLU (gated) and RELU^2 (non-gated) activations.""" + return activation in [MoEActivation.SILU, MoEActivation.RELU2_NO_MUL] @staticmethod def _supports_routing_method( @@ -123,6 +127,8 @@ class TrtLlmBf16Experts(mk.FusedMoEExpertsMonolithic): ) -> torch.Tensor: import flashinfer + assert activation in [MoEActivation.SILU, MoEActivation.RELU2_NO_MUL] + return flashinfer.fused_moe.trtllm_bf16_moe( routing_logits=router_logits, routing_bias=e_score_correction_bias, @@ -138,4 +144,5 @@ class TrtLlmBf16Experts(mk.FusedMoEExpertsMonolithic): local_num_experts=self.local_num_experts, routed_scaling_factor=routed_scaling_factor, routing_method_type=self.routing_method_type, + activation_type=activation_to_flashinfer_int(activation), ) diff --git a/vllm/model_executor/layers/fused_moe/oracle/unquantized.py b/vllm/model_executor/layers/fused_moe/oracle/unquantized.py index 5c156dc9b16..a7dcd801376 100644 --- a/vllm/model_executor/layers/fused_moe/oracle/unquantized.py +++ b/vllm/model_executor/layers/fused_moe/oracle/unquantized.py @@ -19,6 +19,7 @@ from vllm.model_executor.layers.fused_moe.config import ( FusedMoEQuantConfig, ) from vllm.model_executor.layers.quantization.utils.flashinfer_utils import ( + align_moe_weights_for_fi, convert_moe_weights_to_flashinfer_trtllm_block_layout, swap_w13_to_w31, ) @@ -269,11 +270,22 @@ def convert_to_unquantized_kernel_format( w13_weight = swap_w13_to_w31(w13_weight) elif unquantized_backend == UnquantizedMoeBackend.FLASHINFER_TRTLLM: + is_act_and_mul = layer.moe_config.is_act_and_mul + if not is_act_and_mul: + # Kernel requires intermediate_size_per_partition % 128 == 0 (BlockMajorK + # weight layout uses block_k=128). Pad along the intermediate dim when + # the model + TP split don't satisfy the constraint. + w13_weight, w2_weight, padded_intermediate = align_moe_weights_for_fi( + w13_weight, w2_weight, is_act_and_mul, min_alignment=128 + ) + layer.moe_config.intermediate_size_per_partition = padded_intermediate + _cache_permute_indices: dict[torch.Size, torch.Tensor] = {} w13_weight, w2_weight = convert_moe_weights_to_flashinfer_trtllm_block_layout( _cache_permute_indices, w13_weight, w2_weight, + is_gated_act_gemm=is_act_and_mul, ) return w13_weight.contiguous(), w2_weight.contiguous() diff --git a/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py b/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py index 9e29ddaeb6f..1cbfdf69c99 100644 --- a/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py +++ b/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py @@ -108,6 +108,7 @@ def convert_moe_weights_to_flashinfer_trtllm_block_layout( cache_permute_indices: dict[torch.Size, torch.Tensor], w13_weight: torch.Tensor, w2_weight: torch.Tensor, + is_gated_act_gemm: bool = True, ) -> tuple[torch.Tensor, torch.Tensor]: """Convert expert weights to FlashInfer's block layout. @@ -166,9 +167,11 @@ def convert_moe_weights_to_flashinfer_trtllm_block_layout( cache_permute_indices, w13_expert_uint8, epilogue_tile_m, + is_gated_act_gemm=is_gated_act_gemm, ) - rows = w13_expert_uint8.shape[0] - permute_indices = (permute_indices + rows // 2) % rows + if is_gated_act_gemm: + rows = w13_expert_uint8.shape[0] + permute_indices = (permute_indices + rows // 2) % rows _copy_permuted_expert_to_block_layout( w13_weights_shuffled_tensor[i], w13_expert_uint8, @@ -288,12 +291,12 @@ def align_trtllm_fp4_moe_hidden_dim_for_fi( return padded_w13, padded_w13_scale, padded_w2, padded_w2_scale, padded_hidden_size -def align_fp8_moe_weights_for_fi( +def align_moe_weights_for_fi( w13: torch.Tensor, w2: torch.Tensor, is_act_and_mul: bool, min_alignment: int = 16 ) -> tuple[torch.Tensor, torch.Tensor, int]: """Pad intermediate size so FlashInfer kernels' alignment constraints hold. - Some FlashInfer FP8 MoE kernels require the (gated) intermediate size + Some FlashInfer MoE kernels require the (gated) intermediate size used for GEMM to be divisible by a small alignment value. When this is not satisfied (e.g. with certain tensor-parallel sizes), we pad the gate/up and down projection weights along the intermediate dim. @@ -492,7 +495,7 @@ def prepare_fp8_moe_layer_for_fi( # for the gate-up proj. Pad the weights to respect this. if not block_quant: min_alignment = 16 if is_gated else 128 - w13, w2, new_intermediate = align_fp8_moe_weights_for_fi( + w13, w2, new_intermediate = align_moe_weights_for_fi( w13, w2, layer.moe_config.is_act_and_mul, From 0a7bacdcacc5f5c7e511f596b1f6b372da6029b6 Mon Sep 17 00:00:00 2001 From: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Date: Wed, 17 Jun 2026 10:55:48 -0400 Subject: [PATCH 25/37] [DSv4 Perf] DSv4 flashinfer sparse index cache for metadata, 2%~4% TTFT improvement (#45863) Signed-off-by: yewentao256 --- .../kernels/attention/test_flashmla_sparse.py | 147 ++++++++++++++++++ .../deepseek_v4/nvidia/flashinfer_sparse.py | 50 ++++-- vllm/v1/attention/backends/mla/sparse_swa.py | 5 +- 3 files changed, 184 insertions(+), 18 deletions(-) diff --git a/tests/kernels/attention/test_flashmla_sparse.py b/tests/kernels/attention/test_flashmla_sparse.py index ce8b48ac289..010c4479766 100644 --- a/tests/kernels/attention/test_flashmla_sparse.py +++ b/tests/kernels/attention/test_flashmla_sparse.py @@ -143,3 +143,150 @@ def test_deepseek_v4_prefill_chunk_planning_expands_for_short_sequences(): # the adaptive plan keeps all 5 in one chunk assert chunk_plan == [(0, 5, 36, 103)] + + +def test_flashinfer_sparse_indices_cache(monkeypatch): + from vllm.models.deepseek_v4.nvidia import flashinfer_sparse as flashinfer_mod + from vllm.models.deepseek_v4.sparse_mla import DeepseekV4FlashMLAMetadata + from vllm.v1.attention.backends.mla.sparse_swa import DeepseekSparseSWAMetadata + + builder_calls = 0 + + def fake_build(*args, **kwargs): + nonlocal builder_calls + builder_calls += 1 + return ( + torch.tensor([[builder_calls]], dtype=torch.int32), + torch.tensor([builder_calls], dtype=torch.int32), + ) + + monkeypatch.setattr( + flashinfer_mod, "build_flashinfer_mixed_sparse_indices", fake_build + ) + + def make_attn(compress_ratio: int, topk_width: int): + attn = object.__new__(flashinfer_mod.DeepseekV4FlashInferMLAAttention) + attn.compress_ratio = compress_ratio + attn.window_size = 4 + attn.topk_indices_buffer = torch.tensor( + [[0, 1], [2, 3], [4, 5]], dtype=torch.int32 + )[:, :topk_width] + return attn + + def make_swa_metadata(): + return DeepseekSparseSWAMetadata( + block_table=torch.tensor([[0, 1], [2, 3]], dtype=torch.int32), + slot_mapping=torch.tensor([0, 1], dtype=torch.int64), + block_size=64, + seq_lens=torch.tensor([8, 10], dtype=torch.int32), + query_start_loc=torch.tensor([0, 1, 3], dtype=torch.int32), + query_start_loc_cpu=torch.tensor([0, 1, 3], dtype=torch.int32), + token_to_req_indices=torch.tensor([0, 1, 1], dtype=torch.int32), + decode_swa_indices=torch.tensor([[5, 6, -1, -1]], dtype=torch.int32), + decode_swa_lens=torch.tensor([2], dtype=torch.int32), + is_valid_token=torch.tensor([True], dtype=torch.bool), + num_decodes=1, + num_prefills=1, + num_decode_tokens=1, + num_prefill_tokens=2, + ) + + def make_flashmla_metadata(): + return DeepseekV4FlashMLAMetadata( + num_reqs=2, + max_query_len=2, + max_seq_len=10, + num_actual_tokens=3, + query_start_loc=torch.tensor([0, 1, 3], dtype=torch.int32), + slot_mapping=torch.tensor([0, 1, 2], dtype=torch.int64), + block_table=torch.tensor([[0, 1], [2, 3]], dtype=torch.int32), + req_id_per_token=torch.tensor([0, 1, 1], dtype=torch.int32), + block_size=256, + topk_tokens=2, + c128a_global_decode_topk_indices=torch.tensor( + [[[9, 10]]], dtype=torch.int32 + ), + c128a_decode_topk_lens=torch.tensor([2], dtype=torch.int32), + c128a_prefill_topk_indices=torch.tensor( + [[0, 1], [1, 2]], dtype=torch.int32 + ), + ) + + swa_attn = make_attn(1, 0) + swa_metadata = make_swa_metadata() + _, _, sparse_indices_first, sparse_lens_first = ( + swa_attn._build_sparse_index_metadata( + kv_cache=None, + swa_k_cache=torch.empty((1, 64, 512), dtype=torch.bfloat16), + swa_metadata=swa_metadata, + attn_metadata=None, + swa_only=True, + ) + ) + _, _, sparse_indices_second, sparse_lens_second = ( + swa_attn._build_sparse_index_metadata( + kv_cache=None, + swa_k_cache=torch.empty((1, 64, 512), dtype=torch.bfloat16), + swa_metadata=swa_metadata, + attn_metadata=None, + swa_only=True, + ) + ) + assert builder_calls == 1 + assert sparse_indices_first is sparse_indices_second + assert sparse_lens_first is sparse_lens_second + + c128a_attn = make_attn(128, 2) + c128a_metadata = make_swa_metadata() + c128a_flashmla_md = make_flashmla_metadata() + _, _, sparse_indices_first, sparse_lens_first = ( + c128a_attn._build_sparse_index_metadata( + kv_cache=torch.empty((1, 2, 512), dtype=torch.bfloat16), + swa_k_cache=torch.empty((1, 64, 512), dtype=torch.bfloat16), + swa_metadata=c128a_metadata, + attn_metadata=c128a_flashmla_md, + swa_only=False, + ) + ) + _, _, sparse_indices_second, sparse_lens_second = ( + c128a_attn._build_sparse_index_metadata( + kv_cache=torch.empty((1, 2, 512), dtype=torch.bfloat16), + swa_k_cache=torch.empty((1, 64, 512), dtype=torch.bfloat16), + swa_metadata=c128a_metadata, + attn_metadata=c128a_flashmla_md, + swa_only=False, + ) + ) + + assert builder_calls == 2 + assert sparse_indices_first is sparse_indices_second + assert sparse_lens_first is sparse_lens_second + + c4a_attn = make_attn(4, 2) + c4a_metadata = make_swa_metadata() + c4a_flashmla_md = make_flashmla_metadata() + c4a_flashmla_md.c128a_global_decode_topk_indices = None + c4a_flashmla_md.c128a_decode_topk_lens = None + c4a_flashmla_md.c128a_prefill_topk_indices = None + _, _, sparse_indices_third, sparse_lens_third = ( + c4a_attn._build_sparse_index_metadata( + kv_cache=torch.empty((1, 2, 512), dtype=torch.bfloat16), + swa_k_cache=torch.empty((1, 64, 512), dtype=torch.bfloat16), + swa_metadata=c4a_metadata, + attn_metadata=c4a_flashmla_md, + swa_only=False, + ) + ) + _, _, sparse_indices_fourth, sparse_lens_fourth = ( + c4a_attn._build_sparse_index_metadata( + kv_cache=torch.empty((1, 2, 512), dtype=torch.bfloat16), + swa_k_cache=torch.empty((1, 64, 512), dtype=torch.bfloat16), + swa_metadata=c4a_metadata, + attn_metadata=c4a_flashmla_md, + swa_only=False, + ) + ) + + assert builder_calls == 4 + assert sparse_indices_third is not sparse_indices_fourth + assert sparse_lens_third is not sparse_lens_fourth diff --git a/vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py b/vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py index a357edf5548..9b2542450b1 100644 --- a/vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py +++ b/vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py @@ -288,24 +288,40 @@ class DeepseekV4FlashInferMLAAttention(DeepseekV4Attention): query_start_loc = swa_metadata.query_start_loc[: num_reqs + 1] seq_lens = swa_metadata.seq_lens[:num_reqs] assert seq_lens.dtype == torch.int32 - sparse_indices, sparse_topk_lens = build_flashinfer_mixed_sparse_indices( - decode_swa_indices, - decode_compressed_indices, - decode_compressed_topk_lens, - prefill_topk_indices[:num_prefill_tokens], - query_start_loc, - seq_lens, - swa_metadata.token_to_req_indices[:num_tokens], - swa_metadata.block_table[:num_reqs], - swa_metadata.block_size, - compressed_block_table, - compressed_block_size, - self.window_size, - self.compress_ratio, - top_k, - decode_compressed_indices_are_local=decode_compressed_indices_are_local, - decode_is_valid_token=decode_is_valid_token, + # cache for SWA-only and C128A that build the same mixed sparse indices + # C4A stays uncached. + cache_key = ( + "swa_only" + if swa_only + else ("c128a" if self.compress_ratio == 128 else "c4a") ) + cached_sparse = swa_metadata.flashinfer_sparse_index_cache.get(cache_key, None) + if cached_sparse is None: + sparse_indices, sparse_topk_lens = build_flashinfer_mixed_sparse_indices( + decode_swa_indices, + decode_compressed_indices, + decode_compressed_topk_lens, + prefill_topk_indices[:num_prefill_tokens], + query_start_loc, + seq_lens, + swa_metadata.token_to_req_indices[:num_tokens], + swa_metadata.block_table[:num_reqs], + swa_metadata.block_size, + compressed_block_table, + compressed_block_size, + self.window_size, + self.compress_ratio, + top_k, + decode_compressed_indices_are_local=decode_compressed_indices_are_local, + decode_is_valid_token=decode_is_valid_token, + ) + if cache_key != "c4a": + swa_metadata.flashinfer_sparse_index_cache[cache_key] = ( + sparse_indices, + sparse_topk_lens, + ) + else: + sparse_indices, sparse_topk_lens = cached_sparse return compressed_kv_cache, seq_lens, sparse_indices, sparse_topk_lens def _forward( diff --git a/vllm/v1/attention/backends/mla/sparse_swa.py b/vllm/v1/attention/backends/mla/sparse_swa.py index 1774018a8cf..a3fd39bed79 100644 --- a/vllm/v1/attention/backends/mla/sparse_swa.py +++ b/vllm/v1/attention/backends/mla/sparse_swa.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import ClassVar, cast import torch @@ -193,6 +193,9 @@ class DeepseekSparseSWAMetadata: tile_sched_swaonly: "FlashMLASchedMeta | None" = None tile_sched_c4a: "FlashMLASchedMeta | None" = None tile_sched_c128a: "FlashMLASchedMeta | None" = None + flashinfer_sparse_index_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] = field( + default_factory=dict + ) def get_prefill_chunk_plan( self, compress_ratio: int, prefill_chunk_size: int From 46f74e144b33b01684674e3c808babb19835d17b Mon Sep 17 00:00:00 2001 From: "Xiaohong (Sean) Chen" Date: Wed, 17 Jun 2026 11:03:54 -0400 Subject: [PATCH 26/37] [Kernel][Helion][1/N] Add Helion kernel for rms_norm_dynamic_per_token_quant (#34432) Signed-off-by: Sean Chen Co-authored-by: Yanan Cao --- .../test_rms_norm_dynamic_per_token_quant.py | 207 + .../nvidia_b200.json | 2647 ++++++++++++ .../nvidia_h100.json | 3663 +++++++++++++++++ .../ops/rms_norm_dynamic_per_token_quant.py | 233 ++ 4 files changed, 6750 insertions(+) create mode 100644 tests/kernels/helion/test_rms_norm_dynamic_per_token_quant.py create mode 100644 vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_b200.json create mode 100644 vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_h100.json create mode 100644 vllm/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py diff --git a/tests/kernels/helion/test_rms_norm_dynamic_per_token_quant.py b/tests/kernels/helion/test_rms_norm_dynamic_per_token_quant.py new file mode 100644 index 00000000000..3842419562c --- /dev/null +++ b/tests/kernels/helion/test_rms_norm_dynamic_per_token_quant.py @@ -0,0 +1,207 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Tests for the rms_norm_dynamic_per_token_quant helion kernel + +Run `pytest tests/kernels/helion/test_rms_norm_dynamic_per_token_quant.py`. +""" + +from typing import Any + +import pytest +import torch +from torch._subclasses.fake_tensor import FakeTensorMode + +from tests.kernels.helion.utils import skip_if_platform_unsupported +from vllm.kernels.helion.case_key import CaseKey +from vllm.kernels.helion.config_manager import ConfigManager +from vllm.kernels.helion.ops.rms_norm_dynamic_per_token_quant import ( + _pick_cache, + baseline, + pick_config, + rms_norm_dynamic_per_token_quant, +) +from vllm.platforms import current_platform +from vllm.utils.import_utils import has_helion +from vllm.utils.torch_utils import set_random_seed + +if not has_helion(): + pytest.skip( + "Helion is not installed. Install with: pip install vllm[helion]", + allow_module_level=True, + ) + + +def _generate_fake_input(num_tokens: int, hidden_size: int) -> tuple[Any, ...]: + with FakeTensorMode(): + input = torch.randn( + num_tokens, hidden_size, device="cuda", dtype=torch.bfloat16 + ) + result = torch.empty( + input.shape, device=input.device, dtype=current_platform.fp8_dtype() + ) + scale = torch.empty((num_tokens, 1), device=input.device, dtype=torch.float32) + scale_ub = torch.mean(input).to(torch.float32) + residual = torch.randn_like(input) + weight = torch.normal( + mean=1.0, + std=1.0, + size=(hidden_size,), + dtype=input.dtype, + device=input.device, + ) + epsilon = 1e-6 + args = (result, input, weight, scale, epsilon, scale_ub, residual) + return args + + +@pytest.fixture(autouse=True) +def reset_config_manager_singleton(): + ConfigManager.reset_instance() + ConfigManager() + yield + ConfigManager.reset_instance() + + +class TestRmsNormDynamicPerTokenQuantConfigPicker: + def setup_method(self): + _pick_cache.clear() + + def test_config_picker_exact_match(self): + config_keys = [ + CaseKey({"hidden_size": 2048, "num_tokens": 16}), + CaseKey({"hidden_size": 4096, "num_tokens": 16}), + ] + + args = _generate_fake_input(16, 4096) + selected_key = pick_config(args, config_keys) + assert selected_key == CaseKey({"hidden_size": 4096, "num_tokens": 16}) + + def test_config_picker_closest_match(self): + config_keys = [ + CaseKey({"hidden_size": 2048, "num_tokens": 16}), + CaseKey({"hidden_size": 2048, "num_tokens": 32}), + CaseKey({"hidden_size": 4096, "num_tokens": 16}), + CaseKey({"hidden_size": 4096, "num_tokens": 32}), + ] + + args = _generate_fake_input(20, 3000) + selected_key = pick_config(args, config_keys) + assert selected_key == CaseKey({"hidden_size": 2048, "num_tokens": 32}) + + def test_config_picker_no_configs(self): + config_keys: list[dict] = [] + + args = _generate_fake_input(16, 4096) + selected_key = pick_config(args, config_keys) + assert selected_key is None + + def test_config_picker_fallback_to_largest(self): + config_keys = [ + CaseKey({"hidden_size": 2048, "num_tokens": 16}), + CaseKey({"hidden_size": 4096, "num_tokens": 16}), + ] + + args = _generate_fake_input(32, 8192) + selected_key = pick_config(args, config_keys) + assert selected_key == CaseKey({"hidden_size": 4096, "num_tokens": 16}) + + +DTYPES = [torch.bfloat16, torch.float] +QUANT_DTYPES = [torch.int8, current_platform.fp8_dtype()] +VEC_HIDDEN_SIZES = [1024, 1025, 1027, 1029] +# Avoid combinatorial explosion with full Cartesian product +NUM_TOKENS_HIDDEN_SIZES = [ + *[(1, i) for i in [1, 64, *VEC_HIDDEN_SIZES, 5120, 5137]], + *[(2048, i) for i in [1, 64, *VEC_HIDDEN_SIZES, 5137]], + *[(4096, i) for i in [1, 64, 5137]], +] + +ADD_RESIDUAL = [False, True] +SCALE_UBS = [True, False] +SEEDS = [0] + +EPS = 1e-6 + + +class TestRmsNormDynamicPerTokenQuantCorrectness: + @pytest.mark.parametrize("num_tokens, hidden_size", NUM_TOKENS_HIDDEN_SIZES) + @pytest.mark.parametrize("add_residual", ADD_RESIDUAL) + @pytest.mark.parametrize("has_scale_ub", SCALE_UBS) + @pytest.mark.parametrize("dtype", DTYPES) + @pytest.mark.parametrize("quant_dtype", QUANT_DTYPES) + @pytest.mark.parametrize("seed", SEEDS) + def test_rms_norm_dynamic_per_token_quant( + self, + num_tokens: int, + hidden_size: int, + add_residual: bool, + has_scale_ub: bool, + dtype: torch.dtype, + quant_dtype: torch.dtype, + seed: int, + ) -> None: + skip_if_platform_unsupported("rms_norm_dynamic_per_token_quant") + + set_random_seed(seed) + + if has_scale_ub and quant_dtype != current_platform.fp8_dtype(): + # skip + return + + scale = 1 / (hidden_size) + x = torch.randn(num_tokens, hidden_size, dtype=dtype, device="cuda") * scale + weight = torch.normal( + mean=1.0, std=1.0, size=(hidden_size,), dtype=dtype, device=x.device + ) + residual = torch.randn_like(x) * scale if add_residual else None + scale_ub = ( + torch.mean(x).to(dtype=torch.float32, device="cuda") + if has_scale_ub + else None + ) + + ref_out = torch.empty(x.shape, device=x.device, dtype=quant_dtype) + ref_scales = torch.empty((x.shape[0], 1), device=x.device, dtype=torch.float32) + ref_residual = residual.clone() if residual is not None else None + baseline(ref_out, x, weight, ref_scales, EPS, scale_ub, ref_residual) + + ops_out = torch.empty(x.shape, device=x.device, dtype=quant_dtype) + ops_scales = torch.empty((x.shape[0], 1), device=x.device, dtype=torch.float32) + ops_residual = residual.clone() if residual is not None else None + rms_norm_dynamic_per_token_quant( + ops_out, x, weight, ops_scales, EPS, scale_ub, ops_residual + ) + + torch.testing.assert_close(ref_scales, ops_scales) + # allow 1 ULP difference + assert ( + ref_out.view(torch.uint8).to(torch.int16) + - ops_out.view(torch.uint8).to(torch.int16) + ).abs().max() <= 1 + + if add_residual: + torch.testing.assert_close(ref_residual, ops_residual) + + +class TestRmsNormDynamicPerTokenQuantIntegration: + def test_kernel_registration_integration(self): + from vllm.kernels.helion.register import get_registered_kernels + + registered_kernels = get_registered_kernels() + assert "rms_norm_dynamic_per_token_quant" in registered_kernels + + kernel_wrapper = registered_kernels["rms_norm_dynamic_per_token_quant"] + assert kernel_wrapper.op_name == "rms_norm_dynamic_per_token_quant" + assert kernel_wrapper._config_picker is not None + assert kernel_wrapper._mutates_args == ["result", "scale", "residual"] + + def test_fake_impl_functionality(self): + skip_if_platform_unsupported("rms_norm_dynamic_per_token_quant") + from vllm.kernels.helion.register import get_registered_kernels + + registered_kernels = get_registered_kernels() + kernel_wrapper = registered_kernels["rms_norm_dynamic_per_token_quant"] + fake_impl = kernel_wrapper._fake_impl + + args = _generate_fake_input(16, 4096) + assert fake_impl(*args) is None diff --git a/vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_b200.json b/vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_b200.json new file mode 100644 index 00000000000..5be0d7b71b2 --- /dev/null +++ b/vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_b200.json @@ -0,0 +1,2647 @@ +[ + { + "key": { + "hidden_size": 2048, + "num_tokens": 1 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + false, + false, + true + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "", + "first", + "last", + "", + "" + ], + "num_warps": 16, + "num_stages": 3, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 1 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 2, + 3, + 1, + 0 + ], + "range_warp_specializes": [], + "range_multi_buffers": [ + true, + null, + false, + null + ], + "range_flattens": [ + null, + null, + false, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 16, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "persistent_interleaved", + "num_sm_multiplier": 1, + "maxnreg": 128 + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 1 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + false, + true, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "", + "last", + "last", + "first", + "first" + ], + "num_warps": 32, + "num_stages": 6, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 2 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 0, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + null + ], + "range_flattens": [ + null, + true, + null, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 7, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 2 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + false, + null + ], + "range_flattens": [ + null, + false, + true, + true + ], + "load_eviction_policies": [ + "", + "first", + "", + "first", + "last", + "", + "last", + "", + "last" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 2 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 4 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 1, + 2, + 0 + ], + "range_warp_specializes": [ + null, + null, + null, + null + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + false, + true + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "", + "first", + "first", + "", + "last", + "", + "", + "last", + "" + ], + "num_warps": 8, + "num_stages": 3, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 4 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 0, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + null + ], + "range_flattens": [ + null, + true, + null, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 7, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 4 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 8 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 3, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + true, + null + ], + "range_flattens": [ + null, + false, + false, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "last", + "first", + "first", + "", + "" + ], + "num_warps": 8, + "num_stages": 4, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 8 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 1, + 2, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + null, + null + ], + "range_flattens": [ + null, + null, + null, + null + ], + "load_eviction_policies": [ + "first", + "", + "first", + "", + "last", + "first", + "", + "", + "first" + ], + "num_warps": 16, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 8 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 16 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + true, + true + ], + "range_flattens": [ + null, + false, + false, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "", + "", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 16 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 0, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + true, + true, + false + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "last", + "last", + "", + "", + "" + ], + "num_warps": 16, + "num_stages": 2, + "indexing": [ + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 16 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 4, + 1 + ], + "range_warp_specializes": [ + null, + false, + null, + false + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + true, + false, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "first", + "first", + "", + "last", + "", + "last" + ], + "num_warps": 8, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 32 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + true, + true + ], + "range_flattens": [ + null, + false, + false, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "", + "", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 32 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 1, + 2, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + null, + null + ], + "range_flattens": [ + null, + null, + null, + null + ], + "load_eviction_policies": [ + "first", + "", + "first", + "", + "last", + "first", + "", + "", + "first" + ], + "num_warps": 16, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 32 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 8192 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "first", + "first", + "", + "", + "last" + ], + "num_warps": 32, + "num_stages": 2, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 64 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + true, + true + ], + "range_flattens": [ + null, + false, + false, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "", + "", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 64 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 2, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + null + ], + "range_flattens": [ + null, + false, + null, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "first", + "first", + "", + "first", + "" + ], + "num_warps": 16, + "num_stages": 2, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 64 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 128 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + false, + false, + true + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "", + "first", + "last", + "", + "" + ], + "num_warps": 16, + "num_stages": 3, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 128 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 0, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + true, + true, + false + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "last", + "last", + "", + "", + "" + ], + "num_warps": 16, + "num_stages": 2, + "indexing": [ + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 128 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 4, + 1 + ], + "range_warp_specializes": [ + null, + false, + null, + false + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + true, + false, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "first", + "first", + "", + "last", + "", + "last" + ], + "num_warps": 8, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 256 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + true, + true + ], + "range_flattens": [ + null, + false, + false, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "", + "", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 256 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 0, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + null + ], + "range_flattens": [ + null, + true, + null, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 7, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 256 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 512 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 0, + 0, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "first", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 8, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 512 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + true + ], + "range_flattens": [ + null, + false, + null, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 512 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 1024 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + null + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "first", + "", + "last", + "last", + "", + "last", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 1024 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 3, + 0, + 0 + ], + "range_warp_specializes": [ + null, + null, + null, + true + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + true + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "first", + "last", + "first", + "last", + "", + "last", + "", + "", + "last" + ], + "num_warps": 1, + "num_stages": 8, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 1024 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [ + null, + null, + false, + null + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + false, + null, + true + ], + "load_eviction_policies": [ + "", + "last", + "", + "last", + "", + "last", + "", + "", + "" + ], + "num_warps": 8, + "num_stages": 8, + "indexing": [ + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 1024 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 2048 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 0, + 4, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + false, + false + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "first", + "", + "first", + "", + "last", + "last", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 2048 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + true + ], + "range_flattens": [ + null, + false, + null, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 2048 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 1024 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + null + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "first", + "", + "last", + "last", + "", + "last", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 4096 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 1, + 4, + 2 + ], + "range_warp_specializes": [ + null, + false, + false, + null + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + null, + true + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "first", + "", + "last", + "last", + "", + "", + "", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 4096 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [ + null, + null, + false, + null + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + false, + null, + true + ], + "load_eviction_policies": [ + "", + "last", + "", + "last", + "", + "last", + "", + "", + "" + ], + "num_warps": 8, + "num_stages": 8, + "indexing": [ + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 4096 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 8192 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 0, + 4, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + false, + false + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "first", + "", + "first", + "", + "last", + "last", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 8192 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [ + null, + null, + false, + null + ], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + false, + null, + true + ], + "load_eviction_policies": [ + "", + "last", + "", + "last", + "", + "last", + "", + "", + "" + ], + "num_warps": 8, + "num_stages": 8, + "indexing": [ + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 8192 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + } +] \ No newline at end of file diff --git a/vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_h100.json b/vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_h100.json new file mode 100644 index 00000000000..a58a67acc27 --- /dev/null +++ b/vllm/kernels/helion/configs/rms_norm_dynamic_per_token_quant/nvidia_h100.json @@ -0,0 +1,3663 @@ +[ + { + "key": { + "hidden_size": 2048, + "num_tokens": 1 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 0, + 0, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "first", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 8, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 1 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 2, + 3, + 1, + 0 + ], + "range_warp_specializes": [], + "range_multi_buffers": [ + true, + null, + false, + null + ], + "range_flattens": [ + null, + null, + false, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 16, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "persistent_interleaved", + "num_sm_multiplier": 1, + "maxnreg": 128 + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 1 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + false, + true, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "", + "last", + "last", + "first", + "first" + ], + "num_warps": 32, + "num_stages": 6, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 2 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 3, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + true, + null + ], + "range_flattens": [ + null, + false, + false, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "last", + "first", + "first", + "", + "" + ], + "num_warps": 8, + "num_stages": 4, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 2 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 4, + 4, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + true, + false + ], + "range_flattens": [ + null, + true, + true, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "last", + "first", + "first", + "first", + "" + ], + "num_warps": 16, + "num_stages": 3, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 2 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 8192 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + null + ], + "range_flattens": [ + null, + true, + false, + true + ], + "load_eviction_policies": [ + "last", + "first", + "last", + "first", + "first", + "", + "last", + "", + "last" + ], + "num_warps": 32, + "num_stages": 8, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 4 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + null + ], + "range_flattens": [ + null, + true, + null, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "first", + "last", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 4, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 4 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + false, + null + ], + "range_flattens": [ + null, + false, + true, + true + ], + "load_eviction_policies": [ + "", + "first", + "", + "first", + "last", + "", + "last", + "", + "last" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 4 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 4096 + ], + "range_unroll_factors": [ + 0, + 1, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + null + ], + "range_flattens": [ + null, + false, + true, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 8 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + true, + true + ], + "range_flattens": [ + null, + false, + false, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "", + "", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 8 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 0, + 1, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + null, + false + ], + "range_flattens": [ + null, + null, + false, + true + ], + "load_eviction_policies": [ + "first", + "", + "last", + "", + "first", + "last", + "last", + "last", + "first" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 8 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 8192 + ], + "range_unroll_factors": [ + 0, + 2, + 0, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + false + ], + "range_flattens": [ + null, + null, + false, + null + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "", + "", + "first", + "first", + "" + ], + "num_warps": 32, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 16 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + null, + true + ], + "range_flattens": [ + null, + true, + true, + false + ], + "load_eviction_policies": [ + "first", + "first", + "first", + "", + "", + "last", + "", + "", + "first" + ], + "num_warps": 16, + "num_stages": 4, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 16 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + false, + true + ], + "range_flattens": [ + null, + null, + false, + true + ], + "load_eviction_policies": [ + "", + "", + "first", + "", + "last", + "last", + "", + "first", + "" + ], + "num_warps": 16, + "num_stages": 4, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 16 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 8192 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "first", + "first", + "", + "", + "last" + ], + "num_warps": 32, + "num_stages": 2, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 32 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + null, + null + ], + "range_flattens": [ + null, + true, + true, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "last", + "last", + "last", + "last", + "first", + "first" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 32 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 2, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + null + ], + "range_flattens": [ + null, + false, + null, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "first", + "first", + "", + "first", + "" + ], + "num_warps": 16, + "num_stages": 2, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 32 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 8192 + ], + "range_unroll_factors": [ + 0, + 3, + 3, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + false, + false + ], + "range_flattens": [ + null, + true, + null, + true + ], + "load_eviction_policies": [ + "first", + "last", + "first", + "last", + "", + "", + "last", + "", + "" + ], + "num_warps": 32, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 64 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + false, + false, + true + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "", + "first", + "last", + "", + "" + ], + "num_warps": 16, + "num_stages": 3, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 64 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 0, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + true, + true, + false + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "last", + "last", + "last", + "", + "", + "" + ], + "num_warps": 16, + "num_stages": 2, + "indexing": [ + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 64 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 8192 + ], + "range_unroll_factors": [ + 0, + 2, + 2, + 0 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + true, + null + ], + "range_flattens": [ + null, + false, + true, + null + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "", + "first", + "", + "first", + "" + ], + "num_warps": 32, + "num_stages": 6, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 128 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + true, + true + ], + "range_flattens": [ + null, + false, + false, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "", + "", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 128 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 1, + 2, + 1 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + null, + null + ], + "range_flattens": [ + null, + null, + null, + null + ], + "load_eviction_policies": [ + "first", + "", + "first", + "", + "last", + "first", + "", + "", + "first" + ], + "num_warps": 16, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 128 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + true, + false, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "last", + "last", + "", + "", + "last" + ], + "num_warps": 16, + "num_stages": 7, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 256 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 1, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + true, + false + ], + "load_eviction_policies": [ + "", + "", + "first", + "", + "", + "", + "last", + "", + "" + ], + "num_warps": 8, + "num_stages": 8, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 256 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 0, + 1, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + null + ], + "range_flattens": [ + null, + null, + true, + true + ], + "load_eviction_policies": [ + "first", + "", + "first", + "", + "", + "last", + "last", + "last", + "last" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 256 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 4096 + ], + "range_unroll_factors": [ + 0, + 1, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + null + ], + "range_flattens": [ + null, + false, + true, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 512 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + false, + false, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "", + "first", + "first", + "last", + "last" + ], + "num_warps": 8, + "num_stages": 6, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 512 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + true + ], + "range_flattens": [ + null, + false, + null, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 512 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 1024 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 0, + 4, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + false, + false + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "first", + "", + "first", + "", + "last", + "last", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 1024 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 0, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + null + ], + "range_flattens": [ + null, + true, + null, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 7, + "indexing": [ + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 1024 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 4096 + ], + "range_unroll_factors": [ + 0, + 1, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + null + ], + "range_flattens": [ + null, + false, + true, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 2048 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 3, + 0, + 0 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + null, + false, + null + ], + "load_eviction_policies": [ + "", + "", + "first", + "", + "", + "", + "first", + "", + "" + ], + "num_warps": 4, + "num_stages": 4, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 2048 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + true + ], + "range_flattens": [ + null, + false, + null, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 2048 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 4096 + ], + "range_unroll_factors": [ + 0, + 1, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + null + ], + "range_flattens": [ + null, + false, + true, + false + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "last", + "first", + "", + "" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 4096 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 4, + 1, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + null, + true, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "last", + "", + "", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 4096 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + false, + null, + null + ], + "load_eviction_policies": [ + "last", + "", + "", + "last", + "first", + "", + "first", + "", + "last" + ], + "num_warps": 16, + "num_stages": 8, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 4096 + }, + "config": { + "block_sizes": [ + 8192, + 4096, + 4096 + ], + "range_unroll_factors": [ + 1, + 1, + 2, + 1 + ], + "range_warp_specializes": [], + "range_multi_buffers": [ + null, + false, + false, + true + ], + "range_flattens": [ + true, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "last", + "", + "first", + "first", + "last", + "first", + "first", + "" + ], + "num_warps": 16, + "num_stages": 8, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "persistent_interleaved", + "num_sm_multiplier": 32, + "maxnreg": 128 + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 8192 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 1024 + ], + "range_unroll_factors": [ + 0, + 4, + 1, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + null + ], + "range_flattens": [ + null, + null, + true, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "last", + "", + "", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 8192 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 2048 + ], + "range_unroll_factors": [ + 2, + 2, + 3, + 2 + ], + "range_warp_specializes": [], + "range_multi_buffers": [ + true, + false, + false, + null + ], + "range_flattens": [ + false, + null, + false, + false + ], + "load_eviction_policies": [ + "last", + "last", + "", + "last", + "last", + "", + "", + "", + "" + ], + "num_warps": 8, + "num_stages": 7, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "persistent_interleaved", + "num_sm_multiplier": 64, + "maxnreg": 256 + } + }, + { + "key": { + "hidden_size": 8192, + "num_tokens": 8192 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 4096 + ], + "range_unroll_factors": [ + 0, + 3, + 1, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + true, + null + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "last", + "last", + "", + "last", + "last", + "", + "last", + "first", + "" + ], + "num_warps": 16, + "num_stages": 7, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 1 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + false, + false + ], + "range_flattens": [ + null, + false, + true, + false + ], + "load_eviction_policies": [ + "last", + "", + "last", + "", + "", + "last", + "last", + "first", + "first" + ], + "num_warps": 32, + "num_stages": 6, + "indexing": [ + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 2 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 4 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 8 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 16 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + false, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "", + "first", + "first", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 32 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 64 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 128 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 2, + 3 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + null, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "last", + "", + "first", + "last" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 256 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + false, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "", + "first", + "first", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 512 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + false, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "", + "", + "first", + "first", + "first" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 1024 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 1024 + ], + "range_unroll_factors": [ + 0, + 3, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + null + ], + "range_flattens": [ + null, + null, + true, + null + ], + "load_eviction_policies": [ + "first", + "", + "last", + "last", + "", + "last", + "first", + "first", + "last" + ], + "num_warps": 4, + "num_stages": 1, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 2048 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 4096 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 8192 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 4, + 3, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + true, + null + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "first", + "" + ], + "num_warps": 8, + "num_stages": 1, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 5120, + "num_tokens": 16384 + }, + "config": { + "block_sizes": [ + 8192, + 8192, + 2048 + ], + "range_unroll_factors": [ + 0, + 0, + 2, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + true, + null, + true + ], + "range_flattens": [ + null, + true, + null, + true + ], + "load_eviction_policies": [ + "last", + "", + "last", + "last", + "last", + "last", + "first", + "last", + "" + ], + "num_warps": 8, + "num_stages": 3, + "indexing": [ + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 2048, + "num_tokens": 16384 + }, + "config": { + "block_sizes": [ + 2048, + 2048, + 2048 + ], + "range_unroll_factors": [ + 0, + 2, + 4, + 4 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + false, + null, + false + ], + "range_flattens": [ + null, + true, + null, + true + ], + "load_eviction_policies": [ + "", + "", + "", + "", + "last", + "", + "first", + "", + "first" + ], + "num_warps": 8, + "num_stages": 6, + "indexing": [ + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "tensor_descriptor", + "tensor_descriptor" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + }, + { + "key": { + "hidden_size": 4096, + "num_tokens": 16384 + }, + "config": { + "block_sizes": [ + 4096, + 4096, + 4096 + ], + "range_unroll_factors": [ + 0, + 4, + 4, + 2 + ], + "range_warp_specializes": [], + "range_num_stages": [], + "range_multi_buffers": [ + null, + null, + true, + true + ], + "range_flattens": [ + null, + true, + null, + false + ], + "load_eviction_policies": [ + "last", + "last", + "last", + "", + "last", + "first", + "first", + "", + "last" + ], + "num_warps": 16, + "num_stages": 1, + "indexing": [ + "tensor_descriptor", + "tensor_descriptor", + "pointer", + "tensor_descriptor", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "pointer", + "tensor_descriptor", + "pointer" + ], + "atomic_indexing": [], + "pid_type": "flat" + } + } +] \ No newline at end of file diff --git a/vllm/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py b/vllm/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py new file mode 100644 index 00000000000..f15132c27cf --- /dev/null +++ b/vllm/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py @@ -0,0 +1,233 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from itertools import product +from typing import Any + +import torch + +from vllm.kernels.helion.case_key import CaseKey +from vllm.kernels.helion.utils import ( + get_fp8_dtype, + get_int8_min_max, + get_int8_min_scaling_factor, +) +from vllm.logger import init_logger +from vllm.model_executor.layers.quantization.utils.quant_utils import ( + get_fp8_min_max, +) +from vllm.platforms import current_platform +from vllm.utils.import_utils import has_helion + +if not has_helion(): + raise ImportError( + "Helion kernel requires helion to be installed. " + "Install it with: pip install helion" + ) + +import helion +import helion.language as hl + +from vllm.kernels.helion.register import register_kernel + +logger = init_logger(__name__) + + +def generate_inputs() -> dict[CaseKey, tuple[Any, ...]]: + # TODO(xiaohongchen1991): it is difficult for kernel author to cover all + # input property combination. Currently, dtypes are fixed. We need + # optimization to bucket/skip some combinations + num_tokens_list = [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192] + hidden_size_list = [2048, 4096, 5120] + in_dtype: torch.dtype = torch.bfloat16 + out_dtype: torch.dtype = current_platform.fp8_dtype() + scale_dtype: torch.dtype = torch.float32 + inputs = {} + + for num_tokens, hidden_size in product(num_tokens_list, hidden_size_list): + input = torch.randn(num_tokens, hidden_size, device="cuda", dtype=in_dtype) + result = torch.empty(input.shape, device=input.device, dtype=out_dtype) + scale = torch.empty((num_tokens, 1), device=input.device, dtype=scale_dtype) + scale_ub = torch.mean(input).to(scale_dtype) + residual = torch.randn_like(input) + weight = torch.normal( + mean=1.0, + std=1.0, + size=(hidden_size,), + dtype=input.dtype, + device=input.device, + ) + epsilon = 1e-6 + + config_key = CaseKey({"hidden_size": hidden_size, "num_tokens": num_tokens}) + inputs[config_key] = (result, input, weight, scale, epsilon, scale_ub, residual) + + return inputs + + +_pick_cache: dict[tuple[int, int], CaseKey | None] = {} + + +def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | None: + """Pick the best pre-tuned config for the given input shape. + + Selection strategy: + 1. Find the closest hidden_size among available configs + (exact match preferred). + 2. Among the num_tokens values tuned for that hidden_size, pick + the smallest num_tokens >= the input's num_tokens. If the input is + larger than all available num_tokens, fall back to the largest. + """ + + if not config_keys: + return None + + _, input, *_ = args + num_tokens, hidden_size = input.shape + + cache_key = (num_tokens, hidden_size) + cached = _pick_cache.get(cache_key) + if cached is not None: + return cached + + configs: dict[int, list[int]] = {} + for key in config_keys: + if key.is_default(): + continue + configs.setdefault(key["hidden_size"], []).append(key["num_tokens"]) + + if not configs: + return None + + best_hidden_size = min(configs, key=lambda s: abs(s - hidden_size)) + available_num_tokens = sorted(configs[best_hidden_size]) + best_num_tokens = next( + (n for n in available_num_tokens if n >= num_tokens), available_num_tokens[-1] + ) + + result = CaseKey({"hidden_size": best_hidden_size, "num_tokens": best_num_tokens}) + _pick_cache[cache_key] = result + return result + + +def fake_impl( + result: torch.Tensor, # [num_tokens, hidden_size] + input: torch.Tensor, # [num_tokens, hidden_size] + weight: torch.Tensor, # [hidden_size] + scale: torch.Tensor, # [num_tokens, 1] + epsilon: float, + scale_ub: torch.Tensor | None = None, # [] + residual: torch.Tensor | None = None, # [num_tokens, hidden_size] +) -> None: + return + + +def baseline( + result: torch.Tensor, # [num_tokens, hidden_size] + input: torch.Tensor, # [num_tokens, hidden_size] + weight: torch.Tensor, # [num_tokens] + scale: torch.Tensor, # [num_tokens, 1] + epsilon: float, + scale_ub: torch.Tensor | None = None, # [] + residual: torch.Tensor | None = None, # [num_tokens, hidden_size] +) -> None: + torch.ops._C.rms_norm_dynamic_per_token_quant( + result, input, weight, scale, epsilon, scale_ub, residual + ) + + +# Overwrite autotune_baseline_atol and autotune_baseline_rtol +# if too many configs failed due to baseline check during autotuning +@register_kernel( + mutates_args=["result", "scale", "residual"], + config_picker=pick_config, + input_generator=generate_inputs, + fake_impl=fake_impl, + helion_settings=helion.Settings( + autotune_baseline_fn=baseline, + ignore_warnings=[helion.exc.TensorOperationInWrapper], + ), +) # type: ignore[misc] +def rms_norm_dynamic_per_token_quant( + result: torch.Tensor, # [num_tokens, hidden_size] + input: torch.Tensor, # [num_tokens, hidden_size] + weight: torch.Tensor, # [hidden_size] + scale: torch.Tensor, # [num_tokens, 1] + epsilon: float, + scale_ub: torch.Tensor | None = None, # [] + residual: torch.Tensor | None = None, # [num_tokens, hidden_size] +) -> None: + # This code assumes batch_dim and num_tokens are flattened + assert input.ndim == 2 + num_tokens, hidden_size = input.shape + hl.specialize(hidden_size) + + fp8_dtype = get_fp8_dtype() + assert result.dtype in [fp8_dtype, torch.int8] + assert result.is_contiguous() and input.is_contiguous() + + if scale_ub is not None: + assert result.dtype == fp8_dtype + assert scale_ub.dtype == torch.float32 + + assert input.dtype == weight.dtype + assert scale.shape[0] == num_tokens + assert scale.dtype == torch.float32 + + if residual is not None: + assert residual.dtype == input.dtype + + quant_dtype = result.dtype + qtype_traits_min: int | float + qtype_traits_max: int | float + if quant_dtype == torch.int8: + qtype_traits_min, qtype_traits_max = get_int8_min_max() + min_scaling_factor = get_int8_min_scaling_factor() + else: + qtype_traits_min, qtype_traits_max = get_fp8_min_max() + min_scaling_factor = 1.0 / (qtype_traits_max * 512.0) + + qtype_max = float(qtype_traits_max) + + for tile_m in hl.tile(num_tokens, block_size=1): + rms = hl.zeros([tile_m], dtype=torch.float32) + for tile_n in hl.tile(hidden_size): + x_blk = input[tile_m, tile_n].to(torch.float32) + if residual is not None: + x_blk = x_blk + residual[tile_m, tile_n] + rms = rms + x_blk.pow(2).sum(dim=-1) + + rms = torch.rsqrt(rms * (1.0 / hidden_size) + epsilon) + s_blk = hl.zeros([tile_m], dtype=torch.float32) + + for tile_n in hl.tile(hidden_size): + x_blk = input[tile_m, tile_n].to(torch.float32) + if residual is not None: + x_blk = x_blk + residual[tile_m, tile_n] + x_blk = (x_blk * rms[:, None]).to(input.dtype) * weight[None, tile_n] + tmp_blk = torch.amax(torch.abs(x_blk), dim=-1).to(torch.float32) + s_blk = torch.maximum(s_blk, tmp_blk) + + if scale_ub is not None: + scale_ub_s = hl.load(scale_ub, []) + s_blk = s_blk.clamp(max=scale_ub_s) + s_blk = s_blk * (1.0 / qtype_max) + s_blk = s_blk.clamp(min=min_scaling_factor) + scale[tile_m, 0] = s_blk + + for tile_n in hl.tile(hidden_size): + x_blk = input[tile_m, tile_n].to(torch.float32) + if residual is not None: + x_blk = x_blk + residual[tile_m, tile_n] + residual[tile_m, tile_n] = x_blk.to(residual.dtype) + x_blk = (x_blk * rms[:, None]).to(input.dtype) * weight[None, tile_n] + if quant_dtype == torch.int8: + s_inv_blk = 1.0 / s_blk[:, None] + y_blk = x_blk * s_inv_blk + y_blk = y_blk.round() + else: + y_blk = x_blk / s_blk[:, None] + + result[tile_m, tile_n] = y_blk.clamp(qtype_traits_min, qtype_traits_max).to( + result.dtype + ) From eb0fdeb1e83443e1b901cd97dd38cff803488bba Mon Sep 17 00:00:00 2001 From: zhanqiuhu <49648934+ZhanqiuHu@users.noreply.github.com> Date: Wed, 17 Jun 2026 11:17:14 -0400 Subject: [PATCH 27/37] [Bugfix][PD] Fix DSV4 disaggregated serving (#45831) Signed-off-by: ZhanqiuHu --- .../kv_transfer/kv_connector/v1/nixl/base_worker.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py index e587b0cd1fa..74980dabf71 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py @@ -71,6 +71,7 @@ from vllm.v1.kv_cache_interface import ( FullAttentionSpec, MambaSpec, MLAAttentionSpec, + SlidingWindowMLASpec, UniformTypeKVCacheSpecs, ) from vllm.v1.worker.block_table import BlockTable @@ -962,7 +963,9 @@ class NixlBaseConnectorWorker: ) else: self.block_len_per_layer.append(physical_page_size) - is_mla_region = isinstance(layer_spec, MLAAttentionSpec) + is_mla_region = isinstance( + layer_spec, (MLAAttentionSpec, SlidingWindowMLASpec) + ) self._region_is_mla.append(is_mla_region) if not is_mla_region: From 5e27b2baf481518410a9aa5d6b2840cbb2d6ba49 Mon Sep 17 00:00:00 2001 From: danisereb Date: Wed, 17 Jun 2026 18:24:26 +0300 Subject: [PATCH 28/37] [Bugfix] Pass TP group to FlashInfer all-reduce fusion (#45917) Signed-off-by: Daniel Serebrenik --- vllm/config/vllm.py | 9 +-------- .../device_communicators/flashinfer_all_reduce.py | 1 + 2 files changed, 2 insertions(+), 8 deletions(-) diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index b98c74d459a..ba20d75fa11 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -129,13 +129,7 @@ def enable_act_fusion(cfg: "VllmConfig") -> bool: def enable_allreduce_rms_fusion(cfg: "VllmConfig") -> bool: - """Enable if TP > 1, PP == 1, Hopper/Blackwell, and flashinfer installed. - - Gated off for PP > 1: the fused op's GPU-side peer-signal spin-wait - assumes byte-identical kernel launches across TP peers, but concurrent - independent warmup of multiple TP subgroups lets ranks pick divergent - FlashInfer launch configs and deadlock. - """ + """Enable if TP > 1 and Hopper/Blackwell and flashinfer installed.""" from vllm.platforms import current_platform from vllm.utils.flashinfer import has_flashinfer @@ -148,7 +142,6 @@ def enable_allreduce_rms_fusion(cfg: "VllmConfig") -> bool: return ( cfg.parallel_config.tensor_parallel_size > 1 - and cfg.parallel_config.pipeline_parallel_size == 1 and current_platform.is_cuda() and has_flashinfer() and ( diff --git a/vllm/distributed/device_communicators/flashinfer_all_reduce.py b/vllm/distributed/device_communicators/flashinfer_all_reduce.py index 2594c0cf160..38f7bd5ff8d 100644 --- a/vllm/distributed/device_communicators/flashinfer_all_reduce.py +++ b/vllm/distributed/device_communicators/flashinfer_all_reduce.py @@ -61,6 +61,7 @@ def _create_workspace( hidden_dim=hidden_dim, dtype=dtype, comm_backend=comm_backend, + group=group, ) except Exception as e: if "multicast" in str(e).lower(): From 9c7c74bf1023774cf0e247bf77a038989ee272db Mon Sep 17 00:00:00 2001 From: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Date: Wed, 17 Jun 2026 11:34:22 -0400 Subject: [PATCH 29/37] [Log] Update deepgemm log (#45857) Signed-off-by: yewentao256 --- tests/kernels/moe/modular_kernel_tools/common.py | 5 ++++- vllm/model_executor/layers/sparse_attn_indexer.py | 3 ++- vllm/utils/deep_gemm.py | 9 +++++---- 3 files changed, 11 insertions(+), 6 deletions(-) diff --git a/tests/kernels/moe/modular_kernel_tools/common.py b/tests/kernels/moe/modular_kernel_tools/common.py index 646cff4c2d2..8041db68d75 100644 --- a/tests/kernels/moe/modular_kernel_tools/common.py +++ b/tests/kernels/moe/modular_kernel_tools/common.py @@ -328,7 +328,10 @@ class Config: if self.needs_deep_ep_v2() and not has_deep_ep_v2(): return False, "Needs DeepEP v2, but DeepEP v2 not available." if self.needs_deep_gemm() and not has_deep_gemm(): - return False, "Needs DeepGEMM, but DeepGEMM not available." + return ( + False, + "Needs DeepGEMM, but the current vLLM environment does not provide it.", + ) if self.needs_aiter() and not has_aiter(): # noqa: SIM103 return False, "Needs Aiter, but Aiter not available." if self.needs_mori() and not has_mori(): # noqa: SIM103 diff --git a/vllm/model_executor/layers/sparse_attn_indexer.py b/vllm/model_executor/layers/sparse_attn_indexer.py index 9597708b62e..45c5d5f7819 100644 --- a/vllm/model_executor/layers/sparse_attn_indexer.py +++ b/vllm/model_executor/layers/sparse_attn_indexer.py @@ -442,7 +442,8 @@ class SparseAttnIndexer(CustomOp): self.use_fp4_cache = use_fp4_cache if current_platform.is_cuda() and not has_deep_gemm(): raise RuntimeError( - "Sparse Attention Indexer CUDA op requires DeepGEMM to be installed." + "Sparse Attention Indexer CUDA op requires DeepGEMM support in " + "the current vLLM environment." ) def forward_native( diff --git a/vllm/utils/deep_gemm.py b/vllm/utils/deep_gemm.py index 4252ce87754..3c884aad6cd 100644 --- a/vllm/utils/deep_gemm.py +++ b/vllm/utils/deep_gemm.py @@ -120,8 +120,9 @@ def is_deep_gemm_e8m0_used() -> bool: def _missing(*_: Any, **__: Any) -> NoReturn: """Placeholder for unavailable DeepGEMM backend.""" raise RuntimeError( - "DeepGEMM backend is not available or outdated. Please install or " - "update the `deep_gemm` to a newer version to enable FP8 kernels." + "DeepGEMM backend is unavailable in the current vLLM environment, " + "or the available DeepGEMM package does not provide the required APIs " + "for these kernels." ) @@ -156,7 +157,7 @@ def _import_deep_gemm(): logger.debug_once("Imported deep_gemm module from site-packages") return module except ImportError: - logger.debug_once( + logger.info_once( "deep_gemm not found in site-packages, " "trying vendored vllm.third_party.deep_gemm" ) @@ -167,7 +168,7 @@ def _import_deep_gemm(): logger.debug_once("Imported deep_gemm module from vllm.third_party.deep_gemm") return module except ImportError: - logger.debug_once("Vendored deep_gemm not found either") + logger.info_once("Vendored deep_gemm not found either") except Exception as e: # The vendored module may raise RuntimeError during _C.init() # if JIT include files are missing (e.g. incomplete wheel). From 2a47a9ff0f4f302ee64915dfea64642757e31ee4 Mon Sep 17 00:00:00 2001 From: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Date: Wed, 17 Jun 2026 12:34:53 -0400 Subject: [PATCH 30/37] [DSV4 Perf] Optimize dsv4 cudagraph by reducing `eager_break_during_capture`, 26.8% ~ 27.9% E2E TTFT improvement (#45309) Signed-off-by: yewentao256 --- vllm/models/deepseek_v4/attention.py | 67 +++++++++++++++------------- 1 file changed, 37 insertions(+), 30 deletions(-) diff --git a/vllm/models/deepseek_v4/attention.py b/vllm/models/deepseek_v4/attention.py index 29302584880..052874fac03 100644 --- a/vllm/models/deepseek_v4/attention.py +++ b/vllm/models/deepseek_v4/attention.py @@ -14,7 +14,7 @@ import torch.nn.functional as F from transformers import DeepseekV2Config, DeepseekV3Config import vllm.envs as envs -from vllm.compilation.breakable_cudagraph import eager_break_during_capture +from vllm.compilation.breakable_cudagraph import BreakableCUDAGraphCapture from vllm.model_executor.layers.linear import ( ColumnParallelLinear, MergedColumnParallelLinear, @@ -331,8 +331,8 @@ class DeepseekV4Attention(nn.Module, AttentionLayerBase, ABC): ) # Metadata-independent input GEMMs + RMSNorm stay in the captured - # graph; the metadata-dependent rest (q up-proj + kv-insert, indexer, - # compressor, MLA attention) runs in the eager break. + # graph. For C4A layers, the inner sparse_attn_indexer custom op + # runs in the eager break. qr_kv, kv_score, indexer_kv_score, indexer_weights = ( self.attn_gemm_parallel_execute(hidden_states) ) @@ -345,9 +345,6 @@ class DeepseekV4Attention(nn.Module, AttentionLayerBase, ABC): self.eps, ) - # attention_impl is wrapped with @eager_break_during_capture: this is - # where the breakable cudagraph capture breaks (the attention op runs - # eagerly between captured graph segments). self.attention_impl( hidden_states, qr, @@ -423,7 +420,6 @@ class DeepseekV4Attention(nn.Module, AttentionLayerBase, ABC): return qr_kv, kv_score, indexer_kv_score, indexer_weights - @eager_break_during_capture def attention_impl( self, hidden_states: torch.Tensor, @@ -451,31 +447,42 @@ class DeepseekV4Attention(nn.Module, AttentionLayerBase, ABC): def wq_b_kv_insert() -> torch.Tensor: q = self.wq_b(qr).view(-1, self.n_local_heads, self.head_dim) - q = self._fused_qnorm_rope_kv_insert(q, kv, positions, attn_metadata) - return q + return self._fused_qnorm_rope_kv_insert(q, kv, positions, attn_metadata) - # 3-way overlap (matches TRT-LLM PR #14142 Level 1): default runs - # wq_b+kv_insert; slot [0] runs the full indexer; slot [1] runs the - # MLA compressor. Slot [2] is reserved for the indexer's inner - # overlap. ROCm (aux_streams is None) falls back to sequential. - q, _ = execute_in_parallel( - wq_b_kv_insert, - [ - lambda: indexer( - hidden_states, - qr, - indexer_kv_score, - indexer_weights, - positions, - self.indexer_rotary_emb, - ), - lambda: compressor(kv_score, positions, self.rotary_emb), - ], - self.ln_events[0], - [self.ln_events[1], self.ln_events[2]], - [aux_streams[0], aux_streams[1]] if aux_streams is not None else None, - enable=aux_streams is not None, + run_indexer = lambda: indexer( + hidden_states, + qr, + indexer_kv_score, + indexer_weights, + positions, + self.indexer_rotary_emb, ) + run_compressor = lambda: compressor(kv_score, positions, self.rotary_emb) + + if BreakableCUDAGraphCapture.is_active(): + q, _ = maybe_execute_in_parallel( + wq_b_kv_insert, + run_compressor, + self.ln_events[0], + self.ln_events[1], + aux_streams[1] if aux_streams is not None else None, + ) + run_indexer() + else: + # 3-way overlap (matches TRT-LLM PR #14142 Level 1): default runs + # wq_b+kv_insert; slot [0] runs the full indexer; slot [1] runs the + # MLA compressor. Slot [2] is reserved for the indexer's inner + # overlap. ROCm (aux_streams is None) falls back to sequential. + q, _ = execute_in_parallel( + wq_b_kv_insert, + [run_indexer, run_compressor], + self.ln_events[0], + [self.ln_events[1], self.ln_events[2]], + [aux_streams[0], aux_streams[1]] + if aux_streams is not None + else None, + enable=aux_streams is not None, + ) elif self.compressor is not None: # wq_b + kv_insert on default, compressor on aux. aux_stream = ( From d112eb1ac78ede18c60ad98f65da238ca3e052b4 Mon Sep 17 00:00:00 2001 From: qli88 Date: Wed, 17 Jun 2026 13:50:48 -0500 Subject: [PATCH 31/37] [feature] MiniMax-M3-MXFP4 support added (#45896) Signed-off-by: Qiang Li --- .../experts/gpt_oss_triton_kernels_moe.py | 14 ++++++++++++++ .../layers/fused_moe/oracle/mxfp4.py | 13 ++++++++++++- .../layers/quantization/quark/quark_moe.py | 6 ++++++ vllm/models/minimax_m3/amd/model.py | 14 ++++++++++++-- 4 files changed, 44 insertions(+), 3 deletions(-) diff --git a/vllm/model_executor/layers/fused_moe/experts/gpt_oss_triton_kernels_moe.py b/vllm/model_executor/layers/fused_moe/experts/gpt_oss_triton_kernels_moe.py index a7f31afc5ef..31ef144e237 100644 --- a/vllm/model_executor/layers/fused_moe/experts/gpt_oss_triton_kernels_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/gpt_oss_triton_kernels_moe.py @@ -755,6 +755,7 @@ class UnfusedOAITritonExperts(LoRAExpertsMixin, BaseOAITritonExperts): MoEActivation.GELU, MoEActivation.SWIGLUOAI, MoEActivation.SWIGLUSTEP, + MoEActivation.SWIGLUOAI_UNINTERLEAVE, ] @staticmethod @@ -811,6 +812,19 @@ class UnfusedOAITritonExperts(LoRAExpertsMixin, BaseOAITritonExperts): input, quant_config.gemm1_clamp_limit, ) + elif activation == MoEActivation.SWIGLUOAI_UNINTERLEAVE: + assert quant_config.gemm1_clamp_limit is not None + alpha = ( + quant_config.gemm1_alpha + if quant_config.gemm1_alpha is not None + else 1.0 + ) + beta = ( + quant_config.gemm1_beta if quant_config.gemm1_beta is not None else 0.0 + ) + torch.ops._C.silu_and_mul_with_clamp( + output, input, quant_config.gemm1_clamp_limit, alpha, beta + ) else: super().activation(activation, output, input) diff --git a/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py b/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py index 87c44d92fd5..ab76cea1327 100644 --- a/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py +++ b/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py @@ -503,7 +503,18 @@ def select_mxfp4_moe_backend( activation_format, ) - if current_platform.is_cuda() or current_platform.is_rocm(): + if current_platform.is_rocm(): + backend = Mxfp4MoeBackend.TRITON_UNFUSED + logger.info_once(_make_log_backend(backend)) + return _return_or_raise( + Mxfp4MoeBackend.TRITON_UNFUSED, + config, + kMxfp4Static, + None, + activation_format, + ) + + if current_platform.is_cuda(): raise NotImplementedError( "No MXFP4 MoE backend supports the deployment configuration. " f"weight_key=kMxfp4Static, activation_key={activation_key}. " diff --git a/vllm/model_executor/layers/quantization/quark/quark_moe.py b/vllm/model_executor/layers/quantization/quark/quark_moe.py index 9ee901a9910..703fc815015 100644 --- a/vllm/model_executor/layers/quantization/quark/quark_moe.py +++ b/vllm/model_executor/layers/quantization/quark/quark_moe.py @@ -1303,6 +1303,9 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod): w2_bias=getattr(layer, "w2_bias", None), a1_scale=getattr(layer, "w13_input_scale", None), a2_scale=getattr(layer, "w2_input_scale", None), + gemm1_alpha=getattr(layer, "swiglu_alpha", None), + gemm1_beta=getattr(layer, "swiglu_beta", None), + swiglu_limit=getattr(layer, "swiglu_limit", None), ) # Emulation and other schemes @@ -1339,6 +1342,9 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod): a1_scale=None, a2_scale=None, block_shape=None, + gemm1_alpha=getattr(layer, "swiglu_alpha", None), + gemm1_beta=getattr(layer, "swiglu_beta", None), + gemm1_clamp_limit=getattr(layer, "swiglu_limit", None), ) @property diff --git a/vllm/models/minimax_m3/amd/model.py b/vllm/models/minimax_m3/amd/model.py index 4bd92444728..54d1beb7d4d 100644 --- a/vllm/models/minimax_m3/amd/model.py +++ b/vllm/models/minimax_m3/amd/model.py @@ -924,6 +924,11 @@ class MiniMaxM3Model(nn.Module, EagleModelMixin): class MiniMaxM3SparseForCausalLM(nn.Module, SupportsEagle3): """MiniMax M3 (sparse/dense backbone) for causal language modeling.""" + 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_text_config @@ -987,14 +992,19 @@ class MiniMaxM3SparseForConditionalGeneration( # ranks (see ``_process_image_input`` / ``_process_video_input``). supports_encoder_tp_data = True + packed_modules_mapping = { + "qkv_proj": ["q_proj", "k_proj", "v_proj"], + "gate_up_proj": ["gate_proj", "up_proj"], + } + hf_to_vllm_mapper = WeightsMapper( orig_to_new_prefix={ "multi_modal_projector.": "vision_tower.multi_modal_projector.", "patch_merge_mlp.": "vision_tower.patch_merge_mlp.", }, orig_to_new_substr={ - ".mlp.fc1.": ".fc1.", - ".mlp.fc2.": ".fc2.", + ".mlp.fc1": ".fc1", + ".mlp.fc2": ".fc2", }, ) From 091386a99b9542691bb1e935ca44d0efbba6e111 Mon Sep 17 00:00:00 2001 From: wangjiaxin99 Date: Thu, 18 Jun 2026 03:15:46 +0800 Subject: [PATCH 32/37] =?UTF-8?q?[Bugfix]=20MiniMax-M3=20(AMD):=20add=20pa?= =?UTF-8?q?cked=5Fmodules=5Fmapping=20and=20pass=20swiglu=E2=80=A6=20(#457?= =?UTF-8?q?94)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: wangjiaxin99 Co-authored-by: TJian Co-authored-by: Douglas Lehr <91553416+dllehr-amd@users.noreply.github.com> From 2659f60a1a240243baf17e32fce2318f99ac3d49 Mon Sep 17 00:00:00 2001 From: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Date: Wed, 17 Jun 2026 16:12:01 -0400 Subject: [PATCH 33/37] [Refactor] Remove dead quantization code and tests (#45454) Signed-off-by: yewentao256 --- tests/kernels/quantization/test_awq.py | 20 ------ tests/quantization/fp_quant.py | 32 ---------- tests/quantization/test_fp8.py | 64 ------------------- .../layers/quantization/__init__.py | 1 - .../model_executor/layers/quantization/fp8.py | 3 - 5 files changed, 120 deletions(-) delete mode 100644 tests/quantization/fp_quant.py diff --git a/tests/kernels/quantization/test_awq.py b/tests/kernels/quantization/test_awq.py index 3bf59dea309..a8977958023 100644 --- a/tests/kernels/quantization/test_awq.py +++ b/tests/kernels/quantization/test_awq.py @@ -27,23 +27,3 @@ def test_awq_dequantize_opcheck(monkeypatch: pytest.MonkeyPatch): torch.ops._C.awq_dequantize, (qweight, scales, zeros, split_k_iters, thx, thy), ) - - -@pytest.mark.skip(reason="Not working; needs investigation.") -@pytest.mark.skipif( - not hasattr(torch.ops._C, "awq_gemm"), - reason="AWQ is not supported on this GPU type.", -) -def test_awq_gemm_opcheck(monkeypatch: pytest.MonkeyPatch): - with monkeypatch.context() as m: - m.setenv("VLLM_USE_TRITON_AWQ", "0") - input = torch.rand((2, 8192), device="cuda", dtype=torch.float16) - qweight = torch.randint( - -2000000000, 2000000000, (8192, 256), device="cuda", dtype=torch.int32 - ) - scales = torch.empty((64, 2048), device="cuda", dtype=torch.float16) - qzeros = torch.randint( - -2000000000, 2000000000, (64, 256), device="cuda", dtype=torch.int32 - ) - split_k_iters = 8 - opcheck(torch.ops._C.awq_gemm, (input, qweight, scales, qzeros, split_k_iters)) diff --git a/tests/quantization/fp_quant.py b/tests/quantization/fp_quant.py deleted file mode 100644 index 664ce9d111e..00000000000 --- a/tests/quantization/fp_quant.py +++ /dev/null @@ -1,32 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Test model set-up and inference for quantized HF models supported -on the GPU backend using FPQuant. - -Validating the configuration and printing results for manual checking. - -Run `pytest tests/quantization/test_fp_quant.py`. -""" - -import pytest - -from tests.quantization.utils import is_quant_method_supported - -MODELS = [ - "ISTA-DASLab/Qwen3-0.6B-RTN-NVFP4", - "ISTA-DASLab/Qwen3-0.6B-RTN-MXFP4", -] -DTYPE = ["bfloat16"] -EAGER = [True, False] - - -@pytest.mark.skipif( - not is_quant_method_supported("fp_quant"), - reason="FPQuant is not supported on this GPU type.", -) -@pytest.mark.parametrize("model", MODELS) -@pytest.mark.parametrize("eager", EAGER) -def test_fpquant(vllm_runner, model, eager): - with vllm_runner(model, enforce_eager=eager) as llm: - output = llm.generate_greedy(["1 2 3 4 5"], max_tokens=2) - assert output[0][1] == "1 2 3 4 5 6" diff --git a/tests/quantization/test_fp8.py b/tests/quantization/test_fp8.py index 4c789059c7f..499955c9f63 100644 --- a/tests/quantization/test_fp8.py +++ b/tests/quantization/test_fp8.py @@ -67,70 +67,6 @@ def test_model_load_and_run( print(outputs[0][1]) -KV_CACHE_MODELS = [ - # AutoFP8 format using separate .k_scale and .v_scale - # The original checkpoint below was removed from the Hub. To unblock CI and - # until a small replacement with split K/V scales is found, skip this case. - # See PR #27717 for context. - pytest.param( - "nm-testing/Qwen2-1.5B-Instruct-FP8-K-V", - marks=pytest.mark.skip( - reason=( - "Checkpoint removed from HF; temporarily disabling this " - "AutoFP8 split K/V case (PR #27717)." - ) - ), - ), -] - - -@pytest.mark.skipif( - not is_quant_method_supported("fp8"), - reason="FP8 is not supported on this GPU type.", -) -@pytest.mark.parametrize("model_id", KV_CACHE_MODELS) -@pytest.mark.parametrize( - "use_rocm_aiter", [True, False] if current_platform.is_rocm() else [False] -) -def test_kv_cache_model_load_and_run( - vllm_runner, model_id: str, use_rocm_aiter: bool, monkeypatch -): - if use_rocm_aiter: - monkeypatch.setenv("VLLM_ROCM_USE_AITER", "1") - - # `LLM.apply_model` requires pickling a function. - monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1") - with vllm_runner(model_id, kv_cache_dtype="fp8", enforce_eager=True) as llm: - - def check_model(model): - attn = model.model.layers[0].self_attn.attn - - assert isinstance(attn.quant_method, Fp8KVCacheMethod) - - if not current_platform.is_rocm(): - # NOTE: This code path requires validation on Non-CUDA platform - # NOTE: it is valid for scales to be 1.0 (default value), but - # we know these checkpoints have scales < 1.0 - assert 0.0 < attn._k_scale < 1.0 - assert 0.0 < attn._v_scale < 1.0 - else: - # NOTE: This code path is for ROCm platform - # NOTE: it is valid for scales to be 1.0 (default value), but - # we know these checkpoints have scales < 1.0 - # However on ROCm platform, the _k_scale and _v_scale will be - # scaled by a factor of 2 as described in - # vllm/model_executor/layers/quantization/kv_cache.py - assert 0.0 < attn._k_scale < (1.0 * 2.0) - assert 0.0 < attn._v_scale < (1.0 * 2.0) - - llm.apply_model(check_model) - - # note: this does not test accuracy, just that we can run through - # see lm-eval tests for accuracy - outputs = llm.generate_greedy(["Hello my name is"], max_tokens=4) - print(outputs[0][1]) - - @pytest.mark.skipif( not is_quant_method_supported("fp8"), reason="FP8 is not supported on this GPU type.", diff --git a/vllm/model_executor/layers/quantization/__init__.py b/vllm/model_executor/layers/quantization/__init__.py index f47dcae310a..53f4e7d2a8a 100644 --- a/vllm/model_executor/layers/quantization/__init__.py +++ b/vllm/model_executor/layers/quantization/__init__.py @@ -156,7 +156,6 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]: "quark": QuarkConfig, "moe_wna16": MoeWNA16Config, "torchao": TorchAOConfig, - "auto-round": INCConfig, "inc": INCConfig, "mxfp4": Mxfp4Config, "gpt_oss_mxfp4": GptOssMxfp4Config, diff --git a/vllm/model_executor/layers/quantization/fp8.py b/vllm/model_executor/layers/quantization/fp8.py index 6fc9b0d8f65..7cdb04cfbec 100644 --- a/vllm/model_executor/layers/quantization/fp8.py +++ b/vllm/model_executor/layers/quantization/fp8.py @@ -486,9 +486,6 @@ class Fp8LinearMethod(LinearMethodBase): weight_bf16 = weight_fp8 * weight_scale return torch.nn.functional.linear(x, weight_bf16.t(), bias) - if self.use_marlin: - return self.fp8_linear.apply_weights(layer, x, bias) - return self.fp8_linear.apply_weights(layer, x, bias) From 58b2e896423ffb255fa8b3a2c6b283dd99d7dbad Mon Sep 17 00:00:00 2001 From: Luciano Martins <22145370+lucianommartins@users.noreply.github.com> Date: Wed, 17 Jun 2026 17:44:15 -0300 Subject: [PATCH 34/37] [Bugfix][Gemma4] Render reasoning on assistant turns without tool_calls (#45867) Signed-off-by: Luciano Martins Co-authored-by: Luciano Martins --- examples/tool_chat_template_gemma4.jinja | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/examples/tool_chat_template_gemma4.jinja b/examples/tool_chat_template_gemma4.jinja index 9d603aa0b06..6ce01e6479a 100644 --- a/examples/tool_chat_template_gemma4.jinja +++ b/examples/tool_chat_template_gemma4.jinja @@ -231,10 +231,10 @@ {%- endif -%} {%- endif -%} - {#- Render reasoning/reasoning_content as thinking channel (tool-call turns only) -#} + {#- Render reasoning/reasoning_content as thinking channel -#} {%- set thinking_text = message.get('reasoning') or message.get('reasoning_content') -%} {%- set thinking_gate = (loop.index0 > ns_turn.last_user_idx) or preserve_thinking -%} - {%- if thinking_text and thinking_gate and message.get('tool_calls') -%} + {%- if thinking_text and thinking_gate -%} {{- '<|channel>thought\n' + thinking_text + '\n' -}} {%- endif -%} From 9d4b87f4f0bb489c66c466b0b8ae6bbb212f0a2f Mon Sep 17 00:00:00 2001 From: Ting SUN Date: Thu, 18 Jun 2026 05:46:33 +0800 Subject: [PATCH 35/37] [Bugfix][Model] Validate DefaultModelLoader / LoadConfig and fail with clear errors (#45196) Signed-off-by: Ting Sun --- .../model_loader/test_registry.py | 55 +++++++++++++++++++ tests/test_config.py | 11 ++++ vllm/config/load.py | 9 ++- vllm/engine/arg_utils.py | 5 +- .../model_loader/default_loader.py | 35 +++++++++++- 5 files changed, 108 insertions(+), 7 deletions(-) diff --git a/tests/model_executor/model_loader/test_registry.py b/tests/model_executor/model_loader/test_registry.py index 020988ccac1..95b797bb514 100644 --- a/tests/model_executor/model_loader/test_registry.py +++ b/tests/model_executor/model_loader/test_registry.py @@ -8,6 +8,7 @@ from vllm.config import ModelConfig from vllm.config.load import LoadConfig from vllm.model_executor.model_loader import get_model_loader, register_model_loader from vllm.model_executor.model_loader.base_loader import BaseModelLoader +from vllm.model_executor.model_loader.default_loader import DefaultModelLoader @register_model_loader("custom_load_format") @@ -33,3 +34,57 @@ def test_invalid_model_loader(): @register_model_loader("invalid_load_format") class InValidModelLoader: pass + + +def test_default_loader_rejects_zero_num_threads(): + # num_threads=0 used to fail late in ThreadPoolExecutor ("max_workers must be > 0"). + with pytest.raises(ValueError, match="num_threads"): + DefaultModelLoader( + LoadConfig( + model_loader_extra_config={ + "enable_multithread_load": True, + "num_threads": 0, + } + ) + ) + + +def test_default_loader_rejects_multithread_with_non_lazy_strategy(): + # The multi-thread loader ignores safetensors_load_strategy; reject the + # combination instead of silently dropping the requested strategy. + with pytest.raises(ValueError, match="does not support"): + DefaultModelLoader( + LoadConfig( + safetensors_load_strategy="torchao", + model_loader_extra_config={"enable_multithread_load": True}, + ) + ) + + +def test_default_loader_explicit_safetensors_does_not_misread_pt(tmp_path): + # Explicit safetensors must not fall back to a .pt and open it as safetensors. + (tmp_path / "model.pt").write_bytes(b"\x00\x00\x00\x00") + loader = DefaultModelLoader(LoadConfig(load_format="safetensors")) + with pytest.raises(RuntimeError, match="Cannot find any model weights"): + loader._prepare_weights( + str(tmp_path), + None, + None, + fall_back_to_pt=True, + allow_patterns_overrides=None, + ) + + +def test_default_loader_hf_still_falls_back_to_pt(tmp_path): + # Control: load_format="hf" still picks up .pt weights via fallback. + (tmp_path / "model.pt").write_bytes(b"\x00\x00\x00\x00") + loader = DefaultModelLoader(LoadConfig(load_format="hf")) + _, files, use_safetensors = loader._prepare_weights( + str(tmp_path), + None, + None, + fall_back_to_pt=True, + allow_patterns_overrides=None, + ) + assert use_safetensors is False + assert any(f.endswith("model.pt") for f in files) diff --git a/tests/test_config.py b/tests/test_config.py index e6c122de83d..d992ac29696 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1557,3 +1557,14 @@ def test_ir_op_priority_ctx(): # context restored even after exception assert ir.ops.rms_norm.get_priority() == ["vllm_c", "native"] assert ir.ops.fused_add_rms_norm.get_priority() == ["native"] + + +def test_load_config_rejects_invalid_safetensors_load_strategy(): + with pytest.raises(pydantic.ValidationError): + LoadConfig(safetensors_load_strategy="not_a_real_strategy") + + +@pytest.mark.parametrize("bad_load_format", [None, 123]) +def test_load_config_rejects_non_string_load_format(bad_load_format): + with pytest.raises(pydantic.ValidationError): + LoadConfig(load_format=bad_load_format) diff --git a/vllm/config/load.py b/vllm/config/load.py index ed591a2299f..f1066c2b9ad 100644 --- a/vllm/config/load.py +++ b/vllm/config/load.py @@ -1,7 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Literal, TypeAlias from pydantic import Field, field_validator @@ -11,12 +11,11 @@ from vllm.utils.hashing import safe_hash DEFAULT_SAFETENSORS_PREFETCH_NUM_THREADS = 8 DEFAULT_SAFETENSORS_PREFETCH_BLOCK_SIZE = 16 * 1024 * 1024 +SafetensorsLoadStrategy: TypeAlias = Literal["lazy", "eager", "prefetch", "torchao"] if TYPE_CHECKING: - from vllm.model_executor.model_loader import LoadFormats from vllm.model_executor.model_loader.tensorizer import TensorizerConfig else: - LoadFormats = Any TensorizerConfig = Any logger = init_logger(__name__) @@ -26,7 +25,7 @@ logger = init_logger(__name__) class LoadConfig: """Configuration for loading the model weights.""" - load_format: str | LoadFormats = "auto" + load_format: str = "auto" """ The format of the model weights to load. @@ -59,7 +58,7 @@ class LoadConfig: download_dir: str | None = None """Directory to download and load the weights, default to the default cache directory of Hugging Face.""" - safetensors_load_strategy: str | None = None + safetensors_load_strategy: SafetensorsLoadStrategy | None = None """ Specifies the loading strategy for safetensors weights. diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index 85a59bebf8a..9172a8728a0 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -73,6 +73,7 @@ from vllm.config.cache import ( ) from vllm.config.device import Device from vllm.config.kernel import IrOpPriorityConfig, LinearBackend, MoEBackend +from vllm.config.load import SafetensorsLoadStrategy from vllm.config.lora import MaxLoRARanks from vllm.config.mamba import MambaBackendEnum from vllm.config.model import ( @@ -427,7 +428,9 @@ class EngineArgs: allowed_local_media_path: str = ModelConfig.allowed_local_media_path allowed_media_domains: list[str] | None = ModelConfig.allowed_media_domains download_dir: str | None = LoadConfig.download_dir - safetensors_load_strategy: str | None = LoadConfig.safetensors_load_strategy + safetensors_load_strategy: SafetensorsLoadStrategy | None = ( + LoadConfig.safetensors_load_strategy + ) safetensors_prefetch_num_threads: int = LoadConfig.safetensors_prefetch_num_threads safetensors_prefetch_block_size: int = LoadConfig.safetensors_prefetch_block_size load_format: str | LoadFormats = LoadConfig.load_format diff --git a/vllm/model_executor/model_loader/default_loader.py b/vllm/model_executor/model_loader/default_loader.py index 43d5d4a4496..3ea76f4d9b3 100644 --- a/vllm/model_executor/model_loader/default_loader.py +++ b/vllm/model_executor/model_loader/default_loader.py @@ -76,6 +76,11 @@ class DefaultModelLoader(BaseModelLoader): self.local_expert_ids: set[int] | None = None extra_config = load_config.model_loader_extra_config + if not isinstance(extra_config, dict): + raise ValueError( + f"model_loader_extra_config must be a dict for load format " + f"{load_config.load_format}, got {type(extra_config).__name__}" + ) allowed_keys = { "enable_multithread_load", "num_threads", @@ -90,10 +95,36 @@ class DefaultModelLoader(BaseModelLoader): f"{unexpected_keys}" ) + enable_multithread_load = extra_config.get("enable_multithread_load", False) + if not isinstance(enable_multithread_load, bool): + raise ValueError( + f"enable_multithread_load must be a bool, got " + f"{type(enable_multithread_load).__name__}" + ) + num_threads = extra_config.get("num_threads") + if num_threads is not None and not ( + isinstance(num_threads, int) and num_threads > 0 + ): + raise ValueError( + f"num_threads must be a positive integer, got {num_threads!r}" + ) + self.enable_weights_track: bool | None = extra_config.get( "enable_weights_track", None ) + # The multi-thread loader ignores safetensors_load_strategy, so reject + # the combination instead of silently dropping the requested strategy. + if extra_config.get("enable_multithread_load") and ( + load_config.safetensors_load_strategy not in (None, "lazy") + ): + raise ValueError( + "enable_multithread_load does not support " + "safetensors_load_strategy=" + f"{load_config.safetensors_load_strategy!r}; the multi-thread " + "loader only implements the default lazy strategy." + ) + def _prepare_weights( self, model_name_or_path: str, @@ -152,7 +183,9 @@ class DefaultModelLoader(BaseModelLoader): else: raise ValueError(f"Unknown load_format: {load_format}") - if fall_back_to_pt: + # Don't fall back to .pt for explicit safetensors formats; otherwise a + # .pt file is matched and later opened as safetensors. + if fall_back_to_pt and not use_safetensors: allow_patterns += ["*.pt"] if allow_patterns_overrides is not None: From 5fd21eb0b291e16dad66da740b0faf50666a78a5 Mon Sep 17 00:00:00 2001 From: shanjiaz Date: Wed, 17 Jun 2026 18:02:24 -0400 Subject: [PATCH 36/37] [BUG] fix hidden states nan for hybrid attention models (#45849) Signed-off-by: shanjiaz Co-authored-by: shanjiaz --- .../kv_connector/v1/example_hidden_states_connector.py | 7 +++++++ vllm/v1/core/single_type_kv_cache_manager.py | 2 ++ 2 files changed, 9 insertions(+) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/example_hidden_states_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/example_hidden_states_connector.py index 696d3f7fb4c..7e6c95bf8fb 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/example_hidden_states_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/example_hidden_states_connector.py @@ -150,7 +150,14 @@ class ExampleHiddenStatesConnector(KVConnectorBase_V1, SupportsHMA): # Worker-side state (set by register_kv_caches). self._kv_cache: torch.Tensor | None = None + + # Identify which KV cache group holds the hidden-states layer. self._hs_group_idx: int = 0 + if self._kv_cache_config is not None: + for i, group in enumerate(self._kv_cache_config.kv_cache_groups): + if any("cache_only_layers" in n for n in group.layer_names): + self._hs_group_idx = i + break # Only TP rank 0 writes hidden states to disk; other TP ranks no-op. # Set in register_kv_caches (after distributed init). self._is_tp_rank_zero: bool = True diff --git a/vllm/v1/core/single_type_kv_cache_manager.py b/vllm/v1/core/single_type_kv_cache_manager.py index 580b58814b4..c98c59017c5 100644 --- a/vllm/v1/core/single_type_kv_cache_manager.py +++ b/vllm/v1/core/single_type_kv_cache_manager.py @@ -271,6 +271,7 @@ class SingleTypeKVCacheManager(ABC): FullAttentionSpec, TQFullAttentionSpec, MLAAttentionSpec, + HiddenStateCacheSpec, ): self.new_block_ids.extend(b.block_id for b in allocated_blocks) @@ -303,6 +304,7 @@ class SingleTypeKVCacheManager(ABC): FullAttentionSpec, TQFullAttentionSpec, MLAAttentionSpec, + HiddenStateCacheSpec, ): self.new_block_ids.extend(b.block_id for b in new_blocks) return new_blocks From 0d339cf13551bd4e6c8f4c1cfce1f47c806aa5e3 Mon Sep 17 00:00:00 2001 From: Bryan Shan <58582368+Oseltamivir@users.noreply.github.com> Date: Wed, 17 Jun 2026 15:11:29 -0700 Subject: [PATCH 37/37] [Bugfix] Fix NixlConnector handshake block_len validation for GQA-replicated KV heads (#45879) Signed-off-by: Oseltamivir <58582368+Oseltamivir@users.noreply.github.com> Co-authored-by: waynehacking8 --- .../kv_connector/unit/test_nixl_connector.py | 133 ++++++++++++++++-- .../kv_connector/v1/nixl/base_worker.py | 23 ++- 2 files changed, 141 insertions(+), 15 deletions(-) diff --git a/tests/v1/kv_connector/unit/test_nixl_connector.py b/tests/v1/kv_connector/unit/test_nixl_connector.py index 32652118d52..cd13efd4512 100644 --- a/tests/v1/kv_connector/unit/test_nixl_connector.py +++ b/tests/v1/kv_connector/unit/test_nixl_connector.py @@ -517,17 +517,17 @@ class FakeNixlConnectorWorker(NixlConnectorWorker): assert expected_engine_id == self.REMOTE_ENGINE_ID # Adjust remote block length metadata to satisfy heterogeneous TP - # invariants enforced during handshake validation. + # invariants enforced during handshake validation. Use per-rank + # head ratio (not tp_ratio) to account for GQA replication capping. remote_block_lens = list(self.block_len_per_layer) tp_ratio = self.transfer_topo.tp_ratio(remote_tp_size) - if remote_tp_size > self.world_size: - # P TP > D TP case, block_len of remote is smaller + total_kv = self.transfer_topo.total_num_kv_heads + local_heads = self.transfer_topo.local_physical_heads + remote_heads = max(1, total_kv // remote_tp_size) + if remote_tp_size != self.world_size: remote_block_lens = [ - block_len // (-tp_ratio) for block_len in remote_block_lens - ] - elif remote_tp_size < self.world_size: - remote_block_lens = [ - block_len * tp_ratio for block_len in remote_block_lens + block_len * remote_heads // local_heads + for block_len in remote_block_lens ] # When remote tp_size > local tp_size, handshake with multiple @@ -1144,6 +1144,123 @@ class TestNixlHandshake: with pytest.raises(AssertionError): worker2.add_remote_agent(bad_meta, remote_tp_size=1) + @patch( + "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper", + FakeNixlWrapper, + ) + def test_handshake_validates_gqa_replicated_block_len( + self, default_vllm_config, dist_init + ): + """Regression test for #45330. + + When tp_size > total_num_kv_heads, GQA replication caps per-rank + KV heads at 1, so block_len stops scaling with 1/tp. With 8 KV + heads and D_TP=16 pulling from P_TP=8, both sides hold one head + per rank and report the *same* block_len; the old validation + expected local_block_len * tp_ratio and rejected the valid + handshake. + """ + vllm_config = create_vllm_config() + + with patch( + "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size", # noqa: E501 + return_value=16, + ): + connector = NixlConnector( + vllm_config, + KVConnectorRole.WORKER, + make_kv_cache_config(block_size=16), + ) + connector.connector_worker = FakeNixlConnectorWorker( + vllm_config, connector.engine_id, hand_shake_latency=0 + ) + worker = connector.connector_worker + + worker.transfer_topo.total_num_kv_heads = 8 + worker.transfer_topo.local_physical_heads = 1 + worker.kv_cache_layout = "HND" + + worker.slot_size_per_layer = [4096] + worker.block_len_per_layer = [4096 * worker.block_size] + worker.num_blocks = 1 + worker.dst_num_blocks[worker.engine_id] = worker.num_blocks + + # Remote P with TP=8 also has 1 head/rank -> identical + # block_len despite tp_ratio == 2. + meta = NixlAgentMetadata( + engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID, + agent_metadata=FakeNixlWrapper.AGENT_METADATA, + kv_caches_base_addr=[0], + device_id=0, + num_blocks=1, + block_lens=list(worker.block_len_per_layer), + kv_cache_layout="HND", + block_size=worker.block_size, + ssm_sizes=(0, 0), + attn_backend_name=worker.backend_name, + physical_blocks_per_logical_kv_block=1, + ) + + # Must validate cleanly (used to raise AssertionError). + worker.add_remote_agent(meta, remote_tp_size=8) + + @patch( + "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper", + FakeNixlWrapper, + ) + def test_handshake_rejects_wrong_block_len_without_gqa_replication( + self, default_vllm_config, dist_init + ): + """Ensure the head-ratio validation still rejects genuinely wrong + block_lens when GQA replication is NOT in effect (32 KV heads, + D_TP=4, P_TP=2: head_ratio=4, both sides have >1 head/rank). + """ + vllm_config = create_vllm_config() + + with patch( + "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size", # noqa: E501 + return_value=4, + ): + connector = NixlConnector( + vllm_config, + KVConnectorRole.WORKER, + make_kv_cache_config(block_size=16), + ) + connector.connector_worker = FakeNixlConnectorWorker( + vllm_config, connector.engine_id, hand_shake_latency=0 + ) + worker = connector.connector_worker + + worker.transfer_topo.total_num_kv_heads = 32 + worker.transfer_topo.local_physical_heads = 8 # 32 // 4 + worker.kv_cache_layout = "HND" + + slot_size = 4096 + worker.slot_size_per_layer = [slot_size] + worker.block_len_per_layer = [slot_size * worker.block_size] + worker.num_blocks = 1 + worker.dst_num_blocks[worker.engine_id] = worker.num_blocks + + # Remote P_TP=2 has 16 heads/rank -> head_ratio = 16/8 = 2. + # Correct remote block_len = local * 2. Send local * 1 + # (wrong) to verify rejection. + bad_meta = NixlAgentMetadata( + engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID, + agent_metadata=FakeNixlWrapper.AGENT_METADATA, + kv_caches_base_addr=[0], + device_id=0, + num_blocks=1, + block_lens=list(worker.block_len_per_layer), + kv_cache_layout="HND", + block_size=worker.block_size, + ssm_sizes=(0, 0), + attn_backend_name=worker.backend_name, + physical_blocks_per_logical_kv_block=1, + ) + + with pytest.raises(AssertionError): + worker.add_remote_agent(bad_meta, remote_tp_size=2) + # NOTE: resource cleanup in mp backend is a bit finicky, so the order in which # we put here is important. First run ray, it will clean up the resources, then diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py index 74980dabf71..5804732f80f 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py @@ -1599,7 +1599,9 @@ class NixlBaseConnectorWorker: # Per-region block_len validation enforcing the P/D invariant. # REPLICATE regions (MLA, or a whole-model MLA / replicated-KV transfer) # only allow the number of blocks to differ; SPLIT regions scale with - # tp_ratio. Mamba uses the ssm_sizes counterpart, so skip block_len here. + # the per-rank KV head ratio rather than the raw tp_ratio, because GQA + # replication caps per-rank heads at 1 when tp > total_kv_heads + # (issue #45330). Mamba uses the ssm_sizes counterpart, so skip here. if not self._has_mamba: assert len(self.block_len_per_layer) == len(nixl_agent_meta.block_lens), ( "Number of KV layers must match between prefill and decode" @@ -1607,6 +1609,9 @@ class NixlBaseConnectorWorker: model_replicated = self.use_mla or self.transfer_topo.is_kv_replicated( remote_engine_id ) + total_kv_heads = self.transfer_topo.total_num_kv_heads + local_heads = self.transfer_topo.local_physical_heads + remote_heads = max(1, total_kv_heads // remote_tp_size) for i, local_len in enumerate(self.block_len_per_layer): replicated = model_replicated or self._is_region_replicated(i) remote_len = nixl_agent_meta.block_lens[i] @@ -1617,9 +1622,13 @@ class NixlBaseConnectorWorker: f"remote={remote_len}, bsr={block_size_ratio})." ) elif tp_ratio > 0: - assert remote_len == (local_len * tp_ratio) // block_size_ratio, ( + assert ( + remote_len + == (local_len * remote_heads // local_heads) // block_size_ratio + ), ( f"SPLIT region {i}: remote P KV block_len {remote_len} " - f"must equal local {local_len} * tp_ratio {tp_ratio} " + f"must equal local {local_len} * remote_heads " + f"{remote_heads} // local_heads {local_heads} " f"// block_size_ratio {block_size_ratio}." ) else: @@ -1627,10 +1636,10 @@ class NixlBaseConnectorWorker: "Different local/remote block sizes are not supported " "when P TP > D TP." ) - assert remote_len == local_len // (-tp_ratio), ( - f"SPLIT region {i}: remote P KV block_len " - f"{remote_len} must equal local {local_len} " - f"// |tp_ratio| {-tp_ratio}." + assert remote_len == local_len * remote_heads // local_heads, ( + f"SPLIT region {i}: remote P KV block_len {remote_len} " + f"must equal local {local_len} * remote_heads " + f"{remote_heads} // local_heads {local_heads}." ) # TP workers that handhshake with same remote have same #blocks.