pulsatrix
Loading...
Searching...
No Matches
sparse_autoencoder.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <cmath>
9#include <cstdint>
10#include <random>
11#include <stdexcept>
12#include <string>
13#include <vector>
14
19
20namespace pulsatrix {
21
67public:
86 SparseAutoencoder(int64_t dim, int64_t hidden_dim, float l1_lambda, DeviceBackend* backend)
87 : SparseAutoencoder(dim, hidden_dim, l1_lambda, backend, static_cast<unsigned>(next_seed())) {}
88
89 SparseAutoencoder(int64_t dim, int64_t hidden_dim, float l1_lambda, DeviceBackend* backend, unsigned seed)
90 : dim_(dim),
91 hidden_dim_(hidden_dim),
92 l1_lambda_(l1_lambda),
93 backend_(backend),
94 // Clamped only so a rejected dimension can't reach LinearModule/Shape and throw
95 // *their* message before the check below throws this class's own, clearer one --
96 // member initialization necessarily runs before the constructor body.
97 encoder_(dim > 0 ? dim : 1, hidden_dim > 0 ? hidden_dim : 1, backend),
98 relu_(backend),
99 decoder_(hidden_dim > 0 ? hidden_dim : 1, dim > 0 ? dim : 1, backend),
100 loss_(backend) {
101 if (dim <= 0) {
102 throw std::invalid_argument("SparseAutoencoder: dim must be positive");
103 }
104 if (hidden_dim <= 0) {
105 throw std::invalid_argument("SparseAutoencoder: hidden_dim must be positive");
106 }
107 if (l1_lambda < 0.0f) {
108 throw std::invalid_argument("SparseAutoencoder: l1_lambda must be non-negative");
109 }
110
111 std::mt19937 rng(seed);
112 // Fan-in scaling: each layer's init range shrinks with its own input width, so the
113 // pre-activation magnitude at step 0 stays O(1) regardless of dim/hidden_dim rather
114 // than growing with the layer's width (which for a deliberately overcomplete SAE
115 // would otherwise saturate the decoder's input from the first step).
116 init_layer(encoder_, rng, 1.0f / std::sqrt(static_cast<float>(dim)));
117 init_layer(decoder_, rng, 1.0f / std::sqrt(static_cast<float>(hidden_dim)));
118 }
119
159 template <typename OptimizerT>
160 float train_step(const Tensor& input_batch, OptimizerT& optimizer) {
161 validate_batch(input_batch, "SparseAutoencoder::train_step");
162
163 optimizer.zero_grad(encoder_);
164 optimizer.zero_grad(decoder_);
165
166 const Tensor hidden = relu_.forward(encoder_.forward(input_batch));
167 const Tensor reconstruction = decoder_.forward(hidden);
168 const float reconstruction_loss = loss_.forward(reconstruction, input_batch);
169
170 const Tensor grad_reconstruction = loss_.backward();
171 Tensor grad_hidden = decoder_.backward(grad_reconstruction);
172
173 const float batch_size = static_cast<float>(input_batch.shape().dim(0));
174 Tensor l1_grad(grad_hidden.shape(), backend_, grad_hidden.device());
175 l1_grad.fill(l1_lambda_ / batch_size);
176 grad_hidden.accumulate(l1_grad);
177
178 const Tensor grad_pre_activation = relu_.backward(grad_hidden);
179 (void)encoder_.backward(grad_pre_activation); // the input has no upstream to receive this
180
181 optimizer.step(encoder_);
182 optimizer.step(decoder_);
183 // relu_ has no parameters (Module::parameters() default) -- stepping it would be a
184 // safe no-op but would imply there is something to update.
185 return reconstruction_loss;
186 }
187
207 [[nodiscard]] Tensor reconstruct(const Tensor& input_batch) const {
208 validate_batch(input_batch, "SparseAutoencoder::reconstruct");
209 return decoder_.forward(encode(input_batch));
210 }
211
233 [[nodiscard]] float reconstruction_error(const Tensor& input_batch) const {
234 validate_batch(input_batch, "SparseAutoencoder::reconstruction_error");
235
236 const Tensor reconstruction = reconstruct(input_batch);
237 const int64_t n = reconstruction.numel();
238 float sum_squared = 0.0f;
239 for (int64_t i = 0; i < n; ++i) {
240 const float diff = reconstruction.data()[i] - input_batch.data()[i];
241 sum_squared += diff * diff;
242 }
243 return sum_squared / static_cast<float>(n);
244 }
245
262 [[nodiscard]] float mean_hidden_activation(const Tensor& input_batch) const {
263 validate_batch(input_batch, "SparseAutoencoder::mean_hidden_activation");
264
265 const Tensor hidden = encode(input_batch);
266 const int64_t n = hidden.numel();
267 float sum = 0.0f;
268 for (int64_t i = 0; i < n; ++i) {
269 sum += hidden.data()[i];
270 }
271 return sum / static_cast<float>(n);
272 }
273
275 [[nodiscard]] int64_t dim() const { return dim_; }
276
278 [[nodiscard]] int64_t hidden_dim() const { return hidden_dim_; }
279
281 [[nodiscard]] float l1_lambda() const { return l1_lambda_; }
282
287 [[nodiscard]] LinearModule& encoder() { return encoder_; }
288
290 [[nodiscard]] const LinearModule& encoder() const { return encoder_; }
291
293 [[nodiscard]] LinearModule& decoder() { return decoder_; }
294
296 [[nodiscard]] const LinearModule& decoder() const { return decoder_; }
297
298private:
301 [[nodiscard]] Tensor encode(const Tensor& input_batch) const { return relu_.forward(encoder_.forward(input_batch)); }
302
304 static void init_layer(LinearModule& layer, std::mt19937& rng, float scale) {
305 std::vector<float> weights(static_cast<size_t>(layer.weight().numel()));
306 for (float& w : weights) {
307 w = uniform_symmetric(rng, scale);
308 }
309 layer.set_weight(weights);
310 layer.set_bias(std::vector<float>(static_cast<size_t>(layer.bias().numel()), 0.0f));
311 }
312
314 static float uniform_symmetric(std::mt19937& rng, float scale) {
315 // std::uniform_real_distribution's mapping is implementation-defined, so identical
316 // seeds would not give identical weights across standard libraries. This mapping is
317 // fixed here, making a seeded SAE reproducible everywhere -- the same convention
318 // LinearProbe already uses.
319 const float unit = static_cast<float>(rng() - std::mt19937::min()) /
320 static_cast<float>(std::mt19937::max() - std::mt19937::min() + 1ull);
321 return (unit * 2.0f - 1.0f) * scale;
322 }
323
339 void validate_batch(const Tensor& input_batch, const char* method) const {
340 const std::string where(method);
341 if (input_batch.rank() != 2) {
342 throw std::invalid_argument(where + ": input_batch must be rank-2 (N, dim)");
343 }
344 if (input_batch.shape().dim(1) != dim_) {
345 throw std::invalid_argument(where + ": input_batch width must equal dim");
346 }
347 if (input_batch.shape().dim(0) <= 0) {
348 throw std::invalid_argument(where + ": batch must not be empty");
349 }
350 }
351
352 int64_t dim_;
353 int64_t hidden_dim_;
354 float l1_lambda_;
355 DeviceBackend* backend_;
356 mutable LinearModule encoder_;
357 mutable ReluModule relu_;
358 mutable LinearModule decoder_;
359 MSELoss loss_;
360};
361
362} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
y = x @ W + b, batched (x is (N, in_features), y is (N, out_features)) – migrated from the original u...
Definition linear_module.hpp:37
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input, and accumulates the weight/bias gradients internall...
float forward(const Tensor &prediction, const Tensor &target)
Computes the loss value and caches prediction/target for backward().
Tensor backward() const
Computes the gradient w.r.t. the prediction: (2/n) * (prediction - target).
Tensor forward(const Tensor &input)
Runs this module's forward computation.
Definition module.hpp:73
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input.
int64_t dim(size_t index) const
Size of a single dimension.
Definition shape.hpp:93
A sparse autoencoder (SAE): LinearModule(dim, hidden_dim) -> ReluModule -> LinearModule(hidden_dim,...
Definition sparse_autoencoder.hpp:66
float reconstruction_error(const Tensor &input_batch) const
Mean squared reconstruction error over the batch – forward-only, no backward, no parameter update.
Definition sparse_autoencoder.hpp:233
const LinearModule & decoder() const
Const overload of decoder().
Definition sparse_autoencoder.hpp:296
float mean_hidden_activation(const Tensor &input_batch) const
Mean hidden (post-ReLU) activation over the batch – the sparsity metric.
Definition sparse_autoencoder.hpp:262
int64_t dim() const
Width of the activation vectors this SAE reconstructs.
Definition sparse_autoencoder.hpp:275
SparseAutoencoder(int64_t dim, int64_t hidden_dim, float l1_lambda, DeviceBackend *backend, unsigned seed)
Definition sparse_autoencoder.hpp:89
const LinearModule & encoder() const
Const overload of encoder().
Definition sparse_autoencoder.hpp:290
float l1_lambda() const
Coefficient of the L1 penalty on the hidden activation.
Definition sparse_autoencoder.hpp:281
Tensor reconstruct(const Tensor &input_batch) const
The SAE's reconstruction of a batch – the encoder->ReLU->decoder forward path, forward-only: no loss,...
Definition sparse_autoencoder.hpp:207
LinearModule & encoder()
The encoder – inspection (its columns are the learned feature directions, which is the whole point of...
Definition sparse_autoencoder.hpp:287
int64_t hidden_dim() const
Width of the sparse hidden basis.
Definition sparse_autoencoder.hpp:278
SparseAutoencoder(int64_t dim, int64_t hidden_dim, float l1_lambda, DeviceBackend *backend)
Constructs a sparse autoencoder over activations of a given dimension.
Definition sparse_autoencoder.hpp:86
LinearModule & decoder()
The decoder – same rationale as encoder().
Definition sparse_autoencoder.hpp:293
float train_step(const Tensor &input_batch, OptimizerT &optimizer)
Runs one training step: forward, MSE reconstruction loss, backward with the L1 penalty gradient injec...
Definition sparse_autoencoder.hpp:160
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
DeviceType device() const
Which device this tensor's buffer conceptually resides on.
Definition tensor.hpp:122
Tensor & fill(float value)
Sets every element to value. Safe no-op on a zero-element tensor.
Tensor & accumulate(const Tensor &other)
In-place elementwise accumulation: this[i] += other[i] for every element.
int64_t numel() const
Total element count – shape().numel().
Definition tensor.hpp:116
const Shape & shape() const
This tensor's shape.
Definition tensor.hpp:113
const float * data() const
Raw buffer access. nullptr iff numel() == 0.
Definition tensor.hpp:167
One global seed for everything that isn't given its own, and a deterministic mode that forbids nondet...
Dense/fully-connected layer – the reference Module implementation.
Mean squared error loss.
Definition acquisition_functions.hpp:16
uint64_t next_seed()
The next seed in the global stream: a distinct, well-mixed 64-bit value per call, reproducible for a ...
ReLU activation – the second Module subclass, following LinearModule's pattern.