Mamba Plugin#
-
class MambaPlugin : public nvinfer1::IPluginV3, public nvinfer1::IPluginV3OneCore, public nvinfer1::IPluginV3OneBuild, public nvinfer1::IPluginV3OneRuntime#
TensorRT plugin for Mamba Selective State Update (SSM)
Registered as “update_ssm_state” under the “trt_edgellm” ONNX domain.
Implements the selective state space model update: new_state = state * exp(A * dt) + B * dt * x output = sum_i(new_state_i * C_i) + D * x
SiLU gating (z) is handled externally by the ONNX graph (gated_rms_norm).
Input ordering (see constants defined in mambaPlugin.cpp): [0] x [tokens, nheads, dim] FP16 [1] A [nheads] FP32 [2] B [tokens, ngroups, dstate] FP16 [3] C [tokens, ngroups, dstate] FP16 [4] D [nheads] FP16 [5] dt [tokens, nheads] FP16 [6] dt_bias [nheads] FP16 [7] state [resident_rows, nheads, dim, dstate] FP16 [8] query_lengths [batch] INT32 [9] query_start_offsets [batch + 1] INT32 [10] state_indices [batch] INT32 [11] execution_phase_marker [1..8] INT32 [12] context_sequence_count_carrier [0..batch] INT32 shape-only context count [13] tree_parent_ids [tokens] INT32 (DDTree engines only) [14] tree_depths [tokens] INT32 (DDTree engines only)
Outputs: [0] output [tokens, nheads, dim] FP16 [1] state_out [resident_rows, nheads, dim, dstate] aliased state pool [2] replay_da [tokens, nheads] FP32 (spec-verify engines only) [3] replay_u [tokens, nheads, dim] FP32 (spec-verify engines only) [4] replay_b [tokens, ngroups, dstate] FP32 (spec-verify engines only) [5] replay_dt [tokens, nheads] FP32 (spec-verify engines only)
Public Functions
- MambaPlugin(
- std::string const &name,
- int32_t dim,
- int32_t dstate,
- int32_t nheads,
- int32_t ngroups,
- int32_t dtSoftplus,
- int32_t useSpecVerifyState = 0,
- int32_t useDDTree = 0,
- int32_t replayFormatVersion = 0
-
MambaPlugin() = delete#
-
MambaPlugin(MambaPlugin const&) = delete#
-
~MambaPlugin() 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 MambaPluginCreator : public nvinfer1::IPluginCreatorV3One#
Public Functions
-
MambaPluginCreator()#
-
~MambaPluginCreator() override = default#
-
char const *getPluginName() const noexcept override#
-
char const *getPluginVersion() const noexcept override#
- nvinfer1::PluginFieldCollection const *getFieldNames(
-
void setPluginNamespace(char const *pluginNamespace) noexcept#
-
char const *getPluginNamespace() const noexcept override#
- nvinfer1::IPluginV3 *createPlugin(
- char const *name,
- nvinfer1::PluginFieldCollection const *fc,
- nvinfer1::TensorRTPhase phase
-
MambaPluginCreator()#