A sparse autoencoder (SAE): LinearModule(dim, hidden_dim) -> ReluModule -> LinearModule(hidden_dim, dim), trained with MSELoss to reconstruct its own input while an L1 penalty on the hidden ReLU activation pushes most hidden units to zero on any given example.
More...
#include <sparse_autoencoder.hpp>
|
| | SparseAutoencoder (int64_t dim, int64_t hidden_dim, float l1_lambda, DeviceBackend *backend) |
| | Constructs a sparse autoencoder over activations of a given dimension.
|
| |
| | SparseAutoencoder (int64_t dim, int64_t hidden_dim, float l1_lambda, DeviceBackend *backend, unsigned seed) |
| |
| template<typename OptimizerT > |
| float | train_step (const Tensor &input_batch, OptimizerT &optimizer) |
| | Runs one training step: forward, MSE reconstruction loss, backward with the L1 penalty gradient injected at the hidden layer, one optimizer update of the encoder's and decoder's parameters.
|
| |
| Tensor | reconstruct (const Tensor &input_batch) const |
| | The SAE's reconstruction of a batch – the encoder->ReLU->decoder forward path, forward-only: no loss, no backward, no parameter update.
|
| |
| float | reconstruction_error (const Tensor &input_batch) const |
| | Mean squared reconstruction error over the batch – forward-only, no backward, no parameter update.
|
| |
| float | mean_hidden_activation (const Tensor &input_batch) const |
| | Mean hidden (post-ReLU) activation over the batch – the sparsity metric.
|
| |
| int64_t | dim () const |
| | Width of the activation vectors this SAE reconstructs.
|
| |
| int64_t | hidden_dim () const |
| | Width of the sparse hidden basis.
|
| |
| float | l1_lambda () const |
| | Coefficient of the L1 penalty on the hidden activation.
|
| |
| LinearModule & | encoder () |
| | The encoder – inspection (its columns are the learned feature directions, which is the whole point of an SAE) and test-time weight injection.
|
| |
| const LinearModule & | encoder () const |
| | Const overload of encoder().
|
| |
| LinearModule & | decoder () |
| | The decoder – same rationale as encoder().
|
| |
| const LinearModule & | decoder () const |
| | Const overload of decoder().
|
| |
A sparse autoencoder (SAE): LinearModule(dim, hidden_dim) -> ReluModule -> LinearModule(hidden_dim, dim), trained with MSELoss to reconstruct its own input while an L1 penalty on the hidden ReLU activation pushes most hidden units to zero on any given example.
hidden_dim is normally chosen larger than dim – an overcomplete basis, more directions than the activation space being decomposed, which is what distinguishes an SAE from a compressing autoencoder. The activations come from an ActivationSnapshot (this campaign's Phase 1 output) in production use, but nothing here depends on that: an SAE consumes a plain (N, dim) batch, so it is equally usable against hand-built synthetic data – which is exactly what the paired penalty/no-penalty control in tests/sparse_autoencoder_test.cpp needs.
- Note
- Bounded claim (campaign Decision Point 3): this class measures and reports reconstruction fidelity (reconstruction_error) and sparsity (mean_hidden_activation) as numbers. It does not establish that the hidden directions it learns are semantically meaningful or causally compositional features – that is a live debate in the field, and Phase 4 (activation patching) evidence is the minimum prerequisite for even arguing it. Do not let a low reconstruction error be read as a claim about feature semantics.
-
Mirrors LinearProbe/XorNetwork's established shape rather than inventing a training abstraction: the constructor takes a backend,
train_step() drives one loss/backward/optimizer-step cycle, and the epoch loop is the caller's. The SAE deliberately does not own its optimizer – same division of responsibility LinearProbe::train_step already uses, which is what lets a caller choose SGD vs. Adam and keep one optimizer's state across a whole training run.
-
train_step is a template on the optimizer type rather than taking an Optimizer&, for exactly the reason LinearProbe documents: this codebase has no Optimizer base class – SGDOptimizer and AdamOptimizer are unrelated types sharing only the step(Module&) / zero_grad(Module&) shape (a compile-time, duck-typed contract). Introducing an abstract base to give this one method a runtime-polymorphic parameter would be a change to two closed missions' classes for no behavioral gain.
-
Small seeded random weight init, not zero init – and here, unlike LinearProbe, zero init is not merely degenerate but dead: with an all-zero decoder weight the gradient reaching the hidden layer is identically zero, so the encoder never receives any gradient and the autoencoder cannot leave the origin.
-
hidden_dim > dim is not enforced. An undercomplete or square autoencoder is still a well-formed object with well-defined training behavior; rejecting it would turn a modeling choice into a hard error for no correctness gain. Explicit decision (mission exit gate asked for it to be made, not assumed), pinned by ConstructorAcceptsAnUndercompleteHiddenDimension.
◆ SparseAutoencoder() [1/2]
| pulsatrix::SparseAutoencoder::SparseAutoencoder |
( |
int64_t |
dim, |
|
|
int64_t |
hidden_dim, |
|
|
float |
l1_lambda, |
|
|
DeviceBackend * |
backend |
|
) |
| |
|
inline |
Constructs a sparse autoencoder over activations of a given dimension.
- Parameters
-
| dim | Width of the activation vectors this SAE reconstructs. Must be > 0. |
| hidden_dim | Width of the sparse hidden basis. Must be > 0; normally > dim. |
| l1_lambda | Coefficient of the L1 penalty on the hidden activation. Must be >= 0; 0 is a valid degenerate configuration (a plain autoencoder), which is precisely what the mission's no-penalty control trains. |
| backend | Backend to allocate/compute through. Not owned; must outlive this SAE. |
| seed | RNG seed for weight initialization – same seed gives the same SAE. |
- Exceptions
-
| std::invalid_argument | if dim <= 0, hidden_dim <= 0, or l1_lambda < 0. External boundary: all three are caller-supplied (the dimensions typically read off a snapshot tensor's shape) and nothing downstream rejects them – LinearModule(0, h, ...) builds a well-formed zero-element weight and fails only later, confusingly, and a negative lambda never fails at all: it rewards hidden activation without bound, so the run diverges silently rather than erroring. |
Seeded from the global seed stream (next_seed(), FND-7).
◆ SparseAutoencoder() [2/2]
| pulsatrix::SparseAutoencoder::SparseAutoencoder |
( |
int64_t |
dim, |
|
|
int64_t |
hidden_dim, |
|
|
float |
l1_lambda, |
|
|
DeviceBackend * |
backend, |
|
|
unsigned |
seed |
|
) |
| |
|
inline |
◆ decoder() [1/2]
◆ decoder() [2/2]
| const LinearModule & pulsatrix::SparseAutoencoder::decoder |
( |
| ) |
const |
|
inline |
◆ dim()
| int64_t pulsatrix::SparseAutoencoder::dim |
( |
| ) |
const |
|
inline |
Width of the activation vectors this SAE reconstructs.
◆ encoder() [1/2]
The encoder – inspection (its columns are the learned feature directions, which is the whole point of an SAE) and test-time weight injection.
◆ encoder() [2/2]
| const LinearModule & pulsatrix::SparseAutoencoder::encoder |
( |
| ) |
const |
|
inline |
◆ hidden_dim()
| int64_t pulsatrix::SparseAutoencoder::hidden_dim |
( |
| ) |
const |
|
inline |
Width of the sparse hidden basis.
◆ l1_lambda()
| float pulsatrix::SparseAutoencoder::l1_lambda |
( |
| ) |
const |
|
inline |
Coefficient of the L1 penalty on the hidden activation.
◆ mean_hidden_activation()
| float pulsatrix::SparseAutoencoder::mean_hidden_activation |
( |
const Tensor & |
input_batch | ) |
const |
|
inline |
Mean hidden (post-ReLU) activation over the batch – the sparsity metric.
- Parameters
-
| input_batch | Shape (N, dim), N > 0. |
- Returns
- mean over all N * hidden_dim hidden elements. Always >= 0.
- Exceptions
-
| std::invalid_argument | on any malformed batch – see validate_batch(). |
- Note
- Post-ReLU, so every element is already >= 0 and the mean is a direct, valid proxy for the L1 norm this class penalizes – no abs() needed. Lower means sparser: the floor, 0, is every hidden unit clamped off for every example.
-
Reports mean magnitude, not a count of active units. A fraction-nonzero metric (L0) would be the other natural choice and is not offered here: it is not what the penalty term optimizes, and it would be discontinuous in the parameters, making it a poor thing to compare two training runs by. This metric is the one the penalty actually targets, which is the property the mission's paired control tests.
-
const for the same reason reconstruction_error() is – see its note.
◆ reconstruct()
| Tensor pulsatrix::SparseAutoencoder::reconstruct |
( |
const Tensor & |
input_batch | ) |
const |
|
inline |
The SAE's reconstruction of a batch – the encoder->ReLU->decoder forward path, forward-only: no loss, no backward, no parameter update.
- Parameters
-
| input_batch | Shape (N, dim), N > 0. |
- Returns
- Shape (N, dim) – x_hat, the same tensor reconstruction_error() scores against the input. Exposed as a value rather than only as a scalar error because the reconstruction itself is what an activation-patching experiment substitutes (campaign_exai_dl_library_mechanistic_interpretability, Phase 4's exit gate); before this method the only way to obtain one was train_step(), which also mutates the parameters – an observation that changes what it observes.
- Exceptions
-
| std::invalid_argument | on any malformed batch – see validate_batch(). |
- Note
const for the same reason reconstruction_error() is – see its note. Third instance of the mutable-members-for-a-const-scoring-path pattern.
-
No new computation: reconstruction_error() is defined in terms of this method, so the two can never report a reconstruction and an error computed from different forward paths. Pinned by ReconstructOutputAgreesWithReconstructionErrorsInternalComputation, which recomputes the MSE by hand from this method's output and compares.
◆ reconstruction_error()
| float pulsatrix::SparseAutoencoder::reconstruction_error |
( |
const Tensor & |
input_batch | ) |
const |
|
inline |
Mean squared reconstruction error over the batch – forward-only, no backward, no parameter update.
- Parameters
-
| input_batch | Shape (N, dim), N > 0. |
- Returns
- mean over all N * dim elements of (reconstruction - input)^2, i.e. exactly MSELoss's own convention, so this number is directly comparable to the loss train_step() returns.
- Exceptions
-
| std::invalid_argument | on any malformed batch – see validate_batch(). |
- Note
const although Module::forward() is not: the SAE's logical state is its learned parameters, which a forward-only scoring pass cannot change. The modules are mutable purely so forward()'s internal caches (backward()'s working state, not the SAE's observable value) can be written. Flagged as a recurring pattern by LinearProbe's AAR; this is its second instance.
-
Computed in a raw host loop rather than through the member
loss_, and deliberately so: MSELoss::forward caches its operands to arm a subsequent backward(), and a scoring pass has no business arming a backward it never performs – doing so would leave the loss primed with a batch the next train_step() did not produce. The loop itself is the established convention (MSELoss/CrossEntropyLoss/LinearProbe::accuracy all read Tensor::data() directly, since Tensor has no reduction primitive).
◆ train_step()
template<typename OptimizerT >
| float pulsatrix::SparseAutoencoder::train_step |
( |
const Tensor & |
input_batch, |
|
|
OptimizerT & |
optimizer |
|
) |
| |
|
inline |
Runs one training step: forward, MSE reconstruction loss, backward with the L1 penalty gradient injected at the hidden layer, one optimizer update of the encoder's and decoder's parameters.
- Template Parameters
-
- Parameters
-
| input_batch | Shape (N, dim), N > 0. Both the input and the reconstruction target – an autoencoder's target is its input. |
| optimizer | Optimizer to update this SAE's encoder/decoder parameters with. Not owned; its state persists across calls, which is the point of not owning it. |
- Returns
- The reconstruction loss for this batch, measured before the update (mirroring LinearProbe/XorNetwork's train_step return contract). The L1 penalty term is deliberately not folded into this number: the two quantities trade off against each other, and reporting their sum would hide which of them a change in the total came from. The sparsity side is observed via mean_hidden_activation().
- Exceptions
-
| std::invalid_argument | on any malformed batch – see validate_batch(). |
- Note
- Zeroes gradients before accumulating, for the exact reason XorNetwork:: train_step documents: this codebase's modules accumulate parameter gradients across backward() calls until something resets them, so without this every step would silently sum onto every previous step's gradient.
-
The L1 penalty is
(l1_lambda / N) * sum_over(batch, hidden) h_ij – a per-example sum of hidden activations, averaged over the batch. abs() is unnecessary: h is the output of a ReLU and therefore already >= 0. Its gradient w.r.t. every hidden element is the constant l1_lambda / N, which is why a uniform fill()+accumulate() onto the incoming hidden gradient is exactly correct and not an approximation. Adding it before relu_.backward() is also what makes it correct for the clamped-off units: ReLU's own backward zeroes the contribution wherever the pre-activation was <= 0, which is precisely where the penalty has no gradient to give (an already-inactive unit is not pushed further down).
-
The penalty gradient is built and accumulated unconditionally, with no
l1_lambda > 0 fast path. At lambda = 0 it is a tensor of zeros and accumulating it is a no-op, so the no-penalty configuration exercises exactly the same instructions the penalized one does – which is what makes the mission's paired control a controlled comparison rather than a comparison of two code paths.
The documentation for this class was generated from the following file: