Pretraining task diversity and the emergence of non-Bayesian in-context learning for regression

Allan Raventós, Mansheej Paul, Feng Chen, Surya Ganguli

Introduction

Pretrained transformers (PTs) can learn new tasks from just a few examples provided in the prompt without taking any gradient steps on those examples . This ability, called in-context learning (ICL), has unlocked the widspread use of language models by making it efficient to adapt general purpose models to bespoke tasks without explicit training. Though remarkable, what makes ICL mysterious, and potentially harmful , is that the learning algorithm implemented by the PT in its forward pass is not built into its architecture or training process; instead it emerges from pretraining on large-scale data with a next token prediction objective. This raises a foundational question: can ICL really solve fundamentally new tasks that are very different from those seen during pretraining? If so, what learning algorithm does ICL implement? To answer these questions, we need to better understand how the different ingredients that go into pretraining influence this ability.

Towards this end, we explore how the diversity of tasks in the pretraining data affects the emergence of ICL. Prior work has proposed that ICL works by performing Bayesian inference. During pretraining, transformers learn a prior over latent tasks represented in the pretraining data. When prompted with examples at inference time, they "retrieve" relevant pretraining tasks and generate subsequent tokens from the posterior distribution conditioned on the query and inferred tasks. This suggests that ICL performance on a new task is influenced by its similarity to tasks implicitly learned during pretraining. However, the distribution of tasks in our pretraining data, TPretrain\mathcal{T}_{\text{Pretrain}}, is usually a limited and unrepresentative subsample of the ideal distribution of tasks, TTrue\mathcal{T}_{\text{True}}, that we want our model to be capable of learning in-context. For instance, TTrue\mathcal{T}_{\text{True}} could be the set of all instructions we want an A.I. assistant to follow. But, large-scale language modeling datasets used to pretrain these models contain very few examples correctly formatted for ICL. Instruction finetuning (IFT) datasets designed to ameliorate this are expensive to collect and thus contain tasks from just a few domains. Under the Bayesian framework, this distribution mismatch would cause the Bayesian estimator with a prior over the limited pretraining tasks, TPretrain\mathcal{T}_{\text{Pretrain}}, to perform suboptimally on tasks that are very different from those seen during pretraining. This motivates our question: can a model pretrained on a dataset with low task diversity nevertheless learn new, unseen tasks?

For general purpose language modeling, the size and complexity of TPretrain\mathcal{T}_{\text{Pretrain}} and the vague specification of TTrue\mathcal{T}_{\text{True}} make this question challenging to analyze. So, following recent work , we study ICL for linear regression. Here, a task is a linear regression problem with a given latent regression vector; the PT must predict the target for a new data point from examples of data-target pairs provided in the prompt. Prior work has shown that transformers that see an unlimited number of latent regression vectors during pretraining learn to perform ridge regression with the Bayes optimal ridge parameter. We instead consider the case where the pretraining task distribution, TPretrain\mathcal{T}_{\text{Pretrain}}, contains a limited and finite set of latent regression vectors (see Section 2 for details). To evaluate its ability to learn new tasks, the PT is tested on the ideal task distribution, TTrue\mathcal{T}_{\text{True}}, which is a Gaussian distribution over all latent regression vectors. Studying this setting has two advantages: first, we can directly vary the task diversity in TPretrain\mathcal{T}_{\text{Pretrain}} by changing the number of unique latent regression vectors seen during pretraining. Second, we can calculate the optimal estimator that minimizes the pretraining loss—the Bayesian estimator with prior TPretrain\mathcal{T}_{\text{Pretrain}}—as well as the optimal estimator for all tasks—the Bayesian estimator with prior TTrue\mathcal{T}_{\text{True}}. This allows us to interpret the behavior of the PT by comparing its predictions to those of the optimal estimators under either task distribution. In our work, we vary the pretraining task diversity and probe the PT’s ability to learn fundamentally new tasks in-context: does it behave like the optimal estimator for TPretrain\mathcal{T}_{\text{Pretrain}} and perform suboptimally on tasks from TTrue\mathcal{T}_{\text{True}}, or does it align with the optimal estimator for TTrue\mathcal{T}_{\text{True}} which can solve new, unseen tasks?

Contributions. Our contributions are as follows:

We find that a transformer pretrained on data with low task diversity behaves like the Bayesian estimator with prior TPretrain\mathcal{T}_{\text{Pretrain}}; it performs optimally on pretraining tasks but cannot learn new tasks in-context. However, as pretraining task diversity increases, the PT deviates from this Bayesian estimator, significantly outperforming it on new tasks, and at a large but still finite number of pretraining tasks, the PT’s performance closely matches that of the optimal estimator on TTrue\mathcal{T}_{\text{True}}.

We identify a task diversity threshold for the emergence of ICL. Below this threshold, increasing the pretraining dataset size while keeping task diversity constant biases the PT towards the pretraining task distribution. Conversely, beyond this threshold, increasing the dataset size without increasing its task diversity improves the PT’s performance on new, unseen tasks. This suggests that the PT’s behavior undergoes a sharp algorithmic phase transition in the limit of many examples per task, aligning with the optimal estimators on TPretrain\mathcal{T}_{\text{Pretrain}} before the threshold and on TTrue\mathcal{T}_{\text{True}} after it. We also examine this transition from the perspective of learning dynamics.

We empirically show that increasing the task dimension at fixed SNR increases the task diversity threshold. However, the scaling of the PT’s error with dimension is vastly superior to that of the optimal Bayesian estimator for TPretrain\mathcal{T}_{\text{Pretrain}}; at a task diversity that is beyond the threshold at all dimensions we consider, the PT remains near-optimal with increasing dimension, whereas the optimal estimator for TPretrain\mathcal{T}_{\text{Pretrain}} grows progressively less similar to the optimal estimator for TTrue\mathcal{T}_{\text{True}}.

We show that increasing weight decay significantly decreases the task diversity threshold while increasing number of layers or embedding size increases the task diversity threshold. This elucidates the effect of regularization and model capacity on the emergence of ICL.

Overall these contributions suggest that the emergence of in-context learning in pretrained transformers cannot be fully explained by a theory of Bayesian inference on the pretraining distribution.

Problem setup

Pretraining. The transformer is pretrained to minimize the next token prediction mean squared error (MSE) on sequences of data and target pairs. The latent regression vector for each sequence is drawn from the pretraining task distribution, TPretrain\mathcal{T}_{\text{Pretrain}}. This distribution has limited diversity as it is the uniform distribution over a finite set of MM tasks, TPretrain=U{w(1),...,w(M)}\mathcal{T}_{\text{Pretrain}}=\mathcal{U}\{\mathbf{w}^{(1)},...,\mathbf{w}^{(M)}\}. Each task in TPretrain\mathcal{T}_{\text{Pretrain}} is drawn i.i.d from a DD-dimensional standard normal distribution, w(i)∼N(0,ID), i∈1,...,M\mathbf{w}^{(i)}\sim\mathcal{N}(\mathbf{0},\mathbf{I}_{D}),\ i\in{1,...,M}. By increasing the number of tasks, MM, in TPretrain\mathcal{T}_{\text{Pretrain}}, we can increase the diversity of the pretraining data. Since the transformer makes a prediction for every data point in the sequence, its loss, LTPretrainL^{\mathcal{T}_{\text{Pretrain}}}, is just the MSE for each prediction, averaged over the predictions in the sequence:

Evaluation. We evaluate the PT’s performance on tasks seen during pretraining by computing LTPretrainL^{\mathcal{T}_{\text{Pretrain}}} using Eq. 1 but with new samples of data and noise. Since these are new instances of the task with new in-context examples, this evaluation corresponds to the test error of the PT. For a PT to successfully perform ICL of linear regression on new tasks, it must accurately predict the targets from the in-context examples for any task drawn from an ideal task distribution, TTrue\mathcal{T}_{\text{True}}, over all latent regression vectors; in our case TTrue=N(0,ID)\mathcal{T}_{\text{True}}=\mathcal{N}(\mathbf{0},\mathbf{I}_{D}). We evaluate the PT’s performance on new tasks by computing LTTrueL^{\mathcal{T}_{\text{True}}}, which follows Eq. 1 but where the tasks are sampled from the ideal task distribution: w∼TTrue\mathbf{w}\sim\mathcal{T}_{\text{True}} in the expectation.

For task distribution TPretrain=U{w(1),...,w(M)}\mathcal{T}_{\text{Pretrain}}=\mathcal{U}\{\mathbf{w}^{(1)},...,\mathbf{w}^{(M)}\}, the discrete minimum mean squared error (dMMSE) estimator is optimal. It is given by y^kdMMSE=(w^kdMMSE)⊺xk\hat{y}_{k}^{\text{dMMSE}}=\left(\hat{\mathbf{w}}_{k}^{\text{dMMSE}}\right)^{\intercal}\mathbf{x}_{k} where w^1dMMSE=1M∑i=1Mw(i)\hat{\mathbf{w}}_{1}^{\text{dMMSE}}=\frac{1}{M}\sum_{i=1}^{M}\mathbf{w}^{(i)} and for k∈{2,...,K}k\in\{2,...,K\}, (Section A.2)

Intuitively, w^kdMMSE\hat{\mathbf{w}}_{k}^{\text{dMMSE}} is just a weighted sum of the pretraining w(i)\mathbf{w}^{(i)}s with weight governed by the likelihood of observing targets {y1,...,yk−1}\{y_{1},...,y_{k-1}\} conditioned on inputs {x1,...,xk−1}\{\mathbf{x}_{1},...,\mathbf{x}_{k-1}\} and the task being w(i)\mathbf{w}^{(i)}. A PT that minimizes the pretraining loss LTPretrainL^{\mathcal{T}_{\text{Pretrain}}} will behave like this estimator.

For task distribution TTrue=N(0,ID)\mathcal{T}_{\text{True}}=\mathcal{N}(\mathbf{0},\mathbf{I}_{D}), the Ridge regression estimator with the ridge parameter set to the noise scale σ2\sigma^{2} is optimal: y^kRidge=(w^kRidge)⊺xk\hat{y}_{k}^{\text{Ridge}}=\left(\hat{\mathbf{w}}_{k}^{\text{Ridge}}\right)^{\intercal}\mathbf{x}_{k}, where w^1Ridge=0\hat{\mathbf{w}}_{1}^{\text{Ridge}}=\mathbf{0} and for k={2,...,K}k=\{2,...,K\},

Experiments and results

Unless specified otherwise, we study linear regression in D=8D=8 dimensions with up to K=16K=16 in-context examples and observation noise variance σ2=0.25\sigma^{2}=0.25. We use either a base transformer model with the GPT2 architecture with 8 layers, 128-dimensional embeddings, and 2 attention heads or a small model with 4 layers, 64-dimensional embeddings, and 2 attention heads. We train with the Adam optimizer and a one-cycle triangle learning rate schedule with 50% warmup. The base model is trained with batch size 256 for 500K training steps, though these hyperparameters are varied in our experiments. We always sweep over a range of learning rates and choose the largest learning rate at which training is stable. For further details see Appendix B.

To pretrain a randomly initialized transformer on data with task diversity MM, we first construct the pretraining task distribution, TPretrain\mathcal{T}_{\text{Pretrain}}, as described in Section 2. We then minimize the objective LTPretrainL^{\mathcal{T}_{\text{Pretrain}}} in Eq. 1 using minibatch stochastic gradient descent. For each sequence in a minibatch, we sample a single task w\mathbf{w} from TPretrain\mathcal{T}_{\text{Pretrain}}, as well as new samples of data, {xi}i=1K\{\mathbf{x}_{i}\}_{i=1}^{K}, and noise, {εi}i=1K\{\varepsilon_{i}\}_{i=1}^{K}, from their respective continuous distributions, to form a sequence (x1,w⊺x1+ε1,…,xK,w⊺xK+εK)(\mathbf{x}_{1},\mathbf{w}^{\intercal}\mathbf{x}_{1}+\varepsilon_{1},\ldots,\mathbf{x}_{K},\mathbf{w}^{\intercal}\mathbf{x}_{K}+\varepsilon_{K}). If we train for NN steps at batch size BB, the transformer will see a total of NBNB unique sequences and roughly NBM\frac{NB}{M} unique sequences for each latent task in TPretrain\mathcal{T}_{\text{Pretrain}}. By increasing either BB or NN at fixed MM, we can increase the total size of the pretraining dataset (or number of sequences per task) while keeping the dataset diversity—the number of unique w\mathbf{w}s in TPretrain\mathcal{T}_{\text{Pretrain}}—fixed.

For Fig. 2, we pretrain our base transformer on datasets with increasing task diversity (on the x-axis) while keeping the total number of sequences seen during pretraining fixed (B=256,N=500KB=256,N=500K). We evaluate the PTs and both optimal estimators on tasks seen during pretraining drawn from TPretrain\mathcal{T}_{\text{Pretrain}} (Fig. 2 top left) and on new tasks drawn from TTrue\mathcal{T}_{\text{True}} (Fig. 2 bottom left) and plot MSE normalized by task dimension—LT/DL^{\mathcal{T}}/D from Eq. 1). Since dMMSE is optimal on tasks from TPretrain\mathcal{T}_{\text{Pretrain}} (as discussed in Section 2), the green dMMSE markers denote the lowest possible loss the PT could achieve in this setting. In fact, the pretraining objective LTPretrainL^{\mathcal{T}_{\text{Pretrain}}} explicitly encourages the PT to match dMMSE performance. On the other hand, Ridge is optimal on tasks sampled from TTrue\mathcal{T}_{\text{True}} (Fig. 2 bottom left); the blue markers denote the lowest possible MSE the PT could attain on new tasks.

Low task diversity phase: the PT is Bayesian with respect to the pretraining distribution and cannot solve new tasks. At low pretraining task diversity—MM up to about 262^{6}—the PT’s MSE closely tracks that of dMMSE on tasks sampled from TPretrain\mathcal{T}_{\text{Pretrain}} (Fig. 2 top left); the PT performs optimally on tasks seen during pretraining. But it significantly underperforms on new tasks sampled from TTrue\mathcal{T}_{\text{True}}, indicated by the gap in MSE between the PT and Ridge (Fig. 2, bottom left). In this regime, it behaves like the Bayesian estimator with prior TPretrain\mathcal{T}_{\text{Pretrain}}.

High task diversity phase: the PT is non-Bayesian with respect to the pretraining task distribution and can solve new tasks. At higher task diversities—above 262^{6} pretraining tasks—the PT’s MSE deviates from dMMSE and approaches Ridge under both TPretrain\mathcal{T}_{\text{Pretrain}} and TTrue\mathcal{T}_{\text{True}}. Crucially, the PT starts to significantly outperform dMMSE on unseen tasks sampled from TTrue\mathcal{T}_{\text{True}} (Fig. 2 bottom left) at the expense of not fully minimizing its training objective, LTPretrainL^{\mathcal{T}_{\text{Pretrain}}} (gap between PT and dMMSE under TPretrain\mathcal{T}_{\text{Pretrain}}, Fig. 2 top left). This suggests that, a PT trained on a finite but large number of pretraining tasks can learn fundamentally new tasks in-context and this ability depends on it deviating from the optimal Bayesian estimator on the pretraining task distribution.

Finite size scaling of training data suggests an algorithmic phase transition as task-diversity increases. The experiments in Fig. 2 (left column) suggest that, when tested on both task distributions TPretrain\mathcal{T}_{\text{Pretrain}} and TTrue\mathcal{T}_{\text{True}}, the ICL algorithm implemented by a PT exhibits a smooth crossover in performance from dMMSE to Ridge. We next examine how this transition changes as we increase the number of sequences per task seen over pretraining, at fixed task diversity. One might reasonably expect that, if the transformer sees more sequences per latent task in TPretrain\mathcal{T}_{\text{Pretrain}}, both its predictions and performance should become more similar to those of dMMSE, and less similar to those of Ridge, at all values of task diversity. Strikingly, this natural expectation is violated in a manner that facilitates ICL on TTrue\mathcal{T}_{\text{True}}.

At each number of tasks, we increase the number of sequences per task by increasing batch size from 256 to 512 to 1024, while leaving the number of training steps fixed at 500K. We observe that ΔPT, dMMSETPretrain\Delta^{\mathcal{T}_{\text{Pretrain}}}_{\text{PT, dMMSE}}, which quantifies how different the PT and dMMSE estimator’s predictions are when testing on tasks drawn from TPretrain\mathcal{T}_{\text{Pretrain}}, does in fact decrease for M≤210M\leq 2^{10} (Fig. 2 top center) as we train on more sequences per task. Moreover, for each M∈{210,...,214}M\in\{2^{10},...,2^{14}\} the PT’s predictions also become less similar to those of Ridge, both on tasks from TPretrain\mathcal{T}_{\text{Pretrain}} (Fig. 2, top right) and TTrue\mathcal{T}_{\text{True}} (Fig. 2, bottom right). Crucially, this movement in behavior of the PT towards dMMSE and away from Ridge, at least on tasks drawn from TPretrain\mathcal{T}_{\text{Pretrain}}, holds only up to a threshold number of tasks between 2142^{14} and 2152^{15}. Beyond this threshold, pretraining on more sequences per task at a fixed task diversity actually makes the PT more like Ridge, in that both ΔPT,RidgeTPretrain\Delta^{\mathcal{T}_{\text{Pretrain}}}_{\text{PT,Ridge}} and ΔPT,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{PT,Ridge}} decrease (Fig. 2, right top and right bottom respectively). This means that, beyond a task diversity threshold, the PT can not only optimally solve new tasks from TTrue\mathcal{T}_{\text{True}} by matching Ridge performance, but also the PT gets better at doing so if trained on more sequences per task, despite the limited set of tasks experienced in pretraining. Thus, in contrast to the natural expectation stated above, more sequences per task does not promote overspecialization of the PT to the TPretrain\mathcal{T}_{\text{Pretrain}} at task diversities beyond the threshold.

Finally, the motion of the ICL algorithm implemented by PT towards (away) from Ridge above (below) a task diversity threshold (Fig. 2, right top and bottom) indicates that as one increases the number of sequences per task at fixed task diversity, the smooth cross over in performance of the PT between dMMSE and Ridge, shown in Fig. 2, left top and bottom, will become sharper and sharper in task diversity, ultimately exhibiting a sharp phase transition in the limit of infinite number of sequences per task. Remarkably, this phase transition in the ICL algorithm implemented by the PT appears at a moderate task diversity threshold below 2152^{15} pretraining tasks; even though dMMSE significantly underperforms relative to Ridge on TTrue\mathcal{T}_{\text{True}} at this task diversity, the PT nevertheless remains unimpaired by this limited task diversity and can optimally solve new tasks.

Increased training time at fixed batch size further supports an algorithmic phase transition. To confirm the above results, we also increase the number of sequences per task, at each task diversity, by increasing the number of training steps NN from 500K to 1M while keeping batch size fixed at 256. We observe that doubling NN (change from pale blue to red in Fig. 3) and doubling BB (change from pale blue to red in Fig. 2) have very similar effects on ΔPT,dMMSET\Delta^{\mathcal{T}}_{\text{PT,dMMSE}} and ΔPT,RidgeT\Delta^{\mathcal{T}}_{\text{PT,Ridge}}, for both T=TTrue\mathcal{T}=\mathcal{T}_{\text{True}} and T=TPretrain\mathcal{T}=\mathcal{T}_{\text{Pretrain}}. More importantly, the task diversity threshold, which we determined as the cross-over point in ΔPT,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{PT,Ridge}} between batch sizes 256, 512, and 1024 at 500K training steps (Fig. 2 bottom right) happens at the same number of tasks as the crossover point between 500K and 1M steps at batch size 256 (Fig. 3, right). Given that our two approaches for training the baseline transformer on more data yield the same task diversity threshold, and that doubling batch size leads to significantly faster training times than doubling number of steps, from here onward we consider the task diversity threshold to be cross-over point in ΔPT,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{PT,Ridge}} between batch sizes 256 and 512 when training for 500K steps. See Appendix D for more ablations of batch size and training steps that provide further evidence for how the number of sequences seen by the transformer is the key factor determining the similarity of its predictions to those of dMMSE and Ridge at each number of tasks.

Learning dynamics and a break in the scaling of early stopping time further supports an algorithmic phase transition. To probe if the observed transition is merely an effect of under-fitting, we study the learning dynamics of small PTs in the very large number of steps regime. First, in Appendix E, we verify that the small PT also demonstrates an algorithmic phase transition but at lower task diversity threshold between 2112^{11} and 2122^{12} pretraining tasks. In Fig. 4 left, we visualize the learning curves (ΔPT,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{PT,Ridge}} vs training steps) of PTs trained for 500K steps at batch size 512 with pretraining task diversities, MM, below and above the task diversity threshold, M∗M^{*}. For M<M∗M<M^{*}, ΔPT,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{PT,Ridge}} decreases early in training until it reaches a minimum at time step t∗t^{*}, and then increases as the PT approaches dMMSE. We define t∗t^{*} as the early stopping time for Ridge. For M>M∗M>M^{*}, ΔPT,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{PT,Ridge}} decreases throughout training. To evaluate if, in the latter case, models are undertrained and t∗t^{*} is larger than the total training time, we extend training to 2M steps at batch size 512 (4×4\times the training time, see Appendix B). Fig. 4 center, shows these learning curves along with that of the model trained with infinite task diversity; even in this long training regime, the task diversity threshold does not change. For both short and long training durations, models trained with the same MM have similar qualitative behavior (whether distance to Ridge decreases then increases or monotonically decreases). Additionally, learning curves of the models with M>M∗M>M^{*} are very similar to the learning curves for models trained on infinite pretraining task diversities and they achieve similar final accuracy (dahed lines vs markers in Fig. 10), suggesting that these models are approaching the Ridge solution.

In Fig. 4 right, we study how t∗t^{*}, scales with MM. For most M<M∗M<M^{*}, t∗t^{*} obeys a simple scaling behavior t∗∝Mαt^{*}\propto M^{\alpha}, α≈0.47\alpha\approx 0.47. However, for M>210M>2^{10}, the distance to Ridge decreases monotonically through training and t∗=t^{*}= 2M steps. Despite the caveat that our experiments are necessarily in the large but finite training step regime with a decayed learning rate schedule, this stark break in the scaling behavior of t∗t^{*} near the task diversity threshold suggests that the observed transition is not just caused by under-fitting but an underlying difference in the learning dynamics.

The transition along interpolating paths. To obtain an additional description of the algorithmic transition in the PT from dMMSE to Ridge, we compute the ICL performance of the PT, and compare it to both dMMSE and Ridge, on a one parameter family of new tasks wα\mathbf{w}_{\alpha} that interpolate between pairs of seen tasks wi\mathbf{w}_{i} and wj\mathbf{w}_{j} in the support of TPretrain\mathcal{T}_{\text{Pretrain}}. The interpolation path is given by

Here we fix the norm of the interpolated vector wα\mathbf{w}_{\alpha} to the average of the two endpoint norms to avoid ∥wα∥\|\mathbf{w}_{\alpha}\| taking on very small values for α∼12\alpha\sim\frac{1}{2}. Fig. 5 shows the results of this analysis for 252^{5} (left, low task diversity regime), 2102^{10} (center, below task diversity threshold), and 2152^{15} (right, just above the task diversity threshold) tasks. At each value of α\alpha, MSE is averaged over a large number of task pairs (wi,wj)(\mathbf{w}_{i},\mathbf{w}_{j}). Examination of the average performance at the center of many interpolation paths, corresponding to fundamentally new tasks far from tasks seen during pretraining, clearly reveals a transition in PT performance from dMMSE to Ridge, where new tasks can only be optimally learned above, but not below, the task diversity threshold. In contrast, unlike the PT, dMMSE cannot solve new tasks at any task diversity in the range considered.

The PT outperforms a smoothed dMMSE model. We have seen that at an intermediate task diversity the PT significantly outperforms dMMSE on new tasks in TTrue\mathcal{T}_{\text{True}}. It is clear why dMMSE performs poorly on new tasks in TTrue\mathcal{T}_{\text{True}} at low task diversity: its prior over tasks concentrates on MM unique tasks in TPretrain\mathcal{T}_{\text{Pretrain}}, while the prior over tasks in TTrue\mathcal{T}_{\text{True}} is Gaussian. A natural conjecture is that the PT cannot memorize all MM tasks in TPretrain\mathcal{T}_{\text{Pretrain}} for large enough MM. Therefore we also compare PT performance to a smoothed dMMSE estimator in which the discrete point prior over MM tasks seen in pretraining is replaced with a mixture of MM isotropic Gaussians with the same centers but with variance chosen to optimize performance on TTrue\mathcal{T}_{\text{True}} (see Appendix G for details). This smoothed dMMSE outperforms dMMSE as it has a prior over tasks closer to the Gaussian TTrue\mathcal{T}_{\text{True}}. But remarkably, the PT still outperforms the smoothed dMMSE even with optimal smoothing (Fig. 12). This indicates the PT, even at moderate task diversity, implements a more sophisticated algorithm than a simple smoothed dMMSE arising from the PT’s inability to resolve the MM pretraining tasks to high precision.

2 The PT exhibits superior scaling of task diversity threshold with dimension than dMMSE.

We next explore the dependence of the task diversity threshold on the regression problem dimension DD. We vary D∈{8,16,32}D\in\{8,16,32\} while simultaneously scaling up maximal context length as K=2DK=2D, and increasing observation noise σ2\sigma^{2} to match the SNR to that of D=8D=8 and σ2=0.25\sigma^{2}=0.25. We also train a larger model with 12 layers, 256-dimensional embeddings, and 4 attention heads that is sufficiently expressive to match Ridge performance at D=32D=32. Fig. 6, first 3 panels reveal that the task diversity threshold of the PT increases moderately (approximately linearly) with task dimension (i.e. roughly 2142^{14}, 2152^{15}, and 2162^{16} at D=8,16D=8,16, and 3232 respectively). This linear scaling is remarkable considering the volume of all possible tasks scales exponentially with dimension due to the concentration of the Gaussian TTrue\mathcal{T}_{\text{True}} to a sphere for large DD. Thus we expect dMMSE performance to scale much more poorly with dimension DD since the finite number of tasks in TPretrain\mathcal{T}_{\text{Pretrain}} would need to cover a substantial portion of the sphere for dMMSE to approach Ridge. To test this hypothesis, for M=220M=2^{20}, which is the largest task diversity we consider, we explore how the similarity of PT and dMMSE predictions to Ridge on new tasks scales with DD (Fig. 6, right panel). We see that ΔdMMSE,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{dMMSE,Ridge}} grows significantly as we increase DD, while remarkably ΔPT,RidgeTTrue\Delta^{\mathcal{T}_{\text{True}}}_{\text{PT,Ridge}} is largely dimension independent. Overall this indicates that the scaling of PT error with dimension is vastly superior than that of dMMSE; PT remains near optimal and close to Ridge at all DD for M=220M=2^{20}, while dMMSE departs from Ridge as DD increases.

3 Effect of Regularization and model capacity on the task diversity threshold.

We study the dependence of the task diversity threshold on various hyperparameters. First, adding explicit regularization in the form of weight decay (see Appendix B for details), and increasing its value over three orders of magnitude, consistently lowers the threshold task diversity (Fig. 7, left). Note however, the lower task diversity threshold also comes with worse performance (Figure 13, top). This suggests that various forms of implicit regularization could help drive the algorithmic transition in the PT without weight decay. We also explore the effect of model capacity on the task diversity threshold by either increasing the embedding dimension of both small and base PTs or increasing the depth of small PTs. Fig. 7 center shows that increasing embedding dimension over a reasonable range does not affect the task diversity threshold of base PT. However, for small PT, increasing either the embedding dimension (Fig. 7 center) or depth (Fig. 7 right) increases the task diversity threshold. Base PT has a much larger capacity then small PT and also has a larger threshold; we hypothesize that small PT is still in a regime where the threshold is sensitive to capacity while base PT is not. Together, these results suggest that model capacity plays an important role in the emergence of in-context learning: increasing capacity (up to a point) leads to an increase in the task-diversity threshold.

Related work

The Bayesian framework for ICL introduced by Xie et al. 2021, which motivates our work, hypothesizes that PTs "locate" concepts learned during pretraining to solve ICL tasks. A series of empirical work in language models use this framework to select better in-context examples while Min et al. 2022b use it to study the robustness of latent task inference. Our work builds on this framework in the linear regression setting and validates it at low task diversities. However, we find a regime—large but finite number of pretraining tasks—in which the ability to learn new tasks in-context is an emergent phenomenon that cannot be fully explained by Bayesian inference.

Prior work has also shown that transformers can do linear regression in-context. However, they pretrain with unlimited task diversity, sampling a completely new regression vector for each sequence. In contrast, our work considers pretraining datasets with limited task diversity where ICL on new tasks emerges even though the pretraining loss does not explicitly encode it. Another line of work hypothesizes that ICL performs gradient descent in the activations of the forward pass, providing explicit constructions for the weights of the PT to implement this for linear regression or exploring this hypothesis in language models . However more experiments are required to test the hypothesis that trained transformers actually match proposed constructions. Instead of studying the explicit mechanism by which in-context learning is implemented, our work focuses on the impact of the pretraining task diversity. Similar questions pertaining to the role of task diversification have been explored in the meta-learning literature .

Kirsch et al. 2022 also show the emergence of in-context learning with pretraining task diversity on a toy classification task. By studying this question in the controlled setting of linear regression, we can compare to the optimal estimators on TPretrain\mathcal{T}_{\text{Pretrain}} and TTrue\mathcal{T}_{\text{True}}. This allows us to establish that ICL at finite task diversity emerges because the PT departs from the optimal estimator on the pretraining task distribution, and is not just a consequence of the pretraining task distribution becoming similar to the ideal task distribution. Among other important perspectives on ICL, Chan et al. 2022 identify, in a toy setting, several properties of the training distribution—burstiness and occurrence of rare classes—that are necessary for the emergence of ICL. Wei et al. 2023 study how ICL in large language models is affected by semantic priors and input-label mappings, focusing on differences across model scale. Olsson et al. 2022 study inductions heads—circuits responsible for completing patterns by copying tokens—as a mechanism for implementing ICL.

Discussion

Overall, we have extensively explored the impact of pretraining task diversity on the emergence of in-context learning of fundamentally new tasks not seen during pretraining. We found several surprises by working in the controlled setting of linear regression, where we could compare the performance of the PT to Bayesian estimators that are optimal, either for the limited diversity pretraining task distribution TPretrain\mathcal{T}_{\text{Pretrain}} (i.e. dMMSE), or for the diverse ideal task distribution TTrue\mathcal{T}_{\text{True}} (i.e. Ridge). These comparisons reveal an algorithmic phase transition in the PT from the former to the latter at an intermediate task diversity threshold; beyond this threshold, the PT solves fundamentally new tasks not seen during pretraining. Strikingly, this task diversity threshold scales moderately with task dimension, over the range of dimensions considered, despite the exponential growth in the volume of all possible tasks with dimension. Indeed this PT scaling vastly outperforms that of dMMSE. Overall, these results indicate that ICL of new tasks by PTs is an emergent phenomenon that cannot be explained by Bayesian inference on limited diversity pretraining task distributions. Moreover, our experiments suggest some form of implicit regularization in PTs allows them to break free of the pretraining task distribution to solve new tasks, given a moderate pretraining task diversity.

Remarkably, beyond the task diversity threshold, PTs learn the optimal estimator for the underlying generative model for pretraining tasks; this is the case for both Gaussian and Laplace priors over tasks (see Figure 14 for experiments with Laplace prior). This is true even though solutions with lower training loss exist; indeed when trained on more data at fixed diversity, PTs behave more like Ridge at the expense of higher training loss. Our experiments in Fig. 4 suggest that this algorithmic transition is due an underlying change in learning dynamics. We explore this hypothesis by probing the linear mode connectivity of the loss landscape . In Fig. 11 we find that PTs trained with large MM inhabit the same loss basin as PTs trained with M=∞M=\infty: the training loss barrier between PTs trained with M≳213M\gtrsim 2^{13} and PTs with M=∞M=\infty is similar to two PTs trained with M=∞M=\infty. In contrast, there are large loss barriers between PTs trained with M<213M<2^{13} and M=∞M=\infty. Additionally, PTs trained with M=∞M=\infty are closer in weight space to PTs trained with large MM than those trained with small MM (see Appendix F). Overall, these experiments provide further evidence that PTs trained with task diversities beyond the threshold find solutions similar to the optimal model for TTrue\mathcal{T}_{\text{True}}; we leave further exploration of these loss landscapes to future work.

An intriguing question is how these observations carry over to language. A key mystery about the efficacy of ICL in language tasks lies in how different the tasks learned in-context are from the pretraining distribution of large language corpora. It is also less clear how to categorize the contents of such corpora according to tasks and measure their resulting task diversity. Regardless, our observation in linear regression that a moderate threshold in pretraining task diversity can enable PTs to solve new tasks may imply that many language tasks that are quite different from the statistics of large language corpora can still nevertheless be solved in-context.

Our results also suggest that the scale of data alone does not lead to good ICL performance. In fact, below the task diversity threshold, increasing the size of the pretraining dataset without increasing task diversity hurts ICL performance. It is necessary to increase both the diversity and size of the dataset for ICL to emerge. Thus to improve ICL in language settings, our work motivates future studies into uncovering the relevant notion of tasks in language modeling and approaches to increase task diversity in language corpora. More generally, our empirical analysis of the impact of pretraining task diversity on ICL motivates further theoretical studies. Such studies will be key to understanding the mystery of why simple next token prediction during pretraining can lead to in-context learning of so many apparently different tasks.

Acknowledgements

The authors would like to thank the Google TPU Research Cloud (TRC) whose generous support made this project possible. SG thanks an NSF CAREER award and NTT Research for support.

References

Appendix A Bayesian estimators

A.2 dMMSE estimator

For task distribution TPretrain=U{w(1),...,w(M)}\mathcal{T}_{\text{Pretrain}}=\mathcal{U}\{\mathbf{w}^{(1)},...,\mathbf{w}^{(M)}\}, we can directly get the discrete minimum mean squared error (dMMSE) estimator by plugging the uniform discrete distribution into Eq. 6.

A.3 Ridge estimator

For task distribution TTrue=N(0,ID)\mathcal{T}_{\text{True}}=\mathcal{N}(\mathbf{0},\mathbf{I}_{D}), we can directly get the Ridge regression estimator from Eq. 6:

Appendix B Experimental details

For most experiments, we study linear regression in D=8D=8 dimensions with up to K=16K=16 in-context examples and observation noise variance σ2=0.25\sigma^{2}=0.25. Our base model is a transformer with the GPT2 architecture with 8 layers, 128-dimensional embeddings, and 2 attention heads. We train with the Adam optimizer and a one-cycle triangle learning rate schedule with 50% warmup. The base model is trained with batch size 256 for 500K training steps, though these hyperparameters are varied in our experiments. Specifically, for the experiments in Fig. 3 we sweep over 500K and 1M training steps to find the task diversity threshold as a function of training steps. For all other experiments, we sweep batch size 256 and 512 and for the experiments in Fig. 2, we do an additional batch size of 1024.

For the experiments in Fig. 6, we use a larger model with 12 layers, 256-dimensional embeddings, and 4 attention heads. This ensures that the model can learn to perform ridge regression in the D=32D=32 setting. We also train the models on D=8,16,24,32D=8,16,24,32 with up to K=2DK=2D in-context examples in each case. The observation noise variance is scaled at each dimension to keep the signal to noise ratio constant. So σ2=0.5,0.707,0.866,1\sigma^{2}=0.5,0.707,0.866,1 for D=8,16,24,32D=8,16,24,32 respectively.

We always sweep over learning rates in {0.0001, 0.0003, 0.001} and choose the largest learning rate at which training is stable. For almost all experiments, the learning rate is 0.001, the main exception being the experiments in Fig. 6 where the learning rate is 0.0001.

For the 2M training step experiments in Fig. 4 and Fig. 10, the learning rate increases linearly up to 0.001 at 250K steps (i.e. using the same warmup as in the 500K step experiments) and then decreases linearly to 0 at 2M steps, consistent with our choice of decayed learning rate schedules. The experiments in Appendix F are performed at batch size 512 and following the 500K step one-cycle triangle learning rate schedule described above.

The implementation for this paper was done in JAX and all experiments were run on TPU v2-8 and v3-8 provided as part of the Google TPU Research Cloud program. Each training run takes approximately 4 v2-8 TPU hours.

Appendix C Support Figure 2

Fig. 8 provides an additional visualization for Fig. 2 upper middle but on a linear scale on the y-axis. In this figure it is clear that for all task diversities below the threshold—less that 2142^{14}—increasing the batch size aligns the PT with dMMSE, the optimal estimator for TPretrain\mathcal{T}_{\text{Pretrain}}. Thus below the task diversity threshold, the optimal estimator does indeed approach the Bayesian estimator with prior TPretrain\mathcal{T}_{\text{Pretrain}} when trained on more data and is thus unable to learn new tasks as discussed in Section 3.1.

Appendix D Dependence on number of sequences per task

In Section 3.1 we explore how the crossover of the PT from behaving like dMMSE to behaving like Ridge as we increase pretraining task diversity depends on the number of sequences per task seen during pretraining. In Fig. 2 middle and right column, we probe this by increasing the batch size to increase the number of sequences per task seen during pretraining at each level of task diversity. To ensure that that the phase transition and task diversity threshold we find is indeed a consequence of increasing the number of sequences per task and not merely just an effect of increasing batch size, we keep the batch size constant and increase the number of sequences per task in Fig. 3.

In Fig. 9 we provide additional evidence that the number of sequences per task and not batch size nor number of training steps is the primary factor that drives the behavior of the final model. Specifically, we show that the final behavior of the model is invariant to batch size or number of training steps if the number of training sequences per task is kept constant, atleast within a reasonable range of batch size and number of training steps. To do so, at each level of pretraining task diversity, we compare our base model (batch size = 256, 500K training steps, red cicles) to a model trained with batch size 512 (light blue squares) and batch size 128 (purple triangles). For either variant of the batch size, we adjust the number of training steps to keep the total number of pretraining sequences per task constant: 250K training steps for batch size 512 and 1M training steps for batch size 128. In all three cases, the behavior of the PT is identical as measured by ΔPT,dMMSET\Delta^{\mathcal{T}}_{\text{PT,dMMSE}} and ΔPT,RidgeT\Delta^{\mathcal{T}}_{\text{PT,Ridge}} for both TPretrain\mathcal{T}_{\text{Pretrain}} and TTrue\mathcal{T}_{\text{True}} across all numbers of pretraining tasks.

Appendix E Small PT

We run our task diversity threshold finding experiment in a small PT with 64 dimensional embeddings, 4 layers, and 2 attention heads. The existence of a task diversity threshold reproduces in a Small PT: the PTs transition from becoming less like Ridge to becoming more like Ridge when trained on more data (see Fig. 10). We also note that the Small PT with a lower capacity has a lower threshold (∼211.5\sim 2^{11.5} compared to ∼214.5\sim 2^{14.5} for the Base PT in Fig. 2) but also has lower overall performance (compare y-axis range to Fig. 2 lower right).

Appendix F Examining the loss landscape along paths interpolating between finite and infinite task models

For each finite number of tasks, MM, as well as M=∞M=\infty (where each training example samples a new task from the Gaussian distribution), we train four small PTs, each with a different random seed for tasks. This means that all randomness (including model initialization, data, and observation noise) is shared across the four experiments, except for that governing the sampling of w\mathbf{w}’s used for training. We then compare the finite MM models to the M=∞M=\infty models, both in terms of the size of the training loss barrier along a path interpolating between the models, as well as the distance in weight space between the models.

We compute the training loss barrier between a finite MM model and an infinite task model under the objective LPretrainTL^{\mathcal{T}}_{\text{Pretrain}} defined by the specific set of MM tasks the finite task model was trained on. We follow the approach outlined in where the barrier between two models is the loss at a model whose weights are the average of the two models, minus the average of the losses of the two models. Specifically, given two weight configurations w1w_{1} and w2w_{2} and a loss function LL, the loss barrier is calculated as

The distance in weight space, in turn, is simply the L2L_{2} norm of the difference in model weights. For each MM we average barriers and distances across all finite-infinite task pairs of models. We treat the analogous quantities computed between pairs of infinite task models as a baseline. See Fig. 11 and the Discussion section in the main text.

Appendix G Smoothed dMMSE

Here we consider the setting in which the discrete point prior over MM tasks seen in pretraining is replaced with a mixture of MM isotropic Gaussians with variance ϵ2\epsilon^{2}. We call the estimator that does posterior predictive inference under this prior the Smoothed dMMSE estimator (sMMSE). In the limit of ϵ→0\epsilon\rightarrow 0, the sMMSE estimator becomes the dMMSE estimator. On the other hand, if ϵ→∞\epsilon\rightarrow\infty, the prior will look like a Gaussian distribution but will have a variance much larger than the variance of TTrue\mathcal{T}_{\text{True}}. In both edge cases, it won’t perform well on TTrue\mathcal{T}_{\text{True}}, and so there is an optimal ϵ\epsilon that achieves the best performance on TTrue\mathcal{T}_{\text{True}}. In practice, we perform a search by simulation to determine the optimal ϵ\epsilon. We observe in Fig. 12 that for 2132^{13} tasks onward, the PT outperforms the optimally smoothed sMMSE on TTrue\mathcal{T}_{\text{True}}.

Appendix H Supporting figures