pulsatrix
Loading...
Searching...
No Matches
pulsatrix::SparseAutoencoder Class Reference

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>

Public Member Functions

 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().
 

Detailed Description

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.

Constructor & Destructor Documentation

◆ 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
dimWidth of the activation vectors this SAE reconstructs. Must be > 0.
hidden_dimWidth of the sparse hidden basis. Must be > 0; normally > dim.
l1_lambdaCoefficient 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.
backendBackend to allocate/compute through. Not owned; must outlive this SAE.
seedRNG seed for weight initialization – same seed gives the same SAE.
Exceptions
std::invalid_argumentif 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

Member Function Documentation

◆ decoder() [1/2]

LinearModule & pulsatrix::SparseAutoencoder::decoder ( )
inline

The decoder – same rationale as encoder().

◆ decoder() [2/2]

const LinearModule & pulsatrix::SparseAutoencoder::decoder ( ) const
inline

Const overload of decoder().

◆ dim()

int64_t pulsatrix::SparseAutoencoder::dim ( ) const
inline

Width of the activation vectors this SAE reconstructs.

◆ encoder() [1/2]

LinearModule & pulsatrix::SparseAutoencoder::encoder ( )
inline

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

Const overload of encoder().

◆ 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_batchShape (N, dim), N > 0.
Returns
mean over all N * hidden_dim hidden elements. Always >= 0.
Exceptions
std::invalid_argumenton 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_batchShape (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_argumenton 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_batchShape (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_argumenton 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
OptimizerTAny type exposing step(Module&) and zero_grad(Module&) – SGDOptimizer or AdamOptimizer (see the class-level note on why this is a template rather than an Optimizer&).
Parameters
input_batchShape (N, dim), N > 0. Both the input and the reconstruction target – an autoencoder's target is its input.
optimizerOptimizer 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_argumenton 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: