Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
19 commits
Select commit Hold shift + click to select a range
0971d68
[None][feat] Cherry-pick #11143 to Main
leslie-fang25 Jun 12, 2026
289fa1f
Add E2E testing and fix
leslie-fang25 Jun 12, 2026
4ff0d64
[None][feat] Enhance the check to enable this fusion path
leslie-fang25 Jun 12, 2026
9753441
[None][test] bench_moe: shared-expert fusion support
leslie-fang25 Jun 17, 2026
a9b2d25
[None][test] bench_moe: route unfused shared GatedMLP FP8 GEMM via cu…
leslie-fang25 Jun 22, 2026
a8cd40d
[None][fix] Restrict fused shared-expert MoE to tileN>=32 to avoid sm…
leslie-fang25 Jul 11, 2026
9b0a208
[None][chore] Fix yapf formatting
leslie-fang25 Jul 11, 2026
c7345fa
[None][feat] Make TRTLLM-Gen shared-expert fusion opt-in via TLLM_MOE…
leslie-fang25 Jul 13, 2026
f4d6f10
[None][test] Add TRTLLM-Gen shared-expert fusion coverage to modules/…
leslie-fang25 Jul 13, 2026
9432727
[None][test] Drop fusion additions to deprecated thop test_moe.py (co…
leslie-fang25 Jul 13, 2026
73b1a11
[None][test] Drop temporary in-tree fusion E2E test (kept as standalo…
leslie-fang25 Jul 13, 2026
dbc1c1d
[None][chore] Drop unconsumed EP token-sharding plumbing for fused sh…
leslie-fang25 Jul 15, 2026
2e059e9
[None][fix] bench_moe: drop stale cuda-graph max_num_tokens special-case
leslie-fang25 Jul 15, 2026
ff3227e
[None][chore] bench_moe: print traceback on build failures
leslie-fang25 Jul 15, 2026
c7b31f0
[None][fix] Guard quant_config None in TRTLLM-Gen MoE create_weights/…
leslie-fang25 Jul 24, 2026
1756c1a
[None][fix] Align AutoDeploy fp8_block_scale_moe_runner calls with ne…
leslie-fang25 Jul 24, 2026
c5a8e93
[None][fix] Default-init fused-shared-expert routing Data fields
leslie-fang25 Jul 24, 2026
b379112
[None][chore] Address shared-expert fusion review comments
leslie-fang25 Jul 24, 2026
5016db4
[None][chore] Address shared-expert fusion review comments (C++)
leslie-fang25 Jul 24, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -380,17 +380,31 @@ __global__ void routingMainKernel(KernelParams params)
auto finalScore = OutputT{scoreNorm * params.mRouteScale / redNorm};

// write expert idx out already
auto idxTopK = blockIdx.x * params.mTopK + laneIdx;
auto idxTopK = blockIdx.x * params.mTotalExpertsPerToken + laneIdx;
auto idxShared = blockIdx.x * params.mTotalExpertsPerToken + params.mTopK + laneIdx;
if (laneIdx < params.mTopK && params.mPtrTopKPacked != nullptr)
{
PackedScoreIdx<OutputT> packedScore{static_cast<OutputT>(finalScore), static_cast<int16_t>(expertIdx)};
params.mPtrTopKPacked[idxTopK] = packedScore;
}

if (laneIdx < params.mNumFusedSharedExperts && params.mPtrTopKPacked != nullptr)
{
PackedScoreIdx<OutputT> packedScore{
static_cast<OutputT>(1.0F), static_cast<int16_t>(params.mNumExperts + laneIdx)};
params.mPtrTopKPacked[idxShared] = packedScore;
}

if (laneIdx < params.mTopK && params.mPtrTopKWeights != nullptr && params.mPtrTopKIds == nullptr)
{
params.mPtrTopKWeights[idxTopK] = finalScore;
}

// Write score of 1.0 for shared expert if enabled
if (laneIdx < params.mNumFusedSharedExperts && params.mPtrTopKWeights != nullptr)
{
params.mPtrTopKWeights[idxShared] = static_cast<OutputT>(1.0F);
}
}
}
}
Expand Down Expand Up @@ -551,12 +565,32 @@ void run(Data& data, void* stream)
}

int const numBlocks = data.mNumTokens;
int const numThreadsHist = getMaxNumExperts(data.mNumExperts);
// Derive the per-token entry count (topK routed + fused shared) here so callers
// that construct Data directly (unit tests) need not set it; the main kernel uses
// it as the output stride before mTopK is expanded below.
data.mTotalExpertsPerToken = data.mTopK + data.mNumFusedSharedExperts;
// Account for fused shared experts (appended after the routed experts) when sizing the histogram.
int const numThreadsHist = getMaxNumExperts(data.mNumExperts + data.mNumFusedSharedExperts);
static int const smMajor = tensorrt_llm::common::getSMVersion() / 10;
// Step 1: Run DeepSeek-specific topK computation (writes to mPtrTopKPacked)

TLLM_CHECK_WITH_INFO(data.mNumFusedSharedExperts <= WarpSize,
"Number of fused shared experts (%d) must be less than warp size (%d).", data.mNumFusedSharedExperts, WarpSize);

// Step 1: Run DeepSeek-specific topK computation (writes to mPtrTopKPacked).
// When fused shared experts are enabled, the main kernel also appends them (index mNumExperts + i, weight 1.0),
// so it must run before mNumExperts/mTopK are expanded below.
int const numThreadsMain = max(data.mNumExpertGroups * WarpSize, getMaxNumExperts(data.mNumExperts));
launchMainKernel(data, numBlocks, numThreadsMain, stream);

// Fused shared experts are appended after the routed experts; expand the expert/topK counts so the
// permutation pipeline below (which reads data.mNumExperts / data.mTopK) accounts for them.
if (data.mNumFusedSharedExperts > 0)
{
data.mNumExperts += data.mNumFusedSharedExperts;
data.mTopK += data.mNumFusedSharedExperts;
data.mNumLocalExperts += data.mNumFusedSharedExperts;
}

// Step 2: Permutation pipeline (reads from mPtrTopKPacked written by step 1)
if (data.mPtrPermutedIdxSize != nullptr)
{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,12 @@ struct DataBase
int32_t mLocalExpertsStartIdx;
int32_t mLocalExpertsStrideLog2;
int32_t mNumLocalExperts;

/// For fused shared expert. Default to 0 (no fusion) so callers that
/// construct Data directly (e.g. the routing kernel unit tests) stay valid;
/// mTotalExpertsPerToken is derived in routingDeepSeek::run() from mTopK.
int32_t mNumFusedSharedExperts{0};
int32_t mTotalExpertsPerToken{0};
};

template <typename InputT_, typename OutputT_, int MaxNumExperts_, int MaxNumTopExperts_>
Expand Down Expand Up @@ -145,6 +151,9 @@ struct KernelParamsBase
int32_t mLocalExpertsStrideLog2 = 0;
int32_t mNumLocalExperts = 0;

int32_t mNumFusedSharedExperts = 0;
int32_t mTotalExpertsPerToken = 0;

// Public initialization function - make it a template to accept different Data types
template <typename DataType>
void setBaseParams(DataType const& data)
Expand All @@ -171,6 +180,9 @@ struct KernelParamsBase
mLocalExpertsStartIdx = data.mLocalExpertsStartIdx;
mLocalExpertsStrideLog2 = data.mLocalExpertsStrideLog2;
mNumLocalExperts = data.mNumLocalExperts;

mNumFusedSharedExperts = data.mNumFusedSharedExperts;
mTotalExpertsPerToken = data.mTotalExpertsPerToken;
}
};

Expand Down
62 changes: 44 additions & 18 deletions cpp/tensorrt_llm/kernels/trtllmGenKernels/blockScaleMoe/runner.cu
Original file line number Diff line number Diff line change
Expand Up @@ -61,17 +61,20 @@ Runner::Runner(int32_t tileTokensDim, int32_t clusterSizeInBatchDim)
}

void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int32_t numExperts, int32_t topK,
int32_t nGroup, int32_t topkGroup, int32_t localExpertOffset, int32_t localNumExperts, float routedScalingFactor,
int32_t* routingExpertIndexes, int32_t* expertCountHistogram, int32_t* permutedIdxSize,
int32_t numFusedSharedExpert, int32_t nGroup, int32_t topkGroup, int32_t localExpertOffset, int32_t localNumExperts,
float routedScalingFactor, int32_t* routingExpertIndexes, int32_t* expertCountHistogram, int32_t* permutedIdxSize,
int32_t* expandedIdxToPermutedIdx, int32_t* permutedIdxToExpandedIdx, int32_t* permutedIdxToTokenIdx,
void* expertWeights, int32_t* expertIds, int32_t* numTokensPerExpert, int32_t* ctaIdxXyToBatchIdx,
int32_t* ctaIdxXyToMnLimit, int32_t* numNonExitingCtas, btg::Dtype dtypeElt, bool useRoutingScalesOnInput,
bool useDeepSeekFp8, RoutingMethodType routingMethodType, cudaStream_t stream, btg::Dtype dtypeRoutingLogits,
btg::Dtype dtypeRoutingBias)
{
if (routingMethodType == RoutingMethodType::DeepSeekV3 && nGroup <= 1)
if (routingMethodType == RoutingMethodType::DeepSeekV3 && nGroup <= 1 && numFusedSharedExpert == 0)
{
// DeepSeek no-groups case: use routingCustom with SigmoidBias preprocess
// DeepSeek no-groups case: use routingCustom with SigmoidBias preprocess.
// NOTE: routingCustom does not implement fused shared experts; when fusion is
// requested we fall through to the routingDeepSeek path below (which handles it
// for both grouped and non-grouped routing).
// and ScaledSumNormalize postprocess. This is more efficient than the full DeepSeek
// kernel because it uses the warp-level routingTopKExperts flow.
moe::dev::routing::routingCustom::Data routingData;
Expand Down Expand Up @@ -230,10 +233,18 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3
{
TLLM_CHECK_WITH_INFO(topK <= 22, "For DeepSeek routing method, must have topK <= 22");
TLLM_CHECK_WITH_INFO(topkGroup <= 4, "For DeepSeek routing method, must have topkGroup <= 4");
// Fused shared experts assume the full expert set is resident on this rank
// (no expert parallelism): the appended experts live at global indices
// [numExperts, numExperts + numFusedSharedExpert) on every rank.
TLLM_CHECK_WITH_INFO(numFusedSharedExpert == 0 || (localExpertOffset == 0 && localNumExperts == numExperts),
"Fused shared experts do not support expert parallelism yet.");

moe::dev::routing::routingDeepSeek::Data routingData;
routingData.mDtypeOutput = btg::Dtype::Bfloat16;
routingData.mUsePdl = tensorrt_llm::common::getEnvEnablePDL();

int32_t const totalExpertsPerToken = topK + numFusedSharedExpert;

// output:
routingData.mPtrTopKPacked = routingExpertIndexes;
routingData.mPtrExpertCounts = expertCountHistogram;
Expand All @@ -255,9 +266,11 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3
routingData.mPtrTopKIds = expertIds;
routingData.mNumTokens = numTokens;
routingData.mNumExperts = numExperts;
routingData.mNumFusedSharedExperts = numFusedSharedExpert;
routingData.mNumExpertGroups = nGroup;
routingData.mNumLimitedGroups = topkGroup;
routingData.mTopK = topK;
routingData.mTotalExpertsPerToken = totalExpertsPerToken;
routingData.mPaddingLog2 = computeLog2(mTileTokensDim);
routingData.mTileTokensDim = mTileTokensDim;
routingData.mLocalExpertsStartIdx = localExpertOffset;
Expand All @@ -274,6 +287,8 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3
{
TLLM_LOG_WARNING("For Llama routing method, nGroup/topkGroup is ignored, got %d/%d.", nGroup, topkGroup);
}
TLLM_CHECK_WITH_INFO(numFusedSharedExpert == 0, "Llama routing method does not support fusing shared expert");

moe::dev::routing::routingLlama4::Data routingData;
routingData.mDtypeOutput = btg::Dtype::Bfloat16;
routingData.mUsePdl = tensorrt_llm::common::getEnvEnablePDL();
Expand Down Expand Up @@ -318,6 +333,9 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3
else if (routingMethodType == RoutingMethodType::Renormalize
|| routingMethodType == RoutingMethodType::RenormalizeNaive || routingMethodType == RoutingMethodType::Default)
{
TLLM_CHECK_WITH_INFO(
numFusedSharedExpert == 0, "Renormalize routing method does not support fusing shared expert");

moe::dev::routing::routingCustom::Data routingData;

//
Expand Down Expand Up @@ -649,6 +667,9 @@ void Runner::setOpsData(MoERunnerArgs const& args, MoEWorkspace const& workspace
moe::dev::convertsf::Data& convertSfData, moe::dev::activation::Data& activationData,
moe::dev::finalize::Data& finalizeData)
{
int32_t const totalNumExperts = args.num_experts + args.num_fused_shared_experts;
int32_t const totalExpertsPerToken = args.top_k + args.num_fused_shared_experts;

// Setup sf conversion data if needed
convertSfData.inSfPtr = args.hidden_states_scale;
convertSfData.outSfPtr = workspace.hidden_states_scale_linear;
Expand All @@ -667,7 +688,7 @@ void Runner::setOpsData(MoERunnerArgs const& args, MoEWorkspace const& workspace
activationData.inDqSfsPtr = workspace.gemm1_output_scale;
activationData.outDqSfsPtr = workspace.activation_output_scale;
activationData.innerDim = args.intermediate_size * (mActType == ActType::SwiGlu ? 2 : 1);
activationData.topK = args.top_k;
activationData.topK = totalExpertsPerToken; // TODO Rename topK in activation data struct
activationData.numTokens = args.num_tokens;
activationData.expandedIdxToPermutedIdx = workspace.expanded_idx_to_permuted_idx;
// For DeepSeek FP8 the activation runs as a separate kernel rather than
Expand Down Expand Up @@ -699,8 +720,8 @@ void Runner::setOpsData(MoERunnerArgs const& args, MoEWorkspace const& workspace
}
finalizeData.expandedIdxToPermutedIdx = workspace.expanded_idx_to_permuted_idx;
finalizeData.numTokens = args.num_tokens;
finalizeData.numExperts = args.num_experts;
finalizeData.topK = args.top_k;
finalizeData.numExperts = totalNumExperts; // TODO Is this used?
finalizeData.topK = totalExpertsPerToken; // TODO Rename topK in finalize data struct
// We want to fuse unpadding into the finalize kernel, so we need to use the output hidden size.
finalizeData.hiddenDim = args.valid_hidden_size.value_or(args.hidden_size);
finalizeData.hiddenDimPadded = args.output_hidden_size.value_or(args.hidden_size);
Expand All @@ -710,12 +731,15 @@ void Runner::setOpsData(MoERunnerArgs const& args, MoEWorkspace const& workspace

std::tuple<int32_t, int32_t> Runner::getWorkspaceSizeInBytes(MoERunnerArgs const& args, int64_t configIndex) const
{
int32_t const totalLocalExperts = args.local_num_experts + args.num_fused_shared_experts;
int32_t const totalExpertsPerToken = args.top_k + args.num_fused_shared_experts;

auto const& config = mPassingConfigs[configIndex];

auto workspace_size_fc1 = static_cast<int32_t>(mPermuteGemm1.getWorkspaceSizeInBytes(args.top_k, args.hidden_size,
args.intermediate_size, args.local_num_experts, args.num_tokens, config.gemm1Config));
auto workspace_size_fc2 = static_cast<int32_t>(mGemm2.getWorkspaceSizeInBytes(args.top_k, args.hidden_size,
args.intermediate_size, args.local_num_experts, args.num_tokens, config.gemm2Config));
auto workspace_size_fc1 = static_cast<int32_t>(mPermuteGemm1.getWorkspaceSizeInBytes(totalExpertsPerToken,
args.hidden_size, args.intermediate_size, totalLocalExperts, args.num_tokens, config.gemm1Config));
auto workspace_size_fc2 = static_cast<int32_t>(mGemm2.getWorkspaceSizeInBytes(totalExpertsPerToken,
args.hidden_size, args.intermediate_size, totalLocalExperts, args.num_tokens, config.gemm2Config));
return std::make_tuple(workspace_size_fc1, workspace_size_fc2);
}

Expand Down Expand Up @@ -750,7 +774,6 @@ std::vector<int64_t> Runner::getValidConfigIndices(int32_t topK, int32_t hiddenS
int64_t Runner::getDefaultValidConfigIndex(int32_t topK, int32_t hiddenSize, int32_t intermediateSize,
int32_t numLocalExperts, int32_t numTokens, int32_t validHiddenSize, int32_t validIntermediateSize) const
{

int32_t indexGemm1 = mPermuteGemm1.getDefaultValidConfigIndex(
topK, hiddenSize, intermediateSize, numLocalExperts, numTokens, validHiddenSize, validIntermediateSize);
int32_t indexGemm2 = mGemm2.getDefaultValidConfigIndex(
Expand All @@ -773,14 +796,17 @@ void Runner::run(
sync_check_cuda_error(stream);
setOpsData(args, workspace, convertSfData, activationData, finalizeData);

int32_t const totalLocalExperts = args.local_num_experts + args.num_fused_shared_experts;
int32_t const totalExpertsPerToken = args.top_k + args.num_fused_shared_experts;

void* hidden_states_scale_linear{args.hidden_states_scale};

auto const& config = mPassingConfigs[configIndex];

mPermuteGemm1.run(args.hidden_states, hidden_states_scale_linear, args.gemm1_weights, args.gemm1_weights_scale,
workspace.expert_weights, args.output1_scales_scalar, args.output1_scales_gate_scalar, args.gemm1_bias,
args.gemm1_alpha, args.gemm1_beta, args.gemm1_clamp_limit, workspace.gemm1_output, workspace.gemm1_output_scale,
args.top_k, args.hidden_size, args.intermediate_size, args.local_num_experts, args.num_tokens,
totalExpertsPerToken, args.hidden_size, args.intermediate_size, totalLocalExperts, args.num_tokens,
workspace.permuted_idx_to_token_idx, workspace.num_non_exiting_ctas, workspace.total_num_padded_tokens,
workspace.cta_idx_xy_to_batch_idx, workspace.cta_idx_xy_to_mn_limit, workspace.bmm1_workspace,
args.mUseRoutingScalesOnInput, device, stream, config.gemm1Config,
Expand All @@ -801,11 +827,11 @@ void Runner::run(

// Run gemm2
mGemm2.run(gemm2_input, gemm2_input_scale, args.gemm2_weights, args.gemm2_weights_scale, args.output2_scales_scalar,
args.gemm2_bias, workspace.gemm2_output, workspace.gemm2_output_scale, args.top_k,
args.output_hidden_size.value_or(args.hidden_size), args.intermediate_size, args.local_num_experts,
args.num_tokens, workspace.num_non_exiting_ctas, workspace.total_num_padded_tokens,
workspace.cta_idx_xy_to_batch_idx, workspace.cta_idx_xy_to_mn_limit, workspace.bmm2_workspace, device, stream,
config.gemm2Config, args.valid_hidden_size.value_or(args.hidden_size),
args.gemm2_bias, workspace.gemm2_output, workspace.gemm2_output_scale, totalExpertsPerToken,
args.output_hidden_size.value_or(args.hidden_size), args.intermediate_size, totalLocalExperts, args.num_tokens,
workspace.num_non_exiting_ctas, workspace.total_num_padded_tokens, workspace.cta_idx_xy_to_batch_idx,
workspace.cta_idx_xy_to_mn_limit, workspace.bmm2_workspace, device, stream, config.gemm2Config,
args.valid_hidden_size.value_or(args.hidden_size),
args.valid_intermediate_size.value_or(args.intermediate_size));

// Run finalize
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -165,13 +165,13 @@ class Runner
explicit Runner(int32_t tileTokensDim, int32_t clusterSizeInBatchDim = 1);

void run(void* routingLogits, void* routingBias, int32_t numTokens, int32_t numExperts, int32_t topK,
int32_t nGroups, int32_t topkGroups, int32_t localExpertOffset, int32_t localNumExperts,
float routedScalingFactor, int32_t* routingExpertIndexes, int32_t* expertCountHistogram,
int32_t* permutedIdxSize, int32_t* expandedIdxToPermutedIdx, int32_t* permutedIdxToExpandedIdx,
int32_t* permutedIdxToTokenIdx, void* expertWeights, int32_t* expertIds, int32_t* numTokensPerExpert,
int32_t* ctaIdxXyToBatchIdx, int32_t* ctaIdxXyToMnLimit, int32_t* numNonExitingCtas,
batchedGemm::trtllm::gen::Dtype dtypeElt, bool useRoutingScalesOnInput, bool useDeepSeekFp8,
RoutingMethodType routingMethodType, cudaStream_t stream,
int32_t numFusedSharedExpert, int32_t nGroups, int32_t topkGroups, int32_t localExpertOffset,
int32_t localNumExperts, float routedScalingFactor, int32_t* routingExpertIndexes,
int32_t* expertCountHistogram, int32_t* permutedIdxSize, int32_t* expandedIdxToPermutedIdx,
int32_t* permutedIdxToExpandedIdx, int32_t* permutedIdxToTokenIdx, void* expertWeights, int32_t* expertIds,
int32_t* numTokensPerExpert, int32_t* ctaIdxXyToBatchIdx, int32_t* ctaIdxXyToMnLimit,
int32_t* numNonExitingCtas, batchedGemm::trtllm::gen::Dtype dtypeElt, bool useRoutingScalesOnInput,
bool useDeepSeekFp8, RoutingMethodType routingMethodType, cudaStream_t stream,
batchedGemm::trtllm::gen::Dtype dtypeRoutingLogits = batchedGemm::trtllm::gen::Dtype::Bfloat16,
batchedGemm::trtllm::gen::Dtype dtypeRoutingBias = batchedGemm::trtllm::gen::Dtype::Bfloat16);

Expand Down Expand Up @@ -297,6 +297,7 @@ struct MoERunnerArgs

int32_t num_tokens{0};
int32_t num_experts{0};
int32_t num_fused_shared_experts{0};
// Hidden dimension input of MoE block. It might be padded.
int32_t hidden_size{0};
// Hidden dimension output of MoE block. It might be padded.
Expand Down
7 changes: 4 additions & 3 deletions cpp/tensorrt_llm/thop/cuteDslMoeUtilsOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -79,9 +79,10 @@ std::vector<torch::Tensor> moe_topk_sort_impl(torch::optional<torch::Tensor> con
auto const dtypeRoutingLogits = routing_logits.has_value()
? (routing_logits->scalar_type() == at::ScalarType::Float ? btg::Dtype::Fp32 : btg::Dtype::Bfloat16)
: btg::Dtype::Bfloat16;
routing_runner.run(routing_logits_ptr, routing_bias_ptr, num_tokens, num_experts, top_k, n_group.value_or(0),
topk_group.value_or(0), local_expert_offset, local_num_experts, routed_scaling_factor.value_or(1.0),
expert_indexes.data_ptr<int>(), expert_count_histogram.data_ptr<int>(), total_num_padded_tokens.data_ptr<int>(),
routing_runner.run(routing_logits_ptr, routing_bias_ptr, num_tokens, num_experts, top_k,
/* num_fused_shared_expert */ 0, n_group.value_or(0), topk_group.value_or(0), local_expert_offset,
local_num_experts, routed_scaling_factor.value_or(1.0), expert_indexes.data_ptr<int>(),
expert_count_histogram.data_ptr<int>(), total_num_padded_tokens.data_ptr<int>(),
expanded_idx_to_permuted_idx.data_ptr<int>(), permuted_idx_to_expanded_idx.data_ptr<int>(),
nullptr /*permuted_idx_to_token_idx.data_ptr<int>()*/, token_final_scales_ptr, token_selected_experts_ptr,
num_tokens_per_expert.data_ptr<int>(), tile_idx_to_expert_idx.data_ptr<int>(),
Expand Down
Loading
Loading