102 throw std::invalid_argument(
"SparseAutoencoder: dim must be positive");
105 throw std::invalid_argument(
"SparseAutoencoder: hidden_dim must be positive");
108 throw std::invalid_argument(
"SparseAutoencoder: l1_lambda must be non-negative");
111 std::mt19937 rng(seed);
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)));
159 template <
typename OptimizerT>
161 validate_batch(input_batch,
"SparseAutoencoder::train_step");
163 optimizer.zero_grad(encoder_);
164 optimizer.zero_grad(decoder_);
168 const float reconstruction_loss = loss_.
forward(reconstruction, input_batch);
173 const float batch_size =
static_cast<float>(input_batch.
shape().
dim(0));
175 l1_grad.
fill(l1_lambda_ / batch_size);
179 (void)encoder_.
backward(grad_pre_activation);
181 optimizer.step(encoder_);
182 optimizer.step(decoder_);
185 return reconstruction_loss;
208 validate_batch(input_batch,
"SparseAutoencoder::reconstruct");
209 return decoder_.
forward(encode(input_batch));
234 validate_batch(input_batch,
"SparseAutoencoder::reconstruction_error");
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;
243 return sum_squared /
static_cast<float>(n);
263 validate_batch(input_batch,
"SparseAutoencoder::mean_hidden_activation");
265 const Tensor hidden = encode(input_batch);
266 const int64_t n = hidden.
numel();
268 for (int64_t i = 0; i < n; ++i) {
269 sum += hidden.
data()[i];
271 return sum /
static_cast<float>(n);
275 [[nodiscard]] int64_t
dim()
const {
return dim_; }
278 [[nodiscard]] int64_t
hidden_dim()
const {
return hidden_dim_; }
281 [[nodiscard]]
float l1_lambda()
const {
return l1_lambda_; }
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);
309 layer.set_weight(weights);
310 layer.set_bias(std::vector<float>(
static_cast<size_t>(layer.bias().numel()), 0.0f));
314 static float uniform_symmetric(std::mt19937& rng,
float scale) {
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;
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)");
344 if (input_batch.shape().dim(1) != dim_) {
345 throw std::invalid_argument(where +
": input_batch width must equal dim");
347 if (input_batch.shape().dim(0) <= 0) {
348 throw std::invalid_argument(where +
": batch must not be empty");
355 DeviceBackend* backend_;
356 mutable LinearModule encoder_;
357 mutable ReluModule relu_;
358 mutable LinearModule decoder_;
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.
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.