LinearPFN: Amortized Variable Selection for Linear Models with Interactions
Organizations: Institute of Psychology, Humboldt-Universität zu Berlin, Berlin, Germany · Department of Education and Psychology, Freie Universität Berlin, Berlin, Germany · Department for AI in Society, Science, and Technology, Zuse Institute Berlin, Germany · Institute of Mathematics, Technische Universität Berlin, Germany
Abstract
Spike-and-slab regression is a standard Bayesian formulation of variable selection: it returns a posterior distribution over which candidate effects are active rather than a single selected subset, so that every candidate effect carries an inclusion probability. Its cost grows exponentially with the number of candidate effects, so the posterior can be enumerated exactly only when the number of predictors is small. Beyond that reach, the posterior has to be approximated, typically by Markov chain Monte Carlo over the model space, which requires a fresh run for every dataset and, within a fixed budget of steps, may fail to converge. We present LinearPFN, a prior-data fitted transformer network that amortizes spike-and-slab inference for linear models with main effects and pairwise interactions. The network is pretrained once on synthetic datasets, drawn from an explicitly specified prior, and a single forward pass over a new dataset returns posterior inclusion probabilities, posterior-mean coefficients and posterior predictive distributions with no per-dataset fitting. The prior is conjugate by design, so that the posterior for each fixed set of active effects has a closed form, and wherever the exact posterior is still computable by enumeration we verify the network's outputs against it. On real predictor matrices from published social-science datasets, with outcomes drawn from the prior so that the true active set is known, LinearPFN attains a higher per-dataset selection AUC and a higher F1 under the median probability model rule than five classical baselines. The lead holds when the coefficients, the interactions or the noise depart from the prior. Code: https://github.com/schiekiera/LinearPFN. Trained model: https://huggingface.co/schiekiera/LinearPFN.
Figures & tables
| predictive distribution and coefficients | inclusion probabilities | ||||
|---|---|---|---|---|---|
| criterion | threshold | value | criterion | threshold | value |
| headroom | 0.95 | 0.9605 | AUC gap | ||
| coverage error | 0.03 | 0.0005 | ECE gap | 0.01 | 0.0034 |
| coefficient correlation | 0.99 | 0.9931 | resolution ratio | 0.90 / 0.85 | 0.9640 |
| method | AUC/ds, mean | /ds, mean | coef RMSE, mean | fit s, median |
|---|---|---|---|---|
| LinearPFN | 0.96 [0.95, 1.00] | 0.77 [0.67, 1.00] | 0.041 [0.011, 0.051] | 7.40 [1.07, 24.59] |
| hierNet | 0.92 [0.88, 1.00] | 0.53 [0.38, 0.67] | 0.042 [0.011, 0.051] | 25.45 [9.17, 79.28] |
| glinternet | 0.79 [0.68, 0.97] | 0.57 [0.43, 0.71] | 0.044 [0.011, 0.052] | 4.39 [1.18, 15.05] |
| SuSiE | 0.85 [0.76, 1.00] | 0.66 [0.46, 0.89] | 0.051 [0.009, 0.066] | 0.16 [0.08, 0.40] |
| lasso | 0.82 [0.72, 0.96] | 0.48 [0.33, 0.63] | 0.049 [0.012, 0.059] | 0.14 [0.08, 0.71] |
| stability | 0.82 [0.72, 0.95] | 0.55 [0.40, 0.75] | 0.068 [0.014, 0.074] | 0.36 [0.33, 0.49] |
Appendix figures & tables9 assets
Supplementary material from the paper’s appendix.
Appendix
| Predictor prior | |
|---|---|
| rows , predictors | U{20, …, 2,000}, U{2, …, 30} |
| factor share | with probability 0.15, else Beta(1.5, 3) |
| factors | U{1, …, }, the size of the factor block |
| communality: mean, concentration | U(0.1, 0.92), U(4, 40), cap 0.98 |
| high band: probability, mean | 0.15, U(0.9, 0.995), cap 0.995 |
| share of on other factors | 0.15 |
| layers / heads / embedding | 8 / 16 / 512 | parameters | 31.5M (31,541,766) |
| MLP hidden units | 1,024 | GPUs | 2 H200 |
| optimizer steps | 37,500 | wall time | 14.6 h |
| effective batch | 64 (4 16) | GPU-hours | 29.1 |
| datasets seen | 2,400,000 | throughput | 0.72 steps/s |
| learning rate | Muon 0.02, AdamW 0.001 | peak GPU memory | 57.3 GiB |
| reason | entries | reason | entries |
| package outside the frame | 2,908 | described as simulated | 17 |
| index: fewer than 30 rows | 93 | fewer than 3 columns after cleaning | 15 |
| index: fewer than 3 usable columns | 81 | more than 50.0% of rows lost | 13 |
| time-series or contingency-table object | 69 | same title as another table | 8 |
| pure time series | 32 | download failed | 3 |
| same columns as a larger table | 23 | fewer than 30 rows after cleaning | 3 |
| baseline | /ds | AUC/ds | coefficient RMSE | RMSE, hard-zeroed |
|---|---|---|---|---|
| hierNet | +0.244 | +0.044 | 0.000 | |
| [+0.238, +0.249] | [+0.043, +0.046] | [ , 0.000] | [ , +0.001] | |
| glinternet | +0.205 | +0.173 | ||
| [+0.200, +0.210] | [+0.170, +0.176] | [ , ] | [ , ] | |
| SuSiE | +0.118 | +0.110 | ||
| [+0.114, +0.122] | [+0.107, +0.113] | [ , ] | [ , ] |
| control | fixed-magnitude coefficients | |||||
|---|---|---|---|---|---|---|
| method | AUC/ds | /ds | RMSE | AUC/ds | /ds | RMSE |
| LinearPFN | 0.96 | 0.77 | 0.041 | 0.98 (+0.011) | 0.85 (+0.072) | 0.044 (+0.003) |
| hierNet | 0.92 | 0.53 | 0.042 | 0.94 (+0.018) | 0.52 ( ) | 0.045 (+0.003) |
| glinternet | 0.79 | 0.57 | 0.044 | 0.85 (+0.062) | 0.56 ( ) | 0.048 (+0.004) |
| SuSiE | 0.85 | 0.66 | 0.051 | 0.89 (+0.034) | 0.70 (+0.044) | 0.063 (+0.012) |
| lasso | 0.82 | 0.48 | 0.049 | 0.86 (+0.040) | 0.48 (+0.004) | 0.056 (+0.007) |
| PIP | coefficient | MPM | held-out NLL | |||||||
|---|---|---|---|---|---|---|---|---|---|---|
| table | , mains | RMSE | max diff. | RMSE | agree | LinearPFN | MCMC | OLS | ||
| attitude | 30 | 6 | 0.998 | 0.016 | 0.043 | 0.006 | 21/21 | 1.095 | 1.122 | 1.177 |
| swiss | 47 | 5 | 0.854 | 0.074 | 0.157 | 0.019 | 15/15 | 1.277 | 1.312 | 1.765 |
| prostate | 97 | 8 | 0.990 | 0.028 | 0.107 | 0.009 | 36/36 | 0.927 | 0.894 | 0.876 |
| tal_or | 123 | 5 | 0.999 | 0.017 | 0.044 | 0.009 | 15/15 | 1.226 | 1.226 | 1.246 |
| bodyfat | 252 | 13 | 0.780 | 0.083 | 0.600 | 0.019 | 87/91 | 0.651 | 0.663 | 0.667 |