Transformers Can Learn Posterior Predictive Distributions In-Context
AuthorsGyeonghun Kang, Changwoo J. Lee, Xiang Cheng
Resources
This paper shows that transformers can go beyond point estimates and learn full Bayesian predictive distributions in context, with theory explaining why architecture choices like normalization and attention depth matter.
What the paper found
Transformers Can Learn Posterior Predictive Distributions In-Context, by Gyeonghun Kang, Changwoo J. Lee, and Xiang Cheng from Duke University, gives a constructive theory for prior-data fitted networks, or PFNs, showing that a transformer can approximate the full Gaussian process posterior predictive distribution rather than only a point estimate. The key result is an explicit attention-based algorithm: masked self-attention implements a Richardson-style iterative solver for the posterior predictive mean and variance, and a shallow MLP converts those moments into binned probabilities via a softmax over discretized density bins. The approximation error is bounded by three terms: exponential decay in attention depth L, discretization error of order 1/C from the number of bins, and tail truncation outside the support interval. The paper then identifies why PFNs generalize to context sizes n beyond pretraining: without normalization, the admissible step size shrinks like 1/n for both linear and RBF kernels, causing instability or under-convergence, while row-normalized attention acts as Jacobi preconditioning and stabilizes the spectrum. Yet even with normalization, the condition number still grows as Θ(n) for RBF kernels, so deeper stacks are required; the authors prove that to reach fixed accuracy, depth must scale at least on the order of n log(1/ε). Experiments on Bayesian linear regression, RBF regression, and real Sacramento housing and Walker Lake spatial data confirm the theory: increasing depth and bin resolution reduces total variation error, and normalized models maintain calibrated 90% and 95% intervals much farther outside the pretraining range.
Original abstract
Prior-data fitted networks (PFNs) have recently emerged as a powerful approach for Bayesian prediction tasks, approximating the posterior predictive distribution (PPD) through in-context learning. Despite their strong empirical performance and ability to go beyond point predictions, theoretical understandings of the algorithmic capability of transformers to learn distributions in context are still lacking. Focusing on Gaussian process regression problems, we show by construction that transformers can implement a gradient descent algorithm targeting the posterior predictive mean and variance, followed by nonlinear mappings that yield binned probabilities of PPD. We study the error bounds of the approximated PPD in terms of attention depth and bin resolution. Based on these results, we further demonstrate the key role of normalization and the choice of attention depth in enabling the extrapolation abilities of transformers beyond the pretraining sample size range. We conduct simulations that corroborate our findings, providing insight into the expressivity of PFNs targeting PPDs and how architectural choices may influence generalization capabilities.
Read the original paperMore in Transformers
Browse all 42 papers →Pretraining Latent Information Feedback Transformers with Teacher Supervision
Dor Tirosh, Ido Amos, Mor Geva
LIFT teaches Transformers to pass rich hidden-state information across steps, potentially making language models more efficient and capable than standard feed-forward designs.
The Geometry of Inference in Transformer Residual Streams
Timur Mudarisov, Mikhail Burtsev, Radu State
This paper shows how Transformer hidden states gradually geometrically converge toward the correct prediction while eliminating competing possible outcomes.
Transformers Stop Thinking Too Early, and a Tiny LoRA Fixes It
Zehao Jin, Ruixuan Deng, Junran Wang
A small LoRA update appears to make transformers carry information through many more layers, dramatically extending their ability to follow long chains without retraining the full model.