pulsatrix
Loading...
Searching...
No Matches
gru_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <optional>
8#include <initializer_list>
9#include <vector>
10
11#include "pulsatrix/module.hpp"
12
13namespace pulsatrix {
14
74class GRUModule : public Module {
75public:
85 GRUModule(int64_t input_size, int64_t hidden_size, DeviceBackend* backend);
86
104 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
105
108 [[nodiscard]] OpType op_type() const override { return OpType::Recurrent; }
109
111 void set_weight_xz(std::initializer_list<float> values);
113 void set_weight_hz(std::initializer_list<float> values);
115 void set_bias_z(std::initializer_list<float> values);
117 void set_weight_xr(std::initializer_list<float> values);
119 void set_weight_hr(std::initializer_list<float> values);
121 void set_bias_r(std::initializer_list<float> values);
123 void set_weight_xn(std::initializer_list<float> values);
126 void set_weight_hn(std::initializer_list<float> values);
128 void set_bias_n(std::initializer_list<float> values);
129
130 [[nodiscard]] const Tensor& weight_xz() const { return weight_xz_; }
131 [[nodiscard]] const Tensor& weight_hz() const { return weight_hz_; }
132 [[nodiscard]] const Tensor& bias_z() const { return bias_z_; }
133 [[nodiscard]] const Tensor& weight_xr() const { return weight_xr_; }
134 [[nodiscard]] const Tensor& weight_hr() const { return weight_hr_; }
135 [[nodiscard]] const Tensor& bias_r() const { return bias_r_; }
136 [[nodiscard]] const Tensor& weight_xn() const { return weight_xn_; }
137 [[nodiscard]] const Tensor& weight_hn() const { return weight_hn_; }
138 [[nodiscard]] const Tensor& bias_n() const { return bias_n_; }
139
140 [[nodiscard]] const Tensor& weight_xz_grad() const { return weight_xz_grad_; }
141 [[nodiscard]] const Tensor& weight_hz_grad() const { return weight_hz_grad_; }
142 [[nodiscard]] const Tensor& bias_z_grad() const { return bias_z_grad_; }
143 [[nodiscard]] const Tensor& weight_xr_grad() const { return weight_xr_grad_; }
144 [[nodiscard]] const Tensor& weight_hr_grad() const { return weight_hr_grad_; }
145 [[nodiscard]] const Tensor& bias_r_grad() const { return bias_r_grad_; }
146 [[nodiscard]] const Tensor& weight_xn_grad() const { return weight_xn_grad_; }
147 [[nodiscard]] const Tensor& weight_hn_grad() const { return weight_hn_grad_; }
148 [[nodiscard]] const Tensor& bias_n_grad() const { return bias_n_grad_; }
149
163 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
164
165 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
166 return {
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_}},
176 };
177 }
178
179
181 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
182
183protected:
191 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
192
193private:
194 int64_t input_size_;
195 int64_t hidden_size_;
196 DeviceBackend* backend_;
197 Tensor weight_xz_; // (input_size, hidden_size)
198 Tensor weight_hz_; // (hidden_size, hidden_size)
199 Tensor bias_z_; // (hidden_size,)
200 Tensor weight_xr_;
201 Tensor weight_hr_;
202 Tensor bias_r_;
203 Tensor weight_xn_;
204 Tensor weight_hn_;
205 Tensor bias_n_;
206 Tensor weight_xz_grad_;
207 Tensor weight_hz_grad_;
208 Tensor bias_z_grad_;
209 Tensor weight_xr_grad_;
210 Tensor weight_hr_grad_;
211 Tensor bias_r_grad_;
212 Tensor weight_xn_grad_;
213 Tensor weight_hn_grad_;
214 Tensor bias_n_grad_;
215 Tensor last_input_; // (N, L, input_size)
216 Tensor last_hidden_states_; // (N, L+1, hidden_size); index 0 = h_0 = 0
217 Tensor last_gate_z_; // (N, L, hidden_size); update gate
218 Tensor last_gate_r_; // (N, L, hidden_size); reset gate
219 Tensor last_candidate_n_; // (N, L, hidden_size); n_t, the tanh candidate signal
220 // hn_prev_t = h_{t-1} @ W_hn is a DISTINCT intermediate from h_{t-1} itself (the reset
221 // gate multiplies the projection, not the raw state), and it is the node the candidate
222 // path's relevance lands on before being redistributed back onto h_{t-1} -- so it is
223 // cached in its own right, for both backward() and propagate_relevance()'s step 6.
224 Tensor last_hn_prev_; // (N, L, hidden_size)
225 Tensor last_pre_activation_n_; // (N, L, hidden_size); n_t's pre-activation EXCLUDING bias
226 int64_t last_L_ = 0;
227 bool has_forwarded_ = false;
228};
229
230} // namespace pulsatrix
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