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...
|
| | 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).
|
| |
| 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.
|
| |
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:
- 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.
- 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.
- 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).
- 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.
- 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.
- 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.