Mechanistic Interpretability¶
Use this section when you want to see how a model computes its answer, not just which inputs mattered. Interpretability explains input-output behavior. The tools here open the network up: they cache its internal activations, test what those activations encode, and measure which layers the output depends on.
The section also hosts GFlowNets, a training method that learns to sample outcomes in proportion to their reward instead of maximizing it. That makes them a tool for exploring a model's or environment's structure.
Which tool should I use?¶
| Question | Tool |
|---|---|
| What did each layer output for this input? | ExplainerContext::activation_snapshot() → ActivationSnapshot |
| Is concept X linearly readable from layer L? | LinearProbe |
| Can a layer's activations be split into sparser, more interpretable directions? | SparseAutoencoder |
| Which layers does the output depend on for this input? | ExplainerContext::build_circuit_graph() → CircuitGraph |
| What happens to the output if I overwrite one activation? | ExplainerContext::forward_pass_with_patch() |
| What would the model predict if layer L were the last layer? | ExplainerContext::logit_lens() |
| Where does an attention layer look? | ExplainerContext::attention_weights() |
| How do I sample outcomes in proportion to a reward? | GFlowNet types below |
What's inside¶
- Activation access (all on
ExplainerContext):activation_snapshot()returns anActivationSnapshot, a self-contained copy of one forward pass's activations that stays valid after later passes.forward_pass_with_patch()(activation patching),logit_lens(), andattention_weights()cover the rest of the table above. - Probing:
LinearProbetrains a linear classifier to test whether a binary concept is linearly decodable from a layer's activations. - Decomposition:
SparseAutoencoderreconstructs activations through a wider hidden layer with an L1 penalty, so each example uses only a few hidden units.CircuitGraphscores every node by how much zeroing it changes the output. - GFlowNet:
HyperGridEnv,GFlowNetForwardPolicy,sample_gflownet_trajectory(returns aGFlowNetTrajectory),TrajectoryBalanceLoss,DetailedBalanceLoss, andSubTBLoss. For SubTB(λ), you pass each sub-trajectory pair's λ-weight toforward().LearnableScalaris the single trainablelog Zvalue Trajectory Balance needs. It is not aModule.
Full API reference: Doxygen: Mechanistic Interpretability
How to implement¶
Probing for a linearly decodable concept¶
#include "pulsatrix/adam_optimizer.hpp"
#include "pulsatrix/cpu_backend.hpp"
#include "pulsatrix/explainer_context.hpp"
#include "pulsatrix/linear_module.hpp"
#include "pulsatrix/linear_probe.hpp"
#include "pulsatrix/relu_module.hpp"
using namespace pulsatrix;
CPUBackend backend;
// hidden -> relu -> head: your trained network
LinearModule hidden(8, 64, &backend);
ReluModule relu(&backend);
LinearModule head(64, 2, &backend);
ExplainerContext ctx({&hidden, &relu, &head});
// 1. Cache activations for a batch of N inputs. Node 0 is the input; node i+1 is module i's output.
ctx.forward_pass(inputs); // inputs: shape (N, 8)
ActivationSnapshot snapshot = ctx.activation_snapshot();
const Tensor& activation_batch = snapshot.activation(/*relu output=*/2); // shape (N, 64)
// 2. Train a probe on those activations. label_batch: shape (N, 1), values in {0, 1}.
LinearProbe probe(/*activation_dim=*/64, &backend);
AdamOptimizer optimizer(0.01f, &backend);
for (int step = 0; step < 200; ++step) {
probe.train_step(activation_batch, label_batch, optimizer);
}
float acc = probe.accuracy(activation_batch, label_batch);
What's happening: a LinearProbe is a LinearModule(activation_dim, 1) trained with
BCEWithLogitsLoss on (activation, concept label) pairs. High accuracy means the concept is
linearly decodable from that layer. Chance-level accuracy means it is not, at least not
linearly. The probe accepts any (N, activation_dim) batch, so you can also test it on
synthetic data first as a sanity check.
Recipe: Sparse autoencoder + linear probe.
Building a circuit graph¶
#include "pulsatrix/cpu_backend.hpp"
#include "pulsatrix/explainer_context.hpp"
#include "pulsatrix/linear_module.hpp"
#include "pulsatrix/relu_module.hpp"
using namespace pulsatrix;
CPUBackend backend;
// hidden -> relu -> head: your trained network
LinearModule hidden(4, 16, &backend);
ReluModule relu(&backend);
LinearModule head(16, 3, &backend);
ExplainerContext ctx({&hidden, &relu, &head});
Tensor input(Shape({1, 4}), &backend, {0.2f, -1.0f, 0.5f, 0.9f});
CircuitGraph circuit = ctx.build_circuit_graph(input);
for (const CircuitNode& node : circuit.nodes()) {
// node.ablation_effect: how far the output moves when this node is zeroed
}
What's happening: for each node, build_circuit_graph() reruns the forward pass with
that node's activation replaced by zeros. It records the L2 distance between the patched
output and the normal output as ablation_effect. A larger value means the output depends
more on that node for this input. The output node scores 0 by convention. Edges connect each
node to the next and carry the source node's score. CircuitGraph holds data only; to draw it,
use CircuitGraphView from Visualization.
Sampling a GFlowNet trajectory¶
#include "pulsatrix/cpu_backend.hpp"
#include "pulsatrix/gflownet_forward_policy.hpp"
#include "pulsatrix/gflownet_trajectory.hpp"
#include "pulsatrix/hypergrid_env.hpp"
#include "pulsatrix/linear_module.hpp"
using namespace pulsatrix;
CPUBackend backend;
HyperGridEnv env(&backend, /*ndim=*/2, /*side_length=*/5);
LinearModule policy_net(2, 3, &backend); // 2-dim state -> 3 actions (+x, +y, stop)
GFlowNetForwardPolicy policy(&policy_net, /*action_dim=*/3, &backend);
GFlowNetTrajectory traj = sample_gflownet_trajectory(env, policy);
// traj.states / traj.actions: every decision point and the action taken there
// traj.sum_log_pf, traj.sum_log_pb: the log-probability sums Trajectory Balance needs
// traj.terminal_reward: R(x) at the final state
What's happening: sample_gflownet_trajectory resets env. It then samples actions from
the policy, skipping invalid ones, and steps env until the episode ends. An episode ends on
an explicit stop action or at the environment's step cap. Along the way it sums the forward
and backward log-probabilities (Σ log P_F, Σ log P_B).
Feed the trajectory to TrajectoryBalanceLoss, DetailedBalanceLoss, or SubTBLoss. They
train the policy to sample each outcome x with probability proportional to R(x), rather
than always picking the best one.
Recipe: GFlowNet on HyperGrid.