From ce58d298380a99c28300613cc19d0351b25ab48d Mon Sep 17 00:00:00 2001 From: Christina Zhang <83400082+ChristinaZ@users.noreply.github.com> Date: Fri, 13 Feb 2026 02:44:57 -0800 Subject: [PATCH] Add support for expert_number=2048 and K=32 Signed-off-by: Christina Zhang <83400082+ChristinaZ@users.noreply.github.com> --- cpp/tensorrt_llm/kernels/noAuxTcKernels.cu | 14 +- .../blockScaleMoe/CMakeLists.txt | 3 + .../blockScaleMoe/DevKernel.h | 85 +-- .../blockScaleMoe/RoutingDeepSeek.cu | 627 +----------------- .../blockScaleMoe/RoutingKernel.cuh | 413 +++++++----- .../blockScaleMoe/RoutingKernel.h | 16 +- .../blockScaleMoe/RoutingKernelTopK.cuh | 155 +++-- .../blockScaleMoe/RoutingLlama4.cu | 8 +- .../blockScaleMoe/RoutingRenormalize.cu | 493 +------------- .../routingDeepSeek/RoutingDeepSeekCommon.cuh | 115 ++++ .../routingDeepSeek/launchClusterKernel.cu | 64 ++ .../routingDeepSeek/launchCoopKernel.cu | 276 ++++++++ .../routingDeepSeek/launchHistogramKernel.cu | 36 + .../routingDeepSeek/launchInitExpertCounts.cu | 36 + .../routingDeepSeek/launchMainKernel.cu | 289 ++++++++ .../routingDeepSeek/launchOffsetsKernel.cu | 36 + .../RoutingRenormalizeCommon.cuh | 160 +++++ .../routingRenormalize/launchBlockKernel.cu | 294 ++++++++ .../routingRenormalize/launchClusterKernel.cu | 117 ++++ .../launchHistogramKernel.cu | 35 + .../launchHistogramScoresKernel.cu | 106 +++ .../launchInitExpertCounts.cu | 36 + .../routingRenormalize/launchOffsetsKernel.cu | 35 + cpp/tensorrt_llm/thop/fp4BlockScaleMoe.cpp | 6 +- cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp | 4 +- .../thop/fp8PerTensorScaleMoe.cpp | 6 + cpp/tensorrt_llm/thop/mxFp4BlockScaleMoe.cpp | 4 +- .../routing/routingRenormalizeTest.cpp | 10 +- .../kernels/routing/routingTest.cpp | 9 +- tests/unittest/_torch/thop/serial/test_moe.py | 26 +- 30 files changed, 2115 insertions(+), 1399 deletions(-) create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/RoutingDeepSeekCommon.cuh create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchClusterKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchCoopKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchHistogramKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchInitExpertCounts.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchMainKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchOffsetsKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/RoutingRenormalizeCommon.cuh create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchClusterKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramScoresKernel.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchInitExpertCounts.cu create mode 100644 cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchOffsetsKernel.cu diff --git a/cpp/tensorrt_llm/kernels/noAuxTcKernels.cu b/cpp/tensorrt_llm/kernels/noAuxTcKernels.cu index efa69c709886..21f68c71824d 100644 --- a/cpp/tensorrt_llm/kernels/noAuxTcKernels.cu +++ b/cpp/tensorrt_llm/kernels/noAuxTcKernels.cu @@ -208,23 +208,23 @@ __global__ void deepseek_v3_topk_kernel(InputT* scores, OutputT* topkValues, Idx if (warpIdx == 0) { int constexpr NumInterTopKPerThread = (NumInterTopK - 1) / WARP_SIZE + 1; - float intermidiateScore[NumInterTopKPerThread]; - int32_t intermidiateExpert[NumInterTopKPerThread]; + float intermediateScore[NumInterTopKPerThread]; + int32_t intermediateExpert[NumInterTopKPerThread]; for (int i = laneIdx; i < NumInterTopKPerThread * WARP_SIZE; i += WARP_SIZE) { int ii = i / WARP_SIZE; if (i < NumInterTopK) { - intermidiateScore[ii] = smemInterTopScores[i]; - intermidiateExpert[ii] = smemInterTopExperts[i]; + intermediateScore[ii] = smemInterTopScores[i]; + intermediateExpert[ii] = smemInterTopExperts[i]; } else { - intermidiateScore[ii] = invalidScoreFloat; - intermidiateExpert[ii] = MaxNumExperts - 1; + intermediateScore[ii] = invalidScoreFloat; + intermediateExpert[ii] = MaxNumExperts - 1; } } - reduce_topk::reduceTopK(warp, topScores, topExperts, intermidiateScore, intermidiateExpert, + reduce_topk::reduceTopK(warp, topScores, topExperts, intermediateScore, intermediateExpert, /* minValue */ invalidScoreFloat, topk); } } diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/CMakeLists.txt b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/CMakeLists.txt index bf3e92fcc65f..65ff66fbed0a 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/CMakeLists.txt +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/CMakeLists.txt @@ -23,3 +23,6 @@ set_property(TARGET trtllm_gen_fp8_block_scale_moe PROPERTY POSITION_INDEPENDENT_CODE ON) set_property(TARGET trtllm_gen_fp8_block_scale_moe PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) +# Split-compile to parallelize ptxas for files with many kernel instantiations. +target_compile_options(trtllm_gen_fp8_block_scale_moe + PRIVATE $<$:--split-compile=0>) diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/DevKernel.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/DevKernel.h index 8da081b09770..7f8e2b06b00b 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/DevKernel.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/DevKernel.h @@ -34,35 +34,40 @@ namespace moe::dev #define LAUNCH_ESC(...) __VA_ARGS__ #define LAUNCH_PDL(data, coopLaunch, types, kernel, numBlocks, numThreads, smemSize, stream) \ - cudaLaunchConfig_t config{}; \ - config.gridDim = numBlocks; \ - config.blockDim = numThreads; \ - config.dynamicSmemBytes = smemSize; \ - config.stream = (cudaStream_t) stream; \ - \ - cudaLaunchAttribute attributes[2] = {}; \ - attributes[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; \ - attributes[0].val.programmaticStreamSerializationAllowed = int(data.mUsePdl); \ - attributes[1].id = cudaLaunchAttributeCooperative; \ - attributes[1].val.cooperative = int(coopLaunch); \ - config.attrs = attributes; \ - config.numAttrs = 2; \ - if (data.mUsePdl) \ - { \ - auto params = KernelParams::setKernelParams(data); \ - auto kernelTyped = kernel>; \ - if (smemSize > 48 * 1024) \ - TLLM_CUDA_CHECK(cudaFuncSetAttribute(kernelTyped, cudaFuncAttributeMaxDynamicSharedMemorySize, smemSize)); \ - TLLM_CUDA_CHECK(cudaLaunchKernelEx(&config, kernelTyped, params)); \ - } \ - else \ + do \ { \ - auto params = KernelParams::setKernelParams(data); \ - auto kernelTyped = kernel>; \ - if (smemSize > 48 * 1024) \ - TLLM_CUDA_CHECK(cudaFuncSetAttribute(kernelTyped, cudaFuncAttributeMaxDynamicSharedMemorySize, smemSize)); \ - TLLM_CUDA_CHECK(cudaLaunchKernelEx(&config, kernelTyped, params)); \ - } + cudaLaunchConfig_t config{}; \ + config.gridDim = numBlocks; \ + config.blockDim = numThreads; \ + config.dynamicSmemBytes = smemSize; \ + config.stream = (cudaStream_t) stream; \ + \ + cudaLaunchAttribute attributes[2] = {}; \ + attributes[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; \ + attributes[0].val.programmaticStreamSerializationAllowed = int(data.mUsePdl); \ + attributes[1].id = cudaLaunchAttributeCooperative; \ + attributes[1].val.cooperative = int(coopLaunch); \ + config.attrs = attributes; \ + config.numAttrs = 2; \ + if (data.mUsePdl) \ + { \ + auto params = KernelParams::setKernelParams(data); \ + auto kernelTyped = kernel>; \ + if (smemSize > 48 * 1024) \ + TLLM_CUDA_CHECK( \ + cudaFuncSetAttribute(kernelTyped, cudaFuncAttributeMaxDynamicSharedMemorySize, smemSize)); \ + TLLM_CUDA_CHECK(cudaLaunchKernelEx(&config, kernelTyped, params)); \ + } \ + else \ + { \ + auto params = KernelParams::setKernelParams(data); \ + auto kernelTyped = kernel>; \ + if (smemSize > 48 * 1024) \ + TLLM_CUDA_CHECK( \ + cudaFuncSetAttribute(kernelTyped, cudaFuncAttributeMaxDynamicSharedMemorySize, smemSize)); \ + TLLM_CUDA_CHECK(cudaLaunchKernelEx(&config, kernelTyped, params)); \ + } \ + } while (0) #define ADJUST_NUM_BLOCKS(data, kernel, type, numThreads) \ int ctasPerSM = 0; \ @@ -241,12 +246,14 @@ namespace moe::dev #define LAUNCH_ROUTING_LLAMA4(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream) \ if (data.mDtypeExpW == tg::Dtype::Fp32) \ { \ - LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, 128 /* Always 128 for llama4*/), kernel, numBlocks, \ + LAUNCH_TILEN(data, coopLaunch, \ + LAUNCH_ESC(float, float, 128 /* Always 128 for llama4*/, 1 /* Always 1 for llama4*/), kernel, numBlocks, \ numThreads, smemSize, stream); \ } \ else if (data.mDtypeExpW == tg::Dtype::Bfloat16) \ { \ - LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(__nv_bfloat16, __nv_bfloat16, 128 /* Always 128 for llama4*/), \ + LAUNCH_TILEN(data, coopLaunch, \ + LAUNCH_ESC(__nv_bfloat16, __nv_bfloat16, 128 /* Always 128 for llama4*/, 1 /* Always 1 for llama4*/), \ kernel, numBlocks, numThreads, smemSize, stream); \ } \ else \ @@ -294,26 +301,26 @@ namespace moe::dev //////////////////////////////////////////////////////////////////////////////////////////////////// #define LAUNCH_ROUTING_WITH_NUM_EXPERTS( \ - data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, numExperts) \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, numExperts, numTopExperts) \ if (data.mDtypeExpW == tg::Dtype::Fp32 && extraFlag1) \ { \ - LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, true), kernel, numBlocks, numThreads, \ - smemSize, stream); \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, numTopExperts, 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); \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, numTopExperts, false), kernel, numBlocks, \ + numThreads, smemSize, stream); \ } \ else if (data.mDtypeExpW == tg::Dtype::Bfloat16 && extraFlag1) \ { \ - LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(__nv_bfloat16, __nv_bfloat16, numExperts, true), kernel, numBlocks, \ - numThreads, smemSize, stream); \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(__nv_bfloat16, __nv_bfloat16, numExperts, numTopExperts, true), \ + kernel, numBlocks, numThreads, smemSize, stream); \ } \ else if (data.mDtypeExpW == tg::Dtype::Bfloat16) \ { \ - LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(__nv_bfloat16, __nv_bfloat16, numExperts, false), kernel, numBlocks, \ - numThreads, smemSize, stream); \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(__nv_bfloat16, __nv_bfloat16, numExperts, numTopExperts, false), \ + kernel, numBlocks, numThreads, smemSize, stream); \ } \ else \ { \ diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingDeepSeek.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingDeepSeek.cu index 6937a34ccd99..b59580f9f153 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingDeepSeek.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingDeepSeek.cu @@ -14,604 +14,23 @@ * limitations under the License. */ -#include "RoutingKernel.cuh" +#include "routingDeepSeek/RoutingDeepSeekCommon.cuh" namespace moe::dev::routing { - namespace routingDeepSeek { //////////////////////////////////////////////////////////////////////////////////////////////////// -static constexpr int NumNemotronExperts = 512; -static constexpr int NumKimiK2Experts = 384; -static constexpr int NumDeepseekExperts = 256; -static constexpr int MaxSupportedExpertCount = std::max({NumNemotronExperts, NumKimiK2Experts, NumDeepseekExperts}); -static constexpr int NumTopGroupScores = 2; -static constexpr int DefaultMaxNumTopExperts = 8; -static constexpr int MaxSupportedTopExperts = 22; -static constexpr int MaxNumTopGroups = 4; -static constexpr int MaxNumGroups = 8; - -template -__global__ void routingMainKernel(KernelParams params) -{ - // declare types - using OutputT = typename KernelParams::OutputT; - using InputT = typename KernelParams::InputT; - - // declare shared memory structure - // number of experts is bounded by number of threads - __shared__ float __attribute((aligned(128))) smemScoreSigmoid[KernelParams::MaxNumExperts]; - __shared__ float __attribute((aligned(128))) smemScoreBias[KernelParams::MaxNumExperts]; - // number of expert groups is bounded by number of warps - __shared__ float __attribute((aligned(128))) smemGroupScores[MaxNumGroups]; - - // needed for warp reduce - auto block = cg::this_thread_block(); - auto warp = cg::tiled_partition(block); - // for the final reduction of weight norm, only some lanes need to participate - int32_t laneIdx = threadIdx.x % WarpSize; - int32_t warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); - // warps outside the range of expert groups do not participate - if constexpr (KernelParams::UseGroups) - { - if (warpIdx >= params.mNumExpertGroups) - { - return; - } - } - - // note that for invalid scores, we simply use a negative value: - // they work well even with the compacted format used in topK, and - // sigmoid / bias activated scores cannot be negative - static constexpr float invalidScoreFloat = float{-INFINITY}; - const OutputT invalidScore = OutputT{invalidScoreFloat}; - - // load bias already; each warp represents one expert group - auto threadExpert = threadIdx.x; - bool expertSelected = threadExpert < params.mNumExperts; - if constexpr (KernelParams::UseGroups) - { - threadExpert = warpIdx * params.mNumExpertsPerGroup + laneIdx; - expertSelected = laneIdx < params.mNumExpertsPerGroup; - } - auto scoreIdx = int64_t{blockIdx.x} * int64_t{params.mNumExperts} + threadExpert; - auto biasVal = expertSelected ? params.mPtrRoutingBias[threadExpert] : invalidScore; - - // initialize the mPtrExpertCounts - if (params.mPtrExpertCounts) - { - int32_t globalThreadIdx = blockIdx.x * blockDim.x + threadIdx.x; - int32_t globalThreadStride = gridDim.x * blockDim.x; - int32_t expertCountsNum = 2 * params.mNumExperts; - initArr(globalThreadIdx, expertCountsNum, globalThreadStride, params.mPtrExpertCounts, 0); - } - -#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) - // trigger the secondary kernel when using PDL, then wait on primary - if constexpr (KernelParams::UsePdl) - { - cudaTriggerProgrammaticLaunchCompletion(); - cudaGridDependencySynchronize(); - } -#endif - - if (params.mPtrScores != nullptr) - { - // get our assigned thread score; each warp represents one expert group - float score = expertSelected ? static_cast(params.mPtrScores[scoreIdx]) : invalidScoreFloat; - // get the sigmoid score - // note that for invalid values, we simply use a negative value: - // sigmoig scores are always strictly positive - auto scoreSigmoid = sigmoid_accurate(score); - // write the sigmoid score to shared for later use - if (expertSelected) - { - smemScoreSigmoid[threadExpert] = scoreSigmoid; - } - // get the score with bias - // note that with invalid values, because sigmoid is < 1 and bias is -1, - // we must get a negative value, which is smaller than any valid value - auto scoreBias = float{scoreSigmoid + float{biasVal}}; - - if (expertSelected) - { - smemScoreBias[threadExpert] = scoreBias; - } - - // registers for top group score reduction - float topExpGroupScores[NumTopGroupScores]; - [[maybe_unused]] int32_t topExpGroupIdx[NumTopGroupScores]; - float topGroups[MaxNumTopGroups]; // bound of params.mNumLimitedGroups - int32_t topGroupIdx[MaxNumTopGroups]; - float expertScoreGroup[MaxNumTopGroups]; - int32_t expertIdxGroup[MaxNumTopGroups]; - float topScores[KernelParams::MaxNumTopExperts]; // bound of params.mTopK - int32_t topExperts[KernelParams::MaxNumTopExperts]; - - if constexpr (KernelParams::UseGroups) - { - topk::reduceTopK(warp, topExpGroupScores, topExpGroupIdx, scoreBias, threadExpert, - /* minValue */ invalidScoreFloat); - // get the final group score and write it to shared - if (cute::elect_one_sync()) - { - auto groupScore = topExpGroupScores[0] + topExpGroupScores[1]; - smemGroupScores[warpIdx] = groupScore; - } - } - - // make group scores available to all warps - __syncthreads(); - - auto localExpertExtent = params.mNumLocalExperts << params.mLocalExpertsStrideLog2; - if constexpr (KernelParams::UseGroups) - { // a single warp performs the selection of top groups, and goes on to select the final experts - if (warpIdx == 0) - { - float groupScore = laneIdx < params.mNumExpertGroups ? smemGroupScores[laneIdx] : invalidScoreFloat; - topk::reduceTopK(warp, topGroups, topGroupIdx, groupScore, laneIdx, - /* minValue */ invalidScoreFloat); - // final expert selection: get relevant indexes and scores from shared -#pragma unroll - for (int ii = 0; ii < MaxNumTopGroups; ++ii) - { // bound of params.mNumLimitedGroups - auto groupIdx = topGroupIdx[ii]; - expertIdxGroup[ii] = groupIdx * params.mNumExpertsPerGroup + laneIdx; - // note: expertSelected implies laneIdx < params.mNumExpertsPerGroup. - // we have params.mNumExpertsPerGroup == params.mNumExperts / params.mNumExpertGroups, - // thus groupIdx <= params.mNumExpertGroups - 1 => - // groupIdx * params.mNumExpertsPerGroup <= params.mNumExperts - params.mNumExpertsPerGroup - // => expertIdxGroup[ii] < params.mNumExperts <= NumThreads, - // so the access is safe here - expertScoreGroup[ii] - = (ii < params.mNumLimitedGroups) && (groupIdx < params.mNumExpertGroups) && expertSelected - ? smemScoreBias[expertIdxGroup[ii]] - : invalidScoreFloat; - } - - topk::reduceTopK(warp, topScores, topExperts, expertScoreGroup, expertIdxGroup, - /* minValue */ invalidScoreFloat, params.mTopK); - } - } - else if constexpr (KernelParams::MaxNumExperts > topk::MaxNumExpertsUnit) - { - // without groups, each thread just takes `MaxNumTopGroups` experts - int constexpr NumExpertWarps = (KernelParams::MaxNumExperts - 1) / topk::MaxNumExpertsUnit + 1; - int constexpr NumInterTopK = NumExpertWarps * KernelParams::MaxNumTopExperts; - __shared__ float __attribute((aligned(128))) smemInterTopScores[NumInterTopK]; - __shared__ int32_t __attribute((aligned(128))) smemInterTopExperts[NumInterTopK]; - if (warpIdx < NumExpertWarps) - { - int offset = warpIdx * WarpSize * MaxNumTopGroups; -#pragma unroll - for (int ii = 0; ii < MaxNumTopGroups; ++ii) - { - auto expertIdx = ii * WarpSize + laneIdx; - expertIdxGroup[ii] = offset + expertIdx; - expertScoreGroup[ii] = offset + expertIdx < params.mNumExperts ? smemScoreBias[offset + expertIdx] - : invalidScoreFloat; - } - topk::reduceTopK(warp, topScores, topExperts, expertScoreGroup, expertIdxGroup, - /* minValue */ invalidScoreFloat, params.mTopK); - - if (laneIdx < params.mTopK) - { - smemInterTopScores[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] = topScores[laneIdx]; - smemInterTopExperts[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] = topExperts[laneIdx]; - } - else if (laneIdx >= params.mTopK && laneIdx < KernelParams::MaxNumTopExperts) - { - smemInterTopScores[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] = invalidScoreFloat; - smemInterTopExperts[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] - = MaxSupportedExpertCount - 1; - } - } - __syncthreads(); - if (warpIdx == 0) - { - int constexpr NumInterTopKPerThread = (NumInterTopK - 1) / WarpSize + 1; - float intermidiateScore[NumInterTopKPerThread]; - int32_t intermidiateExpert[NumInterTopKPerThread]; - for (int i = laneIdx; i < NumInterTopKPerThread * WarpSize; i += WarpSize) - { - int ii = i / WarpSize; - if (i < NumInterTopK) - { - intermidiateScore[ii] = smemInterTopScores[i]; - intermidiateExpert[ii] = smemInterTopExperts[i]; - } - else - { - intermidiateScore[ii] = invalidScoreFloat; - intermidiateExpert[ii] = KernelParams::MaxNumExperts - 1; - } - } - topk::reduceTopK(warp, topScores, topExperts, intermidiateScore, intermidiateExpert, - /* minValue */ invalidScoreFloat, params.mTopK); - } - } - else - { - if (warpIdx == 0) - { - // without groups, each thread just takes `MaxNumTopGroups` experts -#pragma unroll - for (int ii = 0; ii < MaxNumTopGroups; ++ii) - { - auto expertIdx = ii * WarpSize + laneIdx; - expertIdxGroup[ii] = expertIdx; - expertScoreGroup[ii] - = expertIdx < params.mNumExperts ? smemScoreBias[expertIdx] : invalidScoreFloat; - } - topk::reduceTopK(warp, topScores, topExperts, expertScoreGroup, expertIdxGroup, - /* minValue */ invalidScoreFloat, params.mTopK); - } - } - - if (warpIdx == 0) - { - // determine our lane's expert index and write to output - int32_t expertIdx = 0; -#pragma unroll - for (int ii = 0; ii < params.mTopK; ++ii) - { // bound of params.mTopK - expertIdx = laneIdx == ii ? topExperts[ii] : expertIdx; - } - // determine whether our expert is local to this GPU - auto localExpertIdx = expertIdx - params.mLocalExpertsStartIdx; - auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; - - float scoreNorm = laneIdx < params.mTopK ? smemScoreSigmoid[expertIdx] : 0.F; - auto redNorm = cg::reduce(warp, scoreNorm, cg::plus{}); - auto finalScore = OutputT{scoreNorm * params.mRouteScale / redNorm}; - - // write expert idx out already - auto idxTopK = blockIdx.x * params.mTopK + laneIdx; - if (laneIdx < params.mTopK && params.mPtrTopKPacked != nullptr) - { - PackedScoreIdx packedScore{static_cast(finalScore), static_cast(expertIdx)}; - params.mPtrTopKPacked[idxTopK] = packedScore; - } - - if (laneIdx < params.mTopK && params.mPtrTopKWeights != nullptr && params.mPtrTopKIds == nullptr) - { - params.mPtrTopKWeights[idxTopK] = finalScore; - } - } - } -} - -//////////////////////////////////////////////////////////////////////////////////////////////////// - -template -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) -__global__ void __cluster_dims__(NumBlocksPerCluster, 1, 1) __launch_bounds__(KernelParams::MaxNumExperts) - routingIndicesClusterKernel(KernelParams params) -{ - using OutputT = typename KernelParams::OutputT; - - int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); - int32_t const clusterBlockRank = blockIdx.x; - - //@todo: try to move it into routingPermutation - // then wait on primary grid - if constexpr (KernelParams::UsePdl) - { - cudaGridDependencySynchronize(); - } - routingPermutation(params, nullptr, warpIdx, clusterBlockRank); -} -#else -__global__ void routingIndicesClusterKernel(KernelParams params) -{ - assert(false && "routingIndicesClusterKernel is only supported on SM90+ architectures"); -} -#endif - -//////////////////////////////////////////////////////////////////////////////////////////////////// - -template -#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesCoopKernel(KernelParams params) -{ - // number of experts is bounded by number of threads - int constexpr NumThreads = KernelParams::MaxNumExperts; - __shared__ int32_t __attribute((aligned(128))) smemExpertCount[NumThreads]; - __shared__ int32_t __attribute((aligned(128))) smemExpertOffset[NumThreads]; - // needed for the exclusive sum of token offsets - using Scan = cub::BlockScan; - __shared__ typename Scan::TempStorage tempStorage; - // 64 elements -> 128+ registers. Above that we may start to see spilling to local memory. - static constexpr int MaxExpandedIdxPerThread = 64; - - // Initialize grid. - cg::grid_group grid = cg::this_grid(); - // Note: the following is more efficient than grid.block_index() because we don't use y and z. - int32_t const gridBlockIdx = blockIdx.x; - int32_t const gridThreadIdx = NumThreads * gridBlockIdx + threadIdx.x; - int32_t const numBlocks = gridDim.x; - int32_t const numThreadsPerGrid = numBlocks * NumThreads; - - int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); - - auto expandedIdxSize = params.mNumTokens * params.mTopK; - - // pre-fill the counts with 0 - smemExpertCount[threadIdx.x] = 0; - __syncthreads(); - - // then wait on primary grid - if constexpr (KernelParams::UsePdl) - { - cudaGridDependencySynchronize(); - } - - // each thread keeps has some number of "expanded indexes" assigned to it - // for each of these, we keep the associated expert and offset within expert in registers - int32_t expertIndexes[MaxExpandedIdxPerThread]; - int32_t expertOffsets[MaxExpandedIdxPerThread]; - auto localExpertExtent = params.mNumLocalExperts << params.mLocalExpertsStrideLog2; - // In order to avoid a serialization LDG-ATOMS-LDG-ATOMS-..., we skip multiple iterations at a - // time, and branch between a fast path without bound checks and a slow path with bound checks. - int constexpr IterStride = 4; - static_assert(MaxExpandedIdxPerThread % IterStride == 0); - - // Define a lambda to avoid code duplication in both branches. - auto loopBody = [&](int ii, int expandedIdx) - { - int32_t expertIdx - = params.mPtrTopKIds != nullptr ? params.mPtrTopKIds[expandedIdx] : params.mPtrTopKPacked[expandedIdx].idx; - expertIndexes[ii] = expertIdx; - // check whether this expert is local to our GPU at all and ignore if not - auto localExpertIdx = expertIdx - params.mLocalExpertsStartIdx; - auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; - expertOffsets[ii] = isLocalExpert ? atomicAdd(smemExpertCount + expertIdx, 1) : 0; - }; - -#pragma unroll - for (int32_t ii0 = 0; ii0 < MaxExpandedIdxPerThread; ii0 += IterStride) - { - // Whether it's safe to do multiple iterations without bound checks. - bool const takeFastPath = (ii0 + IterStride) * numThreadsPerGrid <= expandedIdxSize; - if (takeFastPath) - { -#pragma unroll - for (int32_t jj = 0; jj < IterStride; jj++) - { - int const ii = ii0 + jj; - auto expandedIdx = static_cast(gridThreadIdx) + ii * numThreadsPerGrid; - loopBody(ii, expandedIdx); - } - } - else - { - bool doBreak = false; -#pragma unroll - for (int32_t jj = 0; jj < IterStride; jj++) - { - int const ii = ii0 + jj; - auto expandedIdx = static_cast(gridThreadIdx) + ii * numThreadsPerGrid; - if (expandedIdx >= expandedIdxSize) - { - doBreak = true; - break; - } - loopBody(ii, expandedIdx); - } - if (doBreak) - { - break; - } - } - } - - // Make histogram (token counts per expert) available to all threads in the block. - __syncthreads(); - - // - // Each thread now represents one expert - // - - // Add the local bin count to the common bin count and get a per-CTA offset. - int32_t const localExpertCount = smemExpertCount[threadIdx.x]; - - int32_t blockExpertOffset = 0; - if (threadIdx.x < params.mNumExperts) - { - blockExpertOffset = atomicAdd(¶ms.mPtrExpertCounts[threadIdx.x], localExpertCount); - } - - // Sync to wait for completion of the histogram reduction. - grid.sync(); - - // Get total count for this expert. - int32_t count = (threadIdx.x < params.mNumExperts) ? params.mPtrExpertCounts[threadIdx.x] : 0; - - // Note: the scan is redundant in all CTAs, but doing it in only 1 CTA would be worse for latency. - - // Compute the runtime config for projections - // Whether or not an expert is local is taken into account when smemExpertCount is computed - // so we do not need to take it into account here. - - int32_t numCta; - if constexpr (KernelParams::isPow2) - { - numCta = divUpLog2(count, params.mPaddingLog2); - } - else - { - numCta = divUpTileN(count, params.mTileTokensDim); - } - - int32_t ctaOffset; - int32_t numNonExitingCtas; - Scan(tempStorage).ExclusiveSum(numCta, ctaOffset, numNonExitingCtas); - - for (int32_t cta = gridBlockIdx; cta < numCta; cta += numBlocks) - { - const int32_t localExpertIdx = (threadIdx.x - 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) + count; - } - else - { - mnLimit1 = mulTileN(ctaOffset + cta + 1, params.mTileTokensDim); - mnLimit2 = mulTileN(ctaOffset, params.mTileTokensDim) + count; - } - params.mPtrCtaIdxXyToMnLimit[ctaOffset + cta] = min(mnLimit1, mnLimit2); - } - - // get the padded offset associated with this expert - int32_t offset; - if constexpr (KernelParams::isPow2) - { - offset = mulLog2(ctaOffset, params.mPaddingLog2); - } - else - { - offset = mulTileN(ctaOffset, params.mTileTokensDim); - } - int32_t permutedIdxSize; - if constexpr (KernelParams::isPow2) - { - permutedIdxSize = mulLog2(numNonExitingCtas, params.mPaddingLog2); - } - else - { - permutedIdxSize = mulTileN(numNonExitingCtas, params.mTileTokensDim); - } - - // write out padded count - if (gridBlockIdx == 0 && warpIdx == NumThreads / WarpSize - 1 && cute::elect_one_sync()) - { - params.mPtrPermutedIdxSize[0] = permutedIdxSize; - params.mPtrNumNonExitingCtas[0] = numNonExitingCtas; - } - - // write expert offsets to shared - smemExpertOffset[threadIdx.x] = offset + blockExpertOffset; - - // make expert offsets available to all threads - __syncthreads(); - - // trigger the secondary kernel when using PDL - // We can't do it earlier because FC1 depends on the mPtrCtaIdxXyToBatchIdx, - // mPtrCtaIdxXyToMnLimit, mPtrNumNonExitingCtas and mPtrTotalNumPaddedTokens - // TODO: this is not sufficient to ensure visibility in the next kernel! - if constexpr (KernelParams::UsePdl) - { - cudaTriggerProgrammaticLaunchCompletion(); - } - -// each thread has the same "expanded indexes" assigned to it as above -// at this point, we know the final offsets of experts and the offsets within -// experts, which allows writing the final index values -#pragma unroll - for (int32_t ii = 0; ii < MaxExpandedIdxPerThread; ++ii) - { - auto expandedIdx = static_cast(gridThreadIdx) + ii * numThreadsPerGrid; - if (expandedIdx >= expandedIdxSize) - { - break; - } - auto expertIdx = expertIndexes[ii]; - // check whether this expert is local to our GPU at all - auto localExpertIdx = static_cast(expertIdx) - params.mLocalExpertsStartIdx; - auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; - auto tokenIdx = expandedIdx / params.mTopK; - auto permutedIdx = isLocalExpert ? int32_t{smemExpertOffset[expertIdx]} + expertOffsets[ii] : 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; - } - } -} -#else -__global__ void routingIndicesCoopKernel(KernelParams params) -{ - assert(false && "routingIndicesCoopKernel is only supported on SM90+ architectures"); -} -#endif - -int constexpr getMaxNumExperts(int32_t numExperts) -{ - if (numExperts <= topk::MaxNumExpertsUnit) - { - return topk::MaxNumExpertsUnit; - } - else if (numExperts <= NumDeepseekExperts) - { - return NumDeepseekExperts; - } - else if (numExperts <= NumKimiK2Experts) - { - return NumKimiK2Experts; - } - else if (numExperts <= NumNemotronExperts) - { - return NumNemotronExperts; - } - else - { - TLLM_LOG_ERROR("Unsupported numExperts"); - return 0; - } -} +// Forward declarations for split-compiled launch wrappers. +void launchMainKernel(Data& data, int numBlocks, int numThreadsMain, void* stream); +void launchInitExpertCounts(Data& data, int numThreadsHist, void* stream); +void launchClusterKernel(Data& data, int numThreadsHist, void* stream); +void launchCoopKernel(Data& data, int numBlocksCoop, int numThreadsHist, void* stream); +void launchHistogramKernel(Data& data, int numBlocksHistogram, int numThreadsHist, void* stream); +void launchOffsetsKernel(Data& data, int numBlocksOffsets, int numThreadsHist, void* stream); //////////////////////////////////////////////////////////////////////////////////////////////////// -#define LAUNCH_ROUTING_DEEPSEEK( \ - data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, forceFloatInput) \ - if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ - { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, \ - stream, extraFlag1, forceFloatInput, topk::MaxNumExpertsUnit, DefaultMaxNumTopExperts); \ - } \ - else if (data.mNumExperts <= NumDeepseekExperts) \ - { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, \ - stream, extraFlag1, forceFloatInput, NumDeepseekExperts, DefaultMaxNumTopExperts); \ - } \ - else if (data.mNumExperts <= NumKimiK2Experts) \ - { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, \ - stream, extraFlag1, forceFloatInput, NumKimiK2Experts, DefaultMaxNumTopExperts); \ - } \ - else if (data.mNumExperts <= NumNemotronExperts) \ - { \ - if (data.mTopK <= DefaultMaxNumTopExperts) \ - { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, \ - smemSize, stream, extraFlag1, forceFloatInput, NumNemotronExperts, DefaultMaxNumTopExperts); \ - } \ - else if (data.mTopK <= MaxSupportedTopExperts) \ - { \ - LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, \ - smemSize, stream, extraFlag1, forceFloatInput, NumNemotronExperts, MaxSupportedTopExperts); \ - } \ - } \ - else \ - { \ - TLLM_LOG_ERROR("Unsupported numExperts"); \ - } void run(Data& data, void* stream) { @@ -693,35 +112,23 @@ void run(Data& data, void* stream) } int const numThreadsMain = max(data.mNumExpertGroups * WarpSize, getMaxNumExperts(data.mNumExperts)); - LAUNCH_ROUTING_DEEPSEEK(data, - /*coopLaunch=*/false, routingMainKernel, numBlocks, numThreadsMain, - /*smemSize=*/0, // No dynamic smem - stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); + launchMainKernel(data, numBlocks, numThreadsMain, stream); } else { // Reset the global histograms. - LAUNCH_ROUTING_DEEPSEEK(data, false, routingInitExpertCounts, (2 * data.mNumExperts - 1) / numThreadsHist + 1, - numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/false); + launchInitExpertCounts(data, numThreadsHist, stream); } if (data.mPtrPermutedIdxSize != nullptr) { if (useSingleCluster) { - LAUNCH_ROUTING_DEEPSEEK(data, - /*coopLaunch=*/false, routingIndicesClusterKernel, NumBlocksPerCluster, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); + launchClusterKernel(data, numThreadsHist, stream); } else if (data.mNumTokens <= maxTokensCoop) { - LAUNCH_ROUTING_DEEPSEEK(data, - /*coopLaunch=*/true, routingIndicesCoopKernel, numBlocksCoop, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); + launchCoopKernel(data, numBlocksCoop, numThreadsHist, stream); } else { @@ -737,14 +144,8 @@ void run(Data& data, void* stream) int const numBlocksOffsets = std::min((expandedIdxSize + offsetEltsPerBlock - 1) / offsetEltsPerBlock, maxNumBlocks); - LAUNCH_ROUTING_DEEPSEEK(data, - /*coopLaunch=*/false, routingIndicesHistogramKernel, numBlocksHistogram, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); - LAUNCH_ROUTING_DEEPSEEK(data, - /*coopLaunch=*/false, routingIndicesOffsetsKernel, numBlocksOffsets, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); + launchHistogramKernel(data, numBlocksHistogram, numThreadsHist, stream); + launchOffsetsKernel(data, numBlocksOffsets, numThreadsHist, stream); } } } diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh index 82e51ce3e342..4bc7b56aa18b 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.cuh @@ -174,20 +174,24 @@ template __device__ DataType calcSoftmax( cg::thread_block_tile const& warp, DataType score, int32_t laneIdx, int32_t NumTopExperts) { - DataType maxScore = DataType{-INFINITY}; + // Compute in float to support half/bfloat16 inputs safely. + // cg::reduce with cg::greater only supports float/double and integer types; + // using __nv_bfloat16 or __half directly can generate unsupported redux.sync.max instructions. + float maxScore = -INFINITY; if (laneIdx < NumTopExperts) { - maxScore = score >= maxScore ? score : maxScore; + float si = static_cast(score); + maxScore = si >= maxScore ? si : maxScore; } - maxScore = cg::reduce(warp, maxScore, cg::greater()); + maxScore = cg::reduce(warp, maxScore, cg::greater()); - float sumScore = float{0.f}; - float newScore; + float sumScore = 0.f; + float newScore = 0.f; // Get the summation of scores for each token if (laneIdx < NumTopExperts) { - newScore = static_cast(score) - static_cast(maxScore); - newScore = static_cast(exp(newScore)); + newScore = static_cast(score) - maxScore; + newScore = expf(newScore); sumScore += newScore; } sumScore = cg::reduce(warp, sumScore, cg::plus()); @@ -210,6 +214,12 @@ __device__ void routingPermutation(KernelParams params, PackedScoreIdx using OutputT = typename KernelParams::OutputT; using TypePacked = PackedScoreIdx; + // When MaxNumExperts > NumThreads, each thread handles multiple experts. + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + static constexpr int ExpertsPerThread = MaxNumExperts <= NumThreads ? 1 : MaxNumExperts / NumThreads; + static_assert(MaxNumExperts <= NumThreads || MaxNumExperts % NumThreads == 0, + "MaxNumExperts must be <= NumThreads or a multiple of NumThreads"); + static constexpr int MaxNumTokensSingleCluster = NumBlocksPerCluster * NumThreads; // Number of threads in the cluster. static constexpr int NumThreadsPerCluster = NumThreads * NumBlocksPerCluster; @@ -225,14 +235,19 @@ __device__ void routingPermutation(KernelParams params, PackedScoreIdx uint32_t const clusterThreadIdx = NumThreads * clusterBlockRank + threadIdx.x; auto expandedIdxSize = params.mNumTokens * params.mTopK; - // number of experts is bounded by number of threads - __shared__ int32_t __attribute((aligned(128))) smemExpertCount[NumThreads]; - __shared__ int32_t __attribute((aligned(128))) smemExpertOffset[NumThreads]; + // number of experts may exceed number of threads — size by MaxNumExperts + __shared__ int32_t __attribute((aligned(128))) smemExpertCount[MaxNumExperts]; + __shared__ int32_t __attribute((aligned(128))) smemExpertOffset[MaxNumExperts]; - // pre-fill the counts with 0 - if (threadIdx.x < params.mNumExperts) + // pre-fill the counts with 0 — each thread handles ExpertsPerThread experts +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - smemExpertCount[threadIdx.x] = 0; + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) + { + smemExpertCount[expert] = 0; + } } __syncthreads(); @@ -277,7 +292,7 @@ __device__ void routingPermutation(KernelParams params, PackedScoreIdx // check whether this expert is local to our GPU at all and ignore if not auto localExpertIdx = scoreIdx.idx - params.mLocalExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; expertOffsets[ii] = isLocalExpert ? atomicAdd(smemExpertCount + scoreIdx.idx, 1) : 0; if (params.mPtrTopKWeights != nullptr && params.mPtrTopKIds == nullptr) { @@ -327,34 +342,42 @@ __device__ void routingPermutation(KernelParams params, PackedScoreIdx __cluster_barrier_wait(); // - // Each thread now represents one expert + // Each thread now represents ExpertsPerThread experts // - // Total number of tokens for this expert. - int32_t count = 0; + // Total number of tokens for each expert this thread handles. + int32_t count[ExpertsPerThread]; // Per-expert offset for this block. - int32_t blockExpertOffset = 0; + int32_t blockExpertOffset[ExpertsPerThread]; - if (threadIdx.x < params.mNumExperts) - { - // Get the histogram bin from each rank for this expert. - int32_t expertCounts[NumBlocksPerCluster]; #pragma unroll - for (int rank = 0; rank < NumBlocksPerCluster; rank++) + for (int e = 0; e < ExpertsPerThread; e++) + { + int expert = threadIdx.x * ExpertsPerThread + e; + count[e] = 0; + blockExpertOffset[e] = 0; + + if (expert < params.mNumExperts) { - int32_t const* remoteSmem = cg::cluster_group::map_shared_rank(smemExpertCount, rank); - expertCounts[rank] = rank * NumWarps < params.mNumTokens ? remoteSmem[threadIdx.x] : 0; - } + // Get the histogram bin from each rank for this expert. + int32_t expertCounts[NumBlocksPerCluster]; +#pragma unroll + for (int rank = 0; rank < NumBlocksPerCluster; rank++) + { + int32_t const* remoteSmem = cg::cluster_group::map_shared_rank(smemExpertCount, rank); + expertCounts[rank] = rank * NumWarps < params.mNumTokens ? remoteSmem[expert] : 0; + } - // Compute an exclusive prefix sum of the block-local count. + // Compute an exclusive prefix sum of the block-local count. #pragma unroll - for (int rank = 0; rank < NumBlocksPerCluster; rank++) - { - if (rank == clusterBlockRank) + for (int rank = 0; rank < NumBlocksPerCluster; rank++) { - blockExpertOffset = count; + if (rank == clusterBlockRank) + { + blockExpertOffset[e] = count[e]; + } + count[e] += expertCounts[rank]; } - count += expertCounts[rank]; } } @@ -364,56 +387,65 @@ __device__ void routingPermutation(KernelParams params, PackedScoreIdx // Compute the runtime config for projections // Whether or not an expert is local is taken into account when smemExpertCount is computed // so we do not need to take it into account here. - int32_t numCta; - if constexpr (KernelParams::isPow2) - { - numCta = divUpLog2(count, params.mPaddingLog2); - } - else + int32_t numCta[ExpertsPerThread]; +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - numCta = divUpTileN(count, params.mTileTokensDim); + if constexpr (KernelParams::isPow2) + { + numCta[e] = divUpLog2(count[e], params.mPaddingLog2); + } + else + { + numCta[e] = divUpTileN(count[e], params.mTileTokensDim); + } } - int32_t ctaOffset; + int32_t ctaOffset[ExpertsPerThread]; int32_t numNonExitingCtas; Scan(tempStorage).ExclusiveSum(numCta, ctaOffset, numNonExitingCtas); - if (threadIdx.x < params.mNumExperts) +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - // Strided loop to share this work between blocks. - for (int32_t cta = clusterBlockRank; cta < numCta; cta += NumBlocksPerCluster) + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) { - const int32_t localExpertIdx - = (threadIdx.x - params.mLocalExpertsStartIdx) >> params.mLocalExpertsStrideLog2; - params.mPtrCtaIdxXyToBatchIdx[ctaOffset + cta] = localExpertIdx; - int32_t mnLimit1; - int32_t mnLimit2; + // Strided loop to share this work between blocks. + for (int32_t cta = clusterBlockRank; cta < numCta[e]; cta += NumBlocksPerCluster) + { + const int32_t localExpertIdx + = (expert - params.mLocalExpertsStartIdx) >> params.mLocalExpertsStrideLog2; + params.mPtrCtaIdxXyToBatchIdx[ctaOffset[e] + cta] = localExpertIdx; + int32_t mnLimit1; + int32_t mnLimit2; + if constexpr (KernelParams::isPow2) + { + mnLimit1 = mulLog2(ctaOffset[e] + cta + 1, params.mPaddingLog2); + mnLimit2 = mulLog2(ctaOffset[e], params.mPaddingLog2) + count[e]; + } + else + { + mnLimit1 = mulTileN(ctaOffset[e] + cta + 1, params.mTileTokensDim); + mnLimit2 = mulTileN(ctaOffset[e], params.mTileTokensDim) + count[e]; + } + params.mPtrCtaIdxXyToMnLimit[ctaOffset[e] + cta] = min(mnLimit1, mnLimit2); + } + + // get the padded offset associated with this expert + int32_t offset; if constexpr (KernelParams::isPow2) { - mnLimit1 = mulLog2(ctaOffset + cta + 1, params.mPaddingLog2); - mnLimit2 = mulLog2(ctaOffset, params.mPaddingLog2) + count; + offset = mulLog2(ctaOffset[e], params.mPaddingLog2); } else { - mnLimit1 = mulTileN(ctaOffset + cta + 1, params.mTileTokensDim); - mnLimit2 = mulTileN(ctaOffset, params.mTileTokensDim) + count; + offset = mulTileN(ctaOffset[e], params.mTileTokensDim); } - params.mPtrCtaIdxXyToMnLimit[ctaOffset + cta] = min(mnLimit1, mnLimit2); - } - // get the padded offset associated with this expert - int32_t offset; - if constexpr (KernelParams::isPow2) - { - offset = mulLog2(ctaOffset, params.mPaddingLog2); + // write expert offsets to shared + smemExpertOffset[expert] = offset + blockExpertOffset[e]; } - else - { - offset = mulTileN(ctaOffset, params.mTileTokensDim); - } - - // write expert offsets to shared - smemExpertOffset[threadIdx.x] = offset + blockExpertOffset; } // write out padded count @@ -467,7 +499,7 @@ __device__ void routingPermutation(KernelParams params, PackedScoreIdx // check whether this expert is local to our GPU at all auto localExpertIdx = static_cast(expertIdx) - params.mLocalExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; auto tokenIdx = expandedIdx / params.mTopK; auto permutedIdx = isLocalExpert ? int32_t{smemExpertOffset[expertIdx]} + expertOffsets[ii] : int32_t{-1}; if (params.mPtrExpandedIdxToPermutedIdx != nullptr) @@ -495,20 +527,31 @@ __device__ void routingPermutation(KernelParams params, PackedScoreIdx // Note: the histogram calculation could also be fused with routingMainKernel, but this might be // inefficient if we have one CTA per token doing a single global atomic. template -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHistogramKernel(KernelParams params) +__global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelParams::MaxNumExperts : 1024) + routingIndicesHistogramKernel(KernelParams params) { using OutputT = typename KernelParams::OutputT; + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + // Cap actual thread count at 1024 when MaxNumExperts > 1024. + static constexpr int NumThreadsBlock = MaxNumExperts <= 1024 ? MaxNumExperts : 1024; + static constexpr int ExpertsPerThread = MaxNumExperts / NumThreadsBlock; + static_assert(MaxNumExperts % NumThreadsBlock == 0, "MaxNumExperts must be a multiple of NumThreadsBlock"); - // number of experts is bounded by number of threads - __shared__ int32_t __attribute((aligned(128))) smemExpertCount[KernelParams::MaxNumExperts]; + // number of experts is bounded by MaxNumExperts (may exceed thread count) + __shared__ int32_t __attribute((aligned(128))) smemExpertCount[MaxNumExperts]; // For unrolling. uint32_t constexpr NumEltsPerThread = 8; - // Pre-fill the counts with 0 - if (threadIdx.x < params.mNumExperts) + // Pre-fill the counts with 0 — each thread handles ExpertsPerThread experts +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - smemExpertCount[threadIdx.x] = 0; + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) + { + smemExpertCount[expert] = 0; + } } __syncthreads(); @@ -524,8 +567,9 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHis uint32_t const expandedIdxSize = params.mNumTokens * params.mTopK; uint32_t const localExpertExtent = params.mNumLocalExperts << params.mLocalExpertsStrideLog2; - uint32_t const gridBlockOffset = blockIdx.x * KernelParams::MaxNumExperts; - uint32_t const gridStride = gridDim.x * KernelParams::MaxNumExperts; + // Use NumThreadsBlock (actual thread count) for grid-stride addressing + uint32_t const gridBlockOffset = blockIdx.x * NumThreadsBlock; + uint32_t const gridStride = gridDim.x * NumThreadsBlock; // Define a lambda to avoid code duplication in branches. auto loopBody = [&](int expandedIdx) @@ -549,31 +593,31 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHis // check whether this expert is local to our GPU at all and ignore if not auto localExpertIdx = idx - params.mLocalExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; if (isLocalExpert) { atomicAdd(&smemExpertCount[idx], 1); } }; - // Grid-stride loop. + // Grid-stride loop using NumThreadsBlock as block width. for (uint32_t expandedIdx0 = gridBlockOffset * NumEltsPerThread; expandedIdx0 < expandedIdxSize; expandedIdx0 += gridStride * NumEltsPerThread) { // Fast path if bound checks aren't necessary - if (expandedIdx0 + NumEltsPerThread * KernelParams::MaxNumExperts <= expandedIdxSize) + if (expandedIdx0 + NumEltsPerThread * NumThreadsBlock <= expandedIdxSize) { #pragma unroll for (uint32_t ii = 0; ii < NumEltsPerThread; ii++) { - uint32_t expandedIdx = expandedIdx0 + ii * KernelParams::MaxNumExperts + threadIdx.x; + uint32_t expandedIdx = expandedIdx0 + ii * NumThreadsBlock + threadIdx.x; loopBody(expandedIdx); } } else { for (uint32_t expandedIdx = expandedIdx0 + threadIdx.x; expandedIdx < expandedIdxSize; - expandedIdx += KernelParams::MaxNumExperts) + expandedIdx += NumThreadsBlock) { loopBody(expandedIdx); } @@ -582,33 +626,45 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHis __syncthreads(); // - // Each thread now represents one expert + // Each thread now represents ExpertsPerThread experts // // Reduce histograms with atomics. - if (threadIdx.x < params.mNumExperts) +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - int32_t const localExpertCount = smemExpertCount[threadIdx.x]; - atomicAdd(¶ms.mPtrExpertCounts[threadIdx.x], localExpertCount); + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) + { + int32_t const localExpertCount = smemExpertCount[expert]; + atomicAdd(¶ms.mPtrExpertCounts[expert], localExpertCount); + } } } //////////////////////////////////////////////////////////////////////////////////////////////////// template -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOffsetsKernel(KernelParams params) +__global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelParams::MaxNumExperts : 1024) + routingIndicesOffsetsKernel(KernelParams params) { using OutputT = typename KernelParams::OutputT; - - // number of experts is bounded by number of threads - __shared__ int32_t __attribute((aligned(128))) smemExpertOffset[KernelParams::MaxNumExperts]; - __shared__ int32_t __attribute((aligned(128))) smemExpertCount[KernelParams::MaxNumExperts]; - __shared__ int32_t __attribute((aligned(128))) smemExpertTileOffset[KernelParams::MaxNumExperts]; - // needed for the exclusive sum of token offsets - using Scan = cub::BlockScan; + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + // Cap actual thread count at 1024 when MaxNumExperts > 1024. + static constexpr int NumThreadsBlock = MaxNumExperts <= 1024 ? MaxNumExperts : 1024; + static constexpr int ExpertsPerThread = MaxNumExperts / NumThreadsBlock; + static_assert(MaxNumExperts % NumThreadsBlock == 0, "MaxNumExperts must be a multiple of NumThreadsBlock"); + + // number of experts — shared memory sized by MaxNumExperts (may exceed thread count) + __shared__ int32_t __attribute((aligned(128))) smemExpertOffset[MaxNumExperts]; + __shared__ int32_t __attribute((aligned(128))) smemExpertCount[MaxNumExperts]; + __shared__ int32_t __attribute((aligned(128))) smemExpertTileOffset[MaxNumExperts]; + // BlockScan uses actual thread count; array overload handles ExpertsPerThread items per thread + using Scan = cub::BlockScan; __shared__ typename Scan::TempStorage tempStorage; static constexpr int MaxExpandedIdxPerThread = NumEltsPerOffsetTilePerThread; - static constexpr int MaxExpandedIdxPerBlock = KernelParams::MaxNumExperts * MaxExpandedIdxPerThread; + // Tile size uses actual thread count + static constexpr int MaxExpandedIdxPerBlock = NumThreadsBlock * MaxExpandedIdxPerThread; int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); @@ -629,50 +685,65 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff // the scan, with PDL? // - // Each thread represents one expert. + // Each thread represents ExpertsPerThread experts. // - // Get total count for this expert. - int32_t count = (threadIdx.x < params.mNumExperts) ? params.mPtrExpertCounts[threadIdx.x] : 0; + // Get total count for each expert this thread handles. + int32_t count[ExpertsPerThread]; +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) + { + int expert = threadIdx.x * ExpertsPerThread + e; + count[e] = (expert < params.mNumExperts) ? params.mPtrExpertCounts[expert] : 0; + } // Compute the runtime config for projections // Whether or not an expert is local is taken into account when the histogram is computed // so we do not need to take it into account here. - int32_t numCta; - if constexpr (KernelParams::isPow2) - { - numCta = divUpLog2(count, params.mPaddingLog2); - } - else - { - numCta = divUpTileN(count, params.mTileTokensDim); - } - int32_t ctaOffset; - int32_t numNonExitingCtas; - Scan(tempStorage).ExclusiveSum(numCta, ctaOffset, numNonExitingCtas); - - if (threadIdx.x < params.mNumExperts) + int32_t numCta[ExpertsPerThread]; +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - // Get the padded offset associated with this expert - int32_t offset; if constexpr (KernelParams::isPow2) { - offset = mulLog2(ctaOffset, params.mPaddingLog2); + numCta[e] = divUpLog2(count[e], params.mPaddingLog2); } else { - offset = mulTileN(ctaOffset, params.mTileTokensDim); + numCta[e] = divUpTileN(count[e], params.mTileTokensDim); } + } + int32_t ctaOffset[ExpertsPerThread]; + int32_t numNonExitingCtas; + Scan(tempStorage).ExclusiveSum(numCta, ctaOffset, numNonExitingCtas); - // Write expert offsets to shared - smemExpertOffset[threadIdx.x] = offset; +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) + { + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) + { + // Get the padded offset associated with this expert + int32_t offset; + if constexpr (KernelParams::isPow2) + { + offset = mulLog2(ctaOffset[e], params.mPaddingLog2); + } + else + { + offset = mulTileN(ctaOffset[e], params.mTileTokensDim); + } + + // Write expert offsets to shared + smemExpertOffset[expert] = offset; + } } // Sync to make expert offsets available to all threads. __syncthreads(); - // The first block writes out padded count - if (blockIdx.x == 0 && warpIdx == KernelParams::MaxNumExperts / WarpSize - 1 && cute::elect_one_sync()) + // The first block writes out padded count (use last warp of actual thread count) + if (blockIdx.x == 0 && warpIdx == NumThreadsBlock / WarpSize - 1 && cute::elect_one_sync()) { int32_t permutedIdxSize; if constexpr (KernelParams::isPow2) @@ -687,27 +758,32 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff params.mPtrNumNonExitingCtas[0] = numNonExitingCtas; } - if (threadIdx.x < params.mNumExperts) +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - // Strided loop to share this work between blocks. - for (int32_t cta = blockIdx.x; cta < numCta; cta += gridDim.x) + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) { - const int32_t localExpertIdx - = (threadIdx.x - 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) + count; - } - else + // Strided loop to share this work between blocks. + for (int32_t cta = blockIdx.x; cta < numCta[e]; cta += gridDim.x) { - mnLimit1 = mulTileN(ctaOffset + cta + 1, params.mTileTokensDim); - mnLimit2 = mulTileN(ctaOffset, params.mTileTokensDim) + count; + const int32_t localExpertIdx + = (expert - params.mLocalExpertsStartIdx) >> params.mLocalExpertsStrideLog2; + params.mPtrCtaIdxXyToBatchIdx[ctaOffset[e] + cta] = localExpertIdx; + int32_t mnLimit1; + int32_t mnLimit2; + if constexpr (KernelParams::isPow2) + { + mnLimit1 = mulLog2(ctaOffset[e] + cta + 1, params.mPaddingLog2); + mnLimit2 = mulLog2(ctaOffset[e], params.mPaddingLog2) + count[e]; + } + else + { + mnLimit1 = mulTileN(ctaOffset[e] + cta + 1, params.mTileTokensDim); + mnLimit2 = mulTileN(ctaOffset[e], params.mTileTokensDim) + count[e]; + } + params.mPtrCtaIdxXyToMnLimit[ctaOffset[e] + cta] = min(mnLimit1, mnLimit2); } - params.mPtrCtaIdxXyToMnLimit[ctaOffset + cta] = min(mnLimit1, mnLimit2); } } @@ -724,10 +800,15 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff __syncthreads(); } - // Pre-fill the counts with 0 - if (threadIdx.x < params.mNumExperts) + // Pre-fill the counts with 0 — each thread handles ExpertsPerThread experts +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - smemExpertCount[threadIdx.x] = 0; + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) + { + smemExpertCount[expert] = 0; + } } __syncthreads(); @@ -745,7 +826,7 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff // check whether this expert is local to our GPU at all and ignore if not auto localExpertIdx = expertIndexes[ii] - params.mLocalExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; expertOffsets[ii] = isLocalExpert ? atomicAdd(smemExpertCount + expertIndexes[ii], 1) : 0; }; @@ -755,7 +836,7 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff #pragma unroll for (int32_t ii = 0; ii < MaxExpandedIdxPerThread; ii += 1) { - auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * KernelParams::MaxNumExperts + threadIdx.x; + auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * NumThreadsBlock + threadIdx.x; loopBody(ii, expandedIdx); } } @@ -772,16 +853,14 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff { // Whether it's safe to do multiple iterations without bound checks. bool const takeFastPath - = tileIdx * MaxExpandedIdxPerBlock + (ii0 + IterStride) * KernelParams::MaxNumExperts - <= expandedIdxSize; + = tileIdx * MaxExpandedIdxPerBlock + (ii0 + IterStride) * NumThreadsBlock <= expandedIdxSize; if (takeFastPath) { #pragma unroll for (int32_t jj = 0; jj < IterStride; jj++) { int const ii = ii0 + jj; - auto expandedIdx - = tileIdx * MaxExpandedIdxPerBlock + ii * KernelParams::MaxNumExperts + threadIdx.x; + auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * NumThreadsBlock + threadIdx.x; loopBody(ii, expandedIdx); } } @@ -792,8 +871,7 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff for (int32_t jj = 0; jj < IterStride; jj++) { int const ii = ii0 + jj; - auto expandedIdx - = tileIdx * MaxExpandedIdxPerBlock + ii * KernelParams::MaxNumExperts + threadIdx.x; + auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * NumThreadsBlock + threadIdx.x; if (expandedIdx >= expandedIdxSize) { doBreak = true; @@ -813,20 +891,25 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff __syncthreads(); // - // Each thread now represents one expert + // Each thread now represents ExpertsPerThread experts // - if (threadIdx.x < params.mNumExperts) +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) { - // Add the local bin count to the common bin count and get a per-CTA offset. We use the second - // half of the histogram buffer for this histogram, because the first half already holds the - // reduced histogram from the previous kernel. - int32_t const localExpertCount = smemExpertCount[threadIdx.x]; - int32_t const tileExpertOffset - = atomicAdd(¶ms.mPtrExpertCounts[params.mNumExperts + threadIdx.x], localExpertCount); - - // Make per-expert tile offsets available to all threads in the block. - smemExpertTileOffset[threadIdx.x] = tileExpertOffset + smemExpertOffset[threadIdx.x]; + int expert = threadIdx.x * ExpertsPerThread + e; + if (expert < params.mNumExperts) + { + // Add the local bin count to the common bin count and get a per-CTA offset. We use the second + // half of the histogram buffer for this histogram, because the first half already holds the + // reduced histogram from the previous kernel. + int32_t const localExpertCount = smemExpertCount[expert]; + int32_t const tileExpertOffset + = atomicAdd(¶ms.mPtrExpertCounts[params.mNumExperts + expert], localExpertCount); + + // Make per-expert tile offsets available to all threads in the block. + smemExpertTileOffset[expert] = tileExpertOffset + smemExpertOffset[expert]; + } } __syncthreads(); @@ -837,7 +920,7 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff // check whether this expert is local to our GPU at all auto localExpertIdx = static_cast(expertIdx) - params.mLocalExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; auto tokenIdx = expandedIdx / params.mTopK; auto permutedIdx = isLocalExpert ? (expertOffsets[ii] + smemExpertTileOffset[expertIdx]) : int32_t{-1}; if (params.mPtrExpandedIdxToPermutedIdx != nullptr) @@ -859,7 +942,7 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff #pragma unroll for (int32_t ii = 0; ii < MaxExpandedIdxPerThread; ii += 1) { - auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * KernelParams::MaxNumExperts + threadIdx.x; + auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * NumThreadsBlock + threadIdx.x; storeLoopBody(ii, expandedIdx); } } @@ -868,7 +951,7 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff #pragma unroll for (int32_t ii = 0; ii < MaxExpandedIdxPerThread; ii += 1) { - auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * KernelParams::MaxNumExperts + threadIdx.x; + auto expandedIdx = tileIdx * MaxExpandedIdxPerBlock + ii * NumThreadsBlock + threadIdx.x; if (expandedIdx >= expandedIdxSize) { break; @@ -892,12 +975,16 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesOff //////////////////////////////////////////////////////////////////////////////////////////////////// template -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingInitExpertCounts(KernelParams params) +__global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelParams::MaxNumExperts : 1024) + routingInitExpertCounts(KernelParams params) { + // Cap actual thread count at 1024 when MaxNumExperts > 1024. + static constexpr int NumThreadsBlock = KernelParams::MaxNumExperts <= 1024 ? KernelParams::MaxNumExperts : 1024; + // initialize the mPtrExpertCounts int32_t expertCountsNum = 2 * params.mNumExperts; - int32_t globalThreadIdx = blockIdx.x * KernelParams::MaxNumExperts + threadIdx.x; - int32_t globalThreadStride = gridDim.x * KernelParams::MaxNumExperts; + int32_t globalThreadIdx = blockIdx.x * NumThreadsBlock + threadIdx.x; + int32_t globalThreadStride = gridDim.x * NumThreadsBlock; #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) // Wait on primary grid. diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h index 888e04f2541a..3daa1848e5d3 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h @@ -109,12 +109,13 @@ struct DataBase int32_t mNumLocalExperts; }; -template +template struct KernelParamsBase { using InputT = InputT_; using OutputT = OutputT_; static constexpr int MaxNumExperts = MaxNumExperts_; + static constexpr int MaxNumTopExperts = MaxNumTopExperts_; static constexpr bool UsePdl = UsePdl_; static constexpr bool isPow2 = isPow2_; @@ -191,13 +192,12 @@ struct Data : public DataBase template -struct KernelParams : public KernelParamsBase +struct KernelParams : public KernelParamsBase { using InputT = InputT_; using OutputT = OutputT_; static constexpr bool UseGroups = UseGroups_; - static constexpr int MaxNumTopExperts = MaxNumTopExperts_; PackedScoreIdx* mPtrTopKPacked = nullptr; @@ -250,8 +250,8 @@ struct Data : public DataBase tg::Dtype mDtypeExpW{tg::Dtype::Bfloat16}; }; -template -struct KernelParams : public KernelParamsBase +template +struct KernelParams : public KernelParamsBase { using InputT = InputT_; using OutputT = OutputT_; @@ -296,9 +296,9 @@ struct Data : public DataBase bool mApplySoftmaxAfterTopK{true}; }; -template -struct KernelParams : public KernelParamsBase +template +struct KernelParams : public KernelParamsBase { using InputT = InputT_; using OutputT = OutputT_; diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernelTopK.cuh b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernelTopK.cuh index 7eab1c82a117..90cb76c507e1 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernelTopK.cuh +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernelTopK.cuh @@ -35,7 +35,7 @@ namespace cg = cooperative_groups; static constexpr int WarpSize = 32; static constexpr int MaxNumExpertsUnit = 128; -static constexpr int MaxSupportedTopExperts = 22; +static constexpr int MaxSupportedTopExperts = 32; //////////////////////////////////////////////////////////////////////////////////////////////////// @@ -114,10 +114,89 @@ struct TopKRedType topK[J].compVal = pairMin; \ } -//////////////////////////////////////////////////////////////////////////////////////////////////// +// Helper to check if N is a power of 2 +template +struct IsPowerOf2 +{ + static constexpr bool value = (N > 0) && ((N & (N - 1)) == 0); +}; +//////////////////////////////////////////////////////////////////////////////////////////////////// template -struct Sort; +struct Sort +{ + static_assert(N > 0 && N <= 64, "Sort only supports N in range [1, 64]"); + + static __device__ void run(RedType* topK) + { + if constexpr (IsPowerOf2::value) + { +// Bitonic sort for power-of-2 sizes - more efficient +#pragma unroll + for (int k = 2; k <= N; k *= 2) + { +#pragma unroll + for (int j = k / 2; j > 0; j /= 2) + { +#pragma unroll + for (int i = 0; i < N; ++i) + { + int ixj = i ^ j; + if (ixj > i) + { + if ((i & k) == 0) + { + if (topK[i].compVal < topK[ixj].compVal) + { + auto tmp = topK[i].compVal; + topK[i].compVal = topK[ixj].compVal; + topK[ixj].compVal = tmp; + } + } + else + { + if (topK[i].compVal > topK[ixj].compVal) + { + auto tmp = topK[i].compVal; + topK[i].compVal = topK[ixj].compVal; + topK[ixj].compVal = tmp; + } + } + } + } + } + } + } + else + { +// Odd-even transposition sort for non-power-of-2 sizes +#pragma unroll + for (int pass = 0; pass < N; ++pass) + { +#pragma unroll + for (int i = 0; i < N - 1; i += 2) + { + if (topK[i].compVal < topK[i + 1].compVal) + { + auto tmp = topK[i].compVal; + topK[i].compVal = topK[i + 1].compVal; + topK[i + 1].compVal = tmp; + } + } +#pragma unroll + for (int i = 1; i < N - 1; i += 2) + { + if (topK[i].compVal < topK[i + 1].compVal) + { + auto tmp = topK[i].compVal; + topK[i].compVal = topK[i + 1].compVal; + topK[i + 1].compVal = tmp; + } + } + } + } + } +}; template struct Sort<1, RedType> @@ -180,13 +259,13 @@ __forceinline__ __device__ void reduceTopK(cg::thread_block_tile const }; template -__forceinline__ __device__ void reduceTopKFunc(cg::thread_block_tile const& warp, Type (&out)[K], +__forceinline__ __device__ void reduceTopK(cg::thread_block_tile const& warp, Type (&out)[K], int32_t (&outIdx)[K], Type (&value)[N], int32_t (&idx)[N], Type const minValue, int actualK = K) { static_assert(K > 0, "Top K must have K > 0"); - static_assert(K < WarpSize, "Top K must have K < WarpSize"); + static_assert(K <= WarpSize, "Top K must have K <= WarpSize"); static_assert(N > 0, "Top K must have N > 0"); - static_assert(N < 5, "Only support candidates number less than or equal to 128"); + static_assert(N <= 64, "Only support candidates number less than or equal to 64*32=2048"); using RedType = TopKRedType; RedType topK[N]; #pragma unroll @@ -198,8 +277,7 @@ __forceinline__ __device__ void reduceTopKFunc(cg::thread_block_tile c Sort::run(topK); typename RedType::TypeCmp packedMax{}; -#pragma unroll - for (int kk = 0; kk < actualK; ++kk) //@todo: check if actualK is correct + for (int kk = 0; kk < actualK; ++kk) { bool update = kk > 0 && packedMax == topK[0].compVal; #pragma unroll @@ -213,67 +291,6 @@ __forceinline__ __device__ void reduceTopKFunc(cg::thread_block_tile c } }; -template -__forceinline__ __device__ void reduceTopK(cg::thread_block_tile const& warp, Type (&out)[K], - int32_t (&outIdx)[K], Type (&value)[N], int32_t (&idx)[N], Type const minValue, int actualK = K) -{ - static_assert(K > 0, "Top K must have K > 0"); - static_assert(K < WarpSize, "Top K must have K < WarpSize"); - static_assert(N > 0, "Top K must have N > 0"); - static_assert(N <= 16, "Only support candidates number less than or equal to 16*32=512"); - static_assert( - N <= 4 || N % 4 == 0, "Only support candidates number is a multiple of 4*32=128 or less than or equal to 4"); - using RedType = TopKRedType; - - if constexpr (N <= 4) - { - reduceTopKFunc(warp, out, outIdx, value, idx, minValue, actualK); - } - else - { - - constexpr int numLoops = (N - 1) / 4 + 1; - constexpr int numResults = (numLoops * K - 1) / WarpSize + 1; - - Type topKBufferValue[numResults]; - int32_t topKBufferIdx[numResults]; - int32_t laneIdx = threadIdx.x % WarpSize; - - for (int ii = 0; ii < numResults; ++ii) - { - topKBufferValue[ii] = minValue; - topKBufferIdx[ii] = ii * WarpSize - 1; //@todo: check if this is correct - } - for (int loop = 0; loop < numLoops; ++loop) - { - int start = loop * 4; - Type topKValue[K]; - int32_t topKIdx[K]; - Type inValue[4]; - int32_t inIdx[4]; - for (int i = 0; i < 4; ++i) - { - inValue[i] = value[start + i]; - inIdx[i] = idx[start + i]; - } - reduceTopKFunc(warp, topKValue, topKIdx, inValue, inIdx, minValue, actualK); - int inOffset = laneIdx % K; - if (laneIdx >= loop * K && laneIdx < (loop + 1) * K) - { - topKBufferValue[0] = topKValue[inOffset]; - topKBufferIdx[0] = topKIdx[inOffset]; - } - if (loop == numLoops - 1 && (laneIdx < (numLoops * K - WarpSize))) - { - topKBufferValue[1] = topKValue[inOffset]; - topKBufferIdx[1] = topKIdx[inOffset]; - } - } - - reduceTopKFunc(warp, out, outIdx, topKBufferValue, topKBufferIdx, minValue, actualK); - } -}; - #undef TOPK_SWAP } // namespace topk } // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingLlama4.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingLlama4.cu index be4e0e49372b..3362eb80c1b6 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingLlama4.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingLlama4.cu @@ -27,7 +27,7 @@ static constexpr int NumThreads = 1024; static constexpr int NumWarps = NumThreads / WarpSize; static constexpr int MaxNumTopExperts = 1; // static constexpr int MaxNumExperts = 128; -static constexpr int NumExpertsLimit = 128; +static constexpr int MaxSupportedExperts = 128; static constexpr int MaxNumTokensSingleCluster = NumBlocksPerCluster * NumThreads; static constexpr int MaxNumTokensSingleClusterScores = NumBlocksPerCluster * NumWarps; static constexpr int WarpKernelSmemStride = 33; @@ -339,7 +339,7 @@ __global__ void __launch_bounds__(WarpSize) routingIndicesWarpKernel(KernelParam auto expertIdx = threadIdx.x * ExpertsPerThread + ii; auto localExpertIdx = static_cast(expertIdx) - params.mLocalExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; // the permuted index: we add the local offset relative to this expert and token // to the global offset from the scan for this expert auto permutedIdx = isLocalExpert ? finalExpertOffset[ii] + localOffsetToken : int32_t{-1}; @@ -556,8 +556,8 @@ void run(Data const& data, void* stream) "Llama4 routing kernel expects permuted idx and grouped Gemm launch config buffers"); TLLM_CHECK_WITH_INFO(data.mTopK <= MaxNumTopExperts, "Routing kernel expects topK experts <= %d, got %d", MaxNumTopExperts, 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 <= MaxSupportedExperts, + "Routing kernel expects #experts %d to be no more than %d", data.mNumExperts, MaxSupportedExperts); // static_assert(MaxNumExperts <= NumThreads, "#experts must be bounded by #threads"); // static_assert(MaxNumExperts <= numThreadsHist, "#experts must be bounded by #threads"); TLLM_CHECK_WITH_INFO( diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu index 67b6913aaf7a..674758a3b4d3 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingRenormalize.cu @@ -13,465 +13,21 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -#include "RoutingKernel.cuh" +#include "routingRenormalize/RoutingRenormalizeCommon.cuh" namespace moe::dev::routing { namespace routingRenormalize { -//////////////////////////////////////////////////////////////////////////////////////////////////// - -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], 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, bool const normTopkProb, bool const applySoftmaxAfterTopK = true) -{ - DataType minScore = DataType{-INFINITY}; - - for (int i = 0; i < VecSize; i++) - { - auto expertIdx = i * WarpSize + laneIdx; - auto newScore = expertIdx < numExperts ? static_cast(ptrScores[expertIdx]) : minScore; - score[i] = newScore; - idx[i] = expertIdx; - } - if constexpr (DoSoftmaxBeforeTopK) - { - calcSoftmax(warp, score); - } - - // Get the top-k scores and their corresponding expert indices - topk::reduceTopK(warp, warpTopKScore, warpTopKExpertIdx, score, idx, minScore, topK); - - // Normalize the scores - if constexpr (DoSoftmaxBeforeTopK) - { - float sum = float{1.f}; - if (normTopkProb) - { - sum = static_cast(laneIdx < topK ? warpTopKScore[laneIdx] : 0); - sum = cg::reduce(warp, sum, cg::plus()); - } - if (laneIdx < topK) - { - warpTopKScore[laneIdx] = warpTopKScore[laneIdx] / sum; - } - } - else - { - if (applySoftmaxAfterTopK) - { - auto softmaxScore = calcSoftmax(warp, laneIdx < topK ? warpTopKScore[laneIdx] : minScore, laneIdx, topK); - if (laneIdx < topK) - { - warpTopKScore[laneIdx] = softmaxScore; - } - } - } -} - -template -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesBlockKernel(KernelParams params) -{ - // types used in this kernel - using OutputT = typename KernelParams::OutputT; - using InputT = typename KernelParams::InputT; - using BaseType = std::conditional_t; - 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)) - // then wait on primary grid - 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) - { - // in this case, each warp represents a token - BaseType score[VecSize]; - int32_t idx[VecSize]; - - BaseType warpTopKScore[MaxSupportedTopExperts]; - int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; - - BaseType minScore = BaseType{-INFINITY}; - if (validToken) - { - routingTopKExperts(warp, score, idx, - warpTopKScore, warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, - params.mPtrScores + scoreOffset, 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]}; - } - } - } // end if (validToken) - } - __syncthreads(); - - // set local experts - auto localExpertIdx = expert - params.mLocalExpertsStartIdx; - auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < params.mNumLocalExperts - && (localExpertIdx & params.mLocalExpertsStrideLog2) == 0; - // Get the count of each expert and the offset for each token - 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(); - // Get the number of CTAs and the offset for each CTA - 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); - - 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); - } - } - - // at this point, we can write out padded count - 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)) - // we can trigger the next kernel at this point - 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__(NumThreads) - routingIndicesClusterKernel(KernelParams params) -{ - // number of tokens/expanded idx is bounded by total number of warps - using OutputT = typename KernelParams::OutputT; - using InputT = typename KernelParams::InputT; - - using BaseType = std::conditional_t; - using TypePacked = PackedScoreIdx; - - static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; - - __shared__ TypePacked __attribute((aligned(128))) smemPackedScoreIdx[NumWarps * MaxSupportedTopExperts]; - - uint32_t const clusterBlockRank = blockIdx.x; - - int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); - int32_t const laneIdx = cutlass::arch::LaneId(); - - auto warpTokenIdx = clusterBlockRank * NumWarps + warpIdx; - auto scoreOffset = warpTokenIdx * params.mNumExperts; - bool validToken = warpTokenIdx < params.mNumTokens; - - auto block = cg::this_thread_block(); - auto warp = cg::tiled_partition(block); - - // then wait on primary grid - if constexpr (KernelParams::UsePdl) - { - cudaGridDependencySynchronize(); - } - - if (params.mPtrScores != nullptr) - { - // in this case, each warp represents a token - BaseType score[VecSize]; - int32_t idx[VecSize]; - - BaseType warpTopKScore[MaxSupportedTopExperts]; - int32_t warpTopKExpertIdx[MaxSupportedTopExperts]; - - BaseType minScore = BaseType{-INFINITY}; - if (validToken) - { - routingTopKExperts(warp, score, idx, - warpTopKScore, warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, - params.mPtrScores + scoreOffset, params.mNormTopkProb); - - if (laneIdx < params.mTopK) - { - smemPackedScoreIdx[warpIdx * params.mTopK + laneIdx] - = TypePacked{warpTopKScore[laneIdx], static_cast(warpTopKExpertIdx[laneIdx])}; - } - } // end if (validToken) - } - - // make packed scores available to all threads in cluster - __cluster_barrier_arrive(); - __cluster_barrier_wait(); - - if (params.mPtrScores != nullptr) - { - routingPermutation(params, smemPackedScoreIdx, warpIdx, clusterBlockRank); - } - else - { - routingPermutation(params, smemPackedScoreIdx, warpIdx, clusterBlockRank); - } -} -#else -__global__ void __launch_bounds__(NumThreads) routingIndicesClusterKernel(KernelParams /* params */) -{ - assert(false && "routingIndicesClusterKernel is only supported on SM90+ architectures"); -} -#endif // if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) //////////////////////////////////////////////////////////////////////////////////////////////////// - -// this kernel is needed in case we have scores as input for the histogram kernel -template -__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHistogramScoresKernel(KernelParams params) -{ - using OutputT = typename KernelParams::OutputT; - using InputT = typename KernelParams::InputT; - using BaseType = std::conditional_t; - - 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)) - // Wait on primary grid. - if constexpr (KernelParams::UsePdl) - { - cudaGridDependencySynchronize(); - } -#endif // if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - - // initialize the mPtrExpertCounts - 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)) - // Trigger secondary kernel. - if constexpr (KernelParams::UsePdl) - { - cudaTriggerProgrammaticLaunchCompletion(); - } -#endif // if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) - - // in this case, each warp represents a token, and we use a grid-stride loop - // over all warps/tokens - BaseType allScores[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, allExpertIdx, - warpTopKScore, warpTopKExpertIdx, laneIdx, params.mNumExperts, params.mTopK, - params.mPtrScores + scoreOffset, 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_RENORNALIZE(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"); \ - } +// Forward declarations of per-kernel launch wrappers (defined in routingRenormalize/*.cu). +void launchBlockKernel(Data const& data, uint32_t numThreadsHist, void* stream); +void launchClusterKernel(Data const& data, void* stream); +void launchHistogramScoresKernel(Data const& data, uint32_t maxNumBlocks, uint32_t numThreadsHist, void* stream); +void launchInitExpertCounts(Data const& data, uint32_t numThreadsHist, void* stream); +void launchHistogramKernel(Data const& data, int numBlocksHistogram, uint32_t numThreadsHist, void* stream); +void launchOffsetsKernel(Data const& data, int numBlocksOffsets, uint32_t numThreadsHist, void* stream); //////////////////////////////////////////////////////////////////////////////////////////////////// void run(Data const& data, void* stream) @@ -488,10 +44,8 @@ void run(Data const& data, void* stream) "Llama4 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); - // static_assert(MaxNumExperts <= NumThreads, "#experts must be bounded by #threads"); - // static_assert(MaxNumExperts <= NumThreadsHist, "#experts must be bounded by #threads"); //@todo: check how to add + TLLM_CHECK_WITH_INFO(data.mNumExperts <= MaxSupportedExperts, + "Routing kernel expects #experts %d to be no more than %d", data.mNumExperts, MaxSupportedExperts); // similar check TLLM_CHECK_WITH_INFO( data.mNumExperts % 4 == 0, "Routing kernel expects #experts %d to be a multiple of 4.", data.mNumExperts); @@ -509,20 +63,16 @@ void run(Data const& data, void* stream) TLLM_CHECK_WITH_INFO( data.mPtrExpertCounts != nullptr, "When #tokens is large, `mPtrExpertCounts` is a required input."); } - uint32_t const numThreadsHist = getMaxNumExperts(data.mNumExperts); + uint32_t const numThreadsHist = min(1024, getMaxNumExperts(data.mNumExperts)); if (useSingleBlock) { //@TODO: For now we use the single block kernel for cases with token number no larger than 4. // We will future tune this threshold based on the performance. - LAUNCH_ROUTING_RENORNALIZE(data, false, routingIndicesBlockKernel, 1, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mDoSoftmaxBeforeTopK); + launchBlockKernel(data, numThreadsHist, stream); } else if (useSingleCluster) { - LAUNCH_ROUTING_RENORNALIZE(data, false, routingIndicesClusterKernel, NumBlocksPerCluster, NumThreads, - /*smemSize=*/0, // No dynamic smem - stream, data.mDoSoftmaxBeforeTopK); + launchClusterKernel(data, stream); } else { @@ -540,24 +90,15 @@ void run(Data const& data, void* stream) if (data.mPtrScores != nullptr && data.mPtrTopKIds == nullptr) { - LAUNCH_ROUTING_RENORNALIZE(data, false, routingIndicesHistogramScoresKernel, maxNumBlocks, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mDoSoftmaxBeforeTopK); + launchHistogramScoresKernel(data, maxNumBlocks, numThreadsHist, stream); } else { // Reset the global histograms. - LAUNCH_ROUTING_RENORNALIZE(data, false, routingInitExpertCounts, - (2 * data.mNumExperts - 1) / numThreadsHist + 1, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mDoSoftmaxBeforeTopK); + launchInitExpertCounts(data, numThreadsHist, stream); } - LAUNCH_ROUTING_RENORNALIZE(data, false, routingIndicesHistogramKernel, numBlocksHistogram, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mDoSoftmaxBeforeTopK); - LAUNCH_ROUTING_RENORNALIZE(data, false, routingIndicesOffsetsKernel, numBlocksOffsets, numThreadsHist, - /*smemSize=*/0, // No dynamic smem - stream, data.mDoSoftmaxBeforeTopK); + launchHistogramKernel(data, numBlocksHistogram, numThreadsHist, stream); + launchOffsetsKernel(data, numBlocksOffsets, numThreadsHist, stream); } } diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/RoutingDeepSeekCommon.cuh b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/RoutingDeepSeekCommon.cuh new file mode 100644 index 000000000000..b9673be5efe2 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/RoutingDeepSeekCommon.cuh @@ -0,0 +1,115 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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. + */ +#pragma once + +#include "../RoutingKernel.cuh" + +namespace moe::dev::routing +{ +namespace routingDeepSeek +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// +static constexpr int NumNemotronExperts = 512; +static constexpr int NumKimiK2Experts = 384; +static constexpr int NumDeepseekExperts = 256; +static constexpr int MaxSupportedExpertCount = std::max({NumNemotronExperts, NumKimiK2Experts, NumDeepseekExperts}); +static constexpr int NumTopGroupScores = 2; +static constexpr int MaxNumTopGroups = 4; +static constexpr int MaxNumGroups = 8; + +static constexpr int NumTop8Experts = 8; +static constexpr int NumTop22Experts = 22; +static constexpr int MaxSupportedTopExperts = 32; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +int constexpr getMaxNumExperts(int32_t numExperts) +{ + if (numExperts <= topk::MaxNumExpertsUnit) + { + return topk::MaxNumExpertsUnit; + } + else if (numExperts <= NumDeepseekExperts) + { + return NumDeepseekExperts; + } + else if (numExperts <= NumKimiK2Experts) + { + return NumKimiK2Experts; + } + else if (numExperts <= NumNemotronExperts) + { + return NumNemotronExperts; + } + else + { + TLLM_LOG_ERROR("Unsupported numExperts"); + return 0; + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// +// Helper macro: dispatch on topK tier for a given numExperts tier. +#define LAUNCH_DEEPSEEK_WITH_TOPK( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, forceFloatInput, numExperts) \ + if (data.mTopK <= NumTop8Experts) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, \ + stream, extraFlag1, forceFloatInput, numExperts, NumTop8Experts); \ + } \ + else if (data.mTopK <= NumTop22Experts) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, \ + stream, extraFlag1, forceFloatInput, numExperts, NumTop22Experts); \ + } \ + else \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS_FORCE_FLOAT_INPUT(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, \ + stream, extraFlag1, forceFloatInput, numExperts, MaxSupportedTopExperts); \ + } + +#define LAUNCH_ROUTING_DEEPSEEK( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, forceFloatInput) \ + if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ + { \ + LAUNCH_DEEPSEEK_WITH_TOPK(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, \ + forceFloatInput, topk::MaxNumExpertsUnit); \ + } \ + else if (data.mNumExperts <= NumDeepseekExperts) \ + { \ + LAUNCH_DEEPSEEK_WITH_TOPK(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, \ + forceFloatInput, NumDeepseekExperts); \ + } \ + else if (data.mNumExperts <= NumKimiK2Experts) \ + { \ + LAUNCH_DEEPSEEK_WITH_TOPK(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, \ + forceFloatInput, NumKimiK2Experts); \ + } \ + else if (data.mNumExperts <= NumNemotronExperts) \ + { \ + LAUNCH_DEEPSEEK_WITH_TOPK(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, \ + forceFloatInput, NumNemotronExperts); \ + } \ + else \ + { \ + TLLM_LOG_ERROR("Unsupported numExperts"); \ + } + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingDeepSeek +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchClusterKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchClusterKernel.cu new file mode 100644 index 000000000000..14fc591f4e57 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchClusterKernel.cu @@ -0,0 +1,64 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingDeepSeekCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingDeepSeek +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) +__global__ void __cluster_dims__(NumBlocksPerCluster, 1, 1) __launch_bounds__(KernelParams::MaxNumExperts) + routingIndicesClusterKernel(KernelParams params) +{ + using OutputT = typename KernelParams::OutputT; + + int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); + int32_t const clusterBlockRank = blockIdx.x; + + //@todo: try to move it into routingPermutation + // then wait on primary grid + if constexpr (KernelParams::UsePdl) + { + cudaGridDependencySynchronize(); + } + routingPermutation(params, nullptr, warpIdx, clusterBlockRank); +} +#else +__global__ void routingIndicesClusterKernel(KernelParams params) +{ + assert(false && "routingIndicesClusterKernel is only supported on SM90+ architectures"); +} +#endif + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchClusterKernel(Data& data, int numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_DEEPSEEK(data, + /*coopLaunch=*/false, routingIndicesClusterKernel, NumBlocksPerCluster, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingDeepSeek +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchCoopKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchCoopKernel.cu new file mode 100644 index 000000000000..a96db74865dc --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchCoopKernel.cu @@ -0,0 +1,276 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingDeepSeekCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingDeepSeek +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) +__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesCoopKernel(KernelParams params) +{ + // number of experts is bounded by number of threads + int constexpr NumThreads = KernelParams::MaxNumExperts; + __shared__ int32_t __attribute((aligned(128))) smemExpertCount[NumThreads]; + __shared__ int32_t __attribute((aligned(128))) smemExpertOffset[NumThreads]; + // needed for the exclusive sum of token offsets + using Scan = cub::BlockScan; + __shared__ typename Scan::TempStorage tempStorage; + // 64 elements -> 128+ registers. Above that we may start to see spilling to local memory. + static constexpr int MaxExpandedIdxPerThread = 64; + + // Initialize grid. + cg::grid_group grid = cg::this_grid(); + // Note: the following is more efficient than grid.block_index() because we don't use y and z. + int32_t const gridBlockIdx = blockIdx.x; + int32_t const gridThreadIdx = NumThreads * gridBlockIdx + threadIdx.x; + int32_t const numBlocks = gridDim.x; + int32_t const numThreadsPerGrid = numBlocks * NumThreads; + + int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); + + auto expandedIdxSize = params.mNumTokens * params.mTopK; + + // pre-fill the counts with 0 + smemExpertCount[threadIdx.x] = 0; + __syncthreads(); + + // then wait on primary grid + if constexpr (KernelParams::UsePdl) + { + cudaGridDependencySynchronize(); + } + + // each thread keeps has some number of "expanded indexes" assigned to it + // for each of these, we keep the associated expert and offset within expert in registers + int32_t expertIndexes[MaxExpandedIdxPerThread]; + int32_t expertOffsets[MaxExpandedIdxPerThread]; + auto localExpertExtent = params.mNumLocalExperts << params.mLocalExpertsStrideLog2; + // In order to avoid a serialization LDG-ATOMS-LDG-ATOMS-..., we skip multiple iterations at a + // time, and branch between a fast path without bound checks and a slow path with bound checks. + int constexpr IterStride = 4; + static_assert(MaxExpandedIdxPerThread % IterStride == 0); + + // Define a lambda to avoid code duplication in both branches. + auto loopBody = [&](int ii, int expandedIdx) + { + int32_t expertIdx + = params.mPtrTopKIds != nullptr ? params.mPtrTopKIds[expandedIdx] : params.mPtrTopKPacked[expandedIdx].idx; + expertIndexes[ii] = expertIdx; + // check whether this expert is local to our GPU at all and ignore if not + auto localExpertIdx = expertIdx - params.mLocalExpertsStartIdx; + auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; + expertOffsets[ii] = isLocalExpert ? atomicAdd(smemExpertCount + expertIdx, 1) : 0; + }; + +#pragma unroll + for (int32_t ii0 = 0; ii0 < MaxExpandedIdxPerThread; ii0 += IterStride) + { + // Whether it's safe to do multiple iterations without bound checks. + bool const takeFastPath = (ii0 + IterStride) * numThreadsPerGrid <= expandedIdxSize; + if (takeFastPath) + { +#pragma unroll + for (int32_t jj = 0; jj < IterStride; jj++) + { + int const ii = ii0 + jj; + auto expandedIdx = static_cast(gridThreadIdx) + ii * numThreadsPerGrid; + loopBody(ii, expandedIdx); + } + } + else + { + bool doBreak = false; +#pragma unroll + for (int32_t jj = 0; jj < IterStride; jj++) + { + int const ii = ii0 + jj; + auto expandedIdx = static_cast(gridThreadIdx) + ii * numThreadsPerGrid; + if (expandedIdx >= expandedIdxSize) + { + doBreak = true; + break; + } + loopBody(ii, expandedIdx); + } + if (doBreak) + { + break; + } + } + } + + // Make histogram (token counts per expert) available to all threads in the block. + __syncthreads(); + + // + // Each thread now represents one expert + // + + // Add the local bin count to the common bin count and get a per-CTA offset. + int32_t const localExpertCount = smemExpertCount[threadIdx.x]; + + int32_t blockExpertOffset = 0; + if (threadIdx.x < params.mNumExperts) + { + blockExpertOffset = atomicAdd(¶ms.mPtrExpertCounts[threadIdx.x], localExpertCount); + } + + // Sync to wait for completion of the histogram reduction. + grid.sync(); + + // Get total count for this expert. + int32_t count = (threadIdx.x < params.mNumExperts) ? params.mPtrExpertCounts[threadIdx.x] : 0; + + // Note: the scan is redundant in all CTAs, but doing it in only 1 CTA would be worse for latency. + + // Compute the runtime config for projections + // Whether or not an expert is local is taken into account when smemExpertCount is computed + // so we do not need to take it into account here. + + int32_t numCta; + if constexpr (KernelParams::isPow2) + { + numCta = divUpLog2(count, params.mPaddingLog2); + } + else + { + numCta = divUpTileN(count, params.mTileTokensDim); + } + + int32_t ctaOffset; + int32_t numNonExitingCtas; + Scan(tempStorage).ExclusiveSum(numCta, ctaOffset, numNonExitingCtas); + + for (int32_t cta = gridBlockIdx; cta < numCta; cta += numBlocks) + { + const int32_t localExpertIdx = (threadIdx.x - 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) + count; + } + else + { + mnLimit1 = mulTileN(ctaOffset + cta + 1, params.mTileTokensDim); + mnLimit2 = mulTileN(ctaOffset, params.mTileTokensDim) + count; + } + params.mPtrCtaIdxXyToMnLimit[ctaOffset + cta] = min(mnLimit1, mnLimit2); + } + + // get the padded offset associated with this expert + int32_t offset; + if constexpr (KernelParams::isPow2) + { + offset = mulLog2(ctaOffset, params.mPaddingLog2); + } + else + { + offset = mulTileN(ctaOffset, params.mTileTokensDim); + } + int32_t permutedIdxSize; + if constexpr (KernelParams::isPow2) + { + permutedIdxSize = mulLog2(numNonExitingCtas, params.mPaddingLog2); + } + else + { + permutedIdxSize = mulTileN(numNonExitingCtas, params.mTileTokensDim); + } + + // write out padded count + if (gridBlockIdx == 0 && warpIdx == NumThreads / WarpSize - 1 && cute::elect_one_sync()) + { + params.mPtrPermutedIdxSize[0] = permutedIdxSize; + params.mPtrNumNonExitingCtas[0] = numNonExitingCtas; + } + + // write expert offsets to shared + smemExpertOffset[threadIdx.x] = offset + blockExpertOffset; + + // make expert offsets available to all threads + __syncthreads(); + + // trigger the secondary kernel when using PDL + // We can't do it earlier because FC1 depends on the mPtrCtaIdxXyToBatchIdx, + // mPtrCtaIdxXyToMnLimit, mPtrNumNonExitingCtas and mPtrTotalNumPaddedTokens + // TODO: this is not sufficient to ensure visibility in the next kernel! + if constexpr (KernelParams::UsePdl) + { + cudaTriggerProgrammaticLaunchCompletion(); + } + +// each thread has the same "expanded indexes" assigned to it as above +// at this point, we know the final offsets of experts and the offsets within +// experts, which allows writing the final index values +#pragma unroll + for (int32_t ii = 0; ii < MaxExpandedIdxPerThread; ++ii) + { + auto expandedIdx = static_cast(gridThreadIdx) + ii * numThreadsPerGrid; + if (expandedIdx >= expandedIdxSize) + { + break; + } + auto expertIdx = expertIndexes[ii]; + // check whether this expert is local to our GPU at all + auto localExpertIdx = static_cast(expertIdx) - params.mLocalExpertsStartIdx; + auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; + auto tokenIdx = expandedIdx / params.mTopK; + auto permutedIdx = isLocalExpert ? int32_t{smemExpertOffset[expertIdx]} + expertOffsets[ii] : 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; + } + } +} +#else +__global__ void routingIndicesCoopKernel(KernelParams params) +{ + assert(false && "routingIndicesCoopKernel is only supported on SM90+ architectures"); +} +#endif + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchCoopKernel(Data& data, int numBlocksCoop, int numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_DEEPSEEK(data, + /*coopLaunch=*/true, routingIndicesCoopKernel, numBlocksCoop, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingDeepSeek +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchHistogramKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchHistogramKernel.cu new file mode 100644 index 000000000000..1263e289e134 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchHistogramKernel.cu @@ -0,0 +1,36 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingDeepSeekCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingDeepSeek +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchHistogramKernel(Data& data, int numBlocksHistogram, int numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_DEEPSEEK(data, + /*coopLaunch=*/false, routingIndicesHistogramKernel, numBlocksHistogram, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingDeepSeek +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchInitExpertCounts.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchInitExpertCounts.cu new file mode 100644 index 000000000000..5f265878a388 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchInitExpertCounts.cu @@ -0,0 +1,36 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingDeepSeekCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingDeepSeek +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchInitExpertCounts(Data& data, int numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_DEEPSEEK(data, false, routingInitExpertCounts, (2 * data.mNumExperts - 1) / numThreadsHist + 1, + numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/false); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingDeepSeek +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchMainKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchMainKernel.cu new file mode 100644 index 000000000000..1edc469cf70b --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchMainKernel.cu @@ -0,0 +1,289 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingDeepSeekCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingDeepSeek +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +__global__ void routingMainKernel(KernelParams params) +{ + // declare types + using OutputT = typename KernelParams::OutputT; + using InputT = typename KernelParams::InputT; + + // declare shared memory structure + // number of experts is bounded by number of threads + __shared__ float __attribute((aligned(128))) smemScoreSigmoid[KernelParams::MaxNumExperts]; + __shared__ float __attribute((aligned(128))) smemScoreBias[KernelParams::MaxNumExperts]; + // number of expert groups is bounded by number of warps + __shared__ float __attribute((aligned(128))) smemGroupScores[MaxNumGroups]; + + // needed for warp reduce + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + // for the final reduction of weight norm, only some lanes need to participate + int32_t laneIdx = threadIdx.x % WarpSize; + int32_t warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); + // warps outside the range of expert groups do not participate + if constexpr (KernelParams::UseGroups) + { + if (warpIdx >= params.mNumExpertGroups) + { + return; + } + } + + // note that for invalid scores, we simply use a negative value: + // they work well even with the compacted format used in topK, and + // sigmoid / bias activated scores cannot be negative + static constexpr float invalidScoreFloat = float{-INFINITY}; + const OutputT invalidScore = OutputT{invalidScoreFloat}; + + // load bias already; each warp represents one expert group + auto threadExpert = threadIdx.x; + bool expertSelected = threadExpert < params.mNumExperts; + if constexpr (KernelParams::UseGroups) + { + threadExpert = warpIdx * params.mNumExpertsPerGroup + laneIdx; + expertSelected = laneIdx < params.mNumExpertsPerGroup; + } + auto scoreIdx = int64_t{blockIdx.x} * int64_t{params.mNumExperts} + threadExpert; + auto biasVal = expertSelected ? params.mPtrRoutingBias[threadExpert] : invalidScore; + + // initialize the mPtrExpertCounts + if (params.mPtrExpertCounts) + { + int32_t globalThreadIdx = blockIdx.x * blockDim.x + threadIdx.x; + int32_t globalThreadStride = gridDim.x * blockDim.x; + int32_t expertCountsNum = 2 * params.mNumExperts; + initArr(globalThreadIdx, expertCountsNum, globalThreadStride, params.mPtrExpertCounts, 0); + } + +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) + // trigger the secondary kernel when using PDL, then wait on primary + if constexpr (KernelParams::UsePdl) + { + cudaTriggerProgrammaticLaunchCompletion(); + cudaGridDependencySynchronize(); + } +#endif + + if (params.mPtrScores != nullptr) + { + // get our assigned thread score; each warp represents one expert group + float score = expertSelected ? static_cast(params.mPtrScores[scoreIdx]) : invalidScoreFloat; + // get the sigmoid score + // note that for invalid values, we simply use a negative value: + // sigmoig scores are always strictly positive + auto scoreSigmoid = sigmoid_accurate(score); + // write the sigmoid score to shared for later use + if (expertSelected) + { + smemScoreSigmoid[threadExpert] = scoreSigmoid; + } + // get the score with bias + // note that with invalid values, because sigmoid is < 1 and bias is -1, + // we must get a negative value, which is smaller than any valid value + auto scoreBias = float{scoreSigmoid + float{biasVal}}; + + if (expertSelected) + { + smemScoreBias[threadExpert] = scoreBias; + } + + // registers for top group score reduction + float topExpGroupScores[NumTopGroupScores]; + [[maybe_unused]] int32_t topExpGroupIdx[NumTopGroupScores]; + float topGroups[MaxNumTopGroups]; // bound of params.mNumLimitedGroups + int32_t topGroupIdx[MaxNumTopGroups]; + float expertScoreGroup[MaxNumTopGroups]; + int32_t expertIdxGroup[MaxNumTopGroups]; + float topScores[KernelParams::MaxNumTopExperts]; // bound of params.mTopK + int32_t topExperts[KernelParams::MaxNumTopExperts]; + + if constexpr (KernelParams::UseGroups) + { + topk::reduceTopK(warp, topExpGroupScores, topExpGroupIdx, scoreBias, threadExpert, + /* minValue */ invalidScoreFloat); + // get the final group score and write it to shared + if (cute::elect_one_sync()) + { + auto groupScore = topExpGroupScores[0] + topExpGroupScores[1]; + smemGroupScores[warpIdx] = groupScore; + } + } + + // make group scores available to all warps + __syncthreads(); + + auto localExpertExtent = params.mNumLocalExperts << params.mLocalExpertsStrideLog2; + if constexpr (KernelParams::UseGroups) + { // a single warp performs the selection of top groups, and goes on to select the final experts + if (warpIdx == 0) + { + float groupScore = laneIdx < params.mNumExpertGroups ? smemGroupScores[laneIdx] : invalidScoreFloat; + topk::reduceTopK(warp, topGroups, topGroupIdx, groupScore, laneIdx, + /* minValue */ invalidScoreFloat); + // final expert selection: get relevant indexes and scores from shared +#pragma unroll + for (int ii = 0; ii < MaxNumTopGroups; ++ii) + { // bound of params.mNumLimitedGroups + auto groupIdx = topGroupIdx[ii]; + expertIdxGroup[ii] = groupIdx * params.mNumExpertsPerGroup + laneIdx; + // note: expertSelected implies laneIdx < params.mNumExpertsPerGroup. + // we have params.mNumExpertsPerGroup == params.mNumExperts / params.mNumExpertGroups, + // thus groupIdx <= params.mNumExpertGroups - 1 => + // groupIdx * params.mNumExpertsPerGroup <= params.mNumExperts - params.mNumExpertsPerGroup + // => expertIdxGroup[ii] < params.mNumExperts <= NumThreads, + // so the access is safe here + expertScoreGroup[ii] + = (ii < params.mNumLimitedGroups) && (groupIdx < params.mNumExpertGroups) && expertSelected + ? smemScoreBias[expertIdxGroup[ii]] + : invalidScoreFloat; + } + + topk::reduceTopK(warp, topScores, topExperts, expertScoreGroup, expertIdxGroup, + /* minValue */ invalidScoreFloat, params.mTopK); + } + } + else if constexpr (KernelParams::MaxNumExperts > topk::MaxNumExpertsUnit) + { + // without groups, each thread just takes `MaxNumTopGroups` experts + int constexpr NumExpertWarps = (KernelParams::MaxNumExperts - 1) / topk::MaxNumExpertsUnit + 1; + int constexpr NumInterTopK = NumExpertWarps * KernelParams::MaxNumTopExperts; + __shared__ float __attribute((aligned(128))) smemInterTopScores[NumInterTopK]; + __shared__ int32_t __attribute((aligned(128))) smemInterTopExperts[NumInterTopK]; + if (warpIdx < NumExpertWarps) + { + int offset = warpIdx * WarpSize * MaxNumTopGroups; +#pragma unroll + for (int ii = 0; ii < MaxNumTopGroups; ++ii) + { + auto expertIdx = ii * WarpSize + laneIdx; + expertIdxGroup[ii] = offset + expertIdx; + expertScoreGroup[ii] = offset + expertIdx < params.mNumExperts ? smemScoreBias[offset + expertIdx] + : invalidScoreFloat; + } + topk::reduceTopK(warp, topScores, topExperts, expertScoreGroup, expertIdxGroup, + /* minValue */ invalidScoreFloat, params.mTopK); + + if (laneIdx < params.mTopK) + { + smemInterTopScores[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] = topScores[laneIdx]; + smemInterTopExperts[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] = topExperts[laneIdx]; + } + else if (laneIdx >= params.mTopK && laneIdx < KernelParams::MaxNumTopExperts) + { + smemInterTopScores[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] = invalidScoreFloat; + smemInterTopExperts[warpIdx * KernelParams::MaxNumTopExperts + laneIdx] + = MaxSupportedExpertCount - 1; + } + } + __syncthreads(); + if (warpIdx == 0) + { + int constexpr NumInterTopKPerThread = (NumInterTopK - 1) / WarpSize + 1; + float intermediateScore[NumInterTopKPerThread]; + int32_t intermediateExpert[NumInterTopKPerThread]; + for (int i = laneIdx; i < NumInterTopKPerThread * WarpSize; i += WarpSize) + { + int ii = i / WarpSize; + if (i < NumInterTopK) + { + intermediateScore[ii] = smemInterTopScores[i]; + intermediateExpert[ii] = smemInterTopExperts[i]; + } + else + { + intermediateScore[ii] = invalidScoreFloat; + intermediateExpert[ii] = KernelParams::MaxNumExperts - 1; + } + } + topk::reduceTopK(warp, topScores, topExperts, intermediateScore, intermediateExpert, + /* minValue */ invalidScoreFloat, params.mTopK); + } + } + else + { + if (warpIdx == 0) + { + // without groups, each thread just takes `MaxNumTopGroups` experts +#pragma unroll + for (int ii = 0; ii < MaxNumTopGroups; ++ii) + { + auto expertIdx = ii * WarpSize + laneIdx; + expertIdxGroup[ii] = expertIdx; + expertScoreGroup[ii] + = expertIdx < params.mNumExperts ? smemScoreBias[expertIdx] : invalidScoreFloat; + } + topk::reduceTopK(warp, topScores, topExperts, expertScoreGroup, expertIdxGroup, + /* minValue */ invalidScoreFloat, params.mTopK); + } + } + + if (warpIdx == 0) + { + // determine our lane's expert index and write to output + int32_t expertIdx = 0; +#pragma unroll + for (int ii = 0; ii < params.mTopK; ++ii) + { // bound of params.mTopK + expertIdx = laneIdx == ii ? topExperts[ii] : expertIdx; + } + // determine whether our expert is local to this GPU + auto localExpertIdx = expertIdx - params.mLocalExpertsStartIdx; + auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < localExpertExtent + && (localExpertIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; + + float scoreNorm = laneIdx < params.mTopK ? smemScoreSigmoid[expertIdx] : 0.F; + auto redNorm = cg::reduce(warp, scoreNorm, cg::plus{}); + auto finalScore = OutputT{scoreNorm * params.mRouteScale / redNorm}; + + // write expert idx out already + auto idxTopK = blockIdx.x * params.mTopK + laneIdx; + if (laneIdx < params.mTopK && params.mPtrTopKPacked != nullptr) + { + PackedScoreIdx packedScore{static_cast(finalScore), static_cast(expertIdx)}; + params.mPtrTopKPacked[idxTopK] = packedScore; + } + + if (laneIdx < params.mTopK && params.mPtrTopKWeights != nullptr && params.mPtrTopKIds == nullptr) + { + params.mPtrTopKWeights[idxTopK] = finalScore; + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchMainKernel(Data& data, int numBlocks, int numThreadsMain, void* stream) +{ + LAUNCH_ROUTING_DEEPSEEK(data, + /*coopLaunch=*/false, routingMainKernel, numBlocks, numThreadsMain, + /*smemSize=*/0, // No dynamic smem + stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingDeepSeek +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchOffsetsKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchOffsetsKernel.cu new file mode 100644 index 000000000000..0836c21aa5cf --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingDeepSeek/launchOffsetsKernel.cu @@ -0,0 +1,36 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingDeepSeekCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingDeepSeek +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchOffsetsKernel(Data& data, int numBlocksOffsets, int numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_DEEPSEEK(data, + /*coopLaunch=*/false, routingIndicesOffsetsKernel, numBlocksOffsets, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mNumExpertGroups > 1, /*forceFloatInput=*/true); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingDeepSeek +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/RoutingRenormalizeCommon.cuh b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/RoutingRenormalizeCommon.cuh new file mode 100644 index 000000000000..c8bc00e2fb60 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/RoutingRenormalizeCommon.cuh @@ -0,0 +1,160 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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. + */ +#pragma once + +#include "../RoutingKernel.cuh" + +namespace moe::dev::routing +{ +namespace routingRenormalize +{ +//////////////////////////////////////////////////////////////////////////////////////////////////// + +static constexpr int NumExperts128Experts = 128; +static constexpr int NumExperts512Experts = 512; +static constexpr int MaxSupportedExperts = 2048; + +static constexpr int NumTop8Experts = 8; +static constexpr int NumTop16Experts = 16; +static constexpr int MaxSupportedTopExperts = 32; + +static constexpr int NumThreads = 1024; +static constexpr int NumWarps = NumThreads / WarpSize; + +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], int32_t (&idx)[VecSize], DataType (&warpTopKScore)[K], int32_t (&warpTopKExpertIdx)[K], + int32_t const laneIdx, int32_t const numExperts, int32_t topK, InputType const* ptrScores, bool const normTopkProb, + bool const applySoftmaxAfterTopK = true) +{ + DataType minScore = DataType{-INFINITY}; + + for (int i = 0; i < VecSize; i++) + { + auto expertIdx = i * WarpSize + laneIdx; + auto newScore = expertIdx < numExperts ? static_cast(ptrScores[expertIdx]) : minScore; + score[i] = newScore; + idx[i] = expertIdx; + } + if constexpr (DoSoftmaxBeforeTopK) + { + calcSoftmax(warp, score); + } + + // Get the top-k scores and their corresponding expert indices + topk::reduceTopK(warp, warpTopKScore, warpTopKExpertIdx, score, idx, minScore, topK); + + // Normalize the scores + if constexpr (DoSoftmaxBeforeTopK) + { + float sum = float{1.f}; + if (normTopkProb) + { + sum = static_cast(laneIdx < topK ? warpTopKScore[laneIdx] : 0); + sum = cg::reduce(warp, sum, cg::plus()); + } + if (laneIdx < topK) + { + warpTopKScore[laneIdx] = warpTopKScore[laneIdx] / sum; + } + } + else + { + if (applySoftmaxAfterTopK) + { + auto softmaxScore = calcSoftmax(warp, laneIdx < topK ? warpTopKScore[laneIdx] : minScore, laneIdx, topK); + if (laneIdx < topK) + { + warpTopKScore[laneIdx] = softmaxScore; + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +int32_t constexpr getMaxNumExperts(int32_t numExperts) +{ + if (numExperts <= NumExperts128Experts) + { + return NumExperts128Experts; + } + else if (numExperts <= NumExperts512Experts) + { + return NumExperts512Experts; + } + else if (numExperts <= MaxSupportedExperts) + { + return MaxSupportedExperts; + } + else + { + TLLM_LOG_ERROR("Unsupported numExperts"); + return 0; + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Helper macro: dispatch on topK tier for a given numExperts tier. +#define LAUNCH_ROUTING_WITH_TOPK( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, numExperts) \ + if (data.mTopK <= NumTop8Experts) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, \ + numExperts, NumTop8Experts); \ + } \ + else if (data.mTopK <= NumTop16Experts) \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, \ + numExperts, NumTop16Experts); \ + } \ + else \ + { \ + LAUNCH_ROUTING_WITH_NUM_EXPERTS(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, \ + numExperts, MaxSupportedTopExperts); \ + } + +#define LAUNCH_ROUTING_RENORMALIZE(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1) \ + if (data.mNumExperts <= NumExperts128Experts) \ + { \ + LAUNCH_ROUTING_WITH_TOPK( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, NumExperts128Experts); \ + } \ + else if (data.mNumExperts <= NumExperts512Experts) \ + { \ + LAUNCH_ROUTING_WITH_TOPK( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, NumExperts512Experts); \ + } \ + else if (data.mNumExperts <= MaxSupportedExperts) \ + { \ + LAUNCH_ROUTING_WITH_TOPK( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, extraFlag1, MaxSupportedExperts); \ + } \ + else \ + { \ + TLLM_LOG_ERROR("Unsupported numExperts"); \ + } + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingRenormalize +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu new file mode 100644 index 000000000000..93d273f43877 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchBlockKernel.cu @@ -0,0 +1,294 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingRenormalizeCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingRenormalize +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +__global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelParams::MaxNumExperts : 1024) + routingIndicesBlockKernel(KernelParams params) +{ + // types used in this kernel + using OutputT = typename KernelParams::OutputT; + using InputT = typename KernelParams::InputT; + using BaseType = std::conditional_t; + using TypePacked = PackedScoreIdx; + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + // When MaxNumExperts > 1024, cap actual thread count at 1024 and let each thread handle + // multiple experts. This is needed because CUDA blocks support at most 1024 threads. + static constexpr int NumThreadsBlock = MaxNumExperts <= 1024 ? MaxNumExperts : 1024; + static constexpr int ExpertsPerThread = MaxNumExperts / NumThreadsBlock; + static_assert(MaxNumExperts % NumThreadsBlock == 0, "MaxNumExperts must be a multiple of NumThreadsBlock"); + + int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); + int32_t const laneIdx = cutlass::arch::LaneId(); + 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)) + // then wait on primary grid + 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) + { + // in this case, each warp represents a token + BaseType score[VecSize]; + int32_t idx[VecSize]; + + BaseType warpTopKScore[KernelParams::MaxNumTopExperts]; + int32_t warpTopKExpertIdx[KernelParams::MaxNumTopExperts]; + + BaseType minScore = BaseType{-INFINITY}; + if (validToken) + { + routingTopKExperts(warp, score, idx, warpTopKScore, warpTopKExpertIdx, laneIdx, + params.mNumExperts, params.mTopK, params.mPtrScores + scoreOffset, 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]}; + } + } + } // end if (validToken) + } + __syncthreads(); + + // Each thread handles ExpertsPerThread contiguous experts. + // Thread i handles experts [i * ExpertsPerThread, (i+1) * ExpertsPerThread). + // Contiguous assignment ensures prefix sum ordering is correct. + int accExpertCount[ExpertsPerThread]; +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) + { + int expert = threadIdx.x * ExpertsPerThread + e; + auto localExpIdx = expert - params.mLocalExpertsStartIdx; + auto isLocal = localExpIdx >= 0 && localExpIdx < params.mNumLocalExperts + && (localExpIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; + + // Get the count of each expert and the offset for each token + accExpertCount[e] = 0; + if (isLocal) + { + int offset = expert; + for (int j = 0; j < BlockKernelMaxNumTokens; j++) + { + if (smemKIdx[offset] >= 0) + { + smemOffset[offset] = static_cast(accExpertCount[e]); + accExpertCount[e]++; + } + offset += MaxNumExperts; + } + } + } + __syncthreads(); + + // Get the number of CTAs and the offset for each CTA. + // Use cub::BlockScan's array overload: each thread holds ExpertsPerThread items, + // and ExclusiveSum computes the prefix sum across all NumThreadsBlock * ExpertsPerThread + // items in thread order — exactly matching our contiguous expert assignment. + int32_t numCtaPerExpert[ExpertsPerThread]; +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) + { + if constexpr (KernelParams::isPow2) + { + numCtaPerExpert[e] = divUpLog2(accExpertCount[e], params.mPaddingLog2); + } + else + { + numCtaPerExpert[e] = divUpTileN(accExpertCount[e], params.mTileTokensDim); + } + } + int32_t ctaOffsetPerExpert[ExpertsPerThread]; + int32_t numNonExitingCtas; + Scan(tempStorage).ExclusiveSum(numCtaPerExpert, ctaOffsetPerExpert, numNonExitingCtas); + __syncthreads(); // Required barrier before reusing TempStorage for the next BlockScan + + // Compute padded expert scan counts (same array-overload pattern) + int32_t tmpCountPerExpert[ExpertsPerThread]; +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) + { + if constexpr (KernelParams::isPow2) + { + tmpCountPerExpert[e] = divUpMulLog2(accExpertCount[e], params.mPaddingLog2); + } + else + { + tmpCountPerExpert[e] = divUpMulTileN(accExpertCount[e], params.mTileTokensDim); + } + } + int32_t expertScanCountsPerExpert[ExpertsPerThread]; + Scan(tempStorage).ExclusiveSum(tmpCountPerExpert, expertScanCountsPerExpert); + __syncthreads(); + + // Write CTA configs for each expert this thread handles +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) + { + int expert = threadIdx.x * ExpertsPerThread + e; + auto localExpIdx = expert - params.mLocalExpertsStartIdx; + auto isLocal = localExpIdx >= 0 && localExpIdx < params.mNumLocalExperts + && (localExpIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; + + if (isLocal) + { + for (int cta = 0; cta < numCtaPerExpert[e]; ++cta) + { + int32_t const mappedLocalIdx + = (expert - params.mLocalExpertsStartIdx) >> params.mLocalExpertsStrideLog2; + params.mPtrCtaIdxXyToBatchIdx[ctaOffsetPerExpert[e] + cta] = mappedLocalIdx; + int32_t mnLimit1; + int32_t mnLimit2; + if constexpr (KernelParams::isPow2) + { + mnLimit1 = mulLog2(ctaOffsetPerExpert[e] + cta + 1, params.mPaddingLog2); + mnLimit2 = mulLog2(ctaOffsetPerExpert[e], params.mPaddingLog2) + accExpertCount[e]; + } + else + { + mnLimit1 = mulTileN(ctaOffsetPerExpert[e] + cta + 1, params.mTileTokensDim); + mnLimit2 = mulTileN(ctaOffsetPerExpert[e], params.mTileTokensDim) + accExpertCount[e]; + } + params.mPtrCtaIdxXyToMnLimit[ctaOffsetPerExpert[e] + cta] = min(mnLimit1, mnLimit2); + } + } + } + + // at this point, we can write out padded count + 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)) + // we can trigger the next kernel at this point + if constexpr (KernelParams::UsePdl) + { + cudaTriggerProgrammaticLaunchCompletion(); + } +#endif + + for (int tokenIdx = 0; tokenIdx < params.mNumTokens; tokenIdx++) + { +#pragma unroll + for (int e = 0; e < ExpertsPerThread; e++) + { + int expert = threadIdx.x * ExpertsPerThread + e; + int offset = tokenIdx * MaxNumExperts + expert; + if (smemKIdx[offset] >= 0) + { + auto localExpIdx = expert - params.mLocalExpertsStartIdx; + auto isLocal = localExpIdx >= 0 && localExpIdx < params.mNumLocalExperts + && (localExpIdx & ((1 << params.mLocalExpertsStrideLog2) - 1)) == 0; + + int const expandedIdx = tokenIdx * params.mTopK + smemKIdx[offset]; + int const offsetWithinExpert = static_cast(smemOffset[offset]); + int const offsetForExpert = expertScanCountsPerExpert[e]; + int const permutedIdx = isLocal ? offsetForExpert + offsetWithinExpert : int32_t{-1}; + + if (params.mPtrExpandedIdxToPermutedIdx != nullptr) + { + params.mPtrExpandedIdxToPermutedIdx[expandedIdx] = permutedIdx; + } + if (params.mPtrPermutedIdxToExpandedIdx != nullptr && isLocal) + { + params.mPtrPermutedIdxToExpandedIdx[permutedIdx] = expandedIdx; + } + if (params.mPtrPermutedIdxToTokenIdx != nullptr && isLocal) + { + params.mPtrPermutedIdxToTokenIdx[permutedIdx] = tokenIdx; + } + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchBlockKernel(Data const& data, uint32_t numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_RENORMALIZE(data, false, routingIndicesBlockKernel, 1, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mDoSoftmaxBeforeTopK); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingRenormalize +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchClusterKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchClusterKernel.cu new file mode 100644 index 000000000000..b8d7f8b91186 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchClusterKernel.cu @@ -0,0 +1,117 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingRenormalizeCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingRenormalize +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) +__global__ void __cluster_dims__(NumBlocksPerCluster, 1, 1) __launch_bounds__(NumThreads) + routingIndicesClusterKernel(KernelParams params) +{ + // number of tokens/expanded idx is bounded by total number of warps + using OutputT = typename KernelParams::OutputT; + using InputT = typename KernelParams::InputT; + + using BaseType = std::conditional_t; + using TypePacked = PackedScoreIdx; + + static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; + + __shared__ TypePacked __attribute((aligned(128))) smemPackedScoreIdx[NumWarps * KernelParams::MaxNumTopExperts]; + + uint32_t const clusterBlockRank = blockIdx.x; + + int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); + int32_t const laneIdx = cutlass::arch::LaneId(); + + auto warpTokenIdx = clusterBlockRank * NumWarps + warpIdx; + auto scoreOffset = warpTokenIdx * params.mNumExperts; + bool validToken = warpTokenIdx < params.mNumTokens; + + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + + // then wait on primary grid + if constexpr (KernelParams::UsePdl) + { + cudaGridDependencySynchronize(); + } + + if (params.mPtrScores != nullptr) + { + // in this case, each warp represents a token + BaseType score[VecSize]; + int32_t idx[VecSize]; + + BaseType warpTopKScore[KernelParams::MaxNumTopExperts]; + int32_t warpTopKExpertIdx[KernelParams::MaxNumTopExperts]; + + BaseType minScore = BaseType{-INFINITY}; + if (validToken) + { + routingTopKExperts(warp, score, idx, warpTopKScore, warpTopKExpertIdx, laneIdx, + params.mNumExperts, params.mTopK, params.mPtrScores + scoreOffset, params.mNormTopkProb); + + if (laneIdx < params.mTopK) + { + smemPackedScoreIdx[warpIdx * params.mTopK + laneIdx] + = TypePacked{warpTopKScore[laneIdx], static_cast(warpTopKExpertIdx[laneIdx])}; + } + } // end if (validToken) + } + + // make packed scores available to all threads in cluster + __cluster_barrier_arrive(); + __cluster_barrier_wait(); + + if (params.mPtrScores != nullptr) + { + routingPermutation(params, smemPackedScoreIdx, warpIdx, clusterBlockRank); + } + else + { + routingPermutation(params, smemPackedScoreIdx, warpIdx, clusterBlockRank); + } +} +#else +__global__ void __launch_bounds__(NumThreads) routingIndicesClusterKernel(KernelParams /* params */) +{ + assert(false && "routingIndicesClusterKernel is only supported on SM90+ architectures"); +} +#endif // if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchClusterKernel(Data const& data, void* stream) +{ + LAUNCH_ROUTING_RENORMALIZE(data, false, routingIndicesClusterKernel, NumBlocksPerCluster, NumThreads, + /*smemSize=*/0, // No dynamic smem + stream, data.mDoSoftmaxBeforeTopK); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingRenormalize +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramKernel.cu new file mode 100644 index 000000000000..7d6f6177a56e --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramKernel.cu @@ -0,0 +1,35 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingRenormalizeCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingRenormalize +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchHistogramKernel(Data const& data, int numBlocksHistogram, uint32_t numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_RENORMALIZE(data, false, routingIndicesHistogramKernel, numBlocksHistogram, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mDoSoftmaxBeforeTopK); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingRenormalize +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramScoresKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramScoresKernel.cu new file mode 100644 index 000000000000..03bec526f823 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchHistogramScoresKernel.cu @@ -0,0 +1,106 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingRenormalizeCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingRenormalize +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// this kernel is needed in case we have scores as input for the histogram kernel +template +__global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelParams::MaxNumExperts : 1024) + routingIndicesHistogramScoresKernel(KernelParams params) +{ + using OutputT = typename KernelParams::OutputT; + using InputT = typename KernelParams::InputT; + using BaseType = std::conditional_t; + // Cap actual thread count at 1024 when MaxNumExperts > 1024. + static constexpr int NumThreadsBlock = KernelParams::MaxNumExperts <= 1024 ? KernelParams::MaxNumExperts : 1024; + + // VecSize stays based on MaxNumExperts — each warp still processes all experts for one token. + static constexpr int VecSize = KernelParams::MaxNumExperts / WarpSize; + + int32_t const laneIdx = cutlass::arch::LaneId(); + int32_t const warpIdx = threadIdx.x / WarpSize; + // Use NumThreadsBlock (actual thread count) for grid-stride warp/thread addressing + int32_t const globalWarpIdx = blockIdx.x * NumThreadsBlock / WarpSize + warpIdx; + int32_t const globalWarpStride = gridDim.x * NumThreadsBlock / WarpSize; + BaseType minScore = BaseType{-INFINITY}; + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + // Wait on primary grid. + if constexpr (KernelParams::UsePdl) + { + cudaGridDependencySynchronize(); + } +#endif // if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + + // initialize the mPtrExpertCounts — use NumThreadsBlock for grid-stride + int32_t expertCountsNum = 2 * params.mNumExperts; + int32_t globalThreadIdx = blockIdx.x * NumThreadsBlock + threadIdx.x; + int32_t globalThreadStride = gridDim.x * NumThreadsBlock; + initArr(globalThreadIdx, expertCountsNum, globalThreadStride, params.mPtrExpertCounts, 0); + + // in this case, each warp represents a token, and we use a grid-stride loop + // over all warps/tokens + BaseType allScores[VecSize]; + int32_t allExpertIdx[VecSize]; + BaseType warpTopKScore[KernelParams::MaxNumTopExperts]; + int32_t warpTopKExpertIdx[KernelParams::MaxNumTopExperts]; + for (int tokenIdx = globalWarpIdx; tokenIdx < params.mNumTokens; tokenIdx += globalWarpStride) + { + auto scoreOffset = tokenIdx * params.mNumExperts; + + routingTopKExperts(warp, allScores, allExpertIdx, warpTopKScore, warpTopKExpertIdx, laneIdx, + params.mNumExperts, params.mTopK, params.mPtrScores + scoreOffset, params.mNormTopkProb); + + if (laneIdx < params.mTopK) + { + PackedScoreIdx packedScore{ + static_cast(warpTopKScore[laneIdx]), static_cast(warpTopKExpertIdx[laneIdx])}; + params.mPtrTopKPacked[tokenIdx * params.mTopK + laneIdx] = packedScore; + } + } + +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + // Trigger secondary kernel AFTER writing all packed scores, so the next kernel + // (routingIndicesHistogramKernel) sees the completed mPtrTopKPacked writes. + if constexpr (KernelParams::UsePdl) + { + cudaTriggerProgrammaticLaunchCompletion(); + } +#endif // if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchHistogramScoresKernel(Data const& data, uint32_t maxNumBlocks, uint32_t numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_RENORMALIZE(data, false, routingIndicesHistogramScoresKernel, maxNumBlocks, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mDoSoftmaxBeforeTopK); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingRenormalize +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchInitExpertCounts.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchInitExpertCounts.cu new file mode 100644 index 000000000000..807fc89e8978 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchInitExpertCounts.cu @@ -0,0 +1,36 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingRenormalizeCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingRenormalize +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchInitExpertCounts(Data const& data, uint32_t numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_RENORMALIZE(data, false, routingInitExpertCounts, (2 * data.mNumExperts - 1) / numThreadsHist + 1, + numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mDoSoftmaxBeforeTopK); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingRenormalize +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchOffsetsKernel.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchOffsetsKernel.cu new file mode 100644 index 000000000000..fe398c80cdb1 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/routingRenormalize/launchOffsetsKernel.cu @@ -0,0 +1,35 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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 "RoutingRenormalizeCommon.cuh" + +namespace moe::dev::routing +{ +namespace routingRenormalize +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +void launchOffsetsKernel(Data const& data, int numBlocksOffsets, uint32_t numThreadsHist, void* stream) +{ + LAUNCH_ROUTING_RENORMALIZE(data, false, routingIndicesOffsetsKernel, numBlocksOffsets, numThreadsHist, + /*smemSize=*/0, // No dynamic smem + stream, data.mDoSoftmaxBeforeTopK); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingRenormalize +} // namespace moe::dev::routing diff --git a/cpp/tensorrt_llm/thop/fp4BlockScaleMoe.cpp b/cpp/tensorrt_llm/thop/fp4BlockScaleMoe.cpp index d1be58e89b2c..e09ef52a52e4 100644 --- a/cpp/tensorrt_llm/thop/fp4BlockScaleMoe.cpp +++ b/cpp/tensorrt_llm/thop/fp4BlockScaleMoe.cpp @@ -125,8 +125,8 @@ std::vector run_fp4_block_scale_moe_runner(torch::optional(routing_method_type) == RoutingMethodType::Renormalize || static_cast(routing_method_type) == RoutingMethodType::RenormalizeNaive) { - TORCH_CHECK(top_k <= 10 && top_k > 0, - "Current routing kernel (no groups, renormalize) only supports top_k<=10 && top_k>0."); + TORCH_CHECK(top_k <= 32 && top_k > 0, + "Current routing kernel (no groups, renormalize) only supports top_k<=32 && top_k>0."); } else if (static_cast(routing_method_type) == RoutingMethodType::Llama4) { @@ -135,7 +135,7 @@ std::vector run_fp4_block_scale_moe_runner(torch::optional top_k, "num_experts must be greater than top_k"); - TORCH_CHECK(num_experts <= 512, "num_experts must be less than or equal to 512"); + TORCH_CHECK(num_experts <= 2048, "num_experts must be less than or equal to 2048"); // If both routing inputs are provided, they must be on the same device if (routing_logits.has_value() && topk_ids.has_value()) diff --git a/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp b/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp index 2db4e2bf6c5b..762294828917 100644 --- a/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp +++ b/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp @@ -121,8 +121,8 @@ at::Tensor run_fp8_block_scale_moe(at::optional const& routing_logit else if (static_cast(routing_method_type) == RoutingMethodType::Renormalize || static_cast(routing_method_type) == RoutingMethodType::RenormalizeNaive) { - TORCH_CHECK(top_k <= 10 && top_k > 0, - "Current routing kernel (no groups, renormalize) only supports top_k<=8 && top_k>0."); + TORCH_CHECK(top_k <= 32 && top_k > 0, + "Current routing kernel (no groups, renormalize) only supports top_k<=32 && top_k>0."); } else if (static_cast(routing_method_type) == RoutingMethodType::Llama4) { diff --git a/cpp/tensorrt_llm/thop/fp8PerTensorScaleMoe.cpp b/cpp/tensorrt_llm/thop/fp8PerTensorScaleMoe.cpp index efefc0663214..092f8f013620 100644 --- a/cpp/tensorrt_llm/thop/fp8PerTensorScaleMoe.cpp +++ b/cpp/tensorrt_llm/thop/fp8PerTensorScaleMoe.cpp @@ -125,6 +125,12 @@ torch::Tensor fp8_per_tensor_scale_moe_runner(torch::optional con { TORCH_CHECK(top_k == 1, "Current routing kernel (no groups, Llama4) only supports top_k=1."); } + else if (static_cast(routing_method_type) == RoutingMethodType::Renormalize + || static_cast(routing_method_type) == RoutingMethodType::RenormalizeNaive) + { + TORCH_CHECK(top_k <= 32 && top_k > 0, + "Current routing kernel (no groups, renormalize) only supports top_k<=32 && top_k>0."); + } TORCH_CHECK(num_experts % 4 == 0, "Routing kernel expects that num_experts must be divisible by 4"); TORCH_CHECK(num_experts > top_k, "num_experts must be greater than top_k"); diff --git a/cpp/tensorrt_llm/thop/mxFp4BlockScaleMoe.cpp b/cpp/tensorrt_llm/thop/mxFp4BlockScaleMoe.cpp index 08bce0611b06..96fb44286cd8 100644 --- a/cpp/tensorrt_llm/thop/mxFp4BlockScaleMoe.cpp +++ b/cpp/tensorrt_llm/thop/mxFp4BlockScaleMoe.cpp @@ -131,8 +131,8 @@ torch::Tensor dtype_mxe2m1_block_scale_moe_runner(torch::optional else if (static_cast(routing_method_type) == RoutingMethodType::Renormalize || static_cast(routing_method_type) == RoutingMethodType::RenormalizeNaive) { - TORCH_CHECK(top_k <= 10 && top_k > 0, - "Current routing kernel (no groups, renormalize) only supports top_k<=10 && top_k>0."); + TORCH_CHECK(top_k <= 32 && top_k > 0, + "Current routing kernel (no groups, renormalize) only supports top_k<=32 && top_k>0."); } TORCH_CHECK(num_experts % 4 == 0, "Routing kernel expects that num_experts must be divisible by 4"); diff --git a/cpp/tests/unit_tests/kernels/routing/routingRenormalizeTest.cpp b/cpp/tests/unit_tests/kernels/routing/routingRenormalizeTest.cpp index c77b384a7c49..738e2aef1903 100644 --- a/cpp/tests/unit_tests/kernels/routing/routingRenormalizeTest.cpp +++ b/cpp/tests/unit_tests/kernels/routing/routingRenormalizeTest.cpp @@ -278,7 +278,7 @@ TYPED_TEST(RoutingRenormalizeKernelTest, ClusterLevelParallelizationWithRenormal TYPED_TEST(RoutingRenormalizeKernelTest, DeviceLevelParallelization) { RoutingKernelTestParam param(RoutingMethodType::Renormalize, /*numTokens=*/1000, - /*numExperts=*/128, /*topK=*/8, + /*numExperts=*/512, /*topK=*/10, /*expertParallelization=*/1, /*expertParallelizationId=*/0, /*tileTokensDim=*/8, /*paddingLog2=*/3, /*localExpertsStrideLog2=*/0, /*usePdl=*/true, /*getExpWeights=*/true, /*useTopKAsInput=*/false, @@ -344,7 +344,7 @@ TYPED_TEST(RoutingRenormalizeKernelTest, DeviceLevelParallelizationTop4) TYPED_TEST(RoutingRenormalizeKernelTest, BlockLevelParallelizationLargeN) { RoutingKernelTestParam param(RoutingMethodType::Renormalize, /*numTokens=*/4, - /*numExperts=*/512, /*topK=*/10, + /*numExperts=*/2048, /*topK=*/32, /*expertParallelization=*/1, /*expertParallelizationId=*/0, /*tileTokensDim=*/256, /*paddingLog2=*/3, /*localExpertsStrideLog2=*/0, /*usePdl=*/true, /*getExpWeights=*/true, /*useTopKAsInput=*/false, /*hasInvalidTopKInput=*/false, @@ -355,7 +355,7 @@ TYPED_TEST(RoutingRenormalizeKernelTest, BlockLevelParallelizationLargeN) TYPED_TEST(RoutingRenormalizeKernelTest, ClusterLevelParallelizationLargeN) { RoutingKernelTestParam param(RoutingMethodType::Renormalize, /*numTokens=*/100, - /*numExperts=*/512, /*topK=*/10, + /*numExperts=*/2048, /*topK=*/32, /*expertParallelization=*/1, /*expertParallelizationId=*/0, /*tileTokensDim=*/256, /*paddingLog2=*/3, /*localExpertsStrideLog2=*/0, /*usePdl=*/true, /*getExpWeights=*/true, /*useTopKAsInput=*/false, /*hasInvalidTopKInput=*/false, @@ -366,7 +366,7 @@ TYPED_TEST(RoutingRenormalizeKernelTest, ClusterLevelParallelizationLargeN) TYPED_TEST(RoutingRenormalizeKernelTest, DeviceLevelParallelizationLargeN) { RoutingKernelTestParam param(RoutingMethodType::Renormalize, /*numTokens=*/1000, - /*numExperts=*/512, /*topK=*/10, + /*numExperts=*/2048, /*topK=*/32, /*expertParallelization=*/1, /*expertParallelizationId=*/0, /*tileTokensDim=*/256, /*paddingLog2=*/3, /*localExpertsStrideLog2=*/0, /*usePdl=*/true, /*getExpWeights=*/true, /*useTopKAsInput=*/false, /*hasInvalidTopKInput=*/false, @@ -377,7 +377,7 @@ TYPED_TEST(RoutingRenormalizeKernelTest, DeviceLevelParallelizationLargeN) TYPED_TEST(RoutingRenormalizeKernelTest, DeviceLevelParallelizationLargeNWithInvalidTopKInput) { RoutingKernelTestParam param(RoutingMethodType::Renormalize, /*numTokens=*/1000, - /*numExperts=*/512, /*topK=*/10, + /*numExperts=*/2048, /*topK=*/32, /*expertParallelization=*/1, /*expertParallelizationId=*/0, /*tileTokensDim=*/256, /*paddingLog2=*/3, /*localExpertsStrideLog2=*/0, /*usePdl=*/true, /*getExpWeights=*/true, /*useTopKAsInput=*/true, /*hasInvalidTopKInput=*/true, diff --git a/cpp/tests/unit_tests/kernels/routing/routingTest.cpp b/cpp/tests/unit_tests/kernels/routing/routingTest.cpp index b8f387f4b415..a7da5a6d7281 100644 --- a/cpp/tests/unit_tests/kernels/routing/routingTest.cpp +++ b/cpp/tests/unit_tests/kernels/routing/routingTest.cpp @@ -161,7 +161,7 @@ void RoutingKernelTest::computePermutation(RoutingKernelTestParam const& para int32_t localExpertIdx = index - param.localExpertsStartIdx; bool isLocalExpert = localExpertIdx >= 0 && localExpertIdx < param.numLocalExperts - && (localExpertIdx & param.localExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << param.localExpertsStrideLog2) - 1)) == 0; if (index >= 0) { tokenToIdxInExpertHostPtr[it * param.topK + ie] = expertCountsHostPtr[index]; @@ -201,7 +201,7 @@ void RoutingKernelTest::computePermutation(RoutingKernelTestParam const& para int const expert = tokenToExpertHostPtr[expandedIdx]; auto localExpertIdx = expert - param.localExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < param.numLocalExperts - && (localExpertIdx & param.localExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << param.localExpertsStrideLog2) - 1)) == 0; int const offsetWithinExpert = tokenToIdxInExpertHostPtr[expandedIdx]; int const offsetForExpert = expertScanCountsHostPtr[expert]; @@ -274,7 +274,7 @@ void RoutingKernelTest::verifyExpertRoutingIndices(RoutingKernelTestParam con std::set tokenIdx, tokenIdxTest; auto localExpertIdx = ie - param.localExpertsStartIdx; auto isLocalExpert = localExpertIdx >= 0 && localExpertIdx < param.numLocalExperts - && (localExpertIdx & param.localExpertsStrideLog2) == 0; + && (localExpertIdx & ((1 << param.localExpertsStrideLog2) - 1)) == 0; for (int it = 0; it < param.numTokens * param.topK; ++it) { @@ -357,12 +357,10 @@ void RoutingKernelTest::runTest(RoutingKernelTestParam const& param) } // Set seed to time-based seed resetToTimeBasedSeed(); - // Allocate buffers allocateBuffers(param); // Setup buffers setupBuffers(param); - // Call host function callHostFunction(param); if (param.useTopKAsInput) @@ -376,7 +374,6 @@ void RoutingKernelTest::runTest(RoutingKernelTestParam const& param) auto const workspaceSize = getDeviceWorkspaceSize(param); TensorPtr workspaceDevice = mBufferManager->gpu(ITensor::makeShape({static_cast(workspaceSize)}), nvinfer1::DataType::kINT8); - // Call tested function routing callTestedFunction(param, workspaceDevice); // Verify results diff --git a/tests/unittest/_torch/thop/serial/test_moe.py b/tests/unittest/_torch/thop/serial/test_moe.py index b28e695ecece..9c0d00bbe67a 100644 --- a/tests/unittest/_torch/thop/serial/test_moe.py +++ b/tests/unittest/_torch/thop/serial/test_moe.py @@ -1092,6 +1092,17 @@ class TestMoeFp4: "routing_method_type": RoutingMethodType.Renormalize }, id="RoutingRenormalize_qwen_next"), + pytest.param( + { + "num_experts": 2048, + "top_k": 32, + "n_groups": None, + "top_k_groups": None, + "routed_scaling": None, + "has_routing_bias": False, + "routing_method_type": RoutingMethodType.Renormalize + }, + id="RoutingRenormalize_large_experts"), ], ) def test_autotune(self, num_tokens, hidden_size, intermediate_size, @@ -1175,6 +1186,17 @@ def test_autotune_fp8_fp4(self, num_tokens, hidden_size, intermediate_size, "routing_method_type": RoutingMethodType.Renormalize }, id="RoutingRenormalize_qwen_next"), + pytest.param( + { + "num_experts": 2048, + "top_k": 32, + "n_groups": None, + "top_k_groups": None, + "routed_scaling": None, + "has_routing_bias": False, + "routing_method_type": RoutingMethodType.Renormalize + }, + id="RoutingRenormalize_large_experts"), ], ) @pytest.mark.parametrize("use_topk_as_input", [False, True], @@ -1328,7 +1350,7 @@ def run_moe_fp4_test(self, pytest.skip("https://nvbugs/5434352") assert top_k <= num_experts - assert top_k <= 22 + assert top_k <= 32 assert num_experts % 4 == 0 if use_topk_as_input: @@ -2003,7 +2025,7 @@ def test_moe_fp8_per_tensor_scale(num_tokens, hidden_size, intermediate_size, tile_tokens_dim = 8 assert top_k <= num_experts - assert top_k <= 8 + assert top_k <= 32 assert num_experts % 4 == 0 if are_groups_valid(top_k_groups, n_groups):