pulsatrix
Loading...
Searching...
No Matches
rnn_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
37class RNNModule : public Module {
38public:
48 RNNModule(int64_t input_size, int64_t hidden_size, DeviceBackend* backend);
49
62 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
63
66 [[nodiscard]] OpType op_type() const override { return OpType::Recurrent; }
67
69 void set_weight_xh(std::initializer_list<float> values);
71 void set_weight_hh(std::initializer_list<float> values);
73 void set_bias(std::initializer_list<float> values);
74
75 [[nodiscard]] const Tensor& weight_xh() const { return weight_xh_; }
76 [[nodiscard]] const Tensor& weight_hh() const { return weight_hh_; }
77 [[nodiscard]] const Tensor& bias() const { return bias_; }
78 [[nodiscard]] const Tensor& weight_xh_grad() const { return weight_xh_grad_; }
79 [[nodiscard]] const Tensor& weight_hh_grad() const { return weight_hh_grad_; }
80 [[nodiscard]] const Tensor& bias_grad() const { return bias_grad_; }
81
94 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
95
96 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
97 return {
98 {"weight_xh", {&weight_xh_, &weight_xh_grad_}},
99 {"weight_hh", {&weight_hh_, &weight_hh_grad_}},
100 {"bias", {&bias_, &bias_grad_}},
101 };
102 }
103
104
106 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
107
108protected:
116 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
117
118private:
119 int64_t input_size_;
120 int64_t hidden_size_;
121 DeviceBackend* backend_;
122 Tensor weight_xh_; // (input_size, hidden_size)
123 Tensor weight_hh_; // (hidden_size, hidden_size)
124 Tensor bias_; // (hidden_size,)
125 Tensor weight_xh_grad_;
126 Tensor weight_hh_grad_;
127 Tensor bias_grad_;
128 Tensor last_input_; // (N, L, input_size)
129 Tensor last_hidden_states_; // (N, L+1, hidden_size); index 0 = h_0 = 0
130 Tensor last_pre_activation_; // (N, L, hidden_size); z_t EXCLUDING bias, cached for LRP
131 int64_t last_L_ = 0;
132 bool has_forwarded_ = false;
133};
134
135} // 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.
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
h_t = tanh(x_t @ W_xh + h_{t-1} @ W_hh + b_h), h_0 = 0 (zero-initialized, not learnable – a deliberat...
Definition rnn_module.hpp:37
void set_bias(std::initializer_list< float > values)
Overwrites the hidden bias buffer – test/initialization use only.
const Tensor & weight_hh() const
Definition rnn_module.hpp:76
const Tensor & bias() const
Definition rnn_module.hpp:77
const Tensor & weight_xh_grad() const
Definition rnn_module.hpp:78
void set_weight_xh(std::initializer_list< float > values)
Overwrites the input-to-hidden weight buffer – test/initialization use only.
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-timestep tied-weight recurrence.
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition rnn_module.hpp:96
RNNModule(int64_t input_size, int64_t hidden_size, DeviceBackend *backend)
Constructs an RNN layer with zero-initialized weights/bias.
const Tensor & bias_grad() const
Definition rnn_module.hpp:80
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Epsilon/z-rule LRP relevance propagation, generalized to two weighted sources, tanh treated as identi...
const Tensor & weight_xh() const
Definition rnn_module.hpp:75
void set_weight_hh(std::initializer_list< float > values)
Overwrites the hidden-to-hidden weight 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 rnn_module.hpp:106
Tensor backward(const Tensor &grad_output) override
Real backpropagation-through-time (BPTT): accumulates W_xh/W_hh/b_h gradients across every timestep i...
OpType op_type() const override
Recurrent per charter's closed OpType set – a compound accumulate-over-time operation,...
Definition rnn_module.hpp:66
const Tensor & weight_hh_grad() const
Definition rnn_module.hpp:79
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