Gemma4 Audio Attention Plugin#
-
class Gemma4AudioAttentionPlugin : public nvinfer1::IPluginV3, public nvinfer1::IPluginV3OneCore, public nvinfer1::IPluginV3OneBuild, public nvinfer1::IPluginV3OneRuntime#
TensorRT plugin for Gemma 4 audio-encoder chunked local attention (IPluginV3).
Wraps the fused CUDA kernel that computes the full attention body after Q/K/V projection and before the output projection: per-dim learned Q scaling, fixed K scaling, chunked local context gather, content + relative-position scores, tanh soft-cap, local-causal + padding mask, fp32 softmax, value mix.
Inputs (7): 0: qRaw [B, S, H, D] (half/bf16/float) — raw query projections 1: kRaw [B, S, H, D] (same type) — raw key projections 2: v [B, S, H, D] (same type) — value projections 3: gamma [D] (float) — per-dim learned query scale 4: relKey [P, H, D] (same type as q) — projected relative-position embeddings 5: valid [B, S] (bool) — audio validity mask (true = real token) 6: seqLen [1] (int32) — actual sequence length (shape carrier)
Outputs (1): 0: out [B, S, H, D] (same type as q) — attention output
Plugin attributes (serialized): chunk_size (int32) — C, query block size (default 12) left_horizon (int32) — L, effective left context (default 12) context_size (int32) — M, gathered K/V context size (default 24) logit_cap (float) — tanh soft-cap on logits (default 50.0)
Public Functions
- Gemma4AudioAttentionPlugin(
- std::string const &name,
- int32_t chunkSize,
- int32_t leftHorizon,
- int32_t contextSize,
- float logitCap,
- Gemma4AudioAttentionPlugin(
- std::string const &name,
- nvinfer1::PluginFieldCollection const *fc,
-
Gemma4AudioAttentionPlugin() = delete#
-
Gemma4AudioAttentionPlugin(Gemma4AudioAttentionPlugin const&) = delete#
-
~Gemma4AudioAttentionPlugin() override#
- 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#
-
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 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(
-
void setPluginNamespace(char const *pluginNamespace) noexcept#
-
class Gemma4AudioAttentionPluginCreator : public nvinfer1::IPluginCreatorV3One#
Factory class for creating Gemma4AudioAttentionPlugin instances.
Public Functions
-
Gemma4AudioAttentionPluginCreator()#
-
~Gemma4AudioAttentionPluginCreator() override = default#
-
char const *getPluginName() const noexcept override#
-
char const *getPluginVersion() const noexcept override#
- nvinfer1::PluginFieldCollection const *getFieldNames(
-
char const *getPluginNamespace() const noexcept override#
-
void setPluginNamespace(char const *pluginNamespace) noexcept#
- nvinfer1::IPluginV3 *createPlugin(
- char const *name,
- nvinfer1::PluginFieldCollection const *fc,
- nvinfer1::TensorRTPhase phase,
-
Gemma4AudioAttentionPluginCreator()#