Nvfp4 A16 Blackwell Moe Jit Runner#

class Nvfp4A16BlackwellMoeJitRunner#

Loads one compiled MoE JIT bundle (context-keyed module registry, shared between plugin clones) and launches its seven kernels through the driver API. Every launch takes enablePdl: when set, the kernel is launched with programmatic stream serialization (the kernels carry the griddepcontrol wait/trigger unconditionally). Launches never compile or query the device, so they are CUDA-graph capture safe once load() has run.

Public Functions

void load(Nvfp4A16BlackwellMoeJitKernel const &kernel)#
inline bool isLoaded() const noexcept#
inline Nvfp4A16BlackwellMoeJitKey const &getKey() const noexcept#
void launchRoute(
float const *logits,
float const *correctionBias,
int32_t numTokens,
bool normTopkProb,
float routedScalingFactor,
int32_t *topkIndices,
float *topkWeights,
bool enablePdl,
cudaStream_t stream
) const#

Warp-per-token sigmoid top-k routing (ungrouped contract).

void launchLayout(
int32_t const *topkIndices,
int32_t numSlots,
int32_t tokenTile,
int32_t *permutedIdx,
int32_t *tileGroupIdx,
int32_t *numValidTiles,
bool enablePdl,
cudaStream_t stream
) const#

Single-CTA expert-contiguous tile layout of the numSlots routed rows.

void launchGather(
void const *hiddenStates,
int32_t const *permutedIdx,
int32_t const *numValidTiles,
int32_t tokenTile,
int32_t numTokens,
int64_t maxRowsPadded,
void *permutedActivations,
void *output,
bool enablePdl,
cudaStream_t stream
) const#

Permuted-row gather plus output zeroing for the grouped GEMM path.

void launchFc1(
void const *hiddenStates,
int32_t const *topkIndices,
void const *qweights,
void const *blockScales,
float const *globalScales,
void *fc1Output,
float *partials,
int32_t numTokens,
bool enablePdl,
cudaStream_t stream
) const#

Decode FC1 (+ its split-K reduce when the key’s fc1SplitK > 1).

void launchFc2(
void const *fc1Output,
int32_t const *topkIndices,
float const *topkWeights,
void const *qweights,
void const *blockScales,
float const *globalScales,
void *output,
float *partials,
int32_t numTokens,
bool enablePdl,
cudaStream_t stream
) const#

Decode FC2 (+ its split-K reduce when the key’s fc2SplitK > 1).