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

Standard GRU recurrence (Cho et al. 2014), h_0 = 0 (zero-initialized, not learnable – the same deliberate scope cut RNNModule/LSTMModule made, and the same thing that makes this module's conservation exact; see the LRP note below): z_t = sigmoid(x_t @ W_xz + h_{t-1} @ W_hz + b_z) (update gate), r_t = sigmoid(x_t @ W_xr + h_{t-1} @ W_hr + b_r) (reset gate), hn_prev_t = h_{t-1} @ W_hn (internal projection, no bias), n_t = tanh(x_t @ W_xn + r_t * hn_prev_t + b_n) (candidate), h_t = (1 - z_t) * h_{t-1} + z_t * n_t. Note the reset gate multiplies the projected previous hidden state hn_prev_t, not h_{t-1} itself – that projection is a distinct cached intermediate, and it is what gives GRU's LRP rule a different shape from LSTM's. Input (N, L, input_size) -> output (N, L, hidden_size), the full hidden-state sequence (matches RNNModule's/LSTMModule's convention). Single layer, no bidirectional/multi-layer/variable-length support. More...

#include <gru_module.hpp>

Inheritance diagram for pulsatrix::GRUModule:
Collaboration diagram for pulsatrix::GRUModule:

Public Member Functions

 GRUModule (int64_t input_size, int64_t hidden_size, DeviceBackend *backend)
 Constructs a GRU layer with zero-initialized weights/biases.
 
Tensor backward (const Tensor &grad_output) override
 Real backpropagation-through-time (BPTT) across both gates, the reset-gated candidate and the convex (1-z_t)/z_t hidden carry: accumulates every weight and bias gradient across every timestep into the same buffers via Tensor::accumulate(). The gradient w.r.t. h_{t-1} sums three contributions (direct (1-z_t) carry, both gates' recurrent branches, and the candidate's W_hn projection branch), mirroring the two-path relevance structure documented on the class.
 
OpType op_type () const override
 Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule/LSTMModule (no enum change for this mission).
 
void set_weight_xz (std::initializer_list< float > values)
 Overwrites the input-to-update-gate weight buffer – test/initialization use only.
 
void set_weight_hz (std::initializer_list< float > values)
 Overwrites the hidden-to-update-gate weight buffer – test/initialization use only.
 
void set_bias_z (std::initializer_list< float > values)
 Overwrites the update-gate bias buffer – test/initialization use only.
 
void set_weight_xr (std::initializer_list< float > values)
 Overwrites the input-to-reset-gate weight buffer – test/initialization use only.
 
void set_weight_hr (std::initializer_list< float > values)
 Overwrites the hidden-to-reset-gate weight buffer – test/initialization use only.
 
void set_bias_r (std::initializer_list< float > values)
 Overwrites the reset-gate bias buffer – test/initialization use only.
 
void set_weight_xn (std::initializer_list< float > values)
 Overwrites the input-to-candidate weight buffer – test/initialization use only.
 
void set_weight_hn (std::initializer_list< float > values)
 Overwrites the hidden-to-candidate projection weight buffer (the W_hn whose output the reset gate multiplies) – test/initialization use only.
 
void set_bias_n (std::initializer_list< float > values)
 Overwrites the candidate bias buffer – test/initialization use only.
 
const Tensor & weight_xz () const
 
const Tensor & weight_hz () const
 
const Tensor & bias_z () const
 
const Tensor & weight_xr () const
 
const Tensor & weight_hr () const
 
const Tensor & bias_r () const
 
const Tensor & weight_xn () const
 
const Tensor & weight_hn () const
 
const Tensor & bias_n () const
 
const Tensor & weight_xz_grad () const
 
const Tensor & weight_hz_grad () const
 
const Tensor & bias_z_grad () const
 
const Tensor & weight_xr_grad () const
 
const Tensor & weight_hr_grad () const
 
const Tensor & bias_r_grad () const
 
const Tensor & weight_xn_grad () const
 
const Tensor & weight_hn_grad () const
 
const Tensor & bias_n_grad () const
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 Arras et al. 2019 gate-signal LRP, extended to GRU: gates conduct, signals receive, and the carried h_{t-1} relevance sums a direct and a candidate-path contribution. 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 GRU recurrence (Cho et al. 2014), h_0 = 0 (zero-initialized, not learnable – the same deliberate scope cut RNNModule/LSTMModule made, and the same thing that makes this module's conservation exact; see the LRP note below): z_t = sigmoid(x_t @ W_xz + h_{t-1} @ W_hz + b_z) (update gate), r_t = sigmoid(x_t @ W_xr + h_{t-1} @ W_hr + b_r) (reset gate), hn_prev_t = h_{t-1} @ W_hn (internal projection, no bias), n_t = tanh(x_t @ W_xn + r_t * hn_prev_t + b_n) (candidate), h_t = (1 - z_t) * h_{t-1} + z_t * n_t. Note the reset gate multiplies the projected previous hidden state hn_prev_t, not h_{t-1} itself – that projection is a distinct cached intermediate, and it is what gives GRU's LRP rule a different shape from LSTM's. Input (N, L, input_size) -> output (N, L, hidden_size), the full hidden-state sequence (matches RNNModule's/LSTMModule's convention). Single layer, no bidirectional/multi-layer/variable-length 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 plus add/mul/axpby for the gate blend, recurrent_cell(GruBackward / GruLrp), accumulate_rows, lrp_linear, gru_lrp_hprev), reproducing the former host loops' evaluation order so CPU results are bit-identical.
LRP rule (gate-signal principle of Arras et al. 2019, "Explaining Recurrent Neural Network Predictions in Sentiment Analysis", extended here to GRU's reset-gate-inside-the-preactivation structure – Arras et al. cover LSTM explicitly, so the same principle is re-derived below against GRU's equation shape). The gates z_t and r_t are pure conductors, never relevance recipients; every multiplicative interaction is an exact bilinear identity summing to its own result, so relevance splits across the signal operands only, with gate values acting as fixed multiplicative weights. Per timestep, in reverse time order:
  1. h_t = (1-z_t)*h_{t-1} + z_t*n_t is a two-term weighted sum whose denominator is exactly h_t; R(h_t) splits between R(h_{t-1})_direct and R(n_t) in proportion to (1-z_t)*h_{t-1} and z_t*n_t, via the same epsilon/z-rule shape RNNModule uses.
  2. n_t = tanh(n_pre) is identity pass-through (same pointwise-nonlinearity precedent as ReluModule/RNNModule, Montavon et al. 2019). 3/4. n_pre (bias excluded, same "bias has no input feature to redistribute to" convention as every prior module) = x_t@W_xn + r_t*hn_prev_t. R(n_pre) splits between the two terms in proportion to their values; the x_t@W_xn share is then redistributed across x_t's features weighted by W_xn. Those two steps compose into one epsilon-stabilized redistribution over the shared n_pre denominator.
  3. r_t is a pure gate, so ALL of the r_t*hn_prev_t term's relevance passes to R(hn_prev_t) and none to R(r_t).
  4. hn_prev_t = h_{t-1}@W_hn is a standard single-source linear combination; R(hn_prev_t) is redistributed across h_{t-1}'s features weighted by W_hn (epsilon rule). Call this R(h_{t-1})_via_candidate.
  5. The structural feature that distinguishes GRU from LSTM here: the carried relevance accumulator for h_{t-1} is fed by TWO separate paths and must SUM them – R(h_{t-1}) = R(h_{t-1})_direct (step 1) + R(h_{t-1})_via_candidate (step 6). LSTMModule's cell carry is single-path per accumulator; GRU's is not. Letting the second contribution overwrite rather than accumulate onto the first silently breaks conservation.
  6. z_pre/r_pre (the gates' own linear combinations) never receive or emit relevance at all – consistent with LSTMModule's i_t/f_t/o_t. The gate values are consumed purely as precomputed multiplicative coefficients during the backward pass. Because h_0 is zero (this module's own scope cut), both contributions to the relevance that would "leak" into the non-existent state before t=0 are provably exactly zero (both numerators are proportional to h_prev, which is 0 at t=0) – end-to-end conservation holds up to the epsilon stabilizers, verified numerically by dedicated conservation tests.

Constructor & Destructor Documentation

◆ GRUModule()

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

Constructs a GRU layer with zero-initialized weights/biases.

Parameters
input_sizeInput feature dimension.
hidden_sizeHidden 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::GRUModule::backward ( const Tensor &  grad_output)
overridevirtual

Real backpropagation-through-time (BPTT) across both gates, the reset-gated candidate and the convex (1-z_t)/z_t hidden carry: accumulates every weight and bias gradient across every timestep into the same buffers via Tensor::accumulate(). The gradient w.r.t. h_{t-1} sums three contributions (direct (1-z_t) carry, both gates' recurrent branches, and the candidate's W_hn projection branch), mirroring the two-path relevance structure documented on the class.

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_n()

const Tensor & pulsatrix::GRUModule::bias_n ( ) const
inline

◆ bias_n_grad()

const Tensor & pulsatrix::GRUModule::bias_n_grad ( ) const
inline

◆ bias_r()

const Tensor & pulsatrix::GRUModule::bias_r ( ) const
inline

◆ bias_r_grad()

const Tensor & pulsatrix::GRUModule::bias_r_grad ( ) const
inline

◆ bias_z()

const Tensor & pulsatrix::GRUModule::bias_z ( ) const
inline

◆ bias_z_grad()

const Tensor & pulsatrix::GRUModule::bias_z_grad ( ) const
inline

◆ compute_device()

std::optional< DeviceType > pulsatrix::GRUModule::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::GRUModule::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::GRUModule::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::GRUModule::op_type ( ) const
inlineoverridevirtual

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

Implements pulsatrix::Module.

◆ propagate_relevance()

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

Arras et al. 2019 gate-signal LRP, extended to GRU: gates conduct, signals receive, and the carried h_{t-1} relevance sums a direct and a candidate-path contribution. 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 stabilizers – 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_n()

void pulsatrix::GRUModule::set_bias_n ( std::initializer_list< float >  values)

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

◆ set_bias_r()

void pulsatrix::GRUModule::set_bias_r ( std::initializer_list< float >  values)

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

◆ set_bias_z()

void pulsatrix::GRUModule::set_bias_z ( std::initializer_list< float >  values)

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

◆ set_weight_hn()

void pulsatrix::GRUModule::set_weight_hn ( std::initializer_list< float >  values)

Overwrites the hidden-to-candidate projection weight buffer (the W_hn whose output the reset gate multiplies) – test/initialization use only.

◆ set_weight_hr()

void pulsatrix::GRUModule::set_weight_hr ( std::initializer_list< float >  values)

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

◆ set_weight_hz()

void pulsatrix::GRUModule::set_weight_hz ( std::initializer_list< float >  values)

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

◆ set_weight_xn()

void pulsatrix::GRUModule::set_weight_xn ( std::initializer_list< float >  values)

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

◆ set_weight_xr()

void pulsatrix::GRUModule::set_weight_xr ( std::initializer_list< float >  values)

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

◆ set_weight_xz()

void pulsatrix::GRUModule::set_weight_xz ( std::initializer_list< float >  values)

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

◆ weight_hn()

const Tensor & pulsatrix::GRUModule::weight_hn ( ) const
inline

◆ weight_hn_grad()

const Tensor & pulsatrix::GRUModule::weight_hn_grad ( ) const
inline

◆ weight_hr()

const Tensor & pulsatrix::GRUModule::weight_hr ( ) const
inline

◆ weight_hr_grad()

const Tensor & pulsatrix::GRUModule::weight_hr_grad ( ) const
inline

◆ weight_hz()

const Tensor & pulsatrix::GRUModule::weight_hz ( ) const
inline

◆ weight_hz_grad()

const Tensor & pulsatrix::GRUModule::weight_hz_grad ( ) const
inline

◆ weight_xn()

const Tensor & pulsatrix::GRUModule::weight_xn ( ) const
inline

◆ weight_xn_grad()

const Tensor & pulsatrix::GRUModule::weight_xn_grad ( ) const
inline

◆ weight_xr()

const Tensor & pulsatrix::GRUModule::weight_xr ( ) const
inline

◆ weight_xr_grad()

const Tensor & pulsatrix::GRUModule::weight_xr_grad ( ) const
inline

◆ weight_xz()

const Tensor & pulsatrix::GRUModule::weight_xz ( ) const
inline

◆ weight_xz_grad()

const Tensor & pulsatrix::GRUModule::weight_xz_grad ( ) const
inline

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