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

Standard 4-gate LSTM recurrence, h_0 = c_0 = 0 (zero-initialized, not learnable – the same deliberate scope cut RNNModule made, and the same thing that makes this module's conservation exact; see the LRP note below): i_t = sigmoid(x_t @ W_xi + h_{t-1} @ W_hi + b_i), f_t = sigmoid(x_t @ W_xf + h_{t-1} @ W_hf + b_f), g_t = tanh(x_t @ W_xg + h_{t-1} @ W_hg + b_g), o_t = sigmoid(x_t @ W_xo + h_{t-1} @ W_ho + b_o), c_t = f_t * c_{t-1} + i_t * g_t, h_t = o_t * tanh(c_t). Input (N, L, input_size) -> output (N, L, hidden_size), the full hidden-state sequence (matches RNNModule's convention). Single layer, no bidirectional/ multi-layer/variable-length/peephole support. More...

#include <lstm_module.hpp>

Inheritance diagram for pulsatrix::LSTMModule:
Collaboration diagram for pulsatrix::LSTMModule:

Public Member Functions

 LSTMModule (int64_t input_size, int64_t hidden_size, DeviceBackend *backend)
 Constructs an LSTM layer with zero-initialized weights/biases.
 
Tensor backward (const Tensor &grad_output) override
 Real backpropagation-through-time (BPTT) across all four gates and the cell carry: accumulates every per-gate input weight, recurrent weight and bias gradient across every timestep into the same buffers via Tensor::accumulate().
 
OpType op_type () const override
 Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule (no enum change needed for this mission).
 
void set_weight_xi (std::initializer_list< float > values)
 Overwrites the input-to-input-gate weight buffer – test/initialization use only.
 
void set_weight_hi (std::initializer_list< float > values)
 Overwrites the hidden-to-input-gate weight buffer – test/initialization use only.
 
void set_bias_i (std::initializer_list< float > values)
 Overwrites the input-gate bias buffer – test/initialization use only.
 
void set_weight_xf (std::initializer_list< float > values)
 Overwrites the input-to-forget-gate weight buffer – test/initialization use only.
 
void set_weight_hf (std::initializer_list< float > values)
 Overwrites the hidden-to-forget-gate weight buffer – test/initialization use only.
 
void set_bias_f (std::initializer_list< float > values)
 Overwrites the forget-gate bias buffer – test/initialization use only.
 
void set_weight_xg (std::initializer_list< float > values)
 Overwrites the input-to-cell-candidate weight buffer – test/initialization use only.
 
void set_weight_hg (std::initializer_list< float > values)
 Overwrites the hidden-to-cell-candidate weight buffer – test/initialization use only.
 
void set_bias_g (std::initializer_list< float > values)
 Overwrites the cell-candidate bias buffer – test/initialization use only.
 
void set_weight_xo (std::initializer_list< float > values)
 Overwrites the input-to-output-gate weight buffer – test/initialization use only.
 
void set_weight_ho (std::initializer_list< float > values)
 Overwrites the hidden-to-output-gate weight buffer – test/initialization use only.
 
void set_bias_o (std::initializer_list< float > values)
 Overwrites the output-gate bias buffer – test/initialization use only.
 
const Tensor & weight_xi () const
 
const Tensor & weight_hi () const
 
const Tensor & bias_i () const
 
const Tensor & weight_xf () const
 
const Tensor & weight_hf () const
 
const Tensor & bias_f () const
 
const Tensor & weight_xg () const
 
const Tensor & weight_hg () const
 
const Tensor & bias_g () const
 
const Tensor & weight_xo () const
 
const Tensor & weight_ho () const
 
const Tensor & bias_o () const
 
const Tensor & weight_xi_grad () const
 
const Tensor & weight_hi_grad () const
 
const Tensor & bias_i_grad () const
 
const Tensor & weight_xf_grad () const
 
const Tensor & weight_hf_grad () const
 
const Tensor & bias_f_grad () const
 
const Tensor & weight_xg_grad () const
 
const Tensor & weight_hg_grad () const
 
const Tensor & bias_g_grad () const
 
const Tensor & weight_xo_grad () const
 
const Tensor & weight_ho_grad () const
 
const Tensor & bias_o_grad () const
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 Arras et al. 2019 gate-signal LRP: gates conduct, signals receive. 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 gated recurrence.
 

Detailed Description

Standard 4-gate LSTM recurrence, h_0 = c_0 = 0 (zero-initialized, not learnable – the same deliberate scope cut RNNModule made, and the same thing that makes this module's conservation exact; see the LRP note below): i_t = sigmoid(x_t @ W_xi + h_{t-1} @ W_hi + b_i), f_t = sigmoid(x_t @ W_xf + h_{t-1} @ W_hf + b_f), g_t = tanh(x_t @ W_xg + h_{t-1} @ W_hg + b_g), o_t = sigmoid(x_t @ W_xo + h_{t-1} @ W_ho + b_o), c_t = f_t * c_{t-1} + i_t * g_t, h_t = o_t * tanh(c_t). Input (N, L, input_size) -> output (N, L, hidden_size), the full hidden-state sequence (matches RNNModule's convention). Single layer, no bidirectional/ multi-layer/variable-length/peephole support.

Note
Device-generic (GPU-native-kernels Mission 5): forward(), backward() and propagate_relevance() run entirely through DeviceBackend primitives (gemm/gemm_ex, copy_2d timestep slicing, elementwise Sigmoid/Tanh, recurrent_cell(LstmForward / LstmBackward / LstmLrp), accumulate_rows, lrp_linear), reproducing the former host loops' evaluation order so CPU results are bit-identical.
LRP rule (Arras et al. 2019, "Explaining Recurrent Neural Network Predictions in Sentiment Analysis"): the gates i_t, f_t, o_t are pure conductors, never relevance recipients. Every multiplicative interaction in the recurrence is an exact bilinear identity summing to its own result, so relevance splits across the signal operands only, with the gate values acting as fixed multiplicative weights:
  • c_t = f_t*c_{t-1} + i_t*g_t is a two-term weighted sum (weights f_t/i_t, signals c_{t-1}/g_t); R(c_t) splits between R(c_{t-1}) and R(g_t) in proportion to f_t*c_{t-1} and i_t*g_t, via the same epsilon/z-rule redistribution RNNModule applies to its two-source pre-activation.
  • h_t = o_t * tanh(c_t) is a single-signal product, so ALL of R(h_t) passes to R(tanh(c_t)) and none to R(o_t).
  • tanh(c_t) -> c_t is identity pass-through (same pointwise-nonlinearity precedent as ReluModule/RNNModule, Montavon et al. 2019).
  • R(g_t) is then redistributed across x_t and h_{t-1} by one epsilon/z-rule step over g_t's two weighted sources, exactly like RNNModule's. Net flow per timestep, processed in reverse time order threading both a hidden and a cell relevance accumulator (mirroring backward()'s BPTT accumulators): h_t -> c_t -> { c_{t-1} (carried to t-1), g_t -> { x_t, h_{t-1} (carried to t-1) } }. Because h_0 and c_0 are zero (this module's own scope cut), the relevance that would otherwise "leak" into the non-existent state before t=0 is provably exactly zero (both epsilon-rule numerators are proportional to c_prev/h_prev, both == 0 at t=0) – end-to-end conservation is exact up to the epsilon stabilizer, verified numerically by dedicated conservation tests.

Constructor & Destructor Documentation

◆ LSTMModule()

pulsatrix::LSTMModule::LSTMModule ( int64_t  input_size,
int64_t  hidden_size,
DeviceBackend *  backend 
)

Constructs an LSTM layer with zero-initialized weights/biases.

Parameters
input_sizeInput feature dimension.
hidden_sizeHidden (and cell) state dimension.
backendBackend to allocate/compute through. Not owned; must outlive this module.
Exceptions
std::invalid_argumentif input_size <= 0 or hidden_size <= 0 – external boundary (construction arguments can originate from Phase 5's Python bindings with no upstream validation).

Member Function Documentation

◆ backward()

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

Real backpropagation-through-time (BPTT) across all four gates and the cell carry: accumulates every per-gate input weight, recurrent weight and bias gradient across every timestep into the same buffers via Tensor::accumulate().

Parameters
grad_outputGradient w.r.t. this module's output. Must be (N, L, hidden_size) matching the most recent forward() call's output shape.
Returns
Gradient w.r.t. this module's input, shape (N, L, input_size).
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 5): every step runs through DeviceBackend primitives, bit-identical to the former host loops on CPU.

Implements pulsatrix::Module.

◆ bias_f()

const Tensor & pulsatrix::LSTMModule::bias_f ( ) const
inline

◆ bias_f_grad()

const Tensor & pulsatrix::LSTMModule::bias_f_grad ( ) const
inline

◆ bias_g()

const Tensor & pulsatrix::LSTMModule::bias_g ( ) const
inline

◆ bias_g_grad()

const Tensor & pulsatrix::LSTMModule::bias_g_grad ( ) const
inline

◆ bias_i()

const Tensor & pulsatrix::LSTMModule::bias_i ( ) const
inline

◆ bias_i_grad()

const Tensor & pulsatrix::LSTMModule::bias_i_grad ( ) const
inline

◆ bias_o()

const Tensor & pulsatrix::LSTMModule::bias_o ( ) const
inline

◆ bias_o_grad()

const Tensor & pulsatrix::LSTMModule::bias_o_grad ( ) const
inline

◆ compute_device()

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

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

Reimplemented from pulsatrix::Module.

◆ forward_impl()

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

The actual forward computation – per-timestep tied-weight gated recurrence.

Exceptions
std::invalid_argumentif input isn't rank-3 (N, L, input_size), or its last dimension doesn't match input_size.
Note
Device-generic (GPU-native-kernels Mission 5): every step runs through DeviceBackend primitives, bit-identical to the former host loops on CPU.

Implements pulsatrix::Module.

◆ named_parameters()

std::vector< NamedParamRef > pulsatrix::LSTMModule::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::LSTMModule::op_type ( ) const
inlineoverridevirtual

Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule (no enum change needed for this mission).

Implements pulsatrix::Module.

◆ propagate_relevance()

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

Arras et al. 2019 gate-signal LRP: gates conduct, signals receive. See the class-level note for the full per-timestep redistribution.

Parameters
relevance_outRelevance at this module's output. Must be (N, L, hidden_size) matching the most recent forward() call's output shape.
configSelects epsilon.
Returns
Relevance at this module's input, shape (N, L, input_size). Conserves up to the epsilon stabilizer – see the class-level note.
Exceptions
std::logic_errorif forward() has never been called.
std::invalid_argumentif relevance_out's shape doesn't match the cached forward output shape.

Implements pulsatrix::Module.

◆ set_bias_f()

void pulsatrix::LSTMModule::set_bias_f ( std::initializer_list< float >  values)

Overwrites the forget-gate bias buffer – test/initialization use only.

◆ set_bias_g()

void pulsatrix::LSTMModule::set_bias_g ( std::initializer_list< float >  values)

Overwrites the cell-candidate bias buffer – test/initialization use only.

◆ set_bias_i()

void pulsatrix::LSTMModule::set_bias_i ( std::initializer_list< float >  values)

Overwrites the input-gate bias buffer – test/initialization use only.

◆ set_bias_o()

void pulsatrix::LSTMModule::set_bias_o ( std::initializer_list< float >  values)

Overwrites the output-gate bias buffer – test/initialization use only.

◆ set_weight_hf()

void pulsatrix::LSTMModule::set_weight_hf ( std::initializer_list< float >  values)

Overwrites the hidden-to-forget-gate weight buffer – test/initialization use only.

◆ set_weight_hg()

void pulsatrix::LSTMModule::set_weight_hg ( std::initializer_list< float >  values)

Overwrites the hidden-to-cell-candidate weight buffer – test/initialization use only.

◆ set_weight_hi()

void pulsatrix::LSTMModule::set_weight_hi ( std::initializer_list< float >  values)

Overwrites the hidden-to-input-gate weight buffer – test/initialization use only.

◆ set_weight_ho()

void pulsatrix::LSTMModule::set_weight_ho ( std::initializer_list< float >  values)

Overwrites the hidden-to-output-gate weight buffer – test/initialization use only.

◆ set_weight_xf()

void pulsatrix::LSTMModule::set_weight_xf ( std::initializer_list< float >  values)

Overwrites the input-to-forget-gate weight buffer – test/initialization use only.

◆ set_weight_xg()

void pulsatrix::LSTMModule::set_weight_xg ( std::initializer_list< float >  values)

Overwrites the input-to-cell-candidate weight buffer – test/initialization use only.

◆ set_weight_xi()

void pulsatrix::LSTMModule::set_weight_xi ( std::initializer_list< float >  values)

Overwrites the input-to-input-gate weight buffer – test/initialization use only.

◆ set_weight_xo()

void pulsatrix::LSTMModule::set_weight_xo ( std::initializer_list< float >  values)

Overwrites the input-to-output-gate weight buffer – test/initialization use only.

◆ weight_hf()

const Tensor & pulsatrix::LSTMModule::weight_hf ( ) const
inline

◆ weight_hf_grad()

const Tensor & pulsatrix::LSTMModule::weight_hf_grad ( ) const
inline

◆ weight_hg()

const Tensor & pulsatrix::LSTMModule::weight_hg ( ) const
inline

◆ weight_hg_grad()

const Tensor & pulsatrix::LSTMModule::weight_hg_grad ( ) const
inline

◆ weight_hi()

const Tensor & pulsatrix::LSTMModule::weight_hi ( ) const
inline

◆ weight_hi_grad()

const Tensor & pulsatrix::LSTMModule::weight_hi_grad ( ) const
inline

◆ weight_ho()

const Tensor & pulsatrix::LSTMModule::weight_ho ( ) const
inline

◆ weight_ho_grad()

const Tensor & pulsatrix::LSTMModule::weight_ho_grad ( ) const
inline

◆ weight_xf()

const Tensor & pulsatrix::LSTMModule::weight_xf ( ) const
inline

◆ weight_xf_grad()

const Tensor & pulsatrix::LSTMModule::weight_xf_grad ( ) const
inline

◆ weight_xg()

const Tensor & pulsatrix::LSTMModule::weight_xg ( ) const
inline

◆ weight_xg_grad()

const Tensor & pulsatrix::LSTMModule::weight_xg_grad ( ) const
inline

◆ weight_xi()

const Tensor & pulsatrix::LSTMModule::weight_xi ( ) const
inline

◆ weight_xi_grad()

const Tensor & pulsatrix::LSTMModule::weight_xi_grad ( ) const
inline

◆ weight_xo()

const Tensor & pulsatrix::LSTMModule::weight_xo ( ) const
inline

◆ weight_xo_grad()

const Tensor & pulsatrix::LSTMModule::weight_xo_grad ( ) const
inline

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