pulsatrix
Loading...
Searching...
No Matches
rwkv_module.hpp
Go to the documentation of this file.
1
7#pragma once
8
9#include <optional>
10#include <initializer_list>
11#include <vector>
12
13#include "pulsatrix/module.hpp"
14
15namespace pulsatrix {
16
93class RWKVModule : public Module {
94public:
102 RWKVModule(int64_t d_model, DeviceBackend* backend);
103
118 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
119
123 [[nodiscard]] OpType op_type() const override { return OpType::Recurrent; }
124
126 void set_W_r(std::initializer_list<float> values);
128 void set_W_k(std::initializer_list<float> values);
130 void set_W_v(std::initializer_list<float> values);
132 void set_W_o(std::initializer_list<float> values);
134 void set_w(std::initializer_list<float> values);
136 void set_u(std::initializer_list<float> values);
138 void set_mu_r(std::initializer_list<float> values);
140 void set_mu_k(std::initializer_list<float> values);
142 void set_mu_v(std::initializer_list<float> values);
143
145 void set_W_r(const std::vector<float>& values);
147 void set_W_k(const std::vector<float>& values);
149 void set_W_v(const std::vector<float>& values);
151 void set_W_o(const std::vector<float>& values);
153 void set_w(const std::vector<float>& values);
155 void set_u(const std::vector<float>& values);
157 void set_mu_r(const std::vector<float>& values);
159 void set_mu_k(const std::vector<float>& values);
161 void set_mu_v(const std::vector<float>& values);
162
164 [[nodiscard]] const Tensor& W_r() const { return w_r_; }
166 [[nodiscard]] const Tensor& W_k() const { return w_k_; }
168 [[nodiscard]] const Tensor& W_v() const { return w_v_; }
170 [[nodiscard]] const Tensor& W_o() const { return w_o_; }
172 [[nodiscard]] const Tensor& w() const { return w_; }
174 [[nodiscard]] const Tensor& u() const { return u_; }
176 [[nodiscard]] const Tensor& mu_r() const { return mu_r_; }
178 [[nodiscard]] const Tensor& mu_k() const { return mu_k_; }
180 [[nodiscard]] const Tensor& mu_v() const { return mu_v_; }
181
183 [[nodiscard]] const Tensor& W_r_grad() const { return w_r_grad_; }
185 [[nodiscard]] const Tensor& W_k_grad() const { return w_k_grad_; }
187 [[nodiscard]] const Tensor& W_v_grad() const { return w_v_grad_; }
189 [[nodiscard]] const Tensor& W_o_grad() const { return w_o_grad_; }
191 [[nodiscard]] const Tensor& w_grad() const { return w_grad_; }
193 [[nodiscard]] const Tensor& u_grad() const { return u_grad_; }
195 [[nodiscard]] const Tensor& mu_r_grad() const { return mu_r_grad_; }
197 [[nodiscard]] const Tensor& mu_k_grad() const { return mu_k_grad_; }
199 [[nodiscard]] const Tensor& mu_v_grad() const { return mu_v_grad_; }
200
226 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
227
228 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
229 return {
230 {"w_r", {&w_r_, &w_r_grad_}},
231 {"w_k", {&w_k_, &w_k_grad_}},
232 {"w_v", {&w_v_, &w_v_grad_}},
233 {"w_o", {&w_o_, &w_o_grad_}},
234 {"w", {&w_, &w_grad_}},
235 {"u", {&u_, &u_grad_}},
236 {"mu_r", {&mu_r_, &mu_r_grad_}},
237 {"mu_k", {&mu_k_, &mu_k_grad_}},
238 {"mu_v", {&mu_v_, &mu_v_grad_}},
239 };
240 }
241
242
244 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
245
246protected:
253 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
254
255private:
256 int64_t d_model_;
257 DeviceBackend* backend_;
258 Tensor w_r_; // (d_model, d_model)
259 Tensor w_k_; // (d_model, d_model)
260 Tensor w_v_; // (d_model, d_model)
261 Tensor w_o_; // (d_model, d_model)
262 Tensor w_; // (d_model,); decay = exp(-w)
263 Tensor u_; // (d_model,)
264 Tensor mu_r_; // (d_model,)
265 Tensor mu_k_; // (d_model,)
266 Tensor mu_v_; // (d_model,)
267 Tensor w_r_grad_;
268 Tensor w_k_grad_;
269 Tensor w_v_grad_;
270 Tensor w_o_grad_;
271 Tensor w_grad_;
272 Tensor u_grad_;
273 Tensor mu_r_grad_;
274 Tensor mu_k_grad_;
275 Tensor mu_v_grad_;
276 Tensor last_input_; // (N, L, d_model)
277 Tensor last_xr_; // (N, L, d_model); token-shifted receptance input
278 Tensor last_xk_; // (N, L, d_model)
279 Tensor last_xv_; // (N, L, d_model)
280 Tensor last_r_; // (N, L, d_model); sigmoid POST-activation (its own derivative source)
281 Tensor last_k_; // (N, L, d_model)
282 Tensor last_v_; // (N, L, d_model)
283 Tensor last_e_; // (N, L, d_model); exp(u + k_t)
284 Tensor last_kk_; // (N, L, d_model); exp(k_t)
285 Tensor last_num_; // (N, L, d_model)
286 Tensor last_den_; // (N, L, d_model)
287 Tensor last_wkv_; // (N, L, d_model)
288 Tensor last_a_; // (N, L+1, d_model); index 0 along dim 1 = a_0 = 0
289 Tensor last_b_; // (N, L+1, d_model); index 0 along dim 1 = b_0 = 0
290 int64_t last_L_ = 0;
291 bool has_forwarded_ = false;
292};
293
294} // 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 RWKV-4 time-mixing block (Peng et al. 2023, arXiv:2305.13048), input (N, L, d_model) -> output (...
Definition rwkv_module.hpp:93
void set_u(std::initializer_list< float > values)
Overwrites the per-channel current-token bonus u (d_model,) – test/initialization use only.
Tensor backward(const Tensor &grad_output) override
Real backpropagation-through-time across the WKV recurrence: accumulates all nine parameter gradients...
const Tensor & W_v_grad() const
Accumulated gradient w.r.t. W_v.
Definition rwkv_module.hpp:187
const Tensor & mu_r_grad() const
Accumulated gradient w.r.t. mu_r.
Definition rwkv_module.hpp:195
const Tensor & w() const
The per-channel decay rate w (d_model,); the applied decay is exp(-w).
Definition rwkv_module.hpp:172
void set_W_r(const std::vector< float > &values)
std::vector overload of set_W_r() – for callers building values programmatically.
const Tensor & W_r_grad() const
Accumulated gradient w.r.t. W_r.
Definition rwkv_module.hpp:183
const Tensor & W_v() const
The value projection weight (d_model, d_model).
Definition rwkv_module.hpp:168
const Tensor & W_k_grad() const
Accumulated gradient w.r.t. W_k.
Definition rwkv_module.hpp:185
const Tensor & mu_k_grad() const
Accumulated gradient w.r.t. mu_k.
Definition rwkv_module.hpp:197
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition rwkv_module.hpp:244
void set_mu_k(std::initializer_list< float > values)
Overwrites the key token-shift mix ratio (d_model,) – test/initialization use only.
const Tensor & W_o_grad() const
Accumulated gradient w.r.t. W_o.
Definition rwkv_module.hpp:189
void set_mu_r(const std::vector< float > &values)
std::vector overload of set_mu_r().
void set_W_k(std::initializer_list< float > values)
Overwrites the key projection weight (d_model, d_model) – test/initialization use only.
const Tensor & W_o() const
The output projection weight (d_model, d_model).
Definition rwkv_module.hpp:170
const Tensor & u_grad() const
Accumulated gradient w.r.t. u.
Definition rwkv_module.hpp:193
const Tensor & mu_k() const
The key token-shift mix ratio (d_model,).
Definition rwkv_module.hpp:178
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition rwkv_module.hpp:228
const Tensor & W_r() const
The receptance projection weight (d_model, d_model).
Definition rwkv_module.hpp:164
void set_mu_r(std::initializer_list< float > values)
Overwrites the receptance token-shift mix ratio (d_model,) – test/initialization use only.
void set_W_v(const std::vector< float > &values)
std::vector overload of set_W_v().
const Tensor & mu_v() const
The value token-shift mix ratio (d_model,).
Definition rwkv_module.hpp:180
void set_W_o(std::initializer_list< float > values)
Overwrites the output projection weight (d_model, d_model) – test/initialization use only.
void set_u(const std::vector< float > &values)
std::vector overload of set_u().
void set_W_r(std::initializer_list< float > values)
Overwrites the receptance projection weight (d_model, d_model) – test/initialization use only.
const Tensor & mu_r() const
The receptance token-shift mix ratio (d_model,).
Definition rwkv_module.hpp:176
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
The original derived LRP rule (see the class-level note): MambaLRP's detach-the-gate technique adapte...
OpType op_type() const override
Recurrent per charter's closed OpType set – a compound accumulate-over-time operation,...
Definition rwkv_module.hpp:123
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-timestep tied-weight WKV recurrence.
void set_mu_k(const std::vector< float > &values)
std::vector overload of set_mu_k().
void set_mu_v(const std::vector< float > &values)
std::vector overload of set_mu_v().
void set_W_o(const std::vector< float > &values)
std::vector overload of set_W_o().
void set_mu_v(std::initializer_list< float > values)
Overwrites the value token-shift mix ratio (d_model,) – test/initialization use only.
const Tensor & W_k() const
The key projection weight (d_model, d_model).
Definition rwkv_module.hpp:166
const Tensor & u() const
The per-channel current-token bonus u (d_model,).
Definition rwkv_module.hpp:174
const Tensor & w_grad() const
Accumulated gradient w.r.t. w.
Definition rwkv_module.hpp:191
void set_w(std::initializer_list< float > values)
Overwrites the per-channel decay rate w (d_model,), decay = exp(-w) – test/initialization use only.
void set_W_k(const std::vector< float > &values)
std::vector overload of set_W_k().
void set_w(const std::vector< float > &values)
std::vector overload of set_w().
const Tensor & mu_v_grad() const
Accumulated gradient w.r.t. mu_v.
Definition rwkv_module.hpp:199
RWKVModule(int64_t d_model, DeviceBackend *backend)
Constructs an RWKV time-mixing layer with zero-initialized parameters.
void set_W_v(std::initializer_list< float > values)
Overwrites the value projection weight (d_model, d_model) – test/initialization use only.
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