Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
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 @@ -582,11 +582,10 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelPa
}
else
{
// If params.mPtrTopKIds != nullptr, we don't need to store the weights
scoreIdx = params.mPtrTopKPacked[expandedIdx];
idx = scoreIdx.idx;
if (params.mPtrTopKWeights != nullptr)
{
scoreIdx = params.mPtrTopKPacked[expandedIdx];
idx = scoreIdx.idx;
params.mPtrTopKWeights[expandedIdx] = static_cast<OutputT>(scoreIdx.score);
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -64,13 +64,14 @@ struct DataBase

// optional: if `nullptr`, it is not filled
// dim: [mNumTokens, mTopK]
// When mPtrTopKIds is provided, mPtrTopKWeights must be also provided as inputs.
// Otherwise, mPtrTopKWeights is the output scores of the topK experts.
// `routingRenormalize`: when `mPtrScores` is set (logits path), this holds per-slot masses (output).
// When `mPtrScores` is null and ids are precomputed (e.g. MoE thop), this may alias the precomputed
// top-k weight input together with `mPtrTopKIds`.
void* mPtrTopKWeights{nullptr};
// optional: if `nullptr`, it is not filled
// dim: [mNumTokens, mTopK]
// mPtrTopKIds[i] is the index of the expert for the i-th token in the top-k experts
// Together with mPtrTopKWeights, they form the top-k experts for each token
// `routingRenormalize`: output expert indices when computing from `mPtrScores`; input precomputed ids
// when `mPtrScores` is null (MoE thop supplies top-k ids/weights without logits).
int32_t* mPtrTopKIds{nullptr};

// optional: if `nullptr`, scores are used directly as input.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,11 +33,12 @@ void launchOffsetsKernel(Data const& data, int numBlocksOffsets, uint32_t numThr
void run(Data const& data, void* stream)
{
TLLM_CHECK_WITH_INFO(data.mPtrTopKPacked != nullptr || data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr,
"Routing kernel requires at least one input parameter");
if (data.mPtrTopKIds != nullptr)
"Routing kernel requires mPtrScores (logits), mPtrTopKPacked staging, and/or mPtrTopKIds");
if (data.mPtrScores != nullptr && data.mPtrTopKPacked == nullptr)
{
TLLM_CHECK_WITH_INFO(data.mPtrTopKWeights != nullptr,
"When mPtrTopKIds is provided, mPtrTopKWeights must also be provided for Renormalize routing.");
TLLM_CHECK_WITH_INFO(data.mPtrTopKIds != nullptr && data.mPtrTopKWeights != nullptr,
"Renormalize from logits (mPtrScores without mPtrTopKPacked) requires mPtrTopKIds and mPtrTopKWeights "
"outputs.");
}
TLLM_CHECK_WITH_INFO(data.mPtrPermutedIdxSize != nullptr && data.mPtrCtaIdxXyToBatchIdx != nullptr
&& data.mPtrCtaIdxXyToMnLimit != nullptr && data.mPtrNumNonExitingCtas != nullptr,
Expand All @@ -53,14 +54,14 @@ void run(Data const& data, void* stream)
bool const useSingleBlock = data.mNumTokens <= BlockKernelMaxNumTokens
|| (data.mNumTokens <= DynBlockKernelMaxNumTokens && data.mNumExperts <= DynBlockKernelMaxNumExperts);

bool const useSingleCluster = data.mNumTokens <= ((data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr)
? MaxNumTokensSingleClusterScores
: MaxNumTokensSingleCluster);
bool const useSingleCluster = data.mNumTokens
<= ((data.mPtrScores != nullptr || data.mPtrTopKIds != nullptr) ? MaxNumTokensSingleClusterScores
: MaxNumTokensSingleCluster);

if (!useSingleCluster && !useSingleBlock)
{
TLLM_CHECK_WITH_INFO((data.mPtrTopKPacked != nullptr || data.mPtrTopKIds != nullptr),
"When #tokens is large, `mPtrTopKPacked` or `mPtrTopKIds` is a required input.");
"When #tokens is large, `mPtrTopKPacked` staging or `mPtrTopKIds` input is required.");
TLLM_CHECK_WITH_INFO(
data.mPtrExpertCounts != nullptr, "When #tokens is large, `mPtrExpertCounts` is a required input.");
}
Expand Down Expand Up @@ -89,7 +90,7 @@ void run(Data const& data, void* stream)
int const numBlocksOffsets
= std::min((expandedIdxSize + offsetEltsPerBlock - 1) / offsetEltsPerBlock, maxNumBlocks);

if (data.mPtrScores != nullptr && data.mPtrTopKIds == nullptr)
if (data.mPtrScores != nullptr)
{
launchHistogramScoresKernel(data, maxNumBlocks, numThreadsHist, stream);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -69,35 +69,15 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelPa
}
#endif

if (params.mPtrTopKIds != nullptr)
if (params.mPtrScores != nullptr)
{
if (validToken)
{
if (laneIdx < params.mTopK)
{
auto expertIdx = params.mPtrTopKIds[warpIdx * params.mTopK + laneIdx];
if (expertIdx != -1)
{
int offset = warpIdx * MaxNumExperts + expertIdx;
smemKIdx[offset] = static_cast<int8_t>(laneIdx);
}
else
{
params.mPtrExpandedIdxToPermutedIdx[warpIdx * params.mTopK + laneIdx] = int32_t{-1};
}
}
}
}
else if (params.mPtrScores != nullptr)
{
// in this case, each warp represents a token
// Each warp represents a token: compute top-k from logits and write expert ids + weights.
BaseType score[VecSize];
int32_t idx[VecSize];

BaseType warpTopKScore[KernelParams::MaxNumTopExperts];
int32_t warpTopKExpertIdx[KernelParams::MaxNumTopExperts];

BaseType minScore = BaseType{-INFINITY};
if (validToken)
{
routingTopKExperts<BaseType, InputT, VecSize, KernelParams::MaxNumTopExperts,
Expand All @@ -108,12 +88,30 @@ __global__ void __launch_bounds__(KernelParams::MaxNumExperts <= 1024 ? KernelPa
{
int offset = warpIdx * MaxNumExperts + warpTopKExpertIdx[laneIdx];
smemKIdx[offset] = static_cast<int8_t>(laneIdx);
if (params.mPtrTopKWeights != nullptr)
params.mPtrTopKIds[warpIdx * params.mTopK + laneIdx] = warpTopKExpertIdx[laneIdx];
params.mPtrTopKWeights[warpIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]};
}
} // end if (validToken)
}
else if (params.mPtrTopKIds != nullptr)
{
// Precomputed top-k expert ids + weights (e.g. MoE thop); logits path is disabled.
if (validToken)
{
if (laneIdx < params.mTopK)
{
auto expertIdx = params.mPtrTopKIds[warpIdx * params.mTopK + laneIdx];
if (expertIdx != -1)
{
int offset = warpIdx * MaxNumExperts + expertIdx;
smemKIdx[offset] = static_cast<int8_t>(laneIdx);
}
else
{
params.mPtrTopKWeights[warpIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]};
params.mPtrExpandedIdxToPermutedIdx[warpIdx * params.mTopK + laneIdx] = int32_t{-1};
}
}
} // end if (validToken)
}
}
__syncthreads();

Expand Down Expand Up @@ -405,22 +403,7 @@ __global__ void routingIndicesDynBlockKernel(KernelParams params)
// ── Phase 1: TopK — one warp per token (loop only when numTokens > numWarps) ──
for (int tokenIdx = warpIdx; tokenIdx < params.mNumTokens; tokenIdx += numWarps)
{
if (params.mPtrTopKIds != nullptr)
{
if (laneIdx < params.mTopK)
{
auto expertIdx = params.mPtrTopKIds[tokenIdx * params.mTopK + laneIdx];
if (expertIdx > -1 && expertIdx < params.mNumExperts)
{
smemKIdx[tokenIdx * MaxNumExperts + expertIdx] = static_cast<int8_t>(laneIdx);
}
else
{
params.mPtrExpandedIdxToPermutedIdx[tokenIdx * params.mTopK + laneIdx] = int32_t{-1};
}
}
}
else if (params.mPtrScores != nullptr)
if (params.mPtrScores != nullptr)
{
BaseType score[VecSize];
int32_t idx[VecSize];
Expand All @@ -435,9 +418,22 @@ __global__ void routingIndicesDynBlockKernel(KernelParams params)
if (laneIdx < params.mTopK)
{
smemKIdx[tokenIdx * MaxNumExperts + warpTopKExpertIdx[laneIdx]] = static_cast<int8_t>(laneIdx);
if (params.mPtrTopKWeights != nullptr)
params.mPtrTopKIds[tokenIdx * params.mTopK + laneIdx] = warpTopKExpertIdx[laneIdx];
params.mPtrTopKWeights[tokenIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]};
}
}
else if (params.mPtrTopKIds != nullptr)
{
if (laneIdx < params.mTopK)
{
auto expertIdx = params.mPtrTopKIds[tokenIdx * params.mTopK + laneIdx];
if (expertIdx > -1 && expertIdx < params.mNumExperts)
{
smemKIdx[tokenIdx * MaxNumExperts + expertIdx] = static_cast<int8_t>(laneIdx);
}
else
{
params.mPtrTopKWeights[tokenIdx * params.mTopK + laneIdx] = OutputT{warpTopKScore[laneIdx]};
params.mPtrExpandedIdxToPermutedIdx[tokenIdx * params.mTopK + laneIdx] = int32_t{-1};
}
}
}
Expand All @@ -448,6 +444,10 @@ __global__ void routingIndicesDynBlockKernel(KernelParams params)
auto expandedIdx = tokenIdx * params.mTopK + laneIdx;
auto scoreIdx = params.mPtrTopKPacked[expandedIdx];
smemKIdx[tokenIdx * MaxNumExperts + static_cast<int>(scoreIdx.idx)] = static_cast<int8_t>(laneIdx);
if (params.mPtrTopKIds != nullptr)
{
params.mPtrTopKIds[expandedIdx] = static_cast<int32_t>(scoreIdx.idx);
}
if (params.mPtrTopKWeights != nullptr)
{
params.mPtrTopKWeights[expandedIdx] = static_cast<OutputT>(scoreIdx.score);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -172,7 +172,8 @@ void Runner::run(void* routingLogits, void* routingBias, int32_t numTokens, int3
routingData.mDoSoftmaxBeforeTopK = routingMethodType == RoutingMethodType::RenormalizeNaive;
routingData.mNormTopkProb = routingMethodType == RoutingMethodType::RenormalizeNaive;

// Pass-through raw pointer; kernels will cast to the proper InputT based on routing method
// When precomputed top-k ids/weights are supplied (thop), expertIds is non-null and logits are ignored.
// When computing from logits (e.g. TensorRT plugin), expertIds is null and routingLogits drives TopK.
routingData.mPtrScores = expertIds == nullptr ? routingLogits : nullptr;
//
// Outputs
Expand Down