Cute Dsl FMHA Runner#

class CuteDslFMHARunner#

Unified runner for CuTe DSL compiled FMHA kernels (Blackwell SM100/101/110).

Supports two execution modes via AOT-compiled kernel variants:

  1. LLM prefill/chunked-prefill: batched Q [B,S_q,H_q,D] + combined KV cache [B,2,H_kv,Cap,D] with causal masking and optional sliding window.

  2. ViT: packed varlen separate Q/K/V [total_S,H,D] with cu_seqlens [B+1] for ragged batching, bidirectional attention (no causal mask).

Each mode has its own kernel modules and run() overload.

Public Functions

CuteDslFMHARunner(
int32_t numQHeads,
int32_t numKVHeads,
int32_t headDim,
int32_t batchSize = 0,
int32_t seqLenQ = 0,
int32_t kvCacheCapacity = 0
)#
~CuteDslFMHARunner() = default#
CuteDslFMHARunner(CuteDslFMHARunner const&) = delete#
CuteDslFMHARunner &operator=(CuteDslFMHARunner const&) = delete#
bool preflightLlm(
cudaStream_t stream,
int32_t slidingWindowSize = INT_MAX,
bool fp8Input = false,
float skipSoftmaxScaleFactor = 0.0F
)#

Ensures the exact dense LLM variant selected by run() is loaded.

bool preflightPaged(
cudaStream_t stream,
int32_t slidingWindowSize = INT_MAX,
bool fp8Input = false,
bool isCausal = true,
float skipSoftmaxScaleFactor = 0.0F,
bool useBidirectional = false
)#

Ensures the exact paged LLM variant selected by runPaged() is loaded.

bool preflightPagedRagged(
cudaStream_t stream,
int32_t slidingWindowSize = INT_MAX,
nvinfer1::DataType inputDataType = nvinfer1::DataType::kHALF
)#

Ensures the exact ragged paged LLM variant selected by runPagedRagged() is loaded.

bool preflightViT(cudaStream_t stream)#

Ensures the exact packed ViT variant selected by run() is loaded.

bool run(
void const *qPtr,
void const *kvPtr,
void *oPtr,
int32_t const *cuKVSeqLens,
cudaStream_t stream,
float attentionScale,
int32_t slidingWindowSize = INT_MAX,
bool fp8Input = false,
float qScale = 1.0F,
float kScale = 1.0F,
float vScale = 1.0F,
float skipSoftmaxScaleFactor = 0.0F
)#

LLM FMHA: batched Q + combined KV cache with causal masking.

Output is always FP16. Selects kernel variant based on fp8Input:

  • fp8Input=false → FP16 kernels (all scales ignored)

  • fp8Input=true → FP8-input / FP16-output kernels

Parameters:
  • qPtr – Query [B, S_q, H_q, D]

  • kvPtr – Combined KV cache [B, 2, H_kv, Cap, D]

  • oPtr – Output [B, S_q, H_q, D] (always FP16)

  • cuKVSeqLens – Cumulative KV sequence lengths [B+1]

  • stream – CUDA stream

  • attentionScale – Model-defined multiplier applied to QK^T before softmax. For FP8 input, the effective softmax scale is attentionScale * qScale * kScale.

  • slidingWindowSize – Sliding window size (INT_MAX = disabled)

  • fp8Input – Whether Q/KV are FP8 E4M3

  • qScale – Q dequant scale (quant→orig), ignored when fp8Input=false

  • kScale – K dequant scale (quant→orig), ignored when fp8Input=false

  • vScale – V dequant scale (quant→orig), applied to the attention output and ignored when fp8Input=false

  • skipSoftmaxScaleFactor – Skip-softmax (BLASST) calibrated scale factor S; the kernel derives the per-sequence threshold log2(S / seqlen_kv) itself (trtllm-gen parity). A finite positive value (lambda in (0,1)) dispatches the skip-softmax kernel variant, which skips the P*V GEMM of KV tiles whose contribution is negligible — approximate, FP16 causal only. 0.0 (the default, log2 of the degenerate lambda = 1) disables skip, mirroring the slidingWindowSize = INT_MAX sentinel convention.

bool runPaged(
void const *qPtr,
void const *pagedKVPoolPtr,
int32_t const *kvCachePageList,
void *oPtr,
int32_t const *cuKVSeqLens,
int32_t numPages,
int32_t maxPagesPerSeq,
int32_t tokensPerPage,
nvinfer1::DataType kvDataType,
cudaStream_t stream,
float attentionScale,
int32_t slidingWindowSize = INT_MAX,
bool fp8Input = false,
float qScale = 1.0F,
float kScale = 1.0F,
float vScale = 1.0F,
bool isCausal = true,
float skipSoftmaxScaleFactor = 0.0F,
int32_t const *bidirectionalBlockBegin = nullptr,
int32_t const *bidirectionalBlockEnd = nullptr
)#

LLM FMHA over a paged KV cache.

Dispatches a dedicated CuTe DSL AOT variant that reads K/V directly from a paged pool using kvCachePageList. The logical descriptor shape is [numPages, H_kv, tokensPerPage, D], while the physical pool is fixed to NHD [numPages, tokensPerPage, H_kv, D]. This path maps one logical K/V TMA tile to one physical page. The current CuTe DSL variants use a K/V tile width of 128 tokens, so tokensPerPage must be 128 to avoid multi-page tile stitching or a gather workspace. Setting both block range pointers selects the FP16 D512 bidirectional-mask variant; mixed null pointers are invalid.

Parameters:
  • qPtr – Query [B, S_q, H_q, D]

  • pagedKVPoolPtr – Paged KV pool [numPages, tokensPerPage, H_kv, D]

  • kvCachePageList – Page table [B, 2, maxPagesPerSeq], K pages then V pages

  • oPtr – Output [B, S_q, H_q, D] (always FP16)

  • cuKVSeqLens – Cumulative KV sequence lengths [B+1]

  • numPages – Number of pages in the paged KV pool

  • maxPagesPerSeq – Max logical pages per sequence

  • tokensPerPage – Number of tokens per page

  • kvDataType – Paged KV cache dtype (FP16 or FP8)

  • stream – CUDA stream

  • slidingWindowSize – Sliding window size (INT_MAX = disabled)

  • fp8Input – Whether Q/KV are FP8 E4M3

  • qScale – Q dequant scale, ignored when fp8Input=false

  • kScale – K dequant scale, ignored when fp8Input=false

  • vScale – V dequant scale, ignored when fp8Input=false

  • isCausal – Whether to dispatch a causal or dense non-causal variant

  • skipSoftmaxScaleFactor – Skip-softmax calibrated scale factor S (kernel derives per-seq log2(S/seqlen_kv)), or 0 to disable

  • bidirectionalBlockBegin – Optional inclusive bidirectional-block begin positions [B, S_q]. Text/padding rows use -1; every row in a disjoint contiguous vision run must repeat that run’s begin position.

  • bidirectionalBlockEnd – Optional inclusive bidirectional-block end positions [B, S_q]. Text/padding rows use -1; every row in a disjoint contiguous vision run must repeat that run’s end position.

bool runPagedRagged(
void const *qPtr,
void const *pagedKVPoolPtr,
int32_t const *kvCachePageList,
void *oPtr,
int32_t const *cuQSeqLens,
int32_t const *cuKVSeqLens,
int32_t totalQSeqLen,
int32_t maxQSeqLen,
int32_t numFlatPages,
int32_t maxPagesPerSeq,
int32_t tokensPerPage,
cudaStream_t stream,
float attentionScale,
int32_t slidingWindowSize = INT_MAX,
nvinfer1::DataType inputDataType = nvinfer1::DataType::kHALF,
float qScale = 1.0F,
float kScale = 1.0F,
float vScale = 1.0F
)#

Runs causal or sliding-causal FMHA over packed Q/O and a native paged KV pool.

cuQSeqLens and cuKVSeqLens are monotonic prefix arrays of length batchSize + 1. cuQSeqLens starts at zero and ends at totalQSeqLen. cuKVSeqLens contains the complete logical KV history, including prior cached tokens. maxQSeqLen is the maximum adjacent difference in cuQSeqLens, not the packed buffer allocation capacity.

Parameters:
  • qPtr – Query [totalQSeqLen, H_q, D]

  • pagedKVPoolPtr – Flattened K/V pool [numFlatPages, tokensPerPage, H_kv, D]

  • kvCachePageList – Page table [B, 2, maxPagesPerSeq]

  • oPtr – Output [totalQSeqLen, H_q, D] in FP16

  • cuQSeqLens – Cumulative Q sequence lengths [B+1]

  • cuKVSeqLens – Cumulative logical KV sequence lengths [B+1]

  • totalQSeqLen – Number of active packed Q/O tokens

  • maxQSeqLen – Maximum logical Q sequence length in the batch

  • numFlatPages – Number of pages in the flattened K/V pool

  • maxPagesPerSeq – Maximum logical pages per sequence

  • tokensPerPage – Number of tokens per page

  • stream – CUDA stream

  • attentionScale – Model-defined multiplier applied to QK^T before softmax

  • slidingWindowSize – Sliding window size, or INT_MAX when disabled

  • inputDataType – Q/KV input dtype; FP16 and FP8 E4M3 are supported

  • qScale – Q dequant scale, ignored for FP16 input

  • kScale – K dequant scale, ignored for FP16 input

  • vScale – V dequant scale, ignored for FP16 input

bool run(
void const *qPtr,
void const *kPtr,
void const *vPtr,
void *oPtr,
int32_t const *cuSeqLens,
int32_t totalSeqLen,
int32_t maxSeqLen,
int32_t batchSize,
cudaStream_t stream,
float attentionScale,
bool fp8Input = false,
float qScale = 1.0F,
float kScale = 1.0F,
float vScale = 1.0F
)#

ViT FMHA: packed varlen separate Q/K/V, bidirectional.

Output is always FP16. Selects kernel variant based on fp8Input:

  • fp8Input=false → FP16 kernels (all scales ignored)

  • fp8Input=true → FP8-input / FP16-output kernels

Parameters:
  • qPtr – Query [total_S, H, D]

  • kPtr – Key [total_S, H, D]

  • vPtr – Value [total_S, H, D]

  • oPtr – Output [total_S, H, D] (always FP16)

  • cuSeqLens – Cumulative sequence lengths [B+1]

  • totalSeqLen – Sum of all sequence lengths

  • maxSeqLen – Longest individual sequence length

  • batchSize – Number of sequences

  • stream – CUDA stream

  • attentionScale – Absolute multiplier applied to QK^T before softmax. For FP8 input, the effective softmax scale is attentionScale * qScale * kScale.

  • fp8Input – Whether Q/K/V are FP8 E4M3

  • qScale – Q dequant scale (quant→orig), ignored when fp8Input=false

  • kScale – K dequant scale (quant→orig), ignored when fp8Input=false

  • vScale – V dequant scale (quant→orig), applied to the attention output and ignored when fp8Input=false

Public Static Functions

static bool canImplement(int32_t headSize, int32_t smVersion)#
static bool canImplementPagedRagged(
int32_t numQHeads,
int32_t numKVHeads,
int32_t headSize,
int32_t smVersion,
nvinfer1::DataType inputDataType,
CuteDslFMHAMaskType maskType
)#

Returns whether the optimized Blackwell AOT family covers packed-Q/O native paged attention.

static bool canImplementViT(int32_t headSize, int32_t smVersion)#