Skip to content

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 an ActivationSnapshot, a self-contained copy of one forward pass's activations that stays valid after later passes. forward_pass_with_patch() (activation patching), logit_lens(), and attention_weights() cover the rest of the table above.
  • Probing: LinearProbe trains a linear classifier to test whether a binary concept is linearly decodable from a layer's activations.
  • Decomposition: SparseAutoencoder reconstructs activations through a wider hidden layer with an L1 penalty, so each example uses only a few hidden units. CircuitGraph scores every node by how much zeroing it changes the output.
  • GFlowNet: HyperGridEnv, GFlowNetForwardPolicy, sample_gflownet_trajectory (returns a GFlowNetTrajectory), TrajectoryBalanceLoss, DetailedBalanceLoss, and SubTBLoss. For SubTB(λ), you pass each sub-trajectory pair's λ-weight to forward(). LearnableScalar is the single trainable log Z value Trajectory Balance needs. It is not a Module.

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.

Recipes