pulsatrix
Loading...
Searching...
No Matches
retnet_module.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <optional>
9#include <initializer_list>
10#include <vector>
11
12#include "pulsatrix/module.hpp"
13
14namespace pulsatrix {
15
87class RetNetModule : public Module {
88public:
100 RetNetModule(int64_t d_model, int64_t key_dim, float gamma, DeviceBackend* backend);
101
116 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
117
121 [[nodiscard]] OpType op_type() const override { return OpType::Recurrent; }
122
124 void set_W_Q(std::initializer_list<float> values);
126 void set_W_K(std::initializer_list<float> values);
128 void set_W_V(std::initializer_list<float> values);
129
131 void set_W_Q(const std::vector<float>& values);
133 void set_W_K(const std::vector<float>& values);
135 void set_W_V(const std::vector<float>& values);
136
138 [[nodiscard]] const Tensor& W_Q() const { return w_q_; }
140 [[nodiscard]] const Tensor& W_K() const { return w_k_; }
142 [[nodiscard]] const Tensor& W_V() const { return w_v_; }
143
145 [[nodiscard]] const Tensor& W_Q_grad() const { return w_q_grad_; }
147 [[nodiscard]] const Tensor& W_K_grad() const { return w_k_grad_; }
149 [[nodiscard]] const Tensor& W_V_grad() const { return w_v_grad_; }
150
153 [[nodiscard]] float gamma() const { return gamma_; }
154
177 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
178
179 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
180 return {{"w_q", {&w_q_, &w_q_grad_}}, {"w_k", {&w_k_, &w_k_grad_}}, {"w_v", {&w_v_, &w_v_grad_}}};
181 }
182
183
185 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
186
187protected:
194 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
195
196private:
197 int64_t d_model_;
198 int64_t key_dim_;
199 float gamma_;
200 DeviceBackend* backend_;
201 Tensor w_q_; // (d_model, key_dim)
202 Tensor w_k_; // (d_model, key_dim)
203 Tensor w_v_; // (d_model, d_model)
204 Tensor w_q_grad_;
205 Tensor w_k_grad_;
206 Tensor w_v_grad_;
207 Tensor last_input_; // (N, L, d_model)
208 Tensor last_q_; // (N, L, key_dim)
209 Tensor last_k_; // (N, L, key_dim)
210 Tensor last_v_; // (N, L, d_model)
211 Tensor last_states_; // (N, L+1, key_dim, d_model); index 0 along dim 1 = S_0 = 0
212 int64_t last_L_ = 0;
213 bool has_forwarded_ = false;
214};
215
216} // 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
Core RetNet retention block (Sun et al. 2023, arXiv:2307.08621), recurrent mode, input (N,...
Definition retnet_module.hpp:87
const Tensor & W_V_grad() const
Accumulated gradient w.r.t. W_V.
Definition retnet_module.hpp:149
RetNetModule(int64_t d_model, int64_t key_dim, float gamma, DeviceBackend *backend)
Constructs a retention layer with zero-initialized parameters.
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-timestep tied-weight retention recurrence.
Tensor backward(const Tensor &grad_output) override
Real backpropagation-through-time across the retention recurrence: accumulates all three parameter gr...
void set_W_V(std::initializer_list< float > values)
Overwrites the value projection weight (d_model, d_model) – test/initialization use only.
void set_W_Q(std::initializer_list< float > values)
Overwrites the query projection weight (d_model, key_dim) – test/initialization use only.
const Tensor & W_Q_grad() const
Accumulated gradient w.r.t. W_Q.
Definition retnet_module.hpp:145
float gamma() const
The fixed retention decay. A hyperparameter, not a parameter – it has no gradient and is absent from ...
Definition retnet_module.hpp:153
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition retnet_module.hpp:185
void set_W_K(const std::vector< float > &values)
std::vector overload of set_W_K().
const Tensor & W_K_grad() const
Accumulated gradient w.r.t. W_K.
Definition retnet_module.hpp:147
void set_W_V(const std::vector< float > &values)
std::vector overload of set_W_V().
OpType op_type() const override
Recurrent per charter's closed OpType set – a compound accumulate-over-time operation,...
Definition retnet_module.hpp:121
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition retnet_module.hpp:179
const Tensor & W_K() const
The key projection weight (d_model, key_dim).
Definition retnet_module.hpp:140
const Tensor & W_V() const
The value projection weight (d_model, d_model).
Definition retnet_module.hpp:142
void set_W_K(std::initializer_list< float > values)
Overwrites the key projection weight (d_model, key_dim) – test/initialization use only.
void set_W_Q(const std::vector< float > &values)
std::vector overload of set_W_Q() – for callers building values programmatically.
const Tensor & W_Q() const
The query projection weight (d_model, key_dim).
Definition retnet_module.hpp:138
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
The original derived LRP rule (see the class-level note): unrolls the retention recurrence into Y = G...
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