From c3d05b5ec513dc371ba400c2717fcefd991ba2b3 Mon Sep 17 00:00:00 2001 From: Vihang Mehta Date: Sun, 12 Apr 2026 16:10:33 -0700 Subject: [PATCH] Router updates Signed-off-by: Vihang Mehta --- .../blockScaleMoe/RoutingKernel.cuh | 5 +- .../blockScaleMoe/RoutingKernel.h | 9 +- .../blockScaleMoe/RoutingRenormalize.cu | 19 ++-- .../routingRenormalize/launchBlockKernel.cu | 86 +++++++++---------- .../trtllmGenKernels/blockScaleMoe/runner.cu | 3 +- 5 files changed, 62 insertions(+), 60 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh index 4bc7b56aa18b..ca9f24ebdcc7 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh @@ -582,11 +582,10 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelPa } else { - // If params.mPtrTopKIds != nullptr, we don't need to store the weights + scoreIdx = params.mPtrTopKPacked[expandedIdx]; + idx = scoreIdx.idx; if (params.mPtrTopKWeights != nullptr) { - scoreIdx = params.mPtrTopKPacked[expandedIdx]; - idx = scoreIdx.idx; params.mPtrTopKWeights[expandedIdx] = static_cast(scoreIdx.score); } } diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h index 3daa1848e5d3..f720450a63df 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h @@ -64,13 +64,14 @@ struct DataBase // optional: if `nullptr`, it is not filled // dim: [mNumTokens, mTopK] - // When mPtrTopKIds is provided, mPtrTopKWeights must be also provided as inputs. - // Otherwise, mPtrTopKWeights is the output scores of the topK experts. + // `routingRenormalize`: when `mPtrScores` is set (logits path), this holds per-slot masses (output). + // When `mPtrScores` is null and ids are precomputed (e.g. MoE thop), this may alias the precomputed + // top-k weight input together with `mPtrTopKIds`. void* mPtrTopKWeights{nullptr}; // optional: if `nullptr`, it is not filled // dim: [mNumTokens, mTopK] - // mPtrTopKIds[i] is the index of the expert for the i-th token in the top-k experts - // Together with mPtrTopKWeights, they form the top-k experts for each token + // `routingRenormalize`: output expert indices when computing from `mPtrScores`; input precomputed ids + // when `mPtrScores` is null (MoE thop supplies top-k ids/weights without logits). int32_t* mPtrTopKIds{nullptr}; // optional: if `nullptr`, scores are used directly as input. diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu index ff4bb808d92e..463614f0f028 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu @@ -33,11 +33,12 @@ void launchOffsetsKernel(Data const& data, int numBlocksOffsets, uint32_t numThr void run(Data const& data, void* stream) { TLLM_CHECK_WITH_INFO(data.mPtrTopKPacked != nullptr || data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr, - "Routing kernel requires at least one input parameter"); - if (data.mPtrTopKIds != nullptr) + "Routing kernel requires mPtrScores (logits), mPtrTopKPacked staging, and/or mPtrTopKIds"); + if (data.mPtrScores != nullptr && data.mPtrTopKPacked == nullptr) { - TLLM_CHECK_WITH_INFO(data.mPtrTopKWeights != nullptr, - "When mPtrTopKIds is provided, mPtrTopKWeights must also be provided for Renormalize routing."); + TLLM_CHECK_WITH_INFO(data.mPtrTopKIds != nullptr && data.mPtrTopKWeights != nullptr, + "Renormalize from logits (mPtrScores without mPtrTopKPacked) requires mPtrTopKIds and mPtrTopKWeights " + "outputs."); } TLLM_CHECK_WITH_INFO(data.mPtrPermutedIdxSize != nullptr && data.mPtrCtaIdxXyToBatchIdx != nullptr && data.mPtrCtaIdxXyToMnLimit != nullptr && data.mPtrNumNonExitingCtas != nullptr, @@ -53,14 +54,14 @@ void run(Data const& data, void* stream) bool const useSingleBlock = data.mNumTokens <= BlockKernelMaxNumTokens || (data.mNumTokens <= DynBlockKernelMaxNumTokens && data.mNumExperts <= DynBlockKernelMaxNumExperts); - bool const useSingleCluster = data.mNumTokens <= ((data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr) - ? MaxNumTokensSingleClusterScores - : MaxNumTokensSingleCluster); + bool const useSingleCluster = data.mNumTokens + <= ((data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr) ? MaxNumTokensSingleClusterScores + : MaxNumTokensSingleCluster); if (!useSingleCluster && !useSingleBlock) { TLLM_CHECK_WITH_INFO((data.mPtrTopKPacked != nullptr || data.mPtrTopKIds != nullptr), - "When #tokens is large, `mPtrTopKPacked` or `mPtrTopKIds` is a required input."); + "When #tokens is large, `mPtrTopKPacked` staging or `mPtrTopKIds` input is required."); TLLM_CHECK_WITH_INFO( data.mPtrExpertCounts != nullptr, "When #tokens is large, `mPtrExpertCounts` is a required input."); } @@ -89,7 +90,7 @@ void run(Data const& data, void* stream) int const numBlocksOffsets = std::min((expandedIdxSize + offsetEltsPerBlock - 1) / offsetEltsPerBlock, maxNumBlocks); - if (data.mPtrScores != nullptr && data.mPtrTopKIds == nullptr) + if (data.mPtrScores != nullptr) { launchHistogramScoresKernel(data, maxNumBlocks, numThreadsHist, stream); } diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu index 2a4f9257aa9f..d494e7f36aa3 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu @@ -69,35 +69,15 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelPa } #endif - if (params.mPtrTopKIds != nullptr) + if (params.mPtrScores != nullptr) { - if (validToken) - { - if (laneIdx < params.mTopK) - { - auto expertIdx = params.mPtrTopKIds[warpIdx * params.mTopK + laneIdx]; - if (expertIdx != -1) - { - int offset = warpIdx * MaxNumExperts + expertIdx; - smemKIdx[offset] = static_cast(laneIdx); - } - else - { - params.mPtrExpandedIdxToPermutedIdx[warpIdx * params.mTopK + laneIdx] = int32_t{-1}; - } - } - } - } - else if (params.mPtrScores != nullptr) - { - // in this case, each warp represents a token + // Each warp represents a token: compute top-k from logits and write expert ids + weights. BaseType score[VecSize]; int32_t idx[VecSize]; BaseType warpTopKScore[KernelParams::MaxNumTopExperts]; int32_t warpTopKExpertIdx[KernelParams::MaxNumTopExperts]; - BaseType minScore = BaseType{-INFINITY}; if (validToken) { routingTopKExperts(laneIdx); - if (params.mPtrTopKWeights != nullptr) + params.mPtrTopKIds[warpIdx * params.mTopK + laneIdx] = warpTopKExpertIdx[laneIdx]; + params.mPtrTopKWeights[warpIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]}; + } + } // end if (validToken) + } + else if (params.mPtrTopKIds != nullptr) + { + // Precomputed top-k expert ids + weights (e.g. MoE thop); logits path is disabled. + if (validToken) + { + if (laneIdx < params.mTopK) + { + auto expertIdx = params.mPtrTopKIds[warpIdx * params.mTopK + laneIdx]; + if (expertIdx != -1) + { + int offset = warpIdx * MaxNumExperts + expertIdx; + smemKIdx[offset] = static_cast(laneIdx); + } + else { - params.mPtrTopKWeights[warpIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]}; + params.mPtrExpandedIdxToPermutedIdx[warpIdx * params.mTopK + laneIdx] = int32_t{-1}; } } - } // end if (validToken) + } } __syncthreads(); @@ -405,22 +403,7 @@ __global__ void routingIndicesDynBlockKernel(KernelParams params) // ── Phase 1: TopK — one warp per token (loop only when numTokens > numWarps) ── for (int tokenIdx = warpIdx; tokenIdx < params.mNumTokens; tokenIdx += numWarps) { - if (params.mPtrTopKIds != nullptr) - { - if (laneIdx < params.mTopK) - { - auto expertIdx = params.mPtrTopKIds[tokenIdx * params.mTopK + laneIdx]; - if (expertIdx > -1 && expertIdx < params.mNumExperts) - { - smemKIdx[tokenIdx * MaxNumExperts + expertIdx] = static_cast(laneIdx); - } - else - { - params.mPtrExpandedIdxToPermutedIdx[tokenIdx * params.mTopK + laneIdx] = int32_t{-1}; - } - } - } - else if (params.mPtrScores != nullptr) + if (params.mPtrScores != nullptr) { BaseType score[VecSize]; int32_t idx[VecSize]; @@ -435,9 +418,22 @@ __global__ void routingIndicesDynBlockKernel(KernelParams params) if (laneIdx < params.mTopK) { smemKIdx[tokenIdx * MaxNumExperts + warpTopKExpertIdx[laneIdx]] = static_cast(laneIdx); - if (params.mPtrTopKWeights != nullptr) + params.mPtrTopKIds[tokenIdx * params.mTopK + laneIdx] = warpTopKExpertIdx[laneIdx]; + params.mPtrTopKWeights[tokenIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]}; + } + } + else if (params.mPtrTopKIds != nullptr) + { + if (laneIdx < params.mTopK) + { + auto expertIdx = params.mPtrTopKIds[tokenIdx * params.mTopK + laneIdx]; + if (expertIdx > -1 && expertIdx < params.mNumExperts) + { + smemKIdx[tokenIdx * MaxNumExperts + expertIdx] = static_cast(laneIdx); + } + else { - params.mPtrTopKWeights[tokenIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]}; + params.mPtrExpandedIdxToPermutedIdx[tokenIdx * params.mTopK + laneIdx] = int32_t{-1}; } } } @@ -448,6 +444,10 @@ __global__ void routingIndicesDynBlockKernel(KernelParams params) auto expandedIdx = tokenIdx * params.mTopK + laneIdx; auto scoreIdx = params.mPtrTopKPacked[expandedIdx]; smemKIdx[tokenIdx * MaxNumExperts + static_cast(scoreIdx.idx)] = static_cast(laneIdx); + if (params.mPtrTopKIds != nullptr) + { + params.mPtrTopKIds[expandedIdx] = static_cast(scoreIdx.idx); + } if (params.mPtrTopKWeights != nullptr) { params.mPtrTopKWeights[expandedIdx] = static_cast(scoreIdx.score); diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu index 467bca9318ac..a0dcc3d5b2fb 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu @@ -172,7 +172,8 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3 routingData.mDoSoftmaxBeforeTopK = routingMethodType == RoutingMethodType::RenormalizeNaive; routingData.mNormTopkProb = routingMethodType == RoutingMethodType::RenormalizeNaive; - // Pass-through raw pointer; kernels will cast to the proper InputT based on routing method + // When precomputed top-k ids/weights are supplied (thop), expertIds is non-null and logits are ignored. + // When computing from logits (e.g. TensorRT plugin), expertIds is null and routingLogits drives TopK. routingData.mPtrScores = expertIds == nullptr ? routingLogits : nullptr; // // Outputs