8#include <initializer_list>
167 {
"weight_xz", {&weight_xz_, &weight_xz_grad_}},
168 {
"weight_hz", {&weight_hz_, &weight_hz_grad_}},
169 {
"bias_z", {&bias_z_, &bias_z_grad_}},
170 {
"weight_xr", {&weight_xr_, &weight_xr_grad_}},
171 {
"weight_hr", {&weight_hr_, &weight_hr_grad_}},
172 {
"bias_r", {&bias_r_, &bias_r_grad_}},
173 {
"weight_xn", {&weight_xn_, &weight_xn_grad_}},
174 {
"weight_hn", {&weight_hn_, &weight_hn_grad_}},
175 {
"bias_n", {&bias_n_, &bias_n_grad_}},
195 int64_t hidden_size_;
216 Tensor last_hidden_states_;
225 Tensor last_pre_activation_n_;
227 bool has_forwarded_ =
false;
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
virtual DeviceType device() const noexcept=0
Which device this backend's buffers reside on.
Standard GRU recurrence (Cho et al. 2014), h_0 = 0 (zero-initialized, not learnable – the same delibe...
Definition gru_module.hpp:74
const Tensor & weight_xz_grad() const
Definition gru_module.hpp:140
const Tensor & bias_z_grad() const
Definition gru_module.hpp:142
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-timestep tied-weight gated recurrence.
const Tensor & bias_n_grad() const
Definition gru_module.hpp:148
const Tensor & weight_xr_grad() const
Definition gru_module.hpp:143
const Tensor & weight_xz() const
Definition gru_module.hpp:130
void set_bias_r(std::initializer_list< float > values)
Overwrites the reset-gate bias buffer – test/initialization use only.
Tensor backward(const Tensor &grad_output) override
Real backpropagation-through-time (BPTT) across both gates, the reset-gated candidate and the convex ...
const Tensor & bias_r_grad() const
Definition gru_module.hpp:145
const Tensor & bias_z() const
Definition gru_module.hpp:132
const Tensor & weight_xn() const
Definition gru_module.hpp:136
void set_weight_hr(std::initializer_list< float > values)
Overwrites the hidden-to-reset-gate weight buffer – test/initialization use only.
const Tensor & weight_xn_grad() const
Definition gru_module.hpp:146
const Tensor & bias_n() const
Definition gru_module.hpp:138
const Tensor & weight_hn_grad() const
Definition gru_module.hpp:147
const Tensor & bias_r() const
Definition gru_module.hpp:135
void set_weight_xz(std::initializer_list< float > values)
Overwrites the input-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.
OpType op_type() const override
Recurrent per charter's closed OpType set – a compound accumulate-over-time operation,...
Definition gru_module.hpp:108
const Tensor & weight_hz() const
Definition gru_module.hpp:131
const Tensor & weight_xr() const
Definition gru_module.hpp:133
void set_bias_n(std::initializer_list< float > values)
Overwrites the candidate bias buffer – test/initialization use only.
const Tensor & weight_hn() const
Definition gru_module.hpp:137
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 mul...
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,...
const Tensor & weight_hr_grad() const
Definition gru_module.hpp:144
void set_weight_xn(std::initializer_list< float > values)
Overwrites the input-to-candidate weight buffer – test/initialization use only.
const Tensor & weight_hr() const
Definition gru_module.hpp:134
GRUModule(int64_t input_size, int64_t hidden_size, DeviceBackend *backend)
Constructs a GRU layer with zero-initialized weights/biases.
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition gru_module.hpp:181
void set_weight_xr(std::initializer_list< float > values)
Overwrites the input-to-reset-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.
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition gru_module.hpp:165
const Tensor & weight_hz_grad() const
Definition gru_module.hpp:141
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16
OpType
The op-type tag a Node carries. Charter Part 2 §3: nodes are tagged by a small closed set of op types...
Definition op_type.hpp:19
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57