From 855df49086a5a5cfe5b6429c56faff54e4a5c927 Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Fri, 13 Feb 2026 19:05:06 -0800 Subject: [PATCH 01/11] minimax --- .../blockScaleMoe/RoutingKernel.h | 76 ++- .../blockScaleMoe/RoutingMiniMax.cu | 633 ++++++++++++++++++ .../trtllmGenKernels/blockScaleMoe/runner.cu | 33 + .../_torch/modules/moe/moe_test_utils.py | 4 +- 4 files changed, 741 insertions(+), 5 deletions(-) create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h index 888e04f2541a..97195a3daa3d 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h @@ -195,7 +195,11 @@ struct KernelParams : public KernelParamsBase* mPtrTopKPacked = nullptr; @@ -302,7 +311,11 @@ struct KernelParams : public KernelParamsBase* mPtrTopKPacked = nullptr; @@ -329,6 +342,63 @@ void run(Data const& data, void* stream); } // namespace routingRenormalize -//////////////////////////////////////////////////////////////////////////////////////////////////// +////////////////////////////////////////////////////////////////////////////////////////////////// + +namespace routingMiniMax +{ + +////////////////////////////////////////////////////////////////////////////////////////////////// + +struct Data : public DataBase +{ + tg::Dtype mDtypeExpW{tg::Dtype::Fp32}; + + // Used by KernelParams::setKernelParams + void const* mPtrRoutingBias{nullptr}; + + // Used by KernelParams::setKernelParams + bool mNormTopkProb{true}; +}; + +////////////////////////////////////////////////////////////////////////////////////////////////// + +template +struct KernelParams : public KernelParamsBase +{ + using InputT = InputT_; + using OutputT = OutputT_; + + // Add missing static constexpr members for shared utilities + static constexpr int MaxNumExperts = MaxNumExperts_; + static constexpr bool isPow2 = isPow2_; + static constexpr bool UsePdl = UsePdl_; + static constexpr bool DoSoftmaxBeforeTopK = DoSoftmaxBeforeTopK_; + + PackedScoreIdx* mPtrTopKPacked = nullptr; + + OutputT const* mPtrRoutingBias = nullptr; + + int32_t mTopK = 0; + + bool mNormTopkProb = true; + + static KernelParams setKernelParams(Data const& data) + { + KernelParams params; + params.setBaseParams(data); + + params.mPtrTopKPacked = (PackedScoreIdx*) data.mPtrTopKPacked; + params.mPtrRoutingBias = static_cast(data.mPtrRoutingBias); + params.mNormTopkProb = data.mNormTopkProb; + params.mTopK = data.mTopK; + return params; + } +}; + +void run(Data const& data, void* stream); + +} // namespace routingMiniMax + +////////////////////////////////////////////////////////////////////////////////////////////////// } // namespace routing } // namespace moe::dev diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu new file mode 100644 index 000000000000..a4c9b8ec45de --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -0,0 +1,633 @@ +/* + * Copyright (c) 2022-2025, NVIDIA CORPORATION. All rights reserved. + * Contibuted by Baseten.co + * + * 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 "RoutingKernel.cuh" + +namespace moe::dev::routing +{ +namespace routingMiniMax +{ + +////////////////////////////////////////////////////////////////////////////////////////////////// + +static constexpr int NumExpertsLimit = 512; + +static constexpr int NumThreads = 1024; +static constexpr int NumWarps = NumThreads / WarpSize; +static constexpr int MaxSupportedTopExperts = 10; + +static constexpr int MaxNumTokensSingleCluster = NumBlocksPerCluster * NumThreads; +static constexpr int MaxNumTokensSingleClusterScores = NumBlocksPerCluster * NumWarps; + +static constexpr int BlockKernelMaxNumTokens = 4; + +template +__forceinline__ __device__ void routingTopKExperts(cg::thread_block_tile const& warp, + DataType (&score)[VecSize], DataType (&scoreWithBias)[VecSize], int32_t (&idx)[VecSize], + DataType (&warpTopKScore)[MaxSupportedTopExperts], int32_t (&warpTopKExpertIdx)[MaxSupportedTopExperts], + int32_t const laneIdx, int32_t const numExperts, int32_t topK, InputType const* ptrScores, + BiasType const* ptrRoutingBias, bool const normTopkProb) +{ + DataType minScore = DataType{-INFINITY}; + + // Store float scores for accurate shuffle and selection + float score_f[VecSize]; + float scoreWithBias_f[VecSize]; + + // Step 1: Apply sigmoid to logits (NOT softmax - MiniMax2 uses sigmoid) + for (int i = 0; i < VecSize; i++) + { + auto expertIdx = i * WarpSize + laneIdx; + if (expertIdx < numExperts) + { + // Apply sigmoid to get probability (store in float for accuracy) + float logit = static_cast(ptrScores[expertIdx]); + score_f[i] = sigmoid_accurate(logit); + score[i] = static_cast(score_f[i]); + scoreWithBias_f[i] = score_f[i]; // Will add bias below + scoreWithBias[i] = score[i]; // Keep DataType version for compatibility + } + else + { + score_f[i] = 0.0f; + score[i] = DataType{0}; // Invalid experts have 0 probability + scoreWithBias_f[i] = -INFINITY; + scoreWithBias[i] = minScore; // Invalid experts have -inf for selection + } + idx[i] = expertIdx; + } + + // Step 2: Add routing bias for selection (only for valid experts) + for (int i = 0; i < VecSize; i++) + { + auto expertIdx = i * WarpSize + laneIdx; + if (expertIdx < numExperts) + { + float bias = static_cast(ptrRoutingBias[expertIdx]); + scoreWithBias_f[i] = score_f[i] + bias; + scoreWithBias[i] = static_cast(scoreWithBias_f[i]); + } + } + + // Step 3: Top-K selection using bias-added scores (use float for stability) + topk::reduceTopK(warp, warpTopKScore, warpTopKExpertIdx, scoreWithBias, idx, minScore, topK); + + // Step 4: Gather original sigmoid scores (not bias-added) for selected experts + // Use warp shuffle to get the unbiased score from the lane that has it + if (laneIdx < topK) + { + int selectedExpertIdx = warpTopKExpertIdx[laneIdx]; + int targetLane = selectedExpertIdx % WarpSize; + int vecIdx = selectedExpertIdx / WarpSize; + + // Select the per-lane value via unrolled loop, then shuffle + float v = 0.0f; + #pragma unroll + for (int i = 0; i < VecSize; ++i) + { + if (vecIdx == i) v = score_f[i]; + } + float originalScore_f = __shfl_sync(0xffffffff, v, targetLane); + warpTopKScore[laneIdx] = static_cast(originalScore_f); + } + + // Step 5: Renormalize using original sigmoid scores (with epsilon to prevent divide-by-zero) + float sum = 1.0f; + if (normTopkProb) + { + sum = (laneIdx < topK) ? static_cast(warpTopKScore[laneIdx]) : 0.0f; + sum = cg::reduce(warp, sum, cg::plus()) + 1e-20f; // Add epsilon + } + if (laneIdx < topK) + { + float normalized = static_cast(warpTopKScore[laneIdx]) / sum; + warpTopKScore[laneIdx] = static_cast(normalized); + } +} + +template +__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesBlockKernel(KernelParams params) +{ + // Enforce kernel constraints + if (params.mNumTokens > BlockKernelMaxNumTokens) return; + if (params.mTopK > MaxSupportedTopExperts) return; + + // Guard against thread count mismatch + if (blockDim.x != KernelParams::MaxNumExperts) return; + + using OutputT = typename KernelParams::OutputT; + using InputT = typename KernelParams::InputT; + using BaseType = InputT; + using TypePacked = PackedScoreIdx; + int constexpr MaxNumExperts = KernelParams::MaxNumExperts; + + int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); + int32_t const laneIdx = cutlass::arch::LaneId(); + int32_t const expert = threadIdx.x; + auto scoreOffset = warpIdx * params.mNumExperts; + bool validToken = warpIdx < params.mNumTokens; + + static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; + static constexpr int totalExpertCounts = BlockKernelMaxNumTokens * MaxNumExperts; + __shared__ int8_t __attribute((aligned(128))) smemOffset[totalExpertCounts]; + __shared__ int8_t __attribute((aligned(128))) smemKIdx[totalExpertCounts]; + + using Scan = cub::BlockScan; + __shared__ typename Scan::TempStorage tempStorage; + + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + + for (int i = threadIdx.x; i < totalExpertCounts; i += blockDim.x) + { + smemOffset[i] = int8_t{-1}; + smemKIdx[i] = int8_t{-1}; + } + __syncthreads(); + +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + if constexpr (KernelParams::UsePdl) + { + cudaGridDependencySynchronize(); + } +#endif + + if (params.mPtrTopKIds != 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) + { + BaseType score[VecSize]; + BaseType scoreWithBias[VecSize]; + int32_t idx[VecSize]; + + BaseType warpTopKScore[MaxSupportedTopExperts]; + int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; + + BaseType minScore = BaseType{-INFINITY}; + if (validToken) + { + routingTopKExperts(warp, score, scoreWithBias, idx, warpTopKScore, + warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, params.mPtrScores + scoreOffset, + params.mPtrRoutingBias, params.mNormTopkProb); + + if (laneIdx < params.mTopK) + { + int offset = warpIdx * MaxNumExperts + warpTopKExpertIdx[laneIdx]; + smemKIdx[offset] = static_cast(laneIdx); + if (params.mPtrTopKWeights != nullptr) + { + params.mPtrTopKWeights[warpIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]}; + } + } + } + } + __syncthreads(); + + auto localExpertIdx = expert - params.mLocalExpertsStartIdx; + int stride = 1 << params.mLocalExpertsStrideLog2; + auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < (params.mNumLocalExperts * stride) + && (localExpertIdx % stride) == 0; + int accExpertCount = 0; + + if (isLocalExpert) + { + int offset = expert; + for (int j = 0; j < BlockKernelMaxNumTokens; j++) + { + if (smemKIdx[offset] >= 0) + { + smemOffset[offset] = static_cast(accExpertCount); + accExpertCount++; + } + offset += MaxNumExperts; + } + } + __syncthreads(); + + int32_t numCta; + if constexpr (KernelParams::isPow2) + { + numCta = divUpLog2(accExpertCount, params.mPaddingLog2); + } + else + { + numCta = divUpTileN(accExpertCount, params.mTileTokensDim); + } + int32_t ctaOffset = 0; + int32_t numNonExitingCtas; + Scan(tempStorage).ExclusiveSum(numCta, ctaOffset, numNonExitingCtas); + __syncthreads(); + + int32_t expertScanCounts = 0; + int32_t tmpCount; + if constexpr (KernelParams::isPow2) + { + tmpCount = divUpMulLog2(accExpertCount, params.mPaddingLog2); + } + else + { + tmpCount = divUpMulTileN(accExpertCount, params.mTileTokensDim); + } + Scan(tempStorage).ExclusiveSum(tmpCount, expertScanCounts); + __syncthreads(); + + if (isLocalExpert) + { + for (int cta = 0; cta < numCta; ++cta) + { + const int32_t localExpertIdx = (expert - params.mLocalExpertsStartIdx) >> params.mLocalExpertsStrideLog2; + params.mPtrCtaIdxXyToBatchIdx[ctaOffset + cta] = localExpertIdx; + int32_t mnLimit1; + int32_t mnLimit2; + if constexpr (KernelParams::isPow2) + { + mnLimit1 = mulLog2(ctaOffset + cta + 1, params.mPaddingLog2); + mnLimit2 = mulLog2(ctaOffset, params.mPaddingLog2) + accExpertCount; + } + else + { + mnLimit1 = mulTileN(ctaOffset + cta + 1, params.mTileTokensDim); + mnLimit2 = mulTileN(ctaOffset, params.mTileTokensDim) + accExpertCount; + } + params.mPtrCtaIdxXyToMnLimit[ctaOffset + cta] = min(mnLimit1, mnLimit2); + } + } + + if (threadIdx.x == 0) + { + int32_t permutedIdxSize; + if constexpr (KernelParams::isPow2) + { + permutedIdxSize = mulLog2(numNonExitingCtas, params.mPaddingLog2); + } + else + { + permutedIdxSize = mulTileN(numNonExitingCtas, params.mTileTokensDim); + } + params.mPtrPermutedIdxSize[0] = permutedIdxSize; + params.mPtrNumNonExitingCtas[0] = numNonExitingCtas; + } + +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + if constexpr (KernelParams::UsePdl) + { + cudaTriggerProgrammaticLaunchCompletion(); + } +#endif + + for (int tokenIdx = 0; tokenIdx < params.mNumTokens; tokenIdx++) + { + int offset = tokenIdx * MaxNumExperts + threadIdx.x; + if (smemKIdx[offset] >= 0) + { + int const expandedIdx = tokenIdx * params.mTopK + smemKIdx[offset]; + int const offsetWithinExpert = static_cast(smemOffset[offset]); + int const offsetForExpert = expertScanCounts; + int const permutedIdx = isLocalExpert ? offsetForExpert + offsetWithinExpert : int32_t{-1}; + + if (params.mPtrExpandedIdxToPermutedIdx != nullptr) + { + params.mPtrExpandedIdxToPermutedIdx[expandedIdx] = permutedIdx; + } + if (params.mPtrPermutedIdxToExpandedIdx != nullptr && isLocalExpert) + { + params.mPtrPermutedIdxToExpandedIdx[permutedIdx] = expandedIdx; + } + if (params.mPtrPermutedIdxToTokenIdx != nullptr && isLocalExpert) + { + params.mPtrPermutedIdxToTokenIdx[permutedIdx] = tokenIdx; + } + } + } +} + +template +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) +__global__ void __cluster_dims__(NumBlocksPerCluster, 1, 1) + __launch_bounds__(KernelParams::MaxNumExperts) + routingIndicesClusterKernel(KernelParams params) +#else +__global__ void __launch_bounds__(KernelParams::MaxNumExperts) + routingIndicesClusterKernel(KernelParams params) +#endif +{ + // Enforce kernel constraints + if (params.mTopK > MaxSupportedTopExperts) return; + + // Guard against thread count mismatch + if (blockDim.x != KernelParams::MaxNumExperts) return; + + using OutputT = typename KernelParams::OutputT; + using InputT = typename KernelParams::InputT; + + using BaseType = InputT; + using TypePacked = PackedScoreIdx; + + static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; + // routingPermutation expects NumThreads == MaxNumExperts + static constexpr int PermThreads = KernelParams::MaxNumExperts; + static constexpr int PermWarps = PermThreads / WarpSize; + + // Shared memory for topK results - sized for MaxNumExperts threads + __shared__ TypePacked __attribute((aligned(128))) smemPackedScoreIdx[PermWarps * MaxSupportedTopExperts]; + + uint32_t const clusterBlockRank = blockIdx.x; + + int32_t const warpIdx = threadIdx.x / WarpSize; + int32_t const laneIdx = cutlass::arch::LaneId(); + + auto warpTokenIdx = clusterBlockRank * PermWarps + warpIdx; + auto scoreOffset = warpTokenIdx * params.mNumExperts; + bool validToken = warpTokenIdx < params.mNumTokens; + + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + if constexpr (KernelParams::UsePdl) + { + cudaGridDependencySynchronize(); + } +#endif + + if (params.mPtrScores != nullptr) + { + BaseType score[VecSize]; + BaseType scoreWithBias[VecSize]; + int32_t idx[VecSize]; + + BaseType warpTopKScore[MaxSupportedTopExperts]; + int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; + + BaseType minScore = BaseType{-INFINITY}; + if (validToken) + { + routingTopKExperts(warp, score, scoreWithBias, idx, warpTopKScore, + warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, params.mPtrScores + scoreOffset, + params.mPtrRoutingBias, params.mNormTopkProb); + + if (laneIdx < params.mTopK) + { + // Use fixed stride to avoid indexing mistakes + constexpr int STRIDE = MaxSupportedTopExperts; + smemPackedScoreIdx[warpIdx * STRIDE + laneIdx] + = TypePacked{warpTopKScore[laneIdx], static_cast(warpTopKExpertIdx[laneIdx])}; + } + } + } + + // Pack topK results into shared memory for all input types + if (params.mPtrScores != nullptr) + { + // Already packed by routingTopKExperts above + } + else if (params.mPtrTopKIds != nullptr && params.mPtrTopKWeights != nullptr) + { + // Pack from pre-computed TopKIds and TopKWeights + if (validToken && laneIdx < params.mTopK) + { + int id = params.mPtrTopKIds[warpTokenIdx * params.mTopK + laneIdx]; + BaseType w = static_cast(params.mPtrTopKWeights[warpTokenIdx * params.mTopK + laneIdx]); + // Handle invalid expert IDs explicitly + if (id == -1) { + smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = + TypePacked{BaseType{0}, int16_t{-1}}; + } else { + smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = + TypePacked{w, static_cast(id)}; + } + } + } + else if (params.mPtrTopKPacked != nullptr) + { + // Pack from pre-computed packed format with explicit type casting + if (validToken && laneIdx < params.mTopK) + { + auto p = params.mPtrTopKPacked[warpTokenIdx * params.mTopK + laneIdx]; + BaseType w = static_cast(p.score); + int16_t id = p.idx; + smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = TypePacked{w, id}; + } + } + + // Initialize unused lanes to prevent stale data + if (validToken && laneIdx >= params.mTopK && laneIdx < MaxSupportedTopExperts) + { + smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = + TypePacked{BaseType{0}, int16_t{-1}}; + } + + // Synchronize before using shared memory + __syncthreads(); + + // Use fixed stride for shared memory layout + // LoadExpertIdxFromGlobal=false: load from shared memory (smemPackedScoreIdx) + constexpr int STRIDE = MaxSupportedTopExperts; + routingPermutation(params, smemPackedScoreIdx, warpIdx, clusterBlockRank); +} + +template +__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHistogramScoresKernel(KernelParams params) +{ + using OutputT = typename KernelParams::OutputT; + using InputT = typename KernelParams::InputT; + using BaseType = InputT; + + static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; + + int32_t const laneIdx = cutlass::arch::LaneId(); + int32_t const warpIdx = threadIdx.x / WarpSize; + int32_t const globalWarpIdx = blockIdx.x * KernelParams::MaxNumExperts / WarpSize + warpIdx; + int32_t const globalWarpStride = gridDim.x * KernelParams::MaxNumExperts / WarpSize; + BaseType minScore = BaseType{-INFINITY}; + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + if constexpr (KernelParams::UsePdl) + { + cudaGridDependencySynchronize(); + } +#endif + + int32_t expertCountsNum = 2 * params.mNumExperts; + int32_t globalThreadIdx = blockIdx.x * KernelParams::MaxNumExperts + threadIdx.x; + int32_t globalThreadStride = gridDim.x * KernelParams::MaxNumExperts; + initArr(globalThreadIdx, expertCountsNum, globalThreadStride, params.mPtrExpertCounts, 0); + +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + if constexpr (KernelParams::UsePdl) + { + cudaTriggerProgrammaticLaunchCompletion(); + } +#endif + + BaseType allScores[VecSize]; + BaseType allScoresWithBias[VecSize]; + int32_t allExpertIdx[VecSize]; + BaseType warpTopKScore[MaxSupportedTopExperts]; + int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; + for (int tokenIdx = globalWarpIdx; tokenIdx < params.mNumTokens; tokenIdx += globalWarpStride) + { + auto scoreOffset = tokenIdx * params.mNumExperts; + + routingTopKExperts(warp, allScores, allScoresWithBias, allExpertIdx, + warpTopKScore, warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, + params.mPtrScores + scoreOffset, params.mPtrRoutingBias, params.mNormTopkProb); + + if (laneIdx < params.mTopK) + { + PackedScoreIdx packedScore{ + static_cast(warpTopKScore[laneIdx]), static_cast(warpTopKExpertIdx[laneIdx])}; + params.mPtrTopKPacked[tokenIdx * params.mTopK + laneIdx] = packedScore; + } + } +} + +///////////////////////////////////////////////////////////////////////////////////////////////// + +int32_t constexpr getMaxNumExperts(int32_t numExperts) +{ + if (numExperts <= topk::MaxNumExpertsUnit) + { + return topk::MaxNumExpertsUnit; + } + else if (numExperts <= NumExpertsLimit) + { + return NumExpertsLimit; + } + else + { + TLLM_LOG_ERROR("Unsupported numExperts"); + return 0; + } +} + +////////////////////////////////////////////////////////////////////////////////////////////////// + +#define LAUNCH_ROUTING_MINIMAX(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream) \ + if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, false, topk::MaxNumExpertsUnit); \ + } \ + else if (data.mNumExperts <= NumExpertsLimit) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, false, NumExpertsLimit); \ + } \ + else \ + { \ + TLLM_LOG_ERROR("Unsupported numExperts"); \ + } + +////////////////////////////////////////////////////////////////////////////////////////////////// +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) + { + TLLM_CHECK_WITH_INFO(data.mPtrTopKWeights != nullptr, + "When mPtrTopKIds is provided, mPtrTopKWeights must also be provided for MiniMax routing."); + } + TLLM_CHECK_WITH_INFO(data.mPtrPermutedIdxSize != nullptr && data.mPtrCtaIdxXyToBatchIdx != nullptr + && data.mPtrCtaIdxXyToMnLimit != nullptr && data.mPtrNumNonExitingCtas != nullptr, + "MiniMax routing kernel expects permuted idx and grouped Gemm launch config buffers"); + TLLM_CHECK_WITH_INFO(data.mTopK <= MaxSupportedTopExperts, "Routing kernel expects topK experts <= %d, got %d", + MaxSupportedTopExperts, data.mTopK); + TLLM_CHECK_WITH_INFO(data.mNumExperts <= NumExpertsLimit, + "Routing kernel expects #experts %d to be no more than %d", data.mNumExperts, NumExpertsLimit); + TLLM_CHECK_WITH_INFO( + data.mNumExperts % 4 == 0, "Routing kernel expects #experts %d to be a multiple of 4.", data.mNumExperts); + + bool const useSingleBlock = data.mNumTokens <= BlockKernelMaxNumTokens; + + 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."); + TLLM_CHECK_WITH_INFO( + data.mPtrExpertCounts != nullptr, "When #tokens is large, `mPtrExpertCounts` is a required input."); + } + uint32_t const numThreadsHist = getMaxNumExperts(data.mNumExperts); + if (useSingleBlock) + { + LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesBlockKernel, 1, numThreadsHist, + 0, stream); + } + else if (useSingleCluster) + { + LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesClusterKernel, NumBlocksPerCluster, numThreadsHist, + 0, stream); + } + else + { + uint32_t const expandedIdxSize = data.mNumTokens * data.mTopK; + uint32_t const histogramEltsPerBlock = 8 * numThreadsHist; + uint32_t const offsetEltsPerBlock = NumEltsPerOffsetTilePerThread * numThreadsHist; + + uint32_t const maxNumBlocks = 1024; + + int const numBlocksHistogram + = std::min((expandedIdxSize + histogramEltsPerBlock - 1) / histogramEltsPerBlock, maxNumBlocks); + int const numBlocksOffsets + = std::min((expandedIdxSize + offsetEltsPerBlock - 1) / offsetEltsPerBlock, maxNumBlocks); + + if (data.mPtrScores != nullptr && data.mPtrTopKIds == nullptr) + { + LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesHistogramScoresKernel, maxNumBlocks, numThreadsHist, + 0, stream); + } + else + { + LAUNCH_ROUTING_MINIMAX(data, false, routingInitExpertCounts, + (2 * data.mNumExperts - 1) / numThreadsHist + 1, numThreadsHist, + 0, stream); + } + LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesHistogramKernel, numBlocksHistogram, numThreadsHist, + 0, stream); + LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesOffsetsKernel, numBlocksOffsets, numThreadsHist, + 0, stream); + } +} + +////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingMiniMax +} // namespace moe::dev::routing \ No newline at end of file diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu index d750cd8f41e5..3c44753dc20c 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu @@ -205,6 +205,39 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3 moe::dev::routing::routingRenormalize::run(routingData, stream); } + else if (routingMethodType == RoutingMethodType::MiniMax2) + { + moe::dev::routing::routingMiniMax::Data routingData; + routingData.mDtypeExpW = btg::Dtype::Bfloat16; + routingData.mUsePdl = true; + routingData.mNormTopkProb = true; + + routingData.mPtrTopKPacked = routingExpertIndexes; + routingData.mPtrExpertCounts = expertCountHistogram; + routingData.mPtrPermutedIdxSize = permutedIdxSize; + routingData.mPtrExpandedIdxToPermutedIdx = expandedIdxToPermutedIdx; + routingData.mPtrPermutedIdxToExpandedIdx = permutedIdxToExpandedIdx; + routingData.mPtrPermutedIdxToTokenIdx = permutedIdxToTokenIdx; + routingData.mPtrTopKWeights = expertWeights; + routingData.mPtrTopKIds = expertIds; + + routingData.mPtrCtaIdxXyToBatchIdx = ctaIdxXyToBatchIdx; + routingData.mPtrCtaIdxXyToMnLimit = ctaIdxXyToMnLimit; + routingData.mPtrNumNonExitingCtas = numNonExitingCtas; + + routingData.mPtrRoutingBias = routingBias; + routingData.mPtrScores = expertIds == nullptr ? routingLogits : nullptr; + routingData.mNumTokens = numTokens; + routingData.mNumExperts = numExperts; + routingData.mTopK = topK; + routingData.mPaddingLog2 = computeLog2(mTileTokensDim); + routingData.mTileTokensDim = mTileTokensDim; + routingData.mLocalExpertsStartIdx = localExpertOffset; + routingData.mLocalExpertsStrideLog2 = 0; + routingData.mNumLocalExperts = localNumExperts; + + moe::dev::routing::routingMiniMax::run(routingData, stream); + } else { TLLM_CHECK_WITH_INFO(false, "Unimplemented routing method %s of enum %d", diff --git a/tests/unittest/_torch/modules/moe/moe_test_utils.py b/tests/unittest/_torch/modules/moe/moe_test_utils.py index b3987bafaeaa..b4416a0dd9b0 100644 --- a/tests/unittest/_torch/modules/moe/moe_test_utils.py +++ b/tests/unittest/_torch/modules/moe/moe_test_utils.py @@ -121,9 +121,10 @@ def should_skip_trtllm( return None # Routing method compatibility check (used by test_moe_module.py) - # TRTLLMGen C++ routing kernel (runner.cu) only implements: + # TRTLLMGen C++ routing kernel (runner.cu) implements: # - DeepSeekV3 (requires float32 routing_logits) # - Llama4 (requires top_k=1) + # - MiniMaxM2 (sigmoid activation, bias-added selection) # - Renormalize # - RenormalizeNaive # See: cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu:77-212 @@ -138,7 +139,6 @@ def should_skip_trtllm( # Routing methods NOT implemented in C++ kernel trtllm_unimplemented_routing = ( DefaultMoeRoutingMethod, # runner.cu:210 - "Unimplemented routing method" - MiniMaxM2MoeRoutingMethod, # runner.cu:210 - "Unimplemented routing method" ) if routing_method_cls in trtllm_unimplemented_routing: routing_name = routing_method_cls.__name__ From f2f35604a7205d28ee5104832fa3f0188df5c49d Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Fri, 13 Feb 2026 20:42:32 -0800 Subject: [PATCH 02/11] register custom ops --- tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py index ae1b952030e4..97e151df550f 100644 --- a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py @@ -89,6 +89,13 @@ def prepare_dummy_topk_and_hook( lambda: torch.randn( num_experts, dtype=torch.bfloat16, device=hidden_states.device) }) + if routing_method_type == RoutingMethodType.MiniMaxM2: + routing_cls_kwargs.update({ + 'callable_e_score_correction_bias': + lambda: torch.randn( + num_experts, dtype=torch.bfloat16, device=hidden_states.device) + 'num_experts': num_experts, + }) routing_method = ROUTING_METHOD_TYPE_TO_CLASS[routing_method_type]( top_k=top_k, **routing_cls_kwargs) From a5e5dca2a81daad251ae0be05c2ca51b0d7a39bb Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Fri, 13 Feb 2026 21:54:15 -0800 Subject: [PATCH 03/11] minor fixes --- tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py index 97e151df550f..e0ffdbe83414 100644 --- a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py @@ -89,11 +89,11 @@ def prepare_dummy_topk_and_hook( lambda: torch.randn( num_experts, dtype=torch.bfloat16, device=hidden_states.device) }) - if routing_method_type == RoutingMethodType.MiniMaxM2: + if routing_method_type == RoutingMethodType.MiniMax2: routing_cls_kwargs.update({ 'callable_e_score_correction_bias': lambda: torch.randn( - num_experts, dtype=torch.bfloat16, device=hidden_states.device) + num_experts, dtype=torch.bfloat16, device=hidden_states.device), 'num_experts': num_experts, }) routing_method = ROUTING_METHOD_TYPE_TO_CLASS[routing_method_type]( From 2f9e03692e6f06e20167beb4446a09df2c8dab44 Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 12:20:47 -0800 Subject: [PATCH 04/11] fix num experts limit --- .../kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu index a4c9b8ec45de..4d8f7d059209 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -23,7 +23,7 @@ namespace routingMiniMax ////////////////////////////////////////////////////////////////////////////////////////////////// -static constexpr int NumExpertsLimit = 512; +static constexpr int NumExpertsLimit = 256; static constexpr int NumThreads = 1024; static constexpr int NumWarps = NumThreads / WarpSize; @@ -32,7 +32,7 @@ static constexpr int MaxSupportedTopExperts = 10; static constexpr int MaxNumTokensSingleCluster = NumBlocksPerCluster * NumThreads; static constexpr int MaxNumTokensSingleClusterScores = NumBlocksPerCluster * NumWarps; -static constexpr int BlockKernelMaxNumTokens = 4; +static constexpr int BlockKernelMaxNumTokens = 8; template __forceinline__ __device__ void routingTopKExperts(cg::thread_block_tile const& warp, From 117cf549db1f4e0644c5696cd880500fbf58840d Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 14:20:15 -0800 Subject: [PATCH 05/11] fix fused moe autotuner --- .../_torch/modules/fused_moe/fused_moe_trtllm_gen.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py index 8026a7799b44..949e7ee86f5a 100644 --- a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py +++ b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py @@ -522,6 +522,12 @@ def run_moe( n_group = self.routing_method.routing_impl.n_group topk_group = self.routing_method.routing_impl.topk_group routed_scaling_factor = self.routing_method.routing_impl.routed_scaling_factor + elif isinstance(self.routing_method, MiniMaxM2MoeRoutingMethod): + top_k = self.routing_method.top_k + routing_bias = self.routing_method.e_score_correction_bias + n_group = None + topk_group = None + routed_scaling_factor = None else: top_k = self.routing_method.top_k routing_bias = None @@ -1070,8 +1076,9 @@ def forward_fake( else: is_deepseek_v3_routing = isinstance(self.routing_method, DeepSeekV3MoeRoutingMethod) + is_minimax_routing = isinstance(self.routing_method, MiniMaxM2MoeRoutingMethod) top_k = self.routing_method.routing_impl.top_k if is_deepseek_v3_routing else self.routing_method.top_k - routing_bias = self.routing_method.e_score_correction_bias if is_deepseek_v3_routing else None + routing_bias = self.routing_method.e_score_correction_bias if (is_deepseek_v3_routing or is_minimax_routing) else None return fp4_block_scale_fake_output_without_finalize( x, self.num_experts, From 76f6f4df9ef509324f60e6da09e282e4bdcecd86 Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 14:35:27 -0800 Subject: [PATCH 06/11] fixes to python launch method --- tensorrt_llm/_torch/models/modeling_minimaxm2.py | 7 ++++++- .../_torch/modules/fused_moe/fused_moe_trtllm_gen.py | 2 +- 2 files changed, 7 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm2.py b/tensorrt_llm/_torch/models/modeling_minimaxm2.py index 73cd480ee7cc..3380fdced93e 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm2.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm2.py @@ -59,9 +59,14 @@ def __init__( self.gate = Linear( self.hidden_dim, self.num_experts, bias=False, dtype=torch.float32, quant_config=None ) + self.moe_backend = model_config.moe_backend + if self.moe_backend == 'TRTLLM': + bias_dtype = torch.bfloat16 + else: + bias_dtype = torch.float32 self.e_score_correction_bias = nn.Parameter( - torch.empty((self.num_experts), dtype=torch.float32), requires_grad=False + torch.empty((self.num_experts), dtype=bias_dtype), requires_grad=False ) reduce_results = True diff --git a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py index 949e7ee86f5a..a69af42e5b20 100644 --- a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py +++ b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py @@ -42,7 +42,7 @@ W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod, W4A8NVFP4FP8TRTLLMGenFusedMoEMethod, W4A16MXFP4TRTLLMGenFusedMoEMethod) # isort: on -from .routing import BaseMoeRoutingMethod, DeepSeekV3MoeRoutingMethod +from .routing import BaseMoeRoutingMethod, DeepSeekV3MoeRoutingMethod, MiniMaxM2MoeRoutingMethod class TRTLLMGenFusedMoE(MoE): From dc0d51030e1071ad7e92b35e2178d5ded5b7fe6c Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 15:56:57 -0800 Subject: [PATCH 07/11] add simpler kernel --- .../blockScaleMoe/RoutingMiniMax.cu | 740 ++++++------------ .../custom_ops/trtllm_gen_custom_ops.py | 25 +- .../modules/fused_moe/fused_moe_trtllm_gen.py | 2 +- 3 files changed, 269 insertions(+), 498 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu index 4d8f7d059209..7ad2fb1dd270 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -1,9 +1,9 @@ /* * Copyright (c) 2022-2025, NVIDIA CORPORATION. All rights reserved. - * Contibuted by Baseten.co + * Contributed by Baseten.co * * Licensed under the Apache License, Version 2.0 (the "License"); - * You may not use this file except in compliance with 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 @@ -14,6 +14,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + #include "RoutingKernel.cuh" namespace moe::dev::routing @@ -21,501 +22,219 @@ namespace moe::dev::routing namespace routingMiniMax { -////////////////////////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////////////////////// static constexpr int NumExpertsLimit = 256; +static constexpr int MaxSupportedTopExperts = 8; -static constexpr int NumThreads = 1024; -static constexpr int NumWarps = NumThreads / WarpSize; -static constexpr int MaxSupportedTopExperts = 10; - -static constexpr int MaxNumTokensSingleCluster = NumBlocksPerCluster * NumThreads; -static constexpr int MaxNumTokensSingleClusterScores = NumBlocksPerCluster * NumWarps; - -static constexpr int BlockKernelMaxNumTokens = 8; +//////////////////////////////////////////////////////////////////////////////////////////////// -template -__forceinline__ __device__ void routingTopKExperts(cg::thread_block_tile const& warp, - DataType (&score)[VecSize], DataType (&scoreWithBias)[VecSize], int32_t (&idx)[VecSize], - DataType (&warpTopKScore)[MaxSupportedTopExperts], int32_t (&warpTopKExpertIdx)[MaxSupportedTopExperts], - int32_t const laneIdx, int32_t const numExperts, int32_t topK, InputType const* ptrScores, - BiasType const* ptrRoutingBias, bool const normTopkProb) +template +__global__ void routingMainKernel(KernelParams params) { - DataType minScore = DataType{-INFINITY}; - - // Store float scores for accurate shuffle and selection - float score_f[VecSize]; - float scoreWithBias_f[VecSize]; - - // Step 1: Apply sigmoid to logits (NOT softmax - MiniMax2 uses sigmoid) - for (int i = 0; i < VecSize; i++) - { - auto expertIdx = i * WarpSize + laneIdx; - if (expertIdx < numExperts) - { - // Apply sigmoid to get probability (store in float for accuracy) - float logit = static_cast(ptrScores[expertIdx]); - score_f[i] = sigmoid_accurate(logit); - score[i] = static_cast(score_f[i]); - scoreWithBias_f[i] = score_f[i]; // Will add bias below - scoreWithBias[i] = score[i]; // Keep DataType version for compatibility - } - else - { - score_f[i] = 0.0f; - score[i] = DataType{0}; // Invalid experts have 0 probability - scoreWithBias_f[i] = -INFINITY; - scoreWithBias[i] = minScore; // Invalid experts have -inf for selection - } - idx[i] = expertIdx; - } + using OutputT = typename KernelParams::OutputT; - // Step 2: Add routing bias for selection (only for valid experts) - for (int i = 0; i < VecSize; i++) - { - auto expertIdx = i * WarpSize + laneIdx; - if (expertIdx < numExperts) - { - float bias = static_cast(ptrRoutingBias[expertIdx]); - scoreWithBias_f[i] = score_f[i] + bias; - scoreWithBias[i] = static_cast(scoreWithBias_f[i]); - } - } + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + static_assert(MaxNumExperts <= NumExpertsLimit, "MiniMax supports up to 256 experts."); + // MiniMax is configured for topK=8 (enforced at runtime in run()) - // Step 3: Top-K selection using bias-added scores (use float for stability) - topk::reduceTopK(warp, warpTopKScore, warpTopKExpertIdx, scoreWithBias, idx, minScore, topK); + // One token per block + int32_t const tokenIdx = blockIdx.x; + int32_t const expertIdx = threadIdx.x; - // Step 4: Gather original sigmoid scores (not bias-added) for selected experts - // Use warp shuffle to get the unbiased score from the lane that has it - if (laneIdx < topK) - { - int selectedExpertIdx = warpTopKExpertIdx[laneIdx]; - int targetLane = selectedExpertIdx % WarpSize; - int vecIdx = selectedExpertIdx / WarpSize; - - // Select the per-lane value via unrolled loop, then shuffle - float v = 0.0f; - #pragma unroll - for (int i = 0; i < VecSize; ++i) - { - if (vecIdx == i) v = score_f[i]; - } - float originalScore_f = __shfl_sync(0xffffffff, v, targetLane); - warpTopKScore[laneIdx] = static_cast(originalScore_f); - } + // Cooperative groups + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + int32_t const laneIdx = cutlass::arch::LaneId(); + int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); - // Step 5: Renormalize using original sigmoid scores (with epsilon to prevent divide-by-zero) - float sum = 1.0f; - if (normTopkProb) - { - sum = (laneIdx < topK) ? static_cast(warpTopKScore[laneIdx]) : 0.0f; - sum = cg::reduce(warp, sum, cg::plus()) + 1e-20f; // Add epsilon - } - if (laneIdx < topK) - { - float normalized = static_cast(warpTopKScore[laneIdx]) / sum; - warpTopKScore[laneIdx] = static_cast(normalized); - } -} + // Shared memory for per-expert probabilities (DeepSeek style) + __shared__ float __attribute((aligned(128))) smemProb[MaxNumExperts]; + __shared__ float __attribute((aligned(128))) smemSel[MaxNumExperts]; -template -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesBlockKernel(KernelParams params) -{ - // Enforce kernel constraints - if (params.mNumTokens > BlockKernelMaxNumTokens) return; - if (params.mTopK > MaxSupportedTopExperts) return; - - // Guard against thread count mismatch - if (blockDim.x != KernelParams::MaxNumExperts) return; - - using OutputT = typename KernelParams::OutputT; - using InputT = typename KernelParams::InputT; - using BaseType = InputT; - using TypePacked = PackedScoreIdx; - int constexpr MaxNumExperts = KernelParams::MaxNumExperts; + // Invalid handling + static constexpr float kNegInf = float{-INFINITY}; - int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); - int32_t const laneIdx = cutlass::arch::LaneId(); - int32_t const expert = threadIdx.x; - auto scoreOffset = warpIdx * params.mNumExperts; - bool validToken = warpIdx < params.mNumTokens; + // Load and compute per-expert scores for this token + bool const validExpert = expertIdx < params.mNumExperts; + float prob = 0.f; + float sel = kNegInf; - static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; - static constexpr int totalExpertCounts = BlockKernelMaxNumTokens * MaxNumExperts; - __shared__ int8_t __attribute((aligned(128))) smemOffset[totalExpertCounts]; - __shared__ int8_t __attribute((aligned(128))) smemKIdx[totalExpertCounts]; + if (validExpert) + { + // logit layout: token-major, contiguous experts + int64_t scoreIndex = int64_t{tokenIdx} * int64_t{params.mNumExperts} + int64_t{expertIdx}; + float logit = static_cast(params.mPtrScores[scoreIndex]); - using Scan = cub::BlockScan; - __shared__ typename Scan::TempStorage tempStorage; + // MiniMax: sigmoid (not softmax) + prob = sigmoid_accurate(logit); - auto block = cg::this_thread_block(); - auto warp = cg::tiled_partition(block); + float bias = static_cast(params.mPtrRoutingBias[expertIdx]); + sel = prob + bias; // selection score + } - for (int i = threadIdx.x; i < totalExpertCounts; i += blockDim.x) + // Stage to shared so warp0 can index by expert id after topK + if (expertIdx < MaxNumExperts) { - smemOffset[i] = int8_t{-1}; - smemKIdx[i] = int8_t{-1}; + smemProb[expertIdx] = validExpert ? prob : 0.f; // invalid contributes 0 to renorm + smemSel[expertIdx] = validExpert ? sel : kNegInf; // invalid never selected } + __syncthreads(); -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - if constexpr (KernelParams::UsePdl) + // Only warp0 does the final expert selection (DeepSeek style) + if (warpIdx == 0) { - cudaGridDependencySynchronize(); - } -#endif + // Each lane owns VecSize experts: expert = ii*WarpSize + laneIdx + static constexpr int VecSize = MaxNumExperts / WarpSize; // 256 -> 8 + static_assert(MaxNumExperts % WarpSize == 0, "MaxNumExperts must be multiple of 32."); - if (params.mPtrTopKIds != nullptr) - { - if (validToken) + float laneVals[VecSize]; + int32_t laneIdxs[VecSize]; + +#pragma unroll + for (int ii = 0; ii < VecSize; ++ii) { - 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}; - } - } + int e = ii * WarpSize + laneIdx; + laneIdxs[ii] = e; + laneVals[ii] = (e < params.mNumExperts) ? smemSel[e] : kNegInf; } - } - else if (params.mPtrScores != nullptr) - { - BaseType score[VecSize]; - BaseType scoreWithBias[VecSize]; - int32_t idx[VecSize]; - BaseType warpTopKScore[MaxSupportedTopExperts]; - int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; + // TopK outputs + float topScores[MaxSupportedTopExperts]; + int32_t topExperts[MaxSupportedTopExperts]; - BaseType minScore = BaseType{-INFINITY}; - if (validToken) - { - routingTopKExperts(warp, score, scoreWithBias, idx, warpTopKScore, - warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, params.mPtrScores + scoreOffset, - params.mPtrRoutingBias, params.mNormTopkProb); + // Reduce on selection scores (prob+bias) + topk::reduceTopK(warp, topScores, topExperts, laneVals, laneIdxs, kNegInf, params.mTopK); - if (laneIdx < params.mTopK) - { - int offset = warpIdx * MaxNumExperts + warpTopKExpertIdx[laneIdx]; - smemKIdx[offset] = static_cast(laneIdx); - if (params.mPtrTopKWeights != nullptr) - { - params.mPtrTopKWeights[warpIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]}; - } - } - } - } - __syncthreads(); - - auto localExpertIdx = expert - params.mLocalExpertsStartIdx; - int stride = 1 << params.mLocalExpertsStrideLog2; - auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < (params.mNumLocalExperts * stride) - && (localExpertIdx % stride) == 0; - int accExpertCount = 0; + // Convert selection into final weights: + // final = prob (unbiased), optionally renormalized over topK + float w = 0.f; + int32_t chosenExpert = 0; - if (isLocalExpert) - { - int offset = expert; - for (int j = 0; j < BlockKernelMaxNumTokens; j++) +#pragma unroll + for (int ii = 0; ii < MaxSupportedTopExperts; ++ii) { - if (smemKIdx[offset] >= 0) + if (laneIdx == ii) { - smemOffset[offset] = static_cast(accExpertCount); - accExpertCount++; + chosenExpert = topExperts[ii]; + w = (ii < params.mTopK && chosenExpert >= 0 && chosenExpert < params.mNumExperts) + ? smemProb[chosenExpert] + : 0.f; } - offset += MaxNumExperts; } - } - __syncthreads(); - - int32_t numCta; - if constexpr (KernelParams::isPow2) - { - numCta = divUpLog2(accExpertCount, params.mPaddingLog2); - } - else - { - numCta = divUpTileN(accExpertCount, params.mTileTokensDim); - } - int32_t ctaOffset = 0; - int32_t numNonExitingCtas; - Scan(tempStorage).ExclusiveSum(numCta, ctaOffset, numNonExitingCtas); - __syncthreads(); - - int32_t expertScanCounts = 0; - int32_t tmpCount; - if constexpr (KernelParams::isPow2) - { - tmpCount = divUpMulLog2(accExpertCount, params.mPaddingLog2); - } - else - { - tmpCount = divUpMulTileN(accExpertCount, params.mTileTokensDim); - } - Scan(tempStorage).ExclusiveSum(tmpCount, expertScanCounts); - __syncthreads(); - if (isLocalExpert) - { - for (int cta = 0; cta < numCta; ++cta) + // Renormalize within topK if requested + float denom = 1.f; + if (params.mNormTopkProb) { - const int32_t localExpertIdx = (expert - params.mLocalExpertsStartIdx) >> params.mLocalExpertsStrideLog2; - params.mPtrCtaIdxXyToBatchIdx[ctaOffset + cta] = localExpertIdx; - int32_t mnLimit1; - int32_t mnLimit2; - if constexpr (KernelParams::isPow2) - { - mnLimit1 = mulLog2(ctaOffset + cta + 1, params.mPaddingLog2); - mnLimit2 = mulLog2(ctaOffset, params.mPaddingLog2) + accExpertCount; - } - else - { - mnLimit1 = mulTileN(ctaOffset + cta + 1, params.mTileTokensDim); - mnLimit2 = mulTileN(ctaOffset, params.mTileTokensDim) + accExpertCount; - } - params.mPtrCtaIdxXyToMnLimit[ctaOffset + cta] = min(mnLimit1, mnLimit2); + float x = (laneIdx < params.mTopK) ? w : 0.f; + denom = cg::reduce(warp, x, cg::plus{}); + denom += 1e-20f; } - } - if (threadIdx.x == 0) - { - int32_t permutedIdxSize; - if constexpr (KernelParams::isPow2) + float finalW = (laneIdx < params.mTopK) ? (w / denom) : 0.f; + + // Write outputs + int32_t out = tokenIdx * params.mTopK + laneIdx; + + if (laneIdx < params.mTopK && params.mPtrTopKPacked != nullptr) { - permutedIdxSize = mulLog2(numNonExitingCtas, params.mPaddingLog2); + PackedScoreIdx packed{static_cast(finalW), static_cast(chosenExpert)}; + params.mPtrTopKPacked[out] = packed; } - else + if (laneIdx < params.mTopK && params.mPtrTopKWeights != nullptr) { - permutedIdxSize = mulTileN(numNonExitingCtas, params.mTileTokensDim); + params.mPtrTopKWeights[out] = static_cast(finalW); } - params.mPtrPermutedIdxSize[0] = permutedIdxSize; - params.mPtrNumNonExitingCtas[0] = numNonExitingCtas; } +} -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - if constexpr (KernelParams::UsePdl) - { - cudaTriggerProgrammaticLaunchCompletion(); - } -#endif +//////////////////////////////////////////////////////////////////////////////////////////////// - for (int tokenIdx = 0; tokenIdx < params.mNumTokens; tokenIdx++) - { - int offset = tokenIdx * MaxNumExperts + threadIdx.x; - if (smemKIdx[offset] >= 0) - { - int const expandedIdx = tokenIdx * params.mTopK + smemKIdx[offset]; - int const offsetWithinExpert = static_cast(smemOffset[offset]); - int const offsetForExpert = expertScanCounts; - int const permutedIdx = isLocalExpert ? offsetForExpert + offsetWithinExpert : int32_t{-1}; +// Cluster kernel removed for simplification - use histogram path for all token counts - if (params.mPtrExpandedIdxToPermutedIdx != nullptr) - { - params.mPtrExpandedIdxToPermutedIdx[expandedIdx] = permutedIdx; - } - if (params.mPtrPermutedIdxToExpandedIdx != nullptr && isLocalExpert) - { - params.mPtrPermutedIdxToExpandedIdx[permutedIdx] = expandedIdx; - } - if (params.mPtrPermutedIdxToTokenIdx != nullptr && isLocalExpert) - { - params.mPtrPermutedIdxToTokenIdx[permutedIdx] = tokenIdx; - } - } - } -} +//////////////////////////////////////////////////////////////////////////////////////////////// template -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) -__global__ void __cluster_dims__(NumBlocksPerCluster, 1, 1) - __launch_bounds__(KernelParams::MaxNumExperts) - routingIndicesClusterKernel(KernelParams params) -#else -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) - routingIndicesClusterKernel(KernelParams params) -#endif +__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHistogramScoresKernel(KernelParams params) { - // Enforce kernel constraints - if (params.mTopK > MaxSupportedTopExperts) return; - - // Guard against thread count mismatch - if (blockDim.x != KernelParams::MaxNumExperts) return; - using OutputT = typename KernelParams::OutputT; - using InputT = typename KernelParams::InputT; - using BaseType = InputT; - using TypePacked = PackedScoreIdx; - - static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; - // routingPermutation expects NumThreads == MaxNumExperts - static constexpr int PermThreads = KernelParams::MaxNumExperts; - static constexpr int PermWarps = PermThreads / WarpSize; - - // Shared memory for topK results - sized for MaxNumExperts threads - __shared__ TypePacked __attribute((aligned(128))) smemPackedScoreIdx[PermWarps * MaxSupportedTopExperts]; - - uint32_t const clusterBlockRank = blockIdx.x; - - int32_t const warpIdx = threadIdx.x / WarpSize; int32_t const laneIdx = cutlass::arch::LaneId(); - - auto warpTokenIdx = clusterBlockRank * PermWarps + warpIdx; - auto scoreOffset = warpTokenIdx * params.mNumExperts; - bool validToken = warpTokenIdx < params.mNumTokens; + int32_t const warpIdx = threadIdx.x / WarpSize; + int32_t const globalWarpIdx = blockIdx.x * (KernelParams::MaxNumExperts / WarpSize) + warpIdx; + int32_t const globalWarpStride = gridDim.x * (KernelParams::MaxNumExperts / WarpSize); auto block = cg::this_thread_block(); auto warp = cg::tiled_partition(block); -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - if constexpr (KernelParams::UsePdl) - { - cudaGridDependencySynchronize(); - } -#endif + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + static constexpr int VecSize = MaxNumExperts / WarpSize; // 256->8 - if (params.mPtrScores != nullptr) + for (int tokenIdx = globalWarpIdx; tokenIdx < params.mNumTokens; tokenIdx += globalWarpStride) { - BaseType score[VecSize]; - BaseType scoreWithBias[VecSize]; - int32_t idx[VecSize]; - - BaseType warpTopKScore[MaxSupportedTopExperts]; - int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; + // per-lane candidates + float laneVals[VecSize]; + int32_t laneIdxs[VecSize]; - BaseType minScore = BaseType{-INFINITY}; - if (validToken) +// each lane covers VecSize experts +#pragma unroll + for (int ii = 0; ii < VecSize; ++ii) { - routingTopKExperts(warp, score, scoreWithBias, idx, warpTopKScore, - warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, params.mPtrScores + scoreOffset, - params.mPtrRoutingBias, params.mNormTopkProb); + int e = ii * WarpSize + laneIdx; + laneIdxs[ii] = e; - if (laneIdx < params.mTopK) + if (e < params.mNumExperts) { - // Use fixed stride to avoid indexing mistakes - constexpr int STRIDE = MaxSupportedTopExperts; - smemPackedScoreIdx[warpIdx * STRIDE + laneIdx] - = TypePacked{warpTopKScore[laneIdx], static_cast(warpTopKExpertIdx[laneIdx])}; + int64_t scoreIndex = int64_t{tokenIdx} * int64_t{params.mNumExperts} + int64_t{e}; + float logit = static_cast(params.mPtrScores[scoreIndex]); + float prob = sigmoid_accurate(logit); + float bias = static_cast(params.mPtrRoutingBias[e]); + laneVals[ii] = prob + bias; // selection } - } - } - - // Pack topK results into shared memory for all input types - if (params.mPtrScores != nullptr) - { - // Already packed by routingTopKExperts above - } - else if (params.mPtrTopKIds != nullptr && params.mPtrTopKWeights != nullptr) - { - // Pack from pre-computed TopKIds and TopKWeights - if (validToken && laneIdx < params.mTopK) - { - int id = params.mPtrTopKIds[warpTokenIdx * params.mTopK + laneIdx]; - BaseType w = static_cast(params.mPtrTopKWeights[warpTokenIdx * params.mTopK + laneIdx]); - // Handle invalid expert IDs explicitly - if (id == -1) { - smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = - TypePacked{BaseType{0}, int16_t{-1}}; - } else { - smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = - TypePacked{w, static_cast(id)}; + else + { + laneVals[ii] = float{-INFINITY}; } } - } - else if (params.mPtrTopKPacked != nullptr) - { - // Pack from pre-computed packed format with explicit type casting - if (validToken && laneIdx < params.mTopK) - { - auto p = params.mPtrTopKPacked[warpTokenIdx * params.mTopK + laneIdx]; - BaseType w = static_cast(p.score); - int16_t id = p.idx; - smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = TypePacked{w, id}; - } - } - - // Initialize unused lanes to prevent stale data - if (validToken && laneIdx >= params.mTopK && laneIdx < MaxSupportedTopExperts) - { - smemPackedScoreIdx[warpIdx * MaxSupportedTopExperts + laneIdx] = - TypePacked{BaseType{0}, int16_t{-1}}; - } - - // Synchronize before using shared memory - __syncthreads(); - - // Use fixed stride for shared memory layout - // LoadExpertIdxFromGlobal=false: load from shared memory (smemPackedScoreIdx) - constexpr int STRIDE = MaxSupportedTopExperts; - routingPermutation(params, smemPackedScoreIdx, warpIdx, clusterBlockRank); -} -template -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHistogramScoresKernel(KernelParams params) -{ - using OutputT = typename KernelParams::OutputT; - using InputT = typename KernelParams::InputT; - using BaseType = InputT; + float topSel[MaxSupportedTopExperts]; + int32_t topExp[MaxSupportedTopExperts]; + topk::reduceTopK(warp, topSel, topExp, laneVals, laneIdxs, float{-INFINITY}, params.mTopK); - static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; - - int32_t const laneIdx = cutlass::arch::LaneId(); - int32_t const warpIdx = threadIdx.x / WarpSize; - int32_t const globalWarpIdx = blockIdx.x * KernelParams::MaxNumExperts / WarpSize + warpIdx; - int32_t const globalWarpStride = gridDim.x * KernelParams::MaxNumExperts / WarpSize; - BaseType minScore = BaseType{-INFINITY}; - auto block = cg::this_thread_block(); - auto warp = cg::tiled_partition(block); - -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - if constexpr (KernelParams::UsePdl) - { - cudaGridDependencySynchronize(); - } -#endif - - int32_t expertCountsNum = 2 * params.mNumExperts; - int32_t globalThreadIdx = blockIdx.x * KernelParams::MaxNumExperts + threadIdx.x; - int32_t globalThreadStride = gridDim.x * KernelParams::MaxNumExperts; - initArr(globalThreadIdx, expertCountsNum, globalThreadStride, params.mPtrExpertCounts, 0); - -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - if constexpr (KernelParams::UsePdl) - { - cudaTriggerProgrammaticLaunchCompletion(); - } -#endif + // produce packed output weights from *unbiased prob* (sigmoid only) + if (laneIdx < params.mTopK) + { + int e = topExp[laneIdx]; - BaseType allScores[VecSize]; - BaseType allScoresWithBias[VecSize]; - int32_t allExpertIdx[VecSize]; - BaseType warpTopKScore[MaxSupportedTopExperts]; - int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; - for (int tokenIdx = globalWarpIdx; tokenIdx < params.mNumTokens; tokenIdx += globalWarpStride) - { - auto scoreOffset = tokenIdx * params.mNumExperts; + float prob = 0.f; + if (e >= 0 && e < params.mNumExperts) + { + int64_t scoreIndex = int64_t{tokenIdx} * int64_t{params.mNumExperts} + int64_t{e}; + float logit = static_cast(params.mPtrScores[scoreIndex]); + prob = sigmoid_accurate(logit); + } - routingTopKExperts(warp, allScores, allScoresWithBias, allExpertIdx, - warpTopKScore, warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, - params.mPtrScores + scoreOffset, params.mPtrRoutingBias, params.mNormTopkProb); + // renorm if requested + float denom = 1.f; + if (params.mNormTopkProb) + { + float x = (laneIdx < params.mTopK) ? prob : 0.f; + denom = cg::reduce(warp, x, cg::plus{}) + 1e-20f; + } + float finalW = (laneIdx < params.mTopK) ? (prob / denom) : 0.f; - if (laneIdx < params.mTopK) - { - PackedScoreIdx packedScore{ - static_cast(warpTopKScore[laneIdx]), static_cast(warpTopKExpertIdx[laneIdx])}; - params.mPtrTopKPacked[tokenIdx * params.mTopK + laneIdx] = packedScore; + PackedScoreIdx packed{static_cast(finalW), static_cast(e)}; + if (params.mPtrTopKPacked != nullptr && laneIdx < params.mTopK) + { + params.mPtrTopKPacked[tokenIdx * params.mTopK + laneIdx] = packed; + } } } } -///////////////////////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////////////////////// int32_t constexpr getMaxNumExperts(int32_t numExperts) { @@ -534,100 +253,133 @@ int32_t constexpr getMaxNumExperts(int32_t numExperts) } } -////////////////////////////////////////////////////////////////////////////////////////////////// - -#define LAUNCH_ROUTING_MINIMAX(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream) \ - if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ - { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ - data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, false, topk::MaxNumExpertsUnit); \ - } \ - else if (data.mNumExperts <= NumExpertsLimit) \ - { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ - data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, false, NumExpertsLimit); \ - } \ - else \ - { \ - TLLM_LOG_ERROR("Unsupported numExperts"); \ +#define LAUNCH_ROUTING_MINIMAX(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1) \ + if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, topk::MaxNumExpertsUnit); \ + } \ + else if (data.mNumExperts <= NumExpertsLimit) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, NumExpertsLimit); \ + } \ + else \ + { \ + TLLM_LOG_ERROR("Unsupported numExperts"); \ } -////////////////////////////////////////////////////////////////////////////////////////////////// void run(Data const& data, void* stream) { - TLLM_CHECK_WITH_INFO(data.mPtrTopKPacked != nullptr || data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr, + // Create non-const alias for launch macros that may require non-const data + auto& d = const_cast(data); + + TLLM_CHECK_WITH_INFO(d.mPtrTopKPacked != nullptr || d.mPtrScores != nullptr || d.mPtrTopKIds != nullptr, "Routing kernel requires at least one input parameter"); - if (data.mPtrTopKIds != nullptr) + + if (d.mPtrTopKIds != nullptr) { - TLLM_CHECK_WITH_INFO(data.mPtrTopKWeights != nullptr, + TLLM_CHECK_WITH_INFO(d.mPtrTopKWeights != nullptr, "When mPtrTopKIds is provided, mPtrTopKWeights must also be provided for MiniMax routing."); } - TLLM_CHECK_WITH_INFO(data.mPtrPermutedIdxSize != nullptr && data.mPtrCtaIdxXyToBatchIdx != nullptr - && data.mPtrCtaIdxXyToMnLimit != nullptr && data.mPtrNumNonExitingCtas != nullptr, - "MiniMax routing kernel expects permuted idx and grouped Gemm launch config buffers"); - TLLM_CHECK_WITH_INFO(data.mTopK <= MaxSupportedTopExperts, "Routing kernel expects topK experts <= %d, got %d", - MaxSupportedTopExperts, data.mTopK); - TLLM_CHECK_WITH_INFO(data.mNumExperts <= NumExpertsLimit, - "Routing kernel expects #experts %d to be no more than %d", data.mNumExperts, NumExpertsLimit); - TLLM_CHECK_WITH_INFO( - data.mNumExperts % 4 == 0, "Routing kernel expects #experts %d to be a multiple of 4.", data.mNumExperts); - bool const useSingleBlock = data.mNumTokens <= BlockKernelMaxNumTokens; + // Permutation outputs required by grouped GEMM launch + TLLM_CHECK_WITH_INFO(d.mPtrPermutedIdxSize != nullptr && d.mPtrCtaIdxXyToBatchIdx != nullptr + && d.mPtrCtaIdxXyToMnLimit != nullptr && d.mPtrNumNonExitingCtas != nullptr, + "MiniMax routing expects permuted idx and grouped GEMM launch config buffers"); - bool const useSingleCluster = data.mNumTokens <= ((data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr) - ? MaxNumTokensSingleClusterScores - : MaxNumTokensSingleCluster); + TLLM_CHECK_WITH_INFO(d.mTopK == 8, "MiniMax is configured for topK=8, got %d", d.mTopK); + TLLM_CHECK_WITH_INFO(d.mNumExperts == 256, "MiniMax is configured for 256 experts, got %d", d.mNumExperts); - if (!useSingleCluster && !useSingleBlock) + TLLM_CHECK_WITH_INFO(d.mNumExperts % 4 == 0, "Routing expects #experts multiple of 4, got %d", d.mNumExperts); + + // This "DeepSeek-like" path uses a block-per-token routingMainKernel when we need routing + int const numThreadsHist = getMaxNumExperts(d.mNumExperts); + + // Simplified: always use histogram path for permutation building + TLLM_CHECK_WITH_INFO((d.mPtrTopKPacked != nullptr || d.mPtrTopKIds != nullptr), + "MiniMax requires `mPtrTopKPacked` or `mPtrTopKIds` for permutation building."); + TLLM_CHECK_WITH_INFO( + d.mPtrExpertCounts != nullptr, "MiniMax requires `mPtrExpertCounts` for permutation building."); + + // 1) If TopK not provided, compute it - choose path based on token count + if (d.mPtrTopKIds == nullptr) { - TLLM_CHECK_WITH_INFO((data.mPtrTopKPacked != nullptr || data.mPtrTopKIds != nullptr), - "When #tokens is large, `mPtrTopKPacked` or `mPtrTopKIds` is a required input."); + TLLM_CHECK_WITH_INFO(d.mPtrScores != nullptr, "If mPtrTopKIds is null, mPtrScores must be provided."); TLLM_CHECK_WITH_INFO( - data.mPtrExpertCounts != nullptr, "When #tokens is large, `mPtrExpertCounts` is a required input."); - } - uint32_t const numThreadsHist = getMaxNumExperts(data.mNumExperts); - if (useSingleBlock) - { - LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesBlockKernel, 1, numThreadsHist, - 0, stream); - } - else if (useSingleCluster) - { - LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesClusterKernel, NumBlocksPerCluster, numThreadsHist, - 0, stream); + d.mPtrTopKPacked != nullptr, "MiniMax requires mPtrTopKPacked when computing topK from scores."); + + // Choose computation path based on token count to avoid double work + constexpr int SmallTokenThreshold = 256; // Same as MaxNumTokensSingleClusterScores + bool const useSmallTokenPath = d.mNumTokens <= SmallTokenThreshold; + + if (useSmallTokenPath) + { + // Small token count: use block-per-token routingMainKernel + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingMainKernel, + /*numBlocks=*/d.mNumTokens, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream, + /*extraFlag1=*/false); + } + // else: large token count - will use routingIndicesHistogramScoresKernel below } - else - { - uint32_t const expandedIdxSize = data.mNumTokens * data.mTopK; - uint32_t const histogramEltsPerBlock = 8 * numThreadsHist; - uint32_t const offsetEltsPerBlock = NumEltsPerOffsetTilePerThread * numThreadsHist; - uint32_t const maxNumBlocks = 1024; + // 2) Build permutation / CTA schedule (always use histogram path) + if (d.mPtrPermutedIdxSize != nullptr) + { + // Histogram + offsets path + int32_t const expandedIdxSize = d.mNumTokens * d.mTopK; + int32_t const histogramEltsPerBlock = 8 * numThreadsHist; + int32_t const offsetEltsPerBlock = NumEltsPerOffsetTilePerThread * numThreadsHist; + int32_t const maxNumBlocks = 1024; int const numBlocksHistogram = std::min((expandedIdxSize + histogramEltsPerBlock - 1) / histogramEltsPerBlock, maxNumBlocks); int const numBlocksOffsets = std::min((expandedIdxSize + offsetEltsPerBlock - 1) / offsetEltsPerBlock, maxNumBlocks); - if (data.mPtrScores != nullptr && data.mPtrTopKIds == nullptr) - { - LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesHistogramScoresKernel, maxNumBlocks, numThreadsHist, - 0, stream); - } - else + // Always initialize expert counts first (avoid race conditions) + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingInitExpertCounts, + /*numBlocks=*/(2 * d.mNumExperts - 1) / numThreadsHist + 1, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream, + /*extraFlag1=*/false); + + // Only compute topK from scores if we didn't already do it in routingMainKernel + constexpr int SmallTokenThreshold = 256; + bool const usedSmallTokenPath = d.mNumTokens <= SmallTokenThreshold; + + if (d.mPtrScores != nullptr && d.mPtrTopKIds == nullptr && !usedSmallTokenPath) { - LAUNCH_ROUTING_MINIMAX(data, false, routingInitExpertCounts, - (2 * data.mNumExperts - 1) / numThreadsHist + 1, numThreadsHist, - 0, stream); + // produce mPtrTopKPacked from scores (sigmoid+bias selection) for large token counts + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingIndicesHistogramScoresKernel, + /*numBlocks=*/maxNumBlocks, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream, + /*extraFlag1=*/false); } - LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesHistogramKernel, numBlocksHistogram, numThreadsHist, - 0, stream); - LAUNCH_ROUTING_MINIMAX(data, false, routingIndicesOffsetsKernel, numBlocksOffsets, numThreadsHist, - 0, stream); + + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingIndicesHistogramKernel, + /*numBlocks=*/numBlocksHistogram, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream, + /*extraFlag1=*/false); + + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingIndicesOffsetsKernel, + /*numBlocks=*/numBlocksOffsets, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream, + /*extraFlag1=*/false); } } -////////////////////////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////////////////////// } // namespace routingMiniMax } // namespace moe::dev::routing \ No newline at end of file diff --git a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py index e0ffdbe83414..63c871b6f5c0 100644 --- a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py @@ -6,6 +6,7 @@ from tensorrt_llm._torch.modules.fused_moe.routing import ( ROUTING_METHOD_TYPE_TO_CLASS, RoutingMethodType) +from tensorrt_llm._utils import get_sm_version from tensorrt_llm._torch.utils import (ActType_TrtllmGen, Fp4QuantizedTensor, fp4_utils, get_last_power_of_2_num_tokens_buckets, @@ -678,10 +679,28 @@ def _constrain_to_num_tokens(shapes: Tuple[torch.Size]) -> int: return num_tokens HS_SCALE_IDX = 3 - CONSTRAINED_HS_SCALE_DIM = 1 - constraint_hidden_states_scale = ConstraintSpec( - HS_SCALE_IDX, CONSTRAINED_HS_SCALE_DIM, _constrain_to_num_tokens) + if get_sm_version() >= 100: + # SM100+: fp8_quantize_1x128 returns 2D scale (blocked_n, num_tokens) + CONSTRAINED_HS_SCALE_DIM = 1 + constraint_hidden_states_scale = ConstraintSpec( + HS_SCALE_IDX, CONSTRAINED_HS_SCALE_DIM, + _constrain_to_num_tokens) + else: + # SM90: fp8_quantize_1x128 returns 1D scale with layout matching + # the fp8_quantize_1x128 custom op shape formula. + def _constrain_hs_scale_sm90( + shapes: Tuple[torch.Size]) -> int: + num_tokens = shapes[2][0] + hidden_size = shapes[2][1] + pad_m = fp4_utils.pad_up(num_tokens, 4) + blocked_n = (hidden_size + 127) // 128 + return fp4_utils.pad_up(pad_m * blocked_n * 4, 128) // 4 + + CONSTRAINED_HS_SCALE_DIM = 0 + constraint_hidden_states_scale = ConstraintSpec( + HS_SCALE_IDX, CONSTRAINED_HS_SCALE_DIM, + _constrain_hs_scale_sm90) ROUTER_LOGITS_IDX = 0 CONSTRAINED_RL_DIM = 0 diff --git a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py index a69af42e5b20..b286ccae74f0 100644 --- a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py +++ b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py @@ -550,7 +550,7 @@ def run_moe( if self.has_deepseek_fp8_block_scales: assert do_finalize, "fp8_block_scale_moe_runner does not support do_finalize=False" - # fp8_block_scale_moe_runner needs 2D shape for x_sf and only support SM100+ + # fp8_quantize_1x128 returns 2D x_sf on SM100+, 1D on SM90 if x_sf is None: x, x_sf = torch.ops.trtllm.fp8_quantize_1x128(x) From 2619809a9c1d2ddce81292fc1fda6672dd4b370a Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 16:22:57 -0800 Subject: [PATCH 08/11] cuda minimax use fp32 logits --- .../blockScaleMoe/RoutingMiniMax.cu | 33 +++++++++++++++++-- cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp | 3 +- 2 files changed, 33 insertions(+), 3 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu index 7ad2fb1dd270..6ce0a0e95b9a 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -253,15 +253,44 @@ int32_t constexpr getMaxNumExperts(int32_t numExperts) } } +// MiniMax-specific dispatch: InputT is always float (gate is float32). +// Only OutputT varies based on mDtypeExpW (bf16 for bias/weights, or float). +#define LAUNCH_ROUTING_MINIMAX_IMPL( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, numExperts) \ + if (data.mDtypeExpW == tg::Dtype::Fp32 && extraFlag1) \ + { \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, true), kernel, numBlocks, numThreads, \ + smemSize, stream); \ + } \ + else if (data.mDtypeExpW == tg::Dtype::Fp32) \ + { \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, false), kernel, numBlocks, numThreads, \ + smemSize, stream); \ + } \ + else if (data.mDtypeExpW == tg::Dtype::Bfloat16 && extraFlag1) \ + { \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, __nv_bfloat16, numExperts, true), kernel, numBlocks, \ + numThreads, smemSize, stream); \ + } \ + else if (data.mDtypeExpW == tg::Dtype::Bfloat16) \ + { \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, __nv_bfloat16, numExperts, false), kernel, numBlocks, \ + numThreads, smemSize, stream); \ + } \ + else \ + { \ + TLLM_LOG_ERROR("Unsupported dtypeExpW"); \ + } + #define LAUNCH_ROUTING_MINIMAX(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1) \ if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ + LAUNCH_ROUTING_MINIMAX_IMPL( \ data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, topk::MaxNumExpertsUnit); \ } \ else if (data.mNumExperts <= NumExpertsLimit) \ { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ + LAUNCH_ROUTING_MINIMAX_IMPL( \ data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, NumExpertsLimit); \ } \ else \ diff --git a/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp b/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp index 2db4e2bf6c5b..172489c343de 100644 --- a/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp +++ b/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp @@ -63,7 +63,8 @@ at::Tensor run_fp8_block_scale_moe(at::optional const& routing_logit } else if (routing_logits.has_value()) { - if (static_cast(routing_method_type) == RoutingMethodType::DeepSeekV3) + if (static_cast(routing_method_type) == RoutingMethodType::DeepSeekV3 + || static_cast(routing_method_type) == RoutingMethodType::MiniMax2) { TORCH_CHECK(routing_logits.value().scalar_type() == at::ScalarType::Float, "routing_logits must be float"); } From 3a3f1433da11cd794280711c597d91c6649703d2 Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 16:28:01 -0800 Subject: [PATCH 09/11] simplify kernel --- .../blockScaleMoe/RoutingMiniMax.cu | 39 ++++++------------- 1 file changed, 12 insertions(+), 27 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu index 6ce0a0e95b9a..dc3cefa22d54 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -254,24 +254,14 @@ int32_t constexpr getMaxNumExperts(int32_t numExperts) } // MiniMax-specific dispatch: InputT is always float (gate is float32). -// Only OutputT varies based on mDtypeExpW (bf16 for bias/weights, or float). -#define LAUNCH_ROUTING_MINIMAX_IMPL( \ - data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, numExperts) \ - if (data.mDtypeExpW == tg::Dtype::Fp32 && extraFlag1) \ - { \ - LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, true), kernel, numBlocks, numThreads, \ - smemSize, stream); \ - } \ - else if (data.mDtypeExpW == tg::Dtype::Fp32) \ +// OutputT varies based on mDtypeExpW (bf16 for bias/weights, or float). +// DoSoftmaxBeforeTopK is always false (MiniMax uses sigmoid, not softmax). +#define LAUNCH_ROUTING_MINIMAX_IMPL(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, numExperts) \ + if (data.mDtypeExpW == tg::Dtype::Fp32) \ { \ LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, false), kernel, numBlocks, numThreads, \ smemSize, stream); \ } \ - else if (data.mDtypeExpW == tg::Dtype::Bfloat16 && extraFlag1) \ - { \ - LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, __nv_bfloat16, numExperts, true), kernel, numBlocks, \ - numThreads, smemSize, stream); \ - } \ else if (data.mDtypeExpW == tg::Dtype::Bfloat16) \ { \ LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, __nv_bfloat16, numExperts, false), kernel, numBlocks, \ @@ -282,16 +272,16 @@ int32_t constexpr getMaxNumExperts(int32_t numExperts) TLLM_LOG_ERROR("Unsupported dtypeExpW"); \ } -#define LAUNCH_ROUTING_MINIMAX(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1) \ +#define LAUNCH_ROUTING_MINIMAX(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream) \ if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ { \ LAUNCH_ROUTING_MINIMAX_IMPL( \ - data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, topk::MaxNumExpertsUnit); \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, topk::MaxNumExpertsUnit); \ } \ else if (data.mNumExperts <= NumExpertsLimit) \ { \ LAUNCH_ROUTING_MINIMAX_IMPL( \ - data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, NumExpertsLimit); \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, NumExpertsLimit); \ } \ else \ { \ @@ -349,8 +339,7 @@ void run(Data const& data, void* stream) /*coopLaunch=*/false, routingMainKernel, /*numBlocks=*/d.mNumTokens, /*numThreads=*/numThreadsHist, - /*smemSize=*/0, stream, - /*extraFlag1=*/false); + /*smemSize=*/0, stream); } // else: large token count - will use routingIndicesHistogramScoresKernel below } @@ -374,8 +363,7 @@ void run(Data const& data, void* stream) /*coopLaunch=*/false, routingInitExpertCounts, /*numBlocks=*/(2 * d.mNumExperts - 1) / numThreadsHist + 1, /*numThreads=*/numThreadsHist, - /*smemSize=*/0, stream, - /*extraFlag1=*/false); + /*smemSize=*/0, stream); // Only compute topK from scores if we didn't already do it in routingMainKernel constexpr int SmallTokenThreshold = 256; @@ -388,23 +376,20 @@ void run(Data const& data, void* stream) /*coopLaunch=*/false, routingIndicesHistogramScoresKernel, /*numBlocks=*/maxNumBlocks, /*numThreads=*/numThreadsHist, - /*smemSize=*/0, stream, - /*extraFlag1=*/false); + /*smemSize=*/0, stream); } LAUNCH_ROUTING_MINIMAX(d, /*coopLaunch=*/false, routingIndicesHistogramKernel, /*numBlocks=*/numBlocksHistogram, /*numThreads=*/numThreadsHist, - /*smemSize=*/0, stream, - /*extraFlag1=*/false); + /*smemSize=*/0, stream); LAUNCH_ROUTING_MINIMAX(d, /*coopLaunch=*/false, routingIndicesOffsetsKernel, /*numBlocks=*/numBlocksOffsets, /*numThreads=*/numThreadsHist, - /*smemSize=*/0, stream, - /*extraFlag1=*/false); + /*smemSize=*/0, stream); } } From 99b0168c1ab5b7e1e9db83226e1888e92d53b1bb Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 16:39:31 -0800 Subject: [PATCH 10/11] simplify kernel - fix bug --- .../kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu index dc3cefa22d54..53ed4213f69e 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -230,6 +230,10 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHis { params.mPtrTopKPacked[tokenIdx * params.mTopK + laneIdx] = packed; } + if (params.mPtrTopKWeights != nullptr && laneIdx < params.mTopK) + { + params.mPtrTopKWeights[tokenIdx * params.mTopK + laneIdx] = static_cast(finalW); + } } } } From b5cc89084490d192aa1c8556925d6cdaba13592d Mon Sep 17 00:00:00 2001 From: michaelfeil <63565275+michaelfeil@users.noreply.github.com> Date: Sat, 14 Feb 2026 17:04:23 -0800 Subject: [PATCH 11/11] simplify kernel - fix bug warp reduce --- .../blockScaleMoe/RoutingMiniMax.cu | 34 +++++++++++-------- 1 file changed, 20 insertions(+), 14 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu index 53ed4213f69e..4430d6e4d684 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -203,34 +203,40 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHis int32_t topExp[MaxSupportedTopExperts]; topk::reduceTopK(warp, topSel, topExp, laneVals, laneIdxs, float{-INFINITY}, params.mTopK); - // produce packed output weights from *unbiased prob* (sigmoid only) + // Produce final weights from *unbiased prob* (sigmoid only), renormalized over topK. + // All 32 warp lanes must participate in cg::reduce, so compute prob for all lanes + // and mask non-topK lanes to 0. + float prob = 0.f; + int32_t e = -1; if (laneIdx < params.mTopK) { - int e = topExp[laneIdx]; - - float prob = 0.f; + e = topExp[laneIdx]; if (e >= 0 && e < params.mNumExperts) { int64_t scoreIndex = int64_t{tokenIdx} * int64_t{params.mNumExperts} + int64_t{e}; float logit = static_cast(params.mPtrScores[scoreIndex]); prob = sigmoid_accurate(logit); } + } - // renorm if requested - float denom = 1.f; - if (params.mNormTopkProb) - { - float x = (laneIdx < params.mTopK) ? prob : 0.f; - denom = cg::reduce(warp, x, cg::plus{}) + 1e-20f; - } - float finalW = (laneIdx < params.mTopK) ? (prob / denom) : 0.f; + // Renorm: all 32 lanes participate, non-topK lanes contribute 0 + float finalW = prob; + if (params.mNormTopkProb) + { + float x = (laneIdx < params.mTopK) ? prob : 0.f; + float denom = cg::reduce(warp, x, cg::plus{}) + 1e-20f; + finalW = (laneIdx < params.mTopK) ? (prob / denom) : 0.f; + } + // Write outputs (only topK lanes) + if (laneIdx < params.mTopK) + { PackedScoreIdx packed{static_cast(finalW), static_cast(e)}; - if (params.mPtrTopKPacked != nullptr && laneIdx < params.mTopK) + if (params.mPtrTopKPacked != nullptr) { params.mPtrTopKPacked[tokenIdx * params.mTopK + laneIdx] = packed; } - if (params.mPtrTopKWeights != nullptr && laneIdx < params.mTopK) + if (params.mPtrTopKWeights != nullptr) { params.mPtrTopKWeights[tokenIdx * params.mTopK + laneIdx] = static_cast(finalW); }