Cute Dsl Qsa Sparse Runner#
-
class CuteDslQsaSparsePrefillRunner#
Runner for the CuTe DSL QSA sparse-GQA prefill kernels (Qwen3.8-Flash-Next).
One CTA per (query token, KV head); the GQA head group is the MMA M tile and the KV traversal is a per-row cp.async gather over the index list. The AOT family bakes only the head dim and tile tuning; batch, seq, head counts, topK and strides stay runtime-dynamic.
Public Functions
-
CuteDslQsaSparsePrefillRunner() = delete#
Public Static Functions
- static bool canImplement(
- int32_t numQHeads,
- int32_t numKVHeads,
- int32_t headDim,
- int32_t smVersion,
- nvinfer1::DataType dataType
Returns whether the QSA AOT family covers this shape on this SM.
- static bool preflight(
- nvinfer1::DataType dataType,
- cudaStream_t stream
Ensures the variant selected by run() is loaded (CUDA-graph-capture aware).
- static bool run(
- nvinfer1::DataType dataType,
- QsaSparsePrefillParams const ¶ms
Launches sparse prefill attention over the gathered index lists.
-
CuteDslQsaSparsePrefillRunner() = delete#
-
class CuteDslQsaSparseDecodeRunner#
Runner for the single-launch split-K QSA sparse decode kernel.
Public Functions
-
CuteDslQsaSparseDecodeRunner() = delete#
Public Static Functions
- static bool canImplement(
- int32_t numQHeads,
- int32_t numKVHeads,
- int32_t headDim,
- int32_t poolHeadDim,
- int32_t smVersion,
- nvinfer1::DataType dataType
- static bool preflight(
- nvinfer1::DataType dataType,
- cudaStream_t stream
Ensures the decode variant is loaded (CUDA-graph-capture aware).
- static bool run(
- nvinfer1::DataType dataType,
- QsaSparseDecodeParams const ¶ms
-
CuteDslQsaSparseDecodeRunner() = delete#
-
struct QsaSparsePrefillParams#
Per-launch parameters for the QSA sparse-GQA prefill kernel.
Q/K/V/O are dense padded BSND tensors: Q/O are [batchSize, seqLen, numQHeads, headDim] and K/V are [batchSize, seqLen, numKVHeads, headDim], all fp16 (or bf16), contiguous with D innermost.
indicesis the QSA indexer output [batchSize, seqLen, topK] Int32: per query row the selected KV token ids, -1-padded on the right, unsorted, distinct, all < token + 1. Causality lives entirely in the index list — the kernel applies no positional mask.contextLengthsis [batchSize] Int32; output rows at or past the live length store exact zeros, and K/V rows at or past it are never referenced (indices only name live tokens).Public Members
-
void const *qPtr = {}#
-
void const *kPtr = {}#
-
void const *vPtr = {}#
-
void *oPtr = {}#
-
int32_t const *indices = {}#
-
int32_t const *contextLengths = {}#
-
int32_t batchSize = {}#
-
int32_t seqLen = {}#
-
int32_t numQHeads = {}#
-
int32_t numKVHeads = {}#
-
int32_t headDim = {}#
-
int32_t topK = {}#
-
float attentionScale = {}#
-
cudaStream_t stream = {}#
-
void const *qPtr = {}#
-
struct QsaSparseDecodeParams#
Per-launch parameters for the QSA split-K sparse DECODE kernel.
One query token per sequence:
qPtr/oPtrare [batchSize, 1, numQHeads, headDim]. K/V come from the paged pool [2*numPages, 128, numKVHeads, poolHeadDim] throughpageTable[batchSize, 2, maxPagesPerSeq] (V page ids pre-offset by +numPages); only columns [0, headDim) of each row are read — the tail carries the QSA indexer state.indicesis [batchSize, 1, topK] Int32 (-1 padded).contextLengthsis the TOTAL per-sequence length including the new token.partialO(fp32 [batchSize*numKVHeads*kMaxSplits, kPartialRows, headDim]),partialStats(fp32 [batchSize*numKVHeads*kMaxSplits, 2, kPartialRows]) andsplitCounters(int32 [batchSize*numKVHeads]) live in plugin workspace; counters MUST be zero on entry (the kernel release-resets them, but fresh workspaces start as garbage).Public Members
-
void const *qPtr = {}#
-
void const *kvPoolPtr = {}#
-
int32_t const *pageTable = {}#
-
int32_t const *indices = {}#
-
int32_t const *contextLengths = {}#
-
void *oPtr = {}#
-
float *partialO = {}#
-
float *partialStats = {}#
-
int32_t *splitCounters = {}#
-
int32_t batchSize = {}#
-
int32_t numQHeads = {}#
-
int32_t numKVHeads = {}#
-
int32_t headDim = {}#
-
int32_t poolHeadDim = {}#
-
int32_t numFlatPages = {}#
2 * numPages (both planes)
-
int32_t maxPagesPerSeq = {}#
-
int32_t topK = {}#
-
float attentionScale = {}#
-
cudaStream_t stream = {}#
-
void const *qPtr = {}#