Inference and learning in sparse autoencoders as natural gradient flow
Organizations: Astera Institute · UC Berkeley · UCL · Goodfire AI · CSHL
Abstract
Sparse autoencoders are widely used to uncover interpretable features in neural networks, yet reliable recovery remains difficult when features overlap or activate infrequently. These challenges involve both inferring which features explain an input and learning the dictionary that represents them. Here, we unify inference and dictionary learning as natural-gradient flows on a shared variational free energy. We instantiate this framework as BeFOND, an encoder-free sparse coding model with closed-form inference and learning dynamics. We show how recurrent explaining away reduces interference between overlapping features, while Fisher preconditioning can compensate for the slow learning of rare features. On synthetic data, BeFOND improves dictionary recovery and rare-feature detection, with a growing advantage over amortized baselines as superposition increases. On language-model activations, it improves single-feature concept detection and selective intervention, outperforming pretrained reference SAEs with substantially less training data. Its feature quality continues to improve with dictionary width, whereas the evaluated baselines largely plateau. Together, these results show how improving inference and learning within a unified probabilistic framework can make better use of data and dictionary capacity to interpret and intervene on neural representations.
Figures & tables
Appendix figures & tables18 assets
Supplementary material from the paper’s appendix.
Appendix
| Algorithm | Metric | Loss |
| Gradient Descent ( Cauchy, 1847 , GD,) | (Euclidean) | Generic task loss |
| Natural Gradient Descent ( Amari, 1998 , NGD,) | Fisher–Rao | Generic task loss |
| Bayesian Learning Rule ( Khan and Rue, 2023 , BLR,) | Fisher–Rao | (free energy) |
| Experiment | Settings |
| FOND best / continuous | Seed 0; batch 2,048; 88,779 updates; 181.82M decoder and 18.18M prior examples; linear LR , no warmup; |
| Inference / decoder | Bernoulli; ETD1 ; one inner step; uniform training depth 1–10, test depth 200. Prior Fisher, current moments, no damping; ; bias and variance fixed |
| Initialization / constraints | Gaussian scale 0.1443375673; prior logits ; soft decoder norm 3.5 with coefficient 100; prior mean cap and floor |
| Prior / matrix functions | Prior update every 10 batches on 2,048 separate examples; depth , KL threshold 0.01, rejection threshold 1, window 10. Taylor limit 32 applications; Chebyshev tolerance , maximum degree 2,048 |
| Original baselines | BatchTopK/JumpReLU/Matryoshka: target , batch 1,024, LR , 200M decoder examples. MP/ReLU-L1: target , 200,000,512 decoder plus 19M monitoring examples. MF: 20 inference steps in training and testing, batch 1,024, 195,313 updates (200,000,512 examples), cosine schedule with dictionary / other-parameter LR , damping 1 on the first step then 0.2, decoder norm cap 4 |
| Modified baselines | Target for BatchTopK/JumpReLU/Matryoshka; MinFire or resampling and learning-rate decay |
| Neutral roof | Wet roof | |||
| Method | Rain | Sprinkler | Rain | Sprinkler |
| Exact posterior | 0.292138 | 0.292138 | 0.960707 | 0.033017 |
| FOND | 0.228996 | 0.228996 | 0.974111 | 0.019506 |
| Single-projection encoder | 0.217514 | 0.139225 | 0.958493 | 0.064003 |
| MCC | Dead latents | ||||||
| Method | Seed 0 | Seed 1 | Seed 2 | 0 | 1 | 2 | |
| BatchTopK | 1536 | 0.6447949409484863 | 0.6418958902359009 | 0.6367465853691101 | 178 | 217 | 206 |
| 1024 | 0.7172077894210815 | 0.7101552486419678 | 0.715070366859436 | 78 | 80 | 66 | |
| 768 | 0.7237797379493713 | 0.7247610688209534 | 0.7268125414848328 | 49 | 47 | 39 | |
| 512 | 0.6935576796531677 | 0.6925675868988037 | 0.6897946000099182 | 37 | 38 | 55 | |
| 384 | 0.6238665580749512 | 0.6256771087646484 | 0.6227874159812927 | 56 | 82 | 58 | |
| Method | Seed | MCC | Macro F1 | Rare F1 | Dead | |
| BatchTopK | 0 | 0.534523 | 0.445173 | 0.299436 | 34.833798 | 1,577 |
| 1 | 0.535311 | 0.444452 | 0.299666 | 35.031285 | 1,733 | |
| 2 | 0.537442 | 0.446642 | 0.296203 | 34.948331 | 1,707 | |
| Mean | 34.94 | 1672.33 | ||||
| JumpReLU | 0 | 0.483953 | 0.326567 | 0.190439 | 45.148172 | 8,069 |
| 1 | 0.481728 | 0.326672 | 0.184864 | 44.826161 | 8,090 |
| Sparse probing | RAVEL | ||||||||
| Model | |||||||||
| Acc. | AUC | F1 | Acc. | AUC | F1 | Disentangle- ment | Cause | Isolation | |
| Dictionary width: 16k | |||||||||
| BatchTopK (target ) | {.767}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {\mathbf{.798}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.761}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.004}} | {\mathbf{.814}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {\mathbf{.862}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {\mathbf{.812}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.729}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.005}} | {.692}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.019}} | {.767}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.017}} |
| Matryoshka BatchTopK ( probing; RAVEL) | {.768}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.794}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.762}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.810}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.858}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.808}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.731}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.695}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.013}} | {.767}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.015}} |
| Mean-field + Adam ( inference steps) | {.761}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.761}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.749}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.004}} | {.801}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.842}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.001}} | {.798}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.696}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.639}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.006}} | {.753}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.006}} |
| Method / metric | Settings |
| FOND | Widths 16,384–524,288; exponentially sampled training depth (mean 16, cap 100; Appendix C.1 ); FP32 inference/parameters; prior Fisher with current moments and no damping; ; decoder solver matrix_free_phi1_auto |
| Mean-field + Adam | Whitened inputs, 50M tokens, 20 inference steps, readout . After 2,000 updates, clip the dictionary gradient at three times its running-mean norm and skip batches whose negative ELBO exceeds five times its running mean, capped at approximately 20% of recent batches |
| BatchTopK / Matryoshka | Whitened inputs; batch 4,096; 12,207 updates (49,999,872 tokens); constant LR ; FP32 SAEs, BF16 Gemma; SAELens 6.51.1. ZCA fitted on 500M training tokens. Target values in Table 8 were selected on a separate pilot |
| Baseline initialization / losses | Decoder norm initialized to 0.1; decoder bias subtracted before encoding; activations rescaled by decoder norm; auxiliary-loss coefficient 1; threshold EMA rate 0.01 |
| Matryoshka levels | 4,096, 16,384, and the full width (16,384 appears once for the smallest model); Matryoshka auxiliary loss enabled |
| Gemma Scope | Release google/gemma-scope-2b-pt-res , checkpoint paths layer_12/width_{w}/average_l0_{s} , with given by each row of Table 9 . Training: 4B tokens at 16K, 8B at 32K–524K, and 16B at 1M ( Lieberum et al., 2024 ) |
| Sparse probing ( Kantamneni et al., 2025 ) | RAVEL ( Huang et al., 2024 ) | |||||||||
| Width | Target | |||||||||
| Acc. | AUC | F1 | Acc. | AUC | F1 | Overall | Cause | Isolation | ||
| BatchTopK (16k) | ||||||||||
| 16k | 200 | {.765}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.004}} | {.791}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.004}} | {.758}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.005}} | {.810}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.859}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {.808}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {.717}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.009}} | {\mathbf{.698}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.018}} | {.736}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.021}} |
| 16k | 400 | {\mathbf{.767}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {\mathbf{.798}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.003}} | {\mathbf{.761}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.004}} | {\mathbf{.814}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {\mathbf{.862}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {\mathbf{.812}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.002}} | {\mathbf{.729}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.005}} | {.692}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.019}} | {\mathbf{.767}}\scriptscriptstyle{\color[rgb]{0.3984,0.3984,0.3984}\pm{.017}} |
| Matryoshka BatchTopK (16k) | ||||||||||
| Sparse probing | RAVEL | ||||||||||||
| Width | Disent- angle | Cause | Isola- tion | ||||||||||
| Acc. | AUC | F1 | Acc. | AUC | F1 | Acc. | AUC | F1 | |||||
| 16k | 22 | 0.7337 | 0.7476 | 0.7159 | 0.7529 | 0.7861 | 0.7441 | 0.7771 | 0.8243 | 0.7717 | 0.5989 | 0.5121 | 0.6857 |
| 16k | 41 | 0.7519 | 0.7710 | 0.7366 | 0.7686 | 0.8056 | 0.7601 | 0.7893 | 0.8370 | 0.7835 | 0.6644 | 0.6163 | 0.7126 |
| 16k | 82 | 0.7627 | 0.7857 | 0.7545 | 0.7821 | 0.8207 | 0.7768 | 0.8047 | 0.8523 | 0.8016 | 0.6960 | 0.6489 | 0.7431 |
| 16k | 176 | 0.7679 | 0.7938 | 0.7602 | 0.7890 | 0.8262 | 0.7851 | 0.8129 | 0.8587 | 0.8119 | 0.7206 | 0.6508 | 0.7904 |
| Width | FOND | BatchTopK | Matryoshka | Mean-field + Adam |
| 16,384 | 33.125 | 2.719 | 2.991 | 2.714 |
| 65,536 | 118.959 | 9.202 | 10.737 | 10.157 |
| 131,072 | 234.665 | 17.935 | 20.075 | 19.914 |
| 524,288 | 989.321 | 70.121 | 77.382 | 78.331 |