One activation function, one parameter per feature

Senthooran Rajamanoharan, Tom Lieberum, Nicolas Sonnerat, Arthur Conmy, Vikrant Varma, János Kramár, and Neel Nanda posted a sparse autoencoder paper on July 19 whose whole architectural change fits in one line. Replace the ReLU with JumpReLU, defined as z times the Heaviside step of z minus theta, where theta is a learned threshold for each feature. Below the threshold the feature is zero. Above it the feature keeps its full magnitude rather than being pulled down.

The motivation is the same shrinkage problem that has bothered everyone training these things. A ReLU autoencoder with an L1 penalty has to choose between sparsity and faithful magnitudes, because the penalty shrinks every active feature. A threshold separates the two decisions. Whether a feature fires is decided by theta. How strongly it fires is decided by the pre-activation, untouched.

Training through a discontinuity

The loss is L2 reconstruction error plus lambda times the L0 norm of the feature vector. That is the penalty people actually want, the count of active features, rather than the L1 proxy. The problem is that both the step function and the L0 count are piecewise constant in theta, so the gradient with respect to theta is zero almost everywhere and the thresholds would never move.

The paper's fix is to notice that although the loss on one example has no useful gradient, the expected loss over the data distribution does, and that gradient involves the probability density of pre-activations near the threshold. So they define pseudo-derivatives. For the step function the pseudo-derivative with respect to theta is minus one over epsilon times a kernel evaluated at z minus theta over epsilon, and for JumpReLU it is the same scaled by theta. The kernel can be rectangular, triangular, or Gaussian and epsilon is a bandwidth. Averaged over a batch this is a kernel density estimate of the true gradient of the expected loss. That is the straight-through estimator, used here with an argument for why it is the right one rather than as a hack.

The comparison

The experiments run on Gemma 2 9B at layers 9, 20, and 31, on three sites each: the residual stream, the MLP output, and the attention output before the output projection. At matched sparsity, JumpReLU SAEs consistently reach lower delta language model loss than Gated SAEs, and match or slightly beat TopK SAEs. Delta LM loss is the increase in the model's loss when you splice the reconstruction back in, which is the number that matters if you plan to use the features to explain behaviour.

The efficiency argument is that JumpReLU is elementwise, so training costs about what a plain ReLU autoencoder costs, one forward and one backward pass. TopK needs a partial sort over the dictionary on every token. Gated SAEs need extra parameters and a second path through the encoder. For a dictionary with hundreds of thousands of features, the sort is a real cost, and avoiding it while matching the result is the practical case for the method.

Did the threshold hurt interpretability

A fair worry is that a hard threshold might make features fire on cleaner but less meaningful sets. The paper checks this two ways. Five human raters scored 405 feature samples across three sparsity levels for monosemanticity, and the ratings came out similar across JumpReLU, Gated, and TopK. An automated pipeline used Gemini Flash to write an explanation for each feature and then simulate its activations, and the correlation between simulated and real activations was comparable across architectures, with JumpReLU marginally ahead of Gated.

One thing the threshold does change is the tail of very frequent features. JumpReLU, like TopK, produces more features that fire on over 10 percent of tokens than Gated does. In a 131k-width dictionary this is under 0.06 percent of features, and the interpretability cost shows up only for low-sparsity models, but it is there and the authors say so.

Why it became a default

The paper's limitations are honest and small. Evaluation is on one model family, with preliminary Pythia runs pointing the same way. The method adds two hyperparameters, the theta initialisation and the epsilon bandwidth, though the defaults transfer across models once activations are normalised. And the paper notes that a Gated SAE with weight sharing is mathematically the same function as a JumpReLU SAE, so the contribution is the loss and the training procedure rather than the function class.

That last point is why we expect this to be the architecture people train at scale. It needs no sort, no extra parameters, trains the penalty you actually care about, and holds the interpretability line. When DeepMind releases its open dictionaries for Gemma 2, this is the recipe we would expect them to use. The experiment we would run is the STE with other discontinuous objectives, since the paper says the trick generalises beyond L0, and nobody has yet tried a penalty on downstream effect sparsity trained the same way.

Sources

  1. Rajamanoharan et al., Jumping Ahead: Improving Reconstruction Fidelity with JumpReLU Sparse Autoencoders (arXiv 2407.14435)