diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h index 888e04f2541a..97195a3daa3d 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingKernel.h @@ -195,7 +195,11 @@ struct KernelParams : public KernelParamsBase* mPtrTopKPacked = nullptr; @@ -302,7 +311,11 @@ struct KernelParams : public KernelParamsBase* mPtrTopKPacked = nullptr; @@ -329,6 +342,63 @@ void run(Data const& data, void* stream); } // namespace routingRenormalize -//////////////////////////////////////////////////////////////////////////////////////////////////// +////////////////////////////////////////////////////////////////////////////////////////////////// + +namespace routingMiniMax +{ + +////////////////////////////////////////////////////////////////////////////////////////////////// + +struct Data : public DataBase +{ + tg::Dtype mDtypeExpW{tg::Dtype::Fp32}; + + // Used by KernelParams::setKernelParams + void const* mPtrRoutingBias{nullptr}; + + // Used by KernelParams::setKernelParams + bool mNormTopkProb{true}; +}; + +////////////////////////////////////////////////////////////////////////////////////////////////// + +template +struct KernelParams : public KernelParamsBase +{ + using InputT = InputT_; + using OutputT = OutputT_; + + // Add missing static constexpr members for shared utilities + static constexpr int MaxNumExperts = MaxNumExperts_; + static constexpr bool isPow2 = isPow2_; + static constexpr bool UsePdl = UsePdl_; + static constexpr bool DoSoftmaxBeforeTopK = DoSoftmaxBeforeTopK_; + + PackedScoreIdx* mPtrTopKPacked = nullptr; + + OutputT const* mPtrRoutingBias = nullptr; + + int32_t mTopK = 0; + + bool mNormTopkProb = true; + + static KernelParams setKernelParams(Data const& data) + { + KernelParams params; + params.setBaseParams(data); + + params.mPtrTopKPacked = (PackedScoreIdx*) data.mPtrTopKPacked; + params.mPtrRoutingBias = static_cast(data.mPtrRoutingBias); + params.mNormTopkProb = data.mNormTopkProb; + params.mTopK = data.mTopK; + return params; + } +}; + +void run(Data const& data, void* stream); + +} // namespace routingMiniMax + +////////////////////////////////////////////////////////////////////////////////////////////////// } // namespace routing } // namespace moe::dev diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu new file mode 100644 index 000000000000..4430d6e4d684 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/RoutingMiniMax.cu @@ -0,0 +1,409 @@ +/* + * Copyright (c) 2022-2025, NVIDIA CORPORATION. All rights reserved. + * Contributed by Baseten.co + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "RoutingKernel.cuh" + +namespace moe::dev::routing +{ +namespace routingMiniMax +{ + +//////////////////////////////////////////////////////////////////////////////////////////////// + +static constexpr int NumExpertsLimit = 256; +static constexpr int MaxSupportedTopExperts = 8; + +//////////////////////////////////////////////////////////////////////////////////////////////// + +template +__global__ void routingMainKernel(KernelParams params) +{ + using OutputT = typename KernelParams::OutputT; + + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + static_assert(MaxNumExperts <= NumExpertsLimit, "MiniMax supports up to 256 experts."); + // MiniMax is configured for topK=8 (enforced at runtime in run()) + + // One token per block + int32_t const tokenIdx = blockIdx.x; + int32_t const expertIdx = threadIdx.x; + + // Cooperative groups + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + int32_t const laneIdx = cutlass::arch::LaneId(); + int32_t const warpIdx = __shfl_sync(0xffffffff, threadIdx.x / WarpSize, 0); + + // Shared memory for per-expert probabilities (DeepSeek style) + __shared__ float __attribute((aligned(128))) smemProb[MaxNumExperts]; + __shared__ float __attribute((aligned(128))) smemSel[MaxNumExperts]; + + // Invalid handling + static constexpr float kNegInf = float{-INFINITY}; + + // Load and compute per-expert scores for this token + bool const validExpert = expertIdx < params.mNumExperts; + float prob = 0.f; + float sel = kNegInf; + + if (validExpert) + { + // logit layout: token-major, contiguous experts + int64_t scoreIndex = int64_t{tokenIdx} * int64_t{params.mNumExperts} + int64_t{expertIdx}; + float logit = static_cast(params.mPtrScores[scoreIndex]); + + // MiniMax: sigmoid (not softmax) + prob = sigmoid_accurate(logit); + + float bias = static_cast(params.mPtrRoutingBias[expertIdx]); + sel = prob + bias; // selection score + } + + // Stage to shared so warp0 can index by expert id after topK + if (expertIdx < MaxNumExperts) + { + smemProb[expertIdx] = validExpert ? prob : 0.f; // invalid contributes 0 to renorm + smemSel[expertIdx] = validExpert ? sel : kNegInf; // invalid never selected + } + + __syncthreads(); + + // Only warp0 does the final expert selection (DeepSeek style) + if (warpIdx == 0) + { + // Each lane owns VecSize experts: expert = ii*WarpSize + laneIdx + static constexpr int VecSize = MaxNumExperts / WarpSize; // 256 -> 8 + static_assert(MaxNumExperts % WarpSize == 0, "MaxNumExperts must be multiple of 32."); + + float laneVals[VecSize]; + int32_t laneIdxs[VecSize]; + +#pragma unroll + for (int ii = 0; ii < VecSize; ++ii) + { + int e = ii * WarpSize + laneIdx; + laneIdxs[ii] = e; + laneVals[ii] = (e < params.mNumExperts) ? smemSel[e] : kNegInf; + } + + // TopK outputs + float topScores[MaxSupportedTopExperts]; + int32_t topExperts[MaxSupportedTopExperts]; + + // Reduce on selection scores (prob+bias) + topk::reduceTopK(warp, topScores, topExperts, laneVals, laneIdxs, kNegInf, params.mTopK); + + // Convert selection into final weights: + // final = prob (unbiased), optionally renormalized over topK + float w = 0.f; + int32_t chosenExpert = 0; + +#pragma unroll + for (int ii = 0; ii < MaxSupportedTopExperts; ++ii) + { + if (laneIdx == ii) + { + chosenExpert = topExperts[ii]; + w = (ii < params.mTopK && chosenExpert >= 0 && chosenExpert < params.mNumExperts) + ? smemProb[chosenExpert] + : 0.f; + } + } + + // Renormalize within topK if requested + float denom = 1.f; + if (params.mNormTopkProb) + { + float x = (laneIdx < params.mTopK) ? w : 0.f; + denom = cg::reduce(warp, x, cg::plus{}); + denom += 1e-20f; + } + + float finalW = (laneIdx < params.mTopK) ? (w / denom) : 0.f; + + // Write outputs + int32_t out = tokenIdx * params.mTopK + laneIdx; + + if (laneIdx < params.mTopK && params.mPtrTopKPacked != nullptr) + { + PackedScoreIdx packed{static_cast(finalW), static_cast(chosenExpert)}; + params.mPtrTopKPacked[out] = packed; + } + if (laneIdx < params.mTopK && params.mPtrTopKWeights != nullptr) + { + params.mPtrTopKWeights[out] = static_cast(finalW); + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////// + +// Cluster kernel removed for simplification - use histogram path for all token counts + +//////////////////////////////////////////////////////////////////////////////////////////////// + +template +__global__ void __launch_bounds__(KernelParams::MaxNumExperts) routingIndicesHistogramScoresKernel(KernelParams params) +{ + using OutputT = typename KernelParams::OutputT; + + 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); + + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition(block); + + static constexpr int MaxNumExperts = KernelParams::MaxNumExperts; + static constexpr int VecSize = MaxNumExperts / WarpSize; // 256->8 + + for (int tokenIdx = globalWarpIdx; tokenIdx < params.mNumTokens; tokenIdx += globalWarpStride) + { + // per-lane candidates + float laneVals[VecSize]; + int32_t laneIdxs[VecSize]; + +// each lane covers VecSize experts +#pragma unroll + for (int ii = 0; ii < VecSize; ++ii) + { + int e = ii * WarpSize + laneIdx; + laneIdxs[ii] = e; + + if (e < params.mNumExperts) + { + int64_t scoreIndex = int64_t{tokenIdx} * int64_t{params.mNumExperts} + int64_t{e}; + float logit = static_cast(params.mPtrScores[scoreIndex]); + float prob = sigmoid_accurate(logit); + float bias = static_cast(params.mPtrRoutingBias[e]); + laneVals[ii] = prob + bias; // selection + } + else + { + laneVals[ii] = float{-INFINITY}; + } + } + + float topSel[MaxSupportedTopExperts]; + int32_t topExp[MaxSupportedTopExperts]; + topk::reduceTopK(warp, topSel, topExp, laneVals, laneIdxs, float{-INFINITY}, params.mTopK); + + // Produce final weights from *unbiased prob* (sigmoid only), renormalized over topK. + // All 32 warp lanes must participate in cg::reduce, so compute prob for all lanes + // and mask non-topK lanes to 0. + float prob = 0.f; + int32_t e = -1; + if (laneIdx < params.mTopK) + { + e = topExp[laneIdx]; + if (e >= 0 && e < params.mNumExperts) + { + int64_t scoreIndex = int64_t{tokenIdx} * int64_t{params.mNumExperts} + int64_t{e}; + float logit = static_cast(params.mPtrScores[scoreIndex]); + prob = sigmoid_accurate(logit); + } + } + + // Renorm: all 32 lanes participate, non-topK lanes contribute 0 + float finalW = prob; + if (params.mNormTopkProb) + { + float x = (laneIdx < params.mTopK) ? prob : 0.f; + float denom = cg::reduce(warp, x, cg::plus{}) + 1e-20f; + finalW = (laneIdx < params.mTopK) ? (prob / denom) : 0.f; + } + + // Write outputs (only topK lanes) + if (laneIdx < params.mTopK) + { + PackedScoreIdx packed{static_cast(finalW), static_cast(e)}; + if (params.mPtrTopKPacked != nullptr) + { + params.mPtrTopKPacked[tokenIdx * params.mTopK + laneIdx] = packed; + } + if (params.mPtrTopKWeights != nullptr) + { + params.mPtrTopKWeights[tokenIdx * params.mTopK + laneIdx] = static_cast(finalW); + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////// + +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; + } +} + +// MiniMax-specific dispatch: InputT is always float (gate is float32). +// OutputT varies based on mDtypeExpW (bf16 for bias/weights, or float). +// DoSoftmaxBeforeTopK is always false (MiniMax uses sigmoid, not softmax). +#define LAUNCH_ROUTING_MINIMAX_IMPL(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, numExperts) \ + if (data.mDtypeExpW == tg::Dtype::Fp32) \ + { \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, float, numExperts, false), kernel, numBlocks, numThreads, \ + smemSize, stream); \ + } \ + else if (data.mDtypeExpW == tg::Dtype::Bfloat16) \ + { \ + LAUNCH_TILEN(data, coopLaunch, LAUNCH_ESC(float, __nv_bfloat16, numExperts, false), kernel, numBlocks, \ + numThreads, smemSize, stream); \ + } \ + else \ + { \ + TLLM_LOG_ERROR("Unsupported dtypeExpW"); \ + } + +#define LAUNCH_ROUTING_MINIMAX(data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream) \ + if (data.mNumExperts <= topk::MaxNumExpertsUnit) \ + { \ + LAUNCH_ROUTING_MINIMAX_IMPL( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, topk::MaxNumExpertsUnit); \ + } \ + else if (data.mNumExperts <= NumExpertsLimit) \ + { \ + LAUNCH_ROUTING_MINIMAX_IMPL( \ + data, coopLaunch, kernel, numBlocks, numThreads, smemSize, stream, NumExpertsLimit); \ + } \ + else \ + { \ + TLLM_LOG_ERROR("Unsupported numExperts"); \ + } + +void run(Data const& data, void* stream) +{ + // Create non-const alias for launch macros that may require non-const data + auto& d = const_cast(data); + + TLLM_CHECK_WITH_INFO(d.mPtrTopKPacked != nullptr || d.mPtrScores != nullptr || d.mPtrTopKIds != nullptr, + "Routing kernel requires at least one input parameter"); + + if (d.mPtrTopKIds != nullptr) + { + TLLM_CHECK_WITH_INFO(d.mPtrTopKWeights != nullptr, + "When mPtrTopKIds is provided, mPtrTopKWeights must also be provided for MiniMax routing."); + } + + // Permutation outputs required by grouped GEMM launch + TLLM_CHECK_WITH_INFO(d.mPtrPermutedIdxSize != nullptr && d.mPtrCtaIdxXyToBatchIdx != nullptr + && d.mPtrCtaIdxXyToMnLimit != nullptr && d.mPtrNumNonExitingCtas != nullptr, + "MiniMax routing expects permuted idx and grouped GEMM launch config buffers"); + + TLLM_CHECK_WITH_INFO(d.mTopK == 8, "MiniMax is configured for topK=8, got %d", d.mTopK); + TLLM_CHECK_WITH_INFO(d.mNumExperts == 256, "MiniMax is configured for 256 experts, got %d", d.mNumExperts); + + TLLM_CHECK_WITH_INFO(d.mNumExperts % 4 == 0, "Routing expects #experts multiple of 4, got %d", d.mNumExperts); + + // This "DeepSeek-like" path uses a block-per-token routingMainKernel when we need routing + int const numThreadsHist = getMaxNumExperts(d.mNumExperts); + + // Simplified: always use histogram path for permutation building + TLLM_CHECK_WITH_INFO((d.mPtrTopKPacked != nullptr || d.mPtrTopKIds != nullptr), + "MiniMax requires `mPtrTopKPacked` or `mPtrTopKIds` for permutation building."); + TLLM_CHECK_WITH_INFO( + d.mPtrExpertCounts != nullptr, "MiniMax requires `mPtrExpertCounts` for permutation building."); + + // 1) If TopK not provided, compute it - choose path based on token count + if (d.mPtrTopKIds == nullptr) + { + TLLM_CHECK_WITH_INFO(d.mPtrScores != nullptr, "If mPtrTopKIds is null, mPtrScores must be provided."); + TLLM_CHECK_WITH_INFO( + d.mPtrTopKPacked != nullptr, "MiniMax requires mPtrTopKPacked when computing topK from scores."); + + // Choose computation path based on token count to avoid double work + constexpr int SmallTokenThreshold = 256; // Same as MaxNumTokensSingleClusterScores + bool const useSmallTokenPath = d.mNumTokens <= SmallTokenThreshold; + + if (useSmallTokenPath) + { + // Small token count: use block-per-token routingMainKernel + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingMainKernel, + /*numBlocks=*/d.mNumTokens, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream); + } + // else: large token count - will use routingIndicesHistogramScoresKernel below + } + + // 2) Build permutation / CTA schedule (always use histogram path) + if (d.mPtrPermutedIdxSize != nullptr) + { + // Histogram + offsets path + int32_t const expandedIdxSize = d.mNumTokens * d.mTopK; + int32_t const histogramEltsPerBlock = 8 * numThreadsHist; + int32_t const offsetEltsPerBlock = NumEltsPerOffsetTilePerThread * numThreadsHist; + int32_t const maxNumBlocks = 1024; + + int const numBlocksHistogram + = std::min((expandedIdxSize + histogramEltsPerBlock - 1) / histogramEltsPerBlock, maxNumBlocks); + int const numBlocksOffsets + = std::min((expandedIdxSize + offsetEltsPerBlock - 1) / offsetEltsPerBlock, maxNumBlocks); + + // Always initialize expert counts first (avoid race conditions) + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingInitExpertCounts, + /*numBlocks=*/(2 * d.mNumExperts - 1) / numThreadsHist + 1, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream); + + // Only compute topK from scores if we didn't already do it in routingMainKernel + constexpr int SmallTokenThreshold = 256; + bool const usedSmallTokenPath = d.mNumTokens <= SmallTokenThreshold; + + if (d.mPtrScores != nullptr && d.mPtrTopKIds == nullptr && !usedSmallTokenPath) + { + // produce mPtrTopKPacked from scores (sigmoid+bias selection) for large token counts + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingIndicesHistogramScoresKernel, + /*numBlocks=*/maxNumBlocks, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream); + } + + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingIndicesHistogramKernel, + /*numBlocks=*/numBlocksHistogram, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream); + + LAUNCH_ROUTING_MINIMAX(d, + /*coopLaunch=*/false, routingIndicesOffsetsKernel, + /*numBlocks=*/numBlocksOffsets, + /*numThreads=*/numThreadsHist, + /*smemSize=*/0, stream); + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace routingMiniMax +} // namespace moe::dev::routing \ No newline at end of file diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu index d750cd8f41e5..3c44753dc20c 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu @@ -205,6 +205,39 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3 moe::dev::routing::routingRenormalize::run(routingData, stream); } + else if (routingMethodType == RoutingMethodType::MiniMax2) + { + moe::dev::routing::routingMiniMax::Data routingData; + routingData.mDtypeExpW = btg::Dtype::Bfloat16; + routingData.mUsePdl = true; + routingData.mNormTopkProb = true; + + routingData.mPtrTopKPacked = routingExpertIndexes; + routingData.mPtrExpertCounts = expertCountHistogram; + routingData.mPtrPermutedIdxSize = permutedIdxSize; + routingData.mPtrExpandedIdxToPermutedIdx = expandedIdxToPermutedIdx; + routingData.mPtrPermutedIdxToExpandedIdx = permutedIdxToExpandedIdx; + routingData.mPtrPermutedIdxToTokenIdx = permutedIdxToTokenIdx; + routingData.mPtrTopKWeights = expertWeights; + routingData.mPtrTopKIds = expertIds; + + routingData.mPtrCtaIdxXyToBatchIdx = ctaIdxXyToBatchIdx; + routingData.mPtrCtaIdxXyToMnLimit = ctaIdxXyToMnLimit; + routingData.mPtrNumNonExitingCtas = numNonExitingCtas; + + routingData.mPtrRoutingBias = routingBias; + routingData.mPtrScores = expertIds == nullptr ? routingLogits : nullptr; + routingData.mNumTokens = numTokens; + routingData.mNumExperts = numExperts; + routingData.mTopK = topK; + routingData.mPaddingLog2 = computeLog2(mTileTokensDim); + routingData.mTileTokensDim = mTileTokensDim; + routingData.mLocalExpertsStartIdx = localExpertOffset; + routingData.mLocalExpertsStrideLog2 = 0; + routingData.mNumLocalExperts = localNumExperts; + + moe::dev::routing::routingMiniMax::run(routingData, stream); + } else { TLLM_CHECK_WITH_INFO(false, "Unimplemented routing method %s of enum %d", diff --git a/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp b/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp index 2db4e2bf6c5b..172489c343de 100644 --- a/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp +++ b/cpp/tensorrt_llm/thop/fp8BlockScaleMoe.cpp @@ -63,7 +63,8 @@ at::Tensor run_fp8_block_scale_moe(at::optional const& routing_logit } else if (routing_logits.has_value()) { - if (static_cast(routing_method_type) == RoutingMethodType::DeepSeekV3) + if (static_cast(routing_method_type) == RoutingMethodType::DeepSeekV3 + || static_cast(routing_method_type) == RoutingMethodType::MiniMax2) { TORCH_CHECK(routing_logits.value().scalar_type() == at::ScalarType::Float, "routing_logits must be float"); } diff --git a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py index ae1b952030e4..63c871b6f5c0 100644 --- a/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/trtllm_gen_custom_ops.py @@ -6,6 +6,7 @@ from tensorrt_llm._torch.modules.fused_moe.routing import ( ROUTING_METHOD_TYPE_TO_CLASS, RoutingMethodType) +from tensorrt_llm._utils import get_sm_version from tensorrt_llm._torch.utils import (ActType_TrtllmGen, Fp4QuantizedTensor, fp4_utils, get_last_power_of_2_num_tokens_buckets, @@ -89,6 +90,13 @@ def prepare_dummy_topk_and_hook( lambda: torch.randn( num_experts, dtype=torch.bfloat16, device=hidden_states.device) }) + if routing_method_type == RoutingMethodType.MiniMax2: + routing_cls_kwargs.update({ + 'callable_e_score_correction_bias': + lambda: torch.randn( + num_experts, dtype=torch.bfloat16, device=hidden_states.device), + 'num_experts': num_experts, + }) routing_method = ROUTING_METHOD_TYPE_TO_CLASS[routing_method_type]( top_k=top_k, **routing_cls_kwargs) @@ -671,10 +679,28 @@ def _constrain_to_num_tokens(shapes: Tuple[torch.Size]) -> int: return num_tokens HS_SCALE_IDX = 3 - CONSTRAINED_HS_SCALE_DIM = 1 - constraint_hidden_states_scale = ConstraintSpec( - HS_SCALE_IDX, CONSTRAINED_HS_SCALE_DIM, _constrain_to_num_tokens) + if get_sm_version() >= 100: + # SM100+: fp8_quantize_1x128 returns 2D scale (blocked_n, num_tokens) + CONSTRAINED_HS_SCALE_DIM = 1 + constraint_hidden_states_scale = ConstraintSpec( + HS_SCALE_IDX, CONSTRAINED_HS_SCALE_DIM, + _constrain_to_num_tokens) + else: + # SM90: fp8_quantize_1x128 returns 1D scale with layout matching + # the fp8_quantize_1x128 custom op shape formula. + def _constrain_hs_scale_sm90( + shapes: Tuple[torch.Size]) -> int: + num_tokens = shapes[2][0] + hidden_size = shapes[2][1] + pad_m = fp4_utils.pad_up(num_tokens, 4) + blocked_n = (hidden_size + 127) // 128 + return fp4_utils.pad_up(pad_m * blocked_n * 4, 128) // 4 + + CONSTRAINED_HS_SCALE_DIM = 0 + constraint_hidden_states_scale = ConstraintSpec( + HS_SCALE_IDX, CONSTRAINED_HS_SCALE_DIM, + _constrain_hs_scale_sm90) ROUTER_LOGITS_IDX = 0 CONSTRAINED_RL_DIM = 0 diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm2.py b/tensorrt_llm/_torch/models/modeling_minimaxm2.py index 73cd480ee7cc..3380fdced93e 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm2.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm2.py @@ -59,9 +59,14 @@ def __init__( self.gate = Linear( self.hidden_dim, self.num_experts, bias=False, dtype=torch.float32, quant_config=None ) + self.moe_backend = model_config.moe_backend + if self.moe_backend == 'TRTLLM': + bias_dtype = torch.bfloat16 + else: + bias_dtype = torch.float32 self.e_score_correction_bias = nn.Parameter( - torch.empty((self.num_experts), dtype=torch.float32), requires_grad=False + torch.empty((self.num_experts), dtype=bias_dtype), requires_grad=False ) reduce_results = True diff --git a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py index 8026a7799b44..b286ccae74f0 100644 --- a/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py +++ b/tensorrt_llm/_torch/modules/fused_moe/fused_moe_trtllm_gen.py @@ -42,7 +42,7 @@ W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod, W4A8NVFP4FP8TRTLLMGenFusedMoEMethod, W4A16MXFP4TRTLLMGenFusedMoEMethod) # isort: on -from .routing import BaseMoeRoutingMethod, DeepSeekV3MoeRoutingMethod +from .routing import BaseMoeRoutingMethod, DeepSeekV3MoeRoutingMethod, MiniMaxM2MoeRoutingMethod class TRTLLMGenFusedMoE(MoE): @@ -522,6 +522,12 @@ def run_moe( n_group = self.routing_method.routing_impl.n_group topk_group = self.routing_method.routing_impl.topk_group routed_scaling_factor = self.routing_method.routing_impl.routed_scaling_factor + elif isinstance(self.routing_method, MiniMaxM2MoeRoutingMethod): + top_k = self.routing_method.top_k + routing_bias = self.routing_method.e_score_correction_bias + n_group = None + topk_group = None + routed_scaling_factor = None else: top_k = self.routing_method.top_k routing_bias = None @@ -544,7 +550,7 @@ def run_moe( if self.has_deepseek_fp8_block_scales: assert do_finalize, "fp8_block_scale_moe_runner does not support do_finalize=False" - # fp8_block_scale_moe_runner needs 2D shape for x_sf and only support SM100+ + # fp8_quantize_1x128 returns 2D x_sf on SM100+, 1D on SM90 if x_sf is None: x, x_sf = torch.ops.trtllm.fp8_quantize_1x128(x) @@ -1070,8 +1076,9 @@ def forward_fake( else: is_deepseek_v3_routing = isinstance(self.routing_method, DeepSeekV3MoeRoutingMethod) + is_minimax_routing = isinstance(self.routing_method, MiniMaxM2MoeRoutingMethod) top_k = self.routing_method.routing_impl.top_k if is_deepseek_v3_routing else self.routing_method.top_k - routing_bias = self.routing_method.e_score_correction_bias if is_deepseek_v3_routing else None + routing_bias = self.routing_method.e_score_correction_bias if (is_deepseek_v3_routing or is_minimax_routing) else None return fp4_block_scale_fake_output_without_finalize( x, self.num_experts, diff --git a/tests/unittest/_torch/modules/moe/moe_test_utils.py b/tests/unittest/_torch/modules/moe/moe_test_utils.py index b3987bafaeaa..b4416a0dd9b0 100644 --- a/tests/unittest/_torch/modules/moe/moe_test_utils.py +++ b/tests/unittest/_torch/modules/moe/moe_test_utils.py @@ -121,9 +121,10 @@ def should_skip_trtllm( return None # Routing method compatibility check (used by test_moe_module.py) - # TRTLLMGen C++ routing kernel (runner.cu) only implements: + # TRTLLMGen C++ routing kernel (runner.cu) implements: # - DeepSeekV3 (requires float32 routing_logits) # - Llama4 (requires top_k=1) + # - MiniMaxM2 (sigmoid activation, bias-added selection) # - Renormalize # - RenormalizeNaive # See: cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu:77-212 @@ -138,7 +139,6 @@ def should_skip_trtllm( # Routing methods NOT implemented in C++ kernel trtllm_unimplemented_routing = ( DefaultMoeRoutingMethod, # runner.cu:210 - "Unimplemented routing method" - MiniMaxM2MoeRoutingMethod, # runner.cu:210 - "Unimplemented routing method" ) if routing_method_cls in trtllm_unimplemented_routing: routing_name = routing_method_cls.__name__