Qsa Attention Plugin#
-
class QsaAttentionPlugin : public nvinfer1::IPluginV3, public nvinfer1::IPluginV3OneCore, public nvinfer1::IPluginV3OneBuildV2, public nvinfer1::IPluginV3OneRuntime#
QSA (Qwen Sparse Attention) plugin — Qwen3.8-Flash-Next sparse attention layers.
Fuses, in one enqueue:
packed-QKV split + per-head qk-norm (gammas pre-folded to 1+w) + partial RoPE + paged-KV write + dense K/V scratch mirrors (kernel::launchApplyRopeFromPackedToSplit);
the QSA indexer (kernel::runQsaIndexerPrefill): weight-free block-compressed scoring, per-row top-512 blocks, expand + causal tail -> int32 index lists [B, S, 2051];
sparse GQA attention over the dense mirrors (CuteDslQsaSparsePrefillRunner, CuTe DSL AOT group “qsa”). Decode: rope-append of the new token, indexer decode over the pool tails (kernel::runQsaIndexerDecode — B1 fused q-prep + predicated tail write / compress, B2 block scores from the kbar tails, B3 deterministic radix top-512 + expand), then a single-launch split-K sparse attention over the paged pool (CuteDslQsaSparseDecodeRunner).
Contract: the mode is deduced from the kvcache_start_index runtime shape — [0] is normal prefill over padded activations; [B] is single-token decode, which requires S == 1 (chunked prefill and speculative decode are rejected). In decode the start-index values are the per-sequence past lengths and are never read; context_lengths carries the TOTAL lengths including the token being decoded. FP16 surface; all 11 inputs required (no optional slots). The attention output gate (out * sigmoid(gate)) and every projection stay in the graph.
KV pool contract (widened rows): past/present_key_value is [2, numPages, 128, Hkv, head_size + indexer_head_dim] — the pool head dimension is DERIVED from the plugin attributes, never a separate attribute. The leading head_size elements of each row hold roped K / raw V; the tail [head_size, head_size + indexer_head_dim) persists QSA indexer state across steps: block g’s kbar lives in the V-row tail of token 4g (head 0) and, for the trailing incomplete block only, token t’s raw index-K in its K-row tail (head 0). State boundary: the rope/attention kernels are tail-agnostic (the rope kernel writes only the leading head_size elements of each row); only the indexer kernels own the tails — see cpp/kernels/qsaIndexer/qsaIndexerKernels.h.
Public Functions
- QsaAttentionPlugin(
- std::string const &name,
- int32_t numQHeads,
- int32_t numKVHeads,
- int32_t headSize,
- int32_t indexerNumHeads,
- int32_t indexerHeadDim,
- int32_t indexerBudget,
- int32_t indexerCompressRatio,
- float attentionScale,
- float rmsNormEps
- QsaAttentionPlugin(
- std::string const &name,
- nvinfer1::PluginFieldCollection const *fc
-
QsaAttentionPlugin() = delete#
-
~QsaAttentionPlugin() override = default#
- nvinfer1::IPluginCapability *getCapabilityInterface(
- nvinfer1::PluginCapabilityType type
-
nvinfer1::IPluginV3 *clone() noexcept override#
-
char const *getPluginName() const noexcept override#
-
char const *getPluginVersion() const noexcept override#
-
char const *getPluginNamespace() const noexcept override#
-
void setPluginNamespace(char const *pluginNamespace) noexcept#
-
int32_t getNbOutputs() const noexcept override#
- int32_t getOutputDataTypes(
- nvinfer1::DataType *outputTypes,
- int32_t nbOutputs,
- nvinfer1::DataType const *inputTypes,
- int32_t nbInputs
- int32_t getOutputShapes(
- nvinfer1::DimsExprs const *inputs,
- int32_t nbInputs,
- nvinfer1::DimsExprs const *shapeInputs,
- int32_t nbShapeInputs,
- nvinfer1::DimsExprs *outputs,
- int32_t nbOutputs,
- nvinfer1::IExprBuilder &exprBuilder
- bool supportsFormatCombination(
- int32_t pos,
- nvinfer1::DynamicPluginTensorDesc const *inOut,
- int32_t nbInputs,
- int32_t nbOutputs
- int32_t configurePlugin(
- nvinfer1::DynamicPluginTensorDesc const *in,
- int32_t nbInputs,
- nvinfer1::DynamicPluginTensorDesc const *out,
- int32_t nbOutputs
- size_t getWorkspaceSize(
- nvinfer1::DynamicPluginTensorDesc const *inputs,
- int32_t nbInputs,
- nvinfer1::DynamicPluginTensorDesc const *outputs,
- int32_t nbOutputs
-
int32_t getAliasedInput(int32_t outputIndex) noexcept override#
- int32_t enqueue(
- nvinfer1::PluginTensorDesc const *inputDesc,
- nvinfer1::PluginTensorDesc const *outputDesc,
- void const *const *inputs,
- void *const *outputs,
- void *workspace,
- cudaStream_t stream
- int32_t onShapeChange(
- nvinfer1::PluginTensorDesc const *in,
- int32_t nbInputs,
- nvinfer1::PluginTensorDesc const *out,
- int32_t nbOutputs
- nvinfer1::IPluginV3 *attachToContext(
- nvinfer1::IPluginResourceContext *context
- nvinfer1::PluginFieldCollection const *getFieldsToSerialize(
-
class QsaAttentionPluginCreator : public nvinfer1::IPluginCreatorV3One#
Public Functions
-
QsaAttentionPluginCreator()#
-
~QsaAttentionPluginCreator() override = default#
-
char const *getPluginName() const noexcept override#
-
char const *getPluginVersion() const noexcept override#
- nvinfer1::PluginFieldCollection const *getFieldNames(
- nvinfer1::IPluginV3 *createPlugin(
- char const *name,
- nvinfer1::PluginFieldCollection const *fc,
- nvinfer1::TensorRTPhase phase
-
void setPluginNamespace(char const *libNamespace) noexcept#
-
char const *getPluginNamespace() const noexcept override#
-
QsaAttentionPluginCreator()#