pulsatrix
Loading...
Searching...
No Matches
pulsatrix::MambaModule Class Reference

Core Mamba/S6 selective-scan recurrence (Gu & Dao 2023, arXiv:2312.00752), input (N, L, d_model) -> output (N, L, d_model). More...

#include <mamba_module.hpp>

Inheritance diagram for pulsatrix::MambaModule:
Collaboration diagram for pulsatrix::MambaModule:

Public Member Functions

 MambaModule (int64_t d_model, int64_t state_size, DeviceBackend *backend)
 Constructs a selective-scan layer with zero-initialized parameters.
 
Tensor backward (const Tensor &grad_output) override
 Real backpropagation-through-time across the selective scan: accumulates W_delta/bias_delta/W_B/W_C/A/D gradients across every timestep into the same buffers via Tensor::accumulate(). Unlike propagate_relevance, this is the genuine undetached gradient – it differentiates through Delta_t's softplus, through Abar_t = exp(Delta_t*A) and through all three selective projections.
 
OpType op_type () const override
 Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule/LSTMModule/GRUModule (no enum change).
 
void set_W_delta (std::initializer_list< float > values)
 Overwrites the input-to-Delta projection weight (d_model, d_model) – test/initialization use only.
 
void set_bias_delta (std::initializer_list< float > values)
 Overwrites the Delta projection bias (d_model,) – test/initialization use only.
 
void set_W_B (std::initializer_list< float > values)
 Overwrites the input-to-B projection weight (d_model, state_size) – test/initialization use only.
 
void set_W_C (std::initializer_list< float > values)
 Overwrites the input-to-C projection weight (d_model, state_size) – test/initialization use only.
 
void set_A (std::initializer_list< float > values)
 Overwrites the continuous state matrix A (d_model, state_size) – test/initialization use only.
 
void set_D (std::initializer_list< float > values)
 Overwrites the skip/feedthrough vector D (d_model,) – test/initialization use only.
 
void set_W_delta (const std::vector< float > &values)
 std::vector overload of set_W_delta() – for callers building values programmatically.
 
void set_bias_delta (const std::vector< float > &values)
 std::vector overload of set_bias_delta().
 
void set_W_B (const std::vector< float > &values)
 std::vector overload of set_W_B().
 
void set_W_C (const std::vector< float > &values)
 std::vector overload of set_W_C().
 
void set_A (const std::vector< float > &values)
 std::vector overload of set_A().
 
void set_D (const std::vector< float > &values)
 std::vector overload of set_D().
 
const Tensor & W_delta () const
 
const Tensor & bias_delta () const
 
const Tensor & W_B () const
 
const Tensor & W_C () const
 
const Tensor & A () const
 
const Tensor & D () const
 
const Tensor & W_delta_grad () const
 
const Tensor & bias_delta_grad () const
 
const Tensor & W_B_grad () const
 
const Tensor & W_C_grad () const
 
const Tensor & A_grad () const
 
const Tensor & D_grad () const
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 MambaLRP relevance propagation – Abar_t/Bbar_t/C_t/A/D detached and treated as fixed multiplicative constants, epsilon/z-rule over the resulting weighted sums, processed in reverse time order through a carried state-relevance accumulator. See the class-level note for the full per-timestep redistribution.
 
std::vector< NamedParamRef > named_parameters () override
 This module's trainable parameters, each with its hierarchical name – the one place a module declares its parameters (roadmap FND-1).
 
std::optional< DeviceType > compute_device () const override
 Where this layer computes, so forward() rejects an input on another device (FND-8).
 
- Public Member Functions inherited from pulsatrix::Module
virtual ~Module ()=default
 
Tensor forward (const Tensor &input)
 Runs this module's forward computation.
 
std::pair< Tensor, NodeId > forward_traced (const Tensor &input, NodeId input_node, ComputationGraph &graph, Autograd &autograd)
 Runs forward() while also registering a ComputationGraph node (tagged with this module's op_type(), parented to input_node) and wiring an Autograd backward function that reuses this module's own backward() – the opt-in traced/explainable path, per Phase 2 Mission 0.
 
virtual bool supports_lrp_rule (LRPRule rule) const
 Whether propagate_relevance() implements rule (no silent fallback: callers such as ExplainerContext::relevance_pass() throw rather than run a module on a rule it does not implement).
 
virtual std::vector< ParamRef > parameters ()
 This module's trainable parameters and their gradients, for an optimizer to update uniformly across module types.
 
void set_requires_grad (bool requires_grad, const std::string &prefix="")
 Freezes (false) or unfreezes (true) parameters by name (roadmap FND-2).
 
virtual void set_training (bool training)
 Sets this module's training/eval mode. Defaults to training (matches every mainstream framework's Module default).
 
bool is_training () const
 Whether this module is currently in training mode.
 

Protected Member Functions

Tensor forward_impl (const Tensor &input) override
 The actual forward computation – per-timestep tied-weight selective scan.
 

Detailed Description

Core Mamba/S6 selective-scan recurrence (Gu & Dao 2023, arXiv:2312.00752), input (N, L, d_model) -> output (N, L, d_model).

Per timestep t, with batch row b, channel d in [0, d_model) and state index n in [0, state_size): z_delta_t[b,d] = sum_e x_t[b,e]*W_delta[e,d] + bias_delta[d] (Linear, WITH bias) Delta_t[b,d] = softplus(z_delta_t[b,d]) (> 0, stable step) B_t[b,n] = sum_e x_t[b,e]*W_B[e,n] (Linear, NO bias) C_t[b,n] = sum_e x_t[b,e]*W_C[e,n] (Linear, NO bias) Abar_t[b,d,n] = exp(Delta_t[b,d] * A[d,n]) Bbar_t[b,d,n] = Delta_t[b,d] * B_t[b,n] h_t[b,d,n] = Abar_t[b,d,n]*h_{t-1}[b,d,n] + Bbar_t[b,d,n]*x_t[b,d] y_t[b,d] = sum_n( C_t[b,n]*h_t[b,d,n] ) + D[d]*x_t[b,d]

Note
Scope cut, mirroring every prior recurrent module's own: this is the core selective-scan recurrence only, not the full Mamba block. The full block additionally wraps this in an input Conv1D, an outer SiLU-gated multiplicative branch and a final output projection – none of which are built here, for the same reason RNNModule/LSTMModule/GRUModule each implement only their own core recurrence. MultiHeadAttentionModule/SwiGLUModule/TransformerBlock already established that small, independently-correct pieces get composed later.
Discretization scope cut: Bbar uses the first-order/Euler approximation Mamba's own reference implementation uses, not the full ZOH integral (A^{-1}(exp(Delta*A)-I)*Delta*B). Abar is the exact ZOH form exp(Delta*A).
h_0 = 0, zero-initialized and not learnable – same scope cut as every prior recurrent module, and the thing that makes this module's conservation exact rather than merely approximate (see the LRP note below).
Device-generic (GPU-native-kernels Mission 6): forward, backward and propagate_relevance run entirely through DeviceBackend – the projections through gemm/gemm_ex, the selective scan (softplus, exp, the h_t recurrence), its BPTT and its LRP through DeviceBackend::ssm_pass (one lane per (batch, channel), sequential over time) – so the module runs on CPU, CUDA and HIP with no host round-trip.
LRP rule (MambaLRP – Rezaei Jafari, Montavon, Müller, Eberle, "MambaLRP: Explaining Selective State Space Sequence Models", arXiv:2406.07592, NeurIPS 2024). MambaLRP's central finding: naive LRP breaks conservation on the selective scan because the discrete recurrence parameters (Abar_t, Bbar_t, C_t) are themselves functions of the input via their own learned projections, and redistributing "through" that input-dependence corrupts the redistribution. The fix is to treat Abar_t/Bbar_t/C_t (and A, D) as detached fixed multiplicative constants during relevance propagation – exactly the way every existing module here already treats its weight matrices. That collapses the recurrence into the familiar weighted-sum shape, so no new rule shape is needed; the novelty is specifically what gets detached. Per timestep, in reverse time order:
  1. y_t = sum_n(C_t*h_t) + D*x_t is an (state_size + 1)-way weighted sum sharing one output; the epsilon/z-rule (denominator y_t itself, epsilon-stabilized) splits R(y_t) across the state_size state terms and the D skip term in proportion to their values. The skip share lands directly on R(x_t).
  2. h_t = Abar_t*h_{t-1} + Bbar_t*x_t is the two-weighted-source epsilon/z-rule, identical in shape to RNNModule's own, with denominator h_t itself; the h_{t-1} share accumulates into the carried t-1 relevance accumulator and the x_t share adds onto R(x_t).
  3. W_delta/bias_delta/W_B/W_C – the selective projections' own upstream computation – are never touched by propagate_relevance. Delta_t/B_t/C_t are pure conductors, the same convention LSTMModule's/GRUModule's gates use.
  4. Because h_0 = 0 (this module's own scope cut), the relevance that would otherwise leak into the nonexistent state before t=0 is provably exactly zero (the t=0 numerator is Abar_0*h_{-1} with h_{-1} == 0). End-to-end conservation therefore holds up to the epsilon stabilizers only – measured at 2.5e-5 against a sum(R_out) of 5.3 (4.7e-6 relative) in mamba_module_test.cpp's PropagateRelevanceConservationGapIsMeasuredNotAssumed, i.e. the near-exact category (like RoPEModule/SwiGLUModule), NOT SoftmaxModule's large by-design DTD gap. This is the expected outcome: restoring conservation that naive LRP breaks is MambaLRP's whole point, so a large gap here would be a bug, not a property.

Constructor & Destructor Documentation

◆ MambaModule()

pulsatrix::MambaModule::MambaModule ( int64_t  d_model,
int64_t  state_size,
DeviceBackend *  backend 
)

Constructs a selective-scan layer with zero-initialized parameters.

Parameters
d_modelModel/channel dimension (also the input and output feature dimension).
state_sizeLatent SSM state dimension N.
backendBackend to allocate/compute through. Not owned; must outlive this module.
Exceptions
std::invalid_argumentif d_model <= 0 or state_size <= 0 – external boundary (construction arguments can originate from the Python bindings with no upstream validation).

Member Function Documentation

◆ A()

const Tensor & pulsatrix::MambaModule::A ( ) const
inline

◆ A_grad()

const Tensor & pulsatrix::MambaModule::A_grad ( ) const
inline

◆ backward()

Tensor pulsatrix::MambaModule::backward ( const Tensor &  grad_output)
overridevirtual

Real backpropagation-through-time across the selective scan: accumulates W_delta/bias_delta/W_B/W_C/A/D gradients across every timestep into the same buffers via Tensor::accumulate(). Unlike propagate_relevance, this is the genuine undetached gradient – it differentiates through Delta_t's softplus, through Abar_t = exp(Delta_t*A) and through all three selective projections.

Parameters
grad_outputGradient w.r.t. this module's output. Must be (N, L, d_model) matching the most recent forward() call's output shape.
Returns
Gradient w.r.t. this module's input, shape (N, L, d_model).
Exceptions
std::logic_errorif forward() has never been called.
std::invalid_argumentif grad_output's shape doesn't match the cached forward output shape.
Note
Device-generic (GPU-native-kernels Mission 6) – see the class-level note.

Implements pulsatrix::Module.

◆ bias_delta()

const Tensor & pulsatrix::MambaModule::bias_delta ( ) const
inline

◆ bias_delta_grad()

const Tensor & pulsatrix::MambaModule::bias_delta_grad ( ) const
inline

◆ compute_device()

std::optional< DeviceType > pulsatrix::MambaModule::compute_device ( ) const
inlineoverridevirtual

Where this layer computes, so forward() rejects an input on another device (FND-8).

Reimplemented from pulsatrix::Module.

◆ D()

const Tensor & pulsatrix::MambaModule::D ( ) const
inline

◆ D_grad()

const Tensor & pulsatrix::MambaModule::D_grad ( ) const
inline

◆ forward_impl()

Tensor pulsatrix::MambaModule::forward_impl ( const Tensor &  input)
overrideprotectedvirtual

The actual forward computation – per-timestep tied-weight selective scan.

Exceptions
std::invalid_argumentif input isn't rank-3 (N, L, d_model), or its last dimension doesn't match d_model.
Note
Device-generic (GPU-native-kernels Mission 6) – see the class-level note.

Implements pulsatrix::Module.

◆ named_parameters()

std::vector< NamedParamRef > pulsatrix::MambaModule::named_parameters ( )
inlineoverridevirtual

This module's trainable parameters, each with its hierarchical name – the one place a module declares its parameters (roadmap FND-1).

Returns
{name, {value, grad}} entries pointing directly at this module's own members, in a fixed order. Names are unique within the module tree. Default: empty (a parameterless module like ReluModule needs no override).
Note
Override this, not parameters(): saving, loading, freezing by name and optimizer parameter groups all key on these names.

Reimplemented from pulsatrix::Module.

◆ op_type()

OpType pulsatrix::MambaModule::op_type ( ) const
inlineoverridevirtual

Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule/LSTMModule/GRUModule (no enum change).

Implements pulsatrix::Module.

◆ propagate_relevance()

Tensor pulsatrix::MambaModule::propagate_relevance ( const Tensor &  relevance_out,
const LRPRuleConfig &  config 
)
overridevirtual

MambaLRP relevance propagation – Abar_t/Bbar_t/C_t/A/D detached and treated as fixed multiplicative constants, epsilon/z-rule over the resulting weighted sums, processed in reverse time order through a carried state-relevance accumulator. See the class-level note for the full per-timestep redistribution.

Parameters
relevance_outRelevance at this module's output. Must be (N, L, d_model) matching the most recent forward() call's output shape.
configSelects epsilon.
Returns
Relevance at this module's input, shape (N, L, d_model). Conserves up to the epsilon stabilizers – see the class-level note for the measured gap.
Exceptions
std::logic_errorif forward() has never been called.
std::invalid_argumentif relevance_out's shape doesn't match the cached forward output shape.
Note
Device-generic (GPU-native-kernels Mission 6) – see the class-level note.

Implements pulsatrix::Module.

◆ set_A() [1/2]

void pulsatrix::MambaModule::set_A ( const std::vector< float > &  values)

std::vector overload of set_A().

◆ set_A() [2/2]

void pulsatrix::MambaModule::set_A ( std::initializer_list< float >  values)

Overwrites the continuous state matrix A (d_model, state_size) – test/initialization use only.

◆ set_bias_delta() [1/2]

void pulsatrix::MambaModule::set_bias_delta ( const std::vector< float > &  values)

std::vector overload of set_bias_delta().

◆ set_bias_delta() [2/2]

void pulsatrix::MambaModule::set_bias_delta ( std::initializer_list< float >  values)

Overwrites the Delta projection bias (d_model,) – test/initialization use only.

◆ set_D() [1/2]

void pulsatrix::MambaModule::set_D ( const std::vector< float > &  values)

std::vector overload of set_D().

◆ set_D() [2/2]

void pulsatrix::MambaModule::set_D ( std::initializer_list< float >  values)

Overwrites the skip/feedthrough vector D (d_model,) – test/initialization use only.

◆ set_W_B() [1/2]

void pulsatrix::MambaModule::set_W_B ( const std::vector< float > &  values)

std::vector overload of set_W_B().

◆ set_W_B() [2/2]

void pulsatrix::MambaModule::set_W_B ( std::initializer_list< float >  values)

Overwrites the input-to-B projection weight (d_model, state_size) – test/initialization use only.

◆ set_W_C() [1/2]

void pulsatrix::MambaModule::set_W_C ( const std::vector< float > &  values)

std::vector overload of set_W_C().

◆ set_W_C() [2/2]

void pulsatrix::MambaModule::set_W_C ( std::initializer_list< float >  values)

Overwrites the input-to-C projection weight (d_model, state_size) – test/initialization use only.

◆ set_W_delta() [1/2]

void pulsatrix::MambaModule::set_W_delta ( const std::vector< float > &  values)

std::vector overload of set_W_delta() – for callers building values programmatically.

◆ set_W_delta() [2/2]

void pulsatrix::MambaModule::set_W_delta ( std::initializer_list< float >  values)

Overwrites the input-to-Delta projection weight (d_model, d_model) – test/initialization use only.

◆ W_B()

const Tensor & pulsatrix::MambaModule::W_B ( ) const
inline

◆ W_B_grad()

const Tensor & pulsatrix::MambaModule::W_B_grad ( ) const
inline

◆ W_C()

const Tensor & pulsatrix::MambaModule::W_C ( ) const
inline

◆ W_C_grad()

const Tensor & pulsatrix::MambaModule::W_C_grad ( ) const
inline

◆ W_delta()

const Tensor & pulsatrix::MambaModule::W_delta ( ) const
inline

◆ W_delta_grad()

const Tensor & pulsatrix::MambaModule::W_delta_grad ( ) const
inline

The documentation for this class was generated from the following file: