diff --git a/CMakeLists.txt b/CMakeLists.txt index cf215a6f249..3ff367fafa2 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -1143,7 +1143,8 @@ set(VLLM_MOE_EXT_SRC if(VLLM_GPU_LANG STREQUAL "CUDA") list(APPEND VLLM_MOE_EXT_SRC "csrc/libtorch_stable/moe/moe_wna16.cu" - "csrc/libtorch_stable/moe/grouped_topk_kernels.cu") + "csrc/libtorch_stable/moe/grouped_topk_kernels.cu" + "csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels_fast.cu") endif() if(VLLM_GPU_LANG STREQUAL "CUDA") diff --git a/csrc/libtorch_stable/moe/moe_ops.h b/csrc/libtorch_stable/moe/moe_ops.h index 43cbb7f86d3..47efc9e1218 100644 --- a/csrc/libtorch_stable/moe/moe_ops.h +++ b/csrc/libtorch_stable/moe/moe_ops.h @@ -74,6 +74,22 @@ void shuffle_rows(const torch::stable::Tensor& input_tensor, const torch::stable::Tensor& dst2src_map, torch::stable::Tensor& output_tensor); +#ifndef USE_ROCM +// TRT-LLM fused sqrt-softplus routing gate (top-k, renormalize). Ported from +// NVIDIA/TensorRT-LLM PR #15402 (customMoeRoutingKernels.cu). Signature mirrors +// topk_softplus_sqrt for drop-in use: gating_output/topk_weights are float32, +// topk_indices is int32, topk (=topk_weights.size(-1)) must be 6, n_experts in +// {256, 384}, always renormalizes. Hash mode is inferred from tid2eid. +void topk_softplus_sqrt_fast( + torch::stable::Tensor& topk_weights, torch::stable::Tensor& topk_indices, + torch::stable::Tensor& token_expert_indices, + torch::stable::Tensor& gating_output, bool renormalize, + double routed_scaling_factor, + const std::optional& correction_bias, + const std::optional& input_ids, + const std::optional& tid2eid); +#endif + #ifndef USE_ROCM // DeepSeek V3 optimized router GEMM kernel for SM90+ // Computes output = mat_a @ mat_b.T where: diff --git a/csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels_fast.cu b/csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels_fast.cu new file mode 100644 index 00000000000..2172e8d8663 --- /dev/null +++ b/csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels_fast.cu @@ -0,0 +1,263 @@ +/* + * Adapted from + * https://github.com/NVIDIA/TensorRT-LLM/pull/15402 + * cpp/tensorrt_llm/kernels/customMoeRoutingKernels.cu (gate_forward) + * Copyright (c) 2026, The vLLM team. + * SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & + * AFFILIATES. All rights reserved. SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#include +#include +#include +#include +#include + +#include "../../cuda_compat.h" +#include "libtorch_stable/torch_utils.h" + +// CUDA-only: depends on cub-based reduce_topk and a 32-lane warp layout. +#ifndef USE_ROCM + + #include + #include + + #include "moeTopKFuncs.cuh" // vllm::moe::reduce_topk::reduceTopK + +namespace cg = cooperative_groups; + +namespace vllm { +namespace moe { + +// CUDA kernel for gate forward. +// Input: pre-computed scores from linear(x, weight) done outside the kernel. +// Template parameters: +// nExperts: number of experts +// topK: number of top experts to select +// hash: true for hash mode, false for topk mode +// One warp per row (batch element). +template +__global__ void topk_softplus_sqrt_fast_kernel( + float const* __restrict__ scores_in, // [batch_size, nExperts] + float const* __restrict__ bias, // [nExperts] (only when hash=false) + int const* __restrict__ input_ids, // [batch_size] (only when hash=true) + int const* __restrict__ tid2eid, // [vocab_size, topK] (only hash=true) + float* __restrict__ out_weights, // [batch_size, topK] + int* __restrict__ out_indices, // [batch_size, topK] + int batch_size, float route_scale) { + // Compile-time constants + constexpr int kExpertsPerThread = nExperts / WARP_SIZE; + constexpr int kWarpsPerBlock = 4; // Adjust based on occupancy needs + + // Shared memory for original scores (one array per warp in the block) + __shared__ float smem_scores[kWarpsPerBlock][nExperts]; + + // One warp per batch element + int const global_warp_id = (blockIdx.x * blockDim.x + threadIdx.x) / WARP_SIZE; + int const local_warp_id = (threadIdx.x / WARP_SIZE) % kWarpsPerBlock; + int const lane_id = threadIdx.x % WARP_SIZE; + + if (global_warp_id >= batch_size) return; + + auto warp = cg::tiled_partition(cg::this_thread_block()); + + // Pointer to this warp's shared memory and input scores + float* my_smem = smem_scores[local_warp_id]; + float const* scores_row = scores_in + global_warp_id * nExperts; + + // Load scores, apply score function (softplus + sqrt), store to shared mem. +#pragma unroll + for (int e = 0; e < kExpertsPerThread; ++e) { + int expert_id = lane_id + e * WARP_SIZE; + float s = scores_row[expert_id]; + float sp = log1pf(expf(s)); + float score = sqrtf(sp); + my_smem[expert_id] = score; // Store original score to shared memory + } + __syncwarp(); // Ensure all scores are written before reading + + // Output: each of the first K lanes holds one value. + float my_topk_value = 0.0f; + int my_topk_index = 0; + + if constexpr (hash) { + // Hash mode: directly read from shared memory + int token_id = input_ids[global_warp_id]; + int const* expert_ids = tid2eid + token_id * topK; + + if (lane_id < topK) { + int expert_id = expert_ids[lane_id]; + my_topk_index = expert_id; + my_topk_value = my_smem[expert_id]; // Direct lookup from shared memory + } + } else { + // Topk mode: load from shared memory, add bias in registers for topk. + float scores[kExpertsPerThread]; + int indices[kExpertsPerThread]; + +#pragma unroll + for (int e = 0; e < kExpertsPerThread; ++e) { + int expert_id = lane_id + e * WARP_SIZE; + indices[e] = expert_id; + scores[e] = my_smem[expert_id] + bias[expert_id]; // biased for selection + } + + // Use reduceTopK to find the top-k experts (result broadcast to all lanes). + float topk_values[topK]; + int32_t topk_indices[topK]; + constexpr float minValue = -1e30f; + reduce_topk::reduceTopK( + warp, topk_values, topk_indices, scores, indices, minValue, topK); + + // Gather original weights (without bias) from shared memory. + if (lane_id < topK) { + int expert_id = topk_indices[lane_id]; + my_topk_index = expert_id; + my_topk_value = my_smem[expert_id]; // original score (no bias) + } + } + + // Reduce to get the sum (first K lanes have values, others have 0). + float weight_sum = cg::reduce(warp, my_topk_value, cg::plus{}); + + // Normalize weights and write output (first K lanes). + if (lane_id < topK) { + out_weights[global_warp_id * topK + lane_id] = + (my_topk_value / weight_sum) * route_scale; + out_indices[global_warp_id * topK + lane_id] = my_topk_index; + } +} + +// C++ launcher (output tensors passed as parameters). All tensors are float32. +template +void launch_topk_softplus_sqrt_fast_kernel(float* scores_in, float* bias, int* input_ids, + int* tid2eid, float* out_weights, + int* out_indices, int batch_size, + float route_scale, cudaStream_t stream) { + constexpr int warps_per_block = 4; + constexpr int threads_per_block = warps_per_block * WARP_SIZE; + int const blocks = (batch_size + warps_per_block - 1) / warps_per_block; + + topk_softplus_sqrt_fast_kernel + <<>>(scores_in, bias, input_ids, + tid2eid, out_weights, + out_indices, batch_size, + route_scale); +} + +// Dispatch over (n_experts, is_hash). topK and n_experts match the source: +// DeepSeek-V4 uses top-k 6, and n_experts is 256 or 384. +template +void topk_softplus_sqrt_fast_dispatch_experts(float* scores, float* bias, int* input_ids, + int* tid2eid, float* weights, int* indices, + int batch_size, int n_experts, + float route_scale, bool is_hash, + cudaStream_t stream) { + switch (n_experts) { + case 256: + if (is_hash) { + launch_topk_softplus_sqrt_fast_kernel<256, topK, true>( + scores, nullptr, input_ids, tid2eid, weights, indices, batch_size, + route_scale, stream); + } else { + launch_topk_softplus_sqrt_fast_kernel<256, topK, false>( + scores, bias, nullptr, nullptr, weights, indices, batch_size, + route_scale, stream); + } + break; + case 384: + if (is_hash) { + launch_topk_softplus_sqrt_fast_kernel<384, topK, true>( + scores, nullptr, input_ids, tid2eid, weights, indices, batch_size, + route_scale, stream); + } else { + launch_topk_softplus_sqrt_fast_kernel<384, topK, false>( + scores, bias, nullptr, nullptr, weights, indices, batch_size, + route_scale, stream); + } + break; + default: + STD_TORCH_CHECK(false, "topk_softplus_sqrt_fast only supports n_experts 256 or 384, " + "got ", + n_experts); + } +} + +} // namespace moe +} // namespace vllm + +// Stable-ABI entry. Signature mirrors `topk_softplus_sqrt` so it is a drop-in +// replacement. gating_output/topk_weights are float32, topk_indices is int32. +// topk is taken from topk_weights.size(-1) (must be 6), hash mode is inferred +// from tid2eid, and route_scale = routed_scaling_factor. token_expert_indices +// is accepted for ABI parity but unused (this kernel does not permute). +// Selection: sqrt(softplus(score)) (+ bias for topk selection), top-k, then +// renormalize by the sum of the selected (unbiased) weights and scale. +void topk_softplus_sqrt_fast( + torch::stable::Tensor& topk_weights, // [num_tokens, topk] fp32 + torch::stable::Tensor& topk_indices, // [num_tokens, topk] int32 + torch::stable::Tensor& token_expert_indices, // unused (ABI parity) + torch::stable::Tensor& gating_output, // [num_tokens, n_experts] fp32 + bool renormalize, double routed_scaling_factor, + const std::optional& correction_bias, + const std::optional& input_ids, + const std::optional& tid2eid) { + const int n_experts = gating_output.size(-1); + const int batch_size = gating_output.numel() / n_experts; + const int topk = topk_weights.size(-1); + const bool is_hash = tid2eid.has_value(); + + STD_TORCH_CHECK( + renormalize, + "topk_softplus_sqrt_fast always renormalizes; renormalize must be true"); + STD_TORCH_CHECK(topk == 6, "topk_softplus_sqrt_fast only supports topk 6, got ", + topk); + STD_TORCH_CHECK( + gating_output.scalar_type() == torch::headeronly::ScalarType::Float, + "gating_output must be float32"); + STD_TORCH_CHECK( + topk_weights.scalar_type() == torch::headeronly::ScalarType::Float, + "topk_weights must be float32"); + STD_TORCH_CHECK( + topk_indices.scalar_type() == torch::headeronly::ScalarType::Int, + "topk_indices must be int32"); + + float* bias_ptr = nullptr; + int* input_ids_ptr = nullptr; + int* tid2eid_ptr = nullptr; + if (is_hash) { + STD_TORCH_CHECK(input_ids.has_value(), "hash mode requires input_ids"); + input_ids_ptr = input_ids.value().mutable_data_ptr(); + tid2eid_ptr = tid2eid.value().mutable_data_ptr(); + } else { + STD_TORCH_CHECK(correction_bias.has_value(), "topk mode requires bias"); + bias_ptr = correction_bias.value().mutable_data_ptr(); + } + + const torch::stable::accelerator::DeviceGuard guard( + gating_output.get_device_index()); + const cudaStream_t stream = + get_current_cuda_stream(gating_output.get_device_index()); + + float* scores = gating_output.mutable_data_ptr(); + float* weights = topk_weights.mutable_data_ptr(); + int* indices = topk_indices.mutable_data_ptr(); + + vllm::moe::topk_softplus_sqrt_fast_dispatch_experts<6>( + scores, bias_ptr, input_ids_ptr, tid2eid_ptr, weights, indices, + batch_size, n_experts, static_cast(routed_scaling_factor), is_hash, + stream); +} + +#endif // USE_ROCM diff --git a/csrc/libtorch_stable/moe/torch_bindings.cpp b/csrc/libtorch_stable/moe/torch_bindings.cpp index bfcb0074e5b..0f49c5d2c88 100644 --- a/csrc/libtorch_stable/moe/torch_bindings.cpp +++ b/csrc/libtorch_stable/moe/torch_bindings.cpp @@ -122,6 +122,14 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_moe_C, m) { "routed_scaling_factor, Tensor bias, int scoring_func) -> (Tensor, " "Tensor)"); + // TRT-LLM fused sqrt-softplus routing gate. Signature mirrors + // topk_softplus_sqrt for drop-in use (topk fixed at 6, always renormalizes). + m.def( + "topk_softplus_sqrt_fast(Tensor! topk_weights, Tensor! topk_indices, " + "Tensor! token_expert_indices, Tensor gating_output, bool renormalize, " + "float routed_scaling_factor, Tensor? bias, Tensor? input_ids, Tensor? " + "tid2eid) -> ()"); + // DeepSeek V3 optimized router GEMM for SM90+ m.def("dsv3_router_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()"); // conditionally compiled so impl registration is in source file @@ -141,6 +149,7 @@ STABLE_TORCH_LIBRARY_IMPL(_moe_C, CUDA, m) { m.impl("moe_wna16_gemm", TORCH_BOX(&moe_wna16_gemm)); m.impl("shuffle_rows", TORCH_BOX(&shuffle_rows)); m.impl("grouped_topk", TORCH_BOX(&grouped_topk)); + m.impl("topk_softplus_sqrt_fast", TORCH_BOX(&topk_softplus_sqrt_fast)); #endif } diff --git a/vllm/_custom_ops.py b/vllm/_custom_ops.py index a609035560e..be0526dc542 100644 --- a/vllm/_custom_ops.py +++ b/vllm/_custom_ops.py @@ -2373,6 +2373,38 @@ def topk_hash_softplus_sqrt( ) +def topk_softplus_sqrt_fast( + topk_weights: torch.Tensor, + topk_indices: torch.Tensor, + token_expert_indices: torch.Tensor, + gating_output: torch.Tensor, + renormalize: bool = False, + routed_scaling_factor: float = 1.0, + e_score_correction_bias: torch.Tensor | None = None, + input_tokens: torch.Tensor | None = None, + hash_indices_table: torch.Tensor | None = None, +) -> None: + """TRT-LLM fused sqrt-softplus routing gate (ported from TRT-LLM #15402). + + Drop-in replacement for :func:`topk_hash_softplus_sqrt`: same signature. + Constraints: gating_output/topk_weights float32, topk_indices int32, + topk (=topk_weights.shape[-1]) must be 6, n_experts in {256, 384}, and it + always renormalizes (``renormalize`` must be True). Hash mode is inferred + from ``hash_indices_table``; ``token_expert_indices`` is unused. + """ + torch.ops._moe_C.topk_softplus_sqrt_fast( + topk_weights, + topk_indices, + token_expert_indices, + gating_output, + renormalize, + routed_scaling_factor, + e_score_correction_bias, + input_tokens, + hash_indices_table, + ) + + def grouped_topk( scores: torch.Tensor, num_expert_group: int, diff --git a/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py b/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py index d505c5ce4b7..14de82442d1 100644 --- a/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py +++ b/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py @@ -102,6 +102,26 @@ def _topk_softplus_sqrt_torch( return topk_weights, topk_indices +def _use_topk_softplus_sqrt_fast( + topk_weights: torch.Tensor, + topk_indices: torch.Tensor, + gating_output: torch.Tensor, + renormalize: bool, +) -> bool: + """Whether the TRT-LLM fused gate (topk_softplus_sqrt_fast) is applicable. + + It requires CUDA, int32 indices, topk==6, n_experts in {256, 384}, and + always renormalizes. Anything else falls back to topk_hash_softplus_sqrt. + """ + return ( + renormalize + and gating_output.is_cuda + and topk_indices.dtype == torch.int32 + and topk_weights.shape[-1] == 6 + and gating_output.shape[-1] in (256, 384) + ) + + def vllm_topk_softplus_sqrt( topk_weights: torch.Tensor, topk_indices: torch.Tensor, @@ -128,6 +148,24 @@ def vllm_topk_softplus_sqrt( routed_scaling_factor, ) + if _use_topk_softplus_sqrt_fast( + topk_weights, topk_indices, gating_output, renormalize + ): + # Fused TRT-LLM gate kernel; requires fp32 scores (no-op if already + # fp32). Same drop-in signature as topk_hash_softplus_sqrt. + ops.topk_softplus_sqrt_fast( + topk_weights, + topk_indices, + token_expert_indices, + gating_output, + renormalize, + routed_scaling_factor, + e_score_correction_bias, + input_tokens, + hash_indices_table, + ) + return topk_weights, topk_indices + ops.topk_hash_softplus_sqrt( topk_weights, topk_indices,