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 &params
)#

Launches sparse prefill attention over the gathered index lists.

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 &params
)#

Public Static Attributes

static int32_t kMaxSplits = {8}#

Grid split dimension baked into the AOT variants (build_cutedsl.py —max_splits).

static int32_t kPartialRows = {16}#

Accumulator row count of the partial workspaces (m_block_size).

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. indices is 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. contextLengths is [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 = {}#
struct QsaSparseDecodeParams#

Per-launch parameters for the QSA split-K sparse DECODE kernel.

One query token per sequence: qPtr / oPtr are [batchSize, 1, numQHeads, headDim]. K/V come from the paged pool [2*numPages, 128, numKVHeads, poolHeadDim] through pageTable [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. indices is [batchSize, 1, topK] Int32 (-1 padded). contextLengths is the TOTAL per-sequence length including the new token. partialO (fp32 [batchSize*numKVHeads*kMaxSplits, kPartialRows, headDim]), partialStats (fp32 [batchSize*numKVHeads*kMaxSplits, 2, kPartialRows]) and splitCounters (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 = {}#