pulsatrix
Loading...
Searching...
No Matches
lstm_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
57class LSTMModule : public Module {
58public:
68 LSTMModule(int64_t input_size, int64_t hidden_size, DeviceBackend* backend);
69
83 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
84
87 [[nodiscard]] OpType op_type() const override { return OpType::Recurrent; }
88
90 void set_weight_xi(std::initializer_list<float> values);
92 void set_weight_hi(std::initializer_list<float> values);
94 void set_bias_i(std::initializer_list<float> values);
96 void set_weight_xf(std::initializer_list<float> values);
98 void set_weight_hf(std::initializer_list<float> values);
100 void set_bias_f(std::initializer_list<float> values);
102 void set_weight_xg(std::initializer_list<float> values);
104 void set_weight_hg(std::initializer_list<float> values);
106 void set_bias_g(std::initializer_list<float> values);
108 void set_weight_xo(std::initializer_list<float> values);
110 void set_weight_ho(std::initializer_list<float> values);
112 void set_bias_o(std::initializer_list<float> values);
113
114 [[nodiscard]] const Tensor& weight_xi() const { return weight_xi_; }
115 [[nodiscard]] const Tensor& weight_hi() const { return weight_hi_; }
116 [[nodiscard]] const Tensor& bias_i() const { return bias_i_; }
117 [[nodiscard]] const Tensor& weight_xf() const { return weight_xf_; }
118 [[nodiscard]] const Tensor& weight_hf() const { return weight_hf_; }
119 [[nodiscard]] const Tensor& bias_f() const { return bias_f_; }
120 [[nodiscard]] const Tensor& weight_xg() const { return weight_xg_; }
121 [[nodiscard]] const Tensor& weight_hg() const { return weight_hg_; }
122 [[nodiscard]] const Tensor& bias_g() const { return bias_g_; }
123 [[nodiscard]] const Tensor& weight_xo() const { return weight_xo_; }
124 [[nodiscard]] const Tensor& weight_ho() const { return weight_ho_; }
125 [[nodiscard]] const Tensor& bias_o() const { return bias_o_; }
126
127 [[nodiscard]] const Tensor& weight_xi_grad() const { return weight_xi_grad_; }
128 [[nodiscard]] const Tensor& weight_hi_grad() const { return weight_hi_grad_; }
129 [[nodiscard]] const Tensor& bias_i_grad() const { return bias_i_grad_; }
130 [[nodiscard]] const Tensor& weight_xf_grad() const { return weight_xf_grad_; }
131 [[nodiscard]] const Tensor& weight_hf_grad() const { return weight_hf_grad_; }
132 [[nodiscard]] const Tensor& bias_f_grad() const { return bias_f_grad_; }
133 [[nodiscard]] const Tensor& weight_xg_grad() const { return weight_xg_grad_; }
134 [[nodiscard]] const Tensor& weight_hg_grad() const { return weight_hg_grad_; }
135 [[nodiscard]] const Tensor& bias_g_grad() const { return bias_g_grad_; }
136 [[nodiscard]] const Tensor& weight_xo_grad() const { return weight_xo_grad_; }
137 [[nodiscard]] const Tensor& weight_ho_grad() const { return weight_ho_grad_; }
138 [[nodiscard]] const Tensor& bias_o_grad() const { return bias_o_grad_; }
139
152 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
153
154 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
155 return {
156 {"weight_xi", {&weight_xi_, &weight_xi_grad_}},
157 {"weight_hi", {&weight_hi_, &weight_hi_grad_}},
158 {"bias_i", {&bias_i_, &bias_i_grad_}},
159 {"weight_xf", {&weight_xf_, &weight_xf_grad_}},
160 {"weight_hf", {&weight_hf_, &weight_hf_grad_}},
161 {"bias_f", {&bias_f_, &bias_f_grad_}},
162 {"weight_xg", {&weight_xg_, &weight_xg_grad_}},
163 {"weight_hg", {&weight_hg_, &weight_hg_grad_}},
164 {"bias_g", {&bias_g_, &bias_g_grad_}},
165 {"weight_xo", {&weight_xo_, &weight_xo_grad_}},
166 {"weight_ho", {&weight_ho_, &weight_ho_grad_}},
167 {"bias_o", {&bias_o_, &bias_o_grad_}},
168 };
169 }
170
171
173 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
174
175protected:
183 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
184
185private:
186 int64_t input_size_;
187 int64_t hidden_size_;
188 DeviceBackend* backend_;
189 Tensor weight_xi_; // (input_size, hidden_size)
190 Tensor weight_hi_; // (hidden_size, hidden_size)
191 Tensor bias_i_; // (hidden_size,)
192 Tensor weight_xf_;
193 Tensor weight_hf_;
194 Tensor bias_f_;
195 Tensor weight_xg_;
196 Tensor weight_hg_;
197 Tensor bias_g_;
198 Tensor weight_xo_;
199 Tensor weight_ho_;
200 Tensor bias_o_;
201 Tensor weight_xi_grad_;
202 Tensor weight_hi_grad_;
203 Tensor bias_i_grad_;
204 Tensor weight_xf_grad_;
205 Tensor weight_hf_grad_;
206 Tensor bias_f_grad_;
207 Tensor weight_xg_grad_;
208 Tensor weight_hg_grad_;
209 Tensor bias_g_grad_;
210 Tensor weight_xo_grad_;
211 Tensor weight_ho_grad_;
212 Tensor bias_o_grad_;
213 Tensor last_input_; // (N, L, input_size)
214 Tensor last_hidden_states_; // (N, L+1, hidden_size); index 0 = h_0 = 0
215 Tensor last_cell_states_; // (N, L+1, hidden_size); index 0 = c_0 = 0
216 Tensor last_gate_i_; // (N, L, hidden_size)
217 Tensor last_gate_f_; // (N, L, hidden_size)
218 Tensor last_gate_g_; // (N, L, hidden_size); the cell candidate (tanh) signal
219 Tensor last_gate_o_; // (N, L, hidden_size)
220 Tensor last_cell_tanh_; // (N, L, hidden_size); tanh(c_t), cached for backward
221 Tensor last_pre_activation_g_; // (N, L, hidden_size); g_t's z EXCLUDING bias, for LRP
222 int64_t last_L_ = 0;
223 bool has_forwarded_ = false;
224};
225
226} // 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 4-gate LSTM recurrence, h_0 = c_0 = 0 (zero-initialized, not learnable – the same deliberate...
Definition lstm_module.hpp:57
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 f...
const Tensor & weight_xf_grad() const
Definition lstm_module.hpp:130
const Tensor & weight_xo() const
Definition lstm_module.hpp:123
const Tensor & weight_xg_grad() const
Definition lstm_module.hpp:133
const Tensor & bias_f() const
Definition lstm_module.hpp:119
Tensor backward(const Tensor &grad_output) override
Real backpropagation-through-time (BPTT) across all four gates and the cell carry: accumulates every ...
const Tensor & weight_hi() const
Definition lstm_module.hpp:115
void set_weight_hf(std::initializer_list< float > values)
Overwrites the hidden-to-forget-gate weight buffer – test/initialization use only.
const Tensor & bias_o_grad() const
Definition lstm_module.hpp:138
void set_weight_xo(std::initializer_list< float > values)
Overwrites the input-to-output-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.
const Tensor & bias_g() const
Definition lstm_module.hpp:122
const Tensor & bias_o() const
Definition lstm_module.hpp:125
const Tensor & weight_ho_grad() const
Definition lstm_module.hpp:137
const Tensor & bias_g_grad() const
Definition lstm_module.hpp:135
void set_weight_hg(std::initializer_list< float > values)
Overwrites the hidden-to-cell-candidate weight buffer – test/initialization use only.
const Tensor & bias_i_grad() const
Definition lstm_module.hpp:129
const Tensor & weight_xi_grad() const
Definition lstm_module.hpp:127
const Tensor & weight_hf_grad() const
Definition lstm_module.hpp:131
const Tensor & weight_xi() const
Definition lstm_module.hpp:114
const Tensor & bias_i() const
Definition lstm_module.hpp:116
const Tensor & weight_hg() const
Definition lstm_module.hpp:121
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-timestep tied-weight gated recurrence.
void set_weight_xg(std::initializer_list< float > values)
Overwrites the input-to-cell-candidate weight 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.
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition lstm_module.hpp:154
const Tensor & weight_xf() const
Definition lstm_module.hpp:117
LSTMModule(int64_t input_size, int64_t hidden_size, DeviceBackend *backend)
Constructs an LSTM layer with zero-initialized weights/biases.
void set_bias_f(std::initializer_list< float > values)
Overwrites the forget-gate bias buffer – test/initialization use only.
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition lstm_module.hpp:173
void set_bias_o(std::initializer_list< float > values)
Overwrites the output-gate bias buffer – test/initialization use only.
const Tensor & weight_hi_grad() const
Definition lstm_module.hpp:128
const Tensor & weight_xg() const
Definition lstm_module.hpp:120
void set_weight_xi(std::initializer_list< float > values)
Overwrites the input-to-input-gate weight buffer – test/initialization use only.
const Tensor & weight_ho() const
Definition lstm_module.hpp:124
const Tensor & bias_f_grad() const
Definition lstm_module.hpp:132
OpType op_type() const override
Recurrent per charter's closed OpType set – a compound accumulate-over-time operation,...
Definition lstm_module.hpp:87
void set_weight_ho(std::initializer_list< float > values)
Overwrites the hidden-to-output-gate weight buffer – test/initialization use only.
const Tensor & weight_xo_grad() const
Definition lstm_module.hpp:136
void set_bias_g(std::initializer_list< float > values)
Overwrites the cell-candidate bias buffer – test/initialization use only.
const Tensor & weight_hg_grad() const
Definition lstm_module.hpp:134
const Tensor & weight_hf() const
Definition lstm_module.hpp:118
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