pulsatrix
Loading...
Searching...
No Matches
mamba_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
80class MambaModule : public Module {
81public:
91 MambaModule(int64_t d_model, int64_t state_size, DeviceBackend* backend);
92
107 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
108
111 [[nodiscard]] OpType op_type() const override { return OpType::Recurrent; }
112
114 void set_W_delta(std::initializer_list<float> values);
116 void set_bias_delta(std::initializer_list<float> values);
118 void set_W_B(std::initializer_list<float> values);
120 void set_W_C(std::initializer_list<float> values);
122 void set_A(std::initializer_list<float> values);
124 void set_D(std::initializer_list<float> values);
125
127 void set_W_delta(const std::vector<float>& values);
129 void set_bias_delta(const std::vector<float>& values);
131 void set_W_B(const std::vector<float>& values);
133 void set_W_C(const std::vector<float>& values);
135 void set_A(const std::vector<float>& values);
137 void set_D(const std::vector<float>& values);
138
139 [[nodiscard]] const Tensor& W_delta() const { return w_delta_; }
140 [[nodiscard]] const Tensor& bias_delta() const { return bias_delta_; }
141 [[nodiscard]] const Tensor& W_B() const { return w_b_; }
142 [[nodiscard]] const Tensor& W_C() const { return w_c_; }
143 [[nodiscard]] const Tensor& A() const { return a_; }
144 [[nodiscard]] const Tensor& D() const { return d_; }
145
146 [[nodiscard]] const Tensor& W_delta_grad() const { return w_delta_grad_; }
147 [[nodiscard]] const Tensor& bias_delta_grad() const { return bias_delta_grad_; }
148 [[nodiscard]] const Tensor& W_B_grad() const { return w_b_grad_; }
149 [[nodiscard]] const Tensor& W_C_grad() const { return w_c_grad_; }
150 [[nodiscard]] const Tensor& A_grad() const { return a_grad_; }
151 [[nodiscard]] const Tensor& D_grad() const { return d_grad_; }
152
168 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
169
170 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
171 return {
172 {"w_delta", {&w_delta_, &w_delta_grad_}},
173 {"bias_delta", {&bias_delta_, &bias_delta_grad_}},
174 {"w_b", {&w_b_, &w_b_grad_}},
175 {"w_c", {&w_c_, &w_c_grad_}},
176 {"a", {&a_, &a_grad_}},
177 {"d", {&d_, &d_grad_}},
178 };
179 }
180
181
183 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
184
185protected:
192 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
193
194private:
195 int64_t d_model_;
196 int64_t state_size_;
197 DeviceBackend* backend_;
198 Tensor w_delta_; // (d_model, d_model)
199 Tensor bias_delta_; // (d_model,)
200 Tensor w_b_; // (d_model, state_size)
201 Tensor w_c_; // (d_model, state_size)
202 Tensor a_; // (d_model, state_size)
203 Tensor d_; // (d_model,)
204 Tensor w_delta_grad_;
205 Tensor bias_delta_grad_;
206 Tensor w_b_grad_;
207 Tensor w_c_grad_;
208 Tensor a_grad_;
209 Tensor d_grad_;
210 Tensor last_input_; // (N, L, d_model)
211 Tensor last_states_; // (N, L+1, d_model, state_size); index 1 along dim 1 = h_0 = 0
212 Tensor last_abar_; // (N, L, d_model, state_size)
213 Tensor last_bbar_; // (N, L, d_model, state_size)
214 Tensor last_z_delta_; // (N, L, d_model); softplus PRE-activation, for its derivative
215 Tensor last_delta_; // (N, L, d_model)
216 Tensor last_b_; // (N, L, state_size)
217 Tensor last_c_; // (N, L, state_size)
218 Tensor last_output_; // (N, L, d_model); y_t, the LRP epsilon-rule denominator at step 1
219 int64_t last_L_ = 0;
220 bool has_forwarded_ = false;
221};
222
223} // 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.
Core Mamba/S6 selective-scan recurrence (Gu & Dao 2023, arXiv:2312.00752), input (N,...
Definition mamba_module.hpp:80
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition mamba_module.hpp:170
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition mamba_module.hpp:183
const Tensor & D() const
Definition mamba_module.hpp:144
OpType op_type() const override
Recurrent per charter's closed OpType set – a compound accumulate-over-time operation,...
Definition mamba_module.hpp:111
void set_bias_delta(std::initializer_list< float > values)
Overwrites the Delta projection bias (d_model,) – test/initialization use only.
const Tensor & A() const
Definition mamba_module.hpp:143
void set_bias_delta(const std::vector< float > &values)
std::vector overload of set_bias_delta().
void set_A(std::initializer_list< float > values)
Overwrites the continuous state matrix A (d_model, state_size) – test/initialization use only.
void set_W_delta(const std::vector< float > &values)
std::vector overload of set_W_delta() – for callers building values programmatically.
void set_W_C(const std::vector< float > &values)
std::vector overload of set_W_C().
const Tensor & bias_delta() const
Definition mamba_module.hpp:140
const Tensor & W_delta() const
Definition mamba_module.hpp:139
MambaModule(int64_t d_model, int64_t state_size, DeviceBackend *backend)
Constructs a selective-scan layer with zero-initialized parameters.
const Tensor & W_delta_grad() const
Definition mamba_module.hpp:146
Tensor backward(const Tensor &grad_output) override
Real backpropagation-through-time across the selective scan: accumulates W_delta/bias_delta/W_B/W_C/A...
void set_D(std::initializer_list< float > values)
Overwrites the skip/feedthrough vector D (d_model,) – test/initialization use only.
const Tensor & W_B_grad() const
Definition mamba_module.hpp:148
const Tensor & bias_delta_grad() const
Definition mamba_module.hpp:147
const Tensor & W_C() const
Definition mamba_module.hpp:142
void set_D(const std::vector< float > &values)
std::vector overload of set_D().
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-timestep tied-weight selective scan.
void set_W_delta(std::initializer_list< float > values)
Overwrites the input-to-Delta projection weight (d_model, d_model) – test/initialization use only.
void set_W_B(const std::vector< float > &values)
std::vector overload of set_W_B().
void set_W_B(std::initializer_list< float > values)
Overwrites the input-to-B projection weight (d_model, state_size) – test/initialization use only.
const Tensor & D_grad() const
Definition mamba_module.hpp:151
void set_A(const std::vector< float > &values)
std::vector overload of set_A().
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
MambaLRP relevance propagation – Abar_t/Bbar_t/C_t/A/D detached and treated as fixed multiplicative c...
const Tensor & A_grad() const
Definition mamba_module.hpp:150
const Tensor & W_C_grad() const
Definition mamba_module.hpp:149
void set_W_C(std::initializer_list< float > values)
Overwrites the input-to-C projection weight (d_model, state_size) – test/initialization use only.
const Tensor & W_B() const
Definition mamba_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