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, , is usually a limited and unrepresentative subsample of the ideal distribution of tasks, , that we want our model to be capable of learning in-context. For instance, 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, , 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 and the vague specification of 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, , 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, , 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 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 —as well as the optimal estimator for all tasks—the Bayesian estimator with prior . 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 and perform suboptimally on tasks from , or does it align with the optimal estimator for 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 ; 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 .
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 before the threshold and on 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 ; 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 grows progressively less similar to the optimal estimator for .
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, . This distribution has limited diversity as it is the uniform distribution over a finite set of tasks, . Each task in is drawn i.i.d from a -dimensional standard normal distribution, . By increasing the number of tasks, , in , we can increase the diversity of the pretraining data. Since the transformer makes a prediction for every data point in the sequence, its loss, , 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 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, , over all latent regression vectors; in our case . We evaluate the PT’s performance on new tasks by computing , which follows Eq. 1 but where the tasks are sampled from the ideal task distribution: in the expectation.
For task distribution , the discrete minimum mean squared error (dMMSE) estimator is optimal. It is given by where and for , (Section A.2)
Intuitively, is just a weighted sum of the pretraining s with weight governed by the likelihood of observing targets conditioned on inputs and the task being . A PT that minimizes the pretraining loss will behave like this estimator.
For task distribution , the Ridge regression estimator with the ridge parameter set to the noise scale is optimal: , where and for ,
Experiments and results
Unless specified otherwise, we study linear regression in dimensions with up to in-context examples and observation noise variance . 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 , we first construct the pretraining task distribution, , as described in Section 2. We then minimize the objective in Eq. 1 using minibatch stochastic gradient descent. For each sequence in a minibatch, we sample a single task from , as well as new samples of data, , and noise, , from their respective continuous distributions, to form a sequence . If we train for steps at batch size , the transformer will see a total of unique sequences and roughly unique sequences for each latent task in . By increasing either or at fixed , 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 s in —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 (). We evaluate the PTs and both optimal estimators on tasks seen during pretraining drawn from (Fig. 2 top left) and on new tasks drawn from (Fig. 2 bottom left) and plot MSE normalized by task dimension— from Eq. 1). Since dMMSE is optimal on tasks from (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 explicitly encourages the PT to match dMMSE performance. On the other hand, Ridge is optimal on tasks sampled from (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— up to about —the PT’s MSE closely tracks that of dMMSE on tasks sampled from (Fig. 2 top left); the PT performs optimally on tasks seen during pretraining. But it significantly underperforms on new tasks sampled from , 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 .
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 pretraining tasks—the PT’s MSE deviates from dMMSE and approaches Ridge under both and . Crucially, the PT starts to significantly outperform dMMSE on unseen tasks sampled from (Fig. 2 bottom left) at the expense of not fully minimizing its training objective, (gap between PT and dMMSE under , 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 and , 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 , 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 .
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 , which quantifies how different the PT and dMMSE estimator’s predictions are when testing on tasks drawn from , does in fact decrease for (Fig. 2 top center) as we train on more sequences per task. Moreover, for each the PT’s predictions also become less similar to those of Ridge, both on tasks from (Fig. 2, top right) and (Fig. 2, bottom right). Crucially, this movement in behavior of the PT towards dMMSE and away from Ridge, at least on tasks drawn from , holds only up to a threshold number of tasks between and . Beyond this threshold, pretraining on more sequences per task at a fixed task diversity actually makes the PT more like Ridge, in that both and 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 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 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 pretraining tasks; even though dMMSE significantly underperforms relative to Ridge on 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 from 500K to 1M while keeping batch size fixed at 256. We observe that doubling (change from pale blue to red in Fig. 3) and doubling (change from pale blue to red in Fig. 2) have very similar effects on and , for both and . More importantly, the task diversity threshold, which we determined as the cross-over point in 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 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 and pretraining tasks. In Fig. 4 left, we visualize the learning curves ( vs training steps) of PTs trained for 500K steps at batch size 512 with pretraining task diversities, , below and above the task diversity threshold, . For , decreases early in training until it reaches a minimum at time step , and then increases as the PT approaches dMMSE. We define as the early stopping time for Ridge. For , decreases throughout training. To evaluate if, in the latter case, models are undertrained and is larger than the total training time, we extend training to 2M steps at batch size 512 ( 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 have similar qualitative behavior (whether distance to Ridge decreases then increases or monotonically decreases). Additionally, learning curves of the models with 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 , scales with . For most , obeys a simple scaling behavior , . However, for , the distance to Ridge decreases monotonically through training and 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 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 that interpolate between pairs of seen tasks and in the support of . The interpolation path is given by
Here we fix the norm of the interpolated vector to the average of the two endpoint norms to avoid taking on very small values for . Fig. 5 shows the results of this analysis for (left, low task diversity regime), (center, below task diversity threshold), and (right, just above the task diversity threshold) tasks. At each value of , MSE is averaged over a large number of task pairs . 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 . It is clear why dMMSE performs poorly on new tasks in at low task diversity: its prior over tasks concentrates on unique tasks in , while the prior over tasks in is Gaussian. A natural conjecture is that the PT cannot memorize all tasks in for large enough . Therefore we also compare PT performance to a smoothed dMMSE estimator in which the discrete point prior over tasks seen in pretraining is replaced with a mixture of isotropic Gaussians with the same centers but with variance chosen to optimize performance on (see Appendix G for details). This smoothed dMMSE outperforms dMMSE as it has a prior over tasks closer to the Gaussian . 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 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 . We vary while simultaneously scaling up maximal context length as , and increasing observation noise to match the SNR to that of and . 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 . Fig. 6, first 3 panels reveal that the task diversity threshold of the PT increases moderately (approximately linearly) with task dimension (i.e. roughly , , and at , and respectively). This linear scaling is remarkable considering the volume of all possible tasks scales exponentially with dimension due to the concentration of the Gaussian to a sphere for large . Thus we expect dMMSE performance to scale much more poorly with dimension since the finite number of tasks in would need to cover a substantial portion of the sphere for dMMSE to approach Ridge. To test this hypothesis, for , 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 (Fig. 6, right panel). We see that grows significantly as we increase , while remarkably 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 for , while dMMSE departs from Ridge as 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 and . 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 (i.e. dMMSE), or for the diverse ideal task distribution (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 inhabit the same loss basin as PTs trained with : the training loss barrier between PTs trained with and PTs with is similar to two PTs trained with . In contrast, there are large loss barriers between PTs trained with and . Additionally, PTs trained with are closer in weight space to PTs trained with large than those trained with small (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 ; 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 , 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 , we can directly get the Ridge regression estimator from Eq. 6:
Appendix B Experimental details
For most experiments, we study linear regression in dimensions with up to in-context examples and observation noise variance . 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 setting. We also train the models on with up to in-context examples in each case. The observation noise variance is scaled at each dimension to keep the signal to noise ratio constant. So for 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 —increasing the batch size aligns the PT with dMMSE, the optimal estimator for . Thus below the task diversity threshold, the optimal estimator does indeed approach the Bayesian estimator with prior 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 and for both and 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 ( compared to 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, , as well as (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 ’s used for training. We then compare the finite models to the 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 model and an infinite task model under the objective defined by the specific set of 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 and and a loss function , the loss barrier is calculated as
The distance in weight space, in turn, is simply the norm of the difference in model weights. For each 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 tasks seen in pretraining is replaced with a mixture of isotropic Gaussians with variance . We call the estimator that does posterior predictive inference under this prior the Smoothed dMMSE estimator (sMMSE). In the limit of , the sMMSE estimator becomes the dMMSE estimator. On the other hand, if , the prior will look like a Gaussian distribution but will have a variance much larger than the variance of . In both edge cases, it won’t perform well on , and so there is an optimal that achieves the best performance on . In practice, we perform a search by simulation to determine the optimal . We observe in Fig. 12 that for tasks onward, the PT outperforms the optimally smoothed sMMSE on .